Files
MailUI4Agents/server/internal/handler/push_test.go
JianFeeeee 46fa7fa729 feat(push): 可选、配置式、多厂商的推送通道(HMS 为首个实现)
用户要求:推送密钥必须是可选项(自部署后端不能写死推送方式),且要支持
多厂商配置式接入 —— 每个用户各自部署服务器、自己选厂商、自己配凭证。
所以落地成:

· internal/push:通道抽象 + 工厂表(RegisterType),加厂商不改配置层与端点形状;
  HMS 只是第一个实现(internal/push/hms.go)
· 配置在 PUSH_CONFIG(默认 <AGENTMAIL_DATA_DIR>/push.json),一项一个厂商,
  凭证走文件(app_secret_file / files.*,建议 600);环境变量只是可选覆盖
· 没配 = 整条推送路径连一次查库都不发生(shouldDispatch 早退);
  单项配错(未知类型/密钥读不到/enabled:false)只跳过那一条,不影响启动
· push_tokens 表带 provider 维度 + 三个 /me/devices/push-token 端点;
  没配推送时端点照存并回 enabled:false(登记成功 != 服务端开了推送)
· notify.Recipients 末尾异步挂钩:收件人名单直接用 SSE 那份 seen(两条通道
  共用同一份"谁该收到"的判据);失败只记日志,绝不拖住收信

HMS 的形状是拿真凭证打线上接口问出来的(v1 + message.token[] + testMessage;
payload/target 形状 v1 不认、v2 要服务账号 JWT)。未上架应用必须 test_message=true,
单批 ≤10 token(MaxTokensPerRequest 声明)、每日 1000 条兜底(项目级额度)。
实测:App ID + App Secret 能换到 access_token(3600s);形状被线上服务接受。

判据:repo 6 条 + push 12 条 + handler 3 组,全部做过**变异验证** ——
过程中抓出两条假判据(异步分发与 t.Cleanup 赛跑而假绿;密钥文件优先级没被覆盖)
并补掉。Go 全量测试与 go vet 干净。

★ 未验:端到端真机送达(需要真机 token + 客户端按 com.jianf.agentmail 重编并签名,
签名指纹还要在 AGC 登记)—— 从未真正发出过一条能到达设备的推送。
详见 docs/HMS-PUSH-PLAN.md 的「实现状态」一节。
2026-09-15 11:21:00 +08:00

190 lines
6.8 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"
)
/*
推送登记端点的判据(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":"sess-1"}`, "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"] != "sess-1" || 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)
}
}