fix(proxy): defer billing for disconnected tinyfish SSE runs to the background poller

This commit is contained in:
Zachary Lyon 2026-09-15 15:22:51 -07:00
parent 6ddb9441fd
commit ef0bcfb4a0
3 changed files with 71 additions and 1 deletions

View file

@ -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,

View file

@ -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,

View file

@ -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=[{}])