Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
51 changes: 4 additions & 47 deletions internal/secrets/azure_resolver.go
Original file line number Diff line number Diff line change
Expand Up @@ -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.
Expand All @@ -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,
Expand Down
52 changes: 42 additions & 10 deletions internal/secrets/azure_resolver_httptest_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand Down Expand Up @@ -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,
Expand All @@ -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",
Expand Down Expand Up @@ -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")
})
}
}
Loading