proxy-pool/internal/domain/workerruntime/memory.go

212 lines
6.5 KiB
Go

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
}