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 的「实现状态」一节。
This commit is contained in:
2026-09-15 11:21:00 +08:00
parent b806a05bfa
commit 46fa7fa729
14 changed files with 2349 additions and 0 deletions

View File

@ -473,3 +473,19 @@ CREATE TABLE IF NOT EXISTS user_appearance (
updated_at TIMESTAMPTZ DEFAULT NOW(),
PRIMARY KEY (user_id)
);
-- ─── 设备推送 token(可选通道)── 语义与 init_sqlite.sql 里的同名表一致 ──
CREATE TABLE IF NOT EXISTS push_tokens (
token_id TEXT NOT NULL,
provider TEXT NOT NULL,
token TEXT NOT NULL,
owner_name TEXT NOT NULL,
session_id TEXT NOT NULL DEFAULT '',
device_name TEXT NOT NULL DEFAULT '',
created_at TIMESTAMPTZ DEFAULT NOW(),
updated_at TIMESTAMPTZ DEFAULT NOW(),
PRIMARY KEY (token_id)
);
CREATE UNIQUE INDEX IF NOT EXISTS uniq_push_tokens_provider_token ON push_tokens (provider, token);
CREATE INDEX IF NOT EXISTS idx_push_tokens_owner ON push_tokens (owner_name);

View File

@ -533,3 +533,31 @@ CREATE TABLE IF NOT EXISTS user_appearance (
updated_at DATETIME DEFAULT (strftime('%Y-%m-%d %H:%M:%f','now')),
PRIMARY KEY (user_id)
);
-- ─── 设备推送 token(可选通道)──────────────────────────────────────────
--
-- 为什么有 provider 列:**推送是自部署后端的可选项,不是内置依赖**
-- (2026-09-15 用户明确要求:「不能写死推送方式,因为我们是自部署后端」
-- 「即推送密钥应当是可选项」)。一个自部署实例可能一个推送渠道都没配
-- —— 这是常态而不是配置错误;也可能同时接华为 HMS 与别的通道。
-- 加通道不该动 schema、不该动端点形状。
--
-- owner_name 是**注册者**(登录用户)。收件判据与 SSE 同源:一封新邮件推给
-- 谁,推送就发给谁 —— 两条通道不该有两套「谁该收到」的定义。
CREATE TABLE IF NOT EXISTS push_tokens (
token_id TEXT NOT NULL,
provider TEXT NOT NULL,
token TEXT NOT NULL,
owner_name TEXT NOT NULL,
-- 注册时客户端所在的会话:点通知要回到那条会话里的那封信
session_id TEXT NOT NULL DEFAULT '',
device_name TEXT NOT NULL DEFAULT '',
created_at DATETIME DEFAULT (strftime('%Y-%m-%d %H:%M:%f','now')),
updated_at DATETIME DEFAULT (strftime('%Y-%m-%d %H:%M:%f','now')),
PRIMARY KEY (token_id)
);
-- 一个 token 只能属于一个注册者:换人登录是**转移**,不是并存 —— 否则上一任
-- 用户的通知会继续推到同一台设备上(那是隐私事故,不只是脏数据)。
CREATE UNIQUE INDEX IF NOT EXISTS uniq_push_tokens_provider_token ON push_tokens (provider, token);
CREATE INDEX IF NOT EXISTS idx_push_tokens_owner ON push_tokens (owner_name);

View File

@ -0,0 +1,162 @@
package handler
import (
"log"
"net/http"
"regexp"
"strings"
"github.com/agentmail/gateway/internal/middleware"
"github.com/agentmail/gateway/internal/push"
"github.com/agentmail/gateway/internal/repo"
)
/*
设备推送登记 —— /api/v1/me/devices/push-token
# 为什么这三个端点"没配推送也能用"
推送是自部署后端的**可选项**(用户 2026-09-15 明确要求)。所以:
- 登记/注销**不看服务端有没有配凭证**:照存。管理员之后把 `HMS_APP_ID/SECRET`
配上就立刻生效,客户端不必再登记一次(它那时可能已经不在前台了)。
- 响应里回 `enabled`:客户端据此知道「登记成功了,但服务端现在还没开推送」——
而不是把"没收到通知"当成登记失败去反复重试。
- 查不到 provider 凭证时**不返回错误状态码**:那是正常配置,不是故障。
# 为什么 provider 用字符串而不是枚举
加通道(web push、别的厂商)不该动端点形状、不该动 schema。合法性只做**形状**
校验(字符集与长度),不做白名单 —— 白名单会把「服务端还没实现的那个通道」
变成客户端的 400,而那恰恰是最不该拦的时候。
*/
var pushProviderRe = regexp.MustCompile(`^[a-z0-9_-]{1,32}$`)
// pushTokenRequest 是登记/注销共用的请求体。
type pushTokenRequest struct {
Provider string `json:"provider"`
Token string `json:"token"`
// DeviceName 仅供人辨认(如"我的手机"),服务端不做判据。
DeviceName string `json:"device_name"`
// SessionID 是客户端当前所在的会话:点通知要回到那条会话里的那封信。
// 允许为空(客户端还没进任何会话),此时通知只带 mail_id。
SessionID string `json:"session_id"`
}
// pushEnabled 是三个端点共用的响应片段。
func pushStatus() map[string]interface{} {
return map[string]interface{}{
"enabled": push.Enabled(),
"providers": push.Names(),
}
}
// RegisterPushToken 登记(或刷新)一台设备的推送地址。
//
// POST /api/v1/me/devices/push-token
func RegisterPushToken(w http.ResponseWriter, r *http.Request) {
user := middleware.GetUser(r)
if user == nil {
Error(w, http.StatusUnauthorized, "not authenticated")
return
}
var req pushTokenRequest
if !DecodeBody(w, r, &req) {
return
}
req.Provider = strings.TrimSpace(req.Provider)
req.Token = strings.TrimSpace(req.Token)
if !pushProviderRe.MatchString(req.Provider) {
Error(w, http.StatusBadRequest, "provider 非法(只允许小写字母、数字、下划线、连字符,最长 32)")
return
}
if req.Token == "" || len(req.Token) > 512 {
Error(w, http.StatusBadRequest, "token 非法(不能为空,最长 512 字符)")
return
}
if err := repo.UpsertPushToken(r.Context(), req.Provider, req.Token, user.Username,
strings.TrimSpace(req.SessionID), strings.TrimSpace(req.DeviceName)); err != nil {
Error(w, http.StatusInternalServerError, "登记推送地址失败")
return
}
// 顺手清理长期没刷新的登记(90 天)。放在登记路径上是因为它天然低频
// (每台设备只在启动/token 轮换时来一次),不需要定时任务。
// 失败只记日志:清理失败不该让一次正常登记变成 500。
if _, err := repo.PruneStalePushTokens(r.Context()); err != nil {
log.Printf("[push] 清理过期推送登记失败: %v", err)
}
resp := pushStatus()
resp["ok"] = true
JSON(w, http.StatusOK, resp)
}
// DeletePushToken 注销一台设备的推送地址。
//
// DELETE /api/v1/me/devices/push-token
func DeletePushToken(w http.ResponseWriter, r *http.Request) {
user := middleware.GetUser(r)
if user == nil {
Error(w, http.StatusUnauthorized, "not authenticated")
return
}
var req pushTokenRequest
if !DecodeBody(w, r, &req) {
return
}
req.Provider = strings.TrimSpace(req.Provider)
req.Token = strings.TrimSpace(req.Token)
if !pushProviderRe.MatchString(req.Provider) || req.Token == "" {
Error(w, http.StatusBadRequest, "provider 或 token 非法")
return
}
deleted, err := repo.DeletePushToken(r.Context(), req.Provider, req.Token, user.Username)
if err != nil {
Error(w, http.StatusInternalServerError, "注销推送地址失败")
return
}
resp := pushStatus()
resp["ok"] = true
// deleted=false 不是错误(本来就没登记),但客户端可据此避免重复注销。
resp["deleted"] = deleted
JSON(w, http.StatusOK, resp)
}
// ListPushTokens 列出自己的推送登记。
//
// GET /api/v1/me/devices/push-token
func ListPushTokens(w http.ResponseWriter, r *http.Request) {
user := middleware.GetUser(r)
if user == nil {
Error(w, http.StatusUnauthorized, "not authenticated")
return
}
tokens, err := repo.ListPushTokensOf(r.Context(), user.Username)
if err != nil {
Error(w, http.StatusInternalServerError, "读取推送登记失败")
return
}
items := make([]map[string]interface{}, 0, len(tokens))
for _, t := range tokens {
items = append(items, map[string]interface{}{
"id": t.TokenID,
"provider": t.Provider,
// token 只回尾 6 位:客户端不需要全文(它自己刚发过来的),
// 而一个能读到全文的接口等于多一处泄漏面。
"token_tail": tailOf(t.Token, 6),
"device_name": t.DeviceName,
"session_id": t.SessionID,
})
}
resp := pushStatus()
resp["tokens"] = items
JSON(w, http.StatusOK, resp)
}
func tailOf(s string, n int) string {
r := []rune(s)
if len(r) <= n {
return s
}
return string(r[len(r)-n:])
}

View File

@ -0,0 +1,189 @@
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)
}
}

View File

