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:
yucheng-berri 2026-10-08 12:48:16 -07:00 • committed by GitHub
parent 3571b5dc72
commit 29f48e62bf
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
13 changed files with 304 additions and 2 deletions

View file

@ -67,6 +67,7 @@ OPTIONAL_KWARGS_KEYS: Final = (
"itpm",
"otpm",
"use_xai_oauth",
"fireworks_forward_user_id",
PROVIDER_AFFINITY_HEADER_KWARG_KEY,
}
)

View file

@ -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,

View file

@ -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")

View file

@ -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,

View file

@ -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
},

View file

@ -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",

View file

@ -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)

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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

View file

@ -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",

View file

@ -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 */