diff --git a/cmd/nixmsg/serve.go b/cmd/nixmsg/serve.go index 192dfe3..876288a 100644 --- a/cmd/nixmsg/serve.go +++ b/cmd/nixmsg/serve.go @@ -169,6 +169,10 @@ func runServe(ctx context.Context, cfg config.Config) error { Locks: loginLocks, Logger: slog.Default(), TrustedProxies: trustedNets, + Identity: idApp, + Groups: groupApp, + Config: cfg, + Version: Version, KickEndpoint: func(kickCtx context.Context, endpointID string) (bool, error) { if _, found := brk.ConnInfoOf(endpointID); !found { return false, nil diff --git a/docs/DEVIATIONS.md b/docs/DEVIATIONS.md index c06cfcf..6f82f17 100644 --- a/docs/DEVIATIONS.md +++ b/docs/DEVIATIONS.md @@ -645,9 +645,54 @@ - 备选方案:扩展 locks 接口。 - 影响:仅 IP 档锁定时列表可能仍显示未锁定,但 unlock 可解除。 -4. **仍未挂载到 `cmd/nixmsg`** - - 同 A1:可挂载 Handler;接线与 `KickEndpoint` 注入留给总控/后续。 - - 影响:集成测试需自行挂载 Handler。 +4. **A2 时未挂载到 `cmd/nixmsg`;L-WIRE / A3 已接线** + - A2 当时:可挂载 Handler;接线与 `KickEndpoint` 注入留给总控。 + - 现状:`serve` 已挂载 Admin Handler;A3 起注入 `Groups`/`Config`/`Version`/`KickEndpoint`。 + - 影响:无。 + +### A3 2026-09-30 + +1. **概览增加 `endpoints_self`(自助注册数)** + - 原条款:PRD F17「端数量(其中自助注册的数量)」;`admin-api.md` 示例未列该字段。 + - 实际做法:`GET /api/admin/overview` 同时返回契约字段与 `endpoints_self`(`source='self'` 计数)。 + - 原因:以 PRD / 本任务说明为准补齐。 + - 备选方案:只返回 api 文档字段,自助数由前端筛端列表。 + - 影响:W 线可选用该字段;旧 mock 类型可增补。 + +2. **群列表/详情直接读库;变更走 `group.Service`** + - 原条款:群操作调用 `internal/app/group` 已有方法。 + - 实际做法:创建/加人用 `AdminCreate`/`AdminAddMembers`;改名/解散/移除/转让先查群主再以群主为 actor 调 `Rename`/`Dissolve`/`Remove`/`Transfer`。列表与详情(含 `joined_at_ms`)因 `List`/`Get` 按成员可见且 Get 无 joined_at,改为管理侧 SQL。 + - 原因:端协议 API 按成员视角,后台需全局列表。 + - 备选方案:在 group 包增加 AdminList/AdminGet。 + - 影响:列表不依赖 Groups 注入;写操作未注入 Groups 时返回 `503 busy`。 + +3. **后台建群不接受自定义 `id`** + - 原条款:admin-api「id 留空则生成」。 + - 实际做法:`AdminCreate` 无自定义 id 参数;请求带非空 id 返回 `400`。 + - 原因:不改身份线接口签名。 + - 备选方案:扩展 `AdminCreate` 接受可选 id。 + - 影响:后台只能服务器生成群编号。 + +4. **`/metrics` 门禁抽到 `internal/httpx.MetricsGate`** + - 原条款:A 线负责 metrics 访问规则(DEVELOPMENT 4.3)。 + - 实际做法:规则实现放 `httpx.MetricsGate`;`listener.NewMux` RoleShared 调用之(N 目录一行替换)。单独后台监听仍直接挂 Handler 不鉴权。 + - 原因:A 负责目录是 `admin`+`httpx`;listener 仅接线。 + - 备选方案:门禁留在 listener。 + - 影响:`serve` 已传 `MetricsToken`;共用端口无令牌 404、错令牌 401。 + +5. **注册安全码写入审计不含明文** + - 原条款:安全码明文返回已登录管理员,不写日志。 + - 实际做法:GET/PUT 响应含 `code`;`admin_audit` 的 object 为空,不记安全码。存库键 `registration_enabled`=`1`/`0`,与 I1 一致。 + - 原因:D14 / DEVELOPMENT 12 节。 + - 备选:无。 + - 影响:无。 + +6. **消息查询不碰 `message_bodies`** + - 原条款:只读 messages 与 deliveries;响应无正文。 + - 实际做法:SELECT 不含 `body`/`body_enc` 以外的正文列(messages 仍有 `body_enc` 列但查询不选它);不 JOIN `message_bodies`。 + - 原因:任务硬性要求。 + - 备选:无。 + - 影响:无。 ## 后台网页 W diff --git a/internal/admin/a3_test.go b/internal/admin/a3_test.go new file mode 100644 index 0000000..14f55df --- /dev/null +++ b/internal/admin/a3_test.go @@ -0,0 +1,287 @@ +package admin_test + +import ( + "context" + "database/sql" + "encoding/json" + "net/http" + "net/http/cookiejar" + "net/http/httptest" + "path/filepath" + "strings" + "testing" + "time" + + "git.asio.asia/nixevol/NixMsg/internal/admin" + "git.asio.asia/nixevol/NixMsg/internal/app/group" + "git.asio.asia/nixevol/NixMsg/internal/auth" + "git.asio.asia/nixevol/NixMsg/internal/config" + "git.asio.asia/nixevol/NixMsg/internal/store" +) + +func setupA3(t *testing.T) (*store.DB, *httptest.Server, *http.Client) { + t.Helper() + dir := t.TempDir() + db, err := store.Open(filepath.Join(dir, "data"), "FULL") + if err != nil { + t.Fatal(err) + } + t.Cleanup(func() { _ = db.Close() }) + + hash := auth.NewStubHashPool() + if seedErr := admin.SeedAdminPassword(context.Background(), db, hash, testPassword); seedErr != nil { + t.Fatal(seedErr) + } + gApp := group.New(group.Config{DB: db, MaxGroupMembers: 100}) + h := admin.New(admin.Deps{ + DB: db, + Hash: hash, + Tokens: admin.NewRandomAPITokens(), + Locks: admin.NewMemoryLoginLocks(), + Groups: gApp, + Config: config.Default(), + Version: "0.1.0-test", + }) + srv := httptest.NewServer(h) + t.Cleanup(srv.Close) + + jar, err := cookiejar.New(nil) + if err != nil { + t.Fatal(err) + } + client := &http.Client{Jar: jar} + login(t, client, srv.URL) + return db, srv, client +} + +func csrf() map[string]string { + return map[string]string{"X-Nixmsg-Request": "1"} +} + +func insertEndpoint(t *testing.T, db *store.DB, id, source string, enabled bool, online bool) { + t.Helper() + now := time.Now().UnixMilli() + en := 1 + if !enabled { + en = 0 + } + var onlineSince, offlineSince any + if online { + onlineSince = now + } else { + offlineSince = now + } + err := db.Queue.Do(context.Background(), func(tx *sql.Tx) error { + _, e := tx.Exec(`INSERT INTO endpoints( + id, name, remark, source, login_hash, talk_hash, talk_version, + default_delay_ms, enabled, created_at, online_since, offline_since +) VALUES(?,?,?,?,?,?,?,?,?,?,?,?)`, + id, id, "", source, "stub$hash", nil, 0, 0, en, now, onlineSince, offlineSince) + return e + }) + if err != nil { + t.Fatal(err) + } +} + +func TestRegistrationToggleAndGenerate(t *testing.T) { + _, srv, client := setupA3(t) + base := srv.URL + + res := doReq(t, client, http.MethodGet, base+"/api/admin/registration", "", nil) + env := decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("get: %d %+v", res.StatusCode, env) + } + var got map[string]any + _ = json.Unmarshal(env.Data, &got) + if got["enabled"] != false { + t.Fatalf("default enabled=%v", got["enabled"]) + } + + res = doReq(t, client, http.MethodPut, base+"/api/admin/registration", + `{"enabled":true,"code":"abcdefgh"}`, csrf()) + env = decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("put enable: %d %+v", res.StatusCode, env) + } + _ = json.Unmarshal(env.Data, &got) + if got["enabled"] != true || got["code"] != "abcdefgh" { + t.Fatalf("after put: %v", got) + } + + res = doReq(t, client, http.MethodPut, base+"/api/admin/registration", + `{"generate":true}`, csrf()) + env = decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("generate: %d %+v", res.StatusCode, env) + } + _ = json.Unmarshal(env.Data, &got) + code, _ := got["code"].(string) + if len(code) != 16 { + t.Fatalf("generated len=%d code=%q", len(code), code) + } + + res = doReq(t, client, http.MethodPut, base+"/api/admin/registration", + `{"enabled":false}`, csrf()) + env = decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("disable: %d %+v", res.StatusCode, env) + } + _ = json.Unmarshal(env.Data, &got) + if got["enabled"] != false { + t.Fatalf("disabled=%v", got["enabled"]) + } +} + +func TestMessageDetailHasNoBody(t *testing.T) { + db, srv, client := setupA3(t) + base := srv.URL + now := time.Now().UnixMilli() + + err := db.Queue.Do(context.Background(), func(tx *sql.Tx) error { + if _, e := tx.Exec(`INSERT INTO messages( + seq, id, sender_id, dest_kind, dest_id, meta, content_type, body_enc, + send_at, keep, ttl_seconds, receipt, state, reason, created_at +) VALUES(1,'m1','a','endpoint','b','{}','text/plain','enc',?,?,0,1,'dispatched','',?)`, + now, 1, now); e != nil { + return e + } + if _, e := tx.Exec(`INSERT INTO message_bodies(seq, body) VALUES(1, ?)`, + []byte("SECRET_BODY_SHOULD_NOT_APPEAR")); e != nil { + return e + } + _, e := tx.Exec(`INSERT INTO deliveries( + seq, endpoint_id, send_at, keep, state, reason, attempts, updated_at, pushed_at +) VALUES(1,'b',?,1,'pending','',2,?,?)`, now, now, now) + return e + }) + if err != nil { + t.Fatal(err) + } + + res := doReq(t, client, http.MethodGet, base+"/api/admin/messages/1", "", nil) + env := decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("detail: %d %+v", res.StatusCode, env) + } + raw := string(env.Data) + if strings.Contains(raw, "body") || strings.Contains(raw, "SECRET_BODY") { + t.Fatalf("response contains body/secret: %s", raw) + } + var detail map[string]any + _ = json.Unmarshal(env.Data, &detail) + dels, _ := detail["deliveries"].([]any) + if len(dels) != 1 { + t.Fatalf("deliveries=%v", detail["deliveries"]) + } + d0 := dels[0].(map[string]any) + if d0["attempts"].(float64) != 2 { + t.Fatalf("attempts=%v", d0["attempts"]) + } + + res = doReq(t, client, http.MethodGet, base+"/api/admin/messages", "", nil) + env = decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("list: %d %+v", res.StatusCode, env) + } + if strings.Contains(string(env.Data), "SECRET_BODY") || strings.Contains(string(env.Data), `"body"`) { + t.Fatalf("list leaked body: %s", env.Data) + } +} + +func TestGroupsCRUD(t *testing.T) { + db, srv, client := setupA3(t) + base := srv.URL + insertEndpoint(t, db, "alice", "admin", true, false) + insertEndpoint(t, db, "bob", "admin", true, true) + insertEndpoint(t, db, "carol", "self", true, false) + + res := doReq(t, client, http.MethodPost, base+"/api/admin/groups", + `{"name":"一组","owner_id":"alice","member_ids":["bob","carol"]}`, csrf()) + env := decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("create: %d %+v", res.StatusCode, env) + } + var created struct { + ID string `json:"id"` + OwnerID string `json:"owner_id"` + } + _ = json.Unmarshal(env.Data, &created) + if created.ID == "" || created.OwnerID != "alice" { + t.Fatalf("created=%+v", created) + } + + res = doReq(t, client, http.MethodPatch, base+"/api/admin/groups/"+created.ID, + `{"name":"新名"}`, csrf()) + env = decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("rename: %d %+v", res.StatusCode, env) + } + + res = doReq(t, client, http.MethodPost, base+"/api/admin/groups/"+created.ID+"/transfer", + `{"endpoint_id":"bob"}`, csrf()) + env = decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("transfer: %d %+v", res.StatusCode, env) + } + + res = doReq(t, client, http.MethodDelete, + base+"/api/admin/groups/"+created.ID+"/members/carol", "", csrf()) + env = decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("remove: %d %+v", res.StatusCode, env) + } + + res = doReq(t, client, http.MethodDelete, base+"/api/admin/groups/"+created.ID, "", csrf()) + env = decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("dissolve: %d %+v", res.StatusCode, env) + } +} + +func TestOverviewAndSettings(t *testing.T) { + db, srv, client := setupA3(t) + base := srv.URL + insertEndpoint(t, db, "a1", "admin", true, true) + insertEndpoint(t, db, "s1", "self", true, false) + insertEndpoint(t, db, "d1", "admin", false, false) + + res := doReq(t, client, http.MethodGet, base+"/api/admin/overview", "", nil) + env := decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("overview: %d %+v", res.StatusCode, env) + } + var ov map[string]any + _ = json.Unmarshal(env.Data, &ov) + if ov["version"] != "0.1.0-test" { + t.Fatalf("version=%v", ov["version"]) + } + if ov["endpoints_total"].(float64) != 3 { + t.Fatalf("total=%v", ov["endpoints_total"]) + } + if ov["endpoints_self"].(float64) != 1 { + t.Fatalf("self=%v", ov["endpoints_self"]) + } + if ov["endpoints_online"].(float64) != 1 { + t.Fatalf("online=%v", ov["endpoints_online"]) + } + if ov["endpoints_disabled"].(float64) != 1 { + t.Fatalf("disabled=%v", ov["endpoints_disabled"]) + } + + res = doReq(t, client, http.MethodGet, base+"/api/admin/settings", "", nil) + env = decodeEnv(t, res) + if res.StatusCode != 200 || !env.OK { + t.Fatalf("settings: %d %+v", res.StatusCode, env) + } + raw := string(env.Data) + if strings.Contains(raw, "token") || strings.Contains(raw, "password") { + t.Fatalf("settings leaked secret: %s", raw) + } + var st map[string]any + _ = json.Unmarshal(env.Data, &st) + if st["listen"] == nil || st["limits"] == nil { + t.Fatalf("settings=%v", st) + } +} diff --git a/internal/admin/admin_test.go b/internal/admin/admin_test.go index 3237c0b..d999114 100644 --- a/internal/admin/admin_test.go +++ b/internal/admin/admin_test.go @@ -214,11 +214,11 @@ func TestAPITokenAuthAndRestrictions(t *testing.T) { t.Fatalf("auth=%v", me["auth"]) } - // 普通管理接口鉴权通过(业务 501) + // 普通管理接口鉴权通过(A3 overview) res = doReq(t, tokClient, http.MethodGet, base+"/api/admin/overview", "", hdr) env = decodeEnv(t, res) - if res.StatusCode != http.StatusNotImplemented { - t.Fatalf("overview want 501 got %d %+v", res.StatusCode, env) + if res.StatusCode != http.StatusOK || !env.OK { + t.Fatalf("overview want 200 got %d %+v", res.StatusCode, env) } // 禁止 password / tokens diff --git a/internal/admin/groups.go b/internal/admin/groups.go new file mode 100644 index 0000000..14c122d --- /dev/null +++ b/internal/admin/groups.go @@ -0,0 +1,472 @@ +package admin + +import ( + "context" + "database/sql" + "errors" + "net/http" + "strconv" + "strings" + + "git.asio.asia/nixevol/NixMsg/internal/app/group" + "git.asio.asia/nixevol/NixMsg/internal/httpx" + "git.asio.asia/nixevol/NixMsg/internal/protocol" +) + +func (h *Handler) requireGroups(w http.ResponseWriter) bool { + if h.groups == nil { + httpx.WriteError(w, http.StatusServiceUnavailable, "busy", "群服务未注入") + return false + } + return true +} + +func (h *Handler) handleGroupList(w http.ResponseWriter, r *http.Request) { + q := r.URL.Query() + limit, offset, ok := parsePage(w, q.Get("limit"), q.Get("cursor")) + if !ok { + return + } + query := strings.TrimSpace(q.Get("query")) + + where := "1=1" + args := make([]any, 0, 4) + if query != "" { + where += " AND LOWER(g.name) LIKE ?" + args = append(args, "%"+strings.ToLower(query)+"%") + } + + var total int + countSQL := `SELECT COUNT(*) FROM groups g WHERE ` + where + if err := h.db.Read.QueryRowContext(r.Context(), countSQL, args...).Scan(&total); err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + + listArgs := append(append([]any{}, args...), limit, offset) + rows, err := h.db.Read.QueryContext(r.Context(), ` +SELECT g.id, g.name, g.owner_id, g.created_at, + (SELECT COUNT(*) FROM group_members gm WHERE gm.group_id = g.id) AS cnt +FROM groups g +WHERE `+where+` +ORDER BY g.id ASC +LIMIT ? OFFSET ?`, listArgs...) + if err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + defer func() { _ = rows.Close() }() + + items := make([]map[string]any, 0) + for rows.Next() { + var id, name, owner string + var created int64 + var cnt int + if scanErr := rows.Scan(&id, &name, &owner, &created, &cnt); scanErr != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + items = append(items, map[string]any{ + "id": id, + "name": name, + "owner_id": owner, + "member_count": cnt, + "created_at_ms": created, + }) + } + if err := rows.Err(); err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + + next := "" + if offset+len(items) < total { + next = strconv.Itoa(offset + len(items)) + } + httpx.WriteOK(w, map[string]any{"items": items, "next_cursor": next, "total": total}) +} + +func (h *Handler) handleGroupCreate(w http.ResponseWriter, r *http.Request) { + if !h.requireGroups(w) { + return + } + p, _ := principalFrom(r.Context()) + ip := httpx.ClientIP(r, h.trusted) + + var req struct { + ID string `json:"id"` + Name string `json:"name"` + OwnerID string `json:"owner_id"` + MemberIDs []string `json:"member_ids"` + } + if err := httpx.DecodeJSON(r, &req); err != nil { + h.audit(actorString(p), "group_create", "", "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效") + return + } + if strings.TrimSpace(req.ID) != "" { + // AdminCreate 不接受自定义 id;与契约「留空则生成」一致时忽略非空会误导,故拒绝。 + h.audit(actorString(p), "group_create", req.ID, "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "后台创建群请留空 id,由服务器生成") + return + } + if !protocol.ValidName(req.Name) || req.Name == "" { + h.audit(actorString(p), "group_create", "", "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "名称不合法") + return + } + if req.OwnerID == "" { + h.audit(actorString(p), "group_create", "", "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "缺少 owner_id") + return + } + + res, err := h.groups.AdminCreate(r.Context(), req.Name, req.OwnerID, req.MemberIDs) + if err != nil { + h.writeGroupErr(w, p, "group_create", req.OwnerID, ip, err) + return + } + h.audit(actorString(p), "group_create", res.ID, "ok", ip) + failed := res.Failed + if failed == nil { + failed = []group.MemberFail{} + } + httpx.WriteOK(w, map[string]any{ + "id": res.ID, + "name": res.Name, + "owner_id": res.OwnerID, + "failed": failed, + }) +} + +func (h *Handler) handleGroupGet(w http.ResponseWriter, r *http.Request) { + id := r.PathValue("id") + q := r.URL.Query() + limit, offset, ok := parsePage(w, q.Get("limit"), q.Get("cursor")) + if !ok { + return + } + + var name, owner string + var created int64 + err := h.db.Read.QueryRowContext(r.Context(), + `SELECT name, owner_id, created_at FROM groups WHERE id = ?`, id, + ).Scan(&name, &owner, &created) + if isNoRows(err) { + httpx.WriteError(w, http.StatusNotFound, "not_found", "群不存在") + return + } + if err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + + rows, err := h.db.Read.QueryContext(r.Context(), ` +SELECT gm.endpoint_id, COALESCE(e.name, ''), gm.joined_at, + e.online_since, e.offline_since +FROM group_members gm +LEFT JOIN endpoints e ON e.id = gm.endpoint_id +WHERE gm.group_id = ? +ORDER BY gm.endpoint_id ASC +LIMIT ? OFFSET ?`, id, limit, offset) + if err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + defer func() { _ = rows.Close() }() + + members := make([]map[string]any, 0) + for rows.Next() { + var eid, ename string + var joined int64 + var onlineSince, offlineSince sql.NullInt64 + if scanErr := rows.Scan(&eid, &ename, &joined, &onlineSince, &offlineSince); scanErr != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + online, _, _ := presenceFields(onlineSince, offlineSince) + members = append(members, map[string]any{ + "id": eid, + "name": ename, + "online": online, + "joined_at_ms": joined, + }) + } + if err := rows.Err(); err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + + var memberTotal int + _ = h.db.Read.QueryRowContext(r.Context(), + `SELECT COUNT(*) FROM group_members WHERE group_id = ?`, id).Scan(&memberTotal) + next := "" + if offset+len(members) < memberTotal { + next = strconv.Itoa(offset + len(members)) + } + httpx.WriteOK(w, map[string]any{ + "id": id, + "name": name, + "owner_id": owner, + "created_at_ms": created, + "members": members, + "next_cursor": next, + }) +} + +func (h *Handler) handleGroupRename(w http.ResponseWriter, r *http.Request) { + if !h.requireGroups(w) { + return + } + p, _ := principalFrom(r.Context()) + ip := httpx.ClientIP(r, h.trusted) + id := r.PathValue("id") + + var req struct { + Name string `json:"name"` + } + if err := httpx.DecodeJSON(r, &req); err != nil { + h.audit(actorString(p), "group_rename", id, "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效") + return + } + owner, err := h.groupOwner(r.Context(), id) + if isNoRows(err) { + h.audit(actorString(p), "group_rename", id, "not_found", ip) + httpx.WriteError(w, http.StatusNotFound, "not_found", "群不存在") + return + } + if err != nil { + h.audit(actorString(p), "group_rename", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + err = h.groups.Rename(r.Context(), owner, &protocol.GroupRename{ + V: protocol.Version, Type: protocol.TypeGroupRename, RID: "admin", + GroupID: id, Name: req.Name, + }) + if err != nil { + h.writeGroupErr(w, p, "group_rename", id, ip, err) + return + } + summary, err := h.groupSummary(r.Context(), id) + if err != nil { + h.audit(actorString(p), "group_rename", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + h.audit(actorString(p), "group_rename", id, "ok", ip) + httpx.WriteOK(w, summary) +} + +func (h *Handler) handleGroupDissolve(w http.ResponseWriter, r *http.Request) { + if !h.requireGroups(w) { + return + } + p, _ := principalFrom(r.Context()) + ip := httpx.ClientIP(r, h.trusted) + id := r.PathValue("id") + + owner, err := h.groupOwner(r.Context(), id) + if isNoRows(err) { + h.audit(actorString(p), "group_dissolve", id, "not_found", ip) + httpx.WriteError(w, http.StatusNotFound, "not_found", "群不存在") + return + } + if err != nil { + h.audit(actorString(p), "group_dissolve", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + err = h.groups.Dissolve(r.Context(), owner, &protocol.GroupDissolve{ + V: protocol.Version, Type: protocol.TypeGroupDissolve, RID: "admin", GroupID: id, + }) + if err != nil { + h.writeGroupErr(w, p, "group_dissolve", id, ip, err) + return + } + h.audit(actorString(p), "group_dissolve", id, "ok", ip) + httpx.WriteOK(w, map[string]any{}) +} + +func (h *Handler) handleGroupAddMembers(w http.ResponseWriter, r *http.Request) { + if !h.requireGroups(w) { + return + } + p, _ := principalFrom(r.Context()) + ip := httpx.ClientIP(r, h.trusted) + id := r.PathValue("id") + + var req struct { + MemberIDs []string `json:"member_ids"` + } + if err := httpx.DecodeJSON(r, &req); err != nil { + h.audit(actorString(p), "group_add_members", id, "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效") + return + } + if _, err := h.groupOwner(r.Context(), id); isNoRows(err) { + h.audit(actorString(p), "group_add_members", id, "not_found", ip) + httpx.WriteError(w, http.StatusNotFound, "not_found", "群不存在") + return + } else if err != nil { + h.audit(actorString(p), "group_add_members", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + + res, err := h.groups.AdminAddMembers(r.Context(), id, req.MemberIDs) + if err != nil { + h.writeGroupErr(w, p, "group_add_members", id, ip, err) + return + } + failed := res.Failed + if failed == nil { + failed = []group.MemberFail{} + } + h.audit(actorString(p), "group_add_members", id, "ok", ip) + httpx.WriteOK(w, map[string]any{"failed": failed}) +} + +func (h *Handler) handleGroupRemoveMember(w http.ResponseWriter, r *http.Request) { + if !h.requireGroups(w) { + return + } + p, _ := principalFrom(r.Context()) + ip := httpx.ClientIP(r, h.trusted) + id := r.PathValue("id") + endpointID := r.PathValue("endpointId") + + owner, err := h.groupOwner(r.Context(), id) + if isNoRows(err) { + h.audit(actorString(p), "group_remove_member", id, "not_found", ip) + httpx.WriteError(w, http.StatusNotFound, "not_found", "群不存在") + return + } + if err != nil { + h.audit(actorString(p), "group_remove_member", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + err = h.groups.Remove(r.Context(), owner, &protocol.GroupRemove{ + V: protocol.Version, Type: protocol.TypeGroupRemove, RID: "admin", + GroupID: id, EndpointID: endpointID, + }) + if err != nil { + h.writeGroupErr(w, p, "group_remove_member", id, ip, err) + return + } + h.audit(actorString(p), "group_remove_member", id+"/"+endpointID, "ok", ip) + httpx.WriteOK(w, map[string]any{}) +} + +func (h *Handler) handleGroupTransfer(w http.ResponseWriter, r *http.Request) { + if !h.requireGroups(w) { + return + } + p, _ := principalFrom(r.Context()) + ip := httpx.ClientIP(r, h.trusted) + id := r.PathValue("id") + + var req struct { + EndpointID string `json:"endpoint_id"` + } + if err := httpx.DecodeJSON(r, &req); err != nil { + h.audit(actorString(p), "group_transfer", id, "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效") + return + } + owner, err := h.groupOwner(r.Context(), id) + if isNoRows(err) { + h.audit(actorString(p), "group_transfer", id, "not_found", ip) + httpx.WriteError(w, http.StatusNotFound, "not_found", "群不存在") + return + } + if err != nil { + h.audit(actorString(p), "group_transfer", id, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + err = h.groups.Transfer(r.Context(), owner, &protocol.GroupTransfer{ + V: protocol.Version, Type: protocol.TypeGroupTransfer, RID: "admin", + GroupID: id, EndpointID: req.EndpointID, + }) + if err != nil { + h.writeGroupErr(w, p, "group_transfer", id, ip, err) + return + } + h.audit(actorString(p), "group_transfer", id, "ok", ip) + httpx.WriteOK(w, map[string]any{"owner_id": req.EndpointID}) +} + +func (h *Handler) groupOwner(ctx context.Context, groupID string) (string, error) { + var owner string + err := h.db.Read.QueryRowContext(ctx, `SELECT owner_id FROM groups WHERE id = ?`, groupID).Scan(&owner) + return owner, err +} + +func (h *Handler) groupSummary(ctx context.Context, groupID string) (map[string]any, error) { + var name, owner string + var created int64 + var cnt int + err := h.db.Read.QueryRowContext(ctx, ` +SELECT g.name, g.owner_id, g.created_at, + (SELECT COUNT(*) FROM group_members gm WHERE gm.group_id = g.id) +FROM groups g WHERE g.id = ?`, groupID).Scan(&name, &owner, &created, &cnt) + if err != nil { + return nil, err + } + return map[string]any{ + "id": groupID, + "name": name, + "owner_id": owner, + "member_count": cnt, + "created_at_ms": created, + }, nil +} + +func (h *Handler) writeGroupErr(w http.ResponseWriter, p principal, action, object, ip string, err error) { + var pe *protocol.Error + if errors.As(err, &pe) { + status := http.StatusBadRequest + switch pe.Code { + case protocol.CodeNotFound, protocol.CodeInvalidTarget: + status = http.StatusNotFound + case protocol.CodeForbidden, protocol.CodeNotMember, protocol.CodeOwnerCannotLeave: + status = http.StatusForbidden + case protocol.CodeIDTaken: + status = http.StatusConflict + case protocol.CodeBusy: + status = http.StatusServiceUnavailable + } + h.audit(actorString(p), action, object, pe.Code, ip) + httpx.WriteError(w, status, pe.Code, pe.Message) + return + } + h.audit(actorString(p), action, object, "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") +} + +func parsePage(w http.ResponseWriter, limitStr, cursorStr string) (limit, offset int, ok bool) { + limit = defaultListLimit + if limitStr != "" { + n, err := strconv.Atoi(limitStr) + if err != nil || n < 1 { + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "limit 无效") + return 0, 0, false + } + limit = n + } + if limit > maxListLimit { + limit = maxListLimit + } + if cursorStr != "" { + n, err := strconv.Atoi(cursorStr) + if err != nil || n < 0 { + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "cursor 无效") + return 0, 0, false + } + offset = n + } + return limit, offset, true +} diff --git a/internal/admin/handler.go b/internal/admin/handler.go index 445399d..4f1ea89 100644 --- a/internal/admin/handler.go +++ b/internal/admin/handler.go @@ -7,8 +7,10 @@ import ( "sync" "time" + "git.asio.asia/nixevol/NixMsg/internal/app/group" "git.asio.asia/nixevol/NixMsg/internal/app/identity" "git.asio.asia/nixevol/NixMsg/internal/auth" + "git.asio.asia/nixevol/NixMsg/internal/config" "git.asio.asia/nixevol/NixMsg/internal/store" ) @@ -41,20 +43,31 @@ type Deps struct { KickEndpoint EndpointKickFunc // Identity 端停用/启用/删除级联(I5);nil 时回退为仅改 enabled/删行。 Identity identity.Service + + // Groups 群服务;nil 时群写操作返回 busy。列表/详情可只读库。 + Groups group.Service + // Config 只读运行参数来源;零值用 config.Default()。 + Config config.Config + // Version 概览里的版本号;空则 "dev"。 + Version string } // Handler 是可挂载的管理接口(路由前缀 /api/admin/)。 type Handler struct { - db *store.DB - hash auth.HashPool - tokens auth.APITokens - locks auth.LoginLocks - log *slog.Logger - trusted []*net.IPNet - ttl time.Duration - forceSec bool - kick EndpointKickFunc - identity identity.Service + db *store.DB + hash auth.HashPool + tokens auth.APITokens + locks auth.LoginLocks + log *slog.Logger + trusted []*net.IPNet + ttl time.Duration + forceSec bool + kick EndpointKickFunc + identity identity.Service + groups group.Service + cfg config.Config + version string + startedAt time.Time mux *http.ServeMux @@ -74,19 +87,31 @@ func New(d Deps) *Handler { if ttl <= 0 { ttl = defaultSessionTTL } + cfg := d.Config + if cfg.Listen == "" { + cfg = config.Default() + } + ver := d.Version + if ver == "" { + ver = "dev" + } h := &Handler{ - db: d.DB, - hash: d.Hash, - tokens: d.Tokens, - locks: d.Locks, - log: d.Logger, - trusted: d.TrustedProxies, - ttl: ttl, - forceSec: d.SecureCookies, - kick: d.KickEndpoint, - identity: d.Identity, - mux: http.NewServeMux(), - lastUsed: make(map[string]time.Time), + db: d.DB, + hash: d.Hash, + tokens: d.Tokens, + locks: d.Locks, + log: d.Logger, + trusted: d.TrustedProxies, + ttl: ttl, + forceSec: d.SecureCookies, + kick: d.KickEndpoint, + identity: d.Identity, + groups: d.Groups, + cfg: cfg, + version: ver, + startedAt: time.Now(), + mux: http.NewServeMux(), + lastUsed: make(map[string]time.Time), } h.routes() return h @@ -125,25 +150,19 @@ func (h *Handler) routes() { h.mux.Handle("PUT /api/admin/endpoints/{id}/talk-password", h.auth(h.handleEndpointTalkPassword)) h.mux.Handle("POST /api/admin/endpoints/{id}/unlock", h.auth(h.handleEndpointUnlock)) - // 其余管理路由:鉴权生效,业务暂 501(A3) - for _, p := range stubRoutes { - h.mux.Handle(p, h.auth(h.handleNotImplemented)) - } -} - -var stubRoutes = []string{ - "GET /api/admin/overview", - "GET /api/admin/registration", - "PUT /api/admin/registration", - "GET /api/admin/groups", - "POST /api/admin/groups", - "GET /api/admin/groups/{id}", - "PATCH /api/admin/groups/{id}", - "DELETE /api/admin/groups/{id}", - "POST /api/admin/groups/{id}/members", - "DELETE /api/admin/groups/{id}/members/{endpointId}", - "POST /api/admin/groups/{id}/transfer", - "GET /api/admin/messages", - "GET /api/admin/messages/{seq}", - "GET /api/admin/settings", + // A3 + h.mux.Handle("GET /api/admin/overview", h.auth(h.handleOverview)) + h.mux.Handle("GET /api/admin/registration", h.auth(h.handleRegistrationGet)) + h.mux.Handle("PUT /api/admin/registration", h.auth(h.handleRegistrationPut)) + h.mux.Handle("GET /api/admin/groups", h.auth(h.handleGroupList)) + h.mux.Handle("POST /api/admin/groups", h.auth(h.handleGroupCreate)) + h.mux.Handle("GET /api/admin/groups/{id}", h.auth(h.handleGroupGet)) + h.mux.Handle("PATCH /api/admin/groups/{id}", h.auth(h.handleGroupRename)) + h.mux.Handle("DELETE /api/admin/groups/{id}", h.auth(h.handleGroupDissolve)) + h.mux.Handle("POST /api/admin/groups/{id}/members", h.auth(h.handleGroupAddMembers)) + h.mux.Handle("DELETE /api/admin/groups/{id}/members/{endpointId}", h.auth(h.handleGroupRemoveMember)) + h.mux.Handle("POST /api/admin/groups/{id}/transfer", h.auth(h.handleGroupTransfer)) + h.mux.Handle("GET /api/admin/messages", h.auth(h.handleMessageList)) + h.mux.Handle("GET /api/admin/messages/{seq}", h.auth(h.handleMessageGet)) + h.mux.Handle("GET /api/admin/settings", h.auth(h.handleSettingsGet)) } diff --git a/internal/admin/login.go b/internal/admin/login.go index e5779b1..306184e 100644 --- a/internal/admin/login.go +++ b/internal/admin/login.go @@ -4,7 +4,6 @@ import ( "database/sql" "errors" "net/http" - "strings" "git.asio.asia/nixevol/NixMsg/internal/auth" "git.asio.asia/nixevol/NixMsg/internal/httpx" @@ -157,13 +156,3 @@ func (h *Handler) handlePassword(w http.ResponseWriter, r *http.Request) { h.audit(actorString(p), "password_change", "", "ok", ip) httpx.WriteOK(w, map[string]any{}) } - -func (h *Handler) handleNotImplemented(w http.ResponseWriter, r *http.Request) { - p, _ := principalFrom(r.Context()) - ip := httpx.ClientIP(r, h.trusted) - action := strings.ToLower(r.Method) + " " + r.URL.Path - if isMutating(r.Method) { - h.audit(actorString(p), action, "", "not_implemented", ip) - } - httpx.WriteError(w, http.StatusNotImplemented, "not_implemented", "接口尚未实现") -} diff --git a/internal/admin/messages.go b/internal/admin/messages.go new file mode 100644 index 0000000..ae93c6c --- /dev/null +++ b/internal/admin/messages.go @@ -0,0 +1,301 @@ +package admin + +import ( + "context" + "database/sql" + "encoding/json" + "net/http" + "strconv" + "strings" + + "git.asio.asia/nixevol/NixMsg/internal/httpx" +) + +// 消息查询只读 messages / deliveries,不读 message_bodies,响应不含 body/正文。 + +func (h *Handler) handleMessageList(w http.ResponseWriter, r *http.Request) { + q := r.URL.Query() + limit, offset, ok := parsePage(w, q.Get("limit"), q.Get("cursor")) + if !ok { + return + } + + where, args, ok := buildMessageFilter(w, q) + if !ok { + return + } + + var total int + if err := h.db.Read.QueryRowContext(r.Context(), + `SELECT COUNT(*) FROM messages m WHERE `+where, args...).Scan(&total); err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + + listArgs := append(append([]any{}, args...), limit, offset) + rows, err := h.db.Read.QueryContext(r.Context(), ` +SELECT m.seq, m.id, m.sender_id, m.dest_kind, m.dest_id, m.state, m.reason, + m.send_at, m.created_at, m.keep, m.receipt, m.content_type +FROM messages m +WHERE `+where+` +ORDER BY m.seq DESC +LIMIT ? OFFSET ?`, listArgs...) + if err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + defer func() { _ = rows.Close() }() + + items := make([]map[string]any, 0) + seqs := make([]int64, 0) + for rows.Next() { + var seq int64 + var id, sender, destKind, destID, state, reason, contentType string + var sendAt, created int64 + var keep, receipt int + if scanErr := rows.Scan(&seq, &id, &sender, &destKind, &destID, &state, &reason, + &sendAt, &created, &keep, &receipt, &contentType); scanErr != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + seqs = append(seqs, seq) + items = append(items, map[string]any{ + "seq": seq, + "id": id, + "sender_id": sender, + "dest_kind": destKind, + "dest_id": destID, + "state": state, + "reason": reason, + "send_at_ms": sendAt, + "created_at_ms": created, + "keep": keep != 0, + "receipt": receipt != 0, + "content_type": contentType, + }) + } + if rowsErr := rows.Err(); rowsErr != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + + counts, err := h.deliveryCounts(r.Context(), seqs) + if err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + for i, seq := range seqs { + c := counts[seq] + if c == nil { + c = emptyDeliveryCounts() + } + items[i]["delivery_counts"] = c + } + + next := "" + if offset+len(items) < total { + next = strconv.Itoa(offset + len(items)) + } + httpx.WriteOK(w, map[string]any{"items": items, "next_cursor": next, "total": total}) +} + +func (h *Handler) handleMessageGet(w http.ResponseWriter, r *http.Request) { + seqStr := r.PathValue("seq") + seq, err := strconv.ParseInt(seqStr, 10, 64) + if err != nil || seq < 1 { + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "seq 无效") + return + } + q := r.URL.Query() + limit, offset, ok := parsePage(w, q.Get("limit"), q.Get("cursor")) + if !ok { + return + } + + var id, sender, destKind, destID, state, reason, contentType, meta string + var sendAt, created int64 + err = h.db.Read.QueryRowContext(r.Context(), ` +SELECT id, sender_id, dest_kind, dest_id, state, reason, send_at, created_at, + content_type, meta +FROM messages WHERE seq = ?`, seq).Scan( + &id, &sender, &destKind, &destID, &state, &reason, &sendAt, &created, + &contentType, &meta, + ) + if isNoRows(err) { + httpx.WriteError(w, http.StatusNotFound, "not_found", "消息不存在") + return + } + if err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + + var total int + _ = h.db.Read.QueryRowContext(r.Context(), + `SELECT COUNT(*) FROM deliveries WHERE seq = ?`, seq).Scan(&total) + + rows, err := h.db.Read.QueryContext(r.Context(), ` +SELECT endpoint_id, state, reason, attempts, pushed_at, updated_at +FROM deliveries WHERE seq = ? +ORDER BY endpoint_id ASC +LIMIT ? OFFSET ?`, seq, limit, offset) + if err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + defer func() { _ = rows.Close() }() + + deliveries := make([]map[string]any, 0) + for rows.Next() { + var endpointID, dState, dReason string + var attempts int + var pushedAt, updatedAt sql.NullInt64 + if scanErr := rows.Scan(&endpointID, &dState, &dReason, &attempts, &pushedAt, &updatedAt); scanErr != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + deliveries = append(deliveries, map[string]any{ + "endpoint_id": endpointID, + "state": dState, + "reason": dReason, + "attempts": attempts, + "pushed_at_ms": nullInt64API(pushedAt), + "updated_at_ms": updatedAt.Int64, + }) + } + if err := rows.Err(); err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + + next := "" + if offset+len(deliveries) < total { + next = strconv.Itoa(offset + len(deliveries)) + } + + metaObj := any(map[string]any{}) + if strings.TrimSpace(meta) != "" && meta != "{}" { + metaObj = jsonRawOrObject(meta) + } + + httpx.WriteOK(w, map[string]any{ + "seq": seq, + "id": id, + "sender_id": sender, + "dest_kind": destKind, + "dest_id": destID, + "state": state, + "reason": reason, + "send_at_ms": sendAt, + "created_at_ms": created, + "meta": metaObj, + "content_type": contentType, + "deliveries": deliveries, + "next_cursor": next, + }) +} + +func buildMessageFilter(w http.ResponseWriter, q interface{ Get(string) string }) (where string, args []any, ok bool) { + conds := []string{"1=1"} + args = make([]any, 0, 8) + + if v := strings.TrimSpace(q.Get("sender_id")); v != "" { + conds = append(conds, "m.sender_id = ?") + args = append(args, v) + } + if v := strings.TrimSpace(q.Get("group_id")); v != "" { + conds = append(conds, "m.dest_kind = 'group' AND m.dest_id = ?") + args = append(args, v) + } + if v := strings.TrimSpace(q.Get("state")); v != "" { + switch v { + case "scheduled", "dispatched", "completed": + conds = append(conds, "m.state = ?") + args = append(args, v) + default: + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "state 无效") + return "", nil, false + } + } + if v := strings.TrimSpace(q.Get("from_ms")); v != "" { + n, okParse := parseInt64Query(v) + if !okParse { + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "from_ms 无效") + return "", nil, false + } + conds = append(conds, "m.created_at >= ?") + args = append(args, n) + } + if v := strings.TrimSpace(q.Get("to_ms")); v != "" { + n, okParse := parseInt64Query(v) + if !okParse { + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "to_ms 无效") + return "", nil, false + } + conds = append(conds, "m.created_at <= ?") + args = append(args, n) + } + if v := strings.TrimSpace(q.Get("endpoint_id")); v != "" { + conds = append(conds, `EXISTS ( +SELECT 1 FROM deliveries d WHERE d.seq = m.seq AND d.endpoint_id = ?)`) + args = append(args, v) + } + + return strings.Join(conds, " AND "), args, true +} + +func (h *Handler) deliveryCounts(ctx context.Context, seqs []int64) (map[int64]map[string]int, error) { + out := make(map[int64]map[string]int, len(seqs)) + if len(seqs) == 0 { + return out, nil + } + placeholders := make([]string, len(seqs)) + args := make([]any, len(seqs)) + for i, s := range seqs { + placeholders[i] = "?" + args[i] = s + out[s] = emptyDeliveryCounts() + } + rows, err := h.db.Read.QueryContext(ctx, ` +SELECT seq, state, COUNT(*) FROM deliveries +WHERE seq IN (`+strings.Join(placeholders, ",")+`) +GROUP BY seq, state`, args...) + if err != nil { + return nil, err + } + defer func() { _ = rows.Close() }() + for rows.Next() { + var seq int64 + var state string + var n int + if scanErr := rows.Scan(&seq, &state, &n); scanErr != nil { + return nil, scanErr + } + c := out[seq] + switch state { + case "pending", "accepted", "recalled", "expired", "dropped", "rejected": + c[state] = n + } + } + return out, rows.Err() +} + +func emptyDeliveryCounts() map[string]int { + return map[string]int{ + "pending": 0, + "accepted": 0, + "recalled": 0, + "expired": 0, + "dropped": 0, + "rejected": 0, + } +} + +func jsonRawOrObject(s string) any { + // meta 存的是 JSON 文本;解析失败时返回空对象,避免把原始串当正文泄漏路径 + var v any + if err := json.Unmarshal([]byte(s), &v); err != nil || v == nil { + return map[string]any{} + } + return v +} diff --git a/internal/admin/overview.go b/internal/admin/overview.go new file mode 100644 index 0000000..28f500b --- /dev/null +++ b/internal/admin/overview.go @@ -0,0 +1,113 @@ +package admin + +import ( + "database/sql" + "errors" + "net/http" + "strconv" + "time" + + "git.asio.asia/nixevol/NixMsg/internal/httpx" +) + +func (h *Handler) handleOverview(w http.ResponseWriter, r *http.Request) { + ctx := r.Context() + + var total, selfCount, disabled, online int + err := h.db.Read.QueryRowContext(ctx, ` +SELECT + COUNT(*), + COALESCE(SUM(CASE WHEN source = 'self' THEN 1 ELSE 0 END), 0), + COALESCE(SUM(CASE WHEN enabled = 0 THEN 1 ELSE 0 END), 0), + COALESCE(SUM(CASE + WHEN online_since IS NOT NULL + AND (offline_since IS NULL OR online_since > offline_since) THEN 1 + ELSE 0 END), 0) +FROM endpoints`).Scan(&total, &selfCount, &disabled, &online) + if err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + + var groupsTotal int + if err := h.db.Read.QueryRowContext(ctx, `SELECT COUNT(*) FROM groups`).Scan(&groupsTotal); err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + + var pending int + if err := h.db.Read.QueryRowContext(ctx, + `SELECT COUNT(*) FROM deliveries WHERE state = 'pending'`).Scan(&pending); err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + + var scheduled int + if err := h.db.Read.QueryRowContext(ctx, + `SELECT COUNT(*) FROM messages WHERE state = 'scheduled'`).Scan(&scheduled); err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + + httpx.WriteOK(w, map[string]any{ + "version": h.version, + "endpoints_total": total, + "endpoints_self": selfCount, + "endpoints_online": online, + "endpoints_disabled": disabled, + "groups_total": groupsTotal, + "messages_pending": pending, + "messages_scheduled": scheduled, + "uptime_ms": time.Since(h.startedAt).Milliseconds(), + }) +} + +func (h *Handler) handleSettingsGet(w http.ResponseWriter, _ *http.Request) { + cfg := h.cfg + httpx.WriteOK(w, map[string]any{ + "listen": cfg.Listen, + "admin_listen": cfg.AdminListen, + "session_idle_days": cfg.SessionIdleDays, + "record_retention_days": cfg.RecordRetentionDays, + "idempotency_hours": cfg.IdempotencyHours, + "receipt_retention_days": cfg.ReceiptRetentionDays, + "sqlite_synchronous": cfg.SQLiteSynchronous, + "limits": map[string]any{ + "max_body_bytes": cfg.Limits.MaxBodyBytes, + "max_meta_bytes": cfg.Limits.MaxMetaBytes, + "max_frame_bytes": cfg.Limits.MaxFrameBytes, + "max_ttl_seconds": cfg.Limits.MaxTTLSeconds, + "max_schedule_seconds": cfg.Limits.MaxScheduleSeconds, + "max_group_members": cfg.Limits.MaxGroupMembers, + "grace_seconds": cfg.Limits.GraceSeconds, + "ack_timeout_seconds": cfg.Limits.AckTimeoutSeconds, + "delivery_window": cfg.Limits.DeliveryWindow, + "receipt_window": cfg.Limits.ReceiptWindow, + "requests_per_second": cfg.Limits.RequestsPerSecond, + "max_pending_per_sender": cfg.Limits.MaxPendingPerSender, + "max_pending_per_receiver": cfg.Limits.MaxPendingPerReceiver, + }, + }) +} + +func nullInt64API(v sql.NullInt64) any { + if !v.Valid { + return nil + } + return v.Int64 +} + +func parseInt64Query(s string) (int64, bool) { + if s == "" { + return 0, false + } + n, err := strconv.ParseInt(s, 10, 64) + if err != nil { + return 0, false + } + return n, true +} + +func isNoRows(err error) bool { + return errors.Is(err, sql.ErrNoRows) +} diff --git a/internal/admin/registration.go b/internal/admin/registration.go new file mode 100644 index 0000000..583881c --- /dev/null +++ b/internal/admin/registration.go @@ -0,0 +1,173 @@ +package admin + +import ( + "context" + "crypto/rand" + "database/sql" + "errors" + "net/http" + "strings" + "time" + "unicode/utf8" + + "git.asio.asia/nixevol/NixMsg/internal/httpx" +) + +const ( + settingRegistrationEnabled = "registration_enabled" + settingRegistrationCode = "registration_code" + minRegistrationCodeLen = 8 + maxRegistrationCodeLen = 64 + generatedRegistrationLen = 16 +) + +func (h *Handler) handleRegistrationGet(w http.ResponseWriter, r *http.Request) { + data, err := h.loadRegistration(r.Context()) + if err != nil { + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + httpx.WriteOK(w, data) +} + +func (h *Handler) handleRegistrationPut(w http.ResponseWriter, r *http.Request) { + p, _ := principalFrom(r.Context()) + ip := httpx.ClientIP(r, h.trusted) + + var req struct { + Enabled *bool `json:"enabled"` + Code *string `json:"code"` + Generate bool `json:"generate"` + } + if err := httpx.DecodeJSON(r, &req); err != nil { + h.audit(actorString(p), "registration_update", "", "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效") + return + } + if req.Enabled == nil && req.Code == nil && !req.Generate { + h.audit(actorString(p), "registration_update", "", "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", "至少提供一项更新") + return + } + + nowMs := time.Now().UnixMilli() + err := h.db.Queue.Do(r.Context(), func(tx *sql.Tx) error { + if req.Enabled != nil { + val := "0" + if *req.Enabled { + val = "1" + } + if _, e := tx.Exec( + `INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?) + ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at`, + settingRegistrationEnabled, val, nowMs, + ); e != nil { + return e + } + } + if req.Generate { + code, genErr := generateRegistrationCode() + if genErr != nil { + return genErr + } + if _, e := tx.Exec( + `INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?) + ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at`, + settingRegistrationCode, code, nowMs, + ); e != nil { + return e + } + return nil + } + if req.Code != nil { + code := *req.Code + n := utf8.RuneCountInString(code) + if n < minRegistrationCodeLen || n > maxRegistrationCodeLen { + return errBadRequest("安全码须为 8–64 字符") + } + if _, e := tx.Exec( + `INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?) + ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at`, + settingRegistrationCode, code, nowMs, + ); e != nil { + return e + } + } + return nil + }) + if err != nil { + var br badRequestError + if errors.As(err, &br) { + h.audit(actorString(p), "registration_update", "", "bad_request", ip) + httpx.WriteError(w, http.StatusBadRequest, "bad_request", string(br)) + return + } + h.audit(actorString(p), "registration_update", "", "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + + data, err := h.loadRegistration(r.Context()) + if err != nil { + h.audit(actorString(p), "registration_update", "", "error", ip) + httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误") + return + } + // 审计不写安全码明文 + h.audit(actorString(p), "registration_update", "", "ok", ip) + httpx.WriteOK(w, data) +} + +func (h *Handler) loadRegistration(ctx context.Context) (map[string]any, error) { + var enabledVal, codeVal sql.NullString + var enabledAt, codeAt sql.NullInt64 + + err := h.db.Read.QueryRowContext(ctx, + `SELECT value, updated_at FROM settings WHERE key = ?`, settingRegistrationEnabled, + ).Scan(&enabledVal, &enabledAt) + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return nil, err + } + err = h.db.Read.QueryRowContext(ctx, + `SELECT value, updated_at FROM settings WHERE key = ?`, settingRegistrationCode, + ).Scan(&codeVal, &codeAt) + if err != nil && !errors.Is(err, sql.ErrNoRows) { + return nil, err + } + + updated := int64(0) + if enabledAt.Valid && enabledAt.Int64 > updated { + updated = enabledAt.Int64 + } + if codeAt.Valid && codeAt.Int64 > updated { + updated = codeAt.Int64 + } + + return map[string]any{ + "enabled": registrationTruthy(enabledVal.String), + "code": codeVal.String, + "updated_at_ms": updated, + }, nil +} + +func registrationTruthy(v string) bool { + switch strings.TrimSpace(strings.ToLower(v)) { + case "1", "true", "yes", "on": + return true + default: + return false + } +} + +func generateRegistrationCode() (string, error) { + const alphabet = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789" + b := make([]byte, generatedRegistrationLen) + if _, err := rand.Read(b); err != nil { + return "", err + } + out := make([]byte, generatedRegistrationLen) + for i := range b { + out[i] = alphabet[int(b[i])%len(alphabet)] + } + return string(out), nil +} diff --git a/internal/httpx/metrics.go b/internal/httpx/metrics.go new file mode 100644 index 0000000..f22e3e3 --- /dev/null +++ b/internal/httpx/metrics.go @@ -0,0 +1,29 @@ +package httpx + +import ( + "net/http" + "strings" +) + +// MetricsGate 实现 DEVELOPMENT 4.3 节 /metrics 访问规则(共用端口)。 +// token 为空返回 404;Authorization Bearer 不匹配返回 401;匹配则交给 next。 +// 后台单独监听时不要包这层,直接挂 metrics Handler。 +func MetricsGate(token string, next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + if token == "" { + http.NotFound(w, r) + return + } + auth := r.Header.Get("Authorization") + const prefix = "Bearer " + if !strings.HasPrefix(auth, prefix) || auth[len(prefix):] != token { + w.WriteHeader(http.StatusUnauthorized) + return + } + if next == nil { + http.NotFound(w, r) + return + } + next.ServeHTTP(w, r) + }) +} diff --git a/internal/httpx/metrics_test.go b/internal/httpx/metrics_test.go new file mode 100644 index 0000000..b6115e5 --- /dev/null +++ b/internal/httpx/metrics_test.go @@ -0,0 +1,98 @@ +package httpx_test + +import ( + "io" + "net/http" + "net/http/httptest" + "testing" + + "git.asio.asia/nixevol/NixMsg/internal/httpx" + "git.asio.asia/nixevol/NixMsg/internal/listener" + "git.asio.asia/nixevol/NixMsg/internal/metrics" +) + +func TestMetricsGateSharedPort(t *testing.T) { + reg := metrics.New() + okHandler := reg.Handler() + + t.Run("empty token 404", func(t *testing.T) { + mux := listener.NewMux(listener.RoleShared, listener.Handlers{ + Metrics: okHandler, + MetricsToken: "", + }) + res := httptest.NewRecorder() + mux.ServeHTTP(res, httptest.NewRequest(http.MethodGet, "/metrics", nil)) + if res.Code != http.StatusNotFound { + t.Fatalf("want 404 got %d", res.Code) + } + }) + + t.Run("wrong token 401", func(t *testing.T) { + mux := listener.NewMux(listener.RoleShared, listener.Handlers{ + Metrics: okHandler, + MetricsToken: "secret-token", + }) + req := httptest.NewRequest(http.MethodGet, "/metrics", nil) + req.Header.Set("Authorization", "Bearer wrong") + res := httptest.NewRecorder() + mux.ServeHTTP(res, req) + if res.Code != http.StatusUnauthorized { + t.Fatalf("want 401 got %d", res.Code) + } + }) + + t.Run("no auth 401", func(t *testing.T) { + mux := listener.NewMux(listener.RoleShared, listener.Handlers{ + Metrics: okHandler, + MetricsToken: "secret-token", + }) + res := httptest.NewRecorder() + mux.ServeHTTP(res, httptest.NewRequest(http.MethodGet, "/metrics", nil)) + if res.Code != http.StatusUnauthorized { + t.Fatalf("want 401 got %d", res.Code) + } + }) + + t.Run("good token 200", func(t *testing.T) { + mux := listener.NewMux(listener.RoleShared, listener.Handlers{ + Metrics: okHandler, + MetricsToken: "secret-token", + }) + req := httptest.NewRequest(http.MethodGet, "/metrics", nil) + req.Header.Set("Authorization", "Bearer secret-token") + res := httptest.NewRecorder() + mux.ServeHTTP(res, req) + if res.Code != http.StatusOK { + t.Fatalf("want 200 got %d", res.Code) + } + body, _ := io.ReadAll(res.Body) + if len(body) == 0 { + t.Fatal("empty metrics body") + } + }) + + t.Run("admin role no auth", func(t *testing.T) { + mux := listener.NewMux(listener.RoleAdmin, listener.Handlers{ + Metrics: okHandler, + }) + res := httptest.NewRecorder() + mux.ServeHTTP(res, httptest.NewRequest(http.MethodGet, "/metrics", nil)) + if res.Code != http.StatusOK { + t.Fatalf("admin metrics want 200 got %d", res.Code) + } + }) +} + +func TestMetricsGateDirect(t *testing.T) { + h := httpx.MetricsGate("tok", http.HandlerFunc(func(w http.ResponseWriter, _ *http.Request) { + w.WriteHeader(http.StatusOK) + _, _ = w.Write([]byte("ok")) + })) + res := httptest.NewRecorder() + req := httptest.NewRequest(http.MethodGet, "/metrics", nil) + req.Header.Set("Authorization", "Bearer tok") + h.ServeHTTP(res, req) + if res.Code != 200 || res.Body.String() != "ok" { + t.Fatalf("got %d %q", res.Code, res.Body.String()) + } +} diff --git a/internal/listener/routes.go b/internal/listener/routes.go index cd0f7bf..6f3c1d1 100644 --- a/internal/listener/routes.go +++ b/internal/listener/routes.go @@ -2,7 +2,8 @@ package listener import ( "net/http" - "strings" + + "git.asio.asia/nixevol/NixMsg/internal/httpx" ) // RouteRole 区分监听用途,决定哪些路径可用。 @@ -92,7 +93,7 @@ func NewMux(role RouteRole, h Handlers) http.Handler { if h.AdminAPI != nil { mux.Handle("/api/admin/", h.AdminAPI) } - mux.Handle("GET /metrics", metricsGate(h.MetricsToken, h.Metrics)) + mux.Handle("GET /metrics", httpx.MetricsGate(h.MetricsToken, h.Metrics)) if h.Static != nil { mux.Handle("/", h.Static) } else { @@ -102,23 +103,3 @@ func NewMux(role RouteRole, h Handlers) http.Handler { return mux } - -func metricsGate(token string, next http.Handler) http.Handler { - return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { - if token == "" { - http.NotFound(w, r) - return - } - auth := r.Header.Get("Authorization") - const prefix = "Bearer " - if !strings.HasPrefix(auth, prefix) || auth[len(prefix):] != token { - w.WriteHeader(http.StatusUnauthorized) - return - } - if next == nil { - http.NotFound(w, r) - return - } - next.ServeHTTP(w, r) - }) -}