Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
28 changes: 28 additions & 0 deletions internal/engine/engine_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -6,6 +6,7 @@ import (
"math/rand"
"os"
"path/filepath"
"runtime"
"strings"
"testing"

Expand Down Expand Up @@ -424,6 +425,33 @@ func TestVerifyExtraFiles(t *testing.T) {
}
}

func TestExtraFilesWalkError(t *testing.T) {
_, err := extraFiles(filepath.Join(t.TempDir(), "missing"), &manifest.Manifest{}, nil)
if err == nil {
t.Fatal("extraFiles should fail when the root cannot be walked")
}
}

func TestVerifyExtraFailsWhenWalkFails(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("unix permissions")
}
src := sourceDir(t)
dst := filepath.Join(t.TempDir(), "out")
if _, err := Run(context.Background(), baseConfig(src, dst)); err != nil {
t.Fatal(err)
}
blocked := filepath.Join(dst, "blocked")
if err := os.Mkdir(blocked, 0); err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = os.Chmod(blocked, 0o755) })

if _, err := Verify(context.Background(), VerifyConfig{Root: dst, Extra: true}); err == nil {
t.Fatal("Verify --extra should fail when the destination cannot be listed")
}
}

func TestVerifyWithExplicitManifest(t *testing.T) {
src := sourceDir(t)
out := filepath.Join(t.TempDir(), "model.json")
Expand Down
12 changes: 8 additions & 4 deletions internal/engine/verify.go
Original file line number Diff line number Diff line change
Expand Up @@ -181,7 +181,11 @@ func Verify(ctx context.Context, cfg VerifyConfig) (*VerifyResult, error) {
res.OK = len(problems) == 0

if cfg.Extra {
res.Extra = extraFiles(root, m, cfg.Warn)
extra, err := extraFiles(root, m, cfg.Warn)
if err != nil {
return nil, fmt.Errorf("engine: listing extra files: %w", err)
}
res.Extra = extra
if len(res.Extra) > 0 {
res.OK = false
}
Expand Down Expand Up @@ -265,10 +269,10 @@ func badChunks(ctx context.Context, ra *os.File, f *manifest.File) ([]BadChunk,
return out, bytes
}

func extraFiles(root string, m *manifest.Manifest, warn func(string, ...any)) []string {
func extraFiles(root string, m *manifest.Manifest, warn func(string, ...any)) ([]string, error) {
entries, err := scan.Walk(scan.Options{Root: root, IncludeHidden: true, Warn: warn})
if err != nil {
return nil
return nil, err
}
known := make(map[string]struct{}, len(m.Files))
for _, f := range m.Files {
Expand All @@ -281,5 +285,5 @@ func extraFiles(root string, m *manifest.Manifest, warn func(string, ...any)) []
}
}
sort.Strings(extra)
return extra
return extra, nil
}
11 changes: 10 additions & 1 deletion internal/protocol/gzip.go
Original file line number Diff line number Diff line change
Expand Up @@ -23,14 +23,23 @@ func gzipBytes(p []byte) ([]byte, error) {
}

func gunzipBytes(p []byte) ([]byte, error) {
return gunzipBytesLimit(p, MaxFrame)
}

// gunzipBytesLimit decompresses p and rejects output larger than max so a
// small gzip frame cannot expand into an unbounded allocation on the helper.
func gunzipBytesLimit(p []byte, max int64) ([]byte, error) {
r, err := gzip.NewReader(bytes.NewReader(p))
if err != nil {
return nil, fmt.Errorf("protocol: gzip manifest: %w", err)
}
defer r.Close()
out, err := io.ReadAll(r)
out, err := io.ReadAll(io.LimitReader(r, max+1))
if err != nil {
return nil, fmt.Errorf("protocol: gzip manifest: %w", err)
}
if int64(len(out)) > max {
return nil, fmt.Errorf("protocol: gzip manifest exceeds %d bytes", max)
}
return out, nil
}
40 changes: 40 additions & 0 deletions internal/protocol/protocol_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -490,6 +490,28 @@ func TestServerRejectsOutOfOrderFrames(t *testing.T) {
}
}

func TestServerRejectsSecondPlan(t *testing.T) {
src := sourceDir(t)
dst := filepath.Join(t.TempDir(), "out")
m, err := scan.Build(context.Background(), scan.Options{Root: src})
if err != nil {
t.Fatal(err)
}

client, done := pipePair(t, ServerOptions{Root: dst})
if _, err := client.Plan(context.Background(), m, defaultRequest()); err != nil {
t.Fatal(err)
}
_, err = client.Plan(context.Background(), m, defaultRequest())
done()
if err == nil {
t.Fatal("the helper accepted a second plan in the same session")
}
if !bytes.Contains([]byte(err.Error()), []byte("already planned")) {
t.Errorf("error = %v, want it to mention already planned", err)
}
}

func TestCancelledContextStopsClient(t *testing.T) {
client, done := pipePair(t, ServerOptions{Root: t.TempDir()})
defer done()
Expand Down Expand Up @@ -539,6 +561,24 @@ func TestGzipManifestRoundTrip(t *testing.T) {
}
}

func TestGunzipRejectsOversized(t *testing.T) {
payload := bytes.Repeat([]byte("x"), 64)
gz, err := gzipBytes(payload)
if err != nil {
t.Fatal(err)
}
if _, err := gunzipBytesLimit(gz, 16); err == nil {
t.Fatal("gunzipBytesLimit accepted a payload larger than the limit")
}
got, err := gunzipBytesLimit(gz, int64(len(payload)))
if err != nil {
t.Fatal(err)
}
if !bytes.Equal(got, payload) {
t.Fatal("gunzipBytesLimit rejected a payload at the limit")
}
}

func TestManifestGzipOverPipes(t *testing.T) {
src := sourceDir(t)
dst := filepath.Join(t.TempDir(), "out")
Expand Down
4 changes: 4 additions & 0 deletions internal/protocol/server.go
Original file line number Diff line number Diff line change
Expand Up @@ -139,6 +139,10 @@ func (s *server) warn(format string, args ...any) {
}

func (s *server) handlePlan(ctx context.Context, payload []byte) error {
if s.recv != nil {
s.abort()
return fmt.Errorf("protocol: already planned")
}
var req PlanRequest
if err := DecodeJSON(payload, &req); err != nil {
return err
Expand Down
2 changes: 1 addition & 1 deletion internal/receiver/receiver.go
Original file line number Diff line number Diff line change
Expand Up @@ -515,7 +515,7 @@ func (r *Receiver) Finish() (*Summary, error) {
applied := *r.manifest
applied.Model.Root = r.root
if err := manifest.Save(manifest.StatePath(r.root, manifest.ManifestName), &applied, manifest.EncodingJSON); err != nil {
r.opt.warnf("cannot record manifest: %v", err)
return nil, fmt.Errorf("receiver: cannot record manifest: %w", err)
}
}
r.cleanStage()
Expand Down
37 changes: 32 additions & 5 deletions internal/receiver/receiver_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@ import (
"math/rand"
"os"
"path/filepath"
"runtime"
"strings"
"testing"

Expand Down Expand Up @@ -58,6 +59,16 @@ func apply(t *testing.T, src, dst string, opt Options) (*Plan, *Summary) {
}

func applyManifest(t *testing.T, m *manifest.Manifest, src, dst string, opt Options) (*Plan, *Summary) {
t.Helper()
r, plan := applyFiles(t, m, src, dst, opt)
sum, err := r.Finish()
if err != nil {
t.Fatalf("Finish: %v", err)
}
return plan, sum
}

func applyFiles(t *testing.T, m *manifest.Manifest, src, dst string, opt Options) (*Receiver, *Plan) {
t.Helper()
opt.Root = dst
r, err := New(opt)
Expand Down Expand Up @@ -101,11 +112,7 @@ func applyManifest(t *testing.T, m *manifest.Manifest, src, dst string, opt Opti
t.Fatalf("EndFile %s: %v", fp.Path, err)
}
}
sum, err := r.Finish()
if err != nil {
t.Fatalf("Finish: %v", err)
}
return plan, sum
return r, plan
}

func defaults(dst string) Options {
Expand Down Expand Up @@ -718,3 +725,23 @@ func TestUnderRoot(t *testing.T) {
}
}
}

func TestFinishFailsWhenAppliedManifestUnwritable(t *testing.T) {
if runtime.GOOS == "windows" {
t.Skip("unix permissions")
}
src := sourceDir(t)
dst := filepath.Join(t.TempDir(), "out")
m := buildManifest(t, src)
r, _ := applyFiles(t, m, src, dst, defaults(dst))

state := filepath.Join(dst, manifest.StateDir)
if err := os.Chmod(state, 0o555); err != nil {
t.Fatal(err)
}
t.Cleanup(func() { _ = os.Chmod(state, 0o755) })

if _, err := r.Finish(); err == nil {
t.Fatal("Finish succeeded when the applied manifest could not be saved")
}
}