@ -22,6 +22,7 @@ import (
"context"
"github.com/agentmail/gateway/internal/models"
"github.com/agentmail/gateway/internal/push"
"github.com/agentmail/gateway/internal/repo"
"github.com/agentmail/gateway/internal/sse"
"github.com/google/uuid"
@ -224,6 +225,26 @@ func Recipients(ctx context.Context, m Mail) {
if !seen[m.From] {
sse.Default.SendToRecipient(m.From, "session_update", update)
}
// 第二条送达通道:设备推送(可选,没配凭证时这里是空操作)。
//
// 放在 SSE 之后且**异步**:推送慢不拖住收信(一封邮件的送达不能被
// 一个卡住的 HTTP 请求拖住),失败也只记日志。收件人名单直接用上面
// 的 seen —— 两条通道必须共用同一份「谁该收到这封信」的判据,
// 各算一套的话抄送方总有一边收不到(SSE 那边已经因为这个踩过一次)。
pushRecipients := make([]string, 0, len(seen))
for name := range seen {
if name != "" {
pushRecipients = append(pushRecipients, name)
}
}
push.NotifyNewMail(ctx, push.NewMail{
MailID: m.MailID.String(),
SessionID: m.SessionID.String(),
From: m.From,
Subject: m.Subject,
Recipients: pushRecipients,
})
}
// SessionActive 只刷新某一方的会话列表,不推 new_mail。

View File

@ -0,0 +1,304 @@
package push
import (
"encoding/json"
"fmt"
"log"
"os"
"path/filepath"
"strings"
)
/*
推送通道的**配置式**接入(2026-09-15 用户的第二条要求)。
原话:「我们要支持多厂商配置式接入统一推送服务,推送密钥(文件)应当是配置项,
用户可为不同的推送服务配置对应的凭证,因为我们的设计中是每个用户各自部署服务器」。
所以契约是这样定的:
- **多厂商**:配置里是一张表,每一项是一个厂商(`type`)。一个实例可以同时接
HMS 和小米/Web Push/任何东西 —— 只要那个厂商在 factories 里注册过实现。
- **配置式**:加厂商不改代码路径,只加一个 `Factory` 实现 + 一行 `RegisterType`;
用户侧只改配置文件,不重编译。
- **凭证/密钥文件是配置项**:`app_secret_file` 指向密钥文件(推荐),也接受内联
`app_secret`(图省事/做实验);`client_config_file` 指向厂商给的客户端配置
(如华为的 agconnect-services.json)—— 这类文件属于**部署物料**,由部署者提供。
- **谁部署谁配**:每个用户自己部署服务端、自己选厂商、自己填凭证。没配 = 推送
不可用,但服务端一切照常(这就是「可选」的落地)。
配置文件位置:`PUSH_CONFIG` 指定;默认 `<AGENTMAIL_DATA_DIR>/push.json`。
文件不存在**不是错误**(自部署实例默认就没配推送)。
# 单项配错不能拖垮整个服务
某一条配置写坏(类型未知、密钥文件读不到、JSON 写错)时:**只跳过那一条**并打印
一条明确的日志,其余条目照常启用,网关照常启动。理由很直接:推送是可选旁路,
它不该有能力让整个邮件服务起不来 —— 那是把"锦上添花"变成了"单点故障"。
*/
// ProviderConfig 是配置里的一项:一个推送厂商 + 它自己的凭证。
type ProviderConfig struct {
// Type 是厂商实现名("hms"、"webpush"…),必须已 RegisterType。
Type string `json:"type"`
// Name 覆盖推送给客户端看的 provider 名(默认 = Type)。
// 用途:同一个实例接两套同厂商凭证(例如两个应用)时区分开来,
// 客户端登记 token 时用的就是这个值。
Name string `json:"name"`
// Enabled 缺省视为 true;显式 false = 留配置但不启用。
Enabled *bool `json:"enabled"`
// AppID / AppSecret 是厂商的凭证。AppSecret 建议走 AppSecretFile。
AppID string `json:"app_id"`
AppSecret string `json:"app_secret"`
// AppSecretFile 指向**存放密钥的文件**(配置项,不是硬编码)。
AppSecretFile string `json:"app_secret_file"`
// ClientConfigFile 指向厂商给的客户端配置文件(如 agconnect-services.json)。
// 服务端用它核对 app_id/package_name 是否与客户端一致——不一致的推送永远送不到,
// 而症状会表现为"推送静默失效",所以这里宁可启动时就说清楚。
ClientConfigFile string `json:"client_config_file"`
// Files 是**厂商自定义的文件类配置**(键名由厂商实现定义)。
//
// 为什么要有它:不同厂商的密钥形状本就不同(华为是 app_secret,
// Web Push 是 VAPID 密钥对,有的用服务账号 JSON…)。给每个厂商加一个专用字段
// 会让配置层随厂商数量膨胀;一张「名字 → 文件路径」的表则不用改配置层就能接新厂商。
//
// 例:{"app_secret": "/etc/agentmail/hms.secret", "vapid_private_key": "/etc/agentmail/vapid.pem"}
Files map[string]string `json:"files"`
// TestMessage 见 hms.go:未上架应用必须为 true。缺省 true。
TestMessage *bool `json:"test_message"`
// DailyLimit 每日发送上限(条),0 = 用实现的默认值。
DailyLimit int `json:"daily_limit"`
}
// Factory 按配置造一个通道。凭证已在这之前解析好(见 resolveSecret)。
type Factory func(cfg ProviderConfig) (Notifier, error)
var factories = map[string]Factory{}
// RegisterType 注册一个厂商实现。加厂商 = 加一个实现 + 一行这个调用。
func RegisterType(name string, f Factory) {
factories[name] = f
}
func init() {
RegisterType("hms", newHMSFromConfig)
}
// KnownTypes 列出已注册的厂商类型(日志与文档用)。
func KnownTypes() []string {
out := make([]string, 0, len(factories))
for k := range factories {
out = append(out, k)
}
return out
}
type configFile struct {
Providers []ProviderConfig `json:"providers"`
}
// configPath 返回配置文件路径。
func configPath() string {
if p := strings.TrimSpace(os.Getenv("PUSH_CONFIG")); p != "" {
return p
}
dir := strings.TrimSpace(os.Getenv("AGENTMAIL_DATA_DIR"))
if dir == "" {
dir = "data"
}
return filepath.Join(dir, "push.json")
}
// LoadProviders 读配置并造出所有启用的通道。
//
// 单项失败只跳过该项(见包注释):返回值可能少于配置里的条目数。
func LoadProviders() []Notifier {
path := configPath()
entries, err := readConfigEntries(path)
if err != nil {
log.Printf("[push] 配置文件 %s 读取失败,推送不可用(不影响邮件服务): %v", path, err)
return nil
}
// 环境变量是**可选覆盖**:只有在配置里没有同类型的条目时才补一条。
// 保留它是因为临时验证(以及没有配置文件的小部署)很常用;
// 但它不是主路径 —— 主路径是配置文件(用户要求「密钥应当是配置项」)。
if env, ok := hmsConfigFromEnv(); ok && !hasType(entries, env.Type) {
entries = append(entries, env)
}
var out []Notifier
for i, cfg := range entries {
cfg.Type = strings.TrimSpace(cfg.Type)
if cfg.Type == "" {
log.Printf("[push] 第 %d 项缺 type 字段,已跳过", i+1)
continue
}
if cfg.Enabled != nil && !*cfg.Enabled {
log.Printf("[push] %s(%s)在配置里是 disabled,已跳过", nameOf(cfg), cfg.Type)
continue
}
f, ok := factories[cfg.Type]
if !ok {
log.Printf("[push] 不支持的类型 %q(已注册:%s),已跳过该项", cfg.Type, strings.Join(KnownTypes(), ", "))
continue
}
if err := resolveSecret(&cfg); err != nil {
log.Printf("[push] %s(%s)凭证不可用,已跳过: %v", nameOf(cfg), cfg.Type, err)
continue
}
checkClientConfig(&cfg)
n, err := f(cfg)
if err != nil {
log.Printf("[push] %s(%s)初始化失败,已跳过: %v", nameOf(cfg), cfg.Type, err)
continue
}
out = append(out, n)
}
return out
}
// Setup 读配置、注册通道,并把结果打成一行日志(main 调用它)。
//
// 返回已启用的通道名:空表示"这台实例没配推送"—— 那是正常状态,不是错误,
// 所以这里用普通日志而不是告警。
func Setup() []string {
providers := LoadProviders()
for _, p := range providers {
Register(p)
}
if len(providers) == 0 {
log.Printf("[push] 未配置推送通道(%s 不存在或为空)—— 正常状态,SSE 仍是收信主通道", configPath())
return nil
}
return Names()
}
func readConfigEntries(path string) ([]ProviderConfig, error) {
b, err := os.ReadFile(path)
if err != nil {
if os.IsNotExist(err) {
return nil, nil // 没配就是没配
}
return nil, err
}
if strings.TrimSpace(string(b)) == "" {
return nil, nil
}
var f configFile
if err := json.Unmarshal(b, &f); err != nil {
return nil, fmt.Errorf("JSON 解析失败: %w", err)
}
return f.Providers, nil
}
func hasType(entries []ProviderConfig, t string) bool {
for _, e := range entries {
if e.Type == t {
return true
}
}
return false
}
// resolveSecret 把密钥文件的内容解析到 cfg.AppSecret(文件优先于内联)。
//
// **不在这里要求密钥必须存在**:不是每个厂商都用 app_secret(Web Push 用 VAPID
// 密钥对),"这个字段是必需的"是各厂商自己的事,由它的 Factory 判定。
// 这条是判据抓出来的:第一版把"必须有密钥"写在通用层,于是无密钥的厂商条目
// 被默默跳过(测试里 test-echo 就没建起来)。
//
// 但**指明了文件却读不到**是真错误(配置里写了却用不了),所以它照旧让该项被判失败。
func resolveSecret(cfg *ProviderConfig) error {
file := strings.TrimSpace(cfg.AppSecretFile)
if file == "" {
file = strings.TrimSpace(cfg.Files["app_secret"])
}
if file == "" {
return nil
}
fi, err := os.Stat(file)
if err != nil {
return fmt.Errorf("密钥文件不可读 %s: %w", file, err)
}
if fi.Mode().Perm()&0o044 != 0 {
log.Printf("[push] 提醒:密钥文件 %s 权限 %o 对同组/其他人可读,建议 chmod 600", file, fi.Mode().Perm())
}
b, err := os.ReadFile(file)
if err != nil {
return fmt.Errorf("读密钥文件 %s 失败: %w", file, err)
}
secret := strings.TrimSpace(string(b))
if secret == "" {
return fmt.Errorf("密钥文件 %s 是空的", file)
}
cfg.AppSecret = secret
return nil
}
// checkClientConfig 核对客户端配置文件里的 app_id / package_name 与配置是否一致。
//
// 为什么值得做:推送送不到设备的最隐蔽原因是**服务端应用与设备上装的包不是同一个**
// (包名/App ID 对不上),而症状只是"怎么都不来通知"。这里在启动时说清楚,
// 比事后拿着一堆 80300007 猜要便宜得多。只比对能对上的字段,格式不认识就跳过。
func checkClientConfig(cfg *ProviderConfig) {
path := strings.TrimSpace(cfg.ClientConfigFile)
if path == "" {
return
}
b, err := os.ReadFile(path)
if err != nil {
log.Printf("[push] 提醒:客户端配置文件 %s 读不到: %v", path, err)
return
}
var doc struct {
Client struct {
AppID string `json:"app_id"`
PackageName string `json:"package_name"`
} `json:"client"`
}
if err := json.Unmarshal(b, &doc); err != nil {
log.Printf("[push] 提醒:客户端配置文件 %s 不是可识别的 JSON(已跳过核对)", path)
return
}
if doc.Client.AppID != "" && cfg.AppID != "" && doc.Client.AppID != cfg.AppID {
log.Printf("[push] 不一致:客户端配置的 app_id=%s 与服务端配置的 app_id=%s 不是同一个应用 —— 推送送不到设备",
doc.Client.AppID, cfg.AppID)
}
if doc.Client.PackageName != "" {
log.Printf("[push] 客户端包名:%s(设备的包名必须与它一致,且签名指纹要在厂商后台登记过)", doc.Client.PackageName)
}
}
func nameOf(cfg ProviderConfig) string {
if n := strings.TrimSpace(cfg.Name); n != "" {
return n
}
return cfg.Type
}
// hmsConfigFromEnv 把 HMS_* 环境变量转成一条配置(可选覆盖,见 LoadProviders)。
func hmsConfigFromEnv() (ProviderConfig, bool) {
appID := strings.TrimSpace(os.Getenv("HMS_APP_ID"))
secret := strings.TrimSpace(os.Getenv("HMS_APP_SECRET"))
if appID == "" || secret == "" {
return ProviderConfig{}, false
}
cfg := ProviderConfig{Type: "hms", AppID: appID, AppSecret: secret}
if v := strings.TrimSpace(os.Getenv("HMS_TEST_MESSAGE")); v != "" {
b := v == "1" || strings.EqualFold(v, "true") || strings.EqualFold(v, "yes")
cfg.TestMessage = &b
}
if v := strings.TrimSpace(os.Getenv("HMS_DAILY_LIMIT")); v != "" {
var n int
if _, err := fmt.Sscanf(v, "%d", &n); err == nil && n >= 0 {
cfg.DailyLimit = n
}
}
if v := strings.TrimSpace(os.Getenv("HMS_CLIENT_CONFIG_FILE")); v != "" {
cfg.ClientConfigFile = v
}
return cfg, true
}

280
server/internal/push/hms.go Normal file
View File

@ -0,0 +1,280 @@
package push
import (
"bytes"
"context"
"encoding/json"
"fmt"
"io"
"log"
"net/http"
"net/url"
"strings"
"sync"
"time"
"github.com/agentmail/gateway/internal/repo"
)
/*
华为 HMS Push 通道(第一个实现,不是唯一实现)。
# 为什么这些常量长这样:都是**打真接口问出来的**,不是照着文档抄的
2026-09-15 用真凭证(个人开发者账号下的应用 `com.jianf.agentmail`)对线上服务打了
三种形状,用它自己的回答定下实现:
POST v1/{appId}/messages:send {message:{token:[…],notification:{…}}} + testMessage
→ {"code":"80300007","msg":"All the tokens are invalid"} ← 形状被接受,只是假 token 无效 ✓
POST v1/{appId}/messages:send {payload:{…},target:{token:[…]}}
→ {"code":"80300010","msg":"token count should within 1 and 1,000"} ← 这种形状 v1 不认账
POST v2/{appId}/messages:send {payload,target}
→ {"code":"80200001","msg":"Authentication Error"} ← v2 要另一种鉴权(服务账号 JWT)
所以走 v1 + `message.token[]`。v2 / 服务账号密钥那条路**没有**实现,也没验证过:
写进去就是拿没验过的形状冒充能用的代码。
# testMessage 默认开
未上架应用**必须**用测试消息模式才能收到推送(用户 2026-09-15 给的信息):
不开的话未上架应用的限制收紧到约 2 条/天/设备,调试期基本等于收不到。
额度是**项目级**的:1000 条/天,且单次推送最多 10 个 token —— 后面这条由
MaxTokensPerRequest 声明,分批由 push.dispatch 执行。
应用正式上架后要把它改成 false(`HMS_TEST_MESSAGE=false`),否则一直吃测试额度
且受测试消息的频控。
# 成功码
华为回的 `code == "80000000"` 表示成功。这个值来自推送 API 的约定,我**无法在本机
验证成功路径**(需要一台真机产出的 token);失败路径(上面那三个码)是实测的。
所以:成功判据只认 80000000,其余一律当失败并记下 code/msg —— 宁可把成功误判成
失败(记一条日志、少一条通知),也不能把失败当成功(那会静默丢通知且没人查)。
*/
type HMS struct {
// name 是推给客户端看的 provider 名(默认 "hms";同一实例接两套同厂商凭证时用得上)。
name string
AppID string
AppSecret string
// TestMessage 见包注释:未上架应用必须为 true。
TestMessage bool
// DailyLimit 是每日发送上限(条)。华为对未上架应用的测试消息限制是
// **项目级** 1000 条/天,默认按它兜底,避免把额度打光后收到一串失败。
DailyLimit int
// Endpoint 可覆盖,仅用于测试注入(默认走华为线上端点)。
Endpoint string
TokenURL string
Client *http.Client
baseDelay time.Duration
tokenMu sync.Mutex
token string
tokenExp time.Time
dayMu sync.Mutex
day string
dayCount int
}
const (
hmsDefaultEndpoint = "https://push-api.cloud.huawei.com"
hmsTokenURL = "https://oauth-login.cloud.huawei.com/oauth2/v3/token"
hmsMaxTokensPerReq = 10
hmsSuccessCode = "80000000"
hmsAllInvalidCode = "80300007"
)
// newHMSFromConfig 按一项配置建通道(凭证已由 config.go 解析好)。
//
// 没有凭证就**不在配置表里出现** —— 这是「推送可选」的落地点:
// 没配的实例根本不会走到这里,整条推送路径连一次查库都不会发生。
func newHMSFromConfig(cfg ProviderConfig) (Notifier, error) {
appID := strings.TrimSpace(cfg.AppID)
if appID == "" {
return nil, fmt.Errorf("缺 app_id")
}
secret := strings.TrimSpace(cfg.AppSecret)
if secret == "" {
return nil, fmt.Errorf("缺 app_secret(建议用 app_secret_file 指向密钥文件)")
}
name := strings.TrimSpace(cfg.Name)
if name == "" {
name = "hms"
}
test := true
if cfg.TestMessage != nil {
test = *cfg.TestMessage
}
limit := cfg.DailyLimit
if limit <= 0 {
limit = 1000
}
return &HMS{
name: name,
AppID: appID,
AppSecret: secret,
TestMessage: test,
DailyLimit: limit,
Endpoint: hmsDefaultEndpoint,
TokenURL: hmsTokenURL,
Client: &http.Client{Timeout: 15 * time.Second},
}, nil
}
func (h *HMS) Name() string {
if h.name != "" {
return h.name
}
return "hms"
}
// MaxTokensPerRequest 是华为的硬限额(单次推送 ≤10 个 token)。
func (h *HMS) MaxTokensPerRequest() int { return hmsMaxTokensPerReq }
// accessToken 取(并缓存)访问令牌。华为给的有效期是 3600 秒,刷新提前 5 分钟。
func (h *HMS) accessToken(ctx context.Context) (string, error) {
h.tokenMu.Lock()
defer h.tokenMu.Unlock()
if h.token != "" && time.Now().Before(h.tokenExp) {
return h.token, nil
}
form := url.Values{
"grant_type": {"client_credentials"},
"client_id": {h.AppID},
"client_secret": {h.AppSecret},
}
req, err := http.NewRequestWithContext(ctx, http.MethodPost, h.TokenURL, strings.NewReader(form.Encode()))
if err != nil {
return "", err
}
req.Header.Set("Content-Type", "application/x-www-form-urlencoded")
resp, err := h.Client.Do(req)
if err != nil {
return "", err
}
defer resp.Body.Close()
body, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<16))
if resp.StatusCode != http.StatusOK {
// 不把 body 原样吐进日志:它可能含 token 片段。
return "", fmt.Errorf("取 access_token 失败: HTTP %d", resp.StatusCode)
}
var out struct {
AccessToken string `json:"access_token"`
ExpiresIn int `json:"expires_in"`
}
if err := json.Unmarshal(body, &out); err != nil {
return "", fmt.Errorf("解析 access_token 响应失败: %w", err)
}
if out.AccessToken == "" {
return "", fmt.Errorf("access_token 为空")
}
ttl := out.ExpiresIn
if ttl <= 0 {
ttl = 3600
}
h.token = out.AccessToken
h.tokenExp = time.Now().Add(time.Duration(ttl)*time.Second - 5*time.Minute)
return h.token, nil
}
type hmsSendResponse struct {
Code string `json:"code"`
Msg string `json:"msg"`
// IllegalTokens 是华为回的无效率 token 列表(有才用)。字段名按官方响应约定,
// 我这边没有真机 token 因而**未能实测**;因此只在下述两种情况下才据它删表。
IllegalTokens []string `json:"illegal_tokens"`
}
// Send 向一批 token(≤10)投递一条通知。
func (h *HMS) Send(ctx context.Context, tokens []string, n NewMail) error {
if len(tokens) == 0 {
return nil
}
if !h.reserveDaily(len(tokens)) {
return fmt.Errorf("达到每日推送上限 %d 条(HMS_DAILY_LIMIT)", h.DailyLimit)
}
tok, err := h.accessToken(ctx)
if err != nil {
return err
}
data, _ := json.Marshal(map[string]string{
// 与客户端约定的形状(见文档与给 dsh 的契约):点通知按它跳转。
"type": "new_mail",
"mail_id": n.MailID,
"session_id": n.SessionID,
"action": "open_mail",
})
payload := map[string]any{
"validate_only": false,
"message": map[string]any{
"token": tokens,
"notification": map[string]any{
"title": "新邮件:" + truncate(n.Subject, 40),
"body": n.From,
},
// data 必须是**字符串**(华为这套要求 JSON 序列化后的字符串)。
"data": string(data),
},
}
if h.TestMessage {
payload["testMessage"] = true
}
raw, _ := json.Marshal(payload)
req, err := http.NewRequestWithContext(ctx, http.MethodPost,
strings.TrimRight(h.Endpoint, "/")+"/v1/"+url.PathEscape(h.AppID)+"/messages:send", bytes.NewReader(raw))
if err != nil {
return err
}
req.Header.Set("Authorization", "Bearer "+tok)
req.Header.Set("Content-Type", "application/json; charset=UTF-8")
resp, err := h.Client.Do(req)
if err != nil {
return err
}
defer resp.Body.Close()
body, _ := io.ReadAll(io.LimitReader(resp.Body, 1<<16))
var out hmsSendResponse
if err := json.Unmarshal(body, &out); err != nil {
return fmt.Errorf("解析推送响应失败: HTTP %d", resp.StatusCode)
}
if out.Code == hmsSuccessCode {
return nil
}
// 无效 token 自愈:设备卸了 App / token 轮换了。留着它们每次发信都白吃额度
// (测试消息额度是项目级的),所以按值删掉。
if out.Code == hmsAllInvalidCode {
if _, derr := repo.DeletePushTokensByValue(ctx, h.Name(), tokens); derr != nil {
log.Printf("[push] 清理无效 token 失败: %v", derr)
} else {
log.Printf("[push] 已清理 %d 个无效 token(%s)", len(tokens), out.Code)
}
} else if len(out.IllegalTokens) > 0 {
if _, derr := repo.DeletePushTokensByValue(ctx, h.Name(), out.IllegalTokens); derr != nil {
log.Printf("[push] 清理无效 token 失败: %v", derr)
}
}
return fmt.Errorf("华为推送失败: code=%s msg=%s", out.Code, out.Msg)
}
// reserveDaily 记一次每日用量。超上限时返回 false(不把额度打光:
// 打光之后的失败响应刷日志,而且真需要的那条也发不出去)。
func (h *HMS) reserveDaily(n int) bool {
h.dayMu.Lock()
defer h.dayMu.Unlock()
today := time.Now().UTC().Format("2006-01-02")
if h.day != today {
h.day, h.dayCount = today, 0
}
if h.DailyLimit > 0 && h.dayCount+n > h.DailyLimit {
return false
}
h.dayCount += n
return true
}
func truncate(s string, max int) string {
r := []rune(s)
if len(r) <= max {
return s
}
return string(r[:max]) + "…"
}

