#!/usr/bin/env python3
"""Script to validate Dagster definitions and dbt compilation, generating a report."""

import json
import os
import sys
from pathlib import Path
from typing import Dict, Any
os.environ["SNOWFLAKE_PRIVATE_KEY"] = ""
os.environ["DATADOG_API_KEY"] = ""
os.environ["DATADOG_APP_KEY"] = ""
from src.definitions import defs



def count_dagster_definitions() -> Dict[str, Any]:
    """Count and categorize Dagster definitions."""
    if defs is None:
        return {
            "error": "Failed to load Dagster definitions",
            "assets": 0,
            "jobs": 0,
            "schedules": 0,
            "sensors": 0,
            "asset_checks": 0,
            "resources": 0,
        }
    
    stats = {
        "assets": len(defs.assets),
        "jobs": len(defs.jobs),
        "schedules": len(defs.schedules),
        "sensors": len(defs.sensors),
        "asset_checks": len(defs.asset_checks),
        "resources": len(defs.resources),
    }
    
    # Count assets by group
    asset_groups = {}
    for asset in defs.assets:
        group = asset.group_name if hasattr(asset, 'group_name') else "ungrouped"
        asset_groups[group] = asset_groups.get(group, 0) + 1
    
    stats["asset_groups"] = asset_groups
    
    # Count jobs by name
    job_names = [job.name for job in defs.jobs]
    stats["job_names"] = sorted(job_names)
    
    # Count schedules by name
    schedule_names = [schedule.name for schedule in defs.schedules]
    stats["schedule_names"] = sorted(schedule_names)
    
    return stats


def get_dbt_statistics(dbt_project_dir: Path) -> Dict[str, Any]:
    """Get dbt compilation statistics from manifest."""
    stats = {
        "models": 0,
        "sources": 0,
        "tests": 0,
        "seeds": 0,
        "snapshots": 0,
        "exposures": 0,
        "macros": 0,
        "models_by_type": {},
        "models_by_schema": {},
        "compilation_success": False,
        "errors": [],
    }
    
    manifest_path = dbt_project_dir / "target" / "manifest.json"
    
    if not manifest_path.exists():
        stats["errors"].append("Manifest file not found. Ensure dbt compile has been run.")
        return stats
    
    try:
        with open(manifest_path, "r") as f:
            manifest = json.load(f)
        
        stats["compilation_success"] = True
        
        # Count models
        models = manifest.get("nodes", {})
        model_nodes = {k: v for k, v in models.items() if v.get("resource_type") == "model"}
        stats["models"] = len(model_nodes)
        
        # Count by materialization type
        for node in model_nodes.values():
            materialized = node.get("config", {}).get("materialized", "view")
            stats["models_by_type"][materialized] = stats["models_by_type"].get(materialized, 0) + 1
            
            schema = node.get("schema", "unknown")
            stats["models_by_schema"][schema] = stats["models_by_schema"].get(schema, 0) + 1
        
        # Count sources
        sources = manifest.get("sources", {})
        stats["sources"] = len(sources)
        
        # Count tests
        test_nodes = {k: v for k, v in models.items() if v.get("resource_type") == "test"}
        stats["tests"] = len(test_nodes)
        
        # Count seeds
        seed_nodes = {k: v for k, v in models.items() if v.get("resource_type") == "seed"}
        stats["seeds"] = len(seed_nodes)
        
        # Count snapshots
        snapshot_nodes = {k: v for k, v in models.items() if v.get("resource_type") == "snapshot"}
        stats["snapshots"] = len(snapshot_nodes)
        
        # Count exposures
        exposures = manifest.get("exposures", {})
        stats["exposures"] = len(exposures)
        
        # Count macros
        macros = manifest.get("macros", {})
        stats["macros"] = len(macros)
        
    except json.JSONDecodeError as e:
        stats["errors"].append(f"Failed to parse manifest.json: {e}")
    except Exception as e:
        stats["errors"].append(f"Error reading manifest: {e}")
    
    return stats


