feat(fireworks_ai): map litellm session id to x-session-affinity header for prompt caching (#33717)

* feat(fireworks_ai): map litellm session id to x-session-affinity header for prompt caching

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): normalize cached usage in spend logs

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(fireworks_ai): initialize chat config base class

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(fireworks_ai): normalize cached usage for spend logs

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(fireworks_ai): cover cached usage normalization

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* fix(proxy): normalize cached usage in spend logs

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

* test(fireworks_ai): cover session id precedence

Co-Authored-By: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>

---------

Co-authored-by: yuneng-jiang <yuneng@berri.ai>
Co-authored-by: Mateo Wang <277851410+mateo-berri@users.noreply.github.com>
Co-authored-by: Krrish Dholakia <krrishdholakia@berri.ai>
Co-authored-by: Devin AI <158243242+devin-ai-integration[bot]@users.noreply.github.com>
This commit is contained in:
devin-ai-integration[bot] 2026-07-17 21:33:34 -07:00 • committed by GitHub
parent 07e07e6e2b
commit b3d05bd10b
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
5 changed files with 204 additions and 3 deletions

View file

@ -48,7 +48,7 @@ from ...openai.chat.gpt_transformation import (
OpenAIChatCompletionStreamingHandler,
OpenAIGPTConfig,
)
from ..common_utils import FireworksAIException
from ..common_utils import FireworksAIMixin, FireworksAIException
def _extract_fireworks_hidden_params(payload: dict) -> dict:
@ -70,7 +70,7 @@ def _extract_fireworks_hidden_params(payload: dict) -> dict:
return {**top_level, **per_choice}
class FireworksAIConfig(OpenAIGPTConfig):
class FireworksAIConfig(FireworksAIMixin, OpenAIGPTConfig):
"""
Reference: https://docs.fireworks.ai/api-reference/post-chatcompletions
@ -114,6 +114,16 @@ class FireworksAIConfig(OpenAIGPTConfig):
prompt_truncate_len: Optional[int] = None,
context_length_exceeded_behavior: Optional[Literal["error", "truncate"]] = None,
) -> None:
OpenAIGPTConfig.__init__(
self,
frequency_penalty=frequency_penalty,
max_tokens=max_tokens,
n=n,
stop=stop,
temperature=temperature,
top_p=top_p,
response_format=response_format,
)
locals_ = locals().copy()
for key, value in locals_.items():
if key != "self" and value is not None:

View file

@ -12,6 +12,23 @@ class FireworksAIException(BaseLLMException):
pass
def get_fireworks_session_id(litellm_params: dict) -> str | None:
params = litellm_params
for key in ("litellm_session_id", "session_id"):
value = params.get(key)
if value:
return str(value)
metadata = params.get("metadata")
if isinstance(metadata, dict):
value = metadata.get("session_id")
if value:
return str(value)
value = params.get("litellm_trace_id")
if value:
return str(value)
return None
class FireworksAIMixin:
"""
Common Base Config functions across Fireworks AI Endpoints
@ -47,4 +64,9 @@ class FireworksAIMixin:
if api_key is None:
raise ValueError("FIREWORKS_API_KEY is not set")
return {"Authorization": "Bearer {}".format(api_key), **headers}
validated_headers = {"Authorization": "Bearer {}".format(api_key), **headers}
if not any(key.lower() == "x-session-affinity" for key in validated_headers):
session_id = get_fireworks_session_id(litellm_params)
if session_id:
validated_headers["x-session-affinity"] = session_id
return validated_headers

View file

@ -373,6 +373,12 @@ def get_logging_payload(kwargs, response_obj, start_time, end_time) -> SpendLogs
if isinstance(v, BaseModel):
v = v.model_dump()
additional_usage_values.update({k: v})
if "cache_read_input_tokens" not in additional_usage_values:
prompt_tokens_details = additional_usage_values.get("prompt_tokens_details")
if isinstance(prompt_tokens_details, dict):
cached_tokens = prompt_tokens_details.get("cached_tokens")
if isinstance(cached_tokens, int) and cached_tokens > 0:
additional_usage_values["cache_read_input_tokens"] = cached_tokens
clean_metadata["additional_usage_values"] = additional_usage_values
if litellm.cache is not None:

View file

