diff --git a/cmd/exploredns/main.go b/cmd/exploredns/main.go index cadef08..4a57fda 100644 --- a/cmd/exploredns/main.go +++ b/cmd/exploredns/main.go @@ -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] \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)) + } } diff --git a/internal/config/config.go b/internal/config/config.go index d912156..f71ebfc 100644 --- a/internal/config/config.go +++ b/internal/config/config.go @@ -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] \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, + } +} \ No newline at end of file