"""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'])
