feat: persist worker control state in redis
Some checks are pending
ci / proto (push) Waiting to run
ci / test (ubuntu-latest) (push) Waiting to run
ci / test (windows-latest) (push) Waiting to run
ci / race (push) Waiting to run
ci / integration (push) Waiting to run

This commit is contained in:
youfak 2026-07-31 11:25:22 +08:00
parent a79d030c82
commit a463a8cbd2
11 changed files with 472 additions and 106 deletions

View File

@ -96,7 +96,8 @@ func TestNewBuildsClusterSafeKeyspaceAndHashesDynamicTokens(t *testing.T) {
adapter.keys.expiry, adapter.keys.available, adapter.keys.owners, adapter.keys.expiry, adapter.keys.available, adapter.keys.owners,
adapter.keys.ownerExpiry, adapter.keys.epoch, adapter.keys.inventory, adapter.keys.ownerExpiry, adapter.keys.epoch, adapter.keys.inventory,
adapter.keys.stateInventory, adapter.keys.workerSessions, adapter.keys.stateInventory, adapter.keys.workerSessions,
adapter.keys.workerSessionExpiry, adapter.keys.workerRuntime, adapter.keys.workerSessionExpiry, adapter.keys.workerSnapshots,
adapter.keys.workerSnapshotExpiry, adapter.keys.workerRuntime,
adapter.keys.workerRuntimeExpiry, adapter.keys.workerRuntimeExpiry,
} }
for _, key := range staticKeys { for _, key := range staticKeys {

View File

@ -9,41 +9,45 @@ import (
const redisKeyPrefix = "pp:{activity}:" const redisKeyPrefix = "pp:{activity}:"
type keyspace struct { type keyspace struct {
prefix string prefix string
records string records string
unique string unique string
idkeys string idkeys string
expiry string expiry string
available string available string
owners string owners string
ownerExpiry string ownerExpiry string
epoch string epoch string
inventory string inventory string
stateInventory string stateInventory string
workerSessions string workerSessions string
workerSessionExpiry string workerSessionExpiry string
workerRuntime string workerSnapshots string
workerRuntimeExpiry string workerSnapshotExpiry string
workerRuntime string
workerRuntimeExpiry string
} }
func newKeyspace(namespace string) keyspace { func newKeyspace(namespace string) keyspace {
prefix := redisKeyPrefix + namespace prefix := redisKeyPrefix + namespace
return keyspace{ return keyspace{
prefix: prefix, prefix: prefix,
records: prefix + ":records", records: prefix + ":records",
unique: prefix + ":unique", unique: prefix + ":unique",
idkeys: prefix + ":idkeys", idkeys: prefix + ":idkeys",
expiry: prefix + ":expiry", expiry: prefix + ":expiry",
available: prefix + ":available", available: prefix + ":available",
owners: prefix + ":owners", owners: prefix + ":owners",
ownerExpiry: prefix + ":owner-expiry", ownerExpiry: prefix + ":owner-expiry",
epoch: prefix + ":epoch", epoch: prefix + ":epoch",
inventory: prefix + ":inventory", inventory: prefix + ":inventory",
stateInventory: prefix + ":state-inventory", stateInventory: prefix + ":state-inventory",
workerSessions: prefix + ":worker-sessions", workerSessions: prefix + ":worker-sessions",
workerSessionExpiry: prefix + ":worker-session-expiry", workerSessionExpiry: prefix + ":worker-session-expiry",
workerRuntime: prefix + ":worker-runtime", workerSnapshots: prefix + ":worker-snapshots",
workerRuntimeExpiry: prefix + ":worker-runtime-expiry", workerSnapshotExpiry: prefix + ":worker-snapshot-expiry",
workerRuntime: prefix + ":worker-runtime",
workerRuntimeExpiry: prefix + ":worker-runtime-expiry",
} }
} }

View File

@ -29,8 +29,10 @@ func TestRedisOwnershipLifecycle(t *testing.T) {
assertRedisKeysHaveTTL(t, fixture, assertRedisKeysHaveTTL(t, fixture,
fixture.Adapter.keys.owners, fixture.Adapter.keys.owners,
fixture.Adapter.keys.ownerExpiry, fixture.Adapter.keys.ownerExpiry,
fixture.Adapter.keys.epoch,
) )
if ttl, err := fixture.Client.PTTL(context.Background(), fixture.Adapter.keys.epoch).Result(); err != nil || ttl != -1 {
t.Fatalf("PTTL(%s) = %s, %v; want persistent key", fixture.Adapter.keys.epoch, ttl, err)
}
if current, ok, err := fixture.Adapter.Get(context.Background(), "proxy-a"); err != nil || !ok || current != assigned { if current, ok, err := fixture.Adapter.Get(context.Background(), "proxy-a"); err != nil || !ok || current != assigned {
t.Fatalf("Get(assigned) = %+v, %t, %v", current, ok, err) t.Fatalf("Get(assigned) = %+v, %t, %v", current, ok, err)
} }

View File

@ -2,12 +2,9 @@ package redisactivity
import ( import (
"context" "context"
"crypto/sha256"
"encoding/hex" "encoding/hex"
"encoding/json" "encoding/json"
"sort"
"strconv" "strconv"
"strings"
"time" "time"
"proxy-pool/internal/domain/workerruntime" "proxy-pool/internal/domain/workerruntime"
@ -17,17 +14,43 @@ const runtimeWireVersion = 1
const ( const (
runtimeReplaceSession = "replace_session" runtimeReplaceSession = "replace_session"
runtimeCurrentEpoch = "current_epoch"
runtimeOpenSession = "open_session"
runtimeRecordSnapshot = "record_snapshot"
runtimeAcknowledge = "acknowledge_snapshot"
runtimeReplaceReport = "replace_report" runtimeReplaceReport = "replace_report"
runtimeRead = "read" runtimeRead = "read"
) )
type runtimeSessionWire struct { type runtimeSessionWire struct {
Version int `json:"version"` Version int `json:"version"`
WorkerID string `json:"workerId"` WorkerID string `json:"workerId"`
InstanceID string `json:"instanceId"` InstanceID string `json:"instanceId"`
SessionID string `json:"sessionId"` SessionID string `json:"sessionId"`
AckedSnapshotVersion string `json:"ackedSnapshotVersion"` Zone string `json:"zone,omitempty"`
AckedOwnershipEpoch string `json:"ackedOwnershipEpoch"` ProtocolVersion uint32 `json:"protocolVersion,omitempty"`
Labels map[string]string `json:"labels"`
AckedSnapshotVersion string `json:"ackedSnapshotVersion"`
AckedOwnershipEpoch string `json:"ackedOwnershipEpoch"`
AckedChecksum string `json:"ackedChecksum"`
RuntimeEnabled bool `json:"runtimeEnabled"`
}
type runtimeSnapshotReferenceWire struct {
Version int `json:"version"`
WorkerID string `json:"workerId"`
SnapshotVersion string `json:"snapshotVersion"`
OwnershipEpoch string `json:"ownershipEpoch"`
Checksum string `json:"checksum"`
}
type runtimeAcknowledgementWire struct {
Version int `json:"version"`
WorkerID string `json:"workerId"`
SessionID string `json:"sessionId"`
Reference runtimeSnapshotReferenceWire `json:"reference"`
Applied bool `json:"applied"`
ErrorCode string `json:"errorCode"`
} }
type runtimeCounterWire struct { type runtimeCounterWire struct {
@ -63,11 +86,127 @@ type runtimeSnapshotWire struct {
} }
var ( var (
_ workerruntime.ControlStore = (*Adapter)(nil)
_ workerruntime.SessionWriter = (*Adapter)(nil) _ workerruntime.SessionWriter = (*Adapter)(nil)
_ workerruntime.ReportWriter = (*Adapter)(nil) _ workerruntime.ReportWriter = (*Adapter)(nil)
_ workerruntime.RuntimeReader = (*Adapter)(nil) _ workerruntime.RuntimeReader = (*Adapter)(nil)
) )
func (a *Adapter) CurrentOwnershipEpoch(ctx context.Context) (uint64, error) {
if err := validateRuntimeCall(ctx, a); err != nil {
return 0, err
}
reply, err := a.runRuntime(ctx, runtimeCurrentEpoch, 0, nil, "")
if err != nil {
return 0, err
}
if reply.Status != scriptOK || reply.Record == "" {
return 0, invalidScriptReply("unexpected ownership epoch reply")
}
epoch, err := strconv.ParseUint(reply.Record, 10, 64)
if err != nil || epoch == 0 {
return 0, invalidScriptReply("invalid ownership epoch reply")
}
return epoch, nil
}
func (a *Adapter) OpenSession(ctx context.Context, session workerruntime.Session, ttl time.Duration) error {
if err := validateRuntimeCall(ctx, a); err != nil {
return err
}
normalized, err := workerruntime.NormalizeSession(session)
if err != nil || ttl <= 0 {
return workerruntime.ErrInvalidSession
}
payload, err := json.Marshal(runtimeSessionWire{
Version: runtimeWireVersion, WorkerID: normalized.WorkerID, InstanceID: normalized.InstanceID,
SessionID: normalized.SessionID, Zone: normalized.Zone, ProtocolVersion: normalized.ProtocolVersion,
Labels: normalized.Labels, AckedSnapshotVersion: "0", AckedOwnershipEpoch: "0",
AckedChecksum: "", RuntimeEnabled: false,
})
if err != nil {
return workerruntime.ErrInvalidSession
}
reply, err := a.runRuntime(ctx, runtimeOpenSession, durationMillis(ttl), payload, "")
if err != nil {
return err
}
if reply.Status == scriptOK {
return nil
}
if reply.Status == scriptInvalid {
return workerruntime.ErrInvalidSession
}
return invalidScriptReply("unexpected worker open session reply")
}
func (a *Adapter) RecordIssuedSnapshot(ctx context.Context, reference workerruntime.SnapshotReference, ttl time.Duration) error {
if err := validateRuntimeCall(ctx, a); err != nil {
return err
}
normalized, err := workerruntime.NormalizeSnapshotReference(reference)
if err != nil || ttl <= 0 {
return workerruntime.ErrInvalidSnapshotReference
}
payload, err := json.Marshal(referenceWire(normalized))
if err != nil {
return workerruntime.ErrInvalidSnapshotReference
}
reply, err := a.runRuntime(ctx, runtimeRecordSnapshot, durationMillis(ttl), payload, "")
if err != nil {
return err
}
switch reply.Status {
case scriptOK:
return nil
case scriptInvalid:
return workerruntime.ErrInvalidSnapshotReference
case scriptStale:
return workerruntime.ErrStaleSnapshotReference
case scriptConflict:
return workerruntime.ErrConflictingSnapshotReference
case scriptSnapshotMismatch:
return workerruntime.ErrSnapshotMismatch
default:
return invalidScriptReply("unexpected worker snapshot reference reply")
}
}
func (a *Adapter) AcknowledgeSnapshot(ctx context.Context, acknowledgement workerruntime.SnapshotAcknowledgement, ttl time.Duration) error {
if err := validateRuntimeCall(ctx, a); err != nil {
return err
}
normalized, err := workerruntime.NormalizeAcknowledgement(acknowledgement)
if err != nil || ttl <= 0 {
return workerruntime.ErrInvalidAcknowledgement
}
payload, err := json.Marshal(runtimeAcknowledgementWire{
Version: runtimeWireVersion, WorkerID: normalized.WorkerID, SessionID: normalized.SessionID,
Reference: referenceWire(normalized.Reference), Applied: normalized.Applied, ErrorCode: normalized.ErrorCode,
})
if err != nil {
return workerruntime.ErrInvalidAcknowledgement
}
reply, err := a.runRuntime(ctx, runtimeAcknowledge, durationMillis(ttl), payload, "")
if err != nil {
return err
}
switch reply.Status {
case scriptOK:
return nil
case scriptInvalid:
return workerruntime.ErrInvalidAcknowledgement
case scriptUnavailable:
return workerruntime.ErrStaleSession
case scriptStaleAcknowledgement:
return workerruntime.ErrStaleAcknowledgement
case scriptSnapshotMismatch:
return workerruntime.ErrSnapshotMismatch
default:
return invalidScriptReply("unexpected worker acknowledgement reply")
}
}
func (a *Adapter) ReplaceSession(ctx context.Context, session workerruntime.Session, ttl time.Duration) error { func (a *Adapter) ReplaceSession(ctx context.Context, session workerruntime.Session, ttl time.Duration) error {
if err := validateRuntimeCall(ctx, a); err != nil { if err := validateRuntimeCall(ctx, a); err != nil {
return err return err
@ -81,6 +220,7 @@ func (a *Adapter) ReplaceSession(ctx context.Context, session workerruntime.Sess
InstanceID: session.InstanceID, SessionID: session.SessionID, InstanceID: session.InstanceID, SessionID: session.SessionID,
AckedSnapshotVersion: strconv.FormatUint(session.AckedSnapshotVersion, 10), AckedSnapshotVersion: strconv.FormatUint(session.AckedSnapshotVersion, 10),
AckedOwnershipEpoch: strconv.FormatUint(session.AckedOwnershipEpoch, 10), AckedOwnershipEpoch: strconv.FormatUint(session.AckedOwnershipEpoch, 10),
RuntimeEnabled: true,
}) })
if err != nil { if err != nil {
return workerruntime.ErrInvalidSession return workerruntime.ErrInvalidSession
@ -124,6 +264,8 @@ func (a *Adapter) ReplaceRuntime(ctx context.Context, report workerruntime.Repor
return workerruntime.ErrConflictingReport return workerruntime.ErrConflictingReport
case scriptUnavailable: case scriptUnavailable:
return workerruntime.ErrStaleSession return workerruntime.ErrStaleSession
case scriptSnapshotMismatch:
return workerruntime.ErrSnapshotMismatch
default: default:
return invalidScriptReply("unexpected worker runtime reply") return invalidScriptReply("unexpected worker runtime reply")
} }
@ -139,7 +281,7 @@ func (a *Adapter) ReadRuntime(ctx context.Context, proxies []workerruntime.Owned
wires := make([]runtimeOwnedProxyWire, len(proxies)) wires := make([]runtimeOwnedProxyWire, len(proxies))
seen := make(map[string]struct{}, len(proxies)) seen := make(map[string]struct{}, len(proxies))
for index, proxy := range proxies { for index, proxy := range proxies {
if !runtimeClean(proxy.ProxyID) || !runtimeClean(proxy.WorkerID) || proxy.OwnershipEpoch == 0 { if !workerruntime.ValidIdentifier(proxy.ProxyID) || !workerruntime.ValidIdentifier(proxy.WorkerID) || proxy.OwnershipEpoch == 0 {
return nil, workerruntime.ErrInvalidQuery return nil, workerruntime.ErrInvalidQuery
} }
key := proxy.WorkerID + "\x00" + proxy.ProxyID key := proxy.WorkerID + "\x00" + proxy.ProxyID
@ -180,35 +322,27 @@ func (a *Adapter) ReadRuntime(ctx context.Context, proxies []workerruntime.Owned
} }
func (a *Adapter) encodeRuntimeReport(report workerruntime.Report, ttl time.Duration) ([]byte, string, error) { func (a *Adapter) encodeRuntimeReport(report workerruntime.Report, ttl time.Duration) ([]byte, string, error) {
if ttl <= 0 || !runtimeClean(report.WorkerID) || !runtimeClean(report.SessionID) || normalized, digest, err := workerruntime.NormalizeReport(report)
report.Sequence == 0 || report.SnapshotVersion == 0 || report.OwnershipEpoch == 0 || report.ObservedAt.IsZero() || if err != nil || ttl <= 0 || len(normalized.Counters) > a.options.MaxRuntimeCounters {
len(report.Counters) > a.options.MaxRuntimeCounters {
return nil, "", workerruntime.ErrInvalidReport return nil, "", workerruntime.ErrInvalidReport
} }
counters := append([]workerruntime.Counter(nil), report.Counters...) wires := make([]runtimeCounterWire, len(normalized.Counters))
sort.Slice(counters, func(left, right int) bool { return counters[left].ProxyID < counters[right].ProxyID }) for index, counter := range normalized.Counters {
wires := make([]runtimeCounterWire, len(counters))
for index, counter := range counters {
if !runtimeClean(counter.ProxyID) || counter.Active < 0 || counter.Reserved < 0 ||
(index > 0 && counters[index-1].ProxyID == counter.ProxyID) {
return nil, "", workerruntime.ErrInvalidReport
}
wires[index] = runtimeCounterWire{ wires[index] = runtimeCounterWire{
ProxyID: counter.ProxyID, Active: counter.Active, ProxyID: counter.ProxyID, Active: counter.Active,
Reserved: counter.Reserved, Draining: counter.Draining, Reserved: counter.Reserved, Draining: counter.Draining,
} }
} }
payload, err := json.Marshal(runtimeReportWire{ payload, err := json.Marshal(runtimeReportWire{
Version: runtimeWireVersion, WorkerID: report.WorkerID, SessionID: report.SessionID, Version: runtimeWireVersion, WorkerID: normalized.WorkerID, SessionID: normalized.SessionID,
Sequence: strconv.FormatUint(report.Sequence, 10), Sequence: strconv.FormatUint(normalized.Sequence, 10),
SnapshotVersion: strconv.FormatUint(report.SnapshotVersion, 10), SnapshotVersion: strconv.FormatUint(normalized.SnapshotVersion, 10),
OwnershipEpoch: strconv.FormatUint(report.OwnershipEpoch, 10), OwnershipEpoch: strconv.FormatUint(normalized.OwnershipEpoch, 10),
ObservedAtMS: report.ObservedAt.UTC().UnixMilli(), Counters: wires, ObservedAtMS: normalized.ObservedAt.UnixMilli(), Counters: wires,
}) })
if err != nil { if err != nil {
return nil, "", workerruntime.ErrInvalidReport return nil, "", workerruntime.ErrInvalidReport
} }
digest := sha256.Sum256(payload)
return payload, hex.EncodeToString(digest[:]), nil return payload, hex.EncodeToString(digest[:]), nil
} }
@ -221,7 +355,8 @@ func (a *Adapter) runRuntime(
) (runtimeScriptReply, error) { ) (runtimeScriptReply, error) {
result, err := runScript(ctx, a.client, runtimeScript, []string{ result, err := runScript(ctx, a.client, runtimeScript, []string{
a.keys.workerSessions, a.keys.workerSessionExpiry, a.keys.workerSessions, a.keys.workerSessionExpiry,
a.keys.workerRuntime, a.keys.workerRuntimeExpiry, a.keys.owners, a.keys.workerSnapshots, a.keys.workerSnapshotExpiry,
a.keys.workerRuntime, a.keys.workerRuntimeExpiry, a.keys.owners, a.keys.epoch,
}, operation, ttlMS, a.options.CleanupLimit, string(payload), digest) }, operation, ttlMS, a.options.CleanupLimit, string(payload), digest)
if err != nil { if err != nil {
return runtimeScriptReply{}, err return runtimeScriptReply{}, err
@ -243,6 +378,15 @@ func validateRuntimeCall(ctx context.Context, adapter *Adapter) error {
return nil return nil
} }
func runtimeClean(value string) bool { func referenceWire(reference workerruntime.SnapshotReference) runtimeSnapshotReferenceWire {
return value != "" && strings.TrimSpace(value) == value return runtimeSnapshotReferenceWire{
Version: runtimeWireVersion, WorkerID: reference.WorkerID,
SnapshotVersion: strconv.FormatUint(reference.Version, 10),
OwnershipEpoch: strconv.FormatUint(reference.OwnershipEpoch, 10),
Checksum: hex.EncodeToString(reference.Checksum[:]),
}
}
func runtimeClean(value string) bool {
return workerruntime.ValidIdentifier(value)
} }

View File

@ -0,0 +1,27 @@
//go:build integration
package redisactivity
import (
"context"
"testing"
"time"
"proxy-pool/internal/domain/workerruntime/contracttest"
)
func TestRedisWorkerControlStoreContract(t *testing.T) {
contracttest.Run(t, func(*testing.T) contracttest.Fixture {
fixture := newRedisTestFixture(t)
now := redisTestNow()
seedRedisAvailable(t, fixture.Adapter, "provider-a", now, now.Add(time.Second), 2*time.Minute,
testProxy("proxy-a", "192.0.2.10"))
if _, err := fixture.Adapter.Assign(context.Background(), now.Add(2*time.Second), "proxy-a", "worker-a", time.Minute); err != nil {
t.Fatalf("Assign(): %v", err)
}
return contracttest.Fixture{
Store: fixture.Adapter, Reader: fixture.Adapter, TTL: 100 * time.Millisecond,
Advance: time.Sleep,
}
})
}

View File

@ -131,8 +131,8 @@ func TestRedisWorkerRuntimeRejectsEmptyReportBeyondAcknowledgedSnapshot(t *testi
}, },
} { } {
t.Run(name, func(t *testing.T) { t.Run(name, func(t *testing.T) {
if err := fixture.Adapter.ReplaceRuntime(context.Background(), report, time.Minute); !errors.Is(err, workerruntime.ErrStaleReport) { if err := fixture.Adapter.ReplaceRuntime(context.Background(), report, time.Minute); !errors.Is(err, workerruntime.ErrSnapshotMismatch) {
t.Fatalf("ReplaceRuntime() error = %v, want ErrStaleReport", err) t.Fatalf("ReplaceRuntime() error = %v, want ErrSnapshotMismatch", err)
} }
}) })
} }

View File

@ -17,16 +17,18 @@ import (
type scriptStatus string type scriptStatus string
const ( const (
scriptOK scriptStatus = "ok" scriptOK scriptStatus = "ok"
scriptInvalid scriptStatus = "invalid" scriptInvalid scriptStatus = "invalid"
scriptNotFound scriptStatus = "not_found" scriptNotFound scriptStatus = "not_found"
scriptConflict scriptStatus = "conflict" scriptConflict scriptStatus = "conflict"
scriptStale scriptStatus = "stale" scriptStale scriptStatus = "stale"
scriptUnavailable scriptStatus = "unavailable" scriptUnavailable scriptStatus = "unavailable"
scriptInsufficient scriptStatus = "insufficient" scriptInsufficient scriptStatus = "insufficient"
scriptAlreadyOwned scriptStatus = "already_owned" scriptAlreadyOwned scriptStatus = "already_owned"
scriptNotDraining scriptStatus = "not_draining" scriptNotDraining scriptStatus = "not_draining"
scriptDrainNotReady scriptStatus = "drain_not_ready" scriptDrainNotReady scriptStatus = "drain_not_ready"
scriptSnapshotMismatch scriptStatus = "snapshot_mismatch"
scriptStaleAcknowledgement scriptStatus = "stale_acknowledgement"
) )
type upsertScriptReply struct { type upsertScriptReply struct {
@ -77,6 +79,7 @@ type statusScriptInventory struct {
type runtimeScriptReply struct { type runtimeScriptReply struct {
Status scriptStatus `json:"status"` Status scriptStatus `json:"status"`
Snapshots []runtimeSnapshotWire `json:"snapshots"` Snapshots []runtimeSnapshotWire `json:"snapshots"`
Record string `json:"record,omitempty"`
} }
type capacityScriptReply struct { type capacityScriptReply struct {

View File

@ -214,6 +214,7 @@ if operation == 'assign' then
expires_at_ms = tonumber(record.usableUntilMs) expires_at_ms = tonumber(record.usableUntilMs)
end end
local next_epoch = redis.call('INCR', epoch_key) local next_epoch = redis.call('INCR', epoch_key)
redis.call('PERSIST', epoch_key)
local assignment = { local assignment = {
version = 1, version = 1,
proxyId = proxy_id, proxyId = proxy_id,
@ -228,7 +229,6 @@ if operation == 'assign' then
redis.call('ZADD', owner_expiry_key, expires_at_ms, proxy_id) redis.call('ZADD', owner_expiry_key, expires_at_ms, proxy_id)
touch(owners_key, tonumber(record.expiresAtMs)) touch(owners_key, tonumber(record.expiresAtMs))
touch(owner_expiry_key, tonumber(record.expiresAtMs)) touch(owner_expiry_key, tonumber(record.expiresAtMs))
touch(epoch_key, tonumber(record.expiresAtMs))
record.ownerWorkerId = worker_id record.ownerWorkerId = worker_id
redis.call('HSET', records_key, proxy_id, cjson.encode(record)) redis.call('HSET', records_key, proxy_id, cjson.encode(record))
remove_available(proxy_id, record) remove_available(proxy_id, record)
@ -264,7 +264,6 @@ if operation == 'renew' then
redis.call('ZADD', owner_expiry_key, expires_at_ms, proxy_id) redis.call('ZADD', owner_expiry_key, expires_at_ms, proxy_id)
touch(owners_key, tonumber(record.expiresAtMs)) touch(owners_key, tonumber(record.expiresAtMs))
touch(owner_expiry_key, tonumber(record.expiresAtMs)) touch(owner_expiry_key, tonumber(record.expiresAtMs))
touch(epoch_key, tonumber(record.expiresAtMs))
return finish({status = 'ok', record = encoded}) return finish({status = 'ok', record = encoded})
end end

View File

@ -1,8 +1,11 @@
local sessions_key = KEYS[1] local sessions_key = KEYS[1]
local session_expiry_key = KEYS[2] local session_expiry_key = KEYS[2]
local runtime_key = KEYS[3] local snapshots_key = KEYS[3]
local runtime_expiry_key = KEYS[4] local snapshot_expiry_key = KEYS[4]
local owners_key = KEYS[5] local runtime_key = KEYS[5]
local runtime_expiry_key = KEYS[6]
local owners_key = KEYS[7]
local epoch_key = KEYS[8]
local operation = ARGV[1] local operation = ARGV[1]
local ttl_ms = tonumber(ARGV[2]) local ttl_ms = tonumber(ARGV[2])
@ -10,11 +13,15 @@ local cleanup_limit = tonumber(ARGV[3])
local payload = ARGV[4] local payload = ARGV[4]
local digest = ARGV[5] local digest = ARGV[5]
local function reply(status, snapshots) local function reply(status, snapshots, record)
if snapshots then if not snapshots then
return cjson.encode({status = status, snapshots = snapshots}) local suffix = ''
if record then
suffix = ',"record":' .. cjson.encode(record)
end
return '{"status":' .. cjson.encode(status) .. ',"snapshots":[]' .. suffix .. '}'
end end
return '{"status":' .. cjson.encode(status) .. ',"snapshots":[]}' return cjson.encode({status = status, snapshots = snapshots, record = record})
end end
local function now_ms() local function now_ms()
@ -61,13 +68,50 @@ local function cleanup(now)
redis.call('HDEL', runtime_key, worker_id) redis.call('HDEL', runtime_key, worker_id)
redis.call('ZREM', runtime_expiry_key, worker_id) redis.call('ZREM', runtime_expiry_key, worker_id)
end end
local expired_snapshots = redis.call('ZRANGEBYSCORE', snapshot_expiry_key, '-inf', now, 'LIMIT', 0, cleanup_limit)
for _, worker_id in ipairs(expired_snapshots) do
redis.call('HDEL', snapshots_key, worker_id)
redis.call('ZREM', snapshot_expiry_key, worker_id)
end
end end
local function valid_session(value) local function valid_session(value)
return value and value.version == 1 and type(value.workerId) == 'string' and value.workerId ~= '' and return value and value.version == 1 and type(value.workerId) == 'string' and value.workerId ~= '' and
type(value.instanceId) == 'string' and value.instanceId ~= '' and type(value.instanceId) == 'string' and value.instanceId ~= '' and
type(value.sessionId) == 'string' and value.sessionId ~= '' and type(value.sessionId) == 'string' and value.sessionId ~= '' and
valid_uint(value.ackedSnapshotVersion) and valid_uint(value.ackedOwnershipEpoch) type(value.ackedSnapshotVersion) == 'string' and type(value.ackedOwnershipEpoch) == 'string' and
type(value.ackedChecksum) == 'string' and type(value.runtimeEnabled) == 'boolean'
end
local function valid_legacy_session(value)
return valid_session(value) and valid_uint(value.ackedSnapshotVersion) and valid_uint(value.ackedOwnershipEpoch)
end
local function valid_control_session(value)
if not valid_session(value) or type(value.zone) ~= 'string' or value.zone == '' or
type(value.protocolVersion) ~= 'number' or value.protocolVersion <= 0 or type(value.labels) ~= 'table' then
return false
end
if value.ackedSnapshotVersion == '0' and value.ackedOwnershipEpoch == '0' and value.ackedChecksum == '' then
return true
end
return valid_uint(value.ackedSnapshotVersion) and valid_uint(value.ackedOwnershipEpoch) and
string.len(value.ackedChecksum) == 64 and string.match(value.ackedChecksum, '^[0-9a-f]+$') ~= nil
end
local function valid_reference(value)
return value and value.version == 1 and type(value.workerId) == 'string' and value.workerId ~= '' and
valid_uint(value.snapshotVersion) and valid_uint(value.ownershipEpoch) and
type(value.checksum) == 'string' and string.len(value.checksum) == 64 and
string.match(value.checksum, '^[0-9a-f]+$') ~= nil
end
local function compare_reference(left, right)
local epoch_order = compare_uint(left.ownershipEpoch, right.ownershipEpoch)
if epoch_order ~= 0 then
return epoch_order
end
return compare_uint(left.snapshotVersion, right.snapshotVersion)
end end
local function valid_owner(value, worker_id, ownership_epoch, now) local function valid_owner(value, worker_id, ownership_epoch, now)
@ -80,6 +124,139 @@ end
local now = now_ms() local now = now_ms()
cleanup(now) cleanup(now)
if operation == 'current_epoch' then
local epoch = redis.call('GET', epoch_key)
if not epoch then
epoch = '1'
redis.call('SET', epoch_key, epoch)
end
redis.call('PERSIST', epoch_key)
if not valid_uint(epoch) then
return reply('invalid')
end
return reply('ok', nil, epoch)
end
if operation == 'open_session' then
if not ttl_ms or ttl_ms <= 0 then
return reply('invalid')
end
local session = decode_table(payload)
if not valid_control_session(session) or session.ackedSnapshotVersion ~= '0' or
session.ackedOwnershipEpoch ~= '0' or session.ackedChecksum ~= '' or session.runtimeEnabled then
return reply('invalid')
end
redis.call('HDEL', runtime_key, session.workerId)
redis.call('ZREM', runtime_expiry_key, session.workerId)
session.expiresAtMs = now + ttl_ms
redis.call('HSET', sessions_key, session.workerId, cjson.encode(session))
redis.call('ZADD', session_expiry_key, session.expiresAtMs, session.workerId)
return reply('ok')
end
if operation == 'record_snapshot' then
if not ttl_ms or ttl_ms <= 0 then
return reply('invalid')
end
local reference = decode_table(payload)
if not valid_reference(reference) then
return reply('invalid')
end
local epoch = redis.call('GET', epoch_key)
if not epoch then
epoch = '1'
redis.call('SET', epoch_key, epoch)
end
redis.call('PERSIST', epoch_key)
if not valid_uint(epoch) or compare_uint(reference.ownershipEpoch, epoch) ~= 0 then
return reply('snapshot_mismatch')
end
local current = decode_table(redis.call('HGET', snapshots_key, reference.workerId))
if current and type(current.expiresAtMs) == 'number' and current.expiresAtMs > now and valid_reference(current) then
local ordering = compare_reference(reference, current)
if ordering < 0 then
return reply('stale')
end
if ordering == 0 and reference.checksum ~= current.checksum then
return reply('conflict')
end
end
reference.expiresAtMs = now + ttl_ms
redis.call('HSET', snapshots_key, reference.workerId, cjson.encode(reference))
redis.call('ZADD', snapshot_expiry_key, reference.expiresAtMs, reference.workerId)
return reply('ok')
end
if operation == 'acknowledge_snapshot' then
if not ttl_ms or ttl_ms <= 0 then
return reply('invalid')
end
local acknowledgement = decode_table(payload)
if not acknowledgement or acknowledgement.version ~= 1 or type(acknowledgement.workerId) ~= 'string' or
acknowledgement.workerId == '' or type(acknowledgement.sessionId) ~= 'string' or acknowledgement.sessionId == '' or
type(acknowledgement.applied) ~= 'boolean' or type(acknowledgement.errorCode) ~= 'string' or
not valid_reference(acknowledgement.reference) or acknowledgement.reference.workerId ~= acknowledgement.workerId then
return reply('invalid')
end
local session = decode_table(redis.call('HGET', sessions_key, acknowledgement.workerId))
if not valid_control_session(session) or session.sessionId ~= acknowledgement.sessionId or
type(session.expiresAtMs) ~= 'number' or session.expiresAtMs <= now then
return reply('unavailable')
end
if session.ackedSnapshotVersion ~= '0' then
local previous = {ownershipEpoch = session.ackedOwnershipEpoch, snapshotVersion = session.ackedSnapshotVersion}
local acknowledged = compare_reference(acknowledgement.reference, previous)
if acknowledged < 0 then
return reply('stale_acknowledgement')
end
if acknowledged == 0 and acknowledgement.reference.checksum ~= session.ackedChecksum then
return reply('snapshot_mismatch')
end
end
local current = decode_table(redis.call('HGET', snapshots_key, acknowledgement.workerId))
if not valid_reference(current) or type(current.expiresAtMs) ~= 'number' or current.expiresAtMs <= now then
return reply('snapshot_mismatch')
end
local ordering = compare_reference(acknowledgement.reference, current)
if ordering < 0 then
return reply('stale_acknowledgement')
end
if ordering > 0 or acknowledgement.reference.checksum ~= current.checksum then
return reply('snapshot_mismatch')
end
if not acknowledgement.applied then
redis.call('HDEL', runtime_key, acknowledgement.workerId)
redis.call('ZREM', runtime_expiry_key, acknowledgement.workerId)
session.runtimeEnabled = false
session.expiresAtMs = now + ttl_ms
redis.call('HSET', sessions_key, acknowledgement.workerId, cjson.encode(session))
redis.call('ZADD', session_expiry_key, session.expiresAtMs, acknowledgement.workerId)
return reply('ok')
end
if session.ackedSnapshotVersion ~= '0' and
compare_reference(acknowledgement.reference, {ownershipEpoch = session.ackedOwnershipEpoch, snapshotVersion = session.ackedSnapshotVersion}) == 0 then
if not session.runtimeEnabled then
redis.call('HDEL', runtime_key, acknowledgement.workerId)
redis.call('ZREM', runtime_expiry_key, acknowledgement.workerId)
session.runtimeEnabled = true
end
session.expiresAtMs = now + ttl_ms
redis.call('HSET', sessions_key, acknowledgement.workerId, cjson.encode(session))
redis.call('ZADD', session_expiry_key, session.expiresAtMs, acknowledgement.workerId)
return reply('ok')
end
redis.call('HDEL', runtime_key, acknowledgement.workerId)
redis.call('ZREM', runtime_expiry_key, acknowledgement.workerId)
session.ackedSnapshotVersion = acknowledgement.reference.snapshotVersion
session.ackedOwnershipEpoch = acknowledgement.reference.ownershipEpoch
session.ackedChecksum = acknowledgement.reference.checksum
session.runtimeEnabled = true
session.expiresAtMs = now + ttl_ms
redis.call('HSET', sessions_key, acknowledgement.workerId, cjson.encode(session))
redis.call('ZADD', session_expiry_key, session.expiresAtMs, acknowledgement.workerId)
return reply('ok')
end
if operation == 'replace_session' then if operation == 'replace_session' then
if not ttl_ms or ttl_ms <= 0 then if not ttl_ms or ttl_ms <= 0 then
return reply('invalid') return reply('invalid')
@ -89,7 +266,7 @@ if operation == 'replace_session' then
return reply('invalid') return reply('invalid')
end end
local current = decode_table(redis.call('HGET', sessions_key, session.workerId)) local current = decode_table(redis.call('HGET', sessions_key, session.workerId))
if current and valid_session(current) and current.sessionId == session.sessionId and if current and valid_legacy_session(current) and current.sessionId == session.sessionId and
current.instanceId == session.instanceId and type(current.expiresAtMs) == 'number' and current.instanceId == session.instanceId and type(current.expiresAtMs) == 'number' and
current.expiresAtMs > now then current.expiresAtMs > now then
local epoch_order = compare_uint(session.ackedOwnershipEpoch, current.ackedOwnershipEpoch) local epoch_order = compare_uint(session.ackedOwnershipEpoch, current.ackedOwnershipEpoch)
@ -127,9 +304,10 @@ if operation == 'replace_report' then
type(session.expiresAtMs) ~= 'number' or session.expiresAtMs <= now then type(session.expiresAtMs) ~= 'number' or session.expiresAtMs <= now then
return reply('unavailable') return reply('unavailable')
end end
if report.snapshotVersion ~= session.ackedSnapshotVersion or if not valid_uint(session.ackedSnapshotVersion) or not valid_uint(session.ackedOwnershipEpoch) or
not session.runtimeEnabled or report.snapshotVersion ~= session.ackedSnapshotVersion or
report.ownershipEpoch ~= session.ackedOwnershipEpoch then report.ownershipEpoch ~= session.ackedOwnershipEpoch then
return reply('stale') return reply('snapshot_mismatch')
end end
local current = decode_table(redis.call('HGET', runtime_key, report.workerId)) local current = decode_table(redis.call('HGET', runtime_key, report.workerId))
if current and current.sessionId == report.sessionId and valid_uint(current.sequence) then if current and current.sessionId == report.sessionId and valid_uint(current.sequence) then
@ -139,6 +317,12 @@ if operation == 'replace_report' then
end end
if ordering == 0 then if ordering == 0 then
if current.digest == digest then if current.digest == digest then
current.expiresAtMs = now + ttl_ms
redis.call('HSET', runtime_key, report.workerId, cjson.encode(current))
redis.call('ZADD', runtime_expiry_key, current.expiresAtMs, report.workerId)
session.expiresAtMs = now + ttl_ms
redis.call('HSET', sessions_key, report.workerId, cjson.encode(session))
redis.call('ZADD', session_expiry_key, session.expiresAtMs, report.workerId)
return reply('ok') return reply('ok')
end end
return reply('conflict') return reply('conflict')
@ -192,7 +376,8 @@ if operation == 'read' then
local session = decode_table(redis.call('HGET', sessions_key, query.workerId)) local session = decode_table(redis.call('HGET', sessions_key, query.workerId))
local report = decode_table(redis.call('HGET', runtime_key, query.workerId)) local report = decode_table(redis.call('HGET', runtime_key, query.workerId))
cached = {fresh = false, counters = {}} cached = {fresh = false, counters = {}}
if valid_session(session) and session.workerId == query.workerId and report and if valid_session(session) and valid_uint(session.ackedSnapshotVersion) and
valid_uint(session.ackedOwnershipEpoch) and session.runtimeEnabled and session.workerId == query.workerId and report and
report.workerId == query.workerId and report.sessionId == session.sessionId and report.workerId == query.workerId and report.sessionId == session.sessionId and
type(session.expiresAtMs) == 'number' and session.expiresAtMs > now and type(session.expiresAtMs) == 'number' and session.expiresAtMs > now and
type(report.expiresAtMs) == 'number' and report.expiresAtMs > now and type(report.expiresAtMs) == 'number' and report.expiresAtMs > now and

View File

@ -17,7 +17,7 @@ func TestMemoryStoreContract(t *testing.T) {
t.Fatalf("NewMemoryStore(): %v", err) t.Fatalf("NewMemoryStore(): %v", err)
} }
return contracttest.Fixture{ return contracttest.Fixture{
Store: store, Reader: store, Store: store, Reader: store, TTL: time.Minute,
Advance: func(duration time.Duration) { now = now.Add(duration) }, Advance: func(duration time.Duration) { now = now.Add(duration) },
} }
}) })

View File

@ -13,6 +13,7 @@ import (
type Fixture struct { type Fixture struct {
Store workerruntime.ControlStore Store workerruntime.ControlStore
Reader workerruntime.RuntimeReader Reader workerruntime.RuntimeReader
TTL time.Duration
Advance func(time.Duration) Advance func(time.Duration)
} }
@ -28,62 +29,62 @@ func Run(t *testing.T, factory Factory) {
func runLifecycle(t *testing.T, fixture Fixture) { func runLifecycle(t *testing.T, fixture Fixture) {
t.Helper() t.Helper()
ctx := context.Background() ctx := context.Background()
open(t, fixture.Store) open(t, fixture.Store, fixture.TTL)
epoch, err := fixture.Store.CurrentOwnershipEpoch(ctx) epoch, err := fixture.Store.CurrentOwnershipEpoch(ctx)
if err != nil { if err != nil {
t.Fatalf("CurrentOwnershipEpoch(): %v", err) t.Fatalf("CurrentOwnershipEpoch(): %v", err)
} }
reference := snapshot(7, epoch, "snapshot-7") reference := snapshot(7, epoch, "snapshot-7")
report := runtimeReport(1, 7, epoch) report := runtimeReport(1, 7, epoch)
if err := fixture.Store.ReplaceRuntime(ctx, report, time.Minute); !errors.Is(err, workerruntime.ErrSnapshotMismatch) { if err := fixture.Store.ReplaceRuntime(ctx, report, fixture.TTL); !errors.Is(err, workerruntime.ErrSnapshotMismatch) {
t.Fatalf("ReplaceRuntime(before ACK) error = %v", err) t.Fatalf("ReplaceRuntime(before ACK) error = %v", err)
} }
if err := fixture.Store.RecordIssuedSnapshot(ctx, reference, time.Minute); err != nil { if err := fixture.Store.RecordIssuedSnapshot(ctx, reference, fixture.TTL); err != nil {
t.Fatalf("RecordIssuedSnapshot(): %v", err) t.Fatalf("RecordIssuedSnapshot(): %v", err)
} }
ack := workerruntime.SnapshotAcknowledgement{WorkerID: "worker-a", SessionID: "session-a", Reference: reference, Applied: true} ack := workerruntime.SnapshotAcknowledgement{WorkerID: "worker-a", SessionID: "session-a", Reference: reference, Applied: true}
if err := fixture.Store.AcknowledgeSnapshot(ctx, ack, time.Minute); err != nil { if err := fixture.Store.AcknowledgeSnapshot(ctx, ack, fixture.TTL); err != nil {
t.Fatalf("AcknowledgeSnapshot(): %v", err) t.Fatalf("AcknowledgeSnapshot(): %v", err)
} }
if err := fixture.Store.ReplaceRuntime(ctx, report, time.Minute); err != nil { if err := fixture.Store.ReplaceRuntime(ctx, report, fixture.TTL); err != nil {
t.Fatalf("ReplaceRuntime(after ACK): %v", err) t.Fatalf("ReplaceRuntime(after ACK): %v", err)
} }
assertFresh(t, fixture.Reader, epoch, true) assertFresh(t, fixture.Reader, epoch, true)
if err := fixture.Store.AcknowledgeSnapshot(ctx, ack, time.Minute); err != nil { if err := fixture.Store.AcknowledgeSnapshot(ctx, ack, fixture.TTL); err != nil {
t.Fatalf("AcknowledgeSnapshot(replay): %v", err) t.Fatalf("AcknowledgeSnapshot(replay): %v", err)
} }
if err := fixture.Store.ReplaceRuntime(ctx, runtimeReport(0, 7, epoch), time.Minute); !errors.Is(err, workerruntime.ErrInvalidReport) { if err := fixture.Store.ReplaceRuntime(ctx, runtimeReport(0, 7, epoch), fixture.TTL); !errors.Is(err, workerruntime.ErrInvalidReport) {
t.Fatalf("ReplaceRuntime(invalid sequence) error = %v", err) t.Fatalf("ReplaceRuntime(invalid sequence) error = %v", err)
} }
fixture.Advance(2 * time.Minute) fixture.Advance(2 * fixture.TTL)
assertFresh(t, fixture.Reader, epoch, false) assertFresh(t, fixture.Reader, epoch, false)
} }
func runNegativeAck(t *testing.T, fixture Fixture) { func runNegativeAck(t *testing.T, fixture Fixture) {
t.Helper() t.Helper()
ctx := context.Background() ctx := context.Background()
open(t, fixture.Store) open(t, fixture.Store, fixture.TTL)
epoch, err := fixture.Store.CurrentOwnershipEpoch(ctx) epoch, err := fixture.Store.CurrentOwnershipEpoch(ctx)
if err != nil { if err != nil {
t.Fatalf("CurrentOwnershipEpoch(): %v", err) t.Fatalf("CurrentOwnershipEpoch(): %v", err)
} }
first := snapshot(7, epoch, "snapshot-7") first := snapshot(7, epoch, "snapshot-7")
if err := fixture.Store.RecordIssuedSnapshot(ctx, first, time.Minute); err != nil { if err := fixture.Store.RecordIssuedSnapshot(ctx, first, fixture.TTL); err != nil {
t.Fatalf("RecordIssuedSnapshot(first): %v", err) 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 { if err := fixture.Store.AcknowledgeSnapshot(ctx, workerruntime.SnapshotAcknowledgement{WorkerID: "worker-a", SessionID: "session-a", Reference: first, Applied: true}, fixture.TTL); err != nil {
t.Fatalf("AcknowledgeSnapshot(first): %v", err) t.Fatalf("AcknowledgeSnapshot(first): %v", err)
} }
second := snapshot(8, epoch, "snapshot-8") second := snapshot(8, epoch, "snapshot-8")
if err := fixture.Store.RecordIssuedSnapshot(ctx, second, time.Minute); err != nil { if err := fixture.Store.RecordIssuedSnapshot(ctx, second, fixture.TTL); err != nil {
t.Fatalf("RecordIssuedSnapshot(second): %v", err) t.Fatalf("RecordIssuedSnapshot(second): %v", err)
} }
if err := fixture.Store.AcknowledgeSnapshot(ctx, workerruntime.SnapshotAcknowledgement{ if err := fixture.Store.AcknowledgeSnapshot(ctx, workerruntime.SnapshotAcknowledgement{
WorkerID: "worker-a", SessionID: "session-a", Reference: second, ErrorCode: "apply_failed", WorkerID: "worker-a", SessionID: "session-a", Reference: second, ErrorCode: "apply_failed",
}, time.Minute); err != nil { }, fixture.TTL); err != nil {
t.Fatalf("AcknowledgeSnapshot(negative): %v", err) t.Fatalf("AcknowledgeSnapshot(negative): %v", err)
} }
if err := fixture.Store.ReplaceRuntime(ctx, runtimeReport(1, 7, epoch), time.Minute); !errors.Is(err, workerruntime.ErrSnapshotMismatch) { if err := fixture.Store.ReplaceRuntime(ctx, runtimeReport(1, 7, epoch), fixture.TTL); !errors.Is(err, workerruntime.ErrSnapshotMismatch) {
t.Fatalf("ReplaceRuntime(delayed): %v", err) t.Fatalf("ReplaceRuntime(delayed): %v", err)
} }
} }
@ -91,17 +92,17 @@ func runNegativeAck(t *testing.T, fixture Fixture) {
func newFixture(t *testing.T, factory Factory) Fixture { func newFixture(t *testing.T, factory Factory) Fixture {
t.Helper() t.Helper()
fixture := factory(t) fixture := factory(t)
if fixture.Store == nil || fixture.Reader == nil || fixture.Advance == nil { if fixture.Store == nil || fixture.Reader == nil || fixture.TTL <= 0 || fixture.Advance == nil {
t.Fatal("contract fixture is incomplete") t.Fatal("contract fixture is incomplete")
} }
return fixture return fixture
} }
func open(t *testing.T, store workerruntime.ControlStore) { func open(t *testing.T, store workerruntime.ControlStore, ttl time.Duration) {
t.Helper() t.Helper()
err := store.OpenSession(context.Background(), workerruntime.Session{ err := store.OpenSession(context.Background(), workerruntime.Session{
WorkerID: "worker-a", InstanceID: "instance-a", SessionID: "session-a", Zone: "zone-a", ProtocolVersion: 1, WorkerID: "worker-a", InstanceID: "instance-a", SessionID: "session-a", Zone: "zone-a", ProtocolVersion: 1,
}, time.Minute) }, ttl)
if err != nil { if err != nil {
t.Fatalf("OpenSession(): %v", err) t.Fatalf("OpenSession(): %v", err)
} }