feat: validate aiproxy allowlisted domains have aibridge provider mappings at startup (#21577)

## Description

Adds startup validation to ensure all allowlisted domains have
corresponding AI Bridge provider mappings. This prevents a
misconfiguration where a domain could be MITM'd (decrypted) but have no
route to aibridge.

Previously, if a domain was in the allowlist but had no provider
mapping, requests would be decrypted and forwarded to the original
destination, a potential privacy concern. Now the server fails to start
if this misconfiguration is detected.
This commit is contained in:
Susana Ferreira
2026-01-20 17:13:12 +00:00
committed by GitHub
parent ed679bb3da
commit 09f50046cb
2 changed files with 146 additions and 87 deletions
+39 -23
View File
@@ -43,12 +43,13 @@ var loadMitmOnce sync.Once
// - decrypting requests using the configured CA certificate
// - forwarding requests to aibridged for processing
type Server struct {
ctx context.Context
logger slog.Logger
proxy *goproxy.ProxyHttpServer
httpServer *http.Server
listener net.Listener
coderAccessURL *url.URL
ctx context.Context
logger slog.Logger
proxy *goproxy.ProxyHttpServer
httpServer *http.Server
listener net.Listener
coderAccessURL *url.URL
aibridgeProviderFromHost func(host string) string
// caCert is the PEM-encoded CA certificate loaded during initialization.
// This is served to clients who need to trust the proxy.
caCert []byte
@@ -75,6 +76,9 @@ type Options struct {
// Only requests to these domains will be MITM'd and forwarded to aibridged.
// Requests to other domains will be tunneled directly without decryption.
DomainAllowlist []string
// AIBridgeProviderFromHost maps a hostname to a known aibridge provider name.
// If nil, the default provider mapping is used.
AIBridgeProviderFromHost func(host string) string
// UpstreamProxy is the URL of an upstream HTTP proxy to chain tunneled
// (non-allowlisted) requests through. If empty, tunneled requests connect
// directly to their destinations.
@@ -122,6 +126,19 @@ func New(ctx context.Context, logger slog.Logger, opts Options) (*Server, error)
return nil, xerrors.New("domain allowlist is empty, at least one domain is required")
}
// Use custom provider mapper if provided, otherwise use default.
aibridgeProviderFromHost := opts.AIBridgeProviderFromHost
if aibridgeProviderFromHost == nil {
aibridgeProviderFromHost = defaultAIBridgeProvider
}
// Validate that all allowlisted domains have correct aibridge provider mappings.
for _, domain := range opts.DomainAllowlist {
if aibridgeProviderFromHost(domain) == "" {
return nil, xerrors.Errorf("domain %q is in allowlist but has no provider mapping", domain)
}
}
logger.Info(ctx, "configured domain allowlist for MITM",
slog.F("domains", opts.DomainAllowlist),
slog.F("hosts", mitmHosts),
@@ -203,11 +220,12 @@ func New(ctx context.Context, logger slog.Logger, opts Options) (*Server, error)
}
srv := &Server{
ctx: ctx,
logger: logger,
proxy: proxy,
coderAccessURL: coderAccessURL,
caCert: certPEM,
ctx: ctx,
logger: logger,
proxy: proxy,
coderAccessURL: coderAccessURL,
aibridgeProviderFromHost: aibridgeProviderFromHost,
caCert: certPEM,
}
// Reject CONNECT requests to non-standard ports.
@@ -438,15 +456,12 @@ func extractCoderTokenFromProxyAuth(proxyAuth string) string {
return credentials[1]
}
// providerFromHost maps the request host to the aibridge provider name.
// defaultAIBridgeProvider maps the request host to the aibridge provider name.
// - Known AI providers return their provider name, used to route to the
// corresponding aibridge endpoint.
// - Unknown hosts return empty string and are passed through directly.
func providerFromURL(reqURL *url.URL) string {
if reqURL == nil {
return ""
}
switch strings.ToLower(reqURL.Hostname()) {
func defaultAIBridgeProvider(host string) string {
switch strings.ToLower(host) {
case HostAnthropic:
return aibridge.ProviderAnthropic
case HostOpenAI:
@@ -464,17 +479,18 @@ func (s *Server) handleRequest(req *http.Request, ctx *goproxy.ProxyCtx) (*http.
originalPath := req.URL.Path
// Check if this request is for a supported AI provider.
provider := providerFromURL(req.URL)
provider := s.aibridgeProviderFromHost(req.URL.Hostname())
if provider == "" {
// This can happen if a domain is in the allowlist but doesn't have a
// corresponding provider mapping in providerFromURL(). The request was
// decrypted but we don't know how to route it to aibridged.
s.logger.Warn(s.ctx, "decrypted request has no provider mapping, passing through",
// This should never happen: startup validation ensures all allowlisted
// domains have known aibridge provider mappings.
// The request is MITM'd (decrypted) but since there is no mapping,
// there is no known route to aibridge.
// Log error and forward to the original destination as a fallback.
s.logger.Error(s.ctx, "decrypted request has no provider mapping, passing through",
slog.F("host", req.Host),
slog.F("method", req.Method),
slog.F("path", originalPath),
)
// Tunnel (forward) directly to the original destination (no aibridge routing for this host).
return req, nil
}
+107 -64
View File
@@ -97,13 +97,14 @@ func generateSharedTestCA() (certFile, keyFile string, err error) {
}
type testProxyConfig struct {
listenAddr string
coderAccessURL string
allowedPorts []string
certStore *aibridgeproxyd.CertCache
domainAllowlist []string
upstreamProxy string
upstreamProxyCA string
listenAddr string
coderAccessURL string
allowedPorts []string
certStore *aibridgeproxyd.CertCache
domainAllowlist []string
aibridgeProviderFromHost func(string) string
upstreamProxy string
upstreamProxyCA string
}
type testProxyOption func(*testProxyConfig)
@@ -132,6 +133,12 @@ func withDomainAllowlist(domains ...string) testProxyOption {
}
}
func withAIBridgeProviderFromHost(fn func(string) string) testProxyOption {
return func(cfg *testProxyConfig) {
cfg.aibridgeProviderFromHost = fn
}
}
func withUpstreamProxy(upstreamProxy string) testProxyOption {
return func(cfg *testProxyConfig) {
cfg.upstreamProxy = upstreamProxy
@@ -154,6 +161,9 @@ func newTestProxy(t *testing.T, opts ...testProxyOption) *aibridgeproxyd.Server
listenAddr: "127.0.0.1:0",
coderAccessURL: "http://localhost:3000",
domainAllowlist: []string{"127.0.0.1", "localhost"},
aibridgeProviderFromHost: func(host string) string {
return "test-provider"
},
}
for _, opt := range opts {
opt(cfg)
@@ -163,14 +173,15 @@ func newTestProxy(t *testing.T, opts ...testProxyOption) *aibridgeproxyd.Server
logger := slogtest.Make(t, nil)
aibridgeOpts := aibridgeproxyd.Options{
ListenAddr: cfg.listenAddr,
CoderAccessURL: cfg.coderAccessURL,
CertFile: certFile,
KeyFile: keyFile,
AllowedPorts: cfg.allowedPorts,
DomainAllowlist: cfg.domainAllowlist,
UpstreamProxy: cfg.upstreamProxy,
UpstreamProxyCA: cfg.upstreamProxyCA,
ListenAddr: cfg.listenAddr,
CoderAccessURL: cfg.coderAccessURL,
CertFile: certFile,
KeyFile: keyFile,
AllowedPorts: cfg.allowedPorts,
DomainAllowlist: cfg.domainAllowlist,
AIBridgeProviderFromHost: cfg.aibridgeProviderFromHost,
UpstreamProxy: cfg.upstreamProxy,
UpstreamProxyCA: cfg.upstreamProxyCA,
}
if cfg.certStore != nil {
aibridgeOpts.CertStore = cfg.certStore
@@ -277,7 +288,7 @@ func TestNew(t *testing.T) {
CoderAccessURL: "http://localhost:3000",
CertFile: certFile,
KeyFile: keyFile,
DomainAllowlist: []string{"127.0.0.1", "localhost"},
DomainAllowlist: []string{aibridgeproxyd.HostAnthropic, aibridgeproxyd.HostOpenAI},
})
require.Error(t, err)
require.Contains(t, err.Error(), "listen address is required")
@@ -294,7 +305,7 @@ func TestNew(t *testing.T) {
CoderAccessURL: "http://localhost:3000",
CertFile: certFile,
KeyFile: keyFile,
DomainAllowlist: []string{"127.0.0.1", "localhost"},
DomainAllowlist: []string{aibridgeproxyd.HostAnthropic, aibridgeproxyd.HostOpenAI},
})
require.Error(t, err)
require.Contains(t, err.Error(), "listen address is required")
@@ -310,7 +321,7 @@ func TestNew(t *testing.T) {
ListenAddr: "127.0.0.1:0",
CertFile: certFile,
KeyFile: keyFile,
DomainAllowlist: []string{"127.0.0.1", "localhost"},
DomainAllowlist: []string{aibridgeproxyd.HostAnthropic, aibridgeproxyd.HostOpenAI},
})
require.Error(t, err)
require.Contains(t, err.Error(), "coder access URL is required")
@@ -327,7 +338,7 @@ func TestNew(t *testing.T) {
CoderAccessURL: " ",
CertFile: certFile,
KeyFile: keyFile,
DomainAllowlist: []string{"127.0.0.1", "localhost"},
DomainAllowlist: []string{aibridgeproxyd.HostAnthropic, aibridgeproxyd.HostOpenAI},
})
require.Error(t, err)
require.Contains(t, err.Error(), "coder access URL is required")
@@ -344,7 +355,7 @@ func TestNew(t *testing.T) {
CoderAccessURL: "://invalid",
CertFile: certFile,
KeyFile: keyFile,
DomainAllowlist: []string{"127.0.0.1", "localhost"},
DomainAllowlist: []string{aibridgeproxyd.HostAnthropic, aibridgeproxyd.HostOpenAI},
})
require.Error(t, err)
require.Contains(t, err.Error(), "invalid coder access URL")
@@ -359,7 +370,7 @@ func TestNew(t *testing.T) {
ListenAddr: ":0",
CoderAccessURL: "http://localhost:3000",
KeyFile: "key.pem",
DomainAllowlist: []string{"127.0.0.1", "localhost"},
DomainAllowlist: []string{aibridgeproxyd.HostAnthropic, aibridgeproxyd.HostOpenAI},
})
require.Error(t, err)
require.Contains(t, err.Error(), "cert file and key file are required")
@@ -374,7 +385,7 @@ func TestNew(t *testing.T) {
ListenAddr: ":0",
CoderAccessURL: "http://localhost:3000",
CertFile: "cert.pem",
DomainAllowlist: []string{"127.0.0.1", "localhost"},
DomainAllowlist: []string{aibridgeproxyd.HostAnthropic, aibridgeproxyd.HostOpenAI},
})
require.Error(t, err)
require.Contains(t, err.Error(), "cert file and key file are required")
@@ -390,7 +401,7 @@ func TestNew(t *testing.T) {
CoderAccessURL: "http://localhost:3000",
CertFile: "/nonexistent/cert.pem",
KeyFile: "/nonexistent/key.pem",
DomainAllowlist: []string{"127.0.0.1", "localhost"},
DomainAllowlist: []string{aibridgeproxyd.HostAnthropic, aibridgeproxyd.HostOpenAI},
})
require.Error(t, err)
require.Contains(t, err.Error(), "failed to load MITM certificate")
@@ -463,6 +474,23 @@ func TestNew(t *testing.T) {
require.Contains(t, err.Error(), "invalid port in domain")
})
t.Run("AllowlistWithoutProviderMapping", func(t *testing.T) {
t.Parallel()
certFile, keyFile := getSharedTestCA(t)
logger := slogtest.Make(t, nil)
_, err := aibridgeproxyd.New(t.Context(), logger, aibridgeproxyd.Options{
ListenAddr: "127.0.0.1:0",
CoderAccessURL: "http://localhost:3000",
CertFile: certFile,
KeyFile: keyFile,
DomainAllowlist: []string{"unknown.example.com"},
})
require.Error(t, err)
require.Contains(t, err.Error(), `domain "unknown.example.com" is in allowlist but has no provider mapping`)
})
t.Run("InvalidUpstreamProxy", func(t *testing.T) {
t.Parallel()
@@ -474,7 +502,7 @@ func TestNew(t *testing.T) {
CoderAccessURL: "http://localhost:3000",
CertFile: certFile,
KeyFile: keyFile,
DomainAllowlist: []string{"api.anthropic.com"},
DomainAllowlist: []string{aibridgeproxyd.HostAnthropic, aibridgeproxyd.HostOpenAI},
UpstreamProxy: "://invalid-url",
})
require.Error(t, err)
@@ -492,7 +520,7 @@ func TestNew(t *testing.T) {
CoderAccessURL: "http://localhost:3000",
CertFile: certFile,
KeyFile: keyFile,
DomainAllowlist: []string{"api.anthropic.com"},
DomainAllowlist: []string{aibridgeproxyd.HostAnthropic, aibridgeproxyd.HostOpenAI},
UpstreamProxy: "https://proxy.example.com:8080",
UpstreamProxyCA: "/nonexistent/ca.pem",
})
@@ -511,7 +539,7 @@ func TestNew(t *testing.T) {
CoderAccessURL: "http://localhost:3000",
CertFile: certFile,
KeyFile: keyFile,
DomainAllowlist: []string{"api.anthropic.com", "api.openai.com"},
DomainAllowlist: []string{aibridgeproxyd.HostAnthropic, aibridgeproxyd.HostOpenAI},
})
require.NoError(t, err)
require.NotNil(t, srv)
@@ -528,7 +556,7 @@ func TestNew(t *testing.T) {
CoderAccessURL: "http://localhost:3000",
CertFile: certFile,
KeyFile: keyFile,
DomainAllowlist: []string{"api.anthropic.com", "api.openai.com"},
DomainAllowlist: []string{aibridgeproxyd.HostAnthropic, aibridgeproxyd.HostOpenAI},
UpstreamProxy: "http://proxy.example.com:8080",
})
require.NoError(t, err)
@@ -547,7 +575,7 @@ func TestNew(t *testing.T) {
CoderAccessURL: "http://localhost:3000",
CertFile: certFile,
KeyFile: keyFile,
DomainAllowlist: []string{"api.anthropic.com"},
DomainAllowlist: []string{aibridgeproxyd.HostAnthropic, aibridgeproxyd.HostOpenAI},
UpstreamProxy: "https://proxy.example.com:8080",
UpstreamProxyCA: certFile,
})
@@ -567,7 +595,7 @@ func TestClose(t *testing.T) {
CoderAccessURL: "http://localhost:3000",
CertFile: certFile,
KeyFile: keyFile,
DomainAllowlist: []string{"127.0.0.1", "localhost"},
DomainAllowlist: []string{aibridgeproxyd.HostAnthropic, aibridgeproxyd.HostOpenAI},
})
require.NoError(t, err)
@@ -608,6 +636,12 @@ func TestProxy_CertCaching(t *testing.T) {
w.WriteHeader(http.StatusOK)
})
// Create a mock aibridged server for allowlisted (MITM'd) requests.
aibridgedServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
}))
t.Cleanup(func() { aibridgedServer.Close() })
// Create a cert cache so we can inspect it after the request.
certCache := aibridgeproxyd.NewCertCache()
@@ -619,6 +653,7 @@ func TestProxy_CertCaching(t *testing.T) {
// Start the proxy server with the certificate cache.
srv := newTestProxy(t,
withCoderAccessURL(aibridgedServer.URL),
withAllowedPorts(targetURL.Port()),
withCertStore(certCache),
withDomainAllowlist(domainAllowlist...),
@@ -699,8 +734,16 @@ func TestProxy_PortValidation(t *testing.T) {
_, _ = w.Write([]byte("hello from target"))
})
// Start the proxy server on a random port to avoid conflicts when running tests in parallel.
// Create a mock aibridged server for allowlisted (MITM'd) requests.
aibridgedServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte("hello from aibridged"))
}))
t.Cleanup(func() { aibridgedServer.Close() })
// Start the proxy server.
srv := newTestProxy(t,
withCoderAccessURL(aibridgedServer.URL),
withAllowedPorts(tt.allowedPorts(targetURL)...),
withDomainAllowlist(targetURL.Hostname()),
)
@@ -715,15 +758,14 @@ func TestProxy_PortValidation(t *testing.T) {
require.Error(t, err)
return
}
require.NoError(t, err)
defer resp.Body.Close()
// Verify the request was successful and reached the target server.
// Verify the request was successful and routed to aibridged.
body, err := io.ReadAll(resp.Body)
require.NoError(t, err)
require.Equal(t, http.StatusOK, resp.StatusCode)
require.Equal(t, "hello from target", string(body))
require.Equal(t, "hello from aibridged", string(body))
})
}
}
@@ -768,8 +810,16 @@ func TestProxy_Authentication(t *testing.T) {
_, _ = w.Write([]byte("hello from target"))
})
// Start the proxy server on a random port to avoid conflicts when running tests in parallel.
// Create a mock aibridged server for allowlisted (MITM'd) requests.
aibridgedServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
w.WriteHeader(http.StatusOK)
_, _ = w.Write([]byte("hello from aibridged"))
}))
t.Cleanup(func() { aibridgedServer.Close() })
// Start the proxy server.
srv := newTestProxy(t,
withCoderAccessURL(aibridgedServer.URL),
withAllowedPorts(targetURL.Port()),
withDomainAllowlist(targetURL.Hostname()),
)
@@ -786,11 +836,11 @@ func TestProxy_Authentication(t *testing.T) {
require.NoError(t, err)
defer resp.Body.Close()
// Verify the response was successfully proxied.
// Verify the response was successfully routed to aibridged.
body, err := io.ReadAll(resp.Body)
require.NoError(t, err)
require.Equal(t, http.StatusOK, resp.StatusCode)
require.Equal(t, "hello from target", string(body))
require.Equal(t, "hello from aibridged", string(body))
}
})
}
@@ -800,17 +850,16 @@ func TestProxy_MITM(t *testing.T) {
t.Parallel()
tests := []struct {
name string
domainAllowlist []string
allowedPorts []string
buildTargetURL func(tunneledURL *url.URL) (string, error)
tunneled bool
noAIBridgeRouting bool
expectedPath string
name string
domainAllowlist []string
allowedPorts []string
buildTargetURL func(tunneledURL *url.URL) (string, error)
tunneled bool
expectedPath string
}{
{
name: "MitmdAnthropic",
domainAllowlist: []string{"api.anthropic.com"},
domainAllowlist: []string{aibridgeproxyd.HostAnthropic},
allowedPorts: []string{"443"},
buildTargetURL: func(_ *url.URL) (string, error) {
return "https://api.anthropic.com/v1/messages", nil
@@ -819,7 +868,7 @@ func TestProxy_MITM(t *testing.T) {
},
{
name: "MitmdAnthropicNonDefaultPort",
domainAllowlist: []string{"api.anthropic.com"},
domainAllowlist: []string{aibridgeproxyd.HostAnthropic},
allowedPorts: []string{"8443"},
buildTargetURL: func(_ *url.URL) (string, error) {
return "https://api.anthropic.com:8443/v1/messages", nil
@@ -828,7 +877,7 @@ func TestProxy_MITM(t *testing.T) {
},
{
name: "MitmdOpenAI",
domainAllowlist: []string{"api.openai.com"},
domainAllowlist: []string{aibridgeproxyd.HostOpenAI},
allowedPorts: []string{"443"},
buildTargetURL: func(_ *url.URL) (string, error) {
return "https://api.openai.com/v1/chat/completions", nil
@@ -837,7 +886,7 @@ func TestProxy_MITM(t *testing.T) {
},
{
name: "MitmdOpenAINonDefaultPort",
domainAllowlist: []string{"api.openai.com"},
domainAllowlist: []string{aibridgeproxyd.HostOpenAI},
allowedPorts: []string{"8443"},
buildTargetURL: func(_ *url.URL) (string, error) {
return "https://api.openai.com:8443/v1/chat/completions", nil
@@ -846,26 +895,13 @@ func TestProxy_MITM(t *testing.T) {
},
{
name: "TunneledUnknownHost",
domainAllowlist: []string{"other.example.com"},
domainAllowlist: []string{aibridgeproxyd.HostAnthropic, aibridgeproxyd.HostOpenAI},
allowedPorts: nil, // will use tunneledURL.Port()
buildTargetURL: func(tunneledURL *url.URL) (string, error) {
return url.JoinPath(tunneledURL.String(), "/some/path")
},
tunneled: true,
},
// The host is MITM'd but has no provider mapping.
// The request is decrypted but forwarded to the original destination
// instead of being routed to aibridge.
{
name: "MitmdWithoutAIBridgeRouting",
domainAllowlist: nil, // will use tunneledURL.Hostname()
allowedPorts: nil, // will use tunneledURL.Port()
buildTargetURL: func(tunneledURL *url.URL) (string, error) {
return url.JoinPath(tunneledURL.String(), "/some/path")
},
tunneled: false,
noAIBridgeRouting: true,
},
}
for _, tt := range tests {
@@ -908,6 +944,8 @@ func TestProxy_MITM(t *testing.T) {
withCoderAccessURL(aibridgedServer.URL),
withAllowedPorts(allowedPorts...),
withDomainAllowlist(domainAllowlist...),
// Use default provider mapping to test real AI provider routing.
withAIBridgeProviderFromHost(nil),
)
// Build the target URL:
@@ -941,7 +979,7 @@ func TestProxy_MITM(t *testing.T) {
require.NoError(t, err)
require.Equal(t, http.StatusOK, resp.StatusCode)
if tt.tunneled || tt.noAIBridgeRouting {
if tt.tunneled {
// Verify request went to target server, not aibridged.
require.Equal(t, "hello from tunneled", string(body))
require.Empty(t, receivedPath, "aibridged should not receive tunneled requests")
@@ -1030,6 +1068,9 @@ func TestServeCACert_CompoundPEM(t *testing.T) {
CertFile: compoundCertFile,
KeyFile: keyFile,
DomainAllowlist: []string{"127.0.0.1", "localhost"},
AIBridgeProviderFromHost: func(host string) string {
return "test-provider"
},
})
require.NoError(t, err)
t.Cleanup(func() { _ = srv.Close() })
@@ -1248,14 +1289,16 @@ func TestUpstreamProxy(t *testing.T) {
// Configure allowlist based on test case:
// - For tunneled requests, api.anthropic.com is in allowlist, but we target a different host.
// - For MITM, api.anthropic.com must be in the allowlist.
domainAllowlist := []string{"api.anthropic.com"}
domainAllowlist := []string{aibridgeproxyd.HostAnthropic}
// Create aiproxy with upstream proxy configured.
proxyOpts := []testProxyOption{
withCoderAccessURL(aibridgeServer.URL),
withDomainAllowlist(domainAllowlist...),
withUpstreamProxy(upstreamProxy.URL),
withAllowedPorts("80", "443", finalDestinationURL.Port(), parsedTargetURL.Port()),
withAllowedPorts("80", "443", parsedTargetURL.Port()),
// Use default provider mapping to test real AI provider routing.
withAIBridgeProviderFromHost(nil),
}
if upstreamProxyCAFile != "" {
proxyOpts = append(proxyOpts, withUpstreamProxyCA(upstreamProxyCAFile))