From 0fdde1022c29fc0c14aade8b4ca81187a2133938 Mon Sep 17 00:00:00 2001 From: James Xian Date: Wed, 23 Sep 2026 10:29:16 -0700 Subject: [PATCH 1/3] Cache and proactively refresh provider credentials Add worker-local Key Vault, managed identity, and Entra token caches with shared in-flight acquisition, proactive refresh, and fail-closed expiry across JavaScript, Python, and .NET. Include startup/shutdown integration, offline tests, and documentation. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- README.md | 7 + docs/CONTRACT.md | 87 +++++- docs/ONBOARDING.md | 11 + dotnet/Program.cs | 1 + dotnet/README.md | 11 +- dotnet/Src/CredentialRefreshService.cs | 14 + dotnet/Src/DispatchEngine.cs | 84 ++---- dotnet/Src/ISecretResolver.cs | 2 +- dotnet/Src/ProviderCredentials.cs | 139 +++++++++ dotnet/Src/RefreshingCache.cs | 143 +++++++++ dotnet/Src/SecretResolver.cs | 54 ++-- dotnet/tests/CredentialCacheTests.cs | 306 ++++++++++++++++++++ dotnet/tests/EngineTests.cs | 40 ++- javascript/README.md | 8 + javascript/src/functions/SendOtp.js | 5 + javascript/src/functions/credentials.js | 181 ++++++++++++ javascript/src/functions/dispatch.js | 112 ++----- javascript/src/functions/refreshingCache.js | 122 ++++++++ javascript/test/credential-cache.test.js | 252 ++++++++++++++++ javascript/test/credential-sdk.test.js | 108 +++++++ javascript/test/dispatch.test.js | 9 +- javascript/test/sendotp.test.js | 33 ++- python/README.md | 11 +- python/function_app.py | 8 + python/src/credentials.py | 128 ++++++++ python/src/dispatch.py | 90 ++---- python/src/refreshing_cache.py | 136 +++++++++ python/src/secrets.py | 46 ++- python/tests/test_credential_cache.py | 294 +++++++++++++++++++ python/tests/test_credential_sdk.py | 65 +++++ python/tests/test_engine.py | 13 +- python/tests/test_function_app.py | 4 +- 32 files changed, 2223 insertions(+), 301 deletions(-) create mode 100644 dotnet/Src/CredentialRefreshService.cs create mode 100644 dotnet/Src/ProviderCredentials.cs create mode 100644 dotnet/Src/RefreshingCache.cs create mode 100644 dotnet/tests/CredentialCacheTests.cs create mode 100644 javascript/src/functions/credentials.js create mode 100644 javascript/src/functions/refreshingCache.js create mode 100644 javascript/test/credential-cache.test.js create mode 100644 javascript/test/credential-sdk.test.js create mode 100644 python/src/credentials.py create mode 100644 python/src/refreshing_cache.py create mode 100644 python/tests/test_credential_cache.py create mode 100644 python/tests/test_credential_sdk.py diff --git a/README.md b/README.md index de359b1..6736c0c 100644 --- a/README.md +++ b/README.md @@ -178,6 +178,13 @@ login; ordinary local machines have no managed-identity endpoint. Use offline te evaluation locally, or an explicitly injected test resolver for integration work. Never commit local settings, keys or test credentials. +Configured providers are [prepared automatically per worker](docs/CONTRACT.md#credential-caching-and-refresh): +Telesign's Key Vault credentials, Soprano's managed-identity assertion, and its final Entra access +token are cached and refreshed before expiry. Refresh never sends an OTP. Evaluation handling still +skips provider work, but a worker with a configured provider can independently acquire credentials +at startup or during background refresh. Leave `EPP_PROVIDER_NAME` unset for local evaluation-only +work without credential acquisition. No extra refresh app settings are required. + Core Tools does not resolve Azure Key Vault reference expressions locally. Supply the local test PEM or base64 PEM directly; use a reference such as `@Microsoft.KeyVault(SecretUri=https://.vault.azure.net/secrets//)` for `EPP_DECRYPTION_KEY_PEM` in Azure app settings, where the platform resolves it. diff --git a/docs/CONTRACT.md b/docs/CONTRACT.md index 271b202..58f572f 100644 --- a/docs/CONTRACT.md +++ b/docs/CONTRACT.md @@ -113,14 +113,20 @@ there is no API-key fallback. Evaluation skips acquisition. A provider rejection Tokens are treated as opaque: the Function checks SDK expiry metadata, not custom JWT claims. Soprano remains responsible for signature, issuer, audience, expiry, permissions, and account validation. -Credential instances are reused for the configured tenant/application/identity; each acquisition -uses the selected scope. JavaScript and .NET pass one 2.5-second cancellation signal/token through -both exchange stages. Python uses 2.5-second connect/read inactivity timeouts, not a total deadline. -Configured SDK transport retries are disabled. Managed-identity discovery may involve additional -SDK operations; this is not an end-to-end delivery deadline. JavaScript suppresses SDK logs only -in the acquisition's asynchronous context. Python filters Azure Identity/Core/MSAL records on -configured handlers in that context; configure logging sinks before handling requests. .NET disables -credential diagnostics. Keep platform body tracing off and never log credential objects or tokens. +Credential instances are reused for the configured tenant/application/identity. Separate +[worker-local caches](#credential-caching-and-refresh) hold the managed-identity assertion and the +final provider token for each selected scope. A usable final token avoids both SDK acquisition calls +on the delivery path. Each cache owns its refresh independently, so one waiting caller cannot cancel +an assertion refresh another caller needs. JavaScript and .NET bound each acquisition to 2.5 seconds; +JavaScript also links that cancellation to the actual Azure SDK HTTP pipeline rather than relying +only on the SDK's `getToken` option. Python bounds each wait to 2.5 seconds and uses 2.5-second +connect/read inactivity timeouts; its shared refresh may finish after a waiter leaves. + +Credential SDK transport retries are disabled; failed refreshes use the bounded backoff described +below. These are not end-to-end delivery deadlines. JavaScript suppresses SDK logs in the +acquisition's asynchronous context. Python filters Azure Identity/Core/MSAL records on configured +handlers in that context; configure logging sinks before handling requests. .NET disables credential +diagnostics. Keep platform body tracing off and never log credential objects or tokens. When migrating from the earlier optional-JWT branch, replace `EPP_PROVIDER_APPLICATION_ID` with `EPP_OUTBOUND_CLIENT_ID` and `EPP_PROVIDER_MI_CLIENT_ID` with `EPP_OUTBOUND_MI_CLIENT_ID`. @@ -179,6 +185,11 @@ Key Vault reads and outbound provider HTTP are skipped. No provider name, endpoi are needed. Platform authentication and resolution of the decryption-key reference may still require network access. Core Tools has no Easy Auth; local evaluation must remain loopback-only, without tunnels. +This describes the evaluation **request path**. Independently, workers with a configured provider +automatically prewarm and refresh credentials, even if their current traffic is evaluation-only. +No background task dispatches an OTP. A worker without `EPP_PROVIDER_NAME` performs no credential +prewarming, and evaluation does not require that prewarming succeed. + There is no diagnostic environment flag. A live request is not an evaluation request. Adapter-specific wire fields, where required by an API, remain internal and cannot enable a separate non-delivery mode. @@ -321,6 +332,58 @@ the configured adapter, which builds the provider's SMS or voice API call. Purch provider does not install an adapter: add and register that provider's adapter first. Purchase, subscription activation and changing tenant policy belong to provisioning, not this Function. +### Credential caching and refresh + +Provider credential management is automatic for configured providers in every runtime. It changes +when credentials are fetched, not the HTTP/nonce contract, caller authentication, FIC, provider +selection or provider request format. The decryption-key Key Vault reference remains separate and +is still resolved by the platform; this cache does not rotate or replace JWE keys. + +| Cache | Refresh target | Hard usability boundary | +|---|---|---| +| Key Vault API-key bundle | Four minutes after successful retrieval. Required key and customer-ID secrets are fetched concurrently and published together only when both reads succeed and values are nonblank. | Five minutes after retrieval; failed refreshes never extend the old bundle's lifetime. | +| Managed-identity assertion | Five minutes before SDK expiry, or an earlier SDK refresh hint when available. | SDK expiry minus 30 seconds. | +| Final Entra provider token | Five minutes before SDK expiry, or an earlier SDK refresh hint when available, separately for each scope. | SDK expiry minus 30 seconds. | + +These are in-memory caches, one per worker process, not distributed caches or persisted token +stores. API-key entries are isolated by vault, managed identity and manifest secret names. OAuth +state is isolated by provider tenant, application and managed identity, with separate final-token +entries for different scopes. A change of credential configuration stops the old entries; deploy +app-setting changes normally with a worker restart rather than mutating process environment in place. + +Concurrent cache misses share **one in-progress acquisition per entry**. A request with a still-usable +cached credential returns it immediately while a due refresh proceeds separately. Refresh failures +retain that credential only until its original hard expiry. Once expired, callers join the shared +refresh or fail closed with the existing sanitized credential error; no stale-success fallback is +introduced. This does not deduplicate provider deliveries or change caller/provider retry behavior. + +An SDK refresh can return the same token from its own cache. The manager preserves the token's +original expiry instead of treating it as a new token. If the suggested refresh time is already past, +the next check is delayed by half the remaining usable lifetime, bounded to 1-60 seconds, preventing +an immediate refresh loop. Failed refreshes back off exponentially from 5 seconds to a 60-second +base, plus up to 20% jitter. Requests do not bypass that backoff and repeatedly hit a failing +dependency. Successful refresh resets it. + +JavaScript's app-start hook, Python's worker-module initialization thread, and .NET's hosted service +start credential preparation. Each entry then owns its refresh timer; a single distributed timer +trigger would not populate every worker's memory. JavaScript timers are unreferenced and Python +threads are daemon threads; termination hooks/`atexit`/hosted-service shutdown cancel scheduled +work and discard entries. Late completion cannot repopulate a stopped cache. A cold Python caller +can stop waiting without cancelling the shared retrieval. Failed or absent provider configuration +does not prevent evaluation from working. + +**This is not a guarantee that the first request after a cold start meets the caller's budget.** +Initialization can itself be on that first request's critical path, and a worker may receive traffic +before preparation finishes. Existing Always On/minimum-instance settings can help, but readiness, +scale-out and caller-observed latency must be measured in the deployment. A warm final token avoids +the Entra exchange; warming only the managed-identity assertion would not achieve that. + +Background failures emit a compact `credential_refresh_failed` service event with `cacheKind` +(`key_vault`, `managed_identity`, `provider_token`, `configuration`, or `initialization`) and the +fixed reason `credential_unavailable`. They do not carry a request's tracing IDs or secret values. +The existing per-request `providerCredentialElapsedMs` still measures the resolution observed by +that caller. No per-operation MI/Entra timing diagnostics are added. + --- ## 5. Required behaviors @@ -477,10 +540,13 @@ Each language keeps lightweight offline tests covering representative applicatio - Bundled adapter request formats and static provider credentials. - Fail-closed outcomes, missing credentials, HTTPS guards and timeouts. - Envelope validation and real JWE decryption/tamper rejection. -- Evaluation without provider I/O. +- Evaluation handling without provider I/O; configured-provider startup refresh is tested separately. - Awaited delivery, nonce acknowledgement and privacy-safe logging, including the shared service-event order and summary field set in [contract.json](../tests/fixtures/contract.json), identifier provenance, error paths, provider-body timeouts and concurrent request isolation. +- Single-flight credential retrieval, automatic refresh, stale-value expiry, token lifetime + preservation, failure backoff, configuration isolation and cleanup using controlled clocks and + fake dependencies. JavaScript tests also exercise cancellation through the actual SDK pipeline. The sample deliberately omits exhaustive input permutations and SDK internals. These tests use local keys and mocked external services; they do not send SMS and **do not test Easy Auth or platform @@ -501,6 +567,9 @@ not prove handset delivery or support for every provider feature. timing architecture merely because the setup script deploys it. - The outbound timeout is not an end-to-end deadline. Cold starts, platform authentication and Key Vault access can exceed the caller's budget; Python uses connect/read inactivity timeouts. +- Credential caches are process-local and proactive preparation is best-effort, not an Azure + traffic-readiness guarantee. Provider delivery still waits for acceptance; background refresh + is not background OTP delivery, a queue, or protection against duplicate sends. - Voice text is forwarded unchanged. Digit-by-digit rendering required by the setup guide must be verified for the chosen voice integration; unspaced numeric text is not guaranteed to be spoken correctly. - Full body-size/content-type and E.164 validation, subscription provisioning, certification, diff --git a/docs/ONBOARDING.md b/docs/ONBOARDING.md index 4cfe1b7..8462181 100644 --- a/docs/ONBOARDING.md +++ b/docs/ONBOARDING.md @@ -175,3 +175,14 @@ Use [CONTRACT.md](CONTRACT.md) for the full request contract and production limi caller received the response. Credential resolution can use caches. `providerEndpoint` contains the base URL and API path only, without query strings, userinfo or fragments; support IDs remain raw so customers can share the exact reference with Microsoft/provider support. + + Provider credentials are [prewarmed and refreshed per worker](CONTRACT.md#credential-caching-and-refresh) + automatically when a provider is configured. This contacts Key Vault or Entra without sending an OTP. + Evaluation requests still skip those dependencies, but independent background preparation may run + alongside them. For local evaluation-only use without managed identity, leave `EPP_PROVIDER_NAME` + unset. Check for `credential_refresh_failed` warnings before live testing. + + Compare fresh-worker, warm, expiry/rotation and concurrent-request behavior. A warmup or passing + offline test does not prove the first live request fits the caller's timeout. Background refresh + does not retry or deduplicate a provider send. Do not rerun provisioning or change FIC, app + registration, provider settings or decryption keys to deploy this code-only improvement. diff --git a/dotnet/Program.cs b/dotnet/Program.cs index 42cdd60..046a7ad 100644 --- a/dotnet/Program.cs +++ b/dotnet/Program.cs @@ -26,5 +26,6 @@ builder.Services.AddSingleton(); builder.Services.AddSingleton(); +builder.Services.AddHostedService(); builder.Build().Run(); diff --git a/dotnet/README.md b/dotnet/README.md index 05ee4d4..c3de947 100644 --- a/dotnet/README.md +++ b/dotnet/README.md @@ -90,16 +90,25 @@ six-digit numeric run that is not part of a longer number and repeats the comple ## Source +The hosted credential-refresh service automatically prewarms a configured provider. Separate +process-local caches refresh Key Vault bundles, managed-identity assertions and final Entra tokens. +Each acquisition owns its cancellation budget; cancelling a waiter does not cancel another +request's shared retrieval. Shutdown stops timers and drops values. See the +[refresh contract](../docs/CONTRACT.md#credential-caching-and-refresh) for expiry/backoff semantics +and cold-start limitations. Evaluation remains independent from successful credential preparation. + | Source | Purpose | |---|---| | [Program.cs](Program.cs) | Host and adapter registration | | [Functions/SendOtp.cs](Functions/SendOtp.cs) | HTTP handler | | [Src/AppConfig.cs](Src/AppConfig.cs) | Shared deployment settings | | [Src/DispatchEngine.cs](Src/DispatchEngine.cs) | Envelope/JWE handling and dispatch | +| [Src/ProviderCredentials.cs](Src/ProviderCredentials.cs), [RefreshingCache.cs](Src/RefreshingCache.cs) | Single-flight credential bundles and independent token refresh | +| [Src/CredentialRefreshService.cs](Src/CredentialRefreshService.cs) | Per-worker startup and shutdown integration | | [Src/RequestLog.cs](Src/RequestLog.cs) | Request-scoped [service events and summaries](../docs/CONTRACT.md#application-logs) with explicit ID sources | | [Src/ProviderRegistry.cs](Src/ProviderRegistry.cs), [Src/IProviderAdapter.cs](Src/IProviderAdapter.cs) | Adapter lookup and contract | | [Src/Providers/](Src/Providers/) | Adapter manifests and API-specific implementations | -| [Src/SecretResolver.cs](Src/SecretResolver.cs) | Cached Key Vault access via managed identity | +| [Src/SecretResolver.cs](Src/SecretResolver.cs) | Key Vault transport; `ISecretResolver.ResolveAsync` accepts an optional cancellation token and bundle caching belongs to the credential manager | | [Src/OutcomeMapper.cs](Src/OutcomeMapper.cs), [Src/Models.cs](Src/Models.cs) | Outcomes and shared records | Implement `IProviderAdapter` and register it in [Program.cs](Program.cs) without adding provider-specific diff --git a/dotnet/Src/CredentialRefreshService.cs b/dotnet/Src/CredentialRefreshService.cs new file mode 100644 index 0000000..6d49ee7 --- /dev/null +++ b/dotnet/Src/CredentialRefreshService.cs @@ -0,0 +1,14 @@ +using Microsoft.Extensions.Hosting; + +namespace Epp.Otp; + +internal sealed class CredentialRefreshService(DispatchEngine engine) : IHostedService +{ + public Task StartAsync(CancellationToken cancellationToken) => engine.StartCredentialRefreshAsync(cancellationToken); + + public Task StopAsync(CancellationToken cancellationToken) + { + engine.Dispose(); + return Task.CompletedTask; + } +} diff --git a/dotnet/Src/DispatchEngine.cs b/dotnet/Src/DispatchEngine.cs index f07c1bf..e180dab 100644 --- a/dotnet/Src/DispatchEngine.cs +++ b/dotnet/Src/DispatchEngine.cs @@ -210,36 +210,31 @@ private static string NormalizePem(string value) => : Encoding.UTF8.GetString(Convert.FromBase64String(value.Trim())); } -public sealed class DispatchEngine +public sealed class DispatchEngine : IDisposable { public const string ProviderHttpClientName = "otp-provider"; private const int DefaultTimeoutMs = 1500; private const int MaxTimeoutMs = 2500; private readonly ProviderRegistry _registry; - private readonly ISecretResolver _secrets; private readonly IHttpClientFactory _httpFactory; private readonly IEnv _env; - private readonly object _oauthLock = new(); - private TokenCredential? _oauthCredential; - private string? _oauthCredentialConfig; - private readonly Func _createManagedIdentity; - private readonly Func>, TokenCredential> _createOAuthCredential; + private readonly ProviderCredentials _credentials; - public DispatchEngine(ProviderRegistry registry, ISecretResolver secrets, IHttpClientFactory httpFactory, IEnv? env = null) + public DispatchEngine(ProviderRegistry registry, ISecretResolver secrets, IHttpClientFactory httpFactory, + IEnv? env = null, ILogger? log = null) : this(registry, secrets, httpFactory, env, identity => new ManagedIdentityCredential(identity, OAuthOptions()), - (tenant, application, assertion) => new ClientAssertionCredential(tenant, application, assertion, OAuthOptions())) { } + (tenant, application, assertion) => new ClientAssertionCredential(tenant, application, assertion, OAuthOptions()), log) { } internal DispatchEngine(ProviderRegistry registry, ISecretResolver secrets, IHttpClientFactory httpFactory, IEnv? env, Func createManagedIdentity, - Func>, TokenCredential> createOAuthCredential) + Func>, TokenCredential> createOAuthCredential, + ILogger? log = null, TimeProvider? clock = null) { _registry = registry; - _secrets = secrets; _httpFactory = httpFactory; _env = env ?? new ProcessEnv(); - _createManagedIdentity = createManagedIdentity; - _createOAuthCredential = createOAuthCredential; + _credentials = new ProviderCredentials(secrets, _env, createManagedIdentity, createOAuthCredential, log, clock); } private static ClientAssertionCredentialOptions OAuthOptions() @@ -252,8 +247,21 @@ private static ClientAssertionCredentialOptions OAuthOptions() return options; } - private static bool UsableAccessToken(AccessToken token) => - !string.IsNullOrWhiteSpace(token.Token) && token.ExpiresOn > DateTimeOffset.UtcNow.AddSeconds(30); + public async Task StartCredentialRefreshAsync(CancellationToken cancellation = default) + { + var config = AppConfig.Read(_env); + if (string.IsNullOrWhiteSpace(config.ProviderName)) return; + var adapter = _registry.Get(config.ProviderName); + if (adapter is null || (!string.IsNullOrEmpty(config.ProviderAuthMode) && config.ProviderAuthMode != adapter.Manifest.Auth.Mode)) + { + _credentials.ReportFailure("configuration"); + return; + } + try { await _credentials.ResolveAsync(adapter.Manifest.Auth, config, cancellation).ConfigureAwait(false); } + catch (Exception) { _credentials.ReportFailure("initialization"); } + } + + public void Dispose() => _credentials.Dispose(); public async Task DispatchAsync(DispatchRequest dispatch, string requestId, RequestLog? log = null) { @@ -377,50 +385,8 @@ DispatchResult Failure(int status, string stage, string reason, object body) } } - private async Task ResolveCredentialAsync(AuthConfig auth, AppConfig config) - { - if (auth.Mode == "apiKey") - { - var secret = await _secrets.ResolveAsync(auth.KeyVaultSecretName); - var identity = string.IsNullOrEmpty(auth.IdentityKeyVaultSecretName) ? string.Empty : await _secrets.ResolveAsync(auth.IdentityKeyVaultSecretName); - return new ProviderCredential("apiKey", Secret: secret, Identity: identity); - } - if (auth.Mode != "oauth" || string.IsNullOrEmpty(config.ProviderTenantId) - || string.IsNullOrEmpty(config.ProviderScope) || string.IsNullOrEmpty(config.OutboundClientId) - || string.IsNullOrEmpty(config.OutboundManagedIdentityClientId)) - throw new InvalidOperationException("unsupported or incomplete provider authentication"); - - var credentialConfig = string.Join("|", config.ProviderTenantId, config.OutboundClientId, config.OutboundManagedIdentityClientId); - TokenCredential providerCredential; - lock (_oauthLock) - { - if (_oauthCredential is null || _oauthCredentialConfig != credentialConfig) - { - var managedIdentity = _createManagedIdentity(config.OutboundManagedIdentityClientId); - _oauthCredential = _createOAuthCredential( - config.ProviderTenantId, - config.OutboundClientId, - async cancellationToken => - { - var assertion = await managedIdentity.GetTokenAsync( - new TokenRequestContext(new[] { "api://AzureADTokenExchange/.default" }), - cancellationToken); - if (!UsableAccessToken(assertion)) - throw new InvalidOperationException("managed identity assertion unavailable"); - return assertion.Token; - }); - _oauthCredentialConfig = credentialConfig; - } - providerCredential = _oauthCredential; - } - using var cancellation = new CancellationTokenSource(TimeSpan.FromSeconds(2.5)); - var token = await providerCredential.GetTokenAsync( - new TokenRequestContext(new[] { config.ProviderScope }), - cancellation.Token); - if (!UsableAccessToken(token)) - throw new InvalidOperationException("provider OAuth token unavailable"); - return new ProviderCredential("oauth", AccessToken: token.Token); - } + private Task ResolveCredentialAsync(AuthConfig auth, AppConfig config) => + _credentials.ResolveAsync(auth, config); internal static int NormalizeProviderTimeoutMs(string? value) { diff --git a/dotnet/Src/ISecretResolver.cs b/dotnet/Src/ISecretResolver.cs index 5b75ad9..5ae496b 100644 --- a/dotnet/Src/ISecretResolver.cs +++ b/dotnet/Src/ISecretResolver.cs @@ -2,5 +2,5 @@ namespace Epp.Otp; public interface ISecretResolver { - Task ResolveAsync(string? secretName); + Task ResolveAsync(string? secretName, CancellationToken cancellationToken = default); } diff --git a/dotnet/Src/ProviderCredentials.cs b/dotnet/Src/ProviderCredentials.cs new file mode 100644 index 0000000..5cddcb3 --- /dev/null +++ b/dotnet/Src/ProviderCredentials.cs @@ -0,0 +1,139 @@ +using Azure.Core; +using Microsoft.Extensions.Logging; +using Microsoft.Extensions.Logging.Abstractions; + +namespace Epp.Otp; + +internal sealed class ProviderCredentials : IDisposable +{ + private readonly object _gate = new(); + private readonly ISecretResolver _secrets; + private readonly IEnv _env; + private readonly Func _createIdentity; + private readonly Func>, TokenCredential> _createCredential; + private readonly ILogger _log; + private readonly TimeProvider _clock; + private string? _key; + private RefreshingCache? _bundle; + private RefreshingCache? _assertion; + private TokenCredential? _credential; + private readonly Dictionary> _tokens = new(); + + internal ProviderCredentials(ISecretResolver secrets, IEnv env, + Func createIdentity, + Func>, TokenCredential> createCredential, + ILogger? log = null, TimeProvider? clock = null) + { + _secrets = secrets; + _env = env; + _createIdentity = createIdentity; + _createCredential = createCredential; + _log = log ?? NullLogger.Instance; + _clock = clock ?? TimeProvider.System; + } + + internal void ReportFailure(string kind) => + _log.LogWarning("{{\"logType\":\"service\",\"eventName\":\"credential_refresh_failed\",\"cacheKind\":\"{CacheKind}\",\"failureReason\":\"credential_unavailable\"}}", kind); + + private RefreshingCache Cache(string kind, Func>> load) => + new(load, () => ReportFailure(kind), _clock); + + internal async Task ResolveAsync(AuthConfig auth, AppConfig config, CancellationToken cancellation = default) + { + if (auth.Mode == "oauth" && (string.IsNullOrWhiteSpace(config.ProviderTenantId) + || string.IsNullOrWhiteSpace(config.ProviderScope) || string.IsNullOrWhiteSpace(config.OutboundClientId) + || string.IsNullOrWhiteSpace(config.OutboundManagedIdentityClientId))) + { + Dispose(); + throw new InvalidOperationException("provider OAuth token unavailable"); + } + if (auth.Mode is not ("apiKey" or "oauth")) + { + Dispose(); + throw new InvalidOperationException("provider credential unavailable"); + } + var key = System.Text.Json.JsonSerializer.Serialize(auth.Mode == "apiKey" + ? new[] { auth.Mode, _env.Get("KEY_VAULT_URL"), _env.Get("AZURE_CLIENT_ID"), auth.KeyVaultSecretName, auth.IdentityKeyVaultSecretName } + : new[] { auth.Mode, config.ProviderTenantId, config.OutboundClientId, config.OutboundManagedIdentityClientId }); + RefreshingCache? bundle; + RefreshingCache? tokenCache = null; + lock (_gate) + { + if (_key != key) + { + Clear(); + if (auth.Mode == "apiKey") _bundle = CreateBundle(auth); + else CreateOAuth(config); + _key = key; + } + bundle = _bundle; + if (auth.Mode == "oauth") + { + var scope = config.ProviderScope!; + if (!_tokens.TryGetValue(scope, out tokenCache)) + { + var credential = _credential!; + tokenCache = Cache("provider_token", async ct => + TokenEntry(await credential.GetTokenAsync(new TokenRequestContext(new[] { scope }), ct).ConfigureAwait(false))); + _tokens.Add(scope, tokenCache); + } + } + } + if (bundle is not null) return await bundle.GetAsync(cancellation).ConfigureAwait(false); + var token = await tokenCache!.GetAsync(cancellation).ConfigureAwait(false); + return new ProviderCredential("oauth", AccessToken: token.Token); + } + + private RefreshingCache CreateBundle(AuthConfig auth) => Cache("key_vault", async cancellation => + { + if (string.IsNullOrWhiteSpace(auth.KeyVaultSecretName)) throw new InvalidOperationException("provider credential unavailable"); + var secretTask = _secrets.ResolveAsync(auth.KeyVaultSecretName, cancellation); + var identityTask = string.IsNullOrWhiteSpace(auth.IdentityKeyVaultSecretName) + ? Task.FromResult(string.Empty) : _secrets.ResolveAsync(auth.IdentityKeyVaultSecretName, cancellation); + await Task.WhenAll(secretTask, identityTask).ConfigureAwait(false); + var secret = await secretTask.ConfigureAwait(false); + var identity = await identityTask.ConfigureAwait(false); + if (string.IsNullOrWhiteSpace(secret) || (!string.IsNullOrEmpty(auth.IdentityKeyVaultSecretName) && string.IsNullOrWhiteSpace(identity))) + throw new InvalidOperationException("provider credential unavailable"); + var now = _clock.GetUtcNow(); + return new CredentialCacheEntry(new("apiKey", Secret: secret, Identity: identity), + now.AddMinutes(5), now.AddMinutes(4)); + }); + + private void CreateOAuth(AppConfig config) + { + var identity = _createIdentity(config.OutboundManagedIdentityClientId!); + var assertion = Cache("managed_identity", async cancellation => + TokenEntry(await identity.GetTokenAsync( + new TokenRequestContext(new[] { "api://AzureADTokenExchange/.default" }), cancellation).ConfigureAwait(false))); + _assertion = assertion; + _credential = _createCredential(config.ProviderTenantId!, config.OutboundClientId!, + async cancellation => (await assertion.GetAsync(cancellation).ConfigureAwait(false)).Token); + } + + private CredentialCacheEntry TokenEntry(AccessToken token) + { + var now = _clock.GetUtcNow(); + if (string.IsNullOrWhiteSpace(token.Token) || token.ExpiresOn <= now.AddSeconds(30)) + throw new InvalidOperationException("provider credential unavailable"); + var expires = token.ExpiresOn.AddSeconds(-30); + var refresh = token.ExpiresOn.AddMinutes(-5); + if (token.RefreshOn is { } hint && hint < refresh) refresh = hint; + if (refresh <= now) refresh = now.AddSeconds(Math.Max(1, Math.Min(60, (expires - now).TotalSeconds / 2))); + return new(token, expires, refresh); + } + + private void Clear() + { + _bundle?.Dispose(); + _assertion?.Dispose(); + foreach (var cache in _tokens.Values) cache.Dispose(); + _tokens.Clear(); + _bundle = null; + _assertion = null; + _credential = null; + _key = null; + } + + public void Dispose() { lock (_gate) Clear(); } +} diff --git a/dotnet/Src/RefreshingCache.cs b/dotnet/Src/RefreshingCache.cs new file mode 100644 index 0000000..bdb1b2f --- /dev/null +++ b/dotnet/Src/RefreshingCache.cs @@ -0,0 +1,143 @@ +namespace Epp.Otp; + +internal sealed record CredentialCacheEntry(T Value, DateTimeOffset ExpiresAt, DateTimeOffset RefreshAt) +{ + public override string ToString() => nameof(CredentialCacheEntry); +} + +internal sealed class RefreshingCache : IDisposable +{ + private readonly object _gate = new(); + private readonly Func>> _load; + private readonly TimeProvider _clock; + private readonly Func _random; + private readonly Action _onFailure; + private CredentialCacheEntry? _entry; + private Task? _inFlight; + private CancellationTokenSource? _acquisition; + private ITimer? _timer; + private DateTimeOffset _retryAt; + private int _failures; + private bool _closed; + + internal RefreshingCache(Func>> load, + Action onFailure, TimeProvider? clock = null, Func? random = null) + { + _load = load; + _onFailure = onFailure; + _clock = clock ?? TimeProvider.System; + _random = random ?? Random.Shared.NextDouble; + } + + private static Exception Unavailable() => new InvalidOperationException("provider credential unavailable"); + + internal Task GetAsync(CancellationToken cancellationToken = default) + { + Task pending; + lock (_gate) + { + if (_closed) return Task.FromException(Unavailable()); + var now = _clock.GetUtcNow(); + if (_entry is not null && _entry.ExpiresAt > now) + { + if (_entry.RefreshAt <= now && _retryAt <= now) _ = ObserveAsync(StartRefresh()); + return Task.FromResult(_entry.Value); + } + if (_inFlight is null && _retryAt > now) return Task.FromException(Unavailable()); + pending = StartRefresh(); + } + return cancellationToken.CanBeCanceled ? pending.WaitAsync(cancellationToken) : pending; + } + + internal Task RefreshAsync() + { + lock (_gate) + { + if (_closed || _retryAt > _clock.GetUtcNow()) return Task.FromException(Unavailable()); + return StartRefresh(); + } + } + + private Task StartRefresh() + { + if (_inFlight is not null) return _inFlight; + _timer?.Dispose(); + _timer = null; + var completion = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + _inFlight = completion.Task; + _acquisition = new CancellationTokenSource(TimeSpan.FromSeconds(2.5), _clock); + _ = RunRefreshAsync(completion, _acquisition); + return completion.Task; + } + + private async Task RunRefreshAsync(TaskCompletionSource completion, CancellationTokenSource cancellation) + { + CredentialCacheEntry? entry = null; + try + { + entry = await _load(cancellation.Token).WaitAsync(cancellation.Token).ConfigureAwait(false); + if (entry.ExpiresAt <= _clock.GetUtcNow()) entry = null; + } + catch (Exception) + { + // Emit only the fixed failure below; SDK exceptions may contain credentials. + } + lock (_gate) + { + _inFlight = null; + _acquisition = null; + if (_closed || cancellation.IsCancellationRequested) entry = null; + cancellation.Dispose(); + if (!_closed) + { + if (entry is null) + { + _failures++; + var backoff = Math.Min(60, 5 * Math.Pow(2, Math.Min(_failures - 1, 4))); + _retryAt = _clock.GetUtcNow().AddSeconds(backoff * (1 + _random() * 0.2)); + _onFailure(); + } + else + { + _entry = entry; + _failures = 0; + _retryAt = default; + } + var next = entry is null ? _retryAt : entry.RefreshAt; + var wait = Math.Clamp((next - _clock.GetUtcNow()).TotalMilliseconds, 1000, int.MaxValue); + _timer = _clock.CreateTimer(_ => ScheduledRefresh(), null, TimeSpan.FromMilliseconds(wait), Timeout.InfiniteTimeSpan); + } + if (entry is null) completion.TrySetException(Unavailable()); + else completion.TrySetResult(entry.Value); + } + } + + private void ScheduledRefresh() + { + lock (_gate) + { + _timer?.Dispose(); + _timer = null; + if (!_closed) _ = ObserveAsync(StartRefresh()); + } + } + + private static async Task ObserveAsync(Task task) + { + // Refresh failures are already reported; a timer has no request awaiting the result. + try { await task.ConfigureAwait(false); } + catch (Exception) { } + } + + public void Dispose() + { + lock (_gate) + { + _closed = true; + _timer?.Dispose(); + _timer = null; + _acquisition?.Cancel(); + _entry = null; + } + } +} diff --git a/dotnet/Src/SecretResolver.cs b/dotnet/Src/SecretResolver.cs index 3a97f65..f95b4b0 100644 --- a/dotnet/Src/SecretResolver.cs +++ b/dotnet/Src/SecretResolver.cs @@ -1,40 +1,50 @@ -using System.Collections.Concurrent; using Azure.Identity; using Azure.Security.KeyVault.Secrets; namespace Epp.Otp; // Resolves Key Vault secret names to values via the Function's managed identity (user-assigned when -// AZURE_CLIENT_ID is set, else system-assigned), cached briefly so rotations are picked up. +// AZURE_CLIENT_ID is set, else system-assigned). ProviderCredentials caches the complete bundle. public sealed class SecretResolver : ISecretResolver { - private static readonly TimeSpan CacheTtl = TimeSpan.FromMinutes(5); - private readonly ConcurrentDictionary _cache = new(); - private readonly Lazy _client; + private readonly object _gate = new(); + private readonly IEnv _env; + private SecretClient? _client; + private (string Url, string? Identity)? _clientKey; public SecretResolver(IEnv? env = null) { - var environment = env ?? new ProcessEnv(); - _client = new Lazy(() => + _env = env ?? new ProcessEnv(); + } + + private SecretClient GetClient() + { + var url = _env.Get("KEY_VAULT_URL"); + var clientId = _env.Get("AZURE_CLIENT_ID"); + if (string.IsNullOrWhiteSpace(url)) throw new InvalidOperationException("KEY_VAULT_URL not set"); + lock (_gate) { - var url = environment.Get("KEY_VAULT_URL"); - if (string.IsNullOrWhiteSpace(url)) return null; - var clientId = environment.Get("AZURE_CLIENT_ID"); - var credential = string.IsNullOrEmpty(clientId) - ? new ManagedIdentityCredential() - : new ManagedIdentityCredential(clientId); - return new SecretClient(new Uri(url), credential); - }); + if (_client is not null && _clientKey == (url, clientId)) return _client; + var identityOptions = new TokenCredentialOptions(); + identityOptions.Retry.MaxRetries = 0; + identityOptions.Retry.NetworkTimeout = TimeSpan.FromSeconds(2.5); + identityOptions.Diagnostics.IsLoggingEnabled = false; + identityOptions.Diagnostics.IsLoggingContentEnabled = false; + var credential = new ManagedIdentityCredential(clientId, identityOptions); + var options = new SecretClientOptions(); + options.Retry.MaxRetries = 0; + options.Retry.NetworkTimeout = TimeSpan.FromSeconds(2.5); + options.Diagnostics.IsLoggingEnabled = false; + options.Diagnostics.IsLoggingContentEnabled = false; + _client = new SecretClient(new Uri(url), credential, options); + _clientKey = (url, clientId); + return _client; + } } - public async Task ResolveAsync(string? secretName) + public async Task ResolveAsync(string? secretName, CancellationToken cancellationToken = default) { if (string.IsNullOrWhiteSpace(secretName)) return string.Empty; - if (_cache.TryGetValue(secretName, out var cached) && cached.Expires > DateTimeOffset.UtcNow) return cached.Value; - - var client = _client.Value ?? throw new InvalidOperationException("KEY_VAULT_URL not set"); - var value = (await client.GetSecretAsync(secretName)).Value.Value ?? string.Empty; - _cache[secretName] = (value, DateTimeOffset.UtcNow.Add(CacheTtl)); - return value; + return (await GetClient().GetSecretAsync(secretName, cancellationToken: cancellationToken)).Value.Value ?? string.Empty; } } diff --git a/dotnet/tests/CredentialCacheTests.cs b/dotnet/tests/CredentialCacheTests.cs new file mode 100644 index 0000000..649fab8 --- /dev/null +++ b/dotnet/tests/CredentialCacheTests.cs @@ -0,0 +1,306 @@ +using Azure.Core; +using Xunit; + +namespace Epp.Otp.Tests; + +public class CredentialCacheTests +{ + [Fact] + public async Task ColdReadersShareOneFetchAndValidValuesRemainAvailableDuringRefresh() + { + var clock = new ManualClock(); + var release = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var calls = 0; + using var cache = new RefreshingCache(async _ => + { + Interlocked.Increment(ref calls); + await release.Task; + var now = clock.GetUtcNow(); + return new("value", now.AddMinutes(5), now.AddMinutes(4)); + }, () => Assert.Fail("Unexpected refresh failure"), clock); + var readers = Enumerable.Range(0, 20).Select(_ => cache.GetAsync()).ToArray(); + Assert.Equal(1, calls); + release.SetResult(); + Assert.All(await Task.WhenAll(readers), value => Assert.Equal("value", value)); + Assert.Equal(1, calls); + Assert.Equal(1, clock.TimerCount); + release = new(TaskCreationOptions.RunContinuationsAsynchronously); + clock.Advance(TimeSpan.FromMinutes(4)); + Assert.Equal(2, calls); + Assert.Equal("value", await cache.GetAsync()); + release.SetResult(); + await Until(() => clock.TimerCount == 1); + Assert.Equal("value", await cache.GetAsync()); + cache.Dispose(); + Assert.Equal(0, clock.TimerCount); + } + + [Fact] + public async Task RefreshFailuresDoNotExtendExpiryAndUseBackoff() + { + var clock = new ManualClock(); + var fail = false; + var calls = 0; + var failures = 0; + using var cache = new RefreshingCache(_ => + { + calls++; + if (fail) throw new InvalidOperationException("PRIVATE-ERROR"); + var now = clock.GetUtcNow(); + return Task.FromResult(new CredentialCacheEntry("first", now.AddMinutes(5), now.AddMinutes(4))); + }, () => failures++, clock, () => 0); + Assert.Equal("first", await cache.GetAsync()); + fail = true; + clock.Advance(TimeSpan.FromMinutes(4)); + Assert.Equal(1, failures); + for (var i = 0; i < 10; i++) Assert.Equal("first", await cache.GetAsync()); + Assert.Equal(2, calls); + clock.Advance(TimeSpan.FromMinutes(1)); + Assert.Equal(2, failures); + var error = await Assert.ThrowsAsync(() => cache.GetAsync()); + Assert.Equal("provider credential unavailable", error.Message); + fail = false; + clock.Advance(TimeSpan.FromSeconds(10)); + Assert.Equal("first", await cache.GetAsync()); + Assert.Equal(4, calls); + } + + [Fact] + public async Task CancellingAWaiterDoesNotCancelTheSharedRefresh() + { + var clock = new ManualClock(); + var release = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + CancellationToken observed = default; + using var cache = new RefreshingCache(async cancellation => + { + observed = cancellation; + await release.Task; + return new("ready", clock.GetUtcNow().AddMinutes(5), clock.GetUtcNow().AddMinutes(4)); + }, () => Assert.Fail("Unexpected refresh failure"), clock); + using var waiter = new CancellationTokenSource(); + var first = cache.GetAsync(waiter.Token); + var second = cache.GetAsync(); + waiter.Cancel(); + await Assert.ThrowsAnyAsync(() => first); + Assert.False(observed.IsCancellationRequested); + release.SetResult(); + Assert.Equal("ready", await second); + } + + [Fact] + public async Task CacheOwnedDeadlineAndShutdownPreventLatePublication() + { + var clock = new ManualClock(); + var release = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var errors = 0; + CancellationToken observed = default; + using var cache = new RefreshingCache(async cancellation => + { + observed = cancellation; + await release.Task; + return new("late", clock.GetUtcNow().AddMinutes(5), clock.GetUtcNow().AddMinutes(4)); + }, () => errors++, clock); + var pending = cache.GetAsync(); + clock.Advance(TimeSpan.FromSeconds(2.5)); + await Assert.ThrowsAsync(() => pending); + Assert.True(observed.IsCancellationRequested); + Assert.Equal(1, errors); + cache.Dispose(); + release.SetResult(); + await Assert.ThrowsAsync(() => cache.GetAsync()); + Assert.Equal(0, clock.TimerCount); + } + + [Fact] + public async Task ApiKeyPairIsFetchedInParallelAndPublishedAsOneBundle() + { + var clock = new ManualClock(); + var release = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var calls = new List(); + var version = 1; + var failIdentity = false; + var secrets = new Secrets(async (name, _) => + { + calls.Add(name!); + await release.Task; + if (failIdentity && name == "id") throw new InvalidOperationException("PRIVATE-ERROR"); + return name + "-" + version; + }); + using var manager = new ProviderCredentials(secrets, new TestEnv(), _ => throw new Exception(), + (_, _, _) => throw new Exception(), clock: clock); + var auth = new AuthConfig("apiKey", "key", "id"); + var pending = Enumerable.Range(0, 10).Select(_ => manager.ResolveAsync(auth, new AppConfig())).ToArray(); + Assert.Equal(new[] { "key", "id" }, calls); + release.SetResult(); + foreach (var value in await Task.WhenAll(pending)) + { + Assert.Equal("key-1", value.Secret); + Assert.Equal("id-1", value.Identity); + } + version = 2; + failIdentity = true; + clock.Advance(TimeSpan.FromMinutes(4)); + var old = await manager.ResolveAsync(auth, new AppConfig()); + Assert.Equal("key-1", old.Secret); + Assert.Equal("id-1", old.Identity); + failIdentity = false; + clock.Advance(TimeSpan.FromSeconds(7)); + var next = await manager.ResolveAsync(auth, new AppConfig()); + Assert.Equal("key-2", next.Secret); + Assert.Equal("id-2", next.Identity); + } + + [Fact] + public async Task ManagedIdentityAndEntraTokensAreCachedRefreshedAndIsolatedByConfiguration() + { + var clock = new ManualClock(); + var identityCalls = 0; + var providerCalls = 0; + var credentialInstances = 0; + using var manager = new ProviderCredentials(new Secrets((_, _) => throw new Exception("Unexpected Key Vault")), + new TestEnv(), _ => new Token(async (_, _) => + { + Interlocked.Increment(ref identityCalls); + await Task.Yield(); + return new("PRIVATE-ASSERTION", clock.GetUtcNow().AddHours(1)); + }), (_, _, assertion) => + { + credentialInstances++; + return new Token(async (_, cancellation) => + { + Interlocked.Increment(ref providerCalls); + Assert.Equal("PRIVATE-ASSERTION", await assertion(cancellation)); + Assert.Equal("PRIVATE-ASSERTION", await assertion(cancellation)); + return new("PRIVATE-PROVIDER", clock.GetUtcNow().AddHours(1)); + }); + }, clock: clock); + var config = Config(); + var initial = await Task.WhenAll(Enumerable.Range(0, 20).Select(_ => manager.ResolveAsync(new("oauth"), config))); + Assert.All(initial, result => Assert.Equal("PRIVATE-PROVIDER", result.AccessToken)); + Assert.Equal(1, identityCalls); + Assert.Equal(1, providerCalls); + await manager.ResolveAsync(new("oauth"), config); + Assert.Equal(1, providerCalls); + clock.Advance(TimeSpan.FromMinutes(55)); + await Until(() => identityCalls == 2 && providerCalls == 2 && clock.TimerCount == 2); + await manager.ResolveAsync(new("oauth"), Config(scope: "api://second/.default")); + Assert.Equal(2, identityCalls); + Assert.Equal(3, providerCalls); + Assert.Equal(1, credentialInstances); + await manager.ResolveAsync(new("oauth"), Config(application: "different")); + Assert.Equal(3, identityCalls); + Assert.Equal(2, credentialInstances); + manager.Dispose(); + Assert.Equal(0, clock.TimerCount); + } + + [Fact] + public async Task RepeatedSdkTokenDoesNotExtendLifetimeOrCauseATightRefreshLoop() + { + var clock = new ManualClock(); + var expiry = clock.GetUtcNow().AddHours(1); + var calls = 0; + using var manager = new ProviderCredentials(new Secrets((_, _) => throw new Exception()), new TestEnv(), + _ => new Token((_, _) => ValueTask.FromResult(new AccessToken("assertion", expiry))), + (_, _, _) => new Token((_, _) => + { + calls++; + return ValueTask.FromResult(new AccessToken("token", expiry)); + }), clock: clock); + await manager.ResolveAsync(new("oauth"), Config()); + clock.Advance(TimeSpan.FromMinutes(55)); + Assert.Equal(2, calls); + await manager.ResolveAsync(new("oauth"), Config()); + Assert.Equal(2, calls); + clock.Advance(TimeSpan.FromMinutes(1)); + Assert.Equal(3, calls); + clock.Advance(TimeSpan.FromSeconds(210)); + await Assert.ThrowsAsync(() => manager.ResolveAsync(new("oauth"), Config())); + } + + private static AppConfig Config(string scope = "api://provider/.default", string application = "app") => new() + { + ProviderTenantId = "tenant", ProviderScope = scope, OutboundClientId = application, + OutboundManagedIdentityClientId = "identity", + }; + + private static async Task Until(Func condition) + { + for (var i = 0; i < 200; i++) + { + if (condition()) return; + await Task.Delay(5); + } + Assert.True(condition()); + } + + private sealed class Secrets(Func> resolve) : ISecretResolver + { + public Task ResolveAsync(string? name, CancellationToken cancellationToken = default) => resolve(name, cancellationToken); + } + + private sealed class Token(Func> acquire) : TokenCredential + { + public override AccessToken GetToken(TokenRequestContext requestContext, CancellationToken cancellationToken) => + throw new InvalidOperationException("Synchronous acquisition was not expected"); + public override ValueTask GetTokenAsync(TokenRequestContext requestContext, CancellationToken cancellationToken) => + acquire(requestContext, cancellationToken); + } + + private sealed class ManualClock : TimeProvider + { + private readonly object _gate = new(); + private readonly List _timers = new(); + private DateTimeOffset _now = new(2026, 1, 1, 0, 0, 0, TimeSpan.Zero); + public override DateTimeOffset GetUtcNow() { lock (_gate) return _now; } + public int TimerCount { get { lock (_gate) return _timers.Count(timer => !timer.Disposed && timer.Due != DateTimeOffset.MaxValue); } } + public override ITimer CreateTimer(TimerCallback callback, object? state, TimeSpan dueTime, TimeSpan period) + { + lock (_gate) + { + var timer = new ManualTimer(this, callback, state); + timer.Change(dueTime, period); + _timers.Add(timer); + return timer; + } + } + public void Advance(TimeSpan duration) + { + ManualTimer[] due; + lock (_gate) + { + _now += duration; + due = _timers.Where(timer => !timer.Disposed && timer.Due <= _now).ToArray(); + } + foreach (var timer in due) timer.Fire(); + } + + private sealed class ManualTimer(ManualClock clock, TimerCallback callback, object? state) : ITimer + { + internal bool Disposed; + internal DateTimeOffset Due; + private TimeSpan _period; + public bool Change(TimeSpan dueTime, TimeSpan period) + { + lock (clock._gate) + { + if (Disposed) return false; + Due = dueTime == Timeout.InfiniteTimeSpan ? DateTimeOffset.MaxValue : clock._now + dueTime; + _period = period; + return true; + } + } + internal void Fire() + { + lock (clock._gate) + { + if (Disposed) return; + Due = _period == Timeout.InfiniteTimeSpan ? DateTimeOffset.MaxValue : clock._now + _period; + } + callback(state); + } + public void Dispose() { lock (clock._gate) { Disposed = true; clock._timers.Remove(this); } } + public ValueTask DisposeAsync() { Dispose(); return ValueTask.CompletedTask; } + } + } +} diff --git a/dotnet/tests/EngineTests.cs b/dotnet/tests/EngineTests.cs index 94967ed..3727163 100644 --- a/dotnet/tests/EngineTests.cs +++ b/dotnet/tests/EngineTests.cs @@ -34,20 +34,46 @@ private static void ConfigureSoprano(HandlerRig rig) rig.Http.Respond = _ => Task.FromResult(Json(201, "{\"status\":\"ENROUTE\"}")); } + [Fact] + public async Task StartupPreparesOnlyCredentialsAndWarmRequestsReuseTheBundle() + { + using var rig = new HandlerRig(); + await rig.Engine.StartCredentialRefreshAsync(); + Assert.Equal(1, rig.Secrets.Calls); + Assert.Equal(0, rig.Http.Calls); + Assert.Equal(0, rig.Keys.Calls); + AssertAccepted(await rig.Invoke("evaluation")); + Assert.Equal(1, rig.Secrets.Calls); + Assert.Equal(0, rig.Http.Calls); + AssertAccepted(await rig.Invoke()); + Assert.Equal(1, rig.Secrets.Calls); + Assert.Equal(1, rig.Http.Calls); + } + + [Fact] + public async Task StartupWithoutProviderConfigurationKeepsEvaluationIndependent() + { + using var rig = new HandlerRig(); + rig.Env.Clear(); + await rig.Engine.StartCredentialRefreshAsync(); + Assert.Equal(0, rig.Secrets.Calls); + Assert.Equal(0, rig.Http.Calls); + AssertAccepted(await rig.Invoke("evaluation")); + } + [Fact] public async Task SopranoOAuthUsesSetupIdentitiesScopeAndOneBoundedExchange() { var scopes = new List(); var identities = new List(); var applications = new List<(string Tenant, string Application)>(); - CancellationToken outerCancellation = default; using var rig = new HandlerRig(identity => { identities.Add(identity); return new TestTokenCredential((context, cancellation) => { Assert.Equal("api://AzureADTokenExchange/.default", Assert.Single(context.Scopes)); - Assert.Equal(outerCancellation, cancellation); + Assert.True(cancellation.CanBeCanceled); return ValueTask.FromResult(new AccessToken("private-assertion", DateTimeOffset.UtcNow.AddHours(1))); }); }, (tenant, application, assertion) => @@ -56,7 +82,6 @@ public async Task SopranoOAuthUsesSetupIdentitiesScopeAndOneBoundedExchange() return new TestTokenCredential(async (context, cancellation) => { Assert.True(cancellation.CanBeCanceled); - outerCancellation = cancellation; scopes.Add(Assert.Single(context.Scopes)); Assert.Equal("private-assertion", await assertion(cancellation)); return new AccessToken("private-provider-token", DateTimeOffset.UtcNow.AddHours(1)); @@ -93,7 +118,7 @@ public async Task SopranoOAuthUsesSetupIdentitiesScopeAndOneBoundedExchange() rig.Env["EPP_PROVIDER_SCOPE"] = "api://second/.default"; AssertAccepted(await rig.Invoke()); Assert.Single(applications); - Assert.Equal(new[] { "api://provider/.default", "api://provider/.default", "api://second/.default" }, scopes); + Assert.Equal(new[] { "api://provider/.default", "api://second/.default" }, scopes); rig.Env["EPP_OUTBOUND_CLIENT_ID"] = "second-application"; AssertAccepted(await rig.Invoke()); Assert.Equal(new[] { ("provider-tenant", "calling-application"), ("provider-tenant", "second-application") }, applications); @@ -148,6 +173,7 @@ public async Task SopranoOAuthCancellationAndRejectionNeverFallBackOrRetry() Assert.True(observed.IsCancellationRequested); Assert.Equal((0, 0), (rig.Http.Calls, rig.Secrets.Calls)); waitForCancellation = false; + rig.Engine.Dispose(); rig.Http.Respond = _ => Task.FromResult(Json(401, "{\"status\":\"REJECTED\"}")); AssertFailure(rig, await rig.Invoke(), 401); Assert.Equal((1, 0), (rig.Http.Calls, rig.Secrets.Calls)); @@ -800,6 +826,7 @@ private sealed class HandlerRig : IDisposable public TestHttp Http { get; } = new(); public TestKeys Keys { get; } = new(); public CapturingLogger Log { get; } = new(); + public DispatchEngine Engine { get; } public HandlerRig(Func? createIdentity = null, Func>, TokenCredential>? createOAuth = null) { @@ -813,6 +840,7 @@ public HandlerRig(Func? createIdentity = null, { new InfobipProvider(), new TelesignProvider(), new SopranoProvider(), new SinchProvider() }); var engine = createIdentity is null ? new DispatchEngine(registry, Secrets, Http, Env) : new DispatchEngine(registry, Secrets, Http, Env, createIdentity, createOAuth!); + Engine = engine; _function = new SendOtp(engine, new JweDecryptor(Keys), Env, Log); } @@ -843,7 +871,7 @@ public async Task InvokeRaw(string body, Dictionary(await _function.Run(request)); } - public void Dispose() { Keys.Dispose(); Http.Dispose(); } + public void Dispose() { Engine.Dispose(); Keys.Dispose(); Http.Dispose(); } } private sealed class TestSecrets : ISecretResolver @@ -852,7 +880,7 @@ private sealed class TestSecrets : ISecretResolver public string Secret { get; set; } = "private-api-key"; public string Identity { get; set; } = "private-api-id"; public Exception? Error { get; set; } - public Task ResolveAsync(string? name) + public Task ResolveAsync(string? name, CancellationToken cancellationToken = default) { Calls++; if (Error is not null) throw Error; diff --git a/javascript/README.md b/javascript/README.md index 0fa4e04..30d7580 100644 --- a/javascript/README.md +++ b/javascript/README.md @@ -99,12 +99,20 @@ retries. The shared contract defines validation, HTTP outcomes and privacy-safe ## Source and extension points +Configured providers automatically prewarm on the app-start hook. Key Vault credential bundles, +managed-identity assertions and final Entra tokens refresh through separate process-local caches. +Warm requests reuse usable values; concurrent misses share a retrieval, refresh failure never +extends expiry, and termination stops timers. Startup/refresh never sends an OTP. See the +[refresh contract](../docs/CONTRACT.md#credential-caching-and-refresh) for budgets and cold-start +limitations. Leave the provider unset for local evaluation-only use without credential acquisition. + | Source | Purpose | |---|---| | [src/functions/SendOtp.js](src/functions/SendOtp.js) | HTTP handler | | [src/functions/config.js](src/functions/config.js) | Shared deployment settings | | [src/functions/models.js](src/functions/models.js) | Delivery context, normalized `ParsedResponse`, and documented request objects | | [src/functions/dispatch.js](src/functions/dispatch.js) | Envelope/JWE handling, registry and dispatch | +| [src/functions/credentials.js](src/functions/credentials.js), [refreshingCache.js](src/functions/refreshingCache.js) | Provider credential acquisition, single-flight caching and scheduled refresh | | [src/functions/requestLog.js](src/functions/requestLog.js) | Request-scoped [service events and summaries](../docs/CONTRACT.md#application-logs) with explicit ID sources | | [src/functions/providers/](src/functions/providers/) | Adapter manifests and API-specific implementations | | [test/](test/) | Representative offline checks | diff --git a/javascript/src/functions/SendOtp.js b/javascript/src/functions/SendOtp.js index 65a466e..20142cc 100644 --- a/javascript/src/functions/SendOtp.js +++ b/javascript/src/functions/SendOtp.js @@ -11,11 +11,16 @@ const { parseEnvelope, decryptDeliveryContext, contextToDispatch, + startProviderCredentialRefresh, + stopProviderCredentialRefresh, MODE, } = require('./dispatch'); const { readConfig } = require('./config'); const { RequestLog } = require('./requestLog'); +app.hook.appStart(startProviderCredentialRefresh); +app.hook.appTerminate(stopProviderCredentialRefresh); + app.http('SendOtp', { methods: ['POST'], authLevel: 'anonymous', // Protected by platform authentication in Azure. diff --git a/javascript/src/functions/credentials.js b/javascript/src/functions/credentials.js new file mode 100644 index 0000000..9c14952 --- /dev/null +++ b/javascript/src/functions/credentials.js @@ -0,0 +1,181 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +'use strict'; + +const { AsyncLocalStorage } = require('node:async_hooks'); +const { inspect } = require('node:util'); +const { ClientAssertionCredential, ManagedIdentityCredential } = require('@azure/identity'); +const { SecretClient } = require('@azure/keyvault-secrets'); +const { AzureLogger } = require('@azure/logger'); +const { RefreshingCache, tokenEntry } = require('./refreshingCache'); + +const acquisition = new AsyncLocalStorage(); +let filteredLogger; + +function reportRefreshFailure(cacheKind) { + console.warn(JSON.stringify({ logType: 'service', eventName: 'credential_refresh_failed', + cacheKind, failureReason: 'credential_unavailable' })); +} + +function cancellationPolicy() { + return { + name: 'eppCredentialCancellation', + async sendRequest(request, next) { + const signal = acquisition.getStore(); + if (!signal) return next(request); + const controller = new AbortController(); + const existing = request.abortSignal; + const abort = () => controller.abort(); + signal.addEventListener('abort', abort, { once: true }); + existing?.addEventListener('abort', abort, { once: true }); + if (signal.aborted || existing?.aborted) controller.abort(); + request.abortSignal = controller.signal; + try { + controller.signal.throwIfAborted(); + return await next(request); + } finally { + signal.removeEventListener('abort', abort); + existing?.removeEventListener('abort', abort); + } + }, + }; +} + +async function acquireBounded(signal, load) { + if (AzureLogger.log !== filteredLogger) { + const previous = AzureLogger.log; + filteredLogger = (...args) => { if (!acquisition.getStore()) previous(...args); }; + AzureLogger.log = filteredLogger; + } + const controller = new AbortController(); + const abort = () => controller.abort(); + signal.addEventListener('abort', abort, { once: true }); + if (signal.aborted) abort(); + let timeout; + const interrupted = new Promise((_, reject) => { + const fail = () => reject(new Error('provider credential unavailable')); + controller.signal.addEventListener('abort', fail, { once: true }); + if (controller.signal.aborted) fail(); + timeout = setTimeout(abort, 2500); + }); + try { + return await acquisition.run(controller.signal, () => Promise.race([ + Promise.resolve().then(() => { + controller.signal.throwIfAborted(); + return load(controller.signal); + }), + interrupted, + ])); + } finally { + clearTimeout(timeout); + signal.removeEventListener('abort', abort); + } +} + +const sdkOptions = () => ({ retryOptions: { maxRetries: 0 }, + additionalPolicies: [{ policy: cancellationPolicy(), position: 'perCall' }] }); + +class ProviderCredentials { + constructor({ cacheOptions = {}, reportFailure = reportRefreshFailure } = {}) { + this.cacheOptions = cacheOptions; + this.now = cacheOptions.now || Date.now; + this.reportFailure = reportFailure; + this.current = null; + this.currentKey = null; + } + + [inspect.custom]() { return '[ProviderCredentials]'; } + toJSON() { return '[ProviderCredentials]'; } + + cache(kind, load) { + return new RefreshingCache((signal) => acquireBounded(signal, load), + { ...this.cacheOptions, onFailure: () => this.reportFailure(kind) }); + } + + resolve(auth, config) { + const mode = auth?.mode || 'apiKey'; + if (mode === 'oauth' && (!config.providerTenantId || !config.providerScope + || !config.outboundClientId || !config.outboundManagedIdentityClientId)) { + this.close(); + return Promise.reject(new Error('provider OAuth token unavailable')); + } + if (mode !== 'oauth' && mode !== 'apiKey') { + this.close(); + return Promise.reject(new Error('provider credential unavailable')); + } + const key = JSON.stringify(mode === 'apiKey' + ? [mode, config.keyVaultUrl, config.managedIdentityClientId, auth.keyVaultSecretName, auth.identityKeyVaultSecretName] + : [mode, config.providerTenantId, config.outboundClientId, config.outboundManagedIdentityClientId]); + if (this.currentKey !== key) { + this.close(); + this.current = mode === 'apiKey' ? this.apiKeyState(auth, config) : this.oauthState(config); + this.currentKey = key; + } + if (mode === 'apiKey') return this.current.bundle.get(); + const state = this.current; + if (!state.tokens.has(config.providerScope)) { + const scope = config.providerScope; + state.tokens.set(scope, this.cache('provider_token', async (signal) => + tokenEntry(await state.credential.getToken(scope, { abortSignal: signal }), this.now()))); + } + return state.tokens.get(config.providerScope).get().then((token) => + Object.defineProperty({ mode: 'oauth' }, 'accessToken', { value: token.token })) + .catch(() => { throw new Error('provider OAuth token unavailable'); }); + } + + apiKeyState(auth, config) { + let client; + const bundle = this.cache('key_vault', async (signal) => { + if (!auth.keyVaultSecretName || !config.keyVaultUrl) throw new Error('provider credential unavailable'); + if (!client) { + const identity = config.managedIdentityClientId + ? new ManagedIdentityCredential(config.managedIdentityClientId, sdkOptions()) + : new ManagedIdentityCredential(sdkOptions()); + client = new SecretClient(config.keyVaultUrl, identity, sdkOptions()); + } + const [secret, identity] = await Promise.all([ + client.getSecret(auth.keyVaultSecretName, { abortSignal: signal }), + auth.identityKeyVaultSecretName ? client.getSecret(auth.identityKeyVaultSecretName, { abortSignal: signal }) : null, + ]); + if (!secret.value?.trim() || (auth.identityKeyVaultSecretName && !identity?.value?.trim())) { + throw new Error('provider credential unavailable'); + } + const now = this.now(); + let expiresAt = now + 300000; + for (const item of [secret, identity]) { + if (!item) continue; + if (item.properties?.enabled === false + || item.properties?.notBefore?.getTime() > now) throw new Error('provider credential unavailable'); + if (item.properties?.expiresOn) expiresAt = Math.min(expiresAt, item.properties.expiresOn.getTime()); + } + return { value: Object.freeze({ mode: 'apiKey', secret: secret.value, identity: identity?.value || '' }), + expiresAt, refreshAt: Math.min(now + 240000, expiresAt - 30000) }; + }); + return { bundle }; + } + + oauthState(config) { + let identity; + const assertion = this.cache('managed_identity', async (signal) => { + identity ??= new ManagedIdentityCredential({ clientId: config.outboundManagedIdentityClientId, ...sdkOptions() }); + return tokenEntry(await identity.getToken('api://AzureADTokenExchange/.default', { abortSignal: signal }), this.now()); + }); + const credential = new ClientAssertionCredential(config.providerTenantId, config.outboundClientId, + async () => { + const value = await assertion.get(); + acquisition.getStore()?.throwIfAborted(); + return value.token; + }, { authorityHost: 'https://login.microsoftonline.com', ...sdkOptions() }); + return { assertion, credential, tokens: new Map() }; + } + + close() { + this.current?.bundle?.close(); + this.current?.assertion?.close(); + if (this.current?.tokens) for (const cache of this.current.tokens.values()) cache.close(); + this.current = null; + this.currentKey = null; + } +} + +const providerCredentials = new ProviderCredentials(); +module.exports = { ProviderCredentials, providerCredentials, reportRefreshFailure }; diff --git a/javascript/src/functions/dispatch.js b/javascript/src/functions/dispatch.js index 2ae8971..5e08ace 100644 --- a/javascript/src/functions/dispatch.js +++ b/javascript/src/functions/dispatch.js @@ -6,12 +6,9 @@ const crypto = require('crypto'); const { compactDecrypt } = require('jose'); -const { ClientAssertionCredential, ManagedIdentityCredential } = require('@azure/identity'); -const { SecretClient } = require('@azure/keyvault-secrets'); -const { AzureLogger } = require('@azure/logger'); -const { AsyncLocalStorage } = require('node:async_hooks'); const { readConfig } = require('./config'); const { DeliveryContext, TextToVoice } = require('./models'); +const { providerCredentials, reportRefreshFailure } = require('./credentials'); const CHANNEL_BY_CODE = Object.freeze({ 1: 'sms', 2: 'voice' }); const CHANNEL_BY_NAME = Object.freeze({ sms: 1, voice: 2 }); @@ -144,8 +141,6 @@ const OUTCOME = Object.freeze({ STEP_UP: 'StepUp', }); -const SECRET_CACHE_TIME_TO_LIVE_MILLISECONDS = 5 * 60 * 1000; // rotated secrets picked up within this window - const providerRegistry = new Map( [ require('./providers/infobip'), @@ -162,99 +157,28 @@ function getProvider(providerId) { return providerId ? providerRegistry.get(String(providerId).trim().toLowerCase()) || null : null; } -let keyVaultSecretClient = null; -let keyVaultClientConfig; -const secretCache = new Map(); -let oauthCredential = null; -let oauthCredentialConfig; -const oauthRequest = new AsyncLocalStorage(); -let filteredOAuthLogger; - -function usableAccessToken(value) { - return typeof value?.token === 'string' && value.token.trim() - && Number.isFinite(value.expiresOnTimestamp) && value.expiresOnTimestamp > Date.now() + 30000; -} - -function getKeyVaultSecretClient(config) { - const cacheKey = JSON.stringify([config.keyVaultUrl, config.managedIdentityClientId]); - if (!keyVaultSecretClient || keyVaultClientConfig !== cacheKey) { - const credential = config.managedIdentityClientId - ? new ManagedIdentityCredential(config.managedIdentityClientId) - : new ManagedIdentityCredential(); - keyVaultSecretClient = new SecretClient(config.keyVaultUrl, credential); - keyVaultClientConfig = cacheKey; - } - return keyVaultSecretClient; +async function resolveProviderCredential(authConfiguration = {}, config) { + return providerCredentials.resolve(authConfiguration, config); } -async function resolveSecretValue(keyVaultSecretName, config) { - if (!keyVaultSecretName) { - return ''; +async function startProviderCredentialRefresh() { + const config = readConfig(); + if (!config.providerName) return; + const provider = getProvider(config.providerName); + if (!provider || (config.providerAuthMode && config.providerAuthMode !== provider.manifest.auth.mode)) { + reportRefreshFailure('configuration'); + return; } - const cacheKey = JSON.stringify([config.keyVaultUrl, config.managedIdentityClientId, keyVaultSecretName]); - const cachedSecret = secretCache.get(cacheKey); - if (cachedSecret && cachedSecret.expiresAt > Date.now()) { - return cachedSecret.value; + try { + await resolveProviderCredential(provider.manifest.auth, config); + } catch { + // The cache reports acquisition failures; also report configurations rejected before caching. + if (!providerCredentials.current) reportRefreshFailure('configuration'); } - - const secretValue = (await getKeyVaultSecretClient(config).getSecret(keyVaultSecretName)).value || ''; - - secretCache.set(cacheKey, { - value: secretValue, - expiresAt: Date.now() + SECRET_CACHE_TIME_TO_LIVE_MILLISECONDS, - }); - return secretValue; } -async function resolveProviderCredential(authConfiguration = {}, config) { - const { mode = 'apiKey' } = authConfiguration; - if (mode === 'apiKey') { - const [secret, identity] = await Promise.all([ - resolveSecretValue(authConfiguration.keyVaultSecretName, config), - authConfiguration.identityKeyVaultSecretName - ? resolveSecretValue(authConfiguration.identityKeyVaultSecretName, config) - : Promise.resolve(''), - ]); - return { mode: 'apiKey', secret, identity }; - } - if (mode !== 'oauth' || !config.providerTenantId || !config.providerScope - || !config.outboundClientId || !config.outboundManagedIdentityClientId) { - throw new Error('unsupported or incomplete provider authentication'); - } - if (AzureLogger.log !== filteredOAuthLogger) { - const log = AzureLogger.log; - filteredOAuthLogger = (...args) => { if (!oauthRequest.getStore()) log(...args); }; - AzureLogger.log = filteredOAuthLogger; - } - return oauthRequest.run(AbortSignal.timeout(2500), async () => { - try { - const credentialConfig = JSON.stringify([ - config.providerTenantId, config.outboundClientId, config.outboundManagedIdentityClientId, - ]); - if (!oauthCredential || oauthCredentialConfig !== credentialConfig) { - const assertionIdentity = new ManagedIdentityCredential({ - clientId: config.outboundManagedIdentityClientId, retryOptions: { maxRetries: 0 }, - }); - oauthCredential = new ClientAssertionCredential( - config.providerTenantId, - config.outboundClientId, - async () => { - const assertion = await assertionIdentity.getToken('api://AzureADTokenExchange/.default', - { abortSignal: oauthRequest.getStore() }); - if (!usableAccessToken(assertion)) throw new Error('managed identity assertion unavailable'); - return assertion.token; - }, - { authorityHost: 'https://login.microsoftonline.com', retryOptions: { maxRetries: 0 } }, - ); - oauthCredentialConfig = credentialConfig; - } - const accessToken = await oauthCredential.getToken(config.providerScope, { abortSignal: oauthRequest.getStore() }); - if (!usableAccessToken(accessToken)) throw new Error('provider OAuth token unavailable'); - return Object.defineProperty({ mode: 'oauth' }, 'accessToken', { value: accessToken.token }); - } catch { - throw new Error('provider OAuth token unavailable'); - } - }); +function stopProviderCredentialRefresh() { + providerCredentials.close(); } // Status mappings may restrict HTTP success, but cannot turn failed HTTP into Continue. @@ -497,4 +421,6 @@ module.exports = { parseProviderTimeout, isValidProviderUrl, resolveProviderCredential, + startProviderCredentialRefresh, + stopProviderCredentialRefresh, }; diff --git a/javascript/src/functions/refreshingCache.js b/javascript/src/functions/refreshingCache.js new file mode 100644 index 0000000..938a374 --- /dev/null +++ b/javascript/src/functions/refreshingCache.js @@ -0,0 +1,122 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +'use strict'; + +const { inspect } = require('node:util'); + +const unavailable = () => new Error('provider credential unavailable'); + +class RefreshingCache { + constructor(load, { now = Date.now, schedule = setTimeout, cancel = clearTimeout, + random = Math.random, onFailure = () => {} } = {}) { + this.load = load; + this.now = now; + this.schedule = schedule; + this.cancel = cancel; + this.random = random; + this.onFailure = onFailure; + this.entry = null; + this.inFlight = null; + this.timer = null; + this.retryAt = 0; + this.failures = 0; + this.closed = false; + this.controller = null; + } + + [inspect.custom]() { return '[RefreshingCache]'; } + toJSON() { return '[RefreshingCache]'; } + + get() { + if (this.closed) return Promise.reject(unavailable()); + if (this.entry && this.entry.expiresAt > this.now()) { + if (this.entry.refreshAt <= this.now() && this.retryAt <= this.now()) { + this.refreshInBackground(); + } + return Promise.resolve(this.entry.value); + } + if (this.inFlight) return this.inFlight; + if (this.retryAt > this.now()) return Promise.reject(unavailable()); + return this.refresh(); + } + + refresh() { + if (this.closed) return Promise.reject(unavailable()); + if (this.inFlight) return this.inFlight; + if (this.retryAt > this.now()) return Promise.reject(unavailable()); + this.clearTimer(); + this.controller = new AbortController(); + const controller = this.controller; + // Publish the promise before calling a loader that may synchronously reenter the cache. + const operation = Promise.resolve().then(() => this.load(controller.signal)).then((entry) => { + if (this.closed || controller.signal.aborted) throw unavailable(); + if (!entry || !Number.isFinite(entry.expiresAt) || entry.expiresAt <= this.now() + || !Number.isFinite(entry.refreshAt)) throw unavailable(); + this.entry = entry; + this.failures = 0; + this.retryAt = 0; + return entry.value; + }).catch(() => { + if (!this.closed) { + this.failures++; + const backoff = Math.min(60000, 5000 * 2 ** Math.min(this.failures - 1, 4)); + this.retryAt = this.now() + Math.floor(backoff * (1 + this.random() * 0.2)); + this.onFailure(); + } + throw unavailable(); + }).finally(() => { + this.inFlight = null; + this.controller = null; + if (!this.closed) { + const next = this.retryAt || this.entry?.refreshAt; + if (next != null) this.arm(next); + } + }); + this.inFlight = operation; + return operation; + } + + refreshInBackground() { + if (!this.closed && this.retryAt > this.now()) { + this.arm(this.retryAt); + return; + } + // refresh reports its failure; a timer has no request waiting to handle the rejection. + void this.refresh().catch(() => {}); + } + + arm(at) { + this.clearTimer(); + this.timer = this.schedule(() => { + this.timer = null; + this.refreshInBackground(); + }, Math.max(1000, Math.min(2147483647, at - this.now()))); + this.timer?.unref?.(); + } + + clearTimer() { + if (this.timer !== null) this.cancel(this.timer); + this.timer = null; + } + + close() { + this.closed = true; + this.clearTimer(); + this.controller?.abort(); + this.entry = null; + } +} + +function tokenEntry(token, now) { + if (typeof token?.token !== 'string' || !token.token.trim() + || !Number.isFinite(token.expiresOnTimestamp) || token.expiresOnTimestamp <= now + 30000) { + throw unavailable(); + } + const expiresAt = token.expiresOnTimestamp - 30000; + let refreshAt = token.expiresOnTimestamp - 300000; + if (Number.isFinite(token.refreshAfterTimestamp)) refreshAt = Math.min(refreshAt, token.refreshAfterTimestamp); + // An SDK may return its existing token on refresh. Never extend that token's lifetime or spin. + if (refreshAt <= now) refreshAt = now + Math.max(1000, Math.min(60000, (expiresAt - now) / 2)); + return { value: token, expiresAt, refreshAt }; +} + +module.exports = { RefreshingCache, tokenEntry }; diff --git a/javascript/test/credential-cache.test.js b/javascript/test/credential-cache.test.js new file mode 100644 index 0000000..23d2478 --- /dev/null +++ b/javascript/test/credential-cache.test.js @@ -0,0 +1,252 @@ +'use strict'; + +const { test } = require('node:test'); +const assert = require('node:assert/strict'); +const { inspect } = require('node:util'); +const { RefreshingCache, tokenEntry } = require('../src/functions/refreshingCache'); +const { ProviderCredentials } = require('../src/functions/credentials'); +const { ClientAssertionCredential, ManagedIdentityCredential } = require('@azure/identity'); +const { SecretClient } = require('@azure/keyvault-secrets'); +const { readConfig } = require('../src/functions/config'); + +const flush = () => new Promise(setImmediate); +const deferred = () => { + let resolve; + let reject; + const promise = new Promise((yes, no) => { resolve = yes; reject = no; }); + return { promise, resolve, reject }; +}; + +function clock() { + let time = 1700000000000; + const timers = new Set(); + return { + options: { + now: () => time, + random: () => 0, + schedule: (callback, delay) => { + const timer = { at: time + delay, callback, unref() {} }; + timers.add(timer); + return timer; + }, + cancel: (timer) => timers.delete(timer), + }, + get now() { return time; }, + get timerCount() { return timers.size; }, + async advance(ms) { + time += ms; + for (const timer of [...timers]) { + if (timer.at <= time && timers.delete(timer)) timer.callback(); + } + await flush(); + }, + }; +} + +test('concurrent empty-cache readers join one refresh; fresh hits make no loader calls', async () => { + const time = clock(); + const gate = deferred(); + let calls = 0; + const cache = new RefreshingCache(async () => { + calls++; + await gate.promise; + return { value: 'PRIVATE-VALUE', expiresAt: time.now + 300000, refreshAt: time.now + 240000 }; + }, time.options); + try { + const readers = Array.from({ length: 20 }, () => cache.get()); + await flush(); + assert.equal(calls, 1); + gate.resolve(); + assert.deepEqual(await Promise.all(readers), Array(20).fill('PRIVATE-VALUE')); + assert.equal(await cache.get(), 'PRIVATE-VALUE'); + assert.equal(calls, 1); + assert.equal(time.timerCount, 1); + assert.doesNotMatch(inspect(cache) + JSON.stringify(cache), /PRIVATE/); + } finally { cache.close(); } + assert.equal(time.timerCount, 0); +}); + +test('scheduled refresh serves the valid old entry, replaces atomically, and never extends a failed entry', async () => { + const time = clock(); + let next = async () => ({ value: 'first', expiresAt: time.now + 300000, refreshAt: time.now + 240000 }); + const failures = []; + const cache = new RefreshingCache(() => next(), { ...time.options, onFailure: () => failures.push('failed') }); + try { + assert.equal(await cache.get(), 'first'); + const gate = deferred(); + next = () => gate.promise; + await time.advance(240000); + assert.equal(await cache.get(), 'first'); + gate.reject(new Error('PRIVATE-REFRESH-ERROR')); + await flush(); + assert.deepEqual(failures, ['failed']); + assert.equal(await cache.get(), 'first'); + await time.advance(60000); + await assert.rejects(cache.get(), /provider credential unavailable/); + next = async () => ({ value: 'recovered', expiresAt: time.now + 300000, refreshAt: time.now + 240000 }); + await time.advance(10000); + assert.equal(await cache.get(), 'recovered'); + } finally { cache.close(); } +}); + +test('failed cold refresh applies bounded backoff instead of a fetch per request', async () => { + const time = clock(); + let calls = 0; + const cache = new RefreshingCache(async () => { calls++; throw new Error('PRIVATE-FAILURE'); }, time.options); + try { + await assert.rejects(cache.get(), /^Error: provider credential unavailable$/); + for (let i = 0; i < 10; i++) await assert.rejects(cache.get(), /unavailable/); + assert.equal(calls, 1); + await time.advance(4999); + assert.equal(calls, 1); + await time.advance(1); + assert.equal(calls, 2); + await time.advance(9999); + assert.equal(calls, 2); + await time.advance(1); + assert.equal(calls, 3); + } finally { cache.close(); } +}); + +test('stopping a cache cancels refresh and prevents late publication or rescheduling', async () => { + const time = clock(); + const gate = deferred(); + let signal; + const cache = new RefreshingCache(async (abortSignal) => { + signal = abortSignal; + return gate.promise; + }, time.options); + const pending = cache.get(); + await flush(); + cache.close(); + assert.equal(signal.aborted, true); + gate.resolve({ value: 'late', expiresAt: time.now + 300000, refreshAt: time.now + 240000 }); + await assert.rejects(pending, /unavailable/); + await assert.rejects(cache.get(), /unavailable/); + assert.equal(time.timerCount, 0); +}); + +test('token refresh retains original expiry, honors refresh hints, and avoids spinning on SDK cache hits', () => { + const now = 1700000000000; + const token = { token: 'PRIVATE-TOKEN', expiresOnTimestamp: now + 3600000 }; + assert.deepEqual(tokenEntry(token, now), { value: token, expiresAt: now + 3570000, refreshAt: now + 3300000 }); + const repeated = tokenEntry(token, now + 3300000); + assert.equal(repeated.expiresAt, now + 3570000); + assert.equal(repeated.refreshAt, now + 3360000); + assert.equal(tokenEntry({ ...token, refreshAfterTimestamp: now + 600000 }, now).refreshAt, now + 600000); + for (const invalid of [null, { ...token, token: '' }, { ...token, token: ' ' }, + { ...token, expiresOnTimestamp: now + 30000 }, { ...token, expiresOnTimestamp: NaN }, + { ...token, expiresOnTimestamp: Infinity }, { token: 'PRIVATE' }]) { + assert.throws(() => tokenEntry(invalid, now), /unavailable/); + } +}); + +test('Key Vault refresh fetches a parallel credential bundle once and keeps a complete old pair on partial failure', async (t) => { + const time = clock(); + let failIdentity = false; + let version = 1; + const getSecret = t.mock.method(SecretClient.prototype, 'getSecret', async (name) => { + if (failIdentity && name === 'customer-id') throw new Error('PRIVATE-IDENTITY-ERROR'); + return { value: `${name}-${version}` }; + }); + const failures = []; + const manager = new ProviderCredentials({ cacheOptions: time.options, reportFailure: (kind) => failures.push(kind) }); + const auth = { mode: 'apiKey', keyVaultSecretName: 'api-key', identityKeyVaultSecretName: 'customer-id' }; + const config = readConfig({ KEY_VAULT_URL: 'https://unit.vault.azure.net' }); + try { + const results = await Promise.all(Array.from({ length: 10 }, () => manager.resolve(auth, config))); + assert.equal(getSecret.mock.callCount(), 2); + assert.ok(results.every((value) => value.secret === 'api-key-1' && value.identity === 'customer-id-1')); + failIdentity = true; + version = 2; + await time.advance(240000); + const old = await manager.resolve(auth, config); + assert.deepEqual(old, { mode: 'apiKey', secret: 'api-key-1', identity: 'customer-id-1' }); + assert.deepEqual(failures, ['key_vault']); + failIdentity = false; + await time.advance(5000); + assert.deepEqual(await manager.resolve(auth, config), + { mode: 'apiKey', secret: 'api-key-2', identity: 'customer-id-2' }); + assert.equal(getSecret.mock.callCount(), 6); + } finally { manager.close(); } + assert.equal(time.timerCount, 0); +}); + +test('MI and final Entra tokens have independent single-flight refresh and warm requests skip both SDKs', async (t) => { + const time = clock(); + const identity = t.mock.method(ManagedIdentityCredential.prototype, 'getToken', async () => ({ + token: 'PRIVATE-ASSERTION', expiresOnTimestamp: time.now + 3600000, + })); + const provider = t.mock.method(ClientAssertionCredential.prototype, 'getToken', async function () { + await this.getAssertion(); + await this.getAssertion(); + return { token: 'PRIVATE-PROVIDER', expiresOnTimestamp: time.now + 3600000 }; + }); + const manager = new ProviderCredentials({ cacheOptions: time.options }); + const config = readConfig({ + EPP_PROVIDER_TENANT_ID: '11111111-1111-1111-1111-111111111111', + EPP_OUTBOUND_CLIENT_ID: '22222222-2222-2222-2222-222222222222', + EPP_OUTBOUND_MI_CLIENT_ID: '33333333-3333-3333-3333-333333333333', + EPP_PROVIDER_SCOPE: 'api://provider/.default', + }); + try { + const results = await Promise.all(Array.from({ length: 20 }, () => manager.resolve({ mode: 'oauth' }, config))); + assert.ok(results.every((value) => value.accessToken === 'PRIVATE-PROVIDER')); + assert.equal(identity.mock.callCount(), 1); + assert.equal(provider.mock.callCount(), 1); + assert.equal(time.timerCount, 2); + assert.equal(JSON.stringify(results[0]), '{"mode":"oauth"}'); + await manager.resolve({ mode: 'oauth' }, config); + assert.equal(provider.mock.callCount(), 1); + await time.advance(3300000); + assert.equal(identity.mock.callCount(), 2); + assert.equal(provider.mock.callCount(), 2); + await manager.resolve({ mode: 'oauth' }, { ...config, providerScope: 'api://second/.default' }); + assert.equal(identity.mock.callCount(), 2); + assert.equal(provider.mock.callCount(), 3); + assert.equal(provider.mock.calls[0].this, provider.mock.calls[2].this); + await manager.resolve({ mode: 'oauth' }, { ...config, outboundClientId: '44444444-4444-4444-4444-444444444444' }); + assert.equal(identity.mock.callCount(), 3); + assert.notEqual(provider.mock.calls[0].this, provider.mock.calls[3].this); + assert.equal(time.timerCount, 2); + } finally { manager.close(); } + assert.equal(time.timerCount, 0); +}); + +test('bad refresh results never replace a valid Key Vault bundle', async (t) => { + const time = clock(); + const getSecret = t.mock.method(SecretClient.prototype, 'getSecret', async () => ({ value: 'old-key' })); + const manager = new ProviderCredentials({ cacheOptions: time.options, reportFailure: () => {} }); + const auth = { mode: 'apiKey', keyVaultSecretName: 'key' }; + const config = readConfig({ KEY_VAULT_URL: 'https://unit.vault.azure.net' }); + try { + assert.equal((await manager.resolve(auth, config)).secret, 'old-key'); + getSecret.mock.mockImplementation(async () => ({ value: '' })); + await time.advance(240000); + assert.equal((await manager.resolve(auth, config)).secret, 'old-key'); + await time.advance(60000); + await assert.rejects(manager.resolve(auth, config), /unavailable/); + } finally { manager.close(); } +}); + +test('incomplete OAuth reconfiguration clears old timers and never reuses old valid tokens', async (t) => { + const time = clock(); + t.mock.method(ManagedIdentityCredential.prototype, 'getToken', async () => ({ + token: 'assertion', expiresOnTimestamp: time.now + 3600000, + })); + t.mock.method(ClientAssertionCredential.prototype, 'getToken', async function () { + await this.getAssertion(); + return { token: 'token', expiresOnTimestamp: time.now + 3600000 }; + }); + const manager = new ProviderCredentials({ cacheOptions: time.options }); + const config = readConfig({ EPP_PROVIDER_TENANT_ID: 'tenant', EPP_OUTBOUND_CLIENT_ID: 'app', + EPP_OUTBOUND_MI_CLIENT_ID: 'identity', EPP_PROVIDER_SCOPE: 'scope' }); + try { + await manager.resolve({ mode: 'oauth' }, config); + assert.equal(time.timerCount, 2); + await assert.rejects(manager.resolve({ mode: 'oauth' }, { ...config, providerScope: '' }), /unavailable/); + assert.equal(time.timerCount, 0); + await manager.resolve({ mode: 'oauth' }, config); + assert.equal(time.timerCount, 2); + } finally { manager.close(); } +}); diff --git a/javascript/test/credential-sdk.test.js b/javascript/test/credential-sdk.test.js new file mode 100644 index 0000000..fba2d73 --- /dev/null +++ b/javascript/test/credential-sdk.test.js @@ -0,0 +1,108 @@ +'use strict'; + +const { test, mock, beforeEach } = require('node:test'); +const assert = require('node:assert/strict'); +const Module = require('node:module'); +const crypto = require('node:crypto'); +const { setTimeout: delay } = require('node:timers/promises'); +const identity = require('@azure/identity'); +const { createHttpHeaders } = require('@azure/core-rest-pipeline'); +const { readConfig } = require('../src/functions/config'); + +let state; +const originalLoad = Module._load; +const load = mock.method(Module, '_load', function (name, ...args) { + if (name !== '@azure/identity') return originalLoad.call(this, name, ...args); + return { + ...identity, + ManagedIdentityCredential: class extends identity.ManagedIdentityCredential { + constructor(...args) { super(...args); this.testState = state; } + async getToken() { + this.testState.miCalls++; + return { token: 'PRIVATE-ASSERTION', expiresOnTimestamp: Date.now() + 3600000 }; + } + }, + ClientAssertionCredential: class extends identity.ClientAssertionCredential { + constructor(tenant, client, assertion, options) { + super(tenant, client, assertion, { ...options, httpClient: state.transport }); + } + }, + }; +}); +let ProviderCredentials; +try { ({ ProviderCredentials } = require('../src/functions/credentials')); } +finally { load.mock.restore(); } + +beforeEach(() => { + state = { miCalls: 0, tokenCalls: 0, wait: 5, abortObserved: false }; + const current = state; + state.transport = { + async sendRequest(request) { + const url = new URL(request.url); + assert.equal(request.method, 'POST'); + assert.ok(url.pathname.endsWith('/oauth2/v2.0/token')); + current.tokenCalls++; + try { await delay(current.wait, null, { signal: request.abortSignal }); } + catch (error) { current.abortObserved = true; throw error; } + return { request, status: 200, headers: createHttpHeaders({ 'content-type': 'application/json' }), + bodyAsText: JSON.stringify({ access_token: 'PRIVATE-TOKEN', token_type: 'Bearer', + expires_in: 3600, scope: 'api://provider/.default' }) }; + }, + }; +}); + +function config() { + return readConfig({ + EPP_PROVIDER_TENANT_ID: '11111111-1111-1111-1111-111111111111', + EPP_OUTBOUND_CLIENT_ID: crypto.randomUUID(), + EPP_OUTBOUND_MI_CLIENT_ID: '33333333-3333-3333-3333-333333333333', + EPP_PROVIDER_SCOPE: 'api://provider/.default', + }); +} + +test('real SDK sees one initial Entra exchange and none on concurrent or later warm requests', async () => { + const manager = new ProviderCredentials(); + const settings = config(); + try { + const tokens = await Promise.all(Array.from({ length: 20 }, () => manager.resolve({ mode: 'oauth' }, settings))); + assert.ok(tokens.every((token) => token.accessToken === 'PRIVATE-TOKEN')); + assert.equal(state.tokenCalls, 1); + assert.equal(state.miCalls, 1); + await manager.resolve({ mode: 'oauth' }, settings); + assert.equal(state.tokenCalls, 1); + assert.equal(state.miCalls, 1); + } finally { manager.close(); } +}); + +test('the cache-owned acquisition budget aborts actual SDK transport and does not cache a late token', async () => { + const failures = []; + const manager = new ProviderCredentials({ reportFailure: (kind) => failures.push(kind) }); + const settings = config(); + state.wait = 10000; + try { + await assert.rejects(manager.resolve({ mode: 'oauth' }, settings), /^Error: provider OAuth token unavailable$/); + await new Promise(setImmediate); + assert.equal(state.abortObserved, true); + assert.deepEqual(failures, ['provider_token']); + assert.equal(state.tokenCalls, 1); + await assert.rejects(manager.resolve({ mode: 'oauth' }, settings), /unavailable/); + assert.equal(state.tokenCalls, 1); + } finally { manager.close(); } +}); + +test('changing configuration cancels old work without publishing its token into the replacement cache', async () => { + const manager = new ProviderCredentials({ reportFailure: () => {} }); + const settings = config(); + state.wait = 10000; + const first = manager.resolve({ mode: 'oauth' }, settings); + const firstRejected = assert.rejects(first, /unavailable/); + await delay(20); + state.wait = 5; + const result = await manager.resolve({ mode: 'oauth' }, { ...settings, outboundClientId: crypto.randomUUID() }); + await firstRejected; + try { + assert.equal(result.accessToken, 'PRIVATE-TOKEN'); + assert.equal(state.abortObserved, true); + assert.equal(state.tokenCalls, 2); + } finally { manager.close(); } +}); diff --git a/javascript/test/dispatch.test.js b/javascript/test/dispatch.test.js index dfa0a47..28977fd 100644 --- a/javascript/test/dispatch.test.js +++ b/javascript/test/dispatch.test.js @@ -1,6 +1,6 @@ 'use strict'; -const { test } = require('node:test'); +const { test, afterEach } = require('node:test'); const assert = require('node:assert/strict'); const { ClientAssertionCredential, ManagedIdentityCredential } = require('@azure/identity'); const { SecretClient } = require('@azure/keyvault-secrets'); @@ -12,7 +12,9 @@ const { AzureLogger } = require('@azure/logger'); const { dispatchOtp, getProvider, resolveOutcome, outcomeToHttpStatus, parseEnvelope, parseProviderTimeout, isValidProviderUrl, contextToDispatch, resolveProviderCredential, + stopProviderCredentialRefresh, } = require('../src/functions/dispatch'); +afterEach(stopProviderCredentialRefresh); const dispatch = { destination: '+15551234567', message: ' Your code is 918273.\n', channel: 'sms', messageId: 'message-id', correlationId: 'correlation-id' }; const input = { channel: 'sms', endpoint: 'https://provider.example', dispatch, @@ -258,7 +260,7 @@ test('missing API-key or OAuth settings and an unsafe final voice URL make zero assert.equal(getSecret.mock.calls.at(-1).this.vaultUrl, nextConfig.keyVaultUrl); } assert.equal(getSecret.mock.callCount(), calls + 2); - assert.equal(new Set(getSecret.mock.calls.map((call) => call.this)).size, 3); + assert.equal(new Set(getSecret.mock.calls.map((call) => call.this)).size, 4); assert.equal(fetchMock.mock.callCount(), 0); }); @@ -283,7 +285,7 @@ test('Soprano OAuth reuses setup identities and selected scope with private boun assert.equal(identityToken.mock.calls[0].arguments[0], 'api://AzureADTokenExchange/.default'); const signal = providerToken.mock.calls[0].arguments[1].abortSignal; assert.ok(signal instanceof AbortSignal); - assert.equal(identityToken.mock.calls[0].arguments[1].abortSignal, signal); + assert.ok(identityToken.mock.calls[0].arguments[1].abortSignal instanceof AbortSignal); await resolveProviderCredential({ mode: 'oauth' }, { ...config, providerScope: 'api://another/.default' }); assert.equal(providerToken.mock.calls[1].this, providerToken.mock.calls[0].this); assert.equal(providerToken.mock.calls[1].arguments[0], 'api://another/.default'); @@ -296,6 +298,7 @@ test('Soprano OAuth reuses setup identities and selected scope with private boun for (const invalid of [null, { token: '' }, { token: ' ' }, { token: false }, { token: 'stale', expiresOnTimestamp: Date.now() + 10000 }, { token: 'missing-expiry' }]) { const method = stage === 'token' ? providerToken : identityToken; + stopProviderCredentialRefresh(); method.mock.mockImplementation(async () => invalid); if (stage === 'assertion') providerToken.mock.mockImplementation(async function () { await this.getAssertion(); diff --git a/javascript/test/sendotp.test.js b/javascript/test/sendotp.test.js index 14454af..72770c4 100644 --- a/javascript/test/sendotp.test.js +++ b/javascript/test/sendotp.test.js @@ -8,16 +8,20 @@ const { CompactEncrypt } = require('jose'); const { ClientAssertionCredential, ManagedIdentityCredential } = require('@azure/identity'); const { SecretClient } = require('@azure/keyvault-secrets'); const fixtures = require('../../tests/fixtures/contract.json'); -const { getProvider } = require('../src/functions/dispatch'); +const { getProvider, stopProviderCredentialRefresh } = require('../src/functions/dispatch'); const { RequestLog } = require('../src/functions/requestLog'); // Capture the real handler; keys stay in memory and all external I/O is mocked. const { publicKey, privateKey } = crypto.generateKeyPairSync('rsa', { modulusLength: 2048 }); let handler; +let startHook; +let stopHook; const originalLoad = Module._load; const registration = mock.method(Module, '_load', function (name, ...args) { if (name === '@azure/functions') { - return { app: { http: (_name, options) => { handler = options.handler; } } }; + return { app: { hook: { appStart: (callback) => { startHook = callback; }, + appTerminate: (callback) => { stopHook = callback; } }, + http: (_name, options) => { handler = options.handler; } } }; } return originalLoad.call(this, name, ...args); }); @@ -39,6 +43,7 @@ let warnings; let records; let getToken; beforeEach(() => { + stopProviderCredentialRefresh(); savedEnv = Object.fromEntries(envKeys.map((key) => [key, process.env[key]])); for (const key of envKeys) delete process.env[key]; Object.assign(process.env, { EPP_LOG_PLAINTEXT: 'true', @@ -61,6 +66,7 @@ beforeEach(() => { text: async () => JSON.stringify({ status: 'ENROUTE', id: 'provider-reference-id', description: 'PRIVATE-STATUS' }) })); }); afterEach(() => { + stopProviderCredentialRefresh(); mock.restoreAll(); for (const [key, value] of Object.entries(savedEnv)) { if (value === undefined) delete process.env[key]; @@ -127,6 +133,29 @@ function assertFailure(result, status, error = 'provider_delivery_failed') { assert.doesNotMatch(JSON.stringify(result.jsonBody), /PRIVATE|accepted/); } +test('worker startup preloads credentials without delivery and leaves evaluation independent', async () => { + await startHook(); + assert.equal(getToken.mock.callCount(), 1); + assert.equal(fetchMock.mock.callCount(), 0); + assert.equal(getSecret.mock.callCount(), 0); + const response = await invoke(await envelope({ mode: 2 })); + assert.equal(response.status, 200); + assert.equal(getToken.mock.callCount(), 1); + assert.equal(fetchMock.mock.callCount(), 0); + await invoke(await envelope()); + assert.equal(getToken.mock.callCount(), 1); + assert.equal(fetchMock.mock.callCount(), 1); + stopHook(); +}); + +test('worker startup without a configured provider does not acquire any credentials', async () => { + delete process.env.EPP_PROVIDER_NAME; + await startHook(); + assert.equal(getToken.mock.callCount(), 0); + assert.equal(getSecret.mock.callCount(), 0); + assert.equal(fetchMock.mock.callCount(), 0); +}); + test('shared invalid requests return matching safe reasons before provider I/O', async () => { const valid = { type: 'microsoft.mfa.otpDeliver.v1', channel: 1, mode: 1, encryptedDeliveryContext: 'unused' }; for (const fixture of fixtures.badRequests) { diff --git a/python/README.md b/python/README.md index 8702ead..b0f6e92 100644 --- a/python/README.md +++ b/python/README.md @@ -90,15 +90,24 @@ six-digit numeric run that is not part of a longer number and repeats the comple ## Source +Worker initialization starts background credential preparation when a provider is configured. +Key Vault bundles, managed-identity assertions and final Entra tokens use separate process-local +caches with daemon refresh timers. A caller can stop waiting without cancelling shared retrieval; +the HTTP SDK still uses connect/read inactivity timeouts, not a total transport deadline. `atexit` +stops scheduled work and prevents late cache publication. See the +[refresh contract](../docs/CONTRACT.md#credential-caching-and-refresh). Evaluation handling stays +independent; leave the provider unset for local evaluation without background credential acquisition. + | Source | Purpose | |---|---| | [function_app.py](function_app.py) | HTTP handler and adapter registration | | [src/config.py](src/config.py) | Shared deployment settings | | [src/models.py](src/models.py) | Envelope, delivery-context, dispatch and normalized `ParsedResponse` dataclasses | | [src/dispatch.py](src/dispatch.py) | Boundary validation, JWE, provider registry and outcome mapping | +| [src/credentials.py](src/credentials.py), [refreshing_cache.py](src/refreshing_cache.py) | Provider credential bundles, independent token caches and scheduled refresh | | [src/request_log.py](src/request_log.py) | Request-scoped [service events and summaries](../docs/CONTRACT.md#application-logs) with explicit ID sources | | [src/providers/](src/providers/) | Adapter manifests and API-specific implementations | -| [src/secrets.py](src/secrets.py) | Cached Key Vault access via managed identity | +| [src/secrets.py](src/secrets.py) | Key Vault transport; bundle caching belongs to the credential manager | Add and register an adapter without adding provider-specific branches to the shared pipeline. Return `ParsedResponse` from `parse_response` using named fields; the engine reads attributes such as diff --git a/python/function_app.py b/python/function_app.py index 1bdb30e..4733702 100644 --- a/python/function_app.py +++ b/python/function_app.py @@ -1,6 +1,8 @@ +import atexit import json import os import uuid +from threading import Thread import azure.functions as func @@ -106,3 +108,9 @@ def respond(status, body): return respond(500, {"error": "delivery_failed", "correlationId": correlation_id, "requestId": request_id}) finally: log.complete(http_status) + + +# Each worker owns its own memory cache; a timer trigger would only warm one worker. +atexit.register(_engine.close) +if os.environ.get("EPP_PROVIDER_NAME", "").strip(): + Thread(target=_engine.start_credential_refresh, daemon=True).start() diff --git a/python/src/credentials.py b/python/src/credentials.py new file mode 100644 index 0000000..f36c23e --- /dev/null +++ b/python/src/credentials.py @@ -0,0 +1,128 @@ +import json +import logging +import threading +import time +from concurrent.futures import ThreadPoolExecutor +from contextvars import ContextVar + +from azure.identity import ClientAssertionCredential, ManagedIdentityCredential + +from .refreshing_cache import CacheEntry, RefreshingCache, token_entry + +_acquiring = ContextVar("epp_credential_acquisition", default=False) + + +class _CredentialLogFilter(logging.Filter): + def filter(self, record): + return not (_acquiring.get() and record.name.startswith(("azure.identity", "azure.core", "msal"))) + + +_log_filter = _CredentialLogFilter() + + +def report_refresh_failure(kind): + logging.warning("%s", json.dumps({"logType": "service", "eventName": "credential_refresh_failed", + "cacheKind": kind, "failureReason": "credential_unavailable"})) + + +def _private_acquisition(load): + for logger in (logging.getLogger(), *logging.Logger.manager.loggerDict.copy().values()): + if isinstance(logger, logging.Logger): + for handler in logger.handlers: + if _log_filter not in handler.filters: + handler.addFilter(_log_filter) + context = _acquiring.set(True) + try: + return load() + finally: + _acquiring.reset(context) + + +class ProviderCredentials: + def __init__(self, secrets, *, cache_options=None, report_failure=report_refresh_failure): + self._secrets = secrets + self._options = cache_options or {} + self._clock = self._options.get("clock", time.time) + self._report_failure = report_failure + self._lock = threading.RLock() + self._key = None + self._state = None + + def _cache(self, kind, load): + return RefreshingCache(lambda: _private_acquisition(load), **{ + **self._options, "on_failure": lambda: self._report_failure(kind), + }) + + def resolve(self, auth, config): + mode = auth.get("mode") + if mode == "oauth" and not all((config.provider_tenant_id, config.provider_scope, + config.outbound_client_id, config.outbound_managed_identity_client_id)): + self.close() + raise ValueError("provider OAuth token unavailable") + if mode not in ("apiKey", "oauth"): + self.close() + raise ValueError("provider credential unavailable") + key = ((mode, config.env.get("KEY_VAULT_URL"), config.env.get("AZURE_CLIENT_ID"), + auth.get("key_vault_secret_name"), auth.get("identity_key_vault_secret_name")) if mode == "apiKey" else + (mode, config.provider_tenant_id, config.outbound_client_id, config.outbound_managed_identity_client_id)) + with self._lock: + if self._key != key: + self.close() + self._state = self._api_key_state(auth) if mode == "apiKey" else _private_acquisition(lambda: self._oauth_state(config)) + self._key = key + state = self._state + if mode == "apiKey": + cache = state["bundle"] + else: + scope = config.provider_scope + if scope not in state["tokens"]: + state["tokens"][scope] = self._cache("provider_token", lambda: + token_entry(state["credential"].get_token(scope, logging_enable=False), self._clock())) + cache = state["tokens"][scope] + try: + value = cache.get() + except Exception: + raise ValueError("provider OAuth token unavailable" if mode == "oauth" else "provider credential unavailable") from None + return dict(value) if mode == "apiKey" else {"mode": "oauth", "access_token": value.token} + + def _api_key_state(self, auth): + def load(): + key_name = auth.get("key_vault_secret_name") + identity_name = auth.get("identity_key_vault_secret_name") + if not key_name: + raise ValueError("provider credential unavailable") + # Both secrets form one snapshot; do not publish a partial rotation. + with ThreadPoolExecutor(max_workers=2) as pool: + key_future = pool.submit(_private_acquisition, lambda: self._secrets.resolve(key_name)) + identity_future = pool.submit(_private_acquisition, lambda: self._secrets.resolve(identity_name)) if identity_name else None + secret = key_future.result() + identity = identity_future.result() if identity_future else "" + if not isinstance(secret, str) or not secret.strip() or (identity_name and ( + not isinstance(identity, str) or not identity.strip())): + raise ValueError("provider credential unavailable") + now = self._clock() + return CacheEntry({"mode": "apiKey", "secret": secret, "identity": identity}, now + 300, now + 240) + return {"bundle": self._cache("key_vault", load)} + + def _oauth_state(self, config): + identity = ManagedIdentityCredential(client_id=config.outbound_managed_identity_client_id, + retry_total=0, connection_timeout=2.5, read_timeout=2.5, logging_enable=False) + assertion = self._cache("managed_identity", lambda: + token_entry(identity.get_token("api://AzureADTokenExchange/.default", logging_enable=False), self._clock())) + credential = ClientAssertionCredential( + tenant_id=config.provider_tenant_id, client_id=config.outbound_client_id, + func=lambda: assertion.get().token, authority="https://login.microsoftonline.com", + retry_total=0, connection_timeout=2.5, read_timeout=2.5, logging_enable=False) + return {"assertion": assertion, "credential": credential, "tokens": {}} + + def close(self): + with self._lock: + if self._state: + if "bundle" in self._state: + self._state["bundle"].close() + else: + self._state["assertion"].close() + for cache in self._state["tokens"].values(): + cache.close() + self._state = None + self._key = None diff --git a/python/src/dispatch.py b/python/src/dispatch.py index f91191d..d8f637a 100644 --- a/python/src/dispatch.py +++ b/python/src/dispatch.py @@ -3,20 +3,16 @@ import base64 import json import logging -import math import os -import time -from contextvars import ContextVar -from threading import Lock from urllib.parse import urlsplit import requests -from azure.identity import ClientAssertionCredential, ManagedIdentityCredential from jwcrypto import jwe as jwe_module from jwcrypto import jwk from urllib3.exceptions import ReadTimeoutError from .config import read_config +from .credentials import ProviderCredentials, report_refresh_failure from .models import DeliveryContext, DispatchRequest, Envelope, ParsedResponse, TextToVoice from .request_log import RequestLog @@ -28,24 +24,6 @@ BLOCK = "Block" STEP_UP = "StepUp" -_oauth_request = ContextVar("provider_oauth_request", default=False) - - -class _OAuthLogFilter(logging.Filter): - def filter(self, record): - return not (_oauth_request.get() and record.name.startswith(("azure.identity", "azure.core", "msal"))) - - -_oauth_log_filter = _OAuthLogFilter() - - -def _usable_access_token(value): - token = getattr(value, "token", None) - expiry = getattr(value, "expires_on", None) - return (isinstance(token, str) and bool(token.strip()) and type(expiry) in (int, float) - and math.isfinite(expiry) and expiry > time.time() + 30) - - def resolve_outcome(manifest, parsed: ParsedResponse): mapping = manifest["response_mapping"] key = parsed.provider_status_name or parsed.provider_status_code @@ -267,9 +245,23 @@ def __init__(self, registry, secrets, env=None): self.registry = registry self.secrets = secrets self.env = env if env is not None else os.environ - self._oauth_credential = None - self._oauth_credential_config = None - self._oauth_lock = Lock() + self._credentials = ProviderCredentials(secrets) + + def start_credential_refresh(self): + config = read_config(self.env) + if not config.provider_name: + return + adapter = self.registry.get(config.provider_name) + if adapter is None or (config.provider_auth_mode and config.provider_auth_mode != adapter.manifest["auth"]["mode"]): + report_refresh_failure("configuration") + return + try: + self._resolve_credential(adapter.manifest["auth"], config) + except Exception: + report_refresh_failure("initialization") + + def close(self): + self._credentials.close() def dispatch(self, dispatch, request_id, log: RequestLog | None = None): def failure(status, stage, reason, body): @@ -420,51 +412,7 @@ def failure(status, stage, reason, body): log.service("provider_response_cleanup_failed", level=logging.WARNING) def _resolve_credential(self, auth, config): - if auth.get("mode") == "apiKey": - secret = self.secrets.resolve(auth.get("key_vault_secret_name")) - identity = self.secrets.resolve(auth.get("identity_key_vault_secret_name")) if auth.get("identity_key_vault_secret_name") else "" - return {"mode": "apiKey", "secret": secret, "identity": identity} - if auth.get("mode") != "oauth" or not all(( - config.provider_tenant_id, config.provider_scope, - config.outbound_client_id, config.outbound_managed_identity_client_id, - )): - raise ValueError("unsupported or incomplete provider authentication") - for logger in (logging.getLogger(), *logging.Logger.manager.loggerDict.copy().values()): - if isinstance(logger, logging.Logger): - for handler in logger.handlers: - if _oauth_log_filter not in handler.filters: - handler.addFilter(_oauth_log_filter) - context_token = _oauth_request.set(True) - try: - credential_config = ( - config.provider_tenant_id, config.outbound_client_id, config.outbound_managed_identity_client_id, - ) - with self._oauth_lock: - if self._oauth_credential is None or self._oauth_credential_config != credential_config: - assertion_identity = ManagedIdentityCredential(client_id=config.outbound_managed_identity_client_id, - retry_total=0, connection_timeout=2.5, read_timeout=2.5, logging_enable=False) - - def get_assertion(): - token = assertion_identity.get_token("api://AzureADTokenExchange/.default", logging_enable=False) - if not _usable_access_token(token): - raise ValueError("managed identity assertion unavailable") - return token.token - - self._oauth_credential = ClientAssertionCredential( - tenant_id=config.provider_tenant_id, client_id=config.outbound_client_id, func=get_assertion, - authority="https://login.microsoftonline.com", retry_total=0, - connection_timeout=2.5, read_timeout=2.5, logging_enable=False, - ) - self._oauth_credential_config = credential_config - credential = self._oauth_credential - token = credential.get_token(config.provider_scope, logging_enable=False) - if not _usable_access_token(token): - raise ValueError("provider OAuth token unavailable") - return {"mode": "oauth", "access_token": token.token} - except Exception: - raise ValueError("provider OAuth token unavailable") from None - finally: - _oauth_request.reset(context_token) + return self._credentials.resolve(auth, config) def _fail_body(self, provider, channel, reason, dispatch, request_id): return {"status": "failed", "outcome": "Fail", "provider": provider, "channel": channel, "reason": reason, "correlationId": dispatch.correlation_id, "messageId": dispatch.message_id, "requestId": request_id} diff --git a/python/src/refreshing_cache.py b/python/src/refreshing_cache.py new file mode 100644 index 0000000..bd8f8f2 --- /dev/null +++ b/python/src/refreshing_cache.py @@ -0,0 +1,136 @@ +import math +import random +import threading +import time +from concurrent.futures import Future, TimeoutError +from dataclasses import dataclass + + +@dataclass(repr=False) +class CacheEntry: + value: object + expires_at: float + refresh_at: float + + +def _schedule(delay, callback): + timer = threading.Timer(delay, callback) + timer.daemon = True + timer.start() + return timer + + +class RefreshingCache: + def __init__(self, load, *, clock=time.time, schedule=_schedule, jitter=random.random, + on_failure=lambda: None, wait_timeout=2.5): + self._load = load + self._clock = clock + self._schedule = schedule + self._jitter = jitter + self._on_failure = on_failure + self._wait_timeout = wait_timeout + self._lock = threading.RLock() + self._entry = None + self._inflight = None + self._timer = None + self._retry_at = 0 + self._failures = 0 + self._closed = False + + def get(self): + with self._lock: + if self._closed: + raise ValueError("provider credential unavailable") + now = self._clock() + if self._entry and self._entry.expires_at > now: + if self._entry.refresh_at <= now and self._retry_at <= now: + self._begin_refresh() + return self._entry.value + if self._inflight is None and self._retry_at > now: + raise ValueError("provider credential unavailable") + future = self._begin_refresh() + try: + return future.result(timeout=self._wait_timeout) + except TimeoutError: + # A waiter does not cancel the shared refresh needed by other requests. + raise ValueError("provider credential unavailable") from None + + def refresh(self): + with self._lock: + if self._closed or self._retry_at > self._clock(): + raise ValueError("provider credential unavailable") + return self._begin_refresh() + + def _begin_refresh(self): + if self._inflight is not None: + return self._inflight + if self._timer is not None: + self._timer.cancel() + self._timer = None + future = Future() + self._inflight = future + threading.Thread(target=self._run_refresh, args=(future,), daemon=True).start() + return future + + def _run_refresh(self, future): + entry = None + failed = False + try: + entry = self._load() + if (not isinstance(entry, CacheEntry) or not math.isfinite(entry.expires_at) + or entry.expires_at <= self._clock() or not math.isfinite(entry.refresh_at)): + raise ValueError("provider credential unavailable") + except Exception: + failed = True + with self._lock: + if self._closed: + if not future.done(): + future.set_exception(ValueError("provider credential unavailable")) + self._inflight = None + return + if failed: + self._failures += 1 + backoff = min(60, 5 * 2 ** min(self._failures - 1, 4)) + self._retry_at = self._clock() + backoff * (1 + self._jitter() * 0.2) + self._on_failure() + else: + self._entry = entry + self._failures = 0 + self._retry_at = 0 + self._inflight = None + next_refresh = self._retry_at if failed else entry.refresh_at + self._timer = self._schedule(max(1, next_refresh - self._clock()), self._scheduled_refresh) + if failed: + future.set_exception(ValueError("provider credential unavailable")) + else: + future.set_result(entry.value) + + def _scheduled_refresh(self): + with self._lock: + self._timer = None + if not self._closed: + self._begin_refresh() + + def close(self): + with self._lock: + self._closed = True + if self._timer is not None: + self._timer.cancel() + self._timer = None + self._entry = None + + +def token_entry(token, now): + value = getattr(token, "token", None) + expiry = getattr(token, "expires_on", None) + if (not isinstance(value, str) or not value.strip() or type(expiry) not in (int, float) + or not math.isfinite(expiry) or expiry <= now + 30): + raise ValueError("provider credential unavailable") + expires_at = expiry - 30 + refresh_at = expiry - 300 + hint = getattr(token, "refresh_on", None) + if type(hint) in (int, float) and math.isfinite(hint): + refresh_at = min(refresh_at, hint) + if refresh_at <= now: + refresh_at = now + max(1, min(60, (expires_at - now) / 2)) + return CacheEntry(token, expires_at, refresh_at) diff --git a/python/src/secrets.py b/python/src/secrets.py index 036cede..ca06f2b 100644 --- a/python/src/secrets.py +++ b/python/src/secrets.py @@ -1,40 +1,32 @@ import os -import time +from threading import Lock from azure.identity import ManagedIdentityCredential from azure.keyvault.secrets import SecretClient -CACHE_TTL_SECONDS = 5 * 60 - - class SecretResolver: - def __init__(self): + def __init__(self, env=None): + self._env = env if env is not None else os.environ self._client = None - self._cache = {} + self._client_key = None + self._lock = Lock() def _get_client(self): - if self._client is None: - vault_url = os.environ.get("KEY_VAULT_URL") - if not vault_url: - return None - client_id = os.environ.get("AZURE_CLIENT_ID") - credential = ( - ManagedIdentityCredential(client_id=client_id) - if client_id - else ManagedIdentityCredential() - ) - self._client = SecretClient(vault_url=vault_url, credential=credential) - return self._client + vault_url = self._env.get("KEY_VAULT_URL") + client_id = self._env.get("AZURE_CLIENT_ID") + if not vault_url: + raise RuntimeError("KEY_VAULT_URL not set") + with self._lock: + if self._client_key != (vault_url, client_id): + credential = ManagedIdentityCredential(client_id=client_id, logging_enable=False, + retry_total=0, connection_timeout=2.5, read_timeout=2.5) + self._client = SecretClient(vault_url=vault_url, credential=credential, + retry_total=0, connection_timeout=2.5, read_timeout=2.5, logging_enable=False) + self._client_key = (vault_url, client_id) + return self._client def resolve(self, secret_name): if not secret_name: return "" - cached = self._cache.get(secret_name) - if cached and cached[1] > time.time(): - return cached[0] - client = self._get_client() - if client is None: - raise RuntimeError("KEY_VAULT_URL not set") - value = client.get_secret(secret_name).value or "" - self._cache[secret_name] = (value, time.time() + CACHE_TTL_SECONDS) - return value + # ProviderCredentials caches the complete credential bundle and owns refresh. + return self._get_client().get_secret(secret_name, logging_enable=False).value or "" diff --git a/python/tests/test_credential_cache.py b/python/tests/test_credential_cache.py new file mode 100644 index 0000000..4082b26 --- /dev/null +++ b/python/tests/test_credential_cache.py @@ -0,0 +1,294 @@ +import time +from concurrent.futures import ThreadPoolExecutor +from threading import Event, Lock +from types import SimpleNamespace +from unittest.mock import Mock + +import pytest + +import src.credentials as credentials_module +from src.config import read_config +from src.credentials import ProviderCredentials +from src.dispatch import DispatchEngine, ProviderRegistry +from src.providers.telesign import TelesignProvider +from src.refreshing_cache import CacheEntry, RefreshingCache, token_entry + + +def wait_until(predicate): + deadline = time.monotonic() + 3 + while not predicate(): + assert time.monotonic() < deadline, "background refresh did not reach the expected state" + time.sleep(0.005) + + +class Clock: + def __init__(self): + self.now = 1700000000.0 + self.timers = [] + self.lock = Lock() + + def schedule(self, seconds, callback): + timer = SimpleNamespace(at=self.now + seconds, callback=callback, canceled=False) + timer.cancel = lambda: setattr(timer, "canceled", True) + with self.lock: + self.timers.append(timer) + return timer + + @property + def options(self): + return {"clock": lambda: self.now, "schedule": self.schedule, "jitter": lambda: 0} + + @property + def timer_count(self): + with self.lock: + return sum(not timer.canceled for timer in self.timers) + + def advance(self, seconds): + self.now += seconds + with self.lock: + due = [timer for timer in self.timers if not timer.canceled and timer.at <= self.now] + for timer in due: + timer.canceled = True + for timer in due: + timer.callback() + + +def test_single_flight_and_scheduled_refresh_do_not_block_valid_cached_values(): + clock = Clock() + first_started, release_first, refresh_started, release_refresh = Event(), Event(), Event(), Event() + calls = [] + + def load(): + calls.append(clock.now) + if len(calls) == 1: + first_started.set() + assert release_first.wait(3) + value = "PRIVATE-FIRST" + else: + refresh_started.set() + assert release_refresh.wait(3) + value = "PRIVATE-NEXT" + return CacheEntry(value, clock.now + 300, clock.now + 240) + + cache = RefreshingCache(load, **clock.options) + try: + with ThreadPoolExecutor(max_workers=10) as pool: + pending = [pool.submit(cache.get) for _ in range(10)] + assert first_started.wait(3) + assert len(calls) == 1 + release_first.set() + assert all(item.result(3) == "PRIVATE-FIRST" for item in pending) + clock.advance(240) + assert refresh_started.wait(3) + assert cache.get() == "PRIVATE-FIRST" + release_refresh.set() + wait_until(lambda: cache.get() == "PRIVATE-NEXT") + assert len(calls) == 2 + assert clock.timer_count == 1 + assert "PRIVATE" not in repr(cache) + finally: + release_first.set() + release_refresh.set() + cache.close() + assert clock.timer_count == 0 + + +def test_failed_refresh_keeps_only_unexpired_entry_and_uses_backoff(): + clock = Clock() + calls, failures = [], [] + fail = False + + def load(): + calls.append(clock.now) + if fail: + raise ValueError("PRIVATE-ERROR") + return CacheEntry("value", clock.now + 300, clock.now + 240) + + cache = RefreshingCache(load, **clock.options, on_failure=lambda: failures.append("failed")) + try: + assert cache.get() == "value" + fail = True + clock.advance(240) + wait_until(lambda: len(failures) == 1) + assert cache.get() == "value" + for _ in range(10): + assert cache.get() == "value" + assert len(calls) == 2 + clock.advance(60) + wait_until(lambda: len(failures) == 2) + with pytest.raises(ValueError, match="provider credential unavailable"): + cache.get() + fail = False + clock.advance(10) + wait_until(lambda: len(calls) == 4 and cache._inflight is None) + assert cache.get() == "value" + finally: + cache.close() + + +def test_timed_out_waiter_does_not_cancel_refresh_and_shutdown_blocks_late_publication(): + clock = Clock() + started, release = Event(), Event() + + def load(): + started.set() + assert release.wait(3) + return CacheEntry("late", clock.now + 300, clock.now + 240) + + cache = RefreshingCache(load, **clock.options, wait_timeout=0.02) + try: + with pytest.raises(ValueError, match="unavailable"): + cache.get() + assert started.is_set() + future = cache.refresh() + release.set() + assert future.result(3) == "late" + assert cache.get() == "late" + finally: + release.set() + cache.close() + with pytest.raises(ValueError, match="unavailable"): + cache.get() + assert clock.timer_count == 0 + + release.clear() + stopped = RefreshingCache(load, **clock.options) + pending = stopped.refresh() + stopped.close() + release.set() + with pytest.raises(ValueError, match="unavailable"): + pending.result(3) + assert clock.timer_count == 0 + + +def test_token_entry_preserves_real_expiry_and_never_spins_on_sdk_cached_return(): + now = 1700000000 + token = SimpleNamespace(token="PRIVATE-TOKEN", expires_on=now + 3600) + entry = token_entry(token, now) + assert (entry.expires_at, entry.refresh_at) == (now + 3570, now + 3300) + repeated = token_entry(token, now + 3300) + assert repeated.expires_at == now + 3570 + assert repeated.refresh_at == now + 3360 + token.refresh_on = now + 600 + assert token_entry(token, now).refresh_at == now + 600 + for invalid in (None, SimpleNamespace(token=""), SimpleNamespace(token=" "), + SimpleNamespace(token="PRIVATE", expires_on=now + 30), + SimpleNamespace(token="PRIVATE", expires_on=float("inf")), + SimpleNamespace(token="PRIVATE", expires_on=True)): + with pytest.raises(ValueError, match="unavailable"): + token_entry(invalid, now) + + +def oauth_config(): + return read_config({ + "EPP_PROVIDER_TENANT_ID": "tenant", "EPP_PROVIDER_SCOPE": "api://provider/.default", + "EPP_OUTBOUND_CLIENT_ID": "app", "EPP_OUTBOUND_MI_CLIENT_ID": "identity", + }) + + +def test_both_mi_and_provider_token_caches_refresh_independently_and_skip_warm_sdk_calls(monkeypatch): + clock = Clock() + identity = Mock(get_token=Mock(side_effect=lambda *args, **kwargs: + SimpleNamespace(token="PRIVATE-ASSERTION", expires_on=clock.now + 3600))) + monkeypatch.setattr(credentials_module, "ManagedIdentityCredential", Mock(return_value=identity)) + clients = [] + + def create(**kwargs): + def get_token(*args, **options): + assert kwargs["func"]() == "PRIVATE-ASSERTION" + assert kwargs["func"]() == "PRIVATE-ASSERTION" + return SimpleNamespace(token="PRIVATE-TOKEN", expires_on=clock.now + 3600) + client = Mock(get_token=Mock(side_effect=get_token)) + clients.append(client) + return client + + monkeypatch.setattr(credentials_module, "ClientAssertionCredential", create) + manager = ProviderCredentials(Mock(), cache_options=clock.options) + config = oauth_config() + try: + with ThreadPoolExecutor(max_workers=10) as pool: + results = list(pool.map(lambda _: manager.resolve({"mode": "oauth"}, config), range(10))) + assert all(value["access_token"] == "PRIVATE-TOKEN" for value in results) + assert identity.get_token.call_count == 1 and clients[0].get_token.call_count == 1 + assert manager.resolve({"mode": "oauth"}, config)["access_token"] == "PRIVATE-TOKEN" + assert identity.get_token.call_count == 1 and clients[0].get_token.call_count == 1 + clock.advance(3300) + wait_until(lambda: identity.get_token.call_count == 2 and clients[0].get_token.call_count == 2) + wait_until(lambda: clock.timer_count == 2) + config.provider_scope = "api://second/.default" + manager.resolve({"mode": "oauth"}, config) + assert len(clients) == 1 and clients[0].get_token.call_count == 3 + config.outbound_client_id = "different-app" + manager.resolve({"mode": "oauth"}, config) + assert len(clients) == 2 + assert clock.timer_count == 2 + finally: + manager.close() + assert clock.timer_count == 0 + + +def test_keyvault_pair_is_parallel_single_flight_and_failed_partial_refresh_retains_old_pair(): + clock = Clock() + key_started, id_started, release = Event(), Event(), Event() + calls = [] + version = 1 + fail_identity = False + + def resolve(name): + calls.append(name) + (key_started if name == "key" else id_started).set() + assert release.wait(3) + if name == "id" and fail_identity: + raise ValueError("PRIVATE-FAILURE") + return f"{name}-{version}" + + failures = [] + manager = ProviderCredentials(Mock(resolve=Mock(side_effect=resolve)), cache_options=clock.options, + report_failure=lambda kind: failures.append(kind)) + config = read_config({"KEY_VAULT_URL": "https://unit.vault.azure.net"}) + auth = {"mode": "apiKey", "key_vault_secret_name": "key", "identity_key_vault_secret_name": "id"} + try: + with ThreadPoolExecutor(max_workers=2) as pool: + first = pool.submit(manager.resolve, auth, config) + second = pool.submit(manager.resolve, auth, config) + assert key_started.wait(3) and id_started.wait(3) + release.set() + assert first.result(3) == second.result(3) == {"mode": "apiKey", "secret": "key-1", "identity": "id-1"} + assert sorted(calls) == ["id", "key"] + version = 2 + fail_identity = True + clock.advance(240) + wait_until(lambda: failures == ["key_vault"]) + assert manager.resolve(auth, config) == {"mode": "apiKey", "secret": "key-1", "identity": "id-1"} + fail_identity = False + clock.advance(5) + wait_until(lambda: manager.resolve(auth, config)["identity"] == "id-2") + assert manager.resolve(auth, config)["secret"] == "key-2" + finally: + release.set() + manager.close() + + +def test_startup_only_prepares_credentials_and_handles_missing_provider_or_failure(): + secrets = Mock(resolve=Mock(return_value="test-key")) + engine = DispatchEngine(ProviderRegistry([TelesignProvider()]), secrets, + {"EPP_PROVIDER_NAME": "telesign", "EPP_PROVIDER_AUTH_MODE": "apiKey"}) + try: + engine.start_credential_refresh() + assert secrets.resolve.call_count == 2 + engine.start_credential_refresh() + assert secrets.resolve.call_count == 2 + finally: + engine.close() + no_provider = DispatchEngine(ProviderRegistry([TelesignProvider()]), secrets, {}) + try: + no_provider.start_credential_refresh() + assert secrets.resolve.call_count == 2 + finally: + no_provider.close() + broken = DispatchEngine(ProviderRegistry([TelesignProvider()]), Mock(resolve=Mock(side_effect=ValueError("PRIVATE"))), + {"EPP_PROVIDER_NAME": "telesign"}) + try: + broken.start_credential_refresh() + finally: + broken.close() diff --git a/python/tests/test_credential_sdk.py b/python/tests/test_credential_sdk.py new file mode 100644 index 0000000..4f0f534 --- /dev/null +++ b/python/tests/test_credential_sdk.py @@ -0,0 +1,65 @@ +import json +import time +from concurrent.futures import ThreadPoolExecutor +from types import SimpleNamespace +from unittest.mock import Mock + +import requests + +import src.credentials as credentials_module +from src.config import read_config +from src.credentials import ProviderCredentials + + +def test_real_provider_sdk_uses_one_exchange_for_concurrent_requests_and_no_exchange_when_warm(monkeypatch): + token_endpoint_calls = [] + mi_calls = [] + + def response(payload): + result = requests.Response() + result.status_code = 200 + result.headers["Content-Type"] = "application/json" + result._content = json.dumps(payload).encode() + result._content_consumed = True + result.raw = SimpleNamespace(enforce_content_length=False) + return result + + def send(_session, method, url, **kwargs): + if method == "POST" and url.endswith("/oauth2/v2.0/token"): + token_endpoint_calls.append(url) + time.sleep(0.03) + return response({"access_token": "PRIVATE-PROVIDER", "expires_in": 3600, "token_type": "Bearer"}) + if method == "GET" and ".well-known/openid-configuration" in url: + return response({ + "token_endpoint": "https://login.microsoftonline.com/11111111-1111-1111-1111-111111111111/oauth2/v2.0/token", + "authorization_endpoint": "https://login.microsoftonline.com/11111111-1111-1111-1111-111111111111/oauth2/v2.0/authorize", + "issuer": "https://login.microsoftonline.com/11111111-1111-1111-1111-111111111111/v2.0", + }) + raise AssertionError("Unexpected HTTP in fake credential transport") + + def managed(*args, **kwargs): + mi_calls.append(args) + return SimpleNamespace(token="PRIVATE-ASSERTION", expires_on=time.time() + 3600) + + monkeypatch.setattr(requests.Session, "request", send) + monkeypatch.setattr(credentials_module, "ManagedIdentityCredential", Mock( + return_value=Mock(get_token=Mock(side_effect=managed)))) + secrets = Mock() + manager = ProviderCredentials(secrets) + config = read_config({ + "EPP_PROVIDER_TENANT_ID": "11111111-1111-1111-1111-111111111111", + "EPP_OUTBOUND_CLIENT_ID": "22222222-2222-2222-2222-222222222222", + "EPP_OUTBOUND_MI_CLIENT_ID": "33333333-3333-3333-3333-333333333333", + "EPP_PROVIDER_SCOPE": "api://provider/.default", + }) + try: + with ThreadPoolExecutor(max_workers=10) as pool: + results = list(pool.map(lambda _: manager.resolve({"mode": "oauth"}, config), range(10))) + assert all(value["access_token"] == "PRIVATE-PROVIDER" for value in results) + assert len(token_endpoint_calls) == 1 + assert len(mi_calls) == 1 + assert manager.resolve({"mode": "oauth"}, config)["access_token"] == "PRIVATE-PROVIDER" + assert len(token_endpoint_calls) == len(mi_calls) == 1 + secrets.resolve.assert_not_called() + finally: + manager.close() diff --git a/python/tests/test_engine.py b/python/tests/test_engine.py index 936c251..cf25b58 100644 --- a/python/tests/test_engine.py +++ b/python/tests/test_engine.py @@ -11,6 +11,7 @@ from urllib3.exceptions import ReadTimeoutError import src.dispatch as dispatch_module +import src.credentials as credentials_module from src.config import AppConfig, read_config from src.dispatch import DispatchEngine, DispatchRequest, ProviderRegistry from src.models import DeliveryContext, Envelope, TextToVoice @@ -33,7 +34,8 @@ def engine(monkeypatch): "EPP_PROVIDER_CHANNEL": "sms", }) result._resolve_credential = Mock(return_value={"mode": "oauth", "access_token": "provider-token"}) - return result + yield result + result.close() def test_missing_oauth_configuration_never_sends(engine): @@ -75,8 +77,8 @@ def get_token(*args, **options): clients.append(client) return client - monkeypatch.setattr(dispatch_module, "ManagedIdentityCredential", identity_factory) - monkeypatch.setattr(dispatch_module, "ClientAssertionCredential", create_client) + monkeypatch.setattr(credentials_module, "ManagedIdentityCredential", identity_factory) + monkeypatch.setattr(credentials_module, "ClientAssertionCredential", create_client) dispatch_module.requests.request.return_value = Mock(status_code=201, json=Mock(return_value={"status": "ENROUTE"})) for scope in ("api://provider/.default", "api://second/.default"): engine.env["EPP_PROVIDER_SCOPE"] = scope @@ -97,6 +99,7 @@ def get_token(*args, **options): for invalid in (None, SimpleNamespace(token=""), SimpleNamespace(token=" "), SimpleNamespace(token="private-token"), SimpleNamespace(token="private-token", expires_on=time.time() + 5)): + engine.close() if stage == "access": access = invalid else: @@ -125,8 +128,8 @@ def fail(*args, **kwargs): logger.warning("PRIVATE-ACCOUNT-ERROR") raise RuntimeError("PRIVATE-TOKEN-EXCEPTION") - monkeypatch.setattr(dispatch_module, "ManagedIdentityCredential", Mock()) - monkeypatch.setattr(dispatch_module, "ClientAssertionCredential", Mock(return_value=Mock(get_token=fail))) + monkeypatch.setattr(credentials_module, "ManagedIdentityCredential", Mock()) + monkeypatch.setattr(credentials_module, "ClientAssertionCredential", Mock(return_value=Mock(get_token=fail))) try: with ThreadPoolExecutor(max_workers=1) as pool: pending = pool.submit(engine.dispatch, _request(), "request") diff --git a/python/tests/test_function_app.py b/python/tests/test_function_app.py index 5a62ec0..fe5e50a 100644 --- a/python/tests/test_function_app.py +++ b/python/tests/test_function_app.py @@ -43,6 +43,8 @@ def _isolate(monkeypatch): engine._resolve_credential = Mock(return_value={"mode": "oauth", "access_token": "provider-token"}) monkeypatch.setattr(function_app, "_engine", engine) monkeypatch.setattr(dispatch_module.requests, "request", Mock()) + yield + engine.close() def _request(body, headers=None): @@ -67,7 +69,7 @@ def _envelope(**overrides): def _records(caplog): - return [json.loads(record.getMessage()) for record in caplog.records] + return [value for record in caplog.records if (value := json.loads(record.getMessage())).get("functionName")] def _summary(caplog): From 2af194365228946cd5f71b4866cdb1180cf3c3ea Mon Sep 17 00:00:00 2001 From: James Xian Date: Wed, 23 Sep 2026 12:19:38 -0700 Subject: [PATCH 2/3] Fix credential refresh cancellation and lifecycle Propagate cancellation through Azure credential HTTP transports, preserve Python SDK refresh hints, and prevent secret-read workers from blocking shutdown. Separate configuration resets from terminal manager disposal across all runtimes. Name timing policies, clarify cache types and state selection, use structured sanitized .NET logging, and update regression coverage and credential documentation. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- docs/CONTRACT.md | 19 +- dotnet/README.md | 4 +- dotnet/Src/DispatchEngine.cs | 2 +- dotnet/Src/ProviderCredentials.cs | 92 +++++--- dotnet/Src/RefreshingCache.cs | 27 ++- dotnet/Src/SecretResolver.cs | 4 +- dotnet/tests/CredentialCacheTests.cs | 99 ++++++++ dotnet/tests/EngineTests.cs | 19 +- javascript/README.md | 7 +- javascript/package-lock.json | 1 + javascript/package.json | 1 + javascript/src/functions/credentials.js | 165 +++++++++++--- javascript/src/functions/refreshingCache.js | 72 +++++- javascript/test/credential-cache.test.js | 17 ++ javascript/test/credential-sdk.test.js | 213 +++++++++++++++--- javascript/test/dispatch.test.js | 14 +- javascript/test/sendotp.test.js | 13 +- python/README.md | 9 +- python/src/credentials.py | 238 ++++++++++++++------ python/src/refreshing_cache.py | 120 +++++++--- python/src/secrets.py | 30 ++- python/tests/test_credential_cache.py | 156 ++++++++++++- python/tests/test_credential_sdk.py | 21 +- python/tests/test_engine.py | 14 +- 24 files changed, 1110 insertions(+), 247 deletions(-) diff --git a/docs/CONTRACT.md b/docs/CONTRACT.md index 58f572f..7795f29 100644 --- a/docs/CONTRACT.md +++ b/docs/CONTRACT.md @@ -118,10 +118,15 @@ Credential instances are reused for the configured tenant/application/identity. final provider token for each selected scope. A usable final token avoids both SDK acquisition calls on the delivery path. Each cache owns its refresh independently, so one waiting caller cannot cancel an assertion refresh another caller needs. JavaScript and .NET bound each acquisition to 2.5 seconds; -JavaScript also links that cancellation to the actual Azure SDK HTTP pipeline rather than relying -only on the SDK's `getToken` option. Python bounds each wait to 2.5 seconds and uses 2.5-second +JavaScript links that cancellation through an SDK HTTP-client wrapper, including the managed-identity +transport, rather than relying on `getToken` options or policies the SDK can replace. +Python bounds each wait to 2.5 seconds and uses 2.5-second connect/read inactivity timeouts; its shared refresh may finish after a waiter leaves. +Python uses `get_token_info` when the installed SDK supports it, preserving the `refresh_on` hint. +Older supported SDKs expose only expiry metadata through `get_token`; those versions retain the +five-minute pre-expiry refresh target. An acquisition failure never falls back to another token API. + Credential SDK transport retries are disabled; failed refreshes use the bounded backoff described below. These are not end-to-end delivery deadlines. JavaScript suppresses SDK logs in the acquisition's asynchronous context. Python filters Azure Identity/Core/MSAL records on configured @@ -350,6 +355,10 @@ stores. API-key entries are isolated by vault, managed identity and manifest sec state is isolated by provider tenant, application and managed identity, with separate final-token entries for different scopes. A change of credential configuration stops the old entries; deploy app-setting changes normally with a worker restart rather than mutating process environment in place. +Configuration replacement clears old entries without shutting down the manager. Explicit +`close()`/`Dispose()` is terminal: a stopped manager cannot acquire credentials again. Create a +new manager or restart the worker instead of reusing a stopped instance. Timing values are named +policy constants in each runtime's refreshing-cache implementation, not additional app settings. Concurrent cache misses share **one in-progress acquisition per entry**. A request with a still-usable cached credential returns it immediately while a due refresh proceeds separately. Refresh failures @@ -372,6 +381,12 @@ work and discard entries. Late completion cannot repopulate a stopped cache. A c can stop waiting without cancelling the shared retrieval. Failed or absent provider configuration does not prevent evaluation from working. +Python uses explicitly owned daemon threads for parallel secret reads, not executor workers that +are joined before application `atexit` handlers. A bundle keeps ownership of both reads until they +finish, so a waiting caller's timeout or one failed read cannot start overlapping retries. Shutdown +releases pending cache waiters immediately. Synchronous Python SDK I/O is not forcibly cancelled; +unfinished reads cannot publish late values or block normal process exit. + **This is not a guarantee that the first request after a cold start meets the caller's budget.** Initialization can itself be on that first request's critical path, and a worker may receive traffic before preparation finishes. Existing Always On/minimum-instance settings can help, but readiness, diff --git a/dotnet/README.md b/dotnet/README.md index c3de947..9ae5b5c 100644 --- a/dotnet/README.md +++ b/dotnet/README.md @@ -93,7 +93,9 @@ six-digit numeric run that is not part of a longer number and repeats the comple The hosted credential-refresh service automatically prewarms a configured provider. Separate process-local caches refresh Key Vault bundles, managed-identity assertions and final Entra tokens. Each acquisition owns its cancellation budget; cancelling a waiter does not cancel another -request's shared retrieval. Shutdown stops timers and drops values. See the +request's shared retrieval. Shutdown stops timers, cancels acquisitions and drops values. +Disposal is terminal; configuration replacement clears caches without disposing the manager. +Refresh failures use the same structured, sanitized JSON logging pattern as request events. See the [refresh contract](../docs/CONTRACT.md#credential-caching-and-refresh) for expiry/backoff semantics and cold-start limitations. Evaluation remains independent from successful credential preparation. diff --git a/dotnet/Src/DispatchEngine.cs b/dotnet/Src/DispatchEngine.cs index e180dab..5d416c9 100644 --- a/dotnet/Src/DispatchEngine.cs +++ b/dotnet/Src/DispatchEngine.cs @@ -241,7 +241,7 @@ private static ClientAssertionCredentialOptions OAuthOptions() { var options = new ClientAssertionCredentialOptions { AuthorityHost = AzureAuthorityHosts.AzurePublicCloud }; options.Retry.MaxRetries = 0; - options.Retry.NetworkTimeout = TimeSpan.FromSeconds(2.5); + options.Retry.NetworkTimeout = CredentialCachePolicy.AcquisitionTimeout; options.Diagnostics.IsLoggingEnabled = false; options.Diagnostics.IsLoggingContentEnabled = false; return options; diff --git a/dotnet/Src/ProviderCredentials.cs b/dotnet/Src/ProviderCredentials.cs index 5cddcb3..df2ade9 100644 --- a/dotnet/Src/ProviderCredentials.cs +++ b/dotnet/Src/ProviderCredentials.cs @@ -1,3 +1,4 @@ +using System.Text.Json; using Azure.Core; using Microsoft.Extensions.Logging; using Microsoft.Extensions.Logging.Abstractions; @@ -18,6 +19,7 @@ internal sealed class ProviderCredentials : IDisposable private RefreshingCache? _assertion; private TokenCredential? _credential; private readonly Dictionary> _tokens = new(); + private bool _disposed; internal ProviderCredentials(ISecretResolver secrets, IEnv env, Func createIdentity, @@ -32,33 +34,58 @@ internal ProviderCredentials(ISecretResolver secrets, IEnv env, _clock = clock ?? TimeProvider.System; } - internal void ReportFailure(string kind) => - _log.LogWarning("{{\"logType\":\"service\",\"eventName\":\"credential_refresh_failed\",\"cacheKind\":\"{CacheKind}\",\"failureReason\":\"credential_unavailable\"}}", kind); + internal void ReportFailure(string kind) + { + const string eventName = "credential_refresh_failed"; + var record = new Dictionary + { + ["logType"] = "service", + ["eventName"] = eventName, + ["cacheKind"] = kind, + ["failureReason"] = "credential_unavailable", + }; + _log.Log(LogLevel.Warning, new EventId(0, eventName), record, null, + static (state, _) => JsonSerializer.Serialize(state)); + } private RefreshingCache Cache(string kind, Func>> load) => new(load, () => ReportFailure(kind), _clock); internal async Task ResolveAsync(AuthConfig auth, AppConfig config, CancellationToken cancellation = default) { - if (auth.Mode == "oauth" && (string.IsNullOrWhiteSpace(config.ProviderTenantId) - || string.IsNullOrWhiteSpace(config.ProviderScope) || string.IsNullOrWhiteSpace(config.OutboundClientId) - || string.IsNullOrWhiteSpace(config.OutboundManagedIdentityClientId))) - { - Dispose(); - throw new InvalidOperationException("provider OAuth token unavailable"); - } - if (auth.Mode is not ("apiKey" or "oauth")) - { - Dispose(); - throw new InvalidOperationException("provider credential unavailable"); - } - var key = System.Text.Json.JsonSerializer.Serialize(auth.Mode == "apiKey" - ? new[] { auth.Mode, _env.Get("KEY_VAULT_URL"), _env.Get("AZURE_CLIENT_ID"), auth.KeyVaultSecretName, auth.IdentityKeyVaultSecretName } - : new[] { auth.Mode, config.ProviderTenantId, config.OutboundClientId, config.OutboundManagedIdentityClientId }); RefreshingCache? bundle; RefreshingCache? tokenCache = null; lock (_gate) { + ObjectDisposedException.ThrowIf(_disposed, this); + if (auth.Mode == "oauth" && (string.IsNullOrWhiteSpace(config.ProviderTenantId) + || string.IsNullOrWhiteSpace(config.ProviderScope) || string.IsNullOrWhiteSpace(config.OutboundClientId) + || string.IsNullOrWhiteSpace(config.OutboundManagedIdentityClientId))) + { + Clear(); + throw new InvalidOperationException("provider OAuth token unavailable"); + } + if (auth.Mode is not ("apiKey" or "oauth")) + { + Clear(); + throw new InvalidOperationException("provider credential unavailable"); + } + string key; + if (auth.Mode == "apiKey") + { + key = JsonSerializer.Serialize(new[] + { + auth.Mode, _env.Get("KEY_VAULT_URL"), _env.Get("AZURE_CLIENT_ID"), + auth.KeyVaultSecretName, auth.IdentityKeyVaultSecretName, + }); + } + else + { + key = JsonSerializer.Serialize(new[] + { + auth.Mode, config.ProviderTenantId, config.OutboundClientId, config.OutboundManagedIdentityClientId, + }); + } if (_key != key) { Clear(); @@ -69,10 +96,10 @@ internal async Task ResolveAsync(AuthConfig auth, AppConfig bundle = _bundle; if (auth.Mode == "oauth") { - var scope = config.ProviderScope!; + var scope = config.ProviderScope ?? throw new InvalidOperationException("provider OAuth token unavailable"); if (!_tokens.TryGetValue(scope, out tokenCache)) { - var credential = _credential!; + var credential = _credential ?? throw new InvalidOperationException("provider OAuth token unavailable"); tokenCache = Cache("provider_token", async ct => TokenEntry(await credential.GetTokenAsync(new TokenRequestContext(new[] { scope }), ct).ConfigureAwait(false))); _tokens.Add(scope, tokenCache); @@ -80,7 +107,8 @@ internal async Task ResolveAsync(AuthConfig auth, AppConfig } } if (bundle is not null) return await bundle.GetAsync(cancellation).ConfigureAwait(false); - var token = await tokenCache!.GetAsync(cancellation).ConfigureAwait(false); + if (tokenCache is null) throw new InvalidOperationException("provider credential unavailable"); + var token = await tokenCache.GetAsync(cancellation).ConfigureAwait(false); return new ProviderCredential("oauth", AccessToken: token.Token); } @@ -97,7 +125,7 @@ private RefreshingCache CreateBundle(AuthConfig auth) => Cac throw new InvalidOperationException("provider credential unavailable"); var now = _clock.GetUtcNow(); return new CredentialCacheEntry(new("apiKey", Secret: secret, Identity: identity), - now.AddMinutes(5), now.AddMinutes(4)); + now + CredentialCachePolicy.SecretTtl, now + CredentialCachePolicy.SecretRefreshInterval); }); private void CreateOAuth(AppConfig config) @@ -114,12 +142,17 @@ private void CreateOAuth(AppConfig config) private CredentialCacheEntry TokenEntry(AccessToken token) { var now = _clock.GetUtcNow(); - if (string.IsNullOrWhiteSpace(token.Token) || token.ExpiresOn <= now.AddSeconds(30)) + if (string.IsNullOrWhiteSpace(token.Token) || token.ExpiresOn <= now + CredentialCachePolicy.TokenExpirySkew) throw new InvalidOperationException("provider credential unavailable"); - var expires = token.ExpiresOn.AddSeconds(-30); - var refresh = token.ExpiresOn.AddMinutes(-5); + var expires = token.ExpiresOn - CredentialCachePolicy.TokenExpirySkew; + var refresh = token.ExpiresOn - CredentialCachePolicy.TokenRefreshLead; if (token.RefreshOn is { } hint && hint < refresh) refresh = hint; - if (refresh <= now) refresh = now.AddSeconds(Math.Max(1, Math.Min(60, (expires - now).TotalSeconds / 2))); + if (refresh <= now) + { + var delay = Math.Clamp((expires - now).TotalSeconds / 2, + CredentialCachePolicy.MinRefreshDelay.TotalSeconds, CredentialCachePolicy.MaxRefreshDelay.TotalSeconds); + refresh = now.AddSeconds(delay); + } return new(token, expires, refresh); } @@ -135,5 +168,12 @@ private void Clear() _key = null; } - public void Dispose() { lock (_gate) Clear(); } + public void Dispose() + { + lock (_gate) + { + _disposed = true; + Clear(); + } + } } diff --git a/dotnet/Src/RefreshingCache.cs b/dotnet/Src/RefreshingCache.cs index bdb1b2f..782dfa1 100644 --- a/dotnet/Src/RefreshingCache.cs +++ b/dotnet/Src/RefreshingCache.cs @@ -1,5 +1,21 @@ namespace Epp.Otp; +internal static class CredentialCachePolicy +{ + internal static readonly TimeSpan AcquisitionTimeout = TimeSpan.FromSeconds(2.5); + internal static readonly TimeSpan SecretTtl = TimeSpan.FromMinutes(5); + internal static readonly TimeSpan SecretRefreshInterval = TimeSpan.FromMinutes(4); + internal static readonly TimeSpan TokenExpirySkew = TimeSpan.FromSeconds(30); + internal static readonly TimeSpan TokenRefreshLead = TimeSpan.FromMinutes(5); + internal static readonly TimeSpan MinRefreshDelay = TimeSpan.FromSeconds(1); + internal static readonly TimeSpan MaxRefreshDelay = TimeSpan.FromSeconds(60); + internal static readonly TimeSpan InitialRetryDelay = TimeSpan.FromSeconds(5); + internal static readonly TimeSpan MaxRetryDelay = TimeSpan.FromSeconds(60); + internal static readonly TimeSpan MaxTimerDelay = TimeSpan.FromMilliseconds(int.MaxValue); + internal const int MaxRetryExponent = 4; + internal const double RetryJitterRatio = 0.2; +} + internal sealed record CredentialCacheEntry(T Value, DateTimeOffset ExpiresAt, DateTimeOffset RefreshAt) { public override string ToString() => nameof(CredentialCacheEntry); @@ -65,7 +81,7 @@ private Task StartRefresh() _timer = null; var completion = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); _inFlight = completion.Task; - _acquisition = new CancellationTokenSource(TimeSpan.FromSeconds(2.5), _clock); + _acquisition = new CancellationTokenSource(CredentialCachePolicy.AcquisitionTimeout, _clock); _ = RunRefreshAsync(completion, _acquisition); return completion.Task; } @@ -93,8 +109,10 @@ private async Task RunRefreshAsync(TaskCompletionSource completion, Cancellat if (entry is null) { _failures++; - var backoff = Math.Min(60, 5 * Math.Pow(2, Math.Min(_failures - 1, 4))); - _retryAt = _clock.GetUtcNow().AddSeconds(backoff * (1 + _random() * 0.2)); + var exponent = Math.Min(_failures - 1, CredentialCachePolicy.MaxRetryExponent); + var backoff = Math.Min(CredentialCachePolicy.MaxRetryDelay.TotalSeconds, + CredentialCachePolicy.InitialRetryDelay.TotalSeconds * Math.Pow(2, exponent)); + _retryAt = _clock.GetUtcNow().AddSeconds(backoff * (1 + _random() * CredentialCachePolicy.RetryJitterRatio)); _onFailure(); } else @@ -104,7 +122,8 @@ private async Task RunRefreshAsync(TaskCompletionSource completion, Cancellat _retryAt = default; } var next = entry is null ? _retryAt : entry.RefreshAt; - var wait = Math.Clamp((next - _clock.GetUtcNow()).TotalMilliseconds, 1000, int.MaxValue); + var wait = Math.Clamp((next - _clock.GetUtcNow()).TotalMilliseconds, + CredentialCachePolicy.MinRefreshDelay.TotalMilliseconds, CredentialCachePolicy.MaxTimerDelay.TotalMilliseconds); _timer = _clock.CreateTimer(_ => ScheduledRefresh(), null, TimeSpan.FromMilliseconds(wait), Timeout.InfiniteTimeSpan); } if (entry is null) completion.TrySetException(Unavailable()); diff --git a/dotnet/Src/SecretResolver.cs b/dotnet/Src/SecretResolver.cs index f95b4b0..6795144 100644 --- a/dotnet/Src/SecretResolver.cs +++ b/dotnet/Src/SecretResolver.cs @@ -27,13 +27,13 @@ private SecretClient GetClient() if (_client is not null && _clientKey == (url, clientId)) return _client; var identityOptions = new TokenCredentialOptions(); identityOptions.Retry.MaxRetries = 0; - identityOptions.Retry.NetworkTimeout = TimeSpan.FromSeconds(2.5); + identityOptions.Retry.NetworkTimeout = CredentialCachePolicy.AcquisitionTimeout; identityOptions.Diagnostics.IsLoggingEnabled = false; identityOptions.Diagnostics.IsLoggingContentEnabled = false; var credential = new ManagedIdentityCredential(clientId, identityOptions); var options = new SecretClientOptions(); options.Retry.MaxRetries = 0; - options.Retry.NetworkTimeout = TimeSpan.FromSeconds(2.5); + options.Retry.NetworkTimeout = CredentialCachePolicy.AcquisitionTimeout; options.Diagnostics.IsLoggingEnabled = false; options.Diagnostics.IsLoggingContentEnabled = false; _client = new SecretClient(new Uri(url), credential, options); diff --git a/dotnet/tests/CredentialCacheTests.cs b/dotnet/tests/CredentialCacheTests.cs index 649fab8..f0f6d08 100644 --- a/dotnet/tests/CredentialCacheTests.cs +++ b/dotnet/tests/CredentialCacheTests.cs @@ -1,4 +1,6 @@ +using System.Text.Json; using Azure.Core; +using Microsoft.Extensions.Logging; using Xunit; namespace Epp.Otp.Tests; @@ -218,6 +220,87 @@ public async Task RepeatedSdkTokenDoesNotExtendLifetimeOrCauseATightRefreshLoop( await Assert.ThrowsAsync(() => manager.ResolveAsync(new("oauth"), Config())); } + [Fact] + public async Task DisposingManagerIsTerminalAndCancelsPendingAcquisition() + { + var clock = new ManualClock(); + var release = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var calls = 0; + CancellationToken acquisition = default; + var env = new TestEnv { ["KEY_VAULT_URL"] = "https://unit.vault.azure.net" }; + using var manager = new ProviderCredentials(new Secrets((_, cancellation) => + { + calls++; + acquisition = cancellation; + return release.Task; + }), env, _ => throw new Exception("Unexpected managed identity"), + (_, _, _) => throw new Exception("Unexpected OAuth"), clock: clock); + var auth = new AuthConfig("apiKey", "key"); + var pending = manager.ResolveAsync(auth, new AppConfig()); + manager.Dispose(); + manager.Dispose(); + Assert.True(acquisition.IsCancellationRequested); + await Assert.ThrowsAsync(() => pending); + release.SetResult("PRIVATE-LATE-KEY"); + await Assert.ThrowsAsync(() => manager.ResolveAsync(auth, new AppConfig())); + env["KEY_VAULT_URL"] = "https://different.vault.azure.net"; + await Assert.ThrowsAsync(() => manager.ResolveAsync(auth, new AppConfig())); + await Assert.ThrowsAsync(() => manager.ResolveAsync(new("oauth"), Config())); + Assert.Equal(1, calls); + Assert.Equal(0, clock.TimerCount); + } + + [Fact] + public async Task InvalidConfigurationClearsOldStateWithoutDisposingManager() + { + var clock = new ManualClock(); + var instances = 0; + using var manager = new ProviderCredentials(new Secrets((_, _) => throw new Exception("Unexpected Key Vault")), + new TestEnv(), _ => new Token((_, _) => + ValueTask.FromResult(new AccessToken("assertion", clock.GetUtcNow().AddHours(1)))), + (_, _, assertion) => + { + instances++; + return new Token(async (_, cancellation) => + { + await assertion(cancellation); + return new("provider-token", clock.GetUtcNow().AddHours(1)); + }); + }, clock: clock); + await manager.ResolveAsync(new("oauth"), Config()); + Assert.Equal(2, clock.TimerCount); + await Assert.ThrowsAsync(() => manager.ResolveAsync(new("oauth"), Config(scope: ""))); + Assert.Equal(0, clock.TimerCount); + Assert.Equal("provider-token", (await manager.ResolveAsync(new("oauth"), Config())).AccessToken); + Assert.Equal(2, instances); + Assert.Equal(2, clock.TimerCount); + } + + [Fact] + public async Task RefreshFailureUsesStructuredSanitizedLogRecord() + { + var clock = new ManualClock(); + var logger = new CredentialLogger(); + using var manager = new ProviderCredentials(new Secrets((_, _) => + throw new InvalidOperationException("PRIVATE-SDK-ERROR")), new TestEnv(), + _ => throw new Exception("Unexpected managed identity"), + (_, _, _) => throw new Exception("Unexpected OAuth"), logger, clock); + await Assert.ThrowsAsync(() => + manager.ResolveAsync(new("apiKey", "key"), new AppConfig())); + var entry = Assert.Single(logger.Entries); + Assert.Equal(LogLevel.Warning, entry.Level); + Assert.Equal("credential_refresh_failed", entry.EventId.Name); + Assert.Null(entry.Error); + Assert.Equal(4, entry.State.Count); + Assert.Equal("service", entry.State["logType"]); + Assert.Equal("credential_refresh_failed", entry.State["eventName"]); + Assert.Equal("key_vault", entry.State["cacheKind"]); + Assert.Equal("credential_unavailable", entry.State["failureReason"]); + using var json = JsonDocument.Parse(entry.Message); + Assert.Equal("key_vault", json.RootElement.GetProperty("cacheKind").GetString()); + Assert.DoesNotContain("PRIVATE", entry.Message); + } + private static AppConfig Config(string scope = "api://provider/.default", string application = "app") => new() { ProviderTenantId = "tenant", ProviderScope = scope, OutboundClientId = application, @@ -247,6 +330,22 @@ public override ValueTask GetTokenAsync(TokenRequestContext request acquire(requestContext, cancellationToken); } + private sealed record LogEntry(LogLevel Level, EventId EventId, IReadOnlyDictionary State, + Exception? Error, string Message); + + private sealed class CredentialLogger : ILogger + { + internal List Entries { get; } = new(); + public IDisposable? BeginScope(TState state) where TState : notnull => null; + public bool IsEnabled(LogLevel level) => true; + public void Log(LogLevel level, EventId eventId, TState state, Exception? error, + Func formatter) + { + var record = Assert.IsAssignableFrom>(state); + Entries.Add(new(level, eventId, record, error, formatter(state, error))); + } + } + private sealed class ManualClock : TimeProvider { private readonly object _gate = new(); diff --git a/dotnet/tests/EngineTests.cs b/dotnet/tests/EngineTests.cs index 3727163..b182d55 100644 --- a/dotnet/tests/EngineTests.cs +++ b/dotnet/tests/EngineTests.cs @@ -155,7 +155,7 @@ public async Task SopranoOAuthCancellationAndRejectionNeverFallBackOrRetry() { CancellationToken observed = default; var waitForCancellation = true; - using var rig = new HandlerRig(_ => new TestTokenCredential(async (_, cancellation) => + TokenCredential CreateIdentity(string _) => new TestTokenCredential(async (_, cancellation) => { if (waitForCancellation) { @@ -163,20 +163,27 @@ public async Task SopranoOAuthCancellationAndRejectionNeverFallBackOrRetry() await Task.Delay(Timeout.Infinite, cancellation); } return new AccessToken("assertion", DateTimeOffset.UtcNow.AddHours(1)); - }), (_, _, assertion) => new TestTokenCredential(async (_, cancellation) => + }); + TokenCredential CreateProvider(string tenant, string application, Func> assertion) => + new TestTokenCredential(async (_, cancellation) => { await assertion(cancellation); return new AccessToken("token", DateTimeOffset.UtcNow.AddHours(1)); - })); + }); + using var rig = new HandlerRig(CreateIdentity, CreateProvider); ConfigureSoprano(rig); AssertFailure(rig, await rig.Invoke().WaitAsync(TimeSpan.FromSeconds(10)), 502); Assert.True(observed.IsCancellationRequested); Assert.Equal((0, 0), (rig.Http.Calls, rig.Secrets.Calls)); waitForCancellation = false; rig.Engine.Dispose(); - rig.Http.Respond = _ => Task.FromResult(Json(401, "{\"status\":\"REJECTED\"}")); - AssertFailure(rig, await rig.Invoke(), 401); - Assert.Equal((1, 0), (rig.Http.Calls, rig.Secrets.Calls)); + AssertFailure(rig, await rig.Invoke(), 502); + Assert.Equal((0, 0), (rig.Http.Calls, rig.Secrets.Calls)); + using var replacement = new HandlerRig(CreateIdentity, CreateProvider); + ConfigureSoprano(replacement); + replacement.Http.Respond = _ => Task.FromResult(Json(401, "{\"status\":\"REJECTED\"}")); + AssertFailure(replacement, await replacement.Invoke(), 401); + Assert.Equal((1, 0), (replacement.Http.Calls, replacement.Secrets.Calls)); } private sealed class TestTokenCredential(Func> acquire) : TokenCredential diff --git a/javascript/README.md b/javascript/README.md index 30d7580..5772b49 100644 --- a/javascript/README.md +++ b/javascript/README.md @@ -101,8 +101,11 @@ retries. The shared contract defines validation, HTTP outcomes and privacy-safe Configured providers automatically prewarm on the app-start hook. Key Vault credential bundles, managed-identity assertions and final Entra tokens refresh through separate process-local caches. -Warm requests reuse usable values; concurrent misses share a retrieval, refresh failure never -extends expiry, and termination stops timers. Startup/refresh never sends an OTP. See the +Warm requests reuse usable values; concurrent misses share a retrieval and refresh failure never +extends expiry. Acquisition deadlines and termination cancel the actual SDK HTTP transport, +including managed identity, through an HTTP-client wrapper. Termination stops timers and closes +the manager permanently; configuration replacement uses a separate cache reset. Startup/refresh +never sends an OTP. See the [refresh contract](../docs/CONTRACT.md#credential-caching-and-refresh) for budgets and cold-start limitations. Leave the provider unset for local evaluation-only use without credential acquisition. diff --git a/javascript/package-lock.json b/javascript/package-lock.json index 20db06d..d5654ec 100644 --- a/javascript/package-lock.json +++ b/javascript/package-lock.json @@ -8,6 +8,7 @@ "name": "epp-otp-sample", "version": "1.0.0", "dependencies": { + "@azure/core-rest-pipeline": "^1.24.0", "@azure/functions": "^4.0.0", "@azure/identity": "^4.13.1", "@azure/keyvault-secrets": "^4.11.2", diff --git a/javascript/package.json b/javascript/package.json index 88a5420..483d488 100644 --- a/javascript/package.json +++ b/javascript/package.json @@ -8,6 +8,7 @@ "test": "node --test" }, "dependencies": { + "@azure/core-rest-pipeline": "^1.24.0", "@azure/functions": "^4.0.0", "@azure/identity": "^4.13.1", "@azure/keyvault-secrets": "^4.11.2", diff --git a/javascript/src/functions/credentials.js b/javascript/src/functions/credentials.js index 9c14952..cfd93f7 100644 --- a/javascript/src/functions/credentials.js +++ b/javascript/src/functions/credentials.js @@ -3,35 +3,59 @@ const { AsyncLocalStorage } = require('node:async_hooks'); const { inspect } = require('node:util'); +const { createDefaultHttpClient } = require('@azure/core-rest-pipeline'); const { ClientAssertionCredential, ManagedIdentityCredential } = require('@azure/identity'); const { SecretClient } = require('@azure/keyvault-secrets'); const { AzureLogger } = require('@azure/logger'); -const { RefreshingCache, tokenEntry } = require('./refreshingCache'); +const { CACHE_POLICY, RefreshingCache, tokenEntry } = require('./refreshingCache'); +/** + * @template T + * @typedef {InstanceType>} CredentialCache + */ + +/** + * @typedef {ReturnType} AppConfig + * @typedef {{mode?: string, keyVaultSecretName?: string, identityKeyVaultSecretName?: string}} AuthConfig + * @typedef {{mode: 'apiKey', secret: string, identity: string}} ApiKeyCredential + * @typedef {{mode: 'oauth', accessToken: string}} OAuthCredential + * @typedef {import('@azure/core-auth').AccessToken} AccessToken + * @typedef {{mode: 'apiKey', bundle: CredentialCache}} ApiKeyState + * @typedef {object} OAuthState + * @property {'oauth'} mode + * @property {CredentialCache} assertion + * @property {ClientAssertionCredential} credential + * @property {Map>} tokens + */ + +/** @type {AsyncLocalStorage} */ const acquisition = new AsyncLocalStorage(); +/** @type {typeof AzureLogger.log | undefined} */ let filteredLogger; +/** @param {string} cacheKind */ function reportRefreshFailure(cacheKind) { console.warn(JSON.stringify({ logType: 'service', eventName: 'credential_refresh_failed', cacheKind, failureReason: 'credential_unavailable' })); } -function cancellationPolicy() { +/** @returns {import('@azure/core-rest-pipeline').HttpClient} */ +function credentialHttpClient() { + const client = createDefaultHttpClient(); return { - name: 'eppCredentialCancellation', - async sendRequest(request, next) { + async sendRequest(request) { const signal = acquisition.getStore(); - if (!signal) return next(request); + if (!signal) return client.sendRequest(request); const controller = new AbortController(); const existing = request.abortSignal; const abort = () => controller.abort(); signal.addEventListener('abort', abort, { once: true }); - existing?.addEventListener('abort', abort, { once: true }); + existing?.addEventListener('abort', abort); if (signal.aborted || existing?.aborted) controller.abort(); request.abortSignal = controller.signal; try { controller.signal.throwIfAborted(); - return await next(request); + return await client.sendRequest(request); } finally { signal.removeEventListener('abort', abort); existing?.removeEventListener('abort', abort); @@ -40,6 +64,12 @@ function cancellationPolicy() { }; } +/** + * @template T + * @param {AbortSignal} signal + * @param {(signal: AbortSignal) => Promise} load + * @returns {Promise} + */ async function acquireBounded(signal, load) { if (AzureLogger.log !== filteredLogger) { const previous = AzureLogger.log; @@ -51,11 +81,12 @@ async function acquireBounded(signal, load) { signal.addEventListener('abort', abort, { once: true }); if (signal.aborted) abort(); let timeout; + /** @type {Promise} */ const interrupted = new Promise((_, reject) => { const fail = () => reject(new Error('provider credential unavailable')); controller.signal.addEventListener('abort', fail, { once: true }); if (controller.signal.aborted) fail(); - timeout = setTimeout(abort, 2500); + timeout = setTimeout(abort, CACHE_POLICY.acquisitionTimeoutMs); }); try { return await acquisition.run(controller.signal, () => Promise.race([ @@ -71,58 +102,104 @@ async function acquireBounded(signal, load) { } } -const sdkOptions = () => ({ retryOptions: { maxRetries: 0 }, - additionalPolicies: [{ policy: cancellationPolicy(), position: 'perCall' }] }); +// ManagedIdentityCredential replaces additionalPolicies, but preserves the supplied HTTP client. +const sdkOptions = () => ({ + retryOptions: { maxRetries: 0 }, + httpClient: credentialHttpClient(), +}); class ProviderCredentials { + /** + * @param {{cacheOptions?: import('./refreshingCache').CacheOptions, + * reportFailure?: (kind: string) => void}} [options] + */ constructor({ cacheOptions = {}, reportFailure = reportRefreshFailure } = {}) { this.cacheOptions = cacheOptions; this.now = cacheOptions.now || Date.now; this.reportFailure = reportFailure; + /** @type {ApiKeyState | OAuthState | null} */ this.current = null; + /** @type {string | null} */ this.currentKey = null; + this.closed = false; } [inspect.custom]() { return '[ProviderCredentials]'; } toJSON() { return '[ProviderCredentials]'; } + /** + * @template T + * @param {string} kind + * @param {(signal: AbortSignal) => Promise>} load + * @returns {CredentialCache} + */ cache(kind, load) { return new RefreshingCache((signal) => acquireBounded(signal, load), { ...this.cacheOptions, onFailure: () => this.reportFailure(kind) }); } + /** + * @param {AuthConfig} auth + * @param {AppConfig} config + * @returns {Promise} + */ resolve(auth, config) { + if (this.closed) return Promise.reject(new Error('provider credential unavailable')); const mode = auth?.mode || 'apiKey'; if (mode === 'oauth' && (!config.providerTenantId || !config.providerScope || !config.outboundClientId || !config.outboundManagedIdentityClientId)) { - this.close(); + this.#clear(); return Promise.reject(new Error('provider OAuth token unavailable')); } if (mode !== 'oauth' && mode !== 'apiKey') { - this.close(); + this.#clear(); return Promise.reject(new Error('provider credential unavailable')); } - const key = JSON.stringify(mode === 'apiKey' - ? [mode, config.keyVaultUrl, config.managedIdentityClientId, auth.keyVaultSecretName, auth.identityKeyVaultSecretName] - : [mode, config.providerTenantId, config.outboundClientId, config.outboundManagedIdentityClientId]); + let key; + if (mode === 'apiKey') { + key = JSON.stringify([ + mode, config.keyVaultUrl, config.managedIdentityClientId, + auth.keyVaultSecretName, auth.identityKeyVaultSecretName, + ]); + } else { + key = JSON.stringify([ + mode, config.providerTenantId, config.outboundClientId, config.outboundManagedIdentityClientId, + ]); + } if (this.currentKey !== key) { - this.close(); - this.current = mode === 'apiKey' ? this.apiKeyState(auth, config) : this.oauthState(config); + this.#clear(); + if (mode === 'apiKey') { + this.current = this.apiKeyState(auth, config); + } else { + this.current = this.oauthState(config); + } this.currentKey = key; } - if (mode === 'apiKey') return this.current.bundle.get(); const state = this.current; - if (!state.tokens.has(config.providerScope)) { + if (!state) return Promise.reject(new Error('provider credential unavailable')); + if (state.mode === 'apiKey') return state.bundle.get(); + let tokenCache = state.tokens.get(config.providerScope); + if (!tokenCache) { const scope = config.providerScope; - state.tokens.set(scope, this.cache('provider_token', async (signal) => - tokenEntry(await state.credential.getToken(scope, { abortSignal: signal }), this.now()))); + tokenCache = this.cache('provider_token', async (signal) => + tokenEntry(await state.credential.getToken(scope, { abortSignal: signal }), this.now())); + state.tokens.set(scope, tokenCache); } - return state.tokens.get(config.providerScope).get().then((token) => - Object.defineProperty({ mode: 'oauth' }, 'accessToken', { value: token.token })) - .catch(() => { throw new Error('provider OAuth token unavailable'); }); + return tokenCache.get().then((token) => { + /** @type {OAuthCredential} */ + const credential = { mode: 'oauth', accessToken: token.token }; + Object.defineProperty(credential, 'accessToken', { enumerable: false, writable: false, configurable: false }); + return credential; + }).catch(() => { throw new Error('provider OAuth token unavailable'); }); } + /** + * @param {AuthConfig} auth + * @param {AppConfig} config + * @returns {ApiKeyState} + */ apiKeyState(auth, config) { + /** @type {SecretClient | undefined} */ let client; const bundle = this.cache('key_vault', async (signal) => { if (!auth.keyVaultSecretName || !config.keyVaultUrl) throw new Error('provider credential unavailable'); @@ -140,20 +217,31 @@ class ProviderCredentials { throw new Error('provider credential unavailable'); } const now = this.now(); - let expiresAt = now + 300000; + let expiresAt = now + CACHE_POLICY.secretTtlMs; for (const item of [secret, identity]) { if (!item) continue; + const notBefore = item.properties?.notBefore?.getTime(); if (item.properties?.enabled === false - || item.properties?.notBefore?.getTime() > now) throw new Error('provider credential unavailable'); + || (notBefore !== undefined && notBefore > now)) throw new Error('provider credential unavailable'); if (item.properties?.expiresOn) expiresAt = Math.min(expiresAt, item.properties.expiresOn.getTime()); } - return { value: Object.freeze({ mode: 'apiKey', secret: secret.value, identity: identity?.value || '' }), - expiresAt, refreshAt: Math.min(now + 240000, expiresAt - 30000) }; + /** @type {ApiKeyCredential} */ + const credential = { mode: 'apiKey', secret: secret.value, identity: identity?.value || '' }; + return { + value: Object.freeze(credential), + expiresAt, + refreshAt: Math.min(now + CACHE_POLICY.secretRefreshIntervalMs, expiresAt - CACHE_POLICY.secretExpiryRefreshLeadMs), + }; }); - return { bundle }; + return { mode: 'apiKey', bundle }; } + /** + * @param {AppConfig} config + * @returns {OAuthState} + */ oauthState(config) { + /** @type {ManagedIdentityCredential | undefined} */ let identity; const assertion = this.cache('managed_identity', async (signal) => { identity ??= new ManagedIdentityCredential({ clientId: config.outboundManagedIdentityClientId, ...sdkOptions() }); @@ -165,16 +253,25 @@ class ProviderCredentials { acquisition.getStore()?.throwIfAborted(); return value.token; }, { authorityHost: 'https://login.microsoftonline.com', ...sdkOptions() }); - return { assertion, credential, tokens: new Map() }; + return { mode: 'oauth', assertion, credential, tokens: new Map() }; } - close() { - this.current?.bundle?.close(); - this.current?.assertion?.close(); - if (this.current?.tokens) for (const cache of this.current.tokens.values()) cache.close(); + #clear() { + const state = this.current; + if (state?.mode === 'apiKey') { + state.bundle.close(); + } else if (state?.mode === 'oauth') { + state.assertion.close(); + for (const cache of state.tokens.values()) cache.close(); + } this.current = null; this.currentKey = null; } + + close() { + this.closed = true; + this.#clear(); + } } const providerCredentials = new ProviderCredentials(); diff --git a/javascript/src/functions/refreshingCache.js b/javascript/src/functions/refreshingCache.js index 938a374..75fdf67 100644 --- a/javascript/src/functions/refreshingCache.js +++ b/javascript/src/functions/refreshingCache.js @@ -3,9 +3,47 @@ const { inspect } = require('node:util'); +const CACHE_POLICY = Object.freeze({ + acquisitionTimeoutMs: 2500, + secretTtlMs: 5 * 60 * 1000, + secretRefreshIntervalMs: 4 * 60 * 1000, + secretExpiryRefreshLeadMs: 30 * 1000, + tokenExpirySkewMs: 30 * 1000, + tokenRefreshLeadMs: 5 * 60 * 1000, + minRefreshDelayMs: 1000, + maxRefreshDelayMs: 60 * 1000, + initialRetryDelayMs: 5000, + maxRetryDelayMs: 60 * 1000, + maxRetryExponent: 4, + retryJitterRatio: 0.2, + maxTimerDelayMs: 2 ** 31 - 1, +}); + +/** + * @template T + * @typedef {object} CacheEntry + * @property {T} value + * @property {number} expiresAt Absolute Unix time in milliseconds. + * @property {number} refreshAt Absolute Unix time in milliseconds. + */ + +/** + * @typedef {object} CacheOptions + * @property {() => number} [now] Unix time in milliseconds. + * @property {typeof setTimeout} [schedule] + * @property {typeof clearTimeout} [cancel] + * @property {() => number} [random] + * @property {() => void} [onFailure] + */ + const unavailable = () => new Error('provider credential unavailable'); +/** @template T */ class RefreshingCache { + /** + * @param {(signal: AbortSignal) => Promise>} load + * @param {CacheOptions} [options] + */ constructor(load, { now = Date.now, schedule = setTimeout, cancel = clearTimeout, random = Math.random, onFailure = () => {} } = {}) { this.load = load; @@ -14,12 +52,16 @@ class RefreshingCache { this.cancel = cancel; this.random = random; this.onFailure = onFailure; + /** @type {CacheEntry | null} */ this.entry = null; + /** @type {Promise | null} */ this.inFlight = null; + /** @type {ReturnType | null} */ this.timer = null; this.retryAt = 0; this.failures = 0; this.closed = false; + /** @type {AbortController | null} */ this.controller = null; } @@ -58,8 +100,9 @@ class RefreshingCache { }).catch(() => { if (!this.closed) { this.failures++; - const backoff = Math.min(60000, 5000 * 2 ** Math.min(this.failures - 1, 4)); - this.retryAt = this.now() + Math.floor(backoff * (1 + this.random() * 0.2)); + const exponent = Math.min(this.failures - 1, CACHE_POLICY.maxRetryExponent); + const backoff = Math.min(CACHE_POLICY.maxRetryDelayMs, CACHE_POLICY.initialRetryDelayMs * 2 ** exponent); + this.retryAt = this.now() + Math.floor(backoff * (1 + this.random() * CACHE_POLICY.retryJitterRatio)); this.onFailure(); } throw unavailable(); @@ -84,12 +127,13 @@ class RefreshingCache { void this.refresh().catch(() => {}); } + /** @param {number} at Absolute Unix time in milliseconds. */ arm(at) { this.clearTimer(); this.timer = this.schedule(() => { this.timer = null; this.refreshInBackground(); - }, Math.max(1000, Math.min(2147483647, at - this.now()))); + }, Math.max(CACHE_POLICY.minRefreshDelayMs, Math.min(CACHE_POLICY.maxTimerDelayMs, at - this.now()))); this.timer?.unref?.(); } @@ -106,17 +150,27 @@ class RefreshingCache { } } +/** + * @param {import('@azure/core-auth').AccessToken | null | undefined} token + * @param {number} now Unix time in milliseconds. + * @returns {CacheEntry} + */ function tokenEntry(token, now) { if (typeof token?.token !== 'string' || !token.token.trim() - || !Number.isFinite(token.expiresOnTimestamp) || token.expiresOnTimestamp <= now + 30000) { + || !Number.isFinite(token.expiresOnTimestamp) || token.expiresOnTimestamp <= now + CACHE_POLICY.tokenExpirySkewMs) { throw unavailable(); } - const expiresAt = token.expiresOnTimestamp - 30000; - let refreshAt = token.expiresOnTimestamp - 300000; - if (Number.isFinite(token.refreshAfterTimestamp)) refreshAt = Math.min(refreshAt, token.refreshAfterTimestamp); + const expiresAt = token.expiresOnTimestamp - CACHE_POLICY.tokenExpirySkewMs; + let refreshAt = token.expiresOnTimestamp - CACHE_POLICY.tokenRefreshLeadMs; + if (typeof token.refreshAfterTimestamp === 'number' && Number.isFinite(token.refreshAfterTimestamp)) { + refreshAt = Math.min(refreshAt, token.refreshAfterTimestamp); + } // An SDK may return its existing token on refresh. Never extend that token's lifetime or spin. - if (refreshAt <= now) refreshAt = now + Math.max(1000, Math.min(60000, (expiresAt - now) / 2)); + if (refreshAt <= now) { + const delay = Math.min(CACHE_POLICY.maxRefreshDelayMs, (expiresAt - now) / 2); + refreshAt = now + Math.max(CACHE_POLICY.minRefreshDelayMs, delay); + } return { value: token, expiresAt, refreshAt }; } -module.exports = { RefreshingCache, tokenEntry }; +module.exports = { CACHE_POLICY, RefreshingCache, tokenEntry }; diff --git a/javascript/test/credential-cache.test.js b/javascript/test/credential-cache.test.js index 23d2478..8651d31 100644 --- a/javascript/test/credential-cache.test.js +++ b/javascript/test/credential-cache.test.js @@ -250,3 +250,20 @@ test('incomplete OAuth reconfiguration clears old timers and never reuses old va assert.equal(time.timerCount, 2); } finally { manager.close(); } }); + +test('closing a credential manager is terminal, including after a configuration change', async (t) => { + const time = clock(); + const getSecret = t.mock.method(SecretClient.prototype, 'getSecret', async () => ({ value: 'PRIVATE-KEY' })); + const manager = new ProviderCredentials({ cacheOptions: time.options }); + const auth = { mode: 'apiKey', keyVaultSecretName: 'key' }; + const config = readConfig({ KEY_VAULT_URL: 'https://unit.vault.azure.net' }); + await manager.resolve(auth, config); + manager.close(); + manager.close(); + for (const settings of [config, { ...config, keyVaultUrl: 'https://other.vault.azure.net' }]) { + await assert.rejects(manager.resolve(auth, settings), /provider credential unavailable/); + } + await assert.rejects(manager.resolve({ mode: 'oauth' }, config), /provider credential unavailable/); + assert.equal(getSecret.mock.callCount(), 1); + assert.equal(time.timerCount, 0); +}); diff --git a/javascript/test/credential-sdk.test.js b/javascript/test/credential-sdk.test.js index fba2d73..8043ab5 100644 --- a/javascript/test/credential-sdk.test.js +++ b/javascript/test/credential-sdk.test.js @@ -1,61 +1,98 @@ 'use strict'; -const { test, mock, beforeEach } = require('node:test'); +const { test, mock, beforeEach, afterEach } = require('node:test'); const assert = require('node:assert/strict'); const Module = require('node:module'); const crypto = require('node:crypto'); +const { spawnSync } = require('node:child_process'); +const path = require('node:path'); const { setTimeout: delay } = require('node:timers/promises'); -const identity = require('@azure/identity'); -const { createHttpHeaders } = require('@azure/core-rest-pipeline'); +const pipeline = require('@azure/core-rest-pipeline'); +const { createHttpHeaders } = pipeline; const { readConfig } = require('../src/functions/config'); let state; const originalLoad = Module._load; const load = mock.method(Module, '_load', function (name, ...args) { - if (name !== '@azure/identity') return originalLoad.call(this, name, ...args); - return { - ...identity, - ManagedIdentityCredential: class extends identity.ManagedIdentityCredential { - constructor(...args) { super(...args); this.testState = state; } - async getToken() { - this.testState.miCalls++; - return { token: 'PRIVATE-ASSERTION', expiresOnTimestamp: Date.now() + 3600000 }; - } - }, - ClientAssertionCredential: class extends identity.ClientAssertionCredential { - constructor(tenant, client, assertion, options) { - super(tenant, client, assertion, { ...options, httpClient: state.transport }); - } - }, - }; + if (name !== '@azure/core-rest-pipeline') return originalLoad.call(this, name, ...args); + return { ...pipeline, createDefaultHttpClient: () => state.transport }; }); let ProviderCredentials; try { ({ ProviderCredentials } = require('../src/functions/credentials')); } finally { load.mock.restore(); } +const envKeys = [ + 'IDENTITY_ENDPOINT', 'IDENTITY_HEADER', 'IDENTITY_SERVER_THUMBPRINT', + 'MSI_ENDPOINT', 'MSI_SECRET', 'AZURE_FEDERATED_TOKEN_FILE', +]; +let savedEnv; beforeEach(() => { - state = { miCalls: 0, tokenCalls: 0, wait: 5, abortObserved: false }; - const current = state; + savedEnv = Object.fromEntries(envKeys.map((key) => [key, process.env[key]])); + for (const key of envKeys) delete process.env[key]; + process.env.IDENTITY_ENDPOINT = 'http://127.0.0.1:1/synthetic-identity'; + process.env.IDENTITY_HEADER = 'synthetic-header'; + state = { miCalls: 0, tokenCalls: 0, vaultCalls: 0, wait: 5, miWait: 5, requests: [] }; state.transport = { async sendRequest(request) { + // MSAL reuses its first managed-identity transport across credential instances. + const current = state; const url = new URL(request.url); - assert.equal(request.method, 'POST'); - assert.ok(url.pathname.endsWith('/oauth2/v2.0/token')); - current.tokenCalls++; - try { await delay(current.wait, null, { signal: request.abortSignal }); } - catch (error) { current.abortObserved = true; throw error; } + if (url.hostname === 'unit.vault.azure.net') { + current.vaultCalls++; + return { request, status: 401, headers: createHttpHeaders({ + 'www-authenticate': 'Bearer authorization="https://login.microsoftonline.com/11111111-1111-1111-1111-111111111111", resource="https://vault.azure.net"', + }) }; + } + const managedIdentity = url.pathname === '/synthetic-identity'; + if (managedIdentity) { + assert.equal(request.method, 'GET'); + current.miCalls++; + } else { + assert.equal(request.method, 'POST'); + assert.ok(url.pathname.endsWith('/oauth2/v2.0/token')); + current.tokenCalls++; + } + const call = { managedIdentity, aborted: false, completed: false }; + current.requests.push(call); + try { + await delay(managedIdentity ? current.miWait : current.wait, null, { signal: request.abortSignal }); + } catch (error) { + call.aborted = true; + throw error; + } finally { + call.completed = true; + } return { request, status: 200, headers: createHttpHeaders({ 'content-type': 'application/json' }), - bodyAsText: JSON.stringify({ access_token: 'PRIVATE-TOKEN', token_type: 'Bearer', - expires_in: 3600, scope: 'api://provider/.default' }) }; + bodyAsText: JSON.stringify({ + access_token: managedIdentity ? 'PRIVATE-ASSERTION' : 'PRIVATE-TOKEN', + token_type: 'Bearer', expires_in: 3600, + expires_on: String(Math.floor(Date.now() / 1000) + 3600), + resource: url.searchParams.get('resource'), scope: 'api://provider/.default', + }) }; }, }; }); +afterEach(() => { + for (const [key, value] of Object.entries(savedEnv)) { + if (value === undefined) delete process.env[key]; + else process.env[key] = value; + } +}); + +async function waitForRequest(managedIdentity) { + for (let i = 0; i < 200; i++) { + if (state.requests.some((request) => request.managedIdentity === managedIdentity)) return; + await delay(5); + } + assert.fail('SDK request did not reach the transport'); +} + function config() { return readConfig({ EPP_PROVIDER_TENANT_ID: '11111111-1111-1111-1111-111111111111', EPP_OUTBOUND_CLIENT_ID: crypto.randomUUID(), - EPP_OUTBOUND_MI_CLIENT_ID: '33333333-3333-3333-3333-333333333333', + EPP_OUTBOUND_MI_CLIENT_ID: crypto.randomUUID(), EPP_PROVIDER_SCOPE: 'api://provider/.default', }); } @@ -82,7 +119,7 @@ test('the cache-owned acquisition budget aborts actual SDK transport and does no try { await assert.rejects(manager.resolve({ mode: 'oauth' }, settings), /^Error: provider OAuth token unavailable$/); await new Promise(setImmediate); - assert.equal(state.abortObserved, true); + assert.ok(state.requests.some((request) => !request.managedIdentity && request.aborted)); assert.deepEqual(failures, ['provider_token']); assert.equal(state.tokenCalls, 1); await assert.rejects(manager.resolve({ mode: 'oauth' }, settings), /unavailable/); @@ -96,13 +133,125 @@ test('changing configuration cancels old work without publishing its token into state.wait = 10000; const first = manager.resolve({ mode: 'oauth' }, settings); const firstRejected = assert.rejects(first, /unavailable/); - await delay(20); + await waitForRequest(false); state.wait = 5; const result = await manager.resolve({ mode: 'oauth' }, { ...settings, outboundClientId: crypto.randomUUID() }); await firstRejected; try { assert.equal(result.accessToken, 'PRIVATE-TOKEN'); - assert.equal(state.abortObserved, true); + assert.ok(state.requests.some((request) => !request.managedIdentity && request.aborted)); assert.equal(state.tokenCalls, 2); } finally { manager.close(); } }); + +test('the acquisition deadline aborts real managed-identity transport before another refresh starts', async () => { + let now = Date.now(); + const manager = new ProviderCredentials({ + reportFailure: () => {}, + cacheOptions: { now: () => now, random: () => 0, schedule: () => ({ unref() {} }), cancel() {} }, + }); + const settings = config(); + state.miWait = 10000; + try { + await assert.rejects(manager.resolve({ mode: 'oauth' }, settings), /unavailable/); + await new Promise(setImmediate); + assert.equal(state.miCalls, 1); + assert.ok(state.requests.every((request) => request.completed && request.aborted)); + assert.equal(state.tokenCalls, 0); + now += 10000; + state.miWait = 5; + assert.equal((await manager.resolve({ mode: 'oauth' }, settings)).accessToken, 'PRIVATE-TOKEN'); + assert.equal(state.miCalls, 2); + assert.equal(state.tokenCalls, 1); + } finally { manager.close(); } +}); + +test('shutdown aborts managed-identity transport and cannot restart acquisition', async () => { + const manager = new ProviderCredentials({ reportFailure: () => {} }); + const settings = config(); + state.miWait = 10000; + const pending = manager.resolve({ mode: 'oauth' }, settings); + const rejected = assert.rejects(pending, /unavailable/); + try { + await waitForRequest(true); + manager.close(); + await rejected; + await new Promise(setImmediate); + assert.ok(state.requests.every((request) => request.completed && request.aborted)); + await assert.rejects(manager.resolve({ mode: 'oauth' }, settings), /unavailable/); + assert.equal(state.miCalls, 1); + assert.equal(state.tokenCalls, 0); + } finally { + manager.close(); + await rejected; + } +}); + +test('Key Vault acquisition also aborts its real managed-identity transport', async () => { + const manager = new ProviderCredentials({ reportFailure: () => {} }); + const settings = readConfig({ + KEY_VAULT_URL: 'https://unit.vault.azure.net', AZURE_CLIENT_ID: crypto.randomUUID(), + }); + state.miWait = 10000; + try { + await assert.rejects(manager.resolve({ mode: 'apiKey', keyVaultSecretName: 'key' }, settings), /unavailable/); + await new Promise(setImmediate); + assert.equal(state.vaultCalls, 1); + assert.equal(state.miCalls, 1); + assert.ok(state.requests.every((request) => request.completed && request.aborted)); + } finally { manager.close(); } +}); + +test('the default Azure HTTP client closes a stalled managed-identity socket on timeout', () => { + const script = ` + const assert = require('node:assert/strict'); + const { createServer } = require('node:http'); + const { setTimeout: delay } = require('node:timers/promises'); + const { ClientAssertionCredential } = require('@azure/identity'); + const { ProviderCredentials } = require('./src/functions/credentials'); + ClientAssertionCredential.prototype.getToken = async function () { + await this.getAssertion(); + throw new Error('The synthetic managed-identity request must not complete'); + }; + (async () => { + let requests = 0; + let onDisconnect; + const disconnected = new Promise(resolve => { onDisconnect = resolve; }); + const sockets = new Set(); + const server = createServer(() => { requests++; }); + server.on('connection', socket => { + sockets.add(socket); + socket.once('close', () => { sockets.delete(socket); onDisconnect(); }); + }); + await new Promise(resolve => server.listen(0, '127.0.0.1', resolve)); + for (const key of ['MSI_ENDPOINT', 'MSI_SECRET', 'IDENTITY_SERVER_THUMBPRINT', 'AZURE_FEDERATED_TOKEN_FILE']) { + delete process.env[key]; + } + process.env.IDENTITY_ENDPOINT = 'http://127.0.0.1:' + server.address().port + '/identity'; + process.env.IDENTITY_HEADER = 'synthetic-header'; + const manager = new ProviderCredentials({ reportFailure: () => {} }); + try { + await assert.rejects(manager.resolve({ mode: 'oauth' }, { + providerTenantId: '11111111-1111-1111-1111-111111111111', + outboundClientId: '22222222-2222-2222-2222-222222222222', + outboundManagedIdentityClientId: '33333333-3333-3333-3333-333333333333', + providerScope: 'api://provider/.default', + }), /unavailable/); + assert.equal(requests, 1); + const closed = await Promise.race([disconnected.then(() => true), delay(1000, false, { ref: false })]); + assert.equal(closed, true); + console.log('transport-aborted'); + } finally { + manager.close(); + for (const socket of sockets) socket.destroy(); + await new Promise(resolve => server.close(resolve)); + } + })().catch(error => { console.error(error); process.exitCode = 1; }); + `; + const result = spawnSync(process.execPath, ['-e', script], { + cwd: path.resolve(__dirname, '..'), encoding: 'utf8', timeout: 10000, + }); + assert.equal(result.error, undefined); + assert.equal(result.status, 0, result.stderr); + assert.equal(result.stdout.trim(), 'transport-aborted'); +}); diff --git a/javascript/test/dispatch.test.js b/javascript/test/dispatch.test.js index 28977fd..3be56b2 100644 --- a/javascript/test/dispatch.test.js +++ b/javascript/test/dispatch.test.js @@ -1,6 +1,6 @@ 'use strict'; -const { test, afterEach } = require('node:test'); +const { test, beforeEach, afterEach } = require('node:test'); const assert = require('node:assert/strict'); const { ClientAssertionCredential, ManagedIdentityCredential } = require('@azure/identity'); const { SecretClient } = require('@azure/keyvault-secrets'); @@ -9,12 +9,17 @@ const { DeliveryContext, ParsedResponse } = require('../src/functions/models'); const fixtures = require('../../tests/fixtures/contract.json'); const { inspect } = require('node:util'); const { AzureLogger } = require('@azure/logger'); +const { ProviderCredentials, providerCredentials } = require('../src/functions/credentials'); const { dispatchOtp, getProvider, resolveOutcome, outcomeToHttpStatus, parseEnvelope, parseProviderTimeout, isValidProviderUrl, contextToDispatch, resolveProviderCredential, - stopProviderCredentialRefresh, } = require('../src/functions/dispatch'); -afterEach(stopProviderCredentialRefresh); +let credentials; +beforeEach((t) => { + credentials = new ProviderCredentials(); + t.mock.method(providerCredentials, 'resolve', (...args) => credentials.resolve(...args)); +}); +afterEach(() => credentials.close()); const dispatch = { destination: '+15551234567', message: ' Your code is 918273.\n', channel: 'sms', messageId: 'message-id', correlationId: 'correlation-id' }; const input = { channel: 'sms', endpoint: 'https://provider.example', dispatch, @@ -298,7 +303,8 @@ test('Soprano OAuth reuses setup identities and selected scope with private boun for (const invalid of [null, { token: '' }, { token: ' ' }, { token: false }, { token: 'stale', expiresOnTimestamp: Date.now() + 10000 }, { token: 'missing-expiry' }]) { const method = stage === 'token' ? providerToken : identityToken; - stopProviderCredentialRefresh(); + credentials.close(); + credentials = new ProviderCredentials(); method.mock.mockImplementation(async () => invalid); if (stage === 'assertion') providerToken.mock.mockImplementation(async function () { await this.getAssertion(); diff --git a/javascript/test/sendotp.test.js b/javascript/test/sendotp.test.js index 72770c4..5836d9c 100644 --- a/javascript/test/sendotp.test.js +++ b/javascript/test/sendotp.test.js @@ -8,7 +8,8 @@ const { CompactEncrypt } = require('jose'); const { ClientAssertionCredential, ManagedIdentityCredential } = require('@azure/identity'); const { SecretClient } = require('@azure/keyvault-secrets'); const fixtures = require('../../tests/fixtures/contract.json'); -const { getProvider, stopProviderCredentialRefresh } = require('../src/functions/dispatch'); +const { getProvider } = require('../src/functions/dispatch'); +const { ProviderCredentials, providerCredentials } = require('../src/functions/credentials'); const { RequestLog } = require('../src/functions/requestLog'); // Capture the real handler; keys stay in memory and all external I/O is mocked. @@ -42,8 +43,11 @@ let logs; let warnings; let records; let getToken; +let credentials; beforeEach(() => { - stopProviderCredentialRefresh(); + credentials = new ProviderCredentials(); + mock.method(providerCredentials, 'resolve', (...args) => credentials.resolve(...args)); + mock.method(providerCredentials, 'close', () => credentials.close()); savedEnv = Object.fromEntries(envKeys.map((key) => [key, process.env[key]])); for (const key of envKeys) delete process.env[key]; Object.assign(process.env, { EPP_LOG_PLAINTEXT: 'true', @@ -66,7 +70,7 @@ beforeEach(() => { text: async () => JSON.stringify({ status: 'ENROUTE', id: 'provider-reference-id', description: 'PRIVATE-STATUS' }) })); }); afterEach(() => { - stopProviderCredentialRefresh(); + credentials.close(); mock.restoreAll(); for (const [key, value] of Object.entries(savedEnv)) { if (value === undefined) delete process.env[key]; @@ -146,6 +150,9 @@ test('worker startup preloads credentials without delivery and leaves evaluation assert.equal(getToken.mock.callCount(), 1); assert.equal(fetchMock.mock.callCount(), 1); stopHook(); + assertFailure(await invoke(await envelope()), 502); + assert.equal(getToken.mock.callCount(), 1); + assert.equal(fetchMock.mock.callCount(), 1); }); test('worker startup without a configured provider does not acquire any credentials', async () => { diff --git a/python/README.md b/python/README.md index b0f6e92..a1ab622 100644 --- a/python/README.md +++ b/python/README.md @@ -92,9 +92,12 @@ six-digit numeric run that is not part of a longer number and repeats the comple Worker initialization starts background credential preparation when a provider is configured. Key Vault bundles, managed-identity assertions and final Entra tokens use separate process-local -caches with daemon refresh timers. A caller can stop waiting without cancelling shared retrieval; -the HTTP SDK still uses connect/read inactivity timeouts, not a total transport deadline. `atexit` -stops scheduled work and prevents late cache publication. See the +caches with daemon refresh timers and parallel daemon secret reads. A caller can stop waiting +without cancelling shared retrieval or starting overlapping reads; the HTTP SDK still uses +connect/read inactivity timeouts, not a total transport deadline. `get_token_info`, when supported +by the installed SDK, preserves early refresh hints; older SDKs retain the pre-expiry refresh target. +`atexit` stops scheduled work, releases waiters and prevents late cache publication. Pending +synchronous reads do not block process exit. Closing a credential manager is terminal. See the [refresh contract](../docs/CONTRACT.md#credential-caching-and-refresh). Evaluation handling stays independent; leave the provider unset for local evaluation without background credential acquisition. diff --git a/python/src/credentials.py b/python/src/credentials.py index f36c23e..6e4dfeb 100644 --- a/python/src/credentials.py +++ b/python/src/credentials.py @@ -1,31 +1,75 @@ +from __future__ import annotations + import json import logging import threading import time -from concurrent.futures import ThreadPoolExecutor +from collections.abc import Callable, Mapping +from concurrent.futures import Future, wait from contextvars import ContextVar +from dataclasses import dataclass, field +from typing import Literal, Protocol, TypedDict, TypeVar +from azure.core.credentials import TokenCredential from azure.identity import ClientAssertionCredential, ManagedIdentityCredential -from .refreshing_cache import CacheEntry, RefreshingCache, token_entry +from .config import AppConfig +from .refreshing_cache import ( + ACQUISITION_TIMEOUT_SECONDS, + SECRET_REFRESH_INTERVAL_SECONDS, + SECRET_TTL_SECONDS, + CacheEntry, + CacheOptions, + RefreshingCache, + Token, + token_entry, +) _acquiring = ContextVar("epp_credential_acquisition", default=False) +T = TypeVar("T") + + +class SecretReader(Protocol): + def resolve(self, secret_name: str) -> str: ... + + +class ApiKeyCredential(TypedDict): + mode: Literal["apiKey"] + secret: str + identity: str + + +class OAuthCredential(TypedDict): + mode: Literal["oauth"] + access_token: str + + +@dataclass(repr=False) +class _ApiKeyState: + bundle: RefreshingCache[ApiKeyCredential] + + +@dataclass(repr=False) +class _OAuthState: + assertion: RefreshingCache[Token] + credential: TokenCredential + tokens: dict[str, RefreshingCache[Token]] = field(default_factory=dict) class _CredentialLogFilter(logging.Filter): - def filter(self, record): + def filter(self, record: logging.LogRecord) -> bool: return not (_acquiring.get() and record.name.startswith(("azure.identity", "azure.core", "msal"))) _log_filter = _CredentialLogFilter() -def report_refresh_failure(kind): +def report_refresh_failure(kind: str) -> None: logging.warning("%s", json.dumps({"logType": "service", "eventName": "credential_refresh_failed", "cacheKind": kind, "failureReason": "credential_unavailable"})) -def _private_acquisition(load): +def _private_acquisition(load: Callable[[], T]) -> T: for logger in (logging.getLogger(), *logging.Logger.manager.loggerDict.copy().values()): if isinstance(logger, logging.Logger): for handler in logger.handlers: @@ -39,90 +83,152 @@ def _private_acquisition(load): class ProviderCredentials: - def __init__(self, secrets, *, cache_options=None, report_failure=report_refresh_failure): + def __init__( + self, + secrets: SecretReader, + *, + cache_options: CacheOptions | None = None, + report_failure: Callable[[str], None] = report_refresh_failure, + ) -> None: self._secrets = secrets - self._options = cache_options or {} + self._options: CacheOptions = cache_options or {} self._clock = self._options.get("clock", time.time) self._report_failure = report_failure self._lock = threading.RLock() - self._key = None - self._state = None + self._key: tuple[str | None, ...] | None = None + self._state: _ApiKeyState | _OAuthState | None = None + self._closed = False - def _cache(self, kind, load): - return RefreshingCache(lambda: _private_acquisition(load), **{ - **self._options, "on_failure": lambda: self._report_failure(kind), - }) + def _cache(self, kind: str, load: Callable[[], CacheEntry[T]]) -> RefreshingCache[T]: + options = self._options.copy() + options["on_failure"] = lambda: self._report_failure(kind) + return RefreshingCache(lambda: _private_acquisition(load), **options) - def resolve(self, auth, config): + def resolve(self, auth: Mapping[str, str], config: AppConfig) -> ApiKeyCredential | OAuthCredential: mode = auth.get("mode") - if mode == "oauth" and not all((config.provider_tenant_id, config.provider_scope, - config.outbound_client_id, config.outbound_managed_identity_client_id)): - self.close() - raise ValueError("provider OAuth token unavailable") - if mode not in ("apiKey", "oauth"): - self.close() - raise ValueError("provider credential unavailable") - key = ((mode, config.env.get("KEY_VAULT_URL"), config.env.get("AZURE_CLIENT_ID"), - auth.get("key_vault_secret_name"), auth.get("identity_key_vault_secret_name")) if mode == "apiKey" else - (mode, config.provider_tenant_id, config.outbound_client_id, config.outbound_managed_identity_client_id)) + scope = config.provider_scope with self._lock: + if self._closed: + raise ValueError("provider credential unavailable") + if mode == "oauth" and not all(( + config.provider_tenant_id, scope, + config.outbound_client_id, config.outbound_managed_identity_client_id, + )): + self._clear() + raise ValueError("provider OAuth token unavailable") + if mode not in ("apiKey", "oauth"): + self._clear() + raise ValueError("provider credential unavailable") + key: tuple[str | None, ...] + if mode == "apiKey": + key = ( + mode, config.env.get("KEY_VAULT_URL"), config.env.get("AZURE_CLIENT_ID"), + auth.get("key_vault_secret_name"), auth.get("identity_key_vault_secret_name"), + ) + else: + key = ( + mode, config.provider_tenant_id, + config.outbound_client_id, config.outbound_managed_identity_client_id, + ) if self._key != key: - self.close() - self._state = self._api_key_state(auth) if mode == "apiKey" else _private_acquisition(lambda: self._oauth_state(config)) + self._clear() + if mode == "apiKey": + self._state = self._api_key_state(auth) + else: + self._state = _private_acquisition(lambda: self._oauth_state(config)) self._key = key state = self._state - if mode == "apiKey": - cache = state["bundle"] - else: - scope = config.provider_scope - if scope not in state["tokens"]: - state["tokens"][scope] = self._cache("provider_token", lambda: - token_entry(state["credential"].get_token(scope, logging_enable=False), self._clock())) - cache = state["tokens"][scope] + if isinstance(state, _OAuthState) and scope not in state.tokens: + credential = state.credential + state.tokens[scope] = self._cache("provider_token", lambda: self._load_token(credential, scope)) try: - value = cache.get() + if isinstance(state, _ApiKeyState): + return state.bundle.get().copy() + if isinstance(state, _OAuthState): + return {"mode": "oauth", "access_token": state.tokens[scope].get().token} + raise ValueError("provider credential unavailable") except Exception: - raise ValueError("provider OAuth token unavailable" if mode == "oauth" else "provider credential unavailable") from None - return dict(value) if mode == "apiKey" else {"mode": "oauth", "access_token": value.token} - - def _api_key_state(self, auth): - def load(): + reason = "provider OAuth token unavailable" if mode == "oauth" else "provider credential unavailable" + raise ValueError(reason) from None + + def _start_secret_read(self, name: str) -> Future[str]: + future: Future[str] = Future() + + def read() -> None: + try: + with self._lock: + if self._closed: + raise ValueError("provider credential unavailable") + value = _private_acquisition(lambda: self._secrets.resolve(name)) + future.set_result(value) + except Exception: + future.set_exception(ValueError("provider credential unavailable")) + + # Executor workers are joined before atexit, even when their parent is a daemon. + threading.Thread(target=read, daemon=True).start() + return future + + def _api_key_state(self, auth: Mapping[str, str]) -> _ApiKeyState: + def load() -> CacheEntry[ApiKeyCredential]: key_name = auth.get("key_vault_secret_name") identity_name = auth.get("identity_key_vault_secret_name") if not key_name: raise ValueError("provider credential unavailable") - # Both secrets form one snapshot; do not publish a partial rotation. - with ThreadPoolExecutor(max_workers=2) as pool: - key_future = pool.submit(_private_acquisition, lambda: self._secrets.resolve(key_name)) - identity_future = pool.submit(_private_acquisition, lambda: self._secrets.resolve(identity_name)) if identity_name else None - secret = key_future.result() - identity = identity_future.result() if identity_future else "" + key_future = self._start_secret_read(key_name) + identity_future = self._start_secret_read(identity_name) if identity_name else None + pending = [key_future] + if identity_future is not None: + pending.append(identity_future) + # Keep one refresh in flight until both reads finish, even after a caller stops waiting. + wait(pending) + secret = key_future.result() + identity = identity_future.result() if identity_future is not None else "" if not isinstance(secret, str) or not secret.strip() or (identity_name and ( not isinstance(identity, str) or not identity.strip())): raise ValueError("provider credential unavailable") now = self._clock() - return CacheEntry({"mode": "apiKey", "secret": secret, "identity": identity}, now + 300, now + 240) - return {"bundle": self._cache("key_vault", load)} - - def _oauth_state(self, config): - identity = ManagedIdentityCredential(client_id=config.outbound_managed_identity_client_id, - retry_total=0, connection_timeout=2.5, read_timeout=2.5, logging_enable=False) - assertion = self._cache("managed_identity", lambda: - token_entry(identity.get_token("api://AzureADTokenExchange/.default", logging_enable=False), self._clock())) + value: ApiKeyCredential = {"mode": "apiKey", "secret": secret, "identity": identity} + return CacheEntry(value, now + SECRET_TTL_SECONDS, now + SECRET_REFRESH_INTERVAL_SECONDS) + return _ApiKeyState(self._cache("key_vault", load)) + + def _load_token(self, credential: TokenCredential, scope: str) -> CacheEntry[Token]: + get_token_info = getattr(credential, "get_token_info", None) + if callable(get_token_info): + token = get_token_info(scope) + else: + # Older supported SDKs expose only expiry metadata through get_token. + token = credential.get_token(scope, logging_enable=False) + return token_entry(token, self._clock()) + + def _oauth_state(self, config: AppConfig) -> _OAuthState: + identity = ManagedIdentityCredential( + client_id=config.outbound_managed_identity_client_id, + retry_total=0, connection_timeout=ACQUISITION_TIMEOUT_SECONDS, + read_timeout=ACQUISITION_TIMEOUT_SECONDS, logging_enable=False, + ) + assertion = self._cache( + "managed_identity", lambda: self._load_token(identity, "api://AzureADTokenExchange/.default"), + ) credential = ClientAssertionCredential( tenant_id=config.provider_tenant_id, client_id=config.outbound_client_id, func=lambda: assertion.get().token, authority="https://login.microsoftonline.com", - retry_total=0, connection_timeout=2.5, read_timeout=2.5, logging_enable=False) - return {"assertion": assertion, "credential": credential, "tokens": {}} + retry_total=0, connection_timeout=ACQUISITION_TIMEOUT_SECONDS, + read_timeout=ACQUISITION_TIMEOUT_SECONDS, logging_enable=False, + ) + return _OAuthState(assertion, credential) + + def _clear(self) -> None: + state = self._state + if isinstance(state, _ApiKeyState): + state.bundle.close() + elif isinstance(state, _OAuthState): + state.assertion.close() + for cache in state.tokens.values(): + cache.close() + self._state = None + self._key = None - def close(self): + def close(self) -> None: with self._lock: - if self._state: - if "bundle" in self._state: - self._state["bundle"].close() - else: - self._state["assertion"].close() - for cache in self._state["tokens"].values(): - cache.close() - self._state = None - self._key = None + self._closed = True + self._clear() diff --git a/python/src/refreshing_cache.py b/python/src/refreshing_cache.py index bd8f8f2..f2a82ac 100644 --- a/python/src/refreshing_cache.py +++ b/python/src/refreshing_cache.py @@ -1,28 +1,75 @@ +from __future__ import annotations + import math import random import threading import time +from collections.abc import Callable from concurrent.futures import Future, TimeoutError from dataclasses import dataclass +from typing import Generic, Protocol, TypedDict, TypeVar + +ACQUISITION_TIMEOUT_SECONDS = 2.5 +SECRET_TTL_SECONDS = 5 * 60 +SECRET_REFRESH_INTERVAL_SECONDS = 4 * 60 +TOKEN_EXPIRY_SKEW_SECONDS = 30 +TOKEN_REFRESH_LEAD_SECONDS = 5 * 60 +MIN_REFRESH_DELAY_SECONDS = 1 +MAX_REFRESH_DELAY_SECONDS = 60 +INITIAL_RETRY_DELAY_SECONDS = 5 +MAX_RETRY_DELAY_SECONDS = 60 +MAX_RETRY_EXPONENT = 4 +RETRY_JITTER_RATIO = 0.2 +MAX_TIMER_DELAY_SECONDS = (2 ** 31 - 1) / 1000 + +T = TypeVar("T") + + +class Token(Protocol): + @property + def token(self) -> str: ... + + @property + def expires_on(self) -> float: ... + + +class ScheduledCall(Protocol): + def cancel(self) -> None: ... + + +class CacheOptions(TypedDict, total=False): + clock: Callable[[], float] + schedule: Callable[[float, Callable[[], None]], ScheduledCall] + jitter: Callable[[], float] + on_failure: Callable[[], None] + wait_timeout: float @dataclass(repr=False) -class CacheEntry: - value: object +class CacheEntry(Generic[T]): + value: T expires_at: float refresh_at: float -def _schedule(delay, callback): +def _schedule(delay: float, callback: Callable[[], None]) -> ScheduledCall: timer = threading.Timer(delay, callback) timer.daemon = True timer.start() return timer -class RefreshingCache: - def __init__(self, load, *, clock=time.time, schedule=_schedule, jitter=random.random, - on_failure=lambda: None, wait_timeout=2.5): +class RefreshingCache(Generic[T]): + def __init__( + self, + load: Callable[[], CacheEntry[T]], + *, + clock: Callable[[], float] = time.time, + schedule: Callable[[float, Callable[[], None]], ScheduledCall] = _schedule, + jitter: Callable[[], float] = random.random, + on_failure: Callable[[], None] = lambda: None, + wait_timeout: float = ACQUISITION_TIMEOUT_SECONDS, + ) -> None: self._load = load self._clock = clock self._schedule = schedule @@ -30,14 +77,14 @@ def __init__(self, load, *, clock=time.time, schedule=_schedule, jitter=random.r self._on_failure = on_failure self._wait_timeout = wait_timeout self._lock = threading.RLock() - self._entry = None - self._inflight = None - self._timer = None + self._entry: CacheEntry[T] | None = None + self._inflight: Future[T] | None = None + self._timer: ScheduledCall | None = None self._retry_at = 0 self._failures = 0 self._closed = False - def get(self): + def get(self) -> T: with self._lock: if self._closed: raise ValueError("provider credential unavailable") @@ -55,82 +102,91 @@ def get(self): # A waiter does not cancel the shared refresh needed by other requests. raise ValueError("provider credential unavailable") from None - def refresh(self): + def refresh(self) -> Future[T]: with self._lock: if self._closed or self._retry_at > self._clock(): raise ValueError("provider credential unavailable") return self._begin_refresh() - def _begin_refresh(self): + def _begin_refresh(self) -> Future[T]: if self._inflight is not None: return self._inflight if self._timer is not None: self._timer.cancel() self._timer = None - future = Future() + future: Future[T] = Future() self._inflight = future threading.Thread(target=self._run_refresh, args=(future,), daemon=True).start() return future - def _run_refresh(self, future): - entry = None - failed = False + def _run_refresh(self, future: Future[T]) -> None: + with self._lock: + if self._closed: + self._inflight = None + return + entry: CacheEntry[T] | None = None try: entry = self._load() if (not isinstance(entry, CacheEntry) or not math.isfinite(entry.expires_at) or entry.expires_at <= self._clock() or not math.isfinite(entry.refresh_at)): raise ValueError("provider credential unavailable") except Exception: - failed = True + entry = None with self._lock: if self._closed: if not future.done(): future.set_exception(ValueError("provider credential unavailable")) self._inflight = None return - if failed: + if entry is None: self._failures += 1 - backoff = min(60, 5 * 2 ** min(self._failures - 1, 4)) - self._retry_at = self._clock() + backoff * (1 + self._jitter() * 0.2) + exponent = min(self._failures - 1, MAX_RETRY_EXPONENT) + backoff = min(MAX_RETRY_DELAY_SECONDS, INITIAL_RETRY_DELAY_SECONDS * 2 ** exponent) + self._retry_at = self._clock() + backoff * (1 + self._jitter() * RETRY_JITTER_RATIO) self._on_failure() else: self._entry = entry self._failures = 0 self._retry_at = 0 self._inflight = None - next_refresh = self._retry_at if failed else entry.refresh_at - self._timer = self._schedule(max(1, next_refresh - self._clock()), self._scheduled_refresh) - if failed: + next_refresh = self._retry_at if entry is None else entry.refresh_at + delay = min(MAX_TIMER_DELAY_SECONDS, next_refresh - self._clock()) + self._timer = self._schedule(max(MIN_REFRESH_DELAY_SECONDS, delay), self._scheduled_refresh) + if entry is None: future.set_exception(ValueError("provider credential unavailable")) else: future.set_result(entry.value) - def _scheduled_refresh(self): + def _scheduled_refresh(self) -> None: with self._lock: self._timer = None if not self._closed: self._begin_refresh() - def close(self): + def close(self) -> None: with self._lock: self._closed = True if self._timer is not None: self._timer.cancel() self._timer = None self._entry = None + if self._inflight is not None and not self._inflight.done(): + self._inflight.set_exception(ValueError("provider credential unavailable")) -def token_entry(token, now): +def token_entry(token: Token, now: float) -> CacheEntry[Token]: value = getattr(token, "token", None) expiry = getattr(token, "expires_on", None) - if (not isinstance(value, str) or not value.strip() or type(expiry) not in (int, float) - or not math.isfinite(expiry) or expiry <= now + 30): + if (not isinstance(value, str) or not value.strip() + or not isinstance(expiry, (int, float)) or isinstance(expiry, bool) + or not math.isfinite(expiry) or expiry <= now + TOKEN_EXPIRY_SKEW_SECONDS): raise ValueError("provider credential unavailable") - expires_at = expiry - 30 - refresh_at = expiry - 300 + expires_at = expiry - TOKEN_EXPIRY_SKEW_SECONDS + refresh_at = expiry - TOKEN_REFRESH_LEAD_SECONDS hint = getattr(token, "refresh_on", None) - if type(hint) in (int, float) and math.isfinite(hint): + if isinstance(hint, (int, float)) and not isinstance(hint, bool) and math.isfinite(hint): refresh_at = min(refresh_at, hint) if refresh_at <= now: - refresh_at = now + max(1, min(60, (expires_at - now) / 2)) + delay = min(MAX_REFRESH_DELAY_SECONDS, (expires_at - now) / 2) + refresh_at = now + max(MIN_REFRESH_DELAY_SECONDS, delay) return CacheEntry(token, expires_at, refresh_at) diff --git a/python/src/secrets.py b/python/src/secrets.py index ca06f2b..ed97b63 100644 --- a/python/src/secrets.py +++ b/python/src/secrets.py @@ -1,31 +1,41 @@ +from __future__ import annotations + import os +from collections.abc import Mapping from threading import Lock from azure.identity import ManagedIdentityCredential from azure.keyvault.secrets import SecretClient +from .refreshing_cache import ACQUISITION_TIMEOUT_SECONDS + + class SecretResolver: - def __init__(self, env=None): + def __init__(self, env: Mapping[str, str] | None = None) -> None: self._env = env if env is not None else os.environ - self._client = None - self._client_key = None + self._client: SecretClient | None = None + self._client_key: tuple[str, str | None] | None = None self._lock = Lock() - def _get_client(self): + def _get_client(self) -> SecretClient: vault_url = self._env.get("KEY_VAULT_URL") client_id = self._env.get("AZURE_CLIENT_ID") if not vault_url: raise RuntimeError("KEY_VAULT_URL not set") with self._lock: - if self._client_key != (vault_url, client_id): - credential = ManagedIdentityCredential(client_id=client_id, logging_enable=False, - retry_total=0, connection_timeout=2.5, read_timeout=2.5) - self._client = SecretClient(vault_url=vault_url, credential=credential, - retry_total=0, connection_timeout=2.5, read_timeout=2.5, logging_enable=False) + if self._client is None or self._client_key != (vault_url, client_id): + credential = ManagedIdentityCredential( + client_id=client_id, logging_enable=False, retry_total=0, + connection_timeout=ACQUISITION_TIMEOUT_SECONDS, read_timeout=ACQUISITION_TIMEOUT_SECONDS, + ) + self._client = SecretClient( + vault_url=vault_url, credential=credential, retry_total=0, logging_enable=False, + connection_timeout=ACQUISITION_TIMEOUT_SECONDS, read_timeout=ACQUISITION_TIMEOUT_SECONDS, + ) self._client_key = (vault_url, client_id) return self._client - def resolve(self, secret_name): + def resolve(self, secret_name: str | None) -> str: if not secret_name: return "" # ProviderCredentials caches the complete credential bundle and owns refresh. diff --git a/python/tests/test_credential_cache.py b/python/tests/test_credential_cache.py index 4082b26..0aeff67 100644 --- a/python/tests/test_credential_cache.py +++ b/python/tests/test_credential_cache.py @@ -1,6 +1,10 @@ +import subprocess +import sys +import textwrap import time from concurrent.futures import ThreadPoolExecutor from threading import Event, Lock +from pathlib import Path from types import SimpleNamespace from unittest.mock import Mock @@ -188,7 +192,7 @@ def oauth_config(): def test_both_mi_and_provider_token_caches_refresh_independently_and_skip_warm_sdk_calls(monkeypatch): clock = Clock() - identity = Mock(get_token=Mock(side_effect=lambda *args, **kwargs: + identity = Mock(spec=["get_token"], get_token=Mock(side_effect=lambda *args, **kwargs: SimpleNamespace(token="PRIVATE-ASSERTION", expires_on=clock.now + 3600))) monkeypatch.setattr(credentials_module, "ManagedIdentityCredential", Mock(return_value=identity)) clients = [] @@ -198,7 +202,7 @@ def get_token(*args, **options): assert kwargs["func"]() == "PRIVATE-ASSERTION" assert kwargs["func"]() == "PRIVATE-ASSERTION" return SimpleNamespace(token="PRIVATE-TOKEN", expires_on=clock.now + 3600) - client = Mock(get_token=Mock(side_effect=get_token)) + client = Mock(spec=["get_token"], get_token=Mock(side_effect=get_token)) clients.append(client) return client @@ -292,3 +296,151 @@ def test_startup_only_prepares_credentials_and_handles_missing_provider_or_failu broken.start_credential_refresh() finally: broken.close() + + +def test_token_info_refresh_hints_apply_to_both_caches_without_legacy_sdk_calls(monkeypatch): + clock = Clock() + + def token_info(value): + return SimpleNamespace(token=value, expires_on=clock.now + 3600, refresh_on=clock.now + 60) + + identity = Mock( + spec=["get_token", "get_token_info"], + get_token_info=Mock(side_effect=lambda scope: token_info("PRIVATE-ASSERTION")), + ) + monkeypatch.setattr(credentials_module, "ManagedIdentityCredential", Mock(return_value=identity)) + clients = [] + + def create(**kwargs): + def get_token_info(scope): + assert kwargs["func"]() == "PRIVATE-ASSERTION" + return token_info("PRIVATE-PROVIDER") + credential = Mock(spec=["get_token", "get_token_info"], get_token_info=Mock(side_effect=get_token_info)) + clients.append(credential) + return credential + + monkeypatch.setattr(credentials_module, "ClientAssertionCredential", create) + manager = ProviderCredentials(Mock(), cache_options=clock.options) + config = oauth_config() + try: + manager.resolve({"mode": "oauth"}, config) + state = manager._state + assert state.assertion._entry.refresh_at == clock.now + 60 + assert state.tokens[config.provider_scope]._entry.refresh_at == clock.now + 60 + assert "PRIVATE" not in repr(state) + repr(state.assertion._entry) + clock.advance(60) + wait_until(lambda: identity.get_token_info.call_count == clients[0].get_token_info.call_count == 2) + wait_until(lambda: clock.timer_count == 2) + identity.get_token.assert_not_called() + clients[0].get_token.assert_not_called() + config.provider_scope = "" + with pytest.raises(ValueError, match="OAuth token unavailable"): + manager.resolve({"mode": "oauth"}, config) + assert clock.timer_count == 0 + config.provider_scope = "api://provider/.default" + assert manager.resolve({"mode": "oauth"}, config)["access_token"] == "PRIVATE-PROVIDER" + assert len(clients) == 2 + finally: + manager.close() + + +def test_close_releases_waiters_immediately_and_never_reopens_the_manager(): + clock = Clock() + started, release = Event(), Event() + + def load(): + started.set() + assert release.wait(3) + return CacheEntry("late", clock.now + 300, clock.now + 240) + + cache = RefreshingCache(load, **clock.options) + future = cache.refresh() + try: + assert started.wait(3) + cache.close() + with pytest.raises(ValueError, match="unavailable"): + future.result(timeout=0.1) + assert clock.timer_count == 0 + finally: + release.set() + cache.close() + + secrets = Mock(resolve=Mock(return_value="PRIVATE-KEY")) + manager = ProviderCredentials(secrets, cache_options=clock.options) + auth = {"mode": "apiKey", "key_vault_secret_name": "key"} + config = read_config({"KEY_VAULT_URL": "https://unit.vault.azure.net"}) + manager.resolve(auth, config) + manager.close() + manager.close() + for settings in (config, read_config({"KEY_VAULT_URL": "https://other.vault.azure.net"})): + with pytest.raises(ValueError, match="unavailable"): + manager.resolve(auth, settings) + with pytest.raises(ValueError, match="unavailable"): + manager.resolve({"mode": "oauth"}, config) + secrets.resolve.assert_called_once() + assert clock.timer_count == 0 + + +def test_pending_secret_reads_remain_single_flight_after_waiters_leave(): + clock = Clock() + started, release = Event(), Event() + calls = [] + + def resolve(name): + calls.append(name) + if name == "key": + raise ValueError("PRIVATE-FAILURE") + started.set() + assert release.wait(3) + return "PRIVATE-IDENTITY" + + manager = ProviderCredentials( + Mock(resolve=resolve), cache_options={**clock.options, "wait_timeout": 0.02}, + report_failure=lambda kind: None, + ) + auth = {"mode": "apiKey", "key_vault_secret_name": "key", "identity_key_vault_secret_name": "id"} + config = read_config({}) + try: + for _ in range(3): + with pytest.raises(ValueError, match="unavailable"): + manager.resolve(auth, config) + clock.advance(60) + assert started.is_set() + assert sorted(calls) == ["id", "key"] + assert clock.timer_count == 0 + finally: + manager.close() + release.set() + + +def test_pending_secret_reads_do_not_block_process_shutdown(): + script = textwrap.dedent(""" + import atexit + from threading import Event + from types import SimpleNamespace + from src.config import read_config + from src.credentials import ProviderCredentials + + started, blocked = Event(), Event() + def resolve(name): + started.set() + blocked.wait() + return "synthetic-key" + + manager = ProviderCredentials(SimpleNamespace(resolve=resolve), cache_options={"wait_timeout": 0.02}) + atexit.register(lambda: print("shutdown-complete", flush=True)) + atexit.register(manager.close) + try: + manager.resolve({"mode": "apiKey", "key_vault_secret_name": "key"}, read_config({})) + except ValueError: + pass + assert started.wait(1) + manager.close() + print("main-finished", flush=True) + """) + result = subprocess.run( + [sys.executable, "-c", script], cwd=Path(__file__).resolve().parents[1], + capture_output=True, text=True, timeout=10, + ) + assert result.returncode == 0, result.stderr + assert result.stdout.splitlines() == ["main-finished", "shutdown-complete"] diff --git a/python/tests/test_credential_sdk.py b/python/tests/test_credential_sdk.py index 4f0f534..fcf4581 100644 --- a/python/tests/test_credential_sdk.py +++ b/python/tests/test_credential_sdk.py @@ -5,13 +5,16 @@ from unittest.mock import Mock import requests +import pytest +from azure.identity import ClientAssertionCredential import src.credentials as credentials_module from src.config import read_config from src.credentials import ProviderCredentials -def test_real_provider_sdk_uses_one_exchange_for_concurrent_requests_and_no_exchange_when_warm(monkeypatch): +@pytest.mark.parametrize("refresh_in", [None, 60]) +def test_real_provider_sdk_reuses_tokens_and_preserves_refresh_metadata(monkeypatch, refresh_in): token_endpoint_calls = [] mi_calls = [] @@ -28,7 +31,10 @@ def send(_session, method, url, **kwargs): if method == "POST" and url.endswith("/oauth2/v2.0/token"): token_endpoint_calls.append(url) time.sleep(0.03) - return response({"access_token": "PRIVATE-PROVIDER", "expires_in": 3600, "token_type": "Bearer"}) + payload = {"access_token": "PRIVATE-PROVIDER", "expires_in": 3600, "token_type": "Bearer"} + if refresh_in is not None: + payload["refresh_in"] = refresh_in + return response(payload) if method == "GET" and ".well-known/openid-configuration" in url: return response({ "token_endpoint": "https://login.microsoftonline.com/11111111-1111-1111-1111-111111111111/oauth2/v2.0/token", @@ -43,7 +49,7 @@ def managed(*args, **kwargs): monkeypatch.setattr(requests.Session, "request", send) monkeypatch.setattr(credentials_module, "ManagedIdentityCredential", Mock( - return_value=Mock(get_token=Mock(side_effect=managed)))) + return_value=Mock(spec=["get_token"], get_token=Mock(side_effect=managed)))) secrets = Mock() manager = ProviderCredentials(secrets) config = read_config({ @@ -60,6 +66,15 @@ def managed(*args, **kwargs): assert len(mi_calls) == 1 assert manager.resolve({"mode": "oauth"}, config)["access_token"] == "PRIVATE-PROVIDER" assert len(token_endpoint_calls) == len(mi_calls) == 1 + entry = manager._state.tokens[config.provider_scope]._entry + if refresh_in is not None and hasattr(ClientAssertionCredential, "get_token_info"): + info = manager._state.credential.get_token_info(config.provider_scope) + assert info.refresh_on is not None + assert entry.refresh_at == info.refresh_on + assert entry.refresh_at < entry.expires_at - 3000 + else: + assert entry.refresh_at == entry.value.expires_on - 300 + assert len(token_endpoint_calls) == len(mi_calls) == 1 secrets.resolve.assert_not_called() finally: manager.close() diff --git a/python/tests/test_engine.py b/python/tests/test_engine.py index cf25b58..25b258d 100644 --- a/python/tests/test_engine.py +++ b/python/tests/test_engine.py @@ -58,7 +58,7 @@ def test_soprano_oauth_uses_setup_settings_and_rejects_unusable_tokens(engine, m engine._resolve_credential = DispatchEngine._resolve_credential.__get__(engine, DispatchEngine) assertion = SimpleNamespace(token="private-assertion", expires_on=time.time() + 3600) access = SimpleNamespace(token="private-token", expires_on=time.time() + 3600) - managed = Mock(get_token=Mock(side_effect=lambda *args, **kwargs: assertion)) + managed = Mock(spec=["get_token"], get_token=Mock(side_effect=lambda *args, **kwargs: assertion)) identity_factory = Mock(return_value=managed) clients = [] @@ -73,7 +73,7 @@ def get_token(*args, **options): assert kwargs["func"]() == "private-assertion" return access - client = Mock(get_token=Mock(side_effect=get_token)) + client = Mock(spec=["get_token"], get_token=Mock(side_effect=get_token)) clients.append(client) return client @@ -99,13 +99,16 @@ def get_token(*args, **options): for invalid in (None, SimpleNamespace(token=""), SimpleNamespace(token=" "), SimpleNamespace(token="private-token"), SimpleNamespace(token="private-token", expires_on=time.time() + 5)): - engine.close() if stage == "access": access = invalid else: access = SimpleNamespace(token="private-token", expires_on=time.time() + 3600) assertion = invalid - status, body = engine.dispatch(_request(), "request") + candidate = DispatchEngine(engine.registry, engine.secrets, engine.env) + try: + status, body = candidate.dispatch(_request(), "request") + finally: + candidate.close() assert status == 502 and body["reason"] == "provider credential unavailable" assert "private" not in json.dumps(body) engine.secrets.resolve.assert_not_called() @@ -129,7 +132,8 @@ def fail(*args, **kwargs): raise RuntimeError("PRIVATE-TOKEN-EXCEPTION") monkeypatch.setattr(credentials_module, "ManagedIdentityCredential", Mock()) - monkeypatch.setattr(credentials_module, "ClientAssertionCredential", Mock(return_value=Mock(get_token=fail))) + monkeypatch.setattr(credentials_module, "ClientAssertionCredential", + Mock(return_value=Mock(spec=["get_token"], get_token=fail))) try: with ThreadPoolExecutor(max_workers=1) as pool: pending = pool.submit(engine.dispatch, _request(), "request") From 4cac5ae46b37883d5b71e0f7583b6c18be290dfa Mon Sep 17 00:00:00 2001 From: James Xian Date: Thu, 24 Sep 2026 12:31:47 -0700 Subject: [PATCH 3/3] Split API key and access token caches Use provider-selected ApiKeyCache and AccessTokenCache implementations across C#, JavaScript, and Python. Replace generic cache machinery with standard secret caches, SDK token caching, and one fixed refresh loop. Preserve expiry, shared acquisition, cancellation, and privacy protections; update regression coverage and documentation. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> --- docs/CONTRACT.md | 116 ++--- dotnet/README.md | 19 +- dotnet/Src/DispatchEngine.cs | 4 +- dotnet/Src/ProviderCredentials.cs | 304 +++++++------ dotnet/Src/RefreshingCache.cs | 162 ------- dotnet/Src/SecretResolver.cs | 16 +- dotnet/tests/CredentialCacheTests.cs | 142 +++--- dotnet/tests/EngineTests.cs | 16 +- javascript/README.md | 18 +- javascript/package-lock.json | 11 +- javascript/package.json | 3 +- javascript/src/functions/credentials.js | 384 ++++++++--------- javascript/src/functions/refreshingCache.js | 176 -------- javascript/test/credential-cache.test.js | 320 ++++++-------- javascript/test/credential-sdk.test.js | 24 +- javascript/test/dispatch.test.js | 21 +- python/README.md | 22 +- python/requirements.txt | 1 + python/src/credentials.py | 358 +++++++++------- python/src/refreshing_cache.py | 192 --------- python/src/secrets.py | 16 +- python/tests/test_credential_cache.py | 451 +++++++------------- python/tests/test_credential_sdk.py | 25 +- python/tests/test_engine.py | 6 +- 24 files changed, 1092 insertions(+), 1715 deletions(-) delete mode 100644 dotnet/Src/RefreshingCache.cs delete mode 100644 javascript/src/functions/refreshingCache.js delete mode 100644 python/src/refreshing_cache.py diff --git a/docs/CONTRACT.md b/docs/CONTRACT.md index 7795f29..bada50b 100644 --- a/docs/CONTRACT.md +++ b/docs/CONTRACT.md @@ -113,21 +113,16 @@ there is no API-key fallback. Evaluation skips acquisition. A provider rejection Tokens are treated as opaque: the Function checks SDK expiry metadata, not custom JWT claims. Soprano remains responsible for signature, issuer, audience, expiry, permissions, and account validation. -Credential instances are reused for the configured tenant/application/identity. Separate -[worker-local caches](#credential-caching-and-refresh) hold the managed-identity assertion and the -final provider token for each selected scope. A usable final token avoids both SDK acquisition calls -on the delivery path. Each cache owns its refresh independently, so one waiting caller cannot cancel -an assertion refresh another caller needs. JavaScript and .NET bound each acquisition to 2.5 seconds; -JavaScript links that cancellation through an SDK HTTP-client wrapper, including the managed-identity -transport, rather than relying on `getToken` options or policies the SDK can replace. -Python bounds each wait to 2.5 seconds and uses 2.5-second -connect/read inactivity timeouts; its shared refresh may finish after a waiter leaves. - -Python uses `get_token_info` when the installed SDK supports it, preserving the `refresh_on` hint. -Older supported SDKs expose only expiry metadata through `get_token`; those versions retain the -five-minute pre-expiry refresh target. An acquisition failure never falls back to another token API. - -Credential SDK transport retries are disabled; failed refreshes use the bounded backoff described +Credential instances and their SDK caches are reused for the configured tenant/application/identity. +One [worker-local refresh loop](#credential-caching-and-refresh) warms both exchange stages; a +snapshot of the latest provider token keeps refresh off the delivery path. JavaScript and .NET +bound the shared acquisition to 2.5 seconds, independent of individual waiters. JavaScript links +cancellation through an SDK HTTP-client wrapper because `getToken` options alone are insufficient +in the installed SDK. Python bounds caller waits and SDK connect/read inactivity to 2.5 seconds; +shared synchronous retrieval may finish after a waiter leaves. It uses `get_token_info` for refresh +hints when supported, otherwise `get_token`; a failed acquisition never falls back to another API. + +Credential SDK transport retries are disabled; failed refreshes use the fixed polling cadence described below. These are not end-to-end delivery deadlines. JavaScript suppresses SDK logs in the acquisition's asynchronous context. Python filters Azure Identity/Core/MSAL records on configured handlers in that context; configure logging sinks before handling requests. .NET disables credential @@ -344,60 +339,41 @@ when credentials are fetched, not the HTTP/nonce contract, caller authentication selection or provider request format. The decryption-key Key Vault reference remains separate and is still resolved by the platform; this cache does not rotate or replace JWE keys. -| Cache | Refresh target | Hard usability boundary | +The selected provider's manifest determines which of two concrete cache classes is created: + +| Authentication mode | Cache | Acquisition | |---|---|---| -| Key Vault API-key bundle | Four minutes after successful retrieval. Required key and customer-ID secrets are fetched concurrently and published together only when both reads succeed and values are nonblank. | Five minutes after retrieval; failed refreshes never extend the old bundle's lifetime. | -| Managed-identity assertion | Five minutes before SDK expiry, or an earlier SDK refresh hint when available. | SDK expiry minus 30 seconds. | -| Final Entra provider token | Five minutes before SDK expiry, or an earlier SDK refresh hint when available, separately for each scope. | SDK expiry minus 30 seconds. | - -These are in-memory caches, one per worker process, not distributed caches or persisted token -stores. API-key entries are isolated by vault, managed identity and manifest secret names. OAuth -state is isolated by provider tenant, application and managed identity, with separate final-token -entries for different scopes. A change of credential configuration stops the old entries; deploy -app-setting changes normally with a worker restart rather than mutating process environment in place. -Configuration replacement clears old entries without shutting down the manager. Explicit -`close()`/`Dispose()` is terminal: a stopped manager cannot acquire credentials again. Create a -new manager or restart the worker instead of reusing a stopped instance. Timing values are named -policy constants in each runtime's refreshing-cache implementation, not additional app settings. - -Concurrent cache misses share **one in-progress acquisition per entry**. A request with a still-usable -cached credential returns it immediately while a due refresh proceeds separately. Refresh failures -retain that credential only until its original hard expiry. Once expired, callers join the shared -refresh or fail closed with the existing sanitized credential error; no stale-success fallback is -introduced. This does not deduplicate provider deliveries or change caller/provider retry behavior. - -An SDK refresh can return the same token from its own cache. The manager preserves the token's -original expiry instead of treating it as a new token. If the suggested refresh time is already past, -the next check is delayed by half the remaining usable lifetime, bounded to 1-60 seconds, preventing -an immediate refresh loop. Failed refreshes back off exponentially from 5 seconds to a 60-second -base, plus up to 20% jitter. Requests do not bypass that backoff and repeatedly hit a failing -dependency. Successful refresh resets it. - -JavaScript's app-start hook, Python's worker-module initialization thread, and .NET's hosted service -start credential preparation. Each entry then owns its refresh timer; a single distributed timer -trigger would not populate every worker's memory. JavaScript timers are unreferenced and Python -threads are daemon threads; termination hooks/`atexit`/hosted-service shutdown cancel scheduled -work and discard entries. Late completion cannot repopulate a stopped cache. A cold Python caller -can stop waiting without cancelling the shared retrieval. Failed or absent provider configuration -does not prevent evaluation from working. - -Python uses explicitly owned daemon threads for parallel secret reads, not executor workers that -are joined before application `atexit` handlers. A bundle keeps ownership of both reads until they -finish, so a waiting caller's timeout or one failed read cannot start overlapping retries. Shutdown -releases pending cache waiters immediately. Synchronous Python SDK I/O is not forcibly cancelled; -unfinished reads cannot publish late values or block normal process exit. - -**This is not a guarantee that the first request after a cold start meets the caller's budget.** -Initialization can itself be on that first request's critical path, and a worker may receive traffic -before preparation finishes. Existing Always On/minimum-instance settings can help, but readiness, -scale-out and caller-observed latency must be measured in the deployment. A warm final token avoids -the Entra exchange; warming only the managed-identity assertion would not achieve that. - -Background failures emit a compact `credential_refresh_failed` service event with `cacheKind` -(`key_vault`, `managed_identity`, `provider_token`, `configuration`, or `initialization`) and the -fixed reason `credential_unavailable`. They do not carry a request's tracing IDs or secret values. -The existing per-request `providerCredentialElapsedMs` still measures the resolution observed by -that caller. No per-operation MI/Entra timing diagnostics are added. +| `apiKey` | `ApiKeyCache` | Fetch the manifest's Key Vault secrets using managed identity. Cache the complete key/customer-ID bundle in .NET `MemoryCache`, JavaScript `lru-cache`, or Python `cachetools.TTLCache`. | +| `oauth` | `AccessTokenCache` | Reuse Azure Identity's managed-identity and client-assertion credentials and their SDK caches. Retain only the latest usable provider token. No Key Vault access. | + +Only the selected cache starts. Its credential configuration is bound on first use; app-setting +changes require a worker restart, not live cache switching. API-key bundles are published only +after all required reads succeed. Refresh targets four minutes after retrieval and hard expiry +is five minutes; reads or failed refreshes never extend the lifetime. Provider tokens retain their +original SDK expiry and are unusable with 30 seconds or less remaining. Tokens are never persisted. + +The shared coordinator prepares credentials at startup and polls the selected cache every 30 seconds. +`ApiKeyCache` skips retrieval until its refresh target is due; `AccessTokenCache` consults both SDK +credentials and lets the SDK decide whether network acquisition is needed. There is no separate MI +cache, adaptive expiry timer, or exponential retry policy. Failures retry on a later poll; requests +cannot start another acquisition within the same 30-second window. This is best-effort scheduling, +not an exact refresh deadline. Timing policy uses named constants, not extra app settings. + +Concurrent cold requests share one acquisition. Requests with usable cached credentials do not +wait for background refresh. After hard expiry, they join the shared acquisition or fail closed. +JavaScript's app-start hook, Python's initialization thread and .NET's hosted service start only +credential preparation, never provider delivery. Shutdown stops polling and prevents late +publication; `close()`/`Dispose()` is terminal. JavaScript and .NET propagate the 2.5-second +acquisition deadline to SDK HTTP. Python bounds caller waits and SDK connect/read inactivity to +2.5 seconds but cannot forcibly cancel synchronous I/O; daemon secret reads stay shared until +both finish and cannot block process exit. Missing/broken provider configuration does not prevent +evaluation. + +**Prewarming does not guarantee the first request meets the caller's timeout.** Worker readiness, +scale-out and ingress overhead still matter. Refresh does not retry or deduplicate provider sends. +Background failures log only `credential_refresh_failed`, `cacheKind` and the fixed +`credential_unavailable` reason, without request IDs, credential values or SDK exception details. +Per-request `providerCredentialElapsedMs` continues to measure the caller's resolution time. --- @@ -559,8 +535,8 @@ Each language keeps lightweight offline tests covering representative applicatio - Awaited delivery, nonce acknowledgement and privacy-safe logging, including the shared service-event order and summary field set in [contract.json](../tests/fixtures/contract.json), identifier provenance, error paths, provider-body timeouts and concurrent request isolation. -- Single-flight credential retrieval, automatic refresh, stale-value expiry, token lifetime - preservation, failure backoff, configuration isolation and cleanup using controlled clocks and +- Selected-cache-only startup, shared credential retrieval, fixed refresh/retry cadence, hard expiry, + token lifetime preservation and shutdown using controlled clocks and fake dependencies. JavaScript tests also exercise cancellation through the actual SDK pipeline. The sample deliberately omits exhaustive input permutations and SDK internals. These tests use diff --git a/dotnet/README.md b/dotnet/README.md index 9ae5b5c..0bbf2ac 100644 --- a/dotnet/README.md +++ b/dotnet/README.md @@ -90,14 +90,13 @@ six-digit numeric run that is not part of a longer number and repeats the comple ## Source -The hosted credential-refresh service automatically prewarms a configured provider. Separate -process-local caches refresh Key Vault bundles, managed-identity assertions and final Entra tokens. -Each acquisition owns its cancellation budget; cancelling a waiter does not cancel another -request's shared retrieval. Shutdown stops timers, cancels acquisitions and drops values. -Disposal is terminal; configuration replacement clears caches without disposing the manager. -Refresh failures use the same structured, sanitized JSON logging pattern as request events. See the -[refresh contract](../docs/CONTRACT.md#credential-caching-and-refresh) for expiry/backoff semantics -and cold-start limitations. Evaluation remains independent from successful credential preparation. +The hosted service selects `ApiKeyCache` or `AccessTokenCache` from the provider manifest's auth mode. +Only the selected cache starts: API keys use Key Vault and framework `MemoryCache`; access tokens +use the MI/Entra SDKs without Key Vault. One periodic timer polls every 30 seconds. Configuration +changes require restart. Each shared acquisition owns +its cancellation budget; a waiter cannot cancel another request's retrieval. Disposal stops refresh +and prevents late publication. See the [refresh contract](../docs/CONTRACT.md#credential-caching-and-refresh) +for expiry, sanitized failure logs and cold-start limits. Evaluation remains independent. | Source | Purpose | |---|---| @@ -105,12 +104,12 @@ and cold-start limitations. Evaluation remains independent from successful crede | [Functions/SendOtp.cs](Functions/SendOtp.cs) | HTTP handler | | [Src/AppConfig.cs](Src/AppConfig.cs) | Shared deployment settings | | [Src/DispatchEngine.cs](Src/DispatchEngine.cs) | Envelope/JWE handling and dispatch | -| [Src/ProviderCredentials.cs](Src/ProviderCredentials.cs), [RefreshingCache.cs](Src/RefreshingCache.cs) | Single-flight credential bundles and independent token refresh | +| [Src/ProviderCredentials.cs](Src/ProviderCredentials.cs) | `ApiKeyCache`, `AccessTokenCache` and their shared refresh coordinator | | [Src/CredentialRefreshService.cs](Src/CredentialRefreshService.cs) | Per-worker startup and shutdown integration | | [Src/RequestLog.cs](Src/RequestLog.cs) | Request-scoped [service events and summaries](../docs/CONTRACT.md#application-logs) with explicit ID sources | | [Src/ProviderRegistry.cs](Src/ProviderRegistry.cs), [Src/IProviderAdapter.cs](Src/IProviderAdapter.cs) | Adapter lookup and contract | | [Src/Providers/](Src/Providers/) | Adapter manifests and API-specific implementations | -| [Src/SecretResolver.cs](Src/SecretResolver.cs) | Key Vault transport; `ISecretResolver.ResolveAsync` accepts an optional cancellation token and bundle caching belongs to the credential manager | +| [Src/SecretResolver.cs](Src/SecretResolver.cs) | Key Vault transport; `ISecretResolver.ResolveAsync` accepts cancellation and `ApiKeyCache` owns the bundle | | [Src/OutcomeMapper.cs](Src/OutcomeMapper.cs), [Src/Models.cs](Src/Models.cs) | Outcomes and shared records | Implement `IProviderAdapter` and register it in [Program.cs](Program.cs) without adding provider-specific diff --git a/dotnet/Src/DispatchEngine.cs b/dotnet/Src/DispatchEngine.cs index 5d416c9..9eaf53e 100644 --- a/dotnet/Src/DispatchEngine.cs +++ b/dotnet/Src/DispatchEngine.cs @@ -234,14 +234,14 @@ internal DispatchEngine(ProviderRegistry registry, ISecretResolver secrets, IHtt _registry = registry; _httpFactory = httpFactory; _env = env ?? new ProcessEnv(); - _credentials = new ProviderCredentials(secrets, _env, createManagedIdentity, createOAuthCredential, log, clock); + _credentials = new ProviderCredentials(secrets, createManagedIdentity, createOAuthCredential, log, clock); } private static ClientAssertionCredentialOptions OAuthOptions() { var options = new ClientAssertionCredentialOptions { AuthorityHost = AzureAuthorityHosts.AzurePublicCloud }; options.Retry.MaxRetries = 0; - options.Retry.NetworkTimeout = CredentialCachePolicy.AcquisitionTimeout; + options.Retry.NetworkTimeout = ProviderCredentials.AcquisitionTimeout; options.Diagnostics.IsLoggingEnabled = false; options.Diagnostics.IsLoggingContentEnabled = false; return options; diff --git a/dotnet/Src/ProviderCredentials.cs b/dotnet/Src/ProviderCredentials.cs index df2ade9..374cfa0 100644 --- a/dotnet/Src/ProviderCredentials.cs +++ b/dotnet/Src/ProviderCredentials.cs @@ -1,179 +1,241 @@ using System.Text.Json; using Azure.Core; +using Microsoft.Extensions.Caching.Memory; +using Microsoft.Extensions.Internal; using Microsoft.Extensions.Logging; using Microsoft.Extensions.Logging.Abstractions; +using static Epp.Otp.ProviderCredentials; namespace Epp.Otp; +internal interface ICredentialCache : IDisposable +{ + string Stage { get; } + ProviderCredential? Get(); + Task RefreshAsync(CancellationToken cancellation); +} + +internal sealed class ApiKeyCache(ISecretResolver secrets, AuthConfig auth, TimeProvider clock) : ICredentialCache +{ + private const string BundleKey = "bundle"; + private static readonly TimeSpan Ttl = TimeSpan.FromMinutes(5); + private static readonly TimeSpan RefreshInterval = TimeSpan.FromMinutes(4); + private readonly object _gate = new(); + // The fixed key bounds the cache; a size limit can reject replacement before removing the old entry. + private readonly MemoryCache _values = new(new MemoryCacheOptions { Clock = new CacheClock(clock) }); + private DateTimeOffset _refreshAt; + private bool _closed; + public string Stage => "key_vault"; + + private sealed class CacheClock(TimeProvider time) : ISystemClock + { + public DateTimeOffset UtcNow => time.GetUtcNow(); + } + + public ProviderCredential? Get() + { + lock (_gate) return _closed ? null : _values.Get(BundleKey); + } + public async Task RefreshAsync(CancellationToken cancellation) + { + lock (_gate) + { + if (_closed) throw Unavailable(); + if (Get() is not null && _refreshAt > clock.GetUtcNow()) return; + } + if (string.IsNullOrWhiteSpace(auth.KeyVaultSecretName)) throw Unavailable(); + var key = secrets.ResolveAsync(auth.KeyVaultSecretName, cancellation); + var identity = string.IsNullOrWhiteSpace(auth.IdentityKeyVaultSecretName) + ? Task.FromResult("") : secrets.ResolveAsync(auth.IdentityKeyVaultSecretName, cancellation); + await Task.WhenAll(key, identity).ConfigureAwait(false); + if (string.IsNullOrWhiteSpace(key.Result) || + (!string.IsNullOrWhiteSpace(auth.IdentityKeyVaultSecretName) && string.IsNullOrWhiteSpace(identity.Result))) throw Unavailable(); + lock (_gate) + { + cancellation.ThrowIfCancellationRequested(); + if (_closed) throw Unavailable(); + _values.Set(BundleKey, new ProviderCredential(ApiKeyMode, Secret: key.Result, Identity: identity.Result), + new MemoryCacheEntryOptions { AbsoluteExpiration = clock.GetUtcNow() + Ttl }); + _refreshAt = clock.GetUtcNow() + RefreshInterval; + } + } + public void Dispose() + { + lock (_gate) { _closed = true; _values.Dispose(); } + } +} + +internal sealed class AccessTokenCache : ICredentialCache +{ + private const string ExchangeScope = "api://AzureADTokenExchange/.default"; + private static readonly TimeSpan ExpirySkew = TimeSpan.FromSeconds(30); + private readonly object _gate = new(); + private readonly TokenCredential _identity, _credential; + private readonly string _scope; + private readonly TimeProvider _clock; + private AccessToken? _token; + private bool _closed; + public string Stage { get; private set; } = "provider_token"; + + internal AccessTokenCache(AppConfig config, Func createIdentity, + Func>, TokenCredential> createCredential, TimeProvider clock) + { + if (string.IsNullOrWhiteSpace(config.ProviderTenantId) || string.IsNullOrWhiteSpace(config.ProviderScope) + || string.IsNullOrWhiteSpace(config.OutboundClientId) || string.IsNullOrWhiteSpace(config.OutboundManagedIdentityClientId)) throw Unavailable(); + _clock = clock; + _scope = config.ProviderScope; + _identity = createIdentity(config.OutboundManagedIdentityClientId); + _credential = createCredential(config.ProviderTenantId, config.OutboundClientId, + async cancellation => (await Assertion(cancellation).ConfigureAwait(false)).Token); + } + public ProviderCredential? Get() + { + lock (_gate) return !_closed && _token is { } token && token.ExpiresOn > _clock.GetUtcNow() + ExpirySkew + ? new(OAuthMode, AccessToken: token.Token) : null; + } + private AccessToken Check(AccessToken token) + { + if (string.IsNullOrWhiteSpace(token.Token) || token.ExpiresOn <= _clock.GetUtcNow() + ExpirySkew) throw Unavailable(); + return token; + } + private async Task Assertion(CancellationToken cancellation) => + Check(await _identity.GetTokenAsync(new TokenRequestContext(new[] { ExchangeScope }), cancellation).ConfigureAwait(false)); + + public async Task RefreshAsync(CancellationToken cancellation) + { + lock (_gate) { if (_closed) throw Unavailable(); } + Stage = "managed_identity"; + await Assertion(cancellation).ConfigureAwait(false); + Stage = "provider_token"; + var token = Check(await _credential.GetTokenAsync(new TokenRequestContext(new[] { _scope }), cancellation).ConfigureAwait(false)); + lock (_gate) + { + cancellation.ThrowIfCancellationRequested(); + if (_closed) throw Unavailable(); + _token = token; + } + } + public void Dispose() + { + lock (_gate) { _closed = true; _token = null; } + } +} + +// One selected cache and one periodic refresh; configuration changes require a worker restart. internal sealed class ProviderCredentials : IDisposable { + internal const string ApiKeyMode = "apiKey"; + internal const string OAuthMode = "oauth"; + internal static readonly TimeSpan AcquisitionTimeout = TimeSpan.FromSeconds(2.5); + private static readonly TimeSpan PollInterval = TimeSpan.FromSeconds(30); private readonly object _gate = new(); private readonly ISecretResolver _secrets; - private readonly IEnv _env; private readonly Func _createIdentity; private readonly Func>, TokenCredential> _createCredential; private readonly ILogger _log; private readonly TimeProvider _clock; - private string? _key; - private RefreshingCache? _bundle; - private RefreshingCache? _assertion; - private TokenCredential? _credential; - private readonly Dictionary> _tokens = new(); + private ICredentialCache? _cache; + private Task? _pending; + private CancellationTokenSource? _acquisition; + private ITimer? _timer; + private DateTimeOffset _nextAttempt; private bool _disposed; - internal ProviderCredentials(ISecretResolver secrets, IEnv env, - Func createIdentity, + internal ProviderCredentials(ISecretResolver secrets, Func createIdentity, Func>, TokenCredential> createCredential, ILogger? log = null, TimeProvider? clock = null) { _secrets = secrets; - _env = env; _createIdentity = createIdentity; _createCredential = createCredential; _log = log ?? NullLogger.Instance; _clock = clock ?? TimeProvider.System; } - internal void ReportFailure(string kind) { const string eventName = "credential_refresh_failed"; var record = new Dictionary { - ["logType"] = "service", - ["eventName"] = eventName, - ["cacheKind"] = kind, - ["failureReason"] = "credential_unavailable", + ["logType"] = "service", ["eventName"] = eventName, + ["cacheKind"] = kind, ["failureReason"] = "credential_unavailable", }; _log.Log(LogLevel.Warning, new EventId(0, eventName), record, null, static (state, _) => JsonSerializer.Serialize(state)); } - - private RefreshingCache Cache(string kind, Func>> load) => - new(load, () => ReportFailure(kind), _clock); - - internal async Task ResolveAsync(AuthConfig auth, AppConfig config, CancellationToken cancellation = default) + internal Task ResolveAsync(AuthConfig auth, AppConfig config, CancellationToken cancellation = default) { - RefreshingCache? bundle; - RefreshingCache? tokenCache = null; lock (_gate) { ObjectDisposedException.ThrowIf(_disposed, this); - if (auth.Mode == "oauth" && (string.IsNullOrWhiteSpace(config.ProviderTenantId) - || string.IsNullOrWhiteSpace(config.ProviderScope) || string.IsNullOrWhiteSpace(config.OutboundClientId) - || string.IsNullOrWhiteSpace(config.OutboundManagedIdentityClientId))) + if (_cache is null) { - Clear(); - throw new InvalidOperationException("provider OAuth token unavailable"); - } - if (auth.Mode is not ("apiKey" or "oauth")) - { - Clear(); - throw new InvalidOperationException("provider credential unavailable"); - } - string key; - if (auth.Mode == "apiKey") - { - key = JsonSerializer.Serialize(new[] + try { - auth.Mode, _env.Get("KEY_VAULT_URL"), _env.Get("AZURE_CLIENT_ID"), - auth.KeyVaultSecretName, auth.IdentityKeyVaultSecretName, - }); - } - else - { - key = JsonSerializer.Serialize(new[] - { - auth.Mode, config.ProviderTenantId, config.OutboundClientId, config.OutboundManagedIdentityClientId, - }); - } - if (_key != key) - { - Clear(); - if (auth.Mode == "apiKey") _bundle = CreateBundle(auth); - else CreateOAuth(config); - _key = key; - } - bundle = _bundle; - if (auth.Mode == "oauth") - { - var scope = config.ProviderScope ?? throw new InvalidOperationException("provider OAuth token unavailable"); - if (!_tokens.TryGetValue(scope, out tokenCache)) - { - var credential = _credential ?? throw new InvalidOperationException("provider OAuth token unavailable"); - tokenCache = Cache("provider_token", async ct => - TokenEntry(await credential.GetTokenAsync(new TokenRequestContext(new[] { scope }), ct).ConfigureAwait(false))); - _tokens.Add(scope, tokenCache); + _cache = auth.Mode switch + { + ApiKeyMode => new ApiKeyCache(_secrets, auth, _clock), + OAuthMode => new AccessTokenCache(config, _createIdentity, _createCredential, _clock), + _ => throw Unavailable(), + }; + _timer = _clock.CreateTimer(_ => Tick(), null, PollInterval, PollInterval); } + catch (Exception) { ReportFailure("configuration"); return Task.FromException(Unavailable()); } } + var cached = _cache.Get(); + if (cached is not null) return Task.FromResult(cached); + var pending = Refresh(); + return cancellation.CanBeCanceled ? pending.WaitAsync(cancellation) : pending; } - if (bundle is not null) return await bundle.GetAsync(cancellation).ConfigureAwait(false); - if (tokenCache is null) throw new InvalidOperationException("provider credential unavailable"); - var token = await tokenCache.GetAsync(cancellation).ConfigureAwait(false); - return new ProviderCredential("oauth", AccessToken: token.Token); } - - private RefreshingCache CreateBundle(AuthConfig auth) => Cache("key_vault", async cancellation => - { - if (string.IsNullOrWhiteSpace(auth.KeyVaultSecretName)) throw new InvalidOperationException("provider credential unavailable"); - var secretTask = _secrets.ResolveAsync(auth.KeyVaultSecretName, cancellation); - var identityTask = string.IsNullOrWhiteSpace(auth.IdentityKeyVaultSecretName) - ? Task.FromResult(string.Empty) : _secrets.ResolveAsync(auth.IdentityKeyVaultSecretName, cancellation); - await Task.WhenAll(secretTask, identityTask).ConfigureAwait(false); - var secret = await secretTask.ConfigureAwait(false); - var identity = await identityTask.ConfigureAwait(false); - if (string.IsNullOrWhiteSpace(secret) || (!string.IsNullOrEmpty(auth.IdentityKeyVaultSecretName) && string.IsNullOrWhiteSpace(identity))) - throw new InvalidOperationException("provider credential unavailable"); - var now = _clock.GetUtcNow(); - return new CredentialCacheEntry(new("apiKey", Secret: secret, Identity: identity), - now + CredentialCachePolicy.SecretTtl, now + CredentialCachePolicy.SecretRefreshInterval); - }); - - private void CreateOAuth(AppConfig config) + private void Tick() { - var identity = _createIdentity(config.OutboundManagedIdentityClientId!); - var assertion = Cache("managed_identity", async cancellation => - TokenEntry(await identity.GetTokenAsync( - new TokenRequestContext(new[] { "api://AzureADTokenExchange/.default" }), cancellation).ConfigureAwait(false))); - _assertion = assertion; - _credential = _createCredential(config.ProviderTenantId!, config.OutboundClientId!, - async cancellation => (await assertion.GetAsync(cancellation).ConfigureAwait(false)).Token); + lock (_gate) { if (!_disposed) _ = ObserveAsync(Refresh()); } } - - private CredentialCacheEntry TokenEntry(AccessToken token) - { - var now = _clock.GetUtcNow(); - if (string.IsNullOrWhiteSpace(token.Token) || token.ExpiresOn <= now + CredentialCachePolicy.TokenExpirySkew) - throw new InvalidOperationException("provider credential unavailable"); - var expires = token.ExpiresOn - CredentialCachePolicy.TokenExpirySkew; - var refresh = token.ExpiresOn - CredentialCachePolicy.TokenRefreshLead; - if (token.RefreshOn is { } hint && hint < refresh) refresh = hint; - if (refresh <= now) - { - var delay = Math.Clamp((expires - now).TotalSeconds / 2, - CredentialCachePolicy.MinRefreshDelay.TotalSeconds, CredentialCachePolicy.MaxRefreshDelay.TotalSeconds); - refresh = now.AddSeconds(delay); - } - return new(token, expires, refresh); + private static async Task ObserveAsync(Task task) => await task.ConfigureAwait(ConfigureAwaitOptions.SuppressThrowing); + private Task Refresh() + { + if (_pending is not null) return _pending; + if (_disposed || _cache is null || _nextAttempt > _clock.GetUtcNow()) + return Task.FromException(Unavailable()); + _nextAttempt = _clock.GetUtcNow() + PollInterval; + var completion = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + _pending = completion.Task; + var cancellation = _acquisition = new CancellationTokenSource(AcquisitionTimeout, _clock); + _ = RunAsync(_cache, completion, cancellation); + return completion.Task; } - - private void Clear() + private async Task RunAsync(ICredentialCache cache, TaskCompletionSource completion, CancellationTokenSource cancellation) { - _bundle?.Dispose(); - _assertion?.Dispose(); - foreach (var cache in _tokens.Values) cache.Dispose(); - _tokens.Clear(); - _bundle = null; - _assertion = null; - _credential = null; - _key = null; + ProviderCredential? value = null; + try + { + await cache.RefreshAsync(cancellation.Token).WaitAsync(cancellation.Token).ConfigureAwait(false); + value = cache.Get(); + } + catch (Exception) { /* Report only the sanitized failure below. */ } + lock (_gate) + { + if (_disposed || cancellation.IsCancellationRequested) value = null; + _pending = null; + _acquisition = null; + cancellation.Dispose(); + if (value is null) + { + if (!_disposed) ReportFailure(cache.Stage); + completion.TrySetException(Unavailable()); + } + else completion.TrySetResult(value); + } } - + internal static InvalidOperationException Unavailable() => new("provider credential unavailable"); public void Dispose() { lock (_gate) { _disposed = true; - Clear(); + _timer?.Dispose(); + _acquisition?.Cancel(); + _cache?.Dispose(); } } } diff --git a/dotnet/Src/RefreshingCache.cs b/dotnet/Src/RefreshingCache.cs deleted file mode 100644 index 782dfa1..0000000 --- a/dotnet/Src/RefreshingCache.cs +++ /dev/null @@ -1,162 +0,0 @@ -namespace Epp.Otp; - -internal static class CredentialCachePolicy -{ - internal static readonly TimeSpan AcquisitionTimeout = TimeSpan.FromSeconds(2.5); - internal static readonly TimeSpan SecretTtl = TimeSpan.FromMinutes(5); - internal static readonly TimeSpan SecretRefreshInterval = TimeSpan.FromMinutes(4); - internal static readonly TimeSpan TokenExpirySkew = TimeSpan.FromSeconds(30); - internal static readonly TimeSpan TokenRefreshLead = TimeSpan.FromMinutes(5); - internal static readonly TimeSpan MinRefreshDelay = TimeSpan.FromSeconds(1); - internal static readonly TimeSpan MaxRefreshDelay = TimeSpan.FromSeconds(60); - internal static readonly TimeSpan InitialRetryDelay = TimeSpan.FromSeconds(5); - internal static readonly TimeSpan MaxRetryDelay = TimeSpan.FromSeconds(60); - internal static readonly TimeSpan MaxTimerDelay = TimeSpan.FromMilliseconds(int.MaxValue); - internal const int MaxRetryExponent = 4; - internal const double RetryJitterRatio = 0.2; -} - -internal sealed record CredentialCacheEntry(T Value, DateTimeOffset ExpiresAt, DateTimeOffset RefreshAt) -{ - public override string ToString() => nameof(CredentialCacheEntry); -} - -internal sealed class RefreshingCache : IDisposable -{ - private readonly object _gate = new(); - private readonly Func>> _load; - private readonly TimeProvider _clock; - private readonly Func _random; - private readonly Action _onFailure; - private CredentialCacheEntry? _entry; - private Task? _inFlight; - private CancellationTokenSource? _acquisition; - private ITimer? _timer; - private DateTimeOffset _retryAt; - private int _failures; - private bool _closed; - - internal RefreshingCache(Func>> load, - Action onFailure, TimeProvider? clock = null, Func? random = null) - { - _load = load; - _onFailure = onFailure; - _clock = clock ?? TimeProvider.System; - _random = random ?? Random.Shared.NextDouble; - } - - private static Exception Unavailable() => new InvalidOperationException("provider credential unavailable"); - - internal Task GetAsync(CancellationToken cancellationToken = default) - { - Task pending; - lock (_gate) - { - if (_closed) return Task.FromException(Unavailable()); - var now = _clock.GetUtcNow(); - if (_entry is not null && _entry.ExpiresAt > now) - { - if (_entry.RefreshAt <= now && _retryAt <= now) _ = ObserveAsync(StartRefresh()); - return Task.FromResult(_entry.Value); - } - if (_inFlight is null && _retryAt > now) return Task.FromException(Unavailable()); - pending = StartRefresh(); - } - return cancellationToken.CanBeCanceled ? pending.WaitAsync(cancellationToken) : pending; - } - - internal Task RefreshAsync() - { - lock (_gate) - { - if (_closed || _retryAt > _clock.GetUtcNow()) return Task.FromException(Unavailable()); - return StartRefresh(); - } - } - - private Task StartRefresh() - { - if (_inFlight is not null) return _inFlight; - _timer?.Dispose(); - _timer = null; - var completion = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); - _inFlight = completion.Task; - _acquisition = new CancellationTokenSource(CredentialCachePolicy.AcquisitionTimeout, _clock); - _ = RunRefreshAsync(completion, _acquisition); - return completion.Task; - } - - private async Task RunRefreshAsync(TaskCompletionSource completion, CancellationTokenSource cancellation) - { - CredentialCacheEntry? entry = null; - try - { - entry = await _load(cancellation.Token).WaitAsync(cancellation.Token).ConfigureAwait(false); - if (entry.ExpiresAt <= _clock.GetUtcNow()) entry = null; - } - catch (Exception) - { - // Emit only the fixed failure below; SDK exceptions may contain credentials. - } - lock (_gate) - { - _inFlight = null; - _acquisition = null; - if (_closed || cancellation.IsCancellationRequested) entry = null; - cancellation.Dispose(); - if (!_closed) - { - if (entry is null) - { - _failures++; - var exponent = Math.Min(_failures - 1, CredentialCachePolicy.MaxRetryExponent); - var backoff = Math.Min(CredentialCachePolicy.MaxRetryDelay.TotalSeconds, - CredentialCachePolicy.InitialRetryDelay.TotalSeconds * Math.Pow(2, exponent)); - _retryAt = _clock.GetUtcNow().AddSeconds(backoff * (1 + _random() * CredentialCachePolicy.RetryJitterRatio)); - _onFailure(); - } - else - { - _entry = entry; - _failures = 0; - _retryAt = default; - } - var next = entry is null ? _retryAt : entry.RefreshAt; - var wait = Math.Clamp((next - _clock.GetUtcNow()).TotalMilliseconds, - CredentialCachePolicy.MinRefreshDelay.TotalMilliseconds, CredentialCachePolicy.MaxTimerDelay.TotalMilliseconds); - _timer = _clock.CreateTimer(_ => ScheduledRefresh(), null, TimeSpan.FromMilliseconds(wait), Timeout.InfiniteTimeSpan); - } - if (entry is null) completion.TrySetException(Unavailable()); - else completion.TrySetResult(entry.Value); - } - } - - private void ScheduledRefresh() - { - lock (_gate) - { - _timer?.Dispose(); - _timer = null; - if (!_closed) _ = ObserveAsync(StartRefresh()); - } - } - - private static async Task ObserveAsync(Task task) - { - // Refresh failures are already reported; a timer has no request awaiting the result. - try { await task.ConfigureAwait(false); } - catch (Exception) { } - } - - public void Dispose() - { - lock (_gate) - { - _closed = true; - _timer?.Dispose(); - _timer = null; - _acquisition?.Cancel(); - _entry = null; - } - } -} diff --git a/dotnet/Src/SecretResolver.cs b/dotnet/Src/SecretResolver.cs index 6795144..1b51652 100644 --- a/dotnet/Src/SecretResolver.cs +++ b/dotnet/Src/SecretResolver.cs @@ -4,13 +4,12 @@ namespace Epp.Otp; // Resolves Key Vault secret names to values via the Function's managed identity (user-assigned when -// AZURE_CLIENT_ID is set, else system-assigned). ProviderCredentials caches the complete bundle. +// AZURE_CLIENT_ID is set, else system-assigned). ApiKeyCache publishes the complete bundle. public sealed class SecretResolver : ISecretResolver { private readonly object _gate = new(); private readonly IEnv _env; private SecretClient? _client; - private (string Url, string? Identity)? _clientKey; public SecretResolver(IEnv? env = null) { @@ -19,25 +18,24 @@ public SecretResolver(IEnv? env = null) private SecretClient GetClient() { - var url = _env.Get("KEY_VAULT_URL"); - var clientId = _env.Get("AZURE_CLIENT_ID"); - if (string.IsNullOrWhiteSpace(url)) throw new InvalidOperationException("KEY_VAULT_URL not set"); lock (_gate) { - if (_client is not null && _clientKey == (url, clientId)) return _client; + if (_client is not null) return _client; + var url = _env.Get("KEY_VAULT_URL"); + var clientId = _env.Get("AZURE_CLIENT_ID"); + if (string.IsNullOrWhiteSpace(url)) throw new InvalidOperationException("KEY_VAULT_URL not set"); var identityOptions = new TokenCredentialOptions(); identityOptions.Retry.MaxRetries = 0; - identityOptions.Retry.NetworkTimeout = CredentialCachePolicy.AcquisitionTimeout; + identityOptions.Retry.NetworkTimeout = ProviderCredentials.AcquisitionTimeout; identityOptions.Diagnostics.IsLoggingEnabled = false; identityOptions.Diagnostics.IsLoggingContentEnabled = false; var credential = new ManagedIdentityCredential(clientId, identityOptions); var options = new SecretClientOptions(); options.Retry.MaxRetries = 0; - options.Retry.NetworkTimeout = CredentialCachePolicy.AcquisitionTimeout; + options.Retry.NetworkTimeout = ProviderCredentials.AcquisitionTimeout; options.Diagnostics.IsLoggingEnabled = false; options.Diagnostics.IsLoggingContentEnabled = false; _client = new SecretClient(new Uri(url), credential, options); - _clientKey = (url, clientId); return _client; } } diff --git a/dotnet/tests/CredentialCacheTests.cs b/dotnet/tests/CredentialCacheTests.cs index f0f6d08..f49c4ec 100644 --- a/dotnet/tests/CredentialCacheTests.cs +++ b/dotnet/tests/CredentialCacheTests.cs @@ -13,58 +13,66 @@ public async Task ColdReadersShareOneFetchAndValidValuesRemainAvailableDuringRef var clock = new ManualClock(); var release = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); var calls = 0; - using var cache = new RefreshingCache(async _ => + using var manager = new ProviderCredentials(new Secrets(async (_, _) => { Interlocked.Increment(ref calls); await release.Task; - var now = clock.GetUtcNow(); - return new("value", now.AddMinutes(5), now.AddMinutes(4)); - }, () => Assert.Fail("Unexpected refresh failure"), clock); - var readers = Enumerable.Range(0, 20).Select(_ => cache.GetAsync()).ToArray(); + return "value"; + }), _ => throw new Exception(), (_, _, _) => throw new Exception(), clock: clock); + Task Get() => manager.ResolveAsync(new("apiKey", "key"), new AppConfig()); + var readers = Enumerable.Range(0, 20).Select(_ => Get()).ToArray(); Assert.Equal(1, calls); release.SetResult(); - Assert.All(await Task.WhenAll(readers), value => Assert.Equal("value", value)); + Assert.All(await Task.WhenAll(readers), value => Assert.Equal("value", value.Secret)); Assert.Equal(1, calls); Assert.Equal(1, clock.TimerCount); release = new(TaskCreationOptions.RunContinuationsAsynchronously); clock.Advance(TimeSpan.FromMinutes(4)); Assert.Equal(2, calls); - Assert.Equal("value", await cache.GetAsync()); + Assert.Equal("value", (await Get()).Secret); release.SetResult(); await Until(() => clock.TimerCount == 1); - Assert.Equal("value", await cache.GetAsync()); - cache.Dispose(); + Assert.Equal("value", (await Get()).Secret); + manager.Dispose(); Assert.Equal(0, clock.TimerCount); } [Fact] - public async Task RefreshFailuresDoNotExtendExpiryAndUseBackoff() + public async Task RefreshFailuresDoNotExtendExpiryAndUseFixedRetryCadence() { var clock = new ManualClock(); var fail = false; var calls = 0; - var failures = 0; - using var cache = new RefreshingCache(_ => + var log = new CredentialLogger(); + using var manager = new ProviderCredentials(new Secrets((_, _) => { calls++; if (fail) throw new InvalidOperationException("PRIVATE-ERROR"); - var now = clock.GetUtcNow(); - return Task.FromResult(new CredentialCacheEntry("first", now.AddMinutes(5), now.AddMinutes(4))); - }, () => failures++, clock, () => 0); - Assert.Equal("first", await cache.GetAsync()); + return Task.FromResult("first"); + }), _ => throw new Exception(), (_, _, _) => throw new Exception(), log, clock); + Task Get() => manager.ResolveAsync(new("apiKey", "key"), new AppConfig()); + Assert.Equal("first", (await Get()).Secret); fail = true; clock.Advance(TimeSpan.FromMinutes(4)); - Assert.Equal(1, failures); - for (var i = 0; i < 10; i++) Assert.Equal("first", await cache.GetAsync()); + Assert.Single(log.Entries); + for (var i = 0; i < 10; i++) Assert.Equal("first", (await Get()).Secret); Assert.Equal(2, calls); clock.Advance(TimeSpan.FromMinutes(1)); - Assert.Equal(2, failures); - var error = await Assert.ThrowsAsync(() => cache.GetAsync()); + Assert.Equal(2, log.Entries.Count); + var error = await Assert.ThrowsAsync(Get); Assert.Equal("provider credential unavailable", error.Message); + foreach (var delay in new[] { 30, 30, 30, 30, 30 }) + { + var before = calls; + clock.Advance(TimeSpan.FromSeconds(delay - 0.01)); + Assert.Equal(before, calls); + clock.Advance(TimeSpan.FromSeconds(0.01)); + Assert.Equal(before + 1, calls); + } fail = false; - clock.Advance(TimeSpan.FromSeconds(10)); - Assert.Equal("first", await cache.GetAsync()); - Assert.Equal(4, calls); + clock.Advance(TimeSpan.FromSeconds(30)); + Assert.Equal("first", (await Get()).Secret); + Assert.Equal(9, calls); } [Fact] @@ -73,20 +81,20 @@ public async Task CancellingAWaiterDoesNotCancelTheSharedRefresh() var clock = new ManualClock(); var release = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); CancellationToken observed = default; - using var cache = new RefreshingCache(async cancellation => + using var manager = new ProviderCredentials(new Secrets(async (_, cancellation) => { observed = cancellation; await release.Task; - return new("ready", clock.GetUtcNow().AddMinutes(5), clock.GetUtcNow().AddMinutes(4)); - }, () => Assert.Fail("Unexpected refresh failure"), clock); + return "ready"; + }), _ => throw new Exception(), (_, _, _) => throw new Exception(), clock: clock); using var waiter = new CancellationTokenSource(); - var first = cache.GetAsync(waiter.Token); - var second = cache.GetAsync(); + var first = manager.ResolveAsync(new("apiKey", "key"), new AppConfig(), waiter.Token); + var second = manager.ResolveAsync(new("apiKey", "key"), new AppConfig()); waiter.Cancel(); await Assert.ThrowsAnyAsync(() => first); Assert.False(observed.IsCancellationRequested); release.SetResult(); - Assert.Equal("ready", await second); + Assert.Equal("ready", (await second).Secret); } [Fact] @@ -94,22 +102,22 @@ public async Task CacheOwnedDeadlineAndShutdownPreventLatePublication() { var clock = new ManualClock(); var release = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); - var errors = 0; + var log = new CredentialLogger(); CancellationToken observed = default; - using var cache = new RefreshingCache(async cancellation => + using var manager = new ProviderCredentials(new Secrets(async (_, cancellation) => { observed = cancellation; await release.Task; - return new("late", clock.GetUtcNow().AddMinutes(5), clock.GetUtcNow().AddMinutes(4)); - }, () => errors++, clock); - var pending = cache.GetAsync(); + return "late"; + }), _ => throw new Exception(), (_, _, _) => throw new Exception(), log, clock); + var pending = manager.ResolveAsync(new("apiKey", "key"), new AppConfig()); clock.Advance(TimeSpan.FromSeconds(2.5)); await Assert.ThrowsAsync(() => pending); Assert.True(observed.IsCancellationRequested); - Assert.Equal(1, errors); - cache.Dispose(); + Assert.Single(log.Entries); + manager.Dispose(); release.SetResult(); - await Assert.ThrowsAsync(() => cache.GetAsync()); + await Assert.ThrowsAsync(() => manager.ResolveAsync(new("apiKey", "key"), new AppConfig())); Assert.Equal(0, clock.TimerCount); } @@ -128,7 +136,7 @@ public async Task ApiKeyPairIsFetchedInParallelAndPublishedAsOneBundle() if (failIdentity && name == "id") throw new InvalidOperationException("PRIVATE-ERROR"); return name + "-" + version; }); - using var manager = new ProviderCredentials(secrets, new TestEnv(), _ => throw new Exception(), + using var manager = new ProviderCredentials(secrets, _ => throw new Exception(), (_, _, _) => throw new Exception(), clock: clock); var auth = new AuthConfig("apiKey", "key", "id"); var pending = Enumerable.Range(0, 10).Select(_ => manager.ResolveAsync(auth, new AppConfig())).ToArray(); @@ -146,21 +154,21 @@ public async Task ApiKeyPairIsFetchedInParallelAndPublishedAsOneBundle() Assert.Equal("key-1", old.Secret); Assert.Equal("id-1", old.Identity); failIdentity = false; - clock.Advance(TimeSpan.FromSeconds(7)); + clock.Advance(TimeSpan.FromSeconds(30)); var next = await manager.ResolveAsync(auth, new AppConfig()); Assert.Equal("key-2", next.Secret); Assert.Equal("id-2", next.Identity); } [Fact] - public async Task ManagedIdentityAndEntraTokensAreCachedRefreshedAndIsolatedByConfiguration() + public async Task SdkCredentialsAreReusedByOneAccessTokenCacheWithoutKeyVaultCalls() { var clock = new ManualClock(); var identityCalls = 0; var providerCalls = 0; var credentialInstances = 0; using var manager = new ProviderCredentials(new Secrets((_, _) => throw new Exception("Unexpected Key Vault")), - new TestEnv(), _ => new Token(async (_, _) => + _ => new Token(async (_, _) => { Interlocked.Increment(ref identityCalls); await Task.Yield(); @@ -179,19 +187,17 @@ public async Task ManagedIdentityAndEntraTokensAreCachedRefreshedAndIsolatedByCo var config = Config(); var initial = await Task.WhenAll(Enumerable.Range(0, 20).Select(_ => manager.ResolveAsync(new("oauth"), config))); Assert.All(initial, result => Assert.Equal("PRIVATE-PROVIDER", result.AccessToken)); - Assert.Equal(1, identityCalls); + Assert.Equal(3, identityCalls); Assert.Equal(1, providerCalls); await manager.ResolveAsync(new("oauth"), config); Assert.Equal(1, providerCalls); - clock.Advance(TimeSpan.FromMinutes(55)); - await Until(() => identityCalls == 2 && providerCalls == 2 && clock.TimerCount == 2); - await manager.ResolveAsync(new("oauth"), Config(scope: "api://second/.default")); - Assert.Equal(2, identityCalls); - Assert.Equal(3, providerCalls); + clock.Advance(TimeSpan.FromMinutes(1)); + await Until(() => identityCalls == 6 && providerCalls == 2 && clock.TimerCount == 1); + Assert.Equal(1, credentialInstances); + await manager.ResolveAsync(new("oauth"), config); + Assert.Equal(6, identityCalls); + Assert.Equal(2, providerCalls); Assert.Equal(1, credentialInstances); - await manager.ResolveAsync(new("oauth"), Config(application: "different")); - Assert.Equal(3, identityCalls); - Assert.Equal(2, credentialInstances); manager.Dispose(); Assert.Equal(0, clock.TimerCount); } @@ -202,7 +208,7 @@ public async Task RepeatedSdkTokenDoesNotExtendLifetimeOrCauseATightRefreshLoop( var clock = new ManualClock(); var expiry = clock.GetUtcNow().AddHours(1); var calls = 0; - using var manager = new ProviderCredentials(new Secrets((_, _) => throw new Exception()), new TestEnv(), + using var manager = new ProviderCredentials(new Secrets((_, _) => throw new Exception()), _ => new Token((_, _) => ValueTask.FromResult(new AccessToken("assertion", expiry))), (_, _, _) => new Token((_, _) => { @@ -227,13 +233,12 @@ public async Task DisposingManagerIsTerminalAndCancelsPendingAcquisition() var release = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); var calls = 0; CancellationToken acquisition = default; - var env = new TestEnv { ["KEY_VAULT_URL"] = "https://unit.vault.azure.net" }; using var manager = new ProviderCredentials(new Secrets((_, cancellation) => { calls++; acquisition = cancellation; return release.Task; - }), env, _ => throw new Exception("Unexpected managed identity"), + }), _ => throw new Exception("Unexpected managed identity"), (_, _, _) => throw new Exception("Unexpected OAuth"), clock: clock); var auth = new AuthConfig("apiKey", "key"); var pending = manager.ResolveAsync(auth, new AppConfig()); @@ -243,20 +248,18 @@ public async Task DisposingManagerIsTerminalAndCancelsPendingAcquisition() await Assert.ThrowsAsync(() => pending); release.SetResult("PRIVATE-LATE-KEY"); await Assert.ThrowsAsync(() => manager.ResolveAsync(auth, new AppConfig())); - env["KEY_VAULT_URL"] = "https://different.vault.azure.net"; - await Assert.ThrowsAsync(() => manager.ResolveAsync(auth, new AppConfig())); await Assert.ThrowsAsync(() => manager.ResolveAsync(new("oauth"), Config())); Assert.Equal(1, calls); Assert.Equal(0, clock.TimerCount); } [Fact] - public async Task InvalidConfigurationClearsOldStateWithoutDisposingManager() + public async Task InvalidInitialConfigurationDoesNotCreateClientsOrStartATimer() { var clock = new ManualClock(); var instances = 0; using var manager = new ProviderCredentials(new Secrets((_, _) => throw new Exception("Unexpected Key Vault")), - new TestEnv(), _ => new Token((_, _) => + _ => new Token((_, _) => ValueTask.FromResult(new AccessToken("assertion", clock.GetUtcNow().AddHours(1)))), (_, _, assertion) => { @@ -267,13 +270,28 @@ public async Task InvalidConfigurationClearsOldStateWithoutDisposingManager() return new("provider-token", clock.GetUtcNow().AddHours(1)); }); }, clock: clock); - await manager.ResolveAsync(new("oauth"), Config()); - Assert.Equal(2, clock.TimerCount); await Assert.ThrowsAsync(() => manager.ResolveAsync(new("oauth"), Config(scope: ""))); Assert.Equal(0, clock.TimerCount); + Assert.Equal(0, instances); Assert.Equal("provider-token", (await manager.ResolveAsync(new("oauth"), Config())).AccessToken); - Assert.Equal(2, instances); - Assert.Equal(2, clock.TimerCount); + Assert.Equal(1, instances); + Assert.Equal(1, clock.TimerCount); + } + + [Fact] + public async Task ApiKeyCacheReplacesOneCompleteEntryAndStops() + { + var clock = new ManualClock(); + var version = "first"; + using var cache = new ApiKeyCache(new Secrets((_, _) => Task.FromResult(version)), new("apiKey", "key"), clock); + await cache.RefreshAsync(default); + Assert.Equal("first", cache.Get()?.Secret); + version = "second"; + clock.Advance(TimeSpan.FromMinutes(4)); + await cache.RefreshAsync(default); + Assert.Equal("second", cache.Get()?.Secret); + cache.Dispose(); + Assert.Null(cache.Get()); } [Fact] @@ -282,7 +300,7 @@ public async Task RefreshFailureUsesStructuredSanitizedLogRecord() var clock = new ManualClock(); var logger = new CredentialLogger(); using var manager = new ProviderCredentials(new Secrets((_, _) => - throw new InvalidOperationException("PRIVATE-SDK-ERROR")), new TestEnv(), + throw new InvalidOperationException("PRIVATE-SDK-ERROR")), _ => throw new Exception("Unexpected managed identity"), (_, _, _) => throw new Exception("Unexpected OAuth"), logger, clock); await Assert.ThrowsAsync(() => diff --git a/dotnet/tests/EngineTests.cs b/dotnet/tests/EngineTests.cs index b182d55..1f7acbf 100644 --- a/dotnet/tests/EngineTests.cs +++ b/dotnet/tests/EngineTests.cs @@ -61,8 +61,10 @@ public async Task StartupWithoutProviderConfigurationKeepsEvaluationIndependent( AssertAccepted(await rig.Invoke("evaluation")); } - [Fact] - public async Task SopranoOAuthUsesSetupIdentitiesScopeAndOneBoundedExchange() + [Theory] + [InlineData("api://provider/.default")] + [InlineData("api://second/.default")] + public async Task SopranoOAuthUsesSetupIdentitiesScopeAndOneBoundedExchange(string scope) { var scopes = new List(); var identities = new List(); @@ -88,6 +90,7 @@ public async Task SopranoOAuthUsesSetupIdentitiesScopeAndOneBoundedExchange() }); }); ConfigureSoprano(rig); + rig.Env["EPP_PROVIDER_SCOPE"] = scope; AssertAccepted(await rig.Invoke("evaluation")); Assert.Empty(applications); foreach (var channel in new[] { "sms", "voice" }) @@ -115,15 +118,12 @@ public async Task SopranoOAuthUsesSetupIdentitiesScopeAndOneBoundedExchange() } } rig.Env["EPP_PROVIDER_CHANNEL"] = "sms"; - rig.Env["EPP_PROVIDER_SCOPE"] = "api://second/.default"; AssertAccepted(await rig.Invoke()); Assert.Single(applications); - Assert.Equal(new[] { "api://provider/.default", "api://second/.default" }, scopes); - rig.Env["EPP_OUTBOUND_CLIENT_ID"] = "second-application"; - AssertAccepted(await rig.Invoke()); - Assert.Equal(new[] { ("provider-tenant", "calling-application"), ("provider-tenant", "second-application") }, applications); + Assert.Equal(new[] { scope }, scopes); + Assert.Equal(new[] { ("provider-tenant", "calling-application") }, applications); Assert.All(identities, identity => Assert.Equal("outbound-identity", identity)); - Assert.Equal(4, rig.Http.Calls); + Assert.Equal(3, rig.Http.Calls); Assert.Equal(0, rig.Secrets.Calls); Assert.DoesNotContain("private-provider-token", string.Join("\n", rig.Log.Messages)); var credential = new ProviderCredential("oauth", AccessToken: "private-provider-token"); diff --git a/javascript/README.md b/javascript/README.md index 5772b49..81e6c15 100644 --- a/javascript/README.md +++ b/javascript/README.md @@ -99,15 +99,13 @@ retries. The shared contract defines validation, HTTP outcomes and privacy-safe ## Source and extension points -Configured providers automatically prewarm on the app-start hook. Key Vault credential bundles, -managed-identity assertions and final Entra tokens refresh through separate process-local caches. -Warm requests reuse usable values; concurrent misses share a retrieval and refresh failure never -extends expiry. Acquisition deadlines and termination cancel the actual SDK HTTP transport, -including managed identity, through an HTTP-client wrapper. Termination stops timers and closes -the manager permanently; configuration replacement uses a separate cache reset. Startup/refresh -never sends an OTP. See the -[refresh contract](../docs/CONTRACT.md#credential-caching-and-refresh) for budgets and cold-start -limitations. Leave the provider unset for local evaluation-only use without credential acquisition. +The app-start hook selects `ApiKeyCache` or `AccessTokenCache` from the provider manifest's auth mode. +Only that cache starts: API keys use Key Vault and `lru-cache`; access tokens use the MI/Entra SDKs, +without Key Vault. One shared 30-second refresh loop and one in-flight acquisition keep warm reads +nonblocking. Configuration changes require restart; failures never extend expiry. A small HTTP-client +wrapper propagates cancellation to the installed identity SDK. Shutdown prevents late publication. See the +[refresh contract](../docs/CONTRACT.md#credential-caching-and-refresh). Leave the provider unset for +local evaluation-only use without credential acquisition; prewarming never sends an OTP. | Source | Purpose | |---|---| @@ -115,7 +113,7 @@ limitations. Leave the provider unset for local evaluation-only use without cred | [src/functions/config.js](src/functions/config.js) | Shared deployment settings | | [src/functions/models.js](src/functions/models.js) | Delivery context, normalized `ParsedResponse`, and documented request objects | | [src/functions/dispatch.js](src/functions/dispatch.js) | Envelope/JWE handling, registry and dispatch | -| [src/functions/credentials.js](src/functions/credentials.js), [refreshingCache.js](src/functions/refreshingCache.js) | Provider credential acquisition, single-flight caching and scheduled refresh | +| [src/functions/credentials.js](src/functions/credentials.js) | `ApiKeyCache`, `AccessTokenCache` and their shared refresh coordinator | | [src/functions/requestLog.js](src/functions/requestLog.js) | Request-scoped [service events and summaries](../docs/CONTRACT.md#application-logs) with explicit ID sources | | [src/functions/providers/](src/functions/providers/) | Adapter manifests and API-specific implementations | | [test/](test/) | Representative offline checks | diff --git a/javascript/package-lock.json b/javascript/package-lock.json index d5654ec..cd0c455 100644 --- a/javascript/package-lock.json +++ b/javascript/package-lock.json @@ -13,7 +13,8 @@ "@azure/identity": "^4.13.1", "@azure/keyvault-secrets": "^4.11.2", "@azure/logger": "1.3.0", - "jose": "^5.9.6" + "jose": "^5.9.6", + "lru-cache": "^11.1.0" } }, "node_modules/@azure-rest/core-client": { @@ -523,6 +524,14 @@ "integrity": "sha512-Sb487aTOCr9drQVL8pIxOzVhafOjZN9UU54hiN8PU3uAiSV7lx1yYNpbNmex2PK6dSJoNTSJUUswT651yww3Mg==", "license": "MIT" }, + "node_modules/lru-cache": { + "version": "11.5.2", + "integrity": "sha1-AOFmZckMYg+6FKPDaHMql2ST92A=", + "license": "BlueOak-1.0.0", + "engines": { + "node": "20 || >=22" + } + }, "node_modules/ms": { "version": "2.1.3", "integrity": "sha512-6FlzubTLZG3J2a/NVCAleEhjzq5oxgHyaCU9yYXvcLsvoVaHJq/s5xXI6/XXP6tz7R9xAOtHnSO/tXtF3WRTlA==", diff --git a/javascript/package.json b/javascript/package.json index 483d488..0522d49 100644 --- a/javascript/package.json +++ b/javascript/package.json @@ -13,6 +13,7 @@ "@azure/identity": "^4.13.1", "@azure/keyvault-secrets": "^4.11.2", "@azure/logger": "1.3.0", - "jose": "^5.9.6" + "jose": "^5.9.6", + "lru-cache": "^11.1.0" } } diff --git a/javascript/src/functions/credentials.js b/javascript/src/functions/credentials.js index cfd93f7..29b982c 100644 --- a/javascript/src/functions/credentials.js +++ b/javascript/src/functions/credentials.js @@ -3,42 +3,33 @@ const { AsyncLocalStorage } = require('node:async_hooks'); const { inspect } = require('node:util'); +const { LRUCache } = require('lru-cache'); const { createDefaultHttpClient } = require('@azure/core-rest-pipeline'); -const { ClientAssertionCredential, ManagedIdentityCredential } = require('@azure/identity'); +const { AzureAuthorityHosts, ClientAssertionCredential, ManagedIdentityCredential } = require('@azure/identity'); const { SecretClient } = require('@azure/keyvault-secrets'); const { AzureLogger } = require('@azure/logger'); -const { CACHE_POLICY, RefreshingCache, tokenEntry } = require('./refreshingCache'); - -/** - * @template T - * @typedef {InstanceType>} CredentialCache - */ - -/** - * @typedef {ReturnType} AppConfig - * @typedef {{mode?: string, keyVaultSecretName?: string, identityKeyVaultSecretName?: string}} AuthConfig - * @typedef {{mode: 'apiKey', secret: string, identity: string}} ApiKeyCredential - * @typedef {{mode: 'oauth', accessToken: string}} OAuthCredential - * @typedef {import('@azure/core-auth').AccessToken} AccessToken - * @typedef {{mode: 'apiKey', bundle: CredentialCache}} ApiKeyState - * @typedef {object} OAuthState - * @property {'oauth'} mode - * @property {CredentialCache} assertion - * @property {ClientAssertionCredential} credential - * @property {Map>} tokens - */ +const API_KEY_MODE = 'apiKey'; +const OAUTH_MODE = 'oauth'; +const BUNDLE_KEY = 'bundle'; +const TOKEN_EXCHANGE_SCOPE = 'api://AzureADTokenExchange/.default'; +const ACQUISITION_TIMEOUT_MS = 2500; +const REFRESH_POLL_MS = 30000; +const SECRET_TTL_MS = 300000; +const SECRET_REFRESH_MS = 240000; +const TOKEN_SKEW_MS = 30000; +const unavailable = () => new Error('provider credential unavailable'); /** @type {AsyncLocalStorage} */ const acquisition = new AsyncLocalStorage(); /** @type {typeof AzureLogger.log | undefined} */ let filteredLogger; -/** @param {string} cacheKind */ function reportRefreshFailure(cacheKind) { console.warn(JSON.stringify({ logType: 'service', eventName: 'credential_refresh_failed', cacheKind, failureReason: 'credential_unavailable' })); } +// The installed identity SDK does not propagate every getToken abort signal to HTTP. /** @returns {import('@azure/core-rest-pipeline').HttpClient} */ function credentialHttpClient() { const client = createDefaultHttpClient(); @@ -64,215 +55,188 @@ function credentialHttpClient() { }; } +const sdkOptions = () => ({ retryOptions: { maxRetries: 0 }, httpClient: credentialHttpClient() }); +/** @param {import('@azure/core-auth').AccessToken | null} token */ +function checkToken(token, now) { + if (typeof token?.token !== 'string' || !token.token.trim() || !Number.isFinite(token.expiresOnTimestamp) + || token.expiresOnTimestamp <= now + TOKEN_SKEW_MS) throw unavailable(); + return token; +} + /** - * @template T - * @param {AbortSignal} signal - * @param {(signal: AbortSignal) => Promise} load - * @returns {Promise} + * @typedef {{mode: 'apiKey', secret: string, identity: string}} ApiKeyCredential + * @typedef {{mode: 'oauth', accessToken: string}} OAuthCredential + * @typedef {ReturnType} AppConfig + * @typedef {{mode?: string, keyVaultSecretName?: string, identityKeyVaultSecretName?: string}} AuthConfig + * @typedef {{now?: () => number, schedule?: typeof setInterval, cancel?: typeof clearInterval}} RefreshOptions */ -async function acquireBounded(signal, load) { - if (AzureLogger.log !== filteredLogger) { - const previous = AzureLogger.log; - filteredLogger = (...args) => { if (!acquisition.getStore()) previous(...args); }; - AzureLogger.log = filteredLogger; + +class ApiKeyCache { + /** @param {AuthConfig} auth @param {AppConfig} config */ + constructor(auth, config, now = Date.now) { + this.now = now; + this.auth = { ...auth }; + this.vaultUrl = config.keyVaultUrl; + this.identityClientId = config.managedIdentityClientId; + /** @type {SecretClient | undefined} */ + this.client = undefined; + /** @type {LRUCache} */ + this.values = new LRUCache({ max: 1, ttlResolution: 0 }); + this.refreshAt = 0; + this.stage = 'key_vault'; + this.closed = false; } - const controller = new AbortController(); - const abort = () => controller.abort(); - signal.addEventListener('abort', abort, { once: true }); - if (signal.aborted) abort(); - let timeout; - /** @type {Promise} */ - const interrupted = new Promise((_, reject) => { - const fail = () => reject(new Error('provider credential unavailable')); - controller.signal.addEventListener('abort', fail, { once: true }); - if (controller.signal.aborted) fail(); - timeout = setTimeout(abort, CACHE_POLICY.acquisitionTimeoutMs); - }); - try { - return await acquisition.run(controller.signal, () => Promise.race([ - Promise.resolve().then(() => { - controller.signal.throwIfAborted(); - return load(controller.signal); - }), - interrupted, - ])); - } finally { - clearTimeout(timeout); - signal.removeEventListener('abort', abort); + get() { + const entry = this.closed ? undefined : this.values.get(BUNDLE_KEY); + return entry && entry.expiresAt > this.now() ? entry.value : null; + } + async refresh(signal) { + if (this.closed) throw unavailable(); + if (this.get() && this.refreshAt > this.now()) return; + const { keyVaultSecretName: key, identityKeyVaultSecretName: account } = this.auth; + if (!key || !this.vaultUrl) throw unavailable(); + this.client ??= new SecretClient(this.vaultUrl, + new ManagedIdentityCredential({ clientId: this.identityClientId || undefined, ...sdkOptions() }), sdkOptions()); + const [secret, identity] = await Promise.all([ + this.client.getSecret(key, { abortSignal: signal }), + account ? this.client.getSecret(account, { abortSignal: signal }) : null, + ]); + const now = this.now(); + let expiresAt = now + SECRET_TTL_MS; + for (const item of [secret, identity]) { + if (!item) continue; + if (!item.value?.trim() || item.properties?.enabled === false || item.properties?.notBefore?.getTime() > now) throw unavailable(); + if (item.properties?.expiresOn) expiresAt = Math.min(expiresAt, item.properties.expiresOn.getTime()); + } + if (!secret.value?.trim() || (account && !identity?.value?.trim()) + || !Number.isFinite(expiresAt) || expiresAt <= now || this.closed) throw unavailable(); + signal.throwIfAborted(); + this.values.set(BUNDLE_KEY, { value: Object.freeze({ mode: API_KEY_MODE, secret: secret.value, identity: identity?.value || '' }), + expiresAt }, { ttl: expiresAt - now }); + this.refreshAt = Math.min(now + SECRET_REFRESH_MS, expiresAt - TOKEN_SKEW_MS); } + stop() { this.closed = true; this.values.clear(); } + [inspect.custom]() { return '[ApiKeyCache]'; } + toJSON() { return '[ApiKeyCache]'; } } -// ManagedIdentityCredential replaces additionalPolicies, but preserves the supplied HTTP client. -const sdkOptions = () => ({ - retryOptions: { maxRetries: 0 }, - httpClient: credentialHttpClient(), -}); +class AccessTokenCache { + /** @param {AppConfig} config */ + constructor(config, now = Date.now) { + if (!config.providerTenantId || !config.providerScope || !config.outboundClientId + || !config.outboundManagedIdentityClientId) throw unavailable(); + this.now = now; + this.scope = config.providerScope; + this.identity = new ManagedIdentityCredential({ clientId: config.outboundManagedIdentityClientId, ...sdkOptions() }); + this.credential = new ClientAssertionCredential(config.providerTenantId, config.outboundClientId, + async () => (await this.assertion(acquisition.getStore())).token, + { authorityHost: AzureAuthorityHosts.AzurePublicCloud, ...sdkOptions() }); + /** @type {import('@azure/core-auth').AccessToken | null} */ + this.token = null; + this.stage = 'provider_token'; + this.closed = false; + } + get() { + if (this.closed || !this.token || this.token.expiresOnTimestamp <= this.now() + TOKEN_SKEW_MS) return null; + /** @type {OAuthCredential} */ + const value = { mode: OAUTH_MODE, accessToken: this.token.token }; + return Object.defineProperty(value, 'accessToken', { enumerable: false }); + } + async assertion(signal) { + return checkToken(await this.identity.getToken(TOKEN_EXCHANGE_SCOPE, { abortSignal: signal }), this.now()); + } + async refresh(signal) { + if (this.closed) throw unavailable(); + this.stage = 'managed_identity'; + await this.assertion(signal); + this.stage = 'provider_token'; + const token = checkToken(await this.credential.getToken(this.scope, { abortSignal: signal }), this.now()); + signal.throwIfAborted(); + if (this.closed) throw unavailable(); + this.token = token; + } + stop() { this.closed = true; this.token = null; } + [inspect.custom]() { return '[AccessTokenCache]'; } + toJSON() { return '[AccessTokenCache]'; } +} +// Owns one selected cache and one periodic refresh; configuration changes require a worker restart. class ProviderCredentials { - /** - * @param {{cacheOptions?: import('./refreshingCache').CacheOptions, - * reportFailure?: (kind: string) => void}} [options] - */ + /** @param {{cacheOptions?: RefreshOptions, reportFailure?: (kind: string) => void}} [options] */ constructor({ cacheOptions = {}, reportFailure = reportRefreshFailure } = {}) { - this.cacheOptions = cacheOptions; this.now = cacheOptions.now || Date.now; + this.schedule = cacheOptions.schedule || setInterval; + this.cancel = cacheOptions.cancel || clearInterval; this.reportFailure = reportFailure; - /** @type {ApiKeyState | OAuthState | null} */ + /** @type {ApiKeyCache | AccessTokenCache | null} */ this.current = null; - /** @type {string | null} */ - this.currentKey = null; + /** @type {Promise | null} */ + this.pending = null; + /** @type {ReturnType | null} */ + this.timer = null; + /** @type {AbortController | null} */ + this.controller = null; + this.nextAttemptAt = 0; this.closed = false; } - - [inspect.custom]() { return '[ProviderCredentials]'; } - toJSON() { return '[ProviderCredentials]'; } - - /** - * @template T - * @param {string} kind - * @param {(signal: AbortSignal) => Promise>} load - * @returns {CredentialCache} - */ - cache(kind, load) { - return new RefreshingCache((signal) => acquireBounded(signal, load), - { ...this.cacheOptions, onFailure: () => this.reportFailure(kind) }); - } - - /** - * @param {AuthConfig} auth - * @param {AppConfig} config - * @returns {Promise} - */ - resolve(auth, config) { - if (this.closed) return Promise.reject(new Error('provider credential unavailable')); - const mode = auth?.mode || 'apiKey'; - if (mode === 'oauth' && (!config.providerTenantId || !config.providerScope - || !config.outboundClientId || !config.outboundManagedIdentityClientId)) { - this.#clear(); - return Promise.reject(new Error('provider OAuth token unavailable')); - } - if (mode !== 'oauth' && mode !== 'apiKey') { - this.#clear(); - return Promise.reject(new Error('provider credential unavailable')); - } - let key; - if (mode === 'apiKey') { - key = JSON.stringify([ - mode, config.keyVaultUrl, config.managedIdentityClientId, - auth.keyVaultSecretName, auth.identityKeyVaultSecretName, - ]); - } else { - key = JSON.stringify([ - mode, config.providerTenantId, config.outboundClientId, config.outboundManagedIdentityClientId, - ]); - } - if (this.currentKey !== key) { - this.#clear(); - if (mode === 'apiKey') { - this.current = this.apiKeyState(auth, config); - } else { - this.current = this.oauthState(config); - } - this.currentKey = key; - } - const state = this.current; - if (!state) return Promise.reject(new Error('provider credential unavailable')); - if (state.mode === 'apiKey') return state.bundle.get(); - let tokenCache = state.tokens.get(config.providerScope); - if (!tokenCache) { - const scope = config.providerScope; - tokenCache = this.cache('provider_token', async (signal) => - tokenEntry(await state.credential.getToken(scope, { abortSignal: signal }), this.now())); - state.tokens.set(scope, tokenCache); + /** @param {AuthConfig} auth @param {AppConfig} config */ + async resolve(auth, config) { + if (this.closed) throw unavailable(); + if (!this.current) { + try { + switch (auth.mode || API_KEY_MODE) { + case API_KEY_MODE: this.current = new ApiKeyCache(auth, config, this.now); break; + case OAUTH_MODE: this.current = new AccessTokenCache(config, this.now); break; + default: throw unavailable(); + } + this.timer = this.schedule(() => { void this.refresh().catch(() => {}); }, REFRESH_POLL_MS); + this.timer.unref?.(); + } catch { this.reportFailure('configuration'); throw unavailable(); } } - return tokenCache.get().then((token) => { - /** @type {OAuthCredential} */ - const credential = { mode: 'oauth', accessToken: token.token }; - Object.defineProperty(credential, 'accessToken', { enumerable: false, writable: false, configurable: false }); - return credential; - }).catch(() => { throw new Error('provider OAuth token unavailable'); }); - } - - /** - * @param {AuthConfig} auth - * @param {AppConfig} config - * @returns {ApiKeyState} - */ - apiKeyState(auth, config) { - /** @type {SecretClient | undefined} */ - let client; - const bundle = this.cache('key_vault', async (signal) => { - if (!auth.keyVaultSecretName || !config.keyVaultUrl) throw new Error('provider credential unavailable'); - if (!client) { - const identity = config.managedIdentityClientId - ? new ManagedIdentityCredential(config.managedIdentityClientId, sdkOptions()) - : new ManagedIdentityCredential(sdkOptions()); - client = new SecretClient(config.keyVaultUrl, identity, sdkOptions()); - } - const [secret, identity] = await Promise.all([ - client.getSecret(auth.keyVaultSecretName, { abortSignal: signal }), - auth.identityKeyVaultSecretName ? client.getSecret(auth.identityKeyVaultSecretName, { abortSignal: signal }) : null, - ]); - if (!secret.value?.trim() || (auth.identityKeyVaultSecretName && !identity?.value?.trim())) { - throw new Error('provider credential unavailable'); - } - const now = this.now(); - let expiresAt = now + CACHE_POLICY.secretTtlMs; - for (const item of [secret, identity]) { - if (!item) continue; - const notBefore = item.properties?.notBefore?.getTime(); - if (item.properties?.enabled === false - || (notBefore !== undefined && notBefore > now)) throw new Error('provider credential unavailable'); - if (item.properties?.expiresOn) expiresAt = Math.min(expiresAt, item.properties.expiresOn.getTime()); - } - /** @type {ApiKeyCredential} */ - const credential = { mode: 'apiKey', secret: secret.value, identity: identity?.value || '' }; - return { - value: Object.freeze(credential), - expiresAt, - refreshAt: Math.min(now + CACHE_POLICY.secretRefreshIntervalMs, expiresAt - CACHE_POLICY.secretExpiryRefreshLeadMs), - }; - }); - return { mode: 'apiKey', bundle }; - } - - /** - * @param {AppConfig} config - * @returns {OAuthState} - */ - oauthState(config) { - /** @type {ManagedIdentityCredential | undefined} */ - let identity; - const assertion = this.cache('managed_identity', async (signal) => { - identity ??= new ManagedIdentityCredential({ clientId: config.outboundManagedIdentityClientId, ...sdkOptions() }); - return tokenEntry(await identity.getToken('api://AzureADTokenExchange/.default', { abortSignal: signal }), this.now()); - }); - const credential = new ClientAssertionCredential(config.providerTenantId, config.outboundClientId, - async () => { - const value = await assertion.get(); - acquisition.getStore()?.throwIfAborted(); - return value.token; - }, { authorityHost: 'https://login.microsoftonline.com', ...sdkOptions() }); - return { mode: 'oauth', assertion, credential, tokens: new Map() }; + const cached = this.current.get(); + if (cached) return cached; + await this.refresh(); + const value = this.current.get(); + if (!value) throw unavailable(); + return value; } - - #clear() { - const state = this.current; - if (state?.mode === 'apiKey') { - state.bundle.close(); - } else if (state?.mode === 'oauth') { - state.assertion.close(); - for (const cache of state.tokens.values()) cache.close(); + refresh() { + if (this.closed || !this.current) return Promise.reject(unavailable()); + if (this.pending) return this.pending; + if (this.nextAttemptAt > this.now()) return Promise.reject(unavailable()); + this.nextAttemptAt = this.now() + REFRESH_POLL_MS; + const cache = this.current; + const controller = this.controller = new AbortController(); + if (AzureLogger.log !== filteredLogger) { + const previous = AzureLogger.log; + filteredLogger = (...args) => { if (!acquisition.getStore()) previous(...args); }; + AzureLogger.log = filteredLogger; } - this.current = null; - this.currentKey = null; + const timeout = setTimeout(() => controller.abort(), ACQUISITION_TIMEOUT_MS); + const interrupted = new Promise((_, reject) => + controller.signal.addEventListener('abort', () => reject(unavailable()), { once: true })); + this.pending = acquisition.run(controller.signal, () => Promise.race([ + Promise.resolve().then(() => { controller.signal.throwIfAborted(); return cache.refresh(controller.signal); }), interrupted, + ])).catch(() => { + controller.abort(); + if (!this.closed) this.reportFailure(cache.stage); + throw unavailable(); + }).finally(() => { + clearTimeout(timeout); + this.pending = null; + this.controller = null; + }); + return this.pending; } - close() { this.closed = true; - this.#clear(); + if (this.timer) this.cancel(this.timer); + this.controller?.abort(); + this.current?.stop(); } + [inspect.custom]() { return '[ProviderCredentials]'; } + toJSON() { return '[ProviderCredentials]'; } } const providerCredentials = new ProviderCredentials(); -module.exports = { ProviderCredentials, providerCredentials, reportRefreshFailure }; +module.exports = { ApiKeyCache, AccessTokenCache, ProviderCredentials, providerCredentials, reportRefreshFailure }; diff --git a/javascript/src/functions/refreshingCache.js b/javascript/src/functions/refreshingCache.js deleted file mode 100644 index 75fdf67..0000000 --- a/javascript/src/functions/refreshingCache.js +++ /dev/null @@ -1,176 +0,0 @@ -// Copyright (c) Microsoft Corporation. All rights reserved. -'use strict'; - -const { inspect } = require('node:util'); - -const CACHE_POLICY = Object.freeze({ - acquisitionTimeoutMs: 2500, - secretTtlMs: 5 * 60 * 1000, - secretRefreshIntervalMs: 4 * 60 * 1000, - secretExpiryRefreshLeadMs: 30 * 1000, - tokenExpirySkewMs: 30 * 1000, - tokenRefreshLeadMs: 5 * 60 * 1000, - minRefreshDelayMs: 1000, - maxRefreshDelayMs: 60 * 1000, - initialRetryDelayMs: 5000, - maxRetryDelayMs: 60 * 1000, - maxRetryExponent: 4, - retryJitterRatio: 0.2, - maxTimerDelayMs: 2 ** 31 - 1, -}); - -/** - * @template T - * @typedef {object} CacheEntry - * @property {T} value - * @property {number} expiresAt Absolute Unix time in milliseconds. - * @property {number} refreshAt Absolute Unix time in milliseconds. - */ - -/** - * @typedef {object} CacheOptions - * @property {() => number} [now] Unix time in milliseconds. - * @property {typeof setTimeout} [schedule] - * @property {typeof clearTimeout} [cancel] - * @property {() => number} [random] - * @property {() => void} [onFailure] - */ - -const unavailable = () => new Error('provider credential unavailable'); - -/** @template T */ -class RefreshingCache { - /** - * @param {(signal: AbortSignal) => Promise>} load - * @param {CacheOptions} [options] - */ - constructor(load, { now = Date.now, schedule = setTimeout, cancel = clearTimeout, - random = Math.random, onFailure = () => {} } = {}) { - this.load = load; - this.now = now; - this.schedule = schedule; - this.cancel = cancel; - this.random = random; - this.onFailure = onFailure; - /** @type {CacheEntry | null} */ - this.entry = null; - /** @type {Promise | null} */ - this.inFlight = null; - /** @type {ReturnType | null} */ - this.timer = null; - this.retryAt = 0; - this.failures = 0; - this.closed = false; - /** @type {AbortController | null} */ - this.controller = null; - } - - [inspect.custom]() { return '[RefreshingCache]'; } - toJSON() { return '[RefreshingCache]'; } - - get() { - if (this.closed) return Promise.reject(unavailable()); - if (this.entry && this.entry.expiresAt > this.now()) { - if (this.entry.refreshAt <= this.now() && this.retryAt <= this.now()) { - this.refreshInBackground(); - } - return Promise.resolve(this.entry.value); - } - if (this.inFlight) return this.inFlight; - if (this.retryAt > this.now()) return Promise.reject(unavailable()); - return this.refresh(); - } - - refresh() { - if (this.closed) return Promise.reject(unavailable()); - if (this.inFlight) return this.inFlight; - if (this.retryAt > this.now()) return Promise.reject(unavailable()); - this.clearTimer(); - this.controller = new AbortController(); - const controller = this.controller; - // Publish the promise before calling a loader that may synchronously reenter the cache. - const operation = Promise.resolve().then(() => this.load(controller.signal)).then((entry) => { - if (this.closed || controller.signal.aborted) throw unavailable(); - if (!entry || !Number.isFinite(entry.expiresAt) || entry.expiresAt <= this.now() - || !Number.isFinite(entry.refreshAt)) throw unavailable(); - this.entry = entry; - this.failures = 0; - this.retryAt = 0; - return entry.value; - }).catch(() => { - if (!this.closed) { - this.failures++; - const exponent = Math.min(this.failures - 1, CACHE_POLICY.maxRetryExponent); - const backoff = Math.min(CACHE_POLICY.maxRetryDelayMs, CACHE_POLICY.initialRetryDelayMs * 2 ** exponent); - this.retryAt = this.now() + Math.floor(backoff * (1 + this.random() * CACHE_POLICY.retryJitterRatio)); - this.onFailure(); - } - throw unavailable(); - }).finally(() => { - this.inFlight = null; - this.controller = null; - if (!this.closed) { - const next = this.retryAt || this.entry?.refreshAt; - if (next != null) this.arm(next); - } - }); - this.inFlight = operation; - return operation; - } - - refreshInBackground() { - if (!this.closed && this.retryAt > this.now()) { - this.arm(this.retryAt); - return; - } - // refresh reports its failure; a timer has no request waiting to handle the rejection. - void this.refresh().catch(() => {}); - } - - /** @param {number} at Absolute Unix time in milliseconds. */ - arm(at) { - this.clearTimer(); - this.timer = this.schedule(() => { - this.timer = null; - this.refreshInBackground(); - }, Math.max(CACHE_POLICY.minRefreshDelayMs, Math.min(CACHE_POLICY.maxTimerDelayMs, at - this.now()))); - this.timer?.unref?.(); - } - - clearTimer() { - if (this.timer !== null) this.cancel(this.timer); - this.timer = null; - } - - close() { - this.closed = true; - this.clearTimer(); - this.controller?.abort(); - this.entry = null; - } -} - -/** - * @param {import('@azure/core-auth').AccessToken | null | undefined} token - * @param {number} now Unix time in milliseconds. - * @returns {CacheEntry} - */ -function tokenEntry(token, now) { - if (typeof token?.token !== 'string' || !token.token.trim() - || !Number.isFinite(token.expiresOnTimestamp) || token.expiresOnTimestamp <= now + CACHE_POLICY.tokenExpirySkewMs) { - throw unavailable(); - } - const expiresAt = token.expiresOnTimestamp - CACHE_POLICY.tokenExpirySkewMs; - let refreshAt = token.expiresOnTimestamp - CACHE_POLICY.tokenRefreshLeadMs; - if (typeof token.refreshAfterTimestamp === 'number' && Number.isFinite(token.refreshAfterTimestamp)) { - refreshAt = Math.min(refreshAt, token.refreshAfterTimestamp); - } - // An SDK may return its existing token on refresh. Never extend that token's lifetime or spin. - if (refreshAt <= now) { - const delay = Math.min(CACHE_POLICY.maxRefreshDelayMs, (expiresAt - now) / 2); - refreshAt = now + Math.max(CACHE_POLICY.minRefreshDelayMs, delay); - } - return { value: token, expiresAt, refreshAt }; -} - -module.exports = { CACHE_POLICY, RefreshingCache, tokenEntry }; diff --git a/javascript/test/credential-cache.test.js b/javascript/test/credential-cache.test.js index 8651d31..fbf7786 100644 --- a/javascript/test/credential-cache.test.js +++ b/javascript/test/credential-cache.test.js @@ -3,29 +3,26 @@ const { test } = require('node:test'); const assert = require('node:assert/strict'); const { inspect } = require('node:util'); -const { RefreshingCache, tokenEntry } = require('../src/functions/refreshingCache'); -const { ProviderCredentials } = require('../src/functions/credentials'); +const { ApiKeyCache, AccessTokenCache, ProviderCredentials } = require('../src/functions/credentials'); const { ClientAssertionCredential, ManagedIdentityCredential } = require('@azure/identity'); const { SecretClient } = require('@azure/keyvault-secrets'); const { readConfig } = require('../src/functions/config'); const flush = () => new Promise(setImmediate); -const deferred = () => { - let resolve; - let reject; - const promise = new Promise((yes, no) => { resolve = yes; reject = no; }); - return { promise, resolve, reject }; -}; - +const auth = { mode: 'apiKey', keyVaultSecretName: 'key', identityKeyVaultSecretName: 'id' }; +const config = readConfig({ KEY_VAULT_URL: 'https://unit.vault.azure.net' }); +const oauth = readConfig({ + EPP_PROVIDER_TENANT_ID: 'tenant', EPP_OUTBOUND_CLIENT_ID: 'app', + EPP_OUTBOUND_MI_CLIENT_ID: 'identity', EPP_PROVIDER_SCOPE: 'scope', +}); function clock() { let time = 1700000000000; const timers = new Set(); return { options: { now: () => time, - random: () => 0, schedule: (callback, delay) => { - const timer = { at: time + delay, callback, unref() {} }; + const timer = { at: time + delay, period: delay, callback, unref() {} }; timers.add(timer); return timer; }, @@ -36,234 +33,195 @@ function clock() { async advance(ms) { time += ms; for (const timer of [...timers]) { - if (timer.at <= time && timers.delete(timer)) timer.callback(); + if (timer.at <= time && timers.has(timer)) { + timer.at = time + timer.period; + timer.callback(); + } } await flush(); }, }; } -test('concurrent empty-cache readers join one refresh; fresh hits make no loader calls', async () => { +test('library-backed bundle shares concurrent reads, serves during refresh, and publishes pairs atomically', async (t) => { const time = clock(); - const gate = deferred(); - let calls = 0; - const cache = new RefreshingCache(async () => { - calls++; - await gate.promise; - return { value: 'PRIVATE-VALUE', expiresAt: time.now + 300000, refreshAt: time.now + 240000 }; - }, time.options); + let release; + let gate = new Promise((resolve) => { release = resolve; }); + let version = 1; + const getSecret = t.mock.method(SecretClient.prototype, 'getSecret', async (name) => { + await gate; + return { value: `${name}-${version}` }; + }); + const manager = new ProviderCredentials({ cacheOptions: time.options }); try { - const readers = Array.from({ length: 20 }, () => cache.get()); + const pending = Array.from({ length: 20 }, () => manager.resolve(auth, config)); await flush(); - assert.equal(calls, 1); - gate.resolve(); - assert.deepEqual(await Promise.all(readers), Array(20).fill('PRIVATE-VALUE')); - assert.equal(await cache.get(), 'PRIVATE-VALUE'); - assert.equal(calls, 1); + assert.equal(getSecret.mock.callCount(), 2); + release(); + assert.ok((await Promise.all(pending)).every((value) => value.secret === 'key-1' && value.identity === 'id-1')); + assert.ok(manager.current instanceof ApiKeyCache); assert.equal(time.timerCount, 1); - assert.doesNotMatch(inspect(cache) + JSON.stringify(cache), /PRIVATE/); - } finally { cache.close(); } - assert.equal(time.timerCount, 0); -}); - -test('scheduled refresh serves the valid old entry, replaces atomically, and never extends a failed entry', async () => { - const time = clock(); - let next = async () => ({ value: 'first', expiresAt: time.now + 300000, refreshAt: time.now + 240000 }); - const failures = []; - const cache = new RefreshingCache(() => next(), { ...time.options, onFailure: () => failures.push('failed') }); - try { - assert.equal(await cache.get(), 'first'); - const gate = deferred(); - next = () => gate.promise; + assert.doesNotMatch(inspect(manager) + JSON.stringify(manager), /key-1|id-1/); + gate = new Promise((resolve) => { release = resolve; }); + version = 2; await time.advance(240000); - assert.equal(await cache.get(), 'first'); - gate.reject(new Error('PRIVATE-REFRESH-ERROR')); + assert.deepEqual(await manager.resolve(auth, config), { mode: 'apiKey', secret: 'key-1', identity: 'id-1' }); + release(); await flush(); - assert.deepEqual(failures, ['failed']); - assert.equal(await cache.get(), 'first'); - await time.advance(60000); - await assert.rejects(cache.get(), /provider credential unavailable/); - next = async () => ({ value: 'recovered', expiresAt: time.now + 300000, refreshAt: time.now + 240000 }); - await time.advance(10000); - assert.equal(await cache.get(), 'recovered'); - } finally { cache.close(); } -}); - -test('failed cold refresh applies bounded backoff instead of a fetch per request', async () => { - const time = clock(); - let calls = 0; - const cache = new RefreshingCache(async () => { calls++; throw new Error('PRIVATE-FAILURE'); }, time.options); - try { - await assert.rejects(cache.get(), /^Error: provider credential unavailable$/); - for (let i = 0; i < 10; i++) await assert.rejects(cache.get(), /unavailable/); - assert.equal(calls, 1); - await time.advance(4999); - assert.equal(calls, 1); - await time.advance(1); - assert.equal(calls, 2); - await time.advance(9999); - assert.equal(calls, 2); - await time.advance(1); - assert.equal(calls, 3); - } finally { cache.close(); } -}); - -test('stopping a cache cancels refresh and prevents late publication or rescheduling', async () => { - const time = clock(); - const gate = deferred(); - let signal; - const cache = new RefreshingCache(async (abortSignal) => { - signal = abortSignal; - return gate.promise; - }, time.options); - const pending = cache.get(); - await flush(); - cache.close(); - assert.equal(signal.aborted, true); - gate.resolve({ value: 'late', expiresAt: time.now + 300000, refreshAt: time.now + 240000 }); - await assert.rejects(pending, /unavailable/); - await assert.rejects(cache.get(), /unavailable/); + assert.deepEqual(await manager.resolve(auth, config), { mode: 'apiKey', secret: 'key-2', identity: 'id-2' }); + assert.equal(getSecret.mock.callCount(), 4); + } finally { release(); manager.close(); } assert.equal(time.timerCount, 0); }); -test('token refresh retains original expiry, honors refresh hints, and avoids spinning on SDK cache hits', () => { - const now = 1700000000000; - const token = { token: 'PRIVATE-TOKEN', expiresOnTimestamp: now + 3600000 }; - assert.deepEqual(tokenEntry(token, now), { value: token, expiresAt: now + 3570000, refreshAt: now + 3300000 }); - const repeated = tokenEntry(token, now + 3300000); - assert.equal(repeated.expiresAt, now + 3570000); - assert.equal(repeated.refreshAt, now + 3360000); - assert.equal(tokenEntry({ ...token, refreshAfterTimestamp: now + 600000 }, now).refreshAt, now + 600000); - for (const invalid of [null, { ...token, token: '' }, { ...token, token: ' ' }, - { ...token, expiresOnTimestamp: now + 30000 }, { ...token, expiresOnTimestamp: NaN }, - { ...token, expiresOnTimestamp: Infinity }, { token: 'PRIVATE' }]) { - assert.throws(() => tokenEntry(invalid, now), /unavailable/); - } -}); - -test('Key Vault refresh fetches a parallel credential bundle once and keeps a complete old pair on partial failure', async (t) => { +test('partial refresh failure retains the old pair only until hard expiry, with fixed retry cadence', async (t) => { const time = clock(); - let failIdentity = false; - let version = 1; + let fail = false; const getSecret = t.mock.method(SecretClient.prototype, 'getSecret', async (name) => { - if (failIdentity && name === 'customer-id') throw new Error('PRIVATE-IDENTITY-ERROR'); - return { value: `${name}-${version}` }; + if (fail && name === 'id') throw new Error('PRIVATE-FAILURE'); + return { value: name }; }); const failures = []; const manager = new ProviderCredentials({ cacheOptions: time.options, reportFailure: (kind) => failures.push(kind) }); - const auth = { mode: 'apiKey', keyVaultSecretName: 'api-key', identityKeyVaultSecretName: 'customer-id' }; - const config = readConfig({ KEY_VAULT_URL: 'https://unit.vault.azure.net' }); try { - const results = await Promise.all(Array.from({ length: 10 }, () => manager.resolve(auth, config))); - assert.equal(getSecret.mock.callCount(), 2); - assert.ok(results.every((value) => value.secret === 'api-key-1' && value.identity === 'customer-id-1')); - failIdentity = true; - version = 2; + await manager.resolve(auth, config); + fail = true; await time.advance(240000); - const old = await manager.resolve(auth, config); - assert.deepEqual(old, { mode: 'apiKey', secret: 'api-key-1', identity: 'customer-id-1' }); + for (let i = 0; i < 10; i++) assert.equal((await manager.resolve(auth, config)).identity, 'id'); + assert.equal(getSecret.mock.callCount(), 4); assert.deepEqual(failures, ['key_vault']); - failIdentity = false; - await time.advance(5000); - assert.deepEqual(await manager.resolve(auth, config), - { mode: 'apiKey', secret: 'api-key-2', identity: 'customer-id-2' }); + await time.advance(60000); + for (let i = 0; i < 10; i++) await assert.rejects(manager.resolve(auth, config), /unavailable/); assert.equal(getSecret.mock.callCount(), 6); + for (const delay of [30000, 30000, 30000, 30000, 30000]) { + const calls = getSecret.mock.callCount(); + await time.advance(delay - 1); + assert.equal(getSecret.mock.callCount(), calls); + await time.advance(1); + assert.equal(getSecret.mock.callCount(), calls + 2); + } + fail = false; + await time.advance(30000); + assert.equal((await manager.resolve(auth, config)).secret, 'key'); + assert.equal(getSecret.mock.callCount(), 18); } finally { manager.close(); } - assert.equal(time.timerCount, 0); }); -test('MI and final Entra tokens have independent single-flight refresh and warm requests skip both SDKs', async (t) => { +test('invalid or disabled secret values are never published', async (t) => { const time = clock(); + const getSecret = t.mock.method(SecretClient.prototype, 'getSecret', async () => ({ value: 'old' })); + for (const invalid of [{ value: '' }, { value: 'bad', properties: { enabled: false } }, + { value: 'bad', properties: { notBefore: new Date(time.now + 3600000) } }, + { value: 'bad', properties: { expiresOn: new Date(time.now) } }]) { + const manager = new ProviderCredentials({ cacheOptions: time.options, reportFailure() {} }); + getSecret.mock.mockImplementation(async () => invalid); + await assert.rejects(manager.resolve(auth, config), /unavailable/); + manager.close(); + } +}); + +test('SDK credentials are reused, one refresh loop warms both tokens, and reads retain real token expiry', async (t) => { + const time = clock(); + const expiry = time.now + 3600000; const identity = t.mock.method(ManagedIdentityCredential.prototype, 'getToken', async () => ({ - token: 'PRIVATE-ASSERTION', expiresOnTimestamp: time.now + 3600000, + token: 'PRIVATE-ASSERTION', expiresOnTimestamp: expiry, })); const provider = t.mock.method(ClientAssertionCredential.prototype, 'getToken', async function () { await this.getAssertion(); await this.getAssertion(); - return { token: 'PRIVATE-PROVIDER', expiresOnTimestamp: time.now + 3600000 }; - }); - const manager = new ProviderCredentials({ cacheOptions: time.options }); - const config = readConfig({ - EPP_PROVIDER_TENANT_ID: '11111111-1111-1111-1111-111111111111', - EPP_OUTBOUND_CLIENT_ID: '22222222-2222-2222-2222-222222222222', - EPP_OUTBOUND_MI_CLIENT_ID: '33333333-3333-3333-3333-333333333333', - EPP_PROVIDER_SCOPE: 'api://provider/.default', + return { token: 'PRIVATE-PROVIDER', expiresOnTimestamp: expiry, refreshAfterTimestamp: time.now + 10000 }; }); + const manager = new ProviderCredentials({ cacheOptions: time.options, reportFailure() {} }); try { - const results = await Promise.all(Array.from({ length: 20 }, () => manager.resolve({ mode: 'oauth' }, config))); + const results = await Promise.all(Array.from({ length: 20 }, () => manager.resolve({ mode: 'oauth' }, oauth))); assert.ok(results.every((value) => value.accessToken === 'PRIVATE-PROVIDER')); - assert.equal(identity.mock.callCount(), 1); + assert.ok(manager.current instanceof AccessTokenCache); assert.equal(provider.mock.callCount(), 1); - assert.equal(time.timerCount, 2); + assert.equal(identity.mock.callCount(), 3); // One warmup plus SDK assertion callbacks, not network calls. + assert.equal(time.timerCount, 1); assert.equal(JSON.stringify(results[0]), '{"mode":"oauth"}'); - await manager.resolve({ mode: 'oauth' }, config); + await manager.resolve({ mode: 'oauth' }, oauth); assert.equal(provider.mock.callCount(), 1); - await time.advance(3300000); - assert.equal(identity.mock.callCount(), 2); + await time.advance(30000); assert.equal(provider.mock.callCount(), 2); - await manager.resolve({ mode: 'oauth' }, { ...config, providerScope: 'api://second/.default' }); - assert.equal(identity.mock.callCount(), 2); - assert.equal(provider.mock.callCount(), 3); - assert.equal(provider.mock.calls[0].this, provider.mock.calls[2].this); - await manager.resolve({ mode: 'oauth' }, { ...config, outboundClientId: '44444444-4444-4444-4444-444444444444' }); - assert.equal(identity.mock.callCount(), 3); - assert.notEqual(provider.mock.calls[0].this, provider.mock.calls[3].this); - assert.equal(time.timerCount, 2); + assert.equal(provider.mock.calls[0].this, provider.mock.calls[1].this); + await time.advance(3540000); + await assert.rejects(manager.resolve({ mode: 'oauth' }, oauth), /unavailable/); + await time.advance(10000); + await assert.rejects(manager.resolve({ mode: 'oauth' }, oauth), /unavailable/); } finally { manager.close(); } - assert.equal(time.timerCount, 0); }); -test('bad refresh results never replace a valid Key Vault bundle', async (t) => { +test('only the selected cache is created, and new configuration uses a new worker', async (t) => { const time = clock(); - const getSecret = t.mock.method(SecretClient.prototype, 'getSecret', async () => ({ value: 'old-key' })); - const manager = new ProviderCredentials({ cacheOptions: time.options, reportFailure: () => {} }); - const auth = { mode: 'apiKey', keyVaultSecretName: 'key' }; - const config = readConfig({ KEY_VAULT_URL: 'https://unit.vault.azure.net' }); + t.mock.method(ManagedIdentityCredential.prototype, 'getToken', async () => ({ + token: 'assertion', expiresOnTimestamp: time.now + 3600000, + })); + const provider = t.mock.method(ClientAssertionCredential.prototype, 'getToken', async () => ({ + token: 'token', expiresOnTimestamp: time.now + 3600000, + })); + const vault = t.mock.method(SecretClient.prototype, 'getSecret', () => assert.fail('OAuth must not read Key Vault')); + let manager = new ProviderCredentials({ cacheOptions: time.options }); try { - assert.equal((await manager.resolve(auth, config)).secret, 'old-key'); - getSecret.mock.mockImplementation(async () => ({ value: '' })); - await time.advance(240000); - assert.equal((await manager.resolve(auth, config)).secret, 'old-key'); - await time.advance(60000); - await assert.rejects(manager.resolve(auth, config), /unavailable/); + await manager.resolve({ mode: 'oauth' }, oauth); + await manager.resolve({ mode: 'oauth' }, oauth); + assert.equal(provider.mock.callCount(), 1); + assert.ok(manager.current instanceof AccessTokenCache); + manager.close(); + manager = new ProviderCredentials({ cacheOptions: time.options }); + await manager.resolve({ mode: 'oauth' }, { ...oauth, providerScope: 'different-scope' }); + assert.notEqual(provider.mock.calls[0].this, provider.mock.calls[1].this); + manager.close(); + manager = new ProviderCredentials({ cacheOptions: time.options }); + await manager.resolve({ mode: 'oauth' }, { ...oauth, outboundClientId: 'different-app' }); + assert.notEqual(provider.mock.calls[1].this, provider.mock.calls[2].this); + assert.equal(time.timerCount, 1); + assert.equal(vault.mock.callCount(), 0); + manager.close(); + manager = new ProviderCredentials({ cacheOptions: time.options, reportFailure() {} }); + await assert.rejects(manager.resolve({ mode: 'oauth' }, { ...oauth, providerScope: '' }), /unavailable/); + assert.equal(time.timerCount, 0); + await manager.resolve({ mode: 'oauth' }, oauth); + assert.equal(time.timerCount, 1); } finally { manager.close(); } }); -test('incomplete OAuth reconfiguration clears old timers and never reuses old valid tokens', async (t) => { +test('API-key mode never creates an access-token credential and unknown modes start nothing', async (t) => { const time = clock(); - t.mock.method(ManagedIdentityCredential.prototype, 'getToken', async () => ({ - token: 'assertion', expiresOnTimestamp: time.now + 3600000, - })); - t.mock.method(ClientAssertionCredential.prototype, 'getToken', async function () { - await this.getAssertion(); - return { token: 'token', expiresOnTimestamp: time.now + 3600000 }; - }); + t.mock.method(SecretClient.prototype, 'getSecret', async () => ({ value: 'key' })); + const token = t.mock.method(ClientAssertionCredential.prototype, 'getToken', () => assert.fail('Unexpected OAuth')); const manager = new ProviderCredentials({ cacheOptions: time.options }); - const config = readConfig({ EPP_PROVIDER_TENANT_ID: 'tenant', EPP_OUTBOUND_CLIENT_ID: 'app', - EPP_OUTBOUND_MI_CLIENT_ID: 'identity', EPP_PROVIDER_SCOPE: 'scope' }); try { - await manager.resolve({ mode: 'oauth' }, config); - assert.equal(time.timerCount, 2); - await assert.rejects(manager.resolve({ mode: 'oauth' }, { ...config, providerScope: '' }), /unavailable/); - assert.equal(time.timerCount, 0); - await manager.resolve({ mode: 'oauth' }, config); - assert.equal(time.timerCount, 2); + await manager.resolve(auth, config); + assert.ok(manager.current instanceof ApiKeyCache); + assert.equal(token.mock.callCount(), 0); } finally { manager.close(); } + const invalid = new ProviderCredentials({ cacheOptions: time.options, reportFailure() {} }); + await assert.rejects(invalid.resolve({ mode: 'unknown' }, config), /unavailable/); + assert.equal(invalid.current, null); + assert.equal(time.timerCount, 0); + invalid.close(); }); -test('closing a credential manager is terminal, including after a configuration change', async (t) => { +test('shutdown aborts a pending read, prevents late publication, and cannot be reopened', async (t) => { const time = clock(); - const getSecret = t.mock.method(SecretClient.prototype, 'getSecret', async () => ({ value: 'PRIVATE-KEY' })); + let release; + const gate = new Promise((resolve) => { release = resolve; }); + const getSecret = t.mock.method(SecretClient.prototype, 'getSecret', async () => { + await gate; + return { value: 'PRIVATE-LATE-KEY' }; + }); const manager = new ProviderCredentials({ cacheOptions: time.options }); - const auth = { mode: 'apiKey', keyVaultSecretName: 'key' }; - const config = readConfig({ KEY_VAULT_URL: 'https://unit.vault.azure.net' }); - await manager.resolve(auth, config); - manager.close(); + const pending = manager.resolve(auth, config); + const rejected = assert.rejects(pending, /unavailable/); + await flush(); manager.close(); + await rejected; + release(); + await flush(); for (const settings of [config, { ...config, keyVaultUrl: 'https://other.vault.azure.net' }]) { - await assert.rejects(manager.resolve(auth, settings), /provider credential unavailable/); + await assert.rejects(manager.resolve(auth, settings), /unavailable/); } - await assert.rejects(manager.resolve({ mode: 'oauth' }, config), /provider credential unavailable/); - assert.equal(getSecret.mock.callCount(), 1); assert.equal(time.timerCount, 0); + assert.equal(getSecret.mock.callCount(), 2); }); diff --git a/javascript/test/credential-sdk.test.js b/javascript/test/credential-sdk.test.js index 8043ab5..a789b9b 100644 --- a/javascript/test/credential-sdk.test.js +++ b/javascript/test/credential-sdk.test.js @@ -98,7 +98,10 @@ function config() { } test('real SDK sees one initial Entra exchange and none on concurrent or later warm requests', async () => { - const manager = new ProviderCredentials(); + let now = Date.now(); + const manager = new ProviderCredentials({ + cacheOptions: { now: () => now, schedule: () => ({ unref() {} }), cancel() {} }, + }); const settings = config(); try { const tokens = await Promise.all(Array.from({ length: 20 }, () => manager.resolve({ mode: 'oauth' }, settings))); @@ -108,6 +111,11 @@ test('real SDK sees one initial Entra exchange and none on concurrent or later w await manager.resolve({ mode: 'oauth' }, settings); assert.equal(state.tokenCalls, 1); assert.equal(state.miCalls, 1); + now += 60000; + assert.equal((await manager.resolve({ mode: 'oauth' }, settings)).accessToken, 'PRIVATE-TOKEN'); + await manager.refresh(); + assert.equal(state.tokenCalls, 1); + assert.equal(state.miCalls, 1); } finally { manager.close(); } }); @@ -117,7 +125,7 @@ test('the cache-owned acquisition budget aborts actual SDK transport and does no const settings = config(); state.wait = 10000; try { - await assert.rejects(manager.resolve({ mode: 'oauth' }, settings), /^Error: provider OAuth token unavailable$/); + await assert.rejects(manager.resolve({ mode: 'oauth' }, settings), /^Error: provider credential unavailable$/); await new Promise(setImmediate); assert.ok(state.requests.some((request) => !request.managedIdentity && request.aborted)); assert.deepEqual(failures, ['provider_token']); @@ -127,28 +135,30 @@ test('the cache-owned acquisition budget aborts actual SDK transport and does no } finally { manager.close(); } }); -test('changing configuration cancels old work without publishing its token into the replacement cache', async () => { +test('restarting with new configuration cancels old work without publishing into the replacement cache', async () => { const manager = new ProviderCredentials({ reportFailure: () => {} }); const settings = config(); state.wait = 10000; const first = manager.resolve({ mode: 'oauth' }, settings); const firstRejected = assert.rejects(first, /unavailable/); await waitForRequest(false); + manager.close(); state.wait = 5; - const result = await manager.resolve({ mode: 'oauth' }, { ...settings, outboundClientId: crypto.randomUUID() }); + const replacement = new ProviderCredentials({ reportFailure() {} }); + const result = await replacement.resolve({ mode: 'oauth' }, { ...settings, outboundClientId: crypto.randomUUID() }); await firstRejected; try { assert.equal(result.accessToken, 'PRIVATE-TOKEN'); assert.ok(state.requests.some((request) => !request.managedIdentity && request.aborted)); assert.equal(state.tokenCalls, 2); - } finally { manager.close(); } + } finally { manager.close(); replacement.close(); } }); test('the acquisition deadline aborts real managed-identity transport before another refresh starts', async () => { let now = Date.now(); const manager = new ProviderCredentials({ reportFailure: () => {}, - cacheOptions: { now: () => now, random: () => 0, schedule: () => ({ unref() {} }), cancel() {} }, + cacheOptions: { now: () => now, schedule: () => ({ unref() {} }), cancel() {} }, }); const settings = config(); state.miWait = 10000; @@ -158,7 +168,7 @@ test('the acquisition deadline aborts real managed-identity transport before ano assert.equal(state.miCalls, 1); assert.ok(state.requests.every((request) => request.completed && request.aborted)); assert.equal(state.tokenCalls, 0); - now += 10000; + now += 30000; state.miWait = 5; assert.equal((await manager.resolve({ mode: 'oauth' }, settings)).accessToken, 'PRIVATE-TOKEN'); assert.equal(state.miCalls, 2); diff --git a/javascript/test/dispatch.test.js b/javascript/test/dispatch.test.js index 3be56b2..5ed1eba 100644 --- a/javascript/test/dispatch.test.js +++ b/javascript/test/dispatch.test.js @@ -252,6 +252,8 @@ test('missing API-key or OAuth settings and an unsafe final voice URL make zero ['telesign', 'sms', 'provider credential unavailable'], ['sinch', 'voice', 'provider request URL invalid'], ]) { + credentials.close(); + credentials = new ProviderCredentials(); const config = readConfig({ ...settings, EPP_PROVIDER_NAME: providerName }); const result = await dispatchOtp({ ...dispatch, channel }, { config, requestId: 'request-id' }); assert.deepEqual([result.httpStatus, result.body.reason], [502, reason]); @@ -260,12 +262,14 @@ test('missing API-key or OAuth settings and an unsafe final voice URL make zero const calls = getSecret.mock.callCount(); for (const override of [{}, { managedIdentityClientId: '11111111-2222-4333-8444-555555555555' }, { keyVaultUrl: 'https://other-test.vault.azure.net' }]) { + credentials.close(); + credentials = new ProviderCredentials(); const nextConfig = { ...config, ...override }; await dispatchOtp({ ...dispatch, channel: 'voice' }, { config: nextConfig }); assert.equal(getSecret.mock.calls.at(-1).this.vaultUrl, nextConfig.keyVaultUrl); } - assert.equal(getSecret.mock.callCount(), calls + 2); - assert.equal(new Set(getSecret.mock.calls.map((call) => call.this)).size, 4); + assert.equal(getSecret.mock.callCount(), calls + 3); + assert.equal(new Set(getSecret.mock.calls.map((call) => call.this)).size, 5); assert.equal(fetchMock.mock.callCount(), 0); }); @@ -291,10 +295,14 @@ test('Soprano OAuth reuses setup identities and selected scope with private boun const signal = providerToken.mock.calls[0].arguments[1].abortSignal; assert.ok(signal instanceof AbortSignal); assert.ok(identityToken.mock.calls[0].arguments[1].abortSignal instanceof AbortSignal); + credentials.close(); + credentials = new ProviderCredentials(); await resolveProviderCredential({ mode: 'oauth' }, { ...config, providerScope: 'api://another/.default' }); - assert.equal(providerToken.mock.calls[1].this, providerToken.mock.calls[0].this); + assert.notEqual(providerToken.mock.calls[1].this, providerToken.mock.calls[0].this); assert.equal(providerToken.mock.calls[1].arguments[0], 'api://another/.default'); for (const property of ['providerTenantId', 'outboundClientId', 'outboundManagedIdentityClientId']) { + credentials.close(); + credentials = new ProviderCredentials(); await resolveProviderCredential({ mode: 'oauth' }, { ...config, [property]: '44444444-4444-4444-4444-444444444444' }); assert.notEqual(providerToken.mock.calls.at(-1).this, providerToken.mock.calls[0].this); @@ -310,7 +318,7 @@ test('Soprano OAuth reuses setup identities and selected scope with private boun await this.getAssertion(); return { token: 'provider-token', expiresOnTimestamp: Date.now() + 3600000 }; }); - await assert.rejects(resolveProviderCredential({ mode: 'oauth' }, config), /^Error: provider OAuth token unavailable$/); + await assert.rejects(resolveProviderCredential({ mode: 'oauth' }, config), /^Error: provider credential unavailable$/); } } }); @@ -320,6 +328,9 @@ test('Soprano OAuth suppresses SDK diagnostics only during token acquisition', a t.mock.method(AzureLogger, 'log', (...args) => entries.push(args)); let release; const waiting = new Promise(resolve => { release = resolve; }); + t.mock.method(ManagedIdentityCredential.prototype, 'getToken', async () => ({ + token: 'assertion', expiresOnTimestamp: Date.now() + 3600000, + })); t.mock.method(ClientAssertionCredential.prototype, 'getToken', async () => { AzureLogger.log('PRIVATE-TOKEN-AND-ACCOUNT'); await waiting; @@ -334,7 +345,7 @@ test('Soprano OAuth suppresses SDK diagnostics only during token acquisition', a })); AzureLogger.log('unrelated request'); release(); - await assert.rejects(pending, /^Error: provider OAuth token unavailable$/); + await assert.rejects(pending, /^Error: provider credential unavailable$/); AzureLogger.log('after acquisition'); assert.deepEqual(entries, [['unrelated request'], ['after acquisition']]); }); diff --git a/python/README.md b/python/README.md index a1ab622..55dae99 100644 --- a/python/README.md +++ b/python/README.md @@ -90,16 +90,14 @@ six-digit numeric run that is not part of a longer number and repeats the comple ## Source -Worker initialization starts background credential preparation when a provider is configured. -Key Vault bundles, managed-identity assertions and final Entra tokens use separate process-local -caches with daemon refresh timers and parallel daemon secret reads. A caller can stop waiting -without cancelling shared retrieval or starting overlapping reads; the HTTP SDK still uses -connect/read inactivity timeouts, not a total transport deadline. `get_token_info`, when supported -by the installed SDK, preserves early refresh hints; older SDKs retain the pre-expiry refresh target. -`atexit` stops scheduled work, releases waiters and prevents late cache publication. Pending -synchronous reads do not block process exit. Closing a credential manager is terminal. See the -[refresh contract](../docs/CONTRACT.md#credential-caching-and-refresh). Evaluation handling stays -independent; leave the provider unset for local evaluation without background credential acquisition. +Worker initialization selects `ApiKeyCache` or `AccessTokenCache` from the provider manifest's auth mode. +Only the selected cache starts: API keys use Key Vault and `cachetools.TTLCache`; access tokens use +the MI/Entra SDKs without Key Vault. One daemon loop polls every 30 seconds. Configuration changes +require restart. Callers can stop waiting without abandoning shared reads; synchronous SDK I/O uses connect/read +timeouts, not a total transport deadline. `atexit` stops refresh and releases waiters; unfinished +daemon reads cannot publish or block process exit. See the +[refresh contract](../docs/CONTRACT.md#credential-caching-and-refresh). Leave the provider unset for +local evaluation without background credential acquisition. | Source | Purpose | |---|---| @@ -107,10 +105,10 @@ independent; leave the provider unset for local evaluation without background cr | [src/config.py](src/config.py) | Shared deployment settings | | [src/models.py](src/models.py) | Envelope, delivery-context, dispatch and normalized `ParsedResponse` dataclasses | | [src/dispatch.py](src/dispatch.py) | Boundary validation, JWE, provider registry and outcome mapping | -| [src/credentials.py](src/credentials.py), [refreshing_cache.py](src/refreshing_cache.py) | Provider credential bundles, independent token caches and scheduled refresh | +| [src/credentials.py](src/credentials.py) | `ApiKeyCache`, `AccessTokenCache` and their shared refresh coordinator | | [src/request_log.py](src/request_log.py) | Request-scoped [service events and summaries](../docs/CONTRACT.md#application-logs) with explicit ID sources | | [src/providers/](src/providers/) | Adapter manifests and API-specific implementations | -| [src/secrets.py](src/secrets.py) | Key Vault transport; bundle caching belongs to the credential manager | +| [src/secrets.py](src/secrets.py) | Key Vault transport; bundle caching belongs to `ApiKeyCache` | Add and register an adapter without adding provider-specific branches to the shared pipeline. Return `ParsedResponse` from `parse_response` using named fields; the engine reads attributes such as diff --git a/python/requirements.txt b/python/requirements.txt index 41bbff8..a258a23 100644 --- a/python/requirements.txt +++ b/python/requirements.txt @@ -6,3 +6,4 @@ azure-identity>=1.16,<2 azure-keyvault-secrets>=4.8,<5 requests>=2.31,<3 jwcrypto>=1.5,<2 +cachetools>=5.5,<6 diff --git a/python/src/credentials.py b/python/src/credentials.py index 6e4dfeb..fa21628 100644 --- a/python/src/credentials.py +++ b/python/src/credentials.py @@ -2,33 +2,41 @@ import json import logging +import math import threading import time from collections.abc import Callable, Mapping from concurrent.futures import Future, wait from contextvars import ContextVar -from dataclasses import dataclass, field -from typing import Literal, Protocol, TypedDict, TypeVar +from typing import Final, Literal, Protocol, TypedDict, TypeVar from azure.core.credentials import TokenCredential -from azure.identity import ClientAssertionCredential, ManagedIdentityCredential +from azure.identity import AzureAuthorityHosts, ClientAssertionCredential, ManagedIdentityCredential +from cachetools import TTLCache from .config import AppConfig -from .refreshing_cache import ( - ACQUISITION_TIMEOUT_SECONDS, - SECRET_REFRESH_INTERVAL_SECONDS, - SECRET_TTL_SECONDS, - CacheEntry, - CacheOptions, - RefreshingCache, - Token, - token_entry, -) +API_KEY_MODE: Final = "apiKey" +OAUTH_MODE: Final = "oauth" +BUNDLE_KEY = "bundle" +TOKEN_EXCHANGE_SCOPE = "api://AzureADTokenExchange/.default" +CREDENTIAL_ERROR = "provider credential unavailable" +ACQUISITION_TIMEOUT_SECONDS = 2.5 +REFRESH_POLL_SECONDS = 30 +SECRET_TTL_SECONDS = 300 +SECRET_REFRESH_SECONDS = 240 +TOKEN_SKEW_SECONDS = 30 _acquiring = ContextVar("epp_credential_acquisition", default=False) T = TypeVar("T") +class Token(Protocol): + @property + def token(self) -> str: ... + @property + def expires_on(self) -> float: ... + + class SecretReader(Protocol): def resolve(self, secret_name: str) -> str: ... @@ -44,16 +52,9 @@ class OAuthCredential(TypedDict): access_token: str -@dataclass(repr=False) -class _ApiKeyState: - bundle: RefreshingCache[ApiKeyCredential] - - -@dataclass(repr=False) -class _OAuthState: - assertion: RefreshingCache[Token] - credential: TokenCredential - tokens: dict[str, RefreshingCache[Token]] = field(default_factory=dict) +class RefreshOptions(TypedDict, total=False): + clock: Callable[[], float] + wait_timeout: float class _CredentialLogFilter(logging.Filter): @@ -82,153 +83,194 @@ def _private_acquisition(load: Callable[[], T]) -> T: _acquiring.reset(context) -class ProviderCredentials: - def __init__( - self, - secrets: SecretReader, - *, - cache_options: CacheOptions | None = None, - report_failure: Callable[[str], None] = report_refresh_failure, - ) -> None: - self._secrets = secrets - self._options: CacheOptions = cache_options or {} - self._clock = self._options.get("clock", time.time) - self._report_failure = report_failure - self._lock = threading.RLock() - self._key: tuple[str | None, ...] | None = None - self._state: _ApiKeyState | _OAuthState | None = None - self._closed = False +class ApiKeyCache: + stage = "key_vault" - def _cache(self, kind: str, load: Callable[[], CacheEntry[T]]) -> RefreshingCache[T]: - options = self._options.copy() - options["on_failure"] = lambda: self._report_failure(kind) - return RefreshingCache(lambda: _private_acquisition(load), **options) + def __init__(self, secrets: SecretReader, auth: Mapping[str, str], clock=time.time) -> None: + self._secrets, self._auth, self._clock = secrets, dict(auth), clock + self._values: TTLCache[str, ApiKeyCredential] = TTLCache(maxsize=1, ttl=SECRET_TTL_SECONDS, timer=clock) + self._lock = threading.Lock() + self._refresh_at = 0.0 + self._closed = False - def resolve(self, auth: Mapping[str, str], config: AppConfig) -> ApiKeyCredential | OAuthCredential: - mode = auth.get("mode") - scope = config.provider_scope + def get(self) -> ApiKeyCredential | None: with self._lock: - if self._closed: - raise ValueError("provider credential unavailable") - if mode == "oauth" and not all(( - config.provider_tenant_id, scope, - config.outbound_client_id, config.outbound_managed_identity_client_id, - )): - self._clear() - raise ValueError("provider OAuth token unavailable") - if mode not in ("apiKey", "oauth"): - self._clear() - raise ValueError("provider credential unavailable") - key: tuple[str | None, ...] - if mode == "apiKey": - key = ( - mode, config.env.get("KEY_VAULT_URL"), config.env.get("AZURE_CLIENT_ID"), - auth.get("key_vault_secret_name"), auth.get("identity_key_vault_secret_name"), - ) - else: - key = ( - mode, config.provider_tenant_id, - config.outbound_client_id, config.outbound_managed_identity_client_id, - ) - if self._key != key: - self._clear() - if mode == "apiKey": - self._state = self._api_key_state(auth) - else: - self._state = _private_acquisition(lambda: self._oauth_state(config)) - self._key = key - state = self._state - if isinstance(state, _OAuthState) and scope not in state.tokens: - credential = state.credential - state.tokens[scope] = self._cache("provider_token", lambda: self._load_token(credential, scope)) - try: - if isinstance(state, _ApiKeyState): - return state.bundle.get().copy() - if isinstance(state, _OAuthState): - return {"mode": "oauth", "access_token": state.tokens[scope].get().token} - raise ValueError("provider credential unavailable") - except Exception: - reason = "provider OAuth token unavailable" if mode == "oauth" else "provider credential unavailable" - raise ValueError(reason) from None + value = None if self._closed else self._values.get(BUNDLE_KEY) + return value.copy() if value is not None else None - def _start_secret_read(self, name: str) -> Future[str]: - future: Future[str] = Future() + def _read(self, name: str) -> Future[str]: + result: Future[str] = Future() def read() -> None: try: with self._lock: if self._closed: - raise ValueError("provider credential unavailable") - value = _private_acquisition(lambda: self._secrets.resolve(name)) - future.set_result(value) + raise ValueError(CREDENTIAL_ERROR) + result.set_result(_private_acquisition(lambda: self._secrets.resolve(name))) except Exception: - future.set_exception(ValueError("provider credential unavailable")) + result.set_exception(ValueError(CREDENTIAL_ERROR)) - # Executor workers are joined before atexit, even when their parent is a daemon. + # Executor threads are joined at exit; unfinished SDK reads must not block shutdown. threading.Thread(target=read, daemon=True).start() - return future - - def _api_key_state(self, auth: Mapping[str, str]) -> _ApiKeyState: - def load() -> CacheEntry[ApiKeyCredential]: - key_name = auth.get("key_vault_secret_name") - identity_name = auth.get("identity_key_vault_secret_name") - if not key_name: - raise ValueError("provider credential unavailable") - key_future = self._start_secret_read(key_name) - identity_future = self._start_secret_read(identity_name) if identity_name else None - pending = [key_future] - if identity_future is not None: - pending.append(identity_future) - # Keep one refresh in flight until both reads finish, even after a caller stops waiting. - wait(pending) - secret = key_future.result() - identity = identity_future.result() if identity_future is not None else "" - if not isinstance(secret, str) or not secret.strip() or (identity_name and ( - not isinstance(identity, str) or not identity.strip())): - raise ValueError("provider credential unavailable") - now = self._clock() - value: ApiKeyCredential = {"mode": "apiKey", "secret": secret, "identity": identity} - return CacheEntry(value, now + SECRET_TTL_SECONDS, now + SECRET_REFRESH_INTERVAL_SECONDS) - return _ApiKeyState(self._cache("key_vault", load)) - - def _load_token(self, credential: TokenCredential, scope: str) -> CacheEntry[Token]: - get_token_info = getattr(credential, "get_token_info", None) - if callable(get_token_info): - token = get_token_info(scope) - else: - # Older supported SDKs expose only expiry metadata through get_token. - token = credential.get_token(scope, logging_enable=False) - return token_entry(token, self._clock()) - - def _oauth_state(self, config: AppConfig) -> _OAuthState: - identity = ManagedIdentityCredential( - client_id=config.outbound_managed_identity_client_id, - retry_total=0, connection_timeout=ACQUISITION_TIMEOUT_SECONDS, - read_timeout=ACQUISITION_TIMEOUT_SECONDS, logging_enable=False, - ) - assertion = self._cache( - "managed_identity", lambda: self._load_token(identity, "api://AzureADTokenExchange/.default"), - ) - credential = ClientAssertionCredential( - tenant_id=config.provider_tenant_id, client_id=config.outbound_client_id, - func=lambda: assertion.get().token, authority="https://login.microsoftonline.com", - retry_total=0, connection_timeout=ACQUISITION_TIMEOUT_SECONDS, - read_timeout=ACQUISITION_TIMEOUT_SECONDS, logging_enable=False, - ) - return _OAuthState(assertion, credential) - - def _clear(self) -> None: - state = self._state - if isinstance(state, _ApiKeyState): - state.bundle.close() - elif isinstance(state, _OAuthState): - state.assertion.close() - for cache in state.tokens.values(): - cache.close() - self._state = None - self._key = None + return result + + def refresh(self) -> None: + if self.get() is not None and self._clock() < self._refresh_at: + return + key, account = self._auth.get("key_vault_secret_name"), self._auth.get("identity_key_vault_secret_name") + if not key: + raise ValueError(CREDENTIAL_ERROR) + secret = self._read(key) + identity = self._read(account) if account else None + wait([secret, identity] if identity is not None else [secret]) + value, customer = secret.result(), identity.result() if identity is not None else "" + if not value.strip() or (account and not customer.strip()): + raise ValueError(CREDENTIAL_ERROR) + with self._lock: + if self._closed: + raise ValueError(CREDENTIAL_ERROR) + self._values[BUNDLE_KEY] = {"mode": API_KEY_MODE, "secret": value, "identity": customer} + self._refresh_at = self._clock() + SECRET_REFRESH_SECONDS - def close(self) -> None: + def stop(self) -> None: + with self._lock: + self._closed = True + self._values.clear() + + +class AccessTokenCache: + stage = "provider_token" + + def __init__(self, config: AppConfig, clock=time.time) -> None: + if not all((config.provider_tenant_id, config.provider_scope, config.outbound_client_id, + config.outbound_managed_identity_client_id)): + raise ValueError(CREDENTIAL_ERROR) + self._clock, self._scope = clock, config.provider_scope + self._lock = threading.Lock() + self._closed = False + self._token: Token | None = None + self._identity = ManagedIdentityCredential(client_id=config.outbound_managed_identity_client_id, + retry_total=0, logging_enable=False, connection_timeout=ACQUISITION_TIMEOUT_SECONDS, + read_timeout=ACQUISITION_TIMEOUT_SECONDS) + self._credential = ClientAssertionCredential(tenant_id=config.provider_tenant_id, client_id=config.outbound_client_id, + func=lambda: self._load(self._identity, TOKEN_EXCHANGE_SCOPE).token, + authority=AzureAuthorityHosts.AZURE_PUBLIC_CLOUD, retry_total=0, logging_enable=False, + connection_timeout=ACQUISITION_TIMEOUT_SECONDS, read_timeout=ACQUISITION_TIMEOUT_SECONDS) + + def get(self) -> OAuthCredential | None: + with self._lock: + token = self._token + return ({"mode": OAUTH_MODE, "access_token": token.token} + if not self._closed and token and token.expires_on > self._clock() + TOKEN_SKEW_SECONDS else None) + + def _load(self, credential: TokenCredential, scope: str) -> Token: + info = getattr(credential, "get_token_info", None) + token = info(scope) if callable(info) else credential.get_token(scope, logging_enable=False) + if (not isinstance(token.token, str) or not token.token.strip() or not math.isfinite(token.expires_on) + or token.expires_on <= self._clock() + TOKEN_SKEW_SECONDS): + raise ValueError(CREDENTIAL_ERROR) + return token + + def refresh(self) -> None: + with self._lock: + if self._closed: + raise ValueError(CREDENTIAL_ERROR) + self.stage = "managed_identity" + self._load(self._identity, TOKEN_EXCHANGE_SCOPE) + self.stage = "provider_token" + token = self._load(self._credential, self._scope) + with self._lock: + if self._closed: + raise ValueError(CREDENTIAL_ERROR) + self._token = token + + def stop(self) -> None: with self._lock: self._closed = True - self._clear() + self._token = None + + +class ProviderCredentials: + """One selected cache and one refresh loop; configuration changes require a worker restart.""" + + def __init__(self, secrets: SecretReader, *, cache_options: RefreshOptions | None = None, + report_failure: Callable[[str], None] = report_refresh_failure) -> None: + options = cache_options or {} + self._secrets, self._report_failure = secrets, report_failure + self._clock = options.get("clock", time.time) + self._wait_timeout = options.get("wait_timeout", ACQUISITION_TIMEOUT_SECONDS) + self._lock = threading.RLock() + self._stop = threading.Event() + self.cache: ApiKeyCache | AccessTokenCache | None = None + self._pending: Future[ApiKeyCredential | OAuthCredential] | None = None + self._next_attempt = 0.0 + + def resolve(self, auth: Mapping[str, str], config: AppConfig) -> ApiKeyCredential | OAuthCredential: + with self._lock: + if self._stop.is_set(): + raise ValueError(CREDENTIAL_ERROR) + if self.cache is None: + try: + if auth.get("mode") == API_KEY_MODE: + self.cache = ApiKeyCache(self._secrets, auth, self._clock) + elif auth.get("mode") == OAUTH_MODE: + self.cache = _private_acquisition(lambda: AccessTokenCache(config, self._clock)) + else: + raise ValueError(CREDENTIAL_ERROR) + except Exception: + self._report_failure("configuration") + raise ValueError(CREDENTIAL_ERROR) from None + threading.Thread(target=self._loop, daemon=True).start() + value = self.cache.get() + if value is not None: + return value + pending = self.refresh() + try: + return pending.result(timeout=self._wait_timeout) + except Exception: + raise ValueError(CREDENTIAL_ERROR) from None + + def refresh(self) -> Future[ApiKeyCredential | OAuthCredential]: + with self._lock: + if self._stop.is_set() or self.cache is None: + raise ValueError(CREDENTIAL_ERROR) + if self._pending is not None: + return self._pending + if self._next_attempt > self._clock(): + raise ValueError(CREDENTIAL_ERROR) + self._next_attempt = self._clock() + REFRESH_POLL_SECONDS + future: Future[ApiKeyCredential | OAuthCredential] = Future() + self._pending = future + threading.Thread(target=self._run, args=(self.cache, future), daemon=True).start() + return future + + def _run(self, cache: ApiKeyCache | AccessTokenCache, future: Future[ApiKeyCredential | OAuthCredential]) -> None: + value = None + try: + _private_acquisition(cache.refresh) + value = cache.get() + except Exception: + pass # Report only the sanitized failure below. + with self._lock: + self._pending = None + if not self._stop.is_set(): + if value is None: + self._report_failure(cache.stage) + future.set_exception(ValueError(CREDENTIAL_ERROR)) + else: + future.set_result(value) + + def _loop(self) -> None: + while not self._stop.wait(REFRESH_POLL_SECONDS): + try: + self.refresh().result() + except ValueError: + pass # Shared acquisition reports failures; cooldown is intentionally quiet. + + def close(self) -> None: + with self._lock: + self._stop.set() + if self.cache is not None: + self.cache.stop() + if self._pending is not None and not self._pending.done(): + self._pending.set_exception(ValueError(CREDENTIAL_ERROR)) diff --git a/python/src/refreshing_cache.py b/python/src/refreshing_cache.py deleted file mode 100644 index f2a82ac..0000000 --- a/python/src/refreshing_cache.py +++ /dev/null @@ -1,192 +0,0 @@ -from __future__ import annotations - -import math -import random -import threading -import time -from collections.abc import Callable -from concurrent.futures import Future, TimeoutError -from dataclasses import dataclass -from typing import Generic, Protocol, TypedDict, TypeVar - -ACQUISITION_TIMEOUT_SECONDS = 2.5 -SECRET_TTL_SECONDS = 5 * 60 -SECRET_REFRESH_INTERVAL_SECONDS = 4 * 60 -TOKEN_EXPIRY_SKEW_SECONDS = 30 -TOKEN_REFRESH_LEAD_SECONDS = 5 * 60 -MIN_REFRESH_DELAY_SECONDS = 1 -MAX_REFRESH_DELAY_SECONDS = 60 -INITIAL_RETRY_DELAY_SECONDS = 5 -MAX_RETRY_DELAY_SECONDS = 60 -MAX_RETRY_EXPONENT = 4 -RETRY_JITTER_RATIO = 0.2 -MAX_TIMER_DELAY_SECONDS = (2 ** 31 - 1) / 1000 - -T = TypeVar("T") - - -class Token(Protocol): - @property - def token(self) -> str: ... - - @property - def expires_on(self) -> float: ... - - -class ScheduledCall(Protocol): - def cancel(self) -> None: ... - - -class CacheOptions(TypedDict, total=False): - clock: Callable[[], float] - schedule: Callable[[float, Callable[[], None]], ScheduledCall] - jitter: Callable[[], float] - on_failure: Callable[[], None] - wait_timeout: float - - -@dataclass(repr=False) -class CacheEntry(Generic[T]): - value: T - expires_at: float - refresh_at: float - - -def _schedule(delay: float, callback: Callable[[], None]) -> ScheduledCall: - timer = threading.Timer(delay, callback) - timer.daemon = True - timer.start() - return timer - - -class RefreshingCache(Generic[T]): - def __init__( - self, - load: Callable[[], CacheEntry[T]], - *, - clock: Callable[[], float] = time.time, - schedule: Callable[[float, Callable[[], None]], ScheduledCall] = _schedule, - jitter: Callable[[], float] = random.random, - on_failure: Callable[[], None] = lambda: None, - wait_timeout: float = ACQUISITION_TIMEOUT_SECONDS, - ) -> None: - self._load = load - self._clock = clock - self._schedule = schedule - self._jitter = jitter - self._on_failure = on_failure - self._wait_timeout = wait_timeout - self._lock = threading.RLock() - self._entry: CacheEntry[T] | None = None - self._inflight: Future[T] | None = None - self._timer: ScheduledCall | None = None - self._retry_at = 0 - self._failures = 0 - self._closed = False - - def get(self) -> T: - with self._lock: - if self._closed: - raise ValueError("provider credential unavailable") - now = self._clock() - if self._entry and self._entry.expires_at > now: - if self._entry.refresh_at <= now and self._retry_at <= now: - self._begin_refresh() - return self._entry.value - if self._inflight is None and self._retry_at > now: - raise ValueError("provider credential unavailable") - future = self._begin_refresh() - try: - return future.result(timeout=self._wait_timeout) - except TimeoutError: - # A waiter does not cancel the shared refresh needed by other requests. - raise ValueError("provider credential unavailable") from None - - def refresh(self) -> Future[T]: - with self._lock: - if self._closed or self._retry_at > self._clock(): - raise ValueError("provider credential unavailable") - return self._begin_refresh() - - def _begin_refresh(self) -> Future[T]: - if self._inflight is not None: - return self._inflight - if self._timer is not None: - self._timer.cancel() - self._timer = None - future: Future[T] = Future() - self._inflight = future - threading.Thread(target=self._run_refresh, args=(future,), daemon=True).start() - return future - - def _run_refresh(self, future: Future[T]) -> None: - with self._lock: - if self._closed: - self._inflight = None - return - entry: CacheEntry[T] | None = None - try: - entry = self._load() - if (not isinstance(entry, CacheEntry) or not math.isfinite(entry.expires_at) - or entry.expires_at <= self._clock() or not math.isfinite(entry.refresh_at)): - raise ValueError("provider credential unavailable") - except Exception: - entry = None - with self._lock: - if self._closed: - if not future.done(): - future.set_exception(ValueError("provider credential unavailable")) - self._inflight = None - return - if entry is None: - self._failures += 1 - exponent = min(self._failures - 1, MAX_RETRY_EXPONENT) - backoff = min(MAX_RETRY_DELAY_SECONDS, INITIAL_RETRY_DELAY_SECONDS * 2 ** exponent) - self._retry_at = self._clock() + backoff * (1 + self._jitter() * RETRY_JITTER_RATIO) - self._on_failure() - else: - self._entry = entry - self._failures = 0 - self._retry_at = 0 - self._inflight = None - next_refresh = self._retry_at if entry is None else entry.refresh_at - delay = min(MAX_TIMER_DELAY_SECONDS, next_refresh - self._clock()) - self._timer = self._schedule(max(MIN_REFRESH_DELAY_SECONDS, delay), self._scheduled_refresh) - if entry is None: - future.set_exception(ValueError("provider credential unavailable")) - else: - future.set_result(entry.value) - - def _scheduled_refresh(self) -> None: - with self._lock: - self._timer = None - if not self._closed: - self._begin_refresh() - - def close(self) -> None: - with self._lock: - self._closed = True - if self._timer is not None: - self._timer.cancel() - self._timer = None - self._entry = None - if self._inflight is not None and not self._inflight.done(): - self._inflight.set_exception(ValueError("provider credential unavailable")) - - -def token_entry(token: Token, now: float) -> CacheEntry[Token]: - value = getattr(token, "token", None) - expiry = getattr(token, "expires_on", None) - if (not isinstance(value, str) or not value.strip() - or not isinstance(expiry, (int, float)) or isinstance(expiry, bool) - or not math.isfinite(expiry) or expiry <= now + TOKEN_EXPIRY_SKEW_SECONDS): - raise ValueError("provider credential unavailable") - expires_at = expiry - TOKEN_EXPIRY_SKEW_SECONDS - refresh_at = expiry - TOKEN_REFRESH_LEAD_SECONDS - hint = getattr(token, "refresh_on", None) - if isinstance(hint, (int, float)) and not isinstance(hint, bool) and math.isfinite(hint): - refresh_at = min(refresh_at, hint) - if refresh_at <= now: - delay = min(MAX_REFRESH_DELAY_SECONDS, (expires_at - now) / 2) - refresh_at = now + max(MIN_REFRESH_DELAY_SECONDS, delay) - return CacheEntry(token, expires_at, refresh_at) diff --git a/python/src/secrets.py b/python/src/secrets.py index ed97b63..44287a7 100644 --- a/python/src/secrets.py +++ b/python/src/secrets.py @@ -7,23 +7,22 @@ from azure.identity import ManagedIdentityCredential from azure.keyvault.secrets import SecretClient -from .refreshing_cache import ACQUISITION_TIMEOUT_SECONDS +from .credentials import ACQUISITION_TIMEOUT_SECONDS class SecretResolver: def __init__(self, env: Mapping[str, str] | None = None) -> None: self._env = env if env is not None else os.environ self._client: SecretClient | None = None - self._client_key: tuple[str, str | None] | None = None self._lock = Lock() def _get_client(self) -> SecretClient: - vault_url = self._env.get("KEY_VAULT_URL") - client_id = self._env.get("AZURE_CLIENT_ID") - if not vault_url: - raise RuntimeError("KEY_VAULT_URL not set") with self._lock: - if self._client is None or self._client_key != (vault_url, client_id): + if self._client is None: + vault_url = self._env.get("KEY_VAULT_URL") + client_id = self._env.get("AZURE_CLIENT_ID") + if not vault_url: + raise RuntimeError("KEY_VAULT_URL not set") credential = ManagedIdentityCredential( client_id=client_id, logging_enable=False, retry_total=0, connection_timeout=ACQUISITION_TIMEOUT_SECONDS, read_timeout=ACQUISITION_TIMEOUT_SECONDS, @@ -32,11 +31,10 @@ def _get_client(self) -> SecretClient: vault_url=vault_url, credential=credential, retry_total=0, logging_enable=False, connection_timeout=ACQUISITION_TIMEOUT_SECONDS, read_timeout=ACQUISITION_TIMEOUT_SECONDS, ) - self._client_key = (vault_url, client_id) return self._client def resolve(self, secret_name: str | None) -> str: if not secret_name: return "" - # ProviderCredentials caches the complete credential bundle and owns refresh. + # ApiKeyCache publishes the complete credential bundle. return self._get_client().get_secret(secret_name, logging_enable=False).value or "" diff --git a/python/tests/test_credential_cache.py b/python/tests/test_credential_cache.py index 0aeff67..d7a7960 100644 --- a/python/tests/test_credential_cache.py +++ b/python/tests/test_credential_cache.py @@ -3,7 +3,7 @@ import textwrap import time from concurrent.futures import ThreadPoolExecutor -from threading import Event, Lock +from threading import Event from pathlib import Path from types import SimpleNamespace from unittest.mock import Mock @@ -12,10 +12,12 @@ import src.credentials as credentials_module from src.config import read_config -from src.credentials import ProviderCredentials +from src.credentials import ApiKeyCache, AccessTokenCache, ProviderCredentials from src.dispatch import DispatchEngine, ProviderRegistry from src.providers.telesign import TelesignProvider -from src.refreshing_cache import CacheEntry, RefreshingCache, token_entry + +AUTH = {"mode": "apiKey", "key_vault_secret_name": "key", "identity_key_vault_secret_name": "id"} +CONFIG = read_config({"KEY_VAULT_URL": "https://unit.vault.azure.net"}) def wait_until(predicate): @@ -28,262 +30,197 @@ def wait_until(predicate): class Clock: def __init__(self): self.now = 1700000000.0 - self.timers = [] - self.lock = Lock() - - def schedule(self, seconds, callback): - timer = SimpleNamespace(at=self.now + seconds, callback=callback, canceled=False) - timer.cancel = lambda: setattr(timer, "canceled", True) - with self.lock: - self.timers.append(timer) - return timer @property def options(self): - return {"clock": lambda: self.now, "schedule": self.schedule, "jitter": lambda: 0} - - @property - def timer_count(self): - with self.lock: - return sum(not timer.canceled for timer in self.timers) + return {"clock": lambda: self.now} def advance(self, seconds): self.now += seconds - with self.lock: - due = [timer for timer in self.timers if not timer.canceled and timer.at <= self.now] - for timer in due: - timer.canceled = True - for timer in due: - timer.callback() -def test_single_flight_and_scheduled_refresh_do_not_block_valid_cached_values(): +def test_library_cache_shares_parallel_reads_and_serves_a_complete_pair_during_refresh(monkeypatch): clock = Clock() - first_started, release_first, refresh_started, release_refresh = Event(), Event(), Event(), Event() - calls = [] + key_started, id_started, release = Event(), Event(), Event() + version = 1 - def load(): - calls.append(clock.now) - if len(calls) == 1: - first_started.set() - assert release_first.wait(3) - value = "PRIVATE-FIRST" - else: - refresh_started.set() - assert release_refresh.wait(3) - value = "PRIVATE-NEXT" - return CacheEntry(value, clock.now + 300, clock.now + 240) - - cache = RefreshingCache(load, **clock.options) + def read(name): + (key_started if name == "key" else id_started).set() + assert release.wait(3) + return f"PRIVATE-{name}-{version}" + + secrets = Mock(resolve=Mock(side_effect=read)) + oauth = Mock(side_effect=AssertionError("API-key mode must not create an OAuth credential")) + monkeypatch.setattr(credentials_module, "ClientAssertionCredential", oauth) + manager = ProviderCredentials(secrets, cache_options=clock.options) try: with ThreadPoolExecutor(max_workers=10) as pool: - pending = [pool.submit(cache.get) for _ in range(10)] - assert first_started.wait(3) - assert len(calls) == 1 - release_first.set() - assert all(item.result(3) == "PRIVATE-FIRST" for item in pending) + pending = [pool.submit(manager.resolve, AUTH, CONFIG) for _ in range(10)] + assert key_started.wait(3) and id_started.wait(3) + assert secrets.resolve.call_count == 2 + release.set() + assert all(item.result(3)["secret"] == "PRIVATE-key-1" for item in pending) + assert isinstance(manager.cache, ApiKeyCache) + oauth.assert_not_called() + release.clear() + version = 2 clock.advance(240) - assert refresh_started.wait(3) - assert cache.get() == "PRIVATE-FIRST" - release_refresh.set() - wait_until(lambda: cache.get() == "PRIVATE-NEXT") - assert len(calls) == 2 - assert clock.timer_count == 1 - assert "PRIVATE" not in repr(cache) + pending = manager.refresh() + wait_until(lambda: secrets.resolve.call_count == 4) + assert manager.resolve(AUTH, CONFIG)["identity"] == "PRIVATE-id-1" + release.set() + pending.result(timeout=3) + assert manager.resolve(AUTH, CONFIG)["secret"] == "PRIVATE-key-2" + assert manager.resolve(AUTH, CONFIG)["identity"] == "PRIVATE-id-2" + assert "PRIVATE" not in repr(manager) + repr(manager.cache) finally: - release_first.set() - release_refresh.set() - cache.close() - assert clock.timer_count == 0 + release.set() + manager.close() + assert manager.cache.get() is None -def test_failed_refresh_keeps_only_unexpired_entry_and_uses_backoff(): +def test_failed_refresh_never_extends_ttl_or_publishes_a_partial_pair(): clock = Clock() - calls, failures = [], [] fail = False + failures = [] - def load(): - calls.append(clock.now) - if fail: + def read(name): + if fail and name == "id": raise ValueError("PRIVATE-ERROR") - return CacheEntry("value", clock.now + 300, clock.now + 240) + return name - cache = RefreshingCache(load, **clock.options, on_failure=lambda: failures.append("failed")) + secrets = Mock(resolve=Mock(side_effect=read)) + manager = ProviderCredentials(secrets, cache_options=clock.options, report_failure=failures.append) try: - assert cache.get() == "value" + manager.resolve(AUTH, CONFIG) fail = True clock.advance(240) - wait_until(lambda: len(failures) == 1) - assert cache.get() == "value" + with pytest.raises(ValueError, match="unavailable"): + manager.refresh().result(timeout=3) + wait_until(lambda: failures == ["key_vault"]) for _ in range(10): - assert cache.get() == "value" - assert len(calls) == 2 + assert manager.resolve(AUTH, CONFIG)["identity"] == "id" + assert secrets.resolve.call_count == 4 clock.advance(60) + with pytest.raises(ValueError, match="unavailable"): + manager.refresh().result(timeout=3) wait_until(lambda: len(failures) == 2) - with pytest.raises(ValueError, match="provider credential unavailable"): - cache.get() + for _ in range(10): + with pytest.raises(ValueError, match="unavailable"): + manager.resolve(AUTH, CONFIG) + assert secrets.resolve.call_count == 6 + for delay in (30, 30, 30, 30, 30): + calls = secrets.resolve.call_count + clock.advance(delay - 0.5) + with pytest.raises(ValueError, match="unavailable"): + manager.resolve(AUTH, CONFIG) + assert secrets.resolve.call_count == calls + clock.advance(0.5) + with pytest.raises(ValueError, match="unavailable"): + manager.refresh().result(timeout=3) + assert secrets.resolve.call_count == calls + 2 fail = False - clock.advance(10) - wait_until(lambda: len(calls) == 4 and cache._inflight is None) - assert cache.get() == "value" + clock.advance(30) + manager.refresh().result(timeout=3) + assert manager.resolve(AUTH, CONFIG)["secret"] == "key" + assert secrets.resolve.call_count == 18 finally: - cache.close() + manager.close() -def test_timed_out_waiter_does_not_cancel_refresh_and_shutdown_blocks_late_publication(): +def test_waiter_timeouts_and_partial_failures_do_not_start_overlapping_secret_reads(): clock = Clock() started, release = Event(), Event() + calls = [] - def load(): + def read(name): + calls.append(name) + if name == "key": + raise ValueError("PRIVATE-FAILURE") started.set() assert release.wait(3) - return CacheEntry("late", clock.now + 300, clock.now + 240) + return "PRIVATE-IDENTITY" - cache = RefreshingCache(load, **clock.options, wait_timeout=0.02) + manager = ProviderCredentials(Mock(resolve=read), cache_options={**clock.options, "wait_timeout": 0.02}, + report_failure=lambda _: None) try: - with pytest.raises(ValueError, match="unavailable"): - cache.get() + for _ in range(3): + with pytest.raises(ValueError, match="unavailable"): + manager.resolve(AUTH, CONFIG) + clock.advance(60) assert started.is_set() - future = cache.refresh() + assert sorted(calls) == ["id", "key"] + pending = manager.refresh() + manager.close() + with pytest.raises(ValueError, match="unavailable"): + pending.result(timeout=0.1) release.set() - assert future.result(3) == "late" - assert cache.get() == "late" + with pytest.raises(ValueError, match="unavailable"): + manager.resolve(AUTH, CONFIG) + assert manager.cache.get() is None finally: + manager.close() release.set() - cache.close() - with pytest.raises(ValueError, match="unavailable"): - cache.get() - assert clock.timer_count == 0 - - release.clear() - stopped = RefreshingCache(load, **clock.options) - pending = stopped.refresh() - stopped.close() - release.set() - with pytest.raises(ValueError, match="unavailable"): - pending.result(3) - assert clock.timer_count == 0 - - -def test_token_entry_preserves_real_expiry_and_never_spins_on_sdk_cached_return(): - now = 1700000000 - token = SimpleNamespace(token="PRIVATE-TOKEN", expires_on=now + 3600) - entry = token_entry(token, now) - assert (entry.expires_at, entry.refresh_at) == (now + 3570, now + 3300) - repeated = token_entry(token, now + 3300) - assert repeated.expires_at == now + 3570 - assert repeated.refresh_at == now + 3360 - token.refresh_on = now + 600 - assert token_entry(token, now).refresh_at == now + 600 - for invalid in (None, SimpleNamespace(token=""), SimpleNamespace(token=" "), - SimpleNamespace(token="PRIVATE", expires_on=now + 30), - SimpleNamespace(token="PRIVATE", expires_on=float("inf")), - SimpleNamespace(token="PRIVATE", expires_on=True)): - with pytest.raises(ValueError, match="unavailable"): - token_entry(invalid, now) - - -def oauth_config(): - return read_config({ - "EPP_PROVIDER_TENANT_ID": "tenant", "EPP_PROVIDER_SCOPE": "api://provider/.default", - "EPP_OUTBOUND_CLIENT_ID": "app", "EPP_OUTBOUND_MI_CLIENT_ID": "identity", - }) -def test_both_mi_and_provider_token_caches_refresh_independently_and_skip_warm_sdk_calls(monkeypatch): +def test_access_token_cache_uses_sdk_refresh_metadata_and_preserves_original_expiry(monkeypatch): clock = Clock() - identity = Mock(spec=["get_token"], get_token=Mock(side_effect=lambda *args, **kwargs: - SimpleNamespace(token="PRIVATE-ASSERTION", expires_on=clock.now + 3600))) + expiry = clock.now + 3600 + identity = Mock(spec=["get_token", "get_token_info"], get_token_info=Mock( + side_effect=lambda _: SimpleNamespace(token="PRIVATE-ASSERTION", expires_on=expiry, refresh_on=clock.now + 10))) monkeypatch.setattr(credentials_module, "ManagedIdentityCredential", Mock(return_value=identity)) clients = [] def create(**kwargs): - def get_token(*args, **options): + def get_token_info(_): assert kwargs["func"]() == "PRIVATE-ASSERTION" - assert kwargs["func"]() == "PRIVATE-ASSERTION" - return SimpleNamespace(token="PRIVATE-TOKEN", expires_on=clock.now + 3600) - client = Mock(spec=["get_token"], get_token=Mock(side_effect=get_token)) + return SimpleNamespace(token="PRIVATE-TOKEN", expires_on=expiry, refresh_on=clock.now + 10) + client = Mock(spec=["get_token", "get_token_info"], get_token_info=Mock(side_effect=get_token_info)) clients.append(client) return client monkeypatch.setattr(credentials_module, "ClientAssertionCredential", create) - manager = ProviderCredentials(Mock(), cache_options=clock.options) - config = oauth_config() + secrets = Mock(resolve=Mock(side_effect=AssertionError("OAuth must not read Key Vault"))) + manager = ProviderCredentials(secrets, cache_options=clock.options, report_failure=lambda _: None) + config = read_config({"EPP_PROVIDER_TENANT_ID": "tenant", "EPP_OUTBOUND_CLIENT_ID": "app", + "EPP_OUTBOUND_MI_CLIENT_ID": "identity", "EPP_PROVIDER_SCOPE": "scope"}) try: with ThreadPoolExecutor(max_workers=10) as pool: results = list(pool.map(lambda _: manager.resolve({"mode": "oauth"}, config), range(10))) assert all(value["access_token"] == "PRIVATE-TOKEN" for value in results) - assert identity.get_token.call_count == 1 and clients[0].get_token.call_count == 1 - assert manager.resolve({"mode": "oauth"}, config)["access_token"] == "PRIVATE-TOKEN" - assert identity.get_token.call_count == 1 and clients[0].get_token.call_count == 1 - clock.advance(3300) - wait_until(lambda: identity.get_token.call_count == 2 and clients[0].get_token.call_count == 2) - wait_until(lambda: clock.timer_count == 2) - config.provider_scope = "api://second/.default" - manager.resolve({"mode": "oauth"}, config) - assert len(clients) == 1 and clients[0].get_token.call_count == 3 - config.outbound_client_id = "different-app" + assert clients[0].get_token_info.call_count == 1 + assert identity.get_token_info.call_count == 2 # Warmup plus the SDK callback, not HTTP requests. + assert isinstance(manager.cache, AccessTokenCache) + secrets.resolve.assert_not_called() manager.resolve({"mode": "oauth"}, config) - assert len(clients) == 2 - assert clock.timer_count == 2 - finally: - manager.close() - assert clock.timer_count == 0 - - -def test_keyvault_pair_is_parallel_single_flight_and_failed_partial_refresh_retains_old_pair(): - clock = Clock() - key_started, id_started, release = Event(), Event(), Event() - calls = [] - version = 1 - fail_identity = False - - def resolve(name): - calls.append(name) - (key_started if name == "key" else id_started).set() - assert release.wait(3) - if name == "id" and fail_identity: - raise ValueError("PRIVATE-FAILURE") - return f"{name}-{version}" - - failures = [] - manager = ProviderCredentials(Mock(resolve=Mock(side_effect=resolve)), cache_options=clock.options, - report_failure=lambda kind: failures.append(kind)) - config = read_config({"KEY_VAULT_URL": "https://unit.vault.azure.net"}) - auth = {"mode": "apiKey", "key_vault_secret_name": "key", "identity_key_vault_secret_name": "id"} - try: - with ThreadPoolExecutor(max_workers=2) as pool: - first = pool.submit(manager.resolve, auth, config) - second = pool.submit(manager.resolve, auth, config) - assert key_started.wait(3) and id_started.wait(3) - release.set() - assert first.result(3) == second.result(3) == {"mode": "apiKey", "secret": "key-1", "identity": "id-1"} - assert sorted(calls) == ["id", "key"] - version = 2 - fail_identity = True - clock.advance(240) - wait_until(lambda: failures == ["key_vault"]) - assert manager.resolve(auth, config) == {"mode": "apiKey", "secret": "key-1", "identity": "id-1"} - fail_identity = False - clock.advance(5) - wait_until(lambda: manager.resolve(auth, config)["identity"] == "id-2") - assert manager.resolve(auth, config)["secret"] == "key-2" + assert clients[0].get_token_info.call_count == 1 + clock.advance(30) + manager.refresh().result(timeout=3) + assert clients[0].get_token_info.call_count == 2 + identity.get_token.assert_not_called() + clients[0].get_token.assert_not_called() + assert len(clients) == 1 + clock.advance(3540) + with pytest.raises(ValueError, match="unavailable"): + manager.refresh().result(timeout=3) + with pytest.raises(ValueError, match="unavailable"): + manager.resolve({"mode": "oauth"}, config) finally: - release.set() manager.close() -def test_startup_only_prepares_credentials_and_handles_missing_provider_or_failure(): +def test_startup_only_prepares_credentials_and_shutdown_is_terminal(): secrets = Mock(resolve=Mock(return_value="test-key")) engine = DispatchEngine(ProviderRegistry([TelesignProvider()]), secrets, {"EPP_PROVIDER_NAME": "telesign", "EPP_PROVIDER_AUTH_MODE": "apiKey"}) try: engine.start_credential_refresh() - assert secrets.resolve.call_count == 2 engine.start_credential_refresh() assert secrets.resolve.call_count == 2 finally: engine.close() + with pytest.raises(ValueError, match="unavailable"): + engine._credentials.resolve(AUTH, CONFIG) no_provider = DispatchEngine(ProviderRegistry([TelesignProvider()]), secrets, {}) try: no_provider.start_credential_refresh() @@ -298,119 +235,47 @@ def test_startup_only_prepares_credentials_and_handles_missing_provider_or_failu broken.close() -def test_token_info_refresh_hints_apply_to_both_caches_without_legacy_sdk_calls(monkeypatch): - clock = Clock() - - def token_info(value): - return SimpleNamespace(token=value, expires_on=clock.now + 3600, refresh_on=clock.now + 60) - - identity = Mock( - spec=["get_token", "get_token_info"], - get_token_info=Mock(side_effect=lambda scope: token_info("PRIVATE-ASSERTION")), - ) - monkeypatch.setattr(credentials_module, "ManagedIdentityCredential", Mock(return_value=identity)) - clients = [] - - def create(**kwargs): - def get_token_info(scope): - assert kwargs["func"]() == "PRIVATE-ASSERTION" - return token_info("PRIVATE-PROVIDER") - credential = Mock(spec=["get_token", "get_token_info"], get_token_info=Mock(side_effect=get_token_info)) - clients.append(credential) - return credential - - monkeypatch.setattr(credentials_module, "ClientAssertionCredential", create) - manager = ProviderCredentials(Mock(), cache_options=clock.options) - config = oauth_config() - try: - manager.resolve({"mode": "oauth"}, config) - state = manager._state - assert state.assertion._entry.refresh_at == clock.now + 60 - assert state.tokens[config.provider_scope]._entry.refresh_at == clock.now + 60 - assert "PRIVATE" not in repr(state) + repr(state.assertion._entry) - clock.advance(60) - wait_until(lambda: identity.get_token_info.call_count == clients[0].get_token_info.call_count == 2) - wait_until(lambda: clock.timer_count == 2) - identity.get_token.assert_not_called() - clients[0].get_token.assert_not_called() - config.provider_scope = "" - with pytest.raises(ValueError, match="OAuth token unavailable"): - manager.resolve({"mode": "oauth"}, config) - assert clock.timer_count == 0 - config.provider_scope = "api://provider/.default" - assert manager.resolve({"mode": "oauth"}, config)["access_token"] == "PRIVATE-PROVIDER" - assert len(clients) == 2 - finally: - manager.close() - - -def test_close_releases_waiters_immediately_and_never_reopens_the_manager(): - clock = Clock() - started, release = Event(), Event() - - def load(): - started.set() - assert release.wait(3) - return CacheEntry("late", clock.now + 300, clock.now + 240) - - cache = RefreshingCache(load, **clock.options) - future = cache.refresh() - try: - assert started.wait(3) - cache.close() - with pytest.raises(ValueError, match="unavailable"): - future.result(timeout=0.1) - assert clock.timer_count == 0 - finally: - release.set() - cache.close() - - secrets = Mock(resolve=Mock(return_value="PRIVATE-KEY")) - manager = ProviderCredentials(secrets, cache_options=clock.options) - auth = {"mode": "apiKey", "key_vault_secret_name": "key"} - config = read_config({"KEY_VAULT_URL": "https://unit.vault.azure.net"}) - manager.resolve(auth, config) +def test_configuration_changes_require_a_new_worker_and_stopped_cache_cannot_restart(): + secrets = Mock(resolve=Mock(return_value="key")) + manager = ProviderCredentials(secrets) + manager.resolve(AUTH, CONFIG) manager.close() + other = read_config({"KEY_VAULT_URL": "https://other.vault.azure.net"}) + manager = ProviderCredentials(secrets) + manager.resolve(AUTH, other) + assert secrets.resolve.call_count == 4 manager.close() - for settings in (config, read_config({"KEY_VAULT_URL": "https://other.vault.azure.net"})): + manager.close() + for config in [CONFIG, other]: with pytest.raises(ValueError, match="unavailable"): - manager.resolve(auth, settings) - with pytest.raises(ValueError, match="unavailable"): - manager.resolve({"mode": "oauth"}, config) - secrets.resolve.assert_called_once() - assert clock.timer_count == 0 + manager.resolve(AUTH, config) + assert secrets.resolve.call_count == 4 -def test_pending_secret_reads_remain_single_flight_after_waiters_leave(): - clock = Clock() - started, release = Event(), Event() - calls = [] +def test_periodic_refresh_uses_the_selected_cache_and_stops(monkeypatch): + monkeypatch.setattr(credentials_module, "REFRESH_POLL_SECONDS", 0.02) + monkeypatch.setattr(credentials_module, "SECRET_REFRESH_SECONDS", 0.02) + secrets = Mock(resolve=Mock(return_value="key")) + manager = ProviderCredentials(secrets) + manager.resolve(AUTH, CONFIG) + wait_until(lambda: secrets.resolve.call_count >= 4) + manager.close() + count = secrets.resolve.call_count + time.sleep(0.06) + assert secrets.resolve.call_count == count + assert manager.cache.get() is None - def resolve(name): - calls.append(name) - if name == "key": - raise ValueError("PRIVATE-FAILURE") - started.set() - assert release.wait(3) - return "PRIVATE-IDENTITY" - manager = ProviderCredentials( - Mock(resolve=resolve), cache_options={**clock.options, "wait_timeout": 0.02}, - report_failure=lambda kind: None, - ) - auth = {"mode": "apiKey", "key_vault_secret_name": "key", "identity_key_vault_secret_name": "id"} - config = read_config({}) +def test_unknown_auth_mode_does_not_create_a_cache(): + secrets = Mock() + manager = ProviderCredentials(secrets, report_failure=lambda _: None) try: - for _ in range(3): - with pytest.raises(ValueError, match="unavailable"): - manager.resolve(auth, config) - clock.advance(60) - assert started.is_set() - assert sorted(calls) == ["id", "key"] - assert clock.timer_count == 0 + with pytest.raises(ValueError, match="unavailable"): + manager.resolve({"mode": "unknown"}, CONFIG) + assert manager.cache is None + secrets.resolve.assert_not_called() finally: manager.close() - release.set() def test_pending_secret_reads_do_not_block_process_shutdown(): @@ -438,9 +303,7 @@ def resolve(name): manager.close() print("main-finished", flush=True) """) - result = subprocess.run( - [sys.executable, "-c", script], cwd=Path(__file__).resolve().parents[1], - capture_output=True, text=True, timeout=10, - ) + result = subprocess.run([sys.executable, "-c", script], cwd=Path(__file__).resolve().parents[1], + capture_output=True, text=True, timeout=10) assert result.returncode == 0, result.stderr assert result.stdout.splitlines() == ["main-finished", "shutdown-complete"] diff --git a/python/tests/test_credential_sdk.py b/python/tests/test_credential_sdk.py index fcf4581..ed990dc 100644 --- a/python/tests/test_credential_sdk.py +++ b/python/tests/test_credential_sdk.py @@ -6,8 +6,6 @@ import requests import pytest -from azure.identity import ClientAssertionCredential - import src.credentials as credentials_module from src.config import read_config from src.credentials import ProviderCredentials @@ -51,7 +49,10 @@ def managed(*args, **kwargs): monkeypatch.setattr(credentials_module, "ManagedIdentityCredential", Mock( return_value=Mock(spec=["get_token"], get_token=Mock(side_effect=managed)))) secrets = Mock() - manager = ProviderCredentials(secrets) + now = time.time() + manager = ProviderCredentials(secrets, cache_options={ + "clock": lambda: now, + }) config = read_config({ "EPP_PROVIDER_TENANT_ID": "11111111-1111-1111-1111-111111111111", "EPP_OUTBOUND_CLIENT_ID": "22222222-2222-2222-2222-222222222222", @@ -63,18 +64,14 @@ def managed(*args, **kwargs): results = list(pool.map(lambda _: manager.resolve({"mode": "oauth"}, config), range(10))) assert all(value["access_token"] == "PRIVATE-PROVIDER" for value in results) assert len(token_endpoint_calls) == 1 - assert len(mi_calls) == 1 + assert len(mi_calls) == 2 assert manager.resolve({"mode": "oauth"}, config)["access_token"] == "PRIVATE-PROVIDER" - assert len(token_endpoint_calls) == len(mi_calls) == 1 - entry = manager._state.tokens[config.provider_scope]._entry - if refresh_in is not None and hasattr(ClientAssertionCredential, "get_token_info"): - info = manager._state.credential.get_token_info(config.provider_scope) - assert info.refresh_on is not None - assert entry.refresh_at == info.refresh_on - assert entry.refresh_at < entry.expires_at - 3000 - else: - assert entry.refresh_at == entry.value.expires_on - 300 - assert len(token_endpoint_calls) == len(mi_calls) == 1 + assert len(token_endpoint_calls) == 1 and len(mi_calls) == 2 + assert len(token_endpoint_calls) == 1 and len(mi_calls) == 2 + now += 60 + assert manager.resolve({"mode": "oauth"}, config)["access_token"] == "PRIVATE-PROVIDER" + manager.refresh().result(timeout=3) + assert len(token_endpoint_calls) == 1 secrets.resolve.assert_not_called() finally: manager.close() diff --git a/python/tests/test_engine.py b/python/tests/test_engine.py index 25b258d..58b8a71 100644 --- a/python/tests/test_engine.py +++ b/python/tests/test_engine.py @@ -92,6 +92,8 @@ def get_token(*args, **options): connection_timeout=2.5, read_timeout=2.5, logging_enable=False) managed.get_token.assert_called_with("api://AzureADTokenExchange/.default", logging_enable=False) engine.env["EPP_OUTBOUND_CLIENT_ID"] = "second-calling-app" + engine.close() + engine = DispatchEngine(engine.registry, engine.secrets, engine.env) assert engine.dispatch(_request(), "request")[0] == 200 assert len(clients) == 2 dispatch_module.requests.request.reset_mock() @@ -113,6 +115,7 @@ def get_token(*args, **options): assert "private" not in json.dumps(body) engine.secrets.resolve.assert_not_called() dispatch_module.requests.request.assert_not_called() + engine.close() def test_soprano_oauth_sdk_logs_stay_private_without_muting_other_requests(engine, monkeypatch, caplog): @@ -131,7 +134,8 @@ def fail(*args, **kwargs): logger.warning("PRIVATE-ACCOUNT-ERROR") raise RuntimeError("PRIVATE-TOKEN-EXCEPTION") - monkeypatch.setattr(credentials_module, "ManagedIdentityCredential", Mock()) + monkeypatch.setattr(credentials_module, "ManagedIdentityCredential", Mock(return_value=Mock( + spec=["get_token"], get_token=Mock(return_value=SimpleNamespace(token="assertion", expires_on=time.time() + 3600))))) monkeypatch.setattr(credentials_module, "ClientAssertionCredential", Mock(return_value=Mock(spec=["get_token"], get_token=fail))) try: