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:
@ -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);
|
||||
|
||||
|
||||
@ -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);
|
||||
|
||||
|
||||
162
server/internal/handler/push.go
Normal file
162
server/internal/handler/push.go
Normal 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:])
|
||||
}
|
||||
189
server/internal/handler/push_test.go
Normal file
189
server/internal/handler/push_test.go
Normal 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)
|
||||
}
|
||||
}
|
||||
@ -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。
|
||||
|
||||
304
server/internal/push/config.go
Normal file
304
server/internal/push/config.go
Normal 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
280
server/internal/push/hms.go
Normal 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]) + "…"
|
||||
}
|
||||
179
server/internal/push/push.go
Normal file
179
server/internal/push/push.go
Normal 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)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
604
server/internal/push/push_test.go
Normal file
604
server/internal/push/push_test.go
Normal 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
|
||||
}
|
||||
166
server/internal/repo/push_tokens.go
Normal file
166
server/internal/repo/push_tokens.go
Normal 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()
|
||||
}
|
||||
216
server/internal/repo/push_tokens_test.go
Normal file
216
server/internal/repo/push_tokens_test.go
Normal 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("过期登记没被清掉")
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user