diff --git a/internal/engine/engine_test.go b/internal/engine/engine_test.go index 670a015..be753e5 100644 --- a/internal/engine/engine_test.go +++ b/internal/engine/engine_test.go @@ -6,6 +6,7 @@ import ( "math/rand" "os" "path/filepath" + "runtime" "strings" "testing" @@ -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") diff --git a/internal/engine/verify.go b/internal/engine/verify.go index 6d8a58f..15497d0 100644 --- a/internal/engine/verify.go +++ b/internal/engine/verify.go @@ -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 } @@ -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 { @@ -281,5 +285,5 @@ func extraFiles(root string, m *manifest.Manifest, warn func(string, ...any)) [] } } sort.Strings(extra) - return extra + return extra, nil } diff --git a/internal/protocol/gzip.go b/internal/protocol/gzip.go index 05c2615..34be2e3 100644 --- a/internal/protocol/gzip.go +++ b/internal/protocol/gzip.go @@ -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 } diff --git a/internal/protocol/protocol_test.go b/internal/protocol/protocol_test.go index 52e0cbe..b44e300 100644 --- a/internal/protocol/protocol_test.go +++ b/internal/protocol/protocol_test.go @@ -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() @@ -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") diff --git a/internal/protocol/server.go b/internal/protocol/server.go index e7cdca5..8df4063 100644 --- a/internal/protocol/server.go +++ b/internal/protocol/server.go @@ -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 diff --git a/internal/receiver/receiver.go b/internal/receiver/receiver.go index c020946..94f5f69 100644 --- a/internal/receiver/receiver.go +++ b/internal/receiver/receiver.go @@ -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() diff --git a/internal/receiver/receiver_test.go b/internal/receiver/receiver_test.go index 799d890..7e21276 100644 --- a/internal/receiver/receiver_test.go +++ b/internal/receiver/receiver_test.go @@ -5,6 +5,7 @@ import ( "math/rand" "os" "path/filepath" + "runtime" "strings" "testing" @@ -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) @@ -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 { @@ -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") + } +}