feat: persist worker control state in redis
This commit is contained in:
parent
a79d030c82
commit
a463a8cbd2
@ -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 {
|
||||||
|
|||||||
@ -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",
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
@ -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,
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
@ -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)
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@ -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 {
|
||||||
|
|||||||
@ -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
|
||||||
|
|
||||||
|
|||||||
@ -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
|
||||||
|
|||||||
@ -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) },
|
||||||
}
|
}
|
||||||
})
|
})
|
||||||
|
|||||||
@ -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)
|
||||||
}
|
}
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user