Files
keepalive/runner_test.go
T
tiennm99 47afb85093 fix: close remaining secret leaks, timeout gaps and config blind spots
- 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.
2026-10-09 12:36:21 +07:00

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)
}
}