249 lines
8.0 KiB
Go
249 lines
8.0 KiB
Go
package redisactivity
|
|
|
|
import (
|
|
"context"
|
|
"crypto/sha256"
|
|
"encoding/hex"
|
|
"encoding/json"
|
|
"sort"
|
|
"strconv"
|
|
"strings"
|
|
"time"
|
|
|
|
"proxy-pool/internal/domain/workerruntime"
|
|
)
|
|
|
|
const runtimeWireVersion = 1
|
|
|
|
const (
|
|
runtimeReplaceSession = "replace_session"
|
|
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"`
|
|
}
|
|
|
|
type runtimeCounterWire struct {
|
|
ProxyID string `json:"proxyId"`
|
|
Active int64 `json:"active"`
|
|
Reserved int64 `json:"reserved"`
|
|
Draining bool `json:"draining"`
|
|
}
|
|
|
|
type runtimeReportWire struct {
|
|
Version int `json:"version"`
|
|
WorkerID string `json:"workerId"`
|
|
SessionID string `json:"sessionId"`
|
|
Sequence string `json:"sequence"`
|
|
SnapshotVersion string `json:"snapshotVersion"`
|
|
OwnershipEpoch string `json:"ownershipEpoch"`
|
|
ObservedAtMS int64 `json:"observedAtMs"`
|
|
Counters []runtimeCounterWire `json:"counters"`
|
|
}
|
|
|
|
type runtimeOwnedProxyWire struct {
|
|
ProxyID string `json:"proxyId"`
|
|
WorkerID string `json:"workerId"`
|
|
OwnershipEpoch string `json:"ownershipEpoch"`
|
|
}
|
|
|
|
type runtimeSnapshotWire struct {
|
|
ProxyID string `json:"proxyId"`
|
|
Active int64 `json:"active"`
|
|
Reserved int64 `json:"reserved"`
|
|
Draining bool `json:"draining"`
|
|
Fresh bool `json:"fresh"`
|
|
}
|
|
|
|
var (
|
|
_ workerruntime.SessionWriter = (*Adapter)(nil)
|
|
_ workerruntime.ReportWriter = (*Adapter)(nil)
|
|
_ workerruntime.RuntimeReader = (*Adapter)(nil)
|
|
)
|
|
|
|
func (a *Adapter) ReplaceSession(ctx context.Context, session workerruntime.Session, ttl time.Duration) error {
|
|
if err := validateRuntimeCall(ctx, a); err != nil {
|
|
return err
|
|
}
|
|
if !runtimeClean(session.WorkerID) || !runtimeClean(session.InstanceID) || !runtimeClean(session.SessionID) ||
|
|
session.AckedSnapshotVersion == 0 || session.AckedOwnershipEpoch == 0 || ttl <= 0 {
|
|
return workerruntime.ErrInvalidSession
|
|
}
|
|
payload, err := json.Marshal(runtimeSessionWire{
|
|
Version: runtimeWireVersion, WorkerID: session.WorkerID,
|
|
InstanceID: session.InstanceID, SessionID: session.SessionID,
|
|
AckedSnapshotVersion: strconv.FormatUint(session.AckedSnapshotVersion, 10),
|
|
AckedOwnershipEpoch: strconv.FormatUint(session.AckedOwnershipEpoch, 10),
|
|
})
|
|
if err != nil {
|
|
return workerruntime.ErrInvalidSession
|
|
}
|
|
reply, err := a.runRuntime(ctx, runtimeReplaceSession, durationMillis(ttl), payload, "")
|
|
if err != nil {
|
|
return err
|
|
}
|
|
switch reply.Status {
|
|
case scriptOK:
|
|
return nil
|
|
case scriptInvalid:
|
|
return workerruntime.ErrInvalidSession
|
|
case scriptStale:
|
|
return workerruntime.ErrStaleSession
|
|
default:
|
|
return invalidScriptReply("unexpected worker session reply")
|
|
}
|
|
}
|
|
|
|
func (a *Adapter) ReplaceRuntime(ctx context.Context, report workerruntime.Report, ttl time.Duration) error {
|
|
if err := validateRuntimeCall(ctx, a); err != nil {
|
|
return err
|
|
}
|
|
payload, digest, err := a.encodeRuntimeReport(report, ttl)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
reply, err := a.runRuntime(ctx, runtimeReplaceReport, durationMillis(ttl), payload, digest)
|
|
if err != nil {
|
|
return err
|
|
}
|
|
switch reply.Status {
|
|
case scriptOK:
|
|
return nil
|
|
case scriptInvalid:
|
|
return workerruntime.ErrInvalidReport
|
|
case scriptStale:
|
|
return workerruntime.ErrStaleReport
|
|
case scriptConflict:
|
|
return workerruntime.ErrConflictingReport
|
|
case scriptUnavailable:
|
|
return workerruntime.ErrStaleSession
|
|
default:
|
|
return invalidScriptReply("unexpected worker runtime reply")
|
|
}
|
|
}
|
|
|
|
func (a *Adapter) ReadRuntime(ctx context.Context, proxies []workerruntime.OwnedProxy) ([]workerruntime.Snapshot, error) {
|
|
if err := validateRuntimeCall(ctx, a); err != nil {
|
|
return nil, err
|
|
}
|
|
if len(proxies) > a.options.MaxRuntimeCounters {
|
|
return nil, workerruntime.ErrInvalidQuery
|
|
}
|
|
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 {
|
|
return nil, workerruntime.ErrInvalidQuery
|
|
}
|
|
key := proxy.WorkerID + "\x00" + proxy.ProxyID
|
|
if _, exists := seen[key]; exists {
|
|
return nil, workerruntime.ErrInvalidQuery
|
|
}
|
|
seen[key] = struct{}{}
|
|
wires[index] = runtimeOwnedProxyWire{
|
|
ProxyID: proxy.ProxyID, WorkerID: proxy.WorkerID,
|
|
OwnershipEpoch: strconv.FormatUint(proxy.OwnershipEpoch, 10),
|
|
}
|
|
}
|
|
payload, err := json.Marshal(wires)
|
|
if err != nil {
|
|
return nil, workerruntime.ErrInvalidQuery
|
|
}
|
|
reply, err := a.runRuntime(ctx, runtimeRead, 0, payload, "")
|
|
if err != nil {
|
|
return nil, err
|
|
}
|
|
if reply.Status == scriptInvalid {
|
|
return nil, workerruntime.ErrInvalidQuery
|
|
}
|
|
if reply.Status != scriptOK || len(reply.Snapshots) != len(proxies) {
|
|
return nil, invalidScriptReply("unexpected worker runtime read reply")
|
|
}
|
|
result := make([]workerruntime.Snapshot, len(reply.Snapshots))
|
|
for index, snapshot := range reply.Snapshots {
|
|
if snapshot.ProxyID != proxies[index].ProxyID || snapshot.Active < 0 || snapshot.Reserved < 0 {
|
|
return nil, invalidScriptReply("invalid worker runtime snapshot")
|
|
}
|
|
result[index] = workerruntime.Snapshot{
|
|
ProxyID: snapshot.ProxyID, Active: snapshot.Active, Reserved: snapshot.Reserved,
|
|
Draining: snapshot.Draining, Fresh: snapshot.Fresh,
|
|
}
|
|
}
|
|
return result, nil
|
|
}
|
|
|
|
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 {
|
|
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[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,
|
|
})
|
|
if err != nil {
|
|
return nil, "", workerruntime.ErrInvalidReport
|
|
}
|
|
digest := sha256.Sum256(payload)
|
|
return payload, hex.EncodeToString(digest[:]), nil
|
|
}
|
|
|
|
func (a *Adapter) runRuntime(
|
|
ctx context.Context,
|
|
operation string,
|
|
ttlMS int64,
|
|
payload []byte,
|
|
digest string,
|
|
) (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,
|
|
}, operation, ttlMS, a.options.CleanupLimit, string(payload), digest)
|
|
if err != nil {
|
|
return runtimeScriptReply{}, err
|
|
}
|
|
var reply runtimeScriptReply
|
|
if err := decodeScriptResult(result, &reply); err != nil {
|
|
return runtimeScriptReply{}, err
|
|
}
|
|
return reply, nil
|
|
}
|
|
|
|
func validateRuntimeCall(ctx context.Context, adapter *Adapter) error {
|
|
if ctx == nil || adapter == nil {
|
|
return workerruntime.ErrInvalidStore
|
|
}
|
|
if err := ctx.Err(); err != nil {
|
|
return err
|
|
}
|
|
return nil
|
|
}
|
|
|
|
func runtimeClean(value string) bool {
|
|
return value != "" && strings.TrimSpace(value) == value
|
|
}
|