feat: add worker session runtime domain store
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

This commit is contained in:
youfak 2026-07-31 10:57:39 +08:00
parent 43dec7324a
commit a79d030c82
7 changed files with 755 additions and 83 deletions

View 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) },
}
})
}

View 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)
}
}

View File

@ -3,9 +3,6 @@ package workerruntime
import ( import (
"context" "context"
"crypto/sha256" "crypto/sha256"
"encoding/json"
"sort"
"strings"
"sync" "sync"
"time" "time"
) )
@ -13,7 +10,9 @@ import (
type MemoryStore struct { type MemoryStore struct {
mu sync.Mutex mu sync.Mutex
now func() time.Time now func() time.Time
epoch uint64
sessions map[string]memorySession sessions map[string]memorySession
references map[string]memoryReference
reports map[string]memoryReport reports map[string]memoryReport
} }
@ -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
} }

View File

@ -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}},
}
}

View File

@ -2,6 +2,7 @@ package workerruntime
import ( import (
"context" "context"
"crypto/sha256"
"errors" "errors"
"time" "time"
) )
@ -14,14 +15,40 @@ var (
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
} }

View 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,
}
}

View 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)
}
}