From f00ba87b56f234172ddde06ace50c0723e4f053b Mon Sep 17 00:00:00 2001 From: tianrking <10758833+tianrking@users.noreply.github.com> Date: Tue, 11 Aug 2026 04:14:03 +0800 Subject: [PATCH] Honor PostgreSQL host and port environment variables --- pkg/driver/postgres/postgres.go | 33 +++++++++- pkg/driver/postgres/postgres_test.go | 96 ++++++++++++++++++++++++++++ 2 files changed, 126 insertions(+), 3 deletions(-) diff --git a/pkg/driver/postgres/postgres.go b/pkg/driver/postgres/postgres.go index 6f8d54c8..2cfa15b4 100644 --- a/pkg/driver/postgres/postgres.go +++ b/pkg/driver/postgres/postgres.go @@ -6,6 +6,7 @@ import ( "fmt" "io" "net/url" + "os" "os/exec" "regexp" "runtime" @@ -99,8 +100,11 @@ func connectionString(u *url.URL) string { query.Del("socket") } + useEnvHostname := hostname == "" && query.Get("host") == "" && os.Getenv("PGHOST") != "" + useEnvPort := port == "" && query.Get("port") == "" && os.Getenv("PGPORT") != "" + // default hostname - if hostname == "" && query.Get("host") == "" { + if hostname == "" && query.Get("host") == "" && !useEnvHostname { switch runtime.GOOS { case "linux": query.Set("host", "/var/run/postgresql") @@ -121,11 +125,17 @@ func connectionString(u *url.URL) string { port = query.Get("port") query.Del("port") } - if port == "" { + if port == "" && !useEnvPort { switch u.Scheme { case "redshift": port = "5439" default: + // lib/pq supplies PostgreSQL's default port when the hostname comes + // from PGHOST. Keeping it out of the URL also lets PGPORT take + // precedence if both environment variables are configured. + if useEnvHostname { + break + } port = "5432" } } @@ -134,7 +144,24 @@ func connectionString(u *url.URL) string { out, _ := url.Parse(u.String()) // force scheme back to postgres if there was another postgres-compatible scheme out.Scheme = "postgres" - out.Host = fmt.Sprintf("%s:%s", hostname, port) + switch { + case hostname != "" && port != "": + out.Host = fmt.Sprintf("%s:%s", hostname, port) + case hostname != "": + out.Host = hostname + case useEnvHostname: + out.Host = "" + if port != "" { + // A URL authority containing only a port would also supply an empty + // hostname and overwrite PGHOST, so preserve an explicit or + // scheme-specific port as a query parameter instead. + query.Set("port", port) + } + case port != "": + out.Host = fmt.Sprintf(":%s", port) + default: + out.Host = "" + } out.RawQuery = query.Encode() return out.String() diff --git a/pkg/driver/postgres/postgres_test.go b/pkg/driver/postgres/postgres_test.go index 6f67df45..164aac92 100644 --- a/pkg/driver/postgres/postgres_test.go +++ b/pkg/driver/postgres/postgres_test.go @@ -2,18 +2,39 @@ package postgres import ( "database/sql" + "errors" "fmt" + "net" "net/url" "runtime" "testing" + "time" "github.com/amacneil/dbmate/v2/pkg/dbmate" "github.com/amacneil/dbmate/v2/pkg/dbtest" "github.com/amacneil/dbmate/v2/pkg/dbutil" + "github.com/lib/pq" "github.com/stretchr/testify/require" ) +var errDialStopped = errors.New("dial stopped") + +type recordingDialer struct { + network string + address string +} + +func (d *recordingDialer) Dial(network, address string) (net.Conn, error) { + d.network = network + d.address = address + return nil, errDialStopped +} + +func (d *recordingDialer) DialTimeout(network, address string, _ time.Duration) (net.Conn, error) { + return d.Dial(network, address) +} + func testPostgresDriver(t *testing.T) *Driver { u := dbtest.GetenvURLOrSkip(t, "POSTGRES_TEST_URL") drv, err := dbmate.New(u).Driver() @@ -160,6 +181,9 @@ func defaultConnString() string { } func TestConnectionString(t *testing.T) { + t.Setenv("PGHOST", "") + t.Setenv("PGPORT", "") + cases := []struct { input string expected string @@ -193,6 +217,78 @@ func TestConnectionString(t *testing.T) { } } +func TestConnectionStringPreservesPostgresEnvironment(t *testing.T) { + cases := []struct { + name string + input string + pgHost string + pgPort string + expected string + }{ + { + name: "host and port from environment", + input: "postgres:///foo", + pgHost: "database.internal", + pgPort: "6543", + expected: "postgres:///foo", + }, + { + name: "host from environment and default postgres port", + input: "postgres:///foo", + pgHost: "database.internal", + expected: "postgres:///foo", + }, + { + name: "explicit host and port from environment", + input: "postgres://database.internal/foo", + pgPort: "6543", + expected: "postgres://database.internal/foo", + }, + { + name: "host from environment and explicit port", + input: "postgres://:6543/foo", + pgHost: "database.internal", + expected: "postgres:///foo?port=6543", + }, + { + name: "redshift keeps its default port with host from environment", + input: "redshift:///foo", + pgHost: "database.internal", + expected: "postgres:///foo?port=5439", + }, + } + + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + t.Setenv("PGHOST", c.pgHost) + t.Setenv("PGPORT", c.pgPort) + + u, err := url.Parse(c.input) + require.NoError(t, err) + + require.Equal(t, c.expected, connectionString(u)) + }) + } +} + +func TestConnectionStringUsesPostgresEnvironment(t *testing.T) { + t.Setenv("PGHOST", "database.internal") + t.Setenv("PGPORT", "6543") + + u, err := url.Parse("postgres:///foo?sslmode=disable") + require.NoError(t, err) + + connector, err := pq.NewConnector(connectionString(u)) + require.NoError(t, err) + + dialer := &recordingDialer{} + connector.Dialer(dialer) + _, err = connector.Connect(t.Context()) + require.ErrorIs(t, err, errDialStopped) + require.Equal(t, "tcp", dialer.network) + require.Equal(t, "database.internal:6543", dialer.address) +} + func TestConnectionArgsForDump(t *testing.T) { cases := []struct { input string