mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-19 00:01:29 +00:00
refactor(proxy): simplify TypeSafe passthrough pricing lookup and route
Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
parent
2dc9697381
commit
78eb92ca55
3 changed files with 20 additions and 48 deletions
|
|
@ -537,16 +537,6 @@ async def typesafe_proxy_route(
|
|||
user_api_key_dict: UserAPIKeyAuth = Depends(user_api_key_auth),
|
||||
):
|
||||
"""[Docs](https://docs.litellm.ai/docs/pass_through/typesafe)"""
|
||||
if request.method == "POST":
|
||||
try:
|
||||
request_body: Final = await _json_request_body(request)
|
||||
except Exception as e:
|
||||
raise HTTPException(status_code=400, detail=str(e))
|
||||
if not isinstance(request_body, dict):
|
||||
raise HTTPException(status_code=400, detail="Request body must be a JSON object")
|
||||
if "stream" in request_body:
|
||||
raise HTTPException(status_code=400, detail="'stream' is not a TypeSafe request member")
|
||||
|
||||
base_target_url: Final = get_secret_str("TYPESAFE_API_BASE") or "https://api.typesafe.ai"
|
||||
encoded_endpoint: Final = httpx.URL(endpoint).path
|
||||
normalized_endpoint: Final = encoded_endpoint if encoded_endpoint.startswith("/") else f"/{encoded_endpoint}"
|
||||
|
|
|
|||
|
|
@ -1,6 +1,6 @@
|
|||
from collections.abc import Mapping
|
||||
from datetime import datetime
|
||||
from typing import Final, cast
|
||||
from typing import Final
|
||||
|
||||
import httpx
|
||||
from pydantic import BaseModel, TypeAdapter, ValidationError
|
||||
|
|
@ -24,8 +24,13 @@ class _TypeSafeResponse(BaseModel):
|
|||
usage: _TypeSafeUsage | None = None
|
||||
|
||||
|
||||
class _RegistryPricing(BaseModel):
|
||||
input_cost_per_token: float = 0.0
|
||||
output_cost_per_token: float = 0.0
|
||||
|
||||
|
||||
_TYPESAFE_RESPONSE_ADAPTER: Final = TypeAdapter(_TypeSafeResponse)
|
||||
_MODEL_COST_ENTRY_ADAPTER: Final = TypeAdapter(dict[str, object])
|
||||
_REGISTRY_PRICING_ADAPTER: Final = TypeAdapter(_RegistryPricing)
|
||||
|
||||
|
||||
def _parse_typesafe_response(response_body: Mapping[str, object]) -> _TypeSafeResponse:
|
||||
|
|
@ -35,13 +40,17 @@ def _parse_typesafe_response(response_body: Mapping[str, object]) -> _TypeSafeRe
|
|||
return _TypeSafeResponse()
|
||||
|
||||
|
||||
def _get_model_cost_entry(model_key: str) -> Mapping[str, object] | None:
|
||||
model_cost: Final[Mapping[str, object]] = cast(Mapping[str, object], litellm.model_cost)
|
||||
entry: Final[object] = model_cost.get(model_key)
|
||||
try:
|
||||
return _MODEL_COST_ENTRY_ADAPTER.validate_python(entry)
|
||||
except ValidationError:
|
||||
return None
|
||||
def _pricing_for(model_keys: tuple[str, ...]) -> _RegistryPricing:
|
||||
for model_key in model_keys:
|
||||
if model_key not in litellm.model_cost: # pyright: ignore[reportUnknownMemberType] # registry is dynamically typed
|
||||
continue
|
||||
try:
|
||||
return _REGISTRY_PRICING_ADAPTER.validate_python(
|
||||
litellm.model_cost[model_key] # pyright: ignore[reportUnknownMemberType] # registry is dynamically typed
|
||||
)
|
||||
except ValidationError:
|
||||
continue
|
||||
return _RegistryPricing()
|
||||
|
||||
|
||||
class TypeSafePassthroughLoggingHandler:
|
||||
|
|
@ -70,20 +79,9 @@ class TypeSafePassthroughLoggingHandler:
|
|||
candidate_model_keys: Final = tuple(
|
||||
f"typesafe/{model}" for model in (response_model, request_model) if model is not None
|
||||
)
|
||||
cost_entry: Final = next(
|
||||
(entry for model_key in candidate_model_keys if (entry := _get_model_cost_entry(model_key)) is not None),
|
||||
None,
|
||||
)
|
||||
input_cost_per_token: Final = (
|
||||
cost_entry.get("input_cost_per_token", 0.0) if isinstance(cost_entry, Mapping) else 0.0
|
||||
)
|
||||
output_cost_per_token: Final = (
|
||||
cost_entry.get("output_cost_per_token", 0.0) if isinstance(cost_entry, Mapping) else 0.0
|
||||
)
|
||||
pricing: Final = _pricing_for(candidate_model_keys)
|
||||
response_cost: Final = (
|
||||
input_tokens * float(input_cost_per_token) + output_tokens * float(output_cost_per_token)
|
||||
if isinstance(input_cost_per_token, (int, float)) and isinstance(output_cost_per_token, (int, float))
|
||||
else 0.0
|
||||
input_tokens * pricing.input_cost_per_token + output_tokens * pricing.output_cost_per_token
|
||||
)
|
||||
usage_object: Final = Usage(
|
||||
prompt_tokens=input_tokens,
|
||||
|
|
|
|||
|
|
@ -6179,19 +6179,3 @@ class TestTypeSafePassthroughRoute:
|
|||
custom_llm_provider="typesafe",
|
||||
is_streaming_request=False,
|
||||
)
|
||||
assert request.json.await_count == 1
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_rejects_stream_body(self, monkeypatch):
|
||||
monkeypatch.setenv("TYPESAFE_API_KEY", "typesafe-test-key")
|
||||
request = self._request({"stream": True})
|
||||
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await typesafe_proxy_route(
|
||||
endpoint="v1/systemone",
|
||||
request=request,
|
||||
fastapi_response=MagicMock(spec=Response),
|
||||
user_api_key_dict=UserAPIKeyAuth(api_key="virtual-key"),
|
||||
)
|
||||
|
||||
assert exc_info.value.status_code == 400
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue