diff --git a/litellm/integrations/langfuse.py b/litellm/integrations/langfuse.py index d2cfae41adc..59ea1d02ddc 100644 --- a/litellm/integrations/langfuse.py +++ b/litellm/integrations/langfuse.py @@ -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 diff --git a/litellm/litellm_core_utils/llm_request_utils.py b/litellm/litellm_core_utils/llm_request_utils.py new file mode 100644 index 00000000000..557d73b0ab8 --- /dev/null +++ b/litellm/litellm_core_utils/llm_request_utils.py @@ -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 diff --git a/litellm/tests/test_alangfuse.py b/litellm/tests/test_alangfuse.py index 2303dc9e8c9..4496993355e 100644 --- a/litellm/tests/test_alangfuse.py +++ b/litellm/tests/test_alangfuse.py @@ -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"), + } + }, + ) diff --git a/litellm/utils.py b/litellm/utils.py index eac07c55069..795526a3212 100644 --- a/litellm/utils.py +++ b/litellm/utils.py @@ -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():