diff --git a/docs/my-website/docs/completion/token_usage.md b/docs/my-website/docs/completion/token_usage.md
index 807ccfd91ec..0bec6b3f902 100644
--- a/docs/my-website/docs/completion/token_usage.md
+++ b/docs/my-website/docs/completion/token_usage.md
@@ -1,7 +1,21 @@
# Completion Token Usage & Cost
By default LiteLLM returns token usage in all completion requests ([See here](https://litellm.readthedocs.io/en/latest/output/))
-However, we also expose some helper functions + **[NEW]** an API to calculate token usage across providers:
+LiteLLM returns `response_cost` in all calls.
+
+```python
+from litellm import completion
+
+response = litellm.completion(
+ model="gpt-3.5-turbo",
+ messages=[{"role": "user", "content": "Hey, how's it going?"}],
+ mock_response="Hello world",
+ )
+
+print(response._hidden_params["response_cost"])
+```
+
+LiteLLM also exposes some helper functions:
- `encode`: This encodes the text passed in, using the model-specific tokenizer. [**Jump to code**](#1-encode)
@@ -23,7 +37,7 @@ However, we also expose some helper functions + **[NEW]** an API to calculate to
- `api.litellm.ai`: Live token + price count across [all supported models](https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json). [**Jump to code**](#10-apilitellmai)
-📣 This is a community maintained list. Contributions are welcome! ❤️
+📣 [This is a community maintained list](https://github.com/BerriAI/litellm/blob/main/model_prices_and_context_window.json). Contributions are welcome! ❤️
## Example Usage
diff --git a/docs/my-website/docs/proxy/cost_tracking.md b/docs/my-website/docs/proxy/cost_tracking.md
index f01e1042e35..3ccf8f383aa 100644
--- a/docs/my-website/docs/proxy/cost_tracking.md
+++ b/docs/my-website/docs/proxy/cost_tracking.md
@@ -114,6 +114,14 @@ print(response)
**Step3 - Verify Spend Tracked**
That's IT. Now Verify your spend was tracked
+
+
+
+
+
+
+
+
The following spend gets tracked in Table `LiteLLM_SpendLogs`
```json
@@ -144,6 +152,10 @@ Use the `/global/spend/report` endpoint to get daily spend report per
- team
- customer [this is `user` passed to `/chat/completions` request](#how-to-track-spend-with-litellm)
+
+
+
+
diff --git a/docs/my-website/img/response_cost_img.png b/docs/my-website/img/response_cost_img.png
new file mode 100644
index 00000000000..9f466b3fbbc
Binary files /dev/null and b/docs/my-website/img/response_cost_img.png differ
diff --git a/litellm/cost_calculator.py b/litellm/cost_calculator.py
index 2504a95f145..9a7df7ebefb 100644
--- a/litellm/cost_calculator.py
+++ b/litellm/cost_calculator.py
@@ -572,7 +572,6 @@ def completion_cost(
completion_string = litellm.utils.get_response_string(
response_obj=completion_response
)
-
completion_characters = litellm.utils._count_characters(
text=completion_string
)
@@ -610,7 +609,7 @@ def response_cost_calculator(
TextCompletionResponse,
],
model: str,
- custom_llm_provider: str,
+ custom_llm_provider: Optional[str],
call_type: Literal[
"embedding",
"aembedding",
@@ -632,6 +631,10 @@ def response_cost_calculator(
base_model: Optional[str] = None,
custom_pricing: Optional[bool] = None,
) -> Optional[float]:
+ """
+ Returns
+ - float or None: cost of response OR none if error.
+ """
try:
response_cost: float = 0.0
if cache_hit is not None and cache_hit is True:
diff --git a/litellm/proxy/proxy_server.py b/litellm/proxy/proxy_server.py
index 710b3d11d81..553c8a48cb4 100644
--- a/litellm/proxy/proxy_server.py
+++ b/litellm/proxy/proxy_server.py
@@ -433,6 +433,7 @@ def get_custom_headers(
api_base: Optional[str] = None,
version: Optional[str] = None,
model_region: Optional[str] = None,
+ response_cost: Optional[Union[float, str]] = None,
fastest_response_batch_completion: Optional[bool] = None,
**kwargs,
) -> dict:
@@ -443,6 +444,7 @@ def get_custom_headers(
"x-litellm-model-api-base": api_base,
"x-litellm-version": version,
"x-litellm-model-region": model_region,
+ "x-litellm-response-cost": str(response_cost),
"x-litellm-key-tpm-limit": str(user_api_key_dict.tpm_limit),
"x-litellm-key-rpm-limit": str(user_api_key_dict.rpm_limit),
"x-litellm-fastest_response_batch_completion": (
@@ -3048,6 +3050,7 @@ async def chat_completion(
model_id = hidden_params.get("model_id", None) or ""
cache_key = hidden_params.get("cache_key", None) or ""
api_base = hidden_params.get("api_base", None) or ""
+ response_cost = hidden_params.get("response_cost", None) or ""
fastest_response_batch_completion = hidden_params.get(
"fastest_response_batch_completion", None
)
@@ -3066,6 +3069,7 @@ async def chat_completion(
cache_key=cache_key,
api_base=api_base,
version=version,
+ response_cost=response_cost,
model_region=getattr(user_api_key_dict, "allowed_model_region", ""),
fastest_response_batch_completion=fastest_response_batch_completion,
)
@@ -3095,6 +3099,7 @@ async def chat_completion(
cache_key=cache_key,
api_base=api_base,
version=version,
+ response_cost=response_cost,
model_region=getattr(user_api_key_dict, "allowed_model_region", ""),
fastest_response_batch_completion=fastest_response_batch_completion,
**additional_headers,
@@ -3290,6 +3295,7 @@ async def completion(
model_id = hidden_params.get("model_id", None) or ""
cache_key = hidden_params.get("cache_key", None) or ""
api_base = hidden_params.get("api_base", None) or ""
+ response_cost = hidden_params.get("response_cost", None) or ""
### ALERTING ###
data["litellm_status"] = "success" # used for alerting
@@ -3304,6 +3310,7 @@ async def completion(
cache_key=cache_key,
api_base=api_base,
version=version,
+ response_cost=response_cost,
)
selected_data_generator = select_data_generator(
response=response,
@@ -3323,6 +3330,7 @@ async def completion(
cache_key=cache_key,
api_base=api_base,
version=version,
+ response_cost=response_cost,
)
)
@@ -3527,6 +3535,7 @@ async def embeddings(
model_id = hidden_params.get("model_id", None) or ""
cache_key = hidden_params.get("cache_key", None) or ""
api_base = hidden_params.get("api_base", None) or ""
+ response_cost = hidden_params.get("response_cost", None) or ""
fastapi_response.headers.update(
get_custom_headers(
@@ -3535,6 +3544,7 @@ async def embeddings(
cache_key=cache_key,
api_base=api_base,
version=version,
+ response_cost=response_cost,
model_region=getattr(user_api_key_dict, "allowed_model_region", ""),
)
)
@@ -3676,6 +3686,7 @@ async def image_generation(
model_id = hidden_params.get("model_id", None) or ""
cache_key = hidden_params.get("cache_key", None) or ""
api_base = hidden_params.get("api_base", None) or ""
+ response_cost = hidden_params.get("response_cost", None) or ""
fastapi_response.headers.update(
get_custom_headers(
@@ -3684,6 +3695,7 @@ async def image_generation(
cache_key=cache_key,
api_base=api_base,
version=version,
+ response_cost=response_cost,
model_region=getattr(user_api_key_dict, "allowed_model_region", ""),
)
)
@@ -3812,6 +3824,7 @@ async def audio_speech(
model_id = hidden_params.get("model_id", None) or ""
cache_key = hidden_params.get("cache_key", None) or ""
api_base = hidden_params.get("api_base", None) or ""
+ response_cost = hidden_params.get("response_cost", None) or ""
# Printing each chunk size
async def generate(_response: HttpxBinaryResponseContent):
@@ -3825,6 +3838,7 @@ async def audio_speech(
cache_key=cache_key,
api_base=api_base,
version=version,
+ response_cost=response_cost,
model_region=getattr(user_api_key_dict, "allowed_model_region", ""),
fastest_response_batch_completion=None,
)
@@ -3976,6 +3990,7 @@ async def audio_transcriptions(
model_id = hidden_params.get("model_id", None) or ""
cache_key = hidden_params.get("cache_key", None) or ""
api_base = hidden_params.get("api_base", None) or ""
+ response_cost = hidden_params.get("response_cost", None) or ""
fastapi_response.headers.update(
get_custom_headers(
@@ -3984,6 +3999,7 @@ async def audio_transcriptions(
cache_key=cache_key,
api_base=api_base,
version=version,
+ response_cost=response_cost,
model_region=getattr(user_api_key_dict, "allowed_model_region", ""),
)
)
diff --git a/litellm/tests/test_completion_cost.py b/litellm/tests/test_completion_cost.py
index 3a65f729423..bffb68e0e5d 100644
--- a/litellm/tests/test_completion_cost.py
+++ b/litellm/tests/test_completion_cost.py
@@ -712,9 +712,30 @@ def test_vertex_ai_claude_completion_cost():
assert cost == predicted_cost
+
+@pytest.mark.parametrize("sync_mode", [True, False])
+@pytest.mark.asyncio
+async def test_completion_cost_hidden_params(sync_mode):
+ if sync_mode:
+ response = litellm.completion(
+ model="gpt-3.5-turbo",
+ messages=[{"role": "user", "content": "Hey, how's it going?"}],
+ mock_response="Hello world",
+ )
+ else:
+ response = await litellm.acompletion(
+ model="gpt-3.5-turbo",
+ messages=[{"role": "user", "content": "Hey, how's it going?"}],
+ mock_response="Hello world",
+ )
+
+ assert "response_cost" in response._hidden_params
+ assert isinstance(response._hidden_params["response_cost"], float)
+
def test_vertex_ai_gemini_predict_cost():
model = "gemini-1.5-flash"
messages = [{"role": "user", "content": "Hey, hows it going???"}]
predictive_cost = completion_cost(model=model, messages=messages)
assert predictive_cost > 0
+
diff --git a/litellm/utils.py b/litellm/utils.py
index dbc988bb978..53f5f984864 100644
--- a/litellm/utils.py
+++ b/litellm/utils.py
@@ -899,6 +899,17 @@ def client(original_function):
model=model,
optional_params=getattr(logging_obj, "optional_params", {}),
)
+ result._hidden_params["response_cost"] = (
+ litellm.response_cost_calculator(
+ response_object=result,
+ model=getattr(logging_obj, "model", ""),
+ custom_llm_provider=getattr(
+ logging_obj, "custom_llm_provider", None
+ ),
+ call_type=getattr(logging_obj, "call_type", "completion"),
+ optional_params=getattr(logging_obj, "optional_params", {}),
+ )
+ )
result._response_ms = (
end_time - start_time
).total_seconds() * 1000 # return response latency in ms like openai
@@ -1292,6 +1303,17 @@ def client(original_function):
model=model,
optional_params=kwargs,
)
+ result._hidden_params["response_cost"] = (
+ litellm.response_cost_calculator(
+ response_object=result,
+ model=getattr(logging_obj, "model", ""),
+ custom_llm_provider=getattr(
+ logging_obj, "custom_llm_provider", None
+ ),
+ call_type=getattr(logging_obj, "call_type", "completion"),
+ optional_params=getattr(logging_obj, "optional_params", {}),
+ )
+ )
if (
isinstance(result, ModelResponse)
or isinstance(result, EmbeddingResponse)