"""Comprehensive tests for BucketReRanker.""" import numpy as np from unittest.mock import Mock import pytest from suno_recs.worker.bucket_reranker import BucketReRanker from suno_recs.worker.constants import ( EMBEDDING_DIMS_BY_FIELD, HOOK_AUDIO_EMBEDDING_FIELD, HOOK_VIDEO_EMBEDDING_FIELD, CLIP_AUDIO_EMBEDDING_FIELD, ) def create_normalized_embedding(dim): """Create a random normalized embedding vector.""" vec = np.random.randn(dim) return (vec / np.linalg.norm(vec)).tolist() def create_mock_hook(hook_id, score=1.0, has_audio=True, has_video=True): """Create a mock hook with embeddings.""" hook = { 'hook_id': hook_id, '_score': score, } if has_audio: hook[HOOK_AUDIO_EMBEDDING_FIELD] = create_normalized_embedding( EMBEDDING_DIMS_BY_FIELD[HOOK_AUDIO_EMBEDDING_FIELD] ) if has_video: hook[HOOK_VIDEO_EMBEDDING_FIELD] = create_normalized_embedding( EMBEDDING_DIMS_BY_FIELD[HOOK_VIDEO_EMBEDDING_FIELD] ) return hook class TestBucketReRanker: """Test suite for BucketReRanker.""" def setup_method(self): """Set up test fixtures.""" self.mock_es_client = Mock() # Create some seed embeddings self.audio_seed1 = create_normalized_embedding(EMBEDDING_DIMS_BY_FIELD[HOOK_AUDIO_EMBEDDING_FIELD]) self.audio_seed2 = create_normalized_embedding(EMBEDDING_DIMS_BY_FIELD[HOOK_AUDIO_EMBEDDING_FIELD]) self.video_seed1 = create_normalized_embedding(EMBEDDING_DIMS_BY_FIELD[HOOK_VIDEO_EMBEDDING_FIELD]) self.clip_seed = create_normalized_embedding(EMBEDDING_DIMS_BY_FIELD[CLIP_AUDIO_EMBEDDING_FIELD]) # Mock ES client to return seed embeddings def mock_get_docs(ids, index_name, _source): docs = [] for id in ids: if index_name == "hook": if "audio_seed1" in id: docs.append({HOOK_AUDIO_EMBEDDING_FIELD: self.audio_seed1}) elif "audio_seed2" in id: docs.append({HOOK_AUDIO_EMBEDDING_FIELD: self.audio_seed2}) elif "video_seed" in id: docs.append({HOOK_VIDEO_EMBEDDING_FIELD: self.video_seed1}) else: docs.append({}) elif index_name == "clip": docs.append({CLIP_AUDIO_EMBEDDING_FIELD: self.clip_seed}) else: docs.append({}) return docs self.mock_es_client.get_documents_by_ids = mock_get_docs def test_insufficient_seeds(self): """Test that reranking is skipped when there are insufficient seeds.""" # Create reranker with only 2 seeds (threshold is 5) reranker = BucketReRanker( self.mock_es_client, audio_hook_seed_ids=['audio_seed1'], video_hook_seed_ids=['video_seed1'], ) hooks = [create_mock_hook(f'hook_{i}') for i in range(5)] # Should not have sufficient seeds assert not reranker.has_sufficient_seeds # Single bucket rerank should skip reranked, debug = reranker.rerank_bucket('test_bucket', hooks) assert reranked == hooks # No changes # When using batch method internally, debug structure is different # Batch rerank should also skip buckets = {'test_bucket': hooks} rerank_params = {'test_bucket': {'lambda': 0.5}} reranked_buckets, debug = reranker.rerank_all_buckets(buckets, rerank_params) assert reranked_buckets == buckets assert debug['skipped'] == 'insufficient_seeds' def test_empty_bucket(self): """Test handling of empty buckets.""" reranker = BucketReRanker( self.mock_es_client, audio_hook_seed_ids=['audio_seed1'] * 5, # Sufficient seeds ) reranked, debug = reranker.rerank_bucket('empty_bucket', []) assert reranked == [] # With the delegation to rerank_all_buckets, empty bucket is handled differently # The debug structure doesn't have 'skipped' or 'hooks_reranked' at top level assert debug.get('total_unique_hooks') == 0 # No hooks to process def test_basic_reranking(self): """Test basic reranking functionality.""" # Create reranker with sufficient seeds reranker = BucketReRanker( self.mock_es_client, audio_hook_seed_ids=['audio_seed1', 'audio_seed2'] * 2, video_hook_seed_ids=['video_seed1'] * 2, ) # Create hooks with varying scores hooks = [ create_mock_hook('hook_1', score=0.1), create_mock_hook('hook_2', score=0.5), create_mock_hook('hook_3', score=0.9), create_mock_hook('hook_4', score=0.3), create_mock_hook('hook_5', score=0.7), ] # Rerank with some personalization reranked, debug = reranker.rerank_bucket( 'test_bucket', hooks.copy(), blending_lambda=0.5, # 50% personalization preserve_top_n=0, # Don't preserve any ) # Check that all hooks were processed assert len(reranked) == len(hooks) assert all(h.get('reranked', False) for h in reranked) assert all('personalization_score' in h for h in reranked) assert all('blended_score' in h for h in reranked) # Check debug info assert debug['hooks_reranked'] == 5 assert 'elapsed_ms' in debug def test_preserve_top_n(self): """Test that top N items are preserved in their positions.""" reranker = BucketReRanker( self.mock_es_client, audio_hook_seed_ids=['audio_seed1'] * 5, ) # Create hooks with clear ordering hooks = [ create_mock_hook('top_1', score=10.0), create_mock_hook('top_2', score=9.0), create_mock_hook('top_3', score=8.0), create_mock_hook('low_1', score=1.0), create_mock_hook('low_2', score=0.5), ] reranked, debug = reranker.rerank_bucket( 'test_bucket', hooks.copy(), preserve_top_n=3, ) # First 3 should remain in place assert reranked[0]['hook_id'] == 'top_1' assert reranked[1]['hook_id'] == 'top_2' assert reranked[2]['hook_id'] == 'top_3' # Only bottom 2 should be marked as reranked assert not reranked[0].get('reranked', False) assert not reranked[1].get('reranked', False) assert not reranked[2].get('reranked', False) assert reranked[3].get('reranked', False) assert reranked[4].get('reranked', False) def test_mixed_modalities(self): """Test handling of hooks with different modality combinations.""" reranker = BucketReRanker( self.mock_es_client, audio_hook_seed_ids=['audio_seed1'] * 3, video_hook_seed_ids=['video_seed1'] * 2, ) # Create hooks with different modalities hooks = [ create_mock_hook('both', has_audio=True, has_video=True), create_mock_hook('audio_only', has_audio=True, has_video=False), create_mock_hook('video_only', has_audio=False, has_video=True), create_mock_hook('neither', has_audio=False, has_video=False), ] reranked, debug = reranker.rerank_bucket('test_bucket', hooks.copy()) # All should be returned assert len(reranked) == 4 # Check that reranked hooks have personalization scores # Note: hooks may not have personalization_score if they weren't in the rerank pool reranked_hooks = [h for h in reranked if h.get('reranked', False)] assert all('personalization_score' in h for h in reranked_hooks) def test_batch_reranking_multiple_buckets(self): """Test batch reranking of multiple buckets.""" reranker = BucketReRanker( self.mock_es_client, audio_hook_seed_ids=['audio_seed1'] * 5, ) # Create multiple buckets buckets = { 'fresh': [create_mock_hook(f'fresh_{i}', score=i*0.1) for i in range(5)], 'popular': [create_mock_hook(f'popular_{i}', score=i*0.2) for i in range(3)], 'champion': [create_mock_hook(f'champion_{i}', score=i*0.3) for i in range(4)], } rerank_params = { 'fresh': {'lambda': 0.3, 'preserve_top_n': 2}, 'popular': {'lambda': 0.5, 'preserve_top_n': 1}, 'champion': {'lambda': 0.7, 'preserve_top_n': 0}, } reranked_buckets, debug = reranker.rerank_all_buckets( buckets, rerank_params, max_candidates=100, ) # Check all buckets were processed assert set(reranked_buckets.keys()) == set(buckets.keys()) # Check each bucket assert len(reranked_buckets['fresh']) == 5 assert len(reranked_buckets['popular']) == 3 assert len(reranked_buckets['champion']) == 4 # Check debug info assert 'total_unique_hooks' in debug assert 'batch_elapsed_ms' in debug assert 'per_bucket_debug' in debug # Verify preserve_top_n was respected assert not reranked_buckets['fresh'][0].get('reranked', False) assert not reranked_buckets['fresh'][1].get('reranked', False) assert reranked_buckets['fresh'][2].get('reranked', False) assert not reranked_buckets['popular'][0].get('reranked', False) assert reranked_buckets['popular'][1].get('reranked', False) assert all(h.get('reranked', False) for h in reranked_buckets['champion']) def test_embedding_fetching(self): """Test that missing embeddings are fetched correctly.""" reranker = BucketReRanker( self.mock_es_client, audio_hook_seed_ids=['audio_seed1'] * 5, ) # Create hooks without embeddings hooks = [ {'hook_id': 'hook_1', '_score': 1.0}, # No embeddings create_mock_hook('hook_2'), # Has embeddings ] # Save original mock function original_mock = self.mock_es_client.get_documents_by_ids # Mock ES to return embeddings when fetched def mock_get_docs_with_fetch(ids, index_name, _source): if index_name == "hook" and "hook_1" in ids: return [{ HOOK_AUDIO_EMBEDDING_FIELD: create_normalized_embedding( EMBEDDING_DIMS_BY_FIELD[HOOK_AUDIO_EMBEDDING_FIELD] ), HOOK_VIDEO_EMBEDDING_FIELD: create_normalized_embedding( EMBEDDING_DIMS_BY_FIELD[HOOK_VIDEO_EMBEDDING_FIELD] ), }] # Call the original mock function to handle seed fetching return original_mock(ids, index_name, _source) self.mock_es_client.get_documents_by_ids = mock_get_docs_with_fetch # Rerank with preserve_top_n=0 to ensure hooks are actually reranked reranked, debug = reranker.rerank_bucket( 'test_bucket', hooks, preserve_top_n=0 # Don't preserve any ) # Check that reranked hooks have personalization scores reranked_hooks = [h for h in reranked if h.get('reranked', False)] assert len(reranked_hooks) > 0 assert all('personalization_score' in h for h in reranked_hooks) # Check that hook_1 now has embeddings (fetched from ES) hook_1_result = next(h for h in reranked if h['hook_id'] == 'hook_1') assert HOOK_AUDIO_EMBEDDING_FIELD in hook_1_result assert HOOK_VIDEO_EMBEDDING_FIELD in hook_1_result def test_blending_lambda_effects(self): """Test that blending lambda properly weights personalization.""" reranker = BucketReRanker( self.mock_es_client, audio_hook_seed_ids=['audio_seed1'] * 5, ) hooks = [create_mock_hook(f'hook_{i}', score=i*0.2) for i in range(5)] # Test with no personalization reranked_0, _ = reranker.rerank_bucket( 'test', hooks.copy(), blending_lambda=0.0 ) # Test with full personalization reranked_1, _ = reranker.rerank_bucket( 'test', hooks.copy(), blending_lambda=1.0 ) # Test with balanced reranked_5, _ = reranker.rerank_bucket( 'test', hooks.copy(), blending_lambda=0.5 ) # Check that reranked hooks have blended scores for reranked in [reranked_0, reranked_1, reranked_5]: reranked_hooks = [h for h in reranked if h.get('reranked', False)] assert all('blended_score' in h for h in reranked_hooks) assert all('personalization_score' in h for h in reranked_hooks) def test_clip_seeds_integration(self): """Test that clip seeds are properly integrated.""" reranker = BucketReRanker( self.mock_es_client, audio_hook_seed_ids=['audio_seed1'], audio_clip_seed_ids=['clip_seed1'] * 4, # Total 5 seeds ) hooks = [create_mock_hook('hook_1')] # Should have sufficient seeds assert reranker.has_sufficient_seeds reranked, debug = reranker.rerank_bucket('test_bucket', hooks) assert len(reranked) == 1 # With only 1 hook and default preserve_top_n=3, it may be preserved def test_consistency_between_methods(self): """Test that single and batch methods produce identical results.""" reranker = BucketReRanker( self.mock_es_client, audio_hook_seed_ids=['audio_seed1'] * 5, ) hooks = [create_mock_hook(f'hook_{i}', score=i*0.1) for i in range(10)] # Single bucket rerank single_result, single_debug = reranker.rerank_bucket( 'test_bucket', hooks.copy(), blending_lambda=0.4, preserve_top_n=2, ) # Extract personalization scores from reranked hooks only single_scores = {h['hook_id']: h.get('personalization_score', 0) for h in single_result if h.get('reranked', False)} # Reset seeds to ensure same computation reranker2 = BucketReRanker( self.mock_es_client, audio_hook_seed_ids=['audio_seed1'] * 5, ) # Batch rerank buckets = {'test_bucket': hooks.copy()} rerank_params = {'test_bucket': {'lambda': 0.4, 'preserve_top_n': 2}} batch_result, batch_debug = reranker2.rerank_all_buckets(buckets, rerank_params) # Extract personalization scores from reranked hooks only batch_scores = {h['hook_id']: h.get('personalization_score', 0) for h in batch_result['test_bucket'] if h.get('reranked', False)} # Since we can't control random embeddings exactly, we can't compare scores # But we can verify structure assert len(single_result) == len(batch_result['test_bucket']) assert set(h['hook_id'] for h in single_result) == set(h['hook_id'] for h in batch_result['test_bucket']) if __name__ == '__main__': pytest.main([__file__, '-v'])