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("过期登记没被清掉") } }