From 54ae4c5bbfcb93d4f08077a1c2b385f41a77df1e 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:50:15 -0700 Subject: [PATCH] fix(azure_storage): keep the DataLakeServiceClient alive until its TTL elapses (#43082) * fix(azure_storage): keep the DataLakeServiceClient alive until its TTL elapses Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * ci: rerun integrations shard after unrelated gitlab prompt manager timeout Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(integration): audit azure_storage client reuse against a local Data Lake sink Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(azure_storage): cover the exact TTL expiry boundary Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(azure_storage): wait for a rejected write before flipping the sink back Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(e2e): add azure_storage log delivery cells behind an opt-in lane Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(azure_storage): restart the proxy mid burst and bound the loss to the unflushed queue Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(azure_storage): drop the redundant stop after the owned proxy exits Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(e2e): read azure_storage objects at the auth-mode-dependent layout Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(e2e): install the datalake sdk in the e2e lint environment Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(azure_storage): drop the opt-in real Azure e2e cells and their e2e-dev dependency Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> Co-authored-by: yucheng --- .../azure_storage/azure_storage.py | 6 +- .../observability/test_azure_storage_chaos.py | 234 ++++++++++ .../test_azure_storage_client_ttl.py | 401 ++++++++++++++++++ .../azure_storage/test_azure_storage.py | 57 +++ 4 files changed, 696 insertions(+), 2 deletions(-) create mode 100644 tests/integration/observability/test_azure_storage_chaos.py create mode 100644 tests/integration/observability/test_azure_storage_client_ttl.py diff --git a/litellm/integrations/azure_storage/azure_storage.py b/litellm/integrations/azure_storage/azure_storage.py index 1b1093477c8..30e0901c32a 100644 --- a/litellm/integrations/azure_storage/azure_storage.py +++ b/litellm/integrations/azure_storage/azure_storage.py @@ -54,6 +54,7 @@ class AzureBlobStorageLogger(CustomBatchLogger): build_credential_chain_token_provider: Callable[ [], Callable[[], str] ] = _cached_credential_chain_token_provider, + clock: Callable[[], float] = time.time, **kwargs, ): try: @@ -77,6 +78,7 @@ class AzureBlobStorageLogger(CustomBatchLogger): self.azure_storage_endpoint_suffix: str = ( os.getenv("AZURE_STORAGE_ENDPOINT_SUFFIX") or AZURE_STORAGE_DEFAULT_ENDPOINT_SUFFIX ) + self._clock: Callable[[], float] = clock self._service_client = None # Time that the azure service client expires, in order to reset the connection pool and keep it fresh self._service_client_timeout: float | None = None @@ -339,7 +341,7 @@ class AzureBlobStorageLogger(CustomBatchLogger): from azure.storage.filedatalake.aio import DataLakeServiceClient # expire old clients to recover from connection issues - if self._service_client_timeout and self._service_client and self._service_client_timeout > time.time(): + if self._service_client_timeout and self._service_client and self._service_client_timeout <= self._clock(): await self._service_client.close() self._service_client = None if not self._service_client: @@ -347,7 +349,7 @@ class AzureBlobStorageLogger(CustomBatchLogger): account_url=self.azure_storage_dfs_endpoint, credential=self.azure_storage_account_key, ) - self._service_client_timeout = time.time() + _DEFAULT_TTL_FOR_HTTPX_CLIENTS + self._service_client_timeout = self._clock() + _DEFAULT_TTL_FOR_HTTPX_CLIENTS return self._service_client async def upload_to_azure_data_lake_with_azure_account_key(self, payload: StandardLoggingPayload): diff --git a/tests/integration/observability/test_azure_storage_chaos.py b/tests/integration/observability/test_azure_storage_chaos.py new file mode 100644 index 00000000000..079ce72f9ba --- /dev/null +++ b/tests/integration/observability/test_azure_storage_chaos.py @@ -0,0 +1,234 @@ +import os +import signal +import uuid +from pathlib import Path +from typing import Final + +import httpx +from _azure_storage_support import ( + SINK_HOSTS, + RecordingDataLakeSink, + azure_storage_config, + azure_storage_environment, + collect_files, +) +from _s3_v2_support import matched_ids, mixed_burst, surface_reply +from integration._support.client import Gateway, JsonValue, eventually +from integration._support.process import group_members, owned_proxy_process +from integration._support.tls import server_context, write_self_signed_cert +from integration._support.wire import wire_server + +WORKERS: Final = 2 +FLUSH_SECONDS: Final = "1" + + +def _readiness_ok(candidate: Gateway) -> bool: + try: + return candidate.request("GET", "/health/readiness").status_code == 200 + except httpx.TransportError: + return False + + +def _present_count(payloads: tuple[dict[str, JsonValue], ...], answered: tuple[tuple[str, str | None], ...]) -> int: + response_ids: Final = frozenset(response_id for response_id, _ in answered) + call_ids: Final = frozenset(call_id for _, call_id in answered if call_id is not None) + return sum(1 for payload in payloads if payload["id"] in response_ids or payload["litellm_call_id"] in call_ids) + + +def test_sink_outage_mid_burst_loses_only_the_outage_window_and_recovers_exactly_once( + gateway: Gateway, tmp_path: Path +) -> None: + 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": FLUSH_SECONDS, + } + config: Final = azure_storage_config(tmp_path) + with ( + owned_proxy_process(gateway, tmp_path, environment, config=config, workers=WORKERS) as owned, + owned.gateway.scenario() as scenario, + ): + candidate: Final = owned.gateway + openai_model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + anthropic_model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", api_base=provider.url, api_key="synthetic-provider-key" + ) + key: Final = scenario.key(models=[openai_model, anthropic_model]) + first: Final = mixed_burst(candidate, openai_model, anthropic_model, key, f"{marker}-first", per_surface=2) + collect_files(sink, len(first)) + attempts_before_outage: Final = sink.attempts() + sink.fail_status = 503 + outage: Final = mixed_burst( + candidate, openai_model, anthropic_model, key, f"{marker}-outage", per_surface=1 + ) + eventually(sink.attempts, lambda count: count > attempts_before_outage, seconds=30) + readiness: Final = candidate.request("GET", "/health/readiness") + assert readiness.status_code == 200, readiness.text + sink.fail_status = 0 + tail: Final = mixed_burst(candidate, openai_model, anthropic_model, key, f"{marker}-tail", per_surface=1) + answered: Final = first + outage + tail + payloads: Final = eventually( + lambda: tuple(sink.payloads().values()), + lambda stored: _present_count(stored, tail) == len(tail), + seconds=60, + ) + landed: Final = matched_ids(payloads, answered) + assert sink.duplicated() == (), sink.duplicated() + assert len(sink.stored()) == len(landed), f"{len(sink.stored())} files for {len(landed)} matched ids" + assert len(landed) >= len(first) + len(tail), ( + f"lost {len(answered) - len(landed)} of {len(answered)} payloads, " + f"expected at most the {len(outage)} sent during the outage" + ) + assert len(answered) - len(landed) <= len(outage) + + +def test_slow_sink_lands_every_id_once_without_deadlock(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = f"azure-{uuid.uuid4().hex[:8]}" + sink: Final = RecordingDataLakeSink(delay_seconds=0.3) + 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": FLUSH_SECONDS, + } + config: Final = azure_storage_config(tmp_path) + with ( + owned_proxy_process(gateway, tmp_path, environment, config=config, workers=WORKERS) as owned, + owned.gateway.scenario() as scenario, + ): + candidate: Final = owned.gateway + openai_model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + anthropic_model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", api_base=provider.url, api_key="synthetic-provider-key" + ) + key: Final = scenario.key(models=[openai_model, anthropic_model]) + answered: Final = mixed_burst(candidate, openai_model, anthropic_model, key, marker, per_surface=6) + payloads: Final = collect_files(sink, len(answered), seconds=70) + assert len(matched_ids(payloads, answered)) == len(answered), tuple(sink.stored()) + assert len(sink.stored()) == len(answered) + assert sink.duplicated() == (), sink.duplicated() + assert sink.peak >= 1 + assert store.connections() <= 2 * WORKERS, ( + f"{store.connections()} sink connections for {len(answered)} uploads" + ) + + +def test_killing_one_worker_keeps_the_other_serving_and_uploading(gateway: Gateway, tmp_path: Path) -> None: + 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": FLUSH_SECONDS, + } + config: Final = azure_storage_config(tmp_path) + with ( + owned_proxy_process(gateway, tmp_path, environment, config=config, workers=WORKERS) as owned, + owned.gateway.scenario() as scenario, + ): + candidate: Final = owned.gateway + openai_model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + anthropic_model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", api_base=provider.url, api_key="synthetic-provider-key" + ) + key: Final = scenario.key(models=[openai_model, anthropic_model]) + first: Final = mixed_burst(candidate, openai_model, anthropic_model, key, f"{marker}-first", per_surface=2) + collect_files(sink, len(first)) + workers: Final = tuple( + process for process in group_members(owned.process.pid) if process.pid != owned.process.pid + ) + assert workers, "no uvicorn workers in the owned proxy process group" + os.kill(workers[0].pid, signal.SIGKILL) + eventually(lambda: _readiness_ok(candidate), lambda ok: ok, seconds=30) + rest: Final = mixed_burst(candidate, openai_model, anthropic_model, key, f"{marker}-rest", per_surface=4) + payloads: Final = eventually( + lambda: tuple(sink.payloads().values()), + lambda stored: _present_count(stored, rest) == len(rest), + seconds=60, + ) + members_after: Final = eventually( + lambda: len(group_members(owned.process.pid)), + lambda count: count >= 1 + WORKERS, + seconds=30, + return_last_on_timeout=True, + ) + landed: Final = matched_ids(payloads, first + rest) + assert sink.duplicated() == (), sink.duplicated() + assert len(landed) >= len(rest), f"only {len(landed)} payloads landed for {len(rest)} post-kill requests" + assert _present_count(payloads, rest) == len(rest), ( + f"lost {len(rest) - _present_count(payloads, rest)} post-kill payloads; " + f"process group holds {members_after - 1} workers after the kill" + ) + + +def test_restarting_the_proxy_before_the_queue_flushes_bounds_the_loss_to_the_unflushed_queue_and_recovers( + gateway: Gateway, tmp_path: Path +) -> None: + 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": FLUSH_SECONDS, + } + config: Final = azure_storage_config(tmp_path) + with owned_proxy_process(gateway, tmp_path, environment, config=config, workers=WORKERS) as first_owned: + with first_owned.gateway.scenario() as scenario: + openai_model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + anthropic_model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", + api_base=provider.url, + api_key="synthetic-provider-key", + ) + first_key: Final = scenario.key(models=[openai_model, anthropic_model]) + first: Final = mixed_burst( + first_owned.gateway, openai_model, anthropic_model, first_key, f"{marker}-first", per_surface=2 + ) + collect_files(sink, len(first)) + cut: Final = mixed_burst( + first_owned.gateway, openai_model, anthropic_model, first_key, f"{marker}-cut", per_surface=2 + ) + with owned_proxy_process(gateway, tmp_path, environment, config=config, workers=WORKERS) as second_owned: + with second_owned.gateway.scenario() as scenario: + second_openai: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + second_anthropic: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", + api_base=provider.url, + api_key="synthetic-provider-key", + ) + second_key: Final = scenario.key(models=[second_openai, second_anthropic]) + tail: Final = mixed_burst( + second_owned.gateway, second_openai, second_anthropic, second_key, f"{marker}-tail", per_surface=2 + ) + payloads: Final = eventually( + lambda: tuple(sink.payloads().values()), + lambda stored: _present_count(stored, tail) == len(tail), + seconds=60, + ) + answered: Final = first + cut + tail + landed: Final = matched_ids(payloads, answered) + assert sink.duplicated() == (), sink.duplicated() + assert len(sink.stored()) == len(landed), f"{len(sink.stored())} files for {len(landed)} matched ids" + assert _present_count(payloads, first) == len(first) + assert _present_count(payloads, tail) == len(tail) + assert len(answered) - len(landed) <= len(cut), ( + f"lost {len(answered) - len(landed)} of {len(answered)} payloads; the in-memory queue is dropped on " + f"restart by design, so at most the {len(cut)} pre-restart unflushed requests may be lost" + ) diff --git a/tests/integration/observability/test_azure_storage_client_ttl.py b/tests/integration/observability/test_azure_storage_client_ttl.py new file mode 100644 index 00000000000..f8f32820daa --- /dev/null +++ b/tests/integration/observability/test_azure_storage_client_ttl.py @@ -0,0 +1,401 @@ +import json +import uuid +from collections.abc import Callable +from pathlib import Path +from typing import Final + +from _azure_storage_support import ( + SINK_HOSTS, + RecordingDataLakeSink, + azure_storage_config, + azure_storage_environment, + collect_files, +) +from _s3_v2_support import SURFACES, call_surface, matched_ids, surface_reply +from integration._support.client import Gateway, eventually +from integration._support.database import read_rows +from integration._support.process import owned_proxy +from integration._support.tls import server_context, write_self_signed_cert +from integration._support.wire import Reply, Request, wire_server + +WORKERS: Final = 2 +FLUSH_SECONDS: Final = "1" + + +def _chat_completion(candidate: Gateway, model: str, key: str, marker: str) -> tuple[str, str | None]: + response: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": marker}], "cache": {"no-cache": True}}, + key=key, + ) + assert response.status_code == 200, response.text + return str(response.json()["id"]), response.headers.get("x-litellm-call-id") + + +def _marker_of(request: Request) -> str | None: + if request.method != "POST" or not request.body: + return None + body: Final = json.loads(request.body) + messages: Final = body.get("messages") + if isinstance(messages, list) and messages: + content: Final = messages[0].get("content") if isinstance(messages[0], dict) else None + if isinstance(content, str): + return content + input_value: Final = body.get("input") + return input_value if isinstance(input_value, str) else None + + +def upstream_rejecting_fail_markers(status: int) -> Callable[[Request], Reply]: + def respond(request: Request) -> Reply: + marker: Final = _marker_of(request) + if marker is not None and marker.startswith("fail-"): + return Reply(status=status, body=json.dumps({"error": {"message": f"upstream rejected {marker}"}}).encode()) + return surface_reply(request) + + return respond + + +def _spend_row_visible(response_id: str) -> None: + eventually( + lambda: read_rows('SELECT request_id FROM "LiteLLM_SpendLogs" WHERE request_id=%s', (response_id,)), + lambda rows: len(rows) == 1, + seconds=60, + ) + + +def test_every_surface_lands_once_and_the_client_is_reused_across_uploads(gateway: Gateway, tmp_path: Path) -> None: + 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": FLUSH_SECONDS, + } + config: Final = azure_storage_config(tmp_path) + with ( + owned_proxy(gateway, tmp_path, environment, config=config, workers=WORKERS) as candidate, + candidate.scenario() as scenario, + ): + openai_model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + anthropic_model: Final = scenario.model( + model="anthropic/claude-sonnet-4-5-20250929", api_base=provider.url, api_key="synthetic-provider-key" + ) + key: Final = scenario.key(models=[openai_model, anthropic_model]) + answered: Final = tuple( + call_surface(candidate, surface, openai_model, anthropic_model, key, f"{marker}-{surface}-{index}") + for index in range(3) + for surface in SURFACES + ) + payloads: Final = collect_files(sink, len(answered)) + assert len(matched_ids(payloads, answered)) == len(answered), tuple(sink.stored()) + assert sink.duplicated() == (), sink.duplicated() + assert store.connections() <= 2 * WORKERS, ( + f"{store.connections()} sink connections for {len(answered)} uploads" + ) + assert provider.drain() + + +def test_success_callback_mode_uploads_success_and_skips_failure(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = f"azure-{uuid.uuid4().hex[:8]}" + sink: Final = RecordingDataLakeSink() + cert, key = write_self_signed_cert(tmp_path, SINK_HOSTS) + with ( + wire_server(upstream_rejecting_fail_markers(500)) 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": FLUSH_SECONDS, + } + config: Final = azure_storage_config(tmp_path, callback_setting="success_callback") + with ( + owned_proxy(gateway, tmp_path, environment, config=config, workers=WORKERS) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + first_id, _ = _chat_completion(candidate, model, key, f"{marker}-a") + failed: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": f"fail-{marker}-b"}]}, + key=key, + ) + assert failed.status_code >= 500 and f"fail-{marker}-b" in failed.text, failed.text + third_id, _ = _chat_completion(candidate, model, key, f"{marker}-c") + collect_files(sink, 2) + landed: Final = frozenset(str(payload["id"]) for payload in sink.payloads().values()) + assert landed == frozenset({first_id, third_id}), tuple(sink.stored()) + assert all(f"fail-{marker}-b".encode() not in body for body in sink.stored().values()) + + +def test_failure_callback_mode_uploads_only_failures(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = f"azure-{uuid.uuid4().hex[:8]}" + sink: Final = RecordingDataLakeSink() + cert, key = write_self_signed_cert(tmp_path, SINK_HOSTS) + with ( + wire_server(upstream_rejecting_fail_markers(500)) 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": FLUSH_SECONDS, + } + config: Final = azure_storage_config(tmp_path, callback_setting="failure_callback") + with ( + owned_proxy(gateway, tmp_path, environment, config=config, workers=WORKERS) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + _chat_completion(candidate, model, key, f"{marker}-a") + failed: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": f"fail-{marker}-b"}]}, + key=key, + ) + assert failed.status_code >= 500 and f"fail-{marker}-b" in failed.text, failed.text + collect_files(sink, 1) + bodies: Final = tuple(sink.stored().values()) + assert len(bodies) == 1 and f"fail-{marker}-b".encode() in bodies[0], tuple(sink.stored()) + assert f"{marker}-a".encode() not in bodies[0] + + +def _sink_rejection_keeps_the_caller_and_proxy_healthy(gateway: Gateway, tmp_path: Path, status: int) -> None: + marker: Final = f"azure-{uuid.uuid4().hex[:8]}" + sink: Final = RecordingDataLakeSink(fail_status=status) + 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": FLUSH_SECONDS, + } + config: Final = azure_storage_config(tmp_path) + with ( + owned_proxy(gateway, tmp_path, environment, config=config, workers=WORKERS) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + _chat_completion(candidate, model, key, f"{marker}-a") + upload_rejected: Final = ( + (lambda methods: bool(methods)) + if status == 403 + else (lambda methods: any(method != "HEAD" for method in methods)) + ) + eventually(sink.rejected_methods, upload_rejected, seconds=30) + assert not sink.stored(), tuple(sink.stored()) + other_key: Final = scenario.key(models=[model]) + _chat_completion(candidate, model, other_key, f"{marker}-other") + readiness: Final = candidate.request("GET", "/health/readiness") + assert readiness.status_code == 200, readiness.text + sink.fail_status = 0 + third_id, _ = _chat_completion(candidate, model, key, f"{marker}-c") + eventually( + lambda: tuple(sink.payloads().values()), + lambda stored: third_id in {str(payload["id"]) for payload in stored}, + seconds=60, + ) + bodies: Final = tuple(sink.stored().values()) + assert all(f"{marker}-a".encode() not in body for body in bodies), f"{marker}-a should be lost, not retried" + + +def test_sink_403_keeps_the_caller_and_proxy_healthy(gateway: Gateway, tmp_path: Path) -> None: + _sink_rejection_keeps_the_caller_and_proxy_healthy(gateway, tmp_path, 403) + + +def test_sink_404_keeps_the_caller_and_proxy_healthy(gateway: Gateway, tmp_path: Path) -> None: + _sink_rejection_keeps_the_caller_and_proxy_healthy(gateway, tmp_path, 404) + + +def test_upstream_401_reaches_the_caller_and_lands_as_a_failure_payload(gateway: Gateway, tmp_path: Path) -> None: + marker: Final = f"azure-{uuid.uuid4().hex[:8]}" + sink: Final = RecordingDataLakeSink() + cert, key = write_self_signed_cert(tmp_path, SINK_HOSTS) + with ( + wire_server(upstream_rejecting_fail_markers(401)) 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": FLUSH_SECONDS, + } + config: Final = azure_storage_config(tmp_path) + with ( + owned_proxy(gateway, tmp_path, environment, config=config, workers=WORKERS) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + failed: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": model, "messages": [{"role": "user", "content": f"fail-{marker}"}]}, + key=key, + ) + assert failed.status_code == 401 and f"fail-{marker}" in failed.text, failed.text + payloads: Final = collect_files(sink, 1) + assert len(payloads) == 1 and f"fail-{marker}".encode() in next(iter(sink.stored().values())) + assert payloads[0]["status"] == "failure", payloads[0] + + +def test_unknown_model_lands_as_a_failure_payload(gateway: Gateway, tmp_path: Path) -> None: + 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": FLUSH_SECONDS, + } + config: Final = azure_storage_config(tmp_path) + with ( + owned_proxy(gateway, tmp_path, environment, config=config, workers=WORKERS) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + rejected: Final = candidate.request( + "POST", + "/v1/chat/completions", + {"model": "does-not-exist", "messages": [{"role": "user", "content": f"{marker}-unknown"}]}, + key=key, + ) + assert 400 <= rejected.status_code < 500 and "does-not-exist" in rejected.text, rejected.text + success_id, _ = _chat_completion(candidate, model, key, f"{marker}-ok") + payloads: Final = collect_files(sink, 2) + successful: Final = tuple(payload for payload in payloads if str(payload["id"]) == success_id) + failures: Final = tuple(payload for payload in payloads if payload["status"] == "failure") + assert len(successful) == 1 and len(failures) == 1, tuple(sink.stored()) + + +def test_missing_file_system_setting_fails_the_callback_init_and_keeps_the_proxy_serving( + gateway: Gateway, tmp_path: Path +) -> None: + 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 = { + name: value + for name, value in { + **azure_storage_environment(store.url, cert), + "DEFAULT_FLUSH_INTERVAL_SECONDS": FLUSH_SECONDS, + }.items() + if name != "AZURE_STORAGE_FILE_SYSTEM" + } + config: Final = azure_storage_config(tmp_path) + with ( + owned_proxy( + gateway, + tmp_path, + environment, + config=config, + remove_environment=("AZURE_STORAGE_FILE_SYSTEM",), + workers=WORKERS, + ) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + response_id, _ = _chat_completion(candidate, model, key, f"{marker}-ok") + _spend_row_visible(response_id) + assert store.connections() == 0, f"{store.connections()} sink connections without a configured sink" + + +def test_repeated_identical_requests_each_land_exactly_once(gateway: Gateway, tmp_path: Path) -> None: + 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": FLUSH_SECONDS, + } + config: Final = azure_storage_config(tmp_path) + with ( + owned_proxy(gateway, tmp_path, environment, config=config, workers=WORKERS) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + first_id, _ = _chat_completion(candidate, model, key, f"{marker}-a") + second_id, _ = _chat_completion(candidate, model, key, f"{marker}-b") + payloads: Final = collect_files(sink, 2) + landed: Final = frozenset(str(payload["id"]) for payload in payloads) + assert landed == frozenset({first_id, second_id}), tuple(sink.stored()) + assert sink.duplicated() == (), sink.duplicated() + received: Final = tuple(_marker_of(request) for request in provider.drain()) + assert received.count(f"{marker}-a") == 1 and received.count(f"{marker}-b") == 1, received + + +def test_disabled_callback_opens_no_sink_connection(gateway: Gateway, tmp_path: Path) -> None: + 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": FLUSH_SECONDS, + } + with ( + owned_proxy(gateway, tmp_path, environment, workers=WORKERS) as candidate, + candidate.scenario() as scenario, + ): + model: Final = scenario.model(api_base=provider.url + "/v1", api_key="synthetic-provider-key") + key: Final = scenario.key(models=[model]) + response_id, _ = _chat_completion(candidate, model, key, f"{marker}-ok") + _spend_row_visible(response_id) + assert store.connections() == 0, f"{store.connections()} sink connections with the callback disabled" + + +def test_files_upload_to_azure_storage_sibling_path_is_unchanged(gateway: Gateway, tmp_path: Path) -> None: + 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), + wire_server(sink.respond, tls=server_context(cert, key), keep_alive=True) as store, + ): + environment: Final = azure_storage_environment(store.url, cert) + config: Final = azure_storage_config(tmp_path) + with ( + owned_proxy(gateway, tmp_path, environment, config=config, workers=WORKERS) as candidate, + candidate.scenario() as scenario, + ): + key: Final = scenario.key() + content: Final = f'{{"marker": "{marker}"}}\n'.encode() + uploaded: Final = candidate.request_multipart( + "/v1/files", + {"purpose": "user_data", "target_storage": "azure_storage"}, + {"file": ("batch.jsonl", content, "application/jsonl")}, + key=key, + ) + assert uploaded.status_code == 200, uploaded.text + assert uploaded.json()["id"].startswith("file-"), uploaded.text + eventually( + lambda: any(content in body for body in sink.stored().values()), + lambda found: found, + seconds=30, + ) diff --git a/tests/unit/integrations/azure_storage/test_azure_storage.py b/tests/unit/integrations/azure_storage/test_azure_storage.py index dc162b5cd93..0227906a2dd 100644 --- a/tests/unit/integrations/azure_storage/test_azure_storage.py +++ b/tests/unit/integrations/azure_storage/test_azure_storage.py @@ -8,6 +8,7 @@ from unittest.mock import AsyncMock, MagicMock, patch import pytest +from litellm.constants import _DEFAULT_TTL_FOR_HTTPX_CLIENTS from litellm.integrations.azure_storage.azure_storage import ( AzureBlobStorageLogger, _cached_credential_chain_token_provider, @@ -371,6 +372,62 @@ async def test_service_client_defaults_to_commercial_endpoint(mock_env_vars): ) +def _fake_datalake_module() -> MagicMock: + fake_aio_module = MagicMock() + fake_aio_module.DataLakeServiceClient.side_effect = lambda **_: MagicMock(close=AsyncMock()) + return fake_aio_module + + +@pytest.mark.asyncio +async def test_service_client_is_reused_until_its_ttl_elapses(mock_env_vars): + """Within the TTL every upload must share one live client; closing a client + that is still in use by a concurrent upload fails that upload with an Azure + AuthenticationFailed error and drops the audit record""" + fake_aio_module = _fake_datalake_module() + now = 1_000_000.0 + + with patch.dict(sys.modules, {"azure.storage.filedatalake.aio": fake_aio_module}): + logger = AzureBlobStorageLogger(clock=lambda: now) + first = await logger.get_service_client() + second = await logger.get_service_client() + + assert second is first, "a second call inside the TTL must return the same client" + first.close.assert_not_awaited() + assert fake_aio_module.DataLakeServiceClient.call_count == 1 + + +@pytest.mark.asyncio +async def test_service_client_is_replaced_once_its_ttl_elapses(mock_env_vars): + fake_aio_module = _fake_datalake_module() + ticks = iter((1_000_000.0, 1_000_000.0 + _DEFAULT_TTL_FOR_HTTPX_CLIENTS + 1, 2_000_000.0)) + + with patch.dict(sys.modules, {"azure.storage.filedatalake.aio": fake_aio_module}): + logger = AzureBlobStorageLogger(clock=lambda: next(ticks)) + first = await logger.get_service_client() + second = await logger.get_service_client() + + assert second is not first, "an expired client must be closed and rebuilt" + first.close.assert_awaited_once() + second.close.assert_not_awaited() + assert fake_aio_module.DataLakeServiceClient.call_count == 2 + + +@pytest.mark.asyncio +async def test_service_client_is_replaced_at_the_exact_ttl_boundary(mock_env_vars): + fake_aio_module = _fake_datalake_module() + ticks = iter((1_000_000.0, 1_000_000.0 + _DEFAULT_TTL_FOR_HTTPX_CLIENTS, 2_000_000.0)) + + with patch.dict(sys.modules, {"azure.storage.filedatalake.aio": fake_aio_module}): + logger = AzureBlobStorageLogger(clock=lambda: next(ticks)) + first = await logger.get_service_client() + second = await logger.get_service_client() + + assert second is not first, "a call exactly at the TTL must rebuild the client" + first.close.assert_awaited_once() + second.close.assert_not_awaited() + assert fake_aio_module.DataLakeServiceClient.call_count == 2 + + @pytest.mark.parametrize( ("payload_id", "expected"), (