From 47afb850937ae11760af8ac253579cddc0754dc5 Mon Sep 17 00:00:00 2001 From: tiennm99 Date: Fri, 9 Oct 2026 12:36:21 +0700 Subject: [PATCH] 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. --- adapter/adapter.go | 11 ++++++ adapter/couchbase.go | 7 ++++ adapter/mongodb.go | 5 ++- adapter/mysql.go | 1 + adapter/postgresql.go | 34 +++++++++++++----- adapter/postgresql_test.go | 26 ++++++++++++-- adapter/redis.go | 1 + config.go | 46 ++++++++++++++++++++----- config_test.go | 35 +++++++++++++++---- main.go | 23 ++++++++++++- runner.go | 33 ++++++++++++------ runner_test.go | 70 ++++++++++++++++++++++++++++++++++++++ service_name.go | 32 +++++++++++++---- service_name_test.go | 23 +++++++++++++ 14 files changed, 302 insertions(+), 45 deletions(-) create mode 100644 service_name_test.go diff --git a/adapter/adapter.go b/adapter/adapter.go index a392839..44ff7c3 100644 --- a/adapter/adapter.go +++ b/adapter/adapter.go @@ -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 { diff --git a/adapter/couchbase.go b/adapter/couchbase.go index c045913..38ecab5 100644 --- a/adapter/couchbase.go +++ b/adapter/couchbase.go @@ -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 diff --git a/adapter/mongodb.go b/adapter/mongodb.go index ed8db05..cd6f047 100644 --- a/adapter/mongodb.go +++ b/adapter/mongodb.go @@ -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 } diff --git a/adapter/mysql.go b/adapter/mysql.go index 94799cc..17077c0 100644 --- a/adapter/mysql.go +++ b/adapter/mysql.go @@ -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 { diff --git a/adapter/postgresql.go b/adapter/postgresql.go index ed09318..d018441 100644 --- a/adapter/postgresql.go +++ b/adapter/postgresql.go @@ -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") { - return dsn + if strings.HasPrefix(dsn, "postgres://") || strings.HasPrefix(dsn, "postgresql://") { + u, err := url.Parse(dsn) + if err != nil { + // Connect reports the parse error. + return dsn + } + q := u.Query() + if q.Has("connect_timeout") { + return dsn + } + q.Set("connect_timeout", defaultPostgresConnectTimeout) + u.RawQuery = q.Encode() + return u.String() } - 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" + for _, field := range strings.Fields(dsn) { + if strings.HasPrefix(field, "connect_timeout=") { + return dsn + } } + return dsn + " connect_timeout=" + defaultPostgresConnectTimeout } diff --git a/adapter/postgresql_test.go b/adapter/postgresql_test.go index 346cac5..c238ca7 100644 --- a/adapter/postgresql_test.go +++ b/adapter/postgresql_test.go @@ -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) + } +} diff --git a/adapter/redis.go b/adapter/redis.go index 4786b63..77a6146 100644 --- a/adapter/redis.go +++ b/adapter/redis.go @@ -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 { diff --git a/config.go b/config.go index f552073..bd5cfdf 100644 --- a/config.go +++ b/config.go @@ -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) { diff --git a/config_test.go b/config_test.go index 8692d35..6b7f9ff 100644 --- a/config_test.go +++ b/config_test.go @@ -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) } } diff --git a/main.go b/main.go index 676a120..5ec5ee5 100644 --- a/main.go +++ b/main.go @@ -7,6 +7,7 @@ import ( "os/signal" "sync" "syscall" + "time" ) func main() { @@ -31,5 +32,25 @@ func main() { <-sigCh cancel() - wg.Wait() + 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 + } } diff --git a/runner.go b/runner.go index 9454749..fa07689 100644 --- a/runner.go +++ b/runner.go @@ -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,20 +120,23 @@ 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") } - return fmt.Errorf("invalid connection URL: %w", urlErr.Err) + 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 { diff --git a/runner_test.go b/runner_test.go index 11529a1..1793df0 100644 --- a/runner_test.go +++ b/runner_test.go @@ -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) + } +} diff --git a/service_name.go b/service_name.go index af08ead..8252d2c 100644 --- a/service_name.go +++ b/service_name.go @@ -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 "" } diff --git a/service_name_test.go b/service_name_test.go new file mode 100644 index 0000000..487b655 --- /dev/null +++ b/service_name_test.go @@ -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) + } + } +}