Compare commits
120
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
db7b61b90b | ||
|
|
50a572f3e3 | ||
|
|
1d60a91ace | ||
|
|
665b302a1d | ||
|
|
dbf9b8c902 | ||
|
|
23445d0c4b | ||
|
|
3abbe1040d | ||
|
|
e93e2e8202 | ||
|
|
b0b71d5a29 | ||
|
|
acc9f1d964 | ||
|
|
db2b781c8c | ||
|
|
3f09fbff66 | ||
|
|
21628a4dc5 | ||
|
|
1d726dd364 | ||
|
|
e18f7983fd | ||
|
|
ad80a7f23c | ||
|
|
994b6c8054 | ||
|
|
16b2487a2a | ||
|
|
67b427864a | ||
|
|
1bb54927e1 | ||
|
|
0afe9a53b3 | ||
|
|
63bd873c26 | ||
|
|
2e303f3829 | ||
|
|
f97dca95ae | ||
|
|
43fd8ce668 | ||
|
|
44a207df7f | ||
|
|
29fb528823 | ||
|
|
537159445e | ||
|
|
91071d4013 | ||
|
|
60cf05b08b | ||
|
|
fec21ae857 | ||
|
|
3984f399b0 | ||
|
|
251f3d92dd | ||
|
|
5fa8ac6ff2 | ||
|
|
a8570f59af | ||
|
|
ee56249014 | ||
|
|
e552e2e2ca | ||
|
|
e1e534c5d3 | ||
|
|
5a4633fb64 | ||
|
|
5d0fc81891 | ||
|
|
230cfdc697 | ||
|
|
b4ff9096c3 | ||
|
|
ae4b9ef08b | ||
|
|
1bf3acda70 | ||
|
|
e8ca738b33 | ||
|
|
429c2d33db | ||
|
|
a15b6c0911 | ||
|
|
29d1540f7c | ||
|
|
2f1380d084 | ||
|
|
81a44caeef | ||
|
|
757eca39fd | ||
|
|
b754cc9ca9 | ||
|
|
e45d26fe93 | ||
|
|
a09c58beef | ||
|
|
42eac57824 | ||
|
|
c9b39a4e05 | ||
|
|
870e5f861d | ||
|
|
025edaf273 | ||
|
|
584c5601c5 | ||
|
|
3a27225977 | ||
|
|
dfd88c3336 | ||
|
|
20112dbd61 | ||
|
|
404d0d3ca4 | ||
|
|
85e5f1e56e | ||
|
|
6145fc4468 | ||
|
|
f6cd65208d | ||
|
|
7ea1c5349f | ||
|
|
57a5b2e91c | ||
|
|
63c2c59b41 | ||
|
|
cfb238d022 | ||
|
|
31f07ac245 | ||
|
|
488005ac93 | ||
|
|
4186f1315a | ||
|
|
3a0a41db00 | ||
|
|
8600c33391 | ||
|
|
8092431208 | ||
|
|
cd47d0c959 | ||
|
|
ec42640123 | ||
|
|
5f060f60ee | ||
|
|
19f81a5e58 | ||
|
|
c1cee49fb1 | ||
|
|
772e379fd6 | ||
|
|
83b244ff7c | ||
|
|
d9a31cbfa6 | ||
|
|
014405466a | ||
|
|
8c50e1579e | ||
|
|
35d7ce11d6 | ||
|
|
41e54e313e | ||
|
|
e5de76503f | ||
|
|
542b2283f4 | ||
|
|
eb29bcd8b7 | ||
|
|
526af1eb31 | ||
|
|
d9381a644b | ||
|
|
e863a29641 | ||
|
|
5c97954fa1 | ||
|
|
bd4c855ad2 | ||
|
|
ee90d76893 | ||
|
|
4a65138fd7 | ||
|
|
8ec740dcef | ||
|
|
590c84660a | ||
|
|
233188297a | ||
|
|
e674b7f131 | ||
|
|
66355a7f89 | ||
|
|
41881e81f3 | ||
|
|
e3c87fdec5 | ||
|
|
871cd24d68 | ||
|
|
4f8e71b256 | ||
|
|
a6ad560f5d | ||
|
|
6a5547b323 | ||
|
|
46fccc2120 | ||
|
|
89648880bc | ||
|
|
4829be972c | ||
|
|
3b2fa3536e | ||
|
|
9a4d11f74b | ||
|
|
79284576ef | ||
|
|
20aaba0919 | ||
|
|
03a067e29d | ||
|
|
5dcff4b15a | ||
|
|
ce9d6ca285 | ||
|
|
e379d50101 |
@@ -30,11 +30,11 @@ MVP 只做**任务自带商品链接**的情形:管理员先把任务保存为
|
||||
“开始采购(只创建待付款订单)”。该点击创建一次性授权,锁定商品、颜色、尺码、数量和最高总价。
|
||||
|
||||
采购工具领取后在同一次设备会话中完成:打开商品 → 精确选择规格 → 闸门一读 SKU 单价并校验
|
||||
上限 → 设置并复核数量 → 闸门二重读规格与同价 → 进入确认页 → 闸门三校验应付总额 → 服务端
|
||||
上限 → 设置并复核数量 → 闸门二读取顶部总额 → 在同一最终面板读取闸门三最终控件金额 → 服务端
|
||||
原子建立提交围栏 → 精确点击一次“提交订单” → 转待付款。中间不再回网页端等“机器选对了吗”。
|
||||
|
||||
价格**只在规格面板和订单确认页读**——别处的价格文本被拆成多个节点、带券后前缀、
|
||||
实付价与原价混在一起,不可靠。
|
||||
价格只从拼多多 `8.17.0` 已取证合并式最终提交面板的两个独立角色读取:顶部金额承担 Gate1/Gate2,
|
||||
结构化精确文本 `提交订单 ¥{金额}` 中的金额承担 Gate3;后者必须与 Gate2 严格相等且不能回流兜底。
|
||||
|
||||
## 现在处于什么阶段
|
||||
|
||||
@@ -62,11 +62,11 @@ MVP 只做**任务自带商品链接**的情形:管理员先把任务保存为
|
||||
1. 不点击任何支付、免密支付、先用后付或扣款控件
|
||||
2. **提交订单四条件**:授权未消费且服务端提交围栏已建立 + 闸门二通过 + 闸门三通过 +
|
||||
控件唯一,**只点一次**;围栏后只调和,不释放、不重试
|
||||
3. 订单确认页上除提交与返回外零点击
|
||||
3. 合并式最终提交面板上除满足四条件后的最终控件与安全返回外零点击
|
||||
4. 规格按维度精确匹配,防前缀碰撞,找不到即停
|
||||
5. 数量设置后必须读回复核
|
||||
6. **三道价格闸门**,任一道读不到或不通过即停
|
||||
7. **能力分层**——T-103 规格验证路径不得引用数量、确认页或下单函数;后续能力逐段取证
|
||||
7. **能力分层**——T-103 规格验证路径不得引用数量、最终提交面板或下单函数;后续能力逐段取证
|
||||
8. 检测到外部支付交接立即停止,不读取不保存凭据
|
||||
9. 检测到验证码 / 风控 / 人脸 / 短信校验立即停止,不绕过
|
||||
10. 只读非敏感摘要,不提取收货地址原文、手机号、支付凭据
|
||||
@@ -80,7 +80,7 @@ cmbuyer 是 `cmroubao`(Go 后端 + Android AccessibilityService)与 `cmpdd`
|
||||
(Python + uiautomator2)的合并重启,取各自已验证的一半:
|
||||
|
||||
- 保留 cmroubao 的后端任务生命周期、设备侧 API 形状、下单授权状态机、ERP 对接、管理 Web
|
||||
- 保留 cmpdd 的 uiautomator2 真机自动化、按维度精确选规格、订单确认页读取、付款闸门
|
||||
- 保留 cmpdd 的 uiautomator2 真机自动化与按维度精确选规格思路;最终提交面板事实在本项目重新取证
|
||||
- 丢弃自研 Android APK 与 AccessibilityService 感知层
|
||||
|
||||
**前序项目是设计依据,不是事实来源。** 其中的页面判据必须在本项目用真机重新验证。
|
||||
|
||||
@@ -9,6 +9,12 @@
|
||||
| `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`。 |
|
||||
|
||||
示例仅展示变量名,不提供可运行凭据:
|
||||
|
||||
@@ -18,8 +24,45 @@ $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 被读取之前。
|
||||
|
||||
@@ -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
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -7,12 +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 {
|
||||
@@ -34,12 +40,33 @@ func run() error {
|
||||
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
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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 == "" {
|
||||
|
||||
@@ -2,10 +2,15 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"golang.org/x/crypto/bcrypt"
|
||||
)
|
||||
@@ -16,6 +21,12 @@ const (
|
||||
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
|
||||
)
|
||||
|
||||
@@ -26,6 +37,12 @@ type Config struct {
|
||||
SessionSecret []byte
|
||||
CookieSecure bool
|
||||
DatabaseSource string
|
||||
AuthorizationTTL time.Duration
|
||||
MaxTaskQuantity int
|
||||
MaxTotalPrice string
|
||||
EvidenceDirectory string
|
||||
ClaimTokenSecret []byte
|
||||
ClaimLeaseTTL time.Duration
|
||||
}
|
||||
|
||||
// LoadFromEnv 从进程环境读取配置。错误只指出缺失或非法的变量名,绝不回显秘密。
|
||||
@@ -71,6 +88,57 @@ func Load(lookup func(string) (string, bool)) (Config, error) {
|
||||
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,
|
||||
@@ -78,9 +146,28 @@ func Load(lookup func(string) (string, bool)) (Config, error) {
|
||||
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) == "" {
|
||||
|
||||
@@ -3,6 +3,7 @@ package config_test
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"cmbuyer/admin/internal/config"
|
||||
|
||||
@@ -21,13 +22,19 @@ func TestLoad(t *testing.T) {
|
||||
"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)
|
||||
}
|
||||
}
|
||||
@@ -43,6 +50,12 @@ func TestLoadRejectsMissingOrInvalidConfiguration(t *testing.T) {
|
||||
"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 {
|
||||
@@ -55,6 +68,17 @@ func TestLoadRejectsMissingOrInvalidConfiguration(t *testing.T) {
|
||||
{"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 {
|
||||
|
||||
@@ -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:]
|
||||
}
|
||||
@@ -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")
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -26,18 +26,41 @@ func TestUpDownAndIdempotence(t *testing.T) {
|
||||
if err := migrations.Up(context, database, directory); err != nil {
|
||||
t.Fatalf("apply migrations: %v", err)
|
||||
}
|
||||
assertVersion(t, database, 2)
|
||||
assertVersion(t, database, 5)
|
||||
assertTableExists(t, database, "tasks", 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, 5)
|
||||
|
||||
if err := migrations.Down(context, database, directory); err != nil {
|
||||
t.Fatalf("roll back task claim migration: %v", err)
|
||||
}
|
||||
assertVersion(t, database, 4)
|
||||
assertTableExists(t, database, "purchase_attempt_claims", false)
|
||||
assertTableExists(t, database, "device_credentials", true)
|
||||
|
||||
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)
|
||||
@@ -50,7 +73,7 @@ func TestUpDownAndIdempotence(t *testing.T) {
|
||||
if err := migrations.Up(context, database, directory); err != nil {
|
||||
t.Fatalf("reapply v2 after rollback: %v", err)
|
||||
}
|
||||
assertVersion(t, database, 2)
|
||||
assertVersion(t, database, 5)
|
||||
}
|
||||
|
||||
func TestUpgradePreservesManualDraftLosslessly(t *testing.T) {
|
||||
@@ -69,7 +92,7 @@ func TestUpgradePreservesManualDraftLosslessly(t *testing.T) {
|
||||
if err := migrations.Up(context.Background(), database, migrationDirectory(t)); err != nil {
|
||||
t.Fatalf("upgrade v1 draft: %v", err)
|
||||
}
|
||||
assertVersion(t, database, 2)
|
||||
assertVersion(t, database, 5)
|
||||
var got struct {
|
||||
id, source, sourceRef, title, goodsID, color, size, maxPrice, assetID, status, created, updated string
|
||||
quantity, version int
|
||||
@@ -218,6 +241,246 @@ func TestV2SchemaConstraintsAndRelationships(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
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 := 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)
|
||||
}
|
||||
}
|
||||
|
||||
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")
|
||||
}
|
||||
})
|
||||
}
|
||||
|
||||
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)
|
||||
}
|
||||
}
|
||||
|
||||
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")
|
||||
}
|
||||
|
||||
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")
|
||||
}
|
||||
|
||||
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)
|
||||
})
|
||||
}
|
||||
|
||||
func TestDowngradeRejectsV2BusinessDataAtomically(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
@@ -244,9 +507,7 @@ func TestDowngradeRejectsV2BusinessDataAtomically(t *testing.T) {
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
database := openTestDatabase(t)
|
||||
if err := migrations.Up(context.Background(), database, migrationDirectory(t)); err != nil {
|
||||
t.Fatalf("apply migrations: %v", err)
|
||||
}
|
||||
migrateToV2(t, database)
|
||||
test.setup(t, database)
|
||||
before := v2RowCount(t, database)
|
||||
if err := migrations.Down(context.Background(), database, migrationDirectory(t)); err == nil {
|
||||
@@ -271,6 +532,35 @@ func migrateToV1(t *testing.T, database *sql.DB) {
|
||||
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 {
|
||||
@@ -390,12 +680,14 @@ func assertColumnType(t *testing.T, database *sql.DB, table, column, want string
|
||||
}
|
||||
}
|
||||
|
||||
func TestV2MigrationSQLDoesNotDisableForeignKeys(t *testing.T) {
|
||||
contents, err := os.ReadFile(filepath.Join(migrationDirectory(t), "00002_single_pass_model.sql"))
|
||||
if err != nil {
|
||||
t.Fatalf("read migration: %v", err)
|
||||
}
|
||||
if strings.Contains(strings.ToUpper(string(contents)), "PRAGMA FOREIGN_KEYS = OFF") {
|
||||
t.Fatal("migration disables foreign keys")
|
||||
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)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
+160
-14
@@ -2,13 +2,22 @@
|
||||
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"
|
||||
|
||||
@@ -17,6 +26,7 @@ import (
|
||||
)
|
||||
|
||||
const maxFormBytes = 8 << 10
|
||||
const maxJSONBytes = 64 << 10
|
||||
|
||||
// Options 是路由层需要的安全依赖。凭据由启动配置注入,不能在路由中设置默认值。
|
||||
type Options struct {
|
||||
@@ -24,11 +34,15 @@ type Options struct {
|
||||
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 || options.Tasks == 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")
|
||||
}
|
||||
|
||||
@@ -40,12 +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"})
|
||||
}
|
||||
@@ -55,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()
|
||||
}
|
||||
}
|
||||
@@ -126,14 +219,23 @@ func tasksPage(options Options) gin.HandlerFunc {
|
||||
return
|
||||
}
|
||||
|
||||
drafts, err := options.Tasks.ListDrafts(context.Request.Context())
|
||||
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
|
||||
}
|
||||
data := webui.TasksData{CSRFToken: csrfToken, Drafts: drafts}
|
||||
for _, draft := range drafts {
|
||||
if draft.ID == context.Query("created") {
|
||||
for _, row := range data.Tasks {
|
||||
if row.ID == context.Query("created") {
|
||||
data.Success = true
|
||||
break
|
||||
}
|
||||
@@ -169,6 +271,10 @@ func newTaskPage(options Options) gin.HandlerFunc {
|
||||
}
|
||||
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
|
||||
}
|
||||
@@ -185,24 +291,32 @@ func createTask(options Options) gin.HandlerFunc {
|
||||
}
|
||||
fullPage := requestForm.Get("form_mode") == "full"
|
||||
if !validation.Valid() {
|
||||
drafts, err := options.Tasks.ListDrafts(context.Request.Context())
|
||||
if err != nil {
|
||||
context.Status(http.StatusInternalServerError)
|
||||
data, ok := createErrorData(context, options, fullPage)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
renderTasks(context, http.StatusBadRequest, webui.TasksData{CSRFToken: csrfFor(context, options), Drafts: drafts, Form: form, Errors: validation, OpenForm: !fullPage, FullPage: fullPage, FocusField: firstError(validation)})
|
||||
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"] = "该创建请求已用于另一条任务,请重新打开表单。"
|
||||
drafts, listErr := options.Tasks.ListDrafts(context.Request.Context())
|
||||
if listErr != nil {
|
||||
context.Status(http.StatusInternalServerError)
|
||||
data, ok := createErrorData(context, options, fullPage)
|
||||
if !ok {
|
||||
return
|
||||
}
|
||||
renderTasks(context, http.StatusConflict, webui.TasksData{CSRFToken: csrfFor(context, options), Drafts: drafts, Form: form, Errors: validation, OpenForm: !fullPage, FullPage: fullPage, FocusField: firstError(validation)})
|
||||
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)
|
||||
@@ -226,6 +340,38 @@ 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)
|
||||
|
||||
@@ -2,15 +2,21 @@ 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"
|
||||
@@ -190,7 +196,7 @@ func TestTaskCreationRendersSharedFormsAndPersistsOnlyDraft(t *testing.T) {
|
||||
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" disabled`, `disabled>筛选</button>`, `disabled>清除</button>`, `min-height:44px`, `overflow-x:auto`, `prefers-reduced-motion`} {
|
||||
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)
|
||||
}
|
||||
@@ -277,17 +283,116 @@ func TestTaskCreationRendersSharedFormsAndPersistsOnlyDraft(t *testing.T) {
|
||||
t.Fatalf("task list is missing %q", want)
|
||||
}
|
||||
}
|
||||
for _, forbidden := range []string{"utm_source", "试选", "PENDING", "支付", "订单确认", "真机", "提交订单"} {
|
||||
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+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.StatusForbidden {
|
||||
t.Fatalf("anonymous POST /tasks = %d, want 403", response.Code)
|
||||
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 {
|
||||
@@ -327,7 +432,7 @@ func assertSecurityHeaders(t *testing.T, response *httptest.ResponseRecorder) {
|
||||
"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 {
|
||||
@@ -381,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)
|
||||
@@ -392,7 +509,11 @@ func newRouter(t *testing.T) (*gin.Engine, *auth.Manager) {
|
||||
AdminUsername: "admin",
|
||||
AdminPasswordBcrypt: string(hash),
|
||||
Sessions: manager,
|
||||
Tasks: &memoryStore{},
|
||||
Tasks: store,
|
||||
TaskDetails: details,
|
||||
Evidence: evidenceStore,
|
||||
DeviceAuthenticator: deviceAuthenticator,
|
||||
TaskClaims: claims,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("NewRouter: %v", err)
|
||||
@@ -400,7 +521,42 @@ func newRouter(t *testing.T) (*gin.Engine, *auth.Manager) {
|
||||
return router, manager
|
||||
}
|
||||
|
||||
type memoryStore struct{ drafts []tasks.Draft }
|
||||
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 {
|
||||
@@ -415,8 +571,24 @@ func (store *memoryStore) CreateDraft(_ context.Context, draft tasks.Draft) (tas
|
||||
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
|
||||
|
||||
@@ -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,
|
||||
}},
|
||||
}
|
||||
}
|
||||
@@ -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})
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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{"测试<script>", "订单已创建,系统尚未付款", "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
|
||||
}
|
||||
@@ -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[:])
|
||||
}
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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 ©
|
||||
}
|
||||
|
||||
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')
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -13,31 +13,42 @@ 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
|
||||
createGate chan 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 1 FROM tasks LIMIT 1"); err != nil {
|
||||
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)
|
||||
}
|
||||
return &SQLiteStore{database: database, now: time.Now, createGate: make(chan struct{}, 1)}, nil
|
||||
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.createGate <- struct{}{}:
|
||||
defer func() { <-store.createGate }()
|
||||
case store.writeGate <- struct{}{}:
|
||||
defer func() { <-store.writeGate }()
|
||||
case <-writeContext.Done():
|
||||
return Draft{}, writeContext.Err()
|
||||
}
|
||||
|
||||
@@ -9,14 +9,21 @@ import (
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
)
|
||||
|
||||
const (
|
||||
maxTitleLength = 120
|
||||
maxSKUText = 80
|
||||
MaxTitleCodePoints = 120
|
||||
MaxSKUTextCodePoints = 80
|
||||
MaxGoodsIDCharacters = 32
|
||||
MaxMoneyASCIICharacters = 32
|
||||
maxSKUText = MaxSKUTextCodePoints
|
||||
)
|
||||
|
||||
var ErrCreateKeyConflict = errors.New("create key conflicts with a different task")
|
||||
var (
|
||||
ErrCreateKeyConflict = errors.New("create key conflicts with a different task")
|
||||
ErrInvalidDraft = errors.New("invalid draft")
|
||||
)
|
||||
|
||||
type Draft struct {
|
||||
ID string
|
||||
@@ -41,13 +48,13 @@ func Validate(form Form) (Draft, Errors) {
|
||||
if !validUUID(draft.ID) {
|
||||
errors["create_key"] = "创建请求已过期,请重新打开表单。"
|
||||
}
|
||||
if draft.Title == "" || len([]rune(draft.Title)) > maxTitleLength {
|
||||
if !validBoundedText(draft.Title, MaxTitleCodePoints) {
|
||||
errors["title"] = "任务名称不能为空,且不能超过 120 个字符。"
|
||||
}
|
||||
if draft.SKUColor == "" || len([]rune(draft.SKUColor)) > maxSKUText {
|
||||
if !validBoundedText(draft.SKUColor, MaxSKUTextCodePoints) {
|
||||
errors["sku_color"] = "颜色分类不能为空,且不能超过 80 个字符。"
|
||||
}
|
||||
if draft.SKUSize == "" || len([]rune(draft.SKUSize)) > maxSKUText {
|
||||
if !validBoundedText(draft.SKUSize, MaxSKUTextCodePoints) {
|
||||
errors["sku_size"] = "尺码不能为空,且不能超过 80 个字符。"
|
||||
}
|
||||
goodsID, ok := CanonicalGoodsID(strings.TrimSpace(form.ProductURL))
|
||||
@@ -93,6 +100,9 @@ func CanonicalGoodsID(value string) (string, bool) {
|
||||
return "", false
|
||||
}
|
||||
}
|
||||
if !ValidGoodsID(goodsIDs[0]) {
|
||||
return "", false
|
||||
}
|
||||
return goodsIDs[0], true
|
||||
}
|
||||
|
||||
@@ -155,5 +165,54 @@ func normalizeMoney(value string) (string, bool) {
|
||||
if whole == "0" && strings.Trim(fraction, "0") == "" {
|
||||
return "", false
|
||||
}
|
||||
return whole + "." + (fraction + "00")[:2], true
|
||||
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
|
||||
}
|
||||
|
||||
@@ -7,6 +7,7 @@ import (
|
||||
"path/filepath"
|
||||
"regexp"
|
||||
"runtime"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
@@ -38,13 +39,16 @@ func TestValidateNormalizesManualDraft(t *testing.T) {
|
||||
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 = " " },
|
||||
"long color": func(form *Form) { form.SKUColor = string(make([]rune, maxSKUText+1)) },
|
||||
"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" },
|
||||
"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
|
||||
@@ -67,6 +71,7 @@ func TestValidateRejectsInvalidFieldsAndURLs(t *testing.T) {
|
||||
"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)
|
||||
@@ -75,19 +80,62 @@ func TestValidateRejectsInvalidFieldsAndURLs(t *testing.T) {
|
||||
}
|
||||
|
||||
func TestNormalizeMoneyBoundaries(t *testing.T) {
|
||||
for value, want := range map[string]string{"1": "1.00", "1.2": "1.20", "000.01": "0.01", "999999999999999999": "999999999999999999.00"} {
|
||||
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"} {
|
||||
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 {
|
||||
@@ -157,6 +205,37 @@ func TestSQLiteStoreCreatesListsAndHandlesIdempotency(t *testing.T) {
|
||||
}
|
||||
}
|
||||
|
||||
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)
|
||||
|
||||
@@ -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
@@ -3,16 +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").Funcs(template.FuncMap{"list": func(values ...any) []any { return values }}).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 {
|
||||
@@ -22,18 +43,24 @@ type LoginData struct {
|
||||
Error string
|
||||
}
|
||||
|
||||
// TasksData 是受保护的 DRAFT 建单与列表页面所需数据。
|
||||
// TasksData 是受保护的建单与任务工作台页面所需数据。
|
||||
type TasksData struct {
|
||||
CSRFToken string
|
||||
Drafts []tasks.Draft
|
||||
Form tasks.Form
|
||||
Errors tasks.Errors
|
||||
OpenForm bool
|
||||
FullPage bool
|
||||
FocusField string
|
||||
Success bool
|
||||
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)
|
||||
@@ -43,3 +70,93 @@ func RenderLogin(writer io.Writer, data LoginData) error {
|
||||
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,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;
|
||||
@@ -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,109 @@
|
||||
"""采集 T-106 人工准备的最终面板或返回态证据;不执行页面操作。"""
|
||||
|
||||
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.order_confirm_spike import (
|
||||
Android16ForegroundReader,
|
||||
DECLARED_STATES,
|
||||
EXPECTED_GOODS_ID,
|
||||
OrderConfirmEvidenceCapturer,
|
||||
OrderConfirmEvidenceError,
|
||||
)
|
||||
|
||||
|
||||
def parse_arguments(argv: list[str] | None = None) -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="采集 T-106 人工准备的最终面板或返回态本机证据。")
|
||||
parser.add_argument("--serial", required=True, help="ADB device serial;禁止自动选择。")
|
||||
parser.add_argument("--goods-id", required=True, help="T-106 已批准的目标商品标识。")
|
||||
parser.add_argument(
|
||||
"--state",
|
||||
required=True,
|
||||
choices=DECLARED_STATES,
|
||||
help="只允许既有提交前来源态,或人工只按一次返回后的安全页状态。",
|
||||
)
|
||||
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 (
|
||||
type(arguments.serial) is not str
|
||||
or not arguments.serial.strip()
|
||||
or arguments.serial != 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-106 已批准目标。")
|
||||
if type(arguments.state) is not str or arguments.state not in DECLARED_STATES:
|
||||
raise ValueError("--state 必须是批准的人工声明状态。")
|
||||
if not isinstance(arguments.output_dir, Path) or not arguments.output_dir.name:
|
||||
raise ValueError("--output-dir 必须是明确的全新目录。")
|
||||
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
|
||||
|
||||
adb_runner = SubprocessAdbRunner(arguments.adb)
|
||||
capturer = OrderConfirmEvidenceCapturer(
|
||||
AdbClient(adb_runner, timeout_seconds=arguments.timeout),
|
||||
NoReconnectUiautomatorConnector(
|
||||
adbutils.AdbClient(socket_timeout=arguments.timeout).device_list,
|
||||
u2.connect,
|
||||
),
|
||||
Android16ForegroundReader(adb_runner, arguments.timeout),
|
||||
timeout_seconds=arguments.timeout,
|
||||
)
|
||||
try:
|
||||
capturer.capture(
|
||||
arguments.serial,
|
||||
arguments.goods_id,
|
||||
arguments.state,
|
||||
arguments.output_dir,
|
||||
)
|
||||
except (DeviceConnectionError, OrderConfirmEvidenceError):
|
||||
# 第三方异常可能带设备、页面正文或本机目录,命令行只输出固定摘要。
|
||||
print("T-106 只读取证失败:已停止,未发布本机证据目录。", file=sys.stderr)
|
||||
return 1
|
||||
except OSError:
|
||||
print("T-106 只读取证失败:无法发布本机证据目录。", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
print("T-106 只读取证完成。")
|
||||
print("人工复核:请在指定目录检查截图、XML、应用摘要和 manifest。")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,120 @@
|
||||
"""采集 T-105 人工准备的数量 1/2 两态证据;不执行页面操作。"""
|
||||
|
||||
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.quantity_gate2_spike import (
|
||||
Android16TopResumedForegroundReader,
|
||||
DECLARED_QUANTITIES,
|
||||
EXPECTED_GOODS_ID,
|
||||
QuantityGate2EvidenceCapturer,
|
||||
QuantityGate2EvidenceError,
|
||||
)
|
||||
|
||||
|
||||
def parse_arguments(argv: list[str] | None = None) -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="采集 T-105 人工准备的数量两态本机证据。")
|
||||
parser.add_argument("--serial", required=True, help="ADB device serial;禁止自动选择。")
|
||||
parser.add_argument("--goods-id", required=True, help="T-105 已批准的目标商品标识。")
|
||||
parser.add_argument(
|
||||
"--state",
|
||||
required=True,
|
||||
choices=sorted(DECLARED_QUANTITIES),
|
||||
help="人工声明状态:initial=数量1,target=数量2。",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--declared-quantity",
|
||||
required=True,
|
||||
type=int,
|
||||
help="人工看到的数量;必须与状态一致。",
|
||||
)
|
||||
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 (
|
||||
type(arguments.serial) is not str
|
||||
or not arguments.serial.strip()
|
||||
or arguments.serial != 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-105 已批准目标。")
|
||||
if type(arguments.state) is not str or arguments.state not in DECLARED_QUANTITIES:
|
||||
raise ValueError("--state 必须是批准的人工声明状态。")
|
||||
if type(arguments.declared_quantity) is not int or (
|
||||
arguments.declared_quantity != DECLARED_QUANTITIES[arguments.state]
|
||||
):
|
||||
raise ValueError("--declared-quantity 必须与人工声明状态一致。")
|
||||
if not isinstance(arguments.output_dir, Path) or not arguments.output_dir.name:
|
||||
raise ValueError("--output-dir 必须是明确的全新目录。")
|
||||
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
|
||||
|
||||
adb_runner = SubprocessAdbRunner(arguments.adb)
|
||||
capturer = QuantityGate2EvidenceCapturer(
|
||||
AdbClient(adb_runner, timeout_seconds=arguments.timeout),
|
||||
NoReconnectUiautomatorConnector(
|
||||
adbutils.AdbClient(socket_timeout=arguments.timeout).device_list,
|
||||
u2.connect,
|
||||
),
|
||||
Android16TopResumedForegroundReader(adb_runner, arguments.timeout),
|
||||
timeout_seconds=arguments.timeout,
|
||||
)
|
||||
try:
|
||||
capturer.capture(
|
||||
arguments.serial,
|
||||
arguments.goods_id,
|
||||
arguments.state,
|
||||
arguments.declared_quantity,
|
||||
arguments.output_dir,
|
||||
)
|
||||
except (DeviceConnectionError, QuantityGate2EvidenceError):
|
||||
# 第三方异常可能含 serial、Activity、本机路径或页面正文,CLI 只输出固定摘要。
|
||||
print("数量两态只读取证失败:已停止,未发布本机证据目录。", file=sys.stderr)
|
||||
return 1
|
||||
except OSError:
|
||||
print("数量两态只读取证失败:无法发布本机证据目录。", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
print("数量两态只读取证完成。")
|
||||
print("人工复核:请在指定目录检查截图、XML、应用摘要和 manifest。")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,85 @@
|
||||
"""采集 T-104 阶段 A 的一次 Back 后本机证据。"""
|
||||
|
||||
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_selection import SkuSelectionError
|
||||
from cmbuyer_client.pdd.sku_selection_runner import (
|
||||
SkuExitSpikeCapturer,
|
||||
SkuSelectionRunError,
|
||||
)
|
||||
|
||||
|
||||
def parse_arguments(argv: list[str] | None = None) -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="采集 T-104 阶段 A 的一次 Back 后本机证据。")
|
||||
parser.add_argument("--serial", required=True, help="ADB device serial;禁止自动选择。")
|
||||
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 type(arguments.serial) is not str or not arguments.serial.strip():
|
||||
raise ValueError("必须显式提供非空 --serial。")
|
||||
if not isinstance(arguments.output_dir, Path) or not arguments.output_dir.name:
|
||||
raise ValueError("--output-dir 必须是明确的全新目录。")
|
||||
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 = SkuExitSpikeCapturer(
|
||||
AdbClient(SubprocessAdbRunner(arguments.adb), timeout_seconds=arguments.timeout),
|
||||
NoReconnectUiautomatorConnector(
|
||||
adbutils.AdbClient(socket_timeout=arguments.timeout).device_list,
|
||||
u2.connect,
|
||||
),
|
||||
timeout_seconds=arguments.timeout,
|
||||
)
|
||||
try:
|
||||
capturer.capture(arguments.serial, arguments.output_dir)
|
||||
except (DeviceConnectionError, SkuSelectionError, SkuSelectionRunError):
|
||||
# 不回显第三方异常、serial、页面正文或本机路径。
|
||||
print("规格安全退出取证失败:已停止,未发布本地证据目录。", file=sys.stderr)
|
||||
return 1
|
||||
except OSError:
|
||||
print("规格安全退出取证失败:无法发布本地证据目录。", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
print("规格安全退出取证完成。")
|
||||
print("人工复核:请在指定目录检查退出后截图、XML、应用摘要和 manifest。")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,108 @@
|
||||
"""采集 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 (
|
||||
safe_reveal_failure_stage,
|
||||
SkuRevealSpikeCapturer,
|
||||
SkuRevealSpikeError,
|
||||
)
|
||||
from cmbuyer_client.pdd.sku_selection import (
|
||||
EXPECTED_GOODS_ID,
|
||||
_safe_sku_entry_failure_stage,
|
||||
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 _safe_capture_failure_stage(error: BaseException) -> str:
|
||||
# 入口与 reveal 各有不可伪造的正式 marker;入口优先,不能被外层 reveal 标注覆盖。
|
||||
entry_stage = _safe_sku_entry_failure_stage(error)
|
||||
if entry_stage is not None:
|
||||
return entry_stage
|
||||
return safe_reveal_failure_stage(error) or "unknown"
|
||||
|
||||
|
||||
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) as error:
|
||||
# 不回显页面正文、节点、serial、坐标、路径或第三方异常。
|
||||
print(
|
||||
f"规格 reveal 取证失败:stage={_safe_capture_failure_stage(error)};已停止,未发布本地证据目录。",
|
||||
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())
|
||||
@@ -15,7 +15,7 @@ from cmbuyer_client.device.adb import AdbClient, DeviceConnectionError, Subproce
|
||||
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
|
||||
from cmbuyer_client.pdd.sku_selection_runner import SkuSelectionRunError, SkuSelectionRunner, safe_failure_stage
|
||||
|
||||
|
||||
def parse_arguments(argv: list[str] | None = None) -> argparse.Namespace:
|
||||
@@ -66,7 +66,10 @@ def main(argv: list[str] | None = None) -> int:
|
||||
result = runner.run(arguments.serial, arguments.url, arguments.color, arguments.size, arguments.output_dir)
|
||||
except (DeviceConnectionError, SkuSelectionRunError, SkuSelectionError) as error:
|
||||
# Flow 可能来自测试替身或未来实现;CLI 不回显任何异常正文,避免泄露节点树或页面文本。
|
||||
print("规格恢复失败:已停止,未发布本地证据目录。", file=sys.stderr)
|
||||
print(
|
||||
f"规格恢复失败:stage={safe_failure_stage(error)};已停止,未发布本地证据目录。",
|
||||
file=sys.stderr,
|
||||
)
|
||||
return 1
|
||||
except OSError:
|
||||
print("规格恢复失败:无法创建或发布本地证据目录。", file=sys.stderr)
|
||||
|
||||
@@ -0,0 +1,132 @@
|
||||
"""运行 T-105 数量 1→2、Gate2 截图与一次安全退出。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
from datetime import datetime
|
||||
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.quantity_gate2 import (
|
||||
EXPECTED_GATE1_UNIT_PRICE,
|
||||
EXPECTED_GOODS_ID,
|
||||
TASK_COLOR,
|
||||
TASK_SIZE,
|
||||
Gate1Observation,
|
||||
QuantityGate2Error,
|
||||
)
|
||||
from cmbuyer_client.pdd.quantity_gate2_runner import QuantityGate2Runner
|
||||
from cmbuyer_client.pdd.quantity_gate2_spike import Android16TopResumedForegroundReader
|
||||
|
||||
|
||||
def parse_arguments(argv: list[str] | None = None) -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="运行 T-105 已取证数量与 Gate2 闭环。")
|
||||
parser.add_argument("--serial", required=True, help="显式 ADB serial;禁止自动选择。")
|
||||
parser.add_argument("--goods-id", required=True)
|
||||
parser.add_argument("--color", required=True)
|
||||
parser.add_argument("--size", required=True)
|
||||
parser.add_argument("--target-quantity", required=True, type=int)
|
||||
parser.add_argument("--gate1-unit-price", required=True)
|
||||
parser.add_argument("--gate1-screenshot", required=True, type=Path)
|
||||
parser.add_argument("--gate1-captured-at", required=True)
|
||||
parser.add_argument("--max-total-price", required=True)
|
||||
parser.add_argument("--output-dir", required=True, type=Path)
|
||||
parser.add_argument("--timeout", type=float, default=10.0)
|
||||
parser.add_argument("--adb", default="adb")
|
||||
return parser.parse_args(argv)
|
||||
|
||||
|
||||
def validate_arguments(arguments: argparse.Namespace) -> datetime:
|
||||
if type(arguments.serial) is not str or not arguments.serial.strip() or arguments.serial != arguments.serial.strip():
|
||||
raise ValueError("必须显式提供非空 --serial。")
|
||||
if arguments.goods_id != EXPECTED_GOODS_ID:
|
||||
raise ValueError("--goods-id 不是 T-105 已批准目标。")
|
||||
if arguments.color != TASK_COLOR or arguments.size != TASK_SIZE:
|
||||
raise ValueError("颜色或尺码不是 T-105 已批准目标。")
|
||||
if type(arguments.target_quantity) is not int or arguments.target_quantity != 2:
|
||||
raise ValueError("本次真机验收只批准 --target-quantity 2。")
|
||||
if arguments.gate1_unit_price != EXPECTED_GATE1_UNIT_PRICE:
|
||||
raise ValueError("--gate1-unit-price 与已确认 Gate1 不一致。")
|
||||
if not isinstance(arguments.gate1_screenshot, Path) or not arguments.gate1_screenshot.is_file():
|
||||
raise ValueError("--gate1-screenshot 必须是现有显式文件。")
|
||||
try:
|
||||
captured_at = datetime.fromisoformat(arguments.gate1_captured_at)
|
||||
except (TypeError, ValueError) as error:
|
||||
raise ValueError("--gate1-captured-at 必须是带时区 ISO 时间。") from error
|
||||
if captured_at.utcoffset() is None:
|
||||
raise ValueError("--gate1-captured-at 必须带时区。")
|
||||
if arguments.max_total_price != "40.00":
|
||||
raise ValueError("本次真机验收固定 --max-total-price 40.00。")
|
||||
if not isinstance(arguments.output_dir, Path) or not arguments.output_dir.name or arguments.output_dir.exists():
|
||||
raise ValueError("--output-dir 必须是不存在的明确新目录。")
|
||||
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 的有限数值。")
|
||||
return captured_at
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
arguments = parse_arguments(argv)
|
||||
try:
|
||||
captured_at = 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("失败:缺少采购工具真机依赖。", file=sys.stderr)
|
||||
return 2
|
||||
|
||||
adb_runner = SubprocessAdbRunner(arguments.adb)
|
||||
runner = QuantityGate2Runner(
|
||||
AdbClient(adb_runner, timeout_seconds=arguments.timeout),
|
||||
NoReconnectUiautomatorConnector(
|
||||
adbutils.AdbClient(socket_timeout=arguments.timeout).device_list,
|
||||
u2.connect,
|
||||
),
|
||||
Android16TopResumedForegroundReader(adb_runner, arguments.timeout),
|
||||
timeout_seconds=arguments.timeout,
|
||||
)
|
||||
gate1 = Gate1Observation(
|
||||
color=arguments.color,
|
||||
size=arguments.size,
|
||||
quantity=1,
|
||||
gate1_unit_price=arguments.gate1_unit_price,
|
||||
screenshot_path=arguments.gate1_screenshot,
|
||||
captured_at=captured_at,
|
||||
)
|
||||
try:
|
||||
runner.run(
|
||||
arguments.serial,
|
||||
arguments.goods_id,
|
||||
gate1,
|
||||
arguments.target_quantity,
|
||||
arguments.max_total_price,
|
||||
arguments.output_dir,
|
||||
)
|
||||
except (DeviceConnectionError, QuantityGate2Error, OSError):
|
||||
# 不回显第三方异常、serial、本机路径或页面正文。
|
||||
print("T-105 数量/Gate2 运行失败:已停止,未发布证据目录。", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
print("T-105 数量/Gate2 运行完成。")
|
||||
print("人工复核:数量 2、目标规格、12.88/32.76、原始截图和一次安全退出。")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -0,0 +1,138 @@
|
||||
"""运行 T-107 最终面板 Gate3 只读观察与一次安全返回。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import argparse
|
||||
from datetime import datetime
|
||||
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.final_submit_panel import FinalSubmitPanelError
|
||||
from cmbuyer_client.pdd.final_submit_panel_runner import FinalSubmitPanelRunner
|
||||
from cmbuyer_client.pdd.quantity_gate2 import (
|
||||
EXPECTED_GATE1_UNIT_PRICE,
|
||||
EXPECTED_GOODS_ID,
|
||||
TASK_COLOR,
|
||||
TASK_SIZE,
|
||||
Gate2Observation,
|
||||
)
|
||||
from cmbuyer_client.pdd.quantity_gate2_spike import Android16TopResumedForegroundReader
|
||||
|
||||
|
||||
def parse_arguments(argv: list[str] | None = None) -> argparse.Namespace:
|
||||
parser = argparse.ArgumentParser(description="运行 T-107 围栏前最终面板只读 dry-run。")
|
||||
parser.add_argument("--serial", required=True, help="显式 ADB serial;禁止自动选择。")
|
||||
parser.add_argument("--goods-id", required=True)
|
||||
parser.add_argument("--color", required=True)
|
||||
parser.add_argument("--size", required=True)
|
||||
parser.add_argument("--quantity", required=True, type=int)
|
||||
parser.add_argument("--gate1-unit-price", required=True)
|
||||
parser.add_argument("--gate2-panel-total-price", required=True)
|
||||
parser.add_argument("--gate2-screenshot", required=True, type=Path)
|
||||
parser.add_argument("--gate2-captured-at", required=True)
|
||||
parser.add_argument("--max-total-price", required=True)
|
||||
parser.add_argument("--output-dir", required=True, type=Path)
|
||||
parser.add_argument("--timeout", type=float, default=10.0)
|
||||
parser.add_argument("--adb", default="adb")
|
||||
return parser.parse_args(argv)
|
||||
|
||||
|
||||
def validate_arguments(arguments: argparse.Namespace) -> datetime:
|
||||
if type(arguments.serial) is not str or not arguments.serial.strip() or arguments.serial != arguments.serial.strip():
|
||||
raise ValueError("必须显式提供非空 --serial。")
|
||||
if arguments.goods_id != EXPECTED_GOODS_ID:
|
||||
raise ValueError("--goods-id 不是 T-107 已批准目标。")
|
||||
if arguments.color != TASK_COLOR or arguments.size != TASK_SIZE:
|
||||
raise ValueError("颜色或尺码不是 T-107 已批准目标。")
|
||||
if type(arguments.quantity) is not int or arguments.quantity != 2:
|
||||
raise ValueError("本次真机验收只批准 --quantity 2。")
|
||||
if arguments.gate1_unit_price != EXPECTED_GATE1_UNIT_PRICE:
|
||||
raise ValueError("--gate1-unit-price 与已确认 Gate1 不一致。")
|
||||
if arguments.gate2_panel_total_price != "32.76":
|
||||
raise ValueError("--gate2-panel-total-price 与 T-106 已确认值不一致。")
|
||||
if not isinstance(arguments.gate2_screenshot, Path) or not arguments.gate2_screenshot.is_file():
|
||||
raise ValueError("--gate2-screenshot 必须是现有显式文件。")
|
||||
try:
|
||||
captured_at = datetime.fromisoformat(arguments.gate2_captured_at)
|
||||
except (TypeError, ValueError) as error:
|
||||
raise ValueError("--gate2-captured-at 必须是带时区 ISO 时间。") from error
|
||||
if captured_at.utcoffset() is None:
|
||||
raise ValueError("--gate2-captured-at 必须带时区。")
|
||||
if arguments.max_total_price != "40.00":
|
||||
raise ValueError("本次真机验收固定 --max-total-price 40.00。")
|
||||
if not isinstance(arguments.output_dir, Path) or not arguments.output_dir.name or arguments.output_dir.exists():
|
||||
raise ValueError("--output-dir 必须是不存在的明确新目录。")
|
||||
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 的有限数值。")
|
||||
return captured_at
|
||||
|
||||
|
||||
def main(argv: list[str] | None = None) -> int:
|
||||
arguments = parse_arguments(argv)
|
||||
try:
|
||||
captured_at = 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("失败:缺少采购工具真机依赖。", file=sys.stderr)
|
||||
return 2
|
||||
|
||||
adb_runner = SubprocessAdbRunner(arguments.adb)
|
||||
runner = FinalSubmitPanelRunner(
|
||||
AdbClient(adb_runner, timeout_seconds=arguments.timeout),
|
||||
NoReconnectUiautomatorConnector(
|
||||
adbutils.AdbClient(socket_timeout=arguments.timeout).device_list,
|
||||
u2.connect,
|
||||
),
|
||||
Android16TopResumedForegroundReader(adb_runner, arguments.timeout),
|
||||
timeout_seconds=arguments.timeout,
|
||||
)
|
||||
gate2 = Gate2Observation(
|
||||
requested_color=arguments.color,
|
||||
requested_size=arguments.size,
|
||||
actual_color=arguments.color,
|
||||
actual_size=arguments.size,
|
||||
requested_quantity=arguments.quantity,
|
||||
quantity_read=arguments.quantity,
|
||||
gate1_unit_price=arguments.gate1_unit_price,
|
||||
gate2_panel_total_price=arguments.gate2_panel_total_price,
|
||||
max_total_price=arguments.max_total_price,
|
||||
screenshot_path=arguments.gate2_screenshot,
|
||||
captured_at=captured_at,
|
||||
)
|
||||
try:
|
||||
runner.run(
|
||||
arguments.serial,
|
||||
arguments.goods_id,
|
||||
gate2,
|
||||
arguments.output_dir,
|
||||
)
|
||||
except (DeviceConnectionError, FinalSubmitPanelError, OSError):
|
||||
# 不回显第三方异常、serial、本机路径或页面正文。
|
||||
print("T-107 最终面板 dry-run 失败:已停止,未发布证据目录。", file=sys.stderr)
|
||||
return 1
|
||||
|
||||
print("T-107 最终面板 dry-run 完成。")
|
||||
print("人工复核:Gate2/Gate3、最终控件唯一、一次安全返回,且未创建订单或进入付款。")
|
||||
return 0
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
raise SystemExit(main())
|
||||
@@ -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("应用已启动;采购执行功能尚未启用。")
|
||||
window.show()
|
||||
return application.exec()
|
||||
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",
|
||||
]
|
||||
@@ -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):
|
||||
"""服务端要求人工处理的确定性冲突。"""
|
||||
@@ -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=[已隐藏])"
|
||||
@@ -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
|
||||
@@ -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
@@ -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
|
||||
|
||||
@@ -0,0 +1,518 @@
|
||||
"""T-107:合并式最终提交面板的纯 XML Gate3 observer。
|
||||
|
||||
本模块不持有设备,也不返回节点、坐标、selector 或可操作对象。页面结构只来自
|
||||
T-106 在拼多多 8.17.0 / goods_id 937122477375 上由人确认的两态证据。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
from pathlib import Path
|
||||
import re
|
||||
from typing import Iterable
|
||||
from xml.etree import ElementTree
|
||||
|
||||
from ..core.errors import ValidationError
|
||||
from ..core.validation import require_money
|
||||
from ..device.baseline import PDD_PACKAGE
|
||||
from .quantity_gate2 import (
|
||||
TASK_COLOR,
|
||||
TASK_SIZE,
|
||||
UI_COLOR,
|
||||
UI_SIZE,
|
||||
Gate2Observation,
|
||||
)
|
||||
|
||||
|
||||
FINAL_PANEL_QUANTITY = 2
|
||||
|
||||
_PDD_ID = "com.xunmeng.pinduoduo:id/pdd"
|
||||
_QUANTITY_ID = "com.xunmeng.pinduoduo:id/gnl"
|
||||
_TITLE_ID = "com.xunmeng.pinduoduo:id/tv_title"
|
||||
_SUBMIT_TEXT_BOUNDS = "[366,2225][714,2284]"
|
||||
_SUBMIT_TEXT_PARENT_BOUNDS = "[354,2181][726,2328]"
|
||||
_SUBMIT_ACTION_BOUNDS = "[0,2181][1080,2328]"
|
||||
_EXPECTED_SUMMARY = f"已选: {UI_COLOR} {UI_SIZE}"
|
||||
|
||||
_GATE2_TEXT = re.compile(r"^快卖完 ¥(?P<amount>(?:0|[1-9]\d*)\.\d{2})$")
|
||||
_GATE3_TEXT = re.compile(r"^提交订单 ¥(?P<amount>(?:0|[1-9]\d*)\.\d{2})$")
|
||||
|
||||
_PRODUCT_TITLE = "2026年新款高档重工潮流烫钻中长款T恤显瘦宽松上衣淡人穿搭"
|
||||
_PRODUCT_TITLE_BOUNDS = "[36,1571][1044,1689]"
|
||||
_PRODUCT_TITLE_LINES = (
|
||||
("2026年新款高档重工潮流烫钻中长款T恤显瘦宽松", "[36,1571][1023,1624]"),
|
||||
("上衣淡人穿搭", "[36,1636][306,1689]"),
|
||||
)
|
||||
_ENTRY_DESC = "快要抢光¥12.88"
|
||||
_ENTRY_BOUNDS = "[446,2166][1080,2328]"
|
||||
_ENTRY_CHILDREN = (
|
||||
("快要抢光 ¥ 12.88", "[688,2184][1042,2253]"),
|
||||
("免拼购买", "[688,2256][856,2305]"),
|
||||
)
|
||||
_DANGEROUS_RETURN_TERMS = ("提交订单", "确认订单", "立即支付", "去支付", "付款")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _PanelProfile:
|
||||
evidence_id: str
|
||||
panel_bounds: str
|
||||
price_frame_bounds: str
|
||||
price_row_bounds: str
|
||||
price_text_bounds: str
|
||||
summary_bounds: str
|
||||
quantity_bounds: str
|
||||
minus_bounds: str
|
||||
value_bounds: str
|
||||
plus_bounds: str
|
||||
|
||||
|
||||
# 两个配置分别绑定 T-106 原始证据和 T-107 同版本人工五项确认的新证据。
|
||||
# 只能整套精确命中其中一个,禁止把 77px 位移改成范围或混搭坐标。
|
||||
_PANEL_PROFILES = (
|
||||
_PanelProfile(
|
||||
evidence_id="t106-final-panel-8.17.0",
|
||||
panel_bounds="[0,551][1080,938]",
|
||||
price_frame_bounds="[396,575][1053,647]",
|
||||
price_row_bounds="[396,575][740,647]",
|
||||
price_text_bounds="[396,580][722,647]",
|
||||
summary_bounds="[396,659][1053,721]",
|
||||
quantity_bounds="[396,827][645,902]",
|
||||
minus_bounds="[396,827][474,902]",
|
||||
value_bounds="[480,827][561,902]",
|
||||
plus_bounds="[567,827][645,902]",
|
||||
),
|
||||
_PanelProfile(
|
||||
evidence_id="t107-final-panel-shifted-8.17.0",
|
||||
panel_bounds="[0,474][1080,861]",
|
||||
price_frame_bounds="[396,498][1053,570]",
|
||||
price_row_bounds="[396,498][740,570]",
|
||||
price_text_bounds="[396,503][722,570]",
|
||||
summary_bounds="[396,582][1053,644]",
|
||||
quantity_bounds="[396,750][645,825]",
|
||||
minus_bounds="[396,750][474,825]",
|
||||
value_bounds="[480,750][561,825]",
|
||||
plus_bounds="[567,750][645,825]",
|
||||
),
|
||||
)
|
||||
|
||||
|
||||
class FinalSubmitPanelError(RuntimeError):
|
||||
"""最终面板或 Gate3 判据不成立时的脱敏安全停止。"""
|
||||
|
||||
|
||||
class FinalSubmitPanelOverCapError(FinalSubmitPanelError):
|
||||
"""Gate2/Gate3 金额超过管理员授权最高总价。"""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Gate3Observation:
|
||||
requested_color: str
|
||||
requested_size: str
|
||||
actual_color: str
|
||||
actual_size: str
|
||||
requested_quantity: int
|
||||
quantity_read: int
|
||||
gate2_panel_total_price: str
|
||||
gate3_submit_amount: str
|
||||
max_total_price: str
|
||||
submit_control_text: str
|
||||
submit_control_match_count: int
|
||||
submit_control_enabled: bool
|
||||
nearest_clickable_ancestor_unique: bool
|
||||
screenshot_path: Path
|
||||
captured_at: datetime
|
||||
|
||||
|
||||
@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", "")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _VerifiedPanel:
|
||||
gate2_amount: str
|
||||
gate3_amount: str
|
||||
submit_text: str
|
||||
submit_match_count: int
|
||||
submit_enabled: bool
|
||||
nearest_clickable_ancestor_unique: bool
|
||||
|
||||
|
||||
def observe_gate3(
|
||||
raw_hierarchy: str,
|
||||
gate2: Gate2Observation,
|
||||
screenshot_path: Path,
|
||||
captured_at: datetime,
|
||||
) -> Gate3Observation:
|
||||
"""从一棵已读取 XML 中构造不可操作的 Gate3 摘要。"""
|
||||
|
||||
_validate_gate2(gate2)
|
||||
if not isinstance(screenshot_path, Path) or not screenshot_path.name:
|
||||
raise FinalSubmitPanelError("Gate3 截图路径无效,已停止操作。")
|
||||
if not isinstance(captured_at, datetime) or captured_at.utcoffset() is None:
|
||||
raise FinalSubmitPanelError("Gate3 采集时间必须带时区,已停止操作。")
|
||||
|
||||
verified = _verified_final_panel(_parse_nodes(raw_hierarchy))
|
||||
if verified.gate2_amount != gate2.gate2_panel_total_price:
|
||||
raise FinalSubmitPanelError("当前面板 Gate2 顶部总额与可信观察不一致,已停止操作。")
|
||||
if verified.gate3_amount != verified.gate2_amount:
|
||||
raise FinalSubmitPanelError("Gate3 最终控件金额与 Gate2 顶部总额不一致,已停止操作。")
|
||||
if Decimal(verified.gate3_amount) > Decimal(gate2.max_total_price):
|
||||
raise FinalSubmitPanelOverCapError("Gate3 金额超过授权最高总价,已停止操作。")
|
||||
|
||||
return Gate3Observation(
|
||||
requested_color=gate2.requested_color,
|
||||
requested_size=gate2.requested_size,
|
||||
actual_color=gate2.actual_color,
|
||||
actual_size=gate2.actual_size,
|
||||
requested_quantity=gate2.requested_quantity,
|
||||
quantity_read=gate2.quantity_read,
|
||||
gate2_panel_total_price=verified.gate2_amount,
|
||||
gate3_submit_amount=verified.gate3_amount,
|
||||
max_total_price=gate2.max_total_price,
|
||||
submit_control_text=verified.submit_text,
|
||||
submit_control_match_count=verified.submit_match_count,
|
||||
submit_control_enabled=verified.submit_enabled,
|
||||
nearest_clickable_ancestor_unique=verified.nearest_clickable_ancestor_unique,
|
||||
screenshot_path=screenshot_path,
|
||||
captured_at=captured_at,
|
||||
)
|
||||
|
||||
|
||||
def returned_product_projection(raw_hierarchy: str) -> tuple[object, ...]:
|
||||
"""验证 T-106 一次 Back 后的同商品安全页;入口数字只作页面身份,不读价。"""
|
||||
|
||||
nodes = _parse_nodes(raw_hierarchy)
|
||||
if any(
|
||||
_live_clickable(node)
|
||||
and any(term in value for value in (node.text, node.desc) for term in _DANGEROUS_RETURN_TERMS)
|
||||
for node in nodes
|
||||
):
|
||||
raise FinalSubmitPanelError("返回后出现危险动作语义,不能确认安全退出。")
|
||||
|
||||
title = _one(
|
||||
node
|
||||
for node in nodes
|
||||
if _exact(
|
||||
node,
|
||||
"android.view.ViewGroup",
|
||||
_PRODUCT_TITLE_BOUNDS,
|
||||
resource_id=_TITLE_ID,
|
||||
clickable="true",
|
||||
)
|
||||
and node.element.get("long-clickable") == "true"
|
||||
and not node.text
|
||||
and node.desc == _PRODUCT_TITLE
|
||||
)
|
||||
title_children = _direct_children(title)
|
||||
if len(title_children) != len(_PRODUCT_TITLE_LINES) or any(
|
||||
not _exact_text(child, text, bounds)
|
||||
for child, (text, bounds) in zip(title_children, _PRODUCT_TITLE_LINES, strict=True)
|
||||
):
|
||||
raise FinalSubmitPanelError("返回后同商品标题子结构漂移。")
|
||||
|
||||
entry = _one(
|
||||
node
|
||||
for node in nodes
|
||||
if _exact(
|
||||
node,
|
||||
"android.view.ViewGroup",
|
||||
_ENTRY_BOUNDS,
|
||||
resource_id=_PDD_ID,
|
||||
clickable="true",
|
||||
)
|
||||
and not node.text
|
||||
and node.desc == _ENTRY_DESC
|
||||
)
|
||||
all_entry_children = _direct_children(entry)
|
||||
entry_children = [child for child in all_entry_children if child.text]
|
||||
if len(entry_children) != len(_ENTRY_CHILDREN) or any(
|
||||
not _exact_text(child, text, bounds, resource_id=_PDD_ID)
|
||||
for child, (text, bounds) in zip(entry_children, _ENTRY_CHILDREN, strict=True)
|
||||
) or any(
|
||||
child not in entry_children
|
||||
and (child.text or child.desc or _live_clickable(child))
|
||||
for child in all_entry_children
|
||||
):
|
||||
raise FinalSubmitPanelError("返回后商品规格入口结构漂移。")
|
||||
|
||||
return (
|
||||
"final_submit_panel_exit_8_17_0",
|
||||
_projection(title),
|
||||
tuple(_projection(child) for child in title_children),
|
||||
_projection(entry),
|
||||
tuple(_projection(child) for child in entry_children),
|
||||
)
|
||||
|
||||
|
||||
def _validate_gate2(gate2: object) -> None:
|
||||
if not isinstance(gate2, Gate2Observation):
|
||||
raise FinalSubmitPanelError("缺少可信 Gate2Observation,已停止操作。")
|
||||
if (
|
||||
gate2.requested_color != TASK_COLOR
|
||||
or gate2.actual_color != TASK_COLOR
|
||||
or gate2.requested_size != TASK_SIZE
|
||||
or gate2.actual_size != TASK_SIZE
|
||||
):
|
||||
raise FinalSubmitPanelError("Gate2 规格不是 T-107 已取证目标,已停止操作。")
|
||||
if (
|
||||
type(gate2.requested_quantity) is not int
|
||||
or type(gate2.quantity_read) is not int
|
||||
or gate2.requested_quantity != FINAL_PANEL_QUANTITY
|
||||
or gate2.quantity_read != FINAL_PANEL_QUANTITY
|
||||
):
|
||||
raise FinalSubmitPanelError("Gate2 数量不是 T-107 已取证值,已停止操作。")
|
||||
_money(gate2.gate1_unit_price)
|
||||
gate2_amount = _money(gate2.gate2_panel_total_price)
|
||||
maximum = _money(gate2.max_total_price)
|
||||
if not isinstance(gate2.screenshot_path, Path) or not gate2.screenshot_path.name:
|
||||
raise FinalSubmitPanelError("Gate2 截图引用无效,已停止操作。")
|
||||
if not isinstance(gate2.captured_at, datetime) or gate2.captured_at.utcoffset() is None:
|
||||
raise FinalSubmitPanelError("Gate2 采集时间必须带时区,已停止操作。")
|
||||
if Decimal(gate2_amount) > Decimal(maximum):
|
||||
raise FinalSubmitPanelOverCapError("Gate2 顶部总额超过授权最高总价,已停止操作。")
|
||||
|
||||
|
||||
def _verified_final_panel(nodes: list[_Node]) -> _VerifiedPanel:
|
||||
profile_matches = [
|
||||
(profile, node)
|
||||
for profile in _PANEL_PROFILES
|
||||
for node in nodes
|
||||
if _exact(node, "android.view.ViewGroup", profile.panel_bounds, resource_id=_PDD_ID)
|
||||
]
|
||||
if len(profile_matches) != 1:
|
||||
raise FinalSubmitPanelError("最终面板未唯一命中已人工确认的精确布局配置。")
|
||||
profile, panel = profile_matches[0]
|
||||
price_frame = _one(
|
||||
node
|
||||
for node in nodes
|
||||
if node.parent is panel
|
||||
and _exact(node, "android.widget.FrameLayout", profile.price_frame_bounds, resource_id=_PDD_ID)
|
||||
)
|
||||
price_row = _one(
|
||||
node
|
||||
for node in nodes
|
||||
if node.parent is price_frame
|
||||
and _exact(node, "android.widget.LinearLayout", profile.price_row_bounds, resource_id=_PDD_ID)
|
||||
)
|
||||
price = _one(
|
||||
node
|
||||
for node in nodes
|
||||
if node.parent is price_row
|
||||
and _exact_text_node(node, "android.widget.TextView", profile.price_text_bounds, resource_id=_PDD_ID)
|
||||
and _GATE2_TEXT.fullmatch(node.text) is not None
|
||||
)
|
||||
gate2_match = _GATE2_TEXT.fullmatch(price.text)
|
||||
if gate2_match is None:
|
||||
raise FinalSubmitPanelError("最终面板 Gate2 顶部金额不可读。")
|
||||
gate2_amount = _money(gate2_match.group("amount"))
|
||||
|
||||
_one(
|
||||
node
|
||||
for node in nodes
|
||||
if node.parent is panel
|
||||
and _exact_text_node(node, "android.widget.TextView", profile.summary_bounds, resource_id=_PDD_ID)
|
||||
and node.text == _EXPECTED_SUMMARY
|
||||
)
|
||||
quantity_outer = _one(
|
||||
node
|
||||
for node in nodes
|
||||
if node.parent is panel
|
||||
and _exact(node, "android.widget.LinearLayout", profile.quantity_bounds, resource_id=_QUANTITY_ID)
|
||||
)
|
||||
quantity_inner = _one(
|
||||
node
|
||||
for node in nodes
|
||||
if node.parent is quantity_outer
|
||||
and _exact(node, "android.widget.LinearLayout", profile.quantity_bounds, resource_id="")
|
||||
)
|
||||
quantity_children = _direct_children(quantity_inner)
|
||||
if len(quantity_children) != 3:
|
||||
raise FinalSubmitPanelError("最终面板数量子结构漂移。")
|
||||
_one(
|
||||
node
|
||||
for node in quantity_children
|
||||
if _exact(node, "android.widget.ImageView", profile.minus_bounds, resource_id=_PDD_ID, clickable="true")
|
||||
and not node.text
|
||||
and node.desc == "减少数量"
|
||||
)
|
||||
_one(
|
||||
node
|
||||
for node in quantity_children
|
||||
if _exact(node, "android.widget.EditText", profile.value_bounds, resource_id=_PDD_ID, clickable="true")
|
||||
and node.text == str(FINAL_PANEL_QUANTITY)
|
||||
and not node.desc
|
||||
)
|
||||
_one(
|
||||
node
|
||||
for node in quantity_children
|
||||
if _exact(node, "android.widget.ImageView", profile.plus_bounds, resource_id=_PDD_ID, clickable="true")
|
||||
and not node.text
|
||||
and node.desc == "增加数量"
|
||||
)
|
||||
|
||||
submit_like = [
|
||||
node
|
||||
for node in nodes
|
||||
if "提交订单" in node.text or "提交订单" in node.desc
|
||||
]
|
||||
structured = [node for node in submit_like if _GATE3_TEXT.fullmatch(node.text) is not None and not node.desc]
|
||||
if len(submit_like) != 1 or len(structured) != 1:
|
||||
raise FinalSubmitPanelError("最终提交完整结构化文本匹配数不是 1,已停止操作。")
|
||||
submit = structured[0]
|
||||
if not _exact_text_node(submit, "android.widget.TextView", _SUBMIT_TEXT_BOUNDS, resource_id=""):
|
||||
raise FinalSubmitPanelError("最终提交文本节点结构漂移,已停止操作。")
|
||||
submit_parent = submit.parent
|
||||
if (
|
||||
submit_parent is None
|
||||
or not _exact(
|
||||
submit_parent,
|
||||
"android.widget.LinearLayout",
|
||||
_SUBMIT_TEXT_PARENT_BOUNDS,
|
||||
resource_id="",
|
||||
)
|
||||
or _direct_children(submit_parent) != [submit]
|
||||
):
|
||||
raise FinalSubmitPanelError("最终提交文本父结构漂移,已停止操作。")
|
||||
nearest = _nearest_clickable_ancestor(submit)
|
||||
if nearest is None or not _exact(
|
||||
nearest,
|
||||
"android.widget.FrameLayout",
|
||||
_SUBMIT_ACTION_BOUNDS,
|
||||
resource_id=_PDD_ID,
|
||||
clickable="true",
|
||||
):
|
||||
raise FinalSubmitPanelError("最终提交文本的最近可点击祖先不唯一或结构漂移。")
|
||||
gate3_match = _GATE3_TEXT.fullmatch(submit.text)
|
||||
if gate3_match is None:
|
||||
raise FinalSubmitPanelError("Gate3 最终控件金额不可读。")
|
||||
|
||||
return _VerifiedPanel(
|
||||
gate2_amount=gate2_amount,
|
||||
gate3_amount=_money(gate3_match.group("amount")),
|
||||
submit_text=submit.text,
|
||||
submit_match_count=len(structured),
|
||||
submit_enabled=submit.element.get("enabled") == "true" and nearest.element.get("enabled") == "true",
|
||||
nearest_clickable_ancestor_unique=True,
|
||||
)
|
||||
|
||||
|
||||
def _parse_nodes(raw: object) -> list[_Node]:
|
||||
if not isinstance(raw, str) or not raw:
|
||||
raise FinalSubmitPanelError("节点树读取失败,已停止操作。")
|
||||
try:
|
||||
root = ElementTree.fromstring(raw)
|
||||
except ElementTree.ParseError as error:
|
||||
raise FinalSubmitPanelError("节点树格式无效,已停止操作。") from error
|
||||
if root.tag != "hierarchy":
|
||||
raise FinalSubmitPanelError("节点树根节点无效,已停止操作。")
|
||||
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 _money(value: object) -> str:
|
||||
try:
|
||||
return require_money(value, "invalid_money")
|
||||
except ValidationError as error:
|
||||
raise FinalSubmitPanelError("金额不是规范十进制字符串,已停止操作。") from error
|
||||
|
||||
|
||||
def _one(nodes: Iterable[_Node]) -> _Node:
|
||||
matches = list(nodes)
|
||||
if len(matches) != 1:
|
||||
raise FinalSubmitPanelError("证据绑定页面角色不唯一,已停止操作。")
|
||||
return matches[0]
|
||||
|
||||
|
||||
def _exact(
|
||||
node: _Node,
|
||||
class_name: str,
|
||||
bounds: str,
|
||||
*,
|
||||
resource_id: str,
|
||||
clickable: str = "false",
|
||||
) -> bool:
|
||||
return (
|
||||
node.element.get("package") == PDD_PACKAGE
|
||||
and node.element.get("class") == class_name
|
||||
and node.element.get("resource-id", "") == resource_id
|
||||
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") == "false"
|
||||
and node.element.get("scrollable") == "false"
|
||||
)
|
||||
|
||||
|
||||
def _exact_text_node(node: _Node, class_name: str, bounds: str, *, resource_id: str) -> bool:
|
||||
return (
|
||||
_exact(node, class_name, bounds, resource_id=resource_id)
|
||||
and not node.desc
|
||||
)
|
||||
|
||||
|
||||
def _exact_text(node: _Node, text: str, bounds: str, *, resource_id: str = "") -> bool:
|
||||
return (
|
||||
_exact_text_node(node, "android.widget.TextView", bounds, resource_id=resource_id)
|
||||
and node.text == text
|
||||
)
|
||||
|
||||
|
||||
def _direct_children(node: _Node) -> list[_Node]:
|
||||
return [_Node(child, node) for child in node.element if child.tag == "node"]
|
||||
|
||||
|
||||
def _nearest_clickable_ancestor(node: _Node) -> _Node | None:
|
||||
current = node.parent
|
||||
while current is not None:
|
||||
if current.element.get("clickable") == "true":
|
||||
return current
|
||||
current = current.parent
|
||||
return None
|
||||
|
||||
|
||||
def _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 _projection(node: _Node) -> tuple[str, ...]:
|
||||
return (
|
||||
node.element.tag,
|
||||
node.element.get("package", ""),
|
||||
node.element.get("class", ""),
|
||||
node.element.get("resource-id", ""),
|
||||
node.bounds,
|
||||
node.element.get("clickable", ""),
|
||||
node.element.get("enabled", ""),
|
||||
node.element.get("visible-to-user", ""),
|
||||
node.text,
|
||||
node.desc,
|
||||
)
|
||||
@@ -0,0 +1,491 @@
|
||||
"""T-107 真机边界:零页面控件动作观察 Gate3,然后只尝试一次系统 Back。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, replace
|
||||
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, Protocol
|
||||
from uuid import uuid4
|
||||
|
||||
from adbutils.errors import AdbTimeout
|
||||
from PIL import Image, UnidentifiedImageError
|
||||
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 .final_submit_panel import (
|
||||
FinalSubmitPanelError,
|
||||
Gate3Observation,
|
||||
observe_gate3,
|
||||
returned_product_projection,
|
||||
)
|
||||
from .product_open import EXPECTED_PDD_VERSION
|
||||
from .quantity_gate2 import (
|
||||
EXPECTED_ANDROID_VERSION,
|
||||
EXPECTED_DEVICE_MODEL,
|
||||
EXPECTED_GOODS_ID,
|
||||
EXPECTED_SCREEN_SIZE,
|
||||
Gate2Observation,
|
||||
)
|
||||
|
||||
|
||||
class FinalSubmitPanelTimeoutError(FinalSubmitPanelError):
|
||||
"""设备读取或一次 Back 后置条件等待超时。"""
|
||||
|
||||
|
||||
class FinalSubmitPanelAdapterError(FinalSubmitPanelError):
|
||||
"""第三方设备接口失败后的脱敏映射。"""
|
||||
|
||||
|
||||
class ForegroundReader(Protocol):
|
||||
def read(self, serial: str) -> dict[str, str]: ...
|
||||
|
||||
|
||||
class FinalSubmitPanelDevice(Protocol):
|
||||
"""最终面板的窄设备能力;没有任何页面控件动作。"""
|
||||
|
||||
def app_info(self, package_name: str) -> dict[str, Any]: ...
|
||||
|
||||
def current_foreground(self) -> dict[str, str]: ...
|
||||
|
||||
def display_size(self) -> tuple[int, int]: ...
|
||||
|
||||
def dump_window_hierarchy(self) -> str: ...
|
||||
|
||||
def capture_screenshot(self) -> str: ...
|
||||
|
||||
def leave_final_submit_panel_once(self) -> None: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class FinalSubmitPanelRunResult:
|
||||
output_directory: Path
|
||||
screenshot_path: Path
|
||||
manifest_path: Path
|
||||
observation: Gate3Observation
|
||||
|
||||
|
||||
class FinalSubmitPanelFlow:
|
||||
"""读取同一最终面板并在成功观察后执行一次安全返回。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
device: FinalSubmitPanelDevice,
|
||||
wait_timeout_seconds: float = 1.0,
|
||||
poll_interval_seconds: float = 0.2,
|
||||
monotonic_clock: Callable[[], float] = monotonic,
|
||||
sleep_function: Callable[[float], None] = sleep,
|
||||
) -> None:
|
||||
if not _positive_finite(wait_timeout_seconds) or not _positive_finite(poll_interval_seconds):
|
||||
raise ValueError("等待参数必须是大于 0 的有限数值。")
|
||||
self._device = device
|
||||
self._timeout = float(wait_timeout_seconds)
|
||||
self._poll = float(poll_interval_seconds)
|
||||
self._clock = monotonic_clock
|
||||
self._sleep = sleep_function
|
||||
self._observation: Gate3Observation | None = None
|
||||
self._gate2: Gate2Observation | None = None
|
||||
self._terminal = False
|
||||
self._exited = False
|
||||
|
||||
@property
|
||||
def exited(self) -> bool:
|
||||
return self._exited
|
||||
|
||||
def require_ready_for_read(self) -> None:
|
||||
self._require_active()
|
||||
self._require_environment()
|
||||
|
||||
def observe(
|
||||
self,
|
||||
gate2: Gate2Observation,
|
||||
screenshot_path: Path,
|
||||
captured_at: datetime,
|
||||
) -> Gate3Observation:
|
||||
self._require_active()
|
||||
try:
|
||||
self._require_environment()
|
||||
observation = observe_gate3(
|
||||
self._read_hierarchy(),
|
||||
gate2,
|
||||
screenshot_path,
|
||||
captured_at,
|
||||
)
|
||||
self._require_environment()
|
||||
except BaseException:
|
||||
self._terminal = True
|
||||
raise
|
||||
self._gate2 = gate2
|
||||
self._observation = observation
|
||||
return observation
|
||||
|
||||
def exit_final_submit_panel_safely(self) -> None:
|
||||
self._require_active()
|
||||
if self._observation is None or self._gate2 is None:
|
||||
raise FinalSubmitPanelError("没有可信 Gate3 观察,拒绝执行安全返回。")
|
||||
try:
|
||||
self._require_environment()
|
||||
current = observe_gate3(
|
||||
self._read_hierarchy(),
|
||||
self._gate2,
|
||||
self._observation.screenshot_path,
|
||||
self._observation.captured_at,
|
||||
)
|
||||
self._require_environment()
|
||||
if current != self._observation:
|
||||
raise FinalSubmitPanelError("安全返回前最终面板事实漂移,已停止操作。")
|
||||
self._device.leave_final_submit_panel_once()
|
||||
except BaseException:
|
||||
self._terminal = True
|
||||
raise
|
||||
|
||||
deadline = self._clock() + self._timeout
|
||||
stable: tuple[object, ...] | None = None
|
||||
try:
|
||||
while True:
|
||||
self._require_environment()
|
||||
raw = self._read_hierarchy()
|
||||
try:
|
||||
projection = returned_product_projection(raw)
|
||||
except FinalSubmitPanelError:
|
||||
stable = None
|
||||
else:
|
||||
self._require_environment()
|
||||
if stable == projection:
|
||||
self._exited = True
|
||||
self._terminal = True
|
||||
return
|
||||
stable = projection
|
||||
remaining = deadline - self._clock()
|
||||
if remaining <= 0:
|
||||
raise FinalSubmitPanelTimeoutError("一次 Back 后未达到连续稳定同商品判据,未重试返回。")
|
||||
self._sleep(min(self._poll, remaining))
|
||||
except BaseException:
|
||||
self._terminal = True
|
||||
raise
|
||||
|
||||
def _require_environment(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 FinalSubmitPanelError("拼多多版本与 T-106/T-107 证据不一致,已停止操作。")
|
||||
foreground = self._device.current_foreground()
|
||||
if (
|
||||
not isinstance(foreground, dict)
|
||||
or foreground.get("package") != PDD_PACKAGE
|
||||
or not isinstance(foreground.get("activity"), str)
|
||||
or not foreground["activity"].strip()
|
||||
):
|
||||
raise FinalSubmitPanelError("拼多多不是唯一前台应用,已停止操作。")
|
||||
if self._device.display_size() != EXPECTED_SCREEN_SIZE:
|
||||
raise FinalSubmitPanelError("屏幕坐标空间与 T-106/T-107 证据不一致,已停止操作。")
|
||||
|
||||
def _read_hierarchy(self) -> str:
|
||||
value = self._device.dump_window_hierarchy()
|
||||
if not isinstance(value, str) or not value:
|
||||
raise FinalSubmitPanelError("节点树读取失败,已停止操作。")
|
||||
return value
|
||||
|
||||
def _require_active(self) -> None:
|
||||
if self._terminal:
|
||||
raise FinalSubmitPanelError("最终面板流程已进入不可重入终止态。")
|
||||
|
||||
|
||||
class UiautomatorFinalSubmitPanelAdapter(FinalSubmitPanelDevice):
|
||||
"""只暴露截图/XML 读取和一个 Android Back。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
device: Any,
|
||||
foreground_reader: ForegroundReader,
|
||||
serial: str,
|
||||
timeout_seconds: float,
|
||||
) -> None:
|
||||
if type(serial) is not str or not serial.strip() or serial != serial.strip():
|
||||
raise ValueError("serial 必须显式且非空。")
|
||||
if not _positive_finite(timeout_seconds):
|
||||
raise ValueError("timeout_seconds 必须是大于 0 的有限数值。")
|
||||
self._device = device
|
||||
self._foreground_reader = foreground_reader
|
||||
self._serial = serial
|
||||
self._timeout = float(timeout_seconds)
|
||||
self._back_attempted = False
|
||||
self._back_outcome = "not_attempted"
|
||||
|
||||
@property
|
||||
def page_control_action_attempts(self) -> int:
|
||||
return 0
|
||||
|
||||
@property
|
||||
def back_attempts(self) -> int:
|
||||
return int(self._back_attempted)
|
||||
|
||||
@property
|
||||
def back_outcome(self) -> str:
|
||||
return self._back_outcome
|
||||
|
||||
def app_info(self, package_name: str) -> dict[str, Any]:
|
||||
value = self._call("app_info", package_name)
|
||||
if not isinstance(value, dict):
|
||||
raise FinalSubmitPanelAdapterError("无法读取应用版本,已停止操作。")
|
||||
return value
|
||||
|
||||
def current_foreground(self) -> dict[str, str]:
|
||||
try:
|
||||
value = self._foreground_reader.read(self._serial)
|
||||
except FinalSubmitPanelError:
|
||||
raise
|
||||
except Exception as error:
|
||||
raise FinalSubmitPanelAdapterError("无法读取 Android 前台摘要,已停止操作。") from error
|
||||
if not isinstance(value, dict):
|
||||
raise FinalSubmitPanelAdapterError("Android 前台摘要无效,已停止操作。")
|
||||
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(type(item) is not int for item in value):
|
||||
raise FinalSubmitPanelAdapterError("无法读取屏幕坐标空间,已停止操作。")
|
||||
return value
|
||||
|
||||
def dump_window_hierarchy(self) -> str:
|
||||
value = self._call("jsonrpc_call", "dumpWindowHierarchy", [False, 50], timeout=self._timeout)
|
||||
if not isinstance(value, str):
|
||||
raise FinalSubmitPanelAdapterError("节点树读取失败,已停止操作。")
|
||||
return value
|
||||
|
||||
def capture_screenshot(self) -> str:
|
||||
value = self._call("jsonrpc_call", "takeScreenshot", SCREENSHOT_PARAMS, timeout=self._timeout)
|
||||
if not isinstance(value, str):
|
||||
raise FinalSubmitPanelAdapterError("Gate3 原始截图读取失败,已停止操作。")
|
||||
return value
|
||||
|
||||
def leave_final_submit_panel_once(self) -> None:
|
||||
if self._back_attempted:
|
||||
raise FinalSubmitPanelAdapterError("安全返回已尝试过,拒绝重试。")
|
||||
# RPC 超时无法证明 Back 未送达,必须先封存唯一动作机会。
|
||||
self._back_attempted = True
|
||||
self._back_outcome = "ambiguous"
|
||||
self._call("jsonrpc_call", "pressKey", ["back"], timeout=self._timeout)
|
||||
self._back_outcome = "completed"
|
||||
|
||||
def _call(self, method: str, *args: Any, **kwargs: Any) -> Any:
|
||||
try:
|
||||
return getattr(self._device, method)(*args, **kwargs)
|
||||
except (AdbTimeout, HTTPTimeoutError, TimeoutError) as error:
|
||||
raise FinalSubmitPanelTimeoutError("T-107 设备调用超时,已停止操作。") from error
|
||||
except FinalSubmitPanelError:
|
||||
raise
|
||||
except Exception as error:
|
||||
raise FinalSubmitPanelAdapterError("T-107 设备调用失败,已停止操作。") from error
|
||||
|
||||
|
||||
class FinalSubmitPanelRunner:
|
||||
"""从人工停驻的合并式最终面板完成围栏前只读 dry-run。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
adb_client: AdbClient,
|
||||
connector: Callable[[str], Any],
|
||||
foreground_reader: ForegroundReader,
|
||||
timeout_seconds: float,
|
||||
monotonic_clock: Callable[[], float] = monotonic,
|
||||
) -> None:
|
||||
if not _positive_finite(timeout_seconds):
|
||||
raise ValueError("timeout_seconds 必须是大于 0 的有限数值。")
|
||||
self._adb_client = adb_client
|
||||
self._connector = connector
|
||||
self._foreground_reader = foreground_reader
|
||||
self._timeout = float(timeout_seconds)
|
||||
self._clock = monotonic_clock
|
||||
|
||||
def run(
|
||||
self,
|
||||
serial: str,
|
||||
goods_id: str,
|
||||
gate2: Gate2Observation,
|
||||
output_directory: Path,
|
||||
) -> FinalSubmitPanelRunResult:
|
||||
target = Path(output_directory)
|
||||
staging: Path | None = None
|
||||
try:
|
||||
_validate_preflight(serial, goods_id, gate2, target)
|
||||
staging = _prepare_staging(target)
|
||||
inspection = self._adb_client.inspect(serial)
|
||||
_require_expected_device(inspection)
|
||||
adapter = UiautomatorFinalSubmitPanelAdapter(
|
||||
self._connector(serial),
|
||||
self._foreground_reader,
|
||||
serial,
|
||||
self._timeout,
|
||||
)
|
||||
flow = FinalSubmitPanelFlow(
|
||||
adapter,
|
||||
wait_timeout_seconds=self._timeout,
|
||||
monotonic_clock=self._clock,
|
||||
)
|
||||
flow.require_ready_for_read()
|
||||
screenshot_path = staging / "gate3_screenshot.png"
|
||||
_save_base64_screenshot(adapter.capture_screenshot(), screenshot_path)
|
||||
_require_screenshot(screenshot_path)
|
||||
captured_at = datetime.now(UTC)
|
||||
observation = flow.observe(gate2, screenshot_path, captured_at)
|
||||
flow.exit_final_submit_panel_safely()
|
||||
_require_action_audit(adapter, flow)
|
||||
|
||||
manifest_path = staging / "manifest.json"
|
||||
manifest_path.write_text(
|
||||
json.dumps(
|
||||
_manifest(inspection, serial, gate2, observation, screenshot_path, adapter),
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
sort_keys=True,
|
||||
)
|
||||
+ "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
os.rename(staging, target)
|
||||
staging = None
|
||||
except (DeviceConnectionError, FinalSubmitPanelError):
|
||||
_clean_staging(staging)
|
||||
raise
|
||||
except (AdbTimeout, HTTPTimeoutError, TimeoutError) as error:
|
||||
_clean_staging(staging)
|
||||
raise FinalSubmitPanelTimeoutError("T-107 真机运行超时,未发布证据。") from error
|
||||
except OSError as error:
|
||||
_clean_staging(staging)
|
||||
raise FinalSubmitPanelError("T-107 证据无法原子发布。") from error
|
||||
except Exception as error:
|
||||
_clean_staging(staging)
|
||||
raise FinalSubmitPanelError("T-107 真机运行未完成。") from error
|
||||
|
||||
published = replace(observation, screenshot_path=target / "gate3_screenshot.png")
|
||||
return FinalSubmitPanelRunResult(
|
||||
output_directory=target,
|
||||
screenshot_path=target / "gate3_screenshot.png",
|
||||
manifest_path=target / "manifest.json",
|
||||
observation=published,
|
||||
)
|
||||
|
||||
|
||||
def _positive_finite(value: object) -> bool:
|
||||
return isinstance(value, (int, float)) and not isinstance(value, bool) and value > 0 and isfinite(value)
|
||||
|
||||
|
||||
def _validate_preflight(
|
||||
serial: object,
|
||||
goods_id: object,
|
||||
gate2: object,
|
||||
target: Path,
|
||||
) -> None:
|
||||
if type(serial) is not str or not serial.strip() or serial != serial.strip():
|
||||
raise FinalSubmitPanelError("必须显式提供非空设备通道。")
|
||||
if type(goods_id) is not str or goods_id != EXPECTED_GOODS_ID:
|
||||
raise FinalSubmitPanelError("商品不是 T-107 已批准目标。")
|
||||
if not isinstance(gate2, Gate2Observation) or not gate2.screenshot_path.is_file():
|
||||
raise FinalSubmitPanelError("Gate2 原始截图不存在,已停止操作。")
|
||||
if target.exists() or not target.name:
|
||||
raise FinalSubmitPanelError("输出目录必须是不存在的明确新目录。")
|
||||
|
||||
|
||||
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()
|
||||
return staging
|
||||
except OSError as error:
|
||||
_clean_staging(staging)
|
||||
raise FinalSubmitPanelError("输出目录不可写,已停止操作。") from error
|
||||
|
||||
|
||||
def _clean_staging(staging: Path | None) -> None:
|
||||
if staging is not None and staging.exists():
|
||||
shutil.rmtree(staging)
|
||||
|
||||
|
||||
def _require_expected_device(inspection: DeviceInspection) -> None:
|
||||
if inspection.model != EXPECTED_DEVICE_MODEL or inspection.android_version != EXPECTED_ANDROID_VERSION:
|
||||
raise FinalSubmitPanelError("设备型号或 Android 版本与 T-106/T-107 证据不一致。")
|
||||
|
||||
|
||||
def _require_screenshot(path: Path) -> None:
|
||||
try:
|
||||
with Image.open(path) as image:
|
||||
image.load()
|
||||
if image.size != EXPECTED_SCREEN_SIZE or image.format != "PNG":
|
||||
raise FinalSubmitPanelError("Gate3 截图尺寸或格式与证据不一致。")
|
||||
except FinalSubmitPanelError:
|
||||
raise
|
||||
except (OSError, UnidentifiedImageError) as error:
|
||||
raise FinalSubmitPanelError("Gate3 截图不是有效图像。") from error
|
||||
|
||||
|
||||
def _require_action_audit(adapter: UiautomatorFinalSubmitPanelAdapter, flow: FinalSubmitPanelFlow) -> None:
|
||||
if (
|
||||
adapter.page_control_action_attempts != 0
|
||||
or adapter.back_attempts != 1
|
||||
or adapter.back_outcome != "completed"
|
||||
or not flow.exited
|
||||
):
|
||||
raise FinalSubmitPanelError("T-107 动作审计链不完整,拒绝发布。")
|
||||
|
||||
|
||||
def _manifest(
|
||||
inspection: DeviceInspection,
|
||||
serial: str,
|
||||
gate2: Gate2Observation,
|
||||
observation: Gate3Observation,
|
||||
screenshot_path: Path,
|
||||
adapter: UiautomatorFinalSubmitPanelAdapter,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"schema_version": 1,
|
||||
"operation": "t107-final-submit-panel-dry-run",
|
||||
"captured_at": observation.captured_at.isoformat(),
|
||||
"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": EXPECTED_PDD_VERSION,
|
||||
},
|
||||
"selection": {"color": observation.actual_color, "size": observation.actual_size},
|
||||
"quantity": {"requested": observation.requested_quantity, "read": observation.quantity_read},
|
||||
"prices": {
|
||||
"gate2_panel_total_price": observation.gate2_panel_total_price,
|
||||
"gate3_submit_amount": observation.gate3_submit_amount,
|
||||
"max_total_price": observation.max_total_price,
|
||||
},
|
||||
"submit_control": {
|
||||
"text": observation.submit_control_text,
|
||||
"match_count": observation.submit_control_match_count,
|
||||
"enabled": observation.submit_control_enabled,
|
||||
"nearest_clickable_ancestor_unique": observation.nearest_clickable_ancestor_unique,
|
||||
},
|
||||
"gate2_evidence": {
|
||||
"captured_at": gate2.captured_at.isoformat(),
|
||||
"screenshot_sha256": _sha256_file(gate2.screenshot_path),
|
||||
},
|
||||
"gate3_evidence": {
|
||||
"path": screenshot_path.name,
|
||||
"sha256": _sha256_file(screenshot_path),
|
||||
},
|
||||
"action_audit": {
|
||||
"page_control_action_attempts": adapter.page_control_action_attempts,
|
||||
"back_attempts": adapter.back_attempts,
|
||||
"back_rpc_outcome": adapter.back_outcome,
|
||||
},
|
||||
"safe_exit": "completed",
|
||||
"review_status": "human_review_required",
|
||||
}
|
||||
@@ -0,0 +1,340 @@
|
||||
"""T-106 合并式最终提交面板与人工返回态的纯只读取证。"""
|
||||
|
||||
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 re
|
||||
import shutil
|
||||
from typing import Any, Protocol
|
||||
from uuid import uuid4
|
||||
|
||||
from adbutils.errors import AdbTimeout
|
||||
from PIL import Image, UnidentifiedImageError
|
||||
from uiautomator2.exceptions import HTTPTimeoutError
|
||||
|
||||
from ..device.adb import AdbClient, CommandRunner, DeviceConnectionError, DeviceInspection
|
||||
from ..device.baseline import (
|
||||
HIERARCHY_PARAMS,
|
||||
PDD_PACKAGE,
|
||||
SCREENSHOT_PARAMS,
|
||||
_save_base64_screenshot,
|
||||
_sha256_file,
|
||||
_validate_hierarchy,
|
||||
)
|
||||
from .product_open import EXPECTED_PDD_VERSION
|
||||
from .product_url import parse_product_url
|
||||
|
||||
|
||||
EXPECTED_GOODS_ID = "937122477375"
|
||||
EXPECTED_DEVICE_MODEL = "PKG110"
|
||||
EXPECTED_ANDROID_VERSION = "16"
|
||||
EXPECTED_SCREEN_SIZE = (1080, 2376)
|
||||
DECLARED_STATES = ("gate2-navigation-source", "returned-safe-page")
|
||||
|
||||
|
||||
class OrderConfirmEvidenceError(RuntimeError):
|
||||
"""T-106 只读证据未形成完整原子产物。"""
|
||||
|
||||
|
||||
class OrderConfirmEvidenceTimeoutError(OrderConfirmEvidenceError):
|
||||
"""只读取证设备调用超时。"""
|
||||
|
||||
|
||||
class OrderConfirmReadDevice(Protocol):
|
||||
"""T-106 唯一设备边界;故意只有读取能力。"""
|
||||
|
||||
def app_info(self, package_name: str) -> dict[str, Any]: ...
|
||||
|
||||
def window_size(self) -> tuple[int, int]: ...
|
||||
|
||||
def jsonrpc_call(self, method: str, params: Any = None, timeout: float = 10) -> Any: ...
|
||||
|
||||
|
||||
class OrderConfirmForegroundReader(Protocol):
|
||||
"""读取 Android 前台摘要,不暴露通用命令执行。"""
|
||||
|
||||
def read(self, serial: str) -> dict[str, str]: ...
|
||||
|
||||
|
||||
class Android16ForegroundReader:
|
||||
"""读取 Android 16 的唯一 top-resumed Activity。"""
|
||||
|
||||
_TOP_RESUMED_PATTERN = re.compile(
|
||||
r"(?m)^\s*topResumedActivity=ActivityRecord\{[^\r\n}]*?\s+u\d+\s+"
|
||||
r"(?P<package>[^/\s]+)/(?P<activity>[^\s}]+)\s+t\d+\}\s*$"
|
||||
)
|
||||
|
||||
def __init__(self, runner: CommandRunner, timeout_seconds: float) -> None:
|
||||
if not _is_positive_finite(timeout_seconds):
|
||||
raise ValueError("timeout_seconds 必须是大于 0 的有限数值")
|
||||
self._runner = runner
|
||||
self._timeout_seconds = timeout_seconds
|
||||
|
||||
def read(self, serial: str) -> dict[str, str]:
|
||||
if type(serial) is not str or not serial.strip() or serial != serial.strip():
|
||||
raise OrderConfirmEvidenceError("必须显式提供非空设备通道。")
|
||||
result = self._runner.run(
|
||||
("-s", serial, "shell", "dumpsys", "activity", "activities"),
|
||||
self._timeout_seconds,
|
||||
)
|
||||
if result.returncode != 0:
|
||||
raise OrderConfirmEvidenceError("Android 前台摘要读取失败,未发布证据。")
|
||||
matches = list(self._TOP_RESUMED_PATTERN.finditer(result.stdout))
|
||||
if len(matches) != 1:
|
||||
raise OrderConfirmEvidenceError("Android 前台摘要不唯一,未发布证据。")
|
||||
match = matches[0]
|
||||
return {"package": match.group("package"), "activity": match.group("activity")}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class OrderConfirmEvidenceResult:
|
||||
output_directory: Path
|
||||
manifest_path: Path
|
||||
screenshot_path: Path
|
||||
hierarchy_path: Path
|
||||
app_path: Path
|
||||
|
||||
|
||||
class OrderConfirmEvidenceCapturer:
|
||||
"""采集人工准备的单个稳定状态,不解析页面业务字段。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
adb_client: AdbClient,
|
||||
connector: Callable[[str], OrderConfirmReadDevice],
|
||||
foreground_reader: OrderConfirmForegroundReader,
|
||||
timeout_seconds: float,
|
||||
) -> None:
|
||||
if not _is_positive_finite(timeout_seconds):
|
||||
raise ValueError("timeout_seconds 必须是大于 0 的有限数值")
|
||||
self._adb_client = adb_client
|
||||
self._connector = connector
|
||||
self._foreground_reader = foreground_reader
|
||||
self._timeout_seconds = timeout_seconds
|
||||
self._started = False
|
||||
|
||||
def capture(
|
||||
self,
|
||||
serial: str,
|
||||
goods_id: str,
|
||||
human_declared_state: str,
|
||||
output_directory: Path,
|
||||
) -> OrderConfirmEvidenceResult:
|
||||
if self._started:
|
||||
raise OrderConfirmEvidenceError("同一取证器不可重复调用。")
|
||||
self._started = True
|
||||
_validate_inputs(serial, goods_id, human_declared_state)
|
||||
target = Path(output_directory)
|
||||
_validate_new_target(target)
|
||||
|
||||
staging: Path | None = None
|
||||
try:
|
||||
inspection = self._adb_client.inspect(serial)
|
||||
_require_expected_device(inspection)
|
||||
device = self._connector(serial)
|
||||
initial_app = _require_read_precondition(
|
||||
device,
|
||||
self._foreground_reader.read(serial),
|
||||
)
|
||||
|
||||
target.parent.mkdir(parents=True, exist_ok=True)
|
||||
staging = target.parent / f".{target.name}.staging-{uuid4().hex}"
|
||||
staging.mkdir()
|
||||
|
||||
screenshot_path = staging / "screenshot.png"
|
||||
screenshot_payload = _read_rpc(
|
||||
device,
|
||||
"takeScreenshot",
|
||||
SCREENSHOT_PARAMS,
|
||||
self._timeout_seconds,
|
||||
)
|
||||
if not isinstance(screenshot_payload, str):
|
||||
raise OrderConfirmEvidenceError("页面截图无效,未发布证据。")
|
||||
_save_base64_screenshot(screenshot_payload, screenshot_path)
|
||||
_require_screenshot(screenshot_path)
|
||||
|
||||
hierarchy = _read_rpc(
|
||||
device,
|
||||
"dumpWindowHierarchy",
|
||||
HIERARCHY_PARAMS,
|
||||
self._timeout_seconds,
|
||||
)
|
||||
_validate_hierarchy(hierarchy)
|
||||
hierarchy_path = staging / "hierarchy.xml"
|
||||
hierarchy_path.write_text(hierarchy, encoding="utf-8")
|
||||
|
||||
final_app = _require_read_precondition(
|
||||
device,
|
||||
self._foreground_reader.read(serial),
|
||||
)
|
||||
if final_app != initial_app:
|
||||
raise OrderConfirmEvidenceError("取证期间前台页面漂移,未发布证据。")
|
||||
|
||||
app_path = staging / "app.json"
|
||||
app_path.write_text(
|
||||
json.dumps(final_app, ensure_ascii=False, indent=2, sort_keys=True) + "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
manifest_path = staging / "manifest.json"
|
||||
manifest_path.write_text(
|
||||
json.dumps(
|
||||
_manifest(
|
||||
inspection,
|
||||
serial,
|
||||
goods_id,
|
||||
human_declared_state,
|
||||
screenshot_path,
|
||||
hierarchy_path,
|
||||
app_path,
|
||||
),
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
sort_keys=True,
|
||||
)
|
||||
+ "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
os.rename(staging, target)
|
||||
staging = None
|
||||
except (DeviceConnectionError, OrderConfirmEvidenceError):
|
||||
_clean_staging(staging)
|
||||
raise
|
||||
except (AdbTimeout, HTTPTimeoutError, TimeoutError) as error:
|
||||
_clean_staging(staging)
|
||||
raise OrderConfirmEvidenceTimeoutError("T-106 只读取证超时,未发布证据。") from error
|
||||
except (OSError, UnidentifiedImageError, ValueError) as error:
|
||||
_clean_staging(staging)
|
||||
raise OrderConfirmEvidenceError("T-106 证据无法原子发布,未发布证据。") from error
|
||||
except Exception as error:
|
||||
_clean_staging(staging)
|
||||
raise OrderConfirmEvidenceError("T-106 只读取证未完成,未发布证据。") from error
|
||||
|
||||
return OrderConfirmEvidenceResult(
|
||||
output_directory=target,
|
||||
manifest_path=target / "manifest.json",
|
||||
screenshot_path=target / "screenshot.png",
|
||||
hierarchy_path=target / "hierarchy.xml",
|
||||
app_path=target / "app.json",
|
||||
)
|
||||
|
||||
|
||||
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_inputs(serial: object, goods_id: object, state: object) -> None:
|
||||
if type(serial) is not str or not serial.strip() or serial != serial.strip():
|
||||
raise OrderConfirmEvidenceError("必须显式提供非空设备通道。")
|
||||
if type(goods_id) is not str or goods_id != EXPECTED_GOODS_ID:
|
||||
raise OrderConfirmEvidenceError("商品不是 T-106 已批准取证目标。")
|
||||
parse_product_url(f"https://mobile.yangkeduo.com/goods.html?goods_id={goods_id}")
|
||||
if type(state) is not str or state not in DECLARED_STATES:
|
||||
raise OrderConfirmEvidenceError("人工声明状态无效。")
|
||||
|
||||
|
||||
def _validate_new_target(target: Path) -> None:
|
||||
if target.exists() or not target.name:
|
||||
raise OrderConfirmEvidenceError("输出目录必须是不存在的明确新目录。")
|
||||
|
||||
|
||||
def _require_expected_device(inspection: DeviceInspection) -> None:
|
||||
if inspection.model != EXPECTED_DEVICE_MODEL or inspection.android_version != EXPECTED_ANDROID_VERSION:
|
||||
raise OrderConfirmEvidenceError("设备不是已批准取证组合。")
|
||||
|
||||
|
||||
def _require_read_precondition(
|
||||
device: OrderConfirmReadDevice,
|
||||
current: object,
|
||||
) -> dict[str, str]:
|
||||
info = 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 OrderConfirmEvidenceError("拼多多版本不是已批准取证版本。")
|
||||
if not isinstance(current, dict) or current.get("package") != PDD_PACKAGE:
|
||||
raise OrderConfirmEvidenceError("拼多多不在前台。")
|
||||
activity = current.get("activity")
|
||||
if not isinstance(activity, str) or not activity.strip():
|
||||
raise OrderConfirmEvidenceError("前台应用摘要不完整。")
|
||||
if device.window_size() != EXPECTED_SCREEN_SIZE:
|
||||
raise OrderConfirmEvidenceError("屏幕坐标空间不是已批准尺寸。")
|
||||
return {"package": PDD_PACKAGE, "activity": activity, "pdd_version": EXPECTED_PDD_VERSION}
|
||||
|
||||
|
||||
def _read_rpc(
|
||||
device: OrderConfirmReadDevice,
|
||||
method: str,
|
||||
params: object,
|
||||
timeout_seconds: float,
|
||||
) -> object:
|
||||
return device.jsonrpc_call(method, params, timeout=timeout_seconds)
|
||||
|
||||
|
||||
def _require_screenshot(path: Path) -> None:
|
||||
with Image.open(path) as image:
|
||||
image.load()
|
||||
if image.size != EXPECTED_SCREEN_SIZE or image.format != "PNG":
|
||||
raise OrderConfirmEvidenceError("页面截图格式或尺寸无效。")
|
||||
|
||||
|
||||
def _clean_staging(staging: Path | None) -> None:
|
||||
if staging is not None and staging.exists():
|
||||
shutil.rmtree(staging)
|
||||
|
||||
|
||||
def _manifest(
|
||||
inspection: DeviceInspection,
|
||||
serial: str,
|
||||
goods_id: str,
|
||||
state: str,
|
||||
screenshot_path: Path,
|
||||
hierarchy_path: Path,
|
||||
app_path: Path,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"schema_version": 1,
|
||||
"operation": "t106-order-confirm-readonly-evidence",
|
||||
"captured_at": datetime.now(UTC).isoformat(),
|
||||
"product": {
|
||||
"goods_id": goods_id,
|
||||
"canonical_url": f"https://mobile.yangkeduo.com/goods.html?goods_id={goods_id}",
|
||||
},
|
||||
"human_declared_state": state,
|
||||
"review_status": "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,
|
||||
"role": "human_prepared_state_raw_screenshot",
|
||||
"sha256": _sha256_file(screenshot_path),
|
||||
},
|
||||
{
|
||||
"path": hierarchy_path.name,
|
||||
"role": "human_prepared_state_raw_hierarchy_local_only",
|
||||
"sha256": _sha256_file(hierarchy_path),
|
||||
},
|
||||
{
|
||||
"path": app_path.name,
|
||||
"role": "human_prepared_state_app_identity",
|
||||
"sha256": _sha256_file(app_path),
|
||||
},
|
||||
],
|
||||
}
|
||||
@@ -0,0 +1,605 @@
|
||||
"""T-105:证据绑定的数量读回与 Gate2 面板总额。
|
||||
|
||||
本模块只允许数量 1 保持不变,或从数量 1 对唯一加号点击一次到数量 2。
|
||||
它不包含确认页、提交围栏、提交订单或付款能力。
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime
|
||||
from decimal import Decimal
|
||||
from functools import partial
|
||||
from math import isfinite
|
||||
from pathlib import Path
|
||||
import re
|
||||
from time import monotonic, sleep
|
||||
from typing import Any, Callable, Protocol
|
||||
from xml.etree import ElementTree
|
||||
|
||||
from ..core.errors import ValidationError
|
||||
from ..core.validation import require_money
|
||||
from ..device.baseline import PDD_PACKAGE
|
||||
from .product_open import EXPECTED_PDD_VERSION
|
||||
from .sku_selection import SkuSelectionError, _parse_nodes as _parse_sku_nodes
|
||||
from .sku_selection import _product_exit_projection
|
||||
|
||||
|
||||
EXPECTED_DEVICE_MODEL = "PKG110"
|
||||
EXPECTED_ANDROID_VERSION = "16"
|
||||
EXPECTED_SCREEN_SIZE = (1080, 2376)
|
||||
EXPECTED_GOODS_ID = "937122477375"
|
||||
TASK_COLOR = "黑色CHA(纯棉)"
|
||||
TASK_SIZE = "M(建议100-115)"
|
||||
UI_COLOR = "黑色 CHA (纯棉)"
|
||||
UI_SIZE = "M(建议100-115)"
|
||||
EXPECTED_GATE1_UNIT_PRICE = "12.88"
|
||||
|
||||
_PDD_ID = "com.xunmeng.pinduoduo:id/pdd"
|
||||
_QUANTITY_CONTAINER_ID = "com.xunmeng.pinduoduo:id/gnl"
|
||||
_COLOR_ID = "com.xunmeng.pinduoduo:id/tv_content"
|
||||
_TARGET_SUMMARY = f"已选: {UI_COLOR} {UI_SIZE}"
|
||||
_AMOUNT_TEXT = re.compile(r"^快卖完 ¥(?P<amount>(?:0|[1-9]\d*)\.\d{2})$")
|
||||
_BOUNDS = re.compile(r"^\[(\d+),(\d+)\]\[(\d+),(\d+)\]$")
|
||||
|
||||
|
||||
class QuantityGate2Error(RuntimeError):
|
||||
"""数量/Gate2 判据不成立时的脱敏安全停止。"""
|
||||
|
||||
|
||||
class QuantityGate2TimeoutError(QuantityGate2Error):
|
||||
"""设备调用或后置条件等待超时。"""
|
||||
|
||||
|
||||
class QuantityGate2OverCapError(QuantityGate2Error):
|
||||
"""目标数量面板总额超过管理员授权上限。"""
|
||||
|
||||
|
||||
class QuantityGate2Device(Protocol):
|
||||
"""T-105 的窄设备能力;没有通用页面动作。"""
|
||||
|
||||
def app_info(self, package_name: str) -> dict[str, Any]: ...
|
||||
|
||||
def current_foreground(self) -> dict[str, str]: ...
|
||||
|
||||
def display_size(self) -> tuple[int, int]: ...
|
||||
|
||||
def dump_window_hierarchy(self) -> str: ...
|
||||
|
||||
def increment_quantity_once(self, bounds: str) -> None: ...
|
||||
|
||||
def capture_screenshot(self) -> str: ...
|
||||
|
||||
def leave_sku_panel_once(self) -> None: ...
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Gate1Observation:
|
||||
color: str
|
||||
size: str
|
||||
quantity: int
|
||||
gate1_unit_price: str
|
||||
screenshot_path: Path
|
||||
captured_at: datetime
|
||||
|
||||
def __post_init__(self) -> None:
|
||||
if self.color != TASK_COLOR or self.size != TASK_SIZE or type(self.quantity) is not int or self.quantity != 1:
|
||||
raise QuantityGate2Error("Gate1 规格或数量不是已取证前置,已停止操作。")
|
||||
if _money(self.gate1_unit_price) != EXPECTED_GATE1_UNIT_PRICE:
|
||||
raise QuantityGate2Error("Gate1 单价不是已取证值,已停止操作。")
|
||||
if not isinstance(self.screenshot_path, Path) or not self.screenshot_path.name:
|
||||
raise QuantityGate2Error("Gate1 截图路径无效,已停止操作。")
|
||||
if not isinstance(self.captured_at, datetime) or self.captured_at.utcoffset() is None:
|
||||
raise QuantityGate2Error("Gate1 采集时间必须带时区,已停止操作。")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Gate2Observation:
|
||||
requested_color: str
|
||||
requested_size: str
|
||||
actual_color: str
|
||||
actual_size: str
|
||||
requested_quantity: int
|
||||
quantity_read: int
|
||||
gate1_unit_price: str
|
||||
gate2_panel_total_price: str
|
||||
max_total_price: str
|
||||
screenshot_path: Path
|
||||
captured_at: datetime
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _PanelProfile:
|
||||
quantity: int
|
||||
panel_bounds: str
|
||||
price_row_bounds: str
|
||||
price_bounds: str
|
||||
price_text: str
|
||||
summary_bounds: str
|
||||
quantity_bounds: str
|
||||
minus_bounds: str
|
||||
value_bounds: str
|
||||
plus_bounds: str
|
||||
color_bounds: str
|
||||
|
||||
|
||||
_INITIAL = _PanelProfile(
|
||||
1,
|
||||
"[0,474][1080,863]",
|
||||
"[396,498][895,570]",
|
||||
"[396,503][712,570]",
|
||||
"快卖完 ¥12.88",
|
||||
"[396,654][1053,716]",
|
||||
"[396,752][645,827]",
|
||||
"[396,752][474,827]",
|
||||
"[480,752][561,827]",
|
||||
"[567,752][645,827]",
|
||||
"[126,1000][438,1024]",
|
||||
)
|
||||
_TARGET = _PanelProfile(
|
||||
2,
|
||||
"[0,474][1080,861]",
|
||||
"[396,498][740,570]",
|
||||
"[396,503][722,570]",
|
||||
"快卖完 ¥32.76",
|
||||
"[396,582][1053,644]",
|
||||
"[396,750][645,825]",
|
||||
"[396,750][474,825]",
|
||||
"[480,750][561,825]",
|
||||
"[567,750][645,825]",
|
||||
"[126,998][438,1024]",
|
||||
)
|
||||
_PROFILES = {1: _INITIAL, 2: _TARGET}
|
||||
|
||||
|
||||
@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", "")
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class _VerifiedPanel:
|
||||
quantity: int
|
||||
panel_total_price: str
|
||||
plus_bounds: str
|
||||
projection: tuple[object, ...]
|
||||
|
||||
|
||||
class QuantityGate2Flow:
|
||||
"""从已确认数量 1 面板推进至获准数量,并安全退出同一商品。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
device: QuantityGate2Device,
|
||||
wait_timeout_seconds: float = 1.0,
|
||||
poll_interval_seconds: float = 0.2,
|
||||
monotonic_clock: Callable[[], float] = monotonic,
|
||||
sleep_function: Callable[[float], None] = sleep,
|
||||
) -> None:
|
||||
if not _positive_finite(wait_timeout_seconds) or not _positive_finite(poll_interval_seconds):
|
||||
raise ValueError("等待参数必须是大于 0 的有限数值。")
|
||||
self._device = device
|
||||
self._timeout = float(wait_timeout_seconds)
|
||||
self._poll = float(poll_interval_seconds)
|
||||
self._clock = monotonic_clock
|
||||
self._sleep = sleep_function
|
||||
self._increment_attempted = False
|
||||
self._pending: tuple[str, Callable[[list[_Node]], _VerifiedPanel]] | None = None
|
||||
self._terminal = False
|
||||
self._verified_quantity: int | None = None
|
||||
self._verified_total: str | None = None
|
||||
self._exited = False
|
||||
|
||||
@property
|
||||
def increment_attempts(self) -> int:
|
||||
return int(self._increment_attempted)
|
||||
|
||||
@property
|
||||
def exited(self) -> bool:
|
||||
return self._exited
|
||||
|
||||
@property
|
||||
def can_exit_safely(self) -> bool:
|
||||
return self._pending is None and self._verified_quantity in _PROFILES and not self._exited
|
||||
|
||||
def set_quantity_and_verify(
|
||||
self,
|
||||
gate1: Gate1Observation,
|
||||
target_quantity: int,
|
||||
max_total_price: str,
|
||||
) -> _VerifiedPanel:
|
||||
self._require_active()
|
||||
_validate_request(gate1, target_quantity, max_total_price)
|
||||
initial = self._read_verified(_INITIAL)
|
||||
if initial.panel_total_price != gate1.gate1_unit_price:
|
||||
raise QuantityGate2Error("Gate1 当前面板事实已漂移,已停止操作。")
|
||||
|
||||
if target_quantity == 1:
|
||||
verified = initial
|
||||
else:
|
||||
# 动作前 fresh 读取,不能使用上一次节点或缓存 bounds。
|
||||
before = self._read_hierarchy()
|
||||
precondition = _verified_panel(_parse_nodes(before), _INITIAL)
|
||||
_require_unique_action_occupants(_parse_nodes(before), precondition.plus_bounds)
|
||||
self._pending = (before, partial(_verified_panel, profile=_TARGET))
|
||||
self._increment_attempted = True
|
||||
try:
|
||||
self._device.increment_quantity_once(precondition.plus_bounds)
|
||||
verified = self._wait_for_pending()
|
||||
except BaseException:
|
||||
# 点击超时可能已送达;封存后不允许本 Flow 重试或继续。
|
||||
self._terminal = True
|
||||
raise
|
||||
|
||||
self._verified_quantity = verified.quantity
|
||||
self._verified_total = verified.panel_total_price
|
||||
if Decimal(verified.panel_total_price) > Decimal(max_total_price):
|
||||
raise QuantityGate2OverCapError("Gate2 面板总额超过授权最高总价,已停止操作。")
|
||||
return verified
|
||||
|
||||
def build_observation(
|
||||
self,
|
||||
gate1: Gate1Observation,
|
||||
target_quantity: int,
|
||||
max_total_price: str,
|
||||
screenshot_path: Path,
|
||||
captured_at: datetime,
|
||||
) -> Gate2Observation:
|
||||
self._require_active()
|
||||
_validate_request(gate1, target_quantity, max_total_price)
|
||||
if not isinstance(screenshot_path, Path) or not screenshot_path.name:
|
||||
raise QuantityGate2Error("Gate2 截图路径无效,已停止操作。")
|
||||
if not isinstance(captured_at, datetime) or captured_at.utcoffset() is None:
|
||||
raise QuantityGate2Error("Gate2 采集时间必须带时区,已停止操作。")
|
||||
verified = self._read_verified(_PROFILES[target_quantity])
|
||||
if (
|
||||
self._verified_quantity != verified.quantity
|
||||
or self._verified_total != verified.panel_total_price
|
||||
or Decimal(verified.panel_total_price) > Decimal(max_total_price)
|
||||
):
|
||||
raise QuantityGate2Error("截图后 Gate2 事实漂移,已停止操作。")
|
||||
return Gate2Observation(
|
||||
requested_color=gate1.color,
|
||||
requested_size=gate1.size,
|
||||
actual_color=TASK_COLOR,
|
||||
actual_size=TASK_SIZE,
|
||||
requested_quantity=target_quantity,
|
||||
quantity_read=verified.quantity,
|
||||
gate1_unit_price=gate1.gate1_unit_price,
|
||||
gate2_panel_total_price=verified.panel_total_price,
|
||||
max_total_price=max_total_price,
|
||||
screenshot_path=screenshot_path,
|
||||
captured_at=captured_at,
|
||||
)
|
||||
|
||||
def exit_sku_panel_safely(self) -> None:
|
||||
self._require_active()
|
||||
if self._verified_quantity not in _PROFILES:
|
||||
raise QuantityGate2Error("没有可用于安全退出的数量事实,已停止操作。")
|
||||
current = self._read_verified(_PROFILES[self._verified_quantity])
|
||||
if current.panel_total_price != self._verified_total:
|
||||
raise QuantityGate2Error("安全退出前 Gate2 事实漂移,已停止操作。")
|
||||
try:
|
||||
self._device.leave_sku_panel_once()
|
||||
except BaseException:
|
||||
self._terminal = True
|
||||
raise
|
||||
|
||||
deadline = self._clock() + self._timeout
|
||||
stable: tuple[object, ...] | None = None
|
||||
try:
|
||||
while True:
|
||||
self._require_environment()
|
||||
raw = self._read_hierarchy()
|
||||
try:
|
||||
projection = _product_exit_projection(_parse_sku_nodes(raw))
|
||||
except SkuSelectionError:
|
||||
stable = None
|
||||
else:
|
||||
self._require_environment()
|
||||
if stable == projection:
|
||||
self._exited = True
|
||||
self._terminal = True
|
||||
return
|
||||
stable = projection
|
||||
remaining = deadline - self._clock()
|
||||
if remaining <= 0:
|
||||
raise QuantityGate2TimeoutError("安全退出未达到连续稳定同商品判据,未重试返回。")
|
||||
self._sleep(min(self._poll, remaining))
|
||||
except BaseException:
|
||||
self._terminal = True
|
||||
raise
|
||||
|
||||
def reconcile_pending_action(self) -> _VerifiedPanel | None:
|
||||
"""结果不明时只读一次待定后置条件;绝不重发加号。"""
|
||||
|
||||
if self._pending is None:
|
||||
return None
|
||||
return self._wait_for_pending()
|
||||
|
||||
def _wait_for_pending(self) -> _VerifiedPanel:
|
||||
if self._pending is None:
|
||||
raise QuantityGate2Error("没有可调和的数量动作。")
|
||||
previous, condition = self._pending
|
||||
deadline = self._clock() + self._timeout
|
||||
while True:
|
||||
self._require_environment()
|
||||
raw = self._read_hierarchy()
|
||||
if raw != previous:
|
||||
try:
|
||||
verified = condition(_parse_nodes(raw))
|
||||
except QuantityGate2Error:
|
||||
pass
|
||||
else:
|
||||
self._pending = None
|
||||
return verified
|
||||
remaining = deadline - self._clock()
|
||||
if remaining <= 0:
|
||||
raise QuantityGate2TimeoutError("数量动作后置条件未确认,未重试加号。")
|
||||
self._sleep(min(self._poll, remaining))
|
||||
|
||||
def _read_verified(self, profile: _PanelProfile) -> _VerifiedPanel:
|
||||
self._require_environment()
|
||||
return _verified_panel(_parse_nodes(self._read_hierarchy()), profile)
|
||||
|
||||
def _require_environment(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 QuantityGate2Error("拼多多版本与 T-105 证据不一致,已停止操作。")
|
||||
foreground = self._device.current_foreground()
|
||||
if (
|
||||
not isinstance(foreground, dict)
|
||||
or foreground.get("package") != PDD_PACKAGE
|
||||
or not isinstance(foreground.get("activity"), str)
|
||||
or not foreground["activity"].strip()
|
||||
):
|
||||
raise QuantityGate2Error("拼多多不是唯一前台应用,已停止操作。")
|
||||
if self._device.display_size() != EXPECTED_SCREEN_SIZE:
|
||||
raise QuantityGate2Error("屏幕坐标空间与 T-105 证据不一致,已停止操作。")
|
||||
|
||||
def _read_hierarchy(self) -> str:
|
||||
value = self._device.dump_window_hierarchy()
|
||||
if not isinstance(value, str) or not value:
|
||||
raise QuantityGate2Error("节点树读取失败,已停止操作。")
|
||||
return value
|
||||
|
||||
def _require_active(self) -> None:
|
||||
if self._terminal:
|
||||
raise QuantityGate2Error("数量流程已进入不可重入终止态。")
|
||||
|
||||
|
||||
def _positive_finite(value: object) -> bool:
|
||||
return isinstance(value, (int, float)) and not isinstance(value, bool) and value > 0 and isfinite(value)
|
||||
|
||||
|
||||
def _validate_request(gate1: object, target_quantity: object, max_total_price: object) -> None:
|
||||
if not isinstance(gate1, Gate1Observation):
|
||||
raise QuantityGate2Error("缺少可信 Gate1Observation,已停止操作。")
|
||||
if type(target_quantity) is not int or target_quantity not in _PROFILES:
|
||||
raise QuantityGate2Error("目标数量没有本项目真机证据,已停止操作。")
|
||||
_money(max_total_price)
|
||||
|
||||
|
||||
def _money(value: object) -> str:
|
||||
try:
|
||||
return require_money(value, "invalid_money")
|
||||
except ValidationError as error:
|
||||
raise QuantityGate2Error("金额不是规范十进制字符串,已停止操作。") from error
|
||||
|
||||
|
||||
def _parse_nodes(raw: str) -> list[_Node]:
|
||||
try:
|
||||
root = ElementTree.fromstring(raw)
|
||||
except ElementTree.ParseError as error:
|
||||
raise QuantityGate2Error("节点树格式无效,已停止操作。") from error
|
||||
if root.tag != "hierarchy":
|
||||
raise QuantityGate2Error("节点树根节点无效,已停止操作。")
|
||||
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 _verified_panel(nodes: list[_Node], profile: _PanelProfile) -> _VerifiedPanel:
|
||||
panel = _one(
|
||||
node for node in nodes
|
||||
if _exact(node, "android.view.ViewGroup", profile.panel_bounds, resource_id=_PDD_ID, clickable="false")
|
||||
)
|
||||
price_frame = _one(
|
||||
node for node in nodes
|
||||
if node.parent is panel and _exact(node, "android.widget.FrameLayout", "[396,498][1053,570]", resource_id=_PDD_ID, clickable="false")
|
||||
)
|
||||
price_row = _one(
|
||||
node for node in nodes
|
||||
if node.parent is price_frame and _exact(node, "android.widget.LinearLayout", profile.price_row_bounds, resource_id=_PDD_ID, clickable="false")
|
||||
)
|
||||
price = _one(
|
||||
node for node in nodes
|
||||
if node.parent is price_row
|
||||
and _exact(node, "android.widget.TextView", profile.price_bounds, resource_id=_PDD_ID, clickable="false")
|
||||
and node.text == profile.price_text
|
||||
and not node.desc
|
||||
)
|
||||
match = _AMOUNT_TEXT.fullmatch(price.text)
|
||||
if match is None:
|
||||
raise QuantityGate2Error("Gate2 面板总额角色不可读。")
|
||||
amount = _money(match.group("amount"))
|
||||
|
||||
_one(
|
||||
node for node in nodes
|
||||
if node.parent is panel
|
||||
and _exact(node, "android.widget.TextView", profile.summary_bounds, resource_id=_PDD_ID, clickable="false")
|
||||
and node.text == _TARGET_SUMMARY
|
||||
and not node.desc
|
||||
)
|
||||
quantity_outer = _one(
|
||||
node for node in nodes
|
||||
if node.parent is panel
|
||||
and _exact(node, "android.widget.LinearLayout", profile.quantity_bounds, resource_id=_QUANTITY_CONTAINER_ID, clickable="false")
|
||||
)
|
||||
quantity_inner = _one(
|
||||
node for node in nodes
|
||||
if node.parent is quantity_outer
|
||||
and _exact(node, "android.widget.LinearLayout", profile.quantity_bounds, resource_id="", clickable="false")
|
||||
)
|
||||
_one(
|
||||
node for node in nodes
|
||||
if node.parent is quantity_inner
|
||||
and _exact(node, "android.widget.ImageView", profile.minus_bounds, resource_id=_PDD_ID, clickable="true")
|
||||
and node.desc == "减少数量"
|
||||
and not node.text
|
||||
)
|
||||
quantity = _one(
|
||||
node for node in nodes
|
||||
if node.parent is quantity_inner
|
||||
and _exact(node, "android.widget.EditText", profile.value_bounds, resource_id=_PDD_ID, clickable="true")
|
||||
and node.text == str(profile.quantity)
|
||||
and not node.desc
|
||||
)
|
||||
plus = _one(
|
||||
node for node in nodes
|
||||
if node.parent is quantity_inner
|
||||
and _exact(node, "android.widget.ImageView", profile.plus_bounds, resource_id=_PDD_ID, clickable="true")
|
||||
and node.desc == "增加数量"
|
||||
and not node.text
|
||||
)
|
||||
if quantity.element.get("selected") != "false" or plus.element.get("selected") != "false":
|
||||
raise QuantityGate2Error("数量控件选中属性漂移。")
|
||||
|
||||
_require_selected_target(nodes, profile)
|
||||
projection = (
|
||||
"quantity_gate2_8_17_0",
|
||||
profile.quantity,
|
||||
amount,
|
||||
tuple(_projection(node) for node in (panel, price_frame, price_row, price, quantity_outer, quantity_inner, quantity, plus)),
|
||||
)
|
||||
return _VerifiedPanel(profile.quantity, amount, plus.bounds, projection)
|
||||
|
||||
|
||||
def _require_selected_target(nodes: list[_Node], profile: _PanelProfile) -> None:
|
||||
color = _one(
|
||||
node for node in nodes
|
||||
if _exact(node, "android.widget.TextView", profile.color_bounds, resource_id=_COLOR_ID, clickable="true", selected="true")
|
||||
and node.text == UI_COLOR
|
||||
and not node.desc
|
||||
)
|
||||
if color.parent is None or color.parent.element.get("selected") != "true":
|
||||
raise QuantityGate2Error("目标颜色没有精确选中。")
|
||||
size = _one(
|
||||
node for node in nodes
|
||||
if _exact(node, "android.widget.TextView", "[439,1582][831,1667]", resource_id=_PDD_ID, clickable="true", selected="true")
|
||||
and node.text == UI_SIZE
|
||||
and not node.desc
|
||||
)
|
||||
if size.parent is None or size.parent.element.get("clickable") != "true":
|
||||
raise QuantityGate2Error("目标尺码结构漂移。")
|
||||
|
||||
|
||||
def _exact(
|
||||
node: _Node,
|
||||
class_name: str,
|
||||
bounds: str,
|
||||
*,
|
||||
resource_id: str,
|
||||
clickable: str,
|
||||
selected: str = "false",
|
||||
) -> bool:
|
||||
element = node.element
|
||||
return (
|
||||
element.get("package") == PDD_PACKAGE
|
||||
and element.get("class") == class_name
|
||||
and node.bounds == bounds
|
||||
and element.get("resource-id", "") == resource_id
|
||||
and element.get("clickable") == clickable
|
||||
and element.get("selected") == selected
|
||||
and element.get("enabled") == "true"
|
||||
and element.get("visible-to-user") == "true"
|
||||
and element.get("scrollable") == "false"
|
||||
)
|
||||
|
||||
|
||||
def _one(values: Any) -> _Node:
|
||||
matches = list(values)
|
||||
if len(matches) != 1:
|
||||
raise QuantityGate2Error("T-105 页面证据角色缺失或不唯一。")
|
||||
return matches[0]
|
||||
|
||||
|
||||
def _projection(node: _Node) -> tuple[str, ...]:
|
||||
return (
|
||||
node.element.tag,
|
||||
node.element.get("package", ""),
|
||||
node.element.get("class", ""),
|
||||
node.bounds,
|
||||
node.element.get("resource-id", ""),
|
||||
node.element.get("clickable", ""),
|
||||
node.element.get("selected", ""),
|
||||
node.text,
|
||||
node.desc,
|
||||
)
|
||||
|
||||
|
||||
def _bounds_center(bounds: str) -> tuple[int, int]:
|
||||
match = _BOUNDS.fullmatch(bounds)
|
||||
if match is None:
|
||||
raise QuantityGate2Error("数量控件坐标无效。")
|
||||
left, top, right, bottom = (int(value) for value in match.groups())
|
||||
if not (0 <= left < right <= EXPECTED_SCREEN_SIZE[0] and 0 <= top < bottom <= EXPECTED_SCREEN_SIZE[1]):
|
||||
raise QuantityGate2Error("数量控件坐标超出已取证屏幕。")
|
||||
return left + (right - left) // 2, top + (bottom - top) // 2
|
||||
|
||||
|
||||
def _require_unique_action_occupants(nodes: list[_Node], bounds: str) -> None:
|
||||
target = _one(node for node in nodes if node.bounds == bounds and node.desc == "增加数量")
|
||||
x, y = _bounds_center(bounds)
|
||||
occupants = [
|
||||
node for node in nodes
|
||||
if node.element.get("clickable") == "true"
|
||||
and node.element.get("enabled") == "true"
|
||||
and node.element.get("visible-to-user") == "true"
|
||||
and _contains(node.bounds, x, y)
|
||||
]
|
||||
if not occupants or any(not _same_branch(node, target) for node in occupants):
|
||||
raise QuantityGate2Error("数量加号中心存在未知可点击覆盖层,已停止操作。")
|
||||
|
||||
|
||||
def _contains(bounds: str, x: int, y: int) -> bool:
|
||||
match = _BOUNDS.fullmatch(bounds)
|
||||
if match is None:
|
||||
return False
|
||||
left, top, right, bottom = (int(value) for value in match.groups())
|
||||
return left <= x < right and top <= y < bottom
|
||||
|
||||
|
||||
def _same_branch(candidate: _Node, target: _Node) -> bool:
|
||||
current: _Node | None = target
|
||||
while current is not None:
|
||||
if current.element is candidate.element:
|
||||
return True
|
||||
current = current.parent
|
||||
current = candidate
|
||||
while current is not None:
|
||||
if current.element is target.element:
|
||||
return True
|
||||
current = current.parent
|
||||
return False
|
||||
@@ -0,0 +1,397 @@
|
||||
"""T-105 真机运行边界:一次数量加号、Gate2 原图与一次安全退出。"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from collections.abc import Callable
|
||||
from dataclasses import dataclass, replace
|
||||
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, Protocol
|
||||
from uuid import uuid4
|
||||
|
||||
from adbutils.errors import AdbTimeout
|
||||
from PIL import Image, UnidentifiedImageError
|
||||
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 .quantity_gate2 import (
|
||||
EXPECTED_ANDROID_VERSION,
|
||||
EXPECTED_DEVICE_MODEL,
|
||||
EXPECTED_GOODS_ID,
|
||||
EXPECTED_SCREEN_SIZE,
|
||||
Gate1Observation,
|
||||
Gate2Observation,
|
||||
QuantityGate2Device,
|
||||
QuantityGate2Error,
|
||||
QuantityGate2Flow,
|
||||
QuantityGate2TimeoutError,
|
||||
_bounds_center,
|
||||
_money,
|
||||
)
|
||||
|
||||
|
||||
class ForegroundReader(Protocol):
|
||||
def read(self, serial: str) -> dict[str, str]: ...
|
||||
|
||||
|
||||
class QuantityGate2AdapterError(QuantityGate2Error):
|
||||
"""第三方设备接口失败后的脱敏映射。"""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class QuantityGate2RunResult:
|
||||
output_directory: Path
|
||||
screenshot_path: Path
|
||||
manifest_path: Path
|
||||
observation: Gate2Observation
|
||||
|
||||
|
||||
class UiautomatorQuantityGate2Adapter(QuantityGate2Device):
|
||||
"""只暴露 T-105 已批准的一个加号和一个 Back。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
device: Any,
|
||||
foreground_reader: ForegroundReader,
|
||||
serial: str,
|
||||
timeout_seconds: float,
|
||||
) -> None:
|
||||
if type(serial) is not str or not serial.strip() or serial != serial.strip():
|
||||
raise ValueError("serial 必须显式且非空。")
|
||||
if not _positive_finite(timeout_seconds):
|
||||
raise ValueError("timeout_seconds 必须是大于 0 的有限数值。")
|
||||
self._device = device
|
||||
self._foreground_reader = foreground_reader
|
||||
self._serial = serial
|
||||
self._timeout = float(timeout_seconds)
|
||||
self._increment_attempted = False
|
||||
self._increment_bounds: str | None = None
|
||||
self._increment_outcome = "not_attempted"
|
||||
self._back_attempted = False
|
||||
self._back_outcome = "not_attempted"
|
||||
|
||||
@property
|
||||
def increment_attempts(self) -> int:
|
||||
return int(self._increment_attempted)
|
||||
|
||||
@property
|
||||
def increment_bounds(self) -> str | None:
|
||||
return self._increment_bounds
|
||||
|
||||
@property
|
||||
def increment_outcome(self) -> str:
|
||||
return self._increment_outcome
|
||||
|
||||
@property
|
||||
def back_attempts(self) -> int:
|
||||
return int(self._back_attempted)
|
||||
|
||||
@property
|
||||
def back_outcome(self) -> str:
|
||||
return self._back_outcome
|
||||
|
||||
def app_info(self, package_name: str) -> dict[str, Any]:
|
||||
value = self._call("app_info", package_name)
|
||||
if not isinstance(value, dict):
|
||||
raise QuantityGate2AdapterError("无法读取应用版本,已停止操作。")
|
||||
return value
|
||||
|
||||
def current_foreground(self) -> dict[str, str]:
|
||||
try:
|
||||
value = self._foreground_reader.read(self._serial)
|
||||
except QuantityGate2Error:
|
||||
raise
|
||||
except Exception as error:
|
||||
raise QuantityGate2AdapterError("无法读取 Android 前台摘要,已停止操作。") from error
|
||||
if not isinstance(value, dict):
|
||||
raise QuantityGate2AdapterError("Android 前台摘要无效,已停止操作。")
|
||||
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(type(item) is not int for item in value):
|
||||
raise QuantityGate2AdapterError("无法读取屏幕坐标空间,已停止操作。")
|
||||
return value
|
||||
|
||||
def dump_window_hierarchy(self) -> str:
|
||||
value = self._call("jsonrpc_call", "dumpWindowHierarchy", [False, 50], timeout=self._timeout)
|
||||
if not isinstance(value, str):
|
||||
raise QuantityGate2AdapterError("节点树读取失败,已停止操作。")
|
||||
return value
|
||||
|
||||
def increment_quantity_once(self, bounds: str) -> None:
|
||||
if self._increment_attempted:
|
||||
raise QuantityGate2AdapterError("数量加号已尝试过,拒绝重试。")
|
||||
center_x, center_y = _bounds_center(bounds)
|
||||
# RPC 超时无法证明事件未送达,动作机会必须先持久在内存审计状态中。
|
||||
self._increment_attempted = True
|
||||
self._increment_bounds = bounds
|
||||
self._increment_outcome = "ambiguous"
|
||||
self._call("jsonrpc_call", "click", [center_x, center_y], timeout=self._timeout)
|
||||
self._increment_outcome = "completed"
|
||||
|
||||
def capture_screenshot(self) -> str:
|
||||
value = self._call("jsonrpc_call", "takeScreenshot", SCREENSHOT_PARAMS, timeout=self._timeout)
|
||||
if not isinstance(value, str):
|
||||
raise QuantityGate2AdapterError("Gate2 原始截图读取失败,已停止操作。")
|
||||
return value
|
||||
|
||||
def leave_sku_panel_once(self) -> None:
|
||||
if self._back_attempted:
|
||||
raise QuantityGate2AdapterError("安全返回已尝试过,拒绝重试。")
|
||||
self._back_attempted = True
|
||||
self._back_outcome = "ambiguous"
|
||||
self._call("jsonrpc_call", "pressKey", ["back"], timeout=self._timeout)
|
||||
self._back_outcome = "completed"
|
||||
|
||||
def _call(self, method: str, *args: Any, **kwargs: Any) -> Any:
|
||||
try:
|
||||
return getattr(self._device, method)(*args, **kwargs)
|
||||
except (AdbTimeout, HTTPTimeoutError, TimeoutError) as error:
|
||||
raise QuantityGate2TimeoutError("T-105 设备调用超时,已停止操作。") from error
|
||||
except QuantityGate2Error:
|
||||
raise
|
||||
except Exception as error:
|
||||
raise QuantityGate2AdapterError("T-105 设备调用失败,已停止操作。") from error
|
||||
|
||||
|
||||
class QuantityGate2Runner:
|
||||
"""从人工停驻的数量 1 目标面板执行 T-105 已取证闭环。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
adb_client: AdbClient,
|
||||
connector: Callable[[str], Any],
|
||||
foreground_reader: ForegroundReader,
|
||||
timeout_seconds: float,
|
||||
monotonic_clock: Callable[[], float] = monotonic,
|
||||
) -> None:
|
||||
if not _positive_finite(timeout_seconds):
|
||||
raise ValueError("timeout_seconds 必须是大于 0 的有限数值。")
|
||||
self._adb_client = adb_client
|
||||
self._connector = connector
|
||||
self._foreground_reader = foreground_reader
|
||||
self._timeout = float(timeout_seconds)
|
||||
self._clock = monotonic_clock
|
||||
|
||||
def run(
|
||||
self,
|
||||
serial: str,
|
||||
goods_id: str,
|
||||
gate1: Gate1Observation,
|
||||
target_quantity: int,
|
||||
max_total_price: str,
|
||||
output_directory: Path,
|
||||
) -> QuantityGate2RunResult:
|
||||
target = Path(output_directory)
|
||||
staging: Path | None = None
|
||||
adapter: UiautomatorQuantityGate2Adapter | None = None
|
||||
flow: QuantityGate2Flow | None = None
|
||||
try:
|
||||
_validate_preflight(serial, goods_id, gate1, target_quantity, max_total_price, target)
|
||||
staging = _prepare_staging(target)
|
||||
inspection = self._adb_client.inspect(serial)
|
||||
_require_expected_device(inspection)
|
||||
adapter = UiautomatorQuantityGate2Adapter(
|
||||
self._connector(serial),
|
||||
self._foreground_reader,
|
||||
serial,
|
||||
self._timeout,
|
||||
)
|
||||
flow = QuantityGate2Flow(
|
||||
adapter,
|
||||
wait_timeout_seconds=self._timeout,
|
||||
monotonic_clock=self._clock,
|
||||
)
|
||||
flow.set_quantity_and_verify(gate1, target_quantity, max_total_price)
|
||||
|
||||
screenshot_path = staging / "gate2_screenshot.png"
|
||||
_save_base64_screenshot(adapter.capture_screenshot(), screenshot_path)
|
||||
captured_at = datetime.now(UTC)
|
||||
_require_screenshot_size(screenshot_path)
|
||||
observation = flow.build_observation(
|
||||
gate1,
|
||||
target_quantity,
|
||||
max_total_price,
|
||||
screenshot_path,
|
||||
captured_at,
|
||||
)
|
||||
flow.exit_sku_panel_safely()
|
||||
_require_action_audit(adapter, target_quantity)
|
||||
|
||||
manifest_path = staging / "manifest.json"
|
||||
manifest_path.write_text(
|
||||
json.dumps(
|
||||
_manifest(inspection, serial, gate1, observation, screenshot_path, adapter),
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
sort_keys=True,
|
||||
)
|
||||
+ "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
os.rename(staging, target)
|
||||
staging = None
|
||||
except (DeviceConnectionError, QuantityGate2Error):
|
||||
_attempt_known_safe_exit(flow)
|
||||
_clean_staging(staging)
|
||||
raise
|
||||
except (AdbTimeout, HTTPTimeoutError, TimeoutError) as error:
|
||||
_attempt_known_safe_exit(flow)
|
||||
_clean_staging(staging)
|
||||
raise QuantityGate2TimeoutError("T-105 真机运行超时,未发布证据。") from error
|
||||
except OSError as error:
|
||||
_attempt_known_safe_exit(flow)
|
||||
_clean_staging(staging)
|
||||
raise QuantityGate2Error("T-105 证据无法原子发布。") from error
|
||||
except Exception as error:
|
||||
_attempt_known_safe_exit(flow)
|
||||
_clean_staging(staging)
|
||||
raise QuantityGate2Error("T-105 真机运行未完成。") from error
|
||||
|
||||
published_observation = replace(
|
||||
observation,
|
||||
screenshot_path=target / "gate2_screenshot.png",
|
||||
)
|
||||
return QuantityGate2RunResult(
|
||||
output_directory=target,
|
||||
screenshot_path=target / "gate2_screenshot.png",
|
||||
manifest_path=target / "manifest.json",
|
||||
observation=published_observation,
|
||||
)
|
||||
|
||||
|
||||
def _positive_finite(value: object) -> bool:
|
||||
return isinstance(value, (int, float)) and not isinstance(value, bool) and value > 0 and isfinite(value)
|
||||
|
||||
|
||||
def _validate_preflight(
|
||||
serial: object,
|
||||
goods_id: object,
|
||||
gate1: object,
|
||||
target_quantity: object,
|
||||
max_total_price: object,
|
||||
target: Path,
|
||||
) -> None:
|
||||
if type(serial) is not str or not serial.strip() or serial != serial.strip():
|
||||
raise QuantityGate2Error("必须显式提供非空设备通道。")
|
||||
if type(goods_id) is not str or goods_id != EXPECTED_GOODS_ID:
|
||||
raise QuantityGate2Error("商品不是 T-105 已批准目标。")
|
||||
if not isinstance(gate1, Gate1Observation) or not gate1.screenshot_path.is_file():
|
||||
raise QuantityGate2Error("Gate1 原始截图不存在,已停止操作。")
|
||||
if type(target_quantity) is not int or target_quantity not in {1, 2}:
|
||||
raise QuantityGate2Error("目标数量没有 T-105 真机证据。")
|
||||
_money(max_total_price)
|
||||
if target.exists() or not target.name:
|
||||
raise QuantityGate2Error("输出目录必须是不存在的明确新目录。")
|
||||
|
||||
|
||||
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()
|
||||
return staging
|
||||
except OSError as error:
|
||||
_clean_staging(staging)
|
||||
raise QuantityGate2Error("输出目录不可写,已停止操作。") from error
|
||||
|
||||
|
||||
def _clean_staging(staging: Path | None) -> None:
|
||||
if staging is not None and staging.exists():
|
||||
shutil.rmtree(staging)
|
||||
|
||||
|
||||
def _require_expected_device(inspection: DeviceInspection) -> None:
|
||||
if inspection.model != EXPECTED_DEVICE_MODEL or inspection.android_version != EXPECTED_ANDROID_VERSION:
|
||||
raise QuantityGate2Error("设备型号或 Android 版本与 T-105 证据不一致。")
|
||||
|
||||
|
||||
def _require_screenshot_size(path: Path) -> None:
|
||||
try:
|
||||
with Image.open(path) as image:
|
||||
image.load()
|
||||
if image.size != EXPECTED_SCREEN_SIZE:
|
||||
raise QuantityGate2Error("Gate2 截图尺寸与 T-105 证据不一致。")
|
||||
except QuantityGate2Error:
|
||||
raise
|
||||
except (OSError, UnidentifiedImageError) as error:
|
||||
raise QuantityGate2Error("Gate2 截图不是有效图像。") from error
|
||||
|
||||
|
||||
def _require_action_audit(adapter: UiautomatorQuantityGate2Adapter, target_quantity: int) -> None:
|
||||
expected_increment = int(target_quantity == 2)
|
||||
if (
|
||||
adapter.increment_attempts != expected_increment
|
||||
or (expected_increment and adapter.increment_outcome != "completed")
|
||||
or (expected_increment and adapter.increment_bounds != "[567,752][645,827]")
|
||||
or (not expected_increment and adapter.increment_outcome != "not_attempted")
|
||||
or adapter.back_attempts != 1
|
||||
or adapter.back_outcome != "completed"
|
||||
):
|
||||
raise QuantityGate2Error("T-105 动作审计链不完整,拒绝发布。")
|
||||
|
||||
|
||||
def _attempt_known_safe_exit(flow: QuantityGate2Flow | None) -> None:
|
||||
if flow is None or not flow.can_exit_safely:
|
||||
return
|
||||
try:
|
||||
flow.exit_sku_panel_safely()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _manifest(
|
||||
inspection: DeviceInspection,
|
||||
serial: str,
|
||||
gate1: Gate1Observation,
|
||||
observation: Gate2Observation,
|
||||
screenshot_path: Path,
|
||||
adapter: UiautomatorQuantityGate2Adapter,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"schema_version": 1,
|
||||
"operation": "t105-quantity-gate2",
|
||||
"captured_at": observation.captured_at.isoformat(),
|
||||
"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": EXPECTED_PDD_VERSION,
|
||||
},
|
||||
"selection": {"color": observation.actual_color, "size": observation.actual_size},
|
||||
"quantity": {"requested": observation.requested_quantity, "read": observation.quantity_read},
|
||||
"prices": {
|
||||
"gate1_unit_price": observation.gate1_unit_price,
|
||||
"gate2_panel_total_price": observation.gate2_panel_total_price,
|
||||
"max_total_price": observation.max_total_price,
|
||||
},
|
||||
"gate1_evidence": {
|
||||
"captured_at": gate1.captured_at.isoformat(),
|
||||
"screenshot_sha256": _sha256_file(gate1.screenshot_path),
|
||||
},
|
||||
"gate2_evidence": {
|
||||
"path": screenshot_path.name,
|
||||
"sha256": _sha256_file(screenshot_path),
|
||||
},
|
||||
"action_audit": {
|
||||
"increment_attempts": adapter.increment_attempts,
|
||||
"increment_rpc_outcome": adapter.increment_outcome,
|
||||
"back_attempts": adapter.back_attempts,
|
||||
"back_rpc_outcome": adapter.back_outcome,
|
||||
},
|
||||
"safe_exit": "completed",
|
||||
"review_status": "human_review_required",
|
||||
}
|
||||
@@ -0,0 +1,351 @@
|
||||
"""T-105 数量两态的纯只读取证;人工确认前不识别或操作数量控件。"""
|
||||
|
||||
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 re
|
||||
import shutil
|
||||
from typing import Any, Protocol
|
||||
from uuid import uuid4
|
||||
|
||||
from adbutils.errors import AdbTimeout
|
||||
from PIL import Image, UnidentifiedImageError
|
||||
from uiautomator2.exceptions import HTTPTimeoutError
|
||||
|
||||
from ..device.adb import AdbClient, CommandRunner, DeviceConnectionError, DeviceInspection
|
||||
from ..device.baseline import (
|
||||
HIERARCHY_PARAMS,
|
||||
PDD_PACKAGE,
|
||||
SCREENSHOT_PARAMS,
|
||||
_save_base64_screenshot,
|
||||
_sha256_file,
|
||||
_validate_hierarchy,
|
||||
)
|
||||
from .product_open import EXPECTED_PDD_VERSION
|
||||
from .product_url import parse_product_url
|
||||
|
||||
|
||||
EXPECTED_GOODS_ID = "937122477375"
|
||||
EXPECTED_DEVICE_MODEL = "PKG110"
|
||||
EXPECTED_ANDROID_VERSION = "16"
|
||||
EXPECTED_SCREEN_SIZE = (1080, 2376)
|
||||
TARGET_SELECTION = {"color": "黑色CHA(纯棉)", "size": "M(建议100-115)"}
|
||||
DECLARED_QUANTITIES = {"initial": 1, "target": 2}
|
||||
|
||||
|
||||
class QuantityGate2EvidenceError(RuntimeError):
|
||||
"""T-105 两态证据未形成完整原子产物。"""
|
||||
|
||||
|
||||
class QuantityGate2EvidenceTimeoutError(QuantityGate2EvidenceError):
|
||||
"""只读取证设备调用超时。"""
|
||||
|
||||
|
||||
class QuantityGate2ReadDevice(Protocol):
|
||||
"""阶段一唯一设备边界;故意不暴露任何页面动作。"""
|
||||
|
||||
def app_info(self, package_name: str) -> dict[str, Any]: ...
|
||||
|
||||
def window_size(self) -> tuple[int, int]: ...
|
||||
|
||||
def jsonrpc_call(self, method: str, params: Any = None, timeout: float = 10) -> Any: ...
|
||||
|
||||
|
||||
class QuantityGate2ForegroundReader(Protocol):
|
||||
"""只读取 Android 16 的唯一 resumed activity,不暴露通用 shell。"""
|
||||
|
||||
def read(self, serial: str) -> dict[str, str]: ...
|
||||
|
||||
|
||||
class Android16TopResumedForegroundReader:
|
||||
"""绕开 adbutils 2.12.0 对 Android 16 ``topResumedActivity`` 的误解析。"""
|
||||
|
||||
_TOP_RESUMED_PATTERN = re.compile(
|
||||
r"(?m)^\s*topResumedActivity=ActivityRecord\{[^\r\n}]*?\s+u\d+\s+"
|
||||
r"(?P<package>[^/\s]+)/(?P<activity>[^\s}]+)\s+t\d+\}\s*$"
|
||||
)
|
||||
|
||||
def __init__(self, runner: CommandRunner, timeout_seconds: float) -> None:
|
||||
if not _is_positive_finite(timeout_seconds):
|
||||
raise ValueError("timeout_seconds 必须是大于 0 的有限数值")
|
||||
self._runner = runner
|
||||
self._timeout_seconds = timeout_seconds
|
||||
|
||||
def read(self, serial: str) -> dict[str, str]:
|
||||
if type(serial) is not str or not serial.strip() or serial != serial.strip():
|
||||
raise QuantityGate2EvidenceError("必须显式提供非空设备通道。")
|
||||
result = self._runner.run(
|
||||
("-s", serial, "shell", "dumpsys", "activity", "activities"),
|
||||
self._timeout_seconds,
|
||||
)
|
||||
if result.returncode != 0:
|
||||
raise QuantityGate2EvidenceError("Android 前台摘要读取失败,未发布证据。")
|
||||
matches = list(self._TOP_RESUMED_PATTERN.finditer(result.stdout))
|
||||
if len(matches) != 1:
|
||||
raise QuantityGate2EvidenceError("Android 前台摘要不唯一,未发布证据。")
|
||||
match = matches[0]
|
||||
return {
|
||||
"package": match.group("package"),
|
||||
"activity": match.group("activity"),
|
||||
}
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class QuantityGate2EvidenceResult:
|
||||
output_directory: Path
|
||||
manifest_path: Path
|
||||
screenshot_path: Path
|
||||
hierarchy_path: Path
|
||||
app_path: Path
|
||||
|
||||
|
||||
class QuantityGate2EvidenceCapturer:
|
||||
"""记录人工准备的数量状态,不从页面推断声明是否正确。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
adb_client: AdbClient,
|
||||
connector: Callable[[str], QuantityGate2ReadDevice],
|
||||
foreground_reader: QuantityGate2ForegroundReader,
|
||||
timeout_seconds: float,
|
||||
) -> None:
|
||||
if not _is_positive_finite(timeout_seconds):
|
||||
raise ValueError("timeout_seconds 必须是大于 0 的有限数值")
|
||||
self._adb_client = adb_client
|
||||
self._connector = connector
|
||||
self._foreground_reader = foreground_reader
|
||||
self._timeout_seconds = timeout_seconds
|
||||
self._started = False
|
||||
|
||||
def capture(
|
||||
self,
|
||||
serial: str,
|
||||
goods_id: str,
|
||||
human_declared_state: str,
|
||||
human_declared_quantity: int,
|
||||
output_directory: Path,
|
||||
) -> QuantityGate2EvidenceResult:
|
||||
if self._started:
|
||||
raise QuantityGate2EvidenceError("同一取证器不可重复调用。")
|
||||
self._started = True
|
||||
_validate_inputs(serial, goods_id, human_declared_state, human_declared_quantity)
|
||||
target = Path(output_directory)
|
||||
_validate_new_target(target)
|
||||
|
||||
staging: Path | None = None
|
||||
try:
|
||||
inspection = self._adb_client.inspect(serial)
|
||||
_require_expected_device(inspection)
|
||||
device = self._connector(serial)
|
||||
initial_app = _require_read_precondition(
|
||||
device,
|
||||
self._foreground_reader.read(serial),
|
||||
)
|
||||
|
||||
target.parent.mkdir(parents=True, exist_ok=True)
|
||||
staging = target.parent / f".{target.name}.staging-{uuid4().hex}"
|
||||
staging.mkdir()
|
||||
|
||||
screenshot_path = staging / "screenshot.png"
|
||||
screenshot_payload = _read_rpc(
|
||||
device,
|
||||
"takeScreenshot",
|
||||
SCREENSHOT_PARAMS,
|
||||
self._timeout_seconds,
|
||||
)
|
||||
if not isinstance(screenshot_payload, str):
|
||||
raise QuantityGate2EvidenceError("数量状态截图无效,未发布证据。")
|
||||
_save_base64_screenshot(screenshot_payload, screenshot_path)
|
||||
_require_screenshot_size(screenshot_path)
|
||||
|
||||
hierarchy = _read_rpc(
|
||||
device,
|
||||
"dumpWindowHierarchy",
|
||||
HIERARCHY_PARAMS,
|
||||
self._timeout_seconds,
|
||||
)
|
||||
_validate_hierarchy(hierarchy)
|
||||
hierarchy_path = staging / "hierarchy.xml"
|
||||
hierarchy_path.write_text(hierarchy, encoding="utf-8")
|
||||
|
||||
final_app = _require_read_precondition(
|
||||
device,
|
||||
self._foreground_reader.read(serial),
|
||||
)
|
||||
if final_app != initial_app:
|
||||
raise QuantityGate2EvidenceError("数量状态取证期间前台页面漂移,未发布证据。")
|
||||
app_path = staging / "app.json"
|
||||
app_path.write_text(
|
||||
json.dumps(final_app, ensure_ascii=False, indent=2, sort_keys=True) + "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
manifest_path = staging / "manifest.json"
|
||||
manifest_path.write_text(
|
||||
json.dumps(
|
||||
_manifest(
|
||||
inspection,
|
||||
serial,
|
||||
goods_id,
|
||||
human_declared_state,
|
||||
human_declared_quantity,
|
||||
screenshot_path,
|
||||
hierarchy_path,
|
||||
app_path,
|
||||
),
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
sort_keys=True,
|
||||
)
|
||||
+ "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
os.rename(staging, target)
|
||||
staging = None
|
||||
except (DeviceConnectionError, QuantityGate2EvidenceError):
|
||||
_clean_staging(staging)
|
||||
raise
|
||||
except (AdbTimeout, HTTPTimeoutError, TimeoutError) as error:
|
||||
_clean_staging(staging)
|
||||
raise QuantityGate2EvidenceTimeoutError("数量两态只读取证超时,未发布证据。") from error
|
||||
except (OSError, UnidentifiedImageError, ValueError) as error:
|
||||
_clean_staging(staging)
|
||||
raise QuantityGate2EvidenceError("数量两态证据无法原子发布,未发布证据。") from error
|
||||
except Exception as error:
|
||||
_clean_staging(staging)
|
||||
raise QuantityGate2EvidenceError("数量两态只读取证未完成,未发布证据。") from error
|
||||
|
||||
return QuantityGate2EvidenceResult(
|
||||
output_directory=target,
|
||||
manifest_path=target / "manifest.json",
|
||||
screenshot_path=target / "screenshot.png",
|
||||
hierarchy_path=target / "hierarchy.xml",
|
||||
app_path=target / "app.json",
|
||||
)
|
||||
|
||||
|
||||
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_inputs(serial: object, goods_id: object, state: object, quantity: object) -> None:
|
||||
if type(serial) is not str or not serial.strip() or serial != serial.strip():
|
||||
raise QuantityGate2EvidenceError("必须显式提供非空设备通道。")
|
||||
if type(goods_id) is not str or goods_id != EXPECTED_GOODS_ID:
|
||||
raise QuantityGate2EvidenceError("商品不是 T-105 已批准取证目标。")
|
||||
parse_product_url(f"https://mobile.yangkeduo.com/goods.html?goods_id={goods_id}")
|
||||
if type(state) is not str or state not in DECLARED_QUANTITIES:
|
||||
raise QuantityGate2EvidenceError("人工声明状态无效。")
|
||||
if type(quantity) is not int or quantity != DECLARED_QUANTITIES[state]:
|
||||
raise QuantityGate2EvidenceError("人工声明数量与批准状态不一致。")
|
||||
|
||||
|
||||
def _validate_new_target(target: Path) -> None:
|
||||
if target.exists() or not target.name:
|
||||
raise QuantityGate2EvidenceError("输出目录必须是不存在的明确新目录。")
|
||||
|
||||
|
||||
def _require_expected_device(inspection: DeviceInspection) -> None:
|
||||
if inspection.model != EXPECTED_DEVICE_MODEL or inspection.android_version != EXPECTED_ANDROID_VERSION:
|
||||
raise QuantityGate2EvidenceError("设备不是已批准取证组合。")
|
||||
|
||||
|
||||
def _require_read_precondition(
|
||||
device: QuantityGate2ReadDevice,
|
||||
current: object,
|
||||
) -> dict[str, str]:
|
||||
info = 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 QuantityGate2EvidenceError("拼多多版本不是已批准取证版本。")
|
||||
if not isinstance(current, dict) or current.get("package") != PDD_PACKAGE:
|
||||
raise QuantityGate2EvidenceError("拼多多不在前台。")
|
||||
activity = current.get("activity")
|
||||
if not isinstance(activity, str) or not activity.strip():
|
||||
raise QuantityGate2EvidenceError("前台应用摘要不完整。")
|
||||
if device.window_size() != EXPECTED_SCREEN_SIZE:
|
||||
raise QuantityGate2EvidenceError("屏幕坐标空间不是已批准尺寸。")
|
||||
return {"package": PDD_PACKAGE, "activity": activity, "pdd_version": EXPECTED_PDD_VERSION}
|
||||
|
||||
|
||||
def _read_rpc(
|
||||
device: QuantityGate2ReadDevice,
|
||||
method: str,
|
||||
params: object,
|
||||
timeout_seconds: float,
|
||||
) -> object:
|
||||
return device.jsonrpc_call(method, params, timeout=timeout_seconds)
|
||||
|
||||
|
||||
def _require_screenshot_size(path: Path) -> None:
|
||||
with Image.open(path) as image:
|
||||
image.load()
|
||||
if image.size != EXPECTED_SCREEN_SIZE or image.format != "PNG":
|
||||
raise QuantityGate2EvidenceError("数量状态截图格式或尺寸无效。")
|
||||
|
||||
|
||||
def _clean_staging(staging: Path | None) -> None:
|
||||
if staging is not None and staging.exists():
|
||||
shutil.rmtree(staging)
|
||||
|
||||
|
||||
def _manifest(
|
||||
inspection: DeviceInspection,
|
||||
serial: str,
|
||||
goods_id: str,
|
||||
state: str,
|
||||
quantity: int,
|
||||
screenshot_path: Path,
|
||||
hierarchy_path: Path,
|
||||
app_path: Path,
|
||||
) -> dict[str, Any]:
|
||||
return {
|
||||
"schema_version": 1,
|
||||
"operation": "t105-quantity-gate2-readonly-evidence",
|
||||
"captured_at": datetime.now(UTC).isoformat(),
|
||||
"product": {
|
||||
"goods_id": goods_id,
|
||||
"canonical_url": f"https://mobile.yangkeduo.com/goods.html?goods_id={goods_id}",
|
||||
},
|
||||
"human_declared_state": state,
|
||||
"human_declared_quantity": quantity,
|
||||
"human_declared_selection": TARGET_SELECTION,
|
||||
"review_status": "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,
|
||||
"role": "quantity_state_raw_screenshot",
|
||||
"sha256": _sha256_file(screenshot_path),
|
||||
},
|
||||
{
|
||||
"path": hierarchy_path.name,
|
||||
"role": "quantity_state_raw_hierarchy_local_only",
|
||||
"sha256": _sha256_file(hierarchy_path),
|
||||
},
|
||||
{
|
||||
"path": app_path.name,
|
||||
"role": "quantity_state_app_identity",
|
||||
"sha256": _sha256_file(app_path),
|
||||
},
|
||||
],
|
||||
}
|
||||
@@ -0,0 +1,307 @@
|
||||
"""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, 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,
|
||||
_parse_nodes,
|
||||
_require_color_only_panel,
|
||||
_require_safe_reveal_path,
|
||||
_revealed_unselected_projection,
|
||||
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_FAILURE_STAGES = frozenset(
|
||||
(
|
||||
"reveal_precondition",
|
||||
"reveal_attempted",
|
||||
"reveal_candidate",
|
||||
"reveal_after",
|
||||
"reveal_publish",
|
||||
)
|
||||
)
|
||||
_REVEAL_FAILURE_MARKER = object()
|
||||
|
||||
|
||||
class SkuRevealSpikeError(RuntimeError):
|
||||
"""一次性 reveal 取证未形成可发布证据。"""
|
||||
|
||||
|
||||
def _annotate_reveal_failure(error: BaseException, stage: str) -> None:
|
||||
"""只记录本模块实际走到的固定阶段;入口阶段使用独立 marker,不会被覆盖。"""
|
||||
|
||||
if type(stage) is not str or stage not in _REVEAL_FAILURE_STAGES:
|
||||
return
|
||||
try:
|
||||
setattr(error, "_cmbuyer_reveal_failure_stage", stage)
|
||||
# marker 最后写入,任一 setter 失败都不能形成可信阶段。
|
||||
setattr(error, "_cmbuyer_reveal_failure_marker", _REVEAL_FAILURE_MARKER)
|
||||
except BaseException:
|
||||
pass
|
||||
|
||||
|
||||
def safe_reveal_failure_stage(error: BaseException) -> str | None:
|
||||
"""读取可公开的 reveal 控制流阶段;伪造属性或 hostile getter 均失败闭合。"""
|
||||
|
||||
try:
|
||||
marker = getattr(error, "_cmbuyer_reveal_failure_marker", None)
|
||||
stage = getattr(error, "_cmbuyer_reveal_failure_stage", None)
|
||||
if marker is not _REVEAL_FAILURE_MARKER or type(stage) is not str:
|
||||
return None
|
||||
return stage if stage in _REVEAL_FAILURE_STAGES else None
|
||||
except BaseException:
|
||||
return None
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SkuRevealSpikeResult:
|
||||
output_directory: Path
|
||||
manifest_path: Path
|
||||
|
||||
|
||||
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:
|
||||
stage = "reveal_precondition"
|
||||
staging: Path | None = None
|
||||
adapter: UiautomatorSkuPanelAdapter | None = None
|
||||
try:
|
||||
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 = _prepare_staging(target)
|
||||
deadline = self._clock() + self._timeout_seconds
|
||||
inspection = self._adb_client.inspect(serial)
|
||||
_require_expected_device(inspection)
|
||||
adapter = UiautomatorSkuPanelAdapter(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)
|
||||
# spike 只复用生产 Flow 的私有前置准备,不调用完整选择流程,避免重复 reveal/M。
|
||||
flow._prepare_reveal_precondition(
|
||||
resolve_task_selection("黑色CHA(纯棉)", "M(建议100-115)")
|
||||
)
|
||||
|
||||
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"
|
||||
# 此后即属于 attempted:adapter 会在 RPC 前封存唯一机会,结果不明也只能调和。
|
||||
stage = "reveal_attempted"
|
||||
try:
|
||||
adapter.reveal_size_options_once()
|
||||
except SkuSelectionRunError:
|
||||
rpc_outcome = "ambiguous_reconciled"
|
||||
|
||||
stage = "reveal_candidate"
|
||||
projection, after_hierarchy = self._wait_for_candidate(adapter, deadline)
|
||||
stage = "reveal_after"
|
||||
after_directory = staging / "after"
|
||||
_capture_frame(adapter, after_directory, after_hierarchy)
|
||||
reverified = adapter.dump_window_hierarchy()
|
||||
if _revealed_unselected_projection(_parse_nodes(reverified)) != projection:
|
||||
raise SkuRevealSpikeError("截图后候选状态漂移,未发布证据。")
|
||||
(after_directory / "hierarchy.xml").write_text(reverified, encoding="utf-8")
|
||||
|
||||
stage = "reveal_publish"
|
||||
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) as error:
|
||||
_clean_staging(staging)
|
||||
failure_stage = stage
|
||||
if stage == "reveal_attempted" and (adapter is None or not adapter.reveal_attempted):
|
||||
# attempted 只能在 adapter 已于 RPC 前封存唯一机会后成立。
|
||||
failure_stage = "reveal_precondition"
|
||||
_annotate_reveal_failure(error, failure_stage)
|
||||
raise
|
||||
except Exception as error:
|
||||
_clean_staging(staging)
|
||||
mapped = SkuRevealSpikeError("规格 reveal 取证未完成,未发布本地证据目录。")
|
||||
failure_stage = stage
|
||||
if stage == "reveal_attempted" and (adapter is None or not adapter.reveal_attempted):
|
||||
failure_stage = "reveal_precondition"
|
||||
_annotate_reveal_failure(mapped, failure_stage)
|
||||
raise mapped from error
|
||||
|
||||
return SkuRevealSpikeResult(target, target / "manifest.json")
|
||||
|
||||
def _wait_for_candidate(
|
||||
self,
|
||||
adapter: UiautomatorSkuPanelAdapter,
|
||||
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 = _revealed_unselected_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 _capture_frame(adapter: UiautomatorSkuPanelAdapter, 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)
|
||||
)
|
||||
File diff suppressed because it is too large
Load Diff
@@ -21,16 +21,31 @@ 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 ..device.baseline import (
|
||||
PDD_PACKAGE,
|
||||
SCREENSHOT_PARAMS,
|
||||
_save_base64_screenshot,
|
||||
_sha256_file,
|
||||
_validate_hierarchy,
|
||||
)
|
||||
from .product_open import EXPECTED_PDD_VERSION
|
||||
from .product_url import ProductUrl, parse_product_url
|
||||
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,
|
||||
_REVEAL_GESTURE,
|
||||
_ENTRY_TEXT_BOUNDS,
|
||||
_TARGET_COLOR_UNROLLED_BOUNDS,
|
||||
_TARGET_SIZE_BOUNDS,
|
||||
_action_bounds,
|
||||
_annotate_sku_entry_failure,
|
||||
_parse_nodes,
|
||||
_safe_sku_entry_failure_stage,
|
||||
_unit_price,
|
||||
resolve_task_selection,
|
||||
)
|
||||
|
||||
@@ -39,6 +54,25 @@ 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 运行未完整完成;错误文本不携带设备或页面原文。"""
|
||||
@@ -60,6 +94,52 @@ class SkuSelectionDeviceAdapterError(SkuSelectionRunError):
|
||||
"""第三方设备接口失败的脱敏映射。"""
|
||||
|
||||
|
||||
class SkuExitSpikeError(SkuSelectionRunError):
|
||||
"""T-104 阶段 A 未形成完整的本机退出证据。"""
|
||||
|
||||
|
||||
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 摘要。"""
|
||||
@@ -68,12 +148,21 @@ class SkuSelectionRunResult:
|
||||
screenshot_path: Path
|
||||
manifest_path: Path
|
||||
unit_price: str
|
||||
captured_at: datetime
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SkuExitSpikeResult:
|
||||
"""已原子发布、仍需人工判断页面身份的 T-104 证据。"""
|
||||
|
||||
output_directory: Path
|
||||
manifest_path: Path
|
||||
|
||||
|
||||
class UiautomatorSkuPanelAdapter(SkuPanelDevice):
|
||||
"""把 uiautomator2 缩为 T-103 所需的读取与三种命名操作。
|
||||
"""把 uiautomator2 缩为 T-103 所需的读取与四种命名操作。
|
||||
|
||||
``tap_sku_entry``、``tap_sku_option`` 和 ``leave_sku_panel`` 是仅有的状态改变方法;
|
||||
四个命名方法是仅有的状态改变入口;reveal 的手势参数固定且不向 Flow 暴露;
|
||||
坐标由 Flow 和本类双重检查后才计算中心点,每次调用只执行一次底层动作。
|
||||
"""
|
||||
|
||||
@@ -83,7 +172,14 @@ class UiautomatorSkuPanelAdapter(SkuPanelDevice):
|
||||
self._device = device
|
||||
self._timeout_seconds = timeout_seconds
|
||||
self._entry_was_tapped = False
|
||||
self._entry_bounds_attempted: str | None = None
|
||||
self._entry_rpc_outcome = "not_attempted"
|
||||
self._option_bounds_attempted: set[str] = set()
|
||||
self._option_rpc_outcomes: dict[str, str] = {}
|
||||
self._reveal_attempted = False
|
||||
self._reveal_rpc_outcome = "not_attempted"
|
||||
self._left_panel = False
|
||||
self._back_rpc_outcome = "not_attempted"
|
||||
|
||||
@property
|
||||
def entry_was_tapped(self) -> bool:
|
||||
@@ -95,6 +191,38 @@ class UiautomatorSkuPanelAdapter(SkuPanelDevice):
|
||||
def left_panel(self) -> bool:
|
||||
return self._left_panel
|
||||
|
||||
@property
|
||||
def reveal_attempted(self) -> bool:
|
||||
return self._reveal_attempted
|
||||
|
||||
@property
|
||||
def entry_rpc_outcome(self) -> str:
|
||||
return self._entry_rpc_outcome
|
||||
|
||||
@property
|
||||
def entry_bounds_attempted(self) -> str | None:
|
||||
return self._entry_bounds_attempted
|
||||
|
||||
@property
|
||||
def option_rpc_outcomes(self) -> tuple[tuple[str, str], ...]:
|
||||
return tuple(sorted(self._option_rpc_outcomes.items()))
|
||||
|
||||
@property
|
||||
def reveal_rpc_outcome(self) -> str:
|
||||
return self._reveal_rpc_outcome
|
||||
|
||||
@property
|
||||
def option_attempts(self) -> int:
|
||||
return len(self._option_bounds_attempted)
|
||||
|
||||
@property
|
||||
def back_attempts(self) -> int:
|
||||
return int(self._left_panel)
|
||||
|
||||
@property
|
||||
def back_rpc_outcome(self) -> str:
|
||||
return self._back_rpc_outcome
|
||||
|
||||
def app_info(self, package_name: str) -> dict[str, Any]:
|
||||
value = self._call("app_info", package_name)
|
||||
if not isinstance(value, dict):
|
||||
@@ -114,19 +242,46 @@ class UiautomatorSkuPanelAdapter(SkuPanelDevice):
|
||||
return value
|
||||
|
||||
def tap_sku_entry(self, bounds: str) -> None:
|
||||
if self._entry_was_tapped:
|
||||
raise SkuSelectionDeviceAdapterError("规格入口已经尝试过,拒绝重试。")
|
||||
# 超时也可能表示底层事件已经送达;必须先封存 attempt,后续绝不重试该入口。
|
||||
self._entry_was_tapped = True
|
||||
self._entry_bounds_attempted = bounds
|
||||
self._entry_rpc_outcome = "ambiguous"
|
||||
self._tap_bounds_once(bounds)
|
||||
self._entry_rpc_outcome = "completed"
|
||||
|
||||
def tap_sku_option(self, bounds: str) -> None:
|
||||
if bounds in self._option_bounds_attempted:
|
||||
raise SkuSelectionDeviceAdapterError("同一规格选项已经尝试过,拒绝重试。")
|
||||
# 规格 RPC 也可能送达后超时;按 exact bounds 封存本次唯一机会。
|
||||
self._option_bounds_attempted.add(bounds)
|
||||
self._option_rpc_outcomes[bounds] = "ambiguous"
|
||||
self._tap_bounds_once(bounds)
|
||||
self._option_rpc_outcomes[bounds] = "completed"
|
||||
|
||||
def reveal_size_options_once(self) -> None:
|
||||
if self._reveal_attempted:
|
||||
raise SkuSelectionDeviceAdapterError("规格显示动作已经尝试过,拒绝重试。")
|
||||
# 手势唯一机会在 RPC 前封存;生产 API 不接受方向、坐标或步数参数。
|
||||
self._reveal_attempted = True
|
||||
self._reveal_rpc_outcome = "ambiguous"
|
||||
self._call(
|
||||
"jsonrpc_call",
|
||||
"swipe",
|
||||
list(_REVEAL_GESTURE),
|
||||
timeout=self._timeout_seconds,
|
||||
)
|
||||
self._reveal_rpc_outcome = "completed"
|
||||
|
||||
def leave_sku_panel(self) -> None:
|
||||
if self._left_panel:
|
||||
raise SkuSelectionDeviceAdapterError("规格面板已经执行过返回,已停止操作。")
|
||||
# 底层调用即使报错也可能已把返回事件送达;先封存本次机会,finally 不得再次返回。
|
||||
self._left_panel = True
|
||||
self._back_rpc_outcome = "ambiguous"
|
||||
self._call("jsonrpc_call", "pressKey", ["back"], timeout=self._timeout_seconds)
|
||||
self._back_rpc_outcome = "completed"
|
||||
|
||||
def capture_screenshot(self) -> str:
|
||||
value = self._call("jsonrpc_call", "takeScreenshot", SCREENSHOT_PARAMS, timeout=self._timeout_seconds)
|
||||
@@ -158,6 +313,210 @@ class UiautomatorSkuPanelAdapter(SkuPanelDevice):
|
||||
raise SkuSelectionDeviceAdapterError("规格面板设备操作失败,已停止操作。") from error
|
||||
|
||||
|
||||
class UiautomatorSkuExitAdapter:
|
||||
"""T-104 阶段 A 的窄设备边界:只读能力加唯一一次命名 Back。
|
||||
|
||||
取证脚本不能取得 T-103 的入口、规格选项或 reveal 方法。Back 的唯一机会在
|
||||
JSON-RPC 前封存,因为超时无法证明事件没有送达,任何结果不明都不得重发。
|
||||
"""
|
||||
|
||||
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._back_attempted = False
|
||||
self._back_rpc_outcome = "not_attempted"
|
||||
|
||||
@property
|
||||
def back_attempts(self) -> int:
|
||||
return int(self._back_attempted)
|
||||
|
||||
@property
|
||||
def back_rpc_outcome(self) -> str:
|
||||
return self._back_rpc_outcome
|
||||
|
||||
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 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 leave_sku_panel(self) -> None:
|
||||
if self._back_attempted:
|
||||
raise SkuSelectionDeviceAdapterError("本次取证已经尝试过返回,拒绝重试。")
|
||||
self._back_attempted = True
|
||||
self._back_rpc_outcome = "ambiguous"
|
||||
self._call(
|
||||
"jsonrpc_call",
|
||||
"pressKey",
|
||||
["back"],
|
||||
timeout=self._timeout_seconds,
|
||||
)
|
||||
self._back_rpc_outcome = "completed"
|
||||
|
||||
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("T-104 设备操作超时,已停止操作。") from error
|
||||
except SkuSelectionRunError:
|
||||
raise
|
||||
except Exception as error:
|
||||
raise SkuSelectionDeviceAdapterError("T-104 设备操作失败,已停止操作。") from error
|
||||
|
||||
|
||||
class SkuExitSpikeCapturer:
|
||||
"""从人工停驻的已验证目标面板执行一次 Back,再只读采集阶段 A 证据。"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
adb_client: AdbClient,
|
||||
connector: Callable[[str], Any],
|
||||
timeout_seconds: float,
|
||||
) -> 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._started = False
|
||||
|
||||
def capture(self, serial: str, output_directory: Path) -> SkuExitSpikeResult:
|
||||
if self._started:
|
||||
raise SkuExitSpikeError("同一取证器不可重复调用。")
|
||||
self._started = True
|
||||
|
||||
staging: Path | None = None
|
||||
adapter: UiautomatorSkuExitAdapter | None = None
|
||||
try:
|
||||
if type(serial) is not str or not serial.strip():
|
||||
raise SkuExitSpikeError("必须显式提供非空设备通道。")
|
||||
target = Path(output_directory)
|
||||
_validate_new_target(target)
|
||||
staging = _prepare_staging(target)
|
||||
|
||||
inspection = self._adb_client.inspect(serial)
|
||||
_require_expected_device(inspection)
|
||||
adapter = UiautomatorSkuExitAdapter(
|
||||
self._connector(serial),
|
||||
self._timeout_seconds,
|
||||
)
|
||||
|
||||
# 第一次只读核验拒绝把任意页面带入 Back 边界;第二次紧邻 Back,覆盖核验期间漂移。
|
||||
_require_t104_exit_precondition(adapter)
|
||||
_require_t104_exit_precondition(adapter)
|
||||
|
||||
rpc_outcome = "completed"
|
||||
try:
|
||||
adapter.leave_sku_panel()
|
||||
except SkuSelectionRunError:
|
||||
if adapter.back_attempts != 1 or adapter.back_rpc_outcome != "ambiguous":
|
||||
raise
|
||||
# RPC 失败不能证明 Back 未送达。只读采集可供人调和,但绝不再发第二次动作。
|
||||
rpc_outcome = "ambiguous_reconciled"
|
||||
|
||||
if adapter.back_attempts != 1:
|
||||
raise SkuExitSpikeError("返回动作审计不完整,未发布任何证据产物。")
|
||||
if rpc_outcome == "completed" and adapter.back_rpc_outcome != "completed":
|
||||
raise SkuExitSpikeError("返回动作结果不完整,未发布任何证据产物。")
|
||||
|
||||
# Back 后先确认仍是已取证版本的 PDD;不能先把其他前台应用写进本地证据。
|
||||
initial_app = _require_t104_post_app(adapter)
|
||||
|
||||
screenshot_path = staging / "post_exit_screenshot.png"
|
||||
_save_base64_screenshot(adapter.capture_screenshot(), screenshot_path)
|
||||
_require_screenshot_size(screenshot_path)
|
||||
|
||||
hierarchy = adapter.dump_window_hierarchy()
|
||||
_validate_hierarchy(hierarchy)
|
||||
hierarchy_path = staging / "post_exit_hierarchy.xml"
|
||||
hierarchy_path.write_text(hierarchy, encoding="utf-8")
|
||||
|
||||
current_app = _require_t104_post_app(adapter)
|
||||
if current_app != initial_app:
|
||||
raise SkuExitSpikeError("退出后应用摘要在采集期间漂移,未发布任何证据产物。")
|
||||
app_path = staging / "post_exit_app.json"
|
||||
app_path.write_text(
|
||||
json.dumps(current_app, ensure_ascii=False, indent=2, sort_keys=True) + "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
|
||||
manifest_path = staging / "manifest.json"
|
||||
manifest_path.write_text(
|
||||
json.dumps(
|
||||
_t104_exit_manifest(
|
||||
inspection,
|
||||
serial,
|
||||
rpc_outcome,
|
||||
screenshot_path,
|
||||
hierarchy_path,
|
||||
app_path,
|
||||
),
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
sort_keys=True,
|
||||
)
|
||||
+ "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
os.rename(staging, target)
|
||||
staging = None
|
||||
except (DeviceConnectionError, SkuSelectionError, SkuSelectionRunError):
|
||||
_clean_staging(staging)
|
||||
raise
|
||||
except (AdbTimeout, HTTPTimeoutError, TimeoutError) as error:
|
||||
_clean_staging(staging)
|
||||
raise SkuExitSpikeError("T-104 取证超时,未发布任何证据产物。") from error
|
||||
except OSError as error:
|
||||
_clean_staging(staging)
|
||||
raise SkuExitSpikeError("T-104 证据目录无法发布,未发布任何证据产物。") from error
|
||||
except Exception as error:
|
||||
_clean_staging(staging)
|
||||
raise SkuExitSpikeError("T-104 取证未完成,未发布任何证据产物。") from error
|
||||
|
||||
return SkuExitSpikeResult(target, target / "manifest.json")
|
||||
|
||||
|
||||
class SkuSelectionRunner:
|
||||
"""只运行 T-103 目标规格恢复、价格确认、原始截图和一次安全退出。"""
|
||||
|
||||
@@ -183,41 +542,57 @@ class SkuSelectionRunner:
|
||||
task_size: str,
|
||||
output_directory: Path,
|
||||
) -> SkuSelectionRunResult:
|
||||
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)
|
||||
|
||||
stage = "precheck"
|
||||
staging: Path | None = None
|
||||
adapter: UiautomatorSkuPanelAdapter | None = None
|
||||
flow: SkuSelectionFlow | None = None
|
||||
staging = _prepare_staging(target)
|
||||
deadline = self._monotonic_clock() + self._timeout_seconds
|
||||
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)
|
||||
screenshot_payload = adapter.capture_screenshot()
|
||||
captured_at = datetime.now(UTC)
|
||||
_save_base64_screenshot(screenshot_payload, screenshot_path)
|
||||
_require_screenshot_size(screenshot_path)
|
||||
except SkuSelectionRunError:
|
||||
raise
|
||||
@@ -226,36 +601,63 @@ class SkuSelectionRunner:
|
||||
|
||||
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()
|
||||
|
||||
_require_completed_action_audit(adapter)
|
||||
|
||||
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",
|
||||
json.dumps(
|
||||
_manifest(
|
||||
inspection,
|
||||
serial,
|
||||
link,
|
||||
screenshot_path,
|
||||
task_color,
|
||||
task_size,
|
||||
adapter,
|
||||
captured_at,
|
||||
),
|
||||
ensure_ascii=False,
|
||||
indent=2,
|
||||
sort_keys=True,
|
||||
)
|
||||
+ "\n",
|
||||
encoding="utf-8",
|
||||
)
|
||||
# Windows 的 rename 不替换既有目标;并发创建 target 时保留其内容并把本次运行判失败。
|
||||
os.rename(staging, target)
|
||||
staging = None
|
||||
except (DeviceConnectionError, SkuSelectionRunError, SkuSelectionError):
|
||||
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)
|
||||
raise SkuSelectionRunTimeoutError("规格面板运行超时,未发布任何证据产物。") from error
|
||||
mapped = SkuSelectionRunTimeoutError("规格面板运行超时,未发布任何证据产物。")
|
||||
_annotate_mapped_failure(mapped, error, stage)
|
||||
raise mapped from error
|
||||
except OSError as error:
|
||||
_clean_staging(staging)
|
||||
raise SkuSelectionRunError("规格面板证据目录无法创建或发布,未发布任何证据产物。") from error
|
||||
mapped = SkuSelectionRunError("规格面板证据目录无法创建或发布,未发布任何证据产物。")
|
||||
_annotate_mapped_failure(mapped, error, stage)
|
||||
raise mapped from error
|
||||
except Exception as error:
|
||||
_clean_staging(staging)
|
||||
raise SkuSelectionRunError("规格面板运行未完成,未发布任何证据产物。") from error
|
||||
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:
|
||||
# 结果不明只允许读回 pending;失败路径绝不继续后续动作或自动 Back。
|
||||
if flow is not None:
|
||||
try:
|
||||
flow.reconcile_pending_action()
|
||||
flow.exit_sku_panel_safely()
|
||||
except (SkuSelectionRunError, SkuSelectionError):
|
||||
pass
|
||||
|
||||
@@ -264,6 +666,7 @@ class SkuSelectionRunner:
|
||||
screenshot_path=target / "screenshot.png",
|
||||
manifest_path=target / "manifest.json",
|
||||
unit_price=EXPECTED_UNIT_PRICE,
|
||||
captured_at=captured_at,
|
||||
)
|
||||
|
||||
|
||||
@@ -310,6 +713,87 @@ def _require_expected_device(inspection: DeviceInspection) -> None:
|
||||
raise SkuSelectionRunError("设备型号或 Android 版本不是已取证组合,已停止操作。")
|
||||
|
||||
|
||||
def _require_t104_exit_precondition(adapter: UiautomatorSkuExitAdapter) -> None:
|
||||
"""fresh 核验 T-103 已人验的唯一 Back 前置,不产生任何页面动作。"""
|
||||
|
||||
_require_expected_version(adapter.app_info(PDD_PACKAGE))
|
||||
current = adapter.app_current()
|
||||
if current.get("package") != PDD_PACKAGE:
|
||||
raise SkuExitSpikeError("Back 前拼多多不在前台,已停止取证。")
|
||||
if adapter.display_size() != EXPECTED_SCREEN_SIZE:
|
||||
raise SkuExitSpikeError("屏幕坐标空间不是已取证尺寸,已停止取证。")
|
||||
if _unit_price(_parse_nodes(adapter.dump_window_hierarchy())) != EXPECTED_UNIT_PRICE:
|
||||
raise SkuExitSpikeError("Back 前目标规格或当前价不匹配,已停止取证。")
|
||||
|
||||
|
||||
def _post_exit_app_evidence(value: object) -> dict[str, str]:
|
||||
"""只保留页面身份复核需要的应用字段,不把第三方返回整体写盘。"""
|
||||
|
||||
if not isinstance(value, dict):
|
||||
raise SkuExitSpikeError("退出后应用摘要无效,未发布任何证据产物。")
|
||||
package = value.get("package")
|
||||
activity = value.get("activity")
|
||||
if not isinstance(package, str) or not package.strip():
|
||||
raise SkuExitSpikeError("退出后应用包名缺失,未发布任何证据产物。")
|
||||
if not isinstance(activity, str) or not activity.strip():
|
||||
raise SkuExitSpikeError("退出后 Activity 缺失,未发布任何证据产物。")
|
||||
return {"package": package, "activity": activity}
|
||||
|
||||
|
||||
def _require_t104_post_app(adapter: UiautomatorSkuExitAdapter) -> dict[str, str]:
|
||||
current = _post_exit_app_evidence(adapter.app_current())
|
||||
if current["package"] != PDD_PACKAGE:
|
||||
raise SkuExitSpikeError("Back 后拼多多不在前台,未采集页面证据。")
|
||||
_require_expected_version(adapter.app_info(PDD_PACKAGE))
|
||||
return current
|
||||
|
||||
|
||||
def _t104_exit_manifest(
|
||||
inspection: DeviceInspection,
|
||||
serial: str,
|
||||
rpc_outcome: str,
|
||||
screenshot_path: Path,
|
||||
hierarchy_path: Path,
|
||||
app_path: Path,
|
||||
) -> dict[str, Any]:
|
||||
"""阶段 A 只写事实与哈希;不声明退出成功,也不授予价格证据角色。"""
|
||||
|
||||
return {
|
||||
"schema_version": 1,
|
||||
"captured_at": datetime.now(UTC).isoformat(),
|
||||
"operation": "t104-sku-exit-evidence",
|
||||
"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": EXPECTED_PDD_VERSION,
|
||||
},
|
||||
"back_attempts": 1,
|
||||
"rpc_outcome": rpc_outcome,
|
||||
"post_exit_status": "human_review_required",
|
||||
"artifacts": [
|
||||
{
|
||||
"path": screenshot_path.name,
|
||||
"role": "post_exit_human_review_only",
|
||||
"sha256": _sha256_file(screenshot_path),
|
||||
},
|
||||
{
|
||||
"path": hierarchy_path.name,
|
||||
"role": "post_exit_raw_hierarchy_local_only",
|
||||
"sha256": _sha256_file(hierarchy_path),
|
||||
},
|
||||
{
|
||||
"path": app_path.name,
|
||||
"role": "post_exit_app_identity_human_review_only",
|
||||
"sha256": _sha256_file(app_path),
|
||||
},
|
||||
],
|
||||
}
|
||||
|
||||
|
||||
def _require_screenshot_size(screenshot_path: Path) -> None:
|
||||
try:
|
||||
with Image.open(screenshot_path) as image:
|
||||
@@ -322,20 +806,83 @@ def _require_screenshot_size(screenshot_path: Path) -> None:
|
||||
raise SkuSelectionScreenshotError("原始截图无效,未发布任何证据产物。") from error
|
||||
|
||||
|
||||
def _manifest(inspection: DeviceInspection, serial: str, link: ProductUrl, screenshot_path: Path, task_color: str, task_size: str) -> dict[str, Any]:
|
||||
def _require_completed_action_audit(adapter: UiautomatorSkuPanelAdapter) -> None:
|
||||
"""发布只接受完整正常链;测试替身或缺动作结果不能伪造真机闭环。"""
|
||||
|
||||
if (
|
||||
not adapter.entry_was_tapped
|
||||
or adapter.entry_bounds_attempted != _ENTRY_TEXT_BOUNDS
|
||||
or adapter.entry_rpc_outcome != "completed"
|
||||
or adapter.option_rpc_outcomes
|
||||
!= tuple(
|
||||
sorted(
|
||||
(
|
||||
(_TARGET_COLOR_UNROLLED_BOUNDS, "completed"),
|
||||
(_TARGET_SIZE_BOUNDS, "completed"),
|
||||
)
|
||||
)
|
||||
)
|
||||
or not adapter.reveal_attempted
|
||||
or adapter.reveal_rpc_outcome != "completed"
|
||||
or adapter.back_attempts != 1
|
||||
or adapter.back_rpc_outcome != "completed"
|
||||
):
|
||||
raise SkuSelectionRunError("规格动作审计链不完整,未发布任何证据产物。")
|
||||
|
||||
|
||||
def _manifest(
|
||||
inspection: DeviceInspection,
|
||||
serial: str,
|
||||
link: ProductUrl,
|
||||
screenshot_path: Path,
|
||||
task_color: str,
|
||||
task_size: str,
|
||||
adapter: UiautomatorSkuPanelAdapter,
|
||||
captured_at: datetime,
|
||||
) -> dict[str, Any]:
|
||||
"""仅写可审计摘要;原始 serial、节点树、页面文案和实际截图内容均不写入 manifest。"""
|
||||
|
||||
option_outcomes = dict(adapter.option_rpc_outcomes)
|
||||
return {
|
||||
"schema_version": 1,
|
||||
"captured_at": datetime.now(UTC).isoformat(),
|
||||
"captured_at": captured_at.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",
|
||||
"panel_status": "verified_before_back",
|
||||
"back_attempts": adapter.back_attempts,
|
||||
"back_rpc_outcome": adapter.back_rpc_outcome,
|
||||
"actions": {
|
||||
"sku_entry": {
|
||||
"attempts": int(adapter.entry_was_tapped),
|
||||
"rpc_outcome": adapter.entry_rpc_outcome,
|
||||
},
|
||||
"target_color": {
|
||||
"attempts": int(_TARGET_COLOR_UNROLLED_BOUNDS in option_outcomes),
|
||||
"rpc_outcome": option_outcomes.get(
|
||||
_TARGET_COLOR_UNROLLED_BOUNDS, "not_attempted"
|
||||
),
|
||||
},
|
||||
"size_reveal": {
|
||||
"attempts": int(adapter.reveal_attempted),
|
||||
"rpc_outcome": adapter.reveal_rpc_outcome,
|
||||
},
|
||||
"target_size": {
|
||||
"attempts": int(_TARGET_SIZE_BOUNDS in option_outcomes),
|
||||
"rpc_outcome": option_outcomes.get(
|
||||
_TARGET_SIZE_BOUNDS, "not_attempted"
|
||||
),
|
||||
},
|
||||
"back": {
|
||||
"attempts": adapter.back_attempts,
|
||||
"rpc_outcome": adapter.back_rpc_outcome,
|
||||
},
|
||||
},
|
||||
"post_exit_status": "same_product_verified",
|
||||
"safe_exit": "completed",
|
||||
"page_identity": "human_review_required",
|
||||
"page_identity": "same_goods_evidence_bound",
|
||||
"channel": "wifi" if ":" in serial else "usb",
|
||||
"serial_sha256": sha256(serial.encode("utf-8")).hexdigest(),
|
||||
"device": {
|
||||
|
||||
@@ -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
|
||||
@@ -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")
|
||||
@@ -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()
|
||||
|
||||
@@ -0,0 +1,5 @@
|
||||
"""采购工具原生 Qt Widgets 界面。"""
|
||||
|
||||
from .main_window import PurchaseToolWindow
|
||||
|
||||
__all__ = ["PurchaseToolWindow"]
|
||||
@@ -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)
|
||||
@@ -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)
|
||||
@@ -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 "未保存")
|
||||
@@ -0,0 +1 @@
|
||||
"""core tests。"""
|
||||
@@ -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
|
||||
@@ -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 = {
|
||||
|
||||
@@ -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))
|
||||
@@ -0,0 +1,321 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
from datetime import datetime, timezone
|
||||
import hashlib
|
||||
import json
|
||||
from pathlib import Path
|
||||
import sqlite3
|
||||
import tempfile
|
||||
import unittest
|
||||
from unittest import mock
|
||||
|
||||
from cmbuyer_client.core.errors import AmbiguousRemoteError, CredentialRemoteError, ManualRemoteError, StateError
|
||||
from cmbuyer_client.core.models import AssetReceipt, ClaimedTask, RenewResult, ScreenshotAsset, SecretToken
|
||||
from cmbuyer_client.localstate.facade import DurableClientGateway
|
||||
from cmbuyer_client.localstate.models import ProfileSettings
|
||||
from cmbuyer_client.localstate.store import LocalStateStore
|
||||
from cmbuyer_client.remote.evidence_sink import HttpEvidenceSink
|
||||
from cmbuyer_client.remote.http_transport import HttpResponse
|
||||
from cmbuyer_client.remote.task_source import HttpTaskSource
|
||||
from tests.core.test_models import ATTEMPT_ID, TASK_ID, claim_wire
|
||||
from tests.localstate.test_store import DEVICE_TOKEN, FakeProtector, PROFILE
|
||||
from tests.remote.test_task_source import DEVICE_ID, FakeTransport, response
|
||||
|
||||
|
||||
PNG = base64.b64decode(
|
||||
"iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNk+A8AAQUBAScY42YAAAAASUVORK5CYII="
|
||||
)
|
||||
|
||||
|
||||
class InspectingSource:
|
||||
def __init__(self, store: LocalStateStore, profile_id: str) -> None:
|
||||
self.store = store
|
||||
self.profile_id = profile_id
|
||||
self.calls = 0
|
||||
self.mode = "success"
|
||||
self.request_ids: list[str] = []
|
||||
|
||||
def claim_next(self, credentials, request):
|
||||
self.calls += 1
|
||||
self.request_ids.append(request.claim_request_id)
|
||||
# HTTP 适配器被调用时,幂等请求必须已经 durable。
|
||||
self.assert_pending(request.claim_request_id)
|
||||
if self.mode == "ambiguous":
|
||||
raise AmbiguousRemoteError("http_result_unknown")
|
||||
if self.mode == "manual":
|
||||
raise ManualRemoteError("claim_requires_manual")
|
||||
return ClaimedTask.from_wire(claim_wire())
|
||||
|
||||
def renew(self, credentials, request):
|
||||
raise AssertionError("not used")
|
||||
|
||||
def assert_pending(self, request_id: str) -> None:
|
||||
snapshot = LocalStateStore(self.store.database_path, FakeProtector()).recovery_snapshot(self.profile_id)
|
||||
if snapshot.pending_claim is None or snapshot.pending_claim.claim_request_id != request_id:
|
||||
raise AssertionError("HTTP happened before durable prepare")
|
||||
|
||||
|
||||
class InspectingSink:
|
||||
def __init__(self, store: LocalStateStore, profile_id: str) -> None:
|
||||
self.store = store
|
||||
self.profile_id = profile_id
|
||||
self.calls = 0
|
||||
|
||||
def upload(self, credentials, upload):
|
||||
self.calls += 1
|
||||
snapshot = LocalStateStore(self.store.database_path, FakeProtector()).recovery_snapshot(self.profile_id)
|
||||
if not snapshot.pending_evidence or snapshot.pending_evidence[0].upload_key != upload.upload_key:
|
||||
raise AssertionError("HTTP happened before durable evidence slot")
|
||||
return AssetReceipt(
|
||||
"63c9f507-7473-4fa6-8d71-8786c34c6301",
|
||||
upload.task_id,
|
||||
upload.attempt_id,
|
||||
upload.kind,
|
||||
upload.privacy_tier,
|
||||
upload.sha256,
|
||||
len(upload.content),
|
||||
"image/png",
|
||||
1,
|
||||
1,
|
||||
upload.captured_at,
|
||||
)
|
||||
|
||||
|
||||
class InspectingRenewSource:
|
||||
def __init__(self, store: LocalStateStore, profile_id: str) -> None:
|
||||
self.store = store
|
||||
self.profile_id = profile_id
|
||||
self.calls = 0
|
||||
|
||||
def claim_next(self, credentials, request):
|
||||
raise AssertionError("not used")
|
||||
|
||||
def renew(self, credentials, request):
|
||||
self.calls += 1
|
||||
snapshot = LocalStateStore(self.store.database_path, FakeProtector()).recovery_snapshot(self.profile_id)
|
||||
if snapshot.pending_renew is None or snapshot.pending_renew.renew_request_id != request.renew_request_id:
|
||||
raise AssertionError("HTTP happened before durable renew")
|
||||
return RenewResult(request.task_id, request.attempt_id, request.claim_generation, request.expected_lease_expires_at)
|
||||
|
||||
|
||||
class DurableClientGatewayTests(unittest.TestCase):
|
||||
def setUp(self) -> None:
|
||||
self.directory = tempfile.TemporaryDirectory()
|
||||
self.database = Path(self.directory.name) / "client-state.sqlite3"
|
||||
self.store = LocalStateStore(
|
||||
self.database,
|
||||
FakeProtector(),
|
||||
now=lambda: datetime(2026, 8, 4, 9, 0, tzinfo=timezone.utc),
|
||||
)
|
||||
profile = ProfileSettings(
|
||||
PROFILE,
|
||||
"http://127.0.0.1:8080",
|
||||
DEVICE_ID,
|
||||
"D:/Portable/adb/adb.exe",
|
||||
"192.168.0.173:5555",
|
||||
"wifi",
|
||||
)
|
||||
self.store.save_profile(profile, SecretToken(DEVICE_TOKEN))
|
||||
self.store.start_or_resume_polling(PROFILE)
|
||||
self.source = InspectingSource(self.store, PROFILE)
|
||||
self.sink = InspectingSink(self.store, PROFILE)
|
||||
self.gateway = DurableClientGateway(self.store, self.source, self.sink)
|
||||
|
||||
def tearDown(self) -> None:
|
||||
self.directory.cleanup()
|
||||
|
||||
def test_claim_unknown_replays_same_durable_key_then_commits(self) -> None:
|
||||
self.source.mode = "ambiguous"
|
||||
with self.assertRaises(AmbiguousRemoteError):
|
||||
self.gateway.claim_next(PROFILE)
|
||||
self.source.mode = "success"
|
||||
claimed = self.gateway.claim_next(PROFILE)
|
||||
self.assertEqual(claimed.task.id, TASK_ID)
|
||||
self.assertEqual(self.source.request_ids[0], self.source.request_ids[1])
|
||||
self.assertIsNotNone(self.store.active_claim(PROFILE))
|
||||
|
||||
def test_unknown_claim_2xx_keeps_pending_key_for_real_adapter_replay(self) -> None:
|
||||
transport = FakeTransport(response(201, claim_wire()))
|
||||
gateway = DurableClientGateway(self.store, HttpTaskSource(transport), self.sink)
|
||||
with self.assertRaises(AmbiguousRemoteError):
|
||||
gateway.claim_next(PROFILE)
|
||||
pending = self.store.recovery_snapshot(PROFILE).pending_claim
|
||||
self.assertIsNotNone(pending)
|
||||
transport.response = response(200, claim_wire())
|
||||
claimed = gateway.claim_next(PROFILE)
|
||||
self.assertEqual(claimed.task.id, TASK_ID)
|
||||
sent = [call[3] for call in transport.calls]
|
||||
self.assertEqual(sent[0], sent[1])
|
||||
|
||||
def test_claim_401_allows_token_repair_and_same_key_replay(self) -> None:
|
||||
transport = FakeTransport(HttpResponse(401, (), b""))
|
||||
gateway = DurableClientGateway(self.store, HttpTaskSource(transport), self.sink)
|
||||
with self.assertRaises(CredentialRemoteError):
|
||||
gateway.claim_next(PROFILE)
|
||||
request_id = self.store.recovery_snapshot(PROFILE).pending_claim.claim_request_id
|
||||
self.store.save_profile(self.store.load_profile(PROFILE).settings, SecretToken("c" * 64))
|
||||
transport.response = response(200, claim_wire())
|
||||
gateway.claim_next(PROFILE)
|
||||
self.assertEqual(transport.calls[0][3], transport.calls[1][3])
|
||||
self.assertNotEqual(dict(transport.calls[0][2])["Authorization"], dict(transport.calls[1][2])["Authorization"])
|
||||
self.assertEqual(json.loads(transport.calls[1][3])["claim_request_id"], request_id)
|
||||
|
||||
def test_profile_read_sql_failure_after_prepare_is_fixed_error_and_zero_http(self) -> None:
|
||||
original_connect = self.store._connect
|
||||
calls = 0
|
||||
|
||||
def fail_second_connection():
|
||||
nonlocal calls
|
||||
calls += 1
|
||||
connection = original_connect()
|
||||
if calls == 2:
|
||||
connection.set_authorizer(
|
||||
lambda action, table, *_: sqlite3.SQLITE_DENY
|
||||
if action == sqlite3.SQLITE_READ and table == "profiles"
|
||||
else sqlite3.SQLITE_OK
|
||||
)
|
||||
return connection
|
||||
|
||||
with mock.patch.object(self.store, "_connect", side_effect=fail_second_connection):
|
||||
with self.assertRaisesRegex(StateError, "localstate_read_failed") as captured:
|
||||
self.gateway.claim_next(PROFILE)
|
||||
self.assertEqual(self.source.calls, 0)
|
||||
self.assertNotIn(str(self.database), repr(captured.exception))
|
||||
|
||||
def test_renew_is_durable_before_http(self) -> None:
|
||||
self.gateway.claim_next(PROFILE)
|
||||
source = InspectingRenewSource(self.store, PROFILE)
|
||||
gateway = DurableClientGateway(self.store, source, self.sink)
|
||||
result = gateway.renew(PROFILE)
|
||||
self.assertEqual(result.attempt_id, ATTEMPT_ID)
|
||||
self.assertEqual(source.calls, 1)
|
||||
|
||||
def test_unknown_renew_2xx_keeps_pending_payload_for_replay(self) -> None:
|
||||
self.gateway.claim_next(PROFILE)
|
||||
payload = {
|
||||
"task_id": TASK_ID,
|
||||
"attempt_id": ATTEMPT_ID,
|
||||
"claim_generation": 1,
|
||||
"lease_expires_at": "2026-08-04T09:06:00Z",
|
||||
}
|
||||
transport = FakeTransport(response(201, payload))
|
||||
gateway = DurableClientGateway(self.store, HttpTaskSource(transport), self.sink)
|
||||
with self.assertRaises(AmbiguousRemoteError):
|
||||
gateway.renew(PROFILE)
|
||||
pending = self.store.recovery_snapshot(PROFILE).pending_renew
|
||||
self.assertIsNotNone(pending)
|
||||
transport.response = response(200, payload)
|
||||
gateway.renew(PROFILE)
|
||||
self.assertEqual(transport.calls[0][3], transport.calls[1][3])
|
||||
|
||||
def test_renew_401_allows_bearer_repair_without_changing_claim_payload(self) -> None:
|
||||
self.gateway.claim_next(PROFILE)
|
||||
payload = {
|
||||
"task_id": TASK_ID,
|
||||
"attempt_id": ATTEMPT_ID,
|
||||
"claim_generation": 1,
|
||||
"lease_expires_at": "2026-08-04T09:06:00Z",
|
||||
}
|
||||
transport = FakeTransport(HttpResponse(401, (), b""))
|
||||
gateway = DurableClientGateway(self.store, HttpTaskSource(transport), self.sink)
|
||||
with self.assertRaises(CredentialRemoteError):
|
||||
gateway.renew(PROFILE)
|
||||
self.assertIsNotNone(self.store.recovery_snapshot(PROFILE).pending_renew)
|
||||
self.store.save_profile(self.store.load_profile(PROFILE).settings, SecretToken("c" * 64))
|
||||
transport.response = response(200, payload)
|
||||
gateway.renew(PROFILE)
|
||||
self.assertEqual(transport.calls[0][3], transport.calls[1][3])
|
||||
self.assertNotEqual(dict(transport.calls[0][2])["Authorization"], dict(transport.calls[1][2])["Authorization"])
|
||||
|
||||
def test_unknown_evidence_2xx_keeps_pending_multipart_for_replay(self) -> None:
|
||||
self.gateway.claim_next(PROFILE)
|
||||
path = Path(self.directory.name) / "unknown.png"
|
||||
path.write_bytes(PNG)
|
||||
asset = ScreenshotAsset(path, TASK_ID, ATTEMPT_ID, "2026-08-04T09:01:00Z")
|
||||
transport = FakeTransport(HttpResponse(202, (), b""))
|
||||
gateway = DurableClientGateway(self.store, self.source, HttpEvidenceSink(transport))
|
||||
with self.assertRaises(AmbiguousRemoteError):
|
||||
gateway.upload_evidence(PROFILE, asset)
|
||||
pending = self.store.recovery_snapshot(PROFILE).pending_evidence
|
||||
self.assertEqual(len(pending), 1)
|
||||
digest = hashlib.sha256(PNG).hexdigest()
|
||||
receipt = {
|
||||
"asset_id": "63c9f507-7473-4fa6-8d71-8786c34c6301",
|
||||
"task_id": TASK_ID,
|
||||
"attempt_id": ATTEMPT_ID,
|
||||
"kind": "SKU_PANEL_GATE_1",
|
||||
"privacy_tier": "INTERNAL_RAW",
|
||||
"sha256": digest,
|
||||
"byte_size": len(PNG),
|
||||
"content_type": "image/png",
|
||||
"width_px": 1,
|
||||
"height_px": 1,
|
||||
"captured_at": "2026-08-04T09:01:00Z",
|
||||
}
|
||||
transport.response = HttpResponse(201, (("Content-Type", "application/json"),), json.dumps(receipt).encode())
|
||||
gateway.upload_evidence(PROFILE, asset)
|
||||
self.assertEqual(transport.calls[0][3], transport.calls[1][3])
|
||||
|
||||
def test_evidence_401_allows_bearer_repair_with_same_file_and_multipart(self) -> None:
|
||||
self.gateway.claim_next(PROFILE)
|
||||
path = Path(self.directory.name) / "credential.png"
|
||||
path.write_bytes(PNG)
|
||||
asset = ScreenshotAsset(path, TASK_ID, ATTEMPT_ID, "2026-08-04T09:01:00Z")
|
||||
transport = FakeTransport(HttpResponse(401, (), b""))
|
||||
gateway = DurableClientGateway(self.store, self.source, HttpEvidenceSink(transport))
|
||||
with self.assertRaises(CredentialRemoteError):
|
||||
gateway.upload_evidence(PROFILE, asset)
|
||||
self.assertEqual(len(self.store.recovery_snapshot(PROFILE).pending_evidence), 1)
|
||||
self.store.save_profile(self.store.load_profile(PROFILE).settings, SecretToken("c" * 64))
|
||||
digest = hashlib.sha256(PNG).hexdigest()
|
||||
receipt = {
|
||||
"asset_id": "63c9f507-7473-4fa6-8d71-8786c34c6301",
|
||||
"task_id": TASK_ID,
|
||||
"attempt_id": ATTEMPT_ID,
|
||||
"kind": "SKU_PANEL_GATE_1",
|
||||
"privacy_tier": "INTERNAL_RAW",
|
||||
"sha256": digest,
|
||||
"byte_size": len(PNG),
|
||||
"content_type": "image/png",
|
||||
"width_px": 1,
|
||||
"height_px": 1,
|
||||
"captured_at": "2026-08-04T09:01:00Z",
|
||||
}
|
||||
transport.response = HttpResponse(201, (("Content-Type", "application/json"),), json.dumps(receipt).encode())
|
||||
gateway.upload_evidence(PROFILE, asset)
|
||||
self.assertEqual(transport.calls[0][3], transport.calls[1][3])
|
||||
self.assertNotEqual(dict(transport.calls[0][2])["Authorization"], dict(transport.calls[1][2])["Authorization"])
|
||||
|
||||
def test_equivalent_captured_at_replays_exact_original_multipart_bytes(self) -> None:
|
||||
self.gateway.claim_next(PROFILE)
|
||||
path = Path(self.directory.name) / "exact-replay.png"
|
||||
path.write_bytes(PNG)
|
||||
transport = FakeTransport(HttpResponse(202, (), b""))
|
||||
gateway = DurableClientGateway(self.store, self.source, HttpEvidenceSink(transport))
|
||||
first = ScreenshotAsset(path, TASK_ID, ATTEMPT_ID, "2026-08-04T09:01:00.1Z")
|
||||
equivalent = ScreenshotAsset(path, TASK_ID, ATTEMPT_ID, "2026-08-04T09:01:00.100000Z")
|
||||
with self.assertRaises(AmbiguousRemoteError):
|
||||
gateway.upload_evidence(PROFILE, first)
|
||||
with self.assertRaises(AmbiguousRemoteError):
|
||||
gateway.upload_evidence(PROFILE, equivalent)
|
||||
self.assertEqual(transport.calls[0][3], transport.calls[1][3])
|
||||
|
||||
def test_manual_claim_is_durable_and_never_gets_new_key(self) -> None:
|
||||
self.source.mode = "manual"
|
||||
with self.assertRaises(ManualRemoteError):
|
||||
self.gateway.claim_next(PROFILE)
|
||||
with self.assertRaises(Exception):
|
||||
self.gateway.claim_next(PROFILE)
|
||||
self.assertEqual(self.source.calls, 1)
|
||||
|
||||
def test_evidence_slot_exists_before_http_and_success_never_reuploads(self) -> None:
|
||||
self.gateway.claim_next(PROFILE)
|
||||
path = Path(self.directory.name) / "one.png"
|
||||
path.write_bytes(PNG)
|
||||
asset = ScreenshotAsset(path, TASK_ID, ATTEMPT_ID, "2026-08-04T09:01:00.120000Z")
|
||||
first = self.gateway.upload_evidence(PROFILE, asset)
|
||||
path.write_bytes(PNG + b"changed")
|
||||
second = self.gateway.upload_evidence(PROFILE, asset)
|
||||
self.assertEqual(first, second)
|
||||
self.assertEqual(self.sink.calls, 1)
|
||||
@@ -0,0 +1,24 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import unittest
|
||||
|
||||
from cmbuyer_client.core.errors import ProtectionError
|
||||
from cmbuyer_client.localstate.protection import DpapiProtector
|
||||
|
||||
|
||||
@unittest.skipUnless(os.name == "nt", "DPAPI 仅在 Windows 验证")
|
||||
class DpapiProtectorTests(unittest.TestCase):
|
||||
def test_current_user_round_trip_purpose_isolation_and_corruption(self) -> None:
|
||||
protector = DpapiProtector()
|
||||
plaintext = b"a" * 64
|
||||
device_purpose = "device-token:default:33c9f507-7473-4fa6-8d71-8786c34c6301"
|
||||
claim_purpose = "claim-token:default:53c9f507-7473-4fa6-8d71-8786c34c6301"
|
||||
ciphertext = protector.protect(plaintext, purpose=device_purpose)
|
||||
self.assertNotIn(plaintext, ciphertext)
|
||||
self.assertEqual(protector.unprotect(ciphertext, purpose=device_purpose), plaintext)
|
||||
with self.assertRaises(ProtectionError):
|
||||
protector.unprotect(ciphertext, purpose=claim_purpose)
|
||||
damaged = ciphertext[:-1] + bytes((ciphertext[-1] ^ 1,))
|
||||
with self.assertRaises(ProtectionError):
|
||||
protector.unprotect(damaged, purpose=device_purpose)
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user