From a463a8cbd27ed5770b5fec5ac19cd9ce7e5adf33 Mon Sep 17 00:00:00 2001 From: youfak Date: Fri, 31 Jul 2026 11:25:22 +0800 Subject: [PATCH] feat: persist worker control state in redis --- .../adapters/redisactivity/adapter_test.go | 3 +- internal/adapters/redisactivity/keys.go | 64 +++--- .../ownership_integration_test.go | 4 +- internal/adapters/redisactivity/runtime.go | 204 ++++++++++++++--- .../runtime_contract_integration_test.go | 27 +++ .../redisactivity/runtime_integration_test.go | 4 +- internal/adapters/redisactivity/scripts.go | 23 +- .../redisactivity/scripts/ownership.lua | 3 +- .../redisactivity/scripts/runtime.lua | 209 +++++++++++++++++- .../workerruntime/contract_external_test.go | 2 +- .../workerruntime/contracttest/contract.go | 35 +-- 11 files changed, 472 insertions(+), 106 deletions(-) create mode 100644 internal/adapters/redisactivity/runtime_contract_integration_test.go diff --git a/internal/adapters/redisactivity/adapter_test.go b/internal/adapters/redisactivity/adapter_test.go index 6964dfe..67e0088 100644 --- a/internal/adapters/redisactivity/adapter_test.go +++ b/internal/adapters/redisactivity/adapter_test.go @@ -96,7 +96,8 @@ func TestNewBuildsClusterSafeKeyspaceAndHashesDynamicTokens(t *testing.T) { adapter.keys.expiry, adapter.keys.available, adapter.keys.owners, adapter.keys.ownerExpiry, adapter.keys.epoch, adapter.keys.inventory, 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, } for _, key := range staticKeys { diff --git a/internal/adapters/redisactivity/keys.go b/internal/adapters/redisactivity/keys.go index 23e6a2b..c514bad 100644 --- a/internal/adapters/redisactivity/keys.go +++ b/internal/adapters/redisactivity/keys.go @@ -9,41 +9,45 @@ import ( const redisKeyPrefix = "pp:{activity}:" type keyspace struct { - prefix string - records string - unique string - idkeys string - expiry string - available string - owners string - ownerExpiry string - epoch string - inventory string - stateInventory string - workerSessions string - workerSessionExpiry string - workerRuntime string - workerRuntimeExpiry string + prefix string + records string + unique string + idkeys string + expiry string + available string + owners string + ownerExpiry string + epoch string + inventory string + stateInventory string + workerSessions string + workerSessionExpiry string + workerSnapshots string + workerSnapshotExpiry string + workerRuntime string + workerRuntimeExpiry string } func newKeyspace(namespace string) keyspace { prefix := redisKeyPrefix + namespace return keyspace{ - prefix: prefix, - records: prefix + ":records", - unique: prefix + ":unique", - idkeys: prefix + ":idkeys", - expiry: prefix + ":expiry", - available: prefix + ":available", - owners: prefix + ":owners", - ownerExpiry: prefix + ":owner-expiry", - epoch: prefix + ":epoch", - inventory: prefix + ":inventory", - stateInventory: prefix + ":state-inventory", - workerSessions: prefix + ":worker-sessions", - workerSessionExpiry: prefix + ":worker-session-expiry", - workerRuntime: prefix + ":worker-runtime", - workerRuntimeExpiry: prefix + ":worker-runtime-expiry", + prefix: prefix, + records: prefix + ":records", + unique: prefix + ":unique", + idkeys: prefix + ":idkeys", + expiry: prefix + ":expiry", + available: prefix + ":available", + owners: prefix + ":owners", + ownerExpiry: prefix + ":owner-expiry", + epoch: prefix + ":epoch", + inventory: prefix + ":inventory", + stateInventory: prefix + ":state-inventory", + workerSessions: prefix + ":worker-sessions", + workerSessionExpiry: prefix + ":worker-session-expiry", + workerSnapshots: prefix + ":worker-snapshots", + workerSnapshotExpiry: prefix + ":worker-snapshot-expiry", + workerRuntime: prefix + ":worker-runtime", + workerRuntimeExpiry: prefix + ":worker-runtime-expiry", } } diff --git a/internal/adapters/redisactivity/ownership_integration_test.go b/internal/adapters/redisactivity/ownership_integration_test.go index e2881c7..0b82e35 100644 --- a/internal/adapters/redisactivity/ownership_integration_test.go +++ b/internal/adapters/redisactivity/ownership_integration_test.go @@ -29,8 +29,10 @@ func TestRedisOwnershipLifecycle(t *testing.T) { assertRedisKeysHaveTTL(t, fixture, fixture.Adapter.keys.owners, 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 { t.Fatalf("Get(assigned) = %+v, %t, %v", current, ok, err) } diff --git a/internal/adapters/redisactivity/runtime.go b/internal/adapters/redisactivity/runtime.go index 26d1f28..375617d 100644 --- a/internal/adapters/redisactivity/runtime.go +++ b/internal/adapters/redisactivity/runtime.go @@ -2,12 +2,9 @@ package redisactivity import ( "context" - "crypto/sha256" "encoding/hex" "encoding/json" - "sort" "strconv" - "strings" "time" "proxy-pool/internal/domain/workerruntime" @@ -17,17 +14,43 @@ const runtimeWireVersion = 1 const ( runtimeReplaceSession = "replace_session" + runtimeCurrentEpoch = "current_epoch" + runtimeOpenSession = "open_session" + runtimeRecordSnapshot = "record_snapshot" + runtimeAcknowledge = "acknowledge_snapshot" runtimeReplaceReport = "replace_report" runtimeRead = "read" ) type runtimeSessionWire struct { - Version int `json:"version"` - WorkerID string `json:"workerId"` - InstanceID string `json:"instanceId"` - SessionID string `json:"sessionId"` - AckedSnapshotVersion string `json:"ackedSnapshotVersion"` - AckedOwnershipEpoch string `json:"ackedOwnershipEpoch"` + Version int `json:"version"` + WorkerID string `json:"workerId"` + InstanceID string `json:"instanceId"` + SessionID string `json:"sessionId"` + Zone string `json:"zone,omitempty"` + 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 { @@ -63,11 +86,127 @@ type runtimeSnapshotWire struct { } var ( + _ workerruntime.ControlStore = (*Adapter)(nil) _ workerruntime.SessionWriter = (*Adapter)(nil) _ workerruntime.ReportWriter = (*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 { if err := validateRuntimeCall(ctx, a); err != nil { return err @@ -81,6 +220,7 @@ func (a *Adapter) ReplaceSession(ctx context.Context, session workerruntime.Sess InstanceID: session.InstanceID, SessionID: session.SessionID, AckedSnapshotVersion: strconv.FormatUint(session.AckedSnapshotVersion, 10), AckedOwnershipEpoch: strconv.FormatUint(session.AckedOwnershipEpoch, 10), + RuntimeEnabled: true, }) if err != nil { return workerruntime.ErrInvalidSession @@ -124,6 +264,8 @@ func (a *Adapter) ReplaceRuntime(ctx context.Context, report workerruntime.Repor return workerruntime.ErrConflictingReport case scriptUnavailable: return workerruntime.ErrStaleSession + case scriptSnapshotMismatch: + return workerruntime.ErrSnapshotMismatch default: 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)) seen := make(map[string]struct{}, len(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 } 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) { - if ttl <= 0 || !runtimeClean(report.WorkerID) || !runtimeClean(report.SessionID) || - report.Sequence == 0 || report.SnapshotVersion == 0 || report.OwnershipEpoch == 0 || report.ObservedAt.IsZero() || - len(report.Counters) > a.options.MaxRuntimeCounters { + normalized, digest, err := workerruntime.NormalizeReport(report) + if err != nil || ttl <= 0 || len(normalized.Counters) > a.options.MaxRuntimeCounters { return nil, "", workerruntime.ErrInvalidReport } - counters := append([]workerruntime.Counter(nil), report.Counters...) - sort.Slice(counters, func(left, right int) bool { return counters[left].ProxyID < counters[right].ProxyID }) - 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 := make([]runtimeCounterWire, len(normalized.Counters)) + for index, counter := range normalized.Counters { wires[index] = runtimeCounterWire{ ProxyID: counter.ProxyID, Active: counter.Active, Reserved: counter.Reserved, Draining: counter.Draining, } } payload, err := json.Marshal(runtimeReportWire{ - Version: runtimeWireVersion, WorkerID: report.WorkerID, SessionID: report.SessionID, - Sequence: strconv.FormatUint(report.Sequence, 10), - SnapshotVersion: strconv.FormatUint(report.SnapshotVersion, 10), - OwnershipEpoch: strconv.FormatUint(report.OwnershipEpoch, 10), - ObservedAtMS: report.ObservedAt.UTC().UnixMilli(), Counters: wires, + Version: runtimeWireVersion, WorkerID: normalized.WorkerID, SessionID: normalized.SessionID, + Sequence: strconv.FormatUint(normalized.Sequence, 10), + SnapshotVersion: strconv.FormatUint(normalized.SnapshotVersion, 10), + OwnershipEpoch: strconv.FormatUint(normalized.OwnershipEpoch, 10), + ObservedAtMS: normalized.ObservedAt.UnixMilli(), Counters: wires, }) if err != nil { return nil, "", workerruntime.ErrInvalidReport } - digest := sha256.Sum256(payload) return payload, hex.EncodeToString(digest[:]), nil } @@ -221,7 +355,8 @@ func (a *Adapter) runRuntime( ) (runtimeScriptReply, error) { result, err := runScript(ctx, a.client, runtimeScript, []string{ 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) if err != nil { return runtimeScriptReply{}, err @@ -243,6 +378,15 @@ func validateRuntimeCall(ctx context.Context, adapter *Adapter) error { return nil } -func runtimeClean(value string) bool { - return value != "" && strings.TrimSpace(value) == value +func referenceWire(reference workerruntime.SnapshotReference) runtimeSnapshotReferenceWire { + 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) } diff --git a/internal/adapters/redisactivity/runtime_contract_integration_test.go b/internal/adapters/redisactivity/runtime_contract_integration_test.go new file mode 100644 index 0000000..81878ba --- /dev/null +++ b/internal/adapters/redisactivity/runtime_contract_integration_test.go @@ -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, + } + }) +} diff --git a/internal/adapters/redisactivity/runtime_integration_test.go b/internal/adapters/redisactivity/runtime_integration_test.go index 98032a9..178cd7e 100644 --- a/internal/adapters/redisactivity/runtime_integration_test.go +++ b/internal/adapters/redisactivity/runtime_integration_test.go @@ -131,8 +131,8 @@ func TestRedisWorkerRuntimeRejectsEmptyReportBeyondAcknowledgedSnapshot(t *testi }, } { t.Run(name, func(t *testing.T) { - if err := fixture.Adapter.ReplaceRuntime(context.Background(), report, time.Minute); !errors.Is(err, workerruntime.ErrStaleReport) { - t.Fatalf("ReplaceRuntime() error = %v, want ErrStaleReport", err) + if err := fixture.Adapter.ReplaceRuntime(context.Background(), report, time.Minute); !errors.Is(err, workerruntime.ErrSnapshotMismatch) { + t.Fatalf("ReplaceRuntime() error = %v, want ErrSnapshotMismatch", err) } }) } diff --git a/internal/adapters/redisactivity/scripts.go b/internal/adapters/redisactivity/scripts.go index e4007d3..7996c0e 100644 --- a/internal/adapters/redisactivity/scripts.go +++ b/internal/adapters/redisactivity/scripts.go @@ -17,16 +17,18 @@ import ( type scriptStatus string const ( - scriptOK scriptStatus = "ok" - scriptInvalid scriptStatus = "invalid" - scriptNotFound scriptStatus = "not_found" - scriptConflict scriptStatus = "conflict" - scriptStale scriptStatus = "stale" - scriptUnavailable scriptStatus = "unavailable" - scriptInsufficient scriptStatus = "insufficient" - scriptAlreadyOwned scriptStatus = "already_owned" - scriptNotDraining scriptStatus = "not_draining" - scriptDrainNotReady scriptStatus = "drain_not_ready" + scriptOK scriptStatus = "ok" + scriptInvalid scriptStatus = "invalid" + scriptNotFound scriptStatus = "not_found" + scriptConflict scriptStatus = "conflict" + scriptStale scriptStatus = "stale" + scriptUnavailable scriptStatus = "unavailable" + scriptInsufficient scriptStatus = "insufficient" + scriptAlreadyOwned scriptStatus = "already_owned" + scriptNotDraining scriptStatus = "not_draining" + scriptDrainNotReady scriptStatus = "drain_not_ready" + scriptSnapshotMismatch scriptStatus = "snapshot_mismatch" + scriptStaleAcknowledgement scriptStatus = "stale_acknowledgement" ) type upsertScriptReply struct { @@ -77,6 +79,7 @@ type statusScriptInventory struct { type runtimeScriptReply struct { Status scriptStatus `json:"status"` Snapshots []runtimeSnapshotWire `json:"snapshots"` + Record string `json:"record,omitempty"` } type capacityScriptReply struct { diff --git a/internal/adapters/redisactivity/scripts/ownership.lua b/internal/adapters/redisactivity/scripts/ownership.lua index 5106779..2bd33ab 100644 --- a/internal/adapters/redisactivity/scripts/ownership.lua +++ b/internal/adapters/redisactivity/scripts/ownership.lua @@ -214,6 +214,7 @@ if operation == 'assign' then expires_at_ms = tonumber(record.usableUntilMs) end local next_epoch = redis.call('INCR', epoch_key) + redis.call('PERSIST', epoch_key) local assignment = { version = 1, proxyId = proxy_id, @@ -228,7 +229,6 @@ if operation == 'assign' then redis.call('ZADD', owner_expiry_key, expires_at_ms, proxy_id) touch(owners_key, tonumber(record.expiresAtMs)) touch(owner_expiry_key, tonumber(record.expiresAtMs)) - touch(epoch_key, tonumber(record.expiresAtMs)) record.ownerWorkerId = worker_id redis.call('HSET', records_key, proxy_id, cjson.encode(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) touch(owners_key, tonumber(record.expiresAtMs)) touch(owner_expiry_key, tonumber(record.expiresAtMs)) - touch(epoch_key, tonumber(record.expiresAtMs)) return finish({status = 'ok', record = encoded}) end diff --git a/internal/adapters/redisactivity/scripts/runtime.lua b/internal/adapters/redisactivity/scripts/runtime.lua index e2f7c9e..cc7f3a8 100644 --- a/internal/adapters/redisactivity/scripts/runtime.lua +++ b/internal/adapters/redisactivity/scripts/runtime.lua @@ -1,8 +1,11 @@ local sessions_key = KEYS[1] local session_expiry_key = KEYS[2] -local runtime_key = KEYS[3] -local runtime_expiry_key = KEYS[4] -local owners_key = KEYS[5] +local snapshots_key = KEYS[3] +local snapshot_expiry_key = KEYS[4] +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 ttl_ms = tonumber(ARGV[2]) @@ -10,11 +13,15 @@ local cleanup_limit = tonumber(ARGV[3]) local payload = ARGV[4] local digest = ARGV[5] -local function reply(status, snapshots) - if snapshots then - return cjson.encode({status = status, snapshots = snapshots}) +local function reply(status, snapshots, record) + if not snapshots then + local suffix = '' + if record then + suffix = ',"record":' .. cjson.encode(record) + end + return '{"status":' .. cjson.encode(status) .. ',"snapshots":[]' .. suffix .. '}' end - return '{"status":' .. cjson.encode(status) .. ',"snapshots":[]}' + return cjson.encode({status = status, snapshots = snapshots, record = record}) end local function now_ms() @@ -61,13 +68,50 @@ local function cleanup(now) redis.call('HDEL', runtime_key, worker_id) redis.call('ZREM', runtime_expiry_key, worker_id) 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 local function valid_session(value) 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.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 local function valid_owner(value, worker_id, ownership_epoch, now) @@ -80,6 +124,139 @@ end local now = now_ms() 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 not ttl_ms or ttl_ms <= 0 then return reply('invalid') @@ -89,7 +266,7 @@ if operation == 'replace_session' then return reply('invalid') end 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.expiresAtMs > now then 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 return reply('unavailable') 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 - return reply('stale') + return reply('snapshot_mismatch') end local current = decode_table(redis.call('HGET', runtime_key, report.workerId)) if current and current.sessionId == report.sessionId and valid_uint(current.sequence) then @@ -139,6 +317,12 @@ if operation == 'replace_report' then end if ordering == 0 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') end return reply('conflict') @@ -192,7 +376,8 @@ if operation == 'read' then local session = decode_table(redis.call('HGET', sessions_key, query.workerId)) local report = decode_table(redis.call('HGET', runtime_key, query.workerId)) 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 type(session.expiresAtMs) == 'number' and session.expiresAtMs > now and type(report.expiresAtMs) == 'number' and report.expiresAtMs > now and diff --git a/internal/domain/workerruntime/contract_external_test.go b/internal/domain/workerruntime/contract_external_test.go index 2398e35..6be0f62 100644 --- a/internal/domain/workerruntime/contract_external_test.go +++ b/internal/domain/workerruntime/contract_external_test.go @@ -17,7 +17,7 @@ func TestMemoryStoreContract(t *testing.T) { t.Fatalf("NewMemoryStore(): %v", err) } return contracttest.Fixture{ - Store: store, Reader: store, + Store: store, Reader: store, TTL: time.Minute, Advance: func(duration time.Duration) { now = now.Add(duration) }, } }) diff --git a/internal/domain/workerruntime/contracttest/contract.go b/internal/domain/workerruntime/contracttest/contract.go index 584fb5d..a546fbb 100644 --- a/internal/domain/workerruntime/contracttest/contract.go +++ b/internal/domain/workerruntime/contracttest/contract.go @@ -13,6 +13,7 @@ import ( type Fixture struct { Store workerruntime.ControlStore Reader workerruntime.RuntimeReader + TTL time.Duration Advance func(time.Duration) } @@ -28,62 +29,62 @@ func Run(t *testing.T, factory Factory) { func runLifecycle(t *testing.T, fixture Fixture) { t.Helper() ctx := context.Background() - open(t, fixture.Store) + open(t, fixture.Store, fixture.TTL) epoch, err := fixture.Store.CurrentOwnershipEpoch(ctx) if err != nil { t.Fatalf("CurrentOwnershipEpoch(): %v", err) } reference := snapshot(7, epoch, "snapshot-7") 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) } - 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) } 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) } - 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) } 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) } - 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) } - fixture.Advance(2 * time.Minute) + fixture.Advance(2 * fixture.TTL) assertFresh(t, fixture.Reader, epoch, false) } func runNegativeAck(t *testing.T, fixture Fixture) { t.Helper() ctx := context.Background() - open(t, fixture.Store) + open(t, fixture.Store, fixture.TTL) epoch, err := fixture.Store.CurrentOwnershipEpoch(ctx) if err != nil { t.Fatalf("CurrentOwnershipEpoch(): %v", err) } 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) } - 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) } 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) } if err := fixture.Store.AcknowledgeSnapshot(ctx, workerruntime.SnapshotAcknowledgement{ 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) } - 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) } } @@ -91,17 +92,17 @@ func runNegativeAck(t *testing.T, fixture Fixture) { func newFixture(t *testing.T, factory Factory) Fixture { t.Helper() 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") } return fixture } -func open(t *testing.T, store workerruntime.ControlStore) { +func open(t *testing.T, store workerruntime.ControlStore, ttl time.Duration) { t.Helper() err := store.OpenSession(context.Background(), workerruntime.Session{ WorkerID: "worker-a", InstanceID: "instance-a", SessionID: "session-a", Zone: "zone-a", ProtocolVersion: 1, - }, time.Minute) + }, ttl) if err != nil { t.Fatalf("OpenSession(): %v", err) }