diff --git a/litellm/proxy/lens/inference.py b/litellm/proxy/lens/inference.py index b5c4d128c6e..abeda033401 100644 --- a/litellm/proxy/lens/inference.py +++ b/litellm/proxy/lens/inference.py @@ -192,7 +192,7 @@ async def analyze( current, active.model_copy(update=MappingProxyType({"cost": active.cost + estimate})) ).model_copy(update=MappingProxyType({"spent": current.spent + estimate})) - def settle(e: Lens, cost: float) -> Lens: + def settle(e: Lens, cost: float, step: Step | None) -> Lens: charged: Final = next((j for j in e.jobs if j.id == job.id), None) adjusted: Final = ( e.model_copy(update=MappingProxyType({"spent": max(0, e.spent - estimate + cost)})) @@ -201,13 +201,8 @@ async def analyze( ) if charged is None: return adjusted - return replace_job( - adjusted, - add_step( - charged.model_copy(update=MappingProxyType({"cost": max(0, charged.cost - estimate + cost)})), - model_step(response, body, job.settings.model, cost), - ), - ) + refunded: Final = charged.model_copy(update=MappingProxyType({"cost": max(0, charged.cost - estimate + cost)})) + return replace_job(adjusted, add_step(refunded, step) if step is not None else refunded) @asynccontextmanager async def reserve_budget() -> AsyncIterator[None]: @@ -216,7 +211,7 @@ async def analyze( try: yield except BaseException: - await repo.update(lens.id, lambda e: settle(e, 0)) + await repo.update(lens.id, lambda e: settle(e, 0, None)) raise data: Final[dict[str, object]] = { # mutable-ok: proxy processing enriches request data @@ -243,7 +238,8 @@ async def analyze( response, billed_cost = await complete(worker.analysis_key_id, data, reserve_budget, request) cost: Final = billed_cost if billed_cost is not None else completion_charge(deployments, response, estimate) - await repo.update(lens.id, lambda e: settle(e, cost)) + step: Final = model_step(response, body, job.settings.model, cost) + await repo.update(lens.id, lambda e: settle(e, cost, step)) parsed: Final = Completion.model_validate_json(response.model_dump_json()) choice: Final = parsed.choices[0] return ModelResult(