fix(roi): preserve cached estimates across report scope changes

This commit is contained in:
moe-berri 2026-09-30 10:42:27 -07:00
parent 4007f79496
commit 3cf6cd7a5c
24 changed files with 191 additions and 80 deletions

Binary file not shown.

After

Width:  |  Height:  |  Size: 80 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 58 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 63 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 70 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 47 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 76 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 75 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 39 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 70 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 93 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 73 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 39 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 81 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 72 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 76 KiB

Binary file not shown.

After

Width:  |  Height:  |  Size: 50 KiB

View file

@ -256,7 +256,9 @@ def _gateway_http_client() -> AsyncHTTPHandler:
return get_async_httpx_client(
llm_provider="roi_calculator",
params={"transport": _gateway_transport(app), "timeout": 180, "follow_redirects": False},
params=TypeAdapter(dict[str, object]).validate_python(
MappingProxyType({"transport": _gateway_transport(app), "timeout": 180, "follow_redirects": False})
),
)
@ -425,7 +427,7 @@ async def get_roi_calculator_sync_status(
settings: Final = await _load_settings(repository)
report: Final = await _load_report(repository)
next_update: Final = _next_update(settings, status, report)
return status.model_copy(update={"next_update": next_update.isoformat() if next_update else None})
return status.model_copy(update=MappingProxyType({"next_update": next_update.isoformat() if next_update else None}))
@router.post(
@ -512,10 +514,10 @@ async def update_roi_calculator_identity_map(
new_email: Final = normalize_email(update.email)
if not login or (update.email is not None and not new_email):
raise HTTPException(status_code=422, detail="Enter a GitHub login and a valid email address.")
identity_map: Final[Mapping[str, str]] = MappingProxyType(
{key: value for key, value in current.identity_map.items() if update.email is None and key != login}
identity_map: Final[Mapping[str, str]] = (
MappingProxyType({key: value for key, value in current.identity_map.items() if key != login})
if update.email is None
else {**current.identity_map, login: new_email}
else MappingProxyType({**current.identity_map, login: new_email})
)
settings: Final = ROISettings(
github_api_url=current.github_api_url,
@ -622,9 +624,11 @@ async def reset_roi_calculator_setup(
try:
current: Final = await _load_settings(repository)
stored: Final = await _load_stored_settings(repository)
settings: Final = current.model_copy(update={"repos": ()})
settings: Final = current.model_copy(update=MappingProxyType({"repos": ()}))
await _save_settings(repository, settings, stored.github_token, stored.estimator_key)
await store.clear_report()
return _public_settings(settings)
finally:
await store.finish(owner, status.model_copy(update={"running": False, "phase": "idle", "stage": "Idle"}))
await store.finish(
owner, status.model_copy(update=MappingProxyType({"running": False, "phase": "idle", "stage": "Idle"}))
)

View file

@ -3,6 +3,7 @@ import json
from collections.abc import Awaitable
from typing import Final, Literal, Protocol, TypeAlias
import httpx
from pydantic import ValidationError
from typing_extensions import NotRequired, ReadOnly, TypedDict
@ -174,7 +175,7 @@ class Estimator:
if choice.finish_reason not in (None, "stop") or choice.message.content is None:
raise ValueError("incomplete estimator response")
result: Final = ROIEstimatorResult.model_validate_json(extract_classifier_json(choice.message.content))
except Exception:
except (httpx.HTTPError, ValueError, IndexError):
raise SourceError(
"The estimator did not return valid hours and reasoning. Check the selected model and prompt."
) from None

View file

@ -9,7 +9,9 @@ import httpx
from pydantic import BaseModel, ConfigDict, Field, TypeAdapter
from typing_extensions import ReadOnly, TypedDict
from litellm.llms.custom_httpx.http_handler import get_async_httpx_client
from litellm.llms.custom_httpx.http_handler import (
get_async_httpx_client, # pyright: ignore[reportUnknownVariableType] # shared client factory has untyped params
)
from litellm.proxy.roi_calculator.analytics import normalize_email
from litellm.types.llms.custom_http import httpxSpecialProvider
from litellm.types.roi_calculator import ROIPullCommit, ROIPullEvidence, ROIPullFile, ROISettings
@ -273,7 +275,7 @@ async def _fetch_page(
)
try:
parsed: Final[tuple[_T, ...]] = adapter.validate_python(response.json())
except Exception:
except ValueError:
raise SourceError(error_message) from None
return parsed, 'rel="next"' in response.headers.get("link", "")
@ -325,11 +327,9 @@ class GitHub:
else MappingProxyType({"Accept": "application/vnd.github+json"})
)
self._api_url: Final = settings.github_api_url.rstrip("/")
client_params: Final[dict[str, object]] = {
"timeout": 45,
"follow_redirects": False,
**({"transport": transport} if transport is not None else {}),
}
client_params: Final = TypeAdapter(dict[str, object]).validate_python(
MappingProxyType({"timeout": 45, "follow_redirects": False, "transport": transport})
)
self.client: Final[httpx.AsyncClient] = (
client
if client is not None
@ -395,8 +395,8 @@ class GitHub:
later_matches, later_has_more = await search_pages(github_page + 1, pages_remaining - 1)
return (*matches, *later_matches), later_has_more
matches, has_more = await search_pages(first_github_page, _REPOSITORY_SEARCH_PAGES)
return _repository_values(matches), has_more
matches, search_has_more = await search_pages(first_github_page, _REPOSITORY_SEARCH_PAGES)
return _repository_values(matches), search_has_more
async def test_repositories(self, repos: tuple[str, ...]) -> None:
for repo in repos:
@ -439,7 +439,7 @@ class GitHub:
)
try:
detail: Final = _PullDetail.model_validate(detail_response.json())
except Exception:
except ValueError:
raise SourceError("GitHub returned unexpected pull request details.") from None
login: Final = detail.user.login if detail.user and detail.user.login else "deleted-user"
@ -510,7 +510,7 @@ class GitHub:
return ""
profile: Final = _GitHubUserProfile.model_validate(response.json())
return normalize_email(profile.email)
except Exception:
except (httpx.HTTPError, ValueError):
return ""
async def _commit_metadata(
@ -585,7 +585,7 @@ class GitHub:
connection: Final = pull_request.commits
except SourceError:
raise
except Exception:
except ValueError:
raise SourceError("GitHub returned unexpected commit metadata.") from None
new_commits: Final[tuple[ROIPullCommit, ...]] = tuple(
_graphql_commit_evidence(node) for node in connection.nodes

View file

@ -53,10 +53,10 @@ class _DailySpendTable(Protocol):
async def group_by(
self,
*,
by: list[Literal["user_id", "date"]],
sum: dict[str, object],
where: dict[str, object],
order: dict[str, object],
by: Sequence[Literal["user_id", "date"]],
sum: Mapping[str, object],
where: Mapping[str, object],
order: Mapping[str, object],
) -> Sequence[Mapping[str, object]]: ...
@ -121,18 +121,18 @@ async def read_spend(
database: Final = prisma_client.db
daily_table: Final = database.litellm_dailyuserspend
group_by: Final[list[Literal["user_id", "date"]]] = [
"user_id",
"date",
]
sums: Final[dict[str, object]] = {"spend": True, "api_requests": True}
date_filter: Final[dict[str, object]] = {
"date": {
"gte": start.isoformat(),
"lte": end.isoformat(),
}
}
order: Final[dict[str, object]] = {"date": "asc"}
group_by: Final = TypeAdapter(list[Literal["user_id", "date"]]).validate_python(("user_id", "date"))
sums: Final = _JSON_OBJECT_ADAPTER.validate_python(MappingProxyType({"spend": True, "api_requests": True}))
date_filter: Final = _JSON_OBJECT_ADAPTER.validate_python(
MappingProxyType(
{
"date": _JSON_OBJECT_ADAPTER.validate_python(
MappingProxyType({"gte": start.isoformat(), "lte": end.isoformat()})
)
}
)
)
order: Final = _JSON_OBJECT_ADAPTER.validate_python(MappingProxyType({"date": "asc"}))
groups: Final = _DAILY_SPEND_GROUPS.validate_python(
await daily_table.group_by(
by=group_by,
@ -309,6 +309,20 @@ class SyncManager:
await self._coordinator.finish(self._owner, self.status)
return True
async def _heartbeat(
self, task: asyncio.Task[object] | None, coordinator: SyncCoordinator | None, owner: str
) -> None:
if coordinator is None or task is None:
return
try:
while True:
await asyncio.sleep(1)
if not await coordinator.heartbeat(owner, self.status):
task.cancel()
return
except Exception: # noqa: BLE001 - any coordination failure must stop a worker before its lease expires
task.cancel()
async def _run(
self,
settings: ROISettings,
@ -320,21 +334,7 @@ class SyncManager:
coordinator: SyncCoordinator | None,
owner: str,
) -> None:
task: Final = asyncio.current_task()
async def heartbeat() -> None:
if coordinator is None or task is None:
return
try:
while True:
await asyncio.sleep(1)
if not await coordinator.heartbeat(owner, self.status):
task.cancel()
return
except Exception:
task.cancel()
monitor: Final = asyncio.create_task(heartbeat())
monitor: Final = asyncio.create_task(self._heartbeat(asyncio.current_task(), coordinator, owner))
github: Final = self._github_factory(settings, github_transport)
try:
end: Final = self._clock().date()
@ -384,13 +384,17 @@ class SyncManager:
):
profile: Final = await github.profile_email(cached_pull["login"])
cached_record: Final = TypeAdapter(ROIPullRecord).validate_python(
{
**self._cached_record(cached_pull),
"profile_email": profile,
"emails": tuple(
sorted(frozenset(email for email in (*cached_pull["commit_emails"], profile) if email))
),
}
MappingProxyType(
{
**self._cached_record(cached_pull),
"profile_email": profile,
"emails": tuple(
sorted(
frozenset(email for email in (*cached_pull["commit_emails"], profile) if email)
)
),
}
)
)
self._update_estimate_progress(cached_record["estimate"])
return index, cached_record
@ -453,7 +457,7 @@ class SyncManager:
warnings=(),
)
await github.close()
report_json: Final[dict[str, object]] = _JSON_OBJECT_ADAPTER.validate_python(
report_json: Final[Mapping[str, object]] = _JSON_OBJECT_ADAPTER.validate_python(
_REPORT_ADAPTER.dump_python(report, mode="json")
)
monitor.cancel()
@ -482,7 +486,7 @@ class SyncManager:
raise
except SourceError as exc:
self._update_status(phase="error", stage="Sync failed", error=str(exc))
except Exception:
except Exception: # noqa: BLE001 - background job boundary records a safe failure for every source error
self._update_status(
phase="error",
stage="Sync failed",
@ -505,7 +509,10 @@ class SyncManager:
if coordinator is not None and self._status.phase != "complete":
await coordinator.finish(owner, self.status)
def _update_status(self, **update: Unpack[_StatusUpdate]) -> None:
def _update_status(
self,
**update: Unpack[_StatusUpdate], # kwargs-ok: Unpack preserves the typed status update contract
) -> None:
status: Final = ROISyncStatus.model_validate(MappingProxyType({**self._status.model_dump(), **update}))
self._status = status
@ -515,7 +522,7 @@ class SyncManager:
return None
try:
return _REPORT_ADAPTER.validate_python(parameter.param_value)
except Exception:
except ValueError:
return None
def _cached_record(self, pull: ROIPullRecord) -> ROIPullRecord:

View file

@ -1,6 +1,6 @@
import json
from datetime import datetime
from typing import Final, Protocol, cast
from types import MappingProxyType
from typing import Final, Protocol, cast # noqa: TID251 - PrismaWrapper dynamically delegates database methods
from pydantic import BaseModel, ConfigDict, TypeAdapter
@ -11,9 +11,15 @@ _SYNC_KEY: Final = "roi_calculator_sync"
_REPORT_KEY: Final = "roi_calculator_report"
class _SyncState(BaseModel):
owner: str
status: ROISyncStatus
cancel: bool = False
class _StateRow(BaseModel):
model_config = ConfigDict(extra="ignore")
param_value: dict[str, object]
param_value: _SyncState
expired: bool = False
last_run_at: datetime
@ -38,7 +44,7 @@ class SyncStore:
AND ($3::text::double precision = 0 OR "LiteLLM_Config".last_run_at <= NOW() - $3::text::double precision * INTERVAL '1 minute')
RETURNING param_name""",
_SYNC_KEY,
json.dumps({"owner": owner, "status": status.model_dump(), "cancel": False}),
_SyncState(owner=owner, status=status).model_dump_json(),
str(scheduled_interval),
)
return bool(rows)
@ -75,9 +81,12 @@ class SyncStore:
DELETE FROM "LiteLLM_Config" cached
WHERE starts_with(cached.param_name, 'roi_calculator_pull_')
AND EXISTS (SELECT 1 FROM owned) AND $4::text IS NOT NULL
AND NOT EXISTS (
AND EXISTS (
SELECT 1 FROM jsonb_array_elements($4::jsonb->'pulls') pull
WHERE cached.param_name = 'roi_calculator_pull_' || (pull->>'cache_key')
WHERE pull->>'url' = cached.param_value->>'url'
AND pull->'estimate'->>'status' = 'estimated'
AND pull->>'cache_key' IS NOT NULL
AND cached.param_name <> 'roi_calculator_pull_' || (pull->>'cache_key')
)
)
UPDATE "LiteLLM_Config" SET param_value = jsonb_set(param_value, '{status}', $3::jsonb),
@ -101,16 +110,18 @@ class SyncStore:
)
if not rows:
return None
status: Final = ROISyncStatus.model_validate(rows[0].param_value["status"])
status: Final = rows[0].param_value.status
if rows[0].expired and status.running:
return status.model_copy(
update={
"running": False,
"phase": "error",
"finished_at": rows[0].last_run_at.isoformat(),
"stage": "Sync interrupted",
"error": "The worker stopped responding. Run analysis again to resume saved estimates.",
}
update=MappingProxyType(
{
"running": False,
"phase": "error",
"finished_at": rows[0].last_run_at.isoformat(),
"stage": "Sync interrupted",
"error": "The worker stopped responding. Run analysis again to resume saved estimates.",
}
)
)
return status

View file

@ -0,0 +1,87 @@
from datetime import datetime, timezone
from types import MappingProxyType
from typing import Final
import pytest
from pydantic import TypeAdapter
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.types.roi_calculator import ROIPullRecord, ROIReport, ROISyncStatus
from tests.integration._support.database import read_rows, scratch_database, write_rows
@pytest.mark.asyncio
async def test_roi_cache_survives_scope_changes_and_uses_writer(monkeypatch: pytest.MonkeyPatch) -> None:
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())",
(),
database_url=writer_url,
)
monkeypatch.setenv("DATABASE_URL", writer_url)
# The reader deliberately has no table: any accidental replica read fails
monkeypatch.setenv("DATABASE_URL_READ_REPLICA", reader_url)
client: Final = PrismaClient(writer_url, ProxyLogging(UserApiKeyCache()))
await client.connect()
try:
store: Final = SyncStore(client)
report: Final = sample_report(datetime(2026, 9, 30, tzinfo=timezone.utc))
pull: Final[ROIPullRecord] = {
**report["pulls"][0],
"url": "https://github.com/example/repo/pull/1",
"cache_key": "new",
}
for key, url in (("old", pull["url"]), ("new", pull["url"]), ("outside-window", "other-pr")):
value: ROIPullRecord = {**pull, "url": url, "cache_key": key}
write_rows(
'INSERT INTO "LiteLLM_Config" (param_name, param_value) VALUES (%s, %s::jsonb)',
(f"roi_calculator_pull_{key}", TypeAdapter(ROIPullRecord).dump_json(value).decode()),
database_url=writer_url,
)
running: Final = ROISyncStatus(
running=True,
phase="estimates",
stage="Estimating",
done=0,
total=1,
estimated=0,
reused=0,
needs_attention=0,
error=None,
)
complete: Final = running.model_copy(update=MappingProxyType({"running": False, "phase": "complete"}))
narrowed: Final[ROIReport] = {**report, "pulls": (pull,)}
empty: Final[ROIReport] = {**report, "pulls": ()}
assert await store.acquire("worker", running)
assert not await store.acquire("other-worker", running)
observed: Final = await store.status()
assert observed is not None and observed.running
assert await store.heartbeat("worker", running)
assert await store.finish("worker", complete, narrowed)
assert tuple(
row["param_name"]
for row in read_rows(
'SELECT param_name FROM "LiteLLM_Config" WHERE starts_with(param_name, %s) ORDER BY param_name',
("roi_calculator_pull_",),
database_url=writer_url,
)
) == ("roi_calculator_pull_new", "roi_calculator_pull_outside-window")
assert not await store.acquire("scheduled", running, 1440)
assert await store.acquire("manual", running)
assert await store.finish("manual", complete, empty)
assert (
len(
read_rows(
'SELECT param_name FROM "LiteLLM_Config" WHERE starts_with(param_name, %s)',
("roi_calculator_pull_",),
database_url=writer_url,
)
)
== 2
)
finally:
await client.disconnect()

View file

@ -24,6 +24,7 @@ import {
import {
Activity,
BarChart3,
Calculator,
Bell,
Blocks,
Bot,
@ -218,7 +219,7 @@ const menuGroups: MenuGroup[] = [
{
key: "roi-calculator",
page: "roi-calculator",
icon: <BarChart3 {...ICON} />,
icon: <Calculator {...ICON} />,
roles: all_admin_roles,
label: (
<span className="flex items-center gap-2">

6
uv.lock generated
View file

@ -7857,14 +7857,14 @@ wheels = [
[[package]]
name = "pyjwt"
version = "2.14.0"
version = "2.15.0"
source = { registry = "https://pypi.org/simple" }
dependencies = [
{ name = "typing-extensions", marker = "python_full_version < '3.11'" },
]
sdist = { url = "https://files.pythonhosted.org/packages/af/c3/8a3b59c25070cc61dc517fbdfa5dc0904670c96f605cc69759dc09166b99/pyjwt-2.14.0.tar.gz", hash = "sha256:77283c83fb56ecf566a886c757a714bc83668e38156de2cce8263302f42e0b86", size = 113177, upload-time = "2026-09-11T13:11:54.638Z" }
sdist = { url = "https://files.pythonhosted.org/packages/02/a5/5197bfd06417837ac079921c66fa6393f1dea3557272a263cebfef69e432/pyjwt-2.15.0.tar.gz", hash = "sha256:b11c5f9791d7bf51c2b39a81ed669f6b2dbbd669df2942f6c60167e9e3d1abe4", size = 120513, upload-time = "2026-09-23T16:56:00.689Z" }
wheels = [
{ url = "https://files.pythonhosted.org/packages/9c/97/672cb32ce0dfea44b740cb7b4f97038463b9cf7c0ead1aacf595572851d6/pyjwt-2.14.0-py3-none-any.whl", hash = "sha256:ad0cef71c756a56e74863c2919cf0985f72decbcfcb550ee2f422e7c62b5eedc", size = 32896, upload-time = "2026-09-11T13:11:53.409Z" },
{ url = "https://files.pythonhosted.org/packages/e8/55/40e45bf052ee8ee12a4dfd785519660f8effa7b065442b91646ec6828619/pyjwt-2.15.0-py3-none-any.whl", hash = "sha256:7a3742debf6b879e912dbb9819ceec1594be812452b78c5f2e2dfc56564954f8", size = 33680, upload-time = "2026-09-23T16:55:59.241Z" },
]
[package.optional-dependencies]