test(compat): add deterministic local fixtures

Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
naiba
2026-07-14 05:49:15 +00:00
co-authored by naiba/CloudCode
parent b01ff45ef4
commit 92a1cc76be
10 changed files with 1178 additions and 0 deletions
@@ -0,0 +1,156 @@
package fixture
import (
"errors"
"fmt"
"os"
"path/filepath"
"regexp"
"strings"
)
var agentRootNamePattern = regexp.MustCompile(`^[a-z0-9]+(?:-[a-z0-9]+)*$`)
type AgentRoot struct {
absolute string
}
type AgentPath struct {
absolute string
relative string
}
func NewAgentRoot(parent, agentID string) (AgentRoot, error) {
cleanParent := filepath.Clean(parent)
if !filepath.IsAbs(cleanParent) {
return AgentRoot{}, errors.New("agent fixture parent must be absolute")
}
if !agentRootNamePattern.MatchString(agentID) {
return AgentRoot{}, errors.New("invalid agent fixture root name")
}
parentInfo, err := os.Lstat(cleanParent)
if err != nil {
return AgentRoot{}, fmt.Errorf("inspect agent fixture parent: %w", err)
}
if !parentInfo.IsDir() || parentInfo.Mode()&os.ModeSymlink != 0 {
return AgentRoot{}, errors.New("agent fixture parent must be a real directory")
}
absolute := filepath.Join(cleanParent, agentID)
if err := os.Mkdir(absolute, 0o700); err != nil {
return AgentRoot{}, fmt.Errorf("create agent fixture root: %w", err)
}
return AgentRoot{absolute: absolute}, nil
}
func (root AgentRoot) Absolute() string {
return root.absolute
}
func (root AgentRoot) Path(relative string) (AgentPath, error) {
return root.newPath(relative, false)
}
func (root AgentRoot) DestructivePath(relative string) (AgentPath, error) {
return root.newPath(relative, true)
}
func (root AgentRoot) newPath(relative string, destructive bool) (AgentPath, error) {
nativeRelative, err := validateRelativeAgentPath(relative, destructive)
if err != nil {
return AgentPath{}, err
}
absolute := filepath.Clean(filepath.Join(root.absolute, nativeRelative))
containedRelative, err := filepath.Rel(root.absolute, absolute)
if err != nil || filepath.IsAbs(containedRelative) || containedRelative == ".." || strings.HasPrefix(containedRelative, ".."+string(filepath.Separator)) {
return AgentPath{}, rejectPath(PathRejectionEscape)
}
if destructive && containedRelative == "." {
return AgentPath{}, rejectPath(PathRejectionDestructiveRoot)
}
if err := ensureRealParentDirectories(root.absolute, nativeRelative); err != nil {
return AgentPath{}, err
}
return AgentPath{absolute: absolute, relative: nativeRelative}, nil
}
func (path AgentPath) String() string {
return path.absolute
}
func (path AgentPath) Relative() string {
return path.relative
}
func validateRelativeAgentPath(candidate string, destructive bool) (string, error) {
if strings.TrimSpace(candidate) == "" {
return "", rejectPath(PathRejectionEmpty)
}
if hasWindowsVolume(candidate) {
return "", rejectPath(PathRejectionVolume)
}
if filepath.IsAbs(candidate) {
return "", rejectPath(PathRejectionAbsolute)
}
if strings.Contains(candidate, `\`) {
return "", rejectPath(PathRejectionSeparator)
}
if strings.Contains(candidate, ":") {
return "", rejectPath(PathRejectionADS)
}
components := strings.Split(candidate, "/")
for _, component := range components {
if component == "" {
return "", rejectPath(PathRejectionSeparator)
}
if component == ".." {
return "", rejectPath(PathRejectionParent)
}
}
nativeRelative := filepath.FromSlash(candidate)
if destructive && filepath.Clean(nativeRelative) == "." {
return "", rejectPath(PathRejectionDestructiveRoot)
}
return nativeRelative, nil
}
func hasWindowsVolume(candidate string) bool {
if strings.HasPrefix(candidate, `\\`) || strings.HasPrefix(candidate, `//`) {
return true
}
return len(candidate) >= 2 && ((candidate[0] >= 'A' && candidate[0] <= 'Z') || (candidate[0] >= 'a' && candidate[0] <= 'z')) && candidate[1] == ':'
}
func ensureRealParentDirectories(root, relative string) error {
parent := filepath.Dir(relative)
if parent == "." {
return nil
}
rootHandle, err := os.OpenRoot(root)
if err != nil {
return fmt.Errorf("open agent fixture root: %w", err)
}
defer rootHandle.Close()
current := ""
for _, component := range strings.Split(filepath.ToSlash(parent), "/") {
if current == "" {
current = component
} else {
current = filepath.Join(current, component)
}
info, statErr := rootHandle.Lstat(current)
if errors.Is(statErr, os.ErrNotExist) {
if err := rootHandle.Mkdir(current, 0o700); err != nil {
return fmt.Errorf("create agent fixture parent: %w", err)
}
info, statErr = rootHandle.Lstat(current)
}
if statErr != nil {
return fmt.Errorf("inspect agent fixture parent: %w", statErr)
}
if !info.IsDir() || info.Mode()&os.ModeSymlink != 0 {
return rejectPath(PathRejectionSymlinkParent)
}
}
return nil
}
@@ -0,0 +1,219 @@
package fixture
import (
"errors"
"os"
"path/filepath"
"strings"
"testing"
)
func TestFixture_AgentPathContained(t *testing.T) {
// Given
root := newTestAgentRoot(t, "agent-alpha")
// When
path, err := root.Path("documents/report.txt")
// Then
requireNoFixtureError(t, err)
if !filepath.IsAbs(path.String()) {
t.Fatalf("agent path is not absolute: %q", path.String())
}
if path.Relative() != filepath.FromSlash("documents/report.txt") {
t.Fatalf("relative path = %q", path.Relative())
}
assertContainedPath(t, root.Absolute(), path.String())
info, err := os.Lstat(filepath.Join(root.Absolute(), "documents"))
requireNoFixtureError(t, err)
if !info.IsDir() || info.Mode()&os.ModeSymlink != 0 {
t.Fatalf("fixture parent mode = %s", info.Mode())
}
}
func TestFixture_AgentPathRejectsAbsolute(t *testing.T) {
// Given
root := newTestAgentRoot(t, "agent-absolute")
// When
_, err := root.Path(filepath.Join(t.TempDir(), "outside.txt"))
// Then
assertPathRejection(t, err, PathRejectionAbsolute)
}
func TestFixture_AgentPathRejectsParentEscape(t *testing.T) {
root := newTestAgentRoot(t, "agent-parent")
for _, candidate := range []string{"../outside.txt", "inside/../outside.txt"} {
t.Run(candidate, func(t *testing.T) {
_, err := root.Path(candidate)
assertPathRejection(t, err, PathRejectionParent)
})
}
}
func TestFixture_AgentPathRejectsVolumeOrSeparatorEscape(t *testing.T) {
root := newTestAgentRoot(t, "agent-volume")
tests := []struct {
name string
candidate string
reason PathRejectionReason
}{
{name: "drive absolute", candidate: `C:\outside.txt`, reason: PathRejectionVolume},
{name: "drive relative", candidate: `C:outside.txt`, reason: PathRejectionVolume},
{name: "UNC", candidate: `\\server\share\outside.txt`, reason: PathRejectionVolume},
{name: "alternate separator", candidate: `inside\outside.txt`, reason: PathRejectionSeparator},
{name: "empty component", candidate: "inside//outside.txt", reason: PathRejectionSeparator},
}
for _, test := range tests {
t.Run(test.name, func(t *testing.T) {
_, err := root.Path(test.candidate)
assertPathRejection(t, err, test.reason)
})
}
}
func TestFixture_AgentPathRejectsDestructiveRoot(t *testing.T) {
// Given
root := newTestAgentRoot(t, "agent-destructive")
// When
_, err := root.DestructivePath(".")
// Then
assertPathRejection(t, err, PathRejectionDestructiveRoot)
}
func TestFixture_AgentPathRejectsSymlinkParent(t *testing.T) {
// Given
root := newTestAgentRoot(t, "agent-symlink")
outside := t.TempDir()
requireNoFixtureError(t, os.Symlink(outside, filepath.Join(root.Absolute(), "linked")))
// When
_, err := root.Path("linked/file.txt")
// Then
assertPathRejection(t, err, PathRejectionSymlinkParent)
}
func TestFixture_AgentPathRejectsCleanedEscape(t *testing.T) {
// Given
root := newTestAgentRoot(t, "agent-cleaned-escape")
// When
_, err := root.Path("documents/../../outside.txt")
// Then
assertPathRejection(t, err, PathRejectionParent)
}
func TestFixture_AgentPathRejectsEmpty(t *testing.T) {
root := newTestAgentRoot(t, "agent-empty")
for _, candidate := range []string{"", " "} {
t.Run(candidate, func(t *testing.T) {
_, err := root.Path(candidate)
assertPathRejection(t, err, PathRejectionEmpty)
})
}
}
func TestFixture_AgentPathRejectsSymlinkRootParent(t *testing.T) {
// Given
realParent := t.TempDir()
symlinkParent := filepath.Join(t.TempDir(), "fixture-parent")
requireNoFixtureError(t, os.Symlink(realParent, symlinkParent))
// When
_, err := NewAgentRoot(symlinkParent, "agent-symlink-root")
// Then
if err == nil {
t.Fatal("symlink fixture parent was accepted")
}
}
func TestFixture_AgentPathRejectsADS(t *testing.T) {
root := newTestAgentRoot(t, "agent-ads")
candidate := "documents/report.txt:secret-value"
_, err := root.Path(candidate)
assertPathRejection(t, err, PathRejectionADS)
if strings.Contains(err.Error(), candidate) {
t.Fatal("path rejection exposed the candidate path")
}
}
func TestFixture_OutsideRootSentinelUnchanged(t *testing.T) {
// Given
parent := t.TempDir()
sentinel := filepath.Join(parent, "outside-sentinel")
requireNoFixtureError(t, os.WriteFile(sentinel, []byte("unchanged"), 0o600))
root, err := NewAgentRoot(parent, "agent-sentinel")
requireNoFixtureError(t, err)
// When
_, rejectionErr := root.Path("../outside-sentinel")
// Then
assertPathRejection(t, rejectionErr, PathRejectionParent)
content, err := os.ReadFile(sentinel)
requireNoFixtureError(t, err)
if string(content) != "unchanged" {
t.Fatalf("outside-root sentinel changed: %q", content)
}
}
func TestFixture_CreatesDistinctPerAgentRoots(t *testing.T) {
// Given
parent := t.TempDir()
// When
first, err := NewAgentRoot(parent, "agent-first")
requireNoFixtureError(t, err)
second, err := NewAgentRoot(parent, "agent-second")
requireNoFixtureError(t, err)
// Then
if first.Absolute() == second.Absolute() {
t.Fatal("per-Agent fixture roots are shared")
}
assertContainedPath(t, parent, first.Absolute())
assertContainedPath(t, parent, second.Absolute())
}
func newTestAgentRoot(t *testing.T, agentID string) AgentRoot {
t.Helper()
root, err := NewAgentRoot(t.TempDir(), agentID)
requireNoFixtureError(t, err)
return root
}
func assertContainedPath(t *testing.T, root, candidate string) {
t.Helper()
relative, err := filepath.Rel(root, candidate)
requireNoFixtureError(t, err)
if relative == ".." || filepath.IsAbs(relative) || (len(relative) > 3 && relative[:3] == ".."+string(filepath.Separator)) {
t.Fatalf("path %q escaped root %q", candidate, root)
}
}
func assertPathRejection(t *testing.T, err error, reason PathRejectionReason) {
t.Helper()
var pathError *AgentPathError
if !errors.As(err, &pathError) {
t.Fatalf("expected AgentPathError, got %v", err)
}
if pathError.Reason != reason {
t.Fatalf("rejection reason = %q, want %q", pathError.Reason, reason)
}
if pathError.Error() == "" {
t.Fatal("path rejection error is empty")
}
}
func requireNoFixtureError(t *testing.T, err error) {
t.Helper()
if err != nil {
t.Fatal(err)
}
}
@@ -0,0 +1,34 @@
package fixture
import "errors"
var (
ErrPayloadOverrun = errors.New("fixture payload exceeds transfer limit")
ErrPayloadSizeMismatch = errors.New("fixture payload size mismatch")
)
type PathRejectionReason string
const (
PathRejectionEmpty PathRejectionReason = "empty"
PathRejectionAbsolute PathRejectionReason = "absolute"
PathRejectionParent PathRejectionReason = "parent"
PathRejectionVolume PathRejectionReason = "volume"
PathRejectionSeparator PathRejectionReason = "separator"
PathRejectionDestructiveRoot PathRejectionReason = "destructive_root"
PathRejectionEscape PathRejectionReason = "escape"
PathRejectionSymlinkParent PathRejectionReason = "symlink_parent"
PathRejectionADS PathRejectionReason = "ads"
)
type AgentPathError struct {
Reason PathRejectionReason
}
func (e *AgentPathError) Error() string {
return "agent path rejected: " + string(e.Reason)
}
func rejectPath(reason PathRejectionReason) error {
return &AgentPathError{Reason: reason}
}
@@ -0,0 +1,110 @@
package fixture
import (
"crypto/ecdsa"
"crypto/elliptic"
"crypto/rand"
"crypto/tls"
"crypto/x509"
"crypto/x509/pkix"
"encoding/pem"
"errors"
"math/big"
"net"
"time"
)
type LocalTLSFixture struct {
certificate tls.Certificate
rootCAs *x509.CertPool
caPEM []byte
certificatePEM []byte
privateKeyPEM []byte
}
func NewLocalTLSFixture(now time.Time) (LocalTLSFixture, error) {
caKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
return LocalTLSFixture{}, err
}
caTemplate := &x509.Certificate{
SerialNumber: big.NewInt(1),
Subject: pkix.Name{CommonName: "agentcompat local CA"},
NotBefore: now.Add(-time.Hour),
NotAfter: now.Add(24 * time.Hour),
IsCA: true,
BasicConstraintsValid: true,
KeyUsage: x509.KeyUsageCertSign | x509.KeyUsageDigitalSignature,
}
caDER, err := x509.CreateCertificate(rand.Reader, caTemplate, caTemplate, &caKey.PublicKey, caKey)
if err != nil {
return LocalTLSFixture{}, err
}
leafKey, err := ecdsa.GenerateKey(elliptic.P256(), rand.Reader)
if err != nil {
return LocalTLSFixture{}, err
}
leafTemplate := &x509.Certificate{
SerialNumber: big.NewInt(2),
Subject: pkix.Name{CommonName: "localhost"},
NotBefore: now.Add(-time.Hour),
NotAfter: now.Add(24 * time.Hour),
KeyUsage: x509.KeyUsageDigitalSignature,
ExtKeyUsage: []x509.ExtKeyUsage{x509.ExtKeyUsageServerAuth},
DNSNames: []string{"localhost"},
IPAddresses: []net.IP{net.ParseIP("127.0.0.1"), net.ParseIP("::1")},
}
leafDER, err := x509.CreateCertificate(rand.Reader, leafTemplate, caTemplate, &leafKey.PublicKey, caKey)
if err != nil {
return LocalTLSFixture{}, err
}
leafKeyDER, err := x509.MarshalPKCS8PrivateKey(leafKey)
if err != nil {
return LocalTLSFixture{}, err
}
caPEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: caDER})
certificatePEM := pem.EncodeToMemory(&pem.Block{Type: "CERTIFICATE", Bytes: leafDER})
privateKeyPEM := pem.EncodeToMemory(&pem.Block{Type: "PRIVATE KEY", Bytes: leafKeyDER})
certificate, err := tls.X509KeyPair(certificatePEM, privateKeyPEM)
if err != nil {
return LocalTLSFixture{}, err
}
rootCAs := x509.NewCertPool()
if !rootCAs.AppendCertsFromPEM(caPEM) {
return LocalTLSFixture{}, errors.New("append local CA certificate")
}
return LocalTLSFixture{
certificate: certificate,
rootCAs: rootCAs,
caPEM: caPEM,
certificatePEM: certificatePEM,
privateKeyPEM: privateKeyPEM,
}, nil
}
func (fixture LocalTLSFixture) ClientConfig(serverName string) *tls.Config {
return &tls.Config{
MinVersion: tls.VersionTLS13,
RootCAs: fixture.rootCAs.Clone(),
ServerName: serverName,
}
}
func (fixture LocalTLSFixture) Listener(listener net.Listener) net.Listener {
return tls.NewListener(listener, &tls.Config{
MinVersion: tls.VersionTLS13,
Certificates: []tls.Certificate{fixture.certificate},
})
}
func (fixture LocalTLSFixture) CAPEM() []byte {
return append([]byte(nil), fixture.caPEM...)
}
func (fixture LocalTLSFixture) CertificatePEM() []byte {
return append([]byte(nil), fixture.certificatePEM...)
}
func (fixture LocalTLSFixture) PrivateKeyPEM() []byte {
return append([]byte(nil), fixture.privateKeyPEM...)
}
@@ -0,0 +1,130 @@
package fixture
import (
"context"
"crypto/tls"
"crypto/x509"
"encoding/pem"
"errors"
"io"
"net"
"net/http"
"strings"
"testing"
"time"
)
func TestFixture_VerifiesLocalhostTLS(t *testing.T) {
// Given
fixture, err := NewLocalTLSFixture(time.Now())
requireNoFixtureError(t, err)
address, closeServer := startLocalTLSServer(t, fixture)
defer closeServer()
client := localTLSClient(fixture.ClientConfig("localhost"), address)
// When
response, err := client.Get("https://localhost:" + portOf(t, address) + "/ready")
requireNoFixtureError(t, err)
defer response.Body.Close()
body, err := io.ReadAll(response.Body)
requireNoFixtureError(t, err)
// Then
if response.StatusCode != http.StatusOK || string(body) != "tls-ready" {
t.Fatalf("TLS response = %d %q", response.StatusCode, body)
}
if fixture.ClientConfig("localhost").InsecureSkipVerify {
t.Fatal("TLS fixture disabled certificate verification")
}
if len(fixture.CAPEM()) == 0 || len(fixture.CertificatePEM()) == 0 || len(fixture.PrivateKeyPEM()) == 0 {
t.Fatal("TLS fixture did not expose certificate material")
}
assertLocalCertificateProperties(t, fixture)
}
func assertLocalCertificateProperties(t *testing.T, fixture LocalTLSFixture) {
t.Helper()
caBlock, _ := pem.Decode(fixture.CAPEM())
if caBlock == nil {
t.Fatal("fixture CA PEM is invalid")
}
caCertificate, err := x509.ParseCertificate(caBlock.Bytes)
requireNoFixtureError(t, err)
if !caCertificate.IsCA || !caCertificate.BasicConstraintsValid || caCertificate.KeyUsage&x509.KeyUsageCertSign == 0 {
t.Fatalf("fixture CA constraints are invalid: %+v", caCertificate)
}
leafBlock, _ := pem.Decode(fixture.CertificatePEM())
if leafBlock == nil {
t.Fatal("fixture leaf PEM is invalid")
}
leafCertificate, err := x509.ParseCertificate(leafBlock.Bytes)
requireNoFixtureError(t, err)
if err := leafCertificate.VerifyHostname("localhost"); err != nil {
t.Fatalf("verify localhost SAN: %v", err)
}
if err := leafCertificate.VerifyHostname("127.0.0.1"); err != nil {
t.Fatalf("verify loopback SAN: %v", err)
}
}
func TestFixture_RejectsTLSNameMismatch(t *testing.T) {
// Given
fixture, err := NewLocalTLSFixture(time.Now())
requireNoFixtureError(t, err)
address, closeServer := startLocalTLSServer(t, fixture)
defer closeServer()
client := localTLSClient(fixture.ClientConfig("wronghost.invalid"), address)
// When
_, err = client.Get("https://wronghost.invalid:" + portOf(t, address) + "/ready")
// Then
var hostnameError x509.HostnameError
if !errors.As(err, &hostnameError) {
t.Fatalf("TLS mismatch error = %v", err)
}
}
func startLocalTLSServer(t *testing.T, fixture LocalTLSFixture) (string, func()) {
t.Helper()
listener, err := net.Listen("tcp", "127.0.0.1:0")
requireNoFixtureError(t, err)
tlsListener := fixture.Listener(listener)
server := &http.Server{Handler: http.HandlerFunc(func(writer http.ResponseWriter, request *http.Request) {
if request.URL.Path != "/ready" {
http.NotFound(writer, request)
return
}
_, _ = io.WriteString(writer, "tls-ready")
})}
serveDone := make(chan error, 1)
go func() { serveDone <- server.Serve(tlsListener) }()
closeServer := func() {
ctx, cancel := context.WithTimeout(context.Background(), time.Second)
defer cancel()
requireNoFixtureError(t, server.Shutdown(ctx))
serveErr := <-serveDone
if serveErr != nil && !errors.Is(serveErr, http.ErrServerClosed) {
t.Fatalf("serve TLS: %v", serveErr)
}
}
return listener.Addr().String(), closeServer
}
func localTLSClient(config *tls.Config, address string) *http.Client {
transport := &http.Transport{
TLSClientConfig: config,
DialContext: func(ctx context.Context, network, _ string) (net.Conn, error) {
var dialer net.Dialer
return dialer.DialContext(ctx, network, address)
},
}
return &http.Client{Transport: transport, Timeout: 2 * time.Second}
}
func portOf(t *testing.T, address string) string {
t.Helper()
_, port, err := net.SplitHostPort(address)
requireNoFixtureError(t, err)
return strings.TrimSpace(port)
}
@@ -0,0 +1,179 @@
package fixture
import (
"bufio"
"context"
"errors"
"fmt"
"io"
"net"
"net/http"
"sync"
)
type NATEchoRecord struct {
Method string
Path string
Host string
HeaderValue string
Body []byte
RequestHalfClosed bool
}
type natEchoResult struct {
record NATEchoRecord
err error
}
type NATEchoBackend struct {
listener net.Listener
results chan natEchoResult
done chan struct{}
requireHalfClose bool
closeOnce sync.Once
waitGroup sync.WaitGroup
mutex sync.Mutex
connections map[net.Conn]struct{}
}
func StartNATEchoBackend() (*NATEchoBackend, error) {
return startNATEchoBackend(false)
}
func StartNATHalfCloseEchoBackend() (*NATEchoBackend, error) {
return startNATEchoBackend(true)
}
func startNATEchoBackend(requireHalfClose 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)
}
backend := &NATEchoBackend{
listener: listener,
results: make(chan natEchoResult, 16),
done: make(chan struct{}),
requireHalfClose: requireHalfClose,
connections: make(map[net.Conn]struct{}),
}
backend.waitGroup.Add(1)
go backend.accept()
return backend, nil
}
func (backend *NATEchoBackend) Address() string {
return backend.listener.Addr().String()
}
func (backend *NATEchoBackend) 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 echo backend closed")
}
}
func (backend *NATEchoBackend) Close() error {
var closeErr error
backend.closeOnce.Do(func() {
close(backend.done)
closeErr = backend.listener.Close()
backend.mutex.Lock()
connections := make([]net.Conn, 0, len(backend.connections))
for connection := range backend.connections {
connections = append(connections, connection)
}
backend.mutex.Unlock()
for _, connection := range connections {
closeErr = errors.Join(closeErr, connection.Close())
}
backend.waitGroup.Wait()
})
if errors.Is(closeErr, net.ErrClosed) {
return nil
}
return closeErr
}
func (backend *NATEchoBackend) accept() {
defer backend.waitGroup.Done()
for {
connection, err := backend.listener.Accept()
if err != nil {
if errors.Is(err, net.ErrClosed) {
return
}
backend.publish(natEchoResult{err: fmt.Errorf("accept NAT echo connection: %w", err)})
return
}
backend.mutex.Lock()
backend.connections[connection] = struct{}{}
backend.mutex.Unlock()
backend.waitGroup.Add(1)
go backend.handle(connection)
}
}
func (backend *NATEchoBackend) handle(connection net.Conn) {
defer backend.waitGroup.Done()
defer func() {
backend.mutex.Lock()
delete(backend.connections, connection)
backend.mutex.Unlock()
_ = connection.Close()
}()
reader := bufio.NewReader(connection)
request, err := http.ReadRequest(reader)
if err != nil {
backend.publish(natEchoResult{err: fmt.Errorf("read NAT echo request: %w", err)})
return
}
body, err := io.ReadAll(request.Body)
if closeErr := request.Body.Close(); err == nil {
err = closeErr
}
if err != nil {
backend.publish(natEchoResult{err: fmt.Errorf("read NAT echo body: %w", err)})
return
}
halfClosed := false
if backend.requireHalfClose {
_, halfCloseErr := reader.ReadByte()
halfClosed = errors.Is(halfCloseErr, io.EOF)
if halfCloseErr != nil && !halfClosed {
backend.publish(natEchoResult{err: fmt.Errorf("observe NAT request half-close: %w", halfCloseErr)})
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...),
RequestHalfClosed: halfClosed,
}
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(natEchoResult{err: fmt.Errorf("write NAT echo response: %w", err)})
return
}
if tcpConnection, ok := connection.(*net.TCPConn); ok {
if err := tcpConnection.CloseWrite(); err != nil {
backend.publish(natEchoResult{err: fmt.Errorf("half-close NAT echo response: %w", err)})
return
}
}
backend.publish(natEchoResult{record: record})
}
func (backend *NATEchoBackend) publish(result natEchoResult) {
select {
case backend.results <- result:
case <-backend.done:
}
}
@@ -0,0 +1,84 @@
package fixture
import (
"bufio"
"context"
"fmt"
"io"
"net"
"net/http"
"strings"
"testing"
"time"
)
func TestFixture_NATEchoHalfClose(t *testing.T) {
// Given
backend, err := StartNATHalfCloseEchoBackend()
requireNoFixtureError(t, err)
t.Cleanup(func() { requireNoFixtureError(t, backend.Close()) })
connection, err := net.DialTimeout("tcp", backend.Address(), time.Second)
requireNoFixtureError(t, err)
tcpConnection := connection.(*net.TCPConn)
defer tcpConnection.Close()
requireNoFixtureError(t, tcpConnection.SetDeadline(time.Now().Add(2*time.Second)))
requestBody := "half-closed"
request := fmt.Sprintf("POST /echo?case=half-close HTTP/1.1\r\nHost: nat.invalid\r\nX-AgentCompat-Echo: fixture\r\nContent-Length: %d\r\nConnection: close\r\n\r\n%s", len(requestBody), requestBody)
// When
_, err = io.WriteString(tcpConnection, request)
requireNoFixtureError(t, err)
requireNoFixtureError(t, tcpConnection.CloseWrite())
response, err := http.ReadResponse(bufio.NewReader(tcpConnection), nil)
requireNoFixtureError(t, err)
responseBody, 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
const expected = "method=POST\npath=/echo?case=half-close\nhost=nat.invalid\nx-agentcompat-echo=fixture\nbody=half-closed\n"
if response.StatusCode != http.StatusOK || string(responseBody) != expected {
t.Fatalf("NAT echo response = %d %q", response.StatusCode, responseBody)
}
if !record.RequestHalfClosed {
t.Fatal("backend responded before observing request half-close")
}
if record.Method != "POST" || record.Path != "/echo?case=half-close" || record.Host != "nat.invalid" || record.HeaderValue != "fixture" || string(record.Body) != requestBody {
t.Fatalf("NAT request record = %+v", record)
}
}
func TestFixture_NATEcho(t *testing.T) {
// Given
backend, err := StartNATEchoBackend()
requireNoFixtureError(t, err)
t.Cleanup(func() { requireNoFixtureError(t, backend.Close()) })
ctx, cancel := context.WithTimeout(context.Background(), 2*time.Second)
defer cancel()
request, err := http.NewRequestWithContext(ctx, http.MethodPut, "http://"+backend.Address()+"/echo?case=ordinary", strings.NewReader("ordinary"))
requireNoFixtureError(t, err)
request.Host = "nat.invalid"
request.Header.Set("X-AgentCompat-Echo", "fixture")
// When
response, err := http.DefaultClient.Do(request)
requireNoFixtureError(t, err)
responseBody, err := io.ReadAll(response.Body)
requireNoFixtureError(t, err)
requireNoFixtureError(t, response.Body.Close())
record, err := backend.WaitRequest(ctx)
requireNoFixtureError(t, err)
// Then
const expected = "method=PUT\npath=/echo?case=ordinary\nhost=nat.invalid\nx-agentcompat-echo=fixture\nbody=ordinary\n"
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" {
t.Fatalf("NAT request record = %+v", record)
}
}
@@ -0,0 +1,43 @@
package fixture
import (
"crypto/sha256"
"encoding/hex"
"fmt"
"io"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
)
const verifierBufferBytes = 1024 * 1024
type PayloadDigest struct {
Bytes uint64
SHA256 [sha256.Size]byte
}
func (d PayloadDigest) Hex() string {
return hex.EncodeToString(d.SHA256[:])
}
func VerifyPayload(reader io.Reader, expectedBytes uint64) (PayloadDigest, error) {
if expectedBytes > contract.TransferBytes {
return PayloadDigest{}, ErrPayloadOverrun
}
hash := sha256.New()
buffer := make([]byte, verifierBufferBytes)
limited := io.LimitReader(reader, int64(expectedBytes)+1)
written, err := io.CopyBuffer(hash, limited, buffer)
if err != nil {
return PayloadDigest{}, fmt.Errorf("verify payload: %w", err)
}
if uint64(written) > expectedBytes {
return PayloadDigest{}, ErrPayloadOverrun
}
if uint64(written) != expectedBytes {
return PayloadDigest{}, ErrPayloadSizeMismatch
}
digest := PayloadDigest{Bytes: uint64(written)}
copy(digest.SHA256[:], hash.Sum(nil))
return digest, nil
}
@@ -0,0 +1,68 @@
package fixture
import (
"crypto/sha256"
"encoding/binary"
"fmt"
"io"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
)
const payloadBlockSize = sha256.Size
type Payload struct {
seed contract.Seed
size uint64
}
type payloadReader struct {
payload Payload
offset uint64
}
func NewPayload(seed contract.Seed, size uint64) (Payload, error) {
if seed == 0 {
return Payload{}, fmt.Errorf("payload seed must be nonzero")
}
if size > contract.TransferBytes {
return Payload{}, ErrPayloadOverrun
}
return Payload{seed: seed, size: size}, nil
}
func (p Payload) Reader() io.Reader {
return &payloadReader{payload: p}
}
func (r *payloadReader) Read(destination []byte) (int, error) {
if r.offset >= r.payload.size {
return 0, io.EOF
}
remaining := r.payload.size - r.offset
if uint64(len(destination)) > remaining {
destination = destination[:remaining]
}
written := fillPayload(destination, r.payload.seed, r.offset)
r.offset += uint64(written)
if r.offset == r.payload.size {
return written, io.EOF
}
return written, nil
}
func fillPayload(destination []byte, seed contract.Seed, offset uint64) int {
written := 0
for written < len(destination) {
absoluteOffset := offset + uint64(written)
blockIndex := absoluteOffset / payloadBlockSize
blockOffset := absoluteOffset % payloadBlockSize
var input [16]byte
binary.BigEndian.PutUint64(input[:8], uint64(seed))
binary.BigEndian.PutUint64(input[8:], blockIndex)
block := sha256.Sum256(input[:])
copied := copy(destination[written:], block[blockOffset:])
written += copied
}
return written
}
@@ -0,0 +1,155 @@
package fixture
import (
"bytes"
"crypto/sha256"
"encoding/binary"
"errors"
"fmt"
"io"
"runtime"
"testing"
"github.com/nezhahq/nezha/integration/agentcompat/internal/contract"
)
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)
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 {
t.Fatalf("payload bytes = %d", digest.Bytes)
}
if digest.SHA256 != stableDigest.SHA256 {
t.Fatalf("payload SHA changed: %s != %s", digest.Hex(), stableDigest.Hex())
}
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 retained > contract.TransferHeapBytes {
t.Fatalf("retained heap = %d, limit = %d", retained, contract.TransferHeapBytes)
}
t.Logf("bytes=%d sha256=%s retained_heap=%d", digest.Bytes, digest.Hex(), retained)
}
func independentlyHashPayload(seed contract.Seed, size uint64) [sha256.Size]byte {
hash := sha256.New()
var input [16]byte
var remaining = size
for blockIndex := uint64(0); remaining > 0; blockIndex++ {
binary.BigEndian.PutUint64(input[:8], uint64(seed))
binary.BigEndian.PutUint64(input[8:], blockIndex)
block := sha256.Sum256(input[:])
writeBytes := uint64(len(block))
if remaining < writeBytes {
writeBytes = remaining
}
_, _ = hash.Write(block[:writeBytes])
remaining -= writeBytes
}
var digest [sha256.Size]byte
copy(digest[:], hash.Sum(nil))
return digest
}
func TestFixture_RejectsPayloadOverrun(t *testing.T) {
// Given
_, constructorErr := NewPayload(contract.DefaultSeed, contract.TransferBytes+1)
// When
_, verifierErr := VerifyPayload(bytes.NewReader([]byte("overrun")), 1)
_, declaredSizeErr := VerifyPayload(bytes.NewReader(nil), contract.TransferBytes+1)
// Then
if !errors.Is(constructorErr, ErrPayloadOverrun) {
t.Fatalf("constructor error = %v", constructorErr)
}
if !errors.Is(verifierErr, ErrPayloadOverrun) {
t.Fatalf("verifier error = %v", verifierErr)
}
if !errors.Is(declaredSizeErr, ErrPayloadOverrun) {
t.Fatalf("declared size error = %v", declaredSizeErr)
}
}
func TestFixture_PayloadChunkIndependence(t *testing.T) {
// Given
const size = 64 * 1024
payload, err := NewPayload(contract.DefaultSeed, size)
requireNoFixtureError(t, err)
chunkSizes := []int{1, 1024, 1024 * 1024, 7919}
var baseline []byte
for _, chunkSize := range chunkSizes {
// When
content := readPayloadWithChunkSize(t, payload.Reader(), chunkSize)
// Then
if baseline == nil {
baseline = content
continue
}
if !bytes.Equal(content, baseline) {
t.Fatalf("payload changed with chunk size %d", chunkSize)
}
}
}
func TestFixture_PayloadDigestStableAtBoundaries(t *testing.T) {
for _, size := range []uint64{0, 1, 1024 * 1024} {
t.Run(fmt.Sprintf("bytes_%d", size), func(t *testing.T) {
payload, err := NewPayload(contract.DefaultSeed, size)
requireNoFixtureError(t, err)
first, err := VerifyPayload(payload.Reader(), size)
requireNoFixtureError(t, err)
second, err := VerifyPayload(payload.Reader(), size)
requireNoFixtureError(t, err)
if first != second || first.Bytes != size {
t.Fatalf("digest unstable at %d bytes: %+v != %+v", size, first, second)
}
})
}
}
func TestFixture_VerifierRejectsShortPayload(t *testing.T) {
_, err := VerifyPayload(bytes.NewReader([]byte("short")), 6)
if !errors.Is(err, ErrPayloadSizeMismatch) {
t.Fatalf("short payload error = %v", err)
}
}
func readPayloadWithChunkSize(t *testing.T, reader io.Reader, chunkSize int) []byte {
t.Helper()
buffer := make([]byte, chunkSize)
var content bytes.Buffer
for {
readBytes, err := reader.Read(buffer)
if readBytes > 0 {
_, writeErr := content.Write(buffer[:readBytes])
requireNoFixtureError(t, writeErr)
}
if errors.Is(err, io.EOF) {
return content.Bytes()
}
requireNoFixtureError(t, err)
}
}