Compare commits
3
Commits
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
0e068c6ade | ||
|
|
308b0b9edd | ||
|
|
8bcb5c63be |
+102
-3
@@ -278,15 +278,114 @@
|
||||
|
||||
## 消息 M
|
||||
|
||||
暂无。
|
||||
### M1 2026-09-30
|
||||
|
||||
1. **提交时分发做成最小正确版**
|
||||
- 原条款:DEVELOPMENT 7.3 步骤 8 / 7.4:`send_at` 已到则同一写操作内完整分发(停用拒绝、`queue_full`、`expire_at`/宽限、无接收者 `completed`、回执等)。
|
||||
- 实际做法:单聊只插一条 `pending`;群按当时 `group_members` 去掉发送者各插 `pending`;消息改为 `dispatched`。不设 `expire_at`,不检查接收端配额/在线/停用,不因无接收者改为 `completed`,不写回执,不唤醒推送循环。
|
||||
- 原因:M1 范围是提交;完整分发与推送属 M2。
|
||||
- 备选方案:M1 直接实现完整 7.4(抢 M2)。
|
||||
- 影响:到点消息已有投递行,但停用成员仍会有 `pending`;无成员群仍为 `dispatched` 且无投递;推送需等 M2。
|
||||
|
||||
2. **请求频率突发容量写死为 100**
|
||||
- 原条款:DEVELOPMENT 6.10 每端每秒 50、突发 100;配置示例仅有 `requests_per_second`。
|
||||
- 实际做法:`Limits.RequestBurst` 默认 100;`requests_per_second<=0` 时不限速(便于测试)。速率桶挂在 `message.App` 的 `Submit` 入口;`ack`/`receipt_ack` 尚未实现故未接桶。
|
||||
- 原因:配置无独立 burst 字段。
|
||||
- 备选方案:配置增加 `request_burst`;由连接线在上行统一限流。
|
||||
- 影响:改 `requests_per_second` 不改突发;正式接线后若 N 线也限流可能双重计数。
|
||||
|
||||
3. **未接线 `cmd/nixmsg`**
|
||||
- 原条款:可替换 T0.4 假实现。
|
||||
- 实际做法:新增 `message.App` 实现 `Submit`;保留 `Stub`;按任务隔离要求未改 `cmd/nixmsg`/`wire.go`。
|
||||
- 原因:本任务禁止改 `cmd/nixmsg`;总控接线或后续任务再换。
|
||||
- 备选方案:本任务直接改 `wire.go`。
|
||||
- 影响:进程内仍用 Stub,需显式构造 `message.New` 才能用真实提交。
|
||||
|
||||
4. **防重键在、消息行已删时返回 `not_found`**
|
||||
- 原条款:防重命中返回原消息当前状态;未写明消息行已被清理时的提交重试行为(状态查询为 `not_found`)。
|
||||
- 实际做法:`send_keys` 指纹相同但 `messages` 无行时返回 `not_found`。
|
||||
- 原因:无法构造 `send_at`/`state`。
|
||||
- 备选方案:在 `send_keys` 冗余存结果快照。
|
||||
- 影响:保留期过后的重试不再幂等成功。
|
||||
|
||||
## 身份 I
|
||||
|
||||
暂无。
|
||||
### I1 2026-09-30
|
||||
|
||||
1. **注册做成可挂载 Handler,不改 cmd/listener**
|
||||
- 原条款:TASKS I1 / DEVELOPMENT 6.9 在 `listen` 上提供 `POST /api/client/register`;依赖 N1 端口识别。
|
||||
- 实际做法:`identity.NewRegisterHandler` / `identity.NewServer().Handler()` 返回 `http.Handler`,由接线方 `mux.Handle("/api/client/register", h)`;本线不改 `cmd/nixmsg`、`internal/listener`(N 线未合入)。
|
||||
- 原因:隔离交付,避免抢 N/P 接线。
|
||||
- 备选方案:本线直接改 `wire.go` 挂路由。
|
||||
- 影响:合入后需总控或 N/A 接线才对外可访问。
|
||||
|
||||
2. **密码哈希与锁定走 auth 接口,本分支用可替换假实现测**
|
||||
- 原条款:依赖 P3 argon2 池与锁定计数器。
|
||||
- 实际做法:`RegisterConfig.Hash`/`Locks` 注入 `auth.HashPool`、`auth.LoginLocks`;测试用 `auth.NewStubHashPool` + 仅实现 `LockRegisterIP`(5 分钟 10 次)的测试锁定器,不在本线重写 argon2。
|
||||
- 原因:P3 尚未在本分支。
|
||||
- 备选方案:等 P3 合入后再写 I1。
|
||||
- 影响:生产须注入 P3 实现;StubLoginLocks 永不锁定,不能直接用于开放注册。
|
||||
|
||||
3. **settings 开关取值**
|
||||
- 原条款:`settings.registration_enabled`,未规定字符串字面量。
|
||||
- 实际做法:`1`/`true`/`yes`/`on`(大小写不敏感)视为开启,其余(含缺省)关闭;安全码键 `registration_code`。
|
||||
- 原因:与 store 测试写入的 `"0"`/`"1"` 对齐并兼容常见布尔字面量。
|
||||
- 备选方案:仅认 `"1"`。
|
||||
- 影响:A 线写注册设置时宜写 `"1"`/`"0"`。
|
||||
|
||||
4. **客户端 IP**
|
||||
- 原条款:DEVELOPMENT 4.5 受信任代理下用 `X-Forwarded-For`。
|
||||
- 实际做法:Handler 默认取 `RemoteAddr` 的 host;可通过 `RegisterConfig.ClientIP` 注入。本线不做 `trusted_proxies` 解析(属 listener/接线)。
|
||||
- 原因:不改 listener;代理 IP 应由外层在挂载前算好或注入。
|
||||
- 备选方案:在 identity 内复制 4.5 逻辑。
|
||||
- 影响:经代理部署时接线方必须注入真实 IP,否则锁定按直连 IP 计。
|
||||
|
||||
5. **生成登录密码长度**
|
||||
- 原条款:F01 留空则生成,8–128 字符,不以 `nst_` 开头;未规定生成长度。
|
||||
- 实际做法:生成 20 位字母数字;若偶然以 `nst_` 开头则重抽。
|
||||
- 原因:与管理员 init 量级接近,满足规则。
|
||||
- 备选方案:16/32 位。
|
||||
- 影响:无产品行为差异。
|
||||
|
||||
## 后台接口 A
|
||||
|
||||
暂无。
|
||||
### A1 2026-09-30
|
||||
|
||||
1. **A1 仅交付可挂载 Handler,未改 `cmd/nixmsg`**
|
||||
- 原条款:管理接口由服务进程提供;TASKS 要求路由可挂载。
|
||||
- 实际做法:`admin.New(Deps) *Handler` 实现 `http.Handler`,路径为完整 `/api/admin/...`;总控/后续接线在 mux 上 `Handle("/api/admin/", h)` 即可。本任务按隔离要求不改 `cmd/`、`listener`、`protocol`。
|
||||
- 原因:与并行线隔离;A1 验收用 httptest。
|
||||
- 备选方案:本分支同时改 `serve.go` 挂载(易与 N/P 冲突)。
|
||||
- 影响:合入后需在 `serve`/`wire` 挂载并注入真实 DB/Hash/Locks。
|
||||
|
||||
2. **P3 未合入时的假哈希与本地锁定/令牌生成**
|
||||
- 原条款:密码与令牌哈希用 `internal/auth`;锁定与 argon2 池由 P3 实现。
|
||||
- 实际做法:依赖 `auth.HashPool` / `auth.APITokens` / `auth.LoginLocks` 接口;测试注入 `auth.StubHashPool`。提供 `admin.MemoryLoginLocks`(5 分钟 10 次锁 5 分钟)与 `admin.RandomAPITokens`(`nxm_`+32 字节 base64url,SHA-256)供本线与测试使用,不实现 argon2。
|
||||
- 原因:第 1 波允许用假实现;P3 合入后替换注入即可。
|
||||
- 备选方案:阻塞等待 P3。
|
||||
- 影响:生产接线应改用 P3 实现;`MemoryLoginLocks`/`RandomAPITokens` 可保留作测试替身。
|
||||
|
||||
3. **API 令牌 `id` 为字符串**
|
||||
- 原条款:`docs/api/admin-api.md` 示例 `"id": 1`(数字)。
|
||||
- 实际做法:遵循库表 `api_tokens.id TEXT`,响应 `id` 为 16 字节随机十六进制字符串。
|
||||
- 原因:不改已发布迁移;与 schema 一致。
|
||||
- 备选方案:另加 INTEGER 列或把数字存成文本并在 JSON 里发数字。
|
||||
- 影响:W 线类型应按 `string` 解析令牌 id。
|
||||
|
||||
4. **尚未实现的管理路由(鉴权中间件已生效,业务返回 501)**
|
||||
- 原条款:DEVELOPMENT 第 8 节完整路由表。
|
||||
- 实际做法:已实现 `login`/`logout`/`me`/`password` 与 `/api/admin/tokens` 全套。以下路由经鉴权后返回 `501 not_implemented`(属 A2/A3):
|
||||
`GET /overview`;`endpoints` 列表/开通/import/batch/详情/改/删/kick/reset-login-password/talk-password/unlock;`registration` GET/PUT;`groups` 全部(含 `GET /groups/{id}`);`messages` 列表与详情;`GET /settings`。
|
||||
- 原因:A1 范围仅鉴权、令牌、操作日志。
|
||||
- 备选方案:无。
|
||||
- 影响:W/集成测试在 A2/A3 前勿依赖这些业务响应。
|
||||
|
||||
5. **操作日志用 `slog` 结构化字段**
|
||||
- 原条款:写结构化日志(操作者、动作、对象、结果、来源 IP)。
|
||||
- 实际做法:`Logger.Info("admin_audit", "actor", ..., "action", ..., "object", ..., "result", ..., "ip", ...)`;改状态请求写日志;不写密码/令牌/正文。
|
||||
- 原因:文档未规定日志后端。
|
||||
- 备选方案:独立 audit 表。
|
||||
- 影响:日志采集需按 msg=`admin_audit` 过滤。
|
||||
|
||||
## 后台网页 W
|
||||
|
||||
|
||||
@@ -0,0 +1,355 @@
|
||||
package admin_test
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"net/http"
|
||||
"net/http/cookiejar"
|
||||
"net/http/httptest"
|
||||
"net/url"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/admin"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
const testPassword = "admin-password-ok"
|
||||
|
||||
type envelope struct {
|
||||
OK bool `json:"ok"`
|
||||
Data json.RawMessage `json:"data"`
|
||||
Error *struct {
|
||||
Code string `json:"code"`
|
||||
Message string `json:"message"`
|
||||
} `json:"error"`
|
||||
}
|
||||
|
||||
func setup(t *testing.T) (*admin.Handler, *httptest.Server, *http.Client, auth.HashPool) {
|
||||
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)
|
||||
}
|
||||
locks := admin.NewMemoryLoginLocks()
|
||||
h := admin.New(admin.Deps{
|
||||
DB: db,
|
||||
Hash: hash,
|
||||
Tokens: admin.NewRandomAPITokens(),
|
||||
Locks: locks,
|
||||
})
|
||||
srv := httptest.NewServer(h)
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
jar, err := cookiejar.New(nil)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
client := &http.Client{Jar: jar}
|
||||
return h, srv, client, hash
|
||||
}
|
||||
|
||||
func decodeEnv(t *testing.T, res *http.Response) envelope {
|
||||
t.Helper()
|
||||
defer func() { _ = res.Body.Close() }()
|
||||
var env envelope
|
||||
if err := json.NewDecoder(res.Body).Decode(&env); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return env
|
||||
}
|
||||
|
||||
func postJSON(t *testing.T, client *http.Client, url, body string, headers map[string]string) *http.Response {
|
||||
t.Helper()
|
||||
req, err := http.NewRequest(http.MethodPost, url, strings.NewReader(body))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
for k, v := range headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
res, err := client.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return res
|
||||
}
|
||||
|
||||
func doReq(t *testing.T, client *http.Client, method, rawURL, body string, headers map[string]string) *http.Response {
|
||||
t.Helper()
|
||||
req, err := http.NewRequest(method, rawURL, strings.NewReader(body))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if body != "" {
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
}
|
||||
for k, v := range headers {
|
||||
req.Header.Set(k, v)
|
||||
}
|
||||
res, err := client.Do(req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return res
|
||||
}
|
||||
|
||||
func login(t *testing.T, client *http.Client, base string) {
|
||||
t.Helper()
|
||||
res := postJSON(t, client, base+"/api/admin/login",
|
||||
`{"username":"admin","password":"`+testPassword+`"}`, nil)
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("login: status=%d env=%+v", res.StatusCode, env)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginMePasswordLogout(t *testing.T) {
|
||||
_, srv, client, _ := setup(t)
|
||||
base := srv.URL
|
||||
|
||||
login(t, client, base)
|
||||
|
||||
res := doReq(t, client, http.MethodGet, base+"/api/admin/me", "", nil)
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("me: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
var me map[string]any
|
||||
_ = json.Unmarshal(env.Data, &me)
|
||||
if me["username"] != "admin" || me["auth"] != "cookie" {
|
||||
t.Fatalf("me data=%v", me)
|
||||
}
|
||||
|
||||
res = postJSON(t, client, base+"/api/admin/password",
|
||||
`{"old_password":"`+testPassword+`","new_password":"new-password-12"}`,
|
||||
map[string]string{"X-Nixmsg-Request": "1"})
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("password: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
|
||||
res = postJSON(t, client, base+"/api/admin/logout", `{}`,
|
||||
map[string]string{"X-Nixmsg-Request": "1"})
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("logout: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
|
||||
res = doReq(t, client, http.MethodGet, base+"/api/admin/me", "", nil)
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 401 {
|
||||
t.Fatalf("after logout want 401 got %d", res.StatusCode)
|
||||
}
|
||||
|
||||
// 用新密码再登录
|
||||
res = postJSON(t, client, base+"/api/admin/login",
|
||||
`{"username":"admin","password":"new-password-12"}`, nil)
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("relogin: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCSRFRequiredForCookieMutating(t *testing.T) {
|
||||
_, srv, client, _ := setup(t)
|
||||
base := srv.URL
|
||||
login(t, client, base)
|
||||
|
||||
res := postJSON(t, client, base+"/api/admin/password",
|
||||
`{"old_password":"`+testPassword+`","new_password":"new-password-12"}`,
|
||||
nil) // 无 CSRF 头
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 403 || env.Error == nil || env.Error.Code != "forbidden" {
|
||||
t.Fatalf("want 403 forbidden, got %d %+v", res.StatusCode, env)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAPITokenAuthAndRestrictions(t *testing.T) {
|
||||
_, srv, client, _ := setup(t)
|
||||
base := srv.URL
|
||||
login(t, client, base)
|
||||
|
||||
res := postJSON(t, client, base+"/api/admin/tokens",
|
||||
`{"name":"ops"}`,
|
||||
map[string]string{"X-Nixmsg-Request": "1"})
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("create token: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
var created struct {
|
||||
ID string `json:"id"`
|
||||
Token string `json:"token"`
|
||||
Name string `json:"name"`
|
||||
}
|
||||
if err := json.Unmarshal(env.Data, &created); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if !strings.HasPrefix(created.Token, "nxm_") {
|
||||
t.Fatalf("token prefix: %q", created.Token)
|
||||
}
|
||||
|
||||
tokClient := &http.Client{}
|
||||
hdr := map[string]string{"Authorization": "Bearer " + created.Token}
|
||||
|
||||
res = doReq(t, tokClient, http.MethodGet, base+"/api/admin/me", "", hdr)
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("token me: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
var me map[string]any
|
||||
_ = json.Unmarshal(env.Data, &me)
|
||||
if me["auth"] != "token" {
|
||||
t.Fatalf("auth=%v", me["auth"])
|
||||
}
|
||||
|
||||
// 普通管理接口鉴权通过(业务 501)
|
||||
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)
|
||||
}
|
||||
|
||||
// 禁止 password / tokens
|
||||
res = postJSON(t, tokClient, base+"/api/admin/password",
|
||||
`{"old_password":"x","new_password":"new-password-12"}`, hdr)
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 403 {
|
||||
t.Fatalf("token password want 403 got %d", res.StatusCode)
|
||||
}
|
||||
|
||||
res = doReq(t, tokClient, http.MethodGet, base+"/api/admin/tokens", "", hdr)
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 403 {
|
||||
t.Fatalf("token list want 403 got %d", res.StatusCode)
|
||||
}
|
||||
|
||||
// 停用后立即失效
|
||||
res = doReq(t, client, http.MethodPatch, base+"/api/admin/tokens/"+created.ID,
|
||||
`{"enabled":false}`,
|
||||
map[string]string{"X-Nixmsg-Request": "1", "Content-Type": "application/json"})
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 200 || !env.OK {
|
||||
t.Fatalf("disable: %d %+v", res.StatusCode, env)
|
||||
}
|
||||
|
||||
res = doReq(t, tokClient, http.MethodGet, base+"/api/admin/me", "", hdr)
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 401 {
|
||||
t.Fatalf("disabled token want 401 got %d %+v", res.StatusCode, env)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLoginLock(t *testing.T) {
|
||||
_, srv, _, _ := setup(t)
|
||||
base := srv.URL
|
||||
|
||||
for i := 0; i < 9; i++ {
|
||||
client := &http.Client{}
|
||||
res := postJSON(t, client, base+"/api/admin/login",
|
||||
`{"username":"admin","password":"wrong-password!!"}`, nil)
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 401 {
|
||||
t.Fatalf("fail %d: want 401 got %d %+v", i, res.StatusCode, env)
|
||||
}
|
||||
}
|
||||
client := &http.Client{}
|
||||
res := postJSON(t, client, base+"/api/admin/login",
|
||||
`{"username":"admin","password":"wrong-password!!"}`, nil)
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 429 || env.Error == nil || env.Error.Code != "rate_limited" {
|
||||
t.Fatalf("want 429 rate_limited got %d %+v", res.StatusCode, env)
|
||||
}
|
||||
|
||||
res = postJSON(t, client, base+"/api/admin/login",
|
||||
`{"username":"admin","password":"`+testPassword+`"}`, nil)
|
||||
env = decodeEnv(t, res)
|
||||
if res.StatusCode != 429 {
|
||||
t.Fatalf("locked correct login want 429 got %d", res.StatusCode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestBadAPITokenCountsTowardLock(t *testing.T) {
|
||||
_, srv, _, _ := setup(t)
|
||||
base := srv.URL
|
||||
tokClient := &http.Client{}
|
||||
hdr := map[string]string{"Authorization": "Bearer nxm_" + strings.Repeat("a", 43)}
|
||||
|
||||
for i := 0; i < 9; i++ {
|
||||
res := doReq(t, tokClient, http.MethodGet, base+"/api/admin/me", "", hdr)
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 401 {
|
||||
t.Fatalf("bad token %d: want 401 got %d %+v", i, res.StatusCode, env)
|
||||
}
|
||||
}
|
||||
res := doReq(t, tokClient, http.MethodGet, base+"/api/admin/me", "", hdr)
|
||||
env := decodeEnv(t, res)
|
||||
if res.StatusCode != 429 {
|
||||
t.Fatalf("want lock 429 got %d %+v", res.StatusCode, env)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCookieSetAttributes(t *testing.T) {
|
||||
_, srv, client, _ := setup(t)
|
||||
base := srv.URL
|
||||
res := postJSON(t, client, base+"/api/admin/login",
|
||||
`{"username":"admin","password":"`+testPassword+`"}`, nil)
|
||||
_ = decodeEnv(t, res)
|
||||
|
||||
u, _ := url.Parse(base)
|
||||
cookies := client.Jar.Cookies(u)
|
||||
found := false
|
||||
for _, c := range cookies {
|
||||
if c.Name == "nixmsg_admin" {
|
||||
found = true
|
||||
if c.Value == "" {
|
||||
t.Fatal("empty cookie")
|
||||
}
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Fatal("cookie not set")
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountableHandler(t *testing.T) {
|
||||
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()
|
||||
_ = admin.SeedAdminPassword(context.Background(), db, hash, testPassword)
|
||||
h := admin.New(admin.Deps{
|
||||
DB: db,
|
||||
Hash: hash,
|
||||
Tokens: admin.NewRandomAPITokens(),
|
||||
Locks: admin.NewMemoryLoginLocks(),
|
||||
})
|
||||
mux := http.NewServeMux()
|
||||
mux.Handle("/api/admin/", h)
|
||||
srv := httptest.NewServer(mux)
|
||||
t.Cleanup(srv.Close)
|
||||
|
||||
body := `{"username":"admin","password":"` + testPassword + `"}`
|
||||
res, err := http.Post(srv.URL+"/api/admin/login", "application/json", bytes.NewBufferString(body))
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
env := decodeEnv(t, res)
|
||||
if !env.OK {
|
||||
t.Fatalf("mount login failed: %+v", env)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package admin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"strings"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
)
|
||||
|
||||
// RandomAPITokens 生成真实 nxm_ 令牌(P3 合入前供 A1 使用)。
|
||||
// 哈希为 SHA-256 原始字节,与 DEVELOPMENT 一致。
|
||||
type RandomAPITokens struct{}
|
||||
|
||||
// NewRandomAPITokens 返回可注入的 APITokens 实现。
|
||||
func NewRandomAPITokens() *RandomAPITokens { return &RandomAPITokens{} }
|
||||
|
||||
func (RandomAPITokens) Issue(_ context.Context) (string, []byte, error) {
|
||||
buf := make([]byte, 32)
|
||||
if _, err := rand.Read(buf); err != nil {
|
||||
return "", nil, err
|
||||
}
|
||||
tok := "nxm_" + base64.RawURLEncoding.EncodeToString(buf)
|
||||
sum := sha256.Sum256([]byte(tok))
|
||||
return tok, sum[:], nil
|
||||
}
|
||||
|
||||
func (RandomAPITokens) HashToken(token string) []byte {
|
||||
sum := sha256.Sum256([]byte(token))
|
||||
return sum[:]
|
||||
}
|
||||
|
||||
func (RandomAPITokens) LooksLikeAPIToken(credential string) bool {
|
||||
return strings.HasPrefix(credential, "nxm_")
|
||||
}
|
||||
|
||||
var _ auth.APITokens = (*RandomAPITokens)(nil)
|
||||
@@ -0,0 +1,12 @@
|
||||
package admin
|
||||
|
||||
// audit 写结构化操作日志;不写密码、令牌和正文。
|
||||
func (h *Handler) audit(actor, action, object, result, ip string) {
|
||||
h.log.Info("admin_audit",
|
||||
"actor", actor,
|
||||
"action", action,
|
||||
"object", object,
|
||||
"result", result,
|
||||
"ip", ip,
|
||||
)
|
||||
}
|
||||
@@ -0,0 +1,180 @@
|
||||
package admin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"database/sql"
|
||||
"encoding/base64"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/httpx"
|
||||
)
|
||||
|
||||
type authKind string
|
||||
|
||||
const (
|
||||
authCookie authKind = "cookie"
|
||||
authToken authKind = "token"
|
||||
)
|
||||
|
||||
type principal struct {
|
||||
Kind authKind
|
||||
TokenName string
|
||||
TokenID string
|
||||
Session string
|
||||
}
|
||||
|
||||
type ctxKey int
|
||||
|
||||
const principalKey ctxKey = 1
|
||||
|
||||
func principalFrom(ctx context.Context) (principal, bool) {
|
||||
p, ok := ctx.Value(principalKey).(principal)
|
||||
return p, ok
|
||||
}
|
||||
|
||||
func (h *Handler) auth(next http.HandlerFunc) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
ip := httpx.ClientIP(r, h.trusted)
|
||||
|
||||
if locked, retry := h.locks.Check(auth.LockKey{Kind: auth.LockAdminIP, IP: ip}); locked {
|
||||
w.Header().Set("Retry-After", formatRetryAfter(retry))
|
||||
httpx.WriteError(w, http.StatusTooManyRequests, "rate_limited", "登录已锁定,请稍后再试")
|
||||
return
|
||||
}
|
||||
|
||||
p, errCode, errMsg, status := h.authenticate(r, ip)
|
||||
if status != 0 {
|
||||
if status == http.StatusTooManyRequests {
|
||||
w.Header().Set("Retry-After", "300")
|
||||
}
|
||||
httpx.WriteError(w, status, errCode, errMsg)
|
||||
return
|
||||
}
|
||||
|
||||
if p.Kind == authCookie && isMutating(r.Method) {
|
||||
if r.Header.Get(csrfHeader) != csrfValue {
|
||||
httpx.WriteError(w, http.StatusForbidden, "forbidden", "缺少 X-Nixmsg-Request 头")
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
if p.Kind == authToken && isTokenForbiddenPath(r.URL.Path) {
|
||||
httpx.WriteError(w, http.StatusForbidden, "forbidden", "API 令牌无权访问该接口")
|
||||
return
|
||||
}
|
||||
|
||||
ctx := context.WithValue(r.Context(), principalKey, p)
|
||||
next(w, r.WithContext(ctx))
|
||||
})
|
||||
}
|
||||
|
||||
func isMutating(method string) bool {
|
||||
switch method {
|
||||
case http.MethodPost, http.MethodPut, http.MethodPatch, http.MethodDelete:
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func isTokenForbiddenPath(path string) bool {
|
||||
if path == "/api/admin/password" {
|
||||
return true
|
||||
}
|
||||
return path == "/api/admin/tokens" || strings.HasPrefix(path, "/api/admin/tokens/")
|
||||
}
|
||||
|
||||
func (h *Handler) authenticate(r *http.Request, ip string) (principal, string, string, int) {
|
||||
authz := r.Header.Get("Authorization")
|
||||
if strings.HasPrefix(strings.ToLower(authz), "bearer ") {
|
||||
raw := strings.TrimSpace(authz[len("Bearer "):])
|
||||
if raw == "" || !h.tokens.LooksLikeAPIToken(raw) {
|
||||
return authFail(h, ip)
|
||||
}
|
||||
hash := h.tokens.HashToken(raw)
|
||||
info, err := h.lookupAPIToken(r.Context(), hash)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return authFail(h, ip)
|
||||
}
|
||||
return principal{}, "internal", "内部错误", http.StatusInternalServerError
|
||||
}
|
||||
if !info.Enabled {
|
||||
return authFail(h, ip)
|
||||
}
|
||||
h.touchLastUsed(r.Context(), info.ID)
|
||||
return principal{Kind: authToken, TokenName: info.Name, TokenID: info.ID}, "", "", 0
|
||||
}
|
||||
|
||||
c, err := r.Cookie(cookieName)
|
||||
if err != nil || c.Value == "" {
|
||||
return principal{}, "unauthorized", "未登录", http.StatusUnauthorized
|
||||
}
|
||||
hashHex := hashSessionHex(c.Value)
|
||||
ok, err := h.sessionValid(r.Context(), hashHex)
|
||||
if err != nil {
|
||||
return principal{}, "internal", "内部错误", http.StatusInternalServerError
|
||||
}
|
||||
if !ok {
|
||||
return principal{}, "unauthorized", "未登录", http.StatusUnauthorized
|
||||
}
|
||||
return principal{Kind: authCookie, Session: c.Value}, "", "", 0
|
||||
}
|
||||
|
||||
func authFail(h *Handler, ip string) (principal, string, string, int) {
|
||||
if locked, _ := h.locks.Fail(auth.LockKey{Kind: auth.LockAdminIP, IP: ip}); locked {
|
||||
return principal{}, "rate_limited", "登录已锁定,请稍后再试", http.StatusTooManyRequests
|
||||
}
|
||||
return principal{}, "unauthorized", "令牌无效", http.StatusUnauthorized
|
||||
}
|
||||
|
||||
func hashSessionHex(token string) string {
|
||||
sum := sha256.Sum256([]byte(token))
|
||||
return hex.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
func newSessionToken() (plain string, hashHex string, err error) {
|
||||
buf := make([]byte, 32)
|
||||
if _, err = rand.Read(buf); err != nil {
|
||||
return "", "", err
|
||||
}
|
||||
plain = base64.RawURLEncoding.EncodeToString(buf)
|
||||
sum := sha256.Sum256([]byte(plain))
|
||||
return plain, hex.EncodeToString(sum[:]), nil
|
||||
}
|
||||
|
||||
func formatRetryAfter(d time.Duration) string {
|
||||
sec := int(d.Seconds())
|
||||
if sec < 1 {
|
||||
sec = 1
|
||||
}
|
||||
return itoa(sec)
|
||||
}
|
||||
|
||||
func itoa(n int) string {
|
||||
if n == 0 {
|
||||
return "0"
|
||||
}
|
||||
var b [16]byte
|
||||
i := len(b)
|
||||
for n > 0 {
|
||||
i--
|
||||
b[i] = byte('0' + n%10)
|
||||
n /= 10
|
||||
}
|
||||
return string(b[i:])
|
||||
}
|
||||
|
||||
func actorString(p principal) string {
|
||||
if p.Kind == authToken {
|
||||
return "token:" + p.TokenName
|
||||
}
|
||||
return "admin"
|
||||
}
|
||||
@@ -0,0 +1,138 @@
|
||||
package admin
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
const (
|
||||
cookieName = "nixmsg_admin"
|
||||
csrfHeader = "X-Nixmsg-Request"
|
||||
csrfValue = "1"
|
||||
adminUsername = "admin"
|
||||
settingAdminHash = "admin_password_hash"
|
||||
defaultSessionTTL = 12 * time.Hour
|
||||
minPasswordLen = 12
|
||||
lastUsedMinGap = time.Minute
|
||||
)
|
||||
|
||||
// Deps 是管理 Handler 的依赖。
|
||||
type Deps struct {
|
||||
DB *store.DB
|
||||
Hash auth.HashPool
|
||||
Tokens auth.APITokens
|
||||
Locks auth.LoginLocks
|
||||
Logger *slog.Logger
|
||||
|
||||
// TrustedProxies 受信任代理网段。
|
||||
TrustedProxies []*net.IPNet
|
||||
// SessionTTL 会话有效期;零值用 12 小时。
|
||||
SessionTTL time.Duration
|
||||
// SecureCookies 为 true 时 Cookie 始终带 Secure;否则按请求是否 HTTPS 决定。
|
||||
SecureCookies bool
|
||||
}
|
||||
|
||||
// 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
|
||||
|
||||
mux *http.ServeMux
|
||||
|
||||
lastUsedMu sync.Mutex
|
||||
lastUsed map[string]time.Time // api token id -> last DB write
|
||||
}
|
||||
|
||||
// New 构造可挂载的管理 Handler。返回值实现 http.Handler。
|
||||
func New(d Deps) *Handler {
|
||||
if d.Logger == nil {
|
||||
d.Logger = slog.Default()
|
||||
}
|
||||
if d.Locks == nil {
|
||||
d.Locks = NewMemoryLoginLocks()
|
||||
}
|
||||
ttl := d.SessionTTL
|
||||
if ttl <= 0 {
|
||||
ttl = defaultSessionTTL
|
||||
}
|
||||
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,
|
||||
mux: http.NewServeMux(),
|
||||
lastUsed: make(map[string]time.Time),
|
||||
}
|
||||
h.routes()
|
||||
return h
|
||||
}
|
||||
|
||||
// ServeHTTP 实现 http.Handler。
|
||||
func (h *Handler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
h.mux.ServeHTTP(w, r)
|
||||
}
|
||||
|
||||
func (h *Handler) routes() {
|
||||
// 公开
|
||||
h.mux.HandleFunc("POST /api/admin/login", h.handleLogin)
|
||||
|
||||
// 需鉴权
|
||||
h.mux.Handle("POST /api/admin/logout", h.auth(h.handleLogout))
|
||||
h.mux.Handle("GET /api/admin/me", h.auth(h.handleMe))
|
||||
h.mux.Handle("POST /api/admin/password", h.auth(h.handlePassword))
|
||||
|
||||
// 令牌管理:仅 Cookie
|
||||
h.mux.Handle("GET /api/admin/tokens", h.auth(h.handleTokenList))
|
||||
h.mux.Handle("POST /api/admin/tokens", h.auth(h.handleTokenCreate))
|
||||
h.mux.Handle("PATCH /api/admin/tokens/{id}", h.auth(h.handleTokenPatch))
|
||||
h.mux.Handle("DELETE /api/admin/tokens/{id}", h.auth(h.handleTokenDelete))
|
||||
|
||||
// 其余管理路由:鉴权生效,业务暂 501
|
||||
for _, p := range stubRoutes {
|
||||
h.mux.Handle(p, h.auth(h.handleNotImplemented))
|
||||
}
|
||||
}
|
||||
|
||||
var stubRoutes = []string{
|
||||
"GET /api/admin/overview",
|
||||
"GET /api/admin/endpoints",
|
||||
"POST /api/admin/endpoints",
|
||||
"POST /api/admin/endpoints/import",
|
||||
"POST /api/admin/endpoints/batch",
|
||||
"GET /api/admin/endpoints/{id}",
|
||||
"PATCH /api/admin/endpoints/{id}",
|
||||
"DELETE /api/admin/endpoints/{id}",
|
||||
"POST /api/admin/endpoints/{id}/kick",
|
||||
"POST /api/admin/endpoints/{id}/reset-login-password",
|
||||
"PUT /api/admin/endpoints/{id}/talk-password",
|
||||
"POST /api/admin/endpoints/{id}/unlock",
|
||||
"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",
|
||||
}
|
||||
@@ -0,0 +1,169 @@
|
||||
package admin
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/httpx"
|
||||
)
|
||||
|
||||
func (h *Handler) setSessionCookie(w http.ResponseWriter, r *http.Request, value string, maxAge int) {
|
||||
secure := h.forceSec || httpx.IsHTTPS(r, h.trusted)
|
||||
http.SetCookie(w, &http.Cookie{
|
||||
Name: cookieName,
|
||||
Value: value,
|
||||
Path: "/",
|
||||
HttpOnly: true,
|
||||
SameSite: http.SameSiteLaxMode,
|
||||
Secure: secure,
|
||||
MaxAge: maxAge,
|
||||
})
|
||||
}
|
||||
|
||||
func (h *Handler) handleLogin(w http.ResponseWriter, r *http.Request) {
|
||||
ip := httpx.ClientIP(r, h.trusted)
|
||||
if locked, retry := h.locks.Check(auth.LockKey{Kind: auth.LockAdminIP, IP: ip}); locked {
|
||||
w.Header().Set("Retry-After", formatRetryAfter(retry))
|
||||
httpx.WriteError(w, http.StatusTooManyRequests, "rate_limited", "登录已锁定,请稍后再试")
|
||||
return
|
||||
}
|
||||
|
||||
var req struct {
|
||||
Username string `json:"username"`
|
||||
Password string `json:"password"`
|
||||
}
|
||||
if err := httpx.DecodeJSON(r, &req); err != nil {
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效")
|
||||
return
|
||||
}
|
||||
if req.Username != adminUsername {
|
||||
h.failLogin(w, ip)
|
||||
return
|
||||
}
|
||||
|
||||
phc, err := h.getAdminPasswordHash(r.Context())
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
h.failLogin(w, ip)
|
||||
return
|
||||
}
|
||||
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
||||
return
|
||||
}
|
||||
ok, err := h.hash.Verify(r.Context(), auth.PasswordAdmin, req.Password, phc)
|
||||
if err != nil || !ok {
|
||||
h.failLogin(w, ip)
|
||||
return
|
||||
}
|
||||
|
||||
h.locks.Clear(auth.LockKey{Kind: auth.LockAdminIP, IP: ip})
|
||||
|
||||
plain, hashHex, err := newSessionToken()
|
||||
if err != nil {
|
||||
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
||||
return
|
||||
}
|
||||
if err := h.createSession(r.Context(), hashHex, h.ttl); err != nil {
|
||||
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
||||
return
|
||||
}
|
||||
h.setSessionCookie(w, r, plain, int(h.ttl.Seconds()))
|
||||
h.audit("admin", "login", "", "ok", ip)
|
||||
httpx.WriteOK(w, map[string]any{"username": adminUsername})
|
||||
}
|
||||
|
||||
func (h *Handler) failLogin(w http.ResponseWriter, ip string) {
|
||||
locked, retry := h.locks.Fail(auth.LockKey{Kind: auth.LockAdminIP, IP: ip})
|
||||
if locked {
|
||||
w.Header().Set("Retry-After", formatRetryAfter(retry))
|
||||
httpx.WriteError(w, http.StatusTooManyRequests, "rate_limited", "登录已锁定,请稍后再试")
|
||||
return
|
||||
}
|
||||
httpx.WriteError(w, http.StatusUnauthorized, "unauthorized", "用户名或密码错误")
|
||||
}
|
||||
|
||||
func (h *Handler) handleLogout(w http.ResponseWriter, r *http.Request) {
|
||||
p, _ := principalFrom(r.Context())
|
||||
ip := httpx.ClientIP(r, h.trusted)
|
||||
if p.Kind == authCookie && p.Session != "" {
|
||||
_ = h.deleteSession(r.Context(), hashSessionHex(p.Session))
|
||||
}
|
||||
h.setSessionCookie(w, r, "", -1)
|
||||
h.audit(actorString(p), "logout", "", "ok", ip)
|
||||
httpx.WriteOK(w, map[string]any{})
|
||||
}
|
||||
|
||||
func (h *Handler) handleMe(w http.ResponseWriter, r *http.Request) {
|
||||
p, _ := principalFrom(r.Context())
|
||||
authMode := "cookie"
|
||||
if p.Kind == authToken {
|
||||
authMode = "token"
|
||||
}
|
||||
httpx.WriteOK(w, map[string]any{
|
||||
"username": adminUsername,
|
||||
"auth": authMode,
|
||||
})
|
||||
}
|
||||
|
||||
func (h *Handler) handlePassword(w http.ResponseWriter, r *http.Request) {
|
||||
p, _ := principalFrom(r.Context())
|
||||
ip := httpx.ClientIP(r, h.trusted)
|
||||
|
||||
var req struct {
|
||||
OldPassword string `json:"old_password"`
|
||||
NewPassword string `json:"new_password"`
|
||||
}
|
||||
if err := httpx.DecodeJSON(r, &req); err != nil {
|
||||
h.audit(actorString(p), "password_change", "", "bad_request", ip)
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效")
|
||||
return
|
||||
}
|
||||
if len(req.NewPassword) < minPasswordLen {
|
||||
h.audit(actorString(p), "password_change", "", "bad_request", ip)
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "新密码至少 12 位")
|
||||
return
|
||||
}
|
||||
|
||||
phc, err := h.getAdminPasswordHash(r.Context())
|
||||
if err != nil {
|
||||
h.audit(actorString(p), "password_change", "", "error", ip)
|
||||
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
||||
return
|
||||
}
|
||||
ok, err := h.hash.Verify(r.Context(), auth.PasswordAdmin, req.OldPassword, phc)
|
||||
if err != nil || !ok {
|
||||
h.audit(actorString(p), "password_change", "", "unauthorized", ip)
|
||||
httpx.WriteError(w, http.StatusUnauthorized, "unauthorized", "旧密码错误")
|
||||
return
|
||||
}
|
||||
newPHC, err := h.hash.Hash(r.Context(), auth.PasswordAdmin, req.NewPassword)
|
||||
if err != nil {
|
||||
h.audit(actorString(p), "password_change", "", "error", ip)
|
||||
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
||||
return
|
||||
}
|
||||
if err := h.setAdminPasswordHash(r.Context(), newPHC); err != nil {
|
||||
h.audit(actorString(p), "password_change", "", "error", ip)
|
||||
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
||||
return
|
||||
}
|
||||
// 保留当前会话,作废其它会话
|
||||
if p.Session != "" {
|
||||
_ = h.deleteOtherSessions(r.Context(), hashSessionHex(p.Session))
|
||||
}
|
||||
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", "接口尚未实现")
|
||||
}
|
||||
@@ -0,0 +1,125 @@
|
||||
package admin
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
)
|
||||
|
||||
// MemoryLoginLocks 是内存登录锁定(重启清零)。
|
||||
// P3 正式实现合入前,A1 用本实现满足管理员 IP 锁定;参数对齐 DEVELOPMENT 第 5 节。
|
||||
type MemoryLoginLocks struct {
|
||||
mu sync.Mutex
|
||||
entries map[string]*lockState
|
||||
now func() time.Time
|
||||
}
|
||||
|
||||
type lockState struct {
|
||||
fails []time.Time
|
||||
lockedUntil time.Time
|
||||
}
|
||||
|
||||
// NewMemoryLoginLocks 创建内存锁定计数器。
|
||||
func NewMemoryLoginLocks() *MemoryLoginLocks {
|
||||
return &MemoryLoginLocks{
|
||||
entries: make(map[string]*lockState),
|
||||
now: time.Now,
|
||||
}
|
||||
}
|
||||
|
||||
func (l *MemoryLoginLocks) key(k auth.LockKey) string {
|
||||
return string(k.Kind) + "|" + k.EndpointID + "|" + k.PeerID + "|" + k.IP
|
||||
}
|
||||
|
||||
func (l *MemoryLoginLocks) params(kind auth.LockKind) (window time.Duration, threshold int, lockFor time.Duration) {
|
||||
switch kind {
|
||||
case auth.LockLoginEndpoint, auth.LockTalkTarget:
|
||||
return time.Hour, 50, time.Hour
|
||||
default:
|
||||
// LockLoginEndpointIP / LockTalkPair / LockAdminIP / LockRegisterIP
|
||||
return 5 * time.Minute, 10, 5 * time.Minute
|
||||
}
|
||||
}
|
||||
|
||||
// Check 实现 auth.LoginLocks。
|
||||
func (l *MemoryLoginLocks) Check(key auth.LockKey) (bool, time.Duration) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
now := l.now()
|
||||
st := l.entries[l.key(key)]
|
||||
if st == nil {
|
||||
return false, 0
|
||||
}
|
||||
if now.Before(st.lockedUntil) {
|
||||
return true, st.lockedUntil.Sub(now)
|
||||
}
|
||||
return false, 0
|
||||
}
|
||||
|
||||
// Fail 实现 auth.LoginLocks。
|
||||
func (l *MemoryLoginLocks) Fail(key auth.LockKey) (bool, time.Duration) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
now := l.now()
|
||||
k := l.key(key)
|
||||
st := l.entries[k]
|
||||
if st == nil {
|
||||
st = &lockState{}
|
||||
l.entries[k] = st
|
||||
}
|
||||
if now.Before(st.lockedUntil) {
|
||||
return true, st.lockedUntil.Sub(now)
|
||||
}
|
||||
window, threshold, lockFor := l.params(key.Kind)
|
||||
cutoff := now.Add(-window)
|
||||
kept := st.fails[:0]
|
||||
for _, t := range st.fails {
|
||||
if t.After(cutoff) {
|
||||
kept = append(kept, t)
|
||||
}
|
||||
}
|
||||
kept = append(kept, now)
|
||||
st.fails = kept
|
||||
if len(st.fails) >= threshold {
|
||||
st.lockedUntil = now.Add(lockFor)
|
||||
st.fails = nil
|
||||
return true, lockFor
|
||||
}
|
||||
return false, 0
|
||||
}
|
||||
|
||||
// ClearEndpoint 实现 auth.LoginLocks。
|
||||
func (l *MemoryLoginLocks) ClearEndpoint(endpointID string) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
for k := range l.entries {
|
||||
// kind|endpoint|peer|ip
|
||||
parts := splitLockKey(k)
|
||||
if len(parts) >= 2 && parts[1] == endpointID {
|
||||
delete(l.entries, k)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Clear 实现 auth.LoginLocks。
|
||||
func (l *MemoryLoginLocks) Clear(key auth.LockKey) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
delete(l.entries, l.key(key))
|
||||
}
|
||||
|
||||
func splitLockKey(k string) []string {
|
||||
out := make([]string, 0, 4)
|
||||
start := 0
|
||||
for i := 0; i < len(k); i++ {
|
||||
if k[i] == '|' {
|
||||
out = append(out, k[start:i])
|
||||
start = i + 1
|
||||
}
|
||||
}
|
||||
out = append(out, k[start:])
|
||||
return out
|
||||
}
|
||||
|
||||
var _ auth.LoginLocks = (*MemoryLoginLocks)(nil)
|
||||
@@ -0,0 +1,249 @@
|
||||
package admin
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
type apiTokenRow struct {
|
||||
ID string
|
||||
Name string
|
||||
Enabled bool
|
||||
CreatedAt time.Time
|
||||
LastUsedAt *time.Time
|
||||
}
|
||||
|
||||
func (h *Handler) getAdminPasswordHash(ctx context.Context) (string, error) {
|
||||
var v string
|
||||
err := h.db.Read.QueryRowContext(ctx, `SELECT value FROM settings WHERE key = ?`, settingAdminHash).Scan(&v)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
return v, nil
|
||||
}
|
||||
|
||||
func (h *Handler) setAdminPasswordHash(ctx context.Context, phc string) error {
|
||||
now := time.Now().UnixMilli()
|
||||
return h.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, err := tx.Exec(
|
||||
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)
|
||||
ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at`,
|
||||
settingAdminHash, phc, now,
|
||||
)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
func (h *Handler) createSession(ctx context.Context, hashHex string, ttl time.Duration) error {
|
||||
now := time.Now()
|
||||
return h.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, err := tx.Exec(
|
||||
`INSERT INTO admin_sessions(token_hash, created_at, expires_at) VALUES(?, ?, ?)`,
|
||||
hashHex, now.UnixMilli(), now.Add(ttl).UnixMilli(),
|
||||
)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
func (h *Handler) sessionValid(ctx context.Context, hashHex string) (bool, error) {
|
||||
var expires int64
|
||||
err := h.db.Read.QueryRowContext(ctx,
|
||||
`SELECT expires_at FROM admin_sessions WHERE token_hash = ?`, hashHex,
|
||||
).Scan(&expires)
|
||||
if err == sql.ErrNoRows {
|
||||
return false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return false, err
|
||||
}
|
||||
return time.Now().UnixMilli() < expires, nil
|
||||
}
|
||||
|
||||
func (h *Handler) deleteSession(ctx context.Context, hashHex string) error {
|
||||
return h.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, err := tx.Exec(`DELETE FROM admin_sessions WHERE token_hash = ?`, hashHex)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
func (h *Handler) deleteOtherSessions(ctx context.Context, keepHashHex string) error {
|
||||
return h.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, err := tx.Exec(`DELETE FROM admin_sessions WHERE token_hash != ?`, keepHashHex)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
func (h *Handler) lookupAPIToken(ctx context.Context, hash []byte) (apiTokenRow, error) {
|
||||
hashHex := hex.EncodeToString(hash)
|
||||
var (
|
||||
id, name string
|
||||
enabled int
|
||||
created, last sql.NullInt64
|
||||
)
|
||||
err := h.db.Read.QueryRowContext(ctx,
|
||||
`SELECT id, name, enabled, created_at, last_used_at FROM api_tokens WHERE token_hash = ?`,
|
||||
hashHex,
|
||||
).Scan(&id, &name, &enabled, &created, &last)
|
||||
if err != nil {
|
||||
return apiTokenRow{}, err
|
||||
}
|
||||
row := apiTokenRow{
|
||||
ID: id,
|
||||
Name: name,
|
||||
Enabled: enabled != 0,
|
||||
}
|
||||
if created.Valid {
|
||||
row.CreatedAt = time.UnixMilli(created.Int64)
|
||||
}
|
||||
if last.Valid {
|
||||
t := time.UnixMilli(last.Int64)
|
||||
row.LastUsedAt = &t
|
||||
}
|
||||
return row, nil
|
||||
}
|
||||
|
||||
func (h *Handler) listAPITokens(ctx context.Context) ([]apiTokenRow, error) {
|
||||
rows, err := h.db.Read.QueryContext(ctx,
|
||||
`SELECT id, name, enabled, created_at, last_used_at FROM api_tokens ORDER BY created_at ASC`,
|
||||
)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
var out []apiTokenRow
|
||||
for rows.Next() {
|
||||
var (
|
||||
id, name string
|
||||
enabled int
|
||||
created, last sql.NullInt64
|
||||
)
|
||||
if err := rows.Scan(&id, &name, &enabled, &created, &last); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
row := apiTokenRow{ID: id, Name: name, Enabled: enabled != 0}
|
||||
if created.Valid {
|
||||
row.CreatedAt = time.UnixMilli(created.Int64)
|
||||
}
|
||||
if last.Valid {
|
||||
t := time.UnixMilli(last.Int64)
|
||||
row.LastUsedAt = &t
|
||||
}
|
||||
out = append(out, row)
|
||||
}
|
||||
return out, rows.Err()
|
||||
}
|
||||
|
||||
func (h *Handler) insertAPIToken(ctx context.Context, id, name string, hash []byte) (time.Time, error) {
|
||||
now := time.Now()
|
||||
hashHex := hex.EncodeToString(hash)
|
||||
err := h.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, err := tx.Exec(
|
||||
`INSERT INTO api_tokens(id, name, token_hash, enabled, created_at, last_used_at) VALUES(?, ?, ?, 1, ?, NULL)`,
|
||||
id, name, hashHex, now.UnixMilli(),
|
||||
)
|
||||
return err
|
||||
})
|
||||
return now, err
|
||||
}
|
||||
|
||||
func (h *Handler) getAPITokenByID(ctx context.Context, id string) (apiTokenRow, error) {
|
||||
var (
|
||||
name string
|
||||
enabled int
|
||||
created, last sql.NullInt64
|
||||
)
|
||||
err := h.db.Read.QueryRowContext(ctx,
|
||||
`SELECT name, enabled, created_at, last_used_at FROM api_tokens WHERE id = ?`, id,
|
||||
).Scan(&name, &enabled, &created, &last)
|
||||
if err != nil {
|
||||
return apiTokenRow{}, err
|
||||
}
|
||||
row := apiTokenRow{ID: id, Name: name, Enabled: enabled != 0}
|
||||
if created.Valid {
|
||||
row.CreatedAt = time.UnixMilli(created.Int64)
|
||||
}
|
||||
if last.Valid {
|
||||
t := time.UnixMilli(last.Int64)
|
||||
row.LastUsedAt = &t
|
||||
}
|
||||
return row, nil
|
||||
}
|
||||
|
||||
func (h *Handler) updateAPIToken(ctx context.Context, id string, name *string, enabled *bool) error {
|
||||
return h.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
var (
|
||||
curName string
|
||||
curEn int
|
||||
)
|
||||
if err := tx.QueryRow(`SELECT name, enabled FROM api_tokens WHERE id = ?`, id).Scan(&curName, &curEn); err != nil {
|
||||
return err
|
||||
}
|
||||
newName := curName
|
||||
newEn := curEn
|
||||
if name != nil {
|
||||
newName = *name
|
||||
}
|
||||
if enabled != nil {
|
||||
if *enabled {
|
||||
newEn = 1
|
||||
} else {
|
||||
newEn = 0
|
||||
}
|
||||
}
|
||||
_, err := tx.Exec(`UPDATE api_tokens SET name = ?, enabled = ? WHERE id = ?`, newName, newEn, id)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
func (h *Handler) deleteAPIToken(ctx context.Context, id string) error {
|
||||
return h.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
res, err := tx.Exec(`DELETE FROM api_tokens WHERE id = ?`, id)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
n, _ := res.RowsAffected()
|
||||
if n == 0 {
|
||||
return sql.ErrNoRows
|
||||
}
|
||||
return nil
|
||||
})
|
||||
}
|
||||
|
||||
func (h *Handler) touchLastUsed(ctx context.Context, id string) {
|
||||
now := time.Now()
|
||||
h.lastUsedMu.Lock()
|
||||
prev, ok := h.lastUsed[id]
|
||||
if ok && now.Sub(prev) < lastUsedMinGap {
|
||||
h.lastUsedMu.Unlock()
|
||||
return
|
||||
}
|
||||
h.lastUsed[id] = now
|
||||
h.lastUsedMu.Unlock()
|
||||
|
||||
_ = h.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, err := tx.Exec(`UPDATE api_tokens SET last_used_at = ? WHERE id = ?`, now.UnixMilli(), id)
|
||||
return err
|
||||
})
|
||||
}
|
||||
|
||||
// SeedAdminPassword 写入管理员密码哈希(测试与接线辅助);走 HashPool。
|
||||
func SeedAdminPassword(ctx context.Context, db *store.DB, hash auth.HashPool, password string) error {
|
||||
phc, err := hash.Hash(ctx, auth.PasswordAdmin, password)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
now := time.Now().UnixMilli()
|
||||
return db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, err := tx.Exec(
|
||||
`INSERT INTO settings(key, value, updated_at) VALUES(?, ?, ?)
|
||||
ON CONFLICT(key) DO UPDATE SET value = excluded.value, updated_at = excluded.updated_at`,
|
||||
settingAdminHash, phc, now,
|
||||
)
|
||||
return err
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,165 @@
|
||||
package admin
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"net/http"
|
||||
"strings"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/httpx"
|
||||
)
|
||||
|
||||
func (h *Handler) handleTokenList(w http.ResponseWriter, r *http.Request) {
|
||||
items, err := h.listAPITokens(r.Context())
|
||||
if err != nil {
|
||||
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
||||
return
|
||||
}
|
||||
out := make([]map[string]any, 0, len(items))
|
||||
for _, it := range items {
|
||||
row := map[string]any{
|
||||
"id": it.ID,
|
||||
"name": it.Name,
|
||||
"enabled": it.Enabled,
|
||||
"created_at_ms": it.CreatedAt.UnixMilli(),
|
||||
"last_used_at_ms": nil,
|
||||
}
|
||||
if it.LastUsedAt != nil {
|
||||
row["last_used_at_ms"] = it.LastUsedAt.UnixMilli()
|
||||
}
|
||||
out = append(out, row)
|
||||
}
|
||||
httpx.WriteOK(w, map[string]any{
|
||||
"items": out,
|
||||
"next_cursor": "",
|
||||
"total": len(out),
|
||||
})
|
||||
}
|
||||
|
||||
func (h *Handler) handleTokenCreate(w http.ResponseWriter, r *http.Request) {
|
||||
p, _ := principalFrom(r.Context())
|
||||
ip := httpx.ClientIP(r, h.trusted)
|
||||
|
||||
var req struct {
|
||||
Name string `json:"name"`
|
||||
}
|
||||
if err := httpx.DecodeJSON(r, &req); err != nil || strings.TrimSpace(req.Name) == "" {
|
||||
h.audit(actorString(p), "token_create", "", "bad_request", ip)
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "名称不能为空")
|
||||
return
|
||||
}
|
||||
name := strings.TrimSpace(req.Name)
|
||||
|
||||
plain, hash, err := h.tokens.Issue(r.Context())
|
||||
if err != nil {
|
||||
h.audit(actorString(p), "token_create", "", "error", ip)
|
||||
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
||||
return
|
||||
}
|
||||
id, err := newTokenID()
|
||||
if err != nil {
|
||||
h.audit(actorString(p), "token_create", "", "error", ip)
|
||||
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
||||
return
|
||||
}
|
||||
created, err := h.insertAPIToken(r.Context(), id, name, hash)
|
||||
if err != nil {
|
||||
h.audit(actorString(p), "token_create", id, "error", ip)
|
||||
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
||||
return
|
||||
}
|
||||
h.audit(actorString(p), "token_create", id, "ok", ip)
|
||||
httpx.WriteOK(w, map[string]any{
|
||||
"id": id,
|
||||
"name": name,
|
||||
"token": plain,
|
||||
"created_at_ms": created.UnixMilli(),
|
||||
})
|
||||
}
|
||||
|
||||
func (h *Handler) handleTokenPatch(w http.ResponseWriter, r *http.Request) {
|
||||
p, _ := principalFrom(r.Context())
|
||||
ip := httpx.ClientIP(r, h.trusted)
|
||||
id := r.PathValue("id")
|
||||
|
||||
var req struct {
|
||||
Name *string `json:"name"`
|
||||
Enabled *bool `json:"enabled"`
|
||||
}
|
||||
if err := httpx.DecodeJSON(r, &req); err != nil {
|
||||
h.audit(actorString(p), "token_update", id, "bad_request", ip)
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "请求体无效")
|
||||
return
|
||||
}
|
||||
if req.Name == nil && req.Enabled == nil {
|
||||
h.audit(actorString(p), "token_update", id, "bad_request", ip)
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "无更新字段")
|
||||
return
|
||||
}
|
||||
if req.Name != nil {
|
||||
n := strings.TrimSpace(*req.Name)
|
||||
if n == "" {
|
||||
h.audit(actorString(p), "token_update", id, "bad_request", ip)
|
||||
httpx.WriteError(w, http.StatusBadRequest, "bad_request", "名称不能为空")
|
||||
return
|
||||
}
|
||||
req.Name = &n
|
||||
}
|
||||
|
||||
if err := h.updateAPIToken(r.Context(), id, req.Name, req.Enabled); err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
h.audit(actorString(p), "token_update", id, "not_found", ip)
|
||||
httpx.WriteError(w, http.StatusNotFound, "not_found", "令牌不存在")
|
||||
return
|
||||
}
|
||||
h.audit(actorString(p), "token_update", id, "error", ip)
|
||||
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
||||
return
|
||||
}
|
||||
row, err := h.getAPITokenByID(r.Context(), id)
|
||||
if err != nil {
|
||||
h.audit(actorString(p), "token_update", id, "error", ip)
|
||||
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
||||
return
|
||||
}
|
||||
h.audit(actorString(p), "token_update", id, "ok", ip)
|
||||
resp := map[string]any{
|
||||
"id": row.ID,
|
||||
"name": row.Name,
|
||||
"enabled": row.Enabled,
|
||||
"created_at_ms": row.CreatedAt.UnixMilli(),
|
||||
"last_used_at_ms": nil,
|
||||
}
|
||||
if row.LastUsedAt != nil {
|
||||
resp["last_used_at_ms"] = row.LastUsedAt.UnixMilli()
|
||||
}
|
||||
httpx.WriteOK(w, resp)
|
||||
}
|
||||
|
||||
func (h *Handler) handleTokenDelete(w http.ResponseWriter, r *http.Request) {
|
||||
p, _ := principalFrom(r.Context())
|
||||
ip := httpx.ClientIP(r, h.trusted)
|
||||
id := r.PathValue("id")
|
||||
if err := h.deleteAPIToken(r.Context(), id); err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
h.audit(actorString(p), "token_delete", id, "not_found", ip)
|
||||
httpx.WriteError(w, http.StatusNotFound, "not_found", "令牌不存在")
|
||||
return
|
||||
}
|
||||
h.audit(actorString(p), "token_delete", id, "error", ip)
|
||||
httpx.WriteError(w, http.StatusInternalServerError, "internal", "内部错误")
|
||||
return
|
||||
}
|
||||
h.audit(actorString(p), "token_delete", id, "ok", ip)
|
||||
httpx.WriteOK(w, map[string]any{})
|
||||
}
|
||||
|
||||
func newTokenID() (string, error) {
|
||||
b := make([]byte, 16)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return hex.EncodeToString(b), nil
|
||||
}
|
||||
@@ -0,0 +1,417 @@
|
||||
package identity
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/subtle"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
"unicode/utf8"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
const (
|
||||
maxRegisterBodyBytes = 4 * 1024
|
||||
|
||||
settingRegistrationEnabled = "registration_enabled"
|
||||
settingRegistrationCode = "registration_code"
|
||||
|
||||
sourceSelf = "self"
|
||||
|
||||
idAlphabet = "abcdefghijklmnopqrstuvwxyz0123456789"
|
||||
)
|
||||
|
||||
// APIError 是注册 HTTP/业务错误,带 HTTP 状态与协议错误码。
|
||||
type APIError struct {
|
||||
Status int
|
||||
Code string
|
||||
Message string
|
||||
}
|
||||
|
||||
func (e *APIError) Error() string {
|
||||
if e == nil {
|
||||
return ""
|
||||
}
|
||||
if e.Message == "" {
|
||||
return e.Code
|
||||
}
|
||||
return e.Code + ": " + e.Message
|
||||
}
|
||||
|
||||
func apiErr(status int, code, msg string) *APIError {
|
||||
return &APIError{Status: status, Code: code, Message: msg}
|
||||
}
|
||||
|
||||
// RegisterConfig 是可挂载注册处理器的依赖。
|
||||
// Hash / Locks 用 auth 接口;P3 未合入时测试可注入 StubHashPool 与可锁定的 LoginLocks。
|
||||
type RegisterConfig struct {
|
||||
DB *store.DB
|
||||
Hash auth.HashPool
|
||||
Locks auth.LoginLocks
|
||||
Logger *slog.Logger
|
||||
// Now 可测;nil 则用 time.Now。
|
||||
Now func() time.Time
|
||||
// ClientIP 可测;nil 则从 RemoteAddr 取 host。
|
||||
ClientIP func(*http.Request) string
|
||||
}
|
||||
|
||||
// RegisterHandler 处理 POST/OPTIONS /api/client/register(可挂到任意 ServeMux)。
|
||||
type RegisterHandler struct {
|
||||
cfg RegisterConfig
|
||||
}
|
||||
|
||||
// NewRegisterHandler 构造可挂载的注册 Handler。DB/Hash/Locks 必填。
|
||||
func NewRegisterHandler(cfg RegisterConfig) *RegisterHandler {
|
||||
if cfg.Logger == nil {
|
||||
cfg.Logger = slog.Default()
|
||||
}
|
||||
if cfg.Now == nil {
|
||||
cfg.Now = time.Now
|
||||
}
|
||||
if cfg.ClientIP == nil {
|
||||
cfg.ClientIP = clientIPFromRemoteAddr
|
||||
}
|
||||
return &RegisterHandler{cfg: cfg}
|
||||
}
|
||||
|
||||
// ServeHTTP 实现 http.Handler。
|
||||
func (h *RegisterHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) {
|
||||
setCORS(w)
|
||||
switch r.Method {
|
||||
case http.MethodOptions:
|
||||
w.Header().Set("Access-Control-Allow-Methods", "POST, OPTIONS")
|
||||
w.Header().Set("Access-Control-Allow-Headers", "Content-Type")
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
return
|
||||
case http.MethodPost:
|
||||
h.handlePost(w, r)
|
||||
default:
|
||||
writeRegisterError(w, apiErr(http.StatusMethodNotAllowed, protocol.CodeBadRequest, "method not allowed"))
|
||||
}
|
||||
}
|
||||
|
||||
func (h *RegisterHandler) handlePost(w http.ResponseWriter, r *http.Request) {
|
||||
ip := h.cfg.ClientIP(r)
|
||||
r.Body = http.MaxBytesReader(w, r.Body, maxRegisterBodyBytes)
|
||||
body, err := io.ReadAll(r.Body)
|
||||
if err != nil {
|
||||
var maxErr *http.MaxBytesError
|
||||
if errors.As(err, &maxErr) || errors.Is(err, io.ErrUnexpectedEOF) || isBodyTooLarge(err) {
|
||||
h.logResult("bad_request", "", ip)
|
||||
writeRegisterError(w, apiErr(http.StatusBadRequest, protocol.CodeBadRequest, "body too large"))
|
||||
return
|
||||
}
|
||||
h.logResult("bad_request", "", ip)
|
||||
writeRegisterError(w, apiErr(http.StatusBadRequest, protocol.CodeBadRequest, "read body failed"))
|
||||
return
|
||||
}
|
||||
|
||||
req, err := protocol.DecodeRegister(body)
|
||||
if err != nil {
|
||||
h.logResult("bad_request", "", ip)
|
||||
writeRegisterError(w, apiErr(http.StatusBadRequest, protocol.CodeBadRequest, "invalid json"))
|
||||
return
|
||||
}
|
||||
|
||||
result, apiErr := h.register(r.Context(), req, ip)
|
||||
if apiErr != nil {
|
||||
id := ""
|
||||
if req != nil {
|
||||
id = req.ID
|
||||
}
|
||||
h.logResult(apiErr.Code, id, ip)
|
||||
writeRegisterError(w, apiErr)
|
||||
return
|
||||
}
|
||||
h.logResult("ok", result.ID, ip)
|
||||
writeRegisterOK(w, result)
|
||||
}
|
||||
|
||||
func (h *RegisterHandler) register(ctx context.Context, req *protocol.RegisterRequest, ip string) (RegisterResult, *APIError) {
|
||||
enabled, storedCode, err := h.loadRegistrationSettings(ctx)
|
||||
if err != nil {
|
||||
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "settings unavailable")
|
||||
}
|
||||
if !enabled {
|
||||
return RegisterResult{}, apiErr(http.StatusForbidden, protocol.CodeRegistrationClosed, "registration closed")
|
||||
}
|
||||
|
||||
lockKey := auth.LockKey{Kind: auth.LockRegisterIP, IP: ip}
|
||||
if locked, _ := h.cfg.Locks.Check(lockKey); locked {
|
||||
return RegisterResult{}, apiErr(http.StatusTooManyRequests, protocol.CodeRateLimited, "rate limited")
|
||||
}
|
||||
|
||||
if !constantTimeEqual(req.RegistrationCode, storedCode) {
|
||||
h.cfg.Locks.Fail(lockKey)
|
||||
return RegisterResult{}, apiErr(http.StatusForbidden, protocol.CodeRegistrationCodeInvalid, "registration code invalid")
|
||||
}
|
||||
|
||||
if valErr := req.Validate(); valErr != nil {
|
||||
code, msg := protocol.CodeBadRequest, valErr.Error()
|
||||
var pe *protocol.Error
|
||||
if errors.As(valErr, &pe) && pe != nil {
|
||||
code, msg = pe.Code, pe.Message
|
||||
}
|
||||
return RegisterResult{}, apiErr(http.StatusBadRequest, code, msg)
|
||||
}
|
||||
|
||||
id := req.ID
|
||||
loginPassword := req.LoginPassword
|
||||
passwordGenerated := false
|
||||
if loginPassword == "" {
|
||||
pw, genErr := generateLoginPassword()
|
||||
if genErr != nil {
|
||||
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "generate password failed")
|
||||
}
|
||||
loginPassword = pw
|
||||
passwordGenerated = true
|
||||
}
|
||||
|
||||
loginHash, err := h.cfg.Hash.Hash(ctx, auth.PasswordLogin, loginPassword)
|
||||
if err != nil {
|
||||
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "hash failed")
|
||||
}
|
||||
|
||||
var talkHash sql.NullString
|
||||
if req.TalkPassword != "" {
|
||||
th, hashErr := h.cfg.Hash.Hash(ctx, auth.PasswordTalk, req.TalkPassword)
|
||||
if hashErr != nil {
|
||||
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "hash failed")
|
||||
}
|
||||
talkHash = sql.NullString{String: th, Valid: true}
|
||||
}
|
||||
|
||||
nowMs := h.cfg.Now().UnixMilli()
|
||||
const maxIDAttempts = 8
|
||||
for attempt := 0; attempt < maxIDAttempts; attempt++ {
|
||||
useID := id
|
||||
if useID == "" {
|
||||
genID, genErr := generateEndpointID()
|
||||
if genErr != nil {
|
||||
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "generate id failed")
|
||||
}
|
||||
useID = genID
|
||||
}
|
||||
|
||||
insertErr := h.cfg.DB.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, execErr := tx.ExecContext(ctx, `
|
||||
INSERT INTO endpoints(
|
||||
id, name, remark, source, login_hash, talk_hash, talk_version,
|
||||
default_delay_ms, enabled, created_at
|
||||
) VALUES (?, ?, '', ?, ?, ?, 0, 0, 1, ?)`,
|
||||
useID, req.Name, sourceSelf, loginHash, talkHash, nowMs,
|
||||
)
|
||||
return execErr
|
||||
})
|
||||
if insertErr == nil {
|
||||
out := RegisterResult{ID: useID}
|
||||
if passwordGenerated {
|
||||
out.LoginPassword = loginPassword
|
||||
}
|
||||
return out, nil
|
||||
}
|
||||
if isUniqueConstraint(insertErr) {
|
||||
if id != "" {
|
||||
return RegisterResult{}, apiErr(http.StatusConflict, protocol.CodeIDTaken, "id taken")
|
||||
}
|
||||
continue
|
||||
}
|
||||
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "insert failed")
|
||||
}
|
||||
return RegisterResult{}, apiErr(http.StatusServiceUnavailable, protocol.CodeBusy, "generate id exhausted")
|
||||
}
|
||||
|
||||
func (h *RegisterHandler) loadRegistrationSettings(ctx context.Context) (enabled bool, code string, err error) {
|
||||
var enabledVal, codeVal sql.NullString
|
||||
row := h.cfg.DB.Read.QueryRowContext(ctx, `SELECT value FROM settings WHERE key = ?`, settingRegistrationEnabled)
|
||||
if scanErr := row.Scan(&enabledVal); scanErr != nil && !errors.Is(scanErr, sql.ErrNoRows) {
|
||||
return false, "", scanErr
|
||||
}
|
||||
row = h.cfg.DB.Read.QueryRowContext(ctx, `SELECT value FROM settings WHERE key = ?`, settingRegistrationCode)
|
||||
if scanErr := row.Scan(&codeVal); scanErr != nil && !errors.Is(scanErr, sql.ErrNoRows) {
|
||||
return false, "", scanErr
|
||||
}
|
||||
return settingTruthy(enabledVal.String), codeVal.String, nil
|
||||
}
|
||||
|
||||
func (h *RegisterHandler) logResult(result, id, ip string) {
|
||||
h.cfg.Logger.Info("register", "result", result, "id", id, "ip", ip)
|
||||
}
|
||||
|
||||
// Server 实现 identity.Service:I1 只实现 Register,其余仍为未实现。
|
||||
type Server struct {
|
||||
handler *RegisterHandler
|
||||
}
|
||||
|
||||
// NewServer 用同一套依赖构造 Service(Register)与可挂载 Handler。
|
||||
func NewServer(cfg RegisterConfig) *Server {
|
||||
return &Server{handler: NewRegisterHandler(cfg)}
|
||||
}
|
||||
|
||||
// Handler 返回可挂载的注册 HTTP 处理器。
|
||||
func (s *Server) Handler() http.Handler { return s.handler }
|
||||
|
||||
// Register 实现自助注册(source 固定为 self;RemoteIP 用于锁定)。
|
||||
func (s *Server) Register(ctx context.Context, req RegisterRequest) (RegisterResult, error) {
|
||||
preq := &protocol.RegisterRequest{
|
||||
RegistrationCode: req.RegistrationCode,
|
||||
ID: req.ID,
|
||||
LoginPassword: req.LoginPassword,
|
||||
Name: req.Name,
|
||||
TalkPassword: req.TalkPassword,
|
||||
}
|
||||
ip := req.RemoteIP
|
||||
if ip == "" {
|
||||
ip = "0.0.0.0"
|
||||
}
|
||||
result, err := s.handler.register(ctx, preq, ip)
|
||||
if err != nil {
|
||||
return RegisterResult{}, err
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (s *Server) SelfGet(context.Context, string) (SelfInfo, error) {
|
||||
return SelfInfo{}, ErrNotImplemented
|
||||
}
|
||||
func (s *Server) SelfUpdate(context.Context, string, *protocol.SelfUpdate) error {
|
||||
return ErrNotImplemented
|
||||
}
|
||||
func (s *Server) SelfSetTalkPassword(context.Context, string, string) error {
|
||||
return ErrNotImplemented
|
||||
}
|
||||
func (s *Server) SelfChangeLoginPassword(context.Context, string, string, string) (string, error) {
|
||||
return "", ErrNotImplemented
|
||||
}
|
||||
func (s *Server) SelfLogout(context.Context, string) error { return ErrNotImplemented }
|
||||
func (s *Server) UnlockTalk(context.Context, string, string, string) error {
|
||||
return ErrNotImplemented
|
||||
}
|
||||
func (s *Server) HasTalkGrant(context.Context, string, string) (bool, error) {
|
||||
return false, nil
|
||||
}
|
||||
func (s *Server) Disable(context.Context, string) error { return ErrNotImplemented }
|
||||
func (s *Server) Enable(context.Context, string) error { return ErrNotImplemented }
|
||||
func (s *Server) Delete(context.Context, string) error { return ErrNotImplemented }
|
||||
|
||||
var _ Service = (*Server)(nil)
|
||||
var _ http.Handler = (*RegisterHandler)(nil)
|
||||
|
||||
func setCORS(w http.ResponseWriter) {
|
||||
w.Header().Set("Access-Control-Allow-Origin", "*")
|
||||
}
|
||||
|
||||
func writeRegisterOK(w http.ResponseWriter, result RegisterResult) {
|
||||
setCORS(w)
|
||||
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
w.WriteHeader(http.StatusOK)
|
||||
_ = protocol.Encode(w, protocol.RegisterResponse{
|
||||
OK: true,
|
||||
Data: protocol.RegisterData{
|
||||
ID: result.ID,
|
||||
LoginPassword: result.LoginPassword,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
func writeRegisterError(w http.ResponseWriter, err *APIError) {
|
||||
setCORS(w)
|
||||
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
status := http.StatusInternalServerError
|
||||
code, msg := protocol.CodeBusy, "internal error"
|
||||
if err != nil {
|
||||
status = err.Status
|
||||
code, msg = err.Code, err.Message
|
||||
}
|
||||
w.WriteHeader(status)
|
||||
_ = protocol.Encode(w, protocol.RegisterResponse{
|
||||
OK: false,
|
||||
Error: &protocol.ErrorBody{Code: code, Message: msg},
|
||||
})
|
||||
}
|
||||
|
||||
func clientIPFromRemoteAddr(r *http.Request) string {
|
||||
host, _, err := net.SplitHostPort(r.RemoteAddr)
|
||||
if err != nil {
|
||||
return r.RemoteAddr
|
||||
}
|
||||
return host
|
||||
}
|
||||
|
||||
func settingTruthy(v string) bool {
|
||||
switch strings.TrimSpace(strings.ToLower(v)) {
|
||||
case "1", "true", "yes", "on":
|
||||
return true
|
||||
default:
|
||||
return false
|
||||
}
|
||||
}
|
||||
|
||||
func constantTimeEqual(a, b string) bool {
|
||||
// 长度不同时 ConstantTimeCompare 直接失败;先按较短侧对齐比较,再核对长度,避免过早返回。
|
||||
ab := []byte(a)
|
||||
bb := []byte(b)
|
||||
if len(ab) != len(bb) {
|
||||
dummy := make([]byte, len(ab))
|
||||
subtle.ConstantTimeCompare(ab, dummy)
|
||||
return false
|
||||
}
|
||||
return subtle.ConstantTimeCompare(ab, bb) == 1
|
||||
}
|
||||
|
||||
func generateEndpointID() (string, error) {
|
||||
b := make([]byte, 8)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
out := make([]byte, 8)
|
||||
for i := range b {
|
||||
out[i] = idAlphabet[int(b[i])%len(idAlphabet)]
|
||||
}
|
||||
return "e_" + string(out), nil
|
||||
}
|
||||
|
||||
func generateLoginPassword() (string, error) {
|
||||
// 20 字节可读字符,满足 8–128,且不以 nst_ 开头(字母数字混合,冲突概率极低;若撞前缀则重抽)。
|
||||
const alphabet = "ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789"
|
||||
for range 8 {
|
||||
b := make([]byte, 20)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
out := make([]byte, 20)
|
||||
for i := range b {
|
||||
out[i] = alphabet[int(b[i])%len(alphabet)]
|
||||
}
|
||||
pw := string(out)
|
||||
if !strings.HasPrefix(pw, protocol.SessionTokenPrefix) && utf8.RuneCountInString(pw) >= protocol.MinLoginPasswordLen {
|
||||
return pw, nil
|
||||
}
|
||||
}
|
||||
return "", errors.New("identity: generate login password failed")
|
||||
}
|
||||
|
||||
func isUniqueConstraint(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
msg := strings.ToLower(err.Error())
|
||||
return strings.Contains(msg, "unique constraint") || strings.Contains(msg, "constraint failed")
|
||||
}
|
||||
|
||||
func isBodyTooLarge(err error) bool {
|
||||
if err == nil {
|
||||
return false
|
||||
}
|
||||
msg := strings.ToLower(err.Error())
|
||||
return strings.Contains(msg, "request body too large") || strings.Contains(msg, "http: request body too large")
|
||||
}
|
||||
@@ -0,0 +1,446 @@
|
||||
package identity
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"strings"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
// registerIPLocker 仅实现 LockRegisterIP:5 分钟窗口内 10 次失败则锁定 5 分钟。
|
||||
// P3 的完整 LoginLocks 未合入本分支时,测试用此可替换实现覆盖 F23 锁定验收。
|
||||
type registerIPLocker struct {
|
||||
mu sync.Mutex
|
||||
fails map[string][]time.Time
|
||||
lockedUntil map[string]time.Time
|
||||
now func() time.Time
|
||||
window time.Duration
|
||||
limit int
|
||||
lockFor time.Duration
|
||||
}
|
||||
|
||||
func newRegisterIPLocker(now func() time.Time) *registerIPLocker {
|
||||
if now == nil {
|
||||
now = time.Now
|
||||
}
|
||||
return ®isterIPLocker{
|
||||
fails: make(map[string][]time.Time),
|
||||
lockedUntil: make(map[string]time.Time),
|
||||
now: now,
|
||||
window: 5 * time.Minute,
|
||||
limit: 10,
|
||||
lockFor: 5 * time.Minute,
|
||||
}
|
||||
}
|
||||
|
||||
func (l *registerIPLocker) Check(key auth.LockKey) (bool, time.Duration) {
|
||||
if key.Kind != auth.LockRegisterIP {
|
||||
return false, 0
|
||||
}
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
until, ok := l.lockedUntil[key.IP]
|
||||
if !ok {
|
||||
return false, 0
|
||||
}
|
||||
now := l.now()
|
||||
if now.Before(until) {
|
||||
return true, until.Sub(now)
|
||||
}
|
||||
delete(l.lockedUntil, key.IP)
|
||||
return false, 0
|
||||
}
|
||||
|
||||
func (l *registerIPLocker) Fail(key auth.LockKey) (bool, time.Duration) {
|
||||
if key.Kind != auth.LockRegisterIP {
|
||||
return false, 0
|
||||
}
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
now := l.now()
|
||||
if until, ok := l.lockedUntil[key.IP]; ok && now.Before(until) {
|
||||
return true, until.Sub(now)
|
||||
}
|
||||
cutoff := now.Add(-l.window)
|
||||
list := l.fails[key.IP]
|
||||
kept := list[:0]
|
||||
for _, t := range list {
|
||||
if t.After(cutoff) {
|
||||
kept = append(kept, t)
|
||||
}
|
||||
}
|
||||
kept = append(kept, now)
|
||||
l.fails[key.IP] = kept
|
||||
if len(kept) >= l.limit {
|
||||
until := now.Add(l.lockFor)
|
||||
l.lockedUntil[key.IP] = until
|
||||
return true, l.lockFor
|
||||
}
|
||||
return false, 0
|
||||
}
|
||||
|
||||
func (l *registerIPLocker) ClearEndpoint(string) {}
|
||||
func (l *registerIPLocker) Clear(key auth.LockKey) {
|
||||
l.mu.Lock()
|
||||
defer l.mu.Unlock()
|
||||
delete(l.fails, key.IP)
|
||||
delete(l.lockedUntil, key.IP)
|
||||
}
|
||||
|
||||
var _ auth.LoginLocks = (*registerIPLocker)(nil)
|
||||
|
||||
type testEnv struct {
|
||||
db *store.DB
|
||||
hash auth.HashPool
|
||||
locks *registerIPLocker
|
||||
logBuf *bytes.Buffer
|
||||
handler http.Handler
|
||||
fixedIP string
|
||||
}
|
||||
|
||||
func openTestEnv(t *testing.T) *testEnv {
|
||||
t.Helper()
|
||||
db, err := store.Open(t.TempDir(), "FULL")
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
t.Cleanup(func() { _ = db.Close() })
|
||||
|
||||
buf := &bytes.Buffer{}
|
||||
logger := slog.New(slog.NewTextHandler(buf, &slog.HandlerOptions{Level: slog.LevelInfo}))
|
||||
locks := newRegisterIPLocker(time.Now)
|
||||
env := &testEnv{
|
||||
db: db,
|
||||
hash: auth.NewStubHashPool(),
|
||||
locks: locks,
|
||||
logBuf: buf,
|
||||
fixedIP: "203.0.113.10",
|
||||
}
|
||||
env.handler = NewRegisterHandler(RegisterConfig{
|
||||
DB: db,
|
||||
Hash: env.hash,
|
||||
Locks: locks,
|
||||
Logger: logger,
|
||||
ClientIP: func(*http.Request) string {
|
||||
return env.fixedIP
|
||||
},
|
||||
})
|
||||
return env
|
||||
}
|
||||
|
||||
func (e *testEnv) setRegistration(t *testing.T, enabled bool, code string) {
|
||||
t.Helper()
|
||||
en := "0"
|
||||
if enabled {
|
||||
en = "1"
|
||||
}
|
||||
now := time.Now().UnixMilli()
|
||||
err := e.db.Queue.Do(context.Background(), func(tx *sql.Tx) error {
|
||||
if _, err := 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, en, now); err != nil {
|
||||
return err
|
||||
}
|
||||
_, err := 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, now)
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func (e *testEnv) insertEndpoint(t *testing.T, id, loginHash string) {
|
||||
t.Helper()
|
||||
err := e.db.Queue.Do(context.Background(), func(tx *sql.Tx) error {
|
||||
_, err := tx.Exec(`INSERT INTO endpoints(
|
||||
id, name, remark, source, login_hash, talk_hash, talk_version,
|
||||
default_delay_ms, enabled, created_at
|
||||
) VALUES (?, '', '', 'admin', ?, NULL, 0, 0, 1, ?)`, id, loginHash, time.Now().UnixMilli())
|
||||
return err
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func (e *testEnv) getEndpoint(t *testing.T, id string) (source, loginHash string, ok bool) {
|
||||
t.Helper()
|
||||
err := e.db.Read.QueryRow(`SELECT source, login_hash FROM endpoints WHERE id = ?`, id).Scan(&source, &loginHash)
|
||||
if errorsIsNoRows(err) {
|
||||
return "", "", false
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
return source, loginHash, true
|
||||
}
|
||||
|
||||
func errorsIsNoRows(err error) bool {
|
||||
return err == sql.ErrNoRows
|
||||
}
|
||||
|
||||
type registerResp struct {
|
||||
OK bool `json:"ok"`
|
||||
Data struct {
|
||||
ID string `json:"id"`
|
||||
LoginPassword string `json:"login_password"`
|
||||
} `json:"data"`
|
||||
Error *protocol.ErrorBody `json:"error"`
|
||||
}
|
||||
|
||||
func (e *testEnv) doRegister(t *testing.T, body string) (int, registerResp, http.Header) {
|
||||
t.Helper()
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/client/register", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
req.RemoteAddr = e.fixedIP + ":54321"
|
||||
rr := httptest.NewRecorder()
|
||||
e.handler.ServeHTTP(rr, req)
|
||||
var resp registerResp
|
||||
if err := json.Unmarshal(rr.Body.Bytes(), &resp); err != nil {
|
||||
t.Fatalf("decode resp: %v body=%s", err, rr.Body.String())
|
||||
}
|
||||
return rr.Code, resp, rr.Header()
|
||||
}
|
||||
|
||||
func TestRegisterF23_ClosedFails(t *testing.T) {
|
||||
env := openTestEnv(t)
|
||||
env.setRegistration(t, false, "secretcode")
|
||||
|
||||
code, resp, hdr := env.doRegister(t, `{"registration_code":"secretcode","id":"ep_closed","login_password":"password1"}`)
|
||||
if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationClosed {
|
||||
t.Fatalf("status=%d resp=%+v", code, resp)
|
||||
}
|
||||
if hdr.Get("Access-Control-Allow-Origin") != "*" {
|
||||
t.Fatalf("missing CORS: %v", hdr)
|
||||
}
|
||||
if _, _, ok := env.getEndpoint(t, "ep_closed"); ok {
|
||||
t.Fatal("endpoint should not be created when closed")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterF23_WrongCodeFails_RightCodeOK(t *testing.T) {
|
||||
env := openTestEnv(t)
|
||||
env.setRegistration(t, true, "good-code-01")
|
||||
|
||||
code, resp, _ := env.doRegister(t, `{"registration_code":"bad-code-xx","id":"ep_wrong","login_password":"password1"}`)
|
||||
if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationCodeInvalid {
|
||||
t.Fatalf("wrong code: status=%d resp=%+v", code, resp)
|
||||
}
|
||||
|
||||
code, resp, hdr := env.doRegister(t, `{"registration_code":"good-code-01","id":"ep_ok1","login_password":"password1","name":"门口"}`)
|
||||
if code != http.StatusOK || !resp.OK || resp.Data.ID != "ep_ok1" {
|
||||
t.Fatalf("ok register: status=%d resp=%+v", code, resp)
|
||||
}
|
||||
if resp.Data.LoginPassword != "" {
|
||||
t.Fatalf("provided password must not echo: %q", resp.Data.LoginPassword)
|
||||
}
|
||||
if hdr.Get("Access-Control-Allow-Origin") != "*" {
|
||||
t.Fatal("missing CORS on success")
|
||||
}
|
||||
source, loginHash, ok := env.getEndpoint(t, "ep_ok1")
|
||||
if !ok || source != "self" {
|
||||
t.Fatalf("endpoint source=%q ok=%v", source, ok)
|
||||
}
|
||||
match, err := env.hash.Verify(context.Background(), auth.PasswordLogin, "password1", loginHash)
|
||||
if err != nil || !match {
|
||||
t.Fatalf("login hash verify: match=%v err=%v", match, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterF23_ChangeCode_OldFails_ExistingRemains(t *testing.T) {
|
||||
env := openTestEnv(t)
|
||||
env.setRegistration(t, true, "code-old-01")
|
||||
|
||||
code, resp, _ := env.doRegister(t, `{"registration_code":"code-old-01","id":"ep_keep","login_password":"password1"}`)
|
||||
if code != http.StatusOK || resp.Data.ID != "ep_keep" {
|
||||
t.Fatalf("first register: status=%d resp=%+v", code, resp)
|
||||
}
|
||||
_, oldHash, ok := env.getEndpoint(t, "ep_keep")
|
||||
if !ok {
|
||||
t.Fatal("missing endpoint after register")
|
||||
}
|
||||
|
||||
env.setRegistration(t, true, "code-new-02")
|
||||
code, resp, _ = env.doRegister(t, `{"registration_code":"code-old-01","id":"ep_new","login_password":"password1"}`)
|
||||
if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationCodeInvalid {
|
||||
t.Fatalf("old code after rotate: status=%d resp=%+v", code, resp)
|
||||
}
|
||||
code, resp, _ = env.doRegister(t, `{"registration_code":"code-new-02","id":"ep_new","login_password":"password1"}`)
|
||||
if code != http.StatusOK || resp.Data.ID != "ep_new" {
|
||||
t.Fatalf("new code: status=%d resp=%+v", code, resp)
|
||||
}
|
||||
|
||||
_, hashAfter, ok := env.getEndpoint(t, "ep_keep")
|
||||
if !ok || hashAfter != oldHash {
|
||||
t.Fatalf("existing endpoint mutated: ok=%v hashEqual=%v", ok, hashAfter == oldHash)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterF23_WrongCodeLock(t *testing.T) {
|
||||
env := openTestEnv(t)
|
||||
env.setRegistration(t, true, "lock-code-1")
|
||||
|
||||
for i := 0; i < 10; i++ {
|
||||
code, resp, _ := env.doRegister(t, `{"registration_code":"wrong-code","id":"ep_lock","login_password":"password1"}`)
|
||||
if code != http.StatusForbidden || resp.Error == nil || resp.Error.Code != protocol.CodeRegistrationCodeInvalid {
|
||||
t.Fatalf("fail #%d: status=%d resp=%+v", i+1, code, resp)
|
||||
}
|
||||
}
|
||||
code, resp, _ := env.doRegister(t, `{"registration_code":"lock-code-1","id":"ep_lock","login_password":"password1"}`)
|
||||
if code != http.StatusTooManyRequests || resp.Error == nil || resp.Error.Code != protocol.CodeRateLimited {
|
||||
t.Fatalf("locked with good code: status=%d resp=%+v", code, resp)
|
||||
}
|
||||
if _, _, ok := env.getEndpoint(t, "ep_lock"); ok {
|
||||
t.Fatal("must not insert while rate limited")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegisterF23_IDTakenKeepsOriginal(t *testing.T) {
|
||||
env := openTestEnv(t)
|
||||
env.setRegistration(t, true, "taken-code")
|
||||
env.insertEndpoint(t, "ep_taken", "stub$original-password-xx")
|
||||
|
||||
code, resp, _ := env.doRegister(t, `{"registration_code":"taken-code","id":"ep_taken","login_password":"password1"}`)
|
||||
if code != http.StatusConflict || resp.Error == nil || resp.Error.Code != protocol.CodeIDTaken {
|
||||
t.Fatalf("id taken: status=%d resp=%+v", code, resp)
|
||||
}
|
||||
source, loginHash, ok := env.getEndpoint(t, "ep_taken")
|
||||
if !ok || source != "admin" || loginHash != "stub$original-password-xx" {
|
||||
t.Fatalf("original endpoint changed: source=%q hash=%q", source, loginHash)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegister_GenerateIDAndPassword(t *testing.T) {
|
||||
env := openTestEnv(t)
|
||||
env.setRegistration(t, true, "gen-code-01")
|
||||
|
||||
code, resp, _ := env.doRegister(t, `{"registration_code":"gen-code-01","id":"","login_password":""}`)
|
||||
if code != http.StatusOK || !resp.OK {
|
||||
t.Fatalf("status=%d resp=%+v", code, resp)
|
||||
}
|
||||
if !strings.HasPrefix(resp.Data.ID, "e_") || len(resp.Data.ID) != 10 {
|
||||
t.Fatalf("generated id=%q", resp.Data.ID)
|
||||
}
|
||||
if len(resp.Data.LoginPassword) < protocol.MinLoginPasswordLen {
|
||||
t.Fatalf("generated password too short: %q", resp.Data.LoginPassword)
|
||||
}
|
||||
if strings.HasPrefix(resp.Data.LoginPassword, protocol.SessionTokenPrefix) {
|
||||
t.Fatal("generated password starts with nst_")
|
||||
}
|
||||
source, _, ok := env.getEndpoint(t, resp.Data.ID)
|
||||
if !ok || source != "self" {
|
||||
t.Fatalf("source=%q ok=%v", source, ok)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegister_OPTIONS_CORS(t *testing.T) {
|
||||
env := openTestEnv(t)
|
||||
req := httptest.NewRequest(http.MethodOptions, "/api/client/register", nil)
|
||||
rr := httptest.NewRecorder()
|
||||
env.handler.ServeHTTP(rr, req)
|
||||
if rr.Code != http.StatusNoContent {
|
||||
t.Fatalf("status=%d", rr.Code)
|
||||
}
|
||||
if rr.Header().Get("Access-Control-Allow-Origin") != "*" {
|
||||
t.Fatal("missing Allow-Origin")
|
||||
}
|
||||
if !strings.Contains(rr.Header().Get("Access-Control-Allow-Methods"), "POST") {
|
||||
t.Fatalf("methods=%q", rr.Header().Get("Access-Control-Allow-Methods"))
|
||||
}
|
||||
if len(rr.Result().Cookies()) != 0 {
|
||||
t.Fatal("must not set cookies")
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegister_BodyTooLarge(t *testing.T) {
|
||||
env := openTestEnv(t)
|
||||
env.setRegistration(t, true, "big-code-01")
|
||||
body := `{"registration_code":"big-code-01","id":"ep_big","login_password":"password1","name":"` + strings.Repeat("x", 5000) + `"}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/client/register", strings.NewReader(body))
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
rr := httptest.NewRecorder()
|
||||
env.handler.ServeHTTP(rr, req)
|
||||
if rr.Code != http.StatusBadRequest {
|
||||
t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String())
|
||||
}
|
||||
var resp registerResp
|
||||
_ = json.Unmarshal(rr.Body.Bytes(), &resp)
|
||||
if resp.Error == nil || resp.Error.Code != protocol.CodeBadRequest {
|
||||
t.Fatalf("resp=%+v", resp)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegister_LogOmitsSecrets(t *testing.T) {
|
||||
env := openTestEnv(t)
|
||||
env.setRegistration(t, true, "log-secret-code")
|
||||
_, _, _ = env.doRegister(t, `{"registration_code":"log-secret-code","id":"ep_log","login_password":"supersecretpw"}`)
|
||||
logged := env.logBuf.String()
|
||||
if strings.Contains(logged, "log-secret-code") || strings.Contains(logged, "supersecretpw") {
|
||||
t.Fatalf("log leaked secrets: %s", logged)
|
||||
}
|
||||
if !strings.Contains(logged, "ep_log") || !strings.Contains(logged, env.fixedIP) {
|
||||
t.Fatalf("log missing id/ip: %s", logged)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRegister_NSTPasswordRejected(t *testing.T) {
|
||||
env := openTestEnv(t)
|
||||
env.setRegistration(t, true, "nst-code-01")
|
||||
code, resp, _ := env.doRegister(t, `{"registration_code":"nst-code-01","id":"ep_nst","login_password":"nst_notallowed"}`)
|
||||
if code != http.StatusBadRequest || resp.Error == nil || resp.Error.Code != protocol.CodeBadRequest {
|
||||
t.Fatalf("status=%d resp=%+v", code, resp)
|
||||
}
|
||||
}
|
||||
|
||||
func TestMountOnServeMux(t *testing.T) {
|
||||
env := openTestEnv(t)
|
||||
env.setRegistration(t, true, "mux-code-01")
|
||||
mux := http.NewServeMux()
|
||||
mux.Handle("/api/client/register", env.handler)
|
||||
|
||||
body := `{"registration_code":"mux-code-01","id":"ep_mux","login_password":"password1"}`
|
||||
req := httptest.NewRequest(http.MethodPost, "/api/client/register", strings.NewReader(body))
|
||||
rr := httptest.NewRecorder()
|
||||
mux.ServeHTTP(rr, req)
|
||||
if rr.Code != http.StatusOK {
|
||||
t.Fatalf("status=%d body=%s", rr.Code, rr.Body.String())
|
||||
}
|
||||
}
|
||||
|
||||
func TestServerRegisterInterface(t *testing.T) {
|
||||
env := openTestEnv(t)
|
||||
env.setRegistration(t, true, "svc-code-01")
|
||||
svc := NewServer(RegisterConfig{
|
||||
DB: env.db,
|
||||
Hash: env.hash,
|
||||
Locks: env.locks,
|
||||
Logger: slog.New(slog.NewTextHandler(io.Discard, nil)),
|
||||
})
|
||||
res, err := svc.Register(context.Background(), RegisterRequest{
|
||||
RegistrationCode: "svc-code-01",
|
||||
ID: "ep_svc",
|
||||
LoginPassword: "password1",
|
||||
RemoteIP: "198.51.100.1",
|
||||
Source: "self",
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if res.ID != "ep_svc" {
|
||||
t.Fatalf("id=%q", res.ID)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,153 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"context"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/config"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
// 消息状态(DEVELOPMENT 7.1)。
|
||||
const (
|
||||
StateScheduled = "scheduled"
|
||||
StateDispatched = "dispatched"
|
||||
StateCompleted = "completed"
|
||||
)
|
||||
|
||||
// 投递状态。
|
||||
const (
|
||||
DeliveryPending = "pending"
|
||||
)
|
||||
|
||||
// talk_grants.kind。
|
||||
const (
|
||||
GrantKindPassword = "password"
|
||||
GrantKindReply = "reply"
|
||||
)
|
||||
|
||||
// 请求频率桶默认突发容量(DEVELOPMENT 6.10;配置无单独字段)。
|
||||
const defaultRequestBurst = 100
|
||||
|
||||
// Limits 是提交所需的配置上限(来自 config.LimitsConfig)。
|
||||
type Limits struct {
|
||||
MaxBodyBytes int
|
||||
MaxMetaBytes int
|
||||
MaxFrameBytes int
|
||||
MaxTTLSeconds int64
|
||||
MaxScheduleSeconds int64
|
||||
RequestsPerSecond float64
|
||||
RequestBurst int
|
||||
MaxPendingPerSender int
|
||||
MaxPendingPerReceiver int
|
||||
GraceSeconds int64
|
||||
}
|
||||
|
||||
// LimitsFromConfig 从平台配置构造 Limits。
|
||||
func LimitsFromConfig(c config.LimitsConfig) Limits {
|
||||
burst := defaultRequestBurst
|
||||
return Limits{
|
||||
MaxBodyBytes: c.MaxBodyBytes,
|
||||
MaxMetaBytes: c.MaxMetaBytes,
|
||||
MaxFrameBytes: c.MaxFrameBytes,
|
||||
MaxTTLSeconds: int64(c.MaxTTLSeconds),
|
||||
MaxScheduleSeconds: int64(c.MaxScheduleSeconds),
|
||||
RequestsPerSecond: float64(c.RequestsPerSecond),
|
||||
RequestBurst: burst,
|
||||
MaxPendingPerSender: c.MaxPendingPerSender,
|
||||
MaxPendingPerReceiver: c.MaxPendingPerReceiver,
|
||||
GraceSeconds: int64(c.GraceSeconds),
|
||||
}
|
||||
}
|
||||
|
||||
// App 实现 Service 的提交路径(M1);其余方法暂返回未实现或空操作。
|
||||
type App struct {
|
||||
db *store.DB
|
||||
lim Limits
|
||||
hash auth.HashPool
|
||||
locks auth.LoginLocks
|
||||
nowFn func() time.Time
|
||||
rates *rateLimiter
|
||||
}
|
||||
|
||||
// Option 配置 App。
|
||||
type Option func(*App)
|
||||
|
||||
// WithNow 注入时钟(测试用)。
|
||||
func WithNow(now func() time.Time) Option {
|
||||
return func(a *App) { a.nowFn = now }
|
||||
}
|
||||
|
||||
// WithLocks 注入对话密码锁定计数器;nil 表示不锁定。
|
||||
func WithLocks(locks auth.LoginLocks) Option {
|
||||
return func(a *App) { a.locks = locks }
|
||||
}
|
||||
|
||||
// New 创建消息服务实现。hash 用于校验对话密码;locks 可为 nil。
|
||||
func New(db *store.DB, lim Limits, hash auth.HashPool, opts ...Option) *App {
|
||||
if lim.RequestBurst <= 0 {
|
||||
lim.RequestBurst = defaultRequestBurst
|
||||
}
|
||||
a := &App{
|
||||
db: db,
|
||||
lim: lim,
|
||||
hash: hash,
|
||||
nowFn: time.Now,
|
||||
rates: newRateLimiter(lim.RequestsPerSecond, lim.RequestBurst),
|
||||
}
|
||||
for _, opt := range opts {
|
||||
opt(a)
|
||||
}
|
||||
return a
|
||||
}
|
||||
|
||||
func (a *App) now() time.Time {
|
||||
return a.nowFn()
|
||||
}
|
||||
|
||||
func (a *App) protocolLimits() protocol.Limits {
|
||||
return protocol.Limits{
|
||||
MaxBodyBytes: a.lim.MaxBodyBytes,
|
||||
MaxMetaBytes: a.lim.MaxMetaBytes,
|
||||
MaxFrameBytes: a.lim.MaxFrameBytes,
|
||||
}
|
||||
}
|
||||
|
||||
func (a *App) Ack(context.Context, string, *protocol.Ack) (AckResult, error) {
|
||||
return AckResult{}, ErrNotImplemented
|
||||
}
|
||||
|
||||
func (a *App) Recall(context.Context, string, *protocol.Recall) (protocol.RecallData, error) {
|
||||
return protocol.RecallData{}, ErrNotImplemented
|
||||
}
|
||||
|
||||
func (a *App) Status(context.Context, string, *protocol.Status) (any, error) {
|
||||
return nil, ErrNotImplemented
|
||||
}
|
||||
|
||||
func (a *App) ReceiptAck(context.Context, string, *protocol.ReceiptAck) error {
|
||||
return ErrNotImplemented
|
||||
}
|
||||
|
||||
func (a *App) DispatchDue(context.Context, int64, int) (int, error) {
|
||||
return 0, nil
|
||||
}
|
||||
|
||||
func (a *App) PushPending(context.Context, string, port.ConnID) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (a *App) OnPublishDropped(context.Context, string, port.ConnID, []byte) error {
|
||||
return nil
|
||||
}
|
||||
|
||||
func (a *App) CleanupOnce(context.Context, int64) error { return nil }
|
||||
|
||||
func (a *App) RecoverOnStart(context.Context) error { return nil }
|
||||
|
||||
func (a *App) WakePush(string) {}
|
||||
|
||||
var _ Service = (*App)(nil)
|
||||
@@ -0,0 +1,52 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"database/sql"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
)
|
||||
|
||||
// dispatchMinimalTx 是 M1 最小分发:单聊插一条 pending;群按当前成员去掉发送者各插 pending;
|
||||
// 消息改为 dispatched。完整 7.4 规则见 DEVIATIONS「消息 M」。
|
||||
func dispatchMinimalTx(tx *sql.Tx, seq int64, senderID, destKind, destID string, sendAt int64, keep int, nowMs int64) (string, error) {
|
||||
recipients := make([]string, 0, 8)
|
||||
switch destKind {
|
||||
case protocol.TargetEndpoint:
|
||||
recipients = append(recipients, destID)
|
||||
case protocol.TargetGroup:
|
||||
rows, err := tx.Query(`
|
||||
SELECT endpoint_id FROM group_members WHERE group_id = ? AND endpoint_id != ?`, destID, senderID)
|
||||
if err != nil {
|
||||
return "", err
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
for rows.Next() {
|
||||
var id string
|
||||
if err := rows.Scan(&id); err != nil {
|
||||
return "", err
|
||||
}
|
||||
recipients = append(recipients, id)
|
||||
}
|
||||
if err := rows.Err(); err != nil {
|
||||
return "", err
|
||||
}
|
||||
default:
|
||||
return "", errCode(protocol.CodeBadRequest, "invalid dest_kind")
|
||||
}
|
||||
|
||||
for _, ep := range recipients {
|
||||
if _, err := tx.Exec(`
|
||||
INSERT INTO deliveries(seq, endpoint_id, send_at, keep, state, reason, expire_at, pushed_conn, pushed_at, attempts, updated_at)
|
||||
VALUES(?,?,?,?,?,?,NULL,NULL,NULL,0,?)`,
|
||||
seq, ep, sendAt, keep, DeliveryPending, "", nowMs,
|
||||
); err != nil {
|
||||
return "", err
|
||||
}
|
||||
}
|
||||
|
||||
state := StateDispatched
|
||||
if _, err := tx.Exec(`UPDATE messages SET state = ? WHERE seq = ?`, state, seq); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return state, nil
|
||||
}
|
||||
@@ -0,0 +1,7 @@
|
||||
package message
|
||||
|
||||
import "git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
|
||||
func errCode(code, msg string) *protocol.Error {
|
||||
return &protocol.Error{Code: code, Message: msg}
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
// rateLimiter 是每端一个令牌桶:速率 rps、容量 burst。
|
||||
// rps<=0 表示不限速。
|
||||
type rateLimiter struct {
|
||||
rps float64
|
||||
burst float64
|
||||
|
||||
mu sync.Mutex
|
||||
m map[string]*tokenBucket
|
||||
}
|
||||
|
||||
type tokenBucket struct {
|
||||
tokens float64
|
||||
last time.Time
|
||||
}
|
||||
|
||||
func newRateLimiter(rps float64, burst int) *rateLimiter {
|
||||
if burst <= 0 {
|
||||
burst = defaultRequestBurst
|
||||
}
|
||||
return &rateLimiter{
|
||||
rps: rps,
|
||||
burst: float64(burst),
|
||||
m: make(map[string]*tokenBucket),
|
||||
}
|
||||
}
|
||||
|
||||
// allow 消耗 1 个令牌;允许则 true。
|
||||
func (r *rateLimiter) allow(endpointID string, now time.Time) bool {
|
||||
if r == nil || r.rps <= 0 {
|
||||
return true
|
||||
}
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
b := r.m[endpointID]
|
||||
if b == nil {
|
||||
b = &tokenBucket{tokens: r.burst, last: now}
|
||||
r.m[endpointID] = b
|
||||
}
|
||||
elapsed := now.Sub(b.last).Seconds()
|
||||
if elapsed > 0 {
|
||||
b.tokens += elapsed * r.rps
|
||||
if b.tokens > r.burst {
|
||||
b.tokens = r.burst
|
||||
}
|
||||
b.last = now
|
||||
}
|
||||
if b.tokens < 1 {
|
||||
return false
|
||||
}
|
||||
b.tokens--
|
||||
return true
|
||||
}
|
||||
@@ -1,25 +0,0 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"context"
|
||||
"errors"
|
||||
"testing"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
)
|
||||
|
||||
func TestStubSubmitNotImplemented(t *testing.T) {
|
||||
s := NewStub()
|
||||
_, err := s.Submit(context.Background(), "a", port.ConnInfo{}, &protocol.Send{})
|
||||
if !errors.Is(err, ErrNotImplemented) {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStubRecoverNoop(t *testing.T) {
|
||||
s := NewStub()
|
||||
if err := s.RecoverOnStart(context.Background()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,560 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"encoding/hex"
|
||||
"errors"
|
||||
"fmt"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
)
|
||||
|
||||
// Submit 处理发送提交(DEVELOPMENT 7.3):防重 → 校验/配额/授权 → 写入 → 到点则最小分发。
|
||||
func (a *App) Submit(ctx context.Context, senderID string, conn port.ConnInfo, req *protocol.Send) (SubmitResult, error) {
|
||||
if req == nil {
|
||||
return SubmitResult{}, errCode(protocol.CodeBadRequest, "nil send")
|
||||
}
|
||||
if senderID == "" || !protocol.ValidEndpointID(senderID) {
|
||||
return SubmitResult{}, errCode(protocol.CodeBadRequest, "invalid sender")
|
||||
}
|
||||
now := a.now()
|
||||
if !a.rates.allow(senderID, now) {
|
||||
return SubmitResult{}, errCode(protocol.CodeRateLimited, "request rate exceeded")
|
||||
}
|
||||
|
||||
if err := req.Validate(a.protocolLimits()); err != nil {
|
||||
return SubmitResult{}, err
|
||||
}
|
||||
fpHex, err := protocol.RequestFingerprint(req)
|
||||
if err != nil {
|
||||
return SubmitResult{}, err
|
||||
}
|
||||
fp, err := hex.DecodeString(fpHex)
|
||||
if err != nil || len(fp) != 32 {
|
||||
return SubmitResult{}, fmt.Errorf("message: fingerprint decode: %w", err)
|
||||
}
|
||||
|
||||
body, err := protocol.DecodeBody(req.Body)
|
||||
if err != nil {
|
||||
return SubmitResult{}, err
|
||||
}
|
||||
metaJSON, err := protocol.MetaCanonicalJSON(req.Meta)
|
||||
if err != nil {
|
||||
return SubmitResult{}, errCode(protocol.CodeBadRequest, "invalid meta")
|
||||
}
|
||||
contentType := protocol.EffectiveContentType(req.Body)
|
||||
keep := protocol.EffectiveOfflineKeep(req)
|
||||
ttl := protocol.EffectiveOfflineTTL(req)
|
||||
receipt := protocol.EffectiveReceipt(req)
|
||||
if keep && a.lim.MaxTTLSeconds > 0 && ttl > a.lim.MaxTTLSeconds {
|
||||
return SubmitResult{}, errCode(protocol.CodeBadRequest, "ttl_seconds exceeds max_ttl_seconds")
|
||||
}
|
||||
|
||||
nowMs := now.UnixMilli()
|
||||
|
||||
// 防重命中可在读连接快速返回;写路径仍会再查一次以防竞态。
|
||||
if res, hit, lookupErr := a.lookupIdempotent(ctx, senderID, req.ID, fp); lookupErr != nil {
|
||||
return SubmitResult{}, lookupErr
|
||||
} else if hit {
|
||||
return res, nil
|
||||
}
|
||||
|
||||
sender, err := a.loadEndpoint(ctx, senderID)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return SubmitResult{}, errCode(protocol.CodeInvalidTarget, "sender not found")
|
||||
}
|
||||
return SubmitResult{}, err
|
||||
}
|
||||
|
||||
sendAt, err := a.computeSendAt(req, sender.DefaultDelayMs, nowMs)
|
||||
if err != nil {
|
||||
return SubmitResult{}, err
|
||||
}
|
||||
|
||||
var (
|
||||
needPassword bool
|
||||
talkPHC string
|
||||
targetEp *endpointRow
|
||||
)
|
||||
|
||||
switch req.To.Kind {
|
||||
case protocol.TargetEndpoint:
|
||||
targetEp, err = a.loadEndpoint(ctx, req.To.ID)
|
||||
if err != nil {
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return SubmitResult{}, errCode(protocol.CodeInvalidTarget, "target not found")
|
||||
}
|
||||
return SubmitResult{}, err
|
||||
}
|
||||
if targetEp.Enabled == 0 {
|
||||
return SubmitResult{}, errCode(protocol.CodeEndpointDisabled, "target disabled")
|
||||
}
|
||||
if senderID != req.To.ID {
|
||||
needPassword, talkPHC, _, err = a.dmAuthNeeded(ctx, senderID, targetEp)
|
||||
if err != nil {
|
||||
return SubmitResult{}, err
|
||||
}
|
||||
}
|
||||
case protocol.TargetGroup:
|
||||
exists, member, gErr := a.groupMembership(ctx, req.To.ID, senderID)
|
||||
if gErr != nil {
|
||||
return SubmitResult{}, gErr
|
||||
}
|
||||
if !exists {
|
||||
return SubmitResult{}, errCode(protocol.CodeInvalidTarget, "group not found")
|
||||
}
|
||||
if !member {
|
||||
return SubmitResult{}, errCode(protocol.CodeNotMember, "not a group member")
|
||||
}
|
||||
default:
|
||||
return SubmitResult{}, errCode(protocol.CodeBadRequest, "invalid to.kind")
|
||||
}
|
||||
passwordVerified := false
|
||||
if needPassword {
|
||||
if locked, _ := a.talkLocked(senderID, req.To.ID, conn.RemoteIP); locked {
|
||||
return SubmitResult{}, errCode(protocol.CodeRateLimited, "talk password locked")
|
||||
}
|
||||
if req.TalkPassword == "" {
|
||||
return SubmitResult{}, errCode(protocol.CodeTalkPasswordRequired, "talk password required")
|
||||
}
|
||||
if a.hash == nil {
|
||||
return SubmitResult{}, fmt.Errorf("message: hash pool required")
|
||||
}
|
||||
ok, vErr := a.hash.Verify(ctx, auth.PasswordTalk, req.TalkPassword, talkPHC)
|
||||
if vErr != nil {
|
||||
return SubmitResult{}, vErr
|
||||
}
|
||||
if !ok {
|
||||
a.talkFail(senderID, req.To.ID, conn.RemoteIP)
|
||||
return SubmitResult{}, errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid")
|
||||
}
|
||||
passwordVerified = true
|
||||
a.talkClear(senderID, req.To.ID)
|
||||
}
|
||||
|
||||
keepInt := 0
|
||||
if keep {
|
||||
keepInt = 1
|
||||
}
|
||||
receiptInt := 0
|
||||
if receipt {
|
||||
receiptInt = 1
|
||||
}
|
||||
|
||||
var result SubmitResult
|
||||
err = a.db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
if res, hit, e := lookupIdempotentTx(tx, senderID, req.ID, fp); e != nil {
|
||||
return e
|
||||
} else if hit {
|
||||
result = res
|
||||
return nil
|
||||
}
|
||||
|
||||
if e := checkQuotaTx(tx, senderID, a.lim.MaxPendingPerSender); e != nil {
|
||||
return e
|
||||
}
|
||||
|
||||
// 写事务内再确认目标与授权(防并发停用/退群)。
|
||||
switch req.To.Kind {
|
||||
case protocol.TargetEndpoint:
|
||||
ep, e := loadEndpointTx(tx, req.To.ID)
|
||||
if e != nil {
|
||||
if errors.Is(e, sql.ErrNoRows) {
|
||||
return errCode(protocol.CodeInvalidTarget, "target not found")
|
||||
}
|
||||
return e
|
||||
}
|
||||
if ep.Enabled == 0 {
|
||||
return errCode(protocol.CodeEndpointDisabled, "target disabled")
|
||||
}
|
||||
if senderID != req.To.ID {
|
||||
needed, phc, ver, ae := dmAuthNeededTx(tx, senderID, ep)
|
||||
if ae != nil {
|
||||
return ae
|
||||
}
|
||||
if needed {
|
||||
if !passwordVerified {
|
||||
if req.TalkPassword == "" {
|
||||
return errCode(protocol.CodeTalkPasswordRequired, "talk password required")
|
||||
}
|
||||
return errCode(protocol.CodeTalkPasswordInvalid, "talk password invalid")
|
||||
}
|
||||
// 密码版本在校验后变化则拒绝,避免写过期授权。
|
||||
if ep.TalkHash == nil || *ep.TalkHash != phc || ep.TalkVersion != ver {
|
||||
return errCode(protocol.CodeTalkPasswordInvalid, "talk password changed")
|
||||
}
|
||||
if ge := upsertGrantTx(tx, senderID, req.To.ID, ver, GrantKindPassword, nowMs); ge != nil {
|
||||
return ge
|
||||
}
|
||||
} else if passwordVerified {
|
||||
// 已有授权或未设防:带对密码时仍可刷新授权(文档:带对了则写入或更新)。
|
||||
if ep.TalkHash != nil && *ep.TalkHash != "" {
|
||||
if ge := upsertGrantTx(tx, senderID, req.To.ID, ep.TalkVersion, GrantKindPassword, nowMs); ge != nil {
|
||||
return ge
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
// 发送方设了对话密码且发给别人的单聊:给对方写回复授权。
|
||||
snd, se := loadEndpointTx(tx, senderID)
|
||||
if se != nil {
|
||||
return se
|
||||
}
|
||||
if senderID != req.To.ID && snd.TalkHash != nil && *snd.TalkHash != "" {
|
||||
if ge := upsertGrantTx(tx, req.To.ID, senderID, snd.TalkVersion, GrantKindReply, nowMs); ge != nil {
|
||||
return ge
|
||||
}
|
||||
}
|
||||
case protocol.TargetGroup:
|
||||
exists, member, e := groupMembershipTx(tx, req.To.ID, senderID)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
if !exists {
|
||||
return errCode(protocol.CodeInvalidTarget, "group not found")
|
||||
}
|
||||
if !member {
|
||||
return errCode(protocol.CodeNotMember, "not a group member")
|
||||
}
|
||||
}
|
||||
|
||||
state := StateScheduled
|
||||
res, e := tx.Exec(`
|
||||
INSERT INTO messages(
|
||||
id, sender_id, dest_kind, dest_id, meta, content_type, body_enc,
|
||||
send_at, keep, ttl_seconds, receipt, state, reason, created_at
|
||||
) VALUES(?,?,?,?,?,?,?,?,?,?,?,?, '', ?)`,
|
||||
req.ID, senderID, req.To.Kind, req.To.ID, string(metaJSON), contentType, req.Body.Enc,
|
||||
sendAt, keepInt, ttl, receiptInt, state, nowMs,
|
||||
)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
seq, e := res.LastInsertId()
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
if _, e = tx.Exec(`INSERT INTO message_bodies(seq, body) VALUES(?, ?)`, seq, body); e != nil {
|
||||
return e
|
||||
}
|
||||
if _, e = tx.Exec(
|
||||
`INSERT INTO send_keys(sender_id, msg_id, request_sha256, created_at) VALUES(?,?,?,?)`,
|
||||
senderID, req.ID, fp, nowMs,
|
||||
); e != nil {
|
||||
return e
|
||||
}
|
||||
|
||||
finalState := state
|
||||
if sendAt <= nowMs {
|
||||
finalState, e = dispatchMinimalTx(tx, seq, senderID, req.To.Kind, req.To.ID, sendAt, keepInt, nowMs)
|
||||
if e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
result = SubmitResult{ID: req.ID, SendAtMs: sendAt, State: finalState}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
return SubmitResult{}, err
|
||||
}
|
||||
return result, nil
|
||||
}
|
||||
|
||||
func (a *App) computeSendAt(req *protocol.Send, defaultDelayMs, nowMs int64) (int64, error) {
|
||||
if req.SendAtMs != nil && req.DelayMs != nil {
|
||||
return 0, errCode(protocol.CodeBadRequest, "send_at_ms and delay_ms are mutually exclusive")
|
||||
}
|
||||
var sendAt int64
|
||||
switch {
|
||||
case req.SendAtMs != nil:
|
||||
sendAt = *req.SendAtMs
|
||||
case req.DelayMs != nil:
|
||||
if *req.DelayMs < 0 {
|
||||
return 0, errCode(protocol.CodeBadRequest, "delay_ms negative")
|
||||
}
|
||||
sendAt = nowMs + *req.DelayMs
|
||||
default:
|
||||
if defaultDelayMs < 0 {
|
||||
defaultDelayMs = 0
|
||||
}
|
||||
sendAt = nowMs + defaultDelayMs
|
||||
}
|
||||
if a.lim.MaxScheduleSeconds > 0 {
|
||||
maxAt := nowMs + a.lim.MaxScheduleSeconds*1000
|
||||
if sendAt > maxAt {
|
||||
return 0, errCode(protocol.CodeBadRequest, "send time exceeds max_schedule_seconds")
|
||||
}
|
||||
}
|
||||
return sendAt, nil
|
||||
}
|
||||
|
||||
type endpointRow struct {
|
||||
ID string
|
||||
DefaultDelayMs int64
|
||||
TalkHash *string
|
||||
TalkVersion int64
|
||||
Enabled int
|
||||
}
|
||||
|
||||
func (a *App) loadEndpoint(ctx context.Context, id string) (*endpointRow, error) {
|
||||
row := a.db.Read.QueryRowContext(ctx, `
|
||||
SELECT id, default_delay_ms, talk_hash, talk_version, enabled
|
||||
FROM endpoints WHERE id = ?`, id)
|
||||
return scanEndpoint(row)
|
||||
}
|
||||
|
||||
func loadEndpointTx(tx *sql.Tx, id string) (*endpointRow, error) {
|
||||
row := tx.QueryRow(`
|
||||
SELECT id, default_delay_ms, talk_hash, talk_version, enabled
|
||||
FROM endpoints WHERE id = ?`, id)
|
||||
return scanEndpoint(row)
|
||||
}
|
||||
|
||||
func scanEndpoint(row *sql.Row) (*endpointRow, error) {
|
||||
var ep endpointRow
|
||||
var talk sql.NullString
|
||||
if err := row.Scan(&ep.ID, &ep.DefaultDelayMs, &talk, &ep.TalkVersion, &ep.Enabled); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
if talk.Valid {
|
||||
s := talk.String
|
||||
ep.TalkHash = &s
|
||||
}
|
||||
return &ep, nil
|
||||
}
|
||||
|
||||
func (a *App) groupMembership(ctx context.Context, groupID, endpointID string) (exists, member bool, err error) {
|
||||
var one int
|
||||
err = a.db.Read.QueryRowContext(ctx, `SELECT 1 FROM groups WHERE id = ?`, groupID).Scan(&one)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return false, false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return false, false, err
|
||||
}
|
||||
err = a.db.Read.QueryRowContext(ctx,
|
||||
`SELECT 1 FROM group_members WHERE group_id = ? AND endpoint_id = ?`, groupID, endpointID,
|
||||
).Scan(&one)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return true, false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return true, false, err
|
||||
}
|
||||
return true, true, nil
|
||||
}
|
||||
|
||||
func groupMembershipTx(tx *sql.Tx, groupID, endpointID string) (exists, member bool, err error) {
|
||||
var one int
|
||||
err = tx.QueryRow(`SELECT 1 FROM groups WHERE id = ?`, groupID).Scan(&one)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return false, false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return false, false, err
|
||||
}
|
||||
err = tx.QueryRow(
|
||||
`SELECT 1 FROM group_members WHERE group_id = ? AND endpoint_id = ?`, groupID, endpointID,
|
||||
).Scan(&one)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return true, false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return true, false, err
|
||||
}
|
||||
return true, true, nil
|
||||
}
|
||||
|
||||
// dmAuthNeeded 返回是否需要对话密码,以及对方当前 talk_hash / version。
|
||||
func (a *App) dmAuthNeeded(ctx context.Context, senderID string, target *endpointRow) (needed bool, phc string, version int64, err error) {
|
||||
if target.TalkHash == nil || *target.TalkHash == "" {
|
||||
return false, "", target.TalkVersion, nil
|
||||
}
|
||||
ok, err := hasValidGrant(ctx, a.db.Read, senderID, target.ID, target.TalkVersion)
|
||||
if err != nil {
|
||||
return false, "", 0, err
|
||||
}
|
||||
if ok {
|
||||
return false, *target.TalkHash, target.TalkVersion, nil
|
||||
}
|
||||
return true, *target.TalkHash, target.TalkVersion, nil
|
||||
}
|
||||
|
||||
func dmAuthNeededTx(tx *sql.Tx, senderID string, target *endpointRow) (needed bool, phc string, version int64, err error) {
|
||||
if target.TalkHash == nil || *target.TalkHash == "" {
|
||||
return false, "", target.TalkVersion, nil
|
||||
}
|
||||
ok, err := hasValidGrantTx(tx, senderID, target.ID, target.TalkVersion)
|
||||
if err != nil {
|
||||
return false, "", 0, err
|
||||
}
|
||||
if ok {
|
||||
return false, *target.TalkHash, target.TalkVersion, nil
|
||||
}
|
||||
return true, *target.TalkHash, target.TalkVersion, nil
|
||||
}
|
||||
|
||||
func hasValidGrant(ctx context.Context, db *sql.DB, senderID, targetID string, talkVersion int64) (bool, error) {
|
||||
var n int
|
||||
err := db.QueryRowContext(ctx, `
|
||||
SELECT 1 FROM talk_grants
|
||||
WHERE sender_id = ? AND target_id = ? AND target_talk_version = ?
|
||||
LIMIT 1`, senderID, targetID, talkVersion).Scan(&n)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return false, nil
|
||||
}
|
||||
return err == nil, err
|
||||
}
|
||||
|
||||
func hasValidGrantTx(tx *sql.Tx, senderID, targetID string, talkVersion int64) (bool, error) {
|
||||
var n int
|
||||
err := tx.QueryRow(`
|
||||
SELECT 1 FROM talk_grants
|
||||
WHERE sender_id = ? AND target_id = ? AND target_talk_version = ?
|
||||
LIMIT 1`, senderID, targetID, talkVersion).Scan(&n)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return false, nil
|
||||
}
|
||||
return err == nil, err
|
||||
}
|
||||
|
||||
func upsertGrantTx(tx *sql.Tx, senderID, targetID string, talkVersion int64, kind string, nowMs int64) error {
|
||||
_, err := tx.Exec(`
|
||||
INSERT INTO talk_grants(sender_id, target_id, target_talk_version, kind, created_at)
|
||||
VALUES(?,?,?,?,?)
|
||||
ON CONFLICT(sender_id, target_id) DO UPDATE SET
|
||||
target_talk_version = excluded.target_talk_version,
|
||||
kind = excluded.kind,
|
||||
created_at = excluded.created_at`,
|
||||
senderID, targetID, talkVersion, kind, nowMs,
|
||||
)
|
||||
return err
|
||||
}
|
||||
|
||||
func checkQuotaTx(tx *sql.Tx, senderID string, maxPending int) error {
|
||||
if maxPending <= 0 {
|
||||
return nil
|
||||
}
|
||||
var n int
|
||||
err := tx.QueryRow(`
|
||||
SELECT COUNT(*) FROM messages
|
||||
WHERE sender_id = ? AND state IN ('scheduled', 'dispatched')`, senderID).Scan(&n)
|
||||
if err != nil {
|
||||
return err
|
||||
}
|
||||
if n >= maxPending {
|
||||
return errCode(protocol.CodeQuotaExceeded, "max_pending_per_sender exceeded")
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (a *App) lookupIdempotent(ctx context.Context, senderID, msgID string, fp []byte) (SubmitResult, bool, error) {
|
||||
var stored []byte
|
||||
err := a.db.Read.QueryRowContext(ctx, `
|
||||
SELECT request_sha256 FROM send_keys WHERE sender_id = ? AND msg_id = ?`,
|
||||
senderID, msgID,
|
||||
).Scan(&stored)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return SubmitResult{}, false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return SubmitResult{}, false, err
|
||||
}
|
||||
if !bytesEqual(stored, fp) {
|
||||
return SubmitResult{}, false, errCode(protocol.CodeConflict, "message id conflict")
|
||||
}
|
||||
res, err := loadSubmitResult(ctx, a.db.Read, senderID, msgID)
|
||||
if err != nil {
|
||||
return SubmitResult{}, false, err
|
||||
}
|
||||
return res, true, nil
|
||||
}
|
||||
|
||||
func lookupIdempotentTx(tx *sql.Tx, senderID, msgID string, fp []byte) (SubmitResult, bool, error) {
|
||||
var stored []byte
|
||||
err := tx.QueryRow(`
|
||||
SELECT request_sha256 FROM send_keys WHERE sender_id = ? AND msg_id = ?`,
|
||||
senderID, msgID,
|
||||
).Scan(&stored)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return SubmitResult{}, false, nil
|
||||
}
|
||||
if err != nil {
|
||||
return SubmitResult{}, false, err
|
||||
}
|
||||
if !bytesEqual(stored, fp) {
|
||||
return SubmitResult{}, false, errCode(protocol.CodeConflict, "message id conflict")
|
||||
}
|
||||
res, err := loadSubmitResultTx(tx, senderID, msgID)
|
||||
if err != nil {
|
||||
return SubmitResult{}, false, err
|
||||
}
|
||||
return res, true, nil
|
||||
}
|
||||
|
||||
func loadSubmitResult(ctx context.Context, db *sql.DB, senderID, msgID string) (SubmitResult, error) {
|
||||
var res SubmitResult
|
||||
err := db.QueryRowContext(ctx, `
|
||||
SELECT id, send_at, state FROM messages WHERE sender_id = ? AND id = ?`,
|
||||
senderID, msgID,
|
||||
).Scan(&res.ID, &res.SendAtMs, &res.State)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return SubmitResult{}, errCode(protocol.CodeNotFound, "idempotent key without message")
|
||||
}
|
||||
return res, err
|
||||
}
|
||||
|
||||
func loadSubmitResultTx(tx *sql.Tx, senderID, msgID string) (SubmitResult, error) {
|
||||
var res SubmitResult
|
||||
err := tx.QueryRow(`
|
||||
SELECT id, send_at, state FROM messages WHERE sender_id = ? AND id = ?`,
|
||||
senderID, msgID,
|
||||
).Scan(&res.ID, &res.SendAtMs, &res.State)
|
||||
if errors.Is(err, sql.ErrNoRows) {
|
||||
return SubmitResult{}, errCode(protocol.CodeNotFound, "idempotent key without message")
|
||||
}
|
||||
return res, err
|
||||
}
|
||||
|
||||
func bytesEqual(a, b []byte) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
var v byte
|
||||
for i := range a {
|
||||
v |= a[i] ^ b[i]
|
||||
}
|
||||
return v == 0
|
||||
}
|
||||
|
||||
func (a *App) talkLocked(senderID, targetID, ip string) (bool, error) {
|
||||
if a.locks == nil {
|
||||
return false, nil
|
||||
}
|
||||
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: ip}); locked {
|
||||
return true, nil
|
||||
}
|
||||
if locked, _ := a.locks.Check(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID}); locked {
|
||||
return true, nil
|
||||
}
|
||||
return false, nil
|
||||
}
|
||||
|
||||
func (a *App) talkFail(senderID, targetID, ip string) {
|
||||
if a.locks == nil {
|
||||
return
|
||||
}
|
||||
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID, IP: ip})
|
||||
a.locks.Fail(auth.LockKey{Kind: auth.LockTalkTarget, EndpointID: targetID})
|
||||
}
|
||||
|
||||
func (a *App) talkClear(senderID, targetID string) {
|
||||
if a.locks == nil {
|
||||
return
|
||||
}
|
||||
a.locks.Clear(auth.LockKey{Kind: auth.LockTalkPair, EndpointID: senderID, PeerID: targetID})
|
||||
}
|
||||
@@ -0,0 +1,382 @@
|
||||
package message
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"errors"
|
||||
"path/filepath"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/app/port"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/auth"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/config"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/protocol"
|
||||
"git.asio.asia/nixevol/NixMsg/internal/store"
|
||||
)
|
||||
|
||||
func TestStubSubmitNotImplemented(t *testing.T) {
|
||||
s := NewStub()
|
||||
_, err := s.Submit(context.Background(), "a", port.ConnInfo{}, &protocol.Send{})
|
||||
if !errors.Is(err, ErrNotImplemented) {
|
||||
t.Fatalf("got %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStubRecoverNoop(t *testing.T) {
|
||||
s := NewStub()
|
||||
if err := s.RecoverOnStart(context.Background()); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func openTestApp(t *testing.T, lim Limits) (*App, *store.DB) {
|
||||
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() })
|
||||
fixed := time.UnixMilli(1_700_000_000_000)
|
||||
app := New(db, lim, auth.NewStubHashPool(),
|
||||
WithNow(func() time.Time { return fixed }),
|
||||
WithLocks(auth.NewStubLoginLocks()),
|
||||
)
|
||||
return app, db
|
||||
}
|
||||
|
||||
func defaultTestLimits() Limits {
|
||||
cfg := config.Default().Limits
|
||||
lim := LimitsFromConfig(cfg)
|
||||
lim.RequestsPerSecond = 0 // 测试默认不限速
|
||||
return lim
|
||||
}
|
||||
|
||||
func insertEndpoint(t *testing.T, db *store.DB, id string, talkPassword string, enabled int, defaultDelayMs int64) {
|
||||
t.Helper()
|
||||
ctx := context.Background()
|
||||
var talk any
|
||||
var talkVer int64
|
||||
if talkPassword != "" {
|
||||
talk = "stub$" + talkPassword
|
||||
talkVer = 1
|
||||
}
|
||||
err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`
|
||||
INSERT INTO endpoints(id, name, login_hash, talk_hash, talk_version, default_delay_ms, enabled, created_at)
|
||||
VALUES(?,?,?,?,?,?,?,?)`,
|
||||
id, id, "stub$login", talk, talkVer, defaultDelayMs, enabled, 1_700_000_000_000)
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
}
|
||||
|
||||
func baseSend(id, to string) *protocol.Send {
|
||||
return &protocol.Send{
|
||||
V: protocol.Version,
|
||||
Type: protocol.TypeSend,
|
||||
RID: "r1",
|
||||
ID: id,
|
||||
To: protocol.Target{Kind: protocol.TargetEndpoint, ID: to},
|
||||
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hello"},
|
||||
}
|
||||
}
|
||||
|
||||
func protoCode(err error) string {
|
||||
var pe *protocol.Error
|
||||
if errors.As(err, &pe) {
|
||||
return pe.Code
|
||||
}
|
||||
return ""
|
||||
}
|
||||
|
||||
func TestSubmitTable(t *testing.T) {
|
||||
t.Parallel()
|
||||
|
||||
t.Run("idempotent_hit", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
ctx := context.Background()
|
||||
req := baseSend("msg-1", "bob")
|
||||
first, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if first.State != StateDispatched {
|
||||
t.Fatalf("state=%s", first.State)
|
||||
}
|
||||
second, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if second != first {
|
||||
t.Fatalf("want %+v got %+v", first, second)
|
||||
}
|
||||
var n int
|
||||
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM messages WHERE sender_id=? AND id=?`, "alice", "msg-1").Scan(&n); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 1 {
|
||||
t.Fatalf("messages=%d", n)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("conflict", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
ctx := context.Background()
|
||||
req := baseSend("msg-2", "bob")
|
||||
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, req); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
other := baseSend("msg-2", "bob")
|
||||
other.Body.Data = "other"
|
||||
_, err := app.Submit(ctx, "alice", port.ConnInfo{}, other)
|
||||
if protoCode(err) != protocol.CodeConflict {
|
||||
t.Fatalf("want conflict got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("quota_exceeded", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
lim.MaxPendingPerSender = 1
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
ctx := context.Background()
|
||||
delay := int64(60_000)
|
||||
req1 := baseSend("q1", "bob")
|
||||
req1.DelayMs = &delay
|
||||
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, req1); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req2 := baseSend("q2", "bob")
|
||||
req2.DelayMs = &delay
|
||||
_, err := app.Submit(ctx, "alice", port.ConnInfo{}, req2)
|
||||
if protoCode(err) != protocol.CodeQuotaExceeded {
|
||||
t.Fatalf("want quota_exceeded got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("auth_required_and_grant", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "alice-secret", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "secret", 1, 0)
|
||||
ctx := context.Background()
|
||||
|
||||
_, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("a1", "bob"))
|
||||
if protoCode(err) != protocol.CodeTalkPasswordRequired {
|
||||
t.Fatalf("want talk_password_required got %v", err)
|
||||
}
|
||||
|
||||
bad := baseSend("a2", "bob")
|
||||
bad.TalkPassword = "wrong"
|
||||
_, err = app.Submit(ctx, "alice", port.ConnInfo{}, bad)
|
||||
if protoCode(err) != protocol.CodeTalkPasswordInvalid {
|
||||
t.Fatalf("want talk_password_invalid got %v", err)
|
||||
}
|
||||
|
||||
okReq := baseSend("a3", "bob")
|
||||
okReq.TalkPassword = "secret"
|
||||
res, err := app.Submit(ctx, "alice", port.ConnInfo{}, okReq)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if res.State != StateDispatched {
|
||||
t.Fatalf("state=%s", res.State)
|
||||
}
|
||||
// 已有授权后不带密码也可发
|
||||
if _, submitErr := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("a4", "bob")); submitErr != nil {
|
||||
t.Fatal(submitErr)
|
||||
}
|
||||
// 回复授权:bob→alice(因 alice 设了对话密码)
|
||||
var kind string
|
||||
err = db.Read.QueryRow(`
|
||||
SELECT kind FROM talk_grants WHERE sender_id=? AND target_id=?`, "bob", "alice").Scan(&kind)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if kind != GrantKindReply {
|
||||
t.Fatalf("reply grant kind=%s", kind)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("self_skip_talk_password", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "secret", 1, 0)
|
||||
ctx := context.Background()
|
||||
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("self1", "alice")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("delay_and_send_at_mutex", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
ctx := context.Background()
|
||||
delay := int64(1000)
|
||||
sendAt := int64(1_700_000_001_000)
|
||||
req := baseSend("m-mutex", "bob")
|
||||
req.DelayMs = &delay
|
||||
req.SendAtMs = &sendAt
|
||||
_, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
|
||||
if protoCode(err) != protocol.CodeBadRequest {
|
||||
t.Fatalf("want bad_request got %v", err)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("idempotent_before_disabled_check", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
ctx := context.Background()
|
||||
req := baseSend("pre-disable", "bob")
|
||||
first, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
err = db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
_, e := tx.Exec(`UPDATE endpoints SET enabled = 0 WHERE id = ?`, "bob")
|
||||
return e
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
// 新消息应失败
|
||||
_, err = app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("after-disable", "bob"))
|
||||
if protoCode(err) != protocol.CodeEndpointDisabled {
|
||||
t.Fatalf("want endpoint_disabled got %v", err)
|
||||
}
|
||||
// 原请求重试仍返回原结果
|
||||
second, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if second != first {
|
||||
t.Fatalf("want %+v got %+v", first, second)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("scheduled_not_dispatched", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
ctx := context.Background()
|
||||
delay := int64(10_000)
|
||||
req := baseSend("sched-1", "bob")
|
||||
req.DelayMs = &delay
|
||||
res, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if res.State != StateScheduled {
|
||||
t.Fatalf("state=%s", res.State)
|
||||
}
|
||||
var n int
|
||||
if err := db.Read.QueryRow(`SELECT COUNT(*) FROM deliveries`).Scan(&n); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if n != 0 {
|
||||
t.Fatalf("deliveries=%d", n)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("group_dispatch_excludes_sender", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
insertEndpoint(t, db, "carol", "", 1, 0)
|
||||
ctx := context.Background()
|
||||
err := db.Queue.Do(ctx, func(tx *sql.Tx) error {
|
||||
if _, e := tx.Exec(`INSERT INTO groups(id, name, owner_id, created_at) VALUES(?,?,?,?)`,
|
||||
"g1", "g", "alice", 1_700_000_000_000); e != nil {
|
||||
return e
|
||||
}
|
||||
for _, m := range []string{"alice", "bob", "carol"} {
|
||||
if _, e := tx.Exec(`INSERT INTO group_members(group_id, endpoint_id, joined_at) VALUES(?,?,?)`,
|
||||
"g1", m, 1_700_000_000_000); e != nil {
|
||||
return e
|
||||
}
|
||||
}
|
||||
return nil
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
req := &protocol.Send{
|
||||
V: protocol.Version,
|
||||
Type: protocol.TypeSend,
|
||||
RID: "r1",
|
||||
ID: "gmsg-1",
|
||||
To: protocol.Target{Kind: protocol.TargetGroup, ID: "g1"},
|
||||
Body: protocol.Body{Enc: protocol.EncUTF8, Data: "hi"},
|
||||
}
|
||||
res, err := app.Submit(ctx, "alice", port.ConnInfo{}, req)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if res.State != StateDispatched {
|
||||
t.Fatalf("state=%s", res.State)
|
||||
}
|
||||
rows, err := db.Read.Query(`SELECT endpoint_id FROM deliveries ORDER BY endpoint_id`)
|
||||
if err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
defer func() { _ = rows.Close() }()
|
||||
var got []string
|
||||
for rows.Next() {
|
||||
var id string
|
||||
if err := rows.Scan(&id); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
got = append(got, id)
|
||||
}
|
||||
if len(got) != 2 || got[0] != "bob" || got[1] != "carol" {
|
||||
t.Fatalf("recipients=%v", got)
|
||||
}
|
||||
})
|
||||
|
||||
t.Run("rate_limited", func(t *testing.T) {
|
||||
t.Parallel()
|
||||
lim := defaultTestLimits()
|
||||
lim.RequestsPerSecond = 50
|
||||
lim.RequestBurst = 2
|
||||
app, db := openTestApp(t, lim)
|
||||
insertEndpoint(t, db, "alice", "", 1, 0)
|
||||
insertEndpoint(t, db, "bob", "", 1, 0)
|
||||
ctx := context.Background()
|
||||
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r1", "bob")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
if _, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r2", "bob")); err != nil {
|
||||
t.Fatal(err)
|
||||
}
|
||||
_, err := app.Submit(ctx, "alice", port.ConnInfo{}, baseSend("r3", "bob"))
|
||||
if protoCode(err) != protocol.CodeRateLimited {
|
||||
t.Fatalf("want rate_limited got %v", err)
|
||||
}
|
||||
})
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
package httpx
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/http"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ClientIP 按 DEVELOPMENT 4.5:仅当对端在 trusted 网段内时采信
|
||||
// X-Forwarded-For(从右往左第一个不在 trusted 内的地址)。
|
||||
func ClientIP(r *http.Request, trusted []*net.IPNet) string {
|
||||
host, _, err := net.SplitHostPort(r.RemoteAddr)
|
||||
if err != nil {
|
||||
host = r.RemoteAddr
|
||||
}
|
||||
ip := net.ParseIP(host)
|
||||
if ip == nil {
|
||||
return host
|
||||
}
|
||||
if !ipInNets(ip, trusted) {
|
||||
return ip.String()
|
||||
}
|
||||
xff := r.Header.Get("X-Forwarded-For")
|
||||
if xff == "" {
|
||||
return ip.String()
|
||||
}
|
||||
parts := strings.Split(xff, ",")
|
||||
for i := len(parts) - 1; i >= 0; i-- {
|
||||
cand := strings.TrimSpace(parts[i])
|
||||
parsed := net.ParseIP(cand)
|
||||
if parsed == nil {
|
||||
continue
|
||||
}
|
||||
if !ipInNets(parsed, trusted) {
|
||||
return parsed.String()
|
||||
}
|
||||
}
|
||||
return ip.String()
|
||||
}
|
||||
|
||||
// IsHTTPS 判定请求是否视为 HTTPS(直连 TLS 或受信任代理的 X-Forwarded-Proto)。
|
||||
func IsHTTPS(r *http.Request, trusted []*net.IPNet) bool {
|
||||
if r.TLS != nil {
|
||||
return true
|
||||
}
|
||||
host, _, err := net.SplitHostPort(r.RemoteAddr)
|
||||
if err != nil {
|
||||
host = r.RemoteAddr
|
||||
}
|
||||
ip := net.ParseIP(host)
|
||||
if ip == nil || !ipInNets(ip, trusted) {
|
||||
return false
|
||||
}
|
||||
proto := strings.ToLower(strings.TrimSpace(r.Header.Get("X-Forwarded-Proto")))
|
||||
return proto == "https"
|
||||
}
|
||||
|
||||
// ParseCIDRs 解析 CIDR 列表;非法项跳过。
|
||||
func ParseCIDRs(cidrs []string) []*net.IPNet {
|
||||
var out []*net.IPNet
|
||||
for _, c := range cidrs {
|
||||
c = strings.TrimSpace(c)
|
||||
if c == "" {
|
||||
continue
|
||||
}
|
||||
_, n, err := net.ParseCIDR(c)
|
||||
if err != nil {
|
||||
// 允许单 IP 写成无掩码
|
||||
if ip := net.ParseIP(c); ip != nil {
|
||||
if ip.To4() != nil {
|
||||
_, n, err = net.ParseCIDR(ip.String() + "/32")
|
||||
} else {
|
||||
_, n, err = net.ParseCIDR(ip.String() + "/128")
|
||||
}
|
||||
}
|
||||
}
|
||||
if err == nil && n != nil {
|
||||
out = append(out, n)
|
||||
}
|
||||
}
|
||||
return out
|
||||
}
|
||||
|
||||
func ipInNets(ip net.IP, nets []*net.IPNet) bool {
|
||||
for _, n := range nets {
|
||||
if n.Contains(ip) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
package httpx_test
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/http/httptest"
|
||||
"testing"
|
||||
|
||||
"git.asio.asia/nixevol/NixMsg/internal/httpx"
|
||||
)
|
||||
|
||||
func TestClientIPWithoutProxy(t *testing.T) {
|
||||
r := httptest.NewRequest(http.MethodGet, "/", nil)
|
||||
r.RemoteAddr = "203.0.113.9:1234"
|
||||
r.Header.Set("X-Forwarded-For", "198.51.100.1")
|
||||
ip := httpx.ClientIP(r, nil)
|
||||
if ip != "203.0.113.9" {
|
||||
t.Fatalf("got %q", ip)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseCIDRsAndTrustedXFF(t *testing.T) {
|
||||
trusted := httpx.ParseCIDRs([]string{"127.0.0.1/32"})
|
||||
r := httptest.NewRequest(http.MethodGet, "/", nil)
|
||||
r.RemoteAddr = "127.0.0.1:9999"
|
||||
r.Header.Set("X-Forwarded-For", "198.51.100.7, 127.0.0.1")
|
||||
ip := httpx.ClientIP(r, trusted)
|
||||
if ip != "198.51.100.7" {
|
||||
t.Fatalf("got %q", ip)
|
||||
}
|
||||
r.Header.Set("X-Forwarded-Proto", "https")
|
||||
if !httpx.IsHTTPS(r, trusted) {
|
||||
t.Fatal("expected https via proxy")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
// Package httpx 提供管理接口共用的 HTTP 辅助(JSON 信封、客户端 IP、HTTPS 判定)。
|
||||
package httpx
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"errors"
|
||||
"io"
|
||||
"net/http"
|
||||
)
|
||||
|
||||
// ErrorBody 是失败响应里的 error 对象。
|
||||
type ErrorBody struct {
|
||||
Code string `json:"code"`
|
||||
Message string `json:"message"`
|
||||
}
|
||||
|
||||
// Envelope 是管理接口通用响应信封。
|
||||
type Envelope struct {
|
||||
OK bool `json:"ok"`
|
||||
Data any `json:"data,omitempty"`
|
||||
Error *ErrorBody `json:"error,omitempty"`
|
||||
}
|
||||
|
||||
// WriteJSON 写入 JSON 响应。
|
||||
func WriteJSON(w http.ResponseWriter, status int, v any) {
|
||||
w.Header().Set("Content-Type", "application/json; charset=utf-8")
|
||||
w.WriteHeader(status)
|
||||
enc := json.NewEncoder(w)
|
||||
enc.SetEscapeHTML(false)
|
||||
_ = enc.Encode(v)
|
||||
}
|
||||
|
||||
// WriteOK 写入成功信封。
|
||||
func WriteOK(w http.ResponseWriter, data any) {
|
||||
WriteJSON(w, http.StatusOK, Envelope{OK: true, Data: data})
|
||||
}
|
||||
|
||||
// WriteError 写入失败信封。
|
||||
func WriteError(w http.ResponseWriter, status int, code, message string) {
|
||||
WriteJSON(w, status, Envelope{
|
||||
OK: false,
|
||||
Error: &ErrorBody{
|
||||
Code: code,
|
||||
Message: message,
|
||||
},
|
||||
})
|
||||
}
|
||||
|
||||
// DecodeJSON 解码请求 JSON 体;空体对 dst 保持零值。
|
||||
func DecodeJSON(r *http.Request, dst any) error {
|
||||
defer func() { _ = r.Body.Close() }()
|
||||
dec := json.NewDecoder(r.Body)
|
||||
dec.DisallowUnknownFields()
|
||||
if err := dec.Decode(dst); err != nil {
|
||||
if errors.Is(err, io.EOF) {
|
||||
return nil
|
||||
}
|
||||
return err
|
||||
}
|
||||
return nil
|
||||
}
|
||||
Reference in New Issue
Block a user