mirror of
https://github.com/tiennm99/keepalive.git
synced 2026-10-11 03:13:31 +00:00
- Generated service names come only from a URL host, a key=value DSN's host=, or a MySQL tcp() address, so a password in a key=value DSN can no longer end up in the name printed on every log line. - connect_timeout is added only to postgres:// and postgresql:// URLs (parsed, so it never lands after a fragment) or as a key=value token, never glued onto a key=value value containing "://". - Driver parse errors that quote password fragments (mongo escape errors, lib/pq's missing "=" error) are replaced with generic hints. - MongoDB keeps its client only after a successful connect, so a failed connect is not disconnected twice. - Couchbase gets ready_timeout plus 1 minute to connect, so raising ready_timeout takes effect. - Shutdown waits at most 7 seconds, inside Docker's 10-second grace period. - Each adapter declares its config keys; unknown keys anywhere in the file, including under config, log a warning without blocking start.
274 lines
7.5 KiB
Go
274 lines
7.5 KiB
Go
package main
|
|
|
|
import (
|
|
"context"
|
|
"errors"
|
|
"fmt"
|
|
"net/url"
|
|
"strings"
|
|
"sync"
|
|
"sync/atomic"
|
|
"testing"
|
|
"time"
|
|
|
|
"github.com/tiennm99/keepalive/adapter"
|
|
)
|
|
|
|
type retryConnectAdapter struct {
|
|
connects *atomic.Int32
|
|
connected chan<- struct{}
|
|
once *sync.Once
|
|
}
|
|
|
|
func (a *retryConnectAdapter) Connect(context.Context) error {
|
|
if a.connects.Add(1) == 1 {
|
|
return errors.New("connect failed")
|
|
}
|
|
a.once.Do(func() { close(a.connected) })
|
|
return nil
|
|
}
|
|
|
|
func (a *retryConnectAdapter) Increment(context.Context) (int64, error) {
|
|
return 0, nil
|
|
}
|
|
|
|
func (a *retryConnectAdapter) Close(context.Context) error {
|
|
return nil
|
|
}
|
|
|
|
func TestRunServiceRetriesConnectFailure(t *testing.T) {
|
|
oldRetryDelay := retryDelay
|
|
retryDelay = 10 * time.Millisecond
|
|
defer func() { retryDelay = oldRetryDelay }()
|
|
|
|
var connects atomic.Int32
|
|
connected := make(chan struct{})
|
|
var once sync.Once
|
|
adapter.Registry["retry-test"] = func(adapter.Config) (adapter.Adapter, error) {
|
|
return &retryConnectAdapter{connects: &connects, connected: connected, once: &once}, nil
|
|
}
|
|
defer delete(adapter.Registry, "retry-test")
|
|
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
var wg sync.WaitGroup
|
|
runService(ctx, &wg, serviceConfig{
|
|
Name: "retry-test",
|
|
AdapterType: "retry-test",
|
|
Interval: time.Hour,
|
|
Config: adapter.Config{},
|
|
})
|
|
|
|
select {
|
|
case <-connected:
|
|
case <-time.After(time.Second):
|
|
cancel()
|
|
t.Fatal("service did not retry and connect")
|
|
}
|
|
|
|
cancel()
|
|
done := make(chan struct{})
|
|
go func() {
|
|
wg.Wait()
|
|
close(done)
|
|
}()
|
|
|
|
select {
|
|
case <-done:
|
|
case <-time.After(time.Second):
|
|
t.Fatal("service did not stop after context cancellation")
|
|
}
|
|
|
|
if got := connects.Load(); got != 2 {
|
|
t.Fatalf("connect attempts = %d, want 2", got)
|
|
}
|
|
}
|
|
|
|
type failingIncrementAdapter struct {
|
|
connects *atomic.Int32
|
|
closes *atomic.Int32
|
|
}
|
|
|
|
func (a *failingIncrementAdapter) Connect(context.Context) error {
|
|
a.connects.Add(1)
|
|
return nil
|
|
}
|
|
|
|
func (a *failingIncrementAdapter) Increment(context.Context) (int64, error) {
|
|
return 0, errors.New("relation \"keepalive\" does not exist")
|
|
}
|
|
|
|
func (a *failingIncrementAdapter) Close(context.Context) error {
|
|
a.closes.Add(1)
|
|
return nil
|
|
}
|
|
|
|
func TestRunServiceReconnectsAfterIncrementFailure(t *testing.T) {
|
|
oldRetryDelay := retryDelay
|
|
retryDelay = 10 * time.Millisecond
|
|
defer func() { retryDelay = oldRetryDelay }()
|
|
|
|
var connects, closes atomic.Int32
|
|
adapter.Registry["fail-tick-test"] = func(adapter.Config) (adapter.Adapter, error) {
|
|
return &failingIncrementAdapter{connects: &connects, closes: &closes}, nil
|
|
}
|
|
defer delete(adapter.Registry, "fail-tick-test")
|
|
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
var wg sync.WaitGroup
|
|
runService(ctx, &wg, serviceConfig{
|
|
Name: "fail-tick-test",
|
|
AdapterType: "fail-tick-test",
|
|
Interval: 5 * time.Millisecond,
|
|
Config: adapter.Config{},
|
|
})
|
|
|
|
deadline := time.After(time.Second)
|
|
for connects.Load() < 3 {
|
|
select {
|
|
case <-deadline:
|
|
cancel()
|
|
t.Fatalf("connect attempts = %d, want at least 3", connects.Load())
|
|
case <-time.After(5 * time.Millisecond):
|
|
}
|
|
}
|
|
cancel()
|
|
wg.Wait()
|
|
|
|
if closes.Load() < connects.Load()-1 {
|
|
t.Fatalf("closes = %d for %d connects; each failed session must be closed", closes.Load(), connects.Load())
|
|
}
|
|
}
|
|
|
|
func TestRedactURLErrorHidesPassword(t *testing.T) {
|
|
for _, raw := range []string{
|
|
"postgres://user:S3CR%ETpw@db.example.com/keepalive",
|
|
"redis://user:S3CRETpw@cache.example.com:port/0",
|
|
} {
|
|
_, err := url.Parse(raw)
|
|
if err == nil {
|
|
t.Fatalf("%s: want parse error", raw)
|
|
}
|
|
got := redactURLError(fmt.Errorf("connect: %w", err)).Error()
|
|
if strings.Contains(got, "S3CR") {
|
|
t.Fatalf("redacted error leaks password: %s", got)
|
|
}
|
|
if !strings.Contains(got, "invalid connection URL") {
|
|
t.Fatalf("redacted error = %q, want a URL hint", got)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestRedactURLErrorKeepsOtherErrors(t *testing.T) {
|
|
err := errors.New("dial tcp: connection refused")
|
|
if got := redactURLError(err); got != err {
|
|
t.Fatalf("redactURLError changed a non-URL error: %v", got)
|
|
}
|
|
}
|
|
|
|
type countingAdapter struct{ increments *atomic.Int32 }
|
|
|
|
func (a *countingAdapter) Connect(context.Context) error { return nil }
|
|
func (a *countingAdapter) Increment(context.Context) (int64, error) {
|
|
return int64(a.increments.Add(1)), nil
|
|
}
|
|
func (a *countingAdapter) Close(context.Context) error { return nil }
|
|
|
|
func TestRunServiceIncrementsRightAfterConnect(t *testing.T) {
|
|
var increments atomic.Int32
|
|
adapter.Registry["first-tick-test"] = func(adapter.Config) (adapter.Adapter, error) {
|
|
return &countingAdapter{increments: &increments}, nil
|
|
}
|
|
defer delete(adapter.Registry, "first-tick-test")
|
|
|
|
ctx, cancel := context.WithCancel(context.Background())
|
|
var wg sync.WaitGroup
|
|
runService(ctx, &wg, serviceConfig{
|
|
Name: "first-tick-test",
|
|
AdapterType: "first-tick-test",
|
|
Interval: time.Hour,
|
|
Config: adapter.Config{},
|
|
})
|
|
|
|
deadline := time.After(time.Second)
|
|
for increments.Load() == 0 {
|
|
select {
|
|
case <-deadline:
|
|
cancel()
|
|
t.Fatal("no increment before the first interval")
|
|
case <-time.After(5 * time.Millisecond):
|
|
}
|
|
}
|
|
cancel()
|
|
wg.Wait()
|
|
}
|
|
|
|
func TestRedactURLErrorHidesDriverParseErrors(t *testing.T) {
|
|
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
|
|
defer cancel()
|
|
for _, tc := range []struct{ adapterType, key, value string }{
|
|
{"mongodb", "uri", "mongodb://user:S3CR%zzpw@db.example.com"},
|
|
{"postgresql", "url", "host=db.example.com user=u password=abc S3CRETdef dbname=k"},
|
|
} {
|
|
cfg := adapter.Config{tc.key: tc.value, "database": "k", "collection": "c"}
|
|
a, err := adapter.New(tc.adapterType, cfg)
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
err = a.Connect(ctx)
|
|
a.Close(ctx)
|
|
if err == nil {
|
|
t.Fatalf("%s: want connect error", tc.adapterType)
|
|
}
|
|
if got := redactURLError(err).Error(); strings.Contains(got, "S3CR") {
|
|
t.Fatalf("%s: redacted error leaks password: %s", tc.adapterType, got)
|
|
}
|
|
}
|
|
}
|
|
|
|
func TestWaitShutdownGivesUpAfterTimeout(t *testing.T) {
|
|
var wg sync.WaitGroup
|
|
wg.Add(1)
|
|
defer wg.Done()
|
|
if waitShutdown(&wg, 10*time.Millisecond) {
|
|
t.Fatal("waitShutdown reported a clean stop while a service was still running")
|
|
}
|
|
}
|
|
|
|
func TestWaitShutdownReturnsWhenServicesStop(t *testing.T) {
|
|
var wg sync.WaitGroup
|
|
if !waitShutdown(&wg, time.Second) {
|
|
t.Fatal("waitShutdown timed out with no running services")
|
|
}
|
|
}
|
|
|
|
type slowConnectAdapter struct{ timeout time.Duration }
|
|
|
|
func (a slowConnectAdapter) Connect(context.Context) error { return nil }
|
|
func (a slowConnectAdapter) Increment(context.Context) (int64, error) { return 0, nil }
|
|
func (a slowConnectAdapter) Close(context.Context) error { return nil }
|
|
func (a slowConnectAdapter) ConnectTimeout() time.Duration { return a.timeout }
|
|
|
|
func TestConnectTimeoutForUsesLongerAdapterTimeout(t *testing.T) {
|
|
if got := connectTimeoutFor(slowConnectAdapter{timeout: 3 * time.Minute}); got != 3*time.Minute {
|
|
t.Fatalf("connectTimeoutFor = %s, want 3m", got)
|
|
}
|
|
if got := connectTimeoutFor(slowConnectAdapter{timeout: time.Second}); got != connectTimeout {
|
|
t.Fatalf("connectTimeoutFor = %s, want default %s", got, connectTimeout)
|
|
}
|
|
}
|
|
|
|
func TestFailedMongoConnectLeavesNothingToClose(t *testing.T) {
|
|
a, err := adapter.New("mongodb", adapter.Config{"uri": "mongodb://127.0.0.1:1/?serverSelectionTimeoutMS=200", "database": "k", "collection": "c"})
|
|
if err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
|
|
defer cancel()
|
|
if err := a.Connect(ctx); err == nil {
|
|
t.Fatal("want connect error")
|
|
}
|
|
if err := a.Close(ctx); err != nil {
|
|
t.Fatalf("Close after failed Connect = %v, want nil", err)
|
|
}
|
|
}
|