Files
MailUI4Agents/server/internal/handler/push_test.go
JianFeeeee 7e5b11392f fix(push): session_id 非空时必须是合法 UUID —— 它此前静默接受任意字符串,真 bug 就是这么活下来的
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 行,**没有历史脏值要清**。
2026-09-17 19:42:41 +08:00

239 lines
8.9 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

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()