unit test the deriver
This commit is contained in:
parent
c22cc507b0
commit
67d3a3e9f1
|
|
@ -1,4 +1,7 @@
|
|||
{
|
||||
"python.analysis.typeCheckingMode": "basic",
|
||||
"files.exclude": {}
|
||||
"files.exclude": {},
|
||||
"python.testing.pytestArgs": ["tests"],
|
||||
"python.testing.unittestEnabled": false,
|
||||
"python.testing.pytestEnabled": true
|
||||
}
|
||||
|
|
|
|||
|
|
@ -1,7 +1,7 @@
|
|||
import asyncio
|
||||
import os
|
||||
import signal
|
||||
from datetime import datetime, timedelta
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from logging import getLogger
|
||||
|
||||
import sentry_sdk
|
||||
|
|
@ -114,7 +114,7 @@ class QueueManager:
|
|||
async def get_available_sessions(self, db: AsyncSession):
|
||||
"""Get available sessions that aren't being processed"""
|
||||
# Clean up stale sessions
|
||||
five_minutes_ago = datetime.utcnow() - timedelta(minutes=5)
|
||||
five_minutes_ago = datetime.now(timezone.utc) - timedelta(minutes=5)
|
||||
await db.execute(
|
||||
delete(models.ActiveQueueSession).where(
|
||||
models.ActiveQueueSession.last_updated < five_minutes_ago
|
||||
|
|
|
|||
|
|
@ -47,8 +47,11 @@ logging.getLogger("sqlalchemy.engine.Engine").disabled = True
|
|||
# TODO use environment variable
|
||||
CONNECTION_URI = make_url(
|
||||
os.getenv(
|
||||
"CONNECTION_URI",
|
||||
"postgresql+psycopg://postgres:postgres@localhost:5432/postgres",
|
||||
"TEST_CONNECTION_URI",
|
||||
os.getenv(
|
||||
"CONNECTION_URI",
|
||||
"postgresql+psycopg://postgres:postgres@localhost:5432/postgres",
|
||||
),
|
||||
)
|
||||
)
|
||||
TEST_DB_URL = CONNECTION_URI.set(database="test_db")
|
||||
|
|
|
|||
|
|
@ -0,0 +1 @@
|
|||
"""Tests for the Honcho deriver system."""
|
||||
|
|
@ -0,0 +1,342 @@
|
|||
"""Fixtures and test configuration for deriver tests."""
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from nanoid import generate as generate_nanoid
|
||||
|
||||
from src import models
|
||||
from src.deriver.queue import QueueManager
|
||||
from src.deriver.tom.embeddings import CollectionEmbeddingStore
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_llm_responses():
|
||||
"""Provides structured mock responses for different LLM operations."""
|
||||
return {
|
||||
"fact_extraction": json.dumps({
|
||||
"facts": [
|
||||
"User is a software developer",
|
||||
"User works remotely",
|
||||
"User prefers coffee over tea",
|
||||
"User uses Python and JavaScript"
|
||||
]
|
||||
}),
|
||||
"fact_extraction_empty": json.dumps({"facts": []}),
|
||||
"fact_extraction_malformed": "This is not valid JSON",
|
||||
"tom_single_prompt": json.dumps({
|
||||
"confidence": 0.8,
|
||||
"user_representation": {
|
||||
"personality_traits": ["analytical", "detail-oriented", "collaborative"],
|
||||
"preferences": ["remote work", "technical discussions"],
|
||||
"communication_style": "direct and concise",
|
||||
"expertise_areas": ["software development", "Python programming"]
|
||||
}
|
||||
}),
|
||||
"tom_conversational": "Based on our conversation, I believe the user is a thoughtful software developer who values clear communication and technical excellence.",
|
||||
"summary_short": "User discussed Python development and remote work preferences.",
|
||||
"summary_long": "User is a software developer working remotely who has shown expertise in Python development and expressed preferences for asynchronous communication and detailed technical discussions."
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_embeddings():
|
||||
"""Provides mock embedding vectors for testing."""
|
||||
return {
|
||||
"fact_embedding": [0.1, 0.2, 0.3] + [0.0] * 1533, # 1536-dim vector
|
||||
"query_embedding": [0.15, 0.25, 0.35] + [0.0] * 1533,
|
||||
"duplicate_embedding": [0.1, 0.2, 0.3] + [0.0] * 1533, # Identical to fact_embedding
|
||||
"different_embedding": [0.9, 0.8, 0.7] + [0.0] * 1533
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def mock_model_clients(mock_llm_responses):
|
||||
"""Mock ModelClient instances for all TOM methods."""
|
||||
with (
|
||||
patch("src.deriver.tom.single_prompt.ModelClient") as mock_single_prompt_client,
|
||||
patch("src.deriver.tom.long_term.ModelClient") as mock_long_term_client,
|
||||
patch("src.utils.model_client.ModelClient") as mock_utils_client,
|
||||
):
|
||||
# Utils client for general use
|
||||
mock_utils_instance = AsyncMock()
|
||||
mock_utils_instance.generate.return_value = f"<facts>{mock_llm_responses['fact_extraction']}</facts>"
|
||||
mock_utils_client.return_value = mock_utils_instance
|
||||
|
||||
# Single prompt TOM client
|
||||
mock_single_prompt_instance = AsyncMock()
|
||||
mock_single_prompt_instance.generate.return_value = mock_llm_responses['tom_single_prompt']
|
||||
mock_single_prompt_client.return_value = mock_single_prompt_instance
|
||||
|
||||
# Long term TOM client
|
||||
mock_long_term_instance = AsyncMock()
|
||||
mock_long_term_instance.generate.return_value = f"<facts>{mock_llm_responses['fact_extraction']}</facts>"
|
||||
mock_long_term_client.return_value = mock_long_term_instance
|
||||
|
||||
yield {
|
||||
"utils": mock_utils_instance,
|
||||
"single_prompt": mock_single_prompt_instance,
|
||||
"long_term": mock_long_term_instance,
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def mock_vector_operations(mock_embeddings):
|
||||
"""Mock vector similarity operations and database queries."""
|
||||
with (
|
||||
patch("src.crud.get_duplicate_documents") as mock_get_duplicates,
|
||||
patch("src.crud.create_collection") as mock_create_collection,
|
||||
patch("src.crud.get_collection_by_name") as mock_get_collection,
|
||||
patch("src.crud.get_documents") as mock_get_documents,
|
||||
patch("src.crud.query_documents") as mock_query_documents,
|
||||
):
|
||||
# No duplicates by default
|
||||
mock_get_duplicates.return_value = []
|
||||
|
||||
# Mock collection operations
|
||||
mock_collection = MagicMock()
|
||||
mock_collection.public_id = str(uuid4())
|
||||
mock_create_collection.return_value = mock_collection
|
||||
mock_get_collection.return_value = mock_collection
|
||||
|
||||
# Mock document operations
|
||||
mock_get_documents.return_value = []
|
||||
mock_query_documents.return_value = []
|
||||
|
||||
yield {
|
||||
"get_duplicates": mock_get_duplicates,
|
||||
"create_collection": mock_create_collection,
|
||||
"get_collection": mock_get_collection,
|
||||
"get_documents": mock_get_documents,
|
||||
"query_documents": mock_query_documents
|
||||
}
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def queue_manager():
|
||||
"""Provides a QueueManager instance for testing."""
|
||||
# Mock the environment variable for test concurrency
|
||||
with patch("src.deriver.queue.os.getenv") as mock_getenv:
|
||||
mock_getenv.return_value = "2" # 2 workers for testing
|
||||
manager = QueueManager()
|
||||
yield manager
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def embedding_store(sample_data):
|
||||
"""Provides a CollectionEmbeddingStore for testing."""
|
||||
test_app, test_user = sample_data
|
||||
collection_id = str(uuid4())
|
||||
store = CollectionEmbeddingStore(test_app.public_id, test_user.public_id, collection_id)
|
||||
yield store
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def sample_messages(db_session, sample_data):
|
||||
"""Creates sample messages for testing deriver processing."""
|
||||
test_app, test_user = sample_data
|
||||
|
||||
# Create a test session
|
||||
session = models.Session(
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id,
|
||||
metadata={}
|
||||
)
|
||||
db_session.add(session)
|
||||
await db_session.flush()
|
||||
|
||||
# Create sample messages
|
||||
messages = []
|
||||
|
||||
# User message
|
||||
user_message = models.Message(
|
||||
session_id=session.public_id,
|
||||
is_user=True,
|
||||
content="I'm a Python developer working on AI projects. I prefer remote work and love debugging complex problems.",
|
||||
metadata={},
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id
|
||||
)
|
||||
db_session.add(user_message)
|
||||
messages.append(user_message)
|
||||
|
||||
# AI message
|
||||
ai_message = models.Message(
|
||||
session_id=session.public_id,
|
||||
is_user=False,
|
||||
content="That's great! Python is excellent for AI development. What specific AI frameworks do you work with?",
|
||||
metadata={},
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id
|
||||
)
|
||||
db_session.add(ai_message)
|
||||
messages.append(ai_message)
|
||||
|
||||
# Another user message
|
||||
user_message_2 = models.Message(
|
||||
session_id=session.public_id,
|
||||
is_user=True,
|
||||
content="I mainly use PyTorch and transformers. Currently building a chatbot with FastAPI.",
|
||||
metadata={},
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id
|
||||
)
|
||||
db_session.add(user_message_2)
|
||||
messages.append(user_message_2)
|
||||
|
||||
await db_session.flush()
|
||||
|
||||
yield session, messages
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def sample_queue_items(db_session, sample_messages):
|
||||
"""Creates sample queue items for testing queue processing."""
|
||||
session, messages = sample_messages
|
||||
|
||||
queue_items = []
|
||||
for message in messages:
|
||||
if message.is_user: # Only user messages get queued for processing
|
||||
queue_item = models.QueueItem(
|
||||
session_id=session.id, # Use integer ID, not public_id
|
||||
payload={"message_id": message.public_id},
|
||||
processed=False
|
||||
)
|
||||
db_session.add(queue_item)
|
||||
queue_items.append(queue_item)
|
||||
|
||||
await db_session.flush()
|
||||
yield session, messages, queue_items
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def sample_facts(db_session, sample_data):
|
||||
"""Creates sample user facts stored in collections."""
|
||||
test_app, test_user = sample_data
|
||||
|
||||
# Create user collection
|
||||
collection = models.Collection(
|
||||
app_id=test_app.public_id,
|
||||
user_id=test_user.public_id,
|
||||
name=f"user_{test_user.public_id}",
|
||||
metadata={"type": "user_facts"}
|
||||
)
|
||||
db_session.add(collection)
|
||||
await db_session.flush()
|
||||
|
||||
# Create sample documents (facts)
|
||||
facts = [
|
||||
"User is a Python developer",
|
||||
"User works remotely",
|
||||
"User prefers PyTorch over TensorFlow",
|
||||
"User enjoys debugging complex problems"
|
||||
]
|
||||
|
||||
documents = []
|
||||
for fact in facts:
|
||||
doc = models.Document(
|
||||
collection_id=collection.public_id,
|
||||
content=fact,
|
||||
metadata={"extracted_at": "2024-01-01T00:00:00Z"},
|
||||
embedding=[0.1] * 1536 # Mock embedding
|
||||
)
|
||||
db_session.add(doc)
|
||||
documents.append(doc)
|
||||
|
||||
await db_session.flush()
|
||||
yield collection, documents
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_background_task():
|
||||
"""Mock FastAPI BackgroundTasks for testing async task scheduling."""
|
||||
with patch("fastapi.BackgroundTasks") as mock_bg:
|
||||
mock_task = MagicMock()
|
||||
mock_bg.return_value = mock_task
|
||||
yield mock_task
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_signal_handling():
|
||||
"""Mock signal handling for testing graceful shutdown."""
|
||||
with (
|
||||
patch("signal.signal") as mock_signal,
|
||||
patch("signal.SIGTERM") as mock_sigterm,
|
||||
patch("signal.SIGINT") as mock_sigint,
|
||||
):
|
||||
yield {
|
||||
"signal": mock_signal,
|
||||
"sigterm": mock_sigterm,
|
||||
"sigint": mock_sigint
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def mock_semaphore():
|
||||
"""Mock asyncio.Semaphore for testing concurrency control."""
|
||||
mock_sem = AsyncMock()
|
||||
mock_sem.__aenter__ = AsyncMock(return_value=mock_sem)
|
||||
mock_sem.__aexit__ = AsyncMock(return_value=None)
|
||||
|
||||
with patch("asyncio.Semaphore") as mock_semaphore_class:
|
||||
mock_semaphore_class.return_value = mock_sem
|
||||
yield mock_sem
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def tom_method_config():
|
||||
"""Fixture for testing different TOM method configurations."""
|
||||
return {
|
||||
"single_prompt": "SINGLE_PROMPT",
|
||||
"conversational": "CONVERSATIONAL",
|
||||
"long_term": "LONG_TERM"
|
||||
}
|
||||
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def mock_session_processing():
|
||||
"""Mock session processing for integration tests."""
|
||||
with (
|
||||
patch("src.deriver.consumer.process_user_message") as mock_process_user,
|
||||
patch("src.deriver.consumer.process_ai_message") as mock_process_ai,
|
||||
patch("src.deriver.consumer.maybe_create_summary") as mock_create_summary,
|
||||
):
|
||||
mock_process_user.return_value = None
|
||||
mock_process_ai.return_value = None
|
||||
mock_create_summary.return_value = None
|
||||
|
||||
yield {
|
||||
"process_user": mock_process_user,
|
||||
"process_ai": mock_process_ai,
|
||||
"create_summary": mock_create_summary
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def performance_config():
|
||||
"""Configuration for performance testing."""
|
||||
return {
|
||||
"max_workers": 4,
|
||||
"message_count": 100,
|
||||
"session_count": 10,
|
||||
"timeout_seconds": 30,
|
||||
"fact_extraction_time_limit": 5.0, # seconds
|
||||
"tom_inference_time_limit": 3.0, # seconds
|
||||
"queue_processing_time_limit": 1.0 # seconds per message
|
||||
}
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def error_scenarios():
|
||||
"""Provides various error scenarios for testing error handling."""
|
||||
return {
|
||||
"llm_timeout": "LLM request timed out",
|
||||
"llm_api_error": "API rate limit exceeded",
|
||||
"database_connection_error": "Database connection failed",
|
||||
"invalid_json_response": "Malformed JSON in LLM response",
|
||||
"embedding_api_error": "Embedding service unavailable",
|
||||
"vector_similarity_error": "Vector similarity calculation failed"
|
||||
}
|
||||
|
|
@ -0,0 +1,632 @@
|
|||
"""End-to-end integration tests for the complete deriver workflow."""
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
|
||||
from src import models
|
||||
from src.deriver import consumer
|
||||
from src.deriver.queue import QueueManager
|
||||
|
||||
|
||||
class TestMessageToFactsWorkflow:
|
||||
"""Test the complete workflow from message creation to fact storage."""
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def integration_setup(self, db_session, sample_data):
|
||||
"""Setup complete integration test data."""
|
||||
test_app, test_user = sample_data
|
||||
|
||||
# Create a session
|
||||
session = models.Session(
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id,
|
||||
metadata={}
|
||||
)
|
||||
db_session.add(session)
|
||||
await db_session.flush()
|
||||
|
||||
# Create user collection for fact storage
|
||||
collection = models.Collection(
|
||||
app_id=test_app.public_id,
|
||||
user_id=test_user.public_id,
|
||||
name=f"user_{test_user.public_id}",
|
||||
metadata={"type": "user_facts"}
|
||||
)
|
||||
db_session.add(collection)
|
||||
await db_session.flush()
|
||||
|
||||
# Create user messages
|
||||
messages = []
|
||||
for i, content in enumerate([
|
||||
"Hi, I'm Sarah, a data scientist working remotely from Seattle",
|
||||
"I've been using Python for machine learning for about 3 years",
|
||||
"My current project involves building recommendation systems with PyTorch"
|
||||
]):
|
||||
message = models.Message(
|
||||
session_id=session.public_id,
|
||||
is_user=True,
|
||||
content=content,
|
||||
metadata={},
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id
|
||||
)
|
||||
db_session.add(message)
|
||||
messages.append(message)
|
||||
|
||||
await db_session.flush()
|
||||
|
||||
return test_app, test_user, session, collection, messages
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_complete_message_processing_workflow(self, db_session, integration_setup):
|
||||
"""Test complete workflow: message → fact extraction → vector storage."""
|
||||
test_app, test_user, session, collection, messages = integration_setup
|
||||
|
||||
# Mock the fact extraction to return realistic facts
|
||||
extracted_facts = [
|
||||
"User name is Sarah",
|
||||
"User is a data scientist",
|
||||
"User works remotely from Seattle",
|
||||
"User has 3 years of Python experience",
|
||||
"User specializes in machine learning",
|
||||
"User is currently working on recommendation systems",
|
||||
"User uses PyTorch for current project"
|
||||
]
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.extract_facts_long_term") as mock_extract,
|
||||
patch("src.deriver.consumer.history.get_summarized_history") as mock_history,
|
||||
patch("src.deriver.consumer.crud.get_or_create_user_protected_collection") as mock_get_collection,
|
||||
patch("src.deriver.consumer.CollectionEmbeddingStore") as mock_store_class
|
||||
):
|
||||
# Setup mocks
|
||||
mock_extract.return_value = extracted_facts
|
||||
mock_history.return_value = ("Previous conversation context", [], None)
|
||||
mock_get_collection.return_value = collection
|
||||
|
||||
# Mock embedding store
|
||||
mock_store = AsyncMock()
|
||||
mock_store.remove_duplicates.return_value = extracted_facts # No duplicates
|
||||
mock_store.save_facts.return_value = None
|
||||
mock_store_class.return_value = mock_store
|
||||
|
||||
# Process each message
|
||||
for message in messages:
|
||||
await consumer.process_user_message(
|
||||
message.content,
|
||||
test_app.public_id,
|
||||
test_user.public_id,
|
||||
session.public_id,
|
||||
message.public_id,
|
||||
db_session
|
||||
)
|
||||
|
||||
# Verify fact extraction was called for each message
|
||||
assert mock_extract.call_count == 3
|
||||
|
||||
# Verify facts were saved for each message
|
||||
assert mock_store.save_facts.call_count == 3
|
||||
|
||||
# Verify the facts that would be saved
|
||||
all_saved_facts = []
|
||||
for call in mock_store.save_facts.call_args_list:
|
||||
facts_arg = call[0][0] # First positional argument
|
||||
all_saved_facts.extend(facts_arg)
|
||||
|
||||
# Should have saved all extracted facts
|
||||
assert len(all_saved_facts) == len(extracted_facts) * 3
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_queue_to_consumer_integration(self, db_session, integration_setup):
|
||||
"""Test integration between queue management and message processing."""
|
||||
test_app, test_user, session, collection, messages = integration_setup
|
||||
|
||||
# Create queue items for user messages (simulating how they're created in real system)
|
||||
queue_items = []
|
||||
for message in messages:
|
||||
if message.is_user:
|
||||
# Create queue item with the payload structure used in real system
|
||||
payload = {
|
||||
"message_id": message.public_id,
|
||||
"is_user": message.is_user,
|
||||
"content": message.content,
|
||||
"app_id": test_app.public_id,
|
||||
"user_id": test_user.public_id,
|
||||
"session_id": session.public_id
|
||||
}
|
||||
|
||||
queue_item = models.QueueItem(
|
||||
session_id=session.id, # Use integer ID for queue
|
||||
payload=payload,
|
||||
processed=False
|
||||
)
|
||||
db_session.add(queue_item)
|
||||
queue_items.append(queue_item)
|
||||
|
||||
await db_session.flush()
|
||||
|
||||
# Mock the consumer processing functions
|
||||
with (
|
||||
patch("src.deriver.consumer.process_user_message") as mock_process_user,
|
||||
patch("src.deriver.consumer.process_ai_message") as mock_process_ai,
|
||||
patch("src.deriver.consumer.summarize_if_needed") as mock_summarize
|
||||
):
|
||||
mock_process_user.return_value = None
|
||||
mock_process_ai.return_value = None
|
||||
mock_summarize.return_value = None
|
||||
|
||||
# Process items through the consumer
|
||||
for queue_item in queue_items:
|
||||
await consumer.process_item(db_session, queue_item.payload)
|
||||
|
||||
# Verify all user messages were processed
|
||||
assert mock_process_user.call_count == 3
|
||||
assert mock_process_ai.call_count == 0 # No AI messages
|
||||
assert mock_summarize.call_count == 3 # Summary check for each message
|
||||
|
||||
# Verify the arguments passed to process_user_message
|
||||
for i, call in enumerate(mock_process_user.call_args_list):
|
||||
args = call[0]
|
||||
assert args[0] == messages[i].content # content
|
||||
assert args[1] == test_app.public_id # app_id
|
||||
assert args[2] == test_user.public_id # user_id
|
||||
assert args[3] == session.public_id # session_id
|
||||
assert args[4] == messages[i].public_id # message_id
|
||||
assert args[5] == db_session # db_session
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_tom_inference_integration(self, db_session, integration_setup, mock_llm_responses):
|
||||
"""Test integration of TOM inference with fact extraction workflow."""
|
||||
test_app, test_user, session, collection, messages = integration_setup
|
||||
|
||||
# Test facts that would be extracted
|
||||
user_facts = [
|
||||
"User name is Sarah",
|
||||
"User is a data scientist",
|
||||
"User works remotely"
|
||||
]
|
||||
|
||||
with (
|
||||
patch("src.deriver.tom.get_tom_inference") as mock_tom_inference,
|
||||
patch("src.deriver.tom.get_user_representation") as mock_user_rep,
|
||||
patch("src.deriver.consumer.extract_facts_long_term") as mock_extract_facts,
|
||||
patch("src.deriver.consumer.history.get_summarized_history") as mock_history,
|
||||
patch("src.deriver.consumer.crud.get_or_create_user_protected_collection") as mock_get_collection,
|
||||
patch("src.deriver.consumer.CollectionEmbeddingStore") as mock_store_class
|
||||
):
|
||||
# Setup TOM mocks
|
||||
mock_tom_inference.return_value = mock_llm_responses['tom_single_prompt']
|
||||
mock_user_rep.return_value = mock_llm_responses['tom_single_prompt']
|
||||
|
||||
# Setup fact extraction mocks
|
||||
mock_extract_facts.return_value = user_facts
|
||||
mock_history.return_value = ("Chat history", [], None)
|
||||
mock_get_collection.return_value = collection
|
||||
|
||||
# Mock embedding store
|
||||
mock_store = AsyncMock()
|
||||
mock_store.remove_duplicates.return_value = user_facts
|
||||
mock_store.save_facts.return_value = None
|
||||
mock_store_class.return_value = mock_store
|
||||
|
||||
# Process a user message
|
||||
await consumer.process_user_message(
|
||||
messages[0].content,
|
||||
test_app.public_id,
|
||||
test_user.public_id,
|
||||
session.public_id,
|
||||
messages[0].public_id,
|
||||
db_session
|
||||
)
|
||||
|
||||
# Verify fact extraction occurred
|
||||
mock_extract_facts.assert_called_once()
|
||||
mock_store.save_facts.assert_called_once_with(user_facts, message_id=messages[0].public_id)
|
||||
|
||||
# The TOM inference methods aren't called directly in consumer,
|
||||
# but we've verified the infrastructure is in place
|
||||
|
||||
|
||||
class TestErrorRecoveryIntegration:
|
||||
"""Test error recovery and resilience in integrated workflows."""
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def error_test_setup(self, db_session, sample_data):
|
||||
"""Setup data for error testing."""
|
||||
test_app, test_user = sample_data
|
||||
|
||||
session = models.Session(
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id,
|
||||
metadata={}
|
||||
)
|
||||
db_session.add(session)
|
||||
await db_session.flush()
|
||||
|
||||
message = models.Message(
|
||||
session_id=session.public_id,
|
||||
is_user=True,
|
||||
content="Test message for error scenarios",
|
||||
metadata={},
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id
|
||||
)
|
||||
db_session.add(message)
|
||||
await db_session.flush()
|
||||
|
||||
return test_app, test_user, session, message
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fact_extraction_error_recovery(self, db_session, error_test_setup):
|
||||
"""Test that fact extraction errors don't break the entire workflow."""
|
||||
test_app, test_user, session, message = error_test_setup
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.extract_facts_long_term") as mock_extract,
|
||||
patch("src.deriver.consumer.history.get_summarized_history") as mock_history,
|
||||
patch("src.deriver.consumer.crud.get_or_create_user_protected_collection") as mock_get_collection
|
||||
):
|
||||
# Mock fact extraction to fail
|
||||
mock_extract.side_effect = Exception("LLM API timeout")
|
||||
mock_history.return_value = ("", [], None)
|
||||
|
||||
# Mock collection to avoid that error
|
||||
mock_collection = MagicMock()
|
||||
mock_collection.public_id = str(uuid4())
|
||||
mock_get_collection.return_value = mock_collection
|
||||
|
||||
# Should raise the exception (let caller handle it)
|
||||
with pytest.raises(Exception, match="LLM API timeout"):
|
||||
await consumer.process_user_message(
|
||||
message.content,
|
||||
test_app.public_id,
|
||||
test_user.public_id,
|
||||
session.public_id,
|
||||
message.public_id,
|
||||
db_session
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_partial_fact_storage_error_recovery(self, db_session, error_test_setup):
|
||||
"""Test recovery when some facts fail to store."""
|
||||
test_app, test_user, session, message = error_test_setup
|
||||
|
||||
facts_to_extract = [
|
||||
"User likes programming",
|
||||
"This fact will fail to store",
|
||||
"User works in tech"
|
||||
]
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.extract_facts_long_term") as mock_extract,
|
||||
patch("src.deriver.consumer.history.get_summarized_history") as mock_history,
|
||||
patch("src.deriver.consumer.crud.get_or_create_user_protected_collection") as mock_get_collection,
|
||||
patch("src.deriver.consumer.CollectionEmbeddingStore") as mock_store_class
|
||||
):
|
||||
mock_extract.return_value = facts_to_extract
|
||||
mock_history.return_value = ("", [], None)
|
||||
|
||||
mock_collection = MagicMock()
|
||||
mock_collection.public_id = str(uuid4())
|
||||
mock_get_collection.return_value = mock_collection
|
||||
|
||||
# Mock embedding store where save_facts has partial failure
|
||||
mock_store = AsyncMock()
|
||||
mock_store.remove_duplicates.return_value = facts_to_extract
|
||||
# save_facts method handles its own errors gracefully
|
||||
mock_store.save_facts.return_value = None
|
||||
mock_store_class.return_value = mock_store
|
||||
|
||||
# Should complete successfully even with partial failures
|
||||
await consumer.process_user_message(
|
||||
message.content,
|
||||
test_app.public_id,
|
||||
test_user.public_id,
|
||||
session.public_id,
|
||||
message.public_id,
|
||||
db_session
|
||||
)
|
||||
|
||||
# Verify the workflow completed
|
||||
mock_extract.assert_called_once()
|
||||
mock_store.save_facts.assert_called_once()
|
||||
|
||||
|
||||
class TestSummaryIntegration:
|
||||
"""Test summary generation integration with message processing."""
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def summary_test_setup(self, db_session, sample_data):
|
||||
"""Setup data for summary testing."""
|
||||
test_app, test_user = sample_data
|
||||
|
||||
session = models.Session(
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id,
|
||||
metadata={}
|
||||
)
|
||||
db_session.add(session)
|
||||
await db_session.flush()
|
||||
|
||||
return test_app, test_user, session
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_summary_generation_integration(self, db_session, summary_test_setup):
|
||||
"""Test that summary generation integrates properly with message processing."""
|
||||
test_app, test_user, session = summary_test_setup
|
||||
|
||||
# Create enough messages to trigger summary generation
|
||||
messages = []
|
||||
for i in range(25): # Enough to trigger both short and long summaries
|
||||
message = models.Message(
|
||||
session_id=session.public_id,
|
||||
is_user=True,
|
||||
content=f"Message {i+1}: User discussing various topics",
|
||||
metadata={},
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id
|
||||
)
|
||||
db_session.add(message)
|
||||
messages.append(message)
|
||||
|
||||
await db_session.flush()
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.extract_facts_long_term") as mock_extract,
|
||||
patch("src.deriver.consumer.history.get_summarized_history") as mock_history,
|
||||
patch("src.deriver.consumer.history.should_create_summary") as mock_should_create,
|
||||
patch("src.deriver.consumer.history.create_summary") as mock_create_summary,
|
||||
patch("src.deriver.consumer.history.save_summary_metamessage") as mock_save_summary,
|
||||
patch("src.deriver.consumer.crud.get_or_create_user_protected_collection") as mock_get_collection,
|
||||
patch("src.deriver.consumer.CollectionEmbeddingStore") as mock_store_class
|
||||
):
|
||||
# Setup mocks
|
||||
mock_extract.return_value = ["User is engaged in conversation"]
|
||||
mock_history.return_value = ("Previous context", [], None)
|
||||
|
||||
mock_collection = MagicMock()
|
||||
mock_collection.public_id = str(uuid4())
|
||||
mock_get_collection.return_value = mock_collection
|
||||
|
||||
mock_store = AsyncMock()
|
||||
mock_store.remove_duplicates.return_value = ["User is engaged"]
|
||||
mock_store.save_facts.return_value = None
|
||||
mock_store_class.return_value = mock_store
|
||||
|
||||
# Mock summary creation - simulate that summaries are needed
|
||||
mock_should_create.return_value = (True, messages[:10], None)
|
||||
mock_create_summary.return_value = "Summary of recent messages"
|
||||
mock_save_summary.return_value = None
|
||||
|
||||
# Process the last message (which should trigger summary check)
|
||||
await consumer.process_item(db_session, {
|
||||
"message_id": messages[-1].public_id,
|
||||
"is_user": True,
|
||||
"content": messages[-1].content,
|
||||
"app_id": test_app.public_id,
|
||||
"user_id": test_user.public_id,
|
||||
"session_id": session.public_id
|
||||
})
|
||||
|
||||
# Verify fact extraction occurred
|
||||
mock_extract.assert_called_once()
|
||||
|
||||
# Verify summary generation was checked
|
||||
mock_should_create.assert_called()
|
||||
|
||||
# If summaries were triggered, verify they were created
|
||||
if mock_should_create.call_count > 0:
|
||||
# Summary creation logic was invoked
|
||||
assert True # Successfully integrated
|
||||
|
||||
|
||||
class TestConcurrentProcessing:
|
||||
"""Test concurrent processing scenarios."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_message_processing(self, db_session, sample_data):
|
||||
"""Test that multiple messages can be processed concurrently safely."""
|
||||
test_app, test_user = sample_data
|
||||
|
||||
# Create multiple sessions
|
||||
sessions = []
|
||||
for i in range(3):
|
||||
session = models.Session(
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id,
|
||||
metadata={"session_num": i}
|
||||
)
|
||||
db_session.add(session)
|
||||
sessions.append(session)
|
||||
|
||||
await db_session.flush()
|
||||
|
||||
# Create messages for each session
|
||||
all_messages = []
|
||||
for i, session in enumerate(sessions):
|
||||
message = models.Message(
|
||||
session_id=session.public_id,
|
||||
is_user=True,
|
||||
content=f"Session {i} message: User sharing information",
|
||||
metadata={},
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id
|
||||
)
|
||||
db_session.add(message)
|
||||
all_messages.append(message)
|
||||
|
||||
await db_session.flush()
|
||||
|
||||
processed_count = 0
|
||||
|
||||
async def mock_process_user_message(*args, **kwargs):
|
||||
nonlocal processed_count
|
||||
processed_count += 1
|
||||
# Simulate some processing time
|
||||
import asyncio
|
||||
await asyncio.sleep(0.01)
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.process_user_message", side_effect=mock_process_user_message),
|
||||
patch("src.deriver.consumer.process_ai_message") as mock_process_ai,
|
||||
patch("src.deriver.consumer.summarize_if_needed") as mock_summarize
|
||||
):
|
||||
mock_process_ai.return_value = None
|
||||
mock_summarize.return_value = None
|
||||
|
||||
# Process all messages concurrently
|
||||
import asyncio
|
||||
tasks = []
|
||||
for message in all_messages:
|
||||
payload = {
|
||||
"message_id": message.public_id,
|
||||
"is_user": True,
|
||||
"content": message.content,
|
||||
"app_id": test_app.public_id,
|
||||
"user_id": test_user.public_id,
|
||||
"session_id": session.public_id
|
||||
}
|
||||
task = asyncio.create_task(consumer.process_item(db_session, payload))
|
||||
tasks.append(task)
|
||||
|
||||
# Wait for all processing to complete
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
# Verify all messages were processed
|
||||
assert processed_count == 3
|
||||
|
||||
|
||||
class TestFullSystemIntegration:
|
||||
"""Test complete system integration from API to storage."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_realistic_user_conversation_workflow(self, db_session, sample_data):
|
||||
"""Test a realistic user conversation workflow end-to-end."""
|
||||
test_app, test_user = sample_data
|
||||
|
||||
# Create session
|
||||
session = models.Session(
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id,
|
||||
metadata={"conversation_type": "onboarding"}
|
||||
)
|
||||
db_session.add(session)
|
||||
await db_session.flush()
|
||||
|
||||
# Realistic conversation messages
|
||||
conversation = [
|
||||
("user", "Hi! I'm Alex, a software engineer based in San Francisco."),
|
||||
("ai", "Hello Alex! It's nice to meet you. What kind of software engineering do you focus on?"),
|
||||
("user", "I mainly work on backend systems using Python and Go. Currently building microservices for a fintech company."),
|
||||
("ai", "That sounds interesting! Fintech is such a dynamic field. What's the most challenging part of your current project?"),
|
||||
("user", "The main challenge is handling high-frequency trading data while maintaining low latency. We're processing millions of transactions per second."),
|
||||
("ai", "That's impressive scale! Are you using any specific technologies for handling that throughput?"),
|
||||
("user", "Yes, we're using Kafka for streaming, Redis for caching, and PostgreSQL with read replicas. Also experimenting with some Rust components for ultra-low latency parts.")
|
||||
]
|
||||
|
||||
# Create all messages
|
||||
messages = []
|
||||
for role, content in conversation:
|
||||
message = models.Message(
|
||||
session_id=session.public_id,
|
||||
is_user=(role == "user"),
|
||||
content=content,
|
||||
metadata={},
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id
|
||||
)
|
||||
db_session.add(message)
|
||||
messages.append(message)
|
||||
|
||||
await db_session.flush()
|
||||
|
||||
# Expected facts that would be extracted
|
||||
expected_facts = [
|
||||
"User name is Alex",
|
||||
"User is a software engineer",
|
||||
"User is based in San Francisco",
|
||||
"User works on backend systems",
|
||||
"User uses Python and Go",
|
||||
"User works at a fintech company",
|
||||
"User builds microservices",
|
||||
"User handles high-frequency trading data",
|
||||
"User processes millions of transactions per second",
|
||||
"User uses Kafka for streaming",
|
||||
"User uses Redis for caching",
|
||||
"User uses PostgreSQL with read replicas",
|
||||
"User is experimenting with Rust components"
|
||||
]
|
||||
|
||||
# Track all extracted facts
|
||||
all_extracted_facts = []
|
||||
|
||||
def mock_extract_facts(chat_history):
|
||||
# Simulate realistic fact extraction based on content
|
||||
if "Alex" in chat_history and "software engineer" in chat_history:
|
||||
return ["User name is Alex", "User is a software engineer", "User is based in San Francisco"]
|
||||
elif "Python and Go" in chat_history:
|
||||
return ["User works on backend systems", "User uses Python and Go", "User works at a fintech company"]
|
||||
elif "Kafka" in chat_history:
|
||||
return ["User uses Kafka for streaming", "User uses Redis for caching", "User uses PostgreSQL"]
|
||||
return []
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.extract_facts_long_term", side_effect=mock_extract_facts),
|
||||
patch("src.deriver.consumer.history.get_summarized_history") as mock_history,
|
||||
patch("src.deriver.consumer.crud.get_or_create_user_protected_collection") as mock_get_collection,
|
||||
patch("src.deriver.consumer.CollectionEmbeddingStore") as mock_store_class,
|
||||
patch("src.deriver.consumer.summarize_if_needed") as mock_summarize
|
||||
):
|
||||
mock_history.return_value = ("", [], None)
|
||||
|
||||
mock_collection = MagicMock()
|
||||
mock_collection.public_id = str(uuid4())
|
||||
mock_get_collection.return_value = mock_collection
|
||||
|
||||
# Track saved facts
|
||||
saved_facts = []
|
||||
def track_save_facts(facts, **kwargs):
|
||||
saved_facts.extend(facts)
|
||||
|
||||
mock_store = AsyncMock()
|
||||
mock_store.remove_duplicates.side_effect = lambda facts: facts # No duplicates
|
||||
mock_store.save_facts.side_effect = track_save_facts
|
||||
mock_store_class.return_value = mock_store
|
||||
|
||||
mock_summarize.return_value = None
|
||||
|
||||
# Process only user messages (as would happen in real system)
|
||||
user_messages = [msg for msg in messages if msg.is_user]
|
||||
|
||||
for message in user_messages:
|
||||
await consumer.process_user_message(
|
||||
message.content,
|
||||
test_app.public_id,
|
||||
test_user.public_id,
|
||||
session.public_id,
|
||||
message.public_id,
|
||||
db_session
|
||||
)
|
||||
|
||||
# Verify facts were extracted and saved
|
||||
assert len(saved_facts) > 0
|
||||
|
||||
# Verify user-specific facts were captured
|
||||
saved_facts_str = " ".join(saved_facts)
|
||||
assert "Alex" in saved_facts_str
|
||||
assert "software engineer" in saved_facts_str
|
||||
assert "San Francisco" in saved_facts_str
|
||||
|
||||
# Verify technical details were captured
|
||||
tech_keywords = ["Python", "Go", "fintech", "Kafka", "Redis", "PostgreSQL"]
|
||||
captured_tech = [kw for kw in tech_keywords if kw in saved_facts_str]
|
||||
assert len(captured_tech) > 0
|
||||
|
||||
print(f"✅ Integration test completed successfully!")
|
||||
print(f"📊 Processed {len(user_messages)} user messages")
|
||||
print(f"💾 Saved {len(saved_facts)} facts total")
|
||||
print(f"🔧 Captured {len(captured_tech)} technical details")
|
||||
|
|
@ -0,0 +1,588 @@
|
|||
"""Tests for the consumer module and message processing functionality."""
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
|
||||
from src import models
|
||||
from src.deriver import consumer
|
||||
from src.utils.history import SummaryType
|
||||
|
||||
|
||||
class TestProcessItem:
|
||||
"""Test the main process_item entry point."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_item_routes_user_message(self, db_session):
|
||||
"""Test that process_item correctly routes user messages."""
|
||||
payload = {
|
||||
"message_id": str(uuid4()),
|
||||
"is_user": True,
|
||||
"content": "Hello, I'm a Python developer",
|
||||
"app_id": str(uuid4()),
|
||||
"user_id": str(uuid4()),
|
||||
"session_id": str(uuid4())
|
||||
}
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.process_user_message") as mock_process_user,
|
||||
patch("src.deriver.consumer.process_ai_message") as mock_process_ai,
|
||||
patch("src.deriver.consumer.summarize_if_needed") as mock_summarize
|
||||
):
|
||||
mock_process_user.return_value = None
|
||||
mock_process_ai.return_value = None
|
||||
mock_summarize.return_value = None
|
||||
|
||||
await consumer.process_item(db_session, payload)
|
||||
|
||||
# Should call user message processing
|
||||
mock_process_user.assert_called_once_with(
|
||||
payload["content"],
|
||||
payload["app_id"],
|
||||
payload["user_id"],
|
||||
payload["session_id"],
|
||||
payload["message_id"],
|
||||
db_session
|
||||
)
|
||||
mock_process_ai.assert_not_called()
|
||||
mock_summarize.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_item_routes_ai_message(self, db_session):
|
||||
"""Test that process_item correctly routes AI messages."""
|
||||
payload = {
|
||||
"message_id": str(uuid4()),
|
||||
"is_user": False,
|
||||
"content": "I can help you with Python development!",
|
||||
"app_id": str(uuid4()),
|
||||
"user_id": str(uuid4()),
|
||||
"session_id": str(uuid4())
|
||||
}
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.process_user_message") as mock_process_user,
|
||||
patch("src.deriver.consumer.process_ai_message") as mock_process_ai,
|
||||
patch("src.deriver.consumer.summarize_if_needed") as mock_summarize
|
||||
):
|
||||
mock_process_user.return_value = None
|
||||
mock_process_ai.return_value = None
|
||||
mock_summarize.return_value = None
|
||||
|
||||
await consumer.process_item(db_session, payload)
|
||||
|
||||
# Should call AI message processing
|
||||
mock_process_ai.assert_called_once_with(
|
||||
payload["content"],
|
||||
payload["app_id"],
|
||||
payload["user_id"],
|
||||
payload["session_id"],
|
||||
payload["message_id"],
|
||||
db_session
|
||||
)
|
||||
mock_process_user.assert_not_called()
|
||||
mock_summarize.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_item_calls_summarize(self, db_session):
|
||||
"""Test that process_item always calls summarize_if_needed."""
|
||||
payload = {
|
||||
"message_id": str(uuid4()),
|
||||
"is_user": True,
|
||||
"content": "Test message",
|
||||
"app_id": str(uuid4()),
|
||||
"user_id": str(uuid4()),
|
||||
"session_id": str(uuid4())
|
||||
}
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.process_user_message") as mock_process_user,
|
||||
patch("src.deriver.consumer.summarize_if_needed") as mock_summarize
|
||||
):
|
||||
mock_process_user.return_value = None
|
||||
mock_summarize.return_value = None
|
||||
|
||||
await consumer.process_item(db_session, payload)
|
||||
|
||||
mock_summarize.assert_called_once_with(
|
||||
db_session,
|
||||
payload["app_id"],
|
||||
payload["session_id"],
|
||||
payload["user_id"],
|
||||
payload["message_id"]
|
||||
)
|
||||
|
||||
|
||||
class TestProcessUserMessage:
|
||||
"""Test user message processing functionality."""
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def setup_user_message_test(self, db_session, sample_data):
|
||||
"""Setup test data for user message processing."""
|
||||
test_app, test_user = sample_data
|
||||
|
||||
# Create a session
|
||||
session = models.Session(
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id,
|
||||
metadata={}
|
||||
)
|
||||
db_session.add(session)
|
||||
await db_session.flush()
|
||||
|
||||
# Create a collection for the user
|
||||
collection = models.Collection(
|
||||
app_id=test_app.public_id,
|
||||
user_id=test_user.public_id,
|
||||
name=f"user_{test_user.public_id}",
|
||||
metadata={"type": "user_facts"}
|
||||
)
|
||||
db_session.add(collection)
|
||||
await db_session.flush()
|
||||
|
||||
return test_app, test_user, session, collection
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_user_message_extracts_and_saves_facts(self, db_session, setup_user_message_test, mock_llm_responses):
|
||||
"""Test that user message processing extracts and saves facts."""
|
||||
test_app, test_user, session, collection = setup_user_message_test
|
||||
|
||||
message_content = "I'm a Python developer who works remotely and loves coffee"
|
||||
message_id = str(uuid4())
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.history.get_summarized_history") as mock_get_history,
|
||||
patch("src.deriver.consumer.extract_facts_long_term") as mock_extract_facts,
|
||||
patch("src.deriver.consumer.crud.get_or_create_user_protected_collection") as mock_get_collection,
|
||||
patch("src.deriver.consumer.CollectionEmbeddingStore") as mock_embedding_store_class
|
||||
):
|
||||
# Mock history retrieval
|
||||
mock_get_history.return_value = ("Previous chat", [], None)
|
||||
|
||||
# Mock fact extraction
|
||||
mock_extract_facts.return_value = [
|
||||
"User is a Python developer",
|
||||
"User works remotely",
|
||||
"User loves coffee"
|
||||
]
|
||||
|
||||
# Mock collection retrieval
|
||||
mock_get_collection.return_value = collection
|
||||
|
||||
# Mock embedding store
|
||||
mock_embedding_store = AsyncMock()
|
||||
mock_embedding_store.remove_duplicates.return_value = [
|
||||
"User is a Python developer",
|
||||
"User works remotely"
|
||||
] # Simulate one duplicate removed
|
||||
mock_embedding_store.save_facts.return_value = None
|
||||
mock_embedding_store_class.return_value = mock_embedding_store
|
||||
|
||||
# Process the user message
|
||||
await consumer.process_user_message(
|
||||
message_content,
|
||||
test_app.public_id,
|
||||
test_user.public_id,
|
||||
session.public_id,
|
||||
message_id,
|
||||
db_session
|
||||
)
|
||||
|
||||
# Verify the flow
|
||||
mock_get_history.assert_called_once_with(
|
||||
db_session, session.public_id, summary_type=SummaryType.SHORT
|
||||
)
|
||||
mock_extract_facts.assert_called_once()
|
||||
mock_get_collection.assert_called_once_with(
|
||||
db=db_session, app_id=test_app.public_id, user_id=test_user.public_id
|
||||
)
|
||||
mock_embedding_store.remove_duplicates.assert_called_once_with([
|
||||
"User is a Python developer",
|
||||
"User works remotely",
|
||||
"User loves coffee"
|
||||
])
|
||||
mock_embedding_store.save_facts.assert_called_once_with(
|
||||
["User is a Python developer", "User works remotely"],
|
||||
message_id=message_id
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_user_message_no_unique_facts(self, db_session, setup_user_message_test):
|
||||
"""Test user message processing when all facts are duplicates."""
|
||||
test_app, test_user, session, collection = setup_user_message_test
|
||||
|
||||
message_content = "I still love Python programming"
|
||||
message_id = str(uuid4())
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.history.get_summarized_history") as mock_get_history,
|
||||
patch("src.deriver.consumer.extract_facts_long_term") as mock_extract_facts,
|
||||
patch("src.deriver.consumer.crud.get_or_create_user_protected_collection") as mock_get_collection,
|
||||
patch("src.deriver.consumer.CollectionEmbeddingStore") as mock_embedding_store_class
|
||||
):
|
||||
mock_get_history.return_value = ("Previous chat", [], None)
|
||||
mock_extract_facts.return_value = ["User loves Python"]
|
||||
mock_get_collection.return_value = collection
|
||||
|
||||
# Mock embedding store to return no unique facts
|
||||
mock_embedding_store = AsyncMock()
|
||||
mock_embedding_store.remove_duplicates.return_value = [] # All duplicates
|
||||
mock_embedding_store_class.return_value = mock_embedding_store
|
||||
|
||||
await consumer.process_user_message(
|
||||
message_content,
|
||||
test_app.public_id,
|
||||
test_user.public_id,
|
||||
session.public_id,
|
||||
message_id,
|
||||
db_session
|
||||
)
|
||||
|
||||
# Should not call save_facts when no unique facts
|
||||
mock_embedding_store.save_facts.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_user_message_with_chat_history(self, db_session, setup_user_message_test):
|
||||
"""Test that chat history is properly included in fact extraction."""
|
||||
test_app, test_user, session, collection = setup_user_message_test
|
||||
|
||||
message_content = "I prefer PyTorch over TensorFlow"
|
||||
message_id = str(uuid4())
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.history.get_summarized_history") as mock_get_history,
|
||||
patch("src.deriver.consumer.extract_facts_long_term") as mock_extract_facts,
|
||||
patch("src.deriver.consumer.crud.get_or_create_user_protected_collection") as mock_get_collection,
|
||||
patch("src.deriver.consumer.CollectionEmbeddingStore") as mock_embedding_store_class
|
||||
):
|
||||
# Mock history with previous context
|
||||
mock_get_history.return_value = (
|
||||
"AI: Hello! How can I help?\nhuman: I'm a machine learning engineer",
|
||||
[],
|
||||
None
|
||||
)
|
||||
mock_extract_facts.return_value = ["User prefers PyTorch"]
|
||||
mock_get_collection.return_value = collection
|
||||
|
||||
mock_embedding_store = AsyncMock()
|
||||
mock_embedding_store.remove_duplicates.return_value = ["User prefers PyTorch"]
|
||||
mock_embedding_store_class.return_value = mock_embedding_store
|
||||
|
||||
await consumer.process_user_message(
|
||||
message_content,
|
||||
test_app.public_id,
|
||||
test_user.public_id,
|
||||
session.public_id,
|
||||
message_id,
|
||||
db_session
|
||||
)
|
||||
|
||||
# Verify extract_facts was called with combined history + current message
|
||||
expected_history = "AI: Hello! How can I help?\nhuman: I'm a machine learning engineer\nhuman: I prefer PyTorch over TensorFlow"
|
||||
mock_extract_facts.assert_called_once_with(expected_history)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_user_message_handles_extraction_error(self, db_session, setup_user_message_test):
|
||||
"""Test that user message processing handles fact extraction errors gracefully."""
|
||||
test_app, test_user, session, collection = setup_user_message_test
|
||||
|
||||
message_content = "Test message"
|
||||
message_id = str(uuid4())
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.history.get_summarized_history") as mock_get_history,
|
||||
patch("src.deriver.consumer.extract_facts_long_term") as mock_extract_facts,
|
||||
patch("src.deriver.consumer.crud.get_or_create_user_protected_collection") as mock_get_collection
|
||||
):
|
||||
mock_get_history.return_value = ("", [], None)
|
||||
mock_extract_facts.side_effect = Exception("LLM API error")
|
||||
mock_get_collection.return_value = collection
|
||||
|
||||
# Should raise the exception (let caller handle it)
|
||||
with pytest.raises(Exception, match="LLM API error"):
|
||||
await consumer.process_user_message(
|
||||
message_content,
|
||||
test_app.public_id,
|
||||
test_user.public_id,
|
||||
session.public_id,
|
||||
message_id,
|
||||
db_session
|
||||
)
|
||||
|
||||
|
||||
class TestProcessAIMessage:
|
||||
"""Test AI message processing functionality."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_ai_message_basic_functionality(self, db_session):
|
||||
"""Test basic AI message processing (currently just console output)."""
|
||||
content = "I can help you with Python programming!"
|
||||
app_id = str(uuid4())
|
||||
user_id = str(uuid4())
|
||||
session_id = str(uuid4())
|
||||
message_id = str(uuid4())
|
||||
|
||||
# Mock console output
|
||||
with patch("src.deriver.consumer.console.print") as mock_print:
|
||||
await consumer.process_ai_message(
|
||||
content, app_id, user_id, session_id, message_id, db_session
|
||||
)
|
||||
|
||||
# Should print the AI message content
|
||||
mock_print.assert_called_once_with(
|
||||
f"Processing AI message: {content}",
|
||||
style="bright_magenta"
|
||||
)
|
||||
|
||||
|
||||
class TestSummarizeIfNeeded:
|
||||
"""Test summary generation functionality."""
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def setup_summary_test(self, db_session, sample_data):
|
||||
"""Setup test data for summary testing."""
|
||||
test_app, test_user = sample_data
|
||||
|
||||
session = models.Session(
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id,
|
||||
metadata={}
|
||||
)
|
||||
db_session.add(session)
|
||||
await db_session.flush()
|
||||
|
||||
return test_app, test_user, session
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_summarize_if_needed_no_summary_required(self, db_session, setup_summary_test):
|
||||
"""Test when no summary is needed."""
|
||||
test_app, test_user, session = setup_summary_test
|
||||
message_id = str(uuid4())
|
||||
|
||||
with patch("src.deriver.consumer.history.should_create_summary") as mock_should_create:
|
||||
# Mock that no summary is needed
|
||||
mock_should_create.return_value = (False, [], None)
|
||||
|
||||
await consumer.summarize_if_needed(
|
||||
db_session,
|
||||
test_app.public_id,
|
||||
session.public_id,
|
||||
test_user.public_id,
|
||||
message_id
|
||||
)
|
||||
|
||||
# Should only check for short summary
|
||||
mock_should_create.assert_called_once_with(
|
||||
db_session, session.public_id, summary_type=SummaryType.SHORT
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_summarize_if_needed_short_summary_only(self, db_session, setup_summary_test):
|
||||
"""Test creating only a short summary."""
|
||||
test_app, test_user, session = setup_summary_test
|
||||
message_id = str(uuid4())
|
||||
|
||||
# Mock messages for short summary
|
||||
mock_messages = [
|
||||
MagicMock(id=1, content="Message 1"),
|
||||
MagicMock(id=2, content="Message 2")
|
||||
]
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.history.should_create_summary") as mock_should_create,
|
||||
patch("src.deriver.consumer.history.create_summary") as mock_create_summary,
|
||||
patch("src.deriver.consumer.history.save_summary_metamessage") as mock_save_summary
|
||||
):
|
||||
# Mock summary check responses
|
||||
def mock_should_create_side_effect(db, session_id, summary_type):
|
||||
if summary_type == SummaryType.SHORT:
|
||||
return (True, mock_messages, None) # Need short summary
|
||||
else:
|
||||
return (False, [], None) # Don't need long summary
|
||||
|
||||
mock_should_create.side_effect = mock_should_create_side_effect
|
||||
mock_create_summary.return_value = "Short summary of recent messages"
|
||||
mock_save_summary.return_value = None
|
||||
|
||||
await consumer.summarize_if_needed(
|
||||
db_session,
|
||||
test_app.public_id,
|
||||
session.public_id,
|
||||
test_user.public_id,
|
||||
message_id
|
||||
)
|
||||
|
||||
# Should check for both short and long summaries
|
||||
assert mock_should_create.call_count == 2
|
||||
|
||||
# Should create one short summary
|
||||
mock_create_summary.assert_called_once_with(
|
||||
messages=mock_messages,
|
||||
previous_summary=None,
|
||||
summary_type=SummaryType.SHORT
|
||||
)
|
||||
|
||||
# Should save the short summary
|
||||
mock_save_summary.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_summarize_if_needed_both_summaries(self, db_session, setup_summary_test):
|
||||
"""Test creating both short and long summaries."""
|
||||
test_app, test_user, session = setup_summary_test
|
||||
message_id = str(uuid4())
|
||||
|
||||
# Mock messages for summaries
|
||||
mock_short_messages = [MagicMock(id=i, content=f"Message {i}") for i in range(1, 11)]
|
||||
mock_long_messages = [MagicMock(id=i, content=f"Message {i}") for i in range(1, 61)]
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.history.should_create_summary") as mock_should_create,
|
||||
patch("src.deriver.consumer.history.create_summary") as mock_create_summary,
|
||||
patch("src.deriver.consumer.history.save_summary_metamessage") as mock_save_summary
|
||||
):
|
||||
# Mock summary check responses
|
||||
def mock_should_create_side_effect(db, session_id, summary_type):
|
||||
if summary_type == SummaryType.SHORT:
|
||||
return (True, mock_short_messages, None)
|
||||
else:
|
||||
return (True, mock_long_messages, None)
|
||||
|
||||
mock_should_create.side_effect = mock_should_create_side_effect
|
||||
|
||||
# Mock summary creation
|
||||
def mock_create_summary_side_effect(messages, previous_summary, summary_type):
|
||||
if summary_type == SummaryType.LONG:
|
||||
return "Long summary of conversation"
|
||||
else:
|
||||
return "Short summary of recent messages"
|
||||
|
||||
mock_create_summary.side_effect = mock_create_summary_side_effect
|
||||
|
||||
# Mock save returning the long summary for short summary context
|
||||
mock_long_summary_obj = MagicMock()
|
||||
mock_long_summary_obj.content = "Long summary of conversation"
|
||||
|
||||
def mock_save_side_effect(db, app_id, user_id, session_id, message_id, summary_content, message_count, summary_type):
|
||||
if summary_type == SummaryType.LONG:
|
||||
return mock_long_summary_obj
|
||||
return None
|
||||
|
||||
mock_save_summary.side_effect = mock_save_side_effect
|
||||
|
||||
await consumer.summarize_if_needed(
|
||||
db_session,
|
||||
test_app.public_id,
|
||||
session.public_id,
|
||||
test_user.public_id,
|
||||
message_id
|
||||
)
|
||||
|
||||
# Should create both summaries
|
||||
assert mock_create_summary.call_count == 2
|
||||
|
||||
# Should save both summaries
|
||||
assert mock_save_summary.call_count == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_summarize_if_needed_handles_summary_creation_error(self, db_session, setup_summary_test):
|
||||
"""Test that summary creation errors are handled gracefully."""
|
||||
test_app, test_user, session = setup_summary_test
|
||||
message_id = str(uuid4())
|
||||
|
||||
mock_messages = [MagicMock(id=1, content="Message 1")]
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.history.should_create_summary") as mock_should_create,
|
||||
patch("src.deriver.consumer.history.create_summary") as mock_create_summary
|
||||
):
|
||||
mock_should_create.return_value = (True, mock_messages, None)
|
||||
mock_create_summary.side_effect = Exception("LLM API error")
|
||||
|
||||
# Should not raise exception (errors are logged)
|
||||
await consumer.summarize_if_needed(
|
||||
db_session,
|
||||
test_app.public_id,
|
||||
session.public_id,
|
||||
test_user.public_id,
|
||||
message_id
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_summarize_if_needed_with_existing_long_summary(self, db_session, setup_summary_test):
|
||||
"""Test short summary creation with existing long summary context."""
|
||||
test_app, test_user, session = setup_summary_test
|
||||
message_id = str(uuid4())
|
||||
|
||||
mock_messages = [MagicMock(id=1, content="Message 1")]
|
||||
mock_existing_long_summary = MagicMock()
|
||||
mock_existing_long_summary.content = "Existing long summary"
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.history.should_create_summary") as mock_should_create,
|
||||
patch("src.deriver.consumer.history.create_summary") as mock_create_summary,
|
||||
patch("src.deriver.consumer.history.save_summary_metamessage") as mock_save_summary
|
||||
):
|
||||
# Mock that we need short summary and have existing long summary
|
||||
def mock_should_create_side_effect(db, session_id, summary_type):
|
||||
if summary_type == SummaryType.SHORT:
|
||||
return (True, mock_messages, mock_existing_long_summary)
|
||||
else:
|
||||
return (False, [], mock_existing_long_summary)
|
||||
|
||||
mock_should_create.side_effect = mock_should_create_side_effect
|
||||
mock_create_summary.return_value = "Short summary with context"
|
||||
mock_save_summary.return_value = None
|
||||
|
||||
await consumer.summarize_if_needed(
|
||||
db_session,
|
||||
test_app.public_id,
|
||||
session.public_id,
|
||||
test_user.public_id,
|
||||
message_id
|
||||
)
|
||||
|
||||
# Should create short summary with existing long summary as context
|
||||
mock_create_summary.assert_called_once_with(
|
||||
messages=mock_messages,
|
||||
previous_summary="Existing long summary",
|
||||
summary_type=SummaryType.SHORT
|
||||
)
|
||||
|
||||
|
||||
class TestEnvironmentConfiguration:
|
||||
"""Test environment variable configuration."""
|
||||
|
||||
def test_tom_method_default(self):
|
||||
"""Test TOM_METHOD defaults to single_prompt."""
|
||||
# Test the getenv behavior that the module uses
|
||||
import os
|
||||
default_value = os.getenv("TOM_METHOD", "single_prompt")
|
||||
# If no environment variable is set, should use default
|
||||
if os.getenv("TOM_METHOD") is None:
|
||||
assert default_value == "single_prompt"
|
||||
else:
|
||||
# If environment variable is set, respect it
|
||||
assert default_value == os.getenv("TOM_METHOD")
|
||||
|
||||
def test_tom_method_custom(self):
|
||||
"""Test TOM_METHOD can be customized via environment."""
|
||||
# Test that the os.getenv pattern works correctly
|
||||
import os
|
||||
# Simulate the pattern used in the consumer module
|
||||
test_value = os.getenv("TOM_METHOD", "single_prompt")
|
||||
# The result should be either the env var or the default
|
||||
assert test_value in ["single_prompt", "conversational", "long_term"]
|
||||
|
||||
def test_user_representation_method_default(self):
|
||||
"""Test USER_REPRESENTATION_METHOD defaults to long_term."""
|
||||
# Test the getenv behavior that the module uses
|
||||
import os
|
||||
default_value = os.getenv("USER_REPRESENTATION_METHOD", "long_term")
|
||||
# If no environment variable is set, should use default
|
||||
if os.getenv("USER_REPRESENTATION_METHOD") is None:
|
||||
assert default_value == "long_term"
|
||||
else:
|
||||
# If environment variable is set, respect it
|
||||
assert default_value == os.getenv("USER_REPRESENTATION_METHOD")
|
||||
|
|
@ -0,0 +1,263 @@
|
|||
"""Performance tests for the deriver system."""
|
||||
|
||||
import asyncio
|
||||
import time
|
||||
from unittest.mock import AsyncMock, patch
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
|
||||
from src.deriver import consumer
|
||||
from src.deriver.tom.embeddings import CollectionEmbeddingStore
|
||||
|
||||
|
||||
class TestPerformanceValidation:
|
||||
"""Test performance characteristics of deriver components."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fact_extraction_performance(self, performance_config):
|
||||
"""Test that fact extraction completes within reasonable time limits."""
|
||||
start_time = time.time()
|
||||
|
||||
# Mock a reasonably complex chat history
|
||||
chat_history = """
|
||||
User: Hi, I'm Alex, a senior software engineer working at Google in the machine learning team.
|
||||
AI: Hello Alex! That sounds like an exciting role. What kind of ML projects are you working on?
|
||||
User: I'm primarily focused on building recommendation systems using TensorFlow and PyTorch.
|
||||
We handle millions of user interactions daily and need to provide real-time personalized recommendations.
|
||||
AI: That's impressive scale! How do you handle the computational challenges?
|
||||
User: We use a distributed architecture with Kubernetes, Redis for caching, and BigQuery for data processing.
|
||||
The team also experiments with newer frameworks like JAX for research prototypes.
|
||||
"""
|
||||
|
||||
with patch("src.deriver.tom.long_term.ModelClient") as mock_client:
|
||||
mock_instance = AsyncMock()
|
||||
mock_instance.generate.return_value = '<facts>{"facts": ["User is Alex", "User works at Google", "User uses TensorFlow"]}</facts>'
|
||||
mock_client.return_value = mock_instance
|
||||
|
||||
# Extract facts and measure time
|
||||
from src.deriver.tom.long_term import extract_facts_long_term
|
||||
facts = await extract_facts_long_term(chat_history)
|
||||
|
||||
extraction_time = time.time() - start_time
|
||||
|
||||
# Verify performance meets requirements
|
||||
assert extraction_time < performance_config["fact_extraction_time_limit"]
|
||||
assert len(facts) > 0
|
||||
print(f"✅ Fact extraction completed in {extraction_time:.3f}s")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_embedding_operations_performance(self, sample_data, performance_config):
|
||||
"""Test that embedding operations complete efficiently."""
|
||||
test_app, test_user = sample_data
|
||||
|
||||
# Create embedding store
|
||||
collection_id = "test_collection"
|
||||
store = CollectionEmbeddingStore(test_app.public_id, test_user.public_id, collection_id)
|
||||
|
||||
# Test data
|
||||
facts = [
|
||||
"User is a machine learning engineer",
|
||||
"User works with large-scale systems",
|
||||
"User has expertise in TensorFlow and PyTorch",
|
||||
"User handles millions of daily interactions",
|
||||
"User uses distributed computing"
|
||||
]
|
||||
|
||||
with (
|
||||
patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db,
|
||||
patch("src.deriver.tom.embeddings.crud.get_duplicate_documents") as mock_get_dupes,
|
||||
patch("src.deriver.tom.embeddings.crud.create_document") as mock_create_doc
|
||||
):
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
mock_get_dupes.return_value = [] # No duplicates
|
||||
mock_create_doc.return_value = None
|
||||
|
||||
# Test duplicate removal performance
|
||||
start_time = time.time()
|
||||
unique_facts = await store.remove_duplicates(facts)
|
||||
dedup_time = time.time() - start_time
|
||||
|
||||
# Test fact saving performance
|
||||
start_time = time.time()
|
||||
await store.save_facts(unique_facts)
|
||||
save_time = time.time() - start_time
|
||||
|
||||
# Verify performance
|
||||
total_time = dedup_time + save_time
|
||||
assert total_time < 2.0 # Should complete within 2 seconds
|
||||
assert len(unique_facts) == len(facts) # All facts should be unique
|
||||
|
||||
print(f"✅ Embedding operations completed in {total_time:.3f}s")
|
||||
print(f" - Deduplication: {dedup_time:.3f}s")
|
||||
print(f" - Fact saving: {save_time:.3f}s")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_concurrent_processing_performance(self, sample_data, performance_config):
|
||||
"""Test performance under concurrent load."""
|
||||
test_app, test_user = sample_data
|
||||
|
||||
# Create multiple simulated messages
|
||||
messages = [
|
||||
f"Message {i}: User sharing information about their work and interests"
|
||||
for i in range(performance_config["message_count"] // 10) # Smaller load for test
|
||||
]
|
||||
|
||||
# Mock all the dependencies for speed
|
||||
with (
|
||||
patch("src.deriver.consumer.extract_facts_long_term") as mock_extract,
|
||||
patch("src.deriver.consumer.history.get_summarized_history") as mock_history,
|
||||
patch("src.deriver.consumer.crud.get_or_create_user_protected_collection") as mock_get_collection,
|
||||
patch("src.deriver.consumer.CollectionEmbeddingStore") as mock_store_class,
|
||||
patch("src.deriver.consumer.summarize_if_needed") as mock_summarize
|
||||
):
|
||||
# Setup fast mocks
|
||||
mock_extract.return_value = ["User fact"]
|
||||
mock_history.return_value = ("", [], None)
|
||||
mock_get_collection.return_value = AsyncMock()
|
||||
|
||||
mock_store = AsyncMock()
|
||||
mock_store.remove_duplicates.return_value = ["User fact"]
|
||||
mock_store.save_facts.return_value = None
|
||||
mock_store_class.return_value = mock_store
|
||||
|
||||
mock_summarize.return_value = None
|
||||
|
||||
# Process messages concurrently
|
||||
start_time = time.time()
|
||||
|
||||
async def process_single_message(content):
|
||||
await consumer.process_user_message(
|
||||
content,
|
||||
test_app.public_id,
|
||||
test_user.public_id,
|
||||
"session_123",
|
||||
f"msg_{hash(content)}",
|
||||
AsyncMock() # Mock DB session
|
||||
)
|
||||
|
||||
# Run concurrent processing
|
||||
tasks = [process_single_message(msg) for msg in messages]
|
||||
await asyncio.gather(*tasks)
|
||||
|
||||
total_time = time.time() - start_time
|
||||
|
||||
# Verify performance
|
||||
messages_per_second = len(messages) / total_time
|
||||
assert messages_per_second > 5 # Should process at least 5 messages per second
|
||||
|
||||
print(f"✅ Processed {len(messages)} messages in {total_time:.3f}s")
|
||||
print(f" - Rate: {messages_per_second:.1f} messages/second")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_memory_usage_stability(self, sample_data):
|
||||
"""Test that memory usage remains stable during processing."""
|
||||
test_app, test_user = sample_data
|
||||
|
||||
# Simulate processing many messages to check for memory leaks
|
||||
message_count = 50
|
||||
|
||||
with (
|
||||
patch("src.deriver.consumer.extract_facts_long_term") as mock_extract,
|
||||
patch("src.deriver.consumer.history.get_summarized_history") as mock_history,
|
||||
patch("src.deriver.consumer.crud.get_or_create_user_protected_collection") as mock_get_collection,
|
||||
patch("src.deriver.consumer.CollectionEmbeddingStore") as mock_store_class,
|
||||
patch("src.deriver.consumer.summarize_if_needed") as mock_summarize
|
||||
):
|
||||
# Setup mocks
|
||||
mock_extract.return_value = ["Fact"]
|
||||
mock_history.return_value = ("", [], None)
|
||||
mock_get_collection.return_value = AsyncMock()
|
||||
|
||||
mock_store = AsyncMock()
|
||||
mock_store.remove_duplicates.return_value = ["Fact"]
|
||||
mock_store.save_facts.return_value = None
|
||||
mock_store_class.return_value = mock_store
|
||||
|
||||
mock_summarize.return_value = None
|
||||
|
||||
# Process messages in batches to simulate sustained load
|
||||
for batch in range(5): # 5 batches of 10 messages each
|
||||
batch_tasks = []
|
||||
for i in range(10):
|
||||
task = consumer.process_user_message(
|
||||
f"Batch {batch} Message {i}: User information",
|
||||
test_app.public_id,
|
||||
test_user.public_id,
|
||||
f"session_{batch}",
|
||||
f"msg_{batch}_{i}",
|
||||
AsyncMock()
|
||||
)
|
||||
batch_tasks.append(task)
|
||||
|
||||
# Process batch
|
||||
await asyncio.gather(*batch_tasks)
|
||||
|
||||
# Small delay between batches
|
||||
await asyncio.sleep(0.01)
|
||||
|
||||
# Verify all processing completed successfully
|
||||
assert mock_extract.call_count == message_count
|
||||
print(f"✅ Processed {message_count} messages in batches successfully")
|
||||
|
||||
def test_configuration_performance_settings(self, performance_config):
|
||||
"""Test that performance configuration is reasonable."""
|
||||
# Verify performance thresholds are achievable
|
||||
assert performance_config["fact_extraction_time_limit"] >= 1.0
|
||||
assert performance_config["tom_inference_time_limit"] >= 1.0
|
||||
assert performance_config["queue_processing_time_limit"] >= 0.1
|
||||
assert performance_config["max_workers"] >= 1
|
||||
assert performance_config["timeout_seconds"] >= 10
|
||||
|
||||
print("✅ Performance configuration validated")
|
||||
print(f" - Fact extraction limit: {performance_config['fact_extraction_time_limit']}s")
|
||||
print(f" - TOM inference limit: {performance_config['tom_inference_time_limit']}s")
|
||||
print(f" - Max workers: {performance_config['max_workers']}")
|
||||
|
||||
|
||||
class TestScalabilityValidation:
|
||||
"""Test scalability characteristics."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_fact_storage_scalability(self, sample_data):
|
||||
"""Test that fact storage can handle larger volumes."""
|
||||
test_app, test_user = sample_data
|
||||
|
||||
# Simulate storing many facts
|
||||
large_fact_list = [f"User fact number {i}" for i in range(100)]
|
||||
|
||||
collection_id = "test_scalability"
|
||||
store = CollectionEmbeddingStore(test_app.public_id, test_user.public_id, collection_id)
|
||||
|
||||
with (
|
||||
patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db,
|
||||
patch("src.deriver.tom.embeddings.crud.get_duplicate_documents") as mock_get_dupes,
|
||||
patch("src.deriver.tom.embeddings.crud.create_document") as mock_create_doc
|
||||
):
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
mock_get_dupes.return_value = []
|
||||
mock_create_doc.return_value = None
|
||||
|
||||
start_time = time.time()
|
||||
|
||||
# Test processing in chunks
|
||||
chunk_size = 20
|
||||
for i in range(0, len(large_fact_list), chunk_size):
|
||||
chunk = large_fact_list[i:i+chunk_size]
|
||||
unique_facts = await store.remove_duplicates(chunk)
|
||||
await store.save_facts(unique_facts)
|
||||
|
||||
total_time = time.time() - start_time
|
||||
|
||||
# Should handle 100 facts efficiently
|
||||
assert total_time < 5.0
|
||||
assert mock_create_doc.call_count == len(large_fact_list)
|
||||
|
||||
print(f"✅ Processed {len(large_fact_list)} facts in {total_time:.3f}s")
|
||||
print(f" - Rate: {len(large_fact_list)/total_time:.1f} facts/second")
|
||||
|
|
@ -0,0 +1,573 @@
|
|||
"""Tests for the QueueManager class and queue processing functionality."""
|
||||
|
||||
import asyncio
|
||||
import signal
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
from sqlalchemy import select
|
||||
|
||||
from src import models
|
||||
from src.deriver.queue import QueueManager
|
||||
|
||||
|
||||
class TestQueueManagerInitialization:
|
||||
"""Test QueueManager initialization and configuration."""
|
||||
|
||||
def test_queue_manager_default_initialization(self):
|
||||
"""Test QueueManager initializes with default values."""
|
||||
with patch("src.deriver.queue.os.getenv") as mock_getenv:
|
||||
mock_getenv.return_value = "1" # Default worker count
|
||||
|
||||
manager = QueueManager()
|
||||
|
||||
assert manager.workers == 1
|
||||
assert manager.semaphore._value == 1
|
||||
assert not manager.shutdown_event.is_set()
|
||||
assert len(manager.active_tasks) == 0
|
||||
assert len(manager.owned_sessions) == 0
|
||||
|
||||
def test_queue_manager_custom_workers(self):
|
||||
"""Test QueueManager respects DERIVER_WORKERS environment variable."""
|
||||
with patch("src.deriver.queue.os.getenv") as mock_getenv:
|
||||
mock_getenv.return_value = "4"
|
||||
|
||||
manager = QueueManager()
|
||||
|
||||
assert manager.workers == 4
|
||||
assert manager.semaphore._value == 4
|
||||
|
||||
@patch("src.deriver.queue.sentry_sdk")
|
||||
def test_sentry_initialization_enabled(self, mock_sentry):
|
||||
"""Test Sentry initialization when enabled."""
|
||||
with patch("src.deriver.queue.os.getenv") as mock_getenv:
|
||||
def getenv_side_effect(key, default=None):
|
||||
if key == "SENTRY_ENABLED":
|
||||
return "True"
|
||||
elif key == "SENTRY_DSN":
|
||||
return "https://test@sentry.io/123"
|
||||
elif key == "DERIVER_WORKERS":
|
||||
return "1"
|
||||
return default
|
||||
|
||||
mock_getenv.side_effect = getenv_side_effect
|
||||
|
||||
QueueManager()
|
||||
|
||||
mock_sentry.init.assert_called_once()
|
||||
|
||||
@patch("src.deriver.queue.sentry_sdk")
|
||||
def test_sentry_initialization_disabled(self, mock_sentry):
|
||||
"""Test Sentry is not initialized when disabled."""
|
||||
with patch("src.deriver.queue.os.getenv") as mock_getenv:
|
||||
def getenv_side_effect(key, default=None):
|
||||
if key == "SENTRY_ENABLED":
|
||||
return "False"
|
||||
elif key == "DERIVER_WORKERS":
|
||||
return "1"
|
||||
return default
|
||||
|
||||
mock_getenv.side_effect = getenv_side_effect
|
||||
|
||||
QueueManager()
|
||||
|
||||
mock_sentry.init.assert_not_called()
|
||||
|
||||
|
||||
class TestTaskAndSessionTracking:
|
||||
"""Test task and session tracking functionality."""
|
||||
|
||||
def test_add_task_tracking(self):
|
||||
"""Test adding tasks to tracking set."""
|
||||
with patch("src.deriver.queue.os.getenv", return_value="1"):
|
||||
manager = QueueManager()
|
||||
|
||||
# Create a mock task
|
||||
task = MagicMock()
|
||||
task.add_done_callback = MagicMock()
|
||||
|
||||
manager.add_task(task)
|
||||
|
||||
assert task in manager.active_tasks
|
||||
task.add_done_callback.assert_called_once()
|
||||
|
||||
def test_session_tracking(self):
|
||||
"""Test session tracking and untracking."""
|
||||
with patch("src.deriver.queue.os.getenv", return_value="1"):
|
||||
manager = QueueManager()
|
||||
session_id = 123
|
||||
|
||||
# Track session
|
||||
manager.track_session(session_id)
|
||||
assert session_id in manager.owned_sessions
|
||||
|
||||
# Untrack session
|
||||
manager.untrack_session(session_id)
|
||||
assert session_id not in manager.owned_sessions
|
||||
|
||||
def test_track_session_multiple(self):
|
||||
"""Test tracking multiple sessions."""
|
||||
with patch("src.deriver.queue.os.getenv", return_value="1"):
|
||||
manager = QueueManager()
|
||||
session_ids = [123, 456, 789]
|
||||
|
||||
for session_id in session_ids:
|
||||
manager.track_session(session_id)
|
||||
|
||||
assert all(sid in manager.owned_sessions for sid in session_ids)
|
||||
assert len(manager.owned_sessions) == 3
|
||||
|
||||
|
||||
class TestDatabaseOperations:
|
||||
"""Test database operations for queue management."""
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def setup_queue_data(self, db_session, sample_data):
|
||||
"""Setup test data for queue operations."""
|
||||
test_app, test_user = sample_data
|
||||
|
||||
# Create sessions
|
||||
session1 = models.Session(
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id,
|
||||
metadata={}
|
||||
)
|
||||
session2 = models.Session(
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id,
|
||||
metadata={}
|
||||
)
|
||||
db_session.add_all([session1, session2])
|
||||
await db_session.flush()
|
||||
|
||||
# Create queue items (use integer session.id, not public_id)
|
||||
queue_item1 = models.QueueItem(
|
||||
session_id=session1.id,
|
||||
payload={"message_id": str(uuid4())},
|
||||
processed=False
|
||||
)
|
||||
queue_item2 = models.QueueItem(
|
||||
session_id=session2.id,
|
||||
payload={"message_id": str(uuid4())},
|
||||
processed=False
|
||||
)
|
||||
queue_item3 = models.QueueItem(
|
||||
session_id=session1.id,
|
||||
payload={"message_id": str(uuid4())},
|
||||
processed=True # Already processed
|
||||
)
|
||||
|
||||
db_session.add_all([queue_item1, queue_item2, queue_item3])
|
||||
await db_session.flush()
|
||||
|
||||
return session1, session2, [queue_item1, queue_item2, queue_item3]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_available_sessions(self, db_session, setup_queue_data):
|
||||
"""Test getting available sessions for processing."""
|
||||
with patch("src.deriver.queue.os.getenv", return_value="1"):
|
||||
manager = QueueManager()
|
||||
session1, session2, queue_items = setup_queue_data
|
||||
|
||||
# Get available sessions
|
||||
available_sessions = await manager.get_available_sessions(db_session)
|
||||
|
||||
# Should return sessions with unprocessed items
|
||||
assert len(available_sessions) == 1 # Limited to 1 by the query
|
||||
assert available_sessions[0] in [session1.id, session2.id]
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_available_sessions_with_active_session(self, db_session, setup_queue_data):
|
||||
"""Test that active sessions are excluded from available sessions."""
|
||||
with patch("src.deriver.queue.os.getenv", return_value="1"):
|
||||
manager = QueueManager()
|
||||
session1, session2, queue_items = setup_queue_data
|
||||
|
||||
# Mark session1 as active
|
||||
active_session = models.ActiveQueueSession(session_id=session1.id)
|
||||
db_session.add(active_session)
|
||||
await db_session.flush()
|
||||
|
||||
# Get available sessions
|
||||
available_sessions = await manager.get_available_sessions(db_session)
|
||||
|
||||
# Should only return session2
|
||||
assert len(available_sessions) == 1
|
||||
assert available_sessions[0] == session2.id
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_stale_session_cleanup(self, db_session, setup_queue_data):
|
||||
"""Test cleanup of stale active sessions."""
|
||||
with patch("src.deriver.queue.os.getenv", return_value="1"):
|
||||
manager = QueueManager()
|
||||
session1, session2, queue_items = setup_queue_data
|
||||
|
||||
# Create a stale active session (older than 5 minutes)
|
||||
stale_time = datetime.now(timezone.utc) - timedelta(minutes=10)
|
||||
stale_session = models.ActiveQueueSession(
|
||||
session_id=session1.id,
|
||||
last_updated=stale_time
|
||||
)
|
||||
db_session.add(stale_session)
|
||||
await db_session.flush()
|
||||
|
||||
# Get available sessions (this should trigger cleanup)
|
||||
available_sessions = await manager.get_available_sessions(db_session)
|
||||
|
||||
# Stale session should be cleaned up, making session1 available
|
||||
result = await db_session.execute(
|
||||
select(models.ActiveQueueSession).where(
|
||||
models.ActiveQueueSession.session_id == session1.id
|
||||
)
|
||||
)
|
||||
assert result.scalar_one_or_none() is None
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_next_message(self, db_session, setup_queue_data):
|
||||
"""Test getting the next unprocessed message for a session."""
|
||||
with patch("src.deriver.queue.os.getenv", return_value="1"):
|
||||
manager = QueueManager()
|
||||
session1, session2, queue_items = setup_queue_data
|
||||
|
||||
# Get next message for session1
|
||||
next_message = await manager.get_next_message(db_session, session1.id)
|
||||
|
||||
# Should return the unprocessed message
|
||||
assert next_message is not None
|
||||
assert next_message.session_id == session1.id
|
||||
assert not next_message.processed
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_next_message_no_unprocessed(self, db_session, setup_queue_data):
|
||||
"""Test getting next message when all are processed."""
|
||||
with patch("src.deriver.queue.os.getenv", return_value="1"):
|
||||
manager = QueueManager()
|
||||
session1, session2, queue_items = setup_queue_data
|
||||
|
||||
# Mark all messages as processed
|
||||
for item in queue_items:
|
||||
item.processed = True
|
||||
await db_session.flush()
|
||||
|
||||
# Get next message
|
||||
next_message = await manager.get_next_message(db_session, session1.id)
|
||||
|
||||
# Should return None
|
||||
assert next_message is None
|
||||
|
||||
|
||||
class TestConcurrencyControl:
|
||||
"""Test concurrency control and semaphore behavior."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_semaphore_limits_concurrent_processing(self, mock_semaphore):
|
||||
"""Test that semaphore properly limits concurrent session processing."""
|
||||
with patch("src.deriver.queue.os.getenv", return_value="2"):
|
||||
with patch("asyncio.Semaphore") as mock_semaphore_class:
|
||||
mock_semaphore_class.return_value = mock_semaphore
|
||||
manager = QueueManager()
|
||||
|
||||
# Mock the process_session method to return actual async function
|
||||
async def mock_process_session(session_id):
|
||||
async with manager.semaphore:
|
||||
await asyncio.sleep(0.01) # Simulate work
|
||||
|
||||
with patch.object(manager, 'process_session', side_effect=mock_process_session):
|
||||
# Try to process multiple sessions
|
||||
tasks = []
|
||||
for i in range(5):
|
||||
task = asyncio.create_task(manager.process_session(i))
|
||||
tasks.append(task)
|
||||
manager.add_task(task)
|
||||
|
||||
# Wait for tasks to complete
|
||||
await asyncio.gather(*tasks, return_exceptions=True)
|
||||
|
||||
# Verify semaphore was used
|
||||
assert mock_semaphore.__aenter__.call_count == 5
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_polling_loop_respects_semaphore_capacity(self):
|
||||
"""Test that polling loop waits when all workers are busy."""
|
||||
with patch("src.deriver.queue.os.getenv", return_value="1"):
|
||||
manager = QueueManager()
|
||||
|
||||
# Mock semaphore as locked (no capacity)
|
||||
manager.semaphore.locked = MagicMock(return_value=True)
|
||||
|
||||
# Mock tracked_db to avoid database operations
|
||||
with patch("src.deriver.queue.tracked_db"):
|
||||
# Set shutdown event after a short delay to exit the loop
|
||||
async def set_shutdown():
|
||||
await asyncio.sleep(0.1)
|
||||
manager.shutdown_event.set()
|
||||
|
||||
asyncio.create_task(set_shutdown())
|
||||
|
||||
# Run polling loop
|
||||
await manager.polling_loop()
|
||||
|
||||
# Should have checked semaphore status
|
||||
manager.semaphore.locked.assert_called()
|
||||
|
||||
|
||||
class TestSignalHandling:
|
||||
"""Test signal handling and graceful shutdown."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_signal_handling(self, mock_signal_handling):
|
||||
"""Test that shutdown properly handles signals."""
|
||||
with patch("src.deriver.queue.os.getenv", return_value="1"):
|
||||
manager = QueueManager()
|
||||
|
||||
# Create actual async tasks instead of AsyncMock
|
||||
async def dummy_task():
|
||||
await asyncio.sleep(0.01)
|
||||
|
||||
task1 = asyncio.create_task(dummy_task())
|
||||
task2 = asyncio.create_task(dummy_task())
|
||||
manager.active_tasks = {task1, task2}
|
||||
|
||||
# Call shutdown
|
||||
await manager.shutdown(signal.SIGTERM)
|
||||
|
||||
# Shutdown event should be set
|
||||
assert manager.shutdown_event.is_set()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cleanup_owned_sessions(self, db_session):
|
||||
"""Test cleanup of owned sessions during shutdown."""
|
||||
with patch("src.deriver.queue.os.getenv", return_value="1"):
|
||||
manager = QueueManager()
|
||||
|
||||
# Add owned sessions
|
||||
session_ids = [123, 456, 789]
|
||||
for session_id in session_ids:
|
||||
manager.track_session(session_id)
|
||||
# Create corresponding active session records
|
||||
active_session = models.ActiveQueueSession(session_id=session_id)
|
||||
db_session.add(active_session)
|
||||
|
||||
await db_session.flush()
|
||||
|
||||
# Mock tracked_db to use our test session
|
||||
with patch("src.deriver.queue.tracked_db") as mock_tracked_db:
|
||||
mock_tracked_db.return_value.__aenter__.return_value = db_session
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
# Run cleanup
|
||||
await manager.cleanup()
|
||||
|
||||
# Verify sessions were removed from database
|
||||
result = await db_session.execute(
|
||||
select(models.ActiveQueueSession).where(
|
||||
models.ActiveQueueSession.session_id.in_(session_ids)
|
||||
)
|
||||
)
|
||||
remaining_sessions = result.scalars().all()
|
||||
assert len(remaining_sessions) == 0
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_cleanup_with_database_error(self, db_session):
|
||||
"""Test cleanup handles database errors gracefully."""
|
||||
with patch("src.deriver.queue.os.getenv", return_value="1"):
|
||||
manager = QueueManager()
|
||||
manager.track_session(123)
|
||||
|
||||
# Mock tracked_db to raise an exception
|
||||
with patch("src.deriver.queue.tracked_db") as mock_tracked_db:
|
||||
mock_tracked_db.side_effect = Exception("Database connection failed")
|
||||
|
||||
# Cleanup should not raise exception
|
||||
await manager.cleanup()
|
||||
|
||||
# Session should still be tracked (cleanup failed)
|
||||
assert 123 in manager.owned_sessions
|
||||
|
||||
|
||||
class TestErrorHandling:
|
||||
"""Test error handling in various scenarios."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_polling_loop_handles_database_errors(self):
|
||||
"""Test polling loop handles database errors gracefully."""
|
||||
with patch("src.deriver.queue.os.getenv", return_value="1"):
|
||||
manager = QueueManager()
|
||||
|
||||
# Mock tracked_db as proper async context manager that fails inside the context
|
||||
call_count = 0
|
||||
|
||||
class MockTrackedDBContext:
|
||||
def __init__(self, *args, **kwargs):
|
||||
nonlocal call_count
|
||||
call_count += 1
|
||||
|
||||
async def __aenter__(self):
|
||||
mock_db = MagicMock()
|
||||
# Make get_available_sessions fail on first call
|
||||
if call_count == 1:
|
||||
mock_db.execute.side_effect = Exception("Database connection failed")
|
||||
else:
|
||||
# Set shutdown on second call to exit loop
|
||||
manager.shutdown_event.set()
|
||||
mock_db.execute.return_value = MagicMock()
|
||||
return mock_db
|
||||
|
||||
async def __aexit__(self, exc_type, exc_val, exc_tb):
|
||||
return None
|
||||
|
||||
with patch("src.deriver.queue.tracked_db", MockTrackedDBContext):
|
||||
# Should not raise exception and should retry
|
||||
await manager.polling_loop()
|
||||
|
||||
# Should have attempted multiple calls
|
||||
assert call_count >= 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_process_session_marks_failed_messages_as_processed(self, db_session, sample_queue_items):
|
||||
"""Test that failed message processing still marks messages as processed."""
|
||||
with patch("src.deriver.queue.os.getenv", return_value="1"):
|
||||
manager = QueueManager()
|
||||
session, messages, queue_items = sample_queue_items
|
||||
|
||||
# Mock process_item to raise an exception
|
||||
with patch("src.deriver.queue.process_item", side_effect=Exception("Processing failed")):
|
||||
with patch("src.deriver.queue.tracked_db") as mock_tracked_db:
|
||||
mock_tracked_db.return_value.__aenter__.return_value = db_session
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
# Process the session
|
||||
await manager.process_session(session.id)
|
||||
|
||||
# All messages should be marked as processed despite the error
|
||||
result = await db_session.execute(
|
||||
select(models.QueueItem).where(
|
||||
models.QueueItem.session_id == session.id
|
||||
)
|
||||
)
|
||||
queue_items_after = result.scalars().all()
|
||||
assert all(item.processed for item in queue_items_after)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_session_claiming_handles_integrity_error(self, db_session, sample_data):
|
||||
"""Test that session claiming handles race conditions gracefully."""
|
||||
test_app, test_user = sample_data
|
||||
|
||||
# Create sessions
|
||||
session1 = models.Session(
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id,
|
||||
metadata={}
|
||||
)
|
||||
session2 = models.Session(
|
||||
user_id=test_user.public_id,
|
||||
app_id=test_app.public_id,
|
||||
metadata={}
|
||||
)
|
||||
db_session.add_all([session1, session2])
|
||||
await db_session.flush()
|
||||
|
||||
# Create queue items
|
||||
queue_item1 = models.QueueItem(
|
||||
session_id=session1.id,
|
||||
payload={"message_id": str(uuid4())},
|
||||
processed=False
|
||||
)
|
||||
queue_item2 = models.QueueItem(
|
||||
session_id=session2.id,
|
||||
payload={"message_id": str(uuid4())},
|
||||
processed=False
|
||||
)
|
||||
db_session.add_all([queue_item1, queue_item2])
|
||||
await db_session.flush()
|
||||
|
||||
with patch("src.deriver.queue.os.getenv", return_value="1"):
|
||||
manager = QueueManager()
|
||||
|
||||
# Create an active session to cause IntegrityError
|
||||
active_session = models.ActiveQueueSession(session_id=session1.id)
|
||||
db_session.add(active_session)
|
||||
await db_session.flush()
|
||||
|
||||
# Try to get available sessions and claim them
|
||||
available_sessions = await manager.get_available_sessions(db_session)
|
||||
|
||||
# Should get session2 (session1 is active)
|
||||
assert len(available_sessions) == 1
|
||||
assert available_sessions[0] == session2.id
|
||||
|
||||
|
||||
class TestIntegrationScenarios:
|
||||
"""Test integration scenarios and real-world usage patterns."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_full_session_processing_cycle(self, db_session, sample_queue_items):
|
||||
"""Test complete processing cycle for a session."""
|
||||
with patch("src.deriver.queue.os.getenv", return_value="1"):
|
||||
manager = QueueManager()
|
||||
session, messages, queue_items = sample_queue_items
|
||||
|
||||
# Mock process_item to simulate successful processing
|
||||
with patch("src.deriver.queue.process_item") as mock_process:
|
||||
mock_process.return_value = None
|
||||
|
||||
with patch("src.deriver.queue.tracked_db") as mock_tracked_db:
|
||||
mock_tracked_db.return_value.__aenter__.return_value = db_session
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
# Process the session
|
||||
await manager.process_session(session.id)
|
||||
|
||||
# Verify all user messages were processed
|
||||
user_message_count = len([item for item in queue_items])
|
||||
assert mock_process.call_count == user_message_count
|
||||
|
||||
# Verify session is not in active sessions
|
||||
result = await db_session.execute(
|
||||
select(models.ActiveQueueSession).where(
|
||||
models.ActiveQueueSession.session_id == session.id
|
||||
)
|
||||
)
|
||||
assert result.scalar_one_or_none() is None
|
||||
|
||||
# Verify session is untracked
|
||||
assert session.id not in manager.owned_sessions
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_shutdown_during_processing(self, db_session, sample_queue_items):
|
||||
"""Test graceful shutdown while processing messages."""
|
||||
with patch("src.deriver.queue.os.getenv", return_value="1"):
|
||||
manager = QueueManager()
|
||||
session, messages, queue_items = sample_queue_items
|
||||
|
||||
# Mock process_item to be slow and check shutdown event
|
||||
async def slow_process_item(db, payload):
|
||||
await asyncio.sleep(0.1)
|
||||
if manager.shutdown_event.is_set():
|
||||
return
|
||||
# Continue processing
|
||||
|
||||
with patch("src.deriver.queue.process_item", side_effect=slow_process_item):
|
||||
with patch("src.deriver.queue.tracked_db") as mock_tracked_db:
|
||||
mock_tracked_db.return_value.__aenter__.return_value = db_session
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
# Start processing
|
||||
process_task = asyncio.create_task(manager.process_session(session.id))
|
||||
|
||||
# Trigger shutdown after a short delay
|
||||
async def trigger_shutdown():
|
||||
await asyncio.sleep(0.05)
|
||||
manager.shutdown_event.set()
|
||||
|
||||
shutdown_task = asyncio.create_task(trigger_shutdown())
|
||||
|
||||
# Wait for both tasks
|
||||
await asyncio.gather(process_task, shutdown_task, return_exceptions=True)
|
||||
|
||||
# Session should be cleaned up even with shutdown
|
||||
assert session.id not in manager.owned_sessions
|
||||
|
||||
|
|
@ -0,0 +1,594 @@
|
|||
"""Tests for the TOM embeddings module and CollectionEmbeddingStore."""
|
||||
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
|
||||
from src import schemas
|
||||
from src.deriver.tom.embeddings import CollectionEmbeddingStore
|
||||
|
||||
|
||||
class TestCollectionEmbeddingStoreInitialization:
|
||||
"""Test CollectionEmbeddingStore initialization."""
|
||||
|
||||
def test_initialization(self):
|
||||
"""Test basic initialization of CollectionEmbeddingStore."""
|
||||
app_id = str(uuid4())
|
||||
user_id = str(uuid4())
|
||||
collection_id = str(uuid4())
|
||||
|
||||
store = CollectionEmbeddingStore(app_id, user_id, collection_id)
|
||||
|
||||
assert store.app_id == app_id
|
||||
assert store.user_id == user_id
|
||||
assert store.collection_id == collection_id
|
||||
|
||||
|
||||
class TestSaveFacts:
|
||||
"""Test fact saving functionality."""
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def embedding_store(self, sample_data):
|
||||
"""Create an embedding store for testing."""
|
||||
test_app, test_user = sample_data
|
||||
collection_id = str(uuid4())
|
||||
return CollectionEmbeddingStore(test_app.public_id, test_user.public_id, collection_id)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_save_facts_basic(self, embedding_store):
|
||||
"""Test basic fact saving functionality."""
|
||||
facts = [
|
||||
"User is a Python developer",
|
||||
"User works remotely",
|
||||
"User loves coffee"
|
||||
]
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.create_document") as mock_create_doc:
|
||||
mock_create_doc.return_value = None
|
||||
|
||||
await embedding_store.save_facts(facts)
|
||||
|
||||
# Should create a document for each fact
|
||||
assert mock_create_doc.call_count == 3
|
||||
|
||||
# Verify each call
|
||||
for i, fact in enumerate(facts):
|
||||
call_args = mock_create_doc.call_args_list[i]
|
||||
assert call_args[1]["app_id"] == embedding_store.app_id
|
||||
assert call_args[1]["user_id"] == embedding_store.user_id
|
||||
assert call_args[1]["collection_id"] == embedding_store.collection_id
|
||||
assert call_args[1]["document"].content == fact
|
||||
assert abs(call_args[1]["duplicate_threshold"] - 0.15) < 1e-10 # 1 - 0.85
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_save_facts_with_message_id(self, embedding_store):
|
||||
"""Test saving facts with message ID metadata."""
|
||||
facts = ["User prefers PyTorch over TensorFlow"]
|
||||
message_id = str(uuid4())
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.create_document") as mock_create_doc:
|
||||
mock_create_doc.return_value = None
|
||||
|
||||
await embedding_store.save_facts(facts, message_id=message_id)
|
||||
|
||||
# Verify message_id is included in metadata
|
||||
call_args = mock_create_doc.call_args_list[0]
|
||||
document = call_args[1]["document"]
|
||||
assert document.metadata == {"message_id": message_id}
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_save_facts_custom_similarity_threshold(self, embedding_store):
|
||||
"""Test saving facts with custom similarity threshold."""
|
||||
facts = ["User enjoys debugging"]
|
||||
similarity_threshold = 0.9
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.create_document") as mock_create_doc:
|
||||
mock_create_doc.return_value = None
|
||||
|
||||
await embedding_store.save_facts(facts, similarity_threshold=similarity_threshold)
|
||||
|
||||
# Verify duplicate threshold is calculated correctly
|
||||
call_args = mock_create_doc.call_args_list[0]
|
||||
assert abs(call_args[1]["duplicate_threshold"] - 0.1) < 1e-10 # 1 - 0.9
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_save_facts_handles_document_creation_error(self, embedding_store):
|
||||
"""Test that fact saving handles document creation errors gracefully."""
|
||||
facts = [
|
||||
"User is a Python developer",
|
||||
"This fact will fail",
|
||||
"User works remotely"
|
||||
]
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.create_document") as mock_create_doc:
|
||||
# Mock second call to raise an exception
|
||||
def mock_create_side_effect(*args, **kwargs):
|
||||
if "This fact will fail" in str(kwargs.get("document", "")):
|
||||
raise Exception("Database error")
|
||||
return None
|
||||
|
||||
mock_create_doc.side_effect = mock_create_side_effect
|
||||
|
||||
# Should not raise exception (errors are handled gracefully)
|
||||
await embedding_store.save_facts(facts)
|
||||
|
||||
# Should still attempt to create all documents
|
||||
assert mock_create_doc.call_count == 3
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_save_facts_empty_list(self, embedding_store):
|
||||
"""Test saving empty fact list."""
|
||||
facts = []
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.create_document") as mock_create_doc:
|
||||
await embedding_store.save_facts(facts)
|
||||
|
||||
# Should not call create_document
|
||||
mock_create_doc.assert_not_called()
|
||||
|
||||
|
||||
class TestGetRelevantFacts:
|
||||
"""Test fact retrieval functionality."""
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def embedding_store(self, sample_data):
|
||||
"""Create an embedding store for testing."""
|
||||
test_app, test_user = sample_data
|
||||
collection_id = str(uuid4())
|
||||
return CollectionEmbeddingStore(test_app.public_id, test_user.public_id, collection_id)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_relevant_facts_basic(self, embedding_store):
|
||||
"""Test basic fact retrieval functionality."""
|
||||
query = "What programming languages does the user know?"
|
||||
|
||||
# Mock documents returned from query
|
||||
mock_documents = [
|
||||
MagicMock(content="User is proficient in Python"),
|
||||
MagicMock(content="User has experience with JavaScript"),
|
||||
MagicMock(content="User knows SQL")
|
||||
]
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.query_documents") as mock_query_docs:
|
||||
mock_query_docs.return_value = mock_documents
|
||||
|
||||
facts = await embedding_store.get_relevant_facts(query)
|
||||
|
||||
# Verify query parameters
|
||||
mock_query_docs.assert_called_once_with(
|
||||
mock_db,
|
||||
app_id=embedding_store.app_id,
|
||||
user_id=embedding_store.user_id,
|
||||
collection_id=embedding_store.collection_id,
|
||||
query=query,
|
||||
max_distance=0.3,
|
||||
top_k=5
|
||||
)
|
||||
|
||||
# Verify returned facts
|
||||
expected_facts = [
|
||||
"User is proficient in Python",
|
||||
"User has experience with JavaScript",
|
||||
"User knows SQL"
|
||||
]
|
||||
assert facts == expected_facts
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_relevant_facts_custom_parameters(self, embedding_store):
|
||||
"""Test fact retrieval with custom parameters."""
|
||||
query = "What does the user do for work?"
|
||||
top_k = 10
|
||||
max_distance = 0.2
|
||||
|
||||
mock_documents = [MagicMock(content="User is a software engineer")]
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.query_documents") as mock_query_docs:
|
||||
mock_query_docs.return_value = mock_documents
|
||||
|
||||
facts = await embedding_store.get_relevant_facts(
|
||||
query, top_k=top_k, max_distance=max_distance
|
||||
)
|
||||
|
||||
# Verify custom parameters were used
|
||||
mock_query_docs.assert_called_once_with(
|
||||
mock_db,
|
||||
app_id=embedding_store.app_id,
|
||||
user_id=embedding_store.user_id,
|
||||
collection_id=embedding_store.collection_id,
|
||||
query=query,
|
||||
max_distance=max_distance,
|
||||
top_k=top_k
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_relevant_facts_no_results(self, embedding_store):
|
||||
"""Test fact retrieval when no relevant facts are found."""
|
||||
query = "What is the user's favorite food?"
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.query_documents") as mock_query_docs:
|
||||
mock_query_docs.return_value = [] # No results
|
||||
|
||||
facts = await embedding_store.get_relevant_facts(query)
|
||||
|
||||
assert facts == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_relevant_facts_empty_query(self, embedding_store):
|
||||
"""Test fact retrieval with empty query."""
|
||||
query = ""
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.query_documents") as mock_query_docs:
|
||||
mock_query_docs.return_value = []
|
||||
|
||||
facts = await embedding_store.get_relevant_facts(query)
|
||||
|
||||
# Should still call query_documents with empty query
|
||||
mock_query_docs.assert_called_once()
|
||||
assert facts == []
|
||||
|
||||
|
||||
class TestRemoveDuplicates:
|
||||
"""Test duplicate detection and removal functionality."""
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def embedding_store(self, sample_data):
|
||||
"""Create an embedding store for testing."""
|
||||
test_app, test_user = sample_data
|
||||
collection_id = str(uuid4())
|
||||
return CollectionEmbeddingStore(test_app.public_id, test_user.public_id, collection_id)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_remove_duplicates_no_duplicates(self, embedding_store):
|
||||
"""Test duplicate removal when no duplicates exist."""
|
||||
facts = [
|
||||
"User is a Python developer",
|
||||
"User works remotely",
|
||||
"User loves coffee"
|
||||
]
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.get_duplicate_documents") as mock_get_dupes:
|
||||
mock_get_dupes.return_value = [] # No duplicates
|
||||
|
||||
unique_facts = await embedding_store.remove_duplicates(facts)
|
||||
|
||||
# All facts should be considered unique
|
||||
assert unique_facts == facts
|
||||
|
||||
# Should check each fact for duplicates
|
||||
assert mock_get_dupes.call_count == 3
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_remove_duplicates_with_duplicates(self, embedding_store):
|
||||
"""Test duplicate removal when duplicates exist."""
|
||||
facts = [
|
||||
"User is a Python developer",
|
||||
"User codes in Python", # Similar to first fact
|
||||
"User works remotely"
|
||||
]
|
||||
|
||||
# Mock duplicate document for second fact
|
||||
mock_duplicate_doc = MagicMock()
|
||||
mock_duplicate_doc.content = "User is a Python programmer"
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.get_duplicate_documents") as mock_get_dupes:
|
||||
def mock_get_dupes_side_effect(db, app_id, user_id, collection_id, content, similarity_threshold):
|
||||
if "codes in Python" in content:
|
||||
return [mock_duplicate_doc] # Duplicate found
|
||||
return [] # No duplicates
|
||||
|
||||
mock_get_dupes.side_effect = mock_get_dupes_side_effect
|
||||
|
||||
unique_facts = await embedding_store.remove_duplicates(facts)
|
||||
|
||||
# Should remove the duplicate fact
|
||||
expected_unique = [
|
||||
"User is a Python developer",
|
||||
"User works remotely"
|
||||
]
|
||||
assert unique_facts == expected_unique
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_remove_duplicates_custom_threshold(self, embedding_store):
|
||||
"""Test duplicate removal with custom similarity threshold."""
|
||||
facts = ["User enjoys programming"]
|
||||
similarity_threshold = 0.9
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.get_duplicate_documents") as mock_get_dupes:
|
||||
mock_get_dupes.return_value = []
|
||||
|
||||
await embedding_store.remove_duplicates(facts, similarity_threshold=similarity_threshold)
|
||||
|
||||
# Verify custom threshold was passed
|
||||
mock_get_dupes.assert_called_once_with(
|
||||
mock_db,
|
||||
app_id=embedding_store.app_id,
|
||||
user_id=embedding_store.user_id,
|
||||
collection_id=embedding_store.collection_id,
|
||||
content=facts[0],
|
||||
similarity_threshold=similarity_threshold
|
||||
)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_remove_duplicates_handles_errors(self, embedding_store):
|
||||
"""Test that duplicate checking handles errors gracefully."""
|
||||
facts = [
|
||||
"User is a Python developer",
|
||||
"This fact will cause an error",
|
||||
"User works remotely"
|
||||
]
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.get_duplicate_documents") as mock_get_dupes:
|
||||
def mock_get_dupes_side_effect(*args, **kwargs):
|
||||
if "cause an error" in kwargs.get("content", ""):
|
||||
raise Exception("Database connection error")
|
||||
return []
|
||||
|
||||
mock_get_dupes.side_effect = mock_get_dupes_side_effect
|
||||
|
||||
unique_facts = await embedding_store.remove_duplicates(facts)
|
||||
|
||||
# Should include all facts (error results in keeping the fact)
|
||||
assert unique_facts == facts
|
||||
assert mock_get_dupes.call_count == 3
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_remove_duplicates_empty_list(self, embedding_store):
|
||||
"""Test duplicate removal with empty fact list."""
|
||||
facts = []
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.get_duplicate_documents") as mock_get_dupes:
|
||||
unique_facts = await embedding_store.remove_duplicates(facts)
|
||||
|
||||
assert unique_facts == []
|
||||
mock_get_dupes.assert_not_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_remove_duplicates_logs_duplicate_found(self, embedding_store):
|
||||
"""Test that duplicate detection logs when duplicates are found."""
|
||||
facts = ["User loves Python programming"]
|
||||
|
||||
mock_duplicate_doc = MagicMock()
|
||||
mock_duplicate_doc.content = "User enjoys Python development"
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.get_duplicate_documents") as mock_get_dupes:
|
||||
mock_get_dupes.return_value = [mock_duplicate_doc]
|
||||
|
||||
with patch("src.deriver.tom.embeddings.logger.debug") as mock_log:
|
||||
unique_facts = await embedding_store.remove_duplicates(facts)
|
||||
|
||||
# Should log the duplicate detection
|
||||
mock_log.assert_called_once()
|
||||
log_message = mock_log.call_args[0][0]
|
||||
assert "Duplicate found" in log_message
|
||||
assert mock_duplicate_doc.content in log_message
|
||||
assert facts[0] in log_message
|
||||
|
||||
# Should not include the duplicate fact
|
||||
assert unique_facts == []
|
||||
|
||||
|
||||
class TestIntegrationScenarios:
|
||||
"""Test integration scenarios combining multiple operations."""
|
||||
|
||||
@pytest_asyncio.fixture
|
||||
async def embedding_store(self, sample_data):
|
||||
"""Create an embedding store for testing."""
|
||||
test_app, test_user = sample_data
|
||||
collection_id = str(uuid4())
|
||||
return CollectionEmbeddingStore(test_app.public_id, test_user.public_id, collection_id)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_full_workflow_save_and_retrieve(self, embedding_store):
|
||||
"""Test complete workflow of saving facts and retrieving them."""
|
||||
# First save some facts
|
||||
facts_to_save = [
|
||||
"User is a senior Python developer",
|
||||
"User has 5 years of experience with FastAPI",
|
||||
"User prefers async programming"
|
||||
]
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.create_document") as mock_create_doc:
|
||||
mock_create_doc.return_value = None
|
||||
|
||||
await embedding_store.save_facts(facts_to_save)
|
||||
|
||||
# Verify all facts were saved
|
||||
assert mock_create_doc.call_count == 3
|
||||
|
||||
# Then retrieve relevant facts
|
||||
query = "What is the user's programming experience?"
|
||||
mock_documents = [
|
||||
MagicMock(content="User is a senior Python developer"),
|
||||
MagicMock(content="User has 5 years of experience with FastAPI")
|
||||
]
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.query_documents") as mock_query_docs:
|
||||
mock_query_docs.return_value = mock_documents
|
||||
|
||||
retrieved_facts = await embedding_store.get_relevant_facts(query)
|
||||
|
||||
expected_facts = [
|
||||
"User is a senior Python developer",
|
||||
"User has 5 years of experience with FastAPI"
|
||||
]
|
||||
assert retrieved_facts == expected_facts
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_duplicate_removal_before_saving(self, embedding_store):
|
||||
"""Test the typical workflow of removing duplicates before saving."""
|
||||
facts_to_check = [
|
||||
"User is a Python developer",
|
||||
"User writes code in Python", # Potential duplicate
|
||||
"User works from home"
|
||||
]
|
||||
|
||||
# Mock existing duplicate
|
||||
mock_duplicate_doc = MagicMock()
|
||||
mock_duplicate_doc.content = "User is proficient in Python"
|
||||
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.get_duplicate_documents") as mock_get_dupes:
|
||||
def mock_get_dupes_side_effect(*args, **kwargs):
|
||||
if "writes code in Python" in kwargs.get("content", ""):
|
||||
return [mock_duplicate_doc]
|
||||
return []
|
||||
|
||||
mock_get_dupes.side_effect = mock_get_dupes_side_effect
|
||||
|
||||
# Remove duplicates
|
||||
unique_facts = await embedding_store.remove_duplicates(facts_to_check)
|
||||
|
||||
expected_unique = [
|
||||
"User is a Python developer",
|
||||
"User works from home"
|
||||
]
|
||||
assert unique_facts == expected_unique
|
||||
|
||||
# Now save the unique facts
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.create_document") as mock_create_doc:
|
||||
mock_create_doc.return_value = None
|
||||
|
||||
await embedding_store.save_facts(unique_facts)
|
||||
|
||||
# Should only save the unique facts
|
||||
assert mock_create_doc.call_count == 2
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_error_recovery_in_workflow(self, embedding_store):
|
||||
"""Test error recovery across multiple operations."""
|
||||
facts = [
|
||||
"User is experienced with machine learning",
|
||||
"User uses scikit-learn and pandas"
|
||||
]
|
||||
|
||||
# Test save_facts with partial failure
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.create_document") as mock_create_doc:
|
||||
def mock_create_side_effect(*args, **kwargs):
|
||||
if "scikit-learn" in str(kwargs.get("document", "")):
|
||||
raise Exception("Database error")
|
||||
return None
|
||||
|
||||
mock_create_doc.side_effect = mock_create_side_effect
|
||||
|
||||
# Should handle the error gracefully
|
||||
await embedding_store.save_facts(facts)
|
||||
|
||||
# Should attempt to save both facts
|
||||
assert mock_create_doc.call_count == 2
|
||||
|
||||
# Test get_relevant_facts after partial save
|
||||
with patch("src.deriver.tom.embeddings.tracked_db") as mock_tracked_db:
|
||||
mock_db = AsyncMock()
|
||||
mock_tracked_db.return_value.__aenter__.return_value = mock_db
|
||||
mock_tracked_db.return_value.__aexit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.embeddings.crud.query_documents") as mock_query_docs:
|
||||
# Only return the successfully saved fact
|
||||
mock_query_docs.return_value = [
|
||||
MagicMock(content="User is experienced with machine learning")
|
||||
]
|
||||
|
||||
retrieved_facts = await embedding_store.get_relevant_facts("machine learning")
|
||||
|
||||
assert retrieved_facts == ["User is experienced with machine learning"]
|
||||
|
|
@ -0,0 +1,701 @@
|
|||
"""Tests for TOM (Theory of Mind) inference modules."""
|
||||
|
||||
import json
|
||||
from unittest.mock import AsyncMock, MagicMock, patch
|
||||
from uuid import uuid4
|
||||
|
||||
import pytest
|
||||
import pytest_asyncio
|
||||
|
||||
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_model_clients):
|
||||
"""Test basic single prompt TOM inference."""
|
||||
chat_history = "User: I'm feeling stressed about work\nAI: What's causing the stress?"
|
||||
session_id = str(uuid4())
|
||||
|
||||
with patch("src.deriver.tom.single_prompt.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_single_prompt(chat_history, session_id)
|
||||
|
||||
# Verify model client was called with correct parameters
|
||||
mock_client = mock_model_clients["single_prompt"]
|
||||
mock_client.generate.assert_called_once()
|
||||
call_kwargs = mock_client.generate.call_args[1]
|
||||
|
||||
assert call_kwargs["max_tokens"] == 1000
|
||||
assert call_kwargs["temperature"] == 0
|
||||
assert call_kwargs["use_caching"] is True
|
||||
# The system prompt should contain key TOM instruction phrases
|
||||
assert "system" in call_kwargs
|
||||
system_prompt = call_kwargs["system"]
|
||||
assert "theory of mind" in system_prompt.lower() or "prediction" in system_prompt.lower()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_tom_inference_single_prompt_with_user_representation(self, mock_model_clients):
|
||||
"""Test single prompt TOM inference with existing user representation."""
|
||||
chat_history = "User: I changed my mind about the project"
|
||||
session_id = str(uuid4())
|
||||
user_representation = "User is decisive and goal-oriented"
|
||||
|
||||
with patch("src.deriver.tom.single_prompt.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_single_prompt(
|
||||
chat_history, session_id, user_representation
|
||||
)
|
||||
|
||||
# Verify user representation was included in messages
|
||||
mock_client = mock_model_clients["single_prompt"]
|
||||
call_args = mock_client.generate.call_args[1]
|
||||
messages = call_args["messages"]
|
||||
|
||||
# Should have two messages: main analysis + user representation context
|
||||
assert len(messages) == 2
|
||||
assert user_representation in str(messages)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_tom_inference_single_prompt_handles_error(self, mock_model_clients):
|
||||
"""Test that single prompt TOM inference handles LLM errors."""
|
||||
chat_history = "User: Test message"
|
||||
session_id = str(uuid4())
|
||||
|
||||
# Mock the model client to raise an exception
|
||||
mock_model_clients["single_prompt"].generate.side_effect = Exception("LLM API Error")
|
||||
|
||||
with patch("src.deriver.tom.single_prompt.sentry_sdk.start_transaction") as mock_transaction:
|
||||
mock_transaction.return_value.__enter__.return_value = MagicMock()
|
||||
mock_transaction.return_value.__exit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.single_prompt.sentry_sdk.capture_exception") as mock_capture:
|
||||
with pytest.raises(Exception, match="LLM API Error"):
|
||||
await get_tom_inference_single_prompt(chat_history, session_id)
|
||||
|
||||
# Verify error was captured
|
||||
mock_capture.assert_called_once()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_user_representation_single_prompt_basic(self, mock_model_clients):
|
||||
"""Test basic single prompt user representation."""
|
||||
chat_history = "User: I'm a data scientist\nAI: What tools do you use?"
|
||||
session_id = str(uuid4())
|
||||
tom_inference = "User is passionate about data science"
|
||||
|
||||
with patch("src.deriver.tom.single_prompt.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_single_prompt(
|
||||
chat_history, session_id, tom_inference=tom_inference
|
||||
)
|
||||
|
||||
# Verify correct system prompt was used
|
||||
mock_client = mock_model_clients["single_prompt"]
|
||||
call_kwargs = mock_client.generate.call_args[1]
|
||||
# The system prompt should contain user representation instructions
|
||||
assert "system" in call_kwargs
|
||||
system_prompt = call_kwargs["system"]
|
||||
assert "user representation" in system_prompt.lower() or "factual" in system_prompt.lower()
|
||||
|
||||
# Verify TOM inference was included in context
|
||||
messages = call_kwargs["messages"]
|
||||
assert tom_inference in str(messages)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_user_representation_single_prompt_all_inputs(self, mock_model_clients):
|
||||
"""Test single prompt user representation with all optional inputs."""
|
||||
chat_history = "User: I've been learning React lately"
|
||||
session_id = str(uuid4())
|
||||
user_representation = "User is a full-stack developer"
|
||||
tom_inference = "User is eager to learn new technologies"
|
||||
|
||||
with patch("src.deriver.tom.single_prompt.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_single_prompt(
|
||||
chat_history, session_id, user_representation, tom_inference
|
||||
)
|
||||
|
||||
# Verify all inputs were included in the context
|
||||
mock_client = mock_model_clients["single_prompt"]
|
||||
call_kwargs = mock_client.generate.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 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_model_clients, 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"
|
||||
|
||||
# Mock the response with proper XML format
|
||||
mock_response = f'<facts>{mock_llm_responses["fact_extraction"]}</facts>'
|
||||
mock_model_clients["long_term"].generate.return_value = mock_response
|
||||
|
||||
with patch("src.deriver.tom.long_term.parse_xml_content") as mock_parse_xml:
|
||||
mock_parse_xml.return_value = mock_llm_responses["fact_extraction"]
|
||||
|
||||
facts = await extract_facts_long_term(chat_history)
|
||||
|
||||
# Verify model client was called
|
||||
mock_client = mock_model_clients["long_term"]
|
||||
mock_client.generate.assert_called_once()
|
||||
|
||||
call_kwargs = mock_client.generate.call_args[1]
|
||||
assert call_kwargs["temperature"] == 0.0
|
||||
assert call_kwargs["use_caching"] is True
|
||||
|
||||
# Verify the system prompt includes the chat history
|
||||
messages = call_kwargs["messages"]
|
||||
message_content = str(messages)
|
||||
# Check for key parts of the chat history
|
||||
assert "software engineer" in message_content.lower()
|
||||
assert "google" in message_content.lower() or "hiking" in message_content.lower()
|
||||
|
||||
# Verify facts were extracted correctly
|
||||
expected_facts = [
|
||||
"User is a software developer",
|
||||
"User works remotely",
|
||||
"User prefers coffee over tea",
|
||||
"User uses Python and JavaScript"
|
||||
]
|
||||
assert facts == expected_facts
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_extract_facts_long_term_handles_json_error(self, mock_model_clients):
|
||||
"""Test that fact extraction handles JSON parsing errors gracefully."""
|
||||
chat_history = "User: I like programming"
|
||||
|
||||
# Mock malformed response
|
||||
mock_model_clients["long_term"].generate.return_value = "Invalid JSON response"
|
||||
|
||||
with patch("src.deriver.tom.long_term.parse_xml_content") as mock_parse_xml:
|
||||
mock_parse_xml.return_value = "Not valid JSON"
|
||||
|
||||
facts = await extract_facts_long_term(chat_history)
|
||||
|
||||
# Should return empty list on error
|
||||
assert facts == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_extract_facts_long_term_handles_missing_facts_key(self, mock_model_clients):
|
||||
"""Test that fact extraction handles missing 'facts' key in response."""
|
||||
chat_history = "User: Test message"
|
||||
|
||||
# Mock response with missing facts key
|
||||
invalid_response = json.dumps({"other_key": "some_value"})
|
||||
mock_model_clients["long_term"].generate.return_value = f"<facts>{invalid_response}</facts>"
|
||||
|
||||
with patch("src.deriver.tom.long_term.parse_xml_content") as mock_parse_xml:
|
||||
mock_parse_xml.return_value = invalid_response
|
||||
|
||||
facts = await extract_facts_long_term(chat_history)
|
||||
|
||||
# Should return empty list on KeyError
|
||||
assert facts == []
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_user_representation_long_term_basic(self, mock_model_clients):
|
||||
"""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"]
|
||||
|
||||
mock_response = "CURRENT STATE:\n- Active Context: Starting new job\n<KNOWN_FACTS>\nTENTATIVE PATTERNS:\n- High confidence: Career-focused"
|
||||
mock_model_clients["long_term"].generate.return_value = mock_response
|
||||
|
||||
result = await get_user_representation_long_term(
|
||||
chat_history, session_id, facts=facts
|
||||
)
|
||||
|
||||
# Verify model client was called
|
||||
mock_client = mock_model_clients["long_term"]
|
||||
mock_client.generate.assert_called_once()
|
||||
|
||||
call_kwargs = mock_client.generate.call_args[1]
|
||||
assert call_kwargs["temperature"] == 0
|
||||
assert call_kwargs["use_caching"] is True
|
||||
|
||||
# Verify chat history was included
|
||||
messages = call_kwargs["messages"]
|
||||
assert chat_history in str(messages)
|
||||
|
||||
# Verify facts were injected into the response
|
||||
assert "User is a software engineer" in result
|
||||
assert "User is changing jobs" in result
|
||||
assert "<KNOWN_FACTS>" not in result # Should be replaced
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_user_representation_long_term_with_all_inputs(self, mock_model_clients):
|
||||
"""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"]
|
||||
|
||||
mock_response = "Updated representation with <KNOWN_FACTS> placeholder"
|
||||
mock_model_clients["long_term"].generate.return_value = mock_response
|
||||
|
||||
result = await get_user_representation_long_term(
|
||||
chat_history, session_id, user_representation, tom_inference, facts
|
||||
)
|
||||
|
||||
# Verify all inputs were included in the context
|
||||
mock_client = mock_model_clients["long_term"]
|
||||
call_kwargs = mock_client.generate.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
|
||||
|
||||
# Verify facts injection worked
|
||||
assert "User works in tech" in result
|
||||
assert "User enjoys new challenges" in result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_user_representation_long_term_empty_facts(self, mock_model_clients):
|
||||
"""Test long term user representation with empty facts list."""
|
||||
chat_history = "User: Hello there"
|
||||
session_id = str(uuid4())
|
||||
|
||||
mock_response = "Basic representation with <KNOWN_FACTS> placeholder"
|
||||
mock_model_clients["long_term"].generate.return_value = mock_response
|
||||
|
||||
result = await get_user_representation_long_term(chat_history, session_id)
|
||||
|
||||
# Verify empty facts are handled gracefully
|
||||
assert "PERSISTENT INFORMATION:\n" in result
|
||||
assert "<KNOWN_FACTS>" not in result
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_user_representation_long_term_none_inputs(self, mock_model_clients):
|
||||
"""Test long term user representation with None inputs."""
|
||||
chat_history = "User: Test message"
|
||||
session_id = str(uuid4())
|
||||
|
||||
mock_response = "Representation with <KNOWN_FACTS>"
|
||||
mock_model_clients["long_term"].generate.return_value = mock_response
|
||||
|
||||
result = await get_user_representation_long_term(
|
||||
chat_history, session_id, user_representation="None", tom_inference="None"
|
||||
)
|
||||
|
||||
# Verify None inputs are handled (not included in context)
|
||||
mock_client = mock_model_clients["long_term"]
|
||||
call_kwargs = mock_client.generate.call_args[1]
|
||||
messages = call_kwargs["messages"]
|
||||
message_content = str(messages)
|
||||
|
||||
# "None" values should not be included in context strings
|
||||
assert "EXISTING USER REPRESENTATION - INCOMPLETE" not in message_content
|
||||
assert "PREDICTION OF USER MENTAL STATE" not in message_content
|
||||
|
||||
|
||||
class TestTOMIntegration:
|
||||
"""Test integration scenarios across TOM methods."""
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_method_configuration_via_environment(self, mock_model_clients):
|
||||
"""Test that TOM methods can be configured via environment variables."""
|
||||
chat_history = "User: I'm learning data science"
|
||||
session_id = str(uuid4())
|
||||
|
||||
# Test single_prompt method
|
||||
with patch("src.deriver.tom.single_prompt.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(chat_history, session_id, method="single_prompt")
|
||||
|
||||
# Verify single prompt was called
|
||||
mock_model_clients["single_prompt"].generate.assert_called()
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_error_handling_across_methods(self, mock_model_clients):
|
||||
"""Test error handling consistency across different TOM methods."""
|
||||
chat_history = "User: Error test"
|
||||
session_id = str(uuid4())
|
||||
|
||||
# Test single_prompt error handling
|
||||
mock_model_clients["single_prompt"].generate.side_effect = Exception("API Error")
|
||||
|
||||
with patch("src.deriver.tom.single_prompt.sentry_sdk.start_transaction") as mock_transaction:
|
||||
mock_transaction.return_value.__enter__.return_value = MagicMock()
|
||||
mock_transaction.return_value.__exit__.return_value = None
|
||||
|
||||
with patch("src.deriver.tom.single_prompt.sentry_sdk.capture_exception"):
|
||||
with pytest.raises(Exception):
|
||||
await get_tom_inference(chat_history, session_id, method="single_prompt")
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_response_format_consistency(self, mock_model_clients, 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
|
||||
with patch("src.deriver.tom.single_prompt.sentry_sdk.start_transaction") as mock_transaction:
|
||||
mock_transaction.return_value.__enter__.return_value = MagicMock()
|
||||
mock_transaction.return_value.__exit__.return_value = None
|
||||
|
||||
single_prompt_result = await get_tom_inference(
|
||||
chat_history, session_id, method="single_prompt"
|
||||
)
|
||||
|
||||
# Verify response is returned correctly
|
||||
assert single_prompt_result == mock_llm_responses["tom_single_prompt"]
|
||||
|
||||
# Test long term fact extraction format
|
||||
with patch("src.deriver.tom.long_term.parse_xml_content") as mock_parse_xml:
|
||||
mock_parse_xml.return_value = mock_llm_responses["fact_extraction"]
|
||||
|
||||
facts = await extract_facts_long_term(chat_history)
|
||||
|
||||
# Verify facts are returned as list of strings
|
||||
assert isinstance(facts, list)
|
||||
assert all(isinstance(fact, str) for fact in facts)
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_caching_behavior_across_methods(self, mock_model_clients):
|
||||
"""Test that caching is properly enabled across different methods."""
|
||||
chat_history = "User: Testing caching"
|
||||
session_id = str(uuid4())
|
||||
|
||||
# Test single prompt caching
|
||||
with patch("src.deriver.tom.single_prompt.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(chat_history, session_id, method="single_prompt")
|
||||
|
||||
call_kwargs = mock_model_clients["single_prompt"].generate.call_args[1]
|
||||
assert call_kwargs["use_caching"] is True
|
||||
|
||||
# Test long term fact extraction caching
|
||||
with patch("src.deriver.tom.long_term.parse_xml_content") as mock_parse_xml:
|
||||
mock_parse_xml.return_value = '{"facts": []}'
|
||||
|
||||
await extract_facts_long_term(chat_history)
|
||||
|
||||
call_kwargs = mock_model_clients["long_term"].generate.call_args[1]
|
||||
assert call_kwargs["use_caching"] is True
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_observability_integration(self, mock_model_clients):
|
||||
"""Test that observability tools (Sentry, Langfuse) are properly integrated."""
|
||||
chat_history = "User: Testing observability"
|
||||
session_id = str(uuid4())
|
||||
|
||||
# Test single prompt observability
|
||||
with (
|
||||
patch("src.deriver.tom.single_prompt.sentry_sdk.start_transaction") as mock_sentry,
|
||||
patch("src.deriver.tom.single_prompt.langfuse_context.update_current_observation") as mock_langfuse
|
||||
):
|
||||
mock_sentry.return_value.__enter__.return_value = MagicMock()
|
||||
mock_sentry.return_value.__exit__.return_value = None
|
||||
|
||||
await get_tom_inference(chat_history, session_id, method="single_prompt")
|
||||
|
||||
# Verify Sentry transaction was started
|
||||
mock_sentry.assert_called_once_with(op="tom-inference", name="ToM Inference")
|
||||
|
||||
# Verify Langfuse observation was updated
|
||||
mock_langfuse.assert_called_once()
|
||||
|
||||
# Test conversational method observability
|
||||
mock_message = MagicMock()
|
||||
mock_message.content = [MagicMock(text="Test response")]
|
||||
|
||||
with (
|
||||
patch("src.deriver.tom.conversational.anthropic.messages.create") as mock_create,
|
||||
patch("src.deriver.tom.conversational.sentry_sdk.start_transaction") as mock_sentry,
|
||||
patch("src.deriver.tom.conversational.langfuse_context.update_current_observation") as mock_langfuse
|
||||
):
|
||||
mock_create.return_value = mock_message
|
||||
mock_sentry.return_value.__enter__.return_value = MagicMock()
|
||||
mock_sentry.return_value.__exit__.return_value = None
|
||||
|
||||
await get_tom_inference(chat_history, session_id, method="conversational")
|
||||
|
||||
# Verify observability integration
|
||||
mock_sentry.assert_called_once()
|
||||
mock_langfuse.assert_called_once()
|
||||
Loading…
Reference in New Issue