mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-02 02:11:58 +00:00
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:
parent
6997223068
commit
2c3866ebb4
6 changed files with 448 additions and 7 deletions
|
|
@ -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
|
||||
|
|
|
|||
48
tests/integration/_support/tls.py
Normal file
48
tests/integration/_support/tls.py
Normal 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
|
||||
|
|
@ -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()
|
||||
|
|
|
|||
203
tests/integration/observability/_azure_storage_support.py
Normal file
203
tests/integration/observability/_azure_storage_support.py
Normal 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())
|
||||
|
|
@ -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()
|
||||
|
|
@ -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"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue