Files
zcli-playground/valkey/main.go
T
2026-06-09 14:52:14 +02:00

720 lines
20 KiB
Go
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
package main
import (
"context"
"crypto/tls"
"crypto/x509"
"flag"
"fmt"
"log"
"os"
"os/signal"
"strconv"
"syscall"
"time"
"github.com/redis/go-redis/v9"
"gopkg.in/yaml.v3"
)
// tlsRootCAs, when non-nil, is used as the RootCAs pool for every TLS
// connection. Set by --tls-ca for local testing against a Valkey whose
// cert is signed by a non-system CA.
var tlsRootCAs *x509.CertPool
type WorkerConfig struct {
Name string `yaml:"name"`
Role string `yaml:"role"`
Host string `yaml:"host"`
Port int `yaml:"port"`
Password string `yaml:"password"`
DB int `yaml:"db"`
TLS bool `yaml:"tls"`
Interval time.Duration `yaml:"interval"`
Channel string `yaml:"channel"`
ValueBytes int `yaml:"valueBytes"` // bloater role only
}
type Config struct {
Host string `yaml:"host,omitempty"`
Port int `yaml:"port,omitempty"`
Password string `yaml:"password,omitempty"`
DB int `yaml:"db,omitempty"`
TLS bool `yaml:"tls,omitempty"`
Workers []WorkerConfig `yaml:"workers,omitempty"`
}
func loadConfig(path string) (*Config, error) {
data, err := os.ReadFile(path)
if err != nil {
return nil, fmt.Errorf("read config: %w", err)
}
data = []byte(os.ExpandEnv(string(data)))
var cfg Config
if err := yaml.Unmarshal(data, &cfg); err != nil {
return nil, fmt.Errorf("parse config: %w", err)
}
applyZeropsOverrides(&cfg)
return &cfg, nil
}
// applyZeropsOverrides applies env-var overrides on top of the parsed
// config so the same config.yaml works locally and in Zerops:
//
// $VALKEY_PASSWORD → overrides every worker's password (so the secret
// lives in the Zerops env, not in committed YAML).
//
// $ZEROPS_Number → ordinal injected per Zerops replica. Distributes
// 10 containers across distinct DBs and pub/sub
// channels so they don't collide on shared cluster:
// cfg.Workers[*].DB = ordinal mod 16
// cfg.Workers[*].Channel = "<channel>:<ord mod 16>"
// The mod-16 wrap is needed because Valkey defaults
// to DB 0–15 and Zerops's ordinal grows across
// redeploys (a fresh deploy can start well above 16).
//
// Each override is independent: an unset env var is a no-op, leaving the
// YAML value in place.
func applyZeropsOverrides(cfg *Config) {
if pw := os.Getenv("VALKEY_PASSWORD"); pw != "" {
for i := range cfg.Workers {
cfg.Workers[i].Password = pw
}
}
if n := os.Getenv("ZEROPS_Number"); n != "" {
if raw, err := strconv.Atoi(n); err == nil {
dbN := raw % 16
suffix := strconv.Itoa(dbN)
for i := range cfg.Workers {
cfg.Workers[i].DB = dbN
if cfg.Workers[i].Channel != "" {
cfg.Workers[i].Channel = cfg.Workers[i].Channel + ":" + suffix
}
}
}
}
}
// applyFlagOverrides applies --host / --password / --tls-ca on top of the
// loaded config so the local binary can target a remote Valkey without
// editing config.yaml. --tls-ca loads the PEM into the package-level
// tlsRootCAs pool and forces TLS on every worker.
func applyFlagOverrides(cfg *Config, host string, port int, password string, tlsEnable bool, tlsCAPath string) error {
if host != "" {
cfg.Host = host
for i := range cfg.Workers {
cfg.Workers[i].Host = host
}
}
if port != 0 {
cfg.Port = port
for i := range cfg.Workers {
cfg.Workers[i].Port = port
}
}
if password != "" {
cfg.Password = password
for i := range cfg.Workers {
cfg.Workers[i].Password = password
}
}
if tlsCAPath != "" {
pem, err := os.ReadFile(tlsCAPath)
if err != nil {
return fmt.Errorf("read tls-ca: %w", err)
}
pool := x509.NewCertPool()
if !pool.AppendCertsFromPEM(pem) {
return fmt.Errorf("tls-ca: no certificates parsed from %s", tlsCAPath)
}
tlsRootCAs = pool
}
if tlsEnable || tlsCAPath != "" {
cfg.TLS = true
for i := range cfg.Workers {
cfg.Workers[i].TLS = true
}
}
return nil
}
func (c *Config) primary() (*redis.Options, error) {
if c.Host != "" && c.Port != 0 {
return optionsFromConn(c.Host, c.Port, c.Password, c.DB, c.TLS), nil
}
if len(c.Workers) > 0 {
w := c.Workers[0]
return optionsFromConn(w.Host, w.Port, w.Password, w.DB, w.TLS), nil
}
return nil, fmt.Errorf("config has neither top-level host nor workers")
}
func optionsFromConn(host string, port int, password string, db int, useTLS bool) *redis.Options {
opts := &redis.Options{
Addr: fmt.Sprintf("%s:%d", host, port),
Password: password,
DB: db,
}
if useTLS {
opts.TLSConfig = &tls.Config{MinVersion: tls.VersionTLS12, RootCAs: tlsRootCAs}
}
return opts
}
// buildConnOptions resolves the Valkey connection for the --seed/--verify/
// --flush/--inspect modes from either a --conn string or the override flags.
// A conn string carries host/port/password/db/TLS, so it stands in for the
// override flags; redis.ParseURL maps rediss:// to a TLS config. --tls-ca is
// honoured in both paths so a local run can trust a non-system CA.
func buildConnOptions(conn, host string, port int, password string, tlsEnable bool, tlsCAPath string) (*redis.Options, error) {
if conn != "" {
parsed, err := redis.ParseURL(conn)
if err != nil {
return nil, fmt.Errorf("--conn parse error: %w", err)
}
if tlsCAPath != "" {
pem, err := os.ReadFile(tlsCAPath)
if err != nil {
return nil, fmt.Errorf("read tls-ca: %w", err)
}
pool := x509.NewCertPool()
if !pool.AppendCertsFromPEM(pem) {
return nil, fmt.Errorf("tls-ca: no certificates parsed from %s", tlsCAPath)
}
if parsed.TLSConfig == nil {
parsed.TLSConfig = &tls.Config{MinVersion: tls.VersionTLS12}
}
parsed.TLSConfig.RootCAs = pool
}
return parsed, nil
}
cfg := &Config{}
if err := applyFlagOverrides(cfg, host, port, password, tlsEnable, tlsCAPath); err != nil {
return nil, fmt.Errorf("flag override error: %w", err)
}
if cfg.Host == "" {
return nil, fmt.Errorf("requires --host (e.g. --host valkey.zerops over the zcli VPN) or --conn")
}
if cfg.Port == 0 {
cfg.Port = 6379
}
return optionsFromConn(cfg.Host, cfg.Port, cfg.Password, cfg.DB, cfg.TLS), nil
}
func redactPassword(pw string) string {
if pw == "" {
return "(none)"
}
return "(set)"
}
func envInt(name string, def int) int {
if v := os.Getenv(name); v != "" {
if n, err := strconv.Atoi(v); err == nil {
return n
}
}
return def
}
func main() {
configPath := flag.String("config", "config.yaml", "path to YAML config file")
continuous := flag.Bool("continuous", false, "single-worker rolling smoke loop")
chaosLocal := flag.Bool("chaos", false, "multi-worker chaos TUI (local, in-process)")
workerMode := flag.Bool("worker", false, "distributed worker — runs the chaos workload and ships events to --collector-url")
collectorMode := flag.Bool("collector", false, "distributed collector — HTTP dashboard that ingests events from workers")
collectorURL := flag.String("collector-url", os.Getenv("COLLECTOR_URL"), "collector base URL (worker mode); env: COLLECTOR_URL")
port := flag.Int("collector-port", envInt("PORT", 8080), "HTTP port (collector mode); env: PORT")
interval := flag.Duration("interval", time.Second, "tick interval for --continuous")
hostOverride := flag.String("host", "", "override host for all workers (local testing)")
portOverride := flag.Int("port", 0, "override port for all workers (local testing)")
passwordOverride := flag.String("password", "", "override password for all workers (local testing)")
tlsEnable := flag.Bool("tls", false, "enable TLS for all workers using system root CAs (local testing)")
tlsCAPath := flag.String("tls-ca", "", "path to TLS CA cert PEM; enables TLS for all workers (local testing)")
seedConn := flag.String("conn", "", "Valkey connection string for --seed/--verify/--flush/--inspect (e.g. redis://default:pw@valkey1:6379 or rediss://...:6380 for TLS); overrides --host/--port/--password/--tls")
seedMode := flag.Bool("seed", false, "bulk-load one Valkey with deterministic data (local; use --host valkey.zerops over zcli VPN)")
verifyMode := flag.Bool("verify", false, "re-read seed:* keys and check they're 1:1 with the seed (same --seed-* flags)")
flushMode := flag.Bool("flush", false, "FLUSHDB the selected --db (clears the whole database)")
inspectMode := flag.Bool("inspect", false, "print every database and its key count")
seedMB := flag.Int("seed-mb", 500, "total payload in --seed/--verify mode, in MiB")
seedValueBytes := flag.Int("seed-value-bytes", 10*1024, "size of each value in --seed/--verify mode, in bytes")
seedBatch := flag.Int("seed-batch", 1000, "keys per pipeline batch in --seed/--verify mode")
seedSeed := flag.Uint64("seed-seed", 1, "PRNG seed for value generation; must match between --seed and --verify")
seedThrottleMB := flag.Float64("throttle-mb", 0, "throttle --seed write rate in MiB/s (0 = unlimited)")
dbOverride := flag.Int("db", -1, "Valkey DB number for --seed/--verify/--flush/--inspect (-1 = use --conn/config default)")
seedTimeout := flag.Duration("timeout", 5*time.Second, "dial/read/write timeout in --seed/--verify/--flush/--inspect mode (raise for slow VPN links)")
flag.Parse()
switch {
case *collectorMode:
os.Exit(runCollector(*port))
case *seedMode, *verifyMode, *flushMode, *inspectMode:
opts, err := buildConnOptions(*seedConn, *hostOverride, *portOverride, *passwordOverride, *tlsEnable, *tlsCAPath)
if err != nil {
log.Fatalf("connection error: %v", err)
}
// --db overrides whatever the conn string / flags resolved to, so the
// same target can be pointed at any logical database without editing
// the conn URL. -1 means "leave as-is".
if *dbOverride >= 0 {
opts.DB = *dbOverride
}
// Generous timeouts: the zcli VPN adds latency and can stall a dial
// well past go-redis's 5s default mid-run. PoolTimeout > DialTimeout
// so a slow dial doesn't surface as pool exhaustion first.
opts.DialTimeout = *seedTimeout
opts.ReadTimeout = *seedTimeout
opts.WriteTimeout = *seedTimeout
opts.PoolTimeout = *seedTimeout + 5*time.Second
switch {
case *inspectMode:
os.Exit(runInspect(opts))
case *flushMode:
os.Exit(runFlush(opts))
case *verifyMode:
os.Exit(runVerify(opts, *seedSeed, *seedMB*1024*1024, *seedValueBytes, *seedBatch))
default:
os.Exit(runSeed(opts, *seedSeed, *seedMB*1024*1024, *seedValueBytes, *seedBatch, *seedThrottleMB*1024*1024))
}
}
cfg, err := loadConfig(*configPath)
if err != nil {
log.Fatalf("config error: %v", err)
}
if err := applyFlagOverrides(cfg, *hostOverride, *portOverride, *passwordOverride, *tlsEnable, *tlsCAPath); err != nil {
log.Fatalf("flag override error: %v", err)
}
switch {
case *chaosLocal:
if len(cfg.Workers) == 0 {
log.Fatalf("--chaos requires workers: list in %s", *configPath)
}
os.Exit(runChaos(cfg.Workers))
case *workerMode:
if len(cfg.Workers) == 0 {
log.Fatalf("--worker requires workers: list in %s", *configPath)
}
if *collectorURL == "" {
log.Fatalf("--worker requires --collector-url or $COLLECTOR_URL")
}
os.Exit(runWorker(cfg.Workers, *collectorURL))
}
primary, err := cfg.primary()
if err != nil {
log.Fatalf("config error: %v", err)
}
fmt.Printf("redis config: addr=%s db=%d tls=%t password=%s\n",
primary.Addr, primary.DB, primary.TLSConfig != nil, redactPassword(primary.Password))
rdb := redis.NewClient(primary)
defer rdb.Close()
if *continuous {
runContinuous(rdb, *interval)
return
}
ctx, cancel := context.WithTimeout(context.Background(), 30*time.Second)
defer cancel()
tests := []struct {
name string
fn func(context.Context, *redis.Client) error
}{
{"PING", testPing},
{"INFO version", testInfo},
{"STRING SET/GET", testString},
{"INCR/DECR", testCounter},
{"EXPIRE/TTL", testExpire},
{"LIST ops", testList},
{"HASH ops", testHash},
{"SET ops", testSet},
{"SORTED SET ops", testZSet},
{"TRANSACTION", testTxn},
{"PIPELINE", testPipeline},
{"PUB/SUB", testPubSub},
{"FLUSHDB cleanup", testCleanup},
}
failed := 0
for _, t := range tests {
fmt.Printf("[ .. ] %s\n", t.name)
if err := t.fn(ctx, rdb); err != nil {
fmt.Printf("[FAIL] %s: %v\n", t.name, err)
failed++
continue
}
fmt.Printf("[ OK ] %s\n", t.name)
}
fmt.Printf("\n%d/%d tests passed\n", len(tests)-failed, len(tests))
if failed > 0 {
os.Exit(1)
}
}
func testPing(ctx context.Context, r *redis.Client) error {
pong, err := r.Ping(ctx).Result()
if err != nil {
return err
}
if pong != "PONG" {
return fmt.Errorf("expected PONG, got %q", pong)
}
return nil
}
func testInfo(ctx context.Context, r *redis.Client) error {
info, err := r.Info(ctx, "server").Result()
if err != nil {
return err
}
if len(info) == 0 {
return fmt.Errorf("empty INFO response")
}
return nil
}
func testString(ctx context.Context, r *redis.Client) error {
if err := r.Set(ctx, "vk:string", "hello-valkey", 0).Err(); err != nil {
return err
}
got, err := r.Get(ctx, "vk:string").Result()
if err != nil {
return err
}
if got != "hello-valkey" {
return fmt.Errorf("expected hello-valkey, got %q", got)
}
return nil
}
func testCounter(ctx context.Context, r *redis.Client) error {
r.Del(ctx, "vk:counter")
if _, err := r.Incr(ctx, "vk:counter").Result(); err != nil {
return err
}
if _, err := r.IncrBy(ctx, "vk:counter", 9).Result(); err != nil {
return err
}
v, err := r.Decr(ctx, "vk:counter").Result()
if err != nil {
return err
}
if v != 9 {
return fmt.Errorf("expected 9, got %d", v)
}
return nil
}
func testExpire(ctx context.Context, r *redis.Client) error {
if err := r.Set(ctx, "vk:ttl", "x", 5*time.Second).Err(); err != nil {
return err
}
ttl, err := r.TTL(ctx, "vk:ttl").Result()
if err != nil {
return err
}
if ttl <= 0 || ttl > 5*time.Second {
return fmt.Errorf("unexpected TTL: %v", ttl)
}
return nil
}
func testList(ctx context.Context, r *redis.Client) error {
r.Del(ctx, "vk:list")
if _, err := r.RPush(ctx, "vk:list", "a", "b", "c").Result(); err != nil {
return err
}
vals, err := r.LRange(ctx, "vk:list", 0, -1).Result()
if err != nil {
return err
}
if len(vals) != 3 || vals[0] != "a" || vals[2] != "c" {
return fmt.Errorf("unexpected list: %v", vals)
}
return nil
}
func testHash(ctx context.Context, r *redis.Client) error {
r.Del(ctx, "vk:hash")
if err := r.HSet(ctx, "vk:hash", "name", "valkey", "version", "7.2").Err(); err != nil {
return err
}
got, err := r.HGetAll(ctx, "vk:hash").Result()
if err != nil {
return err
}
if got["name"] != "valkey" || got["version"] != "7.2" {
return fmt.Errorf("unexpected hash: %v", got)
}
return nil
}
func testSet(ctx context.Context, r *redis.Client) error {
r.Del(ctx, "vk:set")
if _, err := r.SAdd(ctx, "vk:set", "x", "y", "z", "x").Result(); err != nil {
return err
}
n, err := r.SCard(ctx, "vk:set").Result()
if err != nil {
return err
}
if n != 3 {
return fmt.Errorf("expected 3 members, got %d", n)
}
return nil
}
func testZSet(ctx context.Context, r *redis.Client) error {
r.Del(ctx, "vk:zset")
_, err := r.ZAdd(ctx, "vk:zset",
redis.Z{Score: 1, Member: "one"},
redis.Z{Score: 2, Member: "two"},
redis.Z{Score: 3, Member: "three"},
).Result()
if err != nil {
return err
}
vals, err := r.ZRangeByScore(ctx, "vk:zset", &redis.ZRangeBy{Min: "1", Max: "2"}).Result()
if err != nil {
return err
}
if len(vals) != 2 || vals[0] != "one" {
return fmt.Errorf("unexpected zrange: %v", vals)
}
return nil
}
func testTxn(ctx context.Context, r *redis.Client) error {
r.Del(ctx, "vk:txn")
pipe := r.TxPipeline()
pipe.Set(ctx, "vk:txn", "1", 0)
pipe.Incr(ctx, "vk:txn")
pipe.Incr(ctx, "vk:txn")
if _, err := pipe.Exec(ctx); err != nil {
return err
}
v, err := r.Get(ctx, "vk:txn").Result()
if err != nil {
return err
}
if v != "3" {
return fmt.Errorf("expected 3, got %s", v)
}
return nil
}
func testPipeline(ctx context.Context, r *redis.Client) error {
pipe := r.Pipeline()
for i := 0; i < 5; i++ {
pipe.Set(ctx, fmt.Sprintf("vk:pipe:%d", i), i, 0)
}
if _, err := pipe.Exec(ctx); err != nil {
return err
}
v, err := r.Get(ctx, "vk:pipe:3").Result()
if err != nil {
return err
}
if v != "3" {
return fmt.Errorf("expected 3, got %s", v)
}
return nil
}
func testPubSub(ctx context.Context, r *redis.Client) error {
sub := r.Subscribe(ctx, "vk:channel")
defer sub.Close()
if _, err := sub.Receive(ctx); err != nil {
return fmt.Errorf("subscribe: %w", err)
}
ch := sub.Channel()
if err := r.Publish(ctx, "vk:channel", "ping").Err(); err != nil {
return err
}
select {
case msg := <-ch:
if msg.Payload != "ping" {
return fmt.Errorf("expected ping, got %q", msg.Payload)
}
case <-time.After(3 * time.Second):
return fmt.Errorf("timeout waiting for message")
}
return nil
}
func testCleanup(ctx context.Context, r *redis.Client) error {
keys, err := r.Keys(ctx, "vk:*").Result()
if err != nil {
return err
}
if len(keys) == 0 {
return nil
}
return r.Del(ctx, keys...).Err()
}
func runContinuous(r *redis.Client, interval time.Duration) {
stopCtx, stop := signal.NotifyContext(context.Background(), os.Interrupt, syscall.SIGTERM)
defer stop()
subCtx, cancelSub := context.WithCancel(stopCtx)
defer cancelSub()
sub := r.Subscribe(subCtx, "vk:cont:channel")
defer sub.Close()
if _, err := sub.Receive(subCtx); err != nil {
log.Fatalf("subscribe: %v", err)
}
subCh := sub.Channel()
ops := []struct {
name string
fn func(context.Context, *redis.Client, int64) error
}{
{"PING", opPing},
{"SET/GET", opSetGet},
{"INCR", opIncr},
{"LIST", opList},
{"INFO", opInfo},
{"PUB/SUB", func(ctx context.Context, r *redis.Client, i int64) error {
return opPubSub(ctx, r, i, subCh)
}},
}
var cycle, okCount, failCount int64
ticker := time.NewTicker(interval)
defer ticker.Stop()
fmt.Printf("continuous mode: addr=%s interval=%s (Ctrl+C to stop)\n", r.Options().Addr, interval)
for {
select {
case <-stopCtx.Done():
fmt.Printf("\nstopped: cycles=%d ok=%d failed=%d\n", cycle, okCount, failCount)
return
case <-ticker.C:
}
cycle++
results := make([]string, len(ops))
var failures []string
cycleStart := time.Now()
for i, op := range ops {
ctx, cancel := context.WithTimeout(stopCtx, interval)
start := time.Now()
err := op.fn(ctx, r, cycle)
cancel()
dur := time.Since(start)
if err != nil {
results[i] = fmt.Sprintf("%s=FAIL", op.name)
failures = append(failures, fmt.Sprintf(" %s (%s): %v", op.name, dur, err))
} else {
results[i] = fmt.Sprintf("%s=ok(%s)", op.name, dur.Round(time.Microsecond))
}
}
if len(failures) == 0 {
okCount++
} else {
failCount++
}
fmt.Printf("[%s] #%d %s | total=%s\n",
time.Now().Format("15:04:05"), cycle, joinResults(results), time.Since(cycleStart).Round(time.Microsecond))
for _, f := range failures {
fmt.Println(f)
}
}
}
func joinResults(parts []string) string {
out := ""
for i, p := range parts {
if i > 0 {
out += " "
}
out += p
}
return out
}
func opPing(ctx context.Context, r *redis.Client, _ int64) error {
pong, err := r.Ping(ctx).Result()
if err != nil {
return err
}
if pong != "PONG" {
return fmt.Errorf("expected PONG, got %q", pong)
}
return nil
}
func opSetGet(ctx context.Context, r *redis.Client, i int64) error {
val := fmt.Sprintf("v-%d", i)
if err := r.Set(ctx, "vk:cont:str", val, 10*time.Second).Err(); err != nil {
return err
}
got, err := r.Get(ctx, "vk:cont:str").Result()
if err != nil {
return err
}
if got != val {
return fmt.Errorf("set %q, got %q", val, got)
}
return nil
}
func opIncr(ctx context.Context, r *redis.Client, _ int64) error {
_, err := r.Incr(ctx, "vk:cont:counter").Result()
return err
}
func opList(ctx context.Context, r *redis.Client, i int64) error {
if err := r.RPush(ctx, "vk:cont:list", fmt.Sprintf("item-%d", i)).Err(); err != nil {
return err
}
if err := r.LTrim(ctx, "vk:cont:list", -10, -1).Err(); err != nil {
return err
}
_, err := r.LRange(ctx, "vk:cont:list", 0, -1).Result()
return err
}
func opInfo(ctx context.Context, r *redis.Client, _ int64) error {
info, err := r.Info(ctx, "server").Result()
if err != nil {
return err
}
if len(info) == 0 {
return fmt.Errorf("empty INFO")
}
return nil
}
func opPubSub(ctx context.Context, r *redis.Client, i int64, ch <-chan *redis.Message) error {
payload := fmt.Sprintf("ping-%d", i)
if err := r.Publish(ctx, "vk:cont:channel", payload).Err(); err != nil {
return err
}
for {
select {
case msg := <-ch:
if msg == nil {
return fmt.Errorf("channel closed")
}
if msg.Payload == payload {
return nil
}
case <-ctx.Done():
return fmt.Errorf("timeout waiting for %q", payload)
}
}
}