Compare commits

..
Author SHA1 Message Date
QiuSW b4ff9096c3 test(client): tighten polling test synchronization 2026-08-05 12:47:34 +08:00
QiuSW ae4b9ef08b test(client): accelerate T-103 feedback loop 2026-08-05 12:47:22 +08:00
QiuSW 1bf3acda70 docs(state): record one-shot sku reveal readiness 2026-08-05 12:37:28 +08:00
QiuSW e8ca738b33 feat(client): capture one-shot sku reveal evidence 2026-08-05 12:34:59 +08:00
QiuSW 429c2d33db docs(safety): define one-shot sku reveal evidence 2026-08-05 11:45:58 +08:00
QiuSW a15b6c0911 docs(safety): bind T-103 live panel states 2026-08-05 11:37:19 +08:00
QiuSW 29d1540f7c docs(tasks): record corrected T-103 entry 2026-08-05 10:57:55 +08:00
QiuSW 2f1380d084 fix(client): target verified bottom sku entry 2026-08-05 10:56:55 +08:00
QiuSW 81a44caeef docs(safety): correct verified sku entry boundary 2026-08-05 10:45:28 +08:00
QiuSW 757eca39fd docs(tasks): record T-103 entry substage 2026-08-05 10:34:12 +08:00
QiuSW b754cc9ca9 fix(client): identify sku entry failure substage 2026-08-05 10:32:35 +08:00
QiuSW e45d26fe93 docs(tasks): record T-103 entry evidence 2026-08-05 10:16:05 +08:00
QiuSW a09c58beef fix(client): stabilize verified sku entry 2026-08-05 10:10:56 +08:00
QiuSW 42eac57824 docs(tasks): record safe T-103 diagnostics 2026-08-05 09:24:06 +08:00
QiuSW c9b39a4e05 test(client): remove polling stop race 2026-08-05 09:14:02 +08:00
QiuSW 870e5f861d fix(client): expose safe T-103 failure stages 2026-08-05 09:07:30 +08:00
QiuSW 025edaf273 docs(tasks): record T-304 code-complete review 2026-08-05 02:03:02 +08:00
QiuSW 584c5601c5 task(T-304): record code-complete verification 2026-08-05 01:59:35 +08:00
QiuSW 3a27225977 fix(client): fail closed on polling configuration drift 2026-08-05 01:59:35 +08:00
QiuSW dfd88c3336 feat(client): add safe polling session UI 2026-08-05 01:59:35 +08:00
QiuSW 20112dbd61 feat(client): expose secret-free profile summary 2026-08-05 01:59:35 +08:00
QiuSW 404d0d3ca4 task(T-304): add metadata and stop-latch contract 2026-08-05 01:59:35 +08:00
QiuSW 85e5f1e56e task(T-304): begin polling session UI 2026-08-05 01:59:35 +08:00
QiuSW 6145fc4468 docs(tasks): complete T-211 claim wire bounds 2026-08-05 01:52:49 +08:00
QiuSW f6cd65208d fix(api): align claim bounds across runtime snapshots 2026-08-05 01:39:35 +08:00
QiuSW 7ea1c5349f fix(api): bound claim wire fields end to end 2026-08-05 01:25:23 +08:00
QiuSW 57a5b2e91c chore(tasks): include T-211 claim handler boundary 2026-08-05 01:18:23 +08:00
QiuSW 63c2c59b41 chore(tasks): start T-211 claim wire bounds 2026-08-05 01:11:21 +08:00
QiuSW cfb238d022 docs(tasks): complete T-303 durable client state 2026-08-05 01:04:35 +08:00
QiuSW 31f07ac245 feat(client): add durable HTTP task state 2026-08-05 00:59:32 +08:00
QiuSW 488005ac93 chore(tasks): start T-303 client HTTP state 2026-08-05 00:59:32 +08:00
QiuSW 4186f1315a docs(tasks): tighten exit and evidence contracts 2026-08-04 23:36:01 +08:00
QiuSW 3a0a41db00 docs(tasks): complete MVP safety task graph 2026-08-04 23:20:28 +08:00
QiuSW 8600c33391 docs(tasks): define admin safety closure chain 2026-08-04 22:44:13 +08:00
QiuSW 8092431208 docs(tasks): order evidence before attempt events 2026-08-04 22:41:59 +08:00
QiuSW cd47d0c959 docs(tasks): complete T-302 atomic claim 2026-08-04 22:34:46 +08:00
QiuSW ec42640123 Merge branch 'main' into task/t-302-atomic-claim 2026-08-04 22:33:25 +08:00
QiuSW 5f060f60ee feat(admin): add atomic task claim leases 2026-08-04 22:33:16 +08:00
QiuSW 19f81a5e58 docs(tasks): define T-305 prefence dry run 2026-08-04 22:14:55 +08:00
QiuSW c1cee49fb1 docs(tasks): define T-210 gate evidence kinds 2026-08-04 22:08:58 +08:00
QiuSW 772e379fd6 docs(tasks): define T-107 confirm gate 2026-08-04 22:05:20 +08:00
QiuSW 83b244ff7c docs(tasks): require T-106 navigation-source evidence 2026-08-04 22:02:43 +08:00
QiuSW d9a31cbfa6 docs(tasks): define T-106 confirm evidence 2026-08-04 22:01:45 +08:00
QiuSW 014405466a docs(tasks): define T-105 quantity gate 2026-08-04 21:59:21 +08:00
QiuSW 8c50e1579e docs(tasks): define T-307 attempt sink 2026-08-04 21:50:15 +08:00
QiuSW 35d7ce11d6 docs(tasks): define T-306 evidence publisher 2026-08-04 21:48:22 +08:00
QiuSW 41e54e313e docs(tasks): expose T-104 gate-one artifact 2026-08-04 21:46:54 +08:00
QiuSW e5de76503f docs(tasks): linearize T-205 ordering version 2026-08-04 21:45:31 +08:00
QiuSW 542b2283f4 docs(tasks): freeze T-303 evidence upload slot 2026-08-04 21:42:15 +08:00
QiuSW eb29bcd8b7 docs(tasks): define T-304 polling UI boundary 2026-08-04 21:36:52 +08:00
QiuSW 526af1eb31 docs(tasks): allow T-302 taskdetail fixture repair 2026-08-04 21:33:22 +08:00
QiuSW d9381a644b docs(tasks): define T-303 UI state boundary 2026-08-04 21:28:08 +08:00
QiuSW e863a29641 docs(tasks): allow T-303 roadmap sync 2026-08-04 21:19:58 +08:00
QiuSW 5c97954fa1 docs(tasks): define T-303 HTTP recovery client 2026-08-04 21:19:02 +08:00
QiuSW bd4c855ad2 docs(tasks): harden T-302 replay ownership 2026-08-04 21:15:59 +08:00
QiuSW ee90d76893 docs(tasks): linearize T-302 revocation checks 2026-08-04 21:11:26 +08:00
QiuSW 4a65138fd7 docs(tasks): define T-205 attempt events 2026-08-04 21:07:43 +08:00
QiuSW 8ec740dcef chore(tasks): start T-302 2026-08-04 21:03:38 +08:00
QiuSW 590c84660a docs(tasks): define T-302 atomic claiming 2026-08-04 21:02:57 +08:00
QiuSW 233188297a docs(tasks): complete T-301 2026-08-04 20:58:36 +08:00
QiuSW e674b7f131 merge: T-301 device credential isolation 2026-08-04 20:57:42 +08:00
QiuSW 66355a7f89 feat(admin): add device credential isolation 2026-08-04 20:57:33 +08:00
QiuSW 41881e81f3 docs(tasks): define T-301 auth failure semantics 2026-08-04 20:06:23 +08:00
QiuSW e3c87fdec5 docs(tasks): secure T-301 transport boundary 2026-08-04 20:05:03 +08:00
QiuSW 871cd24d68 fix(tasks): isolate T-301 write paths 2026-08-04 20:01:18 +08:00
QiuSW 4f8e71b256 chore(tasks): start T-301 2026-08-04 20:01:05 +08:00
QiuSW a6ad560f5d docs(tasks): define T-301 device identity 2026-08-04 20:00:29 +08:00
QiuSW 6a5547b323 docs(tasks): complete T-204 2026-08-04 19:51:58 +08:00
QiuSW 46fccc2120 merge: T-204 routed task evidence details 2026-08-04 19:43:20 +08:00
QiuSW 89648880bc feat(admin): add routed task evidence details 2026-08-04 19:42:35 +08:00
QiuSW 4829be972c docs(tasks): record T-103 entry audit 2026-08-04 18:47:47 +08:00
QiuSW 3b2fa3536e docs(tasks): define T-104 safe exit evidence 2026-08-04 18:47:05 +08:00
QiuSW 9a4d11f74b chore(tasks): start T-204 2026-08-04 18:42:09 +08:00
QiuSW 79284576ef docs(tasks): define T-204 detail evidence scope 2026-08-04 18:41:21 +08:00
QiuSW 20aaba0919 docs(tasks): complete T-203 2026-08-04 18:27:21 +08:00
QiuSW 03a067e29d merge: T-203 batch purchase authorization 2026-08-04 18:24:44 +08:00
QiuSW 5dcff4b15a feat(admin): authorize batch purchase starts 2026-08-04 18:24:39 +08:00
QiuSW ce9d6ca285 merge: T-103 verified entry parent chain 2026-08-04 18:21:38 +08:00
QiuSW 44586fef38 fix(client): bind SKU entry to verified parent chain 2026-08-04 18:21:30 +08:00
QiuSW e379d50101 chore(tasks): extend T-203 auth review scope 2026-08-04 18:18:19 +08:00
QiuSW 13547728fc docs(tasks): record T-103 entry parent evidence 2026-08-04 18:01:59 +08:00
QiuSW 0b5c561ed6 docs(tasks): record T-103 offline review 2026-08-04 17:50:46 +08:00
QiuSW 5f650f18b1 merge: T-103 offline SKU selection implementation 2026-08-04 17:48:29 +08:00
QiuSW b49a9b4abe feat(client): implement verified SKU selection flow 2026-08-04 17:47:58 +08:00
QiuSW 96f774ed96 docs(tasks): start T-203 purchase authorization 2026-08-04 17:16:48 +08:00
QiuSW 1f20271366 docs(tasks): sync T-209 completion 2026-08-04 17:15:12 +08:00
QiuSW 27726f4dde docs(state): record T-209 completion 2026-08-04 17:03:33 +08:00
QiuSW f85ef5f714 merge: T-209 single-pass schema 2026-08-04 17:02:29 +08:00
QiuSW e04f05b20b feat(admin): migrate core schema to single-pass model 2026-08-04 17:01:35 +08:00
QiuSW e20457b6db docs(tasks): define T-203 purchase authorization 2026-08-04 16:58:37 +08:00
QiuSW b5f45b87a5 docs(state): record T-202 completion 2026-08-04 16:41:22 +08:00
QiuSW 64a7468cab merge: T-202 draft task creation 2026-08-04 16:38:36 +08:00
QiuSW 442ab88fd7 docs(tasks): define T-209 single-pass schema migration 2026-08-04 16:38:29 +08:00
QiuSW d38cfb61af feat(admin): add draft task creation 2026-08-04 16:33:22 +08:00
QiuSW da540bfdf6 merge: authorized single-pass purchase contract 2026-08-04 16:25:42 +08:00
QiuSW cfd5440ac0 docs(architecture): adopt authorized single-pass purchase 2026-08-04 16:25:34 +08:00
QiuSW 8ba9b231f4 docs(tasks): start single-pass purchase redesign 2026-08-04 15:53:16 +08:00
QiuSW efeb2d958a docs(architecture): allow internal raw evidence screenshots 2026-08-04 15:47:47 +08:00
QiuSW 549099ad24 merge: T-103 observed price prefix 2026-08-04 15:23:10 +08:00
QiuSW 76a522f13f fix(client): bind T-103 observed price prefix 2026-08-04 15:23:02 +08:00
QiuSW 5088d544f2 merge: T-103 refined prefix diagnostics 2026-08-04 15:16:44 +08:00
QiuSW 6e001b2729 fix(client): refine T-103 prefix diagnostics 2026-08-04 15:16:30 +08:00
QiuSW cea27ff7ef docs(tasks): define T-202 draft creation 2026-08-04 15:08:13 +08:00
QiuSW 1c35155c2d merge: T-201 administrator sessions 2026-08-04 15:04:38 +08:00
173 changed files with 26999 additions and 2279 deletions
+7 -4
View File
@@ -10,7 +10,8 @@ cmbuyer 是一个自动化采购系统:**采购服务**(网页端,`admin/`
**系统只创建待付款订单,任何情况下都不自动付款。**
当前状态:仓库只有文档,尚未开始编码。阶段为 Phase 0。
当前状态:两端骨架、基础模型、登录和真机取证脚手架已落地;Phase 1 真机取证与不依赖页面判据的
Phase 2 服务端任务并行。实时快照见 [`docs/current-state.md`](docs/current-state.md)。
## 必读顺序
@@ -33,7 +34,9 @@ cmbuyer 是一个自动化采购系统:**采购服务**(网页端,`admin/`
系统**会**点击「提交订单」创建待付款订单,但必须满足四个前置条件(授权未消费且服务端
提交围栏已建立、闸门二通过、闸门三通过、控件唯一),且**只点一次**、点击后无论结果
都不重试。围栏申请失败或响应不明时不得点击;围栏建立后只能调和同一提交记录。
**第一趟试选的代码路径不得引用任何下单函数**,必须有测试证明不可达。
管理员点击“开始采购(只创建待付款订单)”是唯一的人类授权动作。T-103 的规格选择/读价隔离
验证路径不得引用数量、确认页、提交或付款函数;后续能力必须按真机取证任务逐段开放,并有静态
调用链测试证明未获准能力不可达。
### 2. 安全边界只能收紧
@@ -141,8 +144,8 @@ python scripts/validate_agent_context.py
**跨端契约改动必跑完整门禁**——两端会同时坏。
代码尚未初始化,上述命令在 T-001 / T-002 完成前不可运行;届时由对应任务替换为真实命令
并同步文档。
两端已初始化;跨端契约改动还需在仓库根运行 `./init.ps1`(或 `./init.sh`)完成安装、测试、
vet/build、compileall 和上下文门禁。
## 风格
+12 -20
View File
@@ -8,13 +8,14 @@
## 它做什么
一笔外部订单进来,采购人员需要去拼多多找到同款、选对颜色尺码、下单、把订单号抄回系统。
cmbuyer 把这个过程自动化,人只在两个点介入:**机器选对了吗**和**付不付款**。
cmbuyer 把这个过程自动化:管理员明确“买什么、买多少、最多多少钱”并点击开始采购,系统只
创建待付款订单,**是否付款始终由人决定**。
```text
手工填链接(MVP) / Excel · ERP(V2)
│
v
采购服务(admin/,Go) 建单 · 试选确认 · 下单授权 · 审计
采购服务(admin/,Go) 建单 · 开始采购授权 · 围栏 · 审计
│ HTTP
v
采购工具(client/,Python) 领任务 · 跑流程 · 回传
@@ -23,31 +24,22 @@ cmbuyer 把这个过程自动化,人只在两个点介入:**机器选对了
Android 手机(拼多多 App)
```
## 两趟执行
## 单趟执行
MVP 只做**任务自带商品链接**的情形,分两趟跑完:
MVP 只做**任务自带商品链接**的情形:管理员先把任务保存为 `DRAFT`,再在表格中勾选并点击
“开始采购(只创建待付款订单)”。该点击创建一次性授权,锁定商品、颜色、尺码、数量和最高总价。
| 趟次 | 做什么 |
| --- | --- |
| **第一趟 · 试选** | 开商品 → 精确勾选颜色分类和尺码 → 读单价 → 截图 → **退出释放手机** → 回传 |
| **人工确认** | 人在网页端看「机器选对了吗」→ 确认并**锁定单价** |
| **第二趟 · 下单** | 重新开商品 → 重新选同一规格 → **三道价格闸门** → 提交订单一次 → 转「待付款」 |
为什么分两趟:一台手机是瓶颈,不能停在规格面板上等人。代价是走两遍,换来手机不空闲,
且第二趟能抓住价格变动。
**三道价格闸门**:① 第一趟规格面板读价 ② 第二趟重读必须与授权价一致
③ 订单确认页「实付款」不超上限。任一道读不到或不通过即停,转人工。
采购工具领取后在同一次设备会话中完成:打开商品 → 精确选择规格 → 闸门一读 SKU 单价并校验
上限 → 设置并复核数量 → 闸门二重读规格与同价 → 进入确认页 → 闸门三校验应付总额 → 服务端
原子建立提交围栏 → 精确点击一次“提交订单” → 转待付款。中间不再回网页端等“机器选对了吗”。
价格**只在规格面板和订单确认页读**——别处的价格文本被拆成多个节点、带券后前缀、
实付价与原价混在一起,不可靠。
## 现在处于什么阶段
**Phase 0 · 地基。仓库目前只有文档,尚未开始编码。**
下一步:T-005(网页端 MVP 原型)与 T-006(桌面端 MVP 原型)的文件和自动检查已完成,
请先人工确认 6 个 HTML 原型,再进入 T-001 / T-002 骨架和生产实现。
**Phase 1 真机取证与 Phase 2 的安全服务端工作并行。** 两端骨架、基础模型、登录和真机基线已
落地;T-111 正在冻结单趟采购契约,T-103 随后继续规格精确选择与读价,T-202 等待主审合入。
详见 [`docs/current-state.md`](docs/current-state.md)。
**生死线是 M2**:真机能按链接打开商品、精确勾选颜色分类和尺码、**读到该 SKU 单价**。
@@ -74,7 +66,7 @@ MVP 只做**任务自带商品链接**的情形,分两趟跑完:
4. 规格按维度精确匹配,防前缀碰撞,找不到即停
5. 数量设置后必须读回复核
6. **三道价格闸门**,任一道读不到或不通过即停
7. **第一趟绝不下单**——试选路径不得引用下单函数
7. **能力分层**——T-103 规格验证路径不得引用数量、确认页或下单函数;后续能力逐段取证
8. 检测到外部支付交接立即停止,不读取不保存凭据
9. 检测到验证码 / 风控 / 人脸 / 短信校验立即停止,不绕过
10. 只读非敏感摘要,不提取收货地址原文、手机号、支付凭据
+46
View File
@@ -8,6 +8,13 @@
| `CMBUYER_ADMIN_PASSWORD_BCRYPT` | 非空 bcrypt 密码哈希,不接受明文密码。 |
| `CMBUYER_SESSION_SECRET` | 至少 32 字节的会话签名密钥。 |
| `CMBUYER_COOKIE_SECURE` | 可选;存在时只能精确为 `true` 或 `false`。HTTPS 部署应设为 `true`。 |
| `CMBUYER_DATABASE_SOURCE` | 已迁移 SQLite 的显式 data source。 |
| `CMBUYER_AUTHORIZATION_TTL` | 一次性授权的正 Go duration,例如 `10m`。 |
| `CMBUYER_MAX_TASK_QUANTITY` | 每条任务允许的正整数数量上限。 |
| `CMBUYER_MAX_TOTAL_PRICE` | 每条任务允许的规范正数总价上限,例如 `999.99`。 |
| `CMBUYER_EVIDENCE_DIR` | 内部原始截图的绝对私有目录;不得指向仓库或公开静态目录。 |
| `CMBUYER_CLAIM_TOKEN_SECRET` | claim token 专用 32 字节密钥的 64 位小写十六进制;不得复用 session 或设备 token。 |
| `CMBUYER_CLAIM_LEASE_TTL` | 正 Go duration,且必须严格短于 `CMBUYER_AUTHORIZATION_TTL`。 |
示例仅展示变量名,不提供可运行凭据:
@@ -16,7 +23,46 @@ $env:CMBUYER_ADMIN_USERNAME = '<管理员账号>'
$env:CMBUYER_ADMIN_PASSWORD_BCRYPT = '<bcrypt 密码哈希>'
$env:CMBUYER_SESSION_SECRET = '<至少 32 字节的随机密钥>'
$env:CMBUYER_COOKIE_SECURE = 'true'
$env:CMBUYER_DATABASE_SOURCE = '<SQLite data source>'
$env:CMBUYER_AUTHORIZATION_TTL = '10m'
$env:CMBUYER_MAX_TASK_QUANTITY = '99'
$env:CMBUYER_MAX_TOTAL_PRICE = '999.99'
$env:CMBUYER_EVIDENCE_DIR = '<内部截图绝对目录>'
$env:CMBUYER_CLAIM_TOKEN_SECRET = '<64 位小写十六进制随机值>'
$env:CMBUYER_CLAIM_LEASE_TTL = '1m'
go run ./cmd/migrate -database $env:CMBUYER_DATABASE_SOURCE up
go run ./cmd/server
```
MVP 服务固定绑定 IPv4 回环 `127.0.0.1:8080`,只供同一运营电脑上的采购服务和采购工具使用。不要把监听地址
改成 `0.0.0.0` 或局域网地址;未来若需要非回环访问,必须先单独建立并验收 HTTPS/TLS 终止与代理
信任边界,设备 Bearer 不得经过明文局域网。
数据库迁移完成后,用同一个显式 SQLite data source 管理设备凭据:
```powershell
# 签发:token 只在本次成功输出中显示一次,请立即放入采购工具的受控本机配置。
go run ./cmd/device-credentials -database $env:CMBUYER_DATABASE_SOURCE issue -name '<非秘密设备名称>'
# 仅显示设备 id、名称、状态和时间,不显示 token/hash。
go run ./cmd/device-credentials -database $env:CMBUYER_DATABASE_SOURCE list
# 撤销立即影响之后开始的每次设备请求;重复执行保持 REVOKED,不恢复旧 token。
go run ./cmd/device-credentials -database $env:CMBUYER_DATABASE_SOURCE revoke -device-id '<签发时的设备UUID>'
```
该 CLI 只接受已存在、已迁移的文件型 SQLite 普通路径或 `file:` URI,并强制以 `mode=rw` 打开;
路径拼错或文件缺失时 SQLite 原子拒绝且不会留下空数据库,也不接受内存、只读或可创建模式。CLI
不自动迁移。签发 token 是 32 字节加密随机值的 64 位小写十六进制表示;
SQLite 只保存原始 token 的 32 字节 SHA-256 BLOB。签发输出以外的 list/revoke、日志、错误和 HTTP
响应都不会显示 token 或 hash。
采购服务会话仅保存在当前进程内;进程重启后既有登录会话会安全失效。
管理员的“开始采购(只创建待付款订单)”只签发一次性授权并创建待付款订单的资格;服务不会自动付款,也不包含任何支付操作。
`GET /tasks/{id}` 直接访问时渲染完整详情页,任务列表以同一 URL 加载详情抽屉。内部截图只通过
`GET /evidence/{asset_id}` 向有效管理员会话提供,并始终返回 `no-store`;文件不在静态目录中。
`POST /api/v1/tasks/{id}/evidence` 使用 `Authorization: Bearer <token>` 和
`X-CMBuyer-Device-ID: <小写UUIDv4>` 逐请求查库认证;管理员会话不能代替设备身份。凭据错误统一空
401,SQLite 认证故障为空 503,且两者都发生在上传 body 被读取之前。
+200
View File
@@ -0,0 +1,200 @@
package main
import (
"context"
"encoding/json"
"errors"
"flag"
"fmt"
"io"
"log"
"net/url"
"os"
"path/filepath"
"strings"
"time"
"cmbuyer/admin/internal/deviceauth"
"cmbuyer/admin/internal/storage/sqlite"
)
func main() {
if err := run(context.Background(), os.Args[1:], os.Stdout, os.Stderr); err != nil {
log.Print(err)
os.Exit(1)
}
}
func run(ctx context.Context, args []string, stdout, stderr io.Writer) error {
flags := flag.NewFlagSet("device-credentials", flag.ContinueOnError)
flags.SetOutput(stderr)
databaseSource := flags.String("database", "", "explicit migrated SQLite data source")
if err := flags.Parse(args); err != nil {
return err
}
if *databaseSource == "" {
return errors.New("-database is required")
}
if flags.NArg() < 1 {
return errors.New("usage: device-credentials -database <sqlite-data-source> <issue|list|revoke> [options]")
}
command := flags.Arg(0)
commandArgs := flags.Args()[1:]
var issueName, revokeDeviceID string
switch command {
case "issue":
commandFlags := flag.NewFlagSet("issue", flag.ContinueOnError)
commandFlags.SetOutput(stderr)
commandFlags.StringVar(&issueName, "name", "", "non-secret device display name")
if err := commandFlags.Parse(commandArgs); err != nil {
return err
}
if issueName == "" || commandFlags.NArg() != 0 {
return errors.New("usage: device-credentials -database <sqlite-data-source> issue -name <display-name>")
}
case "list":
if len(commandArgs) != 0 {
return errors.New("usage: device-credentials -database <sqlite-data-source> list")
}
case "revoke":
commandFlags := flag.NewFlagSet("revoke", flag.ContinueOnError)
commandFlags.SetOutput(stderr)
commandFlags.StringVar(&revokeDeviceID, "device-id", "", "canonical device UUID")
if err := commandFlags.Parse(commandArgs); err != nil {
return err
}
if revokeDeviceID == "" || commandFlags.NArg() != 0 {
return errors.New("usage: device-credentials -database <sqlite-data-source> revoke -device-id <uuid>")
}
default:
return fmt.Errorf("unsupported device credential command %q", command)
}
if command == "issue" && !deviceauth.ValidDisplayName(issueName) {
return deviceauth.ErrInvalidCredential
}
if command == "revoke" && !deviceauth.ValidDeviceID(revokeDeviceID) {
return deviceauth.ErrInvalidCredential
}
existingSource, err := existingSQLiteDataSource(*databaseSource)
if err != nil {
return err
}
database, err := sqlite.Open(existingSource)
if err != nil {
return fmt.Errorf("open SQLite database: %w", err)
}
defer database.Close()
store, err := deviceauth.NewCredentialStore(database)
if err != nil {
return err
}
switch command {
case "issue":
issued, err := store.Issue(ctx, issueName)
if err != nil {
return err
}
// The token has json:"-" and is printed only by this explicit post-commit path. Generic
// serialization, list, revoke, errors, and server responses therefore cannot disclose it.
if _, err := fmt.Fprintf(stdout, "device_id=%s\ndisplay_name=%s\ntoken=%s\ncreated_at=%s\n",
issued.DeviceID, issued.DisplayName, issued.Token, issued.CreatedAt.Format(time.RFC3339Nano)); err != nil {
return errors.New("write issued device credential")
}
return nil
case "list":
credentials, err := store.List(ctx)
if err != nil {
return err
}
return writeJSON(stdout, credentials)
case "revoke":
credential, changed, err := store.Revoke(ctx, revokeDeviceID)
if err != nil {
return err
}
return writeJSON(stdout, struct {
Credential deviceauth.Credential `json:"credential"`
RevokedNow bool `json:"revoked_now"`
}{Credential: credential, RevokedNow: changed})
}
return errors.New("unreachable device credential command")
}
func existingSQLiteDataSource(value string) (string, error) {
if value == "" || strings.TrimSpace(value) != value {
return "", errors.New("-database must name an existing file-backed SQLite database")
}
var parsed *url.URL
var query url.Values
if strings.HasPrefix(strings.ToLower(value), "file:") {
var err error
parsed, err = url.Parse(value)
if err != nil || !strings.EqualFold(parsed.Scheme, "file") || parsed.User != nil || parsed.Host != "" || parsed.Fragment != "" {
return "", errors.New("-database file URI is invalid")
}
// go-sqlite3 recognizes URI filenames only with the exact lowercase file: prefix.
// Canonicalize accepted scheme casing before mode=rw reaches the driver, otherwise a
// mixed-case input could be treated as a plain filename and recreate a missing database.
parsed.Scheme = "file"
query, err = url.ParseQuery(parsed.RawQuery)
if err != nil {
return "", errors.New("-database query parameters are invalid")
}
fileName := parsed.Path
if parsed.Opaque != "" {
fileName = parsed.Opaque
}
decodedName, err := url.PathUnescape(fileName)
if err != nil || fileName == "" || strings.EqualFold(decodedName, ":memory:") {
return "", errors.New("-database must name an existing file-backed SQLite database")
}
} else {
pathPart, rawQuery, hasQuery := strings.Cut(value, "?")
if pathPart == "" || strings.EqualFold(pathPart, ":memory:") || strings.Contains(pathPart, "://") {
return "", errors.New("-database must name an existing file-backed SQLite database")
}
var err error
query, err = url.ParseQuery(rawQuery)
if err != nil {
return "", errors.New("-database query parameters are invalid")
}
normalizedPath := filepath.ToSlash(pathPart)
if filepath.VolumeName(pathPart) != "" && !strings.HasPrefix(normalizedPath, "/") {
normalizedPath = "/" + normalizedPath
}
parsed = &url.URL{Scheme: "file", Path: normalizedPath}
if !hasQuery {
query = make(url.Values)
}
}
modes := query["mode"]
if len(modes) > 1 || len(modes) == 1 && modes[0] != "rw" {
return "", errors.New("-database only permits SQLite mode=rw")
}
if len(modes) == 0 {
query.Set("mode", "rw")
}
for _, name := range []string{"immutable", "_query_only"} {
for _, setting := range query[name] {
if setting != "0" && !strings.EqualFold(setting, "false") {
return "", errors.New("-database contains a read-only SQLite option")
}
}
}
parsed.RawQuery = query.Encode()
parsed.ForceQuery = false
return parsed.String(), nil
}
func writeJSON(writer io.Writer, value any) error {
encoder := json.NewEncoder(writer)
encoder.SetEscapeHTML(true)
if err := encoder.Encode(value); err != nil {
return errors.New("write device credential metadata")
}
return nil
}
+293
View File
@@ -0,0 +1,293 @@
package main
import (
"bytes"
"context"
"database/sql"
"errors"
"io"
"net/url"
"os"
"path/filepath"
"runtime"
"strings"
"testing"
"cmbuyer/admin/internal/deviceauth"
"cmbuyer/admin/internal/migrations"
"cmbuyer/admin/internal/storage/sqlite"
)
func TestIssueListAndIdempotentRevokeNeverRediscloseSecret(t *testing.T) {
databaseSource := migratedDatabase(t)
var issued bytes.Buffer
if err := run(context.Background(), []string{"-database", databaseSource, "issue", "-name", "采购工具一号"}, &issued, io.Discard); err != nil {
t.Fatalf("issue: %v", err)
}
fields := outputFields(t, issued.String())
deviceID, token := fields["device_id"], fields["token"]
if len(token) != 64 || strings.Count(issued.String(), token) != 1 {
t.Fatalf("issue token occurrence/length = %d/%d", strings.Count(issued.String(), token), len(token))
}
var listed bytes.Buffer
if err := run(context.Background(), []string{"-database", databaseSource, "list"}, &listed, io.Discard); err != nil {
t.Fatalf("list: %v", err)
}
assertNoSecretMetadata(t, listed.String(), token)
if !strings.Contains(listed.String(), deviceID) || !strings.Contains(listed.String(), "采购工具一号") {
t.Fatalf("list omitted safe metadata: %s", listed.String())
}
var revoked bytes.Buffer
if err := run(context.Background(), []string{"-database", databaseSource, "revoke", "-device-id", deviceID}, &revoked, io.Discard); err != nil {
t.Fatalf("revoke: %v", err)
}
assertNoSecretMetadata(t, revoked.String(), token)
if !strings.Contains(revoked.String(), `"revoked_now":true`) {
t.Fatalf("first revoke output = %s", revoked.String())
}
var repeated bytes.Buffer
if err := run(context.Background(), []string{"-database", databaseSource, "revoke", "-device-id", deviceID}, &repeated, io.Discard); err != nil {
t.Fatalf("repeat revoke: %v", err)
}
assertNoSecretMetadata(t, repeated.String(), token)
if !strings.Contains(repeated.String(), `"revoked_now":false`) {
t.Fatalf("repeat revoke output = %s", repeated.String())
}
}
func TestIssueOutputFailureLeavesCommittedCredentialWithoutSecretInError(t *testing.T) {
databaseSource := migratedDatabase(t)
writer := &recordingFailureWriter{}
err := run(context.Background(), []string{"-database", databaseSource, "issue", "-name", "output failure"}, writer, io.Discard)
if err == nil || err.Error() != "write issued device credential" {
t.Fatalf("issue output failure error = %v", err)
}
fields := outputFields(t, writer.contents.String())
if strings.Contains(err.Error(), fields["token"]) {
t.Fatal("output error disclosed token")
}
database, err := sql.Open("sqlite3", databaseSource)
if err != nil {
t.Fatalf("open database: %v", err)
}
defer database.Close()
var count int
if err := database.QueryRow(`SELECT COUNT(*) FROM device_credentials`).Scan(&count); err != nil || count != 1 {
t.Fatalf("committed credential count = %d, err=%v", count, err)
}
}
func TestCLIRequiresPreMigratedExplicitDatabase(t *testing.T) {
if err := run(context.Background(), []string{"list"}, io.Discard, io.Discard); err == nil {
t.Fatal("command without -database succeeded")
}
missing := filepath.Join(t.TempDir(), "missing.db")
if err := run(context.Background(), []string{"-database", missing, "list"}, io.Discard, io.Discard); err == nil {
t.Fatal("list opened a missing database")
}
if _, err := os.Stat(missing); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("missing database was created: %v", err)
}
unmigrated := filepath.Join(t.TempDir(), "unmigrated.db")
unmigratedDatabase, err := sqlite.Open(unmigrated)
if err != nil {
t.Fatalf("create unmigrated database: %v", err)
}
if _, err := unmigratedDatabase.Exec(`CREATE TABLE unrelated (id INTEGER)`); err != nil {
_ = unmigratedDatabase.Close()
t.Fatalf("initialize unmigrated database: %v", err)
}
if err := unmigratedDatabase.Close(); err != nil {
t.Fatalf("close unmigrated database: %v", err)
}
if err := run(context.Background(), []string{"-database", unmigrated, "list"}, io.Discard, io.Discard); err == nil {
t.Fatal("list accepted an unmigrated database")
}
database, err := sql.Open("sqlite3", unmigrated)
if err != nil {
t.Fatalf("open unmigrated database: %v", err)
}
defer database.Close()
var count int
if err := database.QueryRow(`SELECT COUNT(*) FROM sqlite_master WHERE type='table' AND name='device_credentials'`).Scan(&count); err != nil || count != 0 {
t.Fatalf("device_credentials table count = %d, err=%v", count, err)
}
undeclared := filepath.Join(t.TempDir(), "undeclared.db")
if err := run(context.Background(), []string{"-database", undeclared, "rotate"}, io.Discard, io.Discard); err == nil {
t.Fatal("undeclared command succeeded")
}
if _, err := os.Stat(undeclared); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("undeclared command opened database: %v", err)
}
invalidIssue := filepath.Join(t.TempDir(), "invalid-issue.db")
if err := run(context.Background(), []string{"-database", invalidIssue, "issue", "-name", " padded"}, io.Discard, io.Discard); !errors.Is(err, deviceauth.ErrInvalidCredential) {
t.Fatalf("invalid issue error = %v", err)
}
if _, err := os.Stat(invalidIssue); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("invalid issue opened database: %v", err)
}
invalidRevoke := filepath.Join(t.TempDir(), "invalid-revoke.db")
if err := run(context.Background(), []string{"-database", invalidRevoke, "revoke", "-device-id", "not-a-uuid"}, io.Discard, io.Discard); !errors.Is(err, deviceauth.ErrInvalidCredential) {
t.Fatalf("invalid revoke error = %v", err)
}
if _, err := os.Stat(invalidRevoke); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("invalid revoke opened database: %v", err)
}
migrated := migratedDatabase(t)
unknownID := "13c9f507-7473-4fa6-8d71-8786c34c6301"
var unknownOutput bytes.Buffer
err = run(context.Background(), []string{"-database", migrated, "revoke", "-device-id", unknownID}, &unknownOutput, io.Discard)
if !errors.Is(err, deviceauth.ErrCredentialNotFound) || unknownOutput.Len() != 0 || strings.Contains(err.Error(), unknownID) {
t.Fatalf("unknown revoke = output %q, error %v", unknownOutput.String(), err)
}
}
func TestExistingSQLiteDataSourcePreservesSafeOptionsAndRejectsCreationModes(t *testing.T) {
databaseSource := migratedDatabase(t)
fileURI := (&url.URL{
Scheme: "file",
Path: sqliteURIPath(databaseSource),
RawQuery: "_busy_timeout=5000&cache=shared",
}).String()
normalized, err := existingSQLiteDataSource(fileURI)
if err != nil {
t.Fatalf("normalize file URI: %v", err)
}
parsed, err := url.Parse(normalized)
if err != nil {
t.Fatalf("parse normalized URI: %v", err)
}
query := parsed.Query()
if query.Get("mode") != "rw" || query.Get("_busy_timeout") != "5000" || query.Get("cache") != "shared" {
t.Fatalf("normalized query = %v", query)
}
if err := run(context.Background(), []string{"-database", fileURI, "list"}, io.Discard, io.Discard); err != nil {
t.Fatalf("list existing file URI: %v", err)
}
plainNormalized, err := existingSQLiteDataSource(databaseSource + "?_foreign_keys=on")
if err != nil {
t.Fatalf("normalize ordinary path: %v", err)
}
plainURI, err := url.Parse(plainNormalized)
if err != nil || plainURI.Scheme != "file" || plainURI.Query().Get("mode") != "rw" || plainURI.Query().Get("_foreign_keys") != "on" {
t.Fatalf("ordinary path normalization = %q, err=%v", plainNormalized, err)
}
missing := filepath.Join(t.TempDir(), "missing-uri.db")
missingURI := (&url.URL{Scheme: "file", Path: sqliteURIPath(missing)}).String()
if err := run(context.Background(), []string{"-database", missingURI, "list"}, io.Discard, io.Discard); err == nil {
t.Fatal("missing file URI succeeded")
}
if _, err := os.Stat(missing); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("missing file URI created a file: %v", err)
}
for _, scheme := range []string{"FILE", "File"} {
mixedMissing := filepath.Join(t.TempDir(), strings.ToLower(scheme)+"-missing.db")
canonical := (&url.URL{Scheme: "file", Path: sqliteURIPath(mixedMissing)}).String()
mixedURI := scheme + canonical[len("file"):]
normalized, err := existingSQLiteDataSource(mixedURI)
if err != nil || !strings.HasPrefix(normalized, "file:") {
t.Fatalf("normalize %s URI = %q, err=%v", scheme, normalized, err)
}
if err := run(context.Background(), []string{"-database", mixedURI, "list"}, io.Discard, io.Discard); err == nil {
t.Fatalf("missing %s URI succeeded", scheme)
}
if _, err := os.Stat(mixedMissing); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("missing %s URI created a file: %v", scheme, err)
}
}
for name, source := range map[string]string{
"plain memory": ":memory:",
"URI memory": "file::memory:?cache=shared",
"memory mode": fileURI + "&mode=memory",
"read only mode": fileURI + "&mode=ro",
"create mode": fileURI + "&mode=rwc",
"duplicate mode": fileURI + "&mode=rw&mode=rw",
"immutable": fileURI + "&immutable=1",
"query only": fileURI + "&_query_only=1",
"remote authority": "file://server/share/database.db?mode=rw",
} {
t.Run(name, func(t *testing.T) {
if _, err := existingSQLiteDataSource(source); err == nil {
t.Fatalf("unsafe source accepted: %q", source)
}
})
}
}
func sqliteURIPath(path string) string {
normalized := filepath.ToSlash(path)
if filepath.VolumeName(path) != "" && !strings.HasPrefix(normalized, "/") {
return "/" + normalized
}
return normalized
}
type recordingFailureWriter struct {
contents bytes.Buffer
}
func (writer *recordingFailureWriter) Write(value []byte) (int, error) {
_, _ = writer.contents.Write(value)
return 0, errors.New("injected stdout failure")
}
func assertNoSecretMetadata(t *testing.T, output, token string) {
t.Helper()
if strings.Contains(output, token) || strings.Contains(output, "token") || strings.Contains(output, "hash") || strings.Contains(output, "sha256") {
t.Fatalf("metadata output disclosed secret material: %s", output)
}
}
func outputFields(t *testing.T, output string) map[string]string {
t.Helper()
fields := make(map[string]string)
for _, line := range strings.Split(strings.TrimSpace(output), "\n") {
name, value, found := strings.Cut(line, "=")
if !found || name == "" || value == "" {
t.Fatalf("invalid issue output line %q", line)
}
fields[name] = value
}
for _, required := range []string{"device_id", "display_name", "token", "created_at"} {
if fields[required] == "" {
t.Fatalf("issue output missing %s: %q", required, output)
}
}
return fields
}
func migratedDatabase(t *testing.T) string {
t.Helper()
databaseSource := filepath.Join(t.TempDir(), "credentials.db")
database, err := sqlite.Open(databaseSource)
if err != nil {
t.Fatalf("open database: %v", err)
}
if err := migrations.Up(context.Background(), database, commandMigrationDirectory(t)); err != nil {
_ = database.Close()
t.Fatalf("migrate database: %v", err)
}
if err := database.Close(); err != nil {
t.Fatalf("close migrated database: %v", err)
}
return databaseSource
}
func commandMigrationDirectory(t *testing.T) string {
t.Helper()
_, file, _, ok := runtime.Caller(0)
if !ok {
t.Fatal("locate migrations")
}
return filepath.Join(filepath.Dir(file), "..", "..", "migrations")
}
+40 -1
View File
@@ -7,10 +7,18 @@ import (
"cmbuyer/admin/internal/auth"
"cmbuyer/admin/internal/config"
"cmbuyer/admin/internal/deviceauth"
"cmbuyer/admin/internal/server"
evidencestorage "cmbuyer/admin/internal/storage/evidence"
"cmbuyer/admin/internal/storage/sqlite"
"cmbuyer/admin/internal/taskclaim"
"cmbuyer/admin/internal/taskdetail"
"cmbuyer/admin/internal/tasks"
)
const listenAddress = ":8080"
// Device Bearer credentials must not cross a plaintext LAN. The MVP is a same-computer
// deployment, so widening this address requires a separately reviewed TLS boundary first.
const listenAddress = "127.0.0.1:8080"
func main() {
if err := run(); err != nil {
@@ -23,11 +31,42 @@ func run() error {
if err != nil {
return err
}
database, err := sqlite.Open(configuration.DatabaseSource)
if err != nil {
return err
}
defer database.Close()
taskStore, err := tasks.NewSQLiteStore(database)
if err != nil {
return err
}
taskStore.SetStartPolicy(tasks.StartPolicy{AuthorizationTTL: configuration.AuthorizationTTL, MaxQuantity: configuration.MaxTaskQuantity, MaxTotalPrice: configuration.MaxTotalPrice})
detailStore, err := taskdetail.NewSQLiteStore(database)
if err != nil {
return err
}
evidenceStore, err := evidencestorage.NewStore(database, configuration.EvidenceDirectory)
if err != nil {
return err
}
deviceAuthenticator, err := deviceauth.NewSQLiteAuthenticator(database)
if err != nil {
return err
}
claimStore, err := taskclaim.NewStore(database, configuration.ClaimTokenSecret, configuration.ClaimLeaseTTL)
if err != nil {
return err
}
router, err := server.NewRouter(server.Options{
AdminUsername: configuration.AdminUsername,
AdminPasswordBcrypt: configuration.AdminPasswordBcrypt,
Sessions: auth.NewManager(configuration.SessionSecret, configuration.CookieSecure),
Tasks: taskStore,
TaskDetails: detailStore,
Evidence: evidenceStore,
DeviceAuthenticator: deviceAuthenticator,
TaskClaims: claimStore,
})
if err != nil {
return err
+9
View File
@@ -0,0 +1,9 @@
package main
import "testing"
func TestListenAddressIsIPv4LoopbackOnly(t *testing.T) {
if listenAddress != "127.0.0.1:8080" {
t.Fatalf("listenAddress = %q, want loopback-only endpoint", listenAddress)
}
}
+6
View File
@@ -62,6 +62,12 @@ func (manager *Manager) Ensure(writer http.ResponseWriter, request *http.Request
return current.csrfToken, false
}
// IsAuthenticated 只读检查当前请求是否持有有效管理会话;它不会像 Ensure 一样创建匿名会话。
func (manager *Manager) IsAuthenticated(request *http.Request) bool {
_, current, found := manager.current(request)
return found && current.authenticated
}
// VerifyCSRF 只接受当前未过期会话中以恒定时间比较匹配的 token。
func (manager *Manager) VerifyCSRF(request *http.Request, token string) (authenticated bool, ok bool) {
_, current, found := manager.current(request)
+45
View File
@@ -37,6 +37,51 @@ func TestManagerRejectsTamperedAndExpiredCookies(t *testing.T) {
}
}
func TestIsAuthenticatedDoesNotCreateOrDependOnCSRFValidation(t *testing.T) {
manager := NewManager([]byte(strings.Repeat("s", 32)), false)
missingSession := httptest.NewRequest(http.MethodPost, "/tasks/start-purchases", nil)
if manager.IsAuthenticated(missingSession) {
t.Fatal("missing session was treated as authenticated")
}
if len(manager.sessions) != 0 {
t.Fatalf("read-only authentication check created %d sessions", len(manager.sessions))
}
anonymousRequest := httptest.NewRequest(http.MethodGet, "/login", nil)
anonymousResponse := httptest.NewRecorder()
manager.Ensure(anonymousResponse, anonymousRequest)
anonymousCookie := anonymousResponse.Result().Cookies()[0]
anonymousCheck := httptest.NewRequest(http.MethodPost, "/tasks/start-purchases", nil)
anonymousCheck.AddCookie(anonymousCookie)
if manager.IsAuthenticated(anonymousCheck) {
t.Fatal("anonymous CSRF session was treated as authenticated")
}
loginRequest := httptest.NewRequest(http.MethodPost, "/login", nil)
loginRequest.AddCookie(anonymousCookie)
authenticatedResponse := httptest.NewRecorder()
csrf := manager.RotateAuthenticated(authenticatedResponse, loginRequest)
authenticatedCookie := authenticatedResponse.Result().Cookies()[0]
authenticatedCheck := httptest.NewRequest(http.MethodPost, "/tasks/start-purchases", nil)
authenticatedCheck.AddCookie(authenticatedCookie)
if !manager.IsAuthenticated(authenticatedCheck) {
t.Fatal("valid authenticated session was not recognized")
}
if authenticated, csrfOK := manager.VerifyCSRF(authenticatedCheck, "wrong-token"); authenticated || csrfOK {
t.Fatalf("wrong token result = (%t, %t), want (false, false)", authenticated, csrfOK)
}
validRequest := httptest.NewRequest(http.MethodPost, "/tasks/start-purchases", nil)
validRequest.AddCookie(authenticatedCookie)
if authenticated, csrfOK := manager.VerifyCSRF(validRequest, csrf); !authenticated || !csrfOK {
t.Fatalf("valid token result = (%t, %t), want (true, true)", authenticated, csrfOK)
}
if authenticated, csrfOK := manager.VerifyCSRF(httptest.NewRequest(http.MethodPost, "/tasks/start-purchases", nil), csrf); authenticated || csrfOK {
t.Fatalf("missing session result = (%t, %t), want (false, false)", authenticated, csrfOK)
}
}
func flipCookieValue(t *testing.T, value string) string {
t.Helper()
if value == "" {
+94
View File
@@ -2,10 +2,15 @@
package config
import (
"bytes"
"encoding/hex"
"errors"
"fmt"
"os"
"path/filepath"
"strconv"
"strings"
"time"
"golang.org/x/crypto/bcrypt"
)
@@ -15,6 +20,13 @@ const (
adminPasswordBcryptEnv = "CMBUYER_ADMIN_PASSWORD_BCRYPT"
sessionSecretEnv = "CMBUYER_SESSION_SECRET"
cookieSecureEnv = "CMBUYER_COOKIE_SECURE"
databaseSourceEnv = "CMBUYER_DATABASE_SOURCE"
authorizationTTLEnv = "CMBUYER_AUTHORIZATION_TTL"
maxTaskQuantityEnv = "CMBUYER_MAX_TASK_QUANTITY"
maxTotalPriceEnv = "CMBUYER_MAX_TOTAL_PRICE"
evidenceDirectoryEnv = "CMBUYER_EVIDENCE_DIR"
claimTokenSecretEnv = "CMBUYER_CLAIM_TOKEN_SECRET"
claimLeaseTTLEnv = "CMBUYER_CLAIM_LEASE_TTL"
minimumSecretLength = 32
)
@@ -24,6 +36,13 @@ type Config struct {
AdminPasswordBcrypt string
SessionSecret []byte
CookieSecure bool
DatabaseSource string
AuthorizationTTL time.Duration
MaxTaskQuantity int
MaxTotalPrice string
EvidenceDirectory string
ClaimTokenSecret []byte
ClaimLeaseTTL time.Duration
}
// LoadFromEnv 从进程环境读取配置。错误只指出缺失或非法的变量名,绝不回显秘密。
@@ -65,15 +84,90 @@ func Load(lookup func(string) (string, bool)) (Config, error) {
return Config{}, fmt.Errorf("%s must be exactly true or false", cookieSecureEnv)
}
}
databaseSource, err := required(lookup, databaseSourceEnv)
if err != nil {
return Config{}, err
}
ttlText, err := required(lookup, authorizationTTLEnv)
if err != nil {
return Config{}, err
}
ttl, err := time.ParseDuration(ttlText)
if err != nil || ttl <= 0 {
return Config{}, fmt.Errorf("%s must be a positive duration", authorizationTTLEnv)
}
quantityText, err := required(lookup, maxTaskQuantityEnv)
if err != nil {
return Config{}, err
}
maxQuantity, err := strconv.Atoi(quantityText)
if err != nil || maxQuantity < 1 {
return Config{}, fmt.Errorf("%s must be a positive integer", maxTaskQuantityEnv)
}
maxPrice, err := required(lookup, maxTotalPriceEnv)
if err != nil {
return Config{}, err
}
if !canonicalMoney(maxPrice) {
return Config{}, fmt.Errorf("%s must be a canonical positive decimal", maxTotalPriceEnv)
}
evidenceDirectory, err := required(lookup, evidenceDirectoryEnv)
if err != nil {
return Config{}, err
}
if strings.TrimSpace(evidenceDirectory) != evidenceDirectory || !filepath.IsAbs(evidenceDirectory) {
return Config{}, fmt.Errorf("%s must be an absolute path without surrounding whitespace", evidenceDirectoryEnv)
}
claimSecretText, err := required(lookup, claimTokenSecretEnv)
if err != nil {
return Config{}, err
}
claimSecret, err := hex.DecodeString(claimSecretText)
if err != nil || len(claimSecret) != 32 || hex.EncodeToString(claimSecret) != claimSecretText {
return Config{}, fmt.Errorf("%s must be exactly 64 lowercase hexadecimal characters", claimTokenSecretEnv)
}
// Claim ownership, admin sessions and device authentication are separate security domains.
// Reject both identical configuration text and identical effective key bytes.
if claimSecretText == secret || bytes.Equal(claimSecret, []byte(secret)) {
return Config{}, fmt.Errorf("%s must be isolated from %s", claimTokenSecretEnv, sessionSecretEnv)
}
claimTTLText, err := required(lookup, claimLeaseTTLEnv)
if err != nil {
return Config{}, err
}
claimTTL, err := time.ParseDuration(claimTTLText)
if err != nil || claimTTL <= 0 || claimTTL >= ttl {
return Config{}, fmt.Errorf("%s must be positive and shorter than %s", claimLeaseTTLEnv, authorizationTTLEnv)
}
return Config{
AdminUsername: username,
AdminPasswordBcrypt: passwordHash,
SessionSecret: []byte(secret),
CookieSecure: cookieSecure,
DatabaseSource: databaseSource,
AuthorizationTTL: ttl, MaxTaskQuantity: maxQuantity, MaxTotalPrice: maxPrice,
EvidenceDirectory: evidenceDirectory,
ClaimTokenSecret: claimSecret,
ClaimLeaseTTL: claimTTL,
}, nil
}
func canonicalMoney(value string) bool {
parts := strings.Split(value, ".")
if len(parts) != 2 || len(parts[0]) == 0 || len(parts[1]) != 2 || (len(parts[0]) > 1 && parts[0][0] == '0') {
return false
}
for _, part := range parts {
for _, ch := range part {
if ch < '0' || ch > '9' {
return false
}
}
}
return strings.Trim(parts[0]+parts[1], "0") != ""
}
func required(lookup func(string) (string, bool), name string) (string, error) {
value, present := lookup(name)
if !present || strings.TrimSpace(value) == "" {
+28 -1
View File
@@ -3,6 +3,7 @@ package config_test
import (
"strings"
"testing"
"time"
"cmbuyer/admin/internal/config"
@@ -20,13 +21,20 @@ func TestLoad(t *testing.T) {
"CMBUYER_ADMIN_PASSWORD_BCRYPT": string(hash),
"CMBUYER_SESSION_SECRET": strings.Repeat("s", 32),
"CMBUYER_COOKIE_SECURE": "true",
"CMBUYER_DATABASE_SOURCE": ":memory:",
"CMBUYER_AUTHORIZATION_TTL": "10m",
"CMBUYER_MAX_TASK_QUANTITY": "99",
"CMBUYER_MAX_TOTAL_PRICE": "999.99",
"CMBUYER_EVIDENCE_DIR": t.TempDir(),
"CMBUYER_CLAIM_TOKEN_SECRET": strings.Repeat("ab", 32),
"CMBUYER_CLAIM_LEASE_TTL": "1m",
}
got, err := config.Load(lookup(values))
if err != nil {
t.Fatalf("Load: %v", err)
}
if got.AdminUsername != "admin" || !got.CookieSecure {
if got.AdminUsername != "admin" || !got.CookieSecure || len(got.ClaimTokenSecret) != 32 || got.ClaimLeaseTTL != time.Minute {
t.Fatalf("Load returned unexpected public configuration: %#v", got)
}
}
@@ -41,6 +49,13 @@ func TestLoadRejectsMissingOrInvalidConfiguration(t *testing.T) {
"CMBUYER_ADMIN_USERNAME": "admin",
"CMBUYER_ADMIN_PASSWORD_BCRYPT": string(hash),
"CMBUYER_SESSION_SECRET": strings.Repeat("s", 32),
"CMBUYER_DATABASE_SOURCE": ":memory:",
"CMBUYER_AUTHORIZATION_TTL": "10m",
"CMBUYER_MAX_TASK_QUANTITY": "99",
"CMBUYER_MAX_TOTAL_PRICE": "999.99",
"CMBUYER_EVIDENCE_DIR": t.TempDir(),
"CMBUYER_CLAIM_TOKEN_SECRET": strings.Repeat("ab", 32),
"CMBUYER_CLAIM_LEASE_TTL": "1m",
}
tests := []struct {
@@ -52,6 +67,18 @@ func TestLoadRejectsMissingOrInvalidConfiguration(t *testing.T) {
{"invalid bcrypt", func(values map[string]string) { values["CMBUYER_ADMIN_PASSWORD_BCRYPT"] = "not-a-bcrypt-hash" }, "CMBUYER_ADMIN_PASSWORD_BCRYPT"},
{"short secret", func(values map[string]string) { values["CMBUYER_SESSION_SECRET"] = "short" }, "CMBUYER_SESSION_SECRET"},
{"invalid secure flag", func(values map[string]string) { values["CMBUYER_COOKIE_SECURE"] = "1" }, "CMBUYER_COOKIE_SECURE"},
{"missing database", func(values map[string]string) { delete(values, "CMBUYER_DATABASE_SOURCE") }, "CMBUYER_DATABASE_SOURCE"},
{"invalid authorization ttl", func(values map[string]string) { values["CMBUYER_AUTHORIZATION_TTL"] = "0s" }, "CMBUYER_AUTHORIZATION_TTL"},
{"invalid maximum quantity", func(values map[string]string) { values["CMBUYER_MAX_TASK_QUANTITY"] = "0" }, "CMBUYER_MAX_TASK_QUANTITY"},
{"invalid maximum total price", func(values map[string]string) { values["CMBUYER_MAX_TOTAL_PRICE"] = "1" }, "CMBUYER_MAX_TOTAL_PRICE"},
{"missing evidence directory", func(values map[string]string) { delete(values, "CMBUYER_EVIDENCE_DIR") }, "CMBUYER_EVIDENCE_DIR"},
{"relative evidence directory", func(values map[string]string) { values["CMBUYER_EVIDENCE_DIR"] = "evidence" }, "CMBUYER_EVIDENCE_DIR"},
{"invalid claim secret", func(values map[string]string) { values["CMBUYER_CLAIM_TOKEN_SECRET"] = strings.Repeat("A", 64) }, "CMBUYER_CLAIM_TOKEN_SECRET"},
{"claim secret same raw session secret", func(values map[string]string) {
values["CMBUYER_SESSION_SECRET"] = values["CMBUYER_CLAIM_TOKEN_SECRET"]
}, "CMBUYER_CLAIM_TOKEN_SECRET"},
{"claim secret same decoded session secret", func(values map[string]string) { values["CMBUYER_SESSION_SECRET"] = strings.Repeat("\xab", 32) }, "CMBUYER_CLAIM_TOKEN_SECRET"},
{"invalid claim lease ttl", func(values map[string]string) { values["CMBUYER_CLAIM_LEASE_TTL"] = "10m" }, "CMBUYER_CLAIM_LEASE_TTL"},
}
for _, test := range tests {
+212
View File
@@ -0,0 +1,212 @@
package deviceauth
import (
"context"
"crypto/rand"
"crypto/sha256"
"database/sql"
"encoding/hex"
"errors"
"fmt"
"io"
"strings"
"sync"
"time"
"unicode"
)
const (
StatusActive = "ACTIVE"
StatusRevoked = "REVOKED"
)
var (
ErrInvalidCredential = errors.New("invalid device credential input")
ErrCredentialNotFound = errors.New("device credential not found")
)
type Credential struct {
DeviceID string `json:"device_id"`
DisplayName string `json:"display_name"`
Status string `json:"status"`
CreatedAt time.Time `json:"created_at"`
RevokedAt *time.Time `json:"revoked_at,omitempty"`
}
// IssuedCredential is the only value that can carry the plaintext token. It is returned only
// after SQLite has committed the hash and is intended for the management CLI's one stdout write.
type IssuedCredential struct {
Credential
Token string `json:"-"`
}
type CredentialStore struct {
database *sql.DB
now func() time.Time
random io.Reader
randomMu sync.Mutex
}
func NewCredentialStore(database *sql.DB) (*CredentialStore, error) {
if database == nil {
return nil, errors.New("device credential database is required")
}
if _, err := database.Exec("SELECT device_id FROM device_credentials LIMIT 1"); err != nil {
return nil, errors.New("device credential migration is not available")
}
return &CredentialStore{database: database, now: time.Now, random: rand.Reader}, nil
}
func (store *CredentialStore) Issue(ctx context.Context, displayName string) (IssuedCredential, error) {
if !ValidDisplayName(displayName) {
return IssuedCredential{}, ErrInvalidCredential
}
randomBytes := make([]byte, 16+32)
store.randomMu.Lock()
_, randomErr := io.ReadFull(store.random, randomBytes)
store.randomMu.Unlock()
if randomErr != nil {
return IssuedCredential{}, fmt.Errorf("generate device credential: %w", randomErr)
}
deviceID := formatUUIDv4(randomBytes[:16])
token := hex.EncodeToString(randomBytes[16:])
tokenHash := sha256.Sum256(randomBytes[16:])
createdAt := store.now().UTC()
if createdAt.IsZero() {
return IssuedCredential{}, errors.New("device credential clock is invalid")
}
_, err := store.database.ExecContext(ctx, `INSERT INTO device_credentials
(device_id, display_name, token_sha256, status, created_at, revoked_at)
VALUES (?, ?, ?, ?, ?, NULL)`,
deviceID, displayName, tokenHash[:], StatusActive, createdAt.Format(time.RFC3339Nano))
if err != nil {
return IssuedCredential{}, fmt.Errorf("persist device credential: %w", err)
}
return IssuedCredential{Credential: Credential{
DeviceID: deviceID, DisplayName: displayName, Status: StatusActive, CreatedAt: createdAt,
}, Token: token}, nil
}
func (store *CredentialStore) List(ctx context.Context) ([]Credential, error) {
rows, err := store.database.QueryContext(ctx, `SELECT device_id, display_name, status, created_at, revoked_at
FROM device_credentials ORDER BY created_at, device_id`)
if err != nil {
return nil, fmt.Errorf("list device credentials: %w", err)
}
defer rows.Close()
credentials := make([]Credential, 0)
for rows.Next() {
credential, err := scanCredential(rows)
if err != nil {
return nil, err
}
credentials = append(credentials, credential)
}
if err := rows.Err(); err != nil {
return nil, fmt.Errorf("list device credentials: %w", err)
}
return credentials, nil
}
func (store *CredentialStore) Revoke(ctx context.Context, deviceID string) (Credential, bool, error) {
if !ValidDeviceID(deviceID) {
return Credential{}, false, ErrInvalidCredential
}
revokedAt := store.now().UTC()
if revokedAt.IsZero() {
return Credential{}, false, errors.New("device credential clock is invalid")
}
transaction, err := store.database.BeginTx(ctx, nil)
if err != nil {
return Credential{}, false, fmt.Errorf("begin device credential revocation: %w", err)
}
defer transaction.Rollback()
result, err := transaction.ExecContext(ctx, `UPDATE device_credentials
SET status = ?, revoked_at = ? WHERE device_id = ? AND status = ?`,
StatusRevoked, revokedAt.Format(time.RFC3339Nano), deviceID, StatusActive)
if err != nil {
return Credential{}, false, fmt.Errorf("revoke device credential: %w", err)
}
changedRows, err := result.RowsAffected()
if err != nil {
return Credential{}, false, fmt.Errorf("inspect device credential revocation: %w", err)
}
credential, err := scanCredential(transaction.QueryRowContext(ctx, `SELECT device_id, display_name, status, created_at, revoked_at
FROM device_credentials WHERE device_id = ?`, deviceID))
if errors.Is(err, sql.ErrNoRows) {
return Credential{}, false, ErrCredentialNotFound
}
if err != nil {
return Credential{}, false, err
}
if err := transaction.Commit(); err != nil {
return Credential{}, false, fmt.Errorf("commit device credential revocation: %w", err)
}
return credential, changedRows == 1, nil
}
type rowScanner interface {
Scan(...any) error
}
func scanCredential(row rowScanner) (Credential, error) {
var credential Credential
var created string
var revoked sql.NullString
if err := row.Scan(&credential.DeviceID, &credential.DisplayName, &credential.Status, &created, &revoked); err != nil {
return Credential{}, err
}
if !ValidDeviceID(credential.DeviceID) || !ValidDisplayName(credential.DisplayName) || (credential.Status != StatusActive && credential.Status != StatusRevoked) {
return Credential{}, errors.New("stored device credential metadata is invalid")
}
createdAt, err := parseStoredTime(created)
if err != nil {
return Credential{}, err
}
credential.CreatedAt = createdAt
if revoked.Valid {
revokedAt, err := parseStoredTime(revoked.String)
if err != nil {
return Credential{}, err
}
if revokedAt.Before(createdAt) {
return Credential{}, errors.New("stored device credential status is invalid")
}
credential.RevokedAt = &revokedAt
}
if (credential.Status == StatusActive) != (credential.RevokedAt == nil) {
return Credential{}, errors.New("stored device credential status is invalid")
}
return credential, nil
}
func ValidDisplayName(value string) bool {
if value == "" || len([]rune(value)) > 128 || strings.TrimSpace(value) != value {
return false
}
for _, character := range value {
if unicode.IsControl(character) {
return false
}
}
return true
}
func parseStoredTime(value string) (time.Time, error) {
if strings.TrimSpace(value) != value || !strings.HasSuffix(value, "Z") {
return time.Time{}, errors.New("stored device credential time is invalid")
}
parsed, err := time.Parse(time.RFC3339Nano, value)
if err != nil || parsed.Location() != time.UTC {
return time.Time{}, errors.New("stored device credential time is invalid")
}
return parsed, nil
}
func formatUUIDv4(bytes []byte) string {
copyBytes := append([]byte(nil), bytes...)
copyBytes[6] = (copyBytes[6] & 0x0f) | 0x40
copyBytes[8] = (copyBytes[8] & 0x3f) | 0x80
encoded := hex.EncodeToString(copyBytes)
return encoded[:8] + "-" + encoded[8:12] + "-" + encoded[12:16] + "-" + encoded[16:20] + "-" + encoded[20:]
}
+168
View File
@@ -0,0 +1,168 @@
// Package deviceauth owns the machine identity boundary shared by all device routes.
package deviceauth
import (
"crypto/sha256"
"crypto/subtle"
"database/sql"
"encoding/hex"
"errors"
"net/http"
"strings"
)
const (
AuthorizationHeader = "Authorization"
DeviceIDHeader = "X-CMBuyer-Device-ID"
tokenHexLength = 64
)
var (
// ErrUnauthenticated deliberately covers every credential defect. Callers must not reveal
// whether a device exists, is revoked, or supplied a mismatched token.
ErrUnauthenticated = errors.New("device authentication failed")
// ErrUnavailable is distinct so a storage outage is not disguised as a bad credential.
// HTTP callers still return no diagnostic body because database details are server-only.
ErrUnavailable = errors.New("device authentication unavailable")
)
type Principal struct {
ID string
}
type Authenticator interface {
Authenticate(*http.Request) (Principal, error)
}
// RejectAllAuthenticator is useful for tests and for fail-closed wiring where no credential
// store is available. Production startup uses SQLiteAuthenticator.
type RejectAllAuthenticator struct{}
func (RejectAllAuthenticator) Authenticate(*http.Request) (Principal, error) {
return Principal{}, ErrUnauthenticated
}
type SQLiteAuthenticator struct {
database *sql.DB
}
func NewSQLiteAuthenticator(database *sql.DB) (*SQLiteAuthenticator, error) {
if database == nil {
return nil, errors.New("device credential database is required")
}
if _, err := database.Exec("SELECT device_id FROM device_credentials LIMIT 1"); err != nil {
return nil, errors.New("device credential migration is not available")
}
return &SQLiteAuthenticator{database: database}, nil
}
func (authenticator *SQLiteAuthenticator) Authenticate(request *http.Request) (Principal, error) {
if request == nil {
return Principal{}, ErrUnauthenticated
}
deviceID, token, ok := requestCredentials(request)
if !ok {
return Principal{}, ErrUnauthenticated
}
candidateHash := sha256.Sum256(token)
var storedHash []byte
var hashType string
var hashLength sql.NullInt64
var status sql.NullString
var revokedAt sql.NullString
var found bool
err := authenticator.database.QueryRowContext(
request.Context(),
`SELECT CASE WHEN credentials.device_id IS NULL THEN zeroblob(32) ELSE credentials.token_sha256 END,
typeof(credentials.token_sha256),
length(credentials.token_sha256),
credentials.status,
credentials.revoked_at,
credentials.device_id IS NOT NULL
FROM (SELECT 1) AS singleton
LEFT JOIN device_credentials AS credentials ON credentials.device_id = ?`,
deviceID,
).Scan(&storedHash, &hashType, &hashLength, &status, &revokedAt, &found)
if err != nil {
return Principal{}, ErrUnavailable
}
if len(storedHash) != sha256.Size {
return Principal{}, ErrUnavailable
}
matched := subtle.ConstantTimeCompare(candidateHash[:], storedHash) == 1
if !found {
// The LEFT JOIN supplies a 32-byte dummy hash, so unknown ids take the same compare path
// as known credentials without requiring a plaintext token lookup.
return Principal{}, ErrUnauthenticated
}
if hashType != "blob" || !hashLength.Valid || hashLength.Int64 != sha256.Size || len(storedHash) != sha256.Size || !status.Valid {
return Principal{}, ErrUnavailable
}
switch status.String {
case StatusActive:
if revokedAt.Valid {
return Principal{}, ErrUnavailable
}
case StatusRevoked:
if !revokedAt.Valid {
return Principal{}, ErrUnavailable
}
if _, err := parseStoredTime(revokedAt.String); err != nil {
return Principal{}, ErrUnavailable
}
default:
return Principal{}, ErrUnavailable
}
if !matched || status.String == StatusRevoked {
return Principal{}, ErrUnauthenticated
}
return Principal{ID: deviceID}, nil
}
func requestCredentials(request *http.Request) (string, []byte, bool) {
authorizations := request.Header.Values(AuthorizationHeader)
deviceIDs := request.Header.Values(DeviceIDHeader)
if len(authorizations) != 1 || len(deviceIDs) != 1 {
return "", nil, false
}
authorization := authorizations[0]
if len(authorization) != len("Bearer ")+tokenHexLength || !strings.EqualFold(authorization[:len("Bearer")], "Bearer") || authorization[len("Bearer")] != ' ' {
return "", nil, false
}
tokenHex := authorization[len("Bearer "):]
if !validLowerHex(tokenHex, tokenHexLength) || !ValidDeviceID(deviceIDs[0]) {
return "", nil, false
}
token, err := hex.DecodeString(tokenHex)
if err != nil {
return "", nil, false
}
return deviceIDs[0], token, true
}
func ValidDeviceID(value string) bool {
if len(value) != 36 {
return false
}
for index, character := range value {
if index == 8 || index == 13 || index == 18 || index == 23 {
if character != '-' {
return false
}
continue
}
if !(character >= '0' && character <= '9' || character >= 'a' && character <= 'f') {
return false
}
}
return value[14] == '4' && (value[19] == '8' || value[19] == '9' || value[19] == 'a' || value[19] == 'b')
}
func validLowerHex(value string, length int) bool {
if len(value) != length {
return false
}
decoded, err := hex.DecodeString(value)
return err == nil && hex.EncodeToString(decoded) == value
}
@@ -0,0 +1,407 @@
package deviceauth
import (
"context"
"crypto/rand"
"crypto/sha256"
"database/sql"
"encoding/hex"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"path/filepath"
"runtime"
"strings"
"sync"
"testing"
"time"
"cmbuyer/admin/internal/migrations"
"cmbuyer/admin/internal/storage/sqlite"
)
func TestIssueStoresOnlyRawTokenHashAndGenericJSONOmitsSecret(t *testing.T) {
database, store := newCredentialStore(t)
first, err := store.Issue(context.Background(), "采购工具一号")
if err != nil {
t.Fatalf("Issue first: %v", err)
}
second, err := store.Issue(context.Background(), "采购工具二号")
if err != nil {
t.Fatalf("Issue second: %v", err)
}
if first.DeviceID == second.DeviceID || first.Token == second.Token || !ValidDeviceID(first.DeviceID) || !validLowerHex(first.Token, tokenHexLength) {
t.Fatalf("issued identifiers are not independent canonical values")
}
rawToken, err := hex.DecodeString(first.Token)
if err != nil {
t.Fatalf("decode issued token: %v", err)
}
wantHash := sha256.Sum256(rawToken)
var storedHash []byte
var storageType string
if err := database.QueryRow(`SELECT token_sha256, typeof(token_sha256) FROM device_credentials WHERE device_id = ?`, first.DeviceID).Scan(&storedHash, &storageType); err != nil {
t.Fatalf("read stored hash: %v", err)
}
if storageType != "blob" || len(storedHash) != sha256.Size || !equalBytes(storedHash, wantHash[:]) {
t.Fatalf("stored hash type/length/value = %q/%d/%t", storageType, len(storedHash), equalBytes(storedHash, wantHash[:]))
}
var leakedCopies int
if err := database.QueryRow(`SELECT COUNT(*) FROM device_credentials WHERE CAST(token_sha256 AS TEXT) IN (?, ?)`, first.Token, hex.EncodeToString(wantHash[:])).Scan(&leakedCopies); err != nil {
t.Fatalf("search token copies: %v", err)
}
if leakedCopies != 0 {
t.Fatal("database stored a plaintext or hex-encoded token/hash copy")
}
encoded, err := json.Marshal(first)
if err != nil {
t.Fatalf("marshal issued credential: %v", err)
}
if strings.Contains(string(encoded), first.Token) || strings.Contains(string(encoded), "token") {
t.Fatalf("generic serialization disclosed token field: %s", encoded)
}
}
func TestAuthenticateStrictHeaderMatrixAndBinding(t *testing.T) {
_, store := newCredentialStore(t)
first, err := store.Issue(context.Background(), "one")
if err != nil {
t.Fatalf("issue first: %v", err)
}
second, err := store.Issue(context.Background(), "two")
if err != nil {
t.Fatalf("issue second: %v", err)
}
authenticator := authenticatorForStore(t, store)
for _, scheme := range []string{"Bearer", "bearer", "BEARER"} {
request := credentialRequest(first.DeviceID, scheme+" "+first.Token)
principal, err := authenticator.Authenticate(request)
if err != nil || principal.ID != first.DeviceID {
t.Fatalf("scheme %q Authenticate = (%q, %v)", scheme, principal.ID, err)
}
}
unknownID := newRuntimeUUID(t)
wrongToken := newRuntimeToken(t)
tests := []struct {
name string
mutate func(*http.Request)
}{
{"missing authorization", func(request *http.Request) { request.Header.Del(AuthorizationHeader) }},
{"missing device", func(request *http.Request) { request.Header.Del(DeviceIDHeader) }},
{"empty authorization", func(request *http.Request) { request.Header.Set(AuthorizationHeader, "") }},
{"empty device", func(request *http.Request) { request.Header.Set(DeviceIDHeader, "") }},
{"duplicate authorization", func(request *http.Request) { request.Header.Add(AuthorizationHeader, "Bearer "+first.Token) }},
{"duplicate device", func(request *http.Request) { request.Header.Add(DeviceIDHeader, first.DeviceID) }},
{"combined authorization", func(request *http.Request) {
request.Header.Set(AuthorizationHeader, "Bearer "+first.Token+", Bearer "+first.Token)
}},
{"combined device", func(request *http.Request) { request.Header.Set(DeviceIDHeader, first.DeviceID+", "+first.DeviceID) }},
{"extra separator", func(request *http.Request) { request.Header.Set(AuthorizationHeader, "Bearer "+first.Token) }},
{"tab separator", func(request *http.Request) { request.Header.Set(AuthorizationHeader, "Bearer\t"+first.Token) }},
{"uppercase token", func(request *http.Request) {
request.Header.Set(AuthorizationHeader, "Bearer "+strings.ToUpper(first.Token))
}},
{"short token", func(request *http.Request) { request.Header.Set(AuthorizationHeader, "Bearer "+first.Token[:62]) }},
{"long token", func(request *http.Request) { request.Header.Set(AuthorizationHeader, "Bearer "+first.Token+"00") }},
{"non hex token", func(request *http.Request) { request.Header.Set(AuthorizationHeader, "Bearer "+first.Token[:63]+"g") }},
{"token separator", func(request *http.Request) {
request.Header.Set(AuthorizationHeader, "Bearer "+first.Token[:32]+"-"+first.Token[33:])
}},
{"uppercase device", func(request *http.Request) { request.Header.Set(DeviceIDHeader, strings.ToUpper(first.DeviceID)) }},
{"padded device", func(request *http.Request) { request.Header.Set(DeviceIDHeader, " "+first.DeviceID) }},
{"wrong uuid version", func(request *http.Request) {
request.Header.Set(DeviceIDHeader, first.DeviceID[:14]+"3"+first.DeviceID[15:])
}},
{"wrong uuid variant", func(request *http.Request) {
request.Header.Set(DeviceIDHeader, first.DeviceID[:19]+"7"+first.DeviceID[20:])
}},
{"unknown device", func(request *http.Request) { request.Header.Set(DeviceIDHeader, unknownID) }},
{"wrong token", func(request *http.Request) { request.Header.Set(AuthorizationHeader, "Bearer "+wrongToken) }},
{"token device mismatch", func(request *http.Request) { request.Header.Set(DeviceIDHeader, second.DeviceID) }},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
request := credentialRequest(first.DeviceID, "Bearer "+first.Token)
test.mutate(request)
principal, err := authenticator.Authenticate(request)
if !errors.Is(err, ErrUnauthenticated) || principal != (Principal{}) {
t.Fatalf("Authenticate = (%#v, %v), want empty unauthenticated", principal, err)
}
})
}
if principal, err := authenticator.Authenticate(nil); !errors.Is(err, ErrUnauthenticated) || principal != (Principal{}) {
t.Fatalf("Authenticate(nil) = (%#v, %v)", principal, err)
}
}
func TestRevokeIsImmediateAndIdempotent(t *testing.T) {
_, store := newCredentialStore(t)
issued, err := store.Issue(context.Background(), "device")
if err != nil {
t.Fatalf("Issue: %v", err)
}
authenticator := authenticatorForStore(t, store)
request := credentialRequest(issued.DeviceID, "Bearer "+issued.Token)
if _, err := authenticator.Authenticate(request); err != nil {
t.Fatalf("Authenticate before revoke: %v", err)
}
first, changed, err := store.Revoke(context.Background(), issued.DeviceID)
if err != nil || !changed || first.Status != StatusRevoked || first.RevokedAt == nil {
t.Fatalf("first Revoke = (%#v, %t, %v)", first, changed, err)
}
if principal, err := authenticator.Authenticate(request); !errors.Is(err, ErrUnauthenticated) || principal != (Principal{}) {
t.Fatalf("Authenticate after committed revoke = (%#v, %v)", principal, err)
}
second, changed, err := store.Revoke(context.Background(), issued.DeviceID)
if err != nil || changed || second.RevokedAt == nil || !second.RevokedAt.Equal(*first.RevokedAt) {
t.Fatalf("second Revoke = (%#v, %t, %v)", second, changed, err)
}
listed, err := store.List(context.Background())
if err != nil || len(listed) != 1 || listed[0].Status != StatusRevoked {
t.Fatalf("List = (%#v, %v)", listed, err)
}
}
func TestConcurrentAuthenticationAndRevocation(t *testing.T) {
_, store := newCredentialStore(t)
issued, err := store.Issue(context.Background(), "concurrent")
if err != nil {
t.Fatalf("Issue: %v", err)
}
authenticator := authenticatorForStore(t, store)
request := func() *http.Request { return credentialRequest(issued.DeviceID, "Bearer "+issued.Token) }
start := make(chan struct{})
results := make(chan error, 16)
var wait sync.WaitGroup
for index := 0; index < 16; index++ {
wait.Add(1)
go func() {
defer wait.Done()
<-start
_, err := authenticator.Authenticate(request())
results <- err
}()
}
close(start)
if _, _, err := store.Revoke(context.Background(), issued.DeviceID); err != nil {
t.Fatalf("Revoke: %v", err)
}
wait.Wait()
close(results)
for err := range results {
if err != nil && !errors.Is(err, ErrUnauthenticated) {
t.Fatalf("concurrent Authenticate error = %v", err)
}
}
for index := 0; index < 16; index++ {
if _, err := authenticator.Authenticate(request()); !errors.Is(err, ErrUnauthenticated) {
t.Fatalf("post-commit Authenticate %d error = %v", index, err)
}
}
}
func TestAuthenticationDatabaseFaultAndCorruptionAreUnavailable(t *testing.T) {
database, store := newCredentialStore(t)
issued, err := store.Issue(context.Background(), "device")
if err != nil {
t.Fatalf("Issue: %v", err)
}
authenticator := authenticatorForStore(t, store)
if err := database.Close(); err != nil {
t.Fatalf("close database: %v", err)
}
if _, err := authenticator.Authenticate(credentialRequest(issued.DeviceID, "Bearer "+issued.Token)); !errors.Is(err, ErrUnavailable) {
t.Fatalf("closed database Authenticate error = %v", err)
}
corruptDB, err := sqlite.Open(filepath.Join(t.TempDir(), "corrupt.db"))
if err != nil {
t.Fatalf("open corrupt database: %v", err)
}
t.Cleanup(func() { _ = corruptDB.Close() })
if _, err := corruptDB.Exec(`CREATE TABLE device_credentials (device_id TEXT PRIMARY KEY, token_sha256 BLOB, status TEXT, revoked_at TEXT)`); err != nil {
t.Fatalf("create corrupt table: %v", err)
}
corruptAuthenticator, err := NewSQLiteAuthenticator(corruptDB)
if err != nil {
t.Fatalf("new corrupt authenticator: %v", err)
}
tests := []struct {
name string
hashValue func([sha256.Size]byte) any
status string
revokedAt any
}{
{name: "null hash", hashValue: func([sha256.Size]byte) any { return nil }, status: StatusActive},
{name: "matching text hash", hashValue: func(hash [sha256.Size]byte) any { return string(hash[:]) }, status: StatusActive},
{name: "unknown status", hashValue: func(hash [sha256.Size]byte) any { return hash[:] }, status: "BROKEN"},
{name: "active with revoked time", hashValue: func(hash [sha256.Size]byte) any { return hash[:] }, status: StatusActive, revokedAt: "2026-08-04T00:00:00Z"},
{name: "revoked without time", hashValue: func(hash [sha256.Size]byte) any { return hash[:] }, status: StatusRevoked},
{name: "revoked with invalid time", hashValue: func(hash [sha256.Size]byte) any { return hash[:] }, status: StatusRevoked, revokedAt: "not-a-time"},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
rawToken, token := newRuntimeTokenPair(t)
hash := sha256.Sum256(rawToken)
deviceID := newRuntimeUUID(t)
if _, err := corruptDB.Exec(`INSERT INTO device_credentials VALUES (?, ?, ?, ?)`, deviceID, test.hashValue(hash), test.status, test.revokedAt); err != nil {
t.Fatalf("insert corrupt row: %v", err)
}
principal, err := corruptAuthenticator.Authenticate(credentialRequest(deviceID, "Bearer "+token))
if !errors.Is(err, ErrUnavailable) || principal != (Principal{}) {
t.Fatalf("corrupt Authenticate = (%#v, %v), want unavailable", principal, err)
}
})
}
}
func TestCredentialAuthenticationSurvivesDatabaseReopen(t *testing.T) {
databaseSource := filepath.Join(t.TempDir(), "reopen.db")
database, err := sqlite.Open(databaseSource)
if err != nil {
t.Fatalf("open database: %v", err)
}
if err := migrations.Up(context.Background(), database, deviceMigrationDirectory(t)); err != nil {
_ = database.Close()
t.Fatalf("migrate database: %v", err)
}
store, err := NewCredentialStore(database)
if err != nil {
_ = database.Close()
t.Fatalf("new store: %v", err)
}
issued, err := store.Issue(context.Background(), "reopen")
if err != nil {
_ = database.Close()
t.Fatalf("issue: %v", err)
}
if err := database.Close(); err != nil {
t.Fatalf("close database: %v", err)
}
reopened, err := sqlite.Open(databaseSource)
if err != nil {
t.Fatalf("reopen database: %v", err)
}
defer reopened.Close()
authenticator, err := NewSQLiteAuthenticator(reopened)
if err != nil {
t.Fatalf("new reopened authenticator: %v", err)
}
principal, err := authenticator.Authenticate(credentialRequest(issued.DeviceID, "Bearer "+issued.Token))
if err != nil || principal.ID != issued.DeviceID {
t.Fatalf("Authenticate after reopen = (%#v, %v)", principal, err)
}
}
func TestCredentialInputAndMigrationAreRequired(t *testing.T) {
database, err := sqlite.Open(filepath.Join(t.TempDir(), "unmigrated.db"))
if err != nil {
t.Fatalf("open database: %v", err)
}
t.Cleanup(func() { _ = database.Close() })
if _, err := NewCredentialStore(database); err == nil {
t.Fatal("NewCredentialStore accepted an unmigrated database")
}
if _, err := NewSQLiteAuthenticator(database); err == nil {
t.Fatal("NewSQLiteAuthenticator accepted an unmigrated database")
}
_, store := newCredentialStore(t)
for _, name := range []string{"", " leading", "trailing ", "line\nbreak", strings.Repeat("名", 129)} {
if _, err := store.Issue(context.Background(), name); !errors.Is(err, ErrInvalidCredential) {
t.Fatalf("Issue(%q) error = %v", name, err)
}
}
if _, _, err := store.Revoke(context.Background(), "not-a-uuid"); !errors.Is(err, ErrInvalidCredential) {
t.Fatalf("Revoke invalid id error = %v", err)
}
if _, _, err := store.Revoke(context.Background(), newRuntimeUUID(t)); !errors.Is(err, ErrCredentialNotFound) {
t.Fatalf("Revoke unknown id error = %v", err)
}
}
func newCredentialStore(t *testing.T) (*sql.DB, *CredentialStore) {
t.Helper()
databaseSource := filepath.Join(t.TempDir(), "device-auth.db") + "?_busy_timeout=5000&_journal_mode=WAL"
database, err := sqlite.Open(databaseSource)
if err != nil {
t.Fatalf("open database: %v", err)
}
t.Cleanup(func() { _ = database.Close() })
if err := migrations.Up(context.Background(), database, deviceMigrationDirectory(t)); err != nil {
t.Fatalf("migrate database: %v", err)
}
store, err := NewCredentialStore(database)
if err != nil {
t.Fatalf("NewCredentialStore: %v", err)
}
store.now = func() time.Time { return time.Date(2026, 8, 4, 12, 0, 0, 123, time.UTC) }
return database, store
}
func authenticatorForStore(t *testing.T, store *CredentialStore) *SQLiteAuthenticator {
t.Helper()
authenticator, err := NewSQLiteAuthenticator(store.database)
if err != nil {
t.Fatalf("NewSQLiteAuthenticator: %v", err)
}
return authenticator
}
func credentialRequest(deviceID, authorization string) *http.Request {
request := httptest.NewRequest(http.MethodPost, "/api/v1/tasks/id/evidence", nil)
request.Header.Set(DeviceIDHeader, deviceID)
request.Header.Set(AuthorizationHeader, authorization)
return request
}
func newRuntimeToken(t *testing.T) string {
t.Helper()
_, token := newRuntimeTokenPair(t)
return token
}
func newRuntimeTokenPair(t *testing.T) ([]byte, string) {
t.Helper()
raw := make([]byte, 32)
if _, err := rand.Read(raw); err != nil {
t.Fatalf("generate runtime token: %v", err)
}
return raw, hex.EncodeToString(raw)
}
func newRuntimeUUID(t *testing.T) string {
t.Helper()
raw := make([]byte, 16)
if _, err := rand.Read(raw); err != nil {
t.Fatalf("generate runtime UUID: %v", err)
}
return formatUUIDv4(raw)
}
func equalBytes(left, right []byte) bool {
if len(left) != len(right) {
return false
}
for index := range left {
if left[index] != right[index] {
return false
}
}
return true
}
func deviceMigrationDirectory(t *testing.T) string {
t.Helper()
_, file, _, ok := runtime.Caller(0)
if !ok {
t.Fatal("locate migrations")
}
return filepath.Join(filepath.Dir(file), "..", "..", "migrations")
}
+10 -24
View File
@@ -10,28 +10,24 @@ var ErrInvalidAuthorizationTransition = errors.New("invalid authorization status
type AuthorizationStatus string
const (
AuthorizationStatusPendingDelivery AuthorizationStatus = "PENDING_DELIVERY"
AuthorizationStatusDelivered AuthorizationStatus = "DELIVERED"
AuthorizationStatusAcknowledged AuthorizationStatus = "ACKNOWLEDGED"
AuthorizationStatusExecuting AuthorizationStatus = "EXECUTING"
AuthorizationStatusActive AuthorizationStatus = "ACTIVE"
AuthorizationStatusClaimed AuthorizationStatus = "CLAIMED"
AuthorizationStatusFenced AuthorizationStatus = "FENCED"
AuthorizationStatusConsumed AuthorizationStatus = "CONSUMED"
AuthorizationStatusSuperseded AuthorizationStatus = "SUPERSEDED"
AuthorizationStatusExpired AuthorizationStatus = "EXPIRED"
AuthorizationStatusAbandoned AuthorizationStatus = "ABANDONED"
)
type OrderAuthorization struct {
ID string
TaskID string
SpecTrialID string
Version int
TaskVersion int
StartKey string
GoodsID string
SKUColor string
SKUSize string
Quantity int
AuthorizedUnitPrice string
TotalPriceCap string
Note *string
Status AuthorizationStatus
CreatedBy string
CreatedAt time.Time
@@ -54,25 +50,15 @@ func TransitionAuthorization(current, next AuthorizationStatus) (AuthorizationSt
}
var authorizationTransitions = map[AuthorizationStatus]map[AuthorizationStatus]struct{}{
AuthorizationStatusPendingDelivery: {
AuthorizationStatusDelivered: {},
AuthorizationStatusSuperseded: {},
AuthorizationStatusActive: {
AuthorizationStatusClaimed: {},
AuthorizationStatusExpired: {},
AuthorizationStatusAbandoned: {},
},
AuthorizationStatusDelivered: {
AuthorizationStatusAcknowledged: {},
AuthorizationStatusSuperseded: {},
AuthorizationStatusExpired: {},
},
AuthorizationStatusAcknowledged: {
AuthorizationStatusExecuting: {},
AuthorizationStatusSuperseded: {},
AuthorizationStatusExpired: {},
},
AuthorizationStatusExecuting: {
AuthorizationStatusClaimed: {
AuthorizationStatusFenced: {},
AuthorizationStatusSuperseded: {},
AuthorizationStatusExpired: {},
AuthorizationStatusAbandoned: {},
},
AuthorizationStatusFenced: {
AuthorizationStatusConsumed: {},
+9 -11
View File
@@ -14,19 +14,17 @@ func TestAuthorizationTransitions(t *testing.T) {
next domain.AuthorizationStatus
allowed bool
}{
{"deliver", domain.AuthorizationStatusPendingDelivery, domain.AuthorizationStatusDelivered, true},
{"acknowledge", domain.AuthorizationStatusDelivered, domain.AuthorizationStatusAcknowledged, true},
{"execute", domain.AuthorizationStatusAcknowledged, domain.AuthorizationStatusExecuting, true},
{"fence", domain.AuthorizationStatusExecuting, domain.AuthorizationStatusFenced, true},
{"claim", domain.AuthorizationStatusActive, domain.AuthorizationStatusClaimed, true},
{"fence", domain.AuthorizationStatusClaimed, domain.AuthorizationStatusFenced, true},
{"consume fenced authorization", domain.AuthorizationStatusFenced, domain.AuthorizationStatusConsumed, true},
{"expire pending delivery", domain.AuthorizationStatusPendingDelivery, domain.AuthorizationStatusExpired, true},
{"supersede pending delivery", domain.AuthorizationStatusPendingDelivery, domain.AuthorizationStatusSuperseded, true},
{"expire before fence", domain.AuthorizationStatusExecuting, domain.AuthorizationStatusExpired, true},
{"supersede before fence", domain.AuthorizationStatusDelivered, domain.AuthorizationStatusSuperseded, true},
{"expire active", domain.AuthorizationStatusActive, domain.AuthorizationStatusExpired, true},
{"abandon active", domain.AuthorizationStatusActive, domain.AuthorizationStatusAbandoned, true},
{"expire claimed before fence", domain.AuthorizationStatusClaimed, domain.AuthorizationStatusExpired, true},
{"abandon claimed before fence", domain.AuthorizationStatusClaimed, domain.AuthorizationStatusAbandoned, true},
{"fenced authorization cannot expire", domain.AuthorizationStatusFenced, domain.AuthorizationStatusExpired, false},
{"fenced authorization cannot be superseded", domain.AuthorizationStatusFenced, domain.AuthorizationStatusSuperseded, false},
{"fenced authorization cannot be delivered again", domain.AuthorizationStatusFenced, domain.AuthorizationStatusDelivered, false},
{"consumed authorization cannot restart", domain.AuthorizationStatusConsumed, domain.AuthorizationStatusDelivered, false},
{"fenced authorization cannot be abandoned", domain.AuthorizationStatusFenced, domain.AuthorizationStatusAbandoned, false},
{"fenced authorization cannot be claimed again", domain.AuthorizationStatusFenced, domain.AuthorizationStatusClaimed, false},
{"consumed authorization cannot restart", domain.AuthorizationStatusConsumed, domain.AuthorizationStatusClaimed, false},
}
for _, test := range tests {
+74
View File
@@ -0,0 +1,74 @@
package domain
import (
"errors"
"time"
)
var ErrInvalidAttemptTransition = errors.New("invalid purchase attempt status transition")
// AttemptStatus 只描述单趟领取的可恢复执行。真实提交结果独立由唯一围栏记录调和。
type AttemptStatus string
const (
AttemptStatusClaimed AttemptStatus = "CLAIMED"
AttemptStatusOrdering AttemptStatus = "ORDERING"
AttemptStatusFailed AttemptStatus = "FAILED"
AttemptStatusFenced AttemptStatus = "FENCED"
AttemptStatusAbandoned AttemptStatus = "ABANDONED"
)
// AttemptFailureCode 是服务端可审计的固定失败摘要,不能承载页面正文或其他自由文本。
type AttemptFailureCode string
const (
AttemptFailureAuthorizationExpired AttemptFailureCode = "AUTHORIZATION_EXPIRED"
AttemptFailureLeaseLost AttemptFailureCode = "LEASE_LOST"
AttemptFailureGate1Rejected AttemptFailureCode = "GATE_1_REJECTED"
AttemptFailureQuantityMismatch AttemptFailureCode = "QUANTITY_MISMATCH"
AttemptFailureGate2Rejected AttemptFailureCode = "GATE_2_REJECTED"
AttemptFailureGate3Rejected AttemptFailureCode = "GATE_3_REJECTED"
AttemptFailureFenceRejected AttemptFailureCode = "FENCE_REJECTED"
AttemptFailureSafeAborted AttemptFailureCode = "SAFE_ABORTED"
)
type PurchaseAttempt struct {
ID string
TaskID string
AuthorizationID string
ClaimGeneration int
Status AttemptStatus
Gate1UnitPrice *string
Gate2UnitPrice *string
QuantityRead *int
ConfirmAmount *string
FailureCode *AttemptFailureCode
StartedAt time.Time
FinishedAt *time.Time
}
// CanTransitionTo 只允许围栏前的领取恢复为安全失败;围栏后不再提供回退或重试路径。
func (status AttemptStatus) CanTransitionTo(next AttemptStatus) bool {
_, allowed := attemptTransitions[status][next]
return allowed
}
func TransitionAttempt(current, next AttemptStatus) (AttemptStatus, error) {
if !current.CanTransitionTo(next) {
return current, ErrInvalidAttemptTransition
}
return next, nil
}
var attemptTransitions = map[AttemptStatus]map[AttemptStatus]struct{}{
AttemptStatusClaimed: {
AttemptStatusOrdering: {},
AttemptStatusFailed: {},
AttemptStatusAbandoned: {},
},
AttemptStatusOrdering: {
AttemptStatusFenced: {},
AttemptStatusFailed: {},
AttemptStatusAbandoned: {},
},
}
@@ -0,0 +1,34 @@
package domain_test
import (
"errors"
"testing"
"cmbuyer/admin/internal/domain"
)
func TestPurchaseAttemptTransitions(t *testing.T) {
for _, test := range []struct {
current domain.AttemptStatus
next domain.AttemptStatus
allowed bool
}{
{domain.AttemptStatusClaimed, domain.AttemptStatusOrdering, true},
{domain.AttemptStatusOrdering, domain.AttemptStatusFenced, true},
{domain.AttemptStatusOrdering, domain.AttemptStatusFailed, true},
{domain.AttemptStatusFenced, domain.AttemptStatusOrdering, false},
{domain.AttemptStatusFenced, domain.AttemptStatusAbandoned, false},
{domain.AttemptStatus("UNKNOWN"), domain.AttemptStatusOrdering, false},
} {
got, err := domain.TransitionAttempt(test.current, test.next)
if test.allowed {
if err != nil || got != test.next {
t.Fatalf("TransitionAttempt(%s, %s) = (%s, %v)", test.current, test.next, got, err)
}
continue
}
if !errors.Is(err, domain.ErrInvalidAttemptTransition) || got != test.current {
t.Fatalf("invalid TransitionAttempt(%s, %s) = (%s, %v)", test.current, test.next, got, err)
}
}
}
-16
View File
@@ -1,16 +0,0 @@
package domain
import "time"
type SpecTrial struct {
ID string
TaskID string
Attempt int
ProductTitle string
SelectedColor string
SelectedSize string
UnitPrice string
TotalPrice string
EvidenceSHA256 string
CreatedAt time.Time
}
+4 -4
View File
@@ -20,12 +20,12 @@ type OrderSubmission struct {
ID string
TaskID string
AuthorizationID string
CommandID string
DryRunID string
AttemptID string
Status SubmissionStatus
VerifiedUnitPrice string
Gate1UnitPrice string
Gate2UnitPrice string
QuantityRead int
ConfirmPageAmount string
ConfirmAmount string
CreatedAt time.Time
ResolvedAt *time.Time
}
+1
View File
@@ -20,6 +20,7 @@ func TestSubmissionTransitions(t *testing.T) {
{"cannot reopen fenced submission", domain.SubmissionStatusSubmitted, domain.SubmissionStatusFenced, false},
{"submitted cannot require reconciliation", domain.SubmissionStatusSubmitted, domain.SubmissionStatusReconciliationRequired, false},
{"cannot skip reconciliation", domain.SubmissionStatusFenced, domain.SubmissionStatusManualResolved, false},
{"manual resolution cannot create a second submission", domain.SubmissionStatusManualResolved, domain.SubmissionStatusFenced, false},
}
for _, test := range tests {
+13 -19
View File
@@ -14,15 +14,12 @@ const (
TaskStatusDraft TaskStatus = "DRAFT"
TaskStatusPending TaskStatus = "PENDING"
TaskStatusClaimed TaskStatus = "CLAIMED"
TaskStatusRunning TaskStatus = "RUNNING"
TaskStatusWaitingConfirmation TaskStatus = "WAITING_CONFIRMATION"
TaskStatusPendingRetrial TaskStatus = "PENDING_RETRIAL"
TaskStatusAuthorized TaskStatus = "AUTHORIZED"
TaskStatusOrdering TaskStatus = "ORDERING"
TaskStatusWaitingPayment TaskStatus = "WAITING_PAYMENT"
TaskStatusReconciliationRequired TaskStatus = "RECONCILIATION_REQUIRED"
TaskStatusNeedsManual TaskStatus = "NEEDS_MANUAL"
TaskStatusSucceeded TaskStatus = "SUCCEEDED"
TaskStatusFailed TaskStatus = "FAILED"
TaskStatusCanceled TaskStatus = "CANCELED"
)
@@ -69,34 +66,31 @@ func TransitionTask(current, next TaskStatus) (TaskStatus, error) {
var taskTransitions = map[TaskStatus]map[TaskStatus]struct{}{
TaskStatusDraft: {
TaskStatusPending: {},
TaskStatusCanceled: {},
},
TaskStatusPending: {
TaskStatusClaimed: {},
},
TaskStatusPendingRetrial: {
TaskStatusClaimed: {},
TaskStatusDraft: {},
TaskStatusCanceled: {},
},
TaskStatusClaimed: {
TaskStatusRunning: {},
TaskStatusPending: {},
},
TaskStatusRunning: {
TaskStatusWaitingConfirmation: {},
TaskStatusNeedsManual: {},
},
TaskStatusWaitingConfirmation: {
TaskStatusCanceled: {},
TaskStatusAuthorized: {},
},
TaskStatusAuthorized: {
TaskStatusOrdering: {},
TaskStatusDraft: {},
},
TaskStatusOrdering: {
TaskStatusNeedsManual: {},
TaskStatusWaitingPayment: {},
TaskStatusReconciliationRequired: {},
},
TaskStatusNeedsManual: {
TaskStatusDraft: {},
TaskStatusCanceled: {},
},
TaskStatusWaitingPayment: {
TaskStatusSucceeded: {},
},
TaskStatusReconciliationRequired: {
TaskStatusWaitingPayment: {},
TaskStatusFailed: {},
},
}
+14 -12
View File
@@ -14,22 +14,24 @@ func TestTaskTransitions(t *testing.T) {
next domain.TaskStatus
allowed bool
}{
{"start trial", domain.TaskStatusDraft, domain.TaskStatusPending, true},
{"claim trial", domain.TaskStatusPending, domain.TaskStatusClaimed, true},
{"claim retrial", domain.TaskStatusPendingRetrial, domain.TaskStatusClaimed, true},
{"start trial execution", domain.TaskStatusClaimed, domain.TaskStatusRunning, true},
{"release unstarted claim", domain.TaskStatusClaimed, domain.TaskStatusPending, true},
{"trial completes", domain.TaskStatusRunning, domain.TaskStatusWaitingConfirmation, true},
{"trial needs manual review", domain.TaskStatusRunning, domain.TaskStatusNeedsManual, true},
{"authorize confirmed trial", domain.TaskStatusWaitingConfirmation, domain.TaskStatusAuthorized, true},
{"reject confirmed trial", domain.TaskStatusWaitingConfirmation, domain.TaskStatusCanceled, true},
{"start authorized order leg", domain.TaskStatusAuthorized, domain.TaskStatusOrdering, true},
{"start purchase", domain.TaskStatusDraft, domain.TaskStatusPending, true},
{"cancel draft before fence", domain.TaskStatusDraft, domain.TaskStatusCanceled, true},
{"claim purchase", domain.TaskStatusPending, domain.TaskStatusClaimed, true},
{"release expired authorization", domain.TaskStatusPending, domain.TaskStatusDraft, true},
{"start ordering", domain.TaskStatusClaimed, domain.TaskStatusOrdering, true},
{"release unstarted claim", domain.TaskStatusClaimed, domain.TaskStatusDraft, true},
{"ordering needs manual review", domain.TaskStatusOrdering, domain.TaskStatusNeedsManual, true},
{"order reaches payment", domain.TaskStatusOrdering, domain.TaskStatusWaitingPayment, true},
{"order needs manual review before fence", domain.TaskStatusOrdering, domain.TaskStatusNeedsManual, true},
{"order needs reconciliation", domain.TaskStatusOrdering, domain.TaskStatusReconciliationRequired, true},
{"manual review resets draft", domain.TaskStatusNeedsManual, domain.TaskStatusDraft, true},
{"manual review cancels before fence", domain.TaskStatusNeedsManual, domain.TaskStatusCanceled, true},
{"payment verified", domain.TaskStatusWaitingPayment, domain.TaskStatusSucceeded, true},
{"cannot skip trial", domain.TaskStatusDraft, domain.TaskStatusAuthorized, false},
{"trial cannot enter order leg", domain.TaskStatusRunning, domain.TaskStatusOrdering, false},
{"reconcile confirms waiting payment", domain.TaskStatusReconciliationRequired, domain.TaskStatusWaitingPayment, true},
{"reconcile confirms failed", domain.TaskStatusReconciliationRequired, domain.TaskStatusFailed, true},
{"cannot skip authorization", domain.TaskStatusDraft, domain.TaskStatusOrdering, false},
{"ordering cannot return pending", domain.TaskStatusOrdering, domain.TaskStatusPending, false},
{"ordering cannot bypass manual review to draft", domain.TaskStatusOrdering, domain.TaskStatusDraft, false},
{"terminal task cannot restart", domain.TaskStatusSucceeded, domain.TaskStatusPending, false},
{"unknown status is rejected", domain.TaskStatus("UNKNOWN"), domain.TaskStatusPending, false},
}
+72
View File
@@ -0,0 +1,72 @@
// Package evidence defines the narrow internal screenshot contract shared by HTTP and storage.
package evidence
import (
"context"
"errors"
"io"
"time"
"cmbuyer/admin/internal/deviceauth"
)
const (
KindSKUPanelGate1 = "SKU_PANEL_GATE_1"
PrivacyInternalRaw = "INTERNAL_RAW"
PNGContentType = "image/png"
MaxFileBytes int64 = 10 << 20
MaxImageSide = 8192
MaxImagePixels = 16_777_216
)
var (
ErrInvalid = errors.New("invalid evidence")
ErrConflict = errors.New("evidence upload key conflict")
ErrNotFound = errors.New("evidence not found")
ErrTooLarge = errors.New("evidence file too large")
)
type UploadMetadata struct {
UploadKey string
TaskID string
AttemptID string
Kind string
PrivacyTier string
SHA256 string
CapturedAt time.Time
}
// StagedFile contains only server-generated state. Multipart filenames and client paths never enter this type.
type StagedFile struct {
Path string
SHA256 string
ByteSize int64
ContentType string
Width int
Height int
}
type Asset struct {
ID string `json:"asset_id"`
TaskID string `json:"task_id"`
AttemptID string `json:"attempt_id"`
Kind string `json:"kind"`
PrivacyTier string `json:"privacy_tier"`
SHA256 string `json:"sha256"`
ByteSize int64 `json:"byte_size"`
ContentType string `json:"content_type"`
Width int `json:"width_px"`
Height int `json:"height_px"`
CapturedAt time.Time `json:"captured_at"`
UploadedByDeviceID string `json:"-"`
StorageKey string `json:"-"`
CreatedAt time.Time `json:"-"`
}
// Store separates bounded multipart staging from metadata commit so field order cannot weaken validation.
type Store interface {
Stage(io.Reader, string) (StagedFile, error)
Discard(StagedFile)
Commit(context.Context, deviceauth.Principal, UploadMetadata, StagedFile) (Asset, bool, error)
Open(context.Context, string) (Asset, io.ReadSeekCloser, error)
}
+569 -159
View File
@@ -3,8 +3,11 @@ package migrations_test
import (
"context"
"database/sql"
"os"
"path/filepath"
"runtime"
"strconv"
"strings"
"testing"
"cmbuyer/admin/internal/migrations"
@@ -13,6 +16,8 @@ import (
"github.com/pressly/goose/v3"
)
const migrationTime = "2026-08-04T00:00:00Z"
func TestUpDownAndIdempotence(t *testing.T) {
database := openTestDatabase(t)
directory := migrationDirectory(t)
@@ -21,159 +26,606 @@ func TestUpDownAndIdempotence(t *testing.T) {
if err := migrations.Up(context, database, directory); err != nil {
t.Fatalf("apply migrations: %v", err)
}
assertVersion(t, database, 1)
assertVersion(t, database, 5)
assertTableExists(t, database, "tasks", true)
assertTableExists(t, database, "spec_trials", true)
assertTableExists(t, database, "spec_trials", false)
assertTableExists(t, database, "order_authorizations", true)
assertTableExists(t, database, "purchase_attempts", true)
assertTableExists(t, database, "order_submissions", true)
assertTableExists(t, database, "evidence_assets", true)
assertTableExists(t, database, "device_credentials", true)
assertTableExists(t, database, "purchase_attempt_claims", true)
assertTableExists(t, database, "single_pass_upgrade_guard", false)
if err := migrations.Up(context, database, directory); err != nil {
t.Fatalf("reapply migrations: %v", err)
}
assertVersion(t, database, 1)
assertVersion(t, database, 5)
if err := migrations.Down(context, database, directory); err != nil {
t.Fatalf("roll back migration: %v", err)
t.Fatalf("roll back task claim migration: %v", err)
}
assertVersion(t, database, 0)
assertTableExists(t, database, "tasks", false)
assertTableExists(t, database, "spec_trials", false)
assertTableExists(t, database, "order_authorizations", false)
assertTableExists(t, database, "order_submissions", false)
assertVersion(t, database, 4)
assertTableExists(t, database, "purchase_attempt_claims", false)
assertTableExists(t, database, "device_credentials", true)
if err := migrations.Up(context, database, directory); err != nil {
t.Fatalf("apply migration after rollback: %v", err)
if err := migrations.Down(context, database, directory); err != nil {
t.Fatalf("roll back device credential migration: %v", err)
}
assertVersion(t, database, 3)
assertTableExists(t, database, "device_credentials", false)
assertTableExists(t, database, "evidence_assets", true)
if err := migrations.Down(context, database, directory); err != nil {
t.Fatalf("roll back evidence migration: %v", err)
}
assertVersion(t, database, 2)
assertTableExists(t, database, "evidence_assets", false)
if err := migrations.Down(context, database, directory); err != nil {
t.Fatalf("roll back v2 migration: %v", err)
}
assertVersion(t, database, 1)
assertTableExists(t, database, "spec_trials", true)
assertTableExists(t, database, "purchase_attempts", false)
assertTableExists(t, database, "single_pass_downgrade_guard", false)
if err := migrations.Up(context, database, directory); err != nil {
t.Fatalf("reapply v2 after rollback: %v", err)
}
assertVersion(t, database, 5)
}
func TestSchemaConstraints(t *testing.T) {
func TestUpgradePreservesManualDraftLosslessly(t *testing.T) {
database := openTestDatabase(t)
migrateToV1(t, database)
if _, err := database.Exec(`
INSERT INTO tasks (
id, source, source_ref, title, goods_id, sku_color, sku_size, quantity, max_total_price,
reference_asset_id, status, version, created_at, updated_at
) VALUES ('draft-one', 'MANUAL', 'source-ref', 'title', 'goods', 'white', 'XL', 2, '80.50',
'asset-id', 'DRAFT', 7, '2026-08-03T00:00:00Z', '2026-08-03T01:00:00Z')
`); err != nil {
t.Fatalf("insert v1 draft: %v", err)
}
if err := migrations.Up(context.Background(), database, migrationDirectory(t)); err != nil {
t.Fatalf("upgrade v1 draft: %v", err)
}
assertVersion(t, database, 5)
var got struct {
id, source, sourceRef, title, goodsID, color, size, maxPrice, assetID, status, created, updated string
quantity, version int
}
if err := database.QueryRow(`SELECT id, source, source_ref, title, goods_id, sku_color, sku_size, quantity, max_total_price, reference_asset_id, status, version, created_at, updated_at FROM tasks WHERE id = 'draft-one'`).Scan(
&got.id, &got.source, &got.sourceRef, &got.title, &got.goodsID, &got.color, &got.size, &got.quantity, &got.maxPrice, &got.assetID, &got.status, &got.version, &got.created, &got.updated,
); err != nil {
t.Fatalf("read upgraded draft: %v", err)
}
if got != (struct {
id, source, sourceRef, title, goodsID, color, size, maxPrice, assetID, status, created, updated string
quantity, version int
}{"draft-one", "MANUAL", "source-ref", "title", "goods", "white", "XL", "80.50", "asset-id", "DRAFT", "2026-08-03T00:00:00Z", "2026-08-03T01:00:00Z", 2, 7}) {
t.Fatalf("upgraded draft changed: %#v", got)
}
}
func TestUpgradeRejectsLegacyExecutionDataAtomically(t *testing.T) {
tests := []struct {
name string
setup func(*testing.T, *sql.DB)
}{
{"non-draft task", func(t *testing.T, database *sql.DB) {
insertV1Task(t, database, "pending", "MANUAL", "PENDING", "1.00")
}},
{"non-manual task", func(t *testing.T, database *sql.DB) { insertV1Task(t, database, "excel", "EXCEL", "DRAFT", "1.00") }},
{"invalid v2 money", func(t *testing.T, database *sql.DB) { insertV1Task(t, database, "zero", "MANUAL", "DRAFT", "0.00") }},
{"third decimal place", func(t *testing.T, database *sql.DB) {
insertV1Task(t, database, "third-decimal", "MANUAL", "DRAFT", "1.234")
}},
{"spec trial", func(t *testing.T, database *sql.DB) {
insertV1Task(t, database, "task", "MANUAL", "DRAFT", "1.00")
insertV1SpecTrial(t, database, "trial", "task")
}},
{"authorization", func(t *testing.T, database *sql.DB) {
insertV1Task(t, database, "task", "MANUAL", "DRAFT", "1.00")
insertV1SpecTrial(t, database, "trial", "task")
insertV1Authorization(t, database, "auth", "task", "trial")
}},
{"submission", func(t *testing.T, database *sql.DB) {
insertV1Task(t, database, "task", "MANUAL", "DRAFT", "1.00")
insertV1SpecTrial(t, database, "trial", "task")
insertV1Authorization(t, database, "auth", "task", "trial")
if _, err := database.Exec(`INSERT INTO order_submissions (id, task_id, authorization_id, command_id, dry_run_id, status, verified_unit_price, quantity_read, confirm_page_amount, created_at) VALUES ('submission', 'task', 'auth', 'command', 'dry-run', 'FENCED', '1.00', 1, '1.00', ? )`, migrationTime); err != nil {
t.Fatalf("insert v1 submission: %v", err)
}
}},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
database := openTestDatabase(t)
migrateToV1(t, database)
test.setup(t, database)
before := v1RowCount(t, database)
if err := migrations.Up(context.Background(), database, migrationDirectory(t)); err == nil {
t.Fatal("unsafe legacy data upgraded successfully")
}
assertVersion(t, database, 1)
assertTableExists(t, database, "spec_trials", true)
assertTableExists(t, database, "purchase_attempts", false)
assertTableExists(t, database, "single_pass_upgrade_guard", false)
if after := v1RowCount(t, database); after != before {
t.Fatalf("v1 data changed after rejection: before=%d after=%d", before, after)
}
})
}
}
func TestV2SchemaConstraintsAndRelationships(t *testing.T) {
database := openTestDatabase(t)
if err := migrations.Up(context.Background(), database, migrationDirectory(t)); err != nil {
t.Fatalf("apply migrations: %v", err)
}
for _, column := range []struct {
table string
name string
}{
for _, column := range []struct{ table, name string }{
{"tasks", "max_total_price"},
{"spec_trials", "unit_price"},
{"spec_trials", "total_price"},
{"order_authorizations", "authorized_unit_price"},
{"order_authorizations", "total_price_cap"},
{"order_submissions", "verified_unit_price"},
{"order_submissions", "confirm_page_amount"},
{"purchase_attempts", "gate1_unit_price"},
{"purchase_attempts", "gate2_unit_price"},
{"purchase_attempts", "confirm_amount"},
{"order_submissions", "gate1_unit_price"},
{"order_submissions", "gate2_unit_price"},
{"order_submissions", "confirm_amount"},
} {
assertColumnType(t, database, column.table, column.name, "TEXT")
}
if _, err := database.Exec(`
INSERT INTO tasks (
id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price,
status, created_at, updated_at
) VALUES ('bad-quantity', 'MANUAL', 'title', 'goods', 'white', 'XL', 0, '80.00', 'DRAFT', '2026-08-03T00:00:00Z', '2026-08-03T00:00:00Z')
`); err == nil {
t.Fatal("insert task with quantity 0 succeeded")
for _, legacy := range []string{"spec_trials", "authorized_unit_price", "spec_trial_id", "command_id", "dry_run_id"} {
var count int
if err := database.QueryRow(`SELECT COUNT(*) FROM sqlite_master WHERE sql LIKE '%' || ? || '%'`, legacy).Scan(&count); err != nil {
t.Fatalf("search schema for %s: %v", legacy, err)
}
if count != 0 {
t.Fatalf("legacy identifier %q remains in v2 schema", legacy)
}
}
if _, err := database.Exec(`
INSERT INTO tasks (
id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price,
status, created_at, updated_at
) VALUES ('bad-price', 'MANUAL', 'title', 'goods', 'white', 'XL', 1, '80..00', 'DRAFT', '2026-08-03T00:00:00Z', '2026-08-03T00:00:00Z')
`); err == nil {
t.Fatal("insert task with malformed decimal price succeeded")
insertV2Task(t, database, "task-one", "MANUAL", "DRAFT")
insertV2Task(t, database, "task-two", "MANUAL", "DRAFT")
for index, value := range []string{"", "0", "0.00", "-1.00", "1e2", "1.", "1.234", " 1.00", "one"} {
if _, err := database.Exec(`INSERT INTO tasks (id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price, status, created_at, updated_at) VALUES (?, 'MANUAL', 'title', 'goods', 'white', 'XL', 1, ?, 'DRAFT', ?, ?)`, "bad-price-"+strconv.Itoa(index), value, migrationTime, migrationTime); err == nil {
t.Fatalf("invalid total price %q succeeded", value)
}
}
if _, err := database.Exec(`INSERT INTO tasks (id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price, status, created_at, updated_at) VALUES ('bad-status', 'MANUAL', 'title', 'goods', 'white', 'XL', 1, '1.00', 'UNKNOWN', ?, ?)`, migrationTime, migrationTime); err == nil {
t.Fatal("unknown task status succeeded")
}
insertV2Authorization(t, database, "auth-one", "task-one", 1, "start-one")
insertV2Authorization(t, database, "auth-two", "task-two", 1, "start-two")
if _, err := database.Exec(`INSERT INTO order_authorizations (id, task_id, task_version, start_key, goods_id, sku_color, sku_size, quantity, total_price_cap, status, created_by, created_at, expires_at) VALUES ('bad-auth-price', 'task-one', 2, 'bad-price', 'goods', 'white', 'XL', 1, '1.234', 'ACTIVE', 'admin', ?, ?)`, migrationTime, migrationTime); err == nil {
t.Fatal("third decimal authorization cap succeeded")
}
if _, err := database.Exec(`INSERT INTO order_authorizations (id, task_id, task_version, start_key, goods_id, sku_color, sku_size, quantity, total_price_cap, status, created_by, created_at, expires_at) VALUES ('bad-auth-status', 'task-one', 2, 'bad-status', 'goods', 'white', 'XL', 1, '1.00', 'UNKNOWN', 'admin', ?, ?)`, migrationTime, migrationTime); err == nil {
t.Fatal("unknown authorization status succeeded")
}
insertV2Authorization(t, database, "auth-one-b", "task-one", 2, "start-one-b")
if _, err := database.Exec(`INSERT INTO order_authorizations (id, task_id, task_version, start_key, goods_id, sku_color, sku_size, quantity, total_price_cap, status, created_by, created_at, expires_at) VALUES ('duplicate-version', 'task-one', 1, 'different-start', 'goods', 'white', 'XL', 1, '1.00', 'ACTIVE', 'admin', ?, ?)`, migrationTime, migrationTime); err == nil {
t.Fatal("duplicate task version authorization succeeded")
}
if _, err := database.Exec(`INSERT INTO purchase_attempts (id, task_id, authorization_id, claim_generation, status, started_at) VALUES ('cross-attempt', 'task-one', 'auth-two', 1, 'CLAIMED', ?)`, migrationTime); err == nil {
t.Fatal("attempt using another task authorization succeeded")
}
insertV2Attempt(t, database, "attempt-one", "task-one", "auth-one", 1)
if _, err := database.Exec(`INSERT INTO purchase_attempts (id, task_id, authorization_id, claim_generation, status, gate1_unit_price, started_at) VALUES ('bad-attempt-price', 'task-one', 'auth-one', 2, 'ORDERING', '1.234', ?)`, migrationTime); err == nil {
t.Fatal("third decimal gate price succeeded")
}
if _, err := database.Exec(`INSERT INTO purchase_attempts (id, task_id, authorization_id, claim_generation, status, started_at) VALUES ('bad-attempt-status', 'task-one', 'auth-one', 2, 'UNKNOWN', ?)`, migrationTime); err == nil {
t.Fatal("unknown attempt status succeeded")
}
if _, err := database.Exec(`INSERT INTO purchase_attempts (id, task_id, authorization_id, claim_generation, status, failure_code, started_at) VALUES ('bad-code', 'task-one', 'auth-one', 2, 'FAILED', 'FREE_TEXT', ?)`, migrationTime); err == nil {
t.Fatal("unknown failure code succeeded")
}
if _, err := database.Exec(`INSERT INTO order_submissions (id, task_id, authorization_id, attempt_id, status, gate1_unit_price, gate2_unit_price, quantity_read, confirm_amount, created_at) VALUES ('cross-submission', 'task-one', 'auth-two', 'attempt-one', 'FENCED', '1.00', '1.00', 1, '1.00', ?)`, migrationTime); err == nil {
t.Fatal("submission using another task authorization succeeded")
}
if _, err := database.Exec(`INSERT INTO order_submissions (id, task_id, authorization_id, attempt_id, status, gate1_unit_price, gate2_unit_price, quantity_read, confirm_amount, created_at) VALUES ('cross-authorization-submission', 'task-one', 'auth-one-b', 'attempt-one', 'FENCED', '1.00', '1.00', 1, '1.00', ?)`, migrationTime); err == nil {
t.Fatal("submission combining another same-task authorization and attempt succeeded")
}
if _, err := database.Exec(`INSERT INTO order_submissions (id, task_id, authorization_id, attempt_id, status, gate1_unit_price, gate2_unit_price, quantity_read, confirm_amount, created_at) VALUES ('bad-submission-status', 'task-one', 'auth-one', 'attempt-one', 'UNKNOWN', '1.00', '1.00', 1, '1.00', ?)`, migrationTime); err == nil {
t.Fatal("unknown submission status succeeded")
}
if _, err := database.Exec(`INSERT INTO order_submissions (id, task_id, authorization_id, attempt_id, status, gate1_unit_price, gate2_unit_price, quantity_read, confirm_amount, created_at) VALUES ('bad-submission-price', 'task-one', 'auth-one', 'attempt-one', 'FENCED', '1.234', '1.00', 1, '1.00', ?)`, migrationTime); err == nil {
t.Fatal("third decimal submission price succeeded")
}
insertV2Submission(t, database, "submission-one", "task-one", "auth-one", "attempt-one")
if _, err := database.Exec(`INSERT INTO order_submissions (id, task_id, authorization_id, attempt_id, status, gate1_unit_price, gate2_unit_price, quantity_read, confirm_amount, created_at) VALUES ('duplicate-auth', 'task-one', 'auth-one', 'attempt-one', 'FENCED', '1.00', '1.00', 1, '1.00', ?)`, migrationTime); err == nil {
t.Fatal("second submission for fenced authorization succeeded")
}
}
if _, err := database.Exec(`
INSERT INTO tasks (
id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price,
status, created_at, updated_at
) VALUES ('fractional-quantity', 'MANUAL', 'title', 'goods', 'white', 'XL', 1.5, '80.00', 'DRAFT', '2026-08-03T00:00:00Z', '2026-08-03T00:00:00Z')
`); err == nil {
t.Fatal("insert task with fractional quantity succeeded")
func TestEvidenceSchemaConstraintsAndDowngradeGuard(t *testing.T) {
database := openTestDatabase(t)
migrateToV3(t, database)
insertV2Task(t, database, "task-one", "MANUAL", "DRAFT")
insertV2Authorization(t, database, "auth-one", "task-one", 1, "start-one")
insertV2Attempt(t, database, "attempt-one", "task-one", "auth-one", 1)
insertV2Task(t, database, "task-two", "MANUAL", "DRAFT")
insertV2Authorization(t, database, "auth-two", "task-two", 1, "start-two")
insertV2Attempt(t, database, "attempt-two", "task-two", "auth-two", 1)
hash := strings.Repeat("a", 64)
insert := `INSERT INTO evidence_assets (id, upload_key, task_id, attempt_id, kind, privacy_tier, sha256, byte_size, content_type, width_px, height_px, storage_key, uploaded_by_device_id, captured_at, created_at) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`
validArgs := []any{"asset-one", "upload-one", "task-one", "attempt-one", "SKU_PANEL_GATE_1", "INTERNAL_RAW", hash, 100, "image/png", 100, 100, "aa/" + hash + ".png", "device-one", migrationTime, migrationTime}
if _, err := database.Exec(insert, validArgs...); err != nil {
t.Fatalf("insert valid evidence: %v", err)
}
for name, mutate := range map[string]func([]any){
"attempt from another task": func(values []any) { values[0], values[1], values[3] = "bad-task", "upload-bad-task", "attempt-two" },
"unapproved kind": func(values []any) { values[0], values[1], values[4] = "bad-kind", "upload-bad-kind", "ORDER_CONFIRM" },
"wrong privacy": func(values []any) { values[0], values[1], values[5] = "bad-privacy", "upload-bad-privacy", "PUBLIC" },
"uppercase hash": func(values []any) {
values[0], values[1], values[6], values[11] = "bad-hash", "upload-bad-hash", strings.Repeat("A", 64), "AA/"+strings.Repeat("A", 64)+".png"
},
"too many pixels": func(values []any) {
values[0], values[1], values[9], values[10] = "bad-pixels", "upload-bad-pixels", 8192, 8192
},
"client path": func(values []any) { values[0], values[1], values[11] = "bad-path", "upload-bad-path", `..\secret.png` },
} {
t.Run(name, func(t *testing.T) {
values := append([]any(nil), validArgs...)
mutate(values)
if _, err := database.Exec(insert, values...); err == nil {
t.Fatal("invalid evidence row succeeded")
}
})
}
if _, err := database.Exec(`
INSERT INTO tasks (
id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price,
status, created_at, updated_at
) VALUES ('trailing-decimal', 'MANUAL', 'title', 'goods', 'white', 'XL', 1, '80.', 'DRAFT', '2026-08-03T00:00:00Z', '2026-08-03T00:00:00Z')
`); err == nil {
t.Fatal("insert task with trailing decimal point succeeded")
if err := migrations.Down(context.Background(), database, migrationDirectory(t)); err == nil {
t.Fatal("evidence-bearing schema downgraded successfully")
}
assertVersion(t, database, 3)
assertTableExists(t, database, "evidence_assets", true)
var count int
if err := database.QueryRow("SELECT COUNT(*) FROM evidence_assets").Scan(&count); err != nil || count != 1 {
t.Fatalf("evidence after rejected downgrade = %d, err=%v", count, err)
}
}
if _, err := database.Exec(`
INSERT INTO tasks (
id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price,
status, created_at, updated_at
) VALUES ('bad-status', 'MANUAL', 'title', 'goods', 'white', 'XL', 1, '80.00', 'UNKNOWN', '2026-08-03T00:00:00Z', '2026-08-03T00:00:00Z')
`); err == nil {
t.Fatal("insert task with invalid status succeeded")
func TestDeviceCredentialSchemaConstraintsAndDowngradeGuard(t *testing.T) {
database := openTestDatabase(t)
migrateToV4(t, database)
deviceID := "13c9f507-7473-4fa6-8d71-8786c34c6301"
hash := make([]byte, 32)
for index := range hash {
hash[index] = byte(index + 1)
}
insert := `INSERT INTO device_credentials (device_id, display_name, token_sha256, status, created_at, revoked_at) VALUES (?, ?, ?, ?, ?, ?)`
valid := []any{deviceID, "采购工具一号", hash, "ACTIVE", migrationTime, nil}
if _, err := database.Exec(insert, valid...); err != nil {
t.Fatalf("insert valid credential: %v", err)
}
for name, mutate := range map[string]func([]any){
"uppercase uuid": func(values []any) { values[0], values[2] = strings.ToUpper(deviceID), append([]byte(nil), hash...) },
"wrong uuid version": func(values []any) {
values[0], values[2] = "23c9f507-7473-3fa6-8d71-8786c34c6301", append([]byte(nil), hash...)
},
"blank display name": func(values []any) {
values[0], values[1], values[2] = "33c9f507-7473-4fa6-8d71-8786c34c6301", "", append([]byte(nil), hash...)
},
"padded display name": func(values []any) {
values[0], values[1], values[2] = "43c9f507-7473-4fa6-8d71-8786c34c6301", " padded", append([]byte(nil), hash...)
},
"text hash": func(values []any) {
values[0], values[2] = "53c9f507-7473-4fa6-8d71-8786c34c6301", strings.Repeat("a", 32)
},
"short blob hash": func(values []any) { values[0], values[2] = "63c9f507-7473-4fa6-8d71-8786c34c6301", make([]byte, 31) },
"unknown status": func(values []any) {
values[0], values[2], values[3] = "73c9f507-7473-4fa6-8d71-8786c34c6301", append([]byte(nil), hash...), "UNKNOWN"
},
"active with revoke time": func(values []any) {
values[0], values[2], values[5] = "83c9f507-7473-4fa6-8d71-8786c34c6301", append([]byte(nil), hash...), migrationTime
},
"revoked without time": func(values []any) {
values[0], values[2], values[3] = "93c9f507-7473-4fa6-8d71-8786c34c6301", append([]byte(nil), hash...), "REVOKED"
},
"revoke before creation": func(values []any) {
values[0], values[2], values[3], values[4], values[5] = "b3c9f507-7473-4fa6-8d71-8786c34c6301", append([]byte(nil), hash...), "REVOKED", "2026-08-04T01:00:00Z", "2026-08-04T00:00:00Z"
},
"non UTC created time": func(values []any) {
values[0], values[2], values[4] = "a3c9f507-7473-4fa6-8d71-8786c34c6301", append([]byte(nil), hash...), "2026-08-04T08:00:00+08:00"
},
} {
t.Run(name, func(t *testing.T) {
values := append([]any(nil), valid...)
mutate(values)
if bytesValue, ok := values[2].([]byte); ok && len(bytesValue) == 32 {
bytesValue[0]++
}
if _, err := database.Exec(insert, values...); err == nil {
t.Fatal("invalid device credential row succeeded")
}
})
}
insertTask(t, database, "task-one")
insertTask(t, database, "task-two")
if _, err := database.Exec(`
INSERT INTO spec_trials (
id, task_id, attempt, product_title, selected_color, selected_size, unit_price,
total_price, evidence_sha256, created_at
) VALUES ('orphan-trial', 'missing-task', 1, 'title', 'white', 'XL', '32.50', '65.00', 'hash', '2026-08-03T00:00:00Z')
`); err == nil {
t.Fatal("insert spec trial without task succeeded")
if err := migrations.Down(context.Background(), database, migrationDirectory(t)); err == nil {
t.Fatal("credential-bearing schema downgraded successfully")
}
assertVersion(t, database, 4)
assertTableExists(t, database, "device_credentials", true)
var count int
if err := database.QueryRow(`SELECT COUNT(*) FROM device_credentials`).Scan(&count); err != nil || count != 1 {
t.Fatalf("credentials after rejected downgrade = %d, err=%v", count, err)
}
}
insertSpecTrial(t, database, "trial-one", "task-one")
insertSpecTrial(t, database, "trial-two", "task-two")
if _, err := database.Exec(`
INSERT INTO order_authorizations (
id, task_id, spec_trial_id, version, goods_id, sku_color, sku_size, quantity,
authorized_unit_price, total_price_cap, status, created_by, created_at, expires_at
) VALUES ('authorization-cross-task', 'task-one', 'trial-two', 1, 'goods', 'white', 'XL', 2, '32.50', '80.00', 'PENDING_DELIVERY', 'admin-one', '2026-08-03T00:00:00Z', '2026-08-03T01:00:00Z')
`); err == nil {
t.Fatal("insert authorization with a spec trial from another task succeeded")
func TestTaskClaimMigrationGuardsOwnershipConstraintsAndDowngradeFacts(t *testing.T) {
t.Run("upgrade rejects unmappable execution facts atomically", func(t *testing.T) {
database := openTestDatabase(t)
migrateToV4(t, database)
insertV2Task(t, database, "legacy-task", "MANUAL", "DRAFT")
insertV2Authorization(t, database, "legacy-auth", "legacy-task", 1, "legacy-start")
insertV2Attempt(t, database, "legacy-attempt", "legacy-task", "legacy-auth", 1)
if err := migrations.Up(context.Background(), database, migrationDirectory(t)); err == nil {
t.Fatal("v5 upgrade accepted an attempt without device/session ownership")
}
assertVersion(t, database, 4)
assertTableExists(t, database, "purchase_attempt_claims", false)
var count int
if err := database.QueryRow("SELECT COUNT(*) FROM purchase_attempts").Scan(&count); err != nil || count != 1 {
t.Fatalf("legacy attempt after rejected upgrade = %d, err %v", count, err)
}
})
t.Run("schema binds authorization device session generation and token", func(t *testing.T) {
database := openTestDatabase(t)
if err := migrations.Up(context.Background(), database, migrationDirectory(t)); err != nil {
t.Fatalf("apply migrations: %v", err)
}
deviceA := "13c9f507-7473-4fa6-8d71-8786c34c6301"
deviceB := "23c9f507-7473-4fa6-8d71-8786c34c6301"
sessionA := "33c9f507-7473-4fa6-8d71-8786c34c6301"
sessionB := "43c9f507-7473-4fa6-8d71-8786c34c6301"
taskA := "53c9f507-7473-4fa6-8d71-8786c34c6301"
authA := "63c9f507-7473-4fa6-8d71-8786c34c6301"
attemptA := "73c9f507-7473-4fa6-8d71-8786c34c6301"
taskB := "83c9f507-7473-4fa6-8d71-8786c34c6301"
authB := "93c9f507-7473-4fa6-8d71-8786c34c6301"
attemptB := "a3c9f507-7473-4fa6-8d71-8786c34c6301"
tokenA := make([]byte, 32)
for index := range tokenA {
tokenA[index] = byte(index + 1)
}
for index, device := range []string{deviceA, deviceB} {
hash := make([]byte, 32)
hash[0] = byte(index + 100)
if _, err := database.Exec(`INSERT INTO device_credentials
(device_id,display_name,token_sha256,status,created_at,revoked_at)
VALUES (?, ?, ?, 'ACTIVE', ?, NULL)`, device, "device "+strconv.Itoa(index), hash, migrationTime); err != nil {
t.Fatalf("insert device: %v", err)
}
}
insertV2Task(t, database, taskA, "MANUAL", "DRAFT")
insertV2Authorization(t, database, authA, taskA, 1, "start-a")
insertV2Attempt(t, database, attemptA, taskA, authA, 1)
insertClaim := `INSERT INTO purchase_attempt_claims
(attempt_id,task_id,authorization_id,claimed_by_device_id,session_id,claim_generation,
task_version,task_title,authorization_task_version,goods_id,sku_color,sku_size,quantity,
total_price_cap,authorization_expires_at,claim_nonce,claim_token_sha256,lease_expires_at,claimed_at,closed_at)
VALUES (?, ?, ?, ?, ?, ?, 2, 'task', 1, 'goods', 'white', 'XL', 1, '1.00',
'2026-08-04T01:00:00Z', ?, ?, '2026-08-04T00:05:00Z', ?, NULL)`
if _, err := database.Exec(insertClaim, attemptA, taskA, authA, deviceA, sessionA, 1, make([]byte, 32), tokenA, migrationTime); err != nil {
t.Fatalf("insert valid claim: %v", err)
}
if _, err := database.Exec(`INSERT INTO purchase_attempts
(id,task_id,authorization_id,claim_generation,status,started_at)
VALUES ('b3c9f507-7473-4fa6-8d71-8786c34c6301', ?, ?, 2, 'CLAIMED', ?)`, taskA, authA, migrationTime); err == nil {
t.Fatal("second attempt for one authorization succeeded")
}
insertAuthorization(t, database, "authorization-one", "task-one", "trial-one", 1)
insertAuthorization(t, database, "authorization-task-two", "task-two", "trial-two", 1)
if _, err := database.Exec(`
INSERT INTO order_submissions (
id, task_id, authorization_id, command_id, dry_run_id, status, verified_unit_price,
quantity_read, confirm_page_amount, created_at
) VALUES ('submission-cross-task', 'task-one', 'authorization-task-two', 'command-cross-task', 'dry-run-cross-task', 'FENCED', '32.50', 2, '65.00', '2026-08-03T00:00:00Z')
`); err == nil {
t.Fatal("insert submission with an authorization from another task succeeded")
insertV2Task(t, database, taskB, "MANUAL", "DRAFT")
insertV2Authorization(t, database, authB, taskB, 1, "start-b")
insertV2Attempt(t, database, attemptB, taskB, authB, 1)
if _, err := database.Exec(insertClaim, attemptB, taskB, authB, deviceB, sessionB, 2, make([]byte, 32), make([]byte, 32), migrationTime); err == nil {
t.Fatal("claim with generation different from its attempt succeeded")
}
if _, err := database.Exec(insertClaim, attemptB, taskB, authB, deviceA, sessionB, 1, make([]byte, 32), make([]byte, 32), migrationTime); err == nil {
t.Fatal("second open claim for one device succeeded")
}
if _, err := database.Exec(`
INSERT INTO order_authorizations (
id, task_id, spec_trial_id, version, goods_id, sku_color, sku_size, quantity,
authorized_unit_price, total_price_cap, status, created_by, created_at, expires_at
) VALUES ('authorization-duplicate', 'task-one', 'trial-one', 1, 'goods', 'white', 'XL', 2, '32.50', '80.00', 'PENDING_DELIVERY', 'admin-one', '2026-08-03T00:00:00Z', '2026-08-03T01:00:00Z')
`); err == nil {
t.Fatal("insert authorization with duplicate task version succeeded")
claimRequest := `INSERT INTO task_claim_requests
(claim_request_id,device_id,session_id,outcome,attempt_id,response_lease_expires_at,error_code,created_at)
VALUES (?, ?, ?, 'CLAIMED', ?, '2026-08-04T00:05:00Z', NULL, ?)`
if _, err := database.Exec(claimRequest, "c3c9f507-7473-4fa6-8d71-8786c34c6301", deviceA, sessionB, attemptA, migrationTime); err == nil {
t.Fatal("claim request with another session succeeded")
}
if _, err := database.Exec(claimRequest, "d3c9f507-7473-4fa6-8d71-8786c34c6301", deviceA, sessionA, attemptA, migrationTime); err != nil {
t.Fatalf("insert bound claim request: %v", err)
}
renewal := `INSERT INTO purchase_attempt_lease_renewals
(renew_request_id,task_id,attempt_id,device_id,session_id,claim_generation,
claim_token_sha256,expected_lease_expires_at,lease_expires_at,created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, '2026-08-04T00:05:00Z', '2026-08-04T00:06:00Z', ?)`
if _, err := database.Exec(renewal, "e3c9f507-7473-4fa6-8d71-8786c34c6301", taskA, attemptA, deviceA, sessionA, 2, tokenA, migrationTime); err == nil {
t.Fatal("renewal with another generation succeeded")
}
wrongHash := append([]byte(nil), tokenA...)
wrongHash[0] ^= 0xff
if _, err := database.Exec(renewal, "f3c9f507-7473-4fa6-8d71-8786c34c6301", taskA, attemptA, deviceA, sessionA, 1, wrongHash, migrationTime); err == nil {
t.Fatal("renewal with another token hash succeeded")
}
if err := migrations.Down(context.Background(), database, migrationDirectory(t)); err == nil {
t.Fatal("claim-bearing schema downgraded successfully")
}
assertVersion(t, database, 5)
assertTableExists(t, database, "purchase_attempt_claims", true)
})
t.Run("empty request alone blocks downgrade", func(t *testing.T) {
database := openTestDatabase(t)
if err := migrations.Up(context.Background(), database, migrationDirectory(t)); err != nil {
t.Fatal(err)
}
device := "13c9f507-7473-4fa6-8d71-8786c34c6301"
if _, err := database.Exec(`INSERT INTO device_credentials
(device_id,display_name,token_sha256,status,created_at,revoked_at)
VALUES (?, 'device', ?, 'ACTIVE', ?, NULL)`, device, make([]byte, 32), migrationTime); err != nil {
t.Fatal(err)
}
if _, err := database.Exec(`INSERT INTO task_claim_requests
(claim_request_id,device_id,session_id,outcome,attempt_id,response_lease_expires_at,error_code,created_at)
VALUES ('23c9f507-7473-4fa6-8d71-8786c34c6301', ?,
'33c9f507-7473-4fa6-8d71-8786c34c6301', 'EMPTY', NULL, NULL, NULL, ?)`, device, migrationTime); err != nil {
t.Fatal(err)
}
if err := migrations.Down(context.Background(), database, migrationDirectory(t)); err == nil {
t.Fatal("EMPTY request was silently dropped by downgrade")
}
assertVersion(t, database, 5)
})
}
insertSubmission(t, database, "submission-one", "authorization-one", "command-one")
if _, err := database.Exec(`
INSERT INTO order_submissions (
id, task_id, authorization_id, command_id, dry_run_id, status, verified_unit_price,
quantity_read, confirm_page_amount, created_at
) VALUES ('submission-duplicate-auth', 'task-one', 'authorization-one', 'command-two', 'dry-run-two', 'FENCED', '32.50', 2, '65.00', '2026-08-03T00:00:00Z')
`); err == nil {
t.Fatal("insert submission with duplicate authorization succeeded")
func TestDowngradeRejectsV2BusinessDataAtomically(t *testing.T) {
tests := []struct {
name string
setup func(*testing.T, *sql.DB)
}{
{"authorization", func(t *testing.T, database *sql.DB) {
insertV2Task(t, database, "task", "MANUAL", "DRAFT")
insertV2Authorization(t, database, "auth", "task", 1, "start")
}},
{"attempt", func(t *testing.T, database *sql.DB) {
insertV2Task(t, database, "task", "MANUAL", "DRAFT")
insertV2Authorization(t, database, "auth", "task", 1, "start")
insertV2Attempt(t, database, "attempt", "task", "auth", 1)
}},
{"submission", func(t *testing.T, database *sql.DB) {
insertV2Task(t, database, "task", "MANUAL", "DRAFT")
insertV2Authorization(t, database, "auth", "task", 1, "start")
insertV2Attempt(t, database, "attempt", "task", "auth", 1)
insertV2Submission(t, database, "submission", "task", "auth", "attempt")
}},
{"non-draft task", func(t *testing.T, database *sql.DB) { insertV2Task(t, database, "pending", "MANUAL", "PENDING") }},
{"non-manual task", func(t *testing.T, database *sql.DB) { insertV2Task(t, database, "excel", "EXCEL", "DRAFT") }},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
database := openTestDatabase(t)
migrateToV2(t, database)
test.setup(t, database)
before := v2RowCount(t, database)
if err := migrations.Down(context.Background(), database, migrationDirectory(t)); err == nil {
t.Fatal("unsafe v2 data downgraded successfully")
}
assertVersion(t, database, 2)
assertTableExists(t, database, "purchase_attempts", true)
assertTableExists(t, database, "spec_trials", false)
assertTableExists(t, database, "single_pass_downgrade_guard", false)
if after := v2RowCount(t, database); after != before {
t.Fatalf("v2 data changed after rejected downgrade: before=%d after=%d", before, after)
}
})
}
}
insertAuthorization(t, database, "authorization-two", "task-one", "trial-one", 2)
if _, err := database.Exec(`
INSERT INTO order_submissions (
id, task_id, authorization_id, command_id, dry_run_id, status, verified_unit_price,
quantity_read, confirm_page_amount, created_at
) VALUES ('submission-duplicate-command', 'task-one', 'authorization-two', 'command-one', 'dry-run-three', 'FENCED', '32.50', 2, '65.00', '2026-08-03T00:00:00Z')
`); err == nil {
t.Fatal("insert submission with duplicate command succeeded")
func migrateToV1(t *testing.T, database *sql.DB) {
t.Helper()
if err := migrations.Run(context.Background(), database, migrationDirectory(t), "up-by-one"); err != nil {
t.Fatalf("apply v1: %v", err)
}
assertVersion(t, database, 1)
}
func migrateToV2(t *testing.T, database *sql.DB) {
t.Helper()
if err := migrations.Run(context.Background(), database, migrationDirectory(t), "up-by-one"); err != nil {
t.Fatalf("apply v1: %v", err)
}
if err := migrations.Run(context.Background(), database, migrationDirectory(t), "up-by-one"); err != nil {
t.Fatalf("apply v2: %v", err)
}
assertVersion(t, database, 2)
}
func migrateToV3(t *testing.T, database *sql.DB) {
t.Helper()
migrateToV2(t, database)
if err := migrations.Run(context.Background(), database, migrationDirectory(t), "up-by-one"); err != nil {
t.Fatalf("apply v3: %v", err)
}
assertVersion(t, database, 3)
}
func migrateToV4(t *testing.T, database *sql.DB) {
t.Helper()
migrateToV3(t, database)
if err := migrations.Run(context.Background(), database, migrationDirectory(t), "up-by-one"); err != nil {
t.Fatalf("apply v4: %v", err)
}
assertVersion(t, database, 4)
}
func insertV1Task(t *testing.T, database *sql.DB, id, source, status, price string) {
t.Helper()
if _, err := database.Exec(`INSERT INTO tasks (id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price, status, created_at, updated_at) VALUES (?, ?, 'title', 'goods', 'white', 'XL', 1, ?, ?, ?, ?)`, id, source, price, status, migrationTime, migrationTime); err != nil {
t.Fatalf("insert v1 task: %v", err)
}
}
func insertV1SpecTrial(t *testing.T, database *sql.DB, id, taskID string) {
t.Helper()
if _, err := database.Exec(`INSERT INTO spec_trials (id, task_id, attempt, product_title, selected_color, selected_size, unit_price, total_price, evidence_sha256, created_at) VALUES (?, ?, 1, 'title', 'white', 'XL', '1.00', '1.00', 'hash', ?)`, id, taskID, migrationTime); err != nil {
t.Fatalf("insert v1 spec trial: %v", err)
}
}
func insertV1Authorization(t *testing.T, database *sql.DB, id, taskID, trialID string) {
t.Helper()
if _, err := database.Exec(`INSERT INTO order_authorizations (id, task_id, spec_trial_id, version, goods_id, sku_color, sku_size, quantity, authorized_unit_price, total_price_cap, status, created_by, created_at, expires_at) VALUES (?, ?, ?, 1, 'goods', 'white', 'XL', 1, '1.00', '1.00', 'PENDING_DELIVERY', 'admin', ?, ?)`, id, taskID, trialID, migrationTime, migrationTime); err != nil {
t.Fatalf("insert v1 authorization: %v", err)
}
}
func insertV2Task(t *testing.T, database *sql.DB, id, source, status string) {
t.Helper()
if _, err := database.Exec(`INSERT INTO tasks (id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price, status, created_at, updated_at) VALUES (?, ?, 'title', 'goods', 'white', 'XL', 1, '1.00', ?, ?, ?)`, id, source, status, migrationTime, migrationTime); err != nil {
t.Fatalf("insert v2 task: %v", err)
}
}
func insertV2Authorization(t *testing.T, database *sql.DB, id, taskID string, version int, startKey string) {
t.Helper()
if _, err := database.Exec(`INSERT INTO order_authorizations (id, task_id, task_version, start_key, goods_id, sku_color, sku_size, quantity, total_price_cap, status, created_by, created_at, expires_at) VALUES (?, ?, ?, ?, 'goods', 'white', 'XL', 1, '1.00', 'ACTIVE', 'admin', ?, ?)`, id, taskID, version, startKey, migrationTime, migrationTime); err != nil {
t.Fatalf("insert v2 authorization: %v", err)
}
}
func insertV2Attempt(t *testing.T, database *sql.DB, id, taskID, authorizationID string, generation int) {
t.Helper()
if _, err := database.Exec(`INSERT INTO purchase_attempts (id, task_id, authorization_id, claim_generation, status, started_at) VALUES (?, ?, ?, ?, 'CLAIMED', ?)`, id, taskID, authorizationID, generation, migrationTime); err != nil {
t.Fatalf("insert v2 attempt: %v", err)
}
}
func insertV2Submission(t *testing.T, database *sql.DB, id, taskID, authorizationID, attemptID string) {
t.Helper()
if _, err := database.Exec(`INSERT INTO order_submissions (id, task_id, authorization_id, attempt_id, status, gate1_unit_price, gate2_unit_price, quantity_read, confirm_amount, created_at) VALUES (?, ?, ?, ?, 'FENCED', '1.00', '1.00', 1, '1.00', ?)`, id, taskID, authorizationID, attemptID, migrationTime); err != nil {
t.Fatalf("insert v2 submission: %v", err)
}
}
func v1RowCount(t *testing.T, database *sql.DB) int {
t.Helper()
var count int
if err := database.QueryRow(`SELECT (SELECT COUNT(*) FROM tasks) + (SELECT COUNT(*) FROM spec_trials) + (SELECT COUNT(*) FROM order_authorizations) + (SELECT COUNT(*) FROM order_submissions)`).Scan(&count); err != nil {
t.Fatalf("count v1 rows: %v", err)
}
return count
}
func v2RowCount(t *testing.T, database *sql.DB) int {
t.Helper()
var count int
if err := database.QueryRow(`SELECT (SELECT COUNT(*) FROM tasks) + (SELECT COUNT(*) FROM order_authorizations) + (SELECT COUNT(*) FROM purchase_attempts) + (SELECT COUNT(*) FROM order_submissions)`).Scan(&count); err != nil {
t.Fatalf("count v2 rows: %v", err)
}
return count
}
func openTestDatabase(t *testing.T) *sql.DB {
@@ -182,12 +634,7 @@ func openTestDatabase(t *testing.T) *sql.DB {
if err != nil {
t.Fatalf("open test database: %v", err)
}
t.Cleanup(func() {
if err := database.Close(); err != nil {
t.Errorf("close test database: %v", err)
}
})
t.Cleanup(func() { _ = database.Close() })
return database
}
@@ -197,7 +644,6 @@ func migrationDirectory(t *testing.T) string {
if !ok {
t.Fatal("locate migration test source")
}
return filepath.Join(filepath.Dir(file), "..", "..", "migrations")
}
@@ -234,50 +680,14 @@ func assertColumnType(t *testing.T, database *sql.DB, table, column, want string
}
}
func insertTask(t *testing.T, database *sql.DB, id string) {
t.Helper()
if _, err := database.Exec(`
INSERT INTO tasks (
id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price,
status, created_at, updated_at
) VALUES (?, 'MANUAL', 'title', 'goods', 'white', 'XL', 2, '80.00', 'DRAFT', '2026-08-03T00:00:00Z', '2026-08-03T00:00:00Z')
`, id); err != nil {
t.Fatalf("insert task: %v", err)
func TestMigrationsDoNotDisableForeignKeys(t *testing.T) {
for _, name := range []string{"00002_single_pass_model.sql", "00003_evidence_assets.sql", "00004_device_credentials.sql"} {
contents, err := os.ReadFile(filepath.Join(migrationDirectory(t), name))
if err != nil {
t.Fatalf("read %s: %v", name, err)
}
if strings.Contains(strings.ToUpper(string(contents)), "PRAGMA FOREIGN_KEYS = OFF") {
t.Fatalf("%s disables foreign keys", name)
}
}
func insertSpecTrial(t *testing.T, database *sql.DB, id, taskID string) {
t.Helper()
if _, err := database.Exec(`
INSERT INTO spec_trials (
id, task_id, attempt, product_title, selected_color, selected_size, unit_price,
total_price, evidence_sha256, created_at
) VALUES (?, ?, 1, 'title', 'white', 'XL', '32.50', '65.00', 'hash', '2026-08-03T00:00:00Z')
`, id, taskID); err != nil {
t.Fatalf("insert spec trial: %v", err)
}
}
func insertAuthorization(t *testing.T, database *sql.DB, id, taskID, specTrialID string, version int) {
t.Helper()
if _, err := database.Exec(`
INSERT INTO order_authorizations (
id, task_id, spec_trial_id, version, goods_id, sku_color, sku_size, quantity,
authorized_unit_price, total_price_cap, status, created_by, created_at, expires_at
) VALUES (?, ?, ?, ?, 'goods', 'white', 'XL', 2, '32.50', '80.00', 'PENDING_DELIVERY', 'admin-one', '2026-08-03T00:00:00Z', '2026-08-03T01:00:00Z')
`, id, taskID, specTrialID, version); err != nil {
t.Fatalf("insert authorization: %v", err)
}
}
func insertSubmission(t *testing.T, database *sql.DB, id, authorizationID, commandID string) {
t.Helper()
if _, err := database.Exec(`
INSERT INTO order_submissions (
id, task_id, authorization_id, command_id, dry_run_id, status, verified_unit_price,
quantity_read, confirm_page_amount, created_at
) VALUES (?, 'task-one', ?, ?, 'dry-run-one', 'FENCED', '32.50', 2, '65.00', '2026-08-03T00:00:00Z')
`, id, authorizationID, commandID); err != nil {
t.Fatalf("insert submission: %v", err)
}
}
+205
View File
@@ -0,0 +1,205 @@
package server
import (
"errors"
"io"
"mime"
"mime/multipart"
"net/http"
"strconv"
"strings"
"time"
"unicode/utf8"
"cmbuyer/admin/internal/deviceauth"
"cmbuyer/admin/internal/evidence"
"github.com/gin-gonic/gin"
)
const (
maxEvidenceRequestBytes = evidence.MaxFileBytes + 64<<10
maxEvidenceFieldBytes = 4 << 10
)
var evidenceFieldNames = map[string]struct{}{
"upload_key": {}, "attempt_id": {}, "kind": {}, "privacy_tier": {}, "sha256": {}, "captured_at": {},
}
func uploadEvidence(options Options) gin.HandlerFunc {
return func(context *gin.Context) {
// Authentication deliberately precedes content-type parsing and every body read. A rejected
// device must not make the service spool or inspect a potentially sensitive upload.
principal, err := options.DeviceAuthenticator.Authenticate(context.Request)
if errors.Is(err, deviceauth.ErrUnauthenticated) {
context.Header("WWW-Authenticate", "Bearer")
context.Status(http.StatusUnauthorized)
return
}
if err != nil {
context.Status(http.StatusServiceUnavailable)
return
}
if !deviceauth.ValidDeviceID(principal.ID) {
// A custom authenticator is still an untrusted boundary. Do not defer principal
// validation until Commit because multipart bytes would already have been read.
context.Status(http.StatusServiceUnavailable)
return
}
boundary, ok := multipartBoundary(context.GetHeader("Content-Type"))
if !ok {
context.Status(http.StatusUnsupportedMediaType)
return
}
context.Request.Body = http.MaxBytesReader(context.Writer, context.Request.Body, maxEvidenceRequestBytes)
reader := multipart.NewReader(context.Request.Body, boundary)
fields := make(map[string]string, len(evidenceFieldNames))
var staged evidence.StagedFile
hasFile := false
discard := func() {
if hasFile {
options.Evidence.Discard(staged)
}
}
for {
part, err := reader.NextPart()
if errors.Is(err, io.EOF) {
break
}
if err != nil {
discard()
writeMultipartError(context, err)
return
}
name := part.FormName()
if name == "file" {
if hasFile || part.FileName() == "" || !exactPNGContentType(part.Header.Get("Content-Type")) {
_ = part.Close()
discard()
context.Status(http.StatusUnsupportedMediaType)
return
}
staged, err = options.Evidence.Stage(part, evidence.PNGContentType)
_ = part.Close()
if err != nil {
writeEvidenceStoreError(context, err)
return
}
hasFile = true
continue
}
if _, allowed := evidenceFieldNames[name]; !allowed || part.FileName() != "" {
_ = part.Close()
discard()
context.Status(http.StatusBadRequest)
return
}
if _, duplicate := fields[name]; duplicate {
_ = part.Close()
discard()
context.Status(http.StatusBadRequest)
return
}
value, err := io.ReadAll(io.LimitReader(part, maxEvidenceFieldBytes+1))
_ = part.Close()
if err != nil || len(value) == 0 || len(value) > maxEvidenceFieldBytes || !utf8.Valid(value) {
discard()
context.Status(http.StatusBadRequest)
return
}
fields[name] = string(value)
}
if !hasFile || len(fields) != len(evidenceFieldNames) {
discard()
context.Status(http.StatusBadRequest)
return
}
captured, err := time.Parse(time.RFC3339Nano, fields["captured_at"])
if err != nil || !strings.HasSuffix(fields["captured_at"], "Z") {
discard()
context.Status(http.StatusBadRequest)
return
}
asset, replayed, err := options.Evidence.Commit(context.Request.Context(), principal, evidence.UploadMetadata{
UploadKey: fields["upload_key"], TaskID: context.Param("id"), AttemptID: fields["attempt_id"],
Kind: fields["kind"], PrivacyTier: fields["privacy_tier"], SHA256: fields["sha256"], CapturedAt: captured.UTC(),
}, staged)
if err != nil {
writeEvidenceStoreError(context, err)
return
}
status := http.StatusCreated
if replayed {
status = http.StatusOK
}
context.JSON(status, asset)
}
}
func readEvidence(options Options) gin.HandlerFunc {
return func(context *gin.Context) {
if !options.Sessions.IsAuthenticated(context.Request) {
context.Status(http.StatusUnauthorized)
return
}
asset, file, err := options.Evidence.Open(context.Request.Context(), context.Param("asset_id"))
if errors.Is(err, evidence.ErrNotFound) {
context.Status(http.StatusNotFound)
return
}
if err != nil {
context.Status(http.StatusInternalServerError)
return
}
defer file.Close()
context.Header("Content-Type", evidence.PNGContentType)
context.Header("Content-Length", strconv.FormatInt(asset.ByteSize, 10))
context.Header("Content-Disposition", `inline; filename="evidence.png"`)
context.Header("Cache-Control", "no-store")
context.Header("X-Content-Type-Options", "nosniff")
context.Status(http.StatusOK)
if _, err := io.Copy(context.Writer, file); err != nil {
_ = context.Error(err)
}
}
}
func multipartBoundary(value string) (string, bool) {
mediaType, parameters, err := mime.ParseMediaType(value)
if err != nil || mediaType != "multipart/form-data" || len(parameters) != 1 || parameters["boundary"] == "" {
return "", false
}
return parameters["boundary"], true
}
func exactPNGContentType(value string) bool {
mediaType, parameters, err := mime.ParseMediaType(value)
return err == nil && mediaType == evidence.PNGContentType && len(parameters) == 0
}
func writeMultipartError(context *gin.Context, err error) {
var tooLarge *http.MaxBytesError
if errors.As(err, &tooLarge) {
context.Status(http.StatusRequestEntityTooLarge)
return
}
context.Status(http.StatusBadRequest)
}
func writeEvidenceStoreError(context *gin.Context, err error) {
var tooLarge *http.MaxBytesError
switch {
case errors.As(err, &tooLarge):
context.Status(http.StatusRequestEntityTooLarge)
case errors.Is(err, evidence.ErrTooLarge):
context.Status(http.StatusRequestEntityTooLarge)
case errors.Is(err, evidence.ErrInvalid):
context.Status(http.StatusBadRequest)
case errors.Is(err, evidence.ErrConflict):
context.Status(http.StatusConflict)
default:
context.Status(http.StatusInternalServerError)
}
}
+478
View File
@@ -0,0 +1,478 @@
package server_test
import (
"bytes"
"context"
"crypto/sha256"
"database/sql"
"encoding/hex"
"encoding/json"
"image"
"image/png"
"io"
"mime/multipart"
"net/http"
"net/http/httptest"
"net/textproto"
"path/filepath"
"runtime"
"strings"
"testing"
"cmbuyer/admin/internal/deviceauth"
"cmbuyer/admin/internal/evidence"
"cmbuyer/admin/internal/migrations"
evidencestorage "cmbuyer/admin/internal/storage/evidence"
"cmbuyer/admin/internal/storage/sqlite"
)
const (
evidenceTaskID = "63c9f507-7473-4fa6-8d71-8786c34c6301"
evidenceAuthID = "73c9f507-7473-4fa6-8d71-8786c34c6301"
evidenceAttemptID = "83c9f507-7473-4fa6-8d71-8786c34c6301"
evidenceUploadKey = "93c9f507-7473-4fa6-8d71-8786c34c6301"
evidenceDeviceID = "13c9f507-7473-4fa6-8d71-8786c34c6301"
)
func TestEvidenceUploadAuthenticatesBeforeReadingBody(t *testing.T) {
authenticator := &fakeDeviceAuthenticator{}
router, _ := newRouterWithDependencies(t, &memoryStore{}, emptyDetailStore{}, emptyEvidenceStore{}, authenticator)
poison := &poisonBody{}
request := httptest.NewRequest(http.MethodPost, "/api/v1/tasks/"+evidenceTaskID+"/evidence", nil)
request.Body = poison
request.Header.Set("Content-Type", "text/plain")
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusUnauthorized || response.Body.Len() != 0 || response.Header().Get("WWW-Authenticate") != "Bearer" || poison.reads != 0 || authenticator.calls != 1 {
t.Fatalf("status/reads/auth calls = %d/%d/%d, want 401/0/1", response.Code, poison.reads, authenticator.calls)
}
assertSecurityHeaders(t, response)
}
func TestEvidenceUploadAuthenticationStorageFailureBeforeReadingBody(t *testing.T) {
database, err := sqlite.Open(filepath.Join(t.TempDir(), "authentication-failure.db"))
if err != nil {
t.Fatalf("open database: %v", err)
}
if err := migrations.Up(context.Background(), database, testMigrationDirectory(t)); err != nil {
t.Fatalf("migrate database: %v", err)
}
authenticator, err := deviceauth.NewSQLiteAuthenticator(database)
if err != nil {
t.Fatalf("new authenticator: %v", err)
}
credentialStore, err := deviceauth.NewCredentialStore(database)
if err != nil {
t.Fatalf("new credential store: %v", err)
}
issued, err := credentialStore.Issue(context.Background(), "test device")
if err != nil {
t.Fatalf("issue credential: %v", err)
}
if err := database.Close(); err != nil {
t.Fatalf("close database: %v", err)
}
router, _ := newRouterWithDependencies(t, &memoryStore{}, emptyDetailStore{}, emptyEvidenceStore{}, authenticator)
poison := &poisonBody{}
request := httptest.NewRequest(http.MethodPost, "/api/v1/tasks/"+evidenceTaskID+"/evidence", nil)
request.Body = poison
request.Header.Set(deviceauth.AuthorizationHeader, "Bearer "+issued.Token)
request.Header.Set(deviceauth.DeviceIDHeader, issued.DeviceID)
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusServiceUnavailable || response.Body.Len() != 0 || poison.reads != 0 {
t.Fatalf("storage failure status/body/reads = %d/%q/%d, want 503/empty/0", response.Code, response.Body.String(), poison.reads)
}
}
func TestEvidenceUploadRejectsInvalidSuccessfulPrincipalBeforeReadingBody(t *testing.T) {
router, _ := newRouterWithDependencies(t, &memoryStore{}, emptyDetailStore{}, emptyEvidenceStore{}, uncheckedDeviceAuthenticator{})
poison := &poisonBody{}
request := httptest.NewRequest(http.MethodPost, "/api/v1/tasks/"+evidenceTaskID+"/evidence", nil)
request.Body = poison
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusServiceUnavailable || response.Body.Len() != 0 || poison.reads != 0 {
t.Fatalf("invalid principal status/body/reads = %d/%q/%d, want 503/empty/0", response.Code, response.Body.String(), poison.reads)
}
}
func TestAdminSessionCannotActAsDeviceUploader(t *testing.T) {
router, _ := newRouter(t)
cookie := authenticate(t, router)
request := httptest.NewRequest(http.MethodPost, "/api/v1/tasks/"+evidenceTaskID+"/evidence", nil)
request.Body = &poisonBody{}
request.AddCookie(cookie)
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusUnauthorized {
t.Fatalf("admin upload status = %d, want 401", response.Code)
}
}
func TestRealDeviceCredentialIdentityIsolationAndMixedCredentials(t *testing.T) {
database, err := sqlite.Open(filepath.Join(t.TempDir(), "identity-isolation.db"))
if err != nil {
t.Fatalf("open database: %v", err)
}
t.Cleanup(func() { _ = database.Close() })
if err := migrations.Up(context.Background(), database, testMigrationDirectory(t)); err != nil {
t.Fatalf("migrate database: %v", err)
}
insertEvidenceAttempt(t, database)
assetStore, err := evidencestorage.NewStore(database, filepath.Join(t.TempDir(), "assets"))
if err != nil {
t.Fatalf("new evidence store: %v", err)
}
credentialStore, err := deviceauth.NewCredentialStore(database)
if err != nil {
t.Fatalf("new credential store: %v", err)
}
issued, err := credentialStore.Issue(context.Background(), "采购工具一号")
if err != nil {
t.Fatalf("issue credential: %v", err)
}
insertEvidenceClaim(t, database, issued.DeviceID)
authenticator, err := deviceauth.NewSQLiteAuthenticator(database)
if err != nil {
t.Fatalf("new authenticator: %v", err)
}
taskStore := &memoryStore{}
router, _ := newRouterWithDependencies(t, taskStore, emptyDetailStore{}, assetStore, authenticator)
addDeviceHeaders := func(request *http.Request) {
request.Header.Set(deviceauth.AuthorizationHeader, "Bearer "+issued.Token)
request.Header.Set(deviceauth.DeviceIDHeader, issued.DeviceID)
}
start := newStartRequest(t, validStartBody(), "application/json", "", nil)
addDeviceHeaders(start)
startResponse := httptest.NewRecorder()
router.ServeHTTP(startResponse, start)
create := httptest.NewRequest(http.MethodPost, "/tasks", strings.NewReader("title=device"))
create.Header.Set("Content-Type", "application/x-www-form-urlencoded")
addDeviceHeaders(create)
createResponse := httptest.NewRecorder()
router.ServeHTTP(createResponse, create)
if startResponse.Code != http.StatusUnauthorized || createResponse.Code != http.StatusUnauthorized || taskStore.startCalls != 0 || len(taskStore.drafts) != 0 {
t.Fatalf("device management isolation = start %d/create %d/calls %d/drafts %d", startResponse.Code, createResponse.Code, taskStore.startCalls, len(taskStore.drafts))
}
adminCookie, csrf := authenticatedStartSession(t, router)
mixedWithoutCSRF := newStartRequest(t, validStartBody(), "application/json", "", adminCookie)
addDeviceHeaders(mixedWithoutCSRF)
mixedWithoutCSRFResponse := httptest.NewRecorder()
router.ServeHTTP(mixedWithoutCSRFResponse, mixedWithoutCSRF)
if mixedWithoutCSRFResponse.Code != http.StatusForbidden || taskStore.startCalls != 0 {
t.Fatalf("mixed request bypassed admin CSRF: status/calls=%d/%d", mixedWithoutCSRFResponse.Code, taskStore.startCalls)
}
mixedAdmin := newStartRequest(t, validStartBody(), "application/json", csrf, adminCookie)
addDeviceHeaders(mixedAdmin)
mixedAdminResponse := httptest.NewRecorder()
router.ServeHTTP(mixedAdminResponse, mixedAdmin)
if mixedAdminResponse.Code != http.StatusBadRequest || taskStore.startCalls != 1 {
t.Fatalf("mixed admin request changed identity domain: status/calls=%d/%d", mixedAdminResponse.Code, taskStore.startCalls)
}
pngBytes := serverTestPNG(t, 3, 2)
upload := newEvidenceUploadRequest(t, evidenceTaskID, validEvidenceFields(pngBytes), pngBytes, evidence.PNGContentType, "raw.png", nil)
addDeviceHeaders(upload)
upload.AddCookie(adminCookie)
uploadResponse := httptest.NewRecorder()
router.ServeHTTP(uploadResponse, upload)
if uploadResponse.Code != http.StatusCreated {
t.Fatalf("mixed upload status/body = %d/%q", uploadResponse.Code, uploadResponse.Body.String())
}
var uploadedBy string
if err := database.QueryRow(`SELECT uploaded_by_device_id FROM evidence_assets`).Scan(&uploadedBy); err != nil || uploadedBy != issued.DeviceID {
t.Fatalf("uploaded principal = %q, err=%v", uploadedBy, err)
}
if _, _, err := credentialStore.Revoke(context.Background(), issued.DeviceID); err != nil {
t.Fatalf("revoke credential: %v", err)
}
revokedUpload := newEvidenceUploadRequest(t, evidenceTaskID, validEvidenceFields(pngBytes), pngBytes, evidence.PNGContentType, "raw.png", nil)
addDeviceHeaders(revokedUpload)
revokedResponse := httptest.NewRecorder()
router.ServeHTTP(revokedResponse, revokedUpload)
if revokedResponse.Code != http.StatusUnauthorized || revokedResponse.Body.Len() != 0 {
t.Fatalf("revoked upload = %d/%q", revokedResponse.Code, revokedResponse.Body.String())
}
}
func TestEvidenceUploadReplayConflictAndProtectedRead(t *testing.T) {
router, database := newEvidenceRouter(t, &fakeDeviceAuthenticator{principal: deviceauth.Principal{ID: evidenceDeviceID}})
pngBytes := serverTestPNG(t, 6, 4)
fields := validEvidenceFields(pngBytes)
first := serveEvidenceUpload(t, router, evidenceTaskID, fields, pngBytes, evidence.PNGContentType, `..\private\original.png`, nil)
if first.Code != http.StatusCreated {
t.Fatalf("first upload status/body = %d/%q", first.Code, first.Body.String())
}
var asset evidence.Asset
if err := json.Unmarshal(first.Body.Bytes(), &asset); err != nil {
t.Fatalf("decode upload response: %v", err)
}
if asset.TaskID != evidenceTaskID || asset.AttemptID != evidenceAttemptID || asset.SHA256 != fields["sha256"] || strings.Contains(first.Body.String(), "private") || strings.Contains(first.Body.String(), "original.png") {
t.Fatalf("unsafe upload response = %s", first.Body.String())
}
replay := serveEvidenceUpload(t, router, evidenceTaskID, fields, pngBytes, evidence.PNGContentType, "again.png", nil)
if replay.Code != http.StatusOK {
t.Fatalf("replay status = %d, want 200", replay.Code)
}
var replayed evidence.Asset
if err := json.Unmarshal(replay.Body.Bytes(), &replayed); err != nil || replayed.ID != asset.ID {
t.Fatalf("replay asset = %#v, err %v", replayed, err)
}
conflicting := copyStringMap(fields)
conflicting["captured_at"] = "2026-08-04T09:01:01Z"
if response := serveEvidenceUpload(t, router, evidenceTaskID, conflicting, pngBytes, evidence.PNGContentType, "same.png", nil); response.Code != http.StatusConflict {
t.Fatalf("conflicting replay status = %d, want 409", response.Code)
}
var count int
if err := database.QueryRow("SELECT COUNT(*) FROM evidence_assets").Scan(&count); err != nil || count != 1 {
t.Fatalf("asset count = %d, err %v", count, err)
}
if response := serve(router, http.MethodGet, "/evidence/"+asset.ID, nil, nil); response.Code != http.StatusUnauthorized || response.Body.Len() != 0 {
t.Fatalf("anonymous read = %d/%q", response.Code, response.Body.String())
}
adminCookie := authenticate(t, router)
read := serve(router, http.MethodGet, "/evidence/"+asset.ID, nil, adminCookie)
if read.Code != http.StatusOK || !bytes.Equal(read.Body.Bytes(), pngBytes) {
t.Fatalf("admin read = %d, bytes equal %t", read.Code, bytes.Equal(read.Body.Bytes(), pngBytes))
}
for header, want := range map[string]string{"Content-Type": "image/png", "Cache-Control": "no-store", "X-Content-Type-Options": "nosniff", "Content-Disposition": `inline; filename="evidence.png"`} {
if got := read.Header().Get(header); got != want {
t.Fatalf("%s = %q, want %q", header, got, want)
}
}
missing := serve(router, http.MethodGet, "/evidence/not-a-uuid", nil, adminCookie)
if missing.Code != http.StatusNotFound || missing.Body.Len() != 0 {
t.Fatalf("missing evidence = %d/%q", missing.Code, missing.Body.String())
}
}
func TestEvidenceUploadRejectsStrictMultipartViolations(t *testing.T) {
router, database := newEvidenceRouter(t, &fakeDeviceAuthenticator{principal: deviceauth.Principal{ID: evidenceDeviceID}})
pngBytes := serverTestPNG(t, 2, 2)
base := validEvidenceFields(pngBytes)
wrongHash := copyStringMap(base)
wrongHash["sha256"] = strings.Repeat("b", 64)
uppercaseHash := copyStringMap(base)
uppercaseHash["sha256"] = strings.ToUpper(uppercaseHash["sha256"])
wrongPrivacy := copyStringMap(base)
wrongPrivacy["privacy_tier"] = "PUBLIC"
wrongKind := copyStringMap(base)
wrongKind["kind"] = "ORDER_CONFIRM"
tests := []struct {
name string
fields map[string]string
file []byte
contentType string
extra func(*multipart.Writer) error
want int
}{
{name: "attempt belongs to another task", fields: base, file: pngBytes, contentType: evidence.PNGContentType, want: http.StatusBadRequest},
{name: "xml file", fields: base, file: []byte("<hierarchy/>"), contentType: evidence.PNGContentType, want: http.StatusBadRequest},
{name: "wrong hash", fields: wrongHash, file: pngBytes, contentType: evidence.PNGContentType, want: http.StatusBadRequest},
{name: "uppercase hash", fields: uppercaseHash, file: pngBytes, contentType: evidence.PNGContentType, want: http.StatusBadRequest},
{name: "wrong privacy", fields: wrongPrivacy, file: pngBytes, contentType: evidence.PNGContentType, want: http.StatusBadRequest},
{name: "unapproved kind", fields: wrongKind, file: pngBytes, contentType: evidence.PNGContentType, want: http.StatusBadRequest},
{name: "too large", fields: base, file: make([]byte, evidence.MaxFileBytes+1), contentType: evidence.PNGContentType, want: http.StatusRequestEntityTooLarge},
{name: "wrong file content type", fields: base, file: pngBytes, contentType: "application/xml", want: http.StatusUnsupportedMediaType},
{name: "unknown path field", fields: base, file: pngBytes, contentType: evidence.PNGContentType, extra: func(writer *multipart.Writer) error { return writer.WriteField("path", `C:\secret.xml`) }, want: http.StatusBadRequest},
{name: "duplicate metadata", fields: base, file: pngBytes, contentType: evidence.PNGContentType, extra: func(writer *multipart.Writer) error { return writer.WriteField("sha256", base["sha256"]) }, want: http.StatusBadRequest},
{name: "second file", fields: base, file: pngBytes, contentType: evidence.PNGContentType, extra: func(writer *multipart.Writer) error {
part, err := writer.CreateFormFile("file", "second.png")
if err == nil {
_, err = part.Write(pngBytes)
}
return err
}, want: http.StatusUnsupportedMediaType},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
taskID := evidenceTaskID
if test.name == "attempt belongs to another task" {
taskID = "a3c9f507-7473-4fa6-8d71-8786c34c6301"
}
response := serveEvidenceUpload(t, router, taskID, copyStringMap(test.fields), test.file, test.contentType, "file.png", test.extra)
if response.Code != test.want || response.Body.Len() != 0 {
t.Fatalf("status/body = %d/%q, want %d/empty", response.Code, response.Body.String(), test.want)
}
})
}
var count int
if err := database.QueryRow("SELECT COUNT(*) FROM evidence_assets").Scan(&count); err != nil || count != 0 {
t.Fatalf("invalid requests created %d assets, err %v", count, err)
}
}
type fakeDeviceAuthenticator struct {
principal deviceauth.Principal
err error
calls int
}
func (authenticator *fakeDeviceAuthenticator) Authenticate(*http.Request) (deviceauth.Principal, error) {
authenticator.calls++
if authenticator.err != nil {
return deviceauth.Principal{}, authenticator.err
}
if authenticator.principal.ID == "" {
return deviceauth.Principal{}, deviceauth.ErrUnauthenticated
}
return authenticator.principal, nil
}
type poisonBody struct{ reads int }
type uncheckedDeviceAuthenticator struct{}
func (uncheckedDeviceAuthenticator) Authenticate(*http.Request) (deviceauth.Principal, error) {
return deviceauth.Principal{}, nil
}
func (body *poisonBody) Read([]byte) (int, error) {
body.reads++
return 0, io.ErrUnexpectedEOF
}
func (*poisonBody) Close() error { return nil }
func newEvidenceRouter(t *testing.T, authenticator deviceauth.Authenticator) (http.Handler, *sql.DB) {
t.Helper()
database, err := sqlite.Open(filepath.Join(t.TempDir(), "server-evidence.db"))
if err != nil {
t.Fatalf("open database: %v", err)
}
t.Cleanup(func() { _ = database.Close() })
if err := migrations.Up(context.Background(), database, testMigrationDirectory(t)); err != nil {
t.Fatalf("migrate database: %v", err)
}
insertEvidenceAttempt(t, database)
insertEvidenceClaimDevice(t, database, evidenceDeviceID)
insertEvidenceClaim(t, database, evidenceDeviceID)
store, err := evidencestorage.NewStore(database, filepath.Join(t.TempDir(), "assets"))
if err != nil {
t.Fatalf("new evidence store: %v", err)
}
router, _ := newRouterWithDependencies(t, &memoryStore{}, emptyDetailStore{}, store, authenticator)
return router, database
}
func testMigrationDirectory(t *testing.T) string {
t.Helper()
_, file, _, ok := runtime.Caller(0)
if !ok {
t.Fatal("locate migration directory")
}
return filepath.Join(filepath.Dir(file), "..", "..", "migrations")
}
func insertEvidenceAttempt(t *testing.T, database *sql.DB) {
t.Helper()
timestamp := "2026-08-04T00:00:00Z"
if _, err := database.Exec(`INSERT INTO tasks (id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price, status, version, created_at, updated_at) VALUES (?, 'MANUAL', 'task', '123', 'black', 'M', 1, '1.00', 'DRAFT', 1, ?, ?)`, evidenceTaskID, timestamp, timestamp); err != nil {
t.Fatalf("insert task: %v", err)
}
if _, err := database.Exec(`INSERT INTO order_authorizations (id, task_id, task_version, start_key, goods_id, sku_color, sku_size, quantity, total_price_cap, status, created_by, created_at, expires_at) VALUES (?, ?, 1, 'start', '123', 'black', 'M', 1, '1.00', 'ACTIVE', 'admin', ?, ?)`, evidenceAuthID, evidenceTaskID, timestamp, timestamp); err != nil {
t.Fatalf("insert authorization: %v", err)
}
if _, err := database.Exec(`INSERT INTO purchase_attempts (id, task_id, authorization_id, claim_generation, status, started_at) VALUES (?, ?, ?, 1, 'CLAIMED', ?)`, evidenceAttemptID, evidenceTaskID, evidenceAuthID, timestamp); err != nil {
t.Fatalf("insert attempt: %v", err)
}
}
func insertEvidenceClaimDevice(t *testing.T, database *sql.DB, deviceID string) {
t.Helper()
digest := sha256.Sum256([]byte("fake evidence device"))
if _, err := database.Exec(`INSERT INTO device_credentials
(device_id,display_name,token_sha256,status,created_at,revoked_at)
VALUES (?, 'fake evidence device', ?, 'ACTIVE', '2026-08-04T00:00:00Z', NULL)`, deviceID, digest[:]); err != nil {
t.Fatalf("insert evidence device: %v", err)
}
}
func insertEvidenceClaim(t *testing.T, database *sql.DB, deviceID string) {
t.Helper()
if _, err := database.Exec(`INSERT INTO purchase_attempt_claims
(attempt_id,task_id,authorization_id,claimed_by_device_id,session_id,claim_generation,
task_version,task_title,authorization_task_version,goods_id,sku_color,sku_size,quantity,
total_price_cap,authorization_expires_at,claim_nonce,claim_token_sha256,lease_expires_at,claimed_at,closed_at)
VALUES (?, ?, ?, ?, '23c9f507-7473-4fa6-8d71-8786c34c6301', 1, 1, 'task',
1, '123', 'black', 'M', 1, '1.00', '2026-08-04T00:00:00Z', ?, ?,
'2026-08-04T02:00:00Z', '2026-08-04T00:00:00Z', NULL)`, evidenceAttemptID,
evidenceTaskID, evidenceAuthID, deviceID, bytes.Repeat([]byte{1}, 32), bytes.Repeat([]byte{2}, 32)); err != nil {
t.Fatalf("insert evidence claim: %v", err)
}
}
func serveEvidenceUpload(t *testing.T, router http.Handler, taskID string, fields map[string]string, file []byte, fileContentType, filename string, extra func(*multipart.Writer) error) *httptest.ResponseRecorder {
t.Helper()
request := newEvidenceUploadRequest(t, taskID, fields, file, fileContentType, filename, extra)
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
return response
}
func newEvidenceUploadRequest(t *testing.T, taskID string, fields map[string]string, file []byte, fileContentType, filename string, extra func(*multipart.Writer) error) *http.Request {
t.Helper()
var body bytes.Buffer
writer := multipart.NewWriter(&body)
for _, name := range []string{"upload_key", "attempt_id", "kind", "privacy_tier", "sha256", "captured_at"} {
if err := writer.WriteField(name, fields[name]); err != nil {
t.Fatalf("write field %s: %v", name, err)
}
}
header := make(textproto.MIMEHeader)
header.Set("Content-Disposition", `form-data; name="file"; filename="`+filename+`"`)
header.Set("Content-Type", fileContentType)
part, err := writer.CreatePart(header)
if err != nil {
t.Fatalf("create file part: %v", err)
}
if _, err := part.Write(file); err != nil {
t.Fatalf("write file: %v", err)
}
if extra != nil {
if err := extra(writer); err != nil {
t.Fatalf("write extra part: %v", err)
}
}
if err := writer.Close(); err != nil {
t.Fatalf("close multipart: %v", err)
}
request := httptest.NewRequest(http.MethodPost, "/api/v1/tasks/"+taskID+"/evidence", bytes.NewReader(body.Bytes()))
request.Header.Set("Content-Type", writer.FormDataContentType())
return request
}
func validEvidenceFields(pngBytes []byte) map[string]string {
hash := sha256.Sum256(pngBytes)
return map[string]string{
"upload_key": evidenceUploadKey, "attempt_id": evidenceAttemptID,
"kind": evidence.KindSKUPanelGate1, "privacy_tier": evidence.PrivacyInternalRaw,
"sha256": hex.EncodeToString(hash[:]), "captured_at": "2026-08-04T09:01:00Z",
}
}
func serverTestPNG(t *testing.T, width, height int) []byte {
t.Helper()
var buffer bytes.Buffer
if err := png.Encode(&buffer, image.NewNRGBA(image.Rect(0, 0, width, height))); err != nil {
t.Fatalf("encode PNG: %v", err)
}
return buffer.Bytes()
}
func copyStringMap(values map[string]string) map[string]string {
copy := make(map[string]string, len(values))
for key, value := range values {
copy[key] = value
}
return copy
}
+288 -13
View File
@@ -2,13 +2,23 @@
package server
import (
"bytes"
"crypto/subtle"
"encoding/json"
"errors"
"io"
"mime"
"net/http"
"net/url"
"strings"
"unicode/utf8"
"cmbuyer/admin/internal/auth"
"cmbuyer/admin/internal/deviceauth"
"cmbuyer/admin/internal/evidence"
"cmbuyer/admin/internal/taskclaim"
"cmbuyer/admin/internal/taskdetail"
"cmbuyer/admin/internal/tasks"
"cmbuyer/admin/internal/transport/webui"
"github.com/gin-gonic/gin"
@@ -16,17 +26,23 @@ import (
)
const maxFormBytes = 8 << 10
const maxJSONBytes = 64 << 10
// Options 是路由层需要的安全依赖。凭据由启动配置注入,不能在路由中设置默认值。
type Options struct {
AdminUsername string
AdminPasswordBcrypt string
Sessions *auth.Manager
Tasks tasks.Store
TaskDetails taskdetail.Store
Evidence evidence.Store
DeviceAuthenticator deviceauth.Authenticator
TaskClaims taskclaim.Service
}
// NewRouter 返回当前服务范围内的完整 HTTP 路由。
func NewRouter(options Options) (*gin.Engine, error) {
if options.AdminUsername == "" || options.AdminPasswordBcrypt == "" || options.Sessions == nil {
if options.AdminUsername == "" || options.AdminPasswordBcrypt == "" || options.Sessions == nil || options.Tasks == nil || options.TaskDetails == nil || options.Evidence == nil || options.DeviceAuthenticator == nil || options.TaskClaims == nil {
return nil, errors.New("server authentication options are incomplete")
}
@@ -38,10 +54,91 @@ func NewRouter(options Options) (*gin.Engine, error) {
router.POST("/login", login(options))
router.POST("/logout", logout(options))
router.GET("/tasks", tasksPage(options))
router.GET("/tasks/:id", taskDetailPage(options))
router.GET("/tasks/new", newTaskPage(options))
router.POST("/tasks", createTask(options))
router.POST("/tasks/start-purchases", startPurchases(options))
router.POST("/api/v1/tasks/:id/evidence", uploadEvidence(options))
router.POST("/api/v1/tasks/claim-next", claimNext(options))
router.POST("/api/v1/tasks/:id/lease/renew", renewLease(options))
router.GET("/evidence/:asset_id", readEvidence(options))
router.GET("/static/tasks.js", func(context *gin.Context) {
context.Data(http.StatusOK, "application/javascript; charset=utf-8", webui.TasksScript())
})
return router, nil
}
func startPurchases(options Options) gin.HandlerFunc {
return func(context *gin.Context) {
if !options.Sessions.IsAuthenticated(context.Request) {
context.Status(http.StatusUnauthorized)
return
}
authenticated, csrfOK := options.Sessions.VerifyCSRF(context.Request, context.GetHeader("X-CSRF-Token"))
if !authenticated || !csrfOK {
context.Status(http.StatusForbidden)
return
}
if !isJSONContentType(context.GetHeader("Content-Type")) {
context.Status(http.StatusUnsupportedMediaType)
return
}
context.Request.Body = http.MaxBytesReader(context.Writer, context.Request.Body, maxJSONBytes)
raw, err := io.ReadAll(context.Request.Body)
if err != nil {
var tooLarge *http.MaxBytesError
if errors.As(err, &tooLarge) {
context.Status(http.StatusRequestEntityTooLarge)
} else {
context.Status(http.StatusBadRequest)
}
return
}
if !utf8.Valid(raw) {
context.Status(http.StatusBadRequest)
return
}
decoder := json.NewDecoder(bytes.NewReader(raw))
decoder.DisallowUnknownFields()
var command tasks.StartCommand
if err := decoder.Decode(&command); err != nil {
context.Status(http.StatusBadRequest)
return
}
var extra any
if err := decoder.Decode(&extra); err != io.EOF {
context.Status(http.StatusBadRequest)
return
}
result, err := options.Tasks.StartPurchases(context.Request.Context(), command, options.AdminUsername)
if err != nil {
if errors.Is(err, tasks.ErrInvalidStart) {
context.Status(http.StatusBadRequest)
} else if errors.Is(err, tasks.ErrStartConflict) {
context.Status(http.StatusConflict)
} else {
context.Status(http.StatusInternalServerError)
}
return
}
context.JSON(http.StatusOK, result)
}
}
func isJSONContentType(value string) bool {
mediaType, parameters, err := mime.ParseMediaType(value)
if err != nil || mediaType != "application/json" {
return false
}
for name, value := range parameters {
if name != "charset" || !strings.EqualFold(value, "utf-8") {
return false
}
}
return true
}
func healthz(context *gin.Context) {
context.JSON(http.StatusOK, gin.H{"status": "ok"})
}
@@ -51,7 +148,7 @@ func securityHeaders() gin.HandlerFunc {
context.Header("Cache-Control", "no-store")
context.Header("X-Content-Type-Options", "nosniff")
context.Header("Referrer-Policy", "no-referrer")
context.Header("Content-Security-Policy", "default-src 'self'; style-src 'self' 'unsafe-inline'; script-src 'none'; object-src 'none'; base-uri 'none'; frame-ancestors 'none'; form-action 'self'")
context.Header("Content-Security-Policy", "default-src 'self'; style-src 'self' 'unsafe-inline'; script-src 'self'; object-src 'none'; base-uri 'none'; frame-ancestors 'none'; form-action 'self'")
context.Next()
}
}
@@ -70,11 +167,14 @@ func loginPage(options Options) gin.HandlerFunc {
func login(options Options) gin.HandlerFunc {
return func(context *gin.Context) {
limitFormBody(context)
csrfToken := context.PostForm("csrf_token")
returnPath := returnTo(context.PostForm("return_to"))
username := context.PostForm("username")
password := context.PostForm("password")
if !parseForm(context) {
return
}
form := context.Request.PostForm
csrfToken := form.Get("csrf_token")
returnPath := returnTo(form.Get("return_to"))
username := form.Get("username")
password := form.Get("password")
if _, ok := options.Sessions.VerifyCSRF(context.Request, csrfToken); !ok {
newCSRF, _ := options.Sessions.Ensure(context.Writer, context.Request)
@@ -97,8 +197,10 @@ func login(options Options) gin.HandlerFunc {
func logout(options Options) gin.HandlerFunc {
return func(context *gin.Context) {
limitFormBody(context)
authenticated, ok := options.Sessions.VerifyCSRF(context.Request, context.PostForm("csrf_token"))
if !parseForm(context) {
return
}
authenticated, ok := options.Sessions.VerifyCSRF(context.Request, context.Request.PostForm.Get("csrf_token"))
if !ok || !authenticated {
context.Status(http.StatusForbidden)
return
@@ -117,10 +219,164 @@ func tasksPage(options Options) gin.HandlerFunc {
return
}
context.Header("Content-Type", "text/html; charset=utf-8")
if err := webui.RenderTasks(context.Writer, webui.TasksData{CSRFToken: csrfToken}); err != nil {
_ = context.Error(err)
filter := tasks.TaskFilter{Keyword: context.Query("keyword"), Status: context.Query("status"), CreatedFrom: context.Query("created_from"), CreatedTo: context.Query("created_to")}
if validation := tasks.ValidateTaskFilter(filter); !validation.Valid() {
startKey, err := tasks.NewCreateKey()
if err != nil {
context.Status(http.StatusInternalServerError)
return
}
renderTasks(context, http.StatusBadRequest, webui.TasksData{CSRFToken: csrfToken, Filter: filter, FilterErrors: validation, HasFilter: true, StartKey: startKey})
return
}
data, err := taskListData(context, options, csrfToken, filter)
if err != nil {
context.Status(http.StatusInternalServerError)
return
}
for _, row := range data.Tasks {
if row.ID == context.Query("created") {
data.Success = true
break
}
}
if context.Query("create") == "1" {
form, err := newTaskForm()
if err != nil {
context.Status(http.StatusInternalServerError)
return
}
data.OpenForm = true
data.Form = form
data.FocusField = "title"
}
renderTasks(context, http.StatusOK, data)
}
}
func newTaskPage(options Options) gin.HandlerFunc {
return func(context *gin.Context) {
csrf, authenticated := options.Sessions.Ensure(context.Writer, context.Request)
if !authenticated {
context.Redirect(http.StatusSeeOther, "/login?return_to=%2Ftasks%2Fnew")
return
}
form, err := newTaskForm()
if err != nil {
context.Status(http.StatusInternalServerError)
return
}
renderTasks(context, http.StatusOK, webui.TasksData{CSRFToken: csrf, Form: form, FullPage: true, FocusField: "title"})
}
}
func createTask(options Options) gin.HandlerFunc {
return func(context *gin.Context) {
if !options.Sessions.IsAuthenticated(context.Request) {
context.Status(http.StatusUnauthorized)
return
}
if !parseForm(context) {
return
}
requestForm := context.Request.PostForm
authenticated, csrfOK := options.Sessions.VerifyCSRF(context.Request, requestForm.Get("csrf_token"))
if !csrfOK || !authenticated {
context.Status(http.StatusForbidden)
return
}
form := taskForm(requestForm)
draft, validation := tasks.Validate(form)
if draft.GoodsID != "" {
form.ProductURL = tasks.CanonicalURL(draft.GoodsID)
}
fullPage := requestForm.Get("form_mode") == "full"
if !validation.Valid() {
data, ok := createErrorData(context, options, fullPage)
if !ok {
return
}
data.Form = form
data.Errors = validation
data.OpenForm = !fullPage
data.FullPage = fullPage
data.FocusField = firstError(validation)
renderTasks(context, http.StatusBadRequest, data)
return
}
created, err := options.Tasks.CreateDraft(context.Request.Context(), draft)
if err != nil {
if errors.Is(err, tasks.ErrCreateKeyConflict) {
validation["create_key"] = "该创建请求已用于另一条任务,请重新打开表单。"
data, ok := createErrorData(context, options, fullPage)
if !ok {
return
}
data.Form = form
data.Errors = validation
data.OpenForm = !fullPage
data.FullPage = fullPage
data.FocusField = firstError(validation)
renderTasks(context, http.StatusConflict, data)
return
}
context.Status(http.StatusInternalServerError)
return
}
context.Redirect(http.StatusSeeOther, "/tasks?created="+url.QueryEscape(created.ID))
}
}
func newTaskForm() (tasks.Form, error) {
key, err := tasks.NewCreateKey()
if err != nil {
return tasks.Form{}, err
}
return tasks.Form{CreateKey: key}, nil
}
func taskForm(form url.Values) tasks.Form {
return tasks.Form{CreateKey: form.Get("create_key"), Title: form.Get("title"), ProductURL: form.Get("product_url"), SKUColor: form.Get("sku_color"), SKUSize: form.Get("sku_size"), Quantity: form.Get("quantity"), MaxTotalPrice: form.Get("max_total_price")}
}
func csrfFor(context *gin.Context, options Options) string {
csrf, _ := options.Sessions.Ensure(context.Writer, context.Request)
return csrf
}
func taskListData(context *gin.Context, options Options, csrfToken string, filter tasks.TaskFilter) (webui.TasksData, error) {
rows, err := options.Tasks.ListTasks(context.Request.Context(), filter)
if err != nil {
return webui.TasksData{}, err
}
startKey, err := tasks.NewCreateKey()
if err != nil {
return webui.TasksData{}, err
}
return webui.TasksData{
CSRFToken: csrfToken,
Tasks: rows,
Filter: filter,
HasFilter: filter.Keyword != "" || filter.Status != "" || filter.CreatedFrom != "" || filter.CreatedTo != "",
StartKey: startKey,
}, nil
}
func createErrorData(context *gin.Context, options Options, fullPage bool) (webui.TasksData, bool) {
csrfToken := csrfFor(context, options)
if fullPage {
return webui.TasksData{CSRFToken: csrfToken}, true
}
data, err := taskListData(context, options, csrfToken, tasks.TaskFilter{})
if err != nil {
context.Status(http.StatusInternalServerError)
return webui.TasksData{}, false
}
return data, true
}
func renderTasks(context *gin.Context, status int, data webui.TasksData) {
context.Header("Content-Type", "text/html; charset=utf-8")
context.Status(status)
if err := webui.RenderTasks(context.Writer, data); err != nil {
_ = context.Error(err)
}
}
@@ -137,8 +393,27 @@ func renderLogin(context *gin.Context, status int, csrfToken, returnPath, userna
}
}
func limitFormBody(context *gin.Context) {
func parseForm(context *gin.Context) bool {
context.Request.Body = http.MaxBytesReader(context.Writer, context.Request.Body, maxFormBytes)
if err := context.Request.ParseForm(); err != nil {
var tooLarge *http.MaxBytesError
if errors.As(err, &tooLarge) {
context.Status(http.StatusRequestEntityTooLarge)
} else {
context.Status(http.StatusBadRequest)
}
return false
}
return true
}
func firstError(validation tasks.Errors) string {
for _, field := range []string{"title", "product_url", "sku_color", "sku_size", "quantity", "max_total_price"} {
if _, ok := validation[field]; ok {
return field
}
}
return "title"
}
func returnTo(value string) string {
+361 -1
View File
@@ -1,21 +1,30 @@
package server_test
import (
"context"
"io"
"net/http"
"net/http/httptest"
"net/url"
"regexp"
"strings"
"testing"
"time"
"cmbuyer/admin/internal/auth"
"cmbuyer/admin/internal/deviceauth"
"cmbuyer/admin/internal/evidence"
"cmbuyer/admin/internal/server"
"cmbuyer/admin/internal/taskclaim"
"cmbuyer/admin/internal/taskdetail"
"cmbuyer/admin/internal/tasks"
"github.com/gin-gonic/gin"
"golang.org/x/crypto/bcrypt"
)
var csrfPattern = regexp.MustCompile(`name="csrf_token" value="([^"]+)"`)
var createKeyPattern = regexp.MustCompile(`name="create_key" value="([^"]+)"`)
func TestHealthzIsPublic(t *testing.T) {
router, _ := newRouter(t)
@@ -175,13 +184,255 @@ func TestTamperedCookieCannotAccessTasks(t *testing.T) {
}
func TestTaskCreationRendersSharedFormsAndPersistsOnlyDraft(t *testing.T) {
router, _ := newRouter(t)
cookie := authenticate(t, router)
modal := serve(router, http.MethodGet, "/tasks?create=1", nil, cookie)
if modal.Code != http.StatusOK {
t.Fatalf("GET dialog form status = %d, want 200", modal.Code)
}
fullPage := serve(router, http.MethodGet, "/tasks/new", nil, cookie)
if fullPage.Code != http.StatusOK {
t.Fatalf("GET full form status = %d, want 200", fullPage.Code)
}
for _, want := range []string{`<div class="modal-scrim"`, `<dialog open`, `aria-modal="true"`, `name="title"`, `name="product_url"`, `name="sku_color"`, `name="sku_size"`, `name="quantity"`, `name="max_total_price"`, `type="url" inputmode="url" maxlength="2048"`, `type="number" inputmode="numeric" min="1" step="1"`, `inputmode="decimal" pattern="[0-9]+(\.[0-9]{1,2})?"`, `maxlength="120"`, `maxlength="80"`, `required`, `autofocus`, `导入</button><a class="button primary"`, `type="search"`, `data-start-purchases`, `data-select-all`, `最高总额`, `min-height:44px`, `:focus-visible`, `overflow-x:auto`, `prefers-reduced-motion`} {
if !strings.Contains(modal.Body.String(), want) {
t.Fatalf("dialog form is missing %q", want)
}
}
for _, want := range []string{`name="title"`, `name="product_url"`, `name="sku_color"`, `name="sku_size"`, `name="quantity"`, `name="max_total_price"`, `name="form_mode" value="full"`} {
if !strings.Contains(fullPage.Body.String(), want) {
t.Fatalf("full-page form is missing %q", want)
}
}
invalid := serve(router, http.MethodPost, "/tasks", url.Values{
"csrf_token": {csrfToken(t, modal.Body.String())},
"create_key": {createKey(t, modal.Body.String())},
"title": {`<script>alert(1)</script>`},
"product_url": {"https://mobile.yangkeduo.com/goods.html?goods_id=937122477375&uin=discard"},
"sku_color": {"black"},
"sku_size": {"M"},
"quantity": {"0"},
"max_total_price": {"12.80"},
"form_mode": {"dialog"},
}, cookie)
if invalid.Code != http.StatusBadRequest || !strings.Contains(invalid.Body.String(), `<dialog open`) || !strings.Contains(invalid.Body.String(), "数量必须是正整数") || !strings.Contains(invalid.Body.String(), `role="alert"`) || !strings.Contains(invalid.Body.String(), `href="#quantity"`) || !strings.Contains(invalid.Body.String(), `aria-describedby="quantity-error"`) || !strings.Contains(invalid.Body.String(), `autofocus`) {
t.Fatalf("invalid create = (%d, %q), want dialog validation response", invalid.Code, invalid.Body.String())
}
if strings.Contains(invalid.Body.String(), `<script>alert(1)</script>`) || !strings.Contains(invalid.Body.String(), `&lt;script&gt;alert(1)&lt;/script&gt;`) {
t.Fatalf("invalid create did not safely preserve title: %q", invalid.Body.String())
}
if strings.Contains(invalid.Body.String(), "uin=discard") || !strings.Contains(invalid.Body.String(), `value="https://mobile.yangkeduo.com/goods.html?goods_id=937122477375"`) {
t.Fatalf("invalid create did not canonicalize product URL: %q", invalid.Body.String())
}
createPage := serve(router, http.MethodGet, "/tasks?create=1", nil, cookie)
key := createKey(t, createPage.Body.String())
created := serve(router, http.MethodPost, "/tasks", url.Values{
"csrf_token": {csrfToken(t, createPage.Body.String())},
"create_key": {key},
"title": {"<b>夏季上衣</b>"},
"product_url": {"https://mobile.yangkeduo.com/goods.html?goods_id=937122477375&utm_source=discard"},
"sku_color": {"black"},
"sku_size": {"M"},
"quantity": {"2"},
"max_total_price": {"12.8"},
"form_mode": {"dialog"},
}, cookie)
if created.Code != http.StatusSeeOther || !strings.HasPrefix(created.Header().Get("Location"), "/tasks?created=") {
t.Fatalf("valid create = (%d, %q), want 303 to a created-task acknowledgement", created.Code, created.Header().Get("Location"))
}
replay := serve(router, http.MethodPost, "/tasks", url.Values{
"csrf_token": {csrfToken(t, createPage.Body.String())},
"create_key": {key},
"title": {"<b>夏季上衣</b>"},
"product_url": {"https://mobile.yangkeduo.com/goods.html?goods_id=937122477375&utm_source=discard"},
"sku_color": {"black"},
"sku_size": {"M"},
"quantity": {"2"},
"max_total_price": {"12.8"},
"form_mode": {"dialog"},
}, cookie)
if replay.Code != http.StatusSeeOther {
t.Fatalf("idempotent replay status = %d, want 303", replay.Code)
}
conflict := serve(router, http.MethodPost, "/tasks", url.Values{
"csrf_token": {csrfToken(t, createPage.Body.String())},
"create_key": {key},
"title": {"different task"},
"product_url": {"https://mobile.yangkeduo.com/goods.html?goods_id=937122477375"},
"sku_color": {"black"},
"sku_size": {"M"},
"quantity": {"2"},
"max_total_price": {"12.80"},
"form_mode": {"dialog"},
}, cookie)
if conflict.Code != http.StatusConflict || !strings.Contains(conflict.Body.String(), "该创建请求已用于另一条任务") {
t.Fatalf("conflicting create = (%d, %q), want a 409 form error", conflict.Code, conflict.Body.String())
}
list := serve(router, http.MethodGet, created.Header().Get("Location"), nil, cookie)
if list.Code != http.StatusOK {
t.Fatalf("GET /tasks status = %d, want 200", list.Code)
}
body := list.Body.String()
for _, want := range []string{`任务已创建,已显示在列表首行。`, `&lt;b&gt;夏季上衣&lt;/b&gt;`, `https://mobile.yangkeduo.com/goods.html?goods_id=937122477375`, `target="_blank"`, `rel="noopener noreferrer"`, `¥12.80`, `待开始`, `选择全部任务`, `选择任务`} {
if !strings.Contains(body, want) {
t.Fatalf("task list is missing %q", want)
}
}
for _, forbidden := range []string{"utm_source", "试选", "订单确认", "真机", "提交订单"} {
if strings.Contains(body, forbidden) {
t.Fatalf("task list exposed deferred scope %q", forbidden)
}
}
}
func TestTasksPageKeepsOriginalShellAndRendersFilteredWorkbench(t *testing.T) {
store := &memoryStore{rows: []tasks.TaskRow{
{ID: "b3c9f507-7473-4fa6-8d71-8786c34c6301", Title: "待开始衬衫", GoodsID: "937122477375", SKUColor: "黑色", SKUSize: "M", Quantity: 2, MaxTotalPrice: "12.80", Status: "DRAFT", Version: 3, CreatedAt: time.Date(2026, 8, 4, 1, 2, 3, 0, time.UTC)},
{ID: "c3c9f507-7473-4fa6-8d71-8786c34c6301", Title: "等待领取衬衫", GoodsID: "958756616606", SKUColor: "白色", SKUSize: "L", Quantity: 1, MaxTotalPrice: "20.00", Status: "PENDING", Version: 4, CreatedAt: time.Date(2026, 8, 4, 2, 3, 4, 0, time.UTC)},
}}
router, _ := newRouterWithStore(t, store)
cookie := authenticate(t, router)
query := url.Values{"keyword": {"衬衫"}, "created_from": {"2026-08-04"}, "created_to": {"2026-08-04"}}
response := serve(router, http.MethodGet, "/tasks?"+query.Encode(), nil, cookie)
if response.Code != http.StatusOK {
t.Fatalf("filtered tasks status = %d, want 200", response.Code)
}
body := response.Body.String()
for _, want := range []string{
`<a class="skip" href="#main">`,
`:focus-visible`,
`min-height:44px`,
`@media(max-width:420px)`,
`prefers-reduced-motion`,
`<button class="button" type="button" disabled>导入</button><a class="button primary" href="/tasks?create=1">创建任务</a>`,
`name="keyword" type="search" value="衬衫"`,
`name="created_from" type="date" value="2026-08-04"`,
`name="created_to" type="date" value="2026-08-04"`,
`data-start-purchases`,
`data-selection-summary aria-live="polite"`,
`系统不会付款`,
`开始采购(只创建待付款订单)`,
`采购结果`,
`创建时间(上海)`,
`https://mobile.yangkeduo.com/goods.html?goods_id=937122477375`,
`target="_blank" rel="noopener noreferrer"`,
`data-task-row data-detail-url="/tasks/b3c9f507-7473-4fa6-8d71-8786c34c6301" tabindex="0"`,
`data-open-detail>查看详情</button>`,
`.detail-link-button{display:block;min-height:44px`,
`data-detail-drawer aria-modal="true"`,
`待开始`,
`已授权待领取`,
`datetime="2026-08-04T09:02:03&#43;08:00">2026-08-04 09:02`,
`<script src="/static/tasks.js" defer></script>`,
} {
if !strings.Contains(body, want) {
t.Fatalf("workbench is missing %q", want)
}
}
if strings.Index(body, `name="keyword"`) > strings.Index(body, `data-start-purchases`) || strings.Index(body, `data-start-purchases`) > strings.Index(body, `<div class="table-wrap">`) {
t.Fatal("workbench rows are not ordered as toolbar, filters, batch actions, table")
}
if count := strings.Count(body, `data-task-id=`); count != 1 {
t.Fatalf("selectable row count = %d, want only the DRAFT row", count)
}
for _, forbidden := range []string{`<th scope="col">操作</th>`, `确认开始采购`, `确认机器选对了吗`} {
if strings.Contains(body, forbidden) {
t.Fatalf("workbench exposed forbidden per-row or confirmation UI %q", forbidden)
}
}
if store.listTasksCalls != 1 || store.listDraftsCalls != 0 {
t.Fatalf("GET /tasks calls = (ListTasks %d, ListDrafts %d), want (1, 0)", store.listTasksCalls, store.listDraftsCalls)
}
}
func TestTasksPageRerendersAccessibleFilterErrorsAndKeepsValues(t *testing.T) {
store := &memoryStore{}
router, _ := newRouterWithStore(t, store)
cookie := authenticate(t, router)
query := url.Values{
"keyword": {`保留%_\`},
"status": {"UNKNOWN"},
"created_from": {"2026-02-30"},
"created_to": {"not-a-date"},
}
response := serve(router, http.MethodGet, "/tasks?"+query.Encode(), nil, cookie)
if response.Code != http.StatusBadRequest {
t.Fatalf("invalid filter status = %d, want 400", response.Code)
}
body := response.Body.String()
for _, want := range []string{
`role="alert" aria-live="assertive"`,
`href="#filter-status"`,
`href="#filter-created-from"`,
`href="#filter-created-to"`,
`name="keyword" type="search" value="保留%_\"`,
`<option value="UNKNOWN" selected>无效状态:UNKNOWN</option>`,
`name="created_from" type="date" value="2026-02-30" aria-invalid="true" aria-describedby="filter-created-from-error"`,
`name="created_to" type="date" value="not-a-date" aria-invalid="true" aria-describedby="filter-created-to-error"`,
`id="filter-status-error"`,
`id="filter-created-from-error"`,
`id="filter-created-to-error"`,
`筛选条件有误`,
} {
if !strings.Contains(body, want) {
t.Fatalf("invalid filter page is missing %q", want)
}
}
if store.listTasksCalls != 0 || store.listDraftsCalls != 0 {
t.Fatalf("invalid filter queried stores: ListTasks=%d ListDrafts=%d", store.listTasksCalls, store.listDraftsCalls)
}
assertSecurityHeaders(t, response)
}
func TestTaskCreationRequiresAuthenticationAndCSRF(t *testing.T) {
router, _ := newRouter(t)
if response := serve(router, http.MethodPost, "/tasks", url.Values{}, nil); response.Code != http.StatusUnauthorized {
t.Fatalf("anonymous POST /tasks = %d, want 401", response.Code)
}
cookie := authenticate(t, router)
if response := serve(router, http.MethodPost, "/tasks", url.Values{}, cookie); response.Code != http.StatusForbidden {
t.Fatalf("POST /tasks without CSRF = %d, want 403", response.Code)
}
}
func TestTaskCreationFailsClosedForMalformedOrOversizedForms(t *testing.T) {
router, _ := newRouter(t)
cookie := authenticate(t, router)
page := serve(router, http.MethodGet, "/tasks?create=1", nil, cookie)
base := url.Values{
"csrf_token": {csrfToken(t, page.Body.String())},
"create_key": {createKey(t, page.Body.String())},
"title": {"title"},
"product_url": {"https://mobile.yangkeduo.com/goods.html?goods_id=1;uin=malformed"},
"sku_color": {"black"},
"sku_size": {"M"},
"quantity": {"1"},
"max_total_price": {"1.00"},
"form_mode": {"dialog"},
}
malformed := serve(router, http.MethodPost, "/tasks", base, cookie)
if malformed.Code != http.StatusBadRequest || !strings.Contains(malformed.Body.String(), "canonical 商品链接") {
t.Fatalf("malformed URL create = (%d, %q), want validation failure", malformed.Code, malformed.Body.String())
}
oversized := url.Values{"csrf_token": {csrfToken(t, page.Body.String())}, "title": {strings.Repeat("x", 9<<10)}}
if response := serve(router, http.MethodPost, "/tasks", oversized, cookie); response.Code != http.StatusRequestEntityTooLarge {
t.Fatalf("oversized form status = %d, want 413", response.Code)
}
}
func assertSecurityHeaders(t *testing.T, response *httptest.ResponseRecorder) {
t.Helper()
want := map[string]string{
"Cache-Control": "no-store",
"X-Content-Type-Options": "nosniff",
"Referrer-Policy": "no-referrer",
"Content-Security-Policy": "default-src 'self'; style-src 'self' 'unsafe-inline'; script-src 'none'; object-src 'none'; base-uri 'none'; frame-ancestors 'none'; form-action 'self'",
"Content-Security-Policy": "default-src 'self'; style-src 'self' 'unsafe-inline'; script-src 'self'; object-src 'none'; base-uri 'none'; frame-ancestors 'none'; form-action 'self'",
}
for name, expected := range want {
if got := response.Header().Get(name); got != expected {
@@ -235,6 +486,18 @@ func TestLogoutRequiresCSRFAndRevokesSession(t *testing.T) {
}
func newRouter(t *testing.T) (*gin.Engine, *auth.Manager) {
return newRouterWithStore(t, &memoryStore{})
}
func newRouterWithStore(t *testing.T, store tasks.Store) (*gin.Engine, *auth.Manager) {
return newRouterWithDependencies(t, store, emptyDetailStore{}, emptyEvidenceStore{}, deviceauth.RejectAllAuthenticator{})
}
func newRouterWithDependencies(t *testing.T, store tasks.Store, details taskdetail.Store, evidenceStore evidence.Store, deviceAuthenticator deviceauth.Authenticator) (*gin.Engine, *auth.Manager) {
return newRouterWithClaimService(t, store, details, evidenceStore, deviceAuthenticator, emptyTaskClaimService{})
}
func newRouterWithClaimService(t *testing.T, store tasks.Store, details taskdetail.Store, evidenceStore evidence.Store, deviceAuthenticator deviceauth.Authenticator, claims taskclaim.Service) (*gin.Engine, *auth.Manager) {
t.Helper()
gin.SetMode(gin.TestMode)
hash, err := bcrypt.GenerateFromPassword([]byte("test-password"), bcrypt.MinCost)
@@ -246,6 +509,11 @@ func newRouter(t *testing.T) (*gin.Engine, *auth.Manager) {
AdminUsername: "admin",
AdminPasswordBcrypt: string(hash),
Sessions: manager,
Tasks: store,
TaskDetails: details,
Evidence: evidenceStore,
DeviceAuthenticator: deviceAuthenticator,
TaskClaims: claims,
})
if err != nil {
t.Fatalf("NewRouter: %v", err)
@@ -253,6 +521,75 @@ func newRouter(t *testing.T) (*gin.Engine, *auth.Manager) {
return router, manager
}
type emptyDetailStore struct{}
type emptyTaskClaimService struct{}
func (emptyTaskClaimService) ClaimNext(context.Context, string, taskclaim.ClaimCommand) (taskclaim.ClaimResponse, bool, error) {
return taskclaim.ClaimResponse{}, false, nil
}
func (emptyTaskClaimService) Renew(context.Context, string, taskclaim.RenewCommand) (taskclaim.RenewResponse, error) {
return taskclaim.RenewResponse{}, taskclaim.ErrNotCurrent
}
func (emptyDetailStore) Get(context.Context, string) (taskdetail.Detail, error) {
return taskdetail.Detail{}, taskdetail.ErrNotFound
}
type emptyEvidenceStore struct{}
func (emptyEvidenceStore) Stage(io.Reader, string) (evidence.StagedFile, error) {
return evidence.StagedFile{}, evidence.ErrInvalid
}
func (emptyEvidenceStore) Discard(evidence.StagedFile) {}
func (emptyEvidenceStore) Commit(context.Context, deviceauth.Principal, evidence.UploadMetadata, evidence.StagedFile) (evidence.Asset, bool, error) {
return evidence.Asset{}, false, evidence.ErrInvalid
}
func (emptyEvidenceStore) Open(context.Context, string) (evidence.Asset, io.ReadSeekCloser, error) {
return evidence.Asset{}, nil, evidence.ErrNotFound
}
type memoryStore struct {
drafts []tasks.Draft
rows []tasks.TaskRow
listDraftsCalls int
listTasksCalls int
startCalls int
}
func (store *memoryStore) CreateDraft(_ context.Context, draft tasks.Draft) (tasks.Draft, error) {
for _, existing := range store.drafts {
if existing.ID == draft.ID {
if existing.Title != draft.Title || existing.GoodsID != draft.GoodsID || existing.SKUColor != draft.SKUColor || existing.SKUSize != draft.SKUSize || existing.Quantity != draft.Quantity || existing.MaxTotalPrice != draft.MaxTotalPrice {
return tasks.Draft{}, tasks.ErrCreateKeyConflict
}
return existing, nil
}
}
store.drafts = append(store.drafts, draft)
return draft, nil
}
func (store *memoryStore) ListDrafts(_ context.Context) ([]tasks.Draft, error) {
store.listDraftsCalls++
return append([]tasks.Draft(nil), store.drafts...), nil
}
func (store *memoryStore) ListTasks(_ context.Context, _ tasks.TaskFilter) ([]tasks.TaskRow, error) {
store.listTasksCalls++
if store.rows != nil {
return append([]tasks.TaskRow(nil), store.rows...), nil
}
result := make([]tasks.TaskRow, 0, len(store.drafts))
for _, draft := range store.drafts {
result = append(result, tasks.TaskRow{ID: draft.ID, Title: draft.Title, GoodsID: draft.GoodsID, SKUColor: draft.SKUColor, SKUSize: draft.SKUSize, Quantity: draft.Quantity, MaxTotalPrice: draft.MaxTotalPrice, Status: "DRAFT", Version: 1, CreatedAt: draft.CreatedAt})
}
return result, nil
}
func (store *memoryStore) StartPurchases(_ context.Context, _ tasks.StartCommand, _ string) (tasks.StartResult, error) {
store.startCalls++
return tasks.StartResult{}, tasks.ErrInvalidStart
}
func serve(router http.Handler, method, target string, form url.Values, cookie *http.Cookie) *httptest.ResponseRecorder {
var body *strings.Reader
if form == nil {
@@ -291,3 +628,26 @@ func csrfToken(t *testing.T, body string) string {
}
return matches[1]
}
func createKey(t *testing.T, body string) string {
t.Helper()
matches := createKeyPattern.FindStringSubmatch(body)
if len(matches) != 2 || matches[1] == "" {
t.Fatalf("no create key in response body: %q", body)
}
return matches[1]
}
func authenticate(t *testing.T, router http.Handler) *http.Cookie {
t.Helper()
page := serve(router, http.MethodGet, "/login", nil, nil)
login := serve(router, http.MethodPost, "/login", url.Values{
"csrf_token": {csrfToken(t, page.Body.String())},
"username": {"admin"},
"password": {"test-password"},
}, sessionCookie(t, page))
if login.Code != http.StatusSeeOther {
t.Fatalf("authenticate status = %d, want 303", login.Code)
}
return sessionCookie(t, login)
}
@@ -0,0 +1,328 @@
package server_test
import (
"bytes"
"context"
"encoding/json"
"errors"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"cmbuyer/admin/internal/tasks"
)
const (
startKeyForHTTP = "c3c9f507-7473-4fa6-8d71-8786c34c6301"
taskIDForHTTP = "a3c9f507-7473-4fa6-8d71-8786c34c6301"
)
func TestStartPurchasesAuthenticatesBeforeInspectingRequestBody(t *testing.T) {
store := &startRecordingStore{}
router, _ := newRouterWithStore(t, store)
hugeMalformed := `{"start_key":"` + strings.Repeat("x", 70<<10)
for name, request := range map[string]*http.Request{
"anonymous malformed": newStartRequest(t, hugeMalformed, "text/plain", "", nil),
"device bearer": newStartRequest(t, validStartBody(), "application/json", "", nil),
} {
t.Run(name, func(t *testing.T) {
if name == "device bearer" {
request.Header.Set("Authorization", "Bearer device-token")
}
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusUnauthorized {
t.Fatalf("status = %d, want 401", response.Code)
}
})
}
cookie, csrf := authenticatedStartSession(t, router)
for name, token := range map[string]string{"missing CSRF": "", "wrong CSRF": "wrong-csrf"} {
request := newStartRequest(t, hugeMalformed, "text/plain", token, cookie)
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != http.StatusForbidden {
t.Fatalf("%s status = %d, want 403", name, response.Code)
}
}
if csrf == "" {
t.Fatal("authenticated page did not contain a CSRF token")
}
if store.startCalls != 0 {
t.Fatalf("unauthorized requests called store %d times", store.startCalls)
}
}
func TestStartPurchasesRejectsInvalidUTF8BeforeJSONDecoding(t *testing.T) {
validPrefix := []byte(`{"start_key":"` + startKeyForHTTP + `","tasks":[],"start_key":"`)
duplicateKeyBypass := append(append([]byte(nil), validPrefix...), 0xff)
duplicateKeyBypass = append(duplicateKeyBypass, []byte(`"}`)...)
invalidWhitespace := append([]byte(validStartBody()), 0xfe)
for name, body := range map[string][]byte{
"invalid byte after JSON": invalidWhitespace,
"invalid duplicate-key value": duplicateKeyBypass,
} {
t.Run(name, func(t *testing.T) {
store := &startRecordingStore{}
router, _ := newRouterWithStore(t, store)
cookie, csrf := authenticatedStartSession(t, router)
response := serveStartBytes(t, router, body, "application/json", csrf, cookie)
if response.Code != http.StatusBadRequest || store.startCalls != 0 {
t.Fatalf("status/calls = %d/%d, want 400/0", response.Code, store.startCalls)
}
if response.Body.Len() != 0 {
t.Fatalf("invalid UTF-8 response leaked body %q", response.Body.String())
}
})
}
}
func TestStartPurchasesEnforcesExact64KiBBodyBoundary(t *testing.T) {
const limit = 64 << 10
base := validStartBody()
for name, test := range map[string]struct {
body string
want int
wantCalls int
}{
"exact limit": {body: base + strings.Repeat(" ", limit-len(base)), want: http.StatusOK, wantCalls: 1},
"one over": {body: base + strings.Repeat(" ", limit-len(base)+1), want: http.StatusRequestEntityTooLarge},
} {
t.Run(name, func(t *testing.T) {
store := &startRecordingStore{startResult: successfulStartResult()}
router, _ := newRouterWithStore(t, store)
cookie, csrf := authenticatedStartSession(t, router)
response := serveStartRequest(t, router, test.body, "application/json", csrf, cookie)
if response.Code != test.want || store.startCalls != test.wantCalls {
t.Fatalf("status/calls = %d/%d, want %d/%d", response.Code, store.startCalls, test.want, test.wantCalls)
}
})
}
}
func TestStartPurchasesContentTypeContract(t *testing.T) {
for _, contentType := range []string{
"application/json",
"application/json; charset=utf-8",
"application/json;charset=UTF-8",
} {
t.Run("accept "+contentType, func(t *testing.T) {
store := &startRecordingStore{startResult: successfulStartResult()}
router, _ := newRouterWithStore(t, store)
cookie, csrf := authenticatedStartSession(t, router)
response := serveStartRequest(t, router, validStartBody(), contentType, csrf, cookie)
if response.Code != http.StatusOK || store.startCalls != 1 {
t.Fatalf("status/calls = %d/%d, want 200/1", response.Code, store.startCalls)
}
})
}
for _, contentType := range []string{
"",
"text/plain",
"application/json-patch+json",
"application/json; charset=gbk",
"application/json; profile=unapproved",
"application/json; charset",
} {
t.Run("reject "+contentType, func(t *testing.T) {
store := &startRecordingStore{}
router, _ := newRouterWithStore(t, store)
cookie, csrf := authenticatedStartSession(t, router)
response := serveStartRequest(t, router, validStartBody(), contentType, csrf, cookie)
if response.Code != http.StatusUnsupportedMediaType || store.startCalls != 0 {
t.Fatalf("status/calls = %d/%d, want 415/0", response.Code, store.startCalls)
}
})
}
}
func TestStartPurchasesRejectsMalformedAndOversizedJSON(t *testing.T) {
tests := []struct {
name string
body string
want int
storeErr error
wantCalls int
}{
{name: "empty", body: "", want: http.StatusBadRequest},
{name: "empty object", body: `{}`, want: http.StatusBadRequest, storeErr: tasks.ErrInvalidStart, wantCalls: 1},
{name: "null object", body: `null`, want: http.StatusBadRequest, storeErr: tasks.ErrInvalidStart, wantCalls: 1},
{name: "malformed", body: `{`, want: http.StatusBadRequest},
{name: "wrong top-level type", body: `[]`, want: http.StatusBadRequest},
{name: "unknown field", body: `{"start_key":"` + startKeyForHTTP + `","tasks":[],"created_by":"attacker"}`, want: http.StatusBadRequest},
{name: "wrong field type", body: `{"start_key":"` + startKeyForHTTP + `","tasks":[{"task_id":"` + taskIDForHTTP + `","expected_task_version":"1"}]}`, want: http.StatusBadRequest},
{name: "second JSON value", body: validStartBody() + `{}`, want: http.StatusBadRequest},
{name: "duplicate task ids", body: `{"start_key":"` + startKeyForHTTP + `","tasks":[{"task_id":"` + taskIDForHTTP + `","expected_task_version":1},{"task_id":"` + taskIDForHTTP + `","expected_task_version":1}]}`, want: http.StatusBadRequest, storeErr: tasks.ErrInvalidStart, wantCalls: 1},
{name: "oversized first value", body: `{"start_key":"` + strings.Repeat("x", 70<<10), want: http.StatusRequestEntityTooLarge},
{name: "oversized trailing whitespace", body: validStartBody() + strings.Repeat(" ", 70<<10), want: http.StatusRequestEntityTooLarge},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
store := &startRecordingStore{startErr: test.storeErr}
router, _ := newRouterWithStore(t, store)
cookie, csrf := authenticatedStartSession(t, router)
response := serveStartRequest(t, router, test.body, "application/json", csrf, cookie)
if response.Code != test.want || store.startCalls != test.wantCalls {
t.Fatalf("status/calls = %d/%d, want %d/%d", response.Code, store.startCalls, test.want, test.wantCalls)
}
if response.Body.Len() != 0 {
t.Fatalf("error response leaked body %q", response.Body.String())
}
assertSecurityHeaders(t, response)
})
}
}
func TestStartPurchasesUsesAuthenticatedAdminAndReturnsStableSafeResult(t *testing.T) {
result := successfulStartResult()
store := &startRecordingStore{startResult: result}
router, _ := newRouterWithStore(t, store)
cookie, csrf := authenticatedStartSession(t, router)
first := serveStartRequest(t, router, validStartBody(), "application/json; charset=utf-8", csrf, cookie)
second := serveStartRequest(t, router, validStartBody(), "application/json", csrf, cookie)
for index, response := range []*httptest.ResponseRecorder{first, second} {
if response.Code != http.StatusOK {
t.Fatalf("response %d status = %d, want 200", index, response.Code)
}
if got := response.Header().Get("Content-Type"); got != "application/json; charset=utf-8" {
t.Fatalf("response content type = %q", got)
}
var decoded tasks.StartResult
if err := json.Unmarshal(response.Body.Bytes(), &decoded); err != nil {
t.Fatalf("decode response: %v", err)
}
if decoded.PaymentAutomated || decoded.AuthorizedCount != 1 || decoded.Tasks[0].AuthorizationID != result.Tasks[0].AuthorizationID {
t.Fatalf("unsafe or unstable response = %#v", decoded)
}
assertSecurityHeaders(t, response)
}
if store.startCalls != 2 || len(store.createdBy) != 2 || store.createdBy[0] != "admin" || store.createdBy[1] != "admin" {
t.Fatalf("store calls/created_by = %d/%#v", store.startCalls, store.createdBy)
}
for _, command := range store.commands {
if command.StartKey != startKeyForHTTP || len(command.Tasks) != 1 || command.Tasks[0].TaskID != taskIDForHTTP || command.Tasks[0].ExpectedTaskVersion != 7 {
t.Fatalf("decoded command = %#v", command)
}
}
}
func TestStartPurchasesMapsStoreErrorsWithoutLeakingDetails(t *testing.T) {
for name, test := range map[string]struct {
err error
want int
}{
"invalid": {err: tasks.ErrInvalidStart, want: http.StatusBadRequest},
"conflict": {err: tasks.ErrStartConflict, want: http.StatusConflict},
"internal": {err: errors.New("sqlite secret path and query"), want: http.StatusInternalServerError},
} {
t.Run(name, func(t *testing.T) {
store := &startRecordingStore{startErr: test.err}
router, _ := newRouterWithStore(t, store)
cookie, csrf := authenticatedStartSession(t, router)
response := serveStartRequest(t, router, validStartBody(), "application/json", csrf, cookie)
if response.Code != test.want || store.startCalls != 1 {
t.Fatalf("status/calls = %d/%d, want %d/1", response.Code, store.startCalls, test.want)
}
if response.Body.Len() != 0 || strings.Contains(response.Body.String(), "sqlite") {
t.Fatalf("error leaked details: %q", response.Body.String())
}
})
}
}
type startRecordingStore struct {
startResult tasks.StartResult
startErr error
startCalls int
commands []tasks.StartCommand
createdBy []string
}
func (store *startRecordingStore) CreateDraft(_ context.Context, draft tasks.Draft) (tasks.Draft, error) {
return draft, nil
}
func (store *startRecordingStore) ListDrafts(context.Context) ([]tasks.Draft, error) {
return nil, nil
}
func (store *startRecordingStore) ListTasks(context.Context, tasks.TaskFilter) ([]tasks.TaskRow, error) {
return nil, nil
}
func (store *startRecordingStore) StartPurchases(_ context.Context, command tasks.StartCommand, createdBy string) (tasks.StartResult, error) {
store.startCalls++
store.commands = append(store.commands, command)
store.createdBy = append(store.createdBy, createdBy)
return store.startResult, store.startErr
}
func authenticatedStartSession(t *testing.T, router http.Handler) (*http.Cookie, string) {
t.Helper()
cookie := authenticate(t, router)
page := serve(router, http.MethodGet, "/tasks", nil, cookie)
if page.Code != http.StatusOK {
t.Fatalf("GET /tasks status = %d", page.Code)
}
return cookie, csrfToken(t, page.Body.String())
}
func newStartRequest(t *testing.T, body, contentType, csrf string, cookie *http.Cookie) *http.Request {
t.Helper()
return newStartByteRequest(t, []byte(body), contentType, csrf, cookie)
}
func newStartByteRequest(t *testing.T, body []byte, contentType, csrf string, cookie *http.Cookie) *http.Request {
t.Helper()
request := httptest.NewRequest(http.MethodPost, "/tasks/start-purchases", bytes.NewReader(body))
if contentType != "" {
request.Header.Set("Content-Type", contentType)
}
if csrf != "" {
request.Header.Set("X-CSRF-Token", csrf)
}
if cookie != nil {
request.AddCookie(cookie)
}
return request
}
func serveStartBytes(t *testing.T, router http.Handler, body []byte, contentType, csrf string, cookie *http.Cookie) *httptest.ResponseRecorder {
t.Helper()
response := httptest.NewRecorder()
router.ServeHTTP(response, newStartByteRequest(t, body, contentType, csrf, cookie))
return response
}
func serveStartRequest(t *testing.T, router http.Handler, body, contentType, csrf string, cookie *http.Cookie) *httptest.ResponseRecorder {
t.Helper()
response := httptest.NewRecorder()
router.ServeHTTP(response, newStartRequest(t, body, contentType, csrf, cookie))
return response
}
func validStartBody() string {
return `{"start_key":"` + startKeyForHTTP + `","tasks":[{"task_id":"` + taskIDForHTTP + `","expected_task_version":7}]}`
}
func successfulStartResult() tasks.StartResult {
expires := time.Date(2026, 8, 4, 2, 3, 4, 0, time.UTC)
return tasks.StartResult{
StartKey: startKeyForHTTP,
AuthorizedCount: 1,
PaymentAutomated: false,
Tasks: []tasks.AuthorizedTask{{
TaskID: taskIDForHTTP,
TaskVersion: 8,
AuthorizationID: "d3c9f507-7473-4fa6-8d71-8786c34c6301",
ExpiresAt: expires,
}},
}
}
+176
View File
@@ -0,0 +1,176 @@
package server
import (
"bytes"
"encoding/json"
"errors"
"io"
"net/http"
"unicode/utf8"
"cmbuyer/admin/internal/deviceauth"
"cmbuyer/admin/internal/taskclaim"
"github.com/gin-gonic/gin"
)
const (
maxClaimJSONBytes = 4096
maxClaimResponseJSONBytes = 32 * 1024
)
func claimNext(options Options) gin.HandlerFunc {
return func(context *gin.Context) {
principal, ok := authenticateDevice(context, options)
if !ok {
return
}
var command taskclaim.ClaimCommand
if !decodeClaimJSON(context, &command) {
return
}
response, found, err := options.TaskClaims.ClaimNext(context.Request.Context(), principal.ID, command)
if err != nil {
writeTaskClaimError(context, err)
return
}
if !found {
context.Status(http.StatusNoContent)
return
}
if !taskclaim.ValidClaimResponse(response) {
context.Status(http.StatusServiceUnavailable)
return
}
encoded, err := json.Marshal(response)
if err != nil || len(encoded) > maxClaimResponseJSONBytes {
context.Status(http.StatusServiceUnavailable)
return
}
context.Data(http.StatusOK, "application/json; charset=utf-8", encoded)
}
}
func renewLease(options Options) gin.HandlerFunc {
return func(context *gin.Context) {
principal, ok := authenticateDevice(context, options)
if !ok {
return
}
var command taskclaim.RenewCommand
if !decodeClaimJSON(context, &command) {
return
}
command.TaskID = context.Param("id")
response, err := options.TaskClaims.Renew(context.Request.Context(), principal.ID, command)
if err != nil {
writeTaskClaimError(context, err)
return
}
context.JSON(http.StatusOK, response)
}
}
// Authentication precedes path interpretation, Content-Type parsing and every body read. This
// keeps rejected devices from using parsing differences as an oracle or making the server buffer data.
func authenticateDevice(context *gin.Context, options Options) (deviceauth.Principal, bool) {
principal, err := options.DeviceAuthenticator.Authenticate(context.Request)
if errors.Is(err, deviceauth.ErrUnauthenticated) {
context.Header("WWW-Authenticate", "Bearer")
context.Status(http.StatusUnauthorized)
return deviceauth.Principal{}, false
}
if err != nil || !deviceauth.ValidDeviceID(principal.ID) {
context.Status(http.StatusServiceUnavailable)
return deviceauth.Principal{}, false
}
return principal, true
}
func decodeClaimJSON(context *gin.Context, target any) bool {
if !isJSONContentType(context.GetHeader("Content-Type")) {
writeFixedError(context, http.StatusUnsupportedMediaType, "unsupported_media_type")
return false
}
context.Request.Body = http.MaxBytesReader(context.Writer, context.Request.Body, maxClaimJSONBytes)
raw, err := io.ReadAll(context.Request.Body)
if err != nil {
var tooLarge *http.MaxBytesError
if errors.As(err, &tooLarge) {
writeFixedError(context, http.StatusRequestEntityTooLarge, "request_too_large")
} else {
writeFixedError(context, http.StatusBadRequest, "invalid_request")
}
return false
}
if len(raw) == 0 || !utf8.Valid(raw) {
writeFixedError(context, http.StatusBadRequest, "invalid_request")
return false
}
if !hasUniqueTopLevelJSONFields(raw) {
writeFixedError(context, http.StatusBadRequest, "invalid_request")
return false
}
decoder := json.NewDecoder(bytes.NewReader(raw))
decoder.DisallowUnknownFields()
if err := decoder.Decode(target); err != nil {
writeFixedError(context, http.StatusBadRequest, "invalid_request")
return false
}
var extra any
if err := decoder.Decode(&extra); err != io.EOF {
writeFixedError(context, http.StatusBadRequest, "invalid_request")
return false
}
return true
}
func hasUniqueTopLevelJSONFields(raw []byte) bool {
decoder := json.NewDecoder(bytes.NewReader(raw))
first, err := decoder.Token()
if err != nil || first != json.Delim('{') {
return false
}
seen := make(map[string]struct{})
for decoder.More() {
key, err := decoder.Token()
name, ok := key.(string)
if err != nil || !ok {
return false
}
if _, duplicate := seen[name]; duplicate {
return false
}
seen[name] = struct{}{}
var value json.RawMessage
if err := decoder.Decode(&value); err != nil {
return false
}
}
last, err := decoder.Token()
return err == nil && last == json.Delim('}')
}
func writeTaskClaimError(context *gin.Context, err error) {
switch {
case errors.Is(err, taskclaim.ErrInvalid):
writeFixedError(context, http.StatusBadRequest, "invalid_request")
case errors.Is(err, taskclaim.ErrIdempotencyConflict):
writeFixedError(context, http.StatusConflict, "idempotency_conflict")
case errors.Is(err, taskclaim.ErrRequiresManual):
writeFixedError(context, http.StatusConflict, "claim_requires_manual")
case errors.Is(err, taskclaim.ErrNotCurrent):
writeFixedError(context, http.StatusConflict, "claim_not_current")
case errors.Is(err, taskclaim.ErrDeviceInactive):
context.Header("WWW-Authenticate", "Bearer")
context.Status(http.StatusUnauthorized)
default:
// Storage and transaction failures are intentionally bodyless: SQL, paths and candidate
// details are server-only and must not become a device-facing diagnostic oracle.
context.Status(http.StatusServiceUnavailable)
}
}
func writeFixedError(context *gin.Context, status int, code string) {
context.JSON(status, gin.H{"error": code})
}
+215
View File
@@ -0,0 +1,215 @@
package server_test
import (
"bytes"
"context"
"encoding/json"
"errors"
"io"
"math"
"net/http"
"net/http/httptest"
"strings"
"testing"
"cmbuyer/admin/internal/deviceauth"
"cmbuyer/admin/internal/taskclaim"
)
const (
claimDeviceID = "10000000-0000-4000-8000-000000000001"
claimSessionID = "20000000-0000-4000-8000-000000000001"
claimRequestID = "30000000-0000-4000-8000-000000000001"
claimTaskID = "40000000-0000-4000-8000-000000000001"
claimAttemptID = "50000000-0000-4000-8000-000000000001"
claimRenewID = "60000000-0000-4000-8000-000000000001"
)
func TestTaskClaimEndpointsAuthenticateBeforeBody(t *testing.T) {
for _, authentication := range []struct {
name string
err error
status int
}{
{"unauthenticated", deviceauth.ErrUnauthenticated, http.StatusUnauthorized},
{"authentication storage unavailable", deviceauth.ErrUnavailable, http.StatusServiceUnavailable},
} {
t.Run(authentication.name, func(t *testing.T) {
authenticator := &fakeDeviceAuthenticator{err: authentication.err}
service := &fakeTaskClaimService{}
router, _ := newRouterWithClaimService(t, &memoryStore{}, emptyDetailStore{}, emptyEvidenceStore{}, authenticator, service)
for _, path := range []string{"/api/v1/tasks/claim-next", "/api/v1/tasks/" + claimTaskID + "/lease/renew"} {
body := &poisonBody{}
request := httptest.NewRequest(http.MethodPost, path, nil)
request.Body = body
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code != authentication.status || response.Body.Len() != 0 || body.reads != 0 || service.calls != 0 {
t.Fatalf("%s = status %d, body %q, reads %d, calls %d", path, response.Code, response.Body.String(), body.reads, service.calls)
}
}
})
}
}
func TestClaimNextStrictJSONSuccessEmptyAndErrors(t *testing.T) {
authenticator := &fakeDeviceAuthenticator{principal: deviceauth.Principal{ID: claimDeviceID}}
service := &fakeTaskClaimService{claimResponse: taskclaim.ClaimResponse{
Task: taskclaim.ClaimedTask{ID: claimTaskID, Version: 3, Title: "测试", ProductURL: "https://mobile.yangkeduo.com/goods.html?goods_id=1", GoodsID: "1", SKUColor: "黑色", SKUSize: "M", Quantity: 1, MaxTotalPrice: "1.00"},
Authorization: taskclaim.ClaimedAuthorization{ID: "70000000-0000-4000-8000-000000000001", TaskVersion: 2, ExpiresAt: "2026-08-04T01:10:00Z"},
Attempt: taskclaim.ClaimedAttempt{ID: claimAttemptID, ClaimToken: strings.Repeat("a", 64), ClaimGeneration: 1, LeaseExpiresAt: "2026-08-04T01:03:00Z"},
}, claimFound: true}
router, _ := newRouterWithClaimService(t, &memoryStore{}, emptyDetailStore{}, emptyEvidenceStore{}, authenticator, service)
valid := `{"session_id":"` + claimSessionID + `","claim_request_id":"` + claimRequestID + `"}`
response := serveClaimJSON(router, "/api/v1/tasks/claim-next", valid, "application/json; charset=utf-8")
if response.Code != http.StatusOK || !strings.Contains(response.Body.String(), strings.Repeat("a", 64)) || service.claimCommand.ClaimRequestID != claimRequestID {
t.Fatalf("claim success = %d %q command %#v", response.Code, response.Body.String(), service.claimCommand)
}
service.claimFound = false
response = serveClaimJSON(router, "/api/v1/tasks/claim-next", valid, "application/json")
if response.Code != http.StatusNoContent || response.Body.Len() != 0 {
t.Fatalf("claim empty = %d %q", response.Code, response.Body.String())
}
tests := []struct {
name, body, contentType, code string
status int
}{
{"unsupported type", valid, "text/plain", "unsupported_media_type", http.StatusUnsupportedMediaType},
{"unknown field", strings.TrimSuffix(valid, "}") + `,"device_id":"` + claimDeviceID + `"}`, "application/json", "invalid_request", http.StatusBadRequest},
{"duplicate session", `{"session_id":"` + claimSessionID + `","session_id":"` + claimSessionID + `","claim_request_id":"` + claimRequestID + `"}`, "application/json", "invalid_request", http.StatusBadRequest},
{"duplicate request", `{"session_id":"` + claimSessionID + `","claim_request_id":"` + claimRequestID + `","claim_request_id":"` + claimRequestID + `"}`, "application/json", "invalid_request", http.StatusBadRequest},
{"extra json", valid + `{}`, "application/json", "invalid_request", http.StatusBadRequest},
{"too large", strings.Repeat(" ", 4097), "application/json", "request_too_large", http.StatusRequestEntityTooLarge},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
response := serveClaimJSON(router, "/api/v1/tasks/claim-next", test.body, test.contentType)
if response.Code != test.status || response.Body.String() != `{"error":"`+test.code+`"}` {
t.Fatalf("response = %d %q", response.Code, response.Body.String())
}
})
}
invalidUTF8 := httptest.NewRequest(http.MethodPost, "/api/v1/tasks/claim-next", bytes.NewReader([]byte{'{', 0xff, '}'}))
invalidUTF8.Header.Set("Content-Type", "application/json")
invalidResponse := httptest.NewRecorder()
router.ServeHTTP(invalidResponse, invalidUTF8)
if invalidResponse.Code != http.StatusBadRequest || invalidResponse.Body.String() != `{"error":"invalid_request"}` {
t.Fatalf("invalid UTF-8 = %d %q", invalidResponse.Code, invalidResponse.Body.String())
}
}
func TestClaimResponseWorstLegalFieldsStayBelowCapAndInvalidServiceOutputFailsClosed(t *testing.T) {
goodsID := strings.Repeat("1", 32)
worst := taskclaim.ClaimResponse{
Task: taskclaim.ClaimedTask{
ID: claimTaskID, Version: math.MaxInt, Title: strings.Repeat("<", 120),
ProductURL: "https://mobile.yangkeduo.com/goods.html?goods_id=" + goodsID,
GoodsID: goodsID, SKUColor: strings.Repeat("<", 80), SKUSize: strings.Repeat("<", 80),
Quantity: 9_223_372_036_854_775_807, MaxTotalPrice: strings.Repeat("9", 29) + ".00",
},
Authorization: taskclaim.ClaimedAuthorization{ID: "70000000-0000-4000-8000-000000000001", TaskVersion: math.MaxInt - 1, ExpiresAt: "9999-12-31T23:59:59.999999999Z"},
Attempt: taskclaim.ClaimedAttempt{ID: claimAttemptID, ClaimToken: strings.Repeat("a", 64), ClaimGeneration: 9_223_372_036_854_775_807, LeaseExpiresAt: "9999-12-31T23:59:59.999999999Z"},
}
authenticator := &fakeDeviceAuthenticator{principal: deviceauth.Principal{ID: claimDeviceID}}
service := &fakeTaskClaimService{claimResponse: worst, claimFound: true}
router, _ := newRouterWithClaimService(t, &memoryStore{}, emptyDetailStore{}, emptyEvidenceStore{}, authenticator, service)
request := `{"session_id":"` + claimSessionID + `","claim_request_id":"` + claimRequestID + `"}`
response := serveClaimJSON(router, "/api/v1/tasks/claim-next", request, "application/json")
if response.Code != http.StatusOK || !json.Valid(response.Body.Bytes()) || response.Body.Len() >= 32*1024 {
t.Fatalf("worst legal response = status %d, bytes %d, valid JSON %v", response.Code, response.Body.Len(), json.Valid(response.Body.Bytes()))
}
mutations := map[string]func(*taskclaim.ClaimResponse){
"invalid utf8 title": func(response *taskclaim.ClaimResponse) { response.Task.Title = string([]byte{0xff}) },
"c0 separator title": func(response *taskclaim.ClaimResponse) { response.Task.Title = "visible\u001dhidden" },
"overlong title": func(response *taskclaim.ClaimResponse) { response.Task.Title += "<" },
"overlong goods id": func(response *taskclaim.ClaimResponse) {
response.Task.GoodsID += "1"
response.Task.ProductURL += "1"
},
"overlong color": func(response *taskclaim.ClaimResponse) { response.Task.SKUColor += "<" },
"overlong size": func(response *taskclaim.ClaimResponse) { response.Task.SKUSize += "<" },
"overlong money": func(response *taskclaim.ClaimResponse) { response.Task.MaxTotalPrice = strings.Repeat("9", 30) + ".00" },
}
for name, mutate := range mutations {
t.Run(name, func(t *testing.T) {
invalid := worst
mutate(&invalid)
service := &fakeTaskClaimService{claimResponse: invalid, claimFound: true}
router, _ := newRouterWithClaimService(t, &memoryStore{}, emptyDetailStore{}, emptyEvidenceStore{}, authenticator, service)
response := serveClaimJSON(router, "/api/v1/tasks/claim-next", request, "application/json")
if response.Code != http.StatusServiceUnavailable || response.Body.Len() != 0 {
t.Fatalf("invalid service response = %d %q", response.Code, response.Body.String())
}
})
}
}
func TestRenewStrictBindingResponseAndFixedErrors(t *testing.T) {
authenticator := &fakeDeviceAuthenticator{principal: deviceauth.Principal{ID: claimDeviceID}}
service := &fakeTaskClaimService{renewResponse: taskclaim.RenewResponse{TaskID: claimTaskID, AttemptID: claimAttemptID, ClaimGeneration: 1, LeaseExpiresAt: "2026-08-04T01:04:00Z"}}
router, _ := newRouterWithClaimService(t, &memoryStore{}, emptyDetailStore{}, emptyEvidenceStore{}, authenticator, service)
body := `{"renew_request_id":"` + claimRenewID + `","session_id":"` + claimSessionID + `","attempt_id":"` + claimAttemptID + `","claim_generation":1,"claim_token":"` + strings.Repeat("a", 64) + `","expected_lease_expires_at":"2026-08-04T01:03:00Z"}`
response := serveClaimJSON(router, "/api/v1/tasks/"+claimTaskID+"/lease/renew", body, "application/json")
if response.Code != http.StatusOK || strings.Contains(response.Body.String(), "claim_token") || service.renewCommand.TaskID != claimTaskID {
t.Fatalf("renew response = %d %q command %#v", response.Code, response.Body.String(), service.renewCommand)
}
duplicateToken := strings.Replace(body, `"expected_lease_expires_at"`, `"claim_token":"`+strings.Repeat("a", 64)+`","expected_lease_expires_at"`, 1)
response = serveClaimJSON(router, "/api/v1/tasks/"+claimTaskID+"/lease/renew", duplicateToken, "application/json")
if response.Code != http.StatusBadRequest {
t.Fatalf("duplicate token status = %d", response.Code)
}
errorsToCodes := []struct {
err error
status int
body string
}{
{taskclaim.ErrIdempotencyConflict, http.StatusConflict, `{"error":"idempotency_conflict"}`},
{taskclaim.ErrRequiresManual, http.StatusConflict, `{"error":"claim_requires_manual"}`},
{taskclaim.ErrNotCurrent, http.StatusConflict, `{"error":"claim_not_current"}`},
{taskclaim.ErrDeviceInactive, http.StatusUnauthorized, ""},
{errors.New("database path and SQL must stay private"), http.StatusServiceUnavailable, ""},
}
for _, test := range errorsToCodes {
service.renewErr = test.err
response = serveClaimJSON(router, "/api/v1/tasks/"+claimTaskID+"/lease/renew", body, "application/json")
if response.Code != test.status || response.Body.String() != test.body {
t.Fatalf("error %v = %d %q", test.err, response.Code, response.Body.String())
}
}
}
type fakeTaskClaimService struct {
claimResponse taskclaim.ClaimResponse
claimFound bool
claimErr error
renewResponse taskclaim.RenewResponse
renewErr error
claimCommand taskclaim.ClaimCommand
renewCommand taskclaim.RenewCommand
calls int
}
func (service *fakeTaskClaimService) ClaimNext(_ context.Context, _ string, command taskclaim.ClaimCommand) (taskclaim.ClaimResponse, bool, error) {
service.calls++
service.claimCommand = command
return service.claimResponse, service.claimFound, service.claimErr
}
func (service *fakeTaskClaimService) Renew(_ context.Context, _ string, command taskclaim.RenewCommand) (taskclaim.RenewResponse, error) {
service.calls++
service.renewCommand = command
return service.renewResponse, service.renewErr
}
func serveClaimJSON(router http.Handler, path, body, contentType string) *httptest.ResponseRecorder {
request := httptest.NewRequest(http.MethodPost, path, io.NopCloser(strings.NewReader(body)))
request.Header.Set("Content-Type", contentType)
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
return response
}
+84
View File
@@ -0,0 +1,84 @@
package server
import (
"errors"
"mime"
"net/http"
"net/url"
"strconv"
"strings"
"cmbuyer/admin/internal/taskdetail"
"cmbuyer/admin/internal/transport/webui"
"github.com/gin-gonic/gin"
)
const detailViewHeader = "X-CMBuyer-View"
const detailVaryHeader = "X-CMBuyer-View, Accept, Sec-Fetch-Site"
func taskDetailPage(options Options) gin.HandlerFunc {
return func(context *gin.Context) {
context.Header("Vary", detailVaryHeader)
if !options.Sessions.IsAuthenticated(context.Request) {
context.Redirect(http.StatusSeeOther, "/login?return_to="+url.QueryEscape(context.Request.URL.RequestURI()))
return
}
view := context.GetHeader(detailViewHeader)
if view != "" && view != "drawer" {
context.Status(http.StatusBadRequest)
return
}
if view == "drawer" {
if context.GetHeader("Sec-Fetch-Site") != "same-origin" {
context.Status(http.StatusForbidden)
return
}
if !acceptsHTML(context.GetHeader("Accept")) {
context.Status(http.StatusNotAcceptable)
return
}
}
detail, err := options.TaskDetails.Get(context.Request.Context(), context.Param("id"))
if errors.Is(err, taskdetail.ErrNotFound) {
context.Status(http.StatusNotFound)
return
}
if err != nil {
context.Status(http.StatusInternalServerError)
return
}
context.Header("Content-Type", "text/html; charset=utf-8")
context.Status(http.StatusOK)
data := webui.TaskDetailData{Detail: detail}
if view == "drawer" {
if err := webui.RenderTaskDetailFragment(context.Writer, data); err != nil {
_ = context.Error(err)
}
return
}
if err := webui.RenderTaskDetailPage(context.Writer, data); err != nil {
_ = context.Error(err)
}
}
}
func acceptsHTML(header string) bool {
for _, value := range strings.Split(header, ",") {
mediaType, parameters, err := mime.ParseMediaType(strings.TrimSpace(value))
if err != nil || !strings.EqualFold(mediaType, "text/html") {
continue
}
quality := 1.0
if rawQuality, exists := parameters["q"]; exists {
quality, err = strconv.ParseFloat(rawQuality, 64)
if err != nil || quality < 0 || quality > 1 {
continue
}
}
if quality > 0 {
return true
}
}
return false
}
+110
View File
@@ -0,0 +1,110 @@
package server_test
import (
"context"
"net/http"
"net/http/httptest"
"strings"
"testing"
"time"
"cmbuyer/admin/internal/deviceauth"
"cmbuyer/admin/internal/taskdetail"
)
const detailTaskID = "a3c9f507-7473-4fa6-8d71-8786c34c6301"
func TestTaskDetailRequiresAdminBeforeLookup(t *testing.T) {
details := &recordingDetailStore{detail: taskDetailFixture()}
router, _ := newRouterWithDependencies(t, &memoryStore{}, details, emptyEvidenceStore{}, deviceauth.RejectAllAuthenticator{})
response := serve(router, http.MethodGet, "/tasks/"+detailTaskID, nil, nil)
if response.Code != http.StatusSeeOther || !strings.HasPrefix(response.Header().Get("Location"), "/login?return_to=") || details.calls != 0 {
t.Fatalf("anonymous detail = %d/%q, calls=%d", response.Code, response.Header().Get("Location"), details.calls)
}
}
func TestTaskDetailFullPageAndDrawerShareAuditContent(t *testing.T) {
details := &recordingDetailStore{detail: taskDetailFixture()}
router, _ := newRouterWithDependencies(t, &memoryStore{}, details, emptyEvidenceStore{}, deviceauth.RejectAllAuthenticator{})
cookie := authenticate(t, router)
full := serve(router, http.MethodGet, "/tasks/"+detailTaskID, nil, cookie)
if full.Code != http.StatusOK || !strings.Contains(full.Body.String(), "<!doctype html>") || !strings.Contains(full.Body.String(), `data-task-detail-content`) {
t.Fatalf("full detail = %d/%q", full.Code, full.Body.String())
}
request := httptest.NewRequest(http.MethodGet, "/tasks/"+detailTaskID, nil)
request.AddCookie(cookie)
request.Header.Set("X-CMBuyer-View", "drawer")
request.Header.Set("Accept", "text/html")
request.Header.Set("Sec-Fetch-Site", "same-origin")
fragment := httptest.NewRecorder()
router.ServeHTTP(fragment, request)
if fragment.Code != http.StatusOK || strings.Contains(fragment.Body.String(), "<!doctype html>") || !strings.Contains(fragment.Body.String(), `data-task-detail-content`) {
t.Fatalf("fragment detail = %d/%q", fragment.Code, fragment.Body.String())
}
for _, text := range []string{"测试&lt;script&gt;", "订单已创建,系统尚未付款", "SKU_PANEL_GATE_1", "/evidence/b3c9f507-7473-4fa6-8d71-8786c34c6301", "暂无规格、价格或数量读数", "本页没有重试、再次提交或付款动作"} {
if !strings.Contains(full.Body.String(), text) || !strings.Contains(fragment.Body.String(), text) {
t.Fatalf("shared detail missing %q", text)
}
}
if strings.Contains(full.Body.String(), "<script>") || strings.Contains(fragment.Body.String(), "<script>") {
t.Fatal("task title was not HTML escaped")
}
if got := fragment.Header().Get("Vary"); got != "X-CMBuyer-View, Accept, Sec-Fetch-Site" {
t.Fatalf("fragment Vary = %q", got)
}
if details.calls != 2 {
t.Fatalf("detail store calls = %d, want 2", details.calls)
}
}
func TestTaskDetailRejectsForgedFragmentAndMissingTask(t *testing.T) {
details := &recordingDetailStore{err: taskdetail.ErrNotFound}
router, _ := newRouterWithDependencies(t, &memoryStore{}, details, emptyEvidenceStore{}, deviceauth.RejectAllAuthenticator{})
cookie := authenticate(t, router)
for name, headers := range map[string]map[string]string{
"unknown view": {"X-CMBuyer-View": "xml", "Accept": "text/html"},
"missing fetch site": {"X-CMBuyer-View": "drawer", "Accept": "text/html"},
"cross-site drawer": {"X-CMBuyer-View": "drawer", "Accept": "text/html", "Sec-Fetch-Site": "cross-site"},
"wrong accept": {"X-CMBuyer-View": "drawer", "Accept": "application/json", "Sec-Fetch-Site": "same-origin"},
"html quality zero": {"X-CMBuyer-View": "drawer", "Accept": "text/html;q=0, application/json", "Sec-Fetch-Site": "same-origin"},
"html substring mime": {"X-CMBuyer-View": "drawer", "Accept": "application/nottext/html", "Sec-Fetch-Site": "same-origin"},
} {
t.Run(name, func(t *testing.T) {
request := httptest.NewRequest(http.MethodGet, "/tasks/"+detailTaskID, nil)
request.AddCookie(cookie)
for key, value := range headers {
request.Header.Set(key, value)
}
response := httptest.NewRecorder()
router.ServeHTTP(response, request)
if response.Code < 400 || response.Code >= 500 || response.Body.Len() != 0 {
t.Fatalf("forged fragment = %d/%q", response.Code, response.Body.String())
}
})
}
missing := serve(router, http.MethodGet, "/tasks/not-a-uuid", nil, cookie)
if missing.Code != http.StatusNotFound || missing.Body.Len() != 0 {
t.Fatalf("missing detail = %d/%q", missing.Code, missing.Body.String())
}
}
type recordingDetailStore struct {
detail taskdetail.Detail
err error
calls int
}
func (store *recordingDetailStore) Get(context.Context, string) (taskdetail.Detail, error) {
store.calls++
return store.detail, store.err
}
func taskDetailFixture() taskdetail.Detail {
started := time.Date(2026, 8, 4, 1, 2, 3, 0, time.UTC)
return taskdetail.Detail{
Task: taskdetail.Task{ID: detailTaskID, Source: "MANUAL", Title: "测试<script>", GoodsID: "937122477375", SKUColor: "黑色", SKUSize: "M", Quantity: 2, MaxTotalPrice: "30.00", Status: "WAITING_PAYMENT", Version: 3, CreatedAt: started, UpdatedAt: started},
Authorizations: []taskdetail.Authorization{{ID: "c3c9f507-7473-4fa6-8d71-8786c34c6301", Status: "FENCED", CreatedBy: "admin", TotalPriceCap: "30.00", TaskVersion: 2, CreatedAt: started, ExpiresAt: started.Add(time.Hour)}},
Attempts: []taskdetail.Attempt{{ID: "d3c9f507-7473-4fa6-8d71-8786c34c6301", AuthorizationID: "c3c9f507-7473-4fa6-8d71-8786c34c6301", Status: "CLAIMED", ClaimGeneration: 1, StartedAt: started}},
Evidence: []taskdetail.Evidence{{ID: "b3c9f507-7473-4fa6-8d71-8786c34c6301", AttemptID: "d3c9f507-7473-4fa6-8d71-8786c34c6301", Kind: "SKU_PANEL_GATE_1", PrivacyTier: "INTERNAL_RAW", SHA256: strings.Repeat("a", 64), ByteSize: 100, ContentType: "image/png", Width: 100, Height: 200, CapturedAt: started}},
}
}
@@ -0,0 +1,20 @@
//go:build !windows
package evidence
import (
"fmt"
"os"
)
func syncDirectory(path string) error {
directory, err := os.Open(path)
if err != nil {
return fmt.Errorf("open directory for durability sync: %w", err)
}
defer directory.Close()
if err := directory.Sync(); err != nil {
return fmt.Errorf("sync directory metadata: %w", err)
}
return nil
}
@@ -0,0 +1,35 @@
//go:build windows
package evidence
import (
"fmt"
"syscall"
)
// syncDirectory uses an explicit directory handle because os.Open(...).Sync is not a portable
// Windows directory durability boundary. Any unsupported filesystem or access failure is fatal:
// callers must not make the corresponding evidence row visible in SQLite.
func syncDirectory(path string) error {
pathPointer, err := syscall.UTF16PtrFromString(path)
if err != nil {
return fmt.Errorf("encode directory path for durability sync: %w", err)
}
handle, err := syscall.CreateFile(
pathPointer,
syscall.GENERIC_WRITE,
syscall.FILE_SHARE_READ|syscall.FILE_SHARE_WRITE|syscall.FILE_SHARE_DELETE,
nil,
syscall.OPEN_EXISTING,
syscall.FILE_FLAG_BACKUP_SEMANTICS,
0,
)
if err != nil {
return fmt.Errorf("open directory for durability sync: %w", err)
}
defer syscall.CloseHandle(handle)
if err := syscall.FlushFileBuffers(handle); err != nil {
return fmt.Errorf("flush directory metadata: %w", err)
}
return nil
}
+550
View File
@@ -0,0 +1,550 @@
// Package evidence stores INTERNAL_RAW PNG assets outside the public web tree.
package evidence
import (
"context"
"crypto/rand"
"crypto/sha256"
"database/sql"
"encoding/hex"
"errors"
"fmt"
"image/png"
"io"
"os"
"path/filepath"
"strings"
"sync"
"time"
"cmbuyer/admin/internal/deviceauth"
core "cmbuyer/admin/internal/evidence"
)
var pngSignature = []byte{0x89, 'P', 'N', 'G', 0x0d, 0x0a, 0x1a, 0x0a}
type Store struct {
database *sql.DB
root string
now func() time.Time
random io.Reader
syncDirectory func(string) error
syncFile func(*os.File) error
renameFile func(string, string) error
commitTx func(*sql.Tx) error
mu sync.Mutex
}
func NewStore(database *sql.DB, root string) (*Store, error) {
return newStore(database, root, syncDirectory)
}
func newStore(database *sql.DB, root string, directorySync func(string) error) (*Store, error) {
if database == nil {
return nil, errors.New("evidence database is required")
}
if directorySync == nil {
return nil, errors.New("evidence directory sync is required")
}
if root == "" || !filepath.IsAbs(root) {
return nil, errors.New("evidence root must be an absolute path")
}
absolute, err := filepath.Abs(filepath.Clean(root))
if err != nil {
return nil, fmt.Errorf("resolve evidence root: %w", err)
}
if filepath.Dir(absolute) == absolute {
return nil, errors.New("evidence root cannot be a filesystem root")
}
if err := ensureDurableDirectory(absolute, 0o700, directorySync); err != nil {
return nil, fmt.Errorf("create evidence root: %w", err)
}
// A prior startup may have created the root and then failed its parent sync.
// Existence is therefore never accepted as proof that the directory entry is durable.
if err := directorySync(filepath.Dir(absolute)); err != nil {
return nil, fmt.Errorf("persist evidence root directory: %w", err)
}
if err := os.Chmod(absolute, 0o700); err != nil {
return nil, fmt.Errorf("protect evidence root: %w", err)
}
staging := filepath.Join(absolute, ".staging")
if err := ensureDurableDirectory(staging, 0o700, directorySync); err != nil {
return nil, fmt.Errorf("create evidence staging directory: %w", err)
}
if err := directorySync(absolute); err != nil {
return nil, fmt.Errorf("persist evidence staging directory: %w", err)
}
if err := os.Chmod(staging, 0o700); err != nil {
return nil, fmt.Errorf("protect evidence staging directory: %w", err)
}
if _, err := database.Exec("SELECT storage_key FROM evidence_assets LIMIT 1"); err != nil {
return nil, fmt.Errorf("evidence migration is not available: %w", err)
}
return &Store{
database: database,
root: absolute,
now: time.Now,
random: rand.Reader,
syncDirectory: directorySync,
syncFile: func(file *os.File) error { return file.Sync() },
renameFile: os.Rename,
commitTx: func(transaction *sql.Tx) error { return transaction.Commit() },
}, nil
}
func (store *Store) Stage(reader io.Reader, contentType string) (staged core.StagedFile, resultErr error) {
if reader == nil || contentType != core.PNGContentType {
return core.StagedFile{}, core.ErrInvalid
}
temporary, err := os.CreateTemp(filepath.Join(store.root, ".staging"), "upload-*.png")
if err != nil {
return core.StagedFile{}, err
}
staged.Path = temporary.Name()
defer func() {
if resultErr != nil {
_ = temporary.Close()
_ = os.Remove(staged.Path)
}
}()
if err := temporary.Chmod(0o600); err != nil {
return core.StagedFile{}, err
}
hasher := sha256.New()
written, err := io.Copy(io.MultiWriter(temporary, hasher), io.LimitReader(reader, core.MaxFileBytes+1))
if err != nil {
return core.StagedFile{}, err
}
if written > core.MaxFileBytes {
return core.StagedFile{}, core.ErrTooLarge
}
if written == 0 {
return core.StagedFile{}, core.ErrInvalid
}
if err := temporary.Sync(); err != nil {
return core.StagedFile{}, err
}
if err := temporary.Close(); err != nil {
return core.StagedFile{}, err
}
imageFile, err := os.Open(staged.Path)
if err != nil {
return core.StagedFile{}, err
}
defer imageFile.Close()
width, height, err := validatePNG(imageFile)
if err != nil {
return core.StagedFile{}, err
}
staged.SHA256 = hex.EncodeToString(hasher.Sum(nil))
staged.ByteSize = written
staged.ContentType = core.PNGContentType
staged.Width = width
staged.Height = height
return staged, nil
}
func (store *Store) Discard(staged core.StagedFile) {
if store.isStagedPath(staged.Path) {
_ = os.Remove(staged.Path)
}
}
func (store *Store) Commit(ctx context.Context, principal deviceauth.Principal, metadata core.UploadMetadata, staged core.StagedFile) (core.Asset, bool, error) {
if !store.isStagedPath(staged.Path) || !validPrincipal(principal) || !validMetadata(metadata) || metadata.SHA256 != staged.SHA256 || staged.ContentType != core.PNGContentType || staged.ByteSize < 1 || staged.ByteSize > core.MaxFileBytes || staged.Width < 1 || staged.Height < 1 || staged.Width > core.MaxImageSide || staged.Height > core.MaxImageSide || int64(staged.Width)*int64(staged.Height) > core.MaxImagePixels {
store.Discard(staged)
return core.Asset{}, false, core.ErrInvalid
}
defer store.Discard(staged)
store.mu.Lock()
defer store.mu.Unlock()
transaction, err := store.database.BeginTx(ctx, nil)
if err != nil {
return core.Asset{}, false, err
}
defer transaction.Rollback()
existing, found, err := findByUploadKey(ctx, transaction, principal.ID, metadata.UploadKey)
if err != nil {
return core.Asset{}, false, err
}
if found {
if !sameUpload(existing, principal, metadata, staged) {
return core.Asset{}, false, core.ErrConflict
}
if err := store.verifyStoredFile(existing); err != nil {
return core.Asset{}, false, err
}
if err := store.commitTx(transaction); err != nil {
return core.Asset{}, false, err
}
return existing, true, nil
}
var ownedClaimCount int
if err := transaction.QueryRowContext(ctx, `SELECT COUNT(*) FROM purchase_attempt_claims
WHERE task_id = ? AND attempt_id = ? AND claimed_by_device_id = ? AND closed_at IS NULL`,
metadata.TaskID, metadata.AttemptID, principal.ID).Scan(&ownedClaimCount); err != nil {
return core.Asset{}, false, err
}
// Evidence is auditable only when the authenticated device owns the current attempt. The
// idempotent asset lookup above deliberately remains first so closing a claim later cannot
// destroy stable replay of an already committed screenshot.
if ownedClaimCount != 1 {
return core.Asset{}, false, core.ErrInvalid
}
storageKey := storageKey(metadata.SHA256)
finalPath, err := store.pathForKey(storageKey)
if err != nil {
return core.Asset{}, false, err
}
finalDirectory := filepath.Dir(finalPath)
if err := ensureDurableDirectory(finalDirectory, 0o700, store.syncDirectory); err != nil {
return core.Asset{}, false, err
}
// Always repeat the shard-parent boundary. If an earlier attempt created this
// directory and its parent sync failed, a retry must not trust mere existence.
if err := store.syncDirectory(store.root); err != nil {
return core.Asset{}, false, fmt.Errorf("persist evidence shard directory: %w", err)
}
if err := os.Chmod(finalDirectory, 0o700); err != nil {
return core.Asset{}, false, err
}
if info, statErr := os.Stat(finalPath); statErr == nil {
if !info.Mode().IsRegular() || info.Size() != staged.ByteSize || fileSHA256(finalPath) != staged.SHA256 {
return core.Asset{}, false, errors.New("stored evidence content does not match its key")
}
} else if !errors.Is(statErr, os.ErrNotExist) {
return core.Asset{}, false, statErr
} else {
publishPath, err := store.preparePublishFile(staged, finalDirectory)
if err != nil {
return core.Asset{}, false, err
}
defer os.Remove(publishPath)
if err := store.renameFile(publishPath, finalPath); err != nil {
return core.Asset{}, false, err
}
}
// The publication file was fsynced in this shard before its same-directory rename.
// Persist the final directory entry before SQLite can expose a referencing row.
// A directory sync failure is deliberately fatal; the unreachable file may remain
// as an orphan, but no evidence_assets row may be committed for it.
if err := store.syncDirectory(finalDirectory); err != nil {
return core.Asset{}, false, fmt.Errorf("persist evidence directory entry: %w", err)
}
id, err := newUUID(store.random)
if err != nil {
return core.Asset{}, false, err
}
now := store.now().UTC()
asset := core.Asset{
ID: id, TaskID: metadata.TaskID, AttemptID: metadata.AttemptID,
Kind: metadata.Kind, PrivacyTier: metadata.PrivacyTier, SHA256: staged.SHA256,
ByteSize: staged.ByteSize, ContentType: staged.ContentType, Width: staged.Width, Height: staged.Height,
CapturedAt: metadata.CapturedAt.UTC(), UploadedByDeviceID: principal.ID,
StorageKey: storageKey, CreatedAt: now,
}
_, err = transaction.ExecContext(ctx, `INSERT INTO evidence_assets
(id, upload_key, task_id, attempt_id, kind, privacy_tier, sha256, byte_size, content_type, width_px, height_px, storage_key, uploaded_by_device_id, captured_at, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
asset.ID, metadata.UploadKey, asset.TaskID, asset.AttemptID, asset.Kind, asset.PrivacyTier,
asset.SHA256, asset.ByteSize, asset.ContentType, asset.Width, asset.Height, asset.StorageKey,
asset.UploadedByDeviceID, asset.CapturedAt.Format(time.RFC3339Nano), asset.CreatedAt.Format(time.RFC3339Nano))
if err != nil {
return core.Asset{}, false, err
}
if err := store.commitTx(transaction); err != nil {
return core.Asset{}, false, err
}
return asset, false, nil
}
func (store *Store) Open(ctx context.Context, id string) (core.Asset, io.ReadSeekCloser, error) {
if !validUUID(id) {
return core.Asset{}, nil, core.ErrNotFound
}
asset, found, err := findByID(ctx, store.database, id)
if err != nil {
return core.Asset{}, nil, err
}
if !found || asset.StorageKey != storageKey(asset.SHA256) {
return core.Asset{}, nil, core.ErrNotFound
}
path, err := store.pathForKey(asset.StorageKey)
if err != nil {
return core.Asset{}, nil, core.ErrNotFound
}
file, err := os.Open(path)
if errors.Is(err, os.ErrNotExist) {
return core.Asset{}, nil, core.ErrNotFound
}
if err != nil {
return core.Asset{}, nil, err
}
info, err := file.Stat()
if err != nil || !info.Mode().IsRegular() || info.Size() != asset.ByteSize {
_ = file.Close()
if err != nil {
return core.Asset{}, nil, err
}
return core.Asset{}, nil, core.ErrNotFound
}
return asset, file, nil
}
func (store *Store) verifyStoredFile(asset core.Asset) error {
path, err := store.pathForKey(asset.StorageKey)
if err != nil {
return err
}
info, err := os.Stat(path)
if err != nil || !info.Mode().IsRegular() || info.Size() != asset.ByteSize || fileSHA256(path) != asset.SHA256 {
if err != nil {
return err
}
return errors.New("stored evidence file is invalid")
}
return nil
}
func (store *Store) isStagedPath(path string) bool {
if path == "" {
return false
}
relative, err := filepath.Rel(filepath.Join(store.root, ".staging"), filepath.Clean(path))
return err == nil && relative != "." && relative != "" && relative != ".." && !strings.HasPrefix(relative, ".."+string(filepath.Separator)) && !filepath.IsAbs(relative)
}
func (store *Store) pathForKey(key string) (string, error) {
path := filepath.Join(store.root, filepath.FromSlash(key))
relative, err := filepath.Rel(store.root, path)
if err != nil || relative == "." || relative == "" || relative == ".." || strings.HasPrefix(relative, ".."+string(filepath.Separator)) || filepath.IsAbs(relative) {
return "", errors.New("invalid evidence storage key")
}
return path, nil
}
func (store *Store) preparePublishFile(staged core.StagedFile, directory string) (path string, resultErr error) {
source, err := os.Open(staged.Path)
if err != nil {
return "", err
}
defer source.Close()
temporary, err := os.CreateTemp(directory, ".publish-*.png")
if err != nil {
return "", err
}
temporaryPath := temporary.Name()
path = temporaryPath
defer func() {
if resultErr != nil {
_ = temporary.Close()
_ = os.Remove(temporaryPath)
}
}()
if err := temporary.Chmod(0o600); err != nil {
return "", err
}
hasher := sha256.New()
written, err := io.Copy(io.MultiWriter(temporary, hasher), source)
if err != nil {
return "", err
}
if written != staged.ByteSize || hex.EncodeToString(hasher.Sum(nil)) != staged.SHA256 {
return "", errors.New("staged evidence changed before publication")
}
width, height, err := validatePNG(temporary)
if err != nil {
return "", err
}
if width != staged.Width || height != staged.Height {
return "", errors.New("staged evidence dimensions changed before publication")
}
if err := store.syncFile(temporary); err != nil {
return "", fmt.Errorf("sync evidence publication file: %w", err)
}
if err := temporary.Close(); err != nil {
return "", err
}
return path, nil
}
func validatePNG(reader io.ReadSeeker) (int, int, error) {
if _, err := reader.Seek(0, io.SeekStart); err != nil {
return 0, 0, err
}
signature := make([]byte, len(pngSignature))
if _, err := io.ReadFull(reader, signature); err != nil || string(signature) != string(pngSignature) {
return 0, 0, core.ErrInvalid
}
if _, err := reader.Seek(0, io.SeekStart); err != nil {
return 0, 0, err
}
configuration, err := png.DecodeConfig(reader)
if err != nil || configuration.Width < 1 || configuration.Height < 1 || configuration.Width > core.MaxImageSide || configuration.Height > core.MaxImageSide || int64(configuration.Width)*int64(configuration.Height) > core.MaxImagePixels {
return 0, 0, core.ErrInvalid
}
if _, err := reader.Seek(0, io.SeekStart); err != nil {
return 0, 0, err
}
if _, err := png.Decode(reader); err != nil {
return 0, 0, core.ErrInvalid
}
var trailing [1]byte
if count, err := reader.Read(trailing[:]); count != 0 || !errors.Is(err, io.EOF) {
return 0, 0, core.ErrInvalid
}
return configuration.Width, configuration.Height, nil
}
func ensureDurableDirectory(path string, mode os.FileMode, syncParent func(string) error) error {
info, err := os.Stat(path)
if err == nil {
if !info.IsDir() {
return fmt.Errorf("path exists but is not a directory: %s", path)
}
return nil
}
if !errors.Is(err, os.ErrNotExist) {
return err
}
parent := filepath.Dir(path)
if parent == path {
return fmt.Errorf("cannot create filesystem root as a managed directory: %s", path)
}
if err := ensureDurableDirectory(parent, mode, syncParent); err != nil {
return err
}
if err := os.Mkdir(path, mode); err != nil && !errors.Is(err, os.ErrExist) {
return err
}
info, err = os.Stat(path)
if err != nil {
return err
}
if !info.IsDir() {
return fmt.Errorf("path exists but is not a directory: %s", path)
}
if err := os.Chmod(path, mode); err != nil {
return err
}
// Syncing the parent makes creation of this directory durable. This also covers
// a concurrent creator: returning success without the parent sync could otherwise
// allow the following database transaction to outrun the directory entry.
if err := syncParent(parent); err != nil {
return fmt.Errorf("persist directory creation for %s: %w", path, err)
}
return nil
}
func storageKey(hash string) string { return hash[:2] + "/" + hash + ".png" }
func validMetadata(metadata core.UploadMetadata) bool {
return validUUID(metadata.UploadKey) && validUUID(metadata.TaskID) && validUUID(metadata.AttemptID) && metadata.Kind == core.KindSKUPanelGate1 && metadata.PrivacyTier == core.PrivacyInternalRaw && validSHA256(metadata.SHA256) && !metadata.CapturedAt.IsZero() && metadata.CapturedAt.Location() == time.UTC
}
func validPrincipal(principal deviceauth.Principal) bool {
return deviceauth.ValidDeviceID(principal.ID)
}
func validSHA256(value string) bool {
if len(value) != 64 {
return false
}
for _, character := range value {
if !(character >= '0' && character <= '9' || character >= 'a' && character <= 'f') {
return false
}
}
return true
}
func validUUID(value string) bool {
if len(value) != 36 {
return false
}
for index, character := range value {
if index == 8 || index == 13 || index == 18 || index == 23 {
if character != '-' {
return false
}
continue
}
if !(character >= '0' && character <= '9' || character >= 'a' && character <= 'f') {
return false
}
}
return value[14] == '4' && (value[19] == '8' || value[19] == '9' || value[19] == 'a' || value[19] == 'b')
}
func newUUID(reader io.Reader) (string, error) {
bytes := make([]byte, 16)
if _, err := io.ReadFull(reader, bytes); err != nil {
return "", err
}
bytes[6] = (bytes[6] & 0x0f) | 0x40
bytes[8] = (bytes[8] & 0x3f) | 0x80
encoded := hex.EncodeToString(bytes)
return encoded[:8] + "-" + encoded[8:12] + "-" + encoded[12:16] + "-" + encoded[16:20] + "-" + encoded[20:], nil
}
func fileSHA256(path string) string {
file, err := os.Open(path)
if err != nil {
return ""
}
defer file.Close()
hasher := sha256.New()
if _, err := io.Copy(hasher, file); err != nil {
return ""
}
return hex.EncodeToString(hasher.Sum(nil))
}
type rowScanner interface{ Scan(...any) error }
func findByUploadKey(ctx context.Context, query interface {
QueryRowContext(context.Context, string, ...any) *sql.Row
}, deviceID, uploadKey string) (core.Asset, bool, error) {
return scanAsset(query.QueryRowContext(ctx, `SELECT id, task_id, attempt_id, kind, privacy_tier, sha256, byte_size, content_type, width_px, height_px, storage_key, uploaded_by_device_id, captured_at, created_at FROM evidence_assets WHERE uploaded_by_device_id = ? AND upload_key = ?`, deviceID, uploadKey))
}
func findByID(ctx context.Context, query interface {
QueryRowContext(context.Context, string, ...any) *sql.Row
}, id string) (core.Asset, bool, error) {
return scanAsset(query.QueryRowContext(ctx, `SELECT id, task_id, attempt_id, kind, privacy_tier, sha256, byte_size, content_type, width_px, height_px, storage_key, uploaded_by_device_id, captured_at, created_at FROM evidence_assets WHERE id = ?`, id))
}
func scanAsset(row rowScanner) (core.Asset, bool, error) {
var asset core.Asset
var captured, created string
err := row.Scan(&asset.ID, &asset.TaskID, &asset.AttemptID, &asset.Kind, &asset.PrivacyTier, &asset.SHA256, &asset.ByteSize, &asset.ContentType, &asset.Width, &asset.Height, &asset.StorageKey, &asset.UploadedByDeviceID, &captured, &created)
if errors.Is(err, sql.ErrNoRows) {
return core.Asset{}, false, nil
}
if err != nil {
return core.Asset{}, false, err
}
asset.CapturedAt, err = time.Parse(time.RFC3339Nano, captured)
if err != nil {
return core.Asset{}, false, err
}
asset.CreatedAt, err = time.Parse(time.RFC3339Nano, created)
if err != nil {
return core.Asset{}, false, err
}
return asset, true, nil
}
func sameUpload(asset core.Asset, principal deviceauth.Principal, metadata core.UploadMetadata, staged core.StagedFile) bool {
return asset.TaskID == metadata.TaskID && asset.AttemptID == metadata.AttemptID && asset.Kind == metadata.Kind && asset.PrivacyTier == metadata.PrivacyTier && asset.SHA256 == metadata.SHA256 && asset.ByteSize == staged.ByteSize && asset.ContentType == staged.ContentType && asset.Width == staged.Width && asset.Height == staged.Height && asset.UploadedByDeviceID == principal.ID && asset.CapturedAt.Equal(metadata.CapturedAt)
}
@@ -0,0 +1,604 @@
package evidence
import (
"bytes"
"context"
"crypto/sha256"
"database/sql"
"encoding/hex"
"errors"
"image"
"image/color"
"image/png"
"io"
"os"
"path/filepath"
"runtime"
"strings"
"sync"
"testing"
"time"
"cmbuyer/admin/internal/deviceauth"
core "cmbuyer/admin/internal/evidence"
"cmbuyer/admin/internal/migrations"
"cmbuyer/admin/internal/storage/sqlite"
)
const (
testTaskID = "13c9f507-7473-4fa6-8d71-8786c34c6301"
testAuthID = "23c9f507-7473-4fa6-8d71-8786c34c6301"
testAttemptID = "33c9f507-7473-4fa6-8d71-8786c34c6301"
testUploadKey = "43c9f507-7473-4fa6-8d71-8786c34c6301"
testDeviceID = "53c9f507-7473-4fa6-8d71-8786c34c6301"
)
func TestStageCommitReplayAndOpen(t *testing.T) {
database, store := newTestStore(t)
insertAttemptFixture(t, database)
pngBytes := makePNG(t, 8, 6)
hash := sha256Hex(pngBytes)
metadata := core.UploadMetadata{UploadKey: testUploadKey, TaskID: testTaskID, AttemptID: testAttemptID, Kind: core.KindSKUPanelGate1, PrivacyTier: core.PrivacyInternalRaw, SHA256: hash, CapturedAt: time.Date(2026, 8, 4, 1, 2, 3, 0, time.UTC)}
principal := deviceauth.Principal{ID: testDeviceID}
staged, err := store.Stage(bytes.NewReader(pngBytes), core.PNGContentType)
if err != nil {
t.Fatalf("Stage: %v", err)
}
asset, replayed, err := store.Commit(context.Background(), principal, metadata, staged)
if err != nil || replayed {
t.Fatalf("Commit = replayed %t, err %v", replayed, err)
}
if asset.SHA256 != hash || asset.ByteSize != int64(len(pngBytes)) || asset.Width != 8 || asset.Height != 6 || asset.StorageKey != hash[:2]+"/"+hash+".png" {
t.Fatalf("asset = %#v", asset)
}
opened, reader, err := store.Open(context.Background(), asset.ID)
if err != nil {
t.Fatalf("Open: %v", err)
}
got, err := io.ReadAll(reader)
_ = reader.Close()
if err != nil || !bytes.Equal(got, pngBytes) || opened.ID != asset.ID {
t.Fatalf("opened asset changed: bytes=%t asset=%#v err=%v", bytes.Equal(got, pngBytes), opened, err)
}
replayStage, err := store.Stage(bytes.NewReader(pngBytes), core.PNGContentType)
if err != nil {
t.Fatalf("stage replay: %v", err)
}
replayedAsset, replayed, err := store.Commit(context.Background(), principal, metadata, replayStage)
if err != nil || !replayed || replayedAsset.ID != asset.ID {
t.Fatalf("replay = %#v, %t, %v", replayedAsset, replayed, err)
}
conflictStage, err := store.Stage(bytes.NewReader(pngBytes), core.PNGContentType)
if err != nil {
t.Fatalf("stage conflict: %v", err)
}
conflicting := metadata
conflicting.CapturedAt = conflicting.CapturedAt.Add(time.Second)
if _, _, err := store.Commit(context.Background(), principal, conflicting, conflictStage); !errors.Is(err, core.ErrConflict) {
t.Fatalf("conflicting replay error = %v", err)
}
}
func TestCommitRequiresCurrentClaimOwnerButClosedClaimKeepsHistoricalReplay(t *testing.T) {
database, store := newTestStore(t)
insertAttemptFixture(t, database)
pngBytes := makePNG(t, 4, 3)
metadata := testMetadata(sha256Hex(pngBytes))
stage := func() core.StagedFile {
staged, err := store.Stage(bytes.NewReader(pngBytes), core.PNGContentType)
if err != nil {
t.Fatalf("Stage: %v", err)
}
return staged
}
otherDevice := deviceauth.Principal{ID: "73c9f507-7473-4fa6-8d71-8786c34c6301"}
if _, _, err := store.Commit(context.Background(), otherDevice, metadata, stage()); !errors.Is(err, core.ErrInvalid) {
t.Fatalf("device B upload to device A attempt error = %v", err)
}
principal := deviceauth.Principal{ID: testDeviceID}
asset, replayed, err := store.Commit(context.Background(), principal, metadata, stage())
if err != nil || replayed {
t.Fatalf("owner first Commit = replayed %v, err %v", replayed, err)
}
if _, err := database.Exec("UPDATE purchase_attempt_claims SET closed_at='2026-08-04T03:00:00Z' WHERE attempt_id=?", testAttemptID); err != nil {
t.Fatalf("close claim: %v", err)
}
replayedAsset, replayed, err := store.Commit(context.Background(), principal, metadata, stage())
if err != nil || !replayed || replayedAsset.ID != asset.ID {
t.Fatalf("closed claim historical replay = %#v replayed %v err %v", replayedAsset, replayed, err)
}
newMetadata := metadata
newMetadata.UploadKey = "83c9f507-7473-4fa6-8d71-8786c34c6301"
if _, _, err := store.Commit(context.Background(), principal, newMetadata, stage()); !errors.Is(err, core.ErrInvalid) {
t.Fatalf("closed claim new upload error = %v", err)
}
}
func TestConcurrentReplayCreatesOneAsset(t *testing.T) {
database, store := newTestStore(t)
insertAttemptFixture(t, database)
pngBytes := makePNG(t, 3, 2)
metadata := core.UploadMetadata{UploadKey: testUploadKey, TaskID: testTaskID, AttemptID: testAttemptID, Kind: core.KindSKUPanelGate1, PrivacyTier: core.PrivacyInternalRaw, SHA256: sha256Hex(pngBytes), CapturedAt: time.Date(2026, 8, 4, 1, 2, 3, 0, time.UTC)}
staged := make([]core.StagedFile, 2)
for index := range staged {
var err error
staged[index], err = store.Stage(bytes.NewReader(pngBytes), core.PNGContentType)
if err != nil {
t.Fatalf("Stage %d: %v", index, err)
}
}
var wait sync.WaitGroup
wait.Add(2)
assets := make([]core.Asset, 2)
replays := make([]bool, 2)
errorsSeen := make([]error, 2)
for index := range staged {
go func(index int) {
defer wait.Done()
assets[index], replays[index], errorsSeen[index] = store.Commit(context.Background(), deviceauth.Principal{ID: testDeviceID}, metadata, staged[index])
}(index)
}
wait.Wait()
if errorsSeen[0] != nil || errorsSeen[1] != nil || assets[0].ID != assets[1].ID || replays[0] == replays[1] {
t.Fatalf("concurrent commits assets=%#v replays=%#v errors=%#v", assets, replays, errorsSeen)
}
var count int
if err := database.QueryRow("SELECT COUNT(*) FROM evidence_assets").Scan(&count); err != nil || count != 1 {
t.Fatalf("asset count = %d, err %v", count, err)
}
}
func TestPlatformDirectorySync(t *testing.T) {
if err := syncDirectory(t.TempDir()); err != nil {
t.Fatalf("syncDirectory must either establish the durability boundary or fail closed: %v", err)
}
}
func TestNewStoreRetriesRootParentSyncWhenRootAlreadyExists(t *testing.T) {
database, _ := newTestStore(t)
parent := t.TempDir()
root := filepath.Join(parent, "retry-root")
injected := errors.New("injected root parent sync failure")
if _, err := newStore(database, root, func(string) error { return injected }); !errors.Is(err, injected) {
t.Fatalf("first newStore error = %v, want injected root sync failure", err)
}
if info, err := os.Stat(root); err != nil || !info.IsDir() {
t.Fatalf("failed parent sync must leave root for retry: info=%v err=%v", info, err)
}
var paths []string
store, err := newStore(database, root, func(path string) error {
paths = append(paths, path)
return syncDirectory(path)
})
if err != nil {
t.Fatalf("retry newStore: %v", err)
}
if store == nil || len(paths) == 0 || paths[0] != parent {
t.Fatalf("retry sync paths = %#v, want root parent %q first", paths, parent)
}
}
func TestCommitSyncsShardAndRenameBeforeDatabaseWrite(t *testing.T) {
database, store := newTestStore(t)
insertAttemptFixture(t, database)
pngBytes := makePNG(t, 4, 3)
hash := sha256Hex(pngBytes)
metadata := testMetadata(hash)
staged, err := store.Stage(bytes.NewReader(pngBytes), core.PNGContentType)
if err != nil {
t.Fatalf("Stage: %v", err)
}
finalPath, err := store.pathForKey(storageKey(hash))
if err != nil {
t.Fatalf("final path: %v", err)
}
finalDirectory := filepath.Dir(finalPath)
var events []string
store.syncDirectory = func(path string) error {
var count int
if err := database.QueryRow("SELECT COUNT(*) FROM evidence_assets").Scan(&count); err != nil {
t.Fatalf("count evidence before directory sync: %v", err)
}
if count != 0 {
t.Fatalf("database row became visible before directory sync: %d", count)
}
switch path {
case store.root:
events = append(events, "sync-root")
if path != store.root {
t.Fatalf("shard parent sync path = %q, want evidence root %q", path, store.root)
}
if info, err := os.Stat(finalDirectory); err != nil || !info.IsDir() {
t.Fatalf("shard directory must exist before parent sync: info=%v err=%v", info, err)
}
if _, err := os.Stat(finalPath); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("final file exists before publication: %v", err)
}
case finalDirectory:
events = append(events, "sync-shard")
if info, err := os.Stat(finalPath); err != nil || !info.Mode().IsRegular() {
t.Fatalf("renamed file must exist before shard sync: info=%v err=%v", info, err)
}
default:
t.Fatalf("unexpected extra directory sync: %q", path)
}
return syncDirectory(path)
}
store.syncFile = func(file *os.File) error {
if filepath.Dir(file.Name()) != finalDirectory || !strings.HasPrefix(filepath.Base(file.Name()), ".publish-") {
t.Fatalf("publication temp is not inside shard: %q", file.Name())
}
events = append(events, "sync-file")
return file.Sync()
}
store.renameFile = func(oldPath, newPath string) error {
if filepath.Dir(oldPath) != filepath.Dir(newPath) || newPath != finalPath {
t.Fatalf("rename is not same-directory publication: %q -> %q", oldPath, newPath)
}
events = append(events, "rename")
return os.Rename(oldPath, newPath)
}
if _, replayed, err := store.Commit(context.Background(), deviceauth.Principal{ID: testDeviceID}, metadata, staged); err != nil || replayed {
t.Fatalf("Commit = replayed %t, err %v", replayed, err)
}
if got, want := strings.Join(events, ","), "sync-root,sync-root,sync-file,rename,sync-shard"; got != want {
t.Fatalf("durability order = %q, want %q", got, want)
}
assertEvidenceCount(t, database, 1)
assertNoPublishTemps(t, finalDirectory)
}
func TestCommitDirectorySyncFailuresNeverWriteDatabase(t *testing.T) {
for _, failAt := range []int{1, 2, 3} {
t.Run(map[int]string{1: "new shard parent", 2: "unconditional shard parent", 3: "rename target"}[failAt], func(t *testing.T) {
database, store := newTestStore(t)
insertAttemptFixture(t, database)
pngBytes := makePNG(t, 4, 3)
hash := sha256Hex(pngBytes)
staged, err := store.Stage(bytes.NewReader(pngBytes), core.PNGContentType)
if err != nil {
t.Fatalf("Stage: %v", err)
}
finalPath, err := store.pathForKey(storageKey(hash))
if err != nil {
t.Fatalf("final path: %v", err)
}
injected := errors.New("injected directory sync failure")
calls := 0
store.syncDirectory = func(path string) error {
calls++
if calls == failAt {
return injected
}
return syncDirectory(path)
}
if _, _, err := store.Commit(context.Background(), deviceauth.Principal{ID: testDeviceID}, testMetadata(hash), staged); !errors.Is(err, injected) {
t.Fatalf("Commit error = %v, want injected sync failure", err)
}
if calls != failAt {
t.Fatalf("sync calls = %d, want %d", calls, failAt)
}
assertEvidenceCount(t, database, 0)
_, statErr := os.Stat(finalPath)
if failAt < 3 && !errors.Is(statErr, os.ErrNotExist) {
t.Fatalf("file exists before rename durability boundary: %v", statErr)
}
if failAt == 3 && statErr != nil {
t.Fatalf("post-rename sync failure may leave an orphan file, stat error = %v", statErr)
}
assertNoPublishTemps(t, filepath.Dir(finalPath))
})
}
}
func TestCommitRetriesShardParentSyncAfterPriorFailure(t *testing.T) {
database, store := newTestStore(t)
insertAttemptFixture(t, database)
pngBytes := makePNG(t, 4, 3)
hash := sha256Hex(pngBytes)
finalPath, err := store.pathForKey(storageKey(hash))
if err != nil {
t.Fatalf("final path: %v", err)
}
firstStage, err := store.Stage(bytes.NewReader(pngBytes), core.PNGContentType)
if err != nil {
t.Fatalf("first Stage: %v", err)
}
injected := errors.New("injected first shard parent sync failure")
store.syncDirectory = func(string) error { return injected }
if _, _, err := store.Commit(context.Background(), deviceauth.Principal{ID: testDeviceID}, testMetadata(hash), firstStage); !errors.Is(err, injected) {
t.Fatalf("first Commit error = %v", err)
}
if info, err := os.Stat(filepath.Dir(finalPath)); err != nil || !info.IsDir() {
t.Fatalf("failed first sync must leave the created shard for retry: info=%v err=%v", info, err)
}
assertEvidenceCount(t, database, 0)
secondStage, err := store.Stage(bytes.NewReader(pngBytes), core.PNGContentType)
if err != nil {
t.Fatalf("second Stage: %v", err)
}
var paths []string
store.syncDirectory = func(path string) error {
paths = append(paths, path)
return syncDirectory(path)
}
if _, replayed, err := store.Commit(context.Background(), deviceauth.Principal{ID: testDeviceID}, testMetadata(hash), secondStage); err != nil || replayed {
t.Fatalf("retry Commit = replayed %t, err %v", replayed, err)
}
if len(paths) != 2 || paths[0] != store.root || paths[1] != filepath.Dir(finalPath) {
t.Fatalf("retry sync paths = %#v, want root then shard", paths)
}
assertEvidenceCount(t, database, 1)
assertNoPublishTemps(t, filepath.Dir(finalPath))
}
func TestCommitPublicationFailuresCleanTempAndNeverWriteDatabase(t *testing.T) {
for _, name := range []string{"file sync", "rename"} {
t.Run(name, func(t *testing.T) {
database, store := newTestStore(t)
insertAttemptFixture(t, database)
pngBytes := makePNG(t, 4, 3)
hash := sha256Hex(pngBytes)
staged, err := store.Stage(bytes.NewReader(pngBytes), core.PNGContentType)
if err != nil {
t.Fatalf("Stage: %v", err)
}
finalPath, err := store.pathForKey(storageKey(hash))
if err != nil {
t.Fatalf("final path: %v", err)
}
injected := errors.New("injected publication failure")
if name == "file sync" {
store.syncFile = func(*os.File) error { return injected }
} else {
store.renameFile = func(string, string) error { return injected }
}
if _, _, err := store.Commit(context.Background(), deviceauth.Principal{ID: testDeviceID}, testMetadata(hash), staged); !errors.Is(err, injected) {
t.Fatalf("Commit error = %v, want injected publication failure", err)
}
if _, err := os.Stat(finalPath); !errors.Is(err, os.ErrNotExist) {
t.Fatalf("final file exists after failed publication: %v", err)
}
assertNoPublishTemps(t, filepath.Dir(finalPath))
assertEvidenceCount(t, database, 0)
})
}
}
func TestCommitDatabaseFailuresAfterDurableRenameLeaveOnlyOrphan(t *testing.T) {
for _, test := range []struct {
name string
inject func(*testing.T, *sql.DB, *Store, error)
}{
{
name: "insert",
inject: func(t *testing.T, database *sql.DB, _ *Store, _ error) {
t.Helper()
if _, err := database.Exec(`CREATE TRIGGER fail_evidence_insert BEFORE INSERT ON evidence_assets BEGIN SELECT RAISE(ABORT, 'injected insert failure'); END`); err != nil {
t.Fatalf("create insert failure trigger: %v", err)
}
},
},
{
name: "commit",
inject: func(_ *testing.T, _ *sql.DB, store *Store, injected error) {
store.commitTx = func(*sql.Tx) error { return injected }
},
},
} {
t.Run(test.name, func(t *testing.T) {
database, store := newTestStore(t)
insertAttemptFixture(t, database)
pngBytes := makePNG(t, 4, 3)
hash := sha256Hex(pngBytes)
staged, err := store.Stage(bytes.NewReader(pngBytes), core.PNGContentType)
if err != nil {
t.Fatalf("Stage: %v", err)
}
finalPath, err := store.pathForKey(storageKey(hash))
if err != nil {
t.Fatalf("final path: %v", err)
}
injected := errors.New("injected database failure")
test.inject(t, database, store, injected)
syncCalls := 0
store.syncDirectory = func(path string) error {
syncCalls++
return syncDirectory(path)
}
if _, _, err := store.Commit(context.Background(), deviceauth.Principal{ID: testDeviceID}, testMetadata(hash), staged); err == nil {
t.Fatal("Commit unexpectedly succeeded")
}
if syncCalls != 3 {
t.Fatalf("database failure occurred before both durability syncs: sync calls = %d", syncCalls)
}
if info, err := os.Stat(finalPath); err != nil || !info.Mode().IsRegular() {
t.Fatalf("durable rename may leave only an orphan file: info=%v err=%v", info, err)
}
assertNoPublishTemps(t, filepath.Dir(finalPath))
assertEvidenceCount(t, database, 0)
})
}
}
func TestStageRejectsUnsafeContent(t *testing.T) {
_, store := newTestStore(t)
largePNG := makePNG(t, core.MaxImageSide+1, 1)
pngWithXML := append(makePNG(t, 1, 1), []byte("<hierarchy/>")...)
for name, test := range map[string]struct {
reader io.Reader
contentType string
}{
"wrong content type": {reader: bytes.NewReader(makePNG(t, 1, 1)), contentType: "application/octet-stream"},
"xml": {reader: bytes.NewBufferString("<hierarchy/>"), contentType: core.PNGContentType},
"png with xml tail": {reader: bytes.NewReader(pngWithXML), contentType: core.PNGContentType},
"truncated png": {reader: bytes.NewReader(pngSignature), contentType: core.PNGContentType},
"too wide": {reader: bytes.NewReader(largePNG), contentType: core.PNGContentType},
"too many bytes": {reader: io.LimitReader(zeroReader{}, core.MaxFileBytes+1), contentType: core.PNGContentType},
} {
t.Run(name, func(t *testing.T) {
staged, err := store.Stage(test.reader, test.contentType)
if !errors.Is(err, core.ErrInvalid) && !errors.Is(err, core.ErrTooLarge) {
store.Discard(staged)
t.Fatalf("Stage error = %v", err)
}
})
}
}
func TestCommitRequiresAttemptOwnedByTaskAndLowercaseHash(t *testing.T) {
database, store := newTestStore(t)
insertAttemptFixture(t, database)
pngBytes := makePNG(t, 2, 2)
base := core.UploadMetadata{UploadKey: testUploadKey, TaskID: testTaskID, AttemptID: testAttemptID, Kind: core.KindSKUPanelGate1, PrivacyTier: core.PrivacyInternalRaw, SHA256: sha256Hex(pngBytes), CapturedAt: time.Date(2026, 8, 4, 1, 2, 3, 0, time.UTC)}
for name, mutate := range map[string]func(*core.UploadMetadata){
"unknown attempt": func(value *core.UploadMetadata) { value.AttemptID = "53c9f507-7473-4fa6-8d71-8786c34c6301" },
"uppercase hash": func(value *core.UploadMetadata) {
value.SHA256 = "AAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAAA"
},
"wrong kind": func(value *core.UploadMetadata) { value.Kind = "ORDER_CONFIRM" },
} {
t.Run(name, func(t *testing.T) {
staged, err := store.Stage(bytes.NewReader(pngBytes), core.PNGContentType)
if err != nil {
t.Fatalf("Stage: %v", err)
}
metadata := base
mutate(&metadata)
if _, _, err := store.Commit(context.Background(), deviceauth.Principal{ID: testDeviceID}, metadata, staged); !errors.Is(err, core.ErrInvalid) {
t.Fatalf("Commit error = %v", err)
}
})
}
var count int
if err := database.QueryRow("SELECT COUNT(*) FROM evidence_assets").Scan(&count); err != nil || count != 0 {
t.Fatalf("invalid commits created %d assets, err %v", count, err)
}
}
func TestNewStoreRejectsRelativeAndFilesystemRootPaths(t *testing.T) {
database, _ := newTestStore(t)
if _, err := NewStore(database, "relative-evidence"); err == nil {
t.Fatal("relative evidence root succeeded")
}
volumeRoot := filepath.VolumeName(t.TempDir()) + string(filepath.Separator)
if _, err := NewStore(database, volumeRoot); err == nil {
t.Fatal("filesystem root succeeded")
}
}
type zeroReader struct{}
func (zeroReader) Read(buffer []byte) (int, error) {
for index := range buffer {
buffer[index] = 0
}
return len(buffer), nil
}
func newTestStore(t *testing.T) (*sql.DB, *Store) {
t.Helper()
database, err := sqlite.Open(filepath.Join(t.TempDir(), "evidence.db"))
if err != nil {
t.Fatalf("open database: %v", err)
}
t.Cleanup(func() { _ = database.Close() })
_, file, _, ok := runtime.Caller(0)
if !ok {
t.Fatal("locate test")
}
directory := filepath.Join(filepath.Dir(file), "..", "..", "..", "migrations")
if err := migrations.Up(context.Background(), database, directory); err != nil {
t.Fatalf("migrate database: %v", err)
}
store, err := NewStore(database, filepath.Join(t.TempDir(), "assets"))
if err != nil {
t.Fatalf("NewStore: %v", err)
}
return database, store
}
func insertAttemptFixture(t *testing.T, database *sql.DB) {
t.Helper()
timestamp := "2026-08-04T00:00:00Z"
tokenHash := sha256.Sum256([]byte("evidence-device-token"))
if _, err := database.Exec(`INSERT INTO device_credentials
(device_id,display_name,token_sha256,status,created_at,revoked_at)
VALUES (?, 'evidence device', ?, 'ACTIVE', ?, NULL)`, testDeviceID, tokenHash[:], timestamp); err != nil {
t.Fatalf("insert device: %v", err)
}
if _, err := database.Exec(`INSERT INTO tasks (id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price, status, version, created_at, updated_at) VALUES (?, 'MANUAL', 'task', '123', 'black', 'M', 1, '1.00', 'DRAFT', 1, ?, ?)`, testTaskID, timestamp, timestamp); err != nil {
t.Fatalf("insert task: %v", err)
}
if _, err := database.Exec(`INSERT INTO order_authorizations (id, task_id, task_version, start_key, goods_id, sku_color, sku_size, quantity, total_price_cap, status, created_by, created_at, expires_at) VALUES (?, ?, 1, 'start', '123', 'black', 'M', 1, '1.00', 'ACTIVE', 'admin', ?, ?)`, testAuthID, testTaskID, timestamp, timestamp); err != nil {
t.Fatalf("insert authorization: %v", err)
}
if _, err := database.Exec(`INSERT INTO purchase_attempts (id, task_id, authorization_id, claim_generation, status, started_at) VALUES (?, ?, ?, 1, 'CLAIMED', ?)`, testAttemptID, testTaskID, testAuthID, timestamp); err != nil {
t.Fatalf("insert attempt: %v", err)
}
if _, err := database.Exec(`INSERT INTO purchase_attempt_claims
(attempt_id,task_id,authorization_id,claimed_by_device_id,session_id,claim_generation,
task_version,task_title,authorization_task_version,goods_id,sku_color,sku_size,quantity,
total_price_cap,authorization_expires_at,claim_nonce,claim_token_sha256,lease_expires_at,claimed_at,closed_at)
VALUES (?, ?, ?, ?, '63c9f507-7473-4fa6-8d71-8786c34c6301', 1, 1, 'task',
1, '123', 'black', 'M', 1, '1.00', ?, ?, ?, '2026-08-04T02:00:00Z', ?, NULL)`,
testAttemptID, testTaskID, testAuthID, testDeviceID, timestamp, bytes.Repeat([]byte{1}, 32), bytes.Repeat([]byte{2}, 32), timestamp); err != nil {
t.Fatalf("insert claim: %v", err)
}
}
func testMetadata(hash string) core.UploadMetadata {
return core.UploadMetadata{
UploadKey: testUploadKey, TaskID: testTaskID, AttemptID: testAttemptID,
Kind: core.KindSKUPanelGate1, PrivacyTier: core.PrivacyInternalRaw, SHA256: hash,
CapturedAt: time.Date(2026, 8, 4, 1, 2, 3, 0, time.UTC),
}
}
func assertEvidenceCount(t *testing.T, database *sql.DB, want int) {
t.Helper()
var count int
if err := database.QueryRow("SELECT COUNT(*) FROM evidence_assets").Scan(&count); err != nil {
t.Fatalf("count evidence assets: %v", err)
}
if count != want {
t.Fatalf("evidence asset count = %d, want %d", count, want)
}
}
func assertNoPublishTemps(t *testing.T, directory string) {
t.Helper()
entries, err := os.ReadDir(directory)
if err != nil {
t.Fatalf("read shard directory: %v", err)
}
for _, entry := range entries {
if strings.HasPrefix(entry.Name(), ".publish-") {
t.Fatalf("publication temp leaked: %q", entry.Name())
}
}
}
func makePNG(t *testing.T, width, height int) []byte {
t.Helper()
imageData := image.NewNRGBA(image.Rect(0, 0, width, height))
imageData.Set(0, 0, color.NRGBA{R: 12, G: 34, B: 56, A: 255})
var buffer bytes.Buffer
if err := png.Encode(&buffer, imageData); err != nil {
t.Fatalf("encode PNG: %v", err)
}
return buffer.Bytes()
}
func sha256Hex(value []byte) string {
hash := sha256.Sum256(value)
return hex.EncodeToString(hash[:])
}
+823
View File
@@ -0,0 +1,823 @@
package taskclaim
import (
"context"
"crypto/rand"
"crypto/sha256"
"database/sql"
"encoding/hex"
"errors"
"fmt"
"io"
"math"
"strings"
"sync"
"time"
"cmbuyer/admin/internal/deviceauth"
taskmodel "cmbuyer/admin/internal/tasks"
)
const writeTimeout = 2 * time.Second
type Store struct {
database *sql.DB
secret []byte
leaseTTL time.Duration
now func() time.Time
random io.Reader
randomMu sync.Mutex
writeGate chan struct{}
// The unexported linearization hooks let package tests coordinate real SQLite
// transactions at the first write. Production construction always leaves them nil.
beforeLinearization func()
afterLinearization func()
}
func NewStore(database *sql.DB, secret []byte, leaseTTL time.Duration) (*Store, error) {
if database == nil {
return nil, errors.New("task claim database is required")
}
if len(secret) != sha256.Size {
return nil, errors.New("task claim secret must be 32 bytes")
}
if leaseTTL <= 0 {
return nil, errors.New("task claim lease TTL must be positive")
}
if _, err := database.Exec("SELECT attempt_id, claim_nonce, claim_token_sha256 FROM purchase_attempt_claims LIMIT 1"); err != nil {
return nil, errors.New("task claim migration is not available")
}
store := &Store{
database: database, secret: append([]byte(nil), secret...), leaseTTL: leaseTTL,
now: time.Now, random: rand.Reader, writeGate: make(chan struct{}, 1),
}
if err := store.validateSecretIsolation(); err != nil {
return nil, err
}
if err := store.validateStoredClaims(context.Background()); err != nil {
return nil, err
}
return store, nil
}
// validateSecretIsolation ensures the HMAC key cannot also authenticate a device. The session
// secret comparison is performed while parsing configuration, before either secret is discarded.
func (store *Store) validateSecretIsolation() error {
digest := sha256.Sum256(store.secret)
var count int
if err := store.database.QueryRow(`SELECT COUNT(*) FROM device_credentials WHERE token_sha256 = ?`, digest[:]).Scan(&count); err != nil {
return errors.New("validate task claim secret isolation")
}
if count != 0 {
return errors.New("task claim secret must be isolated from device credentials")
}
return nil
}
// validateStoredClaims covers open and closed claims. Replacing the secret must fail startup;
// silently signing a new token would destroy idempotent recovery and the ownership audit chain.
func (store *Store) validateStoredClaims(ctx context.Context) error {
rows, err := store.database.QueryContext(ctx, `SELECT claims.claimed_by_device_id, claims.task_id, claims.authorization_id,
claims.attempt_id, claims.claim_generation, claims.claim_nonce, typeof(claims.claim_nonce), length(claims.claim_nonce),
claims.claim_token_sha256, typeof(claims.claim_token_sha256), length(claims.claim_token_sha256),
claims.task_title, claims.authorization_task_version, claims.goods_id, claims.sku_color, claims.sku_size,
claims.quantity, claims.total_price_cap, claims.authorization_expires_at, claims.closed_at,
attempts.claim_generation, attempts.status, authorizations.status, tasks.status
FROM purchase_attempt_claims AS claims
LEFT JOIN purchase_attempts AS attempts ON attempts.id = claims.attempt_id
LEFT JOIN order_authorizations AS authorizations ON authorizations.id = claims.authorization_id
LEFT JOIN tasks ON tasks.id = claims.task_id
ORDER BY claims.attempt_id`)
if err != nil {
return errors.New("validate stored task claims")
}
defer rows.Close()
for rows.Next() {
var deviceID, taskID, authorizationID, attemptID string
var generation, authorizationTaskVersion, quantity int
var nonce, storedHash []byte
var nonceType, hashType, title, goodsID, color, size, price, expires string
var nonceLength, hashLength int
var closed, attemptStatus, authorizationStatus, taskStatus sql.NullString
var attemptGeneration sql.NullInt64
if err := rows.Scan(&deviceID, &taskID, &authorizationID, &attemptID, &generation,
&nonce, &nonceType, &nonceLength, &storedHash, &hashType, &hashLength,
&title, &authorizationTaskVersion, &goodsID, &color, &size, &quantity, &price, &expires, &closed,
&attemptGeneration, &attemptStatus, &authorizationStatus, &taskStatus); err != nil {
return errors.New("validate stored task claims")
}
if !deviceauth.ValidDeviceID(deviceID) || !validUUID(taskID) || !validUUID(authorizationID) || !validUUID(attemptID) ||
generation <= 0 || nonceType != "blob" || nonceLength != sha256.Size || len(nonce) != sha256.Size ||
hashType != "blob" || hashLength != sha256.Size || len(storedHash) != sha256.Size ||
authorizationTaskVersion <= 0 || !taskmodel.ValidTaskWireFields(title, goodsID, color, size, price) || quantity <= 0 ||
!validCanonicalTime(expires) || (closed.Valid && !validCanonicalTime(closed.String)) ||
!attemptGeneration.Valid || attemptGeneration.Int64 != int64(generation) ||
!validAttemptStatus(attemptStatus) || !validAuthorizationStatus(authorizationStatus) || !validTaskStatus(taskStatus) {
return errors.New("stored task claim metadata is invalid")
}
token := deriveToken(store.secret, deviceID, taskID, authorizationID, attemptID, generation, nonce)
if !matchingHash(tokenHash(token), storedHash) {
return errors.New("task claim secret does not match stored claims")
}
}
if err := rows.Err(); err != nil {
return errors.New("validate stored task claims")
}
return nil
}
func (store *Store) ClaimNext(ctx context.Context, deviceID string, command ClaimCommand) (ClaimResponse, bool, error) {
if !deviceauth.ValidDeviceID(deviceID) || !validUUID(command.SessionID) || !validUUID(command.ClaimRequestID) {
return ClaimResponse{}, false, ErrInvalid
}
writeCtx, cancel := context.WithTimeout(ctx, writeTimeout)
defer cancel()
select {
case store.writeGate <- struct{}{}:
defer func() { <-store.writeGate }()
case <-writeCtx.Done():
return ClaimResponse{}, false, writeCtx.Err()
}
transaction, err := store.database.BeginTx(writeCtx, nil)
if err != nil {
return ClaimResponse{}, false, err
}
defer transaction.Rollback()
// This must be the transaction's first database statement. The no-op conditional UPDATE takes
// SQLite's write position and linearizes a concurrent credential revocation before any replay,
// EMPTY response, conflict response, candidate read, or other business write is possible.
if store.beforeLinearization != nil {
store.beforeLinearization()
}
active, err := transaction.ExecContext(writeCtx, `UPDATE device_credentials SET status = status
WHERE device_id = ? AND status = 'ACTIVE' AND revoked_at IS NULL`, deviceID)
if err != nil {
return ClaimResponse{}, false, err
}
if ok, err := exactlyOne(active); err != nil {
return ClaimResponse{}, false, err
} else if !ok {
return ClaimResponse{}, false, ErrDeviceInactive
}
if store.afterLinearization != nil {
store.afterLinearization()
}
now, err := store.serverNow()
if err != nil {
return ClaimResponse{}, false, err
}
request, found, err := findClaimRequest(writeCtx, transaction, command.ClaimRequestID)
if err != nil {
return ClaimResponse{}, false, err
}
if found {
if request.DeviceID != deviceID || request.SessionID != command.SessionID {
return ClaimResponse{}, false, ErrIdempotencyConflict
}
switch request.Outcome {
case "EMPTY":
if err := transaction.Commit(); err != nil {
return ClaimResponse{}, false, err
}
return ClaimResponse{}, false, nil
case "BLOCKED":
if err := transaction.Commit(); err != nil {
return ClaimResponse{}, false, err
}
return ClaimResponse{}, false, ErrRequiresManual
case "CLAIMED":
record, found, err := store.loadClaimByAttempt(writeCtx, transaction, request.AttemptID)
if err != nil || !found {
if err == nil {
err = errors.New("stored claim request has no claim")
}
return ClaimResponse{}, false, err
}
response, err := store.responseFor(record, request.ResponseLeaseExpiresAt)
if err != nil {
return ClaimResponse{}, false, err
}
if err := transaction.Commit(); err != nil {
return ClaimResponse{}, false, err
}
return response, true, nil
default:
return ClaimResponse{}, false, errors.New("stored claim request outcome is invalid")
}
}
existing, found, err := store.loadOpenClaimByDevice(writeCtx, transaction, deviceID)
if err != nil {
return ClaimResponse{}, false, err
}
if found {
current := existing.SessionID == command.SessionID && existing.ClosedAt == "" &&
existing.LeaseExpiresAt.After(now) && existing.AuthorizationExpiresAt.After(now) &&
existing.CurrentAuthorizationExpiresAt.After(now) && existing.AuthorizationStatus == "CLAIMED" &&
existing.authorizationConsistent() && existing.recoverableBusinessState()
if !current {
if err := insertClaimRequest(writeCtx, transaction, command.ClaimRequestID, deviceID, command.SessionID, "BLOCKED", "", "", "manual_recovery_required", now); err != nil {
return ClaimResponse{}, false, err
}
if err := transaction.Commit(); err != nil {
return ClaimResponse{}, false, err
}
return ClaimResponse{}, false, ErrRequiresManual
}
response, err := store.responseFor(existing, existing.LeaseExpiresText)
if err != nil {
return ClaimResponse{}, false, err
}
if err := insertClaimRequest(writeCtx, transaction, command.ClaimRequestID, deviceID, command.SessionID, "CLAIMED", existing.AttemptID, existing.LeaseExpiresText, "", now); err != nil {
return ClaimResponse{}, false, err
}
if err := transaction.Commit(); err != nil {
return ClaimResponse{}, false, err
}
return response, true, nil
}
candidate, found, err := findCandidate(writeCtx, transaction, now)
if err != nil {
return ClaimResponse{}, false, err
}
if !found {
if err := insertClaimRequest(writeCtx, transaction, command.ClaimRequestID, deviceID, command.SessionID, "EMPTY", "", "", "", now); err != nil {
return ClaimResponse{}, false, err
}
if err := transaction.Commit(); err != nil {
return ClaimResponse{}, false, err
}
return ClaimResponse{}, false, nil
}
generation, err := nextGeneration(writeCtx, transaction, candidate.TaskID)
if err != nil {
return ClaimResponse{}, false, err
}
attemptID, err := store.newUUID()
if err != nil {
return ClaimResponse{}, false, err
}
nonce, err := store.randomBytes(sha256.Size)
if err != nil {
return ClaimResponse{}, false, err
}
token := deriveToken(store.secret, deviceID, candidate.TaskID, candidate.AuthorizationID, attemptID, generation, nonce)
storedTokenHash := tokenHash(token)
leaseExpires := now.Add(store.leaseTTL)
if candidate.AuthorizationExpiresAt.Before(leaseExpires) {
leaseExpires = candidate.AuthorizationExpiresAt
}
leaseText := formatTime(leaseExpires)
nowText := formatTime(now)
authorizationUpdate, err := transaction.ExecContext(writeCtx, `UPDATE order_authorizations SET status = 'CLAIMED'
WHERE id = ? AND task_id = ? AND status = 'ACTIVE' AND task_version = ?
AND goods_id = ? AND sku_color = ? AND sku_size = ? AND quantity = ?
AND total_price_cap = ? AND expires_at = ?`,
candidate.AuthorizationID, candidate.TaskID, candidate.TaskVersion, candidate.GoodsID,
candidate.SKUColor, candidate.SKUSize, candidate.Quantity, candidate.TotalPriceCap,
candidate.AuthorizationExpiresText)
if err != nil {
return ClaimResponse{}, false, err
}
if ok, err := exactlyOne(authorizationUpdate); err != nil || !ok {
if err == nil {
err = errors.New("authorization changed during claim")
}
return ClaimResponse{}, false, err
}
taskUpdate, err := transaction.ExecContext(writeCtx, `UPDATE tasks SET status = 'CLAIMED', version = version + 1, updated_at = ?
WHERE id = ? AND status = 'PENDING' AND version = ? AND title = ? AND goods_id = ?
AND sku_color = ? AND sku_size = ? AND quantity = ? AND max_total_price = ?`,
nowText, candidate.TaskID, candidate.TaskVersion, candidate.Title, candidate.GoodsID,
candidate.SKUColor, candidate.SKUSize, candidate.Quantity, candidate.TotalPriceCap)
if err != nil {
return ClaimResponse{}, false, err
}
if ok, err := exactlyOne(taskUpdate); err != nil || !ok {
if err == nil {
err = errors.New("task changed during claim")
}
return ClaimResponse{}, false, err
}
if _, err := transaction.ExecContext(writeCtx, `INSERT INTO purchase_attempts
(id, task_id, authorization_id, claim_generation, status, started_at)
VALUES (?, ?, ?, ?, 'CLAIMED', ?)`, attemptID, candidate.TaskID, candidate.AuthorizationID, generation, nowText); err != nil {
return ClaimResponse{}, false, err
}
if _, err := transaction.ExecContext(writeCtx, `INSERT INTO purchase_attempt_claims
(attempt_id, task_id, authorization_id, claimed_by_device_id, session_id, claim_generation,
task_version, task_title, authorization_task_version, goods_id, sku_color, sku_size, quantity,
total_price_cap, authorization_expires_at, claim_nonce, claim_token_sha256,
lease_expires_at, claimed_at, closed_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, NULL)`,
attemptID, candidate.TaskID, candidate.AuthorizationID, deviceID, command.SessionID, generation,
candidate.TaskVersion+1, candidate.Title, candidate.TaskVersion, candidate.GoodsID,
candidate.SKUColor, candidate.SKUSize, candidate.Quantity, candidate.TotalPriceCap,
candidate.AuthorizationExpiresText, nonce, storedTokenHash, leaseText, nowText); err != nil {
return ClaimResponse{}, false, err
}
if err := insertClaimRequest(writeCtx, transaction, command.ClaimRequestID, deviceID, command.SessionID, "CLAIMED", attemptID, leaseText, "", now); err != nil {
return ClaimResponse{}, false, err
}
response := ClaimResponse{
Task: ClaimedTask{ID: candidate.TaskID, Version: candidate.TaskVersion + 1, Title: candidate.Title,
ProductURL: productURL(candidate.GoodsID), GoodsID: candidate.GoodsID, SKUColor: candidate.SKUColor,
SKUSize: candidate.SKUSize, Quantity: candidate.Quantity, MaxTotalPrice: candidate.TotalPriceCap},
Authorization: ClaimedAuthorization{ID: candidate.AuthorizationID, TaskVersion: candidate.TaskVersion, ExpiresAt: candidate.AuthorizationExpiresText},
Attempt: ClaimedAttempt{ID: attemptID, ClaimToken: hex.EncodeToString(token), ClaimGeneration: generation, LeaseExpiresAt: leaseText},
}
if err := transaction.Commit(); err != nil {
return ClaimResponse{}, false, err
}
return response, true, nil
}
func (store *Store) Renew(ctx context.Context, deviceID string, command RenewCommand) (RenewResponse, error) {
providedToken, tokenOK := decodeToken(command.ClaimToken)
if !deviceauth.ValidDeviceID(deviceID) || !validUUID(command.TaskID) || !validUUID(command.RenewRequestID) ||
!validUUID(command.SessionID) || !validUUID(command.AttemptID) || command.ClaimGeneration <= 0 ||
!tokenOK || !validCanonicalTime(command.ExpectedLeaseExpiresAt) {
return RenewResponse{}, ErrInvalid
}
providedHash := tokenHash(providedToken)
writeCtx, cancel := context.WithTimeout(ctx, writeTimeout)
defer cancel()
select {
case store.writeGate <- struct{}{}:
defer func() { <-store.writeGate }()
case <-writeCtx.Done():
return RenewResponse{}, writeCtx.Err()
}
transaction, err := store.database.BeginTx(writeCtx, nil)
if err != nil {
return RenewResponse{}, err
}
defer transaction.Rollback()
// As in ClaimNext, this is deliberately the first database statement in the transaction.
if store.beforeLinearization != nil {
store.beforeLinearization()
}
active, err := transaction.ExecContext(writeCtx, `UPDATE device_credentials SET status = status
WHERE device_id = ? AND status = 'ACTIVE' AND revoked_at IS NULL`, deviceID)
if err != nil {
return RenewResponse{}, err
}
if ok, err := exactlyOne(active); err != nil {
return RenewResponse{}, err
} else if !ok {
return RenewResponse{}, ErrDeviceInactive
}
if store.afterLinearization != nil {
store.afterLinearization()
}
renewal, found, err := findRenewal(writeCtx, transaction, command.RenewRequestID)
if err != nil {
return RenewResponse{}, err
}
if found {
if renewal.TaskID != command.TaskID || renewal.AttemptID != command.AttemptID || renewal.DeviceID != deviceID ||
renewal.SessionID != command.SessionID || renewal.Generation != command.ClaimGeneration ||
renewal.ExpectedLeaseExpiresAt != command.ExpectedLeaseExpiresAt || !matchingHash(renewal.TokenHash, providedHash) {
return RenewResponse{}, ErrIdempotencyConflict
}
response := RenewResponse{TaskID: renewal.TaskID, AttemptID: renewal.AttemptID, ClaimGeneration: renewal.Generation, LeaseExpiresAt: renewal.LeaseExpiresAt}
if err := transaction.Commit(); err != nil {
return RenewResponse{}, err
}
return response, nil
}
now, err := store.serverNow()
if err != nil {
return RenewResponse{}, err
}
record, found, err := store.loadClaimByAttempt(writeCtx, transaction, command.AttemptID)
if err != nil {
return RenewResponse{}, err
}
if !found || record.TaskID != command.TaskID || record.DeviceID != deviceID || record.SessionID != command.SessionID ||
record.Generation != command.ClaimGeneration || !matchingHash(record.TokenHash, providedHash) {
return RenewResponse{}, ErrNotCurrent
}
stateCurrent := record.ClosedAt == "" && record.LeaseExpiresAt.After(now) && record.AuthorizationExpiresAt.After(now) &&
record.CurrentAuthorizationExpiresAt.After(now) && record.AuthorizationStatus == "CLAIMED" &&
record.authorizationConsistent() && record.recoverableBusinessState()
if !stateCurrent || record.LeaseExpiresText != command.ExpectedLeaseExpiresAt {
return RenewResponse{}, ErrNotCurrent
}
leaseExpires := now.Add(store.leaseTTL)
if record.AuthorizationExpiresAt.Before(leaseExpires) {
leaseExpires = record.AuthorizationExpiresAt
}
leaseText := formatTime(leaseExpires)
updated, err := transaction.ExecContext(writeCtx, `UPDATE purchase_attempt_claims SET lease_expires_at = ?
WHERE attempt_id = ? AND lease_expires_at = ? AND closed_at IS NULL`, leaseText, command.AttemptID, command.ExpectedLeaseExpiresAt)
if err != nil {
return RenewResponse{}, err
}
if ok, err := exactlyOne(updated); err != nil || !ok {
if err == nil {
err = ErrNotCurrent
}
return RenewResponse{}, err
}
if _, err := transaction.ExecContext(writeCtx, `INSERT INTO purchase_attempt_lease_renewals
(renew_request_id, task_id, attempt_id, device_id, session_id, claim_generation,
claim_token_sha256, expected_lease_expires_at, lease_expires_at, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?)`,
command.RenewRequestID, command.TaskID, command.AttemptID, deviceID, command.SessionID,
command.ClaimGeneration, record.TokenHash, command.ExpectedLeaseExpiresAt, leaseText, formatTime(now)); err != nil {
return RenewResponse{}, err
}
response := RenewResponse{TaskID: command.TaskID, AttemptID: command.AttemptID, ClaimGeneration: command.ClaimGeneration, LeaseExpiresAt: leaseText}
if err := transaction.Commit(); err != nil {
return RenewResponse{}, err
}
return response, nil
}
type claimRequestRecord struct {
DeviceID, SessionID, Outcome, AttemptID, ResponseLeaseExpiresAt string
}
func findClaimRequest(ctx context.Context, transaction *sql.Tx, requestID string) (claimRequestRecord, bool, error) {
var record claimRequestRecord
var attemptID, responseLease sql.NullString
err := transaction.QueryRowContext(ctx, `SELECT device_id, session_id, outcome, attempt_id, response_lease_expires_at
FROM task_claim_requests WHERE claim_request_id = ?`, requestID).
Scan(&record.DeviceID, &record.SessionID, &record.Outcome, &attemptID, &responseLease)
if errors.Is(err, sql.ErrNoRows) {
return claimRequestRecord{}, false, nil
}
if err != nil {
return claimRequestRecord{}, false, err
}
record.AttemptID, record.ResponseLeaseExpiresAt = attemptID.String, responseLease.String
return record, true, nil
}
func insertClaimRequest(ctx context.Context, transaction *sql.Tx, requestID, deviceID, sessionID, outcome, attemptID, responseLease, errorCode string, now time.Time) error {
var attempt, lease, code any
if attemptID != "" {
attempt = attemptID
}
if responseLease != "" {
lease = responseLease
}
if errorCode != "" {
code = errorCode
}
_, err := transaction.ExecContext(ctx, `INSERT INTO task_claim_requests
(claim_request_id, device_id, session_id, outcome, attempt_id, response_lease_expires_at, error_code, created_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)`, requestID, deviceID, sessionID, outcome, attempt, lease, code, formatTime(now))
return err
}
type renewalRecord struct {
TaskID, AttemptID, DeviceID, SessionID string
Generation int
TokenHash []byte
ExpectedLeaseExpiresAt, LeaseExpiresAt string
}
func findRenewal(ctx context.Context, transaction *sql.Tx, requestID string) (renewalRecord, bool, error) {
var record renewalRecord
err := transaction.QueryRowContext(ctx, `SELECT task_id, attempt_id, device_id, session_id,
claim_generation, claim_token_sha256, expected_lease_expires_at, lease_expires_at
FROM purchase_attempt_lease_renewals WHERE renew_request_id = ?`, requestID).
Scan(&record.TaskID, &record.AttemptID, &record.DeviceID, &record.SessionID, &record.Generation,
&record.TokenHash, &record.ExpectedLeaseExpiresAt, &record.LeaseExpiresAt)
if errors.Is(err, sql.ErrNoRows) {
return renewalRecord{}, false, nil
}
return record, err == nil, err
}
type claimRecord struct {
AttemptID, TaskID, AuthorizationID, DeviceID, SessionID string
Generation, TaskVersion, CurrentTaskVersion int
TaskTitle string
Nonce, TokenHash []byte
LeaseExpiresText, ClaimedAt, ClosedAt string
LeaseExpiresAt time.Time
AuthorizationTaskVersion int
GoodsID, SKUColor, SKUSize, TotalPriceCap string
Quantity int
AuthorizationExpiresText, AuthorizationStatus string
AuthorizationExpiresAt time.Time
CurrentAuthorizationTaskVersion int
CurrentGoodsID, CurrentSKUColor, CurrentSKUSize string
CurrentQuantity int
CurrentTotalPriceCap, CurrentAuthorizationExpiresText string
CurrentAuthorizationExpiresAt time.Time
AttemptStatus, TaskStatus string
CurrentTaskTitle, CurrentTaskGoodsID string
CurrentTaskSKUColor, CurrentTaskSKUSize string
CurrentTaskQuantity int
CurrentTaskMaxTotalPrice string
CurrentAttemptGeneration int
}
const claimSelect = `SELECT claims.attempt_id, claims.task_id, claims.authorization_id,
claims.claimed_by_device_id, claims.session_id, claims.claim_generation, claims.task_version,
claims.task_title, claims.authorization_task_version, claims.goods_id, claims.sku_color,
claims.sku_size, claims.quantity, claims.total_price_cap, claims.authorization_expires_at,
claims.claim_nonce, claims.claim_token_sha256, claims.lease_expires_at,
claims.claimed_at, claims.closed_at, authorizations.task_version, authorizations.goods_id,
authorizations.sku_color, authorizations.sku_size, authorizations.quantity,
authorizations.total_price_cap, authorizations.expires_at, authorizations.status,
attempts.claim_generation, attempts.status, tasks.status, tasks.version, tasks.title, tasks.goods_id,
tasks.sku_color, tasks.sku_size, tasks.quantity, tasks.max_total_price
FROM purchase_attempt_claims AS claims
JOIN order_authorizations AS authorizations
ON authorizations.task_id = claims.task_id AND authorizations.id = claims.authorization_id
JOIN purchase_attempts AS attempts ON attempts.id = claims.attempt_id
JOIN tasks ON tasks.id = claims.task_id `
func (store *Store) loadOpenClaimByDevice(ctx context.Context, transaction *sql.Tx, deviceID string) (claimRecord, bool, error) {
return store.scanClaim(transaction.QueryRowContext(ctx, claimSelect+`WHERE claims.claimed_by_device_id = ? AND claims.closed_at IS NULL`, deviceID))
}
func (store *Store) loadClaimByAttempt(ctx context.Context, transaction *sql.Tx, attemptID string) (claimRecord, bool, error) {
return store.scanClaim(transaction.QueryRowContext(ctx, claimSelect+`WHERE claims.attempt_id = ?`, attemptID))
}
type rowScanner interface{ Scan(...any) error }
func (store *Store) scanClaim(row rowScanner) (claimRecord, bool, error) {
var record claimRecord
var closed sql.NullString
err := row.Scan(&record.AttemptID, &record.TaskID, &record.AuthorizationID, &record.DeviceID,
&record.SessionID, &record.Generation, &record.TaskVersion, &record.TaskTitle,
&record.AuthorizationTaskVersion, &record.GoodsID, &record.SKUColor, &record.SKUSize,
&record.Quantity, &record.TotalPriceCap, &record.AuthorizationExpiresText,
&record.Nonce, &record.TokenHash, &record.LeaseExpiresText, &record.ClaimedAt, &closed,
&record.CurrentAuthorizationTaskVersion, &record.CurrentGoodsID, &record.CurrentSKUColor,
&record.CurrentSKUSize, &record.CurrentQuantity, &record.CurrentTotalPriceCap,
&record.CurrentAuthorizationExpiresText,
&record.AuthorizationStatus, &record.CurrentAttemptGeneration, &record.AttemptStatus,
&record.TaskStatus, &record.CurrentTaskVersion,
&record.CurrentTaskTitle, &record.CurrentTaskGoodsID, &record.CurrentTaskSKUColor,
&record.CurrentTaskSKUSize, &record.CurrentTaskQuantity, &record.CurrentTaskMaxTotalPrice)
if errors.Is(err, sql.ErrNoRows) {
return claimRecord{}, false, nil
}
if err != nil {
return claimRecord{}, false, err
}
record.ClosedAt = closed.String
if !validUUID(record.AttemptID) || !validUUID(record.TaskID) || !validUUID(record.AuthorizationID) ||
!deviceauth.ValidDeviceID(record.DeviceID) || !validUUID(record.SessionID) || record.Generation <= 0 ||
record.CurrentAttemptGeneration != record.Generation ||
record.TaskVersion <= 0 || record.AuthorizationTaskVersion <= 0 ||
!taskmodel.ValidTaskWireFields(record.TaskTitle, record.GoodsID, record.SKUColor, record.SKUSize, record.TotalPriceCap) ||
record.Quantity <= 0 || len(record.Nonce) != sha256.Size || len(record.TokenHash) != sha256.Size {
return claimRecord{}, false, errors.New("stored task claim metadata is invalid")
}
record.LeaseExpiresAt, err = parseCanonicalTime(record.LeaseExpiresText)
if err != nil {
return claimRecord{}, false, errors.New("stored task claim lease is invalid")
}
record.AuthorizationExpiresAt, err = parseCanonicalTime(record.AuthorizationExpiresText)
if err != nil {
return claimRecord{}, false, errors.New("stored authorization expiry is invalid")
}
record.CurrentAuthorizationExpiresAt, err = parseCanonicalTime(record.CurrentAuthorizationExpiresText)
if err != nil {
return claimRecord{}, false, errors.New("current authorization expiry is invalid")
}
derived := deriveToken(store.secret, record.DeviceID, record.TaskID, record.AuthorizationID, record.AttemptID, record.Generation, record.Nonce)
if !matchingHash(tokenHash(derived), record.TokenHash) {
return claimRecord{}, false, errors.New("task claim secret does not match stored claim")
}
return record, true, nil
}
func (record claimRecord) authorizationConsistent() bool {
return taskmodel.ValidAuthorizationFields(record.CurrentGoodsID, record.CurrentSKUColor,
record.CurrentSKUSize, record.CurrentTotalPriceCap) &&
taskmodel.ValidTaskWireFields(record.CurrentTaskTitle, record.CurrentTaskGoodsID,
record.CurrentTaskSKUColor, record.CurrentTaskSKUSize, record.CurrentTaskMaxTotalPrice) &&
record.AuthorizationTaskVersion == record.CurrentAuthorizationTaskVersion &&
record.GoodsID == record.CurrentGoodsID && record.SKUColor == record.CurrentSKUColor &&
record.SKUSize == record.CurrentSKUSize && record.Quantity == record.CurrentQuantity &&
record.TotalPriceCap == record.CurrentTotalPriceCap &&
record.AuthorizationExpiresText == record.CurrentAuthorizationExpiresText &&
record.TaskTitle == record.CurrentTaskTitle && record.GoodsID == record.CurrentTaskGoodsID &&
record.SKUColor == record.CurrentTaskSKUColor && record.SKUSize == record.CurrentTaskSKUSize &&
record.Quantity == record.CurrentTaskQuantity && record.TotalPriceCap == record.CurrentTaskMaxTotalPrice
}
func (record claimRecord) recoverableBusinessState() bool {
if record.TaskStatus == "CLAIMED" && record.AttemptStatus == "CLAIMED" {
return record.CurrentTaskVersion == record.TaskVersion
}
// A later server task may advance this same attempt to ORDERING. A valid lease and identical
// ownership recover that attempt; claim-next still cannot select another task.
return record.TaskStatus == "ORDERING" && record.AttemptStatus == "ORDERING" &&
record.TaskVersion < math.MaxInt && record.CurrentTaskVersion == record.TaskVersion+1
}
func (store *Store) responseFor(record claimRecord, responseLease string) (ClaimResponse, error) {
// Exact idempotent replay is allowed to ignore later source-row drift, but the
// immutable response snapshot itself must still satisfy the current wire bounds.
if !validCanonicalTime(responseLease) ||
!taskmodel.ValidTaskWireFields(record.TaskTitle, record.GoodsID, record.SKUColor, record.SKUSize, record.TotalPriceCap) ||
record.Quantity <= 0 {
return ClaimResponse{}, errors.New("stored claim response snapshot is invalid")
}
token := deriveToken(store.secret, record.DeviceID, record.TaskID, record.AuthorizationID, record.AttemptID, record.Generation, record.Nonce)
return ClaimResponse{
Task: ClaimedTask{ID: record.TaskID, Version: record.TaskVersion, Title: record.TaskTitle,
ProductURL: productURL(record.GoodsID), GoodsID: record.GoodsID, SKUColor: record.SKUColor,
SKUSize: record.SKUSize, Quantity: record.Quantity, MaxTotalPrice: record.TotalPriceCap},
Authorization: ClaimedAuthorization{ID: record.AuthorizationID, TaskVersion: record.AuthorizationTaskVersion, ExpiresAt: record.AuthorizationExpiresText},
Attempt: ClaimedAttempt{ID: record.AttemptID, ClaimToken: hex.EncodeToString(token), ClaimGeneration: record.Generation, LeaseExpiresAt: responseLease},
}, nil
}
type candidate struct {
AuthorizationID, TaskID, Title, GoodsID, SKUColor, SKUSize, TotalPriceCap string
TaskVersion, Quantity int
AuthorizationTaskVersion, AuthorizationQuantity int
AuthorizationGoodsID, AuthorizationSKUColor, AuthorizationSKUSize string
AuthorizationTotalPriceCap string
AuthorizationExpiresText string
AuthorizationExpiresAt time.Time
}
func findCandidate(ctx context.Context, transaction *sql.Tx, now time.Time) (candidate, bool, error) {
rows, err := transaction.QueryContext(ctx, `SELECT authorizations.id, tasks.id, tasks.version,
tasks.title, tasks.goods_id, tasks.sku_color, tasks.sku_size, tasks.quantity,
tasks.max_total_price, authorizations.task_version, authorizations.goods_id,
authorizations.sku_color, authorizations.sku_size, authorizations.quantity,
authorizations.total_price_cap, authorizations.expires_at
FROM order_authorizations AS authorizations
JOIN tasks ON tasks.id = authorizations.task_id
WHERE authorizations.status = 'ACTIVE' AND tasks.status = 'PENDING'
ORDER BY authorizations.created_at, authorizations.rowid, authorizations.id`)
if err != nil {
return candidate{}, false, err
}
defer rows.Close()
for rows.Next() {
var item candidate
if err := rows.Scan(&item.AuthorizationID, &item.TaskID, &item.TaskVersion, &item.Title,
&item.GoodsID, &item.SKUColor, &item.SKUSize, &item.Quantity, &item.TotalPriceCap,
&item.AuthorizationTaskVersion, &item.AuthorizationGoodsID, &item.AuthorizationSKUColor,
&item.AuthorizationSKUSize, &item.AuthorizationQuantity, &item.AuthorizationTotalPriceCap,
&item.AuthorizationExpiresText); err != nil {
return candidate{}, false, err
}
item.AuthorizationExpiresAt, err = parseCanonicalTime(item.AuthorizationExpiresText)
if err != nil {
return candidate{}, false, errors.New("stored authorization expiry is invalid")
}
if !validCandidate(item) {
return candidate{}, false, errors.New("stored claim candidate is invalid")
}
if !candidateSnapshotMatches(item) {
continue
}
if item.AuthorizationExpiresAt.After(now) {
if err := rows.Close(); err != nil {
return candidate{}, false, err
}
return item, true, nil
}
}
if err := rows.Err(); err != nil {
return candidate{}, false, err
}
return candidate{}, false, nil
}
func validCandidate(item candidate) bool {
return validUUID(item.AuthorizationID) && validUUID(item.TaskID) && item.TaskVersion > 0 && item.TaskVersion < math.MaxInt &&
taskmodel.ValidTaskWireFields(item.Title, item.GoodsID, item.SKUColor, item.SKUSize, item.TotalPriceCap) &&
item.Quantity > 0 && item.AuthorizationTaskVersion > 0 && item.AuthorizationTaskVersion < math.MaxInt &&
taskmodel.ValidAuthorizationFields(item.AuthorizationGoodsID, item.AuthorizationSKUColor,
item.AuthorizationSKUSize, item.AuthorizationTotalPriceCap) && item.AuthorizationQuantity > 0
}
func candidateSnapshotMatches(item candidate) bool {
return item.AuthorizationTaskVersion == item.TaskVersion && item.AuthorizationGoodsID == item.GoodsID &&
item.AuthorizationSKUColor == item.SKUColor && item.AuthorizationSKUSize == item.SKUSize &&
item.AuthorizationQuantity == item.Quantity && item.AuthorizationTotalPriceCap == item.TotalPriceCap
}
func validAttemptStatus(value sql.NullString) bool {
return value.Valid && oneOf(value.String, "CLAIMED", "ORDERING", "FAILED", "FENCED", "ABANDONED")
}
func validAuthorizationStatus(value sql.NullString) bool {
return value.Valid && oneOf(value.String, "ACTIVE", "CLAIMED", "FENCED", "CONSUMED", "EXPIRED", "ABANDONED")
}
func validTaskStatus(value sql.NullString) bool {
return value.Valid && oneOf(value.String, "DRAFT", "PENDING", "CLAIMED", "ORDERING", "NEEDS_MANUAL",
"WAITING_PAYMENT", "RECONCILIATION_REQUIRED", "SUCCEEDED", "FAILED", "CANCELED")
}
func oneOf(value string, allowed ...string) bool {
for _, item := range allowed {
if value == item {
return true
}
}
return false
}
func nextGeneration(ctx context.Context, transaction *sql.Tx, taskID string) (int, error) {
var maximum int64
if err := transaction.QueryRowContext(ctx, `SELECT COALESCE(MAX(claim_generation), 0) FROM purchase_attempts WHERE task_id = ?`, taskID).Scan(&maximum); err != nil {
return 0, err
}
if maximum < 0 || maximum >= int64(math.MaxInt) {
return 0, errors.New("task claim generation is exhausted")
}
return int(maximum) + 1, nil
}
func (store *Store) serverNow() (time.Time, error) {
now := store.now().UTC()
if now.IsZero() {
return time.Time{}, errors.New("task claim clock is invalid")
}
return now, nil
}
func (store *Store) randomBytes(size int) ([]byte, error) {
value := make([]byte, size)
store.randomMu.Lock()
_, err := io.ReadFull(store.random, value)
store.randomMu.Unlock()
if err != nil {
return nil, fmt.Errorf("generate task claim randomness: %w", err)
}
return value, nil
}
func (store *Store) newUUID() (string, error) {
value, err := store.randomBytes(16)
if err != nil {
return "", err
}
value[6] = (value[6] & 0x0f) | 0x40
value[8] = (value[8] & 0x3f) | 0x80
encoded := hex.EncodeToString(value)
return encoded[:8] + "-" + encoded[8:12] + "-" + encoded[12:16] + "-" + encoded[16:20] + "-" + encoded[20:], nil
}
func exactlyOne(result sql.Result) (bool, error) {
rows, err := result.RowsAffected()
return rows == 1, err
}
func formatTime(value time.Time) string { return value.UTC().Format(time.RFC3339Nano) }
func parseCanonicalTime(value string) (time.Time, error) {
if !strings.HasSuffix(value, "Z") || strings.TrimSpace(value) != value {
return time.Time{}, ErrInvalid
}
parsed, err := time.Parse(time.RFC3339Nano, value)
if err != nil || parsed.Location() != time.UTC || formatTime(parsed) != value {
return time.Time{}, ErrInvalid
}
return parsed, nil
}
func validCanonicalTime(value string) bool {
_, err := parseCanonicalTime(value)
return err == nil
}
func validUUID(value string) bool {
if len(value) != 36 {
return false
}
for index, character := range value {
if index == 8 || index == 13 || index == 18 || index == 23 {
if character != '-' {
return false
}
continue
}
if !(character >= '0' && character <= '9' || character >= 'a' && character <= 'f') {
return false
}
}
return value[14] == '4' && (value[19] == '8' || value[19] == '9' || value[19] == 'a' || value[19] == 'b')
}
func productURL(goodsID string) string {
return "https://mobile.yangkeduo.com/goods.html?goods_id=" + goodsID
}
File diff suppressed because it is too large Load Diff
+52
View File
@@ -0,0 +1,52 @@
package taskclaim
import (
"crypto/hmac"
"crypto/sha256"
"crypto/subtle"
"encoding/binary"
"encoding/hex"
"hash"
)
const tokenDomain = "cmbuyer/task-claim-token/v1\x00"
func deriveToken(secret []byte, deviceID, taskID, authorizationID, attemptID string, generation int, nonce []byte) []byte {
mac := hmac.New(sha256.New, secret)
_, _ = mac.Write([]byte(tokenDomain))
writeTokenField(mac, deviceID)
writeTokenField(mac, taskID)
writeTokenField(mac, authorizationID)
writeTokenField(mac, attemptID)
var number [8]byte
binary.BigEndian.PutUint64(number[:], uint64(generation))
_, _ = mac.Write(number[:])
writeTokenBytes(mac, nonce)
return mac.Sum(nil)
}
func writeTokenField(writer hash.Hash, value string) { writeTokenBytes(writer, []byte(value)) }
func writeTokenBytes(writer hash.Hash, value []byte) {
var size [4]byte
binary.BigEndian.PutUint32(size[:], uint32(len(value)))
_, _ = writer.Write(size[:])
_, _ = writer.Write(value)
}
func tokenHash(token []byte) []byte {
sum := sha256.Sum256(token)
return sum[:]
}
func matchingHash(left, right []byte) bool {
return len(left) == sha256.Size && len(right) == sha256.Size && subtle.ConstantTimeCompare(left, right) == 1
}
func decodeToken(value string) ([]byte, bool) {
if len(value) != sha256.Size*2 {
return nil, false
}
decoded, err := hex.DecodeString(value)
return decoded, err == nil && hex.EncodeToString(decoded) == value
}
+100
View File
@@ -0,0 +1,100 @@
// Package taskclaim owns the atomic task-claim and lease-renewal boundary.
// A claim token proves only ownership of one attempt; it is never permission to submit an order.
package taskclaim
import (
"context"
"errors"
"math"
taskmodel "cmbuyer/admin/internal/tasks"
)
var (
ErrInvalid = errors.New("invalid task claim request")
ErrIdempotencyConflict = errors.New("task claim idempotency conflict")
ErrRequiresManual = errors.New("task claim requires manual recovery")
ErrNotCurrent = errors.New("task claim is not current")
ErrDeviceInactive = errors.New("task claim device is inactive")
)
type ClaimCommand struct {
SessionID string `json:"session_id"`
ClaimRequestID string `json:"claim_request_id"`
}
type RenewCommand struct {
TaskID string `json:"-"`
RenewRequestID string `json:"renew_request_id"`
SessionID string `json:"session_id"`
AttemptID string `json:"attempt_id"`
ClaimGeneration int `json:"claim_generation"`
ClaimToken string `json:"claim_token"`
ExpectedLeaseExpiresAt string `json:"expected_lease_expires_at"`
}
type ClaimedTask struct {
ID string `json:"id"`
Version int `json:"version"`
Title string `json:"title"`
ProductURL string `json:"product_url"`
GoodsID string `json:"goods_id"`
SKUColor string `json:"sku_color"`
SKUSize string `json:"sku_size"`
Quantity int `json:"quantity"`
MaxTotalPrice string `json:"max_total_price"`
}
type ClaimedAuthorization struct {
ID string `json:"id"`
TaskVersion int `json:"task_version"`
ExpiresAt string `json:"expires_at"`
}
type ClaimedAttempt struct {
ID string `json:"id"`
ClaimToken string `json:"claim_token"`
ClaimGeneration int `json:"claim_generation"`
LeaseExpiresAt string `json:"lease_expires_at"`
}
type ClaimResponse struct {
Task ClaimedTask `json:"task"`
Authorization ClaimedAuthorization `json:"authorization"`
Attempt ClaimedAttempt `json:"attempt"`
}
// ValidClaimResponse closes the service-to-HTTP boundary as well as the SQLite
// boundary. A fake or future Service implementation cannot bypass the same field
// limits enforced while creating and claiming the task.
func ValidClaimResponse(response ClaimResponse) bool {
authorizationExpires, authorizationErr := parseCanonicalTime(response.Authorization.ExpiresAt)
leaseExpires, leaseErr := parseCanonicalTime(response.Attempt.LeaseExpiresAt)
return validUUID(response.Task.ID) && response.Task.Version > 0 &&
response.Authorization.TaskVersion > 0 && response.Authorization.TaskVersion < math.MaxInt &&
response.Task.Version == response.Authorization.TaskVersion+1 &&
taskmodel.ValidTaskWireFields(response.Task.Title, response.Task.GoodsID,
response.Task.SKUColor, response.Task.SKUSize, response.Task.MaxTotalPrice) &&
response.Task.ProductURL == productURL(response.Task.GoodsID) && response.Task.Quantity > 0 &&
validUUID(response.Authorization.ID) && authorizationErr == nil &&
validUUID(response.Attempt.ID) && response.Attempt.ClaimGeneration > 0 &&
len(response.Attempt.ClaimToken) == 64 && tokenTextValid(response.Attempt.ClaimToken) &&
leaseErr == nil && !leaseExpires.After(authorizationExpires)
}
func tokenTextValid(value string) bool {
_, ok := decodeToken(value)
return ok
}
type RenewResponse struct {
TaskID string `json:"task_id"`
AttemptID string `json:"attempt_id"`
ClaimGeneration int `json:"claim_generation"`
LeaseExpiresAt string `json:"lease_expires_at"`
}
type Service interface {
ClaimNext(context.Context, string, ClaimCommand) (ClaimResponse, bool, error)
Renew(context.Context, string, RenewCommand) (RenewResponse, error)
}
+60
View File
@@ -0,0 +1,60 @@
// Package taskdetail provides a read-only audit projection for one task.
package taskdetail
import (
"context"
"errors"
"time"
)
var ErrNotFound = errors.New("task detail not found")
type Store interface {
Get(context.Context, string) (Detail, error)
}
type Detail struct {
Task Task
Authorizations []Authorization
Attempts []Attempt
Submissions []Submission
Evidence []Evidence
}
type Task struct {
ID, Source, Title, GoodsID, SKUColor, SKUSize, MaxTotalPrice, Status string
Quantity, Version int
CreatedAt, UpdatedAt time.Time
}
type Authorization struct {
ID, Status, CreatedBy, TotalPriceCap string
TaskVersion int
CreatedAt, ExpiresAt time.Time
}
type Attempt struct {
ID, AuthorizationID, Status string
ClaimGeneration int
Gate1UnitPrice *string
Gate2UnitPrice *string
QuantityRead *int
ConfirmAmount *string
FailureCode *string
StartedAt time.Time
FinishedAt *time.Time
}
type Submission struct {
ID, AuthorizationID, AttemptID, Status string
Gate1UnitPrice, Gate2UnitPrice, ConfirmAmount string
QuantityRead int
CreatedAt time.Time
ResolvedAt *time.Time
}
type Evidence struct {
ID, AttemptID, Kind, PrivacyTier, SHA256, ContentType string
ByteSize, Width, Height int64
CapturedAt time.Time
}
+204
View File
@@ -0,0 +1,204 @@
package taskdetail
import (
"context"
"database/sql"
"errors"
"fmt"
"time"
)
type SQLiteStore struct{ database *sql.DB }
func NewSQLiteStore(database *sql.DB) (*SQLiteStore, error) {
if database == nil {
return nil, errors.New("task detail database is required")
}
if _, err := database.Exec("SELECT storage_key FROM evidence_assets LIMIT 1"); err != nil {
return nil, fmt.Errorf("task detail migration is not available: %w", err)
}
return &SQLiteStore{database: database}, nil
}
func (store *SQLiteStore) Get(ctx context.Context, id string) (Detail, error) {
if !validUUID(id) {
return Detail{}, ErrNotFound
}
tx, err := store.database.BeginTx(ctx, &sql.TxOptions{ReadOnly: true})
if err != nil {
return Detail{}, err
}
defer tx.Rollback()
var detail Detail
var created, updated string
err = tx.QueryRowContext(ctx, `SELECT id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price, status, version, created_at, updated_at FROM tasks WHERE id = ?`, id).Scan(
&detail.Task.ID, &detail.Task.Source, &detail.Task.Title, &detail.Task.GoodsID, &detail.Task.SKUColor, &detail.Task.SKUSize,
&detail.Task.Quantity, &detail.Task.MaxTotalPrice, &detail.Task.Status, &detail.Task.Version, &created, &updated,
)
if errors.Is(err, sql.ErrNoRows) {
return Detail{}, ErrNotFound
}
if err != nil {
return Detail{}, err
}
if detail.Task.CreatedAt, err = parseTime(created); err != nil {
return Detail{}, err
}
if detail.Task.UpdatedAt, err = parseTime(updated); err != nil {
return Detail{}, err
}
if detail.Authorizations, err = readAuthorizations(ctx, tx, id); err != nil {
return Detail{}, err
}
if detail.Attempts, err = readAttempts(ctx, tx, id); err != nil {
return Detail{}, err
}
if detail.Submissions, err = readSubmissions(ctx, tx, id); err != nil {
return Detail{}, err
}
if detail.Evidence, err = readEvidence(ctx, tx, id); err != nil {
return Detail{}, err
}
if err := tx.Commit(); err != nil {
return Detail{}, err
}
return detail, nil
}
func readAuthorizations(ctx context.Context, tx *sql.Tx, taskID string) ([]Authorization, error) {
rows, err := tx.QueryContext(ctx, `SELECT id, task_version, total_price_cap, status, created_by, created_at, expires_at FROM order_authorizations WHERE task_id = ? ORDER BY created_at DESC, id DESC`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
result := []Authorization{}
for rows.Next() {
var item Authorization
var created, expires string
if err := rows.Scan(&item.ID, &item.TaskVersion, &item.TotalPriceCap, &item.Status, &item.CreatedBy, &created, &expires); err != nil {
return nil, err
}
if item.CreatedAt, err = parseTime(created); err != nil {
return nil, err
}
if item.ExpiresAt, err = parseTime(expires); err != nil {
return nil, err
}
result = append(result, item)
}
return result, rows.Err()
}
func readAttempts(ctx context.Context, tx *sql.Tx, taskID string) ([]Attempt, error) {
rows, err := tx.QueryContext(ctx, `SELECT id, authorization_id, claim_generation, status, gate1_unit_price, gate2_unit_price, quantity_read, confirm_amount, failure_code, started_at, finished_at FROM purchase_attempts WHERE task_id = ? ORDER BY started_at DESC, id DESC`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
result := []Attempt{}
for rows.Next() {
var item Attempt
var gate1, gate2, confirm, failure, finished sql.NullString
var quantity sql.NullInt64
var started string
if err := rows.Scan(&item.ID, &item.AuthorizationID, &item.ClaimGeneration, &item.Status, &gate1, &gate2, &quantity, &confirm, &failure, &started, &finished); err != nil {
return nil, err
}
item.Gate1UnitPrice, item.Gate2UnitPrice, item.ConfirmAmount, item.FailureCode = stringPointer(gate1), stringPointer(gate2), stringPointer(confirm), stringPointer(failure)
if quantity.Valid {
value := int(quantity.Int64)
item.QuantityRead = &value
}
if item.StartedAt, err = parseTime(started); err != nil {
return nil, err
}
if finished.Valid {
value, parseErr := parseTime(finished.String)
if parseErr != nil {
return nil, parseErr
}
item.FinishedAt = &value
}
result = append(result, item)
}
return result, rows.Err()
}
func readSubmissions(ctx context.Context, tx *sql.Tx, taskID string) ([]Submission, error) {
rows, err := tx.QueryContext(ctx, `SELECT id, authorization_id, attempt_id, status, gate1_unit_price, gate2_unit_price, quantity_read, confirm_amount, created_at, resolved_at FROM order_submissions WHERE task_id = ? ORDER BY created_at DESC, id DESC`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
result := []Submission{}
for rows.Next() {
var item Submission
var created string
var resolved sql.NullString
if err := rows.Scan(&item.ID, &item.AuthorizationID, &item.AttemptID, &item.Status, &item.Gate1UnitPrice, &item.Gate2UnitPrice, &item.QuantityRead, &item.ConfirmAmount, &created, &resolved); err != nil {
return nil, err
}
if item.CreatedAt, err = parseTime(created); err != nil {
return nil, err
}
if resolved.Valid {
value, parseErr := parseTime(resolved.String)
if parseErr != nil {
return nil, parseErr
}
item.ResolvedAt = &value
}
result = append(result, item)
}
return result, rows.Err()
}
func readEvidence(ctx context.Context, tx *sql.Tx, taskID string) ([]Evidence, error) {
rows, err := tx.QueryContext(ctx, `SELECT id, attempt_id, kind, privacy_tier, sha256, byte_size, content_type, width_px, height_px, captured_at FROM evidence_assets WHERE task_id = ? ORDER BY captured_at, created_at, id`, taskID)
if err != nil {
return nil, err
}
defer rows.Close()
result := []Evidence{}
for rows.Next() {
var item Evidence
var captured string
if err := rows.Scan(&item.ID, &item.AttemptID, &item.Kind, &item.PrivacyTier, &item.SHA256, &item.ByteSize, &item.ContentType, &item.Width, &item.Height, &captured); err != nil {
return nil, err
}
if item.CapturedAt, err = parseTime(captured); err != nil {
return nil, err
}
result = append(result, item)
}
return result, rows.Err()
}
func parseTime(value string) (time.Time, error) { return time.Parse(time.RFC3339Nano, value) }
func stringPointer(value sql.NullString) *string {
if !value.Valid {
return nil
}
copy := value.String
return &copy
}
func validUUID(value string) bool {
if len(value) != 36 {
return false
}
for index, character := range value {
if index == 8 || index == 13 || index == 18 || index == 23 {
if character != '-' {
return false
}
continue
}
if !(character >= '0' && character <= '9' || character >= 'a' && character <= 'f') {
return false
}
}
return value[14] == '4' && (value[19] == '8' || value[19] == '9' || value[19] == 'a' || value[19] == 'b')
}
+97
View File
@@ -0,0 +1,97 @@
package taskdetail
import (
"context"
"database/sql"
"errors"
"path/filepath"
"runtime"
"strings"
"testing"
"cmbuyer/admin/internal/migrations"
"cmbuyer/admin/internal/storage/sqlite"
)
const (
detailTask = "a3c9f507-7473-4fa6-8d71-8786c34c6301"
detailAuth = "b3c9f507-7473-4fa6-8d71-8786c34c6301"
detailTry = "c3c9f507-7473-4fa6-8d71-8786c34c6301"
detailDevice = "e3c9f507-7473-4fa6-8d71-8786c34c6301"
)
func TestSQLiteStoreReturnsOnlyPersistedAuditFacts(t *testing.T) {
database := openDetailDatabase(t)
timestamp := "2026-08-04T00:00:00Z"
if _, err := database.Exec(`INSERT INTO device_credentials
(device_id,display_name,token_sha256,status,created_at,revoked_at)
VALUES (?, 'detail test device', zeroblob(32), 'ACTIVE', ?, NULL)`, detailDevice, timestamp); err != nil {
t.Fatalf("insert device: %v", err)
}
if _, err := database.Exec(`INSERT INTO tasks (id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price, status, version, created_at, updated_at) VALUES (?, 'MANUAL', 'shirt', '123', 'black', 'M', 2, '30.00', 'CLAIMED', 3, ?, ?)`, detailTask, timestamp, timestamp); err != nil {
t.Fatalf("insert task: %v", err)
}
if _, err := database.Exec(`INSERT INTO order_authorizations (id, task_id, task_version, start_key, goods_id, sku_color, sku_size, quantity, total_price_cap, status, created_by, created_at, expires_at) VALUES (?, ?, 2, 'start', '123', 'black', 'M', 2, '30.00', 'CLAIMED', 'admin', ?, ?)`, detailAuth, detailTask, timestamp, timestamp); err != nil {
t.Fatalf("insert authorization: %v", err)
}
if _, err := database.Exec(`INSERT INTO purchase_attempts (id, task_id, authorization_id, claim_generation, status, started_at) VALUES (?, ?, ?, 1, 'CLAIMED', ?)`, detailTry, detailTask, detailAuth, timestamp); err != nil {
t.Fatalf("insert attempt: %v", err)
}
if _, err := database.Exec(`INSERT INTO purchase_attempt_claims
(attempt_id,task_id,authorization_id,claimed_by_device_id,session_id,claim_generation,
task_version,task_title,authorization_task_version,goods_id,sku_color,sku_size,quantity,
total_price_cap,authorization_expires_at,claim_nonce,claim_token_sha256,lease_expires_at,claimed_at,closed_at)
VALUES (?, ?, ?, ?, 'f3c9f507-7473-4fa6-8d71-8786c34c6301', 1, 3, 'shirt',
2, '123', 'black', 'M', 2, '30.00', ?, zeroblob(32), zeroblob(32),
'2026-08-04T00:05:00Z', ?, NULL)`, detailTry, detailTask, detailAuth, detailDevice, timestamp, timestamp); err != nil {
t.Fatalf("insert claim: %v", err)
}
hash := strings.Repeat("a", 64)
if _, err := database.Exec(`INSERT INTO evidence_assets (id, upload_key, task_id, attempt_id, kind, privacy_tier, sha256, byte_size, content_type, width_px, height_px, storage_key, uploaded_by_device_id, captured_at, created_at) VALUES ('d3c9f507-7473-4fa6-8d71-8786c34c6301', 'upload', ?, ?, 'SKU_PANEL_GATE_1', 'INTERNAL_RAW', ?, 100, 'image/png', 10, 20, ?, ?, ?, ?)`, detailTask, detailTry, hash, "aa/"+hash+".png", detailDevice, timestamp, timestamp); err != nil {
t.Fatalf("insert evidence: %v", err)
}
store, err := NewSQLiteStore(database)
if err != nil {
t.Fatalf("NewSQLiteStore: %v", err)
}
detail, err := store.Get(context.Background(), detailTask)
if err != nil {
t.Fatalf("Get: %v", err)
}
if detail.Task.ID != detailTask || detail.Task.Status != "CLAIMED" || len(detail.Authorizations) != 1 || len(detail.Attempts) != 1 || len(detail.Evidence) != 1 || len(detail.Submissions) != 0 {
t.Fatalf("detail = %#v", detail)
}
if detail.Attempts[0].Gate1UnitPrice != nil || detail.Attempts[0].FailureCode != nil {
t.Fatalf("missing attempt facts were fabricated: %#v", detail.Attempts[0])
}
}
func TestSQLiteStoreFailsClosedForMalformedAndMissingIDs(t *testing.T) {
database := openDetailDatabase(t)
store, err := NewSQLiteStore(database)
if err != nil {
t.Fatalf("NewSQLiteStore: %v", err)
}
for _, id := range []string{"../database", "not-a-uuid", "a3c9f507-7473-1fa6-8d71-8786c34c6301"} {
if _, err := store.Get(context.Background(), id); !errors.Is(err, ErrNotFound) {
t.Fatalf("Get(%q) error = %v", id, err)
}
}
}
func openDetailDatabase(t *testing.T) *sql.DB {
t.Helper()
database, err := sqlite.Open(filepath.Join(t.TempDir(), "details.db"))
if err != nil {
t.Fatalf("open database: %v", err)
}
t.Cleanup(func() { _ = database.Close() })
_, file, _, ok := runtime.Caller(0)
if !ok {
t.Fatal("locate migration directory")
}
if err := migrations.Up(context.Background(), database, filepath.Join(filepath.Dir(file), "..", "..", "migrations")); err != nil {
t.Fatalf("migrate database: %v", err)
}
return database
}
+133
View File
@@ -0,0 +1,133 @@
package tasks
import (
"errors"
"math"
"math/big"
"sort"
"strings"
"time"
_ "time/tzdata"
)
const maxStartItems = 100
var (
ErrStartConflict = errors.New("purchase start conflicts with current task state")
ErrInvalidStart = errors.New("invalid purchase start request")
)
type StartPolicy struct {
AuthorizationTTL time.Duration
MaxQuantity int
MaxTotalPrice string
}
type StartItem struct {
TaskID string `json:"task_id"`
ExpectedTaskVersion int `json:"expected_task_version"`
}
type StartCommand struct {
StartKey string `json:"start_key"`
Tasks []StartItem `json:"tasks"`
}
type AuthorizedTask struct {
TaskID string `json:"task_id"`
TaskVersion int `json:"task_version"`
AuthorizationID string `json:"authorization_id"`
ExpiresAt time.Time `json:"expires_at"`
}
type StartResult struct {
StartKey string `json:"start_key"`
AuthorizedCount int `json:"authorized_count"`
Tasks []AuthorizedTask `json:"tasks"`
PaymentAutomated bool `json:"payment_automated"`
}
type TaskFilter struct{ Keyword, Status, CreatedFrom, CreatedTo string }
type TaskRow struct {
ID, Title, GoodsID, SKUColor, SKUSize, MaxTotalPrice, Status string
Quantity, Version int
CreatedAt time.Time
}
func normalizeCents(value string) (string, *big.Int, bool) {
if value == "" || len(value) > MaxMoneyASCIICharacters || strings.TrimSpace(value) != value {
return "", nil, false
}
parts := strings.Split(value, ".")
if len(parts) != 2 || len(parts[0]) == 0 || len(parts[1]) != 2 || (len(parts[0]) > 1 && parts[0][0] == '0') {
return "", nil, false
}
for _, part := range parts {
for _, ch := range part {
if ch < '0' || ch > '9' {
return "", nil, false
}
}
}
cents := new(big.Int)
if _, ok := cents.SetString(parts[0]+parts[1], 10); !ok || cents.Sign() <= 0 {
return "", nil, false
}
return value, cents, true
}
func ValidCanonicalMoney(value string) bool {
canonical, _, ok := normalizeCents(value)
return ok && canonical == value
}
func startItems(command StartCommand) ([]StartItem, error) {
if !validUUID(command.StartKey) || len(command.Tasks) == 0 || len(command.Tasks) > maxStartItems {
return nil, ErrInvalidStart
}
items := append([]StartItem(nil), command.Tasks...)
sort.Slice(items, func(i, j int) bool { return items[i].TaskID < items[j].TaskID })
for i, item := range items {
if !validUUID(item.TaskID) || item.ExpectedTaskVersion <= 0 || item.ExpectedTaskVersion == math.MaxInt || (i > 0 && item.TaskID == items[i-1].TaskID) {
return nil, ErrInvalidStart
}
}
return items, nil
}
func validTaskStatus(value string) bool {
if value == "" {
return true
}
for _, status := range []string{"DRAFT", "PENDING", "CLAIMED", "ORDERING", "NEEDS_MANUAL", "WAITING_PAYMENT", "RECONCILIATION_REQUIRED", "SUCCEEDED", "FAILED", "CANCELED"} {
if value == status {
return true
}
}
return false
}
func ShanghaiRange(from, to string) (time.Time, time.Time, error) {
if from == "" && to == "" {
return time.Time{}, time.Time{}, nil
}
location, err := time.LoadLocation("Asia/Shanghai")
if err != nil {
return time.Time{}, time.Time{}, err
}
parse := func(value string) (time.Time, error) { return time.ParseInLocation("2006-01-02", value, location) }
var start, end time.Time
if from != "" {
start, err = parse(from)
if err != nil {
return time.Time{}, time.Time{}, ErrInvalidStart
}
start = start.UTC()
}
if to != "" {
end, err = parse(to)
if err != nil {
return time.Time{}, time.Time{}, ErrInvalidStart
}
end = end.AddDate(0, 0, 1).UTC()
}
if !start.IsZero() && !end.IsZero() && !start.Before(end) {
return time.Time{}, time.Time{}, ErrInvalidStart
}
return start, end, nil
}
@@ -0,0 +1,480 @@
package tasks
import (
"context"
"database/sql"
"errors"
"fmt"
"math"
"reflect"
"sort"
"strings"
"sync"
"testing"
"time"
"cmbuyer/admin/internal/migrations"
)
var fixedStartTime = time.Date(2026, 8, 4, 9, 2, 3, 456000000, time.FixedZone("UTC+8", 8*60*60))
func TestStartPurchasesPersistsCompleteSnapshotsForOneAndHundredTasks(t *testing.T) {
for _, count := range []int{1, 100} {
t.Run(fmt.Sprintf("%d tasks", count), func(t *testing.T) {
database := migratedDatabase(t)
store := configuredStartStore(t, database)
store.now = func() time.Time { return fixedStartTime }
items := make([]StartItem, 0, count)
wantDrafts := make(map[string]Draft, count)
for index := 1; index <= count; index++ {
id := startTestUUID(index)
draft := Draft{
ID: id,
Title: fmt.Sprintf("task-%03d", index),
GoodsID: fmt.Sprintf("937122%06d", index),
SKUColor: fmt.Sprintf("color-%03d", index),
SKUSize: fmt.Sprintf("size-%03d", index),
Quantity: index%10 + 1,
MaxTotalPrice: fmt.Sprintf("%d.%02d", index+10, index%100),
}
if _, err := store.CreateDraft(context.Background(), draft); err != nil {
t.Fatalf("create draft %d: %v", index, err)
}
items = append(items, StartItem{TaskID: id, ExpectedTaskVersion: 1})
wantDrafts[id] = draft
}
sort.Slice(items, func(i, j int) bool { return items[i].TaskID > items[j].TaskID })
command := StartCommand{StartKey: startTestUUID(1001 + count), Tasks: items}
result, err := store.StartPurchases(context.Background(), command, "authenticated-admin")
if err != nil {
t.Fatalf("StartPurchases: %v", err)
}
if result.StartKey != command.StartKey || result.AuthorizedCount != count || result.PaymentAutomated || len(result.Tasks) != count {
t.Fatalf("result = %#v", result)
}
wantCreated := fixedStartTime.UTC()
wantExpires := wantCreated.Add(15 * time.Minute)
seenAuthorizationIDs := map[string]bool{}
for index, authorized := range result.Tasks {
if index > 0 && result.Tasks[index-1].TaskID >= authorized.TaskID {
t.Fatalf("result is not in canonical task order: %#v", result.Tasks)
}
if authorized.TaskVersion != 2 || !authorized.ExpiresAt.Equal(wantExpires) || !validUUID(authorized.AuthorizationID) || seenAuthorizationIDs[authorized.AuthorizationID] {
t.Fatalf("authorized task = %#v", authorized)
}
seenAuthorizationIDs[authorized.AuthorizationID] = true
want := wantDrafts[authorized.TaskID]
var taskStatus, taskUpdated, authTaskID, authStartKey, goodsID, color, size, priceCap, authStatus, createdBy, createdAt, expiresAt string
var taskVersion, authTaskVersion, quantity int
err := database.QueryRow(`
SELECT t.status,t.version,t.updated_at,
a.task_id,a.task_version,a.start_key,a.goods_id,a.sku_color,a.sku_size,a.quantity,a.total_price_cap,a.status,a.created_by,a.created_at,a.expires_at
FROM tasks t JOIN order_authorizations a ON a.task_id=t.id WHERE a.id=?`, authorized.AuthorizationID).
Scan(&taskStatus, &taskVersion, &taskUpdated, &authTaskID, &authTaskVersion, &authStartKey, &goodsID, &color, &size, &quantity, &priceCap, &authStatus, &createdBy, &createdAt, &expiresAt)
if err != nil {
t.Fatalf("read authorization snapshot: %v", err)
}
if taskStatus != "PENDING" || taskVersion != 2 || taskUpdated != wantCreated.Format(time.RFC3339Nano) ||
authTaskID != want.ID || authTaskVersion != 2 || authStartKey != command.StartKey ||
goodsID != want.GoodsID || color != want.SKUColor || size != want.SKUSize || quantity != want.Quantity || priceCap != want.MaxTotalPrice ||
authStatus != "ACTIVE" || createdBy != "authenticated-admin" || createdAt != wantCreated.Format(time.RFC3339Nano) || expiresAt != wantExpires.Format(time.RFC3339Nano) {
t.Fatalf("stored task/authorization mismatch for %s", want.ID)
}
}
var distinctCreated, distinctExpires int
if err := database.QueryRow(`SELECT COUNT(DISTINCT created_at), COUNT(DISTINCT expires_at) FROM order_authorizations WHERE start_key=?`, command.StartKey).Scan(&distinctCreated, &distinctExpires); err != nil {
t.Fatal(err)
}
if distinctCreated != 1 || distinctExpires != 1 {
t.Fatalf("batch timestamps are not shared: created=%d expires=%d", distinctCreated, distinctExpires)
}
})
}
}
func TestStartPurchasesRejectsInvalidCommandsAndPolicyWithoutWrites(t *testing.T) {
validItem := StartItem{TaskID: startTestUUID(1), ExpectedTaskVersion: 1}
hundredOne := make([]StartItem, 101)
for index := range hundredOne {
hundredOne[index] = StartItem{TaskID: startTestUUID(index + 1), ExpectedTaskVersion: 1}
}
for name, command := range map[string]StartCommand{
"invalid start key": {StartKey: "not-a-uuid", Tasks: []StartItem{validItem}},
"empty tasks": {StartKey: startTestUUID(1001)},
"over batch limit": {StartKey: startTestUUID(1001), Tasks: hundredOne},
"invalid task id": {StartKey: startTestUUID(1001), Tasks: []StartItem{{TaskID: "1", ExpectedTaskVersion: 1}}},
"duplicate task": {StartKey: startTestUUID(1001), Tasks: []StartItem{validItem, validItem}},
"zero version": {StartKey: startTestUUID(1001), Tasks: []StartItem{{TaskID: validItem.TaskID}}},
"overflow version": {StartKey: startTestUUID(1001), Tasks: []StartItem{{TaskID: validItem.TaskID, ExpectedTaskVersion: math.MaxInt}}},
} {
t.Run(name, func(t *testing.T) {
database := migratedDatabase(t)
store := configuredStartStore(t, database)
_, err := store.StartPurchases(context.Background(), command, "admin")
if !errors.Is(err, ErrInvalidStart) {
t.Fatalf("error = %v, want ErrInvalidStart", err)
}
assertAuthorizationCount(t, database, 0)
})
}
for name, mutate := range map[string]func(*SQLiteStore){
"zero ttl": func(store *SQLiteStore) { store.policy.AuthorizationTTL = 0 },
"zero quantity": func(store *SQLiteStore) { store.policy.MaxQuantity = 0 },
"bad max price": func(store *SQLiteStore) { store.policy.MaxTotalPrice = "999" },
} {
t.Run(name, func(t *testing.T) {
database := migratedDatabase(t)
store := configuredStartStore(t, database)
createStartDraft(t, store, validItem.TaskID)
mutate(store)
_, err := store.StartPurchases(context.Background(), StartCommand{StartKey: startTestUUID(1001), Tasks: []StartItem{validItem}}, "admin")
if !errors.Is(err, ErrInvalidStart) {
t.Fatalf("error = %v, want ErrInvalidStart", err)
}
assertDraftUnchanged(t, database, validItem.TaskID)
assertAuthorizationCount(t, database, 0)
})
}
database := migratedDatabase(t)
store := configuredStartStore(t, database)
createStartDraft(t, store, validItem.TaskID)
_, err := store.StartPurchases(context.Background(), StartCommand{StartKey: startTestUUID(1001), Tasks: []StartItem{validItem}}, "")
if !errors.Is(err, ErrInvalidStart) {
t.Fatalf("empty created_by error = %v", err)
}
assertDraftUnchanged(t, database, validItem.TaskID)
}
func TestStartPurchasesRejectsEveryTaskConflictWithoutAuthorization(t *testing.T) {
for name, mutate := range map[string]func(*testing.T, *SQLiteStore, string, *StartItem){
"missing": func(_ *testing.T, _ *SQLiteStore, _ string, item *StartItem) {
item.TaskID = startTestUUID(99)
},
"not draft": func(t *testing.T, store *SQLiteStore, id string, _ *StartItem) {
execTestSQL(t, store.database, `UPDATE tasks SET status='PENDING' WHERE id=?`, id)
},
"version mismatch": func(_ *testing.T, _ *SQLiteStore, _ string, item *StartItem) {
item.ExpectedTaskVersion = 2
},
"empty goods id": func(t *testing.T, store *SQLiteStore, id string, _ *StartItem) {
execTestSQL(t, store.database, `UPDATE tasks SET goods_id='' WHERE id=?`, id)
},
"nondigit goods id": func(t *testing.T, store *SQLiteStore, id string, _ *StartItem) {
execTestSQL(t, store.database, `UPDATE tasks SET goods_id='937x' WHERE id=?`, id)
},
"overlong goods id": func(t *testing.T, store *SQLiteStore, id string, _ *StartItem) {
execTestSQL(t, store.database, `UPDATE tasks SET goods_id=? WHERE id=?`, strings.Repeat("1", MaxGoodsIDCharacters+1), id)
},
"invalid utf8 title": func(t *testing.T, store *SQLiteStore, id string, _ *StartItem) {
execTestSQL(t, store.database, `UPDATE tasks SET title=? WHERE id=?`, string([]byte{0xff}), id)
},
"overlong title": func(t *testing.T, store *SQLiteStore, id string, _ *StartItem) {
execTestSQL(t, store.database, `UPDATE tasks SET title=? WHERE id=?`, strings.Repeat("😀", MaxTitleCodePoints+1), id)
},
"empty color": func(t *testing.T, store *SQLiteStore, id string, _ *StartItem) {
execTestSQL(t, store.database, `UPDATE tasks SET sku_color='' WHERE id=?`, id)
},
"overlong color": func(t *testing.T, store *SQLiteStore, id string, _ *StartItem) {
execTestSQL(t, store.database, `UPDATE tasks SET sku_color=? WHERE id=?`, strings.Repeat("色", MaxSKUTextCodePoints+1), id)
},
"empty size": func(t *testing.T, store *SQLiteStore, id string, _ *StartItem) {
execTestSQL(t, store.database, `UPDATE tasks SET sku_size='' WHERE id=?`, id)
},
"overlong size": func(t *testing.T, store *SQLiteStore, id string, _ *StartItem) {
execTestSQL(t, store.database, `UPDATE tasks SET sku_size=? WHERE id=?`, strings.Repeat("码", MaxSKUTextCodePoints+1), id)
},
"quantity over policy": func(_ *testing.T, store *SQLiteStore, _ string, _ *StartItem) {
store.policy.MaxQuantity = 1
},
"noncanonical price one decimal": func(t *testing.T, store *SQLiteStore, id string, _ *StartItem) {
execTestSQL(t, store.database, `UPDATE tasks SET max_total_price='12.8' WHERE id=?`, id)
},
"noncanonical leading zero": func(t *testing.T, store *SQLiteStore, id string, _ *StartItem) {
execTestSQL(t, store.database, `UPDATE tasks SET max_total_price='012.80' WHERE id=?`, id)
},
"overlong canonical price": func(t *testing.T, store *SQLiteStore, id string, _ *StartItem) {
execTestSQL(t, store.database, `UPDATE tasks SET max_total_price=? WHERE id=?`, strings.Repeat("1", MaxMoneyASCIICharacters-2)+".00", id)
},
"price over policy": func(_ *testing.T, store *SQLiteStore, _ string, _ *StartItem) {
store.policy.MaxTotalPrice = "12.79"
},
} {
t.Run(name, func(t *testing.T) {
database := migratedDatabase(t)
store := configuredStartStore(t, database)
id := startTestUUID(1)
createStartDraft(t, store, id)
item := StartItem{TaskID: id, ExpectedTaskVersion: 1}
mutate(t, store, id, &item)
_, err := store.StartPurchases(context.Background(), StartCommand{StartKey: startTestUUID(1001), Tasks: []StartItem{item}}, "admin")
if !errors.Is(err, ErrStartConflict) {
t.Fatalf("error = %v, want ErrStartConflict", err)
}
assertAuthorizationCount(t, database, 0)
})
}
}
func TestStartPurchasesReplayRejectsMalformedAuthorizationOrTaskSnapshot(t *testing.T) {
mutations := map[string]func(*testing.T, *sql.DB, string){
"authorization goods id": func(t *testing.T, database *sql.DB, id string) {
execTestSQL(t, database, `UPDATE order_authorizations SET goods_id=? WHERE task_id=?`, strings.Repeat("1", MaxGoodsIDCharacters+1), id)
},
"authorization color": func(t *testing.T, database *sql.DB, id string) {
execTestSQL(t, database, `UPDATE order_authorizations SET sku_color=? WHERE task_id=?`, strings.Repeat("色", MaxSKUTextCodePoints+1), id)
},
"authorization size": func(t *testing.T, database *sql.DB, id string) {
execTestSQL(t, database, `UPDATE order_authorizations SET sku_size=? WHERE task_id=?`, strings.Repeat("码", MaxSKUTextCodePoints+1), id)
},
"authorization money": func(t *testing.T, database *sql.DB, id string) {
execTestSQL(t, database, `UPDATE order_authorizations SET total_price_cap=? WHERE task_id=?`, strings.Repeat("1", MaxMoneyASCIICharacters-2)+".00", id)
},
"task title": func(t *testing.T, database *sql.DB, id string) {
execTestSQL(t, database, `UPDATE tasks SET title=? WHERE id=?`, strings.Repeat("😀", MaxTitleCodePoints+1), id)
},
}
for name, mutate := range mutations {
t.Run(name, func(t *testing.T) {
database := migratedDatabase(t)
store := configuredStartStore(t, database)
id := startTestUUID(1)
createStartDraft(t, store, id)
command := StartCommand{StartKey: startTestUUID(1001), Tasks: []StartItem{{TaskID: id, ExpectedTaskVersion: 1}}}
if _, err := store.StartPurchases(context.Background(), command, "admin"); err != nil {
t.Fatal(err)
}
mutate(t, database, id)
if _, err := store.StartPurchases(context.Background(), command, "admin"); !errors.Is(err, ErrStartConflict) {
t.Fatalf("replay error = %v, want ErrStartConflict", err)
}
assertAuthorizationCount(t, database, 1)
})
}
}
func TestStartPurchasesRollsBackWholeBatchForLateConflictAndSQLFailure(t *testing.T) {
for _, test := range []struct {
name string
breakBatch func(*testing.T, *SQLiteStore, string)
}{
{name: "late validation conflict", breakBatch: func(t *testing.T, store *SQLiteStore, secondID string) {
execTestSQL(t, store.database, `UPDATE tasks SET sku_size='' WHERE id=?`, secondID)
}},
{name: "late SQL failure", breakBatch: func(t *testing.T, store *SQLiteStore, secondID string) {
statement := fmt.Sprintf(`CREATE TRIGGER reject_second_authorization BEFORE INSERT ON order_authorizations WHEN NEW.task_id='%s' BEGIN SELECT RAISE(ABORT, 'test failure'); END`, secondID)
execTestSQL(t, store.database, statement)
}},
} {
t.Run(test.name, func(t *testing.T) {
database := migratedDatabase(t)
store := configuredStartStore(t, database)
firstID, secondID := startTestUUID(1), startTestUUID(2)
createStartDraft(t, store, firstID)
createStartDraft(t, store, secondID)
test.breakBatch(t, store, secondID)
_, err := store.StartPurchases(context.Background(), StartCommand{StartKey: startTestUUID(1001), Tasks: []StartItem{{TaskID: firstID, ExpectedTaskVersion: 1}, {TaskID: secondID, ExpectedTaskVersion: 1}}}, "admin")
if err == nil {
t.Fatal("StartPurchases unexpectedly succeeded")
}
assertDraftUnchanged(t, database, firstID)
var secondStatus string
var secondVersion int
if err := database.QueryRow(`SELECT status,version FROM tasks WHERE id=?`, secondID).Scan(&secondStatus, &secondVersion); err != nil {
t.Fatal(err)
}
if secondStatus != "DRAFT" || secondVersion != 1 {
t.Fatalf("second task = %s/v%d, want DRAFT/v1", secondStatus, secondVersion)
}
assertAuthorizationCount(t, database, 0)
})
}
}
func TestStartPurchasesReplayIsStableAndRejectsDifferentOrIncompleteSets(t *testing.T) {
database := migratedDatabase(t)
store := configuredStartStore(t, database)
firstID, secondID, thirdID := startTestUUID(1), startTestUUID(2), startTestUUID(3)
for _, id := range []string{firstID, secondID, thirdID} {
createStartDraft(t, store, id)
}
command := StartCommand{StartKey: startTestUUID(1001), Tasks: []StartItem{{TaskID: secondID, ExpectedTaskVersion: 1}, {TaskID: firstID, ExpectedTaskVersion: 1}}}
first, err := store.StartPurchases(context.Background(), command, "admin")
if err != nil {
t.Fatal(err)
}
command.Tasks[0], command.Tasks[1] = command.Tasks[1], command.Tasks[0]
replay, err := store.StartPurchases(context.Background(), command, "admin")
if err != nil || !reflect.DeepEqual(replay, first) {
t.Fatalf("replay = (%#v, %v), want %#v", replay, err, first)
}
assertAuthorizationCount(t, database, 2)
conflicting := []StartCommand{
{StartKey: command.StartKey, Tasks: command.Tasks[:1]},
{StartKey: command.StartKey, Tasks: []StartItem{{TaskID: firstID, ExpectedTaskVersion: 2}, {TaskID: secondID, ExpectedTaskVersion: 1}}},
{StartKey: command.StartKey, Tasks: []StartItem{{TaskID: firstID, ExpectedTaskVersion: 1}, {TaskID: secondID, ExpectedTaskVersion: 1}, {TaskID: thirdID, ExpectedTaskVersion: 1}}},
}
for _, changed := range conflicting {
if _, err := store.StartPurchases(context.Background(), changed, "admin"); !errors.Is(err, ErrStartConflict) {
t.Fatalf("different payload error = %v", err)
}
}
assertAuthorizationCount(t, database, 2)
assertDraftUnchanged(t, database, thirdID)
execTestSQL(t, database, `DELETE FROM order_authorizations WHERE task_id=?`, secondID)
if _, err := store.StartPurchases(context.Background(), command, "admin"); !errors.Is(err, ErrStartConflict) {
t.Fatalf("incomplete replay error = %v", err)
}
assertAuthorizationCount(t, database, 1)
}
func TestStartPurchasesConcurrentReplayAndVersionRace(t *testing.T) {
t.Run("same key replays one stable result", func(t *testing.T) {
database := migratedDatabase(t)
store := configuredStartStore(t, database)
id := startTestUUID(1)
createStartDraft(t, store, id)
command := StartCommand{StartKey: startTestUUID(1001), Tasks: []StartItem{{TaskID: id, ExpectedTaskVersion: 1}}}
const callers = 16
start := make(chan struct{})
results := make(chan StartResult, callers)
errorsChannel := make(chan error, callers)
var group sync.WaitGroup
for range callers {
group.Add(1)
go func() {
defer group.Done()
<-start
result, err := store.StartPurchases(context.Background(), command, "admin")
if err != nil {
errorsChannel <- err
return
}
results <- result
}()
}
close(start)
group.Wait()
close(results)
close(errorsChannel)
for err := range errorsChannel {
t.Fatalf("concurrent replay: %v", err)
}
var want StartResult
for result := range results {
if want.StartKey == "" {
want = result
} else if !reflect.DeepEqual(result, want) {
t.Fatalf("unstable replay: %#v != %#v", result, want)
}
}
assertAuthorizationCount(t, database, 1)
var version int
if err := database.QueryRow(`SELECT version FROM tasks WHERE id=?`, id).Scan(&version); err != nil || version != 2 {
t.Fatalf("task version = %d, err=%v", version, err)
}
})
t.Run("different keys race one expected version", func(t *testing.T) {
database := migratedDatabase(t)
store := configuredStartStore(t, database)
id := startTestUUID(1)
createStartDraft(t, store, id)
start := make(chan struct{})
errorsChannel := make(chan error, 2)
var group sync.WaitGroup
for _, key := range []string{startTestUUID(1001), startTestUUID(1002)} {
group.Add(1)
go func(startKey string) {
defer group.Done()
<-start
_, err := store.StartPurchases(context.Background(), StartCommand{StartKey: startKey, Tasks: []StartItem{{TaskID: id, ExpectedTaskVersion: 1}}}, "admin")
errorsChannel <- err
}(key)
}
close(start)
group.Wait()
close(errorsChannel)
successes, conflicts := 0, 0
for err := range errorsChannel {
switch {
case err == nil:
successes++
case errors.Is(err, ErrStartConflict):
conflicts++
default:
t.Fatalf("unexpected race error: %v", err)
}
}
if successes != 1 || conflicts != 1 {
t.Fatalf("success/conflict = %d/%d, want 1/1", successes, conflicts)
}
assertAuthorizationCount(t, database, 1)
})
}
func TestSQLiteStoreRejectsV1SchemaAtStartup(t *testing.T) {
database := openDatabase(t)
if err := migrations.Run(context.Background(), database, migrationDirectory(t), "up-by-one"); err != nil {
t.Fatalf("migrate to v1: %v", err)
}
if _, err := NewSQLiteStore(database); err == nil {
t.Fatal("NewSQLiteStore accepted the v1 two-pass schema")
}
}
func configuredStartStore(t *testing.T, database *sql.DB) *SQLiteStore {
t.Helper()
store, err := NewSQLiteStore(database)
if err != nil {
t.Fatalf("NewSQLiteStore: %v", err)
}
store.SetStartPolicy(StartPolicy{AuthorizationTTL: 15 * time.Minute, MaxQuantity: 10, MaxTotalPrice: "999.99"})
return store
}
func createStartDraft(t *testing.T, store *SQLiteStore, id string) {
t.Helper()
draft := Draft{ID: id, Title: "test", GoodsID: "937122477375", SKUColor: "黑色", SKUSize: "M", Quantity: 2, MaxTotalPrice: "12.80"}
if _, err := store.CreateDraft(context.Background(), draft); err != nil {
t.Fatalf("CreateDraft: %v", err)
}
}
func startTestUUID(number int) string {
return fmt.Sprintf("%08x-1234-4abc-a123-%012x", number, number)
}
func assertAuthorizationCount(t *testing.T, database *sql.DB, want int) {
t.Helper()
var got int
if err := database.QueryRow(`SELECT COUNT(*) FROM order_authorizations`).Scan(&got); err != nil {
t.Fatal(err)
}
if got != want {
t.Fatalf("authorization count = %d, want %d", got, want)
}
}
func assertDraftUnchanged(t *testing.T, database *sql.DB, id string) {
t.Helper()
var status string
var version int
if err := database.QueryRow(`SELECT status,version FROM tasks WHERE id=?`, id).Scan(&status, &version); err != nil {
t.Fatal(err)
}
if status != "DRAFT" || version != 1 {
t.Fatalf("task %s = %s/v%d, want DRAFT/v1", id, status, version)
}
}
func execTestSQL(t *testing.T, database *sql.DB, statement string, arguments ...any) {
t.Helper()
if _, err := database.Exec(statement, arguments...); err != nil {
t.Fatalf("execute test SQL: %v", err)
}
}
+239
View File
@@ -0,0 +1,239 @@
package tasks
import (
"context"
"database/sql"
"errors"
"strings"
"time"
"cmbuyer/admin/internal/domain"
)
// ErrInvalidFilter 表示任务筛选值无效,路由应按字段重新渲染而不是泄露内部错误。
var ErrInvalidFilter = errors.New("invalid task filter")
// SetStartPolicy is called during startup; policy is explicit because authorization limits must not be implicit defaults.
func (store *SQLiteStore) SetStartPolicy(policy StartPolicy) { store.policy = policy }
func (store *SQLiteStore) ListTasks(ctx context.Context, filter TaskFilter) ([]TaskRow, error) {
if !ValidateTaskFilter(filter).Valid() {
return nil, ErrInvalidFilter
}
from, to, err := ShanghaiRange(filter.CreatedFrom, filter.CreatedTo)
if err != nil {
return nil, ErrInvalidFilter
}
clauses, args := []string{"1=1"}, []any{}
if filter.Status != "" {
clauses = append(clauses, "status = ?")
args = append(args, filter.Status)
}
if filter.Keyword != "" {
escaped := strings.NewReplacer("\\", "\\\\", "%", "\\%", "_", "\\_").Replace(filter.Keyword)
clauses = append(clauses, "(title LIKE ? ESCAPE '\\' OR goods_id LIKE ? ESCAPE '\\')")
args = append(args, "%"+escaped+"%", "%"+escaped+"%")
}
if !from.IsZero() {
clauses = append(clauses, "julianday(created_at) >= julianday(?)")
args = append(args, from.Format(time.RFC3339Nano))
}
if !to.IsZero() {
clauses = append(clauses, "julianday(created_at) < julianday(?)")
args = append(args, to.Format(time.RFC3339Nano))
}
rows, err := store.database.QueryContext(ctx, "SELECT id,title,goods_id,sku_color,sku_size,quantity,max_total_price,status,version,created_at FROM tasks WHERE "+strings.Join(clauses, " AND ")+" ORDER BY julianday(created_at) DESC,rowid DESC", args...)
if err != nil {
return nil, err
}
defer rows.Close()
result := []TaskRow{}
for rows.Next() {
var item TaskRow
var created string
if err := rows.Scan(&item.ID, &item.Title, &item.GoodsID, &item.SKUColor, &item.SKUSize, &item.Quantity, &item.MaxTotalPrice, &item.Status, &item.Version, &created); err != nil {
return nil, err
}
item.CreatedAt, err = time.Parse(time.RFC3339Nano, created)
if err != nil {
return nil, err
}
result = append(result, item)
}
return result, rows.Err()
}
// ValidateTaskFilter 返回可关联到字段的错误,使服务端页面拒绝篡改参数时仍能保留输入值。
func ValidateTaskFilter(filter TaskFilter) Errors {
validation := Errors{}
if !validTaskStatus(filter.Status) {
validation["status"] = "请选择有效的任务状态。"
}
location, err := time.LoadLocation("Asia/Shanghai")
if err != nil {
validation["created_from"] = "日期筛选暂不可用,请稍后重试。"
validation["created_to"] = "日期筛选暂不可用,请稍后重试。"
return validation
}
parseDate := func(field, value string) (time.Time, bool) {
if value == "" {
return time.Time{}, true
}
parsed, parseErr := time.ParseInLocation("2006-01-02", value, location)
if parseErr != nil {
validation[field] = "请输入有效日期。"
return time.Time{}, false
}
return parsed, true
}
from, fromOK := parseDate("created_from", filter.CreatedFrom)
to, toOK := parseDate("created_to", filter.CreatedTo)
if fromOK && toOK && !from.IsZero() && !to.IsZero() && from.After(to) {
validation["created_to"] = "结束日期不能早于开始日期。"
}
return validation
}
func (store *SQLiteStore) StartPurchases(ctx context.Context, command StartCommand, createdBy string) (StartResult, error) {
items, err := startItems(command)
if err != nil || createdBy == "" {
return StartResult{}, ErrInvalidStart
}
if store.policy.AuthorizationTTL <= 0 || store.policy.MaxQuantity <= 0 {
return StartResult{}, ErrInvalidStart
}
_, ceiling, ok := normalizeCents(store.policy.MaxTotalPrice)
if !ok {
return StartResult{}, ErrInvalidStart
}
writeCtx, cancel := context.WithTimeout(ctx, sqliteWriteTimeout)
defer cancel()
select {
case store.writeGate <- struct{}{}:
defer func() { <-store.writeGate }()
case <-writeCtx.Done():
return StartResult{}, writeCtx.Err()
}
tx, err := store.database.BeginTx(writeCtx, nil)
if err != nil {
return StartResult{}, err
}
defer tx.Rollback()
// Replay precedes any DRAFT check. One service process serializes this check with creation; SQLite uniqueness remains the cross-transaction backstop.
result, found, err := replayStart(writeCtx, tx, command.StartKey, items)
if err != nil {
return StartResult{}, err
}
if found {
if err := tx.Commit(); err != nil {
return StartResult{}, err
}
return result, nil
}
now := store.now().UTC()
expires := now.Add(store.policy.AuthorizationTTL)
result = StartResult{StartKey: command.StartKey, AuthorizedCount: len(items), Tasks: make([]AuthorizedTask, 0, len(items)), PaymentAutomated: false}
for _, item := range items {
var title, goods, color, size, price, status string
var quantity, version int
if err := tx.QueryRowContext(writeCtx, "SELECT title,goods_id,sku_color,sku_size,quantity,max_total_price,status,version FROM tasks WHERE id=?", item.TaskID).Scan(&title, &goods, &color, &size, &quantity, &price, &status, &version); err != nil {
if err == sql.ErrNoRows {
return StartResult{}, ErrStartConflict
}
return StartResult{}, err
}
if status != "DRAFT" || version != item.ExpectedTaskVersion ||
!ValidTaskWireFields(title, goods, color, size, price) ||
quantity < 1 || quantity > store.policy.MaxQuantity {
return StartResult{}, ErrStartConflict
}
canonical, cents, ok := normalizeCents(price)
if !ok || canonical != price || cents.Cmp(ceiling) > 0 {
return StartResult{}, ErrStartConflict
}
if _, err := domain.TransitionTask(domain.TaskStatusDraft, domain.TaskStatusPending); err != nil {
return StartResult{}, err
}
id, err := NewCreateKey()
if err != nil {
return StartResult{}, err
}
next := version + 1
if _, err = tx.ExecContext(writeCtx, "INSERT INTO order_authorizations (id,task_id,task_version,start_key,goods_id,sku_color,sku_size,quantity,total_price_cap,status,created_by,created_at,expires_at) VALUES (?,?,?,?,?,?,?,?,?,'ACTIVE',?,?,?)", id, item.TaskID, next, command.StartKey, goods, color, size, quantity, price, createdBy, now.Format(time.RFC3339Nano), expires.Format(time.RFC3339Nano)); err != nil {
return StartResult{}, err
}
updated, err := tx.ExecContext(writeCtx, "UPDATE tasks SET status='PENDING',version=version+1,updated_at=? WHERE id=? AND status='DRAFT' AND version=?", now.Format(time.RFC3339Nano), item.TaskID, version)
if err != nil {
return StartResult{}, err
}
affected, err := updated.RowsAffected()
if err != nil {
return StartResult{}, err
}
if affected != 1 {
return StartResult{}, ErrStartConflict
}
result.Tasks = append(result.Tasks, AuthorizedTask{TaskID: item.TaskID, TaskVersion: next, AuthorizationID: id, ExpiresAt: expires})
}
if err := tx.Commit(); err != nil {
return StartResult{}, err
}
return result, nil
}
func replayStart(ctx context.Context, tx *sql.Tx, startKey string, items []StartItem) (StartResult, bool, error) {
rows, err := tx.QueryContext(ctx, `SELECT authorizations.id,authorizations.task_id,
authorizations.task_version,authorizations.expires_at,authorizations.goods_id,
authorizations.sku_color,authorizations.sku_size,authorizations.quantity,
authorizations.total_price_cap,tasks.title,tasks.goods_id,tasks.sku_color,
tasks.sku_size,tasks.quantity,tasks.max_total_price
FROM order_authorizations AS authorizations
JOIN tasks ON tasks.id = authorizations.task_id
WHERE authorizations.start_key=? ORDER BY authorizations.task_id`, startKey)
if err != nil {
return StartResult{}, false, err
}
defer rows.Close()
result := StartResult{StartKey: startKey, PaymentAutomated: false}
for rows.Next() {
var item AuthorizedTask
var expires string
var authorizationGoodsID, authorizationColor, authorizationSize, authorizationPrice string
var taskTitle, taskGoodsID, taskColor, taskSize, taskPrice string
var authorizationQuantity, taskQuantity int
if err := rows.Scan(&item.AuthorizationID, &item.TaskID, &item.TaskVersion, &expires,
&authorizationGoodsID, &authorizationColor, &authorizationSize, &authorizationQuantity,
&authorizationPrice, &taskTitle, &taskGoodsID, &taskColor, &taskSize, &taskQuantity,
&taskPrice); err != nil {
return StartResult{}, false, err
}
if !ValidAuthorizationFields(authorizationGoodsID, authorizationColor, authorizationSize, authorizationPrice) ||
authorizationQuantity <= 0 ||
!ValidTaskWireFields(taskTitle, taskGoodsID, taskColor, taskSize, taskPrice) || taskQuantity <= 0 ||
authorizationGoodsID != taskGoodsID || authorizationColor != taskColor ||
authorizationSize != taskSize || authorizationQuantity != taskQuantity || authorizationPrice != taskPrice {
return StartResult{}, false, ErrStartConflict
}
item.ExpiresAt, err = time.Parse(time.RFC3339Nano, expires)
if err != nil {
return StartResult{}, false, err
}
result.Tasks = append(result.Tasks, item)
}
if err := rows.Err(); err != nil {
return StartResult{}, false, err
}
if len(result.Tasks) == 0 {
return StartResult{}, false, nil
}
if len(result.Tasks) != len(items) {
return StartResult{}, false, ErrStartConflict
}
for i := range items {
if result.Tasks[i].TaskID != items[i].TaskID || result.Tasks[i].TaskVersion-1 != items[i].ExpectedTaskVersion {
return StartResult{}, false, ErrStartConflict
}
}
result.AuthorizedCount = len(result.Tasks)
return result, true, nil
}
+195
View File
@@ -0,0 +1,195 @@
package tasks
import (
"context"
"database/sql"
"errors"
"testing"
"time"
)
func TestListTasksTreatsLikeMetacharactersLiterally(t *testing.T) {
database := migratedDatabase(t)
store, err := NewSQLiteStore(database)
if err != nil {
t.Fatal(err)
}
created := "2026-08-04T01:00:00Z"
insertTaskRow(t, database, "percent", "100%纯棉", "100", "DRAFT", created)
insertTaskRow(t, database, "underscore", "尺码_A", "101", "DRAFT", created)
insertTaskRow(t, database, "backslash", `路径\名称`, "102", "DRAFT", created)
insertTaskRow(t, database, "plain", "普通商品", "103", "DRAFT", created)
for _, test := range []struct {
keyword string
wantID string
}{
{keyword: "%", wantID: "percent"},
{keyword: "_", wantID: "underscore"},
{keyword: `\`, wantID: "backslash"},
} {
t.Run(test.wantID, func(t *testing.T) {
rows, err := store.ListTasks(context.Background(), TaskFilter{Keyword: test.keyword})
if err != nil {
t.Fatal(err)
}
if len(rows) != 1 || rows[0].ID != test.wantID {
t.Fatalf("keyword %q rows = %#v, want only %q", test.keyword, rows, test.wantID)
}
})
}
}
func TestListTasksSupportsEveryStatusAndEmptyMeansAll(t *testing.T) {
database := migratedDatabase(t)
store, err := NewSQLiteStore(database)
if err != nil {
t.Fatal(err)
}
statuses := []string{"DRAFT", "PENDING", "CLAIMED", "ORDERING", "NEEDS_MANUAL", "WAITING_PAYMENT", "RECONCILIATION_REQUIRED", "SUCCEEDED", "FAILED", "CANCELED"}
for index, status := range statuses {
insertTaskRow(t, database, status, status, "200", status, time.Date(2026, 8, 4, 1, 0, index, 0, time.UTC).Format(time.RFC3339Nano))
}
all, err := store.ListTasks(context.Background(), TaskFilter{})
if err != nil {
t.Fatal(err)
}
if len(all) != len(statuses) {
t.Fatalf("all-status rows = %d, want %d", len(all), len(statuses))
}
for _, status := range statuses {
rows, err := store.ListTasks(context.Background(), TaskFilter{Status: status})
if err != nil {
t.Fatalf("status %s: %v", status, err)
}
if len(rows) != 1 || rows[0].Status != status {
t.Fatalf("status %s rows = %#v", status, rows)
}
}
}
func TestListTasksUsesShanghaiHalfOpenDateRange(t *testing.T) {
database := migratedDatabase(t)
store, err := NewSQLiteStore(database)
if err != nil {
t.Fatal(err)
}
insertTaskRow(t, database, "before", "before", "300", "DRAFT", "2026-08-03T15:59:59Z")
insertTaskRow(t, database, "at-start", "at-start", "301", "DRAFT", "2026-08-03T16:00:00Z")
insertTaskRow(t, database, "before-end", "before-end", "302", "DRAFT", "2026-08-04T15:59:59Z")
insertTaskRow(t, database, "at-end", "at-end", "303", "DRAFT", "2026-08-04T16:00:00Z")
rows, err := store.ListTasks(context.Background(), TaskFilter{CreatedFrom: "2026-08-04", CreatedTo: "2026-08-04"})
if err != nil {
t.Fatal(err)
}
if len(rows) != 2 || rows[0].ID != "before-end" || rows[1].ID != "at-start" {
t.Fatalf("Shanghai day rows = %#v, want [before-end at-start]", rows)
}
}
func TestListTasksBreaksEqualTimestampsByDescendingRowID(t *testing.T) {
database := migratedDatabase(t)
store, err := NewSQLiteStore(database)
if err != nil {
t.Fatal(err)
}
created := "2026-08-04T01:02:03Z"
insertTaskRow(t, database, "first", "first", "400", "DRAFT", created)
insertTaskRow(t, database, "second", "second", "401", "DRAFT", created)
rows, err := store.ListTasks(context.Background(), TaskFilter{})
if err != nil {
t.Fatal(err)
}
if len(rows) != 2 || rows[0].ID != "second" || rows[1].ID != "first" {
t.Fatalf("equal-time rows = %#v, want descending rowid", rows)
}
}
func TestListTasksRejectsInvalidStatusAndDates(t *testing.T) {
store, err := NewSQLiteStore(migratedDatabase(t))
if err != nil {
t.Fatal(err)
}
for name, filter := range map[string]TaskFilter{
"status": {Status: "UNKNOWN"},
"from date": {CreatedFrom: "2026-02-30"},
"to date": {CreatedTo: "04/08/2026"},
"reverse range": {CreatedFrom: "2026-08-05", CreatedTo: "2026-08-04"},
} {
t.Run(name, func(t *testing.T) {
rows, err := store.ListTasks(context.Background(), filter)
if !errors.Is(err, ErrInvalidFilter) || rows != nil {
t.Fatalf("ListTasks(%#v) = (%#v, %v), want ErrInvalidFilter", filter, rows, err)
}
})
}
}
func TestStartPurchasesIsAtomicAndReplaysSameSet(t *testing.T) {
database := migratedDatabase(t)
store, err := NewSQLiteStore(database)
if err != nil {
t.Fatal(err)
}
store.SetStartPolicy(StartPolicy{AuthorizationTTL: time.Hour, MaxQuantity: 10, MaxTotalPrice: "999.99"})
store.now = func() time.Time { return time.Date(2026, 8, 4, 1, 2, 3, 0, time.UTC) }
for _, draft := range []Draft{testDraft(testKey, "one"), testDraft("b3c9f507-7473-4fa6-8d71-8786c34c6301", "two")} {
if _, err := store.CreateDraft(context.Background(), draft); err != nil {
t.Fatal(err)
}
}
command := StartCommand{StartKey: "c3c9f507-7473-4fa6-8d71-8786c34c6301", Tasks: []StartItem{{TaskID: "b3c9f507-7473-4fa6-8d71-8786c34c6301", ExpectedTaskVersion: 1}, {TaskID: testKey, ExpectedTaskVersion: 1}}}
first, err := store.StartPurchases(context.Background(), command, "admin")
if err != nil {
t.Fatal(err)
}
if first.AuthorizedCount != 2 || first.PaymentAutomated {
t.Fatalf("start result=%#v", first)
}
command.Tasks[0], command.Tasks[1] = command.Tasks[1], command.Tasks[0]
replay, err := store.StartPurchases(context.Background(), command, "admin")
if err != nil {
t.Fatal(err)
}
if replay.Tasks[0].AuthorizationID != first.Tasks[0].AuthorizationID || replay.Tasks[1].AuthorizationID != first.Tasks[1].AuthorizationID {
t.Fatalf("replay=%#v first=%#v", replay, first)
}
var pending, auths int
if err := database.QueryRow(`SELECT COUNT(*) FROM tasks WHERE status='PENDING' AND version=2`).Scan(&pending); err != nil || pending != 2 {
t.Fatalf("pending=%d err=%v", pending, err)
}
if err := database.QueryRow(`SELECT COUNT(*) FROM order_authorizations WHERE status='ACTIVE' AND created_by='admin'`).Scan(&auths); err != nil || auths != 2 {
t.Fatalf("auths=%d err=%v", auths, err)
}
_, err = store.StartPurchases(context.Background(), StartCommand{StartKey: command.StartKey, Tasks: command.Tasks[:1]}, "admin")
if !errors.Is(err, ErrStartConflict) {
t.Fatalf("subset err=%v", err)
}
}
func TestShanghaiRangeAndMoneyAreFailClosed(t *testing.T) {
start, end, err := ShanghaiRange("2026-08-04", "2026-08-04")
if err != nil || start.Format(time.RFC3339) != "2026-08-03T16:00:00Z" || end.Format(time.RFC3339) != "2026-08-04T16:00:00Z" {
t.Fatalf("range=(%s,%s,%v)", start, end, err)
}
for _, value := range []string{"0.01", "12.80", "999999999999999999999999.99"} {
if _, _, ok := normalizeCents(value); !ok {
t.Fatalf("money %q rejected", value)
}
}
for _, value := range []string{"1", "01.20", "0.00", "1.234", "1.", " 1.00", "1e2"} {
if _, _, ok := normalizeCents(value); ok {
t.Fatalf("money %q accepted", value)
}
}
}
func insertTaskRow(t *testing.T, database *sql.DB, id, title, goodsID, status, createdAt string) {
t.Helper()
if _, err := database.Exec(`INSERT INTO tasks (id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price, status, version, created_at, updated_at) VALUES (?, 'MANUAL', ?, ?, '黑色', 'M', 2, '12.80', ?, 1, ?, ?)`, id, title, goodsID, status, createdAt, createdAt); err != nil {
t.Fatalf("insert task %s: %v", id, err)
}
}
+140
View File
@@ -0,0 +1,140 @@
package tasks
import (
"context"
"database/sql"
"errors"
"fmt"
"time"
)
const sqliteWriteTimeout = 2 * time.Second
type Store interface {
CreateDraft(context.Context, Draft) (Draft, error)
ListDrafts(context.Context) ([]Draft, error)
ListTasks(context.Context, TaskFilter) ([]TaskRow, error)
StartPurchases(context.Context, StartCommand, string) (StartResult, error)
}
type SQLiteStore struct {
database *sql.DB
now func() time.Time
writeGate chan struct{}
policy StartPolicy
}
func NewSQLiteStore(database *sql.DB) (*SQLiteStore, error) {
if database == nil {
return nil, errors.New("database is required")
}
if _, err := database.Exec("SELECT task_version, start_key, total_price_cap FROM order_authorizations LIMIT 1"); err != nil {
return nil, fmt.Errorf("tasks migration is not available: %w", err)
}
if _, err := database.Exec("SELECT 1 FROM purchase_attempts LIMIT 1"); err != nil {
return nil, fmt.Errorf("single-pass migration is not available: %w", err)
}
return &SQLiteStore{database: database, now: time.Now, writeGate: make(chan struct{}, 1)}, nil
}
func (store *SQLiteStore) CreateDraft(ctx context.Context, draft Draft) (Draft, error) {
// Validate again at the persistence boundary. HTTP form validation is not the only
// caller, and a malformed row here would later make an authorized claim unencodable.
if !validUUID(draft.ID) || !ValidTaskWireFields(draft.Title, draft.GoodsID, draft.SKUColor, draft.SKUSize, draft.MaxTotalPrice) || draft.Quantity <= 0 {
return Draft{}, ErrInvalidDraft
}
writeContext, cancel := context.WithTimeout(ctx, sqliteWriteTimeout)
defer cancel()
// SQLite permits one writer at a time. Serializing this store's short create
// transaction prevents concurrent retries of one create key from surfacing as busy.
select {
case store.writeGate <- struct{}{}:
defer func() { <-store.writeGate }()
case <-writeContext.Done():
return Draft{}, writeContext.Err()
}
draft.CreatedAt = store.now().UTC()
transaction, err := store.database.BeginTx(writeContext, nil)
if err != nil {
return Draft{}, err
}
defer transaction.Rollback()
_, err = transaction.ExecContext(writeContext, `INSERT INTO tasks (id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price, status, version, created_at, updated_at) VALUES (?, 'MANUAL', ?, ?, ?, ?, ?, ?, 'DRAFT', 1, ?, ?)`, draft.ID, draft.Title, draft.GoodsID, draft.SKUColor, draft.SKUSize, draft.Quantity, draft.MaxTotalPrice, draft.CreatedAt.Format(time.RFC3339Nano), draft.CreatedAt.Format(time.RFC3339Nano))
if err == nil {
if err := transaction.Commit(); err != nil {
return Draft{}, err
}
return draft, nil
}
existing, found, currentPhase, lookupErr := findDraft(writeContext, transaction, draft.ID)
if lookupErr != nil {
return Draft{}, lookupErr
}
if found && currentPhase && samePayload(existing, draft) {
if err := transaction.Commit(); err != nil {
return Draft{}, err
}
return existing, nil
}
if found {
return Draft{}, ErrCreateKeyConflict
}
return Draft{}, err
}
func (store *SQLiteStore) ListDrafts(ctx context.Context) ([]Draft, error) {
// rowid makes equal timestamps deterministic: SQLite assigns it in insertion order,
// whereas UUID v4 is deliberately not time-sortable.
rows, err := store.database.QueryContext(ctx, `SELECT id, title, goods_id, sku_color, sku_size, quantity, max_total_price, created_at FROM tasks WHERE source = 'MANUAL' AND status = 'DRAFT' ORDER BY created_at DESC, rowid DESC`)
if err != nil {
return nil, err
}
defer rows.Close()
result := []Draft{}
for rows.Next() {
draft, err := scanDraft(rows)
if err != nil {
return nil, err
}
result = append(result, draft)
}
return result, rows.Err()
}
func findDraft(ctx context.Context, transaction *sql.Tx, id string) (Draft, bool, bool, error) {
row := transaction.QueryRowContext(ctx, `SELECT id, title, goods_id, sku_color, sku_size, quantity, max_total_price, created_at, source, status, version FROM tasks WHERE id = ?`, id)
var draft Draft
var created, source, status string
var version int
err := row.Scan(&draft.ID, &draft.Title, &draft.GoodsID, &draft.SKUColor, &draft.SKUSize, &draft.Quantity, &draft.MaxTotalPrice, &created, &source, &status, &version)
if errors.Is(err, sql.ErrNoRows) {
return Draft{}, false, false, nil
}
if err != nil {
return Draft{}, false, false, err
}
parsed, err := time.Parse(time.RFC3339Nano, created)
if err != nil {
return Draft{}, false, false, err
}
draft.CreatedAt = parsed
return draft, true, source == "MANUAL" && status == "DRAFT" && version == 1, nil
}
type scanner interface{ Scan(...any) error }
func scanDraft(row scanner) (Draft, error) {
var draft Draft
var created string
if err := row.Scan(&draft.ID, &draft.Title, &draft.GoodsID, &draft.SKUColor, &draft.SKUSize, &draft.Quantity, &draft.MaxTotalPrice, &created); err != nil {
return Draft{}, err
}
parsed, err := time.Parse(time.RFC3339Nano, created)
if err != nil {
return Draft{}, err
}
draft.CreatedAt = parsed
return draft, nil
}
func samePayload(left, right Draft) bool {
return left.ID == right.ID && left.Title == right.Title && left.GoodsID == right.GoodsID && left.SKUColor == right.SKUColor && left.SKUSize == right.SKUSize && left.Quantity == right.Quantity && left.MaxTotalPrice == right.MaxTotalPrice
}
+218
View File
@@ -0,0 +1,218 @@
// Package tasks 定义手工 DRAFT 任务的校验与窄仓储边界。
package tasks
import (
"crypto/rand"
"encoding/hex"
"errors"
"net/url"
"strconv"
"strings"
"time"
"unicode/utf8"
)
const (
MaxTitleCodePoints = 120
MaxSKUTextCodePoints = 80
MaxGoodsIDCharacters = 32
MaxMoneyASCIICharacters = 32
maxSKUText = MaxSKUTextCodePoints
)
var (
ErrCreateKeyConflict = errors.New("create key conflicts with a different task")
ErrInvalidDraft = errors.New("invalid draft")
)
type Draft struct {
ID string
Title string
GoodsID string
SKUColor string
SKUSize string
Quantity int
MaxTotalPrice string
CreatedAt time.Time
}
type Form struct{ CreateKey, Title, ProductURL, SKUColor, SKUSize, Quantity, MaxTotalPrice string }
type Errors map[string]string
func (errors Errors) Valid() bool { return len(errors) == 0 }
// Validate trims and normalizes a user form. It never reads a product page or derives price data.
func Validate(form Form) (Draft, Errors) {
draft := Draft{ID: strings.TrimSpace(form.CreateKey), Title: strings.TrimSpace(form.Title), SKUColor: strings.TrimSpace(form.SKUColor), SKUSize: strings.TrimSpace(form.SKUSize)}
errors := Errors{}
if !validUUID(draft.ID) {
errors["create_key"] = "创建请求已过期,请重新打开表单。"
}
if !validBoundedText(draft.Title, MaxTitleCodePoints) {
errors["title"] = "任务名称不能为空,且不能超过 120 个字符。"
}
if !validBoundedText(draft.SKUColor, MaxSKUTextCodePoints) {
errors["sku_color"] = "颜色分类不能为空,且不能超过 80 个字符。"
}
if !validBoundedText(draft.SKUSize, MaxSKUTextCodePoints) {
errors["sku_size"] = "尺码不能为空,且不能超过 80 个字符。"
}
goodsID, ok := CanonicalGoodsID(strings.TrimSpace(form.ProductURL))
if !ok {
errors["product_url"] = "请输入唯一的 canonical 商品链接。"
} else {
draft.GoodsID = goodsID
}
quantity, err := strconv.ParseInt(strings.TrimSpace(form.Quantity), 10, 0)
if err != nil || quantity < 1 {
errors["quantity"] = "数量必须是正整数。"
} else {
draft.Quantity = int(quantity)
}
money, ok := normalizeMoney(strings.TrimSpace(form.MaxTotalPrice))
if !ok {
errors["max_total_price"] = "价格上限必须大于零,且最多两位小数。"
} else {
draft.MaxTotalPrice = money
}
return draft, errors
}
// CanonicalGoodsID only accepts the one verified manual-entry URL shape; untrusted query data is discarded.
func CanonicalGoodsID(value string) (string, bool) {
if value == "" || strings.Contains(value, "\\") || strings.Contains(value, "%") {
return "", false
}
parsed, err := url.ParseRequestURI(value)
if err != nil || parsed.Scheme != "https" || parsed.Host != "mobile.yangkeduo.com" || parsed.User != nil || parsed.Port() != "" || parsed.Path != "/goods.html" || parsed.Fragment != "" {
return "", false
}
values, err := url.ParseQuery(parsed.RawQuery)
if err != nil {
return "", false
}
goodsIDs := values["goods_id"]
if len(goodsIDs) != 1 || goodsIDs[0] == "" {
return "", false
}
for _, character := range goodsIDs[0] {
if character < '0' || character > '9' {
return "", false
}
}
if !ValidGoodsID(goodsIDs[0]) {
return "", false
}
return goodsIDs[0], true
}
func CanonicalURL(goodsID string) string {
return "https://mobile.yangkeduo.com/goods.html?goods_id=" + goodsID
}
func NewCreateKey() (string, error) {
bytes := make([]byte, 16)
if _, err := rand.Read(bytes); err != nil {
return "", err
}
bytes[6] = (bytes[6] & 0x0f) | 0x40
bytes[8] = (bytes[8] & 0x3f) | 0x80
hexValue := hex.EncodeToString(bytes)
return hexValue[0:8] + "-" + hexValue[8:12] + "-" + hexValue[12:16] + "-" + hexValue[16:20] + "-" + hexValue[20:32], nil
}
func validUUID(value string) bool {
if len(value) != 36 {
return false
}
for index, character := range value {
if index == 8 || index == 13 || index == 18 || index == 23 {
if character != '-' {
return false
}
continue
}
if !(character >= '0' && character <= '9' || character >= 'a' && character <= 'f') {
return false
}
}
return value[14] == '4' && (value[19] == '8' || value[19] == '9' || value[19] == 'a' || value[19] == 'b')
}
func normalizeMoney(value string) (string, bool) {
parts := strings.Split(value, ".")
if len(parts) > 2 || parts[0] == "" || len(parts) == 2 && (len(parts[1]) == 0 || len(parts[1]) > 2) {
return "", false
}
for _, character := range parts[0] {
if character < '0' || character > '9' {
return "", false
}
}
fraction := ""
if len(parts) == 2 {
fraction = parts[1]
for _, character := range fraction {
if character < '0' || character > '9' {
return "", false
}
}
}
whole := strings.TrimLeft(parts[0], "0")
if whole == "" {
whole = "0"
}
if whole == "0" && strings.Trim(fraction, "0") == "" {
return "", false
}
canonical := whole + "." + (fraction + "00")[:2]
if len(canonical) > MaxMoneyASCIICharacters {
return "", false
}
return canonical, true
}
// ValidTaskWireFields is shared by creation, authorization and claim. Keeping one
// bounded domain prevents a database row from being valid in one stage but impossible
// to encode inside the fixed claim response budget in another stage.
func ValidTaskWireFields(title, goodsID, skuColor, skuSize, maxTotalPrice string) bool {
return validBoundedText(title, MaxTitleCodePoints) &&
ValidAuthorizationFields(goodsID, skuColor, skuSize, maxTotalPrice)
}
func ValidAuthorizationFields(goodsID, skuColor, skuSize, totalPriceCap string) bool {
return ValidGoodsID(goodsID) &&
validBoundedText(skuColor, MaxSKUTextCodePoints) &&
validBoundedText(skuSize, MaxSKUTextCodePoints) &&
ValidCanonicalMoney(totalPriceCap)
}
func ValidGoodsID(value string) bool {
if value == "" || len(value) > MaxGoodsIDCharacters {
return false
}
for index := 0; index < len(value); index++ {
if value[index] < '0' || value[index] > '9' {
return false
}
}
return true
}
func validBoundedText(value string, maximum int) bool {
// RuneCountInString replaces malformed byte sequences with RuneError. Validate first
// so corrupt SQLite text cannot consume the code-point budget as if it were legitimate.
if !utf8.ValidString(value) || value == "" || strings.TrimSpace(value) != value ||
utf8.RuneCountInString(value) > maximum {
return false
}
for _, character := range value {
// Python str.strip treats these four C0 separators as whitespace while Go
// TrimSpace does not. Reject them everywhere so both wire models have one
// explicit persisted-text domain instead of runtime-dependent trimming.
if character >= '\u001c' && character <= '\u001f' {
return false
}
}
return true
}
+379
View File
@@ -0,0 +1,379 @@
package tasks
import (
"context"
"database/sql"
"errors"
"path/filepath"
"regexp"
"runtime"
"strings"
"sync"
"testing"
"time"
"cmbuyer/admin/internal/migrations"
"cmbuyer/admin/internal/storage/sqlite"
)
const testKey = "a3c9f507-7473-4fa6-8d71-8786c34c6301"
func TestValidateNormalizesManualDraft(t *testing.T) {
draft, validation := Validate(Form{
CreateKey: " " + testKey + " ",
Title: " 夏季上衣 ",
ProductURL: "https://mobile.yangkeduo.com/goods.html?goods_id=937122477375&utm_source=untrusted",
SKUColor: " 黑色CHA(纯棉) ",
SKUSize: " M(建议100-115) ",
Quantity: "2",
MaxTotalPrice: "00012.8",
})
if !validation.Valid() {
t.Fatalf("Validate errors = %#v", validation)
}
if draft.ID != testKey || draft.GoodsID != "937122477375" || draft.Title != "夏季上衣" || draft.SKUColor != "黑色CHA(纯棉)" || draft.SKUSize != "M(建议100-115)" || draft.Quantity != 2 || draft.MaxTotalPrice != "12.80" {
t.Fatalf("normalized draft = %#v", draft)
}
}
func TestValidateRejectsInvalidFieldsAndURLs(t *testing.T) {
base := Form{CreateKey: testKey, Title: "title", ProductURL: "https://mobile.yangkeduo.com/goods.html?goods_id=1", SKUColor: "black", SKUSize: "M", Quantity: "1", MaxTotalPrice: "1"}
for name, update := range map[string]func(*Form){
"empty title": func(form *Form) { form.Title = " " },
"invalid utf8 title": func(form *Form) { form.Title = string([]byte{0xff}) },
"long title": func(form *Form) { form.Title = strings.Repeat("😀", MaxTitleCodePoints+1) },
"long color": func(form *Form) { form.SKUColor = string(make([]rune, maxSKUText+1)) },
"invalid utf8 size": func(form *Form) { form.SKUSize = string([]byte{0xff}) },
"fraction quantity": func(form *Form) { form.Quantity = "1.5" },
"zero quantity": func(form *Form) { form.Quantity = "0" },
"too many decimals": func(form *Form) { form.MaxTotalPrice = "1.234" },
"trailing decimal": func(form *Form) { form.MaxTotalPrice = "1." },
"zero money": func(form *Form) { form.MaxTotalPrice = "0.00" },
} {
t.Run(name, func(t *testing.T) {
form := base
update(&form)
if _, validation := Validate(form); validation.Valid() {
t.Fatal("invalid form was accepted")
}
})
}
for _, value := range []string{
"http://mobile.yangkeduo.com/goods.html?goods_id=1",
"https://yangkeduo.com/goods.html?goods_id=1",
"https://mobile.yangkeduo.com:443/goods.html?goods_id=1",
"https://user@mobile.yangkeduo.com/goods.html?goods_id=1",
"https://mobile.yangkeduo.com/goods.html?goods_id=1#fragment",
"https://mobile.yangkeduo.com/goods.html?goods_id=1&goods_id=2",
"https://mobile.yangkeduo.com/goods.html?goods_id=one",
"https://mobile.yangkeduo.com/goods.html?goods_id=%31",
"https://mobile.yangkeduo.com/goods.html?goods_id=1%26goods_id%3D2",
"https://mobile.yangkeduo.com/goods.html?goods_id=1;uin=bad",
"https://mobile.yangkeduo.com/other.html?goods_id=1",
"https://mobile.yangkeduo.com/goods.html?goods_id=" + strings.Repeat("1", MaxGoodsIDCharacters+1),
} {
if _, ok := CanonicalGoodsID(value); ok {
t.Fatalf("CanonicalGoodsID accepted %q", value)
}
}
}
func TestNormalizeMoneyBoundaries(t *testing.T) {
maximum := strings.Repeat("9", MaxMoneyASCIICharacters-3) + ".00"
for value, want := range map[string]string{"1": "1.00", "1.2": "1.20", "000.01": "0.01", "999999999999999999": "999999999999999999.00", maximum: maximum} {
got, ok := normalizeMoney(value)
if !ok || got != want {
t.Fatalf("normalizeMoney(%q) = (%q, %t), want (%q, true)", value, got, ok, want)
}
}
for _, value := range []string{"0", "0.0", "0.00", "1.", ".1", "1.000", "-1", "1e2", " 1", strings.Repeat("9", MaxMoneyASCIICharacters-2) + ".00"} {
if got, ok := normalizeMoney(value); ok {
t.Fatalf("normalizeMoney(%q) = %q, want rejection", value, got)
}
}
}
func TestValidateAcceptsWorstLegalUnicodeFieldBounds(t *testing.T) {
goodsID := strings.Repeat("1", MaxGoodsIDCharacters)
draft, validation := Validate(Form{
CreateKey: testKey, Title: strings.Repeat("😀", MaxTitleCodePoints),
ProductURL: CanonicalURL(goodsID), SKUColor: strings.Repeat("色", MaxSKUTextCodePoints),
SKUSize: strings.Repeat("码", MaxSKUTextCodePoints), Quantity: "1",
MaxTotalPrice: strings.Repeat("9", MaxMoneyASCIICharacters-3) + ".00",
})
if !validation.Valid() || !ValidTaskWireFields(draft.Title, draft.GoodsID, draft.SKUColor, draft.SKUSize, draft.MaxTotalPrice) {
t.Fatalf("worst legal draft = %#v, validation = %#v", draft, validation)
}
}
func TestPersistedTextHasRuntimeIndependentC0AndNBSPDomain(t *testing.T) {
for name, invalid := range map[string]string{
"c0 prefix": "\u001cvalue",
"c0 suffix": "value\u001f",
"c0 interior": "value\u001dinside",
"nbsp prefix": "\u00a0value",
"nbsp suffix": "value\u00a0",
} {
t.Run(name, func(t *testing.T) {
if validBoundedText(invalid, MaxTitleCodePoints) {
t.Fatalf("validBoundedText(%q) accepted runtime-dependent text", invalid)
}
})
}
if !validBoundedText("left\u00a0right", MaxTitleCodePoints) {
t.Fatal("interior NBSP must remain a valid Unicode code point")
}
// Manual form input is normalized with Go TrimSpace before persistence.
draft, validation := Validate(Form{
CreateKey: testKey, Title: "\u00a0title\u00a0",
ProductURL: CanonicalURL("1"), SKUColor: "\u00a0black\u00a0",
SKUSize: "\u00a0M\u00a0", Quantity: "1", MaxTotalPrice: "1",
})
if !validation.Valid() || draft.Title != "title" || draft.SKUColor != "black" || draft.SKUSize != "M" {
t.Fatalf("NBSP form normalization = %#v, errors = %#v", draft, validation)
}
}
func TestNewCreateKeyIsUUIDv4(t *testing.T) {
key, err := NewCreateKey()
if err != nil {
t.Fatalf("NewCreateKey: %v", err)
}
if !regexp.MustCompile(`^[0-9a-f]{8}-[0-9a-f]{4}-4[0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}$`).MatchString(key) {
t.Fatalf("create key %q is not UUID v4", key)
}
}
func TestSQLiteStoreRequiresMigratedDatabase(t *testing.T) {
database := openDatabase(t)
if _, err := NewSQLiteStore(database); err == nil {
t.Fatal("NewSQLiteStore accepted an unmigrated database")
}
}
func TestSQLiteStoreCreatesListsAndHandlesIdempotency(t *testing.T) {
database := migratedDatabase(t)
store, err := NewSQLiteStore(database)
if err != nil {
t.Fatalf("NewSQLiteStore: %v", err)
}
baseTime := time.Date(2026, 8, 4, 9, 0, 0, 0, time.UTC)
call := 0
store.now = func() time.Time {
result := baseTime.Add(time.Duration(call) * time.Minute)
call++
return result
}
first := testDraft(testKey, "first")
created, err := store.CreateDraft(context.Background(), first)
if err != nil {
t.Fatalf("create first draft: %v", err)
}
replayed, err := store.CreateDraft(context.Background(), first)
if err != nil {
t.Fatalf("replay first draft: %v", err)
}
if replayed.CreatedAt != created.CreatedAt {
t.Fatalf("replayed CreatedAt = %s, want original %s", replayed.CreatedAt, created.CreatedAt)
}
second := testDraft("b3c9f507-7473-4fa6-8d71-8786c34c6301", "second")
if _, err := store.CreateDraft(context.Background(), second); err != nil {
t.Fatalf("create second draft: %v", err)
}
drafts, err := store.ListDrafts(context.Background())
if err != nil {
t.Fatalf("list drafts: %v", err)
}
if len(drafts) != 2 || drafts[0].ID != second.ID || drafts[1].ID != first.ID {
t.Fatalf("draft order = %#v, want second then first", drafts)
}
var source, status string
var version int
if err := database.QueryRow(`SELECT source, status, version FROM tasks WHERE id = ?`, first.ID).Scan(&source, &status, &version); err != nil {
t.Fatalf("read stored task: %v", err)
}
if source != "MANUAL" || status != "DRAFT" || version != 1 {
t.Fatalf("stored metadata = (%q, %q, %d)", source, status, version)
}
conflicting := first
conflicting.Title = "different"
if _, err := store.CreateDraft(context.Background(), conflicting); !errors.Is(err, ErrCreateKeyConflict) {
t.Fatalf("conflicting create error = %v, want ErrCreateKeyConflict", err)
}
}
func TestSQLiteStoreRejectsInvalidDraftAtPersistenceBoundary(t *testing.T) {
mutations := map[string]func(*Draft){
"untrimmed title": func(draft *Draft) { draft.Title = " title" },
"c0 interior title": func(draft *Draft) { draft.Title = "title\u001dhidden" },
"invalid utf8 title": func(draft *Draft) { draft.Title = string([]byte{0xff}) },
"long title": func(draft *Draft) { draft.Title = strings.Repeat("😀", MaxTitleCodePoints+1) },
"long color": func(draft *Draft) { draft.SKUColor = strings.Repeat("色", MaxSKUTextCodePoints+1) },
"long size": func(draft *Draft) { draft.SKUSize = strings.Repeat("码", MaxSKUTextCodePoints+1) },
"long goods id": func(draft *Draft) { draft.GoodsID = strings.Repeat("1", MaxGoodsIDCharacters+1) },
"long money": func(draft *Draft) { draft.MaxTotalPrice = strings.Repeat("1", MaxMoneyASCIICharacters-2) + ".00" },
}
for name, mutate := range mutations {
t.Run(name, func(t *testing.T) {
database := migratedDatabase(t)
store, err := NewSQLiteStore(database)
if err != nil {
t.Fatal(err)
}
draft := testDraft(testKey, "title")
mutate(&draft)
if _, err := store.CreateDraft(context.Background(), draft); !errors.Is(err, ErrInvalidDraft) {
t.Fatalf("CreateDraft error = %v, want ErrInvalidDraft", err)
}
var count int
if err := database.QueryRow("SELECT COUNT(*) FROM tasks").Scan(&count); err != nil || count != 0 {
t.Fatalf("tasks after invalid create = %d, err %v", count, err)
}
})
}
}
func TestSQLiteStoreRollsBackFailedCreate(t *testing.T) {
database := migratedDatabase(t)
store, err := NewSQLiteStore(database)
if err != nil {
t.Fatalf("NewSQLiteStore: %v", err)
}
if _, err := database.Exec(`CREATE TRIGGER reject_task BEFORE INSERT ON tasks BEGIN SELECT RAISE(ABORT, 'reject test insert'); END`); err != nil {
t.Fatalf("create trigger: %v", err)
}
if _, err := store.CreateDraft(context.Background(), testDraft(testKey, "blocked")); err == nil {
t.Fatal("CreateDraft succeeded despite rejecting trigger")
}
drafts, err := store.ListDrafts(context.Background())
if err != nil {
t.Fatalf("list after failed create: %v", err)
}
if len(drafts) != 0 {
t.Fatalf("failed create persisted drafts: %#v", drafts)
}
}
func TestSQLiteStoreUsesInsertionOrderForEqualTimesAndFiltersPhase(t *testing.T) {
database := migratedDatabase(t)
store, err := NewSQLiteStore(database)
if err != nil {
t.Fatalf("NewSQLiteStore: %v", err)
}
store.now = func() time.Time { return time.Date(2026, 8, 4, 9, 0, 0, 0, time.UTC) }
first := testDraft(testKey, "first")
second := testDraft("b3c9f507-7473-4fa6-8d71-8786c34c6301", "second")
for _, draft := range []Draft{first, second} {
if _, err := store.CreateDraft(context.Background(), draft); err != nil {
t.Fatalf("create %s: %v", draft.Title, err)
}
}
if _, err := database.Exec(`INSERT INTO tasks (id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price, status, version, created_at, updated_at) VALUES ('excel-draft', 'EXCEL', 'other', '1', 'black', 'M', 1, '1.00', 'DRAFT', 1, '2026-08-04T10:00:00Z', '2026-08-04T10:00:00Z'), ('manual-pending', 'MANUAL', 'other', '2', 'black', 'M', 1, '1.00', 'PENDING', 1, '2026-08-04T10:00:00Z', '2026-08-04T10:00:00Z')`); err != nil {
t.Fatalf("insert out-of-scope tasks: %v", err)
}
drafts, err := store.ListDrafts(context.Background())
if err != nil {
t.Fatalf("list drafts: %v", err)
}
if len(drafts) != 2 || drafts[0].ID != second.ID || drafts[1].ID != first.ID {
t.Fatalf("equal-time draft order/filter = %#v, want second then first only", drafts)
}
if _, err := database.Exec(`UPDATE tasks SET status = 'PENDING' WHERE id = ?`, first.ID); err != nil {
t.Fatalf("move draft outside current phase: %v", err)
}
if _, err := store.CreateDraft(context.Background(), first); !errors.Is(err, ErrCreateKeyConflict) {
t.Fatalf("replay of non-DRAFT record error = %v, want conflict", err)
}
third := testDraft("c3c9f507-7473-4fa6-8d71-8786c34c6301", "third")
if _, err := database.Exec(`INSERT INTO tasks (id, source, title, goods_id, sku_color, sku_size, quantity, max_total_price, status, version, created_at, updated_at) VALUES (?, 'EXCEL', ?, ?, ?, ?, ?, ?, 'DRAFT', 1, '2026-08-04T09:00:00Z', '2026-08-04T09:00:00Z')`, third.ID, third.Title, third.GoodsID, third.SKUColor, third.SKUSize, third.Quantity, third.MaxTotalPrice); err != nil {
t.Fatalf("insert same-payload EXCEL record: %v", err)
}
if _, err := store.CreateDraft(context.Background(), third); !errors.Is(err, ErrCreateKeyConflict) {
t.Fatalf("replay of non-MANUAL record error = %v, want conflict", err)
}
if _, err := database.Exec(`UPDATE tasks SET version = 2, source = 'MANUAL' WHERE id = ?`, third.ID); err != nil {
t.Fatalf("change replay record version: %v", err)
}
if _, err := store.CreateDraft(context.Background(), third); !errors.Is(err, ErrCreateKeyConflict) {
t.Fatalf("replay of non-v1 record error = %v, want conflict", err)
}
}
func TestSQLiteStoreConcurrentIdenticalCreateIsOneDraft(t *testing.T) {
store, err := NewSQLiteStore(migratedDatabase(t))
if err != nil {
t.Fatalf("NewSQLiteStore: %v", err)
}
const callers = 20
start := make(chan struct{})
errors := make(chan error, callers)
results := make(chan Draft, callers)
var group sync.WaitGroup
for range callers {
group.Add(1)
go func() {
defer group.Done()
<-start
draft, err := store.CreateDraft(context.Background(), testDraft(testKey, "same"))
if err != nil {
errors <- err
return
}
results <- draft
}()
}
close(start)
group.Wait()
close(errors)
close(results)
for err := range errors {
t.Fatalf("concurrent create: %v", err)
}
for result := range results {
if result.ID != testKey {
t.Fatalf("concurrent result = %#v", result)
}
}
drafts, err := store.ListDrafts(context.Background())
if err != nil {
t.Fatalf("list after concurrent create: %v", err)
}
if len(drafts) != 1 || drafts[0].ID != testKey {
t.Fatalf("concurrent creates persisted %#v, want exactly one", drafts)
}
}
func testDraft(id, title string) Draft {
return Draft{ID: id, Title: title, GoodsID: "937122477375", SKUColor: "black", SKUSize: "M", Quantity: 2, MaxTotalPrice: "12.80"}
}
func openDatabase(t *testing.T) *sql.DB {
t.Helper()
database, err := sqlite.Open(filepath.Join(t.TempDir(), "tasks.db"))
if err != nil {
t.Fatalf("open database: %v", err)
}
t.Cleanup(func() { _ = database.Close() })
return database
}
func migratedDatabase(t *testing.T) *sql.DB {
t.Helper()
database := openDatabase(t)
if err := migrations.Up(context.Background(), database, migrationDirectory(t)); err != nil {
t.Fatalf("migrate database: %v", err)
}
return database
}
func migrationDirectory(t *testing.T) string {
t.Helper()
_, file, _, ok := runtime.Caller(0)
if !ok {
t.Fatal("locate test source")
}
return filepath.Join(filepath.Dir(file), "..", "..", "migrations")
}
@@ -0,0 +1,211 @@
"use strict";
const test = require("node:test");
const assert = require("node:assert/strict");
const fs = require("node:fs");
const path = require("node:path");
const vm = require("node:vm");
const source = fs.readFileSync(path.join(__dirname, "tasks.js"), "utf8");
test("visible button opens the same routed detail and close restores list state", async () => {
const harness = createDrawerHarness();
harness.button.listeners.click();
await harness.flush();
assert.equal(harness.requests.length, 1);
assert.equal(harness.requests[0].url, "/tasks/a3c9f507-7473-4fa6-8d71-8786c34c6301");
assert.equal(harness.requests[0].options.headers["X-CMBuyer-View"], "drawer");
assert.equal(harness.drawer.open, true);
assert.equal(harness.closeButton.focused, true);
assert.equal(harness.history.pushes.length, 1);
assert.equal(harness.history.pushes[0].url, harness.requests[0].url);
assert.equal(harness.history.pushes[0].state.focusTarget, "button");
harness.closeButton.listeners.click();
assert.equal(harness.history.backCalls, 1);
harness.popstate({state: {cmbuyerList: true}});
assert.equal(harness.drawer.open, false);
assert.equal(harness.button.focused, true);
assert.equal(harness.row.focused, false);
assert.equal(harness.scrolls.length, 1);
assert.equal(harness.scrolls[0].top, 275);
assert.equal(harness.scrolls[0].behavior, "auto");
});
test("button focus target survives back, forward, and back again", async () => {
const harness = createDrawerHarness();
harness.button.listeners.click();
await harness.flush();
const drawerState = harness.history.pushes[0].state;
harness.popstate({state: {cmbuyerList: true}});
assert.equal(harness.button.focusCalls, 1);
assert.equal(harness.row.focusCalls, 0);
harness.popstate({state: drawerState});
await harness.flush();
assert.equal(harness.drawer.open, true);
assert.equal(harness.history.pushes.length, 1);
harness.popstate({state: {cmbuyerList: true}});
assert.equal(harness.drawer.open, false);
assert.equal(harness.button.focusCalls, 2);
assert.equal(harness.row.focusCalls, 0);
});
test("failed forward retry reuses history and closes back to button in one step", async () => {
const harness = createDrawerHarness();
harness.button.listeners.click();
await harness.flush();
const drawerState = harness.history.pushes[0].state;
harness.popstate({state: {cmbuyerList: true}});
harness.failNextRequest();
harness.popstate({state: drawerState});
await harness.flush();
const retry = harness.body.children[1].children[0];
retry.listeners.click();
await harness.flush();
assert.equal(harness.history.pushes.length, 1);
assert.equal(harness.drawer.open, true);
harness.closeButton.listeners.click();
assert.equal(harness.history.backCalls, 1);
harness.popstate({state: {cmbuyerList: true}});
assert.equal(harness.drawer.open, false);
assert.equal(harness.button.focusCalls, 2);
assert.equal(harness.row.focusCalls, 0);
});
test("double click and Enter open rows but nested controls never do", async () => {
const harness = createDrawerHarness();
const ignored = {closest: () => ({})};
const rowTarget = {closest: () => null};
harness.row.listeners.dblclick({target: ignored});
harness.row.listeners.dblclick({target: rowTarget});
await harness.flush();
assert.equal(harness.requests.length, 1);
harness.popstate({state: {cmbuyerList: true}});
let prevented = false;
harness.row.listeners.keydown({key: "Enter", target: harness.row, preventDefault: () => { prevented = true; }});
await harness.flush();
assert.equal(prevented, true);
assert.equal(harness.requests.length, 2);
harness.row.listeners.keydown({key: "Enter", target: ignored, preventDefault: () => assert.fail("nested control Enter was intercepted")});
assert.equal(harness.requests.length, 2);
});
test("browser back and forward close and reopen without duplicating history", async () => {
const harness = createDrawerHarness();
harness.row.listeners.keydown({key: "Enter", target: harness.row, preventDefault() {}});
await harness.flush();
assert.equal(harness.history.pushes.length, 1);
harness.popstate({state: {cmbuyerList: true}});
assert.equal(harness.drawer.open, false);
harness.popstate({state: {cmbuyerDrawer: true, detailURL: harness.row.dataset.detailUrl}});
await harness.flush();
assert.equal(harness.drawer.open, true);
assert.equal(harness.requests.length, 2);
assert.equal(harness.history.pushes.length, 1);
});
test("Escape follows browser history and does not mutate list URL", async () => {
const harness = createDrawerHarness();
harness.button.listeners.click();
await harness.flush();
let prevented = false;
harness.drawer.listeners.cancel({preventDefault: () => { prevented = true; }});
assert.equal(prevented, true);
assert.equal(harness.history.backCalls, 1);
assert.equal(harness.history.replaces[0].url, "/tasks?status=DRAFT");
});
function createDrawerHarness() {
class FakeElement {
constructor() {
this.listeners = {};
this.dataset = {};
this.open = false;
this.focused = false;
this.focusCalls = 0;
this.children = [];
this._innerHTML = "";
}
addEventListener(type, listener) { this.listeners[type] = listener; }
focus() { this.focused = true; this.focusCalls++; }
showModal() { this.open = true; }
close() { this.open = false; }
replaceChildren(...children) { this.children = children; this._innerHTML = ""; }
append(...children) { this.children.push(...children); }
setAttribute() {}
closest() { return null; }
set innerHTML(value) { this._innerHTML = value; }
get innerHTML() { return this._innerHTML; }
}
const body = new FakeElement();
const closeButton = new FakeElement();
const button = new FakeElement();
const row = new FakeElement();
row.dataset.detailUrl = "/tasks/a3c9f507-7473-4fa6-8d71-8786c34c6301";
row.querySelector = (selector) => selector === "[data-open-detail]" ? button : null;
const drawer = new FakeElement();
drawer.querySelector = (selector) => ({"[data-detail-body]": body, "[data-close-detail]": closeButton})[selector] || null;
const requests = [];
const popstateListeners = [];
const scrolls = [];
let failNext = false;
const history = {
state: null,
pushes: [],
replaces: [],
backCalls: 0,
pushState(state, _title, url) { this.state = state; this.pushes.push({state, url}); },
replaceState(state, _title, url) { this.state = state; this.replaces.push({state, url}); },
back() { this.backCalls++; },
};
const document = {
querySelector: (selector) => selector === "[data-start-purchases]" ? null : selector === "[data-detail-drawer]" ? drawer : null,
querySelectorAll: (selector) => selector === "[data-task-row]" ? [row] : [],
createElement: () => new FakeElement(),
contains: (element) => element === row || element === button,
};
const window = {
location: {pathname: "/tasks", search: "?status=DRAFT"},
history,
scrollY: 275,
scrollTo: (value) => scrolls.push(value),
addEventListener(type, listener) { if (type === "popstate") popstateListeners.push(listener); },
};
const context = {
AbortController,
document,
window,
fetch: async (url, options) => {
requests.push({url, options});
if (failNext) {
failNext = false;
return {ok: false, headers: {get: () => "text/html"}, text: async () => ""};
}
return {ok: true, headers: {get: () => "text/html; charset=utf-8"}, text: async () => '<article data-task-detail-content>详情</article>'};
},
};
vm.runInNewContext(source, context, {filename: "tasks.js"});
return {
body, button, closeButton, drawer, history, requests, row, scrolls,
failNextRequest: () => { failNext = true; },
popstate: (event) => { history.state = event.state; popstateListeners.forEach((listener) => listener(event)); },
flush: () => new Promise((resolve) => setImmediate(resolve)),
};
}
@@ -0,0 +1,167 @@
(() => {
"use strict";
const form = document.querySelector("[data-start-purchases]");
if (!form) return;
const all = form.querySelector("[data-select-all]");
const summary = form.querySelector("[data-selection-summary]");
const button = form.querySelector("[data-start-button]");
const feedback = form.querySelector("[data-start-feedback]");
const boxes = () => [...form.querySelectorAll("input[data-task-id]")];
let selectionFrozen = false;
const parseCents = (value) => {
const match = /^(0|[1-9]\d*)\.(\d{2})$/.exec(value);
return match ? BigInt(match[1] + match[2]) : null;
};
const refresh = () => {
const available = boxes();
const selected = available.filter((box) => box.checked);
let cents = 0n;
let pricesValid = true;
selected.forEach((box) => {
const price = parseCents(box.dataset.price);
if (price === null) pricesValid = false;
else cents += price;
});
summary.textContent = `已选 ${selected.length} 条,最高总额 ¥${cents / 100n}.${(cents % 100n).toString().padStart(2, "0")}`;
button.disabled = !selected.length || !pricesValid;
if (!pricesValid) feedback.textContent = "所选任务金额无法安全汇总,请刷新后重选。";
if (all) {
all.checked = selected.length > 0 && selected.length === available.length;
all.indeterminate = selected.length > 0 && selected.length < available.length;
all.disabled = selectionFrozen || available.length === 0;
}
};
const freezeSelection = (frozen) => {
selectionFrozen = frozen;
boxes().forEach((box) => { box.disabled = frozen; });
refresh();
};
boxes().forEach((box) => box.addEventListener("change", refresh));
if (all) all.addEventListener("change", () => { boxes().forEach((box) => { box.checked = all.checked; }); refresh(); });
let frozenPayload = null;
let inFlight = false;
form.addEventListener("submit", async (event) => {
event.preventDefault();
const selected = boxes().filter((box) => box.checked);
if (!selected.length || inFlight) return;
const tasks = selected.map((box) => ({task_id: box.dataset.taskId, expected_task_version: Number(box.dataset.taskVersion)}));
if (tasks.some((item) => !Number.isSafeInteger(item.expected_task_version) || item.expected_task_version < 1)) { feedback.textContent = "任务版本无效,请刷新后重选。"; return; }
frozenPayload = frozenPayload || JSON.stringify({start_key: form.dataset.startKey, tasks});
inFlight = true; freezeSelection(true); button.disabled = true; button.textContent = "正在授权…";
try { const response = await fetch("/tasks/start-purchases", {method:"POST", headers:{"Content-Type":"application/json", "X-CSRF-Token":form.dataset.csrf}, body:frozenPayload});
if (response.ok) { window.location.reload(); return; }
if (response.status === 409) { feedback.textContent = "任务已变化,请刷新后重选。"; frozenPayload = null; freezeSelection(false); boxes().forEach((box) => { box.checked = false; }); refresh(); return; }
if (response.status === 400 || response.status === 401 || response.status === 403) { feedback.textContent = "请求未被接受,请刷新页面后重试。"; frozenPayload = null; freezeSelection(false); return; }
feedback.textContent = "结果暂时不明确,只能使用同一按钮原样重放。";
} catch (_) { feedback.textContent = "网络结果不明确,请使用同一按钮原样重试。"; }
finally { inFlight = false; button.textContent = "开始采购(只创建待付款订单)"; if (frozenPayload) button.disabled = false; }
});
refresh();
})();
(() => {
"use strict";
const drawer = document.querySelector("[data-detail-drawer]");
if (!drawer) return;
const body = drawer.querySelector("[data-detail-body]");
const closeButton = drawer.querySelector("[data-close-detail]");
const rows = [...document.querySelectorAll("[data-task-row]")];
const initialURL = window.location.pathname + window.location.search;
let focusTrigger = null;
let scrollPosition = window.scrollY;
let activeRequest = null;
const isInteractive = (target) => Boolean(target && typeof target.closest === "function" && target.closest("a,button,input,select,textarea,label,[contenteditable=true]"));
const showDrawer = () => {
if (!drawer.open) drawer.showModal();
};
const restoreList = () => {
if (activeRequest) {
activeRequest.abort();
activeRequest = null;
}
if (drawer.open) drawer.close();
window.scrollTo({top: scrollPosition, behavior: "auto"});
if (focusTrigger && document.contains(focusTrigger)) focusTrigger.focus({preventScroll: true});
};
const showError = (url, row, requestedFocus, pushHistory) => {
body.replaceChildren();
const message = document.createElement("p");
message.className = "drawer-feedback";
message.setAttribute("role", "alert");
message.textContent = "任务详情加载失败。请重试,或在完整页打开。";
const actions = document.createElement("p");
const retry = document.createElement("button");
retry.className = "button primary";
retry.type = "button";
retry.textContent = "重试";
retry.addEventListener("click", () => loadDetail(url, row, requestedFocus, pushHistory));
const fallback = document.createElement("a");
fallback.className = "button";
fallback.href = url;
fallback.textContent = "在完整页打开";
actions.className = "actions";
actions.append(retry, fallback);
body.append(message, actions);
};
const loadDetail = async (url, row, requestedFocus, pushHistory) => {
if (activeRequest) activeRequest.abort();
const requestController = new AbortController();
activeRequest = requestController;
focusTrigger = requestedFocus || focusTrigger;
if (pushHistory) scrollPosition = window.scrollY;
body.innerHTML = '<p class="drawer-feedback" role="status">正在加载任务详情…</p>';
showDrawer();
try {
const response = await fetch(url, {headers: {"X-CMBuyer-View": "drawer", "Accept": "text/html"}, credentials: "same-origin", signal: requestController.signal});
if (!response.ok || !String(response.headers.get("Content-Type") || "").toLowerCase().startsWith("text/html")) throw new Error("detail request rejected");
const fragment = await response.text();
if (!fragment.includes("data-task-detail-content")) throw new Error("detail fragment missing");
body.innerHTML = fragment;
if (pushHistory) window.history.pushState({cmbuyerDrawer: true, detailURL: url, focusTarget: requestedFocus === row ? "row" : "button"}, "", url);
closeButton.focus();
} catch (error) {
if (error.name !== "AbortError") showError(url, row, requestedFocus, pushHistory);
} finally {
if (activeRequest === requestController) activeRequest = null;
}
};
const requestClose = () => {
if (window.history.state && window.history.state.cmbuyerDrawer) window.history.back();
else restoreList();
};
window.history.replaceState({cmbuyerList: true, listURL: initialURL}, "", initialURL);
rows.forEach((row) => {
const url = row.dataset.detailUrl;
row.addEventListener("dblclick", (event) => {
if (!isInteractive(event.target)) loadDetail(url, row, row, true);
});
row.addEventListener("keydown", (event) => {
if (event.key === "Enter" && event.target === row) {
event.preventDefault();
loadDetail(url, row, row, true);
}
});
const button = row.querySelector("[data-open-detail]");
if (button) button.addEventListener("click", () => loadDetail(url, row, button, true));
});
closeButton.addEventListener("click", requestClose);
drawer.addEventListener("cancel", (event) => {
event.preventDefault();
requestClose();
});
window.addEventListener("popstate", (event) => {
if (event.state && event.state.cmbuyerDrawer) {
const row = rows.find((candidate) => candidate.dataset.detailUrl === event.state.detailURL);
if (!row) {
restoreList();
return;
}
const requestedFocus = event.state.focusTarget === "button" ? row.querySelector("[data-open-detail]") || row : row;
loadDetail(event.state.detailURL, row, requestedFocus, false);
return;
}
restoreList();
});
})();
@@ -0,0 +1,138 @@
"use strict";
const test = require("node:test");
const assert = require("node:assert/strict");
const fs = require("node:fs");
const path = require("node:path");
const vm = require("node:vm");
const source = fs.readFileSync(path.join(__dirname, "tasks.js"), "utf8");
test("successful authorization sends numeric version and reloads", async () => {
const requests = [];
const harness = createHarness(async (_url, options) => {
requests.push(options);
return {ok: true, status: 200};
});
await harness.submit();
assert.equal(requests.length, 1);
assert.equal(requests[0].headers["Content-Type"], "application/json");
assert.equal(requests[0].headers["X-CSRF-Token"], "csrf-token");
const payload = JSON.parse(requests[0].body);
assert.equal(payload.start_key, "start-key");
assert.equal(typeof payload.tasks[0].expected_task_version, "number");
assert.equal(payload.tasks[0].expected_task_version, 7);
assert.equal(harness.reloads(), 1);
});
test("409 clears stale selection and requires a fresh choice", async () => {
const harness = createHarness(async () => ({ok: false, status: 409}));
await harness.submit();
assert.equal(harness.box.checked, false);
assert.equal(harness.box.disabled, false);
assert.equal(harness.button.disabled, true);
assert.match(harness.feedback.textContent, /任务已变化/);
});
for (const status of [400, 401, 403]) {
test(`${status} releases the frozen payload for a page refresh`, async () => {
const harness = createHarness(async () => ({ok: false, status}));
await harness.submit();
assert.equal(harness.box.checked, true);
assert.equal(harness.box.disabled, false);
assert.equal(harness.button.disabled, false);
assert.match(harness.feedback.textContent, /刷新页面后重试/);
});
}
test("5xx retries the byte-identical frozen payload", async () => {
const bodies = [];
const harness = createHarness(async (_url, options) => {
bodies.push(options.body);
return {ok: false, status: 503};
});
await harness.submit();
assert.equal(harness.box.disabled, true);
assert.equal(harness.button.disabled, false);
assert.match(harness.feedback.textContent, /原样重放/);
await harness.submit();
assert.equal(bodies.length, 2);
assert.equal(bodies[1], bodies[0]);
});
test("network ambiguity retries the same payload and can finish", async () => {
const bodies = [];
let call = 0;
const harness = createHarness(async (_url, options) => {
bodies.push(options.body);
call++;
if (call === 1) throw new Error("network result unknown");
return {ok: true, status: 200};
});
await harness.submit();
assert.equal(harness.box.disabled, true);
assert.match(harness.feedback.textContent, /原样重试/);
await harness.submit();
assert.deepEqual(bodies, [bodies[0], bodies[0]]);
assert.equal(harness.reloads(), 1);
});
function createHarness(fetchImplementation) {
class FakeElement {
constructor() {
this.dataset = {};
this.checked = false;
this.disabled = false;
this.indeterminate = false;
this.textContent = "";
this.listeners = {};
}
addEventListener(type, listener) {
this.listeners[type] = listener;
}
}
const box = new FakeElement();
box.checked = true;
box.dataset = {taskId: "task-id", taskVersion: "7", price: "12.80"};
const selectAll = new FakeElement();
const summary = new FakeElement();
const button = new FakeElement();
const feedback = new FakeElement();
const form = new FakeElement();
form.dataset = {startKey: "start-key", csrf: "csrf-token"};
form.querySelector = (selector) => ({
"[data-select-all]": selectAll,
"[data-selection-summary]": summary,
"[data-start-button]": button,
"[data-start-feedback]": feedback,
})[selector] || null;
form.querySelectorAll = (selector) => selector === "input[data-task-id]" ? [box] : [];
let reloadCount = 0;
const context = {
document: {querySelector: (selector) => selector === "[data-start-purchases]" ? form : null},
fetch: fetchImplementation,
window: {location: {reload: () => { reloadCount++; }}},
};
vm.runInNewContext(source, context, {filename: "tasks.js"});
return {
box,
button,
feedback,
reloads: () => reloadCount,
submit: () => form.listeners.submit({preventDefault() {}}),
};
}
@@ -0,0 +1,38 @@
{{define "task-detail-page.html"}}
<!doctype html>
<html lang="zh-CN">
<head>
<meta charset="utf-8">
<meta name="viewport" content="width=device-width, initial-scale=1">
<title>{{.Detail.Task.Title}} · 任务详情 · 采购服务</title>
<style>
:root{--bg:#f4f7fb;--surface:#fff;--text:#172033;--muted:#526079;--border:#cfd8e6;--primary:#155eef;--danger:#b42318;--success:#067647;--focus:#ffbf47;font-family:"Segoe UI","Microsoft YaHei UI",system-ui,sans-serif}*{box-sizing:border-box}body{margin:0;color:var(--text);background:var(--bg);font-size:16px;line-height:1.55}a{color:#124cc5;text-underline-offset:3px}:focus-visible{outline:3px solid var(--focus);outline-offset:3px}.skip{position:fixed;z-index:100;top:8px;left:8px;padding:10px;color:#fff;background:#172033;transform:translateY(-160%)}.skip:focus{transform:translateY(0)}.topbar{display:flex;align-items:center;justify-content:space-between;gap:16px;min-height:64px;padding:10px clamp(16px,4vw,40px);border-bottom:1px solid var(--border);background:var(--surface)}.brand{color:var(--text);font-weight:700;text-decoration:none}.brand b{display:inline-grid;place-items:center;width:32px;height:32px;margin-right:8px;border-radius:8px;background:var(--primary);color:#fff;font-size:.82rem}.button{display:inline-flex;align-items:center;justify-content:center;min-height:44px;padding:9px 14px;border:1px solid var(--border);border-radius:8px;color:var(--text);background:#fff;font-weight:700;text-decoration:none}.detail-page{width:min(100% - 32px,1120px);margin:28px auto 48px}.detail-shell{display:grid;gap:16px}.detail-head{display:flex;align-items:flex-start;justify-content:space-between;gap:16px}.detail-head h1{margin:0;font-size:clamp(1.45rem,3vw,2rem)}.detail-head p{margin:4px 0;color:var(--muted)}.status{display:inline-block;padding:4px 10px;border-radius:999px;background:#eaf1ff;color:#173d8f;font-size:.88rem;font-weight:700;white-space:nowrap}.safety{margin:0;padding:13px 15px;border:1px solid #a9c3f7;border-left:5px solid var(--primary);border-radius:10px;background:#edf3ff}.detail-grid{display:grid;grid-template-columns:minmax(0,1fr) minmax(250px,320px);gap:16px}.detail-card{overflow:hidden;border:1px solid var(--border);border-radius:12px;background:var(--surface)}.detail-card>header,.detail-card>.detail-body{padding:16px 18px}.detail-card>header{border-bottom:1px solid var(--border)}.detail-card h2,.detail-card h3{margin:0}.detail-card header p,.empty-note{margin:4px 0 0;color:var(--muted)}.facts{display:grid;grid-template-columns:repeat(2,minmax(0,1fr));gap:10px;margin:0}.facts div{min-width:0;padding:11px;border:1px solid var(--border);border-radius:8px;background:#f8fafc}.facts dt{font-size:.82rem;color:var(--muted);font-weight:700}.facts dd{margin:3px 0 0;overflow-wrap:anywhere;font-weight:650}.audit-list{display:grid;gap:10px;margin:0;padding:0;list-style:none}.audit-list li{padding:12px;border:1px solid var(--border);border-radius:8px}.audit-list p{margin:4px 0}.mono{font-family:Consolas,"SFMono-Regular",monospace;overflow-wrap:anywhere}.evidence-grid{display:grid;grid-template-columns:repeat(auto-fit,minmax(220px,1fr));gap:14px}.evidence{margin:0}.evidence img{display:block;width:100%;height:auto;max-height:520px;object-fit:contain;border:1px solid var(--border);border-radius:8px;background:#eef2f7}.evidence figcaption{margin-top:7px;color:var(--muted);font-size:.85rem}.section-stack{display:grid;gap:16px}.privacy-note{margin:12px 0 0;color:var(--muted);font-size:.88rem}@media(max-width:760px){.detail-grid{grid-template-columns:1fr}.detail-head{display:grid}.facts{grid-template-columns:1fr}}@media(prefers-reduced-motion:reduce){*,*::before,*::after{scroll-behavior:auto!important;transition-duration:.01ms!important;animation-duration:.01ms!important}}
</style>
</head>
<body>
<a class="skip" href="#main">跳到主要内容</a>
<header class="topbar"><a class="brand" href="/tasks"><b aria-hidden="true">采</b>采购服务</a><a class="button" href="/tasks">返回任务列表</a></header>
<main class="detail-page" id="main">{{template "task-detail-content" .}}</main>
</body>
</html>
{{end}}
{{define "task-detail-content"}}
<article class="detail-shell" data-task-detail-content data-task-id="{{.Detail.Task.ID}}">
<header class="detail-head"><div><h1>{{.Detail.Task.Title}}</h1><p>任务 <span class="mono">{{.Detail.Task.ID}}</span> · 版本 {{.Detail.Task.Version}}</p></div><span class="status">{{statusLabel .Detail.Task.Status}}</span></header>
<p class="safety"><strong>{{taskSafetyTitle .Detail.Task.Status}}</strong> {{taskSafetyText .Detail.Task.Status}}</p>
<div class="detail-grid">
<div class="section-stack">
<section class="detail-card" aria-labelledby="task-facts-title"><header><h2 id="task-facts-title">任务要求</h2><p>管理员锁定的采购边界;详情页不会触发设备动作。</p></header><div class="detail-body"><dl class="facts"><div><dt>商品</dt><dd><a href="{{canonicalURL .Detail.Task.GoodsID}}" target="_blank" rel="noopener noreferrer">goods_id {{.Detail.Task.GoodsID}}</a></dd></div><div><dt>目标规格</dt><dd>{{.Detail.Task.SKUColor}} / {{.Detail.Task.SKUSize}}</dd></div><div><dt>数量</dt><dd>{{.Detail.Task.Quantity}} 件</dd></div><div><dt>最高总价</dt><dd>¥{{.Detail.Task.MaxTotalPrice}}</dd></div><div><dt>创建时间(上海)</dt><dd><time datetime="{{shanghaiDateTime .Detail.Task.CreatedAt}}">{{shanghaiTime .Detail.Task.CreatedAt}}</time></dd></div><div><dt>更新时间(上海)</dt><dd><time datetime="{{shanghaiDateTime .Detail.Task.UpdatedAt}}">{{shanghaiTime .Detail.Task.UpdatedAt}}</time></dd></div></dl></div></section>
<section class="detail-card" aria-labelledby="execution-title"><header><h2 id="execution-title">设备执行事实</h2><p>只展示数据库中已存在的 attempt;T-204 不创建执行记录。</p></header><div class="detail-body">{{if .Detail.Attempts}}<ol class="audit-list">{{range .Detail.Attempts}}<li><h3>Attempt <span class="mono">{{.ID}}</span></h3><p>状态:{{attemptStatusLabel .Status}} · 领取代次 {{.ClaimGeneration}}</p><p>开始:<time datetime="{{shanghaiDateTime .StartedAt}}">{{shanghaiTime .StartedAt}}</time>{{with .FinishedAt}} · 结束:<time datetime="{{shanghaiDateTime .}}">{{shanghaiTime .}}</time>{{end}}</p>{{with .FailureCode}}<p>失败码:<span class="mono">{{.}}</span></p>{{end}}{{if or .Gate1UnitPrice .Gate2UnitPrice .QuantityRead .ConfirmAmount}}<p>已有读数:{{with .Gate1UnitPrice}}闸门一 ¥{{.}};{{end}}{{with .Gate2UnitPrice}}闸门二 ¥{{.}};{{end}}{{with .QuantityRead}}数量 {{.}};{{end}}{{with .ConfirmAmount}}确认页 ¥{{.}}{{end}}</p>{{else}}<p class="empty-note">暂无规格、价格或数量读数。</p>{{end}}</li>{{end}}</ol>{{else}}<p class="empty-note">暂无设备执行记录。</p>{{end}}</div></section>
<section class="detail-card" aria-labelledby="evidence-title"><header><h2 id="evidence-title">内部截图</h2><p>INTERNAL_RAW 仅供已登录管理员审计,不代表价格闸门通过或人工批准。</p></header><div class="detail-body">{{if .Detail.Evidence}}<div class="evidence-grid">{{range .Detail.Evidence}}<figure class="evidence"><img src="/evidence/{{.ID}}" width="{{.Width}}" height="{{.Height}}" loading="lazy" alt="规格面板内部审计截图,采集于 {{shanghaiTime .CapturedAt}}"><figcaption>{{evidenceKindLabel .Kind}}(<span class="mono">{{.Kind}}</span>)· {{formatBytes .ByteSize}} · <time datetime="{{shanghaiDateTime .CapturedAt}}">{{shanghaiTime .CapturedAt}}</time><br>Attempt <span class="mono">{{.AttemptID}}</span></figcaption></figure>{{end}}</div>{{else}}<p class="empty-note">暂无内部截图。只有已认证设备显式上传的 PNG 会出现在这里。</p>{{end}}<p class="privacy-note">截图可能包含页面已显示的地址或手机号;系统不提取、索引或写入日志。完整 XML、外部支付页和支付凭据不会上传。</p></div></section>
<section class="detail-card" aria-labelledby="submission-title"><header><h2 id="submission-title">提交围栏与结果</h2><p>只读审计;本页没有重试、再次提交或付款动作。</p></header><div class="detail-body">{{if .Detail.Submissions}}<ol class="audit-list">{{range .Detail.Submissions}}<li><h3>Submission <span class="mono">{{.ID}}</span></h3><p>状态:{{submissionStatusLabel .Status}}</p><p>闸门一 ¥{{.Gate1UnitPrice}};闸门二 ¥{{.Gate2UnitPrice}};数量 {{.QuantityRead}};确认页 ¥{{.ConfirmAmount}}</p><p>建立:<time datetime="{{shanghaiDateTime .CreatedAt}}">{{shanghaiTime .CreatedAt}}</time>{{with .ResolvedAt}} · 调和:<time datetime="{{shanghaiDateTime .}}">{{shanghaiTime .}}</time>{{end}}</p></li>{{end}}</ol>{{else}}<p class="empty-note">尚未建立提交围栏;详情页不会创建或释放围栏。</p>{{end}}</div></section>
</div>
<aside class="section-stack" aria-label="任务状态摘要"><section class="detail-card"><header><h2>开始采购授权</h2><p>锁定任务字段和最高总价,不授权付款。</p></header><div class="detail-body">{{if .Detail.Authorizations}}<ol class="audit-list">{{range .Detail.Authorizations}}<li><h3>{{authorizationStatusLabel .Status}}</h3><p class="mono">{{.ID}}</p><p>任务版本 {{.TaskVersion}} · 上限 ¥{{.TotalPriceCap}}</p><p>授权人:{{.CreatedBy}}</p><p><time datetime="{{shanghaiDateTime .CreatedAt}}">{{shanghaiTime .CreatedAt}}</time> 至 <time datetime="{{shanghaiDateTime .ExpiresAt}}">{{shanghaiTime .ExpiresAt}}</time></p></li>{{end}}</ol>{{else}}<p class="empty-note">尚未开始采购,没有授权记录。</p>{{end}}</div></section><section class="detail-card"><header><h2>固定边界</h2></header><div class="detail-body"><ul><li>系统只创建待付款订单,不自动付款。</li><li>截图仅供审计,不替代实时三道价格闸门。</li><li>围栏后只能调和同一提交,禁止再次点击。</li></ul></div></section></aside>
</div>
</article>
{{end}}
File diff suppressed because one or more lines are too long
+129 -3
View File
@@ -3,14 +3,37 @@ package webui
import (
"embed"
"fmt"
"html/template"
"io"
"time"
"cmbuyer/admin/internal/taskdetail"
"cmbuyer/admin/internal/tasks"
)
//go:embed templates/*.html
var templateFiles embed.FS
var templates = template.Must(template.New("webui").ParseFS(templateFiles, "templates/*.html"))
//go:embed static/tasks.js
var tasksScript []byte
var shanghaiLocation = time.FixedZone("Asia/Shanghai", 8*60*60)
var templates = template.Must(template.New("webui").Funcs(template.FuncMap{
"list": func(values ...any) []any { return values },
"statusLabel": statusLabel,
"shanghaiDateTime": func(value time.Time) string { return value.In(shanghaiLocation).Format(time.RFC3339) },
"shanghaiTime": func(value time.Time) string { return value.In(shanghaiLocation).Format("2006-01-02 15:04") },
"canonicalURL": tasks.CanonicalURL,
"formatBytes": formatBytes,
"taskSafetyTitle": taskSafetyTitle,
"taskSafetyText": taskSafetyText,
"authorizationStatusLabel": authorizationStatusLabel,
"attemptStatusLabel": attemptStatusLabel,
"submissionStatusLabel": submissionStatusLabel,
"evidenceKindLabel": evidenceKindLabel,
}).ParseFS(templateFiles, "templates/*.html"))
// LoginData 是登录页面所需的非敏感展示数据。
type LoginData struct {
@@ -20,17 +43,120 @@ type LoginData struct {
Error string
}
// TasksData 是当前受保护任务空壳所需的数据。任务字段将在后续任务实现。
// TasksData 是受保护的建单与任务工作台页面所需数据。
type TasksData struct {
CSRFToken string
Tasks []tasks.TaskRow
Filter tasks.TaskFilter
FilterErrors tasks.Errors
HasFilter bool
StartKey string
Form tasks.Form
Errors tasks.Errors
OpenForm bool
FullPage bool
FocusField string
Success bool
}
type TaskDetailData struct{ Detail taskdetail.Detail }
// RenderLogin 写入登录页。
func RenderLogin(writer io.Writer, data LoginData) error {
return templates.ExecuteTemplate(writer, "login.html", data)
}
// RenderTasks 写入登录后的受保护空壳。
// RenderTasks 写入登录后的受保护任务页。
func RenderTasks(writer io.Writer, data TasksData) error {
return templates.ExecuteTemplate(writer, "tasks.html", data)
}
func RenderTaskDetailPage(writer io.Writer, data TaskDetailData) error {
return templates.ExecuteTemplate(writer, "task-detail-page.html", data)
}
func RenderTaskDetailFragment(writer io.Writer, data TaskDetailData) error {
return templates.ExecuteTemplate(writer, "task-detail-content", data)
}
func TasksScript() []byte { return tasksScript }
func statusLabel(status string) string {
labels := map[string]string{
"DRAFT": "待开始",
"PENDING": "已授权待领取",
"CLAIMED": "已领取",
"ORDERING": "执行中",
"NEEDS_MANUAL": "待人工处理",
"WAITING_PAYMENT": "待付款",
"RECONCILIATION_REQUIRED": "围栏后待调和",
"SUCCEEDED": "已完成",
"FAILED": "失败",
"CANCELED": "已取消",
}
if label, ok := labels[status]; ok {
return label
}
return "未知状态"
}
func taskSafetyTitle(status string) string {
if status == "WAITING_PAYMENT" {
return "订单已创建,系统尚未付款。"
}
if status == "RECONCILIATION_REQUIRED" {
return "订单可能已创建,只能调和同一提交。"
}
return "系统只创建待付款订单,不会自动付款。"
}
func taskSafetyText(status string) string {
if status == "DRAFT" {
return "创建任务不构成授权;请回到列表勾选后开始采购。"
}
if status == "RECONCILIATION_REQUIRED" {
return "围栏保持占用,禁止重新授权、再次提交或释放。"
}
return "截图只供内部审计,不替代实时价格闸门,也不会触发设备动作。"
}
func authorizationStatusLabel(status string) string {
labels := map[string]string{"ACTIVE": "授权有效", "CLAIMED": "已被领取", "FENCED": "提交围栏已建立", "CONSUMED": "授权已消费", "EXPIRED": "授权已过期", "ABANDONED": "授权已关闭"}
if value, ok := labels[status]; ok {
return value
}
return "未知授权状态"
}
func attemptStatusLabel(status string) string {
labels := map[string]string{"CLAIMED": "已领取", "ORDERING": "执行中", "FAILED": "围栏前失败", "FENCED": "已建立围栏", "ABANDONED": "已安全停止"}
if value, ok := labels[status]; ok {
return value
}
return "未知执行状态"
}
func submissionStatusLabel(status string) string {
labels := map[string]string{"FENCED": "围栏已建立", "SUBMITTED": "已创建待付款订单", "RECONCILIATION_REQUIRED": "结果待调和", "MANUAL_RESOLVED": "已人工调和"}
if value, ok := labels[status]; ok {
return value
}
return "未知提交状态"
}
func evidenceKindLabel(kind string) string {
if kind == "SKU_PANEL_GATE_1" {
return "规格面板 · 闸门一"
}
return "内部截图"
}
func formatBytes(value int64) string {
if value >= 1<<20 {
return fmt.Sprintf("%.1f MiB", float64(value)/(1<<20))
}
if value >= 1<<10 {
return fmt.Sprintf("%.1f KiB", float64(value)/(1<<10))
}
return fmt.Sprintf("%d B", value)
}
@@ -0,0 +1,333 @@
-- +goose Up
-- v1 的试选/锁旧价记录无法安全推断为单趟执行事实。先在同一事务中拒绝它们,
-- 避免删除审计数据后再尝试猜测映射。
CREATE TABLE single_pass_upgrade_guard (
valid INTEGER NOT NULL CHECK (valid = 1)
);
INSERT INTO single_pass_upgrade_guard (valid)
SELECT CASE WHEN
(SELECT COUNT(*) FROM spec_trials) = 0
AND (SELECT COUNT(*) FROM order_authorizations) = 0
AND (SELECT COUNT(*) FROM order_submissions) = 0
AND (SELECT COUNT(*) FROM tasks WHERE source <> 'MANUAL' OR status <> 'DRAFT') = 0
-- v2 的金额边界是严格正数;不把 v1 中不能无损纳入该边界的数据悄悄改写。
AND (SELECT COUNT(*) FROM tasks WHERE
max_total_price = ''
OR max_total_price GLOB '*[^0-9.]*'
OR length(max_total_price) - length(replace(max_total_price, '.', '')) > 1
OR max_total_price = '.'
OR (instr(max_total_price, '.') > 0 AND (
instr(max_total_price, '.') = 1
OR length(max_total_price) = instr(max_total_price, '.')
OR length(max_total_price) - instr(max_total_price, '.') > 2
))
OR replace(replace(max_total_price, '.', ''), '0', '') = ''
) = 0
THEN 1 ELSE 0 END;
DROP TABLE single_pass_upgrade_guard;
ALTER TABLE tasks RENAME TO tasks_v1;
DROP TABLE order_submissions;
DROP TABLE order_authorizations;
DROP TABLE spec_trials;
CREATE TABLE tasks (
id TEXT PRIMARY KEY,
source TEXT NOT NULL CHECK (source IN ('MANUAL', 'EXCEL', 'ERP')),
source_ref TEXT,
title TEXT NOT NULL,
goods_id TEXT NOT NULL,
sku_color TEXT NOT NULL,
sku_size TEXT NOT NULL,
quantity INTEGER NOT NULL CHECK (quantity > 0 AND typeof(quantity) = 'integer'),
max_total_price TEXT NOT NULL CHECK (
max_total_price <> ''
AND max_total_price NOT GLOB '*[^0-9.]*'
AND length(max_total_price) - length(replace(max_total_price, '.', '')) <= 1
AND max_total_price <> '.'
AND (instr(max_total_price, '.') = 0 OR (
instr(max_total_price, '.') > 1
AND length(max_total_price) > instr(max_total_price, '.')
AND length(max_total_price) - instr(max_total_price, '.') <= 2
))
AND replace(replace(max_total_price, '.', ''), '0', '') <> ''
),
reference_asset_id TEXT,
status TEXT NOT NULL CHECK (status IN (
'DRAFT', 'PENDING', 'CLAIMED', 'ORDERING', 'NEEDS_MANUAL', 'WAITING_PAYMENT',
'RECONCILIATION_REQUIRED', 'SUCCEEDED', 'FAILED', 'CANCELED'
)),
version INTEGER NOT NULL DEFAULT 1 CHECK (version > 0 AND typeof(version) = 'integer'),
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL
);
INSERT INTO tasks (
id, source, source_ref, title, goods_id, sku_color, sku_size, quantity, max_total_price,
reference_asset_id, status, version, created_at, updated_at
)
SELECT
id, source, source_ref, title, goods_id, sku_color, sku_size, quantity, max_total_price,
reference_asset_id, status, version, created_at, updated_at
FROM tasks_v1;
DROP TABLE tasks_v1;
CREATE TABLE order_authorizations (
id TEXT PRIMARY KEY,
task_id TEXT NOT NULL REFERENCES tasks(id),
task_version INTEGER NOT NULL CHECK (task_version > 0 AND typeof(task_version) = 'integer'),
start_key TEXT NOT NULL,
goods_id TEXT NOT NULL,
sku_color TEXT NOT NULL,
sku_size TEXT NOT NULL,
quantity INTEGER NOT NULL CHECK (quantity > 0 AND typeof(quantity) = 'integer'),
total_price_cap TEXT NOT NULL CHECK (
total_price_cap <> ''
AND total_price_cap NOT GLOB '*[^0-9.]*'
AND length(total_price_cap) - length(replace(total_price_cap, '.', '')) <= 1
AND total_price_cap <> '.'
AND (instr(total_price_cap, '.') = 0 OR (
instr(total_price_cap, '.') > 1
AND length(total_price_cap) > instr(total_price_cap, '.')
AND length(total_price_cap) - instr(total_price_cap, '.') <= 2
))
AND replace(replace(total_price_cap, '.', ''), '0', '') <> ''
),
status TEXT NOT NULL CHECK (status IN ('ACTIVE', 'CLAIMED', 'FENCED', 'CONSUMED', 'EXPIRED', 'ABANDONED')),
created_by TEXT NOT NULL,
created_at TEXT NOT NULL,
expires_at TEXT NOT NULL,
UNIQUE (task_id, task_version),
UNIQUE (start_key, task_id),
UNIQUE (task_id, id)
);
CREATE TABLE purchase_attempts (
id TEXT PRIMARY KEY,
task_id TEXT NOT NULL,
authorization_id TEXT NOT NULL,
claim_generation INTEGER NOT NULL CHECK (claim_generation > 0 AND typeof(claim_generation) = 'integer'),
status TEXT NOT NULL CHECK (status IN ('CLAIMED', 'ORDERING', 'FAILED', 'FENCED', 'ABANDONED')),
gate1_unit_price TEXT CHECK (
gate1_unit_price IS NULL OR (
gate1_unit_price <> ''
AND gate1_unit_price NOT GLOB '*[^0-9.]*'
AND length(gate1_unit_price) - length(replace(gate1_unit_price, '.', '')) <= 1
AND gate1_unit_price <> '.'
AND (instr(gate1_unit_price, '.') = 0 OR (
instr(gate1_unit_price, '.') > 1
AND length(gate1_unit_price) > instr(gate1_unit_price, '.')
AND length(gate1_unit_price) - instr(gate1_unit_price, '.') <= 2
))
AND replace(replace(gate1_unit_price, '.', ''), '0', '') <> ''
)
),
gate2_unit_price TEXT CHECK (
gate2_unit_price IS NULL OR (
gate2_unit_price <> ''
AND gate2_unit_price NOT GLOB '*[^0-9.]*'
AND length(gate2_unit_price) - length(replace(gate2_unit_price, '.', '')) <= 1
AND gate2_unit_price <> '.'
AND (instr(gate2_unit_price, '.') = 0 OR (
instr(gate2_unit_price, '.') > 1
AND length(gate2_unit_price) > instr(gate2_unit_price, '.')
AND length(gate2_unit_price) - instr(gate2_unit_price, '.') <= 2
))
AND replace(replace(gate2_unit_price, '.', ''), '0', '') <> ''
)
),
quantity_read INTEGER CHECK (quantity_read IS NULL OR (quantity_read > 0 AND typeof(quantity_read) = 'integer')),
confirm_amount TEXT CHECK (
confirm_amount IS NULL OR (
confirm_amount <> ''
AND confirm_amount NOT GLOB '*[^0-9.]*'
AND length(confirm_amount) - length(replace(confirm_amount, '.', '')) <= 1
AND confirm_amount <> '.'
AND (instr(confirm_amount, '.') = 0 OR (
instr(confirm_amount, '.') > 1
AND length(confirm_amount) > instr(confirm_amount, '.')
AND length(confirm_amount) - instr(confirm_amount, '.') <= 2
))
AND replace(replace(confirm_amount, '.', ''), '0', '') <> ''
)
),
failure_code TEXT CHECK (failure_code IS NULL OR failure_code IN (
'AUTHORIZATION_EXPIRED', 'LEASE_LOST', 'GATE_1_REJECTED', 'QUANTITY_MISMATCH',
'GATE_2_REJECTED', 'GATE_3_REJECTED', 'FENCE_REJECTED', 'SAFE_ABORTED'
)),
started_at TEXT NOT NULL,
finished_at TEXT,
UNIQUE (task_id, claim_generation),
UNIQUE (task_id, id),
UNIQUE (task_id, authorization_id, id),
FOREIGN KEY (task_id, authorization_id) REFERENCES order_authorizations(task_id, id)
);
CREATE TABLE order_submissions (
id TEXT PRIMARY KEY,
task_id TEXT NOT NULL,
authorization_id TEXT NOT NULL,
attempt_id TEXT NOT NULL,
status TEXT NOT NULL CHECK (status IN ('FENCED', 'SUBMITTED', 'RECONCILIATION_REQUIRED', 'MANUAL_RESOLVED')),
gate1_unit_price TEXT NOT NULL CHECK (
gate1_unit_price <> ''
AND gate1_unit_price NOT GLOB '*[^0-9.]*'
AND length(gate1_unit_price) - length(replace(gate1_unit_price, '.', '')) <= 1
AND gate1_unit_price <> '.'
AND (instr(gate1_unit_price, '.') = 0 OR (
instr(gate1_unit_price, '.') > 1
AND length(gate1_unit_price) > instr(gate1_unit_price, '.')
AND length(gate1_unit_price) - instr(gate1_unit_price, '.') <= 2
))
AND replace(replace(gate1_unit_price, '.', ''), '0', '') <> ''
),
gate2_unit_price TEXT NOT NULL CHECK (
gate2_unit_price <> ''
AND gate2_unit_price NOT GLOB '*[^0-9.]*'
AND length(gate2_unit_price) - length(replace(gate2_unit_price, '.', '')) <= 1
AND gate2_unit_price <> '.'
AND (instr(gate2_unit_price, '.') = 0 OR (
instr(gate2_unit_price, '.') > 1
AND length(gate2_unit_price) > instr(gate2_unit_price, '.')
AND length(gate2_unit_price) - instr(gate2_unit_price, '.') <= 2
))
AND replace(replace(gate2_unit_price, '.', ''), '0', '') <> ''
),
quantity_read INTEGER NOT NULL CHECK (quantity_read > 0 AND typeof(quantity_read) = 'integer'),
confirm_amount TEXT NOT NULL CHECK (
confirm_amount <> ''
AND confirm_amount NOT GLOB '*[^0-9.]*'
AND length(confirm_amount) - length(replace(confirm_amount, '.', '')) <= 1
AND confirm_amount <> '.'
AND (instr(confirm_amount, '.') = 0 OR (
instr(confirm_amount, '.') > 1
AND length(confirm_amount) > instr(confirm_amount, '.')
AND length(confirm_amount) - instr(confirm_amount, '.') <= 2
))
AND replace(replace(confirm_amount, '.', ''), '0', '') <> ''
),
created_at TEXT NOT NULL,
resolved_at TEXT,
UNIQUE (authorization_id),
UNIQUE (attempt_id),
FOREIGN KEY (task_id, authorization_id, attempt_id) REFERENCES purchase_attempts(task_id, authorization_id, id)
);
-- +goose Down
-- 只有尚未产生任何单趟授权或执行事实的纯 MANUAL/DRAFT 数据才能无损回到 v1。
CREATE TABLE single_pass_downgrade_guard (
valid INTEGER NOT NULL CHECK (valid = 1)
);
INSERT INTO single_pass_downgrade_guard (valid)
SELECT CASE WHEN
(SELECT COUNT(*) FROM order_authorizations) = 0
AND (SELECT COUNT(*) FROM purchase_attempts) = 0
AND (SELECT COUNT(*) FROM order_submissions) = 0
AND (SELECT COUNT(*) FROM tasks WHERE source <> 'MANUAL' OR status <> 'DRAFT') = 0
THEN 1 ELSE 0 END;
DROP TABLE single_pass_downgrade_guard;
ALTER TABLE tasks RENAME TO tasks_v2;
DROP TABLE order_submissions;
DROP TABLE purchase_attempts;
DROP TABLE order_authorizations;
CREATE TABLE tasks (
id TEXT PRIMARY KEY,
source TEXT NOT NULL CHECK (source IN ('MANUAL', 'EXCEL', 'ERP')),
source_ref TEXT,
title TEXT NOT NULL,
goods_id TEXT NOT NULL,
sku_color TEXT NOT NULL,
sku_size TEXT NOT NULL,
quantity INTEGER NOT NULL CHECK (quantity > 0 AND typeof(quantity) = 'integer'),
max_total_price TEXT NOT NULL CHECK (
max_total_price <> ''
AND max_total_price NOT GLOB '*[^0-9.]*'
AND length(max_total_price) - length(replace(max_total_price, '.', '')) <= 1
AND max_total_price <> '.'
AND (instr(max_total_price, '.') = 0 OR (
instr(max_total_price, '.') > 1
AND length(max_total_price) > instr(max_total_price, '.')
AND length(max_total_price) - instr(max_total_price, '.') <= 2
))
),
reference_asset_id TEXT,
status TEXT NOT NULL CHECK (status IN (
'DRAFT', 'PENDING', 'CLAIMED', 'RUNNING', 'WAITING_CONFIRMATION',
'PENDING_RETRIAL', 'AUTHORIZED', 'ORDERING', 'WAITING_PAYMENT',
'RECONCILIATION_REQUIRED', 'NEEDS_MANUAL', 'SUCCEEDED', 'CANCELED'
)),
version INTEGER NOT NULL DEFAULT 1 CHECK (version > 0 AND typeof(version) = 'integer'),
created_at TEXT NOT NULL,
updated_at TEXT NOT NULL
);
INSERT INTO tasks (
id, source, source_ref, title, goods_id, sku_color, sku_size, quantity, max_total_price,
reference_asset_id, status, version, created_at, updated_at
)
SELECT
id, source, source_ref, title, goods_id, sku_color, sku_size, quantity, max_total_price,
reference_asset_id, status, version, created_at, updated_at
FROM tasks_v2;
DROP TABLE tasks_v2;
CREATE TABLE spec_trials (
id TEXT PRIMARY KEY,
task_id TEXT NOT NULL REFERENCES tasks(id),
attempt INTEGER NOT NULL CHECK (attempt > 0 AND typeof(attempt) = 'integer'),
product_title TEXT NOT NULL,
selected_color TEXT NOT NULL,
selected_size TEXT NOT NULL,
unit_price TEXT NOT NULL CHECK (unit_price <> '' AND unit_price NOT GLOB '*[^0-9.]*' AND length(unit_price) - length(replace(unit_price, '.', '')) <= 1 AND unit_price <> '.' AND (instr(unit_price, '.') = 0 OR (instr(unit_price, '.') > 1 AND length(unit_price) > instr(unit_price, '.') AND length(unit_price) - instr(unit_price, '.') <= 2))),
total_price TEXT NOT NULL CHECK (total_price <> '' AND total_price NOT GLOB '*[^0-9.]*' AND length(total_price) - length(replace(total_price, '.', '')) <= 1 AND total_price <> '.' AND (instr(total_price, '.') = 0 OR (instr(total_price, '.') > 1 AND length(total_price) > instr(total_price, '.') AND length(total_price) - instr(total_price, '.') <= 2))),
evidence_sha256 TEXT NOT NULL,
created_at TEXT NOT NULL,
UNIQUE (task_id, attempt),
UNIQUE (task_id, id)
);
CREATE TABLE order_authorizations (
id TEXT PRIMARY KEY,
task_id TEXT NOT NULL REFERENCES tasks(id),
spec_trial_id TEXT NOT NULL REFERENCES spec_trials(id),
version INTEGER NOT NULL CHECK (version > 0 AND typeof(version) = 'integer'),
goods_id TEXT NOT NULL,
sku_color TEXT NOT NULL,
sku_size TEXT NOT NULL,
quantity INTEGER NOT NULL CHECK (quantity > 0 AND typeof(quantity) = 'integer'),
authorized_unit_price TEXT NOT NULL CHECK (authorized_unit_price <> '' AND authorized_unit_price NOT GLOB '*[^0-9.]*' AND length(authorized_unit_price) - length(replace(authorized_unit_price, '.', '')) <= 1 AND authorized_unit_price <> '.' AND (instr(authorized_unit_price, '.') = 0 OR (instr(authorized_unit_price, '.') > 1 AND length(authorized_unit_price) > instr(authorized_unit_price, '.') AND length(authorized_unit_price) - instr(authorized_unit_price, '.') <= 2))),
total_price_cap TEXT NOT NULL CHECK (total_price_cap <> '' AND total_price_cap NOT GLOB '*[^0-9.]*' AND length(total_price_cap) - length(replace(total_price_cap, '.', '')) <= 1 AND total_price_cap <> '.' AND (instr(total_price_cap, '.') = 0 OR (instr(total_price_cap, '.') > 1 AND length(total_price_cap) > instr(total_price_cap, '.') AND length(total_price_cap) - instr(total_price_cap, '.') <= 2))),
note TEXT,
status TEXT NOT NULL CHECK (status IN ('PENDING_DELIVERY', 'DELIVERED', 'ACKNOWLEDGED', 'EXECUTING', 'FENCED', 'CONSUMED', 'SUPERSEDED', 'EXPIRED')),
created_by TEXT NOT NULL,
created_at TEXT NOT NULL,
expires_at TEXT NOT NULL,
UNIQUE (task_id, version),
UNIQUE (task_id, id),
FOREIGN KEY (task_id, spec_trial_id) REFERENCES spec_trials(task_id, id)
);
CREATE TABLE order_submissions (
id TEXT PRIMARY KEY,
task_id TEXT NOT NULL REFERENCES tasks(id),
authorization_id TEXT NOT NULL REFERENCES order_authorizations(id),
command_id TEXT NOT NULL,
dry_run_id TEXT NOT NULL,
status TEXT NOT NULL CHECK (status IN ('FENCED', 'SUBMITTED', 'RECONCILIATION_REQUIRED', 'MANUAL_RESOLVED')),
verified_unit_price TEXT NOT NULL CHECK (verified_unit_price <> '' AND verified_unit_price NOT GLOB '*[^0-9.]*' AND length(verified_unit_price) - length(replace(verified_unit_price, '.', '')) <= 1 AND verified_unit_price <> '.' AND (instr(verified_unit_price, '.') = 0 OR (instr(verified_unit_price, '.') > 1 AND length(verified_unit_price) > instr(verified_unit_price, '.') AND length(verified_unit_price) - instr(verified_unit_price, '.') <= 2))),
quantity_read INTEGER NOT NULL CHECK (quantity_read > 0 AND typeof(quantity_read) = 'integer'),
confirm_page_amount TEXT NOT NULL CHECK (confirm_page_amount <> '' AND confirm_page_amount NOT GLOB '*[^0-9.]*' AND length(confirm_page_amount) - length(replace(confirm_page_amount, '.', '')) <= 1 AND confirm_page_amount <> '.' AND (instr(confirm_page_amount, '.') = 0 OR (instr(confirm_page_amount, '.') > 1 AND length(confirm_page_amount) > instr(confirm_page_amount, '.') AND length(confirm_page_amount) - instr(confirm_page_amount, '.') <= 2))),
created_at TEXT NOT NULL,
resolved_at TEXT,
UNIQUE (authorization_id),
UNIQUE (command_id),
FOREIGN KEY (task_id, authorization_id) REFERENCES order_authorizations(task_id, id)
);
@@ -0,0 +1,53 @@
-- +goose Up
CREATE TABLE evidence_assets (
id TEXT PRIMARY KEY,
upload_key TEXT NOT NULL,
task_id TEXT NOT NULL,
attempt_id TEXT NOT NULL,
kind TEXT NOT NULL CHECK (kind = 'SKU_PANEL_GATE_1'),
privacy_tier TEXT NOT NULL CHECK (privacy_tier = 'INTERNAL_RAW'),
sha256 TEXT NOT NULL CHECK (
length(sha256) = 64
AND sha256 NOT GLOB '*[^0-9a-f]*'
),
byte_size INTEGER NOT NULL CHECK (
typeof(byte_size) = 'integer'
AND byte_size > 0
AND byte_size <= 10485760
),
content_type TEXT NOT NULL CHECK (content_type = 'image/png'),
width_px INTEGER NOT NULL CHECK (
typeof(width_px) = 'integer'
AND width_px > 0
AND width_px <= 8192
),
height_px INTEGER NOT NULL CHECK (
typeof(height_px) = 'integer'
AND height_px > 0
AND height_px <= 8192
),
storage_key TEXT NOT NULL CHECK (
storage_key = substr(sha256, 1, 2) || '/' || sha256 || '.png'
),
uploaded_by_device_id TEXT NOT NULL CHECK (trim(uploaded_by_device_id) <> ''),
captured_at TEXT NOT NULL CHECK (trim(captured_at) <> ''),
created_at TEXT NOT NULL CHECK (trim(created_at) <> ''),
CHECK (width_px * height_px <= 16777216),
UNIQUE (uploaded_by_device_id, upload_key),
FOREIGN KEY (task_id, attempt_id) REFERENCES purchase_attempts(task_id, id)
);
CREATE INDEX evidence_assets_task_time_idx
ON evidence_assets (task_id, captured_at, created_at, id);
-- +goose Down
-- 已写入的内部原图是审计事实,回滚迁移不得静默删除它们。
CREATE TABLE evidence_downgrade_guard (
valid INTEGER NOT NULL CHECK (valid = 1)
);
INSERT INTO evidence_downgrade_guard (valid)
SELECT CASE WHEN (SELECT COUNT(*) FROM evidence_assets) = 0 THEN 1 ELSE 0 END;
DROP TABLE evidence_downgrade_guard;
DROP TABLE evidence_assets;
@@ -0,0 +1,62 @@
-- +goose Up
CREATE TABLE device_credentials (
device_id TEXT PRIMARY KEY CHECK (
length(device_id) = 36
AND substr(device_id, 9, 1) = '-'
AND substr(device_id, 14, 1) = '-'
AND substr(device_id, 19, 1) = '-'
AND substr(device_id, 24, 1) = '-'
AND length(replace(device_id, '-', '')) = 32
AND replace(device_id, '-', '') NOT GLOB '*[^0-9a-f]*'
AND substr(device_id, 15, 1) = '4'
AND substr(device_id, 20, 1) IN ('8', '9', 'a', 'b')
),
display_name TEXT NOT NULL CHECK (
display_name = trim(display_name)
AND length(display_name) BETWEEN 1 AND 128
),
token_sha256 BLOB NOT NULL UNIQUE CHECK (
typeof(token_sha256) = 'blob'
AND length(token_sha256) = 32
),
status TEXT NOT NULL CHECK (status IN ('ACTIVE', 'REVOKED')),
created_at TEXT NOT NULL CHECK (
created_at = trim(created_at)
AND length(created_at) >= 20
AND substr(created_at, 11, 1) = 'T'
AND substr(created_at, -1, 1) = 'Z'
AND julianday(created_at) IS NOT NULL
),
revoked_at TEXT CHECK (
revoked_at IS NULL OR (
revoked_at = trim(revoked_at)
AND length(revoked_at) >= 20
AND substr(revoked_at, 11, 1) = 'T'
AND substr(revoked_at, -1, 1) = 'Z'
AND julianday(revoked_at) IS NOT NULL
)
),
CHECK (
(status = 'ACTIVE' AND revoked_at IS NULL)
OR (
status = 'REVOKED'
AND revoked_at IS NOT NULL
AND julianday(revoked_at) >= julianday(created_at)
)
)
);
CREATE INDEX device_credentials_status_created_idx
ON device_credentials (status, created_at, device_id);
-- +goose Down
-- 已签发凭据是安全配置;回滚不得静默删除并让设备身份审计链消失。
CREATE TABLE device_credentials_downgrade_guard (
valid INTEGER NOT NULL CHECK (valid = 1)
);
INSERT INTO device_credentials_downgrade_guard (valid)
SELECT CASE WHEN (SELECT COUNT(*) FROM device_credentials) = 0 THEN 1 ELSE 0 END;
DROP TABLE device_credentials_downgrade_guard;
DROP TABLE device_credentials;
+272
View File
@@ -0,0 +1,272 @@
-- +goose Up
-- v4 中的 attempt、submission 或证据没有设备/session/租约归属,不能安全猜测成 claim。
-- 在同一迁移事务中拒绝这类数据库,避免补出虚假的所有权审计链。
CREATE TABLE task_claim_upgrade_guard (
valid INTEGER NOT NULL CHECK (valid = 1)
);
INSERT INTO task_claim_upgrade_guard (valid)
SELECT CASE WHEN
(SELECT COUNT(*) FROM purchase_attempts) = 0
AND (SELECT COUNT(*) FROM order_submissions) = 0
AND (SELECT COUNT(*) FROM evidence_assets) = 0
THEN 1 ELSE 0 END;
DROP TABLE task_claim_upgrade_guard;
-- 该唯一索引把“一条授权只能产生一个 attempt”下沉到数据库;应用层检查不能替代它。
CREATE UNIQUE INDEX purchase_attempts_one_per_authorization_idx
ON purchase_attempts (authorization_id);
-- claim_generation 是 attempt lineage 的组成部分,不能只在应用层比较。
CREATE UNIQUE INDEX purchase_attempts_claim_lineage_idx
ON purchase_attempts (task_id, authorization_id, id, claim_generation);
CREATE TABLE purchase_attempt_claims (
attempt_id TEXT PRIMARY KEY,
task_id TEXT NOT NULL,
authorization_id TEXT NOT NULL UNIQUE,
claimed_by_device_id TEXT NOT NULL,
session_id TEXT NOT NULL CHECK (
length(session_id) = 36
AND substr(session_id, 9, 1) = '-'
AND substr(session_id, 14, 1) = '-'
AND substr(session_id, 19, 1) = '-'
AND substr(session_id, 24, 1) = '-'
AND length(replace(session_id, '-', '')) = 32
AND replace(session_id, '-', '') NOT GLOB '*[^0-9a-f]*'
AND substr(session_id, 15, 1) = '4'
AND substr(session_id, 20, 1) IN ('8', '9', 'a', 'b')
),
claim_generation INTEGER NOT NULL CHECK (
typeof(claim_generation) = 'integer' AND claim_generation > 0
),
task_version INTEGER NOT NULL CHECK (
typeof(task_version) = 'integer' AND task_version > 0
),
task_title TEXT NOT NULL CHECK (trim(task_title) <> ''),
authorization_task_version INTEGER NOT NULL CHECK (
typeof(authorization_task_version) = 'integer' AND authorization_task_version > 0
),
goods_id TEXT NOT NULL CHECK (trim(goods_id) <> ''),
sku_color TEXT NOT NULL CHECK (trim(sku_color) <> ''),
sku_size TEXT NOT NULL CHECK (trim(sku_size) <> ''),
quantity INTEGER NOT NULL CHECK (typeof(quantity) = 'integer' AND quantity > 0),
total_price_cap TEXT NOT NULL CHECK (trim(total_price_cap) <> ''),
authorization_expires_at TEXT NOT NULL CHECK (
authorization_expires_at = trim(authorization_expires_at)
AND length(authorization_expires_at) >= 20
AND substr(authorization_expires_at, 11, 1) = 'T'
AND substr(authorization_expires_at, -1, 1) = 'Z'
AND julianday(authorization_expires_at) IS NOT NULL
),
claim_nonce BLOB NOT NULL CHECK (
typeof(claim_nonce) = 'blob' AND length(claim_nonce) = 32
),
claim_token_sha256 BLOB NOT NULL CHECK (
typeof(claim_token_sha256) = 'blob' AND length(claim_token_sha256) = 32
),
lease_expires_at TEXT NOT NULL CHECK (
lease_expires_at = trim(lease_expires_at)
AND length(lease_expires_at) >= 20
AND substr(lease_expires_at, 11, 1) = 'T'
AND substr(lease_expires_at, -1, 1) = 'Z'
AND julianday(lease_expires_at) IS NOT NULL
),
claimed_at TEXT NOT NULL CHECK (
claimed_at = trim(claimed_at)
AND length(claimed_at) >= 20
AND substr(claimed_at, 11, 1) = 'T'
AND substr(claimed_at, -1, 1) = 'Z'
AND julianday(claimed_at) IS NOT NULL
),
closed_at TEXT CHECK (
closed_at IS NULL OR (
closed_at = trim(closed_at)
AND length(closed_at) >= 20
AND substr(closed_at, 11, 1) = 'T'
AND substr(closed_at, -1, 1) = 'Z'
AND julianday(closed_at) IS NOT NULL
AND julianday(closed_at) >= julianday(claimed_at)
)
),
UNIQUE (task_id, attempt_id),
UNIQUE (attempt_id, claimed_by_device_id, session_id),
UNIQUE (
task_id, attempt_id, claimed_by_device_id, session_id,
claim_generation, claim_token_sha256
),
UNIQUE (task_id, authorization_id, attempt_id),
FOREIGN KEY (task_id, authorization_id, attempt_id, claim_generation)
REFERENCES purchase_attempts(task_id, authorization_id, id, claim_generation),
FOREIGN KEY (claimed_by_device_id) REFERENCES device_credentials(device_id)
);
-- 过期、撤销或停轮询都不会自动关闭 claim;partial unique 因而阻止另一条开放归属。
CREATE UNIQUE INDEX purchase_attempt_claims_one_open_per_device_idx
ON purchase_attempt_claims (claimed_by_device_id)
WHERE closed_at IS NULL;
CREATE TABLE task_claim_requests (
claim_request_id TEXT PRIMARY KEY CHECK (
length(claim_request_id) = 36
AND substr(claim_request_id, 9, 1) = '-'
AND substr(claim_request_id, 14, 1) = '-'
AND substr(claim_request_id, 19, 1) = '-'
AND substr(claim_request_id, 24, 1) = '-'
AND length(replace(claim_request_id, '-', '')) = 32
AND replace(claim_request_id, '-', '') NOT GLOB '*[^0-9a-f]*'
AND substr(claim_request_id, 15, 1) = '4'
AND substr(claim_request_id, 20, 1) IN ('8', '9', 'a', 'b')
),
device_id TEXT NOT NULL,
session_id TEXT NOT NULL CHECK (
length(session_id) = 36
AND substr(session_id, 9, 1) = '-'
AND substr(session_id, 14, 1) = '-'
AND substr(session_id, 19, 1) = '-'
AND substr(session_id, 24, 1) = '-'
AND length(replace(session_id, '-', '')) = 32
AND replace(session_id, '-', '') NOT GLOB '*[^0-9a-f]*'
AND substr(session_id, 15, 1) = '4'
AND substr(session_id, 20, 1) IN ('8', '9', 'a', 'b')
),
outcome TEXT NOT NULL CHECK (outcome IN ('CLAIMED', 'EMPTY', 'BLOCKED')),
attempt_id TEXT,
response_lease_expires_at TEXT CHECK (
response_lease_expires_at IS NULL OR (
response_lease_expires_at = trim(response_lease_expires_at)
AND length(response_lease_expires_at) >= 20
AND substr(response_lease_expires_at, 11, 1) = 'T'
AND substr(response_lease_expires_at, -1, 1) = 'Z'
AND julianday(response_lease_expires_at) IS NOT NULL
)
),
error_code TEXT CHECK (error_code IS NULL OR error_code = 'manual_recovery_required'),
created_at TEXT NOT NULL CHECK (
created_at = trim(created_at)
AND length(created_at) >= 20
AND substr(created_at, 11, 1) = 'T'
AND substr(created_at, -1, 1) = 'Z'
AND julianday(created_at) IS NOT NULL
),
CHECK (
(outcome = 'CLAIMED' AND attempt_id IS NOT NULL AND response_lease_expires_at IS NOT NULL AND error_code IS NULL)
OR (outcome = 'EMPTY' AND attempt_id IS NULL AND response_lease_expires_at IS NULL AND error_code IS NULL)
OR (outcome = 'BLOCKED' AND attempt_id IS NULL AND response_lease_expires_at IS NULL AND error_code = 'manual_recovery_required')
),
FOREIGN KEY (device_id) REFERENCES device_credentials(device_id),
-- EMPTY/BLOCKED 行的 attempt_id 为 NULL,SQLite 会跳过复合 FK;CLAIMED 行则必须
-- 同时匹配原 claim 的设备和 session,不能由应用 bug 写成跨设备重放。
FOREIGN KEY (attempt_id, device_id, session_id)
REFERENCES purchase_attempt_claims(attempt_id, claimed_by_device_id, session_id)
);
CREATE TABLE purchase_attempt_lease_renewals (
renew_request_id TEXT PRIMARY KEY CHECK (
length(renew_request_id) = 36
AND substr(renew_request_id, 9, 1) = '-'
AND substr(renew_request_id, 14, 1) = '-'
AND substr(renew_request_id, 19, 1) = '-'
AND substr(renew_request_id, 24, 1) = '-'
AND length(replace(renew_request_id, '-', '')) = 32
AND replace(renew_request_id, '-', '') NOT GLOB '*[^0-9a-f]*'
AND substr(renew_request_id, 15, 1) = '4'
AND substr(renew_request_id, 20, 1) IN ('8', '9', 'a', 'b')
),
task_id TEXT NOT NULL,
attempt_id TEXT NOT NULL,
device_id TEXT NOT NULL,
session_id TEXT NOT NULL CHECK (
length(session_id) = 36
AND substr(session_id, 9, 1) = '-'
AND substr(session_id, 14, 1) = '-'
AND substr(session_id, 19, 1) = '-'
AND substr(session_id, 24, 1) = '-'
AND length(replace(session_id, '-', '')) = 32
AND replace(session_id, '-', '') NOT GLOB '*[^0-9a-f]*'
AND substr(session_id, 15, 1) = '4'
AND substr(session_id, 20, 1) IN ('8', '9', 'a', 'b')
),
claim_generation INTEGER NOT NULL CHECK (
typeof(claim_generation) = 'integer' AND claim_generation > 0
),
claim_token_sha256 BLOB NOT NULL CHECK (
typeof(claim_token_sha256) = 'blob' AND length(claim_token_sha256) = 32
),
expected_lease_expires_at TEXT NOT NULL CHECK (
expected_lease_expires_at = trim(expected_lease_expires_at)
AND length(expected_lease_expires_at) >= 20
AND substr(expected_lease_expires_at, 11, 1) = 'T'
AND substr(expected_lease_expires_at, -1, 1) = 'Z'
AND julianday(expected_lease_expires_at) IS NOT NULL
),
lease_expires_at TEXT NOT NULL CHECK (
lease_expires_at = trim(lease_expires_at)
AND length(lease_expires_at) >= 20
AND substr(lease_expires_at, 11, 1) = 'T'
AND substr(lease_expires_at, -1, 1) = 'Z'
AND julianday(lease_expires_at) IS NOT NULL
),
created_at TEXT NOT NULL CHECK (
created_at = trim(created_at)
AND length(created_at) >= 20
AND substr(created_at, 11, 1) = 'T'
AND substr(created_at, -1, 1) = 'Z'
AND julianday(created_at) IS NOT NULL
),
FOREIGN KEY (
task_id, attempt_id, device_id, session_id,
claim_generation, claim_token_sha256
) REFERENCES purchase_attempt_claims(
task_id, attempt_id, claimed_by_device_id, session_id,
claim_generation, claim_token_sha256
)
);
CREATE INDEX order_authorizations_claim_candidate_idx
ON order_authorizations (status, created_at, id);
-- 首次证据写入必须属于认证设备当前未关闭的 claim。历史资产的幂等重放不触发 INSERT,
-- 因而未来人工关闭 claim 后仍可稳定返回原资产。
-- +goose StatementBegin
CREATE TRIGGER evidence_assets_claim_owner_insert
BEFORE INSERT ON evidence_assets
FOR EACH ROW
WHEN NOT EXISTS (
SELECT 1 FROM purchase_attempt_claims AS claims
WHERE claims.task_id = NEW.task_id
AND claims.attempt_id = NEW.attempt_id
AND claims.claimed_by_device_id = NEW.uploaded_by_device_id
AND claims.closed_at IS NULL
)
BEGIN
SELECT RAISE(ABORT, 'evidence claim ownership required');
END;
-- +goose StatementEnd
-- +goose Down
-- 请求、续租、attempt、submission 和证据都是领取或下游审计事实,回滚不得静默删除。
CREATE TABLE task_claim_downgrade_guard (
valid INTEGER NOT NULL CHECK (valid = 1)
);
INSERT INTO task_claim_downgrade_guard (valid)
SELECT CASE WHEN
(SELECT COUNT(*) FROM task_claim_requests) = 0
AND (SELECT COUNT(*) FROM purchase_attempt_lease_renewals) = 0
AND (SELECT COUNT(*) FROM purchase_attempt_claims) = 0
AND (SELECT COUNT(*) FROM purchase_attempts) = 0
AND (SELECT COUNT(*) FROM order_submissions) = 0
AND (SELECT COUNT(*) FROM evidence_assets) = 0
THEN 1 ELSE 0 END;
DROP TABLE task_claim_downgrade_guard;
DROP TRIGGER evidence_assets_claim_owner_insert;
DROP INDEX order_authorizations_claim_candidate_idx;
DROP TABLE purchase_attempt_lease_renewals;
DROP TABLE task_claim_requests;
DROP INDEX purchase_attempt_claims_one_open_per_device_idx;
DROP TABLE purchase_attempt_claims;
DROP INDEX purchase_attempts_claim_lineage_idx;
DROP INDEX purchase_attempts_one_per_authorization_idx;
@@ -0,0 +1,89 @@
"""采集 T-103 一次性尺码 reveal 的 before/after 本机证据。"""
from __future__ import annotations
import argparse
from math import isfinite
from pathlib import Path
import sys
CLIENT_ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(CLIENT_ROOT / "src"))
from cmbuyer_client.device.adb import AdbClient, DeviceConnectionError, SubprocessAdbRunner
from cmbuyer_client.device.baseline import NoReconnectUiautomatorConnector
from cmbuyer_client.pdd.sku_reveal_spike import SkuRevealSpikeCapturer, SkuRevealSpikeError
from cmbuyer_client.pdd.sku_selection import EXPECTED_GOODS_ID, SkuSelectionError
from cmbuyer_client.pdd.sku_selection_runner import SkuSelectionRunError
def parse_arguments(argv: list[str] | None = None) -> argparse.Namespace:
parser = argparse.ArgumentParser(description="采集 T-103 一次性尺码 reveal 的本机证据。")
parser.add_argument("--serial", required=True, help="ADB device serial;禁止自动选择。")
parser.add_argument("--goods-id", required=True, help="固定 T-103 已取证 goods_id。")
parser.add_argument("--output-dir", required=True, type=Path, help="全新本机证据目录;不得覆盖。")
parser.add_argument("--timeout", type=float, default=30.0, help="整趟动作与只读调和时限(秒)。")
parser.add_argument("--adb", default="adb", help="adb 可执行文件路径。")
return parser.parse_args(argv)
def validate_arguments(arguments: argparse.Namespace) -> None:
if type(arguments.serial) is not str or not arguments.serial.strip():
raise ValueError("必须显式提供非空 --serial。")
if type(arguments.goods_id) is not str or arguments.goods_id != EXPECTED_GOODS_ID:
raise ValueError("--goods-id 不是 T-103 已取证商品。")
if (
not isinstance(arguments.timeout, (int, float))
or isinstance(arguments.timeout, bool)
or arguments.timeout <= 0
or not isfinite(arguments.timeout)
):
raise ValueError("--timeout 必须是大于 0 的有限数值。")
def main(argv: list[str] | None = None) -> int:
arguments = parse_arguments(argv)
try:
validate_arguments(arguments)
except ValueError as error:
print(f"失败:{error}", file=sys.stderr)
return 2
try:
import adbutils
import uiautomator2 as u2
except ImportError:
print("失败:缺少 uiautomator2;请在采购工具虚拟环境中运行。", file=sys.stderr)
return 2
capturer = SkuRevealSpikeCapturer(
AdbClient(SubprocessAdbRunner(arguments.adb), timeout_seconds=arguments.timeout),
NoReconnectUiautomatorConnector(
adbutils.AdbClient(socket_timeout=arguments.timeout).device_list,
u2.connect,
),
timeout_seconds=arguments.timeout,
)
try:
result = capturer.capture(
arguments.serial,
arguments.goods_id,
arguments.output_dir,
)
except (DeviceConnectionError, SkuSelectionError, SkuSelectionRunError, SkuRevealSpikeError):
# 不回显页面正文、节点、serial、坐标、路径或第三方异常。
print("规格 reveal 取证失败:已停止,未发布本地证据目录。", file=sys.stderr)
return 1
except OSError:
print("规格 reveal 取证失败:无法创建或发布本地证据目录。", file=sys.stderr)
return 1
print(f"规格 reveal 取证完成:{result.output_directory}")
print(f"manifest:{result.manifest_path}")
print("人工复核:请保持手机不动,确认颜色仍选中、S/M 均未选且未进入提交页。")
return 0
if __name__ == "__main__":
raise SystemExit(main())
+87
View File
@@ -0,0 +1,87 @@
"""恢复 T-103 已取证目标规格、验证现价并保存本地原始截图。"""
from __future__ import annotations
import argparse
from math import isfinite
from pathlib import Path
import sys
CLIENT_ROOT = Path(__file__).resolve().parents[1]
sys.path.insert(0, str(CLIENT_ROOT / "src"))
from cmbuyer_client.device.adb import AdbClient, DeviceConnectionError, SubprocessAdbRunner
from cmbuyer_client.device.baseline import NoReconnectUiautomatorConnector
from cmbuyer_client.pdd.product_url import ProductUrlError, parse_product_url
from cmbuyer_client.pdd.sku_selection import EXPECTED_GOODS_ID, SkuSelectionError, TASK_TO_UI_SELECTION
from cmbuyer_client.pdd.sku_selection_runner import SkuSelectionRunError, SkuSelectionRunner, safe_failure_stage
def parse_arguments(argv: list[str] | None = None) -> argparse.Namespace:
parser = argparse.ArgumentParser(description="恢复 T-103 已取证规格并保存本地原始截图。")
parser.add_argument("--serial", required=True, help="ADB device serial;禁止自动选择。")
parser.add_argument("--url", required=True, help="唯一 canonical goods.html?goods_id= 直链。")
parser.add_argument("--color", required=True, help="T-103 任务颜色值。")
parser.add_argument("--size", required=True, help="T-103 任务尺码值。")
parser.add_argument("--output-dir", required=True, type=Path, help="新建本地目录;不得覆盖已有目录。")
parser.add_argument("--timeout", type=float, default=10.0, help="ADB 与设备 RPC 超时(秒)。")
parser.add_argument("--adb", default="adb", help="adb 可执行文件路径。")
return parser.parse_args(argv)
def validate_arguments(arguments: argparse.Namespace) -> None:
if not isinstance(arguments.serial, str) or not arguments.serial.strip():
raise ValueError("必须显式提供非空 --serial。")
if not isinstance(arguments.timeout, (int, float)) or isinstance(arguments.timeout, bool) or arguments.timeout <= 0 or not isfinite(arguments.timeout):
raise ValueError("--timeout 必须是大于 0 的有限数值。")
link = parse_product_url(arguments.url)
if link.goods_id != EXPECTED_GOODS_ID:
raise ValueError("--url 不是 T-103 已取证商品。")
if (arguments.color, arguments.size) not in TASK_TO_UI_SELECTION:
raise ValueError("--color 与 --size 必须是 T-103 已取证任务值。")
def main(argv: list[str] | None = None) -> int:
arguments = parse_arguments(argv)
try:
validate_arguments(arguments)
except (ValueError, ProductUrlError) as error:
print(f"失败:{error}", file=sys.stderr)
return 2
try:
import adbutils
import uiautomator2 as u2
except ImportError:
print("失败:缺少 uiautomator2;请在采购工具虚拟环境中运行。", file=sys.stderr)
return 2
runner = SkuSelectionRunner(
AdbClient(SubprocessAdbRunner(arguments.adb), timeout_seconds=arguments.timeout),
NoReconnectUiautomatorConnector(adbutils.AdbClient(socket_timeout=arguments.timeout).device_list, u2.connect),
timeout_seconds=arguments.timeout,
)
try:
result = runner.run(arguments.serial, arguments.url, arguments.color, arguments.size, arguments.output_dir)
except (DeviceConnectionError, SkuSelectionRunError, SkuSelectionError) as error:
# Flow 可能来自测试替身或未来实现;CLI 不回显任何异常正文,避免泄露节点树或页面文本。
print(
f"规格恢复失败:stage={safe_failure_stage(error)};已停止,未发布本地证据目录。",
file=sys.stderr,
)
return 1
except OSError:
print("规格恢复失败:无法创建或发布本地证据目录。", file=sys.stderr)
return 1
print(f"规格恢复完成:{result.output_directory}")
print(f"manifest:{result.manifest_path}")
print(f"目标规格:{arguments.color} / {arguments.size}")
print(f"确认单价:{result.unit_price}")
print("页面对应性:请人工核对本地原始截图。")
return 0
if __name__ == "__main__":
raise SystemExit(main())
+43 -15
View File
@@ -5,8 +5,9 @@ from __future__ import annotations
import sys
from collections.abc import Sequence
from .core.errors import StateError
from .logging_policy import configure_application_logger
from .runtime import RuntimePaths
from .runtime import LocalStateRuntime, RuntimePaths
def select_application_argv(argv: Sequence[str] | None) -> list[str]:
@@ -30,8 +31,7 @@ def main(argv: Sequence[str] | None = None) -> int:
return 1
try:
from PySide6.QtCore import Qt
from PySide6.QtWidgets import QApplication, QLabel, QMainWindow
from PySide6.QtWidgets import QApplication
except ImportError:
logger.error("缺少 PySide6,无法启动桌面界面。")
print("无法启动采购工具:缺少 PySide6。请先安装 requirements.txt 中的依赖。", file=sys.stderr)
@@ -39,19 +39,47 @@ def main(argv: Sequence[str] | None = None) -> int:
application = QApplication.instance() or QApplication(select_application_argv(argv))
application.setApplicationName("采购工具")
runtime: LocalStateRuntime | None = None
coordinator = None
try:
runtime = LocalStateRuntime.open(paths)
try:
summary = runtime.store.load_profile_summary("default")
except StateError as error:
if error.reason != "profile_not_found":
raise
summary = None
window = QMainWindow()
window.setWindowTitle("采购工具")
window.setAccessibleName("采购工具")
window.setMinimumSize(420, 240)
window.resize(560, 320)
from .polling.coordinator import PollingCoordinator
from .ui.main_window import PurchaseToolWindow
message = QLabel("应用骨架已初始化。\n采购执行功能尚未启用。")
message.setAlignment(Qt.AlignmentFlag.AlignCenter)
message.setWordWrap(True)
message.setAccessibleName("当前状态")
window.setCentralWidget(message)
logger.info("应用已启动;采购执行功能尚未启用。")
settings = None if summary is None else summary.settings
has_token = False if summary is None else summary.has_stored_device_token
coordinator = PollingCoordinator(
profile_id="default",
store=runtime.store,
gateway_factory=None,
consumer=None,
profile_settings=settings,
poll_interval_seconds=15 if settings is None else settings.poll_interval_seconds,
failure_threshold=3 if settings is None else settings.failure_threshold,
)
window = PurchaseToolWindow(
store=runtime.store,
coordinator=coordinator,
profile_settings=settings,
has_stored_device_token=has_token,
)
logger.info("应用已启动;单趟执行能力尚未接入,真实领取保持禁用。")
window.show()
return application.exec()
except (OSError, RuntimeError, StateError):
logger.error("无法打开采购工具本地安全状态。")
print("无法启动采购工具:本地安全状态不可用。", file=sys.stderr)
return 3
finally:
worker_stopped = True
if coordinator is not None:
worker_stopped = coordinator.shutdown()
if runtime is not None and worker_stopped:
runtime.close()
@@ -0,0 +1,44 @@
"""与 UI、HTTP 和拼多多页面实现无关的客户端核心契约。"""
from .errors import (
AmbiguousRemoteError,
CredentialRemoteError,
ManualRemoteError,
ProtocolRemoteError,
StateError,
ValidationError,
)
from .models import (
AssetReceipt,
AuthorizationSnapshot,
ClaimRequest,
ClaimedTask,
DeviceCredentials,
EvidenceUpload,
PurchaseTask,
RenewRequest,
RenewResult,
SecretToken,
)
from .ports import EvidenceSink, TaskSource
__all__ = [
"AmbiguousRemoteError",
"AssetReceipt",
"AuthorizationSnapshot",
"ClaimRequest",
"ClaimedTask",
"CredentialRemoteError",
"DeviceCredentials",
"EvidenceSink",
"EvidenceUpload",
"ManualRemoteError",
"ProtocolRemoteError",
"PurchaseTask",
"RenewRequest",
"RenewResult",
"SecretToken",
"StateError",
"TaskSource",
"ValidationError",
]
+47
View File
@@ -0,0 +1,47 @@
"""可安全呈现的客户端错误分类。"""
from __future__ import annotations
class ClientError(RuntimeError):
"""错误文本只使用固定 reason code,不携带凭据、响应或本机路径。"""
def __init__(self, reason: str) -> None:
self.reason = reason
super().__init__(reason)
class ValidationError(ClientError):
"""本地输入或 wire schema 不满足固定契约。"""
class StateError(ClientError):
"""本地状态无法安全推进;调用方必须停止而不是绕过。"""
class ProtectionError(ClientError):
"""秘密保护失败。"""
class SingleInstanceError(ClientError):
"""同一配置已经由另一个采购工具进程持有。"""
class RemoteError(ClientError):
"""服务端调用的稳定错误分类。"""
class AmbiguousRemoteError(RemoteError):
"""请求结果不明;只允许以原幂等键、原载荷显式恢复。"""
class CredentialRemoteError(RemoteError):
"""设备凭据无效或已撤销。"""
class ProtocolRemoteError(RemoteError):
"""请求/响应与固定协议不兼容,不得自动重试。"""
class ManualRemoteError(RemoteError):
"""服务端要求人工处理的确定性冲突。"""
+328
View File
@@ -0,0 +1,328 @@
"""任务领取、续租和单张证据上传的不可变值对象。"""
from __future__ import annotations
from dataclasses import dataclass, field
import hashlib
from pathlib import Path
from .errors import ValidationError
from .validation import (
MAX_SKU_TEXT_CODE_POINTS,
MAX_TITLE_CODE_POINTS,
canonical_product_url,
require_exact_fields,
require_goods_id,
require_lower_hex_64,
require_money,
require_persisted_text,
require_positive_int,
require_rfc3339_z,
rfc3339_z_nanoseconds,
require_string,
require_uuid4,
)
EVIDENCE_KIND = "SKU_PANEL_GATE_1"
PRIVACY_TIER = "INTERNAL_RAW"
@dataclass(frozen=True, repr=False)
class SecretToken:
"""64 位小写 token;repr 永不暴露明文。"""
value: str
def __post_init__(self) -> None:
require_lower_hex_64(self.value, "invalid_token")
def __repr__(self) -> str:
return "SecretToken([已隐藏])"
def __str__(self) -> str:
return "[已隐藏]"
@dataclass(frozen=True, repr=False)
class DeviceCredentials:
device_id: str
token: SecretToken
def __post_init__(self) -> None:
require_uuid4(self.device_id, "invalid_device_id")
if not isinstance(self.token, SecretToken):
raise ValidationError("invalid_device_token")
def __repr__(self) -> str:
return f"DeviceCredentials(device_id={self.device_id!r}, token=[已隐藏])"
@dataclass(frozen=True)
class ClaimRequest:
session_id: str
claim_request_id: str
def __post_init__(self) -> None:
require_uuid4(self.session_id, "invalid_session_id")
require_uuid4(self.claim_request_id, "invalid_claim_request_id")
def to_wire(self) -> dict[str, object]:
return {"session_id": self.session_id, "claim_request_id": self.claim_request_id}
@dataclass(frozen=True)
class PurchaseTask:
id: str
version: int
title: str
product_url: str
goods_id: str
sku_color: str
sku_size: str
quantity: int
max_total_price: str
def __post_init__(self) -> None:
require_uuid4(self.id, "invalid_task_id")
require_positive_int(self.version, "invalid_task_version")
require_persisted_text(self.title, "invalid_task_title", maximum=MAX_TITLE_CODE_POINTS)
require_goods_id(self.goods_id)
if self.product_url != canonical_product_url(self.goods_id):
raise ValidationError("invalid_product_url")
require_persisted_text(self.sku_color, "invalid_sku_color", maximum=MAX_SKU_TEXT_CODE_POINTS)
require_persisted_text(self.sku_size, "invalid_sku_size", maximum=MAX_SKU_TEXT_CODE_POINTS)
require_positive_int(self.quantity, "invalid_quantity")
require_money(self.max_total_price, "invalid_max_total_price")
@classmethod
def from_wire(cls, value: object) -> "PurchaseTask":
data = require_exact_fields(
value,
("id", "version", "title", "product_url", "goods_id", "sku_color", "sku_size", "quantity", "max_total_price"),
)
return cls(**data) # type: ignore[arg-type]
@dataclass(frozen=True)
class AuthorizationSnapshot:
id: str
task_version: int
expires_at: str
def __post_init__(self) -> None:
require_uuid4(self.id, "invalid_authorization_id")
require_positive_int(self.task_version, "invalid_authorization_task_version")
require_rfc3339_z(self.expires_at, "invalid_authorization_expiry")
@classmethod
def from_wire(cls, value: object) -> "AuthorizationSnapshot":
data = require_exact_fields(value, ("id", "task_version", "expires_at"))
return cls(**data) # type: ignore[arg-type]
@dataclass(frozen=True)
class AttemptSnapshot:
id: str
claim_token: SecretToken
claim_generation: int
lease_expires_at: str
def __post_init__(self) -> None:
require_uuid4(self.id, "invalid_attempt_id")
if not isinstance(self.claim_token, SecretToken):
object.__setattr__(self, "claim_token", SecretToken(self.claim_token))
require_positive_int(self.claim_generation, "invalid_claim_generation")
require_rfc3339_z(self.lease_expires_at, "invalid_lease_expiry")
@classmethod
def from_wire(cls, value: object) -> "AttemptSnapshot":
data = require_exact_fields(value, ("id", "claim_token", "claim_generation", "lease_expires_at"))
return cls(
id=data["id"], # type: ignore[arg-type]
claim_token=SecretToken(data["claim_token"]), # type: ignore[arg-type]
claim_generation=data["claim_generation"], # type: ignore[arg-type]
lease_expires_at=data["lease_expires_at"], # type: ignore[arg-type]
)
@dataclass(frozen=True)
class ClaimedTask:
task: PurchaseTask
authorization: AuthorizationSnapshot
attempt: AttemptSnapshot = field(repr=False)
def __post_init__(self) -> None:
if self.task.version != self.authorization.task_version + 1:
raise ValidationError("task_authorization_version_mismatch")
if rfc3339_z_nanoseconds(self.attempt.lease_expires_at) > rfc3339_z_nanoseconds(self.authorization.expires_at):
raise ValidationError("claim_lease_exceeds_authorization")
@classmethod
def from_wire(cls, value: object) -> "ClaimedTask":
data = require_exact_fields(value, ("task", "authorization", "attempt"))
return cls(
task=PurchaseTask.from_wire(data["task"]),
authorization=AuthorizationSnapshot.from_wire(data["authorization"]),
attempt=AttemptSnapshot.from_wire(data["attempt"]),
)
@dataclass(frozen=True, repr=False)
class RenewRequest:
task_id: str
renew_request_id: str
session_id: str
attempt_id: str
claim_generation: int
claim_token: SecretToken
expected_lease_expires_at: str
authorization_expires_at: str
def __post_init__(self) -> None:
require_uuid4(self.task_id, "invalid_task_id")
require_uuid4(self.renew_request_id, "invalid_renew_request_id")
require_uuid4(self.session_id, "invalid_session_id")
require_uuid4(self.attempt_id, "invalid_attempt_id")
require_positive_int(self.claim_generation, "invalid_claim_generation")
if not isinstance(self.claim_token, SecretToken):
object.__setattr__(self, "claim_token", SecretToken(self.claim_token))
require_rfc3339_z(self.expected_lease_expires_at, "invalid_expected_lease_expiry")
require_rfc3339_z(self.authorization_expires_at, "invalid_authorization_expiry")
def __repr__(self) -> str:
return (
f"RenewRequest(task_id={self.task_id!r}, renew_request_id={self.renew_request_id!r}, "
"claim_token=[已隐藏])"
)
def to_wire(self) -> dict[str, object]:
return {
"renew_request_id": self.renew_request_id,
"session_id": self.session_id,
"attempt_id": self.attempt_id,
"claim_generation": self.claim_generation,
"claim_token": self.claim_token.value,
"expected_lease_expires_at": self.expected_lease_expires_at,
}
@dataclass(frozen=True)
class RenewResult:
task_id: str
attempt_id: str
claim_generation: int
lease_expires_at: str
def __post_init__(self) -> None:
require_uuid4(self.task_id, "invalid_task_id")
require_uuid4(self.attempt_id, "invalid_attempt_id")
require_positive_int(self.claim_generation, "invalid_claim_generation")
require_rfc3339_z(self.lease_expires_at, "invalid_lease_expiry")
@classmethod
def from_wire(cls, value: object) -> "RenewResult":
data = require_exact_fields(value, ("task_id", "attempt_id", "claim_generation", "lease_expires_at"))
return cls(**data) # type: ignore[arg-type]
@dataclass(frozen=True, repr=False)
class EvidenceUpload:
task_id: str
upload_key: str
attempt_id: str
sha256: str
captured_at: str
content: bytes = field(repr=False)
kind: str = EVIDENCE_KIND
privacy_tier: str = PRIVACY_TIER
width_px: int = field(init=False)
height_px: int = field(init=False)
def __post_init__(self) -> None:
require_uuid4(self.task_id, "invalid_task_id")
require_uuid4(self.upload_key, "invalid_upload_key")
require_uuid4(self.attempt_id, "invalid_attempt_id")
require_lower_hex_64(self.sha256, "invalid_evidence_sha256")
require_rfc3339_z(self.captured_at, "invalid_captured_at")
if self.kind != EVIDENCE_KIND or self.privacy_tier != PRIVACY_TIER:
raise ValidationError("invalid_evidence_metadata")
if not isinstance(self.content, bytes) or not self.content or len(self.content) > 10 * 1024 * 1024:
raise ValidationError("invalid_evidence_size")
if len(self.content) < 24 or not self.content.startswith(b"\x89PNG\r\n\x1a\n") or self.content[12:16] != b"IHDR":
raise ValidationError("invalid_evidence_png")
width = int.from_bytes(self.content[16:20], "big")
height = int.from_bytes(self.content[20:24], "big")
if width <= 0 or height <= 0 or width > 8192 or height > 8192 or width * height > 16_777_216:
raise ValidationError("invalid_evidence_dimensions")
object.__setattr__(self, "width_px", width)
object.__setattr__(self, "height_px", height)
if hashlib.sha256(self.content).hexdigest() != self.sha256:
raise ValidationError("evidence_hash_mismatch")
def __repr__(self) -> str:
return (
f"EvidenceUpload(task_id={self.task_id!r}, upload_key={self.upload_key!r}, "
f"attempt_id={self.attempt_id!r}, byte_size={len(self.content)})"
)
@dataclass(frozen=True)
class AssetReceipt:
asset_id: str
task_id: str
attempt_id: str
kind: str
privacy_tier: str
sha256: str
byte_size: int
content_type: str
width_px: int
height_px: int
captured_at: str
def __post_init__(self) -> None:
require_uuid4(self.asset_id, "invalid_asset_id")
require_uuid4(self.task_id, "invalid_task_id")
require_uuid4(self.attempt_id, "invalid_attempt_id")
if self.kind != EVIDENCE_KIND or self.privacy_tier != PRIVACY_TIER:
raise ValidationError("invalid_asset_metadata")
require_lower_hex_64(self.sha256, "invalid_asset_sha256")
require_positive_int(self.byte_size, "invalid_asset_byte_size")
if self.byte_size > 10 * 1024 * 1024 or self.content_type != "image/png":
raise ValidationError("invalid_asset_content")
width = require_positive_int(self.width_px, "invalid_asset_width")
height = require_positive_int(self.height_px, "invalid_asset_height")
if width > 8192 or height > 8192 or width * height > 16_777_216:
raise ValidationError("invalid_asset_dimensions")
require_rfc3339_z(self.captured_at, "invalid_captured_at")
@classmethod
def from_wire(cls, value: object) -> "AssetReceipt":
data = require_exact_fields(
value,
("asset_id", "task_id", "attempt_id", "kind", "privacy_tier", "sha256", "byte_size", "content_type", "width_px", "height_px", "captured_at"),
)
return cls(**data) # type: ignore[arg-type]
@dataclass(frozen=True, repr=False)
class ScreenshotAsset:
"""调用方显式选择的唯一 PNG;路径不会进入 repr 或 HTTP。"""
path: Path = field(repr=False)
task_id: str
attempt_id: str
captured_at: str
kind: str = EVIDENCE_KIND
privacy_tier: str = PRIVACY_TIER
def __post_init__(self) -> None:
require_uuid4(self.task_id, "invalid_task_id")
require_uuid4(self.attempt_id, "invalid_attempt_id")
require_rfc3339_z(self.captured_at, "invalid_captured_at")
if self.kind != EVIDENCE_KIND or self.privacy_tier != PRIVACY_TIER:
raise ValidationError("invalid_evidence_metadata")
def __repr__(self) -> str:
return f"ScreenshotAsset(task_id={self.task_id!r}, attempt_id={self.attempt_id!r}, path=[已隐藏])"
+17
View File
@@ -0,0 +1,17 @@
"""由 HTTP 适配器实现的窄端口。"""
from __future__ import annotations
from typing import Protocol
from .models import AssetReceipt, ClaimRequest, ClaimedTask, DeviceCredentials, EvidenceUpload, RenewRequest, RenewResult
class TaskSource(Protocol):
def claim_next(self, credentials: DeviceCredentials, request: ClaimRequest) -> ClaimedTask | None: ...
def renew(self, credentials: DeviceCredentials, request: RenewRequest) -> RenewResult: ...
class EvidenceSink(Protocol):
def upload(self, credentials: DeviceCredentials, evidence: EvidenceUpload) -> AssetReceipt: ...
@@ -0,0 +1,197 @@
"""客户端与服务端共享 wire 的严格值校验。"""
from __future__ import annotations
import calendar
from datetime import datetime, timezone
import json
import re
from typing import Any, Iterable, Mapping
from urllib.parse import quote
from .errors import ValidationError
UUID4_RE = re.compile(
r"[0-9a-f]{8}-[0-9a-f]{4}-4[0-9a-f]{3}-[89ab][0-9a-f]{3}-[0-9a-f]{12}"
)
LOWER_HEX_64_RE = re.compile(r"[0-9a-f]{64}")
RFC3339_Z_RE = re.compile(
r"\d{4}-\d{2}-\d{2}T\d{2}:\d{2}:\d{2}(?:\.\d{1,9})?Z"
)
MONEY_RE = re.compile(r"(?:0|[1-9][0-9]*)\.[0-9]{2}")
GOODS_ID_RE = re.compile(r"[0-9]+")
MAX_TITLE_CODE_POINTS = 120
MAX_SKU_TEXT_CODE_POINTS = 80
MAX_GOODS_ID_ASCII_CHARACTERS = 32
MAX_MONEY_ASCII_CHARACTERS = 32
# Go strings.TrimSpace uses Unicode White_Space plus the six ASCII space
# characters below, but unlike Python str.strip it does not include U+001C--
# U+001F. Keep the wire contract independent of either runtime's defaults.
GO_UNICODE_WHITE_SPACE = "\t\n\v\f\r \u0085\u00a0\u1680\u2000\u2001\u2002\u2003\u2004\u2005\u2006\u2007\u2008\u2009\u200a\u2028\u2029\u202f\u205f\u3000"
def require_string(value: object, reason: str, *, maximum: int = 4096) -> str:
if not isinstance(value, str) or not value or len(value) > maximum:
raise ValidationError(reason)
if any(0xD800 <= ord(character) <= 0xDFFF for character in value):
raise ValidationError(reason)
return value
def require_persisted_text(value: object, reason: str, *, maximum: int) -> str:
"""Validate text stored by Go after TrimSpace, without Python trim drift."""
text = require_string(value, reason, maximum=maximum)
if text.strip(GO_UNICODE_WHITE_SPACE) != text:
raise ValidationError(reason)
# Python str.strip treats these C0 separators as whitespace while Go does
# not. Reject them anywhere on both ends instead of assigning them two
# runtime-dependent meanings.
if any(0x1C <= ord(character) <= 0x1F for character in text):
raise ValidationError(reason)
return text
def require_uuid4(value: object, reason: str = "invalid_uuid") -> str:
text = require_string(value, reason, maximum=36)
if UUID4_RE.fullmatch(text) is None:
raise ValidationError(reason)
return text
def require_lower_hex_64(value: object, reason: str = "invalid_hex") -> str:
text = require_string(value, reason, maximum=64)
if LOWER_HEX_64_RE.fullmatch(text) is None:
raise ValidationError(reason)
return text
def require_rfc3339_z(value: object, reason: str = "invalid_timestamp") -> str:
text = require_string(value, reason, maximum=40)
if RFC3339_Z_RE.fullmatch(text) is None:
raise ValidationError(reason)
parsed: datetime | None = None
try:
parsed = datetime.fromisoformat(text[:-1] + "+00:00")
except ValueError:
pass
if parsed is None:
raise ValidationError(reason)
if parsed.utcoffset() is None or parsed.utcoffset().total_seconds() != 0:
raise ValidationError(reason)
return text
def rfc3339_z_nanoseconds(value: object, reason: str = "invalid_timestamp") -> int:
"""无浮点、无微秒截断地把 UTC RFC3339Nano 转成纳秒时间轴。"""
text = require_rfc3339_z(value, reason)
base: datetime | None = None
try:
base = datetime.strptime(text[:19], "%Y-%m-%dT%H:%M:%S").replace(tzinfo=timezone.utc)
except ValueError:
pass
if base is None:
raise ValidationError(reason)
fraction = "" if len(text) == 20 else text[20:-1]
nanoseconds = int(fraction.ljust(9, "0")) if fraction else 0
return calendar.timegm(base.utctimetuple()) * 1_000_000_000 + nanoseconds
def datetime_nanoseconds(value: datetime, reason: str = "invalid_timestamp") -> int:
if not isinstance(value, datetime) or value.utcoffset() is None:
raise ValidationError(reason)
utc = value.astimezone(timezone.utc)
return calendar.timegm(utc.utctimetuple()) * 1_000_000_000 + utc.microsecond * 1_000
def require_positive_int(value: object, reason: str = "invalid_integer") -> int:
# bool 是 int 的子类;wire 中必须显式拒绝 true/false。
if type(value) is not int or value <= 0 or value > 9_223_372_036_854_775_807:
raise ValidationError(reason)
return value
def require_money(value: object, reason: str = "invalid_money") -> str:
text = require_string(value, reason, maximum=MAX_MONEY_ASCII_CHARACTERS)
if MONEY_RE.fullmatch(text) is None or text == "0.00":
raise ValidationError(reason)
return text
def require_goods_id(value: object) -> str:
text = require_string(value, "invalid_goods_id", maximum=MAX_GOODS_ID_ASCII_CHARACTERS)
if GOODS_ID_RE.fullmatch(text) is None:
raise ValidationError("invalid_goods_id")
return text
def canonical_product_url(goods_id: str) -> str:
require_goods_id(goods_id)
return "https://mobile.yangkeduo.com/goods.html?goods_id=" + quote(goods_id, safe="")
def require_exact_fields(
value: object,
required: Iterable[str],
reason: str = "invalid_schema",
) -> Mapping[str, Any]:
if not isinstance(value, dict):
raise ValidationError(reason)
expected = frozenset(required)
if frozenset(value) != expected:
raise ValidationError(reason)
return value
def strict_json_loads(raw: bytes, *, maximum: int) -> object:
if not isinstance(raw, bytes) or len(raw) == 0 or len(raw) > maximum:
raise ValidationError("invalid_json_size")
text: str | None = None
try:
text = raw.decode("utf-8")
except UnicodeDecodeError:
pass
if text is None:
raise ValidationError("invalid_json_utf8")
if text.startswith("\ufeff"):
raise ValidationError("invalid_json_bom")
def pairs_hook(pairs: list[tuple[str, Any]]) -> dict[str, Any]:
result: dict[str, Any] = {}
for key, value in pairs:
if key in result:
raise ValidationError("duplicate_json_key")
result[key] = value
return result
def reject_number(_: str) -> object:
raise ValidationError("invalid_json_number")
def parse_integer(value: str) -> int:
digits = value[1:] if value.startswith("-") else value
if len(digits) > 19:
raise ValidationError("invalid_json_integer")
parsed = int(value)
if parsed < -9_223_372_036_854_775_808 or parsed > 9_223_372_036_854_775_807:
raise ValidationError("invalid_json_integer")
return parsed
parsed_json: object | None = None
failed = False
try:
parsed_json = json.loads(
text,
object_pairs_hook=pairs_hook,
parse_int=parse_integer,
parse_float=reject_number,
parse_constant=reject_number,
)
except ValidationError:
raise
except (json.JSONDecodeError, UnicodeError, ValueError, RecursionError):
failed = True
if failed:
raise ValidationError("invalid_json")
return parsed_json
@@ -22,7 +22,7 @@ from ..pdd.product_url import ProductUrl, ProductUrlError, parse_product_url
from ..pdd.sku_panel_state import HUMAN_DECLARED_STATES
SANITIZER_VERSION = "t103-privacy-v4"
SANITIZER_VERSION = "t103-privacy-v5"
EXPECTED_GOODS_ID = "937122477375"
EXPECTED_PDD_VERSION = "8.17.0"
EXPECTED_DEVICE_MODEL = "PKG110"
@@ -56,9 +56,11 @@ _PRICE_PROJECTION_ATTRIBUTES = (
)
# 仅接受普通 ASCII 空格,且每个可分隔位置最多一个;禁止换行、折扣、支付/提交文案和
# 任何其它字符。前缀捕获组用于区分当前价与至多一个划线/原价候选。
_CROSSING_PRICE_TEXT_RE = re.compile(r" {0,1}(?:(快卖光) {0,1})?[¥¥] {0,1}[1-9]\d*\.\d{2} {0,1}\Z")
_CROSSING_PRICE_PREFIX_RE = re.compile(r" {0,1}(?:快卖光 {0,1})?[¥¥] {0,1}[1-9]\d*\.\d{2} {0,1}")
_CROSSING_PRICE_ALLOWED_CHARACTERS = frozenset(" 快卖光¥¥0123456789.")
# T-103 人工在 live 规格面板确认当前价槽的完整非敏感前缀仅为“快卖完”;不得兼容
# 未取证的“快卖光”或其它相近文案。
_CROSSING_PRICE_TEXT_RE = re.compile(r" {0,1}(?:(快卖完) {0,1})?[¥¥] {0,1}[1-9]\d*\.\d{2} {0,1}\Z")
_CROSSING_PRICE_PREFIX_RE = re.compile(r" {0,1}(?:快卖完 {0,1})?[¥¥] {0,1}[1-9]\d*\.\d{2} {0,1}")
_CROSSING_PRICE_ALLOWED_CHARACTERS = frozenset(" 快卖完¥¥0123456789.")
class SkuEvidenceSanitizationError(RuntimeError):
@@ -423,8 +425,10 @@ def _crossing_price_text_mismatch_reason(text: str) -> str:
return "non_ascii_whitespace"
if any(marker in text for marker in ("提交订单", "支付", "下单", "优惠")):
return "extra_or_order"
without_leading_space = text.lstrip(" ")
if without_leading_space.startswith("快") and not without_leading_space.startswith("快卖光"):
without_one_leading_space = text[1:] if text.startswith(" ") else text
if without_one_leading_space.startswith("快要抢光"):
return "observed_prefix_kuaiyaoqiangguang"
if without_one_leading_space.startswith("快") and not without_one_leading_space.startswith("快卖完"):
return "known_prefix_missing"
if "¥" not in text and "¥" not in text:
return "currency_missing"
@@ -0,0 +1,18 @@
"""Windows 本地恢复、凭据保护和单实例底座。"""
from .models import PollingSession, ProfileSettings, RecoverySnapshot
from .facade import DurableClientGateway
from .protection import DpapiProtector, SecretProtector
from .single_instance import NamedMutex
from .store import LocalStateStore
__all__ = [
"DpapiProtector",
"DurableClientGateway",
"LocalStateStore",
"NamedMutex",
"PollingSession",
"ProfileSettings",
"RecoverySnapshot",
"SecretProtector",
]
@@ -0,0 +1,82 @@
"""把“先持久化,再发一次 HTTP”固化成 T-304/T-306 的唯一集成入口。"""
from __future__ import annotations
from cmbuyer_client.core.errors import (
AmbiguousRemoteError,
CredentialRemoteError,
ManualRemoteError,
ProtocolRemoteError,
)
from cmbuyer_client.core.models import AssetReceipt, ClaimedTask, ScreenshotAsset
from cmbuyer_client.core.ports import EvidenceSink, TaskSource
from .store import LocalStateStore
class DurableClientGateway:
"""不隐藏重试;每次方法调用最多发一次请求,结果不明保留原槽。"""
def __init__(self, store: LocalStateStore, task_source: TaskSource, evidence_sink: EvidenceSink) -> None:
self._store = store
self._task_source = task_source
self._evidence_sink = evidence_sink
def claim_next(self, profile_id: str) -> ClaimedTask | None:
request = self._store.prepare_claim(profile_id)
credentials = self._store.load_profile(profile_id).credentials
try:
claimed = self._task_source.claim_next(credentials, request)
except AmbiguousRemoteError:
raise
except CredentialRemoteError:
raise
except ProtocolRemoteError:
self._store.mark_claim_terminal(profile_id, request, "PROTOCOL")
raise
except ManualRemoteError:
self._store.mark_claim_terminal(profile_id, request, "MANUAL")
raise
if claimed is None:
self._store.commit_claim_empty(profile_id, request)
return None
self._store.commit_claim_success(profile_id, request, claimed)
return claimed
def renew(self, profile_id: str):
request = self._store.prepare_renew(profile_id)
credentials = self._store.load_profile(profile_id).credentials
try:
result = self._task_source.renew(credentials, request)
except AmbiguousRemoteError:
raise
except CredentialRemoteError:
raise
except ProtocolRemoteError:
self._store.mark_renew_terminal(profile_id, request, "PROTOCOL")
raise
except ManualRemoteError:
self._store.mark_renew_terminal(profile_id, request, "MANUAL")
raise
self._store.commit_renew_success(profile_id, request, result)
return result
def upload_evidence(self, profile_id: str, asset: ScreenshotAsset) -> AssetReceipt:
prepared = self._store.prepare_or_resume_evidence(profile_id, asset)
if isinstance(prepared, AssetReceipt):
return prepared
credentials = self._store.load_profile(profile_id).credentials
try:
receipt = self._evidence_sink.upload(credentials, prepared)
except AmbiguousRemoteError:
raise
except CredentialRemoteError:
raise
except ProtocolRemoteError:
self._store.mark_evidence_terminal(profile_id, prepared, "PROTOCOL")
raise
except ManualRemoteError:
self._store.mark_evidence_terminal(profile_id, prepared, "MANUAL")
raise
self._store.commit_evidence_success(profile_id, prepared, receipt)
return receipt
@@ -0,0 +1,98 @@
"""供 T-304 使用的稳定本地配置与恢复快照。"""
from __future__ import annotations
from dataclasses import dataclass, field
import re
from cmbuyer_client.core.models import ClaimRequest, ClaimedTask, DeviceCredentials, RenewRequest
from cmbuyer_client.core.validation import require_string, require_uuid4
LOOPBACK_SERVICE_URL = "http://127.0.0.1:8080"
PROFILE_ID_RE = re.compile(r"[a-z0-9][a-z0-9_-]{0,63}")
@dataclass(frozen=True)
class ProfileSettings:
profile_id: str
service_url: str
device_id: str
adb_path: str
adb_serial: str
transport: str
poll_interval_seconds: int = 15
failure_threshold: int = 3
http_timeout_seconds: int = 10
step_timeout_seconds: int = 45
def __post_init__(self) -> None:
if not isinstance(self.profile_id, str) or PROFILE_ID_RE.fullmatch(self.profile_id) is None:
raise ValueError("invalid_profile_id")
if self.service_url != LOOPBACK_SERVICE_URL:
raise ValueError("service_url_not_allowed")
require_uuid4(self.device_id, "invalid_device_id")
require_string(self.adb_path, "invalid_adb_path", maximum=1024)
require_string(self.adb_serial, "invalid_adb_serial", maximum=200)
if self.transport not in ("usb", "wifi"):
raise ValueError("invalid_transport")
_range(self.poll_interval_seconds, 5, 300, "invalid_poll_interval")
_range(self.failure_threshold, 1, 10, "invalid_failure_threshold")
_range(self.http_timeout_seconds, 1, 120, "invalid_http_timeout")
_range(self.step_timeout_seconds, 5, 300, "invalid_step_timeout")
@dataclass(frozen=True, repr=False)
class LoadedProfile:
settings: ProfileSettings
credentials: DeviceCredentials = field(repr=False)
def __repr__(self) -> str:
return f"LoadedProfile(settings={self.settings!r}, credentials=[已隐藏])"
@dataclass(frozen=True)
class ProfileSummary:
"""不解密、不返回任何 token 数据的配置页只读摘要。"""
settings: ProfileSettings
has_stored_device_token: bool
def __post_init__(self) -> None:
if type(self.has_stored_device_token) is not bool:
raise ValueError("invalid_token_presence")
@dataclass(frozen=True)
class PollingSession:
profile_id: str
session_id: str
accept_new: bool
def __post_init__(self) -> None:
require_uuid4(self.session_id, "invalid_session_id")
if type(self.accept_new) is not bool:
raise ValueError("invalid_accept_new")
@dataclass(frozen=True)
class PendingEvidence:
task_id: str
attempt_id: str
kind: str
upload_key: str
status: str
@dataclass(frozen=True)
class RecoverySnapshot:
session: PollingSession | None
pending_claim: ClaimRequest | None
active_claim: ClaimedTask | None = field(repr=False)
pending_renew: RenewRequest | None = field(repr=False)
pending_evidence: tuple[PendingEvidence, ...]
def _range(value: object, minimum: int, maximum: int, reason: str) -> None:
if type(value) is not int or not minimum <= value <= maximum:
raise ValueError(reason)
@@ -0,0 +1,112 @@
"""Windows 当前用户范围 DPAPI 封装;生产环境绝不降级为明文。"""
from __future__ import annotations
import ctypes
from ctypes import wintypes
import os
import re
from typing import Protocol
from cmbuyer_client.core.errors import ProtectionError
class SecretProtector(Protocol):
def protect(self, plaintext: bytes, *, purpose: str) -> bytes: ...
def unprotect(self, ciphertext: bytes, *, purpose: str) -> bytes: ...
class _DataBlob(ctypes.Structure):
_fields_ = (("cbData", wintypes.DWORD), ("pbData", ctypes.POINTER(ctypes.c_ubyte)))
def _blob(data: bytes) -> tuple[_DataBlob, object]:
buffer = (ctypes.c_ubyte * len(data)).from_buffer_copy(data) if data else (ctypes.c_ubyte * 1)()
return _DataBlob(len(data), ctypes.cast(buffer, ctypes.POINTER(ctypes.c_ubyte))), buffer
class DpapiProtector:
"""使用 CryptProtectData/UI_FORBIDDEN;错误只暴露固定 reason code。"""
_UI_FORBIDDEN = 0x1
_ENTROPY_PREFIX = b"cmbuyer-localstate-v1:"
_PURPOSE_RE = re.compile(
r"(?:device-token:[a-z0-9][a-z0-9_-]{0,63}:[0-9a-f-]{36}|"
r"claim-token:[a-z0-9][a-z0-9_-]{0,63}:[0-9a-f-]{36})",
flags=re.ASCII,
)
def __init__(self) -> None:
if os.name != "nt":
raise ProtectionError("dpapi_requires_windows")
self._crypt32 = ctypes.WinDLL("crypt32", use_last_error=True)
self._kernel32 = ctypes.WinDLL("kernel32", use_last_error=True)
self._crypt32.CryptProtectData.argtypes = (
ctypes.POINTER(_DataBlob),
wintypes.LPCWSTR,
ctypes.POINTER(_DataBlob),
wintypes.LPVOID,
wintypes.LPVOID,
wintypes.DWORD,
ctypes.POINTER(_DataBlob),
)
self._crypt32.CryptProtectData.restype = wintypes.BOOL
self._crypt32.CryptUnprotectData.argtypes = (
ctypes.POINTER(_DataBlob),
ctypes.POINTER(wintypes.LPWSTR),
ctypes.POINTER(_DataBlob),
wintypes.LPVOID,
wintypes.LPVOID,
wintypes.DWORD,
ctypes.POINTER(_DataBlob),
)
self._crypt32.CryptUnprotectData.restype = wintypes.BOOL
self._kernel32.LocalFree.argtypes = (wintypes.HLOCAL,)
self._kernel32.LocalFree.restype = wintypes.HLOCAL
def protect(self, plaintext: bytes, *, purpose: str) -> bytes:
if not isinstance(plaintext, bytes) or not plaintext:
raise ProtectionError("invalid_plaintext")
entropy = self._entropy(purpose)
source, source_buffer = _blob(plaintext)
entropy_blob, entropy_buffer = _blob(entropy)
output = _DataBlob()
if not self._crypt32.CryptProtectData(
ctypes.byref(source), None, ctypes.byref(entropy_blob), None, None, self._UI_FORBIDDEN, ctypes.byref(output)
):
raise ProtectionError("dpapi_protect_failed")
# ctypes 指针不持有底层 Python buffer;局部引用必须活到系统调用返回。
del source_buffer, entropy_buffer
return self._take_output(output, "dpapi_protect_failed")
def unprotect(self, ciphertext: bytes, *, purpose: str) -> bytes:
if not isinstance(ciphertext, bytes) or not ciphertext:
raise ProtectionError("invalid_ciphertext")
entropy = self._entropy(purpose)
source, source_buffer = _blob(ciphertext)
entropy_blob, entropy_buffer = _blob(entropy)
output = _DataBlob()
description = wintypes.LPWSTR()
if not self._crypt32.CryptUnprotectData(
ctypes.byref(source), ctypes.byref(description), ctypes.byref(entropy_blob), None, None, self._UI_FORBIDDEN, ctypes.byref(output)
):
raise ProtectionError("dpapi_unprotect_failed")
del source_buffer, entropy_buffer
if description:
self._kernel32.LocalFree(ctypes.cast(description, wintypes.HLOCAL))
return self._take_output(output, "dpapi_unprotect_failed")
def _take_output(self, output: _DataBlob, reason: str) -> bytes:
if not output.pbData or output.cbData <= 0:
raise ProtectionError(reason)
try:
return ctypes.string_at(output.pbData, output.cbData)
finally:
self._kernel32.LocalFree(ctypes.cast(output.pbData, wintypes.HLOCAL))
@classmethod
def _entropy(cls, purpose: str) -> bytes:
if not isinstance(purpose, str) or cls._PURPOSE_RE.fullmatch(purpose) is None:
raise ProtectionError("invalid_protection_purpose")
return cls._ENTROPY_PREFIX + purpose.encode("ascii")
@@ -0,0 +1,51 @@
"""同一本地数据库的 Windows named mutex。"""
from __future__ import annotations
import ctypes
from ctypes import wintypes
import hashlib
import os
from pathlib import Path
from cmbuyer_client.core.errors import SingleInstanceError
class NamedMutex:
_ALREADY_EXISTS = 183
def __init__(self, database_path: Path) -> None:
if os.name != "nt":
raise SingleInstanceError("named_mutex_requires_windows")
canonical = str(database_path.expanduser().resolve()).casefold().encode("utf-8")
# Global namespace 覆盖同一 Windows 用户的多个交互 session;默认 DACL 不向其他用户泄露句柄。
name = "Global\\cmbuyer-" + hashlib.sha256(canonical).hexdigest()
kernel32 = ctypes.WinDLL("kernel32", use_last_error=True)
kernel32.CreateMutexW.argtypes = (wintypes.LPVOID, wintypes.BOOL, wintypes.LPCWSTR)
kernel32.CreateMutexW.restype = wintypes.HANDLE
kernel32.ReleaseMutex.argtypes = (wintypes.HANDLE,)
kernel32.ReleaseMutex.restype = wintypes.BOOL
kernel32.CloseHandle.argtypes = (wintypes.HANDLE,)
kernel32.CloseHandle.restype = wintypes.BOOL
ctypes.set_last_error(0)
handle = kernel32.CreateMutexW(None, True, name)
if not handle:
raise SingleInstanceError("named_mutex_failed")
if ctypes.get_last_error() == self._ALREADY_EXISTS:
kernel32.CloseHandle(handle)
raise SingleInstanceError("instance_already_running")
self._kernel32 = kernel32
self._handle = handle
def close(self) -> None:
handle = getattr(self, "_handle", None)
if handle:
self._kernel32.ReleaseMutex(handle)
self._kernel32.CloseHandle(handle)
self._handle = None
def __enter__(self) -> "NamedMutex":
return self
def __exit__(self, exc_type: object, exc: object, traceback: object) -> None:
self.close()
File diff suppressed because it is too large Load Diff
+16 -2
View File
@@ -23,6 +23,8 @@ _KEY_VALUE_PATTERN = re.compile(
flags=re.IGNORECASE,
)
_PHONE_PATTERN = re.compile(r"(?<!\d)1[3-9]\d{9}(?!\d)")
_BEARER_PATTERN = re.compile(r"(?i)\bBearer\s+[0-9a-f]{64}\b")
_BARE_TOKEN_PATTERN = re.compile(r"(?<![0-9a-fA-F])[0-9a-fA-F]{64}(?![0-9a-fA-F])")
def redact_text(message: str) -> str:
@@ -31,7 +33,9 @@ def redact_text(message: str) -> str:
def replace_key_value(match: re.Match[str]) -> str:
return f"{match.group('key')}{match.group('separator')}{REDACTED}"
redacted = _KEY_VALUE_PATTERN.sub(replace_key_value, message)
redacted = _BEARER_PATTERN.sub("Bearer " + REDACTED, message)
redacted = _KEY_VALUE_PATTERN.sub(replace_key_value, redacted)
redacted = _BARE_TOKEN_PATTERN.sub(REDACTED, redacted)
return _PHONE_PATTERN.sub(REDACTED, redacted)
@@ -47,6 +51,16 @@ class SensitiveDataFilter(logging.Filter):
return True
class RedactingFormatter(logging.Formatter):
"""再次处理完整格式化文本,覆盖异常 traceback 中的敏感值。"""
def format(self, record: logging.LogRecord) -> str:
return redact_text(super().format(record))
def formatException(self, exc_info: tuple[type[BaseException], BaseException, object]) -> str:
return redact_text(super().formatException(exc_info))
def configure_application_logger(paths: RuntimePaths) -> logging.Logger:
"""配置唯一的 UTF-8 文件日志,并确保其先经过脱敏过滤。"""
@@ -61,6 +75,6 @@ def configure_application_logger(paths: RuntimePaths) -> logging.Logger:
handler = logging.FileHandler(Path(paths.logs) / "client.log", encoding="utf-8")
handler.addFilter(SensitiveDataFilter())
handler.setFormatter(logging.Formatter("%(asctime)s %(levelname)s %(message)s"))
handler.setFormatter(RedactingFormatter("%(asctime)s %(levelname)s %(message)s"))
logger.addHandler(handler)
return logger
+10 -2
View File
@@ -1,15 +1,23 @@
"""拼多多链接的受限打开与只读取证。
"""拼多多链接的受限打开、只读取证与经取证的规格面板选择。
此包不提供页面选择器、输入、滑动、下单或支付能力。
此包不提供通用页面选择器、输入、滑动或任何订单动作。
"""
from .product_open import ProductOpenCapturer, ProductOpenResult
from .product_url import ProductUrl, ProductUrlError, parse_product_url
from .sku_selection import SkuSelection, SkuSelectionError, SkuSelectionFlow
from .sku_selection_runner import SkuSelectionRunError, SkuSelectionRunResult, SkuSelectionRunner
__all__ = [
"ProductOpenCapturer",
"ProductOpenResult",
"ProductUrl",
"ProductUrlError",
"SkuSelection",
"SkuSelectionError",
"SkuSelectionFlow",
"SkuSelectionRunError",
"SkuSelectionRunResult",
"SkuSelectionRunner",
"parse_product_url",
]
@@ -0,0 +1,436 @@
"""T-103 尺码显示动作的一次性真机取证;不属于生产采购 Flow。"""
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
from datetime import UTC, datetime
from hashlib import sha256
import json
from math import isfinite
import os
from pathlib import Path
import shutil
from time import monotonic, sleep
from typing import Any
from uuid import uuid4
from ..device.adb import AdbClient, DeviceConnectionError, DeviceInspection
from ..device.baseline import PDD_PACKAGE, _save_base64_screenshot, _sha256_file
from .product_url import parse_product_url
from .sku_selection import (
EXPECTED_GOODS_ID,
_COLOR_ONLY_SUMMARY,
_PANEL_SURFACE,
_PanelProfile,
_REVEAL_NOT_PROVEN,
_S_SIZE_UI,
_TARGET_COLOR_UI,
_TARGET_SIZE_UI,
_action_bounds,
_clickable_before,
_descendant,
_exact_color_action,
_exact_inert,
_exact_live_layout,
_exact_readonly_text,
_exact_recycler,
_is_size_action_text,
_one,
_parse_nodes,
_require_color_only_panel,
_require_color_action_chain,
_require_color_subtree,
_require_exact_selected_set,
_require_panel_chain,
_require_size_wrapper,
_require_size_action_chain,
_spec_for,
resolve_task_selection,
SkuSelectionError,
SkuSelectionFlow,
)
from .sku_selection_runner import (
EXPECTED_ANDROID_VERSION,
EXPECTED_DEVICE_MODEL,
EXPECTED_SCREEN_SIZE,
SkuSelectionRunError,
UiautomatorSkuPanelAdapter,
_require_expected_device,
_require_expected_version,
_require_screenshot_size,
)
_TARGET_URL = f"https://mobile.yangkeduo.com/goods.html?goods_id={EXPECTED_GOODS_ID}"
_REVEAL_START = (360, 1900)
_REVEAL_END = (360, 1300)
_REVEAL_STEPS = 30
class SkuRevealSpikeError(RuntimeError):
"""一次性 reveal 取证未形成可发布证据。"""
@dataclass(frozen=True)
class SkuRevealSpikeResult:
output_directory: Path
manifest_path: Path
class _RevealEvidenceAdapter(UiautomatorSkuPanelAdapter):
"""只为 spike 增加一个无参数、固定 profile 的一次性手势。"""
def __init__(self, device: Any, timeout_seconds: float) -> None:
super().__init__(device, timeout_seconds)
self._reveal_attempted = False
@property
def reveal_attempted(self) -> bool:
return self._reveal_attempted
def reveal_size_options_once(self) -> None:
if self._reveal_attempted:
raise SkuRevealSpikeError("规格显示动作已经尝试过,拒绝重试。")
# RPC 超时也可能表示手势已经送达,必须在调用前封存唯一机会。
self._reveal_attempted = True
self._call(
"jsonrpc_call",
"swipe",
[*_REVEAL_START, *_REVEAL_END, _REVEAL_STEPS],
timeout=self._timeout_seconds,
)
class SkuRevealSpikeCapturer:
"""打开已取证面板、选一次目标颜色,再采集唯一 reveal 的前后证据。"""
def __init__(
self,
adb_client: AdbClient,
connector: Callable[[str], Any],
timeout_seconds: float,
monotonic_clock: Callable[[], float] = monotonic,
sleep_function: Callable[[float], None] = sleep,
) -> None:
if not _positive_finite(timeout_seconds):
raise ValueError("timeout_seconds 必须是大于 0 的有限数值")
self._adb_client = adb_client
self._connector = connector
self._timeout_seconds = timeout_seconds
self._clock = monotonic_clock
self._sleep = sleep_function
def capture(
self,
serial: str,
goods_id: str,
output_directory: Path,
) -> SkuRevealSpikeResult:
if type(goods_id) is not str or goods_id != EXPECTED_GOODS_ID:
raise SkuRevealSpikeError("商品不是 T-103 已取证目标,已停止取证。")
link = parse_product_url(_TARGET_URL)
target = Path(output_directory)
_validate_new_target(target)
staging: Path | None = None
try:
staging = _prepare_staging(target)
deadline = self._clock() + self._timeout_seconds
inspection = self._adb_client.inspect(serial)
_require_expected_device(inspection)
adapter = _RevealEvidenceAdapter(self._connector(serial), self._timeout_seconds)
_require_expected_version(adapter.app_info(PDD_PACKAGE))
if adapter.display_size() != EXPECTED_SCREEN_SIZE:
raise SkuRevealSpikeError("设备不是已取证的竖屏坐标空间,已停止取证。")
pre_intent = adapter.dump_window_hierarchy()
self._adb_client.start_pdd_view_intent(serial, link.goods_id)
remaining = deadline - self._clock()
if remaining <= 0:
raise SkuRevealSpikeError("规格入口取证超时,未执行 reveal。")
flow = SkuSelectionFlow(
adapter,
entry_wait_timeout_seconds=remaining,
monotonic_clock=self._clock,
sleep_function=self._sleep,
)
flow.open_sku_panel(link.canonical_url, pre_intent)
try:
flow.select_sku_options(resolve_task_selection("黑色CHA(纯棉)", "M(建议100-115)"))
except SkuSelectionError as error:
if error.args != (_REVEAL_NOT_PROVEN,):
raise
else:
raise SkuRevealSpikeError("规格流程未停在已取证的仅颜色状态。")
before_hierarchy = adapter.dump_window_hierarchy()
before_nodes = _parse_nodes(before_hierarchy)
_require_color_only_panel(before_nodes)
_require_safe_reveal_path(before_nodes)
_capture_frame(adapter, staging / "before", before_hierarchy)
# 截图 RPC 期间页面也可能变化;真正发送手势前必须用新树再次证明同一前置与安全通道。
before_hierarchy = adapter.dump_window_hierarchy()
before_nodes = _parse_nodes(before_hierarchy)
_require_color_only_panel(before_nodes)
_require_safe_reveal_path(before_nodes)
_require_expected_version(adapter.app_info(PDD_PACKAGE))
current = adapter.app_current()
if not isinstance(current, dict) or current.get("package") != PDD_PACKAGE:
raise SkuRevealSpikeError("reveal 前拼多多不在前台,未执行手势。")
(staging / "before" / "hierarchy.xml").write_text(before_hierarchy, encoding="utf-8")
rpc_outcome = "completed"
try:
adapter.reveal_size_options_once()
except SkuSelectionRunError:
rpc_outcome = "ambiguous_reconciled"
projection, after_hierarchy = self._wait_for_candidate(adapter, deadline)
after_directory = staging / "after"
_capture_frame(adapter, after_directory, after_hierarchy)
reverified = adapter.dump_window_hierarchy()
if _candidate_projection(_parse_nodes(reverified)) != projection:
raise SkuRevealSpikeError("截图后候选状态漂移,未发布证据。")
(after_directory / "hierarchy.xml").write_text(reverified, encoding="utf-8")
manifest = _manifest(inspection, serial, rpc_outcome, staging)
manifest_path = staging / "manifest.json"
manifest_path.write_text(
json.dumps(manifest, ensure_ascii=False, indent=2, sort_keys=True) + "\n",
encoding="utf-8",
)
os.rename(staging, target)
staging = None
except (DeviceConnectionError, SkuSelectionError, SkuSelectionRunError, SkuRevealSpikeError):
_clean_staging(staging)
raise
except Exception as error:
_clean_staging(staging)
raise SkuRevealSpikeError("规格 reveal 取证未完成,未发布本地证据目录。") from error
return SkuRevealSpikeResult(target, target / "manifest.json")
def _wait_for_candidate(
self,
adapter: _RevealEvidenceAdapter,
deadline: float,
) -> tuple[tuple[tuple[str, ...], ...], str]:
stable: tuple[tuple[str, ...], ...] | None = None
while True:
_require_expected_version(adapter.app_info(PDD_PACKAGE))
current = adapter.app_current()
if not isinstance(current, dict) or current.get("package") != PDD_PACKAGE:
raise SkuRevealSpikeError("reveal 后拼多多不在前台,未发布证据。")
hierarchy = adapter.dump_window_hierarchy()
try:
projection = _candidate_projection(_parse_nodes(hierarchy))
except SkuSelectionError:
projection = None
if projection is not None and projection == stable:
return projection, hierarchy
stable = projection
remaining = deadline - self._clock()
if remaining <= 0:
raise SkuRevealSpikeError("reveal 后未形成稳定候选状态,未发布证据。")
self._sleep(min(0.2, remaining))
def _require_safe_reveal_path(nodes: list[Any]) -> None:
# 当前真机证据证明 x=360 是两列规格卡之间的空隙;必须验证完整线段,离散采样会漏掉窄浮层。
surface = _one(
[node for node in nodes if _exact_inert(node, "android.view.ViewGroup", _PANEL_SURFACE)],
"固定 reveal 通道无法绑定面板内容面。",
)
_require_panel_chain(surface)
content = surface.parent
action_root = content.parent if content is not None else None
action_parent = action_root.parent if action_root is not None else None
if (
content is None
or action_root is None
or action_parent is None
or not _exact_live_layout(action_root, "android.view.ViewGroup", "[0,366][1080,2328]")
or not _exact_live_layout(action_parent, "android.widget.LinearLayout", "[0,120][1080,2328]")
):
raise SkuRevealSpikeError("固定 reveal 通道祖先身份漂移,未执行手势。")
allowed = {id(action_root.element), id(action_parent.element)}
occupants: set[int] = set()
for node in nodes:
if node.element.get("clickable") != "true":
continue
try:
left, top, right, bottom = _action_bounds(node.bounds)
except SkuSelectionError as error:
raise SkuRevealSpikeError("可点击节点坐标不可验证,未执行手势。") from error
if (
left <= _REVEAL_START[0] < right
and top <= _REVEAL_START[1]
and bottom > _REVEAL_END[1]
):
occupants.add(id(node.element))
if occupants != allowed:
raise SkuRevealSpikeError("固定 reveal 通道被未取证可点击节点占用,未执行手势。")
if _REVEAL_START[1] >= 2079 or _REVEAL_END[1] >= 2079:
raise SkuRevealSpikeError("固定 reveal 通道越过规格内容区,未执行手势。")
def _candidate_projection(nodes: list[Any]) -> tuple[tuple[str, ...], ...]:
surface = _one(
[node for node in nodes if _exact_inert(node, "android.view.ViewGroup", _PANEL_SURFACE)],
"候选面板内容面不唯一。",
)
_require_panel_chain(surface)
header = _one(
[node for node in nodes if node.parent is surface and _exact_inert(node, "android.widget.LinearLayout", "[0,366][1080,1000]")],
"候选面板头部不唯一。",
)
outer = _one(
[node for node in nodes if node.parent is surface and _exact_recycler(node, "[0,1000][1080,2079]")],
"候选维度容器不唯一。",
)
price_row = _one(
[node for node in nodes if _descendant(node, header) and _exact_inert(node, "android.widget.LinearLayout", "[396,498][895,570]")],
"候选价格行不唯一。",
)
current = _one(
[node for node in nodes if node.parent is price_row and _exact_readonly_text(node, "快卖完 ¥12.88", "[396,503][712,570]")],
"候选当前价不唯一。",
)
original = _one(
[node for node in nodes if node.parent is price_row and _exact_readonly_text(node, "¥29.88", "[730,503][895,570]")],
"候选原价不唯一。",
)
if _clickable_before(current, surface):
raise SkuSelectionError("候选当前价位于可点击内容祖先下。")
summary = _one(
[node for node in nodes if _descendant(node, header) and _exact_readonly_text(node, _COLOR_ONLY_SUMMARY, "[396,654][1053,716]")],
"候选摘要不唯一。",
)
color_region = _one(
[node for node in nodes if _descendant(node, outer) and _exact_recycler(node, "[36,1000][1080,1483]")],
"候选颜色容器不唯一。",
)
color = _one(
[node for node in nodes if node.parent is color_region and _exact_color_action(node, "[372,1000][684,1024]", True)],
"候选目标颜色不唯一。",
)
rolled_spec = _spec_for(_PanelProfile.TARGETS_SELECTED)
selected_color_nodes = _require_color_subtree(color, rolled_spec)
_require_color_action_chain(color, color_region, outer, surface, rolled_spec)
size_label = _one(
[node for node in nodes if _descendant(node, outer) and _exact_readonly_text(node, "尺码", "[36,1506][114,1552]")],
"候选尺码标题不唯一。",
)
size_header = size_label.parent
if size_header is None or not _exact_live_layout(size_header, "android.widget.LinearLayout", "[36,1489][1044,1570]"):
raise SkuSelectionError("候选尺码标题父结构漂移。")
size_options = _one(
[node for node in nodes if _descendant(node, outer) and _exact_inert(node, "android.view.ViewGroup", "[36,1582][1044,1897]")],
"候选尺码 options 根不唯一。",
)
actions = [node for node in nodes if _descendant(node, size_options) and _is_size_action_text(node)]
s_action = _one(
[node for node in actions if node.text == _S_SIZE_UI and node.bounds == "[36,1582][409,1667]"],
"候选 S action 不唯一。",
)
m_action = _one(
[node for node in actions if node.text == _TARGET_SIZE_UI and node.bounds == "[439,1582][831,1667]"],
"候选 M action 不唯一。",
)
_require_size_wrapper(s_action, size_options)
_require_size_wrapper(m_action, size_options)
_require_size_action_chain(s_action, size_options, outer, surface)
_require_size_action_chain(m_action, size_options, outer, surface)
if any(node.element.get("selected") == "true" for node in actions):
raise SkuSelectionError("reveal 候选态已有尺码被选中。")
_require_exact_selected_set(nodes, surface, selected_color_nodes)
dangerous = [
node for node in nodes
if "提交订单" in node.text
and node.element.get("package") == PDD_PACKAGE
and node.element.get("class") == "android.widget.TextView"
]
if len(dangerous) != 1 or _action_bounds(dangerous[0].bounds)[1] < 2079:
raise SkuSelectionError("提交硬拒绝区位置不唯一。")
return tuple(
(
node.element.get("class", ""),
node.bounds,
node.text,
node.desc,
node.element.get("selected", ""),
node.element.get("clickable", ""),
)
for node in (surface, header, outer, price_row, current, original, summary, color_region, color, size_label, size_options, s_action, m_action, dangerous[0])
)
def _capture_frame(adapter: _RevealEvidenceAdapter, directory: Path, hierarchy: str) -> None:
directory.mkdir()
hierarchy_path = directory / "hierarchy.xml"
hierarchy_path.write_text(hierarchy, encoding="utf-8")
screenshot_path = directory / "screenshot.png"
_save_base64_screenshot(adapter.capture_screenshot(), screenshot_path)
_require_screenshot_size(screenshot_path)
def _manifest(
inspection: DeviceInspection,
serial: str,
rpc_outcome: str,
staging: Path,
) -> dict[str, Any]:
artifacts = []
for relative in (
"before/screenshot.png",
"before/hierarchy.xml",
"after/screenshot.png",
"after/hierarchy.xml",
):
path = staging / relative
artifacts.append({"path": relative, "sha256": _sha256_file(path)})
return {
"schema_version": 1,
"captured_at": datetime.now(UTC).isoformat(),
"operation": "t103-sku-reveal-evidence",
"profile_id": "pdd-8.17.0-size-reveal-gap-v1",
"product": {"goods_id": EXPECTED_GOODS_ID},
"channel": "wifi" if ":" in serial else "usb",
"serial_sha256": sha256(serial.encode("utf-8")).hexdigest(),
"device": {
"model": inspection.model,
"android_version": inspection.android_version,
"pdd_package": PDD_PACKAGE,
"pdd_version": "8.17.0",
},
"reveal_attempts": 1,
"rpc_outcome": rpc_outcome,
"candidate_status": "human_review_required",
"artifacts": artifacts,
}
def _validate_new_target(target: Path) -> None:
if target.exists() or not target.name:
raise SkuRevealSpikeError("输出目录必须是不存在的明确新目录。")
def _prepare_staging(target: Path) -> Path:
target.parent.mkdir(parents=True, exist_ok=True)
staging = target.parent / f".{target.name}.staging-{uuid4().hex}"
staging.mkdir()
return staging
def _clean_staging(staging: Path | None) -> None:
if staging is not None and staging.exists():
shutil.rmtree(staging)
def _positive_finite(value: object) -> bool:
return (
isinstance(value, (int, float))
and not isinstance(value, bool)
and value > 0
and isfinite(value)
)
@@ -0,0 +1,976 @@
"""T-103:仅限已取证 PDD 8.17.0 的规格面板恢复。"""
from __future__ import annotations
from dataclasses import dataclass
from enum import Enum
import re
from time import monotonic, sleep
from typing import Any, Callable, Protocol
from xml.etree import ElementTree
from ..device.baseline import PDD_PACKAGE
from .product_open import EXPECTED_PDD_VERSION
from .product_url import parse_product_url
EXPECTED_GOODS_ID = "937122477375"
EXPECTED_UNIT_PRICE = "12.88"
# 任务值不是页面判据;右侧是 v5 取证的唯一 accessibility 文案(空格/全角括号均有意义)。
TASK_TO_UI_SELECTION = {("黑色CHA(纯棉)", "M(建议100-115)"): ("黑色 CHA (纯棉)", "M(建议100-115)")}
_TARGET_COLOR_UI, _TARGET_SIZE_UI = next(iter(TASK_TO_UI_SELECTION.values()))
_ENTRY = "快要抢光 ¥ 12.88"
_ENTRY_PROMOTION_LABEL = "快要抢光"
_ENTRY_TEXT_BOUNDS = "[688,2184][1042,2253]"
_ENTRY_ACTION_DESC = "快要抢光¥12.88"
_ENTRY_ACTION_BOUNDS = "[446,2166][1080,2328]"
_ENTRY_SIBLING = "免拼购买"
_ENTRY_SIBLING_BOUNDS = "[688,2256][856,2305]"
_FORBIDDEN_ENTRY_ACTION_DESC = (
"购买", "下单", "付款", "订单", "单独购买", "直接拼成", "提交订单", "支付",
"先用后付", "0元下单", "0 元下单",
)
_SIZE = "尺码"
_W, _H = 1080, 2376
_PANEL_SURFACE = "[0,366][1080,2079]"
_EMPTY_SUMMARY = "请选择: 颜色分类 尺码"
_COLOR_ONLY_SUMMARY = "请选择: 尺码"
_S_SIZE_UI = "S(建议80-100)"
_S_SUMMARY = f"已选: {_TARGET_COLOR_UI} {_S_SIZE_UI}"
_TARGET_SUMMARY = f"已选: {_TARGET_COLOR_UI} {_TARGET_SIZE_UI}"
_REVEAL_NOT_PROVEN = "尺码仍在已取证视口外;受控显示动作尚未取证,已停止后续点击。"
_BOUNDS = re.compile(r"^\[(\d+),(\d+)\]\[(\d+),(\d+)\]$")
_ROLLED_CURRENT_PRICE = "快卖完 ¥12.88"
_ROLLED_ORIGINAL_PRICE = "¥29.88"
_BAD_PRICE_ROLE = ("提交订单", "支付", "优惠", "券", "会员", "补贴", "区间", "实付", "到手", "原价", "划线价", "最低", "低至", "起价", "下单", "先用后付", "预估")
class SkuSelectionError(RuntimeError):
"""已取证判据不成立时的脱敏停止。"""
_SKU_ENTRY_FAILURE_STAGES = frozenset(
(
"sku_entry_pre_intent",
"sku_entry_discovery",
"sku_entry_click",
"sku_entry_panel_verify",
)
)
_SKU_ENTRY_FAILURE_MARKER = object()
def _annotate_sku_entry_failure(error: BaseException, stage: str) -> None:
"""把 Flow 实际经过的固定入口子阶段附到原异常,不改变异常类型。"""
if type(stage) is not str or stage not in _SKU_ENTRY_FAILURE_STAGES:
return
try:
# marker 最后写入:若第三方异常拒绝任一属性写入,就不能形成可信诊断。
setattr(error, "_cmbuyer_failure_stage", stage)
setattr(error, "_cmbuyer_sku_entry_failure_marker", _SKU_ENTRY_FAILURE_MARKER)
except BaseException:
pass
def _safe_sku_entry_failure_stage(error: BaseException) -> str | None:
"""只读取由本模块写入的入口子阶段;任意异常自报的值不可信。"""
try:
marker = getattr(error, "_cmbuyer_sku_entry_failure_marker", None)
stage = getattr(error, "_cmbuyer_failure_stage", None)
if marker is not _SKU_ENTRY_FAILURE_MARKER or type(stage) is not str:
return None
return stage if stage in _SKU_ENTRY_FAILURE_STAGES else None
except BaseException:
return None
class SkuPanelDevice(Protocol):
def app_info(self, package_name: str) -> dict[str, Any]: ...
def app_current(self) -> dict[str, Any]: ...
def dump_window_hierarchy(self) -> str: ...
def tap_sku_entry(self, bounds: str) -> None: ...
def tap_sku_option(self, bounds: str) -> None: ...
def leave_sku_panel(self) -> None: ...
@dataclass(frozen=True)
class SkuSelection:
color: str
size: str
def resolve_task_selection(color: str, size: str) -> SkuSelection:
mapped = TASK_TO_UI_SELECTION.get((color, size))
if mapped is None:
raise SkuSelectionError("规格任务值不是已取证的唯一目标,已停止操作。")
return SkuSelection(*mapped)
@dataclass(frozen=True)
class _Node:
element: ElementTree.Element
parent: "_Node | None"
@property
def text(self) -> str: return self.element.get("text", "")
@property
def desc(self) -> str: return self.element.get("content-desc", "")
@property
def bounds(self) -> str: return self.element.get("bounds", "")
class _PanelProfile(Enum):
PANEL_OPEN_EMPTY = "panel_open_empty"
COLOR_SELECTED_SIZE_HIDDEN = "color_selected_size_hidden"
SIZE_VISIBLE_NON_TARGET = "size_visible_non_target"
TARGETS_SELECTED = "targets_selected"
@dataclass(frozen=True)
class _PanelSpec:
profile: _PanelProfile
header_bounds: str
outer_bounds: str
price_row_bounds: str
current_text: str
current_bounds: str
original_text: str
original_bounds: str
summary_text: str
summary_bounds: str
color_label_bounds: str | None
color_region_bounds: str
color_bounds: str
color_selected: bool
size_label_bounds: str
selected_size: str | None
_PANEL_SPECS = (
_PanelSpec(
_PanelProfile.PANEL_OPEN_EMPTY,
"[0,366][1080,1077]", "[0,1077][1080,2079]",
"[396,575][912,647]", "限1件 ¥12.88 ", "[396,580][675,647]",
"券前¥29.88", "[693,580][912,647]",
_EMPTY_SUMMARY, "[396,731][1053,793]", "[36,1106][192,1159]",
"[36,1188][1080,2046]", "[372,1188][684,1587]", False,
"[36,2069][114,2079]", None,
),
_PanelSpec(
_PanelProfile.COLOR_SELECTED_SIZE_HIDDEN,
"[0,366][1080,1077]", "[0,1077][1080,2079]",
"[396,575][912,647]", "限1件 ¥12.88 ", "[396,580][675,647]",
"券前¥29.88", "[693,580][912,647]",
_COLOR_ONLY_SUMMARY, "[396,731][1053,793]", "[36,1106][192,1159]",
"[36,1188][1080,2046]", "[372,1188][684,1587]", True,
"[36,2069][114,2079]", None,
),
_PanelSpec(
_PanelProfile.SIZE_VISIBLE_NON_TARGET,
"[0,366][1080,1000]", "[0,1000][1080,2079]",
"[396,498][895,570]", _ROLLED_CURRENT_PRICE, "[396,503][712,570]",
_ROLLED_ORIGINAL_PRICE, "[730,503][895,570]",
_S_SUMMARY, "[396,654][1053,716]", None,
"[36,1000][1080,1483]", "[372,1000][684,1024]", True,
"[36,1506][114,1552]", _S_SIZE_UI,
),
_PanelSpec(
_PanelProfile.TARGETS_SELECTED,
"[0,366][1080,1000]", "[0,1000][1080,2079]",
"[396,498][895,570]", _ROLLED_CURRENT_PRICE, "[396,503][712,570]",
_ROLLED_ORIGINAL_PRICE, "[730,503][895,570]",
_TARGET_SUMMARY, "[396,654][1053,716]", None,
"[36,1000][1080,1483]", "[372,1000][684,1024]", True,
"[36,1506][114,1552]", _TARGET_SIZE_UI,
),
)
class SkuSelectionFlow:
def __init__(self, device: SkuPanelDevice, entry_wait_timeout_seconds: float = 0.2,
entry_poll_interval_seconds: float = 0.2, monotonic_clock: Callable[[], float] = monotonic,
sleep_function: Callable[[float], None] = sleep) -> None:
if entry_wait_timeout_seconds < 0 or entry_poll_interval_seconds <= 0:
raise ValueError("入口等待参数无效。")
self._device, self._entry_timeout, self._poll = device, entry_wait_timeout_seconds, entry_poll_interval_seconds
self._clock, self._sleep = monotonic_clock, sleep_function
self._pending: tuple[str, Callable[[list[_Node]], Any]] | None = None
def open_sku_panel(self, product_url: str, pre_intent_hierarchy: str | None = None) -> None:
try:
if parse_product_url(product_url).goods_id != EXPECTED_GOODS_ID:
raise SkuSelectionError("商品不是已取证目标,已停止操作。")
if pre_intent_hierarchy is not None:
previous_nodes = _parse_nodes(pre_intent_hierarchy)
# intent 前只判断旧页是否已经存在完整入口链;浮层或额外动作节点不能把旧商品
# 伪装成“不安全所以不存在”,否则 intent 后可能误把旧页当成新目标页。
if _physical_entries(previous_nodes):
raise SkuSelectionError("intent 前页面已出现规格入口,已拒绝旧商品误点。")
except BaseException as error:
_annotate_sku_entry_failure(error, "sku_entry_pre_intent")
raise
try:
entry, before = self._wait_for_entry(pre_intent_hierarchy)
_action_bounds(entry.bounds)
except BaseException as error:
_annotate_sku_entry_failure(error, "sku_entry_discovery")
raise
try:
self._pending = (before, _require_empty_panel)
self._device.tap_sku_entry(entry.bounds)
except BaseException as error:
_annotate_sku_entry_failure(error, "sku_entry_click")
raise
try:
self._wait_after_action(before, _require_empty_panel)
except BaseException as error:
_annotate_sku_entry_failure(error, "sku_entry_panel_verify")
raise
def select_sku_options(self, selection: SkuSelection) -> None:
if selection != SkuSelection(_TARGET_COLOR_UI, _TARGET_SIZE_UI):
raise SkuSelectionError("规格 UI 文案不是获准目标,已停止操作。")
self._require_foreground()
before = self._read_hierarchy()
nodes = _parse_nodes(before)
profile = _classify_panel(nodes)
if profile is _PanelProfile.PANEL_OPEN_EMPTY:
target = _target_color_action(nodes, selected=False)
_action_bounds(target.bounds)
_require_action_occupants(nodes, target)
self._pending = (before, _require_color_only_panel)
self._device.tap_sku_option(target.bounds)
self._wait_after_action(before, _require_color_only_panel)
# 当前证据只证明颜色选择;尺码仍在视口外。没有动作证据时必须在此停住,
# 不能把一次通用 swipe 或下一次规格点击伪装成已验证流程。
raise SkuSelectionError(_REVEAL_NOT_PROVEN)
if profile is _PanelProfile.COLOR_SELECTED_SIZE_HIDDEN:
raise SkuSelectionError(_REVEAL_NOT_PROVEN)
if profile is _PanelProfile.SIZE_VISIBLE_NON_TARGET:
target = _target_size_action(nodes, selected=False)
_action_bounds(target.bounds)
_require_action_occupants(nodes, target)
self._pending = (before, _require_target_panel)
self._device.tap_sku_option(target.bounds)
self._wait_after_action(before, _require_target_panel)
return
if profile is _PanelProfile.TARGETS_SELECTED:
return
raise SkuSelectionError("规格面板状态不属于已取证 profile,已停止操作。")
def read_sku_unit_price(self) -> str:
return _unit_price(self._verified_nodes())
def verify_target_selection_and_read_price(self, selection: SkuSelection) -> str:
if selection != SkuSelection(_TARGET_COLOR_UI, _TARGET_SIZE_UI):
raise SkuSelectionError("规格 UI 文案不是获准目标,已停止读取。")
nodes = self._verified_nodes()
return _unit_price(nodes)
def exit_sku_panel_safely(self) -> None:
self._require_foreground()
before = self._read_hierarchy()
_classify_panel(_parse_nodes(before))
self._device.leave_sku_panel()
deadline = self._clock() + self._entry_timeout
while True:
self._require_foreground()
raw = self._read_hierarchy()
if raw != before:
try:
_classify_panel(_parse_nodes(raw))
except SkuSelectionError:
return
remaining = deadline - self._clock()
if remaining <= 0:
raise SkuSelectionError("安全退出后未确认离开规格面板,未重试返回。")
self._sleep(min(self._poll, remaining))
def reconcile_pending_action(self) -> None:
"""仅只读调和一次已发出但尚未得到后置条件确认的动作。"""
if self._pending is None:
return
before, condition = self._pending
self._wait_after_action(before, condition)
def _wait_for_entry(self, previous: str | None) -> tuple[_Node, str]:
deadline, stable = self._clock() + self._entry_timeout, None
while True:
self._require_version()
current = self._device.app_current()
if isinstance(current, dict) and current.get("package") == PDD_PACKAGE:
raw = self._read_hierarchy()
nodes = _parse_nodes(raw)
entries = _eligible_entries(nodes)
if len(entries) > 1:
raise SkuSelectionError("商品页规格入口不唯一,已停止操作。")
if len(entries) == 1 and raw != previous:
# 商品详情正文包含倒计时等动态节点,全文 XML 稳定不是已取证入口的安全属性。
# 连续两帧只比较证据绑定的底部入口、直接父容器和不可点击兄弟节点;商品正文及
# 父容器内其他非危险动态节点不是入口身份,不能迫使实现退回坐标兜底。
projection = _entry_projection(entries[0], nodes)
if projection is None:
raise SkuSelectionError("商品页规格入口结构失效,已停止操作。")
if stable == projection:
return entries[0], raw
stable = projection
else:
stable = None
else:
stable = None
remaining = deadline - self._clock()
if remaining <= 0:
raise SkuSelectionError("等待已取证规格入口超时,未执行点击。")
self._sleep(min(self._poll, remaining))
def _wait_after_action(self, previous: str, condition: Callable[[list[_Node]], Any]) -> list[_Node]:
deadline = self._clock() + self._entry_timeout
while True:
self._require_foreground()
raw = self._read_hierarchy()
if raw != previous:
nodes = _parse_nodes(raw)
try:
condition(nodes)
self._pending = None
return nodes
except SkuSelectionError:
pass
remaining = deadline - self._clock()
if remaining <= 0:
raise SkuSelectionError("动作后页面未在限定时间内满足已取证后置条件,未重试动作。")
self._sleep(min(self._poll, remaining))
def _verified_nodes(self) -> list[_Node]:
self._require_foreground()
nodes = self._read_nodes()
_classify_panel(nodes)
return nodes
def _require_version(self) -> None:
info = self._device.app_info(PDD_PACKAGE)
version = (info.get("versionName") or info.get("version_name")) if isinstance(info, dict) else None
if version != EXPECTED_PDD_VERSION:
raise SkuSelectionError("拼多多版本与已取证版本不一致,已停止操作。")
def _require_foreground(self) -> None:
self._require_version()
current = self._device.app_current()
if not isinstance(current, dict) or current.get("package") != PDD_PACKAGE:
raise SkuSelectionError("拼多多不在前台,已停止操作。")
def _read_hierarchy(self) -> str:
try: raw = self._device.dump_window_hierarchy()
except Exception as error: raise SkuSelectionError("节点树读取失败,已停止操作。") from error
if not isinstance(raw, str) or not raw: raise SkuSelectionError("节点树不可用,已停止操作。")
return raw
def _read_nodes(self) -> list[_Node]: return _parse_nodes(self._read_hierarchy())
def _parse_nodes(raw: str) -> list[_Node]:
try: root = ElementTree.fromstring(raw)
except ElementTree.ParseError as error: raise SkuSelectionError("节点树格式无效,已停止操作。") from error
if root.tag != "hierarchy": raise SkuSelectionError("节点树根节点无效,已停止操作。")
result: list[_Node] = []
def visit(element: ElementTree.Element, parent: _Node | None) -> None:
node = _Node(element, parent); result.append(node)
for child in element: visit(child, node)
visit(root, None)
return result
def _classify_panel(nodes: list[_Node]) -> _PanelProfile:
matches: list[_PanelProfile] = []
for spec in _PANEL_SPECS:
try:
_match_panel_profile(nodes, spec)
except SkuSelectionError:
continue
matches.append(spec.profile)
if len(matches) != 1:
raise SkuSelectionError("规格面板不符合唯一完整取证 profile,已停止操作。")
return matches[0]
def _require_empty_panel(nodes: list[_Node]) -> None:
_require_profile(nodes, _PanelProfile.PANEL_OPEN_EMPTY)
def _require_color_only_panel(nodes: list[_Node]) -> None:
_require_profile(nodes, _PanelProfile.COLOR_SELECTED_SIZE_HIDDEN)
def _require_target_panel(nodes: list[_Node]) -> None:
_require_profile(nodes, _PanelProfile.TARGETS_SELECTED)
def _require_profile(nodes: list[_Node], expected: _PanelProfile) -> None:
if _classify_panel(nodes) is not expected:
raise SkuSelectionError("规格面板动作后状态与已取证 profile 不一致,已停止操作。")
def _match_panel_profile(nodes: list[_Node], spec: _PanelSpec) -> None:
surface = _one(
[node for node in nodes if _exact_inert(node, "android.view.ViewGroup", _PANEL_SURFACE)],
"规格面板内容面不唯一。",
)
_require_panel_chain(surface)
header = _one(
[node for node in nodes if node.parent is surface and _exact_inert(node, "android.widget.LinearLayout", spec.header_bounds)],
"规格面板头部不唯一。",
)
outer = _one(
[node for node in nodes if node.parent is surface and _exact_recycler(node, spec.outer_bounds)],
"规格面板维度容器不唯一。",
)
price_row = _one(
[node for node in nodes if _descendant(node, header) and _exact_inert(node, "android.widget.LinearLayout", spec.price_row_bounds)],
"规格面板价格行不唯一。",
)
if len([child for child in price_row.element if child.tag == "node"]) != 2:
raise SkuSelectionError("规格面板价格行子节点数量漂移。")
current = _one(
[node for node in nodes if node.parent is price_row and _exact_readonly_text(node, spec.current_text, spec.current_bounds)],
"规格面板当前价角色不唯一。",
)
_one(
[node for node in nodes if node.parent is price_row and _exact_readonly_text(node, spec.original_text, spec.original_bounds)],
"规格面板原价角色不唯一。",
)
if _clickable_before(current, surface):
raise SkuSelectionError("规格面板价格角色位于可点击内容祖先下。")
_one(
[node for node in nodes if _descendant(node, header) and _exact_readonly_text(node, spec.summary_text, spec.summary_bounds)],
"规格面板摘要不唯一。",
)
color_region = _one(
[node for node in nodes if _descendant(node, outer) and _exact_recycler(node, spec.color_region_bounds)],
"规格面板颜色容器不唯一。",
)
if spec.color_label_bounds is None:
if any(node.text == "颜色分类" and _readonly(node) for node in nodes):
raise SkuSelectionError("滚动态出现未取证颜色标题。")
else:
_one(
[node for node in nodes if _descendant(node, outer) and not _descendant(node, color_region) and _exact_readonly_text(node, "颜色分类", spec.color_label_bounds)],
"规格面板颜色标题不唯一。",
)
color = _one(
[node for node in nodes if node.parent is color_region and _exact_color_action(node, spec.color_bounds, spec.color_selected)],
"目标颜色 action 不唯一。",
)
selected_nodes = _require_color_subtree(color, spec)
_require_color_action_chain(color, color_region, outer, surface, spec)
size_label = _one(
[node for node in nodes if _descendant(node, outer) and not _descendant(node, color_region) and _exact_readonly_text(node, _SIZE, spec.size_label_bounds)],
"规格面板尺码标题不唯一。",
)
if spec.selected_size is None:
if any(
node.text in {_S_SIZE_UI, _TARGET_SIZE_UI}
and _is_size_action_text(node)
for node in nodes
):
raise SkuSelectionError("尺码隐藏 profile 出现可点击尺码。")
_require_exact_selected_set(nodes, surface, selected_nodes)
return
size_header = size_label.parent
if size_header is None or not _exact_live_layout(size_header, "android.widget.LinearLayout", "[36,1489][1044,1570]"):
raise SkuSelectionError("规格面板尺码标题父结构漂移。")
size_options = _one(
[node for node in nodes if _descendant(node, outer) and _exact_inert(node, "android.view.ViewGroup", "[36,1582][1044,1897]")],
"规格面板尺码 options 根不唯一。",
)
size_actions = [
node for node in nodes
if _descendant(node, size_options)
and _is_size_action_text(node)
]
s_action = _one([node for node in size_actions if node.text == _S_SIZE_UI and node.bounds == "[36,1582][409,1667]"], "S 尺码 action 不唯一。")
m_action = _one([node for node in size_actions if node.text == _TARGET_SIZE_UI and node.bounds == "[439,1582][831,1667]"], "M 尺码 action 不唯一。")
_require_size_wrapper(s_action, size_options)
_require_size_wrapper(m_action, size_options)
_require_size_action_chain(s_action, size_options, outer, surface)
_require_size_action_chain(m_action, size_options, outer, surface)
selected_sizes = [node for node in size_actions if node.element.get("selected") == "true"]
expected_action = s_action if spec.selected_size == _S_SIZE_UI else m_action
if selected_sizes != [expected_action]:
raise SkuSelectionError("尺码维度 selected 状态不唯一。")
_require_exact_selected_set(nodes, surface, [*selected_nodes, expected_action])
def _unit_price(nodes: list[_Node]) -> str:
_require_profile(nodes, _PanelProfile.TARGETS_SELECTED)
spec = _spec_for(_PanelProfile.TARGETS_SELECTED)
candidates = [
node for node in nodes
if _exact_readonly_text(node, _ROLLED_CURRENT_PRICE, spec.current_bounds)
]
current = _one(candidates, "规格面板现价不唯一,已停止读取。")
surface = _one([node for node in nodes if _exact_inert(node, "android.view.ViewGroup", _PANEL_SURFACE)], "规格面板内容面不唯一。")
if _clickable_before(current, surface) or any(word in current.text for word in _BAD_PRICE_ROLE):
raise SkuSelectionError("规格面板现价角色不可安全读取。")
return EXPECTED_UNIT_PRICE
def _target_color_action(nodes: list[_Node], *, selected: bool) -> _Node:
profile = _classify_panel(nodes)
expected_profile = _PanelProfile.COLOR_SELECTED_SIZE_HIDDEN if selected else _PanelProfile.PANEL_OPEN_EMPTY
if profile is not expected_profile:
raise SkuSelectionError("目标颜色 action 不属于预期 profile。")
spec = _spec_for(profile)
return _one([node for node in nodes if _exact_color_action(node, spec.color_bounds, selected)], "目标颜色 action 不唯一。")
def _target_size_action(nodes: list[_Node], *, selected: bool) -> _Node:
profile = _classify_panel(nodes)
expected_profile = _PanelProfile.TARGETS_SELECTED if selected else _PanelProfile.SIZE_VISIBLE_NON_TARGET
if profile is not expected_profile:
raise SkuSelectionError("目标尺码 action 不属于预期 profile。")
return _one(
[node for node in nodes if _is_size_action_text(node) and node.text == _TARGET_SIZE_UI and node.bounds == "[439,1582][831,1667]" and node.element.get("selected") == str(selected).lower()],
"目标尺码 action 不唯一。",
)
def _spec_for(profile: _PanelProfile) -> _PanelSpec:
return next(spec for spec in _PANEL_SPECS if spec.profile is profile)
def _exact_inert(node: _Node, class_name: str, bounds: str) -> bool:
return _exact_common(node, class_name, bounds, clickable="false", selected="false", scrollable="false") and not node.text and not node.desc
def _exact_recycler(node: _Node, bounds: str) -> bool:
return _exact_common(node, "androidx.recyclerview.widget.RecyclerView", bounds, clickable="false", selected="false", scrollable="true") and not node.text and not node.desc
def _exact_readonly_text(node: _Node, text: str, bounds: str) -> bool:
return _exact_common(node, "android.widget.TextView", bounds, clickable="false", selected="false", scrollable="false") and node.text == text and not node.desc
def _exact_color_action(node: _Node, bounds: str, selected: bool) -> bool:
return (
_exact_common(node, "android.view.ViewGroup", bounds, clickable="true", selected=str(selected).lower(), scrollable="false")
and not node.text
and node.desc == _TARGET_COLOR_UI
)
def _exact_live_layout(node: _Node, class_name: str, bounds: str) -> bool:
return _exact_common(node, class_name, bounds, clickable="true", selected="false", scrollable="false") and not node.text and not node.desc
def _exact_common(node: _Node, class_name: str, bounds: str, *, clickable: str, selected: str, scrollable: str) -> bool:
return (
node.element.get("package") == PDD_PACKAGE
and node.element.get("class") == class_name
and node.bounds == bounds
and node.element.get("clickable") == clickable
and node.element.get("enabled") == "true"
and node.element.get("visible-to-user") == "true"
and node.element.get("selected") == selected
and node.element.get("scrollable") == scrollable
)
def _is_size_action_text(node: _Node) -> bool:
return (
node.element.get("package") == PDD_PACKAGE
and node.element.get("class") == "android.widget.TextView"
and node.element.get("clickable") == "true"
and node.element.get("enabled") == "true"
and node.element.get("visible-to-user") == "true"
and node.element.get("selected") in {"true", "false"}
and node.element.get("scrollable") == "false"
and not node.desc
)
def _require_size_wrapper(action: _Node, size_options: _Node) -> None:
wrapper = action.parent
if (
wrapper is None
or wrapper.parent is not size_options
or not _exact_live_layout(wrapper, "android.view.ViewGroup", action.bounds)
or wrapper.element[0] is not action.element
):
raise SkuSelectionError("尺码 action 父结构漂移。")
children = [child for child in wrapper.element if child.tag == "node"]
if action.element.get("selected") == "true":
if len(children) != 2:
raise SkuSelectionError("已选尺码指示子树数量漂移。")
marker = _Node(children[1], wrapper)
if (
not _exact_common(
marker,
"android.view.View",
action.bounds,
clickable="false",
selected="false",
scrollable="false",
)
or marker.text
or marker.desc
):
raise SkuSelectionError("已选尺码指示节点结构漂移。")
elif len(children) != 1:
raise SkuSelectionError("未选尺码 action 子树数量漂移。")
def _require_panel_chain(surface: _Node) -> None:
expected = (
("android.widget.LinearLayout", "[0,366][1080,2079]", "false"),
("android.view.ViewGroup", "[0,366][1080,2328]", "true"),
("android.widget.LinearLayout", "[0,120][1080,2328]", "true"),
("android.widget.FrameLayout", "[0,120][1080,2328]", "false"),
("android.widget.FrameLayout", "[0,120][1080,2328]", "false"),
("android.widget.LinearLayout", "[0,0][1080,2328]", "false"),
("android.widget.FrameLayout", "[0,0][1080,2376]", "false"),
)
node = surface.parent
for class_name, bounds, clickable in expected:
if (
node is None
or not _exact_common(
node,
class_name,
bounds,
clickable=clickable,
selected="false",
scrollable="false",
)
or node.text
or node.desc
):
raise SkuSelectionError("规格面板祖先链不符合已取证结构。")
node = node.parent
if node is None or node.element.tag != "hierarchy" or node.parent is not None:
raise SkuSelectionError("规格面板根节点结构漂移。")
def _require_color_subtree(color: _Node, spec: _PanelSpec) -> list[_Node]:
selected = str(spec.color_selected).lower()
children = [child for child in color.element if child.tag == "node"]
expected_selected: list[_Node] = [color] if spec.color_selected else []
if spec.color_bounds == "[372,1188][684,1587]":
if len(children) != 4:
raise SkuSelectionError("目标颜色完整卡片子树数量漂移。")
child_nodes = [node for node in _walk_direct_children(color)]
view, image, zoom, label_layout = child_nodes
if not _exact_common(view, "android.view.View", spec.color_bounds, clickable="false", selected=selected, scrollable="false") or view.text or view.desc:
raise SkuSelectionError("目标颜色选中遮罩结构漂移。")
if not _exact_common(image, "android.widget.ImageView", "[372,1188][684,1500]", clickable="true", selected=selected, scrollable="false") or image.text or image.desc != _TARGET_COLOR_UI:
raise SkuSelectionError("目标颜色图片 action 结构漂移。")
if not _exact_common(zoom, "android.widget.ImageView", "[372,1188][483,1299]", clickable="true", selected=selected, scrollable="false") or zoom.text or zoom.desc != "打开大图":
raise SkuSelectionError("目标颜色大图 action 结构漂移。")
if not _exact_common(label_layout, "android.widget.LinearLayout", "[372,1479][684,1587]", clickable="false", selected=selected, scrollable="false") or label_layout.text or label_layout.desc or len(label_layout.element) != 1:
raise SkuSelectionError("目标颜色文字容器结构漂移。")
label = _Node(label_layout.element[0], label_layout)
if not _exact_common(label, "android.widget.TextView", "[372,1479][684,1587]", clickable="true", selected=selected, scrollable="false") or label.text != _TARGET_COLOR_UI or label.desc:
raise SkuSelectionError("目标颜色文字 action 结构漂移。")
if spec.color_selected:
expected_selected.extend((view, image, zoom, label_layout, label))
return expected_selected
if len(children) != 2:
raise SkuSelectionError("目标颜色滚动态子树数量漂移。")
view, label_layout = [node for node in _walk_direct_children(color)]
if not _exact_common(view, "android.view.View", spec.color_bounds, clickable="false", selected=selected, scrollable="false") or view.text or view.desc:
raise SkuSelectionError("目标颜色滚动态遮罩结构漂移。")
if not _exact_common(label_layout, "android.widget.LinearLayout", spec.color_bounds, clickable="false", selected=selected, scrollable="false") or label_layout.text or label_layout.desc or len(label_layout.element) != 1:
raise SkuSelectionError("目标颜色滚动态文字容器漂移。")
label = _Node(label_layout.element[0], label_layout)
if not _exact_common(label, "android.widget.TextView", spec.color_bounds, clickable="true", selected=selected, scrollable="false") or label.text != _TARGET_COLOR_UI or label.desc:
raise SkuSelectionError("目标颜色滚动态文字 action 漂移。")
if spec.color_selected:
expected_selected.extend((view, label_layout, label))
return expected_selected
def _require_color_action_chain(
color: _Node,
color_region: _Node,
outer: _Node,
surface: _Node,
spec: _PanelSpec,
) -> None:
if color.parent is not color_region:
raise SkuSelectionError("目标颜色 action 不属于已取证颜色容器。")
if spec.color_bounds == "[372,1188][684,1587]":
expected = (
("android.widget.FrameLayout", "[0,1188][1080,2046]"),
("android.widget.LinearLayout", "[0,1188][1080,2052]"),
("android.widget.LinearLayout", "[0,1077][1080,2052]"),
)
else:
expected = (
("android.widget.FrameLayout", "[0,1000][1080,1483]"),
("android.widget.LinearLayout", "[0,1000][1080,1489]"),
("android.widget.LinearLayout", "[0,1000][1080,1489]"),
)
node = color_region.parent
for class_name, bounds in expected:
if node is None or not _exact_inert(node, class_name, bounds):
raise SkuSelectionError("目标颜色 action 父链不符合已取证结构。")
node = node.parent
if node is not outer or outer.parent is not surface:
raise SkuSelectionError("目标颜色 action 未沿已取证维度父链回到面板。")
def _require_size_action_chain(
action: _Node,
size_options: _Node,
outer: _Node,
surface: _Node,
) -> None:
wrapper = action.parent
if wrapper is None or wrapper.parent is not size_options:
raise SkuSelectionError("尺码 action 不属于已取证 options 容器。")
node = size_options.parent
for class_name, bounds in (
("android.widget.LinearLayout", "[36,1582][1044,1897]"),
("android.widget.LinearLayout", "[0,1489][1080,1930]"),
):
if node is None or not _exact_inert(node, class_name, bounds):
raise SkuSelectionError("尺码 action 父链不符合已取证结构。")
node = node.parent
if node is not outer or outer.parent is not surface:
raise SkuSelectionError("尺码 action 未沿已取证维度父链回到面板。")
def _walk_direct_children(parent: _Node) -> list[_Node]:
return [_Node(child, parent) for child in parent.element if child.tag == "node"]
def _require_exact_selected_set(nodes: list[_Node], surface: _Node, expected: list[_Node]) -> None:
actual = [
node for node in nodes
if _descendant(node, surface) and node.element.get("selected") == "true"
]
if {id(node.element) for node in actual} != {id(node.element) for node in expected}:
raise SkuSelectionError("规格面板 selected 节点集合与完整 profile 不一致。")
def _require_action_occupants(nodes: list[_Node], target: _Node) -> None:
left, top, right, bottom = _action_bounds(target.bounds)
point = (left + (right - left) // 2, top + (bottom - top) // 2)
allowed: set[int] = {id(target.element)}
parent = target.parent
while parent is not None:
if _is_live_clickable(parent):
if not (
_exact_live_layout(parent, "android.view.ViewGroup", target.bounds)
or _exact_live_layout(parent, "android.view.ViewGroup", "[0,366][1080,2328]")
or _exact_live_layout(parent, "android.widget.LinearLayout", "[0,120][1080,2328]")
):
raise SkuSelectionError("规格 action 存在未取证可点击祖先,已停止操作。")
allowed.add(id(parent.element))
parent = parent.parent
for node in nodes:
if _descendant(node, target) and _is_live_clickable(node):
node_left, node_top, node_right, node_bottom = _action_bounds(node.bounds)
if node_left <= point[0] < node_right and node_top <= point[1] < node_bottom:
allowed.add(id(node.element))
occupants = _live_clickables_covering(nodes, point)
if {id(node.element) for node in occupants} != allowed:
raise SkuSelectionError("规格 action 中心存在未取证可点击占用,已停止操作。")
def _clickable_before(node: _Node, stop: _Node) -> bool:
parent = node.parent
while parent is not None and parent.element is not stop.element:
if parent.element.get("clickable") == "true":
return True
parent = parent.parent
# 当前价必须能沿已取证内容祖先回到面板内容面;不接受另一个树枝上的同坐标文本。
return parent is None
def _descendant(node: _Node, ancestor: _Node) -> bool:
parent = node.parent
while parent is not None:
if parent.element is ancestor.element: return True
parent = parent.parent
return False
def _readonly(node: _Node) -> bool:
return node.element.get("package") == PDD_PACKAGE and node.element.get("class") == "android.widget.TextView" and node.element.get("clickable") == "false" and node.element.get("enabled") == "true" and node.element.get("visible-to-user") == "true"
def _live(node: _Node) -> bool:
return node.element.get("package") == PDD_PACKAGE and node.element.get("clickable") == "true" and node.element.get("enabled") == "true" and node.element.get("visible-to-user") == "true" and bool(node.bounds)
def _eligible_entries(nodes: list[_Node]) -> list[_Node]:
# 入口文本本身不可点击:必须证明它仍是已取证底部父容器的直接子节点,但动作坐标继续
# 使用第一行文本的窄 bounds,避免把父容器中心或第二行“免拼购买”变成坐标兜底。
if any(
node.bounds == _PANEL_SURFACE
and node.element.get("package") == PDD_PACKAGE
and node.element.get("class") == "android.view.ViewGroup"
for node in nodes
):
return []
entries = _physical_entries(nodes)
if len(entries) != 1:
return entries
entry = entries[0]
chain = _physical_entry_chain(entry, nodes)
if chain is None or _entry_projection(entry, nodes) is None:
return []
ancestor = chain[1]
left, top, right, bottom = _action_bounds(entry.bounds)
center = (left + (right - left) // 2, top + (bottom - top) // 2)
occupants = _live_clickables_covering(nodes, center)
# RPC 点的是文本中心而不是祖先对象。只有完整入口链自己的动作祖先占用该坐标时才可点击;
# SystemUI 浮层、额外按钮或任意部分覆盖矩形都可能截获触摸,必须零点击失败关闭。
return entries if len(occupants) == 1 and occupants[0].element is ancestor.element else []
def _physical_entries(nodes: list[_Node]) -> list[_Node]:
entry_labels = [
node for node in nodes
if node.text == _ENTRY
and node.element.get("package") == PDD_PACKAGE
and node.element.get("class") == "android.widget.TextView"
]
# 这里故意只证明物理结构。pre-intent 必须识别旧商品,不能让父容器描述或浮层等
# post-intent 安全条件把已经存在的旧入口伪装成“不存在”。
return [node for node in entry_labels if _physical_entry_chain(node, nodes) is not None]
def _physical_entry_chain(node: _Node, nodes: list[_Node]) -> tuple[_Node, ...] | None:
if node.text != _ENTRY or not _exact_entry_node(
node, "android.widget.TextView", _ENTRY_TEXT_BOUNDS, "false"
) or len(node.element) != 0:
return None
ancestor = node.parent
if ancestor is None or not _exact_entry_node(
ancestor, "android.view.ViewGroup", _ENTRY_ACTION_BOUNDS, "true"
):
return None
siblings = [
candidate for candidate in nodes
if candidate.parent is ancestor
and candidate.text == _ENTRY_SIBLING
and _exact_entry_node(
candidate, "android.widget.TextView", _ENTRY_SIBLING_BOUNDS, "false"
)
and not candidate.desc
and len(candidate.element) == 0
]
if len(siblings) != 1:
return None
return node, ancestor, siblings[0]
def _entry_projection(node: _Node, nodes: list[_Node]) -> tuple[object, ...] | None:
chain = _physical_entry_chain(node, nodes)
if chain is None:
return None
_, ancestor, sibling = chain
if node.desc:
return None
# 金额只作为这个已取证入口的不可变身份。这里既不解析也不返回它,价格闸门仍只能读取规格面板。
if ancestor.text or ancestor.desc != _ENTRY_ACTION_DESC:
return None
sibling_mentions = [
candidate for candidate in nodes
if _ENTRY_SIBLING in candidate.text or _ENTRY_SIBLING in candidate.desc
]
if len(sibling_mentions) != 1 or sibling_mentions[0].element is not sibling.element:
return None
subtree = [
candidate for candidate in nodes
if candidate.element is ancestor.element or _descendant(candidate, ancestor)
]
if any(
forbidden in value
for candidate in subtree
for value in (candidate.text, candidate.desc)
if candidate.element is not sibling.element
for forbidden in _FORBIDDEN_ENTRY_ACTION_DESC
):
return None
if any(
value == _ENTRY_PROMOTION_LABEL
for candidate in subtree
for value in (candidate.text, candidate.desc)
):
return None
entry_text_nodes = [candidate for candidate in subtree if candidate.text == _ENTRY]
if len(entry_text_nodes) != 1 or entry_text_nodes[0].element is not node.element:
return None
return tuple(
_entry_node_projection(candidate)
for candidate in (node, ancestor, sibling)
)
def _entry_node_projection(node: _Node) -> tuple[str, ...]:
return (
node.element.tag,
node.element.get("package", ""),
node.element.get("class", ""),
node.bounds,
node.element.get("clickable", ""),
node.element.get("enabled", ""),
node.element.get("visible-to-user", ""),
node.text,
node.desc,
)
def _is_live_clickable(node: _Node) -> bool:
return (
node.element.get("clickable") == "true"
and node.element.get("enabled") == "true"
and node.element.get("visible-to-user") == "true"
)
def _live_clickables_covering(nodes: list[_Node], point: tuple[int, int]) -> list[_Node]:
occupants: list[_Node] = []
x, y = point
for node in nodes:
if not _is_live_clickable(node):
continue
# 活跃可点击节点的 bounds 无法验证时,无法证明它不会截获入口坐标,因此整体失败关闭。
left, top, right, bottom = _action_bounds(node.bounds)
if left <= x < right and top <= y < bottom:
occupants.append(node)
return occupants
def _exact_entry_node(node: _Node, class_name: str, bounds: str, clickable: str) -> bool:
return (
node.element.get("package") == PDD_PACKAGE
and node.element.get("class") == class_name
and node.bounds == bounds
and node.element.get("clickable") == clickable
and node.element.get("enabled") == "true"
and node.element.get("visible-to-user") == "true"
)
def _action_bounds(bounds: str) -> tuple[int, int, int, int]:
match = _BOUNDS.fullmatch(bounds)
if match is None: raise SkuSelectionError("规格节点坐标格式无效,已停止操作。")
left, top, right, bottom = (int(item) for item in match.groups())
if not (0 <= left < right <= _W and 0 <= top < bottom <= _H):
raise SkuSelectionError("规格节点坐标不在已取证屏幕范围内,已停止操作。")
return left, top, right, bottom
def _one(nodes: list[_Node], message: str) -> _Node:
if len(nodes) != 1: raise SkuSelectionError(message)
return nodes[0]
@@ -0,0 +1,437 @@
"""T-103 真机运行边界:窄适配器、原始截图和无页面正文的摘要。"""
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
from datetime import UTC, datetime
from hashlib import sha256
import json
from math import isfinite
import os
from pathlib import Path
import shutil
from time import monotonic
from typing import Any
from uuid import uuid4
from PIL import Image, UnidentifiedImageError
from adbutils.errors import AdbTimeout
from uiautomator2.exceptions import HTTPTimeoutError
from ..device.adb import AdbClient, DeviceConnectionError, DeviceInspection
from ..device.baseline import PDD_PACKAGE, SCREENSHOT_PARAMS, _save_base64_screenshot, _sha256_file
from .product_open import EXPECTED_PDD_VERSION
from .product_url import ProductUrl, ProductUrlError, parse_product_url
from .sku_selection import (
EXPECTED_GOODS_ID,
EXPECTED_UNIT_PRICE,
SkuPanelDevice,
SkuSelectionError,
SkuSelectionFlow,
_SKU_ENTRY_FAILURE_STAGES,
_action_bounds,
_annotate_sku_entry_failure,
_safe_sku_entry_failure_stage,
resolve_task_selection,
)
EXPECTED_DEVICE_MODEL = "PKG110"
EXPECTED_ANDROID_VERSION = "16"
EXPECTED_SCREEN_SIZE = (1080, 2376)
# CLI 只允许输出这些固定阶段码。阶段码描述运行器自己的控制流,不包含页面
# 文本、节点属性、serial、路径或第三方异常;未知/伪造值统一降级为 unknown。
_FAILURE_STAGES = frozenset(
(
"precheck",
"device_inspection",
"device_session",
"product_open",
"sku_entry",
*_SKU_ENTRY_FAILURE_STAGES,
"sku_selection",
"price_verification",
"screenshot_capture",
"screenshot_reverify",
"safe_exit",
"publish",
)
)
class SkuSelectionRunError(RuntimeError):
"""T-103 运行未完整完成;错误文本不携带设备或页面原文。"""
class SkuSelectionRunTimeoutError(SkuSelectionRunError):
"""设备 RPC 或操作超时。"""
class SkuSelectionScreenshotError(SkuSelectionRunError):
"""原始截图无法作为完整 PNG 原子发布。"""
class SkuSelectionUnexpectedPriceError(SkuSelectionRunError):
"""取证面板现价不是本任务已确认值。"""
class SkuSelectionDeviceAdapterError(SkuSelectionRunError):
"""第三方设备接口失败的脱敏映射。"""
def safe_failure_stage(error: BaseException) -> str:
"""返回允许公开的固定阶段码,绝不回显异常正文。"""
try:
stage = getattr(error, "_cmbuyer_failure_stage", None)
# exact str 避免恶意 str 子类在 hash/eq 中执行任意异常;诊断路径
# 自己也必须失败闭合,不能让异常正文越过 CLI 的统一脱敏出口。
if type(stage) is not str or stage not in _FAILURE_STAGES:
return "unknown"
if stage in _SKU_ENTRY_FAILURE_STAGES:
return stage if _safe_sku_entry_failure_stage(error) == stage else "unknown"
return stage
except BaseException:
return "unknown"
def _annotate_failure(error: BaseException, stage: str) -> None:
"""只给本次异常附加白名单控制流事实;原异常文本仍不对外输出。"""
safe_stage = stage if stage in _FAILURE_STAGES else "unknown"
try:
setattr(error, "_cmbuyer_failure_stage", safe_stage)
except BaseException:
# 极端第三方异常不允许写属性时仍保持原失败闭合语义。
pass
def _failure_stage_for(error: BaseException, runner_stage: str) -> str:
if runner_stage == "sku_entry":
flow_stage = _safe_sku_entry_failure_stage(error)
if flow_stage is not None:
return flow_stage
return runner_stage
def _annotate_mapped_failure(mapped: BaseException, source: BaseException, runner_stage: str) -> None:
stage = _failure_stage_for(source, runner_stage)
if stage in _SKU_ENTRY_FAILURE_STAGES:
_annotate_sku_entry_failure(mapped, stage)
_annotate_failure(mapped, stage)
@dataclass(frozen=True)
class SkuSelectionRunResult:
"""已发布的截图和无页面正文 manifest 摘要。"""
output_directory: Path
screenshot_path: Path
manifest_path: Path
unit_price: str
class UiautomatorSkuPanelAdapter(SkuPanelDevice):
"""把 uiautomator2 缩为 T-103 所需的读取与三种命名操作。
``tap_sku_entry``、``tap_sku_option`` 和 ``leave_sku_panel`` 是仅有的状态改变方法;
坐标由 Flow 和本类双重检查后才计算中心点,每次调用只执行一次底层动作。
"""
def __init__(self, device: Any, timeout_seconds: float) -> None:
if not _is_positive_finite(timeout_seconds):
raise ValueError("timeout_seconds 必须是大于 0 的有限数值")
self._device = device
self._timeout_seconds = timeout_seconds
self._entry_was_tapped = False
self._left_panel = False
@property
def entry_was_tapped(self) -> bool:
"""仅供运行器决定故障后的单次尽力返回,不是页面操作。"""
return self._entry_was_tapped
@property
def left_panel(self) -> bool:
return self._left_panel
def app_info(self, package_name: str) -> dict[str, Any]:
value = self._call("app_info", package_name)
if not isinstance(value, dict):
raise SkuSelectionDeviceAdapterError("无法读取应用版本,已停止操作。")
return value
def app_current(self) -> dict[str, Any]:
value = self._call("app_current")
if not isinstance(value, dict):
raise SkuSelectionDeviceAdapterError("无法读取前台应用,已停止操作。")
return value
def dump_window_hierarchy(self) -> str:
value = self._call("jsonrpc_call", "dumpWindowHierarchy", [False, 50], timeout=self._timeout_seconds)
if not isinstance(value, str):
raise SkuSelectionDeviceAdapterError("节点树读取失败,已停止操作。")
return value
def tap_sku_entry(self, bounds: str) -> None:
# 超时也可能表示底层事件已经送达;必须先封存 attempt,后续绝不重试该入口。
self._entry_was_tapped = True
self._tap_bounds_once(bounds)
def tap_sku_option(self, bounds: str) -> None:
self._tap_bounds_once(bounds)
def leave_sku_panel(self) -> None:
if self._left_panel:
raise SkuSelectionDeviceAdapterError("规格面板已经执行过返回,已停止操作。")
# 底层调用即使报错也可能已把返回事件送达;先封存本次机会,finally 不得再次返回。
self._left_panel = True
self._call("jsonrpc_call", "pressKey", ["back"], timeout=self._timeout_seconds)
def capture_screenshot(self) -> str:
value = self._call("jsonrpc_call", "takeScreenshot", SCREENSHOT_PARAMS, timeout=self._timeout_seconds)
if not isinstance(value, str):
raise SkuSelectionScreenshotError("规格面板原始截图读取失败,未发布任何证据产物。")
return value
def display_size(self) -> tuple[int, int]:
value = self._call("window_size")
if not isinstance(value, tuple) or len(value) != 2 or any(not isinstance(item, int) for item in value):
raise SkuSelectionDeviceAdapterError("无法读取屏幕坐标空间,已停止操作。")
return value
def _tap_bounds_once(self, bounds: str) -> None:
left, top, right, bottom = _action_bounds(bounds)
center_x = left + (right - left) // 2
center_y = top + (bottom - top) // 2
self._call("jsonrpc_call", "click", [center_x, center_y], timeout=self._timeout_seconds)
def _call(self, method: str, *args: Any, **kwargs: Any) -> Any:
try:
operation = getattr(self._device, method)
return operation(*args, **kwargs)
except (AdbTimeout, HTTPTimeoutError, TimeoutError) as error:
raise SkuSelectionRunTimeoutError("规格面板设备操作超时,已停止操作。") from error
except SkuSelectionRunError:
raise
except Exception as error:
raise SkuSelectionDeviceAdapterError("规格面板设备操作失败,已停止操作。") from error
class SkuSelectionRunner:
"""只运行 T-103 目标规格恢复、价格确认、原始截图和一次安全退出。"""
def __init__(
self,
adb_client: AdbClient,
connector: Callable[[str], Any],
timeout_seconds: float,
monotonic_clock: Callable[[], float] = monotonic,
) -> None:
if not _is_positive_finite(timeout_seconds):
raise ValueError("timeout_seconds 必须是大于 0 的有限数值")
self._adb_client = adb_client
self._connector = connector
self._timeout_seconds = timeout_seconds
self._monotonic_clock = monotonic_clock
def run(
self,
serial: str,
product_url: str,
task_color: str,
task_size: str,
output_directory: Path,
) -> SkuSelectionRunResult:
stage = "precheck"
staging: Path | None = None
adapter: UiautomatorSkuPanelAdapter | None = None
flow: SkuSelectionFlow | None = None
try:
link = parse_product_url(product_url)
if link.goods_id != EXPECTED_GOODS_ID:
raise SkuSelectionRunError("商品不是 T-103 已取证目标,已停止操作。")
selection = resolve_task_selection(task_color, task_size)
target = Path(output_directory)
_validate_new_target(target)
staging = _prepare_staging(target)
deadline = self._monotonic_clock() + self._timeout_seconds
stage = "device_inspection"
inspection = self._adb_client.inspect(serial)
_require_expected_device(inspection)
stage = "device_session"
adapter = UiautomatorSkuPanelAdapter(self._connector(serial), self._timeout_seconds)
_require_expected_version(adapter.app_info(PDD_PACKAGE))
if adapter.display_size() != EXPECTED_SCREEN_SIZE:
raise SkuSelectionRunError("设备不是已取证的竖屏坐标空间,已停止操作。")
pre_intent_hierarchy = adapter.dump_window_hierarchy()
# 固定 ACTION_VIEW、固定 PDD package 和 canonical goods_id;不接受任意 URL 或 shell。
stage = "product_open"
self._adb_client.start_pdd_view_intent(serial, link.goods_id)
remaining = deadline - self._monotonic_clock()
if remaining <= 0:
raise SkuSelectionRunTimeoutError("等待规格入口超时,未执行点击。")
stage = "sku_entry"
flow = SkuSelectionFlow(adapter, entry_wait_timeout_seconds=remaining)
flow.open_sku_panel(link.canonical_url, pre_intent_hierarchy)
stage = "sku_selection"
flow.select_sku_options(selection)
stage = "price_verification"
unit_price = flow.verify_target_selection_and_read_price(selection)
if unit_price != EXPECTED_UNIT_PRICE:
raise SkuSelectionUnexpectedPriceError("规格面板现价不是本任务已确认值,已停止操作。")
stage = "screenshot_capture"
screenshot_path = staging / "screenshot.png"
try:
_save_base64_screenshot(adapter.capture_screenshot(), screenshot_path)
_require_screenshot_size(screenshot_path)
except SkuSelectionRunError:
raise
except Exception as error:
raise SkuSelectionScreenshotError("规格面板原始截图保存失败,未发布任何证据产物。") from error
manifest_path = staging / "manifest.json"
# 截图可能落在动态页面切换边界;发布前必须用一棵更新节点树同时重证两维和现价。
stage = "screenshot_reverify"
final_price = flow.verify_target_selection_and_read_price(selection)
if final_price != EXPECTED_UNIT_PRICE:
raise SkuSelectionUnexpectedPriceError("截图后规格面板现价不是本任务已确认值,已停止操作。")
# 正常路径仍经 Flow 做最后一次前台和面板判定;返回操作只发生一次。
stage = "safe_exit"
flow.exit_sku_panel_safely()
stage = "publish"
manifest_path.write_text(
json.dumps(_manifest(inspection, serial, link, screenshot_path, task_color, task_size), ensure_ascii=False, indent=2, sort_keys=True) + "\n",
encoding="utf-8",
)
# Windows 的 rename 不替换既有目标;并发创建 target 时保留其内容并把本次运行判失败。
os.rename(staging, target)
staging = None
except (DeviceConnectionError, ProductUrlError, SkuSelectionRunError, SkuSelectionError) as error:
_clean_staging(staging)
_annotate_failure(error, _failure_stage_for(error, stage))
raise
except (AdbTimeout, HTTPTimeoutError, TimeoutError) as error:
_clean_staging(staging)
mapped = SkuSelectionRunTimeoutError("规格面板运行超时,未发布任何证据产物。")
_annotate_mapped_failure(mapped, error, stage)
raise mapped from error
except OSError as error:
_clean_staging(staging)
mapped = SkuSelectionRunError("规格面板证据目录无法创建或发布,未发布任何证据产物。")
_annotate_mapped_failure(mapped, error, stage)
raise mapped from error
except Exception as error:
_clean_staging(staging)
mapped = SkuSelectionRunError("规格面板运行未完成,未发布任何证据产物。")
_annotate_mapped_failure(mapped, error, stage)
raise mapped from error
finally:
# 失败路径只能复用 Flow 的版本、前台和面板证明;证明不了便停止,绝不盲目返回。
if flow is not None and adapter is not None and adapter.entry_was_tapped and not adapter.left_panel:
try:
flow.reconcile_pending_action()
flow.exit_sku_panel_safely()
except (SkuSelectionRunError, SkuSelectionError):
pass
return SkuSelectionRunResult(
output_directory=target,
screenshot_path=target / "screenshot.png",
manifest_path=target / "manifest.json",
unit_price=EXPECTED_UNIT_PRICE,
)
def _is_positive_finite(value: object) -> bool:
return isinstance(value, (int, float)) and not isinstance(value, bool) and value > 0 and isfinite(value)
def _validate_new_target(target: Path) -> None:
if target.exists():
raise SkuSelectionRunError("输出目录已存在;为防止覆盖旧证据,已停止操作。")
if not target.name:
raise SkuSelectionRunError("输出目录必须是明确的新目录。")
def _prepare_staging(target: Path) -> Path:
staging: Path | None = None
try:
target.parent.mkdir(parents=True, exist_ok=True)
staging = target.parent / f".{target.name}.staging-{uuid4().hex}"
staging.mkdir()
probe = staging / ".write-probe"
probe.write_bytes(b"ok")
probe.unlink()
return staging
except OSError as error:
_clean_staging(staging)
raise SkuSelectionRunError("输出目录不可写,已停止操作。") from error
def _clean_staging(staging: Path | None) -> None:
if staging is not None and staging.exists():
shutil.rmtree(staging)
def _require_expected_version(app_info: object) -> str:
version = (app_info.get("versionName") or app_info.get("version_name")) if isinstance(app_info, dict) else None
if version != EXPECTED_PDD_VERSION:
raise SkuSelectionRunError("拼多多版本与已取证版本不一致,已停止操作。")
return version
def _require_expected_device(inspection: DeviceInspection) -> None:
if inspection.model != EXPECTED_DEVICE_MODEL or inspection.android_version != EXPECTED_ANDROID_VERSION:
raise SkuSelectionRunError("设备型号或 Android 版本不是已取证组合,已停止操作。")
def _require_screenshot_size(screenshot_path: Path) -> None:
try:
with Image.open(screenshot_path) as image:
image.load()
if image.size != EXPECTED_SCREEN_SIZE:
raise SkuSelectionScreenshotError("原始截图坐标空间不是已取证尺寸,未发布任何证据产物。")
except SkuSelectionRunError:
raise
except (UnidentifiedImageError, OSError) as error:
raise SkuSelectionScreenshotError("原始截图无效,未发布任何证据产物。") from error
def _manifest(inspection: DeviceInspection, serial: str, link: ProductUrl, screenshot_path: Path, task_color: str, task_size: str) -> dict[str, Any]:
"""仅写可审计摘要;原始 serial、节点树、页面文案和实际截图内容均不写入 manifest。"""
return {
"schema_version": 1,
"captured_at": datetime.now(UTC).isoformat(),
"operation": "t103-sku-selection",
"product": {"goods_id": link.goods_id, "canonical_url": link.canonical_url},
"target_selection": {"color": task_color, "size": task_size},
"unit_price": EXPECTED_UNIT_PRICE,
"selection_status": "restored",
"panel_status": "verified",
"safe_exit": "completed",
"page_identity": "human_review_required",
"channel": "wifi" if ":" in serial else "usb",
"serial_sha256": sha256(serial.encode("utf-8")).hexdigest(),
"device": {
"model": inspection.model,
"android_version": inspection.android_version,
"pdd_package": PDD_PACKAGE,
"pdd_version": EXPECTED_PDD_VERSION,
},
"artifacts": [{"path": screenshot_path.name, "sha256": _sha256_file(screenshot_path)}],
}
@@ -0,0 +1,5 @@
"""采购工具轮询会话协调器。"""
from .coordinator import ClaimedTaskView, PollingCoordinator, PollingState, RecoveryStatus, StartReadiness
__all__ = ["ClaimedTaskView", "PollingCoordinator", "PollingState", "RecoveryStatus", "StartReadiness"]
@@ -0,0 +1,589 @@
"""在 Qt 事件循环中协调可恢复的领取会话。"""
from __future__ import annotations
from collections.abc import Callable
from dataclasses import dataclass
from enum import Enum
from typing import Protocol
from PySide6.QtCore import QObject, QThread, QTimer, Signal, Slot
from cmbuyer_client.core.errors import (
AmbiguousRemoteError,
ClientError,
CredentialRemoteError,
ManualRemoteError,
ProtocolRemoteError,
StateError,
)
from cmbuyer_client.core.models import ClaimedTask
from cmbuyer_client.localstate.models import PollingSession, ProfileSettings, RecoverySnapshot
from cmbuyer_client.logging_policy import redact_text
SAFE_AMBIGUOUS_REASONS = frozenset(
("http_result_unknown", "server_result_unknown", "truncated_response")
)
class PollingState(str, Enum):
STOPPED = "STOPPED"
STARTING = "STARTING"
BLOCKED = "BLOCKED"
RECOVERING = "RECOVERING"
WAITING = "WAITING"
CLAIMING = "CLAIMING"
ACTIVE = "ACTIVE"
RECOVERY_REQUIRED = "RECOVERY_REQUIRED"
class PollingStore(Protocol):
def recovery_snapshot(self, profile_id: str) -> RecoverySnapshot: ...
def start_or_resume_polling(self, profile_id: str) -> PollingSession: ...
def request_stop(self, profile_id: str) -> PollingSession: ...
class ClaimGateway(Protocol):
def claim_next(self, profile_id: str) -> ClaimedTask | None: ...
class ExecutionConsumer(Protocol):
def accept_claim(self, claimed: ClaimedTask, profile: ProfileSettings) -> None: ...
@dataclass(frozen=True)
class StartReadiness:
"""由后续已取证执行能力注入;T-304 自己不探测网络或设备。"""
ready: bool
reason: str
@dataclass(frozen=True)
class ClaimedTaskView:
"""允许发往 UI 的最小投影,刻意不包含 authorization/claim token。"""
task_id: str
title: str
status: str = "已领取"
@classmethod
def from_claim(cls, claimed: ClaimedTask) -> "ClaimedTaskView":
return cls(task_id=claimed.task.id, title=redact_text(claimed.task.title))
@dataclass(frozen=True)
class RecoveryStatus:
"""可进入 UI 的恢复摘要;不携带 task snapshot、claim token 或密文。"""
has_open_session: bool
session_accept_new: bool
has_pending_claim: bool
has_active_claim: bool
has_pending_active_work: bool
@classmethod
def from_snapshot(cls, snapshot: RecoverySnapshot) -> "RecoveryStatus":
return cls(
has_open_session=snapshot.session is not None,
session_accept_new=bool(snapshot.session and snapshot.session.accept_new),
has_pending_claim=snapshot.pending_claim is not None,
has_active_claim=snapshot.active_claim is not None,
has_pending_active_work=bool(snapshot.pending_renew or snapshot.pending_evidence),
)
@dataclass(frozen=True)
class _BootstrapResult:
recovery: RecoveryStatus
session: PollingSession | None
normalized_stop: PollingSession | None = None
class _PollingWorker(QObject):
bootstrap_finished = Signal(object)
claim_finished = Signal(object)
stop_finished = Signal(object)
failed = Signal(str, object)
def __init__(self, store: PollingStore) -> None:
super().__init__()
self._store = store
self._gateway: ClaimGateway | None = None
@Slot(object)
def configure_gateway(self, gateway: object) -> None:
if not hasattr(gateway, "claim_next"):
self.failed.emit("configure", RuntimeError("invalid_claim_gateway"))
return
self._gateway = gateway # type: ignore[assignment]
@Slot(str)
def bootstrap(self, profile_id: str) -> None:
try:
snapshot = self._store.recovery_snapshot(profile_id)
recovery = RecoveryStatus.from_snapshot(snapshot)
if snapshot.session is not None and snapshot.session.accept_new:
stopped = self._store.request_stop(profile_id)
self.bootstrap_finished.emit(_BootstrapResult(recovery, None, stopped))
return
if snapshot.active_claim is not None:
self.bootstrap_finished.emit(_BootstrapResult(recovery, None))
return
session = self._store.start_or_resume_polling(profile_id)
self.bootstrap_finished.emit(_BootstrapResult(recovery, session))
except Exception as error:
self.failed.emit("bootstrap", error)
@Slot(str)
def inspect_restart(self, profile_id: str) -> None:
try:
snapshot = self._store.recovery_snapshot(profile_id)
recovery = RecoveryStatus.from_snapshot(snapshot)
stopped = None
if snapshot.session is not None and snapshot.session.accept_new:
stopped = self._store.request_stop(profile_id)
self.bootstrap_finished.emit(_BootstrapResult(recovery, None, stopped))
except Exception as error:
self.failed.emit("inspect", error)
@Slot(str)
def claim(self, profile_id: str) -> None:
try:
# DurableClientGateway 在返回前已经提交 EMPTY 或 active claim;UI 不能
# 以 generation 过期为由丢弃这个业务结果。
if self._gateway is None:
raise RuntimeError("claim_gateway_not_configured")
self.claim_finished.emit(self._gateway.claim_next(profile_id))
except Exception as error:
self.failed.emit("claim", error)
@Slot(str)
def stop(self, profile_id: str) -> None:
try:
self.stop_finished.emit(self._store.request_stop(profile_id))
except Exception as error:
self.failed.emit("stop", error)
class PollingCoordinator(QObject):
"""把计时、阻塞 I/O 和可见状态收敛到一个会话边界。"""
state_changed = Signal(object, str, int)
claim_visible = Signal(object)
recovery_status_changed = Signal(object)
configuration_freeze_changed = Signal(bool)
settled = Signal()
_configure_gateway_requested = Signal(object)
_bootstrap_requested = Signal(str)
_inspect_requested = Signal(str)
_claim_requested = Signal(str)
_stop_requested_signal = Signal(str)
def __init__(
self,
*,
profile_id: str,
store: PollingStore | None,
gateway_factory: Callable[[ProfileSettings], ClaimGateway] | None,
consumer: ExecutionConsumer | None,
profile_settings: ProfileSettings | None = None,
readiness: StartReadiness | None = None,
poll_interval_seconds: int = 15,
failure_threshold: int = 3,
timer_interval_ms: int | None = None,
parent: QObject | None = None,
) -> None:
super().__init__(parent)
if not 5 <= poll_interval_seconds <= 300:
raise ValueError("invalid_poll_interval")
if not 1 <= failure_threshold <= 10:
raise ValueError("invalid_failure_threshold")
self.profile_id = profile_id
self._store = store
self._gateway_factory = gateway_factory
self._gateway: ClaimGateway | None = None
self._consumer = consumer
self._profile_settings = profile_settings
self._frozen_profile: ProfileSettings | None = None
self._readiness = readiness
self.recovery_status: RecoveryStatus | None = None
self._timer_interval_override = timer_interval_ms
self._interval_ms = timer_interval_ms or (
profile_settings.poll_interval_seconds * 1000 if profile_settings is not None else poll_interval_seconds * 1000
)
if self._interval_ms <= 0:
raise ValueError("invalid_timer_interval")
self._failure_threshold = (
profile_settings.failure_threshold if profile_settings is not None else failure_threshold
)
self._consecutive_failures = 0
self._operation: str | None = None
self._stop_requested = False
self._epoch = 0
self._scheduled_epoch: int | None = None
self._post_stop_state = PollingState.STOPPED
self._post_stop_reason = "轮询已停止。"
self._thread: QThread | None = None
self._worker: _PollingWorker | None = None
self._timer = QTimer(self)
self._timer.setSingleShot(True)
self._timer.timeout.connect(self._on_timer_timeout)
if consumer is None:
self.state = PollingState.BLOCKED
self.reason = "单趟执行能力尚未接入,不能领取真实任务。"
elif store is None or gateway_factory is None:
self.state = PollingState.BLOCKED
self.reason = "轮询依赖未完整注入,不能领取真实任务。"
elif profile_settings is None:
self.state = PollingState.BLOCKED
self.reason = "尚未保存完整配置,不能开始轮询。"
elif readiness is None or not readiness.ready:
self.state = PollingState.BLOCKED
self.reason = "执行就绪条件未满足,不能开始轮询。" if readiness is None else readiness.reason
else:
self.state = PollingState.STOPPED
self.reason = "轮询已停止。"
if store is not None:
self._thread = QThread(self)
self._worker = _PollingWorker(store)
self._worker.moveToThread(self._thread)
self._configure_gateway_requested.connect(self._worker.configure_gateway)
self._inspect_requested.connect(self._worker.inspect_restart)
self._bootstrap_requested.connect(self._worker.bootstrap)
self._claim_requested.connect(self._worker.claim)
self._stop_requested_signal.connect(self._worker.stop)
self._worker.bootstrap_finished.connect(self._on_bootstrap_finished)
self._worker.claim_finished.connect(self._on_claim_finished)
self._worker.stop_finished.connect(self._on_stop_finished)
self._worker.failed.connect(self._on_worker_failed)
self._thread.start()
self._operation = "inspect"
self._set_state(PollingState.STARTING, "正在读取重启恢复状态并关闭遗留自动领取许可…")
self._inspect_requested.emit(self.profile_id)
@property
def can_start(self) -> bool:
return (
self._consumer is not None
and self._store is not None
and self._gateway_factory is not None
and self._profile_settings is not None
and self._readiness is not None
and self._readiness.ready
and self.state == PollingState.STOPPED
and self._operation is None
)
@property
def consecutive_failures(self) -> int:
return self._consecutive_failures
@property
def operation_in_flight(self) -> bool:
return self._operation is not None
def update_profile_settings(self, settings: ProfileSettings) -> None:
if self._operation is None and self.state in (PollingState.STOPPED, PollingState.BLOCKED):
self._profile_settings = settings
self._refresh_idle_gate()
def update_readiness(self, readiness: StartReadiness) -> None:
self._readiness = readiness
self._refresh_idle_gate()
def _refresh_idle_gate(self) -> None:
if self._operation is not None or self.state == PollingState.RECOVERY_REQUIRED:
return
if self._consumer is None:
self._set_state(PollingState.BLOCKED, "单趟执行能力尚未接入,不能领取真实任务。")
elif self._gateway_factory is None or self._store is None:
self._set_state(PollingState.BLOCKED, "轮询依赖未完整注入,不能领取真实任务。")
elif self._profile_settings is None:
self._set_state(PollingState.BLOCKED, "尚未保存完整配置,不能开始轮询。")
elif self._readiness is None or not self._readiness.ready:
reason = "执行就绪条件未满足,不能开始轮询。" if self._readiness is None else self._readiness.reason
self._set_state(PollingState.BLOCKED, reason)
else:
self._set_state(PollingState.STOPPED, "轮询已停止。")
def start(self) -> None:
# 这道门禁必须早于任何 store/gateway 调用;独立应用没有 consumer,
# 即使调用方绕过禁用按钮直接调用本方法也保持零 HTTP。
if self._consumer is None:
self._set_state(PollingState.BLOCKED, "单趟执行能力尚未接入,不能领取真实任务。")
return
if self._store is None or self._gateway_factory is None or self._worker is None:
self._set_state(PollingState.BLOCKED, "轮询依赖未完整注入,不能领取真实任务。")
return
if self._profile_settings is None:
self._set_state(PollingState.BLOCKED, "尚未保存完整配置,不能开始轮询。")
return
if self._readiness is None or not self._readiness.ready:
reason = "执行就绪条件未满足,不能开始轮询。" if self._readiness is None else self._readiness.reason
self._set_state(PollingState.BLOCKED, reason)
return
if not self.can_start:
return
frozen = self._profile_settings
try:
gateway = self._gateway_factory(frozen)
except Exception:
self._set_state(PollingState.BLOCKED, "领取网关无法按本次冻结配置建立,不能开始轮询。")
return
if gateway is None or not hasattr(gateway, "claim_next"):
self._set_state(PollingState.BLOCKED, "领取网关未完整注入,不能开始轮询。")
return
self._frozen_profile = frozen
self._interval_ms = self._timer_interval_override or frozen.poll_interval_seconds * 1000
self._failure_threshold = frozen.failure_threshold
self._gateway = gateway
self._configure_gateway_requested.emit(gateway)
self._epoch += 1
self._scheduled_epoch = None
self._timer.stop()
self._stop_requested = False
self._post_stop_state = PollingState.STOPPED
self._consecutive_failures = 0
self._operation = "bootstrap"
self._set_state(PollingState.STARTING, "正在读取本地恢复状态并建立轮询会话…")
self._bootstrap_requested.emit(self.profile_id)
def stop(self) -> None:
if self.state == PollingState.STOPPED and self._operation is None:
return
if self._store is None or self._worker is None:
return
# epoch/latch 双保险:stopEvent 先让 timer 队列中已经排队的 timeout
# 失效,再处理持久 stop;之后只有显式 Start 才会获得新 epoch。
self._epoch += 1
self._scheduled_epoch = None
self._stop_requested = True
self._timer.stop()
if self._operation in ("inspect", "bootstrap", "claim", "stop"):
if self._operation == "claim":
self._set_state(PollingState.CLAIMING, "正在等待本次有界领取返回;不会取消或重发请求。")
return
self._request_stop(PollingState.STOPPED, "轮询已停止;只阻止下一次领取。")
def shutdown(self, wait_ms: int = 5000) -> bool:
"""只结束空闲 worker;飞行中 I/O 必须由事件循环等待 settled。"""
self._timer.stop()
if self._operation is not None:
return False
if self._thread is not None and self._thread.isRunning():
self._thread.quit()
return self._thread.wait(wait_ms)
return True
@Slot(object)
def _on_bootstrap_finished(self, raw: object) -> None:
completed_operation = self._operation
self._operation = None
result = raw
if not isinstance(result, _BootstrapResult):
self._block("本地恢复结果无效,轮询已阻止。")
return
self.recovery_status = result.recovery
self.recovery_status_changed.emit(result.recovery)
self.configuration_freeze_changed.emit(
result.recovery.has_pending_claim or result.recovery.has_active_claim
)
if completed_operation == "inspect":
if result.recovery.has_active_claim:
self._set_state(PollingState.RECOVERY_REQUIRED, "遗留会话已停止;必须先安全恢复当前任务。")
elif result.recovery.has_pending_claim:
self._set_state(PollingState.STOPPED, "遗留领取请求已停止;显式开始后只使用原幂等键恢复。")
else:
self._refresh_idle_gate()
self.settled.emit()
return
if result.normalized_stop is not None:
if result.recovery.has_active_claim:
self._set_state(PollingState.RECOVERY_REQUIRED, "遗留会话已停止;必须先安全恢复当前任务。")
else:
self._set_state(PollingState.STOPPED, "遗留会话已停止;请再次显式开始轮询。")
self.settled.emit()
return
if result.recovery.has_active_claim:
self._set_state(
PollingState.RECOVERY_REQUIRED,
"检测到未关闭的采购任务,必须先完成安全恢复,不能领取新任务。",
)
self.settled.emit()
return
if self._stop_requested:
self._request_stop(PollingState.STOPPED, "轮询已停止;未发起领取请求。")
return
if result.recovery.has_pending_claim:
self._set_state(PollingState.RECOVERING, "正在使用原幂等键恢复结果不明的领取请求…")
else:
self._set_state(PollingState.WAITING, "轮询会话已启动,正在等待领取。")
self._schedule_claim(0)
def _schedule_claim(self, delay_ms: int) -> None:
self._scheduled_epoch = self._epoch
self._timer.start(delay_ms)
@Slot()
def _on_timer_timeout(self) -> None:
scheduled_epoch = self._scheduled_epoch
self._scheduled_epoch = None
if scheduled_epoch != self._epoch:
return
self._begin_claim(scheduled_epoch)
def _begin_claim(self, dispatch_epoch: int) -> None:
if dispatch_epoch != self._epoch:
return
if self._operation is not None or self._stop_requested:
return
if self.state not in (PollingState.WAITING, PollingState.RECOVERING):
return
self._operation = "claim"
self.configuration_freeze_changed.emit(True)
self._set_state(PollingState.CLAIMING, "正在领取已授权任务…")
self._claim_requested.emit(self.profile_id)
@Slot(object)
def _on_claim_finished(self, claimed: object) -> None:
self._operation = None
self._consecutive_failures = 0
if claimed is not None and not isinstance(claimed, ClaimedTask):
self._request_stop(PollingState.BLOCKED, "领取结果类型无效,轮询已阻止。")
return
if isinstance(claimed, ClaimedTask):
self.claim_visible.emit(ClaimedTaskView.from_claim(claimed))
if self._stop_requested:
self._request_stop(
PollingState.RECOVERY_REQUIRED,
"停止期间领取已落库;必须先安全恢复该任务,不能领取下一条。",
)
return
try:
consumer = self._consumer
frozen_profile = self._frozen_profile
if consumer is None or frozen_profile is None:
raise RuntimeError("execution_consumer_missing")
# consumer 只能使用本次显式 Start 冻结的不可变配置;不得在
# 已领取后回读可变 UI/store,否则 ADB 身份和超时会发生趟内漂移。
consumer.accept_claim(claimed, frozen_profile)
except Exception:
self._request_stop(
PollingState.RECOVERY_REQUIRED,
"执行 consumer 未接收已落库任务;必须安全恢复,不能重新领取。",
)
return
self._set_state(PollingState.ACTIVE, "任务已安全领取并交给单趟执行能力。")
return
self.configuration_freeze_changed.emit(False)
if self._stop_requested:
self._request_stop(PollingState.STOPPED, "轮询已停止;本次没有可领取任务。")
return
self._set_state(PollingState.WAITING, "暂无已授权任务,等待下一次轮询。")
self._schedule_claim(self._interval_ms)
def _request_stop(self, target: PollingState, reason: str) -> None:
if self._operation == "stop":
return
self._timer.stop()
self._post_stop_state = target
self._post_stop_reason = reason
self._operation = "stop"
self._stop_requested_signal.emit(self.profile_id)
@Slot(object)
def _on_stop_finished(self, session: object) -> None:
self._operation = None
if not isinstance(session, PollingSession) or session.accept_new:
self._block("停止状态未能持久化,轮询已阻止。")
return
self._set_state(self._post_stop_state, self._post_stop_reason)
self.settled.emit()
@Slot(str, object)
def _on_worker_failed(self, operation: str, error: object) -> None:
self._operation = None
if operation == "inspect":
if isinstance(error, StateError) and error.reason == "profile_not_found":
self.recovery_status = RecoveryStatus(False, False, False, False, False)
self.recovery_status_changed.emit(self.recovery_status)
self._refresh_idle_gate()
else:
self._block("本地恢复状态无法安全读取;已停止且不能领取任务。")
self.settled.emit()
return
if (
operation == "claim"
and isinstance(error, AmbiguousRemoteError)
and error.reason in SAFE_AMBIGUOUS_REASONS
):
self._consecutive_failures += 1
if self._stop_requested:
self._request_stop(PollingState.STOPPED, "轮询已停止;结果不明的原领取请求已保留。")
elif self._consecutive_failures >= self._failure_threshold:
self._request_stop(
PollingState.BLOCKED,
"连续领取失败达到阈值;原幂等请求已保留,需排查后重新开始。",
)
else:
self._set_state(
PollingState.RECOVERING,
"领取结果不明;等待使用相同幂等键恢复,不会创建新请求。",
)
self._schedule_claim(self._interval_ms)
return
if operation == "claim" and isinstance(error, AmbiguousRemoteError):
self.configuration_freeze_changed.emit(True)
self._request_stop(
PollingState.BLOCKED,
"领取响应无法证明可安全定时恢复;原槽已保留,需显式开始后同键恢复。",
)
return
if operation == "stop":
self._block("停止状态无法安全落库,轮询已阻止;未清除任何恢复事实。")
return
if isinstance(error, CredentialRemoteError):
reason = "设备凭据无效或已撤销;修复凭据后再手工开始。"
elif isinstance(error, ProtocolRemoteError):
reason = "服务响应与固定协议不兼容;已停止普通重试。"
elif isinstance(error, ManualRemoteError):
reason = "服务端要求人工处理;已停止普通重试。"
elif isinstance(error, ClientError):
reason = "本地安全状态无法推进;已停止普通重试。"
else:
reason = "轮询发生未分类错误;已失败闭合。"
if operation == "claim":
# DurableClientGateway 已把协议错误和 409 人工冲突标成 terminal,
# 二者没有 pending/active;凭据或本地错误则可能保留 pending,继续冻结。
self.configuration_freeze_changed.emit(
not isinstance(error, (ProtocolRemoteError, ManualRemoteError))
)
if operation in ("bootstrap", "configure"):
self._block(reason)
else:
self._request_stop(PollingState.BLOCKED, reason)
def _block(self, reason: str) -> None:
self._timer.stop()
self._scheduled_epoch = None
self._set_state(PollingState.BLOCKED, reason)
if self._operation is None:
self.settled.emit()
def _set_state(self, state: PollingState, reason: str) -> None:
self.state = state
self.reason = reason
self.state_changed.emit(state, reason, self._consecutive_failures)
@@ -0,0 +1,7 @@
"""只连接固定本机采购服务的 HTTP 适配器。"""
from .evidence_sink import HttpEvidenceSink
from .http_transport import HttpTransport, LOOPBACK_SERVICE_URL
from .task_source import HttpTaskSource
__all__ = ["HttpEvidenceSink", "HttpTaskSource", "HttpTransport", "LOOPBACK_SERVICE_URL"]
@@ -0,0 +1,84 @@
"""仅上传调用方显式提供的单个 PNG 的窄 EvidenceSink。"""
from __future__ import annotations
from cmbuyer_client.core.errors import AmbiguousRemoteError, ProtocolRemoteError, ValidationError
from cmbuyer_client.core.models import AssetReceipt, DeviceCredentials, EvidenceUpload
from cmbuyer_client.core.validation import rfc3339_z_nanoseconds
from .http_transport import HttpTransport
from .wire import SMALL_RESPONSE_LIMIT, classify_bodyless_error, common_headers, parse_json_response
class HttpEvidenceSink:
def __init__(self, transport: HttpTransport) -> None:
self._transport = transport
def upload(self, credentials: DeviceCredentials, evidence: EvidenceUpload) -> AssetReceipt:
boundary = "cmbuyer-" + evidence.upload_key.replace("-", "")
marker = ("--" + boundary).encode("ascii")
if marker in evidence.content:
raise ProtocolRemoteError("multipart_boundary_collision")
body = _multipart_body(boundary, evidence)
response = self._transport.request(
"POST",
f"/api/v1/tasks/{evidence.task_id}/evidence",
common_headers(
credentials.device_id,
credentials.token.value,
"multipart/form-data; boundary=" + boundary,
),
body,
response_limit=SMALL_RESPONSE_LIMIT,
)
if response.status not in (200, 201):
classify_bodyless_error(response)
try:
receipt = AssetReceipt.from_wire(parse_json_response(response, maximum=SMALL_RESPONSE_LIMIT))
except ValidationError as error:
raise AmbiguousRemoteError("invalid_evidence_success_response") from error
if (
receipt.task_id != evidence.task_id
or receipt.attempt_id != evidence.attempt_id
or receipt.kind != evidence.kind
or receipt.privacy_tier != evidence.privacy_tier
or receipt.sha256 != evidence.sha256
or receipt.byte_size != len(evidence.content)
or receipt.width_px != evidence.width_px
or receipt.height_px != evidence.height_px
or rfc3339_z_nanoseconds(receipt.captured_at) != rfc3339_z_nanoseconds(evidence.captured_at)
):
raise AmbiguousRemoteError("evidence_response_mismatch")
return receipt
def _multipart_body(boundary: str, evidence: EvidenceUpload) -> bytes:
chunks: list[bytes] = []
def add_field(name: str, value: str) -> None:
chunks.extend(
(
f"--{boundary}\r\n".encode("ascii"),
f'Content-Disposition: form-data; name="{name}"\r\n\r\n'.encode("ascii"),
value.encode("utf-8"),
b"\r\n",
)
)
add_field("upload_key", evidence.upload_key)
add_field("attempt_id", evidence.attempt_id)
add_field("kind", evidence.kind)
add_field("privacy_tier", evidence.privacy_tier)
add_field("sha256", evidence.sha256)
add_field("captured_at", evidence.captured_at)
chunks.extend(
(
f"--{boundary}\r\n".encode("ascii"),
b'Content-Disposition: form-data; name="file"; filename="evidence.png"\r\n',
b"Content-Type: image/png\r\n\r\n",
evidence.content,
b"\r\n",
f"--{boundary}--\r\n".encode("ascii"),
)
)
return b"".join(chunks)
@@ -0,0 +1,116 @@
"""无代理、无重定向、无隐藏重试的 localhost HTTP transport。"""
from __future__ import annotations
from dataclasses import dataclass
import http.client
import re
from typing import Callable, Iterable
from cmbuyer_client.core.errors import AmbiguousRemoteError, ProtocolRemoteError
LOOPBACK_SERVICE_URL = "http://127.0.0.1:8080"
_HOST = "127.0.0.1"
_PORT = 8080
@dataclass(frozen=True)
class HttpResponse:
status: int
headers: tuple[tuple[str, str], ...]
body: bytes
def header_values(self, name: str) -> tuple[str, ...]:
wanted = name.lower()
return tuple(value for key, value in self.headers if key.lower() == wanted)
class HttpTransport:
"""每次调用只创建一个直连 TCP 请求;重试只能由持久化恢复层决定。"""
def __init__(
self,
service_url: str = LOOPBACK_SERVICE_URL,
*,
timeout_seconds: int = 10,
connection_factory: Callable[..., http.client.HTTPConnection] = http.client.HTTPConnection,
) -> None:
if service_url != LOOPBACK_SERVICE_URL:
raise ProtocolRemoteError("service_url_not_allowed")
if type(timeout_seconds) is not int or not 1 <= timeout_seconds <= 120:
raise ProtocolRemoteError("invalid_http_timeout")
self._timeout_seconds = timeout_seconds
self._connection_factory = connection_factory
def request(
self,
method: str,
path: str,
headers: Iterable[tuple[str, str]],
body: bytes,
*,
response_limit: int,
) -> HttpResponse:
if method != "POST" or not path.startswith("/api/v1/") or "?" in path or "#" in path:
raise ProtocolRemoteError("invalid_http_target")
if not isinstance(body, bytes) or type(response_limit) is not int or response_limit <= 0:
raise ProtocolRemoteError("invalid_http_request")
header_items = tuple(headers)
normalized: dict[str, str] = {}
for key, value in header_items:
lowered = key.lower()
if lowered in normalized or "\r" in key or "\n" in key or "\r" in value or "\n" in value:
raise ProtocolRemoteError("invalid_http_headers")
normalized[lowered] = value
connection: http.client.HTTPConnection | None = None
result: HttpResponse | None = None
failure: str | None = None
try:
connection = self._connection_factory(_HOST, _PORT, timeout=self._timeout_seconds)
connection.request(method, path, body=body, headers={key: value for key, value in header_items})
response = connection.getresponse()
response_headers = tuple(response.getheaders())
content_lengths = tuple(value for key, value in response_headers if key.lower() == "content-length")
transfer_encodings = tuple(value for key, value in response_headers if key.lower() == "transfer-encoding")
if len(content_lengths) > 1:
raise AmbiguousRemoteError("invalid_content_length")
if content_lengths and transfer_encodings:
raise AmbiguousRemoteError("ambiguous_response_framing")
if len(transfer_encodings) > 1 or (
transfer_encodings and transfer_encodings[0].lower() != "chunked"
):
raise AmbiguousRemoteError("invalid_transfer_encoding")
declared = content_lengths[0] if content_lengths else None
declared_length: int | None = None
if declared is not None:
if re.fullmatch(r"[0-9]+", declared, flags=re.ASCII) is None:
raise AmbiguousRemoteError("invalid_content_length")
if len(declared) > 10:
raise AmbiguousRemoteError("response_too_large")
declared_length = int(declared)
if declared_length > response_limit:
raise AmbiguousRemoteError("response_too_large")
response_body = response.read(response_limit + 1)
if len(response_body) > response_limit:
raise AmbiguousRemoteError("response_too_large")
if declared_length is not None and len(response_body) != declared_length:
raise AmbiguousRemoteError("truncated_response")
result = HttpResponse(response.status, response_headers, response_body)
except AmbiguousRemoteError as error:
failure = error.reason
except (OSError, TimeoutError, http.client.HTTPException):
failure = "http_result_unknown"
finally:
if connection is not None:
try:
connection.close()
except OSError:
if result is None:
failure = "http_result_unknown"
if failure is not None:
raise AmbiguousRemoteError(failure)
if result is None:
raise AmbiguousRemoteError("http_result_unknown")
return result
@@ -0,0 +1,75 @@
"""领取与续租的固定 localhost HTTP 适配器。"""
from __future__ import annotations
from cmbuyer_client.core.errors import AmbiguousRemoteError, ValidationError
from cmbuyer_client.core.models import ClaimRequest, ClaimedTask, DeviceCredentials, RenewRequest, RenewResult
from cmbuyer_client.core.validation import rfc3339_z_nanoseconds
from .http_transport import HttpTransport
from .wire import (
JSON_RESPONSE_LIMIT,
SMALL_RESPONSE_LIMIT,
classify_json_error,
common_headers,
encode_json,
parse_json_response,
)
class HttpTaskSource:
def __init__(self, transport: HttpTransport) -> None:
self._transport = transport
def claim_next(self, credentials: DeviceCredentials, request: ClaimRequest) -> ClaimedTask | None:
body = encode_json(request.to_wire())
response = self._transport.request(
"POST",
"/api/v1/tasks/claim-next",
common_headers(credentials.device_id, credentials.token.value, "application/json"),
body,
response_limit=JSON_RESPONSE_LIMIT,
)
if response.status == 204:
if response.body or response.header_values("Content-Encoding"):
raise AmbiguousRemoteError("invalid_empty_claim_response")
return None
if response.status != 200:
classify_json_error(
response,
allowed_409=frozenset(("idempotency_conflict", "claim_requires_manual")),
)
try:
claimed = ClaimedTask.from_wire(parse_json_response(response, maximum=JSON_RESPONSE_LIMIT))
except ValidationError as error:
raise AmbiguousRemoteError("invalid_claim_success_response") from error
if rfc3339_z_nanoseconds(claimed.attempt.lease_expires_at) > rfc3339_z_nanoseconds(claimed.authorization.expires_at):
raise AmbiguousRemoteError("invalid_claim_lease")
return claimed
def renew(self, credentials: DeviceCredentials, request: RenewRequest) -> RenewResult:
response = self._transport.request(
"POST",
f"/api/v1/tasks/{request.task_id}/lease/renew",
common_headers(credentials.device_id, credentials.token.value, "application/json"),
encode_json(request.to_wire()),
response_limit=SMALL_RESPONSE_LIMIT,
)
if response.status != 200:
classify_json_error(
response,
allowed_409=frozenset(("idempotency_conflict", "claim_not_current")),
)
try:
result = RenewResult.from_wire(parse_json_response(response, maximum=SMALL_RESPONSE_LIMIT))
except ValidationError as error:
raise AmbiguousRemoteError("invalid_renew_success_response") from error
if (
result.task_id != request.task_id
or result.attempt_id != request.attempt_id
or result.claim_generation != request.claim_generation
or rfc3339_z_nanoseconds(result.lease_expires_at) < rfc3339_z_nanoseconds(request.expected_lease_expires_at)
or rfc3339_z_nanoseconds(result.lease_expires_at) > rfc3339_z_nanoseconds(request.authorization_expires_at)
):
raise AmbiguousRemoteError("renew_response_mismatch")
return result
+106
View File
@@ -0,0 +1,106 @@
"""T-302/T-204 固定 HTTP wire 的编码、解码与错误分类。"""
from __future__ import annotations
import json
from typing import Any, Mapping
from cmbuyer_client.core.errors import (
AmbiguousRemoteError,
CredentialRemoteError,
ManualRemoteError,
ProtocolRemoteError,
ValidationError,
)
from cmbuyer_client.core.validation import require_exact_fields, strict_json_loads
from .http_transport import HttpResponse
JSON_REQUEST_LIMIT = 4096
JSON_RESPONSE_LIMIT = 32 * 1024
SMALL_RESPONSE_LIMIT = 8 * 1024
JSON_CONTENT_TYPES = frozenset(("application/json", "application/json; charset=utf-8"))
def encode_json(value: Mapping[str, object]) -> bytes:
body = json.dumps(value, ensure_ascii=False, separators=(",", ":"), allow_nan=False).encode("utf-8")
if len(body) > JSON_REQUEST_LIMIT:
raise ProtocolRemoteError("request_too_large")
return body
def common_headers(device_id: str, token: str, content_type: str) -> tuple[tuple[str, str], ...]:
return (
("Authorization", "Bearer " + token),
("X-CMBuyer-Device-ID", device_id),
("Accept", "application/json"),
("Content-Type", content_type),
)
def parse_json_response(response: HttpResponse, *, maximum: int) -> object:
encodings = response.header_values("Content-Encoding")
types = response.header_values("Content-Type")
if encodings or len(types) != 1 or types[0].lower() not in JSON_CONTENT_TYPES:
raise ValidationError("invalid_response_content_type")
return strict_json_loads(response.body, maximum=maximum)
def require_empty_response(response: HttpResponse) -> None:
if response.body or response.header_values("Content-Encoding"):
raise ProtocolRemoteError("unexpected_error_body")
def classify_json_error(response: HttpResponse, *, allowed_409: frozenset[str]) -> None:
"""抛出错误,不返回。调用方只在非成功状态使用。"""
if 200 <= response.status <= 299:
# 服务端可能已提交幂等事实;未知 2xx 绝不能终结本地槽或换 key。
raise AmbiguousRemoteError("unknown_success_status")
if response.status == 401:
require_empty_response(response)
raise CredentialRemoteError("device_credential_rejected")
if response.status == 503 or 500 <= response.status <= 599:
# 5xx 无法证明服务端是否在提交响应前完成事务。
raise AmbiguousRemoteError("server_result_unknown")
if response.status == 409:
try:
data = require_exact_fields(parse_json_response(response, maximum=SMALL_RESPONSE_LIMIT), ("error",))
code = data["error"]
except ValidationError as error:
raise ProtocolRemoteError("invalid_conflict_response") from error
if not isinstance(code, str) or code not in allowed_409:
raise ProtocolRemoteError("unknown_conflict")
raise ManualRemoteError(code)
expected = {400: "invalid_request", 413: "request_too_large", 415: "unsupported_media_type"}
if response.status in expected:
try:
data = require_exact_fields(parse_json_response(response, maximum=SMALL_RESPONSE_LIMIT), ("error",))
except ValidationError as error:
raise ProtocolRemoteError("invalid_error_response") from error
if data["error"] != expected[response.status]:
raise ProtocolRemoteError("unexpected_error_code")
raise ProtocolRemoteError(expected[response.status])
if 300 <= response.status <= 399:
raise ProtocolRemoteError("redirect_rejected")
raise ProtocolRemoteError("unexpected_http_status")
def classify_bodyless_error(response: HttpResponse) -> None:
if 200 <= response.status <= 299:
raise AmbiguousRemoteError("unknown_success_status")
if response.status == 401:
require_empty_response(response)
raise CredentialRemoteError("device_credential_rejected")
if response.status == 503 or 500 <= response.status <= 599:
raise AmbiguousRemoteError("server_result_unknown")
if response.status == 409:
require_empty_response(response)
raise ManualRemoteError("evidence_conflict")
if response.status in (400, 403, 413, 415):
require_empty_response(response)
raise ProtocolRemoteError("evidence_request_rejected")
if 300 <= response.status <= 399:
raise ProtocolRemoteError("redirect_rejected")
raise ProtocolRemoteError("unexpected_http_status")
+62 -2
View File
@@ -5,6 +5,7 @@ from __future__ import annotations
from dataclasses import dataclass
import os
from pathlib import Path
from typing import Any, Callable
@dataclass(frozen=True)
@@ -17,21 +18,34 @@ class RuntimePaths:
root: Path
logs: Path
artifacts: Path
state: Path
database: Path
@classmethod
def from_root(cls, root: Path) -> "RuntimePaths":
resolved_root = root.expanduser()
# 路径在进程启动时一次性固化;之后 cwd 改变不能打开第二套数据库或绕过原 mutex。
resolved_root = root.expanduser().resolve(strict=False)
state = resolved_root / "state"
return cls(
root=resolved_root,
logs=resolved_root / "logs",
artifacts=resolved_root / "artifacts",
state=state,
database=state / "client-state.sqlite3",
)
@classmethod
def default(cls) -> "RuntimePaths":
local_app_data = os.environ.get("LOCALAPPDATA")
if local_app_data:
return cls.from_root(Path(local_app_data) / "cmbuyer")
local_root = Path(local_app_data).expanduser()
if not local_root.is_absolute():
raise RuntimeError("local_app_data_must_be_absolute")
return cls.from_root(local_root / "cmbuyer")
if os.name == "nt":
# Windows 上回退到 home 会悄悄创建第二套状态库并绕开同一 mutex,必须失败闭合。
raise RuntimeError("local_app_data_required")
return cls.from_root(Path.home() / ".local" / "share" / "cmbuyer")
@@ -40,3 +54,49 @@ class RuntimePaths:
self.logs.mkdir(parents=True, exist_ok=True)
self.artifacts.mkdir(parents=True, exist_ok=True)
self.state.mkdir(parents=True, exist_ok=True)
@dataclass
class LocalStateRuntime:
"""持有 named mutex 与本地状态库,保证 mutex 总是先取得。"""
paths: RuntimePaths
mutex: Any
store: Any
@classmethod
def open(
cls,
paths: RuntimePaths | None = None,
*,
mutex_factory: Callable[[Path], Any] | None = None,
protector_factory: Callable[[], Any] | None = None,
store_factory: Callable[[Path, Any], Any] | None = None,
) -> "LocalStateRuntime":
from .localstate.protection import DpapiProtector
from .localstate.single_instance import NamedMutex
from .localstate.store import LocalStateStore
selected = paths or RuntimePaths.default()
selected.ensure_exists()
make_mutex = mutex_factory or NamedMutex
make_protector = protector_factory or DpapiProtector
make_store = store_factory or LocalStateStore
mutex = make_mutex(selected.database)
try:
protector = make_protector()
store = make_store(selected.database, protector)
except Exception:
mutex.close()
raise
return cls(selected, mutex, store)
def close(self) -> None:
self.mutex.close()
def __enter__(self) -> "LocalStateRuntime":
return self
def __exit__(self, exc_type: object, exc: object, traceback: object) -> None:
self.close()
+5
View File
@@ -0,0 +1,5 @@
"""采购工具原生 Qt Widgets 界面。"""
from .main_window import PurchaseToolWindow
__all__ = ["PurchaseToolWindow"]
+419
View File
@@ -0,0 +1,419 @@
"""采购执行 Tab:状态、当前任务、滚动日志和历史记录主从视图。"""
from __future__ import annotations
from PySide6.QtCore import QModelIndex, QSize, Qt, Signal, Slot
from PySide6.QtGui import QAction, QKeySequence, QShortcut
from PySide6.QtWidgets import (
QAbstractItemView,
QFrame,
QGroupBox,
QHBoxLayout,
QLabel,
QMenu,
QPlainTextEdit,
QPushButton,
QSizePolicy,
QSplitter,
QStackedWidget,
QTableView,
QToolButton,
QVBoxLayout,
QWidget,
)
from cmbuyer_client.polling.coordinator import ClaimedTaskView, PollingCoordinator, PollingState
from cmbuyer_client.logging_policy import redact_text
from .records import PurchaseRecord, PurchaseRecordModel, PurchaseRecordProvider
class _LogView(QPlainTextEdit):
def __init__(self, parent: QWidget | None = None) -> None:
super().__init__(parent)
self.setObjectName("rollingLog")
self.setReadOnly(True)
self.setPlaceholderText("轮询启动后将在这里显示脱敏日志。")
self.setAccessibleName("滚动日志")
def append_event(self, text: str) -> None:
bar = self.verticalScrollBar()
follow = bar.value() >= bar.maximum() - 2
self.appendPlainText(redact_text(text))
if follow:
bar.setValue(bar.maximum())
class _RecordTableView(QTableView):
"""把双击与 Enter 收敛为唯一 activation 信号,避免平台重复发命令。"""
record_activated = Signal(object)
def mouseDoubleClickEvent(self, event) -> None:
index = self.indexAt(event.position().toPoint())
if index.isValid():
self.setCurrentIndex(index.siblingAtColumn(0))
self.record_activated.emit(index)
event.accept()
return
super().mouseDoubleClickEvent(event)
def keyPressEvent(self, event) -> None:
if event.key() in (Qt.Key.Key_Return, Qt.Key.Key_Enter) and self.currentIndex().isValid():
self.record_activated.emit(self.currentIndex())
event.accept()
return
super().keyPressEvent(event)
class ExecutionPage(QWidget):
LIVE_PAGE = 0
DETAIL_PAGE = 1
COMPACT_DETAIL_WIDTH = 760
def __init__(
self,
coordinator: PollingCoordinator,
record_provider: PurchaseRecordProvider | None = None,
parent: QWidget | None = None,
) -> None:
super().__init__(parent)
self.setObjectName("executionPage")
self._coordinator = coordinator
self._selected_record_id: str | None = None
self._saved_scroll = 0
outer = QVBoxLayout(self)
outer.setContentsMargins(12, 12, 12, 12)
status_row = QHBoxLayout()
self.service_status = self._status("采购服务", "待首次真实领取验证")
self.device_status = self._status("ADB", "待后续执行能力验证")
self.app_status = self._status("拼多多版本", "待后续执行能力验证")
self.session_status = self._status("会话", "已停止")
for widget in (self.service_status, self.device_status, self.app_status, self.session_status):
status_row.addWidget(widget)
status_row.addStretch(1)
self.poll_action = QAction("开始轮询", self)
self.poll_action.setObjectName("pollAction")
self.poll_action.triggered.connect(self._toggle_polling)
self.addAction(self.poll_action)
self.poll_button = QPushButton()
self.poll_button.setObjectName("pollButton")
self.poll_button.clicked.connect(self.poll_action.trigger)
status_row.addWidget(self.poll_button)
outer.addLayout(status_row)
self.banner = QLabel()
self.banner.setObjectName("sessionBanner")
self.banner.setWordWrap(True)
self.banner.setAccessibleName("轮询会话状态")
self.banner.setFrameShape(QFrame.Shape.StyledPanel)
outer.addWidget(self.banner)
self.body_splitter = QSplitter(Qt.Orientation.Horizontal)
self.body_splitter.setObjectName("executionSplitter")
self.body_splitter.setChildrenCollapsible(False)
self.left_stack = QStackedWidget()
self.left_stack.setObjectName("leftWorkspace")
self.left_stack.addWidget(self._build_live_page())
self.left_stack.addWidget(self._build_detail_page())
self.body_splitter.addWidget(self.left_stack)
self.body_splitter.addWidget(self._build_records_page())
self.body_splitter.setStretchFactor(0, 2)
self.body_splitter.setStretchFactor(1, 1)
self.body_splitter.setSizes([760, 380])
outer.addWidget(self.body_splitter, 1)
self.view_record_action = QAction("查看所选记录", self)
self.view_record_action.setObjectName("viewSelectedRecord")
self.view_record_action.setEnabled(False)
self.view_record_action.triggered.connect(self.open_selected_record)
self.addAction(self.view_record_action)
self.view_record_button.setDefaultAction(self.view_record_action)
self.return_action = QAction("返回当前任务", self)
self.return_action.setObjectName("returnToCurrentTask")
self.return_action.setEnabled(False)
self.return_action.triggered.connect(self.return_to_live)
self.addAction(self.return_action)
self.return_button.setDefaultAction(self.return_action)
self.escape_shortcut = QShortcut(QKeySequence(Qt.Key.Key_Escape), self)
self.escape_shortcut.setContext(Qt.ShortcutContext.WidgetWithChildrenShortcut)
self.escape_shortcut.activated.connect(self._escape)
self.record_view.clicked.connect(self._on_record_selected)
self.record_view.record_activated.connect(self._open_index)
self.record_view.setContextMenuPolicy(Qt.ContextMenuPolicy.CustomContextMenu)
self.record_view.customContextMenuRequested.connect(self._show_record_menu)
self.record_view.selectionModel().currentChanged.connect(self._on_current_changed)
coordinator.state_changed.connect(self._on_polling_state)
coordinator.claim_visible.connect(self._show_claimed_task)
self._on_polling_state(coordinator.state, coordinator.reason, coordinator.consecutive_failures)
if record_provider is not None:
# T-304 没有公共历史仓库;只接受调用方准备好的 View DTO 快照,
# 独立应用不注入 provider,模型保持真实空态。
self.set_records(record_provider.snapshot())
@staticmethod
def _status(name: str, value: str) -> QLabel:
label = QLabel(f"{name}\n{value}")
label.setFrameShape(QFrame.Shape.StyledPanel)
label.setMinimumWidth(118)
label.setAccessibleName(name)
return label
def _build_live_page(self) -> QWidget:
page = QWidget()
layout = QVBoxLayout(page)
task_group = QGroupBox("当前任务")
task_layout = QHBoxLayout(task_group)
self.current_task_text = QLabel("当前没有任务。\n启动后只领取已授权任务。")
self.current_task_text.setObjectName("currentTaskText")
self.current_task_text.setWordWrap(True)
self.current_task_text.setAlignment(Qt.AlignmentFlag.AlignTop | Qt.AlignmentFlag.AlignLeft)
self.current_task_text.setTextInteractionFlags(Qt.TextInteractionFlag.TextSelectableByMouse)
self.current_image = QLabel("暂无可信商品图片")
self.current_image.setObjectName("currentTaskImage")
self.current_image.setAlignment(Qt.AlignmentFlag.AlignCenter)
self.current_image.setFrameShape(QFrame.Shape.StyledPanel)
self.current_image.setMinimumSize(QSize(180, 120))
self.current_image.setSizePolicy(QSizePolicy.Policy.Expanding, QSizePolicy.Policy.Expanding)
task_layout.addWidget(self.current_task_text, 2)
task_layout.addWidget(self.current_image, 1)
layout.addWidget(task_group, 1)
log_group = QGroupBox("滚动日志")
log_layout = QVBoxLayout(log_group)
self.log_view = _LogView()
log_layout.addWidget(self.log_view)
layout.addWidget(log_group, 2)
return page
def _build_detail_page(self) -> QWidget:
page = QWidget()
page.setObjectName("recordDetailPage")
layout = QVBoxLayout(page)
header = QHBoxLayout()
self.detail_title = QLabel("采购记录详情")
self.detail_title.setObjectName("recordDetailTitle")
self.detail_title.setStyleSheet("font-size: 18px; font-weight: 600;")
header.addWidget(self.detail_title)
header.addStretch(1)
self.return_button = QToolButton()
self.return_button.setObjectName("returnCurrentTaskButton")
header.addWidget(self.return_button)
layout.addLayout(header)
top = QSplitter(Qt.Orientation.Horizontal)
self.detail_original = QPlainTextEdit()
self.detail_original.setObjectName("recordOriginalText")
self.detail_original.setReadOnly(True)
self.detail_original.setPlaceholderText("没有可显示的原始文字。")
self.detail_image = QLabel("没有可显示的可信图片")
self.detail_image.setObjectName("recordImage")
self.detail_image.setAlignment(Qt.AlignmentFlag.AlignCenter)
self.detail_image.setFrameShape(QFrame.Shape.StyledPanel)
top.addWidget(self.detail_original)
top.addWidget(self.detail_image)
top.setStretchFactor(0, 2)
top.setStretchFactor(1, 1)
layout.addWidget(top, 2)
result_group = QGroupBox("采购结果")
result_layout = QVBoxLayout(result_group)
self.detail_result = QPlainTextEdit()
self.detail_result.setObjectName("recordResult")
self.detail_result.setReadOnly(True)
self.detail_result.setPlaceholderText("暂无采购结果。")
result_layout.addWidget(self.detail_result)
layout.addWidget(result_group, 1)
return page
def _build_records_page(self) -> QWidget:
page = QGroupBox("采购记录")
page.setObjectName("recordsPanel")
self.records_panel = page
layout = QVBoxLayout(page)
header = QHBoxLayout()
self.records_summary = QLabel("暂无记录")
header.addWidget(self.records_summary)
header.addStretch(1)
self.view_record_button = QToolButton()
self.view_record_button.setObjectName("viewSelectedRecordButton")
header.addWidget(self.view_record_button)
layout.addLayout(header)
self.record_model = PurchaseRecordModel(parent=self)
self.record_view = _RecordTableView()
self.record_view.setObjectName("purchaseRecordTable")
self.record_view.setModel(self.record_model)
self.record_view.setSelectionBehavior(QAbstractItemView.SelectionBehavior.SelectRows)
self.record_view.setSelectionMode(QAbstractItemView.SelectionMode.SingleSelection)
self.record_view.setEditTriggers(QAbstractItemView.EditTrigger.NoEditTriggers)
self.record_view.setAlternatingRowColors(True)
self.record_view.setSortingEnabled(False)
self.record_view.horizontalHeader().setStretchLastSection(False)
self.record_view.horizontalHeader().setSectionResizeMode(0, self.record_view.horizontalHeader().ResizeMode.Stretch)
self.record_view.horizontalHeader().setSectionResizeMode(1, self.record_view.horizontalHeader().ResizeMode.ResizeToContents)
self.record_view.verticalHeader().setVisible(False)
layout.addWidget(self.record_view)
return page
def set_records(self, records: list[PurchaseRecord]) -> None:
selected_id = self._selected_record_id
self.record_model.set_records(records)
self.records_summary.setText(f"共 {len(records)} 条" if records else "暂无记录")
if selected_id is not None:
row = self.record_model.row_for_id(selected_id)
if row >= 0:
self.record_view.setCurrentIndex(self.record_model.index(row, 0))
if self.left_stack.currentIndex() == self.DETAIL_PAGE:
self._render_record(self.record_model.record_at(row))
return
self._selected_record_id = None
self.view_record_action.setEnabled(False)
if self.left_stack.currentIndex() == self.DETAIL_PAGE:
self.return_to_live()
@Slot(object, str, int)
def _on_polling_state(self, state: object, reason: str, failures: int) -> None:
polling_state = state if isinstance(state, PollingState) else PollingState.BLOCKED
self.session_status.setText(f"会话\n{self._state_text(polling_state)}")
suffix = f"(连续失败 {failures} 次)" if failures else ""
self.banner.setText(reason + suffix)
running = polling_state in (
PollingState.STARTING,
PollingState.RECOVERING,
PollingState.WAITING,
PollingState.CLAIMING,
PollingState.ACTIVE,
)
self.poll_action.setText("停止轮询" if running else "开始轮询")
self.poll_action.setEnabled(running or self._coordinator.can_start)
self.poll_button.setText(self.poll_action.text())
self.poll_button.setEnabled(self.poll_action.isEnabled())
self.poll_button.setToolTip("" if self.poll_action.isEnabled() else reason)
@staticmethod
def _state_text(state: PollingState) -> str:
return {
PollingState.STOPPED: "已停止",
PollingState.STARTING: "启动中",
PollingState.BLOCKED: "已阻止",
PollingState.RECOVERING: "安全恢复",
PollingState.WAITING: "等待领取",
PollingState.CLAIMING: "正在领取",
PollingState.ACTIVE: "任务执行中",
PollingState.RECOVERY_REQUIRED: "待安全恢复",
}[state]
@Slot()
def _toggle_polling(self) -> None:
if self._coordinator.state in (
PollingState.STARTING,
PollingState.RECOVERING,
PollingState.WAITING,
PollingState.CLAIMING,
PollingState.ACTIVE,
):
self._coordinator.stop()
else:
self._coordinator.start()
@Slot(QModelIndex)
def _on_record_selected(self, index: QModelIndex) -> None:
record = self.record_model.record_at(index.row())
if record is None:
return
self._selected_record_id = record.record_id
self.view_record_action.setEnabled(True)
if self.left_stack.currentIndex() == self.DETAIL_PAGE:
self._saved_scroll = self.record_view.verticalScrollBar().value()
self._render_record(record)
@Slot(object)
def _show_claimed_task(self, raw: object) -> None:
if not isinstance(raw, ClaimedTaskView):
return
self.current_task_text.setText(
f"标题:{raw.title}\n任务 ID:{raw.task_id}\n状态:{raw.status}"
)
self.log_view.append_event(f"已安全领取任务 {raw.task_id};等待单趟执行能力处理。")
@Slot(QModelIndex, QModelIndex)
def _on_current_changed(self, current: QModelIndex, previous: QModelIndex) -> None:
del previous
if current.isValid():
self._on_record_selected(current)
@Slot(QModelIndex)
def _open_index(self, index: QModelIndex) -> None:
if index.isValid():
self.record_view.setCurrentIndex(index.siblingAtColumn(0))
self._on_record_selected(index)
self.view_record_action.trigger()
@Slot()
def open_selected_record(self) -> None:
if self._selected_record_id is None:
return
row = self.record_model.row_for_id(self._selected_record_id)
record = self.record_model.record_at(row)
if record is None:
return
self._saved_scroll = self.record_view.verticalScrollBar().value()
self._render_record(record)
self.left_stack.setCurrentIndex(self.DETAIL_PAGE)
self.return_action.setEnabled(True)
self._apply_compact_detail()
self.return_button.setFocus()
def _render_record(self, record: PurchaseRecord | None) -> None:
if record is None:
self.detail_title.setText("记录不存在")
self.detail_original.clear()
self.detail_result.clear()
self.detail_image.setText("没有可显示的可信图片")
return
self.detail_title.setText(record.title)
self.detail_original.setPlainText(record.original_text)
self.detail_result.setPlainText(record.result_text)
self.detail_image.setText(record.image_description or "没有可显示的可信图片")
@Slot()
def return_to_live(self) -> None:
if self.left_stack.currentIndex() != self.DETAIL_PAGE:
return
self.left_stack.setCurrentIndex(self.LIVE_PAGE)
self.return_action.setEnabled(False)
self.records_panel.setVisible(True)
row = self.record_model.row_for_id(self._selected_record_id or "")
if row >= 0:
self.record_view.setCurrentIndex(self.record_model.index(row, 0))
self.record_view.verticalScrollBar().setValue(self._saved_scroll)
self.record_view.setFocus()
@Slot()
def _escape(self) -> None:
# Qt popup/menu 优先消费 Esc;只有详情态的页面级 shortcut 会执行返回。
if self.left_stack.currentIndex() == self.DETAIL_PAGE:
self.return_action.trigger()
@Slot(object)
def _show_record_menu(self, point: object) -> None:
index = self.record_view.indexAt(point)
if index.isValid():
self.record_view.setCurrentIndex(index.siblingAtColumn(0))
self._on_record_selected(index)
menu = QMenu(self.record_view)
menu.addAction(self.view_record_action)
menu.exec(self.record_view.viewport().mapToGlobal(point))
def resizeEvent(self, event) -> None:
super().resizeEvent(event)
self._apply_compact_detail()
def _apply_compact_detail(self) -> None:
compact_detail = self.width() < self.COMPACT_DETAIL_WIDTH and self.left_stack.currentIndex() == self.DETAIL_PAGE
self.records_panel.setVisible(not compact_detail)
@@ -0,0 +1,95 @@
"""采购工具固定双 Tab 原生窗口。"""
from __future__ import annotations
from PySide6.QtCore import QTimer, Qt, Slot
from PySide6.QtGui import QCloseEvent
from PySide6.QtWidgets import QMainWindow, QTabWidget
from cmbuyer_client.localstate.models import ProfileSettings
from cmbuyer_client.polling.coordinator import PollingCoordinator, PollingState, RecoveryStatus
from .execution import ExecutionPage
from .records import PurchaseRecordProvider
from .settings import ProfileStore, SettingsPage
class PurchaseToolWindow(QMainWindow):
EXECUTION_PAGE_ID = "purchase-execution"
SETTINGS_PAGE_ID = "settings"
def __init__(
self,
*,
store: ProfileStore,
coordinator: PollingCoordinator,
profile_settings: ProfileSettings | None,
has_stored_device_token: bool,
record_provider: PurchaseRecordProvider | None = None,
parent=None,
) -> None:
super().__init__(parent)
self.setWindowTitle("采购工具")
self.setAccessibleName("采购工具")
self.setMinimumSize(720, 520)
self.resize(1180, 760)
self._coordinator = coordinator
self._close_pending = False
recovery = coordinator.recovery_status
frozen = recovery is None or recovery.has_pending_claim or recovery.has_active_claim
self.tabs = QTabWidget()
self.tabs.setObjectName("mainTabs")
self.tabs.setTabsClosable(False)
self.tabs.setMovable(False)
self.execution_page = ExecutionPage(coordinator, record_provider)
self.execution_page.setProperty("pageId", self.EXECUTION_PAGE_ID)
self.settings_page = SettingsPage(
store,
settings=profile_settings,
has_stored_device_token=has_stored_device_token,
identity_frozen=frozen,
)
self.settings_page.setProperty("pageId", self.SETTINGS_PAGE_ID)
self.tabs.addTab(self.execution_page, "采购执行")
self.tabs.addTab(self.settings_page, "配置")
self.tabs.setCurrentWidget(self.execution_page)
self.setCentralWidget(self.tabs)
self.settings_page.settings_saved.connect(
lambda settings, has_token: self._coordinator.update_profile_settings(settings)
)
coordinator.recovery_status_changed.connect(self._on_recovery_status)
coordinator.configuration_freeze_changed.connect(self.settings_page.set_identity_frozen)
coordinator.settled.connect(self._finish_pending_close)
@Slot(object)
def _on_recovery_status(self, raw: object) -> None:
if not isinstance(raw, RecoveryStatus):
self.settings_page.set_identity_frozen(True)
return
self.settings_page.set_identity_frozen(raw.has_pending_claim or raw.has_active_claim)
def closeEvent(self, event: QCloseEvent) -> None:
running = self._coordinator.state in (
PollingState.STARTING,
PollingState.RECOVERING,
PollingState.WAITING,
PollingState.CLAIMING,
PollingState.ACTIVE,
) or self._coordinator.operation_in_flight
if running:
# 不 terminate 飞行中的 QThread。先提升 stop latch,等待 HTTP 自身
# 超时和 DurableClientGateway 落库,再由 settled 重试关闭。
self._close_pending = True
self._coordinator.stop()
event.ignore()
return
event.accept()
@Slot()
def _finish_pending_close(self) -> None:
if not self._close_pending:
return
self._close_pending = False
QTimer.singleShot(0, self.close)
+95
View File
@@ -0,0 +1,95 @@
"""采购记录的只读 Qt Model/View 数据源。"""
from __future__ import annotations
from dataclasses import dataclass, field
from typing import Protocol
from PySide6.QtCore import QAbstractTableModel, QModelIndex, Qt
from cmbuyer_client.logging_policy import redact_text
from cmbuyer_client.core.errors import ValidationError
from cmbuyer_client.core.validation import rfc3339_z_nanoseconds
@dataclass(frozen=True)
class PurchaseRecord:
record_id: str
title: str
status: str
created_at: str
original_text: str = ""
result_text: str = ""
image_description: str = ""
created_at_nanoseconds: int = field(init=False, repr=False)
def __post_init__(self) -> None:
# 记录 provider 只能注入可显示摘要;最终 UI 边界仍统一脱敏,避免
# consumer bug 把 Bearer/裸 token 放入 model、详情或可见日志。
for field in ("title", "status", "original_text", "result_text", "image_description"):
object.__setattr__(self, field, redact_text(getattr(self, field)))
if not isinstance(self.created_at, str):
raise ValueError("noncanonical_record_timestamp")
fraction = self.created_at[20:-1] if len(self.created_at) > 20 and self.created_at.endswith("Z") else ""
if fraction and fraction.endswith("0"):
raise ValueError("noncanonical_record_timestamp")
try:
timestamp = rfc3339_z_nanoseconds(self.created_at)
except ValidationError:
raise ValueError("noncanonical_record_timestamp") from None
object.__setattr__(self, "created_at_nanoseconds", timestamp)
class PurchaseRecordProvider(Protocol):
"""只返回已准备好的无秘密 View DTO;不得在 GUI 线程查询 SQLite/HTTP。"""
def snapshot(self) -> list[PurchaseRecord]: ...
class PurchaseRecordModel(QAbstractTableModel):
RECORD_ID_ROLE = int(Qt.ItemDataRole.UserRole) + 1
def __init__(self, records: list[PurchaseRecord] | None = None, parent=None) -> None:
super().__init__(parent)
self._records: list[PurchaseRecord] = []
self.set_records(records or [])
def rowCount(self, parent: QModelIndex = QModelIndex()) -> int:
return 0 if parent.isValid() else len(self._records)
def columnCount(self, parent: QModelIndex = QModelIndex()) -> int:
return 0 if parent.isValid() else 2
def data(self, index: QModelIndex, role: int = int(Qt.ItemDataRole.DisplayRole)):
if not index.isValid() or not 0 <= index.row() < len(self._records):
return None
record = self._records[index.row()]
if role == int(Qt.ItemDataRole.DisplayRole):
return record.title if index.column() == 0 else record.status
if role == self.RECORD_ID_ROLE:
return record.record_id
if role == int(Qt.ItemDataRole.ToolTipRole):
return f"{record.title}\n{record.created_at}"
if role == int(Qt.ItemDataRole.TextAlignmentRole) and index.column() == 1:
return int(Qt.AlignmentFlag.AlignCenter)
return None
def headerData(self, section: int, orientation: Qt.Orientation, role: int = int(Qt.ItemDataRole.DisplayRole)):
if role != int(Qt.ItemDataRole.DisplayRole) or orientation != Qt.Orientation.Horizontal:
return None
return ("标题", "状态")[section] if 0 <= section < 2 else None
def set_records(self, records: list[PurchaseRecord]) -> None:
self.beginResetModel()
self._records = sorted(
records,
key=lambda item: (item.created_at_nanoseconds, item.record_id),
reverse=True,
)
self.endResetModel()
def record_at(self, row: int) -> PurchaseRecord | None:
return self._records[row] if 0 <= row < len(self._records) else None
def row_for_id(self, record_id: str) -> int:
return next((row for row, item in enumerate(self._records) if item.record_id == record_id), -1)
+249
View File
@@ -0,0 +1,249 @@
"""只做本地校验和显式保存的配置页。"""
from __future__ import annotations
from pathlib import Path
from typing import Protocol
from PySide6.QtCore import Qt, Signal, Slot
from PySide6.QtGui import QAction, QKeySequence
from PySide6.QtWidgets import (
QComboBox,
QFormLayout,
QLabel,
QLineEdit,
QPushButton,
QScrollArea,
QSpinBox,
QVBoxLayout,
QWidget,
)
from cmbuyer_client.core.models import SecretToken
from cmbuyer_client.core.errors import ValidationError
from cmbuyer_client.core.validation import require_uuid4
from cmbuyer_client.localstate.models import LOOPBACK_SERVICE_URL, ProfileSettings
class ProfileStore(Protocol):
def save_profile(self, settings: ProfileSettings, token: SecretToken | None) -> None: ...
class SettingsPage(QScrollArea):
settings_saved = Signal(object, bool)
VALIDATION_HINT = "服务身份将在首次真实领取时验证;设备与 App 状态由后续已取证执行能力验证。"
def __init__(
self,
store: ProfileStore,
*,
profile_id: str = "default",
settings: ProfileSettings | None = None,
has_stored_device_token: bool = False,
identity_frozen: bool = False,
parent: QWidget | None = None,
) -> None:
super().__init__(parent)
self.setObjectName("settingsPage")
self.setWidgetResizable(True)
self._store = store
self._profile_id = profile_id
self._has_stored_device_token = has_stored_device_token
self._loaded_settings = settings
self._identity_frozen = identity_frozen
content = QWidget()
outer = QVBoxLayout(content)
title = QLabel("配置")
title.setObjectName("settingsTitle")
title.setStyleSheet("font-size: 20px; font-weight: 600;")
outer.addWidget(title)
form = QFormLayout()
form.setFieldGrowthPolicy(QFormLayout.FieldGrowthPolicy.ExpandingFieldsGrow)
form.setLabelAlignment(Qt.AlignmentFlag.AlignRight | Qt.AlignmentFlag.AlignVCenter)
outer.addLayout(form)
self.service_url = QLineEdit(LOOPBACK_SERVICE_URL)
self.service_url.setReadOnly(True)
self.service_url.setObjectName("serviceUrl")
form.addRow("采购服务 URL", self.service_url)
self.device_id = QLineEdit()
self.device_id.setObjectName("deviceId")
self.device_id.setPlaceholderText("小写 UUIDv4")
form.addRow("设备 UUID", self.device_id)
self.device_token = QLineEdit()
self.device_token.setObjectName("deviceToken")
self.device_token.setEchoMode(QLineEdit.EchoMode.Password)
# 控件只负责给输入设置合理上限;长度与字符集必须在 save() 中显式
# 验证。若这里限制为 64,粘贴 65 位 token 会被 Qt 静默截成合法
# 64 位并覆盖原凭据。
self.device_token.setMaxLength(256)
self.device_token.setPlaceholderText("首次必填;已有凭据时留空表示保留")
form.addRow("设备 token", self.device_token)
self.token_status = QLabel()
self.token_status.setObjectName("tokenStatus")
form.addRow("凭据状态", self.token_status)
self.adb_path = QLineEdit()
self.adb_path.setObjectName("adbPath")
form.addRow("ADB 路径", self.adb_path)
self.adb_serial = QLineEdit()
self.adb_serial.setObjectName("adbSerial")
form.addRow("设备 serial", self.adb_serial)
self.transport = QComboBox()
self.transport.setObjectName("transport")
self.transport.addItem("USB", "usb")
self.transport.addItem("WiFi", "wifi")
form.addRow("连接方式", self.transport)
self.poll_interval = self._spin(5, 300, 15, "pollInterval", " 秒")
form.addRow("轮询间隔", self.poll_interval)
self.failure_threshold = self._spin(1, 10, 3, "failureThreshold", " 次")
form.addRow("连续失败停止阈值", self.failure_threshold)
self.http_timeout = self._spin(1, 120, 10, "httpTimeout", " 秒")
form.addRow("HTTP 超时", self.http_timeout)
self.step_timeout = self._spin(5, 300, 45, "stepTimeout", " 秒")
form.addRow("真机步骤超时", self.step_timeout)
self.validation_hint = QLabel(self.VALIDATION_HINT)
self.validation_hint.setObjectName("deferredValidationHint")
self.validation_hint.setWordWrap(True)
outer.addWidget(self.validation_hint)
self.feedback = QLabel()
self.feedback.setObjectName("settingsFeedback")
self.feedback.setWordWrap(True)
self.feedback.setAccessibleName("配置保存状态")
outer.addWidget(self.feedback)
self.save_action = QAction("保存配置", self)
self.save_action.setShortcut(QKeySequence.StandardKey.Save)
self.save_action.triggered.connect(self.save)
self.addAction(self.save_action)
self.save_button = QPushButton("保存配置")
self.save_button.setObjectName("saveSettings")
self.save_button.clicked.connect(self.save_action.trigger)
outer.addWidget(self.save_button, 0, Qt.AlignmentFlag.AlignRight)
outer.addStretch(1)
self.setWidget(content)
if settings is not None:
self._load(settings)
self._update_token_status()
self.set_identity_frozen(identity_frozen)
@staticmethod
def _spin(minimum: int, maximum: int, value: int, name: str, suffix: str) -> QSpinBox:
field = QSpinBox()
field.setObjectName(name)
field.setRange(minimum, maximum)
field.setValue(value)
field.setSuffix(suffix)
return field
def _load(self, settings: ProfileSettings) -> None:
self.device_id.setText(settings.device_id)
self.adb_path.setText(settings.adb_path)
self.adb_serial.setText(settings.adb_serial)
self.transport.setCurrentIndex(max(0, self.transport.findData(settings.transport)))
self.poll_interval.setValue(settings.poll_interval_seconds)
self.failure_threshold.setValue(settings.failure_threshold)
self.http_timeout.setValue(settings.http_timeout_seconds)
self.step_timeout.setValue(settings.step_timeout_seconds)
def set_identity_frozen(self, frozen: bool) -> None:
self._identity_frozen = frozen
for field in (
self.device_id,
self.adb_path,
self.adb_serial,
self.transport,
self.poll_interval,
self.failure_threshold,
self.http_timeout,
self.step_timeout,
):
field.setEnabled(not frozen)
if frozen:
self.feedback.setText("存在待恢复或执行中的领取;服务与设备身份参数已冻结。")
elif self.feedback.text().startswith("存在待恢复或执行中的领取"):
self.feedback.clear()
@Slot()
def save(self) -> None:
self.feedback.clear()
token_text = self.device_token.text().strip()
device_id = self.device_id.text().strip()
try:
require_uuid4(device_id, "invalid_device_id")
except (TypeError, ValueError, ValidationError):
self._validation_error(self.device_id, "设备 UUID 必须是小写 UUIDv4。")
return
if not self._has_stored_device_token and not token_text:
self._validation_error(self.device_token, "首次保存必须填写设备 token。")
return
if token_text and (len(token_text) != 64 or any(character not in "0123456789abcdef" for character in token_text)):
self._validation_error(self.device_token, "设备 token 必须为 64 位小写十六进制。")
return
try:
if self._identity_frozen:
if self._loaded_settings is None:
self._validation_error(self.device_id, "冻结配置缺少原始设置,不能保存。")
return
settings = self._loaded_settings
else:
adb_text = self.adb_path.text().strip()
if not adb_text:
self._validation_error(self.adb_path, "请填写 ADB 路径。")
return
adb_file = Path(adb_text).expanduser()
if not adb_file.is_file():
self._validation_error(self.adb_path, "ADB 路径必须指向本机已存在的文件。")
return
adb_file = adb_file.resolve(strict=True)
serial = self.adb_serial.text().strip()
if not serial:
self._validation_error(self.adb_serial, "请填写设备 serial。")
return
settings = ProfileSettings(
profile_id=self._profile_id,
service_url=LOOPBACK_SERVICE_URL,
device_id=device_id,
adb_path=str(adb_file),
adb_serial=serial,
transport=str(self.transport.currentData()),
poll_interval_seconds=self.poll_interval.value(),
failure_threshold=self.failure_threshold.value(),
http_timeout_seconds=self.http_timeout.value(),
step_timeout_seconds=self.step_timeout.value(),
)
token = SecretToken(token_text) if token_text else None
except (TypeError, ValueError):
self._validation_error(self.device_id, "配置格式不正确,请检查设备 UUID 与各项参数。")
return
try:
# None 明确表示保留 T-303 中已有的 DPAPI 密文,绝不是清除凭据。
self._store.save_profile(settings, token)
except Exception:
self.feedback.setText("配置保存失败,本地安全存储未更新。")
(self.device_token if token_text else self.device_id).setFocus()
return
self._has_stored_device_token = True
self._loaded_settings = settings
self.device_token.clear()
self._update_token_status()
self.feedback.setText("配置已保存。本地校验不代表服务、设备或 App 已就绪。")
self.settings_saved.emit(settings, True)
def _validation_error(self, field: QWidget, message: str) -> None:
self.feedback.setText(message)
field.setFocus()
def _update_token_status(self) -> None:
self.token_status.setText("已保存" if self._has_stored_device_token else "未保存")
+1
View File
@@ -0,0 +1 @@
"""core tests。"""
+187
View File
@@ -0,0 +1,187 @@
from __future__ import annotations
import json
import unittest
from cmbuyer_client.core.errors import ValidationError
from cmbuyer_client.core.models import ClaimedTask, SecretToken
from cmbuyer_client.core.validation import rfc3339_z_nanoseconds, strict_json_loads
TASK_ID = "13c9f507-7473-4fa6-8d71-8786c34c6301"
AUTH_ID = "73c9f507-7473-4fa6-8d71-8786c34c6301"
ATTEMPT_ID = "53c9f507-7473-4fa6-8d71-8786c34c6301"
TOKEN = "0123456789abcdef" * 4
def claim_wire() -> dict[str, object]:
return {
"task": {
"id": TASK_ID,
"version": 3,
"title": "纯棉短袖",
"product_url": "https://mobile.yangkeduo.com/goods.html?goods_id=937122477375",
"goods_id": "937122477375",
"sku_color": "黑色CHA(纯棉)",
"sku_size": "M(建议100-115)",
"quantity": 2,
"max_total_price": "30.00",
},
"authorization": {"id": AUTH_ID, "task_version": 2, "expires_at": "2026-08-04T10:00:00Z"},
"attempt": {
"id": ATTEMPT_ID,
"claim_token": TOKEN,
"claim_generation": 1,
"lease_expires_at": "2026-08-04T09:05:00Z",
},
}
class CoreModelsTests(unittest.TestCase):
def test_claim_wire_round_trip_and_secret_repr(self) -> None:
claimed = ClaimedTask.from_wire(claim_wire())
self.assertEqual(claimed.task.quantity, 2)
self.assertNotIn(TOKEN, repr(claimed))
self.assertNotIn(TOKEN, repr(SecretToken(TOKEN)))
def test_rejects_bool_float_wrong_url_and_version_drift(self) -> None:
mutations = []
for mutate in (
lambda value: value["task"].__setitem__("quantity", True),
lambda value: value["task"].__setitem__("max_total_price", "30.0"),
lambda value: value["task"].__setitem__("max_total_price", "0.00"),
lambda value: value["task"].__setitem__("product_url", "https://example.invalid/"),
lambda value: value["task"].__setitem__("version", 2),
):
value = claim_wire()
mutate(value)
mutations.append(value)
for value in mutations:
with self.subTest(value=value), self.assertRaises(ValidationError):
ClaimedTask.from_wire(value)
def test_strict_json_rejects_nested_duplicates_float_nan_bom_and_utf8(self) -> None:
bad_values = (
b'{"task":{"id":1,"id":2}}',
b'{"value":1.0}',
b'{"value":NaN}',
b'\xef\xbb\xbf{}',
b'\xff',
('{"value":' + "9" * 5000 + '}').encode(),
)
for raw in bad_values:
with self.subTest(raw=raw), self.assertRaises(ValidationError):
strict_json_loads(raw, maximum=1024)
self.assertEqual(strict_json_loads(json.dumps({"value": 1}).encode(), maximum=1024), {"value": 1})
def test_rfc3339_nano_comparison_preserves_all_fraction_digits(self) -> None:
equal = (
"2026-08-04T09:01:00.1Z",
"2026-08-04T09:01:00.100000Z",
"2026-08-04T09:01:00.100000000Z",
)
self.assertEqual(len({rfc3339_z_nanoseconds(value) for value in equal}), 1)
ordered = (
"2026-08-04T09:01:00Z",
"2026-08-04T09:01:00.000001Z",
"2026-08-04T09:01:00.0000011Z",
"2026-08-04T09:01:00.000001101Z",
"2026-08-04T09:01:01Z",
)
self.assertEqual([rfc3339_z_nanoseconds(value) for value in ordered], sorted(rfc3339_z_nanoseconds(value) for value in ordered))
def test_money_accepts_positive_subunit_but_rejects_zero_and_noncanonical_forms(self) -> None:
value = claim_wire()
value["task"]["max_total_price"] = "0.01"
self.assertEqual(ClaimedTask.from_wire(value).task.max_total_price, "0.01")
for invalid in ("0.00", "00.01", "1.0", "1.000", "1", 1.0, "1.12", "1.٠٠", "12.00"):
with self.subTest(invalid=invalid), self.assertRaises(ValidationError):
changed = claim_wire()
changed["task"]["max_total_price"] = invalid
ClaimedTask.from_wire(changed)
wide_quantity = claim_wire()
wide_quantity["task"]["quantity"] = 2_147_483_648
self.assertEqual(ClaimedTask.from_wire(wide_quantity).task.quantity, 2_147_483_648)
for invalid_goods in ("123", "1٢3"):
changed = claim_wire()
changed["task"]["goods_id"] = invalid_goods
changed["task"]["product_url"] = "https://mobile.yangkeduo.com/goods.html?goods_id=" + invalid_goods
with self.subTest(invalid_goods=invalid_goods), self.assertRaises(ValidationError):
ClaimedTask.from_wire(changed)
too_large = claim_wire()
too_large["task"]["quantity"] = 9_223_372_036_854_775_808
with self.assertRaises(ValidationError):
ClaimedTask.from_wire(too_large)
def test_claim_fields_share_explicit_server_bounds(self) -> None:
legal = claim_wire()
legal_goods = "1" * 32
legal["task"].update(
version=9_223_372_036_854_775_807,
title="😀" * 120,
goods_id=legal_goods,
product_url="https://mobile.yangkeduo.com/goods.html?goods_id=" + legal_goods,
sku_color="色" * 80,
sku_size="码" * 80,
max_total_price="1" * 29 + ".00",
)
legal["authorization"]["task_version"] = 9_223_372_036_854_775_806
claimed = ClaimedTask.from_wire(legal)
self.assertEqual(len(claimed.task.title), 120)
# Python's default ensure_ascii=True expands astral characters to surrogate
# escape pairs, so this is a conservative parser-budget proof as well.
self.assertLess(len(json.dumps(legal, separators=(",", ":")).encode()), 32 * 1024)
mutations = (
("title", "😀" * 121),
("title", " title"),
("sku_color", "色" * 81),
("sku_color", "black "),
("sku_size", "码" * 81),
("sku_size", " M"),
("max_total_price", "1" * 30 + ".00"),
)
for field, invalid in mutations:
changed = claim_wire()
changed["task"][field] = invalid
with self.subTest(field=field, length=len(invalid)), self.assertRaises(ValidationError):
ClaimedTask.from_wire(changed)
overlong_goods = "1" * 33
changed = claim_wire()
changed["task"].update(
goods_id=overlong_goods,
product_url="https://mobile.yangkeduo.com/goods.html?goods_id=" + overlong_goods,
)
with self.assertRaises(ValidationError):
ClaimedTask.from_wire(changed)
def test_wire_strings_reject_lone_surrogates_but_accept_valid_pair(self) -> None:
for escaped in (r'"\ud800"', r'"\udc00"'):
value = claim_wire()
value["task"]["title"] = json.loads(escaped)
with self.subTest(escaped=escaped), self.assertRaises(ValidationError):
ClaimedTask.from_wire(value)
value = claim_wire()
value["task"]["title"] = json.loads(r'"\ud83d\ude00"')
self.assertEqual(ClaimedTask.from_wire(value).task.title, "😀")
def test_title_rejects_ascii_and_unicode_whitespace_only(self) -> None:
for title in ("", " \t\r\n", "\u3000", " \u3000\t", "\u00a0title", "title\u00a0"):
value = claim_wire()
value["task"]["title"] = title
with self.subTest(title=repr(title)), self.assertRaises(ValidationError):
ClaimedTask.from_wire(value)
def test_persisted_text_has_runtime_independent_c0_and_nbsp_domain(self) -> None:
for field in ("title", "sku_color", "sku_size"):
for invalid in ("\u001cvalue", "value\u001f", "value\u001dinside", "\u00a0value", "value\u00a0"):
value = claim_wire()
value["task"][field] = invalid
with self.subTest(field=field, invalid=repr(invalid)), self.assertRaises(ValidationError):
ClaimedTask.from_wire(value)
value = claim_wire()
value["task"][field] = "left\u00a0right"
self.assertEqual(getattr(ClaimedTask.from_wire(value).task, field), "left\u00a0right")
@@ -2,7 +2,9 @@
from __future__ import annotations
from functools import lru_cache
from hashlib import sha256
from io import BytesIO
import json
from pathlib import Path
import sys
@@ -34,7 +36,7 @@ TEST_ADDRESS = "SYNTHETIC_ADDRESS_NEVER_PUBLISH"
FULL_PHONE = "13800138000"
MASKED_PHONE = "138****0000"
SAFE_TEXT = "synthetic-safe-lower-content"
CURRENT_PRICE = "快卖光 ¥12.88"
CURRENT_PRICE = "快卖完 ¥12.88"
ORIGINAL_PRICE = "¥29.00"
PRICE_CURRENT_BOUNDS = "[396,503][712,570]"
PRICE_ORIGINAL_BOUNDS = "[730,503][895,570]"
@@ -48,6 +50,16 @@ def _hash(path: Path) -> str:
return digest.hexdigest()
@lru_cache(maxsize=None)
def _source_png(size: tuple[int, int]) -> bytes:
image = Image.new("RGB", size, color=(0, 180, 0))
if size == (EXPECTED_SCREENSHOT_WIDTH, EXPECTED_SCREENSHOT_HEIGHT):
image.paste((255, 0, 0), (0, 0, 8, 540))
output = BytesIO()
image.save(output, format="PNG")
return output.getvalue()
def _default_xml() -> str:
return (
"<hierarchy rotation='0'>"
@@ -116,12 +128,7 @@ def _write_raw(
raw = root / "raw"
raw.mkdir(parents=True)
screenshot = raw / "screenshot.png"
image = Image.new("RGB", size, color=(0, 180, 0))
if size == (EXPECTED_SCREENSHOT_WIDTH, EXPECTED_SCREENSHOT_HEIGHT):
for y in range(540):
for x in range(8):
image.putpixel((x, y), (255, 0, 0))
image.save(screenshot, format="PNG")
screenshot.write_bytes(_source_png(size))
hierarchy = raw / "hierarchy.xml"
hierarchy.write_text(_default_xml() if xml is None else xml, encoding="utf-8")
manifest = {
@@ -166,7 +173,7 @@ class SkuEvidenceSanitizerTests(unittest.TestCase):
self.assertNotIn(MASKED_PHONE, derived_xml)
self.assertIn(f'"human_declared_state": "{state}"', manifest)
self.assertIn('"privacy_tier": "SANITIZED"', manifest)
self.assertIn('"sanitizer_version": "t103-privacy-v4"', manifest)
self.assertIn('"sanitizer_version": "t103-privacy-v5"', manifest)
self.assertIn('"screenshot_space": {', manifest)
self.assertIn('"xml_coordinate_space": {', manifest)
self.assertIn('"height": 2376', manifest)
@@ -256,16 +263,16 @@ class SkuEvidenceSanitizerTests(unittest.TestCase):
def test_crossing_price_window_rejects_text_and_structure_drift(self) -> None:
bad_texts = (
f"快卖光 ¥12.88 {TEST_ADDRESS}",
f"快卖光 ¥12.88 {FULL_PHONE}",
"快卖光 ¥12.88 使用微信支付",
"快卖光 ¥12.88 提交订单",
"快卖光 ¥12.88 优惠-11元",
f"快卖完 ¥12.88 {TEST_ADDRESS}",
f"快卖完 ¥12.88 {FULL_PHONE}",
"快卖完 ¥12.88 使用微信支付",
"快卖完 ¥12.88 提交订单",
"快卖完 ¥12.88 优惠-11元",
"快要抢光 ¥12.88",
"快卖光 ¥0.00",
"快卖光 ¥12.8",
"快卖光 ¥12.880",
"快卖光 ¥12.88",
"快卖完 ¥0.00",
"快卖完 ¥12.8",
"快卖完 ¥12.880",
"快卖完 ¥12.88",
)
for text in bad_texts:
with self.subTest(text=text), TemporaryDirectory() as temporary:
@@ -276,31 +283,37 @@ class SkuEvidenceSanitizerTests(unittest.TestCase):
def test_crossing_price_text_mismatch_reports_only_fixed_slot_and_reason(self) -> None:
newline_node = _price_node(CURRENT_PRICE, PRICE_CURRENT_BOUNDS).replace(
"快卖光 ¥12.88", "快卖光&#10;¥12.88"
"快卖完 ¥12.88", "快卖完&#10;¥12.88"
)
cases = (
("newline", _xml_with_prices(current_node=newline_node), "newline", PRICE_CURRENT_BOUNDS),
(
"non-ascii-whitespace",
_xml_with_prices(current="快卖光 ¥12.88"),
_xml_with_prices(current="快卖完 ¥12.88"),
"non_ascii_whitespace",
PRICE_CURRENT_BOUNDS,
),
(
"known-prefix-missing",
_xml_with_prices(current="快要抢光 ¥12.88"),
"observed_prefix_kuaiyaoqiangguang",
PRICE_CURRENT_BOUNDS,
),
(
"other-kuai-prefix-remains-generic",
_xml_with_prices(current="快递地址 ¥12.88"),
"known_prefix_missing",
PRICE_CURRENT_BOUNDS,
),
(
"currency-missing",
_xml_with_prices(current="快卖光 12.88"),
_xml_with_prices(current="快卖完 12.88"),
"currency_missing",
PRICE_CURRENT_BOUNDS,
),
(
"amount-shape",
_xml_with_prices(current="快卖光 ¥12.8"),
_xml_with_prices(current="快卖完 ¥12.8"),
"amount_shape",
PRICE_CURRENT_BOUNDS,
),
@@ -337,10 +350,10 @@ class SkuEvidenceSanitizerTests(unittest.TestCase):
def test_crossing_price_text_mismatch_never_echoes_sensitive_or_order_text(self) -> None:
cases = (
f"快卖光 ¥12.88 {TEST_ADDRESS}",
f"快卖光 ¥12.88 {FULL_PHONE}",
"快卖光 ¥12.88 使用微信支付",
"快卖光 ¥12.88 提交订单",
f"快卖完 ¥12.88 {TEST_ADDRESS}",
f"快卖完 ¥12.88 {FULL_PHONE}",
"快卖完 ¥12.88 使用微信支付",
"快卖完 ¥12.88 提交订单",
)
for text in cases:
with self.subTest(text=text), TemporaryDirectory() as temporary:
@@ -354,14 +367,44 @@ class SkuEvidenceSanitizerTests(unittest.TestCase):
for raw_fragment in (TEST_ADDRESS, FULL_PHONE, "使用微信支付", "提交订单", "¥12.88"):
self.assertNotIn(raw_fragment, message)
def test_observed_prefix_diagnostic_does_not_echo_its_suffix(self) -> None:
suffix = "地址和金额都不得回显"
text = f"快要抢光 ¥12.88 {suffix}"
with TemporaryDirectory() as temporary:
raw = _write_raw(Path(temporary), xml=_xml_with_prices(current=text))
with self.assertRaises(SkuEvidenceSanitizationError) as raised:
sanitize_sku_panel_evidence(raw, raw.parent / "derived")
message = str(raised.exception)
self.assertEqual(
message,
"跨界价格节点文本不匹配:slot=[396,503][712,570];"
"reason=observed_prefix_kuaiyaoqiangguang。",
)
self.assertNotIn(suffix, message)
self.assertNotIn("¥12.88", message)
def test_crossing_price_projection_allows_only_limited_ascii_spaces_and_yen_variants(self) -> None:
for current in ("快卖光 ¥12.88", " 快卖光 ¥ 12.88 ", "快卖光 ¥12.88"):
for current in ("快卖完 ¥12.88", " 快卖完 ¥ 12.88 ", "快卖完 ¥12.88"):
with self.subTest(current=current), TemporaryDirectory() as temporary:
raw = _write_raw(Path(temporary), xml=_xml_with_prices(current=current))
result = sanitize_sku_panel_evidence(raw, raw.parent / "derived")
hierarchy = result.hierarchy_path.read_text(encoding="utf-8")
self.assertIn(current, hierarchy)
def test_unverified_or_missing_current_price_prefixes_fail_closed(self) -> None:
cases = (
("old-prefix", "快卖光 ¥12.88"),
("bottom-button-prefix", "快要抢光 ¥12.88"),
("missing-prefix", "¥12.88"),
)
for name, current in cases:
with self.subTest(name=name), TemporaryDirectory() as temporary:
raw = _write_raw(Path(temporary), xml=_xml_with_prices(current=current))
with self.assertRaises(SkuEvidenceSanitizationError):
sanitize_sku_panel_evidence(raw, raw.parent / "derived")
self.assertFalse((raw.parent / "derived").exists())
def test_unique_current_price_without_original_price_is_published(self) -> None:
with TemporaryDirectory() as temporary:
raw = _write_raw(Path(temporary), xml=_xml_with_prices(include_original=False))
@@ -375,7 +418,7 @@ class SkuEvidenceSanitizerTests(unittest.TestCase):
def test_crossing_price_projection_rejects_newline_and_structure_drift(self) -> None:
encoded_newline = _price_node(CURRENT_PRICE, PRICE_CURRENT_BOUNDS).replace(
"快卖光 ¥12.88", "快卖光&#10;¥12.88"
"快卖完 ¥12.88", "快卖完&#10;¥12.88"
)
with TemporaryDirectory() as temporary:
raw = _write_raw(Path(temporary), xml=_xml_with_prices(current_node=encoded_newline))
@@ -409,7 +452,7 @@ class SkuEvidenceSanitizerTests(unittest.TestCase):
),
(
"duplicate-current",
_xml_with_prices(current=CURRENT_PRICE, original=f"快卖光 {ORIGINAL_PRICE}"),
_xml_with_prices(current=CURRENT_PRICE, original=f"快卖完 {ORIGINAL_PRICE}"),
),
(
"multiple-original",
+1
View File
@@ -0,0 +1 @@
"""localstate tests。"""
@@ -0,0 +1,36 @@
from __future__ import annotations
import ast
from pathlib import Path
import unittest
SRC = Path(__file__).resolve().parents[2] / "src" / "cmbuyer_client"
SCOPED = tuple((SRC / name) for name in ("core", "remote", "localstate"))
class StaticBoundaryTests(unittest.TestCase):
def test_scoped_modules_do_not_import_device_pdd_or_unapproved_capabilities(self) -> None:
forbidden_modules = ("cmbuyer_client.device", "cmbuyer_client.pdd")
forbidden_text = (
"ResultSink",
"/events",
"/fail",
"/submission-fence",
"/result",
"click_permitted",
)
for directory in SCOPED:
for path in directory.glob("*.py"):
text = path.read_text(encoding="utf-8")
tree = ast.parse(text)
imports = []
for node in ast.walk(tree):
if isinstance(node, ast.Import):
imports.extend(alias.name for alias in node.names)
elif isinstance(node, ast.ImportFrom) and node.module:
imports.append(node.module)
for module in forbidden_modules:
self.assertFalse(any(name.startswith(module) for name in imports), (path, module))
for value in forbidden_text:
self.assertNotIn(value, text, (path, value))

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