mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
fix(proxy): defer billing for disconnected tinyfish SSE runs to the background poller
This commit is contained in:
parent
6ddb9441fd
commit
ef0bcfb4a0
3 changed files with 71 additions and 1 deletions
|
|
@ -156,6 +156,25 @@ class TinyFishPassthroughLoggingHandler:
|
|||
verbose_proxy_logger.warning(
|
||||
"TinyFish passthrough: run-async response carried no run_id; logging the request without cost"
|
||||
)
|
||||
TinyFishPassthroughLoggingHandler._spawn_billing_poller(
|
||||
run_id=run_id,
|
||||
logging_obj=logging_obj,
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
cache_hit=cache_hit,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _spawn_billing_poller(
|
||||
run_id: str | None,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
result: str,
|
||||
start_time: datetime,
|
||||
cache_hit: bool,
|
||||
kwargs: Mapping[str, object],
|
||||
client: AsyncHTTPHandler | None = None,
|
||||
) -> None:
|
||||
task: Final = asyncio.create_task(
|
||||
TinyFishPassthroughLoggingHandler._poll_and_log(
|
||||
run_id=run_id,
|
||||
|
|
@ -164,6 +183,7 @@ class TinyFishPassthroughLoggingHandler:
|
|||
start_time=start_time,
|
||||
cache_hit=cache_hit,
|
||||
kwargs=kwargs,
|
||||
client=client,
|
||||
)
|
||||
)
|
||||
_BACKGROUND_BILLING_TASKS.add(task)
|
||||
|
|
@ -286,7 +306,9 @@ class TinyFishPassthroughLoggingHandler:
|
|||
client: AsyncHTTPHandler | None = None,
|
||||
) -> PassThroughEndpointLoggingTypedDict:
|
||||
"""Bill a POST /v1/automation/run-sse stream: SSE events carry no num_of_steps, so the
|
||||
run_id parsed from the buffered events prices the run via one GET /v1/runs/{id}."""
|
||||
run_id parsed from the buffered events prices the run via one GET /v1/runs/{id}. A run
|
||||
still live at stream end (client disconnect) is handed to the background poller instead,
|
||||
signalled by a None result, so its steps are still billed when it finishes."""
|
||||
try:
|
||||
run_id: Final = _run_id_from_sse_chunks(all_chunks)
|
||||
if run_id is None:
|
||||
|
|
@ -294,6 +316,18 @@ class TinyFishPassthroughLoggingHandler:
|
|||
"TinyFish passthrough: no run_id in SSE stream; logging the request without cost"
|
||||
)
|
||||
run: Final = await TinyFishPassthroughLoggingHandler._fetch_run(run_id, client) if run_id else None
|
||||
if run_id is not None and (run is None or (run.get("status") or "") not in TINYFISH_TERMINAL_RUN_STATUSES):
|
||||
TinyFishPassthroughLoggingHandler._spawn_billing_poller(
|
||||
run_id=run_id,
|
||||
logging_obj=litellm_logging_obj,
|
||||
result="",
|
||||
start_time=start_time,
|
||||
cache_hit=False,
|
||||
kwargs=_EMPTY_KWARGS,
|
||||
client=client,
|
||||
)
|
||||
deferred_payload: Final[PassThroughEndpointLoggingTypedDict] = {"result": None, "kwargs": {}}
|
||||
return deferred_payload
|
||||
payload: Final = TinyFishPassthroughLoggingHandler._build_logging_payload(
|
||||
run=run,
|
||||
logging_obj=litellm_logging_obj,
|
||||
|
|
|
|||
|
|
@ -286,6 +286,9 @@ class PassThroughStreamingHandler:
|
|||
end_time=end_time,
|
||||
)
|
||||
)
|
||||
if tinyfish_payload["result"] is None:
|
||||
# run still live upstream (client disconnect); the background poller owns the spend log
|
||||
return
|
||||
await litellm_logging_obj.dispatch_success_handlers(
|
||||
result=tinyfish_payload["result"],
|
||||
start_time=start_time,
|
||||
|
|
|
|||
|
|
@ -248,6 +248,39 @@ class TestSseBilling:
|
|||
assert payload["kwargs"]["model"] == "tinyfish/automation-run"
|
||||
assert fake_client.requested_urls == ["https://agent.tinyfish.ai/v1/runs/run-7?screenshots=none"]
|
||||
|
||||
def test_disconnected_stream_defers_billing_to_poller(self, tinyfish_env):
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.tinyfish_passthrough_logging_handler import (
|
||||
_BACKGROUND_BILLING_TASKS,
|
||||
)
|
||||
|
||||
logging_obj = _make_logging_obj()
|
||||
logging_obj.dispatch_success_handlers = AsyncMock()
|
||||
chunks = ['data: {"type": "STARTED", "run_id": "run-11", "status": "RUNNING"}']
|
||||
fake_client = _FakeClient(
|
||||
payloads=[
|
||||
{"run_id": "run-11", "status": "RUNNING"},
|
||||
{"run_id": "run-11", "status": "COMPLETED", "num_of_steps": 4},
|
||||
]
|
||||
)
|
||||
|
||||
async def _run():
|
||||
payload = await TinyFishPassthroughLoggingHandler.handle_logging_tinyfish_collected_chunks(
|
||||
litellm_logging_obj=logging_obj,
|
||||
url_route="https://agent.tinyfish.ai/v1/automation/run-sse",
|
||||
start_time=datetime.now(),
|
||||
all_chunks=chunks,
|
||||
end_time=datetime.now(),
|
||||
client=fake_client,
|
||||
)
|
||||
assert payload["result"] is None
|
||||
await asyncio.gather(*list(_BACKGROUND_BILLING_TASKS))
|
||||
|
||||
asyncio.run(_run())
|
||||
|
||||
logging_obj.dispatch_success_handlers.assert_awaited_once()
|
||||
assert logging_obj.dispatch_success_handlers.await_args.kwargs["response_cost"] == pytest.approx(0.064)
|
||||
assert len(fake_client.requested_urls) == 2
|
||||
|
||||
def test_stream_without_run_id_logs_without_cost(self, tinyfish_env):
|
||||
fake_client = _FakeClient(payloads=[{}])
|
||||
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue