This commit is contained in:
@@ -0,0 +1,65 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"net"
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"strings"
|
||||
)
|
||||
|
||||
// ClientIP returns the client address for access control, using X-Forwarded-For /
|
||||
// X-Real-IP only when the direct peer is in trusted CIDRs.
|
||||
func ClientIP(r *http.Request, trusted []netip.Prefix) netip.Addr {
|
||||
host, _, err := net.SplitHostPort(r.RemoteAddr)
|
||||
if err != nil {
|
||||
host = r.RemoteAddr
|
||||
}
|
||||
peer, err := netip.ParseAddr(host)
|
||||
if err != nil {
|
||||
return netip.Addr{}
|
||||
}
|
||||
if !containsIP(trusted, peer) {
|
||||
return peer
|
||||
}
|
||||
if xff := r.Header.Get("X-Forwarded-For"); xff != "" {
|
||||
parts := strings.Split(xff, ",")
|
||||
for _, p := range parts {
|
||||
p = strings.TrimSpace(p)
|
||||
if p == "" {
|
||||
continue
|
||||
}
|
||||
if a, err := netip.ParseAddr(p); err == nil {
|
||||
return a
|
||||
}
|
||||
}
|
||||
}
|
||||
if xr := strings.TrimSpace(r.Header.Get("X-Real-IP")); xr != "" {
|
||||
if a, err := netip.ParseAddr(xr); err == nil {
|
||||
return a
|
||||
}
|
||||
}
|
||||
return peer
|
||||
}
|
||||
|
||||
func containsIP(prefixes []netip.Prefix, addr netip.Addr) bool {
|
||||
for _, p := range prefixes {
|
||||
if p.Contains(addr) {
|
||||
return true
|
||||
}
|
||||
}
|
||||
return false
|
||||
}
|
||||
|
||||
// Allowed reports whether addr matches whitelist rules.
|
||||
func Allowed(addr netip.Addr, allowAll bool, whitelist []netip.Prefix) bool {
|
||||
if !addr.IsValid() {
|
||||
return false
|
||||
}
|
||||
if allowAll {
|
||||
return true
|
||||
}
|
||||
if len(whitelist) == 0 {
|
||||
return false
|
||||
}
|
||||
return containsIP(whitelist, addr)
|
||||
}
|
||||
@@ -0,0 +1,34 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"net/netip"
|
||||
"testing"
|
||||
)
|
||||
|
||||
func TestAllowed(t *testing.T) {
|
||||
p, _ := netip.ParsePrefix("127.0.0.1/32")
|
||||
a := netip.MustParseAddr("127.0.0.1")
|
||||
if !Allowed(a, false, []netip.Prefix{p}) {
|
||||
t.Fatal("expected allowed")
|
||||
}
|
||||
if Allowed(a, false, nil) {
|
||||
t.Fatal("empty whitelist should deny")
|
||||
}
|
||||
if !Allowed(a, true, nil) {
|
||||
t.Fatal("allow_all")
|
||||
}
|
||||
}
|
||||
|
||||
func TestClientIPTrustedXFF(t *testing.T) {
|
||||
trusted, _ := netip.ParsePrefix("10.0.0.1/32")
|
||||
r := &http.Request{
|
||||
Header: http.Header{},
|
||||
RemoteAddr: "10.0.0.1:12345",
|
||||
}
|
||||
r.Header.Set("X-Forwarded-For", "203.0.113.5, 10.0.0.1")
|
||||
ip := ClientIP(r, []netip.Prefix{trusted})
|
||||
if ip.String() != "203.0.113.5" {
|
||||
t.Fatalf("got %v", ip)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,221 @@
|
||||
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/config"
|
||||
"github.com/telemt/telemt-api/internal/proxy"
|
||||
)
|
||||
|
||||
// Gateway serves health, metrics, and proxied API routes.
|
||||
type Gateway struct {
|
||||
parsed *config.Parsed
|
||||
proxies map[string]*httputil.ReverseProxy
|
||||
log *slog.Logger
|
||||
transport *http.Transport
|
||||
promHandler http.Handler
|
||||
}
|
||||
|
||||
// NewGateway builds handlers and reverse proxies from parsed config.
|
||||
func NewGateway(p *config.Parsed, log *slog.Logger) (*Gateway, error) {
|
||||
t := &http.Transport{
|
||||
Proxy: http.ProxyFromEnvironment,
|
||||
MaxIdleConns: 64,
|
||||
IdleConnTimeout: 90 * time.Second,
|
||||
TLSHandshakeTimeout: 10 * time.Second,
|
||||
ExpectContinueTimeout: 1 * time.Second,
|
||||
DialContext: (&net.Dialer{Timeout: 5 * time.Second, KeepAlive: 30 * time.Second}).DialContext,
|
||||
ResponseHeaderTimeout: 120 * time.Second,
|
||||
}
|
||||
g := &Gateway{
|
||||
parsed: p,
|
||||
proxies: make(map[string]*httputil.ReverseProxy),
|
||||
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
|
||||
}
|
||||
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.withWhitelist(h)
|
||||
h = g.withAccessLog(h)
|
||||
h = g.withMetrics(h)
|
||||
return h
|
||||
}
|
||||
|
||||
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
|
||||
}
|
||||
const prefix = "/api/"
|
||||
if !strings.HasPrefix(r.URL.Path, prefix) {
|
||||
http.NotFound(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()
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,31 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"strconv"
|
||||
"time"
|
||||
|
||||
"github.com/prometheus/client_golang/prometheus"
|
||||
"github.com/prometheus/client_golang/prometheus/promauto"
|
||||
)
|
||||
|
||||
var (
|
||||
httpInFlight = promauto.NewGauge(prometheus.GaugeOpts{
|
||||
Name: "telemt_gateway_http_in_flight",
|
||||
Help: "Current requests being served.",
|
||||
})
|
||||
httpRequests = promauto.NewCounterVec(prometheus.CounterOpts{
|
||||
Name: "telemt_gateway_http_requests_total",
|
||||
Help: "HTTP requests by status, method, alias.",
|
||||
}, []string{"code", "method", "alias"})
|
||||
httpDuration = promauto.NewHistogramVec(prometheus.HistogramOpts{
|
||||
Name: "telemt_gateway_http_request_duration_seconds",
|
||||
Help: "Request duration in seconds.",
|
||||
Buckets: prometheus.DefBuckets,
|
||||
}, []string{"method", "alias"})
|
||||
)
|
||||
|
||||
func observeRequest(method, alias string, status int, started time.Time) {
|
||||
httpInFlight.Dec()
|
||||
httpRequests.WithLabelValues(strconv.Itoa(status), method, alias).Inc()
|
||||
httpDuration.WithLabelValues(method, alias).Observe(time.Since(started).Seconds())
|
||||
}
|
||||
@@ -0,0 +1,14 @@
|
||||
package server
|
||||
|
||||
import (
|
||||
"crypto/rand"
|
||||
"encoding/hex"
|
||||
)
|
||||
|
||||
func randomID() string {
|
||||
var b [16]byte
|
||||
if _, err := rand.Read(b[:]); err != nil {
|
||||
return "unknown"
|
||||
}
|
||||
return hex.EncodeToString(b[:])
|
||||
}
|
||||
Reference in New Issue
Block a user