diff --git a/conn.go b/conn.go index 229c380c5..1e60af2d8 100644 --- a/conn.go +++ b/conn.go @@ -359,8 +359,13 @@ func (c *Conn) Prepare(ctx context.Context, name, sql string) (sd *pgconn.Statem sd, err = c.pgConn.Prepare(ctx, psName, sql, nil) if err != nil { var pErr *pgconn.PrepareError - if errors.As(err, &pErr) { - c.failedDescribeStatement = psKey + if errors.As(err, &pErr) && pErr.ParseComplete { + // The server-side statement was created under psName — the name sent in + // Parse. In the name == sql case psKey is the SQL text, and deallocating + // by it would close a nonexistent statement while leaking the real one. + // When Parse never completed no statement was created at all, so there + // is nothing to clean up. + c.failedDescribeStatement = psName } return nil, err } diff --git a/conn_test.go b/conn_test.go index db3578b4b..60fd63e88 100644 --- a/conn_test.go +++ b/conn_test.go @@ -3,7 +3,9 @@ package pgx_test import ( "bytes" "context" + "crypto/sha256" "database/sql" + "encoding/hex" "io" "net" "os" @@ -532,6 +534,129 @@ func TestPrepareHandlesTimeoutBetweenParseAndDescribe(t *testing.T) { require.NotNil(t, psd) } +// https://github.com/jackc/pgx/issues/2640 +func TestPrepareWithDigestedNameHandlesTimeoutBetweenParseAndDescribe(t *testing.T) { + // Not parallel because it is a timing sensitive test. + // + // stdlib (and therefore database/sql) calls Prepare(ctx, sql, sql). In that case the statement is prepared on the + // server under stmt_ while the client keys it by the SQL text. Cleanup after a Describe phase failure must + // deallocate the digest name — deallocating by the SQL text closes a nonexistent statement, so the leaked + // statement stays on the server and re-preparing the same SQL fails with 42P05 for the rest of the connection's + // life. + + config, err := pgx.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + var faultyConn *faultyconn.Conn + config.AfterNetConnect = func(ctx context.Context, config *pgconn.Config, conn net.Conn) (net.Conn, error) { + faultyConn = faultyconn.New(conn) + return faultyConn, nil + } + + ctx := context.Background() + conn, err := pgx.ConnectConfig(ctx, config) + require.NoError(t, err) + defer closeConn(t, conn) + require.NotNil(t, faultyConn) + + pgxtest.SkipCockroachDB(t, conn, "Induced error does not occur on CockroachDB") + + _, err = conn.Exec(ctx, "set statement_timeout = '100ms'") + require.NoError(t, err) + + faultyConn.HandleFrontendMessage = func(backendWriter io.Writer, msg pgproto3.FrontendMessage) error { + if _, ok := msg.(*pgproto3.Describe); ok { + time.Sleep(200 * time.Millisecond) + } + buf, err := msg.Encode(nil) + if err != nil { + return err + } + _, err = backendWriter.Write(buf) + return err + } + + sql := "select $1::varchar" + digest := sha256.Sum256([]byte(sql)) + psName := "stmt_" + hex.EncodeToString(digest[0:24]) + + psd, err := conn.Prepare(ctx, sql, sql) + var pgErr *pgconn.PgError + require.ErrorAs(t, err, &pgErr) + require.Equal(t, "57014", pgErr.Code) + require.Nil(t, psd) + + faultyConn.HandleFrontendMessage = nil + + _, err = conn.Exec(ctx, "set statement_timeout = default") + require.NoError(t, err) + + var existsOnServer bool + err = conn.QueryRow( + ctx, + "select exists(select 1 from pg_prepared_statements where name = $1)", + // Avoid using the prepared statement cache or it will clear the broken statement before we can check for its + // existence. + pgx.QueryExecModeExec, + psName, + ).Scan(&existsOnServer) + require.NoError(t, err) + require.True(t, existsOnServer) + + psd, err = conn.Prepare(ctx, sql, sql) + require.NoError(t, err) + require.NotNil(t, psd) +} + +// https://github.com/jackc/pgx/issues/2640 +func TestPrepareFailedParseSchedulesNoCleanup(t *testing.T) { + t.Parallel() + + // A Prepare that fails before Parse completes leaves no statement on the server, so the next Prepare must not + // spend a round trip deallocating anything. + + config, err := pgx.ParseConfig(os.Getenv("PGX_TEST_DATABASE")) + require.NoError(t, err) + + var faultyConn *faultyconn.Conn + config.AfterNetConnect = func(ctx context.Context, config *pgconn.Config, conn net.Conn) (net.Conn, error) { + faultyConn = faultyconn.New(conn) + return faultyConn, nil + } + + ctx := context.Background() + conn, err := pgx.ConnectConfig(ctx, config) + require.NoError(t, err) + defer closeConn(t, conn) + require.NotNil(t, faultyConn) + + sql := "select foo" + psd, err := conn.Prepare(ctx, sql, sql) + require.Error(t, err) + require.Nil(t, psd) + var pErr *pgconn.PrepareError + require.ErrorAs(t, err, &pErr) + require.False(t, pErr.ParseComplete) + + var sentClose bool + faultyConn.HandleFrontendMessage = func(backendWriter io.Writer, msg pgproto3.FrontendMessage) error { + if _, ok := msg.(*pgproto3.Close); ok { + sentClose = true + } + buf, err := msg.Encode(nil) + if err != nil { + return err + } + _, err = backendWriter.Write(buf) + return err + } + + psd, err = conn.Prepare(ctx, "select 1", "select 1") + require.NoError(t, err) + require.NotNil(t, psd) + require.False(t, sentClose) +} + func TestPrepareBadSQLFailure(t *testing.T) { t.Parallel()