Files
pangolin/server/internal/nodes/grpc_test.go
wangjia 721371d806 fix(server): 付费会话「活跃即续期」——根治 macOS 常驻隧道一天多掉线
根因:付费连接凭证硬编码 24h TTL(paidCredentialTTL),而客户端只在 _connect()
(用户/看门狗前台重连)才重签,服务端从不因流量续期。macOS sysext(root)隧道独立于
GUI app 常驻,用户关窗后 Dart 看门狗根本不运行 → 凭证 24h 到期、下次 agent 重注册
用「未过期」快照整表覆盖并重渲染 sing-box → REALITY 会话被剔除、永久黑洞。该逻辑
四端共用同一份 Dart,故为全端共性(macOS 最易现形)。

修法(方案 A,服务端、与客户端生命周期无关,一改修四端):
- ReportUsage 收到某 dp_uuid 有流量,若属付费套餐(!AdGate)即把其凭证 expires_at
  顶到 now+PaidCredentialTTL。活跃会话永不过期;免费凭证 TTL 编码日额度、绝不续期
  (否则击穿日限)。每报按 user 缓存一次 entitlement 查询。
- 新增 NodeStore.RenewCredential(纯 UPDATE,WHERE expires_at>now 不复活已过期会话)。
- 24h 提为 nodes.PaidCredentialTTL 单一真相源,httpapi 引用它消除漂移。
- 纯 DB 续期,无需再 push agent(现有 REALITY 用户仍在,只要 DB 行不过期,下次
  重注册快照仍含它)。可移植 SQL(? 占位 + Go 端算时间,无 MySQL 专属构造)。

测试:handler 层付费续期/免费不续期(mock);store 层真 SQLite 续期/不复活已过期。
go test ./... 全绿、go vet 干净。

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
Claude-Session: https://claude.ai/code/session_01FEVUXAbFT6bF1Qw27RHWoD
2026-09-05 19:34:58 +08:00

1155 lines
36 KiB
Go

