diff --git a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/tinyfish_passthrough_logging_handler.py b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/tinyfish_passthrough_logging_handler.py index e5b83322a9a..70fcf921616 100644 --- a/litellm/proxy/pass_through_endpoints/llm_provider_handlers/tinyfish_passthrough_logging_handler.py +++ b/litellm/proxy/pass_through_endpoints/llm_provider_handlers/tinyfish_passthrough_logging_handler.py @@ -26,6 +26,7 @@ from litellm.types.llms.custom_http import httpxSpecialProvider from litellm.types.passthrough_endpoints.tinyfish import ( TINYFISH_AGENT_DEFAULT_API_BASE, TINYFISH_DEFAULT_COST_PER_STEP, + TINYFISH_MAX_CONSECUTIVE_POLL_FAILURES, TINYFISH_MAX_POLLING_SECONDS, TINYFISH_MODEL_NAME, TINYFISH_POLLING_INTERVAL_SECONDS, @@ -206,21 +207,39 @@ class TinyFishPassthroughLoggingHandler: verbose_proxy_logger.exception("[Non blocking logging error] TinyFish run-async billing failed: %s", e) @staticmethod - async def _poll_until_terminal(run_id: str, client: AsyncHTTPHandler | None = None) -> TinyfishRun | None: + async def _poll_until_terminal( + run_id: str, + client: AsyncHTTPHandler | None = None, + poll_interval_seconds: float = TINYFISH_POLLING_INTERVAL_SECONDS, + ) -> TinyfishRun | None: deadline: Final = time.monotonic() + TINYFISH_MAX_POLLING_SECONDS + last_run: TinyfishRun | None = None # rebind-ok: poll-loop state + consecutive_failures = 0 # rebind-ok: poll-loop state while time.monotonic() < deadline: run = await TinyFishPassthroughLoggingHandler._fetch_run(run_id, client) if run is None: - return None - if (run.get("status") or "") in TINYFISH_TERMINAL_RUN_STATUSES: - return run - await asyncio.sleep(TINYFISH_POLLING_INTERVAL_SECONDS) + # a single transient poll failure must not drop the run's charge + consecutive_failures += 1 + if consecutive_failures >= TINYFISH_MAX_CONSECUTIVE_POLL_FAILURES: + verbose_proxy_logger.warning( + "TinyFish passthrough: giving up on run %s after %s consecutive poll failures; " + "logging the request without cost", + run_id, + consecutive_failures, + ) + return last_run + else: + consecutive_failures = 0 + last_run = run + if (run.get("status") or "") in TINYFISH_TERMINAL_RUN_STATUSES: + return run + await asyncio.sleep(poll_interval_seconds) verbose_proxy_logger.warning( "TinyFish passthrough: run %s not terminal after %ss; logging the request without cost", run_id, TINYFISH_MAX_POLLING_SECONDS, ) - return None + return last_run @staticmethod async def _fetch_run(run_id: str, client: AsyncHTTPHandler | None = None) -> TinyfishRun | None: diff --git a/litellm/types/passthrough_endpoints/tinyfish.py b/litellm/types/passthrough_endpoints/tinyfish.py index a8e2f51818c..78642f0be42 100644 --- a/litellm/types/passthrough_endpoints/tinyfish.py +++ b/litellm/types/passthrough_endpoints/tinyfish.py @@ -9,6 +9,7 @@ TINYFISH_DEFAULT_COST_PER_STEP: Final = 0.016 TINYFISH_MODEL_NAME: Final = "tinyfish/automation-run" TINYFISH_POLLING_INTERVAL_SECONDS: Final = 5.0 TINYFISH_MAX_POLLING_SECONDS: Final = 1200.0 +TINYFISH_MAX_CONSECUTIVE_POLL_FAILURES: Final = 3 TINYFISH_TERMINAL_RUN_STATUSES: Final = frozenset({"COMPLETED", "FAILED", "CANCELLED"}) diff --git a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_tinyfish_passthrough_logging_handler.py b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_tinyfish_passthrough_logging_handler.py index 1b894850de5..5f5dde06c53 100644 --- a/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_tinyfish_passthrough_logging_handler.py +++ b/tests/test_litellm/proxy/pass_through_endpoints/llm_provider_handlers/test_tinyfish_passthrough_logging_handler.py @@ -33,15 +33,18 @@ def _make_response(method: str, url: str, body: dict) -> httpx.Response: class _FakeClient: - def __init__(self, payloads: list[dict], status_code: int = 200): + """Payload items are dicts served with status_code, or (status, dict) tuples for scripted failures.""" + + def __init__(self, payloads: list, status_code: int = 200): self.payloads = payloads self.status_code = status_code self.requested_urls: list[str] = [] async def get(self, url: str, headers: dict) -> httpx.Response: self.requested_urls.append(url) - payload = self.payloads[min(len(self.requested_urls) - 1, len(self.payloads) - 1)] - return httpx.Response(self.status_code, text=json.dumps(payload), request=httpx.Request("GET", url)) + item = self.payloads[min(len(self.requested_urls) - 1, len(self.payloads) - 1)] + status, payload = item if isinstance(item, tuple) else (self.status_code, item) + return httpx.Response(status, text=json.dumps(payload), request=httpx.Request("GET", url)) @pytest.fixture @@ -179,6 +182,29 @@ class TestRunAsyncBilling: assert fake_client.requested_urls == ["https://agent.tinyfish.ai/v1/runs/run-9?screenshots=none"] assert logging_obj.model_call_details["response_cost"] == pytest.approx(0.064) + def test_transient_poll_failure_keeps_polling(self, tinyfish_env): + fake_client = _FakeClient( + payloads=[(500, {}), {"run_id": "run-9", "status": "COMPLETED", "num_of_steps": 3}] + ) + + run = asyncio.run( + TinyFishPassthroughLoggingHandler._poll_until_terminal("run-9", fake_client, poll_interval_seconds=0.0) + ) + + assert run is not None + assert run["num_of_steps"] == 3 + assert len(fake_client.requested_urls) == 2 + + def test_gives_up_after_consecutive_poll_failures(self, tinyfish_env): + fake_client = _FakeClient(payloads=[(500, {})]) + + run = asyncio.run( + TinyFishPassthroughLoggingHandler._poll_until_terminal("run-9", fake_client, poll_interval_seconds=0.0) + ) + + assert run is None + assert len(fake_client.requested_urls) == 3 + def test_traversal_run_id_is_rejected(self, tinyfish_env): fake_client = _FakeClient(payloads=[{}])