feat(gateway): add bind fallback and CIDR policy enforcement
This commit is contained in:
parent
b114dcaeb1
commit
0a775d9fb9
9 changed files with 587 additions and 20 deletions
|
|
@ -465,9 +465,11 @@
|
||||||
},
|
},
|
||||||
"gateway": {
|
"gateway": {
|
||||||
"_comment": "Default log level is set to 'fatal'. Other available options are 'debug', 'info', 'warn' and 'error'.",
|
"_comment": "Default log level is set to 'fatal'. Other available options are 'debug', 'info', 'warn' and 'error'.",
|
||||||
|
"_comment_allowed_cidrs": "Optional CIDR allowlist for gateway HTTP endpoints. Empty means no CIDR restriction unless automatic fallback enables a generated local allowlist.",
|
||||||
"host": "127.0.0.1",
|
"host": "127.0.0.1",
|
||||||
"port": 18790,
|
"port": 18790,
|
||||||
"hot_reload": false,
|
"hot_reload": false,
|
||||||
"log_level": "fatal"
|
"log_level": "fatal",
|
||||||
|
"allowed_cidrs": []
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
|
||||||
|
|
@ -98,6 +98,10 @@ type asyncTask struct {
|
||||||
cancel context.CancelFunc
|
cancel context.CancelFunc
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// HTTPMiddleware wraps an HTTP handler and may return an error if middleware
|
||||||
|
// configuration is invalid.
|
||||||
|
type HTTPMiddleware func(http.Handler) (http.Handler, error)
|
||||||
|
|
||||||
// RecordPlaceholder registers a placeholder message for later editing.
|
// RecordPlaceholder registers a placeholder message for later editing.
|
||||||
// Implements PlaceholderRecorder.
|
// Implements PlaceholderRecorder.
|
||||||
func (m *Manager) RecordPlaceholder(channel, chatID, placeholderID string) {
|
func (m *Manager) RecordPlaceholder(channel, chatID, placeholderID string) {
|
||||||
|
|
@ -436,7 +440,7 @@ func (m *Manager) initChannels(channels *config.ChannelsConfig) error {
|
||||||
// SetupHTTPServer creates a shared HTTP server with the given listen address.
|
// SetupHTTPServer creates a shared HTTP server with the given listen address.
|
||||||
// It registers health endpoints from the health server and discovers channels
|
// It registers health endpoints from the health server and discovers channels
|
||||||
// that implement WebhookHandler and/or HealthChecker to register their handlers.
|
// that implement WebhookHandler and/or HealthChecker to register their handlers.
|
||||||
func (m *Manager) SetupHTTPServer(addr string, healthServer *health.Server) {
|
func (m *Manager) SetupHTTPServer(addr string, healthServer *health.Server, middlewares ...HTTPMiddleware) error {
|
||||||
m.mux = newDynamicServeMux()
|
m.mux = newDynamicServeMux()
|
||||||
|
|
||||||
// Register health endpoints
|
// Register health endpoints
|
||||||
|
|
@ -447,12 +451,26 @@ func (m *Manager) SetupHTTPServer(addr string, healthServer *health.Server) {
|
||||||
// Discover and register webhook handlers and health checkers
|
// Discover and register webhook handlers and health checkers
|
||||||
m.registerHTTPHandlersLocked()
|
m.registerHTTPHandlersLocked()
|
||||||
|
|
||||||
|
handler := http.Handler(m.mux)
|
||||||
|
for idx, mw := range middlewares {
|
||||||
|
if mw == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
wrapped, err := mw(handler)
|
||||||
|
if err != nil {
|
||||||
|
return fmt.Errorf("apply HTTP middleware #%d: %w", idx, err)
|
||||||
|
}
|
||||||
|
handler = wrapped
|
||||||
|
}
|
||||||
|
|
||||||
m.httpServer = &http.Server{
|
m.httpServer = &http.Server{
|
||||||
Addr: addr,
|
Addr: addr,
|
||||||
Handler: m.mux,
|
Handler: handler,
|
||||||
ReadTimeout: 30 * time.Second,
|
ReadTimeout: 30 * time.Second,
|
||||||
WriteTimeout: 30 * time.Second,
|
WriteTimeout: 30 * time.Second,
|
||||||
}
|
}
|
||||||
|
|
||||||
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
// registerHTTPHandlersLocked registers webhook and health-check handlers for
|
// registerHTTPHandlersLocked registers webhook and health-check handlers for
|
||||||
|
|
|
||||||
|
|
@ -358,10 +358,11 @@ func DefaultConfig() *Config {
|
||||||
},
|
},
|
||||||
},
|
},
|
||||||
Gateway: GatewayConfig{
|
Gateway: GatewayConfig{
|
||||||
Host: "127.0.0.1",
|
Host: "127.0.0.1",
|
||||||
Port: 18790,
|
Port: 18790,
|
||||||
HotReload: false,
|
HotReload: false,
|
||||||
LogLevel: DefaultGatewayLogLevel,
|
LogLevel: DefaultGatewayLogLevel,
|
||||||
|
AllowedCIDRs: nil,
|
||||||
},
|
},
|
||||||
Tools: ToolsConfig{
|
Tools: ToolsConfig{
|
||||||
FilterSensitiveData: true,
|
FilterSensitiveData: true,
|
||||||
|
|
|
||||||
|
|
@ -10,10 +10,11 @@ import (
|
||||||
const DefaultGatewayLogLevel = "warn"
|
const DefaultGatewayLogLevel = "warn"
|
||||||
|
|
||||||
type GatewayConfig struct {
|
type GatewayConfig struct {
|
||||||
Host string `json:"host" env:"PICOCLAW_GATEWAY_HOST"`
|
Host string `json:"host" env:"PICOCLAW_GATEWAY_HOST"`
|
||||||
Port int `json:"port" env:"PICOCLAW_GATEWAY_PORT"`
|
Port int `json:"port" env:"PICOCLAW_GATEWAY_PORT"`
|
||||||
HotReload bool `json:"hot_reload" env:"PICOCLAW_GATEWAY_HOT_RELOAD"`
|
HotReload bool `json:"hot_reload" env:"PICOCLAW_GATEWAY_HOT_RELOAD"`
|
||||||
LogLevel string `json:"log_level,omitempty" env:"PICOCLAW_LOG_LEVEL"`
|
LogLevel string `json:"log_level,omitempty" env:"PICOCLAW_LOG_LEVEL"`
|
||||||
|
AllowedCIDRs []string `json:"allowed_cidrs,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
func canonicalGatewayLogLevel(level logger.LogLevel) string {
|
func canonicalGatewayLogLevel(level logger.LogLevel) string {
|
||||||
|
|
|
||||||
|
|
@ -3,10 +3,12 @@ package gateway
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"fmt"
|
"fmt"
|
||||||
|
"net"
|
||||||
"os"
|
"os"
|
||||||
"os/signal"
|
"os/signal"
|
||||||
"path/filepath"
|
"path/filepath"
|
||||||
"sort"
|
"sort"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"sync"
|
"sync"
|
||||||
"sync/atomic"
|
"sync/atomic"
|
||||||
|
|
@ -64,6 +66,9 @@ type services struct {
|
||||||
DeviceService *devices.Service
|
DeviceService *devices.Service
|
||||||
HealthServer *health.Server
|
HealthServer *health.Server
|
||||||
VoiceAgentCancel context.CancelFunc
|
VoiceAgentCancel context.CancelFunc
|
||||||
|
ListenHost string
|
||||||
|
ListenAddr string
|
||||||
|
EffectiveCIDRs []string
|
||||||
manualReloadChan chan struct{}
|
manualReloadChan chan struct{}
|
||||||
reloading atomic.Bool
|
reloading atomic.Bool
|
||||||
authToken string
|
authToken string
|
||||||
|
|
@ -204,7 +209,11 @@ func Run(debug bool, homePath, configPath string, allowEmptyStartup bool) error
|
||||||
runningServices.HealthServer.SetReloadFunc(reloadTrigger)
|
runningServices.HealthServer.SetReloadFunc(reloadTrigger)
|
||||||
agentLoop.SetReloadFunc(reloadTrigger)
|
agentLoop.SetReloadFunc(reloadTrigger)
|
||||||
|
|
||||||
fmt.Printf("✓ Gateway started on %s:%d\n", cfg.Gateway.Host, cfg.Gateway.Port)
|
if runningServices.ListenHost == gatewayFallbackBindHost {
|
||||||
|
fmt.Printf("✓ Gateway started (all interfaces, port %d)\n", cfg.Gateway.Port)
|
||||||
|
} else {
|
||||||
|
fmt.Printf("✓ Gateway started on %s\n", runningServices.ListenAddr)
|
||||||
|
}
|
||||||
fmt.Println("Press Ctrl+C to stop")
|
fmt.Println("Press Ctrl+C to stop")
|
||||||
|
|
||||||
ctx, cancel := context.WithCancel(context.Background())
|
ctx, cancel := context.WithCancel(context.Background())
|
||||||
|
|
@ -265,6 +274,9 @@ func preCheckConfig(cfg *config.Config) error {
|
||||||
if cfg.Gateway.Port <= 0 || cfg.Gateway.Port > 65535 {
|
if cfg.Gateway.Port <= 0 || cfg.Gateway.Port > 65535 {
|
||||||
return fmt.Errorf("invalid gateway port: %d, port must be between 1 and 65535", cfg.Gateway.Port)
|
return fmt.Errorf("invalid gateway port: %d, port must be between 1 and 65535", cfg.Gateway.Port)
|
||||||
}
|
}
|
||||||
|
if _, err := normalizeAndValidateCIDRs(cfg.Gateway.AllowedCIDRs); err != nil {
|
||||||
|
return fmt.Errorf("invalid gateway allowed_cidrs: %w", err)
|
||||||
|
}
|
||||||
return nil
|
return nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|
@ -377,10 +389,39 @@ func setupAndStartServices(
|
||||||
fmt.Println("⚠ Warning: No channels enabled")
|
fmt.Println("⚠ Warning: No channels enabled")
|
||||||
}
|
}
|
||||||
|
|
||||||
addr := fmt.Sprintf("%s:%d", cfg.Gateway.Host, cfg.Gateway.Port)
|
listenDecision, err := resolveGatewayListenDecision(cfg.Gateway.Host, cfg.Gateway.Port, cfg.Gateway.AllowedCIDRs)
|
||||||
|
if err != nil {
|
||||||
|
if fms, ok := runningServices.MediaStore.(*media.FileMediaStore); ok {
|
||||||
|
fms.Stop()
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("error resolving gateway listen host: %w", err)
|
||||||
|
}
|
||||||
|
if listenDecision.AutoFallback {
|
||||||
|
logger.WarnCF("gateway", "Loopback bind failed, fallback to all interfaces with CIDR allowlist", map[string]any{
|
||||||
|
"configured_host": cfg.Gateway.Host,
|
||||||
|
"bind_host": listenDecision.BindHost,
|
||||||
|
"port": cfg.Gateway.Port,
|
||||||
|
"allowed_cidrs": listenDecision.AllowedCIDRs,
|
||||||
|
"reason": listenDecision.FallbackReason,
|
||||||
|
})
|
||||||
|
}
|
||||||
|
|
||||||
|
addr := net.JoinHostPort(listenDecision.BindHost, strconv.Itoa(cfg.Gateway.Port))
|
||||||
|
runningServices.ListenHost = listenDecision.BindHost
|
||||||
|
runningServices.ListenAddr = addr
|
||||||
|
runningServices.EffectiveCIDRs = listenDecision.AllowedCIDRs
|
||||||
runningServices.authToken = authToken
|
runningServices.authToken = authToken
|
||||||
runningServices.HealthServer = health.NewServer(cfg.Gateway.Host, cfg.Gateway.Port, authToken)
|
runningServices.HealthServer = health.NewServer(listenDecision.BindHost, cfg.Gateway.Port, authToken)
|
||||||
runningServices.ChannelManager.SetupHTTPServer(addr, runningServices.HealthServer)
|
if err = runningServices.ChannelManager.SetupHTTPServer(
|
||||||
|
addr,
|
||||||
|
runningServices.HealthServer,
|
||||||
|
newCIDRAllowlistMiddleware(listenDecision.AllowedCIDRs),
|
||||||
|
); err != nil {
|
||||||
|
if fms, ok := runningServices.MediaStore.(*media.FileMediaStore); ok {
|
||||||
|
fms.Stop()
|
||||||
|
}
|
||||||
|
return nil, fmt.Errorf("error setting up shared HTTP server: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
if err = runningServices.ChannelManager.StartAll(context.Background()); err != nil {
|
if err = runningServices.ChannelManager.StartAll(context.Background()); err != nil {
|
||||||
return nil, fmt.Errorf("error starting channels: %w", err)
|
return nil, fmt.Errorf("error starting channels: %w", err)
|
||||||
|
|
@ -396,11 +437,14 @@ func setupAndStartServices(
|
||||||
voiceAgent.Start(vaCtx)
|
voiceAgent.Start(vaCtx)
|
||||||
}
|
}
|
||||||
|
|
||||||
fmt.Printf(
|
if runningServices.ListenHost == gatewayFallbackBindHost {
|
||||||
"✓ Health endpoints available at http://%s:%d/health, /ready and /reload (POST)\n",
|
fmt.Printf("✓ Health endpoints available on all interfaces at port %d (/health, /ready and /reload POST)\n", cfg.Gateway.Port)
|
||||||
cfg.Gateway.Host,
|
} else {
|
||||||
cfg.Gateway.Port,
|
fmt.Printf("✓ Health endpoints available at http://%s/health, /ready and /reload (POST)\n", runningServices.ListenAddr)
|
||||||
)
|
}
|
||||||
|
if len(runningServices.EffectiveCIDRs) > 0 {
|
||||||
|
fmt.Printf("✓ Gateway CIDR allowlist enabled: %s\n", strings.Join(runningServices.EffectiveCIDRs, ", "))
|
||||||
|
}
|
||||||
|
|
||||||
stateManager := state.NewManager(cfg.WorkspacePath())
|
stateManager := state.NewManager(cfg.WorkspacePath())
|
||||||
runningServices.DeviceService = devices.NewService(devices.Config{
|
runningServices.DeviceService = devices.NewService(devices.Config{
|
||||||
|
|
|
||||||
266
pkg/gateway/network_policy.go
Normal file
266
pkg/gateway/network_policy.go
Normal file
|
|
@ -0,0 +1,266 @@
|
||||||
|
package gateway
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"fmt"
|
||||||
|
"net"
|
||||||
|
"net/http"
|
||||||
|
"sort"
|
||||||
|
"strconv"
|
||||||
|
"strings"
|
||||||
|
|
||||||
|
"github.com/sipeed/picoclaw/pkg/channels"
|
||||||
|
)
|
||||||
|
|
||||||
|
const (
|
||||||
|
gatewayDefaultLoopbackHost = "127.0.0.1"
|
||||||
|
gatewayFallbackBindHost = "0.0.0.0"
|
||||||
|
)
|
||||||
|
|
||||||
|
type gatewayListenDecision struct {
|
||||||
|
BindHost string
|
||||||
|
Port int
|
||||||
|
AllowedCIDRs []string
|
||||||
|
AutoFallback bool
|
||||||
|
FallbackReason string
|
||||||
|
}
|
||||||
|
|
||||||
|
var (
|
||||||
|
probeGatewayBind = probeTCPBind
|
||||||
|
discoverGatewayCIDRs = discoverLocalInterfaceCIDRs
|
||||||
|
)
|
||||||
|
|
||||||
|
func resolveGatewayListenDecision(configuredHost string, port int, configuredCIDRs []string) (*gatewayListenDecision, error) {
|
||||||
|
host := strings.TrimSpace(configuredHost)
|
||||||
|
if host == "" {
|
||||||
|
host = gatewayDefaultLoopbackHost
|
||||||
|
}
|
||||||
|
|
||||||
|
normalizedCIDRs, err := normalizeAndValidateCIDRs(configuredCIDRs)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid gateway allowed_cidrs: %w", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
bindErr := probeGatewayBind(host, port)
|
||||||
|
if bindErr == nil {
|
||||||
|
return &gatewayListenDecision{
|
||||||
|
BindHost: host,
|
||||||
|
Port: port,
|
||||||
|
AllowedCIDRs: normalizedCIDRs,
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
if !isLoopbackHost(host) {
|
||||||
|
return nil, fmt.Errorf("bind %s:%d failed: %w", host, port, bindErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
discoveredCIDRs, discoverErr := discoverGatewayCIDRs()
|
||||||
|
if discoverErr != nil {
|
||||||
|
return nil, fmt.Errorf(
|
||||||
|
"loopback bind %s:%d failed: %w; interface discovery failed: %v",
|
||||||
|
host,
|
||||||
|
port,
|
||||||
|
bindErr,
|
||||||
|
discoverErr,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
if len(discoveredCIDRs) == 0 {
|
||||||
|
return nil, fmt.Errorf(
|
||||||
|
"loopback bind %s:%d failed: %w; no non-loopback interface CIDRs discovered",
|
||||||
|
host,
|
||||||
|
port,
|
||||||
|
bindErr,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
fallbackCIDRs := normalizedCIDRs
|
||||||
|
if len(fallbackCIDRs) == 0 {
|
||||||
|
fallbackCIDRs = discoveredCIDRs
|
||||||
|
}
|
||||||
|
if len(fallbackCIDRs) == 0 {
|
||||||
|
return nil, fmt.Errorf("loopback bind %s:%d failed: %w; fallback allowlist is empty", host, port, bindErr)
|
||||||
|
}
|
||||||
|
|
||||||
|
if fallbackBindErr := probeGatewayBind(gatewayFallbackBindHost, port); fallbackBindErr != nil {
|
||||||
|
return nil, fmt.Errorf(
|
||||||
|
"loopback bind %s:%d failed: %w; fallback bind %s:%d failed: %v",
|
||||||
|
host,
|
||||||
|
port,
|
||||||
|
bindErr,
|
||||||
|
gatewayFallbackBindHost,
|
||||||
|
port,
|
||||||
|
fallbackBindErr,
|
||||||
|
)
|
||||||
|
}
|
||||||
|
|
||||||
|
return &gatewayListenDecision{
|
||||||
|
BindHost: gatewayFallbackBindHost,
|
||||||
|
Port: port,
|
||||||
|
AllowedCIDRs: fallbackCIDRs,
|
||||||
|
AutoFallback: true,
|
||||||
|
FallbackReason: fmt.Sprintf("loopback bind %s:%d failed, fallback to %s:%d with CIDR allowlist", host, port, gatewayFallbackBindHost, port),
|
||||||
|
}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func normalizeAndValidateCIDRs(cidrs []string) ([]string, error) {
|
||||||
|
if len(cidrs) == 0 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
seen := make(map[string]struct{}, len(cidrs))
|
||||||
|
out := make([]string, 0, len(cidrs))
|
||||||
|
for _, raw := range cidrs {
|
||||||
|
trimmed := strings.TrimSpace(raw)
|
||||||
|
if trimmed == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
_, ipNet, err := net.ParseCIDR(trimmed)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid CIDR %q: %w", trimmed, err)
|
||||||
|
}
|
||||||
|
canonical := ipNet.String()
|
||||||
|
if _, ok := seen[canonical]; ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
seen[canonical] = struct{}{}
|
||||||
|
out = append(out, canonical)
|
||||||
|
}
|
||||||
|
if len(out) == 0 {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
sort.Strings(out)
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func discoverLocalInterfaceCIDRs() ([]string, error) {
|
||||||
|
ifaces, err := net.Interfaces()
|
||||||
|
if err != nil {
|
||||||
|
return nil, err
|
||||||
|
}
|
||||||
|
|
||||||
|
out := make([]string, 0, len(ifaces))
|
||||||
|
seen := make(map[string]struct{})
|
||||||
|
ifaceErrs := make([]error, 0)
|
||||||
|
|
||||||
|
for _, iface := range ifaces {
|
||||||
|
if iface.Flags&net.FlagUp == 0 || iface.Flags&net.FlagLoopback != 0 {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
addrs, err := iface.Addrs()
|
||||||
|
if err != nil {
|
||||||
|
ifaceErrs = append(ifaceErrs, fmt.Errorf("%s: %w", iface.Name, err))
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, addr := range addrs {
|
||||||
|
ipNet := toIPNet(addr)
|
||||||
|
if ipNet == nil || ipNet.IP == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
ip := ipNet.IP
|
||||||
|
if ip.IsLoopback() || ip.IsUnspecified() || ip.IsMulticast() || ip.IsInterfaceLocalMulticast() || ip.IsLinkLocalMulticast() || ip.IsLinkLocalUnicast() {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
|
||||||
|
masked := ip.Mask(ipNet.Mask)
|
||||||
|
if masked == nil {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
canonical := (&net.IPNet{IP: masked, Mask: ipNet.Mask}).String()
|
||||||
|
if _, ok := seen[canonical]; ok {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
seen[canonical] = struct{}{}
|
||||||
|
out = append(out, canonical)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
sort.Strings(out)
|
||||||
|
if len(out) == 0 && len(ifaceErrs) > 0 {
|
||||||
|
return nil, errors.Join(ifaceErrs...)
|
||||||
|
}
|
||||||
|
return out, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func toIPNet(addr net.Addr) *net.IPNet {
|
||||||
|
switch v := addr.(type) {
|
||||||
|
case *net.IPNet:
|
||||||
|
return v
|
||||||
|
case *net.IPAddr:
|
||||||
|
if v.IP == nil {
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
bits := 128
|
||||||
|
if v.IP.To4() != nil {
|
||||||
|
bits = 32
|
||||||
|
}
|
||||||
|
return &net.IPNet{IP: v.IP, Mask: net.CIDRMask(bits, bits)}
|
||||||
|
default:
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func probeTCPBind(host string, port int) error {
|
||||||
|
addr := net.JoinHostPort(host, strconv.Itoa(port))
|
||||||
|
ln, err := net.Listen("tcp", addr)
|
||||||
|
if err != nil {
|
||||||
|
return err
|
||||||
|
}
|
||||||
|
_ = ln.Close()
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
|
||||||
|
func isLoopbackHost(host string) bool {
|
||||||
|
normalized := strings.TrimSpace(strings.ToLower(host))
|
||||||
|
if normalized == "localhost" {
|
||||||
|
return true
|
||||||
|
}
|
||||||
|
ip := net.ParseIP(normalized)
|
||||||
|
return ip != nil && ip.IsLoopback()
|
||||||
|
}
|
||||||
|
|
||||||
|
func newCIDRAllowlistMiddleware(allowedCIDRs []string) channels.HTTPMiddleware {
|
||||||
|
effectiveCIDRs := append([]string(nil), allowedCIDRs...)
|
||||||
|
return func(next http.Handler) (http.Handler, error) {
|
||||||
|
if len(effectiveCIDRs) == 0 {
|
||||||
|
return next, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
nets := make([]*net.IPNet, 0, len(effectiveCIDRs))
|
||||||
|
for _, cidr := range effectiveCIDRs {
|
||||||
|
_, ipNet, err := net.ParseCIDR(cidr)
|
||||||
|
if err != nil {
|
||||||
|
return nil, fmt.Errorf("invalid CIDR %q: %w", cidr, err)
|
||||||
|
}
|
||||||
|
nets = append(nets, ipNet)
|
||||||
|
}
|
||||||
|
|
||||||
|
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
ip := clientIPFromRemoteAddr(r.RemoteAddr)
|
||||||
|
if ip == nil {
|
||||||
|
http.Error(w, "Forbidden", http.StatusForbidden)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
if ip.IsLoopback() {
|
||||||
|
next.ServeHTTP(w, r)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
for _, ipNet := range nets {
|
||||||
|
if ipNet.Contains(ip) {
|
||||||
|
next.ServeHTTP(w, r)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
}
|
||||||
|
http.Error(w, "Forbidden", http.StatusForbidden)
|
||||||
|
}), nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func clientIPFromRemoteAddr(remoteAddr string) net.IP {
|
||||||
|
host := remoteAddr
|
||||||
|
if h, _, err := net.SplitHostPort(remoteAddr); err == nil {
|
||||||
|
host = h
|
||||||
|
}
|
||||||
|
return net.ParseIP(strings.TrimSpace(host))
|
||||||
|
}
|
||||||
200
pkg/gateway/network_policy_test.go
Normal file
200
pkg/gateway/network_policy_test.go
Normal file
|
|
@ -0,0 +1,200 @@
|
||||||
|
package gateway
|
||||||
|
|
||||||
|
import (
|
||||||
|
"errors"
|
||||||
|
"net/http"
|
||||||
|
"net/http/httptest"
|
||||||
|
"reflect"
|
||||||
|
"strings"
|
||||||
|
"testing"
|
||||||
|
)
|
||||||
|
|
||||||
|
func TestNormalizeAndValidateCIDRs(t *testing.T) {
|
||||||
|
got, err := normalizeAndValidateCIDRs([]string{" 192.168.1.20/24 ", "10.0.0.0/8", "192.168.1.0/24"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("normalizeAndValidateCIDRs() error = %v", err)
|
||||||
|
}
|
||||||
|
want := []string{"10.0.0.0/8", "192.168.1.0/24"}
|
||||||
|
if !reflect.DeepEqual(got, want) {
|
||||||
|
t.Fatalf("normalizeAndValidateCIDRs() = %v, want %v", got, want)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestNormalizeAndValidateCIDRsInvalid(t *testing.T) {
|
||||||
|
_, err := normalizeAndValidateCIDRs([]string{"bad-cidr"})
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("normalizeAndValidateCIDRs() expected error for invalid CIDR")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveGatewayListenDecisionLoopbackFallbackUsesConfiguredCIDRs(t *testing.T) {
|
||||||
|
origProbe := probeGatewayBind
|
||||||
|
origDiscover := discoverGatewayCIDRs
|
||||||
|
t.Cleanup(func() {
|
||||||
|
probeGatewayBind = origProbe
|
||||||
|
discoverGatewayCIDRs = origDiscover
|
||||||
|
})
|
||||||
|
|
||||||
|
probeGatewayBind = func(host string, _ int) error {
|
||||||
|
switch host {
|
||||||
|
case "127.0.0.1":
|
||||||
|
return errors.New("loopback unavailable")
|
||||||
|
case gatewayFallbackBindHost:
|
||||||
|
return nil
|
||||||
|
default:
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
discoverGatewayCIDRs = func() ([]string, error) {
|
||||||
|
return []string{"192.168.50.0/24"}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
decision, err := resolveGatewayListenDecision("127.0.0.1", 18790, []string{"10.0.0.0/8"})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("resolveGatewayListenDecision() error = %v", err)
|
||||||
|
}
|
||||||
|
if !decision.AutoFallback {
|
||||||
|
t.Fatal("decision.AutoFallback = false, want true")
|
||||||
|
}
|
||||||
|
if decision.BindHost != gatewayFallbackBindHost {
|
||||||
|
t.Fatalf("decision.BindHost = %q, want %q", decision.BindHost, gatewayFallbackBindHost)
|
||||||
|
}
|
||||||
|
wantCIDRs := []string{"10.0.0.0/8"}
|
||||||
|
if !reflect.DeepEqual(decision.AllowedCIDRs, wantCIDRs) {
|
||||||
|
t.Fatalf("decision.AllowedCIDRs = %v, want %v", decision.AllowedCIDRs, wantCIDRs)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveGatewayListenDecisionLoopbackFallbackUsesDiscoveredCIDRs(t *testing.T) {
|
||||||
|
origProbe := probeGatewayBind
|
||||||
|
origDiscover := discoverGatewayCIDRs
|
||||||
|
t.Cleanup(func() {
|
||||||
|
probeGatewayBind = origProbe
|
||||||
|
discoverGatewayCIDRs = origDiscover
|
||||||
|
})
|
||||||
|
|
||||||
|
probeGatewayBind = func(host string, _ int) error {
|
||||||
|
switch host {
|
||||||
|
case "localhost":
|
||||||
|
return errors.New("loopback unavailable")
|
||||||
|
case gatewayFallbackBindHost:
|
||||||
|
return nil
|
||||||
|
default:
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
discoverGatewayCIDRs = func() ([]string, error) {
|
||||||
|
return []string{"192.168.1.0/24", "10.0.0.0/8"}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
decision, err := resolveGatewayListenDecision("localhost", 18790, nil)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("resolveGatewayListenDecision() error = %v", err)
|
||||||
|
}
|
||||||
|
if !decision.AutoFallback {
|
||||||
|
t.Fatal("decision.AutoFallback = false, want true")
|
||||||
|
}
|
||||||
|
wantCIDRs := []string{"192.168.1.0/24", "10.0.0.0/8"}
|
||||||
|
if !reflect.DeepEqual(decision.AllowedCIDRs, wantCIDRs) {
|
||||||
|
t.Fatalf("decision.AllowedCIDRs = %v, want %v", decision.AllowedCIDRs, wantCIDRs)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveGatewayListenDecisionNonLoopbackFailure(t *testing.T) {
|
||||||
|
origProbe := probeGatewayBind
|
||||||
|
origDiscover := discoverGatewayCIDRs
|
||||||
|
t.Cleanup(func() {
|
||||||
|
probeGatewayBind = origProbe
|
||||||
|
discoverGatewayCIDRs = origDiscover
|
||||||
|
})
|
||||||
|
|
||||||
|
probeGatewayBind = func(host string, _ int) error {
|
||||||
|
if host == "192.0.2.1" {
|
||||||
|
return errors.New("cannot assign requested address")
|
||||||
|
}
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
discoverGatewayCIDRs = func() ([]string, error) {
|
||||||
|
return []string{"10.0.0.0/8"}, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := resolveGatewayListenDecision("192.0.2.1", 18790, nil)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("resolveGatewayListenDecision() expected error")
|
||||||
|
}
|
||||||
|
if !strings.Contains(err.Error(), "bind 192.0.2.1:18790 failed") {
|
||||||
|
t.Fatalf("error = %q, want bind failure", err.Error())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestResolveGatewayListenDecisionLoopbackFailureNoDiscoveredCIDRs(t *testing.T) {
|
||||||
|
origProbe := probeGatewayBind
|
||||||
|
origDiscover := discoverGatewayCIDRs
|
||||||
|
t.Cleanup(func() {
|
||||||
|
probeGatewayBind = origProbe
|
||||||
|
discoverGatewayCIDRs = origDiscover
|
||||||
|
})
|
||||||
|
|
||||||
|
probeGatewayBind = func(host string, _ int) error {
|
||||||
|
switch host {
|
||||||
|
case "127.0.0.1":
|
||||||
|
return errors.New("loopback unavailable")
|
||||||
|
case gatewayFallbackBindHost:
|
||||||
|
return nil
|
||||||
|
default:
|
||||||
|
return nil
|
||||||
|
}
|
||||||
|
}
|
||||||
|
discoverGatewayCIDRs = func() ([]string, error) {
|
||||||
|
return nil, nil
|
||||||
|
}
|
||||||
|
|
||||||
|
_, err := resolveGatewayListenDecision("127.0.0.1", 18790, nil)
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("resolveGatewayListenDecision() expected error")
|
||||||
|
}
|
||||||
|
if !strings.Contains(err.Error(), "no non-loopback interface CIDRs discovered") {
|
||||||
|
t.Fatalf("error = %q, want no-interface-cidr failure", err.Error())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCIDRAllowlistMiddleware(t *testing.T) {
|
||||||
|
mw := newCIDRAllowlistMiddleware([]string{"192.168.1.0/24"})
|
||||||
|
h, err := mw(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
w.WriteHeader(http.StatusOK)
|
||||||
|
}))
|
||||||
|
if err != nil {
|
||||||
|
t.Fatalf("middleware error = %v", err)
|
||||||
|
}
|
||||||
|
|
||||||
|
tests := []struct {
|
||||||
|
name string
|
||||||
|
remoteAddr string
|
||||||
|
wantStatus int
|
||||||
|
}{
|
||||||
|
{name: "inside cidr", remoteAddr: "192.168.1.99:1234", wantStatus: http.StatusOK},
|
||||||
|
{name: "loopback", remoteAddr: "127.0.0.1:1234", wantStatus: http.StatusOK},
|
||||||
|
{name: "outside cidr", remoteAddr: "10.0.0.7:1234", wantStatus: http.StatusForbidden},
|
||||||
|
{name: "malformed", remoteAddr: "not-an-ip", wantStatus: http.StatusForbidden},
|
||||||
|
}
|
||||||
|
|
||||||
|
for _, tt := range tests {
|
||||||
|
t.Run(tt.name, func(t *testing.T) {
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
req := httptest.NewRequest(http.MethodGet, "/health", nil)
|
||||||
|
req.RemoteAddr = tt.remoteAddr
|
||||||
|
h.ServeHTTP(rec, req)
|
||||||
|
if rec.Code != tt.wantStatus {
|
||||||
|
t.Fatalf("status = %d, want %d", rec.Code, tt.wantStatus)
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func TestCIDRAllowlistMiddlewareInvalidCIDR(t *testing.T) {
|
||||||
|
mw := newCIDRAllowlistMiddleware([]string{"bad-cidr"})
|
||||||
|
_, err := mw(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {}))
|
||||||
|
if err == nil {
|
||||||
|
t.Fatal("middleware expected error for invalid CIDR")
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
@ -4,6 +4,7 @@ import (
|
||||||
"encoding/json"
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"io"
|
"io"
|
||||||
|
"net"
|
||||||
"net/http"
|
"net/http"
|
||||||
"regexp"
|
"regexp"
|
||||||
"strings"
|
"strings"
|
||||||
|
|
@ -279,6 +280,15 @@ func validateConfig(cfg *config.Config) []string {
|
||||||
if cfg.Gateway.Port != 0 && (cfg.Gateway.Port < 1 || cfg.Gateway.Port > 65535) {
|
if cfg.Gateway.Port != 0 && (cfg.Gateway.Port < 1 || cfg.Gateway.Port > 65535) {
|
||||||
errs = append(errs, fmt.Sprintf("gateway.port %d is out of valid range (1-65535)", cfg.Gateway.Port))
|
errs = append(errs, fmt.Sprintf("gateway.port %d is out of valid range (1-65535)", cfg.Gateway.Port))
|
||||||
}
|
}
|
||||||
|
for index, cidr := range cfg.Gateway.AllowedCIDRs {
|
||||||
|
trimmed := strings.TrimSpace(cidr)
|
||||||
|
if trimmed == "" {
|
||||||
|
continue
|
||||||
|
}
|
||||||
|
if _, _, err := net.ParseCIDR(trimmed); err != nil {
|
||||||
|
errs = append(errs, fmt.Sprintf("gateway.allowed_cidrs[%d] is not a valid CIDR: %v", index, err))
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// Pico channel: token required when enabled
|
// Pico channel: token required when enabled
|
||||||
if cfg.Channels.Pico.Enabled && cfg.Channels.Pico.Token.String() == "" {
|
if cfg.Channels.Pico.Enabled && cfg.Channels.Pico.Token.String() == "" {
|
||||||
|
|
|
||||||
|
|
@ -173,6 +173,31 @@ func TestHandlePatchConfig_AllowsInvalidExecRegexPatternsWhenExecDisabled(t *tes
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestHandlePatchConfig_RejectsInvalidGatewayAllowedCIDRs(t *testing.T) {
|
||||||
|
configPath, cleanup := setupOAuthTestEnv(t)
|
||||||
|
defer cleanup()
|
||||||
|
|
||||||
|
h := NewHandler(configPath)
|
||||||
|
mux := http.NewServeMux()
|
||||||
|
h.RegisterRoutes(mux)
|
||||||
|
|
||||||
|
req := httptest.NewRequest(http.MethodPatch, "/api/config", bytes.NewBufferString(`{
|
||||||
|
"gateway": {
|
||||||
|
"allowed_cidrs": ["bad-cidr"]
|
||||||
|
}
|
||||||
|
}`))
|
||||||
|
req.Header.Set("Content-Type", "application/json")
|
||||||
|
|
||||||
|
rec := httptest.NewRecorder()
|
||||||
|
mux.ServeHTTP(rec, req)
|
||||||
|
if rec.Code != http.StatusBadRequest {
|
||||||
|
t.Fatalf("status = %d, want %d, body=%s", rec.Code, http.StatusBadRequest, rec.Body.String())
|
||||||
|
}
|
||||||
|
if !bytes.Contains(rec.Body.Bytes(), []byte("gateway.allowed_cidrs")) {
|
||||||
|
t.Fatalf("expected validation error mentioning gateway.allowed_cidrs, body=%s", rec.Body.String())
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
// setupPicoEnabledEnv creates a test environment with Pico channel enabled and
|
// setupPicoEnabledEnv creates a test environment with Pico channel enabled and
|
||||||
// its token stored only in .security.yml (not in the JSON payload).
|
// its token stored only in .security.yml (not in the JSON payload).
|
||||||
func setupPicoEnabledEnv(t *testing.T) (string, func()) {
|
func setupPicoEnabledEnv(t *testing.T) (string, func()) {
|
||||||
|
|
|
||||||
Loading…
Add table
Reference in a new issue