mirror of
https://github.com/Buriburizaem0n/nezha_domains.git
synced 2026-09-21 02:30:14 +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