View File

@ -0,0 +1,179 @@
/*
Package push 是「新邮件」的第二条送达通道(第一条是 SSE)。
# 它是可选的,且默认关闭
自部署实例通常**一个推送渠道都没配** —— 那是正常状态,不是配置错误:客户端
在线时 SSE 已经够用,推送只解决「App 不在前台 / 被系统杀掉」这一种情形。
因此本包所有入口在没注册任何 provider 时都是**立即返回**:不查库、不建连接、
不刷日志(用户 2026-09-15 的明确要求:「不能写死推送方式,因为我们是自部署后端」
「即推送密钥应当是可选项」)。
# 为什么不写死华为
表与端点都带 `provider` 维度,Notifier 是接口:加一个通道(web push、别的厂商)
只加一个实现 + 一行注册,不动 schema、不动端点形状、不动调用方。
华为 HMS 只是第一个实现(internal/push/hms.go)。
# 为什么发送是异步且会丢
推送发生在**收信路径**上(notify.Recipients),而它已经在库事务之外的下发阶段:
一个慢的推送 HTTP 请求不能拖住邮件送达 —— 收信是主功能,推送是锦上添花。
因此:有界并发 + 超时 + 失败只记日志。**宁可丢一条通知,不可慢一封邮件**。
*/
package push
import (
"context"
"log"
"sync"
"time"
"github.com/agentmail/gateway/internal/repo"
)
// NewMail 是要推送的一封新邮件。
//
// 字段刻意少:推送只负责「通知栏那一行 + 点进去看哪封信」。正文不进通知,
// 否则锁屏上就会露出邮件内容(SSE 是给已解锁的在线客户端用的,两者隐私模型不同)。
type NewMail struct {
MailID string
SessionID string
// From 是发件方名字,用于通知标题。
From string
// Subject 是邮件主题。
Subject string
// Recipients 是**该收到这封信的人名**(= SSE 的收件判据:主收件人 + 抄送方)。
// 两条通道共用同一份名单,不各算一套。
Recipients []string
}
// Notifier 是一个推送通道。
type Notifier interface {
// Name 是 provider 标识,与 push_tokens.provider 的值一致(如 "hms")。
Name() string
// MaxTokensPerRequest 是单次请求能带的最大 token 数(厂商限额,如华为测试消息 ≤10)。
// 由通道自己声明,而不是调用方写死一个「10」—— 限额是通道的属性。
MaxTokensPerRequest() int
// Send 向一批 token 投递同一条通知。返回错误只用于**记日志**。
Send(ctx context.Context, tokens []string, n NewMail) error
}
const (
// maxInFlight 是在途推送任务上限。超了就丢掉这一轮(记日志),不排队:
// 排队的后果是通知在几十秒后集中弹出来,那比丢掉更糟。
maxInFlight = 4
// sendTimeout 是单个 provider 单批的超时。
sendTimeout = 10 * time.Second
)
var (
mu sync.RWMutex
providers []Notifier
slots = make(chan struct{}, maxInFlight)
)
// Register 注册一个推送通道。由 main 按配置调用 —— 没配就不注册。
func Register(n Notifier) {
if n == nil {
return
}
mu.Lock()
defer mu.Unlock()
providers = append(providers, n)
log.Printf("[push] 通道已启用: %s(单批最多 %d 个 token)", n.Name(), n.MaxTokensPerRequest())
}
// Enabled 报告是否配了任何推送通道。
//
// 端点据此回 `enabled`,客户端据此知道自己「登记了也可能收不到」——
// 而不是以为登记失败。
func Enabled() bool {
mu.RLock()
defer mu.RUnlock()
return len(providers) > 0
}
// Names 返回已启用的通道名(端点回给客户端看)。
func Names() []string {
mu.RLock()
defer mu.RUnlock()
out := make([]string, 0, len(providers))
for _, p := range providers {
out = append(out, p.Name())
}
return out
}
// shouldDispatch 报告这封邮件值不值得进推送管线:没配通道、或没有收件人 → 不值。
//
// 为什么单独成一个**纯函数**(不碰库、不起 goroutine):因为“没配通道时零开销”
// 这条判据必须能**同步**验证。2026-09-15 变异验证实测过:把判据写成“调 NotifyNewMail
// 后用 nil 库不 panic”,去掉本函数里的 Enabled() 后测试**仍然绿** —— 分发在
// goroutine 里跑,而 t.Cleanup 已经把真库装回去了,于是判据在错误的理由上通过。
// 纯函数没有这个<E8BF99>赛跑面。
func shouldDispatch(n NewMail) bool {
return Enabled() && len(n.Recipients) > 0
}
// NotifyNewMail 异步把一封新邮件推给收件方登记的设备。
//
// 调用方(notify.Recipients)**永远不因此拿到错误**:推送失败不该影响收信,
// 也不该让调用方写一半成功一半失败的处理逻辑。
func NotifyNewMail(ctx context.Context, n NewMail) {
if !shouldDispatch(n) {
return
}
select {
case slots <- struct{}{}:
default:
log.Printf("[push] 在途任务已达上限 %d,跳过本轮推送(可选通道,丢一条通知不影响收信)", maxInFlight)
return
}
go func() {
defer func() { <-slots }()
// 用 Background 而不是请求的 ctx:发信请求一旦返回,ctx 就被取消,
// 挂在它上面的推送会被立刻掐断(而收信方恰恰是那个已经离线的人)。
sendCtx, cancel := context.WithTimeout(context.Background(), sendTimeout)
defer cancel()
dispatch(sendCtx, n)
}()
}
// dispatch 按 provider 分组投递,每批不超过该通道声明的上限。
func dispatch(ctx context.Context, n NewMail) {
tokens, err := repo.ListPushTokensOfOwners(ctx, n.Recipients)
if err != nil {
log.Printf("[push] 读推送登记失败(不影响收信): %v", err)
return
}
if len(tokens) == 0 {
return
}
grouped := map[string][]string{}
for _, t := range tokens {
grouped[t.Provider] = append(grouped[t.Provider], t.Token)
}
mu.RLock()
list := append([]Notifier(nil), providers...)
mu.RUnlock()
for _, p := range list {
ts := grouped[p.Name()]
if len(ts) == 0 {
continue
}
batch := p.MaxTokensPerRequest()
if batch <= 0 {
batch = 1
}
for i := 0; i < len(ts); i += batch {
end := i + batch
if end > len(ts) {
end = len(ts)
}
if err := p.Send(ctx, ts[i:end], n); err != nil {
log.Printf("[push] %s 投递失败(不影响收信): %v", p.Name(), err)
}
}
}
}

