mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
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:
parent
152e233722
commit
978d1ece6b
7 changed files with 94 additions and 34 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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))
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue