356 lines
10 KiB
Go
356 lines
10 KiB
Go
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)
|
|
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 && 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
|
|
}
|