Merge pull request #4275 from BerriAI/litellm_fix_langfuse_log_prompts

[Fix] Use Langfuse prompt Object with LiteLLM Proxy
This commit is contained in:
Ishaan Jaff 2024-06-18 20:09:13 -07:00 committed by GitHub
commit e1b646c304
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
4 changed files with 141 additions and 35 deletions

View file

@ -1,11 +1,13 @@
#### What this does ####
# On success, logs events to Langfuse
import os
import copy
import os
import traceback
from packaging.version import Version
from litellm._logging import verbose_logger
import litellm
from litellm._logging import verbose_logger
class LangFuseLogger:
@ -14,8 +16,8 @@ class LangFuseLogger:
self, langfuse_public_key=None, langfuse_secret=None, flush_interval=1
):
try:
from langfuse import Langfuse
import langfuse
from langfuse import Langfuse
except Exception as e:
raise Exception(
f"\033[91mLangfuse not installed, try running 'pip install langfuse' to fix this error: {e}\n{traceback.format_exc()}\033[0m"
@ -251,7 +253,7 @@ class LangFuseLogger:
input,
response_obj,
):
from langfuse.model import CreateTrace, CreateGeneration
from langfuse.model import CreateGeneration, CreateTrace
verbose_logger.warning(
"Please upgrade langfuse to v2.0.0 or higher: https://github.com/langfuse/langfuse-python/releases/tag/v2.0.1"
@ -533,30 +535,9 @@ class LangFuseLogger:
generation_params["parent_observation_id"] = parent_observation_id
if supports_prompt:
user_prompt = clean_metadata.pop("prompt", None)
if user_prompt is None:
pass
elif isinstance(user_prompt, dict):
from langfuse.model import (
TextPromptClient,
ChatPromptClient,
Prompt_Text,
Prompt_Chat,
)
if user_prompt.get("type", "") == "chat":
_prompt_chat = Prompt_Chat(**user_prompt)
generation_params["prompt"] = ChatPromptClient(
prompt=_prompt_chat
)
elif user_prompt.get("type", "") == "text":
_prompt_text = Prompt_Text(**user_prompt)
generation_params["prompt"] = TextPromptClient(
prompt=_prompt_text
)
else:
generation_params["prompt"] = user_prompt
generation_params = _add_prompt_to_generation_params(
generation_params=generation_params, clean_metadata=clean_metadata
)
if output is not None and isinstance(output, str) and level == "ERROR":
generation_params["status_message"] = output
@ -569,5 +550,58 @@ class LangFuseLogger:
return generation_client.trace_id, generation_id
except Exception as e:
verbose_logger.debug(f"Langfuse Layer Error - {traceback.format_exc()}")
verbose_logger.error(f"Langfuse Layer Error - {traceback.format_exc()}")
return None, None
def _add_prompt_to_generation_params(
generation_params: dict, clean_metadata: dict
) -> dict:
from langfuse.model import (
ChatPromptClient,
Prompt_Chat,
Prompt_Text,
TextPromptClient,
)
user_prompt = clean_metadata.pop("prompt", None)
if user_prompt is None:
pass
elif isinstance(user_prompt, dict):
if user_prompt.get("type", "") == "chat":
_prompt_chat = Prompt_Chat(**user_prompt)
generation_params["prompt"] = ChatPromptClient(prompt=_prompt_chat)
elif user_prompt.get("type", "") == "text":
_prompt_text = Prompt_Text(**user_prompt)
generation_params["prompt"] = TextPromptClient(prompt=_prompt_text)
elif "version" in user_prompt and "prompt" in user_prompt:
# prompts
if isinstance(user_prompt["prompt"], str):
_prompt_obj = Prompt_Text(
name=user_prompt["name"],
prompt=user_prompt["prompt"],
version=user_prompt["version"],
config=user_prompt.get("config", None),
)
generation_params["prompt"] = TextPromptClient(prompt=_prompt_obj)
elif isinstance(user_prompt["prompt"], list):
_prompt_obj = Prompt_Chat(
name=user_prompt["name"],
prompt=user_prompt["prompt"],
version=user_prompt["version"],
config=user_prompt.get("config", None),
)
generation_params["prompt"] = ChatPromptClient(prompt=_prompt_obj)
else:
verbose_logger.error(
"[Non-blocking] Langfuse Logger: Invalid prompt format"
)
else:
verbose_logger.error(
"[Non-blocking] Langfuse Logger: Invalid prompt format. No prompt logged to Langfuse"
)
else:
generation_params["prompt"] = user_prompt
return generation_params

View file

@ -0,0 +1,28 @@
from typing import Dict, Optional
def _ensure_extra_body_is_safe(extra_body: Optional[Dict]) -> Optional[Dict]:
"""
Ensure that the extra_body sent in the request is safe, otherwise users will see this error
"Object of type TextPromptClient is not JSON serializable
Relevant Issue: https://github.com/BerriAI/litellm/issues/4140
"""
if extra_body is None:
return None
if not isinstance(extra_body, dict):
return extra_body
if "metadata" in extra_body and isinstance(extra_body["metadata"], dict):
if "prompt" in extra_body["metadata"]:
_prompt = extra_body["metadata"].get("prompt")
# users can send Langfuse TextPromptClient objects, so we need to convert them to dicts
# Langfuse TextPromptClients have .__dict__ attribute
if _prompt is not None and hasattr(_prompt, "__dict__"):
extra_body["metadata"]["prompt"] = _prompt.__dict__
return extra_body

View file

@ -1,22 +1,22 @@
import asyncio
import copy
import json
import sys
import os
import asyncio
import logging
import os
import sys
from unittest.mock import MagicMock, patch
logging.basicConfig(level=logging.DEBUG)
sys.path.insert(0, os.path.abspath("../.."))
from litellm import completion
import litellm
from litellm import completion
litellm.num_retries = 3
litellm.success_callback = ["langfuse"]
os.environ["LANGFUSE_DEBUG"] = "True"
import time
import pytest
@ -551,7 +551,9 @@ def test_aaalangfuse_existing_trace_id():
Assert no changes to the trace
"""
# Test - if the logs were sent to the correct team on langfuse
import litellm, datetime
import datetime
import litellm
from litellm.integrations.langfuse import LangFuseLogger
langfuse_Logger = LangFuseLogger(
@ -827,3 +829,40 @@ def test_langfuse_logging_tool_calling():
# test_langfuse_logging_tool_calling()
def get_langfuse_prompt(name: str):
import langfuse
from langfuse import Langfuse
try:
langfuse = Langfuse(
public_key=os.environ["LANGFUSE_DEV_PUBLIC_KEY"],
secret_key=os.environ["LANGFUSE_DEV_SK_KEY"],
host=os.environ["LANGFUSE_HOST"],
)
# Get current production version of a text prompt
prompt = langfuse.get_prompt(name=name)
return prompt
except Exception as e:
raise Exception(f"Error getting prompt: {e}")
@pytest.mark.asyncio
@pytest.mark.skip(
reason="local only test, use this to verify if we can send request to litellm proxy server"
)
async def test_make_request():
response = await litellm.acompletion(
model="openai/llama3",
api_key="sk-1234",
base_url="http://localhost:4000",
messages=[{"role": "user", "content": "Hi 👋 - i'm claude"}],
extra_body={
"metadata": {
"tags": ["openai"],
"prompt": get_langfuse_prompt("test-chat"),
}
},
)

View file

@ -50,6 +50,7 @@ import litellm._service_logger # for storing API inputs, outputs, and metadata
import litellm.litellm_core_utils
from litellm.caching import DualCache
from litellm.litellm_core_utils.core_helpers import map_finish_reason
from litellm.litellm_core_utils.llm_request_utils import _ensure_extra_body_is_safe
from litellm.litellm_core_utils.redact_messages import (
redact_message_input_output_from_logging,
)
@ -3256,6 +3257,10 @@ def get_optional_params(
extra_body[k] = passed_params[k]
optional_params.setdefault("extra_body", {})
optional_params["extra_body"] = {**optional_params["extra_body"], **extra_body}
optional_params["extra_body"] = _ensure_extra_body_is_safe(
extra_body=optional_params["extra_body"]
)
else:
# if user passed in non-default kwargs for specific providers/models, pass them along
for k in passed_params.keys():