mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
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:
@@ -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
|
||||
}
|
||||
|
||||
|
||||
@@ -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))
|
||||
|
||||
Reference in New Issue
Block a user