package config import ( "errors" "testing" ) func TestParseQueryTypeValid(t *testing.T) { cases := []struct { input string wantType uint16 }{ {"A", 1}, {"aaaa", 28}, {"NS", 2}, {"CNAME", 5}, {"MX", 15}, {"TXT", 16}, {"SOA", 6}, {"PTR", 12}, {"ANY", 255}, } for _, tc := range cases { got, err := ParseQueryType(tc.input) if err != nil { t.Errorf("ParseQueryType(%q) unexpected error: %v", tc.input, err) } if got != tc.wantType { t.Errorf("ParseQueryType(%q) = %d, want %d", tc.input, got, tc.wantType) } } } func TestParseQueryTypeInvalid(t *testing.T) { _, err := ParseQueryType("BOGUS") if err == nil { t.Fatal("expected error for invalid query type") } if !errors.Is(err, ErrInvalidQueryType) { t.Errorf("expected ErrInvalidQueryType, got %v", err) } } func TestParseUDPSizeValid(t *testing.T) { cases := []string{"512", "2048", "4096"} for _, s := range cases { if _, err := ParseUDPSize(s); err != nil { t.Errorf("ParseUDPSize(%q) unexpected error: %v", s, err) } } } func TestParseUDPSizeInvalid(t *testing.T) { cases := []string{"0", "511", "4097", "notanumber"} for _, s := range cases { _, err := ParseUDPSize(s) if err == nil { t.Errorf("ParseUDPSize(%q): expected error", s) } if !errors.Is(err, ErrInvalidUDPSize) { t.Errorf("ParseUDPSize(%q): expected ErrInvalidUDPSize, got %v", s, err) } } } func TestValidateOK(t *testing.T) { cfg := DefaultConfig() if err := cfg.Validate(); err != nil { t.Fatalf("DefaultConfig should be valid, got: %v", err) } } func TestValidateAlwaysTCPRequiresAllowTCP(t *testing.T) { cfg := DefaultConfig() cfg.AlwaysTCP = true cfg.AllowTCP = false err := cfg.Validate() if err == nil { t.Fatal("expected error when AlwaysTCP=true and AllowTCP=false") } if !errors.Is(err, ErrAlwaysTCPRequiresTCP) { t.Errorf("expected ErrAlwaysTCPRequiresTCP, got %v", err) } } func TestValidateAlwaysTCPWithAllowTCP(t *testing.T) { cfg := DefaultConfig() cfg.AlwaysTCP = true cfg.AllowTCP = true if err := cfg.Validate(); err != nil { t.Errorf("AlwaysTCP=true, AllowTCP=true should be valid, got: %v", err) } } func TestValidateBadUDPSize(t *testing.T) { cfg := DefaultConfig() cfg.UDPSize = 100 if err := cfg.Validate(); !errors.Is(err, ErrInvalidUDPSize) { t.Errorf("expected ErrInvalidUDPSize, got %v", err) } } func TestValidateBadMaxDepth(t *testing.T) { cfg := DefaultConfig() cfg.MaxDepth = 0 if err := cfg.Validate(); !errors.Is(err, ErrInvalidMaxDepth) { t.Errorf("expected ErrInvalidMaxDepth, got %v", err) } } func TestValidateBadRetries(t *testing.T) { cfg := DefaultConfig() cfg.Retries = 11 if err := cfg.Validate(); !errors.Is(err, ErrInvalidRetries) { t.Errorf("expected ErrInvalidRetries, got %v", err) } } func TestGetDomainOK(t *testing.T) { cfg := DefaultConfig() d, err := cfg.GetDomain([]string{"example.com"}) if err != nil { t.Fatalf("unexpected error: %v", err) } if d != "example.com" { t.Errorf("got %q, want %q", d, "example.com") } } func TestGetDomainMissing(t *testing.T) { cfg := DefaultConfig() _, err := cfg.GetDomain([]string{}) if err == nil { t.Fatal("expected error for missing domain") } if !errors.Is(err, ErrMissingDomain) { t.Errorf("expected ErrMissingDomain, got %v", err) } } func TestParseRootServerEmpty(t *testing.T) { cfg := DefaultConfig() ip, err := cfg.ParseRootServer() if err != nil { t.Fatalf("unexpected error for empty root server: %v", err) } if ip != nil { t.Errorf("expected nil IP for empty root server, got %v", ip) } } func TestParseRootServerValidIPv4(t *testing.T) { cfg := DefaultConfig() cfg.RootServer = "198.41.0.4" ip, err := cfg.ParseRootServer() if err != nil { t.Fatalf("unexpected error: %v", err) } if ip == nil || ip.String() != "198.41.0.4" { t.Errorf("expected 198.41.0.4, got %v", ip) } } func TestParseRootServerValidIPv6(t *testing.T) { cfg := DefaultConfig() cfg.RootServer = "2001:503:ba3e::2:30" ip, err := cfg.ParseRootServer() if err != nil { t.Fatalf("unexpected error: %v", err) } if ip == nil { t.Error("expected non-nil IP for valid IPv6 address") } } func TestParseRootServerInvalid(t *testing.T) { cfg := DefaultConfig() cfg.RootServer = "not-an-ip" _, err := cfg.ParseRootServer() if err == nil { t.Fatal("expected error for invalid root server IP") } if !errors.Is(err, ErrInvalidRootServer) { t.Errorf("expected ErrInvalidRootServer, got %v", err) } } func TestParseRootServerHostname(t *testing.T) { cfg := DefaultConfig() cfg.RootServer = "a.root-servers.net" _, err := cfg.ParseRootServer() if err == nil { t.Fatal("expected error for hostname (not IP) root server") } if !errors.Is(err, ErrInvalidRootServer) { t.Errorf("expected ErrInvalidRootServer, got %v", err) } } func TestParseDebugLevel(t *testing.T) { cases := []struct { d, dd bool want int }{ {false, false, 0}, {true, false, 1}, {false, true, 2}, {true, true, 2}, // dd takes precedence } for _, tc := range cases { got := ParseDebugLevel(tc.d, tc.dd) if got != tc.want { t.Errorf("ParseDebugLevel(d=%v, dd=%v) = %d, want %d", tc.d, tc.dd, got, tc.want) } } } func TestParseMaxDepthValid(t *testing.T) { cases := []string{"1", "20", "100"} for _, s := range cases { v, err := ParseMaxDepth(s) if err != nil { t.Errorf("ParseMaxDepth(%q) unexpected error: %v", s, err) } if v < 1 || v > 100 { t.Errorf("ParseMaxDepth(%q) = %d, out of range", s, v) } } } func TestParseMaxDepthInvalid(t *testing.T) { cases := []string{"0", "101", "notanumber"} for _, s := range cases { _, err := ParseMaxDepth(s) if err == nil { t.Errorf("ParseMaxDepth(%q): expected error", s) } if !errors.Is(err, ErrInvalidMaxDepth) { t.Errorf("ParseMaxDepth(%q): expected ErrInvalidMaxDepth, got %v", s, err) } } } func TestParseRetriesValid(t *testing.T) { cases := []string{"0", "5", "10"} for _, s := range cases { v, err := ParseRetries(s) if err != nil { t.Errorf("ParseRetries(%q) unexpected error: %v", s, err) } if v < 0 || v > 10 { t.Errorf("ParseRetries(%q) = %d, out of range", s, v) } } } func TestParseRetriesInvalid(t *testing.T) { cases := []string{"-1", "11", "notanumber"} for _, s := range cases { _, err := ParseRetries(s) if err == nil { t.Errorf("ParseRetries(%q): expected error", s) } if !errors.Is(err, ErrInvalidRetries) { t.Errorf("ParseRetries(%q): expected ErrInvalidRetries, got %v", s, err) } } } func TestValidateBadQueryType(t *testing.T) { cfg := DefaultConfig() cfg.QueryType = "BOGUS" if err := cfg.Validate(); !errors.Is(err, ErrInvalidQueryType) { t.Errorf("expected ErrInvalidQueryType, got %v", err) } } func TestDefaultConfigIsValid(t *testing.T) { cfg := DefaultConfig() if cfg.QueryType != "A" { t.Errorf("QueryType = %q, want A", cfg.QueryType) } if cfg.MaxDepth != 20 { t.Errorf("MaxDepth = %d, want 20", cfg.MaxDepth) } if cfg.Retries != 2 { t.Errorf("Retries = %d, want 2", cfg.Retries) } }