feat: add worker session runtime domain store
This commit is contained in:
parent
43dec7324a
commit
a79d030c82
24
internal/domain/workerruntime/contract_external_test.go
Normal file
24
internal/domain/workerruntime/contract_external_test.go
Normal file
@ -0,0 +1,24 @@
|
|||||||
|
package workerruntime_test
|
||||||
|
|
||||||
|
import (
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"proxy-pool/internal/domain/workerruntime"
|
||||||
|
"proxy-pool/internal/domain/workerruntime/contracttest"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestMemoryStoreContract(t *testing.T) {
|
||||||
|
now := time.Date(2026, 7, 31, 9, 0, 0, 0, time.UTC)
|
||||||
|
contracttest.Run(t, func(t *testing.T) contracttest.Fixture {
|
||||||
|
t.Helper()
|
||||||
|
store, err := workerruntime.NewMemoryStore(func() time.Time { return now })
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NewMemoryStore(): %v", err)
|
||||||
|
}
|
||||||
|
return contracttest.Fixture{
|
||||||
|
Store: store, Reader: store,
|
||||||
|
Advance: func(duration time.Duration) { now = now.Add(duration) },
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
131
internal/domain/workerruntime/contracttest/contract.go
Normal file
131
internal/domain/workerruntime/contracttest/contract.go
Normal file
@ -0,0 +1,131 @@
|
|||||||
|
package contracttest
|
||||||
|
|
||||||
|
import (
|
||||||
|
"context"
|
||||||
|
"crypto/sha256"
|
||||||
|
"errors"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
|
||||||
|
"proxy-pool/internal/domain/workerruntime"
|
||||||
|
)
|
||||||
|
|
||||||
|
type Fixture struct {
|
||||||
|
Store workerruntime.ControlStore
|
||||||
|
Reader workerruntime.RuntimeReader
|
||||||
|
Advance func(time.Duration)
|
||||||
|
}
|
||||||
|
|
||||||
|
type Factory func(*testing.T) Fixture
|
||||||
|
|
||||||
|
// Run exercises the public control-store behavior shared by Memory and Redis.
|
||||||
|
func Run(t *testing.T, factory Factory) {
|
||||||
|
t.Helper()
|
||||||
|
t.Run("acknowledged runtime lifecycle", func(t *testing.T) { runLifecycle(t, newFixture(t, factory)) })
|
||||||
|
t.Run("negative acknowledgement fences runtime", func(t *testing.T) { runNegativeAck(t, newFixture(t, factory)) })
|
||||||
|
}
|
||||||
|
|
||||||
|
func runLifecycle(t *testing.T, fixture Fixture) {
|
||||||
|
t.Helper()
|
||||||
|
ctx := context.Background()
|
||||||
|
open(t, fixture.Store)
|
||||||
|
epoch, err := fixture.Store.CurrentOwnershipEpoch(ctx)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CurrentOwnershipEpoch(): %v", err)
|
||||||
|
}
|
||||||
|
reference := snapshot(7, epoch, "snapshot-7")
|
||||||
|
report := runtimeReport(1, 7, epoch)
|
||||||
|
if err := fixture.Store.ReplaceRuntime(ctx, report, time.Minute); !errors.Is(err, workerruntime.ErrSnapshotMismatch) {
|
||||||
|
t.Fatalf("ReplaceRuntime(before ACK) error = %v", err)
|
||||||
|
}
|
||||||
|
if err := fixture.Store.RecordIssuedSnapshot(ctx, reference, time.Minute); err != nil {
|
||||||
|
t.Fatalf("RecordIssuedSnapshot(): %v", err)
|
||||||
|
}
|
||||||
|
ack := workerruntime.SnapshotAcknowledgement{WorkerID: "worker-a", SessionID: "session-a", Reference: reference, Applied: true}
|
||||||
|
if err := fixture.Store.AcknowledgeSnapshot(ctx, ack, time.Minute); err != nil {
|
||||||
|
t.Fatalf("AcknowledgeSnapshot(): %v", err)
|
||||||
|
}
|
||||||
|
if err := fixture.Store.ReplaceRuntime(ctx, report, time.Minute); err != nil {
|
||||||
|
t.Fatalf("ReplaceRuntime(after ACK): %v", err)
|
||||||
|
}
|
||||||
|
assertFresh(t, fixture.Reader, epoch, true)
|
||||||
|
if err := fixture.Store.AcknowledgeSnapshot(ctx, ack, time.Minute); err != nil {
|
||||||
|
t.Fatalf("AcknowledgeSnapshot(replay): %v", err)
|
||||||
|
}
|
||||||
|
if err := fixture.Store.ReplaceRuntime(ctx, runtimeReport(0, 7, epoch), time.Minute); !errors.Is(err, workerruntime.ErrInvalidReport) {
|
||||||
|
t.Fatalf("ReplaceRuntime(invalid sequence) error = %v", err)
|
||||||
|
}
|
||||||
|
fixture.Advance(2 * time.Minute)
|
||||||
|
assertFresh(t, fixture.Reader, epoch, false)
|
||||||
|
}
|
||||||
|
|
||||||
|
func runNegativeAck(t *testing.T, fixture Fixture) {
|
||||||
|
t.Helper()
|
||||||
|
ctx := context.Background()
|
||||||
|
open(t, fixture.Store)
|
||||||
|
epoch, err := fixture.Store.CurrentOwnershipEpoch(ctx)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CurrentOwnershipEpoch(): %v", err)
|
||||||
|
}
|
||||||
|
first := snapshot(7, epoch, "snapshot-7")
|
||||||
|
if err := fixture.Store.RecordIssuedSnapshot(ctx, first, time.Minute); err != nil {
|
||||||
|
t.Fatalf("RecordIssuedSnapshot(first): %v", err)
|
||||||
|
}
|
||||||
|
if err := fixture.Store.AcknowledgeSnapshot(ctx, workerruntime.SnapshotAcknowledgement{WorkerID: "worker-a", SessionID: "session-a", Reference: first, Applied: true}, time.Minute); err != nil {
|
||||||
|
t.Fatalf("AcknowledgeSnapshot(first): %v", err)
|
||||||
|
}
|
||||||
|
second := snapshot(8, epoch, "snapshot-8")
|
||||||
|
if err := fixture.Store.RecordIssuedSnapshot(ctx, second, time.Minute); err != nil {
|
||||||
|
t.Fatalf("RecordIssuedSnapshot(second): %v", err)
|
||||||
|
}
|
||||||
|
if err := fixture.Store.AcknowledgeSnapshot(ctx, workerruntime.SnapshotAcknowledgement{
|
||||||
|
WorkerID: "worker-a", SessionID: "session-a", Reference: second, ErrorCode: "apply_failed",
|
||||||
|
}, time.Minute); err != nil {
|
||||||
|
t.Fatalf("AcknowledgeSnapshot(negative): %v", err)
|
||||||
|
}
|
||||||
|
if err := fixture.Store.ReplaceRuntime(ctx, runtimeReport(1, 7, epoch), time.Minute); !errors.Is(err, workerruntime.ErrSnapshotMismatch) {
|
||||||
|
t.Fatalf("ReplaceRuntime(delayed): %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func newFixture(t *testing.T, factory Factory) Fixture {
|
||||||
|
t.Helper()
|
||||||
|
fixture := factory(t)
|
||||||
|
if fixture.Store == nil || fixture.Reader == nil || fixture.Advance == nil {
|
||||||
|
t.Fatal("contract fixture is incomplete")
|
||||||
|
}
|
||||||
|
return fixture
|
||||||
|
}
|
||||||
|
|
||||||
|
func open(t *testing.T, store workerruntime.ControlStore) {
|
||||||
|
t.Helper()
|
||||||
|
err := store.OpenSession(context.Background(), workerruntime.Session{
|
||||||
|
WorkerID: "worker-a", InstanceID: "instance-a", SessionID: "session-a", Zone: "zone-a", ProtocolVersion: 1,
|
||||||
|
}, time.Minute)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("OpenSession(): %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func snapshot(version, epoch uint64, value string) workerruntime.SnapshotReference {
|
||||||
|
return workerruntime.SnapshotReference{
|
||||||
|
WorkerID: "worker-a", Version: version, OwnershipEpoch: epoch, Checksum: sha256.Sum256([]byte(value)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func runtimeReport(sequence, version, epoch uint64) workerruntime.Report {
|
||||||
|
return workerruntime.Report{
|
||||||
|
WorkerID: "worker-a", SessionID: "session-a", Sequence: sequence, SnapshotVersion: version,
|
||||||
|
OwnershipEpoch: epoch, ObservedAt: time.Date(2026, 7, 31, 9, 0, 0, 0, time.UTC),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func assertFresh(t *testing.T, reader workerruntime.RuntimeReader, epoch uint64, want bool) {
|
||||||
|
t.Helper()
|
||||||
|
snapshots, err := reader.ReadRuntime(context.Background(), []workerruntime.OwnedProxy{{
|
||||||
|
ProxyID: "proxy-a", WorkerID: "worker-a", OwnershipEpoch: epoch,
|
||||||
|
}})
|
||||||
|
if err != nil || len(snapshots) != 1 || snapshots[0].Fresh != want {
|
||||||
|
t.Fatalf("ReadRuntime() = %+v, %v; want Fresh=%t", snapshots, err, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
@ -3,18 +3,17 @@ package workerruntime
|
|||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"crypto/sha256"
|
"crypto/sha256"
|
||||||
"encoding/json"
|
|
||||||
"sort"
|
|
||||||
"strings"
|
|
||||||
"sync"
|
"sync"
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
type MemoryStore struct {
|
type MemoryStore struct {
|
||||||
mu sync.Mutex
|
mu sync.Mutex
|
||||||
now func() time.Time
|
now func() time.Time
|
||||||
sessions map[string]memorySession
|
epoch uint64
|
||||||
reports map[string]memoryReport
|
sessions map[string]memorySession
|
||||||
|
references map[string]memoryReference
|
||||||
|
reports map[string]memoryReport
|
||||||
}
|
}
|
||||||
|
|
||||||
type memorySession struct {
|
type memorySession struct {
|
||||||
@ -22,6 +21,11 @@ type memorySession struct {
|
|||||||
expiresAt time.Time
|
expiresAt time.Time
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type memoryReference struct {
|
||||||
|
value SnapshotReference
|
||||||
|
expiresAt time.Time
|
||||||
|
}
|
||||||
|
|
||||||
type memoryReport struct {
|
type memoryReport struct {
|
||||||
value Report
|
value Report
|
||||||
digest [sha256.Size]byte
|
digest [sha256.Size]byte
|
||||||
@ -30,6 +34,7 @@ type memoryReport struct {
|
|||||||
}
|
}
|
||||||
|
|
||||||
var (
|
var (
|
||||||
|
_ ControlStore = (*MemoryStore)(nil)
|
||||||
_ SessionWriter = (*MemoryStore)(nil)
|
_ SessionWriter = (*MemoryStore)(nil)
|
||||||
_ ReportWriter = (*MemoryStore)(nil)
|
_ ReportWriter = (*MemoryStore)(nil)
|
||||||
_ RuntimeReader = (*MemoryStore)(nil)
|
_ RuntimeReader = (*MemoryStore)(nil)
|
||||||
@ -40,33 +45,184 @@ func NewMemoryStore(now func() time.Time) (*MemoryStore, error) {
|
|||||||
return nil, ErrInvalidStore
|
return nil, ErrInvalidStore
|
||||||
}
|
}
|
||||||
return &MemoryStore{
|
return &MemoryStore{
|
||||||
now: now, sessions: make(map[string]memorySession), reports: make(map[string]memoryReport),
|
now: now, epoch: 1, sessions: make(map[string]memorySession),
|
||||||
|
references: make(map[string]memoryReference), reports: make(map[string]memoryReport),
|
||||||
}, nil
|
}, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func (store *MemoryStore) ReplaceSession(ctx context.Context, session Session, ttl time.Duration) error {
|
func (store *MemoryStore) CurrentOwnershipEpoch(ctx context.Context) (uint64, error) {
|
||||||
if ctx == nil || store == nil || !validSession(session) || ttl <= 0 {
|
if ctx == nil || store == nil {
|
||||||
|
return 0, ErrInvalidStore
|
||||||
|
}
|
||||||
|
if err := ctx.Err(); err != nil {
|
||||||
|
return 0, err
|
||||||
|
}
|
||||||
|
store.mu.Lock()
|
||||||
|
defer store.mu.Unlock()
|
||||||
|
if store.epoch == 0 {
|
||||||
|
return 0, ErrInvalidStore
|
||||||
|
}
|
||||||
|
return store.epoch, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// OpenSession always replaces the previous Worker session and clears Runtime.
|
||||||
|
func (store *MemoryStore) OpenSession(ctx context.Context, session Session, ttl time.Duration) error {
|
||||||
|
if ctx == nil || store == nil || ttl <= 0 {
|
||||||
return ErrInvalidSession
|
return ErrInvalidSession
|
||||||
}
|
}
|
||||||
if err := ctx.Err(); err != nil {
|
if err := ctx.Err(); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
now := store.now().UTC()
|
normalized, err := NormalizeSession(session)
|
||||||
if now.IsZero() {
|
if err != nil {
|
||||||
return ErrInvalidStore
|
return err
|
||||||
|
}
|
||||||
|
now, err := store.currentTime()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
store.mu.Lock()
|
||||||
|
defer store.mu.Unlock()
|
||||||
|
delete(store.reports, normalized.WorkerID)
|
||||||
|
store.sessions[normalized.WorkerID] = memorySession{value: normalized, expiresAt: now.Add(ttl)}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (store *MemoryStore) RecordIssuedSnapshot(ctx context.Context, reference SnapshotReference, ttl time.Duration) error {
|
||||||
|
if ctx == nil || store == nil || ttl <= 0 {
|
||||||
|
return ErrInvalidSnapshotReference
|
||||||
|
}
|
||||||
|
if err := ctx.Err(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
normalized, err := NormalizeSnapshotReference(reference)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
now, err := store.currentTime()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
store.mu.Lock()
|
||||||
|
defer store.mu.Unlock()
|
||||||
|
if normalized.OwnershipEpoch != store.epoch {
|
||||||
|
return ErrSnapshotMismatch
|
||||||
|
}
|
||||||
|
current, exists := store.references[normalized.WorkerID]
|
||||||
|
if exists && !current.expiresAt.After(now) {
|
||||||
|
delete(store.references, normalized.WorkerID)
|
||||||
|
exists = false
|
||||||
|
}
|
||||||
|
if exists {
|
||||||
|
switch compareSnapshotTuple(normalized, current.value) {
|
||||||
|
case -1:
|
||||||
|
return ErrStaleSnapshotReference
|
||||||
|
case 0:
|
||||||
|
if normalized.Checksum != current.value.Checksum {
|
||||||
|
return ErrConflictingSnapshotReference
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
store.references[normalized.WorkerID] = memoryReference{value: normalized, expiresAt: now.Add(ttl)}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func (store *MemoryStore) AcknowledgeSnapshot(ctx context.Context, acknowledgement SnapshotAcknowledgement, ttl time.Duration) error {
|
||||||
|
if ctx == nil || store == nil || ttl <= 0 {
|
||||||
|
return ErrInvalidAcknowledgement
|
||||||
|
}
|
||||||
|
if err := ctx.Err(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
normalized, err := NormalizeAcknowledgement(acknowledgement)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
now, err := store.currentTime()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
store.mu.Lock()
|
||||||
|
defer store.mu.Unlock()
|
||||||
|
session, exists := store.sessions[normalized.WorkerID]
|
||||||
|
if !exists || !session.expiresAt.After(now) || session.value.SessionID != normalized.SessionID {
|
||||||
|
return ErrStaleSession
|
||||||
|
}
|
||||||
|
if session.value.AckedSnapshotVersion != 0 {
|
||||||
|
acknowledged := sessionReference(session.value)
|
||||||
|
switch compareSnapshotTuple(normalized.Reference, acknowledged) {
|
||||||
|
case -1:
|
||||||
|
return ErrStaleAcknowledgement
|
||||||
|
case 0:
|
||||||
|
if normalized.Reference.Checksum != acknowledged.Checksum {
|
||||||
|
return ErrSnapshotMismatch
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
reference, exists := store.references[normalized.WorkerID]
|
||||||
|
if !exists || !reference.expiresAt.After(now) {
|
||||||
|
return ErrSnapshotMismatch
|
||||||
|
}
|
||||||
|
switch compareSnapshotTuple(normalized.Reference, reference.value) {
|
||||||
|
case -1:
|
||||||
|
return ErrStaleAcknowledgement
|
||||||
|
case 1:
|
||||||
|
return ErrSnapshotMismatch
|
||||||
|
}
|
||||||
|
if normalized.Reference.Checksum != reference.value.Checksum {
|
||||||
|
return ErrSnapshotMismatch
|
||||||
|
}
|
||||||
|
if !normalized.Applied {
|
||||||
|
delete(store.reports, normalized.WorkerID)
|
||||||
|
session.value.RuntimeEnabled = false
|
||||||
|
session.expiresAt = now.Add(ttl)
|
||||||
|
store.sessions[normalized.WorkerID] = session
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
if session.value.AckedSnapshotVersion != 0 && compareSnapshotTuple(normalized.Reference, sessionReference(session.value)) == 0 {
|
||||||
|
if !session.value.RuntimeEnabled {
|
||||||
|
delete(store.reports, normalized.WorkerID)
|
||||||
|
session.value.RuntimeEnabled = true
|
||||||
|
}
|
||||||
|
session.expiresAt = now.Add(ttl)
|
||||||
|
store.sessions[normalized.WorkerID] = session
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
delete(store.reports, normalized.WorkerID)
|
||||||
|
session.value.AckedSnapshotVersion = normalized.Reference.Version
|
||||||
|
session.value.AckedOwnershipEpoch = normalized.Reference.OwnershipEpoch
|
||||||
|
session.value.AckedChecksum = normalized.Reference.Checksum
|
||||||
|
session.value.RuntimeEnabled = true
|
||||||
|
session.expiresAt = now.Add(ttl)
|
||||||
|
store.sessions[normalized.WorkerID] = session
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// ReplaceSession is retained temporarily for the pre-control-plane adapters.
|
||||||
|
func (store *MemoryStore) ReplaceSession(ctx context.Context, session Session, ttl time.Duration) error {
|
||||||
|
if ctx == nil || store == nil || !validLegacySession(session) || ttl <= 0 {
|
||||||
|
return ErrInvalidSession
|
||||||
|
}
|
||||||
|
if err := ctx.Err(); err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
now, err := store.currentTime()
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
}
|
}
|
||||||
store.mu.Lock()
|
store.mu.Lock()
|
||||||
defer store.mu.Unlock()
|
defer store.mu.Unlock()
|
||||||
current, exists := store.sessions[session.WorkerID]
|
current, exists := store.sessions[session.WorkerID]
|
||||||
identityChanged := exists && (current.value.SessionID != session.SessionID || current.value.InstanceID != session.InstanceID)
|
identityChanged := exists && (current.value.SessionID != session.SessionID || current.value.InstanceID != session.InstanceID)
|
||||||
expired := exists && !current.expiresAt.After(now)
|
expired := exists && !current.expiresAt.After(now)
|
||||||
if exists && !identityChanged && !expired && sessionBefore(session, current.value) {
|
if exists && !identityChanged && !expired && legacySessionBefore(session, current.value) {
|
||||||
return ErrStaleSession
|
return ErrStaleSession
|
||||||
}
|
}
|
||||||
ackAdvanced := exists && !identityChanged && !expired && sessionAfter(session, current.value)
|
ackAdvanced := exists && !identityChanged && !expired && legacySessionAfter(session, current.value)
|
||||||
if identityChanged || expired || ackAdvanced {
|
if identityChanged || expired || ackAdvanced {
|
||||||
delete(store.reports, session.WorkerID)
|
delete(store.reports, session.WorkerID)
|
||||||
}
|
}
|
||||||
|
session.RuntimeEnabled = true
|
||||||
store.sessions[session.WorkerID] = memorySession{value: session, expiresAt: now.Add(ttl)}
|
store.sessions[session.WorkerID] = memorySession{value: session, expiresAt: now.Add(ttl)}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
@ -78,44 +234,48 @@ func (store *MemoryStore) ReplaceRuntime(ctx context.Context, report Report, ttl
|
|||||||
if err := ctx.Err(); err != nil {
|
if err := ctx.Err(); err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
normalized, counterIndex, err := normalizeReport(report)
|
normalized, digest, err := NormalizeReport(report)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return err
|
return err
|
||||||
}
|
}
|
||||||
payload, err := json.Marshal(normalized)
|
now, err := store.currentTime()
|
||||||
if err != nil {
|
if err != nil {
|
||||||
return ErrInvalidReport
|
return err
|
||||||
}
|
|
||||||
digest := sha256.Sum256(payload)
|
|
||||||
now := store.now().UTC()
|
|
||||||
if now.IsZero() {
|
|
||||||
return ErrInvalidStore
|
|
||||||
}
|
}
|
||||||
store.mu.Lock()
|
store.mu.Lock()
|
||||||
defer store.mu.Unlock()
|
defer store.mu.Unlock()
|
||||||
session, exists := store.sessions[report.WorkerID]
|
session, exists := store.sessions[normalized.WorkerID]
|
||||||
if !exists || !session.expiresAt.After(now) || session.value.SessionID != report.SessionID {
|
if !exists || !session.expiresAt.After(now) || session.value.SessionID != normalized.SessionID {
|
||||||
return ErrStaleSession
|
return ErrStaleSession
|
||||||
}
|
}
|
||||||
if report.SnapshotVersion != session.value.AckedSnapshotVersion ||
|
if !session.value.RuntimeEnabled || normalized.SnapshotVersion != session.value.AckedSnapshotVersion ||
|
||||||
report.OwnershipEpoch != session.value.AckedOwnershipEpoch {
|
normalized.OwnershipEpoch != session.value.AckedOwnershipEpoch {
|
||||||
return ErrStaleReport
|
return ErrSnapshotMismatch
|
||||||
}
|
}
|
||||||
if current, exists := store.reports[report.WorkerID]; exists && current.value.SessionID == report.SessionID {
|
current, exists := store.reports[normalized.WorkerID]
|
||||||
|
if exists && !current.expiresAt.After(now) {
|
||||||
|
delete(store.reports, normalized.WorkerID)
|
||||||
|
exists = false
|
||||||
|
}
|
||||||
|
if exists && current.value.SessionID == normalized.SessionID {
|
||||||
switch {
|
switch {
|
||||||
case normalized.Sequence < current.value.Sequence:
|
case normalized.Sequence < current.value.Sequence:
|
||||||
return ErrStaleReport
|
return ErrStaleReport
|
||||||
case normalized.Sequence == current.value.Sequence && digest != current.digest:
|
case normalized.Sequence == current.value.Sequence && digest != current.digest:
|
||||||
return ErrConflictingReport
|
return ErrConflictingReport
|
||||||
case normalized.Sequence == current.value.Sequence:
|
case normalized.Sequence == current.value.Sequence:
|
||||||
|
current.expiresAt = now.Add(ttl)
|
||||||
|
store.reports[normalized.WorkerID] = current
|
||||||
|
session.expiresAt = now.Add(ttl)
|
||||||
|
store.sessions[normalized.WorkerID] = session
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
store.reports[report.WorkerID] = memoryReport{
|
store.reports[normalized.WorkerID] = memoryReport{
|
||||||
value: normalized, digest: digest, expiresAt: now.Add(ttl), counters: counterIndex,
|
value: normalized, digest: digest, expiresAt: now.Add(ttl), counters: countersByProxy(normalized.Counters),
|
||||||
}
|
}
|
||||||
session.expiresAt = now.Add(ttl)
|
session.expiresAt = now.Add(ttl)
|
||||||
store.sessions[report.WorkerID] = session
|
store.sessions[normalized.WorkerID] = session
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
@ -128,7 +288,7 @@ func (store *MemoryStore) ReadRuntime(ctx context.Context, proxies []OwnedProxy)
|
|||||||
}
|
}
|
||||||
seen := make(map[string]struct{}, len(proxies))
|
seen := make(map[string]struct{}, len(proxies))
|
||||||
for _, proxy := range proxies {
|
for _, proxy := range proxies {
|
||||||
if !clean(proxy.ProxyID) || !clean(proxy.WorkerID) || proxy.OwnershipEpoch == 0 {
|
if !ValidIdentifier(proxy.ProxyID) || !ValidIdentifier(proxy.WorkerID) || proxy.OwnershipEpoch == 0 {
|
||||||
return nil, ErrInvalidQuery
|
return nil, ErrInvalidQuery
|
||||||
}
|
}
|
||||||
key := proxy.WorkerID + "\x00" + proxy.ProxyID
|
key := proxy.WorkerID + "\x00" + proxy.ProxyID
|
||||||
@ -137,9 +297,9 @@ func (store *MemoryStore) ReadRuntime(ctx context.Context, proxies []OwnedProxy)
|
|||||||
}
|
}
|
||||||
seen[key] = struct{}{}
|
seen[key] = struct{}{}
|
||||||
}
|
}
|
||||||
now := store.now().UTC()
|
now, err := store.currentTime()
|
||||||
if now.IsZero() {
|
if err != nil {
|
||||||
return nil, ErrInvalidStore
|
return nil, err
|
||||||
}
|
}
|
||||||
store.mu.Lock()
|
store.mu.Lock()
|
||||||
defer store.mu.Unlock()
|
defer store.mu.Unlock()
|
||||||
@ -149,7 +309,7 @@ func (store *MemoryStore) ReadRuntime(ctx context.Context, proxies []OwnedProxy)
|
|||||||
session, sessionExists := store.sessions[proxy.WorkerID]
|
session, sessionExists := store.sessions[proxy.WorkerID]
|
||||||
report, reportExists := store.reports[proxy.WorkerID]
|
report, reportExists := store.reports[proxy.WorkerID]
|
||||||
if !sessionExists || !reportExists || !session.expiresAt.After(now) || !report.expiresAt.After(now) ||
|
if !sessionExists || !reportExists || !session.expiresAt.After(now) || !report.expiresAt.After(now) ||
|
||||||
report.value.SessionID != session.value.SessionID ||
|
!session.value.RuntimeEnabled || report.value.SessionID != session.value.SessionID ||
|
||||||
report.value.SnapshotVersion != session.value.AckedSnapshotVersion ||
|
report.value.SnapshotVersion != session.value.AckedSnapshotVersion ||
|
||||||
report.value.OwnershipEpoch != session.value.AckedOwnershipEpoch ||
|
report.value.OwnershipEpoch != session.value.AckedOwnershipEpoch ||
|
||||||
report.value.OwnershipEpoch < proxy.OwnershipEpoch {
|
report.value.OwnershipEpoch < proxy.OwnershipEpoch {
|
||||||
@ -165,47 +325,31 @@ func (store *MemoryStore) ReadRuntime(ctx context.Context, proxies []OwnedProxy)
|
|||||||
return result, nil
|
return result, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
func normalizeReport(report Report) (Report, map[string]Counter, error) {
|
func (store *MemoryStore) currentTime() (time.Time, error) {
|
||||||
if !clean(report.WorkerID) || !clean(report.SessionID) || report.Sequence == 0 ||
|
now := store.now().UTC()
|
||||||
report.SnapshotVersion == 0 || report.OwnershipEpoch == 0 || report.ObservedAt.IsZero() {
|
if now.IsZero() {
|
||||||
return Report{}, nil, ErrInvalidReport
|
return time.Time{}, ErrInvalidStore
|
||||||
}
|
}
|
||||||
normalized := report
|
return now, nil
|
||||||
normalized.ObservedAt = report.ObservedAt.UTC()
|
|
||||||
normalized.Counters = append([]Counter(nil), report.Counters...)
|
|
||||||
sort.Slice(normalized.Counters, func(left, right int) bool {
|
|
||||||
return normalized.Counters[left].ProxyID < normalized.Counters[right].ProxyID
|
|
||||||
})
|
|
||||||
index := make(map[string]Counter, len(normalized.Counters))
|
|
||||||
for _, counter := range normalized.Counters {
|
|
||||||
if !clean(counter.ProxyID) || counter.Active < 0 || counter.Reserved < 0 {
|
|
||||||
return Report{}, nil, ErrInvalidReport
|
|
||||||
}
|
|
||||||
if _, exists := index[counter.ProxyID]; exists {
|
|
||||||
return Report{}, nil, ErrInvalidReport
|
|
||||||
}
|
|
||||||
index[counter.ProxyID] = counter
|
|
||||||
}
|
|
||||||
return normalized, index, nil
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func validSession(session Session) bool {
|
func countersByProxy(counters []Counter) map[string]Counter {
|
||||||
return clean(session.WorkerID) && clean(session.InstanceID) && clean(session.SessionID) &&
|
indexed := make(map[string]Counter, len(counters))
|
||||||
|
for _, counter := range counters {
|
||||||
|
indexed[counter.ProxyID] = counter
|
||||||
|
}
|
||||||
|
return indexed
|
||||||
|
}
|
||||||
|
|
||||||
|
func validLegacySession(session Session) bool {
|
||||||
|
return ValidIdentifier(session.WorkerID) && ValidIdentifier(session.InstanceID) && ValidIdentifier(session.SessionID) &&
|
||||||
session.AckedSnapshotVersion > 0 && session.AckedOwnershipEpoch > 0
|
session.AckedSnapshotVersion > 0 && session.AckedOwnershipEpoch > 0
|
||||||
}
|
}
|
||||||
|
|
||||||
func sessionBefore(left, right Session) bool {
|
func legacySessionBefore(left, right Session) bool {
|
||||||
return left.AckedOwnershipEpoch < right.AckedOwnershipEpoch ||
|
return compareSnapshotTuple(sessionReference(left), sessionReference(right)) < 0
|
||||||
(left.AckedOwnershipEpoch == right.AckedOwnershipEpoch &&
|
|
||||||
left.AckedSnapshotVersion < right.AckedSnapshotVersion)
|
|
||||||
}
|
}
|
||||||
|
|
||||||
func sessionAfter(left, right Session) bool {
|
func legacySessionAfter(left, right Session) bool {
|
||||||
return left.AckedOwnershipEpoch > right.AckedOwnershipEpoch ||
|
return compareSnapshotTuple(sessionReference(left), sessionReference(right)) > 0
|
||||||
(left.AckedOwnershipEpoch == right.AckedOwnershipEpoch &&
|
|
||||||
left.AckedSnapshotVersion > right.AckedSnapshotVersion)
|
|
||||||
}
|
|
||||||
|
|
||||||
func clean(value string) bool {
|
|
||||||
return value != "" && strings.TrimSpace(value) == value
|
|
||||||
}
|
}
|
||||||
|
|||||||
@ -2,11 +2,125 @@ package workerruntime
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"crypto/sha256"
|
||||||
"errors"
|
"errors"
|
||||||
"testing"
|
"testing"
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
func TestMemoryStoreRequiresAcknowledgedSnapshotForRuntime(t *testing.T) {
|
||||||
|
now := time.Date(2026, 7, 31, 9, 0, 0, 0, time.UTC)
|
||||||
|
store := newRuntimeStore(t, &now)
|
||||||
|
ctx := context.Background()
|
||||||
|
openControlSession(t, store, "session-a", time.Minute)
|
||||||
|
epoch, err := store.CurrentOwnershipEpoch(ctx)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CurrentOwnershipEpoch(): %v", err)
|
||||||
|
}
|
||||||
|
report := controlReport(now, "session-a", 1, 7, epoch)
|
||||||
|
if err := store.ReplaceRuntime(ctx, report, time.Minute); !errors.Is(err, ErrSnapshotMismatch) {
|
||||||
|
t.Fatalf("ReplaceRuntime(before ACK) error = %v, want ErrSnapshotMismatch", err)
|
||||||
|
}
|
||||||
|
reference := controlReference(7, epoch, "snapshot-7")
|
||||||
|
if err := store.RecordIssuedSnapshot(ctx, reference, time.Minute); err != nil {
|
||||||
|
t.Fatalf("RecordIssuedSnapshot(): %v", err)
|
||||||
|
}
|
||||||
|
if err := store.AcknowledgeSnapshot(ctx, SnapshotAcknowledgement{
|
||||||
|
WorkerID: "worker-a", SessionID: "session-a", Reference: reference, Applied: true,
|
||||||
|
}, time.Minute); err != nil {
|
||||||
|
t.Fatalf("AcknowledgeSnapshot(): %v", err)
|
||||||
|
}
|
||||||
|
if err := store.ReplaceRuntime(ctx, report, time.Minute); err != nil {
|
||||||
|
t.Fatalf("ReplaceRuntime(after ACK): %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMemoryStoreNegativeAcknowledgementFencesDelayedRuntime(t *testing.T) {
|
||||||
|
now := time.Date(2026, 7, 31, 9, 0, 0, 0, time.UTC)
|
||||||
|
store := newRuntimeStore(t, &now)
|
||||||
|
ctx := context.Background()
|
||||||
|
openControlSession(t, store, "session-a", time.Minute)
|
||||||
|
epoch, err := store.CurrentOwnershipEpoch(ctx)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CurrentOwnershipEpoch(): %v", err)
|
||||||
|
}
|
||||||
|
first := controlReference(7, epoch, "snapshot-7")
|
||||||
|
if err := store.RecordIssuedSnapshot(ctx, first, time.Minute); err != nil {
|
||||||
|
t.Fatalf("RecordIssuedSnapshot(first): %v", err)
|
||||||
|
}
|
||||||
|
if err := store.AcknowledgeSnapshot(ctx, SnapshotAcknowledgement{WorkerID: "worker-a", SessionID: "session-a", Reference: first, Applied: true}, time.Minute); err != nil {
|
||||||
|
t.Fatalf("AcknowledgeSnapshot(first): %v", err)
|
||||||
|
}
|
||||||
|
if err := store.ReplaceRuntime(ctx, controlReport(now, "session-a", 1, 7, epoch), time.Minute); err != nil {
|
||||||
|
t.Fatalf("ReplaceRuntime(first): %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
second := controlReference(8, epoch, "snapshot-8")
|
||||||
|
if err := store.RecordIssuedSnapshot(ctx, second, time.Minute); err != nil {
|
||||||
|
t.Fatalf("RecordIssuedSnapshot(second): %v", err)
|
||||||
|
}
|
||||||
|
if err := store.AcknowledgeSnapshot(ctx, SnapshotAcknowledgement{WorkerID: "worker-a", SessionID: "session-a", Reference: second, Applied: false, ErrorCode: "apply_failed"}, time.Minute); err != nil {
|
||||||
|
t.Fatalf("AcknowledgeSnapshot(negative): %v", err)
|
||||||
|
}
|
||||||
|
if err := store.ReplaceRuntime(ctx, controlReport(now, "session-a", 2, 7, epoch), time.Minute); !errors.Is(err, ErrSnapshotMismatch) {
|
||||||
|
t.Fatalf("ReplaceRuntime(delayed): %v, want ErrSnapshotMismatch", err)
|
||||||
|
}
|
||||||
|
if err := store.AcknowledgeSnapshot(ctx, SnapshotAcknowledgement{WorkerID: "worker-a", SessionID: "session-a", Reference: second, Applied: true}, time.Minute); err != nil {
|
||||||
|
t.Fatalf("AcknowledgeSnapshot(recover): %v", err)
|
||||||
|
}
|
||||||
|
if err := store.ReplaceRuntime(ctx, controlReport(now, "session-a", 2, 8, epoch), time.Minute); err != nil {
|
||||||
|
t.Fatalf("ReplaceRuntime(recovered): %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMemoryStoreAcknowledgementReplayPreservesRuntimeFence(t *testing.T) {
|
||||||
|
now := time.Date(2026, 7, 31, 9, 0, 0, 0, time.UTC)
|
||||||
|
store := newRuntimeStore(t, &now)
|
||||||
|
ctx := context.Background()
|
||||||
|
openControlSession(t, store, "session-a", time.Minute)
|
||||||
|
epoch, err := store.CurrentOwnershipEpoch(ctx)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CurrentOwnershipEpoch(): %v", err)
|
||||||
|
}
|
||||||
|
reference := controlReference(7, epoch, "snapshot-7")
|
||||||
|
if err := store.RecordIssuedSnapshot(ctx, reference, time.Minute); err != nil {
|
||||||
|
t.Fatalf("RecordIssuedSnapshot(): %v", err)
|
||||||
|
}
|
||||||
|
ack := SnapshotAcknowledgement{WorkerID: "worker-a", SessionID: "session-a", Reference: reference, Applied: true}
|
||||||
|
if err := store.AcknowledgeSnapshot(ctx, ack, time.Minute); err != nil {
|
||||||
|
t.Fatalf("AcknowledgeSnapshot(first): %v", err)
|
||||||
|
}
|
||||||
|
if err := store.ReplaceRuntime(ctx, controlReport(now, "session-a", 100, 7, epoch), time.Minute); err != nil {
|
||||||
|
t.Fatalf("ReplaceRuntime(): %v", err)
|
||||||
|
}
|
||||||
|
if err := store.AcknowledgeSnapshot(ctx, ack, time.Minute); err != nil {
|
||||||
|
t.Fatalf("AcknowledgeSnapshot(replay): %v", err)
|
||||||
|
}
|
||||||
|
got, err := store.ReadRuntime(ctx, []OwnedProxy{{ProxyID: "proxy-a", WorkerID: "worker-a", OwnershipEpoch: epoch}})
|
||||||
|
if err != nil || len(got) != 1 || !got[0].Fresh {
|
||||||
|
t.Fatalf("ReadRuntime(after ACK replay) = %+v, %v", got, err)
|
||||||
|
}
|
||||||
|
if err := store.ReplaceRuntime(ctx, controlReport(now, "session-a", 99, 7, epoch), time.Minute); !errors.Is(err, ErrStaleReport) {
|
||||||
|
t.Fatalf("ReplaceRuntime(stale): %v, want ErrStaleReport", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestMemoryStoreRejectsConflictingSnapshotReference(t *testing.T) {
|
||||||
|
now := time.Date(2026, 7, 31, 9, 0, 0, 0, time.UTC)
|
||||||
|
store := newRuntimeStore(t, &now)
|
||||||
|
ctx := context.Background()
|
||||||
|
epoch, err := store.CurrentOwnershipEpoch(ctx)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("CurrentOwnershipEpoch(): %v", err)
|
||||||
|
}
|
||||||
|
if err := store.RecordIssuedSnapshot(ctx, controlReference(7, epoch, "first"), time.Minute); err != nil {
|
||||||
|
t.Fatalf("RecordIssuedSnapshot(first): %v", err)
|
||||||
|
}
|
||||||
|
if err := store.RecordIssuedSnapshot(ctx, controlReference(7, epoch, "second"), time.Minute); !errors.Is(err, ErrConflictingSnapshotReference) {
|
||||||
|
t.Fatalf("RecordIssuedSnapshot(conflict): %v, want ErrConflictingSnapshotReference", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
func TestMemoryStoreReplacesSparseRuntimeAndClearsMissingCounters(t *testing.T) {
|
func TestMemoryStoreReplacesSparseRuntimeAndClearsMissingCounters(t *testing.T) {
|
||||||
now := time.Date(2026, 7, 30, 15, 0, 0, 0, time.UTC)
|
now := time.Date(2026, 7, 30, 15, 0, 0, 0, time.UTC)
|
||||||
store := newRuntimeStore(t, &now)
|
store := newRuntimeStore(t, &now)
|
||||||
@ -103,13 +217,13 @@ func TestMemoryStoreRejectsRuntimeBeyondAcknowledgedSnapshot(t *testing.T) {
|
|||||||
WorkerID: "worker-a", SessionID: "session-a", Sequence: 1,
|
WorkerID: "worker-a", SessionID: "session-a", Sequence: 1,
|
||||||
SnapshotVersion: 4, OwnershipEpoch: 9, ObservedAt: now,
|
SnapshotVersion: 4, OwnershipEpoch: 9, ObservedAt: now,
|
||||||
}
|
}
|
||||||
if err := store.ReplaceRuntime(context.Background(), report, time.Minute); !errors.Is(err, ErrStaleReport) {
|
if err := store.ReplaceRuntime(context.Background(), report, time.Minute); !errors.Is(err, ErrSnapshotMismatch) {
|
||||||
t.Fatalf("ReplaceRuntime(ahead snapshot) error = %v, want ErrStaleReport", err)
|
t.Fatalf("ReplaceRuntime(ahead snapshot) error = %v, want ErrSnapshotMismatch", err)
|
||||||
}
|
}
|
||||||
report.SnapshotVersion = 3
|
report.SnapshotVersion = 3
|
||||||
report.OwnershipEpoch = 10
|
report.OwnershipEpoch = 10
|
||||||
if err := store.ReplaceRuntime(context.Background(), report, time.Minute); !errors.Is(err, ErrStaleReport) {
|
if err := store.ReplaceRuntime(context.Background(), report, time.Minute); !errors.Is(err, ErrSnapshotMismatch) {
|
||||||
t.Fatalf("ReplaceRuntime(ahead epoch) error = %v, want ErrStaleReport", err)
|
t.Fatalf("ReplaceRuntime(ahead epoch) error = %v, want ErrSnapshotMismatch", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@ -154,3 +268,26 @@ func registerRuntimeSession(t *testing.T, store *MemoryStore, sessionID string,
|
|||||||
t.Fatalf("ReplaceSession(): %v", err)
|
t.Fatalf("ReplaceSession(): %v", err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func openControlSession(t *testing.T, store *MemoryStore, sessionID string, ttl time.Duration) {
|
||||||
|
t.Helper()
|
||||||
|
if err := store.OpenSession(context.Background(), Session{
|
||||||
|
WorkerID: "worker-a", InstanceID: "instance-a", SessionID: sessionID, Zone: "zone-a", ProtocolVersion: 1,
|
||||||
|
Labels: map[string]string{"region": "test"},
|
||||||
|
}, ttl); err != nil {
|
||||||
|
t.Fatalf("OpenSession(): %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func controlReference(version, epoch uint64, content string) SnapshotReference {
|
||||||
|
return SnapshotReference{
|
||||||
|
WorkerID: "worker-a", Version: version, OwnershipEpoch: epoch, Checksum: sha256.Sum256([]byte(content)),
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func controlReport(now time.Time, sessionID string, sequence, version, epoch uint64) Report {
|
||||||
|
return Report{
|
||||||
|
WorkerID: "worker-a", SessionID: sessionID, Sequence: sequence, SnapshotVersion: version,
|
||||||
|
OwnershipEpoch: epoch, ObservedAt: now, Counters: []Counter{{ProxyID: "proxy-a", Active: 1}},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|||||||
@ -2,26 +2,53 @@ package workerruntime
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"crypto/sha256"
|
||||||
"errors"
|
"errors"
|
||||||
"time"
|
"time"
|
||||||
)
|
)
|
||||||
|
|
||||||
var (
|
var (
|
||||||
ErrInvalidStore = errors.New("invalid worker runtime store")
|
ErrInvalidStore = errors.New("invalid worker runtime store")
|
||||||
ErrInvalidSession = errors.New("invalid worker runtime session")
|
ErrInvalidSession = errors.New("invalid worker runtime session")
|
||||||
ErrInvalidReport = errors.New("invalid worker runtime report")
|
ErrInvalidReport = errors.New("invalid worker runtime report")
|
||||||
ErrInvalidQuery = errors.New("invalid worker runtime query")
|
ErrInvalidQuery = errors.New("invalid worker runtime query")
|
||||||
ErrStaleSession = errors.New("stale worker runtime session")
|
ErrStaleSession = errors.New("stale worker runtime session")
|
||||||
ErrStaleReport = errors.New("stale worker runtime report")
|
ErrStaleReport = errors.New("stale worker runtime report")
|
||||||
ErrConflictingReport = errors.New("conflicting worker runtime report")
|
ErrConflictingReport = errors.New("conflicting worker runtime report")
|
||||||
|
ErrInvalidSnapshotReference = errors.New("invalid worker snapshot reference")
|
||||||
|
ErrInvalidAcknowledgement = errors.New("invalid worker snapshot acknowledgement")
|
||||||
|
ErrSnapshotMismatch = errors.New("worker snapshot does not match acknowledged state")
|
||||||
|
ErrStaleSnapshotReference = errors.New("stale worker snapshot reference")
|
||||||
|
ErrConflictingSnapshotReference = errors.New("conflicting worker snapshot reference")
|
||||||
|
ErrStaleAcknowledgement = errors.New("stale worker snapshot acknowledgement")
|
||||||
)
|
)
|
||||||
|
|
||||||
type Session struct {
|
type Session struct {
|
||||||
WorkerID string
|
WorkerID string
|
||||||
InstanceID string
|
InstanceID string
|
||||||
SessionID string
|
SessionID string
|
||||||
|
Zone string
|
||||||
|
ProtocolVersion uint32
|
||||||
|
Labels map[string]string
|
||||||
AckedSnapshotVersion uint64
|
AckedSnapshotVersion uint64
|
||||||
AckedOwnershipEpoch uint64
|
AckedOwnershipEpoch uint64
|
||||||
|
AckedChecksum [sha256.Size]byte
|
||||||
|
RuntimeEnabled bool
|
||||||
|
}
|
||||||
|
|
||||||
|
type SnapshotReference struct {
|
||||||
|
WorkerID string
|
||||||
|
Version uint64
|
||||||
|
OwnershipEpoch uint64
|
||||||
|
Checksum [sha256.Size]byte
|
||||||
|
}
|
||||||
|
|
||||||
|
type SnapshotAcknowledgement struct {
|
||||||
|
WorkerID string
|
||||||
|
SessionID string
|
||||||
|
Reference SnapshotReference
|
||||||
|
Applied bool
|
||||||
|
ErrorCode string
|
||||||
}
|
}
|
||||||
|
|
||||||
type Counter struct {
|
type Counter struct {
|
||||||
@ -61,6 +88,14 @@ type SessionWriter interface {
|
|||||||
ReplaceSession(context.Context, Session, time.Duration) error
|
ReplaceSession(context.Context, Session, time.Duration) error
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type ControlStore interface {
|
||||||
|
CurrentOwnershipEpoch(context.Context) (uint64, error)
|
||||||
|
OpenSession(context.Context, Session, time.Duration) error
|
||||||
|
RecordIssuedSnapshot(context.Context, SnapshotReference, time.Duration) error
|
||||||
|
AcknowledgeSnapshot(context.Context, SnapshotAcknowledgement, time.Duration) error
|
||||||
|
ReplaceRuntime(context.Context, Report, time.Duration) error
|
||||||
|
}
|
||||||
|
|
||||||
type ReportWriter interface {
|
type ReportWriter interface {
|
||||||
ReplaceRuntime(context.Context, Report, time.Duration) error
|
ReplaceRuntime(context.Context, Report, time.Duration) error
|
||||||
}
|
}
|
||||||
|
|||||||
138
internal/domain/workerruntime/validation.go
Normal file
138
internal/domain/workerruntime/validation.go
Normal file
@ -0,0 +1,138 @@
|
|||||||
|
package workerruntime
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/sha256"
|
||||||
|
"encoding/json"
|
||||||
|
"regexp"
|
||||||
|
"sort"
|
||||||
|
"strings"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
maximumLabels = 32
|
||||||
|
maximumLabelKeyBytes = 64
|
||||||
|
maximumLabelValueBytes = 256
|
||||||
|
maximumLabelTotalBytes = 4 << 10
|
||||||
|
)
|
||||||
|
|
||||||
|
var identifierPattern = regexp.MustCompile(`^[A-Za-z0-9][A-Za-z0-9._:-]{0,127}$`)
|
||||||
|
|
||||||
|
// ValidIdentifier accepts stable Worker, Session and Proxy identifiers.
|
||||||
|
func ValidIdentifier(value string) bool {
|
||||||
|
return identifierPattern.MatchString(value)
|
||||||
|
}
|
||||||
|
|
||||||
|
// NormalizeLabels validates and deep-copies the bounded Worker label set.
|
||||||
|
func NormalizeLabels(labels map[string]string) (map[string]string, error) {
|
||||||
|
if len(labels) > maximumLabels {
|
||||||
|
return nil, ErrInvalidSession
|
||||||
|
}
|
||||||
|
normalized := make(map[string]string, len(labels))
|
||||||
|
total := 0
|
||||||
|
for key, value := range labels {
|
||||||
|
if !ValidIdentifier(key) || len(key) > maximumLabelKeyBytes || value == "" ||
|
||||||
|
strings.TrimSpace(value) != value || strings.IndexByte(value, 0) >= 0 ||
|
||||||
|
len(value) > maximumLabelValueBytes {
|
||||||
|
return nil, ErrInvalidSession
|
||||||
|
}
|
||||||
|
total += len(key) + len(value)
|
||||||
|
if total > maximumLabelTotalBytes {
|
||||||
|
return nil, ErrInvalidSession
|
||||||
|
}
|
||||||
|
normalized[key] = value
|
||||||
|
}
|
||||||
|
return normalized, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// NormalizeSession prepares a new, not-yet-acknowledged session for storage.
|
||||||
|
func NormalizeSession(session Session) (Session, error) {
|
||||||
|
if !ValidIdentifier(session.WorkerID) || !ValidIdentifier(session.InstanceID) ||
|
||||||
|
!ValidIdentifier(session.SessionID) || !ValidIdentifier(session.Zone) ||
|
||||||
|
session.ProtocolVersion == 0 || session.AckedSnapshotVersion != 0 ||
|
||||||
|
session.AckedOwnershipEpoch != 0 || !checksumIsZero(session.AckedChecksum) ||
|
||||||
|
session.RuntimeEnabled {
|
||||||
|
return Session{}, ErrInvalidSession
|
||||||
|
}
|
||||||
|
labels, err := NormalizeLabels(session.Labels)
|
||||||
|
if err != nil {
|
||||||
|
return Session{}, err
|
||||||
|
}
|
||||||
|
session.Labels = labels
|
||||||
|
return session, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func NormalizeSnapshotReference(reference SnapshotReference) (SnapshotReference, error) {
|
||||||
|
if !ValidIdentifier(reference.WorkerID) || reference.Version == 0 || reference.OwnershipEpoch == 0 ||
|
||||||
|
checksumIsZero(reference.Checksum) {
|
||||||
|
return SnapshotReference{}, ErrInvalidSnapshotReference
|
||||||
|
}
|
||||||
|
return reference, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func NormalizeAcknowledgement(acknowledgement SnapshotAcknowledgement) (SnapshotAcknowledgement, error) {
|
||||||
|
if !ValidIdentifier(acknowledgement.WorkerID) || !ValidIdentifier(acknowledgement.SessionID) ||
|
||||||
|
(acknowledgement.ErrorCode != "" && !ValidIdentifier(acknowledgement.ErrorCode)) {
|
||||||
|
return SnapshotAcknowledgement{}, ErrInvalidAcknowledgement
|
||||||
|
}
|
||||||
|
reference, err := NormalizeSnapshotReference(acknowledgement.Reference)
|
||||||
|
if err != nil || reference.WorkerID != acknowledgement.WorkerID {
|
||||||
|
return SnapshotAcknowledgement{}, ErrInvalidAcknowledgement
|
||||||
|
}
|
||||||
|
acknowledgement.Reference = reference
|
||||||
|
return acknowledgement, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
// NormalizeReport returns the canonical sparse replacement and its digest.
|
||||||
|
func NormalizeReport(report Report) (Report, [sha256.Size]byte, error) {
|
||||||
|
if !ValidIdentifier(report.WorkerID) || !ValidIdentifier(report.SessionID) || report.Sequence == 0 ||
|
||||||
|
report.SnapshotVersion == 0 || report.OwnershipEpoch == 0 || report.ObservedAt.IsZero() {
|
||||||
|
return Report{}, [sha256.Size]byte{}, ErrInvalidReport
|
||||||
|
}
|
||||||
|
normalized := report
|
||||||
|
normalized.ObservedAt = report.ObservedAt.UTC()
|
||||||
|
normalized.Counters = append([]Counter(nil), report.Counters...)
|
||||||
|
sort.Slice(normalized.Counters, func(left, right int) bool {
|
||||||
|
return normalized.Counters[left].ProxyID < normalized.Counters[right].ProxyID
|
||||||
|
})
|
||||||
|
seen := make(map[string]struct{}, len(normalized.Counters))
|
||||||
|
for _, counter := range normalized.Counters {
|
||||||
|
if !ValidIdentifier(counter.ProxyID) || counter.Active < 0 || counter.Reserved < 0 {
|
||||||
|
return Report{}, [sha256.Size]byte{}, ErrInvalidReport
|
||||||
|
}
|
||||||
|
if _, exists := seen[counter.ProxyID]; exists {
|
||||||
|
return Report{}, [sha256.Size]byte{}, ErrInvalidReport
|
||||||
|
}
|
||||||
|
seen[counter.ProxyID] = struct{}{}
|
||||||
|
}
|
||||||
|
payload, err := json.Marshal(normalized)
|
||||||
|
if err != nil {
|
||||||
|
return Report{}, [sha256.Size]byte{}, ErrInvalidReport
|
||||||
|
}
|
||||||
|
return normalized, sha256.Sum256(payload), nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func checksumIsZero(checksum [sha256.Size]byte) bool {
|
||||||
|
return checksum == [sha256.Size]byte{}
|
||||||
|
}
|
||||||
|
|
||||||
|
func compareSnapshotTuple(left, right SnapshotReference) int {
|
||||||
|
switch {
|
||||||
|
case left.OwnershipEpoch < right.OwnershipEpoch:
|
||||||
|
return -1
|
||||||
|
case left.OwnershipEpoch > right.OwnershipEpoch:
|
||||||
|
return 1
|
||||||
|
case left.Version < right.Version:
|
||||||
|
return -1
|
||||||
|
case left.Version > right.Version:
|
||||||
|
return 1
|
||||||
|
default:
|
||||||
|
return 0
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func sessionReference(session Session) SnapshotReference {
|
||||||
|
return SnapshotReference{
|
||||||
|
WorkerID: session.WorkerID, Version: session.AckedSnapshotVersion,
|
||||||
|
OwnershipEpoch: session.AckedOwnershipEpoch, Checksum: session.AckedChecksum,
|
||||||
|
}
|
||||||
|
}
|
||||||
63
internal/domain/workerruntime/validation_test.go
Normal file
63
internal/domain/workerruntime/validation_test.go
Normal file
@ -0,0 +1,63 @@
|
|||||||
|
package workerruntime
|
||||||
|
|
||||||
|
import (
|
||||||
|
"crypto/sha256"
|
||||||
|
"errors"
|
||||||
|
"testing"
|
||||||
|
"time"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNormalizeLabelsClonesAndBoundsValues(t *testing.T) {
|
||||||
|
source := map[string]string{"region": "cn-north"}
|
||||||
|
labels, err := NormalizeLabels(source)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("NormalizeLabels(): %v", err)
|
||||||
|
}
|
||||||
|
labels["region"] = "changed"
|
||||||
|
if source["region"] != "cn-north" {
|
||||||
|
t.Fatal("NormalizeLabels() aliases the source map")
|
||||||
|
}
|
||||||
|
if !ValidIdentifier("worker-a:1") || ValidIdentifier("worker a") || ValidIdentifier("") {
|
||||||
|
t.Fatal("ValidIdentifier() accepted or rejected an invalid value")
|
||||||
|
}
|
||||||
|
if _, err := NormalizeLabels(map[string]string{" region": "cn"}); !errors.Is(err, ErrInvalidSession) {
|
||||||
|
t.Fatalf("NormalizeLabels(invalid key) error = %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeReportUsesStableCounterOrdering(t *testing.T) {
|
||||||
|
now := time.Date(2026, 7, 31, 9, 0, 0, 0, time.FixedZone("UTC+8", 8*60*60))
|
||||||
|
base := Report{
|
||||||
|
WorkerID: "worker-a", SessionID: "session-a", Sequence: 1, SnapshotVersion: 7, OwnershipEpoch: 3,
|
||||||
|
ObservedAt: now, Counters: []Counter{{ProxyID: "proxy-b", Active: 2}, {ProxyID: "proxy-a", Reserved: 1}},
|
||||||
|
}
|
||||||
|
normalized, digest, err := NormalizeReport(base)
|
||||||
|
if err != nil || normalized.ObservedAt.Location() != time.UTC || normalized.Counters[0].ProxyID != "proxy-a" {
|
||||||
|
t.Fatalf("NormalizeReport() = %+v, %x, %v", normalized, digest, err)
|
||||||
|
}
|
||||||
|
base.Counters[0], base.Counters[1] = base.Counters[1], base.Counters[0]
|
||||||
|
_, replayDigest, err := NormalizeReport(base)
|
||||||
|
if err != nil || digest != replayDigest {
|
||||||
|
t.Fatalf("NormalizeReport(reordered) digest = %x, %v; want %x", replayDigest, err, digest)
|
||||||
|
}
|
||||||
|
if _, _, err := NormalizeReport(Report{}); !errors.Is(err, ErrInvalidReport) {
|
||||||
|
t.Fatalf("NormalizeReport(invalid) error = %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeSnapshotReferenceAndAcknowledgement(t *testing.T) {
|
||||||
|
reference := SnapshotReference{
|
||||||
|
WorkerID: "worker-a", Version: 7, OwnershipEpoch: 3, Checksum: sha256.Sum256([]byte("snapshot")),
|
||||||
|
}
|
||||||
|
if _, err := NormalizeSnapshotReference(reference); err != nil {
|
||||||
|
t.Fatalf("NormalizeSnapshotReference(): %v", err)
|
||||||
|
}
|
||||||
|
acknowledgement := SnapshotAcknowledgement{WorkerID: "worker-a", SessionID: "session-a", Reference: reference, ErrorCode: "apply_failed"}
|
||||||
|
if _, err := NormalizeAcknowledgement(acknowledgement); err != nil {
|
||||||
|
t.Fatalf("NormalizeAcknowledgement(): %v", err)
|
||||||
|
}
|
||||||
|
acknowledgement.Reference.WorkerID = "worker-b"
|
||||||
|
if _, err := NormalizeAcknowledgement(acknowledgement); !errors.Is(err, ErrInvalidAcknowledgement) {
|
||||||
|
t.Fatalf("NormalizeAcknowledgement(worker mismatch) error = %v", err)
|
||||||
|
}
|
||||||
|
}
|
||||||
Loading…
Reference in New Issue
Block a user