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..bada50b 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 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 +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,49 @@ 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. + +The selected provider's manifest determines which of two concrete cache classes is created: + +| Authentication mode | Cache | Acquisition | +|---|---|---| +| `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. + --- ## 5. Required behaviors @@ -477,10 +531,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. +- 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 local keys and mocked external services; they do not send SMS and **do not test Easy Auth or platform @@ -501,6 +558,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..0bbf2ac 100644 --- a/dotnet/README.md +++ b/dotnet/README.md @@ -90,16 +90,26 @@ six-digit numeric run that is not part of a longer number and repeats the comple ## Source +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 | |---|---| | [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) | `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) | Cached Key Vault access via managed identity | +| [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/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..9eaf53e 100644 --- a/dotnet/Src/DispatchEngine.cs +++ b/dotnet/Src/DispatchEngine.cs @@ -210,50 +210,58 @@ 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, createManagedIdentity, createOAuthCredential, log, clock); } 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 = ProviderCredentials.AcquisitionTimeout; options.Diagnostics.IsLoggingEnabled = false; options.Diagnostics.IsLoggingContentEnabled = false; 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..374cfa0 --- /dev/null +++ b/dotnet/Src/ProviderCredentials.cs @@ -0,0 +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 Func _createIdentity; + private readonly Func>, TokenCredential> _createCredential; + private readonly ILogger _log; + private readonly TimeProvider _clock; + private ICredentialCache? _cache; + private Task? _pending; + private CancellationTokenSource? _acquisition; + private ITimer? _timer; + private DateTimeOffset _nextAttempt; + private bool _disposed; + + internal ProviderCredentials(ISecretResolver secrets, Func createIdentity, + Func>, TokenCredential> createCredential, + ILogger? log = null, TimeProvider? clock = null) + { + _secrets = secrets; + _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", + }; + _log.Log(LogLevel.Warning, new EventId(0, eventName), record, null, + static (state, _) => JsonSerializer.Serialize(state)); + } + internal Task ResolveAsync(AuthConfig auth, AppConfig config, CancellationToken cancellation = default) + { + lock (_gate) + { + ObjectDisposedException.ThrowIf(_disposed, this); + if (_cache is null) + { + try + { + _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; + } + } + private void Tick() + { + lock (_gate) { if (!_disposed) _ = ObserveAsync(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 async Task RunAsync(ICredentialCache cache, TaskCompletionSource completion, CancellationTokenSource cancellation) + { + 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; + _timer?.Dispose(); + _acquisition?.Cancel(); + _cache?.Dispose(); + } + } +} diff --git a/dotnet/Src/SecretResolver.cs b/dotnet/Src/SecretResolver.cs index 3a97f65..1b51652 100644 --- a/dotnet/Src/SecretResolver.cs +++ b/dotnet/Src/SecretResolver.cs @@ -1,40 +1,48 @@ -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). ApiKeyCache publishes 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; public SecretResolver(IEnv? env = null) { - var environment = env ?? new ProcessEnv(); - _client = new Lazy(() => + _env = env ?? new ProcessEnv(); + } + + private SecretClient GetClient() + { + 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) 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 = 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 = ProviderCredentials.AcquisitionTimeout; + options.Diagnostics.IsLoggingEnabled = false; + options.Diagnostics.IsLoggingContentEnabled = false; + _client = new SecretClient(new Uri(url), credential, options); + 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..f49c4ec --- /dev/null +++ b/dotnet/tests/CredentialCacheTests.cs @@ -0,0 +1,423 @@ +using System.Text.Json; +using Azure.Core; +using Microsoft.Extensions.Logging; +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 manager = new ProviderCredentials(new Secrets(async (_, _) => + { + Interlocked.Increment(ref calls); + await release.Task; + 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.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 Get()).Secret); + release.SetResult(); + await Until(() => clock.TimerCount == 1); + Assert.Equal("value", (await Get()).Secret); + manager.Dispose(); + Assert.Equal(0, clock.TimerCount); + } + + [Fact] + public async Task RefreshFailuresDoNotExtendExpiryAndUseFixedRetryCadence() + { + var clock = new ManualClock(); + var fail = false; + var calls = 0; + var log = new CredentialLogger(); + using var manager = new ProviderCredentials(new Secrets((_, _) => + { + calls++; + if (fail) throw new InvalidOperationException("PRIVATE-ERROR"); + 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.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, 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(30)); + Assert.Equal("first", (await Get()).Secret); + Assert.Equal(9, calls); + } + + [Fact] + public async Task CancellingAWaiterDoesNotCancelTheSharedRefresh() + { + var clock = new ManualClock(); + var release = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + CancellationToken observed = default; + using var manager = new ProviderCredentials(new Secrets(async (_, cancellation) => + { + observed = cancellation; + await release.Task; + return "ready"; + }), _ => throw new Exception(), (_, _, _) => throw new Exception(), clock: clock); + using var waiter = new CancellationTokenSource(); + 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).Secret); + } + + [Fact] + public async Task CacheOwnedDeadlineAndShutdownPreventLatePublication() + { + var clock = new ManualClock(); + var release = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var log = new CredentialLogger(); + CancellationToken observed = default; + using var manager = new ProviderCredentials(new Secrets(async (_, cancellation) => + { + observed = cancellation; + await release.Task; + 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.Single(log.Entries); + manager.Dispose(); + release.SetResult(); + await Assert.ThrowsAsync(() => manager.ResolveAsync(new("apiKey", "key"), new AppConfig())); + 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, _ => 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(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 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 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(3, identityCalls); + Assert.Equal(1, providerCalls); + await manager.ResolveAsync(new("oauth"), config); + Assert.Equal(1, 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); + 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 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())); + } + + [Fact] + public async Task DisposingManagerIsTerminalAndCancelsPendingAcquisition() + { + var clock = new ManualClock(); + var release = new TaskCompletionSource(TaskCreationOptions.RunContinuationsAsynchronously); + var calls = 0; + CancellationToken acquisition = default; + using var manager = new ProviderCredentials(new Secrets((_, cancellation) => + { + calls++; + acquisition = cancellation; + return release.Task; + }), _ => 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())); + await Assert.ThrowsAsync(() => manager.ResolveAsync(new("oauth"), Config())); + Assert.Equal(1, calls); + Assert.Equal(0, clock.TimerCount); + } + + [Fact] + 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 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 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(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] + 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")), + _ => 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, + 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 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(); + 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..1f7acbf 100644 --- a/dotnet/tests/EngineTests.cs +++ b/dotnet/tests/EngineTests.cs @@ -35,19 +35,47 @@ private static void ConfigureSoprano(HandlerRig rig) } [Fact] - public async Task SopranoOAuthUsesSetupIdentitiesScopeAndOneBoundedExchange() + 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")); + } + + [Theory] + [InlineData("api://provider/.default")] + [InlineData("api://second/.default")] + public async Task SopranoOAuthUsesSetupIdentitiesScopeAndOneBoundedExchange(string scope) { 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,13 +84,13 @@ 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)); }); }); ConfigureSoprano(rig); + rig.Env["EPP_PROVIDER_SCOPE"] = scope; AssertAccepted(await rig.Invoke("evaluation")); Assert.Empty(applications); foreach (var channel in new[] { "sms", "voice" }) @@ -90,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://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"); @@ -130,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) { @@ -138,19 +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.Http.Respond = _ => Task.FromResult(Json(401, "{\"status\":\"REJECTED\"}")); - AssertFailure(rig, await rig.Invoke(), 401); - Assert.Equal((1, 0), (rig.Http.Calls, rig.Secrets.Calls)); + rig.Engine.Dispose(); + 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 @@ -800,6 +833,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 +847,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 +878,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 +887,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..81e6c15 100644 --- a/javascript/README.md +++ b/javascript/README.md @@ -99,12 +99,21 @@ retries. The shared contract defines validation, HTTP outcomes and privacy-safe ## Source and extension points +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 | |---|---| | [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) | `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 20db06d..cd0c455 100644 --- a/javascript/package-lock.json +++ b/javascript/package-lock.json @@ -8,11 +8,13 @@ "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", "@azure/logger": "1.3.0", - "jose": "^5.9.6" + "jose": "^5.9.6", + "lru-cache": "^11.1.0" } }, "node_modules/@azure-rest/core-client": { @@ -522,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 88a5420..0522d49 100644 --- a/javascript/package.json +++ b/javascript/package.json @@ -8,10 +8,12 @@ "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", "@azure/logger": "1.3.0", - "jose": "^5.9.6" + "jose": "^5.9.6", + "lru-cache": "^11.1.0" } } 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..29b982c --- /dev/null +++ b/javascript/src/functions/credentials.js @@ -0,0 +1,242 @@ +// Copyright (c) Microsoft Corporation. All rights reserved. +'use strict'; + +const { AsyncLocalStorage } = require('node:async_hooks'); +const { inspect } = require('node:util'); +const { LRUCache } = require('lru-cache'); +const { createDefaultHttpClient } = require('@azure/core-rest-pipeline'); +const { AzureAuthorityHosts, ClientAssertionCredential, ManagedIdentityCredential } = require('@azure/identity'); +const { SecretClient } = require('@azure/keyvault-secrets'); +const { AzureLogger } = require('@azure/logger'); + +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; + +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(); + return { + async sendRequest(request) { + const signal = acquisition.getStore(); + 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); + if (signal.aborted || existing?.aborted) controller.abort(); + request.abortSignal = controller.signal; + try { + controller.signal.throwIfAborted(); + return await client.sendRequest(request); + } finally { + signal.removeEventListener('abort', abort); + existing?.removeEventListener('abort', abort); + } + }, + }; +} + +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; +} + +/** + * @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 + */ + +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; + } + 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]'; } +} + +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?: RefreshOptions, reportFailure?: (kind: string) => void}} [options] */ + constructor({ cacheOptions = {}, reportFailure = reportRefreshFailure } = {}) { + this.now = cacheOptions.now || Date.now; + this.schedule = cacheOptions.schedule || setInterval; + this.cancel = cacheOptions.cancel || clearInterval; + this.reportFailure = reportFailure; + /** @type {ApiKeyCache | AccessTokenCache | null} */ + this.current = 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; + } + /** @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(); } + } + const cached = this.current.get(); + if (cached) return cached; + await this.refresh(); + const value = this.current.get(); + if (!value) throw unavailable(); + return value; + } + 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; + } + 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; + 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 = { ApiKeyCache, AccessTokenCache, 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/test/credential-cache.test.js b/javascript/test/credential-cache.test.js new file mode 100644 index 0000000..fbf7786 --- /dev/null +++ b/javascript/test/credential-cache.test.js @@ -0,0 +1,227 @@ +'use strict'; + +const { test } = require('node:test'); +const assert = require('node:assert/strict'); +const { inspect } = require('node:util'); +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 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, + schedule: (callback, delay) => { + const timer = { at: time + delay, period: 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.has(timer)) { + timer.at = time + timer.period; + timer.callback(); + } + } + await flush(); + }, + }; +} + +test('library-backed bundle shares concurrent reads, serves during refresh, and publishes pairs atomically', async (t) => { + const time = clock(); + 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 pending = Array.from({ length: 20 }, () => manager.resolve(auth, config)); + await flush(); + 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(manager) + JSON.stringify(manager), /key-1|id-1/); + gate = new Promise((resolve) => { release = resolve; }); + version = 2; + await time.advance(240000); + assert.deepEqual(await manager.resolve(auth, config), { mode: 'apiKey', secret: 'key-1', identity: 'id-1' }); + release(); + await flush(); + 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('partial refresh failure retains the old pair only until hard expiry, with fixed retry cadence', async (t) => { + const time = clock(); + let fail = false; + const getSecret = t.mock.method(SecretClient.prototype, 'getSecret', async (name) => { + 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) }); + try { + await manager.resolve(auth, config); + fail = true; + await time.advance(240000); + 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']); + 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(); } +}); + +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: expiry, + })); + const provider = t.mock.method(ClientAssertionCredential.prototype, 'getToken', async function () { + await this.getAssertion(); + await this.getAssertion(); + 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' }, oauth))); + assert.ok(results.every((value) => value.accessToken === 'PRIVATE-PROVIDER')); + assert.ok(manager.current instanceof AccessTokenCache); + assert.equal(provider.mock.callCount(), 1); + 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' }, oauth); + assert.equal(provider.mock.callCount(), 1); + await time.advance(30000); + assert.equal(provider.mock.callCount(), 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(); } +}); + +test('only the selected cache is created, and new configuration uses a new worker', async (t) => { + const time = clock(); + 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 { + 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('API-key mode never creates an access-token credential and unknown modes start nothing', async (t) => { + const time = clock(); + 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 }); + try { + 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('shutdown aborts a pending read, prevents late publication, and cannot be reopened', async (t) => { + const time = clock(); + 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 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), /unavailable/); + } + 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 new file mode 100644 index 0000000..a789b9b --- /dev/null +++ b/javascript/test/credential-sdk.test.js @@ -0,0 +1,267 @@ +'use strict'; + +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 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/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(() => { + 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); + 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: 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: crypto.randomUUID(), + EPP_PROVIDER_SCOPE: 'api://provider/.default', + }); +} + +test('real SDK sees one initial Entra exchange and none on concurrent or later warm requests', async () => { + 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))); + 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); + 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(); } +}); + +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 credential unavailable$/); + await new Promise(setImmediate); + 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/); + assert.equal(state.tokenCalls, 1); + } finally { manager.close(); } +}); + +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 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(); 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, 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 += 30000; + 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 dfa0a47..5ed1eba 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, beforeEach, afterEach } = require('node:test'); const assert = require('node:assert/strict'); const { ClientAssertionCredential, ManagedIdentityCredential } = require('@azure/identity'); const { SecretClient } = require('@azure/keyvault-secrets'); @@ -9,10 +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, } = require('../src/functions/dispatch'); +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, @@ -245,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]); @@ -253,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, 3); + 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); }); @@ -283,11 +294,15 @@ 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); + 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); @@ -296,12 +311,14 @@ 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; + credentials.close(); + credentials = new ProviderCredentials(); method.mock.mockImplementation(async () => invalid); if (stage === 'assertion') providerToken.mock.mockImplementation(async function () { 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$/); } } }); @@ -311,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; @@ -325,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/javascript/test/sendotp.test.js b/javascript/test/sendotp.test.js index 14454af..5836d9c 100644 --- a/javascript/test/sendotp.test.js +++ b/javascript/test/sendotp.test.js @@ -9,15 +9,20 @@ const { ClientAssertionCredential, ManagedIdentityCredential } = require('@azure const { SecretClient } = require('@azure/keyvault-secrets'); const fixtures = require('../../tests/fixtures/contract.json'); 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. 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); }); @@ -38,7 +43,11 @@ let logs; let warnings; let records; let getToken; +let credentials; beforeEach(() => { + 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', @@ -61,6 +70,7 @@ beforeEach(() => { text: async () => JSON.stringify({ status: 'ENROUTE', id: 'provider-reference-id', description: 'PRIVATE-STATUS' }) })); }); afterEach(() => { + credentials.close(); mock.restoreAll(); for (const [key, value] of Object.entries(savedEnv)) { if (value === undefined) delete process.env[key]; @@ -127,6 +137,32 @@ 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(); + 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 () => { + 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..55dae99 100644 --- a/python/README.md +++ b/python/README.md @@ -90,15 +90,25 @@ six-digit numeric run that is not part of a longer number and repeats the comple ## Source +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 | |---|---| | [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) | `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) | Cached Key Vault access via managed identity | +| [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/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/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 new file mode 100644 index 0000000..fa21628 --- /dev/null +++ b/python/src/credentials.py @@ -0,0 +1,276 @@ +from __future__ import annotations + +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 typing import Final, Literal, Protocol, TypedDict, TypeVar + +from azure.core.credentials import TokenCredential +from azure.identity import AzureAuthorityHosts, ClientAssertionCredential, ManagedIdentityCredential +from cachetools import TTLCache + +from .config import AppConfig + +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: ... + + +class ApiKeyCredential(TypedDict): + mode: Literal["apiKey"] + secret: str + identity: str + + +class OAuthCredential(TypedDict): + mode: Literal["oauth"] + access_token: str + + +class RefreshOptions(TypedDict, total=False): + clock: Callable[[], float] + wait_timeout: float + + +class _CredentialLogFilter(logging.Filter): + 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: str) -> None: + logging.warning("%s", json.dumps({"logType": "service", "eventName": "credential_refresh_failed", + "cacheKind": kind, "failureReason": "credential_unavailable"})) + + +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: + if _log_filter not in handler.filters: + handler.addFilter(_log_filter) + context = _acquiring.set(True) + try: + return load() + finally: + _acquiring.reset(context) + + +class ApiKeyCache: + stage = "key_vault" + + 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 get(self) -> ApiKeyCredential | None: + with self._lock: + value = None if self._closed else self._values.get(BUNDLE_KEY) + return value.copy() if value is not None else None + + def _read(self, name: str) -> Future[str]: + result: Future[str] = Future() + + def read() -> None: + try: + with self._lock: + if self._closed: + raise ValueError(CREDENTIAL_ERROR) + result.set_result(_private_acquisition(lambda: self._secrets.resolve(name))) + except Exception: + result.set_exception(ValueError(CREDENTIAL_ERROR)) + + # Executor threads are joined at exit; unfinished SDK reads must not block shutdown. + threading.Thread(target=read, daemon=True).start() + 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 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._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/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/secrets.py b/python/src/secrets.py index 036cede..44287a7 100644 --- a/python/src/secrets.py +++ b/python/src/secrets.py @@ -1,40 +1,40 @@ +from __future__ import annotations + import os -import time +from collections.abc import Mapping +from threading import Lock from azure.identity import ManagedIdentityCredential from azure.keyvault.secrets import SecretClient -CACHE_TTL_SECONDS = 5 * 60 +from .credentials import ACQUISITION_TIMEOUT_SECONDS class SecretResolver: - def __init__(self): - self._client = None - self._cache = {} + 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._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 + def _get_client(self) -> SecretClient: + with self._lock: + 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, + ) + 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, + ) + return self._client - def resolve(self, secret_name): + def resolve(self, secret_name: str | None) -> str: 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 + # 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 new file mode 100644 index 0000000..d7a7960 --- /dev/null +++ b/python/tests/test_credential_cache.py @@ -0,0 +1,309 @@ +import subprocess +import sys +import textwrap +import time +from concurrent.futures import ThreadPoolExecutor +from threading import Event +from pathlib import Path +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 ApiKeyCache, AccessTokenCache, ProviderCredentials +from src.dispatch import DispatchEngine, ProviderRegistry +from src.providers.telesign import TelesignProvider + +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): + 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 + + @property + def options(self): + return {"clock": lambda: self.now} + + def advance(self, seconds): + self.now += seconds + + +def test_library_cache_shares_parallel_reads_and_serves_a_complete_pair_during_refresh(monkeypatch): + clock = Clock() + key_started, id_started, release = Event(), Event(), Event() + version = 1 + + 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(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) + 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.set() + manager.close() + assert manager.cache.get() is None + + +def test_failed_refresh_never_extends_ttl_or_publishes_a_partial_pair(): + clock = Clock() + fail = False + failures = [] + + def read(name): + if fail and name == "id": + raise ValueError("PRIVATE-ERROR") + return name + + secrets = Mock(resolve=Mock(side_effect=read)) + manager = ProviderCredentials(secrets, cache_options=clock.options, report_failure=failures.append) + try: + manager.resolve(AUTH, CONFIG) + fail = True + clock.advance(240) + with pytest.raises(ValueError, match="unavailable"): + manager.refresh().result(timeout=3) + wait_until(lambda: failures == ["key_vault"]) + for _ in range(10): + 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) + 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(30) + manager.refresh().result(timeout=3) + assert manager.resolve(AUTH, CONFIG)["secret"] == "key" + assert secrets.resolve.call_count == 18 + finally: + manager.close() + + +def test_waiter_timeouts_and_partial_failures_do_not_start_overlapping_secret_reads(): + clock = Clock() + started, release = Event(), Event() + calls = [] + + def read(name): + calls.append(name) + if name == "key": + raise ValueError("PRIVATE-FAILURE") + started.set() + assert release.wait(3) + return "PRIVATE-IDENTITY" + + manager = ProviderCredentials(Mock(resolve=read), cache_options={**clock.options, "wait_timeout": 0.02}, + 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"] + pending = manager.refresh() + manager.close() + with pytest.raises(ValueError, match="unavailable"): + pending.result(timeout=0.1) + release.set() + with pytest.raises(ValueError, match="unavailable"): + manager.resolve(AUTH, CONFIG) + assert manager.cache.get() is None + finally: + manager.close() + release.set() + + +def test_access_token_cache_uses_sdk_refresh_metadata_and_preserves_original_expiry(monkeypatch): + clock = Clock() + 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_info(_): + assert kwargs["func"]() == "PRIVATE-ASSERTION" + 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) + 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 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 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: + manager.close() + + +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() + 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() + 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() + + +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() + manager.close() + for config in [CONFIG, other]: + with pytest.raises(ValueError, match="unavailable"): + manager.resolve(AUTH, config) + assert secrets.resolve.call_count == 4 + + +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 test_unknown_auth_mode_does_not_create_a_cache(): + secrets = Mock() + manager = ProviderCredentials(secrets, report_failure=lambda _: None) + try: + with pytest.raises(ValueError, match="unavailable"): + manager.resolve({"mode": "unknown"}, CONFIG) + assert manager.cache is None + secrets.resolve.assert_not_called() + finally: + manager.close() + + +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 new file mode 100644 index 0000000..ed990dc --- /dev/null +++ b/python/tests/test_credential_sdk.py @@ -0,0 +1,77 @@ +import json +import time +from concurrent.futures import ThreadPoolExecutor +from types import SimpleNamespace +from unittest.mock import Mock + +import requests +import pytest +import src.credentials as credentials_module +from src.config import read_config +from src.credentials import ProviderCredentials + + +@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 = [] + + 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) + 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", + "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(spec=["get_token"], get_token=Mock(side_effect=managed)))) + secrets = Mock() + 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", + "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) == 2 + assert manager.resolve({"mode": "oauth"}, config)["access_token"] == "PRIVATE-PROVIDER" + 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 936c251..58b8a71 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): @@ -56,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 = [] @@ -71,12 +73,12 @@ 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 - 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 @@ -90,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() @@ -102,11 +106,16 @@ def get_token(*args, **options): 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() dispatch_module.requests.request.assert_not_called() + engine.close() def test_soprano_oauth_sdk_logs_stay_private_without_muting_other_requests(engine, monkeypatch, caplog): @@ -125,8 +134,10 @@ 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(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: 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):