mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-09 03:18:44 +00:00
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:
parent
07e07e6e2b
commit
b3d05bd10b
5 changed files with 204 additions and 3 deletions
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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?"}],
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue