from typing import Collection, Iterable from .... import logger from ....constants import DBAuditLogActions, RowTypes, Tables from ....models import NewAuditFlag from ....typings import AuditGroupID, Rows, UserID from ....utils.sql import placeholders logger = logger.new_logger(__name__) class Flags: """Class for interacting with audit flags data.""" def __init__(self, client): self.client = client async def add( self, audit_id: int, user: UserID, flags: Iterable[NewAuditFlag] ) -> None: """Add audit flags. Args: audit_id: Audit ID. user: User ID performing the action. flags: Flags to add. """ query = f""" INSERT INTO {Tables.AUDITS_FLAGS} (audit, row_idx, upc, isrc, asset_id, video_id, flag, details, date_created_utc, created_by, resolution, resolution_subtype, date_resolved_utc) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, UTC_TIMESTAMP(), %s, %s, %s, IF(%s IS NOT NULL, UTC_TIMESTAMP(), NULL)); """ params = [ ( audit_id, flag.row_idx, flag.upc, flag.isrc, flag.asset_id, flag.video_id, flag.text, flag.details, user, flag.resolution, # For auto-resolved flags only flag.resolution_subtype, # For auto-resolved flags only flag.resolution, # Merely to determine date_resolved_utc ) for flag in flags ] conn = await self.client.db.get_connection() with conn: with conn.cursor() as cur: cur.executemany(query, params) conn.commit() async def resolve( self, flags: Collection[int], user: UserID, resolution: str, resolution_subtype: str | None = None, ) -> Rows: """Resolve flags. Args: flags: Flag IDs. user: User ID performing the action. resolution: Resolution type. resolution_subtype: Resolution subtype, if applicable. """ conn = await self.client.db.get_connection() with conn: try: with conn.cursor() as cur: # Execute the UPDATE query cur.execute( f""" UPDATE {Tables.AUDITS_FLAGS} SET date_resolved_utc = UTC_TIMESTAMP(), resolved_by = %s, resolution = %s, resolution_subtype = %s WHERE id IN ({placeholders(flags)}); """, (user, resolution, resolution_subtype, *flags), ) # Execute the SELECT query immediately after query = f""" SELECT {self.client._audit_flag_columns}, CREATED_BY.nickname AS created_by_nickname, RESOLVED_BY.nickname AS resolved_by_nickname FROM {Tables.AUDITS_FLAGS} AF LEFT JOIN {Tables.USERS} CREATED_BY ON AF.created_by = CREATED_BY.id LEFT JOIN {Tables.USERS} RESOLVED_BY ON AF.resolved_by = RESOLVED_BY.id WHERE AF.id IN ({placeholders(flags)}); """ cur.execute(query, (*flags,)) # Fetch and return the results resp = cur.fetchall() conn.commit() # Commit the transaction after both queries return resp except Exception as ex: logger.error("Exception occurred: {}", ex) conn.rollback() # Rollback in case of an exception async def unique_rows( self, audit_group_ids: Collection[AuditGroupID], ) -> Rows: """Get counts of unique flagged rows for audits in one or more audit groups. Args: audit_group_ids: Audit group IDs. """ audit_group_ids = list(set(audit_group_ids)) query = f""" WITH AUDITS_IN_GROUP AS ( SELECT `group`, `id`, type FROM {Tables.AUDITS} WHERE `group` IN ({placeholders(audit_group_ids)}) ), FILTERED_FLAGS AS ( SELECT AUDITS_IN_GROUP.`group`, AUDITS_IN_GROUP.`id` AS audit, row_idx, resolution, resolution_subtype, (resolution IS NULL) AS is_pending FROM {Tables.AUDITS_FLAGS} AF RIGHT JOIN AUDITS_IN_GROUP ON AF.audit = AUDITS_IN_GROUP.`id` ), ANALYZED_ROWS AS ( SELECT audit, COUNT(*) AS row_count FROM {Tables.AUDITS_ROWS} AR RIGHT JOIN AUDITS_IN_GROUP ON AR.audit = AUDITS_IN_GROUP.`id` LEFT JOIN {Tables.AUDITS_ROWS_ROWTYPES} ARRT ON AR.id = ARRT.row_id LEFT JOIN {Tables.AUDITS_ROWTYPES} ART ON ARRT.rowtype_id = ART.id WHERE ART.type = '{RowTypes.ANALYZED}' GROUP BY audit ) SELECT FILTERED_FLAGS.`group`, FILTERED_FLAGS.audit AS audit, ANY_VALUE(AL.row_count) AS rows_total, IFNULL(ANALYZED_ROWS.row_count, 0) AS rows_analyzed, COUNT(DISTINCT FILTERED_FLAGS.row_idx) AS rows_flagged, COUNT(DISTINCT IF(FILTERED_FLAGS.is_pending = 1, FILTERED_FLAGS.row_idx, NULL)) AS rows_pending FROM FILTERED_FLAGS LEFT JOIN {Tables.AUDITS_LOG} AL ON AL.audit = FILTERED_FLAGS.audit LEFT JOIN ANALYZED_ROWS ON ANALYZED_ROWS.audit = FILTERED_FLAGS.audit WHERE AL.action = '{DBAuditLogActions.COMPLETED_FETCH}' GROUP BY 1, 2, 3, 4 """ data = await self.client.db.execute_query_fetchall(query, (audit_group_ids,)) return data or []