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
|
// - 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
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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))
|
||||||
|
|||||||
Reference in New Issue
Block a user