Files
2026-06-23 05:02:15 +08:00

102 lines
2.3 KiB
Go

package files
import (
"os"
"path/filepath"
"testing"
)
func setupEnvDir(t *testing.T) string {
t.Helper()
dir := t.TempDir()
files := []string{".env", ".env.local", ".env.production"}
for _, f := range files {
if err := os.WriteFile(filepath.Join(dir, f), []byte("KEY=val"), 0o644); err != nil {
t.Fatal(err)
}
}
// A file that should NOT be matched.
if err := os.WriteFile(filepath.Join(dir, "env.txt"), []byte("not env"), 0o644); err != nil {
t.Fatal(err)
}
// Skipped directory.
nm := filepath.Join(dir, "node_modules")
if err := os.MkdirAll(nm, 0o755); err != nil {
t.Fatal(err)
}
if err := os.WriteFile(filepath.Join(nm, ".env"), []byte("skip"), 0o644); err != nil {
t.Fatal(err)
}
return dir
}
func TestFindEnvFiles_MatchesEnvPrefix(t *testing.T) {
dir := setupEnvDir(t)
found, err := FindEnvFiles(dir)
if err != nil {
t.Fatalf("FindEnvFiles: %v", err)
}
if len(found) != 3 {
t.Errorf("found %d files, want 3: %v", len(found), found)
}
}
func TestFindEnvFiles_SkipsNodeModules(t *testing.T) {
dir := setupEnvDir(t)
found, err := FindEnvFiles(dir)
if err != nil {
t.Fatalf("FindEnvFiles: %v", err)
}
for _, f := range found {
if filepath.Dir(f) == filepath.Join(dir, "node_modules") {
t.Errorf("node_modules .env should be skipped, got %s", f)
}
}
}
func TestLockEnvFiles_SetsReadOnly(t *testing.T) {
dir := setupEnvDir(t)
result, err := LockEnvFiles(dir)
if err != nil {
t.Fatalf("LockEnvFiles: %v", err)
}
if result.Total != 3 {
t.Errorf("Total = %d, want 3", result.Total)
}
if len(result.Failed) != 0 {
t.Errorf("unexpected failures: %v", result.Failed)
}
for _, p := range result.OK {
info, err := os.Stat(p)
if err != nil {
t.Fatal(err)
}
if info.Mode()&0o222 != 0 {
t.Errorf("file %s should be read-only, got mode %v", p, info.Mode())
}
}
}
func TestUnlockEnvFiles_SetsReadWrite(t *testing.T) {
dir := setupEnvDir(t)
if _, err := LockEnvFiles(dir); err != nil {
t.Fatalf("LockEnvFiles: %v", err)
}
result, err := UnlockEnvFiles(dir)
if err != nil {
t.Fatalf("UnlockEnvFiles: %v", err)
}
if len(result.Failed) != 0 {
t.Errorf("unexpected failures: %v", result.Failed)
}
for _, p := range result.OK {
info, err := os.Stat(p)
if err != nil {
t.Fatal(err)
}
if info.Mode()&0o200 == 0 {
t.Errorf("file %s should be writable, got mode %v", p, info.Mode())
}
}
}