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 } type memorySession struct { value Session expiresAt time.Time } type memoryReport struct { value Report digest [sha256.Size]byte expiresAt time.Time counters map[string]Counter } var ( _ 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, sessions: make(map[string]memorySession), 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 { return ErrInvalidSession } if err := ctx.Err(); err != nil { return err } now := store.now().UTC() if now.IsZero() { return ErrInvalidStore } 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) { return ErrStaleSession } ackAdvanced := exists && !identityChanged && !expired && sessionAfter(session, current.value) if identityChanged || expired || ackAdvanced { delete(store.reports, session.WorkerID) } 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, counterIndex, err := normalizeReport(report) if err != nil { return err } payload, err := json.Marshal(normalized) if err != nil { return ErrInvalidReport } digest := sha256.Sum256(payload) now := store.now().UTC() if now.IsZero() { return ErrInvalidStore } store.mu.Lock() defer store.mu.Unlock() session, exists := store.sessions[report.WorkerID] if !exists || !session.expiresAt.After(now) || session.value.SessionID != report.SessionID { return ErrStaleSession } if report.SnapshotVersion != session.value.AckedSnapshotVersion || report.OwnershipEpoch != session.value.AckedOwnershipEpoch { return ErrStaleReport } if current, exists := store.reports[report.WorkerID]; exists && current.value.SessionID == report.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: return nil } } store.reports[report.WorkerID] = memoryReport{ value: normalized, digest: digest, expiresAt: now.Add(ttl), counters: counterIndex, } session.expiresAt = now.Add(ttl) store.sessions[report.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 !clean(proxy.ProxyID) || !clean(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 := store.now().UTC() if now.IsZero() { return nil, ErrInvalidStore } 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) || 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 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 } 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 } func validSession(session Session) bool { return clean(session.WorkerID) && clean(session.InstanceID) && clean(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 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 }