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.
This commit is contained in:
tiennm99 committed 2026-10-09 12:36:21 +07:00
1 parent e646af19b8
commit 47afb85093
14 files changed
+299 -42

No files matched your search

+11
View File
@@ -75,6 +75,17 @@ type Factory func(Config) (Adapter, error)
var Registry = map[string]Factory{}
// ConfigKeys lists the config keys each adapter reads, so the config loader
// can warn about misspelled or unsupported keys. counter_key is set by the
// loader and is not listed.
var ConfigKeys = map[string][]string{}
// ConnectTimeouter is implemented by adapters whose Connect may legitimately
// need longer than the runner's default connect timeout.
type ConnectTimeouter interface {
ConnectTimeout() time.Duration
}
func New(adapterType string, cfg Config) (Adapter, error) {
f, ok := Registry[adapterType]
if !ok {
+7
View File
@@ -8,6 +8,7 @@ import (
)
func init() {
ConfigKeys["couchbase"] = []string{"connection_string", "username", "password", "bucket_name", "scope_name", "collection_name", "ready_timeout", "bucket_ram_quota_mb"}
Registry["couchbase"] = func(cfg Config) (Adapter, error) {
conn, err := cfg.Required("connection_string")
if err != nil {
@@ -118,6 +119,12 @@ func (a *couchbaseAdapter) Increment(ctx context.Context) (int64, error) {
return int64(res.Content()), nil
}
// ConnectTimeout leaves room for ready_timeout, which bounds each of the
// bucket, scope, collection and document setup steps, plus a margin.
func (a *couchbaseAdapter) ConnectTimeout() time.Duration {
return a.readyTimeout + time.Minute
}
func (a *couchbaseAdapter) Close(_ context.Context) error {
if a.cluster == nil {
return nil
+4 -1
View File
@@ -9,6 +9,8 @@ import (
)
func init() {
ConfigKeys["mongodb"] = []string{"uri", "database", "collection"}
ConfigKeys["mongo"] = ConfigKeys["mongodb"]
Registry["mongodb"] = func(cfg Config) (Adapter, error) {
uri, err := cfg.Required("uri")
if err != nil {
@@ -46,7 +48,6 @@ func (a *mongoAdapter) Connect(ctx context.Context) error {
if err != nil {
return err
}
a.client = client
a.coll = client.Database(a.dbName).Collection(a.collName)
if err := client.Ping(ctx, nil); err != nil {
client.Disconnect(ctx)
@@ -56,6 +57,8 @@ func (a *mongoAdapter) Connect(ctx context.Context) error {
client.Disconnect(ctx)
return err
}
// Set only on success so Close does not disconnect a failed client again.
a.client = client
return nil
}
+1
View File
@@ -9,6 +9,7 @@ import (
)
func init() {
ConfigKeys["mysql"] = []string{"dsn"}
Registry["mysql"] = func(cfg Config) (Adapter, error) {
dsn, err := cfg.Required("dsn")
if err != nil {
+24 -8
View File
@@ -3,12 +3,17 @@ package adapter
import (
"context"
"database/sql"
"net/url"
"strings"
_ "github.com/lib/pq"
)
const defaultPostgresConnectTimeout = "30"
func init() {
ConfigKeys["postgresql"] = []string{"url"}
ConfigKeys["postgres"] = ConfigKeys["postgresql"]
Registry["postgresql"] = func(cfg Config) (Adapter, error) {
url, err := cfg.Required("url")
if err != nil {
@@ -86,16 +91,27 @@ func (a *postgresAdapter) Close(_ context.Context) error {
// withDefaultConnectTimeout adds connect_timeout when the DSN has none:
// lib/pq honours the context only while dialing, so a server that accepts the
// connection but never answers the startup handshake would hang forever.
// Like lib/pq, it treats only postgres:// and postgresql:// strings as URLs
// and everything else as a key=value DSN.
func withDefaultConnectTimeout(dsn string) string {
if strings.Contains(dsn, "connect_timeout") {
if strings.HasPrefix(dsn, "postgres://") || strings.HasPrefix(dsn, "postgresql://") {
u, err := url.Parse(dsn)
if err != nil {
// Connect reports the parse error.
return dsn
}
switch {
case !strings.Contains(dsn, "://"):
return dsn + " connect_timeout=30"
case strings.Contains(dsn, "?"):
return dsn + "&connect_timeout=30"
default:
return dsn + "?connect_timeout=30"
q := u.Query()
if q.Has("connect_timeout") {
return dsn
}
q.Set("connect_timeout", defaultPostgresConnectTimeout)
u.RawQuery = q.Encode()
return u.String()
}
for _, field := range strings.Fields(dsn) {
if strings.HasPrefix(field, "connect_timeout=") {
return dsn
}
}
return dsn + " connect_timeout=" + defaultPostgresConnectTimeout
}
+24 -2
View File
@@ -5,12 +5,34 @@ import "testing"
func TestWithDefaultConnectTimeout(t *testing.T) {
for _, tc := range []struct{ in, want string }{
{"postgres://u:p@db.example.com/k", "postgres://u:p@db.example.com/k?connect_timeout=30"},
{"postgres://u:p@db.example.com/k?sslmode=require", "postgres://u:p@db.example.com/k?sslmode=require&connect_timeout=30"},
{"postgresql://u:p@db.example.com/k?sslmode=require", "postgresql://u:p@db.example.com/k?connect_timeout=30&sslmode=require"},
{"postgres://db.example.com/k?connect_timeout=5", "postgres://db.example.com/k?connect_timeout=5"},
// The timeout belongs in the query, never after a fragment.
{"postgres://u:p@db.example.com/k#x", "postgres://u:p@db.example.com/k?connect_timeout=30#x"},
{"host=db.example.com dbname=k", "host=db.example.com dbname=k connect_timeout=30"},
{"host=db.example.com connect_timeout=5", "host=db.example.com connect_timeout=5"},
// "://" inside a key=value value does not make it a URL.
{"host=h user=u password=a://b dbname=k", "host=h user=u password=a://b dbname=k connect_timeout=30"},
// Unparseable URLs are left for Connect to report.
{"postgres://u:p%zz@db.example.com/k", "postgres://u:p%zz@db.example.com/k"},
} {
if got := withDefaultConnectTimeout(tc.in); got != tc.want {
t.Fatalf("withDefaultConnectTimeout(%q) = %q, want %q", tc.in, got, tc.want)
t.Errorf("withDefaultConnectTimeout(%q) = %q, want %q", tc.in, got, tc.want)
}
}
}
func TestConfigKeysCoverEveryAdapter(t *testing.T) {
for name := range Registry {
if _, ok := ConfigKeys[name]; !ok {
t.Errorf("adapter %q has no ConfigKeys entry", name)
}
}
}
func TestCouchbaseConnectTimeoutExceedsReadyTimeout(t *testing.T) {
a := &couchbaseAdapter{readyTimeout: 2 * 60 * 1e9}
if a.ConnectTimeout() <= a.readyTimeout {
t.Fatalf("ConnectTimeout %s does not exceed ready_timeout %s", a.ConnectTimeout(), a.readyTimeout)
}
}
+1
View File
@@ -16,6 +16,7 @@ func (discardRedisLogger) Printf(context.Context, string, ...interface{}) {}
func init() {
redis.SetLogger(discardRedisLogger{})
ConfigKeys["redis"] = []string{"url", "namespace"}
Registry["redis"] = func(cfg Config) (Adapter, error) {
url, err := cfg.Required("url")
if err != nil {
+37 -9
View File
@@ -6,7 +6,9 @@ import (
"fmt"
"io"
"log"
"maps"
"os"
"slices"
"strconv"
"strings"
"time"
@@ -58,22 +60,48 @@ func loadConfigFile(path string) ([]serviceConfig, error) {
if err := yaml.Unmarshal(data, &raw); err != nil {
return nil, fmt.Errorf("parse config: %w", err)
}
if err := checkUnknownFields(data); err != nil {
log.Printf("warning: %s: %v; unknown keys are ignored", path, err)
for _, warning := range schemaWarnings(data, raw) {
log.Printf("warning: %s: %s; ignoring it", path, warning)
}
return normalizeConfig(raw)
}
// checkUnknownFields reports keys that appConfig does not define, such as a
// misspelled interval. Keys under a service's config map are not checked.
func checkUnknownFields(data []byte) error {
// schemaWarnings lists every key the config schema does not define: keys
// appConfig and serviceFileConfig lack (such as a misspelled interval), and
// keys under a service's config map that its adapter does not read (such as
// a misspelled namespace). Unknown keys are ignored, not fatal.
func schemaWarnings(data []byte, raw appConfig) []string {
var warnings []string
decoder := yaml.NewDecoder(bytes.NewReader(data))
decoder.KnownFields(true)
var raw appConfig
if err := decoder.Decode(&raw); err != nil && !errors.Is(err, io.EOF) {
return err
var strict appConfig
if err := decoder.Decode(&strict); err != nil && !errors.Is(err, io.EOF) {
var typeErr *yaml.TypeError
if errors.As(err, &typeErr) {
warnings = append(warnings, typeErr.Errors...)
} else {
warnings = append(warnings, err.Error())
}
return nil
}
for i, service := range raw.Services {
adapterType := strings.TrimSpace(service.Adapter)
known, ok := adapter.ConfigKeys[adapterType]
if !ok {
// normalizeConfig rejects the unknown adapter itself.
continue
}
for _, key := range slices.Sorted(maps.Keys(service.Config)) {
if key == "counter_key" || slices.Contains(known, key) {
// normalizeConfig rejects config.counter_key with its own error.
continue
}
warnings = append(warnings, fmt.Sprintf("services[%d].config: unknown key %q for adapter %s (known: %s)",
i, key, adapterType, strings.Join(known, ", ")))
}
}
return warnings
}
func defaultConfigFile() (string, error) {
+29 -6
View File
@@ -6,6 +6,8 @@ import (
"strings"
"testing"
"time"
"gopkg.in/yaml.v3"
)
func TestNormalizeConfigGeneratesNamesFromAdapterAndHost(t *testing.T) {
@@ -265,13 +267,34 @@ func TestFirstExistingConfigFileRejectsDirectory(t *testing.T) {
}
}
func TestCheckUnknownFieldsReportsTypos(t *testing.T) {
err := checkUnknownFields([]byte("intervall: 1m\nservices:\n - adapter: redis\n couter_key: x\n config:\n url: redis://cache.example.com\n anything: goes\n"))
if err == nil || !strings.Contains(err.Error(), "intervall") || !strings.Contains(err.Error(), "couter_key") {
t.Fatalf("err = %v, want both unknown keys reported", err)
func TestSchemaWarningsReportsUnknownKeysEverywhere(t *testing.T) {
data := []byte("intervall: 1m\nservices:\n - adapter: redis\n couter_key: x\n config:\n url: redis://cache.example.com\n namspace: keepalive\n - adapter: postgres\n config:\n url: postgres://db.example.com/k\n")
var raw appConfig
if err := yaml.Unmarshal(data, &raw); err != nil {
t.Fatal(err)
}
if err := checkUnknownFields([]byte("interval: 1m\nservices:\n - adapter: redis\n config:\n url: redis://cache.example.com\n")); err != nil {
t.Fatalf("valid config reported: %v", err)
got := strings.Join(schemaWarnings(data, raw), "\n")
for _, want := range []string{"intervall", "couter_key", `services[0].config: unknown key "namspace" for adapter redis`} {
if !strings.Contains(got, want) {
t.Fatalf("warnings %q do not mention %q", got, want)
}
}
if strings.Contains(got, "services[1]") {
t.Fatalf("valid service reported: %q", got)
}
}
func TestSchemaWarningsAcceptsValidConfig(t *testing.T) {
data, err := os.ReadFile("config.example.yml")
if err != nil {
t.Fatal(err)
}
var raw appConfig
if err := yaml.Unmarshal(data, &raw); err != nil {
t.Fatal(err)
}
if got := schemaWarnings(data, raw); len(got) != 0 {
t.Fatalf("config.example.yml warnings: %q", got)
}
}
+21
View File
@@ -7,6 +7,7 @@ import (
"os/signal"
"sync"
"syscall"
"time"
)
func main() {
@@ -31,5 +32,25 @@ func main() {
<-sigCh
cancel()
waitShutdown(&wg, shutdownTimeout)
}
// shutdownTimeout stays under Docker's default 10s stop grace period. A
// driver stuck in a handshake that ignores cancellation (lib/pq) would
// otherwise hold the process until Docker sends SIGKILL.
const shutdownTimeout = 7 * time.Second
func waitShutdown(wg *sync.WaitGroup, timeout time.Duration) bool {
done := make(chan struct{})
go func() {
wg.Wait()
close(done)
}()
select {
case <-done:
return true
case <-time.After(timeout):
log.Printf("shutdown: services still stopping after %s; exiting", timeout)
return false
}
}
+21 -10
View File
@@ -6,6 +6,7 @@ import (
"fmt"
"log"
"net/url"
"strings"
"sync"
"time"
@@ -37,7 +38,7 @@ func runService(ctx context.Context, wg *sync.WaitGroup, config serviceConfig) {
return
}
connectCtx, cancel := context.WithTimeout(ctx, connectTimeout)
connectCtx, cancel := context.WithTimeout(ctx, connectTimeoutFor(a))
err = a.Connect(connectCtx)
cancel()
if err != nil {
@@ -101,6 +102,13 @@ func incrementOnce(ctx context.Context, svc runningService) error {
return nil
}
func connectTimeoutFor(a adapter.Adapter) time.Duration {
if t, ok := a.(adapter.ConnectTimeouter); ok && t.ConnectTimeout() > connectTimeout {
return t.ConnectTimeout()
}
return connectTimeout
}
func closeService(_ context.Context, name string, a adapter.Adapter) {
if a == nil {
return
@@ -112,21 +120,24 @@ func closeService(_ context.Context, name string, a adapter.Adapter) {
}
}
// redactURLError drops the raw URL from a *url.Error. Drivers return one when
// a connection URL fails to parse, and its message repeats the URL, password
// included.
// redactURLError hides connection-string parse errors that quote part of
// the string, password included: *url.Error repeats the whole URL, an escape
// error quotes the bad escape, and lib/pq quotes the token after a stray space
// in a key=value DSN.
func redactURLError(err error) error {
var urlErr *url.Error
if !errors.As(err, &urlErr) {
return err
}
var escapeErr url.EscapeError
if errors.As(urlErr.Err, &escapeErr) {
// The message quotes the bad escape, which may be part of a password.
if errors.As(err, &escapeErr) {
return errors.New("invalid connection URL: invalid percent-escape; percent-encode special characters in the user name and password")
}
if strings.Contains(err.Error(), `missing "=" after`) {
return errors.New("invalid key=value DSN: quote values that contain spaces, like password='a b'")
}
var urlErr *url.Error
if errors.As(err, &urlErr) {
return fmt.Errorf("invalid connection URL: %w", urlErr.Err)
}
return err
}
func waitContext(ctx context.Context, d time.Duration) bool {
timer := time.NewTimer(d)
+70
View File
@@ -201,3 +201,73 @@ func TestRunServiceIncrementsRightAfterConnect(t *testing.T) {
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)
}
}
+26 -6
View File
@@ -61,23 +61,43 @@ func hostFromEndpoint(endpoint string) string {
return ""
}
// Never fall back to splitting arbitrary text: a malformed URL or a
// key=value DSN can carry a password, and the name is logged on every line.
if strings.Contains(endpoint, "://") {
if u, err := url.Parse(endpoint); err == nil {
if host := u.Hostname(); host != "" {
return host
}
return u.Hostname()
}
return ""
}
if host := hostFromMySQLDSN(endpoint); host != "" {
return host
}
host, _, err := net.SplitHostPort(endpoint)
if err == nil && host != "" {
return host
if strings.Contains(endpoint, "=") {
return hostFromKeyValueDSN(endpoint)
}
if strings.ContainsAny(endpoint, " \t@/") {
return ""
}
host, _, err := net.SplitHostPort(endpoint)
if err == nil {
return host
}
return ""
}
// hostFromKeyValueDSN reads host= from a PostgreSQL key=value DSN.
func hostFromKeyValueDSN(dsn string) string {
for _, field := range strings.Fields(dsn) {
if value, ok := strings.CutPrefix(field, "host="); ok {
host := strings.Trim(value, "'")
// host may list several hosts; the first names the service.
host, _, _ = strings.Cut(host, ",")
return host
}
}
return ""
}
+23
View File
@@ -0,0 +1,23 @@
package main
import (
"strings"
"testing"
)
func TestHostFromEndpointNeverLeaksDSNSecrets(t *testing.T) {
for _, tc := range []struct{ in, want string }{
{"postgres://u:p@db.example.com:5432/k", "db.example.com"},
{"host=db.example.com user=u password=SeCrEt:x dbname=k", "db.example.com"},
{"host='db.example.com,db2.example.com' password=SeCrEt:x", "db.example.com"},
{"user=u password=SeCrEt:x dbname=k", ""},
{"postgres://u:SeCrEt%zz@db.example.com:5432/k", ""},
{"u:SeCrEt@tcp(db.example.com:3306)/k", "db.example.com"},
{"cache.example.com:6379", "cache.example.com"},
} {
got := hostFromEndpoint(tc.in)
if got != tc.want || strings.Contains(strings.ToLower(got), "secret") {
t.Errorf("hostFromEndpoint(%q) = %q, want %q", tc.in, got, tc.want)
}
}
}