47 lines
1.1 KiB
Go
47 lines
1.1 KiB
Go
//go:build linux
|
|
|
|
package safeio
|
|
|
|
import (
|
|
"errors"
|
|
"os"
|
|
"path/filepath"
|
|
"testing"
|
|
)
|
|
|
|
func TestReadCanonicalRegularRejectsReplacementDuringRead(t *testing.T) {
|
|
root := t.TempDir()
|
|
path := filepath.Join(root, "schema.sql")
|
|
replacement := filepath.Join(root, "replacement.sql")
|
|
parked := filepath.Join(root, "parked.sql")
|
|
contents := make([]byte, 64<<20)
|
|
if err := os.WriteFile(path, contents, 0o600); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := os.WriteFile(replacement, contents, 0o600); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
|
|
entered := make(chan struct{})
|
|
proceed := make(chan struct{})
|
|
beforeBoundedRead = func() {
|
|
close(entered)
|
|
<-proceed
|
|
}
|
|
t.Cleanup(func() { beforeBoundedRead = nil })
|
|
result := make(chan error, 1)
|
|
go func() { _, err := ReadCanonicalRegular(path, int64(len(contents))); result <- err }()
|
|
<-entered
|
|
if err := os.Rename(path, parked); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
if err := os.Rename(replacement, path); err != nil {
|
|
t.Fatal(err)
|
|
}
|
|
close(proceed)
|
|
err := <-result
|
|
if !errors.Is(err, ErrUnsafeFile) {
|
|
t.Fatalf("replacement during read error = %v, want ErrUnsafeFile", err)
|
|
}
|
|
}
|