def generate_markdown_report(dagster_stats: Dict[str, Any], dbt_stats: Dict[str, Any]) -> str:
    """Generate a markdown report from statistics."""
    report = []
    report.append("## 📊 Dagster & dbt Validation Report\n")
    
    # Dagster section
    report.append("### 🎯 Dagster Definitions\n")
    if dagster_stats.get("error"):
        report.append(f"❌ **Error**: {dagster_stats['error']}\n")
    else:
        report.append("✅ **Status**: Definitions loaded successfully\n")
    report.append("| Component | Count |")
    report.append("|-----------|-------|")
    report.append(f"| Assets | {dagster_stats.get('assets', 0)} |")
    report.append(f"| Jobs | {dagster_stats.get('jobs', 0)} |")
    report.append(f"| Schedules | {dagster_stats.get('schedules', 0)} |")
    report.append(f"| Sensors | {dagster_stats.get('sensors', 0)} |")
    report.append(f"| Asset Checks | {dagster_stats.get('asset_checks', 0)} |")
    report.append(f"| Resources | {dagster_stats.get('resources', 0)} |")
    report.append("")
    
    # Asset groups
    if dagster_stats.get("asset_groups") and not dagster_stats.get("error"):
        report.append("#### Assets by Group\n")
        report.append("| Group | Count |")
        report.append("|-------|-------|")
        for group, count in sorted(dagster_stats["asset_groups"].items()):
            report.append(f"| \`{group}\` | {count} |")
        report.append("")
    
    # Jobs
    if dagster_stats.get("job_names") and not dagster_stats.get("error"):
        report.append("#### Jobs\n")
        report.append(f"Jobs found: {', '.join([f'{name}' for name in dagster_stats['job_names'][:10]])}")
        if len(dagster_stats["job_names"]) > 10:
            report.append(f"*... and {len(dagster_stats['job_names']) - 10} more*")
        report.append("")
    
    # Schedules
    if dagster_stats.get("schedule_names") and not dagster_stats.get("error"):
        report.append("#### Schedules\n")
        report.append(f"Schedules found: {', '.join([f'{name}' for name in dagster_stats['schedule_names'][:10]])}")
        if len(dagster_stats["schedule_names"]) > 10:
            report.append(f"*... and {len(dagster_stats['schedule_names']) - 10} more*")
        report.append("")
    
    # dbt section
    report.append("### 📦 dbt Compilation Report\n")
    
    if dbt_stats.get("compilation_success"):
        report.append("✅ **Compilation Status**: Success\n")
        report.append("| Component | Count |")
        report.append("|-----------|-------|")
        report.append(f"| Models | {dbt_stats['models']} |")
        report.append(f"| Sources | {dbt_stats['sources']} |")
        report.append(f"| Tests | {dbt_stats['tests']} |")
        report.append(f"| Seeds | {dbt_stats['seeds']} |")
        report.append(f"| Snapshots | {dbt_stats['snapshots']} |")
        report.append(f"| Exposures | {dbt_stats['exposures']} |")
        report.append(f"| Macros | {dbt_stats['macros']} |")
        report.append("")
        
        # Models by materialization
        if dbt_stats.get("models_by_type"):
            report.append("#### Models by Materialization\n")
            report.append("| Type | Count |")
            report.append("|------|-------|")
            for materialized, count in sorted(dbt_stats["models_by_type"].items()):
                report.append(f"| \`{materialized}\` | {count} |")
            report.append("")
        
        # Models by schema
        if dbt_stats.get("models_by_schema"):
            report.append("#### Models by Schema\n")
            report.append("| Schema | Count |")
            report.append("|--------|-------|")
            for schema, count in sorted(dbt_stats["models_by_schema"].items(), key=lambda x: -x[1])[:10]:
                report.append(f"| \`{schema}\` | {count} |")
            if len(dbt_stats["models_by_schema"]) > 10:
                report.append(f"*... and {len(dbt_stats['models_by_schema']) - 10} more schemas*")
            report.append("")
    else:
        report.append("❌ **Compilation Status**: Failed or incomplete\n")
        if dbt_stats.get("errors"):
            report.append("#### Errors\n")
            for error in dbt_stats["errors"]:
                report.append(f"- {error}")
            report.append("")
    
    return "\n".join(report)


def main():
    """Main function to generate validation report."""
    # Determine project root
    script_dir = Path(__file__).parent
    project_root = script_dir.parent
    
    dbt_project_dir = Path(os.getenv("DBT_PROJECT_DIR", "src/assets/dbt/analytics"))
    if not dbt_project_dir.is_absolute():
        # Make it relative to project root
        dbt_project_dir = project_root / dbt_project_dir
    
    # Get output file path from environment or use default
    output_file = os.getenv("VALIDATION_OUTPUT_FILE", "validation_output.json")
    output_path = Path(output_file)
    
    # If relative path, make it relative to project root
    if not output_path.is_absolute():
        output_path = project_root / output_path
    
    print("Validating Dagster definitions...", file=sys.stderr)
    dagster_stats = count_dagster_definitions()
    
    print("Analyzing dbt compilation...", file=sys.stderr)
    dbt_stats = get_dbt_statistics(dbt_project_dir)
    
    report = generate_markdown_report(dagster_stats, dbt_stats)
    
    # Output JSON for programmatic use
    output = {
        "dagster": dagster_stats,
        "dbt": dbt_stats,
        "report": report,
    }
    
    print(json.dumps(output, indent=2))
    
    # Write JSON to file instead of stdout to avoid contamination
    with open(output_path, "w") as f:
        json.dump(output, f, indent=2)


if __name__ == "__main__":
    main()

