Files
pangolin/server/internal/nodes/grpc_test.go
T
wangjia f2499a64e1
ci-pangolin / Lint — shellcheck (push) Successful in 7s
ci-pangolin / OpenAPI Sync Check (push) Successful in 16s
ci-pangolin / Redline Scan — 脱敏 (UI 文案) (push) Successful in 6s
ci-pangolin / Flutter — analyze + test (push) Successful in 24s
ci-pangolin / Portable SQL — 可移植性 (mysql/sqlite) (push) Successful in 4s
ci-pangolin / Codegen Drift — token 生成物未漂移 (push) Successful in 5s
ci-pangolin / Go — build + test (push) Successful in 10s
ci-pangolin / E2E Smoke — L4 进程级端到端 (push) Successful in 14s
ci-pangolin / Go — integration (mysql/redis testcontainers) (push) Failing after 4m8s
ci-pangolin / Golden — 视觉回归 (components + auth) (push) Successful in 14s
fix(usage): 统计按客户端本地日聚合 —— 修「29号却显示28号数据」时区 bug
根因:服务端用量全按 UTC 预切天,UTC+8 用户在中国 0-8 点时 UTC 还是前一天,
「今天」就显示前一天数据。

改法(用户拍板:小时桶+查询带时区偏移):
- migration 000017 `usage_hourly`(UTC 纪元小时整数桶,稀疏存储,空小时不落行)。
- ReportUsage 同时累加 usage_hourly(按 windowEnd 的 UTC 小时);usage_daily 保留作配额/今日。
- `UsageCurve(days, offsetMin)`:读 UTC 小时桶,按客户端时区偏移**重新聚合成本地日**曲线。
- `/v1/usage?tz_offset=分钟`;客户端 `usage()` 带 `DateTime.now().timeZoneOffset.inMinutes`。
- 客户端早已「取曲线最后一点当今日」,故只改服务端聚合 + 带偏移即可。

测试:store 时区分桶证明(UTC 18:00 在 UTC vs UTC+8 落不同本地日,锁住 bug 修复)+
migration v17 + usage_hourly 表;集成测试改喂 hourly;全量 go/flutter test 绿。

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
2026-06-29 07:02:04 +08:00

973 lines
30 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
}
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 nil, nil
}
func (m *mockNodeStore) PersistCredential(_ context.Context, _ int64, _ *agentv1.Credential, _ time.Time) error {
return nil
}
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) 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) 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_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])
}
}