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).
This commit is contained in:
Yucheng Zhu 2026-07-20 18:57:32 -07:00
parent 152e233722
commit 978d1ece6b
7 changed files with 94 additions and 34 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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