mirror of
https://github.com/BerriAI/litellm.git
synced 2026-10-10 03:28:53 +00:00
Update mlflow logger usage span attributes (#13561)
* test: sync mlflow request tags * fix test
This commit is contained in:
parent
aaf9c38a10
commit
f6e53deacd
2 changed files with 42 additions and 5 deletions
|
|
@ -189,9 +189,9 @@ class MlflowLogger(CustomLogger):
|
|||
{
|
||||
"api_base": standard_obj.get("api_base"),
|
||||
"cache_hit": standard_obj.get("cache_hit"),
|
||||
"usage": {
|
||||
"completion_tokens": standard_obj.get("completion_tokens"),
|
||||
"prompt_tokens": standard_obj.get("prompt_tokens"),
|
||||
"mlflow.chat.tokenUsage": {
|
||||
"input_tokens": standard_obj.get("prompt_tokens"),
|
||||
"output_tokens": standard_obj.get("completion_tokens"),
|
||||
"total_tokens": standard_obj.get("total_tokens"),
|
||||
},
|
||||
"raw_llm_response": standard_obj.get("response"),
|
||||
|
|
|
|||
|
|
@ -71,5 +71,42 @@ async def test_mlflow_request_tags_functionality():
|
|||
tags_param = call_args.kwargs.get('tags', {})
|
||||
expected_tags = {"tag1": "", "tag2": "", "production": ""}
|
||||
assert tags_param == expected_tags, f"Expected tags {expected_tags}, got {tags_param}"
|
||||
|
||||
print("✅ Request tags properly transformed and passed to MLflow trace")
|
||||
|
||||
|
||||
|
||||
def test_mlflow_token_usage_attribute_structure():
|
||||
"""Ensure token usage attributes are formatted with mlflow.chat.tokenUsage."""
|
||||
|
||||
mock_mlflow_tracking = MagicMock()
|
||||
mock_mlflow_tracking.MlflowClient = MagicMock()
|
||||
|
||||
with patch.dict(
|
||||
"sys.modules",
|
||||
{
|
||||
"mlflow": MagicMock(),
|
||||
"mlflow.tracking": mock_mlflow_tracking,
|
||||
"mlflow.tracing.utils": MagicMock(),
|
||||
},
|
||||
):
|
||||
from litellm.integrations.mlflow import MlflowLogger
|
||||
|
||||
mlflow_logger = MlflowLogger()
|
||||
|
||||
attrs = mlflow_logger._extract_attributes( # type: ignore
|
||||
{
|
||||
"litellm_call_id": "123",
|
||||
"call_type": "completion",
|
||||
"model": "gpt-3.5-turbo",
|
||||
"standard_logging_object": {
|
||||
"prompt_tokens": 5,
|
||||
"completion_tokens": 7,
|
||||
"total_tokens": 12,
|
||||
},
|
||||
}
|
||||
)
|
||||
|
||||
assert attrs["mlflow.chat.tokenUsage"] == {
|
||||
"input_tokens": 5,
|
||||
"output_tokens": 7,
|
||||
"total_tokens": 12,
|
||||
}
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue