diff --git a/Makefile b/Makefile index 3e29803f..3e63e2dd 100644 --- a/Makefile +++ b/Makefile @@ -6,7 +6,8 @@ test-entra-node test-entra-node-local \ test-apim-java \ test-cosmos test-cosmos-mongo test-cosmos-postgresql test-cosmos-cassandra test-cosmos-gremlin test-cosmos-table test-cosmos-nosql test-cosmos-all \ - test-sql test-mysql test-mariadb test-terraform-compat test-opentofu-compat test-azcli test-iac-compat compat-docker test-compat clean + test-sql test-mysql test-mariadb test-terraform-compat test-opentofu-compat test-azcli test-iac-compat compat-docker test-compat clean \ + smoke-native-crypto MVN = ./mvnw PORT = 4577 @@ -470,6 +471,30 @@ test: build $(MVN) test $(MAKE) compat-docker +# ── Native-Image Crypto Smoke Gate ──────────────────────────────────────────── + +smoke-native-crypto: + $(MVN) package -Dnative -DskipTests -B -Dquarkus.native.additional-build-args-append="-Ob" -q + ./target/*-runner & echo $$! > /tmp/floci-az-native.pid + @echo "Waiting for floci-az native runner on port $(PORT)..." + @EXIT=0; \ + ATTEMPTS=0; \ + until curl -sf http://localhost:$(PORT)/health > /dev/null 2>&1; do \ + ATTEMPTS=$$((ATTEMPTS + 1)); \ + if [ $$ATTEMPTS -ge 120 ]; then \ + echo "floci-az native runner did not become healthy after 120s" >&2; \ + EXIT=1; \ + break; \ + fi; \ + sleep 1; \ + done; \ + if [ $$EXIT -eq 0 ]; then \ + bash scripts/native-crypto-smoke.sh || EXIT=$$?; \ + fi; \ + kill $$(cat /tmp/floci-az-native.pid 2>/dev/null) 2>/dev/null || true; \ + rm -f /tmp/floci-az-native.pid; \ + exit $$EXIT + # ── Cleanup ─────────────────────────────────────────────────────────────────── clean: diff --git a/README.md b/README.md index 659a7ed6..b0320365 100644 --- a/README.md +++ b/README.md @@ -311,7 +311,7 @@ flowchart LR | **App Configuration** | `/{account}-appconfig/` | Key-values, labels, feature flags, snapshots (async provisioning), revisions, locks, ETags; pagination (`@nextLink`), `$select`, `tags` filtering, `Accept-Datetime` time-travel, `Sync-Token` | | **Cosmos DB (NoSQL)** | `/{account}-cosmos/` | Databases, containers, documents CRUD + full SQL queries: always-on, no Docker. PATCH; transactional batch. | | **Cosmos DB NoSQL (embedded)** | `/{account}-cosmos-nosql/` | Same embedded SQL engine as above, exposed as a named engine endpoint. Opt-in with `FLOCI_AZ_SERVICES_COSMOS_ENGINES_NOSQL_ENABLED=true`; no Docker required. | -| **Key Vault** | `/{account}-keyvault/` | Secrets CRUD, versioning, soft-delete, properties update | +| **Key Vault** | `/{account}-keyvault/` | Secrets CRUD, versioning, soft-delete, properties update; Keys CRUD, backup/restore, rotation, RSA/EC/oct crypto (encrypt/decrypt/sign/verify/wrap/unwrap), `/rng`; Managed HSM (`/{account}-managedhsm/`) | | **Event Hubs** | AMQP `:5672` / Kafka `:9093` | AMQP 1.0 (Artemis sidecar), Kafka-compatible (Redpanda, opt-in) | | **Service Bus** | `/{account}-servicebus/` + AMQP `:5673` | Queues, topics, subscriptions (created dynamically); AMQP 1.0 data plane via Artemis sidecar, or mocked (management plane only) | | **Azure SQL Database** | ARM path + `/{account}-sql/` | Servers, databases, firewall rules; ARM-only by default, optional managed SQL Server 2025 containers | diff --git a/compatibility-tests/compat-azcli/test/keyvault.bats b/compatibility-tests/compat-azcli/test/keyvault.bats index bb467f67..51dccd9b 100644 --- a/compatibility-tests/compat-azcli/test/keyvault.bats +++ b/compatibility-tests/compat-azcli/test/keyvault.bats @@ -33,3 +33,92 @@ setup() { assert_success assert_equal "$(echo "$output" | jq -r '.value')" "hello-from-az-cli" } + +@test "az keyvault: key create/show/list/delete round-trip (data-plane)" { + run az keyvault key create --vault-name "$KV_NAME" -n "$KEY_NAME" \ + --kty RSA --size 2048 -o none + if [ "$status" -ne 0 ]; then + skip "key vault data-plane not reachable: $output" + fi + + run az_json keyvault key show --vault-name "$KV_NAME" -n "$KEY_NAME" + assert_success + assert_equal "$(echo "$output" | jq -r '.key.kty')" "RSA" + + run az_json keyvault key list --vault-name "$KV_NAME" + assert_success + [ -n "$(echo "$output" | jq -r '.[].kid')" ] + + run az keyvault key delete --vault-name "$KV_NAME" -n "$KEY_NAME" -o none + assert_success +} + +@test "az keyvault: key set-attributes + encrypt/decrypt + rotation-policy (data-plane)" { + run az keyvault key create --vault-name "$KV_NAME" -n "$CRYPTO_KEY_NAME" \ + --kty RSA --size 2048 -o none + if [ "$status" -ne 0 ]; then + skip "key vault data-plane not reachable: $output" + fi + + # set-attributes: disable the key and verify the flag is reflected on show. + run az keyvault key set-attributes --vault-name "$KV_NAME" -n "$CRYPTO_KEY_NAME" \ + --enabled false -o none + assert_success + run az_json keyvault key show --vault-name "$KV_NAME" -n "$CRYPTO_KEY_NAME" + assert_success + assert_equal "$(echo "$output" | jq -r '.attributes.enabled')" "false" + + # re-enable so encrypt/decrypt can proceed. + run az keyvault key set-attributes --vault-name "$KV_NAME" -n "$CRYPTO_KEY_NAME" \ + --enabled true -o none + assert_success + + # encrypt/decrypt round-trip. + local plaintext="hello-az-cli" + local b64 + b64=$(printf '%s' "$plaintext" | base64) + run az_json keyvault key encrypt --vault-name "$KV_NAME" -n "$CRYPTO_KEY_NAME" \ + --algorithm RSA-OAEP-256 --value "$b64" + assert_success + local ciphertext + ciphertext=$(echo "$output" | jq -r '.result') + [ -n "$ciphertext" ] + + run az_json keyvault key decrypt --vault-name "$KV_NAME" -n "$CRYPTO_KEY_NAME" \ + --algorithm RSA-OAEP-256 --value "$ciphertext" + assert_success + assert_equal "$(echo "$output" | jq -r '.result' | base64 -d)" "$plaintext" + + # rotation-policy show + update round-trip. + run az_json keyvault key rotation-policy show --vault-name "$KV_NAME" -n "$CRYPTO_KEY_NAME" + assert_success + run az_json keyvault key rotation-policy update --vault-name "$KV_NAME" -n "$CRYPTO_KEY_NAME" \ + --value '{"lifetimeActions":[{"trigger":{"timeAfterCreate":"P30D"},"action":{"type":"Rotate"}}],"attributes":{"expiryTime":"P90D"}}' + assert_success + + run az keyvault key delete --vault-name "$KV_NAME" -n "$CRYPTO_KEY_NAME" -o none + assert_success +} + +@test "az keyvault: key soft-delete lifecycle (data-plane)" { + run az keyvault key create --vault-name "$KV_NAME" -n "$LIFECYCLE_KEY_NAME" \ + --kty RSA --size 2048 -o none + if [ "$status" -ne 0 ]; then + skip "key vault data-plane not reachable: $output" + fi + + run az keyvault key delete --vault-name "$KV_NAME" -n "$LIFECYCLE_KEY_NAME" -o none + assert_success + + run az_json keyvault key list-deleted --vault-name "$KV_NAME" + assert_success + [ -n "$(echo "$output" | jq -r '.[].kid')" ] + + run az keyvault key recover --vault-name "$KV_NAME" -n "$LIFECYCLE_KEY_NAME" -o none + assert_success + + run az keyvault key delete --vault-name "$KV_NAME" -n "$LIFECYCLE_KEY_NAME" -o none + assert_success + run az keyvault key purge --vault-name "$KV_NAME" -n "$LIFECYCLE_KEY_NAME" -o none + assert_success +} diff --git a/compatibility-tests/compat-azcli/test/test_helper/common-setup.bash b/compatibility-tests/compat-azcli/test/test_helper/common-setup.bash index 814a72b6..5a1cfdd7 100644 --- a/compatibility-tests/compat-azcli/test/test_helper/common-setup.bash +++ b/compatibility-tests/compat-azcli/test/test_helper/common-setup.bash @@ -23,6 +23,9 @@ export CONTAINER_NAME="floci-test-container" export BLOB_NAME="hello.txt" export KV_NAME="floci-test-kv" export SECRET_NAME="floci-test-secret" +export KEY_NAME="floci-test-key" +export CRYPTO_KEY_NAME="floci-test-crypto-key" +export LIFECYCLE_KEY_NAME="floci-test-lifecycle-key" export VNET_NAME="floci-test-vnet" export SUBNET_NAME="floci-test-subnet" export NIC_NAME="floci-test-nic" diff --git a/compatibility-tests/sdk-test-java/pom.xml b/compatibility-tests/sdk-test-java/pom.xml index 8533cff1..4442f835 100644 --- a/compatibility-tests/sdk-test-java/pom.xml +++ b/compatibility-tests/sdk-test-java/pom.xml @@ -62,6 +62,10 @@ com.azure azure-security-keyvault-secrets + + com.azure + azure-security-keyvault-keys + com.azure azure-identity diff --git a/compatibility-tests/sdk-test-java/src/test/java/io/floci/az/compat/EmulatorConfig.java b/compatibility-tests/sdk-test-java/src/test/java/io/floci/az/compat/EmulatorConfig.java index cc68408b..cbe29c04 100644 --- a/compatibility-tests/sdk-test-java/src/test/java/io/floci/az/compat/EmulatorConfig.java +++ b/compatibility-tests/sdk-test-java/src/test/java/io/floci/az/compat/EmulatorConfig.java @@ -14,6 +14,10 @@ import com.azure.messaging.servicebus.ServiceBusClientBuilder; import com.azure.security.keyvault.secrets.SecretClient; import com.azure.security.keyvault.secrets.SecretClientBuilder; +import com.azure.security.keyvault.keys.KeyClient; +import com.azure.security.keyvault.keys.KeyClientBuilder; +import com.azure.security.keyvault.keys.cryptography.CryptographyClient; +import com.azure.security.keyvault.keys.cryptography.CryptographyClientBuilder; import com.fasterxml.jackson.core.type.TypeReference; import com.fasterxml.jackson.databind.ObjectMapper; import org.apache.qpid.jms.JmsConnectionFactory; @@ -184,6 +188,35 @@ public Mono process(HttpPipelineCallContext context, HttpPipelineN } } + /** + * The Azure SDKs require a {@code https://{account}.vault.azure.net/keys/...} key identifier for + * cryptography clients and derive the request URL from its host, dropping any path component. The + * emulator is reached over path-based routing instead, so this policy rewrites each crypto request + * from the host-based URL to {@code http://{endpoint}/{account}-keyvault/keys/...}. + */ + static final class KeyVaultDataPlanePolicy implements HttpPipelinePolicy { + private final String accountPath; + + KeyVaultDataPlanePolicy(String accountPath) { + this.accountPath = accountPath; + } + + @Override + public Mono process(HttpPipelineCallContext context, HttpPipelineNextPolicy next) { + URL url = context.getHttpRequest().getUrl(); + URI endpoint = URI.create(BASE); + try { + String path = "/" + accountPath + url.getPath(); + String query = url.getQuery(); + context.getHttpRequest().setUrl(new URL("http", endpoint.getHost(), endpoint.getPort(), + path + (query != null ? "?" + query : ""))); + } catch (MalformedURLException e) { + return Mono.error(e); + } + return next.process(); + } + } + // ── Event Hubs / AMQP ──────────────────────────────────────────────────── static final String EVENTHUB_HOST = @@ -318,6 +351,53 @@ static SecretClient buildKeyVaultClient() { .buildClient(); } + static KeyClient buildKeyClient() { + String vaultUrl = keyVaultUrl(); + return new KeyClientBuilder() + .vaultUrl(vaultUrl) + .credential(req -> Mono.just(new AccessToken("fake-token", OffsetDateTime.now().plusHours(1)))) + .addPolicy(new ForceHttpPolicy()) + .disableChallengeResourceVerification() + .buildClient(); + } + + /** Managed HSM data-plane client: same handler, but routed via the {@code -managedhsm} suffix. */ + static KeyClient buildManagedHsmKeyClient() { + String vaultUrl = BASE.replace("http://", "https://") + "/" + ACCOUNT + "-managedhsm"; + return new KeyClientBuilder() + .vaultUrl(vaultUrl) + .credential(req -> Mono.just(new AccessToken("fake-token", OffsetDateTime.now().plusHours(1)))) + .addPolicy(new ForceHttpPolicy()) + .disableChallengeResourceVerification() + .buildClient(); + } + + static CryptographyClient buildCryptographyClient(String keyName, String keyVersion) { + return buildCryptographyClient(keyName, keyVersion, false); + } + + /** Managed HSM cryptography client: same handler, routed via the {@code -managedhsm} suffix. */ + static CryptographyClient buildManagedHsmCryptographyClient(String keyName, String keyVersion) { + return buildCryptographyClient(keyName, keyVersion, true); + } + + private static CryptographyClient buildCryptographyClient(String keyName, String keyVersion, boolean hsm) { + String host = hsm + ? "https://" + ACCOUNT + ".managedhsm.azure.net" + : "https://" + ACCOUNT + ".vault.azure.net"; + String keyIdentifier = host + "/keys/" + keyName + "/" + keyVersion; + return new CryptographyClientBuilder() + .keyIdentifier(keyIdentifier) + .credential(req -> Mono.just(new AccessToken("fake-token", OffsetDateTime.now().plusHours(1)))) + .addPolicy(new KeyVaultDataPlanePolicy(ACCOUNT + (hsm ? "-managedhsm" : "-keyvault"))) + .disableChallengeResourceVerification() + .buildClient(); + } + + private static String keyVaultUrl() { + return BASE.replace("http://", "https://") + "/" + ACCOUNT + "-keyvault"; + } + // ── Service Bus ─────────────────────────────────────────────────────────── static final String SERVICEBUS_HOST = diff --git a/compatibility-tests/sdk-test-java/src/test/java/io/floci/az/compat/KeyVaultKeysCompatibilityTest.java b/compatibility-tests/sdk-test-java/src/test/java/io/floci/az/compat/KeyVaultKeysCompatibilityTest.java new file mode 100644 index 00000000..c53eeb13 --- /dev/null +++ b/compatibility-tests/sdk-test-java/src/test/java/io/floci/az/compat/KeyVaultKeysCompatibilityTest.java @@ -0,0 +1,309 @@ +package io.floci.az.compat; + +import com.azure.core.exception.ResourceNotFoundException; +import com.azure.core.util.Context; +import com.azure.security.keyvault.keys.KeyClient; +import com.azure.security.keyvault.keys.cryptography.CryptographyClient; +import com.azure.security.keyvault.keys.cryptography.models.DecryptParameters; +import com.azure.security.keyvault.keys.cryptography.models.EncryptParameters; +import com.azure.security.keyvault.keys.cryptography.models.EncryptionAlgorithm; +import com.azure.security.keyvault.keys.cryptography.models.KeyWrapAlgorithm; +import com.azure.security.keyvault.keys.cryptography.models.SignatureAlgorithm; +import com.azure.security.keyvault.keys.models.CreateEcKeyOptions; +import com.azure.security.keyvault.keys.models.CreateOctKeyOptions; +import com.azure.security.keyvault.keys.models.CreateRsaKeyOptions; +import com.azure.security.keyvault.keys.models.DeletedKey; +import com.azure.security.keyvault.keys.models.KeyCurveName; +import com.azure.security.keyvault.keys.models.KeyProperties; +import com.azure.security.keyvault.keys.models.KeyRotationPolicy; +import com.azure.security.keyvault.keys.models.KeyVaultKey; +import org.junit.jupiter.api.*; + +import java.util.List; +import java.util.UUID; + +import static org.junit.jupiter.api.Assertions.*; + +@TestInstance(TestInstance.Lifecycle.PER_CLASS) +@TestMethodOrder(MethodOrderer.DisplayName.class) +@DisplayName("Key Vault Keys + Crypto Compatibility") +class KeyVaultKeysCompatibilityTest { + + private KeyClient client; + + @BeforeAll + void setup() { + EmulatorConfig.assumeEmulatorRunning(); + client = EmulatorConfig.buildKeyClient(); + } + + private String name(String prefix) { + return prefix + UUID.randomUUID().toString().replace("-", "").substring(0, 8); + } + + private CryptographyClient crypto(KeyVaultKey key) { + return EmulatorConfig.buildCryptographyClient(key.getName(), key.getProperties().getVersion()); + } + + // ------------------------------------------------------------------------- + // Key CRUD + lifecycle + // ------------------------------------------------------------------------- + + @Test + @DisplayName("create/get RSA key") + void createAndGetRsaKey() { + String n = name("rsa"); + KeyVaultKey created = client.createRsaKey(new CreateRsaKeyOptions(n).setKeySize(2048)); + assertEquals(n, created.getName()); + assertNotNull(created.getKey().getN()); + + KeyVaultKey fetched = client.getKey(n); + assertEquals(n, fetched.getName()); + client.beginDeleteKey(n); + } + + @Test + @DisplayName("create/get EC key") + void createAndGetEcKey() { + String n = name("ec"); + KeyVaultKey created = client.createEcKey(new CreateEcKeyOptions(n).setCurveName(KeyCurveName.P_256)); + assertNotNull(created.getKey().getX()); + assertNotNull(created.getKey().getY()); + + assertEquals(n, client.getKey(n).getName()); + client.beginDeleteKey(n); + } + + @Test + @DisplayName("create/get oct key") + void createAndGetOctKey() { + String n = name("oct"); + KeyVaultKey created = client.createOctKey(new CreateOctKeyOptions(n).setKeySize(256)); + // Symmetric key material is never released in a Key Vault response. + assertNull(created.getKey().getK()); + + assertEquals(n, client.getKey(n).getName()); + client.beginDeleteKey(n); + } + + @Test + @DisplayName("get nonexistent key throws 404") + void getNonexistentThrows() { + ResourceNotFoundException ex = assertThrows(ResourceNotFoundException.class, + () -> client.getKey("no-such-key-xyz-" + UUID.randomUUID())); + assertEquals(404, ex.getResponse().getStatusCode()); + } + + @Test + @DisplayName("list keys includes created ones") + void listKeys() { + String n = name("list"); + client.createRsaKey(new CreateRsaKeyOptions(n).setKeySize(2048)); + List names = client.listPropertiesOfKeys().stream().map(KeyProperties::getName).toList(); + assertTrue(names.contains(n)); + client.beginDeleteKey(n); + } + + @Test + @DisplayName("list key versions") + void listKeyVersions() { + String n = name("ver"); + KeyVaultKey v1 = client.createRsaKey(new CreateRsaKeyOptions(n).setKeySize(2048)); + KeyVaultKey v2 = client.createRsaKey(new CreateRsaKeyOptions(n).setKeySize(2048)); + assertNotEquals(v1.getProperties().getVersion(), v2.getProperties().getVersion()); + assertEquals(2, client.listPropertiesOfKeyVersions(n).stream().count()); + client.beginDeleteKey(n); + } + + @Test + @DisplayName("get specific version") + void getSpecificVersion() { + String n = name("sv"); + KeyVaultKey v1 = client.createRsaKey(new CreateRsaKeyOptions(n).setKeySize(2048)); + String version = v1.getProperties().getVersion(); + client.createRsaKey(new CreateRsaKeyOptions(n).setKeySize(2048)); + assertEquals(version, client.getKey(n, version).getProperties().getVersion()); + client.beginDeleteKey(n); + } + + @Test + @DisplayName("delete/recover/purge lifecycle") + void deleteRecoverPurge() { + String n = name("lc"); + client.createRsaKey(new CreateRsaKeyOptions(n).setKeySize(2048)); + client.beginDeleteKey(n).poll(); + assertNotNull(client.getDeletedKey(n).getName()); + + client.beginRecoverDeletedKey(n).poll(); + assertEquals(n, client.getKey(n).getName()); + + client.beginDeleteKey(n).poll(); + client.purgeDeletedKey(n); + List deleted = client.listDeletedKeys().stream().map(DeletedKey::getName).toList(); + assertFalse(deleted.contains(n)); + } + + @Test + @DisplayName("backup and restore key") + void backupRestoreKey() { + String n = name("bak"); + client.createRsaKey(new CreateRsaKeyOptions(n).setKeySize(2048)); + byte[] backup = client.backupKey(n); + assertNotNull(backup); + assertTrue(backup.length > 0); + + client.beginDeleteKey(n).poll(); + client.purgeDeletedKey(n); + + KeyVaultKey restored = client.restoreKeyBackup(backup); + assertEquals(n, restored.getName()); + client.beginDeleteKey(n); + } + + @Test + @DisplayName("rotate key yields a fresh version") + void rotateKey() { + String n = name("rot"); + KeyVaultKey created = client.createRsaKey(new CreateRsaKeyOptions(n).setKeySize(2048)); + String v1 = created.getProperties().getVersion(); + KeyVaultKey rotated = client.rotateKey(n); + assertNotEquals(v1, rotated.getProperties().getVersion()); + client.beginDeleteKey(n); + } + + @Test + @DisplayName("get and update key rotation policy") + void rotationPolicyGetAndUpdate() { + String n = name("rotpol"); + client.createRsaKey(new CreateRsaKeyOptions(n).setKeySize(2048)); + + KeyRotationPolicy policy = client.getKeyRotationPolicy(n); + assertNotNull(policy); + assertNotNull(policy.getLifetimeActions()); + + // Round-trip the default policy back through update (PUT) and verify it persists. + KeyRotationPolicy updated = client.updateKeyRotationPolicy(n, policy); + assertNotNull(updated); + client.beginDeleteKey(n); + } + + // ------------------------------------------------------------------------- + // Crypto operations + // ------------------------------------------------------------------------- + + @Test + @DisplayName("RSA-OAEP-256 encrypt/decrypt round-trip") + void rsaOaep256RoundTrip() { + String n = name("crypt-oaep"); + KeyVaultKey key = client.createRsaKey(new CreateRsaKeyOptions(n).setKeySize(2048)); + CryptographyClient c = crypto(key); + byte[] plaintext = "hello-rsa-oaep".getBytes(); + + byte[] ciphertext = c.encrypt(EncryptParameters.createRsaOaep256Parameters(plaintext), Context.NONE).getCipherText(); + byte[] decrypted = c.decrypt(DecryptParameters.createRsaOaep256Parameters(ciphertext), Context.NONE).getPlainText(); + assertArrayEquals(plaintext, decrypted); + client.beginDeleteKey(n); + } + + @Test + @DisplayName("RSA1_5 encrypt/decrypt round-trip") + void rsa15RoundTrip() { + String n = name("crypt-rsa15"); + KeyVaultKey key = client.createRsaKey(new CreateRsaKeyOptions(n).setKeySize(2048)); + CryptographyClient c = crypto(key); + byte[] plaintext = "hello-rsa15".getBytes(); + + byte[] ciphertext = c.encrypt(EncryptParameters.createRsa15Parameters(plaintext), Context.NONE).getCipherText(); + byte[] decrypted = c.decrypt(DecryptParameters.createRsa15Parameters(ciphertext), Context.NONE).getPlainText(); + assertArrayEquals(plaintext, decrypted); + client.beginDeleteKey(n); + } + + @Test + @DisplayName("A256GCM encrypt/decrypt round-trip") + void a256GcmRoundTrip() { + String n = name("crypt-gcm"); + KeyVaultKey key = client.createOctKey(new CreateOctKeyOptions(n).setKeySize(256)); + CryptographyClient c = crypto(key); + byte[] plaintext = "hello-aes-gcm".getBytes(); + + var result = c.encrypt(EncryptParameters.createA256GcmParameters(plaintext), Context.NONE); + byte[] decrypted = c.decrypt(DecryptParameters.createA256GcmParameters( + result.getCipherText(), result.getIv(), result.getAuthenticationTag()), Context.NONE).getPlainText(); + assertArrayEquals(plaintext, decrypted); + client.beginDeleteKey(n); + } + + @Test + @DisplayName("RS256 signData/verifyData round-trip") + void rs256SignVerify() { + String n = name("sig-rs256"); + KeyVaultKey key = client.createRsaKey(new CreateRsaKeyOptions(n).setKeySize(2048)); + CryptographyClient c = crypto(key); + byte[] data = "sign-me-rs256".getBytes(); + + byte[] signature = c.signData(SignatureAlgorithm.RS256, data).getSignature(); + assertTrue(c.verifyData(SignatureAlgorithm.RS256, data, signature).isValid()); + client.beginDeleteKey(n); + } + + @Test + @DisplayName("PS256 signData/verifyData round-trip") + void ps256SignVerify() { + String n = name("sig-ps256"); + KeyVaultKey key = client.createRsaKey(new CreateRsaKeyOptions(n).setKeySize(2048)); + CryptographyClient c = crypto(key); + byte[] data = "sign-me-ps256".getBytes(); + + byte[] signature = c.signData(SignatureAlgorithm.PS256, data).getSignature(); + assertTrue(c.verifyData(SignatureAlgorithm.PS256, data, signature).isValid()); + client.beginDeleteKey(n); + } + + @Test + @DisplayName("ES256/ES384/ES512 signData/verifyData round-trips") + void ecSignVerify() { + KeyCurveName[] curves = {KeyCurveName.P_256, KeyCurveName.P_384, KeyCurveName.P_521}; + SignatureAlgorithm[] algs = {SignatureAlgorithm.ES256, SignatureAlgorithm.ES384, SignatureAlgorithm.ES512}; + for (int i = 0; i < curves.length; i++) { + String n = name("sig-ec"); + KeyVaultKey key = client.createEcKey(new CreateEcKeyOptions(n).setCurveName(curves[i])); + CryptographyClient c = crypto(key); + byte[] data = "sign-me-ec".getBytes(); + byte[] signature = c.signData(algs[i], data).getSignature(); + assertTrue(c.verifyData(algs[i], data, signature).isValid()); + client.beginDeleteKey(n); + } + } + + @Test + @DisplayName("RSA-OAEP-256 wrapKey/unwrapKey round-trip") + void wrapUnwrapKey() { + String n = name("wrap"); + KeyVaultKey key = client.createRsaKey(new CreateRsaKeyOptions(n).setKeySize(2048)); + CryptographyClient c = crypto(key); + byte[] keyMaterial = "wrapped-key-material".getBytes(); + + byte[] wrapped = c.wrapKey(KeyWrapAlgorithm.RSA_OAEP_256, keyMaterial).getEncryptedKey(); + byte[] unwrapped = c.unwrapKey(KeyWrapAlgorithm.RSA_OAEP_256, wrapped).getKey(); + assertArrayEquals(keyMaterial, unwrapped); + client.beginDeleteKey(n); + } + + // ------------------------------------------------------------------------- + // Managed HSM + // ------------------------------------------------------------------------- + + @Test + @DisplayName("managed HSM basic key CRUD") + void managedHsmBasicCrud() { + KeyClient hsm = EmulatorConfig.buildManagedHsmKeyClient(); + String n = name("hsm"); + KeyVaultKey created = hsm.createRsaKey(new CreateRsaKeyOptions(n).setKeySize(2048)); + assertTrue(created.getKey().getId().contains("managedhsm")); + + assertEquals(n, hsm.getKey(n).getName()); + assertTrue(hsm.listPropertiesOfKeys().stream().anyMatch(p -> p.getName().equals(n))); + hsm.beginDeleteKey(n); + } +} diff --git a/compatibility-tests/sdk-test-node/package.json b/compatibility-tests/sdk-test-node/package.json index b44ccbb9..b6730f99 100644 --- a/compatibility-tests/sdk-test-node/package.json +++ b/compatibility-tests/sdk-test-node/package.json @@ -12,6 +12,7 @@ "@azure/data-tables": "^13.3.0", "@azure/identity": "^4.5.0", "@azure/keyvault-secrets": "^4.8.0", + "@azure/keyvault-keys": "^4.9.0", "@azure/msal-node": "^5.4.2", "@azure/service-bus": "7.9.5", "@azure/storage-blob": "^12.23.0", diff --git a/compatibility-tests/sdk-test-node/tests/keyvault-keys.test.ts b/compatibility-tests/sdk-test-node/tests/keyvault-keys.test.ts new file mode 100644 index 00000000..ce09ca99 --- /dev/null +++ b/compatibility-tests/sdk-test-node/tests/keyvault-keys.test.ts @@ -0,0 +1,257 @@ +import { KeyClient, CryptographyClient } from "@azure/keyvault-keys"; +import { TokenCredential } from "@azure/core-auth"; +import { PipelinePolicy } from "@azure/core-rest-pipeline"; +import { ACCOUNT } from "./config"; +import * as crypto from "crypto"; + +const BASE = process.env.FLOCI_AZ_ENDPOINT ?? "http://localhost:4577"; +// SDK requires https:// vault URL; ForceHttpPolicy rewrites it back before sending +const VAULT_URL = BASE.replace("http://", "https://") + `/${ACCOUNT}-keyvault`; + +const MHSM_URL = BASE.replace("http://", "https://") + `/${ACCOUNT}-managedhsm`; + +// The cryptography client parses a key id with parseKeyVaultKeyIdentifier, which +// only accepts the real Azure shape https://.vault.azure.net/keys//. +// We build a host-based kid and rewrite the request back onto the emulator's +// path-based route: https://.vault.azure.net/... becomes +// {endpoint}/devstoreaccount1-keyvault/.... +const VAULT_HOST = "devstoreaccount1.vault.azure.net"; + +const fakeCredential: TokenCredential = { + getToken: async () => ({ + token: "fake-token-for-local-emulator", + expiresOnTimestamp: Date.now() + 3_600_000, + }), +}; + +const forceHttpPolicy: PipelinePolicy = { + name: "ForceHttpPolicy", + sendRequest(request, next) { + request.url = request.url.replace(/^https:\/\//, "http://"); + request.allowInsecureConnection = true; + return next(request); + }, +}; + +function keyClient(url: string = VAULT_URL): KeyClient { + return new KeyClient(url, fakeCredential, { + disableChallengeResourceVerification: true, + additionalPolicies: [{ policy: forceHttpPolicy, position: "perCall" }], + }); +} + +const cryptoHttpPolicy: PipelinePolicy = { + name: "ForceEmulatorRoutePolicy", + sendRequest(request, next) { + request.url = request.url.replace( + `https://${VAULT_HOST}/`, + `${BASE}/${ACCOUNT}-keyvault/`, + ); + request.allowInsecureConnection = true; + return next(request); + }, +}; + +// CRITICAL: build the crypto kid ourselves — the emulator returns ids with host +// devstoreaccount1.vault.azure.net which does NOT route back to the emulator. +function cryptoClient(keyName: string, keyVersion: string): CryptographyClient { + const kid = `https://${VAULT_HOST}/keys/${keyName}/${keyVersion}`; + return new CryptographyClient(kid, fakeCredential, { + disableChallengeResourceVerification: true, + additionalPolicies: [{ policy: cryptoHttpPolicy, position: "perCall" }], + }); +} + +function uid(prefix: string): string { + return `${prefix}-${Math.random().toString(36).substring(2, 10)}`; +} + +const client = keyClient(); +const hsm = keyClient(MHSM_URL); + +// --- Key CRUD + lifecycle --- + +test("create and get RSA key", async () => { + const n = uid("rsa"); + const key = await client.createRsaKey(n, { keySize: 2048 }); + expect(key.name).toBe(n); + expect(key.key?.n).toBeTruthy(); + expect((await client.getKey(n)).name).toBe(n); + await client.beginDeleteKey(n); +}); + +test("create and get EC key", async () => { + const n = uid("ec"); + const key = await client.createEcKey(n, { curve: "P-256" }); + expect(key.key?.x).toBeTruthy(); + expect(key.key?.y).toBeTruthy(); + expect((await client.getKey(n)).name).toBe(n); + await client.beginDeleteKey(n); +}); + +test("create and get oct key", async () => { + const n = uid("oct"); + const key = await client.createOctKey(n, { keySize: 256 }); + // Symmetric key material is never released in a Key Vault response. + expect(key.key?.k).toBeUndefined(); + expect((await client.getKey(n)).name).toBe(n); + await client.beginDeleteKey(n); +}); + +test("get nonexistent key throws", async () => { + await expect(client.getKey("no-such-" + uid("x"))).rejects.toThrow(); +}); + +test("list keys includes created one", async () => { + const n = uid("list"); + await client.createRsaKey(n, { keySize: 2048 }); + const names: string[] = []; + for await (const p of client.listPropertiesOfKeys()) names.push(p.name); + expect(names).toContain(n); + await client.beginDeleteKey(n); +}); + +test("list key versions", async () => { + const n = uid("ver"); + await client.createRsaKey(n, { keySize: 2048 }); + await client.createRsaKey(n, { keySize: 2048 }); + const versions: string[] = []; + for await (const p of client.listPropertiesOfKeyVersions(n)) versions.push(p.version!); + expect(versions.length).toBeGreaterThanOrEqual(2); + await client.beginDeleteKey(n); +}); + +test("delete recover purge lifecycle", async () => { + const n = uid("lc"); + await client.createRsaKey(n, { keySize: 2048 }); + + await client.beginDeleteKey(n); + expect((await client.getDeletedKey(n)).name).toBe(n); + + await client.beginRecoverDeletedKey(n).then((p) => p.pollUntilDone()); + expect((await client.getKey(n)).name).toBe(n); + + await client.beginDeleteKey(n); + await client.purgeDeletedKey(n); + + const deleted: string[] = []; + for await (const p of client.listDeletedKeys()) deleted.push(p.name); + expect(deleted).not.toContain(n); +}); + +test("backup and restore key", async () => { + const n = uid("bak"); + await client.createRsaKey(n, { keySize: 2048 }); + + const backup = await client.backupKey(n); + expect(backup).toBeInstanceOf(Uint8Array); + expect(backup!.length).toBeGreaterThan(0); + + await client.beginDeleteKey(n); + await client.purgeDeletedKey(n); + + const restored = await client.restoreKeyBackup(backup!); + expect(restored.name).toBe(n); + await client.beginDeleteKey(n); +}); + +test("rotate key yields a fresh version", async () => { + const n = uid("rot"); + const v1 = (await client.createRsaKey(n, { keySize: 2048 })).properties.version!; + expect((await client.rotateKey(n)).properties.version).not.toBe(v1); + expect(await client.getKeyRotationPolicy(n)).toBeTruthy(); + await client.beginDeleteKey(n); +}); + +// --- Cryptography --- + +test("RSA-OAEP-256 encrypt/decrypt round-trip", async () => { + const n = uid("crypt-oaep"); + const key = await client.createRsaKey(n, { keySize: 2048 }); + const c = cryptoClient(n, key.properties.version!); + + const ct = (await c.encrypt("RSA-OAEP-256", Buffer.from("hello"))).result; + expect(Buffer.from((await c.decrypt("RSA-OAEP-256", ct)).result).toString()).toBe("hello"); + + await client.beginDeleteKey(n); +}); + +test("RSA-OAEP encrypt/decrypt round-trip", async () => { + const n = uid("crypt-oaep1"); + const key = await client.createRsaKey(n, { keySize: 2048 }); + const c = cryptoClient(n, key.properties.version!); + + const ct = (await c.encrypt("RSA-OAEP", Buffer.from("hello"))).result; + expect(Buffer.from((await c.decrypt("RSA-OAEP", ct)).result).toString()).toBe("hello"); + + await client.beginDeleteKey(n); +}); + +test("RSA-OAEP-256 wrap/unwrap round-trip", async () => { + const n = uid("wrap"); + const key = await client.createRsaKey(n, { keySize: 2048 }); + const c = cryptoClient(n, key.properties.version!); + + const material = Buffer.from("material"); + const wrapped = (await c.wrapKey("RSA-OAEP-256", material)).result; + expect(Buffer.from((await c.unwrapKey("RSA-OAEP-256", wrapped)).result).equals(material)).toBe(true); + + await client.beginDeleteKey(n); +}); + +test("RS256 sign/verify round-trip", async () => { + const n = uid("sig-rs256"); + const key = await client.createRsaKey(n, { keySize: 2048 }); + const c = cryptoClient(n, key.properties.version!); + + const digest = crypto.createHash("sha256").update("sign-me").digest(); + const sig = (await c.sign("RS256", digest)).result; + expect((await c.verify("RS256", digest, sig)).result).toBe(true); + + await client.beginDeleteKey(n); +}); + +test("ES256 sign/verify round-trip", async () => { + const n = uid("sig-es256"); + const key = await client.createEcKey(n, { curve: "P-256" }); + const c = cryptoClient(n, key.properties.version!); + + const digest = crypto.createHash("sha256").update("sign-me").digest(); + const sig = (await c.sign("ES256", digest)).result; + expect((await c.verify("ES256", digest, sig)).result).toBe(true); + + await client.beginDeleteKey(n); +}); + +test("A256GCM encrypt/decrypt round-trip", async () => { + const n = uid("crypt-gcm"); + const key = await client.createOctKey(n, { keySize: 256 }); + const c = cryptoClient(n, key.properties.version!); + + const res = await c.encrypt("A256GCM", Buffer.from("gcm-plaintext")); + const decrypted = await c.decrypt({ + algorithm: "A256GCM", + ciphertext: res.result, + iv: res.iv!, + authenticationTag: res.authenticationTag!, + }); + expect(Buffer.from(decrypted.result).toString()).toBe("gcm-plaintext"); + + await client.beginDeleteKey(n); +}); + +// --- Managed HSM --- + +test("managed HSM basic key CRUD", async () => { + const n = uid("hsm"); + const created = await hsm.createRsaKey(n, { keySize: 2048 }); + expect(created.name).toBe(n); + + expect((await hsm.getKey(n)).name).toBe(n); + + const names: string[] = []; + for await (const p of hsm.listPropertiesOfKeys()) names.push(p.name); + expect(names).toContain(n); + + await hsm.beginDeleteKey(n); +}); diff --git a/compatibility-tests/sdk-test-python/requirements.txt b/compatibility-tests/sdk-test-python/requirements.txt index a909341c..b5c29d6e 100644 --- a/compatibility-tests/sdk-test-python/requirements.txt +++ b/compatibility-tests/sdk-test-python/requirements.txt @@ -5,6 +5,7 @@ azure-data-tables==12.6.0 azure-cosmos==4.7.0 azure-appconfiguration==1.7.1 azure-keyvault-secrets>=4.7.0 +azure-keyvault-keys>=4.7.0 azure-identity>=1.15.0 azure-core>=1.30.0 redis==5.0.8 diff --git a/compatibility-tests/sdk-test-python/tests/keyvault/conftest.py b/compatibility-tests/sdk-test-python/tests/keyvault/conftest.py index ac1b4dd2..d5210043 100644 --- a/compatibility-tests/sdk-test-python/tests/keyvault/conftest.py +++ b/compatibility-tests/sdk-test-python/tests/keyvault/conftest.py @@ -5,6 +5,7 @@ from azure.core.credentials import AccessToken, TokenCredential from azure.core.pipeline.transport import RequestsTransport from azure.keyvault.secrets import SecretClient +from azure.keyvault.keys import KeyClient class FakeCredential(TokenCredential): @@ -35,3 +36,24 @@ def client(): transport=ForceHttpTransport(), verify_challenge_resource=False, ) + + +@pytest.fixture +def keys_client(): + return KeyClient( + vault_url=VAULT_URL, + credential=FakeCredential(), + transport=ForceHttpTransport(), + verify_challenge_resource=False, + ) + + +@pytest.fixture +def hsm_keys_client(): + hsm_url = re.sub(r"^http://", "https://", ENDPOINT) + "/devstoreaccount1-managedhsm" + return KeyClient( + vault_url=hsm_url, + credential=FakeCredential(), + transport=ForceHttpTransport(), + verify_challenge_resource=False, + ) diff --git a/compatibility-tests/sdk-test-python/tests/keyvault/test_keys.py b/compatibility-tests/sdk-test-python/tests/keyvault/test_keys.py new file mode 100644 index 00000000..e1c620ce --- /dev/null +++ b/compatibility-tests/sdk-test-python/tests/keyvault/test_keys.py @@ -0,0 +1,260 @@ +""" +Compatibility tests for Azure Key Vault Keys and Cryptography. +""" +import os +import time +import uuid +import hashlib + +import pytest +from azure.core.credentials import AccessToken, TokenCredential +from azure.core.exceptions import ResourceNotFoundError +from azure.core.pipeline.transport import RequestsTransport +from azure.keyvault.keys.crypto import ( + CryptographyClient, + EncryptionAlgorithm, + KeyWrapAlgorithm, + SignatureAlgorithm, +) + + +def unique(prefix="key"): + return f"{prefix}-{uuid.uuid4().hex[:8]}" + + +# The emulator returns key ids with host `devstoreaccount1.vault.azure.net`, which +# does not resolve back to the emulator. We build a host-based key id (the only shape +# parse_key_vault_id accepts) and have the transport rewrite the request back onto the +# emulator's path-based route: https://devstoreaccount1.vault.azure.net/keys/... becomes +# {endpoint}/devstoreaccount1-keyvault/keys/.... +_ENDPOINT = os.environ.get("FLOCI_AZ_ENDPOINT", "http://localhost:4577") +_VAULT_HOST = "devstoreaccount1.vault.azure.net" + + +class _FakeCredential(TokenCredential): + def get_token(self, *scopes, **kwargs): + return AccessToken("fake-token-for-local-emulator", int(time.time()) + 3600) + + +class _ForceHttpTransport(RequestsTransport): + def send(self, request, **kwargs): + request.url = request.url.replace( + f"https://{_VAULT_HOST}/", f"{_ENDPOINT}/devstoreaccount1-keyvault/", 1 + ) + return super().send(request, **kwargs) + + +def _crypto(key_name, key_version): + kid = f"https://{_VAULT_HOST}/keys/{key_name}/{key_version}" + return CryptographyClient( + kid, + credential=_FakeCredential(), + transport=_ForceHttpTransport(), + verify_challenge_resource=False, + ) + + +# --------------------------------------------------------------------------- +# Key CRUD + lifecycle +# --------------------------------------------------------------------------- + +def test_create_and_get_rsa_key(keys_client): + name = unique("rsa") + created = keys_client.create_rsa_key(name, size=2048) + assert created.name == name + assert created.key.kty == "RSA" + assert created.key.n is not None + + fetched = keys_client.get_key(name) + assert fetched.name == name + keys_client.begin_delete_key(name).result() + + +def test_create_and_get_ec_key(keys_client): + name = unique("ec") + created = keys_client.create_ec_key(name, curve="P-256") + assert created.key.x is not None + assert created.key.y is not None + + assert keys_client.get_key(name).name == name + keys_client.begin_delete_key(name).result() + + +def test_create_and_get_oct_key(keys_client): + name = unique("oct") + created = keys_client.create_oct_key(name, size=256) + # Symmetric key material is never released in a Key Vault response. + assert created.key.k is None + + assert keys_client.get_key(name).name == name + keys_client.begin_delete_key(name).result() + + +def test_get_nonexistent_key_raises(keys_client): + with pytest.raises(ResourceNotFoundError): + keys_client.get_key(unique("nonexistent")) + + +def test_list_properties_of_keys(keys_client): + name = unique("list") + keys_client.create_rsa_key(name, size=2048) + + listed = [p.name for p in keys_client.list_properties_of_keys()] + assert name in listed + keys_client.begin_delete_key(name).result() + + +def test_list_properties_of_key_versions(keys_client): + name = unique("ver") + keys_client.create_rsa_key(name, size=2048) + keys_client.create_rsa_key(name, size=2048) + + versions = list(keys_client.list_properties_of_key_versions(name)) + assert len(versions) >= 2 + keys_client.begin_delete_key(name).result() + + +def test_soft_delete_lifecycle(keys_client): + name = unique("lifecycle") + keys_client.create_rsa_key(name, size=2048) + + keys_client.begin_delete_key(name).result() + deleted = keys_client.get_deleted_key(name) + assert deleted.name == name + + keys_client.begin_recover_deleted_key(name).result() + assert keys_client.get_key(name).name == name + + keys_client.begin_delete_key(name).result() + keys_client.purge_deleted_key(name) + + with pytest.raises(ResourceNotFoundError): + keys_client.get_deleted_key(name) + + +def test_backup_and_restore_key(keys_client): + name = unique("backup") + keys_client.create_rsa_key(name, size=2048) + + backup = keys_client.backup_key(name) + assert backup is not None + assert len(backup) > 0 + + keys_client.begin_delete_key(name).result() + keys_client.purge_deleted_key(name) + + restored = keys_client.restore_key_backup(backup) + assert restored.name == name + keys_client.begin_delete_key(name).result() + + +def test_key_rotation(keys_client): + name = unique("rotate") + created = keys_client.create_rsa_key(name, size=2048) + original_version = created.properties.version + + policy = keys_client.get_key_rotation_policy(name) + assert policy is not None + assert hasattr(policy, "lifetime_actions") + + rotated = keys_client.rotate_key(name) + assert rotated.properties.version != original_version + keys_client.begin_delete_key(name).result() + + +# --------------------------------------------------------------------------- +# Cryptography +# --------------------------------------------------------------------------- + +def test_rsa_oaep_256_encrypt_decrypt(keys_client): + name = unique("crypt-oaep256") + key = keys_client.create_rsa_key(name, size=2048) + crypto = _crypto(name, key.properties.version) + + encrypted = crypto.encrypt(EncryptionAlgorithm.rsa_oaep_256, b"hello") + decrypted = crypto.decrypt(EncryptionAlgorithm.rsa_oaep_256, encrypted.ciphertext) + assert decrypted.plaintext == b"hello" + keys_client.begin_delete_key(name).result() + + +def test_rsa_oaep_encrypt_decrypt(keys_client): + name = unique("crypt-oaep") + key = keys_client.create_rsa_key(name, size=2048) + crypto = _crypto(name, key.properties.version) + + encrypted = crypto.encrypt(EncryptionAlgorithm.rsa_oaep, b"hello-oaep") + decrypted = crypto.decrypt(EncryptionAlgorithm.rsa_oaep, encrypted.ciphertext) + assert decrypted.plaintext == b"hello-oaep" + keys_client.begin_delete_key(name).result() + + +def test_rsa_oaep_256_wrap_unwrap(keys_client): + name = unique("wrap") + key = keys_client.create_rsa_key(name, size=2048) + crypto = _crypto(name, key.properties.version) + + wrapped = crypto.wrap_key(KeyWrapAlgorithm.rsa_oaep_256, b"wrapped-material") + unwrapped = crypto.unwrap_key(KeyWrapAlgorithm.rsa_oaep_256, wrapped.encrypted_key) + assert unwrapped.key == b"wrapped-material" + keys_client.begin_delete_key(name).result() + + +def test_rs256_sign_verify(keys_client): + name = unique("sig-rs256") + key = keys_client.create_rsa_key(name, size=2048) + crypto = _crypto(name, key.properties.version) + + digest = hashlib.sha256(b"sign-me").digest() + signature = crypto.sign(SignatureAlgorithm.rs256, digest).signature + verified = crypto.verify(SignatureAlgorithm.rs256, digest, signature) + assert verified.is_valid is True + keys_client.begin_delete_key(name).result() + + +def test_es256_sign_verify(keys_client): + name = unique("sig-es256") + key = keys_client.create_ec_key(name, curve="P-256") + crypto = _crypto(name, key.properties.version) + + digest = hashlib.sha256(b"sign-me").digest() + signature = crypto.sign(SignatureAlgorithm.es256, digest).signature + verified = crypto.verify(SignatureAlgorithm.es256, digest, signature) + assert verified.is_valid is True + keys_client.begin_delete_key(name).result() + + +def test_a256_gcm_encrypt_decrypt(keys_client): + name = unique("crypt-gcm") + key = keys_client.create_oct_key(name, size=256) + crypto = _crypto(name, key.properties.version) + + result = crypto.encrypt(EncryptionAlgorithm.a256_gcm, b"gcm-plaintext") + assert result.ciphertext is not None + assert result.iv is not None + assert result.tag is not None + + decrypted = crypto.decrypt( + EncryptionAlgorithm.a256_gcm, + result.ciphertext, + iv=result.iv, + authentication_tag=result.tag, + ) + assert decrypted.plaintext == b"gcm-plaintext" + keys_client.begin_delete_key(name).result() + + +# --------------------------------------------------------------------------- +# Managed HSM +# --------------------------------------------------------------------------- + +def test_managed_hsm_basic_crud(hsm_keys_client): + name = unique("hsm") + created = hsm_keys_client.create_rsa_key(name, size=2048) + assert created.name == name + + assert hsm_keys_client.get_key(name).name == name + + listed = [p.name for p in hsm_keys_client.list_properties_of_keys()] + assert name in listed + + hsm_keys_client.begin_delete_key(name).result() diff --git a/docs/services/index.md b/docs/services/index.md index a3b3a746..1cdb23e8 100644 --- a/docs/services/index.md +++ b/docs/services/index.md @@ -12,7 +12,7 @@ Floci-AZ provides emulation for several core Azure services. | **App Configuration** | `/{account}-appconfig/` | ✅ Key-values, labels, feature flags, snapshots, revisions, locks, pagination, `$select`, tags filtering, `Accept-Datetime`, `Sync-Token` | | **Cosmos DB (SQL API)** | `/{account}-cosmos/` | ✅ Databases, containers, documents CRUD, SQL queries, partition keys | | **Cosmos DB multi-API** | _(engine sidecars)_ | ✅ MongoDB, PostgreSQL, Cassandra, Gremlin, Table, NoSQL (opt-in Docker engines) | -| **Key Vault** | `/{account}-keyvault/` | ✅ Secrets CRUD, versioning, soft-delete, properties update | +| **Key Vault** | `/{account}-keyvault/` | ✅ Secrets & keys CRUD, versioning, soft-delete, properties update, backup/restore, rotation, RSA/EC/oct crypto, `/rng`, Managed HSM | | **Event Hubs** | AMQP `:5672` / Kafka `:9093` | ✅ AMQP 1.0 (Artemis), Kafka-compatible (Redpanda, opt-in) | | **Service Bus** | `/{account}-servicebus/` + AMQP `:5673` | ✅ Queues, topics, subscriptions (dynamic); AMQP 1.0 via Artemis sidecar or mocked | | **Azure SQL Database** | ARM path + `/{account}-sql/` | ✅ Servers, databases, firewall rules; ARM-only by default, managed SQL Server opt-in | diff --git a/docs/services/key-vault.md b/docs/services/key-vault.md index 9b5bec68..1114cbd0 100644 --- a/docs/services/key-vault.md +++ b/docs/services/key-vault.md @@ -1,9 +1,11 @@ # Key Vault -Compatible with `azure-keyvault-secrets` SDKs (Python, Java, JavaScript, .NET). +Compatible with the `azure-keyvault-secrets` and `azure-keyvault-keys` SDKs (Python, Java, JavaScript, .NET). ## Features +### Secrets + - **Secrets CRUD** — set, get, delete, list secrets - **Versioning** — each `set_secret` creates a new immutable version; latest pointer tracks the most recent - **Soft-delete lifecycle** — delete moves a secret to the deleted namespace; recover or purge it @@ -15,6 +17,18 @@ Compatible with `azure-keyvault-secrets` SDKs (Python, Java, JavaScript, .NET). - **Backup** — backup a secret (base64-encoded blob) - **32-char hex version IDs** — matches Azure's version ID format +### Keys & Cryptography + +- **Keys CRUD** — create/import RSA (`RSA`, `RSA-HSM`), EC (`EC`, `EC-HSM`, P-256/P-384/P-521), and + oct (`oct`, `oct-HSM`) keys; get/list/list-versions; PATCH attributes +- **Soft-delete lifecycle** — delete → `deletedkeys` namespace → recover or purge +- **Backup/restore** — `POST /keys/{name}/backup` and `POST /keys/restore` (see deviations below) +- **Rotation** — `POST /keys/{name}/rotate` and rotation-policy management (`/keys/{name}/rotationpolicy`) +- **Crypto ops** — `encrypt`/`decrypt` (RSA-OAEP, RSA-OAEP-256, RSA1_5; AES-GCM A128/A192/A256), + `sign`/`verify` (RS256/384/512, PS256/384/512, ES256/384/512), `wrapkey`/`unwrapkey` +- **`/rng`** — random bytes for client-side key material (see deviations) +- **Managed HSM** — the same data plane served under the `/{account}-managedhsm/` suffix + ## Endpoint ``` @@ -242,6 +256,61 @@ All endpoints sit under `/{accountName}-keyvault/` with an `api-version` query p | `DELETE` | `/deletedsecrets/{name}` | Purge (permanently delete) | | `POST` | `/deletedsecrets/{name}/recover` | Recover a deleted secret | +### Keys + +| Method | Path | Description | +|---|---|---| +| `POST` | `/keys/{name}/create` | Create a key (`kty`, `key_size`, `curve`, `key_ops`) | +| `PUT` | `/keys/{name}` | Import a key (`key` JWK) | +| `GET` | `/keys` | List keys | +| `GET` | `/keys/{name}` | Get latest version | +| `GET` | `/keys/{name}/{version}` | Get a specific version | +| `PATCH` | `/keys/{name}/{version}` | Update key attributes (`enabled`, `nbf`, `exp`, `tags`) | +| `DELETE` | `/keys/{name}` | Soft-delete a key | +| `GET` | `/keys/{name}/versions` | List all versions | +| `POST` | `/keys/{name}/backup` | Backup a key | +| `POST` | `/keys/restore` | Restore a key from a backup blob | +| `POST` | `/keys/{name}/rotate` | Rotate (regenerate) a key | +| `GET`/`PUT` | `/keys/{name}/rotationpolicy` | Get/put the key rotation policy | + +### Deleted Keys + +| Method | Path | Description | +|---|---|---| +| `GET` | `/deletedkeys` | List deleted keys | +| `GET` | `/deletedkeys/{name}` | Get a deleted key | +| `DELETE` | `/deletedkeys/{name}` | Purge (permanently delete) | +| `POST` | `/deletedkeys/{name}/recover` | Recover a deleted key | + +### Crypto Operations + +| Method | Path | Description | +|---|---|---| +| `POST` | `/keys/{name}[/{version}]/encrypt` | Encrypt (`RSA-OAEP`, `RSA-OAEP-256`, `RSA1_5`, `A128GCM`, `A192GCM`, `A256GCM`) | +| `POST` | `/keys/{name}[/{version}]/decrypt` | Decrypt | +| `POST` | `/keys/{name}[/{version}]/sign` | Sign (`RS256/384/512`, `PS256/384/512`, `ES256/384/512`) | +| `POST` | `/keys/{name}[/{version}]/verify` | Verify a signature | +| `POST` | `/keys/{name}[/{version}]/wrapkey` | Wrap a key (`RSA-OAEP-256` for RSA keys) | +| `POST` | `/keys/{name}[/{version}]/unwrapkey` | Unwrap a key | +| `POST` | `/rng` | Random bytes (`{"count": 1..128}`) | + +--- + +## Intentional deviations + +These are deliberate differences from real Azure Key Vault. They are stable, documented behavior — not bugs: + +- **Backup blobs are unencrypted plaintext.** Real Azure returns HSM-encrypted opaque blobs that can only be + restored into the same vault. floci-az emits a readable JSON snapshot (JWK + metadata) so backups are + portable and inspectable. Do not treat backup blobs as secrets. +- **`/rng` caps at 128 bytes per request.** This matches real Azure (which caps at 128). An earlier draft of the + plan stated 1024; that was incorrect. +- **Key material is stored in the clear in the storage backend** (memory/persistent), not HSM-encrypted. +- **`key_size` may appear in the returned public JWK** even though it is not part of Azure's `JsonWebKey` + schema. Azure SDKs tolerate unknown fields; kept for convenience. +- **`nbf`/`exp` must be numeric Unix epoch timestamps.** Whitespace is tolerated and normalized; any other + form is rejected with `400 BadParameter` at create/import/PATCH/restore time. + --- ## Storage Configuration diff --git a/scripts/native-crypto-smoke.sh b/scripts/native-crypto-smoke.sh new file mode 100755 index 00000000..67e09e42 --- /dev/null +++ b/scripts/native-crypto-smoke.sh @@ -0,0 +1,67 @@ +#!/usr/bin/env bash +set -euo pipefail + +BASE="http://localhost:4577/devstoreaccount1-keyvault" +# Array so curl receives each header as a single argument (unquoted $AUTH word-splits +# "Bearer x" into separate URL args). +AUTH=(-H "Authorization: Bearer x" -H "Content-Type: application/json") + +b64url() { base64 | tr '+/' '-_' | tr -d '=\n'; } +digest_sha256() { openssl dgst -sha256 -binary | b64url; } + +echo "==> Running native crypto smoke checks against $BASE..." + +# The emulator resolves version-less crypto paths to the latest key version, so we +# address operations by key name — the returned kid is a *.vault.azure.net URL that +# does not route back to the emulator. + +# 1. RSA-OAEP-256 +curl -sf "${AUTH[@]}" -X POST "$BASE/keys/smoke-rsa/create?api-version=7.4" \ + -d '{"kty":"RSA","key_size":2048}' >/dev/null +pt=$(echo -n "hello" | b64url) +ct=$(curl -sf "${AUTH[@]}" -X POST "$BASE/keys/smoke-rsa/encrypt?api-version=7.4" \ + -d "{\"alg\":\"RSA-OAEP-256\",\"value\":\"$pt\"}" | jq -r '.value') +dec=$(curl -sf "${AUTH[@]}" -X POST "$BASE/keys/smoke-rsa/decrypt?api-version=7.4" \ + -d "{\"alg\":\"RSA-OAEP-256\",\"value\":\"$ct\"}" | jq -r '.value') +[ "$dec" = "$pt" ] || { echo "FAIL: RSA-OAEP-256 round-trip"; exit 1; } +echo " ✓ RSA-OAEP-256 encrypt/decrypt passed" + +# 2. RS256 +dg=$(echo -n "smoke-msg" | digest_sha256) +sig=$(curl -sf "${AUTH[@]}" -X POST "$BASE/keys/smoke-rsa/sign?api-version=7.4" \ + -d "{\"alg\":\"RS256\",\"value\":\"$dg\"}" | jq -r '.value') +ok=$(curl -sf "${AUTH[@]}" -X POST "$BASE/keys/smoke-rsa/verify?api-version=7.4" \ + -d "{\"alg\":\"RS256\",\"digest\":\"$dg\",\"value\":\"$sig\"}" | jq -r '.value') +[ "$ok" = "true" ] || { echo "FAIL: RS256 verify"; exit 1; } +echo " ✓ RS256 sign/verify passed" + +# 3. PS256 +sig=$(curl -sf "${AUTH[@]}" -X POST "$BASE/keys/smoke-rsa/sign?api-version=7.4" \ + -d "{\"alg\":\"PS256\",\"value\":\"$dg\"}" | jq -r '.value') +ok=$(curl -sf "${AUTH[@]}" -X POST "$BASE/keys/smoke-rsa/verify?api-version=7.4" \ + -d "{\"alg\":\"PS256\",\"digest\":\"$dg\",\"value\":\"$sig\"}" | jq -r '.value') +[ "$ok" = "true" ] || { echo "FAIL: PS256 verify"; exit 1; } +echo " ✓ PS256 sign/verify passed" + +# 4. ES256 +curl -sf "${AUTH[@]}" -X POST "$BASE/keys/smoke-ec/create?api-version=7.4" \ + -d '{"kty":"EC","crv":"P-256"}' >/dev/null +sig=$(curl -sf "${AUTH[@]}" -X POST "$BASE/keys/smoke-ec/sign?api-version=7.4" \ + -d "{\"alg\":\"ES256\",\"value\":\"$dg\"}" | jq -r '.value') +ok=$(curl -sf "${AUTH[@]}" -X POST "$BASE/keys/smoke-ec/verify?api-version=7.4" \ + -d "{\"alg\":\"ES256\",\"digest\":\"$dg\",\"value\":\"$sig\"}" | jq -r '.value') +[ "$ok" = "true" ] || { echo "FAIL: ES256 verify"; exit 1; } +echo " ✓ ES256 sign/verify passed" + +# 5. AES-GCM +curl -sf "${AUTH[@]}" -X POST "$BASE/keys/smoke-oct/create?api-version=7.4" \ + -d '{"kty":"oct","key_size":256}' >/dev/null +resp=$(curl -sf "${AUTH[@]}" -X POST "$BASE/keys/smoke-oct/encrypt?api-version=7.4" \ + -d "{\"alg\":\"A256GCM\",\"value\":\"$pt\"}") +ct=$(echo "$resp" | jq -r '.value'); iv=$(echo "$resp" | jq -r '.iv'); tag=$(echo "$resp" | jq -r '.tag') +dec=$(curl -sf "${AUTH[@]}" -X POST "$BASE/keys/smoke-oct/decrypt?api-version=7.4" \ + -d "{\"alg\":\"A256GCM\",\"value\":\"$ct\",\"iv\":\"$iv\",\"tag\":\"$tag\"}" | jq -r '.value') +[ "$dec" = "$pt" ] || { echo "FAIL: AES-GCM round-trip"; exit 1; } +echo " ✓ AES-GCM (A256GCM) encrypt/decrypt passed" + +echo "ALL 5 NATIVE CRYPTO SMOKE CHECKS PASSED" diff --git a/src/main/java/io/floci/az/core/AzureRoutingFilter.java b/src/main/java/io/floci/az/core/AzureRoutingFilter.java index 08be8ec0..6d0ed7d7 100644 --- a/src/main/java/io/floci/az/core/AzureRoutingFilter.java +++ b/src/main/java/io/floci/az/core/AzureRoutingFilter.java @@ -137,7 +137,7 @@ String resourcePath() { /** Well-known Key Vault data-plane path prefixes, as sent to the ARM base URL by azurerm v3. */ private static final Set KEY_VAULT_COLLECTIONS = Set.of( - "secrets", "certificates", "keys", "deletedsecrets", "deletedcertificates", "deletedkeys" + "secrets", "certificates", "keys", "deletedsecrets", "deletedcertificates", "deletedkeys", "rng" ); /** @@ -352,6 +352,13 @@ public Uni filter(ContainerRequestContext requestContext, @Context Htt private Response doFilter(ContainerRequestContext requestContext, String decodedPath, String rawPath, HttpHeaders headers, String capturedHost, String remoteAddress) { + // Never trust a client-supplied account-suffix header: only dispatchByAccountSuffix may set it. + // Header names are case-insensitive on the wire, so match keys case-insensitively. + for (String header : new ArrayList<>(requestContext.getHeaders().keySet())) { + if ("x-floci-account-suffix".equalsIgnoreCase(header)) { + requestContext.getHeaders().remove(header); + } + } String path = trimLeadingSlash(decodedPath); String encodedPath = trimLeadingSlash(rawPath); @@ -650,6 +657,7 @@ private Response dispatchByAccountSuffix(RoutingContext ctx) { if (route != null) { serviceType = route.serviceType(); accountName = stripSuffix(accountName, route); + ctx.requestContext().getHeaders().putSingle("x-floci-account-suffix", route.suffix()); } else { serviceType = resolveStorageServiceType(ctx.requestContext(), resourcePath); } diff --git a/src/main/java/io/floci/az/core/tls/CertificateGenerator.java b/src/main/java/io/floci/az/core/tls/CertificateGenerator.java index 5d141784..644479ee 100644 --- a/src/main/java/io/floci/az/core/tls/CertificateGenerator.java +++ b/src/main/java/io/floci/az/core/tls/CertificateGenerator.java @@ -4,11 +4,14 @@ import org.bouncycastle.asn1.DEROctetString; import org.bouncycastle.asn1.x500.X500Name; import org.bouncycastle.asn1.x509.BasicConstraints; +import org.bouncycastle.asn1.x509.ExtendedKeyUsage; import org.bouncycastle.asn1.x509.Extension; import org.bouncycastle.asn1.x509.GeneralName; import org.bouncycastle.asn1.x509.GeneralNames; +import org.bouncycastle.asn1.x509.KeyPurposeId; import org.bouncycastle.asn1.x509.KeyUsage; import org.bouncycastle.cert.jcajce.JcaX509CertificateConverter; +import org.bouncycastle.cert.jcajce.JcaX509ExtensionUtils; import org.bouncycastle.cert.jcajce.JcaX509v3CertificateBuilder; import org.bouncycastle.jce.provider.BouncyCastleProvider; import org.bouncycastle.openssl.jcajce.JcaPEMWriter; @@ -39,14 +42,17 @@ public class CertificateGenerator { "^\\[?([0-9a-fA-F:]+)]?$|^(\\d{1,3}\\.\\d{1,3}\\.\\d{1,3}\\.\\d{1,3})$" ); - public record GeneratedCertificate(String certificatePem, String privateKeyPem) {} + /** + * @param chainPem leaf followed by CA, in PEM (Quarkus "chain file" format) + * @param caPem the CA certificate alone, in PEM (the client trust anchor) + * @param privateKeyPem the leaf's private key, in PEM + */ + public record GeneratedCertificate(String chainPem, String caPem, String privateKeyPem) {} /** - * Generates a genuinely self-signed certificate (issuer == subject, marked as a CA) suitable - * for use as a trust anchor: a client that adds this certificate to its CA store can - * verify a TLS connection that presents it. Used for floci-az's own HTTPS server certificate - * so that containers making HTTPS calls back to floci-az can trust it once the certificate is - * installed in their CA bundle. + * Generates a CA → leaf chain: a self-signed CA (CN=floci-az-ca) signs a non-CA leaf + * (CN=floci-az) carrying the SANs and serverAuth EKU. The leaf must not be a CA, because strict + * validators (e.g. rustls) reject a CA-flagged certificate in the end-entity position. */ public GeneratedCertificate generateCertificate(List sans) { try { @@ -55,39 +61,67 @@ public GeneratedCertificate generateCertificate(List sans) { // certificate builder and signer which use BC's internal implementations. KeyPairGenerator keyGen = KeyPairGenerator.getInstance("RSA"); keyGen.initialize(2048, SECURE_RANDOM); - KeyPair keyPair = keyGen.generateKeyPair(); - Instant now = Instant.now(); - X500Name name = new X500Name("CN=floci-az"); - BigInteger serial = new BigInteger(128, SECURE_RANDOM); + JcaX509ExtensionUtils extensionUtils = new JcaX509ExtensionUtils(); + Instant now = Instant.now(); - var certBuilder = new JcaX509v3CertificateBuilder( - name, serial, + // --- Self-signed CA certificate (the trust anchor) --- + KeyPair caKeyPair = keyGen.generateKeyPair(); + X500Name caName = new X500Name("CN=floci-az-ca"); + BigInteger caSerial = new BigInteger(128, SECURE_RANDOM); + + var caBuilder = new JcaX509v3CertificateBuilder( + caName, caSerial, + Date.from(now), Date.from(now.plus(3650, ChronoUnit.DAYS)), + caName, caKeyPair.getPublic()); + caBuilder.addExtension(Extension.basicConstraints, true, new BasicConstraints(true)); + caBuilder.addExtension(Extension.keyUsage, true, + new KeyUsage(KeyUsage.keyCertSign | KeyUsage.cRLSign)); + caBuilder.addExtension(Extension.subjectKeyIdentifier, false, + extensionUtils.createSubjectKeyIdentifier(caKeyPair.getPublic())); + + ContentSigner caSigner = new JcaContentSignerBuilder("SHA512WithRSA") + .build(caKeyPair.getPrivate()); + X509Certificate caCert = new JcaX509CertificateConverter() + .getCertificate(caBuilder.build(caSigner)); + + // --- End-entity leaf certificate signed by the CA --- + KeyPair leafKeyPair = keyGen.generateKeyPair(); + X500Name leafName = new X500Name("CN=floci-az"); + BigInteger leafSerial = new BigInteger(128, SECURE_RANDOM); + + var leafBuilder = new JcaX509v3CertificateBuilder( + caName, leafSerial, Date.from(now), Date.from(now.plus(365, ChronoUnit.DAYS)), - name, keyPair.getPublic()); + leafName, leafKeyPair.getPublic()); GeneralName[] sanEntries = sans.stream() .map(CertificateGenerator::toGeneralName) .filter(gn -> gn != null) .toArray(GeneralName[]::new); - - certBuilder.addExtension(Extension.subjectAlternativeName, false, new GeneralNames(sanEntries)); - // A trust anchor must be a CA so clients accept it as one, and it needs keyCertSign - // so it can act as its own issuer. - certBuilder.addExtension(Extension.basicConstraints, true, new BasicConstraints(true)); - certBuilder.addExtension(Extension.keyUsage, true, - new KeyUsage(KeyUsage.digitalSignature | KeyUsage.keyEncipherment | KeyUsage.keyCertSign)); - - ContentSigner signer = new JcaContentSignerBuilder("SHA512WithRSA") - .build(keyPair.getPrivate()); - - X509Certificate cert = new JcaX509CertificateConverter() - .getCertificate(certBuilder.build(signer)); - - return new GeneratedCertificate(toPem(cert), toPem(keyPair.getPrivate())); + leafBuilder.addExtension(Extension.subjectAlternativeName, false, new GeneralNames(sanEntries)); + leafBuilder.addExtension(Extension.basicConstraints, true, new BasicConstraints(false)); + leafBuilder.addExtension(Extension.keyUsage, true, + new KeyUsage(KeyUsage.digitalSignature | KeyUsage.keyEncipherment)); + leafBuilder.addExtension(Extension.extendedKeyUsage, false, + new ExtendedKeyUsage(KeyPurposeId.id_kp_serverAuth)); + leafBuilder.addExtension(Extension.authorityKeyIdentifier, false, + extensionUtils.createAuthorityKeyIdentifier(caKeyPair.getPublic())); + leafBuilder.addExtension(Extension.subjectKeyIdentifier, false, + extensionUtils.createSubjectKeyIdentifier(leafKeyPair.getPublic())); + + ContentSigner leafSigner = new JcaContentSignerBuilder("SHA512WithRSA") + .build(caKeyPair.getPrivate()); + X509Certificate leafCert = new JcaX509CertificateConverter() + .getCertificate(leafBuilder.build(leafSigner)); + + return new GeneratedCertificate( + toPem(leafCert) + toPem(caCert), + toPem(caCert), + toPem(leafKeyPair.getPrivate())); } catch (Exception e) { - throw new IllegalStateException("Failed to generate self-signed TLS certificate", e); + throw new IllegalStateException("Failed to generate TLS certificate chain", e); } } diff --git a/src/main/java/io/floci/az/core/tls/TlsConfigSource.java b/src/main/java/io/floci/az/core/tls/TlsConfigSource.java index f485c0a8..cc6409e3 100644 --- a/src/main/java/io/floci/az/core/tls/TlsConfigSource.java +++ b/src/main/java/io/floci/az/core/tls/TlsConfigSource.java @@ -42,6 +42,7 @@ public class TlsConfigSource implements ConfigSource { private static final Logger LOG = Logger.getLogger(TlsConfigSource.class); private static final String SELF_SIGNED_CERT_NAME = "floci-az-selfsigned.crt"; + private static final String SELF_SIGNED_CA_NAME = "floci-az-selfsigned-ca.crt"; private static final String SELF_SIGNED_KEY_NAME = "floci-az-selfsigned.key"; private static final String SELF_SIGNED_METADATA_NAME = "floci-az-selfsigned.metadata.json"; private static final String TLS_DIR = "tls"; @@ -74,30 +75,37 @@ public TlsConfigSource() { String selfSigned = resolveProperty("floci-az.tls.self-signed", "true"); String persistentPath = resolveProperty("floci-az.storage.persistent-path", "./data"); + // The PEM served at GET /_floci/tls-cert (the client trust anchor): the generated CA in + // self-signed mode, or the user-provided cert file itself. + String anchorPath = null; + if (!certPath.isBlank() && !keyPath.isBlank()) { validateFileExists(certPath, "TLS certificate"); validateFileExists(keyPath, "TLS private key"); + anchorPath = certPath; LOG.infov("TLS: using user-provided certificate: {0}", certPath); } else if ("true".equalsIgnoreCase(selfSigned)) { Path tlsDir = Path.of(persistentPath, TLS_DIR); Path certFile = tlsDir.resolve(SELF_SIGNED_CERT_NAME); + Path caFile = tlsDir.resolve(SELF_SIGNED_CA_NAME); Path keyFile = tlsDir.resolve(SELF_SIGNED_KEY_NAME); List customHostnames = extractCustomHostnames(); List allSans = buildSanList(customHostnames); - if (Files.exists(certFile) && Files.exists(keyFile)) { + if (Files.exists(certFile) && Files.exists(keyFile) && Files.exists(caFile)) { if (hostnameConfigChanged(tlsDir, allSans)) { - generateSelfSignedCert(tlsDir, certFile, keyFile, allSans); + generateSelfSignedCert(tlsDir, certFile, caFile, keyFile, allSans); } else { LOG.infov("TLS: reusing existing self-signed certificate: {0}", certFile); } } else { - generateSelfSignedCert(tlsDir, certFile, keyFile, allSans); + generateSelfSignedCert(tlsDir, certFile, caFile, keyFile, allSans); } certPath = certFile.toAbsolutePath().toString(); keyPath = keyFile.toAbsolutePath().toString(); + anchorPath = caFile.toAbsolutePath().toString(); } else { throw new IllegalStateException( "TLS enabled but no certificate provided and self-signed generation disabled. " @@ -114,7 +122,7 @@ public TlsConfigSource() { properties.put("quarkus.http.ssl-port", String.valueOf(HTTPS_INTERNAL_PORT)); try { - currentCertPem = Files.readString(Path.of(certPath)); + currentCertPem = Files.readString(Path.of(anchorPath)); } catch (IOException e) { LOG.warnv("TLS: could not read cert PEM for /_floci/tls-cert endpoint: {0}", e.getMessage()); } @@ -187,7 +195,7 @@ private static List buildSanList(List customHostnames) { // (not in a container). all.addAll(List.of("localhost", "127.0.0.1", "0.0.0.0", "*.localhost", "localhost.floci-az.io", "*.localhost.floci-az.io", - "*.vault.azure.net", "host.docker.internal")); + "*.vault.azure.net", "*.managedhsm.azure.net", "host.docker.internal")); all.addAll(customHostnames); return all; } @@ -232,7 +240,7 @@ private static void ensureBouncyCastleRegistered() { } } - private void generateSelfSignedCert(Path tlsDir, Path certFile, Path keyFile, List sans) { + private void generateSelfSignedCert(Path tlsDir, Path certFile, Path caFile, Path keyFile, List sans) { try { Files.createDirectories(tlsDir); ensureBouncyCastleRegistered(); @@ -240,7 +248,9 @@ private void generateSelfSignedCert(Path tlsDir, Path certFile, Path keyFile, Li CertificateGenerator.GeneratedCertificate generated = new CertificateGenerator().generateCertificate(sans); - Files.writeString(certFile, generated.certificatePem()); + // certFile holds the leaf+CA chain (same filename as before); caFile holds the CA alone. + Files.writeString(certFile, generated.chainPem()); + Files.writeString(caFile, generated.caPem()); Files.writeString(keyFile, generated.privateKeyPem()); LOG.infov("TLS: generated self-signed certificate: {0}", certFile); diff --git a/src/main/java/io/floci/az/services/arm/ArmHandler.java b/src/main/java/io/floci/az/services/arm/ArmHandler.java index 95781b96..ceb34faa 100644 --- a/src/main/java/io/floci/az/services/arm/ArmHandler.java +++ b/src/main/java/io/floci/az/services/arm/ArmHandler.java @@ -57,6 +57,7 @@ public class ArmHandler implements AzureServiceHandler { private final Map> resourceGroups = new ConcurrentHashMap<>(); private final Map> storageAccounts = new ConcurrentHashMap<>(); private final Map> keyVaults = new ConcurrentHashMap<>(); + private final Map> managedHsms = new ConcurrentHashMap<>(); private final Map> webApps = new ConcurrentHashMap<>(); private final EmulatorConfig config; private final BlobServiceHandler blobHandler; @@ -201,6 +202,16 @@ private Response dispatch(AzureRequest req) { return Response.ok(Map.of("value", vaults)).build(); } + // ── Cross-subscription managed HSM listing ───────────────────────────── + if (path.matches("subscriptions/[^/?]+/providers/Microsoft\\.KeyVault/managedHSMs([?].*)?")) { + String sub = extractSub(path); + List> hsms = managedHsms.values().stream() + .filter(h -> sub.equals(h.get("_sub"))) + .map(ArmHandler::stripInternal) + .toList(); + return Response.ok(Map.of("value", hsms)).build(); + } + // ── Provider registration check (skip_provider_registration=true still calls this) ── // Only matches /subscriptions/{sub}/providers or /subscriptions/{sub}/providers/{namespace} // (no resource type segment), so the more-specific handlers above take precedence. @@ -217,6 +228,10 @@ private Response dispatch(AzureRequest req) { .filter(v -> sub.equals(v.get("_sub"))) .map(ArmHandler::stripInternal) .toList()); + resources.addAll(managedHsms.values().stream() + .filter(h -> sub.equals(h.get("_sub"))) + .map(ArmHandler::stripInternal) + .toList()); resources.addAll(apiManagementHandler.listSubscriptionServices(sub)); return Response.ok(Map.of("value", resources)).build(); } @@ -265,6 +280,10 @@ private Response handleResourceGroupBranch(AzureRequest req, String path, String .filter(v -> sub.equals(v.get("_sub")) && rg.equals(v.get("_rg"))) .map(ArmHandler::stripInternal) .forEach(resources::add); + managedHsms.values().stream() + .filter(h -> sub.equals(h.get("_sub")) && rg.equals(h.get("_rg"))) + .map(ArmHandler::stripInternal) + .forEach(resources::add); webApps.values().stream() .filter(v -> sub.equals(v.get("_sub")) && rg.equals(v.get("_rg"))) .map(ArmHandler::stripInternal) @@ -641,6 +660,26 @@ private Response handleKeyVault(AzureRequest req, String path, String method, St }; } + // Managed HSM list + if (path.matches(".*/providers/Microsoft\\.KeyVault/managedHSMs([?].*)?")) { + List> hsms = managedHsms.values().stream() + .filter(h -> sub.equals(h.get("_sub")) && rg.equals(h.get("_rg"))) + .map(ArmHandler::stripInternal) + .toList(); + return Response.ok(Map.of("value", hsms)).build(); + } + + // Single managed HSM + if (path.contains("/managedHSMs/")) { + String hsmName = extractResourceName(path, "managedHSMs"); + return switch (method) { + case "PUT" -> createOrUpdateManagedHsm(req, sub, rg, hsmName); + case "GET" -> getManagedHsm(sub, rg, hsmName); + case "DELETE" -> { managedHsms.remove(hsmKey(sub, rg, hsmName)); yield Response.ok().build(); } + default -> Response.status(405).build(); + }; + } + return armNotFound(path); } @@ -684,6 +723,56 @@ private Response getKeyVault(String sub, String rg, String vaultName) { return Response.ok(stripInternal(resource)).build(); } + private Response createOrUpdateManagedHsm(AzureRequest req, String sub, String rg, String hsmName) { + Map body = parseBody(req); + String location = bodyString(body, "location", "eastus"); + String hsmUri = "https://" + hsmName + ".managedhsm.azure.net/"; + Map bodyProps = body.containsKey("properties") + ? cast(body.get("properties")) : Map.of(); + String tenantId = bodyString(bodyProps, "tenantId", TENANT_ID); + @SuppressWarnings("unchecked") + List initialAdminObjectIds = bodyProps.get("initialAdminObjectIds") instanceof List l + ? (List) l : List.of(); + + Map properties = new LinkedHashMap<>(); + properties.put("tenantId", tenantId); + properties.put("initialAdminObjectIds", initialAdminObjectIds); + properties.put("enableSoftDelete", true); + properties.put("softDeleteRetentionInDays", 90); + properties.put("enablePurgeProtection", true); + properties.put("hsmUri", hsmUri); + properties.put("provisioningState", "Succeeded"); + + Map sku = body.get("sku") instanceof Map m + ? cast(m) : Map.of(); + Map skuOut = Map.of( + "family", bodyString(sku, "family", "B"), + "name", bodyString(sku, "name", "Standard_B1")); + + Map resource = new LinkedHashMap<>(); + resource.put("_sub", sub); + resource.put("_rg", rg); + resource.put("id", "/subscriptions/" + sub + "/resourceGroups/" + rg + + "/providers/Microsoft.KeyVault/managedHSMs/" + hsmName); + resource.put("name", hsmName); + resource.put("type", "Microsoft.KeyVault/managedHSMs"); + resource.put("location", location); + resource.put("sku", skuOut); + resource.put("properties", properties); + + managedHsms.put(hsmKey(sub, rg, hsmName), resource); + LOG.infof("ARM: created managed HSM %s (hsmUri=%s)", hsmName, hsmUri); + return Response.ok(stripInternal(resource)).build(); + } + + private Response getManagedHsm(String sub, String rg, String hsmName) { + Map resource = managedHsms.get(hsmKey(sub, rg, hsmName)); + if (resource == null) { + return armNotFound("managedHSMs/" + hsmName); + } + return Response.ok(stripInternal(resource)).build(); + } + // ── Resource Groups ─────────────────────────────────────────────────────── private Response createOrUpdateResourceGroup(String sub, String rg, AzureRequest req) { @@ -747,6 +836,7 @@ private static String extractAfter(String path, String marker) { private static String rgKey(String sub, String rg) { return sub + "/" + rg; } private static String saKey(String sub, String rg, String name) { return sub + "/" + rg + "/" + name; } private static String kvKey(String sub, String rg, String name) { return sub + "/" + rg + "/kv/" + name; } + private static String hsmKey(String sub, String rg, String name) { return sub + "/" + rg + "/hsm/" + name; } private static String webAppKey(String sub, String rg, String name) { return sub + "/" + rg + "/web/" + name; } private static Map cast(Object o) { diff --git a/src/main/java/io/floci/az/services/keyvault/KeyVaultCrypto.java b/src/main/java/io/floci/az/services/keyvault/KeyVaultCrypto.java new file mode 100644 index 00000000..9f41d8c5 --- /dev/null +++ b/src/main/java/io/floci/az/services/keyvault/KeyVaultCrypto.java @@ -0,0 +1,984 @@ +package io.floci.az.services.keyvault; + +import javax.crypto.Cipher; +import javax.crypto.SecretKey; +import javax.crypto.spec.GCMParameterSpec; +import javax.crypto.spec.OAEPParameterSpec; +import javax.crypto.spec.PSource; +import javax.crypto.spec.SecretKeySpec; +import java.math.BigInteger; +import java.security.AlgorithmParameters; +import java.security.KeyFactory; +import java.security.KeyPair; +import java.security.KeyPairGenerator; +import java.security.MessageDigest; +import java.security.PrivateKey; +import java.security.PublicKey; +import java.security.SecureRandom; +import java.security.Signature; +import java.security.interfaces.ECPrivateKey; +import java.security.interfaces.ECPublicKey; +import java.security.interfaces.RSAKey; +import java.security.interfaces.RSAPrivateCrtKey; +import java.security.interfaces.RSAPublicKey; +import java.security.spec.ECGenParameterSpec; +import java.security.spec.ECParameterSpec; +import java.security.spec.ECPoint; +import java.security.spec.ECPrivateKeySpec; +import java.security.spec.ECPublicKeySpec; +import java.security.spec.MGF1ParameterSpec; +import java.security.spec.RSAPrivateCrtKeySpec; +import java.security.spec.RSAPrivateKeySpec; +import java.security.spec.RSAPublicKeySpec; +import java.util.ArrayList; +import java.util.Arrays; +import java.util.Base64; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; + +/** + * Pure-JDK crypto engine for the Key Vault keys data plane: static helpers using only built-in + * JDK providers (SunJCE/SunEC/SunRsaSign), avoiding BouncyCastle for GraalVM native image. + * The {@code value}/{@code digest} given to {@link #sign}/{@link #verify} is the client's + * pre-hashed digest, never the raw message, so signing must never double-hash. + */ +final class KeyVaultCrypto { + + private static final SecureRandom SECURE_RANDOM = new SecureRandom(); + private static final Base64.Encoder URL = Base64.getUrlEncoder().withoutPadding(); + private static final Base64.Decoder URL_DECODER = Base64.getUrlDecoder(); + + // SunJCE's OAEPWithSHA-256AndMGF1Padding uses MGF1 SHA-1 unless an explicit OAEPParameterSpec + // pins both the digest and the MGF. Clients (Python/Node/az) use MGF1 SHA-256, so we must too. + private static final OAEPParameterSpec OAEP_SHA1 = + new OAEPParameterSpec("SHA-1", "MGF1", MGF1ParameterSpec.SHA1, PSource.PSpecified.DEFAULT); + private static final OAEPParameterSpec OAEP_SHA256 = + new OAEPParameterSpec("SHA-256", "MGF1", MGF1ParameterSpec.SHA256, PSource.PSpecified.DEFAULT); + + private KeyVaultCrypto() { + } + + /** Raised on unsupported algorithm / wrong key type / malformed input. Maps to a 400. */ + static final class CryptoException extends RuntimeException { + CryptoException(String message) { + super(message); + } + + CryptoException(String message, Throwable cause) { + super(message, cause); + } + } + + // ── base64url helpers ────────────────────────────────────────────────────── + + static String b64Url(byte[] data) { + return URL.encodeToString(data); + } + + static byte[] b64UrlDecode(String value) { + if (value == null || value.isEmpty()) { + return new byte[0]; + } + try { + String padded = value; + int rem = value.length() % 4; + if (rem != 0) { + padded = value + "=".repeat(4 - rem); + } + return URL_DECODER.decode(padded); + } catch (IllegalArgumentException e) { + throw new CryptoException("Invalid base64url value", e); + } + } + + /** Big-endian, sign-bit-free encoding required by JWK {@code n}/{@code e}/{@code d}/... */ + static byte[] toUnsignedBytes(BigInteger value) { + byte[] bytes = value.toByteArray(); + if (bytes.length > 1 && bytes[0] == 0) { + return Arrays.copyOfRange(bytes, 1, bytes.length); + } + return bytes; + } + + private static BigInteger bigint(Object value) { + if (!(value instanceof String s)) { + throw new CryptoException("JWK field must be a base64url string"); + } + return new BigInteger(1, b64UrlDecode(s)); + } + + private static byte[] fixedLengthUnsigned(BigInteger value, int len) { + byte[] raw = toUnsignedBytes(value); + if (raw.length > len) { + throw new CryptoException("JWK field does not fit curve field size"); + } + if (raw.length == len) { + return raw; + } + byte[] out = new byte[len]; + System.arraycopy(raw, 0, out, len - raw.length, raw.length); + return out; + } + + // ── Key type helpers ─────────────────────────────────────────────────────── + + /** Strips the {@code -HSM} suffix, e.g. {@code RSA-HSM → RSA}. */ + static String baseKty(String kty) { + if (kty == null) { + return ""; + } + int dash = kty.indexOf('-'); + return dash < 0 ? kty : kty.substring(0, dash); + } + + static boolean isRsaKty(String kty) { + return "RSA".equals(baseKty(kty)); + } + + static boolean isEcKty(String kty) { + return "EC".equals(baseKty(kty)); + } + + static boolean isOctKty(String kty) { + return "oct".equals(baseKty(kty)); + } + + static List defaultKeyOps(String baseKty) { + return switch (baseKty) { + case "RSA" -> List.of("encrypt", "decrypt", "sign", "verify", "wrapKey", "unwrapKey"); + case "EC" -> List.of("sign", "verify"); + case "oct" -> List.of("encrypt", "decrypt", "wrapKey", "unwrapKey"); + default -> List.of(); + }; + } + + private static String curveToStd(String crv) { + if (crv == null) { + throw new CryptoException("EC keys require a curve"); + } + return switch (crv) { + case "P-256" -> "secp256r1"; + case "P-384" -> "secp384r1"; + case "P-521" -> "secp521r1"; + default -> throw new CryptoException("Unsupported curve: " + crv); + }; + } + + // ── JWK generation / import / public strip ───────────────────────────────── + + static Map generateJwk(String kty, int keySize, String crv, List keyOps) { + String base = baseKty(kty); + Map jwk = new LinkedHashMap<>(); + jwk.put("kty", kty); + jwk.put("key_ops", keyOps); + try { + switch (base) { + case "RSA" -> { + int size = keySize > 0 ? keySize : 2048; + KeyPairGenerator gen = KeyPairGenerator.getInstance("RSA"); + gen.initialize(size, SECURE_RANDOM); + KeyPair pair = gen.generateKeyPair(); + RSAPublicKey pub = (RSAPublicKey) pair.getPublic(); + RSAPrivateCrtKey priv = (RSAPrivateCrtKey) pair.getPrivate(); + jwk.put("n", b64Url(toUnsignedBytes(pub.getModulus()))); + jwk.put("e", b64Url(toUnsignedBytes(pub.getPublicExponent()))); + jwk.put("d", b64Url(toUnsignedBytes(priv.getPrivateExponent()))); + jwk.put("p", b64Url(toUnsignedBytes(priv.getPrimeP()))); + jwk.put("q", b64Url(toUnsignedBytes(priv.getPrimeQ()))); + jwk.put("dp", b64Url(toUnsignedBytes(priv.getPrimeExponentP()))); + jwk.put("dq", b64Url(toUnsignedBytes(priv.getPrimeExponentQ()))); + jwk.put("qi", b64Url(toUnsignedBytes(priv.getCrtCoefficient()))); + jwk.put("keySize", size); + } + case "EC" -> { + String curve = curveToStd(crv); + KeyPairGenerator gen = KeyPairGenerator.getInstance("EC"); + gen.initialize(new ECGenParameterSpec(curve), SECURE_RANDOM); + KeyPair pair = gen.generateKeyPair(); + ECPublicKey pub = (ECPublicKey) pair.getPublic(); + ECPrivateKey priv = (ECPrivateKey) pair.getPrivate(); + int fieldLen = (pub.getParams().getCurve().getField().getFieldSize() + 7) / 8; + jwk.put("crv", crv != null ? crv : "P-256"); + jwk.put("x", b64Url(fixedLengthUnsigned(pub.getW().getAffineX(), fieldLen))); + jwk.put("y", b64Url(fixedLengthUnsigned(pub.getW().getAffineY(), fieldLen))); + jwk.put("d", b64Url(fixedLengthUnsigned(priv.getS(), fieldLen))); + } + case "oct" -> { + int bits = keySize > 0 ? keySize : 256; + if (bits != 128 && bits != 192 && bits != 256) { + throw new CryptoException("Unsupported oct key size: " + bits); + } + byte[] k = new byte[bits / 8]; + SECURE_RANDOM.nextBytes(k); + jwk.put("k", b64Url(k)); + jwk.put("keySize", bits); + } + default -> throw new CryptoException("Unsupported key type: " + kty); + } + } catch (CryptoException e) { + throw e; + } catch (Exception e) { + throw new CryptoException("Key generation failed for " + kty, e); + } + return jwk; + } + + /** + * Validates and normalizes a client-supplied JWK (import). Rebuilds the full key map from the + * supplied fields so it can be stored exactly like a generated key. + */ + static Map importJwk(Map supplied) { + Object ktyVal = supplied.get("kty"); + if (!(ktyVal instanceof String kty) || kty.isBlank()) { + throw new CryptoException("key type (kty) is required"); + } + String base = baseKty(kty); + Map jwk = new LinkedHashMap<>(); + jwk.put("kty", kty); + jwk.put("key_ops", supplied.getOrDefault("key_ops", defaultKeyOps(base))); + switch (base) { + case "RSA" -> { + requireStringField(supplied, "n"); + requireStringField(supplied, "e"); + requireStringIfPresent(supplied, "d"); + requireStringIfPresent(supplied, "p"); + requireStringIfPresent(supplied, "q"); + requireStringIfPresent(supplied, "dp"); + requireStringIfPresent(supplied, "dq"); + requireStringIfPresent(supplied, "qi"); + copyFields(jwk, supplied, "n", "e", "d", "p", "q", "dp", "dq", "qi"); + jwk.put("keySize", b64UrlDecode((String) supplied.get("n")).length * 8); + } + case "EC" -> { + requireStringField(supplied, "crv"); + requireStringField(supplied, "x"); + requireStringField(supplied, "y"); + requireStringIfPresent(supplied, "d"); + curveToStd((String) supplied.get("crv")); + copyFields(jwk, supplied, "crv", "x", "y", "d"); + } + case "oct" -> { + requireStringField(supplied, "k"); + copyFields(jwk, supplied, "k"); + jwk.put("keySize", b64UrlDecode((String) supplied.get("k")).length * 8); + } + default -> throw new CryptoException("Unsupported key type: " + kty); + } + return jwk; + } + + private static void requireStringField(Map supplied, String field) { + Object value = supplied.get(field); + if (!(value instanceof String)) { + throw new CryptoException("JWK field '" + field + "' is required and must be a string"); + } + } + + private static void requireStringIfPresent(Map supplied, String field) { + Object value = supplied.get(field); + if (value != null && !(value instanceof String)) { + throw new CryptoException("JWK field '" + field + "' must be a string"); + } + } + + private static void copyFields(Map target, Map source, String... fields) { + for (String field : fields) { + if (source.get(field) != null) { + target.put(field, source.get(field)); + } + } + } + + /** + * Strips private fields for the wire. Symmetric ({@code oct}) key material is never released in + * any response; the full JWK (including {@code k}) stays server-side for crypto operations. + */ + static Map publicJwk(Map full) { + String kty = (String) full.get("kty"); + String base = baseKty(kty); + Map pub = new LinkedHashMap<>(); + pub.put("kty", kty); + if (full.containsKey("key_ops")) { + pub.put("key_ops", full.get("key_ops")); + } + switch (base) { + case "RSA" -> { + pub.put("n", full.get("n")); + pub.put("e", full.get("e")); + if (full.containsKey("keySize")) { + pub.put("key_size", full.get("keySize")); + } + } + case "EC" -> { + pub.put("crv", full.get("crv")); + pub.put("x", full.get("x")); + pub.put("y", full.get("y")); + } + case "oct" -> { + if (full.containsKey("keySize")) { + pub.put("key_size", full.get("keySize")); + } + } + default -> { + } + } + return pub; + } + + // ── Key reconstruction ───────────────────────────────────────────────────── + + static PrivateKey reconstructRsaPrivate(Map jwk) { + try { + KeyFactory kf = KeyFactory.getInstance("RSA"); + BigInteger n = bigint(jwk.get("n")); + BigInteger e = bigint(jwk.get("e")); + if (jwk.get("d") != null && jwk.get("p") != null) { + return kf.generatePrivate(new RSAPrivateCrtKeySpec(n, e, bigint(jwk.get("d")), + bigint(jwk.get("p")), bigint(jwk.get("q")), bigint(jwk.get("dp")), + bigint(jwk.get("dq")), bigint(jwk.get("qi")))); + } + if (jwk.get("d") != null) { + return kf.generatePrivate(new RSAPrivateKeySpec(n, bigint(jwk.get("d")))); + } + throw new CryptoException("private key material missing"); + } catch (CryptoException e) { + throw e; + } catch (Exception e) { + throw new CryptoException("RSA private key reconstruction failed", e); + } + } + + static PublicKey reconstructRsaPublic(Map jwk) { + try { + KeyFactory kf = KeyFactory.getInstance("RSA"); + return kf.generatePublic(new RSAPublicKeySpec(bigint(jwk.get("n")), bigint(jwk.get("e")))); + } catch (Exception e) { + throw new CryptoException("RSA public key reconstruction failed", e); + } + } + + static PrivateKey reconstructEcPrivate(Map jwk) { + try { + KeyFactory kf = KeyFactory.getInstance("EC"); + return kf.generatePrivate(new ECPrivateKeySpec(bigint(jwk.get("d")), ecParams((String) jwk.get("crv")))); + } catch (Exception e) { + throw new CryptoException("EC private key reconstruction failed", e); + } + } + + static PublicKey reconstructEcPublic(Map jwk) { + try { + KeyFactory kf = KeyFactory.getInstance("EC"); + ECPoint w = new ECPoint(bigint(jwk.get("x")), bigint(jwk.get("y"))); + return kf.generatePublic(new ECPublicKeySpec(w, ecParams((String) jwk.get("crv")))); + } catch (Exception e) { + throw new CryptoException("EC public key reconstruction failed", e); + } + } + + static SecretKey reconstructSecretKey(Map jwk) { + return new SecretKeySpec(b64UrlDecode((String) jwk.get("k")), "AES"); + } + + private static ECParameterSpec ecParams(String crv) { + try { + AlgorithmParameters ap = AlgorithmParameters.getInstance("EC"); + ap.init(new ECGenParameterSpec(curveToStd(crv))); + return ap.getParameterSpec(ECParameterSpec.class); + } catch (Exception e) { + throw new CryptoException("Unsupported curve: " + crv, e); + } + } + + // ── Encrypt / decrypt / wrap / unwrap ────────────────────────────────────── + + record CipherResult(byte[] value, byte[] iv, byte[] tag) {} + + static CipherResult encrypt(Map jwk, String alg, byte[] value, byte[] iv, byte[] aad) { + String base = baseKty((String) jwk.get("kty")); + return switch (alg) { + case "RSA1_5" -> { + requireRsa(base, alg); + yield new CipherResult(rsaCrypt(jwk, "RSA/ECB/PKCS1Padding", Cipher.ENCRYPT_MODE, value, null), null, null); + } + case "RSA-OAEP" -> { + requireRsa(base, alg); + yield new CipherResult(rsaCrypt(jwk, "RSA/ECB/OAEPWithSHA-1AndMGF1Padding", Cipher.ENCRYPT_MODE, value, OAEP_SHA1), null, null); + } + case "RSA-OAEP-256" -> { + requireRsa(base, alg); + yield new CipherResult(rsaCrypt(jwk, "RSA/ECB/OAEPWithSHA-256AndMGF1Padding", Cipher.ENCRYPT_MODE, value, OAEP_SHA256), null, null); + } + case "A128GCM", "A192GCM", "A256GCM" -> { + requireOct(base, alg); + requireOctSize(jwk, alg); + yield gcmEncrypt(jwk, value, iv, aad); + } + default -> throw new CryptoException("Unsupported algorithm: " + alg); + }; + } + + static byte[] decrypt(Map jwk, String alg, byte[] value, byte[] iv, byte[] aad, byte[] tag) { + String base = baseKty((String) jwk.get("kty")); + return switch (alg) { + case "RSA1_5" -> { + requireRsa(base, alg); + yield rsaCrypt(jwk, "RSA/ECB/PKCS1Padding", Cipher.DECRYPT_MODE, value, null); + } + case "RSA-OAEP" -> { + requireRsa(base, alg); + yield rsaCrypt(jwk, "RSA/ECB/OAEPWithSHA-1AndMGF1Padding", Cipher.DECRYPT_MODE, value, OAEP_SHA1); + } + case "RSA-OAEP-256" -> { + requireRsa(base, alg); + yield rsaCrypt(jwk, "RSA/ECB/OAEPWithSHA-256AndMGF1Padding", Cipher.DECRYPT_MODE, value, OAEP_SHA256); + } + case "A128GCM", "A192GCM", "A256GCM" -> { + requireOct(base, alg); + requireOctSize(jwk, alg); + yield gcmDecrypt(jwk, value, iv, aad, tag); + } + default -> throw new CryptoException("Unsupported algorithm: " + alg); + }; + } + + private static byte[] rsaCrypt(Map jwk, String transform, int mode, byte[] input, + OAEPParameterSpec oaep) { + try { + Cipher cipher = Cipher.getInstance(transform); + if (oaep != null) { + cipher.init(mode, mode == Cipher.ENCRYPT_MODE ? reconstructRsaPublic(jwk) : reconstructRsaPrivate(jwk), oaep); + } else { + cipher.init(mode, mode == Cipher.ENCRYPT_MODE ? reconstructRsaPublic(jwk) : reconstructRsaPrivate(jwk)); + } + return cipher.doFinal(input); + } catch (CryptoException e) { + throw e; + } catch (Exception e) { + throw new CryptoException("RSA operation failed", e); + } + } + + private static CipherResult gcmEncrypt(Map jwk, byte[] plaintext, byte[] ivIgnored, byte[] aad) { + byte[] ivBytes = randomIv(12); + try { + Cipher cipher = Cipher.getInstance("AES/GCM/NoPadding"); + cipher.init(Cipher.ENCRYPT_MODE, reconstructSecretKey(jwk), new GCMParameterSpec(128, ivBytes)); + if (aad != null && aad.length > 0) { + cipher.updateAAD(aad); + } + byte[] out = cipher.doFinal(plaintext); + byte[] ct = Arrays.copyOfRange(out, 0, out.length - 16); + byte[] tag = Arrays.copyOfRange(out, out.length - 16, out.length); + return new CipherResult(ct, ivBytes, tag); + } catch (Exception e) { + throw new CryptoException("AES-GCM encryption failed", e); + } + } + + private static byte[] gcmDecrypt(Map jwk, byte[] ciphertext, byte[] iv, byte[] aad, byte[] tag) { + if (tag == null || tag.length == 0) { + throw new CryptoException("AES-GCM decryption requires a tag"); + } + byte[] combined = new byte[ciphertext.length + tag.length]; + System.arraycopy(ciphertext, 0, combined, 0, ciphertext.length); + System.arraycopy(tag, 0, combined, ciphertext.length, tag.length); + try { + Cipher cipher = Cipher.getInstance("AES/GCM/NoPadding"); + cipher.init(Cipher.DECRYPT_MODE, reconstructSecretKey(jwk), new GCMParameterSpec(128, iv)); + if (aad != null && aad.length > 0) { + cipher.updateAAD(aad); + } + return cipher.doFinal(combined); + } catch (Exception e) { + throw new CryptoException("AES-GCM decryption failed", e); + } + } + + private static byte[] randomIv(int length) { + byte[] iv = new byte[length]; + SECURE_RANDOM.nextBytes(iv); + return iv; + } + + // ── Sign / verify ────────────────────────────────────────────────────────── + + static byte[] sign(Map jwk, String alg, byte[] digest) { + validateDigestLength(alg, digest); + String base = baseKty((String) jwk.get("kty")); + return switch (alg) { + case "RS256", "RS384", "RS512" -> { + requireRsa(base, alg); + yield rsSign(jwk, digest, digestInfoPrefix(alg)); + } + case "PS256", "PS384", "PS512" -> { + requireRsa(base, alg); + yield pssSign(jwk, digest, hashName(alg)); + } + case "ES256", "ES384", "ES512" -> { + requireEc(base, alg); + requireCurve(jwk, alg); + yield ecSign(jwk, digest, curveByteLength(alg)); + } + default -> throw new CryptoException("Unsupported algorithm: " + alg); + }; + } + + static boolean verify(Map jwk, String alg, byte[] digest, byte[] signature) { + validateDigestLength(alg, digest); + String base = baseKty((String) jwk.get("kty")); + return switch (alg) { + case "RS256", "RS384", "RS512" -> { + requireRsa(base, alg); + yield rsVerify(jwk, digest, digestInfoPrefix(alg), signature); + } + case "PS256", "PS384", "PS512" -> { + requireRsa(base, alg); + yield pssVerify(jwk, digest, signature, hashName(alg)); + } + case "ES256", "ES384", "ES512" -> { + requireEc(base, alg); + requireCurve(jwk, alg); + yield ecVerify(jwk, digest, signature, curveByteLength(alg)); + } + default -> throw new CryptoException("Unsupported algorithm: " + alg); + }; + } + + private static void requireRsa(String base, String alg) { + if (!"RSA".equals(base)) { + throw new CryptoException("Algorithm " + alg + " requires an RSA key"); + } + } + + private static void requireEc(String base, String alg) { + if (!"EC".equals(base)) { + throw new CryptoException("Algorithm " + alg + " requires an EC key"); + } + } + + private static void requireOct(String base, String alg) { + if (!"oct".equals(base)) { + throw new CryptoException("Algorithm " + alg + " requires an oct key"); + } + } + + /** Rejects an AES-GCM algorithm when the stored oct key's size does not match the algorithm. */ + private static void requireOctSize(Map jwk, String alg) { + int bits = jwk.get("keySize") instanceof Number n ? n.intValue() : 0; + if (bits == 0 && jwk.get("k") instanceof String k) { + bits = b64UrlDecode(k).length * 8; + } + int expected = switch (alg) { + case "A128GCM" -> 128; + case "A192GCM" -> 192; + case "A256GCM" -> 256; + default -> 0; + }; + if (bits != expected) { + throw new CryptoException("Algorithm " + alg + " requires a " + expected + "-bit key"); + } + } + + /** The key type (RSA/EC/oct) required by an algorithm, or {@code null} for an unknown algorithm. */ + static String keyTypeForAlg(String alg) { + return switch (alg) { + case "RSA1_5", "RSA-OAEP", "RSA-OAEP-256", + "RS256", "RS384", "RS512", "PS256", "PS384", "PS512" -> "RSA"; + case "A128GCM", "A192GCM", "A256GCM" -> "oct"; + case "ES256", "ES384", "ES512" -> "EC"; + default -> null; + }; + } + + // ── RS* (EMSA-PKCS1-v1_5 by hand, raw digest input) ──────────────────────── + + private static byte[] rsSign(Map jwk, byte[] digest, byte[] prefix) { + byte[] em = concat(prefix, digest); + try { + Cipher cipher = Cipher.getInstance("RSA/ECB/PKCS1Padding"); + cipher.init(Cipher.ENCRYPT_MODE, reconstructRsaPrivate(jwk)); + return cipher.doFinal(em); + } catch (CryptoException e) { + throw e; + } catch (Exception e) { + throw new CryptoException("RSA sign failed", e); + } + } + + private static boolean rsVerify(Map jwk, byte[] digest, byte[] prefix, byte[] signature) { + byte[] em = concat(prefix, digest); + try { + Cipher cipher = Cipher.getInstance("RSA/ECB/PKCS1Padding"); + cipher.init(Cipher.DECRYPT_MODE, reconstructRsaPublic(jwk)); + byte[] recovered = cipher.doFinal(signature); + return MessageDigest.isEqual(em, recovered); + } catch (Exception e) { + return false; + } + } + + // ── PS* (manual EMSA-PSS, RFC 8017 §9.1.1, salt length = hash length) ────── + + private static byte[] pssSign(Map jwk, byte[] mHash, String hashName) { + PrivateKey priv = reconstructRsaPrivate(jwk); + int modBits = rsaModulusBits(priv); + int emBits = modBits - 1; + int emLen = (emBits + 7) / 8; + int hLen = mHash.length; + int sLen = hLen; + + byte[] salt = new byte[sLen]; + SECURE_RANDOM.nextBytes(salt); + + byte[] mPrime = new byte[8 + hLen + sLen]; + System.arraycopy(mHash, 0, mPrime, 8, hLen); + System.arraycopy(salt, 0, mPrime, 8 + hLen, sLen); + byte[] h = hash(hashName, mPrime); + + int psLen = emLen - sLen - hLen - 2; + if (psLen < 0) { + throw new CryptoException("RSA key too small for " + hashName + " PSS signature"); + } + byte[] db = new byte[psLen + 1 + sLen]; + db[psLen] = 1; + System.arraycopy(salt, 0, db, psLen + 1, sLen); + + byte[] dbMask = mgf1(h, emLen - hLen - 1, hashName); + byte[] maskedDb = xor(db, dbMask); + maskedDb[0] &= (byte) (0xFF >>> (8 * emLen - emBits)); + + byte[] em = new byte[emLen]; + System.arraycopy(maskedDb, 0, em, 0, emLen - hLen - 1); + System.arraycopy(h, 0, em, emLen - hLen - 1, hLen); + em[emLen - 1] = (byte) 0xbc; + + return rawRsaPrivateOp(priv, em, emLen); + } + + private static boolean pssVerify(Map jwk, byte[] mHash, byte[] signature, String hashName) { + PublicKey pub = reconstructRsaPublic(jwk); + int modBits = rsaModulusBits(pub); + int emBits = modBits - 1; + int emLen = (emBits + 7) / 8; + int hLen = mHash.length; + int sLen = hLen; + + byte[] em; + try { + Cipher cipher = Cipher.getInstance("RSA/ECB/NoPadding"); + cipher.init(Cipher.DECRYPT_MODE, pub); + em = leftPad(cipher.doFinal(signature), emLen); + } catch (Exception e) { + return false; + } + if (em.length != emLen || em[emLen - 1] != (byte) 0xbc) { + return false; + } + + int psLen = emLen - sLen - hLen - 2; + if (psLen < 0) { + return false; + } + + byte[] maskedDb = Arrays.copyOfRange(em, 0, emLen - hLen - 1); + byte[] h = Arrays.copyOfRange(em, emLen - hLen - 1, emLen - 1); + maskedDb[0] &= (byte) (0xFF >>> (8 * emLen - emBits)); + + byte[] dbMask = mgf1(h, emLen - hLen - 1, hashName); + byte[] db = xor(maskedDb, dbMask); + db[0] &= (byte) (0xFF >>> (8 * emLen - emBits)); + + for (int i = 0; i < psLen; i++) { + if (db[i] != 0) { + return false; + } + } + if (db[psLen] != 1) { + return false; + } + + byte[] salt = Arrays.copyOfRange(db, psLen + 1, db.length); + byte[] mPrime = new byte[8 + hLen + sLen]; + System.arraycopy(mHash, 0, mPrime, 8, hLen); + System.arraycopy(salt, 0, mPrime, 8 + hLen, sLen); + byte[] h2 = hash(hashName, mPrime); + return MessageDigest.isEqual(h, h2); + } + + // ── ES* (NONEwithECDSA DER + manual DER ↔ raw R‖S) ───────────────────────── + + private static byte[] ecSign(Map jwk, byte[] digest, int curveByteLength) { + try { + Signature signature = Signature.getInstance("NONEwithECDSA"); + signature.initSign(reconstructEcPrivate(jwk)); + signature.update(digest); + return derToRawRs(signature.sign(), curveByteLength); + } catch (CryptoException e) { + throw e; + } catch (Exception e) { + throw new CryptoException("ECDSA sign failed", e); + } + } + + private static boolean ecVerify(Map jwk, byte[] digest, byte[] raw, int curveByteLength) { + try { + Signature signature = Signature.getInstance("NONEwithECDSA"); + signature.initVerify(reconstructEcPublic(jwk)); + signature.update(digest); + return signature.verify(rawRsToDer(raw, curveByteLength)); + } catch (Exception e) { + return false; + } + } + + /** ASN.1 DER {@code SEQUENCE { INTEGER r, INTEGER s }} → fixed-size {@code R‖S} (IEEE P1363). */ + static byte[] derToRawRs(byte[] der, int curveByteLength) { + BigInteger[] rs = decodeDerSignature(der); + byte[] raw = new byte[2 * curveByteLength]; + writeFixedLength(raw, 0, rs[0], curveByteLength); + writeFixedLength(raw, curveByteLength, rs[1], curveByteLength); + return raw; + } + + /** Fixed-size {@code R‖S} (IEEE P1363) → ASN.1 DER {@code SEQUENCE { INTEGER r, INTEGER s }}. */ + static byte[] rawRsToDer(byte[] raw, int curveByteLength) { + if (raw.length != 2 * curveByteLength) { + throw new CryptoException("invalid raw R||S signature length"); + } + BigInteger r = new BigInteger(1, Arrays.copyOfRange(raw, 0, curveByteLength)); + BigInteger s = new BigInteger(1, Arrays.copyOfRange(raw, curveByteLength, 2 * curveByteLength)); + byte[] rDer = derInteger(r); + byte[] sDer = derInteger(s); + byte[] len = derLength(rDer.length + sDer.length); + byte[] out = new byte[1 + len.length + rDer.length + sDer.length]; + int pos = 0; + out[pos++] = 0x30; + System.arraycopy(len, 0, out, pos, len.length); + pos += len.length; + System.arraycopy(rDer, 0, out, pos, rDer.length); + pos += rDer.length; + System.arraycopy(sDer, 0, out, pos, sDer.length); + return out; + } + + private static BigInteger[] decodeDerSignature(byte[] der) { + int pos = 0; + if (pos >= der.length || der[pos++] != 0x30) { + throw new CryptoException("invalid DER signature"); + } + int[] l1 = readDerLength(der, pos); + pos = l1[1]; + if (pos >= der.length || der[pos++] != 0x02) { + throw new CryptoException("invalid DER signature"); + } + int[] l2 = readDerLength(der, pos); + pos = l2[1]; + byte[] rBytes = Arrays.copyOfRange(der, pos, pos + l2[0]); + pos += l2[0]; + if (pos >= der.length || der[pos++] != 0x02) { + throw new CryptoException("invalid DER signature"); + } + int[] l3 = readDerLength(der, pos); + pos = l3[1]; + byte[] sBytes = Arrays.copyOfRange(der, pos, pos + l3[0]); + return new BigInteger[]{new BigInteger(rBytes), new BigInteger(sBytes)}; + } + + private static int[] readDerLength(byte[] der, int pos) { + int b = der[pos] & 0xFF; + if (b < 0x80) { + return new int[]{b, pos + 1}; + } + int numBytes = b & 0x7F; + int len = 0; + for (int i = 0; i < numBytes; i++) { + len = (len << 8) | (der[pos + 1 + i] & 0xFF); + } + return new int[]{len, pos + 1 + numBytes}; + } + + private static byte[] derInteger(BigInteger v) { + byte[] b = v.toByteArray(); + byte[] out = new byte[2 + b.length]; + out[0] = 0x02; + out[1] = (byte) b.length; + System.arraycopy(b, 0, out, 2, b.length); + return out; + } + + private static byte[] derLength(int len) { + if (len < 0x80) { + return new byte[]{(byte) len}; + } else if (len <= 0xFF) { + return new byte[]{(byte) 0x81, (byte) len}; + } else { + return new byte[]{(byte) 0x82, (byte) (len >> 8), (byte) len}; + } + } + + private static void writeFixedLength(byte[] out, int off, BigInteger v, int len) { + byte[] b = v.toByteArray(); + int start = 0; + if (b.length > len) { + if (b.length == len + 1 && b[0] == 0) { + start = 1; + } else { + throw new CryptoException("R||S component too large for curve"); + } + } + int pad = len - (b.length - start); + Arrays.fill(out, off, off + pad, (byte) 0); + System.arraycopy(b, start, out, off + pad, b.length - start); + } + + // ── Shared primitives ────────────────────────────────────────────────────── + + private static byte[] rawRsaPrivateOp(PrivateKey priv, byte[] em, int modLen) { + try { + Cipher cipher = Cipher.getInstance("RSA/ECB/NoPadding"); + cipher.init(Cipher.ENCRYPT_MODE, priv); + return leftPad(cipher.doFinal(em), modLen); + } catch (Exception e) { + throw new CryptoException("RSA private operation failed", e); + } + } + + private static int rsaModulusBits(java.security.Key key) { + return ((RSAKey) key).getModulus().bitLength(); + } + + private static byte[] hash(String name, byte[] data) { + try { + return MessageDigest.getInstance(name).digest(data); + } catch (Exception e) { + throw new CryptoException("Hash algorithm unavailable: " + name, e); + } + } + + private static byte[] mgf1(byte[] seed, int maskLen, String hashName) { + try { + int hLen = hashLen(hashName); + byte[] out = new byte[maskLen]; + MessageDigest md = MessageDigest.getInstance(hashName); + byte[] counter = new byte[4]; + for (int i = 0, done = 0; done < maskLen; i++) { + counter[0] = (byte) (i >>> 24); + counter[1] = (byte) (i >>> 16); + counter[2] = (byte) (i >>> 8); + counter[3] = (byte) i; + md.update(seed); + md.update(counter); + byte[] t = md.digest(); + int toCopy = Math.min(hLen, maskLen - done); + System.arraycopy(t, 0, out, done, toCopy); + done += toCopy; + } + return out; + } catch (Exception e) { + throw new CryptoException("MGF1 unavailable for " + hashName, e); + } + } + + private static int hashLen(String hashName) { + return switch (hashName) { + case "SHA-256" -> 32; + case "SHA-384" -> 48; + case "SHA-512" -> 64; + default -> throw new CryptoException("Unsupported hash: " + hashName); + }; + } + + private static byte[] xor(byte[] a, byte[] b) { + byte[] out = new byte[a.length]; + for (int i = 0; i < a.length; i++) { + out[i] = (byte) (a[i] ^ b[i]); + } + return out; + } + + private static byte[] concat(byte[] a, byte[] b) { + byte[] out = new byte[a.length + b.length]; + System.arraycopy(a, 0, out, 0, a.length); + System.arraycopy(b, 0, out, a.length, b.length); + return out; + } + + private static byte[] leftPad(byte[] in, int len) { + if (in.length >= len) { + return in; + } + byte[] out = new byte[len]; + System.arraycopy(in, 0, out, len - in.length, in.length); + return out; + } + + private static byte[] digestInfoPrefix(String alg) { + return switch (alg) { + case "RS256" -> hex("3031300d060960864801650304020105000420"); + case "RS384" -> hex("3041300d060960864801650304020205000430"); + case "RS512" -> hex("3051300d060960864801650304020305000440"); + default -> throw new CryptoException("Unsupported algorithm: " + alg); + }; + } + + private static String hashName(String alg) { + return switch (alg) { + case "PS256" -> "SHA-256"; + case "PS384" -> "SHA-384"; + case "PS512" -> "SHA-512"; + default -> throw new CryptoException("Unsupported algorithm: " + alg); + }; + } + + private static int curveByteLength(String alg) { + return switch (alg) { + case "ES256" -> 32; + case "ES384" -> 48; + case "ES512" -> 66; + default -> throw new CryptoException("Unsupported algorithm: " + alg); + }; + } + + private static int digestByteLength(String alg) { + return switch (alg) { + case "RS256", "PS256", "ES256" -> 32; + case "RS384", "PS384", "ES384" -> 48; + case "RS512", "PS512", "ES512" -> 64; + default -> throw new CryptoException("Unsupported algorithm: " + alg); + }; + } + + private static String curveForAlg(String alg) { + return switch (alg) { + case "ES256" -> "P-256"; + case "ES384" -> "P-384"; + case "ES512" -> "P-521"; + default -> null; + }; + } + + private static void validateDigestLength(String alg, byte[] digest) { + if (digest == null || digest.length != digestByteLength(alg)) { + throw new CryptoException("Digest length " + (digest == null ? 0 : digest.length) + + " does not match algorithm " + alg); + } + } + + private static void requireCurve(Map jwk, String alg) { + String expected = curveForAlg(alg); + String crv = (String) jwk.get("crv"); + if (expected != null && (crv == null || !expected.equals(crv))) { + throw new CryptoException("Algorithm " + alg + " requires curve " + expected); + } + } + + private static byte[] hex(String s) { + byte[] out = new byte[s.length() / 2]; + for (int i = 0; i < out.length; i++) { + out[i] = (byte) Integer.parseInt(s.substring(2 * i, 2 * i + 2), 16); + } + return out; + } +} diff --git a/src/main/java/io/floci/az/services/keyvault/KeyVaultHandler.java b/src/main/java/io/floci/az/services/keyvault/KeyVaultHandler.java index d2e83f2f..81efc446 100644 --- a/src/main/java/io/floci/az/services/keyvault/KeyVaultHandler.java +++ b/src/main/java/io/floci/az/services/keyvault/KeyVaultHandler.java @@ -33,11 +33,13 @@ public class KeyVaultHandler implements AzureServiceHandler, Resettable { private final StorageBackend store; private final EmulatorConfig config; + private final KeyVaultKeys keys; @Inject public KeyVaultHandler(StorageFactory factory, EmulatorConfig config) { this.store = factory.create("keyvault"); this.config = config; + this.keys = new KeyVaultKeys(store); } @Override @@ -56,7 +58,9 @@ public boolean enabled(String serviceType) { public ServiceRoutes routes() { return ServiceRoutes.builder() .host(".vault.azure.net") + .host(".managedhsm.azure.net") .account("-keyvault", "keyvault") + .account("-managedhsm", "keyvault") .build(); } @@ -76,24 +80,32 @@ public Response handle(AzureRequest req) { LOG.debugf("KeyVault %s /%s", method, path); + // Managed HSM flavor: strip the port from Host so {account}.managedhsm.azure.net:4577 still matches. + String host = req.headers().getHeaderString("Host"); + String hostWithoutPort = host != null + ? (host.contains(":") ? host.substring(0, host.indexOf(':')) : host) : null; + String accountSuffix = req.headers().getHeaderString("x-floci-account-suffix"); + boolean hsm = (hostWithoutPort != null && hostWithoutPort.endsWith(".managedhsm.azure.net")) + || "-managedhsm".equals(accountSuffix); + // The Azure SDK challenge_auth_policy sends a bodiless probe to elicit a challenge, then // retries with the real body. Return 401 with a Bearer challenge so the SDK caches it and // sends subsequent requests with the Authorization header and their original bodies. String auth = req.headers().getHeaderString("Authorization"); if (auth == null || auth.isEmpty()) { + String challengeResource = hsm ? "https://managedhsm.azure.net" : "https://vault.azure.net"; return Response.status(401) .header("WWW-Authenticate", "Bearer authorization=\"https://login.microsoftonline.com/common\", " - + "resource=\"https://vault.azure.net\"") + + "resource=\"" + challengeResource + "\"") .build(); } // Root probe — azurerm provider polls this to confirm the vault is reachable. if (routePath.isEmpty()) { - return Response.ok(java.util.Map.of( - "type", "Microsoft.KeyVault/vaults", - "id", "https://" + account + ".vault.azure.net/" - )).build(); + String probeType = hsm ? "Microsoft.KeyVault/managedHSMs" : "Microsoft.KeyVault/vaults"; + String probeId = "https://" + account + (hsm ? ".managedhsm.azure.net/" : ".vault.azure.net/"); + return Response.ok(java.util.Map.of("type", probeType, "id", probeId)).build(); } if ("secrets".equals(routePath)) { @@ -109,6 +121,23 @@ public Response handle(AzureRequest req) { return handleDeletedSecrets(req, method, account, path.substring("deletedsecrets/".length())); } + // Keys data plane + if ("keys".equals(path)) { + return "GET".equals(method) ? keys.listKeys(account, hsm) : methodNotAllowed(); + } + if (path.startsWith("keys/")) { + return handleKeys(req, method, account, path.substring("keys/".length()), hsm); + } + if ("deletedkeys".equals(path)) { + return "GET".equals(method) ? keys.listDeletedKeys(account, hsm) : methodNotAllowed(); + } + if (path.startsWith("deletedkeys/")) { + return handleDeletedKeys(req, method, account, path.substring("deletedkeys/".length()), hsm); + } + if ("rng".equals(path)) { + return "POST".equals(method) ? handleRng(req) : methodNotAllowed(); + } + // Certificate contacts — azurerm provider reads this after key vault creation. // Return an empty contacts list so the provider sees no contacts configured. if ("certificates/contacts".equals(routePath)) { @@ -121,7 +150,7 @@ public Response handle(AzureRequest req) { return methodNotAllowed(); } - return kvError(404, "BadRequest", "Resource not found: " + path); + return kvError(404, "KeyNotFound", "Resource not found: " + path); } // ------------------------------------------------------------------------- @@ -195,6 +224,150 @@ private Response handleDeletedSecrets(AzureRequest req, String method, String ac }; } + // ------------------------------------------------------------------------- + // Keys sub-routers + // ------------------------------------------------------------------------- + + private Response handleKeys(AzureRequest req, String method, String account, String rest, boolean hsm) { + // /keys/{name}/backup + int backupIdx = rest.indexOf("/backup"); + if (backupIdx != -1) { + String name = rest.substring(0, backupIdx); + return "POST".equals(method) ? keys.backupKey(account, name, hsm) : methodNotAllowed(); + } + + // /keys/restore — POST restores; any other verb treats "restore" as a key name. + if ("restore".equals(rest) && "POST".equals(method)) { + return keys.restoreKey(req, account, hsm); + } + + // /keys/{name}/create (POST; a PUT /keys/{name} is an import) + int createIdx = rest.indexOf("/create"); + if (createIdx != -1) { + String name = rest.substring(0, createIdx); + return "POST".equals(method) ? keys.createKey(req, account, name, hsm) : methodNotAllowed(); + } + + // /keys/{name}/versions[/{version}] + int versionsIdx = rest.indexOf("/versions"); + if (versionsIdx != -1) { + String name = rest.substring(0, versionsIdx); + String afterVersions = rest.substring(versionsIdx + "/versions".length()); + if (afterVersions.isEmpty() || "/".equals(afterVersions)) { + return "GET".equals(method) ? keys.listKeyVersions(account, name, hsm) : methodNotAllowed(); + } + String version = afterVersions.startsWith("/") ? afterVersions.substring(1) : afterVersions; + return switch (method) { + case "GET" -> keys.getKeyVersion(account, name, version, hsm); + case "PATCH" -> keys.updateKeyProperties(req, account, name, version, hsm); + default -> methodNotAllowed(); + }; + } + + // /keys/{name}/rotationpolicy + int rotationIdx = rest.indexOf("/rotationpolicy"); + if (rotationIdx != -1) { + String name = rest.substring(0, rotationIdx); + return switch (method) { + case "GET" -> keys.getRotationPolicy(account, name, hsm); + case "PUT" -> keys.putRotationPolicy(req, account, name, hsm); + default -> methodNotAllowed(); + }; + } + + // /keys/{name}/rotate + int rotateIdx = rest.indexOf("/rotate"); + if (rotateIdx != -1) { + String name = rest.substring(0, rotateIdx); + return "POST".equals(method) ? keys.rotateKey(account, name, hsm) : methodNotAllowed(); + } + + // /keys/{name}/{version}[/{op}] or /keys/{name}/{op} (empty/omitted version) + int slash = rest.indexOf('/'); + if (slash != -1) { + String name = rest.substring(0, slash); + String remainder = rest.substring(slash + 1); + int slash2 = remainder.indexOf('/'); + if (slash2 != -1) { + String version = remainder.substring(0, slash2); + String op = remainder.substring(slash2 + 1); + if (isCryptoOp(op)) { + return "POST".equals(method) ? keys.cryptoOp(req, account, name, version, op, hsm) + : methodNotAllowed(); + } + return kvError(404, "KeyNotFound", "Resource not found: " + pathLabel(name, remainder)); + } + if (isCryptoOp(remainder)) { + return "POST".equals(method) ? keys.cryptoOp(req, account, name, "", remainder, hsm) + : methodNotAllowed(); + } + if (remainder.isEmpty()) { + return switch (method) { + case "GET" -> keys.getKey(account, name, hsm); + case "PUT" -> keys.importKey(req, account, name, hsm); + case "DELETE" -> keys.deleteKey(account, name, hsm); + // az keyvault key set-attributes PATCHes the latest version with an empty + // version segment (/keys/{name}/). + case "PATCH" -> keys.updateKeyPropertiesLatest(req, account, name, hsm); + default -> methodNotAllowed(); + }; + } + String version = remainder; + return switch (method) { + case "GET" -> keys.getKeyVersion(account, name, version, hsm); + case "PATCH" -> keys.updateKeyProperties(req, account, name, version, hsm); + default -> methodNotAllowed(); + }; + } + + // /keys/{name} (latest) + return switch (method) { + case "GET" -> keys.getKey(account, rest, hsm); + case "PUT" -> keys.importKey(req, account, rest, hsm); + case "DELETE" -> keys.deleteKey(account, rest, hsm); + case "PATCH" -> keys.updateKeyPropertiesLatest(req, account, rest, hsm); + default -> methodNotAllowed(); + }; + } + + private Response handleDeletedKeys(AzureRequest req, String method, String account, String rest, boolean hsm) { + int recoverIdx = rest.indexOf("/recover"); + if (recoverIdx != -1) { + String name = rest.substring(0, recoverIdx); + return "POST".equals(method) ? keys.recoverDeletedKey(account, name, hsm) : methodNotAllowed(); + } + return switch (method) { + case "GET" -> keys.getDeletedKey(account, rest, hsm); + case "DELETE" -> keys.purgeDeletedKey(account, rest, hsm); + default -> methodNotAllowed(); + }; + } + + private Response handleRng(AzureRequest req) { + Map body = parseBody(req); + int count = 32; + if (body.get("count") instanceof Number n) { + count = n.intValue(); + } + if (count < 1 || count > 128) { + return kvError(400, "BadParameter", "The count parameter must be between 1 and 128."); + } + byte[] bytes = new byte[count]; + new java.security.SecureRandom().nextBytes(bytes); + Map response = new LinkedHashMap<>(); + response.put("value", java.util.Base64.getUrlEncoder().withoutPadding().encodeToString(bytes)); + return Response.ok(toJson(response), "application/json").build(); + } + + private static boolean isCryptoOp(String op) { + return op.equals("encrypt") || op.equals("decrypt") || op.equals("wrapkey") + || op.equals("unwrapkey") || op.equals("sign") || op.equals("verify"); + } + + private static String pathLabel(String name, String remainder) { + return "keys/" + name + "/" + remainder; + } + // ------------------------------------------------------------------------- // Secret CRUD // ------------------------------------------------------------------------- diff --git a/src/main/java/io/floci/az/services/keyvault/KeyVaultKeys.java b/src/main/java/io/floci/az/services/keyvault/KeyVaultKeys.java new file mode 100644 index 00000000..74845996 --- /dev/null +++ b/src/main/java/io/floci/az/services/keyvault/KeyVaultKeys.java @@ -0,0 +1,1040 @@ +package io.floci.az.services.keyvault; + +import com.fasterxml.jackson.core.JsonProcessingException; +import com.fasterxml.jackson.core.type.TypeReference; +import com.fasterxml.jackson.databind.ObjectMapper; +import io.floci.az.core.AzureRequest; +import io.floci.az.core.StoredObject; +import io.floci.az.core.arm.ArmJson; +import io.floci.az.core.storage.StorageBackend; +import jakarta.ws.rs.core.Response; +import org.jboss.logging.Logger; + +import java.io.IOException; +import java.time.Instant; +import java.util.ArrayList; +import java.util.Base64; +import java.util.HashMap; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; +import java.util.Optional; +import java.util.UUID; +import java.util.stream.Collectors; + +/** + * Key-object business logic for the Key Vault data plane: CRUD, versions, soft-delete / recover / + * purge, backup / restore, rotation policy, and crypto dispatch. + * + *

