""" Hive Segment Model. This hive_segment model is used to store information about hive_segment table. """ from typing import Any from sqlalchemy import ( TIMESTAMP, Column, Float, ForeignKey, Index, Integer, and_, case, delete, func, insert, text, ) from sqlalchemy.exc import IntegrityError from assets.connectors import mysql from assets.constants import api, asset_types from assets.constants.ai_detection import ( AI_MAX_SCORE_THRESHOLD, AI_PERCENTAGE_THRESHOLD, AI_SEGMENT_SCORE_THRESHOLD, AIDetectionResultCode, ) from assets.exceptions import HiveSegmentExists from assets.models import hive_task as hive_task_model from assets.models.asset_final import AssetFinal from assets.models.asset_upload import AssetUpload class HiveSegment(mysql.AuModel): """Table definition for hive_segment table.""" __tablename__ = "hive_segment" asset_final_id = Column(Integer, nullable=False, primary_key=True) time = Column(Integer, nullable=False, primary_key=True) ai_generated_music = Column(Float, nullable=True) ai_generated_music_vocal = Column(Float, nullable=True) mubert = Column(Float, nullable=True) musicgen = Column(Float, nullable=True) riffusion = Column(Float, nullable=True) stable_audio = Column(Float, nullable=True) suno = Column(Float, nullable=True) udio = Column(Float, nullable=True) yue = Column(Float, nullable=True) minimax = Column(Float, nullable=True) mureka = Column(Float, nullable=True) ace_step = Column(Float, nullable=True) duobao = Column(Float, nullable=True) google = Column(Float, nullable=True) heartmula = Column(Float, nullable=True) loudly = Column(Float, nullable=True) task_id = Column(Integer, ForeignKey("hive_task.id")) created_at = Column( TIMESTAMP, nullable=False, server_default=text("CURRENT_TIMESTAMP") ) __table_args__ = (Index("asset_final_id", "time"),) def as_dict(self) -> dict[str, Any]: """Return object as dict. Returns: dict: Dictionary representation of object """ hive_segment_dict = { "asset_final_id": self.asset_final_id, "time": self.time, "ai_generated_music": self.ai_generated_music, "ai_generated_music_vocal": self.ai_generated_music_vocal, "mubert": self.mubert, "musicgen": self.musicgen, "riffusion": self.riffusion, "stable_audio": self.stable_audio, "suno": self.suno, "udio": self.udio, "yue": self.yue, "minimax": self.minimax, "mureka": self.mureka, "ace_step": self.ace_step, "duobao": self.duobao, "google": self.google, "heartmula": self.heartmula, "loudly": self.loudly, "task_id": self.task_id, } return hive_segment_dict def create_hive_segment( asset_final_id: int, task: dict[str, Any], segments: list[dict[str, Any]], overwrite: bool = False, ) -> None: """Create hive_segment item in the table. Args: asset_final_id (int): Asset final id. task (dict): hive_task data. task_id (str): task id model (str): model_version (int): segments (list): List of segment dicts. time (int): time. ai_generated_music (float): ai_generated_music_vocal (float): mubert (float): musicgen (float): riffusion (float): stable_audio (float): suno (float): udio (float): overwrite (bool): When True, atomically replace any existing scan for this asset (delete then insert) instead of raising HiveSegmentExists. Used by the re-scan/backfill path; defaults to False (first-scan-wins). """ try: with mysql.au_db_session() as session: if overwrite: # fk_hive_segment_task blocks deleting a hive_task while its # hive_segment children still reference it, so delete the children # first, then the parent, before re-inserting. session.execute( delete(HiveSegment).where( HiveSegment.asset_final_id == asset_final_id ) ) hive_task_model._delete_hive_task(session, asset_final_id) task_id = hive_task_model._create_hive_task(session, asset_final_id, task) hive_segments = [ HiveSegment( asset_final_id=asset_final_id, task_id=task_id, time=x["time"], ai_generated_music=x["ai_generated_music"], ai_generated_music_vocal=x["ai_generated_music_vocal"], mubert=x["mubert"], musicgen=x["musicgen"], riffusion=x["riffusion"], stable_audio=x["stable_audio"], suno=x["suno"], udio=x["udio"], yue=x["yue"], minimax=x["minimax"], mureka=x["mureka"], ace_step=x["ace_step"], duobao=x["duobao"], google=x["google"], heartmula=x["heartmula"], loudly=x["loudly"], ) for x in segments ] session.execute(insert(HiveSegment), [x.as_dict() for x in hive_segments]) except IntegrityError as e: if overwrite: # Existing rows were deleted above, so a duplicate is impossible here: # any IntegrityError is a real constraint failure (FK, etc.). Let it # surface as a 500 (logged + Sentry) rather than a misleading # "already exists" 409 that hides the true cause. raise raise HiveSegmentExists( f"hive segment data already exist for asset_final_id {asset_final_id}" ) from e def _classify_ai_detection_result( segment_count: int, max_score: float, percent_above: float ) -> AIDetectionResultCode: """Classify AI detection result based on threshold criteria. Args: max_score: Maximum AI detection score across all segments percent_above: Percentage of segments above threshold Returns: AIDetectionResultCode: AI detection result code (SUSPECTED or NOT_SUSPECTED) """ if segment_count == 0: return AIDetectionResultCode.NO_ANALYSIS_DATA if (max_score >= AI_MAX_SCORE_THRESHOLD) or ( percent_above >= AI_PERCENTAGE_THRESHOLD ): return AIDetectionResultCode.AI_GENERATED_AUDIO_SUSPECTED return AIDetectionResultCode.NO_AI_GENERATED_AUDIO_SUSPECTED def get_ai_generated_audio_results_by_product( product_id: int, track_ids: list[int], ) -> list[dict[str, int | AIDetectionResultCode]]: """Query AI detection data from database and classify results. Args: product_id (int): Product ID to filter asset uploads track_ids (list[int]): Track IDs to filter asset uploads Returns: list[dict]: List of dicts with track_id, asset_final_id, and code """ with mysql.au_db_session() as session: # Subquery to get latest upload_id per (track, upload_type) latest_uploads = ( session.query( AssetUpload.track_unique_id, AssetUpload.asset_upload_type_id.label("asset_upload_type_id"), func.max(AssetUpload.asset_upload_id).label("latest_upload_id"), ) .filter( AssetUpload.product_id == product_id, AssetUpload.deleted == 0, AssetUpload.track_unique_id.in_(track_ids), AssetUpload.api_version == api.API_VERSION_V2, ) .group_by(AssetUpload.track_unique_id, AssetUpload.asset_upload_type_id) .subquery() ) query_results = ( session.query( AssetUpload.track_unique_id.label("track_unique_id"), AssetFinal.asset_final_id.label("asset_final_id"), func.count(HiveSegment.asset_final_id).label("segment_count"), func.coalesce(func.max(HiveSegment.ai_generated_music), 0).label( "max_score" ), func.coalesce( case( ( func.count(HiveSegment.asset_final_id) > 0, func.count( case( ( HiveSegment.ai_generated_music >= AI_SEGMENT_SCORE_THRESHOLD, 1, ), else_=None, ) ) / func.count(HiveSegment.asset_final_id), ), else_=None, ), 0, ).label("percent_above_threshold"), ) .join( latest_uploads, and_( AssetUpload.track_unique_id == latest_uploads.c.track_unique_id, AssetUpload.asset_upload_type_id == latest_uploads.c.asset_upload_type_id, ), ) .filter(AssetUpload.asset_upload_id == latest_uploads.c.latest_upload_id) .join(AssetFinal, AssetUpload.asset_upload_id == AssetFinal.asset_upload_id) .outerjoin( HiveSegment, AssetFinal.asset_final_id == HiveSegment.asset_final_id ) .filter(AssetFinal.asset_type == asset_types.TYPE_FILE_FLAC) .group_by(AssetUpload.track_unique_id, AssetFinal.asset_final_id) .all() ) return [ { "track_id": row.track_unique_id, "asset_final_id": row.asset_final_id, "code": _classify_ai_detection_result( row.segment_count, row.max_score, row.percent_above_threshold, ), } for row in query_results ]