mirror of
https://github.com/tiennm99/keepalive.git
synced 2026-10-11 03:13:31 +00:00
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:
1 parent
e646af19b8
commit
47afb85093
14 files changed
+302
-45
No files matched your search
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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 {
|
||||
|
||||
+25
-9
@@ -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
|
||||
}
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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)
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
@@ -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
|
||||
}
|
||||
}
|
||||
@@ -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 {
|
||||
|
||||
@@ -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
@@ -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 ""
|
||||
}
|
||||
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user