feat: Implement platform-specific hosts file path retrieval
- Created a new `platform` package to handle OS-specific hosts file paths. - Added implementations for macOS, Linux, and Windows to retrieve the hosts file path. - Introduced tests for the `GetHostsFilePath` function to ensure correct behavior across platforms. feat: Add DNS resolver for IP lookups - Implemented a `resolver` package for performing DNS lookups. - Created a `DNSResolver` struct with a `LookupIP` method to handle A-record lookups and CNAME resolution. - Added comprehensive tests for various DNS scenarios, including CNAME chains and error handling. chore: Refactor project structure and update dependencies - Restructured the project to follow Go's standard package layout. - Updated `go.mod` to Go 1.20 and added `golang.org/x/net` dependency. - Removed old source files and ensured a clean build with all tests passing.
This commit is contained in:
@@ -0,0 +1,31 @@
|
||||
package hostfile
|
||||
|
||||
import "os"
|
||||
|
||||
// FileSystem abstracts filesystem operations for testability.
|
||||
type FileSystem interface {
|
||||
ReadFile(path string) ([]byte, error)
|
||||
WriteFile(path string, data []byte, perm os.FileMode) error
|
||||
Stat(path string) (os.FileInfo, error)
|
||||
CreateTemp(dir, pattern string) (*os.File, error)
|
||||
Rename(oldpath, newpath string) error
|
||||
Remove(path string) error
|
||||
Chmod(path string, mode os.FileMode) error
|
||||
}
|
||||
|
||||
// OSFileSystem implements FileSystem using real OS calls.
|
||||
type OSFileSystem struct{}
|
||||
|
||||
func (OSFileSystem) ReadFile(path string) ([]byte, error) { return os.ReadFile(path) }
|
||||
func (OSFileSystem) WriteFile(path string, data []byte, perm os.FileMode) error {
|
||||
return os.WriteFile(path, data, perm)
|
||||
}
|
||||
func (OSFileSystem) Stat(path string) (os.FileInfo, error) { return os.Stat(path) }
|
||||
func (OSFileSystem) CreateTemp(dir, pattern string) (*os.File, error) {
|
||||
return os.CreateTemp(dir, pattern)
|
||||
}
|
||||
func (OSFileSystem) Rename(oldpath, newpath string) error { return os.Rename(oldpath, newpath) }
|
||||
func (OSFileSystem) Remove(path string) error { return os.Remove(path) }
|
||||
func (OSFileSystem) Chmod(path string, mode os.FileMode) error {
|
||||
return os.Chmod(path, mode)
|
||||
}
|
||||
@@ -0,0 +1,201 @@
|
||||
package hostfile
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"math/rand"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
const (
|
||||
StartMarker = "# DNSHelper <<-> START CONFIG"
|
||||
EndMarker = "# DNSHelper <->> END CONFIG"
|
||||
)
|
||||
|
||||
// DNSEntry represents a single IP-to-hostname mapping.
|
||||
type DNSEntry struct {
|
||||
IP string
|
||||
Hostname string
|
||||
}
|
||||
|
||||
// String returns the hosts file line representation: "IP\tHostname".
|
||||
func (e DNSEntry) String() string {
|
||||
return e.IP + "\t" + e.Hostname
|
||||
}
|
||||
|
||||
// HostsFile represents the parsed content of a hosts file.
|
||||
type HostsFile struct {
|
||||
Path string
|
||||
OriginalContent []string
|
||||
PrefixContent []string
|
||||
ManagedContent []string
|
||||
PostfixContent []string
|
||||
HasManagedBlock bool
|
||||
}
|
||||
|
||||
// Manager handles hosts file read and atomic write operations.
|
||||
type Manager struct {
|
||||
fs FileSystem
|
||||
}
|
||||
|
||||
// NewManager creates a Manager with the given FileSystem implementation.
|
||||
func NewManager(fs FileSystem) *Manager {
|
||||
return &Manager{fs: fs}
|
||||
}
|
||||
|
||||
// Read reads and parses the hosts file at the given path.
|
||||
// Returns an error if the file cannot be read or the managed block is corrupt.
|
||||
func (m *Manager) Read(path string) (*HostsFile, error) {
|
||||
data, err := m.fs.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("reading hosts file %s: %w", path, err)
|
||||
}
|
||||
|
||||
lines := strings.Split(string(data), "\n")
|
||||
// Remove trailing empty line from Split if file ends with newline.
|
||||
if len(lines) > 0 && lines[len(lines)-1] == "" {
|
||||
lines = lines[:len(lines)-1]
|
||||
}
|
||||
// Trim \r from each line to handle Windows \r\n line endings.
|
||||
// strings.Split on "\n" does not strip \r, unlike bufio.Scanner.
|
||||
for i, line := range lines {
|
||||
lines[i] = strings.TrimRight(line, "\r")
|
||||
}
|
||||
|
||||
prefix, managed, postfix, hasBlock, err := ParseManagedBlock(lines)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return &HostsFile{
|
||||
Path: path,
|
||||
OriginalContent: lines,
|
||||
PrefixContent: prefix,
|
||||
ManagedContent: managed,
|
||||
PostfixContent: postfix,
|
||||
HasManagedBlock: hasBlock,
|
||||
}, nil
|
||||
}
|
||||
|
||||
// Write atomically writes the hosts file content back to disk.
|
||||
// It creates a backup before writing and removes it on success.
|
||||
// Uses the temp-file-and-rename pattern for atomicity (R-002).
|
||||
// backupDir is the directory for the backup file (typically the exe directory).
|
||||
func (m *Manager) Write(hf *HostsFile, backupDir string) error {
|
||||
// 1. Assemble new content.
|
||||
newContent := AssembleContent(hf.PrefixContent, hf.ManagedContent, hf.PostfixContent)
|
||||
|
||||
// 2. Skip write if content is unchanged.
|
||||
if slicesEqual(newContent, hf.OriginalContent) {
|
||||
return nil
|
||||
}
|
||||
|
||||
// 3. Build output bytes (LF line endings).
|
||||
var buf strings.Builder
|
||||
for _, line := range newContent {
|
||||
buf.WriteString(line)
|
||||
buf.WriteString("\n")
|
||||
}
|
||||
output := []byte(buf.String())
|
||||
|
||||
// 4. Get original file permissions so we can restore them on the temp file.
|
||||
info, err := m.fs.Stat(hf.Path)
|
||||
if err != nil {
|
||||
return fmt.Errorf("checking hosts file permissions: %w", err)
|
||||
}
|
||||
originalMode := info.Mode()
|
||||
|
||||
// 5. Create backup (copy of current hosts file in backupDir).
|
||||
backupPath, err := m.createBackup(hf.Path, backupDir)
|
||||
if err != nil {
|
||||
return fmt.Errorf("creating backup: %w", err)
|
||||
}
|
||||
|
||||
// 6. Write to a temp file in the SAME directory as the hosts file.
|
||||
// This is required so the subsequent rename stays on the same filesystem
|
||||
// partition (avoids EXDEV on Unix).
|
||||
success := false
|
||||
tempFile, err := m.fs.CreateTemp(filepath.Dir(hf.Path), ".dns-helper-tmp-*")
|
||||
if err != nil {
|
||||
m.fs.Remove(backupPath) //nolint:errcheck
|
||||
return fmt.Errorf("creating temp file: %w", err)
|
||||
}
|
||||
tempPath := tempFile.Name()
|
||||
defer func() {
|
||||
if !success {
|
||||
m.fs.Remove(tempPath) //nolint:errcheck
|
||||
m.fs.Remove(backupPath) //nolint:errcheck
|
||||
}
|
||||
}()
|
||||
|
||||
if _, err := tempFile.Write(output); err != nil {
|
||||
tempFile.Close() //nolint:errcheck
|
||||
return fmt.Errorf("writing temp file: %w", err)
|
||||
}
|
||||
// Sync before close for durability (data survives power loss).
|
||||
if err := tempFile.Sync(); err != nil {
|
||||
tempFile.Close() //nolint:errcheck
|
||||
return fmt.Errorf("syncing temp file: %w", err)
|
||||
}
|
||||
// Close before rename — required on Windows.
|
||||
tempFile.Close() //nolint:errcheck
|
||||
|
||||
// 7. Set permissions on temp file to match original.
|
||||
if err := m.fs.Chmod(tempPath, originalMode); err != nil {
|
||||
return fmt.Errorf("setting temp file permissions: %w", err)
|
||||
}
|
||||
|
||||
// 8. Atomic rename (with retry on Windows sharing violations).
|
||||
if err := m.renameWithRetry(tempPath, hf.Path); err != nil {
|
||||
return fmt.Errorf("renaming temp file to hosts file: %w", err)
|
||||
}
|
||||
|
||||
// 9. Success — delete the backup (original was safely replaced).
|
||||
success = true
|
||||
m.fs.Remove(backupPath) //nolint:errcheck
|
||||
return nil
|
||||
}
|
||||
|
||||
// createBackup copies the hosts file to <backupDir>/hosts.bak.<YYYYMMDD>-<4hex>.
|
||||
func (m *Manager) createBackup(hostsPath, backupDir string) (string, error) {
|
||||
data, err := m.fs.ReadFile(hostsPath)
|
||||
if err != nil {
|
||||
return "", fmt.Errorf("reading hosts file for backup: %w", err)
|
||||
}
|
||||
|
||||
timestamp := time.Now().Format("20060102")
|
||||
suffix := fmt.Sprintf("%04x", rand.Uint32()&0xFFFF)
|
||||
backupName := fmt.Sprintf("hosts.bak.%s-%s", timestamp, suffix)
|
||||
backupPath := filepath.Join(backupDir, backupName)
|
||||
|
||||
if err := m.fs.WriteFile(backupPath, data, 0600); err != nil {
|
||||
return "", fmt.Errorf("writing backup file %s: %w", backupPath, err)
|
||||
}
|
||||
return backupPath, nil
|
||||
}
|
||||
|
||||
// renameWithRetry calls Rename and retries once on failure (for Windows
|
||||
// sharing violations where the hosts file may be briefly locked).
|
||||
func (m *Manager) renameWithRetry(src, dst string) error {
|
||||
err := m.fs.Rename(src, dst)
|
||||
if err == nil {
|
||||
return nil
|
||||
}
|
||||
// Single retry after a short delay.
|
||||
time.Sleep(200 * time.Millisecond)
|
||||
return m.fs.Rename(src, dst)
|
||||
}
|
||||
|
||||
// slicesEqual reports whether a and b contain the same strings in the same order.
|
||||
func slicesEqual(a, b []string) bool {
|
||||
if len(a) != len(b) {
|
||||
return false
|
||||
}
|
||||
for i := range a {
|
||||
if a[i] != b[i] {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return true
|
||||
}
|
||||
@@ -0,0 +1,369 @@
|
||||
package hostfile_test
|
||||
|
||||
import (
|
||||
"errors"
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"dns-helper/hostfile"
|
||||
)
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// fakeFileSystem — implements hostfile.FileSystem for unit tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
type fakeFileInfo struct {
|
||||
mode os.FileMode
|
||||
}
|
||||
|
||||
func (f fakeFileInfo) Name() string { return "" }
|
||||
func (f fakeFileInfo) Size() int64 { return 0 }
|
||||
func (f fakeFileInfo) Mode() os.FileMode { return f.mode }
|
||||
func (f fakeFileInfo) ModTime() time.Time { return time.Time{} }
|
||||
func (f fakeFileInfo) IsDir() bool { return false }
|
||||
func (f fakeFileInfo) Sys() any { return nil }
|
||||
|
||||
type writeCall struct {
|
||||
path string
|
||||
data []byte
|
||||
perm os.FileMode
|
||||
}
|
||||
|
||||
type renameCall struct {
|
||||
src string
|
||||
dst string
|
||||
}
|
||||
|
||||
type fakeFileSystem struct {
|
||||
readFiles map[string][]byte
|
||||
readErrors map[string]error
|
||||
statMode os.FileMode
|
||||
statErr error
|
||||
renameErr error
|
||||
removeErr error
|
||||
|
||||
// tracking
|
||||
createTempDir string
|
||||
lastTempFile *os.File
|
||||
writeCalls []writeCall
|
||||
renameCalls []renameCall
|
||||
removedPaths []string
|
||||
chmodPath string
|
||||
chmodMode os.FileMode
|
||||
|
||||
// real temp dir on disk for CreateTemp to use
|
||||
realTempDir string
|
||||
}
|
||||
|
||||
func newFakeFS(t *testing.T) *fakeFileSystem {
|
||||
t.Helper()
|
||||
return &fakeFileSystem{
|
||||
readFiles: make(map[string][]byte),
|
||||
readErrors: make(map[string]error),
|
||||
realTempDir: t.TempDir(),
|
||||
statMode: 0644,
|
||||
}
|
||||
}
|
||||
|
||||
func (f *fakeFileSystem) ReadFile(path string) ([]byte, error) {
|
||||
if err, ok := f.readErrors[path]; ok {
|
||||
return nil, err
|
||||
}
|
||||
if data, ok := f.readFiles[path]; ok {
|
||||
return data, nil
|
||||
}
|
||||
return nil, fmt.Errorf("file not found in fake fs: %s", path)
|
||||
}
|
||||
|
||||
func (f *fakeFileSystem) WriteFile(path string, data []byte, perm os.FileMode) error {
|
||||
f.writeCalls = append(f.writeCalls, writeCall{path: path, data: data, perm: perm})
|
||||
return nil
|
||||
}
|
||||
|
||||
func (f *fakeFileSystem) Stat(path string) (os.FileInfo, error) {
|
||||
if f.statErr != nil {
|
||||
return nil, f.statErr
|
||||
}
|
||||
return fakeFileInfo{mode: f.statMode}, nil
|
||||
}
|
||||
|
||||
func (f *fakeFileSystem) CreateTemp(dir, pattern string) (*os.File, error) {
|
||||
f.createTempDir = dir
|
||||
tf, err := os.CreateTemp(f.realTempDir, pattern)
|
||||
if err == nil {
|
||||
f.lastTempFile = tf
|
||||
}
|
||||
return tf, err
|
||||
}
|
||||
|
||||
func (f *fakeFileSystem) Rename(oldpath, newpath string) error {
|
||||
f.renameCalls = append(f.renameCalls, renameCall{src: oldpath, dst: newpath})
|
||||
return f.renameErr
|
||||
}
|
||||
|
||||
func (f *fakeFileSystem) Remove(path string) error {
|
||||
f.removedPaths = append(f.removedPaths, path)
|
||||
return f.removeErr
|
||||
}
|
||||
|
||||
func (f *fakeFileSystem) Chmod(path string, mode os.FileMode) error {
|
||||
f.chmodPath = path
|
||||
f.chmodMode = mode
|
||||
return nil
|
||||
}
|
||||
|
||||
// containsRemoved returns true if path was passed to Remove.
|
||||
func (f *fakeFileSystem) containsRemoved(path string) bool {
|
||||
for _, p := range f.removedPaths {
|
||||
if p == path {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Read tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestRead_ParsesFileCorrectly(t *testing.T) {
|
||||
fs := newFakeFS(t)
|
||||
content := strings.Join([]string{
|
||||
"127.0.0.1 localhost",
|
||||
hostfile.StartMarker,
|
||||
"1.1.1.1\thost1",
|
||||
hostfile.EndMarker,
|
||||
"# end",
|
||||
}, "\n") + "\n"
|
||||
hostsPath := "/fake/hosts"
|
||||
fs.readFiles[hostsPath] = []byte(content)
|
||||
|
||||
mgr := hostfile.NewManager(fs)
|
||||
hf, err := mgr.Read(hostsPath)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if hf.Path != hostsPath {
|
||||
t.Errorf("expected Path=%q, got %q", hostsPath, hf.Path)
|
||||
}
|
||||
if !hf.HasManagedBlock {
|
||||
t.Error("expected HasManagedBlock=true")
|
||||
}
|
||||
if len(hf.ManagedContent) != 1 || hf.ManagedContent[0] != "1.1.1.1\thost1" {
|
||||
t.Errorf("unexpected ManagedContent: %v", hf.ManagedContent)
|
||||
}
|
||||
if len(hf.PrefixContent) != 1 {
|
||||
t.Errorf("expected 1 prefix line, got %v", hf.PrefixContent)
|
||||
}
|
||||
if len(hf.PostfixContent) != 1 {
|
||||
t.Errorf("expected 1 postfix line, got %v", hf.PostfixContent)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRead_HandlesMissingFile(t *testing.T) {
|
||||
fs := newFakeFS(t)
|
||||
hostsPath := "/fake/hosts"
|
||||
fs.readErrors[hostsPath] = errors.New("no such file")
|
||||
|
||||
mgr := hostfile.NewManager(fs)
|
||||
_, err := mgr.Read(hostsPath)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for missing file")
|
||||
}
|
||||
if !strings.Contains(err.Error(), hostsPath) {
|
||||
t.Errorf("error should mention the path, got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// Write tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
// buildTestHostsFile builds a HostsFile where OriginalContent differs from
|
||||
// the assembled content (so Write proceeds).
|
||||
func buildTestHostsFile(hostsPath string) *hostfile.HostsFile {
|
||||
prefix := []string{"127.0.0.1 localhost"}
|
||||
managed := []string{"1.1.1.1\thost1"}
|
||||
postfix := []string{"# end"}
|
||||
// OriginalContent is different (e.g. empty managed block originally)
|
||||
original := prefix
|
||||
return &hostfile.HostsFile{
|
||||
Path: hostsPath,
|
||||
OriginalContent: original,
|
||||
PrefixContent: prefix,
|
||||
ManagedContent: managed,
|
||||
PostfixContent: postfix,
|
||||
HasManagedBlock: false,
|
||||
}
|
||||
}
|
||||
|
||||
func TestWrite_CreatesBackupBeforeWriting(t *testing.T) {
|
||||
fs := newFakeFS(t)
|
||||
hostsPath := filepath.Join(t.TempDir(), "hosts")
|
||||
backupDir := t.TempDir()
|
||||
hostsContent := "127.0.0.1 localhost\n"
|
||||
fs.readFiles[hostsPath] = []byte(hostsContent)
|
||||
|
||||
mgr := hostfile.NewManager(fs)
|
||||
hf := buildTestHostsFile(hostsPath)
|
||||
if err := mgr.Write(hf, backupDir); err != nil {
|
||||
t.Fatalf("Write failed: %v", err)
|
||||
}
|
||||
|
||||
// Verify backup was created via WriteFile
|
||||
found := false
|
||||
for _, wc := range fs.writeCalls {
|
||||
if strings.HasPrefix(wc.path, backupDir) && strings.Contains(wc.path, "hosts.bak.") {
|
||||
found = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Errorf("expected backup WriteFile call in backupDir %q, write calls: %v", backupDir, fs.writeCalls)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWrite_UsesTempFileInSameDirectory(t *testing.T) {
|
||||
fs := newFakeFS(t)
|
||||
hostsDir := t.TempDir()
|
||||
hostsPath := filepath.Join(hostsDir, "hosts")
|
||||
backupDir := t.TempDir()
|
||||
fs.readFiles[hostsPath] = []byte("127.0.0.1 localhost\n")
|
||||
|
||||
mgr := hostfile.NewManager(fs)
|
||||
hf := buildTestHostsFile(hostsPath)
|
||||
if err := mgr.Write(hf, backupDir); err != nil {
|
||||
t.Fatalf("Write failed: %v", err)
|
||||
}
|
||||
|
||||
if fs.createTempDir != hostsDir {
|
||||
t.Errorf("expected CreateTemp dir=%q, got %q", hostsDir, fs.createTempDir)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWrite_RenamesTempOverHostsFile(t *testing.T) {
|
||||
fs := newFakeFS(t)
|
||||
hostsPath := filepath.Join(t.TempDir(), "hosts")
|
||||
backupDir := t.TempDir()
|
||||
fs.readFiles[hostsPath] = []byte("127.0.0.1 localhost\n")
|
||||
|
||||
mgr := hostfile.NewManager(fs)
|
||||
hf := buildTestHostsFile(hostsPath)
|
||||
if err := mgr.Write(hf, backupDir); err != nil {
|
||||
t.Fatalf("Write failed: %v", err)
|
||||
}
|
||||
|
||||
if len(fs.renameCalls) == 0 {
|
||||
t.Fatal("expected Rename to be called")
|
||||
}
|
||||
lastRename := fs.renameCalls[len(fs.renameCalls)-1]
|
||||
if lastRename.dst != hostsPath {
|
||||
t.Errorf("expected Rename dst=%q, got %q", hostsPath, lastRename.dst)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWrite_PreservesOriginalPermissions(t *testing.T) {
|
||||
fs := newFakeFS(t)
|
||||
fs.statMode = 0640
|
||||
hostsPath := filepath.Join(t.TempDir(), "hosts")
|
||||
backupDir := t.TempDir()
|
||||
fs.readFiles[hostsPath] = []byte("127.0.0.1 localhost\n")
|
||||
|
||||
mgr := hostfile.NewManager(fs)
|
||||
hf := buildTestHostsFile(hostsPath)
|
||||
if err := mgr.Write(hf, backupDir); err != nil {
|
||||
t.Fatalf("Write failed: %v", err)
|
||||
}
|
||||
|
||||
if fs.chmodMode != 0640 {
|
||||
t.Errorf("expected Chmod mode=0640, got %v", fs.chmodMode)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWrite_DeletesBackupOnSuccess(t *testing.T) {
|
||||
fs := newFakeFS(t)
|
||||
hostsPath := filepath.Join(t.TempDir(), "hosts")
|
||||
backupDir := t.TempDir()
|
||||
fs.readFiles[hostsPath] = []byte("127.0.0.1 localhost\n")
|
||||
|
||||
mgr := hostfile.NewManager(fs)
|
||||
hf := buildTestHostsFile(hostsPath)
|
||||
if err := mgr.Write(hf, backupDir); err != nil {
|
||||
t.Fatalf("Write failed: %v", err)
|
||||
}
|
||||
|
||||
// Find the backup path from writeCalls
|
||||
var backupPath string
|
||||
for _, wc := range fs.writeCalls {
|
||||
if strings.HasPrefix(wc.path, backupDir) && strings.Contains(wc.path, "hosts.bak.") {
|
||||
backupPath = wc.path
|
||||
break
|
||||
}
|
||||
}
|
||||
if backupPath == "" {
|
||||
t.Fatal("backup not found in writeCalls")
|
||||
}
|
||||
if !fs.containsRemoved(backupPath) {
|
||||
t.Errorf("expected backup %q to be removed on success, removed paths: %v", backupPath, fs.removedPaths)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWrite_CleansUpTempOnFailure(t *testing.T) {
|
||||
fs := newFakeFS(t)
|
||||
fs.renameErr = errors.New("rename failed")
|
||||
hostsPath := filepath.Join(t.TempDir(), "hosts")
|
||||
backupDir := t.TempDir()
|
||||
fs.readFiles[hostsPath] = []byte("127.0.0.1 localhost\n")
|
||||
|
||||
mgr := hostfile.NewManager(fs)
|
||||
hf := buildTestHostsFile(hostsPath)
|
||||
err := mgr.Write(hf, backupDir)
|
||||
if err == nil {
|
||||
t.Fatal("expected Write to fail when Rename fails")
|
||||
}
|
||||
|
||||
if fs.lastTempFile == nil {
|
||||
t.Fatal("expected CreateTemp to have been called")
|
||||
}
|
||||
tempPath := fs.lastTempFile.Name()
|
||||
if !fs.containsRemoved(tempPath) {
|
||||
t.Errorf("expected temp file %q to be removed on failure, removed: %v", tempPath, fs.removedPaths)
|
||||
}
|
||||
}
|
||||
|
||||
func TestWrite_SkipsIfContentUnchanged(t *testing.T) {
|
||||
fs := newFakeFS(t)
|
||||
hostsPath := filepath.Join(t.TempDir(), "hosts")
|
||||
backupDir := t.TempDir()
|
||||
|
||||
prefix := []string{"127.0.0.1 localhost"}
|
||||
managed := []string{"1.1.1.1\thost1"}
|
||||
postfix := []string{"# end"}
|
||||
// OriginalContent matches assembled content exactly
|
||||
assembled := hostfile.AssembleContent(prefix, managed, postfix)
|
||||
|
||||
hf := &hostfile.HostsFile{
|
||||
Path: hostsPath,
|
||||
OriginalContent: assembled,
|
||||
PrefixContent: prefix,
|
||||
ManagedContent: managed,
|
||||
PostfixContent: postfix,
|
||||
HasManagedBlock: true,
|
||||
}
|
||||
|
||||
mgr := hostfile.NewManager(fs)
|
||||
if err := mgr.Write(hf, backupDir); err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
|
||||
if fs.lastTempFile != nil {
|
||||
t.Error("expected no CreateTemp call when content is unchanged")
|
||||
}
|
||||
if len(fs.writeCalls) != 0 {
|
||||
t.Errorf("expected no WriteFile calls, got: %v", fs.writeCalls)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,148 @@
|
||||
package hostfile
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ParseManagedBlock splits raw file lines into prefix, managed, and postfix sections.
|
||||
// Returns an error if the managed block markers are corrupt.
|
||||
//
|
||||
// Uses -1 as sentinel for "not found" because the start marker can appear at line 0.
|
||||
// The current code's zero-index bug (treating index 0 as "not found") is avoided here.
|
||||
func ParseManagedBlock(lines []string) (prefix, managed, postfix []string, hasManagedBlock bool, err error) {
|
||||
startIdx := -1 // sentinel: not found
|
||||
endIdx := -1 // sentinel: not found
|
||||
startCount := 0
|
||||
endCount := 0
|
||||
|
||||
for i, line := range lines {
|
||||
if strings.Contains(line, "DNSHelper <<-> START CONFIG") {
|
||||
startCount++
|
||||
if startIdx == -1 {
|
||||
startIdx = i
|
||||
}
|
||||
}
|
||||
if strings.Contains(line, "DNSHelper <->> END CONFIG") {
|
||||
endCount++
|
||||
if endIdx == -1 {
|
||||
endIdx = i
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Corruption detection (R-008): scan ALL lines first, then evaluate.
|
||||
if startCount > 1 || endCount > 1 {
|
||||
return nil, nil, nil, false, fmt.Errorf("managed block is corrupt: duplicate start/end markers found")
|
||||
}
|
||||
if startIdx >= 0 && endIdx < 0 {
|
||||
return nil, nil, nil, false, fmt.Errorf("managed block is corrupt: start marker found at line %d but no end marker", startIdx+1)
|
||||
}
|
||||
if endIdx >= 0 && startIdx < 0 {
|
||||
return nil, nil, nil, false, fmt.Errorf("managed block is corrupt: end marker found at line %d but no start marker", endIdx+1)
|
||||
}
|
||||
if startIdx >= 0 && endIdx >= 0 && endIdx < startIdx {
|
||||
return nil, nil, nil, false, fmt.Errorf("managed block is corrupt: end marker (line %d) appears before start marker (line %d)", endIdx+1, startIdx+1)
|
||||
}
|
||||
|
||||
if startIdx < 0 && endIdx < 0 {
|
||||
// No managed block — all lines are prefix.
|
||||
return lines, nil, nil, false, nil
|
||||
}
|
||||
|
||||
// Valid managed block.
|
||||
prefix = lines[:startIdx]
|
||||
managed = lines[startIdx+1 : endIdx]
|
||||
postfix = lines[endIdx+1:]
|
||||
return prefix, managed, postfix, true, nil
|
||||
}
|
||||
|
||||
// AddEntries removes existing entries for the same hostnames as the new entries,
|
||||
// then appends the new entries. Returns deduplicated managed content lines.
|
||||
func AddEntries(existing []string, entries []DNSEntry) []string {
|
||||
// Collect hostnames being added.
|
||||
hostnames := make(map[string]bool)
|
||||
for _, e := range entries {
|
||||
hostnames[e.Hostname] = true
|
||||
}
|
||||
|
||||
// Keep existing lines that don't match any new hostname.
|
||||
var result []string
|
||||
for _, line := range existing {
|
||||
keep := true
|
||||
for h := range hostnames {
|
||||
if strings.Contains(line, h) {
|
||||
keep = false
|
||||
break
|
||||
}
|
||||
}
|
||||
if keep {
|
||||
result = append(result, line)
|
||||
}
|
||||
}
|
||||
|
||||
// Append new entries.
|
||||
for _, e := range entries {
|
||||
result = append(result, e.String())
|
||||
}
|
||||
|
||||
// Deduplicate preserving order.
|
||||
return deduplicateLines(result)
|
||||
}
|
||||
|
||||
// RemoveByHostname removes all managed entries matching any of the specified
|
||||
// hostnames. Returns the remaining content.
|
||||
//
|
||||
// The correct algorithm iterates lines on the outer loop (not hosts), which
|
||||
// avoids the duplication bug in the legacy removeHostsFromExistingContent
|
||||
// function (dataprep.go).
|
||||
func RemoveByHostname(existing []string, hostnames []string) []string {
|
||||
var result []string
|
||||
for _, line := range existing {
|
||||
shouldRemove := false
|
||||
for _, host := range hostnames {
|
||||
if strings.Contains(line, host) {
|
||||
shouldRemove = true
|
||||
break
|
||||
}
|
||||
}
|
||||
if !shouldRemove {
|
||||
result = append(result, line)
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
|
||||
// RemoveAll returns nil, indicating all managed entries should be removed.
|
||||
// When AssembleContent receives nil/empty managed content, it omits the
|
||||
// managed block entirely (markers included).
|
||||
func RemoveAll() []string {
|
||||
return nil
|
||||
}
|
||||
|
||||
// AssembleContent builds the full file content from sections.
|
||||
// If managed is empty, the managed block (markers included) is omitted entirely.
|
||||
func AssembleContent(prefix, managed, postfix []string) []string {
|
||||
var result []string
|
||||
result = append(result, prefix...)
|
||||
if len(managed) > 0 {
|
||||
result = append(result, StartMarker)
|
||||
result = append(result, managed...)
|
||||
result = append(result, EndMarker)
|
||||
}
|
||||
result = append(result, postfix...)
|
||||
return result
|
||||
}
|
||||
|
||||
// deduplicateLines removes duplicate lines preserving order.
|
||||
func deduplicateLines(lines []string) []string {
|
||||
seen := make(map[string]bool)
|
||||
var result []string
|
||||
for _, line := range lines {
|
||||
if !seen[line] {
|
||||
seen[line] = true
|
||||
result = append(result, line)
|
||||
}
|
||||
}
|
||||
return result
|
||||
}
|
||||
@@ -0,0 +1,360 @@
|
||||
package hostfile_test
|
||||
|
||||
import (
|
||||
"strings"
|
||||
"testing"
|
||||
|
||||
"dns-helper/hostfile"
|
||||
)
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// ParseManagedBlock tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestParseManagedBlock_NoManagedBlock(t *testing.T) {
|
||||
lines := []string{
|
||||
"127.0.0.1 localhost",
|
||||
"# some comment",
|
||||
"1.2.3.4 example.com",
|
||||
}
|
||||
prefix, managed, postfix, has, err := hostfile.ParseManagedBlock(lines)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if has {
|
||||
t.Fatal("expected HasManagedBlock=false")
|
||||
}
|
||||
if len(managed) != 0 {
|
||||
t.Errorf("expected empty managed, got %v", managed)
|
||||
}
|
||||
if len(postfix) != 0 {
|
||||
t.Errorf("expected empty postfix, got %v", postfix)
|
||||
}
|
||||
if len(prefix) != len(lines) {
|
||||
t.Errorf("expected all lines in prefix, got %v", prefix)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseManagedBlock_ValidBlock(t *testing.T) {
|
||||
lines := []string{
|
||||
"127.0.0.1 localhost",
|
||||
hostfile.StartMarker,
|
||||
"1.1.1.1\thost1",
|
||||
"2.2.2.2\thost2",
|
||||
hostfile.EndMarker,
|
||||
"# trailing comment",
|
||||
}
|
||||
prefix, managed, postfix, has, err := hostfile.ParseManagedBlock(lines)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !has {
|
||||
t.Fatal("expected HasManagedBlock=true")
|
||||
}
|
||||
if len(prefix) != 1 || prefix[0] != "127.0.0.1 localhost" {
|
||||
t.Errorf("unexpected prefix: %v", prefix)
|
||||
}
|
||||
if len(managed) != 2 {
|
||||
t.Errorf("expected 2 managed lines, got %v", managed)
|
||||
}
|
||||
if len(postfix) != 1 || postfix[0] != "# trailing comment" {
|
||||
t.Errorf("unexpected postfix: %v", postfix)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseManagedBlock_StartMarkerAtLine0(t *testing.T) {
|
||||
lines := []string{
|
||||
hostfile.StartMarker,
|
||||
"1.1.1.1\thost1",
|
||||
hostfile.EndMarker,
|
||||
}
|
||||
prefix, managed, postfix, has, err := hostfile.ParseManagedBlock(lines)
|
||||
if err != nil {
|
||||
t.Fatalf("unexpected error: %v", err)
|
||||
}
|
||||
if !has {
|
||||
t.Fatal("expected HasManagedBlock=true")
|
||||
}
|
||||
if len(prefix) != 0 {
|
||||
t.Errorf("expected empty prefix when start marker is at line 0, got %v", prefix)
|
||||
}
|
||||
if len(managed) != 1 || managed[0] != "1.1.1.1\thost1" {
|
||||
t.Errorf("unexpected managed: %v", managed)
|
||||
}
|
||||
if len(postfix) != 0 {
|
||||
t.Errorf("expected empty postfix, got %v", postfix)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseManagedBlock_CorruptStartWithoutEnd(t *testing.T) {
|
||||
lines := []string{
|
||||
"127.0.0.1 localhost",
|
||||
hostfile.StartMarker,
|
||||
"1.1.1.1\thost1",
|
||||
}
|
||||
_, _, _, _, err := hostfile.ParseManagedBlock(lines)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for start marker without end marker")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "start marker found") {
|
||||
t.Errorf("error should mention 'start marker found', got: %v", err)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "no end marker") {
|
||||
t.Errorf("error should mention 'no end marker', got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseManagedBlock_CorruptEndWithoutStart(t *testing.T) {
|
||||
lines := []string{
|
||||
"127.0.0.1 localhost",
|
||||
"1.1.1.1\thost1",
|
||||
hostfile.EndMarker,
|
||||
}
|
||||
_, _, _, _, err := hostfile.ParseManagedBlock(lines)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for end marker without start marker")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "end marker found") {
|
||||
t.Errorf("error should mention 'end marker found', got: %v", err)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "no start marker") {
|
||||
t.Errorf("error should mention 'no start marker', got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseManagedBlock_CorruptEndBeforeStart(t *testing.T) {
|
||||
lines := []string{
|
||||
"127.0.0.1 localhost",
|
||||
hostfile.EndMarker,
|
||||
"1.1.1.1\thost1",
|
||||
hostfile.StartMarker,
|
||||
}
|
||||
_, _, _, _, err := hostfile.ParseManagedBlock(lines)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for end marker before start marker")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "end marker") {
|
||||
t.Errorf("error should mention 'end marker', got: %v", err)
|
||||
}
|
||||
if !strings.Contains(err.Error(), "before start marker") {
|
||||
t.Errorf("error should mention 'before start marker', got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestParseManagedBlock_CorruptDuplicateMarkers(t *testing.T) {
|
||||
lines := []string{
|
||||
hostfile.StartMarker,
|
||||
"1.1.1.1\thost1",
|
||||
hostfile.EndMarker,
|
||||
hostfile.StartMarker, // duplicate
|
||||
"2.2.2.2\thost2",
|
||||
hostfile.EndMarker,
|
||||
}
|
||||
_, _, _, _, err := hostfile.ParseManagedBlock(lines)
|
||||
if err == nil {
|
||||
t.Fatal("expected error for duplicate markers")
|
||||
}
|
||||
if !strings.Contains(err.Error(), "duplicate") {
|
||||
t.Errorf("error should mention 'duplicate', got: %v", err)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// AddEntries tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestAddEntries_AddsNewEntries(t *testing.T) {
|
||||
existing := []string{"1.1.1.1\thost1"}
|
||||
entries := []hostfile.DNSEntry{{IP: "2.2.2.2", Hostname: "host2"}}
|
||||
result := hostfile.AddEntries(existing, entries)
|
||||
if len(result) != 2 {
|
||||
t.Fatalf("expected 2 entries, got %d: %v", len(result), result)
|
||||
}
|
||||
if result[0] != "1.1.1.1\thost1" {
|
||||
t.Errorf("unexpected entry[0]: %q", result[0])
|
||||
}
|
||||
if result[1] != "2.2.2.2\thost2" {
|
||||
t.Errorf("unexpected entry[1]: %q", result[1])
|
||||
}
|
||||
}
|
||||
|
||||
func TestAddEntries_ReplacesExistingHostname(t *testing.T) {
|
||||
existing := []string{"1.1.1.1\thost1", "2.2.2.2\thost2"}
|
||||
// Replace host1 with a new IP
|
||||
entries := []hostfile.DNSEntry{{IP: "9.9.9.9", Hostname: "host1"}}
|
||||
result := hostfile.AddEntries(existing, entries)
|
||||
// host1 should be replaced, host2 preserved
|
||||
if len(result) != 2 {
|
||||
t.Fatalf("expected 2 entries, got %d: %v", len(result), result)
|
||||
}
|
||||
found := false
|
||||
for _, line := range result {
|
||||
if strings.Contains(line, "1.1.1.1") && strings.Contains(line, "host1") {
|
||||
t.Errorf("old entry for host1 should have been replaced: %v", result)
|
||||
}
|
||||
if strings.Contains(line, "9.9.9.9") && strings.Contains(line, "host1") {
|
||||
found = true
|
||||
}
|
||||
}
|
||||
if !found {
|
||||
t.Errorf("new entry for host1 not found: %v", result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAddEntries_Deduplicates(t *testing.T) {
|
||||
existing := []string{"1.1.1.1\thost1"}
|
||||
entries := []hostfile.DNSEntry{
|
||||
{IP: "2.2.2.2", Hostname: "host2"},
|
||||
{IP: "2.2.2.2", Hostname: "host2"}, // duplicate
|
||||
}
|
||||
result := hostfile.AddEntries(existing, entries)
|
||||
count := 0
|
||||
for _, line := range result {
|
||||
if line == "2.2.2.2\thost2" {
|
||||
count++
|
||||
}
|
||||
}
|
||||
if count != 1 {
|
||||
t.Errorf("expected exactly 1 host2 entry after dedup, got %d: %v", count, result)
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// AssembleContent tests
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestAssembleContent_WithManagedBlock(t *testing.T) {
|
||||
prefix := []string{"127.0.0.1 localhost"}
|
||||
managed := []string{"1.1.1.1\thost1"}
|
||||
postfix := []string{"# end"}
|
||||
result := hostfile.AssembleContent(prefix, managed, postfix)
|
||||
expected := []string{
|
||||
"127.0.0.1 localhost",
|
||||
hostfile.StartMarker,
|
||||
"1.1.1.1\thost1",
|
||||
hostfile.EndMarker,
|
||||
"# end",
|
||||
}
|
||||
if len(result) != len(expected) {
|
||||
t.Fatalf("expected %d lines, got %d: %v", len(expected), len(result), result)
|
||||
}
|
||||
for i, line := range expected {
|
||||
if result[i] != line {
|
||||
t.Errorf("line %d: expected %q, got %q", i, line, result[i])
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func TestAssembleContent_EmptyManagedOmitsBlock(t *testing.T) {
|
||||
prefix := []string{"127.0.0.1 localhost"}
|
||||
postfix := []string{"# end"}
|
||||
result := hostfile.AssembleContent(prefix, nil, postfix)
|
||||
if len(result) != 2 {
|
||||
t.Fatalf("expected 2 lines (no managed block), got %d: %v", len(result), result)
|
||||
}
|
||||
for _, line := range result {
|
||||
if strings.Contains(line, "DNSHelper") {
|
||||
t.Errorf("managed block markers should be absent when managed is empty: %v", result)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// RemoveByHostname tests (T018)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestRemoveByHostname_SingleHost(t *testing.T) {
|
||||
existing := []string{
|
||||
"1.1.1.1\thost1",
|
||||
"2.2.2.2\thost2",
|
||||
}
|
||||
result := hostfile.RemoveByHostname(existing, []string{"host1"})
|
||||
if len(result) != 1 {
|
||||
t.Fatalf("expected 1 entry, got %d: %v", len(result), result)
|
||||
}
|
||||
if result[0] != "2.2.2.2\thost2" {
|
||||
t.Errorf("expected host2 entry to remain, got %q", result[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoveByHostname_MultipleHosts(t *testing.T) {
|
||||
existing := []string{
|
||||
"1.1.1.1\thost1",
|
||||
"2.2.2.2\thost2",
|
||||
"3.3.3.3\thost3",
|
||||
}
|
||||
result := hostfile.RemoveByHostname(existing, []string{"host1", "host2"})
|
||||
if len(result) != 1 {
|
||||
t.Fatalf("expected 1 entry, got %d: %v", len(result), result)
|
||||
}
|
||||
if result[0] != "3.3.3.3\thost3" {
|
||||
t.Errorf("expected host3 entry to remain, got %q", result[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoveByHostname_HostNotFound(t *testing.T) {
|
||||
existing := []string{
|
||||
"1.1.1.1\thost1",
|
||||
}
|
||||
result := hostfile.RemoveByHostname(existing, []string{"host2"})
|
||||
if len(result) != 1 {
|
||||
t.Fatalf("expected content unchanged (1 entry), got %d: %v", len(result), result)
|
||||
}
|
||||
if result[0] != "1.1.1.1\thost1" {
|
||||
t.Errorf("expected host1 entry unchanged, got %q", result[0])
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoveByHostname_AllHostsRemoved(t *testing.T) {
|
||||
existing := []string{
|
||||
"1.1.1.1\thost1",
|
||||
}
|
||||
result := hostfile.RemoveByHostname(existing, []string{"host1"})
|
||||
if len(result) != 0 {
|
||||
t.Errorf("expected empty slice, got %d entries: %v", len(result), result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoveByHostname_RegressionNoDuplicates(t *testing.T) {
|
||||
// This is the R-004 regression test. The legacy code's loop-order bug
|
||||
// produced duplicates when removing multiple hosts at once.
|
||||
existing := []string{
|
||||
"1.1.1.1\thost1",
|
||||
"2.2.2.2\thost2",
|
||||
"3.3.3.3\thost3",
|
||||
}
|
||||
result := hostfile.RemoveByHostname(existing, []string{"host1", "host2"})
|
||||
if len(result) != 1 {
|
||||
t.Fatalf("expected exactly 1 entry after multi-host removal, got %d: %v", len(result), result)
|
||||
}
|
||||
if result[0] != "3.3.3.3\thost3" {
|
||||
t.Errorf("expected only host3 to remain, got %q", result[0])
|
||||
}
|
||||
}
|
||||
|
||||
// ---------------------------------------------------------------------------
|
||||
// RemoveAll tests (T018)
|
||||
// ---------------------------------------------------------------------------
|
||||
|
||||
func TestRemoveAll_ReturnEmptySlice(t *testing.T) {
|
||||
result := hostfile.RemoveAll()
|
||||
if len(result) != 0 {
|
||||
t.Errorf("expected empty result, got %d entries: %v", len(result), result)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRemoveAll_AssembleContentOmitsBlock(t *testing.T) {
|
||||
// When RemoveAll result is passed to AssembleContent, no managed block markers
|
||||
// should appear in the assembled output.
|
||||
prefix := []string{"127.0.0.1 localhost"}
|
||||
postfix := []string{"# trailing"}
|
||||
result := hostfile.AssembleContent(prefix, hostfile.RemoveAll(), postfix)
|
||||
if len(result) != 2 {
|
||||
t.Fatalf("expected 2 lines with no managed block, got %d: %v", len(result), result)
|
||||
}
|
||||
for _, line := range result {
|
||||
if strings.Contains(line, "DNSHelper") {
|
||||
t.Errorf("managed block markers should be absent after RemoveAll: %v", result)
|
||||
}
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user