From 36ebb51c1076fc9905397958d280b96e01da244a Mon Sep 17 00:00:00 2001 From: Nisheet Jain Date: Thu, 1 Oct 2026 14:34:59 -0700 Subject: [PATCH 1/3] Simplify Python provider flow Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 0769c5c6-97a7-4645-ab83-5cda60c73b16 --- docs/CONTRACT.md | 66 +-- python/README.md | 33 +- python/function_app.py | 294 ++++++++--- python/src/credentials.py | 20 +- python/src/dispatch.py | 418 --------------- python/src/jwe.py | 68 +++ python/src/models.py | 113 ++-- python/src/otp_log.py | 118 +++++ python/src/provider.py | 212 ++++++++ python/src/providers/infobip.py | 61 ++- python/src/providers/sinch.py | 80 +-- python/src/providers/soprano.py | 63 ++- python/src/providers/telesign.py | 74 +-- python/src/request_log.py | 210 -------- python/tests/test_contract.py | 285 +++++----- python/tests/test_credential_cache.py | 49 +- python/tests/test_credential_sdk.py | 4 +- python/tests/test_engine.py | 344 ++++-------- python/tests/test_function_app.py | 733 +++++++++----------------- 19 files changed, 1415 insertions(+), 1830 deletions(-) delete mode 100644 python/src/dispatch.py create mode 100644 python/src/jwe.py create mode 100644 python/src/otp_log.py create mode 100644 python/src/provider.py delete mode 100644 python/src/request_log.py diff --git a/docs/CONTRACT.md b/docs/CONTRACT.md index 57ec15c..f16c3b2 100644 --- a/docs/CONTRACT.md +++ b/docs/CONTRACT.md @@ -224,9 +224,10 @@ success-looking status. Explicit `Block`/`StepUp` outcomes remain non-success re ## 3. Provider adapter contract -Each provider is one unit exposing three things: +Each provider is one unit that owns authentication requirements, request construction and response +mapping. JavaScript exposes these through: -- **`manifest`**: protocol facts only: +- **`manifest`**: protocol facts: - `id`: provider id selected by `EPP_PROVIDER_NAME`; its complete request URL is `EPP_PROVIDER_ENDPOINT` - `auth`: either `{ mode: 'apiKey', keyVaultSecretName, identityKeyVaultSecretName? }` or `{ mode: 'oauth' }`; unsupported modes fail closed @@ -236,12 +237,12 @@ Each provider is one unit exposing three things: `providerHttpStatus`, optional `providerMessageId`, `providerStatusName`, `providerStatusCode` and `providerStatusDescription` (snake_case attributes in Python, PascalCase in .NET). -The adapter reads its API-specific JSON and constructs a normalized `ParsedResponse` object: -[JavaScript](../javascript/src/functions/models.js), [Python](../python/src/models.py), -[.NET](../dotnet/Src/Models.cs). The engine reads named properties/attributes rather than provider JSON -or string-key response dictionaries. Optional values default to null/None; a status name takes precedence -over a code during outcome mapping, as before. Custom Python adapters must return `ParsedResponse`, -not the former dictionary. +Python and .NET use provider classes instead of a manifest-driven engine. Each class declares its +provider id and authentication mode, owns its credential secret names or OAuth acquisition, builds +its private outbound request, and maps provider JSON directly to `ProviderResult`: +[Python](../python/src/models.py), [.NET](../dotnet/Src/Models.cs). `ProviderResult` keeps `Outcome` +coarse and stable while `FailureReason` carries one fixed safe diagnostic classification. Optional +values default to null/None and raw provider JSON never enters the shared orchestration layer. This model is internal: do not serialize it into the endpoint response or logs. The [request logger](#application-logs) selects only the provider HTTP status, a status found in the @@ -252,12 +253,12 @@ Provider requests are serialized only when building the outbound HTTP body; inco is parsed once and normalized inside its adapter. No serialization framework or provider-specific class hierarchy is required. -Adapters require registration in the chosen runtime. Consult the selected adapter and its manifest -for required credentials and options: the manifest declares authentication and protocol mappings; -the implementation reads adapter-specific options from app settings. Individual API contracts remain +Providers require registration in the chosen runtime. Consult the selected provider for required +credentials and options. JavaScript declares them in its manifest; Python and .NET declare them on +the provider implementation. Individual API contracts remain in the adapters; the [onboarding credential naming table](ONBOARDING.md#provider-credential-names) -lists the exact manifest secret names for provisioning and authorized local tests. Keep that table -aligned with the manifests; never include secret values in documentation or the settings sample. +lists the exact secret names for provisioning and authorized local tests. Keep that table aligned +with the providers; never include secret values in documentation or the settings sample. ### Telesign EPP integration @@ -306,7 +307,7 @@ Set by provisioning. **Identical names across all languages.** | `KEY_VAULT_URL` | Key Vault URI for API-key providers | | `AZURE_CLIENT_ID` | set for a user-assigned managed identity | -Telesign credentials live in **Key Vault**, under the names in its manifest, and are fetched via +Telesign credentials live in **Key Vault**, under the names owned by its provider implementation, and are fetched via managed identity. Soprano exchanges an outbound managed-identity assertion for a token in the configured provider tenant/scope. Do not put provider secrets in code or app settings. @@ -345,11 +346,11 @@ when credentials are fetched, not the HTTP/nonce contract, caller authentication selection or provider request format. The decryption-key Key Vault reference remains separate and is still resolved by the platform; this cache does not rotate or replace JWE keys. -The selected provider's manifest determines which of two concrete cache classes is created: +The selected provider's authentication requirements determine which concrete cache 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`. | +| `apiKey` | `ApiKeyCache` | Fetch the provider'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 @@ -390,8 +391,9 @@ Per-request `providerCredentialElapsedMs` continues to measure the caller's reso set, else system-assigned). No static credentials. - **Privacy**: never log phone numbers, passcodes, nonce values, bearer tokens, API keys, JWE headers/payloads, raw exceptions, provider descriptions/responses or endpoint query strings. There is no plaintext diagnostic - override. Each handler emits separate service events. JavaScript and Python also emit one - [request summary](#application-logs); .NET uses standard structured `ILogger` events and scopes instead. + override. Each handler emits separate service events. JavaScript also emits one + [request summary](#application-logs); Python and .NET use standard structured logging events and + immutable request context/scopes instead. Generated Function IDs remain distinguished from raw Microsoft/provider support IDs. Original wire IDs and the required nonce echo remain unchanged. Support IDs can correlate customer activity; restrict log access and retention. Endpoint logs contain only scheme, host/port and API path, never userinfo, @@ -414,18 +416,20 @@ Per-request `providerCredentialElapsedMs` continues to measure the caller's reso ### Application logs -JavaScript and Python emit JSON records with the shared fields described below. Service events have -`logType: "service"` and an individual `eventName`; each invocation ends with one -`logType: "request"`, `eventName: "request_completed"` summary. +JavaScript emits JSON service records and one `request_completed` request summary. Python uses the +standard `logging` pipeline: `otp_log.py` defines fixed event IDs, names, levels and fields, while an +immutable `LoggerAdapter` context supplies the Function and Microsoft trace identifiers. The +configured logging provider owns Python output formatting and export; Python does not manually +serialize a mutable request summary. -.NET uses the standard `ILogger` pipeline instead of manually serializing JSON. `OtpLog` defines -source-generated events with stable IDs and names, while `ILogger.BeginScope` supplies +Likewise, .NET uses the standard `ILogger` pipeline instead of manually serializing JSON. `OtpLog` +defines source-generated events with stable IDs and names, while `ILogger.BeginScope` supplies `FunctionName`, `FunctionRequestId`, `FunctionInvocationId`, `MsClientRequestId`, `MsCorrelationId` and `MsCorrelationIdSource`. The configured logging provider owns output formatting and export. .NET emits `request_completed` as an ordinary typed event rather than a mutable comprehensive summary. -A successful .NET live request emits: +A successful Python or .NET live request emits: `request_received`, `payload_validated`, `delivery_context_decrypted`, `provider_selected`, `provider_credential_resolution_started`, `provider_credential_resolved`, @@ -445,13 +449,13 @@ bodies, decrypted delivery fields, credentials, provider response bodies, query exception messages. Endpoint values contain only scheme, host/port and path. Evaluation omits all provider events and emits `evaluation_completed`. -A successful JavaScript or Python live request emits these separate service events, followed by the -request summary: +A successful live request emits these separate service events. JavaScript then emits its request +summary; Python and .NET finish with the structured `request_completed` event: | Service event | Safe information recorded | |---|---| | `request_received` | Function invocation and available raw Microsoft trace IDs under their `x-ms-*` names; no raw body or arbitrary headers. | -| `envelope_validated` | Allowlisted body metadata: validated `envelopeType`, normalized `channel`, `evaluation`, optional `ttlSeconds`, and `encryptedDeliveryContextPresent: true`. | +| `envelope_validated` / `payload_validated` | Allowlisted body metadata: validated payload type, normalized `channel`, `evaluation` and optional `ttlSeconds`. | | `delivery_context_decrypted` | Decryption completed; no plaintext fields, JWE or key ID. | | `provider_selected` | Registered provider and its authentication mode. | | `provider_credential_resolution_started` | OAuth client-assertion or Key Vault credential source, with explicitly named raw OAuth application/identity/tenant IDs. | @@ -460,7 +464,7 @@ request summary: | `provider_request_built` | Allowlisted HTTP method, final endpoint scheme/host/port/API path, HTTPS and disabled redirects; no query string, authorization headers or body. | | `provider_request_started` | The outbound send is beginning, with method, sanitized endpoint and timeout. | | `provider_response_received` | Actual upstream HTTP status; emitted before response-body reading completes. | -| `provider_response_processed` | Mapped provider status/outcome, raw provider message/reference ID, duration and resulting Function HTTP status. | +| `provider_response_processed` | Mapped provider status/outcome, fixed failure classification and duration. Raw provider descriptions and bodies are excluded. | | `response_prepared` | Response status and booleans indicating nonce/correlation inclusion, not their values or the response body. | Body metadata is built from validated fields, **not** from a body dump with a few sensitive @@ -485,9 +489,9 @@ values are represented as `other` without changing the request sent. `response_prepared` is emitted on success **and failure** immediately before returning the handler response. It does not claim the host has serialized/transmitted that response or Microsoft received it; consult platform request telemetry for transport completion. Evaluation emits -`evaluation_completed` instead of provider events, then `response_prepared` and the summary, +`evaluation_completed` instead of provider events, then `response_prepared` and `request_completed`, without resolving provider configuration, credentials or HTTP. Failures emit their own stage event, -such as `decryption_failed`, `provider_credentials_failed` or `provider_transport_failed`. +`request_failed`, with a fixed failure stage and reason. A parsed provider rejection uses `provider_response_processed` with its non-success outcome and fixed failure reason. @@ -523,7 +527,7 @@ These are tracing fields, not authentication assertions. In particular, an incom does not become a trusted tenant identity in logs. The existing wire correlation precedence, provider request IDs and public responses are unchanged. -The JavaScript and Python request summary contains: +The JavaScript request summary contains: | Fields | Purpose | |---|---| diff --git a/python/README.md b/python/README.md index 55dae99..8ba93a0 100644 --- a/python/README.md +++ b/python/README.md @@ -1,14 +1,14 @@ # External Phone Provider Function: Python (v2 model) -Implements the shared [contract](../docs/CONTRACT.md) with one dispatch engine and one selected -provider per deployment. Target: Python 3.11, Azure Functions v4, Python v2 programming model. +Implements the shared [contract](../docs/CONTRACT.md) with a direct Azure Function flow and one +selected provider per deployment. Target: Python 3.11, Azure Functions v4, Python v2 programming model. ## Setup and deployment 1. Follow [customer onboarding](../docs/ONBOARDING.md). Set `EPP_PROVIDER_NAME` to the selected - adapter's registered manifest id (`` is only a placeholder). -2. Consult the selected adapter and its manifest in [src/providers/](src/providers/) for required - credentials and options. Store credentials in Key Vault under the declared secret names, grant + provider id (`` is only a placeholder). +2. Consult the selected provider in [src/providers/](src/providers/) for required credentials and + options. Store credentials in Key Vault under the provider-owned secret names, grant the Function's managed identity *Key Vault Secrets User*, and configure the matching endpoint/options. 3. Base private local settings on [../docs/local.settings.sample.json](../docs/local.settings.sample.json), replacing placeholders and selecting `FUNCTIONS_WORKER_RUNTIME=python`. Put settings at the @@ -49,7 +49,7 @@ For live delivery, add `EPP_PROVIDER_NAME`, the complete selected `EPP_PROVIDER_ matching provider authentication settings to `Values`. Add `EPP_PROVIDER_ACCOUNT_NAME` and any adapter-specific options only when required. Keep values as strings, including optional `EPP_PROVIDER_TIMEOUT_MS: "1500"`. Replace placeholders; provider API -keys belong in the manifest-named Key Vault secrets, not this file. See the +keys belong in the provider-named Key Vault secrets, not this file. See the [complete variable table](../README.md#configure-environment-variables). Core Tools loads `Values` into `os.environ`. Direct Python execution and pytest do not automatically @@ -90,7 +90,7 @@ 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. +Worker initialization selects `ApiKeyCache` or `AccessTokenCache` from the provider's credential specification. 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 @@ -101,16 +101,17 @@ local evaluation without background credential acquisition. | Source | Purpose | |---|---| -| [function_app.py](function_app.py) | HTTP handler and adapter registration | +| [function_app.py](function_app.py) | Typed request orchestration and direct provider selection | | [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/models.py](src/models.py) | Typed Entra payload, delivery context, provider request and result dataclasses | +| [src/jwe.py](src/jwe.py) | Pinned JWE decryption and typed delivery-context conversion | +| [src/provider.py](src/provider.py) | Shared HTTPS transport, timeout handling and endpoint status mapping | +| [src/credentials.py](src/credentials.py) | `CredentialTokenService`, `ApiKeyCache` and `AccessTokenCache` | +| [src/otp_log.py](src/otp_log.py) | Fixed standard-logging event definitions and immutable request context | +| [src/providers/](src/providers/) | Provider-owned credentials, requests and response mapping | | [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 -`parsed.provider_status_name`. Raw provider JSON remains local to the adapter, not a shared model hierarchy. +Add a provider by subclassing `PhoneProviderBase`, declaring its credential specification, and +implementing `build_request` and `map_response`. Return `ProviderResult` with the coarse endpoint +outcome and a fixed safe failure classification. Raw provider JSON remains local to the provider. See [production limitations](../docs/CONTRACT.md#production-limitations) before production use. diff --git a/python/function_app.py b/python/function_app.py index 4733702..adb4d1a 100644 --- a/python/function_app.py +++ b/python/function_app.py @@ -1,116 +1,258 @@ import atexit import json +import logging import os +import time import uuid from threading import Thread import azure.functions as func +from src import otp_log from src.config import read_config -from src.dispatch import ( - MODE_EVALUATION, - DispatchEngine, - ProviderRegistry, - context_to_dispatch, - decrypt_delivery_context, - make_key_provider, - parse_envelope, +from src.credentials import CredentialTokenService, report_refresh_failure +from src.jwe import JweDecryptor +from src.models import EntraSendOtpPayload, OtpDelivery +from src.provider import ( + ProviderSendError, + is_https_endpoint, + provider_timeout_ms, + to_endpoint_http_status, ) from src.providers.infobip import InfobipProvider from src.providers.sinch import SinchProvider from src.providers.soprano import SopranoProvider from src.providers.telesign import TelesignProvider -from src.request_log import RequestLog from src.secrets import SecretResolver app = func.FunctionApp() -_registry = ProviderRegistry([InfobipProvider(), TelesignProvider(), SopranoProvider(), SinchProvider()]) -_secrets = SecretResolver() -_engine = DispatchEngine(_registry, _secrets) -_key_provider = make_key_provider(os.environ) +_providers = { + provider.name: provider + for provider in (InfobipProvider(), TelesignProvider(), SopranoProvider(), SinchProvider()) +} +_credentials = CredentialTokenService(SecretResolver()) +_decryptor = JweDecryptor(os.environ) +_logger = logging.getLogger("epp.send_otp") + + +class InvalidRequest(Exception): + def __init__(self, status_code, error, reason=None, correlation_id=None): + super().__init__(error) + self.status_code = status_code + self.error = error + self.reason = reason + self.correlation_id = correlation_id + + +def _select_provider(name): + return _providers.get(name.lower()) if isinstance(name, str) and name else None + + +def _failure(log, status, stage, reason): + otp_log.request_failed( + log, + logging.ERROR if status >= 500 else logging.WARNING, + stage, + reason, + status, + ) + return status + + +def _send_to_provider(delivery, log): + config = read_config() + provider = _select_provider(config.provider_name) + if provider is None: + return _failure(log, 400, "provider_selection", "unknown_provider"), log + + log = log.with_context({"providerName": provider.name}) + otp_log.provider_selected(log, provider.name, provider.authentication_mode) + channel = (delivery.channel or "sms").lower() + if channel not in ("sms", "voice"): + return _failure(log, 400, "provider_configuration", "unsupported_channel"), log + if config.provider_channel and config.provider_channel != channel: + return _failure(log, 400, "provider_configuration", "channel_not_configured"), log + if config.provider_auth_mode and config.provider_auth_mode != provider.authentication_mode: + return _failure(log, 502, "provider_configuration", "authentication_mode_mismatch"), log + + credential_source = ( + "managed_identity_client_assertion" + if provider.authentication_mode == "oauth" + else "key_vault" + ) + credential_started = time.monotonic() + try: + otp_log.credential_resolution_started(log, provider.name, credential_source) + credential = _credentials.get_credentials(provider, config) + except Exception: + return _failure(log, 502, "provider_credentials", "credential_unavailable"), log + otp_log.credential_resolved( + log, provider.name, int((time.monotonic() - credential_started) * 1000)) + + if not is_https_endpoint(config.provider_endpoint): + return _failure(log, 502, "provider_configuration", "invalid_provider_endpoint"), log + + try: + result = provider.send_otp( + channel, + config.provider_endpoint, + delivery, + credential, + config.env, + provider_timeout_ms(config.provider_timeout_ms), + log, + ) + except ProviderSendError as error: + return error.status_code, log + status = to_endpoint_http_status(result) + if status >= 400: + _failure( + log, + status, + "provider_response", + result.failure_reason or "provider_rejected", + ) + return status, log # Azure Easy Auth must enforce authentication; local handler calls are anonymous. @app.route(route="SendOtp", methods=["POST"], auth_level=func.AuthLevel.ANONYMOUS) def send_otp(req: func.HttpRequest, context: func.Context = None) -> func.HttpResponse: - request_id = str(uuid.uuid4()) - ms_request_id = req.headers.get("x-ms-client-request-id") - client_request_id = ms_request_id or request_id - header_correlation_id = req.headers.get("x-ms-correlation-id") - log = RequestLog(request_id, context.invocation_id if context else None, - ms_request_id, header_correlation_id, context.function_name if context else "send_otp") - correlation_id = header_correlation_id or request_id - http_status = 500 - evaluation = False - - def respond(status, body): - nonlocal http_status - response = func.HttpResponse(json.dumps(body), status_code=status, mimetype="application/json") - http_status = status - log.response_prepared(status, "nonce" in body, "correlationId" in body) - return response + started = time.monotonic() + request_id = uuid.uuid4().hex + ms_request_id = otp_log.safe_identifier(req.headers.get("x-ms-client-request-id")) + header_correlation_id = otp_log.safe_identifier(req.headers.get("x-ms-correlation-id")) + log = otp_log.RequestLogger(_logger, { + "functionName": context.function_name if context else "send_otp", + "functionRequestId": request_id, + "functionInvocationId": context.invocation_id if context else None, + "x-ms-client-request-id": ms_request_id, + "x-ms-correlation-id": header_correlation_id, + "msCorrelationIdSource": "header" if header_correlation_id else "none", + "channel": None, + "evaluation": None, + "providerName": None, + }) + status_code = 500 + body = None + payload = None try: - log.service("request_received") - config = read_config() + otp_log.request_received(log) try: - payload = req.get_json() + raw_payload = req.get_json() except ValueError: - log.failure("request_validation", "invalid JSON body", 400) - return respond(400, {"error": "bad_request", "reason": "invalid JSON body", "requestId": request_id}) + _failure(log, 400, "request_validation", "invalid JSON body") + raise InvalidRequest(400, "bad_request", "invalid JSON body") - envelope, error = parse_envelope(payload) - if error: - log.failure("request_validation", error, 400) - return respond(400, {"error": "bad_request", "reason": error, "requestId": request_id}) + payload, payload_error = EntraSendOtpPayload.from_payload(raw_payload) + if payload_error: + _failure(log, 400, "request_validation", payload_error) + raise InvalidRequest(400, "bad_request", payload_error) - correlation_id = envelope.correlation_id or header_correlation_id or request_id - log.envelope_validated(envelope, envelope.correlation_id or header_correlation_id, - "envelope" if envelope.correlation_id else "header") - envelope.correlation_id = correlation_id - evaluation = envelope.mode == MODE_EVALUATION + payload_correlation_id = otp_log.safe_identifier(payload.correlation_id) + correlation_id = payload.correlation_id or header_correlation_id or request_id + log = log.with_context({ + "x-ms-correlation-id": payload_correlation_id or header_correlation_id, + "msCorrelationIdSource": ( + "envelope" if payload_correlation_id + else "header" if header_correlation_id + else "none" + ), + "channel": payload.channel_name, + "evaluation": payload.is_evaluation, + }) + otp_log.payload_validated( + log, payload.type, payload.channel_name, payload.is_evaluation, payload.ttl_seconds) + config = read_config() try: - header, delivery = decrypt_delivery_context(envelope.encrypted_delivery_context, _key_provider) + decrypted = _decryptor.decrypt(payload.encrypted_delivery_context) except Exception: - log.failure("decryption", "decryption_failed", 400) - return respond(400, {"error": "decryption_failed", "correlationId": correlation_id, "requestId": request_id}) - log.service("delivery_context_decrypted") - - if config.expected_key_id and header.get("kid") != config.expected_key_id: - log.key_id_mismatch() - - if delivery is None or not delivery.is_complete: - log.failure("delivery_context_validation", "incomplete delivery context", 400) - return respond(400, {"error": "bad_request", "reason": "incomplete delivery context", - "correlationId": correlation_id, "requestId": request_id}) - - # Evaluation skips provider lookup, configuration, secrets and HTTP. - if not evaluation: - dispatch = context_to_dispatch(delivery, envelope, client_request_id) - status, _ = _engine.dispatch(dispatch, request_id, log) - if status != 200: - return respond(status, {"error": "provider_delivery_failed", - "correlationId": correlation_id, "requestId": request_id}) + _failure(log, 400, "decryption", "decryption_failed") + raise InvalidRequest(400, "decryption_failed", correlation_id=correlation_id) + otp_log.delivery_context_decrypted(log) + + if config.expected_key_id and config.expected_key_id != decrypted.key_id: + otp_log.encryption_key_id_mismatch(log) + delivery_context = decrypted.value + if delivery_context is None or not delivery_context.is_complete: + _failure(log, 400, "delivery_context_validation", "incomplete delivery context") + raise InvalidRequest( + 400, "bad_request", "incomplete delivery context", correlation_id) + + if payload.is_evaluation: + otp_log.evaluation_completed(log) else: - log.service("evaluation_completed") + delivery = OtpDelivery( + delivery_context.phone_number, + delivery_context.message, + payload.channel_name, + ms_request_id or request_id, + correlation_id, + delivery_context.locale, + ) + provider_status, log = _send_to_provider(delivery, log) + if provider_status >= 400: + raise InvalidRequest( + provider_status, "provider_delivery_failed", correlation_id=correlation_id) - # Live delivery must finish before nonce acceptance. - return respond(200, { - "nonce": delivery.nonce, + status_code = 200 + body = { + "nonce": delivery_context.nonce, "correlationId": correlation_id, "providerStatus": "accepted", - }) + } + except InvalidRequest as error: + status_code = error.status_code + body = {"error": error.error, "requestId": request_id} + if error.reason is not None: + body["reason"] = error.reason + if error.correlation_id is not None: + body["correlationId"] = error.correlation_id + except Exception: + otp_log.unexpected_error(log) + status_code = 500 + body = { + "error": "delivery_failed", + "correlationId": ( + payload.correlation_id if payload is not None + else header_correlation_id or request_id + ), + "requestId": request_id, + } + + otp_log.response_prepared( + log, status_code, "nonce" in body, "correlationId" in body) + otp_log.request_completed( + log, + status_code, + "evaluated" if status_code == 200 and payload and payload.is_evaluation + else "accepted" if status_code == 200 + else "failed", + int((time.monotonic() - started) * 1000), + ) + return func.HttpResponse(json.dumps(body), status_code=status_code, mimetype="application/json") + + +def _warm_selected_credentials(): + config = read_config() + if not config.provider_name: + return + provider = _select_provider(config.provider_name) + if provider is None or ( + config.provider_auth_mode + and config.provider_auth_mode != provider.authentication_mode + ): + report_refresh_failure("configuration") + return + try: + _credentials.get_credentials(provider, config) except Exception: - if not log.data["failureStage"]: - log.failure("handler", "unexpected_error", 500) - return respond(500, {"error": "delivery_failed", "correlationId": correlation_id, "requestId": request_id}) - finally: - log.complete(http_status) + report_refresh_failure("initialization") -# Each worker owns its own memory cache; a timer trigger would only warm one worker. -atexit.register(_engine.close) +atexit.register(_credentials.close) if os.environ.get("EPP_PROVIDER_NAME", "").strip(): - Thread(target=_engine.start_credential_refresh, daemon=True).start() + Thread(target=_warm_selected_credentials, daemon=True).start() diff --git a/python/src/credentials.py b/python/src/credentials.py index fa21628..b271ffc 100644 --- a/python/src/credentials.py +++ b/python/src/credentials.py @@ -1,6 +1,5 @@ from __future__ import annotations -import json import logging import math import threading @@ -15,6 +14,7 @@ from cachetools import TTLCache from .config import AppConfig +from . import otp_log API_KEY_MODE: Final = "apiKey" OAUTH_MODE: Final = "oauth" @@ -66,8 +66,7 @@ def filter(self, record: logging.LogRecord) -> bool: def report_refresh_failure(kind: str) -> None: - logging.warning("%s", json.dumps({"logType": "service", "eventName": "credential_refresh_failed", - "cacheKind": kind, "failureReason": "credential_unavailable"})) + otp_log.credential_refresh_failed(logging.getLogger(__name__), kind) def _private_acquisition(load: Callable[[], T]) -> T: @@ -190,8 +189,8 @@ def stop(self) -> None: self._token = None -class ProviderCredentials: - """One selected cache and one refresh loop; configuration changes require a worker restart.""" +class CredentialTokenService: + """Caches credentials for the provider selected by this worker.""" def __init__(self, secrets: SecretReader, *, cache_options: RefreshOptions | None = None, report_failure: Callable[[str], None] = report_refresh_failure) -> None: @@ -205,15 +204,18 @@ def __init__(self, secrets: SecretReader, *, cache_options: RefreshOptions | Non self._pending: Future[ApiKeyCredential | OAuthCredential] | None = None self._next_attempt = 0.0 - def resolve(self, auth: Mapping[str, str], config: AppConfig) -> ApiKeyCredential | OAuthCredential: + def get_credentials(self, provider, config: AppConfig) -> ApiKeyCredential | OAuthCredential: + return self.resolve(provider.credential_spec, config) + + def resolve(self, spec: 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: + if spec.get("mode") == API_KEY_MODE: + self.cache = ApiKeyCache(self._secrets, spec, self._clock) + elif spec.get("mode") == OAUTH_MODE: self.cache = _private_acquisition(lambda: AccessTokenCache(config, self._clock)) else: raise ValueError(CREDENTIAL_ERROR) diff --git a/python/src/dispatch.py b/python/src/dispatch.py deleted file mode 100644 index d8f637a..0000000 --- a/python/src/dispatch.py +++ /dev/null @@ -1,418 +0,0 @@ -from __future__ import annotations - -import base64 -import json -import logging -import os -from urllib.parse import urlsplit - -import requests -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 - -DEFAULT_TIMEOUT_MS = 1500 -DEFAULT_CHANNELS = ["sms", "voice"] - -CONTINUE = "Continue" -FAIL = "Fail" -BLOCK = "Block" -STEP_UP = "StepUp" - -def resolve_outcome(manifest, parsed: ParsedResponse): - mapping = manifest["response_mapping"] - key = parsed.provider_status_name or parsed.provider_status_code - if key: - outcome = mapping.get(key) or mapping.get("default", FAIL) - else: - outcome = CONTINUE if parsed.success else mapping.get("default", FAIL) - return FAIL if outcome == CONTINUE and not parsed.success else outcome - - -def to_http_status(outcome, provider_http_status): - if outcome == CONTINUE: - return 200 - if outcome == BLOCK: - return 403 - if outcome == STEP_UP: - return 409 - if outcome == FAIL: - if provider_http_status == 429: - return 429 - if provider_http_status in (401, 403): - return 401 - if 400 <= provider_http_status < 500: - return 400 - return 502 - - -class ProviderRegistry: - def __init__(self, adapters): - self._by_id = {adapter.manifest["id"].lower(): adapter for adapter in adapters} - - def get(self, provider_id): - if not provider_id: - return None - return self._by_id.get(provider_id.lower()) - - -CHANNEL_BY_CODE = {1: "sms", 2: "voice"} -CHANNEL_BY_NAME = {"sms": 1, "voice": 2} -MODE_LIVE = 1 -MODE_EVALUATION = 2 -MODE_BY_NAME = {"live": MODE_LIVE, "evaluation": MODE_EVALUATION} - -MAX_JWE_LENGTH = 16384 - - -def _normalize_channel(channel): - if type(channel) is int: - return channel if channel in CHANNEL_BY_CODE else None - if isinstance(channel, str): - return CHANNEL_BY_NAME.get(channel.lower()) - return None - - -def _normalize_mode(mode): - if type(mode) is int: - return mode if mode in (MODE_LIVE, MODE_EVALUATION) else None - if isinstance(mode, str): - return MODE_BY_NAME.get(mode.lower()) - return None - - -def parse_envelope(payload) -> tuple[Envelope | None, str | None]: - if not isinstance(payload, dict): - return None, "invalid envelope" - if payload.get("type") != "microsoft.mfa.otpDeliver.v1": - return None, "unsupported envelope type" - encrypted = payload.get("encryptedDeliveryContext") - if not isinstance(encrypted, str) or not encrypted.strip(): - return None, "encryptedDeliveryContext is required" - channel = _normalize_channel(payload.get("channel")) - if channel is None: - return None, "unsupported channel" - mode = _normalize_mode(payload.get("mode")) - if mode is None: - return None, "unsupported mode" - ttl_seconds = payload.get("ttlSeconds") - if "ttlSeconds" in payload: - if type(ttl_seconds) is not int: - return None, "invalid ttlSeconds" - if ttl_seconds <= 0: - return None, "ttlSeconds expired" - if ttl_seconds > 2147483647: - return None, "invalid ttlSeconds" - return Envelope( - type=payload.get("type"), - tenant_id=payload.get("tenantId"), - correlation_id=payload.get("correlationId"), - channel=channel, - mode=mode, - ttl_seconds=ttl_seconds, - encrypted_delivery_context=encrypted, - ), None - - -def read_protected_header(compact_jwe): - header_segment = compact_jwe.split(".")[0] - header_segment += "=" * (-len(header_segment) % 4) - return json.loads(base64.urlsafe_b64decode(header_segment)) - - -def make_key_provider(env): - def key_provider(_kid): - return read_config(env).decryption_key_pem - - return key_provider - - -def _assert_well_formed_jwe(compact_jwe): - if not isinstance(compact_jwe, str) or not compact_jwe: - raise ValueError("malformed JWE") - if len(compact_jwe) > MAX_JWE_LENGTH: - raise ValueError("delivery context exceeds size limit") - segments = compact_jwe.split(".") - if len(segments) != 5 or not all(segments): - raise ValueError("malformed JWE: expected five non-empty segments") - - -# Cache only the configured key to avoid repeated RSA imports. -_key_cache = {} - - -def _normalize_pem(value): - # Base64 preserves PEM newlines in app settings. - text = value if isinstance(value, str) else value.decode("utf-8") - if "-----BEGIN" in text: - return text - return base64.b64decode(text).decode("utf-8") - - -def _load_private_key(pem): - if not pem: - raise ValueError("private key unavailable (EPP_DECRYPTION_KEY_PEM is not set)") - cached = _key_cache.get(pem) - if cached is None: - cached = jwk.JWK.from_pem(_normalize_pem(pem).encode("utf-8")) - _key_cache.clear() - _key_cache[pem] = cached - return cached - - -def decrypt_delivery_context(compact_jwe, key_provider): - _assert_well_formed_jwe(compact_jwe) - header = read_protected_header(compact_jwe) - key = _load_private_key(key_provider(header.get("kid"))) - token = jwe_module.JWE(algs=["RSA-OAEP-256", "A256GCM"]) - token.deserialize(compact_jwe, key=key) - payload = json.loads(token.payload.decode("utf-8")) - return header, DeliveryContext.from_payload(payload) - - -def context_to_dispatch(context, envelope, message_id): - return DispatchRequest( - destination=context.phone_number, - message=context.message, - channel=CHANNEL_BY_CODE[envelope.channel], - message_id=message_id, - correlation_id=envelope.correlation_id, - locale=context.locale, - text_to_voice=context.text_to_voice, - ) - - -def _valid_provider_url(value): - if not isinstance(value, str) or not value or "#" in value: - return False - # urlsplit strips controls; reject them before parsing. - if any(character.isspace() or ord(character) < 32 or ord(character) == 127 for character in value): - return False - try: - parsed = urlsplit(value) - port = parsed.port # Access validates the port's syntax and range. - return ( - parsed.scheme == "https" - and bool(parsed.hostname) - and port != 0 - and parsed.username is None - and parsed.password is None - and not parsed.fragment - and not parsed.netloc.endswith(":") - ) - except ValueError: - return False - - -def _provider_timeout_ms(value): - digits = value.strip() if isinstance(value, str) else "" - if not digits or not digits.isascii() or not digits.isdecimal(): - return DEFAULT_TIMEOUT_MS - digits = digits.lstrip("0") - if not digits: - return DEFAULT_TIMEOUT_MS - # Clamp before int() to avoid its digit limit. - if len(digits) > 4 or (len(digits) == 4 and digits > "2500"): - return 2500 - return int(digits) - - -def _has_read_timeout(error): - # requests wraps urllib3 body-read timeouts in ConnectionError. - pending = [error] - seen = set() - while pending: - current = pending.pop() - if id(current) in seen: - continue - seen.add(id(current)) - if isinstance(current, ReadTimeoutError): - return True - pending.extend( - nested for nested in (current.__cause__, current.__context__, *current.args) - if isinstance(nested, Exception) - ) - return False - - -class DispatchEngine: - 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._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): - if log: - log.failure(stage, reason, status) - return status, body - - config = read_config(self.env) - adapter = self.registry.get(config.provider_name) - if adapter is None: - return failure(400, "provider_selection", "unknown_provider", - {"status": "error", "reason": "unknown provider", "requestId": request_id}) - - manifest = adapter.manifest - if log: - log.provider_selected(manifest) - provider_id = manifest["id"] - channel = dispatch.channel if dispatch.channel is not None else "sms" - if not isinstance(channel, str): - return failure(400, "provider_configuration", "unsupported_channel", - {"status": "error", "provider": provider_id, "reason": "unsupported channel", "requestId": request_id}) - channel = channel.lower() - - if channel not in DEFAULT_CHANNELS: - return failure(400, "provider_configuration", "unsupported_channel", - {"status": "error", "provider": provider_id, "reason": "unsupported channel", "requestId": request_id}) - if config.provider_channel and config.provider_channel != channel: - return failure(400, "provider_configuration", "channel_not_configured", - {"status": "error", "provider": provider_id, "reason": "channel not configured", "requestId": request_id}) - - if channel == "voice" and manifest.get("requires_text_to_voice") and ( - not isinstance(dispatch.text_to_voice, TextToVoice) or not dispatch.text_to_voice.is_complete - ): - return failure(400, "provider_request_build", "incomplete_voice_context", - self._fail_body(provider_id, channel, "incomplete voice context", dispatch, request_id)) - - auth = manifest["auth"] - if config.provider_auth_mode and config.provider_auth_mode != auth.get("mode"): - return failure(502, "provider_configuration", "authentication_mode_mismatch", - self._fail_body(provider_id, channel, "provider authentication mismatch", dispatch, request_id)) - try: - if log: - log.credential_resolution_started(config) - credential = self._resolve_credential(auth, config) - except Exception: - return failure(502, "provider_credentials", "credential_unavailable", - self._fail_body(provider_id, channel, "provider credential unavailable", dispatch, request_id)) - credential_unavailable = ( - credential.get("mode") == "apiKey" - and (not credential.get("secret") or (auth.get("identity_key_vault_secret_name") and not credential.get("identity"))) - ) or (credential.get("mode") == "oauth" and not credential.get("access_token")) - if credential_unavailable: - return failure(502, "provider_credentials", "credential_unavailable", - self._fail_body(provider_id, channel, "provider credential unavailable", dispatch, request_id)) - if log: - log.credential_resolved() - - endpoint = config.provider_endpoint - if not endpoint: - return failure(502, "provider_configuration", "invalid_provider_endpoint", - self._fail_body(provider_id, channel, "provider endpoint not configured", dispatch, request_id)) - if not _valid_provider_url(endpoint): - return failure(502, "provider_configuration", "invalid_provider_endpoint", - self._fail_body(provider_id, channel, "invalid provider endpoint", dispatch, request_id)) - - try: - if log: - log.service("provider_request_build_started") - provider_request = adapter.build_request(channel, endpoint, dispatch, credential, config.env) - except Exception: - return failure(502, "provider_request_build", "request_build_failed", - self._fail_body(provider_id, channel, "provider request failed", dispatch, request_id)) - if not _valid_provider_url(provider_request.get("url")): - return failure(502, "provider_request_build", "invalid_provider_request_url", - self._fail_body(provider_id, channel, "invalid provider request URL", dispatch, request_id)) - if log: - log.provider_request_built(provider_request.get("method"), provider_request["url"]) - - timeout_ms = _provider_timeout_ms(config.provider_timeout_ms) - response = None - stage = "provider_transport" - try: - if log: - log.provider_request_started(timeout_ms) - response = requests.request( - provider_request["method"], - provider_request["url"], - headers=provider_request["headers"], - data=provider_request["body"], - # Connect/read inactivity, not a total delivery deadline. - timeout=timeout_ms / 1000, - allow_redirects=False, # Never forward credentials to a redirect target. - stream=True, # Own the response for cleanup if body reading fails. - ) - if log: - log.provider_response_received(response.status_code) - - valid_json = True - try: - body_json = response.json() - except ValueError: - valid_json = False - if log: - log.service("provider_response_invalid_json", level=logging.WARNING) - body_json = {} - if log: - log.provider_request_finished() - - stage = "provider_response" - ok = 200 <= response.status_code < 300 - parsed = adapter.parse_response(response.status_code, ok, body_json) - outcome = resolve_outcome(manifest, parsed) - http_status = to_http_status(outcome, parsed.provider_http_status or response.status_code) - if log: - log.provider_response_processed(manifest, parsed, outcome, http_status, valid_json) - - return http_status, { - "status": "accepted" if outcome == CONTINUE else "failed", - "outcome": outcome, - "provider": provider_id, - "channel": channel, - "messageId": dispatch.message_id, - "correlationId": dispatch.correlation_id, - "requestId": request_id, - } - except requests.exceptions.RequestException as error: - if response is None: - response = getattr(error, "response", None) - if isinstance(error, requests.exceptions.Timeout) or ( - isinstance(error, requests.exceptions.ConnectionError) and _has_read_timeout(error) - ): - return failure(504, "provider_transport", "provider_timeout", - self._fail_body(provider_id, channel, "provider timeout", dispatch, request_id)) - return failure(502, "provider_transport", "provider_network_error", - self._fail_body(provider_id, channel, "provider request failed", dispatch, request_id)) - except Exception: - return failure(502, stage, "response_parse_failed" if stage == "provider_response" else "provider_network_error", - self._fail_body(provider_id, channel, "provider response failed", dispatch, request_id)) - finally: - if log: - log.provider_request_finished() - close = getattr(response, "close", None) - if callable(close): - try: - close() - except Exception: - if log: - log.service("provider_response_cleanup_failed", level=logging.WARNING) - - def _resolve_credential(self, auth, config): - 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/jwe.py b/python/src/jwe.py new file mode 100644 index 0000000..03acce7 --- /dev/null +++ b/python/src/jwe.py @@ -0,0 +1,68 @@ +from __future__ import annotations + +import base64 +import json +from dataclasses import dataclass + +from jwcrypto import jwe as jwe_module +from jwcrypto import jwk + +from .config import read_config +from .models import DeliveryContext + +MAX_JWE_LENGTH = 16384 + + +@dataclass(frozen=True) +class DecryptedPayload: + key_id: str | None + value: DeliveryContext | None + + +class JweDecryptor: + def __init__(self, env=None) -> None: + self._env = env + self._cached_pem: str | None = None + self._cached_key: jwk.JWK | None = None + + def decrypt(self, encrypted_content: str) -> DecryptedPayload: + self._assert_well_formed(encrypted_content) + header = self._read_protected_header(encrypted_content) + token = jwe_module.JWE(algs=["RSA-OAEP-256", "A256GCM"]) + token.deserialize(encrypted_content, key=self._get_private_key()) + payload = json.loads(token.payload.decode("utf-8")) + return DecryptedPayload(header.get("kid"), DeliveryContext.from_payload(payload)) + + def _get_private_key(self) -> jwk.JWK: + pem = read_config(self._env).decryption_key_pem + if not pem: + raise ValueError("private key unavailable") + if self._cached_key is None or self._cached_pem != pem: + self._cached_key = jwk.JWK.from_pem(self._normalize_pem(pem).encode("utf-8")) + self._cached_pem = pem + return self._cached_key + + @staticmethod + def _assert_well_formed(encrypted_content: str) -> None: + if not isinstance(encrypted_content, str) or not encrypted_content: + raise ValueError("malformed JWE") + if len(encrypted_content) > MAX_JWE_LENGTH: + raise ValueError("delivery context exceeds size limit") + segments = encrypted_content.split(".") + if len(segments) != 5 or not all(segments): + raise ValueError("malformed JWE") + + @staticmethod + def _read_protected_header(encrypted_content: str) -> dict: + segment = encrypted_content.split(".", 1)[0] + segment += "=" * (-len(segment) % 4) + value = json.loads(base64.urlsafe_b64decode(segment)) + if not isinstance(value, dict): + raise ValueError("invalid protected header") + return value + + @staticmethod + def _normalize_pem(value: str) -> str: + if "-----BEGIN" in value: + return value + return base64.b64decode(value).decode("utf-8") diff --git a/python/src/models.py b/python/src/models.py index ee2bf14..9305627 100644 --- a/python/src/models.py +++ b/python/src/models.py @@ -1,39 +1,87 @@ from __future__ import annotations from dataclasses import dataclass +from enum import Enum -@dataclass(repr=False) -class Envelope: +CHANNEL_BY_CODE = {1: "sms", 2: "voice"} +CHANNEL_BY_NAME = {"sms": 1, "voice": 2} +MODE_LIVE = 1 +MODE_EVALUATION = 2 +MODE_BY_NAME = {"live": MODE_LIVE, "evaluation": MODE_EVALUATION} + + +def _normalize_channel(channel: object) -> int | None: + if type(channel) is int: + return channel if channel in CHANNEL_BY_CODE else None + if isinstance(channel, str): + return CHANNEL_BY_NAME.get(channel.lower()) + return None + + +def _normalize_mode(mode: object) -> int | None: + if type(mode) is int: + return mode if mode in (MODE_LIVE, MODE_EVALUATION) else None + if isinstance(mode, str): + return MODE_BY_NAME.get(mode.lower()) + return None + + +@dataclass(frozen=True, repr=False) +class EntraSendOtpPayload: type: str - tenant_id: object - correlation_id: object + tenant_id: str | None + correlation_id: str | None channel: int mode: int ttl_seconds: int | None encrypted_delivery_context: str - -@dataclass(repr=False) -class TextToVoice: - before_password_text: object - password: object - language: object - @classmethod - def from_payload(cls, payload: object) -> "TextToVoice | None": + def from_payload(cls, payload: object) -> tuple["EntraSendOtpPayload | None", str | None]: if not isinstance(payload, dict): - return None - return cls(payload.get("beforePasswordText"), payload.get("password"), payload.get("language")) + return None, "invalid envelope" + if payload.get("type") != "microsoft.mfa.otpDeliver.v1": + return None, "unsupported envelope type" + encrypted = payload.get("encryptedDeliveryContext") + if not isinstance(encrypted, str) or not encrypted.strip(): + return None, "encryptedDeliveryContext is required" + channel = _normalize_channel(payload.get("channel")) + if channel is None: + return None, "unsupported channel" + mode = _normalize_mode(payload.get("mode")) + if mode is None: + return None, "unsupported mode" + ttl_seconds = payload.get("ttlSeconds") + if "ttlSeconds" in payload: + if type(ttl_seconds) is not int or ttl_seconds > 2147483647: + return None, "invalid ttlSeconds" + if ttl_seconds <= 0: + return None, "ttlSeconds expired" + return cls( + type=payload.get("type"), + tenant_id=payload.get("tenantId") if isinstance(payload.get("tenantId"), str) else None, + correlation_id=( + payload.get("correlationId") + if isinstance(payload.get("correlationId"), str) and payload.get("correlationId") + else None + ), + channel=channel, + mode=mode, + ttl_seconds=ttl_seconds, + encrypted_delivery_context=encrypted, + ), None @property - def is_complete(self) -> bool: - return isinstance(self.before_password_text, str) and all( - isinstance(value, str) and value.strip() for value in (self.password, self.language) - ) + def channel_name(self) -> str: + return CHANNEL_BY_CODE[self.channel] + + @property + def is_evaluation(self) -> bool: + return self.mode == MODE_EVALUATION -@dataclass(repr=False) +@dataclass(frozen=True, repr=False) class DeliveryContext: # Keep raw JSON values until is_complete validates the required strings. nonce: object @@ -42,7 +90,6 @@ class DeliveryContext: locale: object = None extension: object = None risk_context: object = None - text_to_voice: TextToVoice | None = None @classmethod def from_payload(cls, payload: object) -> "DeliveryContext | None": @@ -52,10 +99,9 @@ def from_payload(cls, payload: object) -> "DeliveryContext | None": nonce=payload.get("nonce"), phone_number=payload.get("phoneNumber"), message=payload.get("message"), - locale=payload.get("locale"), + locale=payload.get("locale") if isinstance(payload.get("locale"), str) else None, extension=payload.get("extension"), risk_context=payload.get("riskContext"), - text_to_voice=TextToVoice.from_payload(payload.get("textToVoice")), ) @property @@ -66,24 +112,29 @@ def is_complete(self) -> bool: ) -@dataclass(repr=False) -class DispatchRequest: - destination: str +@dataclass(frozen=True, repr=False) +class OtpDelivery: + phone_number: str message: str | None channel: str message_id: str correlation_id: str | None locale: str | None - text_to_voice: TextToVoice | None = None -@dataclass(repr=False) -class ParsedResponse: - """Adapter-normalized result for outcome mapping, not a public HTTP response.""" +class Outcome(str, Enum): + CONTINUE = "Continue" + FAIL = "Fail" + BLOCK = "Block" + - success: bool +@dataclass(frozen=True, repr=False) +class ProviderResult: + outcome: Outcome + status_recognized: bool provider_http_status: int provider_message_id: str | None = None provider_status_name: str | None = None provider_status_code: str | None = None - provider_status_description: str | None = None \ No newline at end of file + provider_status_description: str | None = None + failure_reason: str | None = None \ No newline at end of file diff --git a/python/src/otp_log.py b/python/src/otp_log.py new file mode 100644 index 0000000..07f723f --- /dev/null +++ b/python/src/otp_log.py @@ -0,0 +1,118 @@ +from __future__ import annotations + +import logging +import re +from collections.abc import Mapping + +_IDENTIFIER_PATTERN = re.compile(r"[A-Za-z0-9][A-Za-z0-9._:-]{0,127}") + + +def safe_identifier(value: object) -> str | None: + return value if isinstance(value, str) and _IDENTIFIER_PATTERN.fullmatch(value) else None + + +class RequestLogger(logging.LoggerAdapter): + def process(self, msg, kwargs): + extra = dict(self.extra) + extra.update(kwargs.pop("extra", {})) + kwargs["extra"] = extra + return msg, kwargs + + def with_context(self, values: Mapping[str, object]) -> "RequestLogger": + return RequestLogger(self.logger, {**self.extra, **values}) + + +def _emit(log: logging.LoggerAdapter, event_id: int, event_name: str, message: str, + *, level: int = logging.INFO, **fields: object) -> None: + log.log(level, message, extra={"event_id": event_id, "event_name": event_name, **fields}) + + +def request_received(log): _emit(log, 1000, "request_received", "OTP request received") + + +def payload_validated(log, payload_type, channel, evaluation, ttl_seconds): + _emit(log, 1001, "payload_validated", "OTP payload validated", + payloadType=payload_type, channel=channel, evaluation=evaluation, ttlSeconds=ttl_seconds) + + +def delivery_context_decrypted(log): + _emit(log, 1002, "delivery_context_decrypted", "Delivery context decrypted") + + +def encryption_key_id_mismatch(log): + _emit(log, 1003, "encryption_key_id_mismatch", + "Encrypted payload key identifier did not match the configured identifier", + level=logging.WARNING) + + +def evaluation_completed(log): _emit(log, 1004, "evaluation_completed", "Evaluation request completed") + + +def provider_selected(log, provider_name, authentication_mode): + _emit(log, 1100, "provider_selected", "Provider selected", + providerName=provider_name, authenticationMode=authentication_mode) + + +def credential_resolution_started(log, provider_name, credential_source): + _emit(log, 1101, "provider_credential_resolution_started", "Resolving provider credentials", + providerName=provider_name, credentialSource=credential_source) + + +def credential_resolved(log, provider_name, elapsed_ms): + _emit(log, 1102, "provider_credential_resolved", "Provider credentials resolved", + providerName=provider_name, elapsedMs=elapsed_ms) + + +def provider_request_build_started(log): + _emit(log, 1200, "provider_request_build_started", "Building provider request") + + +def provider_request_built(log, http_method, provider_endpoint): + _emit(log, 1201, "provider_request_built", "Provider request built", + httpMethod=http_method, providerEndpoint=provider_endpoint, redirectsAllowed=False) + + +def provider_request_started(log, timeout_ms): + _emit(log, 1202, "provider_request_started", "Sending provider request", timeoutMs=timeout_ms) + + +def provider_response_received(log, provider_http_status): + _emit(log, 1203, "provider_response_received", "Provider response received", + providerHttpStatus=provider_http_status) + + +def provider_response_invalid_json(log): + _emit(log, 1204, "provider_response_invalid_json", "Provider response was not valid JSON", + level=logging.WARNING) + + +def provider_response_processed(log, level, provider_http_status, provider_status, + provider_outcome, failure_reason, elapsed_ms): + _emit(log, 1205, "provider_response_processed", "Provider response processed", level=level, + providerHttpStatus=provider_http_status, providerStatus=provider_status, + providerOutcome=provider_outcome, failureReason=failure_reason, elapsedMs=elapsed_ms) + + +def request_failed(log, level, failure_stage, failure_reason, http_status): + _emit(log, 1300, "request_failed", "OTP request failed", level=level, + failureStage=failure_stage, failureReason=failure_reason, httpStatus=http_status) + + +def unexpected_error(log): + _emit(log, 1301, "unexpected_error", "Unexpected OTP request failure", level=logging.ERROR) + + +def response_prepared(log, http_status, contains_nonce, contains_correlation_id): + _emit(log, 1400, "response_prepared", "OTP response prepared", + httpStatus=http_status, containsNonce=contains_nonce, + containsCorrelationId=contains_correlation_id) + + +def request_completed(log, http_status, result, elapsed_ms): + _emit(log, 1401, "request_completed", "OTP request completed", + httpStatus=http_status, result=result, elapsedMs=elapsed_ms) + + +def credential_refresh_failed(log, cache_kind): + _emit(log, 1500, "credential_refresh_failed", "Credential refresh failed", + level=logging.WARNING, cacheKind=cache_kind, failureReason="credential_unavailable") diff --git a/python/src/provider.py b/python/src/provider.py new file mode 100644 index 0000000..43137b1 --- /dev/null +++ b/python/src/provider.py @@ -0,0 +1,212 @@ +from __future__ import annotations + +import logging +import time +from abc import ABC, abstractmethod +from dataclasses import dataclass +from urllib.parse import urlsplit, urlunsplit + +import requests +from urllib3.exceptions import ReadTimeoutError + +from . import otp_log +from .models import Outcome, OtpDelivery, ProviderResult + +DEFAULT_TIMEOUT_MS = 1500 + + +@dataclass(frozen=True) +class OutboundRequest: + method: str + url: str + headers: dict[str, str] + body: str + + +class ProviderSendError(Exception): + def __init__(self, status_code: int) -> None: + super().__init__("provider delivery failed") + self.status_code = status_code + + +class PhoneProviderBase(ABC): + name: str + authentication_mode: str + + @abstractmethod + def build_request(self, channel, endpoint, delivery, credential, env) -> OutboundRequest: + raise NotImplementedError + + @abstractmethod + def map_response(self, payload: object, http_status: int) -> ProviderResult: + raise NotImplementedError + + @property + @abstractmethod + def credential_spec(self) -> dict[str, str]: + raise NotImplementedError + + def send_otp(self, channel: str, endpoint: str, delivery: OtpDelivery, credential: dict, + env, timeout_ms: int, log) -> ProviderResult: + stage = "provider_request_build" + response = None + + def fail(status: int, reason: str, failure_stage: str | None = None): + current_stage = failure_stage or stage + otp_log.request_failed( + log, logging.ERROR if status >= 500 else logging.WARNING, + current_stage, reason, status) + raise ProviderSendError(status) + + try: + otp_log.provider_request_build_started(log) + request = self.build_request(channel, endpoint, delivery, credential, env) + if not is_https_endpoint(request.url): + fail(502, "invalid_provider_request_url") + otp_log.provider_request_built(log, normalize_http_method(request.method), sanitize_endpoint(request.url)) + + stage = "provider_transport" + started = time.monotonic() + otp_log.provider_request_started(log, timeout_ms) + response = requests.request( + request.method, + request.url, + headers=request.headers, + data=request.body, + timeout=timeout_ms / 1000, + allow_redirects=False, + stream=True, + ) + otp_log.provider_response_received(log, response.status_code) + try: + payload = response.json() + except ValueError: + otp_log.provider_response_invalid_json(log) + result = ProviderResult( + Outcome.FAIL, False, response.status_code, + failure_reason="invalid_provider_json") + else: + stage = "provider_response" + try: + result = self.map_response(payload, response.status_code) + except Exception: + fail(502, "response_parse_failed") + + elapsed_ms = int((time.monotonic() - started) * 1000) + status = to_endpoint_http_status(result) + provider_status = ( + result.provider_status_name or result.provider_status_code or "unmapped" + if result.status_recognized else "unmapped" + ) + otp_log.provider_response_processed( + log, + logging.ERROR if status >= 500 else logging.INFO if status == 200 else logging.WARNING, + result.provider_http_status, + provider_status, + result.outcome.value, + result.failure_reason, + elapsed_ms, + ) + return result + except requests.exceptions.RequestException as error: + if isinstance(error, requests.exceptions.Timeout) or ( + isinstance(error, requests.exceptions.ConnectionError) and has_read_timeout(error) + ): + fail(504, "provider_timeout", "provider_transport") + fail(502, "provider_network_error", "provider_transport") + except ProviderSendError: + raise + except Exception: + reason = "request_build_failed" if stage == "provider_request_build" else "provider_network_error" + fail(502, reason) + finally: + close = getattr(response, "close", None) + if callable(close): + try: + close() + except Exception: + pass + + +def classify_failure(http_status: int, outcome: Outcome, status_recognized: bool) -> str | None: + if not 200 <= http_status < 300: + return "provider_http_error" + if outcome != Outcome.FAIL: + return None + return "provider_rejected" if status_recognized else "unrecognized_provider_status" + + +def to_endpoint_http_status(result: ProviderResult) -> int: + if result.outcome == Outcome.CONTINUE: + return 200 + if result.outcome == Outcome.BLOCK: + return 403 + if result.provider_http_status == 429: + return 429 + if result.provider_http_status in (401, 403): + return 401 + if 400 <= result.provider_http_status < 500: + return 400 + return 502 + + +def is_https_endpoint(value: object) -> bool: + if not isinstance(value, str) or not value or "#" in value: + return False + if any(character.isspace() or ord(character) < 32 or ord(character) == 127 for character in value): + return False + try: + parsed = urlsplit(value) + port = parsed.port + return ( + parsed.scheme == "https" + and bool(parsed.hostname) + and port != 0 + and parsed.username is None + and parsed.password is None + and not parsed.fragment + and not parsed.netloc.endswith(":") + ) + except ValueError: + return False + + +def sanitize_endpoint(value: str) -> str: + parsed = urlsplit(value) + host = f"[{parsed.hostname}]" if ":" in parsed.hostname else parsed.hostname + authority = host if parsed.port in (None, 443) else f"{host}:{parsed.port}" + return urlunsplit((parsed.scheme, authority, parsed.path or "/", "", "")) + + +def normalize_http_method(method: object) -> str: + value = method.upper() if isinstance(method, str) else None + return value if value in {"GET", "HEAD", "POST", "PUT", "DELETE", "CONNECT", "OPTIONS", "TRACE", "PATCH"} else "other" + + +def provider_timeout_ms(value: object) -> int: + digits = value.strip() if isinstance(value, str) else "" + if not digits or not digits.isascii() or not digits.isdecimal(): + return DEFAULT_TIMEOUT_MS + digits = digits.lstrip("0") + if not digits: + return DEFAULT_TIMEOUT_MS + if len(digits) > 4 or (len(digits) == 4 and digits > "2500"): + return 2500 + return int(digits) + + +def has_read_timeout(error: Exception) -> bool: + pending = [error] + seen = set() + while pending: + current = pending.pop() + if id(current) in seen: + continue + seen.add(id(current)) + if isinstance(current, ReadTimeoutError): + return True + pending.extend( + nested for nested in (current.__cause__, current.__context__, *current.args) + if isinstance(nested, Exception) + ) + return False diff --git a/python/src/providers/infobip.py b/python/src/providers/infobip.py index eb8de22..b285bd7 100644 --- a/python/src/providers/infobip.py +++ b/python/src/providers/infobip.py @@ -1,55 +1,54 @@ import json -from ..models import ParsedResponse - - -class InfobipProvider: - manifest = { - "id": "infobip", - "auth": {"mode": "apiKey", "key_vault_secret_name": "infobip-api-key"}, - "response_mapping": { - "ACCEPTED": "Continue", - "PENDING": "Continue", - "DELIVERED": "Continue", - "REJECTED": "Fail", - "EXPIRED": "Fail", - "UNDELIVERABLE": "Fail", - "default": "Fail", - }, - } - - def build_request(self, channel, endpoint, dispatch, credential, env): +from ..models import Outcome, ProviderResult +from ..provider import OutboundRequest, PhoneProviderBase, classify_failure + + +class InfobipProvider(PhoneProviderBase): + name = "infobip" + authentication_mode = "apiKey" + + @property + def credential_spec(self): + return {"mode": self.authentication_mode, "key_vault_secret_name": "infobip-api-key"} + + def build_request(self, channel, endpoint, delivery, credential, env): sender_id = env.get("EPP_PROVIDER_ACCOUNT_NAME") or "Verify" authorization = f"App {credential['secret']}" headers = {"Authorization": authorization, "Content-Type": "application/json", "Accept": "application/json"} - message_id = dispatch.correlation_id or dispatch.message_id + message_id = delivery.correlation_id or delivery.message_id if channel == "voice": body = {"messages": [{ "from": sender_id, - "destinations": [{"to": dispatch.destination, "messageId": message_id}], - "text": dispatch.message, - "language": dispatch.locale or "en", + "destinations": [{"to": delivery.phone_number, "messageId": message_id}], + "text": delivery.message, + "language": delivery.locale or "en", "voice": {"name": "Joanna", "gender": "female"}, }]} - return {"url": f"{endpoint}/tts/3/advanced", "method": "POST", "headers": headers, "body": json.dumps(body)} + return OutboundRequest("POST", f"{endpoint}/tts/3/advanced", headers, json.dumps(body)) body = {"messages": [{ "sender": sender_id, - "destinations": [{"to": dispatch.destination, "messageId": message_id}], - "content": {"text": dispatch.message}, + "destinations": [{"to": delivery.phone_number, "messageId": message_id}], + "content": {"text": delivery.message}, }]} - return {"url": f"{endpoint}/sms/3/messages", "method": "POST", "headers": headers, "body": json.dumps(body)} + return OutboundRequest("POST", f"{endpoint}/sms/3/messages", headers, json.dumps(body)) - def parse_response(self, http_status, ok, json_body): - messages = json_body.get("messages") if isinstance(json_body, dict) else None + def map_response(self, payload, http_status): + messages = payload.get("messages") if isinstance(payload, dict) else None first_message = messages[0] if messages else {} status = first_message.get("status") or {} status_name = (status.get("groupName") or status.get("name") or "").upper() or None - return ParsedResponse( - success=ok, + recognized = status_name in {"ACCEPTED", "PENDING", "DELIVERED", "REJECTED", "EXPIRED", "UNDELIVERABLE"} + outcome = Outcome.CONTINUE if status_name in {"ACCEPTED", "PENDING", "DELIVERED"} else Outcome.FAIL + final_outcome = outcome if 200 <= http_status < 300 else Outcome.FAIL + return ProviderResult( + outcome=final_outcome, + status_recognized=recognized, provider_http_status=http_status, provider_message_id=first_message.get("messageId"), provider_status_name=status_name, provider_status_description=status.get("description"), + failure_reason=classify_failure(http_status, final_outcome, recognized), ) diff --git a/python/src/providers/sinch.py b/python/src/providers/sinch.py index cb6e3fa..d19ee1f 100644 --- a/python/src/providers/sinch.py +++ b/python/src/providers/sinch.py @@ -1,50 +1,66 @@ import json -from ..models import ParsedResponse +from ..models import Outcome, ProviderResult +from ..provider import OutboundRequest, PhoneProviderBase, classify_failure -class SinchProvider: - manifest = { - "id": "sinch", - "auth": {"mode": "apiKey", "key_vault_secret_name": "sinch-api-token"}, - "response_mapping": { - "Dispatched": "Continue", "Delivered": "Continue", "Queued": "Continue", - "Failed": "Fail", "Rejected": "Fail", "default": "Fail", - }, - } +class SinchProvider(PhoneProviderBase): + name = "sinch" + authentication_mode = "apiKey" - def build_request(self, channel, endpoint, dispatch, credential, env): - bearer = credential["secret"] - headers = {"Authorization": f"Bearer {bearer}", "Content-Type": "application/json", "Accept": "application/json"} - reference = dispatch.correlation_id or dispatch.message_id + @property + def credential_spec(self): + return {"mode": self.authentication_mode, "key_vault_secret_name": "sinch-api-token"} + + def build_request(self, channel, endpoint, delivery, credential, env): + headers = { + "Authorization": f"Bearer {credential['secret']}", + "Content-Type": "application/json", + "Accept": "application/json", + } + reference = delivery.correlation_id or delivery.message_id if channel == "voice": voice_base = env.get("SINCH_VOICE_ENDPOINT") or "https://calling.api.sinch.com" body = {"method": "ttsCallout", "ttsCallout": { - "destination": {"type": "number", "endpoint": dispatch.destination}, - "text": dispatch.message, - "locale": dispatch.locale or "en-US", + "destination": {"type": "number", "endpoint": delivery.phone_number}, + "text": delivery.message, + "locale": delivery.locale or "en-US", "custom": reference, }} - return {"url": f"{voice_base}/calling/v1/callouts", "method": "POST", "headers": headers, "body": json.dumps(body)} + return OutboundRequest("POST", f"{voice_base}/calling/v1/callouts", headers, json.dumps(body)) service_plan_id = env.get("SINCH_SERVICE_PLAN_ID") or "" body = { "from": env.get("EPP_PROVIDER_ACCOUNT_NAME") or "Verify", - "to": [dispatch.destination], - "body": dispatch.message, + "to": [delivery.phone_number], + "body": delivery.message, "client_reference": reference, } - return {"url": f"{endpoint}/xms/v1/{service_plan_id}/batches", "method": "POST", "headers": headers, "body": json.dumps(body)} - - def parse_response(self, http_status, ok, json_body): - identifier = None - if isinstance(json_body, dict): - identifier = json_body.get("id") or json_body.get("callId") - return ParsedResponse( - success=ok, - provider_http_status=http_status, - provider_message_id=str(identifier) if identifier is not None else None, - provider_status_name="Dispatched" if ok else None, - provider_status_description=json_body.get("text") if isinstance(json_body, dict) else None, + return OutboundRequest("POST", f"{endpoint}/xms/v1/{service_plan_id}/batches", headers, json.dumps(body)) + + def map_response(self, payload, http_status): + identifier = payload.get("id") or payload.get("callId") if isinstance(payload, dict) else None + status = payload.get("status") if isinstance(payload, dict) else None + description = payload.get("text") if isinstance(payload, dict) else None + successful = 200 <= http_status < 300 + if status is None and successful and isinstance(identifier, str) and identifier.strip(): + status = "Dispatched" + recognized = status in {"Dispatched", "Delivered", "Queued", "Failed", "Rejected"} + outcome = Outcome.CONTINUE if status in {"Dispatched", "Delivered", "Queued"} else Outcome.FAIL + has_message_id = isinstance(identifier, str) and bool(identifier.strip()) + final_outcome = outcome if successful and has_message_id else Outcome.FAIL + failure_reason = ( + "missing_provider_message_id" + if successful and not has_message_id + else classify_failure(http_status, final_outcome, recognized) + ) + return ProviderResult( + final_outcome, + recognized, + http_status, + identifier if isinstance(identifier, str) else None, + status if isinstance(status, str) else None, + provider_status_description=description if isinstance(description, str) else None, + failure_reason=failure_reason, ) diff --git a/python/src/providers/soprano.py b/python/src/providers/soprano.py index 5e874e6..084b4c7 100644 --- a/python/src/providers/soprano.py +++ b/python/src/providers/soprano.py @@ -1,7 +1,8 @@ import json import re -from ..models import ParsedResponse +from ..models import Outcome, ProviderResult +from ..provider import OutboundRequest, PhoneProviderBase, classify_failure DEFAULT_VOICE_LANGUAGE = "en-US" VOICE_GENDER = 1 @@ -23,18 +24,15 @@ def _build_text_to_voice(message, locale): } -class SopranoProvider: - manifest = { - "id": "soprano", - "auth": {"mode": "oauth"}, - "response_mapping": { - "ENROUTE": "Continue", "ACCEPTED": "Continue", "SUBMITTED": "Continue", - "SENT": "Continue", "DELIVERED": "Continue", "QUEUED": "Continue", - "FAILED": "Fail", "REJECTED": "Fail", "FILTERED": "Fail", "BLOCKED": "Block", "default": "Fail", - }, - } +class SopranoProvider(PhoneProviderBase): + name = "soprano" + authentication_mode = "oauth" + + @property + def credential_spec(self): + return {"mode": self.authentication_mode} - def build_request(self, channel, endpoint, dispatch, credential, env): + def build_request(self, channel, endpoint, delivery, credential, env): message_type = "voice" if channel == "voice" else "sms" headers = { "Authorization": f"Bearer {credential.get('access_token') or ''}", @@ -42,29 +40,44 @@ def build_request(self, channel, endpoint, dispatch, credential, env): "Accept": "application/json", } body = { - "destination": str(dispatch.destination).lstrip("+"), + "destination": str(delivery.phone_number).lstrip("+"), "messageTypes": [message_type], - "correlationId": dispatch.correlation_id or dispatch.message_id, + "correlationId": delivery.correlation_id or delivery.message_id, "shutterMode": False, } if channel == "voice": - body["voice"] = {"text2voice": _build_text_to_voice(dispatch.message, dispatch.locale)} + body["voice"] = {"text2voice": _build_text_to_voice(delivery.message, delivery.locale)} else: - body["text"] = dispatch.message - return {"url": endpoint, "method": "POST", "headers": headers, "body": json.dumps(body)} + body["text"] = delivery.message + return OutboundRequest("POST", endpoint, headers, json.dumps(body)) - def parse_response(self, http_status, ok, json_body): - payload = json_body[0] if isinstance(json_body, list) and json_body else json_body + def map_response(self, body, http_status): + payload = body[0] if isinstance(body, list) and body else body payload = payload if isinstance(payload, dict) else {} identifier = payload.get("id") - identifier = str(identifier) if identifier is not None else payload.get("messageId") + identifier = ( + str(identifier) + if isinstance(identifier, (str, int)) and not isinstance(identifier, bool) + else payload.get("messageId") + ) value = payload.get("status") if value is None: value = payload.get("state") status = value.upper() if isinstance(value, str) and value else "UNKNOWN" - return ParsedResponse( - success=ok, - provider_http_status=http_status, - provider_message_id=identifier, - provider_status_name=status, + continue_statuses = {"ENROUTE", "ACCEPTED", "SUBMITTED", "SENT", "DELIVERED", "QUEUED"} + fail_statuses = {"FAILED", "REJECTED", "FILTERED"} + recognized = status in continue_statuses | fail_statuses | {"BLOCKED"} + outcome = ( + Outcome.CONTINUE if status in continue_statuses + else Outcome.BLOCK if status == "BLOCKED" + else Outcome.FAIL + ) + final_outcome = outcome if 200 <= http_status < 300 else Outcome.FAIL + return ProviderResult( + final_outcome, + recognized, + http_status, + identifier if isinstance(identifier, str) else None, + status, + failure_reason=classify_failure(http_status, final_outcome, recognized), ) diff --git a/python/src/providers/telesign.py b/python/src/providers/telesign.py index e3c9fd2..b293eae 100644 --- a/python/src/providers/telesign.py +++ b/python/src/providers/telesign.py @@ -2,7 +2,8 @@ import json import re -from ..models import ParsedResponse +from ..models import Outcome, ProviderResult +from ..provider import OutboundRequest, PhoneProviderBase, classify_failure VOICE_PASSCODE_PATTERN = re.compile(r"(?= 500 else logging.INFO if http_status == 200 else logging.WARNING) - - def failure(self, stage, reason, http_status): - self._credential_resolution_finished() - self.provider_request_finished() - self.data["failureStage"] = stage - self.data["failureReason"] = reason - self.service(f"{stage}_failed", {"failureReason": reason, "httpStatus": http_status}, - logging.ERROR if http_status >= 500 else logging.WARNING) - - def response_prepared(self, http_status, contains_nonce, contains_correlation_id): - self.data["responseContainsNonce"] = contains_nonce - self.data["responseContainsCorrelationId"] = contains_correlation_id - self.service("response_prepared", { - "httpStatus": http_status, - "responseContainsNonce": contains_nonce, - "responseContainsCorrelationId": contains_correlation_id, - }) - - def complete(self, http_status): - self._credential_resolution_finished() - self.provider_request_finished() - logging.info("%s", json.dumps({ - "logType": "request", "eventName": "request_completed", - **self.data, - "httpStatus": http_status, - "result": ("evaluated" if self.data["evaluation"] else "accepted") if http_status == 200 else "failed", - "elapsedMs": int((time.monotonic() - self.started) * 1000), - })) diff --git a/python/tests/test_contract.py b/python/tests/test_contract.py index 3e903f8..35988e9 100644 --- a/python/tests/test_contract.py +++ b/python/tests/test_contract.py @@ -4,8 +4,7 @@ import pytest -from src.dispatch import DispatchRequest, ProviderRegistry, context_to_dispatch, parse_envelope, resolve_outcome -from src.models import DeliveryContext, Envelope, ParsedResponse, TextToVoice +from src.models import EntraSendOtpPayload, Outcome, OtpDelivery, ProviderResult from src.providers.infobip import InfobipProvider from src.providers.sinch import SinchProvider from src.providers.soprano import SopranoProvider @@ -14,198 +13,164 @@ MESSAGE = " Use 918273; then 1234.\nDo not rewrite + or café. " -def _dispatch(channel="sms"): - return DispatchRequest("+15551234567", MESSAGE, channel, "message-id", "correlation-id", "en-US", - TextToVoice("Your code is", "001234", "en-US") if channel == "voice" else None) +def _delivery(channel="sms"): + return OtpDelivery( + "+15551234567", MESSAGE, channel, "message-id", "correlation-id", "en-US") @pytest.mark.parametrize("channel", ["sms", "voice"]) -def test_soprano_selected_endpoint_and_oauth_contract(channel): - dispatch = _dispatch(channel) - if channel == "voice": - dispatch.locale = "fr-FR" - dispatch.text_to_voice = TextToVoice("ignored", "001234", "override") - request = ProviderRegistry([SopranoProvider()]).get("SOPRANO").build_request( - channel, "https://qa4.example/oauth/messages", dispatch, +def test_soprano_owns_request_and_response_contract(channel): + delivery = _delivery(channel) + request = SopranoProvider().build_request( + channel, + "https://qa4.example/oauth/messages", + delivery, {"mode": "oauth", "access_token": "provider-token"}, {}, ) - assert request["url"] == "https://qa4.example/oauth/messages" and request["method"] == "POST" - assert request["headers"] == { - "Authorization": "Bear" + "er provider-token", - "Content-Type": "application/json", "Accept": "application/json", - } - expected = { - "destination": "15551234567", "messageTypes": [channel], - "correlationId": "correlation-id", "shutterMode": False, - } + assert request.url == "https://qa4.example/oauth/messages" + assert request.method == "POST" + assert request.headers["Authorization"] == "Bearer provider-token" + body = json.loads(request.body) + assert body["destination"] == "15551234567" + assert body["messageTypes"] == [channel] + assert body["correlationId"] == "correlation-id" if channel == "voice": - expected["voice"] = {"text2voice": { + assert body["voice"]["text2voice"] == { "beforePasswordText": " Use ", "password": "918273", "afterPasswordText": "; then 1234.\nDo not rewrite + or café. ", - "language": "fr-FR", + "language": "en-US", "gender": 1, "loop": 2, - }} + } + assert "text" not in body else: - expected["text"] = MESSAGE - assert json.loads(request["body"]) == expected - response = SopranoProvider().parse_response(201, True, {"id": 123, "status": "ENROUTE"}) - assert response == ParsedResponse(True, 201, provider_message_id="123", provider_status_name="ENROUTE") - assert "ENROUTE" not in repr(response) + assert body["text"] == MESSAGE - -@pytest.mark.parametrize("locale", [None, "", " ", {"untrusted": True}]) -def test_soprano_voice_defaults_language_without_valid_locale(locale): - dispatch = _dispatch("voice") - dispatch.locale = locale - request = SopranoProvider().build_request( - "voice", "https://qa4.example/oauth/messages", dispatch, - {"mode": "oauth", "access_token": "provider-token"}, {}, - ) - assert json.loads(request["body"])["voice"]["text2voice"]["language"] == "en-US" + result = SopranoProvider().map_response({"id": 123, "status": "ENROUTE"}, 201) + assert result == ProviderResult( + Outcome.CONTINUE, True, 201, "123", "ENROUTE") -def test_soprano_voice_requires_six_digit_passcode(): - dispatch = _dispatch("voice") - dispatch.message = "Your code is unavailable." - with pytest.raises(ValueError, match="six-digit passcode"): - SopranoProvider().build_request( - "voice", "https://qa4.example/oauth/messages", dispatch, - {"mode": "oauth", "access_token": "provider-token"}, {}, - ) +@pytest.mark.parametrize("body,expected_status,recognized", [ + ({"status": "accepted", "id": 1}, "ACCEPTED", True), + ([{"state": "queued", "messageId": "m"}], "QUEUED", True), + ({"status": False, "state": "ACCEPTED"}, "UNKNOWN", False), + ({}, "UNKNOWN", False), +]) +def test_soprano_preserves_protocol_specific_mixed_response_shapes(body, expected_status, recognized): + result = SopranoProvider().map_response(body, 200) + assert result.provider_status_name == expected_status + assert result.status_recognized is recognized + assert result.failure_reason == (None if recognized else "unrecognized_provider_status") -def test_infobip_sms_request_and_response_contract(): +def test_infobip_contract_is_strict_and_provider_owned(): request = InfobipProvider().build_request( - "sms", "https://infobip.example", _dispatch(), - {"mode": "apiKey", "secret": "ib"}, {"EPP_PROVIDER_ACCOUNT_NAME": "EPP"}, - ) - assert request["method"] == "POST" and request["url"] == "https://infobip.example/sms/3/messages" - assert request["headers"]["Authorization"] == "App ib" - assert json.loads(request["body"])["messages"] == [{ - "sender": "EPP", "destinations": [{"to": "+15551234567", "messageId": "correlation-id"}], + "sms", "https://infobip.example", _delivery(), + {"mode": "apiKey", "secret": "ib"}, {"EPP_PROVIDER_ACCOUNT_NAME": "EPP"}) + assert request.url == "https://infobip.example/sms/3/messages" + assert request.headers["Authorization"] == "App ib" + assert json.loads(request.body)["messages"] == [{ + "sender": "EPP", + "destinations": [{"to": "+15551234567", "messageId": "correlation-id"}], "content": {"text": MESSAGE}, }] - response = InfobipProvider().parse_response(200, True, { + result = InfobipProvider().map_response({ "messages": [{"messageId": "message-id", "status": {"groupName": "PENDING"}}], - }) - assert response == ParsedResponse(True, 200, provider_message_id="message-id", provider_status_name="PENDING") + }, 200) + assert result == ProviderResult( + Outcome.CONTINUE, True, 200, "message-id", "PENDING") -@pytest.mark.parametrize("channel,locale", [ - ("sms", "en"), ("voice", "en"), ("sms", None), ("sms", ""), ("sms", {"untrusted": True}), -]) -def test_telesign_epp_request_contract(channel, locale): - dispatch = _dispatch(channel) - dispatch.locale = locale - request = TelesignProvider().build_request( - channel, f"https://verify.telesign.com/epp/{channel}", dispatch, - {"mode": "apiKey", "secret": "key", "identity": "customer"}, {}, - ) - assert request["method"] == "POST" and request["url"] == f"https://verify.telesign.com/epp/{channel}" - assert request["headers"] == {"Authorization": "Basic " + base64.b64encode(b"customer:key").decode(), - "Content-Type": "application/json", "Accept": "application/json"} - expected_text = ( - " Use 9, 1, 8, 2, 7, 3; then 1234.\nDo not rewrite + or café. " - " Use 9, 1, 8, 2, 7, 3; then 1234.\nDo not rewrite + or café. " - if channel == "voice" else MESSAGE - ) - assert json.loads(request["body"]) == { - "recipient": {"phone_number": "+15551234567"}, - "message": {"text": expected_text, "language": "en"} if locale == "en" else {"text": expected_text}, - "channels": [{"channel": channel}], "correlation_id": "correlation-id", - } +def test_sinch_requires_a_provider_message_id_for_success(): + provider = SinchProvider() + accepted = provider.map_response({"id": "message-id"}, 200) + assert accepted.outcome == Outcome.CONTINUE + assert accepted.provider_status_name == "Dispatched" + missing = provider.map_response({}, 200) + assert missing.outcome == Outcome.FAIL + assert missing.failure_reason == "missing_provider_message_id" -def test_telesign_voice_paces_only_six_digit_numeric_runs_and_repeats_message(): - dispatch = _dispatch("voice") - dispatch.message = "Code 001234; ref 1234567; alternate 654321." +@pytest.mark.parametrize("channel", ["sms", "voice"]) +def test_telesign_request_contract(channel): request = TelesignProvider().build_request( - "voice", "https://verify.telesign.com/epp/voice", dispatch, - {"mode": "apiKey", "secret": "key", "identity": "customer"}, {}, - ) - assert json.loads(request["body"])["message"]["text"] == ( - "Code 0, 0, 1, 2, 3, 4; ref 1234567; alternate 6, 5, 4, 3, 2, 1. " - "Code 0, 0, 1, 2, 3, 4; ref 1234567; alternate 6, 5, 4, 3, 2, 1." + channel, + f"https://verify.telesign.com/epp/{channel}", + _delivery(channel), + {"mode": "apiKey", "secret": "key", "identity": "customer"}, + {}, ) + assert request.headers["Authorization"] == ( + "Basic " + base64.b64encode(b"customer:key").decode()) + body = json.loads(request.body) + assert body["recipient"] == {"phone_number": "+15551234567"} + assert body["channels"] == [{"channel": channel}] + assert body["correlation_id"] == "correlation-id" + if channel == "voice": + assert "9, 1, 8, 2, 7, 3" in body["message"]["text"] + + +def test_telesign_enforces_integer_status_codes(): + provider = TelesignProvider() + for payload in ({}, {"status": {"code": True}}, {"status": {"code": "290"}}): + result = provider.map_response(payload, 200) + assert result.outcome == Outcome.FAIL + assert result.status_recognized is False + assert result.failure_reason == "unrecognized_provider_status" + accepted = provider.map_response( + {"reference_id": "message-id", "status": {"code": 290}}, 200) + assert accepted == ProviderResult( + Outcome.CONTINUE, True, 200, "message-id", + provider_status_code="290") -def test_telesign_epp_validates_recipients_and_status(): - adapter = TelesignProvider() - response = adapter.parse_response(200, True, {"reference_id": "message-id", "status": {"code": 290}}) - assert response == ParsedResponse(True, 200, provider_message_id="message-id", provider_status_code="290") +def test_http_errors_override_provider_acceptance(): + result = SopranoProvider().map_response({"status": "ACCEPTED"}, 500) + assert result.outcome == Outcome.FAIL + assert result.failure_reason == "provider_http_error" + + +def test_telesign_validates_recipient_and_channel(): + provider = TelesignProvider() credential = {"identity": "customer", "secret": "key"} - for destination in ("15551234567", "+0123", "+1", "+1234567890123456", "+123\n", "+123\r", "+12 34", None): - dispatch = _dispatch() - dispatch.destination = destination + for destination in ("15551234567", "+0123", "+1", "+1234567890123456", None): + delivery = _delivery() + delivery = OtpDelivery( + destination, delivery.message, delivery.channel, delivery.message_id, + delivery.correlation_id, delivery.locale) with pytest.raises(ValueError, match="invalid recipient"): - adapter.build_request("sms", "https://verify.telesign.com", dispatch, credential, {}) + provider.build_request("sms", "https://verify.telesign.com", delivery, credential, {}) with pytest.raises(ValueError, match="unsupported channel"): - adapter.build_request("email", "https://verify.telesign.com", _dispatch(), credential, {}) - dispatch = _dispatch() - for correlation_id in (None, "", 123, True, [], {"invalid": True}): - dispatch.correlation_id = correlation_id - request = adapter.build_request("sms", "https://verify.telesign.com", dispatch, credential, {}) - assert json.loads(request["body"])["correlation_id"] == dispatch.message_id - for payload in (None, {}, {"status": []}, {"status": {"code": True}}, {"status": {"code": "290"}}, {"status": {"code": 999}}): - assert resolve_outcome(adapter.manifest, adapter.parse_response(200, True, payload)) == "Fail" - for code, ok, outcome in ((290, True, "Continue"), (100, True, "Continue"), (290, False, "Fail"), - (3001, True, "Continue"), (3001, False, "Fail")): - parsed = adapter.parse_response(200 if ok else 500, ok, {"status": {"code": code, "description": "status detail"}}) - assert parsed.provider_status_description == "status detail" - assert resolve_outcome(adapter.manifest, parsed) == outcome - - -def test_sinch_sms_request_and_response_contract(): - request = SinchProvider().build_request( - "sms", "https://sinch.example", _dispatch(), - {"mode": "apiKey", "secret": "static-api-token"}, - {"SINCH_SERVICE_PLAN_ID": "plan", "EPP_PROVIDER_ACCOUNT_NAME": "EPP"}, - ) - assert request["method"] == "POST" and request["url"] == "https://sinch.example/xms/v1/plan/batches" - assert request["headers"]["Authorization"] == "Bearer static-api-token" - assert json.loads(request["body"]) == { - "from": "EPP", "to": ["+15551234567"], "body": MESSAGE, "client_reference": "correlation-id", + provider.build_request("email", "https://verify.telesign.com", _delivery(), credential, {}) + + +def test_typed_payload_accepts_contract_values_and_rejects_bad_requests(): + valid = { + "type": "microsoft.mfa.otpDeliver.v1", + "channel": 1, + "mode": 1, + "encryptedDeliveryContext": "jwe", } - response = SinchProvider().parse_response(200, True, {"id": "message-id"}) - assert response == ParsedResponse(True, 200, provider_message_id="message-id", provider_status_name="Dispatched") - - -def test_request_models_preserve_content_and_accept_valid_routing_and_ttl(): - payload = {"type": "microsoft.mfa.otpDeliver.v1", "channel": 1, "mode": 1, "encryptedDeliveryContext": "jwe"} - for channel, mode, expected in ((1, 1, (1, 1)), ("VOICE", "Evaluation", (2, 2))): - envelope, error = parse_envelope({**payload, "channel": channel, "mode": mode}) - assert error is None and isinstance(envelope, Envelope) - assert (envelope.channel, envelope.mode) == expected - for ttl in (1, 2147483647): - envelope, error = parse_envelope({**payload, "ttlSeconds": ttl}) - assert error is None and envelope.ttl_seconds == ttl - envelope, error = parse_envelope(payload) - assert error is None and envelope.ttl_seconds is None - - context = DeliveryContext.from_payload({ - "nonce": " nonce ", "phoneNumber": "+15551234567", "message": MESSAGE, - "locale": {"opaque": "metadata"}, - }) - assert isinstance(context, DeliveryContext) and context.is_complete - dispatch = context_to_dispatch(context, envelope, "message-id") - assert isinstance(dispatch, DispatchRequest) - assert context.nonce == " nonce " and dispatch.message == MESSAGE - assert dispatch.destination == context.phone_number and dispatch.locale is context.locale - assert MESSAGE not in repr(context) + repr(dispatch) - assert "encrypted_delivery_context" not in repr(envelope) - assert DeliveryContext.from_payload(None) is None - - -def test_envelope_parser_rejects_invalid_inputs_with_the_contract_reason(): - fixtures = json.loads((Path(__file__).resolve().parents[2] / "tests/fixtures/contract.json").read_text()) - valid = {"type": "microsoft.mfa.otpDeliver.v1", "channel": 1, "mode": 1, "encryptedDeliveryContext": "jwe"} + for channel, mode, expected in ((1, 1, ("sms", False)), ("VOICE", "Evaluation", ("voice", True))): + payload, error = EntraSendOtpPayload.from_payload( + {**valid, "channel": channel, "mode": mode}) + assert error is None + assert (payload.channel_name, payload.is_evaluation) == expected + + fixtures = json.loads( + (Path(__file__).resolve().parents[2] / "tests/fixtures/contract.json") + .read_text(encoding="utf-8")) for fixture in fixtures["badRequests"]: - # Malformed JSON is handled before the parser receives an object. if fixture["reason"] == "invalid JSON body": continue - payload = json.loads(fixture["rawBody"]) if "rawBody" in fixture else {**valid, **fixture["overrides"]} - envelope, error = parse_envelope(payload) - assert envelope is None and error == fixture["reason"], fixture["name"] + value = ( + json.loads(fixture["rawBody"]) + if "rawBody" in fixture + else {**valid, **fixture["overrides"]} + ) + payload, error = EntraSendOtpPayload.from_payload(value) + assert payload is None + assert error == fixture["reason"] diff --git a/python/tests/test_credential_cache.py b/python/tests/test_credential_cache.py index d7a7960..467852b 100644 --- a/python/tests/test_credential_cache.py +++ b/python/tests/test_credential_cache.py @@ -12,8 +12,7 @@ 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.credentials import ApiKeyCache, AccessTokenCache, CredentialTokenService from src.providers.telesign import TelesignProvider AUTH = {"mode": "apiKey", "key_vault_secret_name": "key", "identity_key_vault_secret_name": "id"} @@ -52,7 +51,7 @@ def read(name): 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) + manager = CredentialTokenService(secrets, cache_options=clock.options) try: with ThreadPoolExecutor(max_workers=10) as pool: pending = [pool.submit(manager.resolve, AUTH, CONFIG) for _ in range(10)] @@ -90,7 +89,7 @@ def read(name): return name secrets = Mock(resolve=Mock(side_effect=read)) - manager = ProviderCredentials(secrets, cache_options=clock.options, report_failure=failures.append) + manager = CredentialTokenService(secrets, cache_options=clock.options, report_failure=failures.append) try: manager.resolve(AUTH, CONFIG) fail = True @@ -141,7 +140,7 @@ def read(name): assert release.wait(3) return "PRIVATE-IDENTITY" - manager = ProviderCredentials(Mock(resolve=read), cache_options={**clock.options, "wait_timeout": 0.02}, + manager = CredentialTokenService(Mock(resolve=read), cache_options={**clock.options, "wait_timeout": 0.02}, report_failure=lambda _: None) try: for _ in range(3): @@ -181,7 +180,7 @@ def get_token_info(_): 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) + manager = CredentialTokenService(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: @@ -209,39 +208,27 @@ def get_token_info(_): manager.close() -def test_startup_only_prepares_credentials_and_shutdown_is_terminal(): +def test_selected_provider_credentials_are_cached_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"}) + manager = CredentialTokenService(secrets) + provider = TelesignProvider() try: - engine.start_credential_refresh() - engine.start_credential_refresh() + manager.get_credentials(provider, CONFIG) + manager.get_credentials(provider, CONFIG) assert secrets.resolve.call_count == 2 finally: - engine.close() + manager.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() + manager.get_credentials(provider, CONFIG) def test_configuration_changes_require_a_new_worker_and_stopped_cache_cannot_restart(): secrets = Mock(resolve=Mock(return_value="key")) - manager = ProviderCredentials(secrets) + manager = CredentialTokenService(secrets) manager.resolve(AUTH, CONFIG) manager.close() other = read_config({"KEY_VAULT_URL": "https://other.vault.azure.net"}) - manager = ProviderCredentials(secrets) + manager = CredentialTokenService(secrets) manager.resolve(AUTH, other) assert secrets.resolve.call_count == 4 manager.close() @@ -256,7 +243,7 @@ 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 = CredentialTokenService(secrets) manager.resolve(AUTH, CONFIG) wait_until(lambda: secrets.resolve.call_count >= 4) manager.close() @@ -268,7 +255,7 @@ def test_periodic_refresh_uses_the_selected_cache_and_stops(monkeypatch): def test_unknown_auth_mode_does_not_create_a_cache(): secrets = Mock() - manager = ProviderCredentials(secrets, report_failure=lambda _: None) + manager = CredentialTokenService(secrets, report_failure=lambda _: None) try: with pytest.raises(ValueError, match="unavailable"): manager.resolve({"mode": "unknown"}, CONFIG) @@ -284,7 +271,7 @@ def test_pending_secret_reads_do_not_block_process_shutdown(): from threading import Event from types import SimpleNamespace from src.config import read_config - from src.credentials import ProviderCredentials + from src.credentials import CredentialTokenService started, blocked = Event(), Event() def resolve(name): @@ -292,7 +279,7 @@ def resolve(name): blocked.wait() return "synthetic-key" - manager = ProviderCredentials(SimpleNamespace(resolve=resolve), cache_options={"wait_timeout": 0.02}) + manager = CredentialTokenService(SimpleNamespace(resolve=resolve), cache_options={"wait_timeout": 0.02}) atexit.register(lambda: print("shutdown-complete", flush=True)) atexit.register(manager.close) try: diff --git a/python/tests/test_credential_sdk.py b/python/tests/test_credential_sdk.py index ed990dc..8f15d6c 100644 --- a/python/tests/test_credential_sdk.py +++ b/python/tests/test_credential_sdk.py @@ -8,7 +8,7 @@ import pytest import src.credentials as credentials_module from src.config import read_config -from src.credentials import ProviderCredentials +from src.credentials import CredentialTokenService @pytest.mark.parametrize("refresh_in", [None, 60]) @@ -50,7 +50,7 @@ def managed(*args, **kwargs): return_value=Mock(spec=["get_token"], get_token=Mock(side_effect=managed)))) secrets = Mock() now = time.time() - manager = ProviderCredentials(secrets, cache_options={ + manager = CredentialTokenService(secrets, cache_options={ "clock": lambda: now, }) config = read_config({ diff --git a/python/tests/test_engine.py b/python/tests/test_engine.py index 58b8a71..45531c3 100644 --- a/python/tests/test_engine.py +++ b/python/tests/test_engine.py @@ -1,258 +1,126 @@ -import json -import io import logging -import time -from concurrent.futures import ThreadPoolExecutor -from threading import Event -from types import SimpleNamespace from unittest.mock import Mock import pytest 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 -from src.providers.sinch import SinchProvider +import src.provider as provider_module +from src.models import Outcome, OtpDelivery +from src.otp_log import RequestLogger +from src.provider import ( + ProviderSendError, + is_https_endpoint, + provider_timeout_ms, + to_endpoint_http_status, +) from src.providers.soprano import SopranoProvider -def _request(channel="sms"): - return DispatchRequest("+15551234567", "Your code is 918273", channel, "message", "correlation", "en-US") +def _delivery(channel="sms"): + return OtpDelivery( + "+15551234567", "Your code is 918273", channel, + "message", "correlation", "en-US") -@pytest.fixture -def engine(monkeypatch): - registry = ProviderRegistry([SopranoProvider(), SinchProvider()]) - monkeypatch.setattr(dispatch_module.requests, "request", Mock()) - result = DispatchEngine(registry, Mock(resolve=Mock(return_value="test-key")), { - "EPP_PROVIDER_NAME": " SOPRANO ", - "EPP_PROVIDER_ENDPOINT": "https://qa4.example/oauth/messages", - "EPP_PROVIDER_AUTH_MODE": "oauth", - "EPP_PROVIDER_CHANNEL": "sms", +def _log(): + return RequestLogger(logging.getLogger("test.provider"), { + "function_request_id": "request", }) - result._resolve_credential = Mock(return_value={"mode": "oauth", "access_token": "provider-token"}) - yield result - result.close() -def test_missing_oauth_configuration_never_sends(engine): - engine._resolve_credential = DispatchEngine._resolve_credential.__get__(engine, DispatchEngine) - status, body = engine.dispatch(_request(), "r") - assert status == 502 and body["reason"] == "provider credential unavailable" - dispatch_module.requests.request.assert_not_called() - - -def _oauth_settings(): - return {"EPP_PROVIDER_NAME": "soprano", "EPP_PROVIDER_AUTH_MODE": "oauth", "EPP_PROVIDER_CHANNEL": "sms", - "EPP_PROVIDER_ENDPOINT": "https://provider.example/full/sms/url/", - "EPP_PROVIDER_TENANT_ID": "provider-tenant", "EPP_PROVIDER_SCOPE": "api://provider/.default", - "EPP_OUTBOUND_CLIENT_ID": "calling-app", "EPP_OUTBOUND_MI_CLIENT_ID": "outbound-identity", - "AZURE_CLIENT_ID": "different-vault-identity"} - - -def test_soprano_oauth_uses_setup_settings_and_rejects_unusable_tokens(engine, monkeypatch): - engine.env = _oauth_settings() - 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(spec=["get_token"], get_token=Mock(side_effect=lambda *args, **kwargs: assertion)) - identity_factory = Mock(return_value=managed) - clients = [] - - def create_client(**kwargs): - assert kwargs["tenant_id"] == engine.env["EPP_PROVIDER_TENANT_ID"] - assert kwargs["client_id"] == engine.env["EPP_OUTBOUND_CLIENT_ID"] - assert kwargs["retry_total"] == 0 and kwargs["connection_timeout"] == kwargs["read_timeout"] == 2.5 - assert kwargs["logging_enable"] is False - - def get_token(*args, **options): - assert args == (engine.env["EPP_PROVIDER_SCOPE"],) and options == {"logging_enable": False} - assert kwargs["func"]() == "private-assertion" - return access - - client = Mock(spec=["get_token"], get_token=Mock(side_effect=get_token)) - clients.append(client) - return 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 - assert engine.dispatch(_request(), "request")[0] == 200 - assert len(clients) == 1 - sent = dispatch_module.requests.request.call_args - assert sent.args[1] == engine.env["EPP_PROVIDER_ENDPOINT"] - assert sent.kwargs["headers"] == {"Content-Type": "application/json", "Accept": "application/json", - "Authorization": "Bearer private-token"} - identity_factory.assert_called_once_with(client_id="outbound-identity", retry_total=0, - 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() - for stage in ("access", "assertion"): - for invalid in (None, SimpleNamespace(token=""), SimpleNamespace(token=" "), - SimpleNamespace(token="private-token"), - SimpleNamespace(token="private-token", expires_on=time.time() + 5)): - if stage == "access": - access = invalid - else: - access = SimpleNamespace(token="private-token", expires_on=time.time() + 3600) - assertion = invalid - 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): - engine.env = _oauth_settings() - engine._resolve_credential = DispatchEngine._resolve_credential.__get__(engine, DispatchEngine) - started, release = Event(), Event() - logger = logging.getLogger("azure.identity.test_setup_oauth") - output = io.StringIO() - handler = logging.StreamHandler(output) - logger.addHandler(handler) - - def fail(*args, **kwargs): - logger.warning("PRIVATE-SDK-TOKEN") - started.set() - assert release.wait(5) - logger.warning("PRIVATE-ACCOUNT-ERROR") - raise RuntimeError("PRIVATE-TOKEN-EXCEPTION") - - 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") - try: - assert started.wait(5) - logger.warning("unrelated request") - finally: - release.set() - status, body = pending.result(timeout=5) - assert status == 502 and body["reason"] == "provider credential unavailable" - logger.warning("after acquisition") - assert "PRIVATE" not in output.getvalue() + caplog.text + json.dumps(body) - assert "unrelated request" in output.getvalue() and "after acquisition" in output.getvalue() - engine.secrets.resolve.assert_not_called() - dispatch_module.requests.request.assert_not_called() - finally: - logger.removeHandler(handler) - - -def test_soprano_voice_payload_uses_oauth(engine): - engine.env["EPP_PROVIDER_CHANNEL"] = "voice" - speech = {"beforePasswordText": "Your code is", "password": "001234", "language": "en-US"} - context = DeliveryContext.from_payload({"nonce": "n", "phoneNumber": "+15551234567", - "message": "Your code is 001234", "locale": "fr-FR", - "textToVoice": speech}) - envelope = Envelope("microsoft.mfa.otpDeliver.v1", "tenant", "correlation", 2, 1, None, "encrypted") - request = dispatch_module.context_to_dispatch(context, envelope, "message") - dispatch_module.requests.request.return_value = Mock(status_code=200, json=Mock(return_value={"status": "ACCEPTED"})) - status, body = engine.dispatch(request, "r") - assert status == 200 and body["outcome"] == "Continue" - sent = dispatch_module.requests.request.call_args.kwargs - payload = json.loads(sent["data"]) - assert payload["voice"] == {"text2voice": { - "beforePasswordText": "Your code is ", - "password": "001234", - "afterPasswordText": "", - "language": "fr-FR", - "gender": 1, - "loop": 2, - }} - assert payload["messageTypes"] == ["voice"] and payload["destination"] == "15551234567" - assert "text" not in payload - assert sent["headers"]["Authorization"] == "Bear" + "er provider-token" - assert "001234" not in repr(request.text_to_voice) - - -def test_soprano_voice_without_six_digit_passcode_never_sends(engine): - engine.env["EPP_PROVIDER_CHANNEL"] = "voice" - request = _request("voice") - request.message = "Your code is unavailable." - status, body = engine.dispatch(request, "r") - assert status == 502 and body["reason"] == "provider request failed" - engine.secrets.resolve.assert_not_called() - dispatch_module.requests.request.assert_not_called() - - -def test_base_and_sinch_voice_final_url_guards(engine): - for url in ("http://api.example", "https://api.example:0"): - engine.env["EPP_PROVIDER_ENDPOINT"] = url - status, body = engine.dispatch(_request(), "r") - assert status == 502 and body["reason"] == "invalid provider endpoint" - engine.env["EPP_PROVIDER_ENDPOINT"] = "https://api.example" - engine.env["EPP_PROVIDER_NAME"] = "sinch" - engine.env.pop("EPP_PROVIDER_CHANNEL", None) - engine.env.pop("EPP_PROVIDER_AUTH_MODE", None) - engine._resolve_credential = Mock(return_value={"mode": "apiKey", "secret": "test-key", "identity": ""}) - for url in ("http://voice.example", "https://voice.example:0"): - engine.env["SINCH_VOICE_ENDPOINT"] = url - status, body = engine.dispatch(_request("voice"), "r") - assert status == 502 and body["reason"] == "invalid provider request URL" - dispatch_module.requests.request.assert_not_called() - - -def test_provider_outcomes_fail_closed(engine, monkeypatch): - monkeypatch.setenv("EPP_PROVIDER_NAME", "sinch") # The injected provider setting must win. - engine.env["EPP_DECRYPTION_KEY_PEM"] = "test-private-pem" - config = read_config(engine.env) - assert isinstance(config, AppConfig) and config.provider_name == "soprano" - assert config.env is engine.env and config.decryption_key_pem == "test-private-pem" - assert "test-private-pem" not in repr(config) and "EPP_PROVIDER_NAME" not in repr(config) - assert engine.registry.get(None) is None - cases = ( - (202, {"state": "accepted"}, 200, "Continue"), - (500, {"status": "ACCEPTED"}, 502, "Fail"), - (200, {"status": "FAILED"}, 502, "Fail"), - (200, {"status": "FILTERED"}, 502, "Fail"), - (200, {}, 502, "Fail"), - (200, {"status": False, "state": "ACCEPTED"}, 502, "Fail"), - (200, {"status": "BLOCKED"}, 403, "Block"), +def _send(provider, monkeypatch, response): + monkeypatch.setattr(provider_module.requests, "request", Mock(return_value=response)) + return provider.send_otp( + "sms", + "https://provider.example/messages", + _delivery(), + {"mode": "oauth", "access_token": "provider-token"}, + {}, + 1500, + _log(), ) - for upstream_status, payload, expected, outcome in cases: - response = Mock(status_code=upstream_status, json=Mock(return_value=payload)) - send = Mock(return_value=response) - monkeypatch.setattr(dispatch_module.requests, "request", send) - status, body = engine.dispatch(_request(), "r") - assert (status, body["outcome"], body["provider"]) == (expected, outcome, "soprano") - send.assert_called_once() - response.close.assert_called_once() -def test_transport_failures_and_wrapped_read_timeout(engine, monkeypatch): - errors = dispatch_module.requests.exceptions - for error, expected in ((errors.Timeout("offline"), 504), (errors.ConnectionError("offline"), 502)): - send = Mock(side_effect=error) - monkeypatch.setattr(dispatch_module.requests, "request", send) - status, body = engine.dispatch(_request(), "r") - assert status == expected and body["outcome"] == "Fail" - send.assert_called_once() +@pytest.mark.parametrize("value,valid", [ + ("https://api.example/path", True), + ("https://api.example:443/path?key=private", True), + ("http://api.example", False), + ("https://api.example:0", False), + ("https://user:password@api.example", False), + ("https://api.example/path#fragment", False), +]) +def test_https_endpoint_guard(value, valid): + assert is_https_endpoint(value) is valid + + +@pytest.mark.parametrize("value,expected", [ + (None, 1500), ("", 1500), ("1500", 1500), ("0001", 1), + ("2501", 2500), ("999999999999999999", 2500), ("1.5", 1500), +]) +def test_provider_timeout_normalization(value, expected): + assert provider_timeout_ms(value) == expected + + +def test_shared_transport_maps_response_and_closes_it(monkeypatch): + response = Mock(status_code=201, json=Mock(return_value={"status": "ENROUTE", "id": 123})) + result = _send(SopranoProvider(), monkeypatch, response) + assert result.outcome == Outcome.CONTINUE + assert result.provider_message_id == "123" + request = provider_module.requests.request.call_args + assert request.args[:2] == ("POST", "https://provider.example/messages") + assert request.kwargs["allow_redirects"] is False + assert request.kwargs["stream"] is True + response.close.assert_called_once() + - # requests can wrap a streamed body-read timeout in ConnectionError. - wrapped = errors.ConnectionError(ReadTimeoutError(None, "https://provider.example", "offline")) +@pytest.mark.parametrize("upstream,payload,expected_status,reason", [ + (500, {"status": "ACCEPTED"}, 502, "provider_http_error"), + (200, {"status": "FAILED"}, 502, "provider_rejected"), + (200, {}, 502, "unrecognized_provider_status"), + (200, {"status": "BLOCKED"}, 403, None), +]) +def test_provider_outcomes_fail_closed(monkeypatch, upstream, payload, expected_status, reason): + response = Mock(status_code=upstream, json=Mock(return_value=payload)) + result = _send(SopranoProvider(), monkeypatch, response) + assert to_endpoint_http_status(result) == expected_status + assert result.failure_reason == reason + + +def test_invalid_json_is_a_classified_provider_failure(monkeypatch, caplog): + caplog.set_level(logging.INFO) + response = Mock(status_code=200, json=Mock(side_effect=ValueError("PRIVATE-BODY"))) + result = _send(SopranoProvider(), monkeypatch, response) + assert result.outcome == Outcome.FAIL + assert result.failure_reason == "invalid_provider_json" + assert to_endpoint_http_status(result) == 502 + assert [record.event_name for record in caplog.records if hasattr(record, "event_name")][-2:] == [ + "provider_response_invalid_json", + "provider_response_processed", + ] + assert "PRIVATE-BODY" not in caplog.text + + +@pytest.mark.parametrize("error,status", [ + (provider_module.requests.exceptions.Timeout("PRIVATE"), 504), + (provider_module.requests.exceptions.ConnectionError("PRIVATE"), 502), +]) +def test_transport_failures_are_sanitized(monkeypatch, error, status): + monkeypatch.setattr(provider_module.requests, "request", Mock(side_effect=error)) + with pytest.raises(ProviderSendError) as raised: + SopranoProvider().send_otp( + "sms", "https://provider.example/messages", _delivery(), + {"mode": "oauth", "access_token": "provider-token"}, {}, 1500, _log()) + assert raised.value.status_code == status + + +def test_wrapped_stream_read_timeout_is_504_and_response_is_closed(monkeypatch): + wrapped = provider_module.requests.exceptions.ConnectionError( + ReadTimeoutError(None, "https://provider.example", "PRIVATE")) response = Mock(status_code=200, json=Mock(side_effect=wrapped)) - send = Mock(return_value=response) - monkeypatch.setattr(dispatch_module.requests, "request", send) - status, body = engine.dispatch(_request(), "r") - assert status == 504 and body["reason"] == "provider timeout" - send.assert_called_once() + monkeypatch.setattr(provider_module.requests, "request", Mock(return_value=response)) + with pytest.raises(ProviderSendError) as raised: + SopranoProvider().send_otp( + "sms", "https://provider.example/messages", _delivery(), + {"mode": "oauth", "access_token": "provider-token"}, {}, 1500, _log()) + assert raised.value.status_code == 504 response.close.assert_called_once() diff --git a/python/tests/test_function_app.py b/python/tests/test_function_app.py index 2ef50e2..c00ae6b 100644 --- a/python/tests/test_function_app.py +++ b/python/tests/test_function_app.py @@ -1,579 +1,342 @@ -import base64 import json import logging -from datetime import datetime, timedelta, timezone from concurrent.futures import ThreadPoolExecutor from pathlib import Path from threading import Event -from types import SimpleNamespace from unittest.mock import Mock import azure.functions as func import pytest from jwcrypto import jwe, jwk -from cryptography import x509 -from cryptography.hazmat.primitives import hashes, serialization -from cryptography.x509.oid import NameOID import function_app -import src.dispatch as dispatch_module -from src.request_log import RequestLog +import src.provider as provider_module +from src.jwe import JweDecryptor _KEY = jwk.JWK.generate(kty="RSA", size=2048) -_PRIVATE_PEM = _KEY.export_to_pem(private_key=True, password=None).decode() +_PRIVATE_PEM = _KEY.export_to_pem(True, None).decode() _CORRELATION = "2b65f5e5-9628-4894-8ba6-8785c3a9c010" _NONCE = "test-nonce" _PHONE = "+14255551234" _MESSAGE = " Your code is 123456; keep 7890 unchanged.\nCafé. " _CONTEXT = {"nonce": _NONCE, "phoneNumber": _PHONE, "message": _MESSAGE, "locale": "en-US"} -_FIXTURES = json.loads((Path(__file__).resolve().parents[2] / "tests/fixtures/contract.json").read_text(encoding="utf-8")) +_FIXTURES = json.loads( + (Path(__file__).resolve().parents[2] / "tests/fixtures/contract.json") + .read_text(encoding="utf-8")) -# get_functions() cannot be called twice on the same app. _FUNCTION = function_app.app.get_functions()[0] _HANDLER = _FUNCTION.get_user_function() -_FUNCTION_NAME = _FUNCTION.get_function_name() @pytest.fixture(autouse=True) def _isolate(monkeypatch): + monkeypatch.setenv("EPP_DECRYPTION_KEY_PEM", _PRIVATE_PEM) monkeypatch.delenv("EPP_ENCRYPTION_KEY_ID", raising=False) - monkeypatch.setenv("EPP_PROVIDER_NAME", "SOPRANO") - monkeypatch.setattr(function_app, "_key_provider", Mock(return_value=_PRIVATE_PEM)) - engine = dispatch_module.DispatchEngine( - function_app._registry, Mock(resolve=Mock(return_value="test-key")), - {"EPP_PROVIDER_NAME": "soprano", "EPP_PROVIDER_ENDPOINT": "https://qa4.example/oauth/messages", - "EPP_PROVIDER_AUTH_MODE": "oauth"}, - ) - 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() + monkeypatch.setenv("EPP_PROVIDER_NAME", "soprano") + monkeypatch.setenv("EPP_PROVIDER_ENDPOINT", "https://qa4.example/oauth/messages") + monkeypatch.setenv("EPP_PROVIDER_AUTH_MODE", "oauth") + monkeypatch.delenv("EPP_PROVIDER_CHANNEL", raising=False) + monkeypatch.setattr(function_app, "_decryptor", JweDecryptor(function_app.os.environ)) + monkeypatch.setattr(function_app, "_credentials", Mock( + get_credentials=Mock(return_value={"mode": "oauth", "access_token": "provider-token"}))) + monkeypatch.setattr(provider_module.requests, "request", Mock()) def _request(body, headers=None): raw = body if isinstance(body, bytes) else json.dumps(body).encode() - return func.HttpRequest(method="POST", url="/api/SendOtp", headers=headers or {}, params={}, body=raw) + return func.HttpRequest( + method="POST", url="/api/SendOtp", headers=headers or {}, params={}, body=raw) -def _encrypt(alg="RSA-OAEP-256", kid="test-kid", enc="A256GCM", context=None): - token = jwe.JWE(json.dumps(_CONTEXT if context is None else context).encode(), - protected=json.dumps({"alg": alg, "enc": enc, "kid": kid})) +def _encrypt(alg="RSA-OAEP-256", enc="A256GCM", kid="test-kid", context=None): + token = jwe.JWE( + json.dumps(_CONTEXT if context is None else context).encode(), + protected=json.dumps({"alg": alg, "enc": enc, "kid": kid}), + ) token.add_recipient(_KEY) return token.serialize(compact=True) def _envelope(**overrides): - payload = {"type": "microsoft.mfa.otpDeliver.v1", - "correlationId": _CORRELATION, "channel": 1, "mode": 1, "ttlSeconds": 60} + payload = { + "type": "microsoft.mfa.otpDeliver.v1", + "correlationId": _CORRELATION, + "channel": 1, + "mode": 1, + "ttlSeconds": 60, + "encryptedDeliveryContext": _encrypt(), + } payload.update(overrides) - if "encryptedDeliveryContext" not in payload: - payload["encryptedDeliveryContext"] = _encrypt() return payload -def _records(caplog): - return [value for record in caplog.records if (value := json.loads(record.getMessage())).get("functionName")] - - -def _summary(caplog): - records = _records(caplog) - summaries = [record for record in records if record["logType"] == "request"] - assert len(summaries) == 1 - summary = summaries[0] - assert summary == records[-1] - assert summary["eventName"] == "request_completed" - assert set(summary) == set(_FIXTURES["logging"]["summaryFields"]) - assert all(record["logType"] == "service" for record in records[:-1]) - assert all(record["functionRequestId"] == summary["functionRequestId"] for record in records) - assert all(record["functionInvocationId"] == summary["functionInvocationId"] for record in records) - assert all(record["functionName"] == _FUNCTION_NAME for record in records) - assert summary["elapsedMs"] >= 0 - prepared = [record for record in records if record["eventName"] == "response_prepared"] - assert len(prepared) == 1 and prepared[0] == records[-2] - assert prepared[0]["httpStatus"] == summary["httpStatus"] - assert prepared[0]["responseContainsNonce"] is summary["responseContainsNonce"] - assert prepared[0]["responseContainsCorrelationId"] is summary["responseContainsCorrelationId"] - assert summary["responseContainsNonce"] is (summary["httpStatus"] == 200) - for private in (_NONCE, _PHONE, _MESSAGE, "PRIVATE", "test-key", "provider-token"): - assert private not in caplog.text - return summary - - -def test_jwe_tag_tampering_and_missing_segments_fail_before_provider_io(monkeypatch, caplog): - monkeypatch.setenv("EPP_ENCRYPTION_KEY_ID", "configured-key-id") - segments = _encrypt().split(".") - tag = segments[-1] - segments[-1] = ("A" if tag[0] != "A" else "B") + tag[1:] - for compact in (".".join(segments), ".".join(segments[:4])): - response = _HANDLER(_request(_envelope(encryptedDeliveryContext=compact))) - assert response.status_code == 400 and json.loads(response.get_body())["error"] == "decryption_failed" - assert "encryption_key_id_mismatch" not in caplog.text - function_app._engine.secrets.resolve.assert_not_called() - dispatch_module.requests.request.assert_not_called() +def _events(caplog): + return [record for record in caplog.records if hasattr(record, "event_name")] -def test_shared_invalid_requests_return_safe_reasons_before_provider_io(): +def test_invalid_requests_return_contract_reasons_before_provider_work(caplog): + caplog.set_level(logging.INFO) valid = _envelope(encryptedDeliveryContext="unused") for fixture in _FIXTURES["badRequests"]: - payload = fixture["rawBody"].encode() if "rawBody" in fixture else {**valid, **fixture["overrides"]} - response = _HANDLER(_request(payload)) - result = json.loads(response.get_body()) - assert response.status_code == 400 and result["requestId"] - assert result == {"error": "bad_request", "reason": fixture["reason"], - "requestId": result["requestId"]}, fixture["name"] - contexts = [{**_CONTEXT, **changes} for changes in _FIXTURES["incompleteContexts"]] - for context in (*contexts, False, []): - payload = _envelope(mode=2, encryptedDeliveryContext=_encrypt(context=context)) + caplog.clear() + payload = ( + fixture["rawBody"].encode() + if "rawBody" in fixture + else {**valid, **fixture["overrides"]} + ) response = _HANDLER(_request(payload)) result = json.loads(response.get_body()) assert response.status_code == 400 - assert result == {"error": "bad_request", "reason": "incomplete delivery context", - "correlationId": _CORRELATION, "requestId": result["requestId"]} - function_app._engine.secrets.resolve.assert_not_called() - dispatch_module.requests.request.assert_not_called() - - -def test_shared_jwe_policy_permits_only_rsa_oaep_256_with_a256gcm(): - for fixture in _FIXTURES["jwe"]: - compact = _encrypt(alg=fixture["alg"], enc=fixture["enc"]) - response = _HANDLER(_request(_envelope(mode=2, encryptedDeliveryContext=compact))) - result = json.loads(response.get_body()) - assert response.status_code == (200 if fixture["accepted"] else 400) - if fixture["accepted"]: - assert result["nonce"] == _NONCE - else: - assert result == {"error": "decryption_failed", "correlationId": _CORRELATION, - "requestId": result["requestId"]} - function_app._engine.secrets.resolve.assert_not_called() - dispatch_module.requests.request.assert_not_called() - - -def test_jwe_authenticates_original_protected_header_bytes(): - header = '{ "kid" : "test-key", "enc" : "A256GCM", "alg" : "RSA-OAEP-256" }' - token = jwe.JWE(json.dumps(_CONTEXT).encode(), protected=header) - token.add_recipient(_KEY) - segments = token.serialize(compact=True).split('.') - original = base64.urlsafe_b64encode(header.encode()).decode().rstrip('=') - assert segments[0] == original - response = _HANDLER(_request(_envelope(mode=2, encryptedDeliveryContext='.'.join(segments)))) - assert response.status_code == 200 and json.loads(response.get_body())["nonce"] == _NONCE - segments[0] = base64.urlsafe_b64encode(json.dumps(json.loads(header), separators=(',', ':')).encode()).decode().rstrip('=') - response = _HANDLER(_request(_envelope(mode=2, encryptedDeliveryContext='.'.join(segments)))) - assert response.status_code == 400 and json.loads(response.get_body())["error"] == "decryption_failed" - function_app._engine.secrets.resolve.assert_not_called() - dispatch_module.requests.request.assert_not_called() - - -@pytest.mark.parametrize("certificate_first", [True, False]) -@pytest.mark.parametrize("base64_encoded", [True, False]) -def test_evaluation_accepts_key_vault_pem_bundle(monkeypatch, certificate_first, base64_encoded): - private_pem = _PRIVATE_PEM.encode() - private_key = serialization.load_pem_private_key(private_pem, password=None) - subject = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, "EPP-test")]) - certificate = (x509.CertificateBuilder().subject_name(subject).issuer_name(subject) - .public_key(private_key.public_key()).serial_number(x509.random_serial_number()) - .not_valid_before(datetime.now(timezone.utc) - timedelta(minutes=1)) - .not_valid_after(datetime.now(timezone.utc) + timedelta(days=1)) - .sign(private_key, hashes.SHA256()).public_bytes(serialization.Encoding.PEM)) - bundle = certificate + private_pem if certificate_first else private_pem + certificate - value = base64.b64encode(bundle).decode() if base64_encoded else bundle.decode() - monkeypatch.setattr(function_app, "_key_provider", - dispatch_module.make_key_provider({"EPP_DECRYPTION_KEY_PEM": value})) - response = _HANDLER(_request(_envelope(mode="evaluation"))) - assert response.status_code == 200 - assert json.loads(response.get_body())["nonce"] == _NONCE - dispatch_module.requests.request.assert_not_called() - function_app._engine._resolve_credential.assert_not_called() - - -def test_evaluation_decrypts_without_provider_configuration_or_work(monkeypatch, caplog): + assert result == { + "error": "bad_request", + "reason": fixture["reason"], + "requestId": result["requestId"], + } + assert _events(caplog)[-2].event_name == "response_prepared" + function_app._credentials.get_credentials.assert_not_called() + provider_module.requests.request.assert_not_called() + + +@pytest.mark.parametrize("context", [ + {"nonce": "", "phoneNumber": _PHONE, "message": _MESSAGE}, + {"nonce": _NONCE, "phoneNumber": "", "message": _MESSAGE}, + {"nonce": _NONCE, "phoneNumber": _PHONE, "message": ""}, + False, + [], +]) +def test_incomplete_delivery_context_fails_before_provider_work(context): + response = _HANDLER(_request(_envelope( + mode=2, encryptedDeliveryContext=_encrypt(context=context)))) + result = json.loads(response.get_body()) + assert response.status_code == 400 + assert result["error"] == "bad_request" + assert result["reason"] == "incomplete delivery context" + assert "nonce" not in result + function_app._credentials.get_credentials.assert_not_called() + + +def test_evaluation_validates_and_decrypts_without_provider_work(monkeypatch, caplog): caplog.set_level(logging.INFO) - monkeypatch.setenv("EPP_ENCRYPTION_KEY_ID", "configured-key-id") monkeypatch.delenv("EPP_PROVIDER_NAME") - function_app._engine.env.clear() - lookup = Mock() - monkeypatch.setattr(function_app._registry, "get", lookup) - response = _HANDLER(_request(_envelope(mode="Evaluation", provider="untrusted-body-provider"))) + monkeypatch.setenv("EPP_ENCRYPTION_KEY_ID", "configured-key") + selection = Mock(side_effect=AssertionError("evaluation must not select a provider")) + monkeypatch.setattr(function_app, "_select_provider", selection) + response = _HANDLER(_request(_envelope(mode="Evaluation"))) assert response.status_code == 200 assert json.loads(response.get_body()) == { - "nonce": _NONCE, "correlationId": _CORRELATION, "providerStatus": "accepted", + "nonce": _NONCE, + "correlationId": _CORRELATION, + "providerStatus": "accepted", } - function_app._key_provider.assert_called_once_with("test-kid") - warnings = [json.loads(record.getMessage())["eventName"] for record in caplog.records if record.levelno == logging.WARNING] - assert warnings == ["encryption_key_id_mismatch"] - assert [record["eventName"] for record in _records(caplog) - if record["eventName"] != "encryption_key_id_mismatch"] == _FIXTURES["logging"]["evaluationEvents"] - summary = _summary(caplog) - assert summary["encryptionKeyIdMismatch"] is True - assert summary["evaluation"] is True and summary["result"] == "evaluated" - assert summary["providerName"] is None and summary["providerAttempted"] is False - assert summary["providerHttpStatus"] is None and summary["providerElapsedMs"] is None - assert summary["providerCredentialSource"] is None and summary["providerCredentialElapsedMs"] is None - assert summary["providerEndpoint"] is None - assert all(value not in caplog.text for value in ("configured-key-id", "test-kid", "untrusted-body-provider")) - lookup.assert_not_called() - function_app._engine.secrets.resolve.assert_not_called() - dispatch_module.requests.request.assert_not_called() - - -def test_live_acceptance_waits_and_preserves_wire_data_but_not_plaintext_logs(monkeypatch, caplog): + selection.assert_not_called() + function_app._credentials.get_credentials.assert_not_called() + provider_module.requests.request.assert_not_called() + assert [record.event_name for record in _events(caplog)] == [ + "request_received", + "payload_validated", + "delivery_context_decrypted", + "encryption_key_id_mismatch", + "evaluation_completed", + "response_prepared", + "request_completed", + ] + + +@pytest.mark.parametrize("alg,enc,accepted", [ + ("RSA-OAEP-256", "A256GCM", True), + ("RSA-OAEP", "A256GCM", False), + ("RSA-OAEP-256", "A128GCM", False), +]) +def test_jwe_algorithms_are_pinned(alg, enc, accepted): + compact = _encrypt(alg=alg, enc=enc) + response = _HANDLER(_request(_envelope(mode=2, encryptedDeliveryContext=compact))) + assert response.status_code == (200 if accepted else 400) + result = json.loads(response.get_body()) + assert ("nonce" in result) is accepted + if not accepted: + assert result["error"] == "decryption_failed" + + +def test_jwe_tampering_fails_before_provider_work(): + segments = _encrypt().split(".") + segments[-1] = ("A" if segments[-1][0] != "A" else "B") + segments[-1][1:] + for compact in (".".join(segments), ".".join(segments[:4])): + response = _HANDLER(_request(_envelope( + encryptedDeliveryContext=compact))) + assert response.status_code == 400 + assert json.loads(response.get_body())["error"] == "decryption_failed" + function_app._credentials.get_credentials.assert_not_called() + provider_module.requests.request.assert_not_called() + + +def test_live_acceptance_waits_for_provider_and_preserves_wire_message(monkeypatch, caplog): caplog.set_level(logging.INFO) - monkeypatch.setenv("EPP_LOG_PLAINTEXT", "true") # Must not bypass privacy. entered, release = Event(), Event() - upstream = Mock(status_code=202, json=Mock(return_value={"status": "ENROUTE"})) + upstream = Mock( + status_code=202, + json=Mock(return_value={"status": "ENROUTE", "id": "provider-id"}), + ) def wait_for_acceptance(*args, **kwargs): entered.set() - assert release.wait(5), "test did not release provider acceptance" + assert release.wait(5) return upstream send = Mock(side_effect=wait_for_acceptance) - monkeypatch.setattr(dispatch_module.requests, "request", send) - speech = {"beforePasswordText": "Your code is", "password": "001234", "language": "en-US"} + monkeypatch.setattr(provider_module.requests, "request", send) with ThreadPoolExecutor(max_workers=1) as executor: - request = _request(_envelope(channel=2, encryptedDeliveryContext=_encrypt( - context={**_CONTEXT, "textToVoice": speech})), {"x-ms-client-request-id": "wire-message"}) - pending = executor.submit(_HANDLER, request) + pending = executor.submit( + _HANDLER, + _request( + _envelope(channel=2), + {"x-ms-client-request-id": "wire-message"}, + ), + ) try: - assert entered.wait(5), "handler did not reach provider" + assert entered.wait(5) assert not pending.done() - assert not any(record["logType"] == "request" for record in _records(caplog)) - assert _records(caplog)[-1]["eventName"] == "provider_request_started" finally: release.set() response = pending.result(timeout=5) + assert response.status_code == 200 - assert json.loads(response.get_body()) == { - "nonce": _NONCE, "correlationId": _CORRELATION, "providerStatus": "accepted", - } - send.assert_called_once() - upstream.close.assert_called_once() wire = json.loads(send.call_args.kwargs["data"]) - assert wire["voice"] == {"text2voice": { + assert wire["voice"]["text2voice"] == { "beforePasswordText": " Your code is ", "password": "123456", "afterPasswordText": "; keep 7890 unchanged.\nCafé. ", "language": "en-US", "gender": 1, "loop": 2, - }} - assert "text" not in wire and wire["messageTypes"] == ["voice"] and wire["correlationId"] == _CORRELATION - summary = _summary(caplog) - assert [record["eventName"] for record in _records(caplog)] == _FIXTURES["logging"]["liveEvents"] - assert summary["x-ms-correlation-id"] == _CORRELATION - assert summary["x-ms-client-request-id"] == "wire-message" - assert summary["providerName"] == "soprano" and summary["providerAuthMode"] == "oauth" - assert summary["providerHttpStatus"] == 202 and summary["providerStatus"] == "ENROUTE" - assert summary["providerOutcome"] == "Continue" and summary["providerAttempted"] is True - assert summary["channel"] == "voice" and summary["failureStage"] is None - assert summary["providerTimeoutMs"] == 1500 - assert 0 <= summary["providerElapsedMs"] <= summary["elapsedMs"] - for private in (_NONCE, _PHONE, _MESSAGE, "123456", "001234", "test-key"): + } + assert wire["correlationId"] == _CORRELATION + assert [record.event_name for record in _events(caplog)] == [ + "request_received", + "payload_validated", + "delivery_context_decrypted", + "provider_selected", + "provider_credential_resolution_started", + "provider_credential_resolved", + "provider_request_build_started", + "provider_request_built", + "provider_request_started", + "provider_response_received", + "provider_response_processed", + "response_prepared", + "request_completed", + ] + for private in (_NONCE, _PHONE, _MESSAGE, "123456", "provider-token"): assert private not in caplog.text -def test_provider_failure_preserves_status_without_retry_or_nonce(monkeypatch): - upstream = Mock(status_code=429, json=Mock(return_value={"status": "ENROUTE"})) - send = Mock(return_value=upstream) - monkeypatch.setattr(dispatch_module.requests, "request", send) - response = _HANDLER(_request(_envelope())) - body = json.loads(response.get_body()) - assert response.status_code == 429 and body["error"] == "provider_delivery_failed" - assert "nonce" not in body - send.assert_called_once() - assert send.call_args.kwargs["allow_redirects"] is False - upstream.close.assert_called_once() - - -def test_unexpected_handler_error_is_generic_and_does_not_send(monkeypatch, caplog): - caplog.set_level(logging.INFO) - monkeypatch.setattr(function_app, "read_config", Mock(side_effect=RuntimeError("PRIVATE-ERROR"))) - response = _HANDLER(_request({})) - body = json.loads(response.get_body()) - assert response.status_code == 500 and body["error"] == "delivery_failed" - assert set(body) == {"error", "correlationId", "requestId"} - assert "PRIVATE-ERROR" not in response.get_body().decode() + caplog.text - assert _summary(caplog)["failureStage"] == "handler" - assert _summary(caplog)["failureReason"] == "unexpected_error" - dispatch_module.requests.request.assert_not_called() - - -def test_identifier_sources_are_explicit_and_missing_microsoft_ids_are_not_generated(monkeypatch, caplog): - caplog.set_level(logging.INFO) - upstream = Mock(status_code=201, json=Mock(return_value={ - "status": "ENROUTE", "id": "provider-reference-id", "description": "PRIVATE-DESCRIPTION", - })) - monkeypatch.setattr(dispatch_module.requests, "request", Mock(return_value=upstream)) - headers = {"x-ms-client-request-id": "ms-request-id", "x-ms-correlation-id": "ms-header-correlation-id"} - for correlation_id in ("ms-envelope-correlation-id", None): - caplog.clear() - context = SimpleNamespace(invocation_id="function-invocation-id", function_name=_FUNCTION_NAME) - response = _HANDLER(_request(_envelope(correlationId=correlation_id), headers), context) - assert response.status_code == 200 - summary = _summary(caplog) - assert summary["x-ms-client-request-id"] == headers["x-ms-client-request-id"] - assert summary["x-ms-correlation-id"] == (correlation_id or headers["x-ms-correlation-id"]) - assert summary["msCorrelationIdSource"] == ("envelope" if correlation_id else "header") - assert _records(caplog)[0]["msCorrelationIdSource"] == "header" - assert summary["functionInvocationId"] == "function-invocation-id" - assert summary["providerMessageId"] == "provider-reference-id" - assert summary["functionRequestId"] not in (summary["x-ms-client-request-id"], summary["functionInvocationId"]) - caplog.clear() - response = _HANDLER(_request(_envelope(correlationId=None))) - summary = _summary(caplog) - assert json.loads(response.get_body())["correlationId"] == summary["functionRequestId"] - assert summary["x-ms-client-request-id"] is None and summary["x-ms-correlation-id"] is None - assert summary["msCorrelationIdSource"] == "none" and summary["functionInvocationId"] is None - - -@pytest.mark.parametrize("scenario,status,stage,reason,attempted", [ - ("invalid_json", 400, "request_validation", "invalid JSON body", False), - ("invalid_envelope", 400, "request_validation", "unsupported envelope type", False), - ("decryption", 400, "decryption", "decryption_failed", False), - ("incomplete_context", 400, "delivery_context_validation", "incomplete delivery context", False), - ("unknown_provider", 400, "provider_selection", "unknown_provider", False), - ("wrong_channel", 400, "provider_configuration", "channel_not_configured", False), - ("authentication_mismatch", 502, "provider_configuration", "authentication_mode_mismatch", False), - ("invalid_endpoint", 502, "provider_configuration", "invalid_provider_endpoint", False), - ("credentials", 502, "provider_credentials", "credential_unavailable", False), - ("request_build", 502, "provider_request_build", "request_build_failed", False), - ("timeout", 504, "provider_transport", "provider_timeout", True), - ("body_timeout", 504, "provider_transport", "provider_timeout", True), - ("network", 502, "provider_transport", "provider_network_error", True), - ("response_parse", 502, "provider_response", "response_parse_failed", True), - ("http_rejection", 429, "provider_response", "provider_rejected", True), +@pytest.mark.parametrize("status,payload,expected,reason", [ + (429, {"status": "ENROUTE"}, 429, "provider_http_error"), + (200, {"status": "FAILED"}, 502, "provider_rejected"), + (200, {}, 502, "unrecognized_provider_status"), ]) -def test_failures_emit_separate_events_and_complete_summaries(monkeypatch, caplog, scenario, status, stage, reason, attempted): +def test_provider_failures_preserve_safe_classification( + monkeypatch, caplog, status, payload, expected, reason +): caplog.set_level(logging.INFO) - payload = _envelope() - upstream = Mock(status_code=201, json=Mock(return_value={"status": "ENROUTE"})) - send = Mock(return_value=upstream) - monkeypatch.setattr(dispatch_module.requests, "request", send) - engine = function_app._engine - adapter = engine.registry.get("soprano") - if scenario == "invalid_json": - payload = b"{" - elif scenario == "invalid_envelope": - payload = {} - elif scenario == "decryption": - payload["encryptedDeliveryContext"] = "PRIVATE-NOT-A-JWE" - elif scenario == "incomplete_context": - payload["encryptedDeliveryContext"] = _encrypt(context={**_CONTEXT, "nonce": ""}) - elif scenario == "unknown_provider": - engine.env["EPP_PROVIDER_NAME"] = "PRIVATE-UNKNOWN-PROVIDER" - elif scenario == "wrong_channel": - engine.env["EPP_PROVIDER_CHANNEL"] = "voice" - elif scenario == "authentication_mismatch": - engine.env["EPP_PROVIDER_AUTH_MODE"] = "apiKey" - elif scenario == "invalid_endpoint": - engine.env["EPP_PROVIDER_ENDPOINT"] = "http://PRIVATE-ENDPOINT" - elif scenario == "credentials": - engine._resolve_credential.side_effect = RuntimeError("PRIVATE-CREDENTIAL-ERROR") - elif scenario == "request_build": - monkeypatch.setattr(adapter, "build_request", Mock(side_effect=RuntimeError("PRIVATE-BUILD-ERROR"))) - elif scenario == "timeout": - send.side_effect = dispatch_module.requests.exceptions.Timeout("PRIVATE-TIMEOUT") - elif scenario == "body_timeout": - upstream.json.side_effect = dispatch_module.requests.exceptions.Timeout("PRIVATE-BODY-TIMEOUT") - elif scenario == "network": - send.side_effect = dispatch_module.requests.exceptions.ConnectionError("PRIVATE-NETWORK-ERROR") - elif scenario == "response_parse": - monkeypatch.setattr(adapter, "parse_response", Mock(side_effect=RuntimeError("PRIVATE-PARSE-ERROR"))) - elif scenario == "http_rejection": - upstream.status_code = 429 - headers = {"x-ms-client-request-id": "ms-request-id", "x-ms-correlation-id": "ms-header-correlation-id"} - response = _HANDLER(_request(payload, headers)) - summary = _summary(caplog) - assert response.status_code == status == summary["httpStatus"] - assert summary["functionRequestId"] == json.loads(response.get_body())["requestId"] - assert summary["failureStage"] == stage and summary["failureReason"] == reason - assert summary["x-ms-client-request-id"] == headers["x-ms-client-request-id"] - assert summary["msCorrelationIdSource"] == ("header" if stage == "request_validation" else "envelope") - assert summary["x-ms-correlation-id"] == ( - headers["x-ms-correlation-id"] if stage == "request_validation" else _CORRELATION) - assert summary["providerAttempted"] is attempted and send.call_count == int(attempted) - assert summary["result"] == "failed" - assert _records(caplog)[-3]["eventName"] == ("provider_response_processed" if scenario == "http_rejection" else f"{stage}_failed") - assert summary["responseContainsNonce"] is False - assert summary["responseContainsCorrelationId"] is ("correlationId" in json.loads(response.get_body())) - if scenario == "credentials": - assert summary["providerCredentialSource"] == "managed_identity_client_assertion" - assert 0 <= summary["providerCredentialElapsedMs"] <= summary["elapsedMs"] - assert not any(record["eventName"] == "provider_credential_resolved" for record in _records(caplog)) - assert any(record.levelno == (logging.ERROR if status >= 500 else logging.WARNING) for record in caplog.records) - if attempted: - assert summary["providerTimeoutMs"] == 1500 - assert 0 <= summary["providerElapsedMs"] <= summary["elapsedMs"] - if scenario == "body_timeout": - assert summary["providerHttpStatus"] == 201 and summary["providerStatus"] is None - - -@pytest.mark.parametrize("correlation", [42, {"detail": "support-correlation-id"}, ["support-correlation-id"], ""]) -def test_invalid_correlation_metadata_cannot_leak_or_prevent_summary(caplog, correlation): + monkeypatch.setattr(provider_module.requests, "request", Mock(return_value=Mock( + status_code=status, json=Mock(return_value=payload)))) + response = _HANDLER(_request(_envelope())) + result = json.loads(response.get_body()) + assert response.status_code == expected + assert result["error"] == "provider_delivery_failed" + assert "nonce" not in result + processed = next( + record for record in _events(caplog) + if record.event_name == "provider_response_processed") + assert processed.failureReason == reason + assert processed.providerStatus in ("ENROUTE", "FAILED", "unmapped") + failure = next( + record for record in _events(caplog) + if record.event_name == "request_failed") + assert failure.failureStage == "provider_response" + assert failure.failureReason == reason + + +def test_invalid_provider_json_has_specific_failure_reason(monkeypatch, caplog): caplog.set_level(logging.INFO) - response = _HANDLER(_request(_envelope(mode=2, correlationId=correlation))) - assert response.status_code == 200 - summary = _summary(caplog) - assert summary["x-ms-correlation-id"] is None and summary["msCorrelationIdSource"] == "none" + monkeypatch.setattr(provider_module.requests, "request", Mock(return_value=Mock( + status_code=200, json=Mock(side_effect=ValueError("PRIVATE-RESPONSE"))))) + response = _HANDLER(_request(_envelope())) + assert response.status_code == 502 + processed = next( + record for record in _events(caplog) + if record.event_name == "provider_response_processed") + assert processed.failureReason == "invalid_provider_json" + assert "PRIVATE-RESPONSE" not in caplog.text -@pytest.mark.parametrize("valid_json", [True, False]) -def test_provider_diagnostics_never_log_unknown_statuses_or_response_bodies(monkeypatch, caplog, valid_json): +@pytest.mark.parametrize("provider,status,stage,reason", [ + ("unknown", 400, "provider_selection", "unknown_provider"), + ("soprano", 502, "provider_credentials", "credential_unavailable"), +]) +def test_selection_and_credential_failures_are_safe( + monkeypatch, caplog, provider, status, stage, reason +): caplog.set_level(logging.INFO) - upstream = Mock(status_code=200, json=Mock(return_value={ - "status": "PRIVATE-STATUS\nFORGED", "id": "provider-reference-id", "description": "PRIVATE-DESCRIPTION", - })) - if not valid_json: - upstream.json.side_effect = ValueError("PRIVATE-RESPONSE") - monkeypatch.setattr(dispatch_module.requests, "request", Mock(return_value=upstream)) + monkeypatch.setenv("EPP_PROVIDER_NAME", provider) + if stage == "provider_credentials": + function_app._credentials.get_credentials.side_effect = RuntimeError("PRIVATE") response = _HANDLER(_request(_envelope())) - summary = _summary(caplog) - assert response.status_code == 502 - assert summary["providerStatus"] == "unmapped" and summary["providerOutcome"] == "Fail" - assert summary["failureReason"] == ("provider_rejected" if valid_json else "invalid_provider_json") - assert "FORGED" not in caplog.text + assert response.status_code == status + failure = next( + record for record in _events(caplog) + if record.event_name == "request_failed") + assert (failure.failureStage, failure.failureReason) == (stage, reason) + assert "PRIVATE" not in caplog.text -def test_interleaved_invocations_keep_separate_log_contexts(monkeypatch, caplog): +def test_request_context_is_isolated_between_concurrent_invocations(monkeypatch, caplog): caplog.set_level(logging.INFO) - monkeypatch.setattr(dispatch_module.requests, "request", Mock( - side_effect=lambda *args, **kwargs: Mock(status_code=201, json=Mock(return_value={"status": "ENROUTE"})))) - requests = [_request(_envelope(correlationId=value)) for value in ("correlation-first", "correlation-second")] + monkeypatch.setattr(provider_module.requests, "request", Mock( + side_effect=lambda *args, **kwargs: Mock( + status_code=201, + json=Mock(return_value={"status": "ENROUTE", "id": "provider-id"})))) + requests = [ + _request(_envelope(correlationId=value)) + for value in ("correlation-first", "correlation-second") + ] with ThreadPoolExecutor(max_workers=2) as executor: - results = list(executor.map(_HANDLER, requests)) - assert all(result.status_code == 200 for result in results) - records = _records(caplog) - summaries = [record for record in records if record["logType"] == "request"] - assert len(summaries) == 2 - assert len({record["functionRequestId"] for record in summaries}) == 2 - assert {record["x-ms-correlation-id"] for record in summaries} == {"correlation-first", "correlation-second"} - for summary in summaries: - events = [record for record in records if record["functionRequestId"] == summary["functionRequestId"]] - assert [record["eventName"] for record in events] == _FIXTURES["logging"]["liveEvents"] - assert all(record["x-ms-correlation-id"] == summary["x-ms-correlation-id"] for record in events[1:]) - assert "PRIVATE" not in caplog.text + responses = list(executor.map(_HANDLER, requests)) + assert all(response.status_code == 200 for response in responses) + completed = [ + record for record in _events(caplog) + if record.event_name == "request_completed" + ] + assert len(completed) == 2 + assert len({record.functionRequestId for record in completed}) == 2 + assert {record.__dict__["x-ms-correlation-id"] for record in completed} == { + "correlation-first", "correlation-second"} + + +def test_unexpected_handler_error_is_generic(monkeypatch, caplog): + caplog.set_level(logging.INFO) + monkeypatch.setattr( + function_app, "read_config", Mock(side_effect=RuntimeError("PRIVATE-ERROR"))) + response = _HANDLER(_request(_envelope(mode=2))) + result = json.loads(response.get_body()) + assert response.status_code == 500 + assert result["error"] == "delivery_failed" + assert "PRIVATE-ERROR" not in response.get_body().decode() + caplog.text + assert any(record.event_name == "unexpected_error" for record in _events(caplog)) -def test_successful_lifecycle_logs_only_allowed_body_fields_oauth_ids_and_final_endpoint(monkeypatch, caplog): +def test_identifier_logging_accepts_only_safe_support_ids(monkeypatch, caplog): caplog.set_level(logging.INFO) - engine = function_app._engine - engine.env.update({ - "EPP_PROVIDER_ENDPOINT": "https://provider.example/api/send?key=PRIVATE-QUERY", - "EPP_PROVIDER_TENANT_ID": "provider-tenant-id", - "EPP_OUTBOUND_CLIENT_ID": "outbound-client-id", - "EPP_OUTBOUND_MI_CLIENT_ID": "outbound-mi-client-id", - }) - monkeypatch.setattr(dispatch_module.requests, "request", Mock(return_value=Mock( - status_code=201, json=Mock(return_value={"status": "ENROUTE"})))) - payload = _envelope(tenantId="PRIVATE-TENANT", diagnosticData={"token": "PRIVATE-UNKNOWN-FIELD"}) - response = _HANDLER(_request(payload, {"authorization": "PRIVATE-INBOUND-AUTH"})) + response = _HANDLER(_request( + _envelope(mode=2, correlationId={"PRIVATE": "VALUE"}), + {"x-ms-client-request-id": "safe-request-id", + "x-ms-correlation-id": "safe-header-correlation"}, + )) assert response.status_code == 200 - summary = _summary(caplog) - records = _records(caplog) - validated = next(record for record in records if record["eventName"] == "envelope_validated") - assert validated["envelopeType"] == summary["envelopeType"] == payload["type"] - assert validated["ttlSeconds"] == summary["ttlSeconds"] == 60 - assert validated["encryptedDeliveryContextPresent"] is True - credentials = [record for record in records if record["eventName"] in ( - "provider_credential_resolution_started", "provider_credential_resolved")] - assert len(credentials) == 2 - for record in [summary, *credentials]: - assert record["providerCredentialSource"] == "managed_identity_client_assertion" - assert record["providerTenantId"] == engine.env["EPP_PROVIDER_TENANT_ID"] - assert record["functionOutboundClientId"] == engine.env["EPP_OUTBOUND_CLIENT_ID"] - assert record["functionOutboundManagedIdentityClientId"] == engine.env["EPP_OUTBOUND_MI_CLIENT_ID"] - assert 0 <= summary["providerCredentialElapsedMs"] <= summary["elapsedMs"] - for record in [summary, *[record for record in records if record["eventName"] in ( - "provider_request_built", "provider_request_started")]]: - assert record["providerHttpMethod"] == "POST" - assert record["providerEndpoint"] == "https://provider.example/api/send" - built = next(record for record in records if record["eventName"] == "provider_request_built") - assert built["providerScheme"] == "https" and built["redirectsAllowed"] is False - assert payload["encryptedDeliveryContext"] not in caplog.text - assert summary["responseContainsNonce"] is True and summary["responseContainsCorrelationId"] is True - assert [record["eventName"] for record in records] == _FIXTURES["logging"]["liveEvents"] - - -def test_api_key_lifecycle_identifies_key_vault_resolution_without_logging_credentials(monkeypatch, caplog): - caplog.set_level(logging.INFO) - engine = function_app._engine - engine.env.update({"EPP_PROVIDER_NAME": "telesign", "EPP_PROVIDER_AUTH_MODE": "apiKey"}) - monkeypatch.setattr(engine, "_resolve_credential", - dispatch_module.DispatchEngine._resolve_credential.__get__(engine)) - monkeypatch.setattr(dispatch_module.requests, "request", Mock(return_value=Mock( - status_code=200, json=Mock(return_value={"status": {"code": 3001}})))) - assert _HANDLER(_request(_envelope())).status_code == 200 - summary = _summary(caplog) - assert summary["providerCredentialSource"] == "key_vault" and summary["providerAuthMode"] == "apiKey" - assert summary["providerTenantId"] is None - assert summary["functionOutboundClientId"] is None and summary["functionOutboundManagedIdentityClientId"] is None - assert 0 <= summary["providerCredentialElapsedMs"] <= summary["elapsedMs"] - assert [call.args[0] for call in engine.secrets.resolve.call_args_list] == ["telesign-api-key", "telesign-customer-id"] - assert [record["eventName"] for record in _records(caplog)] == _FIXTURES["logging"]["liveEvents"] - - -def test_optional_ttl_stays_null_and_invalid_body_values_never_enter_metadata(caplog): - caplog.set_level(logging.INFO) - payload = _envelope(mode=2) - del payload["ttlSeconds"] - assert _HANDLER(_request(payload)).status_code == 200 - assert _summary(caplog)["ttlSeconds"] is None - caplog.clear() - assert _HANDLER(_request({**payload, "ttlSeconds": "PRIVATE-INVALID-TTL"})).status_code == 400 - summary = _summary(caplog) - assert summary["ttlSeconds"] is None and summary["envelopeType"] is None - assert not any(record["eventName"] == "envelope_validated" for record in _records(caplog)) - - -def test_request_preparation_uses_the_adapter_final_url_and_allowlisted_method(monkeypatch, caplog): - caplog.set_level(logging.INFO) - adapter = function_app._engine.registry.get("soprano") - build_request = adapter.build_request - final_url = "https://different-provider.example/api/final?token=PRIVATE-TOKEN" - monkeypatch.setattr(adapter, "build_request", lambda *args: { - **build_request(*args), "url": final_url, "method": "PRIVATE-METHOD", - }) - send = Mock(return_value=Mock(status_code=201, json=Mock(return_value={"status": "ENROUTE"}))) - monkeypatch.setattr(dispatch_module.requests, "request", send) - assert _HANDLER(_request(_envelope())).status_code == 200 - summary = _summary(caplog) - assert summary["providerEndpoint"] == "https://different-provider.example/api/final" and summary["providerHttpMethod"] == "other" - assert summary["providerEndpoint"] != function_app._engine.env["EPP_PROVIDER_ENDPOINT"] - assert send.call_args.args[:2] == ("PRIVATE-METHOD", final_url) - - -def test_shared_id_cases_preserve_raw_values_or_explicitly_omit_invalid_metadata(caplog): - caplog.set_level(logging.INFO) - fields = ["x-ms-client-request-id", "x-ms-correlation-id", "providerTenantId", - "functionOutboundClientId", "functionOutboundManagedIdentityClientId", "providerMessageId"] - manifest = function_app._registry.get("soprano").manifest - for fixture in _FIXTURES["logging"]["identifiers"]: - caplog.clear() - value = "A" * fixture["length"] if "length" in fixture else fixture["value"] - log = RequestLog("function-request", None, value, value) - log.provider_selected(manifest) - log.credential_resolution_started(SimpleNamespace( - provider_tenant_id=value, outbound_client_id=value, outbound_managed_identity_client_id=value)) - log.provider_response_processed(manifest, SimpleNamespace( - provider_status_name="ENROUTE", provider_status_code=None, provider_message_id=value), "Continue", 200, True) - log.complete(200) - records = _records(caplog) - summary = records[-1] - for field in fields: - assert summary[field] == (value if fixture["accepted"] else None) - assert summary["omittedIdFields"] == (fields if fixture.get("omitted") else []) - assert not any(key.endswith("Hash") for record in records for key in record) - assert "PRIVATE" not in caplog.text - - -def test_shared_endpoint_cases_keep_only_scheme_host_port_and_api_path(caplog): - caplog.set_level(logging.INFO) - for fixture in _FIXTURES["logging"]["endpoints"]: - caplog.clear() - log = RequestLog("function-request", None, None, None) - log.provider_request_built("POST", fixture["url"]) - log.provider_request_started(1500) - log.complete(200) - assert all(record["providerEndpoint"] == fixture["logged"] for record in _records(caplog)) - assert "PRIVATE" not in caplog.text + completed = next( + record for record in _events(caplog) + if record.event_name == "request_completed") + assert completed.__dict__["x-ms-client-request-id"] == "safe-request-id" + assert completed.__dict__["x-ms-correlation-id"] == "safe-header-correlation" + assert completed.msCorrelationIdSource == "header" + assert "VALUE" not in caplog.text From c0fbd8260e7ec867d6e46fe6e8be470956d77ed1 Mon Sep 17 00:00:00 2001 From: Nisheet Jain Date: Fri, 2 Oct 2026 16:11:16 -0700 Subject: [PATCH 2/3] Fix Python source packaging Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 0769c5c6-97a7-4645-ab83-5cda60c73b16 --- package-python.ps1 | 6 +++++- 1 file changed, 5 insertions(+), 1 deletion(-) diff --git a/package-python.ps1 b/package-python.ps1 index d36582d..67b226f 100644 --- a/package-python.ps1 +++ b/package-python.ps1 @@ -23,7 +23,11 @@ try { New-Item -ItemType Directory -Path (Split-Path $destination) -Force | Out-Null Copy-Item -LiteralPath $file.FullName -Destination $destination } - if (-not (Test-Path -LiteralPath (Join-Path $stage 'src/dispatch.py'))) { throw 'Missing Python application source.' } + foreach ($name in @('function_app.py', 'src/jwe.py', 'src/provider.py')) { + if (-not (Test-Path -LiteralPath (Join-Path $stage $name))) { + throw "Missing Python application source: $name." + } + } $zip = Join-Path $temporary 'app.zip' [IO.Compression.ZipFile]::CreateFromDirectory($stage, $zip) New-Item -ItemType Directory -Path (Split-Path $archive) -Force | Out-Null From 28266a8371830d1eee338ad9becd1fe727eb4c62 Mon Sep 17 00:00:00 2001 From: Nisheet Jain Date: Fri, 2 Oct 2026 20:55:33 -0700 Subject: [PATCH 3/3] Update Python package verification Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com> Copilot-Session: 0769c5c6-97a7-4645-ab83-5cda60c73b16 --- .github/workflows/packages.yml | 2 +- 1 file changed, 1 insertion(+), 1 deletion(-) diff --git a/.github/workflows/packages.yml b/.github/workflows/packages.yml index 6f22c73..e3eb93f 100644 --- a/.github/workflows/packages.yml +++ b/.github/workflows/packages.yml @@ -75,7 +75,7 @@ jobs: $required = @{ 'epp-javascript.zip' = @('host.json', 'src/functions/SendOtp.js', 'node_modules/@azure/functions/package.json') 'epp-dotnet-source.zip' = @('host.json', 'dotnet.csproj', 'Program.cs', 'Functions/SendOtp.cs', 'Src/PhoneProviderBase.cs') - 'epp-python-source.zip' = @('host.json', 'function_app.py', 'requirements.txt', 'src/dispatch.py') + 'epp-python-source.zip' = @('host.json', 'function_app.py', 'requirements.txt', 'src/jwe.py', 'src/provider.py') } $checksums = foreach ($name in ($required.Keys | Sort-Object)) { $path = Join-Path 'artifacts' $name