Plain object (not a CDI bean) built by {@link KeyVaultHandler} with the shared + * {@code "keyvault"} {@link StorageBackend}; keys are stored under {@code {account}/keys/...} to + * avoid colliding with secrets. Crypto delegates to the pure-JDK {@link KeyVaultCrypto} engine. + */ +final class KeyVaultKeys { + + private static final Logger LOG = Logger.getLogger(KeyVaultKeys.class); + private static final ObjectMapper MAPPER = new ObjectMapper(); + private static final Base64.Encoder URL = Base64.getUrlEncoder().withoutPadding(); + private static final String RECOVERY_LEVEL = "Purgeable"; + private static final long PURGE_RETENTION_SECONDS = 7L * 24 * 3600; + + private final StorageBackend store; + + KeyVaultKeys(StorageBackend store) { + this.store = store; + } + + // ── Create / import ──────────────────────────────────────────────────────── + + Response createKey(AzureRequest req, String account, String name, boolean hsm) { + Response invalid = validateKeyName(name); + if (invalid != null) { + return invalid; + } + Map body = parseBody(req); + String kty = bodyString(body, "kty", "RSA"); + int keySize = bodyInt(body, "key_size", 0); + String curve = body.containsKey("curve") ? bodyString(body, "curve", null) + : bodyString(body, "crv", null); + @SuppressWarnings("unchecked") + List keyOps = body.containsKey("key_ops") + ? (body.get("key_ops") instanceof List l ? (List) l : List.of()) + : KeyVaultCrypto.defaultKeyOps(KeyVaultCrypto.baseKty(kty)); + + Map jwk; + try { + jwk = KeyVaultCrypto.generateJwk(kty, keySize, curve, keyOps); + } catch (KeyVaultCrypto.CryptoException e) { + return kvError(400, "BadParameter", e.getMessage()); + } + jwk.put("hsm", kty.endsWith("-HSM")); + jwk.put("tags", stringMap(body.get("tags"))); + + return storeNewKey(req, account, name, jwk, body, hsm); + } + + Response importKey(AzureRequest req, String account, String name, boolean hsm) { + Response invalid = validateKeyName(name); + if (invalid != null) { + return invalid; + } + Map body = parseBody(req); + @SuppressWarnings("unchecked") + Map supplied = body.get("key") instanceof Map m + ? (Map) m : new LinkedHashMap<>(); + + Map jwk; + try { + jwk = KeyVaultCrypto.importJwk(supplied); + } catch (KeyVaultCrypto.CryptoException e) { + return kvError(400, "BadParameter", e.getMessage()); + } + jwk.put("hsm", Boolean.parseBoolean(String.valueOf(body.getOrDefault("HSM", false)))); + jwk.put("tags", stringMap(body.get("tags"))); + + return storeNewKey(req, account, name, jwk, body, hsm); + } + + private Response storeNewKey(AzureRequest req, String account, String name, + Map jwk, Map body, boolean hsm) { + // Recreating a name that is soft-deleted (not purged) is a 409 until it is recovered or purged. + if (store.get(deletedKeyKey(account, name, hsm)).isPresent()) { + return kvError(409, "Conflict", + "A key with (name/id) " + name + " was recently deleted and must be recovered or purged first."); + } + @SuppressWarnings("unchecked") + Map attrs = body.get("attributes") instanceof Map m + ? (Map) m : new LinkedHashMap<>(); + + long now = Instant.now().getEpochSecond(); + String versionId = newVersionId(); + + Map meta = new HashMap<>(); + meta.put("version", versionId); + meta.put("created", String.valueOf(now)); + meta.put("updated", String.valueOf(now)); + try { + normalizeAttrMeta(meta, attrs); + } catch (IllegalArgumentException e) { + return kvError(400, "BadParameter", e.getMessage()); + } + + StoredObject versionObj = new StoredObject(keyVersionKey(account, name, versionId, hsm), + toBytes(jwk), meta, Instant.now(), newVersionId().substring(0, 16)); + store.put(versionObj.key(), versionObj); + + Map latestMeta = new HashMap<>(versionObj.metadata()); + latestMeta.put("latestVersion", versionId); + String latestKey = keyLatestKey(account, name, hsm); + store.put(latestKey, new StoredObject(latestKey, versionObj.data(), latestMeta, + versionObj.lastModified(), versionObj.etag())); + + return Response.ok(toJson(keyBundle(account, name, versionId, versionObj, hsm)), + "application/json").build(); + } + + // ── Read ─────────────────────────────────────────────────────────────────── + + Response getKey(String account, String name, boolean hsm) { + Optional opt = store.get(keyLatestKey(account, name, hsm)); + if (opt.isEmpty()) { + return keyNotFound(name); + } + StoredObject obj = opt.get(); + String versionId = obj.metadata().getOrDefault("latestVersion", obj.metadata().get("version")); + return Response.ok(toJson(keyBundle(account, name, versionId, obj, hsm)), "application/json").build(); + } + + Response getKeyVersion(String account, String name, String version, boolean hsm) { + if (store.get(deletedKeyKey(account, name, hsm)).isPresent()) { + return keyNotFound(name + "/" + version); + } + Optional opt = store.get(keyVersionKey(account, name, version, hsm)); + if (opt.isEmpty()) { + return keyNotFound(name + "/" + version); + } + return Response.ok(toJson(keyBundle(account, name, version, opt.get(), hsm)), "application/json").build(); + } + + Response listKeys(String account, boolean hsm) { + String prefix = keysPrefix(account, hsm); + List> items = store.scan(k -> k.startsWith(prefix)) + .stream() + .filter(obj -> !obj.key().substring(prefix.length()).contains("/")) + .map(obj -> { + String name = obj.key().substring(prefix.length()); + return keyItem(account, name, obj, hsm, false); + }) + .collect(Collectors.toCollection(ArrayList::new)); + return jsonList(items); + } + + Response listKeyVersions(String account, String name, boolean hsm) { + if (store.get(deletedKeyKey(account, name, hsm)).isPresent()) { + return keyNotFound(name); + } + String prefix = keyVersionsPrefix(account, name, hsm); + List> items = store.scan(k -> k.startsWith(prefix)) + .stream() + .map(obj -> { + String version = obj.key().substring(prefix.length()); + return keyItem(account, name + "/" + version, obj, hsm, true); + }) + .collect(Collectors.toCollection(ArrayList::new)); + return jsonList(items); + } + + // ── Update (PATCH) ───────────────────────────────────────────────────────── + + Response updateKeyProperties(AzureRequest req, String account, String name, String version, boolean hsm) { + if (store.get(deletedKeyKey(account, name, hsm)).isPresent()) { + return keyNotFound(name + "/" + version); + } + String versionKey = keyVersionKey(account, name, version, hsm); + Optional opt = store.get(versionKey); + if (opt.isEmpty()) { + return keyNotFound(name + "/" + version); + } + StoredObject existing = opt.get(); + Map body = parseBody(req); + @SuppressWarnings("unchecked") + Map attrs = body.get("attributes") instanceof Map m + ? (Map) m : new LinkedHashMap<>(); + + Map data = parseStoredData(existing); + if (body.containsKey("tags")) { + data.put("tags", body.get("tags")); + } + + Map meta = new HashMap<>(existing.metadata()); + try { + normalizeAttrMeta(meta, attrs); + } catch (IllegalArgumentException e) { + return kvError(400, "BadParameter", e.getMessage()); + } + long now = Instant.now().getEpochSecond(); + meta.put("updated", String.valueOf(now)); + + Instant updatedAt = Instant.now(); + String newEtag = newVersionId().substring(0, 16); + StoredObject updated = new StoredObject(versionKey, toBytes(data), meta, updatedAt, newEtag); + store.put(versionKey, updated); + + String latestKey = keyLatestKey(account, name, hsm); + store.get(latestKey).ifPresent(latest -> { + if (version.equals(latest.metadata().get("latestVersion"))) { + Map latestMeta = new HashMap<>(meta); + latestMeta.put("latestVersion", version); + store.put(latestKey, new StoredObject(latestKey, updated.data(), latestMeta, updatedAt, newEtag)); + } + }); + + return Response.ok(toJson(keyBundle(account, name, version, updated, hsm)), "application/json").build(); + } + + /** + * {@code PATCH /keys/{name}} with no version: the Azure CLI's {@code key set-attributes} PATCHes + * the latest version with an empty version segment. Resolves the latest version and delegates. + */ + Response updateKeyPropertiesLatest(AzureRequest req, String account, String name, boolean hsm) { + Optional opt = store.get(keyLatestKey(account, name, hsm)); + if (opt.isEmpty()) { + return keyNotFound(name); + } + String version = opt.get().metadata().getOrDefault("latestVersion", opt.get().metadata().get("version")); + return updateKeyProperties(req, account, name, version, hsm); + } + + // ── Delete / deleted / recover / purge ───────────────────────────────────── + + Response deleteKey(String account, String name, boolean hsm) { + String latestKey = keyLatestKey(account, name, hsm); + Optional opt = store.get(latestKey); + if (opt.isEmpty()) { + return keyNotFound(name); + } + StoredObject obj = opt.get(); + long now = Instant.now().getEpochSecond(); + long purge = now + PURGE_RETENTION_SECONDS; + + Map deletedMeta = new HashMap<>(obj.metadata()); + deletedMeta.put("deletedDate", String.valueOf(now)); + deletedMeta.put("scheduledPurgeDate", String.valueOf(purge)); + String deletedKey = deletedKeyKey(account, name, hsm); + store.put(deletedKey, new StoredObject(deletedKey, obj.data(), deletedMeta, + obj.lastModified(), obj.etag())); + store.delete(latestKey); + + String versionId = obj.metadata().getOrDefault("latestVersion", obj.metadata().get("version")); + return Response.ok(toJson(deletedKeyBundle(account, name, versionId, + store.get(deletedKey).get(), now, purge, hsm)), "application/json").build(); + } + + Response getDeletedKey(String account, String name, boolean hsm) { + Optional opt = store.get(deletedKeyKey(account, name, hsm)); + if (opt.isEmpty()) { + return deletedKeyNotFound(name); + } + StoredObject obj = opt.get(); + String versionId = obj.metadata().getOrDefault("latestVersion", obj.metadata().get("version")); + long deletedDate = parseLong(obj.metadata().get("deletedDate"), 0L); + long purgeDate = parseLong(obj.metadata().get("scheduledPurgeDate"), 0L); + return Response.ok(toJson(deletedKeyBundle(account, name, versionId, obj, deletedDate, purgeDate, hsm)), + "application/json").build(); + } + + Response listDeletedKeys(String account, boolean hsm) { + String prefix = deletedKeysPrefix(account, hsm); + List> items = store.scan(k -> k.startsWith(prefix)) + .stream() + .map(obj -> { + String name = obj.key().substring(prefix.length()); + String versionId = obj.metadata().getOrDefault("latestVersion", obj.metadata().get("version")); + long deleted = parseLong(obj.metadata().get("deletedDate"), 0L); + long purge = parseLong(obj.metadata().get("scheduledPurgeDate"), 0L); + return deletedKeyItem(account, name, versionId, obj, deleted, purge, hsm); + }) + .collect(Collectors.toCollection(ArrayList::new)); + return jsonList(items); + } + + Response recoverDeletedKey(String account, String name, boolean hsm) { + String deletedKey = deletedKeyKey(account, name, hsm); + Optional opt = store.get(deletedKey); + if (opt.isEmpty()) { + return deletedKeyNotFound(name); + } + StoredObject obj = opt.get(); + + Map restoredMeta = new HashMap<>(obj.metadata()); + restoredMeta.remove("deletedDate"); + restoredMeta.remove("scheduledPurgeDate"); + + String latestKey = keyLatestKey(account, name, hsm); + store.put(latestKey, new StoredObject(latestKey, obj.data(), restoredMeta, + obj.lastModified(), obj.etag())); + + String versionId = restoredMeta.getOrDefault("latestVersion", restoredMeta.get("version")); + String versionKey = keyVersionKey(account, name, versionId, hsm); + if (store.get(versionKey).isEmpty()) { + Map versionMeta = new HashMap<>(restoredMeta); + versionMeta.remove("latestVersion"); + store.put(versionKey, new StoredObject(versionKey, obj.data(), versionMeta, + obj.lastModified(), obj.etag())); + } + + store.delete(deletedKey); + + return Response.ok(toJson(keyBundle(account, name, versionId, store.get(latestKey).get(), hsm)), + "application/json").build(); + } + + Response purgeDeletedKey(String account, String name, boolean hsm) { + String deletedKey = deletedKeyKey(account, name, hsm); + if (store.get(deletedKey).isEmpty()) { + return deletedKeyNotFound(name); + } + store.delete(deletedKey); + String versionPrefix = keyVersionsPrefix(account, name, hsm); + store.scan(k -> k.startsWith(versionPrefix)).forEach(obj -> store.delete(obj.key())); + // Also remove any rotation policy so a recreated key does not inherit it. + store.delete(rotationPolicyKey(account, name, hsm)); + return Response.noContent().build(); + } + + // ── Backup / restore ─────────────────────────────────────────────────────── + + Response backupKey(String account, String name, boolean hsm) { + Optional opt = store.get(keyLatestKey(account, name, hsm)); + if (opt.isEmpty()) { + return keyNotFound(name); + } + StoredObject obj = opt.get(); + + Map snapshot = new LinkedHashMap<>(); + List> versions = new ArrayList<>(); + String versionPrefix = keyVersionsPrefix(account, name, hsm); + store.scan(k -> k.startsWith(versionPrefix)).forEach(v -> { + Map entry = new LinkedHashMap<>(); + entry.put("version", v.key().substring(versionPrefix.length())); + entry.put("data", parseStoredData(v)); + entry.put("meta", v.metadata()); + versions.add(entry); + }); + snapshot.put("name", name); + snapshot.put("versions", versions); + snapshot.put("latestVersion", obj.metadata().getOrDefault("latestVersion", obj.metadata().get("version"))); + snapshot.put("tags", parseStoredData(obj).get("tags")); + snapshot.put("rotationPolicy", store.get(rotationPolicyKey(account, name, hsm)) + .map(this::parseStoredData).orElse(null)); + + Map response = new LinkedHashMap<>(); + response.put("value", URL.encodeToString(toBytes(snapshot))); + return Response.ok(toJson(response), "application/json").build(); + } + + Response restoreKey(AzureRequest req, String account, boolean hsm) { + Map body = parseBody(req); + String value = bodyString(body, "value", null); + if (value == null) { + return kvError(400, "BadParameter", "A backup value is required."); + } + String name = null; + try { + Map snapshot; + try { + snapshot = MAPPER.readValue(KeyVaultCrypto.b64UrlDecode(value), new TypeReference<>() {}); + } catch (Exception e) { + return kvError(400, "BadParameter", "Invalid backup value."); + } + + name = bodyString(snapshot, "name", null); + if (name == null || !name.matches("^[a-zA-Z0-9-]+$")) { + return kvError(400, "BadParameter", "Invalid key name in backup snapshot."); + } + if (store.get(keyLatestKey(account, name, hsm)).isPresent() + || store.get(deletedKeyKey(account, name, hsm)).isPresent()) { + return kvError(409, "Conflict", "A key with (name/id) " + name + " already exists in this key vault."); + } + + @SuppressWarnings("unchecked") + List> versions = snapshot.get("versions") instanceof List l + ? (List>) l : List.of(); + if (versions.size() > 100) { + return kvError(400, "BadParameter", "Backup snapshot exceeds maximum version limit of 100."); + } + String latestVersion = bodyString(snapshot, "latestVersion", null); + + // Validate pass: every entry must be a well-formed JWK before anything is written, + // so a malformed later entry cannot leave orphaned version objects behind. + List validated = new ArrayList<>(); + for (Map entry : versions) { + String version = bodyString(entry, "version", null); + if (version == null || version.isEmpty()) { + continue; + } + @SuppressWarnings("unchecked") + Map data = entry.get("data") instanceof Map m + ? (Map) m : null; + if (data == null) { + continue; + } + Map validatedJwk = KeyVaultCrypto.importJwk(data); + if (data.get("hsm") != null) { + validatedJwk.put("hsm", data.get("hsm")); + } + if (data.get("tags") != null) { + validatedJwk.put("tags", data.get("tags")); + } + @SuppressWarnings("unchecked") + Map rawMeta = entry.get("meta") instanceof Map m + ? (Map) m : new LinkedHashMap<>(); + Map safeMeta = new HashMap<>(); + for (Map.Entry mEntry : rawMeta.entrySet()) { + if (mEntry.getKey() != null && mEntry.getValue() != null) { + safeMeta.put(String.valueOf(mEntry.getKey()), String.valueOf(mEntry.getValue())); + } + } + try { + normalizeAttrMeta(safeMeta, rawMeta); + } catch (IllegalArgumentException e) { + return kvError(400, "BadParameter", e.getMessage()); + } + validated.add(new RestoredVersion(version, validatedJwk, safeMeta)); + } + + if (validated.isEmpty()) { + return kvError(404, "KeyNotFound", "No key versions were restored for " + name + "."); + } + + String chosenLatest = latestVersion; + boolean latestFound = false; + for (RestoredVersion restored : validated) { + if (restored.version().equals(chosenLatest)) { + latestFound = true; + break; + } + } + if (!latestFound) { + return kvError(400, "BadParameter", "Backup snapshot latestVersion does not match any restored version."); + } + + // Write pass: only after every entry has validated successfully. + StoredObject latest = null; + for (RestoredVersion restored : validated) { + StoredObject versionObj = new StoredObject(keyVersionKey(account, name, restored.version(), hsm), + toBytes(restored.jwk()), restored.meta(), Instant.now(), newVersionId().substring(0, 16)); + store.put(versionObj.key(), versionObj); + if (restored.version().equals(chosenLatest)) { + latest = versionObj; + } + } + + Map latestMeta = new HashMap<>(latest.metadata()); + latestMeta.put("latestVersion", chosenLatest); + store.put(keyLatestKey(account, name, hsm), new StoredObject(keyLatestKey(account, name, hsm), + latest.data(), latestMeta, latest.lastModified(), latest.etag())); + + if (snapshot.get("rotationPolicy") instanceof Map policy) { + @SuppressWarnings("unchecked") + Map policyMap = (Map) policy; + long now = Instant.now().getEpochSecond(); + Map meta = new HashMap<>(); + meta.put("created", String.valueOf(now)); + meta.put("updated", String.valueOf(now)); + store.put(rotationPolicyKey(account, name, hsm), new StoredObject(rotationPolicyKey(account, name, hsm), + toBytes(policyMap), meta, Instant.now(), newVersionId().substring(0, 16))); + } + + return Response.ok(toJson(keyBundle(account, name, chosenLatest, + store.get(keyLatestKey(account, name, hsm)).get(), hsm)), "application/json").build(); + } catch (KeyVaultCrypto.CryptoException e) { + return kvError(400, "BadParameter", e.getMessage()); + } catch (Exception e) { + LOG.warnf("Key restore failed for %s: %s", name, sanitizeLogValue(e.getMessage())); + return kvError(400, "BadParameter", "Invalid backup snapshot."); + } + } + + // ── Rotation policy ──────────────────────────────────────────────────────── + + Response getRotationPolicy(String account, String name, boolean hsm) { + if (store.get(keyLatestKey(account, name, hsm)).isEmpty()) { + return keyNotFound(name); + } + Map policy = store.get(rotationPolicyKey(account, name, hsm)) + .map(this::parseStoredData).orElseGet(this::defaultPolicy); + return Response.ok(toJson(policyBundle(account, name, policy, hsm)), "application/json").build(); + } + + Response putRotationPolicy(AzureRequest req, String account, String name, boolean hsm) { + if (store.get(keyLatestKey(account, name, hsm)).isEmpty()) { + return keyNotFound(name); + } + Map body = parseBody(req); + Map policy = new LinkedHashMap<>(); + policy.put("lifetimeActions", body.getOrDefault("lifetimeActions", List.of())); + policy.put("attributes", body.getOrDefault("attributes", new LinkedHashMap<>())); + + long now = Instant.now().getEpochSecond(); + Map meta = new HashMap<>(); + meta.put("created", String.valueOf(now)); + meta.put("updated", String.valueOf(now)); + String key = rotationPolicyKey(account, name, hsm); + store.put(key, new StoredObject(key, toBytes(policy), meta, Instant.now(), + newVersionId().substring(0, 16))); + + return Response.ok(toJson(policyBundle(account, name, policy, hsm)), "application/json").build(); + } + + Response rotateKey(String account, String name, boolean hsm) { + Optional opt = store.get(keyLatestKey(account, name, hsm)); + if (opt.isEmpty()) { + return keyNotFound(name); + } + StoredObject current = opt.get(); + Map data = parseStoredData(current); + String kty = (String) data.get("kty"); + String crv = (String) data.get("crv"); + int keySize = data.get("keySize") instanceof Number n ? n.intValue() : 0; + @SuppressWarnings("unchecked") + List keyOps = data.get("key_ops") instanceof List l ? (List) l : null; + + Map newJwk; + try { + newJwk = KeyVaultCrypto.generateJwk(kty, keySize, crv, keyOps); + } catch (KeyVaultCrypto.CryptoException e) { + return kvError(400, "BadParameter", e.getMessage()); + } + if (data.get("hsm") != null) { + newJwk.put("hsm", data.get("hsm")); + } + if (data.get("tags") != null) { + newJwk.put("tags", data.get("tags")); + } + + long now = Instant.now().getEpochSecond(); + String versionId = newVersionId(); + + Map versionMeta = new HashMap<>(); + versionMeta.put("version", versionId); + versionMeta.put("enabled", current.metadata().getOrDefault("enabled", "true")); + versionMeta.put("nbf", current.metadata().getOrDefault("nbf", "")); + versionMeta.put("exp", current.metadata().getOrDefault("exp", "")); + versionMeta.put("created", String.valueOf(now)); + versionMeta.put("updated", String.valueOf(now)); + + StoredObject versionObj = new StoredObject(keyVersionKey(account, name, versionId, hsm), + toBytes(newJwk), versionMeta, Instant.now(), newVersionId().substring(0, 16)); + store.put(versionObj.key(), versionObj); + + Map latestMeta = new HashMap<>(versionMeta); + latestMeta.put("latestVersion", versionId); + store.put(keyLatestKey(account, name, hsm), new StoredObject(keyLatestKey(account, name, hsm), + toBytes(newJwk), latestMeta, Instant.now(), newVersionId().substring(0, 16))); + + return Response.ok(toJson(keyBundle(account, name, versionId, + store.get(keyVersionKey(account, name, versionId, hsm)).get(), hsm)), "application/json").build(); + } + + // ── Cryptographic operations ─────────────────────────────────────────────── + + Response cryptoOp(AzureRequest req, String account, String name, String version, String op, boolean hsm) { + StoredObject keyObj = resolveKeyForCrypto(account, name, version, hsm); + if (keyObj == null) { + return keyNotFound(name + (version != null && !version.isEmpty() ? "/" + version : "")); + } + if (!"true".equals(keyObj.metadata().getOrDefault("enabled", "true"))) { + return kvError(403, "Forbidden", "The key " + name + " is disabled."); + } + long nowEpoch = Instant.now().getEpochSecond(); + Long exp = parseLongNullable(keyObj.metadata().get("exp")); + if (exp != null && nowEpoch > exp) { + return kvError(403, "Forbidden", "The key " + name + " is expired."); + } + Long nbf = parseLongNullable(keyObj.metadata().get("nbf")); + if (nbf != null && nowEpoch < nbf) { + return kvError(403, "Forbidden", "The key " + name + " is not yet valid."); + } + + Map data = parseStoredData(keyObj); + + Map body = parseBody(req); + String alg = bodyString(body, "alg", null); + if (alg == null) { + return kvError(400, "BadParameter", "An algorithm is required."); + } + + // Key-type/algorithm mismatch (or unknown algorithm) is a 400, before the key_ops check. + String requiredKty = KeyVaultCrypto.keyTypeForAlg(alg); + if (requiredKty == null || !requiredKty.equals(KeyVaultCrypto.baseKty((String) data.get("kty")))) { + return kvError(400, "BadParameter", "Algorithm " + alg + " is not supported for this key."); + } + + String requiredOp = requiredKeyOp(op); + if (requiredOp != null && !permits(data.get("key_ops"), requiredOp)) { + return kvError(403, "Forbidden", "The operation " + op + " is not permitted on this key."); + } + + String resolvedVersion = version != null && !version.isEmpty() + ? version : keyObj.metadata().getOrDefault("latestVersion", keyObj.metadata().get("version")); + String kid = keyVersionId(account, name, resolvedVersion, hsm); + + try { + switch (op) { + case "encrypt", "wrapkey" -> { + byte[] value = decodeField(body, "value"); + if (value == null) { + return kvError(400, "BadParameter", "A value is required."); + } + byte[] iv = decodeField(body, "iv"); + byte[] aad = decodeField(body, "aad"); + KeyVaultCrypto.CipherResult result = KeyVaultCrypto.encrypt(data, alg, value, iv, aad); + Map resp = new LinkedHashMap<>(); + resp.put("kid", kid); + resp.put("value", KeyVaultCrypto.b64Url(result.value())); + if (result.iv() != null) { + resp.put("iv", KeyVaultCrypto.b64Url(result.iv())); + } + if (result.tag() != null) { + resp.put("tag", KeyVaultCrypto.b64Url(result.tag())); + } + return Response.ok(toJson(resp), "application/json").build(); + } + case "decrypt", "unwrapkey" -> { + byte[] value = decodeField(body, "value"); + if (value == null) { + return kvError(400, "BadParameter", "A value is required."); + } + byte[] iv = decodeField(body, "iv"); + byte[] aad = decodeField(body, "aad"); + byte[] tag = decodeField(body, "tag"); + byte[] plaintext = KeyVaultCrypto.decrypt(data, alg, value, iv, aad, tag); + Map resp = new LinkedHashMap<>(); + resp.put("kid", kid); + resp.put("value", KeyVaultCrypto.b64Url(plaintext)); + return Response.ok(toJson(resp), "application/json").build(); + } + case "sign" -> { + byte[] digest = decodeField(body, "value"); + if (digest == null) { + return kvError(400, "BadParameter", "A value is required."); + } + byte[] signature = KeyVaultCrypto.sign(data, alg, digest); + Map resp = new LinkedHashMap<>(); + resp.put("kid", kid); + resp.put("value", KeyVaultCrypto.b64Url(signature)); + return Response.ok(toJson(resp), "application/json").build(); + } + case "verify" -> { + byte[] digest = decodeField(body, "digest"); + if (digest == null) { + return kvError(400, "BadParameter", "A digest is required."); + } + byte[] signature = decodeField(body, "value"); + if (signature == null) { + return kvError(400, "BadParameter", "A value is required."); + } + boolean valid = KeyVaultCrypto.verify(data, alg, digest, signature); + // Strictly {"value": bool} — no kid field, or SDK VerifyResult deserialization breaks. + return Response.ok(toJson(Map.of("value", valid)), "application/json").build(); + } + default -> { + return kvError(400, "BadParameter", "Unsupported operation: " + op); + } + } + } catch (KeyVaultCrypto.CryptoException e) { + return kvError(400, "BadParameter", e.getMessage()); + } + } + + private StoredObject resolveKeyForCrypto(String account, String name, String version, boolean hsm) { + if (store.get(deletedKeyKey(account, name, hsm)).isPresent()) { + return null; + } + if (version != null && !version.isEmpty()) { + return store.get(keyVersionKey(account, name, version, hsm)).orElse(null); + } + return store.get(keyLatestKey(account, name, hsm)).orElse(null); + } + + private static String requiredKeyOp(String op) { + return switch (op) { + case "encrypt" -> "encrypt"; + case "decrypt" -> "decrypt"; + case "wrapkey" -> "wrapKey"; + case "unwrapkey" -> "unwrapKey"; + case "sign" -> "sign"; + case "verify" -> "verify"; + default -> null; + }; + } + + private static boolean permits(Object keyOps, String op) { + if (keyOps == null) { + return false; + } + if (keyOps instanceof List list) { + return list.contains(op); + } + return false; + } + + // ── Response builders ────────────────────────────────────────────────────── + + private Map keyBundle(String account, String name, String version, + StoredObject obj, boolean hsm) { + Map data = parseStoredData(obj); + Map key = KeyVaultCrypto.publicJwk(data); + key.put("kid", keyVersionId(account, name, version, hsm)); + Map bundle = new LinkedHashMap<>(); + bundle.put("key", key); + bundle.put("attributes", buildKeyAttributes(obj.metadata())); + bundle.put("tags", data.getOrDefault("tags", new LinkedHashMap<>())); + return bundle; + } + + private Map deletedKeyBundle(String account, String name, String version, + StoredObject obj, long deletedDate, long purgeDate, boolean hsm) { + Map bundle = keyBundle(account, name, version, obj, hsm); + bundle.put("recoveryId", recoveryId(account, name, hsm)); + bundle.put("deletedDate", deletedDate); + bundle.put("scheduledPurgeDate", purgeDate); + return bundle; + } + + private Map keyItem(String account, String idPath, StoredObject obj, + boolean hsm, boolean includeVersion) { + Map data = parseStoredData(obj); + Map item = new LinkedHashMap<>(); + item.put("kid", includeVersion ? keyVersionIdFromPath(account, idPath, hsm) + : keyBaseId(account, idPath, hsm)); + item.put("attributes", buildKeyAttributes(obj.metadata())); + item.put("tags", data.getOrDefault("tags", new LinkedHashMap<>())); + return item; + } + + private Map deletedKeyItem(String account, String name, String version, + StoredObject obj, long deletedDate, long purgeDate, boolean hsm) { + Map item = keyItem(account, name, obj, hsm, false); + item.put("recoveryId", recoveryId(account, name, hsm)); + item.put("deletedDate", deletedDate); + item.put("scheduledPurgeDate", purgeDate); + return item; + } + + private Map buildKeyAttributes(Map meta) { + Map attrs = new LinkedHashMap<>(); + attrs.put("enabled", Boolean.parseBoolean(meta.getOrDefault("enabled", "true"))); + String nbf = meta.get("nbf"); + attrs.put("nbf", (nbf != null && !nbf.isEmpty() && !"null".equals(nbf)) ? parseLong(nbf, 0L) : null); + String exp = meta.get("exp"); + attrs.put("exp", (exp != null && !exp.isEmpty() && !"null".equals(exp)) ? parseLong(exp, 0L) : null); + attrs.put("created", parseLong(meta.get("created"), 0L)); + attrs.put("updated", parseLong(meta.get("updated"), 0L)); + attrs.put("recoveryLevel", RECOVERY_LEVEL); + attrs.put("recoverableDays", 7); + return attrs; + } + + private Map defaultPolicy() { + Map attributes = new LinkedHashMap<>(); + attributes.put("expiryTime", null); + Map policy = new LinkedHashMap<>(); + policy.put("lifetimeActions", List.of()); + policy.put("attributes", attributes); + return policy; + } + + private Map policyBundle(String account, String name, Map policy, boolean hsm) { + Map out = new LinkedHashMap<>(); + out.put("id", keyBaseId(account, name, hsm) + "/rotationpolicy"); + out.put("lifetimeActions", policy.getOrDefault("lifetimeActions", List.of())); + @SuppressWarnings("unchecked") + Map attrs = new LinkedHashMap<>( + policy.get("attributes") instanceof Map m ? (Map) m : Map.of()); + attrs.putIfAbsent("created", Instant.now().getEpochSecond()); + attrs.putIfAbsent("updated", Instant.now().getEpochSecond()); + out.put("attributes", attrs); + return out; + } + + private Response jsonList(List> items) { + Map result = new LinkedHashMap<>(); + result.put("value", items); + result.put("nextLink", null); + return Response.ok(toJson(result), "application/json").build(); + } + + // ── Storage helpers ──────────────────────────────────────────────────────── + + // Managed HSM and Key Vault share the same handler, so storage keys are scoped per flavor to + // keep {account}-keyvault and {account}-managedhsm namespaces fully isolated. The non-HSM + // layout ({account}/keys/...) is unchanged from before; HSM keys live under {account}/hsm/. + + private String keyLatestKey(String account, String name, boolean hsm) { + return (hsm ? account + "/hsm/keys/" : account + "/keys/") + name; + } + + private String keyVersionKey(String account, String name, String version, boolean hsm) { + return keyLatestKey(account, name, hsm) + "/versions/" + version; + } + + private String deletedKeyKey(String account, String name, boolean hsm) { + return (hsm ? account + "/hsm/deletedkeys/" : account + "/deletedkeys/") + name; + } + + private String rotationPolicyKey(String account, String name, boolean hsm) { + return keyLatestKey(account, name, hsm) + "/rotationpolicy"; + } + + private String keysPrefix(String account, boolean hsm) { + return hsm ? account + "/hsm/keys/" : account + "/keys/"; + } + + private String keyVersionsPrefix(String account, String name, boolean hsm) { + return keyLatestKey(account, name, hsm) + "/versions/"; + } + + private String deletedKeysPrefix(String account, boolean hsm) { + return hsm ? account + "/hsm/deletedkeys/" : account + "/deletedkeys/"; + } + + // ── URL builders ─────────────────────────────────────────────────────────── + + private String vaultHost(String account, boolean hsm) { + return "https://" + account + (hsm ? ".managedhsm.azure.net" : ".vault.azure.net"); + } + + private String keyVersionId(String account, String name, String version, boolean hsm) { + return vaultHost(account, hsm) + "/keys/" + name + "/" + version; + } + + private String keyVersionIdFromPath(String account, String path, boolean hsm) { + return vaultHost(account, hsm) + "/keys/" + path; + } + + private String keyBaseId(String account, String name, boolean hsm) { + return vaultHost(account, hsm) + "/keys/" + name; + } + + private String recoveryId(String account, String name, boolean hsm) { + return vaultHost(account, hsm) + "/deletedkeys/" + name; + } + + // ── Errors ───────────────────────────────────────────────────────────────── + + private Response keyNotFound(String name) { + return kvError(404, "KeyNotFound", + "A key with (name/id) " + name + " was not found in this key vault. " + + "If you recently deleted this key, it may still be recoverable. " + + "For more information, see https://docs.microsoft.com/en-us/rest/api/keyvault/getkey."); + } + + private Response deletedKeyNotFound(String name) { + return kvError(404, "KeyNotFound", + "A deleted key with (name/id) " + name + " was not found in this key vault."); + } + + private Response kvError(int status, String code, String message) { + Map body = new LinkedHashMap<>(); + Map detail = new LinkedHashMap<>(); + detail.put("code", code); + detail.put("message", message); + body.put("error", detail); + return Response.status(status).entity(toJson(body)).type("application/json").build(); + } + + // ── Utilities ────────────────────────────────────────────────────────────── + + private String newVersionId() { + return UUID.randomUUID().toString().replace("-", ""); + } + + private byte[] decodeField(Map body, String field) { + Object value = body.get(field); + if (value == null) { + return null; + } + return KeyVaultCrypto.b64UrlDecode(String.valueOf(value)); + } + + @SuppressWarnings("unchecked") + private Map parseStoredData(StoredObject obj) { + try { + return MAPPER.readValue(obj.data(), new TypeReference<>() {}); + } catch (IOException e) { + return new LinkedHashMap<>(); + } + } + + private Map parseBody(AzureRequest req) { + return ArmJson.parseBodyMutable(req); + } + + private byte[] toBytes(Object obj) { + try { + return MAPPER.writeValueAsBytes(obj); + } catch (JsonProcessingException e) { + return new byte[0]; + } + } + + private String toJson(Object obj) { + try { + return MAPPER.writeValueAsString(obj); + } catch (JsonProcessingException e) { + return "{}"; + } + } + + private long parseLong(String val, long fallback) { + if (val == null || val.isEmpty()) { + return fallback; + } + try { + return Long.parseLong(val); + } catch (NumberFormatException e) { + return fallback; + } + } + + private Long parseLongNullable(String val) { + if (val == null || val.isEmpty() || "null".equals(val)) { + return null; + } + try { + return Long.parseLong(val); + } catch (NumberFormatException e) { + LOG.warnf("Ignoring malformed nbf/exp metadata value: %s", sanitizeLogValue(val)); + return null; + } + } + + /** Neutralizes CR/LF and caps length so attacker-controlled values cannot forge log lines. */ + private static String sanitizeLogValue(String value) { + if (value == null) { + return null; + } + String cleaned = value.replace("\r", "\\r").replace("\n", "\\n"); + return cleaned.length() > 200 ? cleaned.substring(0, 200) + "..." : cleaned; + } + + private static Long parseEpochAttribute(Map attrs, String field) { + Object value = attrs.get(field); + if (value == null) { + return null; + } + if (value instanceof Number n) { + return n.longValue(); + } + if (value instanceof String s) { + try { + return Long.parseLong(s.trim()); + } catch (NumberFormatException e) { + // fall through to the invalid-value error below + } + } + throw new IllegalArgumentException("The '" + field + "' attribute must be a numeric Unix epoch timestamp."); + } + + /** + * Sole writer of enabled/nbf/exp metadata: only fields present in {@code attrs} are written and + * always canonicalized, so partial updates and restore snapshots inherit absent fields. + * + * @throws IllegalArgumentException if nbf/exp is present but not a numeric epoch timestamp. + */ + private static void normalizeAttrMeta(Map meta, Map attrs) { + if (attrs == null) { + return; + } + if (attrs.containsKey("enabled")) { + meta.put("enabled", String.valueOf(Boolean.parseBoolean(String.valueOf(attrs.get("enabled"))))); + } + if (attrs.containsKey("nbf")) { + meta.put("nbf", normalizeEpoch(attrs, "nbf")); + } + if (attrs.containsKey("exp")) { + meta.put("exp", normalizeEpoch(attrs, "exp")); + } + } + + /** Canonicalizes one epoch attribute: null/"null"/blank → "" (unset); otherwise parsed (throws if invalid). */ + private static String normalizeEpoch(Map attrs, String field) { + Object value = attrs.get(field); + if (value == null || "null".equals(String.valueOf(value))) { + return ""; + } + if (value instanceof String s && s.trim().isEmpty()) { + return ""; + } + return String.valueOf(parseEpochAttribute(attrs, field)); + } + + private static String bodyString(Map map, String key, String defaultValue) { + Object v = map.get(key); + return v instanceof String s ? s : defaultValue; + } + + private static int bodyInt(Map map, String key, int defaultValue) { + Object v = map.get(key); + return v instanceof Number n ? n.intValue() : defaultValue; + } + + @SuppressWarnings("unchecked") + private static Map stringMap(Object value) { + if (value instanceof Map m) { + return (Map) m; + } + return new LinkedHashMap<>(); + } + + private Response validateKeyName(String name) { + if (name == null || !name.matches("^[a-zA-Z0-9-]+$")) { + return kvError(400, "BadParameter", "Invalid key name: " + name); + } + return null; + } + + private record RestoredVersion(String version, Map jwk, Map meta) {} +} diff --git a/src/main/resources/application.yml b/src/main/resources/application.yml index a8f9c59a..d9edf4de 100644 --- a/src/main/resources/application.yml +++ b/src/main/resources/application.yml @@ -33,6 +33,7 @@ quarkus: --initialize-at-run-time=io.floci.az.services.redis.RedisHandler, --initialize-at-run-time=io.floci.az.services.entra.TokenIssuer, --initialize-at-run-time=io.floci.az.services.entra.SigningKeyProvider, + --initialize-at-run-time=io.floci.az.services.keyvault.KeyVaultCrypto, --initialize-at-run-time=io.floci.az.services.eventgrid.EventGridService floci-az: diff --git a/src/test/java/io/floci/az/core/AzureRoutingFilterTest.java b/src/test/java/io/floci/az/core/AzureRoutingFilterTest.java index 7ec0da4a..5463148e 100644 --- a/src/test/java/io/floci/az/core/AzureRoutingFilterTest.java +++ b/src/test/java/io/floci/az/core/AzureRoutingFilterTest.java @@ -90,6 +90,45 @@ void keyvaultSuffixRoutesToKeyVault() { .header("WWW-Authenticate", containsString("Bearer")); } + @Test + void managedhsmSuffixRoutesToKeyVault() { + given().when().get("/devstoreaccount1-managedhsm/secrets/foo?api-version=7.4") + .then().statusCode(401) + .header("WWW-Authenticate", containsString("Bearer")); + } + + @Test + void managedhsmSuffixActivatesHsmFlavor() { + given().when().get("/devstoreaccount1-managedhsm/secrets/foo?api-version=7.4") + .then().statusCode(401) + .header("WWW-Authenticate", containsString("resource=\"https://managedhsm.azure.net\"")); + } + + @Test + void clientSuppliedAccountSuffixHeaderIsIgnored() { + // A client-crafted x-floci-account-suffix must not flip an ordinary vault request to HSM flavor. + given().header("x-floci-account-suffix", "-managedhsm") + .when().get("/devstoreaccount1-keyvault/secrets/foo?api-version=7.4") + .then().statusCode(401) + .header("WWW-Authenticate", containsString("resource=\"https://vault.azure.net\"")); + } + + @Test + void clientSuppliedSuffixHeaderIgnoredOnHostRoute() { + // The host-based route does NOT overwrite the suffix header, so only the strip protects it. + given().header("Host", "myvault.vault.azure.net") + .header("x-floci-account-suffix", "-managedhsm") + .when().get("/secrets/foo?api-version=7.4") + .then().statusCode(401) + .header("WWW-Authenticate", containsString("resource=\"https://vault.azure.net\"")); + } + + // NOTE: a mixed-case variant of the spoof header is intentionally not pinned here. RESTEasy's + // header map is already case-insensitive at the current Quarkus version, so a test asserting + // "X-Floci-Account-Suffix" is stripped cannot distinguish the case-insensitive strip loop in + // AzureRoutingFilter from a plain single-key remove. That loop is defense-in-depth (kept), and + // its behavior is already covered by the two suffix-spoofing tests above. + @Test void blobDefaultAccountRoutesToBlob() { given().when().get("/devstoreaccount1/?comp=list") @@ -185,7 +224,7 @@ private static void assertDisabledEcho(String path, String expectedServiceType) @Test void keyVaultCollectionsAtArmBaseRouteToKeyVault() { for (String collection : new String[] { - "secrets", "certificates", "keys", "deletedsecrets", "deletedcertificates", "deletedkeys"}) { + "secrets", "certificates", "keys", "deletedsecrets", "deletedcertificates", "deletedkeys", "rng"}) { assertKeyVaultChallenge("/" + collection + "?api-version=7.4"); assertKeyVaultChallenge("/" + collection + "/foo?api-version=7.4"); } diff --git a/src/test/java/io/floci/az/core/RoutingTableAssemblyTest.java b/src/test/java/io/floci/az/core/RoutingTableAssemblyTest.java index 5ab564cf..717c112e 100644 --- a/src/test/java/io/floci/az/core/RoutingTableAssemblyTest.java +++ b/src/test/java/io/floci/az/core/RoutingTableAssemblyTest.java @@ -35,6 +35,7 @@ class RoutingTableAssemblyTest { /** A4's HOST_ROUTES plus routes introduced by later services. */ private static final Set> GOLDEN_HOST_ROUTES = Set.of( Map.entry(".vault.azure.net", "keyvault"), + Map.entry(".managedhsm.azure.net", "keyvault"), Map.entry(".communication.azure.com", "email"), Map.entry(".blob.core.windows.net", "blob"), Map.entry(".dfs.core.windows.net", "blob"), @@ -58,6 +59,7 @@ class RoutingTableAssemblyTest { Map.entry("-functions", "functions"), Map.entry("-appconfig", "appconfig"), Map.entry("-keyvault", "keyvault"), + Map.entry("-managedhsm", "keyvault"), Map.entry("-eventgrid", "eventgrid"), Map.entry("-eventhub", "eventhub"), Map.entry("-sql", "sql"), diff --git a/src/test/java/io/floci/az/core/tls/TlsConfigSourceCertificateGenerationTest.java b/src/test/java/io/floci/az/core/tls/TlsConfigSourceCertificateGenerationTest.java index 3b8c3587..d74072de 100644 --- a/src/test/java/io/floci/az/core/tls/TlsConfigSourceCertificateGenerationTest.java +++ b/src/test/java/io/floci/az/core/tls/TlsConfigSourceCertificateGenerationTest.java @@ -8,8 +8,10 @@ import java.io.ByteArrayInputStream; import java.nio.file.Files; import java.nio.file.Path; +import java.security.cert.Certificate; import java.security.cert.CertificateFactory; import java.security.cert.X509Certificate; +import java.util.ArrayList; import java.util.Collection; import java.util.List; @@ -86,8 +88,53 @@ void certificateWithDefaultConfigHasDefaultSans() throws Exception { assertTrue(sans.contains("0.0.0.0")); assertTrue(sans.contains("host.docker.internal"), "SANs should include 'host.docker.internal' so function containers can reach floci-az on the host"); - assertEquals(8, sans.size(), - "Default cert should have exactly 8 SANs (localhost, 127.0.0.1, 0.0.0.0, *.localhost, localhost.floci-az.io, *.localhost.floci-az.io, *.vault.azure.net, host.docker.internal)"); + assertTrue(sans.contains("*.managedhsm.azure.net"), + "SANs should include '*.managedhsm.azure.net' for Managed-HSM-flavored vault URLs"); + assertEquals(9, sans.size(), + "Default cert should have exactly 9 SANs (localhost, 127.0.0.1, 0.0.0.0, *.localhost, localhost.floci-az.io, *.localhost.floci-az.io, *.vault.azure.net, *.managedhsm.azure.net, host.docker.internal)"); + } + + @Test + void certificateChainHasNonCaLeafSignedByCa() throws Exception { + new TlsConfigSource(); + + Path certFile = tempDir.resolve("tls/floci-az-selfsigned.crt"); + Path caFile = tempDir.resolve("tls/floci-az-selfsigned-ca.crt"); + + // The .crt file is a 2-certificate chain (leaf immediately followed by CA), so parse it + // with generateCertificates (plural), not generateCertificate. + String pem = Files.readString(certFile); + CertificateFactory cf = CertificateFactory.getInstance("X.509"); + Collection parsed = cf.generateCertificates( + new ByteArrayInputStream(pem.getBytes())); + List certs = new ArrayList<>(); + for (Certificate c : parsed) { + certs.add((X509Certificate) c); + } + assertEquals(2, certs.size(), "Cert file should contain a leaf+CA chain (2 certificates)"); + + X509Certificate leaf = certs.get(0); + X509Certificate ca = certs.get(1); + + // Leaf must be a genuine end-entity certificate: strict path validators (e.g. + // rustls-platform-verifier's CaUsedAsEndEntity check) reject a CA-flagged cert in the + // leaf position. + assertEquals(-1, leaf.getBasicConstraints(), "Leaf must not be a CA"); + assertNotEquals(leaf.getIssuerX500Principal(), leaf.getSubjectX500Principal(), + "Leaf must not be self-signed (issuer != subject)"); + + // Second certificate is the signing CA. + assertTrue(ca.getBasicConstraints() >= 0, "CA must be a CA (getBasicConstraints() >= 0)"); + assertEquals(ca.getSubjectX500Principal(), leaf.getIssuerX500Principal(), + "Leaf issuer should be the CA's subject"); + + // The leaf really is signed by the CA's key. + leaf.verify(ca.getPublicKey()); + + // The standalone CA file matches what GET /_floci/tls-cert serves as the trust anchor. + assertTrue(Files.exists(caFile), "Standalone CA file should be written alongside the chain"); + assertEquals(Files.readString(caFile), TlsConfigSource.currentCertPem, + "GET /_floci/tls-cert should serve the CA certificate, not the leaf"); } @Test diff --git a/src/test/java/io/floci/az/services/arm/ArmHandlerTest.java b/src/test/java/io/floci/az/services/arm/ArmHandlerTest.java new file mode 100644 index 00000000..5375727f --- /dev/null +++ b/src/test/java/io/floci/az/services/arm/ArmHandlerTest.java @@ -0,0 +1,54 @@ +package io.floci.az.services.arm; + +import io.quarkus.test.junit.QuarkusTest; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +import static io.restassured.RestAssured.given; +import static org.hamcrest.Matchers.equalTo; +import static org.hamcrest.Matchers.hasEntry; +import static org.hamcrest.Matchers.hasItem; + +@QuarkusTest +@DisplayName("ARM Managed HSM") +class ArmHandlerTest { + + @Test + @DisplayName("create, get by name, list by subscription/RG, and delete a Managed HSM") + void managedHsmCreateGetListDelete() { + String sub = "hsm-sub-1"; + String rg = "hsm-rg-1"; + String name = "hsm-1"; + String base = "/subscriptions/" + sub + "/resourceGroups/" + rg + + "/providers/Microsoft.KeyVault/managedHSMs/" + name; + + given().contentType("application/json") + .body("{\"location\":\"eastus\",\"sku\":{\"family\":\"B\",\"name\":\"Standard_B1\"}," + + "\"properties\":{\"initialAdminObjectIds\":[\"obj-1\"]}}") + .when().put(base + "?api-version=2023-07-01") + .then().statusCode(200) + .body("name", equalTo(name)) + .body("type", equalTo("Microsoft.KeyVault/managedHSMs")) + .body("properties.hsmUri", equalTo("https://" + name + ".managedhsm.azure.net/")) + .body("sku.name", equalTo("Standard_B1")); + + given().when().get(base + "?api-version=2023-07-01") + .then().statusCode(200) + .body("name", equalTo(name)); + + given().when().get("/subscriptions/" + sub + "/resourceGroups/" + rg + + "/providers/Microsoft.KeyVault/managedHSMs?api-version=2023-07-01") + .then().statusCode(200) + .body("value", hasItem(hasEntry("name", name))); + + given().when().get("/subscriptions/" + sub + "/providers/Microsoft.KeyVault/managedHSMs?api-version=2023-07-01") + .then().statusCode(200) + .body("value", hasItem(hasEntry("name", name))); + + given().when().delete(base + "?api-version=2023-07-01") + .then().statusCode(200); + + given().when().get(base + "?api-version=2023-07-01") + .then().statusCode(404); + } +} diff --git a/src/test/java/io/floci/az/services/keyvault/KeyVaultCryptoTest.java b/src/test/java/io/floci/az/services/keyvault/KeyVaultCryptoTest.java new file mode 100644 index 00000000..571f7c79 --- /dev/null +++ b/src/test/java/io/floci/az/services/keyvault/KeyVaultCryptoTest.java @@ -0,0 +1,533 @@ +package io.floci.az.services.keyvault; + +import io.quarkus.test.junit.QuarkusTest; +import io.restassured.http.ContentType; +import io.restassured.response.Response; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +import java.math.BigInteger; +import java.nio.charset.StandardCharsets; +import java.security.AlgorithmParameters; +import java.security.KeyFactory; +import java.security.MessageDigest; +import java.security.PublicKey; +import java.security.Signature; +import java.security.SecureRandom; +import java.security.spec.ECGenParameterSpec; +import java.security.spec.ECParameterSpec; +import java.security.spec.ECPoint; +import java.security.spec.ECPublicKeySpec; +import java.security.spec.MGF1ParameterSpec; +import java.security.spec.PSSParameterSpec; +import java.security.spec.RSAPublicKeySpec; +import java.util.Arrays; +import java.util.Base64; + +import javax.crypto.Cipher; +import javax.crypto.spec.OAEPParameterSpec; +import javax.crypto.spec.PSource; + +import static io.restassured.RestAssured.given; +import static org.hamcrest.Matchers.*; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertFalse; +import static org.junit.jupiter.api.Assertions.assertTrue; + +/** + * Quarkus-level crypto round-trip tests: RSA-OAEP(-256)/RSA1_5, AES-GCM, RS256/PS256/ES256 + * sign/verify (with independent JDK cross-checks), wrap/unwrap, and error mapping. + */ +@QuarkusTest +@DisplayName("KeyVaultCrypto — cryptographic operations") +class KeyVaultCryptoTest { + + private static final String BASE = "/devstoreaccount1-keyvault"; + private static final String API = "?api-version=7.4"; + private static final String AUTH = "Bearer fake"; + private static final Base64.Encoder URL = Base64.getUrlEncoder().withoutPadding(); + private static final Base64.Decoder URL_DECODER = Base64.getUrlDecoder(); + + @BeforeEach + void reset() { + given().post("/_admin/reset").then().statusCode(204); + } + + // ── RSA encrypt/decrypt ──────────────────────────────────────────────────── + + @Test + @DisplayName("RSA-OAEP encrypt/decrypt round-trip") + void rsaOaepRoundTrip() { + assertRsaRoundTrip("RSA-OAEP"); + } + + @Test + @DisplayName("RSA-OAEP-256 encrypt/decrypt round-trip") + void rsaOaep256RoundTrip() { + assertRsaRoundTrip("RSA-OAEP-256"); + } + + @Test + @DisplayName("RSA1_5 encrypt/decrypt round-trip") + void rsa15RoundTrip() { + assertRsaRoundTrip("RSA1_5"); + } + + private void assertRsaRoundTrip(String alg) { + Response created = createKey("rsa-" + alg.replace("_", "-"), "RSA", null, 2048); + String kid = created.jsonPath().getString("key.kid"); + String pt = b64Url("hello " + alg); + + String ct = given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"" + alg + "\",\"value\":\"" + pt + "\"}") + .when().post(kidPath(kid) + "/encrypt" + API) + .then().statusCode(200).extract().jsonPath().getString("value"); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"" + alg + "\",\"value\":\"" + ct + "\"}") + .when().post(kidPath(kid) + "/decrypt" + API) + .then().statusCode(200) + .body("value", equalTo(pt)); + } + + @Test + @DisplayName("RSA-OAEP-256 decrypt accepts ciphertext from an independent MGF1 SHA-256 encryptor") + void rsaOaep256InteropDecrypt() { + Response created = createKey("oaep256interop", "RSA", null, 2048); + String kid = created.jsonPath().getString("key.kid"); + PublicKey pub = rsaPublicKey(created.jsonPath().getString("key.n"), + created.jsonPath().getString("key.e")); + byte[] plaintext = "interop-oaep256".getBytes(StandardCharsets.UTF_8); + + String ct; + try { + Cipher enc = Cipher.getInstance("RSA/ECB/OAEPWithSHA-256AndMGF1Padding"); + enc.init(Cipher.ENCRYPT_MODE, pub, + new OAEPParameterSpec("SHA-256", "MGF1", MGF1ParameterSpec.SHA256, PSource.PSpecified.DEFAULT)); + ct = URL.encodeToString(enc.doFinal(plaintext)); + } catch (Exception e) { + throw new IllegalStateException(e); + } + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"RSA-OAEP-256\",\"value\":\"" + ct + "\"}") + .when().post(kidPath(kid) + "/decrypt" + API) + .then().statusCode(200) + .body("value", equalTo(b64Url(plaintext))); + } + + // ── AES-GCM ──────────────────────────────────────────────────────────────── + + @Test + @DisplayName("A256GCM encrypt/decrypt round-trip with explicit iv/aad/tag") + void aesGcmRoundTrip() { + Response created = createKey("gcm1", "oct", null, 256); + String kid = created.jsonPath().getString("key.kid"); + String pt = b64Url("secret-data"); + byte[] iv = new byte[12]; + new SecureRandom().nextBytes(iv); + String ivB64 = URL.encodeToString(iv); + String aad = b64Url("authenticated-header"); + + Response enc = given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"A256GCM\",\"value\":\"" + pt + "\",\"iv\":\"" + ivB64 + "\",\"aad\":\"" + aad + "\"}") + .when().post(kidPath(kid) + "/encrypt" + API); + enc.then().statusCode(200) + .body("tag", notNullValue()); + String serverIv = enc.jsonPath().getString("iv"); + assertTrue(serverIv != null); + assertEquals(12, URL_DECODER.decode(serverIv).length); + assertFalse(Arrays.equals(iv, URL_DECODER.decode(serverIv))); + String ct = enc.jsonPath().getString("value"); + String tag = enc.jsonPath().getString("tag"); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"A256GCM\",\"value\":\"" + ct + "\",\"iv\":\"" + serverIv + "\",\"aad\":\"" + aad + "\",\"tag\":\"" + tag + "\"}") + .when().post(kidPath(kid) + "/decrypt" + API) + .then().statusCode(200) + .body("value", equalTo(pt)); + } + + @Test + @DisplayName("AES-GCM ciphertext/tag tamper is rejected with 400") + void aesGcmTagTamperRejected() { + Response created = createKey("gcmtamper", "oct", null, 256); + String kid = created.jsonPath().getString("key.kid"); + String pt = b64Url("tamper-target"); + String aad = b64Url("aad"); + + Response enc = given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"A256GCM\",\"value\":\"" + pt + "\",\"aad\":\"" + aad + "\"}") + .when().post(kidPath(kid) + "/encrypt" + API); + enc.then().statusCode(200); + String iv = enc.jsonPath().getString("iv"); + String ct = enc.jsonPath().getString("value"); + String tag = enc.jsonPath().getString("tag"); + + byte[] ctBytes = URL_DECODER.decode(ct); + ctBytes[0] ^= 1; + String tamperedCt = URL.encodeToString(ctBytes); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"A256GCM\",\"value\":\"" + tamperedCt + "\",\"iv\":\"" + iv + "\",\"aad\":\"" + aad + "\",\"tag\":\"" + tag + "\"}") + .when().post(kidPath(kid) + "/decrypt" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + + byte[] tagBytes = URL_DECODER.decode(tag); + tagBytes[0] ^= 1; + String tamperedTag = URL.encodeToString(tagBytes); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"A256GCM\",\"value\":\"" + ct + "\",\"iv\":\"" + iv + "\",\"aad\":\"" + aad + "\",\"tag\":\"" + tamperedTag + "\"}") + .when().post(kidPath(kid) + "/decrypt" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + } + + @Test + @DisplayName("AES-GCM encrypt ignores caller IV and generates its own 12-byte IV") + void gcmEncryptIgnoresCallerIv() { + Response created = createKey("gcmiv", "oct", null, 256); + String kid = created.jsonPath().getString("key.kid"); + String pt = b64Url("ignore-my-iv"); + byte[] callerIv = new byte[12]; + new SecureRandom().nextBytes(callerIv); + + Response enc = given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"A256GCM\",\"value\":\"" + pt + "\",\"iv\":\"" + URL.encodeToString(callerIv) + "\"}") + .when().post(kidPath(kid) + "/encrypt" + API); + enc.then().statusCode(200); + + String serverIv = enc.jsonPath().getString("iv"); + assertTrue(serverIv != null); + assertEquals(12, URL_DECODER.decode(serverIv).length); + assertFalse(Arrays.equals(callerIv, URL_DECODER.decode(serverIv))); + } + + @Test + @DisplayName("AES-GCM algorithm/key-size mismatch returns 400 BadParameter") + void gcmAlgorithmKeySizeMismatch400() { + // A128GCM against a 256-bit key and A256GCM against a 128-bit key must both be rejected. + Response small = createKey("gcmsmall", "oct", null, 128); + Response large = createKey("gcmlarge", "oct", null, 256); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"A256GCM\",\"value\":\"" + b64Url("x") + "\"}") + .when().post(kidPath(small.jsonPath().getString("key.kid")) + "/encrypt" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"A128GCM\",\"value\":\"" + b64Url("x") + "\"}") + .when().post(kidPath(large.jsonPath().getString("key.kid")) + "/encrypt" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"A192GCM\",\"value\":\"" + b64Url("x") + "\"}") + .when().post(kidPath(large.jsonPath().getString("key.kid")) + "/encrypt" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + } + + // ── Sign / verify ────────────────────────────────────────────────────────── + @Test + @DisplayName("RS256 sign/verify; verify returns {\"value\":true} with no kid") + void rs256SignVerify() { + Response created = createKey("rs256", "RSA", null, 2048); + String kid = created.jsonPath().getString("key.kid"); + byte[] message = "hello world".getBytes(StandardCharsets.UTF_8); + String digest = b64Url(sha256(message)); + + String sig = sign(kid, "RS256", digest); + + // Verify via the service — strictly {"value": true} with no kid field. + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"RS256\",\"digest\":\"" + digest + "\",\"value\":\"" + sig + "\"}") + .when().post(kidPath(kid) + "/verify" + API) + .then().statusCode(200) + .body("value", equalTo(true)) + .body("kid", nullValue()); + + // Independent cross-check with the JDK. + PublicKey pub = rsaPublicKey(created.jsonPath().getString("key.n"), + created.jsonPath().getString("key.e")); + assertTrue(jdkRsaVerify("SHA256withRSA", pub, message, URL_DECODER.decode(sig)), + "JDK SHA256withRSA must accept the service signature"); + } + + @Test + @DisplayName("PS256 sign/verify (manual EMSA-PSS) cross-checked with the JDK") + void ps256SignVerify() { + Response created = createKey("ps256", "RSA", null, 2048); + String kid = created.jsonPath().getString("key.kid"); + byte[] message = "pss message".getBytes(StandardCharsets.UTF_8); + String digest = b64Url(sha256(message)); + + String sig = sign(kid, "PS256", digest); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"PS256\",\"digest\":\"" + digest + "\",\"value\":\"" + sig + "\"}") + .when().post(kidPath(kid) + "/verify" + API) + .then().statusCode(200) + .body("value", equalTo(true)) + .body("kid", nullValue()); + + PublicKey pub = rsaPublicKey(created.jsonPath().getString("key.n"), + created.jsonPath().getString("key.e")); + assertTrue(jdkPssVerify(pub, message, URL_DECODER.decode(sig)), + "JDK RSASSA-PSS must accept the service signature"); + } + + @Test + @DisplayName("ES256 sign/verify (NONEwithECDSA + DER↔raw) cross-checked with the JDK") + void es256SignVerify() { + Response created = createKey("es256", "EC", "P-256", 0); + String kid = created.jsonPath().getString("key.kid"); + byte[] message = "ecdsa message".getBytes(StandardCharsets.UTF_8); + String digest = b64Url(sha256(message)); + + String sig = sign(kid, "ES256", digest); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"ES256\",\"digest\":\"" + digest + "\",\"value\":\"" + sig + "\"}") + .when().post(kidPath(kid) + "/verify" + API) + .then().statusCode(200) + .body("value", equalTo(true)) + .body("kid", nullValue()); + + // Cross-check: raw R||S → DER → JDK SHA256withECDSA over the original message. + byte[] raw = URL_DECODER.decode(sig); + byte[] der = KeyVaultCrypto.rawRsToDer(raw, 32); + PublicKey pub = ecPublicKey(created.jsonPath().getString("key.crv"), + created.jsonPath().getString("key.x"), created.jsonPath().getString("key.y")); + try { + Signature jdk = Signature.getInstance("SHA256withECDSA"); + jdk.initVerify(pub); + jdk.update(message); + assertTrue(jdk.verify(der), "JDK SHA256withECDSA must accept the DER-transformed signature"); + } catch (Exception e) { + throw new IllegalStateException(e); + } + } + + @Test + @DisplayName("ES512 (P-521) sign/verify round-trip cross-checked with the JDK") + void es512SignVerifyRoundTrip() { + Response created = createKey("es512", "EC", "P-521", 0); + String kid = created.jsonPath().getString("key.kid"); + byte[] message = "es512 message".getBytes(StandardCharsets.UTF_8); + String digest = b64Url(sha512(message)); + + String sig = sign(kid, "ES512", digest); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"ES512\",\"digest\":\"" + digest + "\",\"value\":\"" + sig + "\"}") + .when().post(kidPath(kid) + "/verify" + API) + .then().statusCode(200) + .body("value", equalTo(true)) + .body("kid", nullValue()); + + byte[] der = KeyVaultCrypto.rawRsToDer(URL_DECODER.decode(sig), 66); + PublicKey pub = ecPublicKey(created.jsonPath().getString("key.crv"), + created.jsonPath().getString("key.x"), created.jsonPath().getString("key.y")); + try { + Signature jdk = Signature.getInstance("SHA512withECDSA"); + jdk.initVerify(pub); + jdk.update(message); + assertTrue(jdk.verify(der), "JDK SHA512withECDSA must accept the DER-transformed signature"); + } catch (Exception e) { + throw new IllegalStateException(e); + } + } + + // ── Wrap / unwrap ────────────────────────────────────────────────────────── + + @Test + @DisplayName("wrapKey/unwrapKey round-trip (RSA-OAEP-256)") + void wrapUnwrapRoundTrip() { + Response created = createKey("wrap1", "RSA", null, 2048); + String kid = created.jsonPath().getString("key.kid"); + String pt = b64Url("wrap-me"); + + String wrapped = given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"RSA-OAEP-256\",\"value\":\"" + pt + "\"}") + .when().post(kidPath(kid) + "/wrapkey" + API) + .then().statusCode(200).extract().jsonPath().getString("value"); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"RSA-OAEP-256\",\"value\":\"" + wrapped + "\"}") + .when().post(kidPath(kid) + "/unwrapkey" + API) + .then().statusCode(200) + .body("value", equalTo(pt)); + } + + // ── Error mapping ────────────────────────────────────────────────────────── + + @Test + @DisplayName("unknown algorithm returns 400 BadParameter") + void unknownAlg400() { + Response created = createKey("unknown", "RSA", null, 2048); + String kid = created.jsonPath().getString("key.kid"); + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"NOT-AN-ALG\",\"value\":\"" + b64Url("x") + "\"}") + .when().post(kidPath(kid) + "/sign" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + } + + @Test + @DisplayName("crypto on a disabled key returns 403 Forbidden") + void disabledKeyCrypto403() { + Response created = createKey("disabled", "RSA", null, 2048); + String kid = created.jsonPath().getString("key.kid"); + String version = kid.substring(kid.lastIndexOf('/') + 1); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"attributes\":{\"enabled\":false}}") + .when().patch(BASE + "/keys/disabled/" + version + API) + .then().statusCode(200); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"RSA-OAEP-256\",\"value\":\"" + b64Url("x") + "\"}") + .when().post(kidPath(kid) + "/encrypt" + API) + .then().statusCode(403) + .body("error.code", equalTo("Forbidden")); + } + + @Test + @DisplayName("RSA algorithm on an EC/oct key returns 400 BadParameter") + void signWithWrongKeyType400() { + Response created = createKey("wrongtype", "oct", null, 256); + String kid = created.jsonPath().getString("key.kid"); + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"RS256\",\"value\":\"" + b64Url("digest") + "\"}") + .when().post(kidPath(kid) + "/sign" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + } + + @Test + @DisplayName("wrong digest length returns 400 BadParameter") + void invalidDigestLengthRejected() { + String shortDigest = b64Url(new byte[]{1, 2, 3, 4}); + assertSignRejectsDigestLength("badlen-rs256", "RSA", null, 2048, "RS256", shortDigest); + assertSignRejectsDigestLength("badlen-ps256", "RSA", null, 2048, "PS256", shortDigest); + assertSignRejectsDigestLength("badlen-es256", "EC", "P-256", 0, "ES256", shortDigest); + } + + private void assertSignRejectsDigestLength(String name, String kty, String crv, int keySize, + String alg, String digest) { + Response created = createKey(name, kty, crv, keySize); + String kid = created.jsonPath().getString("key.kid"); + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"" + alg + "\",\"value\":\"" + digest + "\"}") + .when().post(kidPath(kid) + "/sign" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + } + + // ── Helpers ──────────────────────────────────────────────────────────────── + + private Response createKey(String name, String kty, String crv, int keySize) { + String body = "{\"kty\":\"" + kty + "\"" + + (crv != null ? ",\"crv\":\"" + crv + "\"" : "") + + (keySize > 0 ? ",\"key_size\":" + keySize : "") + + "}"; + Response created = given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body(body) + .when().post(BASE + "/keys/" + name + "/create" + API); + created.then().statusCode(200); + return created; + } + + private String sign(String kid, String alg, String digest) { + Response resp = given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"" + alg + "\",\"value\":\"" + digest + "\"}") + .when().post(kidPath(kid) + "/sign" + API); + resp.then().statusCode(200); + return resp.jsonPath().getString("value"); + } + + /** Rewrites an absolute {@code https://.../keys/{name}/{version}} kid to a local emulator path. */ + private static String kidPath(String kid) { + return BASE + kid.substring(kid.indexOf("/keys/")); + } + + private static String b64Url(byte[] bytes) { + return URL.encodeToString(bytes); + } + + private static String b64Url(String value) { + return b64Url(value.getBytes(StandardCharsets.UTF_8)); + } + + private static byte[] sha256(byte[] data) { + try { + return MessageDigest.getInstance("SHA-256").digest(data); + } catch (Exception e) { + throw new IllegalStateException(e); + } + } + + private static byte[] sha512(byte[] data) { + try { + return MessageDigest.getInstance("SHA-512").digest(data); + } catch (Exception e) { + throw new IllegalStateException(e); + } + } + + private static PublicKey rsaPublicKey(String nB64, String eB64) { + try { + BigInteger n = new BigInteger(1, URL_DECODER.decode(nB64)); + BigInteger e = new BigInteger(1, URL_DECODER.decode(eB64)); + return KeyFactory.getInstance("RSA").generatePublic(new RSAPublicKeySpec(n, e)); + } catch (Exception e) { + throw new IllegalStateException(e); + } + } + + private static boolean jdkRsaVerify(String alg, PublicKey pub, byte[] message, byte[] signature) { + try { + Signature sig = Signature.getInstance(alg); + sig.initVerify(pub); + sig.update(message); + return sig.verify(signature); + } catch (Exception e) { + throw new IllegalStateException(e); + } + } + + private static boolean jdkPssVerify(PublicKey pub, byte[] message, byte[] signature) { + try { + Signature sig = Signature.getInstance("RSASSA-PSS"); + sig.setParameter(new PSSParameterSpec("SHA-256", "MGF1", MGF1ParameterSpec.SHA256, 32, 1)); + sig.initVerify(pub); + sig.update(message); + return sig.verify(signature); + } catch (Exception e) { + throw new IllegalStateException(e); + } + } + + private static PublicKey ecPublicKey(String crv, String xB64, String yB64) { + try { + String curve = switch (crv) { + case "P-384" -> "secp384r1"; + case "P-521" -> "secp521r1"; + default -> "secp256r1"; + }; + AlgorithmParameters ap = AlgorithmParameters.getInstance("EC"); + ap.init(new ECGenParameterSpec(curve)); + ECParameterSpec params = ap.getParameterSpec(ECParameterSpec.class); + ECPoint w = new ECPoint(new BigInteger(1, URL_DECODER.decode(xB64)), + new BigInteger(1, URL_DECODER.decode(yB64))); + return KeyFactory.getInstance("EC").generatePublic(new ECPublicKeySpec(w, params)); + } catch (Exception e) { + throw new IllegalStateException(e); + } + } +} diff --git a/src/test/java/io/floci/az/services/keyvault/KeyVaultKeysTest.java b/src/test/java/io/floci/az/services/keyvault/KeyVaultKeysTest.java new file mode 100644 index 00000000..a639d84c --- /dev/null +++ b/src/test/java/io/floci/az/services/keyvault/KeyVaultKeysTest.java @@ -0,0 +1,1062 @@ +package io.floci.az.services.keyvault; + +import io.quarkus.test.junit.QuarkusTest; +import io.restassured.http.ContentType; +import io.restassured.response.Response; +import org.junit.jupiter.api.BeforeEach; +import org.junit.jupiter.api.DisplayName; +import org.junit.jupiter.api.Test; + +import java.math.BigInteger; +import java.security.KeyPair; +import java.security.KeyPairGenerator; +import java.security.interfaces.RSAPrivateCrtKey; +import java.security.interfaces.RSAPublicKey; +import java.util.ArrayList; +import java.util.Base64; +import java.util.LinkedHashMap; +import java.util.List; +import java.util.Map; + +import static io.restassured.RestAssured.given; +import static org.hamcrest.Matchers.*; +import static org.junit.jupiter.api.Assertions.assertEquals; +import static org.junit.jupiter.api.Assertions.assertNotEquals; +import static org.junit.jupiter.api.Assertions.assertTrue; + +/** + * Quarkus-level tests for the Key Vault keys data plane (CRUD, versions, soft-delete, + * backup/restore, rotation policy, rng, Managed HSM flavor detection, crypto version resolution). + */ +@QuarkusTest +@DisplayName("KeyVaultKeys — key CRUD and lifecycle") +class KeyVaultKeysTest { + + private static final String BASE = "/devstoreaccount1-keyvault"; + private static final String API = "?api-version=7.4"; + private static final String AUTH = "Bearer fake"; + + @BeforeEach + void reset() { + given().post("/_admin/reset").then().statusCode(204); + } + + // ── Creation ─────────────────────────────────────────────────────────────── + + @Test + @DisplayName("POST create RSA key returns public JWK without private fields") + void createRsaKeyReturnsJwk() { + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(BASE + "/keys/rsa1/create" + API) + .then().statusCode(200) + .body("key.kty", equalTo("RSA")) + .body("key.n", notNullValue()) + .body("key.e", notNullValue()) + .body("key.d", nullValue()) + .body("key.p", nullValue()) + .body("key.q", nullValue()) + .body("key.kid", containsString(".vault.azure.net/keys/rsa1/")); + } + + @Test + @DisplayName("POST create EC P-256 key returns crv/x/y") + void createEcKeyP256() { + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"EC\",\"crv\":\"P-256\"}") + .when().post(BASE + "/keys/ec1/create" + API) + .then().statusCode(200) + .body("key.kty", equalTo("EC")) + .body("key.crv", equalTo("P-256")) + .body("key.x", notNullValue()) + .body("key.y", notNullValue()) + .body("key.d", nullValue()); + } + + @Test + @DisplayName("POST create oct key strips symmetric key material") + void createOctKeyStripsK() { + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"oct\",\"key_size\":256}") + .when().post(BASE + "/keys/oct1/create" + API) + .then().statusCode(200) + .body("key.kty", equalTo("oct")) + .body("key.k", nullValue()); + } + + @Test + @DisplayName("PUT import returns the same public RSA fields") + void importRsaKey() { + Map jwk = fullRsaJwk(); + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"key\":" + toJson(jwk) + "}") + .when().put(BASE + "/keys/imported" + API) + .then().statusCode(200) + .body("key.kty", equalTo("RSA")) + .body("key.n", equalTo(jwk.get("n"))) + .body("key.e", equalTo(jwk.get("e"))) + .body("key.d", nullValue()); + } + + @Test + @DisplayName("create echoes key_ops, nbf and exp") + void createWithKeyOpsAndAttributes() { + long nbf = 1000000000L; + long exp = 2000000000L; + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048," + + "\"key_ops\":[\"encrypt\",\"decrypt\"]," + + "\"attributes\":{\"nbf\":" + nbf + ",\"exp\":" + exp + "}}") + .when().post(BASE + "/keys/ops1/create" + API) + .then().statusCode(200) + .body("key.key_ops", hasItems("encrypt", "decrypt")) + .body("attributes.nbf", equalTo((int) nbf)) + .body("attributes.exp", equalTo((int) exp)); + } + + // ── Read ─────────────────────────────────────────────────────────────────── + + @Test + @DisplayName("GET latest and version, list keys omits versions, list versions includes them") + void getKeyLatestAndVersion() { + Response created = given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(BASE + "/keys/g1/create" + API); + created.then().statusCode(200); + String kid = created.jsonPath().getString("key.kid"); + String version = kid.substring(kid.lastIndexOf('/') + 1); + + given().header("Authorization", AUTH) + .when().get(BASE + "/keys/g1" + API) + .then().statusCode(200).body("key.kid", equalTo(kid)); + + given().header("Authorization", AUTH) + .when().get(BASE + "/keys/g1/" + version + API) + .then().statusCode(200).body("key.kid", equalTo(kid)); + + given().header("Authorization", AUTH) + .when().get(BASE + "/keys" + API) + .then().statusCode(200) + .body("value.size()", greaterThanOrEqualTo(1)) + .body("value[0].kid", equalTo("https://devstoreaccount1.vault.azure.net/keys/g1")); + + given().header("Authorization", AUTH) + .when().get(BASE + "/keys/g1/versions" + API) + .then().statusCode(200) + .body("value.size()", equalTo(1)) + .body("value[0].kid", equalTo(kid)); + } + + @Test + @DisplayName("missing key returns 404 KeyNotFound") + void getMissingKey404() { + given().header("Authorization", AUTH) + .when().get(BASE + "/keys/nope" + API) + .then().statusCode(404) + .body("error.code", equalTo("KeyNotFound")); + } + + // ── Update / disable ─────────────────────────────────────────────────────── + + @Test + @DisplayName("PATCH attributes and tags; crypto on a disabled key is 403") + void patchKeyAttributesAndTags() { + Response created = given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(BASE + "/keys/p1/create" + API); + created.then().statusCode(200); + String version = created.jsonPath().getString("key.kid"); + version = version.substring(version.lastIndexOf('/') + 1); + + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"attributes\":{\"enabled\":false},\"tags\":{\"env\":\"test\"}}") + .when().patch(BASE + "/keys/p1/" + version + API) + .then().statusCode(200) + .body("attributes.enabled", equalTo(false)) + .body("tags.env", equalTo("test")); + + // Crypto on the disabled key must be forbidden. + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"alg\":\"RSA-OAEP-256\",\"value\":\"" + b64Url("hello") + "\"}") + .when().post(BASE + "/keys/p1/" + version + "/encrypt" + API) + .then().statusCode(403) + .body("error.code", equalTo("Forbidden")); + } + + // ── Soft delete / recover / purge ────────────────────────────────────────── + + @Test + @DisplayName("delete/recover/purge lifecycle; purge removes the rotation policy") + void deleteKeySoftDeleteRecoverPurge() { + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(BASE + "/keys/lc1/create" + API) + .then().statusCode(200); + + // Set a rotation policy. + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"lifetimeActions\":[{\"trigger\":{\"timeAfterCreate\":\"P30D\"},\"action\":{\"type\":\"Rotate\"}}],\"attributes\":{}}") + .when().put(BASE + "/keys/lc1/rotationpolicy" + API) + .then().statusCode(200) + .body("lifetimeActions.size()", equalTo(1)); + + // Soft delete. + given().header("Authorization", AUTH) + .when().delete(BASE + "/keys/lc1" + API) + .then().statusCode(200) + .body("recoveryId", containsString("/deletedkeys/lc1")); + + // Rotation policy on a soft-deleted key must 404. + given().header("Authorization", AUTH) + .when().get(BASE + "/keys/lc1/rotationpolicy" + API) + .then().statusCode(404) + .body("error.code", equalTo("KeyNotFound")); + + // Recover. + given().header("Authorization", AUTH) + .when().post(BASE + "/deletedkeys/lc1/recover" + API) + .then().statusCode(200) + .body("key.kid", containsString("/keys/lc1/")); + + // Delete again and purge. + given().header("Authorization", AUTH) + .when().delete(BASE + "/keys/lc1" + API) + .then().statusCode(200); + given().header("Authorization", AUTH) + .when().delete(BASE + "/deletedkeys/lc1" + API) + .then().statusCode(204); + + // Recreating the key must NOT inherit the purged rotation policy. + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(BASE + "/keys/lc1/create" + API) + .then().statusCode(200); + given().header("Authorization", AUTH) + .when().get(BASE + "/keys/lc1/rotationpolicy" + API) + .then().statusCode(200) + .body("lifetimeActions.size()", equalTo(0)); + } + + // ── Backup / restore ─────────────────────────────────────────────────────── + + @Test + @DisplayName("recreating a soft-deleted key name returns 409 until purged") + void recreateSoftDeletedKeyConflicts() { + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(BASE + "/keys/recreate1/create" + API) + .then().statusCode(200); + + given().header("Authorization", AUTH) + .when().delete(BASE + "/keys/recreate1" + API) + .then().statusCode(200); + + // Create over the soft-deleted name must conflict. + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(BASE + "/keys/recreate1/create" + API) + .then().statusCode(409) + .body("error.code", equalTo("Conflict")); + + // Purge frees the name for recreation. + given().header("Authorization", AUTH) + .when().delete(BASE + "/deletedkeys/recreate1" + API) + .then().statusCode(204); + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(BASE + "/keys/recreate1/create" + API) + .then().statusCode(200); + } + + @Test + @DisplayName("backup → restore round-trip; restore over existing is 409") + void backupRestoreRoundTrip() { + Response created = given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(BASE + "/keys/bk1/create" + API); + created.then().statusCode(200); + String kid = created.jsonPath().getString("key.kid"); + String version = kid.substring(kid.lastIndexOf('/') + 1); + + Response backup = given().header("Authorization", AUTH) + .when().post(BASE + "/keys/bk1/backup" + API); + backup.then().statusCode(200); + String backupValue = backup.jsonPath().getString("value"); + assertTrue(backupValue != null && !backupValue.isEmpty()); + + // Restore over an existing key → 409 Conflict. + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"value\":\"" + backupValue + "\"}") + .when().post(BASE + "/keys/restore" + API) + .then().statusCode(409) + .body("error.code", equalTo("Conflict")); + + // Delete + purge, then restore. + given().header("Authorization", AUTH).when().delete(BASE + "/keys/bk1" + API).then().statusCode(200); + given().header("Authorization", AUTH).when().delete(BASE + "/deletedkeys/bk1" + API).then().statusCode(204); + + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"value\":\"" + backupValue + "\"}") + .when().post(BASE + "/keys/restore" + API) + .then().statusCode(200) + .body("key.kid", equalTo(kid)); + + given().header("Authorization", AUTH) + .when().get(BASE + "/keys/bk1/" + version + API) + .then().statusCode(200) + .body("key.kid", equalTo(kid)); + } + + // ── Rotation policy ──────────────────────────────────────────────────────── + + @Test + @DisplayName("rotation policy PUT/GET round-trip; unset returns default") + void rotationPolicyPutGet() { + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(BASE + "/keys/rp1/create" + API) + .then().statusCode(200); + + // Unset → default (empty lifetimeActions). + given().header("Authorization", AUTH) + .when().get(BASE + "/keys/rp1/rotationpolicy" + API) + .then().statusCode(200) + .body("lifetimeActions.size()", equalTo(0)) + .body("id", equalTo("https://devstoreaccount1.vault.azure.net/keys/rp1/rotationpolicy")); + + String policy = "{\"lifetimeActions\":[" + + "{\"trigger\":{\"timeAfterCreate\":\"P30D\"},\"action\":{\"type\":\"Rotate\"}}," + + "{\"trigger\":{\"timeBeforeExpiry\":\"P7D\"},\"action\":{\"type\":\"Notify\"}}]," + + "\"attributes\":{\"expiryTime\":\"P1Y\"}}"; + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body(policy) + .when().put(BASE + "/keys/rp1/rotationpolicy" + API) + .then().statusCode(200) + .body("lifetimeActions.size()", equalTo(2)); + + given().header("Authorization", AUTH) + .when().get(BASE + "/keys/rp1/rotationpolicy" + API) + .then().statusCode(200) + .body("lifetimeActions.size()", equalTo(2)) + .body("lifetimeActions[0].action.type", equalTo("Rotate")) + .body("attributes.expiryTime", equalTo("P1Y")); + } + + // ── rng ──────────────────────────────────────────────────────────────────── + + @Test + @DisplayName("POST /rng returns the requested number of random bytes") + void rngReturnsRequestedCount() { + for (String path : new String[]{BASE + "/rng", "/rng"}) { + Response resp = given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"count\":64}") + .when().post(path + API); + resp.then().statusCode(200); + byte[] bytes = Base64.getUrlDecoder().decode(resp.jsonPath().getString("value")); + assertEquals(64, bytes.length, "rng byte count via " + path); + } + } + + // ── Managed HSM flavor detection ─────────────────────────────────────────── + + @Test + @DisplayName("Host with port still detects managed HSM flavor for kid URLs") + void managedHsmKidHostWithPort() { + given().header("Authorization", AUTH) + .header("Host", "hsm1.managedhsm.azure.net:4577") + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post("/keys/hsmkey/create" + API) + .then().statusCode(200) + .body("key.kid", startsWith("https://hsm1.managedhsm.azure.net/keys/hsmkey/")); + } + + @Test + @DisplayName("challenge resource and root probe are flavor-aware") + void unauthorizedChallengeAndRootProbeFlavor() { + given().when().get(BASE + "/secrets/foo" + API) + .then().statusCode(401) + .header("WWW-Authenticate", containsString("resource=\"https://vault.azure.net\"")); + + given().header("Host", "hsm1.managedhsm.azure.net:4577") + .when().get("/secrets/foo" + API) + .then().statusCode(401) + .header("WWW-Authenticate", containsString("resource=\"https://managedhsm.azure.net\"")); + + given().header("Authorization", AUTH) + .when().get(BASE + API) + .then().statusCode(200) + .body("type", equalTo("Microsoft.KeyVault/vaults")) + .body("id", equalTo("https://devstoreaccount1.vault.azure.net/")); + + given().header("Authorization", AUTH) + .header("Host", "hsm1.managedhsm.azure.net:4577") + .when().get("/" + API) + .then().statusCode(200) + .body("type", equalTo("Microsoft.KeyVault/managedHSMs")) + .body("id", equalTo("https://hsm1.managedhsm.azure.net/")); + } + + // ── Managed HSM / Key Vault isolation ────────────────────────────────────── + + @Test + @DisplayName("Key Vault and Managed HSM key namespaces are isolated") + void vaultAndManagedHsmNamespacesIsolated() { + String mhsmBase = "/devstoreaccount1-managedhsm"; + + // A key created in the vault must not be visible through the Managed HSM flavor. + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(BASE + "/keys/isov/create" + API) + .then().statusCode(200); + given().header("Authorization", AUTH) + .when().get(mhsmBase + "/keys/isov" + API) + .then().statusCode(404) + .body("error.code", equalTo("KeyNotFound")); + given().header("Authorization", AUTH) + .when().get(mhsmBase + "/keys" + API) + .then().statusCode(200) + .body("value.size()", equalTo(0)); + + // And a key created in the HSM must not be visible through the vault flavor. + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(mhsmBase + "/keys/isoh/create" + API) + .then().statusCode(200) + .body("key.kid", containsString("managedhsm.azure.net/keys/isoh/")); + given().header("Authorization", AUTH) + .when().get(BASE + "/keys/isoh" + API) + .then().statusCode(404) + .body("error.code", equalTo("KeyNotFound")); + } + + // ── Crypto version resolution ────────────────────────────────────────────── @Test + @DisplayName("crypto ops with empty or omitted version resolve to the latest version") + void cryptoOpWithEmptyVersionResolvesLatest() { + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(BASE + "/keys/vr1/create" + API) + .then().statusCode(200); + + String pt = b64Url("resolve-latest"); + for (String opPath : new String[]{"/keys/vr1//encrypt", "/keys/vr1/encrypt"}) { + Response enc = given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"alg\":\"RSA-OAEP-256\",\"value\":\"" + pt + "\"}") + .when().post(BASE + opPath + API); + enc.then().statusCode(200); + String ct = enc.jsonPath().getString("value"); + + Response dec = given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"alg\":\"RSA-OAEP-256\",\"value\":\"" + ct + "\"}") + .when().post(BASE + "/keys/vr1//decrypt" + API); + dec.then().statusCode(200); + assertEquals(pt, dec.jsonPath().getString("value"), "round-trip via " + opPath); + } + } + + // ── Soft-delete tombstoning ───────────────────────────────────────────────── + + @Test + @DisplayName("soft-deleted key rejects version GET and crypto with 404") + void softDeletedKeyRejectsVersionAccess() { + Response created = given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(BASE + "/keys/sd1/create" + API); + created.then().statusCode(200); + String kid = created.jsonPath().getString("key.kid"); + String version = kid.substring(kid.lastIndexOf('/') + 1); + + given().header("Authorization", AUTH) + .when().delete(BASE + "/keys/sd1" + API) + .then().statusCode(200); + + given().header("Authorization", AUTH) + .when().get(BASE + "/keys/sd1/" + version + API) + .then().statusCode(404) + .body("error.code", equalTo("KeyNotFound")); + + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"alg\":\"RSA-OAEP-256\",\"value\":\"" + b64Url("hello") + "\"}") + .when().post(BASE + "/keys/sd1/" + version + "/encrypt" + API) + .then().statusCode(404) + .body("error.code", equalTo("KeyNotFound")); + + given().header("Authorization", AUTH) + .when().get(BASE + "/keys/sd1/versions" + API) + .then().statusCode(404) + .body("error.code", equalTo("KeyNotFound")); + + given().header("Authorization", AUTH) + .when().post(BASE + "/deletedkeys/sd1/recover" + API) + .then().statusCode(200); + } + + @Test + @DisplayName("expired and not-yet-valid keys return 403 on crypto") + void expiredAndNotYetValidKeysForbidCrypto() { + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048,\"attributes\":{\"exp\":1000000000}}") + .when().post(BASE + "/keys/exp1/create" + API) + .then().statusCode(200); + + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"alg\":\"RS256\",\"value\":\"" + b64Url("digest") + "\"}") + .when().post(BASE + "/keys/exp1/sign" + API) + .then().statusCode(403) + .body("error.code", equalTo("Forbidden")); + + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048,\"attributes\":{\"nbf\":9999999999}}") + .when().post(BASE + "/keys/nbf1/create" + API) + .then().statusCode(200); + + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"alg\":\"RS256\",\"value\":\"" + b64Url("digest") + "\"}") + .when().post(BASE + "/keys/nbf1/sign" + API) + .then().statusCode(403) + .body("error.code", equalTo("Forbidden")); + } + + @Test + @DisplayName("malformed/attacker-crafted restore is rejected with 400") + void restoreRejectsMalformedBackup() { + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{}") + .when().post(BASE + "/keys/restore" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"value\":\"!!!\"}") + .when().post(BASE + "/keys/restore" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + + Map badName = new LinkedHashMap<>(); + badName.put("name", "bad/name"); + badName.put("versions", new ArrayList<>()); + String badNameValue = URL.encodeToString( + toJson(badName).getBytes(java.nio.charset.StandardCharsets.UTF_8)); + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"value\":\"" + badNameValue + "\"}") + .when().post(BASE + "/keys/restore" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + + Map badJwkData = new LinkedHashMap<>(); + badJwkData.put("kty", "RSA"); + badJwkData.put("n", "not-base64!"); + Map badJwkVersion = new LinkedHashMap<>(); + badJwkVersion.put("version", "v1"); + badJwkVersion.put("data", badJwkData); + List> badJwkVersions = new ArrayList<>(); + badJwkVersions.add(badJwkVersion); + Map badJwk = new LinkedHashMap<>(); + badJwk.put("name", "badkey1"); + badJwk.put("versions", badJwkVersions); + String badJwkValue = URL.encodeToString( + toJson(badJwk).getBytes(java.nio.charset.StandardCharsets.UTF_8)); + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"value\":\"" + badJwkValue + "\"}") + .when().post(BASE + "/keys/restore" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + + Map full = fullRsaJwk(); + List> manyVersions = new ArrayList<>(); + for (int i = 0; i < 101; i++) { + Map entry = new LinkedHashMap<>(); + entry.put("version", "v" + i); + Map data = new LinkedHashMap<>(); + data.put("kty", "RSA"); + data.put("n", full.get("n")); + data.put("e", full.get("e")); + entry.put("data", data); + manyVersions.add(entry); + } + Map tooMany = new LinkedHashMap<>(); + tooMany.put("name", "badkey2"); + tooMany.put("versions", manyVersions); + String tooManyValue = URL.encodeToString( + toJson(tooMany).getBytes(java.nio.charset.StandardCharsets.UTF_8)); + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"value\":\"" + tooManyValue + "\"}") + .when().post(BASE + "/keys/restore" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + } + + @Test + @DisplayName("rotateKey generates fresh key material") + void rotateKeyGeneratesFreshMaterial() { + Response created = given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(BASE + "/keys/rot1/create" + API); + created.then().statusCode(200); + String v1 = created.jsonPath().getString("key.n"); + String v1Kid = created.jsonPath().getString("key.kid"); + String v1Version = v1Kid.substring(v1Kid.lastIndexOf('/') + 1); + + Response rotated = given().header("Authorization", AUTH) + .when().post(BASE + "/keys/rot1/rotate" + API); + rotated.then().statusCode(200); + String v2 = rotated.jsonPath().getString("key.n"); + String v2Version = rotated.jsonPath().getString("key.kid"); + v2Version = v2Version.substring(v2Version.lastIndexOf('/') + 1); + + assertNotEquals(v1, v2, "rotated key must have fresh modulus"); + assertNotEquals(v1Version, v2Version, "rotated key must get a new version id"); + assertTrue(rotated.jsonPath().getLong("attributes.created") > 0, "rotated version needs a created timestamp"); + assertTrue(rotated.jsonPath().getLong("attributes.updated") > 0, "rotated version needs an updated timestamp"); + } + + @Test + @DisplayName("key_ops [] forbids all operations") + void emptyKeyOpsForbidsAllOperations() { + Map jwk = fullRsaJwk(); + jwk.put("key_ops", List.of()); + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"key\":" + toJson(jwk) + "}") + .when().put(BASE + "/keys/emptyops" + API) + .then().statusCode(200); + + given().header("Authorization", AUTH) + .contentType(ContentType.JSON) + .body("{\"alg\":\"RS256\",\"value\":\"" + b64Url("digest") + "\"}") + .when().post(BASE + "/keys/emptyops/sign" + API) + .then().statusCode(403) + .body("error.code", equalTo("Forbidden")); + } + + // ── Regression: security/robustness fixes ────────────────────────────────── + + @Test + @DisplayName("non-numeric nbf/exp attributes are rejected with 400") + void nonNumericAttributesRejected() { + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048,\"attributes\":{\"exp\":\"not-a-number\"}}") + .when().post(BASE + "/keys/nn1/create" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + + Response created = given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(BASE + "/keys/nn2/create" + API); + created.then().statusCode(200); + String kid = created.jsonPath().getString("key.kid"); + String version = kid.substring(kid.lastIndexOf('/') + 1); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"attributes\":{\"nbf\":\"not-a-number\"}}") + .when().patch(BASE + "/keys/nn2/" + version + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + } + + @Test + @DisplayName("numeric-string exp is enforced, not silently ignored") + void numericStringExpIsEnforced() { + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048,\"attributes\":{\"exp\":\"1000000000\"}}") + .when().post(BASE + "/keys/nn3/create" + API) + .then().statusCode(200); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"RS256\",\"value\":\"" + b64Url("digest") + "\"}") + .when().post(BASE + "/keys/nn3/sign" + API) + .then().statusCode(403) + .body("error.code", equalTo("Forbidden")); + } + + @Test + @DisplayName("non-string kty on import is rejected with 400") + void importNonStringKtyRejected() { + Map jwk = fullRsaJwk(); + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"key\":{\"kty\":42,\"n\":\"" + jwk.get("n") + "\",\"e\":\"" + jwk.get("e") + "\"}}") + .when().put(BASE + "/keys/badkty" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + } + + @Test + @DisplayName("restore with a partially invalid backup leaves no orphaned key") + void restorePartialFailureLeavesNoOrphans() { + Map full = fullRsaJwk(); + Map validData = new LinkedHashMap<>(); + validData.put("kty", "RSA"); + validData.put("n", full.get("n")); + validData.put("e", full.get("e")); + Map validVersion = new LinkedHashMap<>(); + validVersion.put("version", "v1"); + validVersion.put("data", validData); + + Map invalidData = new LinkedHashMap<>(); + invalidData.put("kty", "RSA"); + Map invalidVersion = new LinkedHashMap<>(); + invalidVersion.put("version", "v2"); + invalidVersion.put("data", invalidData); + + List> versions = new ArrayList<>(); + versions.add(validVersion); + versions.add(invalidVersion); + Map snapshot = new LinkedHashMap<>(); + snapshot.put("name", "orphan1"); + snapshot.put("versions", versions); + String value = URL.encodeToString(toJson(snapshot).getBytes(java.nio.charset.StandardCharsets.UTF_8)); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"value\":\"" + value + "\"}") + .when().post(BASE + "/keys/restore" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + + given().header("Authorization", AUTH) + .when().get(BASE + "/keys/orphan1/v1" + API) + .then().statusCode(404) + .body("error.code", equalTo("KeyNotFound")); + } + + @Test + @DisplayName("key_ops allow-list permits only the listed operations") + void keyOpsAllowList() { + Map jwk = fullRsaJwk(); + jwk.put("key_ops", List.of("encrypt")); + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"key\":" + toJson(jwk) + "}") + .when().put(BASE + "/keys/allowops" + API) + .then().statusCode(200); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"RSA-OAEP-256\",\"value\":\"" + b64Url("hello") + "\"}") + .when().post(BASE + "/keys/allowops/encrypt" + API) + .then().statusCode(200); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"RS256\",\"value\":\"" + b64Url("digest") + "\"}") + .when().post(BASE + "/keys/allowops/sign" + API) + .then().statusCode(403) + .body("error.code", equalTo("Forbidden")); + } + + @Test + @DisplayName("invalid key names (containing /) are rejected with 400") + void invalidKeyNameRejected() { + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(BASE + "/keys/bad/name/create" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + } + + @Test + @DisplayName("rng count outside 1..128 is rejected with 400") + void rngCountBounds() { + for (int count : new int[]{0, -5, 129, 1024}) { + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"count\":" + count + "}") + .when().post(BASE + "/rng" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + } + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"count\":128}") + .when().post(BASE + "/rng" + API) + .then().statusCode(200); + } + + @Test + @DisplayName("rotating a key with an unsupported stored size returns 400") + void rotateUnsupportedKeySize400() { + byte[] k = new byte[20]; + new java.security.SecureRandom().nextBytes(k); + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"key\":{\"kty\":\"oct\",\"k\":\"" + URL.encodeToString(k) + "\"}}") + .when().put(BASE + "/keys/weirdoct" + API) + .then().statusCode(200); + + given().header("Authorization", AUTH) + .when().post(BASE + "/keys/weirdoct/rotate" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + } + + @Test + @DisplayName("non-string private JWK fields on import are rejected with 400") + void importNonStringPrivateFieldRejected() { + // RSA d + Map rsaD = fullRsaJwk(); + rsaD.put("d", 123); + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"key\":" + toJson(rsaD) + "}") + .when().put(BASE + "/keys/badd" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + + // RSA p (a second RSA private field, distinct code path from d) + Map rsaP = fullRsaJwk(); + rsaP.put("p", 123); + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"key\":" + toJson(rsaP) + "}") + .when().put(BASE + "/keys/badp" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + + // EC d — harvest a valid EC public JWK from the emulator, then corrupt the private field. + Response ecResp = given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"kty\":\"EC\",\"crv\":\"P-256\"}") + .when().post(BASE + "/keys/ecsrc/create" + API); + ecResp.then().statusCode(200); + Map ecJwk = new LinkedHashMap<>(); + ecJwk.put("kty", "EC"); + ecJwk.put("crv", ecResp.jsonPath().getString("key.crv")); + ecJwk.put("x", ecResp.jsonPath().getString("key.x")); + ecJwk.put("y", ecResp.jsonPath().getString("key.y")); + ecJwk.put("d", 123); + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"key\":" + toJson(ecJwk) + "}") + .when().put(BASE + "/keys/badecd" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + } + + @Test + @DisplayName("restore normalizes whitespace-padded exp; expired restored key is 403, not bypassed") + void restoreNormalizesWhitespaceExp() { + Map full = fullRsaJwk(); + Map data = new LinkedHashMap<>(); + data.put("kty", "RSA"); + data.put("n", full.get("n")); + data.put("e", full.get("e")); + data.put("d", full.get("d")); + data.put("p", full.get("p")); + data.put("q", full.get("q")); + data.put("dp", full.get("dp")); + data.put("dq", full.get("dq")); + data.put("qi", full.get("qi")); + + Map meta = new LinkedHashMap<>(); + meta.put("exp", " 1000000000 "); // long in the past; only enforced if trimmed, not if dropped + + Map version = new LinkedHashMap<>(); + version.put("version", "v1"); + version.put("data", data); + version.put("meta", meta); + List> versions = new ArrayList<>(); + versions.add(version); + Map snapshot = new LinkedHashMap<>(); + snapshot.put("name", "wsexp"); + snapshot.put("latestVersion", "v1"); + snapshot.put("versions", versions); + String value = URL.encodeToString(toJson(snapshot).getBytes(java.nio.charset.StandardCharsets.UTF_8)); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"value\":\"" + value + "\"}") + .when().post(BASE + "/keys/restore" + API) + .then().statusCode(200); + + // If exp were dropped (unset), this encrypt would 200; trimmed-and-enforced it must 403. + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"RSA-OAEP-256\",\"value\":\"" + b64Url("hello") + "\"}") + .when().post(BASE + "/keys/wsexp/encrypt" + API) + .then().statusCode(403) + .body("error.code", equalTo("Forbidden")); + } + + @Test + @DisplayName("restore with a latestVersion absent from versions is rejected with 400") + void restoreLatestVersionMismatch400() { + Map full = fullRsaJwk(); + Map data = new LinkedHashMap<>(); + data.put("kty", "RSA"); + data.put("n", full.get("n")); + data.put("e", full.get("e")); + Map version = new LinkedHashMap<>(); + version.put("version", "v1"); + version.put("data", data); + List> versions = new ArrayList<>(); + versions.add(version); + Map snapshot = new LinkedHashMap<>(); + snapshot.put("name", "mismatch"); + snapshot.put("latestVersion", "v999"); + snapshot.put("versions", versions); + String value = URL.encodeToString(toJson(snapshot).getBytes(java.nio.charset.StandardCharsets.UTF_8)); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"value\":\"" + value + "\"}") + .when().post(BASE + "/keys/restore" + API) + .then().statusCode(400) + .body("error.code", equalTo("BadParameter")); + } + + @Test + @DisplayName("PATCH enabled=\"TRUE\" normalizes to true and leaves the key operable") + void patchEnabledUppercaseNormalized() { + Response created = given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048}") + .when().post(BASE + "/keys/enup/create" + API); + created.then().statusCode(200); + String kid = created.jsonPath().getString("key.kid"); + String version = kid.substring(kid.lastIndexOf('/') + 1); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"attributes\":{\"enabled\":\"TRUE\"}}") + .when().patch(BASE + "/keys/enup/" + version + API) + .then().statusCode(200) + .body("attributes.enabled", equalTo(true)); + + given().header("Authorization", AUTH) + .when().get(BASE + "/keys/enup" + API) + .then().statusCode(200) + .body("attributes.enabled", equalTo(true)); + + // The normalized value must satisfy the strict reader too: crypto op succeeds, not 403. + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"RS256\",\"value\":\"" + b64Url("0123456789abcdef0123456789abcdef") + "\"}") + .when().post(BASE + "/keys/enup/sign" + API) + .then().statusCode(200); + } + + @Test + @DisplayName("rotating a key with deny-all key_ops preserves the restriction") + void rotatePreservesKeyOpsRestriction() { + Map jwk = fullRsaJwk(); + jwk.put("key_ops", List.of()); + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"key\":" + toJson(jwk) + "}") + .when().put(BASE + "/keys/rotops" + API) + .then().statusCode(200); + + given().header("Authorization", AUTH) + .when().post(BASE + "/keys/rotops/rotate" + API) + .then().statusCode(200); + + // The rotated version must still deny everything. + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"RS256\",\"value\":\"" + b64Url("digest") + "\"}") + .when().post(BASE + "/keys/rotops/sign" + API) + .then().statusCode(403) + .body("error.code", equalTo("Forbidden")); + } + + @Test + @DisplayName("create with explicit empty key_ops denies all operations") + void createWithExplicitEmptyKeyOps() { + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"kty\":\"RSA\",\"key_size\":2048,\"key_ops\":[]}") + .when().post(BASE + "/keys/emptyops/create" + API) + .then().statusCode(200); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"RS256\",\"value\":\"" + b64Url("digest") + "\"}") + .when().post(BASE + "/keys/emptyops/sign" + API) + .then().statusCode(403) + .body("error.code", equalTo("Forbidden")); + } + + @Test + @DisplayName("key_ops null forbids all operations") + void nullKeyOpsForbidsAllOperations() { + Map jwk = fullRsaJwk(); + jwk.put("key_ops", null); + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"key\":" + toJson(jwk) + "}") + .when().put(BASE + "/keys/nullops" + API) + .then().statusCode(200); + + given().header("Authorization", AUTH).contentType(ContentType.JSON) + .body("{\"alg\":\"RS256\",\"value\":\"" + b64Url("digest") + "\"}") + .when().post(BASE + "/keys/nullops/sign" + API) + .then().statusCode(403) + .body("error.code", equalTo("Forbidden")); + } + + // ── Helpers ──────────────────────────────────────────────────────────────── + + private static final Base64.Encoder URL = Base64.getUrlEncoder().withoutPadding(); + + private static String b64Url(String value) { + return URL.encodeToString(value.getBytes(java.nio.charset.StandardCharsets.UTF_8)); + } + + /** Builds a full RSA JWK map (private CRT fields included) from a fresh JDK keypair. */ + private static Map fullRsaJwk() { + try { + KeyPairGenerator gen = KeyPairGenerator.getInstance("RSA"); + gen.initialize(2048); + KeyPair pair = gen.generateKeyPair(); + RSAPublicKey pub = (RSAPublicKey) pair.getPublic(); + RSAPrivateCrtKey priv = (RSAPrivateCrtKey) pair.getPrivate(); + Map jwk = new LinkedHashMap<>(); + jwk.put("kty", "RSA"); + jwk.put("n", URL.encodeToString(unsigned(pub.getModulus()))); + jwk.put("e", URL.encodeToString(unsigned(pub.getPublicExponent()))); + jwk.put("d", URL.encodeToString(unsigned(priv.getPrivateExponent()))); + jwk.put("p", URL.encodeToString(unsigned(priv.getPrimeP()))); + jwk.put("q", URL.encodeToString(unsigned(priv.getPrimeQ()))); + jwk.put("dp", URL.encodeToString(unsigned(priv.getPrimeExponentP()))); + jwk.put("dq", URL.encodeToString(unsigned(priv.getPrimeExponentQ()))); + jwk.put("qi", URL.encodeToString(unsigned(priv.getCrtCoefficient()))); + return jwk; + } catch (Exception e) { + throw new IllegalStateException(e); + } + } + + private static byte[] unsigned(BigInteger value) { + byte[] bytes = value.toByteArray(); + if (bytes.length > 1 && bytes[0] == 0) { + byte[] trimmed = new byte[bytes.length - 1]; + System.arraycopy(bytes, 1, trimmed, 0, trimmed.length); + return trimmed; + } + return bytes; + } + + private static String toJson(Object obj) { + try { + return new com.fasterxml.jackson.databind.ObjectMapper().writeValueAsString(obj); + } catch (Exception e) { + throw new IllegalStateException(e); + } + } +}