diff --git a/internal/domain/workerruntime/contract_external_test.go b/internal/domain/workerruntime/contract_external_test.go new file mode 100644 index 0000000..2398e35 --- /dev/null +++ b/internal/domain/workerruntime/contract_external_test.go @@ -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) }, + } + }) +} diff --git a/internal/domain/workerruntime/contracttest/contract.go b/internal/domain/workerruntime/contracttest/contract.go new file mode 100644 index 0000000..584fb5d --- /dev/null +++ b/internal/domain/workerruntime/contracttest/contract.go @@ -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) + } +} diff --git a/internal/domain/workerruntime/memory.go b/internal/domain/workerruntime/memory.go index bdcd89f..c72e4de 100644 --- a/internal/domain/workerruntime/memory.go +++ b/internal/domain/workerruntime/memory.go @@ -3,18 +3,17 @@ package workerruntime import ( "context" "crypto/sha256" - "encoding/json" - "sort" - "strings" "sync" "time" ) type MemoryStore struct { - mu sync.Mutex - now func() time.Time - sessions map[string]memorySession - reports map[string]memoryReport + mu sync.Mutex + now func() time.Time + epoch uint64 + sessions map[string]memorySession + references map[string]memoryReference + reports map[string]memoryReport } type memorySession struct { @@ -22,6 +21,11 @@ type memorySession struct { expiresAt time.Time } +type memoryReference struct { + value SnapshotReference + expiresAt time.Time +} + type memoryReport struct { value Report digest [sha256.Size]byte @@ -30,6 +34,7 @@ type memoryReport struct { } var ( + _ ControlStore = (*MemoryStore)(nil) _ SessionWriter = (*MemoryStore)(nil) _ ReportWriter = (*MemoryStore)(nil) _ RuntimeReader = (*MemoryStore)(nil) @@ -40,33 +45,184 @@ func NewMemoryStore(now func() time.Time) (*MemoryStore, error) { return nil, ErrInvalidStore } 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 } -func (store *MemoryStore) ReplaceSession(ctx context.Context, session Session, ttl time.Duration) error { - if ctx == nil || store == nil || !validSession(session) || ttl <= 0 { +func (store *MemoryStore) CurrentOwnershipEpoch(ctx context.Context) (uint64, error) { + 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 } if err := ctx.Err(); err != nil { return err } - now := store.now().UTC() - if now.IsZero() { - return ErrInvalidStore + normalized, err := NormalizeSession(session) + if err != nil { + 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() defer store.mu.Unlock() current, exists := store.sessions[session.WorkerID] identityChanged := exists && (current.value.SessionID != session.SessionID || current.value.InstanceID != session.InstanceID) expired := exists && !current.expiresAt.After(now) - if exists && !identityChanged && !expired && sessionBefore(session, current.value) { + if exists && !identityChanged && !expired && legacySessionBefore(session, current.value) { return ErrStaleSession } - ackAdvanced := exists && !identityChanged && !expired && sessionAfter(session, current.value) + ackAdvanced := exists && !identityChanged && !expired && legacySessionAfter(session, current.value) if identityChanged || expired || ackAdvanced { delete(store.reports, session.WorkerID) } + session.RuntimeEnabled = true store.sessions[session.WorkerID] = memorySession{value: session, expiresAt: now.Add(ttl)} return nil } @@ -78,44 +234,48 @@ func (store *MemoryStore) ReplaceRuntime(ctx context.Context, report Report, ttl if err := ctx.Err(); err != nil { return err } - normalized, counterIndex, err := normalizeReport(report) + normalized, digest, err := NormalizeReport(report) if err != nil { return err } - payload, err := json.Marshal(normalized) + now, err := store.currentTime() if err != nil { - return ErrInvalidReport - } - digest := sha256.Sum256(payload) - now := store.now().UTC() - if now.IsZero() { - return ErrInvalidStore + return err } store.mu.Lock() defer store.mu.Unlock() - session, exists := store.sessions[report.WorkerID] - if !exists || !session.expiresAt.After(now) || session.value.SessionID != report.SessionID { + session, exists := store.sessions[normalized.WorkerID] + if !exists || !session.expiresAt.After(now) || session.value.SessionID != normalized.SessionID { return ErrStaleSession } - if report.SnapshotVersion != session.value.AckedSnapshotVersion || - report.OwnershipEpoch != session.value.AckedOwnershipEpoch { - return ErrStaleReport + if !session.value.RuntimeEnabled || normalized.SnapshotVersion != session.value.AckedSnapshotVersion || + normalized.OwnershipEpoch != session.value.AckedOwnershipEpoch { + 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 { case normalized.Sequence < current.value.Sequence: return ErrStaleReport case normalized.Sequence == current.value.Sequence && digest != current.digest: return ErrConflictingReport 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 } } - store.reports[report.WorkerID] = memoryReport{ - value: normalized, digest: digest, expiresAt: now.Add(ttl), counters: counterIndex, + store.reports[normalized.WorkerID] = memoryReport{ + value: normalized, digest: digest, expiresAt: now.Add(ttl), counters: countersByProxy(normalized.Counters), } session.expiresAt = now.Add(ttl) - store.sessions[report.WorkerID] = session + store.sessions[normalized.WorkerID] = session return nil } @@ -128,7 +288,7 @@ func (store *MemoryStore) ReadRuntime(ctx context.Context, proxies []OwnedProxy) } seen := make(map[string]struct{}, len(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 } key := proxy.WorkerID + "\x00" + proxy.ProxyID @@ -137,9 +297,9 @@ func (store *MemoryStore) ReadRuntime(ctx context.Context, proxies []OwnedProxy) } seen[key] = struct{}{} } - now := store.now().UTC() - if now.IsZero() { - return nil, ErrInvalidStore + now, err := store.currentTime() + if err != nil { + return nil, err } store.mu.Lock() defer store.mu.Unlock() @@ -149,7 +309,7 @@ func (store *MemoryStore) ReadRuntime(ctx context.Context, proxies []OwnedProxy) session, sessionExists := store.sessions[proxy.WorkerID] report, reportExists := store.reports[proxy.WorkerID] 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.OwnershipEpoch != session.value.AckedOwnershipEpoch || report.value.OwnershipEpoch < proxy.OwnershipEpoch { @@ -165,47 +325,31 @@ func (store *MemoryStore) ReadRuntime(ctx context.Context, proxies []OwnedProxy) return result, nil } -func normalizeReport(report Report) (Report, map[string]Counter, error) { - if !clean(report.WorkerID) || !clean(report.SessionID) || report.Sequence == 0 || - report.SnapshotVersion == 0 || report.OwnershipEpoch == 0 || report.ObservedAt.IsZero() { - return Report{}, nil, ErrInvalidReport +func (store *MemoryStore) currentTime() (time.Time, error) { + now := store.now().UTC() + if now.IsZero() { + return time.Time{}, ErrInvalidStore } - 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 - }) - 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 + return now, nil } -func validSession(session Session) bool { - return clean(session.WorkerID) && clean(session.InstanceID) && clean(session.SessionID) && +func countersByProxy(counters []Counter) map[string]Counter { + 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 } -func sessionBefore(left, right Session) bool { - return left.AckedOwnershipEpoch < right.AckedOwnershipEpoch || - (left.AckedOwnershipEpoch == right.AckedOwnershipEpoch && - left.AckedSnapshotVersion < right.AckedSnapshotVersion) +func legacySessionBefore(left, right Session) bool { + return compareSnapshotTuple(sessionReference(left), sessionReference(right)) < 0 } -func sessionAfter(left, right Session) bool { - return left.AckedOwnershipEpoch > right.AckedOwnershipEpoch || - (left.AckedOwnershipEpoch == right.AckedOwnershipEpoch && - left.AckedSnapshotVersion > right.AckedSnapshotVersion) -} - -func clean(value string) bool { - return value != "" && strings.TrimSpace(value) == value +func legacySessionAfter(left, right Session) bool { + return compareSnapshotTuple(sessionReference(left), sessionReference(right)) > 0 } diff --git a/internal/domain/workerruntime/memory_test.go b/internal/domain/workerruntime/memory_test.go index 64760d4..2acbfab 100644 --- a/internal/domain/workerruntime/memory_test.go +++ b/internal/domain/workerruntime/memory_test.go @@ -2,11 +2,125 @@ package workerruntime import ( "context" + "crypto/sha256" "errors" "testing" "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) { now := time.Date(2026, 7, 30, 15, 0, 0, 0, time.UTC) store := newRuntimeStore(t, &now) @@ -103,13 +217,13 @@ func TestMemoryStoreRejectsRuntimeBeyondAcknowledgedSnapshot(t *testing.T) { WorkerID: "worker-a", SessionID: "session-a", Sequence: 1, SnapshotVersion: 4, OwnershipEpoch: 9, ObservedAt: now, } - if err := store.ReplaceRuntime(context.Background(), report, time.Minute); !errors.Is(err, ErrStaleReport) { - t.Fatalf("ReplaceRuntime(ahead snapshot) error = %v, want ErrStaleReport", err) + if err := store.ReplaceRuntime(context.Background(), report, time.Minute); !errors.Is(err, ErrSnapshotMismatch) { + t.Fatalf("ReplaceRuntime(ahead snapshot) error = %v, want ErrSnapshotMismatch", err) } report.SnapshotVersion = 3 report.OwnershipEpoch = 10 - if err := store.ReplaceRuntime(context.Background(), report, time.Minute); !errors.Is(err, ErrStaleReport) { - t.Fatalf("ReplaceRuntime(ahead epoch) error = %v, want ErrStaleReport", err) + if err := store.ReplaceRuntime(context.Background(), report, time.Minute); !errors.Is(err, ErrSnapshotMismatch) { + 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) } } + +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}}, + } +} diff --git a/internal/domain/workerruntime/runtime.go b/internal/domain/workerruntime/runtime.go index f4861ef..a7e423b 100644 --- a/internal/domain/workerruntime/runtime.go +++ b/internal/domain/workerruntime/runtime.go @@ -2,26 +2,53 @@ package workerruntime import ( "context" + "crypto/sha256" "errors" "time" ) var ( - ErrInvalidStore = errors.New("invalid worker runtime store") - ErrInvalidSession = errors.New("invalid worker runtime session") - ErrInvalidReport = errors.New("invalid worker runtime report") - ErrInvalidQuery = errors.New("invalid worker runtime query") - ErrStaleSession = errors.New("stale worker runtime session") - ErrStaleReport = errors.New("stale worker runtime report") - ErrConflictingReport = errors.New("conflicting worker runtime report") + ErrInvalidStore = errors.New("invalid worker runtime store") + ErrInvalidSession = errors.New("invalid worker runtime session") + ErrInvalidReport = errors.New("invalid worker runtime report") + ErrInvalidQuery = errors.New("invalid worker runtime query") + ErrStaleSession = errors.New("stale worker runtime session") + ErrStaleReport = errors.New("stale 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 { WorkerID string InstanceID string SessionID string + Zone string + ProtocolVersion uint32 + Labels map[string]string AckedSnapshotVersion 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 { @@ -61,6 +88,14 @@ type SessionWriter interface { 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 { ReplaceRuntime(context.Context, Report, time.Duration) error } diff --git a/internal/domain/workerruntime/validation.go b/internal/domain/workerruntime/validation.go new file mode 100644 index 0000000..0ea62b6 --- /dev/null +++ b/internal/domain/workerruntime/validation.go @@ -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, + } +} diff --git a/internal/domain/workerruntime/validation_test.go b/internal/domain/workerruntime/validation_test.go new file mode 100644 index 0000000..0f30a39 --- /dev/null +++ b/internal/domain/workerruntime/validation_test.go @@ -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) + } +}