package nodes_test
import (
"context"
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"encoding/pem"
"math/big"
"net"
"sync"
"testing"
"time"
"github.com/alicebob/miniredis/v2"
"github.com/redis/go-redis/v9"
"google.golang.org/grpc"
"google.golang.org/grpc/codes"
"google.golang.org/grpc/credentials"
"google.golang.org/grpc/status"
"google.golang.org/grpc/test/bufconn"
"github.com/wangjia/pangolin/server/internal/mtls"
"github.com/wangjia/pangolin/server/internal/nodes"
agentv1 "github.com/wangjia/pangolin/server/internal/pb/agentv1"
)
// ─── mock NodeStore ───────────────────────────────────────────────────────────
// mockNodeStore is an in-memory NodeStore used by gRPC integration tests.
type mockNodeStore struct {
mu sync.Mutex
nodeUUID string
configVer int64
usageAccum []mockUsageEntry
hourlyAccum []mockUsageEntry
usersByDpUUID map[string]int64
// devicesByDpUUID maps a per-device dp_uuid → {userID, deviceID}.
devicesByDpUUID map[string][2]int64
deviceUsageAccum []mockDeviceUsageEntry
lastSeenTouched []int64
// ent is returned by EntitlementForUser (nil = no subscription = free plan).
ent *nodes.Entitlement
// renewed logs RenewCredential calls (renew-on-activity assertions).
renewed []mockRenewEntry
}
type mockRenewEntry struct {
DpUUID string
ExpiresAt time.Time
}
type mockUsageEntry struct {
UserID int64
Date time.Time
BytesUp int64
BytesDown int64
Minutes int64
}
type mockDeviceUsageEntry struct {
UserID int64
DeviceID int64
Date time.Time
BytesUp int64
BytesDown int64
Minutes int64
}
func (m *mockNodeStore) NodeByUUID(_ context.Context, uuid string) (*nodes.NodeRow, error) {
if uuid == m.nodeUUID {
return &nodes.NodeRow{
ID: 1,
UUID: uuid,
Status: "up",
RealityPBK: "test-private-key",
RealitySNI: "www.example.com",
Endpoint: "1.2.3.4:443",
}, nil
}
return nil, nil
}
func (m *mockNodeStore) ConfigVersion(_ context.Context) (int64, error) {
return m.configVer, nil
}
func (m *mockNodeStore) ActiveNodeUUIDs(_ context.Context) ([]string, error) {
return []string{m.nodeUUID}, nil
}
func (m *mockNodeStore) CredentialsForNode(_ context.Context, _ string) ([]*agentv1.Credential, error) {
return nil, nil
}
func (m *mockNodeStore) UserIDByDpUUID(_ context.Context, dpUUID string) (int64, bool, error) {
if uid, ok := m.usersByDpUUID[dpUUID]; ok {
return uid, true, nil
}
return 0, false, nil
}
func (m *mockNodeStore) ListUp(_ context.Context) ([]*nodes.NodeRow, error) {
return []*nodes.NodeRow{{ID: 1, UUID: m.nodeUUID, Status: "up"}}, nil
}
func (m *mockNodeStore) EntitlementForUser(_ context.Context, _ int64) (*nodes.Entitlement, error) {
return m.ent, nil
}
func (m *mockNodeStore) PersistCredential(_ context.Context, _ int64, _ *agentv1.Credential, _ time.Time) error {
return nil
}
func (m *mockNodeStore) RenewCredential(_ context.Context, dpUUID string, newExpiresAt time.Time) error {
m.mu.Lock()
defer m.mu.Unlock()
m.renewed = append(m.renewed, mockRenewEntry{dpUUID, newExpiresAt})
return nil
}
func (m *mockNodeStore) renewLog() []mockRenewEntry {
m.mu.Lock()
defer m.mu.Unlock()
out := make([]mockRenewEntry, len(m.renewed))
copy(out, m.renewed)
return out
}
func (m *mockNodeStore) DeleteCredential(_ context.Context, _ int64, _ string) error {
return nil
}
func (m *mockNodeStore) NodesHoldingCredential(_ context.Context, _ string) ([]nodes.CredentialLocation, error) {
return nil, nil
}
func (m *mockNodeStore) AccumulateUsage(_ context.Context, userID int64, date time.Time,
bytesUp, bytesDown, minutes int64,
) error {
m.mu.Lock()
defer m.mu.Unlock()
m.usageAccum = append(m.usageAccum, mockUsageEntry{userID, date, bytesUp, bytesDown, minutes})
return nil
}
func (m *mockNodeStore) AccumulateHourly(_ context.Context, userID int64, hour time.Time,
bytesUp, bytesDown, minutes int64,
) error {
m.mu.Lock()
defer m.mu.Unlock()
m.hourlyAccum = append(m.hourlyAccum, mockUsageEntry{userID, hour, bytesUp, bytesDown, minutes})
return nil
}
func (m *mockNodeStore) AccumulateDeviceHourly(_ context.Context, _, _ int64, _ time.Time,
_, _, _ int64,
) error {
return nil
}
func (m *mockNodeStore) EnsureDeviceDpUUID(_ context.Context, _ int64, _ string) (string, int64, error) {
return "", 0, nil
}
func (m *mockNodeStore) AccountDayBytes(_ context.Context, _ int64, _ time.Time) (int64, error) {
return 0, nil
}
func (m *mockNodeStore) AccountDayMinutes(_ context.Context, _ int64, _ time.Time) (int, int, error) {
return 0, 0, nil
}
func (m *mockNodeStore) CountActiveDevices(_ context.Context, _ int64, _ time.Time) (int, error) {
return 0, nil
}
func (m *mockNodeStore) UserDeviceByDpUUID(_ context.Context, dpUUID string) (int64, int64, bool, error) {
if m.devicesByDpUUID != nil {
if ud, ok := m.devicesByDpUUID[dpUUID]; ok {
return ud[0], ud[1], true, nil
}
}
if uid, ok := m.usersByDpUUID[dpUUID]; ok {
return uid, 0, true, nil // legacy account credential, no device dimension
}
return 0, 0, false, nil
}
func (m *mockNodeStore) AccumulateDeviceUsage(_ context.Context, userID, deviceID int64, date time.Time,
bytesUp, bytesDown, minutes int64,
) error {
m.mu.Lock()
defer m.mu.Unlock()
m.deviceUsageAccum = append(m.deviceUsageAccum,
mockDeviceUsageEntry{userID, deviceID, date, bytesUp, bytesDown, minutes})
return nil
}
func (m *mockNodeStore) TouchDeviceLastSeen(_ context.Context, deviceID int64) error {
m.mu.Lock()
defer m.mu.Unlock()
m.lastSeenTouched = append(m.lastSeenTouched, deviceID)
return nil
}
func (m *mockNodeStore) deviceUsageLog() []mockDeviceUsageEntry {
m.mu.Lock()
defer m.mu.Unlock()
out := make([]mockDeviceUsageEntry, len(m.deviceUsageAccum))
copy(out, m.deviceUsageAccum)
return out
}
func (m *mockNodeStore) usageLog() []mockUsageEntry {
m.mu.Lock()
defer m.mu.Unlock()
out := make([]mockUsageEntry, len(m.usageAccum))
copy(out, m.usageAccum)
return out
}
// ─── test infrastructure helpers ─────────────────────────────────────────────
// makeTestServerCert generates a self-signed ECDSA P-256 TLS certificate
// for use as the gRPC server's transport credential.
func makeTestServerCert(t *testing.T) tls.Certificate {
t.Helper()
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatalf("gen server key: %v", err)
}
tmpl := &x509.Certificate{
SerialNumber: big.NewInt(1),
Subject: pkix.Name{CommonName: "pangolin-grpc-test"},
NotBefore: time.Now().Add(-time.Minute),
NotAfter: time.Now().Add(time.Hour),
KeyUsage: x509.KeyUsageDigitalSignature,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
BasicConstraintsValid: true,
}
certDER, err := x509.CreateCertificate(rand.Reader, tmpl, tmpl, &key.PublicKey, key)
if err != nil {
t.Fatalf("create server cert: %v", err)
}
keyDER, err := x509.MarshalECPrivateKey(key)
if err != nil {
t.Fatalf("marshal server key: %v", err)
}
certPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: certDER})
keyPEM := pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER})
cert, err := tls.X509KeyPair(certPEM, keyPEM)
if err != nil {
t.Fatalf("x509 keypair server: %v", err)
}
return cert
}
// genKeyAndCSR generates an ECDSA P-256 private key and a PKCS#10 CSR (both PEM).
func genKeyAndCSR(t *testing.T) (keyPEM, csrPEM []byte) {
t.Helper()
key, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
t.Fatalf("gen node key: %v", err)
}
keyDER, err := x509.MarshalECPrivateKey(key)
if err != nil {
t.Fatalf("marshal node key: %v", err)
}
keyPEM = pem.EncodeToMemory(&pem.Block{Type: "EC PRIVATE KEY", Bytes: keyDER})
tmpl := &x509.CertificateRequest{
Subject: pkix.Name{CommonName: "pangolin-node-pending"},
SignatureAlgorithm: x509.ECDSAWithSHA256,
}
csrDER, err := x509.CreateCertificateRequest(rand.Reader, tmpl, key)
if err != nil {
t.Fatalf("create CSR: %v", err)
}
csrPEM = pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE REQUEST", Bytes: csrDER})
return
}
// testServerBundle holds all components for a single test gRPC server instance.
type testServerBundle struct {
t *testing.T
lis *bufconn.Listener
srv *grpc.Server
hub *nodes.Hub
rdb *redis.Client
store *mockNodeStore
tokens *mtls.BootstrapTokenManager
load *nodes.LoadCache
}
// newTestServer creates a fully wired gRPC server (bufconn, miniredis, mTLS).
func newTestServer(t *testing.T, configVer int64, nodeUUID string) *testServerBundle {
t.Helper()
// miniredis
mr := miniredis.RunT(t)
rdb := redis.NewClient(&redis.Options{Addr: mr.Addr()})
t.Cleanup(func() { rdb.Close() })
// CA + CRL + token manager
dir := t.TempDir()
ca, err := mtls.NewCA(mtls.CAConfig{
KeyPath: dir + "/ca.key",
CertPath: dir + "/ca.crt",
})
if err != nil {
t.Fatalf("NewCA: %v", err)
}
crl := mtls.NewCRL(rdb, nil)
tokens := mtls.NewBootstrapTokenManager(rdb)
// Mock store
store := &mockNodeStore{
nodeUUID: nodeUUID,
configVer: configVer,
usersByDpUUID: map[string]int64{"dp-user1": 101},
}
// Hub + load cache
hub := nodes.NewHub(rdb)
load := nodes.NewLoadCache(rdb)
ctx, cancel := context.WithCancel(context.Background())
hub.Start(ctx)
t.Cleanup(cancel)
// Handler
handler := nodes.NewHandler(ca, tokens, hub, store, load)
// gRPC server TLS config
serverCert := makeTestServerCert(t)
tlsCfg := mtls.NewServerTLSConfig(ca, crl)
tlsCfg.Certificates = []tls.Certificate{serverCert}
// gRPC server
lis := bufconn.Listen(1 << 20)
srv := grpc.NewServer(
grpc.Creds(credentials.NewTLS(tlsCfg)),
grpc.ChainUnaryInterceptor(mtls.UnaryServerInterceptor()),
grpc.ChainStreamInterceptor(mtls.StreamServerInterceptor()),
)
agentv1.RegisterAgentServiceServer(srv, handler)
go func() { _ = srv.Serve(lis) }()
t.Cleanup(srv.Stop)
return &testServerBundle{
t: t, lis: lis, srv: srv, hub: hub,
rdb: rdb, store: store, tokens: tokens, load: load,
}
}
// bufDialer returns a context dialer that routes through the bufconn listener.
func bufDialer(lis *bufconn.Listener) func(context.Context, string) (net.Conn, error) {
return func(ctx context.Context, _ string) (net.Conn, error) {
return lis.DialContext(ctx)
}
}
// dialNoClientCert connects to the test server without presenting a client cert.
// Used for the Enroll RPC (no mTLS cert yet).
func dialNoClientCert(t *testing.T, lis *bufconn.Listener) *grpc.ClientConn {
t.Helper()
conn, err := grpc.NewClient("passthrough:///bufnet",
grpc.WithContextDialer(bufDialer(lis)),
grpc.WithTransportCredentials(credentials.NewTLS(&tls.Config{
InsecureSkipVerify: true, //nolint:gosec // test-only
})),
)
if err != nil {
t.Fatalf("dialNoClientCert: %v", err)
}
t.Cleanup(func() { conn.Close() })
return conn
}
// dialWithNodeCert connects to the test server presenting certPEM / keyPEM as
// the client certificate. Used after Enroll to authenticate as the enrolled node.
func dialWithNodeCert(t *testing.T, lis *bufconn.Listener, certPEM, keyPEM []byte) *grpc.ClientConn {
t.Helper()
nodeCert, err := tls.X509KeyPair(certPEM, keyPEM)
if err != nil {
t.Fatalf("X509KeyPair: %v", err)
}
conn, err := grpc.NewClient("passthrough:///bufnet",
grpc.WithContextDialer(bufDialer(lis)),
grpc.WithTransportCredentials(credentials.NewTLS(&tls.Config{
Certificates: []tls.Certificate{nodeCert},
InsecureSkipVerify: true, //nolint:gosec // test-only
})),
)
if err != nil {
t.Fatalf("dialWithNodeCert: %v", err)
}
t.Cleanup(func() { conn.Close() })
return conn
}
// eventually polls fn until it returns true or timeout elapses.
func eventually(t *testing.T, timeout time.Duration, fn func() bool, msg string) {
t.Helper()
deadline := time.Now().Add(timeout)
for time.Now().Before(deadline) {
if fn() {
return
}
time.Sleep(10 * time.Millisecond)
}
t.Fatalf("timeout waiting for: %s", msg)
}
// ─── Enroll tests ─────────────────────────────────────────────────────────────
// TestEnroll_Valid exercises the happy path: issue a bootstrap token, submit a
// valid CSR, and assert the response contains the node UUID and a valid cert.
func TestEnroll_Valid(t *testing.T) {
const testNodeUUID = "test-node-enroll-valid"
b := newTestServer(t, 1, testNodeUUID)
ctx := context.Background()
token, err := b.tokens.IssueToken(ctx, testNodeUUID)
if err != nil {
t.Fatalf("IssueToken: %v", err)
}
keyPEM, csrPEM := genKeyAndCSR(t)
_ = keyPEM
conn := dialNoClientCert(t, b.lis)
client := agentv1.NewAgentServiceClient(conn)
resp, err := client.Enroll(ctx, &agentv1.EnrollRequest{
BootstrapToken: token,
CSRPEM: csrPEM,
AgentVersion: "test-1.0",
})
if err != nil {
t.Fatalf("Enroll: %v", err)
}
if resp.NodeUUID != testNodeUUID {
t.Errorf("NodeUUID=%q, want %q", resp.NodeUUID, testNodeUUID)
}
if len(resp.CertPEM) == 0 {
t.Error("CertPEM is empty")
}
if len(resp.CAPEM) == 0 {
t.Error("CAPEM is empty")
}
if resp.NotAfterUnix == 0 {
t.Error("NotAfterUnix is 0")
}
// Verify the returned cert has the correct CN.
block, _ := pem.Decode(resp.CertPEM)
if block == nil {
t.Fatal("CertPEM has no PEM block")
}
cert, err := x509.ParseCertificate(block.Bytes)
if err != nil {
t.Fatalf("parse cert: %v", err)
}
if cert.Subject.CommonName != testNodeUUID {
t.Errorf("cert CN=%q, want %q", cert.Subject.CommonName, testNodeUUID)
}
}
// TestEnroll_BadToken verifies that an invalid bootstrap token returns
// Unauthenticated and that a second call with the same token also fails.
func TestEnroll_BadToken(t *testing.T) {
const testNodeUUID = "test-node-enroll-badtoken"
b := newTestServer(t, 1, testNodeUUID)
ctx := context.Background()
_, csrPEM := genKeyAndCSR(t)
conn := dialNoClientCert(t, b.lis)
client := agentv1.NewAgentServiceClient(conn)
_, err := client.Enroll(ctx, &agentv1.EnrollRequest{
BootstrapToken: "wrong-token-xxxxxxxx",
CSRPEM: csrPEM,
})
if err == nil {
t.Fatal("expected Unauthenticated error, got nil")
}
if st, _ := status.FromError(err); st.Code() != codes.Unauthenticated {
t.Errorf("error code = %v, want Unauthenticated", st.Code())
}
}
// TestEnroll_TokenOneShot verifies that a bootstrap token can only be used once.
func TestEnroll_TokenOneShot(t *testing.T) {
const testNodeUUID = "test-node-oneshot"
b := newTestServer(t, 1, testNodeUUID)
ctx := context.Background()
token, _ := b.tokens.IssueToken(ctx, testNodeUUID)
keyPEM, csrPEM := genKeyAndCSR(t)
_ = keyPEM
conn := dialNoClientCert(t, b.lis)
client := agentv1.NewAgentServiceClient(conn)
// First call succeeds.
if _, err := client.Enroll(ctx, &agentv1.EnrollRequest{
BootstrapToken: token, CSRPEM: csrPEM,
}); err != nil {
t.Fatalf("first Enroll: %v", err)
}
// Second call must fail (token consumed).
_, err := client.Enroll(ctx, &agentv1.EnrollRequest{
BootstrapToken: token, CSRPEM: csrPEM,
})
if err == nil {
t.Fatal("second Enroll: expected error, got nil")
}
if st, _ := status.FromError(err); st.Code() != codes.Unauthenticated {
t.Errorf("second Enroll code = %v, want Unauthenticated", st.Code())
}
}
// ─── helpers: full flow setup ─────────────────────────────────────────────────
// enrollNode runs the full Enroll flow and returns the gRPC connection that
// presents the issued node cert (ready for Register/Heartbeat/Subscribe/Ack).
func enrollNode(t *testing.T, b *testServerBundle, nodeUUID string) (certPEM, keyPEM []byte, authedConn *grpc.ClientConn) {
t.Helper()
ctx := context.Background()
token, err := b.tokens.IssueToken(ctx, nodeUUID)
if err != nil {
t.Fatalf("IssueToken: %v", err)
}
keyPEM, csrPEM := genKeyAndCSR(t)
conn := dialNoClientCert(t, b.lis)
client := agentv1.NewAgentServiceClient(conn)
resp, err := client.Enroll(ctx, &agentv1.EnrollRequest{
BootstrapToken: token,
CSRPEM: csrPEM,
})
if err != nil {
t.Fatalf("Enroll: %v", err)
}
conn.Close()
authedConn = dialWithNodeCert(t, b.lis, resp.CertPEM, keyPEM)
return resp.CertPEM, keyPEM, authedConn
}
// ─── Register tests ───────────────────────────────────────────────────────────
// TestRegister_OK verifies Register returns a valid ConfigSnapshot for a known node.
func TestRegister_OK(t *testing.T) {
const nodeUUID = "test-node-register"
b := newTestServer(t, 7, nodeUUID)
ctx := context.Background()
_, _, conn := enrollNode(t, b, nodeUUID)
client := agentv1.NewAgentServiceClient(conn)
snap, err := client.Register(ctx, &agentv1.RegisterRequest{
NodeUUID: nodeUUID,
AgentVersion: "test-1.0",
})
if err != nil {
t.Fatalf("Register: %v", err)
}
if snap.ConfigVersion != 7 {
t.Errorf("ConfigVersion=%d, want 7", snap.ConfigVersion)
}
// The mock store returns a node with RealityPBK set (endpoint 1.2.3.4:443,
// SNI www.example.com). The snapshot must carry a renderable inbound: a
// non-zero listen port (parsed from endpoint) and a handshake target, or the
// agent's sing-box config is invalid.
if snap.Reality == nil {
t.Fatal("Reality is nil, want non-nil")
}
if snap.Reality.ListenPort != 443 {
t.Errorf("Reality.ListenPort=%d, want 443 (from endpoint)", snap.Reality.ListenPort)
}
if snap.Reality.ServerName != "www.example.com" {
t.Errorf("Reality.ServerName=%q, want www.example.com", snap.Reality.ServerName)
}
if snap.Reality.HandshakeServer != "www.example.com" {
t.Errorf("Reality.HandshakeServer=%q, want www.example.com", snap.Reality.HandshakeServer)
}
if snap.Reality.HandshakePort != 443 {
t.Errorf("Reality.HandshakePort=%d, want 443", snap.Reality.HandshakePort)
}
if snap.Reality.PrivateKey == "" {
t.Error("Reality.PrivateKey is empty, want the node's REALITY key")
}
}
// ─── Heartbeat tests (includes Redis node:load assertion) ────────────────────
// TestHeartbeat_LoadWritten verifies that a Heartbeat call writes to the Redis
// node:load:{nodeUUID} hash with the reported metrics.
func TestHeartbeat_LoadWritten(t *testing.T) {
const nodeUUID = "test-node-heartbeat"
b := newTestServer(t, 5, nodeUUID)
ctx := context.Background()
_, _, conn := enrollNode(t, b, nodeUUID)
client := agentv1.NewAgentServiceClient(conn)
// Register first (as the agent would do in practice).
if _, err := client.Register(ctx, &agentv1.RegisterRequest{
NodeUUID: nodeUUID, AgentVersion: "1.0",
}); err != nil {
t.Fatalf("Register: %v", err)
}
// Send a heartbeat with known metrics.
resp, err := client.Heartbeat(ctx, &agentv1.HeartbeatRequest{
NodeUUID: nodeUUID,
ConfigVersion: 5,
OnlinePeers: 42,
BandwidthUpBps: 1_000_000,
BandwidthDownBps: 5_000_000,
CPUPercent: 12.5,
TimestampUnix: time.Now().Unix(),
})
if err != nil {
t.Fatalf("Heartbeat: %v", err)
}
if resp.ServerTimeUnix == 0 {
t.Error("ServerTimeUnix is 0")
}
if resp.NeedFullResync {
t.Error("NeedFullResync=true for matching config version, want false")
}
// Assert Redis hash was written.
eventually(t, time.Second, func() bool {
l, ok, _ := b.load.Get(ctx, nodeUUID)
return ok && l != nil && l.OnlineCount == 42
}, "node:load hash to contain online_count=42")
l, _, _ := b.load.Get(ctx, nodeUUID)
if l.BandwidthUpBps != 1_000_000 {
t.Errorf("bw_up=%d, want 1000000", l.BandwidthUpBps)
}
if l.BandwidthDownBps != 5_000_000 {
t.Errorf("bw_down=%d, want 5000000", l.BandwidthDownBps)
}
}
// TestHeartbeat_ResyncFlag verifies that a Heartbeat with a stale config_version
// triggers NeedFullResync = true.
func TestHeartbeat_ResyncFlag(t *testing.T) {
const nodeUUID = "test-node-resync"
b := newTestServer(t, 10, nodeUUID) // server is at config version 10
ctx := context.Background()
_, _, conn := enrollNode(t, b, nodeUUID)
client := agentv1.NewAgentServiceClient(conn)
if _, err := client.Register(ctx, &agentv1.RegisterRequest{NodeUUID: nodeUUID}); err != nil {
t.Fatalf("Register: %v", err)
}
resp, err := client.Heartbeat(ctx, &agentv1.HeartbeatRequest{
NodeUUID: nodeUUID,
ConfigVersion: 5, // stale: server is at 10
TimestampUnix: time.Now().Unix(),
})
if err != nil {
t.Fatalf("Heartbeat: %v", err)
}
if !resp.NeedFullResync {
t.Error("NeedFullResync=false for stale config_version, want true")
}
}
// ─── Subscribe / Ack tests ────────────────────────────────────────────────────
// TestSubscribe_PushAck exercises the full Subscribe→Push→receive→Ack cycle and
// verifies the command is removed from the Redis queue after Ack.
func TestSubscribe_PushAck(t *testing.T) {
const nodeUUID = "test-node-subscribe"
b := newTestServer(t, 1, nodeUUID)
ctx, cancel := context.WithCancel(context.Background())
defer cancel()
_, _, conn := enrollNode(t, b, nodeUUID)
client := agentv1.NewAgentServiceClient(conn)
if _, err := client.Register(ctx, &agentv1.RegisterRequest{NodeUUID: nodeUUID}); err != nil {
t.Fatalf("Register: %v", err)
}
// Open Subscribe stream.
stream, err := client.Subscribe(ctx, &agentv1.SubscribeRequest{
NodeUUID: nodeUUID,
LastCommandID: 0,
})
if err != nil {
t.Fatalf("Subscribe: %v", err)
}
// Push a command via the hub.
pushed := &agentv1.Command{
Type: agentv1.CommandTypeUpsert,
Upsert: &agentv1.UpsertPayload{Credential: &agentv1.Credential{DpUUID: "dp-sub"}},
}
if err := b.hub.Push(ctx, nodeUUID, pushed); err != nil {
t.Fatalf("hub.Push: %v", err)
}
if pushed.CommandID == 0 {
t.Fatal("pushed.CommandID == 0: hub did not assign an ID")
}
// Receive the command.
received, err := stream.Recv()
if err != nil {
t.Fatalf("stream.Recv: %v", err)
}
if received.CommandID != pushed.CommandID {
t.Errorf("received command_id=%d, want %d", received.CommandID, pushed.CommandID)
}
if received.Upsert.Credential.DpUUID != "dp-sub" {
t.Errorf("dp_uuid=%q, want dp-sub", received.Upsert.Credential.DpUUID)
}
// Ack the command.
if _, err := client.Ack(ctx, &agentv1.AckRequest{
NodeUUID: nodeUUID,
CommandID: received.CommandID,
}); err != nil {
t.Fatalf("Ack: %v", err)
}
// Verify the command was removed from the Redis queue.
eventually(t, time.Second, func() bool {
cmds, _ := b.hub.Replay(context.Background(), nodeUUID, 0)
return len(cmds) == 0
}, "command queue to be empty after Ack")
}
// TestSubscribe_ReconnectResume verifies the at-least-once resume invariant:
// 1. Push commands A and B; agent receives and acks A.
// 2. Drop the stream before ack of B.
// 3. Push command C while offline.
// 4. Reconnect with last_command_id = A.CommandID.
// 5. Assert: B and C are replayed; A is NOT replayed.
func TestSubscribe_ReconnectResume(t *testing.T) {
const nodeUUID = "test-node-reconnect"
b := newTestServer(t, 1, nodeUUID)
bgCtx := context.Background()
// Helper: enroll and return a connection with the node cert.
enrollRaw := func() (certPEM, keyPEM []byte) {
tok, err := b.tokens.IssueToken(bgCtx, nodeUUID)
if err != nil {
t.Fatalf("IssueToken: %v", err)
}
kp, csr := genKeyAndCSR(t)
c0 := dialNoClientCert(t, b.lis)
resp, err := agentv1.NewAgentServiceClient(c0).Enroll(bgCtx,
&agentv1.EnrollRequest{BootstrapToken: tok, CSRPEM: csr})
if err != nil {
t.Fatalf("Enroll: %v", err)
}
c0.Close()
return resp.CertPEM, kp
}
cert1, key1 := enrollRaw()
// ─── session 1 ─────────────────────────────────────────────────────────
ctx1, cancel1 := context.WithCancel(bgCtx)
conn1 := dialWithNodeCert(t, b.lis, cert1, key1)
defer conn1.Close()
client1 := agentv1.NewAgentServiceClient(conn1)
if _, err := client1.Register(ctx1, &agentv1.RegisterRequest{NodeUUID: nodeUUID}); err != nil {
t.Fatalf("Register session1: %v", err)
}
stream1, err := client1.Subscribe(ctx1, &agentv1.SubscribeRequest{
NodeUUID: nodeUUID,
LastCommandID: 0,
})
if err != nil {
t.Fatalf("Subscribe session1: %v", err)
}
// Push command A.
cmdA := &agentv1.Command{
Type: agentv1.CommandTypeUpsert,
Upsert: &agentv1.UpsertPayload{Credential: &agentv1.Credential{DpUUID: "dp-A"}},
}
if err := b.hub.Push(bgCtx, nodeUUID, cmdA); err != nil {
t.Fatalf("push A: %v", err)
}
// Receive A.
recvA, err := stream1.Recv()
if err != nil {
t.Fatalf("recv A: %v", err)
}
if recvA.Upsert.Credential.DpUUID != "dp-A" {
t.Errorf("recv A dp_uuid=%q, want dp-A", recvA.Upsert.Credential.DpUUID)
}
// Ack A.
if _, err := client1.Ack(bgCtx, &agentv1.AckRequest{
NodeUUID: nodeUUID, CommandID: recvA.CommandID,
}); err != nil {
t.Fatalf("ack A: %v", err)
}
// Push command B.
cmdB := &agentv1.Command{
Type: agentv1.CommandTypeUpsert,
Upsert: &agentv1.UpsertPayload{Credential: &agentv1.Credential{DpUUID: "dp-B"}},
}
if err := b.hub.Push(bgCtx, nodeUUID, cmdB); err != nil {
t.Fatalf("push B: %v", err)
}
// Receive B (but don't ack — simulate drop before ack).
recvB, err := stream1.Recv()
if err != nil {
t.Fatalf("recv B: %v", err)
}
if recvB.Upsert.Credential.DpUUID != "dp-B" {
t.Errorf("recv B dp_uuid=%q, want dp-B", recvB.Upsert.Credential.DpUUID)
}
// Drop stream by cancelling session 1's context.
cancel1()
time.Sleep(50 * time.Millisecond) // let server observe the cancellation
// Push command C while node is "offline" (no active stream).
cmdC := &agentv1.Command{
Type: agentv1.CommandTypeRevoke,
Revoke: &agentv1.RevokePayload{DpUUID: "dp-A"},
}
if err := b.hub.Push(bgCtx, nodeUUID, cmdC); err != nil {
t.Fatalf("push C: %v", err)
}
// ─── session 2: reconnect ──────────────────────────────────────────────
// Enroll a second time with a fresh cert (each cert is one-time use for enroll).
cert2, key2 := enrollRaw()
ctx2, cancel2 := context.WithCancel(bgCtx)
defer cancel2()
conn2 := dialWithNodeCert(t, b.lis, cert2, key2)
defer conn2.Close()
client2 := agentv1.NewAgentServiceClient(conn2)
if _, err := client2.Register(ctx2, &agentv1.RegisterRequest{NodeUUID: nodeUUID}); err != nil {
t.Fatalf("Register session2: %v", err)
}
// Reconnect with last_command_id = A.CommandID (A was acked, B/C were not).
stream2, err := client2.Subscribe(ctx2, &agentv1.SubscribeRequest{
NodeUUID: nodeUUID,
LastCommandID: recvA.CommandID,
})
if err != nil {
t.Fatalf("Subscribe session2: %v", err)
}
// Collect resumed commands (expect exactly B and C).
type recvResult struct {
cmd *agentv1.Command
err error
}
recvCh := make(chan recvResult, 3)
go func() {
for {
cmd, err := stream2.Recv()
recvCh <- recvResult{cmd, err}
if err != nil {
return
}
}
}()
var resumed []*agentv1.Command
timeout := time.After(3 * time.Second)
for len(resumed) < 2 {
select {
case r := <-recvCh:
if r.err != nil {
t.Fatalf("stream2.Recv: %v", r.err)
}
resumed = append(resumed, r.cmd)
case <-timeout:
t.Fatalf("timeout: got %d resumed commands, want 2", len(resumed))
}
}
// Resumed commands must be B and C in order.
if resumed[0].Upsert.Credential.DpUUID != "dp-B" {
t.Errorf("resumed[0] dp_uuid=%q, want dp-B", resumed[0].Upsert.Credential.DpUUID)
}
if resumed[1].Revoke.DpUUID != "dp-A" {
t.Errorf("resumed[1] revoke dp_uuid=%q, want dp-A", resumed[1].Revoke.DpUUID)
}
// Command A must NOT be in the resumed batch (it was acked).
for _, c := range resumed {
if c.CommandID == recvA.CommandID {
t.Errorf("command A (id=%d) was re-delivered, but it was acked", recvA.CommandID)
}
}
}
// ─── ReportUsage test ─────────────────────────────────────────────────────────
// TestReportUsage_Accumulates verifies that ReportUsage resolves dp_uuid → user_id
// and records usage in the mock store.
func TestReportUsage_Accumulates(t *testing.T) {
const nodeUUID = "test-node-usage"
b := newTestServer(t, 1, nodeUUID)
ctx := context.Background()
_, _, conn := enrollNode(t, b, nodeUUID)
client := agentv1.NewAgentServiceClient(conn)
if _, err := client.Register(ctx, &agentv1.RegisterRequest{NodeUUID: nodeUUID}); err != nil {
t.Fatalf("Register: %v", err)
}
now := time.Now()
_, err := client.ReportUsage(ctx, &agentv1.UsageReport{
NodeUUID: nodeUUID,
WindowStartUnix: now.Add(-time.Minute).Unix(),
WindowEndUnix: now.Unix(),
Entries: []*agentv1.UsageEntry{
{DpUUID: "dp-user1", BytesUp: 100, BytesDown: 200, SessionMinutes: 3},
{DpUUID: "dp-unknown", BytesUp: 50, BytesDown: 50, SessionMinutes: 1}, // skipped
},
})
if err != nil {
t.Fatalf("ReportUsage: %v", err)
}
log := b.store.usageLog()
if len(log) != 1 {
t.Fatalf("usage log has %d entries, want 1 (dp-unknown skipped)", len(log))
}
if log[0].UserID != 101 {
t.Errorf("user_id=%d, want 101", log[0].UserID)
}
if log[0].BytesUp != 100 || log[0].BytesDown != 200 {
t.Errorf("bytes_up=%d bytes_down=%d, want 100/200",
log[0].BytesUp, log[0].BytesDown)
}
// Account-level dp_uuid → no per-device attribution (deviceID 0).
if dev := b.store.deviceUsageLog(); len(dev) != 0 {
t.Errorf("device usage should be empty for account credential, got %d", len(dev))
}
}
// TestReportUsage_AccountMinutesDedupPerUser verifies the account-level minutes are
// deduped to WALL-CLOCK per user within a window: two devices of the same user active
// in one window count as 1 minute (not 2), while bytes still sum and each device keeps
// its own minute. Guards the多设备并发超计 fix.
func TestReportUsage_AccountMinutesDedupPerUser(t *testing.T) {
const nodeUUID = "test-node-usage-dedup"
b := newTestServer(t, 1, nodeUUID)
ctx := context.Background()
// Two devices of user 101 (dev 55 & 66) + one device of user 202 (dev 77).
b.store.devicesByDpUUID = map[string][2]int64{
"dp-u1-d1": {101, 55},
"dp-u1-d2": {101, 66},
"dp-u2-d1": {202, 77},
}
_, _, conn := enrollNode(t, b, nodeUUID)
client := agentv1.NewAgentServiceClient(conn)
if _, err := client.Register(ctx, &agentv1.RegisterRequest{NodeUUID: nodeUUID}); err != nil {
t.Fatalf("Register: %v", err)
}
now := time.Now()
if _, err := client.ReportUsage(ctx, &agentv1.UsageReport{
NodeUUID: nodeUUID,
WindowStartUnix: now.Add(-time.Minute).Unix(),
WindowEndUnix: now.Unix(),
Entries: []*agentv1.UsageEntry{
{DpUUID: "dp-u1-d1", BytesUp: 100, BytesDown: 200, SessionMinutes: 1},
{DpUUID: "dp-u1-d2", BytesUp: 10, BytesDown: 20, SessionMinutes: 1},
{DpUUID: "dp-u2-d1", BytesUp: 5, BytesDown: 7, SessionMinutes: 1},
},
}); err != nil {
t.Fatalf("ReportUsage: %v", err)
}
// Account rollup: one entry per user; user 101 minutes deduped to 1 (not 2),
// bytes summed across its two devices; user 202 = 1 minute.
byUser := map[int64]mockUsageEntry{}
for _, e := range b.store.usageLog() {
byUser[e.UserID] = e
}
if len(byUser) != 2 {
t.Fatalf("account rollup should have 2 users, got %d", len(byUser))
}
if u := byUser[101]; u.Minutes != 1 || u.BytesUp != 110 || u.BytesDown != 220 {
t.Errorf("user101 account: minutes=%d bytesUp=%d bytesDown=%d, want 1/110/220 (墙上时钟去重)",
u.Minutes, u.BytesUp, u.BytesDown)
}
if u := byUser[202]; u.Minutes != 1 || u.BytesUp != 5 {
t.Errorf("user202 account: minutes=%d bytesUp=%d, want 1/5", u.Minutes, u.BytesUp)
}
// Per-device: each of user101's two devices keeps its own minute (2 rows, 2 min total).
var u1Rows int
var u1Min int64
for _, d := range b.store.deviceUsageLog() {
if d.UserID == 101 {
u1Rows++
u1Min += d.Minutes
}
}
if u1Rows != 2 || u1Min != 2 {
t.Errorf("user101 per-device: rows=%d totalMinutes=%d, want 2/2 (每设备各计 1)", u1Rows, u1Min)
}
}
// TestReportUsage_PerDevice verifies a per-device dp_uuid is dual-written: account
// rollup (usage_daily) AND per-device attribution (usage_device_daily).
func TestReportUsage_PerDevice(t *testing.T) {
const nodeUUID = "test-node-usage-dev"
b := newTestServer(t, 1, nodeUUID)
ctx := context.Background()
// dp-dev1 resolves to user 101, device 55 (per-device credential).
b.store.devicesByDpUUID = map[string][2]int64{"dp-dev1": {101, 55}}
_, _, conn := enrollNode(t, b, nodeUUID)
client := agentv1.NewAgentServiceClient(conn)
if _, err := client.Register(ctx, &agentv1.RegisterRequest{NodeUUID: nodeUUID}); err != nil {
t.Fatalf("Register: %v", err)
}
now := time.Now()
if _, err := client.ReportUsage(ctx, &agentv1.UsageReport{
NodeUUID: nodeUUID,
WindowStartUnix: now.Add(-time.Minute).Unix(),
WindowEndUnix: now.Unix(),
Entries: []*agentv1.UsageEntry{
{DpUUID: "dp-dev1", BytesUp: 100, BytesDown: 200, SessionMinutes: 3},
},
}); err != nil {
t.Fatalf("ReportUsage: %v", err)
}
// Account rollup still recorded.
acc := b.store.usageLog()
if len(acc) != 1 || acc[0].UserID != 101 || acc[0].BytesUp != 100 {
t.Fatalf("account usage wrong: %+v", acc)
}
// Per-device usage recorded with device attribution.
dev := b.store.deviceUsageLog()
if len(dev) != 1 {
t.Fatalf("device usage has %d entries, want 1", len(dev))
}
if dev[0].UserID != 101 || dev[0].DeviceID != 55 ||
dev[0].BytesUp != 100 || dev[0].BytesDown != 200 || dev[0].Minutes != 3 {
t.Errorf("device usage wrong: %+v", dev[0])
}
}
// TestReportUsage_RenewsPaidCredential verifies renew-on-activity: a PAID user's
// usage report bumps its data-plane credential expiry ~PaidCredentialTTL into the
// future — so a persistent session never hits the TTL wall mid-connection (the
// root cause of "connected ~1 day then silently drops" on always-on macOS).
func TestReportUsage_RenewsPaidCredential(t *testing.T) {
const nodeUUID = "test-node-renew-paid"
b := newTestServer(t, 1, nodeUUID)
ctx := context.Background()
b.store.devicesByDpUUID = map[string][2]int64{"dp-paid": {101, 55}}
b.store.ent = &nodes.Entitlement{AdGate: false} // paid plan
_, _, conn := enrollNode(t, b, nodeUUID)
client := agentv1.NewAgentServiceClient(conn)
if _, err := client.Register(ctx, &agentv1.RegisterRequest{NodeUUID: nodeUUID}); err != nil {
t.Fatalf("Register: %v", err)
}
before := time.Now().UTC()
now := time.Now()
if _, err := client.ReportUsage(ctx, &agentv1.UsageReport{
NodeUUID: nodeUUID,
WindowStartUnix: now.Add(-time.Minute).Unix(),
WindowEndUnix: now.Unix(),
Entries: []*agentv1.UsageEntry{
{DpUUID: "dp-paid", BytesUp: 100, BytesDown: 200, SessionMinutes: 1},
},
}); err != nil {
t.Fatalf("ReportUsage: %v", err)
}
renewed := b.store.renewLog()
if len(renewed) != 1 {
t.Fatalf("renew calls = %d, want 1 (paid session should renew)", len(renewed))
}
if renewed[0].DpUUID != "dp-paid" {
t.Errorf("renewed dp_uuid = %q, want dp-paid", renewed[0].DpUUID)
}
// Expiry should land ~PaidCredentialTTL from now (allow the test's own runtime slack).
wantMin := before.Add(nodes.PaidCredentialTTL)
wantMax := time.Now().UTC().Add(nodes.PaidCredentialTTL + time.Minute)
if renewed[0].ExpiresAt.Before(wantMin) || renewed[0].ExpiresAt.After(wantMax) {
t.Errorf("renewed expiry %v out of [%v, %v]", renewed[0].ExpiresAt, wantMin, wantMax)
}
}
// TestReportUsage_DoesNotRenewFreeCredential verifies free credentials are NEVER
// renewed: their TTL encodes the daily-minute quota, and renewing would bypass the
// data-plane hard cut-off that enforces it.
func TestReportUsage_DoesNotRenewFreeCredential(t *testing.T) {
const nodeUUID = "test-node-renew-free"
b := newTestServer(t, 1, nodeUUID)
ctx := context.Background()
b.store.devicesByDpUUID = map[string][2]int64{"dp-free": {202, 66}}
b.store.ent = &nodes.Entitlement{AdGate: true} // free plan (minute-quota-gated)
_, _, conn := enrollNode(t, b, nodeUUID)
client := agentv1.NewAgentServiceClient(conn)
if _, err := client.Register(ctx, &agentv1.RegisterRequest{NodeUUID: nodeUUID}); err != nil {
t.Fatalf("Register: %v", err)
}
now := time.Now()
if _, err := client.ReportUsage(ctx, &agentv1.UsageReport{
NodeUUID: nodeUUID,
WindowStartUnix: now.Add(-time.Minute).Unix(),
WindowEndUnix: now.Unix(),
Entries: []*agentv1.UsageEntry{
{DpUUID: "dp-free", BytesUp: 100, BytesDown: 200, SessionMinutes: 1},
},
}); err != nil {
t.Fatalf("ReportUsage: %v", err)
}
if renewed := b.store.renewLog(); len(renewed) != 0 {
t.Fatalf("renew calls = %d, want 0 (free credential must NOT be renewed)", len(renewed))
}
}