195 lines
5.7 KiB
Go
195 lines
5.7 KiB
Go
package store
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"net"
|
|
"sync/atomic"
|
|
"syscall"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/jackc/pgx/v5/pgconn"
|
|
"github.com/jackc/pgx/v5/pgproto3"
|
|
)
|
|
|
|
// servePG answers every connection on ln like a PostgreSQL server would: the i-th
|
|
// connection (from 1) gets a FATAL error with the SQLSTATE answer(i) returns, or,
|
|
// for "", a trust login and an empty reply to each simple query, which is all a
|
|
// ping sends. It counts the connections in accepts.
|
|
func servePG(ln net.Listener, answer func(i int32) string, accepts *atomic.Int32) {
|
|
for {
|
|
conn, err := ln.Accept()
|
|
if err != nil {
|
|
return
|
|
}
|
|
code := answer(accepts.Add(1))
|
|
go func() {
|
|
defer conn.Close()
|
|
be := pgproto3.NewBackend(conn, conn)
|
|
if _, err := be.ReceiveStartupMessage(); err != nil {
|
|
return
|
|
}
|
|
if code != "" {
|
|
be.Send(&pgproto3.ErrorResponse{Severity: "FATAL", Code: code, Message: "answered " + code})
|
|
_ = be.Flush()
|
|
return
|
|
}
|
|
be.Send(&pgproto3.AuthenticationOk{})
|
|
be.Send(&pgproto3.BackendKeyData{ProcessID: 1, SecretKey: []byte{0, 0, 0, 1}})
|
|
be.Send(&pgproto3.ReadyForQuery{TxStatus: 'I'})
|
|
if be.Flush() != nil {
|
|
return
|
|
}
|
|
for {
|
|
msg, err := be.Receive()
|
|
if err != nil {
|
|
return
|
|
}
|
|
switch msg.(type) {
|
|
case *pgproto3.Query:
|
|
be.Send(&pgproto3.EmptyQueryResponse{})
|
|
be.Send(&pgproto3.ReadyForQuery{TxStatus: 'I'})
|
|
if be.Flush() != nil {
|
|
return
|
|
}
|
|
case *pgproto3.Terminate:
|
|
return
|
|
}
|
|
}
|
|
}()
|
|
}
|
|
}
|
|
|
|
// freeAddr returns a loopback address nothing listens on.
|
|
func freeAddr(t *testing.T) string {
|
|
t.Helper()
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
addr := ln.Addr().String()
|
|
ln.Close()
|
|
return addr
|
|
}
|
|
|
|
func testDSN(addr string) string {
|
|
return fmt.Sprintf("postgres://felis@%s/felis?sslmode=disable", addr)
|
|
}
|
|
|
|
// A pod's first dials are refused until the network policy admits it; the open
|
|
// waits that out and connects once the server is reachable.
|
|
func TestOpenRetryingWaitsOutARefusedDial(t *testing.T) {
|
|
addr := freeAddr(t)
|
|
var accepts atomic.Int32
|
|
go func() {
|
|
time.Sleep(150 * time.Millisecond)
|
|
ln, err := net.Listen("tcp", addr)
|
|
if err != nil {
|
|
return
|
|
}
|
|
t.Cleanup(func() { ln.Close() })
|
|
servePG(ln, func(int32) string { return "" }, &accepts)
|
|
}()
|
|
var retries []error
|
|
drv, err := OpenRetrying(context.Background(), testDSN(addr), 5*time.Second, 20*time.Millisecond, func(err error) { retries = append(retries, err) })
|
|
if err != nil {
|
|
t.Fatalf("open: %v (after %d retries)", err, len(retries))
|
|
}
|
|
drv.Close()
|
|
if len(retries) == 0 {
|
|
t.Fatal("connected without a retry: the listener was up too early for this test to mean anything")
|
|
}
|
|
for _, r := range retries {
|
|
if !errors.Is(r, syscall.ECONNREFUSED) {
|
|
t.Fatalf("retried on %v, want only refused dials", r)
|
|
}
|
|
}
|
|
if n := accepts.Load(); n != 1 {
|
|
t.Fatalf("server saw %d connections, want 1", n)
|
|
}
|
|
}
|
|
|
|
// The window bounds the wait, and Open itself tries once.
|
|
func TestOpenRetryingGivesUpAtTheWindow(t *testing.T) {
|
|
addr := freeAddr(t)
|
|
retries := 0
|
|
start := time.Now()
|
|
_, err := OpenRetrying(context.Background(), testDSN(addr), 200*time.Millisecond, 20*time.Millisecond, func(error) { retries++ })
|
|
if !errors.Is(err, syscall.ECONNREFUSED) {
|
|
t.Fatalf("err = %v, want a refused dial", err)
|
|
}
|
|
if elapsed := time.Since(start); elapsed < 200*time.Millisecond || elapsed > 3*time.Second {
|
|
t.Fatalf("gave up after %s, want the 200ms window", elapsed)
|
|
}
|
|
if retries < 2 || retries > 11 {
|
|
t.Fatalf("%d retries in a 200ms window at 20ms, want several and at most one per interval", retries)
|
|
}
|
|
|
|
// Open has no window: one refused dial is its answer (a retry would call the nil
|
|
// onRetry and panic).
|
|
if _, err := Open(context.Background(), testDSN(addr)); !errors.Is(err, syscall.ECONNREFUSED) {
|
|
t.Fatalf("Open err = %v, want a refused dial", err)
|
|
}
|
|
}
|
|
|
|
// What the server says is final, except that it is still starting.
|
|
func TestOpenRetryingTakesTheServersAnswer(t *testing.T) {
|
|
for _, tc := range []struct {
|
|
name string
|
|
answer func(i int32) string
|
|
wantErr string // SQLSTATE of the final error; "" for a connection
|
|
accepts int32
|
|
}{
|
|
{"bad password", func(int32) string { return "28P01" }, "28P01", 1},
|
|
{"no such database", func(int32) string { return "3D000" }, "3D000", 1},
|
|
{"starting up", func(i int32) string {
|
|
if i <= 2 {
|
|
return "57P03"
|
|
}
|
|
return ""
|
|
}, "", 3},
|
|
} {
|
|
t.Run(tc.name, func(t *testing.T) {
|
|
ln, err := net.Listen("tcp", "127.0.0.1:0")
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
defer ln.Close()
|
|
var accepts atomic.Int32
|
|
go servePG(ln, tc.answer, &accepts)
|
|
retries := 0
|
|
drv, err := OpenRetrying(context.Background(), testDSN(ln.Addr().String()), 2*time.Second, 10*time.Millisecond, func(error) { retries++ })
|
|
if tc.wantErr == "" {
|
|
if err != nil {
|
|
t.Fatalf("open: %v", err)
|
|
}
|
|
drv.Close()
|
|
} else {
|
|
var pgErr *pgconn.PgError
|
|
if !errors.As(err, &pgErr) || pgErr.Code != tc.wantErr {
|
|
t.Fatalf("err = %v, want SQLSTATE %s", err, tc.wantErr)
|
|
}
|
|
}
|
|
if n := accepts.Load(); n != tc.accepts || retries != int(tc.accepts)-1 {
|
|
t.Fatalf("server saw %d connections after %d retries, want %d", n, retries, tc.accepts)
|
|
}
|
|
})
|
|
}
|
|
}
|
|
|
|
// A cancelled context ends the wait without another retry.
|
|
func TestOpenRetryingStopsWhenCancelled(t *testing.T) {
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
cancel()
|
|
retries := 0
|
|
start := time.Now()
|
|
if _, err := OpenRetrying(ctx, testDSN(freeAddr(t)), 5*time.Second, 20*time.Millisecond, func(error) { retries++ }); err == nil {
|
|
t.Fatal("opened on a cancelled context")
|
|
}
|
|
if retries != 0 || time.Since(start) > time.Second {
|
|
t.Fatalf("%d retries over %s on a cancelled context, want none", retries, time.Since(start))
|
|
}
|
|
}
|