package server import ( "context" "encoding/json" "log/slog" "net" "net/http" "net/http/httputil" "net/url" "strings" "time" "github.com/prometheus/client_golang/prometheus/promhttp" "github.com/telemt/telemt-api/internal/aggregate" "github.com/telemt/telemt-api/internal/config" "github.com/telemt/telemt-api/internal/geoip" "github.com/telemt/telemt-api/internal/proxy" "github.com/telemt/telemt-api/internal/webui" ) // Gateway serves health, metrics, and proxied API routes. type Gateway struct { parsed *config.Parsed proxies map[string]*httputil.ReverseProxy agg *aggregate.Handler geo *geoip.Service log *slog.Logger transport *http.Transport promHandler http.Handler corsAllowed []string webUI http.Handler } // NewGateway builds handlers and reverse proxies from parsed config. func NewGateway(p *config.Parsed, log *slog.Logger, geo *geoip.Service) (*Gateway, error) { t := proxy.DirectTransport() t.MaxIdleConns = 64 t.IdleConnTimeout = 90 * time.Second t.TLSHandshakeTimeout = 10 * time.Second t.ExpectContinueTimeout = 1 * time.Second t.DialContext = (&net.Dialer{Timeout: 5 * time.Second, KeepAlive: 30 * time.Second}).DialContext t.ResponseHeaderTimeout = 120 * time.Second g := &Gateway{ parsed: p, proxies: make(map[string]*httputil.ReverseProxy), geo: geo, log: log, transport: t, promHandler: promhttp.Handler(), } for i := range p.Config.Servers { s := &p.Config.Servers[i] u, err := url.Parse(s.BaseURL) if err != nil { return nil, err } auth := p.AuthByAlias[s.Alias] strip := "/api/" + s.Alias rp := proxy.NewReverseProxy(u, strip, s.PathPrefix, auth) rp.Transport = t rp.ErrorHandler = func(w http.ResponseWriter, r *http.Request, err error) { log.Error("upstream error", "alias", s.Alias, "err", err) w.Header().Set("Content-Type", "application/json; charset=utf-8") w.WriteHeader(http.StatusBadGateway) _ = json.NewEncoder(w).Encode(map[string]any{ "ok": false, "error": map[string]string{"code": "bad_gateway", "message": "upstream unreachable"}, }) } g.proxies[s.Alias] = rp } var aggCacheTTL time.Duration if p.Config.Aggregate != nil && p.Config.Aggregate.CacheTTLMs > 0 { aggCacheTTL = time.Duration(p.Config.Aggregate.CacheTTLMs) * time.Millisecond } g.corsAllowed = append([]string(nil), p.Config.CorsAllowedOrigins...) g.agg = aggregate.NewHandler(p, &http.Client{Transport: t}, geo, aggCacheTTL) g.webUI = webui.Handler() return g, nil } // Handler returns the root HTTP handler with middleware. func (g *Gateway) Handler() http.Handler { var h http.Handler = http.HandlerFunc(g.serve) h = g.withCORS(h) h = g.withWhitelist(h) h = g.withAccessLog(h) h = g.withMetrics(h) return h } func (g *Gateway) withCORS(next http.Handler) http.Handler { if len(g.corsAllowed) == 0 { return next } return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Add("Vary", "Origin") origin := r.Header.Get("Origin") ok, allowOrigin := corsMatch(g.corsAllowed, origin) if ok { w.Header().Set("Access-Control-Allow-Origin", allowOrigin) w.Header().Set("Access-Control-Allow-Methods", "GET, HEAD, POST, PUT, PATCH, DELETE, OPTIONS") w.Header().Set("Access-Control-Allow-Headers", "Authorization, Content-Type, X-Request-Id") w.Header().Set("Access-Control-Max-Age", "86400") } if r.Method == http.MethodOptions { if ok { w.WriteHeader(http.StatusNoContent) return } } next.ServeHTTP(w, r) }) } func corsMatch(allowed []string, origin string) (ok bool, allowOrigin string) { if origin == "" { return false, "" } for _, a := range allowed { a = strings.TrimSpace(a) if a == "" { continue } if a == "*" { return true, "*" } if strings.EqualFold(a, origin) { return true, origin } } return false, "" } func (g *Gateway) withWhitelist(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path == "/health" { next.ServeHTTP(w, r) return } ip := ClientIP(r, g.parsed.Trusted) if !Allowed(ip, g.parsed.Config.AllowAll, g.parsed.Whitelist) { w.Header().Set("Content-Type", "application/json; charset=utf-8") w.WriteHeader(http.StatusForbidden) _ = json.NewEncoder(w).Encode(map[string]any{ "ok": false, "error": map[string]string{"code": "forbidden", "message": "source address not allowed"}, }) return } next.ServeHTTP(w, r) }) } func (g *Gateway) withAccessLog(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { rid := r.Header.Get("X-Request-Id") if rid == "" { rid = randomID() r.Header.Set("X-Request-Id", rid) } w.Header().Set("X-Request-Id", rid) start := time.Now() lw := &statusWriter{ResponseWriter: w, status: http.StatusOK} next.ServeHTTP(lw, r) g.log.Info("request", "request_id", rid, "method", r.Method, "path", r.URL.Path, "status", lw.status, "duration_ms", time.Since(start).Milliseconds(), "remote", r.RemoteAddr, ) }) } func (g *Gateway) withMetrics(next http.Handler) http.Handler { return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { if r.URL.Path == "/health" { next.ServeHTTP(w, r) return } httpInFlight.Inc() start := time.Now() alias := routeAlias(r.URL.Path) lw := &statusWriter{ResponseWriter: w, status: http.StatusOK} defer observeRequest(r.Method, alias, lw.status, start) next.ServeHTTP(lw, r) }) } func routeAlias(path string) string { const pfx = "/api/" if !strings.HasPrefix(path, pfx) { if path == "/metrics" { return "metrics" } return "_" } rest := strings.TrimPrefix(path, pfx) if rest == "" { return "_" } i := strings.IndexByte(rest, '/') if i < 0 { return rest } return rest[:i] } func (g *Gateway) serve(w http.ResponseWriter, r *http.Request) { switch r.URL.Path { case "/health": if r.Method != http.MethodGet { http.Error(w, "method not allowed", http.StatusMethodNotAllowed) return } w.Header().Set("Content-Type", "application/json; charset=utf-8") _ = json.NewEncoder(w).Encode(map[string]any{"status": "ok"}) return case "/metrics": if r.Method != http.MethodGet { http.Error(w, "method not allowed", http.StatusMethodNotAllowed) return } g.promHandler.ServeHTTP(w, r) return } // So /api//mtg/health and /api/../api/agg/... route like /api/mtg/health and /api/agg/... if strings.HasPrefix(r.URL.Path, "/api") { proxy.NormalizeRequestURLPath(r) } const prefix = "/api/" if r.URL.Path == "/api/agg" || strings.HasPrefix(r.URL.Path, "/api/agg/") { g.agg.ServeHTTP(w, r) return } if !strings.HasPrefix(r.URL.Path, prefix) { g.webUI.ServeHTTP(w, r) return } trim := strings.TrimPrefix(r.URL.Path, prefix) if trim == "" { http.NotFound(w, r) return } var alias string if i := strings.IndexByte(trim, '/'); i >= 0 { alias = trim[:i] } else { alias = trim } if alias == "" { http.NotFound(w, r) return } rp, ok := g.proxies[alias] if !ok { w.Header().Set("Content-Type", "application/json; charset=utf-8") w.WriteHeader(http.StatusNotFound) _ = json.NewEncoder(w).Encode(map[string]any{ "ok": false, "error": map[string]string{"code": "not_found", "message": "unknown alias"}, }) return } rp.ServeHTTP(w, r) } type statusWriter struct { http.ResponseWriter status int } func (s *statusWriter) WriteHeader(code int) { s.status = code s.ResponseWriter.WriteHeader(code) } // Shutdown idle connections on the shared transport. func (g *Gateway) Shutdown(ctx context.Context) error { g.transport.CloseIdleConnections() if g.geo != nil { _ = g.geo.Close() } return nil }