package workerruntime import ( "context" "crypto/sha256" "sync" "time" ) type MemoryStore struct { mu sync.Mutex now func() time.Time epoch uint64 sessions map[string]memorySession references map[string]memoryReference reports map[string]memoryReport } type memorySession struct { value Session expiresAt time.Time } type memoryReference struct { value SnapshotReference expiresAt time.Time } type memoryReport struct { value Report digest [sha256.Size]byte expiresAt time.Time counters map[string]Counter } var ( _ ControlStore = (*MemoryStore)(nil) _ SessionWriter = (*MemoryStore)(nil) _ ReportWriter = (*MemoryStore)(nil) _ RuntimeReader = (*MemoryStore)(nil) ) func NewMemoryStore(now func() time.Time) (*MemoryStore, error) { if now == nil { return nil, ErrInvalidStore } return &MemoryStore{ now: now, epoch: 1, sessions: make(map[string]memorySession), references: make(map[string]memoryReference), reports: make(map[string]memoryReport), }, nil } 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 } 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) delete(store.references, normalized.WorkerID) store.sessions[normalized.WorkerID] = memorySession{value: normalized, expiresAt: now.Add(ttl)} return nil } func (store *MemoryStore) ValidateSession(ctx context.Context, workerID, sessionID string) error { if ctx == nil || store == nil || !ValidIdentifier(workerID) || !ValidIdentifier(sessionID) { 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() session, exists := store.sessions[workerID] if !exists || !session.expiresAt.After(now) || session.value.SessionID != sessionID { return ErrStaleSession } return nil } func (store *MemoryStore) RecordIssuedSnapshot(ctx context.Context, sessionID string, 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 || !ValidIdentifier(sessionID) { return ErrInvalidSnapshotReference } 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 != sessionID { return ErrStaleSession } 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 && legacySessionBefore(session, current.value) { return ErrStaleSession } 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 } func (store *MemoryStore) ReplaceRuntime(ctx context.Context, report Report, ttl time.Duration) error { if ctx == nil || store == nil || ttl <= 0 { return ErrInvalidReport } if err := ctx.Err(); err != nil { return err } normalized, digest, err := NormalizeReport(report) 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.RuntimeEnabled || normalized.SnapshotVersion != session.value.AckedSnapshotVersion || normalized.OwnershipEpoch != session.value.AckedOwnershipEpoch { return ErrSnapshotMismatch } 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[normalized.WorkerID] = memoryReport{ value: normalized, digest: digest, expiresAt: now.Add(ttl), counters: countersByProxy(normalized.Counters), } session.expiresAt = now.Add(ttl) store.sessions[normalized.WorkerID] = session return nil } func (store *MemoryStore) ReadRuntime(ctx context.Context, proxies []OwnedProxy) ([]Snapshot, error) { if ctx == nil || store == nil { return nil, ErrInvalidQuery } if err := ctx.Err(); err != nil { return nil, err } seen := make(map[string]struct{}, len(proxies)) for _, proxy := range proxies { if !ValidIdentifier(proxy.ProxyID) || !ValidIdentifier(proxy.WorkerID) || proxy.OwnershipEpoch == 0 { return nil, ErrInvalidQuery } key := proxy.WorkerID + "\x00" + proxy.ProxyID if _, exists := seen[key]; exists { return nil, ErrInvalidQuery } seen[key] = struct{}{} } now, err := store.currentTime() if err != nil { return nil, err } store.mu.Lock() defer store.mu.Unlock() result := make([]Snapshot, len(proxies)) for index, proxy := range proxies { result[index].ProxyID = proxy.ProxyID session, sessionExists := store.sessions[proxy.WorkerID] report, reportExists := store.reports[proxy.WorkerID] if !sessionExists || !reportExists || !session.expiresAt.After(now) || !report.expiresAt.After(now) || !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 { continue } result[index].Fresh = true if counter, exists := report.counters[proxy.ProxyID]; exists { result[index].Active = counter.Active result[index].Reserved = counter.Reserved result[index].Draining = counter.Draining } } return result, nil } func (store *MemoryStore) currentTime() (time.Time, error) { now := store.now().UTC() if now.IsZero() { return time.Time{}, ErrInvalidStore } return now, nil } 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 legacySessionBefore(left, right Session) bool { return compareSnapshotTuple(sessionReference(left), sessionReference(right)) < 0 } func legacySessionAfter(left, right Session) bool { return compareSnapshotTuple(sessionReference(left), sessionReference(right)) > 0 }