diff --git a/litellm/proxy/_lazy_openapi_snapshot.json b/litellm/proxy/_lazy_openapi_snapshot.json index 5372aeed84f..e78bf6ebb3d 100644 --- a/litellm/proxy/_lazy_openapi_snapshot.json +++ b/litellm/proxy/_lazy_openapi_snapshot.json @@ -47779,6 +47779,10 @@ "title": "Github Api Url", "type": "string" }, + "has_estimator_key": { + "title": "Has Estimator Key", + "type": "boolean" + }, "has_github_token": { "title": "Has Github Token", "type": "boolean" @@ -47800,6 +47804,10 @@ }, "title": "Repos", "type": "array" + }, + "update_interval_minutes": { + "title": "Update Interval Minutes", + "type": "number" } }, "required": [ @@ -47808,6 +47816,8 @@ "estimator_model", "estimator_prompt", "backfill_days", + "update_interval_minutes", + "has_estimator_key", "identity_map", "has_github_token", "default_prompt", @@ -47833,6 +47843,17 @@ ], "title": "Backfill Days" }, + "estimator_key": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Estimator Key" + }, "estimator_model": { "anyOf": [ { @@ -47890,6 +47911,19 @@ } ], "title": "Repos" + }, + "update_interval_minutes": { + "anyOf": [ + { + "maximum": 43200.0, + "minimum": 0.0, + "type": "number" + }, + { + "type": "null" + } + ], + "title": "Update Interval Minutes" } }, "title": "ROISettingsUpdate", @@ -48007,6 +48041,11 @@ "title": "Done", "type": "integer" }, + "elapsed_seconds": { + "default": 0, + "title": "Elapsed Seconds", + "type": "integer" + }, "error": { "anyOf": [ { @@ -48022,10 +48061,32 @@ "title": "Estimated", "type": "integer" }, + "finished_at": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Finished At" + }, "needs_attention": { "title": "Needs Attention", "type": "integer" }, + "next_update": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Next Update" + }, "phase": { "enum": [ "idle", @@ -48039,6 +48100,17 @@ "title": "Phase", "type": "string" }, + "remaining_seconds": { + "anyOf": [ + { + "type": "integer" + }, + { + "type": "null" + } + ], + "title": "Remaining Seconds" + }, "reused": { "title": "Reused", "type": "integer" @@ -48051,6 +48123,17 @@ "title": "Stage", "type": "string" }, + "started_at": { + "anyOf": [ + { + "type": "string" + }, + { + "type": "null" + } + ], + "title": "Started At" + }, "total": { "title": "Total", "type": "integer" @@ -48141,6 +48224,32 @@ } }, "paths": { + "/roi-calculator/connections/test": { + "post": { + "operationId": "test_roi_calculator_connections_roi_calculator_connections_test_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ROISettingsResponse" + } + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Test Roi Calculator Connections", + "tags": [ + "roi_calculator" + ] + } + }, "/roi-calculator/identity-map": { "put": { "operationId": "update_roi_calculator_identity_map_roi_calculator_identity_map_put", @@ -48190,6 +48299,22 @@ "/roi-calculator/report": { "get": { "operationId": "get_roi_calculator_report_roi_calculator_report_get", + "parameters": [ + { + "in": "query", + "name": "mode", + "required": false, + "schema": { + "default": "live", + "enum": [ + "live", + "demo" + ], + "title": "Mode", + "type": "string" + } + } + ], "responses": { "200": { "content": { @@ -48200,6 +48325,16 @@ } }, "description": "Successful Response" + }, + "422": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/HTTPValidationError" + } + } + }, + "description": "Validation Error" } }, "security": [ @@ -48344,6 +48479,32 @@ ] } }, + "/roi-calculator/setup/reset": { + "post": { + "operationId": "reset_roi_calculator_setup_roi_calculator_setup_reset_post", + "responses": { + "200": { + "content": { + "application/json": { + "schema": { + "$ref": "#/components/schemas/ROISettingsResponse" + } + } + }, + "description": "Successful Response" + } + }, + "security": [ + { + "APIKeyHeader": [] + } + ], + "summary": "Reset Roi Calculator Setup", + "tags": [ + "roi_calculator" + ] + } + }, "/roi-calculator/sync": { "delete": { "operationId": "cancel_roi_calculator_sync_roi_calculator_sync_delete", diff --git a/litellm/proxy/management_endpoints/roi_calculator_endpoints.py b/litellm/proxy/management_endpoints/roi_calculator_endpoints.py index 5f4835f0cac..d35da2870c1 100644 --- a/litellm/proxy/management_endpoints/roi_calculator_endpoints.py +++ b/litellm/proxy/management_endpoints/roi_calculator_endpoints.py @@ -1,13 +1,14 @@ from collections.abc import Mapping, Sequence -from datetime import date +from datetime import date, datetime, timedelta, timezone from enum import Enum from types import MappingProxyType -from typing import Annotated, Final +from typing import Annotated, Final, Literal import httpx from fastapi import APIRouter, Depends, HTTPException, Query from pydantic import BaseModel, ConfigDict, Field, SecretStr, TypeAdapter, ValidationError +from litellm.llms.custom_httpx.http_handler import AsyncHTTPHandler from litellm.proxy._types import CommonProxyErrors, LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.auth.user_api_key_auth import user_api_key_auth from litellm.proxy.common_utils.encrypt_decrypt_utils import decrypt_value_helper, encrypt_value_helper @@ -15,8 +16,8 @@ from litellm.proxy.roi_calculator.analytics import normalize_email, summarize from litellm.proxy.roi_calculator.estimator import CompletionCaller, EstimatorModel from litellm.proxy.roi_calculator.github import GitHub, SourceError from litellm.proxy.roi_calculator.sync import SpendReader, SyncManager, read_spend, spend_prisma_client +from litellm.proxy.roi_calculator.sync_store import SyncStore from litellm.repositories.config_repository import ConfigRepository -from litellm.types.llms.openai import AllMessageValues from litellm.types.roi_calculator import ( DEFAULT_PROMPT, ROICompletionRequest, @@ -46,10 +47,12 @@ class _StoredSettings(BaseModel): github_api_url: str = "https://api.github.com" github_token: str = "" + estimator_key: str = "" repos: tuple[str, ...] = () estimator_model: str = "" estimator_prompt: str = DEFAULT_PROMPT backfill_days: int = Field(default=7, ge=1, le=3650) + update_interval_minutes: float = Field(default=1440, ge=0, le=43200) identity_map: Mapping[str, str] = Field(default_factory=lambda: MappingProxyType({})) @@ -116,7 +119,6 @@ def get_github_transport() -> httpx.AsyncBaseTransport | None: _ROUTER_ESTIMATOR_DEPLOYMENTS: Final = TypeAdapter(tuple[_RouterEstimatorDeployment, ...]) _MODEL_NAMES: Final = TypeAdapter(tuple[str, ...]) -_ROUTER_MESSAGES: Final = TypeAdapter(list[AllMessageValues]) def _estimator_models_from_deployments(deployments: Sequence[object]) -> tuple[EstimatorModel, ...]: @@ -174,6 +176,10 @@ async def _load_settings(repository: ConfigRepository) -> ROISettings: return ROISettings( github_api_url=stored.github_api_url, github_token=SecretStr(token or ""), + estimator_key=SecretStr(decrypt_value_helper(stored.estimator_key, _SETTINGS_KEY) or "") + if stored.estimator_key + else SecretStr(""), + update_interval_minutes=stored.update_interval_minutes, repos=stored.repos, estimator_model=stored.estimator_model, estimator_prompt=stored.estimator_prompt, @@ -188,10 +194,13 @@ async def _save_settings( repository: ConfigRepository, settings: ROISettings, encrypted_token: str, + encrypted_estimator_key: str, ) -> None: stored: Final = _StoredSettings( github_api_url=settings.github_api_url, github_token=encrypted_token, + estimator_key=encrypted_estimator_key, + update_interval_minutes=settings.update_interval_minutes, repos=settings.repos, estimator_model=settings.estimator_model, estimator_prompt=settings.estimator_prompt, @@ -221,50 +230,73 @@ def _public_settings(settings: ROISettings) -> ROISettingsResponse: backfill_days=settings.backfill_days, identity_map=settings.identity_map, has_github_token=bool(settings.github_token.get_secret_value()), + has_estimator_key=bool(settings.estimator_key.get_secret_value()), + update_interval_minutes=settings.update_interval_minutes, default_prompt=DEFAULT_PROMPT, available_models=models, ready=bool(settings.repos and settings.estimator_model and settings.estimator_model in models), ) -def _completion_caller() -> CompletionCaller: - from litellm.proxy.proxy_server import llm_router +def _gateway_key(settings: ROISettings) -> str: + from litellm.proxy.proxy_server import master_key - if llm_router is None: - raise HTTPException(status_code=503, detail="The proxy model router is not ready.") + credential: Final = settings.estimator_key.get_secret_value() or master_key + if not credential: + raise HTTPException(status_code=409, detail="Add an estimator API key in Advanced settings.") + return credential + + +def _completion_caller(settings: ROISettings) -> CompletionCaller: + from litellm.proxy.proxy_server import app + + credential: Final = _gateway_key(settings) + transport: Final = httpx.ASGITransport(app=app) async def complete(request: ROICompletionRequest) -> object: - messages: Final = _ROUTER_MESSAGES.validate_python(request.messages) - tags: Final[list[str]] = list( # mutable-ok: the router requires list-valued tags - request.metadata["tags"] - ) - metadata: Final[dict[str, object]] = { # mutable-ok: the router requires dict metadata - "tags": tags, - "litellm_roi_estimator": request.metadata["litellm_roi_estimator"], - } - if request.reasoning_effort is None: - return await llm_router.acompletion( - model=request.model, - messages=messages, - temperature=request.temperature, - response_format=request.response_format, - max_tokens=request.max_tokens, - metadata=metadata, + client: Final = AsyncHTTPHandler(transport=transport, timeout=180, follow_redirects=False) + try: + response: Final = await client.client.post( + "http://litellm.internal/v1/chat/completions", + headers=MappingProxyType({"authorization": f"Bearer {credential}", "content-type": "application/json"}), + content=request.model_dump_json(exclude_none=True), ) - response_with_reasoning: Final[object] = await llm_router.acompletion( - model=request.model, - messages=messages, - temperature=request.temperature, - response_format=request.response_format, - max_tokens=request.max_tokens, - metadata=metadata, - reasoning_effort=request.reasoning_effort, - ) - return response_with_reasoning + response.raise_for_status() + return TypeAdapter(object).validate_python(response.json()) + finally: + await client.close() return complete +class _GatewayModel(BaseModel): + id: str + + +class _GatewayModels(BaseModel): + data: tuple[_GatewayModel, ...] + + +async def _test_estimator_access(settings: ROISettings) -> None: + from litellm.proxy.proxy_server import app + + credential: Final = _gateway_key(settings) + client: Final = AsyncHTTPHandler(transport=httpx.ASGITransport(app=app), timeout=30, follow_redirects=False) + try: + response: Final = await client.client.get( + "http://litellm.internal/v1/models", + headers=MappingProxyType({"authorization": f"Bearer {credential}"}), + ) + response.raise_for_status() + models: Final = _GatewayModels.model_validate(response.json()) + if not any(model.id == settings.estimator_model for model in models.data): + raise HTTPException(status_code=409, detail="The estimator key cannot access the selected model.") + except (httpx.HTTPError, ValidationError): + raise HTTPException(status_code=409, detail="The estimator key could not connect to the gateway.") from None + finally: + await client.close() + + def _spend_reader(repository: ConfigRepository) -> SpendReader: async def get_spend(start: date, end: date) -> tuple[ROISpendRecord, ...]: prisma_client: Final = spend_prisma_client(repository.prisma_client) @@ -314,10 +346,22 @@ async def update_roi_calculator_settings( if github_url_changed else (current.github_token.get_secret_value(), stored.github_token) ) + estimator_key: Final = ( + patch.estimator_key or "" + if "estimator_key" in patch.model_fields_set + else current.estimator_key.get_secret_value() + ) + encrypted_estimator_key: Final = ( + TypeAdapter(str).validate_python(encrypt_value_helper(estimator_key)) if estimator_key else "" + ) try: settings: Final = ROISettings( github_api_url=github_api_url, github_token=SecretStr(plaintext_token), + estimator_key=SecretStr(estimator_key), + update_interval_minutes=patch.update_interval_minutes + if patch.update_interval_minutes is not None + else current.update_interval_minutes, repos=patch.repos if patch.repos is not None else current.repos, estimator_model=(patch.estimator_model if patch.estimator_model is not None else current.estimator_model), estimator_prompt=( @@ -328,7 +372,7 @@ async def update_roi_calculator_settings( ) except ValidationError as exc: raise HTTPException(status_code=422, detail=exc.errors(include_context=False)) from None - await _save_settings(repository, settings, encrypted_token) + await _save_settings(repository, settings, encrypted_token, encrypted_estimator_key) return _public_settings(settings) @@ -367,9 +411,14 @@ async def get_roi_calculator_repositories( ) async def get_roi_calculator_sync_status( _user: Annotated[UserAPIKeyAuth, Depends(_read_admin)], + repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], manager: Annotated[SyncManager, Depends(get_roi_sync_manager)], ) -> ROISyncStatus: - return manager.status + status: Final = await SyncStore(repository.prisma_client).status() or manager.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}) @router.post( @@ -388,13 +437,14 @@ async def start_roi_calculator_sync( public: Final = _public_settings(settings) if not public.ready: raise HTTPException(status_code=409, detail="Connect GitHub, select repositories, and choose a router model.") - if not manager.start( + if not await manager.start( settings, repository, _spend_reader(repository), - _completion_caller(), + _completion_caller(settings), transport, _router_estimator_models(settings.estimator_model), + SyncStore(repository.prisma_client), ): raise HTTPException(status_code=409, detail="A sync is already running.") return manager.status @@ -407,10 +457,13 @@ async def start_roi_calculator_sync( ) async def cancel_roi_calculator_sync( _user: Annotated[UserAPIKeyAuth, Depends(_write_admin)], + repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], manager: Annotated[SyncManager, Depends(get_roi_sync_manager)], ) -> ROISyncStatus: + store: Final = SyncStore(repository.prisma_client) + await store.cancel() await manager.cancel() - return manager.status + return await store.status() or manager.status @router.get( @@ -421,7 +474,13 @@ async def cancel_roi_calculator_sync( async def get_roi_calculator_report( _user: Annotated[UserAPIKeyAuth, Depends(_read_admin)], repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], + mode: Literal["live", "demo"] = "live", ) -> ROIReportResponse: + if mode == "demo": + from litellm.proxy.roi_calculator.sample import sample_report + + sample: Final = summarize(sample_report(datetime.now(timezone.utc)), MappingProxyType({})) + return ROIReportResponse(report=ROISummaryResponse.model_validate(sample)) report: Final = await _load_report(repository) if report is None: return ROIReportResponse(report=None) @@ -454,16 +513,111 @@ async def update_roi_calculator_identity_map( settings: Final = ROISettings( github_api_url=current.github_api_url, github_token=current.github_token, + estimator_key=current.estimator_key, + update_interval_minutes=current.update_interval_minutes, repos=current.repos, estimator_model=current.estimator_model, estimator_prompt=current.estimator_prompt, backfill_days=current.backfill_days, identity_map=identity_map, ) - await _save_settings(repository, settings, current_stored.github_token) + await _save_settings(repository, settings, current_stored.github_token, current_stored.estimator_key) report: Final = await _load_report(repository) summary: Final = summarize(report, settings.identity_map) if report is not None else None return ROIIdentityMapResponse( report=ROISummaryResponse.model_validate(summary) if summary is not None else None, identity_map=settings.identity_map, ) + + +def _next_update(settings: ROISettings, status: ROISyncStatus, report: ROIReport | None) -> datetime | None: + if ( + not report + or not settings.repos + or not settings.estimator_model + or not settings.update_interval_minutes + or status.running + ): + return None + anchor: Final = status.finished_at or status.started_at or report["synced_at"] + return datetime.fromisoformat(anchor.replace("Z", "+00:00")) + timedelta(minutes=settings.update_interval_minutes) + + +async def run_scheduled_sync() -> None: + from litellm.proxy.proxy_server import prisma_client + + if prisma_client is None: + return + repository: Final = ConfigRepository(prisma_client) + settings: Final = await _load_settings(repository) + if not settings.update_interval_minutes or not _public_settings(settings).ready: + return + store: Final = SyncStore(prisma_client) + status: Final = await store.status() or _SYNC_MANAGER.status + report: Final = await _load_report(repository) + next_update: Final = _next_update(settings, status, report) + if next_update is None or next_update > datetime.now(timezone.utc): + return + await _SYNC_MANAGER.start( + settings, + repository, + _spend_reader(repository), + _completion_caller(settings), + estimator_models=_router_estimator_models(settings.estimator_model), + coordinator=store, + scheduled_interval=settings.update_interval_minutes, + ) + + +@router.post("/roi-calculator/connections/test", tags=_ROI_TAGS) +async def test_roi_calculator_connections( + _user: Annotated[UserAPIKeyAuth, Depends(_write_admin)], + repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], + transport: Annotated[httpx.AsyncBaseTransport | None, Depends(get_github_transport)], +) -> ROISettingsResponse: + settings: Final = await _load_settings(repository) + public: Final = _public_settings(settings) + if not public.ready: + raise HTTPException(status_code=409, detail="Choose repositories and an available estimator model first.") + await _test_estimator_access(settings) + github: Final = GitHub(settings, transport) + try: + await github.test_repositories(settings.repos) + except SourceError as exc: + raise HTTPException(status_code=502, detail=str(exc)) from None + finally: + await github.close() + return public + + +@router.post("/roi-calculator/setup/reset", tags=_ROI_TAGS) +async def reset_roi_calculator_setup( + _user: Annotated[UserAPIKeyAuth, Depends(_write_admin)], + repository: Annotated[ConfigRepository, Depends(get_roi_config_repository)], +) -> ROISettingsResponse: + from uuid import uuid4 + + store: Final = SyncStore(repository.prisma_client) + owner: Final = str(uuid4()) + status: Final = ROISyncStatus( + running=True, + phase="spend", + stage="Restarting setup", + done=0, + total=0, + estimated=0, + reused=0, + needs_attention=0, + error=None, + ) + if not await store.acquire(owner, status): + raise HTTPException(status_code=409, detail="Cancel the running analysis before restarting setup.") + try: + current: Final = await _load_settings(repository) + stored: Final = await _load_stored_settings(repository) + settings: Final = current.model_copy(update={"repos": ()}) + await _save_settings(repository, settings, stored.github_token, stored.estimator_key) + await repository.prisma_client.db.execute_raw('DELETE FROM "LiteLLM_Config" WHERE param_name = $1', _REPORT_KEY) + return _public_settings(settings) + finally: + await store.finish(owner, status.model_copy(update={"running": False, "phase": "idle", "stage": "Idle"})) diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py index 7cd116c23d7..a3269914ab9 100644 --- a/litellm/proxy/proxy_server.py +++ b/litellm/proxy/proxy_server.py @@ -1539,6 +1539,11 @@ async def proxy_startup_event(app: FastAPI) -> AsyncGenerator[None, None]: if not model_info_scheduler.running: model_info_scheduler.start() + if scheduler is not None and prisma_client is not None: + from litellm.proxy.management_endpoints.roi_calculator_endpoints import run_scheduled_sync + + scheduler.add_job(run_scheduled_sync, "interval", seconds=30, id="roi_calculator_refresh", max_instances=1) + # End of startup event yield diff --git a/litellm/proxy/roi_calculator/github.py b/litellm/proxy/roi_calculator/github.py index f4a38450eb7..f2ddcef3ecf 100644 --- a/litellm/proxy/roi_calculator/github.py +++ b/litellm/proxy/roi_calculator/github.py @@ -312,6 +312,7 @@ class GitHub: ) -> None: if client is not None and transport is not None: raise ValueError("Pass either an injected GitHub client or a transport.") + self._profiles: Mapping[str, str] = MappingProxyType({}) token: Final = settings.github_token.get_secret_value() self._headers: Final[Mapping[str, str]] = ( MappingProxyType( @@ -397,6 +398,17 @@ class GitHub: matches, has_more = await search_pages(first_github_page, _REPOSITORY_SEARCH_PAGES) return _repository_values(matches), has_more + async def test_repositories(self, repos: tuple[str, ...]) -> None: + for repo in repos: + await _request(self.client, "GET", self._url(f"repos/{repo}"), headers=self._headers) + await _request( + self.client, + "GET", + self._url(f"repos/{repo}/pulls"), + params=MappingProxyType({"per_page": 1, "state": "closed"}), + headers=self._headers, + ) + async def pulls(self, repo: str, start: date, end: date) -> tuple[GitHubPullListItem, ...]: async def pull_pages() -> AsyncIterator[GitHubPullListItem]: async for page in _pages( @@ -443,13 +455,18 @@ class GitHub: yield item files: Final = tuple(item.evidence() for item in await _collect(file_pages())) - profile_email: Final = await self._profile_email(login) + profile_email: Final = await self.profile_email(login) commits, authors, commit_count = await self._commit_metadata(repo, pull.number, detail) + commit_emails: Final = tuple( + sorted( + frozenset(normalize_email(author[1]) for author in authors if author[0].casefold() == login.casefold()) + ) + ) email_candidates: Final = frozenset( address for address in ( profile_email, - *(normalize_email(author[1]) for author in authors if author[0].casefold() == login.casefold()), + *commit_emails, ) if address ) @@ -463,6 +480,7 @@ class GitHub: "login": login, "emails": tuple(sorted(email_candidates)), "profile_email": profile_email, + "commit_emails": commit_emails, "merged_at": detail.merged_at, "head_sha": detail.head.sha, "additions": detail.additions, @@ -475,7 +493,14 @@ class GitHub: } return evidence - async def _profile_email(self, login: str) -> str: + async def profile_email(self, login: str) -> str: + if login.casefold() in self._profiles: + return self._profiles[login.casefold()] + address: Final = await self._load_profile_email(login) + self._profiles = MappingProxyType({**self._profiles, login.casefold(): address}) + return address + + async def _load_profile_email(self, login: str) -> str: try: response: Final = await self.client.get( self._url(f"users/{quote(login, safe='')}"), diff --git a/litellm/proxy/roi_calculator/sample.py b/litellm/proxy/roi_calculator/sample.py new file mode 100644 index 00000000000..fe5fbbaa866 --- /dev/null +++ b/litellm/proxy/roi_calculator/sample.py @@ -0,0 +1,64 @@ +from datetime import datetime, timedelta +from typing import Final + +from litellm.types.roi_calculator import DEFAULT_PROMPT, ROIEstimate, ROIPullRecord, ROIReport, ROISpendRecord + + +def sample_report(now: datetime) -> ROIReport: + start: Final = now.date() - timedelta(days=29) + examples: Final = ( + ("alex", "alex@example.com", "Add usage breakdown by model", 6.5, 18.2), + ("jordan", "jordan@example.com", "Fix streaming response cancellation", 4.0, 12.8), + ("casey", "", "Add integration tests for billing", 5.5, 0.0), + ) + + def pull(index: int, login: str, email: str, title: str, hours: float) -> ROIPullRecord: + estimate: Final[ROIEstimate] = { + "status": "estimated", + "hours": hours, + "reasoning": "Sample estimate of engineering effort without AI assistance. Live estimates use PR descriptions, file change counts, and commit metadata.", + "model": "your-estimator-model", + "effort_basis": "without_ai", + "evidence_source": "pr_metadata", + "cached": False, + } + return ROIPullRecord( + repo="example/gateway", + number=142 + index, + title=title, + url="", + login=login, + emails=(email,) if email else (), + profile_email=email, + merged_at=(start + timedelta(days=2 + index * 2)).isoformat() + "T14:20:00Z", + head_sha=f"sample-{index}", + additions=47 + index * 23, + deletions=12 + index * 4, + changed_files=3, + commit_count=1, + incomplete_metadata=False, + estimate=estimate, + cache_key=None, + ) + + pulls: Final = tuple( + pull(index, login, email, title, hours) for index, (login, email, title, hours, _) in enumerate(examples) + ) + spend: Final = tuple( + ROISpendRecord(date=pulls[index]["merged_at"][:10], user_id=login, email=email, spend=cost, requests=150) + for index, (login, email, _, _, cost) in enumerate(examples) + if email + ) + return ROIReport( + mode="demo", + start=start.isoformat(), + end=now.date().isoformat(), + synced_at=now.isoformat(), + repos=("example/gateway",), + estimator_model="your-estimator-model", + estimator_prompt=DEFAULT_PROMPT, + effort_basis="without_ai", + spend=spend, + pulls=pulls, + settings_fingerprint="sample", + ) diff --git a/litellm/proxy/roi_calculator/sync.py b/litellm/proxy/roi_calculator/sync.py index 2a6fe249502..61c9d1523cb 100644 --- a/litellm/proxy/roi_calculator/sync.py +++ b/litellm/proxy/roi_calculator/sync.py @@ -5,6 +5,7 @@ from datetime import date, datetime, timedelta, timezone from itertools import chain from types import MappingProxyType from typing import Final, Literal, Protocol, runtime_checkable +from uuid import uuid4 import httpx from pydantic import BaseModel, ConfigDict, Field, TypeAdapter @@ -41,6 +42,12 @@ class _ReportRepository(Protocol): async def set_param(self, param_name: str, param_value: object) -> object: ... +class SyncCoordinator(Protocol): + async def acquire(self, owner: str, status: ROISyncStatus, scheduled_interval: float = 0) -> bool: ... + async def heartbeat(self, owner: str, status: ROISyncStatus) -> bool: ... + async def finish(self, owner: str, status: ROISyncStatus, report: ROIReport | None = None) -> bool: ... + + class _DailySpendTable(Protocol): async def group_by( self, @@ -222,12 +229,26 @@ class SyncManager: error=None, ) self._task: asyncio.Task[None] | None = None + self._coordinator: SyncCoordinator | None = None + self._owner: str = "" @property def status(self) -> ROISyncStatus: - return self._status + if self._status.started_at is None: + return self._status + start: Final = datetime.fromisoformat(self._status.started_at) + finish: Final = datetime.fromisoformat(self._status.finished_at) if self._status.finished_at else self._clock() + elapsed: Final = max(0, int((finish - start).total_seconds())) + remaining: Final = ( + max(0, round(elapsed / self._status.done * (self._status.total - self._status.done))) + if self._status.running and self._status.done >= PR_CONCURRENCY + else None + ) + return self._status.model_copy( + update=MappingProxyType({"elapsed_seconds": elapsed, "remaining_seconds": remaining}) + ) - def start( + async def start( self, settings: ROISettings, repository: _ReportRepository, @@ -235,11 +256,14 @@ class SyncManager: complete: CompletionCaller, github_transport: httpx.AsyncBaseTransport | None = None, estimator_models: tuple[EstimatorModel, ...] | None = None, + coordinator: SyncCoordinator | None = None, + scheduled_interval: float = 0, ) -> bool: if self._status.running or not settings.repos or not settings.estimator_model: return False - self._status = ROISyncStatus( + initial_status: Final = ROISyncStatus( running=True, + started_at=self._clock().isoformat(), phase="spend", stage="Reading gateway spend", done=0, @@ -249,8 +273,16 @@ class SyncManager: needs_attention=0, error=None, ) + owner: Final = str(uuid4()) + if coordinator is not None and not await coordinator.acquire(owner, initial_status, scheduled_interval): + return False + self._status = initial_status + self._coordinator = coordinator + self._owner = owner self._task = asyncio.create_task( - self._run(settings, repository, spend_reader, complete, github_transport, estimator_models) + self._run( + settings, repository, spend_reader, complete, github_transport, estimator_models, coordinator, owner + ) ) return True @@ -262,6 +294,9 @@ class SyncManager: with suppress(asyncio.CancelledError): await task self._update_status(running=False, phase="cancelled", stage="Sync cancelled") + self._status = self.status.model_copy(update=MappingProxyType({"finished_at": self._clock().isoformat()})) + if self._coordinator is not None: + await self._coordinator.finish(self._owner, self.status) return True async def _run( @@ -272,7 +307,24 @@ class SyncManager: complete: CompletionCaller, github_transport: httpx.AsyncBaseTransport | None, estimator_models: tuple[EstimatorModel, ...] | None, + 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()) github: Final = self._github_factory(settings, github_transport) try: end: Final = self._clock().date() @@ -298,54 +350,84 @@ class SyncManager: (index, repo, pull, cache_key(settings, context, repo, pull)) for index, (repo, pull) in enumerate(queue) ) - cached: Final = tuple( - (index, previous_pulls[key]) - for index, _, _, key in indexed_queue - if key is not None - and (previous_pull := previous_pulls.get(key)) is not None - and previous_pull["estimate"]["status"] == "estimated" - ) - cached_by_index: Final[Mapping[int, ROIPullRecord]] = MappingProxyType( - {index: self._cached_record(pull) for index, pull in cached} - ) - pending: Final = tuple(item for item in indexed_queue if item[0] not in cached_by_index) - reused_count: Final = len(cached) self._update_status( phase="estimates", stage="Estimating new or changed pull requests", - done=reused_count, total=len(queue), - estimated=reused_count, - reused=reused_count, ) - semaphore: Final = asyncio.Semaphore(PR_CONCURRENCY) estimator: Final = Estimator(settings, complete, estimator_models) async def process( item: tuple[int, str, GitHubPullListItem, str | None], ) -> tuple[int, ROIPullRecord]: - async with semaphore: - index, repo, pull, key = item - evidence: Final = await github.evidence(repo, pull) - estimate: Final = await _estimate_with_fallback(estimator, evidence) - record: Final = self._report_record(evidence, estimate, key) - self._update_estimate_progress(estimate) - return index, record + index, repo, pull, key = item + saved: Final = await repository.get_param("roi_calculator_pull_" + key) if key is not None else None + cached_pull: Final = ( + TypeAdapter(ROIPullRecord).validate_python(saved.param_value) + if saved is not None + else previous_pulls.get(key or "") + ) + if ( + cached_pull is not None + and cached_pull["estimate"]["status"] == "estimated" + and "commit_emails" in cached_pull + ): + 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)) + ), + } + ) + self._update_estimate_progress(cached_record["estimate"]) + return index, cached_record + evidence: Final = await github.evidence(repo, pull) + estimate: Final = await _estimate_with_fallback(estimator, evidence) + evidence_item: Final = GitHubPullListItem.model_validate( + MappingProxyType( + { + "number": evidence["number"], + "title": evidence["title"], + "body": evidence["body"], + "head": MappingProxyType({"sha": evidence["head_sha"]}), + "user": MappingProxyType({"login": evidence["login"]}), + "merged_at": evidence["merged_at"], + "updated_at": evidence["merged_at"], + } + ) + ) + fetched_key: Final = cache_key(settings, context, repo, evidence_item) + record: Final = self._report_record(evidence, estimate, fetched_key) + if fetched_key is not None and estimate["status"] == "estimated": + await repository.set_param( + "roi_calculator_pull_" + fetched_key, + _JSON_OBJECT_ADAPTER.validate_python( + TypeAdapter(ROIPullRecord).dump_python(record, mode="json") + ), + ) + self._update_estimate_progress(estimate) + return index, record - workers: Final = tuple(asyncio.create_task(process(item)) for item in pending) + async def worker(offset: int) -> tuple[tuple[int, ROIPullRecord], ...]: + return tuple( + [await process(indexed_queue[index]) for index in range(offset, len(indexed_queue), PR_CONCURRENCY)] + ) + + workers: Final = tuple(asyncio.create_task(worker(offset)) for offset in range(PR_CONCURRENCY)) try: - processed: Final = await asyncio.gather(*workers) + groups: Final = await asyncio.gather(*workers) + processed: Final = tuple(chain.from_iterable(groups)) finally: - for worker in workers: - if not worker.done(): - worker.cancel() + for worker_task in workers: + if not worker_task.done(): + worker_task.cancel() await asyncio.gather(*workers, return_exceptions=True) processed_by_index: Final[Mapping[int, ROIPullRecord]] = MappingProxyType( {index: pull for index, pull in processed} ) - report_pulls: Final[Mapping[int, ROIPullRecord]] = MappingProxyType( - {**cached_by_index, **processed_by_index} - ) report: Final = ROIReport( mode="live", start=start.isoformat(), @@ -356,7 +438,7 @@ class SyncManager: estimator_prompt=settings.estimator_prompt, effort_basis="without_ai", spend=spend, - pulls=tuple(report_pulls[index] for index in range(len(queue))), + pulls=tuple(processed_by_index[index] for index in range(len(queue))), settings_fingerprint=settings_fingerprint(settings), warnings=(), ) @@ -364,8 +446,27 @@ class SyncManager: report_json: Final[dict[str, object]] = _JSON_OBJECT_ADAPTER.validate_python( _REPORT_ADAPTER.dump_python(report, mode="json") ) - await repository.set_param("roi_calculator_report", report_json) - self._update_status(phase="complete", stage="Up to date") + monitor.cancel() + with suppress(asyncio.CancelledError): + await monitor + completed_status: Final = self.status.model_copy( + update=MappingProxyType( + { + "running": False, + "phase": "complete", + "stage": "Analysis complete", + "finished_at": self._clock().isoformat(), + } + ) + ) + if coordinator is not None: + if not await coordinator.finish(owner, completed_status, report): + raise SourceError( + "This sync was cancelled or replaced. Run analysis again to resume saved estimates." + ) + else: + await repository.set_param("roi_calculator_report", report_json) + self._status = completed_status except asyncio.CancelledError: self._update_status(phase="cancelled", stage="Sync cancelled") raise @@ -381,11 +482,18 @@ class SyncManager: ), ) finally: + monitor.cancel() + with suppress(asyncio.CancelledError): + await monitor try: if self._status.phase != "complete": await github.close() finally: - self._update_status(running=False) + self._status = self._status.model_copy( + update=MappingProxyType({"running": False, "finished_at": self._clock().isoformat()}) + ) + if coordinator is not None and self._status.phase != "complete": + await coordinator.finish(owner, self.status) def _update_status(self, **update: Unpack[_StatusUpdate]) -> None: status: Final = ROISyncStatus.model_validate(MappingProxyType({**self._status.model_dump(), **update})) @@ -410,6 +518,7 @@ class SyncManager: login=pull["login"], emails=pull["emails"], profile_email=pull["profile_email"], + commit_emails=pull.get("commit_emails", ()), merged_at=pull["merged_at"], head_sha=pull["head_sha"], additions=pull["additions"], @@ -435,6 +544,7 @@ class SyncManager: login=evidence["login"], emails=evidence["emails"], profile_email=evidence["profile_email"], + commit_emails=evidence.get("commit_emails", ()), merged_at=evidence["merged_at"], head_sha=evidence["head_sha"], additions=evidence["additions"], diff --git a/litellm/proxy/roi_calculator/sync_store.py b/litellm/proxy/roi_calculator/sync_store.py new file mode 100644 index 00000000000..760991d3219 --- /dev/null +++ b/litellm/proxy/roi_calculator/sync_store.py @@ -0,0 +1,121 @@ +import json +from typing import Final, Protocol, runtime_checkable + +from pydantic import BaseModel, ConfigDict, TypeAdapter + +from litellm.proxy.utils import PrismaClient +from litellm.types.roi_calculator import ROIReport, ROISyncStatus + +_SYNC_KEY: Final = "roi_calculator_sync" +_REPORT_KEY: Final = "roi_calculator_report" + + +class _StateRow(BaseModel): + model_config = ConfigDict(extra="ignore") + param_value: dict[str, object] + expired: bool = False + + +@runtime_checkable +class _SyncDatabase(Protocol): + async def query_raw(self, query: str, *args: object) -> object: ... + async def execute_raw(self, query: str, *args: object) -> int: ... + + +def _sync_database(database: object) -> _SyncDatabase: + if not isinstance(database, _SyncDatabase): + raise TypeError("The database does not support sync coordination.") + return database + + +class SyncStore: + def __init__(self, prisma: PrismaClient) -> None: + self._db: Final = _sync_database(prisma.db) + + async def acquire(self, owner: str, status: ROISyncStatus, scheduled_interval: float = 0) -> bool: + rows: Final = await self._db.query_raw( + """INSERT INTO "LiteLLM_Config" (param_name, param_value, last_run_at) + VALUES ($1, $2::jsonb, NOW()) + ON CONFLICT (param_name) DO UPDATE + SET param_value = EXCLUDED.param_value, last_run_at = NOW() + WHERE ("LiteLLM_Config".last_run_at < NOW() - INTERVAL '60 seconds' + OR "LiteLLM_Config".param_value->'status'->>'running' = 'false') + AND ($3::float = 0 OR "LiteLLM_Config".last_run_at <= NOW() - $3::float * INTERVAL '1 minute') + RETURNING param_name""", + _SYNC_KEY, + json.dumps({"owner": owner, "status": status.model_dump(), "cancel": False}), + scheduled_interval, + ) + return bool(rows) + + async def heartbeat(self, owner: str, status: ROISyncStatus) -> bool: + rows: Final = await self._db.query_raw( + """UPDATE "LiteLLM_Config" + SET param_value = jsonb_set(param_value, '{status}', $3::jsonb), last_run_at = NOW() + WHERE param_name = $1 AND param_value->>'owner' = $2 + AND param_value->>'cancel' = 'false' + AND param_value->'status'->>'running' = 'true' + AND last_run_at >= NOW() - INTERVAL '60 seconds' + RETURNING param_name""", + _SYNC_KEY, + owner, + status.model_dump_json(), + ) + return bool(rows) + + async def finish(self, owner: str, status: ROISyncStatus, report: ROIReport | None = None) -> bool: + report_json: Final = TypeAdapter(ROIReport).dump_json(report).decode() if report is not None else None + rows: Final = await self._db.query_raw( + """WITH owned AS ( + SELECT param_name FROM "LiteLLM_Config" + WHERE param_name = $1 AND param_value->>'owner' = $2 + AND last_run_at >= NOW() - INTERVAL '60 seconds' + AND ($4::text IS NULL OR param_value->>'cancel' = 'false') + FOR UPDATE + ), report_write AS ( + INSERT INTO "LiteLLM_Config" (param_name, param_value) + SELECT $5, $4::jsonb FROM owned WHERE $4::text IS NOT NULL + ON CONFLICT (param_name) DO UPDATE SET param_value = EXCLUDED.param_value + ) + UPDATE "LiteLLM_Config" SET param_value = jsonb_set(param_value, '{status}', $3::jsonb), + last_run_at = NOW() + WHERE param_name IN (SELECT param_name FROM owned) RETURNING param_name""", + _SYNC_KEY, + owner, + status.model_dump_json(), + report_json, + _REPORT_KEY, + ) + return bool(rows) + + async def status(self) -> ROISyncStatus | None: + rows: Final = TypeAdapter(tuple[_StateRow, ...]).validate_python( + await self._db.query_raw( + """SELECT param_value, last_run_at < NOW() - INTERVAL '60 seconds' AS expired + FROM "LiteLLM_Config" WHERE param_name = $1""", + _SYNC_KEY, + ) + ) + if not rows: + return None + status: Final = ROISyncStatus.model_validate(rows[0].param_value["status"]) + if rows[0].expired and status.running: + return status.model_copy( + update={ + "running": False, + "phase": "error", + "stage": "Sync interrupted", + "error": "The worker stopped responding. Run analysis again to resume saved estimates.", + } + ) + return status + + async def cancel(self) -> None: + await self._db.execute_raw( + """UPDATE "LiteLLM_Config" SET param_value = jsonb_set(param_value, '{cancel}', 'true'::jsonb) + WHERE param_name = $1 AND param_value->'status'->>'running' = 'true' """, + _SYNC_KEY, + ) + + async def clear_report(self) -> None: + await self._db.execute_raw('DELETE FROM "LiteLLM_Config" WHERE param_name = $1', _REPORT_KEY) diff --git a/litellm/types/roi_calculator.py b/litellm/types/roi_calculator.py index b57a169c0c8..7e3b31732df 100644 --- a/litellm/types/roi_calculator.py +++ b/litellm/types/roi_calculator.py @@ -16,12 +16,21 @@ class ROISettings(BaseModel): github_api_url: str = "https://api.github.com" github_token: SecretStr = SecretStr("") + estimator_key: SecretStr = SecretStr("") repos: tuple[str, ...] = () estimator_model: str = "" estimator_prompt: str = DEFAULT_PROMPT backfill_days: int = Field(default=7, ge=1, le=3650) + update_interval_minutes: float = Field(default=1440, ge=0, le=43200, allow_inf_nan=False) identity_map: Mapping[str, str] = Field(default_factory=lambda: MappingProxyType({})) + @field_validator("update_interval_minutes") + @classmethod + def validate_update_interval(cls, value: float) -> float: + if 0 < value < 5: + raise ValueError("Choose manual updates (0), or an interval of at least 5 minutes.") + return value + @field_validator("github_api_url") @classmethod def normalize_github_api_url(cls, value: str) -> str: @@ -93,10 +102,12 @@ class ROISettingsUpdate(BaseModel): github_api_url: str | None = None github_token: str | None = None + estimator_key: str | None = None repos: tuple[str, ...] | None = None estimator_model: str | None = None estimator_prompt: str | None = None backfill_days: int | None = Field(default=None, ge=1, le=3650) + update_interval_minutes: float | None = Field(default=None, ge=0, le=43200, allow_inf_nan=False) class ROISettingsResponse(BaseModel): @@ -105,6 +116,8 @@ class ROISettingsResponse(BaseModel): estimator_model: str estimator_prompt: str backfill_days: int + update_interval_minutes: float + has_estimator_key: bool identity_map: Mapping[str, str] has_github_token: bool default_prompt: str @@ -134,6 +147,11 @@ class ROISyncStatus(BaseModel): reused: int needs_attention: int error: str | None + started_at: str | None = None + finished_at: str | None = None + next_update: str | None = None + elapsed_seconds: int = 0 + remaining_seconds: int | None = None class ROISpendRecord(TypedDict): @@ -162,6 +180,7 @@ class ROIPullRecord(TypedDict): login: ReadOnly[str] emails: ReadOnly[tuple[str, ...]] profile_email: ReadOnly[str] + commit_emails: NotRequired[ReadOnly[tuple[str, ...]]] merged_at: ReadOnly[str] head_sha: ReadOnly[str] additions: ReadOnly[int] @@ -213,6 +232,7 @@ class ROIPullEvidence(TypedDict): login: ReadOnly[str] emails: ReadOnly[tuple[str, ...]] profile_email: ReadOnly[str] + commit_emails: NotRequired[ReadOnly[tuple[str, ...]]] merged_at: ReadOnly[str] head_sha: ReadOnly[str] additions: ReadOnly[int] diff --git a/tests/unit/interactions/test_openapi_compliance.py b/tests/unit/interactions/test_openapi_compliance.py index d3f1183cea6..4378c4d9739 100644 --- a/tests/unit/interactions/test_openapi_compliance.py +++ b/tests/unit/interactions/test_openapi_compliance.py @@ -9,7 +9,7 @@ Run with: pytest tests/unit/interactions/test_openapi_compliance.py -v import json import os -from typing import Any, Dict +from typing import Any, Dict, Final from unittest.mock import MagicMock, patch import httpx @@ -32,8 +32,7 @@ def _load_openapi_spec_dict() -> Dict[str, Any]: return response.json() except Exception as e: # pragma: no cover - defensive, env-dependent pytest.skip( - f"Skipping Google Interactions OpenAPI compliance tests - " - f"unable to load spec from {OPENAPI_SPEC_URL}: {e}" + f"Skipping Google Interactions OpenAPI compliance tests - unable to load spec from {OPENAPI_SPEC_URL}: {e}" ) @@ -44,6 +43,20 @@ def _declared_type_value(variant_schema: Dict[str, Any]) -> Any: return type_property.get("const") or (enum_values[0] if len(enum_values) == 1 else None) +def _model_request_schema(spec: Dict[str, Any]) -> Dict[str, Any]: + create: Final = next( + methods["post"] + for path, methods in spec["paths"].items() + if path.endswith("/interactions") and "post" in methods + ) + schema: Final = create["requestBody"]["content"]["application/json"]["schema"] + variants: Final = schema.get("oneOf", [schema]) + resolved: Final = tuple( + spec["components"]["schemas"][item["$ref"].split("/")[-1]] if "$ref" in item else item for item in variants + ) + return next(item for item in resolved if "model" in item.get("properties", {})) + + @pytest.fixture(scope="module") def spec_dict() -> Dict[str, Any]: """Load raw spec dict for manual validation.""" @@ -61,11 +74,10 @@ class TestRequestCompliance: def test_create_model_interaction_request_schema(self, spec_dict): """Verify CreateModelInteractionParams schema fields.""" - schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"] + schema = _model_request_schema(spec_dict) - # Required fields per spec - assert "model" in schema["required"] - assert "input" in schema["required"] + assert "model" in schema["properties"] + assert "input" in schema["properties"] # Check our supported optional fields exist in spec our_optional_fields = [ @@ -88,7 +100,7 @@ class TestRequestCompliance: def test_input_types_match_spec(self, spec_dict): """Verify input field supports string, Content, Content[], Turn[].""" - schema = spec_dict["components"]["schemas"]["CreateModelInteractionParams"] + schema = _model_request_schema(spec_dict) input_schema = schema["properties"]["input"] # The input property may be inline oneOf or a $ref to InteractionsInput @@ -125,22 +137,18 @@ class TestRequestCompliance: discriminator = content_schema.get("discriminator") if discriminator is not None: - assert ( - discriminator.get("propertyName") == "type" - ), f"Content is discriminated on {discriminator.get('propertyName')!r}, not 'type'" + assert discriminator.get("propertyName") == "type", ( + f"Content is discriminated on {discriminator.get('propertyName')!r}, not 'type'" + ) variant_names = [ - option["$ref"].split("/")[-1] - for option in content_schema.get("oneOf", []) - if "$ref" in option + option["$ref"].split("/")[-1] for option in content_schema.get("oneOf", []) if "$ref" in option ] assert variant_names, f"Content is not a union of named variants: {content_schema}" mapping = (discriminator or {}).get("mapping") or {} type_values = { - variant: mapping_value - for mapping_value, ref in mapping.items() - for variant in [ref.split("/")[-1]] + variant: mapping_value for mapping_value, ref in mapping.items() for variant in [ref.split("/")[-1]] } or { variant: _declared_type_value(spec_dict["components"]["schemas"].get(variant, {})) for variant in variant_names @@ -191,7 +199,9 @@ class TestRequestCompliance: for option in spec_dict["components"]["schemas"]["Step"]["oneOf"] if "$ref" in option } - assert {"UserInputStep", "ModelOutputStep"} <= step_variants, f"Step union is missing role steps: {step_variants}" + assert {"UserInputStep", "ModelOutputStep"} <= step_variants, ( + f"Step union is missing role steps: {step_variants}" + ) for step_name, type_value in [("UserInputStep", "user_input"), ("ModelOutputStep", "model_output")]: step_schema = spec_dict["components"]["schemas"][step_name] @@ -261,9 +271,7 @@ class TestResponseCompliance: expected_fields = ["total_input_tokens", "total_output_tokens", "total_tokens"] for field in expected_fields: - assert ( - field in usage_schema["properties"] - ), f"Usage field '{field}' not in spec" + assert field in usage_schema["properties"], f"Usage field '{field}' not in spec" print(f"✓ Usage field '{field}' exists") @@ -282,9 +290,7 @@ class TestToolsCompliance: """Verify FunctionDeclaration schema for function tools.""" if "FunctionDeclaration" in spec_dict["components"]["schemas"]: func_schema = spec_dict["components"]["schemas"]["FunctionDeclaration"] - assert "name" in func_schema.get( - "properties", {} - ) or "name" in func_schema.get("required", []) + assert "name" in func_schema.get("properties", {}) or "name" in func_schema.get("required", []) print("✓ FunctionDeclaration schema found") else: print("⚠ FunctionDeclaration schema not found (may be nested)") @@ -313,7 +319,7 @@ class TestEndpointCompliance: get_path = None for path, methods in paths.items(): - if "{id}" in path and "interactions" in path and "get" in methods: + if "/interactions/{" in path and path.endswith("}") and "get" in methods: get_path = path break @@ -326,7 +332,7 @@ class TestEndpointCompliance: delete_path = None for path, methods in paths.items(): - if "{id}" in path and "interactions" in path and "delete" in methods: + if "/interactions/{" in path and path.endswith("}") and "delete" in methods: delete_path = path break @@ -350,6 +356,4 @@ if __name__ == "__main__": if method in ["get", "post", "delete", "put", "patch"]: print(f" {method.upper()} {path}") - print( - f"\nSchemas: {list(spec.get('components', {}).get('schemas', {}).keys())[:10]}..." - ) + print(f"\nSchemas: {list(spec.get('components', {}).get('schemas', {}).keys())[:10]}...") 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 ff131fff32e..bd635193ba6 100644 --- a/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py +++ b/tests/unit/proxy/management_endpoints/test_roi_calculator_endpoints.py @@ -137,3 +137,53 @@ def test_github_api_url_must_use_https() -> None: assert response.status_code == 422 assert not repository.values + + +@pytest.mark.parametrize("role", [LitellmUserRoles.INTERNAL_USER, LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY]) +@pytest.mark.parametrize( + "method,path,body", + [ + ("POST", "/roi-calculator/sync", {}), + ("DELETE", "/roi-calculator/sync", {}), + ("POST", "/roi-calculator/setup/reset", {}), + ("POST", "/roi-calculator/connections/test", {}), + ("PUT", "/roi-calculator/identity-map", {"github_login": "alice", "email": "alice@example.com"}), + ], +) +def test_all_writes_require_full_admin(role: LitellmUserRoles, method: str, path: str, body: Mapping[str, str]) -> None: + client: Final = _client(role, _ConfigRepository()) + assert client.request(method, path, json=body).status_code == 403 + + +def test_schedule_and_estimator_key_persist_without_exposing_secrets(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("LITELLM_SALT_KEY", "roi-calculator-test-salt-key-0123456789") + repository: Final = _ConfigRepository() + client: Final = _client(LitellmUserRoles.PROXY_ADMIN, repository) + saved: Final = client.put( + "/roi-calculator/settings", json={"estimator_key": "sk-test-secret", "update_interval_minutes": 60} + ) + assert saved.status_code == 200 + assert saved.json()["has_estimator_key"] is True + assert saved.json()["update_interval_minutes"] == 60 + assert "sk-test-secret" not in saved.text + assert "sk-test-secret" not in str(repository.values) + updated: Final = client.put("/roi-calculator/settings", json={"estimator_key": None, "update_interval_minutes": 0}) + assert updated.json()["has_estimator_key"] is False + assert updated.json()["update_interval_minutes"] == 0 + + +def test_sample_preview_does_not_change_live_settings_or_report() -> None: + repository: Final = _ConfigRepository() + client: Final = _client(LitellmUserRoles.PROXY_ADMIN_VIEW_ONLY, repository) + response: Final = client.get("/roi-calculator/report", params={"mode": "demo"}) + assert response.status_code == 200 + assert response.json()["report"]["mode"] == "demo" + assert response.json()["report"]["metrics"]["cost_per_hour"] > 0 + assert not repository.values + assert client.get("/roi-calculator/report").json()["report"] is None + + +@pytest.mark.parametrize("interval", [0.1, 1, 4.99]) +def test_schedule_rejects_intervals_under_five_minutes(interval: float) -> None: + client: Final = _client(LitellmUserRoles.PROXY_ADMIN, _ConfigRepository()) + assert client.put("/roi-calculator/settings", json={"update_interval_minutes": interval}).status_code == 422 diff --git a/tests/unit/proxy/roi_calculator/test_sync.py b/tests/unit/proxy/roi_calculator/test_sync.py index aeaaa72eb91..6576a995ba3 100644 --- a/tests/unit/proxy/roi_calculator/test_sync.py +++ b/tests/unit/proxy/roi_calculator/test_sync.py @@ -151,6 +151,7 @@ def _settings(estimator_prompt: str = "Estimate effort.") -> ROISettings: def _transport( pull_detail_status: int = 200, unexpected_details: bool = False, + profile_email: str = "alice@example.com", ) -> httpx.MockTransport: def respond(request: httpx.Request) -> httpx.Response: path = request.url.path @@ -166,7 +167,7 @@ def _transport( content=_PULL_FILES_JSON, ) if path == "/users/alice": - return httpx.Response(200, content=_USER_JSON) + return httpx.Response(200, json={"email": profile_email}) if path == "/repos/org/repo/pulls/42/commits": return httpx.Response(200, content=_COMMITS_JSON) raise AssertionError(f"Unexpected GitHub request: {request.method} {path}") @@ -211,23 +212,23 @@ async def _wait_until_finished(manager: SyncManager) -> None: @pytest.mark.asyncio -async def test_unchanged_estimated_pull_skips_github_details_and_model_call() -> None: +async def test_unchanged_estimated_pull_refreshes_identity_without_model_call() -> None: repository: Final = _ReportRepository() manager: Final = SyncManager(clock=_fixed_now) complete: Final = _completion() - assert manager.start(_settings(), repository, _spend_reader(), complete, _transport()) + assert await manager.start(_settings(), repository, _spend_reader(), complete, _transport()) await _wait_until_finished(manager) async def unexpected_completion(request: ROICompletionRequest) -> object: raise AssertionError("A reused estimate must not call the estimator.") - assert manager.start( + assert await manager.start( _settings(), repository, _spend_reader(), unexpected_completion, - _transport(unexpected_details=True), + _transport(unexpected_details=True, profile_email="new@example.com"), ) await _wait_until_finished(manager) @@ -235,6 +236,8 @@ async def test_unchanged_estimated_pull_skips_github_details_and_model_call() -> assert manager.status.reused == 1 report: Final = TypeAdapter(ROIReport).validate_python(repository.values["roi_calculator_report"]) assert report["pulls"][0]["estimate"].get("cached") is True + assert report["pulls"][0]["profile_email"] == "new@example.com" + assert report["pulls"][0]["emails"] == ("alice@example.com", "new@example.com") @pytest.mark.asyncio @@ -274,7 +277,7 @@ async def test_sync_does_not_persist_a_report_when_github_fails() -> None: repository: Final = _ReportRepository() manager: Final = SyncManager(clock=_fixed_now) - assert manager.start( + assert await manager.start( _settings(), repository, _spend_reader(), @@ -293,7 +296,7 @@ async def test_cancelling_estimation_leaves_the_previous_report_unchanged() -> N repository: Final = _ReportRepository() manager: Final = SyncManager(clock=_fixed_now) - assert manager.start(_settings(), repository, _spend_reader(), _completion(), _transport()) + assert await manager.start(_settings(), repository, _spend_reader(), _completion(), _transport()) await _wait_until_finished(manager) previous_report: Final = repository.values["roi_calculator_report"] @@ -302,7 +305,7 @@ async def test_cancelling_estimation_leaves_the_previous_report_unchanged() -> N entered_estimator.set() await asyncio.Event().wait() - assert manager.start( + assert await manager.start( _settings(estimator_prompt="Different estimator instructions."), repository, _spend_reader(), @@ -314,3 +317,38 @@ async def test_cancelling_estimation_leaves_the_previous_report_unchanged() -> N assert await manager.cancel() assert manager.status.phase == "cancelled" assert repository.values["roi_calculator_report"] is previous_report + + +@pytest.mark.asyncio +async def test_immediate_cancel_allows_another_run() -> None: + repository: Final = _ReportRepository() + manager: Final = SyncManager(clock=_fixed_now) + assert await manager.start(_settings(), repository, _spend_reader(), _completion(), _transport()) + assert await manager.cancel() + assert manager.status.phase == "cancelled" + assert manager.status.finished_at is not None + assert await manager.start(_settings(), repository, _spend_reader(), _completion(), _transport()) + await _wait_until_finished(manager) + assert manager.status.phase == "complete" + + +@pytest.mark.asyncio +async def test_saved_estimates_survive_report_reset() -> None: + repository: Final = _ReportRepository() + manager: Final = SyncManager(clock=_fixed_now) + assert await manager.start(_settings(), repository, _spend_reader(), _completion(), _transport()) + await _wait_until_finished(manager) + repository.values = MappingProxyType( + {key: value for key, value in repository.values.items() if key != "roi_calculator_report"} + ) + + async def unexpected_completion(request: ROICompletionRequest) -> object: + raise AssertionError("Saved estimates should survive report reset") + + restarted: Final = SyncManager(clock=_fixed_now) + assert await restarted.start( + _settings(), repository, _spend_reader(), unexpected_completion, _transport(unexpected_details=True) + ) + await _wait_until_finished(restarted) + assert restarted.status.phase == "complete" + assert restarted.status.reused == 1 diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorDialogs.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorDialogs.tsx index 56bd60ea1a7..d7f73c84cb1 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorDialogs.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorDialogs.tsx @@ -73,9 +73,11 @@ export function PullReasoningDialog({ )} - + {pull.url && ( + + )} )} @@ -87,11 +89,13 @@ export function PullReasoningDialog({ export function IdentityMatchDialog({ selection, identityMap, + gatewayEmails, onClose, onSave, }: { selection: PersonMatchSelection | null; identityMap: Record; + gatewayEmails: string[]; onClose: () => void; onSave: (payload: ROIIdentityMapUpdate) => Promise; }) { @@ -137,12 +141,18 @@ export function IdentityMatchDialog({ setEmail(event.target.value)} required /> + + {Array.from(new Set(gatewayEmails)).map((address) => ( + {error && (

{error} diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.integration.test.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.integration.test.tsx index ab6848766bf..3c7965e917f 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.integration.test.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.integration.test.tsx @@ -133,6 +133,7 @@ describe("ROICalculatorView", () => { beforeEach(() => { vi.mocked(apiClient.get).mockReset(); vi.mocked(apiClient.put).mockReset(); + vi.mocked(apiClient.post).mockReset(); vi.mocked(apiClient.get).mockImplementation((path: string) => { if (path === "/roi-calculator/settings") return Promise.resolve(settings); if (path === "/roi-calculator/report") return Promise.resolve({ report: summary }); @@ -192,7 +193,7 @@ describe("ROICalculatorView", () => { fireEvent.click(await screen.findByRole("tab", { name: "People" })); fireEvent.click(await screen.findByRole("button", { name: "alice-work" })); - fireEvent.change(screen.getByRole("textbox", { name: "Gateway email" }), { + fireEvent.change(screen.getByLabelText("Gateway email"), { target: { value: "alice+work@example.com" }, }); fireEvent.click(screen.getByRole("button", { name: "Save match" })); @@ -206,8 +207,9 @@ describe("ROICalculatorView", () => { }); it("presents onboarding settings once when no report exists", async () => { + const emptySettings = { ...settings, has_github_token: false, ready: false, repos: [], estimator_model: "" }; vi.mocked(apiClient.get).mockImplementation((path: string) => { - if (path === "/roi-calculator/settings") return Promise.resolve(settings); + if (path === "/roi-calculator/settings") return Promise.resolve(emptySettings); if (path === "/roi-calculator/report") return Promise.resolve({ report: null }); return Promise.resolve(idleStatus); }); @@ -243,10 +245,9 @@ describe("ROICalculatorView", () => { render(); expect(await screen.findByRole("progressbar", { name: "Sync progress" })).toBeInTheDocument(); - fireEvent.click(screen.getByRole("tab", { name: "Settings" })); expect(await screen.findByText("Spend per estimated engineering hour", {}, { timeout: 5000 })).toBeInTheDocument(); expect(screen.queryByRole("heading", { name: "Connect GitHub to get started" })).not.toBeInTheDocument(); - expect(screen.getByRole("status")).toHaveTextContent("Up to date · Last synced Sep 30, 2026, 12:00 PM UTC"); + expect(screen.getByRole("status")).toHaveTextContent("Last synced Sep 30, 2026, 12:00 PM UTC"); expect(screen.getByRole("status")).toHaveTextContent("57 of 57 estimates reused"); }); @@ -328,4 +329,22 @@ describe("ROICalculatorView", () => { expect(await screen.findByText("Spend per estimated engineering hour", {}, { timeout: 7000 })).toBeInTheDocument(); expect(screen.queryByText("The sync status could not be loaded.")).not.toBeInTheDocument(); }); + it("saves the edited schedule before running from Settings", async () => { + vi.mocked(apiClient.put).mockResolvedValue(settings); + vi.mocked(apiClient.post).mockResolvedValue({ ...idleStatus, running: true }); + render(); + fireEvent.click(await screen.findByRole("tab", { name: "Settings" })); + fireEvent.change(screen.getByLabelText("Update interval (hours)"), { target: { value: "6" } }); + fireEvent.click(screen.getByRole("button", { name: "Save and run analysis" })); + await waitFor(() => expect(apiClient.post).toHaveBeenCalledWith("/roi-calculator/sync", { accessToken: "token" })); + expect(apiClient.put).toHaveBeenCalledWith( + "/roi-calculator/settings", + expect.objectContaining({ + body: expect.objectContaining({ update_interval_minutes: 360, estimator_model: "estimator" }), + }), + ); + expect(vi.mocked(apiClient.put).mock.invocationCallOrder[0]).toBeLessThan( + vi.mocked(apiClient.post).mock.invocationCallOrder[0], + ); + }); }); diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.tsx index 57d8f4d5c2c..b9bf304447c 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorView.tsx @@ -49,14 +49,19 @@ export default function ROICalculatorView({ userRole?: string | null; isViewOnly?: boolean; }) { - const readOnly = isViewOnly && isProxyAdminTierRole(userRole ?? ""); + const [sampleSummary, setSampleSummary] = React.useState(null); + const adminReadOnly = isViewOnly && isProxyAdminTierRole(userRole ?? ""); + const readOnly = adminReadOnly || sampleSummary !== null; const [view, setView] = React.useState("overview"); const [settings, setSettings] = React.useState(null); - const [summary, setSummary] = React.useState(null); + const [liveSummary, setSummary] = React.useState(null); + const summary = sampleSummary ?? liveSummary; const [status, setStatus] = React.useState(IDLE_STATUS); const [selectedPull, setSelectedPull] = React.useState(null); const [matchingPerson, setMatchingPerson] = React.useState(null); const [error, setError] = React.useState(null); + const statusRef = React.useRef(IDLE_STATUS); + const settingsLoaded = settings !== null; const [query, setQuery] = React.useState(""); const loadReport = React.useCallback(async () => { @@ -78,6 +83,7 @@ export default function ROICalculatorView({ setSettings(nextSettings); setSummary(reportResponse.report); setStatus(syncStatus); + statusRef.current = syncStatus; setError(null); }) .catch((reason: unknown) => { @@ -89,9 +95,10 @@ export default function ROICalculatorView({ }, [accessToken]); React.useEffect(() => { - if (!accessToken || !status.running) return; + if (!accessToken || !settingsLoaded) return; let cancelled = false; let requestInFlight = false; + let reportNeedsRefresh = false; const interval = window.setInterval(() => { if (requestInFlight) return; requestInFlight = true; @@ -99,19 +106,20 @@ export default function ROICalculatorView({ .get("/roi-calculator/sync", { accessToken }) .then(async (nextStatus) => { if (cancelled) return; - setError(null); - if (!nextStatus.running && nextStatus.phase === "complete") { - try { - const report = await loadReport(); - if (cancelled) return; - setSummary(report); - setError(null); - if (view === "settings") setView("overview"); - } catch (reason) { - if (!cancelled) setError(extractErrorMessage(reason)); - } + const previousStatus = statusRef.current; + statusRef.current = nextStatus; + setStatus(nextStatus); + const finished = !nextStatus.running && nextStatus.phase === "complete"; + const reportChanged = previousStatus.running || nextStatus.finished_at !== previousStatus.finished_at; + if (finished && (reportChanged || reportNeedsRefresh)) { + reportNeedsRefresh = true; + const report = await loadReport(); + if (cancelled) return; + setSummary(report); + reportNeedsRefresh = false; + setView((current) => (current === "settings" ? "overview" : current)); } - if (!cancelled) setStatus(nextStatus); + if (!cancelled) setError(null); }) .catch((reason: unknown) => { if (!cancelled) setError(extractErrorMessage(reason)); @@ -124,13 +132,15 @@ export default function ROICalculatorView({ cancelled = true; window.clearInterval(interval); }; - }, [accessToken, loadReport, status.running, view]); + }, [accessToken, loadReport, settingsLoaded]); const startSync = React.useCallback(async () => { if (!accessToken || readOnly) return; try { setError(null); - setStatus(await apiClient.post("/roi-calculator/sync", { accessToken })); + const nextStatus = await apiClient.post("/roi-calculator/sync", { accessToken }); + statusRef.current = nextStatus; + setStatus(nextStatus); } catch (reason) { setError(extractErrorMessage(reason)); } @@ -180,6 +190,27 @@ export default function ROICalculatorView({ ); } + const previewSample = async () => { + try { + const response = await apiClient.get("/roi-calculator/report", { + accessToken, + query: { mode: "demo" }, + }); + setSampleSummary(response.report); + setView("overview"); + } catch (reason) { + setError(extractErrorMessage(reason)); + } + }; + const resetView = (updated: ROISettings) => { + setSettings(updated); + setSummary(null); + setView("overview"); + setStatus(IDLE_STATUS); + statusRef.current = IDLE_STATUS; + }; + const showLiveStatus = !sampleSummary && !status.running; + const scheduleLabel = settings.update_interval_minutes ? "Automatic updates enabled" : "Manual updates"; const progress = status.total > 0 ? Math.min(100, (status.done / status.total) * 100) : 0; const statusIsIdleOrComplete = status.phase === "idle" || status.phase === "complete"; const syncIsUpToDate = !status.running && statusIsIdleOrComplete; @@ -197,7 +228,7 @@ export default function ROICalculatorView({ : "Compare gateway spend with estimated engineering effort for merged pull requests"} {syncedAt && ( - Up to date · Last synced {formatSyncedAt(syncedAt)} + Last synced {formatSyncedAt(syncedAt)} {!status.running && status.phase === "complete" && status.reused > 0 ? ` · ${status.reused} of ${status.total} estimates reused` : ""} @@ -206,27 +237,50 @@ export default function ROICalculatorView({ } /> - {readOnly && ( + {!liveSummary && showLiveStatus && ( + + )} + {sampleSummary && ( + + Sample report + + Example data only. No GitHub or model requests were made. + + + + )} + {liveSummary && showLiveStatus && ( +

+ {status.next_update ? `Next update ${formatSyncedAt(status.next_update)}` : scheduleLabel} +

+ )} + {adminReadOnly && (

Read-only access. Settings, analysis runs, and email matches are unavailable.

)} -
- setView(value as View)}> - - Overview - People - Settings - - - {view !== "settings" && !readOnly && ( - - )} -
+ {summary && ( +
+ setView(value as View)}> + + Overview + People + {!sampleSummary && Settings} + + + {view !== "settings" && !readOnly && ( + + )} +
+ )} {error && ( @@ -263,6 +317,8 @@ export default function ROICalculatorView({

{status.done} of {status.total} pull requests processed · {status.reused} reused + {` · ${status.elapsed_seconds ?? 0}s elapsed`} + {status.remaining_seconds != null ? ` · about ${status.remaining_seconds}s remaining` : ""}

{!readOnly && ( @@ -276,11 +332,11 @@ export default function ROICalculatorView({ {view === "settings" || (!summary && !status.running) ? ( (person.email ? [person.email] : [])) ?? []} onClose={() => setMatchingPerson(null)} onSave={updateIdentity} /> diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorViews.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorViews.tsx index dff2ef6ffca..75df18752de 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorViews.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROICalculatorViews.tsx @@ -15,7 +15,7 @@ import { import type { ChartConfig } from "@/components/ui/chart"; import { Input } from "@/components/ui/input"; import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from "@/components/ui/table"; -import { coverageLabel, effortNote, estimateLabel, formatMoney, formatNumber } from "./roiCalculatorData"; +import { coverageLabel, peopleCsv, effortNote, estimateLabel, formatMoney, formatNumber } from "./roiCalculatorData"; import type { ROIPerson, ROIPull, ROISummary } from "./roiCalculatorData"; const CHART_CONFIG = { @@ -207,8 +207,21 @@ export function ROIPeopleView({ onMatch: (person: ROIPerson, login: string) => void; readOnly?: boolean; }) { + const exportCsv = () => { + const url = URL.createObjectURL(new Blob([peopleCsv(summary)], { type: "text/csv;charset=utf-8" })); + const link = document.createElement("a"); + link.href = url; + link.download = "litellm-roi.csv"; + link.click(); + window.setTimeout(() => URL.revokeObjectURL(url), 1000); + }; return (
+
+ +

{effortNote(summary.effort_basis)} Spend includes each person’s full gateway usage for this period. This does not measure hours saved by AI or financial returns. @@ -247,8 +260,9 @@ export function ROIPeopleView({ ) : ( Unassigned gateway spend )} - {person.match_methods.some((method) => - ["manual", "commit email", "profile email"].includes(method), + {person.match_methods.some( + (method) => + ["manual", "commit email", "profile email"].includes(method) && person.spend != null, ) ? ( Matched ) : ( diff --git a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROISettingsPanel.tsx b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROISettingsPanel.tsx index c37d71d7d7a..977d0dbc760 100644 --- a/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROISettingsPanel.tsx +++ b/ui/litellm-dashboard/src/app/(dashboard)/roi-calculator/_components/ROISettingsPanel.tsx @@ -8,6 +8,14 @@ import { Button } from "@/components/ui/button"; import { Card, CardContent, CardDescription, CardHeader } from "@/components/ui/card"; import { Input } from "@/components/ui/input"; import { Label } from "@/components/ui/label"; +import { + Dialog, + DialogContent, + DialogHeader, + DialogTitle, + DialogDescription, + DialogFooter, +} from "@/components/ui/dialog"; import { Textarea } from "@/components/ui/textarea"; import type { ROIRepository, ROIRepositoriesResponse, ROISettings, ROISettingsUpdate } from "./roiCalculatorData"; @@ -16,6 +24,7 @@ export default function ROISettingsPanel({ initialSettings, onboarding, onSaved, + onReset, onStartSync, readOnly, syncDisabled, @@ -24,10 +33,13 @@ export default function ROISettingsPanel({ initialSettings: ROISettings; onboarding: boolean; onSaved: (settings: ROISettings) => void; + onReset: (settings: ROISettings) => void; onStartSync: () => Promise; readOnly: boolean; syncDisabled: boolean; }) { + const initialStep = initialSettings.has_github_token ? 1 : 0; + const [step, setStep] = React.useState(initialSettings.ready ? 2 : initialStep); const [apiUrl, setApiUrl] = React.useState(initialSettings.github_api_url); const [token, setToken] = React.useState(""); const [clearToken, setClearToken] = React.useState(false); @@ -35,6 +47,13 @@ export default function ROISettingsPanel({ const [model, setModel] = React.useState(initialSettings.estimator_model); const [prompt, setPrompt] = React.useState(initialSettings.estimator_prompt); const [backfillDays, setBackfillDays] = React.useState(String(initialSettings.backfill_days)); + const [intervalHours, setIntervalHours] = React.useState( + String((initialSettings.update_interval_minutes ?? 1440) / 60), + ); + const [estimatorKey, setEstimatorKey] = React.useState(""); + const [clearEstimatorKey, setClearEstimatorKey] = React.useState(false); + const [repositoryName, setRepositoryName] = React.useState(""); + const [resetOpen, setResetOpen] = React.useState(false); const [repositoryQuery, setRepositoryQuery] = React.useState(""); const [repositoryPage, setRepositoryPage] = React.useState(1); const [availableRepos, setAvailableRepos] = React.useState([]); @@ -65,15 +84,17 @@ export default function ROISettingsPanel({ } }; - const saveSettings = async (event: React.FormEvent) => { - event.preventDefault(); - if (!accessToken || readOnly) return; + const saveSettings = async () => { + if (!accessToken || readOnly) return false; const body: ROISettingsUpdate = { github_api_url: apiUrl, repos, estimator_model: model, estimator_prompt: prompt, backfill_days: Number(backfillDays), + update_interval_minutes: Number(intervalHours) * 60, + ...(clearEstimatorKey ? { estimator_key: null } : {}), + ...(estimatorKey.trim() ? { estimator_key: estimatorKey.trim() } : {}), ...(clearToken ? { github_token: null } : {}), ...(token.trim() ? { github_token: token.trim() } : {}), }; @@ -82,12 +103,65 @@ export default function ROISettingsPanel({ const updated: ROISettings = await apiClient.put("/roi-calculator/settings", { accessToken, body }); onSaved(updated); setToken(""); + setEstimatorKey(""); + setClearEstimatorKey(false); setClearToken(false); setMessage("Settings saved."); setError(null); + return true; } catch (reason) { setError(extractErrorMessage(reason)); setMessage(null); + return false; + } finally { + setBusy(false); + } + }; + + const submit = async (event: React.FormEvent) => { + event.preventDefault(); + if (!(await saveSettings())) return; + if (onboarding && step === 0) { + try { + const result = await apiClient.get("/roi-calculator/repositories", { accessToken }); + setAvailableRepos(result.repositories); + setHasMoreRepos(result.has_more); + setRepositoryPage(1); + setStep(1); + } catch (reason) { + setError(extractErrorMessage(reason)); + } + } else if (onboarding && step === 1) setStep(2); + else if (onboarding) await onStartSync(); + }; + + const saveAndRun = async () => { + if (await saveSettings()) await onStartSync(); + }; + + const testConnections = async () => { + if (!(await saveSettings())) return; + setBusy(true); + try { + await apiClient.post("/roi-calculator/connections/test", { accessToken }); + setMessage("Gateway model and selected repositories are available."); + } catch (reason) { + setError(extractErrorMessage(reason)); + } finally { + setBusy(false); + } + }; + + const resetSetup = async () => { + setBusy(true); + try { + const updated = await apiClient.post("/roi-calculator/setup/reset", { accessToken }); + setRepos([]); + setStep(updated.has_github_token ? 1 : 0); + setResetOpen(false); + onReset(updated); + } catch (reason) { + setError(extractErrorMessage(reason)); } finally { setBusy(false); } @@ -97,15 +171,25 @@ export default function ROISettingsPanel({ setRepos((current) => (current.includes(name) ? current.filter((repo) => repo !== name) : [...current, name])); }; + const formDisabled = busy || syncDisabled; + const runDisabled = formDisabled || !repos.length || !model; + const githubUrlChanged = apiUrl !== initialSettings.github_api_url; + const missingReplacementToken = initialSettings.has_github_token && githubUrlChanged && !token.trim(); + const stepReady = [Boolean(token.trim() || initialSettings.has_github_token), repos.length > 0, Boolean(model)][step]; + const onboardingLabel = step < 2 ? "Continue" : "Start backfill"; + const submitLabel = onboarding ? onboardingLabel : "Save settings"; + return (

- {onboarding ? "Connect GitHub to get started" : "ROI Calculator settings"} + {onboarding + ? ["Connect GitHub to get started", "Choose repositories", "Choose an estimator"][step] + : "ROI Calculator settings"}

{onboarding - ? "Save a GitHub token, choose repositories and a router model, then run the analysis." + ? "Your gateway is already connected. Set up GitHub and an estimator to see your first report." : "Choose GitHub repositories and the router model used for metadata-only estimates."} @@ -120,168 +204,317 @@ export default function ROISettingsPanel({ {message}

)} -
void saveSettings(event)}> -
- - setApiUrl(event.target.value)} - /> -
-
- - { - setToken(event.target.value); - setClearToken(false); - }} - placeholder={initialSettings.has_github_token ? "Token saved" : "Enter a GitHub token"} - /> -

- {initialSettings.has_github_token - ? "A token is saved securely and is never shown here." - : "Save a token to list repositories and read private repository metadata."} -

- {initialSettings.has_github_token && apiUrl !== initialSettings.github_api_url && !token.trim() && ( -

- Changing the GitHub API URL clears the saved token. Enter a replacement token to keep access. -

- )} - {initialSettings.has_github_token && ( - - )} -
-
- -
- setRepositoryQuery(event.target.value)} - placeholder="Search repositories" - /> - -
- {!canLoadRepositories && ( -

- Save the GitHub token and API URL before loading repositories. -

- )} - {repos.length > 0 &&

Selected: {repos.join(", ")}

} -
- {availableRepos.map((repository) => ( -
+ )} -
-
- - -
-
- -