diff --git a/integration/agentcompat/internal/fixture/agent_path.go b/integration/agentcompat/internal/fixture/agent_path.go new file mode 100644 index 00000000..6e1bef6e --- /dev/null +++ b/integration/agentcompat/internal/fixture/agent_path.go @@ -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 +} diff --git a/integration/agentcompat/internal/fixture/agent_path_test.go b/integration/agentcompat/internal/fixture/agent_path_test.go new file mode 100644 index 00000000..d274b273 --- /dev/null +++ b/integration/agentcompat/internal/fixture/agent_path_test.go @@ -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) + } +} diff --git a/integration/agentcompat/internal/fixture/fixture_errors.go b/integration/agentcompat/internal/fixture/fixture_errors.go new file mode 100644 index 00000000..790e0683 --- /dev/null +++ b/integration/agentcompat/internal/fixture/fixture_errors.go @@ -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} +} diff --git a/integration/agentcompat/internal/fixture/local_ca.go b/integration/agentcompat/internal/fixture/local_ca.go new file mode 100644 index 00000000..15aeed6c --- /dev/null +++ b/integration/agentcompat/internal/fixture/local_ca.go @@ -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...) +} diff --git a/integration/agentcompat/internal/fixture/local_ca_test.go b/integration/agentcompat/internal/fixture/local_ca_test.go new file mode 100644 index 00000000..c4b7748f --- /dev/null +++ b/integration/agentcompat/internal/fixture/local_ca_test.go @@ -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) +} diff --git a/integration/agentcompat/internal/fixture/nat_echo.go b/integration/agentcompat/internal/fixture/nat_echo.go new file mode 100644 index 00000000..f11e79fb --- /dev/null +++ b/integration/agentcompat/internal/fixture/nat_echo.go @@ -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: + } +} diff --git a/integration/agentcompat/internal/fixture/nat_echo_test.go b/integration/agentcompat/internal/fixture/nat_echo_test.go new file mode 100644 index 00000000..bab246a3 --- /dev/null +++ b/integration/agentcompat/internal/fixture/nat_echo_test.go @@ -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) + } +} diff --git a/integration/agentcompat/internal/fixture/payload_hash.go b/integration/agentcompat/internal/fixture/payload_hash.go new file mode 100644 index 00000000..f98450db --- /dev/null +++ b/integration/agentcompat/internal/fixture/payload_hash.go @@ -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 +} diff --git a/integration/agentcompat/internal/fixture/payload_reader.go b/integration/agentcompat/internal/fixture/payload_reader.go new file mode 100644 index 00000000..e6688e33 --- /dev/null +++ b/integration/agentcompat/internal/fixture/payload_reader.go @@ -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 +} diff --git a/integration/agentcompat/internal/fixture/payload_test.go b/integration/agentcompat/internal/fixture/payload_test.go new file mode 100644 index 00000000..ab940032 --- /dev/null +++ b/integration/agentcompat/internal/fixture/payload_test.go @@ -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) + } +}