@ -13,6 +13,7 @@ sys.path.insert(
from litellm import get_model_info, supports_reasoning, supports_vision
from litellm.llms.fireworks_ai.chat.transformation import FireworksAIConfig
from litellm.llms.fireworks_ai.common_utils import get_fireworks_session_id
from litellm.types.utils import (
ChatCompletionMessageToolCall,
Function,
@ -32,6 +33,105 @@ def force_local_model_cost(monkeypatch):
litellm.model_cost = get_model_cost_map(url=litellm.model_cost_map_url)
def test_validate_environment_sets_session_affinity_from_litellm_session_id():
config = FireworksAIConfig()
headers = config.validate_environment(
headers={},
model="accounts/fireworks/models/test-model",
messages=[],
optional_params={},
litellm_params={"litellm_session_id": "session-123"},
api_key="test-key",
)
assert headers["x-session-affinity"] == "session-123"
def test_validate_environment_sets_session_affinity_from_metadata_session_id():
config = FireworksAIConfig()
headers = config.validate_environment(
headers={},
model="accounts/fireworks/models/test-model",
messages=[],
optional_params={},
litellm_params={"metadata": {"session_id": "metadata-session-123"}},
api_key="test-key",
)
assert headers["x-session-affinity"] == "metadata-session-123"
def test_validate_environment_sets_session_affinity_from_session_id():
config = FireworksAIConfig()
headers = config.validate_environment(
headers={},
model="accounts/fireworks/models/test-model",
messages=[],
optional_params={},
litellm_params={"session_id": "session-id-123"},
api_key="test-key",
)
assert headers["x-session-affinity"] == "session-id-123"
def test_validate_environment_sets_session_affinity_from_trace_id():
config = FireworksAIConfig()
headers = config.validate_environment(
headers={},
model="accounts/fireworks/models/test-model",
messages=[],
optional_params={},
litellm_params={"litellm_trace_id": "trace-id-123"},
api_key="test-key",
)
assert headers["x-session-affinity"] == "trace-id-123"
def test_validate_environment_does_not_set_session_affinity_without_session_id():
config = FireworksAIConfig()
headers = config.validate_environment(
headers={},
model="accounts/fireworks/models/test-model",
messages=[],
optional_params={},
litellm_params={},
api_key="test-key",
)
assert "x-session-affinity" not in headers
def test_validate_environment_preserves_explicit_session_affinity_header():
config = FireworksAIConfig()
headers = config.validate_environment(
headers={"x-session-affinity": "explicit-session"},
model="accounts/fireworks/models/test-model",
messages=[],
optional_params={},
litellm_params={"litellm_session_id": "session-123"},
api_key="test-key",
)
assert headers["x-session-affinity"] == "explicit-session"
def test_get_fireworks_session_id_prefers_litellm_session_id_over_trace_id():
assert (
get_fireworks_session_id(
{"litellm_session_id": "session-123", "litellm_trace_id": "trace-123"}
)
== "session-123"
)
def test_handle_message_content_with_tool_calls():
config = FireworksAIConfig()
message = Message(

View file

@ -46,6 +46,69 @@ from litellm.types.utils import (
)
def _get_additional_usage_values_for_usage(usage: litellm.Usage) -> dict:
payload = get_logging_payload(
kwargs={
"model": "gpt-4o-mini",
"litellm_params": {"metadata": {"user_api_key": "test-key"}},
},
response_obj=litellm.ModelResponse(
id="chatcmpl-test",
choices=[],
usage=usage,
),
start_time=datetime.datetime.now(timezone.utc),
end_time=datetime.datetime.now(timezone.utc),
)
metadata = json.loads(payload["metadata"])
return metadata["additional_usage_values"]
def test_get_logging_payload_maps_openai_cached_tokens_to_cache_read_input_tokens():
additional_usage_values = _get_additional_usage_values_for_usage(
litellm.Usage(
prompt_tokens=10,
completion_tokens=2,
total_tokens=12,
prompt_tokens_details={"cached_tokens": 123},
)
)
assert additional_usage_values["cache_read_input_tokens"] == 123
assert additional_usage_values["prompt_tokens_details"]["cached_tokens"] == 123
def test_get_logging_payload_preserves_anthropic_cache_read_input_tokens():
additional_usage_values = _get_additional_usage_values_for_usage(
litellm.Usage(
prompt_tokens=10,
completion_tokens=2,
total_tokens=12,
prompt_tokens_details={"cached_tokens": 123},
cache_read_input_tokens=456,
)
)
assert additional_usage_values["cache_read_input_tokens"] == 456
@pytest.mark.parametrize(
"prompt_tokens_details",
[None, {"cached_tokens": 0}],
)
def test_get_logging_payload_does_not_map_missing_or_zero_cached_tokens(prompt_tokens_details):
additional_usage_values = _get_additional_usage_values_for_usage(
litellm.Usage(
prompt_tokens=10,
completion_tokens=2,
total_tokens=12,
prompt_tokens_details=prompt_tokens_details,
)
)
assert "cache_read_input_tokens" not in additional_usage_values
def test_sanitize_request_body_for_spend_logs_payload_basic():
request_body = {
"messages": [{"role": "user", "content": "Hello, how are you?"}],