import asyncio import json from collections import deque from datetime import datetime from decimal import Decimal from .... import logger from ....constants import ( CREATED_PREFIX, DB, AuditStatus, AuditTypes, Auth0Users, DBAuditLogActions, DBColumns, Tables, ) from ....typings import AuditGroupID, DateType, LabelID, LabelName, Rows, UserID from ....utils.sql import placeholders logger = logger.new_logger(__name__) class AuditGroups: """Class for interacting with audits data.""" def __init__(self, client): self.client = client async def new( self, user: UserID, label_id: LabelID, label_name: LabelName, *, include_audio: bool = True, include_video: bool = False, include_art_track: bool = False, scheduled_for: datetime | None = None, ) -> int: """Create audit for user. Args: user: User ID. label_id: ID of the label to audit. label_name: Name of the label to audit. include_audio: Whether the audit includes sound recordings. include_video: Whether the audit includes video. include_art_track: Whether the audit includes art track. scheduled_for: Date to schedule the audit for. If None, the audit is created for immediate execution. Returns: Created audit group ID. """ conn = await self.client.db.get_connection() with conn: with conn.cursor() as cur: try: cur.execute( f""" INSERT INTO {Tables.AUDITS_GROUPS} (user, label_id, label_name) VALUES (%s, %s, %s) """, (user, label_id, label_name), ) created_audit_group_id = cur.lastrowid audit_types = [ (AuditTypes.SR.value, include_audio), (AuditTypes.MV.value, include_video), (AuditTypes.AT.value, include_art_track), ] audits = { audit_type: None for audit_type, include in audit_types if include } for audit in audits.keys(): cur.execute( f""" INSERT INTO {Tables.AUDITS} (`group`, `type`) VALUES (%s, %s) """, (created_audit_group_id, audit), ) audits[audit] = cur.lastrowid action = ( DBAuditLogActions.CREATED_IMMEDIATE.value if scheduled_for is None else DBAuditLogActions.CREATED_SCHEDULED.value ) cur.executemany( f""" INSERT INTO {Tables.AUDITS_LOG} (audit, action, date_created_utc, date_scheduled_utc, user) VALUES (%s, %s, UTC_TIMESTAMP(), %s, %s) """, [ (audit_id, action, scheduled_for, user) for audit_id in audits.values() ], ) conn.commit() except Exception as ex: conn.rollback() logger.error( f"Transaction rolled back due to error creating audit: {ex}" ) return created_audit_group_id async def get(self, audit_group_id: AuditGroupID) -> dict | None: """Return the audit group data for a given audit group ID. Args: audit_group_id: Audit group ID to get. """ query = f"""SELECT * FROM {Tables.AUDITS_GROUPS} WHERE id = %s""" return await self.client.db.execute_query_fetchone(query, (audit_group_id,)) async def archive(self, audit_group_ids: list[AuditGroupID]) -> None: """Archive one or more audit groups. Args: audit_group_ids: Audit group IDs to archive. """ unique_audit_group_ids = set(audit_group_ids) query = f""" UPDATE {Tables.AUDITS_GROUPS} SET {DBColumns.DATE_ARCHIVED} = UTC_TIMESTAMP() WHERE id IN ({placeholders(unique_audit_group_ids)}) """ await self.client.db.execute_query_nofetch(query, tuple(unique_audit_group_ids)) async def delete(self, audit_group_ids: list[AuditGroupID]) -> None: """Permanently delete one or more audit groups, including all associated audits, flags and logs. Args: audit_group_ids: Audit group IDs to delete. """ unique_audit_group_ids = set(audit_group_ids) query = f""" DELETE FROM {Tables.AUDITS_GROUPS} WHERE id IN ({placeholders(unique_audit_group_ids)}) """ await self.client.db.execute_query_nofetch(query, tuple(unique_audit_group_ids)) async def purge(self, days_ago: int = None) -> None: """Permanently delete all audit groups that were archived at or more than `days_ago` days ago. Args: days_ago: Positive integer; number of days to look back for archived audit groups. If 0 or None, all archived audit groups are deleted. """ days_ago = days_ago or 0 query = f""" DELETE FROM {Tables.AUDITS_GROUPS} WHERE {DBColumns.DATE_ARCHIVED} < (UTC_TIMESTAMP() - INTERVAL %s - 1 DAY) """ await self.client.db.execute_query_nofetch(query, (max(0, days_ago),)) async def list( self, *, created_by: UserID | None = None, status: AuditStatus | None = None, text: str | None = None, audit_types: list[AuditTypes] | None = None, scheduled_date_start: str | DateType | None = None, scheduled_date_end: str | DateType | None = None, score_min: int | None = None, score_max: int | None = None, order_by: str = "date_created_utc", order_desc: bool = True, offset: int | None = 0, limit: int | None = 200, archived: bool = True, ) -> tuple[list[dict], int]: """List audit groups with aggregated data. Args: created_by: User ID that created the audit. If None, all users are included. status: Audit status to filter by. If None, all statuses are included. text: Text to search for in the label name, label ID, or group ID. If provided, only audits that partially match the text will be included. audit_types: List of audit types to filter by. If None, all types are included. scheduled_date_start: Start date for the scheduled date range. Must be in the format "YYYY-MM-DD". scheduled_date_end: End date for the scheduled date range. Must be in the format "YYYY-MM-DD". score_min: Minimum score to filter by. If None, all scores are included. score_max: Maximum score to filter by. If None, all scores are included. order_by: Column to order by. order_desc: Whether to order in descending order (True) or ascending order (False). offset: Offset for the query. limit: Limit for the query. archived: Whether to include archived audit groups in the results or not. Returns: List of audit groups included in the specified range (offset, limit) and the total number of rows matching the query, without the limit. """ if created_by is not None and not isinstance(created_by, int): raise ValueError("`created_by` must be an integer.") query = f""" WITH SELECTED_GROUPS AS ( SELECT AG.id AS group_id FROM {Tables.AUDITS_GROUPS} AG WHERE 1 = 1 {"" if archived else f"AND AG.{DBColumns.DATE_ARCHIVED} IS NULL"} AND EXISTS ( SELECT 1 FROM {Tables.AUDITS} A JOIN {Tables.AUDITS_LOG} AL ON AL.audit = A.id WHERE A.`group` = AG.id AND AL.action LIKE '{CREATED_PREFIX}%%' ) ), AUDITS_IN_SCOPE AS ( SELECT A.id AS audit_id, A.`group` AS group_id, A.type AS audit_type FROM {Tables.AUDITS} A JOIN SELECTED_GROUPS SG ON SG.group_id = A.`group` ), AUDIT_FACTS AS ( SELECT S.audit_id, S.group_id AS audit_group, S.audit_type, LS.date_created_utc, CB.created_by, LS.date_scheduled_utc, LS.date_started_utc, LS.date_completed_utc, LS.row_count, COALESCE(FS.flags_total, 0) AS flags_total, COALESCE(FS.flags_resolved, 0) AS flags_resolved, FS.date_last_resolved_utc, COALESCE(FS.rows_flagged, 0) AS rows_flagged, COALESCE(FS.rows_flagged, 0) - COALESCE(FS.rows_pending, 0) AS rows_resolved FROM AUDITS_IN_SCOPE S LEFT JOIN LATERAL ( SELECT MIN( CASE WHEN AL.action LIKE '{CREATED_PREFIX}%%' THEN AL.date_created_utc END ) AS date_created_utc, MIN( CASE WHEN AL.action = '{DBAuditLogActions.CREATED_SCHEDULED}' THEN AL.date_created_utc END ) AS date_scheduled_utc, MIN( CASE WHEN AL.action = '{DBAuditLogActions.STARTED}' THEN AL.date_created_utc END ) AS date_started_utc, MAX( CASE WHEN AL.action = '{DBAuditLogActions.COMPLETED}' THEN AL.date_created_utc END ) AS date_completed_utc, MAX( CASE WHEN AL.action = '{DBAuditLogActions.COMPLETED_FETCH}' THEN AL.row_count END ) AS row_count FROM {Tables.AUDITS_LOG} AL WHERE AL.audit = S.audit_id AND ( AL.action LIKE '{CREATED_PREFIX}%%' OR AL.action IN ( '{DBAuditLogActions.STARTED}', '{DBAuditLogActions.COMPLETED}', '{DBAuditLogActions.COMPLETED_FETCH}' ) ) ) LS ON TRUE LEFT JOIN LATERAL ( SELECT AL.user AS created_by FROM {Tables.AUDITS_LOG} AL WHERE AL.audit = S.audit_id AND AL.action LIKE '{CREATED_PREFIX}%%' ORDER BY AL.date_created_utc, AL.id LIMIT 1 ) CB ON TRUE LEFT JOIN LATERAL ( SELECT COUNT(*) AS flags_total, SUM(AF.resolution IS NOT NULL) AS flags_resolved, MAX( CASE WHEN AF.resolution IS NOT NULL THEN AF.date_resolved_utc END ) AS date_last_resolved_utc, COUNT(DISTINCT AF.row_idx) AS rows_flagged, COUNT(DISTINCT CASE WHEN AF.resolution IS NULL THEN AF.row_idx END) AS rows_pending FROM {Tables.AUDITS_FLAGS} AF WHERE AF.audit = S.audit_id ) FS ON TRUE ), FINAL_TABLE AS ( SELECT FT.*, JSON_OBJECT( '{Auth0Users.ID}', U.id, '{Auth0Users.NICKNAME}', U.nickname ) AS _created_by, FT.flags_total - FT.flags_resolved AS flags_pending, IF( FT.flags_total IS NULL OR FT.flags_total = 0, NULL, ROUND(COALESCE(FT.flags_resolved, 0) / FT.flags_total, 4) ) AS flags_resolved_pct, FT.rows_flagged - FT.rows_resolved AS rows_pending, IFNULL(date_scheduled_utc, date_created_utc) AS date_scheduled_or_created_utc, IF( date_completed_utc IS NULL OR row_count IS NULL OR row_count = 0, NULL, CAST( ROUND( ((row_count - (rows_flagged - rows_resolved)) / row_count) * 100, 0 ) AS UNSIGNED ) ) AS score FROM ( SELECT AF.audit_group AS `group`, AG.label_id, AG.label_name, AG.error, JSON_ARRAYAGG(AF.audit_type) AS types, MIN(AF.date_created_utc) AS date_created_utc, ANY_VALUE(AF.created_by) AS created_by, MIN(AF.date_scheduled_utc) AS date_scheduled_utc, MIN(AF.date_started_utc) AS date_started_utc, MAX(AF.date_last_resolved_utc) AS date_last_resolved_utc, CASE WHEN SUM(AF.date_completed_utc IS NULL) > 0 THEN NULL ELSE MAX(AF.date_completed_utc) END AS date_completed_utc, CAST(SUM(AF.row_count) AS UNSIGNED) AS row_count, CAST(SUM(AF.flags_total) AS UNSIGNED) AS flags_total, CAST(SUM(AF.flags_resolved) AS UNSIGNED) AS flags_resolved, CAST(SUM(AF.rows_flagged) AS UNSIGNED) AS rows_flagged, CAST(SUM(AF.rows_resolved) AS UNSIGNED) AS rows_resolved FROM AUDIT_FACTS AF JOIN {Tables.AUDITS_GROUPS} AG ON AG.id = AF.audit_group GROUP BY AF.audit_group, AG.label_id, AG.label_name, AG.error ) FT LEFT JOIN {Tables.USERS} U ON U.id = FT.created_by ) """ whitelisted_columns = { DBColumns.GROUP, DBColumns.LABEL_ID, DBColumns.LABEL_NAME, DBColumns.ERROR, DBColumns.TYPES, DBColumns.DATE_CREATED, DBColumns.CREATED_BY, DBColumns.DATE_SCHEDULED, DBColumns.DATE_STARTED, DBColumns.DATE_COMPLETED, DBColumns.DATE_LAST_RESOLVED, DBColumns.ROW_COUNT, DBColumns.FLAGS_TOTAL, DBColumns.FLAGS_RESOLVED, DBColumns.FLAGS_PENDING, DBColumns.FLAGS_RESOLVED_PCT, DBColumns.ROWS_FLAGGED, DBColumns.ROWS_RESOLVED, DBColumns.ROWS_PENDING, DBColumns.SCORE, DBColumns.DATE_SCHEDULED_OR_CREATED, } placeholders_values = [] clauses = [] if created_by: placeholders_values.append(created_by) clauses.append(f"{DBColumns.CREATED_BY} = %s") if text: placeholders_values.extend([f"%{text.lower()}%"] * 3) clauses.append(f""" (LOWER({DBColumns.LABEL_NAME}) LIKE %s OR LOWER({DBColumns.LABEL_ID}) LIKE %s OR CAST(`{DBColumns.GROUP}` AS CHAR) LIKE %s) """) if status: statuses = { AuditStatus.IN_PROGRESS: f"{DBColumns.DATE_STARTED} IS NOT NULL AND " f"{DBColumns.DATE_COMPLETED} IS NULL AND " f"{DBColumns.ERROR} IS NULL", AuditStatus.COMPLETED: f"{DBColumns.DATE_COMPLETED} IS NOT NULL " f"AND {DBColumns.ERROR} IS NULL", AuditStatus.FAILED: f"{DBColumns.ERROR} IS NOT NULL", } if status_clause := statuses.get(status.upper()): clauses.append(status_clause) else: logger.warning(f"Invalid status: {status}") return [], 0 if audit_types: audit_types_clauses = [] for audit_type in audit_types: placeholders_values.append(f'"{audit_type}"') audit_types_clauses.append(f"JSON_CONTAINS({DBColumns.TYPES}, %s)") clauses.append(f"({' AND '.join(audit_types_clauses)})") if scheduled_date_start: scheduled_date_start = _stringify_date(scheduled_date_start) placeholders_values.append(scheduled_date_start) clauses.append(f"CAST({DBColumns.DATE_SCHEDULED_OR_CREATED} AS DATE) >= %s") if scheduled_date_end: scheduled_date_end = _stringify_date(scheduled_date_end) placeholders_values.append(scheduled_date_end) clauses.append(f"CAST({DBColumns.DATE_SCHEDULED_OR_CREATED} AS DATE) <= %s") if score_min is not None: placeholders_values.append(score_min) clauses.append(f"{DBColumns.SCORE} >= %s") if score_max is not None: placeholders_values.append(score_max) clauses.append(f"({DBColumns.SCORE} IS NULL OR {DBColumns.SCORE} <= %s)") # Combine all WHERE clauses into a single WHERE statement (if any) if clauses: clauses_concatenated = " AND ".join(clauses) clauses = ["WHERE " + clauses_concatenated] # Can't use parametrization for ORDER BY, LIMIT and OFFSET, so to prevent # SQL injection, only allow whitelisted columns to be used for ordering # and force the order direction to be either ASC or DESC, as well as cast # LIMIT and OFFSET to integers. if order_by not in whitelisted_columns: raise ValueError(f"Invalid order_by column: {order_by}") clauses.append(f"ORDER BY `{order_by}` {'DESC' if order_desc else 'ASC'}") clauses.append(f"LIMIT {int(offset)}, {int(limit)}") clauses = " ".join(clauses) query_data = query + """ SELECT * FROM FINAL_TABLE """ + clauses query_count = query + """ SELECT COUNT(*) AS total_rows FROM FINAL_TABLE """ + clauses.split("ORDER BY", maxsplit=1)[0] # Ignore ORDER BY, LIMIT/OFFSET placeholders_values = tuple(placeholders_values) db = self.client.db data, _row_count = await asyncio.gather( db.execute_query_fetchall(query_data, placeholders_values), db.execute_query_fetchone(query_count, placeholders_values), ) types = DB.TYPES for row in data: # Convert stringified JSON list to actual list row[types] = json.loads(row[types]) # Convert all Decimals (percentages) to floats for preventing exceptions in # JSON serialization for key, value in row.items(): if isinstance(value, Decimal): row[key] = float(value) row_count = _row_count["total_rows"] if _row_count else 0 return data, row_count async def list_children(self, audit_group_id: AuditGroupID) -> Rows: """Get audits belonging to a given audit group. Args: audit_group_id: Audit group ID. """ query = f""" SELECT id, type FROM {Tables.AUDITS} WHERE `group` = %s """ data = await self.client.db.execute_query_fetchall(query, (audit_group_id,)) return data or [] async def get_flags(self, audit_group_id: AuditGroupID) -> Rows: """Get audit flags details for audits in a given audit group.""" query = f""" SELECT A.type, {self.client._audit_flag_columns}, CREATED_BY.nickname AS created_by_nickname, RESOLVED_BY.nickname AS resolved_by_nickname FROM {Tables.AUDITS} A LEFT JOIN {Tables.AUDITS_FLAGS} AF ON AF.audit = A.id 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 A.group = %s """ flags = deque( await self.client.db.execute_query_fetchall(query, (audit_group_id,)) ) # Filter out rows with no flags. This can happen when LEFT JOIN finds no # matching rows in the joined tables (thus an incorrect flag row is created). valid_flags = [] while flags: flag = flags.popleft() if flag["id"]: valid_flags.append(flag) return valid_flags async def get_logs(self, audit_group_id: AuditGroupID) -> Rows: """Get audit log entries for audits in a given audit group.""" query = f""" SELECT AG.id AS `group`, A.id AS audit, AL.id, AL.action, AL.date_created_utc, AL.date_scheduled_utc, AL.details, AL.row_count, AL.user, U.nickname AS user_nickname FROM {Tables.AUDITS_GROUPS} AG LEFT JOIN {Tables.AUDITS} A ON AG.id = A.group LEFT JOIN {Tables.AUDITS_LOG} AL ON A.id = AL.audit LEFT JOIN {Tables.USERS} U ON AL.user = U.id WHERE AG.id = %s """ data = await self.client.db.execute_query_fetchall(query, (audit_group_id,)) return data async def mark_run_as_failed( self, audit_group_id: AuditGroupID, error: str ) -> None: """Mark audit group as failed to run (e.g. due to errors during run execution). Args: audit_group_id: Audit group ID. error: Error message to log. """ query = f""" UPDATE {Tables.AUDITS_GROUPS} SET error = %s WHERE id = %s """ await self.client.db.execute_query_nofetch( query, ( error, audit_group_id, ), ) def _stringify_date(asof: str | DateType | None) -> str | None: """If arg is a date or datetime object, convert it to a string in the format 'YYYY-MM-DD'.""" if isinstance(asof, DateType): return asof.strftime("%Y-%m-%d") return asof