From 0543995ec34dab6565ae0039b1c8d89d33cdc092 Mon Sep 17 00:00:00 2001 From: naiba Date: Mon, 20 Jul 2026 04:42:03 +0000 Subject: [PATCH] test(agentcompat): harden deterministic fixtures Co-authored-by: naiba/CloudCode --- .../agentcompat/internal/fixture/nat_echo.go | 100 +++++++----- .../internal/fixture/nat_echo_test.go | 101 +++++++++++- .../agentcompat/internal/fixture/nat_hold.go | 147 ++++++++++++++++++ .../internal/fixture/nat_hold_test.go | 126 +++++++++++++++ .../internal/fixture/payload_heap_test.go | 65 ++++++++ .../internal/fixture/payload_measurement.go | 89 +++++++++++ .../internal/fixture/payload_test.go | 20 +-- .../agentcompat/internal/testpaths/source.go | 62 ++++++++ .../internal/testpaths/source_test.go | 61 ++++++++ 9 files changed, 717 insertions(+), 54 deletions(-) create mode 100644 integration/agentcompat/internal/fixture/nat_hold.go create mode 100644 integration/agentcompat/internal/fixture/nat_hold_test.go create mode 100644 integration/agentcompat/internal/fixture/payload_heap_test.go create mode 100644 integration/agentcompat/internal/fixture/payload_measurement.go create mode 100644 integration/agentcompat/internal/testpaths/source.go create mode 100644 integration/agentcompat/internal/testpaths/source_test.go diff --git a/integration/agentcompat/internal/fixture/nat_echo.go b/integration/agentcompat/internal/fixture/nat_echo.go index d0548216..5f6847eb 100644 --- a/integration/agentcompat/internal/fixture/nat_echo.go +++ b/integration/agentcompat/internal/fixture/nat_echo.go @@ -9,15 +9,19 @@ import ( "net" "net/http" "sync" + + "github.com/nezhahq/nezha/pkg/agentcompatcontract" ) type NATEchoRecord struct { - Method string - Path string - Host string - HeaderValue string - Body []byte - RequestHalfClosed bool + Method string + Path string + Host string + HeaderValue string + Body []byte + RequestHalfClosed bool + ResponseHalfClosed bool + SensitiveHeadersPresent bool } type natEchoResult struct { @@ -26,42 +30,53 @@ type natEchoResult struct { } type NATEchoBackend struct { - listener net.Listener - results chan natEchoResult - done chan struct{} - connectionsReady chan struct{} - requireHalfClose bool - closeOnce sync.Once - waitGroup sync.WaitGroup - mutex sync.Mutex - connections map[net.Conn]struct{} - closing bool + listener net.Listener + results chan natEchoResult + done chan struct{} + connectionsReady chan struct{} + requireRequestHalfClose bool + halfCloseResponse bool + closeOnce sync.Once + closeErr error + waitGroup sync.WaitGroup + mutex sync.Mutex + connections map[net.Conn]struct{} + closing bool } func StartNATEchoBackend() (*NATEchoBackend, error) { - return startNATEchoBackend(false) + return startNATEchoBackend(false, false) } func StartNATHalfCloseEchoBackend() (*NATEchoBackend, error) { - return startNATEchoBackend(true) + return startNATEchoBackend(true, false) } -func startNATEchoBackend(requireHalfClose bool) (*NATEchoBackend, error) { +func StartNATResponseHalfCloseEchoBackend() (*NATEchoBackend, error) { + return startNATEchoBackend(false, true) +} + +func startNATEchoBackend(requireRequestHalfClose, halfCloseResponse bool) (*NATEchoBackend, error) { listener, err := net.Listen("tcp", "127.0.0.1:0") if err != nil { return nil, fmt.Errorf("listen for NAT echo: %w", err) } + return newNATEchoBackend(listener, requireRequestHalfClose, halfCloseResponse), nil +} + +func newNATEchoBackend(listener net.Listener, requireRequestHalfClose, halfCloseResponse bool) *NATEchoBackend { backend := &NATEchoBackend{ - listener: listener, - results: make(chan natEchoResult, 16), - done: make(chan struct{}), - connectionsReady: make(chan struct{}, 16), - requireHalfClose: requireHalfClose, - connections: make(map[net.Conn]struct{}), + listener: listener, + results: make(chan natEchoResult, 16), + done: make(chan struct{}), + connectionsReady: make(chan struct{}, 16), + requireRequestHalfClose: requireRequestHalfClose, + halfCloseResponse: halfCloseResponse, + connections: make(map[net.Conn]struct{}), } backend.waitGroup.Add(1) go backend.accept() - return backend, nil + return backend } func (backend *NATEchoBackend) Address() string { @@ -91,10 +106,9 @@ func (backend *NATEchoBackend) WaitConnection(ctx context.Context) error { } func (backend *NATEchoBackend) Close() error { - var closeErr error backend.closeOnce.Do(func() { close(backend.done) - closeErr = backend.listener.Close() + backend.closeErr = normalizeNATEchoCloseError(backend.listener.Close()) backend.mutex.Lock() backend.closing = true connections := make([]net.Conn, 0, len(backend.connections)) @@ -103,14 +117,20 @@ func (backend *NATEchoBackend) Close() error { } backend.mutex.Unlock() for _, connection := range connections { - closeErr = errors.Join(closeErr, connection.Close()) + backend.closeErr = errors.Join(backend.closeErr, normalizeNATEchoCloseError(connection.Close())) } backend.waitGroup.Wait() }) - if errors.Is(closeErr, net.ErrClosed) { + // sync.Once publishes the first shutdown result after cleanup completes, so + // concurrent and repeated callers observe the same outcome. + return backend.closeErr +} + +func normalizeNATEchoCloseError(err error) error { + if errors.Is(err, net.ErrClosed) { return nil } - return closeErr + return err } func (backend *NATEchoBackend) accept() { @@ -164,7 +184,7 @@ func (backend *NATEchoBackend) handle(connection net.Conn) { return } halfClosed := false - if backend.requireHalfClose { + if backend.requireRequestHalfClose { _, halfCloseErr := reader.ReadByte() halfClosed = errors.Is(halfCloseErr, io.EOF) if halfCloseErr != nil && !halfClosed { @@ -173,12 +193,13 @@ func (backend *NATEchoBackend) handle(connection net.Conn) { } } record := NATEchoRecord{ - Method: request.Method, - Path: request.URL.RequestURI(), - Host: request.Host, - HeaderValue: request.Header.Get("X-AgentCompat-Echo"), - Body: append([]byte(nil), body...), - RequestHalfClosed: halfClosed, + Method: request.Method, + Path: request.URL.RequestURI(), + Host: request.Host, + HeaderValue: request.Header.Get("X-AgentCompat-Echo"), + Body: append([]byte(nil), body...), + RequestHalfClosed: halfClosed, + SensitiveHeadersPresent: request.Header.Get(agentcompatcontract.IOStreamCapabilityHeader) != "" || request.Header.Get("Authorization") != "", } responseBody := fmt.Sprintf("method=%s\npath=%s\nhost=%s\nx-agentcompat-echo=%s\nbody=%s\n", record.Method, record.Path, record.Host, record.HeaderValue, record.Body) response := fmt.Sprintf("HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: %d\r\nConnection: close\r\n\r\n%s", len(responseBody), responseBody) @@ -186,11 +207,12 @@ func (backend *NATEchoBackend) handle(connection net.Conn) { backend.publish(natEchoResult{err: fmt.Errorf("write NAT echo response: %w", err)}) return } - if tcpConnection, ok := connection.(*net.TCPConn); ok { + if tcpConnection, ok := connection.(*net.TCPConn); ok && backend.halfCloseResponse { if err := tcpConnection.CloseWrite(); err != nil { backend.publish(natEchoResult{err: fmt.Errorf("half-close NAT echo response: %w", err)}) return } + record.ResponseHalfClosed = true } backend.publish(natEchoResult{record: record}) } diff --git a/integration/agentcompat/internal/fixture/nat_echo_test.go b/integration/agentcompat/internal/fixture/nat_echo_test.go index 57d1dc6f..cde9da72 100644 --- a/integration/agentcompat/internal/fixture/nat_echo_test.go +++ b/integration/agentcompat/internal/fixture/nat_echo_test.go @@ -3,11 +3,13 @@ package fixture import ( "bufio" "context" + "errors" "fmt" "io" "net" "net/http" "strings" + "sync" "testing" "time" ) @@ -78,11 +80,40 @@ func TestFixture_NATEcho(t *testing.T) { if response.StatusCode != http.StatusOK || string(responseBody) != expected { t.Fatalf("NAT echo response = %d %q", response.StatusCode, responseBody) } - if record.RequestHalfClosed || record.Method != http.MethodPut || record.Host != "nat.invalid" { + if record.RequestHalfClosed || record.ResponseHalfClosed || record.Method != http.MethodPut || record.Host != "nat.invalid" { t.Fatalf("NAT request record = %+v", record) } } +func TestFixture_NATResponseHalfCloseEcho(t *testing.T) { + // Given + backend, err := StartNATResponseHalfCloseEchoBackend() + requireNoFixtureError(t, err) + t.Cleanup(func() { requireNoFixtureError(t, backend.Close()) }) + connection, err := net.DialTimeout("tcp", backend.Address(), time.Second) + requireNoFixtureError(t, err) + defer connection.Close() + requireNoFixtureError(t, connection.SetDeadline(time.Now().Add(2*time.Second))) + + // When + _, err = io.WriteString(connection, "GET /response-half-close HTTP/1.1\r\nHost: nat.invalid\r\nContent-Length: 0\r\n\r\n") + requireNoFixtureError(t, err) + response, err := http.ReadResponse(bufio.NewReader(connection), nil) + requireNoFixtureError(t, err) + _, err = io.ReadAll(response.Body) + requireNoFixtureError(t, err) + requireNoFixtureError(t, response.Body.Close()) + ctx, cancel := context.WithTimeout(context.Background(), time.Second) + defer cancel() + record, err := backend.WaitRequest(ctx) + requireNoFixtureError(t, err) + + // Then + if !record.ResponseHalfClosed || record.RequestHalfClosed { + t.Fatalf("NAT response half-close record = %+v", record) + } +} + func TestFixture_NATEchoCloseInterruptsIncompleteRequest(t *testing.T) { // Given backend, err := StartNATEchoBackend() @@ -132,3 +163,71 @@ func TestFixture_NATEchoCloseTerminatesSockets(t *testing.T) { t.Fatal("NAT listener accepted a connection after backend close") } } + +func TestFixture_NATEchoCloseIsConcurrentAndIdempotent(t *testing.T) { + // Given + listener, err := net.Listen("tcp", "127.0.0.1:0") + requireNoFixtureError(t, err) + closeFailure := errors.New("injected listener close failure") + backend := newNATEchoBackend(closeErrorListener{Listener: listener, err: closeFailure}, false, false) + connection, err := net.DialTimeout("tcp", backend.Address(), time.Second) + requireNoFixtureError(t, err) + defer connection.Close() + _, err = io.WriteString(connection, "GET /incomplete HTTP/1.1\r\nHost: nat.invalid\r\n") + requireNoFixtureError(t, err) + ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second) + defer cancel() + requireNoFixtureError(t, backend.WaitConnection(ctx)) + + // When + const callers = 16 + start := make(chan struct{}) + closeErrors := make(chan error, callers) + var waitGroup sync.WaitGroup + waitGroup.Add(callers) + for range callers { + go func() { + defer waitGroup.Done() + <-start + closeErrors <- backend.Close() + }() + } + close(start) + waitGroup.Wait() + close(closeErrors) + + // Then + for closeErr := range closeErrors { + if !errors.Is(closeErr, closeFailure) { + t.Fatalf("concurrent close error = %v, want %v", closeErr, closeFailure) + } + } + if closeErr := backend.Close(); !errors.Is(closeErr, closeFailure) { + t.Fatalf("repeated close error = %v, want %v", closeErr, closeFailure) + } +} + +type closeErrorListener struct { + net.Listener + err error +} + +func (listener closeErrorListener) Close() error { + _ = listener.Listener.Close() + return listener.err +} + +func (listener closeErrorListener) Accept() (net.Conn, error) { + connection, err := listener.Listener.Accept() + if err != nil { + return nil, err + } + return alreadyClosedErrorConn{Conn: connection}, nil +} + +type alreadyClosedErrorConn struct{ net.Conn } + +func (connection alreadyClosedErrorConn) Close() error { + _ = connection.Conn.Close() + return net.ErrClosed +} diff --git a/integration/agentcompat/internal/fixture/nat_hold.go b/integration/agentcompat/internal/fixture/nat_hold.go new file mode 100644 index 00000000..46a4bc0b --- /dev/null +++ b/integration/agentcompat/internal/fixture/nat_hold.go @@ -0,0 +1,147 @@ +package fixture + +import ( + "bufio" + "context" + "errors" + "fmt" + "io" + "net" + "net/http" + "sync" + + "github.com/nezhahq/nezha/pkg/agentcompatcontract" +) + +type NATHoldBackend struct { + listener net.Listener + results chan natHoldResult + done chan struct{} + released chan struct{} + observed chan struct{} + closeOnce sync.Once + releaseOnce sync.Once + closeErr error + waitGroup sync.WaitGroup + connectionMu sync.Mutex + connection net.Conn +} + +type natHoldResult struct { + record NATEchoRecord + err error +} + +func StartNATHoldBackend() (*NATHoldBackend, error) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + return nil, fmt.Errorf("listen for NAT hold: %w", err) + } + backend := &NATHoldBackend{listener: listener, results: make(chan natHoldResult, 1), done: make(chan struct{}), released: make(chan struct{}), observed: make(chan struct{})} + backend.waitGroup.Add(2) + go backend.accept() + return backend, nil +} + +func (backend *NATHoldBackend) Address() string { return backend.listener.Addr().String() } + +func (backend *NATHoldBackend) RequestObserved() <-chan struct{} { return backend.observed } + +func (backend *NATHoldBackend) ResponseReleased() <-chan struct{} { return backend.released } + +func (backend *NATHoldBackend) WaitRequest(ctx context.Context) (NATEchoRecord, error) { + select { + case result := <-backend.results: + return result.record, result.err + case <-ctx.Done(): + return NATEchoRecord{}, ctx.Err() + case <-backend.done: + return NATEchoRecord{}, errors.New("NAT hold backend closed") + } +} + +func (backend *NATHoldBackend) Release() error { + backend.releaseOnce.Do(func() { close(backend.released) }) + return nil +} + +func (backend *NATHoldBackend) Close() error { + backend.closeOnce.Do(func() { + close(backend.done) + backend.closeErr = normalizeNATHoldCloseError(backend.listener.Close()) + backend.connectionMu.Lock() + connection := backend.connection + backend.connectionMu.Unlock() + if connection != nil { + backend.closeErr = errors.Join(backend.closeErr, normalizeNATHoldCloseError(connection.Close())) + } + backend.releaseOnce.Do(func() { close(backend.released) }) + backend.waitGroup.Wait() + }) + return backend.closeErr +} + +func normalizeNATHoldCloseError(err error) error { + if errors.Is(err, net.ErrClosed) { + return nil + } + return err +} + +func (backend *NATHoldBackend) accept() { + defer backend.waitGroup.Done() + connection, err := backend.listener.Accept() + if err != nil { + backend.waitGroup.Done() + if !errors.Is(err, net.ErrClosed) { + backend.publish(natHoldResult{err: fmt.Errorf("accept NAT hold connection: %w", err)}) + } + return + } + backend.connectionMu.Lock() + backend.connection = connection + backend.connectionMu.Unlock() + select { + case <-backend.done: + _ = connection.Close() + default: + } + go backend.handle(connection) +} + +func (backend *NATHoldBackend) handle(connection net.Conn) { + defer backend.waitGroup.Done() + defer func() { _ = connection.Close() }() + request, err := http.ReadRequest(bufio.NewReader(connection)) + if err != nil { + backend.publish(natHoldResult{err: fmt.Errorf("read NAT hold request: %w", err)}) + return + } + body, err := io.ReadAll(request.Body) + if closeErr := request.Body.Close(); err == nil { + err = closeErr + } + if err != nil { + backend.publish(natHoldResult{err: fmt.Errorf("read NAT hold body: %w", err)}) + return + } + record := NATEchoRecord{Method: request.Method, Path: request.URL.RequestURI(), Host: request.Host, HeaderValue: request.Header.Get("X-AgentCompat-Echo"), Body: append([]byte(nil), body...), SensitiveHeadersPresent: request.Header.Get(agentcompatcontract.IOStreamCapabilityHeader) != "" || request.Header.Get("Authorization") != ""} + backend.publish(natHoldResult{record: record}) + close(backend.observed) + select { + case <-backend.released: + responseBody := fmt.Sprintf("method=%s\npath=%s\nhost=%s\nx-agentcompat-echo=%s\nbody=%s\n", record.Method, record.Path, record.Host, record.HeaderValue, record.Body) + response := fmt.Sprintf("HTTP/1.1 200 OK\r\nContent-Type: text/plain\r\nContent-Length: %d\r\nConnection: close\r\n\r\n%s", len(responseBody), responseBody) + if _, err := io.WriteString(connection, response); err != nil { + backend.publish(natHoldResult{err: fmt.Errorf("write NAT hold response: %w", err)}) + } + case <-backend.done: + } +} + +func (backend *NATHoldBackend) publish(result natHoldResult) { + select { + case backend.results <- result: + case <-backend.done: + } +} diff --git a/integration/agentcompat/internal/fixture/nat_hold_test.go b/integration/agentcompat/internal/fixture/nat_hold_test.go new file mode 100644 index 00000000..cd98332b --- /dev/null +++ b/integration/agentcompat/internal/fixture/nat_hold_test.go @@ -0,0 +1,126 @@ +package fixture + +import ( + "bufio" + "context" + "errors" + "io" + "net" + "net/http" + "sync" + "testing" +) + +func TestNATHoldBackendObservesBeforeReleaseAndRespondsAfterRelease(t *testing.T) { + backend, err := StartNATHoldBackend() + if err != nil { + t.Fatal(err) + } + defer func() { + if err := backend.Close(); err != nil { + t.Fatal(err) + } + }() + connection, err := net.Dial("tcp", backend.Address()) + if err != nil { + t.Fatal(err) + } + defer connection.Close() + _, err = io.WriteString(connection, "PATCH /hold HTTP/1.1\r\nHost: hold.invalid\r\nX-AgentCompat-Echo: exact\r\nContent-Length: 4\r\n\r\nbody") + if err != nil { + t.Fatal(err) + } + record, err := backend.WaitRequest(context.Background()) + if err != nil || record.Method != "PATCH" || record.Path != "/hold" || record.Host != "hold.invalid" || record.HeaderValue != "exact" || string(record.Body) != "body" { + t.Fatalf("record=%+v err=%v", record, err) + } + select { + case <-backend.ResponseReleased(): + t.Fatal("response released before Release") + default: + } + if err := backend.Release(); err != nil { + t.Fatal(err) + } + response, err := httpReadResponse(connection) + if err != nil || response != "method=PATCH\npath=/hold\nhost=hold.invalid\nx-agentcompat-echo=exact\nbody=body\n" { + t.Fatalf("response=%q err=%v", response, err) + } +} + +func TestNATHoldBackendCloseInterruptsIncompleteRequest(t *testing.T) { + backend, err := StartNATHoldBackend() + if err != nil { + t.Fatal(err) + } + connection, err := net.Dial("tcp", backend.Address()) + if err != nil { + t.Fatal(err) + } + if _, err := io.WriteString(connection, "GET /incomplete HTTP/1.1\r\nHost: hold.invalid\r\n"); err != nil { + t.Fatal(err) + } + if err := backend.Close(); err != nil { + t.Fatal(err) + } + if err := connection.Close(); err != nil { + t.Fatal(err) + } +} + +func TestNATHoldBackendCloseIsConcurrentAndRetainsListenerError(t *testing.T) { + listener, err := net.Listen("tcp", "127.0.0.1:0") + if err != nil { + t.Fatal(err) + } + closeFailure := errors.New("hold listener close failed") + backend := newNATHoldBackend(closeErrorListener{Listener: listener, err: closeFailure}) + connection, err := net.Dial("tcp", backend.Address()) + if err != nil { + t.Fatal(err) + } + defer connection.Close() + if _, err := io.WriteString(connection, "GET /hold HTTP/1.1\r\nHost: hold.invalid\r\n"); err != nil { + t.Fatal(err) + } + + const callers = 8 + errorsSeen := make(chan error, callers) + var group sync.WaitGroup + group.Add(callers) + for range callers { + go func() { + defer group.Done() + errorsSeen <- backend.Close() + }() + } + group.Wait() + + for range callers { + if err := <-errorsSeen; !errors.Is(err, closeFailure) { + t.Fatalf("Close error=%v, want %v", err, closeFailure) + } + } + if err := backend.Close(); !errors.Is(err, closeFailure) { + t.Fatalf("repeated Close error=%v, want %v", err, closeFailure) + } +} + +func newNATHoldBackend(listener net.Listener) *NATHoldBackend { + backend := &NATHoldBackend{listener: listener, results: make(chan natHoldResult, 1), done: make(chan struct{}), released: make(chan struct{}), observed: make(chan struct{})} + backend.waitGroup.Add(2) + go backend.accept() + return backend +} + +func httpReadResponse(connection net.Conn) (string, error) { + response, err := http.ReadResponse(bufio.NewReader(connection), nil) + if err != nil { + return "", err + } + body, err := io.ReadAll(response.Body) + if closeErr := response.Body.Close(); err == nil { + err = closeErr + } + return string(body), err +} diff --git a/integration/agentcompat/internal/fixture/payload_heap_test.go b/integration/agentcompat/internal/fixture/payload_heap_test.go new file mode 100644 index 00000000..d70c0632 --- /dev/null +++ b/integration/agentcompat/internal/fixture/payload_heap_test.go @@ -0,0 +1,65 @@ +package fixture + +import ( + "bytes" + "io" + "runtime" + "testing" + + "github.com/nezhahq/nezha/integration/agentcompat/internal/contract" +) + +type retainedHeapReader struct { + reader io.Reader + baseline uint64 + peak uint64 +} + +func verifyPayloadPeakRetainedHeap(reader io.Reader, expectedBytes uint64) (PayloadDigest, uint64, error) { + runtime.GC() + var baseline runtime.MemStats + runtime.ReadMemStats(&baseline) + measuredReader := &retainedHeapReader{reader: reader, baseline: baseline.HeapAlloc} + digest, err := VerifyPayload(measuredReader, expectedBytes) + return digest, measuredReader.peak, err +} + +func (reader *retainedHeapReader) Read(destination []byte) (int, error) { + readBytes, err := reader.reader.Read(destination) + // Sampling after a forced collection measures the live heap retained by each + // streaming checkpoint instead of scheduler-dependent allocation churn. + runtime.GC() + var sample runtime.MemStats + runtime.ReadMemStats(&sample) + runtime.KeepAlive(destination) + if sample.HeapAlloc > reader.baseline { + reader.peak = max(reader.peak, sample.HeapAlloc-reader.baseline) + } + return readBytes, err +} + +func TestFixture_PeakRetainedHeapDetectsRetainedAllocation(t *testing.T) { + // Given + reader := &allocationRetainingReader{reader: bytes.NewReader([]byte("x"))} + + // When + _, peakRetainedHeap, err := verifyPayloadPeakRetainedHeap(reader, 1) + requireNoFixtureError(t, err) + + // Then + if peakRetainedHeap <= contract.TransferHeapBytes { + t.Fatalf("peak retained heap = %d, want greater than %d", peakRetainedHeap, contract.TransferHeapBytes) + } +} + +type allocationRetainingReader struct { + reader io.Reader + retained []byte +} + +func (reader *allocationRetainingReader) Read(destination []byte) (int, error) { + if reader.retained == nil { + reader.retained = make([]byte, contract.TransferHeapBytes+1024*1024) + } + return reader.reader.Read(destination) +} diff --git a/integration/agentcompat/internal/fixture/payload_measurement.go b/integration/agentcompat/internal/fixture/payload_measurement.go new file mode 100644 index 00000000..37225d7a --- /dev/null +++ b/integration/agentcompat/internal/fixture/payload_measurement.go @@ -0,0 +1,89 @@ +package fixture + +import ( + "crypto/sha256" + "hash" + "io" + "runtime" +) + +type PayloadMeasurement struct { + Digest PayloadDigest + RetainedHeapBytes uint64 + Chunks uint64 +} + +type RetainedHeapProbe struct { + baseline uint64 +} + +func NewRetainedHeapProbe() RetainedHeapProbe { + runtime.GC() + var sample runtime.MemStats + runtime.ReadMemStats(&sample) + return RetainedHeapProbe{baseline: sample.HeapAlloc} +} + +func (probe RetainedHeapProbe) RetainedBytes() uint64 { + runtime.GC() + var sample runtime.MemStats + runtime.ReadMemStats(&sample) + if sample.HeapAlloc <= probe.baseline { + return 0 + } + return sample.HeapAlloc - probe.baseline +} + +type MeasuredReader struct { + reader io.Reader + hash hash.Hash + bytes uint64 + chunks uint64 +} + +func NewMeasuredReader(reader io.Reader) *MeasuredReader { + return &MeasuredReader{reader: reader, hash: sha256.New()} +} + +func (reader *MeasuredReader) Read(destination []byte) (int, error) { + readBytes, err := reader.reader.Read(destination) + if readBytes > 0 { + _, _ = reader.hash.Write(destination[:readBytes]) + reader.bytes += uint64(readBytes) + reader.chunks++ + } + return readBytes, err +} + +func (reader *MeasuredReader) Measurement() PayloadMeasurement { + return newPayloadMeasurement(reader.hash, reader.bytes, reader.chunks) +} + +type MeasuredWriter struct { + hash hash.Hash + bytes uint64 + chunks uint64 +} + +func NewMeasuredWriter() *MeasuredWriter { + return &MeasuredWriter{hash: sha256.New()} +} + +func (writer *MeasuredWriter) Write(payload []byte) (int, error) { + written, err := writer.hash.Write(payload) + if written > 0 { + writer.bytes += uint64(written) + writer.chunks++ + } + return written, err +} + +func (writer *MeasuredWriter) Measurement() PayloadMeasurement { + return newPayloadMeasurement(writer.hash, writer.bytes, writer.chunks) +} + +func newPayloadMeasurement(payloadHash hash.Hash, bytes, chunks uint64) PayloadMeasurement { + digest := PayloadDigest{Bytes: bytes} + copy(digest.SHA256[:], payloadHash.Sum(nil)) + return PayloadMeasurement{Digest: digest, Chunks: chunks} +} diff --git a/integration/agentcompat/internal/fixture/payload_test.go b/integration/agentcompat/internal/fixture/payload_test.go index ab940032..7c8503f0 100644 --- a/integration/agentcompat/internal/fixture/payload_test.go +++ b/integration/agentcompat/internal/fixture/payload_test.go @@ -7,7 +7,6 @@ import ( "errors" "fmt" "io" - "runtime" "testing" "github.com/nezhahq/nezha/integration/agentcompat/internal/contract" @@ -17,19 +16,13 @@ func TestFixture_StreamsExact100MiB(t *testing.T) { // Given payload, err := NewPayload(contract.DefaultSeed, contract.TransferBytes) requireNoFixtureError(t, err) - runtime.GC() - var before runtime.MemStats - runtime.ReadMemStats(&before) // When - digest, err := VerifyPayload(payload.Reader(), contract.TransferBytes) + digest, peakRetainedHeap, err := verifyPayloadPeakRetainedHeap(payload.Reader(), contract.TransferBytes) requireNoFixtureError(t, err) stableDigest, err := VerifyPayload(payload.Reader(), contract.TransferBytes) requireNoFixtureError(t, err) independentDigest := independentlyHashPayload(contract.DefaultSeed, contract.TransferBytes) - runtime.GC() - var after runtime.MemStats - runtime.ReadMemStats(&after) // Then if digest.Bytes != contract.TransferBytes { @@ -41,14 +34,13 @@ func TestFixture_StreamsExact100MiB(t *testing.T) { if digest.SHA256 != independentDigest { t.Fatalf("payload SHA does not match independent generator: %s", digest.Hex()) } - var retained uint64 - if after.HeapAlloc > before.HeapAlloc { - retained = after.HeapAlloc - before.HeapAlloc + if peakRetainedHeap == 0 { + t.Fatal("peak retained heap measurement did not observe live streaming allocations") } - if retained > contract.TransferHeapBytes { - t.Fatalf("retained heap = %d, limit = %d", retained, contract.TransferHeapBytes) + if peakRetainedHeap > contract.TransferHeapBytes { + t.Fatalf("peak retained heap = %d, limit = %d", peakRetainedHeap, contract.TransferHeapBytes) } - t.Logf("bytes=%d sha256=%s retained_heap=%d", digest.Bytes, digest.Hex(), retained) + t.Logf("bytes=%d sha256=%s peak_retained_heap=%d", digest.Bytes, digest.Hex(), peakRetainedHeap) } func independentlyHashPayload(seed contract.Seed, size uint64) [sha256.Size]byte { diff --git a/integration/agentcompat/internal/testpaths/source.go b/integration/agentcompat/internal/testpaths/source.go new file mode 100644 index 00000000..21810162 --- /dev/null +++ b/integration/agentcompat/internal/testpaths/source.go @@ -0,0 +1,62 @@ +package testpaths + +import ( + "errors" + "os" + "path/filepath" +) + +func NezhaSource(start string) (string, error) { + if configured := os.Getenv("NEZHA_SOURCE"); configured != "" { + return absoluteDirectory(configured) + } + return findModuleRoot(start) +} + +func AgentSource(nezhaSource string) (string, error) { + if configured := os.Getenv("AGENT_SOURCE"); configured != "" { + return absoluteDirectory(configured) + } + root, err := absoluteDirectory(nezhaSource) + if err != nil { + return "", err + } + return absoluteDirectory(filepath.Join(filepath.Dir(root), "agent")) +} + +func absoluteDirectory(raw string) (string, error) { + if raw == "" || !filepath.IsAbs(raw) { + return "", errors.New("source path must be absolute") + } + clean := filepath.Clean(raw) + info, err := os.Stat(clean) + if err != nil { + return "", err + } + if !info.IsDir() { + return "", errors.New("source path must be a directory") + } + return clean, nil +} + +func findModuleRoot(start string) (string, error) { + if start == "" { + start, _ = os.Getwd() + } + absolute, err := filepath.Abs(start) + if err != nil { + return "", err + } + if info, statErr := os.Stat(absolute); statErr == nil && !info.IsDir() { + absolute = filepath.Dir(absolute) + } + for directory := absolute; ; directory = filepath.Dir(directory) { + if _, err := os.Stat(filepath.Join(directory, "go.mod")); err == nil { + return directory, nil + } + parent := filepath.Dir(directory) + if parent == directory { + return "", errors.New("module root not found") + } + } +} diff --git a/integration/agentcompat/internal/testpaths/source_test.go b/integration/agentcompat/internal/testpaths/source_test.go new file mode 100644 index 00000000..7a693f1e --- /dev/null +++ b/integration/agentcompat/internal/testpaths/source_test.go @@ -0,0 +1,61 @@ +package testpaths + +import ( + "os" + "path/filepath" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestNezhaSource_UsesModuleRootFromNonRepositoryCWD(t *testing.T) { + t.Setenv("NEZHA_SOURCE", "") + root := t.TempDir() + require.NoError(t, os.WriteFile(filepath.Join(root, "go.mod"), []byte("module example.test\n"), 0o600)) + nonRepository := t.TempDir() + original, err := os.Getwd() + require.NoError(t, err) + require.NoError(t, os.Chdir(nonRepository)) + t.Cleanup(func() { require.NoError(t, os.Chdir(original)) }) + + resolved, err := NezhaSource(root) + + require.NoError(t, err) + require.Equal(t, root, resolved) +} + +func TestSourceResolversHonorExplicitEnvironmentPaths(t *testing.T) { + // Given + nezha := t.TempDir() + agent := t.TempDir() + t.Setenv("NEZHA_SOURCE", nezha) + t.Setenv("AGENT_SOURCE", agent) + + // When + resolvedNezha, nezhaErr := NezhaSource(t.TempDir()) + resolvedAgent, agentErr := AgentSource(nezha) + + // Then + require.NoError(t, nezhaErr) + require.NoError(t, agentErr) + require.Equal(t, nezha, resolvedNezha) + require.Equal(t, agent, resolvedAgent) +} + +func TestAgentSource_ResolvesAdjacentCheckoutThroughSymlink(t *testing.T) { + t.Setenv("AGENT_SOURCE", "") + parent := t.TempDir() + nezha := filepath.Join(parent, "nezha") + agent := filepath.Join(parent, "agent") + require.NoError(t, os.MkdirAll(nezha, 0o700)) + require.NoError(t, os.MkdirAll(agent, 0o700)) + require.NoError(t, os.WriteFile(filepath.Join(nezha, "go.mod"), []byte("module nezha.test\n"), 0o600)) + require.NoError(t, os.WriteFile(filepath.Join(agent, "go.mod"), []byte("module agent.test\n"), 0o600)) + linkParent := filepath.Join(t.TempDir(), "checkout") + require.NoError(t, os.Symlink(parent, linkParent)) + + resolved, err := AgentSource(filepath.Join(linkParent, "nezha")) + + require.NoError(t, err) + require.Equal(t, filepath.Join(linkParent, "agent"), resolved) +}