mirror of
https://github.com/netbirdio/netbird.git
synced 2026-10-02 11:39:06 +02:00
trying embedded caddy reverse proxy
This commit is contained in:
@@ -0,0 +1,626 @@
|
||||
package reverseproxy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"sort"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/caddyserver/caddy/v2"
|
||||
"github.com/caddyserver/caddy/v2/caddyconfig"
|
||||
"github.com/caddyserver/caddy/v2/modules/caddyhttp"
|
||||
"github.com/caddyserver/caddy/v2/modules/caddyhttp/reverseproxy"
|
||||
"github.com/caddyserver/caddy/v2/modules/logging"
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
// CaddyProxy wraps Caddy's reverse proxy functionality
|
||||
type CaddyProxy struct {
|
||||
config Config
|
||||
mu sync.RWMutex
|
||||
isRunning bool
|
||||
routes map[string]*RouteConfig // key is route ID
|
||||
requestCallback RequestDataCallback
|
||||
// customHandlers stores handlers with custom transports that can't be JSON-serialized
|
||||
// key is "routeID:path" to uniquely identify each handler
|
||||
customHandlers map[string]*reverseproxy.Handler
|
||||
}
|
||||
|
||||
// Config holds the reverse proxy configuration
|
||||
type Config struct {
|
||||
// ListenAddress is the address to listen on
|
||||
ListenAddress string
|
||||
|
||||
// EnableHTTPS enables automatic HTTPS with Let's Encrypt
|
||||
EnableHTTPS bool
|
||||
|
||||
// TLSEmail is the email for Let's Encrypt registration
|
||||
TLSEmail string
|
||||
|
||||
// RequestDataCallback is called for each proxied request with metrics
|
||||
RequestDataCallback RequestDataCallback
|
||||
}
|
||||
|
||||
// RouteConfig defines a routing configuration
|
||||
type RouteConfig struct {
|
||||
// ID is a unique identifier for this route
|
||||
ID string
|
||||
|
||||
// Domain is the domain to listen on (e.g., "example.com" or "*" for all)
|
||||
Domain string
|
||||
|
||||
// PathMappings defines paths that should be forwarded to specific ports
|
||||
// Key is the path prefix (e.g., "/", "/api", "/admin")
|
||||
// Value is the target IP:port (e.g., "192.168.1.100:3000")
|
||||
// Must have at least one entry. Use "/" or "" for the default/catch-all route.
|
||||
PathMappings map[string]string
|
||||
|
||||
// Conn is an optional existing network connection to use for this route
|
||||
// This allows routing through specific tunnels (e.g., WireGuard) per route
|
||||
// If set, this connection will be reused for all requests to this route
|
||||
Conn net.Conn
|
||||
|
||||
// CustomDialer is an optional custom dialer for this specific route
|
||||
// This is used if Conn is not set. It allows using different network connections per route
|
||||
CustomDialer func(ctx context.Context, network, address string) (net.Conn, error)
|
||||
}
|
||||
|
||||
// New creates a new Caddy-based reverse proxy
|
||||
func New(config Config) (*CaddyProxy, error) {
|
||||
// Default to port 443 if not specified
|
||||
if config.ListenAddress == "" {
|
||||
config.ListenAddress = ":443"
|
||||
}
|
||||
|
||||
cp := &CaddyProxy{
|
||||
config: config,
|
||||
isRunning: false,
|
||||
routes: make(map[string]*RouteConfig),
|
||||
requestCallback: config.RequestDataCallback,
|
||||
customHandlers: make(map[string]*reverseproxy.Handler),
|
||||
}
|
||||
|
||||
return cp, nil
|
||||
}
|
||||
|
||||
// Start starts the Caddy reverse proxy server
|
||||
func (cp *CaddyProxy) Start() error {
|
||||
cp.mu.Lock()
|
||||
if cp.isRunning {
|
||||
cp.mu.Unlock()
|
||||
return fmt.Errorf("reverse proxy already running")
|
||||
}
|
||||
cp.isRunning = true
|
||||
cp.mu.Unlock()
|
||||
|
||||
// Build Caddy configuration
|
||||
cfg, err := cp.buildCaddyConfig()
|
||||
if err != nil {
|
||||
cp.mu.Lock()
|
||||
cp.isRunning = false
|
||||
cp.mu.Unlock()
|
||||
return fmt.Errorf("failed to build Caddy config: %w", err)
|
||||
}
|
||||
|
||||
// Run Caddy with the configuration
|
||||
err = caddy.Run(cfg)
|
||||
if err != nil {
|
||||
cp.mu.Lock()
|
||||
cp.isRunning = false
|
||||
cp.mu.Unlock()
|
||||
return fmt.Errorf("failed to run Caddy: %w", err)
|
||||
}
|
||||
|
||||
log.Infof("Caddy reverse proxy started on %s", cp.config.ListenAddress)
|
||||
log.Infof("Configured %d route(s)", len(cp.routes))
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Stop gracefully stops the Caddy reverse proxy
|
||||
func (cp *CaddyProxy) Stop(ctx context.Context) error {
|
||||
cp.mu.Lock()
|
||||
if !cp.isRunning {
|
||||
cp.mu.Unlock()
|
||||
return fmt.Errorf("reverse proxy not running")
|
||||
}
|
||||
cp.mu.Unlock()
|
||||
|
||||
log.Info("Stopping Caddy reverse proxy...")
|
||||
|
||||
// Stop Caddy
|
||||
if err := caddy.Stop(); err != nil {
|
||||
return fmt.Errorf("failed to stop Caddy: %w", err)
|
||||
}
|
||||
|
||||
cp.mu.Lock()
|
||||
cp.isRunning = false
|
||||
cp.mu.Unlock()
|
||||
|
||||
log.Info("Caddy reverse proxy stopped")
|
||||
return nil
|
||||
}
|
||||
|
||||
// buildCaddyConfig builds the Caddy configuration
|
||||
func (cp *CaddyProxy) buildCaddyConfig() (*caddy.Config, error) {
|
||||
cp.mu.RLock()
|
||||
defer cp.mu.RUnlock()
|
||||
|
||||
if len(cp.routes) == 0 {
|
||||
// Create a default empty server that returns 404
|
||||
httpServer := &caddyhttp.Server{
|
||||
Listen: []string{cp.config.ListenAddress},
|
||||
Routes: caddyhttp.RouteList{},
|
||||
}
|
||||
|
||||
httpApp := &caddyhttp.App{
|
||||
Servers: map[string]*caddyhttp.Server{
|
||||
"proxy": httpServer,
|
||||
},
|
||||
}
|
||||
|
||||
cfg := &caddy.Config{
|
||||
Admin: &caddy.AdminConfig{
|
||||
Disabled: true,
|
||||
},
|
||||
AppsRaw: caddy.ModuleMap{
|
||||
"http": caddyconfig.JSON(httpApp, nil),
|
||||
},
|
||||
}
|
||||
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
// Build routes grouped by domain
|
||||
domainRoutes := make(map[string][]caddyhttp.Route)
|
||||
// Track unique service IDs for logger configuration
|
||||
serviceIDs := make(map[string]bool)
|
||||
|
||||
for _, routeConfig := range cp.routes {
|
||||
domain := routeConfig.Domain
|
||||
if domain == "" {
|
||||
domain = "*" // wildcard for all domains
|
||||
}
|
||||
|
||||
// Register callback for this service ID
|
||||
if cp.requestCallback != nil {
|
||||
RegisterCallback(routeConfig.ID, cp.requestCallback)
|
||||
serviceIDs[routeConfig.ID] = true
|
||||
}
|
||||
|
||||
// Sort path mappings by path length (longest first) for proper matching
|
||||
// This ensures more specific paths match before catch-all paths
|
||||
paths := make([]string, 0, len(routeConfig.PathMappings))
|
||||
for path := range routeConfig.PathMappings {
|
||||
paths = append(paths, path)
|
||||
}
|
||||
sort.Slice(paths, func(i, j int) bool {
|
||||
// Sort by length descending, but put empty string last (catch-all)
|
||||
if paths[i] == "" || paths[i] == "/" {
|
||||
return false
|
||||
}
|
||||
if paths[j] == "" || paths[j] == "/" {
|
||||
return true
|
||||
}
|
||||
return len(paths[i]) > len(paths[j])
|
||||
})
|
||||
|
||||
// Create routes for each path mapping
|
||||
for _, path := range paths {
|
||||
target := routeConfig.PathMappings[path]
|
||||
route := cp.createRoute(routeConfig, path, target)
|
||||
domainRoutes[domain] = append(domainRoutes[domain], route)
|
||||
}
|
||||
}
|
||||
|
||||
// Build Caddy routes
|
||||
var caddyRoutes caddyhttp.RouteList
|
||||
for domain, routes := range domainRoutes {
|
||||
if domain != "*" {
|
||||
// Add host matcher for specific domains
|
||||
for i := range routes {
|
||||
routes[i].MatcherSetsRaw = []caddy.ModuleMap{
|
||||
{
|
||||
"host": caddyconfig.JSON(caddyhttp.MatchHost{domain}, nil),
|
||||
},
|
||||
}
|
||||
}
|
||||
}
|
||||
caddyRoutes = append(caddyRoutes, routes...)
|
||||
}
|
||||
|
||||
// Create HTTP server with access logging if callback is set
|
||||
httpServer := &caddyhttp.Server{
|
||||
Listen: []string{cp.config.ListenAddress},
|
||||
Routes: caddyRoutes,
|
||||
}
|
||||
|
||||
// Configure server logging if callback is set
|
||||
if cp.requestCallback != nil {
|
||||
httpServer.Logs = &caddyhttp.ServerLogConfig{
|
||||
// Use our custom logger for access logs
|
||||
LoggerNames: map[string]caddyhttp.StringArray{
|
||||
"http.log.access": {"http_access"},
|
||||
},
|
||||
// Disable default access logging (only use custom logger)
|
||||
ShouldLogCredentials: false,
|
||||
}
|
||||
}
|
||||
|
||||
// Disable automatic HTTPS if not enabled
|
||||
if !cp.config.EnableHTTPS {
|
||||
// Explicitly disable automatic HTTPS for the server
|
||||
httpServer.AutoHTTPS = &caddyhttp.AutoHTTPSConfig{
|
||||
Disabled: true,
|
||||
}
|
||||
}
|
||||
|
||||
// Build HTTP app
|
||||
httpApp := &caddyhttp.App{
|
||||
Servers: map[string]*caddyhttp.Server{
|
||||
"proxy": httpServer,
|
||||
},
|
||||
}
|
||||
|
||||
// Provision the HTTP app to set up handlers from JSON
|
||||
ctx, cancel := caddy.NewContext(caddy.Context{Context: context.Background()})
|
||||
defer cancel()
|
||||
|
||||
if err := httpApp.Provision(ctx); err != nil {
|
||||
return nil, fmt.Errorf("failed to provision HTTP app: %w", err)
|
||||
}
|
||||
|
||||
// After provisioning, inject custom transports into handlers
|
||||
// This is done post-provisioning so the Transport field is preserved
|
||||
if err := cp.injectCustomTransports(httpApp); err != nil {
|
||||
return nil, fmt.Errorf("failed to inject custom transports: %w", err)
|
||||
}
|
||||
|
||||
// Create Caddy config with the provisioned app
|
||||
// IMPORTANT: We pass the already-provisioned app, not JSON
|
||||
// This preserves the Transport fields we set
|
||||
cfg := &caddy.Config{
|
||||
Admin: &caddy.AdminConfig{
|
||||
Disabled: true,
|
||||
},
|
||||
// Apps field takes already-provisioned apps
|
||||
Apps: map[string]caddy.App{
|
||||
"http": httpApp,
|
||||
},
|
||||
}
|
||||
|
||||
// Configure logging if callback is set
|
||||
if cp.requestCallback != nil {
|
||||
// Register the callback for the proxy service ID
|
||||
RegisterCallback("proxy", cp.requestCallback)
|
||||
|
||||
// Build logging config with proper module names
|
||||
cfg.Logging = &caddy.Logging{
|
||||
Logs: map[string]*caddy.CustomLog{
|
||||
"http_access": {
|
||||
BaseLog: caddy.BaseLog{
|
||||
WriterRaw: caddyconfig.JSONModuleObject(&CallbackWriter{ServiceID: "proxy"}, "output", "callback", nil),
|
||||
EncoderRaw: caddyconfig.JSONModuleObject(&logging.JSONEncoder{}, "format", "json", nil),
|
||||
Level: "INFO",
|
||||
},
|
||||
Include: []string{"http.log.access"},
|
||||
},
|
||||
},
|
||||
}
|
||||
|
||||
log.Infof("Configured custom logging with callback writer for service: proxy")
|
||||
}
|
||||
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
// createRoute creates a Caddy route for a path and target with service ID tracking
|
||||
func (cp *CaddyProxy) createRoute(routeConfig *RouteConfig, path, target string) caddyhttp.Route {
|
||||
// Check if this route needs a custom transport
|
||||
hasCustomTransport := routeConfig.Conn != nil || routeConfig.CustomDialer != nil
|
||||
|
||||
if hasCustomTransport {
|
||||
// For routes with custom transports, store them separately
|
||||
// and configure the upstream to use a special dial address that we'll intercept
|
||||
handlerKey := fmt.Sprintf("%s:%s", routeConfig.ID, path)
|
||||
|
||||
// Create upstream with custom dial configuration
|
||||
upstream := &reverseproxy.Upstream{
|
||||
Dial: target,
|
||||
}
|
||||
|
||||
// Create the reverse proxy handler with custom transport
|
||||
handler := &reverseproxy.Handler{
|
||||
Upstreams: reverseproxy.UpstreamPool{upstream},
|
||||
}
|
||||
|
||||
// Configure the custom transport
|
||||
if routeConfig.Conn != nil {
|
||||
// Use the provided connection directly
|
||||
transport := &http.Transport{
|
||||
DialContext: func(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
log.Debugf("Reusing existing connection for route %s to %s", routeConfig.ID, address)
|
||||
return routeConfig.Conn, nil
|
||||
},
|
||||
MaxIdleConns: 1,
|
||||
MaxIdleConnsPerHost: 1,
|
||||
IdleConnTimeout: 0,
|
||||
DisableKeepAlives: false,
|
||||
TLSHandshakeTimeout: 10 * time.Second,
|
||||
ExpectContinueTimeout: 1 * time.Second,
|
||||
}
|
||||
handler.Transport = transport
|
||||
log.Infof("Configured net.Conn transport for route %s (path: %s)", routeConfig.ID, path)
|
||||
} else if routeConfig.CustomDialer != nil {
|
||||
// Use the custom dialer function
|
||||
transport := &http.Transport{
|
||||
DialContext: routeConfig.CustomDialer,
|
||||
MaxIdleConns: 100,
|
||||
IdleConnTimeout: 90 * time.Second,
|
||||
TLSHandshakeTimeout: 10 * time.Second,
|
||||
ExpectContinueTimeout: 1 * time.Second,
|
||||
}
|
||||
handler.Transport = transport
|
||||
log.Infof("Configured custom dialer transport for route %s (path: %s)", routeConfig.ID, path)
|
||||
}
|
||||
|
||||
// Store the handler for later injection
|
||||
cp.customHandlers[handlerKey] = handler
|
||||
|
||||
// Create route using HandlersRaw with a placeholder that will be replaced
|
||||
// We'll use JSON serialization here, but inject the real handler after Caddy loads
|
||||
route := caddyhttp.Route{
|
||||
HandlersRaw: []json.RawMessage{
|
||||
caddyconfig.JSONModuleObject(handler, "handler", "reverse_proxy", nil),
|
||||
},
|
||||
}
|
||||
|
||||
if path != "" {
|
||||
route.MatcherSetsRaw = []caddy.ModuleMap{
|
||||
{
|
||||
"path": caddyconfig.JSON(caddyhttp.MatchPath{path + "*"}, nil),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
return route
|
||||
}
|
||||
|
||||
// Standard route without custom transport
|
||||
upstream := &reverseproxy.Upstream{
|
||||
Dial: target,
|
||||
}
|
||||
|
||||
handler := &reverseproxy.Handler{
|
||||
Upstreams: reverseproxy.UpstreamPool{upstream},
|
||||
}
|
||||
|
||||
route := caddyhttp.Route{
|
||||
HandlersRaw: []json.RawMessage{
|
||||
caddyconfig.JSONModuleObject(handler, "handler", "reverse_proxy", nil),
|
||||
},
|
||||
}
|
||||
|
||||
if path != "" {
|
||||
route.MatcherSetsRaw = []caddy.ModuleMap{
|
||||
{
|
||||
"path": caddyconfig.JSON(caddyhttp.MatchPath{path + "*"}, nil),
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
return route
|
||||
}
|
||||
|
||||
// IsRunning returns whether the proxy is running
|
||||
func (cp *CaddyProxy) IsRunning() bool {
|
||||
cp.mu.RLock()
|
||||
defer cp.mu.RUnlock()
|
||||
return cp.isRunning
|
||||
}
|
||||
|
||||
// GetConfig returns the proxy configuration
|
||||
func (cp *CaddyProxy) GetConfig() Config {
|
||||
return cp.config
|
||||
}
|
||||
|
||||
// AddRoute adds a new route configuration to the proxy
|
||||
// If the proxy is running, it will reload the configuration
|
||||
func (cp *CaddyProxy) AddRoute(route *RouteConfig) error {
|
||||
if route == nil {
|
||||
return fmt.Errorf("route cannot be nil")
|
||||
}
|
||||
if route.ID == "" {
|
||||
return fmt.Errorf("route ID is required")
|
||||
}
|
||||
if len(route.PathMappings) == 0 {
|
||||
return fmt.Errorf("route must have at least one path mapping")
|
||||
}
|
||||
|
||||
cp.mu.Lock()
|
||||
// Check if route already exists
|
||||
if _, exists := cp.routes[route.ID]; exists {
|
||||
cp.mu.Unlock()
|
||||
return fmt.Errorf("route with ID %s already exists", route.ID)
|
||||
}
|
||||
|
||||
// Add new route
|
||||
cp.routes[route.ID] = route
|
||||
isRunning := cp.isRunning
|
||||
cp.mu.Unlock()
|
||||
|
||||
log.WithFields(log.Fields{
|
||||
"route_id": route.ID,
|
||||
"domain": route.Domain,
|
||||
"paths": len(route.PathMappings),
|
||||
}).Info("Added route")
|
||||
|
||||
// Reload configuration if proxy is running
|
||||
if isRunning {
|
||||
if err := cp.reloadConfig(); err != nil {
|
||||
// Rollback: remove the route
|
||||
cp.mu.Lock()
|
||||
delete(cp.routes, route.ID)
|
||||
cp.mu.Unlock()
|
||||
return fmt.Errorf("failed to reload config after adding route: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// RemoveRoute removes a route from the proxy
|
||||
// If the proxy is running, it will reload the configuration
|
||||
func (cp *CaddyProxy) RemoveRoute(routeID string) error {
|
||||
cp.mu.Lock()
|
||||
// Check if route exists
|
||||
route, exists := cp.routes[routeID]
|
||||
if !exists {
|
||||
cp.mu.Unlock()
|
||||
return fmt.Errorf("route %s not found", routeID)
|
||||
}
|
||||
|
||||
// Remove route
|
||||
delete(cp.routes, routeID)
|
||||
isRunning := cp.isRunning
|
||||
cp.mu.Unlock()
|
||||
|
||||
log.Infof("Removed route: %s", routeID)
|
||||
|
||||
// Reload configuration if proxy is running
|
||||
if isRunning {
|
||||
if err := cp.reloadConfig(); err != nil {
|
||||
// Rollback: add the route back
|
||||
cp.mu.Lock()
|
||||
cp.routes[routeID] = route
|
||||
cp.mu.Unlock()
|
||||
return fmt.Errorf("failed to reload config after removing route: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// UpdateRoute updates an existing route configuration
|
||||
// If the proxy is running, it will reload the configuration
|
||||
func (cp *CaddyProxy) UpdateRoute(route *RouteConfig) error {
|
||||
if route == nil {
|
||||
return fmt.Errorf("route cannot be nil")
|
||||
}
|
||||
if route.ID == "" {
|
||||
return fmt.Errorf("route ID is required")
|
||||
}
|
||||
|
||||
cp.mu.Lock()
|
||||
// Check if route exists
|
||||
oldRoute, exists := cp.routes[route.ID]
|
||||
if !exists {
|
||||
cp.mu.Unlock()
|
||||
return fmt.Errorf("route %s not found", route.ID)
|
||||
}
|
||||
|
||||
// Update route
|
||||
cp.routes[route.ID] = route
|
||||
isRunning := cp.isRunning
|
||||
cp.mu.Unlock()
|
||||
|
||||
log.WithFields(log.Fields{
|
||||
"route_id": route.ID,
|
||||
"domain": route.Domain,
|
||||
"paths": len(route.PathMappings),
|
||||
}).Info("Updated route")
|
||||
|
||||
// Reload configuration if proxy is running
|
||||
if isRunning {
|
||||
if err := cp.reloadConfig(); err != nil {
|
||||
// Rollback: restore old route
|
||||
cp.mu.Lock()
|
||||
cp.routes[route.ID] = oldRoute
|
||||
cp.mu.Unlock()
|
||||
return fmt.Errorf("failed to reload config after updating route: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// ListRoutes returns a list of all configured route IDs
|
||||
func (cp *CaddyProxy) ListRoutes() []string {
|
||||
cp.mu.RLock()
|
||||
defer cp.mu.RUnlock()
|
||||
|
||||
routes := make([]string, 0, len(cp.routes))
|
||||
for id := range cp.routes {
|
||||
routes = append(routes, id)
|
||||
}
|
||||
return routes
|
||||
}
|
||||
|
||||
// GetRoute returns a route configuration by ID
|
||||
func (cp *CaddyProxy) GetRoute(routeID string) (*RouteConfig, error) {
|
||||
cp.mu.RLock()
|
||||
defer cp.mu.RUnlock()
|
||||
|
||||
route, exists := cp.routes[routeID]
|
||||
if !exists {
|
||||
return nil, fmt.Errorf("route %s not found", routeID)
|
||||
}
|
||||
|
||||
return route, nil
|
||||
}
|
||||
|
||||
// injectCustomTransports injects custom transports into provisioned handlers
|
||||
// This must be called after httpApp.Provision() but before passing to Caddy.Run()
|
||||
func (cp *CaddyProxy) injectCustomTransports(httpApp *caddyhttp.App) error {
|
||||
// Iterate through all servers
|
||||
for serverName, server := range httpApp.Servers {
|
||||
log.Debugf("Injecting custom transports for server: %s", serverName)
|
||||
|
||||
// Iterate through all routes
|
||||
for routeIdx, route := range server.Routes {
|
||||
// Iterate through all handlers in the route
|
||||
for handlerIdx, handler := range route.Handlers {
|
||||
// Check if this is a reverse proxy handler
|
||||
if rpHandler, ok := handler.(*reverseproxy.Handler); ok {
|
||||
// Try to find a matching custom handler for this route
|
||||
// We need to match by handler configuration since we don't have route metadata here
|
||||
for handlerKey, customHandler := range cp.customHandlers {
|
||||
// Check if the upstream configuration matches
|
||||
if len(rpHandler.Upstreams) > 0 && len(customHandler.Upstreams) > 0 {
|
||||
if rpHandler.Upstreams[0].Dial == customHandler.Upstreams[0].Dial {
|
||||
// Match found! Inject the custom transport
|
||||
rpHandler.Transport = customHandler.Transport
|
||||
log.Infof("Injected custom transport for route %d, handler %d (key: %s)", routeIdx, handlerIdx, handlerKey)
|
||||
break
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// reloadConfig rebuilds and reloads the Caddy configuration
|
||||
// Must be called without holding the lock
|
||||
func (cp *CaddyProxy) reloadConfig() error {
|
||||
log.Info("Reloading Caddy configuration...")
|
||||
|
||||
cfg, err := cp.buildCaddyConfig()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to build config: %w", err)
|
||||
}
|
||||
|
||||
if err := caddy.Run(cfg); err != nil {
|
||||
return fmt.Errorf("failed to load config: %w", err)
|
||||
}
|
||||
|
||||
log.Info("Caddy configuration reloaded successfully")
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,225 @@
|
||||
package reverseproxy
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"strings"
|
||||
"sync"
|
||||
|
||||
"github.com/caddyserver/caddy/v2"
|
||||
"github.com/caddyserver/caddy/v2/caddyconfig/caddyfile"
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
var (
|
||||
// Global map to store callbacks per service ID
|
||||
callbackRegistry = make(map[string]RequestDataCallback)
|
||||
callbackMu sync.RWMutex
|
||||
)
|
||||
|
||||
// RegisterCallback registers a callback for a specific service ID
|
||||
func RegisterCallback(serviceID string, callback RequestDataCallback) {
|
||||
callbackMu.Lock()
|
||||
defer callbackMu.Unlock()
|
||||
callbackRegistry[serviceID] = callback
|
||||
}
|
||||
|
||||
// UnregisterCallback removes a callback for a specific service ID
|
||||
func UnregisterCallback(serviceID string) {
|
||||
callbackMu.Lock()
|
||||
defer callbackMu.Unlock()
|
||||
delete(callbackRegistry, serviceID)
|
||||
}
|
||||
|
||||
// getCallback retrieves the callback for a service ID
|
||||
func getCallback(serviceID string) RequestDataCallback {
|
||||
callbackMu.RLock()
|
||||
defer callbackMu.RUnlock()
|
||||
return callbackRegistry[serviceID]
|
||||
}
|
||||
|
||||
func init() {
|
||||
caddy.RegisterModule(CallbackWriter{})
|
||||
}
|
||||
|
||||
// CallbackWriter is a Caddy log writer module that sends request data via callback
|
||||
type CallbackWriter struct {
|
||||
ServiceID string `json:"service_id,omitempty"`
|
||||
}
|
||||
|
||||
// CaddyModule returns the Caddy module information
|
||||
func (CallbackWriter) CaddyModule() caddy.ModuleInfo {
|
||||
return caddy.ModuleInfo{
|
||||
ID: "caddy.logging.writers.callback",
|
||||
New: func() caddy.Module { return new(CallbackWriter) },
|
||||
}
|
||||
}
|
||||
|
||||
// Provision sets up the callback writer
|
||||
func (cw *CallbackWriter) Provision(ctx caddy.Context) error {
|
||||
log.Infof("CallbackWriter.Provision called for service_id: %s", cw.ServiceID)
|
||||
return nil
|
||||
}
|
||||
|
||||
// String returns a human-readable representation of the writer
|
||||
func (cw *CallbackWriter) String() string {
|
||||
return fmt.Sprintf("callback writer for service %s", cw.ServiceID)
|
||||
}
|
||||
|
||||
// WriterKey returns a unique key for this writer configuration
|
||||
func (cw *CallbackWriter) WriterKey() string {
|
||||
return "callback_" + cw.ServiceID
|
||||
}
|
||||
|
||||
// OpenWriter opens the writer
|
||||
func (cw *CallbackWriter) OpenWriter() (io.WriteCloser, error) {
|
||||
log.Infof("CallbackWriter.OpenWriter called for service_id: %s", cw.ServiceID)
|
||||
writer := &LogWriter{
|
||||
serviceID: cw.ServiceID,
|
||||
}
|
||||
log.Infof("Created LogWriter instance: %p for service_id: %s", writer, cw.ServiceID)
|
||||
return writer, nil
|
||||
}
|
||||
|
||||
// UnmarshalCaddyfile implements caddyfile.Unmarshaler
|
||||
func (cw *CallbackWriter) UnmarshalCaddyfile(d *caddyfile.Dispenser) error {
|
||||
for d.Next() {
|
||||
if !d.NextArg() {
|
||||
return d.ArgErr()
|
||||
}
|
||||
cw.ServiceID = d.Val()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// Ensure CallbackWriter implements the required interfaces
|
||||
var (
|
||||
_ caddy.Provisioner = (*CallbackWriter)(nil)
|
||||
_ caddy.WriterOpener = (*CallbackWriter)(nil)
|
||||
_ caddyfile.Unmarshaler = (*CallbackWriter)(nil)
|
||||
)
|
||||
|
||||
// LogWriter is a custom io.Writer that parses Caddy's structured JSON logs
|
||||
// and extracts request metrics to send via callback
|
||||
type LogWriter struct {
|
||||
serviceID string
|
||||
}
|
||||
|
||||
// NewLogWriter creates a new log writer with the given service ID
|
||||
func NewLogWriter(serviceID string) *LogWriter {
|
||||
return &LogWriter{
|
||||
serviceID: serviceID,
|
||||
}
|
||||
}
|
||||
|
||||
// Write implements io.Writer
|
||||
func (lw *LogWriter) Write(p []byte) (n int, err error) {
|
||||
// DEBUG: Log that we received data
|
||||
log.Infof("LogWriter.Write called with %d bytes for service_id: %s", len(p), lw.serviceID)
|
||||
log.Debugf("LogWriter content: %s", string(p))
|
||||
|
||||
// Caddy writes one JSON object per line
|
||||
// Parse the JSON to extract request metrics
|
||||
var logEntry map[string]interface{}
|
||||
if err := json.Unmarshal(p, &logEntry); err != nil {
|
||||
// Not JSON or malformed, skip
|
||||
log.Debugf("Failed to unmarshal JSON: %v", err)
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
// Caddy access logs have a nested "request" object
|
||||
// Check if this is an access log entry by looking for "request" field
|
||||
requestObj, hasRequest := logEntry["request"]
|
||||
if !hasRequest {
|
||||
log.Debugf("Not an access log entry (no 'request' field)")
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
request, ok := requestObj.(map[string]interface{})
|
||||
if !ok {
|
||||
log.Debugf("'request' field is not a map")
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
// Extract fields
|
||||
data := &RequestData{
|
||||
ServiceID: lw.serviceID,
|
||||
}
|
||||
|
||||
// Extract method from request object
|
||||
if method, ok := request["method"].(string); ok {
|
||||
data.Method = method
|
||||
}
|
||||
|
||||
// Extract host from request object and strip port
|
||||
if host, ok := request["host"].(string); ok {
|
||||
// Strip port from host (e.g., "test.netbird.io:54321" -> "test.netbird.io")
|
||||
if idx := strings.LastIndex(host, ":"); idx != -1 {
|
||||
data.Host = host[:idx]
|
||||
} else {
|
||||
data.Host = host
|
||||
}
|
||||
}
|
||||
|
||||
// Extract path (uri field) from request object
|
||||
if uri, ok := request["uri"].(string); ok {
|
||||
data.Path = uri
|
||||
}
|
||||
|
||||
// Extract status code from top-level
|
||||
if status, ok := logEntry["status"].(float64); ok {
|
||||
data.ResponseCode = int32(status)
|
||||
}
|
||||
|
||||
// Extract duration (in seconds, convert to milliseconds) from top-level
|
||||
if duration, ok := logEntry["duration"].(float64); ok {
|
||||
data.DurationMs = int64(duration * 1000)
|
||||
}
|
||||
|
||||
// Extract source IP from request object - try multiple fields
|
||||
if clientIP, ok := request["client_ip"].(string); ok {
|
||||
data.SourceIP = clientIP
|
||||
} else if remoteIP, ok := request["remote_ip"].(string); ok {
|
||||
data.SourceIP = remoteIP
|
||||
} else if remoteAddr, ok := request["remote_addr"].(string); ok {
|
||||
// remote_addr is in "IP:port" format
|
||||
if idx := strings.LastIndex(remoteAddr, ":"); idx != -1 {
|
||||
data.SourceIP = remoteAddr[:idx]
|
||||
} else {
|
||||
data.SourceIP = remoteAddr
|
||||
}
|
||||
}
|
||||
|
||||
// Call callback if set and we have valid data
|
||||
callback := getCallback(lw.serviceID)
|
||||
if callback != nil && data.Method != "" {
|
||||
log.Infof("Calling callback for request: %s %s", data.Method, data.Path)
|
||||
go func() {
|
||||
// Run in goroutine to avoid blocking log writes
|
||||
callback(data)
|
||||
}()
|
||||
} else {
|
||||
log.Warnf("No callback registered for service_id: %s", lw.serviceID)
|
||||
}
|
||||
|
||||
log.WithFields(log.Fields{
|
||||
"service_id": data.ServiceID,
|
||||
"method": data.Method,
|
||||
"host": data.Host,
|
||||
"path": data.Path,
|
||||
"status": data.ResponseCode,
|
||||
"duration_ms": data.DurationMs,
|
||||
"source_ip": data.SourceIP,
|
||||
}).Info("Request logged via callback writer")
|
||||
|
||||
return len(p), nil
|
||||
}
|
||||
|
||||
// Close implements io.Closer (no-op for our use case)
|
||||
func (lw *LogWriter) Close() error {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Ensure LogWriter implements io.WriteCloser
|
||||
var _ io.WriteCloser = (*LogWriter)(nil)
|
||||
@@ -0,0 +1,251 @@
|
||||
package reverseproxy
|
||||
|
||||
import (
|
||||
"encoding/json"
|
||||
"sync"
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestLogWriter_Write(t *testing.T) {
|
||||
// Create a channel to receive callback data
|
||||
callbackChan := make(chan *RequestData, 1)
|
||||
var callbackMu sync.Mutex
|
||||
var callbackCalled bool
|
||||
|
||||
// Register a test callback
|
||||
testServiceID := "test-service"
|
||||
RegisterCallback(testServiceID, func(data *RequestData) {
|
||||
callbackMu.Lock()
|
||||
callbackCalled = true
|
||||
callbackMu.Unlock()
|
||||
callbackChan <- data
|
||||
})
|
||||
defer UnregisterCallback(testServiceID)
|
||||
|
||||
// Create a log writer
|
||||
writer := NewLogWriter(testServiceID)
|
||||
|
||||
// Create a sample Caddy access log entry (matching the structure from your logs)
|
||||
logEntry := map[string]interface{}{
|
||||
"level": "info",
|
||||
"ts": 1768352053.7900746,
|
||||
"logger": "http.log.access",
|
||||
"msg": "handled request",
|
||||
"request": map[string]interface{}{
|
||||
"remote_ip": "::1",
|
||||
"remote_port": "51972",
|
||||
"client_ip": "::1",
|
||||
"proto": "HTTP/1.1",
|
||||
"method": "GET",
|
||||
"host": "test.netbird.io:54321",
|
||||
"uri": "/test/path",
|
||||
},
|
||||
"bytes_read": 0,
|
||||
"user_id": "",
|
||||
"duration": 0.004779453,
|
||||
"size": 615,
|
||||
"status": 200,
|
||||
}
|
||||
|
||||
// Marshal to JSON
|
||||
logJSON, err := json.Marshal(logEntry)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to marshal log entry: %v", err)
|
||||
}
|
||||
|
||||
// Write to the log writer
|
||||
n, err := writer.Write(logJSON)
|
||||
if err != nil {
|
||||
t.Fatalf("Write failed: %v", err)
|
||||
}
|
||||
|
||||
if n != len(logJSON) {
|
||||
t.Errorf("Expected to write %d bytes, wrote %d", len(logJSON), n)
|
||||
}
|
||||
|
||||
// Wait for callback to be called (with timeout)
|
||||
select {
|
||||
case data := <-callbackChan:
|
||||
// Verify the extracted data
|
||||
if data.ServiceID != testServiceID {
|
||||
t.Errorf("Expected service_id %s, got %s", testServiceID, data.ServiceID)
|
||||
}
|
||||
if data.Method != "GET" {
|
||||
t.Errorf("Expected method GET, got %s", data.Method)
|
||||
}
|
||||
if data.Host != "test.netbird.io" {
|
||||
t.Errorf("Expected host test.netbird.io, got %s", data.Host)
|
||||
}
|
||||
if data.Path != "/test/path" {
|
||||
t.Errorf("Expected path /test/path, got %s", data.Path)
|
||||
}
|
||||
if data.ResponseCode != 200 {
|
||||
t.Errorf("Expected status 200, got %d", data.ResponseCode)
|
||||
}
|
||||
if data.SourceIP != "::1" {
|
||||
t.Errorf("Expected source_ip ::1, got %s", data.SourceIP)
|
||||
}
|
||||
// Duration should be ~4.78ms (0.004779453 * 1000)
|
||||
if data.DurationMs < 4 || data.DurationMs > 5 {
|
||||
t.Errorf("Expected duration ~4-5ms, got %dms", data.DurationMs)
|
||||
}
|
||||
case <-time.After(1 * time.Second):
|
||||
t.Fatal("Callback was not called within timeout")
|
||||
}
|
||||
|
||||
// Verify callback was called
|
||||
callbackMu.Lock()
|
||||
defer callbackMu.Unlock()
|
||||
if !callbackCalled {
|
||||
t.Error("Callback was never called")
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogWriter_Write_NonAccessLog(t *testing.T) {
|
||||
// Create a channel to receive callback data
|
||||
callbackChan := make(chan *RequestData, 1)
|
||||
|
||||
// Register a test callback
|
||||
testServiceID := "test-service-2"
|
||||
RegisterCallback(testServiceID, func(data *RequestData) {
|
||||
callbackChan <- data
|
||||
})
|
||||
defer UnregisterCallback(testServiceID)
|
||||
|
||||
// Create a log writer
|
||||
writer := NewLogWriter(testServiceID)
|
||||
|
||||
// Create a non-access log entry (e.g., a TLS log)
|
||||
logEntry := map[string]interface{}{
|
||||
"level": "info",
|
||||
"ts": 1768352032.12347,
|
||||
"logger": "tls",
|
||||
"msg": "storage cleaning happened too recently",
|
||||
}
|
||||
|
||||
// Marshal to JSON
|
||||
logJSON, err := json.Marshal(logEntry)
|
||||
if err != nil {
|
||||
t.Fatalf("Failed to marshal log entry: %v", err)
|
||||
}
|
||||
|
||||
// Write to the log writer
|
||||
n, err := writer.Write(logJSON)
|
||||
if err != nil {
|
||||
t.Fatalf("Write failed: %v", err)
|
||||
}
|
||||
|
||||
if n != len(logJSON) {
|
||||
t.Errorf("Expected to write %d bytes, wrote %d", len(logJSON), n)
|
||||
}
|
||||
|
||||
// Callback should NOT be called for non-access logs
|
||||
select {
|
||||
case data := <-callbackChan:
|
||||
t.Errorf("Callback should not be called for non-access log, but got: %+v", data)
|
||||
case <-time.After(100 * time.Millisecond):
|
||||
// Expected - callback not called
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogWriter_Write_MalformedJSON(t *testing.T) {
|
||||
// Create a log writer
|
||||
writer := NewLogWriter("test-service-3")
|
||||
|
||||
// Write malformed JSON
|
||||
malformedJSON := []byte("{this is not valid json")
|
||||
|
||||
// Should not fail, just skip the entry
|
||||
n, err := writer.Write(malformedJSON)
|
||||
if err != nil {
|
||||
t.Fatalf("Write should not fail on malformed JSON: %v", err)
|
||||
}
|
||||
|
||||
if n != len(malformedJSON) {
|
||||
t.Errorf("Expected to write %d bytes, wrote %d", len(malformedJSON), n)
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallbackRegistry(t *testing.T) {
|
||||
serviceID := "test-registry"
|
||||
var called bool
|
||||
|
||||
// Test registering a callback
|
||||
callback := func(data *RequestData) {
|
||||
called = true
|
||||
}
|
||||
RegisterCallback(serviceID, callback)
|
||||
|
||||
// Test retrieving the callback
|
||||
retrievedCallback := getCallback(serviceID)
|
||||
if retrievedCallback == nil {
|
||||
t.Fatal("Expected to retrieve callback, got nil")
|
||||
}
|
||||
|
||||
// Call the retrieved callback to verify it works
|
||||
retrievedCallback(&RequestData{})
|
||||
if !called {
|
||||
t.Error("Callback was not called")
|
||||
}
|
||||
|
||||
// Test unregistering
|
||||
UnregisterCallback(serviceID)
|
||||
retrievedCallback = getCallback(serviceID)
|
||||
if retrievedCallback != nil {
|
||||
t.Error("Expected nil after unregistering, got a callback")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallbackWriter_Module(t *testing.T) {
|
||||
// Test that the module is properly configured
|
||||
cw := CallbackWriter{ServiceID: "test"}
|
||||
|
||||
moduleInfo := cw.CaddyModule()
|
||||
if moduleInfo.ID != "caddy.logging.writers.callback" {
|
||||
t.Errorf("Expected module ID 'caddy.logging.writers.callback', got '%s'", moduleInfo.ID)
|
||||
}
|
||||
|
||||
if moduleInfo.New == nil {
|
||||
t.Error("Expected New function to be set")
|
||||
}
|
||||
|
||||
// Test creating a new instance via the New function
|
||||
newModule := moduleInfo.New()
|
||||
if newModule == nil {
|
||||
t.Error("Expected New() to return a module instance")
|
||||
}
|
||||
|
||||
_, ok := newModule.(*CallbackWriter)
|
||||
if !ok {
|
||||
t.Error("Expected New() to return a *CallbackWriter")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallbackWriter_WriterKey(t *testing.T) {
|
||||
cw := &CallbackWriter{ServiceID: "my-service"}
|
||||
|
||||
expectedKey := "callback_my-service"
|
||||
if cw.WriterKey() != expectedKey {
|
||||
t.Errorf("Expected writer key '%s', got '%s'", expectedKey, cw.WriterKey())
|
||||
}
|
||||
}
|
||||
|
||||
func TestCallbackWriter_String(t *testing.T) {
|
||||
cw := &CallbackWriter{ServiceID: "my-service"}
|
||||
|
||||
str := cw.String()
|
||||
if str != "callback writer for service my-service" {
|
||||
t.Errorf("Unexpected string representation: %s", str)
|
||||
}
|
||||
}
|
||||
|
||||
func TestLogWriter_Close(t *testing.T) {
|
||||
writer := NewLogWriter("test")
|
||||
|
||||
// Close should not fail
|
||||
err := writer.Close()
|
||||
if err != nil {
|
||||
t.Errorf("Close should not fail: %v", err)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,131 @@
|
||||
package reverseproxy
|
||||
|
||||
import (
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/caddyserver/caddy/v2/modules/caddyhttp"
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
// RequestDataCallback is called for each request that passes through the proxy
|
||||
type RequestDataCallback func(data *RequestData)
|
||||
|
||||
// RequestData contains metadata about a proxied request
|
||||
type RequestData struct {
|
||||
ServiceID string
|
||||
Host string
|
||||
Path string
|
||||
DurationMs int64
|
||||
Method string
|
||||
ResponseCode int32
|
||||
SourceIP string
|
||||
}
|
||||
|
||||
// MetricsMiddleware wraps a handler to capture request metrics
|
||||
type MetricsMiddleware struct {
|
||||
Next caddyhttp.Handler
|
||||
ServiceID string
|
||||
Callback RequestDataCallback
|
||||
}
|
||||
|
||||
// ServeHTTP implements caddyhttp.MiddlewareHandler
|
||||
func (m *MetricsMiddleware) ServeHTTP(w http.ResponseWriter, r *http.Request, next caddyhttp.Handler) error {
|
||||
// Record start time
|
||||
startTime := time.Now()
|
||||
|
||||
// Wrap the response writer to capture status code
|
||||
wrappedWriter := &responseWriterWrapper{
|
||||
ResponseWriter: w,
|
||||
statusCode: http.StatusOK, // Default to 200
|
||||
}
|
||||
|
||||
// Call the next handler (Caddy's reverse proxy)
|
||||
err := next.ServeHTTP(wrappedWriter, r)
|
||||
|
||||
// Calculate duration
|
||||
duration := time.Since(startTime)
|
||||
|
||||
// Extract source IP (handle X-Forwarded-For or direct connection)
|
||||
sourceIP := extractSourceIP(r)
|
||||
|
||||
// Create request data
|
||||
data := &RequestData{
|
||||
ServiceID: m.ServiceID,
|
||||
Path: r.URL.Path,
|
||||
DurationMs: duration.Milliseconds(),
|
||||
Method: r.Method,
|
||||
ResponseCode: int32(wrappedWriter.statusCode),
|
||||
SourceIP: sourceIP,
|
||||
}
|
||||
|
||||
// Call callback if set
|
||||
if m.Callback != nil {
|
||||
go func() {
|
||||
// Run callback in goroutine to avoid blocking response
|
||||
m.Callback(data)
|
||||
}()
|
||||
}
|
||||
|
||||
log.WithFields(log.Fields{
|
||||
"service_id": data.ServiceID,
|
||||
"method": data.Method,
|
||||
"path": data.Path,
|
||||
"status": data.ResponseCode,
|
||||
"duration_ms": data.DurationMs,
|
||||
"source_ip": data.SourceIP,
|
||||
}).Debug("Request proxied")
|
||||
|
||||
return err
|
||||
}
|
||||
|
||||
// responseWriterWrapper wraps http.ResponseWriter to capture status code
|
||||
type responseWriterWrapper struct {
|
||||
http.ResponseWriter
|
||||
statusCode int
|
||||
written bool
|
||||
}
|
||||
|
||||
// WriteHeader captures the status code
|
||||
func (w *responseWriterWrapper) WriteHeader(statusCode int) {
|
||||
if !w.written {
|
||||
w.statusCode = statusCode
|
||||
w.written = true
|
||||
}
|
||||
w.ResponseWriter.WriteHeader(statusCode)
|
||||
}
|
||||
|
||||
// Write ensures we capture status if WriteHeader wasn't called explicitly
|
||||
func (w *responseWriterWrapper) Write(b []byte) (int, error) {
|
||||
if !w.written {
|
||||
w.written = true
|
||||
// Status code defaults to 200 if not explicitly set
|
||||
}
|
||||
return w.ResponseWriter.Write(b)
|
||||
}
|
||||
|
||||
// extractSourceIP extracts the real client IP from the request
|
||||
func extractSourceIP(r *http.Request) string {
|
||||
// Check X-Forwarded-For header first (if behind a proxy)
|
||||
if xff := r.Header.Get("X-Forwarded-For"); xff != "" {
|
||||
// X-Forwarded-For can be a comma-separated list, take the first one
|
||||
parts := strings.Split(xff, ",")
|
||||
if len(parts) > 0 {
|
||||
return strings.TrimSpace(parts[0])
|
||||
}
|
||||
}
|
||||
|
||||
// Check X-Real-IP header
|
||||
if xri := r.Header.Get("X-Real-IP"); xri != "" {
|
||||
return xri
|
||||
}
|
||||
|
||||
// Fall back to RemoteAddr
|
||||
// RemoteAddr is in format "IP:port", so we need to strip the port
|
||||
if idx := strings.LastIndex(r.RemoteAddr, ":"); idx != -1 {
|
||||
return r.RemoteAddr[:idx]
|
||||
}
|
||||
|
||||
return r.RemoteAddr
|
||||
}
|
||||
@@ -0,0 +1,139 @@
|
||||
package reverseproxy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"net"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
log "github.com/sirupsen/logrus"
|
||||
)
|
||||
|
||||
// customTransportRegistry stores custom dialers and connections globally
|
||||
// This allows them to be accessed after Caddy deserializes the configuration from JSON
|
||||
var customTransportRegistry = &transportRegistry{
|
||||
transports: make(map[string]*customTransport),
|
||||
}
|
||||
|
||||
// transportRegistry manages custom transports for routes
|
||||
type transportRegistry struct {
|
||||
mu sync.RWMutex
|
||||
transports map[string]*customTransport // key is "routeID:path"
|
||||
}
|
||||
|
||||
// customTransport wraps either a net.Conn or a custom dialer
|
||||
type customTransport struct {
|
||||
routeID string
|
||||
path string
|
||||
conn net.Conn
|
||||
customDialer func(ctx context.Context, network, address string) (net.Conn, error)
|
||||
defaultDialer *net.Dialer
|
||||
}
|
||||
|
||||
// Register registers a custom transport for a route
|
||||
func (r *transportRegistry) Register(routeID, path string, conn net.Conn, dialer func(ctx context.Context, network, address string) (net.Conn, error)) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
key := fmt.Sprintf("%s:%s", routeID, path)
|
||||
r.transports[key] = &customTransport{
|
||||
routeID: routeID,
|
||||
path: path,
|
||||
conn: conn,
|
||||
customDialer: dialer,
|
||||
defaultDialer: &net.Dialer{Timeout: 30 * time.Second},
|
||||
}
|
||||
|
||||
if conn != nil {
|
||||
log.Infof("Registered net.Conn transport for route %s (path: %s)", routeID, path)
|
||||
} else if dialer != nil {
|
||||
log.Infof("Registered custom dialer transport for route %s (path: %s)", routeID, path)
|
||||
}
|
||||
}
|
||||
|
||||
// Get retrieves a custom transport for a route
|
||||
func (r *transportRegistry) Get(routeID, path string) *customTransport {
|
||||
r.mu.RLock()
|
||||
defer r.mu.RUnlock()
|
||||
|
||||
key := fmt.Sprintf("%s:%s", routeID, path)
|
||||
return r.transports[key]
|
||||
}
|
||||
|
||||
// Unregister removes a custom transport
|
||||
func (r *transportRegistry) Unregister(routeID, path string) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
key := fmt.Sprintf("%s:%s", routeID, path)
|
||||
delete(r.transports, key)
|
||||
log.Infof("Unregistered transport for route %s (path: %s)", routeID, path)
|
||||
}
|
||||
|
||||
// Clear removes all custom transports
|
||||
func (r *transportRegistry) Clear() {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
r.transports = make(map[string]*customTransport)
|
||||
log.Info("Cleared all custom transports")
|
||||
}
|
||||
|
||||
// DialContext implements the DialContext function for custom transports
|
||||
func (ct *customTransport) DialContext(ctx context.Context, network, address string) (net.Conn, error) {
|
||||
// If we have a pre-existing connection, return it
|
||||
if ct.conn != nil {
|
||||
log.Debugf("Reusing existing connection for route %s (path: %s) to %s", ct.routeID, ct.path, address)
|
||||
return ct.conn, nil
|
||||
}
|
||||
|
||||
// If we have a custom dialer, use it
|
||||
if ct.customDialer != nil {
|
||||
log.Debugf("Using custom dialer for route %s (path: %s) to %s", ct.routeID, ct.path, address)
|
||||
return ct.customDialer(ctx, network, address)
|
||||
}
|
||||
|
||||
// Fallback to default dialer (this shouldn't happen if registered correctly)
|
||||
log.Warnf("No custom transport found for route %s (path: %s), using default dialer", ct.routeID, ct.path)
|
||||
return ct.defaultDialer.DialContext(ctx, network, address)
|
||||
}
|
||||
|
||||
// NewCustomHTTPTransport creates an HTTP transport that uses the custom dialer
|
||||
func NewCustomHTTPTransport(routeID, path string) *http.Transport {
|
||||
transport := customTransportRegistry.Get(routeID, path)
|
||||
if transport == nil {
|
||||
// No custom transport registered, return standard transport
|
||||
log.Warnf("No custom transport found for route %s (path: %s), using standard transport", routeID, path)
|
||||
return &http.Transport{
|
||||
MaxIdleConns: 100,
|
||||
IdleConnTimeout: 90 * time.Second,
|
||||
TLSHandshakeTimeout: 10 * time.Second,
|
||||
ExpectContinueTimeout: 1 * time.Second,
|
||||
}
|
||||
}
|
||||
|
||||
// Configure transport based on whether we're using a connection or dialer
|
||||
if transport.conn != nil {
|
||||
// Using a pre-existing connection - disable pooling
|
||||
return &http.Transport{
|
||||
DialContext: transport.DialContext,
|
||||
MaxIdleConns: 1,
|
||||
MaxIdleConnsPerHost: 1,
|
||||
IdleConnTimeout: 0, // Keep alive indefinitely
|
||||
DisableKeepAlives: false,
|
||||
TLSHandshakeTimeout: 10 * time.Second,
|
||||
ExpectContinueTimeout: 1 * time.Second,
|
||||
}
|
||||
}
|
||||
|
||||
// Using a custom dialer - use normal pooling
|
||||
return &http.Transport{
|
||||
DialContext: transport.DialContext,
|
||||
MaxIdleConns: 100,
|
||||
IdleConnTimeout: 90 * time.Second,
|
||||
TLSHandshakeTimeout: 10 * time.Second,
|
||||
ExpectContinueTimeout: 1 * time.Second,
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user