feat(gateway): add bind fallback and CIDR policy enforcement

This commit is contained in:
Sakurapainting 2026-04-02 18:02:31 +08:00
parent b114dcaeb1
commit 0a775d9fb9
9 changed files with 587 additions and 20 deletions

View file

@ -465,9 +465,11 @@
},
"gateway": {
"_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",
"port": 18790,
"hot_reload": false,
"log_level": "fatal"
"log_level": "fatal",
"allowed_cidrs": []
}
}

View file

@ -98,6 +98,10 @@ type asyncTask struct {
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.
// Implements PlaceholderRecorder.
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.
// It registers health endpoints from the health server and discovers channels
// 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()
// Register health endpoints
@ -447,12 +451,26 @@ func (m *Manager) SetupHTTPServer(addr string, healthServer *health.Server) {
// Discover and register webhook handlers and health checkers
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{
Addr: addr,
Handler: m.mux,
Handler: handler,
ReadTimeout: 30 * time.Second,
WriteTimeout: 30 * time.Second,
}
return nil
}
// registerHTTPHandlersLocked registers webhook and health-check handlers for

View file

@ -362,6 +362,7 @@ func DefaultConfig() *Config {
Port: 18790,
HotReload: false,
LogLevel: DefaultGatewayLogLevel,
AllowedCIDRs: nil,
},
Tools: ToolsConfig{
FilterSensitiveData: true,

View file

@ -14,6 +14,7 @@ type GatewayConfig struct {
Port int `json:"port" env:"PICOCLAW_GATEWAY_PORT"`
HotReload bool `json:"hot_reload" env:"PICOCLAW_GATEWAY_HOT_RELOAD"`
LogLevel string `json:"log_level,omitempty" env:"PICOCLAW_LOG_LEVEL"`
AllowedCIDRs []string `json:"allowed_cidrs,omitempty"`
}
func canonicalGatewayLogLevel(level logger.LogLevel) string {

View file

@ -3,10 +3,12 @@ package gateway
import (
"context"
"fmt"
"net"
"os"
"os/signal"
"path/filepath"
"sort"
"strconv"
"strings"
"sync"
"sync/atomic"
@ -64,6 +66,9 @@ type services struct {
DeviceService *devices.Service
HealthServer *health.Server
VoiceAgentCancel context.CancelFunc
ListenHost string
ListenAddr string
EffectiveCIDRs []string
manualReloadChan chan struct{}
reloading atomic.Bool
authToken string
@ -204,7 +209,11 @@ func Run(debug bool, homePath, configPath string, allowEmptyStartup bool) error
runningServices.HealthServer.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")
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 {
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
}
@ -377,10 +389,39 @@ func setupAndStartServices(
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.HealthServer = health.NewServer(cfg.Gateway.Host, cfg.Gateway.Port, authToken)
runningServices.ChannelManager.SetupHTTPServer(addr, runningServices.HealthServer)
runningServices.HealthServer = health.NewServer(listenDecision.BindHost, cfg.Gateway.Port, authToken)
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 {
return nil, fmt.Errorf("error starting channels: %w", err)
@ -396,11 +437,14 @@ func setupAndStartServices(
voiceAgent.Start(vaCtx)
}
fmt.Printf(
"✓ Health endpoints available at http://%s:%d/health, /ready and /reload (POST)\n",
cfg.Gateway.Host,
cfg.Gateway.Port,
)
if runningServices.ListenHost == gatewayFallbackBindHost {
fmt.Printf("✓ Health endpoints available on all interfaces at port %d (/health, /ready and /reload POST)\n", cfg.Gateway.Port)
} else {
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())
runningServices.DeviceService = devices.NewService(devices.Config{

View 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))
}

View 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")
}
}

View file

@ -4,6 +4,7 @@ import (
"encoding/json"
"fmt"
"io"
"net"
"net/http"
"regexp"
"strings"
@ -279,6 +280,15 @@ func validateConfig(cfg *config.Config) []string {
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))
}
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
if cfg.Channels.Pico.Enabled && cfg.Channels.Pico.Token.String() == "" {

View file

@ -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
// its token stored only in .security.yml (not in the JSON payload).
func setupPicoEnabledEnv(t *testing.T) (string, func()) {