"""Tests for HiveSegment model.""" from typing import Any import pytest from flexmock import flexmock from sqlalchemy.exc import IntegrityError, SQLAlchemyError from assets.connectors import mysql from assets.constants import field_const from assets.constants.ai_detection import AIDetectionResultCode from assets.exceptions import HiveSegmentExists from assets.models import hive_segment, hive_task from tests.unit.models import au_operations @pytest.fixture def valid_hive_segment_data() -> dict[str, Any]: """Fixture for valid data for create_hive_segment.""" return { "asset_final_id": 555, field_const.TASK: { field_const.TASK_ID: "a22414c0-0cba-11f1-bc98-c7fc13f309f3", field_const.MODEL: "ai_music_classifier_DORIAN_2025_04_02_v00", field_const.MODEL_VERSION: 1, }, "segments": [ { "time": 700, "ai_generated_music": 1.99, "ai_generated_music_vocal": 0.55, "mubert": 1, "musicgen": 1, "riffusion": 1, "stable_audio": 1, "suno": 1, "udio": 1, "yue": 1, "minimax": 1, "mureka": 1, "ace_step": 1, "duobao": 1, "google": 1, "heartmula": 1, "loudly": 1, } ], # 'created_at': '2025-08-25 15:58:50', } @pytest.fixture def db_fixture( asset_upload_data_for_hive: list[dict[str, Any]], asset_final_data: list[dict[str, Any]], hive_task_data: dict[str, Any], hive_segment_data: list[dict[str, Any]], ) -> None: """Set up the hive_segment table.""" au_operations.truncate_tables() au_operations.seed_asset_upload_table(asset_upload_data_for_hive) au_operations.seed_asset_final_table(asset_final_data) au_operations.seed_hive_task_table(hive_task_data) au_operations.seed_hive_segment_table(hive_segment_data) def test_create_hive_segment( db_fixture: None, valid_hive_segment_data: dict[str, Any] ) -> None: """Test creating a new hive_segment.""" hive_segment.create_hive_segment(**valid_hive_segment_data) assert True def test_create_hive_segment_error( db_fixture: None, valid_hive_segment_data: dict[str, Any] ) -> None: """Test creating a new HiveSegment with sql error.""" error_message = "Query error" ( flexmock(mysql) .should_receive("au_db_session") .and_raise(SQLAlchemyError(error_message)) ) with pytest.raises(SQLAlchemyError) as exc: hive_segment.create_hive_segment(**valid_hive_segment_data) assert str(exc.value) == error_message def test_create_hive_segment_integrity_error( db_fixture: None, valid_hive_segment_data: dict[str, Any] ) -> None: """Test creating a new HiveSegment with sql error.""" ( flexmock(mysql) .should_receive("au_db_session") .and_raise(IntegrityError("foo", None, Exception())) ) with pytest.raises(HiveSegmentExists) as exc: hive_segment.create_hive_segment(**valid_hive_segment_data) assert ( str(exc.value) == f"409 Conflict: hive segment data already exist for asset_final_id {valid_hive_segment_data['asset_final_id']}" ) def _full_segment(time: int, score: float = 0.1) -> dict[str, Any]: """Segment dict with all score keys present (mostly None); create_hive_segment reads every key.""" return { "time": time, "ai_generated_music": score, "ai_generated_music_vocal": score, "mubert": None, "musicgen": None, "riffusion": None, "stable_audio": None, "suno": None, "udio": None, "yue": None, "minimax": None, "mureka": None, "ace_step": None, "duobao": None, "google": None, "heartmula": None, "loudly": None, } def test_create_hive_segment_overwrite_replaces_existing( db_fixture: None, valid_hive_segment_data: dict[str, Any] ) -> None: """overwrite=True replaces an existing scan instead of raising HiveSegmentExists.""" hive_segment.create_hive_segment(**valid_hive_segment_data) rescan = { "asset_final_id": valid_hive_segment_data["asset_final_id"], field_const.TASK: { field_const.TASK_ID: "b33525d1-1dcb-42e2-cd09-d8fd24e410a4", field_const.MODEL: "ai_music_classifier_HYMN_2026_04_24_v00", field_const.MODEL_VERSION: 2, }, "segments": [_full_segment(3), _full_segment(6)], } hive_segment.create_hive_segment(**rescan, overwrite=True) with mysql.au_db_session() as session: segment_times = { row.time for row in session.query(hive_segment.HiveSegment.time) .filter(hive_segment.HiveSegment.asset_final_id == 555) .all() } task_models = [ row.model for row in session.query(hive_task.HiveTask.model) .filter(hive_task.HiveTask.asset_final_id == 555) .all() ] assert segment_times == {3, 6} # new segments; old time 700 is gone assert task_models == [ "ai_music_classifier_HYMN_2026_04_24_v00" ] # 1 task, replaced def test_create_hive_segment_overwrite_inserts_when_none_exists( db_fixture: None, valid_hive_segment_data: dict[str, Any] ) -> None: """overwrite=True with no prior scan just inserts (the delete is a no-op).""" hive_segment.create_hive_segment(**valid_hive_segment_data, overwrite=True) with mysql.au_db_session() as session: segment_count = ( session.query(hive_segment.HiveSegment) .filter(hive_segment.HiveSegment.asset_final_id == 555) .count() ) task_count = ( session.query(hive_task.HiveTask) .filter(hive_task.HiveTask.asset_final_id == 555) .count() ) assert segment_count == 1 assert task_count == 1 def test_create_hive_segment_overwrite_only_touches_target_asset( db_fixture: None, valid_hive_segment_data: dict[str, Any] ) -> None: """overwrite deletes only the target asset's rows; a sibling asset's scan survives.""" # Seed data gives asset 111 one hive_segment and one hive_task (id 11). # Overwrite asset 555 (valid_hive_segment_data); assert 111 is untouched below. hive_segment.create_hive_segment(**valid_hive_segment_data, overwrite=True) with mysql.au_db_session() as session: sibling_segments = ( session.query(hive_segment.HiveSegment) .filter(hive_segment.HiveSegment.asset_final_id == 111) .count() ) sibling_tasks = ( session.query(hive_task.HiveTask) .filter(hive_task.HiveTask.asset_final_id == 111) .count() ) assert sibling_segments == 1 # asset 111's segment untouched assert sibling_tasks == 1 # asset 111's hive_task untouched def test_create_hive_segment_overwrite_legacy_asset_without_task_row( db_fixture: None, ) -> None: """overwrite replaces a pre-task-era scan (segment with NULL task_id, no hive_task row).""" # Seed data gives asset 222 one segment (time 300, task_id NULL) and no hive_task row. rescan: dict[str, Any] = { "asset_final_id": 222, field_const.TASK: { field_const.TASK_ID: "b33525d1-1dcb-42e2-cd09-d8fd24e410a4", field_const.MODEL: "ai_music_classifier_HYMN_2026_04_24_v00", field_const.MODEL_VERSION: 2, }, "segments": [_full_segment(9)], } hive_segment.create_hive_segment(**rescan, overwrite=True) with mysql.au_db_session() as session: segment_times = { row.time for row in session.query(hive_segment.HiveSegment.time) .filter(hive_segment.HiveSegment.asset_final_id == 222) .all() } task_count = ( session.query(hive_task.HiveTask) .filter(hive_task.HiveTask.asset_final_id == 222) .count() ) assert segment_times == {9} # legacy time-300 segment replaced assert task_count == 1 # backfill created the task row def test_create_hive_segment_overwrite_reraises_integrity_error( db_fixture: None, ) -> None: """In overwrite mode a real integrity failure propagates, not masked as a 409 HiveSegmentExists.""" # Two segments share a time -> collide on the (asset_final_id, time) PK. Since overwrite # deleted first, this is a genuine constraint failure, not a pre-existing duplicate, so it # must surface as IntegrityError (-> 500 + Sentry), not HiveSegmentExists. payload: dict[str, Any] = { "asset_final_id": 555, field_const.TASK: { field_const.TASK_ID: "b33525d1-1dcb-42e2-cd09-d8fd24e410a4", field_const.MODEL: "ai_music_classifier_HYMN_2026_04_24_v00", field_const.MODEL_VERSION: 2, }, "segments": [_full_segment(5), _full_segment(5)], } with pytest.raises(IntegrityError): hive_segment.create_hive_segment(**payload, overwrite=True) def test_create_hive_segment_default_conflicts_and_preserves_existing( db_fixture: None, valid_hive_segment_data: dict[str, Any] ) -> None: """Without overwrite, a re-scan raises HiveSegmentExists and leaves the original scan intact.""" # First scan: asset 555, time 700. hive_segment.create_hive_segment(**valid_hive_segment_data) # A second scan without overwrite must conflict (overwrite defaults to False). with pytest.raises(HiveSegmentExists): hive_segment.create_hive_segment(**valid_hive_segment_data) with mysql.au_db_session() as session: segment_times = { row.time for row in session.query(hive_segment.HiveSegment.time) .filter(hive_segment.HiveSegment.asset_final_id == 555) .all() } task_count = ( session.query(hive_task.HiveTask) .filter(hive_task.HiveTask.asset_final_id == 555) .count() ) assert segment_times == {700} # original preserved, not clobbered assert task_count == 1 # the rejected attempt's task insert rolled back def test__classify_ai_detection_result() -> None: """Test AI detection result classification logic.""" result = hive_segment._classify_ai_detection_result( segment_count=0, max_score=0.9, percent_above=0.3 ) assert result == AIDetectionResultCode.NO_ANALYSIS_DATA result = hive_segment._classify_ai_detection_result( segment_count=10, max_score=0.999999, percent_above=0.3 ) assert result == AIDetectionResultCode.AI_GENERATED_AUDIO_SUSPECTED result = hive_segment._classify_ai_detection_result( segment_count=5, max_score=1.5, percent_above=0.3 ) assert result == AIDetectionResultCode.AI_GENERATED_AUDIO_SUSPECTED result = hive_segment._classify_ai_detection_result( segment_count=8, max_score=0.9, percent_above=0.5 ) assert result == AIDetectionResultCode.AI_GENERATED_AUDIO_SUSPECTED result = hive_segment._classify_ai_detection_result( segment_count=12, max_score=0.9, percent_above=0.75 ) assert result == AIDetectionResultCode.AI_GENERATED_AUDIO_SUSPECTED result = hive_segment._classify_ai_detection_result( segment_count=7, max_score=0.9, percent_above=0.3 ) assert result == AIDetectionResultCode.NO_AI_GENERATED_AUDIO_SUSPECTED def test_get_ai_generated_audio_results_by_product_error(db_fixture: None) -> None: """Test creating a new HiveSegment with sql error.""" error_message = "Query error" ( flexmock(mysql) .should_receive("au_db_session") .and_raise(SQLAlchemyError(error_message)) ) with pytest.raises(SQLAlchemyError) as exc: hive_segment.get_ai_generated_audio_results_by_product( product_id=1001, track_ids=[123, 234] ) assert str(exc.value) == error_message def test_get_ai_generated_audio_results_by_product(db_fixture: None) -> None: """Test getting AI generated audio results by product.""" product_id = 1001 track_ids = [123, 234, 345, 456, 567] results = hive_segment.get_ai_generated_audio_results_by_product( product_id, track_ids ) expected = [ { "track_id": 123, "asset_final_id": 111, "code": AIDetectionResultCode.AI_GENERATED_AUDIO_SUSPECTED, }, { "track_id": 234, "asset_final_id": 222, "code": AIDetectionResultCode.AI_GENERATED_AUDIO_SUSPECTED, }, { "track_id": 345, "asset_final_id": 333, "code": AIDetectionResultCode.NO_AI_GENERATED_AUDIO_SUSPECTED, }, { "track_id": 456, "asset_final_id": 444, "code": AIDetectionResultCode.AI_GENERATED_AUDIO_SUSPECTED, }, { "track_id": 567, "asset_final_id": 555, "code": AIDetectionResultCode.NO_ANALYSIS_DATA, }, ] assert len(results) == 5 assert results == expected