mirror of
https://github.com/Buriburizaem0n/nezha_domains.git
synced 2026-09-19 09:40:12 +00:00
test(agentcompat): harden deterministic fixtures
Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
@@ -9,6 +9,8 @@ import (
|
||||
"net"
|
||||
"net/http"
|
||||
"sync"
|
||||
|
||||
"github.com/nezhahq/nezha/pkg/agentcompatcontract"
|
||||
)
|
||||
|
||||
type NATEchoRecord struct {
|
||||
@@ -18,6 +20,8 @@ type NATEchoRecord struct {
|
||||
HeaderValue string
|
||||
Body []byte
|
||||
RequestHalfClosed bool
|
||||
ResponseHalfClosed bool
|
||||
SensitiveHeadersPresent bool
|
||||
}
|
||||
|
||||
type natEchoResult struct {
|
||||
@@ -30,8 +34,10 @@ type NATEchoBackend struct {
|
||||
results chan natEchoResult
|
||||
done chan struct{}
|
||||
connectionsReady chan struct{}
|
||||
requireHalfClose bool
|
||||
requireRequestHalfClose bool
|
||||
halfCloseResponse bool
|
||||
closeOnce sync.Once
|
||||
closeErr error
|
||||
waitGroup sync.WaitGroup
|
||||
mutex sync.Mutex
|
||||
connections map[net.Conn]struct{}
|
||||
@@ -39,29 +45,38 @@ type NATEchoBackend struct {
|
||||
}
|
||||
|
||||
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,
|
||||
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 {
|
||||
@@ -179,6 +199,7 @@ func (backend *NATEchoBackend) handle(connection net.Conn) {
|
||||
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})
|
||||
}
|
||||
|
||||
@@ -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
|
||||
}
|
||||
|
||||
@@ -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"
|
||||
"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 {
|
||||
|
||||
@@ -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)
|
||||
}
|
||||
Reference in New Issue
Block a user