View File

@ -0,0 +1,604 @@
package push
import (
"bytes"
"context"
"encoding/json"
"io"
"log"
"net/http"
"net/http/httptest"
"os"
"path/filepath"
"strings"
"sync"
"testing"
"github.com/agentmail/gateway/internal/db"
"github.com/agentmail/gateway/internal/repo"
)
/*
推送通道的判据(2026-09-15)。
用户对这件事的要求是**可选**:「不能写死推送方式,因为我们是自部署后端」
「即推送密钥应当是可选项」。所以第一组判据钉的不是"推得出去",而是
**没配凭证时它必须彻底不存在**(不查库、不占 goroutine、不刷日志)。
第二组钉华为那条路的具体形状 —— 那些常量是拿真凭证打真接口问出来的
(见 hms.go 的注释),判据把形状钉住,避免以后"顺手改一下"就悄悄失效。
*/
// ─── 夹具 ────────────────────────────────────────────────────────────────
type fakeNotifier struct {
name string
max int
mu sync.Mutex
calls [][]string
last NewMail
}
func (f *fakeNotifier) Name() string { return f.name }
func (f *fakeNotifier) MaxTokensPerRequest() int { return f.max }
func (f *fakeNotifier) Send(_ context.Context, tokens []string, n NewMail) error {
f.mu.Lock()
defer f.mu.Unlock()
f.calls = append(f.calls, append([]string(nil), tokens...))
f.last = n
return nil
}
func (f *fakeNotifier) sizes() []int {
f.mu.Lock()
defer f.mu.Unlock()
out := make([]int, 0, len(f.calls))
for _, c := range f.calls {
out = append(out, len(c))
}
return out
}
// withProviders 临时替换全局通道表(测试之间互不影响)。
func withProviders(t *testing.T, ps ...Notifier) {
t.Helper()
mu.Lock()
saved := providers
providers = nil
mu.Unlock()
for _, p := range ps {
Register(p)
}
t.Cleanup(func() {
mu.Lock()
providers = saved
mu.Unlock()
})
}
func setupPushDB(t *testing.T) {
t.Helper()
dir := t.TempDir()
if err := db.Connect(context.Background(), 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)
}
// ─── 一、可选性:没配凭证 = 彻底不存在 ──────────────────────────────────
// 没注册任何通道时,推送必须彻底不存在:不分发、不查库、不起 goroutine。
//
// 判据分两侧(只验一侧的判据是假的):
// - 坏样本:没通道时 shouldDispatch 必须为 false;
// - 干净样本:注了通道且有收件人时必须为 true —— 否则“永远返回 false”也能骗过上面一条。
//
// 另外真的拿 nil 库调一次 NotifyNewMail 做冒烟(任何 repo 调用都会 panic)。
// 它**不是**主判据:分发在 goroutine 里,与 t.Cleanup 赛跑 —— 2026-09-15 变异验证
// 实测过:只写这一点时,去掉 push.go 的 Enabled() 早退仍然绿(判据在错误的理由上通过)。
func TestNoProvidersTouchesNothing(t *testing.T) {
withProviders(t) // 一个都不注册
saved := db.DB
db.DB = nil
t.Cleanup(func() { db.DB = saved })
if Enabled() {
t.Fatal("一个通道都没注册时 Enabled() 必须是 false")
}
if names := Names(); len(names) != 0 {
t.Fatalf("没注册通道时不该有名字,实际 %v", names)
}
if shouldDispatch(NewMail{Recipients: []string{"alice"}}) {
t.Fatal("没配任何推送通道时分发必须被跳过(零开销)")
}
if shouldDispatch(NewMail{}) {
t.Fatal("没有收件人时不该分发")
}
// 冒烟:走完 NotifyNewMail 不该碰库
NotifyNewMail(context.Background(), NewMail{
MailID: "m1", SessionID: "s1", From: "bob", Subject: "你好",
Recipients: []string{"alice"},
})
NotifyNewMail(context.Background(), NewMail{MailID: "m2"})
// 干净样本:注了通道 + 有收件人 → 必须分发
withProviders(t, &fakeNotifier{name: "fake", max: 10})
if !shouldDispatch(NewMail{Recipients: []string{"alice"}}) {
t.Fatal("配了通道且有收件人时必须分发(否则这条判据挡不住“永远不分发”的实现)")
}
if shouldDispatch(NewMail{}) {
t.Fatal("注了通道但没有收件人时仍不该分发")
}
}
/*
配置式多厂商接入的判据(2026-09-15 用户的第二条要求)。
原话:「我们要支持多厂商配置式接入统一推送服务,推送密钥(文件)应当是配置项,
用户可为不同的推送服务配置对应的凭证,因为我们的设计中是每个用户各自部署服务器」。
所以这里钉四件事:
1. 配置能真的造出多个通道(而不是只有一个 HMS 硬编码路径);
2. **密钥文件**能作为配置项用(app_secret_file),且权限过松会提醒;
3. **单项配错不能拖垮服务**:未知类型 / 密钥读不到 / 被 disabled → 只跳过那一条;
4. 没配文件 = 没推送,且不是错误(自部署实例的默认形态)。
*/
func writePushConfig(t *testing.T, body string) string {
t.Helper()
path := filepath.Join(t.TempDir(), "push.json")
if err := os.WriteFile(path, []byte(body), 0o600); err != nil {
t.Fatal(err)
}
t.Setenv("PUSH_CONFIG", path)
return path
}
func writeSecret(t *testing.T, mode os.FileMode) string {
t.Helper()
p := filepath.Join(t.TempDir(), "hms.secret")
if err := os.WriteFile(p, []byte("sec-from-file\n"), mode); err != nil {
t.Fatal(err)
}
return p
}
// 未知的 type 只跳过自己;同实例可以同时接多个厂商。
func TestConfigFileDrivesMultipleProviders(t *testing.T) {
RegisterType("test-echo", func(cfg ProviderConfig) (Notifier, error) {
return &fakeNotifier{name: nameOf(cfg), max: 3}, nil
})
secret := writeSecret(t, 0o600)
writePushConfig(t, `{
"providers": [
{"type":"hms","app_id":"app-1","app_secret_file":`+jsonStr(secret)+`},
{"type":"test-echo","name":"echo-a"},
{"type":"vendor-not-implemented","app_id":"x"},
{"type":"hms","name":"hms-disabled","enabled":false,"app_id":"a","app_secret":"s"}
]
}`)
ps := LoadProviders()
names := map[string]Notifier{}
for _, p := range ps {
names[p.Name()] = p
}
if len(ps) != 2 {
t.Fatalf("应启用 2 个通道(hms + test-echo),实际 %d: %v", len(ps), names)
}
if _, ok := names["hms"]; !ok {
t.Fatalf("缺 hms 通道: %v", names)
}
if _, ok := names["echo-a"]; !ok {
t.Fatalf("Name 覆盖没生效(应叫 echo-a): %v", names)
}
// 密钥来自**文件**(配置项),而不是内联
h, _ := names["hms"].(*HMS)
if h == nil || h.AppSecret != "sec-from-file" {
t.Fatalf("app_secret_file 没被读进来: %+v", h)
}
// 缺省 = 测试消息开(未上架应用只有这个模式收得到)
if !h.TestMessage {
t.Fatal("test_message 缺省必须为 true(未上架应用)")
}
if h.DailyLimit != 1000 {
t.Fatalf("daily_limit 缺省应为 1000(华为测试消息的项目级限制),实际 %d", h.DailyLimit)
}
}
// 密钥文件权限过松要提醒(一个 0644 的密钥文件是真实的配置错误),但不拦启动。
func TestConfigWarnsOnLooseSecretFile(t *testing.T) {
secret := writeSecret(t, 0o644)
writePushConfig(t, `{"providers":[{"type":"hms","app_id":"a","app_secret_file":`+jsonStr(secret)+`}]}`)
var buf bytes.Buffer
log.SetOutput(&buf)
t.Cleanup(func() { log.SetOutput(os.Stderr) })
if ps := LoadProviders(); len(ps) != 1 {
t.Fatalf("权限过松只该提醒不该拒绝,实际 %d 个通道", len(ps))
}
log.SetOutput(os.Stderr)
if !strings.Contains(buf.String(), "权限") {
t.Fatalf("应提醒密钥文件权限过松,实际日志:%s", buf.String())
}
}
// 单项配错不能拖垮其他项:第一个密钥文件读不到,第二个仍必须启用。
func TestConfigBadEntryDoesNotKillOthers(t *testing.T) {
RegisterType("test-echo", func(cfg ProviderConfig) (Notifier, error) {
return &fakeNotifier{name: nameOf(cfg), max: 3}, nil
})
writePushConfig(t, `{
"providers": [
{"type":"hms","app_id":"a","app_secret_file":"/nonexistent/secret"},
{"type":"hms","name":"hms-ok","app_id":"b","app_secret":"inline"}
]
}`)
ps := LoadProviders()
if len(ps) != 1 || ps[0].Name() != "hms-ok" {
t.Fatalf("坏条目应只跳自己,实际 %v", ps)
}
}
// 没配配置文件 = 没推送,而且不是错误。
func TestNoConfigMeansNoProviders(t *testing.T) {
t.Setenv("PUSH_CONFIG", filepath.Join(t.TempDir(), "does-not-exist.json"))
if ps := LoadProviders(); len(ps) != 0 {
t.Fatalf("没配置文件时应 0 个通道,实际 %d", len(ps))
}
}
// 环境变量是**可选覆盖**:配置里没有同类型条目时补一条;已有则以配置为准。
func TestEnvIsOverrideNotTheMainPath(t *testing.T) {
t.Setenv("HMS_APP_ID", "env-app")
t.Setenv("HMS_APP_SECRET", "env-secret")
t.Setenv("HMS_TEST_MESSAGE", "false")
// 1) 配置文件里没有 hms 条目 → 用环境变量补一条
writePushConfig(t, `{"providers":[]}`)
ps := LoadProviders()
if len(ps) != 1 || ps[0].Name() != "hms" {
t.Fatalf("环境变量应能补出一条 hms 通道,实际 %v", ps)
}
if h := ps[0].(*HMS); h.TestMessage {
t.Fatal("HMS_TEST_MESSAGE=false 没生效")
}
// 2) 配置文件里已有 hms 条目 → 环境变量不得再补一条(避免两个同名通道)
writePushConfig(t, `{"providers":[{"type":"hms","name":"hms-from-file","app_id":"f","app_secret":"s"}]}`)
ps = LoadProviders()
if len(ps) != 1 || ps[0].Name() != "hms-from-file" {
t.Fatalf("配置文件优先,环境变量不该再加一条,实际 %v", ps)
}
}
// 密钥文件优先于内联密钥(用户明确要求「推送密钥(文件)应当是配置项」)。
//
// 为什么这条值得单独写:两者都给是**常见**情况(配置里留着旧的内联密钥做参考,
// 同时切到文件)。不明确优先关系的结果是「改了文件却没生效」这种最难查的静默失效。
// 第一版这条判据缺失,变异验证(把优先级反转)居然全绿 —— 因此补上。
func TestSecretFileBeatsInlineSecret(t *testing.T) {
secret := writeSecret(t, 0o600)
writePushConfig(t, `{"providers":[{"type":"hms","app_id":"a","app_secret":"inline-wrong","app_secret_file":`+jsonStr(secret)+`}]}`)
ps := LoadProviders()
if len(ps) != 1 {
t.Fatalf("应有 1 个通道,实际 %d", len(ps))
}
if h := ps[0].(*HMS); h.AppSecret != "sec-from-file" {
t.Fatalf("密钥文件应优先于内联密钥,实际用了 %q", h.AppSecret)
}
}
// 厂自定义文件表(files)也能当密钥配置项用:不是每个厂商都叫 app_secret。
func TestGenericFilesMapWorksAsSecretSource(t *testing.T) {
secret := writeSecret(t, 0o600)
writePushConfig(t, `{"providers":[{"type":"hms","app_id":"a","files":{"app_secret":`+jsonStr(secret)+`}}]}`)
ps := LoadProviders()
if len(ps) != 1 {
t.Fatalf("files.app_secret 应被当作密钥文件配置项,实际 %d 个通道", len(ps))
}
if h := ps[0].(*HMS); h.AppSecret != "sec-from-file" {
t.Fatalf("files 里的密钥文件没被读进来: %q", h.AppSecret)
}
}
// 缺密钥的厂商条目由**它自己的 Factory** 判失败(而不是通用配置层):
// 通用层要求密钥就会把 Web Push 这类不用 app_secret 的厂商误杀(判据抓出来过)。
func TestMissingSecretIsFactoryBusiness(t *testing.T) {
RegisterType("test-nosecret", func(cfg ProviderConfig) (Notifier, error) {
return &fakeNotifier{name: nameOf(cfg), max: 1}, nil
})
writePushConfig(t, `{"providers":[
{"type":"test-nosecret","name":"no-secret-needed"},
{"type":"hms","app_id":"a"}
]}`)
ps := LoadProviders()
if len(ps) != 1 || ps[0].Name() != "no-secret-needed" {
t.Fatalf("不需要密钥的厂商应能建起来;需要密钥而没给的 hms 应被自己的 Factory 判失败。实际 %v", ps)
}
}
func jsonStr(s string) string {
b, _ := json.Marshal(s)
return string(b)
}
// ─── 二、分发:按 provider 分组、按通道声明的上限分批 ────────────────────
func TestDispatchBatchesByProviderLimit(t *testing.T) {
setupPushDB(t)
ctx := context.Background()
// 23 个 token,通道声明的上限是 10 → 必须切成 10/10/3
for i := 0; i < 23; i++ {
if err := repo.UpsertPushToken(ctx, "fake", tokName(i), "alice", "", ""); err != nil {
t.Fatal(err)
}
}
f := &fakeNotifier{name: "fake", max: 10}
withProviders(t, f)
dispatch(ctx, NewMail{MailID: "m-1", SessionID: "s-1", From: "bob", Subject: "主题", Recipients: []string{"alice"}})
if got := f.sizes(); len(got) != 3 || got[0] != 10 || got[1] != 10 || got[2] != 3 {
t.Fatalf("分批不对:期望 [10 10 3],实际 %v(超限会被华为拒:token count should within 1 and 1,000;测试消息另限 10)", got)
}
if f.last.MailID != "m-1" || f.last.SessionID != "s-1" || f.last.From != "bob" {
t.Fatalf("推给通道的邮件内容不对: %+v", f.last)
}
}
func TestDispatchGroupsByProvider(t *testing.T) {
setupPushDB(t)
ctx := context.Background()
if err := repo.UpsertPushToken(ctx, "a", "t-a1", "alice", "", ""); err != nil {
t.Fatal(err)
}
if err := repo.UpsertPushToken(ctx, "b", "t-b1", "bob", "", ""); err != nil {
t.Fatal(err)
}
fa := &fakeNotifier{name: "a", max: 10}
fb := &fakeNotifier{name: "b", max: 10}
withProviders(t, fa, fb)
dispatch(ctx, NewMail{MailID: "m", Recipients: []string{"alice", "bob"}})
if len(fa.calls) != 1 || len(fa.calls[0]) != 1 || fa.calls[0][0] != "t-a1" {
t.Fatalf("通道 a 应只拿到自己的 token,实际 %v", fa.calls)
}
if len(fb.calls) != 1 || len(fb.calls[0]) != 1 || fb.calls[0][0] != "t-b1" {
t.Fatalf("通道 b 应只拿到自己的 token,实际 %v", fb.calls)
}
}
func TestDispatchSkipsRecipientsWithoutTokens(t *testing.T) {
setupPushDB(t)
f := &fakeNotifier{name: "fake", max: 10}
withProviders(t, f)
// 谁都没登记过:不该有任何发送,也不该报错
dispatch(context.Background(), NewMail{MailID: "m", Recipients: []string{"nobody"}})
if len(f.calls) != 0 {
t.Fatalf("没人登记过就不该发,实际 %v", f.calls)
}
}
func tokName(i int) string {
return "tok-" + strings.Repeat("0", 2-len(itoa(i))) + itoa(i)
}
func itoa(i int) string {
if i == 0 {
return "0"
}
var b []byte
for i > 0 {
b = append([]byte{byte('0' + i%10)}, b...)
i /= 10
}
return string(b)
}
// ─── 三、华为通道的具体形状 ─────────────────────────────────────────────
// hmsStub 起一个假的华为端点:/token 发令牌,其余路径收推送。
type hmsStub struct {
srv *httptest.Server
mu sync.Mutex
tokenReq int
pushReqs []map[string]any
authHdrs []string
code string
}
func newHMSStub(t *testing.T) *hmsStub {
t.Helper()
s := &hmsStub{code: "80000000"}
s.srv = httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
if r.URL.Path == "/token" {
s.mu.Lock()
s.tokenReq++
s.mu.Unlock()
w.Header().Set("Content-Type", "application/json")
_, _ = io.WriteString(w, `{"access_token":"tok-abc","expires_in":3600}`)
return
}
body, _ := io.ReadAll(r.Body)
var m map[string]any
_ = json.Unmarshal(body, &m)
s.mu.Lock()
s.pushReqs = append(s.pushReqs, m)
s.authHdrs = append(s.authHdrs, r.Header.Get("Authorization"))
code := s.code
s.mu.Unlock()
w.Header().Set("Content-Type", "application/json")
_, _ = io.WriteString(w, `{"code":"`+code+`","msg":"stub"}`)
}))
t.Cleanup(s.srv.Close)
return s
}
func (s *hmsStub) hms() *HMS {
return &HMS{
AppID: "app-1", AppSecret: "sec-1", TestMessage: true, DailyLimit: 1000,
Endpoint: s.srv.URL, TokenURL: s.srv.URL + "/token", Client: s.srv.Client(),
}
}
func (s *hmsStub) count() int {
s.mu.Lock()
defer s.mu.Unlock()
return len(s.pushReqs)
}
// 推送请求的形状:v1 端点 + message.token + notification + data(字符串) + testMessage。
// 形状不对时华为回的是参数类错误(实测 80300010),而这条判据把它钉死在本地。
func TestHMSSendShape(t *testing.T) {
s := newHMSStub(t)
err := s.hms().Send(context.Background(), []string{"tok-1", "tok-2"}, NewMail{
MailID: "mail-1", SessionID: "sess-1", From: "bob", Subject: "标题",
})
if err != nil {
t.Fatalf("成功码应被当成功: %v", err)
}
if s.count() != 1 {
t.Fatalf("应发出 1 个请求,实际 %d", s.count())
}
req := s.pushReqs[0]
if req["testMessage"] != true {
t.Fatalf("未上架应用必须带 testMessage=true,实际 %v", req["testMessage"])
}
msg, ok := req["message"].(map[string]any)
if !ok {
t.Fatalf("缺 message 字段: %v", req)
}
tokens, _ := msg["token"].([]any)
if len(tokens) != 2 || tokens[0] != "tok-1" || tokens[1] != "tok-2" {
t.Fatalf("token 列表不对: %v", msg["token"])
}
notif, _ := msg["notification"].(map[string]any)
if notif == nil || !strings.Contains(str(notif["title"]), "标题") || notif["body"] != "bob" {
t.Fatalf("通知内容不对: %v", notif)
}
// data 必须是**字符串**(华为这套要求序列化后的 JSON 字符串)
dataStr, ok := msg["data"].(string)
if !ok {
t.Fatalf("data 必须是字符串,实际 %T", msg["data"])
}
var data map[string]string
if err := json.Unmarshal([]byte(dataStr), &data); err != nil {
t.Fatalf("data 不是合法 JSON 字符串: %v", err)
}
if data["type"] != "new_mail" || data["mail_id"] != "mail-1" || data["session_id"] != "sess-1" || data["action"] != "open_mail" {
t.Fatalf("data 的跳转契约不对(客户端按它跳转): %v", data)
}
if s.authHdrs[0] != "Bearer tok-abc" {
t.Fatalf("Authorization 头不对: %q", s.authHdrs[0])
}
}
// 访问令牌要缓存:两次发送只换一次令牌(华为给 3600 秒,我们的实现提前 5 分钟刷新)。
func TestHMSCachesAccessToken(t *testing.T) {
s := newHMSStub(t)
h := s.hms()
for i := 0; i < 2; i++ {
if err := h.Send(context.Background(), []string{"tok-1"}, NewMail{MailID: "m"}); err != nil {
t.Fatal(err)
}
}
s.mu.Lock()
defer s.mu.Unlock()
if s.tokenReq != 1 {
t.Fatalf("两次发送应只换一次令牌,实际换了 %d 次", s.tokenReq)
}
}
func TestHMSTestMessageCanBeDisabled(t *testing.T) {
s := newHMSStub(t)
h := s.hms()
h.TestMessage = false // 上架之后
if err := h.Send(context.Background(), []string{"tok-1"}, NewMail{MailID: "m"}); err != nil {
t.Fatal(err)
}
if _, present := s.pushReqs[0]["testMessage"]; present {
t.Fatal("HMS_TEST_MESSAGE=false 时不该带 testMessage")
}
}
// 华为回「这些 token 全无效」时必须把它们从库里清掉:
// 否则每次发信都向死 token 发(白吃项目级额度),而用户永远收不到。
func TestHMSSendPrunesAllInvalidTokens(t *testing.T) {
setupPushDB(t)
ctx := context.Background()
if err := repo.UpsertPushToken(ctx, "hms", "dead-1", "alice", "", ""); err != nil {
t.Fatal(err)
}
if err := repo.UpsertPushToken(ctx, "hms", "dead-2", "alice", "", ""); err != nil {
t.Fatal(err)
}
s := newHMSStub(t)
s.code = "80300007" // 实测码:All the tokens are invalid
h := s.hms()
err := h.Send(ctx, []string{"dead-1", "dead-2"}, NewMail{MailID: "m"})
if err == nil {
t.Fatal("无效 token 必须报错(不能静默当成功)")
}
if !strings.Contains(err.Error(), "80300007") {
t.Fatalf("错误里应带上华为的 code,便于排查: %v", err)
}
left, _ := repo.ListPushTokensOf(ctx, "alice")
if len(left) != 0 {
t.Fatalf("无效 token 没被清理: %+v", left)
}
}
// 认证/参数类错误不该清 token(那是我们自己的问题,不是设备的问题)。
func TestHMSSendKeepsTokensOnOtherErrors(t *testing.T) {
setupPushDB(t)
ctx := context.Background()
if err := repo.UpsertPushToken(ctx, "hms", "good-1", "alice", "", ""); err != nil {
t.Fatal(err)
}
s := newHMSStub(t)
s.code = "80200001" // Authentication Error
if err := s.hms().Send(ctx, []string{"good-1"}, NewMail{MailID: "m"}); err == nil {
t.Fatal("认证失败必须报错")
}
left, _ := repo.ListPushTokensOf(ctx, "alice")
if len(left) != 1 {
t.Fatal("认证类错误不该删 token(设备是好的,错在我们)")
}
}
// 每日上限只是**兜底**(华为对未上架应用的测试消息限 1000 条/天/项目)。
// 到线就停手,不把额度打光换来一串失败响应。
func TestHMSDailyLimitStopsSending(t *testing.T) {
s := newHMSStub(t)
h := s.hms()
h.DailyLimit = 1
if err := h.Send(context.Background(), []string{"tok-1"}, NewMail{MailID: "m1"}); err != nil {
t.Fatal(err)
}
err := h.Send(context.Background(), []string{"tok-2"}, NewMail{MailID: "m2"})
if err == nil {
t.Fatal("超过每日上限必须拒绝发送")
}
if !strings.Contains(err.Error(), "上限") {
t.Fatalf("错误信息应说明是上限问题: %v", err)
}
if s.count() != 1 {
t.Fatalf("到线后不该再打网络,实际打了 %d 次", s.count())
}
}
func str(v any) string {
s, _ := v.(string)
return s
}

