fix(azure_storage): name Data Lake objects without base64 padding or slashes (#43914)

Co-authored-by: yucheng <yucheng@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-09-30 18:17:13 -07:00 • committed by GitHub
parent 6997223068
commit 2c3866ebb4
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
6 changed files with 448 additions and 7 deletions

View file

@ -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

View file

@ -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

View file

@ -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()

View file

@ -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())

View file

@ -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()

View file

@ -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"