dsh(f07daac1 之后的 5fab8734)报:客户端一度把这个字段填成**服务器地址** (`AccountInfo.server` = `https://…/api/v1`),而这里只 TrimSpace、不看格式 ⇒ 不 400、不影响收信、不进日志;而**投递路径也不看它**(通知里的 session_id 来自 邮件自己,见 notify/mail.go;dispatch 只用 token 的 Provider/Token)。 两条路都不看 ⇒ 存了假值没有任何机制会报警,可它存在的唯一目的就是 「点通知回到那条会话」。 改动只有一条校验(非空才校验,空串仍合法 —— 契约允许"客户端还没进任何会话"), 400 文案自带药方(写明"要么 UUID、要么留空")。 判据 `TestPushTokenSessionIDMustBeUUID` 四个用例:服务器地址顶替 / 随便一个词 / 空串 / 真 UUID。**变体验证两个方向**: · 去掉校验(连 import 一起去)⇒ 「非法」两条当场红(实际 200); · 把判断写反(`err == nil` 才报错)⇒ 合法的真 UUID 也红 ⇒ 证明它不是"恒 400"。 恢复后 `go test ./...` 全绿、`go vet` 干净。 取舍:老客户端带旧值上报会吃 400,按契约 400 归静默 ⇒ 不打扰用户、只是那次不上报。 `push_tokens` 现为 0 行,**没有历史脏值要清**。
239 lines
8.9 KiB
Go
239 lines
8.9 KiB
Go
package handler
|
||
|
||
import (
|
||
"context"
|
||
"encoding/json"
|
||
"net/http"
|
||
"net/http/httptest"
|
||
"path/filepath"
|
||
"strings"
|
||
"testing"
|
||
|
||
"github.com/agentmail/gateway/internal/db"
|
||
"github.com/agentmail/gateway/internal/middleware"
|
||
"github.com/agentmail/gateway/internal/models"
|
||
"github.com/agentmail/gateway/internal/push"
|
||
"github.com/google/uuid"
|
||
)
|
||
|
||
/*
|
||
推送登记端点的判据(2026-09-15)。
|
||
|
||
用户的要求是「推送密钥应当是可选项」—— 自部署实例**没配推送**是常态。
|
||
所以这里钉住的是:**没配凭证时端点照样能用**(登记照存、回 enabled=false),
|
||
而不是"没配就报错"。客户端据此知道「登记成功了,但服务端现在没开推送」,
|
||
而不是把收不到通知当成登记失败去反复重试。
|
||
|
||
(本文件里的两个用例有先后依赖:先验「没配 = enabled:false」,再验「配了 = enabled:true」。
|
||
push 包的全局通道表只增不减,顺序反了前者会假红。)
|
||
*/
|
||
|
||
func setupPushHandlerDB(t *testing.T) {
|
||
t.Helper()
|
||
dir := t.TempDir()
|
||
if err := db.Connect(context.Background(), "sqlite://"+filepath.Join(dir, "t.db")); err != nil {
|
||
t.Fatalf("connect: %v", err)
|
||
}
|
||
if err := db.Migrate(context.Background()); err != nil {
|
||
t.Fatalf("migrate: %v", err)
|
||
}
|
||
t.Cleanup(db.Close)
|
||
}
|
||
|
||
func pushReq(method, body, user string) *http.Request {
|
||
var r *http.Request
|
||
if body == "" {
|
||
r = httptest.NewRequest(method, "/api/v1/me/devices/push-token", nil)
|
||
} else {
|
||
r = httptest.NewRequest(method, "/api/v1/me/devices/push-token", strings.NewReader(body))
|
||
}
|
||
if user != "" {
|
||
r = r.WithContext(context.WithValue(r.Context(), middleware.UserKey, &models.User{Username: user}))
|
||
}
|
||
return r
|
||
}
|
||
|
||
func decodeBody(t *testing.T, rec *httptest.ResponseRecorder) map[string]any {
|
||
t.Helper()
|
||
var out map[string]any
|
||
if err := json.Unmarshal(rec.Body.Bytes(), &out); err != nil {
|
||
t.Fatalf("响应不是 JSON: %v(%s)", err, rec.Body.String())
|
||
}
|
||
return out
|
||
}
|
||
|
||
// 没配任何推送通道时:登记照存、回 enabled=false,且**不是错误**。
|
||
func TestPushTokenEndpointsWorkWithoutProviderCreds(t *testing.T) {
|
||
setupPushHandlerDB(t)
|
||
ctx := context.Background()
|
||
|
||
rec := httptest.NewRecorder()
|
||
RegisterPushToken(rec, pushReq(http.MethodPost,
|
||
`{"provider":"hms","token":"tok-abc123456","device_name":"我的手机","session_id":"`+testSID+`"}`, "alice"))
|
||
if rec.Code != http.StatusOK {
|
||
t.Fatalf("没配推送时登记也必须是 200(那是正常配置,不是故障),实际 %d: %s", rec.Code, rec.Body.String())
|
||
}
|
||
body := decodeBody(t, rec)
|
||
if body["enabled"] != false {
|
||
t.Fatalf("没配推送时应回 enabled=false,实际 %v", body["enabled"])
|
||
}
|
||
if body["ok"] != true {
|
||
t.Fatalf("登记本身必须成功,实际 %v", body)
|
||
}
|
||
|
||
// GET 能列出来,且 token 只回尾 6 位(接口不该能读到 token 全文)
|
||
rec = httptest.NewRecorder()
|
||
ListPushTokens(rec, pushReq(http.MethodGet, "", "alice"))
|
||
if rec.Code != http.StatusOK {
|
||
t.Fatalf("GET 失败: %d", rec.Code)
|
||
}
|
||
body = decodeBody(t, rec)
|
||
items, _ := body["tokens"].([]any)
|
||
if len(items) != 1 {
|
||
t.Fatalf("应列出 1 条登记,实际 %v", body["tokens"])
|
||
}
|
||
item, _ := items[0].(map[string]any)
|
||
if item["token_tail"] != "123456" {
|
||
t.Fatalf("token 应只回尾 6 位,实际 %v", item["token_tail"])
|
||
}
|
||
if _, leaked := item["token"]; leaked {
|
||
t.Fatal("不该回 token 全文(多一处泄漏面)")
|
||
}
|
||
if item["session_id"] != testSID || item["provider"] != "hms" {
|
||
t.Fatalf("登记字段不对: %v", item)
|
||
}
|
||
|
||
// 别人看不到我的登记
|
||
rec = httptest.NewRecorder()
|
||
ListPushTokens(rec, pushReq(http.MethodGet, "", "bob"))
|
||
if items, _ := decodeBody(t, rec)["tokens"].([]any); len(items) != 0 {
|
||
t.Fatalf("bob 不该看到 alice 的登记: %v", items)
|
||
}
|
||
|
||
// 注销
|
||
rec = httptest.NewRecorder()
|
||
DeletePushToken(rec, pushReq(http.MethodDelete, `{"provider":"hms","token":"tok-abc123456"}`, "alice"))
|
||
if rec.Code != http.StatusOK {
|
||
t.Fatalf("注销失败: %d %s", rec.Code, rec.Body.String())
|
||
}
|
||
if removed, _ := decodeBody(t, rec)["deleted"].(bool); !removed {
|
||
t.Fatal("注销应报告删到了行")
|
||
}
|
||
rec = httptest.NewRecorder()
|
||
ListPushTokens(rec, pushReq(http.MethodGet, "", "alice"))
|
||
if items, _ := decodeBody(t, rec)["tokens"].([]any); len(items) != 0 {
|
||
t.Fatalf("注销后不该还有登记: %v", items)
|
||
}
|
||
_ = ctx
|
||
}
|
||
|
||
// 配了通道之后,同一个端点回 enabled=true 并报出通道名(客户端据此判断"现在能收到")。
|
||
func TestPushTokenEndpointReportsEnabledProviders(t *testing.T) {
|
||
setupPushHandlerDB(t)
|
||
push.Register(fakeProviderForHandler{name: "hms"})
|
||
|
||
rec := httptest.NewRecorder()
|
||
RegisterPushToken(rec, pushReq(http.MethodPost, `{"provider":"hms","token":"tok-xyz"}`, "alice"))
|
||
body := decodeBody(t, rec)
|
||
if body["enabled"] != true {
|
||
t.Fatalf("配了通道应回 enabled=true,实际 %v", body["enabled"])
|
||
}
|
||
providers, _ := body["providers"].([]any)
|
||
if len(providers) != 1 || providers[0] != "hms" {
|
||
t.Fatalf("providers 应含 hms,实际 %v", providers)
|
||
}
|
||
}
|
||
|
||
type fakeProviderForHandler struct{ name string }
|
||
|
||
func (f fakeProviderForHandler) Name() string { return f.name }
|
||
func (f fakeProviderForHandler) MaxTokensPerRequest() int { return 10 }
|
||
func (f fakeProviderForHandler) Send(context.Context, []string, push.NewMail) error {
|
||
return nil
|
||
}
|
||
|
||
// 入参校验与鉴权。provider 只做**形状**校验(不做白名单:白名单会把"服务端还没实现的
|
||
// 那个通道"变成客户端的 400,而那恰恰是最不该拦的时候)。
|
||
func TestPushTokenEndpointInputValidation(t *testing.T) {
|
||
setupPushHandlerDB(t)
|
||
|
||
cases := []struct {
|
||
name string
|
||
body string
|
||
user string
|
||
want int
|
||
}{
|
||
{"未登录", `{"provider":"hms","token":"t"}`, "", http.StatusUnauthorized},
|
||
{"provider 含大写", `{"provider":"HMS","token":"t"}`, "alice", http.StatusBadRequest},
|
||
{"provider 为空", `{"provider":"","token":"t"}`, "alice", http.StatusBadRequest},
|
||
{"provider 太长", `{"provider":"` + strings.Repeat("a", 33) + `","token":"t"}`, "alice", http.StatusBadRequest},
|
||
{"token 为空", `{"provider":"hms","token":""}`, "alice", http.StatusBadRequest},
|
||
{"token 超长", `{"provider":"hms","token":"` + strings.Repeat("a", 513) + `"}`, "alice", http.StatusBadRequest},
|
||
{"未实现的新通道也收下", `{"provider":"webpush","token":"t"}`, "alice", http.StatusOK},
|
||
}
|
||
for _, c := range cases {
|
||
t.Run(c.name, func(t *testing.T) {
|
||
rec := httptest.NewRecorder()
|
||
RegisterPushToken(rec, pushReq(http.MethodPost, c.body, c.user))
|
||
if rec.Code != c.want {
|
||
t.Fatalf("%s: 期望 %d,实际 %d(%s)", c.name, c.want, rec.Code, rec.Body.String())
|
||
}
|
||
})
|
||
}
|
||
|
||
// DELETE 也要鉴权
|
||
rec := httptest.NewRecorder()
|
||
DeletePushToken(rec, pushReq(http.MethodDelete, `{"provider":"hms","token":"t"}`, ""))
|
||
if rec.Code != http.StatusUnauthorized {
|
||
t.Fatalf("未登录的注销应 401,实际 %d", rec.Code)
|
||
}
|
||
}
|
||
|
||
/*
|
||
★ session_id 的格式契约(pi 2026-09-17,由 dsh 报的真 bug 催生)。
|
||
|
||
那个 bug:客户端一度把这个字段填成**服务器地址**(`https://…/api/v1`),而服务端
|
||
只 TrimSpace、允许任何字符串 ⇒ 不报错、不影响收信、不进日志;而**投递路径不看它**
|
||
(通知里的 session_id 来自邮件自己,见 notify/mail.go),所以它错了也没有任何机制会报警。
|
||
但这个字段存在的唯一目的就是「点通知回到那条会话」—— 存错值 = 将来路由到错会话。
|
||
|
||
所以这里钉三条:
|
||
① 非空且不是 UUID ⇒ **400**(把"静默存错值"换成"这次静默不上报");
|
||
② 空串仍然合法(契约:客户端还没进任何会话);
|
||
③ 合法 UUID ⇒ 照存(登记成功)。
|
||
*/
|
||
func TestPushTokenSessionIDMustBeUUID(t *testing.T) {
|
||
setupPushHandlerDB(t)
|
||
user := "bob"
|
||
|
||
cases := []struct {
|
||
name string
|
||
sid string
|
||
status int
|
||
}{
|
||
{"非法:服务器地址顶替(就是这个 bug 的形状)", "https://mail.jianfgit.xyz/api/v1", http.StatusBadRequest},
|
||
{"非法:随便一个词", "sess-1", http.StatusBadRequest},
|
||
{"合法:空串(还没进任何会话)", "", http.StatusOK},
|
||
{"合法:真 UUID", testSID, http.StatusOK},
|
||
}
|
||
for _, c := range cases {
|
||
t.Run(c.name, func(t *testing.T) {
|
||
body := `{"provider":"hms","token":"tok-fmt-` + c.name + `","session_id":"` + c.sid + `"}`
|
||
rec := httptest.NewRecorder()
|
||
RegisterPushToken(rec, pushReq(http.MethodPost, body, user))
|
||
if rec.Code != c.status {
|
||
t.Fatalf("session_id=%q 期望 %d,实际 %d(%s)", c.sid, c.status, rec.Code, rec.Body.String())
|
||
}
|
||
// 错误报文要自带药方:说清"要么是 UUID、要么留空",而不是只说"非法"。
|
||
if c.status == http.StatusBadRequest {
|
||
if msg, _ := decodeBody(t, rec)["error"].(string); !strings.Contains(msg, "UUID") {
|
||
t.Fatalf("400 的文案没告诉人怎么改:%q", msg)
|
||
}
|
||
}
|
||
})
|
||
}
|
||
}
|
||
|
||
// testSID 是一个合法会话 UUID(handler 层不校验它是否真的存在 —— 那是会话解析的事)
|
||
var testSID = uuid.New().String()
|