mirror of
https://github.com/Buriburizaem0n/nezha_domains.git
synced 2026-09-19 09:40:12 +00:00
test(compat): add deterministic local fixtures
Co-authored-by: naiba/CloudCode <hi+cloudcode@nai.ba>
This commit is contained in:
@@ -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)
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user