mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-24 00:52:24 +00:00
fix(fal_ai): align /fal_ai queue gate with the pricer and normalise resolution type (#42505)
* fix(fal_ai): align /fal_ai queue gate with the pricer and normalise resolution type The /fal_ai gate admitted catalog keys the Fal pricer cannot price, so those jobs were forwarded and logged at 0.0 spend. The gate now reuses the pricer as its eligibility predicate. Resolution is normalised to a string before the keyed price lookup so int and str spellings bill the same. Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(fal_ai): assert passthrough pricing invariants on synthetic catalog entries Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * fix(fal_ai): price the submitted body in the /fal_ai queue gate so keyed-only entries are admitted Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> * test(fal_ai): move the queue gate repro to tests/integration Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com> --------- Co-authored-by: kerry <kerry@berri.ai> Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
3b9c9e0523
commit
5d3b31fb02
6 changed files with 204 additions and 10 deletions
|
|
@ -136,12 +136,18 @@ def _entry(key: str) -> Mapping[str, object] | None:
|
|||
return _OBJECT_MAP.validate_python(raw_entry)
|
||||
|
||||
|
||||
def _resolution_key(resolution: object) -> str | None:
|
||||
if isinstance(resolution, bool) or not isinstance(resolution, (int, str)):
|
||||
return None
|
||||
return str(resolution)
|
||||
|
||||
|
||||
def fal_ai_passthrough_cost(model: str, request_body: Mapping[str, object]) -> float | None:
|
||||
entry: Final = _entry(f"{litellm.LlmProviders.FAL_AI.value}/{model}")
|
||||
if entry is None:
|
||||
return None
|
||||
resolution: Final = request_body.get("resolution")
|
||||
keyed_cost: Final = entry.get(f"output_cost_per_image_{resolution}") if isinstance(resolution, int) else None
|
||||
resolution: Final = _resolution_key(request_body.get("resolution"))
|
||||
keyed_cost: Final = entry.get(f"output_cost_per_image_{resolution}") if resolution is not None else None
|
||||
cost: Final = keyed_cost if isinstance(keyed_cost, (int, float)) else entry.get("output_cost_per_image")
|
||||
return float(cost) if isinstance(cost, (int, float)) else None
|
||||
|
||||
|
|
|
|||
|
|
@ -53,7 +53,7 @@ from litellm.llms.deepgram.common_utils import (
|
|||
deepgram_listen_requested_model,
|
||||
deepgram_listen_websocket_target,
|
||||
)
|
||||
from litellm.llms.fal_ai.cost_calculator import fal_ai_queue_base
|
||||
from litellm.llms.fal_ai.cost_calculator import fal_ai_passthrough_cost, fal_ai_queue_base
|
||||
from litellm.llms.nvidia_nim.passthrough.transformation import nvidia_nim_model_group_in_path
|
||||
from litellm.llms.vertex_ai.vertex_llm_base import VertexBase
|
||||
from litellm.passthrough.main import AsyncPassthroughStreamingResponse
|
||||
|
|
@ -460,13 +460,11 @@ async def fal_ai_proxy_route(
|
|||
status_code=401,
|
||||
detail="FAL_AI_API_KEY is not set and no fal_ai pass-through deployment credentials are configured",
|
||||
)
|
||||
if "/requests/" not in endpoint:
|
||||
priced_model: Final = f"fal_ai/{endpoint}"
|
||||
if priced_model not in (litellm.model_cost or {}):
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"{priced_model} has no pricing entry; only priced Fal endpoints can be submitted through /fal_ai",
|
||||
)
|
||||
if "/requests/" not in endpoint and fal_ai_passthrough_cost(endpoint, await _read_request_body(request)) is None:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
detail=f"fal_ai/{endpoint} has no pricing entry for this request; only priced Fal requests can be submitted through /fal_ai",
|
||||
)
|
||||
endpoint_func: Final = create_pass_through_route(
|
||||
endpoint=endpoint,
|
||||
target=str(updated_url),
|
||||
|
|
|
|||
|
|
@ -190,6 +190,12 @@
|
|||
"tests/integration/providers/test_fal_ai_passthrough_wire.py::test_fal_queue_submit_charges_and_polls_pass_through_free": [
|
||||
"other.provider_wire.fal_ai.passthrough_queue_submit_charges_and_polls_do_not"
|
||||
],
|
||||
"tests/integration/providers/test_fal_ai_passthrough_wire.py::test_fal_queue_submit_prices_string_resolution_like_the_integer": [
|
||||
"other.provider_wire.fal_ai.passthrough_queue_submit_prices_string_resolution_like_integer"
|
||||
],
|
||||
"tests/integration/providers/test_fal_ai_passthrough_wire.py::test_fal_queue_submit_to_catalog_key_the_pricer_cannot_price_is_rejected_not_forwarded": [
|
||||
"other.provider_wire.fal_ai.passthrough_queue_submit_rejects_unpriceable_catalog_key"
|
||||
],
|
||||
"tests/integration/providers/test_fal_ai_image_wire.py::test_fal_gpt_image_25_edit_inlines_upload_as_data_url_and_charges_keyed_row": [
|
||||
"other.provider_wire.fal_ai.image_edit_json_data_urls_and_keyed_pricing"
|
||||
],
|
||||
|
|
|
|||
|
|
@ -1,4 +1,5 @@
|
|||
import json
|
||||
from pathlib import Path
|
||||
from typing import Final
|
||||
|
||||
import pytest
|
||||
|
|
@ -84,3 +85,107 @@ def test_fal_queue_submit_charges_and_polls_pass_through_free(gateway: Gateway,
|
|||
("GET", f"/{_MODEL}/requests/req-1/status"),
|
||||
("GET", f"/{_MODEL}/requests/req-1"),
|
||||
]
|
||||
|
||||
|
||||
@pytest.mark.covers("other.provider_wire.fal_ai.passthrough_queue_submit_rejects_unpriceable_catalog_key")
|
||||
def test_fal_queue_submit_to_catalog_key_the_pricer_cannot_price_is_rejected_not_forwarded(
|
||||
gateway: Gateway, tmp_path: Path
|
||||
) -> None:
|
||||
def respond(request: Request) -> Reply:
|
||||
return Reply(body=json.dumps({"request_id": "req-1", "status": "IN_QUEUE"}).encode())
|
||||
|
||||
config: Final = tmp_path / "proxy_config.yaml"
|
||||
config.write_text(
|
||||
"model_list: []\n"
|
||||
"general_settings:\n"
|
||||
" master_key: os.environ/LITELLM_MASTER_KEY\n"
|
||||
" database_url: os.environ/DATABASE_URL\n"
|
||||
" store_model_in_db: true\n"
|
||||
" disable_spend_logs: false\n"
|
||||
" proxy_batch_write_at: 1\n"
|
||||
"router_settings:\n"
|
||||
" disable_cooldowns: true\n"
|
||||
)
|
||||
with wire_server(respond) as wire:
|
||||
with owned_proxy(
|
||||
gateway,
|
||||
tmp_path,
|
||||
{"FAL_AI_QUEUE_API_BASE": wire.url, "FAL_AI_API_KEY": "synthetic-fal-key"},
|
||||
config=config,
|
||||
) as candidate:
|
||||
submit: Final = candidate.request(
|
||||
"POST",
|
||||
"/fal_ai/fal-ai/moondream3-preview/query",
|
||||
{"image_url": "https://example.com/in.png", "prompt": "one word"},
|
||||
)
|
||||
assert submit.status_code == 400, submit.text
|
||||
assert "pricing" in submit.text
|
||||
assert wire.drain() == ()
|
||||
|
||||
|
||||
@pytest.mark.covers("other.provider_wire.fal_ai.passthrough_queue_submit_prices_string_resolution_like_integer")
|
||||
def test_fal_queue_submit_prices_string_resolution_like_the_integer(gateway: Gateway, tmp_path: Path) -> None:
|
||||
def respond(request: Request) -> Reply:
|
||||
assert request.headers["authorization"] == "Key synthetic-fal-key"
|
||||
return Reply(body=json.dumps({"request_id": "req-1", "status": "IN_QUEUE"}).encode())
|
||||
|
||||
numeric_body: Final = {"image_url": "https://example.com/in.png", "resolution": 512}
|
||||
string_body: Final = {"image_url": "https://example.com/in.png", "resolution": "512"}
|
||||
config: Final = tmp_path / "proxy_config.yaml"
|
||||
config.write_text(
|
||||
"model_list: []\n"
|
||||
"general_settings:\n"
|
||||
" master_key: os.environ/LITELLM_MASTER_KEY\n"
|
||||
" database_url: os.environ/DATABASE_URL\n"
|
||||
" store_model_in_db: true\n"
|
||||
" disable_spend_logs: false\n"
|
||||
" proxy_batch_write_at: 1\n"
|
||||
"router_settings:\n"
|
||||
" disable_cooldowns: true\n"
|
||||
)
|
||||
with wire_server(respond) as wire:
|
||||
with owned_proxy(
|
||||
gateway,
|
||||
tmp_path,
|
||||
{"FAL_AI_QUEUE_API_BASE": wire.url, "FAL_AI_API_KEY": "synthetic-fal-key"},
|
||||
config=config,
|
||||
) as candidate:
|
||||
numeric: Final = candidate.request("POST", f"/fal_ai/{_MODEL}", numeric_body)
|
||||
assert numeric.status_code == 200, numeric.text
|
||||
assert json.loads(numeric.content) == {"request_id": "req-1", "status": "IN_QUEUE"}
|
||||
string: Final = candidate.request("POST", f"/fal_ai/{_MODEL}", string_body)
|
||||
assert string.status_code == 200, string.text
|
||||
assert json.loads(string.content) == {"request_id": "req-1", "status": "IN_QUEUE"}
|
||||
numeric_rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT spend FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
|
||||
(numeric.headers["x-litellm-call-id"],),
|
||||
),
|
||||
lambda values: len(values) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
string_rows: Final = eventually(
|
||||
lambda: read_rows(
|
||||
'SELECT spend FROM "LiteLLM_SpendLogs" WHERE request_id=%s',
|
||||
(string.headers["x-litellm-call-id"],),
|
||||
),
|
||||
lambda values: len(values) == 1,
|
||||
seconds=70,
|
||||
)
|
||||
numeric_spend_value: Final = numeric_rows[0]["spend"]
|
||||
string_spend_value: Final = string_rows[0]["spend"]
|
||||
assert isinstance(numeric_spend_value, (int, float))
|
||||
assert isinstance(string_spend_value, (int, float))
|
||||
numeric_spend: Final = float(numeric_spend_value)
|
||||
string_spend: Final = float(string_spend_value)
|
||||
assert numeric_spend > 0, f"resolution 512 logged {numeric_spend} spend"
|
||||
assert string_spend > 0, f'resolution "512" logged {string_spend} spend'
|
||||
assert numeric_spend == string_spend, (
|
||||
f'resolution 512 was billed {numeric_spend} but resolution "512" was billed {string_spend}'
|
||||
)
|
||||
forwarded: Final = wire.drain()
|
||||
assert [(request.method, request.target) for request in forwarded] == [
|
||||
("POST", f"/{_MODEL}"),
|
||||
("POST", f"/{_MODEL}"),
|
||||
]
|
||||
assert [json.loads(request.body) for request in forwarded] == [numeric_body, string_body]
|
||||
|
|
|
|||
|
|
@ -7,6 +7,7 @@ from litellm.litellm_core_utils.llm_cost_calc.utils import CostCalculatorUtils
|
|||
from litellm.llms.fal_ai.cost_calculator import cost_calculator, fal_ai_passthrough_cost
|
||||
from litellm.types.utils import ImageObject, ImageResponse
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _use_local_model_cost_map(monkeypatch):
|
||||
monkeypatch.setenv("LITELLM_LOCAL_MODEL_COST_MAP", "True")
|
||||
|
|
@ -203,3 +204,41 @@ def test_passthrough_trellis_2_without_resolution_falls_back_to_default_rate():
|
|||
|
||||
def test_passthrough_unknown_model_returns_none():
|
||||
assert fal_ai_passthrough_cost("fal-ai/no-such-model", {"resolution": 512}) is None
|
||||
|
||||
|
||||
def test_passthrough_string_resolution_is_priced_like_the_integer(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"fal_ai/fal-ai/keyed-model",
|
||||
{
|
||||
"litellm_provider": "fal_ai",
|
||||
"mode": "image_generation",
|
||||
"output_cost_per_image": 0.3,
|
||||
"output_cost_per_image_512": 0.25,
|
||||
"output_cost_per_image_1536": 0.35,
|
||||
},
|
||||
)
|
||||
assert fal_ai_passthrough_cost("fal-ai/keyed-model", {"resolution": "512"}) == 0.25
|
||||
assert fal_ai_passthrough_cost("fal-ai/keyed-model", {"resolution": 512}) == 0.25
|
||||
assert fal_ai_passthrough_cost("fal-ai/keyed-model", {"resolution": "1536"}) == 0.35
|
||||
assert fal_ai_passthrough_cost("fal-ai/keyed-model", {"resolution": True}) == 0.3
|
||||
assert fal_ai_passthrough_cost("fal-ai/keyed-model", {"resolution": 512.0}) == 0.3
|
||||
|
||||
|
||||
def test_passthrough_cost_is_none_only_when_no_price_applies_to_the_request(monkeypatch: pytest.MonkeyPatch):
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"fal_ai/fal-ai/priceless-model",
|
||||
{"litellm_provider": "fal_ai", "mode": "image_generation"},
|
||||
)
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"fal_ai/fal-ai/keyed-only-model",
|
||||
{"litellm_provider": "fal_ai", "mode": "image_generation", "output_cost_per_image_512": 0.02},
|
||||
)
|
||||
assert fal_ai_passthrough_cost("fal-ai/priceless-model", {}) is None
|
||||
assert fal_ai_passthrough_cost("fal-ai/priceless-model", {"resolution": 512}) is None
|
||||
assert fal_ai_passthrough_cost("fal-ai/no-such-model", {}) is None
|
||||
assert fal_ai_passthrough_cost("fal-ai/keyed-only-model", {}) is None
|
||||
assert fal_ai_passthrough_cost("fal-ai/keyed-only-model", {"resolution": 1024}) is None
|
||||
assert fal_ai_passthrough_cost("fal-ai/keyed-only-model", {"resolution": "512"}) == 0.02
|
||||
|
|
|
|||
|
|
@ -7200,6 +7200,46 @@ class TestFalAIPassthroughRoute:
|
|||
assert "no pricing entry" in response.text
|
||||
assert not route.calls
|
||||
|
||||
def test_submit_to_catalog_key_the_pricer_cannot_price_returns_400_without_upstream_call(
|
||||
self, client: TestClient, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"fal_ai/fal-ai/priceless-model",
|
||||
{"litellm_provider": "fal_ai", "mode": "image_generation"},
|
||||
)
|
||||
with respx.mock(assert_all_called=False) as upstream:
|
||||
route = upstream.post("https://queue.fal.run/fal-ai/priceless-model").mock(
|
||||
return_value=httpx.Response(200, json={"request_id": "req-1"})
|
||||
)
|
||||
response = client.post("/fal_ai/fal-ai/priceless-model", json={"image_url": "https://example.com/in.png"})
|
||||
|
||||
assert response.status_code == 400, response.text
|
||||
assert "no pricing entry" in response.text
|
||||
assert not route.calls
|
||||
|
||||
def test_submit_gate_prices_the_request_body_not_an_empty_one(
|
||||
self, client: TestClient, monkeypatch: pytest.MonkeyPatch
|
||||
) -> None:
|
||||
monkeypatch.setitem(
|
||||
litellm.model_cost,
|
||||
"fal_ai/fal-ai/keyed-only-model",
|
||||
{"litellm_provider": "fal_ai", "mode": "image_generation", "output_cost_per_image_512": 0.02},
|
||||
)
|
||||
with respx.mock(assert_all_called=False) as upstream:
|
||||
route = upstream.post("https://queue.fal.run/fal-ai/keyed-only-model").mock(
|
||||
return_value=httpx.Response(200, json={"request_id": "req-1"})
|
||||
)
|
||||
priced = client.post(
|
||||
"/fal_ai/fal-ai/keyed-only-model", json={"image_url": "https://example.com/in.png", "resolution": "512"}
|
||||
)
|
||||
unpriced = client.post("/fal_ai/fal-ai/keyed-only-model", json={"image_url": "https://example.com/in.png"})
|
||||
|
||||
assert priced.status_code == 200, priced.text
|
||||
assert unpriced.status_code == 400, unpriced.text
|
||||
assert "no pricing entry" in unpriced.text
|
||||
assert len(route.calls) == 1
|
||||
|
||||
def test_status_get_on_unpriced_endpoint_forwards(self, client: TestClient) -> None:
|
||||
with respx.mock(assert_all_called=True) as upstream:
|
||||
upstream.get("https://queue.fal.run/fal-ai/unpriced-model/requests/req-9/status").mock(
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue