Files
ExploreDNS/web/api/handler.go
T
196eaa4a7a
CI / test (pull_request) Failing after 1m21s
feat: HTTP API server with SSE and SPA serving (HAN-389)
- web/api/server.go: HTTP server with CORS middleware, graceful shutdown,
  configurable listen address via NewServer(addr); uses net.Listen so the
  bound port is available via Addr() for tests.

- web/api/handler.go: REST handlers over Go 1.22 mux patterns:
    POST /api/traverse        start async traversal -> {id, status}
    GET  /api/traverse/{id}   poll status + results
    GET  /api/traverse/{id}/stream  SSE live progress events
    GET  /api/health          health check
  In-memory job store with 1-hour TTL and background cleanup goroutine.
  Background traversal via traverse.Traverser with hook-driven progress
  events fanned out to SSE subscribers.

- web/api/static/index.html: placeholder SPA shell (embedded via
  //go:embed static). Unknown URL paths fall back to index.html for
  client-side routing.

- cmd/server/main.go: runnable server binary with -addr flag and
  SIGINT/SIGTERM graceful shutdown.

All tests pass: go test -race ./...  |  go vet ./...

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Co-authored-by: multica-agent <github@multica.ai>
2026-06-08 04:19:02 +10:00

482 lines
12 KiB
Go

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:])
}