diff --git a/internal/secrets/azure_resolver.go b/internal/secrets/azure_resolver.go index 540a3901..bae8af19 100644 --- a/internal/secrets/azure_resolver.go +++ b/internal/secrets/azure_resolver.go @@ -4,59 +4,18 @@ import ( "context" "encoding/json" "fmt" - "net" "net/http" "strings" - "time" "github.com/Azure/azure-sdk-for-go/sdk/azcore/policy" "github.com/Azure/azure-sdk-for-go/sdk/azidentity" "github.com/Azure/azure-sdk-for-go/sdk/security/keyvault/azsecrets" -) - -// imdsAddresses are the well-known metadata service endpoints that must never -// be reachable from application-level HTTP clients. A request that reaches -// these addresses from user-supplied input is an SSRF attack that can leak -// managed-identity credentials. -var imdsAddresses = map[string]bool{ - "169.254.169.254": true, // Azure/AWS/GCP link-local IMDS (IPv4) - "fd00:ec2::254": true, // AWS IMDS (IPv6) -} -// blockIMDSDialer wraps net.Dialer and rejects outbound connections to IMDS -// addresses before a TCP handshake is attempted. -type blockIMDSDialer struct { - inner net.Dialer -} - -func (d *blockIMDSDialer) DialContext(ctx context.Context, network, addr string) (net.Conn, error) { - host, _, err := net.SplitHostPort(addr) - if err != nil { - host = addr - } - if imdsAddresses[host] { - return nil, fmt.Errorf("connection to metadata endpoint %s is blocked", host) - } - return d.inner.DialContext(ctx, network, addr) -} + "github.com/LeanerCloud/cloud-commitments-go/pkg/httpclient" +) -// imdsBlockingTransport returns an *http.Client whose transport rejects -// connections to Azure/AWS/GCP IMDS link-local addresses, mitigating SSRF -// attacks that could exfiltrate managed-identity credentials. func imdsBlockingTransport() *http.Client { - dialer := &blockIMDSDialer{ - inner: net.Dialer{ - Timeout: 10 * time.Second, - KeepAlive: 30 * time.Second, - }, - } - return &http.Client{ - Timeout: 30 * time.Second, - Transport: &http.Transport{ - DialContext: dialer.DialContext, - TLSHandshakeTimeout: 10 * time.Second, - }, - } + return httpclient.New() } // AzureResolver implements Resolver for Azure Key Vault. @@ -65,9 +24,7 @@ type AzureResolver struct { vaultURL string } -// NewAzureResolver creates a new Azure Key Vault resolver. -// The underlying HTTP client blocks connections to IMDS link-local addresses -// (169.254.169.254, fd00:ec2::254) to mitigate SSRF attacks. +// NewAzureResolver creates a Key Vault resolver using the shared metadata-blocking transport. func NewAzureResolver(ctx context.Context, vaultURL string) (*AzureResolver, error) { // Create a credential using DefaultAzureCredential. // Note: azidentity.NewDefaultAzureCredential does not accept a context parameter, diff --git a/internal/secrets/azure_resolver_httptest_test.go b/internal/secrets/azure_resolver_httptest_test.go index 44c69b44..543b931d 100644 --- a/internal/secrets/azure_resolver_httptest_test.go +++ b/internal/secrets/azure_resolver_httptest_test.go @@ -2,10 +2,12 @@ package secrets import ( "context" + "crypto/tls" + "crypto/x509" "encoding/json" + "net" "net/http" "net/http/httptest" - "strings" "sync/atomic" "testing" "time" @@ -46,22 +48,33 @@ func azureTestHandler(handler http.HandlerFunc) http.HandlerFunc { } } -// newTestAzureResolver creates an AzureResolver backed by a mock HTTPS server. -// TLS is required because the GA Key Vault challenge policy only attaches -// credentials over TLS-protected connections; server.Client() supplies a -// transport that trusts the test server's certificate. -// The handler must NOT reference the server variable (it is created inside this function). +// The Key Vault challenge policy requires TLS before attaching credentials. func newTestAzureResolver(t *testing.T, handler http.HandlerFunc) (*AzureResolver, *httptest.Server) { t.Helper() - server := httptest.NewTLSServer(azureTestHandler(handler)) + listener, err := net.Listen("tcp", "127.0.0.1:0") + require.NoError(t, err) + server := &httptest.Server{ + Listener: listener, + Config: &http.Server{Handler: azureTestHandler(handler)}, + } + server.StartTLS() + t.Cleanup(server.Close) + + httpClient := imdsBlockingTransport() + transport, ok := httpClient.Transport.(*http.Transport) + require.True(t, ok) + roots := x509.NewCertPool() + roots.AddCert(server.Certificate()) + transport.TLSClientConfig = &tls.Config{RootCAs: roots, MinVersion: tls.VersionTLS12} + t.Cleanup(httpClient.CloseIdleConnections) cred := &fakeTokenCredential{} client, err := azsecrets.NewClient(server.URL, cred, &azsecrets.ClientOptions{ ClientOptions: policy.ClientOptions{ - Transport: server.Client(), + Transport: httpClient, Retry: policy.RetryOptions{ - MaxRetries: 0, + MaxRetries: -1, }, }, DisableChallengeResourceVerification: true, @@ -76,7 +89,8 @@ func newTestAzureResolver(t *testing.T, handler http.HandlerFunc) (*AzureResolve func TestAzureResolverReal_GetSecret_Success(t *testing.T) { resolver, server := newTestAzureResolver(t, func(w http.ResponseWriter, r *http.Request) { - assert.True(t, strings.HasPrefix(r.URL.Path, "/secrets/")) + assert.Equal(t, "/secrets/test-secret/", r.URL.Path) + assert.Equal(t, "Bearer fake-access-token", r.Header.Get("Authorization")) resp := map[string]interface{}{ "value": "my-azure-secret-value", "id": "https://myvault.vault.azure.net/secrets/test-secret/abc123", @@ -308,3 +322,21 @@ func TestAzureResolver_IMDSBlocked(t *testing.T) { assert.Contains(t, err.Error(), "blocked", "error message should indicate the connection was blocked") } + +func TestNewAzureResolver_GetSecret_BlocksMetadata(t *testing.T) { + for name, vaultURL := range map[string]string{ + "literal": "https://169.254.169.254", + "mapped": "https://[::ffff:169.254.169.254]", + } { + t.Run(name, func(t *testing.T) { + ctx := context.Background() + resolver, err := NewAzureResolver(ctx, vaultURL) + require.NoError(t, err) + + secret, err := resolver.GetSecret(ctx, "test-secret") + require.Error(t, err) + assert.Empty(t, secret) + assert.Contains(t, err.Error(), "connection to metadata endpoint 169.254.169.254 is blocked") + }) + } +}