diff --git a/litellm/proxy/management_endpoints/roi_calculator_endpoints.py b/litellm/proxy/management_endpoints/roi_calculator_endpoints.py index 22ee17bbc5f..593d14bfaac 100644 --- a/litellm/proxy/management_endpoints/roi_calculator_endpoints.py +++ b/litellm/proxy/management_endpoints/roi_calculator_endpoints.py @@ -110,7 +110,7 @@ async def get_roi_config_repository( status_code=500, detail=CommonProxyErrors.db_not_connected_error.value, ) - return ConfigRepository(prisma_client) + return ConfigRepository(prisma_client, use_writer=True) def get_roi_sync_manager() -> SyncManager: @@ -561,7 +561,7 @@ async def run_scheduled_sync() -> None: if prisma_client is None: return - repository: Final = ConfigRepository(prisma_client) + repository: Final = ConfigRepository(prisma_client, use_writer=True) settings: Final = await _load_settings(repository) if not settings.update_interval_minutes or not _public_settings(settings).ready: return diff --git a/litellm/proxy/roi_calculator/github.py b/litellm/proxy/roi_calculator/github.py index 4ee837498dc..6fa2ce24d63 100644 --- a/litellm/proxy/roi_calculator/github.py +++ b/litellm/proxy/roi_calculator/github.py @@ -37,6 +37,7 @@ class _GitHubHead(_GitHubModel): class GitHubPullListItem(_GitHubModel): number: int + html_url: str = "" merged_at: str | None = None updated_at: str title: str diff --git a/litellm/proxy/roi_calculator/sync.py b/litellm/proxy/roi_calculator/sync.py index 893b11ac007..645d9ea60a8 100644 --- a/litellm/proxy/roi_calculator/sync.py +++ b/litellm/proxy/roi_calculator/sync.py @@ -210,6 +210,35 @@ async def _estimate_with_fallback( return estimate +async def _unavailable_record(github: GitHub, repo: str, pull: GitHubPullListItem, error: SourceError) -> ROIPullRecord: + login: Final = pull.user.login if pull.user and pull.user.login else "deleted-user" + profile: Final = await github.profile_email(login) + estimate: Final[ROIEstimate] = { + "status": "needs_review", + "hours": None, + "reasoning": f"PR metadata could not be read: {error} Run analysis again to retry this PR.", + } + return ROIPullRecord( + repo=repo, + number=pull.number, + title=pull.title, + url=pull.html_url, + login=login, + emails=(profile,) if profile else (), + profile_email=profile, + commit_emails=(), + merged_at=pull.merged_at or pull.updated_at, + head_sha=pull.head.sha if pull.head else "", + additions=0, + deletions=0, + changed_files=0, + commit_count=0, + incomplete_metadata=True, + estimate=estimate, + cache_key=None, + ) + + class SyncManager: def __init__( self, @@ -398,7 +427,12 @@ class SyncManager: ) self._update_estimate_progress(cached_record["estimate"]) return index, cached_record - evidence: Final = await github.evidence(repo, pull) + try: + evidence: Final = await github.evidence(repo, pull) + except SourceError as exc: + unavailable: Final = await _unavailable_record(github, repo, pull, exc) + self._update_estimate_progress(unavailable["estimate"]) + return index, unavailable estimate: Final = await _estimate_with_fallback(estimator, evidence) evidence_item: Final = GitHubPullListItem.model_validate( MappingProxyType( diff --git a/litellm/proxy/roi_calculator/sync_store.py b/litellm/proxy/roi_calculator/sync_store.py index c074b0f16d5..43a2533eb59 100644 --- a/litellm/proxy/roi_calculator/sync_store.py +++ b/litellm/proxy/roi_calculator/sync_store.py @@ -127,7 +127,14 @@ class SyncStore: async def cancel(self) -> None: await self._db.execute_raw( - """UPDATE "LiteLLM_Config" SET param_value = jsonb_set(param_value, '{cancel}', 'true'::jsonb) + """UPDATE "LiteLLM_Config" + SET param_value = param_value || jsonb_build_object( + 'cancel', true, 'owner', '', + 'status', (param_value->'status') || jsonb_build_object( + 'running', false, 'phase', 'cancelled', 'stage', 'Sync cancelled', + 'finished_at', to_char(NOW() AT TIME ZONE 'UTC', 'YYYY-MM-DD"T"HH24:MI:SS.US"+00:00"') + ) + ), last_run_at = NOW() WHERE param_name = $1 AND param_value->'status'->>'running' = 'true' """, _SYNC_KEY, ) diff --git a/litellm/repositories/config_repository.py b/litellm/repositories/config_repository.py index 8b8280622fd..c5674a4b398 100644 --- a/litellm/repositories/config_repository.py +++ b/litellm/repositories/config_repository.py @@ -44,8 +44,9 @@ class ConfigParam: class ConfigRepository: """Repository for config database operations.""" - def __init__(self, prisma_client: PrismaClient | None): + def __init__(self, prisma_client: PrismaClient | None, *, use_writer: bool = False): self._prisma_client: Final = prisma_client + self._use_writer: Final = use_writer @property def prisma_client(self) -> PrismaClient: @@ -55,7 +56,8 @@ class ConfigRepository: @property def _config_table(self) -> _ConfigTable: - return cast(_ConfigTable, self.prisma_client.db.litellm_config) + database: Final = self.prisma_client.writer_db if self._use_writer else self.prisma_client.db + return cast(_ConfigTable, database.litellm_config) @property def table(self) -> _ConfigTable: diff --git a/tests/integration/database/test_roi_sync_store.py b/tests/integration/database/test_roi_sync_store.py index 7f3762af8a7..8caab0fa2ad 100644 --- a/tests/integration/database/test_roi_sync_store.py +++ b/tests/integration/database/test_roi_sync_store.py @@ -9,6 +9,7 @@ from litellm.proxy.common_utils.user_api_key_cache import UserApiKeyCache from litellm.proxy.roi_calculator.sample import sample_report from litellm.proxy.roi_calculator.sync_store import SyncStore from litellm.proxy.utils import PrismaClient, ProxyLogging +from litellm.repositories.config_repository import ConfigRepository from litellm.types.roi_calculator import ROIPullRecord, ROIReport, ROISyncStatus from tests.integration._support.database import read_rows, scratch_database, write_rows @@ -18,7 +19,7 @@ async def test_roi_cache_survives_scope_changes_and_uses_writer(monkeypatch: pyt with scratch_database() as writer_url, scratch_database() as reader_url: write_rows( 'CREATE TABLE "LiteLLM_Config" (param_name TEXT PRIMARY KEY, param_value JSONB NOT NULL, ' - "last_run_at TIMESTAMP NOT NULL DEFAULT NOW())", + "last_run_at TIMESTAMP NOT NULL DEFAULT NOW(), reload_revision BIGINT NOT NULL DEFAULT 0)", (), database_url=writer_url, ) @@ -29,6 +30,13 @@ async def test_roi_cache_survives_scope_changes_and_uses_writer(monkeypatch: pyt await client.connect() try: store: Final = SyncStore(client) + repository: Final = ConfigRepository(client, use_writer=True) + await repository.set_param("roi_calculator_settings", '{"repos":["example/repo"]}') + settings_row: Final = await repository.get_param("roi_calculator_settings") + assert settings_row is not None + assert TypeAdapter(dict[str, tuple[str, ...]]).validate_python(settings_row.param_value)["repos"] == ( + "example/repo", + ) report: Final = sample_report(datetime(2026, 9, 30, tzinfo=timezone.utc)) pull: Final[ROIPullRecord] = { **report["pulls"][0], @@ -70,6 +78,12 @@ async def test_roi_cache_survives_scope_changes_and_uses_writer(monkeypatch: pyt database_url=writer_url, ) ) == ("roi_calculator_pull_new", "roi_calculator_pull_outside-window") + published: Final = await repository.get_param("roi_calculator_report") + assert published is not None + assert TypeAdapter(ROIReport).validate_python(published.param_value)["pulls"] == (pull,) + cached: Final = await repository.get_param("roi_calculator_pull_new") + assert cached is not None + assert TypeAdapter(ROIPullRecord).validate_python(cached.param_value)["cache_key"] == "new" assert not await store.acquire("scheduled", running, 1440) assert await store.acquire("manual", running) write_rows( @@ -94,5 +108,12 @@ async def test_roi_cache_survives_scope_changes_and_uses_writer(monkeypatch: pyt ) == 2 ) + assert await store.acquire("remote", running) + await store.cancel() + cancelled: Final = await store.status() + assert cancelled is not None and cancelled.phase == "cancelled" and not cancelled.running + assert not await store.heartbeat("remote", running) + assert not await store.finish("remote", complete, narrowed) + assert await store.acquire("after-cancel", running) finally: await client.disconnect() diff --git a/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py b/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py index f3b4808926c..eaf896831b1 100644 --- a/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py @@ -1,3 +1,4 @@ +import asyncio import json from collections.abc import Mapping from datetime import datetime, timezone @@ -19,7 +20,7 @@ from litellm.proxy.management_endpoints.roi_calculator_endpoints import ( ) from litellm.proxy.roi_calculator.estimator import estimator_options from litellm.proxy.roi_calculator.sample import sample_report -from litellm.types.roi_calculator import ROISettings, ROISyncStatus +from litellm.types.roi_calculator import ROIReport, ROISettings, ROISyncStatus _JSON_HEADERS: Final = MappingProxyType({"content-type": "application/json"}) @@ -209,3 +210,29 @@ def test_schedule_normalizes_legacy_and_offset_timestamps(anchor: str) -> None: ) report: Final = sample_report(datetime(2026, 9, 30, tzinfo=timezone.utc)) assert _next_update(settings, status, report) == datetime(2026, 9, 30, 13, tzinfo=timezone.utc) + + +def test_manual_match_recalculates_saved_report_and_removal_restores_cohort() -> None: + repository: Final = _ConfigRepository() + report: Final[ROIReport] = {**sample_report(datetime(2026, 9, 30, tzinfo=timezone.utc)), "mode": "live"} + serialized: Final = TypeAdapter(dict[str, object]).validate_json(TypeAdapter(ROIReport).dump_json(report)) + asyncio.run(repository.set_param("roi_calculator_report", serialized)) + client: Final = _client(LitellmUserRoles.PROXY_ADMIN, repository) + before: Final = client.get("/roi-calculator/report") + assert before.status_code == 200 + assert before.json()["report"]["metrics"]["output_hours"] == 10.5 + matched: Final = client.put( + "/roi-calculator/identity-map", + content='{"github_login":" CASEY ","email":"Alex@Example.com"}', + headers=_JSON_HEADERS, + ) + assert matched.status_code == 200 + assert matched.json()["identity_map"]["casey"] == "alex@example.com" + assert matched.json()["report"]["metrics"]["output_hours"] == 16 + assert matched.json()["report"]["metrics"]["cost_per_hour"] == pytest.approx(31 / 16) + removed: Final = client.put( + "/roi-calculator/identity-map", content='{"github_login":"casey","email":null}', headers=_JSON_HEADERS + ) + assert removed.status_code == 200 + assert not removed.json()["identity_map"] + assert removed.json()["report"]["metrics"] == before.json()["report"]["metrics"] diff --git a/tests/unit/proxy/roi_calculator/test_sync.py b/tests/unit/proxy/roi_calculator/test_sync.py index 0d3ddffa3a2..9b80ae54e5a 100644 --- a/tests/unit/proxy/roi_calculator/test_sync.py +++ b/tests/unit/proxy/roi_calculator/test_sync.py @@ -10,6 +10,7 @@ import pytest from pydantic import TypeAdapter from litellm.proxy.roi_calculator.estimator import CompletionCaller +from litellm.proxy.roi_calculator.github import GitHubPullListItem from litellm.proxy.roi_calculator.sync import SpendReader, SyncManager, read_spend from litellm.types.roi_calculator import ( ROICompletionRequest, @@ -274,7 +275,7 @@ async def test_read_spend_joins_user_emails_and_preserves_unmatched_identities() @pytest.mark.asyncio -async def test_sync_does_not_persist_a_report_when_github_fails() -> None: +async def test_unreadable_pr_is_reported_and_retried_on_next_run() -> None: repository: Final = _ReportRepository() manager: Final = SyncManager(clock=_fixed_now) @@ -287,8 +288,19 @@ async def test_sync_does_not_persist_a_report_when_github_fails() -> None: ) await _wait_until_finished(manager) - assert not repository.values - assert manager.status.phase == "error" + assert manager.status.phase == "complete" + assert manager.status.needs_attention == 1 + failed: Final = TypeAdapter(ROIReport).validate_python(repository.values["roi_calculator_report"]) + assert failed["pulls"][0]["estimate"]["status"] == "needs_review" + assert failed["pulls"][0]["estimate"]["hours"] is None + assert failed["pulls"][0]["incomplete_metadata"] is True + assert failed["pulls"][0]["cache_key"] is None + assert await manager.start(_settings(), repository, _spend_reader(), _completion(), _transport()) + await _wait_until_finished(manager) + recovered: Final = TypeAdapter(ROIReport).validate_python(repository.values["roi_calculator_report"]) + assert recovered["pulls"][0]["estimate"]["status"] == "estimated" + assert recovered["pulls"][0]["estimate"]["hours"] == 4 + assert manager.status.reused == 0 @pytest.mark.asyncio @@ -411,3 +423,31 @@ async def test_expired_lease_can_restart_without_restarting_the_gateway() -> Non assert cancelled.is_set() assert manager.status.phase == "complete" assert manager.status.estimated == 1 + + +@pytest.mark.asyncio +async def test_one_unreadable_pr_preserves_other_estimates_in_report() -> None: + baseline: Final = _transport() + listed: Final = TypeAdapter(tuple[GitHubPullListItem, ...]).validate_json(_PULL_LIST_JSON)[0] + second: Final = listed.model_copy(update=MappingProxyType({"number": 43})) + listing: Final = TypeAdapter(tuple[GitHubPullListItem, ...]).dump_json((listed, second)) + + def respond(request: httpx.Request) -> httpx.Response: + if request.url.path == "/repos/org/repo/pulls": + return httpx.Response(200, content=listing) + if request.url.path == "/repos/org/repo/pulls/43": + return httpx.Response(404) + return baseline.handle_request(request) + + repository: Final = _ReportRepository() + manager: Final = SyncManager(clock=_fixed_now) + assert await manager.start(_settings(), repository, _spend_reader(), _completion(), httpx.MockTransport(respond)) + await _wait_until_finished(manager) + report: Final = TypeAdapter(ROIReport).validate_python(repository.values["roi_calculator_report"]) + assert tuple((pull["number"], pull["estimate"]["status"]) for pull in report["pulls"]) == ( + (42, "estimated"), + (43, "needs_review"), + ) + assert manager.status.phase == "complete" + assert manager.status.estimated == 1 + assert manager.status.needs_attention == 1