package api import ( "context" "crypto/rand" "encoding/json" "fmt" "io/fs" "net" "net/http" "os" "sort" "strconv" "strings" "sync" "time" "gitea.hansenits.com.au/hits/ExploreDNS/internal/config" idns "gitea.hansenits.com.au/hits/ExploreDNS/internal/dns" "gitea.hansenits.com.au/hits/ExploreDNS/internal/fingerprint" "gitea.hansenits.com.au/hits/ExploreDNS/internal/traverse" ) // Job status values. const ( statusRunning = "running" statusComplete = "complete" statusError = "error" ) // jobTTL is how long completed jobs are retained in memory. const jobTTL = time.Hour // Public-exposure guards. Every traversal gets a hard deadline so jobs // always reach a terminal state (the TTL cleanup only purges finished // jobs), and the number of in-flight traversals is capped. Overridable // via EXPLOREDNS_JOB_TIMEOUT (Go duration) and EXPLOREDNS_MAX_JOBS. const ( defaultJobTimeout = 5 * time.Minute defaultMaxRunningJobs = 8 ) // Fingerprinting runs after a traversal reaches a terminal state: bounded // concurrency across the unique server IPs, with its own overall deadline so // a timed-out or cancelled job context never blocks the server list. const ( fingerprintConcurrency = 8 fingerprintTimeout = 15 * time.Second ) // TraverseRequest is the JSON body for POST /api/traverse. type TraverseRequest struct { Domain string `json:"domain"` Type string `json:"type"` AllRoots bool `json:"all_roots"` } // TraverseStartResponse is returned by POST /api/traverse. type TraverseStartResponse struct { ID string `json:"id"` Status string `json:"status"` } // ProgressEvent carries a single traversal hook event. type ProgressEvent struct { Stage string `json:"stage"` RefID string `json:"refid"` Depth int `json:"depth"` Name string `json:"name"` QType string `json:"qtype"` Server string `json:"server,omitempty"` IPs string `json:"ips,omitempty"` Bailiwick string `json:"bailiwick,omitempty"` Status string `json:"status,omitempty"` IsResolve bool `json:"is_resolve,omitempty"` CompletedEarlier string `json:"completed_earlier,omitempty"` } // ResultItem is one aggregated leaf outcome for API consumers. Parent and // ParentIP identify the referring server so clients can render the noglue and // lame-referral wordings; Qname/Qclass/Qtype are the failing query so clients // can render the "While querying" line when it differs from the original. type ResultItem struct { RefID string `json:"refid,omitempty"` Depth int `json:"depth"` Probability float64 `json:"probability"` Status string `json:"status"` Server string `json:"server,omitempty"` IP string `json:"ip,omitempty"` Parent string `json:"parent,omitempty"` ParentIP string `json:"parent_ip,omitempty"` Qname string `json:"qname,omitempty"` Qclass string `json:"qclass,omitempty"` Qtype string `json:"qtype,omitempty"` Answers []string `json:"answers,omitempty"` Message string `json:"message,omitempty"` } // SummaryAnswer is one distinct answered RRset with its accumulated // probability (traverse.SummaryStats). type SummaryAnswer struct { Probability float64 `json:"probability"` Records []string `json:"records"` } // SummaryStatus is the accumulated probability of one non-answered status. type SummaryStatus struct { Status string `json:"status"` Probability float64 `json:"probability"` } // Summary is the grouped view of the aggregated leaves; probabilities across // Answers plus ByStatus sum to 1.0. type Summary struct { Answers []SummaryAnswer `json:"answers,omitempty"` ByStatus []SummaryStatus `json:"by_status,omitempty"` } // ServerInfo is one (server name, IP) pair queried during a traversal plus // its version.bind fingerprint ("" when the server didn't answer the probe). type ServerInfo struct { Name string `json:"name"` IP string `json:"ip"` Version string `json:"version"` } // TraversalJob holds all state for a single asynchronous traversal. type TraversalJob struct { ID string `json:"id"` Status string `json:"status"` Domain string `json:"domain"` QueryType string `json:"query_type"` Results []ResultItem `json:"results,omitempty"` Summary *Summary `json:"summary,omitempty"` Progress []ProgressEvent `json:"progress,omitempty"` Servers []ServerInfo `json:"servers,omitempty"` Error string `json:"error,omitempty"` StartedAt time.Time `json:"started_at"` DoneAt *time.Time `json:"done_at,omitempty"` mu sync.RWMutex subs []chan ProgressEvent cancel context.CancelFunc clientIP string // serversDone flips once the post-traversal fingerprinting step has // stored Servers (or was skipped); until then GET …/servers is pending. serversDone bool } // subscribeSnapshot atomically registers a subscriber and snapshots the // progress recorded so far. Publishing appends to Progress and sends to // subscribers under the same lock, so every event lands either in the // returned snapshot or on the channel — never both, never neither. func (j *TraversalJob) subscribeSnapshot() (sub <-chan ProgressEvent, past []ProgressEvent, done bool) { ch := make(chan ProgressEvent, 32) j.mu.Lock() defer j.mu.Unlock() j.subs = append(j.subs, ch) past = make([]ProgressEvent, len(j.Progress)) copy(past, j.Progress) return ch, past, j.Status != statusRunning } // publishLocked sends ev to all current subscribers. Caller must hold j.mu. func (j *TraversalJob) publishLocked(ev ProgressEvent) { for _, ch := range j.subs { select { case ch <- ev: default: // subscriber too slow — drop rather than block traversal } } } // unsubscribe removes ch from the subscriber list. func (j *TraversalJob) unsubscribe(ch <-chan ProgressEvent) { j.mu.Lock() defer j.mu.Unlock() for i, s := range j.subs { if s == ch { j.subs = append(j.subs[:i], j.subs[i+1:]...) return } } } // closeSubscribers drains and closes all subscriber channels. func (j *TraversalJob) closeSubscribers() { j.mu.Lock() defer j.mu.Unlock() for _, ch := range j.subs { close(ch) } j.subs = nil } // store is a thread-safe in-memory job registry. type store struct { mu sync.RWMutex jobs map[string]*TraversalJob } func newStore() *store { return &store{jobs: make(map[string]*TraversalJob)} } func (s *store) get(id string) (*TraversalJob, bool) { s.mu.RLock() defer s.mu.RUnlock() j, ok := s.jobs[id] return j, ok } func (s *store) set(j *TraversalJob) { s.mu.Lock() defer s.mu.Unlock() s.jobs[j.ID] = j } // runningCount reports how many jobs have not yet reached a terminal state. func (s *store) runningCount() int { s.mu.RLock() defer s.mu.RUnlock() n := 0 for _, j := range s.jobs { j.mu.RLock() if j.DoneAt == nil { n++ } j.mu.RUnlock() } return n } // cleanup removes completed jobs older than jobTTL. func (s *store) cleanup() { cutoff := time.Now().Add(-jobTTL) s.mu.Lock() defer s.mu.Unlock() for id, j := range s.jobs { j.mu.RLock() done := j.DoneAt != nil && j.DoneAt.Before(cutoff) j.mu.RUnlock() if done { delete(s.jobs, id) } } } // versionQuerier is the subset of fingerprint.Fingerprinter the handler // uses; tests substitute a fake so no probes leave the process. type versionQuerier interface { Query(ctx context.Context, ip net.IP) string } // Handler wires together the HTTP routes and the job store. type Handler struct { st *store mux *http.ServeMux jobTimeout time.Duration maxRunning int version string limiter *rateLimiter webhook *webhookReporter // fp fingerprints server IPs after each traversal; shared across jobs // so its per-IP cache is reused. fp versionQuerier } func newHandler(ctx context.Context) *Handler { limit, window := parseRateLimit(os.Getenv("EXPLOREDNS_RATE_LIMIT")) h := &Handler{ st: newStore(), mux: http.NewServeMux(), jobTimeout: envDuration("EXPLOREDNS_JOB_TIMEOUT", defaultJobTimeout), maxRunning: envInt("EXPLOREDNS_MAX_JOBS", defaultMaxRunningJobs), version: "dev", limiter: newRateLimiter(limit, window), webhook: newWebhookReporter(os.Getenv("EXPLOREDNS_WEBHOOK_URL"), os.Getenv("EXPLOREDNS_WEBHOOK_TOKEN")), fp: fingerprint.New(), } h.mux.HandleFunc("POST /api/traverse", h.startTraversal) h.mux.HandleFunc("GET /api/traverse/{id}/stream", h.streamTraversal) h.mux.HandleFunc("GET /api/traverse/{id}/servers", h.getServers) h.mux.HandleFunc("GET /api/traverse/{id}", h.getTraversal) h.mux.HandleFunc("GET /api/health", h.health) // periodic cleanup; exits when ctx is cancelled (e.g. on Server.Shutdown). go func() { t := time.NewTicker(10 * time.Minute) defer t.Stop() for { select { case <-ctx.Done(): return case <-t.C: h.st.cleanup() h.limiter.sweep() } } }() return h } // registerStatic adds SPA static file serving to the mux. func (h *Handler) registerStatic(sub fs.FS) { fileServer := http.FileServer(http.FS(sub)) h.mux.Handle("/", spaHandler{fileServer: fileServer, fs: sub}) } // health handles GET /api/health. func (h *Handler) health(w http.ResponseWriter, _ *http.Request) { body := map[string]string{"status": "ok", "version": h.version} // On Fly.io this identifies which machine served the request — // useful for observing anycast routing and auto-start behaviour. if region := os.Getenv("FLY_REGION"); region != "" { body["region"] = region } writeJSON(w, http.StatusOK, body) } // startTraversal handles POST /api/traverse. func (h *Handler) startTraversal(w http.ResponseWriter, r *http.Request) { ip := clientIP(r) if h.limiter != nil && !rateLimitExempt(r) && !h.limiter.allow(ip) { writeError(w, http.StatusTooManyRequests, fmt.Sprintf("rate limit exceeded: %s per client IP", h.limiter)) return } r.Body = http.MaxBytesReader(w, r.Body, 1<<20) // 1 MB limit var req TraverseRequest if err := json.NewDecoder(r.Body).Decode(&req); err != nil { writeError(w, http.StatusBadRequest, "invalid request body: "+err.Error()) return } if req.Domain == "" { writeError(w, http.StatusBadRequest, "domain is required") return } queryType := req.Type if queryType == "" { queryType = "A" } qtype, err := config.ParseQueryType(queryType) if err != nil { writeError(w, http.StatusBadRequest, err.Error()) return } if h.st.runningCount() >= h.maxRunning { writeError(w, http.StatusTooManyRequests, "too many concurrent traversals, try again shortly") return } id := newUUID() ctx, cancel := context.WithTimeout(context.Background(), h.jobTimeout) job := &TraversalJob{ ID: id, Status: statusRunning, Domain: req.Domain, QueryType: queryType, StartedAt: time.Now(), cancel: cancel, clientIP: ip, } h.st.set(job) h.webhook.send(webhookEventStart, webhookStartEvent{ Event: webhookEventStart, ID: job.ID, Domain: job.Domain, QueryType: job.QueryType, AllRoots: req.AllRoots, ClientIP: ip, StartedAt: job.StartedAt, }) go h.runTraversal(ctx, job, req.Domain, qtype, req.AllRoots) writeJSON(w, http.StatusAccepted, TraverseStartResponse{ID: id, Status: statusRunning}) } // getTraversal handles GET /api/traverse/{id}. func (h *Handler) getTraversal(w http.ResponseWriter, r *http.Request) { id := r.PathValue("id") job, ok := h.st.get(id) if !ok { writeError(w, http.StatusNotFound, "traversal not found") return } job.mu.RLock() // Copy the fields we need while holding the read lock. snapshot := struct { ID string `json:"id"` Status string `json:"status"` Domain string `json:"domain"` QueryType string `json:"query_type"` Results []ResultItem `json:"results,omitempty"` Summary *Summary `json:"summary,omitempty"` Progress []ProgressEvent `json:"progress,omitempty"` Servers []ServerInfo `json:"servers,omitempty"` Error string `json:"error,omitempty"` StartedAt time.Time `json:"started_at"` DoneAt *time.Time `json:"done_at,omitempty"` }{ ID: job.ID, Status: job.Status, Domain: job.Domain, QueryType: job.QueryType, Results: job.Results, Summary: job.Summary, Progress: job.Progress, Servers: job.Servers, Error: job.Error, StartedAt: job.StartedAt, DoneAt: job.DoneAt, } job.mu.RUnlock() writeJSON(w, http.StatusOK, snapshot) } // getServers handles GET /api/traverse/{id}/servers. It answers 202 with a // pending body while the traversal or the post-traversal fingerprinting is // still in flight, then the fingerprinted server list. func (h *Handler) getServers(w http.ResponseWriter, r *http.Request) { id := r.PathValue("id") job, ok := h.st.get(id) if !ok { writeError(w, http.StatusNotFound, "traversal not found") return } job.mu.RLock() done := job.serversDone servers := job.Servers job.mu.RUnlock() if !done { writeJSON(w, http.StatusAccepted, map[string]string{"status": "pending"}) return } if servers == nil { servers = []ServerInfo{} } writeJSON(w, http.StatusOK, struct { Status string `json:"status"` Servers []ServerInfo `json:"servers"` }{Status: "complete", Servers: servers}) } // streamTraversal handles GET /api/traverse/{id}/stream (SSE). func (h *Handler) streamTraversal(w http.ResponseWriter, r *http.Request) { id := r.PathValue("id") job, ok := h.st.get(id) if !ok { writeError(w, http.StatusNotFound, "traversal not found") return } // Set SSE headers before any write. w.Header().Set("Content-Type", "text/event-stream") w.Header().Set("Cache-Control", "no-cache") w.Header().Set("Connection", "keep-alive") w.Header().Set("X-Accel-Buffering", "no") flusher, ok := w.(http.Flusher) if !ok { writeError(w, http.StatusInternalServerError, "streaming not supported") return } // Subscribe and snapshot atomically so events published in between are // neither missed nor delivered twice. Unsubscribe when the client // disconnects so stale channels don't accumulate. sub, past, alreadyDone := job.subscribeSnapshot() defer job.unsubscribe(sub) sendSSE := func(ev ProgressEvent) bool { b, err := json.Marshal(ev) if err != nil { return false } _, err = fmt.Fprintf(w, "data: %s\n\n", b) flusher.Flush() return err == nil } for _, ev := range past { if !sendSSE(ev) { return } } if alreadyDone { fmt.Fprintf(w, "event: done\ndata: {}\n\n") flusher.Flush() return } ctx := r.Context() for { select { case <-ctx.Done(): return case ev, open := <-sub: if !open { fmt.Fprintf(w, "event: done\ndata: {}\n\n") flusher.Flush() return } if !sendSSE(ev) { return } } } } // runTraversal executes a traversal in the background and updates the job. func (h *Handler) runTraversal(ctx context.Context, job *TraversalJob, domain string, qtype uint16, allRoots bool) { defer func() { if r := recover(); r != nil { now := time.Now() job.mu.Lock() if job.DoneAt == nil { job.Status = statusError job.Error = fmt.Sprintf("panic: %v", r) job.DoneAt = &now } job.mu.Unlock() } // Never leave GET …/servers pending: fingerprinting is skipped on // the panic path, so flip the flag here (idempotent otherwise). job.mu.Lock() job.serversDone = true job.mu.Unlock() job.cancel() job.closeSubscribers() h.reportCompletion(job) }() cfg := traverse.DefaultTraverserConfig() cfg.QueryType = qtype if allRoots { cfg.RootConfig = &idns.RootDiscoveryConfig{AllRoots: true} } cfg.Hooks = &traverse.TraverserHooks{ OnEvent: func(event traverse.TraversalEvent) { ref := event.Referral if ref == nil || ref.IsRootRoot() { return } ev := ProgressEvent{ Stage: event.Stage.String(), RefID: event.RefID, Depth: ref.Depth(), Name: ref.Qname, QType: traverse.TypeToString(ref.Qtype), Server: ref.Server, IPs: ref.TxtIPs(), Bailiwick: ref.Bailiwick, Status: string(event.Status), IsResolve: event.IsResolve, CompletedEarlier: event.CompletedEarlier, } job.mu.Lock() job.Progress = append(job.Progress, ev) job.publishLocked(ev) // inside lock: no race with subscribe+replay job.mu.Unlock() }, } tr := traverse.NewTraverser(cfg) root, err := tr.Run(ctx, domain) h.commitResult(ctx, job, root, err) // Tell streaming clients the traversal reached a terminal state so they // can fetch results now; the stream stays open for the servers event // published once fingerprinting (below) finishes. job.mu.Lock() ev := ProgressEvent{Stage: "complete", Status: job.Status} job.Progress = append(job.Progress, ev) job.publishLocked(ev) job.mu.Unlock() // Fingerprint the servers queried during the traversal. This runs after // the terminal status is committed, so results never wait on versions. h.fingerprintServers(job, tr.ServersEncountered()) } // commitResult stores the traversal outcome and moves the job to its // terminal status. func (h *Handler) commitResult(ctx context.Context, job *TraversalJob, root *traverse.Referral, err error) { now := time.Now() job.mu.Lock() defer job.mu.Unlock() job.DoneAt = &now timedOut := ctx.Err() == context.DeadlineExceeded if err != nil && ctx.Err() == nil { job.Status = statusError job.Error = err.Error() return } var items []ResultItem if root != nil { leaves := root.StatsList() items = make([]ResultItem, 0, len(leaves)) for _, leaf := range leaves { items = append(items, toResultItem(leaf)) } job.Summary = toSummary(root.SummaryStats()) } job.Results = items if timedOut { // Keep any partial results, but do not report a truncated // traversal as a successful one. job.Status = statusError job.Error = fmt.Sprintf("traversal timed out after %s", h.jobTimeout) return } job.Status = statusComplete } // fingerprintServers turns the traversal's (server, ip) pairs into // job.Servers, probing each unique IP's version.bind with bounded // concurrency, then publishes a {"stage":"servers"} event so streaming // clients know the list is ready without polling. Pseudo "key:" entries and // non-address entries are skipped. func (h *Handler) fingerprintServers(job *TraversalJob, seen map[string][]string) { type pair struct{ name, ip string } var pairs []pair uniq := make(map[string]bool) var ips []net.IP for name, addrs := range seen { for _, addr := range addrs { if strings.HasPrefix(addr, "key:") { continue } ip := net.ParseIP(addr) if ip == nil { continue } pairs = append(pairs, pair{name: name, ip: addr}) if !uniq[addr] { uniq[addr] = true ips = append(ips, ip) } } } versions := make(map[string]string, len(ips)) if len(ips) > 0 && h.fp != nil { ctx, cancel := context.WithTimeout(context.Background(), fingerprintTimeout) defer cancel() var ( wg sync.WaitGroup mu sync.Mutex sem = make(chan struct{}, fingerprintConcurrency) ) for _, ip := range ips { wg.Add(1) go func(ip net.IP) { defer wg.Done() sem <- struct{}{} defer func() { <-sem }() v := h.fp.Query(ctx, ip) mu.Lock() versions[ip.String()] = v mu.Unlock() }(ip) } wg.Wait() } servers := make([]ServerInfo, 0, len(pairs)) for _, p := range pairs { servers = append(servers, ServerInfo{Name: p.name, IP: p.ip, Version: versions[p.ip]}) } sort.Slice(servers, func(i, j int) bool { if servers[i].Name != servers[j].Name { return servers[i].Name < servers[j].Name } return servers[i].IP < servers[j].IP }) ev := ProgressEvent{Stage: "servers"} job.mu.Lock() job.Servers = servers job.serversDone = true job.Progress = append(job.Progress, ev) job.publishLocked(ev) job.mu.Unlock() } // reportCompletion posts the webhook "complete" event for a job that has // reached a terminal state. Fire-and-forget; never blocks the caller. func (h *Handler) reportCompletion(job *TraversalJob) { if h.webhook == nil { return } job.mu.RLock() ev := webhookCompleteEvent{ Event: webhookEventComplete, ID: job.ID, Domain: job.Domain, QueryType: job.QueryType, ClientIP: job.clientIP, StartedAt: job.StartedAt, Status: job.Status, Error: job.Error, ResultCount: len(job.Results), Summary: job.Summary, } if job.DoneAt != nil { ev.DoneAt = *job.DoneAt ev.DurationMS = job.DoneAt.Sub(job.StartedAt).Milliseconds() } job.mu.RUnlock() h.webhook.send(webhookEventComplete, ev) } // envDuration reads a Go duration from the environment, falling back to // def when unset or unparsable. func envDuration(name string, def time.Duration) time.Duration { v := os.Getenv(name) if v == "" { return def } d, err := time.ParseDuration(v) if err != nil || d <= 0 { return def } return d } // envInt reads a positive integer from the environment, falling back to // def when unset or unparsable. func envInt(name string, def int) int { v := os.Getenv(name) if v == "" { return def } n, err := strconv.Atoi(v) if err != nil || n <= 0 { return def } return n } // toSummary converts the engine's grouped stats to the API representation: // answers sorted by RRset key (as SummaryStats returns them), remaining // statuses sorted lexically like the CLI Summary Results section. func toSummary(stats *traverse.SummaryStats) *Summary { if stats == nil { return nil } summary := &Summary{} for _, answer := range stats.Answers { item := SummaryAnswer{Probability: answer.Prob} for _, rr := range answer.RRs { item.Records = append(item.Records, collapseWhitespace(rr.String())) } summary.Answers = append(summary.Answers, item) } statuses := make([]traverse.Status, 0, len(stats.ByStatus)) for status := range stats.ByStatus { if status != traverse.StatusAnswered { statuses = append(statuses, status) } } sort.Slice(statuses, func(i, j int) bool { return statuses[i] < statuses[j] }) for _, status := range statuses { summary.ByStatus = append(summary.ByStatus, SummaryStatus{ Status: string(status), Probability: stats.ByStatus[status], }) } return summary } // collapseWhitespace renders an RR on one line with runs of whitespace // collapsed to single spaces, matching the CLI summary records. func collapseWhitespace(s string) string { return strings.Join(strings.Fields(s), " ") } // toResultItem converts one aggregated leaf to its API representation. func toResultItem(leaf *traverse.StatsEntry) ResultItem { resp := leaf.Response item := ResultItem{ Probability: leaf.Prob, Status: string(resp.Status), IP: resp.IP, Server: resp.Server, ParentIP: resp.ParentIP, Qname: trimFQDN(resp.Qname), Qclass: traverse.ClassToString(resp.Qclass), Qtype: traverse.TypeToString(resp.Qtype), } if leaf.Referral != nil { item.RefID = leaf.Referral.RefID item.Depth = leaf.Referral.Depth() item.Server = leaf.Referral.Server item.ParentIP = leaf.Referral.ParentIP if leaf.Referral.Parent != nil { item.Parent = leaf.Referral.Parent.Server } } if resp.DQ != nil { for _, rr := range resp.DQ.Answers { item.Answers = append(item.Answers, idns.FormatRecord(rr)) } switch resp.Status { case traverse.StatusError: item.Message = resp.DQ.ErrorMessage case traverse.StatusException: item.Message = resp.DQ.ExceptionMessage } } return item } // trimFQDN removes a trailing dot from an FQDN for nicer output. func trimFQDN(s string) string { if len(s) > 0 && s[len(s)-1] == '.' { return s[:len(s)-1] } return s } // spaHandler serves static files and falls back to index.html for unknown paths. type spaHandler struct { fileServer http.Handler fs fs.FS } func (h spaHandler) ServeHTTP(w http.ResponseWriter, r *http.Request) { // If the file exists, serve it directly. if _, err := fs.Stat(h.fs, r.URL.Path[1:]); err == nil { h.fileServer.ServeHTTP(w, r) return } // Serve index.html for SPA client-side routing. r2 := r.Clone(r.Context()) r2.URL.Path = "/" h.fileServer.ServeHTTP(w, r2) } // writeJSON encodes v as JSON and writes it to w with the given status code. func writeJSON(w http.ResponseWriter, status int, v any) { w.Header().Set("Content-Type", "application/json") w.WriteHeader(status) _ = json.NewEncoder(w).Encode(v) } // writeError writes a JSON error response. func writeError(w http.ResponseWriter, status int, msg string) { writeJSON(w, status, map[string]string{"error": msg}) } // newUUID returns a random UUID v4 string. func newUUID() string { var b [16]byte _, _ = rand.Read(b[:]) b[6] = (b[6] & 0x0f) | 0x40 // version 4 b[8] = (b[8] & 0x3f) | 0x80 // variant bits return fmt.Sprintf("%08x-%04x-%04x-%04x-%012x", b[0:4], b[4:6], b[6:8], b[8:10], b[10:]) }