proxy-pool/internal/controller/worker/snapshot_source.go
youfak b9b323af9b
Some checks are pending
ci / proto (push) Waiting to run
ci / test (ubuntu-latest) (push) Waiting to run
ci / test (windows-latest) (push) Waiting to run
ci / race (push) Waiting to run
ci / integration (push) Waiting to run
feat: refresh worker snapshots before expiry
2026-08-02 10:27:56 +08:00

184 lines
5.2 KiB
Go

package worker
import (
"context"
"errors"
"time"
controlplanev1 "proxy-pool/gen/controlplane/v1"
"proxy-pool/internal/controlplane/snapshotwire"
"proxy-pool/internal/domain/workerruntime"
"google.golang.org/protobuf/types/known/timestamppb"
)
var ErrSnapshotsUnavailable = errors.New("worker snapshots are unavailable")
type SnapshotWatchRequest struct {
WorkerID string
SessionID string
LastAppliedVersion uint64
LastChecksum []byte
}
type SnapshotSource interface {
Watch(context.Context, SnapshotWatchRequest) (<-chan *controlplanev1.WorkerSnapshot, error)
}
// RefreshingSnapshotSource keeps a Worker snapshot stream current without
// requiring the Gateway to reconnect after every validity window. It only
// forwards complete snapshots, so Gateway delta handling remains fail-closed.
type RefreshingSnapshotSource struct {
source SnapshotSource
refreshEvery time.Duration
}
func NewRefreshingSnapshotSource(source SnapshotSource, refreshEvery time.Duration) (*RefreshingSnapshotSource, error) {
if source == nil || refreshEvery <= 0 {
return nil, ErrSnapshotsUnavailable
}
return &RefreshingSnapshotSource{source: source, refreshEvery: refreshEvery}, nil
}
func snapshotRefreshEvery(validFor time.Duration) time.Duration {
if validFor <= 0 {
return 0
}
if half := validFor / 2; half > 0 {
return half
}
return validFor
}
func (source *RefreshingSnapshotSource) Watch(
ctx context.Context,
request SnapshotWatchRequest,
) (<-chan *controlplanev1.WorkerSnapshot, error) {
if source == nil || source.source == nil || ctx == nil {
return nil, ErrSnapshotsUnavailable
}
first, err := source.next(ctx, request)
if err != nil {
return nil, err
}
updates := make(chan *controlplanev1.WorkerSnapshot, 1)
updates <- first
go source.refresh(ctx, updates, request, first)
return updates, nil
}
func (source *RefreshingSnapshotSource) refresh(
ctx context.Context,
updates chan<- *controlplanev1.WorkerSnapshot,
request SnapshotWatchRequest,
current *controlplanev1.WorkerSnapshot,
) {
defer close(updates)
version := current.GetVersion()
checksum := append([]byte(nil), current.GetChecksum()...)
timer := time.NewTimer(source.delayFor(current))
defer timer.Stop()
for {
select {
case <-ctx.Done():
return
case <-timer.C:
request.LastAppliedVersion = version
request.LastChecksum = append(request.LastChecksum[:0], checksum...)
next, err := source.next(ctx, request)
if err != nil {
return
}
select {
case <-ctx.Done():
return
case updates <- next:
version = next.GetVersion()
checksum = append(checksum[:0], next.GetChecksum()...)
timer.Reset(source.delayFor(next))
}
}
}
}
func (source *RefreshingSnapshotSource) next(
ctx context.Context,
request SnapshotWatchRequest,
) (*controlplanev1.WorkerSnapshot, error) {
stream, err := source.source.Watch(ctx, request)
if err != nil {
return nil, err
}
select {
case <-ctx.Done():
return nil, ctx.Err()
case snapshot, ok := <-stream:
if !ok || snapshot == nil || snapshot.GetVersion() <= request.LastAppliedVersion {
return nil, ErrSnapshotsUnavailable
}
return snapshot, nil
}
}
func (source *RefreshingSnapshotSource) delayFor(snapshot *controlplanev1.WorkerSnapshot) time.Duration {
delay := source.refreshEvery
if snapshot == nil || snapshot.GetGeneratedAt() == nil || snapshot.GetValidUntil() == nil {
return delay
}
lifetime := snapshot.GetValidUntil().AsTime().Sub(snapshot.GetGeneratedAt().AsTime())
if half := lifetime / 2; half > 0 && half < delay {
return half
}
return delay
}
type OwnershipEpochReader interface {
CurrentOwnershipEpoch(context.Context) (uint64, error)
}
type InitialSnapshotSource struct {
epochs OwnershipEpochReader
validFor time.Duration
now func() time.Time
}
func NewInitialSnapshotSource(epochs OwnershipEpochReader, validFor time.Duration, now func() time.Time) (*InitialSnapshotSource, error) {
if epochs == nil || validFor <= 0 || now == nil {
return nil, ErrSnapshotsUnavailable
}
return &InitialSnapshotSource{epochs: epochs, validFor: validFor, now: now}, nil
}
func (source *InitialSnapshotSource) Watch(ctx context.Context, request SnapshotWatchRequest) (<-chan *controlplanev1.WorkerSnapshot, error) {
if source == nil || ctx == nil || !workerruntime.ValidIdentifier(request.WorkerID) ||
!workerruntime.ValidIdentifier(request.SessionID) || request.LastAppliedVersion == ^uint64(0) {
return nil, ErrSnapshotsUnavailable
}
epoch, err := source.epochs.CurrentOwnershipEpoch(ctx)
if err != nil {
return nil, err
}
if epoch == 0 {
return nil, ErrSnapshotsUnavailable
}
now := source.now().UTC()
if now.IsZero() {
return nil, ErrSnapshotsUnavailable
}
full := &controlplanev1.WorkerSnapshot{
Version: request.LastAppliedVersion + 1, OwnershipEpoch: epoch,
GeneratedAt: timestamppb.New(now), ValidUntil: timestamppb.New(now.Add(source.validFor)),
}
checksum, err := snapshotwire.Checksum(full)
if err != nil {
return nil, err
}
full.Checksum = append([]byte(nil), checksum[:]...)
updates := make(chan *controlplanev1.WorkerSnapshot, 1)
updates <- full
return updates, nil
}
var _ SnapshotSource = (*InitialSnapshotSource)(nil)
var _ SnapshotSource = (*RefreshingSnapshotSource)(nil)