test(agentcompat): harden deterministic fixtures

Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
naiba
2026-07-20 04:42:03 +00:00
co-authored by naiba/CloudCode
parent fcf58b0b4b
commit 0543995ec3
9 changed files with 717 additions and 54 deletions
@@ -9,15 +9,19 @@ import (
"net" "net"
"net/http" "net/http"
"sync" "sync"
"github.com/nezhahq/nezha/pkg/agentcompatcontract"
) )
type NATEchoRecord struct { type NATEchoRecord struct {
Method string Method string
Path string Path string
Host string Host string
HeaderValue string HeaderValue string
Body []byte Body []byte
RequestHalfClosed bool RequestHalfClosed bool
ResponseHalfClosed bool
SensitiveHeadersPresent bool
} }
type natEchoResult struct { type natEchoResult struct {
@@ -26,42 +30,53 @@ type natEchoResult struct {
} }
type NATEchoBackend struct { type NATEchoBackend struct {
listener net.Listener listener net.Listener
results chan natEchoResult results chan natEchoResult
done chan struct{} done chan struct{}
connectionsReady chan struct{} connectionsReady chan struct{}
requireHalfClose bool requireRequestHalfClose bool
closeOnce sync.Once halfCloseResponse bool
waitGroup sync.WaitGroup closeOnce sync.Once
mutex sync.Mutex closeErr error
connections map[net.Conn]struct{} waitGroup sync.WaitGroup
closing bool mutex sync.Mutex
connections map[net.Conn]struct{}
closing bool
} }
func StartNATEchoBackend() (*NATEchoBackend, error) { func StartNATEchoBackend() (*NATEchoBackend, error) {
return startNATEchoBackend(false) return startNATEchoBackend(false, false)
} }
func StartNATHalfCloseEchoBackend() (*NATEchoBackend, error) { 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") listener, err := net.Listen("tcp", "127.0.0.1:0")
if err != nil { if err != nil {
return nil, fmt.Errorf("listen for NAT echo: %w", err) 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{ backend := &NATEchoBackend{
listener: listener, listener: listener,
results: make(chan natEchoResult, 16), results: make(chan natEchoResult, 16),
done: make(chan struct{}), done: make(chan struct{}),
connectionsReady: make(chan struct{}, 16), connectionsReady: make(chan struct{}, 16),
requireHalfClose: requireHalfClose, requireRequestHalfClose: requireRequestHalfClose,
connections: make(map[net.Conn]struct{}), halfCloseResponse: halfCloseResponse,
connections: make(map[net.Conn]struct{}),
} }
backend.waitGroup.Add(1) backend.waitGroup.Add(1)
go backend.accept() go backend.accept()
return backend, nil return backend
} }
func (backend *NATEchoBackend) Address() string { func (backend *NATEchoBackend) Address() string {
@@ -91,10 +106,9 @@ func (backend *NATEchoBackend) WaitConnection(ctx context.Context) error {
} }
func (backend *NATEchoBackend) Close() error { func (backend *NATEchoBackend) Close() error {
var closeErr error
backend.closeOnce.Do(func() { backend.closeOnce.Do(func() {
close(backend.done) close(backend.done)
closeErr = backend.listener.Close() backend.closeErr = normalizeNATEchoCloseError(backend.listener.Close())
backend.mutex.Lock() backend.mutex.Lock()
backend.closing = true backend.closing = true
connections := make([]net.Conn, 0, len(backend.connections)) connections := make([]net.Conn, 0, len(backend.connections))
@@ -103,14 +117,20 @@ func (backend *NATEchoBackend) Close() error {
} }
backend.mutex.Unlock() backend.mutex.Unlock()
for _, connection := range connections { for _, connection := range connections {
closeErr = errors.Join(closeErr, connection.Close()) backend.closeErr = errors.Join(backend.closeErr, normalizeNATEchoCloseError(connection.Close()))
} }
backend.waitGroup.Wait() 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 nil
} }
return closeErr return err
} }
func (backend *NATEchoBackend) accept() { func (backend *NATEchoBackend) accept() {
@@ -164,7 +184,7 @@ func (backend *NATEchoBackend) handle(connection net.Conn) {
return return
} }
halfClosed := false halfClosed := false
if backend.requireHalfClose { if backend.requireRequestHalfClose {
_, halfCloseErr := reader.ReadByte() _, halfCloseErr := reader.ReadByte()
halfClosed = errors.Is(halfCloseErr, io.EOF) halfClosed = errors.Is(halfCloseErr, io.EOF)
if halfCloseErr != nil && !halfClosed { if halfCloseErr != nil && !halfClosed {
@@ -173,12 +193,13 @@ func (backend *NATEchoBackend) handle(connection net.Conn) {
} }
} }
record := NATEchoRecord{ record := NATEchoRecord{
Method: request.Method, Method: request.Method,
Path: request.URL.RequestURI(), Path: request.URL.RequestURI(),
Host: request.Host, Host: request.Host,
HeaderValue: request.Header.Get("X-AgentCompat-Echo"), HeaderValue: request.Header.Get("X-AgentCompat-Echo"),
Body: append([]byte(nil), body...), Body: append([]byte(nil), body...),
RequestHalfClosed: halfClosed, 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) 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) 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)}) backend.publish(natEchoResult{err: fmt.Errorf("write NAT echo response: %w", err)})
return return
} }
if tcpConnection, ok := connection.(*net.TCPConn); ok { if tcpConnection, ok := connection.(*net.TCPConn); ok && backend.halfCloseResponse {
if err := tcpConnection.CloseWrite(); err != nil { if err := tcpConnection.CloseWrite(); err != nil {
backend.publish(natEchoResult{err: fmt.Errorf("half-close NAT echo response: %w", err)}) backend.publish(natEchoResult{err: fmt.Errorf("half-close NAT echo response: %w", err)})
return return
} }
record.ResponseHalfClosed = true
} }
backend.publish(natEchoResult{record: record}) backend.publish(natEchoResult{record: record})
} }
@@ -3,11 +3,13 @@ package fixture
import ( import (
"bufio" "bufio"
"context" "context"
"errors"
"fmt" "fmt"
"io" "io"
"net" "net"
"net/http" "net/http"
"strings" "strings"
"sync"
"testing" "testing"
"time" "time"
) )
@@ -78,11 +80,40 @@ func TestFixture_NATEcho(t *testing.T) {
if response.StatusCode != http.StatusOK || string(responseBody) != expected { if response.StatusCode != http.StatusOK || string(responseBody) != expected {
t.Fatalf("NAT echo response = %d %q", response.StatusCode, responseBody) 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) 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) { func TestFixture_NATEchoCloseInterruptsIncompleteRequest(t *testing.T) {
// Given // Given
backend, err := StartNATEchoBackend() backend, err := StartNATEchoBackend()
@@ -132,3 +163,71 @@ func TestFixture_NATEchoCloseTerminatesSockets(t *testing.T) {
t.Fatal("NAT listener accepted a connection after backend close") 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
}
@@ -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:
}
}
@@ -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
}
@@ -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)
}
@@ -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}
}
@@ -7,7 +7,6 @@ import (
"errors" "errors"
"fmt" "fmt"
"io" "io"
"runtime"
"testing" "testing"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract" "github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
@@ -17,19 +16,13 @@ func TestFixture_StreamsExact100MiB(t *testing.T) {
// Given // Given
payload, err := NewPayload(contract.DefaultSeed, contract.TransferBytes) payload, err := NewPayload(contract.DefaultSeed, contract.TransferBytes)
requireNoFixtureError(t, err) requireNoFixtureError(t, err)
runtime.GC()
var before runtime.MemStats
runtime.ReadMemStats(&before)
// When // When
digest, err := VerifyPayload(payload.Reader(), contract.TransferBytes) digest, peakRetainedHeap, err := verifyPayloadPeakRetainedHeap(payload.Reader(), contract.TransferBytes)
requireNoFixtureError(t, err) requireNoFixtureError(t, err)
stableDigest, err := VerifyPayload(payload.Reader(), contract.TransferBytes) stableDigest, err := VerifyPayload(payload.Reader(), contract.TransferBytes)
requireNoFixtureError(t, err) requireNoFixtureError(t, err)
independentDigest := independentlyHashPayload(contract.DefaultSeed, contract.TransferBytes) independentDigest := independentlyHashPayload(contract.DefaultSeed, contract.TransferBytes)
runtime.GC()
var after runtime.MemStats
runtime.ReadMemStats(&after)
// Then // Then
if digest.Bytes != contract.TransferBytes { if digest.Bytes != contract.TransferBytes {
@@ -41,14 +34,13 @@ func TestFixture_StreamsExact100MiB(t *testing.T) {
if digest.SHA256 != independentDigest { if digest.SHA256 != independentDigest {
t.Fatalf("payload SHA does not match independent generator: %s", digest.Hex()) t.Fatalf("payload SHA does not match independent generator: %s", digest.Hex())
} }
var retained uint64 if peakRetainedHeap == 0 {
if after.HeapAlloc > before.HeapAlloc { t.Fatal("peak retained heap measurement did not observe live streaming allocations")
retained = after.HeapAlloc - before.HeapAlloc
} }
if retained > contract.TransferHeapBytes { if peakRetainedHeap > contract.TransferHeapBytes {
t.Fatalf("retained heap = %d, limit = %d", retained, 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 { func independentlyHashPayload(seed contract.Seed, size uint64) [sha256.Size]byte {
@@ -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")
}
}
}
@@ -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)
}