diff --git a/.github/workflows/build.yaml b/.github/workflows/build.yaml index ae4faac..a87090c 100644 --- a/.github/workflows/build.yaml +++ b/.github/workflows/build.yaml @@ -24,7 +24,7 @@ jobs: - name: Set up Go uses: actions/setup-go@v2 with: - go-version: 1.17 + go-version: 1.25.3 - name: Unit tests run: | diff --git a/.github/workflows/fuzz.yaml b/.github/workflows/fuzz.yaml new file mode 100644 index 0000000..a6dc53f --- /dev/null +++ b/.github/workflows/fuzz.yaml @@ -0,0 +1,77 @@ +name: Fuzz tests + +on: + pull_request: + branches: + - main + - develop + schedule: + - cron: '17 2 * * *' + workflow_dispatch: + +permissions: + contents: read + +jobs: + regression: + name: Regression corpus + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4.4.0 + - uses: actions/setup-go@d35c59abb061a4a6fb18e82ac0862c26744d6ab5 # v5.5.0 + with: + go-version: '1.25.3' + - name: Run regression corpus + run: go test ./... + + fuzz: + name: ${{ matrix.target }} + needs: regression + runs-on: ubuntu-latest + strategy: + fail-fast: false + matrix: + include: + - target: FuzzReader + package: ./internal/buffer + corpus: internal/buffer/testdata/fuzz/FuzzReader + - target: FuzzSplitCompoundQuery + package: ./pkg/sqlbackend + corpus: pkg/sqlbackend/testdata/fuzz/FuzzSplitCompoundQuery + - target: FuzzStartup + package: . + corpus: testdata/fuzz/FuzzStartup + - target: FuzzSessionRaw + package: . + corpus: testdata/fuzz/FuzzSessionRaw + - target: FuzzSessionStructured + package: . + corpus: testdata/fuzz/FuzzSessionStructured + - target: FuzzRowEncoding + package: . + corpus: testdata/fuzz/FuzzRowEncoding + env: + FUZZ_TARGET: ${{ matrix.target }} + FUZZ_PACKAGE: ${{ matrix.package }} + FUZZ_TIME: ${{ github.event_name == 'schedule' && '10m' || '60s' }} + GOCACHE: ${{ github.workspace }}/.cache/go-build + steps: + - uses: actions/checkout@11d5960a326750d5838078e36cf38b85af677262 # v4.4.0 + - uses: actions/setup-go@d35c59abb061a4a6fb18e82ac0862c26744d6ab5 # v5.5.0 + with: + go-version: '1.25.3' + - uses: actions/cache@0400d5f644dc74513175e3cd8d07132dd4860809 # v4.2.4 + with: + path: ${{ env.GOCACHE }}/fuzz + key: fuzz-${{ runner.os }}-${{ matrix.target }}-${{ hashFiles('go.sum') }} + restore-keys: | + fuzz-${{ runner.os }}-${{ matrix.target }}- + - name: Fuzz target + id: fuzz + run: make fuzz FUZZ_TARGET="$FUZZ_TARGET" FUZZ_PACKAGE="$FUZZ_PACKAGE" TIME="$FUZZ_TIME" + - uses: actions/upload-artifact@5d5d22a31266ced268874388b861e4b58bb5c2f3 # v4.3.1 + if: failure() + with: + name: fuzz-failure-${{ matrix.target }} + path: ${{ matrix.corpus }}/ + if-no-files-found: ignore diff --git a/.github/workflows/lint.yaml b/.github/workflows/lint.yaml index 9a9b2f5..4256a89 100644 --- a/.github/workflows/lint.yaml +++ b/.github/workflows/lint.yaml @@ -29,7 +29,7 @@ jobs: - name: Set up Go uses: actions/setup-go@v2 with: - go-version: 1.17 + go-version: 1.25.3 - name: Restore bin uses: actions/cache@v4 diff --git a/Makefile b/Makefile index 2bd7081..42d5fbb 100644 --- a/Makefile +++ b/Makefile @@ -6,6 +6,7 @@ BUILD_DIR = $(CURDIR)/build GOPATH = $(HOME)/go GOBIN = $(GOPATH)/bin GO ?= GOGC=off $(shell which go) +FUZZ_GO = $(subst GOGC=off ,,$(GO)) GOLANGCI_LINT_VERSION = v2.5.0 @@ -39,6 +40,11 @@ lint: | $(GOLANGCI_LINT) ; $(info $(M) running golint…) @ ## Run the project l test: ## Run all tests $Q $(GO) test ./... +.PHONY: fuzz +fuzz: ## Run one fuzz target (FUZZ_TARGET=FuzzXxx FUZZ_PACKAGE=./pkg TIME=60s) + @if [ -z "$(FUZZ_TARGET)" ] || [ -z "$(FUZZ_PACKAGE)" ]; then echo "set FUZZ_TARGET and FUZZ_PACKAGE"; exit 2; fi + $Q GOGC=100 $(FUZZ_GO) test -run=^$$ -fuzz=^$(FUZZ_TARGET)$$ -fuzztime=$(TIME) -parallel=1 $(FUZZ_PACKAGE) + .PHONY: fmt fmt: ; $(info $(M) running gofmt…) @ ## Run gofmt on all source files $Q $(GO) fmt $(PKGS) diff --git a/command.go b/command.go index 1be9b26..8098f24 100644 --- a/command.go +++ b/command.go @@ -192,7 +192,7 @@ func (srv *Server) handleCommand(ctx context.Context, conn SQLConnection, t type } func (srv *Server) handleSimpleQuery(ctx context.Context, cn SQLConnection) error { - if srv.SimpleQuery == nil && srv.SQLBackendFactory == nil { + if srv.SimpleQuery == nil && !cn.HasSQLBackend() { ErrorCode(cn, NewErrUnimplementedMessageType(types.ClientSimpleQuery)) return readyForQuery(cn, types.ServerIdle) } diff --git a/docs/fuzzing.md b/docs/fuzzing.md new file mode 100644 index 0000000..ef6bc96 --- /dev/null +++ b/docs/fuzzing.md @@ -0,0 +1,57 @@ +# Fuzzing + +The project uses Go's native fuzzing support. Fuzz tests and deterministic +regression seeds live beside the packages they exercise; the ordinary +`go test ./...` run executes every seed. + +## Run a target + +Go runs one fuzz target per invocation. Run a target for a time limit with +`make fuzz`, naming both the target and its package: + +```sh +make fuzz FUZZ_TARGET=FuzzSessionStructured FUZZ_PACKAGE=. TIME=5m +make fuzz FUZZ_TARGET=FuzzReader FUZZ_PACKAGE=./internal/buffer TIME=5m +``` + +The available targets are `FuzzReader` (`./internal/buffer`), +`FuzzSplitCompoundQuery` (`./pkg/sqlbackend`), and the root-package targets +`FuzzStartup`, `FuzzSessionRaw`, `FuzzSessionStructured`, and `FuzzRowEncoding`. +The CI workflow runs the regression corpus on every pull request, then fuzzes +each target for one minute. Its nightly run fuzzes each target for ten minutes +and caches Go's generated corpus outside the repository. + +## Add a target or seed + +Add a `FuzzXxx(*testing.F)` test beside the package code. Seed it with +`f.Add` or add a Go fuzz corpus file under `testdata/fuzz/FuzzXxx/`. Keep +inputs synthetic and deterministic. For protocol targets, exercise complete +message sequences as well as malformed framing, and check responses with an +independent decoder. State protocol invariants in the test with a link to the +relevant PostgreSQL protocol section. + +The root integration tests can optionally capture client-to-server bytes and +write sanitized startup and command seeds: + +```sh +PSQL_WIRE_RECORD_FUZZ_SEEDS=1 go test -run '^TestClientConnect$' . +``` + +The capture code replaces client startup parameters with synthetic values +before writing a seed. Review every resulting corpus file before committing it; +never commit passwords, hostnames, real query text, or other environment data. + +## Triage a failure + +Go writes a minimized failing input under the target's `testdata/fuzz/FuzzXxx/` +directory. Reproduce it with the command printed by Go, or run all deterministic +regressions with: + +```sh +go test ./... +``` + +Keep the failing input as a regression seed after fixing the underlying defect. +Do not delete a seed or narrow a target just to make the fuzz run pass. Review +the minimized bytes, identify the violated invariant or panic, make the +smallest correct fix, and rerun both the specific target and the full tests. diff --git a/format.go b/format.go index 6b1d5be..c1d27ae 100644 --- a/format.go +++ b/format.go @@ -11,11 +11,23 @@ type FormatCode int16 // Encoder returns the format encoder for the given data type func (code FormatCode) Encoder(t *pgtype.DataType) FormatEncoder { + if t == nil || t.Value == nil { + return unknownEncoderfunc(fmt.Errorf("format %d has no data type value", code)) + } + switch code { case TextFormat: - return t.Value.(pgtype.TextEncoder).EncodeText + encoder, ok := t.Value.(pgtype.TextEncoder) + if !ok { + return unknownEncoderfunc(fmt.Errorf("data type %q does not support text encoding", t.Name)) + } + return encoder.EncodeText case BinaryFormat: - return t.Value.(pgtype.BinaryEncoder).EncodeBinary + encoder, ok := t.Value.(pgtype.BinaryEncoder) + if !ok { + return unknownEncoderfunc(fmt.Errorf("data type %q does not support binary encoding", t.Name)) + } + return encoder.EncodeBinary default: return unknownEncoderfunc(fmt.Errorf("unknown format encoder %d", code)) } diff --git a/go.mod b/go.mod index 2726995..f3198e0 100644 --- a/go.mod +++ b/go.mod @@ -1,6 +1,6 @@ module github.com/stackql/psql-wire -go 1.16 +go 1.25.3 require ( github.com/jackc/pgtype v1.8.1 @@ -10,3 +10,20 @@ require ( go.uber.org/zap v1.19.1 golang.org/x/tools v0.1.5 ) + +require ( + github.com/jackc/chunkreader/v2 v2.0.1 // indirect + github.com/jackc/pgconn v1.9.1-0.20210724152538-d89c8390a530 // indirect + github.com/jackc/pgio v1.0.0 // indirect + github.com/jackc/pgpassfile v1.0.0 // indirect + github.com/jackc/pgproto3/v2 v2.1.1 // indirect + github.com/jackc/pgservicefile v0.0.0-20200714003250-2b9c44734f2b // indirect + github.com/konsorten/go-windows-terminal-sequences v1.0.2 // indirect + go.uber.org/atomic v1.7.0 // indirect + go.uber.org/multierr v1.6.0 // indirect + golang.org/x/crypto v0.0.0-20210711020723-a769d52b0f97 // indirect + golang.org/x/mod v0.4.2 // indirect + golang.org/x/sys v0.0.0-20210615035016-665e8c7367d1 // indirect + golang.org/x/text v0.3.6 // indirect + golang.org/x/xerrors v0.0.0-20200804184101-5ec99f83aff1 // indirect +) diff --git a/go.sum b/go.sum index ba8949c..d67dbbc 100644 --- a/go.sum +++ b/go.sum @@ -16,7 +16,6 @@ github.com/go-stack/stack v1.8.0/go.mod h1:v0f6uXyyMGvRgIKkXu+yp6POWl0qKG85gN/me github.com/gofrs/uuid v4.0.0+incompatible h1:1SD/1F5pU8p29ybwgQSwpQk+mwdRrXCYuPhW6m+TnJw= github.com/gofrs/uuid v4.0.0+incompatible/go.mod h1:b2aQJv3Z4Fp6yNu3cdSllBxTCLRxnplIgP/c0N/04lM= github.com/google/renameio v0.1.0/go.mod h1:KWCgfxg9yswjAJkECMjeO8J8rahYeXnNhOm40UhjYkI= -github.com/jackc/chunkreader v1.0.0 h1:4s39bBR8ByfqH+DKm8rQA3E1LHZWB9XWcrz8fqaZbe0= github.com/jackc/chunkreader v1.0.0/go.mod h1:RT6O25fNZIuasFJRyZ4R/Y2BbhasbmZXF9QQ7T3kePo= github.com/jackc/chunkreader/v2 v2.0.0/go.mod h1:odVSm741yZoC3dpHEUXIqA9tQRhFrgOHwnPIn9lDKlk= github.com/jackc/chunkreader/v2 v2.0.1 h1:i+RDz65UE+mmpjTfyz0MoVTnzeYxroil2G82ki7MGG8= @@ -36,7 +35,6 @@ github.com/jackc/pgmock v0.0.0-20210724152146-4ad1a8207f65 h1:DadwsjnMwFjfWc9y5W github.com/jackc/pgmock v0.0.0-20210724152146-4ad1a8207f65/go.mod h1:5R2h2EEX+qri8jOWMbJCtaPWkrrNc7OHwsp2TCqp7ak= github.com/jackc/pgpassfile v1.0.0 h1:/6Hmqy13Ss2zCq62VdNG8tM1wchn8zjSGOBJ6icpsIM= github.com/jackc/pgpassfile v1.0.0/go.mod h1:CEx0iS5ambNFdcRtxPj5JhEz+xB6uRky5eyVu/W2HEg= -github.com/jackc/pgproto3 v1.1.0 h1:FYYE4yRw+AgI8wXIinMlNjBbp/UitDJwfj5LqqewP1A= github.com/jackc/pgproto3 v1.1.0/go.mod h1:eR5FA3leWg7p9aeAqi37XOTgTIbkABlvcPB3E5rlc78= github.com/jackc/pgproto3/v2 v2.0.0-alpha1.0.20190420180111-c116219b62db/go.mod h1:bhq50y+xrl9n5mRYyCBFKkpRVTLYJVWeCc+mEAI3yXA= github.com/jackc/pgproto3/v2 v2.0.0-alpha1.0.20190609003834-432c2951c711/go.mod h1:uH0AWtUmuShn0bcesswc4aBTWGvw0cAxIJp+6OB//Wg= diff --git a/internal/buffer/reader.go b/internal/buffer/reader.go index 73856b0..7524bd9 100644 --- a/internal/buffer/reader.go +++ b/internal/buffer/reader.go @@ -4,6 +4,7 @@ import ( "bufio" "bytes" "encoding/binary" + "fmt" "io" "unsafe" @@ -191,6 +192,10 @@ func (reader *simpleReader) GetPrepareType() (PrepareType, error) { // GetBytes returns the buffer's contents as a []byte. func (reader *simpleReader) GetBytes(n int) ([]byte, error) { + if n < 0 { + return nil, fmt.Errorf("negative byte count %d", n) + } + if len(reader.Msg) < n { return nil, NewInsufficientData(len(reader.Msg)) } diff --git a/internal/buffer/reader_fuzz_test.go b/internal/buffer/reader_fuzz_test.go new file mode 100644 index 0000000..06292fd --- /dev/null +++ b/internal/buffer/reader_fuzz_test.go @@ -0,0 +1,40 @@ +package buffer + +import ( + "bytes" + "testing" +) + +func FuzzReader(f *testing.F) { + f.Add([]byte{0xff}) + f.Add([]byte{0, 1, 2, 3, 4, 5, 6, 7}) + f.Add([]byte{0, 0, 0, 4, 'Q', 0, 0, 0}) + + f.Fuzz(func(t *testing.T, input []byte) { + reader := CreateTestReader(input, nil) + if len(input) > 0 { + _, _ = reader.GetBytes(int(int8(input[0]))) + } + for i, op := range input { + switch op % 4 { + case 0: + n := int(int8(op)) + _, _ = reader.GetBytes(n) + case 1: + _, _ = reader.GetString() + case 2: + _, _ = reader.GetUint16() + case 3: + _, _ = reader.GetUint32() + } + if i >= 127 { + break + } + } + + untyped := NewReader(bytes.NewReader(input), 256) + _, _ = untyped.ReadUntypedMsg() + typed := NewReader(bytes.NewReader(input), 256) + _, _, _ = typed.ReadTypedMsg() + }) +} diff --git a/internal/buffer/reader_test.go b/internal/buffer/reader_test.go index e926469..3285684 100644 --- a/internal/buffer/reader_test.go +++ b/internal/buffer/reader_test.go @@ -174,6 +174,21 @@ func TestGetStringNulTerminatorNotfound(t *testing.T) { } } +func TestGetBytesNegativeCount(t *testing.T) { + reader := CreateTestReader([]byte("data"), nil) + + value, err := reader.GetBytes(-1) + if err == nil { + t.Fatal("expected an error for a negative byte count") + } + if value != nil { + t.Fatalf("unexpected result for a negative byte count: %q", value) + } + if got := string(reader.PeekMsg()); got != "data" { + t.Fatalf("negative byte count consumed data: got %q", got) + } +} + func TestGetInsufficientData(t *testing.T) { buffer := bytes.NewBuffer([]byte{}) reader := CreateTestReader( diff --git a/internal/buffer/testdata/fuzz/FuzzReader/negative-byte-count b/internal/buffer/testdata/fuzz/FuzzReader/negative-byte-count new file mode 100644 index 0000000..910cd46 --- /dev/null +++ b/internal/buffer/testdata/fuzz/FuzzReader/negative-byte-count @@ -0,0 +1,2 @@ +go test fuzz v1 +[]byte("\xff") diff --git a/pkg/sqlbackend/pgsqlbackend_fuzz_test.go b/pkg/sqlbackend/pgsqlbackend_fuzz_test.go new file mode 100644 index 0000000..88db998 --- /dev/null +++ b/pkg/sqlbackend/pgsqlbackend_fuzz_test.go @@ -0,0 +1,24 @@ +package sqlbackend + +import ( + "strings" + "testing" +) + +func FuzzSplitCompoundQuery(f *testing.F) { + f.Add("") + f.Add("select 1;select 2") + f.Add(`select "a;b";select "c\";d"`) + f.Add(";;;") + + f.Fuzz(func(t *testing.T, query string) { + backend := NewSimpleSQLBackend(nil) + parts, err := backend.SplitCompoundQuery(query) + if err != nil { + t.Fatal(err) + } + if got := strings.Join(parts, ";"); got != query { + t.Fatalf("joining split query changed it: got %q, want %q", got, query) + } + }) +} diff --git a/row.go b/row.go index 28eb0c1..8386218 100644 --- a/row.go +++ b/row.go @@ -4,6 +4,7 @@ import ( "context" "errors" "fmt" + "strconv" "strings" "github.com/jackc/pgtype" @@ -123,6 +124,14 @@ func (column Column) Write(ctx context.Context, writer buffer.Writer, src interf src = s2 } } + case *pgtype.Polygon: + if polygonText, ok := src.(string); ok { + points, err := parsePolygonText(polygonText) + if err != nil { + return err + } + src = points + } } err = typed.Value.Set(src) if err != nil { @@ -141,6 +150,31 @@ func (column Column) Write(ctx context.Context, writer buffer.Writer, src interf return nil } +func parsePolygonText(src string) ([]pgtype.Vec2, error) { + if len(src) < 7 || !strings.HasPrefix(src, "((") || !strings.HasSuffix(src, "))") { + return nil, fmt.Errorf("invalid polygon representation") + } + + parts := strings.Split(src[2:len(src)-2], "),(") + points := make([]pgtype.Vec2, 0, len(parts)) + for _, part := range parts { + coordinates := strings.Split(part, ",") + if len(coordinates) != 2 { + return nil, fmt.Errorf("invalid polygon point %q", part) + } + x, err := strconv.ParseFloat(coordinates[0], 64) + if err != nil { + return nil, fmt.Errorf("invalid polygon x coordinate: %w", err) + } + y, err := strconv.ParseFloat(coordinates[1], 64) + if err != nil { + return nil, fmt.Errorf("invalid polygon y coordinate: %w", err) + } + points = append(points, pgtype.Vec2{X: x, Y: y}) + } + return points, nil +} + // asTextBytes extracts raw bytes from string or []byte sources for text-format // passthrough. Returns (nil, true) for NULL-representing strings. func asTextBytes(src interface{}) ([]byte, bool) { diff --git a/row_fuzz_test.go b/row_fuzz_test.go new file mode 100644 index 0000000..750ddb9 --- /dev/null +++ b/row_fuzz_test.go @@ -0,0 +1,73 @@ +package wire + +import ( + "bytes" + "context" + "fmt" + "reflect" + "testing" + + "github.com/jackc/pgtype" + "github.com/lib/pq/oid" + "github.com/stackql/psql-wire/internal/buffer" +) + +func FuzzRowEncoding(f *testing.F) { + for _, registeredOID := range fuzzRegisteredOIDs() { + f.Add(registeredOID, uint8(TextFormat), []byte{1}) + f.Add(registeredOID, uint8(BinaryFormat), []byte{1}) + } + f.Add(uint32(oid.T_text), uint8(TextFormat), []byte("fuzz")) + f.Add(uint32(oid.T_int4), uint8(BinaryFormat), []byte{0xff, 0xff, 0xff, 0xff}) + + f.Fuzz(func(t *testing.T, oidValue uint32, format uint8, input []byte) { + ctx := setTypeInfo(context.Background()) + column := Column{Oid: oid.Oid(oidValue), Format: FormatCode(format % 2)} + var output bytes.Buffer + writer := buffer.NewWriter(&output) + var value interface{} + if len(input) > 0 { + switch input[0] % 5 { + case 0: + value = nil + case 1: + value = string(input) + case 2: + value = input + case 3: + value = int64(len(input)) + case 4: + value = struct{}{} + } + } + _ = column.Write(ctx, writer, value) + _ = (Columns{column}).Write(ctx, writer, []interface{}{value}) + + dataType := &pgtype.DataType{Name: "unsupported", OID: oidValue, Value: &fuzzUnsupportedType{}} + encoder := column.Format.Encoder(dataType) + _, _ = encoder(TypeInfo(ctx), nil) + + encoder = column.Format.Encoder(nil) + _, _ = encoder(TypeInfo(ctx), nil) + }) +} + +type fuzzUnsupportedType struct{} + +func (*fuzzUnsupportedType) Set(interface{}) error { return nil } +func (*fuzzUnsupportedType) Get() interface{} { return nil } +func (*fuzzUnsupportedType) AssignTo(interface{}) error { return nil } + +func fuzzRegisteredOIDs() []uint32 { + connInfo := pgtype.NewConnInfo() + oids := reflect.ValueOf(connInfo).Elem().FieldByName("oidToDataType") + if !oids.IsValid() || oids.Kind() != reflect.Map { + panic(fmt.Sprintf("pgtype ConnInfo OID registry unavailable: %T", connInfo)) + } + keys := oids.MapKeys() + values := make([]uint32, 0, len(keys)) + for _, key := range keys { + values = append(values, uint32(key.Uint())) + } + return values +} diff --git a/row_test.go b/row_test.go index 1250866..d9bc2f1 100644 --- a/row_test.go +++ b/row_test.go @@ -237,6 +237,17 @@ func TestColumnWrite_NonString_UsesStandardEncoder(t *testing.T) { } } +func TestColumnWrite_PolygonBinaryRejectsMalformedText(t *testing.T) { + ctx := setTypeInfo(context.Background()) + column := Column{Oid: oid.T_polygon, Format: BinaryFormat} + writer := buffer.NewWriter(&bytes.Buffer{}) + + err := column.Write(ctx, writer, "8000000") + if err == nil { + t.Fatal("expected malformed polygon text to return an error") + } +} + func TestResolveResultFormat(t *testing.T) { tests := []struct { name string diff --git a/testdata/fuzz/FuzzRowEncoding/415eef8ff2ade5c6 b/testdata/fuzz/FuzzRowEncoding/415eef8ff2ade5c6 new file mode 100644 index 0000000..01517fa --- /dev/null +++ b/testdata/fuzz/FuzzRowEncoding/415eef8ff2ade5c6 @@ -0,0 +1,4 @@ +go test fuzz v1 +uint32(604) +byte('\x05') +[]byte("8000000") diff --git a/testdata/fuzz/FuzzSessionRaw/recorded-4babf41ae431e912 b/testdata/fuzz/FuzzSessionRaw/recorded-4babf41ae431e912 new file mode 100644 index 0000000..547aea6 --- /dev/null +++ b/testdata/fuzz/FuzzSessionRaw/recorded-4babf41ae431e912 @@ -0,0 +1,2 @@ +go test fuzz v1 +[]byte("X\x00\x00\x00\x04") diff --git a/testdata/fuzz/FuzzSessionRaw/recorded-8a2e315b22b05895 b/testdata/fuzz/FuzzSessionRaw/recorded-8a2e315b22b05895 new file mode 100644 index 0000000..5d53d26 --- /dev/null +++ b/testdata/fuzz/FuzzSessionRaw/recorded-8a2e315b22b05895 @@ -0,0 +1,2 @@ +go test fuzz v1 +[]byte("Q\x00\x00\x00\x06;\x00X\x00\x00\x00\x04") diff --git a/testdata/fuzz/FuzzSessionRaw/simple-query-terminate b/testdata/fuzz/FuzzSessionRaw/simple-query-terminate new file mode 100644 index 0000000..e1181e0 --- /dev/null +++ b/testdata/fuzz/FuzzSessionRaw/simple-query-terminate @@ -0,0 +1,2 @@ +go test fuzz v1 +[]byte("Q\x00\x00\x00\x0dselect 1\x00X\x00\x00\x00\x04") diff --git a/testdata/fuzz/FuzzSessionStructured/backend-factory-failure b/testdata/fuzz/FuzzSessionStructured/backend-factory-failure new file mode 100644 index 0000000..74987fa --- /dev/null +++ b/testdata/fuzz/FuzzSessionStructured/backend-factory-failure @@ -0,0 +1,2 @@ +go test fuzz v1 +[]byte("\x06\x07") diff --git a/testdata/fuzz/FuzzSessionStructured/complete-extended-cycle b/testdata/fuzz/FuzzSessionStructured/complete-extended-cycle new file mode 100644 index 0000000..c1c5b91 --- /dev/null +++ b/testdata/fuzz/FuzzSessionStructured/complete-extended-cycle @@ -0,0 +1,2 @@ +go test fuzz v1 +[]byte("\x00\x00\x01\x02\x03\x04\x05\x06\x07\x08") diff --git a/testdata/fuzz/FuzzStartup/recorded-94a681ddf692e469 b/testdata/fuzz/FuzzStartup/recorded-94a681ddf692e469 new file mode 100644 index 0000000..72f8968 --- /dev/null +++ b/testdata/fuzz/FuzzStartup/recorded-94a681ddf692e469 @@ -0,0 +1,2 @@ +go test fuzz v1 +[]byte("\x00\x00\x00\x00#\x00\x03\x00\x00user\x00fuzzer\x00database\x00fuzz\x00\x00") diff --git a/testdata/fuzz/FuzzStartup/ssl-request b/testdata/fuzz/FuzzStartup/ssl-request new file mode 100644 index 0000000..aa5b4bf --- /dev/null +++ b/testdata/fuzz/FuzzStartup/ssl-request @@ -0,0 +1,2 @@ +go test fuzz v1 +[]byte("\x00\x00\x00\x00\x08\x04\xd2\x16\x2f") diff --git a/testdata/fuzz/FuzzStartup/startup-v3 b/testdata/fuzz/FuzzStartup/startup-v3 new file mode 100644 index 0000000..084e668 --- /dev/null +++ b/testdata/fuzz/FuzzStartup/startup-v3 @@ -0,0 +1,2 @@ +go test fuzz v1 +[]byte("\x00\x00\x00\x23\x00\x03\x00\x00user\x00fuzzer\x00database\x00fuzz\x00\x00") diff --git a/wire_fuzz_test.go b/wire_fuzz_test.go new file mode 100644 index 0000000..8b25453 --- /dev/null +++ b/wire_fuzz_test.go @@ -0,0 +1,685 @@ +package wire + +import ( + "bytes" + "context" + "crypto/sha256" + "crypto/tls" + "encoding/binary" + "errors" + "fmt" + "io" + "net" + "os" + "path/filepath" + "runtime" + "strconv" + "sync" + "testing" + "time" + + "github.com/jackc/pgproto3/v2" + "github.com/lib/pq/oid" + "github.com/sirupsen/logrus" + "github.com/stackql/psql-wire/internal/types" + "github.com/stackql/psql-wire/pkg/sqlbackend" + "github.com/stackql/psql-wire/pkg/sqldata" +) + +const fuzzMessageBufferSize = 256 + +var fuzzLogger = func() *logrus.Logger { + logger := logrus.New() + logger.SetOutput(io.Discard) + return logger +}() + +func FuzzStartup(f *testing.F) { + f.Add(append([]byte{0}, fuzzStartupPacket(uint32(types.Version30))...)) + f.Add(append([]byte{0}, fuzzStartupPacket(uint32(types.VersionSSLRequest))...)) + f.Add(append([]byte{2}, fuzzStartupPacket(uint32(types.Version30))...)) + f.Add(append([]byte{2}, fuzzStartupPacket(uint32(types.VersionSSLRequest))...)) + f.Add(append([]byte{1}, append(fuzzStartupPacket(uint32(types.Version30)), fuzzTypedMessage(byte(types.ClientPassword), []byte("pw\x00"))...)...)) + f.Add(append([]byte{0}, fuzzStartupPacket(uint32(types.VersionCancel))...)) + + f.Fuzz(func(t *testing.T, input []byte) { + mode, startup := byte(0), input + if len(input) > 0 { + mode, startup = input[0]%3, input[1:] + } + + srv := newFuzzServer(nil) + switch mode { + case 1: + srv.Auth = ClearTextPassword(func(_, password string) (bool, error) { + return password == "pw", nil + }) + case 2: + srv.ClientAuth = tls.RequireAndVerifyClientCert + } + runFuzzServer(t, srv, startup, nil) + }) +} + +func FuzzSessionRaw(f *testing.F) { + f.Add(fuzzTypedMessage(byte(types.ClientSimpleQuery), []byte("select 1\x00"))) + f.Add(append( + fuzzTypedMessage(byte(types.ClientSimpleQuery), []byte("select 1\x00")), + fuzzTypedMessage(byte(types.ClientTerminate), nil)..., + )) + f.Add(fuzzTypedMessage(byte(types.ClientSync), nil)) + f.Add([]byte{byte(types.ClientSimpleQuery), 0xff, 0xff, 0xff, 0xff}) + + f.Fuzz(func(t *testing.T, commands []byte) { + input := append(fuzzStartupPacket(uint32(types.Version30)), commands...) + srv := newFuzzServer(newFuzzBackendFactory(2, false)) + runFuzzServer(t, srv, input, nil) + }) +} + +func FuzzSessionStructured(f *testing.F) { + f.Add([]byte{0, 0, 1, 2, 3, 5}) + f.Add([]byte{2, 7, 0, 1, 2, 3, 4, 5, 6, 7}) + f.Add([]byte{6, 0, 1, 3, 5}) + f.Add([]byte{6, 7}) + f.Add([]byte{1, 0, 1, 2, 3, 5}) + + f.Fuzz(func(t *testing.T, script []byte) { + mode := byte(0) + if len(script) > 0 { + mode = script[0] % 7 + script = script[1:] + } + commands, requests := fuzzStructuredMessages(script) + srv := newFuzzServer(newFuzzBackendFactory(mode, mode == 6)) + messages := runFuzzServer(t, srv, append(fuzzStartupPacket(uint32(types.Version30)), commands...), requests) + verifyStructuredResponses(t, requests, messages) + }) +} + +func newFuzzServer(factory sqlbackend.SQLBackendFactory) *Server { + options := []OptionFn{ + MessageBufferSize(fuzzMessageBufferSize), + Logger(fuzzLogger), + } + if factory != nil { + options = append(options, SQLBackendFactory(factory)) + } + srv, err := NewServer(options...) + if err != nil { + panic(err) + } + return srv +} + +type fuzzMemoryConn struct { + input *bytes.Reader + output bytes.Buffer + maxOutput int +} + +func newFuzzMemoryConn(input []byte) *fuzzMemoryConn { + maxOutput := len(input)*16 + 64*1024 + return &fuzzMemoryConn{ + input: bytes.NewReader(input), + maxOutput: maxOutput, + } +} + +func (c *fuzzMemoryConn) Read(p []byte) (int, error) { return c.input.Read(p) } + +func (c *fuzzMemoryConn) Write(p []byte) (int, error) { + if len(p) > c.maxOutput-c.output.Len() { + return 0, errors.New("fuzz output exceeded its input-relative memory budget") + } + return c.output.Write(p) +} + +func (c *fuzzMemoryConn) Close() error { return nil } +func (c *fuzzMemoryConn) LocalAddr() net.Addr { return fuzzAddr("local") } +func (c *fuzzMemoryConn) RemoteAddr() net.Addr { return fuzzAddr("remote") } +func (c *fuzzMemoryConn) SetDeadline(time.Time) error { return nil } +func (c *fuzzMemoryConn) SetReadDeadline(time.Time) error { return nil } +func (c *fuzzMemoryConn) SetWriteDeadline(time.Time) error { return nil } + +type fuzzAddr string + +func (a fuzzAddr) Network() string { return string(a) } +func (a fuzzAddr) String() string { return string(a) } + +func runFuzzServer(t *testing.T, srv *Server, input []byte, requests []byte) []pgproto3.BackendMessage { + t.Helper() + conn := newFuzzMemoryConn(input) + before := runtime.MemStats{} + runtime.ReadMemStats(&before) + started := time.Now() + _ = srv.serve(context.Background(), conn) + elapsed := time.Since(started) + after := runtime.MemStats{} + runtime.ReadMemStats(&after) + + if elapsed > time.Second { + t.Fatalf("serving an exhausted in-memory input took %s", elapsed) + } + allocated := after.TotalAlloc - before.TotalAlloc + allocationLimit := uint64(len(input))*512 + 2*1024*1024 + if allocated > allocationLimit { + t.Fatalf("serving %d input bytes allocated %d bytes (limit %d)", len(input), allocated, allocationLimit) + } + + messages := decodeFuzzBackendMessages(t, conn.output.Bytes()) + if requests != nil && len(messages) > 0 { + readyCount := 0 + for _, message := range messages { + if _, ok := message.(*pgproto3.ReadyForQuery); ok { + readyCount++ + } + } + maxReady := 1 + for _, request := range requests { + if request == byte(types.ClientSimpleQuery) || request == byte(types.ClientSync) { + maxReady++ + } + } + if readyCount == 0 || readyCount > maxReady { + t.Fatalf("ReadyForQuery count %d is outside [1, %d]", readyCount, maxReady) + } + } + return messages +} + +func decodeFuzzBackendMessages(t *testing.T, output []byte) []pgproto3.BackendMessage { + t.Helper() + // PostgreSQL 16 Message Formats frame backend responses independently from + // this library's writer; pgproto3 decodes every complete response: + // https://www.postgresql.org/docs/16/protocol-message-formats.html + // The same section specifies the unframed 'S' or 'N' SSLRequest response. + if len(output) > 0 && output[0] == 'N' && (len(output) == 1 || output[1] == 'R') { + output = output[1:] + } + chunks := &fuzzChunkReader{reader: bytes.NewReader(output), remaining: len(output)} + decoder := pgproto3.NewFrontend(chunks, io.Discard) + var messages []pgproto3.BackendMessage + for chunks.remaining > 0 { + message, err := decoder.Receive() + if err != nil { + t.Fatalf("server output is not a sequence of complete backend messages: %v", err) + } + messages = append(messages, message) + } + return messages +} + +type fuzzChunkReader struct { + reader *bytes.Reader + remaining int +} + +func (r *fuzzChunkReader) Next(n int) ([]byte, error) { + if n < 0 || n > r.remaining { + return nil, io.ErrUnexpectedEOF + } + chunk := make([]byte, n) + if _, err := io.ReadFull(r.reader, chunk); err != nil { + return nil, err + } + r.remaining -= n + return chunk, nil +} + +func fuzzStartupPacket(version uint32) []byte { + body := make([]byte, 4, 64) + binary.BigEndian.PutUint32(body, version) + body = append(body, []byte("user\x00fuzzer\x00database\x00fuzz\x00\x00")...) + packet := make([]byte, 4, len(body)+4) + binary.BigEndian.PutUint32(packet, uint32(len(body)+4)) + return append(packet, body...) +} + +func fuzzTypedMessage(kind byte, body []byte) []byte { + message := make([]byte, 5, len(body)+5) + message[0] = kind + binary.BigEndian.PutUint32(message[1:5], uint32(len(body)+4)) + return append(message, body...) +} + +func fuzzCString(dst []byte, value string) []byte { + dst = append(dst, value...) + return append(dst, 0) +} + +func fuzzStructuredMessages(script []byte) ([]byte, []byte) { + var commands []byte + var requests []byte + for _, op := range script { + var kind byte + var body []byte + switch op % 9 { + case 0: + kind = byte(types.ClientParse) + body = fuzzCString(body, "s") + body = fuzzCString(body, fmt.Sprintf("select %d", op)) + body = append(body, 0, 0) + case 1: + kind = byte(types.ClientBind) + body = fuzzCString(body, "p") + body = fuzzCString(body, "s") + body = append(body, 0, 0, 0, 0, 0, 0) + case 2: + kind = byte(types.ClientDescribe) + if op&0x10 == 0 { + body = append(body, byte('S')) + body = fuzzCString(body, "s") + } else { + body = append(body, byte('P')) + body = fuzzCString(body, "p") + } + case 3: + kind = byte(types.ClientExecute) + body = fuzzCString(body, "p") + var maxRows [4]byte + binary.BigEndian.PutUint32(maxRows[:], uint32(op)) + body = append(body, maxRows[:]...) + case 4: + kind = byte(types.ClientClose) + if op&0x10 == 0 { + body = append(body, byte('S')) + body = fuzzCString(body, "s") + } else { + body = append(body, byte('P')) + body = fuzzCString(body, "p") + } + case 5: + kind = byte(types.ClientSync) + case 6: + kind = byte(types.ClientFlush) + case 7: + kind = byte(types.ClientSimpleQuery) + body = fuzzCString(body, "select "+strconv.Itoa(int(op))) + case 8: + kind = byte(types.ClientTerminate) + } + commands = append(commands, fuzzTypedMessage(kind, body)...) + requests = append(requests, kind) + if kind == byte(types.ClientTerminate) { + break + } + } + return commands, requests +} + +func verifyStructuredResponses(t *testing.T, requests []byte, messages []pgproto3.BackendMessage) { + t.Helper() + if len(messages) == 0 { + t.Fatal("completed structured startup produced no backend messages") + } + // PostgreSQL 16 Protocol Flow specifies the initial ReadyForQuery before + // the first query cycle: https://www.postgresql.org/docs/16/protocol-flow.html + initialReady := -1 + for i, message := range messages { + if _, ok := message.(*pgproto3.ReadyForQuery); ok { + initialReady = i + break + } + } + if initialReady < 0 { + t.Fatal("completed structured startup did not send ReadyForQuery") + } + for i := 0; i < initialReady; i++ { + if _, ok := messages[i].(*pgproto3.ReadyForQuery); ok { + t.Fatal("startup emitted more than one initial ReadyForQuery") + } + } + + // PostgreSQL 16 Protocol Flow specifies that an extended-protocol error + // suppresses responses until Sync; a simple-query error is followed by + // ReadyForQuery: https://www.postgresql.org/docs/16/protocol-flow.html + for i, message := range messages { + if _, ok := message.(*pgproto3.ErrorResponse); ok && i+1 < len(messages) { + if _, ready := messages[i+1].(*pgproto3.ReadyForQuery); !ready { + t.Fatal("server emitted a response after ErrorResponse before Sync") + } + } + } + + // PostgreSQL 16 Protocol Flow specifies one ReadyForQuery for each processed + // Query and Sync; Terminate is not followed by a backend response: + // https://www.postgresql.org/docs/16/protocol-flow.html + verifyReadyResponseGroups(t, requests, messages[initialReady+1:]) +} + +func verifyReadyResponseGroups(t *testing.T, requests []byte, messages []pgproto3.BackendMessage) { + t.Helper() + cursor := 0 + inErrorState := false + for _, request := range requests { + if request == byte(types.ClientTerminate) { + break + } + if inErrorState { + if request == byte(types.ClientSync) { + consumeFuzzResponse(t, messages, &cursor, func(message pgproto3.BackendMessage) bool { + _, ok := message.(*pgproto3.ReadyForQuery) + return ok + }, "Sync ReadyForQuery") + inErrorState = false + } + continue + } + + switch types.ClientMessage(request) { + case types.ClientSimpleQuery, types.ClientSync: + consumeReadyForQuery(t, messages, &cursor) + case types.ClientParse: + inErrorState = consumeFuzzResponse(t, messages, &cursor, isParseCompleteOrError, "ParseComplete or ErrorResponse") + case types.ClientBind: + inErrorState = consumeFuzzResponse(t, messages, &cursor, isBindCompleteOrError, "BindComplete or ErrorResponse") + case types.ClientDescribe: + inErrorState = consumeFuzzResponse(t, messages, &cursor, isDescribeCompleteOrError, "Describe response") + case types.ClientExecute: + inErrorState = consumeFuzzResponse(t, messages, &cursor, isExecuteCompleteOrError, "Execute response") + case types.ClientClose: + inErrorState = consumeFuzzResponse(t, messages, &cursor, isCloseCompleteOrError, "CloseComplete or ErrorResponse") + case types.ClientFlush: + default: + t.Fatalf("unexpected structured request type %q", request) + } + } + if cursor != len(messages) { + t.Fatalf("unexpected trailing backend responses: consumed %d of %d", cursor, len(messages)) + } +} + +func consumeReadyForQuery(t *testing.T, messages []pgproto3.BackendMessage, cursor *int) { + t.Helper() + for *cursor < len(messages) { + message := messages[*cursor] + *cursor++ + if _, ok := message.(*pgproto3.ReadyForQuery); ok { + return + } + } + t.Fatal("missing ReadyForQuery") +} + +func consumeFuzzResponse(t *testing.T, messages []pgproto3.BackendMessage, cursor *int, done func(pgproto3.BackendMessage) bool, description string) bool { + t.Helper() + for *cursor < len(messages) { + message := messages[*cursor] + *cursor++ + if _, failed := message.(*pgproto3.ErrorResponse); failed { + return true + } + if done(message) { + return false + } + } + t.Fatalf("missing %s in server output", description) + return false +} + +func isParseCompleteOrError(message pgproto3.BackendMessage) bool { + _, ok := message.(*pgproto3.ParseComplete) + return ok +} + +func isBindCompleteOrError(message pgproto3.BackendMessage) bool { + _, ok := message.(*pgproto3.BindComplete) + return ok +} + +func isDescribeCompleteOrError(message pgproto3.BackendMessage) bool { + switch message.(type) { + case *pgproto3.NoData, *pgproto3.RowDescription: + return true + default: + return false + } +} + +func isExecuteCompleteOrError(message pgproto3.BackendMessage) bool { + _, ok := message.(*pgproto3.CommandComplete) + return ok +} + +func isCloseCompleteOrError(message pgproto3.BackendMessage) bool { + _, ok := message.(*pgproto3.CloseComplete) + return ok +} + +type fuzzBackendFactory struct { + mode byte + fail bool +} + +func newFuzzBackendFactory(mode byte, fail bool) sqlbackend.SQLBackendFactory { + return fuzzBackendFactory{mode: mode, fail: fail} +} + +func (f fuzzBackendFactory) NewSQLBackend() (sqlbackend.ISQLBackend, error) { + if f.fail { + return nil, errors.New("fuzz backend factory failure") + } + return &fuzzBackend{mode: f.mode}, nil +} + +type fuzzBackend struct { + mode byte +} + +func (b *fuzzBackend) HandleSimpleQuery(context.Context, string) (sqldata.ISQLResultStream, error) { + return b.stream() +} + +func (b *fuzzBackend) SplitCompoundQuery(query string) ([]string, error) { + return sqlbackend.NewSimpleSQLBackend(nil).SplitCompoundQuery(query) +} + +func (b *fuzzBackend) GetDebugStr() string { return "" } + +func (b *fuzzBackend) HandleParse(_ context.Context, _ string, _ string, oids []uint32) ([]uint32, error) { + if b.mode == 1 { + return nil, errors.New("fuzz parse failure") + } + return oids, nil +} + +func (b *fuzzBackend) HandleBind(context.Context, string, string, []int16, [][]byte, []int16) error { + if b.mode == 1 { + return errors.New("fuzz bind failure") + } + return nil +} + +func (b *fuzzBackend) HandleDescribeStatement(_ context.Context, _ string, _ string, oids []uint32) ([]uint32, []sqldata.ISQLColumn, error) { + if b.mode == 1 { + return nil, nil, errors.New("fuzz describe failure") + } + return oids, nil, nil +} + +func (b *fuzzBackend) HandleDescribePortal(context.Context, string, string, string, []uint32) ([]sqldata.ISQLColumn, error) { + if b.mode == 1 { + return nil, errors.New("fuzz describe failure") + } + return nil, nil +} + +func (b *fuzzBackend) HandleExecute(context.Context, string, string, string, []int16, [][]byte, []int16, int32) (sqldata.ISQLResultStream, error) { + return b.stream() +} + +func (b *fuzzBackend) HandleCloseStatement(context.Context, string) error { return nil } +func (b *fuzzBackend) HandleClosePortal(context.Context, string) error { return nil } + +func (b *fuzzBackend) stream() (sqldata.ISQLResultStream, error) { + if b.mode == 1 { + return nil, errors.New("fuzz query failure") + } + if b.mode == 2 { + return nil, nil + } + + columnCount := 1 + oidValue := uint32(oid.T_text) + rowValues := []interface{}{"fuzz"} + if b.mode == 3 { + rowValues = []interface{}{"fuzz", "extra"} + } else if b.mode == 4 { + oidValue = uint32(oid.T_int4) + rowValues = []interface{}{struct{}{}} + } + columns := make([]sqldata.ISQLColumn, columnCount) + for i := range columns { + columns[i] = sqldata.NewSQLColumn(sqldata.NewSQLTable(0, ""), "value", int16(i+1), oidValue, -1, -1, "TextFormat") + } + result := sqldata.NewSQLResult(columns, 0, 0, []sqldata.ISQLRow{ + sqldata.NewSQLRow(rowValues), + }) + return &fuzzResultStream{result: result, failAfterFirst: b.mode == 5}, nil +} + +type fuzzResultStream struct { + result sqldata.ISQLResult + failAfterFirst bool + read bool +} + +func (s *fuzzResultStream) GetColumns() []sqldata.ISQLColumn { return s.result.GetColumns() } + +func (s *fuzzResultStream) Read() (sqldata.ISQLResult, error) { + if !s.read { + s.read = true + if s.failAfterFirst { + return s.result, nil + } + return s.result, io.EOF + } + if s.failAfterFirst { + return nil, errors.New("fuzz stream failure") + } + return nil, io.EOF +} + +func (*fuzzResultStream) Write(sqldata.ISQLResult) error { return errors.New("write unsupported") } +func (*fuzzResultStream) Close() error { return nil } + +type fuzzRecordingListener struct { + net.Listener + mu sync.Mutex + conns []*fuzzRecordingConn +} + +func (l *fuzzRecordingListener) Accept() (net.Conn, error) { + conn, err := l.Listener.Accept() + if err != nil { + return nil, err + } + recording := &fuzzRecordingConn{Conn: conn, done: make(chan struct{})} + l.mu.Lock() + l.conns = append(l.conns, recording) + l.mu.Unlock() + return recording, nil +} + +type fuzzRecordingConn struct { + net.Conn + mu sync.Mutex + received bytes.Buffer + done chan struct{} + close sync.Once +} + +func (c *fuzzRecordingConn) Read(p []byte) (int, error) { + n, err := c.Conn.Read(p) + if n > 0 { + c.mu.Lock() + _, _ = c.received.Write(p[:n]) + c.mu.Unlock() + } + return n, err +} + +func (c *fuzzRecordingConn) Close() error { + err := c.Conn.Close() + c.close.Do(func() { + close(c.done) + }) + return err +} + +func (l *fuzzRecordingListener) recordSeeds(t *testing.T) { + t.Helper() + if os.Getenv("PSQL_WIRE_RECORD_FUZZ_SEEDS") != "1" { + return + } + l.mu.Lock() + conns := append([]*fuzzRecordingConn(nil), l.conns...) + l.mu.Unlock() + + for index, conn := range conns { + select { + case <-conn.done: + case <-time.After(5 * time.Second): + t.Fatalf("timed out waiting for client connection %d to finish", index) + } + + conn.mu.Lock() + clientBytes := append([]byte(nil), conn.received.Bytes()...) + conn.mu.Unlock() + recordFuzzClientSession(t, clientBytes) + } +} + +func recordFuzzClientSession(t *testing.T, clientBytes []byte) { + t.Helper() + if os.Getenv("PSQL_WIRE_RECORD_FUZZ_SEEDS") != "1" { + return + } + if len(clientBytes) < 8 { + return + } + + offset := 0 + var sslRequest []byte + firstLength := int(binary.BigEndian.Uint32(clientBytes[:4])) + if firstLength == 8 { + version := binary.BigEndian.Uint32(clientBytes[4:8]) + if version == uint32(types.VersionSSLRequest) { + sslRequest = clientBytes[:8] + offset = 8 + } else if version != uint32(types.Version30) { + return + } + } + if len(clientBytes)-offset < 8 { + return + } + startupLength := int(binary.BigEndian.Uint32(clientBytes[offset : offset+4])) + if startupLength < 8 || startupLength > len(clientBytes)-offset { + return + } + + startupSeed := append([]byte{0}, sslRequest...) + startupSeed = append(startupSeed, fuzzStartupPacket(uint32(types.Version30))...) + writeFuzzCorpusSeed(t, "FuzzStartup", startupSeed) + + commands := clientBytes[offset+startupLength:] + if len(commands) != 0 { + writeFuzzCorpusSeed(t, "FuzzSessionRaw", commands) + } +} + +func writeFuzzCorpusSeed(t *testing.T, target string, input []byte) { + t.Helper() + digest := sha256.Sum256(input) + name := fmt.Sprintf("recorded-%x", digest[:8]) + path := filepath.Join("testdata", "fuzz", target, name) + if err := os.MkdirAll(filepath.Dir(path), 0o755); err != nil { + t.Fatal(err) + } + content := []byte("go test fuzz v1\n[]byte(" + strconv.QuoteToASCII(string(input)) + ")\n") + if err := os.WriteFile(path, content, 0o644); err != nil { + t.Fatal(err) + } +} diff --git a/wire_test.go b/wire_test.go index 8f988c8..132a49e 100644 --- a/wire_test.go +++ b/wire_test.go @@ -37,6 +37,23 @@ func TListenAndServe(t *testing.T, server *Server) *net.TCPAddr { return listener.Addr().(*net.TCPAddr) } +func TListenAndServeWithCapture(t *testing.T, server *Server) (*net.TCPAddr, *fuzzRecordingListener) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + recording := &fuzzRecordingListener{Listener: listener} + + t.Cleanup(func() { + if err := server.Close(); err != nil { + t.Fatal(err) + } + }) + + go server.Serve(recording) //nolint:errcheck + return listener.Addr().(*net.TCPAddr), recording +} + func TestClientConnect(t *testing.T) { t.Parallel() @@ -49,7 +66,7 @@ func TestClientConnect(t *testing.T) { t.Fatal(err) } - address := TListenAndServe(t, server) + address, captured := TListenAndServeWithCapture(t, server) t.Run("mock", func(t *testing.T) { conn, err := net.Dial("tcp", address.String()) @@ -100,6 +117,8 @@ func TestClientConnect(t *testing.T) { t.Fatal(err) } }) + + captured.recordSeeds(t) } func TestZeroRowResultRetainsColumnMetadata(t *testing.T) {