mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
feat(fireworks_ai): forward the LiteLLM user id as user behind fireworks_forward_user_id (#45265)
* feat(fireworks_ai): forward the LiteLLM user id as user behind fireworks_forward_user_id * fix(fireworks_ai): keep the forwarded user id when the caller sends extra_body.user
This commit is contained in:
parent
3571b5dc72
commit
29f48e62bf
13 changed files with 304 additions and 2 deletions
|
|
@ -67,6 +67,7 @@ OPTIONAL_KWARGS_KEYS: Final = (
|
|||
"itpm",
|
||||
"otpm",
|
||||
"use_xai_oauth",
|
||||
"fireworks_forward_user_id",
|
||||
PROVIDER_AFFINITY_HEADER_KWARG_KEY,
|
||||
}
|
||||
)
|
||||
|
|
|
|||
|
|
@ -43,7 +43,9 @@ from ..common_utils import (
|
|||
FIREROUTER,
|
||||
FireworksAIException,
|
||||
FireworksAIMixin,
|
||||
get_fireworks_forwarded_user_id,
|
||||
resolve_fireworks_resource_name,
|
||||
without_caller_user,
|
||||
)
|
||||
|
||||
if TYPE_CHECKING:
|
||||
|
|
@ -674,13 +676,24 @@ class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig):
|
|||
**stream_options,
|
||||
"include_usage": True,
|
||||
}
|
||||
return super().transform_request(
|
||||
request: Final = super().transform_request(
|
||||
model=resolved_model,
|
||||
messages=messages,
|
||||
optional_params=optional_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
forwarded_user_id: Final = get_fireworks_forwarded_user_id(litellm_params)
|
||||
return request if forwarded_user_id is None else {**request, "user": forwarded_user_id}
|
||||
|
||||
def transform_extra_body(
|
||||
self,
|
||||
extra_body: Mapping[str, object],
|
||||
request: Mapping[str, object],
|
||||
model: str,
|
||||
litellm_params: Mapping[str, object],
|
||||
) -> Mapping[str, object]:
|
||||
return without_caller_user(extra_body, get_fireworks_forwarded_user_id(litellm_params))
|
||||
|
||||
def _handle_message_content_with_tool_calls(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -49,6 +49,33 @@ def with_fireworks_session_affinity(
|
|||
return MappingProxyType({**headers, "x-session-affinity": session_id})
|
||||
|
||||
|
||||
FIREWORKS_FORWARD_USER_ID_PARAM: Final = "fireworks_forward_user_id"
|
||||
|
||||
|
||||
def _authenticated_user_id(metadata: object) -> str | None:
|
||||
user_id: Final = metadata.get("user_api_key_user_id") if isinstance(metadata, Mapping) else None
|
||||
return user_id if isinstance(user_id, str) and user_id else None
|
||||
|
||||
|
||||
def get_fireworks_forwarded_user_id(litellm_params: Mapping[str, object]) -> str | None:
|
||||
if litellm_params.get(FIREWORKS_FORWARD_USER_ID_PARAM) is not True:
|
||||
return None
|
||||
return next(
|
||||
(
|
||||
user_id
|
||||
for key in ("metadata", "litellm_metadata")
|
||||
if (user_id := _authenticated_user_id(litellm_params.get(key))) is not None
|
||||
),
|
||||
None,
|
||||
)
|
||||
|
||||
|
||||
def without_caller_user(extra_body: Mapping[str, object], forwarded_user_id: str | None) -> Mapping[str, object]:
|
||||
if forwarded_user_id is None:
|
||||
return extra_body
|
||||
return MappingProxyType({key: value for key, value in extra_body.items() if key != "user"})
|
||||
|
||||
|
||||
def resolve_fireworks_api_key(api_key: str | None) -> str | None:
|
||||
return api_key or (
|
||||
get_secret_str("FIREWORKS_API_KEY")
|
||||
|
|
|
|||
|
|
@ -7,9 +7,12 @@ import httpx
|
|||
from openai.types.responses import EasyInputMessageParam, ResponseInputContentParam, ResponseInputItemParam
|
||||
|
||||
from litellm.llms.fireworks_ai.common_utils import (
|
||||
FIREWORKS_FORWARD_USER_ID_PARAM,
|
||||
get_fireworks_forwarded_user_id,
|
||||
resolve_fireworks_api_key,
|
||||
resolve_fireworks_resource_name,
|
||||
with_fireworks_session_affinity,
|
||||
without_caller_user,
|
||||
)
|
||||
from litellm.llms.openai.responses.transformation import OpenAIResponsesAPIConfig
|
||||
from litellm.secret_managers.main import get_secret_str
|
||||
|
|
@ -31,6 +34,16 @@ def _session_params(litellm_params: GenericLiteLLMParams) -> Mapping[str, object
|
|||
)
|
||||
|
||||
|
||||
def _forwarded_user_params(litellm_params: GenericLiteLLMParams) -> Mapping[str, object]:
|
||||
extras: Final[Mapping[str, object]] = litellm_params.model_extra or MappingProxyType({})
|
||||
return MappingProxyType(
|
||||
{
|
||||
FIREWORKS_FORWARD_USER_ID_PARAM: litellm_params.fireworks_forward_user_id,
|
||||
"litellm_metadata": extras.get("litellm_metadata"),
|
||||
}
|
||||
)
|
||||
|
||||
|
||||
_INSTRUCTION_ROLES: Final = frozenset({"system", "developer"})
|
||||
|
||||
|
||||
|
|
@ -163,13 +176,24 @@ class FireworksAIResponsesAPIConfig(OpenAIResponsesAPIConfig):
|
|||
*instruction_entries,
|
||||
)
|
||||
}
|
||||
return super().transform_responses_api_request(
|
||||
request: Final = super().transform_responses_api_request(
|
||||
model=resolve_fireworks_resource_name(model),
|
||||
input=folded_input,
|
||||
response_api_optional_request_params=folded_params,
|
||||
litellm_params=litellm_params,
|
||||
headers=headers,
|
||||
)
|
||||
forwarded_user_id: Final = get_fireworks_forwarded_user_id(_forwarded_user_params(litellm_params))
|
||||
return request if forwarded_user_id is None else {**request, "user": forwarded_user_id}
|
||||
|
||||
def transform_extra_body(
|
||||
self,
|
||||
extra_body: Mapping[str, object],
|
||||
request: Mapping[str, object],
|
||||
model: str,
|
||||
litellm_params: GenericLiteLLMParams,
|
||||
) -> Mapping[str, object]:
|
||||
return without_caller_user(extra_body, get_fireworks_forwarded_user_id(_forwarded_user_params(litellm_params)))
|
||||
|
||||
def transform_delete_response_api_response(
|
||||
self,
|
||||
|
|
|
|||
|
|
@ -5737,6 +5737,7 @@ def completion(
|
|||
*ANTHROPIC_WIF_KWARGS_KEYS,
|
||||
*OPENAI_WIF_KWARGS_KEYS,
|
||||
PROVIDER_AFFINITY_HEADER_KWARG_KEY,
|
||||
"fireworks_forward_user_id",
|
||||
)
|
||||
if key in kwargs
|
||||
},
|
||||
|
|
|
|||
|
|
@ -431,6 +431,9 @@ _BANNED_REQUEST_BODY_PARAMS: Final[tuple[str, ...]] = (
|
|||
# so a caller-supplied value picks a transport and a callback surface the
|
||||
# admin did not choose.
|
||||
"rust",
|
||||
# Deployment opt-in: a caller-supplied false would switch off identity
|
||||
# forwarding and let the caller choose the `user` Fireworks sees.
|
||||
"fireworks_forward_user_id",
|
||||
# SDK-only field; also rejected outright in is_request_body_safe.
|
||||
"model_list",
|
||||
"vertex_ai_credentials",
|
||||
|
|
|
|||
|
|
@ -87,6 +87,7 @@ class ProviderConnection:
|
|||
litellm_credential_name: str | None = None
|
||||
configurable_clientside_auth_params: "Sequence[str | ConfigurableClientsideParamsCustomAuth] | None" = None
|
||||
use_xai_oauth: bool | None = None
|
||||
fireworks_forward_user_id: bool | None = None
|
||||
|
||||
|
||||
@dataclass(frozen=True, slots=True, kw_only=True)
|
||||
|
|
|
|||
|
|
@ -504,6 +504,10 @@ class GenericLiteLLMParams(CredentialLiteLLMParams, CustomPricingLiteLLMParams):
|
|||
default=False,
|
||||
description="Use stored xAI OAuth credentials when no xAI API key is configured.",
|
||||
)
|
||||
fireworks_forward_user_id: bool | None = Field(
|
||||
default=None,
|
||||
description="Send the LiteLLM user id of the calling key as the `user` field on Fireworks AI chat, responses and messages requests.",
|
||||
)
|
||||
model_config = ConfigDict(extra="allow", arbitrary_types_allowed=True)
|
||||
merge_reasoning_content_in_choices: bool | None = False
|
||||
model_info: dict | None = None
|
||||
|
|
|
|||
|
|
@ -2008,3 +2008,126 @@ VISION_MODEL = next(
|
|||
for key, info in litellm.model_cost.items()
|
||||
if key.startswith("fireworks_ai/accounts/fireworks/models/") and info.get("supports_vision") is True
|
||||
)
|
||||
|
||||
|
||||
def _fireworks_chat_client() -> MagicMock:
|
||||
body: Final = {
|
||||
"id": "chat-user-attribution",
|
||||
"object": "chat.completion",
|
||||
"created": 1,
|
||||
"model": "accounts/fireworks/models/kimi-k3",
|
||||
"choices": [{"index": 0, "message": {"role": "assistant", "content": "ok"}, "finish_reason": "stop"}],
|
||||
"usage": {"prompt_tokens": 1, "completion_tokens": 1, "total_tokens": 2},
|
||||
}
|
||||
raw_response: Final = MagicMock()
|
||||
raw_response.status_code = 200
|
||||
raw_response.headers = {}
|
||||
raw_response.text = json.dumps(body)
|
||||
raw_response.json = lambda: body
|
||||
client: Final = MagicMock(spec=HTTPHandler)
|
||||
client.post.return_value = raw_response
|
||||
return client
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"call_kwargs, expected_user",
|
||||
[
|
||||
pytest.param(
|
||||
{"fireworks_forward_user_id": True, "metadata": {"user_api_key_user_id": "dev-alice"}},
|
||||
"dev-alice",
|
||||
id="opted-in-sends-litellm-user-id",
|
||||
),
|
||||
pytest.param(
|
||||
{"fireworks_forward_user_id": True, "metadata": {"user_api_key_user_id": "dev-alice"}, "user": "caller"},
|
||||
"dev-alice",
|
||||
id="litellm-user-id-replaces-caller-user",
|
||||
),
|
||||
pytest.param(
|
||||
{"fireworks_forward_user_id": True, "litellm_metadata": {"user_api_key_user_id": "dev-bob"}},
|
||||
"dev-bob",
|
||||
id="reads-litellm-metadata",
|
||||
),
|
||||
pytest.param(
|
||||
{
|
||||
"fireworks_forward_user_id": True,
|
||||
"metadata": {"tags": ["caller-tag"]},
|
||||
"litellm_metadata": {"user_api_key_user_id": "dev-bob"},
|
||||
},
|
||||
"dev-bob",
|
||||
id="reads-litellm-metadata-next-to-caller-metadata",
|
||||
),
|
||||
pytest.param(
|
||||
{"fireworks_forward_user_id": True, "metadata": {"user_api_key_user_id": ""}, "user": "caller"},
|
||||
"caller",
|
||||
id="empty-litellm-user-id-keeps-caller-user",
|
||||
),
|
||||
pytest.param(
|
||||
{"fireworks_forward_user_id": True, "metadata": {"user_api_key_user_id": None}, "user": "caller"},
|
||||
"caller",
|
||||
id="no-litellm-user-id-keeps-caller-user",
|
||||
),
|
||||
pytest.param(
|
||||
{"fireworks_forward_user_id": True, "metadata": {"user_api_key_user_id": None}},
|
||||
None,
|
||||
id="no-litellm-user-id-sends-no-user",
|
||||
),
|
||||
pytest.param(
|
||||
{"metadata": {"user_api_key_user_id": "dev-alice"}, "user": "caller"},
|
||||
"caller",
|
||||
id="not-opted-in-keeps-caller-user",
|
||||
),
|
||||
pytest.param(
|
||||
{"metadata": {"user_api_key_user_id": "dev-alice"}},
|
||||
None,
|
||||
id="not-opted-in-sends-no-user",
|
||||
),
|
||||
pytest.param(
|
||||
{"fireworks_forward_user_id": "true", "metadata": {"user_api_key_user_id": "dev-alice"}},
|
||||
None,
|
||||
id="non-bool-flag-is-off",
|
||||
),
|
||||
pytest.param(
|
||||
{
|
||||
"fireworks_forward_user_id": True,
|
||||
"metadata": {"user_api_key_user_id": "dev-alice"},
|
||||
"extra_body": {"user": "dev-bob", "prompt_cache_max_len": 1},
|
||||
},
|
||||
"dev-alice",
|
||||
id="litellm-user-id-replaces-extra-body-user",
|
||||
),
|
||||
pytest.param(
|
||||
{"metadata": {"user_api_key_user_id": "dev-alice"}, "extra_body": {"user": "dev-bob"}},
|
||||
"dev-bob",
|
||||
id="not-opted-in-keeps-extra-body-user",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_completion_forwards_litellm_user_id_as_user(call_kwargs: dict[str, object], expected_user: str | None) -> None:
|
||||
client: Final = _fireworks_chat_client()
|
||||
litellm.completion(
|
||||
model="fireworks_ai/accounts/fireworks/models/kimi-k3",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
api_key="fw-test-key",
|
||||
client=client,
|
||||
**call_kwargs,
|
||||
)
|
||||
request_body: Final = json.loads(client.post.call_args.kwargs["data"])
|
||||
assert request_body.get("user") == expected_user
|
||||
assert "fireworks_forward_user_id" not in request_body
|
||||
|
||||
|
||||
def test_completion_forwards_litellm_user_id_when_streaming() -> None:
|
||||
client: Final = _fireworks_chat_client()
|
||||
client.post.return_value.iter_lines = lambda: iter(())
|
||||
litellm.completion(
|
||||
model="fireworks_ai/accounts/fireworks/models/kimi-k3",
|
||||
messages=[{"role": "user", "content": "hi"}],
|
||||
api_key="fw-test-key",
|
||||
client=client,
|
||||
stream=True,
|
||||
fireworks_forward_user_id=True,
|
||||
metadata={"user_api_key_user_id": "dev-alice"},
|
||||
)
|
||||
request_body: Final = json.loads(client.post.call_args.kwargs["data"])
|
||||
assert request_body["user"] == "dev-alice"
|
||||
assert request_body["stream"] is True
|
||||
|
|
|
|||
|
|
@ -608,3 +608,62 @@ def test_streaming_responses_call_hits_native_endpoint_and_yields_every_firework
|
|||
assert tuple(event.type for event in received) == tuple(event["type"] for event in FIREWORKS_SSE_EVENTS)
|
||||
assert "".join(event.delta for event in received if event.type == "response.output_text.delta") == "pong"
|
||||
assert received[-1].response.usage.output_tokens == 89
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"call_kwargs, expected_user",
|
||||
[
|
||||
pytest.param(
|
||||
{"fireworks_forward_user_id": True, "litellm_metadata": {"user_api_key_user_id": "dev-alice"}},
|
||||
"dev-alice",
|
||||
id="opted-in-sends-litellm-user-id",
|
||||
),
|
||||
pytest.param(
|
||||
{
|
||||
"fireworks_forward_user_id": True,
|
||||
"litellm_metadata": {"user_api_key_user_id": "dev-alice"},
|
||||
"user": "caller",
|
||||
},
|
||||
"dev-alice",
|
||||
id="litellm-user-id-replaces-caller-user",
|
||||
),
|
||||
pytest.param(
|
||||
{"fireworks_forward_user_id": True, "metadata": {"user_api_key_user_id": "dev-alice"}},
|
||||
None,
|
||||
id="ignores-caller-responses-metadata",
|
||||
),
|
||||
pytest.param(
|
||||
{"fireworks_forward_user_id": True, "user": "caller"},
|
||||
"caller",
|
||||
id="no-litellm-user-id-keeps-caller-user",
|
||||
),
|
||||
pytest.param(
|
||||
{"litellm_metadata": {"user_api_key_user_id": "dev-alice"}},
|
||||
None,
|
||||
id="not-opted-in-sends-no-user",
|
||||
),
|
||||
pytest.param(
|
||||
{
|
||||
"fireworks_forward_user_id": True,
|
||||
"litellm_metadata": {"user_api_key_user_id": "dev-alice"},
|
||||
"extra_body": {"user": "dev-bob"},
|
||||
},
|
||||
"dev-alice",
|
||||
id="litellm-user-id-replaces-extra-body-user",
|
||||
),
|
||||
pytest.param(
|
||||
{"litellm_metadata": {"user_api_key_user_id": "dev-alice"}, "extra_body": {"user": "dev-bob"}},
|
||||
"dev-bob",
|
||||
id="not-opted-in-keeps-extra-body-user",
|
||||
),
|
||||
],
|
||||
)
|
||||
def test_responses_call_forwards_litellm_user_id_as_user(
|
||||
call_kwargs: Mapping[str, object], expected_user: str | None
|
||||
) -> None:
|
||||
client: Final = _mock_http_client(_fireworks_response("accounts/fireworks/models/kimi-k3"))
|
||||
with patch(HTTPX_CLIENT_FACTORY, return_value=client):
|
||||
litellm.responses(model="fireworks_ai/kimi-k3", input="hi", api_key="fw-test-key", **call_kwargs)
|
||||
_, _, body = _sent_request(client)
|
||||
assert body.get("user") == expected_user
|
||||
assert "fireworks_forward_user_id" not in body
|
||||
|
|
|
|||
|
|
@ -4,6 +4,7 @@ Unit tests for auth_utils functions related to rate limiting and customer ID ext
|
|||
|
||||
import base64
|
||||
import logging
|
||||
from collections.abc import Callable
|
||||
from typing import Final, Optional
|
||||
from unittest.mock import MagicMock, patch
|
||||
|
||||
|
|
@ -2877,6 +2878,40 @@ class TestIsRequestBodySafeBlocksClaudePlatformWorkspaceOverride:
|
|||
is True
|
||||
)
|
||||
|
||||
|
||||
class TestIsRequestBodySafeBlocksFireworksForwardUserId:
|
||||
@pytest.mark.parametrize("value", [True, False])
|
||||
@pytest.mark.parametrize(
|
||||
"body_for",
|
||||
[
|
||||
pytest.param(lambda value: {"fireworks_forward_user_id": value}, id="root"),
|
||||
pytest.param(lambda value: {"extra_body": {"fireworks_forward_user_id": value}}, id="extra_body"),
|
||||
pytest.param(lambda value: {"metadata": {"fireworks_forward_user_id": value}}, id="metadata"),
|
||||
],
|
||||
)
|
||||
def test_fireworks_forward_user_id_in_request_body_is_rejected(
|
||||
self, body_for: Callable[[bool], dict[str, object]], value: bool
|
||||
) -> None:
|
||||
with pytest.raises(ValueError, match="fireworks_forward_user_id"):
|
||||
is_request_body_safe(
|
||||
request_body={"model": "fireworks-model", "user": "someone-else", **body_for(value)},
|
||||
general_settings={},
|
||||
llm_router=None,
|
||||
model="fireworks-model",
|
||||
)
|
||||
|
||||
def test_admin_opt_in_proxy_wide_allows_fireworks_forward_user_id(self) -> None:
|
||||
assert (
|
||||
is_request_body_safe(
|
||||
request_body={"model": "fireworks-model", "fireworks_forward_user_id": False},
|
||||
general_settings={"allow_client_side_credentials": True},
|
||||
llm_router=None,
|
||||
model="fireworks-model",
|
||||
)
|
||||
is True
|
||||
)
|
||||
|
||||
|
||||
class TestIsRequestBodySafeBlocksRustOptIn:
|
||||
"""``rust`` hands the whole call to the Rust core, which signs and sends
|
||||
with its own HTTP client rather than the one the deployment configured, and
|
||||
|
|
|
|||
|
|
@ -85,6 +85,7 @@ CONNECTION_NAMES: Final = (
|
|||
"litellm_credential_name",
|
||||
"configurable_clientside_auth_params",
|
||||
"use_xai_oauth",
|
||||
"fireworks_forward_user_id",
|
||||
"aws_batch_role_arn",
|
||||
"s3_bucket_name",
|
||||
"s3_region_name",
|
||||
|
|
|
|||
10
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
10
ui/litellm-dashboard/src/lib/http/schema.d.ts
generated
vendored
|
|
@ -35997,6 +35997,11 @@ export interface components {
|
|||
default_api_key_tpm_limit?: number | null;
|
||||
/** Drop Params */
|
||||
drop_params?: boolean | string | null;
|
||||
/**
|
||||
* Fireworks Forward User Id
|
||||
* @description Send the LiteLLM user id of the calling key as the `user` field on Fireworks AI chat, responses and messages requests.
|
||||
*/
|
||||
fireworks_forward_user_id?: boolean | null;
|
||||
/** Gcs Bucket Name */
|
||||
gcs_bucket_name?: string | null;
|
||||
/** Google Maps Grounding Cost Per Query */
|
||||
|
|
@ -51194,6 +51199,11 @@ export interface components {
|
|||
default_api_key_tpm_limit?: number | null;
|
||||
/** Drop Params */
|
||||
drop_params?: boolean | string | null;
|
||||
/**
|
||||
* Fireworks Forward User Id
|
||||
* @description Send the LiteLLM user id of the calling key as the `user` field on Fireworks AI chat, responses and messages requests.
|
||||
*/
|
||||
fireworks_forward_user_id?: boolean | null;
|
||||
/** Gcs Bucket Name */
|
||||
gcs_bucket_name?: string | null;
|
||||
/** Google Maps Grounding Cost Per Query */
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue