import json import os from abc import ABC from datetime import datetime, timedelta, timezone from typing import Any, Dict, List, Optional import uuid # Add to imports at top import pandas as pd from dagster import ( AssetSelection, OpExecutionContext, ScheduleDefinition, asset, define_asset_job, ) from src.utils.database import invoke_postgres_lambda, query_snowflake from .constants import Capability, EntityType, JobGroup, JobName, Status class EnforcementAsset(ABC): """Base class for enforcement assets.""" def __init__( self, name: JobName, description: str, entity_type: EntityType, capability: Capability, status: Status, reason: str, shadow_mode: bool = False, schedule: Optional[str] = None, query_file: Optional[str] = None, max_entities_per_run: int = 100, ): self.name = name self.description = description self.entity_type = entity_type.value # Convert enum to string self.capability = capability.value # Convert enum to string self.status = status.value # Convert enum to string self.reason = reason self.shadow_mode = shadow_mode self.schedule = schedule self.query_file = query_file or f"queries/{name}.sql" # Default to name if not provided self.max_entities_per_run = max_entities_per_run # Create asset and job self.asset = self.create_asset() self.job = self.create_job() if schedule else None self.schedule_def = self.create_schedule() if schedule else None def get_query(self) -> str: """Load and return the SQL query from file. Override this method if you need custom query loading logic. """ query_path = os.path.join(os.path.dirname(__file__), self.query_file) with open(query_path, "r") as f: return f.read() def get_query_params(self, context: OpExecutionContext) -> Dict[str, Any]: """Define parameters for the Snowflake query.""" _now = datetime.now(tz=timezone.utc) params = { "start_date": _now.strftime("%Y-%m-%d"), "end_date": _now.strftime("%Y-%m-%d"), "start_hour": (_now - timedelta(hours=2)).strftime("%H"), "end_hour": (_now - timedelta(hours=1)).strftime("%H"), } context.log.debug(f"Generated query parameters: {params}") return params def create_asset(self): """Create the Dagster asset.""" instance = self @asset( name=self.name.value, description=self.description, group_name=JobGroup.BOT_ENFORCEMENTS.value, deps=[ "bot_status_changes", ], owners=["team:core-pod"], metadata={ "slack": "#tech-anti-bots", }, tags={"team": "core-pod", "monitored": "true", "tech-alerts": "true"}, ) def enforcement_asset(context: OpExecutionContext) -> Optional[pd.DataFrame]: # Get query and params query = instance.get_query() query_params = instance.get_query_params(context) try: df = query_snowflake(query, query_params) # Log DataFrame details context.log.info(f"DataFrame columns: {df.columns}") context.log.info(f"Preview of first 10 rows:\n{df.head(10).to_markdown()}") context.log.info(f"Total rows in df: {len(df)}") # Add early return if no data if len(df) == 0: context.log.info( f"No data returned from Snowflake query for {instance.name.value}" ) return pd.DataFrame() # Return empty DataFrame # Validate required fields if "ENTITY_ID" not in df: raise ValueError( "process_data must return a dict with 'entity_id' keys. " f"Got: {list(df.keys())}" ) entity_ids = df["ENTITY_ID"].tolist() if not entity_ids: context.log.info("No entity_ids to process for enforcement") return None # Changed to return None if len(entity_ids) > instance.max_entities_per_run: # Changed self to instance context.log.warning( f"Circuit breaker: Found {len(entity_ids)} entities for enforcement, limiting to first {instance.max_entities_per_run}" ) entity_ids = entity_ids[: instance.max_entities_per_run] # Write to Postgres in batches of 10K BATCH_SIZE = 10000 for i in range(0, len(entity_ids), BATCH_SIZE): batch = entity_ids[i : i + BATCH_SIZE] instance.upsert_to_postgres(batch, context) context.log.info( f"Processed batch {i // BATCH_SIZE + 1} of {(len(entity_ids) + BATCH_SIZE - 1) // BATCH_SIZE}" ) return df # Return the DataFrame except Exception: context.log.error(f"Query failed with parameters: {query_params}") raise return enforcement_asset def create_job(self): """Create a job for this feature.""" return define_asset_job( name=f"{self.name.value}_job", description=self.description, selection=AssetSelection.assets(self.asset.key), ) def create_schedule(self): """Create a schedule with parameters.""" return ScheduleDefinition( job=self.job, cron_schedule=self.schedule, name=f"{self.name.value}_schedule", ) def upsert_to_postgres(self, entity_ids: List[str], context: OpExecutionContext): """Shared method to handle Postgres upserts.""" current_time = datetime.now(tz=timezone.utc) try: context.log.info( f"Upserting bot features to Postgres with dagster run id: {context.run_id}" ) upsert_stmt = """ INSERT INTO moderation_rules AS m ( id, entity_type, entity_id, capability, status, reason, shadow_mode, created_by, created_at, updated_at ) VALUES ( uuid_generate_v4(), %s, %s, %s, %s, %s, %s, %s, %s, %s ) ON CONFLICT (entity_type, entity_id, capability) DO UPDATE SET capability = EXCLUDED.capability, status = EXCLUDED.status, reason = EXCLUDED.reason, created_by = EXCLUDED.created_by, updated_at = EXCLUDED.updated_at WHERE m.status != 'granted' """ values = [ [ self.entity_type, str(entity_id), self.capability, self.status, self.reason, self.shadow_mode, "dagster", current_time.isoformat(), current_time.isoformat(), ] for entity_id in entity_ids ] context.log.info(f"Upserting {len(values)} rows to Postgres") context.log.info(f"Upsert statement: {upsert_stmt}") context.log.info(f"Values: {values}") response = invoke_postgres_lambda(upsert_stmt, values, is_reader=False) if response.get("statusCode") != 200: raise Exception(f"Error saving bot records to Postgres: {response.get('body')}") context.log.info(f"Save results:\n{json.dumps(response, indent=2)}") except Exception as e: context.log.error(f"Error writing to Postgres: {e!s}") raise