fix(roi): fence cancelled syncs and read reports from writer

This commit is contained in:
moe-berri 2026-09-30 11:03:19 -07:00
parent b3aacaf49d
commit bb3a380877
8 changed files with 143 additions and 11 deletions

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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

View file

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