package driver import ( "bufio" "context" "fmt" "io" "net" "net/url" "strconv" "strings" "git.unkin.net/unkin/waitfordb/internal/config" ) func init() { register(valkey{}) } type valkey struct{} func (valkey) Name() string { return "valkey" } func (valkey) Open(cfg config.Config) (Pinger, error) { p := &valkeyPinger{ addr: net.JoinHostPort(cfg.Host, cfg.Port), user: cfg.User, password: cfg.Password, } if cfg.DSN != "" { u, err := url.Parse(cfg.DSN) if err != nil { return nil, fmt.Errorf("invalid valkey DSN: %w", err) } if u.Scheme != "redis" && u.Scheme != "valkey" { return nil, fmt.Errorf("unsupported valkey DSN scheme %q (use redis:// or valkey://)", u.Scheme) } port := u.Port() if port == "" { port = "6379" } p.addr = net.JoinHostPort(u.Hostname(), port) if u.User != nil { p.user = u.User.Username() if pw, ok := u.User.Password(); ok { p.password = pw } } } return p, nil } // valkeyPinger speaks just enough RESP to AUTH and PING. Each Ping uses a // fresh connection so a half-open socket from a failed attempt cannot linger. type valkeyPinger struct { addr string user string password string } func (p *valkeyPinger) Ping(ctx context.Context) error { var d net.Dialer conn, err := d.DialContext(ctx, "tcp", p.addr) if err != nil { return err } defer conn.Close() if dl, ok := ctx.Deadline(); ok { _ = conn.SetDeadline(dl) } r := bufio.NewReader(conn) if p.password != "" { args := []string{"AUTH"} if p.user != "" && p.user != "default" { args = append(args, p.user) } args = append(args, p.password) if _, err := roundTrip(conn, r, args...); err != nil { return fmt.Errorf("AUTH %s: %w", p.user, err) } } reply, err := roundTrip(conn, r, "PING") if err != nil { return fmt.Errorf("PING: %w", err) } if reply != "PONG" { return fmt.Errorf("PING: unexpected reply %q", reply) } return nil } func (p *valkeyPinger) Close() error { return nil } func roundTrip(w io.Writer, r *bufio.Reader, args ...string) (string, error) { var b strings.Builder fmt.Fprintf(&b, "*%d\r\n", len(args)) for _, a := range args { fmt.Fprintf(&b, "$%d\r\n%s\r\n", len(a), a) } if _, err := io.WriteString(w, b.String()); err != nil { return "", err } return readReply(r) } func readReply(r *bufio.Reader) (string, error) { line, err := r.ReadString('\n') if err != nil { return "", err } line = strings.TrimRight(line, "\r\n") if line == "" { return "", fmt.Errorf("empty reply") } switch line[0] { case '+': return line[1:], nil case '-': return "", fmt.Errorf("server error: %s", line[1:]) case '$': n, err := strconv.Atoi(line[1:]) if err != nil || n < 0 { return "", fmt.Errorf("unexpected reply %q", line) } buf := make([]byte, n+2) // payload + trailing CRLF if _, err := io.ReadFull(r, buf); err != nil { return "", err } return string(buf[:n]), nil default: return "", fmt.Errorf("unexpected reply %q", line) } }