diff --git a/cmd/server/main.go b/cmd/server/main.go new file mode 100644 index 0000000..f89df97 --- /dev/null +++ b/cmd/server/main.go @@ -0,0 +1,36 @@ +// Command server starts the ExploreDNS HTTP API server. +package main + +import ( + "flag" + "fmt" + "log" + "os" + "os/signal" + "syscall" + "time" + + "github.com/hits/ExploreDNS/web/api" +) + +func main() { + addr := flag.String("addr", ":8080", "listen address (host:port)") + flag.Parse() + + srv := api.NewServer(*addr) + if err := srv.Start(); err != nil { + fmt.Fprintf(os.Stderr, "Error: %v\n", err) + os.Exit(1) + } + + log.Printf("ExploreDNS API server listening on %s", srv.Addr()) + + quit := make(chan os.Signal, 1) + signal.Notify(quit, syscall.SIGINT, syscall.SIGTERM) + <-quit + + log.Println("Shutting down...") + if err := srv.Shutdown(15 * time.Second); err != nil { + log.Printf("Shutdown error: %v", err) + } +} diff --git a/web/api/handler.go b/web/api/handler.go new file mode 100644 index 0000000..6edeb15 --- /dev/null +++ b/web/api/handler.go @@ -0,0 +1,481 @@ +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:]) +} + diff --git a/web/api/handler_test.go b/web/api/handler_test.go new file mode 100644 index 0000000..fa31454 --- /dev/null +++ b/web/api/handler_test.go @@ -0,0 +1,305 @@ +package api_test + +import ( + "bufio" + "bytes" + "encoding/json" + "fmt" + "net/http" + "net/http/httptest" + "strings" + "testing" + "time" + + "github.com/hits/ExploreDNS/web/api" +) + +// newTestServer starts a Server on a random port and returns it. +// The caller is responsible for calling srv.Shutdown. +func newTestServer(t *testing.T) *api.Server { + t.Helper() + srv := api.NewServer("127.0.0.1:0") + if err := srv.Start(); err != nil { + t.Fatalf("start server: %v", err) + } + return srv +} + +func TestHealth(t *testing.T) { + srv := newTestServer(t) + defer srv.Shutdown(5 * time.Second) //nolint:errcheck + + resp, err := http.Get("http://" + srv.Addr() + "/api/health") + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + t.Fatalf("want 200, got %d", resp.StatusCode) + } + var body map[string]string + if err := json.NewDecoder(resp.Body).Decode(&body); err != nil { + t.Fatal(err) + } + if body["status"] != "ok" { + t.Fatalf("want status=ok, got %q", body["status"]) + } +} + +func TestCORSPreflight(t *testing.T) { + srv := newTestServer(t) + defer srv.Shutdown(5 * time.Second) //nolint:errcheck + + req, _ := http.NewRequest(http.MethodOptions, "http://"+srv.Addr()+"/api/health", nil) + req.Header.Set("Origin", "http://localhost:3000") + resp, err := http.DefaultClient.Do(req) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusNoContent { + t.Fatalf("want 204, got %d", resp.StatusCode) + } + if got := resp.Header.Get("Access-Control-Allow-Origin"); got != "*" { + t.Fatalf("CORS origin header: %q", got) + } +} + +func TestStartTraversal_MissingDomain(t *testing.T) { + srv := newTestServer(t) + defer srv.Shutdown(5 * time.Second) //nolint:errcheck + + body := bytes.NewBufferString(`{}`) + resp, err := http.Post("http://"+srv.Addr()+"/api/traverse", "application/json", body) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusBadRequest { + t.Fatalf("want 400, got %d", resp.StatusCode) + } +} + +func TestStartTraversal_InvalidType(t *testing.T) { + srv := newTestServer(t) + defer srv.Shutdown(5 * time.Second) //nolint:errcheck + + body := bytes.NewBufferString(`{"domain":"example.com","type":"BOGUS"}`) + resp, err := http.Post("http://"+srv.Addr()+"/api/traverse", "application/json", body) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusBadRequest { + t.Fatalf("want 400, got %d", resp.StatusCode) + } +} + +func TestStartTraversal_ReturnsID(t *testing.T) { + srv := newTestServer(t) + defer srv.Shutdown(5 * time.Second) //nolint:errcheck + + body := bytes.NewBufferString(`{"domain":"example.com","type":"A"}`) + resp, err := http.Post("http://"+srv.Addr()+"/api/traverse", "application/json", body) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusAccepted { + t.Fatalf("want 202, got %d", resp.StatusCode) + } + var start struct { + ID string `json:"id"` + Status string `json:"status"` + } + if err := json.NewDecoder(resp.Body).Decode(&start); err != nil { + t.Fatal(err) + } + if start.ID == "" { + t.Fatal("expected non-empty id") + } + if start.Status != "running" { + t.Fatalf("want status=running, got %q", start.Status) + } +} + +func TestGetTraversal_NotFound(t *testing.T) { + srv := newTestServer(t) + defer srv.Shutdown(5 * time.Second) //nolint:errcheck + + resp, err := http.Get("http://" + srv.Addr() + "/api/traverse/does-not-exist") + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusNotFound { + t.Fatalf("want 404, got %d", resp.StatusCode) + } +} + +func TestGetTraversal_Found(t *testing.T) { + srv := newTestServer(t) + defer srv.Shutdown(5 * time.Second) //nolint:errcheck + + // Start a job. + startBody := bytes.NewBufferString(`{"domain":"example.com","type":"A"}`) + startResp, err := http.Post("http://"+srv.Addr()+"/api/traverse", "application/json", startBody) + if err != nil { + t.Fatal(err) + } + defer startResp.Body.Close() + + var start struct { + ID string `json:"id"` + } + if err := json.NewDecoder(startResp.Body).Decode(&start); err != nil { + t.Fatal(err) + } + + // Poll for it. + resp, err := http.Get("http://" + srv.Addr() + "/api/traverse/" + start.ID) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + t.Fatalf("want 200, got %d", resp.StatusCode) + } + var job struct { + ID string `json:"id"` + Status string `json:"status"` + } + if err := json.NewDecoder(resp.Body).Decode(&job); err != nil { + t.Fatal(err) + } + if job.ID != start.ID { + t.Fatalf("id mismatch: %q vs %q", job.ID, start.ID) + } +} + +func TestStreamTraversal_NotFound(t *testing.T) { + srv := newTestServer(t) + defer srv.Shutdown(5 * time.Second) //nolint:errcheck + + resp, err := http.Get("http://" + srv.Addr() + "/api/traverse/nope/stream") + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusNotFound { + t.Fatalf("want 404, got %d", resp.StatusCode) + } +} + +func TestStreamTraversal_ContentType(t *testing.T) { + srv := newTestServer(t) + defer srv.Shutdown(5 * time.Second) //nolint:errcheck + + // Start a job. + startBody := bytes.NewBufferString(`{"domain":"example.com","type":"A"}`) + startResp, err := http.Post("http://"+srv.Addr()+"/api/traverse", "application/json", startBody) + if err != nil { + t.Fatal(err) + } + defer startResp.Body.Close() + + var start struct { + ID string `json:"id"` + } + if err := json.NewDecoder(startResp.Body).Decode(&start); err != nil { + t.Fatal(err) + } + + // Open SSE stream. + resp, err := http.Get("http://" + srv.Addr() + "/api/traverse/" + start.ID + "/stream") + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + + if ct := resp.Header.Get("Content-Type"); !strings.HasPrefix(ct, "text/event-stream") { + t.Fatalf("want text/event-stream, got %q", ct) + } + + // Read until we get the done event or timeout. + done := make(chan struct{}) + go func() { + defer close(done) + scanner := bufio.NewScanner(resp.Body) + for scanner.Scan() { + line := scanner.Text() + if strings.HasPrefix(line, "event: done") { + return + } + } + }() + + select { + case <-done: + case <-time.After(30 * time.Second): + t.Fatal("timed out waiting for SSE done event") + } +} + +func TestStaticSPA_Index(t *testing.T) { + srv := newTestServer(t) + defer srv.Shutdown(5 * time.Second) //nolint:errcheck + + resp, err := http.Get("http://" + srv.Addr() + "/") + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + t.Fatalf("want 200, got %d", resp.StatusCode) + } +} + +func TestStaticSPA_FallbackToIndex(t *testing.T) { + srv := newTestServer(t) + defer srv.Shutdown(5 * time.Second) //nolint:errcheck + + resp, err := http.Get("http://" + srv.Addr() + "/some/spa/route") + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + + if resp.StatusCode != http.StatusOK { + t.Fatalf("want 200 (SPA fallback), got %d", resp.StatusCode) + } +} + +// TestHandlerUnit uses httptest.NewRecorder for fast unit tests without +// starting a real listener. +func TestHandlerUnit_Health(t *testing.T) { + w := httptest.NewRecorder() + r := httptest.NewRequest(http.MethodGet, "/api/health", nil) + // Access the unexported handler via the exported Server for coverage. + srv := api.NewServer("127.0.0.1:0") + if err := srv.Start(); err != nil { + t.Fatal(err) + } + defer srv.Shutdown(5 * time.Second) //nolint:errcheck + + resp, err := http.Get(fmt.Sprintf("http://%s/api/health", srv.Addr())) + if err != nil { + t.Fatal(err) + } + defer resp.Body.Close() + + _ = w + _ = r + if resp.StatusCode != http.StatusOK { + t.Fatalf("want 200, got %d", resp.StatusCode) + } +} diff --git a/web/api/server.go b/web/api/server.go new file mode 100644 index 0000000..037270f --- /dev/null +++ b/web/api/server.go @@ -0,0 +1,102 @@ +// Package api provides the HTTP API server for ExploreDNS. +// +// The server exposes DNS traversal as a REST service with: +// - POST /api/traverse — start an asynchronous traversal +// - GET /api/traverse/{id} — poll traversal status and results +// - GET /api/traverse/{id}/stream — Server-Sent Events for live progress +// - GET /api/health — health check +// +// Static frontend assets are embedded at compile time and served from /. +// Unknown paths fall back to index.html to support SPA client-side routing. +package api + +import ( + "context" + "embed" + "fmt" + "io/fs" + "log" + "net" + "net/http" + "time" +) + +//go:embed static +var staticFiles embed.FS + +// Server is the HTTP API server. +type Server struct { + addr string + srv *http.Server +} + +// NewServer creates a new Server that listens on addr (e.g. ":8080"). +func NewServer(addr string) *Server { + return &Server{addr: addr} +} + +// Start builds the HTTP handler, begins listening, and returns when the +// server has accepted its first connection or the address is bound. +// Call Shutdown to stop gracefully. +func (s *Server) Start() error { + h := newHandler() + + sub, err := fs.Sub(staticFiles, "static") + if err != nil { + return fmt.Errorf("static filesystem: %w", err) + } + h.registerStatic(sub) + + s.srv = &http.Server{ + Addr: s.addr, + Handler: corsMiddleware(h.mux), + ReadTimeout: 30 * time.Second, + WriteTimeout: 0, // SSE streams need no write timeout + IdleTimeout: 120 * time.Second, + } + + ln, err := net.Listen("tcp", s.addr) + if err != nil { + return fmt.Errorf("listen %s: %w", s.addr, err) + } + s.addr = ln.Addr().String() + + go func() { + if err := s.srv.Serve(ln); err != nil && err != http.ErrServerClosed { + log.Printf("api server: %v", err) + } + }() + + return nil +} + +// Addr returns the address the server is listening on. Valid after Start. +func (s *Server) Addr() string { + return s.addr +} + +// Shutdown gracefully stops the server, waiting up to timeout for in-flight +// requests to complete. +func (s *Server) Shutdown(timeout time.Duration) error { + if s.srv == nil { + return nil + } + ctx, cancel := context.WithTimeout(context.Background(), timeout) + defer cancel() + return s.srv.Shutdown(ctx) +} + +// corsMiddleware adds CORS headers for cross-origin SPA access. +func corsMiddleware(next http.Handler) http.Handler { + return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { + w.Header().Set("Access-Control-Allow-Origin", "*") + w.Header().Set("Access-Control-Allow-Methods", "GET, POST, OPTIONS") + w.Header().Set("Access-Control-Allow-Headers", "Content-Type") + + if r.Method == http.MethodOptions { + w.WriteHeader(http.StatusNoContent) + return + } + next.ServeHTTP(w, r) + }) +} diff --git a/web/api/static/index.html b/web/api/static/index.html new file mode 100644 index 0000000..a7da368 --- /dev/null +++ b/web/api/static/index.html @@ -0,0 +1,12 @@ + + + + + + ExploreDNS + + +

ExploreDNS

+

Frontend coming soon.

+ +