From 27035947dd9677877f8e70818f3dad697fcacc16 Mon Sep 17 00:00:00 2001 From: youfak Date: Fri, 31 Jul 2026 13:29:26 +0800 Subject: [PATCH] feat: coordinate gateway control plane session --- internal/gateway/controlplane/reporter.go | 23 +++- .../gateway/controlplane/session_runner.go | 52 +++++++++ .../controlplane/session_runner_test.go | 109 ++++++++++++++++++ 3 files changed, 182 insertions(+), 2 deletions(-) create mode 100644 internal/gateway/controlplane/session_runner.go create mode 100644 internal/gateway/controlplane/session_runner_test.go diff --git a/internal/gateway/controlplane/reporter.go b/internal/gateway/controlplane/reporter.go index 48d4237..5fa9143 100644 --- a/internal/gateway/controlplane/reporter.go +++ b/internal/gateway/controlplane/reporter.go @@ -138,10 +138,29 @@ func (reporter *RuntimeReporter) Report(ctx context.Context) error { // reports. The snapshot stream may start after Register, so an absent local // snapshot is a normal waiting state rather than a process failure. func (reporter *RuntimeReporter) Run(ctx context.Context) error { - registration, err := reporter.Register(ctx) - if err != nil { + if _, err := reporter.Register(ctx); err != nil { return err } + return reporter.RunRegistered(ctx) +} + +// RunRegistered reports Runtime for the session established by Register. +// It is useful when a SnapshotWatcher and a RuntimeReporter share one worker +// lifecycle and registration must occur exactly once. +func (reporter *RuntimeReporter) RunRegistered(ctx context.Context) error { + if reporter == nil || ctx == nil { + return ErrInvalidOptions + } + reporter.mu.Lock() + registration := reporter.session + reporter.mu.Unlock() + if registration.SessionID == "" { + return ErrNotRegistered + } + return reporter.runRegistered(ctx, registration) +} + +func (reporter *RuntimeReporter) runRegistered(ctx context.Context, registration Registration) error { ticker := time.NewTicker(registration.HeartbeatInterval) defer ticker.Stop() for { diff --git a/internal/gateway/controlplane/session_runner.go b/internal/gateway/controlplane/session_runner.go new file mode 100644 index 0000000..b0b72dd --- /dev/null +++ b/internal/gateway/controlplane/session_runner.go @@ -0,0 +1,52 @@ +package controlplane + +import ( + "context" + "errors" +) + +var ErrSnapshotStreamClosed = errors.New("worker snapshot stream closed") + +type SessionRunner struct { + reporter *RuntimeReporter + watcher *SnapshotWatcher +} + +func NewSessionRunner(reporter *RuntimeReporter, watcher *SnapshotWatcher) (*SessionRunner, error) { + if reporter == nil || watcher == nil { + return nil, ErrInvalidOptions + } + return &SessionRunner{reporter: reporter, watcher: watcher}, nil +} + +// Run registers exactly once, then runs snapshot intake and Runtime reporting +// against that session. A stream ending without context cancellation is a +// failure so the owning process can apply its reconnect policy. +func (runner *SessionRunner) Run(ctx context.Context) error { + if runner == nil || ctx == nil { + return ErrInvalidOptions + } + registration, err := runner.reporter.Register(ctx) + if err != nil { + return err + } + groupCtx, cancel := context.WithCancel(ctx) + defer cancel() + results := make(chan error, 2) + go func() { results <- runner.watcher.Watch(groupCtx, registration.SessionID) }() + go func() { results <- runner.reporter.RunRegistered(groupCtx) }() + + first := <-results + cancel() + second := <-results + if ctx.Err() != nil { + return ctx.Err() + } + if first == nil { + return ErrSnapshotStreamClosed + } + if errors.Is(first, context.Canceled) && second != nil { + return second + } + return first +} diff --git a/internal/gateway/controlplane/session_runner_test.go b/internal/gateway/controlplane/session_runner_test.go new file mode 100644 index 0000000..cee03eb --- /dev/null +++ b/internal/gateway/controlplane/session_runner_test.go @@ -0,0 +1,109 @@ +package controlplane + +import ( + "context" + "errors" + "sync" + "testing" + "time" + + controlplanev1 "proxy-pool/gen/controlplane/v1" + "proxy-pool/internal/gateway/snapshot" + + "google.golang.org/protobuf/types/known/durationpb" + "google.golang.org/protobuf/types/known/emptypb" + "google.golang.org/protobuf/types/known/timestamppb" +) + +func TestSessionRunnerRegistersOnceWatchesAndReports(t *testing.T) { + store := snapshot.NewStore("cluster-a", "worker-a") + full := &controlplanev1.WorkerSnapshot{ + Version: 1, OwnershipEpoch: 7, GeneratedAt: timestamppb.New(time.Now()), ValidUntil: timestamppb.New(time.Now().Add(time.Minute)), + } + setSnapshotChecksum(t, full) + runtime := make(chan *controlplanev1.ReportRuntimeRequest, 8) + reporterClient := &sessionReporterClient{runtime: runtime} + reporter, err := NewRuntimeReporter(reporterClient, store, Options{ + WorkerID: "worker-a", InstanceID: "instance-a", Zone: "zone-a", ProtocolVersion: 1, Now: time.Now, + }) + if err != nil { + t.Fatalf("NewRuntimeReporter(): %v", err) + } + watcherClient := &sessionWatcherClient{full: full} + watcher, err := NewSnapshotWatcher(watcherClient, store, SnapshotWatcherOptions{ClusterID: "cluster-a", WorkerID: "worker-a"}) + if err != nil { + t.Fatalf("NewSnapshotWatcher(): %v", err) + } + runner, err := NewSessionRunner(reporter, watcher) + if err != nil { + t.Fatalf("NewSessionRunner(): %v", err) + } + ctx, cancel := context.WithCancel(context.Background()) + result := make(chan error, 1) + go func() { result <- runner.Run(ctx) }() + select { + case request := <-runtime: + if request.GetSessionId() != "session-a" || request.GetSnapshotVersion() != 1 || request.GetReportSequence() != 1 { + t.Fatalf("runtime request = %+v", request) + } + case <-time.After(time.Second): + t.Fatal("runtime report was not sent") + } + cancel() + if err := <-result; !errors.Is(err, context.Canceled) { + t.Fatalf("Run() error = %v, want context canceled", err) + } + if reporterClient.registers != 1 || watcherClient.ack == nil || !watcherClient.ack.GetApplied() { + t.Fatalf("registers=%d ack=%+v", reporterClient.registers, watcherClient.ack) + } +} + +type sessionReporterClient struct { + mu sync.Mutex + registers int + runtime chan *controlplanev1.ReportRuntimeRequest +} + +func (client *sessionReporterClient) RegisterWorker(context.Context, *controlplanev1.RegisterWorkerRequest) (*controlplanev1.RegisterWorkerResponse, error) { + client.mu.Lock() + client.registers++ + client.mu.Unlock() + return &controlplanev1.RegisterWorkerResponse{ + WorkerId: "worker-a", SessionId: "session-a", OwnershipEpoch: 7, + HeartbeatInterval: durationpb.New(5 * time.Millisecond), MaxStaleAge: durationpb.New(time.Second), + }, nil +} + +func (client *sessionReporterClient) ReportRuntime(_ context.Context, request *controlplanev1.ReportRuntimeRequest) (*controlplanev1.ReportRuntimeResponse, error) { + client.runtime <- request + return &controlplanev1.ReportRuntimeResponse{AcceptedOwnershipEpoch: 7}, nil +} + +type sessionWatcherClient struct { + full *controlplanev1.WorkerSnapshot + ack *controlplanev1.AcknowledgeSnapshotRequest +} + +func (client *sessionWatcherClient) Watch(ctx context.Context, _ *controlplanev1.WatchSnapshotsRequest) (SnapshotStream, error) { + return &sessionSnapshotStream{ctx: ctx, full: client.full}, nil +} + +func (client *sessionWatcherClient) Acknowledge(_ context.Context, acknowledgement *controlplanev1.AcknowledgeSnapshotRequest) (*emptypb.Empty, error) { + client.ack = acknowledgement + return &emptypb.Empty{}, nil +} + +type sessionSnapshotStream struct { + ctx context.Context + full *controlplanev1.WorkerSnapshot + sent bool +} + +func (stream *sessionSnapshotStream) Recv() (*controlplanev1.SnapshotEnvelope, error) { + if !stream.sent { + stream.sent = true + return &controlplanev1.SnapshotEnvelope{Payload: &controlplanev1.SnapshotEnvelope_Full{Full: stream.full}}, nil + } + <-stream.ctx.Done() + return nil, stream.ctx.Err() +}