Improving langfuse usage details unit test
This commit is contained in:
parent
f4f89cbae1
commit
118d15c1c6
@ -9,10 +9,8 @@ from litellm.integrations.langfuse.langfuse import LangFuseLogger
|
||||
# Import LangfuseUsageDetails directly from the module where it's defined
|
||||
from litellm.types.integrations.langfuse import *
|
||||
|
||||
|
||||
|
||||
class TestLangfuseUsageDetails(unittest.TestCase):
|
||||
|
||||
|
||||
def setUp(self):
|
||||
# Set up environment variables for testing
|
||||
self.env_patcher = patch.dict('os.environ', {
|
||||
@ -21,50 +19,50 @@ class TestLangfuseUsageDetails(unittest.TestCase):
|
||||
'LANGFUSE_HOST': 'https://test.langfuse.com'
|
||||
})
|
||||
self.env_patcher.start()
|
||||
|
||||
|
||||
# Create mock objects
|
||||
self.mock_langfuse_client = MagicMock()
|
||||
self.mock_langfuse_trace = MagicMock()
|
||||
self.mock_langfuse_generation = MagicMock()
|
||||
|
||||
|
||||
# Setup the trace and generation chain
|
||||
self.mock_langfuse_trace.generation.return_value = self.mock_langfuse_generation
|
||||
self.mock_langfuse_client.trace.return_value = self.mock_langfuse_trace
|
||||
|
||||
|
||||
# Mock the langfuse module that's imported locally in methods
|
||||
self.langfuse_module_patcher = patch.dict('sys.modules', {'langfuse': MagicMock()})
|
||||
self.mock_langfuse_module = self.langfuse_module_patcher.start()
|
||||
|
||||
|
||||
# Create a mock for the langfuse module with version
|
||||
self.mock_langfuse = MagicMock()
|
||||
self.mock_langfuse.version = MagicMock()
|
||||
self.mock_langfuse.version.__version__ = "3.0.0" # Set a version that supports all features
|
||||
|
||||
|
||||
# Mock the Langfuse class
|
||||
self.mock_langfuse_class = MagicMock()
|
||||
self.mock_langfuse_class.return_value = self.mock_langfuse_client
|
||||
|
||||
|
||||
# Set up the sys.modules['langfuse'] mock
|
||||
sys.modules['langfuse'] = self.mock_langfuse
|
||||
sys.modules['langfuse'].Langfuse = self.mock_langfuse_class
|
||||
|
||||
|
||||
# Mock the Langfuse client
|
||||
self.mock_langfuse_client = MagicMock()
|
||||
self.mock_langfuse_trace = MagicMock()
|
||||
self.mock_langfuse_generation = MagicMock()
|
||||
|
||||
|
||||
# Setup the trace and generation chain
|
||||
self.mock_langfuse_trace.generation.return_value = self.mock_langfuse_generation
|
||||
self.mock_langfuse_client.trace.return_value = self.mock_langfuse_trace
|
||||
|
||||
|
||||
# Mock the Langfuse class
|
||||
self.mock_langfuse_class = MagicMock()
|
||||
self.mock_langfuse_class.return_value = self.mock_langfuse_client
|
||||
self.mock_langfuse.Langfuse = self.mock_langfuse_class
|
||||
|
||||
|
||||
# Create the logger
|
||||
self.logger = LangFuseLogger()
|
||||
|
||||
|
||||
# Add the log_event_on_langfuse method to the instance
|
||||
def log_event_on_langfuse(self, kwargs, response_obj, start_time=None, end_time=None, user_id=None, level="DEFAULT", status_message=None):
|
||||
# This implementation calls _log_langfuse_v2 directly
|
||||
@ -83,21 +81,21 @@ class TestLangfuseUsageDetails(unittest.TestCase):
|
||||
litellm_call_id=kwargs.get("litellm_call_id", None),
|
||||
print_verbose=True # Add the missing parameter
|
||||
)
|
||||
|
||||
|
||||
# Bind the method to the instance
|
||||
import types
|
||||
self.logger.log_event_on_langfuse = types.MethodType(log_event_on_langfuse, self.logger)
|
||||
|
||||
|
||||
# Make sure _is_langfuse_v2 returns True
|
||||
def mock_is_langfuse_v2(self):
|
||||
return True
|
||||
|
||||
|
||||
self.logger._is_langfuse_v2 = types.MethodType(mock_is_langfuse_v2, self.logger)
|
||||
|
||||
|
||||
def tearDown(self):
|
||||
self.env_patcher.stop()
|
||||
self.langfuse_module_patcher.stop()
|
||||
|
||||
|
||||
def test_langfuse_usage_details_type(self):
|
||||
"""Test that LangfuseUsageDetails TypedDict is properly defined with the correct fields"""
|
||||
# Create an instance of LangfuseUsageDetails
|
||||
@ -107,13 +105,13 @@ class TestLangfuseUsageDetails(unittest.TestCase):
|
||||
"cache_creation_input_tokens": 5,
|
||||
"cache_read_input_tokens": 3
|
||||
}
|
||||
|
||||
|
||||
# Verify all fields are present
|
||||
self.assertEqual(usage_details["input"], 10)
|
||||
self.assertEqual(usage_details["output"], 20)
|
||||
self.assertEqual(usage_details["cache_creation_input_tokens"], 5)
|
||||
self.assertEqual(usage_details["cache_read_input_tokens"], 3)
|
||||
|
||||
|
||||
# Test with all fields (all fields are required in TypedDict by default)
|
||||
minimal_usage_details: LangfuseUsageDetails = {
|
||||
"input": 10,
|
||||
@ -121,10 +119,10 @@ class TestLangfuseUsageDetails(unittest.TestCase):
|
||||
"cache_creation_input_tokens": 0,
|
||||
"cache_read_input_tokens": 0
|
||||
}
|
||||
|
||||
|
||||
self.assertEqual(minimal_usage_details["input"], 10)
|
||||
self.assertEqual(minimal_usage_details["output"], 20)
|
||||
|
||||
|
||||
def test_log_langfuse_v2_usage_details(self):
|
||||
"""Test that usage_details in _log_langfuse_v2 is correctly typed and assigned"""
|
||||
# Create a mock response object with usage information
|
||||
@ -132,7 +130,7 @@ class TestLangfuseUsageDetails(unittest.TestCase):
|
||||
response_obj.usage = MagicMock()
|
||||
response_obj.usage.prompt_tokens = 15
|
||||
response_obj.usage.completion_tokens = 25
|
||||
|
||||
|
||||
# Add the cache token attributes using get method
|
||||
def mock_get(key, default=None):
|
||||
if key == 'cache_creation_input_tokens':
|
||||
@ -140,20 +138,20 @@ class TestLangfuseUsageDetails(unittest.TestCase):
|
||||
elif key == 'cache_read_input_tokens':
|
||||
return 4
|
||||
return default
|
||||
|
||||
|
||||
response_obj.usage.get = mock_get
|
||||
|
||||
|
||||
# Create kwargs for the log_event method
|
||||
kwargs = {
|
||||
"model": "gpt-4",
|
||||
"messages": [{"role": "user", "content": "Hello"}],
|
||||
"litellm_params": {"metadata": {}}
|
||||
}
|
||||
|
||||
|
||||
# Create start and end times
|
||||
start_time = datetime.datetime.now()
|
||||
end_time = start_time + datetime.timedelta(seconds=1)
|
||||
|
||||
|
||||
# Call the log_event method
|
||||
with patch.object(self.logger, '_log_langfuse_v2') as mock_log_langfuse_v2:
|
||||
self.logger.log_event_on_langfuse(
|
||||
@ -162,16 +160,16 @@ class TestLangfuseUsageDetails(unittest.TestCase):
|
||||
start_time=start_time,
|
||||
end_time=end_time
|
||||
)
|
||||
|
||||
|
||||
# Check if _log_langfuse_v2 was called
|
||||
mock_log_langfuse_v2.assert_called_once()
|
||||
|
||||
|
||||
# Get the arguments passed to _log_langfuse_v2
|
||||
call_args = mock_log_langfuse_v2.call_args[1]
|
||||
|
||||
|
||||
# Verify response_obj was passed correctly
|
||||
self.assertEqual(call_args["response_obj"], response_obj)
|
||||
|
||||
|
||||
def test_langfuse_usage_details_optional_fields(self):
|
||||
"""Test that LangfuseUsageDetails fields are properly defined as Optional"""
|
||||
# Create an instance with None values for optional fields
|
||||
@ -181,18 +179,18 @@ class TestLangfuseUsageDetails(unittest.TestCase):
|
||||
"cache_creation_input_tokens": None,
|
||||
"cache_read_input_tokens": None
|
||||
}
|
||||
|
||||
|
||||
# Verify fields can be None
|
||||
self.assertEqual(usage_details["input"], 10)
|
||||
self.assertEqual(usage_details["output"], 20)
|
||||
self.assertIsNone(usage_details["cache_creation_input_tokens"])
|
||||
self.assertIsNone(usage_details["cache_read_input_tokens"])
|
||||
|
||||
|
||||
def test_langfuse_usage_details_structure(self):
|
||||
"""Test that LangfuseUsageDetails has the correct structure as defined in the commit"""
|
||||
# This test directly verifies the structure of the TypedDict
|
||||
# without relying on the LangFuseLogger class
|
||||
|
||||
|
||||
# Create a dictionary that matches the LangfuseUsageDetails structure
|
||||
usage_details = {
|
||||
"input": 15,
|
||||
@ -200,19 +198,15 @@ class TestLangfuseUsageDetails(unittest.TestCase):
|
||||
"cache_creation_input_tokens": 7,
|
||||
"cache_read_input_tokens": 4
|
||||
}
|
||||
|
||||
|
||||
# Verify the structure matches what we expect
|
||||
self.assertIn("input", usage_details)
|
||||
self.assertIn("output", usage_details)
|
||||
self.assertIn("cache_creation_input_tokens", usage_details)
|
||||
self.assertIn("cache_read_input_tokens", usage_details)
|
||||
|
||||
|
||||
# Verify the values
|
||||
self.assertEqual(usage_details["input"], 15)
|
||||
self.assertEqual(usage_details["output"], 25)
|
||||
self.assertEqual(usage_details["cache_creation_input_tokens"], 7)
|
||||
self.assertEqual(usage_details["cache_read_input_tokens"], 4)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
|
||||
Loading…
Reference in New Issue
Block a user