mirror of
https://github.com/BerriAI/litellm.git
synced 2026-09-13 23:11:40 +00:00
fix(proxy): split rust control-plane auth and logging endpoints
This commit is contained in:
parent
d1c371c337
commit
cbc81ac7be
7 changed files with 366 additions and 320 deletions
|
|
@ -346,7 +346,7 @@ from litellm.proxy.hooks.prompt_injection_detection import (
|
|||
from litellm.proxy.hooks.proxy_track_cost_callback import _ProxyDBLogger
|
||||
from litellm.proxy.image_endpoints.endpoints import router as image_router
|
||||
from litellm.proxy.litellm_pre_call_utils import add_litellm_data_to_request
|
||||
from litellm.proxy.rust_control_plane_endpoints.callback_logs_endpoints import (
|
||||
from litellm.proxy.rust_control_plane_endpoints.router import (
|
||||
rust_control_plane_router,
|
||||
)
|
||||
from litellm.proxy.management_endpoints.budget_management_endpoints import (
|
||||
|
|
|
|||
127
litellm/proxy/rust_control_plane_endpoints/auth_endpoints.py
Normal file
127
litellm/proxy/rust_control_plane_endpoints/auth_endpoints.py
Normal file
|
|
@ -0,0 +1,127 @@
|
|||
"""
|
||||
Authentication endpoints consumed by the Rust data-plane gateway.
|
||||
|
||||
The Rust gateway terminates client connections and needs to validate virtual
|
||||
keys without reimplementing LiteLLM's proxy auth logic. It calls this internal
|
||||
control-plane route to verify the key against the requested data-plane route
|
||||
and model.
|
||||
"""
|
||||
|
||||
import hmac
|
||||
import json
|
||||
import os
|
||||
from typing import Any, Optional
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from pydantic import BaseModel
|
||||
|
||||
from litellm.proxy._types import ProxyException
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
|
||||
DATA_PLANE_KEY_ENV_VAR = "LITELLM_DATA_PLANE_KEY"
|
||||
DATA_PLANE_KEY_HEADER = "X-LiteLLM-Data-Plane-Key"
|
||||
|
||||
router = APIRouter(prefix="/v1/rust_control_plane", tags=["rust control plane"])
|
||||
|
||||
|
||||
def require_data_plane_key(request: Request) -> None:
|
||||
"""
|
||||
Authenticate requests from the Rust data plane with a dedicated secret.
|
||||
|
||||
This intentionally uses ``LITELLM_DATA_PLANE_KEY`` instead of the proxy
|
||||
master key: the data plane is a separate trust boundary and must not get
|
||||
admin privileges.
|
||||
"""
|
||||
expected_key: Optional[str] = os.getenv(DATA_PLANE_KEY_ENV_VAR)
|
||||
if not expected_key:
|
||||
raise HTTPException(status_code=500, detail="data-plane auth not configured")
|
||||
|
||||
provided_key: Optional[str] = request.headers.get(DATA_PLANE_KEY_HEADER)
|
||||
if not provided_key or not hmac.compare_digest(provided_key, expected_key):
|
||||
raise HTTPException(status_code=401, detail="invalid data-plane key")
|
||||
|
||||
|
||||
class VerifyKeyRequest(BaseModel):
|
||||
api_key: str
|
||||
# The actual route the gateway is serving this key on (e.g. "/v1/realtime").
|
||||
# REQUIRED and not defaulted: the gateway always sends its own request path,
|
||||
# so validation runs route/model restrictions against the real route.
|
||||
route: str
|
||||
# Forwarded so user_api_key_auth's model access checks run for key, team,
|
||||
# and access-group restrictions.
|
||||
model: Optional[str] = None
|
||||
|
||||
|
||||
def _synthetic_request(
|
||||
route: str, authorization_header: str, model: Optional[str]
|
||||
) -> Request:
|
||||
"""
|
||||
Build a minimal ASGI request for user_api_key_auth to validate the key
|
||||
against the data-plane route and model instead of this internal endpoint.
|
||||
"""
|
||||
body = json.dumps({"model": model} if model is not None else {}).encode()
|
||||
|
||||
async def receive() -> dict[str, Any]:
|
||||
return {"type": "http.request", "body": body, "more_body": False}
|
||||
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": route,
|
||||
"raw_path": route.encode(),
|
||||
"headers": [
|
||||
(b"authorization", authorization_header.encode()),
|
||||
(b"content-type", b"application/json"),
|
||||
],
|
||||
"query_string": b"",
|
||||
"scheme": "http",
|
||||
"client": ("127.0.0.1", 0),
|
||||
"server": ("127.0.0.1", 4000),
|
||||
}
|
||||
request = Request(scope, receive)
|
||||
# Admission checks only. Realtime spend is reconciled later through callback
|
||||
# logs, so optimistic reservation here would have no request lifecycle to
|
||||
# release it.
|
||||
request.state.skip_budget_reservation = True
|
||||
return request
|
||||
|
||||
|
||||
@router.post(
|
||||
"/authentication",
|
||||
dependencies=[Depends(require_data_plane_key)],
|
||||
include_in_schema=False,
|
||||
)
|
||||
async def verify_key(body: VerifyKeyRequest) -> dict[str, Any]:
|
||||
"""
|
||||
Verify a virtual key on behalf of the Rust ai-gateway (data plane).
|
||||
|
||||
The Rust gateway forwards the client's virtual key here; this delegates to
|
||||
the proxy's existing ``user_api_key_auth`` validation and returns the
|
||||
resolved ``UserAPIKeyAuth`` as JSON so the data plane can enforce the same
|
||||
limits the control plane would.
|
||||
|
||||
Gated by ``require_data_plane_key`` (the dedicated data-plane secret, not
|
||||
the master key).
|
||||
|
||||
On any auth failure the response is a 401 with a minimal body so internals
|
||||
are not leaked to the caller.
|
||||
"""
|
||||
bearer_key = (
|
||||
body.api_key if body.api_key.startswith("Bearer ") else f"Bearer {body.api_key}"
|
||||
)
|
||||
synthetic_request = _synthetic_request(
|
||||
route=body.route, authorization_header=bearer_key, model=body.model
|
||||
)
|
||||
try:
|
||||
auth = await user_api_key_auth(request=synthetic_request, api_key=bearer_key)
|
||||
except (ProxyException, HTTPException) as exc:
|
||||
status_code = getattr(exc, "status_code", None) or getattr(exc, "code", None)
|
||||
try:
|
||||
is_server_error = status_code is not None and int(status_code) >= 500
|
||||
except (TypeError, ValueError):
|
||||
is_server_error = False
|
||||
if is_server_error:
|
||||
raise
|
||||
raise HTTPException(status_code=401, detail="invalid api key")
|
||||
|
||||
return auth.model_dump(exclude_none=True, mode="json")
|
||||
|
|
@ -14,19 +14,15 @@ The endpoint is generic: realtime is the first producer, but the contract is the
|
|||
self-describing `StandardLoggingPayload`, so completions/responses can use it too.
|
||||
"""
|
||||
|
||||
import hmac
|
||||
import json
|
||||
import os
|
||||
import uuid
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Optional
|
||||
from typing import Any
|
||||
|
||||
from fastapi import APIRouter, Depends, HTTPException, Request
|
||||
from pydantic import BaseModel
|
||||
from fastapi import APIRouter, Depends, HTTPException
|
||||
|
||||
from litellm._logging import verbose_proxy_logger
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
from litellm.proxy._types import LitellmUserRoles, ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.auth.user_api_key_auth import user_api_key_auth
|
||||
from litellm.types.proxy.callback_logs_endpoints import (
|
||||
CallbackLogFailure,
|
||||
|
|
@ -35,118 +31,7 @@ from litellm.types.proxy.callback_logs_endpoints import (
|
|||
CallbackLogsResponse,
|
||||
)
|
||||
|
||||
# Routes the Python proxy exposes for the Rust data-plane gateway to call into
|
||||
# (auth + logging today; budgets later). Namespaced under /v1/rust_control_plane
|
||||
# so they're clearly distinct from the proxy's own control-plane/management routes.
|
||||
rust_control_plane_router = APIRouter(
|
||||
prefix="/v1/rust_control_plane", tags=["rust control plane"]
|
||||
)
|
||||
|
||||
DATA_PLANE_KEY_ENV_VAR = "LITELLM_DATA_PLANE_KEY"
|
||||
DATA_PLANE_KEY_HEADER = "X-LiteLLM-Data-Plane-Key"
|
||||
|
||||
|
||||
def require_data_plane_key(request: Request) -> None:
|
||||
"""
|
||||
Authenticate requests from the Rust data plane with a dedicated secret.
|
||||
|
||||
This intentionally uses ``LITELLM_DATA_PLANE_KEY`` instead of the proxy
|
||||
master key: the data plane is a separate trust boundary and must not get
|
||||
admin privileges.
|
||||
"""
|
||||
expected_key: Optional[str] = os.getenv(DATA_PLANE_KEY_ENV_VAR)
|
||||
if not expected_key:
|
||||
raise HTTPException(status_code=500, detail="data-plane auth not configured")
|
||||
|
||||
provided_key: Optional[str] = request.headers.get(DATA_PLANE_KEY_HEADER)
|
||||
if not provided_key or not hmac.compare_digest(provided_key, expected_key):
|
||||
raise HTTPException(status_code=401, detail="invalid data-plane key")
|
||||
|
||||
|
||||
class VerifyKeyRequest(BaseModel):
|
||||
api_key: str
|
||||
# The actual route the gateway is serving this key on (e.g. "/v1/realtime").
|
||||
# REQUIRED and not defaulted: the gateway always sends its own request path,
|
||||
# so validation runs route/model restrictions against the real route.
|
||||
route: str
|
||||
# Forwarded so user_api_key_auth's model access checks run for key, team,
|
||||
# and access-group restrictions.
|
||||
model: Optional[str] = None
|
||||
|
||||
|
||||
def _synthetic_request(
|
||||
route: str, authorization_header: str, model: Optional[str]
|
||||
) -> Request:
|
||||
"""
|
||||
Build a minimal ASGI request for user_api_key_auth to validate the key
|
||||
against the data-plane route and model instead of this internal endpoint.
|
||||
"""
|
||||
body = json.dumps({"model": model} if model is not None else {}).encode()
|
||||
|
||||
async def receive() -> dict[str, Any]:
|
||||
return {"type": "http.request", "body": body, "more_body": False}
|
||||
|
||||
scope = {
|
||||
"type": "http",
|
||||
"method": "POST",
|
||||
"path": route,
|
||||
"raw_path": route.encode(),
|
||||
"headers": [
|
||||
(b"authorization", authorization_header.encode()),
|
||||
(b"content-type", b"application/json"),
|
||||
],
|
||||
"query_string": b"",
|
||||
"scheme": "http",
|
||||
"client": ("127.0.0.1", 0),
|
||||
"server": ("127.0.0.1", 4000),
|
||||
}
|
||||
request = Request(scope, receive)
|
||||
# Admission checks only. Realtime spend is reconciled later through callback
|
||||
# logs, so optimistic reservation here would have no request lifecycle to
|
||||
# release it.
|
||||
request.state.skip_budget_reservation = True
|
||||
return request
|
||||
|
||||
|
||||
@rust_control_plane_router.post(
|
||||
"/authentication",
|
||||
dependencies=[Depends(require_data_plane_key)],
|
||||
include_in_schema=False,
|
||||
)
|
||||
async def verify_key(body: VerifyKeyRequest) -> dict[str, Any]:
|
||||
"""
|
||||
Verify a virtual key on behalf of the Rust ai-gateway (data plane).
|
||||
|
||||
The Rust gateway forwards the client's virtual key here; this delegates to
|
||||
the proxy's existing ``user_api_key_auth`` validation and returns the
|
||||
resolved ``UserAPIKeyAuth`` as JSON so the data plane can enforce the same
|
||||
limits the control plane would.
|
||||
|
||||
Gated by ``require_data_plane_key`` (the dedicated data-plane secret, not
|
||||
the master key).
|
||||
|
||||
On any auth failure the response is a 401 with a minimal body so internals
|
||||
are not leaked to the caller.
|
||||
"""
|
||||
bearer_key = (
|
||||
body.api_key if body.api_key.startswith("Bearer ") else f"Bearer {body.api_key}"
|
||||
)
|
||||
synthetic_request = _synthetic_request(
|
||||
route=body.route, authorization_header=bearer_key, model=body.model
|
||||
)
|
||||
try:
|
||||
auth = await user_api_key_auth(request=synthetic_request, api_key=bearer_key)
|
||||
except (ProxyException, HTTPException) as exc:
|
||||
status_code = getattr(exc, "status_code", None) or getattr(exc, "code", None)
|
||||
try:
|
||||
is_server_error = status_code is not None and int(status_code) >= 500
|
||||
except (TypeError, ValueError):
|
||||
is_server_error = False
|
||||
if is_server_error:
|
||||
raise
|
||||
raise HTTPException(status_code=401, detail="invalid api key")
|
||||
|
||||
return auth.model_dump(exclude_none=True, mode="json")
|
||||
router = APIRouter(prefix="/v1/rust_control_plane", tags=["rust control plane"])
|
||||
|
||||
|
||||
class CallbackLogsReplayer:
|
||||
|
|
@ -295,7 +180,7 @@ class CallbackLogsReplayer:
|
|||
)
|
||||
|
||||
|
||||
@rust_control_plane_router.post(
|
||||
@router.post(
|
||||
"/logs",
|
||||
dependencies=[Depends(user_api_key_auth)],
|
||||
)
|
||||
14
litellm/proxy/rust_control_plane_endpoints/router.py
Normal file
14
litellm/proxy/rust_control_plane_endpoints/router.py
Normal file
|
|
@ -0,0 +1,14 @@
|
|||
"""Combined router for Python endpoints consumed by the Rust data plane."""
|
||||
|
||||
from fastapi import APIRouter
|
||||
|
||||
from litellm.proxy.rust_control_plane_endpoints.auth_endpoints import (
|
||||
router as auth_router,
|
||||
)
|
||||
from litellm.proxy.rust_control_plane_endpoints.logging_endpoints import (
|
||||
router as logging_router,
|
||||
)
|
||||
|
||||
rust_control_plane_router = APIRouter()
|
||||
rust_control_plane_router.include_router(auth_router)
|
||||
rust_control_plane_router.include_router(logging_router)
|
||||
|
|
@ -1,29 +1,18 @@
|
|||
"""Unit tests for POST /v1/callbacks/logs (replay logging payloads → callbacks)."""
|
||||
|
||||
import time
|
||||
"""Unit tests for POST /v1/rust_control_plane/authentication."""
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException, Request
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
from litellm.proxy._types import LitellmUserRoles, ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.rust_control_plane_endpoints.callback_logs_endpoints import (
|
||||
CallbackLogsReplayer,
|
||||
from litellm.proxy._types import ProxyException, UserAPIKeyAuth
|
||||
from litellm.proxy.rust_control_plane_endpoints.auth_endpoints import (
|
||||
DATA_PLANE_KEY_ENV_VAR,
|
||||
DATA_PLANE_KEY_HEADER,
|
||||
VerifyKeyRequest,
|
||||
_synthetic_request,
|
||||
ingest_callback_logs,
|
||||
require_data_plane_key,
|
||||
rust_control_plane_router,
|
||||
router,
|
||||
verify_key,
|
||||
)
|
||||
from litellm.types.proxy.callback_logs_endpoints import (
|
||||
CallbackLogRecord,
|
||||
CallbackLogsRequest,
|
||||
)
|
||||
|
||||
REQ_ID = "cb-logs-unit-test-1"
|
||||
|
||||
|
||||
def _make_request(headers: dict) -> Request:
|
||||
|
|
@ -40,31 +29,6 @@ def _make_request(headers: dict) -> Request:
|
|||
return Request(scope)
|
||||
|
||||
|
||||
def _sample_payload(**overrides):
|
||||
payload = {
|
||||
"id": REQ_ID,
|
||||
"litellm_call_id": REQ_ID,
|
||||
"call_type": "acompletion",
|
||||
"stream": False,
|
||||
"response_cost": 0.0123,
|
||||
"custom_llm_provider": "openai",
|
||||
"total_tokens": 42,
|
||||
"prompt_tokens": 30,
|
||||
"completion_tokens": 12,
|
||||
"startTime": time.time() - 2,
|
||||
"endTime": time.time(),
|
||||
"model": "gpt-4o-mini",
|
||||
"metadata": {
|
||||
"user_api_key_hash": "rust-gateway-test-key",
|
||||
"user_api_key_user_id": "user-cb-logs-test",
|
||||
"user_api_key_team_id": "team-cb-logs-test",
|
||||
},
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
}
|
||||
payload.update(overrides)
|
||||
return payload
|
||||
|
||||
|
||||
def test_require_data_plane_key_500_when_env_unset(monkeypatch):
|
||||
monkeypatch.delenv(DATA_PLANE_KEY_ENV_VAR, raising=False)
|
||||
request = _make_request({DATA_PLANE_KEY_HEADER: "anything"})
|
||||
|
|
@ -118,7 +82,7 @@ def test_require_data_plane_key_passes_when_correct(monkeypatch):
|
|||
def test_router_mounts_auth_verify_under_rust_control_plane():
|
||||
assert any(
|
||||
getattr(route, "path", None) == "/v1/rust_control_plane/authentication"
|
||||
for route in rust_control_plane_router.routes
|
||||
for route in router.routes
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -149,7 +113,7 @@ async def test_verify_key_returns_model_dump(monkeypatch):
|
|||
return expected_auth
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.rust_control_plane_endpoints.callback_logs_endpoints.user_api_key_auth",
|
||||
"litellm.proxy.rust_control_plane_endpoints.auth_endpoints.user_api_key_auth",
|
||||
fake_user_api_key_auth,
|
||||
)
|
||||
|
||||
|
|
@ -178,7 +142,7 @@ async def test_verify_key_omits_model_when_absent(monkeypatch):
|
|||
return UserAPIKeyAuth(api_key="hashed-key")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.rust_control_plane_endpoints.callback_logs_endpoints.user_api_key_auth",
|
||||
"litellm.proxy.rust_control_plane_endpoints.auth_endpoints.user_api_key_auth",
|
||||
fake_user_api_key_auth,
|
||||
)
|
||||
|
||||
|
|
@ -198,7 +162,7 @@ async def test_verify_key_does_not_double_prefix_existing_bearer(monkeypatch):
|
|||
return UserAPIKeyAuth(api_key="hashed-key")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.rust_control_plane_endpoints.callback_logs_endpoints.user_api_key_auth",
|
||||
"litellm.proxy.rust_control_plane_endpoints.auth_endpoints.user_api_key_auth",
|
||||
fake_user_api_key_auth,
|
||||
)
|
||||
|
||||
|
|
@ -220,7 +184,7 @@ async def test_verify_key_401_on_proxy_exception(monkeypatch):
|
|||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.rust_control_plane_endpoints.callback_logs_endpoints.user_api_key_auth",
|
||||
"litellm.proxy.rust_control_plane_endpoints.auth_endpoints.user_api_key_auth",
|
||||
fake_user_api_key_auth,
|
||||
)
|
||||
|
||||
|
|
@ -237,7 +201,7 @@ async def test_verify_key_401_on_http_exception(monkeypatch):
|
|||
raise HTTPException(status_code=403, detail="forbidden internals")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.rust_control_plane_endpoints.callback_logs_endpoints.user_api_key_auth",
|
||||
"litellm.proxy.rust_control_plane_endpoints.auth_endpoints.user_api_key_auth",
|
||||
fake_user_api_key_auth,
|
||||
)
|
||||
|
||||
|
|
@ -256,7 +220,7 @@ async def test_verify_key_propagates_http_5xx(monkeypatch):
|
|||
raise HTTPException(status_code=503, detail="db unavailable")
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.rust_control_plane_endpoints.callback_logs_endpoints.user_api_key_auth",
|
||||
"litellm.proxy.rust_control_plane_endpoints.auth_endpoints.user_api_key_auth",
|
||||
fake_user_api_key_auth,
|
||||
)
|
||||
|
||||
|
|
@ -275,161 +239,10 @@ async def test_verify_key_propagates_proxy_5xx(monkeypatch):
|
|||
)
|
||||
|
||||
monkeypatch.setattr(
|
||||
"litellm.proxy.rust_control_plane_endpoints.callback_logs_endpoints.user_api_key_auth",
|
||||
"litellm.proxy.rust_control_plane_endpoints.auth_endpoints.user_api_key_auth",
|
||||
fake_user_api_key_auth,
|
||||
)
|
||||
|
||||
body = VerifyKeyRequest(api_key="sk-key", route="/v1/realtime")
|
||||
with pytest.raises(ProxyException):
|
||||
await verify_key(body=body)
|
||||
|
||||
|
||||
def test_epoch_to_datetime_handles_float_and_fallback():
|
||||
dt = CallbackLogsReplayer._epoch_to_datetime(1_700_000_000.5)
|
||||
assert dt.year == 2023
|
||||
# Non-numeric input must not raise — falls back to "now".
|
||||
assert CallbackLogsReplayer._epoch_to_datetime(None) is not None
|
||||
|
||||
|
||||
def test_build_logging_obj_seeds_model_call_details():
|
||||
obj = CallbackLogsReplayer._build_logging_obj(_sample_payload())
|
||||
details = obj.model_call_details
|
||||
# Prebuilt payload is set so the handler skips rebuilding it.
|
||||
assert details["standard_logging_object"]["id"] == REQ_ID
|
||||
assert details["response_cost"] == 0.0123
|
||||
assert details["call_type"] == "acompletion"
|
||||
# Metadata is mapped to the keys the cost-tracking callback reads.
|
||||
md = details["litellm_params"]["metadata"]
|
||||
assert md["user_api_key"] == "rust-gateway-test-key"
|
||||
assert md["user_api_key_user_id"] == "user-cb-logs-test"
|
||||
assert md["user_api_key_team_id"] == "team-cb-logs-test"
|
||||
|
||||
|
||||
def test_response_obj_carries_usage():
|
||||
obj = CallbackLogsReplayer._response_obj_from_payload(_sample_payload())
|
||||
assert obj["usage"]["total_tokens"] == 42
|
||||
assert obj["usage"]["prompt_tokens"] == 30
|
||||
assert obj["usage"]["completion_tokens"] == 12
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_success_record_invokes_success_handler(monkeypatch):
|
||||
captured = {}
|
||||
|
||||
async def fake_success(self, result=None, start_time=None, end_time=None, **kwargs):
|
||||
captured["standard_logging_object"] = self.model_call_details.get(
|
||||
"standard_logging_object"
|
||||
)
|
||||
captured["result"] = result
|
||||
|
||||
monkeypatch.setattr(LiteLLMLogging, "async_success_handler", fake_success)
|
||||
|
||||
body = CallbackLogsRequest(
|
||||
records=[
|
||||
CallbackLogRecord(
|
||||
status="success", standard_logging_payload=_sample_payload()
|
||||
)
|
||||
]
|
||||
)
|
||||
resp = await ingest_callback_logs(
|
||||
body, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
)
|
||||
assert resp.processed == 1 and resp.failed == 0
|
||||
assert captured["standard_logging_object"]["id"] == REQ_ID
|
||||
assert captured["result"]["usage"]["total_tokens"] == 42
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failure_record_invokes_failure_handler(monkeypatch):
|
||||
captured = {}
|
||||
|
||||
async def fake_failure(
|
||||
self, exception, traceback_exception, start_time=None, end_time=None
|
||||
):
|
||||
captured["exception"] = str(exception)
|
||||
|
||||
monkeypatch.setattr(LiteLLMLogging, "async_failure_handler", fake_failure)
|
||||
|
||||
body = CallbackLogsRequest(
|
||||
records=[
|
||||
CallbackLogRecord(
|
||||
status="failure",
|
||||
standard_logging_payload=_sample_payload(),
|
||||
error="upstream exploded",
|
||||
)
|
||||
]
|
||||
)
|
||||
resp = await ingest_callback_logs(
|
||||
body, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
)
|
||||
assert resp.processed == 1 and resp.failed == 0
|
||||
assert captured["exception"] == "upstream exploded"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_admin_is_rejected(monkeypatch):
|
||||
async def fake_success(self, **kwargs):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(LiteLLMLogging, "async_success_handler", fake_success)
|
||||
|
||||
body = CallbackLogsRequest(
|
||||
records=[
|
||||
CallbackLogRecord(
|
||||
status="success", standard_logging_payload=_sample_payload()
|
||||
)
|
||||
]
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await ingest_callback_logs(
|
||||
body,
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER),
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_one_bad_record_does_not_sink_the_batch(monkeypatch):
|
||||
calls = {"n": 0}
|
||||
|
||||
async def flaky_success(
|
||||
self, result=None, start_time=None, end_time=None, **kwargs
|
||||
):
|
||||
calls["n"] += 1
|
||||
if calls["n"] == 1:
|
||||
raise ValueError("boom on first record")
|
||||
|
||||
monkeypatch.setattr(LiteLLMLogging, "async_success_handler", flaky_success)
|
||||
|
||||
body = CallbackLogsRequest(
|
||||
records=[
|
||||
CallbackLogRecord(
|
||||
status="success", standard_logging_payload=_sample_payload()
|
||||
),
|
||||
CallbackLogRecord(
|
||||
status="success", standard_logging_payload=_sample_payload()
|
||||
),
|
||||
]
|
||||
)
|
||||
resp = await ingest_callback_logs(
|
||||
body, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
)
|
||||
assert resp.processed == 1 and resp.failed == 1
|
||||
# The failed record is reported back by index + error, not silently dropped.
|
||||
assert len(resp.failures) == 1
|
||||
assert resp.failures[0].index == 0
|
||||
assert "boom on first record" in resp.failures[0].error
|
||||
|
||||
|
||||
def test_batch_over_limit_is_rejected():
|
||||
from litellm.constants import MAX_CALLBACK_LOG_RECORDS
|
||||
from pydantic import ValidationError
|
||||
|
||||
# One over the cap must fail validation (422 at the API boundary), bounding
|
||||
# the callback/DB fan-out a single POST can trigger.
|
||||
too_many = [
|
||||
CallbackLogRecord(status="success", standard_logging_payload=_sample_payload())
|
||||
for _ in range(MAX_CALLBACK_LOG_RECORDS + 1)
|
||||
]
|
||||
with pytest.raises(ValidationError):
|
||||
CallbackLogsRequest(records=too_many)
|
||||
|
|
@ -0,0 +1,195 @@
|
|||
"""Unit tests for POST /v1/rust_control_plane/logs."""
|
||||
|
||||
import time
|
||||
|
||||
import pytest
|
||||
from fastapi import HTTPException
|
||||
|
||||
from litellm.litellm_core_utils.litellm_logging import Logging as LiteLLMLogging
|
||||
from litellm.proxy._types import LitellmUserRoles, UserAPIKeyAuth
|
||||
from litellm.proxy.rust_control_plane_endpoints.logging_endpoints import (
|
||||
CallbackLogsReplayer,
|
||||
ingest_callback_logs,
|
||||
)
|
||||
from litellm.types.proxy.callback_logs_endpoints import (
|
||||
CallbackLogRecord,
|
||||
CallbackLogsRequest,
|
||||
)
|
||||
|
||||
REQ_ID = "cb-logs-unit-test-1"
|
||||
|
||||
|
||||
def _sample_payload(**overrides):
|
||||
payload = {
|
||||
"id": REQ_ID,
|
||||
"litellm_call_id": REQ_ID,
|
||||
"call_type": "acompletion",
|
||||
"stream": False,
|
||||
"response_cost": 0.0123,
|
||||
"custom_llm_provider": "openai",
|
||||
"total_tokens": 42,
|
||||
"prompt_tokens": 30,
|
||||
"completion_tokens": 12,
|
||||
"startTime": time.time() - 2,
|
||||
"endTime": time.time(),
|
||||
"model": "gpt-4o-mini",
|
||||
"metadata": {
|
||||
"user_api_key_hash": "rust-gateway-test-key",
|
||||
"user_api_key_user_id": "user-cb-logs-test",
|
||||
"user_api_key_team_id": "team-cb-logs-test",
|
||||
},
|
||||
"messages": [{"role": "user", "content": "hi"}],
|
||||
}
|
||||
payload.update(overrides)
|
||||
return payload
|
||||
|
||||
|
||||
def test_epoch_to_datetime_handles_float_and_fallback():
|
||||
dt = CallbackLogsReplayer._epoch_to_datetime(1_700_000_000.5)
|
||||
assert dt.year == 2023
|
||||
# Non-numeric input must not raise -- falls back to "now".
|
||||
assert CallbackLogsReplayer._epoch_to_datetime(None) is not None
|
||||
|
||||
|
||||
def test_build_logging_obj_seeds_model_call_details():
|
||||
obj = CallbackLogsReplayer._build_logging_obj(_sample_payload())
|
||||
details = obj.model_call_details
|
||||
# Prebuilt payload is set so the handler skips rebuilding it.
|
||||
assert details["standard_logging_object"]["id"] == REQ_ID
|
||||
assert details["response_cost"] == 0.0123
|
||||
assert details["call_type"] == "acompletion"
|
||||
# Metadata is mapped to the keys the cost-tracking callback reads.
|
||||
md = details["litellm_params"]["metadata"]
|
||||
assert md["user_api_key"] == "rust-gateway-test-key"
|
||||
assert md["user_api_key_user_id"] == "user-cb-logs-test"
|
||||
assert md["user_api_key_team_id"] == "team-cb-logs-test"
|
||||
|
||||
|
||||
def test_response_obj_carries_usage():
|
||||
obj = CallbackLogsReplayer._response_obj_from_payload(_sample_payload())
|
||||
assert obj["usage"]["total_tokens"] == 42
|
||||
assert obj["usage"]["prompt_tokens"] == 30
|
||||
assert obj["usage"]["completion_tokens"] == 12
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_success_record_invokes_success_handler(monkeypatch):
|
||||
captured = {}
|
||||
|
||||
async def fake_success(self, result=None, start_time=None, end_time=None, **kwargs):
|
||||
captured["standard_logging_object"] = self.model_call_details.get(
|
||||
"standard_logging_object"
|
||||
)
|
||||
captured["result"] = result
|
||||
|
||||
monkeypatch.setattr(LiteLLMLogging, "async_success_handler", fake_success)
|
||||
|
||||
body = CallbackLogsRequest(
|
||||
records=[
|
||||
CallbackLogRecord(
|
||||
status="success", standard_logging_payload=_sample_payload()
|
||||
)
|
||||
]
|
||||
)
|
||||
resp = await ingest_callback_logs(
|
||||
body, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
)
|
||||
assert resp.processed == 1 and resp.failed == 0
|
||||
assert captured["standard_logging_object"]["id"] == REQ_ID
|
||||
assert captured["result"]["usage"]["total_tokens"] == 42
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_failure_record_invokes_failure_handler(monkeypatch):
|
||||
captured = {}
|
||||
|
||||
async def fake_failure(
|
||||
self, exception, traceback_exception, start_time=None, end_time=None
|
||||
):
|
||||
captured["exception"] = str(exception)
|
||||
|
||||
monkeypatch.setattr(LiteLLMLogging, "async_failure_handler", fake_failure)
|
||||
|
||||
body = CallbackLogsRequest(
|
||||
records=[
|
||||
CallbackLogRecord(
|
||||
status="failure",
|
||||
standard_logging_payload=_sample_payload(),
|
||||
error="upstream exploded",
|
||||
)
|
||||
]
|
||||
)
|
||||
resp = await ingest_callback_logs(
|
||||
body, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
)
|
||||
assert resp.processed == 1 and resp.failed == 0
|
||||
assert captured["exception"] == "upstream exploded"
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_non_admin_is_rejected(monkeypatch):
|
||||
async def fake_success(self, **kwargs):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(LiteLLMLogging, "async_success_handler", fake_success)
|
||||
|
||||
body = CallbackLogsRequest(
|
||||
records=[
|
||||
CallbackLogRecord(
|
||||
status="success", standard_logging_payload=_sample_payload()
|
||||
)
|
||||
]
|
||||
)
|
||||
with pytest.raises(HTTPException) as exc_info:
|
||||
await ingest_callback_logs(
|
||||
body,
|
||||
user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.INTERNAL_USER),
|
||||
)
|
||||
assert exc_info.value.status_code == 403
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_one_bad_record_does_not_sink_the_batch(monkeypatch):
|
||||
calls = {"n": 0}
|
||||
|
||||
async def flaky_success(
|
||||
self, result=None, start_time=None, end_time=None, **kwargs
|
||||
):
|
||||
calls["n"] += 1
|
||||
if calls["n"] == 1:
|
||||
raise ValueError("boom on first record")
|
||||
|
||||
monkeypatch.setattr(LiteLLMLogging, "async_success_handler", flaky_success)
|
||||
|
||||
body = CallbackLogsRequest(
|
||||
records=[
|
||||
CallbackLogRecord(
|
||||
status="success", standard_logging_payload=_sample_payload()
|
||||
),
|
||||
CallbackLogRecord(
|
||||
status="success", standard_logging_payload=_sample_payload()
|
||||
),
|
||||
]
|
||||
)
|
||||
resp = await ingest_callback_logs(
|
||||
body, user_api_key_dict=UserAPIKeyAuth(user_role=LitellmUserRoles.PROXY_ADMIN)
|
||||
)
|
||||
assert resp.processed == 1 and resp.failed == 1
|
||||
# The failed record is reported back by index + error, not silently dropped.
|
||||
assert len(resp.failures) == 1
|
||||
assert resp.failures[0].index == 0
|
||||
assert "boom on first record" in resp.failures[0].error
|
||||
|
||||
|
||||
def test_batch_over_limit_is_rejected():
|
||||
from litellm.constants import MAX_CALLBACK_LOG_RECORDS
|
||||
from pydantic import ValidationError
|
||||
|
||||
# One over the cap must fail validation (422 at the API boundary), bounding
|
||||
# the callback/DB fan-out a single POST can trigger.
|
||||
too_many = [
|
||||
CallbackLogRecord(status="success", standard_logging_payload=_sample_payload())
|
||||
for _ in range(MAX_CALLBACK_LOG_RECORDS + 1)
|
||||
]
|
||||
with pytest.raises(ValidationError):
|
||||
CallbackLogsRequest(records=too_many)
|
||||
|
|
@ -0,0 +1,12 @@
|
|||
"""Unit tests for the combined Rust control-plane router."""
|
||||
|
||||
from litellm.proxy.rust_control_plane_endpoints.router import rust_control_plane_router
|
||||
|
||||
|
||||
def test_router_mounts_rust_consumed_endpoints():
|
||||
route_paths = {
|
||||
getattr(route, "path", None) for route in rust_control_plane_router.routes
|
||||
}
|
||||
|
||||
assert "/v1/rust_control_plane/authentication" in route_paths
|
||||
assert "/v1/rust_control_plane/logs" in route_paths
|
||||
Loading…
Add table
Reference in a new issue