"
+```
+
+The app first proves token A with binding A returns HTTP 200, then requires
+token A with binding B to be rejected.
+
+See
+[`msal4j-sdk/docs/managed-identity-v2-mtls-pop.md`](../msal4j-sdk/docs/managed-identity-v2-mtls-pop.md)
+for architecture, protocol reconciliation, and troubleshooting.
diff --git a/msal4j-mtls-extensions/pom.xml b/msal4j-mtls-extensions/pom.xml
new file mode 100644
index 00000000..1a2389a8
--- /dev/null
+++ b/msal4j-mtls-extensions/pom.xml
@@ -0,0 +1,93 @@
+
+
+ 4.0.0
+
+ com.microsoft.azure
+ msal4j-mtls-extensions
+ 1.0.0
+ jar
+
+ Microsoft Authentication Library for Java - mTLS Extensions
+
+ Extension package that enables mTLS Proof-of-Possession (mTLS PoP) token acquisition
+ for Azure Managed Identity scenarios requiring KeyGuard-bound certificates. Uses JNA
+ to call Windows CNG (ncrypt.dll) and AttestationClientLib.dll directly from Java,
+ implementing a java.security.Provider that allows JSSE to use a non-exportable
+ KeyGuard RSA key during the TLS handshake. No .NET runtime or subprocess required.
+
+
+
+ 8
+ 8
+ UTF-8
+ 1.1.5
+ 90dfcce20e1a74519b49796eeee17e6e59a257c3acf754f454a49380d28a568b
+
+
+
+
+ com.microsoft.azure
+ msal4j
+ 1.25.1
+
+
+
+
+ net.java.dev.jna
+ jna
+ 5.14.0
+
+
+
+
+ org.junit.jupiter
+ junit-jupiter-api
+ 5.10.0
+ test
+
+
+ org.junit.jupiter
+ junit-jupiter-engine
+ 5.10.0
+ test
+
+
+ org.mockito
+ mockito-inline
+ 4.11.0
+ test
+
+
+ org.mockito
+ mockito-junit-jupiter
+ 4.11.0
+ test
+
+
+
+
+
+
+ org.apache.maven.plugins
+ maven-jar-plugin
+ 3.3.0
+
+
+
+ ${keyguard.attestation.version}
+ ${keyguard.attestation.sha256}
+
+
+
+
+
+ org.apache.maven.plugins
+ maven-surefire-plugin
+ 3.1.2
+
+
+
+
+
diff --git a/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/AttestationLibrary.java b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/AttestationLibrary.java
new file mode 100644
index 00000000..19ee93e6
--- /dev/null
+++ b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/AttestationLibrary.java
@@ -0,0 +1,109 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import com.sun.jna.Callback;
+import com.sun.jna.Library;
+import com.sun.jna.Pointer;
+import com.sun.jna.Structure;
+import com.sun.jna.ptr.PointerByReference;
+
+import java.util.Arrays;
+import java.util.List;
+
+/**
+ * JNA binding for {@code AttestationClientLib.dll} — the Windows DLL shipped by Azure
+ * that produces a MAA (Microsoft Azure Attestation) JWT proving a CNG KeyGuard key is
+ * hardware-protected.
+ *
+ * Function signatures (ANSI cdecl, x64 Windows) documented in MSAL.NET's
+ * {@code KeyGuardMaa/AttestationInterop.cs} and also used by msal-go's
+ * {@code cng_windows.go}:
+ *
+ * int InitAttestationLib(AttestationLogInfo*)
+ * int AttestKeyGuardImportKey(char* endpoint, char* authToken, char* clientPayload,
+ * NCRYPT_KEY_HANDLE keyHandle, char** token, char* clientId)
+ * void FreeAttestationToken(char* token)
+ * void UninitAttestationLib()
+ *
+ *
+ * This interface is loaded lazily via {@link AttestationLibraryLoader} from the
+ * Microsoft-signed DLL bundled in the optional extension. It is only extracted and loaded
+ * when MAA attestation is requested.
+ */
+interface AttestationLibrary extends Library {
+
+ /**
+ * No-op log callback that satisfies the DLL's requirement for a non-null LogFunc.
+ *
+ * The DLL requires a non-null log function pointer in {@link AttestationLogInfo}.
+ * Passing {@code Pointer.NULL} causes {@code InitAttestationLib} to return an error
+ * (0xFFFFFFF8 = -8). Mirrors msal-go's {@code dummyLogCallback}.
+ *
+ * Signature (cdecl, x64 Windows):
+ * {@code void LogFunc(void* ctx, char* tag, int lvl, char* func, int line, char* msg)}
+ */
+ interface LogCallback extends Callback {
+ void log(Pointer ctx, Pointer tag, int level, Pointer func, int line, Pointer msg);
+ }
+
+ /** Shared no-op log callback instance — kept alive to prevent GC. */
+ LogCallback NOOP_LOG = (ctx, tag, level, func, line, msg) -> {};
+
+ /**
+ * Mirrors the {@code AttestationLogInfo} struct:
+ * struct AttestationLogInfo { LogFunc Log; void* Ctx; }
+ *
+ * The {@code logFunc} field MUST be a non-null function pointer — the DLL validates
+ * this and returns an error if it is null. Use {@link #NOOP_LOG} for no-op logging.
+ */
+ class AttestationLogInfo extends Structure {
+ /** Function pointer for the log callback. MUST NOT be null. */
+ public LogCallback logFunc;
+ /** Caller context pointer, passed as first arg to logFunc. */
+ public Pointer ctx;
+
+ public AttestationLogInfo() {
+ logFunc = NOOP_LOG; // DLL requires a non-null log function pointer
+ ctx = Pointer.NULL;
+ }
+
+ @Override
+ protected List getFieldOrder() {
+ return Arrays.asList("logFunc", "ctx");
+ }
+ }
+
+ /**
+ * Initializes the attestation library.
+ *
+ * @param logInfo logging configuration; {@code logFunc} MUST be non-null
+ * @return 0 on success, non-zero on failure
+ */
+ int InitAttestationLib(AttestationLogInfo logInfo);
+
+ /**
+ * Produces a MAA JWT proving the given CNG key is VBS/KeyGuard-protected.
+ *
+ * @param endpoint MAA endpoint URL (ANSI string, e.g. "https://sharedcuse.cuse.attest.azure.net")
+ * @param authToken unused, pass null
+ * @param clientPayload unused, pass null
+ * @param keyHandle the {@code NCRYPT_KEY_HANDLE} from NCrypt* operations
+ * @param tokenOut receives the pointer to the MAA JWT string (caller must free with FreeAttestationToken)
+ * @param clientId managed identity client ID (ANSI string)
+ * @return 0 on success, non-zero on failure
+ */
+ int AttestKeyGuardImportKey(String endpoint, String authToken, String clientPayload,
+ Pointer keyHandle, PointerByReference tokenOut, String clientId);
+
+ /**
+ * Frees a MAA JWT string allocated by {@link #AttestKeyGuardImportKey}.
+ *
+ * @param token the pointer returned in {@code tokenOut} by AttestKeyGuardImportKey
+ */
+ void FreeAttestationToken(Pointer token);
+
+ /** Uninitializes the attestation library. Call after all attestation operations. */
+ void UninitAttestationLib();
+}
diff --git a/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/AttestationLibraryLoader.java b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/AttestationLibraryLoader.java
new file mode 100644
index 00000000..9c8df0c1
--- /dev/null
+++ b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/AttestationLibraryLoader.java
@@ -0,0 +1,136 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import com.sun.jna.Native;
+
+import java.io.File;
+import java.io.IOException;
+import java.io.InputStream;
+import java.nio.file.Files;
+import java.nio.file.Path;
+import java.security.DigestInputStream;
+import java.security.MessageDigest;
+import java.security.NoSuchAlgorithmException;
+
+final class AttestationLibraryLoader {
+
+ static final String VERSION = "1.1.5";
+ static final String RESOURCE_PATH =
+ "/META-INF/native/win-x64/AttestationClientLib.dll";
+ static final String SHA256 =
+ "90dfcce20e1a74519b49796eeee17e6e59a257c3acf754f454a49380d28a568b";
+
+ private static final Object LOAD_LOCK = new Object();
+ private static volatile AttestationLibrary loadedLibrary;
+
+ private AttestationLibraryLoader() {
+ }
+
+ static AttestationLibrary load() throws MtlsMsiException {
+ AttestationLibrary library = loadedLibrary;
+ if (library != null) {
+ return library;
+ }
+
+ synchronized (LOAD_LOCK) {
+ library = loadedLibrary;
+ if (library != null) {
+ return library;
+ }
+
+ Path extractedLibrary = extractBundledLibrary();
+ try {
+ library = Native.load(
+ extractedLibrary.toAbsolutePath().toString(),
+ AttestationLibrary.class);
+ } catch (UnsatisfiedLinkError e) {
+ throw new MtlsMsiException(
+ "Could not load bundled Microsoft.Azure.Security.KeyGuardAttestation " +
+ VERSION + " native library: " + e.getMessage(),
+ e);
+ }
+ loadedLibrary = library;
+ return library;
+ }
+ }
+
+ static Path extractBundledLibrary() throws MtlsMsiException {
+ String architecture = System.getProperty("os.arch", "");
+ if (!"amd64".equalsIgnoreCase(architecture)
+ && !"x86_64".equalsIgnoreCase(architecture)) {
+ throw new MtlsMsiException(
+ "Microsoft.Azure.Security.KeyGuardAttestation " + VERSION +
+ " is bundled only for Windows x64; detected architecture: " +
+ architecture);
+ }
+
+ MessageDigest digest;
+ try {
+ digest = MessageDigest.getInstance("SHA-256");
+ } catch (NoSuchAlgorithmException e) {
+ throw new MtlsMsiException("SHA-256 is unavailable in this Java runtime.", e);
+ }
+
+ Path directory = null;
+ Path library = null;
+ try (InputStream resource =
+ AttestationLibraryLoader.class.getResourceAsStream(RESOURCE_PATH)) {
+ if (resource == null) {
+ throw new MtlsMsiException(
+ "Bundled Microsoft.Azure.Security.KeyGuardAttestation " + VERSION +
+ " native library is missing from the extension JAR.");
+ }
+
+ directory = Files.createTempDirectory("msal4j-keyguard-attestation-");
+ library = directory.resolve("AttestationClientLib.dll");
+ try (DigestInputStream verifiedResource =
+ new DigestInputStream(resource, digest)) {
+ Files.copy(verifiedResource, library);
+ }
+
+ String actualHash = toHex(digest.digest());
+ if (!SHA256.equals(actualHash)) {
+ Files.deleteIfExists(library);
+ Files.deleteIfExists(directory);
+ throw new MtlsMsiException(
+ "Bundled Microsoft.Azure.Security.KeyGuardAttestation " + VERSION +
+ " failed SHA-256 verification.");
+ }
+
+ File directoryFile = directory.toFile();
+ File libraryFile = library.toFile();
+ directoryFile.deleteOnExit();
+ libraryFile.deleteOnExit();
+ return library;
+ } catch (IOException e) {
+ if (library != null) {
+ try {
+ Files.deleteIfExists(library);
+ } catch (IOException ignored) {
+ // Preserve the original extraction failure.
+ }
+ }
+ if (directory != null) {
+ try {
+ Files.deleteIfExists(directory);
+ } catch (IOException ignored) {
+ // Preserve the original extraction failure.
+ }
+ }
+ throw new MtlsMsiException(
+ "Could not extract bundled Microsoft.Azure.Security.KeyGuardAttestation " +
+ VERSION + " native library.",
+ e);
+ }
+ }
+
+ private static String toHex(byte[] bytes) {
+ StringBuilder result = new StringBuilder(bytes.length * 2);
+ for (byte value : bytes) {
+ result.append(String.format("%02x", value & 0xff));
+ }
+ return result.toString();
+ }
+}
diff --git a/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/AttestationTokenCache.java b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/AttestationTokenCache.java
new file mode 100644
index 00000000..2305cf7b
--- /dev/null
+++ b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/AttestationTokenCache.java
@@ -0,0 +1,117 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import java.nio.charset.StandardCharsets;
+import java.util.Base64;
+import java.util.Map;
+import java.util.concurrent.ConcurrentHashMap;
+
+final class AttestationTokenCache {
+
+ private static final long FRESHNESS_BUFFER_SECONDS = 5L * 60L;
+ private final Map entries = new ConcurrentHashMap<>();
+ private final Map locks = new ConcurrentHashMap<>();
+
+ String getOrAttest(
+ String endpoint,
+ String keyId,
+ AttestationOperation operation) {
+ String cacheKey = normalizeEndpoint(endpoint) + "|" + keyId;
+ Entry cached = entries.get(cacheKey);
+ if (isFresh(cached)) {
+ return cached.jwt;
+ }
+
+ Object lock = locks.computeIfAbsent(cacheKey, ignored -> new Object());
+ synchronized (lock) {
+ cached = entries.get(cacheKey);
+ if (isFresh(cached)) {
+ return cached.jwt;
+ }
+ String jwt = operation.attest();
+ if (jwt == null || jwt.trim().isEmpty()) {
+ throw new MtlsMsiException(
+ "KeyGuard attestation failed; no attestation token was produced.");
+ }
+ long expiresOn = readExpiry(jwt);
+ if (expiresOn <= currentEpochSeconds() + FRESHNESS_BUFFER_SECONDS) {
+ throw new MtlsMsiException(
+ "KeyGuard attestation returned an expired or insufficiently fresh token.");
+ }
+ entries.put(cacheKey, new Entry(jwt, expiresOn));
+ return jwt;
+ }
+ }
+
+ private static boolean isFresh(Entry entry) {
+ return entry != null
+ && entry.expiresOn > currentEpochSeconds() + FRESHNESS_BUFFER_SECONDS;
+ }
+
+ private static long readExpiry(String jwt) throws MtlsMsiException {
+ try {
+ String[] segments = jwt.split("\\.");
+ if (segments.length < 2) {
+ throw new IllegalArgumentException("JWT has fewer than two segments");
+ }
+ String payload = new String(
+ Base64.getUrlDecoder().decode(padBase64(segments[1])),
+ StandardCharsets.UTF_8);
+ String marker = "\"exp\"";
+ int index = payload.indexOf(marker);
+ int colon = index < 0 ? -1 : payload.indexOf(':', index + marker.length());
+ if (colon < 0) {
+ throw new IllegalArgumentException("JWT has no exp claim");
+ }
+ int start = colon + 1;
+ while (start < payload.length()
+ && Character.isWhitespace(payload.charAt(start))) {
+ start++;
+ }
+ int end = start;
+ while (end < payload.length() && Character.isDigit(payload.charAt(end))) {
+ end++;
+ }
+ return Long.parseLong(payload.substring(start, end));
+ } catch (Exception e) {
+ throw new MtlsMsiException(
+ "Unable to determine KeyGuard attestation token expiry.", e);
+ }
+ }
+
+ private static String padBase64(String value) {
+ int remainder = value.length() % 4;
+ if (remainder == 0) {
+ return value;
+ }
+ return value + (remainder == 2 ? "==" : "=");
+ }
+
+ private static String normalizeEndpoint(String endpoint) {
+ String normalized = endpoint.trim().toLowerCase();
+ while (normalized.endsWith("/")) {
+ normalized = normalized.substring(0, normalized.length() - 1);
+ }
+ return normalized;
+ }
+
+ private static long currentEpochSeconds() {
+ return System.currentTimeMillis() / 1000L;
+ }
+
+ interface AttestationOperation {
+ String attest();
+ }
+
+ private static final class Entry {
+ final String jwt;
+ final long expiresOn;
+
+ Entry(String jwt, long expiresOn) {
+ this.jwt = jwt;
+ this.expiresOn = expiresOn;
+ }
+ }
+}
diff --git a/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/CngKeyGuard.java b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/CngKeyGuard.java
new file mode 100644
index 00000000..97bdae30
--- /dev/null
+++ b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/CngKeyGuard.java
@@ -0,0 +1,398 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import com.sun.jna.Memory;
+import com.sun.jna.Pointer;
+import com.sun.jna.WString;
+import com.sun.jna.ptr.IntByReference;
+import com.sun.jna.ptr.PointerByReference;
+
+import java.math.BigInteger;
+import java.nio.ByteBuffer;
+import java.nio.ByteOrder;
+
+/**
+ * Windows CNG key operations for mTLS PoP Managed Identity.
+ *
+ * Mirrors msal-go's {@code cng_windows.go} and MSAL.NET's
+ * {@code WindowsManagedIdentityKeyProvider}: creates or opens a persisted RSA key in
+ * the Microsoft Software Key Storage Provider, using the same 3-level priority:
+ *
+ * - KeyGuard — Software KSP + USER scope + VBS Virtual Isolation flags.
+ * Requires Credential Guard / Core Isolation on the VM.
+ * - Hardware — Software KSP + USER scope, no VBS flags.
+ * - mTLS PoP requires KeyGuard and throws {@link MtlsMsiException} for Hardware keys.
+ *
+ */
+final class CngKeyGuard {
+
+ private static final Object ATTESTATION_LIBRARY_LOCK = new Object();
+ private static final String MS_SOFTWARE_KSP = "Microsoft Software Key Storage Provider";
+ private static final String RSA_ALG = "RSA";
+ private static final String EXPORT_POLICY = "Export Policy";
+ private static final String KEY_LENGTH = "Length";
+ private static final String VIRTUAL_ISO = "Virtual Iso";
+ private static final String RSAPUBLICBLOB = "RSAPUBLICBLOB";
+
+ private CngKeyGuard() {}
+
+ /**
+ * Gets or creates the mTLS PoP binding key, attempting KeyGuard first.
+ *
+ * @param keyName persisted key name in the KSP (e.g. {@code "MSALMtlsKey_"})
+ * @return a {@link CngRsaPrivateKey} backed by the CNG handle
+ * @throws MtlsMsiException if the system is not Windows, the key cannot be created,
+ * or KeyGuard protection is unavailable
+ */
+ static CngRsaPrivateKey getOrCreateKey(String keyName) throws MtlsMsiException {
+ if (!isWindows()) {
+ throw new MtlsMsiException("mTLS PoP Managed Identity is only supported on Windows Azure VMs.");
+ }
+
+ // 1. Try KeyGuard (USER scope + VBS Virtual Isolation flags).
+ int kgCreateFlags = NCryptLibrary.NCRYPT_OVERWRITE_KEY_FLAG
+ | NCryptLibrary.NCRYPT_USE_VIRTUAL_ISOLATION_FLAG
+ | NCryptLibrary.NCRYPT_USE_PER_BOOT_KEY_FLAG;
+ try {
+ CngRsaPrivateKey key = tryGetOrCreateKey(keyName, NCryptLibrary.NCRYPT_SILENT_FLAG, kgCreateFlags, NCryptLibrary.NCRYPT_SILENT_FLAG);
+ if (isKeyGuardProtected(key.nativeHandle())) {
+ return key;
+ }
+ // Created but VBS protection not active — delete and retry once (mirrors MSAL.NET).
+ NCryptLibrary.INSTANCE.NCryptDeleteKey(key.nativeHandle(), 0);
+ key = tryGetOrCreateKey(keyName, NCryptLibrary.NCRYPT_SILENT_FLAG, kgCreateFlags, NCryptLibrary.NCRYPT_SILENT_FLAG);
+ if (isKeyGuardProtected(key.nativeHandle())) {
+ return key;
+ }
+ key.close();
+ } catch (MtlsMsiException ignored) {
+ // KeyGuard not available on this VM; fall through to error below.
+ }
+
+ throw new MtlsMsiException(
+ "mTLS PoP requires a VBS KeyGuard-protected RSA key, but KeyGuard is not available " +
+ "on this VM. Ensure Credential Guard / Core Isolation is enabled: the VM must be " +
+ "Trusted Launch (Secure Boot + vTPM) with VBS active " +
+ "(check msinfo32.exe: 'Virtualization-based security' = Running).");
+ }
+
+ /**
+ * Produces a MAA JWT by calling {@code AttestationClientLib.dll}.
+ *
+ * @param keyHandle CNG key handle owned by the binding generation
+ * @param endpoint MAA attestation endpoint URL (from IMDS platform metadata)
+ * @param clientId managed identity client ID (from IMDS platform metadata)
+ * @return the MAA JWT string
+ * @throws MtlsMsiException if the DLL is not present, or attestation fails
+ */
+ static String getAttestationToken(Pointer keyHandle, String endpoint, String clientId)
+ throws MtlsMsiException {
+ synchronized (ATTESTATION_LIBRARY_LOCK) {
+ return getAttestationTokenSynchronized(keyHandle, endpoint, clientId);
+ }
+ }
+
+ private static String getAttestationTokenSynchronized(
+ Pointer keyHandle,
+ String endpoint,
+ String clientId) throws MtlsMsiException {
+
+ AttestationLibrary attestLib = AttestationLibraryLoader.load();
+
+ AttestationLibrary.AttestationLogInfo logInfo = new AttestationLibrary.AttestationLogInfo();
+ int ret = attestLib.InitAttestationLib(logInfo);
+ if (ret != 0) {
+ throw new MtlsMsiException(
+ String.format("InitAttestationLib failed: 0x%x", ret));
+ }
+
+ try {
+ PointerByReference tokenRef = new PointerByReference();
+ ret = attestLib.AttestKeyGuardImportKey(endpoint, null, null, keyHandle, tokenRef, clientId);
+ if (ret != 0) {
+ throw new MtlsMsiException(String.format(
+ "AttestKeyGuardImportKey failed (rc=0x%x). This usually means the VM's vTPM " +
+ "is not provisioned for attestation. mTLS PoP requires a Trusted Launch Azure VM " +
+ "(Secure Boot + vTPM) with an EK certificate. " +
+ "Check 'tpmtool.exe getdeviceinformation': 'Is Capable For Attestation' must be true.", ret));
+ }
+
+ Pointer tokenPtr = tokenRef.getValue();
+ if (tokenPtr == null || tokenPtr == Pointer.NULL) {
+ throw new MtlsMsiException("AttestKeyGuardImportKey returned null token");
+ }
+
+ try {
+ String jwt = tokenPtr.getString(0); // ANSI (null-terminated)
+ if (jwt == null || jwt.isEmpty()) {
+ throw new MtlsMsiException("AttestKeyGuardImportKey returned empty token");
+ }
+ return jwt;
+ } finally {
+ attestLib.FreeAttestationToken(tokenPtr);
+ }
+ } finally {
+ attestLib.UninitAttestationLib();
+ }
+ }
+
+ /**
+ * Signs a digest using {@code NCryptSignHash} with PKCS#1 v1.5 padding.
+ *
+ * @param keyHandle CNG key handle
+ * @param digest the hash bytes to sign
+ * @param hashAlgCng CNG hash algorithm name (e.g. {@code "SHA256"})
+ * @return DER-encoded signature bytes
+ * @throws MtlsMsiException if signing fails
+ */
+ static byte[] signPkcs1(Pointer keyHandle, byte[] digest, String hashAlgCng)
+ throws MtlsMsiException {
+ NCryptLibrary.BcryptPkcs1PaddingInfo padding =
+ new NCryptLibrary.BcryptPkcs1PaddingInfo(hashAlgCng);
+ return ncryptSign(keyHandle, padding.getPointer(), NCryptLibrary.NCRYPT_PAD_PKCS1_FLAG, digest, "PKCS1v15");
+ }
+
+ /**
+ * Signs a digest using {@code NCryptSignHash} with RSASSA-PSS padding.
+ *
+ * @param keyHandle CNG key handle
+ * @param digest the hash bytes to sign
+ * @param hashAlgCng CNG hash algorithm name (e.g. {@code "SHA256"})
+ * @param saltLen PSS salt length in bytes
+ * @return DER-encoded signature bytes
+ * @throws MtlsMsiException if signing fails
+ */
+ static byte[] signPss(Pointer keyHandle, byte[] digest, String hashAlgCng, int saltLen)
+ throws MtlsMsiException {
+ NCryptLibrary.BcryptPssPaddingInfo padding =
+ new NCryptLibrary.BcryptPssPaddingInfo(hashAlgCng, saltLen);
+ return ncryptSign(keyHandle, padding.getPointer(), NCryptLibrary.NCRYPT_PAD_PSS_FLAG, digest, "PSS");
+ }
+
+ private static byte[] ncryptSign(Pointer hKey, Pointer paddingPtr, int paddingFlag,
+ byte[] digest, String label) throws MtlsMsiException {
+ IntByReference sigLen = new IntByReference(0);
+ // First call: query the signature buffer size.
+ int ret = NCryptLibrary.INSTANCE.NCryptSignHash(
+ hKey, paddingPtr, digest, digest.length,
+ null, 0, sigLen, paddingFlag);
+ if (ret != NCryptLibrary.ERROR_SUCCESS) {
+ throw new MtlsMsiException(
+ String.format("NCryptSignHash %s (size query) failed: 0x%x", label, ret));
+ }
+
+ Memory sigBuf = new Memory(sigLen.getValue());
+ ret = NCryptLibrary.INSTANCE.NCryptSignHash(
+ hKey, paddingPtr, digest, digest.length,
+ sigBuf, sigLen.getValue(), sigLen, paddingFlag);
+ if (ret != NCryptLibrary.ERROR_SUCCESS) {
+ throw new MtlsMsiException(
+ String.format("NCryptSignHash %s failed: 0x%x", label, ret));
+ }
+
+ return sigBuf.getByteArray(0, sigLen.getValue());
+ }
+
+ // ─── Internal helpers ─────────────────────────────────────────────────────
+
+ private static CngRsaPrivateKey tryGetOrCreateKey(String keyName,
+ int openFlags,
+ int createFlags,
+ int finalizeFlags) throws MtlsMsiException {
+ Pointer hProvider = openProvider();
+ Pointer hKey = null;
+ boolean created = false;
+ boolean ownershipTransferred = false;
+ try {
+ WString keyNameW = new WString(keyName);
+
+ // 1. Try to open an existing key.
+ hKey = openExistingUsableKey(
+ hProvider,
+ keyNameW,
+ openFlags,
+ NCryptLibrary.INSTANCE,
+ CngKeyGuard::probeOpenedKey);
+ PointerByReference phKey = new PointerByReference();
+
+ // 2. Create a new key if open failed.
+ if (hKey == null) {
+ int ret = NCryptLibrary.INSTANCE.NCryptCreatePersistedKey(
+ hProvider, phKey,
+ new WString(RSA_ALG),
+ keyNameW,
+ 0, createFlags);
+ if (ret != NCryptLibrary.ERROR_SUCCESS) {
+ throw new MtlsMsiException(
+ String.format("NCryptCreatePersistedKey failed: 0x%x", ret));
+ }
+ hKey = phKey.getValue();
+ created = true;
+
+ // Set key length to 2048.
+ setDwordProperty(hKey, KEY_LENGTH, 2048);
+ // Set non-exportable.
+ setDwordProperty(hKey, EXPORT_POLICY, NCryptLibrary.NCRYPT_ALLOW_EXPORT_NONE);
+
+ ret = NCryptLibrary.INSTANCE.NCryptFinalizeKey(hKey, finalizeFlags);
+ if (ret != NCryptLibrary.ERROR_SUCCESS) {
+ NCryptLibrary.INSTANCE.NCryptDeleteKey(hKey, 0);
+ hKey = null;
+ throw new MtlsMsiException(
+ String.format("NCryptFinalizeKey failed: 0x%x. " +
+ "VBS isolation flags are not supported on this machine " +
+ "(Credential Guard / Core Isolation not active).", ret));
+ }
+ }
+
+ BigInteger[] pubKey = exportPublicKey(hKey);
+ CngRsaPrivateKey privateKey =
+ new CngRsaPrivateKey(hKey, pubKey[0], pubKey[1].intValue());
+ ownershipTransferred = true;
+ return privateKey;
+
+ } finally {
+ if (!ownershipTransferred && hKey != null) {
+ if (created) {
+ NCryptLibrary.INSTANCE.NCryptDeleteKey(hKey, 0);
+ } else {
+ NCryptLibrary.INSTANCE.NCryptFreeObject(hKey);
+ }
+ }
+ NCryptLibrary.INSTANCE.NCryptFreeObject(hProvider);
+ }
+ }
+
+ static boolean isKeyGuardProtected(Pointer hKey) {
+ WString propW = new WString(VIRTUAL_ISO);
+ byte[] buf = new byte[4];
+ IntByReference pcbResult = new IntByReference(0);
+ int ret = NCryptLibrary.INSTANCE.NCryptGetProperty(
+ hKey, propW, buf, buf.length, pcbResult, 0);
+ if (ret != NCryptLibrary.ERROR_SUCCESS || pcbResult.getValue() < 4) {
+ return false;
+ }
+ int val = ByteBuffer.wrap(buf).order(ByteOrder.LITTLE_ENDIAN).getInt();
+ return val != 0;
+ }
+
+ private static void probeKeyLiveness(Pointer hKey) {
+ signPkcs1(hKey, new byte[32], "SHA256");
+ }
+
+ private static void probeOpenedKey(Pointer hKey) {
+ exportPublicKeyBytes(hKey);
+ probeKeyLiveness(hKey);
+ }
+
+ static Pointer openExistingUsableKey(
+ Pointer provider,
+ WString keyName,
+ int openFlags,
+ NCryptLibrary nativeApi,
+ KeyProbe probe) {
+ PointerByReference keyReference = new PointerByReference();
+ int result = nativeApi.NCryptOpenKey(
+ provider, keyReference, keyName, 0, openFlags);
+ if (result != NCryptLibrary.ERROR_SUCCESS) {
+ return null;
+ }
+
+ Pointer key = keyReference.getValue();
+ try {
+ probe.probe(key);
+ return key;
+ } catch (MtlsMsiException e) {
+ nativeApi.NCryptDeleteKey(key, 0);
+ return null;
+ }
+ }
+
+ interface KeyProbe {
+ void probe(Pointer key);
+ }
+
+ /** Returns byte[] of the RSAPUBLICBLOB for use in CSR SubjectPublicKeyInfo. */
+ static byte[] exportPublicKeyBytes(Pointer hKey) throws MtlsMsiException {
+ WString blobType = new WString(RSAPUBLICBLOB);
+ IntByReference pcbResult = new IntByReference(0);
+
+ // Query size.
+ int ret = NCryptLibrary.INSTANCE.NCryptExportKey(
+ hKey, null, blobType, null, null, 0, pcbResult, 0);
+ if (ret != NCryptLibrary.ERROR_SUCCESS) {
+ throw new MtlsMsiException(
+ String.format("NCryptExportKey (size query) failed: 0x%x", ret));
+ }
+
+ Memory blob = new Memory(pcbResult.getValue());
+ ret = NCryptLibrary.INSTANCE.NCryptExportKey(
+ hKey, null, blobType, null, blob, pcbResult.getValue(), pcbResult, 0);
+ if (ret != NCryptLibrary.ERROR_SUCCESS) {
+ throw new MtlsMsiException(
+ String.format("NCryptExportKey failed: 0x%x", ret));
+ }
+
+ return blob.getByteArray(0, pcbResult.getValue());
+ }
+
+ /**
+ * Returns [modulus, publicExponent] parsed from the RSAPUBLICBLOB.
+ * BCRYPT_RSAKEY_BLOB format (24-byte header):
+ * Magic(4) BitLength(4) cbPublicExp(4) cbModulus(4) cbPrime1(4) cbPrime2(4)
+ * followed by PublicExponent bytes then Modulus bytes.
+ */
+ static BigInteger[] exportPublicKey(Pointer hKey) throws MtlsMsiException {
+ byte[] blob = exportPublicKeyBytes(hKey);
+ if (blob.length < 24) {
+ throw new MtlsMsiException("RSAPUBLICBLOB too short: " + blob.length);
+ }
+
+ ByteBuffer bb = ByteBuffer.wrap(blob).order(ByteOrder.LITTLE_ENDIAN);
+ bb.getInt(); // magic
+ bb.getInt(); // bitLength
+ int cbPublicExp = bb.getInt();
+ int cbModulus = bb.getInt();
+ // skip cbPrime1, cbPrime2
+ bb.position(24);
+
+ byte[] expBytes = new byte[cbPublicExp];
+ bb.get(expBytes);
+ byte[] modBytes = new byte[cbModulus];
+ bb.get(modBytes);
+
+ return new BigInteger[] {
+ new BigInteger(1, modBytes), // [0] = modulus
+ new BigInteger(1, expBytes) // [1] = publicExponent
+ };
+ }
+
+ private static Pointer openProvider() throws MtlsMsiException {
+ PointerByReference phProvider = new PointerByReference();
+ int ret = NCryptLibrary.INSTANCE.NCryptOpenStorageProvider(
+ phProvider, new WString(MS_SOFTWARE_KSP), 0);
+ if (ret != NCryptLibrary.ERROR_SUCCESS) {
+ throw new MtlsMsiException(
+ String.format("NCryptOpenStorageProvider failed: 0x%x", ret));
+ }
+ return phProvider.getValue();
+ }
+
+ private static void setDwordProperty(Pointer hKey, String propName, int value)
+ throws MtlsMsiException {
+ byte[] buf = ByteBuffer.allocate(4).order(ByteOrder.LITTLE_ENDIAN).putInt(value).array();
+ int ret = NCryptLibrary.INSTANCE.NCryptSetProperty(
+ hKey, new WString(propName), buf, buf.length, NCryptLibrary.NCRYPT_SILENT_FLAG);
+ if (ret != NCryptLibrary.ERROR_SUCCESS) {
+ throw new MtlsMsiException(
+ String.format("NCryptSetProperty(%s) failed: 0x%x", propName, ret));
+ }
+ }
+
+ private static boolean isWindows() {
+ String os = System.getProperty("os.name", "").toLowerCase();
+ return os.contains("windows");
+ }
+}
diff --git a/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/CngProvider.java b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/CngProvider.java
new file mode 100644
index 00000000..0ba55a19
--- /dev/null
+++ b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/CngProvider.java
@@ -0,0 +1,83 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import java.security.Provider;
+import java.security.Security;
+import java.util.Arrays;
+import java.util.Collections;
+
+/**
+ * A {@link Provider} that routes RSA signature operations for {@link CngRsaPrivateKey}
+ * keys through Windows CNG ({@code NCryptSignHash}) via JNA.
+ *
+ * Install once per JVM via {@link #installIfAbsent()} before creating an
+ * {@code SSLContext} that uses a {@link CngRsaPrivateKey}. JSSE will call
+ * {@code Signature.getInstance("SHA256withRSA")} or
+ * {@code Signature.getInstance("RSASSA-PSS")}; with this provider at
+ * high priority, {@link CngSignatureSpi} intercepts the call and signs via
+ * {@code NCryptSignHash} instead of requiring an exportable private exponent.
+ *
+ * For non-{@link CngRsaPrivateKey} keys, {@link CngSignatureSpi} automatically
+ * delegates to the next available provider, so installing this provider does not
+ * break other RSA signing in the same JVM.
+ */
+public final class CngProvider extends Provider {
+
+ private static final long serialVersionUID = 1L;
+ private static final String PROVIDER_NAME = "CNG";
+ private static final double PROVIDER_VERSION = 1.0;
+ private static final String PROVIDER_INFO = "Windows CNG JNA provider for JSSE mTLS";
+
+ public CngProvider() {
+ super(PROVIDER_NAME, PROVIDER_VERSION, PROVIDER_INFO);
+
+ putSignature("SHA256withRSA", CngSignatureSpi.Sha256WithRsa.class,
+ "SHA-256withRSA");
+ putSignature("SHA384withRSA", CngSignatureSpi.Sha384WithRsa.class,
+ "SHA-384withRSA");
+ putSignature("SHA512withRSA", CngSignatureSpi.Sha512WithRsa.class,
+ "SHA-512withRSA");
+ putSignature("RSASSA-PSS", CngSignatureSpi.RsaSsaPss.class,
+ "SHA256withRSAandMGF1", "SHA384withRSAandMGF1", "SHA512withRSAandMGF1");
+ }
+
+ private void putSignature(
+ String algorithm,
+ Class extends CngSignatureSpi> implementation,
+ String... aliases) {
+ putService(new CngSignatureService(
+ this,
+ algorithm,
+ implementation.getName(),
+ Arrays.asList(aliases)));
+ }
+
+ /**
+ * Installs this provider at position 1 (highest priority) if it is not already
+ * registered. Safe to call multiple times.
+ */
+ public static void installIfAbsent() {
+ if (Security.getProvider(PROVIDER_NAME) == null) {
+ Security.insertProviderAt(new CngProvider(), 1);
+ }
+ }
+
+ private static final class CngSignatureService extends Provider.Service {
+
+ CngSignatureService(
+ Provider provider,
+ String algorithm,
+ String className,
+ java.util.List aliases) {
+ super(provider, "Signature", algorithm, className, aliases,
+ Collections.emptyMap());
+ }
+
+ @Override
+ public boolean supportsParameter(Object parameter) {
+ return parameter instanceof CngRsaPrivateKey;
+ }
+ }
+}
diff --git a/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/CngRsaPrivateKey.java b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/CngRsaPrivateKey.java
new file mode 100644
index 00000000..71fcb3c8
--- /dev/null
+++ b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/CngRsaPrivateKey.java
@@ -0,0 +1,107 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import com.sun.jna.Pointer;
+
+import java.math.BigInteger;
+import java.security.interfaces.RSAPrivateKey;
+import java.util.concurrent.atomic.AtomicBoolean;
+
+/**
+ * A non-exportable RSA private key backed by a Windows CNG {@code NCRYPT_KEY_HANDLE}.
+ *
+ * This key implements {@link RSAPrivateKey} so that JSSE recognizes it as an RSA key
+ * and selects RSA cipher suites. The private exponent is {@code null} and
+ * {@link #getEncoded()} returns {@code null} because the key material never leaves the
+ * CNG key storage provider (KeyGuard VBS isolation).
+ *
+ * Signing is performed by {@link CngKeyGuard#signPkcs1} / {@link CngKeyGuard#signPss},
+ * dispatched from {@link CngSignatureSpi}.
+ *
+ * Callers must call {@link #close()} when done to free the CNG handle.
+ */
+public final class CngRsaPrivateKey implements RSAPrivateKey, AutoCloseable {
+
+ private static final long serialVersionUID = 1L;
+
+ private final Pointer handle;
+ private final BigInteger modulus;
+ private final int publicExponent;
+ private final AtomicBoolean closed = new AtomicBoolean();
+ private final Runnable releaser;
+
+ CngRsaPrivateKey(Pointer handle, BigInteger modulus, int publicExponent) {
+ this(handle, modulus, publicExponent,
+ () -> NCryptLibrary.INSTANCE.NCryptFreeObject(handle));
+ }
+
+ CngRsaPrivateKey(
+ Pointer handle,
+ BigInteger modulus,
+ int publicExponent,
+ Runnable releaser) {
+ this.handle = handle;
+ this.modulus = modulus;
+ this.publicExponent = publicExponent;
+ this.releaser = releaser;
+ }
+
+ Pointer nativeHandle() {
+ if (closed.get()) throw new IllegalStateException("CNG key handle has been closed");
+ return handle;
+ }
+
+ // ─── RSAKey ───────────────────────────────────────────────────────────────
+
+ /** Returns the RSA modulus (from the exported RSAPUBLICBLOB — public information). */
+ @Override
+ public BigInteger getModulus() {
+ return modulus;
+ }
+
+ /**
+ * Always returns {@code null}. The private exponent is non-exportable from the
+ * KeyGuard-protected CNG key; signing is delegated to {@code NCryptSignHash}.
+ */
+ @Override
+ public BigInteger getPrivateExponent() {
+ return null;
+ }
+
+ // ─── Key ──────────────────────────────────────────────────────────────────
+
+ @Override
+ public String getAlgorithm() { return "RSA"; }
+
+ /** Returns {@code null} — non-exportable key has no serializable encoding. */
+ @Override
+ public String getFormat() { return null; }
+
+ /** Returns {@code null} — non-exportable key has no serializable encoding. */
+ @Override
+ public byte[] getEncoded() { return null; }
+
+ // ─── AutoCloseable ────────────────────────────────────────────────────────
+
+ /**
+ * Frees the underlying CNG key handle via {@code NCryptFreeObject}.
+ * The key remains persisted in the KSP; only the in-process handle is released.
+ */
+ @Override
+ public void close() {
+ if (closed.compareAndSet(false, true)) {
+ releaser.run();
+ }
+ }
+
+ boolean isClosed() {
+ return closed.get();
+ }
+
+ /** The public exponent (e.g. 65537 = 0x10001). */
+ public int getPublicExponent() {
+ return publicExponent;
+ }
+}
diff --git a/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/CngSignatureSpi.java b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/CngSignatureSpi.java
new file mode 100644
index 00000000..556f293f
--- /dev/null
+++ b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/CngSignatureSpi.java
@@ -0,0 +1,303 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import com.sun.jna.Pointer;
+
+import java.security.AlgorithmParameters;
+import java.security.InvalidAlgorithmParameterException;
+import java.security.InvalidKeyException;
+import java.security.MessageDigest;
+import java.security.NoSuchAlgorithmException;
+import java.security.PrivateKey;
+import java.security.Provider;
+import java.security.Security;
+import java.security.SignatureException;
+import java.security.SignatureSpi;
+import java.security.spec.AlgorithmParameterSpec;
+import java.security.spec.MGF1ParameterSpec;
+import java.security.spec.PSSParameterSpec;
+
+/**
+ * {@link SignatureSpi} implementations that delegate signing to Windows CNG via JNA.
+ *
+ * Two families are provided:
+ *
+ * - {@link Sha256WithRsa} / {@link Sha1WithRsa} — PKCS#1 v1.5 padding
+ * - {@link RsaSsaPss} — RSASSA-PSS with configurable parameters
+ *
+ *
+ * For non-{@link CngRsaPrivateKey} keys, each SPI delegates to the next available
+ * provider so that installing {@link CngProvider} at high priority does not break other
+ * code in the same JVM that signs with regular (exportable) RSA keys.
+ */
+abstract class CngSignatureSpi extends SignatureSpi {
+
+ // ─── Concrete algorithms ───────────────────────────────────────────────────
+
+ /** SHA-256 with RSA PKCS#1 v1.5 */
+ public static class Sha256WithRsa extends CngSignatureSpi {
+ public Sha256WithRsa() { super("SHA-256", "SHA256", false, 32); }
+ }
+
+ /** SHA-384 with RSA PKCS#1 v1.5 */
+ public static class Sha384WithRsa extends CngSignatureSpi {
+ public Sha384WithRsa() { super("SHA-384", "SHA384", false, 48); }
+ }
+
+ /** SHA-512 with RSA PKCS#1 v1.5 */
+ public static class Sha512WithRsa extends CngSignatureSpi {
+ public Sha512WithRsa() { super("SHA-512", "SHA512", false, 64); }
+ }
+
+ /** RSASSA-PSS — algorithm parameters set via {@link #engineSetParameter(AlgorithmParameterSpec)} */
+ public static class RsaSsaPss extends CngSignatureSpi {
+ public RsaSsaPss() { super("SHA-256", "SHA256", true, 32); }
+ }
+
+ // ─── State ────────────────────────────────────────────────────────────────
+
+ private final boolean pss;
+
+ // CNG mode
+ private Pointer cngHandle;
+ private MessageDigest digest;
+ private String hashJce; // Java algorithm name (e.g. "SHA-256")
+ private String hashCng; // CNG algorithm name (e.g. "SHA256")
+ private int saltLen;
+
+ // Delegation mode (non-CNG keys)
+ private java.security.Signature delegate;
+
+ CngSignatureSpi(String hashJce, String hashCng, boolean pss, int saltLen) {
+ this.hashJce = hashJce;
+ this.hashCng = hashCng;
+ this.pss = pss;
+ this.saltLen = saltLen;
+ }
+
+ // ─── SignatureSpi ─────────────────────────────────────────────────────────
+
+ @Override
+ protected void engineInitVerify(java.security.PublicKey publicKey)
+ throws InvalidKeyException {
+ // CNG only handles signing (NCryptSignHash). For verification (server cert
+ // validation, etc.) we deliberately throw InvalidKeyException so that
+ // Signature.Delegate.chooseProvider() skips this SPI and falls through to
+ // SunRsaSign or another standard provider that handles RSA/ECDSA verification.
+ throw new InvalidKeyException(
+ "CngSignatureSpi does not support verification; use SunRsaSign");
+ }
+
+ @Override
+ protected void engineInitSign(PrivateKey key) throws InvalidKeyException {
+ if (key instanceof CngRsaPrivateKey) {
+ Pointer h;
+ try {
+ h = ((CngRsaPrivateKey) key).nativeHandle();
+ } catch (IllegalStateException e) {
+ throw new InvalidKeyException("CNG key is closed: " + e.getMessage(), e);
+ }
+ if (h == null) {
+ throw new InvalidKeyException("CNG key handle is null (key may be closed or invalid)");
+ }
+ cngHandle = h;
+ delegate = null;
+ try {
+ digest = MessageDigest.getInstance(hashJce);
+ } catch (NoSuchAlgorithmException e) {
+ throw new InvalidKeyException("MessageDigest " + hashJce + " not available", e);
+ }
+ } else {
+ // Delegate to the next provider that handles this algorithm.
+ cngHandle = null;
+ delegate = null;
+ try {
+ delegate = getNextProviderSignature();
+ delegate.initSign(key);
+ } catch (NoSuchAlgorithmException e) {
+ throw new InvalidKeyException("No fallback provider: " + e.getMessage(), e);
+ }
+ }
+ }
+
+ @Override
+ protected void engineUpdate(byte b) throws SignatureException {
+ if (cngHandle != null) {
+ digest.update(b);
+ } else if (delegate != null) {
+ delegate.update(b);
+ } else {
+ throw new SignatureException(
+ "CngSignatureSpi.engineUpdate called before engineInitSign — " +
+ "Signature object was not properly initialized");
+ }
+ }
+
+ @Override
+ protected void engineUpdate(byte[] b, int off, int len) throws SignatureException {
+ if (cngHandle != null) {
+ digest.update(b, off, len);
+ } else if (delegate != null) {
+ delegate.update(b, off, len);
+ } else {
+ throw new SignatureException(
+ "CngSignatureSpi.engineUpdate called before engineInitSign — " +
+ "Signature object was not properly initialized");
+ }
+ }
+
+ @Override
+ protected byte[] engineSign() throws SignatureException {
+ if (cngHandle != null) {
+ byte[] hash = digest.digest();
+ try {
+ if (pss) {
+ return CngKeyGuard.signPss(cngHandle, hash, hashCng, saltLen);
+ } else {
+ return CngKeyGuard.signPkcs1(cngHandle, hash, hashCng);
+ }
+ } catch (MtlsMsiException e) {
+ throw new SignatureException("CNG signing failed: " + e.getMessage(), e);
+ }
+ } else {
+ return delegate.sign();
+ }
+ }
+
+ @Override
+ protected boolean engineVerify(byte[] sigBytes) throws SignatureException {
+ if (delegate != null) {
+ return delegate.verify(sigBytes);
+ }
+ // Verification is not needed for client-auth TLS or CSR generation.
+ throw new SignatureException("CngSignatureSpi does not support verify (CNG-backed keys)");
+ }
+
+ @Override
+ protected void engineSetParameter(AlgorithmParameterSpec params)
+ throws InvalidAlgorithmParameterException {
+ if (!pss) {
+ if (params != null) {
+ throw new InvalidAlgorithmParameterException(
+ "Parameters are not supported for PKCS#1 RSA signatures.");
+ }
+ return;
+ }
+ if (params instanceof PSSParameterSpec) {
+ PSSParameterSpec pssSpec = (PSSParameterSpec) params;
+ validatePssParameters(pssSpec);
+ hashJce = pssSpec.getDigestAlgorithm();
+ hashCng = toCngHashName(pssSpec.getDigestAlgorithm());
+ saltLen = pssSpec.getSaltLength();
+ if (cngHandle != null) {
+ // Re-initialize the digest with the new hash algorithm.
+ try {
+ digest = MessageDigest.getInstance(hashJce);
+ } catch (NoSuchAlgorithmException e) {
+ throw new InvalidAlgorithmParameterException(
+ "MessageDigest " + hashJce + " not available", e);
+ }
+ } else if (delegate != null) {
+ // Forward PSS params to the delegating provider's Signature instance.
+ try {
+ delegate.setParameter(params);
+ } catch (Exception e) {
+ throw new InvalidAlgorithmParameterException(e.getMessage(), e);
+ }
+ }
+ } else {
+ throw new InvalidAlgorithmParameterException(
+ "RSASSA-PSS requires PSSParameterSpec.");
+ }
+ }
+
+ @Override
+ @SuppressWarnings("deprecation")
+ protected void engineSetParameter(String param, Object value) {
+ // Legacy method — no-op, required by abstract superclass.
+ }
+
+ @Override
+ @SuppressWarnings("deprecation")
+ protected Object engineGetParameter(String param) {
+ return null;
+ }
+
+ @Override
+ protected AlgorithmParameters engineGetParameters() {
+ if (pss && delegate == null) {
+ try {
+ AlgorithmParameters ap = AlgorithmParameters.getInstance("RSASSA-PSS");
+ ap.init(new PSSParameterSpec(hashJce, "MGF1",
+ new MGF1ParameterSpec(hashJce), saltLen, 1));
+ return ap;
+ } catch (Exception e) {
+ return null;
+ }
+ }
+ if (delegate != null) {
+ return delegate.getParameters();
+ }
+ return null;
+ }
+
+ // ─── Helpers ──────────────────────────────────────────────────────────────
+
+ private java.security.Signature getNextProviderSignature() throws NoSuchAlgorithmException {
+ String algName = pss ? "RSASSA-PSS" : hashCng + "withRSA";
+ for (Provider p : Security.getProviders()) {
+ if (p instanceof CngProvider) continue;
+ if (p.getService("Signature", algName) != null) {
+ return java.security.Signature.getInstance(algName, p);
+ }
+ }
+ throw new NoSuchAlgorithmException(
+ "No provider for " + algName + " besides CngProvider");
+ }
+
+ private static String toCngHashName(String jceHashName)
+ throws InvalidAlgorithmParameterException {
+ if (jceHashName == null) {
+ throw new InvalidAlgorithmParameterException("PSS digest must be specified.");
+ }
+ switch (jceHashName.toUpperCase().replace("-", "")) {
+ case "SHA256": return "SHA256";
+ case "SHA384": return "SHA384";
+ case "SHA512": return "SHA512";
+ default:
+ throw new InvalidAlgorithmParameterException(
+ "Unsupported RSA digest: " + jceHashName);
+ }
+ }
+
+ private static void validatePssParameters(PSSParameterSpec spec)
+ throws InvalidAlgorithmParameterException {
+ String hash = toCngHashName(spec.getDigestAlgorithm());
+ if (!"MGF1".equalsIgnoreCase(spec.getMGFAlgorithm())) {
+ throw new InvalidAlgorithmParameterException(
+ "Only MGF1 is supported for RSASSA-PSS.");
+ }
+ if (!(spec.getMGFParameters() instanceof MGF1ParameterSpec)) {
+ throw new InvalidAlgorithmParameterException(
+ "MGF1 parameters must be MGF1ParameterSpec.");
+ }
+ String mgfHash = toCngHashName(
+ ((MGF1ParameterSpec) spec.getMGFParameters()).getDigestAlgorithm());
+ if (!hash.equals(mgfHash)) {
+ throw new InvalidAlgorithmParameterException(
+ "PSS digest and MGF1 digest must match.");
+ }
+ int expectedSaltLength = "SHA256".equals(hash) ? 32
+ : "SHA384".equals(hash) ? 48 : 64;
+ if (spec.getSaltLength() != expectedSaltLength) {
+ throw new InvalidAlgorithmParameterException(
+ "PSS salt length must equal the digest length.");
+ }
+ if (spec.getTrailerField() != 1) {
+ throw new InvalidAlgorithmParameterException(
+ "Only trailer field 1 is supported for RSASSA-PSS.");
+ }
+ }
+}
diff --git a/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/CngX509ExtendedKeyManager.java b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/CngX509ExtendedKeyManager.java
new file mode 100644
index 00000000..77a993d3
--- /dev/null
+++ b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/CngX509ExtendedKeyManager.java
@@ -0,0 +1,95 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import javax.net.ssl.SSLEngine;
+import javax.net.ssl.X509ExtendedKeyManager;
+import java.net.Socket;
+import java.security.Principal;
+import java.security.PrivateKey;
+import java.security.cert.X509Certificate;
+
+final class CngX509ExtendedKeyManager extends X509ExtendedKeyManager {
+
+ private static final String ALIAS = "msal-keyguard-mtls";
+ private final CngRsaPrivateKey privateKey;
+ private final X509Certificate[] certificateChain;
+
+ CngX509ExtendedKeyManager(
+ CngRsaPrivateKey privateKey,
+ X509Certificate certificate) {
+ this.privateKey = privateKey;
+ this.certificateChain = new X509Certificate[]{certificate};
+ }
+
+ @Override
+ public String[] getClientAliases(String keyType, Principal[] issuers) {
+ return supportsKeyType(keyType) ? new String[]{ALIAS} : null;
+ }
+
+ @Override
+ public String chooseClientAlias(
+ String[] keyTypes,
+ Principal[] issuers,
+ Socket socket) {
+ return supportsAnyKeyType(keyTypes) ? ALIAS : null;
+ }
+
+ @Override
+ public String chooseEngineClientAlias(
+ String[] keyTypes,
+ Principal[] issuers,
+ SSLEngine engine) {
+ return supportsAnyKeyType(keyTypes) ? ALIAS : null;
+ }
+
+ @Override
+ public X509Certificate[] getCertificateChain(String alias) {
+ return ALIAS.equals(alias) ? certificateChain.clone() : null;
+ }
+
+ @Override
+ public PrivateKey getPrivateKey(String alias) {
+ return ALIAS.equals(alias) ? privateKey : null;
+ }
+
+ @Override
+ public String[] getServerAliases(String keyType, Principal[] issuers) {
+ return null;
+ }
+
+ @Override
+ public String chooseServerAlias(
+ String keyType,
+ Principal[] issuers,
+ Socket socket) {
+ return null;
+ }
+
+ @Override
+ public String chooseEngineServerAlias(
+ String keyType,
+ Principal[] issuers,
+ SSLEngine engine) {
+ return null;
+ }
+
+ private static boolean supportsAnyKeyType(String[] keyTypes) {
+ if (keyTypes == null) {
+ return false;
+ }
+ for (String keyType : keyTypes) {
+ if (supportsKeyType(keyType)) {
+ return true;
+ }
+ }
+ return false;
+ }
+
+ private static boolean supportsKeyType(String keyType) {
+ return keyType != null
+ && ("RSA".equalsIgnoreCase(keyType)
+ || "RSASSA-PSS".equalsIgnoreCase(keyType));
+ }
+}
diff --git a/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/ImdsV2Client.java b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/ImdsV2Client.java
new file mode 100644
index 00000000..60418efe
--- /dev/null
+++ b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/ImdsV2Client.java
@@ -0,0 +1,264 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import com.microsoft.aad.msal4j.ManagedIdentityMtlsHttpRequest;
+import com.microsoft.aad.msal4j.ManagedIdentityMtlsHttpResponse;
+import com.microsoft.aad.msal4j.ManagedIdentityMtlsRequest;
+
+import java.net.URLEncoder;
+import java.util.LinkedHashMap;
+import java.util.List;
+import java.util.Map;
+
+final class ImdsV2Client {
+
+ private static final String IMDS_BASE = "http://169.254.169.254";
+ private static final String PLATFORM_METADATA_PATH =
+ "/metadata/identity/getplatformmetadata";
+ private static final String ISSUE_CREDENTIAL_PATH =
+ "/metadata/identity/issuecredential";
+ private static final String API_VERSION_QUERY = "cred-api-version=2.0";
+
+ private ImdsV2Client() {
+ }
+
+ static PlatformMetadata getPlatformMetadata(ManagedIdentityMtlsRequest request) {
+ ManagedIdentityMtlsHttpResponse httpResponse = execute(
+ request,
+ "GET",
+ buildUrl(PLATFORM_METADATA_PATH, request),
+ null);
+ validateImdsOrigin(httpResponse);
+ String response = httpResponse.body();
+ PlatformMetadata metadata = new PlatformMetadata(
+ extractString(response, "clientId"),
+ extractString(response, "tenantId"),
+ extractNestedString(response, "cuId", "vmId"),
+ extractNestedString(response, "cuId", "vmssId"),
+ extractString(response, "attestationEndpoint"));
+ if (isBlank(metadata.clientId)
+ || isBlank(metadata.tenantId)
+ || isBlank(metadata.cuId())
+ || (request.attestationEnabled()
+ && isBlank(metadata.attestationEndpoint))) {
+ throw new MtlsMsiException(
+ "IMDS getplatformmetadata returned an incomplete attested KeyGuard contract.");
+ }
+ return metadata;
+ }
+
+ static CredentialResponse issueCredential(
+ ManagedIdentityMtlsRequest request,
+ String csr,
+ String attestationToken) {
+ if (request.attestationEnabled() && isBlank(attestationToken)) {
+ throw new MtlsMsiException(
+ "KeyGuard attestation failed; no attestation token was produced.");
+ }
+ String body = "{\"csr\":\"" + escapeJson(csr) + "\""
+ + (isBlank(attestationToken)
+ ? ""
+ : ",\"attestation_token\":\""
+ + escapeJson(attestationToken) + "\"")
+ + "}";
+ String response = execute(
+ request,
+ "POST",
+ buildUrl(ISSUE_CREDENTIAL_PATH, request),
+ body).body();
+ CredentialResponse credential = new CredentialResponse(
+ extractString(response, "certificate"),
+ extractString(response, "mtls_authentication_endpoint"),
+ extractString(response, "client_id"),
+ extractString(response, "tenant_id"),
+ extractString(response, "identity_type"));
+ if (isBlank(credential.certificate)
+ || isBlank(credential.mtlsAuthenticationEndpoint)
+ || isBlank(credential.clientId)
+ || isBlank(credential.tenantId)
+ || isBlank(credential.identityType)) {
+ throw new MtlsMsiException(
+ "IMDS issuecredential returned an incomplete binding credential.");
+ }
+ return credential;
+ }
+
+ private static ManagedIdentityMtlsHttpResponse execute(
+ ManagedIdentityMtlsRequest request,
+ String method,
+ String url,
+ String body) {
+ Map headers = new LinkedHashMap<>();
+ headers.put("Metadata", "true");
+ headers.put("x-ms-client-request-id", request.correlationId());
+ if (body != null) {
+ headers.put("Content-Type", "application/json");
+ }
+ ManagedIdentityMtlsHttpResponse response = request.httpClient().execute(
+ new ManagedIdentityMtlsHttpRequest(method, url, headers, body));
+ if (response.statusCode() != 200) {
+ throw new MtlsMsiException(
+ "IMDS " + method + " " + url + " failed with HTTP "
+ + response.statusCode() + ".");
+ }
+ return response;
+ }
+
+ private static void validateImdsOrigin(
+ ManagedIdentityMtlsHttpResponse response) {
+ for (Map.Entry> header
+ : response.headers().entrySet()) {
+ if (header.getKey() == null
+ || !"server".equalsIgnoreCase(header.getKey())) {
+ continue;
+ }
+ if (header.getValue() == null) {
+ continue;
+ }
+ for (String value : header.getValue()) {
+ String normalized = value == null
+ ? "" : value.trim().toUpperCase();
+ if ("IMDS".equals(normalized)
+ || normalized.startsWith("IMDS/")) {
+ return;
+ }
+ }
+ }
+ throw new MtlsMsiException(
+ "IMDS getplatformmetadata response did not contain the expected Server header.");
+ }
+
+ private static String buildUrl(String path, ManagedIdentityMtlsRequest request) {
+ StringBuilder url = new StringBuilder(IMDS_BASE)
+ .append(path)
+ .append('?')
+ .append(API_VERSION_QUERY);
+ if (!isBlank(request.identityQueryParameter())) {
+ url.append('&')
+ .append(encode(request.identityQueryParameter()))
+ .append('=')
+ .append(encode(request.identityQueryValue()));
+ }
+ return url.toString();
+ }
+
+ private static String encode(String value) {
+ try {
+ return URLEncoder.encode(value, "UTF-8");
+ } catch (Exception e) {
+ throw new MtlsMsiException("Unable to encode IMDS query parameter.", e);
+ }
+ }
+
+ static String extractString(String json, String key) {
+ if (json == null) {
+ return null;
+ }
+ String marker = "\"" + key + "\"";
+ int keyIndex = json.indexOf(marker);
+ if (keyIndex < 0) {
+ return null;
+ }
+ int colon = json.indexOf(':', keyIndex + marker.length());
+ if (colon < 0) {
+ return null;
+ }
+ int quote = skipWhitespace(json, colon + 1);
+ if (quote >= json.length() || json.charAt(quote) != '"') {
+ return null;
+ }
+ StringBuilder value = new StringBuilder();
+ for (int i = quote + 1; i < json.length(); i++) {
+ char c = json.charAt(i);
+ if (c == '"') {
+ return value.toString();
+ }
+ if (c == '\\' && i + 1 < json.length()) {
+ char escaped = json.charAt(++i);
+ value.append(escaped == 'n' ? '\n'
+ : escaped == 'r' ? '\r'
+ : escaped == 't' ? '\t' : escaped);
+ } else {
+ value.append(c);
+ }
+ }
+ return null;
+ }
+
+ private static String extractNestedString(String json, String objectKey, String key) {
+ String marker = "\"" + objectKey + "\"";
+ int objectIndex = json == null ? -1 : json.indexOf(marker);
+ if (objectIndex < 0) {
+ return null;
+ }
+ int start = json.indexOf('{', objectIndex + marker.length());
+ if (start < 0) {
+ return null;
+ }
+ int end = json.indexOf('}', start + 1);
+ return end < 0 ? null : extractString(json.substring(start, end + 1), key);
+ }
+
+ private static int skipWhitespace(String value, int index) {
+ while (index < value.length() && Character.isWhitespace(value.charAt(index))) {
+ index++;
+ }
+ return index;
+ }
+
+ private static String escapeJson(String value) {
+ return value.replace("\\", "\\\\").replace("\"", "\\\"");
+ }
+
+ private static boolean isBlank(String value) {
+ return value == null || value.trim().isEmpty();
+ }
+
+ static final class PlatformMetadata {
+ final String clientId;
+ final String tenantId;
+ final String vmId;
+ final String vmssId;
+ final String attestationEndpoint;
+
+ PlatformMetadata(
+ String clientId,
+ String tenantId,
+ String vmId,
+ String vmssId,
+ String attestationEndpoint) {
+ this.clientId = clientId;
+ this.tenantId = tenantId;
+ this.vmId = vmId;
+ this.vmssId = vmssId;
+ this.attestationEndpoint = attestationEndpoint;
+ }
+
+ String cuId() {
+ return isBlank(vmId) ? vmssId : vmId;
+ }
+ }
+
+ static final class CredentialResponse {
+ final String certificate;
+ final String mtlsAuthenticationEndpoint;
+ final String clientId;
+ final String tenantId;
+ final String identityType;
+
+ CredentialResponse(
+ String certificate,
+ String mtlsAuthenticationEndpoint,
+ String clientId,
+ String tenantId,
+ String identityType) {
+ this.certificate = certificate;
+ this.mtlsAuthenticationEndpoint = mtlsAuthenticationEndpoint;
+ this.clientId = clientId;
+ this.tenantId = tenantId;
+ this.identityType = identityType;
+ }
+ }
+}
diff --git a/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/KeyGuardManagedIdentityMtlsProvider.java b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/KeyGuardManagedIdentityMtlsProvider.java
new file mode 100644
index 00000000..4de5f3f8
--- /dev/null
+++ b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/KeyGuardManagedIdentityMtlsProvider.java
@@ -0,0 +1,282 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import com.microsoft.aad.msal4j.IManagedIdentityMtlsProvider;
+import com.microsoft.aad.msal4j.ManagedIdentityMtlsBinding;
+import com.microsoft.aad.msal4j.ManagedIdentityMtlsRequest;
+
+import java.io.ByteArrayInputStream;
+import java.math.BigInteger;
+import java.nio.charset.StandardCharsets;
+import java.security.MessageDigest;
+import java.security.interfaces.RSAPublicKey;
+import java.security.cert.CertificateFactory;
+import java.security.cert.X509Certificate;
+import java.util.ArrayList;
+import java.util.Base64;
+import java.util.List;
+import java.util.Map;
+import java.util.concurrent.ConcurrentHashMap;
+
+/**
+ * Windows KeyGuard/MAA implementation of the managed identity mTLS provider SPI.
+ */
+public final class KeyGuardManagedIdentityMtlsProvider
+ implements IManagedIdentityMtlsProvider {
+
+ private static final long ROTATION_BUFFER_MILLIS = 24L * 60L * 60L * 1000L;
+ private static final AttestationTokenCache ATTESTATION_CACHE =
+ new AttestationTokenCache();
+ private static final Map CURRENT =
+ new ConcurrentHashMap<>();
+ private static final Map> RETIRED =
+ new ConcurrentHashMap<>();
+ private static final Map LOCKS = new ConcurrentHashMap<>();
+
+ @Override
+ public ManagedIdentityMtlsBinding getOrCreateBinding(
+ ManagedIdentityMtlsRequest request) {
+ validateRequest(request);
+ String cacheKey = request.bindingCacheKey();
+ BindingGeneration cached = CURRENT.get(cacheKey);
+ if (canReturnWithoutCleanup(
+ cached,
+ RETIRED.containsKey(cacheKey))) {
+ return cached.binding;
+ }
+
+ Object lock = LOCKS.computeIfAbsent(cacheKey, ignored -> new Object());
+ synchronized (lock) {
+ cleanupRetired(cacheKey);
+ cached = CURRENT.get(cacheKey);
+ if (isCurrent(cached)) {
+ return cached.binding;
+ }
+ BindingGeneration created = createBinding(request);
+ BindingGeneration previous = CURRENT.put(cacheKey, created);
+ if (previous != null) {
+ RETIRED.computeIfAbsent(cacheKey, ignored -> new ArrayList<>())
+ .add(previous);
+ }
+ return created.binding;
+ }
+ }
+
+ private static BindingGeneration createBinding(
+ ManagedIdentityMtlsRequest request) {
+ ImdsV2Client.PlatformMetadata metadata =
+ ImdsV2Client.getPlatformMetadata(request);
+ validateSelectedIdentity(request, metadata);
+
+ String keyName = "MSALJavaMtls_" + shortHash(
+ request.bindingCacheKey() + "|" + metadata.cuId());
+ CngRsaPrivateKey privateKey = CngKeyGuard.getOrCreateKey(keyName);
+ try {
+ BigInteger[] publicKey =
+ CngKeyGuard.exportPublicKey(privateKey.nativeHandle());
+ String attestationKeyId = hashBytes(
+ CngKeyGuard.exportPublicKeyBytes(privateKey.nativeHandle()));
+ String csr = Pkcs10Builder.generate(
+ privateKey.nativeHandle(),
+ publicKey[0],
+ publicKey[1].intValue(),
+ metadata.clientId,
+ metadata.tenantId,
+ metadata.vmId,
+ metadata.vmssId);
+ String attestationToken = request.attestationEnabled()
+ ? ATTESTATION_CACHE.getOrAttest(
+ metadata.attestationEndpoint,
+ attestationKeyId,
+ () -> CngKeyGuard.getAttestationToken(
+ privateKey.nativeHandle(),
+ metadata.attestationEndpoint,
+ metadata.clientId))
+ : null;
+ ImdsV2Client.CredentialResponse credential =
+ ImdsV2Client.issueCredential(request, csr, attestationToken);
+ validateCredential(metadata, request, credential);
+
+ X509Certificate certificate = parseCertificate(credential.certificate);
+ validateCertificateMatchesKey(certificate, publicKey);
+ KeyGuardMtlsBindingContext context =
+ new KeyGuardMtlsBindingContext(privateKey, certificate);
+ String endpoint = trimTrailingSlash(
+ credential.mtlsAuthenticationEndpoint)
+ + "/" + trimSlashes(credential.tenantId)
+ + "/oauth2/v2.0/token";
+ ManagedIdentityMtlsBinding binding = new ManagedIdentityMtlsBinding(
+ context,
+ credential.clientId,
+ endpoint);
+ return new BindingGeneration(binding, context, certificate.getNotAfter().getTime());
+ } catch (RuntimeException e) {
+ privateKey.close();
+ throw e;
+ }
+ }
+
+ private static void validateRequest(ManagedIdentityMtlsRequest request) {
+ if (request == null
+ || request.httpClient() == null
+ || isBlank(request.bindingCacheKey())
+ || isBlank(request.correlationId())) {
+ throw new MtlsMsiException(
+ "Managed identity mTLS provider request is incomplete.");
+ }
+ if ((request.identityQueryParameter() == null)
+ != (request.identityQueryValue() == null)) {
+ throw new MtlsMsiException(
+ "Managed identity selector name and value must be supplied together.");
+ }
+ }
+
+ private static void validateSelectedIdentity(
+ ManagedIdentityMtlsRequest request,
+ ImdsV2Client.PlatformMetadata metadata) {
+ if ("client_id".equals(request.identityQueryParameter())
+ && !metadata.clientId.equalsIgnoreCase(request.identityQueryValue())) {
+ throw new MtlsMsiException(
+ "IMDS returned a different managed identity than the requested client ID.");
+ }
+ }
+
+ private static void validateCredential(
+ ImdsV2Client.PlatformMetadata metadata,
+ ManagedIdentityMtlsRequest request,
+ ImdsV2Client.CredentialResponse credential) {
+ if (!metadata.clientId.equalsIgnoreCase(credential.clientId)
+ || !metadata.tenantId.equalsIgnoreCase(credential.tenantId)) {
+ throw new MtlsMsiException(
+ "IMDS issuecredential identity does not match platform metadata.");
+ }
+ boolean systemAssigned = request.identityQueryParameter() == null;
+ String expectedIdentityType =
+ systemAssigned ? "SystemAssigned" : "UserAssigned";
+ if (!expectedIdentityType.equalsIgnoreCase(credential.identityType)) {
+ throw new MtlsMsiException(
+ "IMDS issuecredential returned unexpected identity_type '"
+ + credential.identityType + "'.");
+ }
+ }
+
+ private static X509Certificate parseCertificate(String encoded) {
+ try {
+ byte[] der = Base64.getDecoder().decode(encoded);
+ return (X509Certificate) CertificateFactory.getInstance("X.509")
+ .generateCertificate(new ByteArrayInputStream(der));
+ } catch (Exception e) {
+ throw new MtlsMsiException(
+ "Unable to parse the IMDS binding certificate.", e);
+ }
+ }
+
+ private static void validateCertificateMatchesKey(
+ X509Certificate certificate,
+ BigInteger[] publicKey) {
+ if (!(certificate.getPublicKey() instanceof RSAPublicKey)) {
+ throw new MtlsMsiException(
+ "IMDS binding certificate does not contain an RSA public key.");
+ }
+ RSAPublicKey certificateKey = (RSAPublicKey) certificate.getPublicKey();
+ if (!publicKey[0].equals(certificateKey.getModulus())
+ || !publicKey[1].equals(certificateKey.getPublicExponent())) {
+ throw new MtlsMsiException(
+ "IMDS binding certificate does not match the KeyGuard key.");
+ }
+ }
+
+ private static boolean isCurrent(BindingGeneration generation) {
+ return generation != null
+ && isCertificateCurrent(
+ generation.notAfterMillis,
+ System.currentTimeMillis());
+ }
+
+ static boolean canReturnWithoutCleanup(
+ BindingGeneration generation,
+ boolean hasRetiredGenerations) {
+ return isCurrent(generation) && !hasRetiredGenerations;
+ }
+
+ static boolean isCertificateCurrent(long notAfterMillis, long nowMillis) {
+ return nowMillis < notAfterMillis - ROTATION_BUFFER_MILLIS;
+ }
+
+ private static void cleanupRetired(String cacheKey) {
+ List generations = RETIRED.get(cacheKey);
+ if (generations == null) {
+ return;
+ }
+ long now = System.currentTimeMillis();
+ List retained = new ArrayList<>();
+ for (BindingGeneration generation : generations) {
+ if (now >= generation.notAfterMillis) {
+ generation.context.closeNativeKey();
+ } else {
+ retained.add(generation);
+ }
+ }
+ if (retained.isEmpty()) {
+ RETIRED.remove(cacheKey);
+ } else {
+ RETIRED.put(cacheKey, retained);
+ }
+ }
+
+ private static String shortHash(String value) {
+ try {
+ byte[] hash = MessageDigest.getInstance("SHA-256")
+ .digest(value.getBytes(StandardCharsets.UTF_8));
+ return Base64.getUrlEncoder().withoutPadding()
+ .encodeToString(hash).substring(0, 32);
+ } catch (Exception e) {
+ throw new MtlsMsiException("Unable to derive the KeyGuard key name.", e);
+ }
+ }
+
+ private static String hashBytes(byte[] value) {
+ try {
+ return Base64.getUrlEncoder().withoutPadding().encodeToString(
+ MessageDigest.getInstance("SHA-256").digest(value));
+ } catch (Exception e) {
+ throw new MtlsMsiException(
+ "Unable to derive the KeyGuard attestation cache key.", e);
+ }
+ }
+
+ private static String trimTrailingSlash(String value) {
+ while (value.endsWith("/")) {
+ value = value.substring(0, value.length() - 1);
+ }
+ return value;
+ }
+
+ private static String trimSlashes(String value) {
+ while (value.startsWith("/")) {
+ value = value.substring(1);
+ }
+ return trimTrailingSlash(value);
+ }
+
+ private static boolean isBlank(String value) {
+ return value == null || value.trim().isEmpty();
+ }
+
+ static final class BindingGeneration {
+ final ManagedIdentityMtlsBinding binding;
+ final KeyGuardMtlsBindingContext context;
+ final long notAfterMillis;
+
+ BindingGeneration(
+ ManagedIdentityMtlsBinding binding,
+ KeyGuardMtlsBindingContext context,
+ long notAfterMillis) {
+ this.binding = binding;
+ this.context = context;
+ this.notAfterMillis = notAfterMillis;
+ }
+ }
+}
diff --git a/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/KeyGuardMtlsBindingContext.java b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/KeyGuardMtlsBindingContext.java
new file mode 100644
index 00000000..2c8b2d28
--- /dev/null
+++ b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/KeyGuardMtlsBindingContext.java
@@ -0,0 +1,77 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import com.microsoft.aad.msal4j.IMtlsBindingContext;
+
+import javax.net.ssl.KeyManager;
+import javax.net.ssl.SSLContext;
+import java.security.MessageDigest;
+import java.security.cert.X509Certificate;
+import java.util.Base64;
+
+final class KeyGuardMtlsBindingContext implements IMtlsBindingContext {
+
+ private final CngRsaPrivateKey privateKey;
+ private final X509Certificate certificate;
+ private final SSLContext sslContext;
+ private final String keyId;
+
+ KeyGuardMtlsBindingContext(
+ CngRsaPrivateKey privateKey,
+ X509Certificate certificate) {
+ this.privateKey = privateKey;
+ this.certificate = certificate;
+ this.keyId = calculateKeyId(certificate);
+ this.sslContext = createSslContext(privateKey, certificate);
+ }
+
+ @Override
+ public SSLContext sslContext() {
+ return sslContext;
+ }
+
+ @Override
+ public X509Certificate bindingCertificate() {
+ return certificate;
+ }
+
+ @Override
+ public String keyId() {
+ return keyId;
+ }
+
+ void closeNativeKey() {
+ privateKey.close();
+ }
+
+ private static SSLContext createSslContext(
+ CngRsaPrivateKey privateKey,
+ X509Certificate certificate) {
+ try {
+ CngProvider.installIfAbsent();
+ SSLContext context = SSLContext.getInstance("TLSv1.2");
+ context.init(
+ new KeyManager[]{
+ new CngX509ExtendedKeyManager(privateKey, certificate)
+ },
+ null,
+ null);
+ return context;
+ } catch (Exception e) {
+ privateKey.close();
+ throw new MtlsMsiException("Unable to create KeyGuard JSSE SSLContext.", e);
+ }
+ }
+
+ static String calculateKeyId(X509Certificate certificate) {
+ try {
+ byte[] digest = MessageDigest.getInstance("SHA-256")
+ .digest(certificate.getEncoded());
+ return Base64.getUrlEncoder().withoutPadding().encodeToString(digest);
+ } catch (Exception e) {
+ throw new MtlsMsiException("Unable to calculate binding certificate key ID.", e);
+ }
+ }
+}
diff --git a/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/MtlsMsiException.java b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/MtlsMsiException.java
new file mode 100644
index 00000000..a788f1de
--- /dev/null
+++ b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/MtlsMsiException.java
@@ -0,0 +1,18 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+/**
+ * Thrown when the KeyGuard-backed managed identity mTLS binding flow fails.
+ */
+public class MtlsMsiException extends RuntimeException {
+
+ public MtlsMsiException(String message) {
+ super(message);
+ }
+
+ public MtlsMsiException(String message, Throwable cause) {
+ super(message, cause);
+ }
+}
diff --git a/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/NCryptLibrary.java b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/NCryptLibrary.java
new file mode 100644
index 00000000..b1ab7148
--- /dev/null
+++ b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/NCryptLibrary.java
@@ -0,0 +1,118 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import com.sun.jna.Library;
+import com.sun.jna.Native;
+import com.sun.jna.Pointer;
+import com.sun.jna.Structure;
+import com.sun.jna.WString;
+import com.sun.jna.ptr.IntByReference;
+import com.sun.jna.ptr.PointerByReference;
+
+import java.util.Arrays;
+import java.util.List;
+
+/**
+ * JNA binding for {@code ncrypt.dll} — Windows CNG (Cryptography Next Generation) key
+ * storage and signing operations.
+ *
+ * Function signatures mirror MSAL.NET's {@code WindowsCngKeyOperations} and msal-go's
+ * {@code cng_windows.go}. All NCrypt functions follow the Windows x64 calling convention
+ * (which equals cdecl on x64).
+ */
+interface NCryptLibrary extends Library {
+
+ NCryptLibrary INSTANCE = Native.load("ncrypt", NCryptLibrary.class);
+
+ // ─── NCrypt constants ──────────────────────────────────────────────────────
+
+ int ERROR_SUCCESS = 0;
+
+ int NCRYPT_SILENT_FLAG = 0x00000040;
+ int NCRYPT_OVERWRITE_KEY_FLAG = 0x00000080;
+ int NCRYPT_MACHINE_KEY_FLAG = 0x00000020; // not used (USER scope only)
+ int NCRYPT_USE_VIRTUAL_ISOLATION_FLAG = 0x00020000; // VBS KeyGuard
+ int NCRYPT_USE_PER_BOOT_KEY_FLAG = 0x00040000; // ephemeral per boot
+ int NCRYPT_ALLOW_EXPORT_NONE = 0; // non-exportable
+
+ int NCRYPT_PAD_PKCS1_FLAG = 0x00000002;
+ int NCRYPT_PAD_PSS_FLAG = 0x00000008;
+
+ // ─── Padding info structures ───────────────────────────────────────────────
+
+ /** Maps to {@code BCRYPT_PKCS1_PADDING_INFO} — used with NCRYPT_PAD_PKCS1_FLAG. */
+ class BcryptPkcs1PaddingInfo extends Structure {
+ /** Algorithm name for the hash (e.g. L"SHA256"). LPCWSTR in C. */
+ public WString pszAlgId;
+
+ public BcryptPkcs1PaddingInfo(String algName) {
+ pszAlgId = new WString(algName);
+ write();
+ }
+
+ @Override
+ protected List getFieldOrder() {
+ return Arrays.asList("pszAlgId");
+ }
+ }
+
+ /** Maps to {@code BCRYPT_PSS_PADDING_INFO} — used with NCRYPT_PAD_PSS_FLAG. */
+ class BcryptPssPaddingInfo extends Structure {
+ /** Algorithm name for the hash (e.g. L"SHA256"). LPCWSTR in C. */
+ public WString pszAlgId;
+ /** Salt length in bytes. Typically equals hash output length for RSASSA-PSS. */
+ public int cbSalt;
+
+ public BcryptPssPaddingInfo(String algName, int saltLen) {
+ pszAlgId = new WString(algName);
+ cbSalt = saltLen;
+ write();
+ }
+
+ @Override
+ protected List getFieldOrder() {
+ return Arrays.asList("pszAlgId", "cbSalt");
+ }
+ }
+
+ // ─── NCrypt API ────────────────────────────────────────────────────────────
+
+ int NCryptOpenStorageProvider(PointerByReference phProvider, WString pszProviderName, int dwFlags);
+
+ int NCryptOpenKey(Pointer hProvider, PointerByReference phKey, WString pszKeyName,
+ int dwLegacyKeySpec, int dwFlags);
+
+ int NCryptCreatePersistedKey(Pointer hProvider, PointerByReference phKey,
+ WString pszAlgId, WString pszKeyName,
+ int dwLegacyKeySpec, int dwFlags);
+
+ int NCryptSetProperty(Pointer hObject, WString pszProperty,
+ byte[] pbInput, int cbInput, int dwFlags);
+
+ int NCryptGetProperty(Pointer hObject, WString pszProperty,
+ byte[] pbOutput, int cbOutput,
+ IntByReference pcbResult, int dwFlags);
+
+ int NCryptFinalizeKey(Pointer hKey, int dwFlags);
+
+ /** First call: pass {@code pbOutput=null, cbOutput=0} to query required buffer size. */
+ int NCryptExportKey(Pointer hKey, Pointer hExportKey, WString pszBlobType,
+ Pointer pParameterList, Pointer pbOutput, int cbOutput,
+ IntByReference pcbResult, int dwFlags);
+
+ /**
+ * First call: pass {@code pbSignature=null, cbSignature=0} to get required buffer size
+ * (returned in {@code pcbResult}).
+ * Second call: pass a {@code Memory} buffer of that size.
+ */
+ int NCryptSignHash(Pointer hKey, Pointer pPaddingInfo,
+ byte[] pbHashValue, int cbHashValue,
+ Pointer pbSignature, int cbSignature,
+ IntByReference pcbResult, int dwFlags);
+
+ int NCryptFreeObject(Pointer hObject);
+
+ int NCryptDeleteKey(Pointer hKey, int dwFlags);
+}
diff --git a/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/Pkcs10Builder.java b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/Pkcs10Builder.java
new file mode 100644
index 00000000..3cc8e919
--- /dev/null
+++ b/msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/Pkcs10Builder.java
@@ -0,0 +1,326 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import com.sun.jna.Pointer;
+
+import java.io.ByteArrayOutputStream;
+import java.math.BigInteger;
+import java.security.MessageDigest;
+import java.security.NoSuchAlgorithmException;
+import java.util.Base64;
+
+/**
+ * Builds a PKCS#10 Certification Request (CSR) that matches the format produced by
+ * MSAL.NET and msal-go for the Azure IMDSv2 {@code /issuecredential} endpoint.
+ *
+ * CSR Structure
+ *
+ * CertificationRequest ::= SEQUENCE {
+ * certificationRequestInfo CertificationRequestInfo,
+ * signatureAlgorithm AlgorithmIdentifier, -- RSASSA-PSS with SHA-256 params
+ * signature BIT STRING
+ * }
+ *
+ * CertificationRequestInfo ::= SEQUENCE {
+ * version INTEGER { v1(0) }
+ * subject Name -- CN={clientId}, DC={tenantId}
+ * subjectPKInfo SubjectPublicKeyInfo
+ * attributes [0] IMPLICIT SET OF -- OID 1.3.6.1.4.1.311.90.2.10 = cuId JSON
+ * }
+ *
+ *
+ * Signing
+ * Signature: RSASSA-PSS with SHA-256, salt length = 32 bytes (hash output length).
+ * Signing is delegated to {@link CngKeyGuard#signPss} so the non-exportable KeyGuard key
+ * never leaves CNG.
+ *
+ * This is a pure-Java port of msal-go's {@code generateCSR()} in {@code imdsv2.go},
+ * using manual DER encoding to avoid adding external ASN.1 library dependencies.
+ */
+final class Pkcs10Builder {
+
+ private Pkcs10Builder() {}
+
+ // ─── OIDs (pre-encoded DER) ────────────────────────────────────────────────
+
+ // rsaEncryption: 1.2.840.113549.1.1.1
+ private static final byte[] OID_RSA_ENCRYPTION = hexToBytes("2a864886f70d010101");
+ // sha256: 2.16.840.1.101.3.4.2.1
+ private static final byte[] OID_SHA256 = hexToBytes("608648016503040201");
+ // mgf1: 1.2.840.113549.1.1.8
+ private static final byte[] OID_MGF1 = hexToBytes("2a864886f70d010108");
+ // id-RSASSA-PSS: 1.2.840.113549.1.1.10
+ private static final byte[] OID_RSASSA_PSS = hexToBytes("2a864886f70d01010a");
+ // commonName: 2.5.4.3
+ private static final byte[] OID_COMMON_NAME = hexToBytes("5504 03".replace(" ", ""));
+ // domainComponent: 0.9.2342.19200300.100.1.25
+ private static final byte[] OID_DOMAIN_COMPONENT = hexToBytes("0992268993f22c6401 19".replace(" ", ""));
+ // cuId attribute: 1.3.6.1.4.1.311.90.2.10
+ private static final byte[] OID_CU_ID = hexToBytes("2b060104018237 5a02 0a".replace(" ", ""));
+
+ // ─── Public API ───────────────────────────────────────────────────────────
+
+ /**
+ * Generates a PKCS#10 CSR and returns it as standard Base64-encoded DER
+ * (no PEM headers), ready to be placed in the {@code csr} field of the
+ * IMDS {@code /issuecredential} JSON request.
+ *
+ * @param keyHandle CNG key handle (the private key — signs the CSR TBS bytes)
+ * @param modulus RSA public key modulus (from {@link CngKeyGuard#exportPublicKey})
+ * @param publicExp RSA public exponent
+ * @param clientId managed identity client ID → CN in subject
+ * @param tenantId tenant GUID → DC in subject
+ * @param vmId VM ID for the cuId attribute ({@code cuId.vmId}); may be null
+ * @param vmssId VMSS ID for the cuId attribute; may be null
+ * @return Base64-encoded DER of the PKCS#10 CSR
+ */
+ static String generate(Pointer keyHandle, BigInteger modulus, int publicExp,
+ String clientId, String tenantId, String vmId, String vmssId)
+ throws MtlsMsiException {
+
+ // --- SubjectPublicKeyInfo ------------------------------------------
+ byte[] spki = buildSpki(modulus, publicExp);
+
+ // --- Subject: CN={clientId}, DC={tenantId} -------------------------
+ byte[] subject = buildSubject(clientId, tenantId);
+
+ // --- cuId attribute ------------------------------------------------
+ byte[] cuIdJson = buildCuIdJson(vmId, vmssId);
+ byte[] attributes = buildCuIdAttribute(cuIdJson);
+
+ // --- CertificationRequestInfo SEQUENCE -----------------------------
+ byte[] version = derInteger(new byte[]{0x00}); // INTEGER v1(0)
+ byte[] certReqInfo = derSequence(concat(version, subject, spki, attributes));
+
+ // --- Sign with RSASSA-PSS SHA-256 (salt=32) -----------------------
+ byte[] tbs;
+ try {
+ tbs = MessageDigest.getInstance("SHA-256").digest(certReqInfo);
+ } catch (NoSuchAlgorithmException e) {
+ throw new MtlsMsiException("SHA-256 not available: " + e.getMessage(), e);
+ }
+ byte[] sig = CngKeyGuard.signPss(keyHandle, tbs, "SHA256", 32);
+
+ // --- AlgorithmIdentifier for RSASSA-PSS ----------------------------
+ byte[] sigAlgId = buildPssAlgorithmIdentifier();
+
+ // --- BIT STRING wrapping the signature -----------------------------
+ byte[] sigBitString = derBitString(sig);
+
+ // --- Final CertificationRequest SEQUENCE ---------------------------
+ byte[] csr = derSequence(concat(certReqInfo, sigAlgId, sigBitString));
+
+ return Base64.getEncoder().encodeToString(csr);
+ }
+
+ // ─── DER building blocks ──────────────────────────────────────────────────
+
+ /** DER SEQUENCE */
+ static byte[] derSequence(byte[] content) {
+ return derTagLen(0x30, content);
+ }
+
+ /** DER SET */
+ private static byte[] derSet(byte[] content) {
+ return derTagLen(0x31, content);
+ }
+
+ /** DER INTEGER from raw bytes (big-endian, with sign byte if high bit set) */
+ private static byte[] derInteger(byte[] value) {
+ // Add leading 0x00 if high bit is set (unsigned → signed two's complement).
+ byte[] content = (value[0] & 0x80) != 0
+ ? concat(new byte[]{0x00}, value)
+ : value;
+ return derTagLen(0x02, content);
+ }
+
+ /** DER OBJECT IDENTIFIER from pre-encoded OID value bytes */
+ private static byte[] derOid(byte[] oidBytes) {
+ return derTagLen(0x06, oidBytes);
+ }
+
+ /** DER UTF8String */
+ private static byte[] derUtf8String(String s) {
+ byte[] bytes = s.getBytes(java.nio.charset.StandardCharsets.UTF_8);
+ return derTagLen(0x0C, bytes);
+ }
+
+ /** DER BIT STRING — prepend 0x00 (zero unused bits) */
+ static byte[] derBitString(byte[] data) {
+ byte[] content = new byte[data.length + 1];
+ content[0] = 0x00;
+ System.arraycopy(data, 0, content, 1, data.length);
+ return derTagLen(0x03, content);
+ }
+
+ /** DER NULL */
+ private static final byte[] DER_NULL = {0x05, 0x00};
+
+ /** Context-specific explicit tag [N] wrapping content */
+ private static byte[] contextExplicit(int n, byte[] content) {
+ return derTagLen(0xA0 | n, content);
+ }
+
+ /** Context-specific implicit tag [N] wrapping content */
+ private static byte[] contextImplicit(int n, byte[] content) {
+ return derTagLen(0x80 | n, content);
+ }
+
+ /** Writes tag + DER length + content */
+ private static byte[] derTagLen(int tag, byte[] content) {
+ ByteArrayOutputStream out = new ByteArrayOutputStream();
+ out.write(tag);
+ int len = content.length;
+ if (len < 0x80) {
+ out.write(len);
+ } else if (len < 0x100) {
+ out.write(0x81);
+ out.write(len);
+ } else if (len < 0x10000) {
+ out.write(0x82);
+ out.write((len >> 8) & 0xFF);
+ out.write(len & 0xFF);
+ } else {
+ out.write(0x83);
+ out.write((len >> 16) & 0xFF);
+ out.write((len >> 8) & 0xFF);
+ out.write(len & 0xFF);
+ }
+ try { out.write(content); } catch (java.io.IOException ignored) {}
+ return out.toByteArray();
+ }
+
+ // ─── Component builders ───────────────────────────────────────────────────
+
+ /**
+ * SubjectPublicKeyInfo ::= SEQUENCE { algorithm AlgorithmIdentifier, subjectPublicKey BIT STRING }
+ * AlgorithmIdentifier for RSA: SEQUENCE { OID rsaEncryption, NULL }
+ * Public key: BIT STRING containing RSAPublicKey SEQUENCE { modulus INTEGER, publicExp INTEGER }
+ */
+ private static byte[] buildSpki(BigInteger modulus, int publicExp) {
+ // RSAPublicKey SEQUENCE { modulus INTEGER, publicExp INTEGER }
+ byte[] modBytes = modulus.toByteArray();
+ byte[] expBytes = BigInteger.valueOf(publicExp).toByteArray();
+ byte[] rsaPublicKey = derSequence(concat(derInteger(modBytes), derInteger(expBytes)));
+
+ // AlgorithmIdentifier for rsaEncryption
+ byte[] algId = derSequence(concat(derOid(OID_RSA_ENCRYPTION), DER_NULL));
+
+ // SubjectPublicKeyInfo
+ return derSequence(concat(algId, derBitString(rsaPublicKey)));
+ }
+
+ /**
+ * Name ::= SEQUENCE { RDN SEQUENCE { AttributeTypeAndValue SEQUENCE { OID, value } } }
+ * Subject: CN={clientId}, DC={tenantId}
+ * Matches msal-go: pkix.Name{CommonName: clientId, ExtraNames: []pkix.AttributeTypeAndValue{{Type: dcOID, Value: tenantId}}}
+ */
+ private static byte[] buildSubject(String clientId, String tenantId) {
+ // AttributeTypeAndValue SEQUENCE { OID commonName, UTF8String clientId }
+ byte[] cnAttr = derSequence(concat(derOid(OID_COMMON_NAME), derUtf8String(clientId)));
+ byte[] cnRdn = derSet(cnAttr);
+
+ // AttributeTypeAndValue SEQUENCE { OID domainComponent, UTF8String tenantId }
+ byte[] dcAttr = derSequence(concat(derOid(OID_DOMAIN_COMPONENT), derUtf8String(tenantId)));
+ byte[] dcRdn = derSet(dcAttr);
+
+ // Name = SEQUENCE of RDNs
+ return derSequence(concat(cnRdn, dcRdn));
+ }
+
+ /**
+ * Builds the cuId JSON string. Matches msal-go's json.Marshal(cuID):
+ * {@code {"vmId":"","vmssId":""}} with omitempty semantics.
+ */
+ private static byte[] buildCuIdJson(String vmId, String vmssId) {
+ StringBuilder sb = new StringBuilder("{");
+ boolean first = true;
+ if (vmId != null && !vmId.isEmpty()) {
+ sb.append("\"vmId\":\"").append(vmId).append("\"");
+ first = false;
+ }
+ if (vmssId != null && !vmssId.isEmpty()) {
+ if (!first) sb.append(",");
+ sb.append("\"vmssId\":\"").append(vmssId).append("\"");
+ }
+ sb.append("}");
+ return sb.toString().getBytes(java.nio.charset.StandardCharsets.UTF_8);
+ }
+
+ /**
+ * CertificationRequestInfo attributes [0]:
+ * [0] CONSTRUCTED {
+ * SEQUENCE { OID 1.3.6.1.4.1.311.90.2.10, SET { UTF8String(cuIdJson) } }
+ * }
+ *
+ * Per PKCS#10, {@code [0] IMPLICIT Attributes} — IMPLICIT tagging of a constructed type
+ * keeps the constructed bit, so the tag byte is {@code 0xA0} (context-specific, constructed).
+ * Mirrors msal-go's {@code buildCuIDAttribute()} which uses
+ * {@code asn1.RawValue{Class: ClassContextSpecific, Tag: 0, IsCompound: true}}.
+ */
+ private static byte[] buildCuIdAttribute(byte[] cuIdJsonBytes) {
+ byte[] utf8Str = derTagLen(0x0C, cuIdJsonBytes); // UTF8String
+ byte[] valueSet = derSet(utf8Str); // SET { UTF8String }
+ byte[] attrSeq = derSequence(concat(derOid(OID_CU_ID), valueSet)); // SEQUENCE { OID, SET }
+ return contextExplicit(0, attrSeq); // [0] CONSTRUCTED { SEQUENCE } — 0xA0 tag
+ }
+
+ /**
+ * AlgorithmIdentifier for RSASSA-PSS with SHA-256:
+ * SEQUENCE {
+ * OID id-RSASSA-PSS,
+ * SEQUENCE { -- RSASSA-PSS-params
+ * [0] SEQUENCE { OID sha-256, NULL }, -- hashAlgorithm
+ * [1] SEQUENCE { OID mgf1, SEQUENCE { OID sha-256, NULL } }, -- maskGenAlgorithm
+ * [2] INTEGER 32 -- saltLength
+ * }
+ * }
+ * Matches msal-go's explicit PSS AlgorithmIdentifier.
+ */
+ private static byte[] buildPssAlgorithmIdentifier() {
+ // sha256AlgID: SEQUENCE { OID sha256, NULL }
+ byte[] sha256AlgId = derSequence(concat(derOid(OID_SHA256), DER_NULL));
+
+ // hashAlgorithm [0]: sha256AlgID
+ byte[] hashAlgorithm = contextExplicit(0, sha256AlgId);
+
+ // mgf1AlgID: SEQUENCE { OID mgf1, sha256AlgID }
+ byte[] mgf1AlgId = derSequence(concat(derOid(OID_MGF1), sha256AlgId));
+ // maskGenAlgorithm [1]: mgf1AlgID
+ byte[] maskGenAlgorithm = contextExplicit(1, mgf1AlgId);
+
+ // saltLength [2]: INTEGER 32
+ byte[] saltLength = contextExplicit(2, derInteger(new byte[]{32}));
+
+ // RSASSA-PSS-params SEQUENCE
+ byte[] pssParams = derSequence(concat(hashAlgorithm, maskGenAlgorithm, saltLength));
+
+ // AlgorithmIdentifier SEQUENCE { OID id-RSASSA-PSS, pssParams }
+ return derSequence(concat(derOid(OID_RSASSA_PSS), pssParams));
+ }
+
+ // ─── Utility ──────────────────────────────────────────────────────────────
+
+ private static byte[] concat(byte[]... arrays) {
+ int total = 0;
+ for (byte[] a : arrays) total += a.length;
+ byte[] result = new byte[total];
+ int offset = 0;
+ for (byte[] a : arrays) {
+ System.arraycopy(a, 0, result, offset, a.length);
+ offset += a.length;
+ }
+ return result;
+ }
+
+ private static byte[] hexToBytes(String hex) {
+ hex = hex.replace(" ", "");
+ byte[] result = new byte[hex.length() / 2];
+ for (int i = 0; i < result.length; i++) {
+ result[i] = (byte) Integer.parseInt(hex.substring(i * 2, i * 2 + 2), 16);
+ }
+ return result;
+ }
+}
diff --git a/msal4j-mtls-extensions/src/main/resources/META-INF/LICENSE-Microsoft.Azure.Security.KeyGuardAttestation.txt b/msal4j-mtls-extensions/src/main/resources/META-INF/LICENSE-Microsoft.Azure.Security.KeyGuardAttestation.txt
new file mode 100644
index 00000000..88040d88
--- /dev/null
+++ b/msal4j-mtls-extensions/src/main/resources/META-INF/LICENSE-Microsoft.Azure.Security.KeyGuardAttestation.txt
@@ -0,0 +1,21 @@
+MIT License
+
+Copyright (c) Microsoft Corporation. All rights reserved.
+
+Permission is hereby granted, free of charge, to any person obtaining a copy
+of this software and associated documentation files (the "Software"), to deal
+in the Software without restriction, including without limitation the rights
+to use, copy, modify, merge, publish, distribute, sublicense, and/or sell
+copies of the Software, and to permit persons to whom the Software is
+furnished to do so, subject to the following conditions:
+
+The above copyright notice and this permission notice shall be included in all
+copies or substantial portions of the Software.
+
+THE SOFTWARE IS PROVIDED "AS IS", WITHOUT WARRANTY OF ANY KIND, EXPRESS OR
+IMPLIED, INCLUDING BUT NOT LIMITED TO THE WARRANTIES OF MERCHANTABILITY,
+FITNESS FOR A PARTICULAR PURPOSE AND NONINFRINGEMENT. IN NO EVENT SHALL THE
+AUTHORS OR COPYRIGHT HOLDERS BE LIABLE FOR ANY CLAIM, DAMAGES OR OTHER
+LIABILITY, WHETHER IN AN ACTION OF CONTRACT, TORT OR OTHERWISE, ARISING FROM,
+OUT OF OR IN CONNECTION WITH THE SOFTWARE OR THE USE OR OTHER DEALINGS IN THE
+SOFTWARE.
diff --git a/msal4j-mtls-extensions/src/main/resources/META-INF/NOTICE-Microsoft.Azure.Security.KeyGuardAttestation.txt b/msal4j-mtls-extensions/src/main/resources/META-INF/NOTICE-Microsoft.Azure.Security.KeyGuardAttestation.txt
new file mode 100644
index 00000000..dcf47bfd
--- /dev/null
+++ b/msal4j-mtls-extensions/src/main/resources/META-INF/NOTICE-Microsoft.Azure.Security.KeyGuardAttestation.txt
@@ -0,0 +1,5 @@
+This artifact includes AttestationClientLib.dll from
+Microsoft.Azure.Security.KeyGuardAttestation 1.1.5.
+
+Copyright (c) Microsoft Corporation. All rights reserved.
+Licensed under the MIT License.
diff --git a/msal4j-mtls-extensions/src/main/resources/META-INF/native/win-x64/AttestationClientLib.dll b/msal4j-mtls-extensions/src/main/resources/META-INF/native/win-x64/AttestationClientLib.dll
new file mode 100644
index 00000000..bb2695ef
Binary files /dev/null and b/msal4j-mtls-extensions/src/main/resources/META-INF/native/win-x64/AttestationClientLib.dll differ
diff --git a/msal4j-mtls-extensions/src/main/resources/META-INF/services/com.microsoft.aad.msal4j.IManagedIdentityMtlsProvider b/msal4j-mtls-extensions/src/main/resources/META-INF/services/com.microsoft.aad.msal4j.IManagedIdentityMtlsProvider
new file mode 100644
index 00000000..66b24582
--- /dev/null
+++ b/msal4j-mtls-extensions/src/main/resources/META-INF/services/com.microsoft.aad.msal4j.IManagedIdentityMtlsProvider
@@ -0,0 +1 @@
+com.microsoft.aad.msal4j.mtls.KeyGuardManagedIdentityMtlsProvider
diff --git a/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/AttestationLibraryLoaderTest.java b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/AttestationLibraryLoaderTest.java
new file mode 100644
index 00000000..b4d8169c
--- /dev/null
+++ b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/AttestationLibraryLoaderTest.java
@@ -0,0 +1,43 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import org.junit.jupiter.api.Test;
+
+import java.io.InputStream;
+import java.security.DigestInputStream;
+import java.security.MessageDigest;
+
+import static org.junit.jupiter.api.Assertions.assertEquals;
+import static org.junit.jupiter.api.Assertions.assertNotNull;
+
+class AttestationLibraryLoaderTest {
+
+ @Test
+ void bundledLibraryMatchesMsalDotNetVersionAndHash() throws Exception {
+ assertEquals("1.1.5", AttestationLibraryLoader.VERSION);
+
+ MessageDigest digest = MessageDigest.getInstance("SHA-256");
+ try (InputStream resource = AttestationLibraryLoader.class.getResourceAsStream(
+ AttestationLibraryLoader.RESOURCE_PATH)) {
+ assertNotNull(resource);
+ try (DigestInputStream input = new DigestInputStream(resource, digest)) {
+ byte[] buffer = new byte[8192];
+ while (input.read(buffer) != -1) {
+ // Consume the complete resource.
+ }
+ }
+ }
+
+ assertEquals(AttestationLibraryLoader.SHA256, toHex(digest.digest()));
+ }
+
+ private static String toHex(byte[] bytes) {
+ StringBuilder result = new StringBuilder(bytes.length * 2);
+ for (byte value : bytes) {
+ result.append(String.format("%02x", value & 0xff));
+ }
+ return result.toString();
+ }
+}
diff --git a/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/AttestationTokenCacheTest.java b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/AttestationTokenCacheTest.java
new file mode 100644
index 00000000..6357ab8d
--- /dev/null
+++ b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/AttestationTokenCacheTest.java
@@ -0,0 +1,112 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import org.junit.jupiter.api.Test;
+
+import java.nio.charset.StandardCharsets;
+import java.time.Instant;
+import java.util.Base64;
+import java.util.Collections;
+import java.util.List;
+import java.util.concurrent.*;
+import java.util.concurrent.atomic.AtomicInteger;
+
+import static org.junit.jupiter.api.Assertions.*;
+
+class AttestationTokenCacheTest {
+
+ @Test
+ void cacheNormalizesEndpointAndIsScopedByKeyId() throws Exception {
+ AttestationTokenCache cache = new AttestationTokenCache();
+ AtomicInteger loads = new AtomicInteger();
+
+ String first = cache.getOrAttest("HTTPS://Example.COM/", "key-a",
+ () -> token(loads.incrementAndGet(), 3600));
+ String second = cache.getOrAttest("https://example.com", "key-a",
+ () -> token(loads.incrementAndGet(), 3600));
+ String otherKey = cache.getOrAttest("https://example.com", "key-b",
+ () -> token(loads.incrementAndGet(), 3600));
+
+ assertEquals(first, second);
+ assertNotEquals(first, otherKey);
+ assertEquals(2, loads.get());
+ }
+
+ @Test
+ void staleTokenInsideFreshnessBufferIsNotReused() throws Exception {
+ AttestationTokenCache cache = new AttestationTokenCache();
+ AtomicInteger loads = new AtomicInteger();
+
+ assertThrows(MtlsMsiException.class,
+ () -> cache.getOrAttest("https://example.com", "key",
+ () -> token(loads.incrementAndGet(), 299)));
+ cache.getOrAttest("https://example.com", "key",
+ () -> token(loads.incrementAndGet(), 3600));
+
+ assertEquals(2, loads.get());
+ }
+
+ @Test
+ void concurrentCallsAreSingleFlightPerKey() throws Exception {
+ AttestationTokenCache cache = new AttestationTokenCache();
+ AtomicInteger loads = new AtomicInteger();
+ ExecutorService executor = Executors.newFixedThreadPool(8);
+ try {
+ Callable task = () -> cache.getOrAttest(
+ "https://example.com",
+ "key",
+ () -> {
+ loads.incrementAndGet();
+ try {
+ Thread.sleep(25);
+ } catch (InterruptedException e) {
+ Thread.currentThread().interrupt();
+ throw new IllegalStateException(e);
+ }
+ return token(1, 3600);
+ });
+ List> futures = executor.invokeAll(
+ Collections.nCopies(8, task));
+ for (Future future : futures) {
+ assertEquals(futures.get(0).get(), future.get());
+ }
+ } finally {
+ executor.shutdownNow();
+ }
+
+ assertEquals(1, loads.get());
+ }
+
+ @Test
+ void failedLoadIsNotCached() {
+ AttestationTokenCache cache = new AttestationTokenCache();
+ AtomicInteger attempts = new AtomicInteger();
+
+ assertThrows(Exception.class, () -> cache.getOrAttest(
+ "https://example.com", "key", () -> {
+ attempts.incrementAndGet();
+ throw new IllegalStateException("attestation failed");
+ }));
+ assertThrows(Exception.class, () -> cache.getOrAttest(
+ "https://example.com", "key", () -> {
+ attempts.incrementAndGet();
+ throw new IllegalStateException("attestation failed");
+ }));
+
+ assertEquals(2, attempts.get());
+ }
+
+ private static String token(int marker, long validForSeconds) {
+ String header = encode("{\"alg\":\"none\"}");
+ String payload = encode("{\"marker\":" + marker + ",\"exp\":"
+ + (Instant.now().getEpochSecond() + validForSeconds) + "}");
+ return header + "." + payload + ".signature";
+ }
+
+ private static String encode(String value) {
+ return Base64.getUrlEncoder().withoutPadding().encodeToString(
+ value.getBytes(StandardCharsets.UTF_8));
+ }
+}
diff --git a/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/CngKeyGuardTest.java b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/CngKeyGuardTest.java
new file mode 100644
index 00000000..c34b7a7f
--- /dev/null
+++ b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/CngKeyGuardTest.java
@@ -0,0 +1,50 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import com.sun.jna.Pointer;
+import com.sun.jna.WString;
+import com.sun.jna.ptr.PointerByReference;
+import org.junit.jupiter.api.Test;
+
+import static org.junit.jupiter.api.Assertions.assertNull;
+import static org.mockito.ArgumentMatchers.any;
+import static org.mockito.ArgumentMatchers.anyInt;
+import static org.mockito.Mockito.mock;
+import static org.mockito.Mockito.verify;
+import static org.mockito.Mockito.when;
+
+class CngKeyGuardTest {
+
+ @Test
+ void staleOpenedKeyIsDeletedSoCallerRecreatesIt() {
+ NCryptLibrary nativeApi = mock(NCryptLibrary.class);
+ Pointer provider = Pointer.createConstant(11);
+ Pointer staleKey = Pointer.createConstant(42);
+ when(nativeApi.NCryptOpenKey(
+ any(),
+ any(),
+ any(),
+ anyInt(),
+ anyInt()))
+ .thenAnswer(invocation -> {
+ PointerByReference reference = invocation.getArgument(1);
+ reference.setValue(staleKey);
+ return NCryptLibrary.ERROR_SUCCESS;
+ });
+
+ Pointer opened = CngKeyGuard.openExistingUsableKey(
+ provider,
+ new WString("stale-key"),
+ NCryptLibrary.NCRYPT_SILENT_FLAG,
+ nativeApi,
+ key -> {
+ // Public export may have succeeded before this private-operation failure.
+ throw new MtlsMsiException("NCryptSignHash failed");
+ });
+
+ assertNull(opened);
+ verify(nativeApi).NCryptDeleteKey(staleKey, 0);
+ }
+}
diff --git a/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/CngProviderTest.java b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/CngProviderTest.java
new file mode 100644
index 00000000..bb8fd990
--- /dev/null
+++ b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/CngProviderTest.java
@@ -0,0 +1,119 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import org.junit.jupiter.api.AfterEach;
+import org.junit.jupiter.api.Test;
+import org.junit.jupiter.api.condition.EnabledOnOs;
+import org.junit.jupiter.api.condition.OS;
+
+import java.security.Provider;
+import java.security.Security;
+import java.security.Signature;
+import java.security.KeyPair;
+import java.security.KeyPairGenerator;
+import java.util.Arrays;
+
+import static org.junit.jupiter.api.Assertions.*;
+
+/**
+ * Unit tests for {@link CngProvider} — provider registration and service declarations.
+ *
+ * Requires Windows because loading {@link CngProvider} transitively loads
+ * {@link CngSignatureSpi} → {@link CngRsaPrivateKey} → {@link NCryptLibrary}
+ * ({@code ncrypt.dll}).
+ */
+@EnabledOnOs(OS.WINDOWS)
+class CngProviderTest {
+
+ @AfterEach
+ void removeCngProvider() {
+ // Remove after each test to prevent state from bleeding between tests
+ Security.removeProvider("CNG");
+ }
+
+ // ─── installIfAbsent ─────────────────────────────────────────────────────
+
+ @Test
+ void installIfAbsent_registersProviderByName() {
+ CngProvider.installIfAbsent();
+ assertNotNull(Security.getProvider("CNG"),
+ "CNG provider must be registered in the JVM Security list after installIfAbsent()");
+ }
+
+ @Test
+ void installIfAbsent_isIdempotent() {
+ CngProvider.installIfAbsent();
+ CngProvider.installIfAbsent(); // second call must be a no-op
+
+ long cngCount = Arrays.stream(Security.getProviders())
+ .filter(p -> "CNG".equals(p.getName()))
+ .count();
+ assertEquals(1, cngCount,
+ "CNG provider must appear exactly once even after multiple installIfAbsent() calls");
+ }
+
+ @Test
+ void installIfAbsent_insertsAtHighestPriority() {
+ CngProvider.installIfAbsent();
+ Provider[] providers = Security.getProviders();
+ // Security position 1 = index 0 in the array
+ assertEquals("CNG", providers[0].getName(),
+ "CNG must be at Security position 1 (highest priority) so JSSE uses it first");
+ }
+
+ // ─── Service registrations ────────────────────────────────────────────────
+
+ @Test
+ void provider_registersSha256WithRsa() {
+ Provider p = new CngProvider();
+ assertNotNull(p.getService("Signature", "SHA256withRSA"),
+ "CNG provider must advertise SHA256withRSA (used by TLS 1.2 client cert verify)");
+ }
+
+ @Test
+ void provider_registersSha384WithRsa() {
+ Provider p = new CngProvider();
+ assertNotNull(p.getService("Signature", "SHA384withRSA"));
+ }
+
+ @Test
+ void provider_registersSha512WithRsa() {
+ Provider p = new CngProvider();
+ assertNotNull(p.getService("Signature", "SHA512withRSA"));
+ }
+
+ @Test
+ void provider_registersRsaSsaPss() {
+ Provider p = new CngProvider();
+ assertNotNull(p.getService("Signature", "RSASSA-PSS"),
+ "CNG provider must advertise RSASSA-PSS");
+ }
+
+ @Test
+ void provider_name_isCng() {
+ assertEquals("CNG", new CngProvider().getName());
+ }
+
+ @Test
+ void provider_sha256Alias_resolves() {
+ Provider p = new CngProvider();
+ // Alias "SHA-256withRSA" must resolve to "SHA256withRSA"
+ assertNotNull(p.getService("Signature", "SHA-256withRSA"),
+ "Alias SHA-256withRSA must resolve via the CNG provider");
+ }
+
+ @Test
+ void installedProvider_bypassesOrdinaryRsaKeys() throws Exception {
+ CngProvider.installIfAbsent();
+ KeyPairGenerator generator = KeyPairGenerator.getInstance("RSA");
+ generator.initialize(2048);
+ KeyPair keyPair = generator.generateKeyPair();
+
+ Signature signature = Signature.getInstance("SHA256withRSA");
+ signature.initSign(keyPair.getPrivate());
+
+ assertNotEquals("CNG", signature.getProvider().getName());
+ }
+}
diff --git a/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/CngRsaPrivateKeyTest.java b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/CngRsaPrivateKeyTest.java
new file mode 100644
index 00000000..e66490d0
--- /dev/null
+++ b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/CngRsaPrivateKeyTest.java
@@ -0,0 +1,45 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import com.sun.jna.Pointer;
+import org.junit.jupiter.api.Test;
+
+import java.math.BigInteger;
+import java.util.concurrent.atomic.AtomicInteger;
+
+import static org.junit.jupiter.api.Assertions.*;
+
+class CngRsaPrivateKeyTest {
+
+ @Test
+ void keyMaterialIsNotExportable() {
+ CngRsaPrivateKey key = new CngRsaPrivateKey(
+ Pointer.createConstant(42),
+ BigInteger.valueOf(17),
+ 65537,
+ () -> { });
+
+ assertNull(key.getEncoded());
+ assertNull(key.getFormat());
+ assertEquals("RSA", key.getAlgorithm());
+ assertFalse(key.toString().contains("42"));
+ }
+
+ @Test
+ void closeReleasesNativeHandleExactlyOnce() {
+ AtomicInteger releases = new AtomicInteger();
+ CngRsaPrivateKey key = new CngRsaPrivateKey(
+ Pointer.createConstant(42),
+ BigInteger.valueOf(17),
+ 65537,
+ releases::incrementAndGet);
+
+ key.close();
+ key.close();
+
+ assertEquals(1, releases.get());
+ assertTrue(key.isClosed());
+ }
+}
diff --git a/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/CngSignatureParametersTest.java b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/CngSignatureParametersTest.java
new file mode 100644
index 00000000..78e3a1e5
--- /dev/null
+++ b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/CngSignatureParametersTest.java
@@ -0,0 +1,54 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import org.junit.jupiter.api.Test;
+
+import java.security.InvalidAlgorithmParameterException;
+import java.security.spec.MGF1ParameterSpec;
+import java.security.spec.PSSParameterSpec;
+
+import static org.junit.jupiter.api.Assertions.*;
+
+class CngSignatureParametersTest {
+
+ @Test
+ void pssAcceptsSupportedTlsParameters() throws Exception {
+ CngSignatureSpi spi = new CngSignatureSpi.RsaSsaPss();
+
+ spi.engineSetParameter(new PSSParameterSpec(
+ "SHA-256", "MGF1", MGF1ParameterSpec.SHA256, 32, 1));
+ spi.engineSetParameter(new PSSParameterSpec(
+ "SHA-384", "MGF1", MGF1ParameterSpec.SHA384, 48, 1));
+ spi.engineSetParameter(new PSSParameterSpec(
+ "SHA-512", "MGF1", MGF1ParameterSpec.SHA512, 64, 1));
+ }
+
+ @Test
+ void pssRejectsMismatchedMgfDigest() {
+ CngSignatureSpi spi = new CngSignatureSpi.RsaSsaPss();
+
+ assertThrows(InvalidAlgorithmParameterException.class,
+ () -> spi.engineSetParameter(new PSSParameterSpec(
+ "SHA-256", "MGF1", MGF1ParameterSpec.SHA384, 32, 1)));
+ }
+
+ @Test
+ void pssRejectsUnexpectedSaltLength() {
+ CngSignatureSpi spi = new CngSignatureSpi.RsaSsaPss();
+
+ assertThrows(InvalidAlgorithmParameterException.class,
+ () -> spi.engineSetParameter(new PSSParameterSpec(
+ "SHA-256", "MGF1", MGF1ParameterSpec.SHA256, 20, 1)));
+ }
+
+ @Test
+ void pssRejectsUnsupportedDigestInsteadOfFallingBack() {
+ CngSignatureSpi spi = new CngSignatureSpi.RsaSsaPss();
+
+ assertThrows(InvalidAlgorithmParameterException.class,
+ () -> spi.engineSetParameter(new PSSParameterSpec(
+ "SHA-1", "MGF1", MGF1ParameterSpec.SHA1, 20, 1)));
+ }
+}
diff --git a/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/CngX509ExtendedKeyManagerTest.java b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/CngX509ExtendedKeyManagerTest.java
new file mode 100644
index 00000000..6bf99a82
--- /dev/null
+++ b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/CngX509ExtendedKeyManagerTest.java
@@ -0,0 +1,41 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import com.sun.jna.Pointer;
+import org.junit.jupiter.api.Test;
+
+import javax.net.ssl.SSLContext;
+import javax.net.ssl.SSLEngine;
+import java.math.BigInteger;
+import java.security.cert.X509Certificate;
+
+import static org.junit.jupiter.api.Assertions.*;
+import static org.mockito.Mockito.mock;
+
+class CngX509ExtendedKeyManagerTest {
+
+ @Test
+ void selectsOnlyRsaClientAliasForSocketAndEngine() throws Exception {
+ CngRsaPrivateKey key = new CngRsaPrivateKey(
+ Pointer.createConstant(42),
+ BigInteger.valueOf(17),
+ 65537,
+ () -> { });
+ X509Certificate certificate = mock(X509Certificate.class);
+ CngX509ExtendedKeyManager manager =
+ new CngX509ExtendedKeyManager(key, certificate);
+ SSLEngine engine = SSLContext.getDefault().createSSLEngine();
+ String alias = manager.chooseClientAlias(new String[]{"RSA"}, null, null);
+
+ assertNotNull(alias);
+ assertEquals(alias,
+ manager.chooseEngineClientAlias(new String[]{"EC", "RSA"}, null, engine));
+ assertNull(manager.chooseClientAlias(new String[]{"EC"}, null, null));
+ assertNull(manager.chooseEngineClientAlias(new String[]{"EC"}, null, engine));
+ assertSame(key, manager.getPrivateKey(alias));
+ assertArrayEquals(new X509Certificate[]{certificate},
+ manager.getCertificateChain(alias));
+ }
+}
diff --git a/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/ImdsV2ClientTest.java b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/ImdsV2ClientTest.java
new file mode 100644
index 00000000..5c21ef8c
--- /dev/null
+++ b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/ImdsV2ClientTest.java
@@ -0,0 +1,127 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import com.microsoft.aad.msal4j.ManagedIdentityMtlsHttpResponse;
+import com.microsoft.aad.msal4j.ManagedIdentityMtlsRequest;
+import org.junit.jupiter.api.Test;
+
+import java.util.Collections;
+import java.util.HashMap;
+import java.util.List;
+import java.util.Map;
+import java.util.concurrent.atomic.AtomicReference;
+
+import static org.junit.jupiter.api.Assertions.*;
+
+class ImdsV2ClientTest {
+
+ @Test
+ void metadataUsesCurrentV2ContractAndIdentitySelector() {
+ AtomicReference url = new AtomicReference<>();
+ ManagedIdentityMtlsRequest request = request("client_id", "client id", http -> {
+ url.set(http.url());
+ return response("{\"clientId\":\"client id\",\"tenantId\":\"tenant\","
+ + "\"cuId\":{\"vmId\":\"vm\",\"vmssId\":\"set\"},"
+ + "\"attestationEndpoint\":\"https://maa.example\"}");
+ });
+
+ ImdsV2Client.PlatformMetadata metadata =
+ ImdsV2Client.getPlatformMetadata(request);
+
+ assertTrue(url.get().contains("/metadata/identity/getplatformmetadata"));
+ assertTrue(url.get().contains("cred-api-version=2.0"));
+ assertTrue(url.get().contains("client_id=client+id"));
+ assertEquals("client id", metadata.clientId);
+ assertEquals("vm", metadata.cuId());
+ }
+
+ @Test
+ void issueCredentialRequiresAttestationAndCurrentFields() {
+ ManagedIdentityMtlsRequest request = request(null, null, http -> {
+ assertEquals("POST", http.method());
+ assertTrue(http.body().contains("\"attestation_token\":\"jwt\""));
+ return response("{\"certificate\":\"cert\","
+ + "\"mtls_authentication_endpoint\":\"https://login.example/token\","
+ + "\"client_id\":\"client\",\"tenant_id\":\"tenant\","
+ + "\"identity_type\":\"SystemAssigned\"}");
+ });
+
+ assertThrows(MtlsMsiException.class,
+ () -> ImdsV2Client.issueCredential(request, "csr", ""));
+ ImdsV2Client.CredentialResponse credential =
+ ImdsV2Client.issueCredential(request, "csr", "jwt");
+
+ assertEquals("client", credential.clientId);
+ assertEquals("https://login.example/token",
+ credential.mtlsAuthenticationEndpoint);
+ }
+
+ @Test
+ void issueCredentialOmitsAttestationWhenNotRequested() {
+ ManagedIdentityMtlsRequest request = request(
+ null,
+ null,
+ http -> {
+ assertFalse(http.body().contains("attestation_token"));
+ return response("{\"certificate\":\"cert\","
+ + "\"mtls_authentication_endpoint\":\"https://login.example/token\","
+ + "\"client_id\":\"client\",\"tenant_id\":\"tenant\","
+ + "\"identity_type\":\"SystemAssigned\"}");
+ },
+ false);
+
+ assertDoesNotThrow(
+ () -> ImdsV2Client.issueCredential(request, "csr", null));
+ }
+
+ @Test
+ void incompleteOrFailedImdsResponseFailsClosed() {
+ assertThrows(MtlsMsiException.class,
+ () -> ImdsV2Client.getPlatformMetadata(
+ request(null, null, http -> response("{}"))));
+ assertThrows(MtlsMsiException.class,
+ () -> ImdsV2Client.getPlatformMetadata(
+ request(null, null, http ->
+ new ManagedIdentityMtlsHttpResponse(
+ 500, "failure", Collections.emptyMap()))));
+ }
+
+ @Test
+ void metadataRejectsResponsesWithoutImdsServerMarker() {
+ assertThrows(MtlsMsiException.class,
+ () -> ImdsV2Client.getPlatformMetadata(
+ request(null, null, http ->
+ new ManagedIdentityMtlsHttpResponse(
+ 200,
+ "{\"clientId\":\"client\",\"tenantId\":\"tenant\","
+ + "\"cuId\":{\"vmId\":\"vm\"},"
+ + "\"attestationEndpoint\":\"https://maa.example\"}",
+ Collections.emptyMap()))));
+ }
+
+ private static ManagedIdentityMtlsRequest request(
+ String selector,
+ String value,
+ com.microsoft.aad.msal4j.IManagedIdentityMtlsHttpClient client) {
+ return request(selector, value, client, true);
+ }
+
+ private static ManagedIdentityMtlsRequest request(
+ String selector,
+ String value,
+ com.microsoft.aad.msal4j.IManagedIdentityMtlsHttpClient client,
+ boolean attestationEnabled) {
+ return new ManagedIdentityMtlsRequest(
+ selector, value, "binding", "correlation", client,
+ attestationEnabled);
+ }
+
+ private static ManagedIdentityMtlsHttpResponse response(String body) {
+ Map> headers = new HashMap<>();
+ headers.put("Server", Collections.singletonList("IMDS/150.0"));
+ return new ManagedIdentityMtlsHttpResponse(
+ 200, body, headers);
+ }
+}
diff --git a/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/KeyGuardManagedIdentityMtlsProviderTest.java b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/KeyGuardManagedIdentityMtlsProviderTest.java
new file mode 100644
index 00000000..a66d4c53
--- /dev/null
+++ b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/KeyGuardManagedIdentityMtlsProviderTest.java
@@ -0,0 +1,41 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import org.junit.jupiter.api.Test;
+
+import java.time.Duration;
+import java.time.Instant;
+
+import static org.junit.jupiter.api.Assertions.*;
+
+class KeyGuardManagedIdentityMtlsProviderTest {
+
+ @Test
+ void certificateRotatesInsideTwentyFourHourWindow() {
+ long now = Instant.now().toEpochMilli();
+
+ assertTrue(KeyGuardManagedIdentityMtlsProvider.isCertificateCurrent(
+ now + Duration.ofHours(25).toMillis(), now));
+ assertFalse(KeyGuardManagedIdentityMtlsProvider.isCertificateCurrent(
+ now + Duration.ofHours(24).toMillis(), now));
+ assertFalse(KeyGuardManagedIdentityMtlsProvider.isCertificateCurrent(
+ now + Duration.ofHours(1).toMillis(), now));
+ }
+
+ @Test
+ void retiredGenerationPreventsUnlockedCacheHit() {
+ long now = Instant.now().toEpochMilli();
+ KeyGuardManagedIdentityMtlsProvider.BindingGeneration current =
+ new KeyGuardManagedIdentityMtlsProvider.BindingGeneration(
+ null,
+ null,
+ now + Duration.ofHours(25).toMillis());
+
+ assertTrue(KeyGuardManagedIdentityMtlsProvider
+ .canReturnWithoutCleanup(current, false));
+ assertFalse(KeyGuardManagedIdentityMtlsProvider
+ .canReturnWithoutCleanup(current, true));
+ }
+}
diff --git a/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/KeyGuardMtlsBindingContextTest.java b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/KeyGuardMtlsBindingContextTest.java
new file mode 100644
index 00000000..7050914b
--- /dev/null
+++ b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/KeyGuardMtlsBindingContextTest.java
@@ -0,0 +1,59 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import com.sun.jna.Pointer;
+import org.junit.jupiter.api.Test;
+
+import java.math.BigInteger;
+import java.security.MessageDigest;
+import java.security.cert.X509Certificate;
+import java.util.Base64;
+
+import static org.junit.jupiter.api.Assertions.*;
+import static org.mockito.Mockito.*;
+
+class KeyGuardMtlsBindingContextTest {
+
+ @Test
+ void keyIdUsesFullLeafCertificateDer() throws Exception {
+ byte[] certificateDer = new byte[]{1, 2, 3, 4, 5};
+ X509Certificate certificate = mock(X509Certificate.class);
+ when(certificate.getEncoded()).thenReturn(certificateDer);
+
+ KeyGuardMtlsBindingContext context =
+ new KeyGuardMtlsBindingContext(key(), certificate);
+ String expected = Base64.getUrlEncoder().withoutPadding().encodeToString(
+ MessageDigest.getInstance("SHA-256").digest(certificateDer));
+
+ assertEquals(expected, context.keyId());
+ assertNotNull(context.sslContext());
+ assertEquals("TLSv1.2", context.sslContext().getProtocol());
+ assertArrayEquals(
+ new String[]{"TLSv1.2"},
+ context.sslContext().getDefaultSSLParameters().getProtocols());
+ assertSame(certificate, context.bindingCertificate());
+ }
+
+ @Test
+ void renewedCertificateWithSameKeyChangesBindingKeyId() throws Exception {
+ X509Certificate first = mock(X509Certificate.class);
+ X509Certificate second = mock(X509Certificate.class);
+ when(first.getEncoded()).thenReturn(new byte[]{1});
+ when(second.getEncoded()).thenReturn(new byte[]{2});
+ CngRsaPrivateKey key = key();
+
+ assertNotEquals(
+ new KeyGuardMtlsBindingContext(key, first).keyId(),
+ new KeyGuardMtlsBindingContext(key, second).keyId());
+ }
+
+ private static CngRsaPrivateKey key() {
+ return new CngRsaPrivateKey(
+ Pointer.createConstant(42),
+ BigInteger.valueOf(17),
+ 65537,
+ () -> { });
+ }
+}
diff --git a/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/Pkcs10BuilderTest.java b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/Pkcs10BuilderTest.java
new file mode 100644
index 00000000..0af24613
--- /dev/null
+++ b/msal4j-mtls-extensions/src/test/java/com/microsoft/aad/msal4j/mtls/Pkcs10BuilderTest.java
@@ -0,0 +1,219 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j.mtls;
+
+import com.sun.jna.Pointer;
+import org.junit.jupiter.api.Test;
+import org.junit.jupiter.api.condition.EnabledOnOs;
+import org.junit.jupiter.api.condition.OS;
+import org.mockito.MockedStatic;
+import org.mockito.Mockito;
+
+import java.math.BigInteger;
+import java.util.Base64;
+
+import static org.junit.jupiter.api.Assertions.*;
+import static org.mockito.ArgumentMatchers.*;
+import static org.mockito.ArgumentMatchers.nullable;
+
+/**
+ * Unit tests for {@link Pkcs10Builder} DER encoding.
+ *
+ * Tests are split into two groups:
+ *
+ * - Pure DER primitives — no CNG required; run on all platforms.
+ * - Full CSR generation — requires Windows (loads NCryptLibrary via
+ * CngKeyGuard); {@link CngKeyGuard#signPss} is mocked so no real CNG key is needed.
+ *
+ *
+ * The CSR format must match msal-go's {@code generateCSR()} and MSAL.NET's
+ * {@code Csr.Generate()} exactly so that the Azure IMDS {@code /issuecredential}
+ * endpoint can parse and validate it.
+ */
+class Pkcs10BuilderTest {
+
+ // ─── DER primitives (pure Java, cross-platform) ───────────────────────────
+
+ @Test
+ void derSequence_short_wrapsWithTag30() {
+ byte[] content = {0x01, 0x02, 0x03};
+ byte[] seq = Pkcs10Builder.derSequence(content);
+
+ assertEquals(0x30, seq[0] & 0xFF, "SEQUENCE tag must be 0x30");
+ assertEquals(3, seq[1] & 0xFF, "Short-form length must equal content length");
+ assertEquals(0x01, seq[2]);
+ assertEquals(0x02, seq[3]);
+ assertEquals(0x03, seq[4]);
+ assertEquals(5, seq.length);
+ }
+
+ @Test
+ void derSequence_shortFormMaxLength() {
+ // 127 bytes is the maximum for single-byte short-form length
+ byte[] content = new byte[127];
+ byte[] seq = Pkcs10Builder.derSequence(content);
+
+ assertEquals(0x30, seq[0] & 0xFF);
+ assertEquals(127, seq[1] & 0xFF);
+ assertEquals(2 + 127, seq.length);
+ }
+
+ @Test
+ void derSequence_longForm1Byte_length128() {
+ // 128 bytes requires 0x81 long-form header
+ byte[] content = new byte[128];
+ byte[] seq = Pkcs10Builder.derSequence(content);
+
+ assertEquals(0x30, seq[0] & 0xFF);
+ assertEquals(0x81, seq[1] & 0xFF, "Long-form header byte for lengths 128-255 must be 0x81");
+ assertEquals(128, seq[2] & 0xFF);
+ assertEquals(3 + 128, seq.length);
+ }
+
+ @Test
+ void derSequence_longForm2Byte_length256() {
+ // 256 bytes requires 0x82 two-byte length
+ byte[] content = new byte[256];
+ byte[] seq = Pkcs10Builder.derSequence(content);
+
+ assertEquals(0x30, seq[0] & 0xFF);
+ assertEquals(0x82, seq[1] & 0xFF, "Long-form header for lengths 256+ must be 0x82");
+ assertEquals(1, seq[2] & 0xFF, "High byte of length 256 (0x0100)");
+ assertEquals(0, seq[3] & 0xFF, "Low byte of length 256");
+ assertEquals(4 + 256, seq.length);
+ }
+
+ @Test
+ void derBitString_prependsZeroUnusedBitsByte() {
+ byte[] data = {0x01, 0x02};
+ byte[] bs = Pkcs10Builder.derBitString(data);
+
+ assertEquals(0x03, bs[0] & 0xFF, "BIT STRING tag must be 0x03");
+ assertEquals(3, bs[1] & 0xFF, "Length must cover the unused-bits byte + data");
+ assertEquals(0x00, bs[2], "Unused bits must be 0x00 (byte-aligned content)");
+ assertEquals(0x01, bs[3]);
+ assertEquals(0x02, bs[4]);
+ }
+
+ @Test
+ void derBitString_empty_hasOnlyUnusedBitsByte() {
+ byte[] bs = Pkcs10Builder.derBitString(new byte[0]);
+ assertEquals(0x03, bs[0] & 0xFF);
+ assertEquals(1, bs[1] & 0xFF);
+ assertEquals(0x00, bs[2]);
+ }
+
+ // ─── Full CSR generation (Windows only — CngKeyGuard is mocked) ───────────
+
+ @Test
+ @EnabledOnOs(OS.WINDOWS)
+ void generate_outputIsBase64EncodedDerSequence() throws Exception {
+ BigInteger modulus = BigInteger.valueOf(2).pow(2047).add(BigInteger.ONE);
+ byte[] fakeSignature = new byte[256]; // 2048-bit RSA output size
+
+ try (MockedStatic mockCng = Mockito.mockStatic(CngKeyGuard.class)) {
+ mockCng.when(() -> CngKeyGuard.signPss(any(), any(), anyString(), anyInt()))
+ .thenReturn(fakeSignature);
+
+ String b64 = Pkcs10Builder.generate(
+ Pointer.NULL, modulus, 65537,
+ "test-client-id", "test-tenant-id", "vm-id-1", null);
+
+ assertNotNull(b64);
+ byte[] der = Base64.getDecoder().decode(b64);
+ assertEquals(0x30, der[0] & 0xFF,
+ "Outermost CSR element must be a DER SEQUENCE (0x30)");
+ }
+ }
+
+ @Test
+ @EnabledOnOs(OS.WINDOWS)
+ void generate_cuIdAttributeContainsVmIdAndVmssId() throws Exception {
+ BigInteger modulus = BigInteger.valueOf(2).pow(2047).add(BigInteger.ONE);
+ byte[] fakeSignature = new byte[256];
+
+ try (MockedStatic mockCng = Mockito.mockStatic(CngKeyGuard.class)) {
+ mockCng.when(() -> CngKeyGuard.signPss(any(), any(), anyString(), anyInt()))
+ .thenReturn(fakeSignature);
+
+ String b64 = Pkcs10Builder.generate(
+ Pointer.NULL, modulus, 65537,
+ "client-a", "tenant-b", "my-vm-id", "my-vmss-id");
+
+ byte[] der = Base64.getDecoder().decode(b64);
+ String derText = new String(der, java.nio.charset.StandardCharsets.UTF_8);
+ assertTrue(derText.contains("my-vm-id"),
+ "CSR DER must embed vmId in the cuId attribute");
+ assertTrue(derText.contains("my-vmss-id"),
+ "CSR DER must embed vmssId in the cuId attribute");
+ }
+ }
+
+ @Test
+ @EnabledOnOs(OS.WINDOWS)
+ void generate_cuIdAttributeEmptyObject_whenBothIdsNull() throws Exception {
+ BigInteger modulus = BigInteger.valueOf(2).pow(2047).add(BigInteger.ONE);
+ byte[] fakeSignature = new byte[256];
+
+ try (MockedStatic mockCng = Mockito.mockStatic(CngKeyGuard.class)) {
+ mockCng.when(() -> CngKeyGuard.signPss(any(), any(), anyString(), anyInt()))
+ .thenReturn(fakeSignature);
+
+ String b64 = Pkcs10Builder.generate(
+ Pointer.NULL, modulus, 65537,
+ "client-a", "tenant-b", null, null);
+
+ byte[] der = Base64.getDecoder().decode(b64);
+ String derText = new String(der, java.nio.charset.StandardCharsets.UTF_8);
+ assertTrue(derText.contains("{}"),
+ "cuId JSON must be '{}' when both vmId and vmssId are null (omitempty)");
+ }
+ }
+
+ @Test
+ @EnabledOnOs(OS.WINDOWS)
+ void generate_subjectContainsClientIdAndTenantId() throws Exception {
+ BigInteger modulus = BigInteger.valueOf(2).pow(2047).add(BigInteger.ONE);
+ byte[] fakeSignature = new byte[256];
+
+ try (MockedStatic mockCng = Mockito.mockStatic(CngKeyGuard.class)) {
+ mockCng.when(() -> CngKeyGuard.signPss(any(), any(), anyString(), anyInt()))
+ .thenReturn(fakeSignature);
+
+ String b64 = Pkcs10Builder.generate(
+ Pointer.NULL, modulus, 65537,
+ "subject-client-id", "subject-tenant-id", null, null);
+
+ byte[] der = Base64.getDecoder().decode(b64);
+ String derText = new String(der, java.nio.charset.StandardCharsets.UTF_8);
+ assertTrue(derText.contains("subject-client-id"),
+ "CSR subject (CN) must contain the clientId");
+ assertTrue(derText.contains("subject-tenant-id"),
+ "CSR subject (DC) must contain the tenantId");
+ }
+ }
+
+ @Test
+ @EnabledOnOs(OS.WINDOWS)
+ void generate_signPssCalledWithSha256AndSalt32() throws Exception {
+ BigInteger modulus = BigInteger.valueOf(2).pow(2047).add(BigInteger.ONE);
+ byte[] fakeSignature = new byte[256];
+
+ try (MockedStatic mockCng = Mockito.mockStatic(CngKeyGuard.class)) {
+ mockCng.when(() -> CngKeyGuard.signPss(any(), any(), anyString(), anyInt()))
+ .thenReturn(fakeSignature);
+
+ Pkcs10Builder.generate(
+ Pointer.NULL, modulus, 65537,
+ "c", "t", null, null);
+
+ // Verify the exact signature algorithm parameters (must match msal-go and MSAL.NET)
+ mockCng.verify(() -> CngKeyGuard.signPss(
+ nullable(Pointer.class),
+ any(byte[].class),
+ eq("SHA256"),
+ eq(32)));
+ }
+ }
+}
diff --git a/msal4j-sdk/docs/managed-identity-v2-mtls-pop-review-guide.md b/msal4j-sdk/docs/managed-identity-v2-mtls-pop-review-guide.md
new file mode 100644
index 00000000..f2248803
--- /dev/null
+++ b/msal4j-sdk/docs/managed-identity-v2-mtls-pop-review-guide.md
@@ -0,0 +1,785 @@
+# Reviewer guide: Managed Identity v2 KeyGuard mTLS PoP
+
+This guide is the recommended entry point for reviewing the Managed Identity v2
+mTLS Proof-of-Possession change. It explains the intended architecture, security
+boundaries, protocol flow, review order, test coverage, and manual validation.
+
+The implementation is intentionally split between the portable MSAL core and an
+optional Windows extension. Reviewers should verify that this separation remains
+intact: native code performs only platform cryptographic operations, while Java
+continues to own OAuth, HTTP, caching, certificate parsing, and TLS.
+
+## Review goals
+
+The change is successful only if all of the following remain true:
+
+- The KeyGuard private key is never exported into Java memory.
+- Java JSSE performs both token-endpoint and downstream TLS.
+- Native interop is limited to Windows CNG/KeyGuard and attestation operations.
+- Attestation is optional, but fails closed whenever it is requested.
+- Tokens are cached only with the exact certificate binding that produced them.
+- A normal bearer token cannot satisfy an mTLS PoP request.
+- A token bound to certificate A cannot be used with certificate B.
+- Application developers receive a reusable standard Java `SSLContext`.
+- The standard MSAL OAuth pipeline remains responsible for token requests.
+- Custom HTTP clients cannot silently discard the mTLS configuration.
+- Credential-bound HTTP requests cannot follow redirects.
+- The optional native extension does not affect applications that do not use it.
+- Java 8 source and bytecode compatibility are preserved.
+
+## Recommended review order
+
+Reviewing the files in this order minimizes context switching:
+
+1. Public API and result surface.
+2. Managed Identity request orchestration.
+3. OAuth and HTTP integration.
+4. Token-cache partitioning.
+5. Optional provider loading.
+6. KeyGuard and signing bridge.
+7. Certificate and binding lifecycle.
+8. Attestation and IMDS v2.
+9. Native packaging.
+10. Unit tests and the manual E2E.
+
+## Architecture summary
+
+```mermaid
+flowchart TD
+ App[Application] --> MIA[ManagedIdentityApplication]
+ MIA --> Core[MSAL core]
+ Core --> SPI[IManagedIdentityMtlsProvider]
+ SPI --> Ext[Optional Windows extension]
+ Ext --> IMDS[IMDS v2]
+ Ext --> KG[Windows CNG / KeyGuard]
+ Ext --> ATT[AttestationClientLib]
+ KG --> Key[Non-exportable RSA key]
+ ATT --> MAA[MAA attestation JWT]
+ IMDS --> Cert[Binding certificate]
+ Ext --> Context[IMtlsBindingContext]
+ Context --> JSSE[Java SSLContext / JSSE]
+ Core --> ESTS[Normal MSAL OAuth pipeline]
+ JSSE --> ESTS
+ ESTS --> Result[mtls_pop token + binding context]
+ Result --> Resource[Independent Java HTTPS resource call]
+```
+
+### Ownership boundary
+
+| Component | Owns | Must not own |
+| --- | --- | --- |
+| MSAL core | OAuth, claims, capabilities, retries, telemetry, response parsing, token cache | KeyGuard handles, CSR construction, attestation implementation |
+| Windows extension | CNG key operations, attestation, CSR, binding certificate lifecycle | Bespoke OAuth token client, persistent token cache |
+| JCA/JSSE bridge | TLS signatures through `PrivateKey` and `SignatureSpi` | Private-key export |
+| Application | Independent downstream HTTP call through returned `SSLContext` | Native key-handle management |
+
+## End-to-end protocol flow
+
+```mermaid
+sequenceDiagram
+ participant App as Application
+ participant Core as MSAL core
+ participant Ext as KeyGuard extension
+ participant IMDS as IMDS v2
+ participant MAA as Attestation library / MAA
+ participant ESTS as ESTS token endpoint
+ participant KV as Token-bound resource
+
+ App->>Core: acquireTokenForManagedIdentity(parameters)
+ Core->>Ext: acquireBinding(request, HTTP callback)
+ Ext->>IMDS: GET getPlatformMetadata
+ IMDS-->>Ext: identity, CUID, regional token URL, attestation endpoint
+ Ext->>Ext: open or create KeyGuard key
+ Ext->>Ext: private signing liveness probe
+ opt Attestation requested
+ Ext->>MAA: attest KeyGuard handle
+ MAA-->>Ext: attestation JWT
+ end
+ Ext->>Ext: build RSA-PSS PKCS#10 CSR
+ Ext->>IMDS: POST issuecredential
+ IMDS-->>Ext: binding certificate
+ Ext->>Ext: verify certificate public key
+ Ext-->>Core: certificate + binding SSLContext + key ID
+ Core->>Core: binding-aware token-cache lookup
+ alt Cache miss
+ Core->>ESTS: normal OAuth request over binding SSLContext
+ ESTS-->>Core: token_type=mtls_pop + access token
+ Core->>Core: validate token type
+ Core->>Core: binding-aware cache write
+ end
+ Core-->>App: token + IMtlsBindingContext
+ App->>KV: Java HTTPS with token and returned SSLContext
+ KV-->>App: protected resource response
+```
+
+## Public API review
+
+### Parameter combinations
+
+The intended combinations are:
+
+```java
+ManagedIdentityParameters.builder(resource)
+ .withMtlsProofOfPossession()
+ .build();
+```
+
+```java
+ManagedIdentityParameters.builder(resource)
+ .withMtlsProofOfPossession()
+ .withAttestationSupport()
+ .build();
+```
+
+The following must be rejected:
+
+```java
+ManagedIdentityParameters.builder(resource)
+ .withAttestationSupport()
+ .build();
+```
+
+Attestation is a strengthening option for an mTLS binding. It is not a
+standalone token-acquisition mode.
+
+### Result surface
+
+An mTLS result exposes:
+
+- `tokenType()`, which must be exactly `mtls_pop`;
+- `bindingCertificate()`, the public leaf certificate;
+- `mtlsBindingContext()`, a process-local reusable binding context;
+- `mtlsBindingContext().sslContext()`, the standard JSSE context;
+- `mtlsBindingContext().keyId()`, the full-certificate binding identity.
+
+The binding context is intentionally not serializable. Native handles and
+`SSLContext` instances must be reconstructed by each process.
+
+## Core request orchestration
+
+Primary file:
+
+`msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/AcquireTokenByManagedIdentitySupplier.java`
+
+Review these invariants:
+
+- The extension is invoked only for mTLS PoP requests.
+- Provider exceptions are normalized to MSAL exceptions.
+- Existing MSAL exceptions retain their original error codes.
+- IMDS calls use request-scoped IMDS retry behavior.
+- The external ESTS request does not inherit IMDS retry behavior.
+- The resolved binding cache key is stored on request-scoped state.
+- Public `ManagedIdentityParameters` are not mutated during acquisition.
+- Cache lookup and cache write use the same resolved extended cache hash.
+- The final result preserves token metadata when the binding context is attached.
+
+## OAuth pipeline integration
+
+Primary files:
+
+- `TokenRequestExecutor.java`
+- `OAuthHttpRequest.java`
+- `HttpRequest.java`
+- `DefaultHttpClient.java`
+- `TokenResponse.java`
+
+The ESTS exchange must remain a specialization of the normal OAuth pipeline,
+not a second token client.
+
+Review that the request-specific endpoint and socket factory do not bypass:
+
+- claims merging;
+- client capabilities;
+- telemetry headers;
+- correlation IDs;
+- retry and error parsing;
+- token response deserialization;
+- refresh metadata;
+- cache writes.
+
+The token response must explicitly contain `token_type=mtls_pop`. Missing or
+different token types fail closed.
+
+## HTTP security review
+
+### Redirects
+
+Any request carrying a request-specific client-certificate socket factory must
+disable automatic redirects. A redirect could disclose proof-of-possession to
+an unintended destination.
+
+### HTTPS-only endpoint
+
+The regional token endpoint returned by platform metadata must use HTTPS.
+Reject non-HTTPS endpoints before sending OAuth parameters or presenting the
+binding certificate.
+
+### Custom HTTP clients
+
+MSAL Java supports application-provided `IHttpClient` implementations. Existing
+implementations predate request-specific socket factories.
+
+The `IMtlsCapableHttpClient` marker is an explicit compatibility contract:
+
+- the client understands `HttpRequest.sslSocketFactory()`;
+- the client applies that factory to the exact request;
+- the client preserves no-redirect behavior;
+- a custom client lacking the capability fails fast.
+
+Silent fallback to a non-mTLS request is a security failure.
+
+## Binding-aware token cache
+
+An mTLS token is reusable only with the certificate to which it was issued.
+
+The extended cache identity includes:
+
+| Dimension | Value |
+| --- | --- |
+| Token type | `mtls_pop` |
+| Binding identity | Base64UrlNoPadding(SHA-256(full leaf certificate DER)) |
+| Attestation mode | Attested and unattested requests are isolated |
+| Existing MSAL dimensions | Authority, tenant, client, scopes/resource, claims, account, flow |
+
+Review these cases:
+
+- bearer cache entries cannot satisfy mTLS requests;
+- certificate A cannot satisfy certificate B;
+- same-key certificate renewal changes the cache partition;
+- attested and unattested bindings never cross-hit;
+- force refresh bypasses the access-token cache;
+- cache hits still return a live binding context.
+
+## KeyGuard private-key bridge
+
+```mermaid
+flowchart LR
+ JSSE[JSSE handshake] --> Signature[Signature API]
+ Signature --> SPI[CngSignatureSpi]
+ SPI --> JNA[JNA]
+ JNA --> NCrypt[NCryptSignHash]
+ NCrypt --> KG[VBS KeyGuard]
+ KG --> Signature
+```
+
+### `CngRsaPrivateKey`
+
+Review that:
+
+- `getEncoded()` returns `null`;
+- `getFormat()` returns `null`;
+- private exponent access is unavailable;
+- only public modulus and exponent are represented in Java;
+- native handle cleanup is idempotent;
+- accidental Java serialization cannot expose private key material.
+
+### `CngSignatureSpi`
+
+Review that:
+
+- supported hashes are explicit;
+- unsupported algorithms fail rather than defaulting;
+- PSS parameters are validated;
+- MGF must be MGF1;
+- digest and MGF digest must match;
+- salt lengths must be supported;
+- trailer field must be one;
+- non-KeyGuard keys delegate to another provider;
+- provider delegation cannot recurse into `CngProvider`;
+- signatures are produced only through `NCryptSignHash`.
+
+### `CngX509ExtendedKeyManager`
+
+Review socket and engine paths:
+
+- `chooseClientAlias`;
+- `chooseEngineClientAlias`;
+- certificate chain lookup;
+- private key lookup;
+- RSA key-type filtering.
+
+Both `SSLSocket` and `SSLEngine` consumers must be supported.
+
+## Key lifecycle
+
+### Per-boot stale keys
+
+KeyGuard KSP metadata can survive a reboot even when the VBS-protected private
+material is no longer usable.
+
+Opening the key and exporting its public key is not a sufficient liveness test.
+The extension performs a private signing probe after reopening an existing key.
+
+```mermaid
+flowchart TD
+ Open[Open persisted key] --> Export[Export public key]
+ Export --> Probe[Private signing probe]
+ Probe -->|Success| Use[Use existing key]
+ Probe -->|Failure| Delete[Delete stale key]
+ Delete --> Create[Create new KeyGuard key]
+ Create --> Attest[Create new attestation and certificate]
+```
+
+Review that stale-key recovery:
+
+- deletes the unusable key;
+- recreates it before CSR generation;
+- does not reuse attestation evidence for old key material;
+- closes failed native handles;
+- fails closed if recreation is unsuccessful.
+
+### Certificate rotation
+
+Review that:
+
+- certificates rotate before expiry;
+- the current generation remains available during safe handoff;
+- old generations are retained only while needed;
+- expired retired generations close native handles;
+- normal cache hits still perform retired-generation cleanup;
+- rotation changes the token-cache binding partition.
+
+## Attestation
+
+Attestation is requested only when `.withAttestationSupport()` is present.
+
+When selected:
+
+- missing DLL loading fails;
+- empty attestation output fails;
+- malformed JWTs fail;
+- expired evidence fails;
+- stale evidence inside the freshness buffer is not reused;
+- failures never downgrade to unattested issuance.
+
+### Attestation cache
+
+The cache identity is:
+
+```text
+normalized attestation endpoint + fingerprint of current public key material
+```
+
+Review:
+
+- normalized endpoint handling;
+- five-minute freshness buffer;
+- successful-result-only caching;
+- no caching of failures;
+- per-key single-flight synchronization;
+- no coalescing between distinct keys.
+
+## IMDS v2 review
+
+Primary file:
+
+`msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/ImdsV2Client.java`
+
+Protocol values:
+
+| Area | Value |
+| --- | --- |
+| API query | `cred-api-version=2.0` |
+| Metadata path | `/metadata/identity/getplatformmetadata` |
+| Credential path | `/metadata/identity/issuecredential` |
+| Identity types | `SystemAssigned`, `UserAssigned` |
+| CUID CSR OID | `1.3.6.1.4.1.311.90.2.10` |
+| CSR signature | RSA-PSS SHA-256 |
+
+Review that:
+
+- requests include the IMDS metadata header;
+- metadata response includes the expected IMDS server marker;
+- identity selection matches the requested managed identity;
+- required fields are validated before use;
+- optional attestation is omitted when not requested;
+- missing attestation fails only when attestation was requested;
+- certificate response is validated before constructing the binding context.
+
+## CSR review
+
+Primary file:
+
+`msal4j-mtls-extensions/src/main/java/com/microsoft/aad/msal4j/mtls/Pkcs10Builder.java`
+
+Review:
+
+- DER encoding, not text concatenation;
+- subject and public-key encoding;
+- CUID attribute OID;
+- UTF-8 JSON attribute value;
+- RSA-PSS SHA-256 signature;
+- correct MGF1 parameters;
+- salt length;
+- signature BIT STRING encoding;
+- CSR uses the current KeyGuard key.
+
+## Native DLL packaging
+
+The optional extension bundles the x64 native DLL from:
+
+```text
+Microsoft.Azure.Security.KeyGuardAttestation 1.1.5
+```
+
+This matches the version pinned by current MSAL.NET.
+
+The DLL is stored at:
+
+```text
+META-INF/native/win-x64/AttestationClientLib.dll
+```
+
+Expected SHA-256:
+
+```text
+90dfcce20e1a74519b49796eeee17e6e59a257c3acf754f454a49380d28a568b
+```
+
+Review:
+
+- resource exists in the production extension JAR;
+- resource also survives E2E shading;
+- runtime extraction uses a unique temporary directory;
+- extracted bytes are verified before loading;
+- architecture mismatch fails clearly;
+- manual `PATH` or `java.library.path` configuration is unnecessary;
+- package license and notice are included;
+- only the optional extension carries the native payload.
+
+## Test inventory
+
+### Core SDK
+
+| Test area | Purpose |
+| --- | --- |
+| Parameter tests | Valid and invalid PoP/attestation combinations |
+| Provider loader tests | Missing, unique, and ambiguous providers |
+| Binding tests | Endpoint and certificate-binding validation |
+| Result tests | Token type and transient binding context |
+| HTTP tests | Socket factory application and redirect prevention |
+| Token executor tests | Normal OAuth pipeline with request-specific TLS |
+| Managed Identity supplier tests | Exception normalization and IMDS retry behavior |
+| Cache tests | Binding-aware extended cache partitions |
+
+### Windows extension
+
+| Test area | Purpose |
+| --- | --- |
+| KeyGuard tests | Stale-key detection and recreation |
+| Private-key tests | Non-exportability and cleanup |
+| Provider tests | Signature registrations and delegation |
+| Signature tests | PKCS#1 and PSS parameter validation |
+| Key-manager tests | Socket and engine alias selection |
+| CSR tests | DER structure and PSS signing |
+| Attestation cache tests | Expiry, freshness, and single flight |
+| IMDS tests | Contracts, optional attestation, origin validation |
+| Binding-context tests | Full-DER key ID and configured SSL context |
+| Provider lifecycle tests | Rotation and retired-generation cleanup |
+| Native-loader tests | Version and packaged DLL hash |
+
+## Build commands
+
+From the repository root:
+
+```powershell
+mvn -pl msal4j-mtls-extensions -am test
+```
+
+Build the profile-only E2E:
+
+```powershell
+mvn -pl msal4j-mtls-extensions-e2e -am `
+ -Pe2e `
+ -DskipTests `
+ -Dmaven.javadoc.skip=true `
+ package
+```
+
+The production extension JAR must not contain the E2E main class.
+
+The shaded E2E artifact is:
+
+```text
+msal4j-mtls-extensions-e2e/target/*-e2e.jar
+```
+
+## Manual VM prerequisites
+
+Use a disposable Windows x64 Trusted Launch Azure VM or VMSS instance with:
+
+- Secure Boot;
+- vTPM;
+- VBS / Credential Guard;
+- Managed Identity;
+- Java 8 or later;
+- network access to IMDS, attestation, ESTS, and the test resource.
+
+For attestation, the TPM must report that it is capable of attestation.
+
+Do not use production secrets for the manual test.
+
+## Key Vault test configuration
+
+Use a dedicated test vault and secret.
+
+The VM managed identity needs only secret `get` permission.
+
+For an access-policy vault:
+
+```powershell
+az keyvault set-policy `
+ --resource-group `
+ --name `
+ --object-id `
+ --secret-permissions get
+```
+
+The test vault must be configured for token-bound authentication. Apply service
+configuration only to a dedicated vault because enforcement can reject ordinary
+clients.
+
+Example ARM body:
+
+```json
+{
+ "properties": {
+ "tokenBindingParameters": {
+ "mode": "Enforced",
+ "minimumTokenBindingStrength": "Unattested"
+ }
+ }
+}
+```
+
+Reviewers should use the currently supported Key Vault ARM API version available
+in their test environment.
+
+## Run the E2E
+
+Set:
+
+```powershell
+$env:MSAL_JAVA_MTLS_AKV_URL = "https://.vault.azure.net"
+$env:MSAL_JAVA_MTLS_AKV_SECRET_NAME = ""
+```
+
+Optional identity selection:
+
+```powershell
+$env:MSAL_JAVA_MTLS_IDENTITY_CLIENT_ID = ""
+```
+
+Optional force refresh:
+
+```powershell
+$env:MSAL_JAVA_MTLS_FORCE_REFRESH = "true"
+```
+
+Optional second identity for token-A/binding-B rejection:
+
+```powershell
+$env:MSAL_JAVA_MTLS_MISMATCH_IDENTITY_CLIENT_ID = ""
+```
+
+Run:
+
+```powershell
+.\run-java-msi-v2-mtls-devapp.ps1
+```
+
+No standalone attestation DLL or native library path is required.
+
+## Expected positive output
+
+The app must confirm:
+
+```text
+PASS: token_type = mtls_pop
+PASS: binding certificate returned
+PASS: reusable JSSE binding context returned
+PASS: cnf.x5t#S256 matches binding certificate
+PASS: HTTP 200
+PASS: AKV response validated
+PASS: TokenSource = CACHE
+PASS: matching binding context available
+RESULT: PASS
+```
+
+## Expected negative output: no certificate
+
+The app reuses the same valid `mtls_pop` token but creates a client connection
+without the binding key manager.
+
+Expected result:
+
+```text
+PASS: token without certificate rejected with HTTP 401 Unauthorized
+```
+
+This proves the resource is not accepting the token merely because it is
+otherwise valid.
+
+## Expected negative output: mismatched certificate
+
+When a distinct second UAMI is configured:
+
+```text
+token A + binding A -> HTTP 200
+token A + binding B -> rejected
+```
+
+The app verifies that the two key IDs differ before making the negative call.
+
+## Live validation completed
+
+The current implementation has been validated on a Windows Server 2025 Trusted
+Launch Azure VM:
+
+- the extension JAR loaded the bundled attestation DLL;
+- no standalone DLL was placed beside the app;
+- attestation completed;
+- ESTS returned `mtls_pop`;
+- token `cnf.x5t#S256` matched the full certificate DER hash;
+- the returned Java `SSLContext` performed the downstream request;
+- the test Key Vault returned HTTP 200 with the correct certificate;
+- the same valid token without the certificate returned HTTP 401;
+- the second acquisition returned `TokenSource.CACHE`.
+
+Environment-specific subscription, tenant, identity, vault, and secret
+identifiers are intentionally omitted.
+
+## Troubleshooting
+
+| Symptom | Likely cause | Check |
+| --- | --- | --- |
+| Extension provider not found | Extension JAR absent | Application dependencies and ServiceLoader resource |
+| Multiple providers found | Duplicate extension implementations | Classpath |
+| KeyGuard unavailable | VBS or Trusted Launch missing | Secure Boot, vTPM, VBS status |
+| Stale key after reboot | Per-boot private material lost | Signing liveness probe and recreation logs |
+| Attestation DLL load failure | Corrupt or wrong architecture resource | JAR resource, hash, Windows x64 |
+| Attestation failure | TPM not provisioned | TPM attestation capability |
+| IMDS metadata rejected | Missing IMDS response marker | Response headers and endpoint |
+| Credential issuance rejected | CSR, CUID, identity, or attestation mismatch | IMDS response body and correlation ID |
+| ESTS token type is not `mtls_pop` | Service not enrolled or request invalid | Token response and endpoint |
+| `cnf` mismatch | Wrong certificate or cache partition | Full-DER key ID |
+| Resource HTTP 401 without certificate | Expected negative result | Confirm positive call still returns 200 |
+| Resource rejects correct certificate | Resource enrollment or identity permission | Resource configuration and access policy |
+| Custom HTTP client failure | Client does not honor socket factory | `IMtlsCapableHttpClient` implementation |
+| Unexpected redirect | Credential endpoint redirected | Redirect policy and configured endpoint |
+| Second acquisition hits IDP | Cache identity changed or force refresh enabled | Token source and key ID |
+
+## Threat-model checklist
+
+### Private key
+
+- [ ] No Java API exposes private key bytes.
+- [ ] No export flags permit private-key export.
+- [ ] Every TLS signature reaches `NCryptSignHash`.
+- [ ] Handles are closed exactly once.
+- [ ] Stale handles are deleted and recreated.
+
+### Attestation
+
+- [ ] Optional unless explicitly requested.
+- [ ] Fail closed when requested.
+- [ ] Cache is key-bound and endpoint-bound.
+- [ ] Failures are not cached.
+- [ ] Expiry and freshness buffer are enforced.
+
+### Certificate
+
+- [ ] Issued certificate public key matches KeyGuard key.
+- [ ] Full DER determines binding identity.
+- [ ] Rotation creates a new cache partition.
+- [ ] Old native handles are retired and closed.
+
+### OAuth
+
+- [ ] Normal token pipeline is used.
+- [ ] Claims and capabilities are preserved.
+- [ ] Token endpoint is HTTPS.
+- [ ] Token type is explicitly validated.
+- [ ] Provider errors become MSAL errors.
+
+### HTTP
+
+- [ ] Request-specific socket factory is applied.
+- [ ] Redirects are disabled.
+- [ ] Custom clients fail fast without mTLS capability.
+- [ ] Downstream calls can use standard Java clients.
+
+### Cache
+
+- [ ] Bearer and mTLS entries cannot cross-hit.
+- [ ] Certificate A and B cannot cross-hit.
+- [ ] Attested and unattested entries cannot cross-hit.
+- [ ] Lookup and write use the same request-scoped key.
+
+### Packaging
+
+- [ ] Native DLL version is recorded.
+- [ ] Native DLL hash is tested.
+- [ ] License and notice are included.
+- [ ] E2E code is absent from the production JAR.
+- [ ] No manual DLL deployment is required.
+
+## File-focused checklist
+
+### Core API
+
+- [ ] `ManagedIdentityParameters` validates option combinations.
+- [ ] `IAuthenticationResult` exposes binding metadata without breaking old callers.
+- [ ] `AuthenticationResult` preserves existing metadata.
+- [ ] New interfaces are minimal and documented.
+
+### Supplier
+
+- [ ] Binding acquisition happens before mTLS cache lookup.
+- [ ] Extended cache hash is immutable request state.
+- [ ] IMDS and ESTS retry policies remain separated.
+- [ ] Errors do not leak extension implementation types.
+
+### HTTP stack
+
+- [ ] Socket factory remains request-scoped.
+- [ ] Default HTTP client applies it only to the intended request.
+- [ ] No redirects occur for credential-bound traffic.
+
+### Extension
+
+- [ ] Provider lifecycle is concurrency-safe.
+- [ ] Attestation cache is concurrency-safe.
+- [ ] Native loader is concurrency-safe.
+- [ ] Certificate rotation is concurrency-safe.
+- [ ] Failure paths close handles.
+
+### E2E
+
+- [ ] Uses `IAuthenticationResult` and `IMtlsBindingContext`.
+- [ ] Does not call an MSAL downstream-resource helper.
+- [ ] Verifies `cnf.x5t#S256`.
+- [ ] Requires exact HTTP 200 for the positive resource call.
+- [ ] Requires HTTP 401 when no certificate is presented.
+- [ ] Verifies cache reuse.
+- [ ] Supports force refresh.
+- [ ] Supports token-A/binding-B rejection with a second identity.
+
+## Review completion criteria
+
+The PR is ready only when reviewers can answer yes to each question:
+
+1. Is private key material always non-exportable?
+2. Does JSSE perform TLS without WinHTTP or Schannel as the Java HTTP stack?
+3. Does the standard OAuth token pipeline remain intact?
+4. Are token cache entries bound to the complete certificate identity?
+5. Does attestation fail closed when requested?
+6. Are stale per-boot KeyGuard keys recovered safely?
+7. Are credential-bound redirects prevented?
+8. Can custom HTTP clients fail safely?
+9. Is the native DLL packaged, verified, and licensed?
+10. Does the positive Key Vault call return HTTP 200?
+11. Does the same token without the certificate return HTTP 401?
+12. Does a cache hit retain the correct live binding context?
+13. Are tests and production code Java 8 compatible?
+14. Is the branch still a single coherent commit?
diff --git a/msal4j-sdk/docs/managed-identity-v2-mtls-pop.md b/msal4j-sdk/docs/managed-identity-v2-mtls-pop.md
new file mode 100644
index 00000000..5f1041aa
--- /dev/null
+++ b/msal4j-sdk/docs/managed-identity-v2-mtls-pop.md
@@ -0,0 +1,151 @@
+# Managed Identity v2 attested mTLS PoP
+
+## Architecture
+
+MSAL core owns request validation, identity selection, correlation IDs, HTTP
+policy, proxy behavior, retries, telemetry, OAuth form construction, response
+parsing, and token caching. The optional Windows extension owns only:
+
+1. KeyGuard RSA key creation and NCrypt signing.
+2. CSR construction.
+3. `AttestationClientLib.dll` invocation and MAA JWT caching.
+4. IMDS v2 binding-certificate lifecycle.
+5. A JCA `PrivateKey`, key-selective `Provider`, `SignatureSpi`, and
+ `X509ExtendedKeyManager`.
+
+JSSE performs the token-endpoint and downstream mTLS handshakes:
+
+```text
+JSSE
+ -> X509ExtendedKeyManager
+ -> CngRsaPrivateKey
+ -> CngSignatureSpi
+ -> NCryptSignHash
+ -> VBS KeyGuard
+```
+
+The provider advertises the required RSA signature algorithms, but its services
+accept only `CngRsaPrivateKey`. Ordinary Java RSA keys continue to use the
+platform's normal providers.
+
+## Current protocol contract
+
+Current MSAL.NET product behavior takes precedence over older design and
+prototype material where they differ:
+
+| Area | Current behavior |
+| --- | --- |
+| IMDS API | `cred-api-version=2.0` |
+| Metadata path | `/metadata/identity/getplatformmetadata` |
+| Credential path | `/metadata/identity/issuecredential` |
+| Metadata fields | `clientId`, `tenantId`, `cuId`, `attestationEndpoint` |
+| Credential fields | `certificate`, `client_id`, `tenant_id`, `identity_type`, `mtls_authentication_endpoint` |
+| CSR subject | `CN={clientId}, DC={tenantId}` |
+| CUID attribute | OID `1.3.6.1.4.1.311.90.2.10`, DER UTF8 JSON |
+| CSR signature | RSA-PSS with SHA-256 and 32-byte salt |
+| Certificate rotation | 24 hours before expiry |
+| Token request | client credentials with `token_type=mtls_pop` over JSSE mTLS |
+
+The older design's `api-version=2025-05-01`, challenge-password CUID encoding,
+and three-day rotation window are not used.
+
+## Public API and Java 8 compatibility
+
+Request mTLS PoP and opt into attestation separately:
+
+```java
+ManagedIdentityParameters.builder("https://vault.azure.net")
+ .withMtlsProofOfPossession()
+ .withAttestationSupport()
+ .build();
+```
+
+`withMtlsProofOfPossession()` can be used without attestation.
+`withAttestationSupport()` requires mTLS PoP and makes any attestation failure
+fail closed.
+
+The returned `IMtlsBindingContext` contains an `SSLContext`, leaf certificate,
+and complete-certificate key ID. It contains no token-acquisition helper and no
+native HTTP surface.
+
+Production source remains Java 8 compatible. Use `HttpsURLConnection` for the
+primary compatibility proof. Java 11 `java.net.http.HttpClient` may use the same
+`SSLContext` in an application compiled separately for Java 11 or later.
+
+## Cache safety
+
+MSAL resolves the current binding generation before cache lookup. The extended
+access-token cache key includes:
+
+```text
+token_type = mtls_pop
+key_id = Base64UrlNoPadding(SHA256(leafCertificate.getEncoded()))
+attestation = att0 | att1
+```
+
+This prevents:
+
+- Bearer and mTLS PoP token collisions.
+- Token reuse across certificate renewal.
+- Returning an mTLS token without a matching live binding context.
+
+Native key state and `SSLContext` are process-local and excluded from serialized
+cache data.
+
+## Failure behavior
+
+The flow fails closed when:
+
+- the optional extension is absent or ambiguous;
+- the source is not IMDS VM/VMSS;
+- platform metadata or credential responses are incomplete;
+- identity selection does not match IMDS metadata;
+- KeyGuard or CNG operations fail;
+- the bundled `Microsoft.Azure.Security.KeyGuardAttestation` 1.1.5 native
+ library is missing, corrupt, or cannot be loaded;
+- attestation is empty, malformed, expired, or insufficiently fresh;
+- the issued certificate does not match the KeyGuard public key;
+- the token endpoint does not explicitly return `token_type=mtls_pop`;
+- a configured custom HTTP client does not implement
+ `IMtlsCapableHttpClient`.
+
+Credential-bound token requests do not follow redirects. A custom HTTP client
+must honor `HttpRequest.sslSocketFactory()` and preserve that no-redirect
+behavior.
+
+Errors and logs must not contain access tokens, attestation JWTs, private-key
+material, or native handles.
+
+## Manual Key Vault validation
+
+Set:
+
+```powershell
+$env:MSAL_JAVA_MTLS_AKV_URL = "https://.vault.azure.net"
+$env:MSAL_JAVA_MTLS_AKV_SECRET_NAME = ""
+$env:MSAL_JAVA_MTLS_IDENTITY_CLIENT_ID = ""
+$env:MSAL_JAVA_MTLS_MISMATCH_IDENTITY_CLIENT_ID = ""
+$env:MSAL_JAVA_MTLS_EXPECTED_SECRET_VALUE = ""
+$env:MSAL_JAVA_MTLS_FORCE_REFRESH = "true" # optional
+$env:MSAL_JAVA_MTLS_TOKEN_ONLY = "true" # optional: skip Key Vault call
+```
+
+Run:
+
+```powershell
+.\run-java-msi-v2-mtls-devapp.ps1
+```
+
+The app verifies:
+
+- explicit `mtls_pop`;
+- returned certificate and binding context;
+- JWT `cnf.x5t#S256` equals the complete-certificate key ID;
+- independent Java 8 `HttpsURLConnection` receives Key Vault HTTP 200;
+- the same valid `mtls_pop` token without its binding certificate receives
+ HTTP 401 `Unauthorized`;
+- token A with binding B is rejected when a distinct second managed identity
+ client ID is supplied;
+- second acquisition is `TokenSource.CACHE` with the same binding generation;
+- optional force refresh is `TokenSource.IDENTITY_PROVIDER` and still receives
+ HTTP 200.
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/AcquireTokenByManagedIdentitySupplier.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/AcquireTokenByManagedIdentitySupplier.java
index c6545cf7..34eb1bed 100644
--- a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/AcquireTokenByManagedIdentitySupplier.java
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/AcquireTokenByManagedIdentitySupplier.java
@@ -6,8 +6,12 @@
import org.slf4j.Logger;
import org.slf4j.LoggerFactory;
+import java.net.MalformedURLException;
+import java.net.URL;
import java.time.Instant;
+import java.util.HashMap;
import java.util.HashSet;
+import java.util.Map;
import java.util.Set;
class AcquireTokenByManagedIdentitySupplier extends AuthenticationResultSupplier {
@@ -39,10 +43,22 @@ AuthenticationResult execute() throws Exception {
);
CacheRefreshReason cacheRefreshReason = CacheRefreshReason.NOT_APPLICABLE;
+ ManagedIdentityMtlsBinding mtlsBinding = null;
+
+ if (managedIdentityParameters.mtlsProofOfPossession()) {
+ mtlsBinding = resolveMtlsBinding();
+ String extCacheKeyHash = managedIdentityParameters
+ .computeMtlsExtCacheKeyHash(
+ mtlsBinding.bindingContext().keyId());
+ msalRequest.extCacheKeyHash(extCacheKeyHash);
+ }
if (managedIdentityParameters.forceRefresh) {
LOG.debug("ForceRefresh set to true. Skipping cache lookup and attempting to acquire new token");
- return fetchNewAccessTokenAndSaveToCache(tokenRequestExecutor, CacheRefreshReason.FORCE_REFRESH);
+ return fetchNewAccessTokenAndSaveToCache(
+ tokenRequestExecutor,
+ CacheRefreshReason.FORCE_REFRESH,
+ mtlsBinding);
}
@@ -66,6 +82,7 @@ AuthenticationResult execute() throws Exception {
this.clientApplication,
context,
null);
+ silentRequest.extCacheKeyHash(msalRequest.extCacheKeyHash());
AcquireTokenSilentSupplier supplier = new AcquireTokenSilentSupplier(
this.clientApplication,
@@ -82,21 +99,32 @@ AuthenticationResult execute() throws Exception {
if (cacheRefreshReason == CacheRefreshReason.NOT_APPLICABLE) {
LOG.debug("Returning token from cache");
result.metadata().tokenSource(TokenSource.CACHE);
- return result;
+ return mtlsBinding == null
+ ? result
+ : result.withMtlsBindingContext(mtlsBinding.bindingContext());
} else {
if (cacheRefreshReason == CacheRefreshReason.CLAIMS) {
LOG.debug("Claims are passed, creating token hash and refreshing the token");
managedIdentityParameters.revokedTokenHash = StringHelper.createSha256HashHexString(result.accessToken());
- return fetchNewAccessTokenAndSaveToCache(tokenRequestExecutor, CacheRefreshReason.CLAIMS);
+ return fetchNewAccessTokenAndSaveToCache(
+ tokenRequestExecutor,
+ CacheRefreshReason.CLAIMS,
+ mtlsBinding);
}
LOG.debug("Refreshing access token. Cache refresh reason: {}", cacheRefreshReason);
- return fetchNewAccessTokenAndSaveToCache(tokenRequestExecutor, cacheRefreshReason);
+ return fetchNewAccessTokenAndSaveToCache(
+ tokenRequestExecutor,
+ cacheRefreshReason,
+ mtlsBinding);
}
} catch (MsalClientException ex) {
if (ex.errorCode().equals(AuthenticationErrorCode.CACHE_MISS)) {
LOG.debug("Cache lookup failed: {}", ex.getMessage());
- return fetchNewAccessTokenAndSaveToCache(tokenRequestExecutor, cacheRefreshReason);
+ return fetchNewAccessTokenAndSaveToCache(
+ tokenRequestExecutor,
+ cacheRefreshReason,
+ mtlsBinding);
} else {
LOG.error("Error occurred while cache lookup: {}", ex.getMessage());
throw ex;
@@ -104,24 +132,40 @@ AuthenticationResult execute() throws Exception {
}
}
- private AuthenticationResult fetchNewAccessTokenAndSaveToCache(TokenRequestExecutor tokenRequestExecutor, CacheRefreshReason cacheRefreshReason) {
-
- ManagedIdentityClient managedIdentityClient = new ManagedIdentityClient(msalRequest, tokenRequestExecutor.getServiceBundle());
-
- LOG.debug("[Managed Identity] Managed Identity source and ID type identified and set successfully, request will use Managed Identity for {}",
- managedIdentityClient.managedIdentitySource.managedIdentitySourceType.name());
-
- ManagedIdentityResponse managedIdentityResponse = managedIdentityClient
- .getManagedIdentityResponse(managedIdentityParameters);
+ private AuthenticationResult fetchNewAccessTokenAndSaveToCache(
+ TokenRequestExecutor tokenRequestExecutor,
+ CacheRefreshReason cacheRefreshReason,
+ ManagedIdentityMtlsBinding mtlsBinding) {
+
+ AuthenticationResult authenticationResult;
+ if (mtlsBinding != null) {
+ authenticationResult = acquireMtlsPopToken(
+ mtlsBinding,
+ tokenRequestExecutor);
+ } else {
+ ManagedIdentityClient managedIdentityClient =
+ new ManagedIdentityClient(msalRequest, tokenRequestExecutor.getServiceBundle());
+
+ LOG.debug("[Managed Identity] Managed Identity source and ID type identified and set successfully, request will use Managed Identity for {}",
+ managedIdentityClient.managedIdentitySource.managedIdentitySourceType.name());
+
+ ManagedIdentityResponse managedIdentityResponse = managedIdentityClient
+ .getManagedIdentityResponse(managedIdentityParameters);
+ authenticationResult =
+ createFromManagedIdentityResponse(
+ managedIdentityResponse,
+ null);
+ }
- AuthenticationResult authenticationResult = createFromManagedIdentityResponse(managedIdentityResponse);
clientApplication.tokenCache.saveTokens(tokenRequestExecutor, authenticationResult, clientApplication.authenticationAuthority.host);
authenticationResult.metadata().tokenSource(TokenSource.IDENTITY_PROVIDER);
authenticationResult.metadata().cacheRefreshReason(cacheRefreshReason);
return authenticationResult;
}
- private AuthenticationResult createFromManagedIdentityResponse(ManagedIdentityResponse managedIdentityResponse) {
+ private AuthenticationResult createFromManagedIdentityResponse(
+ ManagedIdentityResponse managedIdentityResponse,
+ ManagedIdentityMtlsBinding mtlsBinding) {
long expiresOn = getExpiresOnFromManagedIdentityTimestamp(managedIdentityResponse.expiresOn);
long refreshOn = calculateRefreshOn(expiresOn);
AuthenticationResultMetadata metadata = AuthenticationResultMetadata.builder()
@@ -136,9 +180,168 @@ private AuthenticationResult createFromManagedIdentityResponse(ManagedIdentityRe
.extExpiresOn(0)
.refreshOn(refreshOn)
.metadata(metadata)
+ .tokenType(managedIdentityResponse.getTokenType())
+ .isPopAuthorization(mtlsBinding == null ? null : Boolean.TRUE)
+ .mtlsBindingContext(mtlsBinding == null ? null : mtlsBinding.bindingContext())
.build();
}
+ private ManagedIdentityMtlsBinding resolveMtlsBinding() {
+ ManagedIdentitySourceType source =
+ ManagedIdentityClient.getManagedIdentitySource();
+ if (source != ManagedIdentitySourceType.DEFAULT_TO_IMDS
+ && source != ManagedIdentitySourceType.IMDS) {
+ throw new MsalClientException(
+ "Managed identity mTLS PoP is supported only on the IMDS v2 VM/VMSS source.",
+ MsalError.MANAGED_IDENTITY_MTLS_REQUEST_FAILED);
+ }
+
+ ManagedIdentityApplication application =
+ (ManagedIdentityApplication) msalRequest.application();
+ ManagedIdentityId identity = application.getManagedIdentityId();
+ String queryName = null;
+ String queryValue = identity.getUserAssignedId();
+ switch (identity.getIdType()) {
+ case CLIENT_ID:
+ queryName = Constants.MANAGED_IDENTITY_CLIENT_ID;
+ break;
+ case RESOURCE_ID:
+ queryName = Constants.MANAGED_IDENTITY_RESOURCE_ID_IMDS;
+ break;
+ case OBJECT_ID:
+ queryName = Constants.MANAGED_IDENTITY_OBJECT_ID;
+ break;
+ case SYSTEM_ASSIGNED:
+ queryValue = null;
+ break;
+ default:
+ throw new MsalClientException(
+ "Unsupported managed identity selector for mTLS PoP.",
+ MsalError.MANAGED_IDENTITY_MTLS_REQUEST_FAILED);
+ }
+
+ final ServiceBundle serviceBundle = clientApplication.serviceBundle();
+ final RequestContext requestContext = msalRequest.requestContext();
+ final HttpHelper imdsHttpHelper = new HttpHelper(
+ clientApplication,
+ new IMDSRetryPolicy());
+ IManagedIdentityMtlsHttpClient httpClient = createMtlsProviderHttpClient(
+ imdsHttpHelper,
+ serviceBundle,
+ requestContext);
+
+ String bindingCacheKey = identity.getIdType().name() + ":"
+ + (queryValue == null ? "" : queryValue)
+ + (managedIdentityParameters.attestationSupport()
+ ? ":att1" : ":att0");
+ ManagedIdentityMtlsRequest request = new ManagedIdentityMtlsRequest(
+ queryName,
+ queryValue,
+ bindingCacheKey,
+ requestContext.correlationId(),
+ httpClient,
+ managedIdentityParameters.attestationSupport());
+ return getMtlsProviderBinding(
+ ManagedIdentityMtlsProviderLoader.load(),
+ request);
+ }
+
+ static ManagedIdentityMtlsBinding getMtlsProviderBinding(
+ IManagedIdentityMtlsProvider provider,
+ ManagedIdentityMtlsRequest request) {
+ try {
+ return provider.getOrCreateBinding(request);
+ } catch (MsalException e) {
+ throw e;
+ } catch (RuntimeException e) {
+ MsalClientException wrapped = new MsalClientException(
+ "The managed identity mTLS provider failed.",
+ MsalError.MANAGED_IDENTITY_MTLS_REQUEST_FAILED);
+ wrapped.initCause(e);
+ throw wrapped;
+ }
+ }
+
+ static IManagedIdentityMtlsHttpClient createMtlsProviderHttpClient(
+ HttpHelper imdsHttpHelper,
+ ServiceBundle serviceBundle,
+ RequestContext requestContext) {
+ return request -> {
+ HttpMethod method;
+ if ("GET".equalsIgnoreCase(request.method())) {
+ method = HttpMethod.GET;
+ } else if ("POST".equalsIgnoreCase(request.method())) {
+ method = HttpMethod.POST;
+ } else {
+ throw new MsalClientException(
+ "Unsupported IMDS mTLS provider HTTP method: " + request.method(),
+ MsalError.MANAGED_IDENTITY_MTLS_REQUEST_FAILED);
+ }
+
+ HttpRequest httpRequest = new HttpRequest(
+ method,
+ request.url(),
+ request.headers(),
+ request.body());
+ IHttpResponse response = imdsHttpHelper
+ .executeHttpRequest(httpRequest, requestContext, serviceBundle);
+ return new ManagedIdentityMtlsHttpResponse(
+ response.statusCode(),
+ response.body(),
+ response.headers());
+ };
+ }
+
+ private AuthenticationResult acquireMtlsPopToken(
+ ManagedIdentityMtlsBinding binding,
+ TokenRequestExecutor tokenRequestExecutor) {
+ if (!(clientApplication.httpClient() instanceof IMtlsCapableHttpClient)) {
+ throw new MsalClientException(
+ "The configured custom HTTP client does not declare support for request-specific mTLS. "
+ + "Implement IMtlsCapableHttpClient and honor HttpRequest.sslSocketFactory().",
+ MsalError.MANAGED_IDENTITY_MTLS_REQUEST_FAILED);
+ }
+
+ String scope = managedIdentityParameters.resource().endsWith("/.default")
+ ? managedIdentityParameters.resource()
+ : managedIdentityParameters.resource().replaceAll("/+$", "") + "/.default";
+ Map body = new HashMap<>();
+ body.put("grant_type", "client_credentials");
+ body.put("client_id", binding.clientId());
+ body.put("scope", scope);
+ body.put("token_type", "mtls_pop");
+ AuthenticationResult result;
+ try {
+ result = tokenRequestExecutor.executeTokenRequest(
+ new URL(binding.tokenEndpoint()),
+ binding.bindingContext().sslContext().getSocketFactory(),
+ body);
+ } catch (MalformedURLException e) {
+ throw new MsalClientException(
+ "The managed identity mTLS token endpoint is invalid.",
+ MsalError.MANAGED_IDENTITY_MTLS_REQUEST_FAILED);
+ } catch (java.io.IOException e) {
+ throw new MsalClientException(e);
+ }
+
+ validateMtlsTokenResponse(result);
+ return result.withMtlsBindingContext(
+ binding.bindingContext(),
+ managedIdentityParameters.resource());
+ }
+
+ static void validateMtlsTokenResponse(
+ IAuthenticationResult tokenResponse) {
+ if (tokenResponse == null
+ || StringHelper.isBlank(tokenResponse.accessToken())
+ || !"mtls_pop".equals(tokenResponse.tokenType())) {
+ throw new MsalServiceException(
+ "The managed identity mTLS endpoint did not explicitly return token_type=mtls_pop.",
+ MsalError.MANAGED_IDENTITY_MTLS_TOKEN_TYPE_INVALID,
+ ManagedIdentitySourceType.IMDS);
+ }
+ }
+
static long getExpiresOnFromManagedIdentityTimestamp(String dateTimeStamp) {
if (dateTimeStamp == null || dateTimeStamp.isEmpty()) {
return 0;
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/AuthenticationResult.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/AuthenticationResult.java
index d87dfc4b..876f443d 100644
--- a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/AuthenticationResult.java
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/AuthenticationResult.java
@@ -5,6 +5,7 @@
import java.util.Date;
import java.util.Objects;
+import java.security.cert.X509Certificate;
final class AuthenticationResult implements IAuthenticationResult {
private static final long serialVersionUID = 1L;
@@ -25,8 +26,10 @@ final class AuthenticationResult implements IAuthenticationResult {
private final String scopes;
private final AuthenticationResultMetadata metadata;
private final Boolean isPopAuthorization;
+ private final String tokenType;
+ private final transient IMtlsBindingContext mtlsBindingContext;
- AuthenticationResult(String accessToken, long expiresOn, long extExpiresOn, String refreshToken, Long refreshOn, String familyId, String idToken, AccountCacheEntity accountCacheEntity, String environment, String scopes, AuthenticationResultMetadata metadata, Boolean isPopAuthorization) {
+ AuthenticationResult(String accessToken, long expiresOn, long extExpiresOn, String refreshToken, Long refreshOn, String familyId, String idToken, AccountCacheEntity accountCacheEntity, String environment, String scopes, AuthenticationResultMetadata metadata, Boolean isPopAuthorization, String tokenType, IMtlsBindingContext mtlsBindingContext) {
this.accessToken = accessToken;
this.expiresOn = expiresOn;
this.extExpiresOn = extExpiresOn;
@@ -39,6 +42,8 @@ final class AuthenticationResult implements IAuthenticationResult {
this.scopes = scopes;
this.metadata = metadata == null ? AuthenticationResultMetadata.builder().build() : metadata;
this.isPopAuthorization = isPopAuthorization;
+ this.tokenType = StringHelper.isBlank(tokenType) ? "Bearer" : tokenType;
+ this.mtlsBindingContext = mtlsBindingContext;
this.expiresOnDate = new Date(expiresOn * 1000);
}
@@ -129,6 +134,59 @@ Boolean isPopAuthorization() {
return this.isPopAuthorization;
}
+ AuthenticationResult withMtlsBindingContext(IMtlsBindingContext bindingContext) {
+ return withMtlsBindingContext(
+ bindingContext,
+ null,
+ "mtls_pop");
+ }
+
+ AuthenticationResult withMtlsBindingContext(
+ IMtlsBindingContext bindingContext,
+ String defaultScopes) {
+ return withMtlsBindingContext(
+ bindingContext,
+ defaultScopes,
+ tokenType);
+ }
+
+ private AuthenticationResult withMtlsBindingContext(
+ IMtlsBindingContext bindingContext,
+ String defaultScopes,
+ String resultTokenType) {
+ return AuthenticationResult.builder()
+ .accessToken(accessToken)
+ .expiresOn(expiresOn)
+ .extExpiresOn(extExpiresOn)
+ .refreshToken(refreshToken)
+ .refreshOn(refreshOn)
+ .familyId(familyId)
+ .idToken(idToken)
+ .accountCacheEntity(accountCacheEntity)
+ .environment(environment)
+ .scopes(StringHelper.isBlank(scopes) ? defaultScopes : scopes)
+ .metadata(metadata)
+ .isPopAuthorization(Boolean.TRUE)
+ .tokenType(resultTokenType)
+ .mtlsBindingContext(bindingContext)
+ .build();
+ }
+
+ @Override
+ public String tokenType() {
+ return StringHelper.isBlank(tokenType) ? "Bearer" : tokenType;
+ }
+
+ @Override
+ public X509Certificate bindingCertificate() {
+ return mtlsBindingContext == null ? null : mtlsBindingContext.bindingCertificate();
+ }
+
+ @Override
+ public IMtlsBindingContext mtlsBindingContext() {
+ return mtlsBindingContext;
+ }
+
static AuthenticationResultBuilder builder() {
return new AuthenticationResultBuilder();
}
@@ -146,6 +204,8 @@ static class AuthenticationResultBuilder {
private String scopes;
private AuthenticationResultMetadata metadata;
private Boolean isPopAuthorization;
+ private String tokenType;
+ private IMtlsBindingContext mtlsBindingContext;
AuthenticationResultBuilder() {
}
@@ -210,12 +270,28 @@ public AuthenticationResultBuilder isPopAuthorization(Boolean isPopAuthorization
return this;
}
+ public AuthenticationResultBuilder tokenType(String tokenType) {
+ this.tokenType = tokenType;
+ return this;
+ }
+
+ public AuthenticationResultBuilder mtlsBindingContext(IMtlsBindingContext mtlsBindingContext) {
+ this.mtlsBindingContext = mtlsBindingContext;
+ return this;
+ }
+
public AuthenticationResult build() {
- return new AuthenticationResult(this.accessToken, this.expiresOn, this.extExpiresOn, this.refreshToken, this.refreshOn, this.familyId, this.idToken, this.accountCacheEntity, this.environment, this.scopes, this.metadata, this.isPopAuthorization);
+ return new AuthenticationResult(this.accessToken, this.expiresOn, this.extExpiresOn, this.refreshToken, this.refreshOn, this.familyId, this.idToken, this.accountCacheEntity, this.environment, this.scopes, this.metadata, this.isPopAuthorization, this.tokenType, this.mtlsBindingContext);
}
public String toString() {
- return "AuthenticationResult.AuthenticationResultBuilder(accessToken=" + this.accessToken + ", expiresOn=" + this.expiresOn + ", extExpiresOn=" + this.extExpiresOn + ", refreshToken=" + this.refreshToken + ", refreshOn=" + this.refreshOn + ", familyId=" + this.familyId + ", idToken=" + this.idToken + ", accountCacheEntity=" + this.accountCacheEntity + ", environment=" + this.environment + ", scopes=" + this.scopes + ", metadata=" + this.metadata + ", isPopAuthorization=" + this.isPopAuthorization + ")";
+ return "AuthenticationResult.AuthenticationResultBuilder(expiresOn=" + this.expiresOn
+ + ", extExpiresOn=" + this.extExpiresOn
+ + ", environment=" + this.environment
+ + ", scopes=" + this.scopes
+ + ", metadata=" + this.metadata
+ + ", isPopAuthorization=" + this.isPopAuthorization
+ + ", tokenType=" + this.tokenType + ")";
}
}
@@ -243,6 +319,7 @@ public boolean equals(Object o) {
if (!Objects.equals(environment, other.environment)) return false;
if (!Objects.equals(expiresOnDate, other.expiresOnDate)) return false;
if (!Objects.equals(scopes, other.scopes)) return false;
+ if (!Objects.equals(tokenType, other.tokenType)) return false;
return Objects.equals(metadata, other.metadata);
}
@@ -264,6 +341,7 @@ public int hashCode() {
result = result * 59 + (this.environment == null ? 43 : this.environment.hashCode());
result = result * 59 + (this.expiresOnDate == null ? 43 : this.expiresOnDate.hashCode());
result = result * 59 + (this.scopes == null ? 43 : this.scopes.hashCode());
+ result = result * 59 + (this.tokenType == null ? 43 : this.tokenType.hashCode());
result = result * 59 + (this.metadata == null ? 43 : this.metadata.hashCode());
return result;
}
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/DefaultHttpClient.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/DefaultHttpClient.java
index 7a88a2a4..6f0fbe8b 100644
--- a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/DefaultHttpClient.java
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/DefaultHttpClient.java
@@ -20,7 +20,7 @@
import java.nio.charset.StandardCharsets;
import java.util.Map;
-class DefaultHttpClient implements IHttpClient {
+class DefaultHttpClient implements IMtlsCapableHttpClient {
private static final Logger LOG = LoggerFactory.getLogger(DefaultHttpClient.class);
final Proxy proxy;
@@ -50,7 +50,7 @@ public IHttpResponse send(HttpRequest httpRequest) throws Exception {
private HttpResponse executeHttpGet(HttpRequest httpRequest) throws Exception {
- final HttpURLConnection conn = openConnection(httpRequest.url());
+ final HttpURLConnection conn = openConnection(httpRequest.url(), httpRequest.sslSocketFactory());
configureAdditionalHeaders(conn, httpRequest);
return readResponseFromConnection(conn);
@@ -58,7 +58,7 @@ private HttpResponse executeHttpGet(HttpRequest httpRequest) throws Exception {
private HttpResponse executeHttpPost(HttpRequest httpRequest) throws Exception {
- final HttpURLConnection conn = openConnection(httpRequest.url());
+ final HttpURLConnection conn = openConnection(httpRequest.url(), httpRequest.sslSocketFactory());
configureAdditionalHeaders(conn, httpRequest);
conn.setRequestMethod("POST");
conn.setDoOutput(true);
@@ -79,6 +79,11 @@ private HttpResponse executeHttpPost(HttpRequest httpRequest) throws Exception {
HttpURLConnection openConnection(final URL finalURL)
throws IOException {
+ return openConnection(finalURL, null);
+ }
+
+ HttpURLConnection openConnection(final URL finalURL, SSLSocketFactory requestSslSocketFactory)
+ throws IOException {
URLConnection connection;
if (proxy != null) {
@@ -93,8 +98,13 @@ HttpURLConnection openConnection(final URL finalURL)
if (connection instanceof HttpsURLConnection) {
HttpsURLConnection httpsConnection = (HttpsURLConnection) connection;
- if (sslSocketFactory != null) {
- httpsConnection.setSSLSocketFactory(sslSocketFactory);
+ SSLSocketFactory effectiveSslSocketFactory =
+ requestSslSocketFactory != null ? requestSslSocketFactory : sslSocketFactory;
+ if (effectiveSslSocketFactory != null) {
+ httpsConnection.setSSLSocketFactory(effectiveSslSocketFactory);
+ }
+ if (requestSslSocketFactory != null) {
+ httpsConnection.setInstanceFollowRedirects(false);
}
return httpsConnection;
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/HttpRequest.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/HttpRequest.java
index 2923f799..b0cfb903 100644
--- a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/HttpRequest.java
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/HttpRequest.java
@@ -7,6 +7,7 @@
import java.net.URL;
import java.util.Map;
import java.util.Objects;
+import javax.net.ssl.SSLSocketFactory;
/**
* Contains information about outgoing HTTP request. Should be adapted to HTTP request for HTTP
@@ -34,6 +35,8 @@ public class HttpRequest {
*/
private String body;
+ private transient SSLSocketFactory sslSocketFactory;
+
HttpRequest(HttpMethod httpMethod, String url) {
this.httpMethod = httpMethod;
this.url = createUrlFromString(url);
@@ -100,6 +103,18 @@ public String body() {
return this.body;
}
+ /**
+ * Returns a request-specific JSSE socket factory, when the request requires mTLS.
+ */
+ public SSLSocketFactory sslSocketFactory() {
+ return sslSocketFactory;
+ }
+
+ HttpRequest sslSocketFactory(SSLSocketFactory sslSocketFactory) {
+ this.sslSocketFactory = sslSocketFactory;
+ return this;
+ }
+
//These methods are based on those generated by Lombok's @EqualsAndHashCode annotation.
//They have the same functionality as the generated methods, but were refactored for readability.
@Override
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/IAuthenticationResult.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/IAuthenticationResult.java
index 934a2d2c..1bca463c 100644
--- a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/IAuthenticationResult.java
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/IAuthenticationResult.java
@@ -4,6 +4,7 @@
package com.microsoft.aad.msal4j;
import java.io.Serializable;
+import java.security.cert.X509Certificate;
/**
* Interface representing the results of token acquisition operation.
@@ -51,4 +52,25 @@ public interface IAuthenticationResult extends Serializable {
default AuthenticationResultMetadata metadata() {
return AuthenticationResultMetadata.builder().build();
}
+
+ /**
+ * @return token type returned by the identity provider
+ */
+ default String tokenType() {
+ return "Bearer";
+ }
+
+ /**
+ * @return binding certificate for an mTLS PoP result, otherwise null
+ */
+ default X509Certificate bindingCertificate() {
+ return null;
+ }
+
+ /**
+ * @return live process-local binding context for an mTLS PoP result, otherwise null
+ */
+ default IMtlsBindingContext mtlsBindingContext() {
+ return null;
+ }
}
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/IManagedIdentityMtlsHttpClient.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/IManagedIdentityMtlsHttpClient.java
new file mode 100644
index 00000000..bfc88ea5
--- /dev/null
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/IManagedIdentityMtlsHttpClient.java
@@ -0,0 +1,12 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j;
+
+/**
+ * MSAL-owned HTTP callback used by the optional mTLS provider for IMDS v2 requests.
+ */
+public interface IManagedIdentityMtlsHttpClient {
+
+ ManagedIdentityMtlsHttpResponse execute(ManagedIdentityMtlsHttpRequest request);
+}
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/IManagedIdentityMtlsProvider.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/IManagedIdentityMtlsProvider.java
new file mode 100644
index 00000000..d103706f
--- /dev/null
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/IManagedIdentityMtlsProvider.java
@@ -0,0 +1,15 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j;
+
+/**
+ * Optional platform extension for attested KeyGuard managed identity mTLS PoP bindings.
+ *
+ * Implementations are discovered with {@link java.util.ServiceLoader}. MSAL core remains
+ * loadable when the optional Windows extension is absent.
+ */
+public interface IManagedIdentityMtlsProvider {
+
+ ManagedIdentityMtlsBinding getOrCreateBinding(ManagedIdentityMtlsRequest request);
+}
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/IMtlsBindingContext.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/IMtlsBindingContext.java
new file mode 100644
index 00000000..300952ac
--- /dev/null
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/IMtlsBindingContext.java
@@ -0,0 +1,25 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j;
+
+import javax.net.ssl.SSLContext;
+import java.security.cert.X509Certificate;
+
+/**
+ * Process-local mTLS binding capability associated with an mTLS PoP access token.
+ *
+ * The private key is not exportable. Applications use the returned {@link SSLContext}
+ * with a Java JSSE HTTP stack when calling the downstream resource.
+ */
+public interface IMtlsBindingContext {
+
+ SSLContext sslContext();
+
+ X509Certificate bindingCertificate();
+
+ /**
+ * Base64URL-without-padding SHA-256 digest of the complete leaf certificate DER.
+ */
+ String keyId();
+}
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/IMtlsCapableHttpClient.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/IMtlsCapableHttpClient.java
new file mode 100644
index 00000000..1f0242d3
--- /dev/null
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/IMtlsCapableHttpClient.java
@@ -0,0 +1,13 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j;
+
+/**
+ * Marker contract for custom HTTP clients that honor request-specific mTLS settings.
+ *
+ * Implementations must use {@link HttpRequest#sslSocketFactory()} when it is present
+ * and must not automatically follow redirects for that credential-bound request.
+ */
+public interface IMtlsCapableHttpClient extends IHttpClient {
+}
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsBinding.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsBinding.java
new file mode 100644
index 00000000..92e07cad
--- /dev/null
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsBinding.java
@@ -0,0 +1,57 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j;
+
+import java.net.MalformedURLException;
+import java.net.URL;
+
+/**
+ * Immutable binding generation returned by the optional managed identity mTLS provider.
+ */
+public final class ManagedIdentityMtlsBinding {
+
+ private final IMtlsBindingContext bindingContext;
+ private final String clientId;
+ private final String tokenEndpoint;
+
+ public ManagedIdentityMtlsBinding(
+ IMtlsBindingContext bindingContext,
+ String clientId,
+ String tokenEndpoint) {
+ if (bindingContext == null) {
+ throw new NullPointerException("bindingContext");
+ }
+ if (clientId == null || clientId.trim().isEmpty()) {
+ throw new IllegalArgumentException("clientId must not be blank");
+ }
+ if (tokenEndpoint == null || tokenEndpoint.trim().isEmpty()) {
+ throw new IllegalArgumentException("tokenEndpoint must not be blank");
+ }
+ try {
+ URL endpoint = new URL(tokenEndpoint);
+ if (!"https".equalsIgnoreCase(endpoint.getProtocol())) {
+ throw new IllegalArgumentException(
+ "tokenEndpoint must use HTTPS");
+ }
+ } catch (MalformedURLException e) {
+ throw new IllegalArgumentException(
+ "tokenEndpoint must be a valid HTTPS URL", e);
+ }
+ this.bindingContext = bindingContext;
+ this.clientId = clientId;
+ this.tokenEndpoint = tokenEndpoint;
+ }
+
+ public IMtlsBindingContext bindingContext() {
+ return bindingContext;
+ }
+
+ public String clientId() {
+ return clientId;
+ }
+
+ public String tokenEndpoint() {
+ return tokenEndpoint;
+ }
+}
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsHttpRequest.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsHttpRequest.java
new file mode 100644
index 00000000..d08347b5
--- /dev/null
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsHttpRequest.java
@@ -0,0 +1,48 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j;
+
+import java.util.Collections;
+import java.util.LinkedHashMap;
+import java.util.Map;
+
+/**
+ * Immutable IMDS request issued through MSAL's HTTP, retry, proxy and telemetry pipeline.
+ */
+public final class ManagedIdentityMtlsHttpRequest {
+
+ private final String method;
+ private final String url;
+ private final Map headers;
+ private final String body;
+
+ public ManagedIdentityMtlsHttpRequest(
+ String method,
+ String url,
+ Map headers,
+ String body) {
+ this.method = method;
+ this.url = url;
+ this.headers = headers == null
+ ? Collections.emptyMap()
+ : Collections.unmodifiableMap(new LinkedHashMap<>(headers));
+ this.body = body;
+ }
+
+ public String method() {
+ return method;
+ }
+
+ public String url() {
+ return url;
+ }
+
+ public Map headers() {
+ return headers;
+ }
+
+ public String body() {
+ return body;
+ }
+}
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsHttpResponse.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsHttpResponse.java
new file mode 100644
index 00000000..1b33c9b4
--- /dev/null
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsHttpResponse.java
@@ -0,0 +1,41 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j;
+
+import java.util.Collections;
+import java.util.List;
+import java.util.Map;
+
+/**
+ * Response returned to the optional mTLS provider by MSAL's HTTP pipeline.
+ */
+public final class ManagedIdentityMtlsHttpResponse {
+
+ private final int statusCode;
+ private final String body;
+ private final Map> headers;
+
+ public ManagedIdentityMtlsHttpResponse(
+ int statusCode,
+ String body,
+ Map> headers) {
+ this.statusCode = statusCode;
+ this.body = body;
+ this.headers = headers == null
+ ? Collections.>emptyMap()
+ : Collections.unmodifiableMap(headers);
+ }
+
+ public int statusCode() {
+ return statusCode;
+ }
+
+ public String body() {
+ return body;
+ }
+
+ public Map> headers() {
+ return headers;
+ }
+}
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsProviderLoader.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsProviderLoader.java
new file mode 100644
index 00000000..e0acaf82
--- /dev/null
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsProviderLoader.java
@@ -0,0 +1,33 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j;
+
+import java.util.Iterator;
+import java.util.ServiceLoader;
+
+final class ManagedIdentityMtlsProviderLoader {
+
+ private ManagedIdentityMtlsProviderLoader() {
+ }
+
+ static IManagedIdentityMtlsProvider load() {
+ ServiceLoader loader =
+ ServiceLoader.load(IManagedIdentityMtlsProvider.class);
+ Iterator providers = loader.iterator();
+ if (!providers.hasNext()) {
+ throw new MsalClientException(
+ "Managed identity mTLS PoP requires the optional "
+ + "com.microsoft.azure:msal4j-mtls-extensions dependency.",
+ MsalError.MANAGED_IDENTITY_MTLS_PROVIDER_UNAVAILABLE);
+ }
+
+ IManagedIdentityMtlsProvider provider = providers.next();
+ if (providers.hasNext()) {
+ throw new MsalClientException(
+ "Multiple managed identity mTLS providers were found. Configure exactly one provider.",
+ MsalError.MANAGED_IDENTITY_MTLS_PROVIDER_UNAVAILABLE);
+ }
+ return provider;
+ }
+}
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsRequest.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsRequest.java
new file mode 100644
index 00000000..da92692a
--- /dev/null
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsRequest.java
@@ -0,0 +1,66 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j;
+
+/**
+ * Request passed from MSAL core to the optional KeyGuard mTLS provider.
+ */
+public final class ManagedIdentityMtlsRequest {
+
+ private final String identityQueryParameter;
+ private final String identityQueryValue;
+ private final String bindingCacheKey;
+ private final String correlationId;
+ private final IManagedIdentityMtlsHttpClient httpClient;
+ private final boolean attestationEnabled;
+
+ public ManagedIdentityMtlsRequest(
+ String identityQueryParameter,
+ String identityQueryValue,
+ String bindingCacheKey,
+ String correlationId,
+ IManagedIdentityMtlsHttpClient httpClient) {
+ this(identityQueryParameter, identityQueryValue, bindingCacheKey,
+ correlationId, httpClient, true);
+ }
+
+ public ManagedIdentityMtlsRequest(
+ String identityQueryParameter,
+ String identityQueryValue,
+ String bindingCacheKey,
+ String correlationId,
+ IManagedIdentityMtlsHttpClient httpClient,
+ boolean attestationEnabled) {
+ this.identityQueryParameter = identityQueryParameter;
+ this.identityQueryValue = identityQueryValue;
+ this.bindingCacheKey = bindingCacheKey;
+ this.correlationId = correlationId;
+ this.httpClient = httpClient;
+ this.attestationEnabled = attestationEnabled;
+ }
+
+ public String identityQueryParameter() {
+ return identityQueryParameter;
+ }
+
+ public String identityQueryValue() {
+ return identityQueryValue;
+ }
+
+ public String bindingCacheKey() {
+ return bindingCacheKey;
+ }
+
+ public String correlationId() {
+ return correlationId;
+ }
+
+ public IManagedIdentityMtlsHttpClient httpClient() {
+ return httpClient;
+ }
+
+ public boolean attestationEnabled() {
+ return attestationEnabled;
+ }
+}
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityParameters.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityParameters.java
index 21335802..51d44592 100644
--- a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityParameters.java
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityParameters.java
@@ -5,6 +5,8 @@
import java.util.Map;
import java.util.Set;
+import java.util.SortedMap;
+import java.util.TreeMap;
/**
* Object containing parameters for managed identity flow. Can be used as parameter to
@@ -16,11 +18,20 @@ public class ManagedIdentityParameters implements IAcquireTokenParameters {
boolean forceRefresh;
String claims;
String revokedTokenHash;
+ boolean mtlsProofOfPossession;
+ boolean attestationSupport;
- private ManagedIdentityParameters(String resource, boolean forceRefresh, String claims) {
+ private ManagedIdentityParameters(
+ String resource,
+ boolean forceRefresh,
+ String claims,
+ boolean mtlsProofOfPossession,
+ boolean attestationSupport) {
this.resource = resource;
this.forceRefresh = forceRefresh;
this.claims = claims;
+ this.mtlsProofOfPossession = mtlsProofOfPossession;
+ this.attestationSupport = attestationSupport;
}
@Override
@@ -83,10 +94,36 @@ public String revokedTokenHash() {
return this.revokedTokenHash;
}
+ public boolean mtlsProofOfPossession() {
+ return mtlsProofOfPossession;
+ }
+
+ public boolean attestationSupport() {
+ return attestationSupport;
+ }
+
+ @Override
+ public String computeExtCacheKeyHash() {
+ return "";
+ }
+
+ String computeMtlsExtCacheKeyHash(String bindingKeyId) {
+ if (!mtlsProofOfPossession || StringHelper.isBlank(bindingKeyId)) {
+ return "";
+ }
+ SortedMap components = new TreeMap<>();
+ components.put("token_type", "mtls_pop");
+ components.put("key_id", bindingKeyId);
+ components.put("attestation", attestationSupport ? "att1" : "att0");
+ return StringHelper.computeExtCacheKeyHash(components);
+ }
+
public static class ManagedIdentityParametersBuilder {
private String resource;
private boolean forceRefresh;
private String claims;
+ private boolean mtlsProofOfPossession;
+ private boolean attestationSupport;
ManagedIdentityParametersBuilder() {
}
@@ -118,12 +155,41 @@ public ManagedIdentityParametersBuilder claims(String claims) {
return this;
}
+ /**
+ * Requests the KeyGuard managed identity v2 mTLS PoP flow.
+ */
+ public ManagedIdentityParametersBuilder withMtlsProofOfPossession() {
+ this.mtlsProofOfPossession = true;
+ return this;
+ }
+
+ /**
+ * Requires MAA attestation for the KeyGuard binding key.
+ * This option requires {@link #withMtlsProofOfPossession()}.
+ */
+ public ManagedIdentityParametersBuilder withAttestationSupport() {
+ this.attestationSupport = true;
+ return this;
+ }
+
public ManagedIdentityParameters build() {
- return new ManagedIdentityParameters(this.resource, this.forceRefresh, this.claims);
+ if (attestationSupport && !mtlsProofOfPossession) {
+ throw new IllegalArgumentException(
+ "Attestation support requires managed identity mTLS PoP.");
+ }
+ return new ManagedIdentityParameters(
+ this.resource,
+ this.forceRefresh,
+ this.claims,
+ this.mtlsProofOfPossession,
+ this.attestationSupport);
}
public String toString() {
- return "ManagedIdentityParameters.ManagedIdentityParametersBuilder(resource=" + this.resource + ", forceRefresh=" + this.forceRefresh + ")";
+ return "ManagedIdentityParameters.ManagedIdentityParametersBuilder(resource=" + this.resource
+ + ", forceRefresh=" + this.forceRefresh
+ + ", mtlsProofOfPossession=" + this.mtlsProofOfPossession
+ + ", attestationSupport=" + this.attestationSupport + ")";
}
}
}
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityResponse.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityResponse.java
index f6eb4e7e..fa44f57b 100644
--- a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityResponse.java
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/ManagedIdentityResponse.java
@@ -37,6 +37,10 @@ public static ManagedIdentityResponse fromJson(JsonReader jsonReader) throws IOE
case "expires_on":
response.expiresOn = reader.getString();
break;
+ case "expires_in":
+ response.expiresOn = String.valueOf(
+ (System.currentTimeMillis() / 1000) + reader.getLong());
+ break;
case "resource":
response.resource = reader.getString();
break;
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/MsalError.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/MsalError.java
index 23d4e546..edda89b8 100644
--- a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/MsalError.java
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/MsalError.java
@@ -38,4 +38,13 @@ public class MsalError {
public static final String MANAGED_IDENTITY_FILE_READ_ERROR = "managed_identity_file_read_error";
public static final String MANAGED_IDENTITY_RESPONSE_PARSE_FAILURE = "managed_identity_response_parse_failure";
+
+ public static final String MANAGED_IDENTITY_MTLS_PROVIDER_UNAVAILABLE =
+ "managed_identity_mtls_provider_unavailable";
+
+ public static final String MANAGED_IDENTITY_MTLS_TOKEN_TYPE_INVALID =
+ "managed_identity_mtls_token_type_invalid";
+
+ public static final String MANAGED_IDENTITY_MTLS_REQUEST_FAILED =
+ "managed_identity_mtls_request_failed";
}
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/MsalRequest.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/MsalRequest.java
index f730f5bd..e5134ac3 100644
--- a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/MsalRequest.java
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/MsalRequest.java
@@ -9,6 +9,7 @@ abstract class MsalRequest {
private final AbstractApplicationBase application;
private final RequestContext requestContext;
private final HttpHeaders headers;
+ private String extCacheKeyHash;
MsalRequest(AbstractApplicationBase clientApplicationBase, AbstractMsalAuthorizationGrant abstractMsalAuthorizationGrant, RequestContext requestContext) {
this.application = clientApplicationBase;
@@ -44,4 +45,12 @@ RequestContext requestContext() {
HttpHeaders headers() {
return this.headers;
}
+
+ String extCacheKeyHash() {
+ return extCacheKeyHash;
+ }
+
+ void extCacheKeyHash(String extCacheKeyHash) {
+ this.extCacheKeyHash = extCacheKeyHash;
+ }
}
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/OAuthHttpRequest.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/OAuthHttpRequest.java
index 49ecc2fc..4a4a0ed4 100644
--- a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/OAuthHttpRequest.java
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/OAuthHttpRequest.java
@@ -10,6 +10,7 @@
import java.util.HashMap;
import java.util.List;
import java.util.Map;
+import javax.net.ssl.SSLSocketFactory;
class OAuthHttpRequest {
@@ -19,6 +20,7 @@ class OAuthHttpRequest {
private final Map extraHeaderParams;
private final ServiceBundle serviceBundle;
private final RequestContext requestContext;
+ private SSLSocketFactory sslSocketFactory;
OAuthHttpRequest(final HttpMethod method,
final URL url,
@@ -39,7 +41,8 @@ public HttpResponse send() throws IOException {
HttpMethod.POST,
this.url.toString(),
httpHeaders,
- this.query);
+ this.query)
+ .sslSocketFactory(sslSocketFactory);
IHttpResponse httpResponse = serviceBundle.getHttpHelper().executeHttpRequest(
httpRequest,
@@ -107,4 +110,13 @@ void setQuery(String query) {
Map getExtraHeaderParams() {
return this.extraHeaderParams;
}
+
+ OAuthHttpRequest sslSocketFactory(SSLSocketFactory sslSocketFactory) {
+ this.sslSocketFactory = sslSocketFactory;
+ return this;
+ }
+
+ SSLSocketFactory sslSocketFactory() {
+ return sslSocketFactory;
+ }
}
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/TokenCache.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/TokenCache.java
index 59a7b903..3afd5323 100644
--- a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/TokenCache.java
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/TokenCache.java
@@ -346,6 +346,10 @@ private static AccessTokenCacheEntity createAccessTokenCacheEntity(TokenRequestE
* The algorithm uses sorted key-value concatenation → SHA-256 → Base64URL (cross-SDK compatible).
*/
private static String computeExtCacheKeyHashForRequest(MsalRequest msalRequest) {
+ if (!StringHelper.isBlank(msalRequest.extCacheKeyHash())) {
+ return msalRequest.extCacheKeyHash();
+ }
+
// A RefreshTokenRequest inherits the parent silent request's RequestContext, whose
// apiParameters (SilentParameters) carries no client-originated claims and would therefore
// return an empty hash. Prefer the hash threaded onto the parent silent request so a refreshed
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/TokenRequestExecutor.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/TokenRequestExecutor.java
index 8c3cc62b..460ba083 100644
--- a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/TokenRequestExecutor.java
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/TokenRequestExecutor.java
@@ -8,7 +8,9 @@
import java.io.IOException;
import java.net.MalformedURLException;
+import java.net.URL;
import java.util.*;
+import javax.net.ssl.SSLSocketFactory;
class TokenRequestExecutor {
private static final Logger LOG = LoggerFactory.getLogger(TokenRequestExecutor.class);
@@ -35,6 +37,18 @@ AuthenticationResult executeTokenRequest() throws IOException {
return createAuthenticationResultFromOauthHttpResponse(oauthHttpResponse);
}
+ AuthenticationResult executeTokenRequest(
+ URL tokenEndpoint,
+ SSLSocketFactory sslSocketFactory,
+ Map parameters) throws IOException {
+ LOG.debug("Sending token request to: {}", tokenEndpoint);
+ OAuthHttpRequest request = createOauthHttpRequest(
+ tokenEndpoint,
+ sslSocketFactory,
+ parameters);
+ return createAuthenticationResultFromOauthHttpResponse(request.send());
+ }
+
OAuthHttpRequest createOauthHttpRequest() throws MalformedURLException {
if (requestAuthority.tokenEndpointUrl() == null) {
@@ -97,6 +111,65 @@ OAuthHttpRequest createOauthHttpRequest() throws MalformedURLException {
return oauthHttpRequest;
}
+ OAuthHttpRequest createOauthHttpRequest(
+ URL tokenEndpoint,
+ SSLSocketFactory sslSocketFactory,
+ Map parameters) {
+ if (tokenEndpoint == null) {
+ throw new MsalClientException("The endpoint URI is not specified",
+ AuthenticationErrorCode.INVALID_ENDPOINT_URI);
+ }
+
+ OAuthHttpRequest request = new OAuthHttpRequest(
+ HttpMethod.POST,
+ tokenEndpoint,
+ msalRequest.headers().getReadonlyHeaderMap(),
+ msalRequest.requestContext(),
+ serviceBundle)
+ .sslSocketFactory(sslSocketFactory);
+
+ Map params = new HashMap<>(parameters);
+ mergeClaimsAndCapabilities(params);
+ request.setQuery(StringHelper.serializeQueryParameters(params));
+ return request;
+ }
+
+ private void mergeClaimsAndCapabilities(Map params) {
+ String claims = params.get("claims");
+ if (msalRequest.application() instanceof AbstractClientApplicationBase
+ && ((AbstractClientApplicationBase) msalRequest.application()).clientCapabilities() != null) {
+ claims = mergeClaims(
+ claims,
+ ((AbstractClientApplicationBase) msalRequest.application()).clientCapabilities());
+ } else if (msalRequest.application() instanceof ManagedIdentityApplication) {
+ List capabilities =
+ ((ManagedIdentityApplication) msalRequest.application()).getClientCapabilities();
+ if (capabilities != null && !capabilities.isEmpty()) {
+ claims = mergeClaims(
+ claims,
+ JsonHelper.formCapabilitiesJson(new HashSet<>(capabilities)));
+ }
+ }
+
+ ClaimsRequest requestClaims = msalRequest.requestContext().apiParameters().claims();
+ if (requestClaims != null) {
+ claims = mergeClaims(claims, requestClaims.formatAsJSONString());
+ }
+ if (!StringHelper.isBlank(claims)) {
+ params.put("claims", claims);
+ }
+ }
+
+ private static String mergeClaims(String first, String second) {
+ if (StringHelper.isBlank(first)) {
+ return second;
+ }
+ if (StringHelper.isBlank(second)) {
+ return first;
+ }
+ return JsonHelper.mergeJSONString(first, second);
+ }
+
private void addQueryParameters(OAuthHttpRequest oauthHttpRequest) {
Map queryParameters = StringHelper.parseQueryParameters(oauthHttpRequest.query);
String clientID = msalRequest.application().clientId();
@@ -240,6 +313,7 @@ private AuthenticationResult createAuthenticationResultFromOauthHttpResponse(Htt
refreshOn(response.getRefreshIn() > 0 ? currTimestampSec + response.getRefreshIn() : 0).
accountCacheEntity(accountCacheEntity).
scopes(response.getScope()).
+ tokenType(response.tokenType()).
metadata(AuthenticationResultMetadata.builder()
.tokenSource(TokenSource.IDENTITY_PROVIDER)
.refreshOn(response.getRefreshIn() > 0 ? currTimestampSec + response.getRefreshIn() : 0)
diff --git a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/TokenResponse.java b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/TokenResponse.java
index b314bb77..8a543081 100644
--- a/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/TokenResponse.java
+++ b/msal4j-sdk/src/main/java/com/microsoft/aad/msal4j/TokenResponse.java
@@ -16,11 +16,13 @@ class TokenResponse {
private String accessToken;
private String idToken;
private String refreshToken;
+ private String tokenType;
TokenResponse(Map jsonMap) {
this.accessToken = jsonMap.get("access_token");
this.idToken = jsonMap.get("id_token");
this.refreshToken = jsonMap.get("refresh_token");
+ this.tokenType = jsonMap.get("token_type");
this.scope = jsonMap.get("scope");
this.clientInfo = jsonMap.get("client_info");
this.expiresIn = StringHelper.isNullOrBlank(jsonMap.get("expires_in")) ? 0 : Long.parseLong(jsonMap.get("expires_in"));
@@ -73,4 +75,8 @@ public String idToken() {
public String refreshToken() {
return refreshToken;
}
+
+ String tokenType() {
+ return tokenType;
+ }
}
diff --git a/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/AcquireTokenByManagedIdentitySupplierMtlsTest.java b/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/AcquireTokenByManagedIdentitySupplierMtlsTest.java
new file mode 100644
index 00000000..86395a7e
--- /dev/null
+++ b/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/AcquireTokenByManagedIdentitySupplierMtlsTest.java
@@ -0,0 +1,123 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j;
+
+import org.junit.jupiter.api.Test;
+
+import java.util.Collections;
+
+import static org.junit.jupiter.api.Assertions.assertEquals;
+import static org.junit.jupiter.api.Assertions.assertSame;
+import static org.junit.jupiter.api.Assertions.assertThrows;
+import static org.mockito.ArgumentMatchers.any;
+import static org.mockito.Mockito.mock;
+import static org.mockito.Mockito.verify;
+import static org.mockito.Mockito.when;
+
+class AcquireTokenByManagedIdentitySupplierMtlsTest {
+
+ @Test
+ void providerRuntimeFailureIsNormalized() {
+ RuntimeException providerFailure =
+ new RuntimeException("extension-specific failure");
+ IManagedIdentityMtlsProvider provider = request -> {
+ throw providerFailure;
+ };
+
+ MsalClientException exception = assertThrows(
+ MsalClientException.class,
+ () -> AcquireTokenByManagedIdentitySupplier
+ .getMtlsProviderBinding(provider, request()));
+
+ assertEquals(MsalError.MANAGED_IDENTITY_MTLS_REQUEST_FAILED,
+ exception.errorCode());
+ assertSame(providerFailure, exception.getCause());
+ }
+
+ @Test
+ void providerMsalFailureIsPreserved() {
+ MsalClientException providerFailure = new MsalClientException(
+ "known failure",
+ MsalError.MANAGED_IDENTITY_MTLS_PROVIDER_UNAVAILABLE);
+ IManagedIdentityMtlsProvider provider = request -> {
+ throw providerFailure;
+ };
+
+ assertSame(providerFailure, assertThrows(
+ MsalClientException.class,
+ () -> AcquireTokenByManagedIdentitySupplier
+ .getMtlsProviderBinding(provider, request())));
+ }
+
+ @Test
+ void imdsCallbackUsesImdsRetryPolicyForHttp410() throws Exception {
+ DefaultHttpClient httpClient = mock(DefaultHttpClient.class);
+ HttpResponse gone = response(HttpStatus.HTTP_GONE, "updating");
+ HttpResponse success = response(HttpStatus.HTTP_OK, "{}");
+ when(httpClient.send(any(HttpRequest.class)))
+ .thenReturn(gone, success);
+
+ IMDSRetryPolicy.setRetryDelayMs(0);
+ try {
+ ManagedIdentityApplication application =
+ ManagedIdentityApplication
+ .builder(ManagedIdentityId.systemAssigned())
+ .httpClient(httpClient)
+ .build();
+ ManagedIdentityParameters parameters =
+ ManagedIdentityParameters
+ .builder("https://vault.azure.net")
+ .withMtlsProofOfPossession()
+ .build();
+ RequestContext context = new RequestContext(
+ application,
+ PublicApi.ACQUIRE_TOKEN_BY_SYSTEM_ASSIGNED_MANAGED_IDENTITY,
+ parameters);
+ HttpHelper imdsHelper =
+ new HttpHelper(httpClient, new IMDSRetryPolicy());
+ ServiceBundle serviceBundle = new ServiceBundle(
+ null,
+ new TelemetryManager(null, false),
+ imdsHelper);
+ IManagedIdentityMtlsHttpClient callback =
+ AcquireTokenByManagedIdentitySupplier
+ .createMtlsProviderHttpClient(
+ imdsHelper,
+ serviceBundle,
+ context);
+
+ ManagedIdentityMtlsHttpResponse result = callback.execute(
+ new ManagedIdentityMtlsHttpRequest(
+ "GET",
+ "http://169.254.169.254/metadata/identity/getplatformmetadata",
+ Collections.singletonMap("Metadata", "true"),
+ null));
+
+ assertEquals(HttpStatus.HTTP_OK, result.statusCode());
+ verify(httpClient,
+ org.mockito.Mockito.times(2)).send(any(HttpRequest.class));
+ } finally {
+ IMDSRetryPolicy.resetToDefaults();
+ }
+ }
+
+ private static ManagedIdentityMtlsRequest request() {
+ return new ManagedIdentityMtlsRequest(
+ null,
+ null,
+ "binding",
+ "correlation",
+ httpRequest -> new ManagedIdentityMtlsHttpResponse(
+ HttpStatus.HTTP_OK,
+ "{}",
+ Collections.emptyMap()));
+ }
+
+ private static HttpResponse response(int status, String body) {
+ HttpResponse response = new HttpResponse();
+ response.statusCode(status);
+ response.body(body);
+ return response;
+ }
+}
diff --git a/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/AuthenticationResultMtlsTest.java b/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/AuthenticationResultMtlsTest.java
new file mode 100644
index 00000000..8e538a04
--- /dev/null
+++ b/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/AuthenticationResultMtlsTest.java
@@ -0,0 +1,83 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j;
+
+import org.junit.jupiter.api.Test;
+
+import javax.net.ssl.SSLContext;
+import java.io.*;
+import java.security.cert.X509Certificate;
+
+import static org.junit.jupiter.api.Assertions.*;
+
+class AuthenticationResultMtlsTest {
+
+ @Test
+ void bindingContextIsProcessLocalAndNotSerialized() throws Exception {
+ IMtlsBindingContext context = new TestBindingContext();
+ AuthenticationResult result = AuthenticationResult.builder()
+ .accessToken("secret")
+ .expiresOn(System.currentTimeMillis() / 1000 + 3600)
+ .tokenType("mtls_pop")
+ .mtlsBindingContext(context)
+ .build();
+
+ byte[] serialized;
+ try (ByteArrayOutputStream bytes = new ByteArrayOutputStream();
+ ObjectOutputStream output = new ObjectOutputStream(bytes)) {
+ output.writeObject(result);
+ serialized = bytes.toByteArray();
+ }
+ AuthenticationResult restored;
+ try (ObjectInputStream input = new ObjectInputStream(
+ new ByteArrayInputStream(serialized))) {
+ restored = (AuthenticationResult) input.readObject();
+ }
+
+ assertEquals("mtls_pop", restored.tokenType());
+ assertNull(restored.mtlsBindingContext());
+ assertNull(restored.bindingCertificate());
+ }
+
+ @Test
+ void equalityIncludesTokenTypeButExcludesLiveBindingContext() {
+ AuthenticationResult bearer = result("Bearer", null);
+ AuthenticationResult first = result("mtls_pop", new TestBindingContext());
+ AuthenticationResult second =
+ first.withMtlsBindingContext(new TestBindingContext());
+
+ assertNotEquals(bearer, first);
+ assertEquals(first, second);
+ assertEquals(first.hashCode(), second.hashCode());
+ }
+
+ private static AuthenticationResult result(
+ String tokenType,
+ IMtlsBindingContext context) {
+ return AuthenticationResult.builder()
+ .accessToken("secret")
+ .expiresOn(123)
+ .tokenType(tokenType)
+ .isPopAuthorization(context == null ? null : Boolean.TRUE)
+ .mtlsBindingContext(context)
+ .build();
+ }
+
+ private static final class TestBindingContext implements IMtlsBindingContext {
+ @Override
+ public SSLContext sslContext() {
+ return null;
+ }
+
+ @Override
+ public X509Certificate bindingCertificate() {
+ return null;
+ }
+
+ @Override
+ public String keyId() {
+ return "key";
+ }
+ }
+}
diff --git a/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/DefaultHttpClientMtlsTest.java b/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/DefaultHttpClientMtlsTest.java
new file mode 100644
index 00000000..a15871f0
--- /dev/null
+++ b/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/DefaultHttpClientMtlsTest.java
@@ -0,0 +1,29 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j;
+
+import org.junit.jupiter.api.Test;
+
+import javax.net.ssl.HttpsURLConnection;
+import javax.net.ssl.SSLSocketFactory;
+import java.net.URL;
+
+import static org.junit.jupiter.api.Assertions.assertFalse;
+import static org.mockito.Mockito.mock;
+
+class DefaultHttpClientMtlsTest {
+
+ @Test
+ void requestSpecificMtlsDisablesRedirects() throws Exception {
+ DefaultHttpClient client =
+ new DefaultHttpClient(null, null, null, null);
+
+ HttpsURLConnection connection = (HttpsURLConnection)
+ client.openConnection(
+ new URL("https://localhost/token"),
+ mock(SSLSocketFactory.class));
+
+ assertFalse(connection.getInstanceFollowRedirects());
+ }
+}
diff --git a/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsBindingTest.java b/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsBindingTest.java
new file mode 100644
index 00000000..925b6c22
--- /dev/null
+++ b/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsBindingTest.java
@@ -0,0 +1,41 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j;
+
+import org.junit.jupiter.api.Test;
+
+import javax.net.ssl.SSLContext;
+import java.security.cert.X509Certificate;
+
+import static org.junit.jupiter.api.Assertions.assertThrows;
+
+class ManagedIdentityMtlsBindingTest {
+
+ @Test
+ void tokenEndpointMustUseHttps() {
+ IMtlsBindingContext context = new IMtlsBindingContext() {
+ @Override
+ public SSLContext sslContext() {
+ return null;
+ }
+
+ @Override
+ public X509Certificate bindingCertificate() {
+ return null;
+ }
+
+ @Override
+ public String keyId() {
+ return "key";
+ }
+ };
+
+ assertThrows(IllegalArgumentException.class,
+ () -> new ManagedIdentityMtlsBinding(
+ context, "client", "http://login.example/token"));
+ assertThrows(IllegalArgumentException.class,
+ () -> new ManagedIdentityMtlsBinding(
+ context, "client", "not-a-url"));
+ }
+}
diff --git a/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsParametersTest.java b/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsParametersTest.java
new file mode 100644
index 00000000..dfe67dea
--- /dev/null
+++ b/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsParametersTest.java
@@ -0,0 +1,105 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j;
+
+import org.junit.jupiter.api.Test;
+
+import static org.junit.jupiter.api.Assertions.*;
+
+class ManagedIdentityMtlsParametersTest {
+
+ @Test
+ void attestationRequiresMtlsButMtlsCanBeRequestedAlone() {
+ assertDoesNotThrow(
+ () -> ManagedIdentityParameters.builder("https://vault.azure.net")
+ .withMtlsProofOfPossession()
+ .build());
+ assertThrows(IllegalArgumentException.class,
+ () -> ManagedIdentityParameters.builder("https://vault.azure.net")
+ .withAttestationSupport()
+ .build());
+ }
+
+ @Test
+ void bearerAndMtlsCachePartitionsCannotCollide() {
+ ManagedIdentityParameters bearer =
+ ManagedIdentityParameters.builder("https://vault.azure.net").build();
+ ManagedIdentityParameters mtls = ManagedIdentityParameters
+ .builder("https://vault.azure.net")
+ .withMtlsProofOfPossession()
+ .withAttestationSupport()
+ .build();
+
+ assertEquals("", bearer.computeExtCacheKeyHash());
+ assertFalse(mtls.computeMtlsExtCacheKeyHash("certificate-a").isEmpty());
+ assertNotEquals(bearer.computeExtCacheKeyHash(),
+ mtls.computeMtlsExtCacheKeyHash("certificate-a"));
+ }
+
+ @Test
+ void renewedCertificateCreatesNewCachePartition() {
+ ManagedIdentityParameters parameters = ManagedIdentityParameters
+ .builder("https://vault.azure.net")
+ .withMtlsProofOfPossession()
+ .withAttestationSupport()
+ .build();
+ String first = parameters.computeMtlsExtCacheKeyHash("certificate-a");
+
+ assertNotEquals(first,
+ parameters.computeMtlsExtCacheKeyHash("certificate-b"));
+ }
+
+ @Test
+ void attestationModeCreatesDistinctCachePartition() {
+ ManagedIdentityParameters unattested = ManagedIdentityParameters
+ .builder("https://vault.azure.net")
+ .withMtlsProofOfPossession()
+ .build();
+ ManagedIdentityParameters attested = ManagedIdentityParameters
+ .builder("https://vault.azure.net")
+ .withMtlsProofOfPossession()
+ .withAttestationSupport()
+ .build();
+
+ assertNotEquals(
+ unattested.computeMtlsExtCacheKeyHash("certificate-a"),
+ attested.computeMtlsExtCacheKeyHash("certificate-a"));
+ }
+
+ @Test
+ void buildersDoNotExposeClaimsOrTokenMaterial() {
+ ManagedIdentityParameters.ManagedIdentityParametersBuilder builder =
+ ManagedIdentityParameters.builder("https://vault.azure.net")
+ .claims("{\"access_token\":\"secret\"}")
+ .withMtlsProofOfPossession()
+ .withAttestationSupport();
+
+ assertFalse(builder.toString().contains("secret"));
+ }
+
+ @Test
+ void tokenEndpointMustExplicitlyReturnMtlsPop() {
+ AuthenticationResult bearer = AuthenticationResult.builder()
+ .accessToken("token")
+ .tokenType("Bearer")
+ .build();
+ assertThrows(MsalServiceException.class,
+ () -> AcquireTokenByManagedIdentitySupplier
+ .validateMtlsTokenResponse(bearer));
+
+ AuthenticationResult missingTokenType = AuthenticationResult.builder()
+ .accessToken("token")
+ .build();
+ assertThrows(MsalServiceException.class,
+ () -> AcquireTokenByManagedIdentitySupplier
+ .validateMtlsTokenResponse(missingTokenType));
+
+ AuthenticationResult mtlsPop = AuthenticationResult.builder()
+ .accessToken("token")
+ .tokenType("mtls_pop")
+ .build();
+ assertDoesNotThrow(() -> AcquireTokenByManagedIdentitySupplier
+ .validateMtlsTokenResponse(mtlsPop));
+ }
+}
diff --git a/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsProviderLoaderTest.java b/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsProviderLoaderTest.java
new file mode 100644
index 00000000..355a28f6
--- /dev/null
+++ b/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/ManagedIdentityMtlsProviderLoaderTest.java
@@ -0,0 +1,21 @@
+// Copyright (c) Microsoft Corporation. All rights reserved.
+// Licensed under the MIT License.
+
+package com.microsoft.aad.msal4j;
+
+import org.junit.jupiter.api.Test;
+
+import static org.junit.jupiter.api.Assertions.*;
+
+class ManagedIdentityMtlsProviderLoaderTest {
+
+ @Test
+ void missingOptionalProviderFailsClosed() {
+ MsalClientException exception = assertThrows(
+ MsalClientException.class,
+ ManagedIdentityMtlsProviderLoader::load);
+
+ assertEquals(MsalError.MANAGED_IDENTITY_MTLS_PROVIDER_UNAVAILABLE,
+ exception.errorCode());
+ }
+}
diff --git a/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/TokenRequestExecutorTest.java b/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/TokenRequestExecutorTest.java
index ac743755..2137446e 100644
--- a/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/TokenRequestExecutorTest.java
+++ b/msal4j-sdk/src/test/java/com/microsoft/aad/msal4j/TokenRequestExecutorTest.java
@@ -26,12 +26,54 @@
import java.util.Base64;
import java.util.Collections;
import java.util.HashMap;
+import java.util.Map;
import java.util.concurrent.ExecutionException;
+import javax.net.ssl.SSLSocketFactory;
@ExtendWith(MockitoExtension.class)
@TestInstance(TestInstance.Lifecycle.PER_CLASS)
class TokenRequestExecutorTest {
+ @Test
+ void managedIdentityMtlsRequestUsesNormalClaimsPipelineAndSocketFactory()
+ throws Exception {
+ ManagedIdentityApplication app = ManagedIdentityApplication
+ .builder(ManagedIdentityId.systemAssigned())
+ .clientCapabilities(Collections.singletonList("cp1"))
+ .build();
+ ManagedIdentityParameters parameters = ManagedIdentityParameters
+ .builder("https://vault.azure.net")
+ .claims("{\"access_token\":{\"custom\":{\"essential\":true}}}")
+ .withMtlsProofOfPossession()
+ .build();
+ ManagedIdentityRequest managedIdentityRequest =
+ new ManagedIdentityRequest(
+ app,
+ new RequestContext(
+ app,
+ PublicApi.ACQUIRE_TOKEN_BY_SYSTEM_ASSIGNED_MANAGED_IDENTITY,
+ parameters));
+ TokenRequestExecutor executor = new TokenRequestExecutor(
+ new AADAuthority(new URL(TestConstants.ORGANIZATIONS_AUTHORITY)),
+ managedIdentityRequest,
+ app.serviceBundle());
+ SSLSocketFactory socketFactory = mock(SSLSocketFactory.class);
+ Map body = new HashMap<>();
+ body.put("grant_type", "client_credentials");
+ body.put("token_type", "mtls_pop");
+
+ OAuthHttpRequest request = executor.createOauthHttpRequest(
+ new URL("https://login.example/tenant/oauth2/v2.0/token"),
+ socketFactory,
+ body);
+ Map query = StringHelper.parseQueryParameters(request.query);
+
+ assertEquals(socketFactory, request.sslSocketFactory());
+ assertEquals("mtls_pop", query.get("token_type"));
+ assertTrue(query.get("claims").contains("\"xms_cc\""));
+ assertTrue(query.get("claims").contains("\"custom\""));
+ }
+
@Test
void executeOAuthRequest_SCBadRequestErrorInvalidGrant_InteractionRequiredException()
throws MsalException,
diff --git a/pom.xml b/pom.xml
index 373e77e7..d457aed3 100644
--- a/pom.xml
+++ b/pom.xml
@@ -7,7 +7,16 @@
pom
msal4j-sdk
+ msal4j-mtls-extensions
msal4j-brokers
msal4j-persistence-extension
+
+
+ e2e
+
+ msal4j-mtls-extensions-e2e
+
+
+
diff --git a/run-java-msi-v2-mtls-devapp.ps1 b/run-java-msi-v2-mtls-devapp.ps1
new file mode 100644
index 00000000..d620e1dd
--- /dev/null
+++ b/run-java-msi-v2-mtls-devapp.ps1
@@ -0,0 +1,88 @@
+[CmdletBinding()]
+param(
+ [string]$Maven,
+ [switch]$SkipBuild
+)
+
+$ErrorActionPreference = "Stop"
+$repoRoot = $PSScriptRoot
+
+function Assert-EnvironmentVariable {
+ param([string]$Name)
+ if ([string]::IsNullOrWhiteSpace([Environment]::GetEnvironmentVariable($Name))) {
+ throw "Required environment variable is missing: $Name"
+ }
+}
+
+if (-not $IsWindows -and $PSVersionTable.PSEdition -eq "Core") {
+ throw "Managed Identity v2 KeyGuard mTLS PoP requires Windows."
+}
+
+$tokenOnly = [bool]::Parse(
+ $(if ([string]::IsNullOrWhiteSpace($env:MSAL_JAVA_MTLS_TOKEN_ONLY)) {
+ "false"
+ } else {
+ $env:MSAL_JAVA_MTLS_TOKEN_ONLY
+ }))
+if (-not $tokenOnly) {
+ Assert-EnvironmentVariable "MSAL_JAVA_MTLS_AKV_URL"
+ Assert-EnvironmentVariable "MSAL_JAVA_MTLS_AKV_SECRET_NAME"
+}
+
+$java = Get-Command java -ErrorAction SilentlyContinue
+if ($null -eq $java) {
+ throw "java.exe was not found on PATH."
+}
+
+$tpm = Get-Tpm
+if (-not $tpm.TpmPresent -or -not $tpm.TpmReady) {
+ throw "A present and ready TPM is required."
+}
+
+try {
+ if (-not (Confirm-SecureBootUEFI)) {
+ throw "Secure Boot is not enabled."
+ }
+} catch [System.PlatformNotSupportedException] {
+ throw "Secure Boot status could not be verified on this platform."
+}
+
+if (-not $SkipBuild) {
+ if ([string]::IsNullOrWhiteSpace($Maven)) {
+ $mavenCommand = Get-Command mvn.cmd, mvn -ErrorAction SilentlyContinue |
+ Select-Object -First 1
+ if ($null -ne $mavenCommand) {
+ $Maven = $mavenCommand.Source
+ }
+ }
+ if ([string]::IsNullOrWhiteSpace($Maven) -or
+ -not (Test-Path -LiteralPath $Maven -PathType Leaf)) {
+ throw "Maven was not found on PATH. Pass -Maven with the path to mvn.cmd."
+ }
+ Push-Location $repoRoot
+ try {
+ & $Maven -q -pl msal4j-mtls-extensions-e2e -am `
+ '-Pe2e' '-DskipTests' '-Dmaven.javadoc.skip=true' package
+ if ($LASTEXITCODE -ne 0) {
+ throw "Maven build failed with exit code $LASTEXITCODE."
+ }
+ } finally {
+ Pop-Location
+ }
+}
+
+$jar = Get-ChildItem `
+ (Join-Path $repoRoot "msal4j-mtls-extensions-e2e\target") `
+ -Filter "*-e2e.jar" |
+ Sort-Object LastWriteTime -Descending |
+ Select-Object -First 1
+if ($null -eq $jar) {
+ throw "The manual validation JAR was not found. Run without -SkipBuild."
+}
+
+Write-Host "Attestation DLL: bundled Microsoft.Azure.Security.KeyGuardAttestation 1.1.5"
+Write-Host "Validation JAR: $($jar.FullName)"
+& $java.Source -jar $jar.FullName
+if ($LASTEXITCODE -ne 0) {
+ throw "Managed Identity v2 mTLS PoP validation failed with exit code $LASTEXITCODE."
+}