From 2c3866ebb4cc337b04beb1dac9535b836d356a44 Mon Sep 17 00:00:00 2001 From: "devin-ai-integration[bot]" <158243242+devin-ai-integration[bot]@users.noreply.github.com> Date: Wed, 30 Sep 2026 18:17:13 -0700 Subject: [PATCH] fix(azure_storage): name Data Lake objects without base64 padding or slashes (#43914) Co-authored-by: yucheng Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --- .../azure_storage/azure_storage.py | 12 +- tests/integration/_support/tls.py | 48 +++++ tests/integration/_support/wire.py | 29 ++- .../observability/_azure_storage_support.py | 203 ++++++++++++++++++ .../test_azure_storage_file_names.py | 61 ++++++ .../azure_storage/test_azure_storage.py | 102 +++++++++ 6 files changed, 448 insertions(+), 7 deletions(-) create mode 100644 tests/integration/_support/tls.py create mode 100644 tests/integration/observability/_azure_storage_support.py create mode 100644 tests/integration/observability/test_azure_storage_file_names.py diff --git a/litellm/integrations/azure_storage/azure_storage.py b/litellm/integrations/azure_storage/azure_storage.py index 13058bf4f22..1b1093477c8 100644 --- a/litellm/integrations/azure_storage/azure_storage.py +++ b/litellm/integrations/azure_storage/azure_storage.py @@ -30,6 +30,14 @@ from litellm.types.secret_managers.get_azure_ad_token_provider import ( from litellm.types.utils import StandardLoggingPayload AZURE_STORAGE_TOKEN_SCOPE: Final = "https://storage.azure.com/.default" +_ADLS_SAFE_NAME: Final = str.maketrans("/", "_", "=") + + +def adls_safe_file_name(payload_id: str | None) -> str: + """`=` padding and `/` in a base64 payload id are what the Data Lake service rejects, so the name drops the + padding and maps `/` to `_`. Standard base64 has no `_` and its padding is fixed by the length, so ids from + that alphabet stay distinct; anything else is left as is.""" + return f"{(payload_id or str(uuid.uuid4())).translate(_ADLS_SAFE_NAME)}.json" @cache @@ -182,7 +190,7 @@ class AzureBlobStorageLogger(CustomBatchLogger): async_client: Final = get_async_httpx_client(llm_provider=httpxSpecialProvider.LoggingCallback) json_payload: Final = safe_dumps(payload) + "\n" # Add newline for each log entry payload_bytes: Final = json_payload.encode("utf-8") - filename: Final = f"{payload.get('id') or str(uuid.uuid4())}.json" + filename: Final = adls_safe_file_name(payload.get("id")) base_url = f"{self.azure_storage_dfs_endpoint}/{self.azure_storage_file_system}/{filename}" # Execute the 3-step upload process @@ -368,7 +376,7 @@ class AzureBlobStorageLogger(CustomBatchLogger): verbose_logger.debug("Created directory: %s", today) # Create a file client - file_name: Final = f"{payload.get('id') or str(uuid.uuid4())}.json" + file_name: Final = adls_safe_file_name(payload.get("id")) file_client: Final = directory_client.get_file_client(file_name) # Create the file diff --git a/tests/integration/_support/tls.py b/tests/integration/_support/tls.py new file mode 100644 index 00000000000..39b98bcbf6c --- /dev/null +++ b/tests/integration/_support/tls.py @@ -0,0 +1,48 @@ +import datetime +import ipaddress +import ssl +from pathlib import Path +from typing import Final + +from cryptography import x509 +from cryptography.hazmat.primitives import hashes, serialization +from cryptography.hazmat.primitives.asymmetric import rsa +from cryptography.x509.oid import NameOID + + +def write_self_signed_cert(cert_dir: Path, names: tuple[str, ...] = ("localhost",)) -> tuple[Path, Path]: + """Write a loopback certificate valid for `names` and 127.0.0.1; returns (cert path, key path).""" + key: Final = rsa.generate_private_key(public_exponent=65537, key_size=2048) + now: Final = datetime.datetime.now(datetime.timezone.utc) + subject: Final = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, names[0])]) + alternatives: Final[tuple[x509.GeneralName, ...]] = tuple(x509.DNSName(name) for name in names) + ( + x509.IPAddress(ipaddress.ip_address("127.0.0.1")), + ) + cert: Final = ( + x509.CertificateBuilder() + .subject_name(subject) + .issuer_name(subject) + .public_key(key.public_key()) + .serial_number(x509.random_serial_number()) + .not_valid_before(now - datetime.timedelta(days=1)) + .not_valid_after(now + datetime.timedelta(days=7)) + .add_extension(x509.SubjectAlternativeName(alternatives), critical=False) + .sign(key, hashes.SHA256()) + ) + cert_file: Final = cert_dir / "cert.pem" + key_file: Final = cert_dir / "key.pem" + cert_file.write_bytes(cert.public_bytes(serialization.Encoding.PEM)) + key_file.write_bytes( + key.private_bytes( + serialization.Encoding.PEM, + serialization.PrivateFormat.TraditionalOpenSSL, + serialization.NoEncryption(), + ) + ) + return cert_file, key_file + + +def server_context(cert_file: Path, key_file: Path) -> ssl.SSLContext: + context: Final = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER) + context.load_cert_chain(certfile=cert_file, keyfile=key_file) + return context diff --git a/tests/integration/_support/wire.py b/tests/integration/_support/wire.py index ed96d4e4e83..1201a156c00 100644 --- a/tests/integration/_support/wire.py +++ b/tests/integration/_support/wire.py @@ -37,24 +37,37 @@ class Wire: url: str received: SimpleQueue[Request] disconnected: SimpleQueue[str] + connected: SimpleQueue[str] def drain(self) -> tuple[Request, ...]: return tuple(self.received.get_nowait() for _ in range(self.received.qsize())) + def connections(self) -> int: + return self.connected.qsize() + @contextmanager def wire_server( - respond: Callable[[Request], Reply], tls: ssl.SSLContext | None = None, port: int = 0 + respond: Callable[[Request], Reply], + tls: ssl.SSLContext | None = None, + port: int = 0, + keep_alive: bool = False, ) -> Generator[Wire, None, None]: - """Owned TCP peer; requests traverse the real HTTP client and serialization.""" + """Owned TCP peer; requests traverse the real HTTP client and serialization. With `keep_alive` the + peer honours HTTP/1.1 persistent connections so `connections()` counts the client's TCP sessions.""" received: Final[SimpleQueue[Request]] = SimpleQueue() errors: Final[SimpleQueue[Exception]] = SimpleQueue() disconnected: Final[SimpleQueue[str]] = SimpleQueue() + connected: Final[SimpleQueue[str]] = SimpleQueue() class Handler(BaseHTTPRequestHandler): protocol_version = "HTTP/1.1" timeout = 5 + def setup(self) -> None: + super().setup() + connected.put(f"{self.client_address[0]}:{self.client_address[1]}") + def respond(self) -> None: request: Final = Request( self.command, @@ -76,10 +89,13 @@ def wire_server( self.send_header("content-length", str(len(reply.body))) else: self.send_header("transfer-encoding", "chunked") - self.send_header("connection", "close") + if not keep_alive: + self.send_header("connection", "close") self.end_headers() try: - if reply.chunks is None: + if self.command == "HEAD": + self.wfile.flush() + elif reply.chunks is None: self.wfile.write(reply.body) else: for index, chunk in enumerate(reply.chunks): @@ -98,12 +114,14 @@ def wire_server( disconnected.put(request.target) except Exception as error: errors.put(error) - self.close_connection = True + self.close_connection = not keep_alive do_POST = respond do_PUT = respond do_GET = respond do_DELETE = respond + do_PATCH = respond + do_HEAD = respond def log_message(self, format: str, *args: object) -> None: pass @@ -124,6 +142,7 @@ def wire_server( f"{'https' if tls is not None else 'http'}://127.0.0.1:{server.server_port}", received, disconnected, + connected, ) finally: server.shutdown() diff --git a/tests/integration/observability/_azure_storage_support.py b/tests/integration/observability/_azure_storage_support.py new file mode 100644 index 00000000000..74bbb8e82fe --- /dev/null +++ b/tests/integration/observability/_azure_storage_support.py @@ -0,0 +1,203 @@ +import base64 +import hashlib +import hmac +import json +import threading +import time +from collections.abc import Mapping +from dataclasses import dataclass, field +from pathlib import Path +from types import MappingProxyType +from typing import Final +from urllib.parse import parse_qs, parse_qsl, quote, unquote, urlsplit + +import yaml +from integration._support.client import JsonValue, eventually, object_value +from integration._support.wire import Reply, Request + +ACCOUNT: Final = "litellmaudit" +FILE_SYSTEM: Final = "litellm-logs" +SINK_HOSTS: Final = (f"{ACCOUNT}.dfs.core.localhost", f"{ACCOUNT}.blob.core.localhost") +ACCOUNT_KEY: Final = base64.b64encode(b"synthetic-account-key-for-integration-tests").decode() +AUTHENTICATION_FAILED: Final = ( + b'{"error":{"code":"AuthenticationFailed","message":"Server failed to authenticate the request. ' + b'Make sure the value of Authorization header is formed correctly including the signature."}}' +) +_SIGNED_HEADERS: Final = ( + "content-encoding", + "content-language", + "content-length", + "content-md5", + "content-type", + "date", + "if-modified-since", + "if-match", + "if-none-match", + "if-unmodified-since", + "byte_range", +) + + +def shared_key_signature(request: Request) -> str: + """The SharedKey signature the service computes for a request: canonical headers, the account plus the + path exactly as sent on the wire, then the decoded query. The aio client signs a directory-scoped file + path with `%3D` but sends a bare `=`, so a padded name fails here the way it fails on the service.""" + headers: Final = {name.lower(): value for name, value in request.headers.items() if value} + standard: Final = tuple( + "" if name == "content-length" and headers.get(name) == "0" else headers.get(name, "") + for name in _SIGNED_HEADERS + ) + canonical_headers: Final = "".join( + f"{name}:{value}\n" for name, value in sorted(headers.items()) if name.startswith("x-ms-") + ) + parts: Final = urlsplit(request.target) + canonical_resource: Final = f"/{ACCOUNT}{parts.path}" + canonical_query: Final = "".join( + f"\n{name.lower()}:{unquote(value)}" for name, value in sorted(parse_qsl(parts.query, keep_blank_values=True)) + ) + string_to_sign: Final = ( + f"{request.method}\n" + "\n".join(standard) + "\n" + canonical_headers + canonical_resource + canonical_query + ) + digest: Final = hmac.new(base64.b64decode(ACCOUNT_KEY), string_to_sign.encode(), hashlib.sha256).digest() + return f"SharedKey {ACCOUNT}:{base64.b64encode(digest).decode()}" + + +@dataclass(slots=True) +class RecordingDataLakeSink: + """Speaks enough of the Azure Data Lake Gen2 REST surface for the SDK's account-key upload: filesystem + HEAD/PUT, blob HEAD for `exists`, PUT ?resource=directory|file, PATCH ?action=append|flush. Flushed + files are kept by path and can be failed, delayed or served slowly for the chaos cells.""" + + fail_status: int = 0 + delay_seconds: float = 0.0 + lock: threading.Lock = field(default_factory=threading.Lock) + directories: set[str] = field(default_factory=set) # mutable-ok: the sink is the durable store for the run + pending: dict[str, bytearray] = field(default_factory=dict) # mutable-ok: append lands before flush + files: dict[str, bytes] = field(default_factory=dict) # mutable-ok: flushed files must be readable later + flush_count: dict[str, int] = field(default_factory=dict) # mutable-ok: re-flush of one path means double upload + rejected: list[str] = field(default_factory=list) # mutable-ok: rejected request methods seen while failing + unauthenticated: list[str] = field( + default_factory=list + ) # mutable-ok: targets whose SharedKey signature did not verify + in_flight: int = 0 + peak: int = 0 + attempt_count: int = 0 + + def respond(self, request: Request) -> Reply: + parts: Final = urlsplit(request.target) + query: Final = {name: values[-1] for name, values in parse_qs(parts.query).items()} + path: Final = unquote(parts.path) + with self.lock: + self.attempt_count += 1 + if self.fail_status: + self.rejected.append(request.method) + return Reply(status=self.fail_status, body=b'{"error":{"code":"SinkFailure"}}') + presented: Final = next( + (value for name, value in request.headers.items() if name.lower() == "authorization"), "" + ) + if presented != shared_key_signature(request): + self.unauthenticated.append(request.target) + return Reply( + status=403, headers={"x-ms-error-code": "AuthenticationFailed"}, body=AUTHENTICATION_FAILED + ) + if path != f"/{FILE_SYSTEM}" and not path.startswith(f"/{FILE_SYSTEM}/"): + return Reply(status=400, body=b'{"error":{"code":"InvalidUri"}}') + self.in_flight += 1 + self.peak = max(self.peak, self.in_flight) + try: + if self.delay_seconds: + time.sleep(self.delay_seconds) + with self.lock: + return self._apply(request, path, query) + finally: + with self.lock: + self.in_flight -= 1 + + def _apply(self, request: Request, path: str, query: Mapping[str, str]) -> Reply: + stamp: Final = {"etag": '"0x1"', "last-modified": "Thu, 01 Jan 2026 00:00:00 GMT", "x-ms-request-id": "sink"} + empty: Final = "text/plain" + if path == f"/{FILE_SYSTEM}": + if request.method in ("HEAD", "GET"): + return Reply(headers={**stamp, "x-ms-namespace-enabled": "true"}, body=b"{}", content_type=empty) + if request.method == "PUT" and query.get("resource") == "filesystem": + return Reply(status=201, headers=stamp, body=b"", content_type=empty) + return Reply(status=400, body=b'{"error":{"code":"InvalidUri"}}') + if request.method == "HEAD": + if path in self.directories: + return Reply(headers={**stamp, "x-ms-meta-hdi_isfolder": "true"}, body=b"", content_type=empty) + if path in self.files: + return Reply(headers=stamp, body=b"", content_type=empty) + return Reply(status=404, headers={"x-ms-error-code": "PathNotFound"}, body=b"", content_type=empty) + if request.method == "GET": + if path in self.files: + return Reply(headers=stamp, body=self.files[path]) + return Reply(status=404, headers={"x-ms-error-code": "PathNotFound"}, body=b"", content_type=empty) + if request.method == "PUT": + if query.get("resource") == "directory": + self.directories.add(path) + return Reply(status=201, headers=stamp, body=b"", content_type=empty) + assert query.get("resource") == "file", request.target + self.pending[path] = bytearray() + return Reply(status=201, headers=stamp, body=b"", content_type=empty) + assert request.method == "PATCH", request.method + if query.get("action") == "append": + assert int(query["position"]) == len(self.pending[path]), request.target + self.pending[path].extend(request.body) + return Reply(status=202, headers=stamp, body=b"", content_type=empty) + assert query.get("action") == "flush", request.target + assert int(query["position"]) == len(self.pending[path]), request.target + self.files[path] = bytes(self.pending.pop(path)) + self.flush_count[path] = self.flush_count.get(path, 0) + 1 + return Reply(status=200, headers=stamp, body=b"", content_type=empty) + + def attempts(self) -> int: + with self.lock: + return self.attempt_count + + def rejected_methods(self) -> tuple[str, ...]: + with self.lock: + return tuple(self.rejected) + + def unauthenticated_targets(self) -> tuple[str, ...]: + with self.lock: + return tuple(self.unauthenticated) + + def duplicated(self) -> tuple[str, ...]: + with self.lock: + return tuple(path for path, count in self.flush_count.items() if count > 1) + + def stored(self) -> Mapping[str, bytes]: + with self.lock: + return MappingProxyType(dict(self.files)) + + def payloads(self) -> Mapping[str, dict[str, JsonValue]]: + return MappingProxyType({path: object_value(json.loads(body)) for path, body in self.stored().items()}) + + +def azure_storage_config( + path: Path, settings: Mapping[str, JsonValue] | None = None, *, callback_setting: str = "callbacks" +) -> Path: + config: Final = yaml.safe_load(Path("tests/integration/proxy_config.yaml").read_text()) + config["litellm_settings"].update({callback_setting: ["azure_storage"], **(settings or {})}) + target: Final = path / "azure_storage.yaml" + target.write_text(yaml.safe_dump(config)) + return target + + +def azure_storage_environment(sink_url: str, cert_file: Path) -> Mapping[str, str]: + port: Final = urlsplit(sink_url).port + return MappingProxyType( + { + "AZURE_STORAGE_ACCOUNT_NAME": ACCOUNT, + "AZURE_STORAGE_FILE_SYSTEM": FILE_SYSTEM, + "AZURE_STORAGE_ACCOUNT_KEY": ACCOUNT_KEY, + "AZURE_STORAGE_ENDPOINT_SUFFIX": f"core.localhost:{port}", + "SSL_CERT_FILE": str(cert_file), + } + ) + + +def collect_files(sink: RecordingDataLakeSink, count: int, seconds: float = 60) -> tuple[dict[str, JsonValue], ...]: + """Wait until `count` flushed files exist, then return every stored payload.""" + eventually(lambda: len(sink.stored()), lambda total: total >= count, seconds=seconds) + return tuple(sink.payloads().values()) diff --git a/tests/integration/observability/test_azure_storage_file_names.py b/tests/integration/observability/test_azure_storage_file_names.py new file mode 100644 index 00000000000..5009dba1d53 --- /dev/null +++ b/tests/integration/observability/test_azure_storage_file_names.py @@ -0,0 +1,61 @@ +import re +import uuid +from pathlib import Path +from typing import Final + +from _azure_storage_support import ( + SINK_HOSTS, + RecordingDataLakeSink, + azure_storage_config, + azure_storage_environment, +) +from _s3_v2_support import surface_reply +from integration._support.client import Gateway, eventually +from integration._support.process import owned_proxy +from integration._support.tls import server_context, write_self_signed_cert +from integration._support.wire import wire_server + +ADLS_SAFE_FILE_NAME: Final = re.compile(r"^[A-Za-z0-9._+-]+\.json$") + + +def _responses_id(candidate: Gateway, model: str, key: str, marker: str) -> str: + response: Final = candidate.request("POST", "/v1/responses", {"model": model, "input": marker}, key=key) + assert response.status_code == 200, response.text + return str(response.json()["id"]) + + +def test_responses_ids_with_base64_padding_land_under_adls_safe_names(gateway: Gateway, tmp_path: Path) -> None: + """A /v1/responses id is `resp_` plus base64 with `=` padding decided by the encoded length, so upstream ids + of several lengths yield both `=` and `==` padded ids; each must land as a file the service accepts.""" + marker: Final = f"azure-{uuid.uuid4().hex[:8]}" + sink: Final = RecordingDataLakeSink() + cert, key = write_self_signed_cert(tmp_path, SINK_HOSTS) + with ( + wire_server(surface_reply) as provider, + wire_server(sink.respond, tls=server_context(cert, key), keep_alive=True) as store, + ): + environment: Final = {**azure_storage_environment(store.url, cert), "DEFAULT_FLUSH_INTERVAL_SECONDS": "1"} + config: Final = azure_storage_config(tmp_path) + with ( + owned_proxy(gateway, tmp_path, environment, config=config, workers=1) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + api_key: Final = scenario.key(models=[model]) + answered: Final = tuple( + _responses_id(candidate, model, api_key, f"{marker}-{'x' * extra}") for extra in range(6) + ) + assert {response_id.count("=") for response_id in answered} >= {1, 2}, answered + eventually( + lambda: len(sink.stored()) + len(sink.unauthenticated_targets()), + lambda settled: settled >= len(answered), + seconds=60, + ) + assert sink.unauthenticated_targets() == (), sink.unauthenticated_targets() + assert frozenset(str(payload["id"]) for payload in sink.payloads().values()) == frozenset(answered), tuple( + sink.stored() + ) + names: Final = tuple(path.rsplit("/", 1)[1] for path in sink.stored()) + assert all(ADLS_SAFE_FILE_NAME.match(name) for name in names), names + assert len(frozenset(names)) == len(answered), names + assert provider.drain() diff --git a/tests/unit/integrations/azure_storage/test_azure_storage.py b/tests/unit/integrations/azure_storage/test_azure_storage.py index 6e1dab4a71a..dc162b5cd93 100644 --- a/tests/unit/integrations/azure_storage/test_azure_storage.py +++ b/tests/unit/integrations/azure_storage/test_azure_storage.py @@ -1,4 +1,7 @@ import asyncio +import base64 +import json +import re import sys import threading from unittest.mock import AsyncMock, MagicMock, patch @@ -8,6 +11,7 @@ import pytest from litellm.integrations.azure_storage.azure_storage import ( AzureBlobStorageLogger, _cached_credential_chain_token_provider, + adls_safe_file_name, ) from litellm.types.secret_managers.get_azure_ad_token_provider import AzureCredentialType from litellm.types.utils import StandardLoggingPayload @@ -365,3 +369,101 @@ async def test_service_client_defaults_to_commercial_endpoint(mock_env_vars): fake_aio_module.DataLakeServiceClient.call_args.kwargs["account_url"] == "https://test-account.dfs.core.windows.net" ) + + +@pytest.mark.parametrize( + ("payload_id", "expected"), + ( + ("resp_YWJj", "resp_YWJj.json"), + ("resp_YWJjZA==", "resp_YWJjZA.json"), + ("resp_YWJjZGU=", "resp_YWJjZGU.json"), + ("resp_+/8=", "resp_+_8.json"), + ("resp_a+b", "resp_a+b.json"), + ("chatcmpl-abc123", "chatcmpl-abc123.json"), + ), +) +def test_adls_safe_file_name_rewrites_base64_padding_and_reserved_characters(payload_id, expected): + name = adls_safe_file_name(payload_id) + assert name == expected, f"{payload_id!r} must map to {expected!r}, got {name!r}" + assert re.fullmatch(r"[A-Za-z0-9._+-]+\.json", name), ( + f"{name!r} must contain no characters Data Lake treats as path separators or signing input" + ) + + +def test_adls_safe_file_name_is_deterministic_and_distinct_per_id(): + ids = ( + "resp_" + base64.b64encode(b"a").decode(), + "resp_" + base64.b64encode(b"ab").decode(), + "resp_" + base64.b64encode(b"abc").decode(), + "resp_" + base64.b64encode(b"abcd").decode(), + "resp_" + base64.b64encode(b"\xfb\xff").decode(), + ) + names = tuple(adls_safe_file_name(payload_id) for payload_id in ids) + again = tuple(adls_safe_file_name(payload_id) for payload_id in ids) + assert names == again, "the rewrite must be deterministic for a given id" + assert len(set(names)) == len(ids), f"distinct ids must map to distinct names, got {names}" + + +def test_adls_safe_file_name_without_an_id_is_a_uuid_json(): + name = adls_safe_file_name(None) + assert re.fullmatch(r"[0-9a-f-]{36}\.json", name), ( + f"an id-less payload must fall back to a uuid-named file, got {name!r}" + ) + + +@pytest.mark.asyncio +async def test_account_key_upload_names_the_file_adls_safe_and_keeps_the_original_id( + workload_identity_env_vars, monkeypatch +): + monkeypatch.setenv("AZURE_STORAGE_ACCOUNT_KEY", "dGVzdC1rZXk=") + + file_client = MagicMock() + file_client.create_file = AsyncMock() + file_client.append_data = AsyncMock() + file_client.flush_data = AsyncMock() + directory_client = MagicMock() + directory_client.exists = AsyncMock(return_value=True) + directory_client.get_file_client = MagicMock(return_value=file_client) + file_system_client = MagicMock() + file_system_client.get_directory_client = MagicMock(return_value=directory_client) + service_client = MagicMock() + service_client.get_file_system_client = MagicMock(return_value=file_system_client) + fake_aio_module = MagicMock() + fake_aio_module.DataLakeServiceClient = MagicMock(return_value=service_client) + + with patch.dict(sys.modules, {"azure.storage.filedatalake.aio": fake_aio_module}): + logger = AzureBlobStorageLogger() + await logger.async_upload_payload_to_azure_blob_storage({"id": "resp_YWJjZA=="}) + + directory_client.get_file_client.assert_called_once_with("resp_YWJjZA.json") + body = json.loads(file_client.append_data.call_args.kwargs["data"]) + assert body["id"] == "resp_YWJjZA==", "the stored payload must keep the original id byte for byte" + + +@pytest.mark.asyncio +async def test_entra_upload_names_the_file_adls_safe_and_keeps_the_original_id(mock_env_vars): + with ( + patch("litellm.integrations.azure_storage.azure_storage.get_async_httpx_client") as mock_get_client, + patch("litellm.integrations.azure_storage.azure_storage.get_azure_ad_token_from_entra_id") as mock_get_token, + ): + mock_http_client = AsyncMock() + mock_response = MagicMock() + mock_http_client.put.return_value = mock_response + mock_http_client.patch.return_value = mock_response + mock_get_client.return_value = mock_http_client + mock_token_provider = MagicMock() + mock_token_provider.return_value = "mock-azure-ad-token" + mock_get_token.return_value = mock_token_provider + + logger = AzureBlobStorageLogger() + logger.azure_auth_token = "mock-azure-ad-token" + logger.token_expiry = None + + await logger.async_upload_payload_to_azure_blob_storage({"id": "resp_YWJjZA=="}) + + put_call_args = mock_http_client.put.call_args + assert put_call_args[0][0] == ( + "https://test-account.dfs.core.windows.net/test-container/resp_YWJjZA.json?resource=file" + ), f"the Entra path must be the rewritten name, got {put_call_args[0][0]!r}" + append_call = mock_http_client.patch.call_args_list[0] + assert "resp_YWJjZA==" in append_call[1]["data"], "the stored payload must keep the original id byte for byte"