diff --git a/admin/config/config.go b/admin/config/config.go index 54ca16d..1650dd9 100644 --- a/admin/config/config.go +++ b/admin/config/config.go @@ -9,8 +9,17 @@ import ( "os" "path/filepath" "strings" + "time" ) +// OnlineThreshold 是判定客户端"在线"的时间窗: +// 最近一次调接口在这个时间之内就算在线,否则离线。 +// +// 取值要**宽松**:客户端执行长任务期间不会调 claim, +// 太短会被误判成离线。取「轮询周期 + 最长任务时长」量级。 +// 本项目有意不做心跳,见 docs/admin/04-client-api.md §3。 +const OnlineThreshold = 10 * time.Minute + // DataDir 返回可写数据目录,不存在就创建。 // // 打包成 exe 后 = exe 旁边的 data/ diff --git a/admin/handler/api/client_api.go b/admin/handler/api/client_api.go index dbac946..7269deb 100644 --- a/admin/handler/api/client_api.go +++ b/admin/handler/api/client_api.go @@ -9,9 +9,14 @@ package api import ( "database/sql" + "encoding/json" + "log" "net/http" "github.com/gin-gonic/gin" + + "cmautobuy/admin/model" + "cmautobuy/admin/service" ) // Handler 持有接口共用的依赖。 @@ -75,17 +80,81 @@ func (h *Handler) Claim(c *gin.Context) { return } - // TODO(骨架): 1. service.RegisterClient(h.db, clientID, req) - // 新 client_id 就新增,已有就更新 device/capabilities/last_seen_at。 - // name 若已被人工改过,不要覆盖。 + // 1. 注册/更新客户端。**注册就在这里做,没有单独的注册接口。** + // capabilities 原样存起来,将来查问题时能看到客户端当时声明了什么。 + caps, _ := json.Marshal(req.Capabilities) + err := service.RegisterClient(h.db, model.Client{ + ClientID: clientID, + Name: req.Client.Name, // 为空时 repository 会用 clientID 兜底 + DeviceAddress: req.Device.Address, + Platform: req.Device.Platform, + PddPackage: req.Device.PddPackage, + Capabilities: string(caps), + }) + if err != nil { + log.Printf("client_register_failed client_id=%s err=%v", clientID, err) + apiError(c, http.StatusInternalServerError, "CLIENT_REGISTER_FAILED", + "登记客户端失败", true) + return + } - // TODO(骨架): 2. service.ClaimNextTask(h.db, clientID) - // 只返回分配给这个客户端的任务,一次一条。 - // UPDATE ... WHERE task_id = ? AND status = 'assigned' - // 检查影响行数,为 0 说明被抢先了,取下一条。 + // 2. 领取一个任务,只会拿到分配给这个客户端的 + task, err := service.ClaimNextTask(h.db, clientID, req.SupportedTypes) + if err != nil { + log.Printf("task_claim_failed client_id=%s err=%v", clientID, err) + apiError(c, http.StatusInternalServerError, "TASK_CLAIM_FAILED", + "领取任务失败", true) + return + } - // 新客户端第一次来必然拿不到任务(还没人给它分配),返回 204 是正常的。 - c.Status(http.StatusNoContent) + // 3. 没有可领的任务 -> 204 No Content,**不是 200 加空对象**。 + // 新客户端第一次来必然走到这里(还没人给它分配),是正常的。 + if task == nil { + c.Status(http.StatusNoContent) + return + } + + log.Printf("task_claimed client_id=%s task_id=%s type=%s", + clientID, task.TaskID, task.TaskType) + c.JSON(http.StatusOK, gin.H{"task": taskPayload(task)}) +} + +// taskPayload 把任务转成契约里的响应结构。 +// +// **不含租约,也不含 Admin 侧状态**——Client 不关心这些, +// 见 docs/client/04-admin-api-contract.md §5。 +func taskPayload(t *model.Task) gin.H { + payload := gin.H{ + "goods_url": t.PddGoodsURL, + } + if t.PddGoodsID != "" { + payload["goods_id"] = t.PddGoodsID + } + if t.PddOptions != "" { + // pdd_options 存的是 JSON 字符串,要还原成对象再嵌进去, + // 否则 Client 拿到的是一个字符串而不是 {"color":...} + var options map[string]any + if err := json.Unmarshal([]byte(t.PddOptions), &options); err == nil { + payload["options"] = options + } + } + // 采购任务必须带数量和价格上限,这是 Client 的价格保护 + if t.Quantity > 0 { + payload["quantity"] = t.Quantity + } + if t.MaxPriceCent > 0 { + payload["max_price_cent"] = t.MaxPriceCent + } + + return gin.H{ + "id": t.TaskID, + "type": t.TaskType, + "version": t.Version, + "priority": t.Priority, + "payload": payload, + "created_at": t.CreatedAt, + "updated_at": t.UpdatedAt, + } } // SubmitResult 接收成功结果。 diff --git a/admin/handler/web/csrf.go b/admin/handler/web/csrf.go new file mode 100644 index 0000000..0c2f2b1 --- /dev/null +++ b/admin/handler/web/csrf.go @@ -0,0 +1,92 @@ +package web + +import ( + "crypto/rand" + "crypto/subtle" + "encoding/base64" + "net/http" + + "github.com/gin-gonic/gin" +) + +// CSRF 防护,用的是「双提交 Cookie」这个最简单的方案: +// +// 1. 给浏览器种一个随机 token 的 Cookie; +// 2. 每个表单里带一个同值的隐藏字段; +// 3. 提交时比对两者,不一致就拒绝。 +// +// 原理:攻击者的页面可以诱导浏览器带上 Cookie 发请求, +// 但**读不到** Cookie 的值,所以拼不出正确的隐藏字段。 +// +// 没有引第三方库,是因为这个方案本身就几十行, +// 而且多一个依赖就多一个可能要求 Go >= 1.25 的风险。 +// +// **只给页面路由用。** 给 Client 的 /api/v1/client/* 绝不能加—— +// 它不是浏览器、没有 Cookie,加了会直接把它挡在门外。 +const ( + csrfCookieName = "cmautobuy_csrf" + csrfFieldName = "csrf_token" + csrfTokenBytes = 32 +) + +// CSRFMiddleware 返回 Gin 中间件。 +// +// GET 等安全方法只负责发 token;POST 等写操作要校验。 +func CSRFMiddleware() gin.HandlerFunc { + return func(c *gin.Context) { + token, err := c.Cookie(csrfCookieName) + if err != nil || token == "" { + token, err = newCSRFToken() + if err != nil { + c.AbortWithStatus(http.StatusInternalServerError) + return + } + // HttpOnly 必须为 false —— 双提交方案要让页面把值填进表单。 + // 本项目只监听本机,Secure 先留 false,将来上 HTTPS 再打开。 + c.SetCookie(csrfCookieName, token, 12*3600, "/", "", false, false) + } + // 交给模板渲染成隐藏字段 + c.Set(csrfFieldName, token) + + switch c.Request.Method { + case http.MethodGet, http.MethodHead, http.MethodOptions: + c.Next() + return + } + + submitted := c.PostForm(csrfFieldName) + if submitted == "" { + submitted = c.GetHeader("X-CSRF-Token") + } + // 用常数时间比较,避免通过响应快慢猜 token + if subtle.ConstantTimeCompare([]byte(submitted), []byte(token)) != 1 { + c.HTML(http.StatusForbidden, "partials/error", gin.H{ + "Title": "请求被拒绝", + "Message": "表单校验失败(CSRF token 不匹配)。" + + "通常是页面开太久过期了,返回上一页刷新后重试即可。", + }) + c.Abort() + return + } + c.Next() + } +} + +// newCSRFToken 生成一个随机 token。 +func newCSRFToken() (string, error) { + buf := make([]byte, csrfTokenBytes) + if _, err := rand.Read(buf); err != nil { + return "", err + } + return base64.RawURLEncoding.EncodeToString(buf), nil +} + +// csrfToken 从上下文取出当前 token,供页面渲染隐藏字段。 +func csrfToken(c *gin.Context) string { + if v, ok := c.Get(csrfFieldName); ok { + if s, ok := v.(string); ok { + return s + } + } + return "" +} diff --git a/admin/handler/web/others.go b/admin/handler/web/others.go index 8e463cb..dbdf948 100644 --- a/admin/handler/web/others.go +++ b/admin/handler/web/others.go @@ -1,9 +1,13 @@ package web import ( + "fmt" + "log" "net/http" "github.com/gin-gonic/gin" + + "cmautobuy/admin/service" ) // ---------- 2. 顺运宝数据 ---------- @@ -16,7 +20,7 @@ func (h *Handler) SybList(c *gin.Context) { // TODO(骨架): 查货运单,并左联 sku_mappings 得出匹配状态 var rows []gin.H - c.HTML(http.StatusOK, "syb/list", page("syb", "顺运宝数据", gin.H{ + c.HTML(http.StatusOK, "syb/list", page(c, "syb", "顺运宝数据", gin.H{ "Keyword": keyword, "Rows": rows, "Status": "尚未实现:同步或手工录入后这里显示货运单", @@ -78,7 +82,7 @@ func (h *Handler) TaskList(c *gin.Context) { // TODO(骨架): 查 tasks,task_type = 'purchase' var rows []gin.H - c.HTML(http.StatusOK, "task/list", page("tasks", "采购任务", gin.H{ + c.HTML(http.StatusOK, "task/list", page(c, "tasks", "采购任务", gin.H{ "Keyword": keyword, "Rows": rows, "Status": "尚未实现:创建任务后这里显示执行进度", @@ -103,18 +107,49 @@ func (h *Handler) TaskDelete(c *gin.Context) { func (h *Handler) ClientList(c *gin.Context) { keyword := c.Query("name") - // TODO(骨架): 查 clients,并按 last_seen_at 算在线状态 - var rows []gin.H + views, err := service.ListClientViews(h.db, keyword, h.onlineThreshold) + if err != nil { + fail(c, http.StatusInternalServerError, + "读取客户端列表失败,数据没有被改动。请稍后重试,或查看 data/logs/ 里的日志。") + return + } - c.HTML(http.StatusOK, "client/list", page("clients", "客户端列表", gin.H{ + online := 0 + for _, v := range views { + if v.Status == "在线" { + online++ + } + } + status := fmt.Sprintf("共 %d 台客户端 · 在线 %d · 离线 %d", + len(views), online, len(views)-online) + if len(views) == 0 { + status = "还没有客户端。客户端第一次调用领取接口时会自动登记。" + } + + c.HTML(http.StatusOK, "client/list", page(c, "clients", "客户端列表", gin.H{ "Keyword": keyword, - "Rows": rows, - "Status": "尚未实现:客户端第一次调领取接口后会自动出现在这里", + "Rows": views, + "Status": status, })) } // ClientDelete 批量删除客户端。 +// +// 二次确认在前端做(见 static/js/app.js),这里直接删。 +// 删掉之后客户端下次调 claim 会重新登记,属于正常行为。 func (h *Handler) ClientDelete(c *gin.Context) { - // TODO(骨架) - fail(c, http.StatusNotImplemented, "删除功能尚未实现。") + ids := c.PostFormArray("ids") + if len(ids) == 0 { + fail(c, http.StatusBadRequest, "没有选中任何客户端,请勾选后再删除。") + return + } + + n, err := service.DeleteClients(h.db, ids) + if err != nil { + fail(c, http.StatusInternalServerError, + "删除失败,数据没有被改动。请稍后重试,或查看 data/logs/ 里的日志。") + return + } + log.Printf("clients_deleted count=%d", n) + c.Redirect(http.StatusSeeOther, "/clients") } diff --git a/admin/handler/web/shopee.go b/admin/handler/web/shopee.go index ca7f938..8ea8f90 100644 --- a/admin/handler/web/shopee.go +++ b/admin/handler/web/shopee.go @@ -18,7 +18,7 @@ func (h *Handler) ShopeeList(c *gin.Context) { // 注意:搜索要走数据库查询,不要一次查全量再在内存里过滤。 var rows []gin.H - c.HTML(http.StatusOK, "shopee/list", page("shopee", "蝦皮数据", gin.H{ + c.HTML(http.StatusOK, "shopee/list", page(c, "shopee", "蝦皮数据", gin.H{ "Keyword": keyword, "Rows": rows, "Status": "尚未实现:导入后这里显示商品和规格", diff --git a/admin/handler/web/web.go b/admin/handler/web/web.go index dbf3a02..f738f37 100644 --- a/admin/handler/web/web.go +++ b/admin/handler/web/web.go @@ -12,6 +12,7 @@ package web import ( "database/sql" "net/http" + "time" "github.com/gin-gonic/gin" ) @@ -19,48 +20,57 @@ import ( // Handler 持有各页面共用的依赖。 type Handler struct { db *sql.DB + // onlineThreshold 是判定客户端"在线"的时间窗。 + // 取「轮询周期 + 最长任务时长」,宽松一点, + // 免得客户端执行长任务期间被误判成离线。 + onlineThreshold time.Duration } // Register 把四个模块的页面路由挂上去。 // // 四个模块的页面结构完全一致(顶部工具条 / 中间带勾选的表格 / 底部状态条), // 这是有意的,见 docs/admin/05-ui-specification.md §2。 -func Register(r *gin.Engine, db *sql.DB) { - h := &Handler{db: db} +func Register(r *gin.Engine, db *sql.DB, onlineThreshold time.Duration) { + h := &Handler{db: db, onlineThreshold: onlineThreshold} + + // CSRF 只挂在页面路由上。 + // 给 Client 的 /api/v1/client/* 绝不能加——它不是浏览器、没有 Cookie。 + pages := r.Group("/", CSRFMiddleware()) // 打开根路径直接进第一个模块 - r.GET("/", func(c *gin.Context) { + pages.GET("/", func(c *gin.Context) { c.Redirect(http.StatusFound, "/shopee") }) // 1. 蝦皮数据 - r.GET("/shopee", h.ShopeeList) - r.POST("/shopee/import", h.ShopeeImport) - r.POST("/shopee/save", h.ShopeeSave) - r.POST("/shopee/delete", h.ShopeeDelete) - r.POST("/shopee/collect", h.ShopeeCollect) + pages.GET("/shopee", h.ShopeeList) + pages.POST("/shopee/import", h.ShopeeImport) + pages.POST("/shopee/save", h.ShopeeSave) + pages.POST("/shopee/delete", h.ShopeeDelete) + pages.POST("/shopee/collect", h.ShopeeCollect) // 2. 顺运宝数据 - r.GET("/syb", h.SybList) - r.POST("/syb/sync", h.SybSync) - r.POST("/syb/match", h.SybMatch) - r.POST("/syb/create-task", h.SybCreateTask) - r.POST("/syb/delete", h.SybDelete) + pages.GET("/syb", h.SybList) + pages.POST("/syb/sync", h.SybSync) + pages.POST("/syb/match", h.SybMatch) + pages.POST("/syb/create-task", h.SybCreateTask) + pages.POST("/syb/delete", h.SybDelete) // 3. 采购任务 - r.GET("/tasks", h.TaskList) - r.POST("/tasks/delete", h.TaskDelete) + pages.GET("/tasks", h.TaskList) + pages.POST("/tasks/delete", h.TaskDelete) // 4. 客户端列表 - r.GET("/clients", h.ClientList) - r.POST("/clients/delete", h.ClientDelete) + pages.GET("/clients", h.ClientList) + pages.POST("/clients/delete", h.ClientDelete) } -// page 组装每个页面都要的公共数据(导航高亮、标题)。 -func page(active, title string, extra gin.H) gin.H { +// page 组装每个页面都要的公共数据(导航高亮、标题、CSRF token)。 +func page(c *gin.Context, active, title string, extra gin.H) gin.H { data := gin.H{ - "Active": active, - "Title": title, + "Active": active, + "Title": title, + "CSRFToken": csrfToken(c), } for k, v := range extra { data[k] = v diff --git a/admin/main.go b/admin/main.go index 3c72fc3..856fcae 100644 --- a/admin/main.go +++ b/admin/main.go @@ -75,8 +75,8 @@ func main() { r.StaticFS("/static", http.FS(staticSub)) // 4. 路由 - web.Register(r, db) // 给浏览器的页面 - api.Register(r, db) // 给 Client 的接口 + web.Register(r, db, config.OnlineThreshold) // 给浏览器的页面 + api.Register(r, db) // 给 Client 的接口 log.Printf("Admin 已启动: http://%s", *addr) if err := r.Run(*addr); err != nil { diff --git a/admin/model/model.go b/admin/model/model.go index 7bd9ae2..7b2e1a8 100644 --- a/admin/model/model.go +++ b/admin/model/model.go @@ -4,6 +4,32 @@ // 字段含义的权威定义在 docs/admin/03-data-model.md。 package model +import "time" + +// ---------- 时间 ---------- + +// TimeLayout 是全项目统一的时间格式:带时区的 ISO 8601。 +// 库里存 UTC,页面上再转本地时区显示。 +const TimeLayout = time.RFC3339 + +// NowISO 返回当前 UTC 时间的字符串形式。 +// 所有写库的时间戳都要用它,不要各写各的格式。 +func NowISO() string { + return time.Now().UTC().Format(TimeLayout) +} + +// ParseISO 解析库里存的时间字符串。解析不了返回零值和 false。 +func ParseISO(s string) (time.Time, bool) { + if s == "" { + return time.Time{}, false + } + ts, err := time.Parse(TimeLayout, s) + if err != nil { + return time.Time{}, false + } + return ts, true +} + // ---------- 蝦皮 ---------- // CollectStatus 是某个蝦皮商品对应的 PDD 商品数据采到没有。 @@ -154,8 +180,7 @@ type Task struct { // Client 是一台执行任务的客户端。 // -// 注意**没有 Status 字段**:在线状态是算出来的, -// LastSeenAt 在 N 分钟内算在线,否则离线。 +// 注意**没有 Status 字段**:在线状态是算出来的(见 IsOnline), // 存成字段会和真实情况不同步。 type Client struct { ClientID string @@ -168,3 +193,28 @@ type Client struct { CreatedAt string UpdatedAt string } + +// IsOnline 判断客户端此刻算不算在线。 +// +// 规则很简单:最近活动时间在 threshold 之内就算在线。 +// **没有心跳是有意的**,所以客户端执行长任务期间不调接口, +// 可能显示成离线——这是已知且接受的取舍, +// 见 docs/admin/04-client-api.md §3。 +// +// LastSeenAt 解析不了时一律当离线,不要当在线。 +func (c Client) IsOnline(now time.Time, threshold time.Duration) bool { + seen, ok := ParseISO(c.LastSeenAt) + if !ok { + return false + } + return now.Sub(seen) < threshold +} + +// StatusText 返回给页面显示的中文状态。 +// 状态不能只靠颜色区分,必须有文字。 +func (c Client) StatusText(now time.Time, threshold time.Duration) string { + if c.IsOnline(now, threshold) { + return "在线" + } + return "离线" +} diff --git a/admin/repository/client.go b/admin/repository/client.go new file mode 100644 index 0000000..60eb61d --- /dev/null +++ b/admin/repository/client.go @@ -0,0 +1,121 @@ +package repository + +import ( + "database/sql" + "fmt" + "strings" + + "cmautobuy/admin/model" +) + +// UpsertClient 登记或更新一台客户端。 +// +// 注意 name 的处理:**只在第一次注册时写入,之后不再更新**。 +// 这样操作员在 Admin 界面上改成好记的名字后,客户端每次 claim +// 都不会把它覆盖回去。做法是 ON CONFLICT 的 DO UPDATE 里不含 name。 +// +// name 为空时用 clientID 当显示名,保证列表里不出现空白行。 +func UpsertClient(db *sql.DB, c model.Client) error { + if c.ClientID == "" { + return fmt.Errorf("client_id 不能为空") + } + name := c.Name + if name == "" { + name = c.ClientID + } + now := model.NowISO() + + _, err := db.Exec(` + INSERT INTO clients (client_id, name, device_address, platform, + pdd_package, capabilities, + last_seen_at, created_at, updated_at) + VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?) + ON CONFLICT(client_id) DO UPDATE SET + device_address = excluded.device_address, + platform = excluded.platform, + pdd_package = excluded.pdd_package, + capabilities = excluded.capabilities, + last_seen_at = excluded.last_seen_at, + updated_at = excluded.updated_at`, + c.ClientID, name, c.DeviceAddress, c.Platform, + c.PddPackage, c.Capabilities, now, now, now) + if err != nil { + return fmt.Errorf("登记客户端 %s 失败: %w", c.ClientID, err) + } + return nil +} + +// TouchClient 只刷新 last_seen_at。 +// +// claim / result / failure 三个接口都要调。只在 claim 里调的话, +// 客户端执行长任务期间不调 claim,会被误判成离线。 +func TouchClient(db *sql.DB, clientID string) error { + now := model.NowISO() + _, err := db.Exec( + `UPDATE clients SET last_seen_at = ?, updated_at = ? WHERE client_id = ?`, + now, now, clientID) + if err != nil { + return fmt.Errorf("刷新客户端 %s 活动时间失败: %w", clientID, err) + } + return nil +} + +// ListClients 按名称模糊查询客户端,keyword 为空时返回全部。 +func ListClients(db *sql.DB, keyword string) ([]model.Client, error) { + query := `SELECT client_id, name, device_address, platform, pdd_package, + capabilities, last_seen_at, created_at, updated_at + FROM clients` + args := []any{} + + if kw := strings.TrimSpace(keyword); kw != "" { + // 参数化查询,通配符拼在值里而不是 SQL 里 + query += ` WHERE name LIKE ? OR client_id LIKE ?` + like := "%" + kw + "%" + args = append(args, like, like) + } + query += ` ORDER BY last_seen_at DESC, client_id` + + rows, err := db.Query(query, args...) + if err != nil { + return nil, fmt.Errorf("查询客户端列表失败: %w", err) + } + defer rows.Close() + + var out []model.Client + for rows.Next() { + var c model.Client + var name, addr, platform, pkg, caps sql.NullString + if err := rows.Scan(&c.ClientID, &name, &addr, &platform, &pkg, + &caps, &c.LastSeenAt, &c.CreatedAt, &c.UpdatedAt); err != nil { + return nil, fmt.Errorf("读取客户端行失败: %w", err) + } + c.Name = name.String + c.DeviceAddress = addr.String + c.Platform = platform.String + c.PddPackage = pkg.String + c.Capabilities = caps.String + out = append(out, c) + } + return out, rows.Err() +} + +// DeleteClients 批量删除客户端。返回实际删除的条数。 +func DeleteClients(db *sql.DB, clientIDs []string) (int64, error) { + if len(clientIDs) == 0 { + return 0, nil + } + + // 占位符按数量生成,值仍然是参数化传入,不存在注入 + placeholders := strings.TrimSuffix(strings.Repeat("?,", len(clientIDs)), ",") + args := make([]any, len(clientIDs)) + for i, id := range clientIDs { + args[i] = id + } + + res, err := db.Exec( + `DELETE FROM clients WHERE client_id IN (`+placeholders+`)`, args...) + if err != nil { + return 0, fmt.Errorf("删除客户端失败: %w", err) + } + return res.RowsAffected() +} diff --git a/admin/repository/db.go b/admin/repository/db.go index fba2d26..41a4f68 100644 --- a/admin/repository/db.go +++ b/admin/repository/db.go @@ -20,26 +20,45 @@ import ( _ "modernc.org/sqlite" ) -// Open 打开 data/admin.db,并设置必要的 PRAGMA。 +// Open 打开 data/admin.db。 +// +// **PRAGMA 必须写在 DSN 里,不能用 db.Exec("PRAGMA ...") 设置。** +// +// 原因是 Go 的 database/sql 是一个**连接池**:db.Exec 只作用于当时 +// 拿到的那一条连接,池子后来新开的连接完全没执行过那些 PRAGMA。 +// 并发写的时候,没有 busy_timeout 的那些连接会直接报 +// "database is locked (SQLITE_BUSY)",而不是等锁释放。 +// +// 写进 DSN 后,驱动会对**每一条新连接**都应用一遍,这才是对的。 func Open(dataDir string) (*sql.DB, error) { path := filepath.Join(dataDir, "admin.db") - db, err := sql.Open("sqlite", path) + // 三条设置的含义见 docs/admin/03-data-model.md §2: + // busy_timeout 拿不到锁时最多等 5 秒,而不是立刻报错 + // journal_mode WAL 模式,读和写可以同时进行 + // foreign_keys 打开外键约束(SQLite 默认是关的) + dsn := "file:" + path + + "?_pragma=busy_timeout(5000)" + + "&_pragma=journal_mode(WAL)" + + "&_pragma=foreign_keys(1)" + + db, err := sql.Open("sqlite", dsn) if err != nil { return nil, fmt.Errorf("打开数据库 %s 失败: %w", path, err) } - // 这三条见 docs/admin/03-data-model.md §2 - pragmas := []string{ - "PRAGMA foreign_keys = ON", - "PRAGMA journal_mode = WAL", - "PRAGMA busy_timeout = 5000", - } - for _, p := range pragmas { - if _, err := db.Exec(p); err != nil { - db.Close() - return nil, fmt.Errorf("执行 %s 失败: %w", p, err) - } + // SQLite 同一时刻只允许一个写事务。连接数放太开, + // 大量连接会互相抢锁、把 busy_timeout 耗光。 + // 本项目是单机内部工具,并发量很小,限制在个位数足够。 + // + // 注意**不要设成 1**:那样在一个事务里再调用需要连接的代码会死锁。 + db.SetMaxOpenConns(4) + db.SetMaxIdleConns(4) + + // sql.Open 是懒加载的,这里主动连一次,好让配置错误立刻暴露 + if err := db.Ping(); err != nil { + db.Close() + return nil, fmt.Errorf("连接数据库 %s 失败: %w", path, err) } return db, nil } diff --git a/admin/repository/task.go b/admin/repository/task.go new file mode 100644 index 0000000..46d1b9d --- /dev/null +++ b/admin/repository/task.go @@ -0,0 +1,127 @@ +package repository + +import ( + "database/sql" + "fmt" + "strings" + + "cmautobuy/admin/model" +) + +// claimCandidateLimit 是一次最多尝试抢多少条。 +// 抢不到说明被别的客户端拿走了,再试下一条;都抢不到就当作没任务。 +const claimCandidateLimit = 10 + +// ClaimNextTask 为指定客户端领取一个任务。 +// +// 没有可领的任务时返回 (nil, nil) —— 调用方据此返回 204。 +// +// 防并发的做法是**条件更新 + 检查影响行数**:先查出候选, +// 再用 `WHERE task_id = ? AND status = 'assigned'` 去更新, +// 影响行数为 0 就说明被别人抢先了,换下一条。 +// 不用 SELECT ... FOR UPDATE,SQLite 没有那个。 +func ClaimNextTask(db *sql.DB, clientID string, supportedTypes []string) (*model.Task, error) { + if clientID == "" { + return nil, fmt.Errorf("client_id 不能为空") + } + + query := `SELECT task_id FROM tasks + WHERE assigned_client = ? AND status = 'assigned'` + args := []any{clientID} + + // 客户端只声明支持某些类型时,不要给它别的类型 + if len(supportedTypes) > 0 { + placeholders := strings.TrimSuffix(strings.Repeat("?,", len(supportedTypes)), ",") + query += ` AND task_type IN (` + placeholders + `)` + for _, t := range supportedTypes { + args = append(args, t) + } + } + query += ` ORDER BY priority DESC, created_at LIMIT ?` + args = append(args, claimCandidateLimit) + + rows, err := db.Query(query, args...) + if err != nil { + return nil, fmt.Errorf("查询可领任务失败: %w", err) + } + var candidates []string + for rows.Next() { + var id string + if err := rows.Scan(&id); err != nil { + rows.Close() + return nil, fmt.Errorf("读取候选任务失败: %w", err) + } + candidates = append(candidates, id) + } + rows.Close() + if err := rows.Err(); err != nil { + return nil, err + } + + now := model.NowISO() + for _, taskID := range candidates { + res, err := db.Exec(` + UPDATE tasks SET status = 'claimed', claimed_at = ?, updated_at = ? + WHERE task_id = ? AND status = 'assigned'`, + now, now, taskID) + if err != nil { + return nil, fmt.Errorf("领取任务 %s 失败: %w", taskID, err) + } + n, err := res.RowsAffected() + if err != nil { + return nil, err + } + if n == 0 { + continue // 被别的客户端抢先了,换下一条 + } + return GetTask(db, taskID) + } + return nil, nil // 没有可领的任务 +} + +// GetTask 按编号读一条任务。 +func GetTask(db *sql.DB, taskID string) (*model.Task, error) { + var t model.Task + var assigned, claimedAt, sybID, orderNo, goodsID, skuID sql.NullString + var pddGoodsID, pddOptions, resultData, errCode, errMsg, finishedAt sql.NullString + var quantity, maxPrice sql.NullInt64 + + err := db.QueryRow(` + SELECT task_id, task_type, status, version, priority, + assigned_client, claimed_at, + syb_id, order_no, goods_id, shopee_sku_id, + pdd_goods_url, pdd_goods_id, pdd_options, + quantity, max_price_cent, + result_data, error_code, error_message, finished_at, + created_at, updated_at + FROM tasks WHERE task_id = ?`, taskID).Scan( + &t.TaskID, &t.TaskType, &t.Status, &t.Version, &t.Priority, + &assigned, &claimedAt, + &sybID, &orderNo, &goodsID, &skuID, + &t.PddGoodsURL, &pddGoodsID, &pddOptions, + &quantity, &maxPrice, + &resultData, &errCode, &errMsg, &finishedAt, + &t.CreatedAt, &t.UpdatedAt) + if err == sql.ErrNoRows { + return nil, nil + } + if err != nil { + return nil, fmt.Errorf("读取任务 %s 失败: %w", taskID, err) + } + + t.AssignedClient = assigned.String + t.ClaimedAt = claimedAt.String + t.SybID = sybID.String + t.OrderNo = orderNo.String + t.GoodsID = goodsID.String + t.ShopeeSKUID = skuID.String + t.PddGoodsID = pddGoodsID.String + t.PddOptions = pddOptions.String + t.Quantity = int(quantity.Int64) + t.MaxPriceCent = maxPrice.Int64 + t.ResultData = resultData.String + t.ErrorCode = errCode.String + t.ErrorMessage = errMsg.String + t.FinishedAt = finishedAt.String + return &t, nil +} diff --git a/admin/service/client_test.go b/admin/service/client_test.go new file mode 100644 index 0000000..e920d51 --- /dev/null +++ b/admin/service/client_test.go @@ -0,0 +1,358 @@ +package service + +import ( + "database/sql" + "path/filepath" + "sync" + "testing" + "time" + + "cmautobuy/admin/model" + "cmautobuy/admin/repository" +) + +// newTestDB 建一个临时库,测试结束自动删。 +// 用真实的 migrations,这样表结构变了测试会跟着失败。 +func newTestDB(t *testing.T) *sql.DB { + t.Helper() + + db, err := repository.Open(t.TempDir()) + if err != nil { + t.Fatalf("打开测试库失败: %v", err) + } + if err := repository.Migrate(db); err != nil { + t.Fatalf("迁移失败: %v", err) + } + t.Cleanup(func() { db.Close() }) + return db +} + +// insertTask 插一条待领取的任务,供领取相关的测试使用。 +func insertTask(t *testing.T, db *sql.DB, taskID, client string) { + t.Helper() + now := model.NowISO() + _, err := db.Exec(` + INSERT INTO tasks (task_id, task_type, status, assigned_client, + order_no, pdd_goods_url, pdd_goods_id, pdd_options, + quantity, max_price_cent, created_at, updated_at) + VALUES (?, 'purchase', 'assigned', ?, 'TEST-ORDER', + 'https://mobile.yangkeduo.com/goods.html?goods_id=1', '1', + '{"color":"黑色","size":"M码"}', 2, 4200, ?, ?)`, + taskID, client, now, now) + if err != nil { + t.Fatalf("插入测试任务失败: %v", err) + } +} + +// ── 注册 ─────────────────────────────────────────────── + +func TestRegisterClient_新客户端被登记(t *testing.T) { + db := newTestDB(t) + + err := RegisterClient(db, model.Client{ + ClientID: "client-001", Name: "办公室-01", + DeviceAddress: "192.168.0.173:5555", Platform: "android", + }) + if err != nil { + t.Fatalf("注册失败: %v", err) + } + + views, err := ListClientViews(db, "", time.Minute) + if err != nil { + t.Fatalf("查询失败: %v", err) + } + if len(views) != 1 { + t.Fatalf("期望 1 台客户端,实际 %d", len(views)) + } + if views[0].Name != "办公室-01" { + t.Errorf("名称错误:期望 办公室-01,实际 %s", views[0].Name) + } + if views[0].DeviceAddress != "192.168.0.173:5555" { + t.Errorf("设备地址错误:%s", views[0].DeviceAddress) + } +} + +func TestRegisterClient_没上报名称时用编号兜底(t *testing.T) { + db := newTestDB(t) + + if err := RegisterClient(db, model.Client{ClientID: "client-002"}); err != nil { + t.Fatalf("注册失败: %v", err) + } + + views, _ := ListClientViews(db, "", time.Minute) + if views[0].Name != "client-002" { + t.Errorf("期望用 client-002 兜底,实际 %q", views[0].Name) + } +} + +// 这条是本工单的重点:操作员改过名字后,客户端再来注册不能覆盖它。 +func TestRegisterClient_人工改过的名称不被覆盖(t *testing.T) { + db := newTestDB(t) + + // 客户端第一次注册,上报名字是 "默认名" + if err := RegisterClient(db, model.Client{ + ClientID: "client-003", Name: "默认名", Platform: "android", + }); err != nil { + t.Fatalf("首次注册失败: %v", err) + } + + // 操作员在界面上改成好记的名字 + if _, err := db.Exec( + `UPDATE clients SET name = ? WHERE client_id = ?`, + "仓库那台", "client-003"); err != nil { + t.Fatalf("人工改名失败: %v", err) + } + + // 客户端再次 claim,又上报了 "默认名" + if err := RegisterClient(db, model.Client{ + ClientID: "client-003", Name: "默认名", Platform: "android", + DeviceAddress: "10.0.0.9:5555", + }); err != nil { + t.Fatalf("再次注册失败: %v", err) + } + + views, _ := ListClientViews(db, "", time.Minute) + if views[0].Name != "仓库那台" { + t.Errorf("人工改的名字被覆盖了:期望 仓库那台,实际 %s", views[0].Name) + } + // 但设备信息应该被更新 + if views[0].DeviceAddress != "10.0.0.9:5555" { + t.Errorf("设备地址没更新:%s", views[0].DeviceAddress) + } +} + +// ── 在线状态 ─────────────────────────────────────────── + +func TestIsOnline_阈值边界(t *testing.T) { + now := time.Date(2026, 8, 6, 12, 0, 0, 0, time.UTC) + threshold := 10 * time.Minute + + cases := []struct { + name string + lastSeen string + want bool + }{ + {"刚刚活动", now.Add(-1 * time.Second).Format(model.TimeLayout), true}, + {"阈值内", now.Add(-9 * time.Minute).Format(model.TimeLayout), true}, + {"刚好到阈值", now.Add(-10 * time.Minute).Format(model.TimeLayout), false}, + {"超过阈值", now.Add(-11 * time.Minute).Format(model.TimeLayout), false}, + {"时间解析不了一律算离线", "不是时间", false}, + {"空值一律算离线", "", false}, + } + for _, tc := range cases { + t.Run(tc.name, func(t *testing.T) { + c := model.Client{LastSeenAt: tc.lastSeen} + if got := c.IsOnline(now, threshold); got != tc.want { + t.Errorf("期望 %v,实际 %v", tc.want, got) + } + }) + } +} + +func TestTouchClient_刷新活动时间(t *testing.T) { + db := newTestDB(t) + + // 造一台很久没活动的客户端 + if err := RegisterClient(db, model.Client{ClientID: "client-004"}); err != nil { + t.Fatalf("注册失败: %v", err) + } + old := "2020-01-01T00:00:00Z" + if _, err := db.Exec( + `UPDATE clients SET last_seen_at = ? WHERE client_id = ?`, + old, "client-004"); err != nil { + t.Fatalf("改时间失败: %v", err) + } + + views, _ := ListClientViews(db, "", 10*time.Minute) + if views[0].Status != "离线" { + t.Fatalf("改时间后应为离线,实际 %s", views[0].Status) + } + + if err := TouchClient(db, "client-004"); err != nil { + t.Fatalf("TouchClient 失败: %v", err) + } + + views, _ = ListClientViews(db, "", 10*time.Minute) + if views[0].Status != "在线" { + t.Errorf("刷新后应为在线,实际 %s", views[0].Status) + } +} + +// ── 搜索与删除 ───────────────────────────────────────── + +func TestListClientViews_按名称搜索(t *testing.T) { + db := newTestDB(t) + RegisterClient(db, model.Client{ClientID: "c-1", Name: "办公室-01"}) + RegisterClient(db, model.Client{ClientID: "c-2", Name: "仓库-01"}) + + views, err := ListClientViews(db, "办公室", time.Minute) + if err != nil { + t.Fatalf("搜索失败: %v", err) + } + if len(views) != 1 || views[0].ClientID != "c-1" { + t.Errorf("搜索结果不对:%+v", views) + } +} + +func TestDeleteClients_批量删除(t *testing.T) { + db := newTestDB(t) + RegisterClient(db, model.Client{ClientID: "c-1"}) + RegisterClient(db, model.Client{ClientID: "c-2"}) + RegisterClient(db, model.Client{ClientID: "c-3"}) + + n, err := DeleteClients(db, []string{"c-1", "c-3"}) + if err != nil { + t.Fatalf("删除失败: %v", err) + } + if n != 2 { + t.Errorf("期望删除 2 条,实际 %d", n) + } + + views, _ := ListClientViews(db, "", time.Minute) + if len(views) != 1 || views[0].ClientID != "c-2" { + t.Errorf("剩余客户端不对:%+v", views) + } +} + +// ── 领取任务 ─────────────────────────────────────────── + +func TestClaimNextTask_没有任务返回nil(t *testing.T) { + db := newTestDB(t) + RegisterClient(db, model.Client{ClientID: "client-001"}) + + task, err := ClaimNextTask(db, "client-001", []string{"collect", "purchase"}) + if err != nil { + t.Fatalf("领取失败: %v", err) + } + if task != nil { + t.Errorf("没有任务时应返回 nil,实际拿到 %+v", task) + } +} + +func TestClaimNextTask_只领分配给自己的(t *testing.T) { + db := newTestDB(t) + insertTask(t, db, "TASK-A", "client-001") + insertTask(t, db, "TASK-B", "client-002") + + task, err := ClaimNextTask(db, "client-002", []string{"purchase"}) + if err != nil { + t.Fatalf("领取失败: %v", err) + } + if task == nil { + t.Fatal("应该领到 TASK-B,实际没领到") + } + if task.TaskID != "TASK-B" { + t.Errorf("领错了任务:期望 TASK-B,实际 %s", task.TaskID) + } +} + +func TestClaimNextTask_同一任务不会被领两次(t *testing.T) { + db := newTestDB(t) + insertTask(t, db, "TASK-A", "client-001") + + first, _ := ClaimNextTask(db, "client-001", []string{"purchase"}) + if first == nil { + t.Fatal("第一次应该领到任务") + } + + second, err := ClaimNextTask(db, "client-001", []string{"purchase"}) + if err != nil { + t.Fatalf("第二次领取报错: %v", err) + } + if second != nil { + t.Errorf("同一任务被领了两次:%s", second.TaskID) + } +} + +// 并发领取:多个 goroutine 同时抢一条任务,只能有一个拿到。 +func TestClaimNextTask_并发只有一个拿到(t *testing.T) { + db := newTestDB(t) + insertTask(t, db, "TASK-ONLY-ONE", "client-001") + + const workers = 8 + var ( + wg sync.WaitGroup + mu sync.Mutex + gotIt int + lastErr error + ) + for i := 0; i < workers; i++ { + wg.Add(1) + go func() { + defer wg.Done() + task, err := ClaimNextTask(db, "client-001", []string{"purchase"}) + mu.Lock() + defer mu.Unlock() + if err != nil { + lastErr = err + return + } + if task != nil { + gotIt++ + } + }() + } + wg.Wait() + + if lastErr != nil { + t.Fatalf("并发领取出错: %v", lastErr) + } + if gotIt != 1 { + t.Errorf("同一任务被 %d 个客户端领到,期望正好 1 个", gotIt) + } +} + +func TestClaimNextTask_不领不支持的类型(t *testing.T) { + db := newTestDB(t) + insertTask(t, db, "TASK-P", "client-001") // 这是 purchase 类型 + + task, err := ClaimNextTask(db, "client-001", []string{"collect"}) + if err != nil { + t.Fatalf("领取失败: %v", err) + } + if task != nil { + t.Errorf("只支持 collect 的客户端不该拿到 purchase 任务:%s", task.TaskID) + } +} + +func TestClaimNextTask_返回的任务带齐必填字段(t *testing.T) { + db := newTestDB(t) + insertTask(t, db, "TASK-A", "client-001") + + task, _ := ClaimNextTask(db, "client-001", []string{"purchase"}) + if task == nil { + t.Fatal("应该领到任务") + } + + // Client 契约要求:goods_url 必有值 + if task.PddGoodsURL == "" { + t.Error("pdd_goods_url 为空,Client 无法执行") + } + // 采购任务必须带价格保护 + if task.Quantity <= 0 { + t.Errorf("采购任务的数量必须 > 0,实际 %d", task.Quantity) + } + if task.MaxPriceCent <= 0 { + t.Errorf("采购任务的价格上限必须 > 0,实际 %d", task.MaxPriceCent) + } + if task.Status != model.TaskClaimed { + t.Errorf("领取后状态应为 claimed,实际 %s", task.Status) + } +} + +// 确认迁移文件路径拼接没问题(顺带覆盖 repository.Open) +func TestOpen_数据库文件落在指定目录(t *testing.T) { + dir := t.TempDir() + db, err := repository.Open(dir) + if err != nil { + t.Fatalf("打开失败: %v", err) + } + defer db.Close() + + if _, err := repository.Migrate(db), error(nil); err != nil { + t.Fatal(err) + } + if _, err := filepath.Abs(filepath.Join(dir, "admin.db")); err != nil { + t.Fatal(err) + } +} diff --git a/admin/service/service.go b/admin/service/service.go index b43919c..156e90c 100644 --- a/admin/service/service.go +++ b/admin/service/service.go @@ -9,6 +9,10 @@ package service import ( "database/sql" "errors" + "time" + + "cmautobuy/admin/model" + "cmautobuy/admin/repository" ) // ErrNotImplemented 表示该功能还没实现。 @@ -114,13 +118,11 @@ func CreatePurchaseTasks(db *sql.DB, sybIDs []string, clientID string) (created // **没有单独的注册接口,也没有心跳**——注册就在 claim 里做, // 理由见 docs/admin/04-client-api.md §3。 // -// 规则: -// - 新 client_id 就新增,已有就更新 device/capabilities/last_seen_at; -// - name 若已被人工改过,**不要用客户端上报的覆盖**; -// - 客户端没上报 name 时,用 client_id 当显示名。 -func RegisterClient(db *sql.DB, clientID, name, address, platform, pddPackage, capabilitiesJSON string) error { - // TODO(骨架): upsert clients - return ErrNotImplemented +// 关于 name:**只在第一次注册时写入,之后不再更新**。 +// 这样操作员在界面上改成好记的名字后,客户端每次 claim +// 都不会把它覆盖回去。客户端没上报 name 时用 clientID 当显示名。 +func RegisterClient(db *sql.DB, c model.Client) error { + return repository.UpsertClient(db, c) } // TouchClient 刷新 last_seen_at。 @@ -129,6 +131,44 @@ func RegisterClient(db *sql.DB, clientID, name, address, platform, pddPackage, c // 只在 claim 里调的话,客户端执行长任务期间不调 claim, // 会被误判成离线。 func TouchClient(db *sql.DB, clientID string) error { - // TODO(骨架) - return ErrNotImplemented + return repository.TouchClient(db, clientID) +} + +// ClientView 是客户端列表页要显示的一行。 +// Status 是**算出来的**,数据库里没有这个字段。 +type ClientView struct { + model.Client + Status string +} + +// ListClientViews 查客户端列表,并把在线状态算出来。 +func ListClientViews(db *sql.DB, keyword string, threshold time.Duration) ([]ClientView, error) { + clients, err := repository.ListClients(db, keyword) + if err != nil { + return nil, err + } + now := time.Now().UTC() + + views := make([]ClientView, 0, len(clients)) + for _, c := range clients { + views = append(views, ClientView{ + Client: c, + Status: c.StatusText(now, threshold), + }) + } + return views, nil +} + +// DeleteClients 批量删除,返回实际删除条数。 +func DeleteClients(db *sql.DB, clientIDs []string) (int64, error) { + return repository.DeleteClients(db, clientIDs) +} + +// ---------- 领取任务 ---------- + +// ClaimNextTask 为客户端领取一个任务,没有可领的返回 (nil, nil)。 +// +// 调用方拿到 nil 要返回 204 No Content,**不是 200 加空对象**。 +func ClaimNextTask(db *sql.DB, clientID string, supportedTypes []string) (*model.Task, error) { + return repository.ClaimNextTask(db, clientID, supportedTypes) } diff --git a/admin/templates/client/list.html b/admin/templates/client/list.html index f06a571..4f99491 100644 --- a/admin/templates/client/list.html +++ b/admin/templates/client/list.html @@ -8,7 +8,8 @@ -
@@ -27,10 +28,17 @@ {{range .Rows}} - {{/* TODO(骨架): 行渲染。 - 「状态」是算出来的:最近活动在 N 分钟内为在线,否则离线。 - 数据库里没有 status 字段,存成字段会和真实情况不同步。 - 名称可以人工改成好记的,改过之后客户端上报的名称不再覆盖它。 */}} +