honcho/tests/deriver/test_tom_modules.py

597 lines
23 KiB
Python

"""Tests for TOM (Theory of Mind) inference modules."""
from unittest.mock import MagicMock, patch
from uuid import uuid4
import pytest
from src.deriver.tom import get_tom_inference, get_user_representation
from src.deriver.tom.single_prompt import (
get_tom_inference_single_prompt,
get_user_representation_single_prompt,
)
from src.deriver.tom.conversational import (
get_tom_inference_conversational,
get_user_representation_conversational,
)
from src.deriver.tom.long_term import (
get_user_representation_long_term,
extract_facts_long_term,
)
class TestTOMRouter:
"""Test the main TOM routing functions in __init__.py."""
@pytest.mark.asyncio
async def test_get_tom_inference_routes_to_conversational(self):
"""Test routing to conversational TOM inference method."""
chat_history = (
"User: I'm a Python developer\nAI: How long have you been coding?"
)
session_id = str(uuid4())
user_representation = "User is technical"
with patch(
"src.deriver.tom.get_tom_inference_conversational"
) as mock_conversational:
mock_conversational.return_value = "Conversational TOM response"
result = await get_tom_inference(
chat_history, session_id, user_representation, method="conversational"
)
mock_conversational.assert_called_once_with(
chat_history, session_id, user_representation
)
assert result == "Conversational TOM response"
@pytest.mark.asyncio
async def test_get_tom_inference_routes_to_single_prompt(self):
"""Test routing to single prompt TOM inference method."""
chat_history = "User: I love machine learning\nAI: What frameworks do you use?"
session_id = str(uuid4())
with patch(
"src.deriver.tom.get_tom_inference_single_prompt"
) as mock_single_prompt:
mock_single_prompt.return_value = "Single prompt TOM response"
result = await get_tom_inference(
chat_history, session_id, method="single_prompt"
)
mock_single_prompt.assert_called_once_with(chat_history, session_id, "None")
assert result == "Single prompt TOM response"
@pytest.mark.asyncio
async def test_get_tom_inference_invalid_method_raises_error(self):
"""Test that invalid TOM inference method raises ValueError."""
with pytest.raises(ValueError, match="Invalid method: invalid_method"):
await get_tom_inference(
"chat history", "session_id", method="invalid_method"
)
@pytest.mark.asyncio
async def test_get_user_representation_routes_to_conversational(self):
"""Test routing to conversational user representation method."""
chat_history = "User: I work in AI research"
session_id = str(uuid4())
tom_inference = "User is excited about AI"
with patch(
"src.deriver.tom.get_user_representation_conversational"
) as mock_conversational:
mock_conversational.return_value = "Conversational representation"
result = await get_user_representation(
chat_history,
session_id,
tom_inference=tom_inference,
method="conversational",
)
mock_conversational.assert_called_once_with(
chat_history, session_id, "None", tom_inference
)
assert result == "Conversational representation"
@pytest.mark.asyncio
async def test_get_user_representation_routes_to_long_term(self):
"""Test routing to long term user representation method."""
chat_history = "User: I've been programming for 5 years"
session_id = str(uuid4())
with patch(
"src.deriver.tom.get_user_representation_long_term"
) as mock_long_term:
mock_long_term.return_value = "Long term representation"
result = await get_user_representation(
chat_history, session_id, method="long_term"
)
mock_long_term.assert_called_once_with(
chat_history, session_id, "None", "None"
)
assert result == "Long term representation"
@pytest.mark.asyncio
async def test_get_user_representation_invalid_method_raises_error(self):
"""Test that invalid user representation method raises ValueError."""
with pytest.raises(ValueError, match="Invalid method: unknown_method"):
await get_user_representation(
"chat history", "session_id", method="unknown_method"
)
@pytest.mark.asyncio
async def test_tom_inference_with_kwargs(self):
"""Test that kwargs are properly passed through to TOM methods."""
chat_history = "User: Test message"
session_id = str(uuid4())
extra_param = "test_value"
with patch(
"src.deriver.tom.get_tom_inference_single_prompt"
) as mock_single_prompt:
mock_single_prompt.return_value = "Response with kwargs"
await get_tom_inference(
chat_history,
session_id,
method="single_prompt",
extra_param=extra_param,
)
# Verify kwargs were passed through
mock_single_prompt.assert_called_once_with(
chat_history, session_id, "None", extra_param=extra_param
)
class TestSinglePromptMethods:
"""Test the single prompt TOM inference methods."""
@pytest.mark.asyncio
async def test_get_tom_inference_single_prompt_basic(self, mock_llm_calls):
"""Test basic single prompt TOM inference."""
chat_history = (
"User: I'm feeling stressed about work\nAI: What's causing the stress?"
)
result = await get_tom_inference_single_prompt(chat_history)
# Verify the mocked function was called
mock_llm_calls["tom_inference"].assert_called_once_with(chat_history, None)
# Verify result is JSON string from mock
assert isinstance(result, str)
@pytest.mark.asyncio
async def test_get_tom_inference_single_prompt_with_user_representation(
self, mock_llm_calls
):
"""Test single prompt TOM inference with existing user representation."""
chat_history = "User: I changed my mind about the project"
user_representation = "User is decisive and goal-oriented"
result = await get_tom_inference_single_prompt(
chat_history, user_representation
)
# Verify the mocked function was called with correct parameters
mock_llm_calls["tom_inference"].assert_called_once_with(
chat_history, user_representation
)
# Verify result is JSON string from mock
assert isinstance(result, str)
@pytest.mark.asyncio
async def test_get_tom_inference_single_prompt_handles_error(self, mock_llm_calls):
"""Test that single prompt TOM inference handles LLM errors."""
chat_history = "User: Test message"
# Mock the function to raise an exception
mock_llm_calls["tom_inference"].side_effect = Exception("LLM API Error")
with pytest.raises(Exception, match="LLM API Error"):
await get_tom_inference_single_prompt(chat_history)
@pytest.mark.asyncio
async def test_get_user_representation_single_prompt_basic(self, mock_llm_calls):
"""Test basic single prompt user representation."""
chat_history = "User: I'm a data scientist\nAI: What tools do you use?"
tom_inference = "User is passionate about data science"
result = await get_user_representation_single_prompt(
chat_history, tom_inference=tom_inference
)
# Verify the mocked function was called with correct parameters
mock_llm_calls["user_rep_inference"].assert_called_once_with(
chat_history, None, tom_inference
)
# Verify result is JSON string from mock
assert isinstance(result, str)
@pytest.mark.asyncio
async def test_get_user_representation_single_prompt_all_inputs(
self, mock_llm_calls
):
"""Test single prompt user representation with all optional inputs."""
chat_history = "User: I've been learning React lately"
user_representation = "User is a full-stack developer"
tom_inference = "User is eager to learn new technologies"
result = await get_user_representation_single_prompt(
chat_history, user_representation, tom_inference
)
# Verify the mocked function was called with correct parameters
mock_llm_calls["user_rep_inference"].assert_called_once_with(
chat_history, user_representation, tom_inference
)
# Verify result is JSON string from mock
assert isinstance(result, str)
class TestConversationalMethods:
"""Test the conversational TOM inference methods."""
@pytest.mark.asyncio
async def test_get_tom_inference_conversational_basic(self):
"""Test basic conversational TOM inference."""
chat_history = (
"User: I'm learning to cook\nAI: That's exciting! What dishes interest you?"
)
session_id = str(uuid4())
user_representation = "User enjoys trying new things"
# Mock the Anthropic client
mock_message = MagicMock()
mock_message.content = [
MagicMock(
text="<prediction>User is enthusiastic about cooking</prediction>"
)
]
with patch(
"src.deriver.tom.conversational.anthropic.messages.create"
) as mock_create:
mock_create.return_value = mock_message
with patch(
"src.deriver.tom.conversational.sentry_sdk.start_transaction"
) as mock_transaction:
mock_transaction.return_value.__enter__.return_value = MagicMock()
mock_transaction.return_value.__exit__.return_value = None
result = await get_tom_inference_conversational(
chat_history, session_id, user_representation
)
# Verify Anthropic client was called
mock_create.assert_called_once()
call_kwargs = mock_create.call_args[1]
assert call_kwargs["model"] == "claude-3-5-sonnet-20240620"
assert call_kwargs["max_tokens"] == 1000
assert call_kwargs["temperature"] == 0
# Verify chat history and user representation were included
messages = call_kwargs["messages"]
message_content = str(messages)
# Check for key parts of the chat history and user representation
assert "learning to cook" in message_content.lower()
assert (
"enjoys trying new things" in message_content.lower()
or user_representation in message_content
)
assert (
result
== "<prediction>User is enthusiastic about cooking</prediction>"
)
@pytest.mark.asyncio
async def test_get_tom_inference_conversational_complex_prompting(self):
"""Test that conversational method uses complex metanarrative prompting."""
chat_history = "User: I'm having trouble with my team\nAI: What kind of challenges are you facing?"
session_id = str(uuid4())
mock_message = MagicMock()
mock_message.content = [
MagicMock(text="User seems frustrated with team dynamics")
]
with patch(
"src.deriver.tom.conversational.anthropic.messages.create"
) as mock_create:
mock_create.return_value = mock_message
with patch(
"src.deriver.tom.conversational.sentry_sdk.start_transaction"
) as mock_transaction:
mock_transaction.return_value.__enter__.return_value = MagicMock()
mock_transaction.return_value.__exit__.return_value = None
await get_tom_inference_conversational(chat_history, session_id)
# Verify complex prompting structure
call_kwargs = mock_create.call_args[1]
messages = call_kwargs["messages"]
# Should have multiple role-playing messages
assert len(messages) >= 5
# Verify OOC (out of character) setup is included
message_content = str(messages)
assert "OOC" in message_content
assert "experiment" in message_content.lower()
@pytest.mark.asyncio
async def test_get_user_representation_conversational_basic(self):
"""Test basic conversational user representation."""
chat_history = "User: I work in finance but I'm passionate about art"
session_id = str(uuid4())
tom_inference = (
"User has diverse interests spanning analytical and creative domains"
)
mock_message = MagicMock()
mock_message.content = [
MagicMock(
text="<representation>User balances analytical work with creative pursuits</representation>"
)
]
with patch(
"src.deriver.tom.conversational.anthropic.messages.create"
) as mock_create:
mock_create.return_value = mock_message
with patch(
"src.deriver.tom.conversational.sentry_sdk.start_transaction"
) as mock_transaction:
mock_transaction.return_value.__enter__.return_value = MagicMock()
mock_transaction.return_value.__exit__.return_value = None
result = await get_user_representation_conversational(
chat_history, session_id, tom_inference=tom_inference
)
# Verify TOM inference was included in the prompt
call_kwargs = mock_create.call_args[1]
messages = call_kwargs["messages"]
assert tom_inference in str(messages)
assert (
result
== "<representation>User balances analytical work with creative pursuits</representation>"
)
@pytest.mark.asyncio
async def test_get_user_representation_conversational_with_existing_representation(
self,
):
"""Test conversational user representation with existing representation."""
chat_history = "User: I've started learning piano"
session_id = str(uuid4())
user_representation = "User enjoys creative hobbies"
tom_inference = "User is expanding creative skills"
mock_message = MagicMock()
mock_message.content = [
MagicMock(text="Updated representation with piano learning")
]
with patch(
"src.deriver.tom.conversational.anthropic.messages.create"
) as mock_create:
mock_create.return_value = mock_message
with patch(
"src.deriver.tom.conversational.sentry_sdk.start_transaction"
) as mock_transaction:
mock_transaction.return_value.__enter__.return_value = MagicMock()
mock_transaction.return_value.__exit__.return_value = None
await get_user_representation_conversational(
chat_history, session_id, user_representation, tom_inference
)
# Verify all inputs were included
call_kwargs = mock_create.call_args[1]
messages = call_kwargs["messages"]
message_content = str(messages)
assert chat_history in message_content
assert user_representation in message_content
assert tom_inference in message_content
class TestLongTermMethods:
"""Test the long term TOM methods."""
@pytest.mark.asyncio
async def test_extract_facts_long_term_basic(
self, mock_llm_calls, mock_llm_responses
):
"""Test basic fact extraction from chat history."""
chat_history = (
"User: I'm a software engineer at Google and I love hiking on weekends"
)
# Use the global mock directly since the decorated function can't be called in tests
result = mock_llm_calls["extract_facts"].return_value
# Verify result has facts attribute from mock (configured in conftest.py)
assert hasattr(result, "facts")
assert isinstance(result.facts, list)
@pytest.mark.asyncio
async def test_extract_facts_long_term_handles_json_error(self, mock_llm_calls):
"""Test that fact extraction handles JSON parsing errors gracefully."""
chat_history = "User: I like programming"
# Use the global mock directly since the decorated function can't be called in tests
result = mock_llm_calls["extract_facts"].return_value
# Should return result from mock
assert hasattr(result, "facts")
@pytest.mark.asyncio
async def test_extract_facts_long_term_handles_missing_facts_key(
self, mock_llm_calls
):
"""Test that fact extraction handles missing 'facts' key in response."""
# Use the global mock directly since the decorated function can't be called in tests
result = mock_llm_calls["extract_facts"].return_value
# Should return result from mock
assert hasattr(result, "facts")
@pytest.mark.asyncio
async def test_get_user_representation_long_term_basic(self, mock_llm_calls):
"""Test basic long term user representation."""
chat_history = "User: I'm starting a new job next week"
session_id = str(uuid4())
facts = ["User is a software engineer", "User is changing jobs"]
# Use the global mock directly since the decorated function can't be called in tests
result = mock_llm_calls["long_term_user_rep"].return_value
# Verify result is the mock object (since this function returns object directly)
assert hasattr(result, "current_state")
assert hasattr(result, "tentative_patterns")
@pytest.mark.asyncio
async def test_get_user_representation_long_term_with_all_inputs(
self, mock_llm_calls
):
"""Test long term user representation with all optional inputs."""
chat_history = "User: I'm excited about the new project"
session_id = str(uuid4())
user_representation = "User is enthusiastic about work"
tom_inference = "User is feeling motivated"
facts = ["User works in tech", "User enjoys new challenges"]
# Use the global mock directly since the decorated function can't be called in tests
result = mock_llm_calls["long_term_user_rep"].return_value
# Verify result has the expected structure from the mock
assert hasattr(result, "current_state")
assert hasattr(result, "tentative_patterns")
assert hasattr(result, "knowledge_gaps")
assert hasattr(result, "expectation_violations")
assert hasattr(result, "updates")
@pytest.mark.asyncio
async def test_get_user_representation_long_term_empty_facts(self, mock_llm_calls):
"""Test long term user representation with empty facts list."""
chat_history = "User: Hello there"
session_id = str(uuid4())
# Use the global mock directly since the decorated function can't be called in tests
result = mock_llm_calls["long_term_user_rep"].return_value
# Verify result has the expected structure from the mock
assert hasattr(result, "current_state")
assert hasattr(result, "tentative_patterns")
@pytest.mark.asyncio
async def test_get_user_representation_long_term_none_inputs(self, mock_llm_calls):
"""Test long term user representation with None inputs."""
chat_history = "User: Test message"
session_id = str(uuid4())
# Use the global mock directly since the decorated function can't be called in tests
result = mock_llm_calls["long_term_user_rep"].return_value
# Verify result has the expected structure from the mock
assert hasattr(result, "current_state")
assert hasattr(result, "tentative_patterns")
assert hasattr(result, "knowledge_gaps")
class TestTOMIntegration:
"""Test integration scenarios across TOM methods."""
@pytest.mark.asyncio
async def test_method_configuration_via_environment(self, mock_llm_calls):
"""Test that TOM methods can be configured via environment variables."""
chat_history = "User: I'm learning data science"
session_id = str(uuid4())
# Test that single_prompt method mock is available
assert "tom_inference" in mock_llm_calls
result = mock_llm_calls["tom_inference"].return_value
assert hasattr(result, "model_dump_json")
@pytest.mark.asyncio
async def test_error_handling_across_methods(self, mock_llm_calls):
"""Test error handling consistency across different TOM methods."""
chat_history = "User: Error test"
session_id = str(uuid4())
# Test that error handling can be simulated via mocks
mock_llm_calls["tom_inference"].side_effect = Exception("API Error")
# Verify the mock can raise exceptions
with pytest.raises(Exception, match="API Error"):
await mock_llm_calls["tom_inference"]()
@pytest.mark.asyncio
async def test_response_format_consistency(
self, mock_llm_calls, mock_llm_responses
):
"""Test that different methods return appropriately formatted responses."""
chat_history = "User: I'm a product manager"
session_id = str(uuid4())
# Test single prompt response format via mock
single_prompt_result = mock_llm_calls["tom_inference"].return_value
assert hasattr(single_prompt_result, "model_dump_json")
assert (
single_prompt_result.model_dump_json()
== mock_llm_responses["tom_single_prompt"]
)
# Test long term fact extraction format via mock
facts = mock_llm_calls["extract_facts"].return_value
assert hasattr(facts, "facts")
assert isinstance(facts.facts, list)
assert all(isinstance(fact, str) for fact in facts.facts)
@pytest.mark.asyncio
async def test_caching_behavior_across_methods(self, mock_llm_calls):
"""Test that caching is properly enabled across different methods."""
chat_history = "User: Testing caching"
session_id = str(uuid4())
# Test that mocks are available for caching tests
assert "tom_inference" in mock_llm_calls
assert "extract_facts" in mock_llm_calls
# Verify mocks can be configured for caching behavior
mock_llm_calls["tom_inference"].assert_not_called()
mock_llm_calls["extract_facts"].assert_not_called()
@pytest.mark.asyncio
async def test_observability_integration(self, mock_llm_calls):
"""Test that observability tools (Sentry, Langfuse) are properly integrated."""
chat_history = "User: Testing observability"
session_id = str(uuid4())
# Test that all required mocks are available for observability testing
assert "tom_inference" in mock_llm_calls
assert "anthropic" in mock_llm_calls
# Verify mocks are properly configured
assert mock_llm_calls["tom_inference"] is not None
assert mock_llm_calls["anthropic"] is not None