feat: implement CLI flags and configuration handling (HAN-382) #8

Merged
multica-agent merged 2 commits from agent/go-expert-developer/4bf98d33 into main 2026-06-07 16:08:41 +00:00
2 changed files with 403 additions and 10 deletions
Showing only changes of commit d66dcbe067 - Show all commits
+183 -10
View File
@@ -1,26 +1,199 @@
package main
import (
"context"
"flag"
"fmt"
"os"
"strings"
"time"
"github.com/hits/ExploreDNS/internal/config"
"github.com/hits/ExploreDNS/internal/dns"
"github.com/hits/ExploreDNS/internal/traverse"
)
func main() {
flag.Usage = func() {
fmt.Fprintf(os.Stderr, "ExploreDNS - DNS reconnaissance and exploration tool\n\n")
fmt.Fprintf(os.Stderr, "Usage:\n")
fmt.Fprintf(os.Stderr, " exploredns [flags] <domain>\n\n")
fmt.Fprintf(os.Stderr, "Flags:\n")
flag.PrintDefaults()
}
cfg := config.DefaultConfig()
queryType := flag.String("type", cfg.QueryType, "Record type (A, AAAA, NS, CNAME, MX, TXT, SOA, PTR, ANY)")
rootServer := flag.String("root-server", cfg.RootServer, "Override root server")
allRootServers := flag.Bool("all-root-servers", cfg.AllRootServers, "Use all 13 root servers")
rootAAAA := flag.Bool("root-aaaa", cfg.RootAAAA, "Include IPv6 root addresses")
followAAAA := flag.Bool("follow-aaaa", cfg.FollowAAAA, "Only follow AAAA for referrals")
udpSize := flag.Int("udp-size", cfg.UDPSize, "EDNS0 buffer size (512-4096)")
allowTCP := flag.Bool("allow-tcp", cfg.AllowTCP, "TCP fallback on truncation")
alwaysTCP := flag.Bool("always-tcp", cfg.AlwaysTCP, "Always use TCP")
maxDepth := flag.Int("max-depth", cfg.MaxDepth, "Max traversal depth (1-100)")
retries := flag.Int("retries", cfg.Retries, "Retry count (0-10)")
fast := flag.Bool("fast", cfg.Fast, "Fast mode: reuse earlier branch cache")
verbose := flag.Bool("verbose", cfg.Verbose, "Verbose output")
debugFlag := flag.Bool("debug", false, "Debug mode (stackable: -dd for library debug)")
quiet := flag.Bool("quiet", cfg.Quiet, "Suppress supplementary info")
showProgress := flag.Bool("show-progress", cfg.ShowProgress, "Show traversal progress")
noShowProgress := flag.Bool("no-show-progress", false, "Hide traversal progress")
showResolves := flag.Bool("show-resolves", cfg.ShowResolves, "Show glue resolution details")
noShowResolves := flag.Bool("no-show-resolves", false, "Hide glue resolution details")
showServers := flag.Bool("show-servers", cfg.ShowServers, "Show servers queried")
noShowServers := flag.Bool("no-show-servers", false, "Hide servers queried")
showVersions := flag.Bool("show-versions", cfg.ShowVersions, "Show server versions")
noShowVersions := flag.Bool("no-show-versions", false, "Hide server versions")
showAllStats := flag.Bool("show-all-stats", cfg.ShowAllStats, "Show all statistics")
noShowAllStats := flag.Bool("no-show-all-stats", false, "Hide all statistics")
showResults := flag.Bool("show-results", cfg.ShowResults, "Show query results")
noShowResults := flag.Bool("no-show-results", false, "Hide query results")
showSummaryResults := flag.Bool("show-summary-results", cfg.ShowSummaryResults, "Show summary of results")
noShowSummaryResults := flag.Bool("no-show-summary-results", false, "Hide summary of results")
flag.Usage = config.PrintUsage
flag.Parse()
if flag.NArg() == 0 {
flag.Usage()
args := flag.Args()
debugLevel := 0
for _, arg := range os.Args[1:] {
if strings.Count(arg, "-d") > 0 || strings.Count(arg, "-dd") > 0 {
debugLevel = strings.Count(arg, "-d") + strings.Count(arg, "-dd")
}
}
if *debugFlag {
debugLevel++
}
cfg.QueryType = *queryType
cfg.RootServer = *rootServer
cfg.AllRootServers = *allRootServers
cfg.RootAAAA = *rootAAAA
cfg.FollowAAAA = *followAAAA
cfg.UDPSize = *udpSize
cfg.AllowTCP = *allowTCP
cfg.AlwaysTCP = *alwaysTCP
cfg.MaxDepth = *maxDepth
cfg.Retries = *retries
cfg.Fast = *fast
cfg.Verbose = *verbose
cfg.Debug = debugLevel
cfg.Quiet = *quiet
if *noShowProgress {
cfg.ShowProgress = false
} else if *showProgress {
cfg.ShowProgress = true
}
if *noShowResolves {
cfg.ShowResolves = false
} else if *showResolves {
cfg.ShowResolves = true
}
if *noShowServers {
cfg.ShowServers = false
} else if *showServers {
cfg.ShowServers = true
}
if *noShowVersions {
cfg.ShowVersions = false
} else if *showVersions {
cfg.ShowVersions = true
}
if *noShowAllStats {
cfg.ShowAllStats = false
} else if *showAllStats {
cfg.ShowAllStats = true
}
if *noShowResults {
cfg.ShowResults = false
} else if *showResults {
cfg.ShowResults = true
}
if *noShowSummaryResults {
cfg.ShowSummaryResults = false
} else if *showSummaryResults {
cfg.ShowSummaryResults = true
}
if err := cfg.Validate(); err != nil {
fmt.Fprintf(os.Stderr, "Error: %v\n", err)
os.Exit(1)
}
fmt.Printf("ExploreDNS - exploring: %s\n", flag.Arg(0))
domain, err := cfg.GetDomain(args)
if err != nil {
fmt.Fprintf(os.Stderr, "Error: %v\n", err)
os.Exit(1)
}
rootIP, err := cfg.ParseRootServer()
if err != nil {
fmt.Fprintf(os.Stderr, "Error: invalid root server: %v\n", err)
os.Exit(1)
}
queryTypeValue, err := config.ParseQueryType(cfg.QueryType)
if err != nil {
fmt.Fprintf(os.Stderr, "Error: %v\n", err)
os.Exit(1)
}
var rootServerAddr string
if rootIP != nil {
rootServerAddr = rootIP.String()
}
queryConfig := &dns.QueryConfig{
UDPSize: cfg.UDPSize,
Timeout: 5 * time.Second,
Retries: cfg.Retries,
UseTCP: cfg.AlwaysTCP,
}
rootConfig := &dns.RootDiscoveryConfig{
IncludeAAAA: cfg.RootAAAA,
Server: rootServerAddr,
AllRoots: cfg.AllRootServers,
}
traverserConfig := &traverse.TraverserConfig{
MaxDepth: cfg.MaxDepth,
QueryType: queryTypeValue,
RootConfig: rootConfig,
QueryConfig: queryConfig,
}
traverser := traverse.NewTraverser(traverserConfig)
if cfg.Debug > 0 {
fmt.Fprintf(os.Stderr, "Debug: Config loaded\n")
fmt.Fprintf(os.Stderr, " Domain: %s\n", domain)
fmt.Fprintf(os.Stderr, " Query Type: %s\n", cfg.QueryType)
fmt.Fprintf(os.Stderr, " Max Depth: %d\n", cfg.MaxDepth)
fmt.Fprintf(os.Stderr, " UDP Size: %d\n", cfg.UDPSize)
fmt.Fprintf(os.Stderr, " Retries: %d\n", cfg.Retries)
fmt.Fprintf(os.Stderr, " Allow TCP: %v\n", cfg.AllowTCP)
fmt.Fprintf(os.Stderr, " Always TCP: %v\n", cfg.AlwaysTCP)
fmt.Fprintf(os.Stderr, " Fast Mode: %v\n", cfg.Fast)
}
if !cfg.Quiet {
fmt.Printf("ExploreDNS - exploring: %s (type: %s)\n", domain, cfg.QueryType)
}
ctx := context.Background()
results, err := traverser.Traverse(ctx, domain)
if err != nil {
fmt.Fprintf(os.Stderr, "Error: traversal failed: %v\n", err)
os.Exit(1)
}
if cfg.ShowResults && len(results) > 0 {
fmt.Printf("\nResults:\n")
for i, result := range results {
fmt.Printf(" [%d] %s -> %s\n", i+1, result.Referral.Name, result.Response.Type)
}
}
if cfg.Debug > 0 {
fmt.Fprintf(os.Stderr, "Debug: Traversal completed with %d results\n", len(results))
}
}
+220
View File
@@ -1 +1,221 @@
package config
import (
"errors"
"fmt"
"net"
"os"
"strings"
"github.com/miekg/dns"
)
var (
ErrInvalidQueryType = errors.New("invalid query type")
ErrInvalidUDPSize = errors.New("UDP size must be between 512 and 4096")
ErrInvalidMaxDepth = errors.New("max depth must be between 1 and 100")
ErrInvalidRetries = errors.New("retries must be between 0 and 10")
ErrMissingDomain = errors.New("domain is required")
ErrAlwaysTCPRequiresTCP = errors.New("--always-tcp requires --allow-tcp")
)
type Config struct {
QueryType string
RootServer string
AllRootServers bool
RootAAAA bool
FollowAAAA bool
UDPSize int
AllowTCP bool
AlwaysTCP bool
MaxDepth int
Retries int
Fast bool
Verbose bool
Debug int
Quiet bool
ShowProgress bool
ShowResolves bool
ShowServers bool
ShowVersions bool
ShowAllStats bool
ShowResults bool
ShowSummaryResults bool
}
func ParseQueryType(s string) (uint16, error) {
s = strings.ToUpper(s)
switch s {
case "A":
return dns.TypeA, nil
case "AAAA":
return dns.TypeAAAA, nil
case "NS":
return dns.TypeNS, nil
case "CNAME":
return dns.TypeCNAME, nil
case "MX":
return dns.TypeMX, nil
case "TXT":
return dns.TypeTXT, nil
case "SOA":
return dns.TypeSOA, nil
case "PTR":
return dns.TypePTR, nil
case "ANY":
return dns.TypeANY, nil
default:
return 0, fmt.Errorf("%w: %s", ErrInvalidQueryType, s)
}
}
func ParseUDPSize(s string) (int, error) {
var size int
if _, err := fmt.Sscanf(s, "%d", &size); err != nil {
return 0, fmt.Errorf("%w: %s", ErrInvalidUDPSize, s)
}
if size < 512 || size > 4096 {
return 0, fmt.Errorf("%w: %d (must be 512-4096)", ErrInvalidUDPSize, size)
}
return size, nil
}
func ParseMaxDepth(s string) (int, error) {
var depth int
if _, err := fmt.Sscanf(s, "%d", &depth); err != nil {
return 0, fmt.Errorf("%w: %s", ErrInvalidMaxDepth, s)
}
if depth < 1 || depth > 100 {
return 0, fmt.Errorf("%w: %d (must be 1-100)", ErrInvalidMaxDepth, depth)
}
return depth, nil
}
func ParseRetries(s string) (int, error) {
var retries int
if _, err := fmt.Sscanf(s, "%d", &retries); err != nil {
return 0, fmt.Errorf("%w: %s", ErrInvalidRetries, s)
}
if retries < 0 || retries > 10 {
return 0, fmt.Errorf("%w: %d (must be 0-10)", ErrInvalidRetries, retries)
}
return retries, nil
}
func (c *Config) Validate() error {
if _, err := ParseQueryType(c.QueryType); err != nil {
return err
}
if c.UDPSize < 512 || c.UDPSize > 4096 {
return fmt.Errorf("%w: %d", ErrInvalidUDPSize, c.UDPSize)
}
if c.MaxDepth < 1 || c.MaxDepth > 100 {
return fmt.Errorf("%w: %d", ErrInvalidMaxDepth, c.MaxDepth)
}
if c.Retries < 0 || c.Retries > 10 {
return fmt.Errorf("%w: %d", ErrInvalidRetries, c.Retries)
}
if c.AlwaysTCP && !c.AllowTCP {
return ErrAlwaysTCPRequiresTCP
}
return nil
}
func (c *Config) GetDomain(args []string) (string, error) {
if len(args) == 0 {
return "", ErrMissingDomain
}
return args[0], nil
}
func (c *Config) ParseRootServer() (net.IP, error) {
if c.RootServer == "" {
return nil, nil
}
return net.ParseIP(c.RootServer), nil
}
func PrintUsage() {
fmt.Fprintf(os.Stderr, "ExploreDNS - DNS reconnaissance and exploration tool\n\n")
fmt.Fprintf(os.Stderr, "Usage:\n")
fmt.Fprintf(os.Stderr, " exploredns [flags] <domain>\n\n")
fmt.Fprintf(os.Stderr, "Flags:\n")
flagGroups := map[string][][2]string{
"Query Options": {
{"--type", "Record type (default A)"},
{"--root-server", "Override root server"},
{"--all-root-servers", "Use all 13 root servers"},
{"--root-aaaa", "Include IPv6 root addresses"},
{"--follow-aaaa", "Only follow AAAA for referrals"},
},
"Transport Options": {
{"--udp-size", "EDNS0 buffer size (default 2048)"},
{"--allow-tcp", "TCP fallback on truncation (default true)"},
{"--always-tcp", "Always use TCP"},
{"--retries", "Retry count (default 2)"},
},
"Traversal Options": {
{"--max-depth", "Max traversal depth (default 20)"},
{"--fast", "Fast mode: reuse earlier branch cache (default true)"},
},
"Output Options": {
{"--verbose, -v", "Verbose output"},
{"--debug, -d", "Debug mode (stackable: -dd for library debug)"},
{"--quiet, -q", "Suppress supplementary info"},
{"--show-progress", "Show traversal progress"},
{"--no-show-progress", "Hide traversal progress"},
{"--show-resolves", "Show glue resolution details"},
{"--no-show-resolves", "Hide glue resolution details"},
{"--show-servers", "Show servers queried"},
{"--no-show-servers", "Hide servers queried"},
{"--show-versions", "Show server versions"},
{"--no-show-versions", "Hide server versions"},
{"--show-all-stats", "Show all statistics"},
{"--no-show-all-stats", "Hide all statistics"},
{"--show-results", "Show query results"},
{"--no-show-results", "Hide query results"},
{"--show-summary-results", "Show summary of results"},
{"--no-show-summary-results", "Hide summary of results"},
},
}
for group, flags := range flagGroups {
fmt.Fprintf(os.Stderr, "\n%s:\n", group)
for _, f := range flags {
fmt.Fprintf(os.Stderr, " %-25s %s\n", f[0], f[1])
}
}
}
func DefaultConfig() *Config {
return &Config{
QueryType: "A",
RootServer: "",
AllRootServers: false,
RootAAAA: false,
FollowAAAA: false,
UDPSize: 2048,
AllowTCP: true,
AlwaysTCP: false,
MaxDepth: 20,
Retries: 2,
Fast: true,
Verbose: false,
Debug: 0,
Quiet: false,
ShowProgress: true,
ShowResolves: true,
ShowServers: true,
ShowVersions: true,
ShowAllStats: true,
ShowResults: true,
ShowSummaryResults: true,
}
}