mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-03 02:22:24 +00:00
feat(proxy): add TinyFish Agent API passthrough with per-step billing
This commit is contained in:
parent
dab7f6a86a
commit
13d4d2c5e2
13 changed files with 1076 additions and 0 deletions
|
|
@ -100,6 +100,7 @@ GATEWAY_PATH_PREFIXES: tuple[str, ...] = (
|
|||
"/cursor/",
|
||||
"/milvus/",
|
||||
"/openai_passthrough/",
|
||||
"/tinyfish/",
|
||||
# Dynamic provider / toolset passthrough (path templates)
|
||||
"/{provider}/",
|
||||
"/toolset/",
|
||||
|
|
|
|||
|
|
@ -207,6 +207,7 @@ LAZY_FEATURES: Final[tuple[LazyFeature, ...]] = (
|
|||
"/mistral/",
|
||||
"/openai/",
|
||||
"/openai_passthrough/",
|
||||
"/tinyfish/",
|
||||
"/vertex-ai/",
|
||||
"/vertex_ai/",
|
||||
"/vllm/",
|
||||
|
|
|
|||
|
|
@ -19941,6 +19941,96 @@
|
|||
]
|
||||
}
|
||||
},
|
||||
"/tinyfish/{endpoint}": {
|
||||
"get": {
|
||||
"description": "Pass-through for the TinyFish Agent API (goal-based web automation).\n\nForwarded endpoints:\n- POST /v1/automation/run \u2014 run to completion (blocking)\n- POST /v1/automation/run-async \u2014 submit a run, poll GET /v1/runs/{id} for the result\n- POST /v1/automation/run-sse \u2014 run with SSE progress events\n- GET /v1/runs \u2014 list runs\n- GET /v1/runs/{id} \u2014 run status / result\n- POST /v1/runs/{id}/cancel \u2014 cancel a run\n\nEvery other Agent API endpoint (vault, wallet, browser profiles) returns 403: all\nproxy callers share one upstream key.\n\nCredential lookup order:\n1. passthrough_endpoint_router (config.yaml deployments with use_in_pass_through)\n2. TINYFISH_API_KEY environment variable\n\n[Docs](https://docs.litellm.ai/docs/pass_through/tinyfish)",
|
||||
"operationId": "tinyfish_proxy_route_tinyfish__endpoint__get",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
"name": "endpoint",
|
||||
"required": true,
|
||||
"schema": {
|
||||
"title": "Endpoint",
|
||||
"type": "string"
|
||||
}
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
},
|
||||
"422": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/HTTPValidationError"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Validation Error"
|
||||
}
|
||||
},
|
||||
"security": [
|
||||
{
|
||||
"APIKeyHeader": []
|
||||
}
|
||||
],
|
||||
"summary": "Tinyfish Proxy Route",
|
||||
"tags": [
|
||||
"llm_passthrough"
|
||||
]
|
||||
},
|
||||
"post": {
|
||||
"description": "Pass-through for the TinyFish Agent API (goal-based web automation).\n\nForwarded endpoints:\n- POST /v1/automation/run \u2014 run to completion (blocking)\n- POST /v1/automation/run-async \u2014 submit a run, poll GET /v1/runs/{id} for the result\n- POST /v1/automation/run-sse \u2014 run with SSE progress events\n- GET /v1/runs \u2014 list runs\n- GET /v1/runs/{id} \u2014 run status / result\n- POST /v1/runs/{id}/cancel \u2014 cancel a run\n\nEvery other Agent API endpoint (vault, wallet, browser profiles) returns 403: all\nproxy callers share one upstream key.\n\nCredential lookup order:\n1. passthrough_endpoint_router (config.yaml deployments with use_in_pass_through)\n2. TINYFISH_API_KEY environment variable\n\n[Docs](https://docs.litellm.ai/docs/pass_through/tinyfish)",
|
||||
"operationId": "tinyfish_proxy_route_tinyfish__endpoint__post",
|
||||
"parameters": [
|
||||
{
|
||||
"in": "path",
|
||||
"name": "endpoint",
|
||||
"required": true,
|
||||
"schema": {
|
||||
"title": "Endpoint",
|
||||
"type": "string"
|
||||
}
|
||||
}
|
||||
],
|
||||
"responses": {
|
||||
"200": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {}
|
||||
}
|
||||
},
|
||||
"description": "Successful Response"
|
||||
},
|
||||
"422": {
|
||||
"content": {
|
||||
"application/json": {
|
||||
"schema": {
|
||||
"$ref": "#/components/schemas/HTTPValidationError"
|
||||
}
|
||||
}
|
||||
},
|
||||
"description": "Validation Error"
|
||||
}
|
||||
},
|
||||
"security": [
|
||||
{
|
||||
"APIKeyHeader": []
|
||||
}
|
||||
],
|
||||
"summary": "Tinyfish Proxy Route",
|
||||
"tags": [
|
||||
"llm_passthrough"
|
||||
]
|
||||
}
|
||||
},
|
||||
"/vertex_ai/discovery/{endpoint}": {
|
||||
"delete": {
|
||||
"description": "Call any vertex discovery endpoint using the proxy.\n\nJust use `{PROXY_BASE_URL}/vertex_ai/discovery/{endpoint:path}`\n\nTarget url: `https://discoveryengine.googleapis.com`",
|
||||
|
|
|
|||
|
|
@ -472,6 +472,7 @@ class LiteLLMRoutes(enum.Enum):
|
|||
"/openai_passthrough",
|
||||
"/assemblyai",
|
||||
"/eu.assemblyai",
|
||||
"/tinyfish",
|
||||
"/vllm",
|
||||
"/mistral",
|
||||
"/milvus",
|
||||
|
|
|
|||
|
|
@ -79,6 +79,10 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
|||
LITELLM_PASS_THROUGH_CUSTOM_BODY_STATE_KEY,
|
||||
LITELLM_PASS_THROUGH_RAW_BODY_STATE_KEY,
|
||||
)
|
||||
from litellm.types.passthrough_endpoints.tinyfish import (
|
||||
TINYFISH_AUTHENTICATED_RUN_FIELDS,
|
||||
is_allowed_tinyfish_endpoint,
|
||||
)
|
||||
from litellm.types.passthrough_endpoints.vertex_ai import VertexPassThroughCredentials
|
||||
from litellm.types.router import LiteLLMParamsTypedDict
|
||||
from litellm.types.utils import LlmProviders
|
||||
|
|
@ -2713,6 +2717,106 @@ async def cursor_proxy_route(
|
|||
return received_value
|
||||
|
||||
|
||||
async def _tinyfish_blocked_body_fields(request: Request) -> tuple[str, ...]:
|
||||
raw_body: Final = await request.body()
|
||||
if not raw_body:
|
||||
return ()
|
||||
try:
|
||||
parsed: Final[object] = json.loads(raw_body) # any-ok: json.loads -> Any
|
||||
except (json.JSONDecodeError, UnicodeDecodeError):
|
||||
return ()
|
||||
if not isinstance(parsed, dict):
|
||||
return ()
|
||||
return tuple(sorted(key for key in parsed if key in TINYFISH_AUTHENTICATED_RUN_FIELDS))
|
||||
|
||||
|
||||
@router.api_route(
|
||||
"/tinyfish/{endpoint:path}",
|
||||
methods=["GET", "POST"], # mutable-ok: fastapi api_route requires List[str]
|
||||
tags=["TinyFish Pass-through", "pass-through"], # mutable-ok: fastapi api_route requires a list
|
||||
)
|
||||
async def tinyfish_proxy_route(
|
||||
endpoint: str,
|
||||
request: Request,
|
||||
fastapi_response: Response,
|
||||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""
|
||||
Pass-through for the TinyFish Agent API (goal-based web automation).
|
||||
|
||||
Forwarded endpoints:
|
||||
- POST /v1/automation/run — run to completion (blocking)
|
||||
- POST /v1/automation/run-async — submit a run, poll GET /v1/runs/{id} for the result
|
||||
- POST /v1/automation/run-sse — run with SSE progress events
|
||||
- GET /v1/runs — list runs
|
||||
- GET /v1/runs/{id} — run status / result
|
||||
- POST /v1/runs/{id}/cancel — cancel a run
|
||||
|
||||
Every other Agent API endpoint (vault, wallet, browser profiles) returns 403: all
|
||||
proxy callers share one upstream key.
|
||||
|
||||
Credential lookup order:
|
||||
1. passthrough_endpoint_router (config.yaml deployments with use_in_pass_through)
|
||||
2. TINYFISH_API_KEY environment variable
|
||||
|
||||
[Docs](https://docs.litellm.ai/docs/pass_through/tinyfish)
|
||||
"""
|
||||
from .llm_provider_handlers.tinyfish_passthrough_logging_handler import (
|
||||
resolve_tinyfish_agent_api_base,
|
||||
)
|
||||
|
||||
raw_endpoint_path: Final = httpx.URL(endpoint).path
|
||||
encoded_endpoint: Final = raw_endpoint_path if raw_endpoint_path.startswith("/") else f"/{raw_endpoint_path}"
|
||||
|
||||
if not is_allowed_tinyfish_endpoint(request.method, encoded_endpoint):
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"{request.method} {encoded_endpoint} is not an allowed TinyFish Agent passthrough endpoint. "
|
||||
"Allowed: POST /v1/automation/run, POST /v1/automation/run-async, POST /v1/automation/run-sse, "
|
||||
"GET /v1/runs, GET /v1/runs/{id}, POST /v1/runs/{id}/cancel.",
|
||||
)
|
||||
|
||||
if request.method == "POST" and encoded_endpoint.startswith("/v1/automation/"):
|
||||
blocked_fields: Final = await _tinyfish_blocked_body_fields(request)
|
||||
if blocked_fields and str_to_bool(os.getenv("TINYFISH_ALLOW_AUTHENTICATED_RUNS")) is not True:
|
||||
raise HTTPException(
|
||||
status_code=403,
|
||||
detail=f"Request fields [{', '.join(blocked_fields)}] run with the shared TinyFish account's saved "
|
||||
"credentials and are disabled on this proxy. Ask the proxy admin to set "
|
||||
"TINYFISH_ALLOW_AUTHENTICATED_RUNS=true to allow them.",
|
||||
)
|
||||
|
||||
tinyfish_api_key: Final = passthrough_endpoint_router.get_credentials(
|
||||
custom_llm_provider="tinyfish",
|
||||
region_name=None,
|
||||
)
|
||||
if tinyfish_api_key is None:
|
||||
raise HTTPException(
|
||||
status_code=401,
|
||||
detail="TinyFish API key not found. Set the TINYFISH_API_KEY environment variable or add a "
|
||||
"deployment with use_in_pass_through: true.",
|
||||
)
|
||||
|
||||
base_url: Final = httpx.URL(resolve_tinyfish_agent_api_base())
|
||||
updated_url: Final = base_url.copy_with(
|
||||
path=HttpPassThroughEndpointHelpers.join_base_and_endpoint_path(base_url, encoded_endpoint)
|
||||
)
|
||||
|
||||
endpoint_func: Final = create_pass_through_route(
|
||||
endpoint=endpoint,
|
||||
target=str(updated_url),
|
||||
custom_headers=MappingProxyType({"X-API-Key": tinyfish_api_key}),
|
||||
custom_llm_provider="tinyfish",
|
||||
)
|
||||
received_value: Final = await endpoint_func(
|
||||
request,
|
||||
fastapi_response,
|
||||
user_api_key_dict,
|
||||
)
|
||||
|
||||
return received_value
|
||||
|
||||
|
||||
VERTEX_LIVE_UNCONFIGURED_CLOSE_REASON: Final = (
|
||||
"Vertex AI auth failed: set a use_in_pass_through vertex model, default_vertex_config, or DEFAULT_VERTEXAI_* env"
|
||||
)
|
||||
|
|
|
|||
|
|
@ -0,0 +1,337 @@
|
|||
import asyncio
|
||||
import json
|
||||
import os
|
||||
import time
|
||||
import urllib.parse
|
||||
from collections.abc import Mapping, Sequence
|
||||
from datetime import datetime
|
||||
from types import MappingProxyType
|
||||
from typing import Final, cast
|
||||
from urllib.parse import urlparse
|
||||
|
||||
import httpx
|
||||
from pydantic import TypeAdapter, ValidationError
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLoggingObj
|
||||
from litellm.litellm_core_utils.litellm_logging import (
|
||||
get_standard_logging_object_payload,
|
||||
)
|
||||
from litellm.llms.custom_httpx.http_handler import (
|
||||
AsyncHTTPHandler,
|
||||
get_async_httpx_client,
|
||||
)
|
||||
from litellm.proxy._types import PassThroughEndpointLoggingTypedDict
|
||||
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_POLLING_SECONDS,
|
||||
TINYFISH_MODEL_NAME,
|
||||
TINYFISH_POLLING_INTERVAL_SECONDS,
|
||||
TINYFISH_TERMINAL_RUN_STATUSES,
|
||||
TinyfishRun,
|
||||
)
|
||||
from litellm.types.utils import StandardPassThroughResponseObject
|
||||
|
||||
_RUN_ADAPTER: Final = TypeAdapter(TinyfishRun)
|
||||
|
||||
_EMPTY_KWARGS: Final[Mapping[str, object]] = MappingProxyType({})
|
||||
|
||||
# asyncio tasks are weakly referenced by the loop; hold them until done or they can vanish mid-poll
|
||||
_BACKGROUND_BILLING_TASKS: Final[set["asyncio.Task[None]"]] = set() # mutable-ok: task registry
|
||||
|
||||
|
||||
def resolve_tinyfish_agent_api_base() -> str:
|
||||
return (os.getenv("TINYFISH_AGENT_API_BASE") or TINYFISH_AGENT_DEFAULT_API_BASE).rstrip("/")
|
||||
|
||||
|
||||
def resolve_tinyfish_cost_per_step() -> float:
|
||||
raw: Final = os.getenv("TINYFISH_COST_PER_STEP")
|
||||
if raw is None:
|
||||
return TINYFISH_DEFAULT_COST_PER_STEP
|
||||
try:
|
||||
return float(raw)
|
||||
except ValueError:
|
||||
verbose_proxy_logger.warning(
|
||||
"TINYFISH_COST_PER_STEP=%r is not a number; using the default rate %s",
|
||||
raw,
|
||||
TINYFISH_DEFAULT_COST_PER_STEP,
|
||||
)
|
||||
return TINYFISH_DEFAULT_COST_PER_STEP
|
||||
|
||||
|
||||
def is_tinyfish_agent_url(url: str) -> bool:
|
||||
hostname: Final = urlparse(url).hostname
|
||||
return hostname is not None and hostname == urlparse(resolve_tinyfish_agent_api_base()).hostname
|
||||
|
||||
|
||||
def _parse_run(payload: object) -> TinyfishRun | None:
|
||||
try:
|
||||
return _RUN_ADAPTER.validate_python(payload)
|
||||
except ValidationError as e:
|
||||
verbose_proxy_logger.warning("TinyFish passthrough: unexpected run object shape: %s", e)
|
||||
return None
|
||||
|
||||
|
||||
def _run_cost(run: TinyfishRun | None) -> float | None:
|
||||
if run is None:
|
||||
return None
|
||||
num_of_steps: Final = run.get("num_of_steps")
|
||||
if num_of_steps is None:
|
||||
return None
|
||||
return num_of_steps * resolve_tinyfish_cost_per_step()
|
||||
|
||||
|
||||
class TinyFishPassthroughLoggingHandler:
|
||||
@staticmethod
|
||||
def _should_log_request(request_method: str, url_route: str) -> bool:
|
||||
"""Only run submissions are billed; GET /v1/runs* polling and cancels never write spend rows."""
|
||||
return request_method == "POST" and "/v1/automation/" in urlparse(url_route).path
|
||||
|
||||
@staticmethod
|
||||
def is_run_async_route(url_route: str) -> bool:
|
||||
return urlparse(url_route).path.endswith("/v1/automation/run-async")
|
||||
|
||||
@staticmethod
|
||||
def tinyfish_passthrough_handler(
|
||||
httpx_response: httpx.Response,
|
||||
response_body: Mapping[str, object] | None,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
url_route: str,
|
||||
result: str,
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
cache_hit: bool,
|
||||
request_body: Mapping[str, object],
|
||||
**kwargs: object, # kwargs-ok: the passthrough logging dispatch forwards shared logging kwargs to every handler
|
||||
) -> PassThroughEndpointLoggingTypedDict:
|
||||
"""Bill a blocking POST /v1/automation/run: the response is the terminal run object."""
|
||||
try:
|
||||
run: Final = _parse_run(response_body) if response_body is not None else None
|
||||
handler_payload: Final = TinyFishPassthroughLoggingHandler._build_logging_payload(
|
||||
run=run,
|
||||
logging_obj=logging_obj,
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error in TinyFish passthrough logging handler: %s", e)
|
||||
fallback_payload: Final[PassThroughEndpointLoggingTypedDict] = {
|
||||
"result": StandardPassThroughResponseObject(response=result),
|
||||
"kwargs": kwargs,
|
||||
}
|
||||
return fallback_payload
|
||||
return handler_payload
|
||||
|
||||
@staticmethod
|
||||
def start_async_run_billing(
|
||||
response_body: Mapping[str, object] | None,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
result: str,
|
||||
start_time: datetime,
|
||||
cache_hit: bool,
|
||||
**kwargs: object, # kwargs-ok: shared logging kwargs, replayed into _handle_logging when the run finishes
|
||||
) -> None:
|
||||
"""Bill POST /v1/automation/run-async once, when the polled run turns terminal."""
|
||||
submitted: Final = _parse_run(response_body) if response_body is not None else None
|
||||
run_id: Final = submitted.get("run_id") if submitted is not None else None
|
||||
if not run_id:
|
||||
verbose_proxy_logger.warning(
|
||||
"TinyFish passthrough: run-async response carried no run_id; logging the request without cost"
|
||||
)
|
||||
task: Final = asyncio.create_task(
|
||||
TinyFishPassthroughLoggingHandler._poll_and_log(
|
||||
run_id=run_id,
|
||||
logging_obj=logging_obj,
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
cache_hit=cache_hit,
|
||||
kwargs=kwargs,
|
||||
)
|
||||
)
|
||||
_BACKGROUND_BILLING_TASKS.add(task)
|
||||
task.add_done_callback(_BACKGROUND_BILLING_TASKS.discard)
|
||||
|
||||
@staticmethod
|
||||
async def _poll_and_log(
|
||||
run_id: str | None,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
result: str,
|
||||
start_time: datetime,
|
||||
cache_hit: bool,
|
||||
kwargs: Mapping[str, object],
|
||||
client: AsyncHTTPHandler | None = None,
|
||||
) -> None:
|
||||
from ..pass_through_endpoints import pass_through_endpoint_logging
|
||||
|
||||
try:
|
||||
run: Final = (
|
||||
await TinyFishPassthroughLoggingHandler._poll_until_terminal(run_id, client) if run_id else None
|
||||
)
|
||||
payload: Final = TinyFishPassthroughLoggingHandler._build_logging_payload(
|
||||
run=run,
|
||||
logging_obj=logging_obj,
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=datetime.now(),
|
||||
kwargs=kwargs,
|
||||
)
|
||||
logged_result: Final = payload["result"] or StandardPassThroughResponseObject(response=result)
|
||||
logging_kwargs: Final = cast("Mapping[str, object]", payload["kwargs"])
|
||||
await pass_through_endpoint_logging._handle_logging( # pyright: ignore[reportPrivateUsage] # shared passthrough logging dispatcher, same access as the assemblyai handler
|
||||
logging_obj=logging_obj,
|
||||
standard_logging_response_object=logged_result,
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=datetime.now(),
|
||||
cache_hit=cache_hit,
|
||||
**logging_kwargs,
|
||||
)
|
||||
except Exception as e:
|
||||
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:
|
||||
deadline: Final = time.monotonic() + TINYFISH_MAX_POLLING_SECONDS
|
||||
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)
|
||||
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
|
||||
|
||||
@staticmethod
|
||||
async def _fetch_run(run_id: str, client: AsyncHTTPHandler | None = None) -> TinyfishRun | None:
|
||||
from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
||||
passthrough_endpoint_router,
|
||||
)
|
||||
|
||||
api_key: Final = passthrough_endpoint_router.get_credentials(custom_llm_provider="tinyfish", region_name=None)
|
||||
if api_key is None:
|
||||
verbose_proxy_logger.warning("TinyFish passthrough: no API key available to poll run %s", run_id)
|
||||
return None
|
||||
if any(c in run_id for c in ("/", "\\", "#", "?")) or ".." in run_id:
|
||||
verbose_proxy_logger.warning("TinyFish passthrough: invalid run_id %r", run_id)
|
||||
return None
|
||||
safe_run_id: Final = urllib.parse.quote(run_id, safe="")
|
||||
resolved_client: Final = client or get_async_httpx_client(
|
||||
llm_provider=httpxSpecialProvider.PassThroughEndpoint,
|
||||
params={"timeout": 30.0}, # mutable-ok: get_async_httpx_client takes a plain dict of client params
|
||||
)
|
||||
try:
|
||||
# screenshots=none keeps the poll payload small (no per-step screenshot URLs needed)
|
||||
response: Final = await resolved_client.get(
|
||||
f"{resolve_tinyfish_agent_api_base()}/v1/runs/{safe_run_id}?screenshots=none",
|
||||
headers={"X-API-Key": api_key}, # mutable-ok: httpx headers= takes a plain dict
|
||||
)
|
||||
if not (200 <= response.status_code < 300):
|
||||
verbose_proxy_logger.warning(
|
||||
"TinyFish passthrough: GET /v1/runs/%s returned %s", safe_run_id, response.status_code
|
||||
)
|
||||
return None
|
||||
payload: Final[object] = response.json() # any-ok: httpx Response.json() -> Any
|
||||
return _parse_run(payload)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.warning("[Non blocking logging error] TinyFish run fetch failed: %s", e)
|
||||
return None
|
||||
|
||||
@staticmethod
|
||||
async def _handle_logging_tinyfish_collected_chunks(
|
||||
litellm_logging_obj: LiteLLMLoggingObj,
|
||||
url_route: str,
|
||||
start_time: datetime,
|
||||
all_chunks: Sequence[str],
|
||||
end_time: datetime,
|
||||
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}."""
|
||||
try:
|
||||
run_id: Final = _run_id_from_sse_chunks(all_chunks)
|
||||
if run_id is None:
|
||||
verbose_proxy_logger.warning(
|
||||
"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
|
||||
payload: Final = TinyFishPassthroughLoggingHandler._build_logging_payload(
|
||||
run=run,
|
||||
logging_obj=litellm_logging_obj,
|
||||
result="",
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
kwargs=_EMPTY_KWARGS,
|
||||
)
|
||||
except Exception as e:
|
||||
verbose_proxy_logger.exception("Error in TinyFish SSE passthrough logging handler: %s", e)
|
||||
fallback_payload: Final[PassThroughEndpointLoggingTypedDict] = {
|
||||
"result": StandardPassThroughResponseObject(response=""),
|
||||
"kwargs": {},
|
||||
}
|
||||
return fallback_payload
|
||||
return payload
|
||||
|
||||
@staticmethod
|
||||
def _build_logging_payload(
|
||||
run: TinyfishRun | None,
|
||||
logging_obj: LiteLLMLoggingObj,
|
||||
result: str,
|
||||
start_time: datetime,
|
||||
end_time: datetime,
|
||||
kwargs: Mapping[str, object],
|
||||
) -> PassThroughEndpointLoggingTypedDict:
|
||||
response_cost: Final = _run_cost(run)
|
||||
updated_kwargs: Final = { # mutable-ok: the logging pipeline requires a plain kwargs dict
|
||||
**kwargs,
|
||||
"model": TINYFISH_MODEL_NAME,
|
||||
"custom_llm_provider": "tinyfish",
|
||||
"response_cost": response_cost,
|
||||
}
|
||||
logging_obj.model_call_details.update(
|
||||
model=TINYFISH_MODEL_NAME,
|
||||
custom_llm_provider="tinyfish",
|
||||
response_cost=response_cost,
|
||||
)
|
||||
|
||||
logged_response: Final = StandardPassThroughResponseObject(
|
||||
response=json.dumps(run) if run is not None else result
|
||||
)
|
||||
standard_logging_object: Final = get_standard_logging_object_payload(
|
||||
kwargs=updated_kwargs,
|
||||
init_response_obj=logged_response,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
logging_obj=logging_obj,
|
||||
status="success",
|
||||
)
|
||||
handler_payload: Final[PassThroughEndpointLoggingTypedDict] = {
|
||||
"result": logged_response,
|
||||
"kwargs": {**updated_kwargs, "standard_logging_object": standard_logging_object},
|
||||
}
|
||||
return handler_payload
|
||||
|
||||
|
||||
def _run_id_from_sse_chunks(all_chunks: Sequence[str]) -> str | None:
|
||||
for line in all_chunks:
|
||||
if not line.startswith("data:"):
|
||||
continue
|
||||
try:
|
||||
event_payload: object = json.loads(line[5:].strip()) # any-ok: json.loads -> Any
|
||||
except json.JSONDecodeError:
|
||||
continue
|
||||
event = _parse_run(event_payload)
|
||||
if event is None:
|
||||
continue
|
||||
run_id = event.get("run_id")
|
||||
if run_id:
|
||||
return run_id
|
||||
return None
|
||||
|
|
@ -103,6 +103,9 @@ from litellm.types.passthrough_endpoints.pass_through_endpoints import (
|
|||
)
|
||||
from litellm.types.utils import TRUSTED_CALLBACK_VARS_FIELD, Usage
|
||||
|
||||
from .llm_provider_handlers.tinyfish_passthrough_logging_handler import (
|
||||
is_tinyfish_agent_url,
|
||||
)
|
||||
from .streaming_handler import PassThroughStreamingHandler
|
||||
from .success_handler import PassThroughEndpointLogging
|
||||
from .upstream_usage_headers import (
|
||||
|
|
@ -374,6 +377,8 @@ class HttpPassThroughEndpointHelpers(BasePassthroughUtils):
|
|||
or (parsed_url.hostname and "openai.com" in parsed_url.hostname)
|
||||
):
|
||||
return EndpointType.OPENAI
|
||||
elif is_tinyfish_agent_url(url):
|
||||
return EndpointType.TINYFISH
|
||||
return EndpointType.GENERIC
|
||||
|
||||
@staticmethod
|
||||
|
|
|
|||
|
|
@ -25,6 +25,9 @@ from .llm_provider_handlers.gemini_passthrough_logging_handler import (
|
|||
from .llm_provider_handlers.openai_passthrough_logging_handler import (
|
||||
OpenAIPassthroughLoggingHandler,
|
||||
)
|
||||
from .llm_provider_handlers.tinyfish_passthrough_logging_handler import (
|
||||
TinyFishPassthroughLoggingHandler,
|
||||
)
|
||||
from .llm_provider_handlers.vertex_passthrough_logging_handler import (
|
||||
VertexPassthroughLoggingHandler,
|
||||
)
|
||||
|
|
@ -271,6 +274,27 @@ class PassThroughStreamingHandler:
|
|||
- OpenAI
|
||||
"""
|
||||
try:
|
||||
# TinyFish is dispatched before the sync builder: its SSE events carry no
|
||||
# num_of_steps, so pricing needs an async GET /v1/runs/{id} after the stream.
|
||||
if endpoint_type == EndpointType.TINYFISH:
|
||||
tinyfish_payload: Final = (
|
||||
await TinyFishPassthroughLoggingHandler._handle_logging_tinyfish_collected_chunks(
|
||||
litellm_logging_obj=litellm_logging_obj,
|
||||
url_route=url_route,
|
||||
start_time=start_time,
|
||||
all_chunks=PassThroughStreamingHandler._convert_raw_bytes_to_str_lines(raw_bytes),
|
||||
end_time=end_time,
|
||||
)
|
||||
)
|
||||
await litellm_logging_obj.dispatch_success_handlers(
|
||||
result=tinyfish_payload["result"],
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
cache_hit=litellm_logging_obj.model_call_details.get("cache_hit") is True,
|
||||
prefer_async_handlers=True,
|
||||
**tinyfish_payload["kwargs"],
|
||||
)
|
||||
return
|
||||
(
|
||||
standard_logging_response_object,
|
||||
kwargs,
|
||||
|
|
|
|||
|
|
@ -27,6 +27,10 @@ from .llm_provider_handlers.cursor_passthrough_logging_handler import (
|
|||
from .llm_provider_handlers.gemini_passthrough_logging_handler import (
|
||||
GeminiPassthroughLoggingHandler,
|
||||
)
|
||||
from .llm_provider_handlers.tinyfish_passthrough_logging_handler import (
|
||||
TinyFishPassthroughLoggingHandler,
|
||||
is_tinyfish_agent_url,
|
||||
)
|
||||
from .llm_provider_handlers.vertex_passthrough_logging_handler import (
|
||||
VertexPassthroughLoggingHandler,
|
||||
)
|
||||
|
|
@ -256,6 +260,21 @@ class PassThroughEndpointLogging:
|
|||
)
|
||||
standard_logging_response_object = comprehend_medical_handler_result["result"] # rebind-ok: elif-chain
|
||||
kwargs = comprehend_medical_handler_result["kwargs"] # rebind-ok: elif-chain contract
|
||||
elif self.is_tinyfish_route(url_route, custom_llm_provider):
|
||||
tinyfish_handler_result: Final = TinyFishPassthroughLoggingHandler.tinyfish_passthrough_handler(
|
||||
httpx_response=httpx_response,
|
||||
response_body=response_body if isinstance(response_body, dict) else None,
|
||||
logging_obj=logging_obj,
|
||||
url_route=url_route,
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
end_time=end_time,
|
||||
cache_hit=cache_hit,
|
||||
request_body=request_body,
|
||||
**kwargs,
|
||||
)
|
||||
standard_logging_response_object = tinyfish_handler_result["result"] # rebind-ok: elif-chain
|
||||
kwargs = tinyfish_handler_result["kwargs"] # rebind-ok: elif-chain contract
|
||||
elif self.is_vertex_ai_live_route(url_route):
|
||||
from .llm_provider_handlers.vertex_ai_live_passthrough_logging_handler import (
|
||||
VertexAILivePassthroughLoggingHandler,
|
||||
|
|
@ -300,6 +319,21 @@ class PassThroughEndpointLogging:
|
|||
):
|
||||
standard_logging_response_object: PassThroughEndpointLoggingResultValues | None = None
|
||||
logging_obj.model_call_details["passthrough_logging_payload"] = passthrough_logging_payload
|
||||
if self.is_tinyfish_route(url_route, custom_llm_provider):
|
||||
# GET /v1/runs* polling and cancels never write spend rows; run-async bills once,
|
||||
# from a background poller that re-enters _handle_logging at run completion.
|
||||
if not TinyFishPassthroughLoggingHandler._should_log_request(httpx_response.request.method, url_route):
|
||||
return
|
||||
if TinyFishPassthroughLoggingHandler.is_run_async_route(url_route):
|
||||
TinyFishPassthroughLoggingHandler.start_async_run_billing(
|
||||
response_body=response_body if isinstance(response_body, dict) else None,
|
||||
logging_obj=logging_obj,
|
||||
result=result,
|
||||
start_time=start_time,
|
||||
cache_hit=cache_hit,
|
||||
**kwargs,
|
||||
)
|
||||
return
|
||||
if self.is_assemblyai_route(url_route):
|
||||
if AssemblyAIPassthroughLoggingHandler._should_log_request(httpx_response.request.method) is not True:
|
||||
return
|
||||
|
|
@ -387,6 +421,9 @@ class PassThroughEndpointLogging:
|
|||
def is_comprehend_medical_route(self, custom_llm_provider: str | None) -> bool:
|
||||
return custom_llm_provider == "comprehendmedical"
|
||||
|
||||
def is_tinyfish_route(self, url_route: str, custom_llm_provider: str | None) -> bool:
|
||||
return custom_llm_provider == "tinyfish" or is_tinyfish_agent_url(url_route)
|
||||
|
||||
def is_langfuse_route(self, url_route: str):
|
||||
parsed_url: Final = urlparse(url_route)
|
||||
for route in self.TRACKED_LANGFUSE_ROUTES:
|
||||
|
|
|
|||
|
|
@ -25,6 +25,7 @@ class EndpointType(str, Enum):
|
|||
GEMINI = "gemini"
|
||||
ANTHROPIC = "anthropic"
|
||||
OPENAI = "openai"
|
||||
TINYFISH = "tinyfish"
|
||||
GENERIC = "generic"
|
||||
|
||||
|
||||
|
|
|
|||
55
litellm/types/passthrough_endpoints/tinyfish.py
Normal file
55
litellm/types/passthrough_endpoints/tinyfish.py
Normal file
|
|
@ -0,0 +1,55 @@
|
|||
from typing import Final
|
||||
|
||||
from typing_extensions import ReadOnly, TypedDict
|
||||
|
||||
TINYFISH_AGENT_DEFAULT_API_BASE: Final = "https://agent.tinyfish.ai"
|
||||
TINYFISH_AGENT_DOCS_URL: Final = "https://docs.tinyfish.ai/agent-api"
|
||||
# TinyFish's published Agent API rate (USD per run step); override with env TINYFISH_COST_PER_STEP
|
||||
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_TERMINAL_RUN_STATUSES: Final = frozenset({"COMPLETED", "FAILED", "CANCELLED"})
|
||||
|
||||
# Fields that run with the TinyFish account's saved logins/vault; all proxy callers share
|
||||
# one upstream key, so these are rejected unless TINYFISH_ALLOW_AUTHENTICATED_RUNS=true.
|
||||
TINYFISH_AUTHENTICATED_RUN_FIELDS: Final = frozenset({"use_profile", "profile_id", "use_vault", "credential_item_ids"})
|
||||
|
||||
_RUN_SUBMIT_PATHS: Final = frozenset(
|
||||
{("v1", "automation", "run"), ("v1", "automation", "run-async"), ("v1", "automation", "run-sse")}
|
||||
)
|
||||
|
||||
|
||||
class TinyfishRunError(TypedDict, total=False):
|
||||
code: ReadOnly[str | None]
|
||||
message: ReadOnly[str | None]
|
||||
category: ReadOnly[str | None]
|
||||
retry_after: ReadOnly[float | None]
|
||||
help_url: ReadOnly[str | None]
|
||||
|
||||
|
||||
class TinyfishRun(TypedDict, total=False):
|
||||
"""Run objects are null-heavy until terminal, so every field must tolerate None."""
|
||||
|
||||
run_id: ReadOnly[str | None]
|
||||
status: ReadOnly[str | None]
|
||||
num_of_steps: ReadOnly[int | None]
|
||||
result: ReadOnly[object]
|
||||
error: ReadOnly[TinyfishRunError | None]
|
||||
type: ReadOnly[str | None]
|
||||
|
||||
|
||||
def is_allowed_tinyfish_endpoint(method: str, path: str) -> bool:
|
||||
"""agent.tinyfish.ai also serves vault/wallet/browser-profile management under the
|
||||
same key, so only the run endpoints may be forwarded."""
|
||||
segments: Final = tuple(part for part in path.split("/") if part)
|
||||
if any(segment in (".", "..") for segment in segments):
|
||||
return False
|
||||
if method == "POST" and segments in _RUN_SUBMIT_PATHS:
|
||||
return True
|
||||
if method == "GET" and segments == ("v1", "runs"):
|
||||
return True
|
||||
if method == "GET" and len(segments) == 3 and segments[:2] == ("v1", "runs"):
|
||||
return True
|
||||
return method == "POST" and len(segments) == 4 and segments[:2] == ("v1", "runs") and segments[3] == "cancel"
|
||||
|
|
@ -0,0 +1,280 @@
|
|||
import asyncio
|
||||
import json
|
||||
from datetime import datetime
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import httpx
|
||||
import pytest
|
||||
|
||||
from litellm.proxy.pass_through_endpoints.llm_provider_handlers.tinyfish_passthrough_logging_handler import (
|
||||
TinyFishPassthroughLoggingHandler,
|
||||
is_tinyfish_agent_url,
|
||||
resolve_tinyfish_cost_per_step,
|
||||
)
|
||||
from litellm.proxy.pass_through_endpoints.success_handler import (
|
||||
PassThroughEndpointLogging,
|
||||
)
|
||||
from litellm.types.passthrough_endpoints.tinyfish import is_allowed_tinyfish_endpoint
|
||||
|
||||
RUN_URL = "https://agent.tinyfish.ai/v1/automation/run"
|
||||
RUN_ASYNC_URL = "https://agent.tinyfish.ai/v1/automation/run-async"
|
||||
|
||||
|
||||
def _make_logging_obj() -> MagicMock:
|
||||
logging_obj = MagicMock()
|
||||
logging_obj.litellm_call_id = "test-call-id"
|
||||
logging_obj.model_call_details = {}
|
||||
return logging_obj
|
||||
|
||||
|
||||
def _make_response(method: str, url: str, body: dict) -> httpx.Response:
|
||||
request = httpx.Request(method, url)
|
||||
return httpx.Response(200, request=request, text=json.dumps(body))
|
||||
|
||||
|
||||
class _FakeClient:
|
||||
def __init__(self, payloads: list[dict], 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))
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def tinyfish_env(monkeypatch):
|
||||
monkeypatch.setenv("TINYFISH_API_KEY", "sk-tf-test")
|
||||
monkeypatch.delenv("TINYFISH_COST_PER_STEP", raising=False)
|
||||
monkeypatch.delenv("TINYFISH_AGENT_API_BASE", raising=False)
|
||||
|
||||
|
||||
class TestCostResolution:
|
||||
def test_default_rate(self, tinyfish_env):
|
||||
assert resolve_tinyfish_cost_per_step() == pytest.approx(0.016)
|
||||
|
||||
def test_env_override(self, tinyfish_env, monkeypatch):
|
||||
monkeypatch.setenv("TINYFISH_COST_PER_STEP", "0.02")
|
||||
assert resolve_tinyfish_cost_per_step() == pytest.approx(0.02)
|
||||
|
||||
def test_invalid_env_falls_back_to_default(self, tinyfish_env, monkeypatch):
|
||||
monkeypatch.setenv("TINYFISH_COST_PER_STEP", "free")
|
||||
assert resolve_tinyfish_cost_per_step() == pytest.approx(0.016)
|
||||
|
||||
|
||||
class TestBillingGate:
|
||||
@pytest.mark.parametrize(
|
||||
"method,url,expected",
|
||||
[
|
||||
("POST", RUN_URL, True),
|
||||
("POST", RUN_ASYNC_URL, True),
|
||||
("POST", "https://agent.tinyfish.ai/v1/automation/run-sse", True),
|
||||
("GET", "https://agent.tinyfish.ai/v1/runs", False),
|
||||
("GET", "https://agent.tinyfish.ai/v1/runs/run-123?screenshots=none", False),
|
||||
("POST", "https://agent.tinyfish.ai/v1/runs/run-123/cancel", False),
|
||||
],
|
||||
)
|
||||
def test_only_run_submissions_are_billed(self, method, url, expected):
|
||||
assert TinyFishPassthroughLoggingHandler._should_log_request(method, url) is expected
|
||||
|
||||
def test_polling_writes_no_spend_row(self, tinyfish_env):
|
||||
logging_obj = _make_logging_obj()
|
||||
logging_obj.dispatch_success_handlers = AsyncMock()
|
||||
poll_url = "https://agent.tinyfish.ai/v1/runs/run-123"
|
||||
|
||||
asyncio.run(
|
||||
PassThroughEndpointLogging().pass_through_async_success_handler(
|
||||
httpx_response=_make_response("GET", poll_url, {"run_id": "run-123", "status": "RUNNING"}),
|
||||
response_body={"run_id": "run-123", "status": "RUNNING"},
|
||||
logging_obj=logging_obj,
|
||||
url_route=poll_url,
|
||||
result="",
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
cache_hit=False,
|
||||
request_body={},
|
||||
passthrough_logging_payload={"url": poll_url},
|
||||
custom_llm_provider="tinyfish",
|
||||
)
|
||||
)
|
||||
|
||||
logging_obj.dispatch_success_handlers.assert_not_awaited()
|
||||
|
||||
|
||||
class TestBlockingRunBilling:
|
||||
def _handle(self, response_body: dict, logging_obj: MagicMock):
|
||||
return TinyFishPassthroughLoggingHandler.tinyfish_passthrough_handler(
|
||||
httpx_response=_make_response("POST", RUN_URL, response_body),
|
||||
response_body=response_body,
|
||||
logging_obj=logging_obj,
|
||||
url_route=RUN_URL,
|
||||
result=json.dumps(response_body),
|
||||
start_time=datetime.now(),
|
||||
end_time=datetime.now(),
|
||||
cache_hit=False,
|
||||
request_body={"url": "https://scrapeme.live/shop", "goal": "extract products"},
|
||||
)
|
||||
|
||||
def test_bills_steps_times_rate(self, tinyfish_env):
|
||||
logging_obj = _make_logging_obj()
|
||||
run = {"run_id": "run-1", "status": "COMPLETED", "num_of_steps": 3, "result": {"products": []}}
|
||||
|
||||
handler_result = self._handle(run, logging_obj)
|
||||
|
||||
assert handler_result["kwargs"]["model"] == "tinyfish/automation-run"
|
||||
assert handler_result["kwargs"]["custom_llm_provider"] == "tinyfish"
|
||||
assert handler_result["kwargs"]["response_cost"] == pytest.approx(0.048)
|
||||
assert "standard_logging_object" in handler_result["kwargs"]
|
||||
assert logging_obj.model_call_details["response_cost"] == pytest.approx(0.048)
|
||||
|
||||
def test_env_rate_override_applies(self, tinyfish_env, monkeypatch):
|
||||
monkeypatch.setenv("TINYFISH_COST_PER_STEP", "0.5")
|
||||
run = {"run_id": "run-1", "status": "COMPLETED", "num_of_steps": 2}
|
||||
|
||||
handler_result = self._handle(run, _make_logging_obj())
|
||||
|
||||
assert handler_result["kwargs"]["response_cost"] == pytest.approx(1.0)
|
||||
|
||||
def test_failed_run_still_bills_steps_taken(self, tinyfish_env):
|
||||
run = {"run_id": "run-1", "status": "FAILED", "num_of_steps": 2, "error": {"code": "AGENT_FAILURE"}}
|
||||
|
||||
handler_result = self._handle(run, _make_logging_obj())
|
||||
|
||||
assert handler_result["kwargs"]["response_cost"] == pytest.approx(0.032)
|
||||
|
||||
def test_null_steps_logs_without_cost(self, tinyfish_env):
|
||||
run = {"run_id": "run-1", "status": "RUNNING", "num_of_steps": None}
|
||||
|
||||
handler_result = self._handle(run, _make_logging_obj())
|
||||
|
||||
assert handler_result["kwargs"]["response_cost"] is None
|
||||
|
||||
|
||||
class TestRunAsyncBilling:
|
||||
def test_poll_and_log_bills_once_terminal(self, tinyfish_env):
|
||||
logging_obj = _make_logging_obj()
|
||||
logging_obj.dispatch_success_handlers = AsyncMock()
|
||||
fake_client = _FakeClient(
|
||||
payloads=[{"run_id": "run-9", "status": "COMPLETED", "num_of_steps": 4, "result": "ok"}]
|
||||
)
|
||||
|
||||
asyncio.run(
|
||||
TinyFishPassthroughLoggingHandler._poll_and_log(
|
||||
run_id="run-9",
|
||||
logging_obj=logging_obj,
|
||||
result="",
|
||||
start_time=datetime.now(),
|
||||
cache_hit=False,
|
||||
kwargs={},
|
||||
client=fake_client,
|
||||
)
|
||||
)
|
||||
|
||||
logging_obj.dispatch_success_handlers.assert_awaited_once()
|
||||
awaited_kwargs = logging_obj.dispatch_success_handlers.await_args.kwargs
|
||||
assert awaited_kwargs["response_cost"] == pytest.approx(0.064)
|
||||
assert awaited_kwargs["model"] == "tinyfish/automation-run"
|
||||
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_traversal_run_id_is_rejected(self, tinyfish_env):
|
||||
fake_client = _FakeClient(payloads=[{}])
|
||||
|
||||
run = asyncio.run(TinyFishPassthroughLoggingHandler._fetch_run("../vault/items", fake_client))
|
||||
|
||||
assert run is None
|
||||
assert fake_client.requested_urls == []
|
||||
|
||||
def test_upstream_error_status_returns_none(self, tinyfish_env):
|
||||
fake_client = _FakeClient(payloads=[{"error": {"code": "NOT_FOUND"}}], status_code=404)
|
||||
|
||||
run = asyncio.run(TinyFishPassthroughLoggingHandler._fetch_run("run-1", fake_client))
|
||||
|
||||
assert run is None
|
||||
|
||||
|
||||
class TestSseBilling:
|
||||
def test_collected_chunks_price_via_run_fetch(self, tinyfish_env):
|
||||
logging_obj = _make_logging_obj()
|
||||
chunks = [
|
||||
'data: {"type": "STARTED", "run_id": "run-7", "status": "RUNNING"}',
|
||||
'data: {"type": "PROGRESS", "run_id": "run-7"}',
|
||||
'data: {"type": "COMPLETE", "run_id": "run-7", "status": "COMPLETED", "result": "done"}',
|
||||
]
|
||||
fake_client = _FakeClient(
|
||||
payloads=[{"run_id": "run-7", "status": "COMPLETED", "num_of_steps": 5, "result": "done"}]
|
||||
)
|
||||
|
||||
payload = asyncio.run(
|
||||
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["kwargs"]["response_cost"] == pytest.approx(0.08)
|
||||
assert payload["kwargs"]["model"] == "tinyfish/automation-run"
|
||||
assert fake_client.requested_urls == ["https://agent.tinyfish.ai/v1/runs/run-7?screenshots=none"]
|
||||
|
||||
def test_stream_without_run_id_logs_without_cost(self, tinyfish_env):
|
||||
fake_client = _FakeClient(payloads=[{}])
|
||||
|
||||
payload = asyncio.run(
|
||||
TinyFishPassthroughLoggingHandler._handle_logging_tinyfish_collected_chunks(
|
||||
litellm_logging_obj=_make_logging_obj(),
|
||||
url_route="https://agent.tinyfish.ai/v1/automation/run-sse",
|
||||
start_time=datetime.now(),
|
||||
all_chunks=["data: not-json", ": keepalive"],
|
||||
end_time=datetime.now(),
|
||||
client=fake_client,
|
||||
)
|
||||
)
|
||||
|
||||
assert payload["kwargs"]["response_cost"] is None
|
||||
assert fake_client.requested_urls == []
|
||||
|
||||
|
||||
class TestRouteDetection:
|
||||
def test_provider_tag_claims_route(self):
|
||||
assert PassThroughEndpointLogging().is_tinyfish_route("https://example.com/x", "tinyfish")
|
||||
|
||||
def test_agent_host_claims_route(self):
|
||||
assert PassThroughEndpointLogging().is_tinyfish_route("https://agent.tinyfish.ai/v1/runs", None)
|
||||
|
||||
def test_other_providers_do_not_claim(self):
|
||||
assert not PassThroughEndpointLogging().is_tinyfish_route("https://api.openai.com/v1", "openai")
|
||||
|
||||
def test_env_base_override_claims_route(self, monkeypatch):
|
||||
monkeypatch.setenv("TINYFISH_AGENT_API_BASE", "https://agent.staging.tinyfish.ai")
|
||||
assert is_tinyfish_agent_url("https://agent.staging.tinyfish.ai/v1/runs/x")
|
||||
assert not is_tinyfish_agent_url("https://agent.tinyfish.ai/v1/runs/x")
|
||||
|
||||
|
||||
class TestEndpointAllowlist:
|
||||
@pytest.mark.parametrize(
|
||||
"method,path,expected",
|
||||
[
|
||||
("POST", "/v1/automation/run", True),
|
||||
("POST", "/v1/automation/run-async", True),
|
||||
("POST", "/v1/automation/run-sse", True),
|
||||
("GET", "/v1/runs", True),
|
||||
("GET", "/v1/runs/run-abc-123", True),
|
||||
("POST", "/v1/runs/run-abc-123/cancel", True),
|
||||
("GET", "/v1/vault/items", False),
|
||||
("GET", "/v1/wallet", False),
|
||||
("POST", "/v1/browser-profiles", False),
|
||||
("DELETE", "/v1/runs/run-abc-123", False),
|
||||
("GET", "/v1/automation/run", False),
|
||||
("POST", "/v1/runs", False),
|
||||
("GET", "/v1/runs/..", False),
|
||||
("POST", "/v1/runs/../automation/run/cancel", False),
|
||||
],
|
||||
)
|
||||
def test_allowlist(self, method, path, expected):
|
||||
assert is_allowed_tinyfish_endpoint(method, path) is expected
|
||||
|
|
@ -41,6 +41,7 @@ from litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints import (
|
|||
milvus_proxy_route,
|
||||
mistral_proxy_route,
|
||||
openai_proxy_route,
|
||||
tinyfish_proxy_route,
|
||||
vertex_discovery_proxy_route,
|
||||
vertex_proxy_route,
|
||||
vllm_proxy_route,
|
||||
|
|
@ -5492,3 +5493,142 @@ class TestAzureRelayDeploymentSegment:
|
|||
)
|
||||
|
||||
assert [call["model"] for call in captured] == ["gpt", "gpt"]
|
||||
|
||||
|
||||
class TestTinyFishProxyRoute:
|
||||
"""Tests for the TinyFish Agent pass-through route."""
|
||||
|
||||
def _mock_request(self, method: str, body: bytes = b"") -> MagicMock:
|
||||
mock_request = MagicMock(spec=Request)
|
||||
mock_request.method = method
|
||||
mock_request.query_params = {}
|
||||
mock_request.headers = {}
|
||||
mock_request.body = AsyncMock(return_value=body)
|
||||
return mock_request
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_forwards_allowed_run_endpoint_with_server_key(self):
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
|
||||
return_value="sk-tf-server",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route"
|
||||
) as mock_create_route,
|
||||
):
|
||||
mock_endpoint_func = AsyncMock(return_value={"run_id": "run-1", "status": "COMPLETED"})
|
||||
mock_create_route.return_value = mock_endpoint_func
|
||||
|
||||
result = await tinyfish_proxy_route(
|
||||
endpoint="v1/automation/run",
|
||||
request=self._mock_request("POST", b'{"url": "https://scrapeme.live/shop", "goal": "extract"}'),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
|
||||
call_args = mock_create_route.call_args[1]
|
||||
assert call_args["target"] == "https://agent.tinyfish.ai/v1/automation/run"
|
||||
assert dict(call_args["custom_headers"]) == {"X-API-Key": "sk-tf-server"}
|
||||
assert call_args["custom_llm_provider"] == "tinyfish"
|
||||
assert result == {"run_id": "run-1", "status": "COMPLETED"}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
@pytest.mark.parametrize(
|
||||
"method,endpoint",
|
||||
[
|
||||
("GET", "v1/vault/items"),
|
||||
("GET", "v1/wallet"),
|
||||
("POST", "v1/browser-profiles"),
|
||||
("GET", "v1/automation/run"),
|
||||
],
|
||||
)
|
||||
async def test_blocks_endpoints_outside_allowlist(self, method, endpoint):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await tinyfish_proxy_route(
|
||||
endpoint=endpoint,
|
||||
request=self._mock_request(method),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rejects_authenticated_run_fields_by_default(self, monkeypatch):
|
||||
monkeypatch.delenv("TINYFISH_ALLOW_AUTHENTICATED_RUNS", raising=False)
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await tinyfish_proxy_route(
|
||||
endpoint="v1/automation/run",
|
||||
request=self._mock_request("POST", b'{"url": "https://x.com", "goal": "g", "use_vault": true}'),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 403
|
||||
assert "use_vault" in exc_info.value.detail
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_env_opt_in_allows_authenticated_run_fields(self, monkeypatch):
|
||||
monkeypatch.setenv("TINYFISH_ALLOW_AUTHENTICATED_RUNS", "true")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
|
||||
return_value="sk-tf-server",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route"
|
||||
) as mock_create_route,
|
||||
):
|
||||
mock_create_route.return_value = AsyncMock(return_value={"ok": True})
|
||||
|
||||
result = await tinyfish_proxy_route(
|
||||
endpoint="v1/automation/run",
|
||||
request=self._mock_request("POST", b'{"url": "https://x.com", "goal": "g", "use_vault": true}'),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
|
||||
assert result == {"ok": True}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_raises_401_on_missing_api_key(self):
|
||||
with patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
|
||||
return_value=None,
|
||||
):
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await tinyfish_proxy_route(
|
||||
endpoint="v1/runs",
|
||||
request=self._mock_request("GET"),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 401
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_env_base_override_changes_target(self, monkeypatch):
|
||||
monkeypatch.setenv("TINYFISH_AGENT_API_BASE", "https://agent.staging.tinyfish.ai")
|
||||
|
||||
with (
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.passthrough_endpoint_router.get_credentials",
|
||||
return_value="sk-tf-server",
|
||||
),
|
||||
patch(
|
||||
"litellm.proxy.pass_through_endpoints.llm_passthrough_endpoints.create_pass_through_route"
|
||||
) as mock_create_route,
|
||||
):
|
||||
mock_create_route.return_value = AsyncMock(return_value={})
|
||||
|
||||
await tinyfish_proxy_route(
|
||||
endpoint="v1/runs/run-123",
|
||||
request=self._mock_request("GET"),
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=MagicMock(),
|
||||
)
|
||||
|
||||
assert mock_create_route.call_args[1]["target"] == "https://agent.staging.tinyfish.ai/v1/runs/run-123"
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue