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