View File

@ -0,0 +1,166 @@
package repo
import (
"context"
"fmt"
"strings"
"time"
"github.com/agentmail/gateway/internal/db"
"github.com/google/uuid"
)
/*
设备推送 token 的读写(可选通道,见 internal/push)。
# 语义是「设备 + 注册者」
同一个 push token 只属于一个注册者:换人登录是**转移**,不是并存。这不是洁癖 ——
同一台设备上如果两个账号各存一份同一个 token,前一个人没退干净时,他的新邮件
会推到后一个人手里的那台设备上。
# 为什么没配推送也要能写
推送是自部署后端的可选项(用户 2026-09-15 明确要求)。所以登记 token 不依赖
「当前是否配了推送渠道」:端点照存,管理员之后把凭证配上就立刻生效,
不需要客户端重新登记一遍(客户端那时可能已经不在前台了)。
*/
// PushToken 是一台设备为某个注册者登记的推送地址。
type PushToken struct {
TokenID string
Provider string
Token string
OwnerName string
SessionID string
DeviceName string
}
// UpsertPushToken 登记/刷新一台设备的推送地址。
//
// 同一个 (provider, token) 重复登记时**改归属**并刷新 session/device_name:
// 客户端每次启动都会登记一次,若这里报错或插重复行,表会随启动次数膨胀。
func UpsertPushToken(ctx context.Context, provider, token, owner, sessionID, deviceName string) error {
_, err := db.DB.ExecContext(ctx, `
INSERT INTO push_tokens (token_id, provider, token, owner_name, session_id, device_name, created_at, updated_at)
VALUES ($1, $2, $3, $4, $5, $6, NOW(), NOW())
ON CONFLICT (provider, token) DO UPDATE
SET owner_name = $4, session_id = $5, device_name = $6, updated_at = NOW()`,
uuid.New().String(), provider, token, owner, sessionID, deviceName)
return err
}
// DeletePushToken 注销一个 token。只允许**注册者本人**注销:不能凭一个 token
// 字符串把别人的设备推送掐掉。
//
// 返回是否真的删到了行 —— 客户端拿它区分「已注销」与「本来就没登记」,
// 但两者对客户端都不是错误(推送是可选通道,注销失败不该弹提示)。
func DeletePushToken(ctx context.Context, provider, token, owner string) (bool, error) {
res, err := db.DB.ExecContext(ctx, `
DELETE FROM push_tokens WHERE provider = $1 AND token = $2 AND owner_name = $3`,
provider, token, owner)
if err != nil {
return false, err
}
n, err := res.RowsAffected()
if err != nil {
return false, err
}
return n > 0, nil
}
// ListPushTokensOf 列出某个注册者的全部推送地址(各 provider 都有)。
func ListPushTokensOf(ctx context.Context, owner string) ([]PushToken, error) {
rows, err := db.DB.QueryContext(ctx, `
SELECT token_id, provider, token, owner_name, session_id, device_name
FROM push_tokens WHERE owner_name = $1 ORDER BY provider, created_at`, owner)
if err != nil {
return nil, err
}
defer rows.Close()
return scanPushTokens(rows)
}
// ListPushTokensOfOwners 一次取回多个注册者的推送地址。
//
// 为什么批量:一封邮件可能同时推给收件人 + 若干抄送方,逐个查库就是每个参与方
// 一次往返(`new_mail` 的分发已经因为同类原因重排过一次:会话级字段不许每人查一次)。
// 参与方为空时直接返回,不拼 `IN ()` 这种非法 SQL。
func ListPushTokensOfOwners(ctx context.Context, owners []string) ([]PushToken, error) {
if len(owners) == 0 {
return nil, nil
}
ph := make([]string, len(owners))
args := make([]any, 0, len(owners))
for i, o := range owners {
ph[i] = fmt.Sprintf("$%d", i+1)
args = append(args, o)
}
rows, err := db.DB.QueryContext(ctx, `
SELECT token_id, provider, token, owner_name, session_id, device_name
FROM push_tokens WHERE owner_name IN (`+strings.Join(ph, ",")+`)
ORDER BY provider, created_at`, args...)
if err != nil {
return nil, err
}
defer rows.Close()
return scanPushTokens(rows)
}
// PruneStalePushTokens 删掉 90 天没刷新过的登记。
//
// 客户端只在启动/换 token 时登记,所以「长期不刷新」就等于「这台设备不再用了」
// (App 卸载、token 轮换、换机)。留着它们的代价是每次发信都向一批死 token 发推送,
// 而华为对测试消息的额度是**项目级**的(1000 条/天,未上架应用),死 token 会白吃额度。
func PruneStalePushTokens(ctx context.Context) (int64, error) {
// 截止时间在 Go 里算,不用 SQL 的日期运算:两种方言的写法不同
// (PG 是 INTERVAL,SQLite 没有),而库里存的就是 NOW() 写的文本
// "YYYY-MM-DD HH:MM:SS.ffffff"(见 db.go 的 now 注册),UTC 字符串比较即正确。
cutoff := time.Now().UTC().Add(-90 * 24 * time.Hour).Format("2006-01-02 15:04:05.000000")
res, err := db.DB.ExecContext(ctx,
`DELETE FROM push_tokens WHERE updated_at < $1`, cutoff)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
// DeletePushTokensByValue 按 token 值删除(不看注册者)。
//
// 只给**系统自愈**用:厂商回「这些 token 无效」(设备卸了 App、token 轮换)时,
// 留着它们的结果是每次发信都白吃额度(华为测试消息是**项目级** 1000 条/天),
// 而且用户那边永远收不到。用户主动注销走 DeletePushToken(要带 owner)。
func DeletePushTokensByValue(ctx context.Context, provider string, tokens []string) (int64, error) {
if len(tokens) == 0 {
return 0, nil
}
ph := make([]string, len(tokens))
args := make([]any, 0, len(tokens)+1)
args = append(args, provider)
for i, t := range tokens {
ph[i] = fmt.Sprintf("$%d", i+2)
args = append(args, t)
}
res, err := db.DB.ExecContext(ctx,
`DELETE FROM push_tokens WHERE provider = $1 AND token IN (`+strings.Join(ph, ",")+`)`, args...)
if err != nil {
return 0, err
}
return res.RowsAffected()
}
func scanPushTokens(rows interface {
Next() bool
Scan(...any) error
Err() error
}) ([]PushToken, error) {
var out []PushToken
for rows.Next() {
var t PushToken
if err := rows.Scan(&t.TokenID, &t.Provider, &t.Token, &t.OwnerName, &t.SessionID, &t.DeviceName); err != nil {
return nil, err
}
out = append(out, t)
}
return out, rows.Err()
}

View File

@ -0,0 +1,216 @@
package repo
import (
"context"
"testing"
"time"
"github.com/agentmail/gateway/internal/db"
)
/*
推送登记的判据(2026-09-15)。
推送是**可选通道**(用户要求「不能写死推送方式…推送密钥应当是可选项」),
所以这里钉住的不是"能不能推",而是登记本身的四个语义:
1. 重复登记是**刷新**,不是插入新行 —— 客户端每次启动都会登记,表不能随之膨胀;
2. 同一个 token 换人登录是**转移**(否则上一任用户的通知推到同一台设备上);
3. 注销要认**注册者**(不能凭一个 token 字符串掐掉别人的设备推送);
4. 无效 token 能按值清理(否则每次发信白吃华为的项目级额度)。
*/
func TestPushTokenUpsertIsRefreshNotInsert(t *testing.T) {
setupTestDB(t)
ctx := context.Background()
for i := 0; i < 3; i++ {
if err := UpsertPushToken(ctx, "hms", "tok-1", "alice", "s-1", "我的手机"); err != nil {
t.Fatalf("第 %d 次登记失败: %v", i+1, err)
}
}
got, err := ListPushTokensOf(ctx, "alice")
if err != nil {
t.Fatal(err)
}
if len(got) != 1 {
t.Fatalf("同一 token 登记 3 次应只有 1 行,实际 %d 行(表会随客户端启动次数膨胀)", len(got))
}
if got[0].Provider != "hms" || got[0].Token != "tok-1" || got[0].DeviceName != "我的手机" {
t.Fatalf("登记内容不对: %+v", got[0])
}
// 刷新要更新 session_id(客户端换了会话,点通知该回到新会话)
if err := UpsertPushToken(ctx, "hms", "tok-1", "alice", "s-2", "我的手机"); err != nil {
t.Fatal(err)
}
got, _ = ListPushTokensOf(ctx, "alice")
if len(got) != 1 || got[0].SessionID != "s-2" {
t.Fatalf("重复登记应刷新 session_id,实际 %+v", got)
}
}
func TestPushTokenIsTransferredOnRelogin(t *testing.T) {
setupTestDB(t)
ctx := context.Background()
if err := UpsertPushToken(ctx, "hms", "shared-device", "alice", "", "同一台手机"); err != nil {
t.Fatal(err)
}
// bob 在同一台设备上登录并登记同一个 token
if err := UpsertPushToken(ctx, "hms", "shared-device", "bob", "", "同一台手机"); err != nil {
t.Fatal(err)
}
aliceTokens, _ := ListPushTokensOf(ctx, "alice")
if len(aliceTokens) != 0 {
t.Fatalf("换人登录后 alice 不该还持有这台设备:%+v(否则 alice 的新邮件会推到 bob 手里的设备上)", aliceTokens)
}
bobTokens, _ := ListPushTokensOf(ctx, "bob")
if len(bobTokens) != 1 {
t.Fatalf("bob 应持有这台设备,实际 %+v", bobTokens)
}
}
func TestPushTokenPerProviderSameValueCoexist(t *testing.T) {
setupTestDB(t)
ctx := context.Background()
// 不同 provider 的 token 空间是独立的,同一个字符串不该互相覆盖。
if err := UpsertPushToken(ctx, "hms", "same-string", "alice", "", ""); err != nil {
t.Fatal(err)
}
if err := UpsertPushToken(ctx, "webpush", "same-string", "alice", "", ""); err != nil {
t.Fatal(err)
}
got, _ := ListPushTokensOf(ctx, "alice")
if len(got) != 2 {
t.Fatalf("两个 provider 的同名 token 应并存(provider 是维度的一部分),实际 %d 行", len(got))
}
}
func TestDeletePushTokenRequiresOwner(t *testing.T) {
setupTestDB(t)
ctx := context.Background()
if err := UpsertPushToken(ctx, "hms", "tok-1", "alice", "", ""); err != nil {
t.Fatal(err)
}
// bob 拿着同一个 token 字符串来注销:必须无效
deleted, err := DeletePushToken(ctx, "hms", "tok-1", "bob")
if err != nil {
t.Fatal(err)
}
if deleted {
t.Fatal("bob 不该能注销 alice 的设备(token 字符串不是凭证)")
}
if got, _ := ListPushTokensOf(ctx, "alice"); len(got) != 1 {
t.Fatal("alice 的登记被误删了")
}
deleted, err = DeletePushToken(ctx, "hms", "tok-1", "alice")
if err != nil || !deleted {
t.Fatalf("本人注销应成功,deleted=%v err=%v", deleted, err)
}
// 再注销一次:不是错误,只是 deleted=false
deleted, err = DeletePushToken(ctx, "hms", "tok-1", "alice")
if err != nil {
t.Fatal(err)
}
if deleted {
t.Fatal("第二次注销不该报告删到了行")
}
}
func TestListPushTokensOfOwnersBatches(t *testing.T) {
setupTestDB(t)
ctx := context.Background()
if err := UpsertPushToken(ctx, "hms", "t-a", "alice", "", ""); err != nil {
t.Fatal(err)
}
if err := UpsertPushToken(ctx, "hms", "t-b", "bob", "", ""); err != nil {
t.Fatal(err)
}
if err := UpsertPushToken(ctx, "hms", "t-c", "carol", "", ""); err != nil {
t.Fatal(err)
}
got, err := ListPushTokensOfOwners(ctx, []string{"alice", "bob"})
if err != nil {
t.Fatal(err)
}
if len(got) != 2 {
t.Fatalf("应取回 alice+bob 两个 token,实际 %d", len(got))
}
for _, tk := range got {
if tk.OwnerName == "carol" {
t.Fatal("不该取回未请求的 carol 的登记")
}
}
// 空名单不该拼出 `IN ()` 这种非法 SQL
if got, err := ListPushTokensOfOwners(ctx, nil); err != nil || got != nil {
t.Fatalf("空名单应直接返回 nil, nil,实际 %v / %v", got, err)
}
}
func TestDeletePushTokensByValueIgnoresOwner(t *testing.T) {
setupTestDB(t)
ctx := context.Background()
for _, o := range []string{"alice", "bob"} {
if err := UpsertPushToken(ctx, "hms", "dead-"+o, o, "", ""); err != nil {
t.Fatal(err)
}
}
n, err := DeletePushTokensByValue(ctx, "hms", []string{"dead-alice", "dead-bob", "not-there"})
if err != nil {
t.Fatal(err)
}
if n != 2 {
t.Fatalf("应删掉 2 行,实际 %d", n)
}
if left, _ := ListPushTokensOfOwners(ctx, []string{"alice", "bob"}); len(left) != 0 {
t.Fatalf("清理不干净: %+v", left)
}
// 别的 provider 不该被误删
if err := UpsertPushToken(ctx, "webpush", "dead-alice", "alice", "", ""); err != nil {
t.Fatal(err)
}
if n, _ := DeletePushTokensByValue(ctx, "hms", []string{"dead-alice"}); n != 0 {
t.Fatal("按 hms 清理时误删了 webpush 的登记")
}
}
func TestPruneStalePushTokens(t *testing.T) {
setupTestDB(t)
ctx := context.Background()
if err := UpsertPushToken(ctx, "hms", "fresh", "alice", "", ""); err != nil {
t.Fatal(err)
}
if err := UpsertPushToken(ctx, "hms", "stale", "bob", "", ""); err != nil {
t.Fatal(err)
}
// 把 stale 那条改老:UPDATE 直接写 91 天前的时刻(与 db.go 的 NOW() 同格式)
old := time.Now().UTC().Add(-91 * 24 * time.Hour).Format("2006-01-02 15:04:05.000000")
if _, err := db.DB.ExecContext(ctx, `UPDATE push_tokens SET updated_at = $1 WHERE token = 'stale'`, old); err != nil {
t.Fatal(err)
}
n, err := PruneStalePushTokens(ctx)
if err != nil {
t.Fatal(err)
}
if n != 1 {
t.Fatalf("应清掉 1 条过期登记,实际 %d", n)
}
left, _ := ListPushTokensOf(ctx, "alice")
if len(left) != 1 {
t.Fatal("新鲜的登记被误删了")
}
if gone, _ := ListPushTokensOf(ctx, "bob"); len(gone) != 0 {
t.Fatal("过期登记没被清掉")
}
}