mirror of
https://github.com/coder/coder.git
synced 2026-09-24 15:04:27 +08:00
refactor: add aiproxy test helpers to reduce boilerplate (#21404)
## Description Adds test helper functions to reduce boilerplate in `aibridgeproxyd` tests: * `newTestProxy`: creates a proxy server with functional options, waits for it to be ready * `newProxyClient`: creates an HTTP client configured to use the proxy * `newTargetServer`: creates a mock HTTPS server and returns its URL Related to: https://github.com/coder/coder/pull/21344#discussion_r2638930199
This commit is contained in:
@@ -96,6 +96,131 @@ func generateSharedTestCA() (certFile, keyFile string, err error) {
|
|||||||
return certPath, keyPath, nil
|
return certPath, keyPath, nil
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type testProxyConfig struct {
|
||||||
|
listenAddr string
|
||||||
|
coderAccessURL string
|
||||||
|
allowedPorts []string
|
||||||
|
certStore *aibridgeproxyd.CertCache
|
||||||
|
}
|
||||||
|
|
||||||
|
type testProxyOption func(*testProxyConfig)
|
||||||
|
|
||||||
|
func withAllowedPorts(ports ...string) testProxyOption {
|
||||||
|
return func(cfg *testProxyConfig) {
|
||||||
|
cfg.allowedPorts = ports
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func withCoderAccessURL(coderAccessURL string) testProxyOption {
|
||||||
|
return func(cfg *testProxyConfig) {
|
||||||
|
cfg.coderAccessURL = coderAccessURL
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
func withCertStore(store *aibridgeproxyd.CertCache) testProxyOption {
|
||||||
|
return func(cfg *testProxyConfig) {
|
||||||
|
cfg.certStore = store
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
// newTestProxy creates a new AI Bridge Proxy server for testing.
|
||||||
|
// It uses the shared test CA and registers cleanup automatically.
|
||||||
|
// It waits for the proxy server to be ready before returning.
|
||||||
|
func newTestProxy(t *testing.T, opts ...testProxyOption) *aibridgeproxyd.Server {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
cfg := &testProxyConfig{
|
||||||
|
listenAddr: "127.0.0.1:0",
|
||||||
|
coderAccessURL: "http://localhost:3000",
|
||||||
|
}
|
||||||
|
for _, opt := range opts {
|
||||||
|
opt(cfg)
|
||||||
|
}
|
||||||
|
|
||||||
|
certFile, keyFile := getSharedTestCA(t)
|
||||||
|
logger := slogtest.Make(t, nil)
|
||||||
|
|
||||||
|
aibridgeOpts := aibridgeproxyd.Options{
|
||||||
|
ListenAddr: cfg.listenAddr,
|
||||||
|
CoderAccessURL: cfg.coderAccessURL,
|
||||||
|
CertFile: certFile,
|
||||||
|
KeyFile: keyFile,
|
||||||
|
AllowedPorts: cfg.allowedPorts,
|
||||||
|
}
|
||||||
|
if cfg.certStore != nil {
|
||||||
|
aibridgeOpts.CertStore = cfg.certStore
|
||||||
|
}
|
||||||
|
|
||||||
|
srv, err := aibridgeproxyd.New(t.Context(), logger, aibridgeOpts)
|
||||||
|
require.NoError(t, err)
|
||||||
|
t.Cleanup(func() { _ = srv.Close() })
|
||||||
|
|
||||||
|
// Wait for the proxy server to be ready.
|
||||||
|
proxyAddr := srv.Addr()
|
||||||
|
require.NotEmpty(t, proxyAddr)
|
||||||
|
require.Eventually(t, func() bool {
|
||||||
|
conn, err := net.Dial("tcp", proxyAddr)
|
||||||
|
if err != nil {
|
||||||
|
return false
|
||||||
|
}
|
||||||
|
_ = conn.Close()
|
||||||
|
return true
|
||||||
|
}, testutil.WaitShort, testutil.IntervalFast)
|
||||||
|
|
||||||
|
return srv
|
||||||
|
}
|
||||||
|
|
||||||
|
// newProxyClient creates an HTTP client configured to use the proxy and trust its CA.
|
||||||
|
// It adds a Proxy-Authorization header with the provided token for authentication.
|
||||||
|
func newProxyClient(t *testing.T, srv *aibridgeproxyd.Server, proxyAuth string) *http.Client {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
certFile, _ := getSharedTestCA(t)
|
||||||
|
|
||||||
|
// Load the CA certificate so the client trusts the proxy's MITM certificate.
|
||||||
|
certPEM, err := os.ReadFile(certFile)
|
||||||
|
require.NoError(t, err)
|
||||||
|
certPool := x509.NewCertPool()
|
||||||
|
ok := certPool.AppendCertsFromPEM(certPEM)
|
||||||
|
require.True(t, ok)
|
||||||
|
|
||||||
|
// Create an HTTP client configured to use the proxy.
|
||||||
|
proxyURL, err := url.Parse("http://" + srv.Addr())
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
transport := &http.Transport{
|
||||||
|
Proxy: http.ProxyURL(proxyURL),
|
||||||
|
TLSClientConfig: &tls.Config{
|
||||||
|
MinVersion: tls.VersionTLS12,
|
||||||
|
RootCAs: certPool,
|
||||||
|
},
|
||||||
|
}
|
||||||
|
|
||||||
|
// Only set the header if proxyAuth is provided. This allows tests to
|
||||||
|
// verify behavior when the Proxy-Authorization header is missing.
|
||||||
|
if proxyAuth != "" {
|
||||||
|
transport.ProxyConnectHeader = http.Header{
|
||||||
|
"Proxy-Authorization": []string{proxyAuth},
|
||||||
|
}
|
||||||
|
}
|
||||||
|
|
||||||
|
return &http.Client{Transport: transport}
|
||||||
|
}
|
||||||
|
|
||||||
|
// newTargetServer creates a mock HTTPS server that will be the target of proxied requests.
|
||||||
|
// It returns the server's parsed URL. The server is automatically closed when the test ends.
|
||||||
|
func newTargetServer(t *testing.T, handler http.HandlerFunc) *url.URL {
|
||||||
|
t.Helper()
|
||||||
|
|
||||||
|
srv := httptest.NewTLSServer(handler)
|
||||||
|
t.Cleanup(srv.Close)
|
||||||
|
|
||||||
|
srvURL, err := url.Parse(srv.URL)
|
||||||
|
require.NoError(t, err)
|
||||||
|
|
||||||
|
return srvURL
|
||||||
|
}
|
||||||
|
|
||||||
// makeProxyAuthHeader creates a Proxy-Authorization header value with the given token.
|
// makeProxyAuthHeader creates a Proxy-Authorization header value with the given token.
|
||||||
// Format: "Basic base64(username:token)" where username is "ignored".
|
// Format: "Basic base64(username:token)" where username is "ignored".
|
||||||
func makeProxyAuthHeader(token string) string {
|
func makeProxyAuthHeader(token string) string {
|
||||||
@@ -257,72 +382,24 @@ func TestClose(t *testing.T) {
|
|||||||
func TestProxy_CertCaching(t *testing.T) {
|
func TestProxy_CertCaching(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
// Create a mock HTTPS server that will be the target of our proxied request.
|
// Create a mock HTTPS server that will be the target of the proxied request.
|
||||||
targetServer := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
targetURL := newTargetServer(t, func(w http.ResponseWriter, r *http.Request) {
|
||||||
w.WriteHeader(http.StatusOK)
|
w.WriteHeader(http.StatusOK)
|
||||||
}))
|
})
|
||||||
defer targetServer.Close()
|
|
||||||
|
|
||||||
targetURL, err := url.Parse(targetServer.URL)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
certFile, keyFile := getSharedTestCA(t)
|
|
||||||
logger := slogtest.Make(t, nil)
|
|
||||||
|
|
||||||
// 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()
|
||||||
|
|
||||||
// Start the proxy server with the certificate cache.
|
// Start the proxy server with the certificate cache.
|
||||||
srv, err := aibridgeproxyd.New(t.Context(), logger, aibridgeproxyd.Options{
|
srv := newTestProxy(t,
|
||||||
ListenAddr: "127.0.0.1:0",
|
withAllowedPorts(targetURL.Port()),
|
||||||
CoderAccessURL: "http://localhost:3000",
|
withCertStore(certCache),
|
||||||
CertFile: certFile,
|
)
|
||||||
KeyFile: keyFile,
|
|
||||||
CertStore: certCache,
|
|
||||||
AllowedPorts: []string{targetURL.Port()},
|
|
||||||
})
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = srv.Close() })
|
|
||||||
|
|
||||||
proxyAddr := srv.Addr()
|
|
||||||
require.NotEmpty(t, proxyAddr)
|
|
||||||
|
|
||||||
// Wait for the proxy server to be ready.
|
|
||||||
require.Eventually(t, func() bool {
|
|
||||||
conn, err := net.Dial("tcp", proxyAddr)
|
|
||||||
if err != nil {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
_ = conn.Close()
|
|
||||||
return true
|
|
||||||
}, testutil.WaitShort, testutil.IntervalFast)
|
|
||||||
|
|
||||||
// Load the CA certificate so the client trusts the proxy's MITM certificate.
|
|
||||||
certPEM, err := os.ReadFile(certFile)
|
|
||||||
require.NoError(t, err)
|
|
||||||
certPool := x509.NewCertPool()
|
|
||||||
certPool.AppendCertsFromPEM(certPEM)
|
|
||||||
|
|
||||||
// Create an HTTP client configured to use the proxy.
|
|
||||||
proxyURL, err := url.Parse("http://" + proxyAddr)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
client := &http.Client{
|
|
||||||
Transport: &http.Transport{
|
|
||||||
Proxy: http.ProxyURL(proxyURL),
|
|
||||||
ProxyConnectHeader: http.Header{
|
|
||||||
"Proxy-Authorization": []string{makeProxyAuthHeader("test-session-token")},
|
|
||||||
},
|
|
||||||
TLSClientConfig: &tls.Config{
|
|
||||||
MinVersion: tls.VersionTLS12,
|
|
||||||
RootCAs: certPool,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
// Make a request through the proxy to the target server.
|
// Make a request through the proxy to the target server.
|
||||||
// This triggers MITM and caches the generated certificate.
|
// This triggers MITM and caches the generated certificate.
|
||||||
req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, targetServer.URL, nil)
|
client := newProxyClient(t, srv, makeProxyAuthHeader("test-session-token"))
|
||||||
|
req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, targetURL.String(), nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
resp, err := client.Do(req)
|
resp, err := client.Do(req)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
@@ -343,17 +420,23 @@ func TestProxy_PortValidation(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
tests := []struct {
|
tests := []struct {
|
||||||
name string
|
name string
|
||||||
allowPort bool
|
allowedPorts func(targetURL *url.URL) []string
|
||||||
expectError bool
|
expectError bool
|
||||||
}{
|
}{
|
||||||
{
|
{
|
||||||
name: "AllowedPort",
|
name: "AllowedPort",
|
||||||
allowPort: true,
|
// Include the target's random port so the request is allowed.
|
||||||
|
allowedPorts: func(targetURL *url.URL) []string {
|
||||||
|
return []string{targetURL.Port()}
|
||||||
|
},
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
name: "RejectedPort",
|
name: "RejectedPort",
|
||||||
allowPort: false,
|
// Only allow port 443 which doesn't match the target.
|
||||||
|
allowedPorts: func(_ *url.URL) []string {
|
||||||
|
return []string{"443"}
|
||||||
|
},
|
||||||
expectError: true,
|
expectError: true,
|
||||||
},
|
},
|
||||||
}
|
}
|
||||||
@@ -363,77 +446,17 @@ func TestProxy_PortValidation(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
// Create a target HTTPS server that will be the destination of our proxied request.
|
// Create a target HTTPS server that will be the destination of our proxied request.
|
||||||
targetServer := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
targetURL := newTargetServer(t, func(w http.ResponseWriter, r *http.Request) {
|
||||||
w.WriteHeader(http.StatusOK)
|
w.WriteHeader(http.StatusOK)
|
||||||
_, _ = w.Write([]byte("hello from target"))
|
_, _ = w.Write([]byte("hello from target"))
|
||||||
}))
|
})
|
||||||
t.Cleanup(func() { targetServer.Close() })
|
|
||||||
|
|
||||||
targetURL, err := url.Parse(targetServer.URL)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
certFile, keyFile := getSharedTestCA(t)
|
|
||||||
logger := slogtest.Make(t, nil)
|
|
||||||
|
|
||||||
// Configure allowed ports based on test case.
|
|
||||||
// For allowed case, include the target's random port.
|
|
||||||
// For rejected case, only allow port 443 which doesn't match the target.
|
|
||||||
var allowedPorts []string
|
|
||||||
if tt.allowPort {
|
|
||||||
allowedPorts = []string{targetURL.Port()}
|
|
||||||
} else {
|
|
||||||
allowedPorts = []string{"443"}
|
|
||||||
}
|
|
||||||
|
|
||||||
// Start the proxy server on a random port to avoid conflicts when running tests in parallel.
|
// Start the proxy server on a random port to avoid conflicts when running tests in parallel.
|
||||||
srv, err := aibridgeproxyd.New(t.Context(), logger, aibridgeproxyd.Options{
|
srv := newTestProxy(t, withAllowedPorts(tt.allowedPorts(targetURL)...))
|
||||||
ListenAddr: "127.0.0.1:0",
|
|
||||||
CoderAccessURL: "http://localhost:3000",
|
|
||||||
CertFile: certFile,
|
|
||||||
KeyFile: keyFile,
|
|
||||||
AllowedPorts: allowedPorts,
|
|
||||||
})
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = srv.Close() })
|
|
||||||
|
|
||||||
proxyAddr := srv.Addr()
|
|
||||||
require.NotEmpty(t, proxyAddr)
|
|
||||||
|
|
||||||
// Wait for the proxy server to be ready.
|
|
||||||
require.Eventually(t, func() bool {
|
|
||||||
conn, err := net.Dial("tcp", proxyAddr)
|
|
||||||
if err != nil {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
_ = conn.Close()
|
|
||||||
return true
|
|
||||||
}, testutil.WaitShort, testutil.IntervalFast)
|
|
||||||
|
|
||||||
// Load the CA certificate so the client trusts the proxy's MITM certificate.
|
|
||||||
certPEM, err := os.ReadFile(certFile)
|
|
||||||
require.NoError(t, err)
|
|
||||||
certPool := x509.NewCertPool()
|
|
||||||
certPool.AppendCertsFromPEM(certPEM)
|
|
||||||
|
|
||||||
// Create an HTTP client configured to use the proxy.
|
|
||||||
proxyURL, err := url.Parse("http://" + proxyAddr)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
client := &http.Client{
|
|
||||||
Transport: &http.Transport{
|
|
||||||
Proxy: http.ProxyURL(proxyURL),
|
|
||||||
ProxyConnectHeader: http.Header{
|
|
||||||
"Proxy-Authorization": []string{makeProxyAuthHeader("test-session-token")},
|
|
||||||
},
|
|
||||||
TLSClientConfig: &tls.Config{
|
|
||||||
MinVersion: tls.VersionTLS12,
|
|
||||||
RootCAs: certPool,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
// Make a request through the proxy to the target server.
|
// Make a request through the proxy to the target server.
|
||||||
req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, targetServer.URL, nil)
|
client := newProxyClient(t, srv, makeProxyAuthHeader("test-session-token"))
|
||||||
|
req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, targetURL.String(), nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
|
|
||||||
resp, err := client.Do(req)
|
resp, err := client.Do(req)
|
||||||
@@ -445,6 +468,7 @@ func TestProxy_PortValidation(t *testing.T) {
|
|||||||
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.
|
||||||
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)
|
||||||
@@ -488,71 +512,17 @@ func TestProxy_Authentication(t *testing.T) {
|
|||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
// Create a mock HTTPS server that will be the target of our proxied request.
|
// Create a mock HTTPS server that will be the target of our proxied request.
|
||||||
targetServer := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
targetURL := newTargetServer(t, func(w http.ResponseWriter, _ *http.Request) {
|
||||||
w.WriteHeader(http.StatusOK)
|
w.WriteHeader(http.StatusOK)
|
||||||
_, _ = w.Write([]byte("hello from target"))
|
_, _ = w.Write([]byte("hello from target"))
|
||||||
}))
|
})
|
||||||
t.Cleanup(func() { targetServer.Close() })
|
|
||||||
|
|
||||||
targetURL, err := url.Parse(targetServer.URL)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
certFile, keyFile := getSharedTestCA(t)
|
|
||||||
logger := slogtest.Make(t, nil)
|
|
||||||
|
|
||||||
// Start the proxy server on a random port to avoid conflicts when running tests in parallel.
|
// Start the proxy server on a random port to avoid conflicts when running tests in parallel.
|
||||||
// The actual port is accessible via srv.Addr().
|
srv := newTestProxy(t, withAllowedPorts(targetURL.Port()))
|
||||||
srv, err := aibridgeproxyd.New(t.Context(), logger, aibridgeproxyd.Options{
|
|
||||||
ListenAddr: "127.0.0.1:0",
|
|
||||||
CoderAccessURL: "http://localhost:3000",
|
|
||||||
CertFile: certFile,
|
|
||||||
KeyFile: keyFile,
|
|
||||||
AllowedPorts: []string{targetURL.Port()},
|
|
||||||
})
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = srv.Close() })
|
|
||||||
|
|
||||||
proxyAddr := srv.Addr()
|
|
||||||
require.NotEmpty(t, proxyAddr)
|
|
||||||
|
|
||||||
// Wait for the proxy server to be ready.
|
|
||||||
require.Eventually(t, func() bool {
|
|
||||||
conn, err := net.Dial("tcp", proxyAddr)
|
|
||||||
if err != nil {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
_ = conn.Close()
|
|
||||||
return true
|
|
||||||
}, testutil.WaitShort, testutil.IntervalFast)
|
|
||||||
|
|
||||||
// Load the CA certificate so the client trusts the proxy's MITM certificate.
|
|
||||||
certPEM, err := os.ReadFile(certFile)
|
|
||||||
require.NoError(t, err)
|
|
||||||
certPool := x509.NewCertPool()
|
|
||||||
certPool.AppendCertsFromPEM(certPEM)
|
|
||||||
|
|
||||||
// Create an HTTP client configured to use the proxy.
|
|
||||||
proxyURL, err := url.Parse("http://" + proxyAddr)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
transport := &http.Transport{
|
|
||||||
Proxy: http.ProxyURL(proxyURL),
|
|
||||||
TLSClientConfig: &tls.Config{
|
|
||||||
MinVersion: tls.VersionTLS12,
|
|
||||||
RootCAs: certPool,
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
if tt.proxyAuth != "" {
|
|
||||||
transport.ProxyConnectHeader = http.Header{
|
|
||||||
"Proxy-Authorization": []string{tt.proxyAuth},
|
|
||||||
}
|
|
||||||
}
|
|
||||||
|
|
||||||
client := &http.Client{Transport: transport}
|
|
||||||
|
|
||||||
// Make a request through the proxy to the target server.
|
// Make a request through the proxy to the target server.
|
||||||
req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, targetServer.URL, nil)
|
client := newProxyClient(t, srv, tt.proxyAuth)
|
||||||
|
req, err := http.NewRequestWithContext(t.Context(), http.MethodGet, targetURL.String(), nil)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
resp, err := client.Do(req)
|
resp, err := client.Do(req)
|
||||||
|
|
||||||
@@ -624,7 +594,7 @@ func TestProxy_MITM(t *testing.T) {
|
|||||||
var receivedPath string
|
var receivedPath string
|
||||||
var receivedAuth string
|
var receivedAuth string
|
||||||
|
|
||||||
// Create a mock aibridged server.
|
// Create a mock aibridged server that captures requests.
|
||||||
aibridgedServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
aibridgedServer := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
receivedPath = r.URL.Path
|
receivedPath = r.URL.Path
|
||||||
receivedAuth = r.Header.Get("Authorization")
|
receivedAuth = r.Header.Get("Authorization")
|
||||||
@@ -634,14 +604,10 @@ func TestProxy_MITM(t *testing.T) {
|
|||||||
t.Cleanup(func() { aibridgedServer.Close() })
|
t.Cleanup(func() { aibridgedServer.Close() })
|
||||||
|
|
||||||
// Create a mock target server for passthrough tests.
|
// Create a mock target server for passthrough tests.
|
||||||
targetServer := httptest.NewTLSServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
passthroughURL := newTargetServer(t, func(w http.ResponseWriter, _ *http.Request) {
|
||||||
w.WriteHeader(http.StatusOK)
|
w.WriteHeader(http.StatusOK)
|
||||||
_, _ = w.Write([]byte("hello from passthrough"))
|
_, _ = w.Write([]byte("hello from passthrough"))
|
||||||
}))
|
})
|
||||||
t.Cleanup(func() { targetServer.Close() })
|
|
||||||
|
|
||||||
certFile, keyFile := getSharedTestCA(t)
|
|
||||||
logger := slogtest.Make(t, nil)
|
|
||||||
|
|
||||||
// Configure allowed ports based on test case.
|
// Configure allowed ports based on test case.
|
||||||
// AI provider tests connect to the specified port, or 443 if not specified.
|
// AI provider tests connect to the specified port, or 443 if not specified.
|
||||||
@@ -649,9 +615,7 @@ func TestProxy_MITM(t *testing.T) {
|
|||||||
var allowedPorts []string
|
var allowedPorts []string
|
||||||
switch {
|
switch {
|
||||||
case tt.passthrough:
|
case tt.passthrough:
|
||||||
parsedTargetURL, err := url.Parse(targetServer.URL)
|
allowedPorts = []string{passthroughURL.Port()}
|
||||||
require.NoError(t, err)
|
|
||||||
allowedPorts = []string{parsedTargetURL.Port()}
|
|
||||||
case tt.targetPort != "":
|
case tt.targetPort != "":
|
||||||
allowedPorts = []string{tt.targetPort}
|
allowedPorts = []string{tt.targetPort}
|
||||||
default:
|
default:
|
||||||
@@ -659,60 +623,20 @@ func TestProxy_MITM(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Start the proxy server pointing to our mock aibridged.
|
// Start the proxy server pointing to our mock aibridged.
|
||||||
srv, err := aibridgeproxyd.New(t.Context(), logger, aibridgeproxyd.Options{
|
srv := newTestProxy(t,
|
||||||
ListenAddr: "127.0.0.1:0",
|
withCoderAccessURL(aibridgedServer.URL),
|
||||||
CoderAccessURL: aibridgedServer.URL,
|
withAllowedPorts(allowedPorts...),
|
||||||
CertFile: certFile,
|
)
|
||||||
KeyFile: keyFile,
|
|
||||||
AllowedPorts: allowedPorts,
|
|
||||||
})
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = srv.Close() })
|
|
||||||
|
|
||||||
proxyAddr := srv.Addr()
|
|
||||||
require.NotEmpty(t, proxyAddr)
|
|
||||||
|
|
||||||
// Wait for the proxy server to be ready.
|
|
||||||
require.Eventually(t, func() bool {
|
|
||||||
conn, err := net.Dial("tcp", proxyAddr)
|
|
||||||
if err != nil {
|
|
||||||
return false
|
|
||||||
}
|
|
||||||
_ = conn.Close()
|
|
||||||
return true
|
|
||||||
}, testutil.WaitShort, testutil.IntervalFast)
|
|
||||||
|
|
||||||
// Load the CA certificate.
|
|
||||||
certPEM, err := os.ReadFile(certFile)
|
|
||||||
require.NoError(t, err)
|
|
||||||
certPool := x509.NewCertPool()
|
|
||||||
certPool.AppendCertsFromPEM(certPEM)
|
|
||||||
|
|
||||||
// Create an HTTP client configured to use the proxy.
|
|
||||||
proxyURL, err := url.Parse("http://" + proxyAddr)
|
|
||||||
require.NoError(t, err)
|
|
||||||
|
|
||||||
client := &http.Client{
|
|
||||||
Transport: &http.Transport{
|
|
||||||
Proxy: http.ProxyURL(proxyURL),
|
|
||||||
ProxyConnectHeader: http.Header{
|
|
||||||
"Proxy-Authorization": []string{makeProxyAuthHeader("test-session-token")},
|
|
||||||
},
|
|
||||||
TLSClientConfig: &tls.Config{
|
|
||||||
MinVersion: tls.VersionTLS12,
|
|
||||||
RootCAs: certPool,
|
|
||||||
},
|
|
||||||
},
|
|
||||||
}
|
|
||||||
|
|
||||||
// Build the target URL:
|
// Build the target URL:
|
||||||
// - For passthrough, target the local mock TLS server.
|
// - For passthrough, target the local mock TLS server.
|
||||||
// - For AI providers, use their real hostnames to trigger routing.
|
// - For AI providers, use their real hostnames to trigger routing.
|
||||||
// Non-default ports are included explicitly; default port (443) is omitted.
|
// Non-default ports are included explicitly; default port (443) is omitted.
|
||||||
var targetURL string
|
var targetURL string
|
||||||
|
var err error
|
||||||
switch {
|
switch {
|
||||||
case tt.passthrough:
|
case tt.passthrough:
|
||||||
targetURL, err = url.JoinPath(targetServer.URL, tt.targetPath)
|
targetURL, err = url.JoinPath(passthroughURL.String(), tt.targetPath)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
case tt.targetPort != "":
|
case tt.targetPort != "":
|
||||||
targetURL, err = url.JoinPath("https://"+tt.targetHost+":"+tt.targetPort, tt.targetPath)
|
targetURL, err = url.JoinPath("https://"+tt.targetHost+":"+tt.targetPort, tt.targetPath)
|
||||||
@@ -723,6 +647,7 @@ func TestProxy_MITM(t *testing.T) {
|
|||||||
}
|
}
|
||||||
|
|
||||||
// Make a request through the proxy to the target URL.
|
// Make a request through the proxy to the target URL.
|
||||||
|
client := newProxyClient(t, srv, makeProxyAuthHeader("test-session-token"))
|
||||||
req, err := http.NewRequestWithContext(t.Context(), http.MethodPost, targetURL, strings.NewReader(`{}`))
|
req, err := http.NewRequestWithContext(t.Context(), http.MethodPost, targetURL, strings.NewReader(`{}`))
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
req.Header.Set("Content-Type", "application/json")
|
req.Header.Set("Content-Type", "application/json")
|
||||||
@@ -761,17 +686,7 @@ func TestServeCACert(t *testing.T) {
|
|||||||
t.Run("Success", func(t *testing.T) {
|
t.Run("Success", func(t *testing.T) {
|
||||||
t.Parallel()
|
t.Parallel()
|
||||||
|
|
||||||
certFile, keyFile := getSharedTestCA(t)
|
srv := newTestProxy(t)
|
||||||
logger := slogtest.Make(t, nil)
|
|
||||||
|
|
||||||
srv, err := aibridgeproxyd.New(t.Context(), logger, aibridgeproxyd.Options{
|
|
||||||
ListenAddr: "127.0.0.1:0",
|
|
||||||
CoderAccessURL: "http://localhost:3000",
|
|
||||||
CertFile: certFile,
|
|
||||||
KeyFile: keyFile,
|
|
||||||
})
|
|
||||||
require.NoError(t, err)
|
|
||||||
t.Cleanup(func() { _ = srv.Close() })
|
|
||||||
|
|
||||||
// Create a request to the CA cert endpoint via the Handler.
|
// Create a request to the CA cert endpoint via the Handler.
|
||||||
req := httptest.NewRequest(http.MethodGet, "/ca-cert.pem", nil)
|
req := httptest.NewRequest(http.MethodGet, "/ca-cert.pem", nil)
|
||||||
@@ -795,6 +710,7 @@ func TestServeCACert(t *testing.T) {
|
|||||||
require.NotNil(t, cert)
|
require.NotNil(t, cert)
|
||||||
|
|
||||||
// Verify it matches the original certificate.
|
// Verify it matches the original certificate.
|
||||||
|
certFile, _ := getSharedTestCA(t)
|
||||||
expectedCertPEM, err := os.ReadFile(certFile)
|
expectedCertPEM, err := os.ReadFile(certFile)
|
||||||
require.NoError(t, err)
|
require.NoError(t, err)
|
||||||
require.Equal(t, expectedCertPEM, body)
|
require.Equal(t, expectedCertPEM, body)
|
||||||
|
|||||||
Reference in New Issue
Block a user