package api import ( "context" "crypto/rand" "encoding/json" "fmt" "io/fs" "net/http" "sync" "time" "github.com/hits/ExploreDNS/internal/config" idns "github.com/hits/ExploreDNS/internal/dns" "github.com/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 // 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"` Depth int `json:"depth"` Name string `json:"name"` QType string `json:"qtype"` Server string `json:"server,omitempty"` Bailiwick string `json:"bailiwick,omitempty"` IsResolve bool `json:"is_resolve,omitempty"` } // ResultItem is a single traversal step result for API consumers. type ResultItem struct { Depth int `json:"depth"` Probability float64 `json:"probability"` ResponseType string `json:"response_type"` Server string `json:"server,omitempty"` Answers []string `json:"answers,omitempty"` CNAMEChain []string `json:"cname_chain,omitempty"` } // 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"` Progress []ProgressEvent `json:"progress,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 } // subscribe returns a channel that receives future progress events. // The channel is closed when the job finishes. func (j *TraversalJob) subscribe() <-chan ProgressEvent { ch := make(chan ProgressEvent, 32) j.mu.Lock() j.subs = append(j.subs, ch) j.mu.Unlock() return ch } // publish sends an event to all current subscribers. func (j *TraversalJob) publish(ev ProgressEvent) { j.mu.Lock() defer j.mu.Unlock() for _, ch := range j.subs { select { case ch <- ev: default: // subscriber too slow — drop rather than block traversal } } } // 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 } // 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) } } } // Handler wires together the HTTP routes and the job store. type Handler struct { st *store mux *http.ServeMux } func newHandler() *Handler { h := &Handler{ st: newStore(), mux: http.NewServeMux(), } 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}", h.getTraversal) h.mux.HandleFunc("GET /api/health", h.health) // periodic cleanup go func() { t := time.NewTicker(10 * time.Minute) defer t.Stop() for range t.C { h.st.cleanup() } }() 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) { writeJSON(w, http.StatusOK, map[string]string{"status": "ok"}) } // startTraversal handles POST /api/traverse. func (h *Handler) startTraversal(w http.ResponseWriter, r *http.Request) { 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 } id := newUUID() ctx, cancel := context.WithCancel(context.Background()) job := &TraversalJob{ ID: id, Status: statusRunning, Domain: req.Domain, QueryType: queryType, StartedAt: time.Now(), cancel: cancel, } h.st.set(job) 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"` Progress []ProgressEvent `json:"progress,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, Progress: job.Progress, Error: job.Error, StartedAt: job.StartedAt, DoneAt: job.DoneAt, } job.mu.RUnlock() writeJSON(w, http.StatusOK, snapshot) } // 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 before snapshotting progress so we don't miss events between // the two operations. sub := job.subscribe() // Replay events already recorded. job.mu.RLock() past := make([]ProgressEvent, len(job.Progress)) copy(past, job.Progress) alreadyDone := job.Status != statusRunning job.mu.RUnlock() 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() job.Status = statusError job.Error = fmt.Sprintf("panic: %v", r) job.DoneAt = &now job.mu.Unlock() } job.cancel() job.closeSubscribers() }() 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.Result.Referral if ref == nil { return } stage := "start" if event.Stage == traverse.EventComplete { stage = "complete" } server := "" if event.Result.Response != nil && event.Result.Response.Server != nil { server = event.Result.Response.Server.String() } ev := ProgressEvent{ Stage: stage, Depth: ref.Depth, Name: trimFQDN(ref.Name), QType: idns.QNameType(ref.Qtype), Server: server, Bailiwick: trimFQDN(ref.Bailiwick), IsResolve: event.IsResolve, } job.mu.Lock() job.Progress = append(job.Progress, ev) job.mu.Unlock() job.publish(ev) }, } tr := traverse.NewTraverser(cfg) rawResults, err := tr.Traverse(ctx, domain) now := time.Now() job.mu.Lock() defer job.mu.Unlock() job.DoneAt = &now if err != nil && ctx.Err() == nil { job.Status = statusError job.Error = err.Error() return } items := make([]ResultItem, 0, len(rawResults)) for _, r := range rawResults { items = append(items, toResultItem(r)) } job.Results = items job.Status = statusComplete } // toResultItem converts a TraversalResult to its API representation. func toResultItem(r traverse.TraversalResult) ResultItem { item := ResultItem{} if r.Referral != nil { item.Depth = r.Referral.Depth item.Probability = r.Referral.Prob } if r.Response != nil { item.ResponseType = r.Response.Type.String() if r.Response.Server != nil { item.Server = r.Response.Server.String() } if r.Response.Decoded != nil { for _, rr := range r.Response.Decoded.Answers { item.Answers = append(item.Answers, idns.FormatRecord(rr)) } item.CNAMEChain = append(item.CNAMEChain, r.Response.Decoded.CNAMEChain...) } } 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:]) }