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
175 changes: 167 additions & 8 deletions credentials/vault/aws.go
Original file line number Diff line number Diff line change
Expand Up @@ -5,11 +5,24 @@ package vault

import (
"context"
"crypto/sha256"
"encoding/base64"
"encoding/hex"
"encoding/json"
"fmt"
"net/http"
"net/url"
"strconv"

"strings"
"time"

"github.com/aws/aws-sdk-go-v2/aws"
v4 "github.com/aws/aws-sdk-go-v2/aws/signer/v4"
"github.com/aws/aws-sdk-go-v2/service/iam"
"github.com/aws/aws-sdk-go-v2/service/sts"
smithyendpoints "github.com/aws/smithy-go/endpoints"
"github.com/hashicorp/go-hclog"
"github.com/hashicorp/go-secure-stdlib/awsutil"
awsutil "github.com/hashicorp/go-secure-stdlib/awsutil/v2"
corev1 "k8s.io/api/core/v1"
"k8s.io/apimachinery/pkg/types"
ctrlclient "sigs.k8s.io/controller-runtime/pkg/client"
Expand All @@ -21,6 +34,152 @@ import (
"github.com/hashicorp/vault-secrets-operator/helpers"
)

const (
iamServerIDHeader = "X-Vault-AWS-IAM-Server-ID"
stsGetCallerIdentityBody = "Action=GetCallerIdentity&Version=2011-06-15"
stsContentType = "application/x-www-form-urlencoded; charset=utf-8"
stsSigningName = "sts"
)

// customSTSEndpointResolver implements sts.EndpointResolverV2 for a fixed endpoint URL.
type customSTSEndpointResolver struct {
endpointURL string
}

func (r *customSTSEndpointResolver) ResolveEndpoint(_ context.Context, _ sts.EndpointParameters) (smithyendpoints.Endpoint, error) {
uri, err := url.Parse(r.endpointURL)
if err != nil {
return smithyendpoints.Endpoint{}, fmt.Errorf("failed to parse custom STS endpoint URL: %w", err)
}
return smithyendpoints.Endpoint{URI: *uri}, nil
}

// customIAMEndpointResolver implements iam.EndpointResolverV2 for a fixed endpoint URL.
type customIAMEndpointResolver struct {
endpointURL string
}

func (r *customIAMEndpointResolver) ResolveEndpoint(_ context.Context, _ iam.EndpointParameters) (smithyendpoints.Endpoint, error) {
uri, err := url.Parse(r.endpointURL)
if err != nil {
return smithyendpoints.Endpoint{}, fmt.Errorf("failed to parse custom IAM endpoint URL: %w", err)
}
return smithyendpoints.Endpoint{URI: *uri}, nil
}

type stsSigningEndpoint struct {
requestURL string
signingName string
signingRegion string
}

// generateLoginData builds the Vault AWS IAM login payload by constructing and
// signing a sts:GetCallerIdentity HTTP request using AWS SDK v2.
// It replicates the behaviour of the now-removed awsutil.GenerateLoginData from
// go-secure-stdlib/awsutil v0.
func generateLoginData(ctx context.Context, awsConfig *aws.Config, headerValue, stsEndpoint string) (map[string]interface{}, error) {
if awsConfig == nil || awsConfig.Credentials == nil {
return nil, fmt.Errorf("AWS credentials are not configured")
}

credentials, err := awsConfig.Credentials.Retrieve(ctx)
if err != nil {
return nil, fmt.Errorf("failed to retrieve AWS credentials: %w", err)
}

region := awsConfig.Region
if region == "" {
region = awsutil.DefaultRegion
}

endpoint, err := resolveSTSSigningEndpoint(region, stsEndpoint)
if err != nil {
return nil, err
}

req, body, err := buildSignedGetCallerIdentityRequest(ctx, credentials, endpoint, region, headerValue)
if err != nil {
return nil, err
}

headers := req.Header.Clone()
if headers.Get("Host") == "" {
headers.Set("Host", req.URL.Host)
}

headersJSON, err := json.Marshal(headers)
if err != nil {
return nil, fmt.Errorf("failed to marshal request headers: %w", err)
}

loginData := map[string]interface{}{
"iam_http_request_method": req.Method,
"iam_request_url": base64.StdEncoding.EncodeToString([]byte(req.URL.String())),
"iam_request_headers": base64.StdEncoding.EncodeToString(headersJSON),
"iam_request_body": base64.StdEncoding.EncodeToString([]byte(body)),
}
return loginData, nil
}

// resolveSTSSigningEndpoint returns the STS URL and signing metadata for the given region.
// When a custom endpoint is provided it is used directly; otherwise a regional
// endpoint of the form https://sts.<region>.amazonaws.com is used, ensuring
// the signed Host header always matches the target region.
func resolveSTSSigningEndpoint(region, endpointURL string) (stsSigningEndpoint, error) {
if endpointURL != "" {
uri, err := url.Parse(endpointURL)
if err != nil {
return stsSigningEndpoint{}, fmt.Errorf("failed to parse custom STS endpoint URL: %w", err)
}
if uri.Scheme == "" || uri.Host == "" {
return stsSigningEndpoint{}, fmt.Errorf("invalid custom STS endpoint URL %q", endpointURL)
}
return stsSigningEndpoint{
requestURL: uri.String(),
signingName: stsSigningName,
signingRegion: region,
}, nil
}
return stsSigningEndpoint{
requestURL: fmt.Sprintf("https://sts.%s.amazonaws.com", region),
signingName: stsSigningName,
signingRegion: region,
}, nil
}

func buildSignedGetCallerIdentityRequest(ctx context.Context, credentials aws.Credentials, endpoint stsSigningEndpoint, region, headerValue string) (*http.Request, string, error) {
body := stsGetCallerIdentityBody
req, err := http.NewRequestWithContext(ctx, http.MethodPost, endpoint.requestURL, strings.NewReader(body))
if err != nil {
return nil, "", fmt.Errorf("failed to build GetCallerIdentity request: %w", err)
}

req.Header.Set("Content-Type", stsContentType)
if headerValue != "" {
req.Header.Set(iamServerIDHeader, headerValue)
}

signingRegion := endpoint.signingRegion
if signingRegion == "" {
signingRegion = region
}
if signingRegion == "" {
signingRegion = awsutil.DefaultRegion
}
signingName := endpoint.signingName
if signingName == "" {
signingName = stsSigningName
}

payloadHash := sha256.Sum256([]byte(body))
signer := v4.NewSigner()
if err := signer.SignHTTP(ctx, credentials, req, hex.EncodeToString(payloadHash[:]), signingName, signingRegion, time.Now().UTC()); err != nil {
return nil, "", fmt.Errorf("failed to sign GetCallerIdentity request: %w", err)
}

return req, body, nil
}

const (
AWSAnnotationRole = "eks.amazonaws.com/role-arn"
AWSAnnotationAudience = "eks.amazonaws.com/audience"
Expand All @@ -30,6 +189,7 @@ const (
K8sRootCA = "kube-root-ca.crt"
)

// Compile-time assertion that AWSCredentialProvider implements CredentialProvider.
var _ CredentialProvider = (*AWSCredentialProvider)(nil)

type AWSCredentialProvider struct {
Expand Down Expand Up @@ -144,19 +304,18 @@ func (l *AWSCredentialProvider) GetCreds(ctx context.Context, client ctrlclient.
return nil, err
}

// TODO: convert logr to something compatible with hclog for use in the
// awsutil functions
config.Logger = hclog.Default()
config.Logger.SetLevel(hclog.Debug)

creds, err := config.GenerateCredentialChain(awsutil.WithSkipWebIdentityValidity(true))
// GenerateCredentialChain returns *aws.Config (SDK v2).
awsCfg, err := config.GenerateCredentialChain(ctx)
if err != nil {
return nil, err
}

headerValue := l.authObj.Spec.AWS.HeaderValue

loginData, err := awsutil.GenerateLoginData(creds, headerValue, config.Region, config.Logger)
loginData, err := generateLoginData(ctx, awsCfg, headerValue, l.authObj.Spec.AWS.STSEndpoint)
if err != nil {
return nil, err
}
Expand All @@ -177,10 +336,10 @@ func (l *AWSCredentialProvider) getCredentialsConfig(credsSecret *corev1.Secret,
config.RoleSessionName = l.authObj.Spec.AWS.SessionName
}
if l.authObj.Spec.AWS.STSEndpoint != "" {
config.STSEndpoint = l.authObj.Spec.AWS.STSEndpoint
config.STSEndpointResolver = &customSTSEndpointResolver{endpointURL: l.authObj.Spec.AWS.STSEndpoint}
}
if l.authObj.Spec.AWS.IAMEndpoint != "" {
config.IAMEndpoint = l.authObj.Spec.AWS.IAMEndpoint
config.IAMEndpointResolver = &customIAMEndpointResolver{endpointURL: l.authObj.Spec.AWS.IAMEndpoint}
}

if credsSecret != nil {
Expand Down
Loading