mirror of
https://github.com/seaweedfs/seaweedfs.git
synced 2026-10-11 16:57:45 +02:00
kms/azure: fix encrypt/decrypt round trip for Azure Key Vault (#11692)
* kms/azure: fix encrypt/decrypt round trip for Azure Key Vault The provider stored the wrapped data key as string(encryptResult.Result) in the JSON envelope, sent the encryption context as AAD to an RSA-OAEP key, and passed a full key URL to the azkeys client in the name position. Each of those breaks a round trip on its own. - split the key URL into the (name, version) pair azkeys expects before Encrypt, Decrypt and GetKey; a plain name still resolves to the latest version, and decrypt keeps the stored version so objects stay readable after a key rotation - store the wrapped key base64-encoded and decode it strictly, so the JSON envelope no longer replaces the raw bytes with U+FFFD - stop sending AAD: Key Vault rejects it on RSA-OAEP with BadParameter, so the encryption context is logged and dropped instead * kms/azure: reject a key URL that names another vault splitKeyID dropped the host of a full Key Vault URL, so a key ID from a different vault silently resolved to this vault's same-named key. Return an error instead of encrypting under a key the caller did not ask for. * kms/azure: bind encryption context via an envelope digest RSA-OAEP rejects AAD, so removing it left the encryption context unauthenticated: a wrapped key could be decrypted under a different object's context. Record a sha256 digest of the marshaled context in the envelope's provider_specific field on encrypt and verify it before calling Decrypt, restoring the binding without AAD. * kms/azure: treat an explicit :443 port as the same vault * kms/azure: accept an absent context digest only for empty contexts * kms/azure: normalize both hosts when comparing key URLs to the vault splitKeyID stripped :443 and a trailing dot from the configured vault but only :443 from the key URL, so a key URL naming the same vault with a trailing DNS dot was rejected before Azure was ever called. Compare both sides through vaultHost so they are normalized identically. * kms/azure: reject non-key vault URLs in splitKeyID, document digest limits A Key Vault URL that does not name a key under /keys/, has an empty key name, or carries extra path segments now fails fast instead of being passed to the client as a key name, where it would surface as a confusing vault-side error. Also note that the envelope context digest is a client-side mismatch check, not vault-authenticated AAD, and only allocate providerSpecific when a context is present. * ci: compile and test azurekms-gated code weed/kms/azure is excluded from every default build, so nothing in CI compiled it; that is how the provider shipped unregistered. Build the tree and run the kms tests with -tags azurekms on every Go change. --------- Co-authored-by: Yi-111-a <> Co-authored-by: Chris Lu <chrislusf@users.noreply.github.com> Co-authored-by: Chris Lu <chris.lu@gmail.com>
This commit is contained in:
3 files changed
+503
-22
No files matched your search
@@ -148,3 +148,21 @@ jobs:
|
||||
# -short skips the e2e suites already covered on amd64.
|
||||
- name: Test linux/386
|
||||
run: cd weed; GOOS=linux GOARCH=386 go test -short ./...
|
||||
|
||||
build-azurekms:
|
||||
name: Build and test with azurekms tag
|
||||
runs-on: ubuntu-latest
|
||||
steps:
|
||||
- name: Check out code into the Go module directory
|
||||
uses: actions/checkout@v7
|
||||
- name: Set up Go
|
||||
uses: actions/setup-go@v7
|
||||
with:
|
||||
go-version-file: 'go.mod'
|
||||
# weed/kms/azure is excluded from every default build, so without this
|
||||
# nothing compiles it and tag-gated code silently rots.
|
||||
- name: Build and test with azurekms
|
||||
run: |
|
||||
cd weed
|
||||
go build -tags azurekms ./...
|
||||
go test -tags azurekms ./kms/...
|
||||
+122
-22
@@ -5,9 +5,12 @@ package azure
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"crypto/sha256"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"net/http"
|
||||
"net/url"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
@@ -26,6 +29,93 @@ func init() {
|
||||
seaweedkms.RegisterProvider("azure", NewAzureKMSProvider)
|
||||
}
|
||||
|
||||
// vaultHost returns the host of a configured Key Vault URL, or "" when the URL
|
||||
// carries none. Hosts are compared case-insensitively, without a trailing dot,
|
||||
// and without an explicit :443, which names the same vault as no port.
|
||||
func vaultHost(vaultURL string) string {
|
||||
parsed, err := url.Parse(vaultURL)
|
||||
if err != nil {
|
||||
return ""
|
||||
}
|
||||
host := strings.TrimSuffix(strings.ToLower(parsed.Host), ":443")
|
||||
return strings.TrimSuffix(host, ".")
|
||||
}
|
||||
|
||||
// splitKeyID turns a Key Vault key identifier into the (name, version) pair the
|
||||
// azkeys client expects. A plain key name, or anything that is not a URL, is
|
||||
// returned unchanged together with an empty version, which the client resolves
|
||||
// to the latest version.
|
||||
//
|
||||
// A URL naming a different vault than the one this provider is configured for
|
||||
// is rejected: the client addresses only its own vault, so dropping the host
|
||||
// would silently encrypt under this vault's same-named key. A URL for this
|
||||
// vault that does not name a key under /keys/ is rejected too: passing the URL
|
||||
// on as a key name would only surface as a confusing vault-side 404.
|
||||
func (p *AzureKMSProvider) splitKeyID(keyID string) (string, string, error) {
|
||||
parsed, err := url.Parse(keyID)
|
||||
if err != nil || parsed.Scheme == "" || parsed.Host == "" {
|
||||
return keyID, "", nil
|
||||
}
|
||||
if host := vaultHost(p.vaultURL); host != "" && vaultHost(keyID) != host {
|
||||
return "", "", fmt.Errorf("key ID %q names vault %q, but this provider is configured for %q", keyID, parsed.Host, host)
|
||||
}
|
||||
parts := strings.Split(strings.Trim(parsed.Path, "/"), "/")
|
||||
if len(parts) < 2 || parts[0] != "keys" || parts[1] == "" {
|
||||
return "", "", fmt.Errorf("key ID %q is a Key Vault URL but does not name a key", keyID)
|
||||
}
|
||||
if len(parts) == 2 {
|
||||
return parts[1], "", nil
|
||||
}
|
||||
if len(parts) == 3 {
|
||||
return parts[1], parts[2], nil
|
||||
}
|
||||
return "", "", fmt.Errorf("key ID %q has too many path segments for a key URL", keyID)
|
||||
}
|
||||
|
||||
// encodeCiphertext stores the wrapped data key in the JSON envelope as base64.
|
||||
// The envelope is JSON, so raw binary would be replaced by U+FFFD on the way in.
|
||||
func encodeCiphertext(ciphertext []byte) string {
|
||||
return base64.StdEncoding.EncodeToString(ciphertext)
|
||||
}
|
||||
|
||||
// contextDigest digests the encryption context for storage in the envelope.
|
||||
// RSA-OAEP takes no AAD, so the context is bound to the wrapped key by
|
||||
// recording its hash and checking it on decrypt instead. Unlike KMS-checked
|
||||
// AAD this is a client-side mismatch check only — it is not authenticated by
|
||||
// the vault, and an actor who can rewrite an object's envelope could swap in
|
||||
// a ciphertext and digest from another object.
|
||||
func contextDigest(context map[string]string) string {
|
||||
encoded, _ := json.Marshal(context) // map keys marshal in sorted order
|
||||
sum := sha256.Sum256(encoded)
|
||||
return base64.StdEncoding.EncodeToString(sum[:])
|
||||
}
|
||||
|
||||
// checkContext verifies the request's encryption context matches the context
|
||||
// recorded when the key was wrapped.
|
||||
func checkContext(envelope *seaweedkms.CiphertextEnvelope, context map[string]string) error {
|
||||
recorded, _ := envelope.ProviderSpecific["encryption_context_sha256"].(string)
|
||||
if recorded == "" {
|
||||
if len(context) == 0 {
|
||||
return nil
|
||||
}
|
||||
} else if recorded == contextDigest(context) {
|
||||
return nil
|
||||
}
|
||||
return fmt.Errorf("encryption context does not match the wrapped key")
|
||||
}
|
||||
|
||||
// decodeCiphertext reads back what encodeCiphertext wrote.
|
||||
func decodeCiphertext(ciphertext string) ([]byte, error) {
|
||||
decoded, err := base64.StdEncoding.Strict().DecodeString(ciphertext)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("wrapped data key is not valid base64: %w", err)
|
||||
}
|
||||
if len(decoded) == 0 {
|
||||
return nil, fmt.Errorf("wrapped data key is empty")
|
||||
}
|
||||
return decoded, nil
|
||||
}
|
||||
|
||||
// AzureKMSProvider implements the KMSProvider interface using Azure Key Vault
|
||||
type AzureKMSProvider struct {
|
||||
client *azkeys.Client
|
||||
@@ -155,19 +245,23 @@ func (p *AzureKMSProvider) GenerateDataKey(ctx context.Context, req *seaweedkms.
|
||||
Value: dataKey,
|
||||
}
|
||||
|
||||
// Add encryption context as Additional Authenticated Data (AAD) if provided
|
||||
// RSA-OAEP takes no additional authenticated data: Key Vault rejects the
|
||||
// request with BadParameter when AAD is set on an RSA-OAEP key. The
|
||||
// context is bound to the wrapped key via a digest in the envelope
|
||||
// instead, which Decrypt verifies.
|
||||
var providerSpecific map[string]interface{}
|
||||
if len(req.EncryptionContext) > 0 {
|
||||
// Marshal encryption context to JSON for deterministic AAD
|
||||
aadBytes, err := json.Marshal(req.EncryptionContext)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal encryption context: %w", err)
|
||||
providerSpecific = map[string]interface{}{
|
||||
"encryption_context_sha256": contextDigest(req.EncryptionContext),
|
||||
}
|
||||
encryptParams.AAD = aadBytes
|
||||
glog.V(4).Infof("Azure KMS: Using encryption context as AAD for key %s", req.KeyID)
|
||||
}
|
||||
|
||||
// Call Azure Key Vault to encrypt the data key
|
||||
encryptResult, err := p.client.Encrypt(ctx, req.KeyID, "", encryptParams, nil)
|
||||
keyName, keyVersion, err := p.splitKeyID(req.KeyID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
encryptResult, err := p.client.Encrypt(ctx, keyName, keyVersion, encryptParams, nil)
|
||||
if err != nil {
|
||||
return nil, p.convertAzureError(err, req.KeyID)
|
||||
}
|
||||
@@ -179,7 +273,7 @@ func (p *AzureKMSProvider) GenerateDataKey(ctx context.Context, req *seaweedkms.
|
||||
}
|
||||
|
||||
// Create standardized envelope format for consistent API behavior
|
||||
envelopeBlob, err := seaweedkms.CreateEnvelope("azure", actualKeyID, string(encryptResult.Result), nil)
|
||||
envelopeBlob, err := seaweedkms.CreateEnvelope("azure", actualKeyID, encodeCiphertext(encryptResult.Result), providerSpecific)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create ciphertext envelope: %w", err)
|
||||
}
|
||||
@@ -215,8 +309,11 @@ func (p *AzureKMSProvider) Decrypt(ctx context.Context, req *seaweedkms.DecryptR
|
||||
return nil, fmt.Errorf("envelope missing key ID")
|
||||
}
|
||||
|
||||
// Convert string back to bytes
|
||||
ciphertext := []byte(envelope.Ciphertext)
|
||||
// Convert the base64 envelope field back to the raw wrapped key
|
||||
ciphertext, err := decodeCiphertext(envelope.Ciphertext)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("invalid Azure ciphertext envelope: %w", err)
|
||||
}
|
||||
|
||||
// Prepare decryption parameters
|
||||
decryptAlgorithm := azkeys.JSONWebKeyEncryptionAlgorithmRSAOAEP256
|
||||
@@ -225,20 +322,19 @@ func (p *AzureKMSProvider) Decrypt(ctx context.Context, req *seaweedkms.DecryptR
|
||||
Value: ciphertext,
|
||||
}
|
||||
|
||||
// Add encryption context as Additional Authenticated Data (AAD) if provided
|
||||
if len(req.EncryptionContext) > 0 {
|
||||
// Marshal encryption context to JSON for deterministic AAD (must match encryption)
|
||||
aadBytes, err := json.Marshal(req.EncryptionContext)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal encryption context: %w", err)
|
||||
}
|
||||
decryptParams.AAD = aadBytes
|
||||
glog.V(4).Infof("Azure KMS: Using encryption context as AAD for decryption of key %s", keyID)
|
||||
// RSA-OAEP takes no AAD, so the encryption context is bound via a digest
|
||||
// recorded in the envelope rather than authenticated by the vault.
|
||||
if err := checkContext(envelope, req.EncryptionContext); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
// Call Azure Key Vault to decrypt the data key
|
||||
glog.V(4).Infof("Azure KMS: Decrypting data key using key %s", keyID)
|
||||
decryptResult, err := p.client.Decrypt(ctx, keyID, "", decryptParams, nil)
|
||||
decryptName, decryptVersion, err := p.splitKeyID(keyID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
decryptResult, err := p.client.Decrypt(ctx, decryptName, decryptVersion, decryptParams, nil)
|
||||
if err != nil {
|
||||
return nil, p.convertAzureError(err, keyID)
|
||||
}
|
||||
@@ -270,7 +366,11 @@ func (p *AzureKMSProvider) DescribeKey(ctx context.Context, req *seaweedkms.Desc
|
||||
|
||||
// Get key from Azure Key Vault
|
||||
glog.V(4).Infof("Azure KMS: Describing key %s", req.KeyID)
|
||||
result, err := p.client.GetKey(ctx, req.KeyID, "", nil)
|
||||
describeName, describeVersion, err := p.splitKeyID(req.KeyID)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
result, err := p.client.GetKey(ctx, describeName, describeVersion, nil)
|
||||
if err != nil {
|
||||
return nil, p.convertAzureError(err, req.KeyID)
|
||||
}
|
||||
|
||||
@@ -0,0 +1,363 @@
|
||||
//go:build azurekms
|
||||
|
||||
package azure
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"encoding/json"
|
||||
"io"
|
||||
"net/http"
|
||||
"strings"
|
||||
"testing"
|
||||
"time"
|
||||
|
||||
"github.com/Azure/azure-sdk-for-go/sdk/azcore"
|
||||
"github.com/Azure/azure-sdk-for-go/sdk/azcore/policy"
|
||||
"github.com/Azure/azure-sdk-for-go/sdk/keyvault/azkeys"
|
||||
|
||||
seaweedkms "github.com/seaweedfs/seaweedfs/weed/kms"
|
||||
)
|
||||
|
||||
func TestSplitKeyID(t *testing.T) {
|
||||
provider := &AzureKMSProvider{vaultURL: "https://myvault.vault.azure.net"}
|
||||
|
||||
tests := []struct {
|
||||
name string
|
||||
keyID string
|
||||
wantName string
|
||||
wantVersion string
|
||||
wantErr bool
|
||||
}{
|
||||
{
|
||||
name: "plain key name",
|
||||
keyID: "my-key",
|
||||
wantName: "my-key",
|
||||
},
|
||||
{
|
||||
name: "key url without version",
|
||||
keyID: "https://myvault.vault.azure.net/keys/my-key",
|
||||
wantName: "my-key",
|
||||
},
|
||||
{
|
||||
name: "key url with version",
|
||||
keyID: "https://myvault.vault.azure.net/keys/my-key/abc123",
|
||||
wantName: "my-key",
|
||||
wantVersion: "abc123",
|
||||
},
|
||||
{
|
||||
name: "key url with trailing slash",
|
||||
keyID: "https://myvault.vault.azure.net/keys/my-key/",
|
||||
wantName: "my-key",
|
||||
},
|
||||
{
|
||||
name: "key url with a mixed case host",
|
||||
keyID: "https://MyVault.vault.azure.net/keys/my-key",
|
||||
wantName: "my-key",
|
||||
},
|
||||
{
|
||||
name: "key url with a non keys path",
|
||||
keyID: "https://myvault.vault.azure.net/secrets/my-secret",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "key url with extra path segments",
|
||||
keyID: "https://myvault.vault.azure.net/keys/my-key/abc123/extra",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "key url with an empty key name",
|
||||
keyID: "https://myvault.vault.azure.net/keys//abc123",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "name that only looks like a url",
|
||||
keyID: "my-key:abc123",
|
||||
wantName: "my-key:abc123",
|
||||
},
|
||||
{
|
||||
name: "key url with explicit default port",
|
||||
keyID: "https://myvault.vault.azure.net:443/keys/my-key",
|
||||
wantName: "my-key",
|
||||
},
|
||||
{
|
||||
name: "key url with non-default port",
|
||||
keyID: "https://myvault.vault.azure.net:8443/keys/my-key",
|
||||
wantErr: true,
|
||||
},
|
||||
{
|
||||
name: "key url from another vault",
|
||||
keyID: "https://othervault.vault.azure.net/keys/my-key",
|
||||
wantErr: true,
|
||||
},
|
||||
}
|
||||
|
||||
for _, test := range tests {
|
||||
t.Run(test.name, func(t *testing.T) {
|
||||
name, version, err := provider.splitKeyID(test.keyID)
|
||||
if test.wantErr {
|
||||
if err == nil {
|
||||
t.Fatalf("splitKeyID(%q) succeeded, want error", test.keyID)
|
||||
}
|
||||
return
|
||||
}
|
||||
if err != nil {
|
||||
t.Fatalf("splitKeyID(%q): %v", test.keyID, err)
|
||||
}
|
||||
if name != test.wantName {
|
||||
t.Fatalf("name = %q, want %q", name, test.wantName)
|
||||
}
|
||||
if version != test.wantVersion {
|
||||
t.Fatalf("version = %q, want %q", version, test.wantVersion)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// The client addresses only the vault it was built for, so a key URL naming a
|
||||
// different vault must not resolve to this vault's same-named key.
|
||||
func TestSplitKeyIDExplicitDefaultPort(t *testing.T) {
|
||||
provider := &AzureKMSProvider{vaultURL: "https://myvault.vault.azure.net:443"}
|
||||
name, _, err := provider.splitKeyID("https://myvault.vault.azure.net/keys/my-key")
|
||||
if err != nil || name != "my-key" {
|
||||
t.Fatalf("splitKeyID = %q, %v; want my-key, nil", name, err)
|
||||
}
|
||||
}
|
||||
|
||||
func TestSplitKeyIDRejectsForeignVault(t *testing.T) {
|
||||
provider := &AzureKMSProvider{vaultURL: "https://myvault.vault.azure.net/"}
|
||||
|
||||
if _, _, err := provider.splitKeyID("https://evil.vault.azure.net/keys/my-key/abc123"); err == nil {
|
||||
t.Fatal("splitKeyID accepted a key URL from another vault")
|
||||
}
|
||||
}
|
||||
|
||||
// A trailing DNS dot names the same vault, so both sides of the comparison are
|
||||
// normalized the same way.
|
||||
func TestSplitKeyIDTrailingDot(t *testing.T) {
|
||||
tests := []struct {
|
||||
name string
|
||||
vaultURL string
|
||||
keyID string
|
||||
}{
|
||||
{
|
||||
name: "trailing dot on the key url",
|
||||
vaultURL: "https://myvault.vault.azure.net",
|
||||
keyID: "https://myvault.vault.azure.net./keys/my-key/abc123",
|
||||
},
|
||||
{
|
||||
name: "trailing dot on the configured vault",
|
||||
vaultURL: "https://myvault.vault.azure.net./",
|
||||
keyID: "https://myvault.vault.azure.net/keys/my-key/abc123",
|
||||
},
|
||||
}
|
||||
|
||||
for _, tt := range tests {
|
||||
t.Run(tt.name, func(t *testing.T) {
|
||||
provider := &AzureKMSProvider{vaultURL: tt.vaultURL}
|
||||
name, version, err := provider.splitKeyID(tt.keyID)
|
||||
if err != nil || name != "my-key" || version != "abc123" {
|
||||
t.Fatalf("splitKeyID = %q, %q, %v; want my-key, abc123, nil", name, version, err)
|
||||
}
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
// The envelope is JSON, so a raw binary wrapped key would be replaced by
|
||||
// U+FFFD on the way in and could never be decrypted.
|
||||
func TestCiphertextEnvelopeRoundTrip(t *testing.T) {
|
||||
wrapped := make([]byte, 256)
|
||||
if _, err := rand.Read(wrapped); err != nil {
|
||||
t.Fatalf("generate wrapped key: %v", err)
|
||||
}
|
||||
|
||||
keyID := "https://myvault.vault.azure.net/keys/my-key/abc123"
|
||||
envelopeBlob, err := seaweedkms.CreateEnvelope("azure", keyID, encodeCiphertext(wrapped), nil)
|
||||
if err != nil {
|
||||
t.Fatalf("create envelope: %v", err)
|
||||
}
|
||||
|
||||
envelope, err := seaweedkms.ParseEnvelope(envelopeBlob)
|
||||
if err != nil {
|
||||
t.Fatalf("parse envelope: %v", err)
|
||||
}
|
||||
if envelope.KeyID != keyID {
|
||||
t.Fatalf("key id = %q, want %q", envelope.KeyID, keyID)
|
||||
}
|
||||
|
||||
decoded, err := decodeCiphertext(envelope.Ciphertext)
|
||||
if err != nil {
|
||||
t.Fatalf("decode ciphertext: %v", err)
|
||||
}
|
||||
if !bytes.Equal(decoded, wrapped) {
|
||||
t.Fatal("decoded wrapped key differs from the encrypted one")
|
||||
}
|
||||
}
|
||||
|
||||
func TestDecodeCiphertextRejectsInvalidInput(t *testing.T) {
|
||||
for _, ciphertext := range []string{"", "not base64!!", "a"} {
|
||||
if _, err := decodeCiphertext(ciphertext); err == nil {
|
||||
t.Fatalf("decodeCiphertext(%q) succeeded, want error", ciphertext)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// RSA-OAEP cannot authenticate the encryption context, so the context is
|
||||
// bound to the wrapped key via a digest in the envelope.
|
||||
func TestEncryptionContextBinding(t *testing.T) {
|
||||
context := map[string]string{"aws:s3:bucket": "bucket", "aws:s3:object": "key"}
|
||||
envelopeBlob, err := seaweedkms.CreateEnvelope("azure", "my-key", "d3JhcHBlZA==", map[string]interface{}{
|
||||
"encryption_context_sha256": contextDigest(context),
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("create envelope: %v", err)
|
||||
}
|
||||
envelope, err := seaweedkms.ParseEnvelope(envelopeBlob)
|
||||
if err != nil {
|
||||
t.Fatalf("parse envelope: %v", err)
|
||||
}
|
||||
|
||||
if err := checkContext(envelope, context); err != nil {
|
||||
t.Fatalf("matching context rejected: %v", err)
|
||||
}
|
||||
if err := checkContext(envelope, map[string]string{"aws:s3:object": "other"}); err == nil {
|
||||
t.Fatal("mismatched context accepted")
|
||||
}
|
||||
if err := checkContext(envelope, nil); err == nil {
|
||||
t.Fatal("missing context accepted for a context-bound key")
|
||||
}
|
||||
}
|
||||
|
||||
type fakeCredential struct{}
|
||||
|
||||
func (fakeCredential) GetToken(_ context.Context, _ policy.TokenRequestOptions) (azcore.AccessToken, error) {
|
||||
return azcore.AccessToken{Token: "token", ExpiresOn: time.Now().Add(time.Hour)}, nil
|
||||
}
|
||||
|
||||
type fakeTransport struct {
|
||||
do func(*http.Request) (*http.Response, error)
|
||||
}
|
||||
|
||||
func (t fakeTransport) Do(req *http.Request) (*http.Response, error) {
|
||||
return t.do(req)
|
||||
}
|
||||
|
||||
func jsonResponse(status int, body string) *http.Response {
|
||||
return &http.Response{
|
||||
StatusCode: status,
|
||||
Header: http.Header{"Content-Type": []string{"application/json"}},
|
||||
Body: io.NopCloser(strings.NewReader(body)),
|
||||
}
|
||||
}
|
||||
|
||||
// Drive GenerateDataKey and Decrypt through a fake Key Vault transport so the
|
||||
// whole path — key ID split, request body shape, envelope encoding — is
|
||||
// exercised, not just the helpers.
|
||||
func TestGenerateDataKeyDecryptRoundTrip(t *testing.T) {
|
||||
vaultURL := "https://myvault.vault.azure.net"
|
||||
var wrapped []byte
|
||||
var sawAAD bool
|
||||
|
||||
client, err := azkeys.NewClient(vaultURL, fakeCredential{}, &azkeys.ClientOptions{
|
||||
ClientOptions: azcore.ClientOptions{
|
||||
Transport: fakeTransport{do: func(req *http.Request) (*http.Response, error) {
|
||||
if req.Header.Get("Authorization") == "" {
|
||||
// Key Vault answers an unauthenticated request with the
|
||||
// challenge the client then satisfies.
|
||||
resp := jsonResponse(401, `{"error":{"code":"Unauthorized"}}`)
|
||||
resp.Header.Set("WWW-Authenticate", `Bearer authorization="https://login.windows.net/tenant", resource="https://vault.azure.net"`)
|
||||
return resp, nil
|
||||
}
|
||||
var body []byte
|
||||
if rc, err := req.GetBody(); err == nil {
|
||||
body, _ = io.ReadAll(rc)
|
||||
rc.Close()
|
||||
} else if req.Body != nil {
|
||||
body, _ = io.ReadAll(req.Body)
|
||||
}
|
||||
sawAAD = bytes.Contains(body, []byte(`"aad"`))
|
||||
var params struct {
|
||||
Value string `json:"value"`
|
||||
}
|
||||
if err := json.Unmarshal(body, ¶ms); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
value, err := base64.RawURLEncoding.DecodeString(params.Value)
|
||||
if err != nil {
|
||||
return nil, err
|
||||
}
|
||||
switch {
|
||||
case strings.HasSuffix(req.URL.Path, "/encrypt"):
|
||||
wrapped = value
|
||||
return jsonResponse(200, `{"kid":"`+vaultURL+`/keys/my-key/abc123","value":"`+base64.RawURLEncoding.EncodeToString(wrapped)+`"}`), nil
|
||||
case strings.HasSuffix(req.URL.Path, "/decrypt"):
|
||||
if !bytes.Equal(value, wrapped) {
|
||||
return jsonResponse(400, `{"error":{"code":"BadParameter"}}`), nil
|
||||
}
|
||||
return jsonResponse(200, `{"kid":"`+vaultURL+`/keys/my-key/abc123","value":"`+params.Value+`"}`), nil
|
||||
default:
|
||||
return jsonResponse(404, `{"error":{"code":"NotFound"}}`), nil
|
||||
}
|
||||
}},
|
||||
},
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("new client: %v", err)
|
||||
}
|
||||
|
||||
provider := &AzureKMSProvider{client: client, vaultURL: vaultURL}
|
||||
contextMap := map[string]string{"aws:s3:bucket": "bucket"}
|
||||
|
||||
resp, err := provider.GenerateDataKey(context.Background(), &seaweedkms.GenerateDataKeyRequest{
|
||||
KeyID: vaultURL + "/keys/my-key",
|
||||
KeySpec: seaweedkms.KeySpecAES256,
|
||||
EncryptionContext: contextMap,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("GenerateDataKey: %v", err)
|
||||
}
|
||||
if sawAAD {
|
||||
t.Fatal("encrypt request carried AAD")
|
||||
}
|
||||
|
||||
decrypted, err := provider.Decrypt(context.Background(), &seaweedkms.DecryptRequest{
|
||||
CiphertextBlob: resp.CiphertextBlob,
|
||||
EncryptionContext: contextMap,
|
||||
})
|
||||
if err != nil {
|
||||
t.Fatalf("Decrypt: %v", err)
|
||||
}
|
||||
if !bytes.Equal(decrypted.Plaintext, resp.Plaintext) {
|
||||
t.Fatal("decrypted key differs from generated key")
|
||||
}
|
||||
|
||||
if _, err := provider.Decrypt(context.Background(), &seaweedkms.DecryptRequest{
|
||||
CiphertextBlob: resp.CiphertextBlob,
|
||||
EncryptionContext: map[string]string{"aws:s3:bucket": "other"},
|
||||
}); err == nil {
|
||||
t.Fatal("Decrypt succeeded with a different encryption context")
|
||||
}
|
||||
}
|
||||
|
||||
func TestCheckContextEmpty(t *testing.T) {
|
||||
// Keys wrapped without a context carry no digest; decrypting them with
|
||||
// no context must succeed, and with a context must fail.
|
||||
blob, err := seaweedkms.CreateEnvelope("azure", "my-key", "d3JhcHBlZA==", nil)
|
||||
if err != nil {
|
||||
t.Fatalf("create envelope: %v", err)
|
||||
}
|
||||
envelope, err := seaweedkms.ParseEnvelope(blob)
|
||||
if err != nil {
|
||||
t.Fatalf("parse envelope: %v", err)
|
||||
}
|
||||
for _, context := range []map[string]string{nil, {}} {
|
||||
if err := checkContext(envelope, context); err != nil {
|
||||
t.Fatalf("empty context rejected for an unbound key: %v", err)
|
||||
}
|
||||
}
|
||||
if err := checkContext(envelope, map[string]string{"k": "v"}); err == nil {
|
||||
t.Fatal("context accepted for an unbound key")
|
||||
}
|
||||
}
|
||||
Reference in new issue
Block a user