From 978d1ece6b9d1c948234d21292e189de7a568527 Mon Sep 17 00:00:00 2001 From: Yucheng Zhu Date: Mon, 20 Jul 2026 18:57:32 -0700 Subject: [PATCH] fix(ptu): address greptile p1/p2 + async-client CI check Currency guard (P1): the rollup now inspects azure_fetcher.last_currency after each call and skips writes when the response reports a non-USD currency, so a customer on a non-USD Azure subscription cannot land a raw foreign-currency amount in ptu_flat_cost as if it were USD. The warning names the reservation, day, and currency so operators can spot the skip and file a follow-up for currency conversion. Client hardening (P2): - client_secret is field(repr=False) so repr(config) no longer exposes it - httpx.RequestError family (ConnectError, ReadTimeout, RemoteProtocolError) is wrapped into AzureCostManagementError so the client honors its documented interface for every failure path, not just HTTPStatusError - Removed the unreachable AttributeError branch in _parse_cost_and_currency - Removed the unused AzureCostManagementError import in proxy_server.py CI (ensure_async_clients_test): the client no longer instantiates AsyncHTTPHandler directly; it goes through get_async_httpx_client with a new httpxSpecialProvider.AzureCostManagement enum value so it shares the same cached connection pool contract as the rest of the integrations. Behavior changes: azure_billing reservations that receive a non-USD Cost Management response are now skipped with a warning instead of writing the raw amount; POST /ptu_reservation/new with cost_source=azure_billing and no azure_resource_id continues to return 400 (stage 1 test updated to assert the new error message rather than the old blanket rejection). --- .../azure_cost_management_client.py | 22 ++++++------ litellm/proxy/proxy_server.py | 5 +-- .../spend_tracking/ptu_reservation_rollup.py | 27 +++++++++----- litellm/types/llms/custom_http.py | 1 + .../test_azure_cost_management_client.py | 32 +++++++++++++++-- .../test_ptu_reservation_endpoints.py | 5 +-- .../test_ptu_reservation_rollup.py | 36 +++++++++++++++---- 7 files changed, 94 insertions(+), 34 deletions(-) diff --git a/litellm/integrations/azure_cost_management/azure_cost_management_client.py b/litellm/integrations/azure_cost_management/azure_cost_management_client.py index c7e778fd227..f92e6d6de27 100644 --- a/litellm/integrations/azure_cost_management/azure_cost_management_client.py +++ b/litellm/integrations/azure_cost_management/azure_cost_management_client.py @@ -8,18 +8,19 @@ one public method the rollup needs, not a general Cost Management SDK. from __future__ import annotations import os -from dataclasses import dataclass +from dataclasses import dataclass, field from datetime import date, datetime, time, timedelta, timezone from typing import Any, Callable, Optional import httpx from litellm._logging import verbose_logger -from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler, get_async_httpx_client +from litellm.types.llms.custom_http import httpxSpecialProvider class AzureCostManagementError(Exception): - """Raised when the Cost Management API returns a non-2xx or an unexpected payload.""" + """Raised when the Cost Management API returns a non-2xx, network failure, or an unexpected payload.""" @dataclass(frozen=True, slots=True) @@ -27,7 +28,7 @@ class AzureCostManagementConfig: subscription_id: str tenant_id: str client_id: str - client_secret: str + client_secret: str = field(repr=False) api_version: str = "2023-11-01" @classmethod @@ -88,7 +89,7 @@ class AzureCostManagementClient: token_provider: Optional[TokenProvider] = None, ) -> None: self._config = config - self._http = http_handler or AsyncHTTPHandler() + self._http = http_handler or get_async_httpx_client(httpxSpecialProvider.AzureCostManagement) self._token_provider = token_provider or _default_token_provider_factory(config) self._last_currency: Optional[str] = None @@ -145,15 +146,14 @@ class AzureCostManagementClient: raise AzureCostManagementError( f"Azure Cost Management HTTP {exc.response.status_code}: {exc.response.text}" ) from exc + except httpx.RequestError as exc: + raise AzureCostManagementError(f"Azure Cost Management network error: {exc}") from exc return self._parse_cost_and_currency(response.json()) def _parse_cost_and_currency(self, payload: Any) -> float: - try: - properties = payload.get("properties", {}) if isinstance(payload, dict) else {} - rows = properties.get("rows") or [] - columns = properties.get("columns") or [] - except AttributeError as exc: - raise AzureCostManagementError(f"Unexpected payload shape: {payload!r}") from exc + properties = payload.get("properties", {}) if isinstance(payload, dict) else {} + rows = properties.get("rows") or [] + columns = properties.get("columns") or [] if not rows: self._last_currency = None diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 07e810ff648..63ab2085b20 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -7480,10 +7480,7 @@ def _build_azure_cost_fetcher_if_enabled() -> Optional[Any]: ) return None try: - from litellm.integrations.azure_cost_management import ( - AzureCostManagementClient, - AzureCostManagementError, - ) + from litellm.integrations.azure_cost_management import AzureCostManagementClient from litellm.integrations.azure_cost_management.azure_cost_management_client import ( AzureCostManagementConfig, ) diff --git a/litellm/proxy/spend_tracking/ptu_reservation_rollup.py b/litellm/proxy/spend_tracking/ptu_reservation_rollup.py index c1d3986b5e8..e9f91ef1d00 100644 --- a/litellm/proxy/spend_tracking/ptu_reservation_rollup.py +++ b/litellm/proxy/spend_tracking/ptu_reservation_rollup.py @@ -64,14 +64,6 @@ async def _compute_daily_flat_cost( return 0.0 try: fetched = await azure_fetcher.get_daily_cost(reservation.azure_resource_id, day) - verbose_proxy_logger.info( - "PTU rollup: azure_billing reservation=%s day=%s resource=%s returned $%.4f", - getattr(reservation, "id", "?"), - day.isoformat(), - reservation.azure_resource_id, - fetched, - ) - return fetched except Exception as exc: # noqa: BLE001 # log and continue; one bad reservation must not stop the batch verbose_proxy_logger.error( "PTU rollup: azure fetch failed for reservation=%s day=%s: %s", @@ -80,6 +72,25 @@ async def _compute_daily_flat_cost( exc, ) return 0.0 + currency = getattr(azure_fetcher, "last_currency", None) + if currency is not None and currency.upper() != "USD": + verbose_proxy_logger.warning( + "PTU rollup: azure_billing reservation=%s day=%s returned currency=%s; " + "skipping write to avoid storing non-USD amount as USD (LiteLLM_DailyTeamSpend " + "assumes USD). File a follow-up for currency conversion.", + getattr(reservation, "id", "?"), + day.isoformat(), + currency, + ) + return 0.0 + verbose_proxy_logger.info( + "PTU rollup: azure_billing reservation=%s day=%s resource=%s returned $%.4f", + getattr(reservation, "id", "?"), + day.isoformat(), + reservation.azure_resource_id, + fetched, + ) + return fetched verbose_proxy_logger.warning( "PTU rollup: unknown cost_source=%s on reservation=%s; skipping", reservation.cost_source, diff --git a/litellm/types/llms/custom_http.py b/litellm/types/llms/custom_http.py index d7dfc0e486b..4d0b6ff6cd1 100644 --- a/litellm/types/llms/custom_http.py +++ b/litellm/types/llms/custom_http.py @@ -30,6 +30,7 @@ class httpxSpecialProvider(str, Enum): PromptManagement = "prompt_management" UI = "ui" Sandbox = "sandbox" + AzureCostManagement = "azure_cost_management" VerifyTypes = Union[str, bool, ssl.SSLContext] diff --git a/tests/test_litellm/integrations/azure_cost_management/test_azure_cost_management_client.py b/tests/test_litellm/integrations/azure_cost_management/test_azure_cost_management_client.py index 3d4dae94567..385ae15f3cc 100644 --- a/tests/test_litellm/integrations/azure_cost_management/test_azure_cost_management_client.py +++ b/tests/test_litellm/integrations/azure_cost_management/test_azure_cost_management_client.py @@ -97,9 +97,7 @@ async def test_get_daily_cost_raises_on_http_error(): response = MagicMock() response.status_code = 403 response.text = "forbidden" - http.post = AsyncMock( - side_effect=httpx.HTTPStatusError("forbidden", request=MagicMock(), response=response) - ) + http.post = AsyncMock(side_effect=httpx.HTTPStatusError("forbidden", request=MagicMock(), response=response)) client = AzureCostManagementClient( config=_config(), @@ -143,3 +141,31 @@ def test_config_from_env_raises_when_creds_missing(monkeypatch): with pytest.raises(AzureCostManagementError) as exc: AzureCostManagementConfig.from_env(subscription_id="sub-x") assert "AZURE_TENANT_ID" in str(exc.value) + + +def test_client_secret_not_in_repr(): + """Regression: client_secret must not surface in repr(config) — the field is repr=False.""" + cfg = AzureCostManagementConfig( + subscription_id="sub-x", + tenant_id="tenant", + client_id="client", + client_secret="super-secret-value", + ) + assert "super-secret-value" not in repr(cfg) + assert "client_secret" not in repr(cfg) + + +@pytest.mark.asyncio +async def test_get_daily_cost_wraps_network_errors(): + """httpx.RequestError family (ConnectError, ReadTimeout, RemoteProtocolError) must surface as AzureCostManagementError.""" + http = MagicMock() + http.post = AsyncMock(side_effect=httpx.ReadTimeout("upstream timeout")) + client = AzureCostManagementClient( + config=_config(), + http_handler=http, + token_provider=lambda: "fake-token", + ) + + with pytest.raises(AzureCostManagementError) as exc: + await client.get_daily_cost("/subs/x/deploy/y", date(2026, 7, 15)) + assert "network error" in str(exc.value).lower() diff --git a/tests/test_litellm/proxy/management_endpoints/test_ptu_reservation_endpoints.py b/tests/test_litellm/proxy/management_endpoints/test_ptu_reservation_endpoints.py index 0fcba1074dd..f95ecee1fb9 100644 --- a/tests/test_litellm/proxy/management_endpoints/test_ptu_reservation_endpoints.py +++ b/tests/test_litellm/proxy/management_endpoints/test_ptu_reservation_endpoints.py @@ -171,7 +171,8 @@ async def test_new_reservation_rejects_effective_to_equal_from(client_and_mocks) @pytest.mark.asyncio -async def test_new_reservation_rejects_azure_billing_mode(client_and_mocks): +async def test_new_reservation_rejects_azure_billing_without_resource_id(client_and_mocks): + """Stage 6 relaxed the blanket azure_billing reject; azure_resource_id is now the required field.""" client, _, mock_table = client_and_mocks resp = client.post( "/ptu_reservation/new", @@ -179,11 +180,11 @@ async def test_new_reservation_rejects_azure_billing_mode(client_and_mocks): "team_id": "team_x", "model": "gpt-4", "cost_source": "azure_billing", - "azure_resource_id": "/subscriptions/x/deployments/gpt-4-ptu", "effective_from": _iso(), }, ) assert resp.status_code == 400, resp.text + assert "azure_resource_id" in resp.text mock_table.create.assert_not_awaited() diff --git a/tests/test_litellm/proxy/spend_tracking/test_ptu_reservation_rollup.py b/tests/test_litellm/proxy/spend_tracking/test_ptu_reservation_rollup.py index 27324d91ea6..02bbb455422 100644 --- a/tests/test_litellm/proxy/spend_tracking/test_ptu_reservation_rollup.py +++ b/tests/test_litellm/proxy/spend_tracking/test_ptu_reservation_rollup.py @@ -119,6 +119,7 @@ async def test_compute_flat_cost_azure_billing_uses_fetcher_result(): r = _r(cost_source="azure_billing", ptu_count=None, cost_per_ptu=None, azure_resource_id="/subs/x/deploy/y") fetcher = MagicMock() fetcher.get_daily_cost = AsyncMock(return_value=42.5) + fetcher.last_currency = "USD" result = await _compute_daily_flat_cost(r, date(2026, 7, 15), azure_fetcher=fetcher) @@ -126,6 +127,32 @@ async def test_compute_flat_cost_azure_billing_uses_fetcher_result(): fetcher.get_daily_cost.assert_awaited_once_with("/subs/x/deploy/y", date(2026, 7, 15)) +@pytest.mark.asyncio +async def test_compute_flat_cost_azure_billing_returns_zero_on_non_usd_currency(): + """Non-USD Azure response must not be persisted as USD (silent data corruption).""" + r = _r(cost_source="azure_billing", ptu_count=None, cost_per_ptu=None, azure_resource_id="/subs/x/deploy/y") + fetcher = MagicMock() + fetcher.get_daily_cost = AsyncMock(return_value=99.9) + fetcher.last_currency = "EUR" + + result = await _compute_daily_flat_cost(r, date(2026, 7, 15), azure_fetcher=fetcher) + + assert result == 0.0 + + +@pytest.mark.asyncio +async def test_compute_flat_cost_azure_billing_accepts_missing_currency(): + """Empty Azure response (no rows) sets last_currency=None; treat as USD path (0.0 is written by the outer skip).""" + r = _r(cost_source="azure_billing", ptu_count=None, cost_per_ptu=None, azure_resource_id="/subs/x/deploy/y") + fetcher = MagicMock() + fetcher.get_daily_cost = AsyncMock(return_value=0.0) + fetcher.last_currency = None + + result = await _compute_daily_flat_cost(r, date(2026, 7, 15), azure_fetcher=fetcher) + + assert result == 0.0 + + @pytest.mark.asyncio async def test_compute_flat_cost_azure_billing_fetcher_error_returns_zero(): r = _r(cost_source="azure_billing", ptu_count=None, cost_per_ptu=None, azure_resource_id="/subs/x/deploy/y") @@ -158,10 +185,9 @@ async def test_rollup_azure_billing_reservation_writes_fetched_amount(mock_prism mock_reservation.find_many = AsyncMock(return_value=[reservation]) fetcher = MagicMock() fetcher.get_daily_cost = AsyncMock(return_value=150.0) + fetcher.last_currency = "USD" - result = await run_ptu_reservation_rollup( - prisma, target_date=date(2026, 7, 12), azure_fetcher=fetcher - ) + result = await run_ptu_reservation_rollup(prisma, target_date=date(2026, 7, 12), azure_fetcher=fetcher) assert result.rows_written == 1 fetcher.get_daily_cost.assert_awaited_once_with("/subs/x/deploy/y", date(2026, 7, 12)) @@ -297,9 +323,7 @@ async def test_rollup_queries_active_reservations_at_day_start(mock_prisma): @pytest.mark.asyncio async def test_rollup_idempotent_second_run_upserts_same_row(mock_prisma): prisma, mock_daily, mock_reservation = mock_prisma - mock_reservation.find_many = AsyncMock( - return_value=[_r(id="res_1", ptu_count=1, cost_per_ptu=200.0)] - ) + mock_reservation.find_many = AsyncMock(return_value=[_r(id="res_1", ptu_count=1, cost_per_ptu=200.0)]) r1 = await run_ptu_reservation_rollup(prisma, target_date=date(2026, 7, 12)) r2 = await run_ptu_reservation_rollup(prisma, target_date=date(2026, 7, 12))