diff --git a/cmd/room-semgrep/adapter.go b/cmd/room-semgrep/adapter.go index a3dc60d..9710d10 100644 --- a/cmd/room-semgrep/adapter.go +++ b/cmd/room-semgrep/adapter.go @@ -29,10 +29,6 @@ func newAdapter(semgrepCore, config, repositoryRoot string, covered []string) (* return nil, fmt.Errorf("%s must be an absolute path", name) } } - configInfo, err := os.Stat(config) - if err != nil || !configInfo.Mode().IsRegular() { - return nil, errors.New("Semgrep config must be a regular file") - } rootInfo, err := os.Stat(repositoryRoot) if err != nil || !rootInfo.IsDir() { return nil, errors.New("repository root must be a directory") @@ -48,9 +44,9 @@ func newAdapter(semgrepCore, config, repositoryRoot string, covered []string) (* } coveredSet[signal] = true } - configData, err := os.ReadFile(config) + configData, err := readRegularFile(config) if err != nil { - return nil, errors.New("read Semgrep config") + return nil, errors.New("Semgrep config must be a regular file") } if err := validateRuleCoverage(configData, coveredSet); err != nil { return nil, err @@ -235,7 +231,7 @@ func failed(response analyzerResponse, code string) analyzerResponse { } func (a *adapter) configMatches(expected string) bool { - config, err := os.ReadFile(a.config) + config, err := readRegularFile(a.config) if err != nil { return false } diff --git a/cmd/room-semgrep/main_test.go b/cmd/room-semgrep/main_test.go index fb114ff..0d0eeec 100644 --- a/cmd/room-semgrep/main_test.go +++ b/cmd/room-semgrep/main_test.go @@ -272,6 +272,41 @@ func TestAdapterRejectsSymlinkTargets(t *testing.T) { } } +func TestAdapterRejectsSymlinkedConfig(t *testing.T) { + _, repository := workspace(t) + target := writeFile(t, "real-rules.yml", testRules) + linked := filepath.Join(t.TempDir(), "rules.yml") + if err := os.Symlink(target, linked); err != nil { + t.Fatal(err) + } + if _, err := newAdapter("/missing/semgrep", linked, repository, []string{sqlSignal}); err == nil { + t.Fatal("expected symlinked config to be rejected") + } +} + +func TestAdapterFailsClosedWhenConfigBecomesSymlink(t *testing.T) { + _, repository := workspace(t) + if err := os.WriteFile(filepath.Join(repository, "main.go"), []byte("package main\n"), 0o600); err != nil { + t.Fatal(err) + } + config := writeFile(t, "rules.yml", testRules) + adapter, err := newAdapter("/missing/semgrep", config, repository, []string{sqlSignal}) + if err != nil { + t.Fatal(err) + } + request := requestFor(repository, config, []byte("diff --git a/main.go b/main.go\n--- /dev/null\n+++ b/main.go\n@@ -0,0 +1 @@\n+package main\n")) + target := writeFile(t, "same-rules.yml", testRules) + if err := os.Remove(config); err != nil { + t.Fatal(err) + } + if err := os.Symlink(target, config); err != nil { + t.Fatal(err) + } + if response := adapter.analyze(t.Context(), request); response.FailureCode != "config_digest_mismatch" { + t.Fatalf("response = %+v", response) + } +} + func TestAdapterRejectsUnimplementedCoverageAndTraversalDiff(t *testing.T) { _, repository := workspace(t) emptyConfig := writeFile(t, "empty.yml", "rules: []\n") diff --git a/cmd/room-semgrep/snapshot.go b/cmd/room-semgrep/snapshot.go index 3dafbdc..a39d69a 100644 --- a/cmd/room-semgrep/snapshot.go +++ b/cmd/room-semgrep/snapshot.go @@ -6,6 +6,7 @@ import ( "encoding/hex" "encoding/json" "errors" + "fmt" "io" "os" "path/filepath" @@ -46,7 +47,7 @@ func (a *adapter) createSnapshot(request analyzerRequest, artifact diffArtifact) if unix.Fstat(rootFD, &rootStat) != nil || unix.Fstat(workingFD, &workingStat) != nil || rootStat.Dev != workingStat.Dev || rootStat.Ino != workingStat.Ino { return snapshot{}, errors.New("working directory does not match repository root") } - config, err := os.ReadFile(a.config) + config, err := readRegularFile(a.config) if err != nil { return snapshot{}, err } @@ -138,18 +139,32 @@ func readRegularBeneath(rootFD int, path string) ([]byte, error) { if err != nil { return nil, errors.New("diff target cannot be opened safely") } + return readRegularFD(fileFD, path, "diff target") +} + +// readRegularFile reads an operator-owned absolute path without following a +// symlink at its final component. +func readRegularFile(path string) ([]byte, error) { + fileFD, err := unix.Open(path, unix.O_RDONLY|unix.O_CLOEXEC|unix.O_NOFOLLOW|unix.O_NONBLOCK, 0) + if err != nil { + return nil, err + } + return readRegularFD(fileFD, path, path) +} + +func readRegularFD(fileFD int, path, what string) ([]byte, error) { file := os.NewFile(uintptr(fileFD), path) defer file.Close() var before, after unix.Stat_t if err := unix.Fstat(fileFD, &before); err != nil || before.Mode&unix.S_IFMT != unix.S_IFREG || before.Size < 0 || before.Size > 64<<20 { - return nil, errors.New("diff target must be a regular file of at most 64 MiB") + return nil, fmt.Errorf("%s must be a regular file of at most 64 MiB", what) } data, err := io.ReadAll(io.LimitReader(file, before.Size+1)) if err != nil || int64(len(data)) != before.Size { - return nil, errors.New("diff target changed while being read") + return nil, fmt.Errorf("%s changed while being read", what) } if err := unix.Fstat(fileFD, &after); err != nil || before.Ino != after.Ino || before.Size != after.Size || before.Mtim != after.Mtim || before.Ctim != after.Ctim { - return nil, errors.New("diff target changed while being read") + return nil, fmt.Errorf("%s changed while being read", what) } return data, nil }