diff --git a/litellm/proxy/lens/endpoints.py b/litellm/proxy/lens/endpoints.py index d11a6b4fad6..04f178a3107 100644 --- a/litellm/proxy/lens/endpoints.py +++ b/litellm/proxy/lens/endpoints.py @@ -22,6 +22,7 @@ from litellm.proxy.lens.inference import Deployment, deployment_prices from litellm.proxy.lens.models import ( ActivitySelection, Claim, + Coverage, Execution, ExecutionContent, FindingDraft, @@ -632,7 +633,7 @@ async def result(lens_id: str, job_id: str, body: Result, worker: WorkerAuth, st end_job(active, "failed" if body.error else "completed", now).model_copy( update=MappingProxyType( { - "coverage": active.coverage if body.error else body.coverage, + "coverage": active.coverage if body.error and body.coverage == Coverage() else body.coverage, "error": body.error, "assessments": body.assessments, "findings": tuple(snapshot_finding(e, f, job.revision, now) for f in body.findings), diff --git a/tests/unit/proxy/lens/test_endpoints.py b/tests/unit/proxy/lens/test_endpoints.py index ec09cce00f4..922bcb08cca 100644 --- a/tests/unit/proxy/lens/test_endpoints.py +++ b/tests/unit/proxy/lens/test_endpoints.py @@ -1,4 +1,5 @@ from datetime import datetime, timedelta, timezone +from types import SimpleNamespace from typing import Final import pytest @@ -11,6 +12,7 @@ from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth from litellm.proxy.lens.endpoints import ( list_agents, read_reviews, + result, run_settings, run_window, user_scope, @@ -19,7 +21,68 @@ from litellm.proxy.lens.endpoints import ( watching, worker_supports_model, ) -from litellm.proxy.lens.models import ActivitySelection, Lens, LensSettings, RunRequest, Scope +from litellm.proxy.lens.models import ActivitySelection, Coverage, Lens, LensSettings, Result, RunRequest, Scope +from litellm.proxy.lens.repository import Row +from litellm.proxy.lens.state import claim_job, queue_job, replace_job +from tests.unit.proxy.lens.test_state import NOW, lens, worker + + +class ResultDatabase: + def __init__(self, stored: Lens) -> None: + self.stored = stored + + async def query_raw(self, query: str, *args: object) -> tuple[Row, ...]: + if query.startswith("SELECT data FROM"): + return (Row(data=self.stored.model_dump(mode="json")),) + payload: Final = args[0] + assert isinstance(payload, str) + self.stored = Lens.model_validate_json(payload) + return (Row(data=1),) + + +@pytest.mark.parametrize( + "final_coverage,error,expected", + ( + ( + Coverage(eligible=2, selected=2, screened=2, partial=1, unassessable=1), + "Source unavailable during session review", + Coverage(eligible=2, selected=2, screened=2, partial=1, unassessable=1), + ), + ( + Coverage(eligible=2, selected=2, screened=2, investigated=1, candidates=1, partial=1), + "Source unavailable during investigation", + Coverage(eligible=2, selected=2, screened=2, investigated=1, candidates=1, partial=1), + ), + ( + Coverage(), + "Worker interrupted", + Coverage(eligible=2, selected=2, screened=1), + ), + (Coverage(), "", Coverage()), + ), + ids=("review-diagnostic", "investigation-diagnostic", "interrupted-worker", "empty-success"), +) +@pytest.mark.asyncio +async def test_result_persists_final_coverage_but_keeps_progress_when_worker_is_interrupted( + monkeypatch: pytest.MonkeyPatch, final_coverage: Coverage, error: str, expected: Coverage +) -> None: + from litellm.proxy import proxy_server + + assigned: Final = claim_job(queue_job(lens(), NOW, "job"), worker(), NOW) + active: Final = assigned.jobs[0].model_copy( + update={ + "lease_until": datetime.max.replace(tzinfo=timezone.utc), + "coverage": Coverage(eligible=2, selected=2, screened=1), + } + ) + db: Final = ResultDatabase(replace_job(assigned, active)) + monkeypatch.setattr(proxy_server, "prisma_client", SimpleNamespace(db=db)) + saved: Final = await result("lens", "job", Result(coverage=final_coverage, error=error), worker(), None) + + assert saved == db.stored + assert saved.jobs[0].coverage == expected + assert saved.jobs[0].error == error + assert saved.jobs[0].status == ("failed" if error else "completed") @pytest.fixture