mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-06 02:48:13 +00:00
fix(roi): fence cancelled syncs and read reports from writer
This commit is contained in:
parent
b3aacaf49d
commit
bb3a380877
8 changed files with 143 additions and 11 deletions
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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"]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue