#!/usr/bin/env python3
"""List Airflow DAGs that failed within a time window.

Examples:
    ./scripts/list-failed-dags                    # Last 24 hours (default)
    ./scripts/list-failed-dags --since 12h        # Last 12 hours
    ./scripts/list-failed-dags --since 3d         # Last 3 days
    ./scripts/list-failed-dags --since 2025-01-30 # Since specific date
"""

import argparse
import subprocess
import sys
import re
from datetime import datetime, timedelta, timezone


def parse_since(since_str: str) -> datetime:
    """Parse a time specification into a datetime.

    Supports:
        - Relative: 24h, 12h, 3d, 7d
        - Absolute: 2025-01-30, 2025-01-30T12:00:00
    """
    now = datetime.now(timezone.utc)

    # Try relative format (e.g., 24h, 3d)
    match = re.match(r'^(\d+)([hd])$', since_str.lower())
    if match:
        value, unit = int(match.group(1)), match.group(2)
        if unit == 'h':
            return now - timedelta(hours=value)
        elif unit == 'd':
            return now - timedelta(days=value)

    # Try absolute date format
    for fmt in ['%Y-%m-%d', '%Y-%m-%dT%H:%M:%S', '%Y-%m-%dT%H:%M:%S%z']:
        try:
            dt = datetime.strptime(since_str, fmt)
            if dt.tzinfo is None:
                dt = dt.replace(tzinfo=timezone.utc)
            return dt
        except ValueError:
            continue

    raise ValueError(f"Cannot parse time: {since_str}. Use format like '24h', '3d', or '2025-01-30'")


def query_failed_dags(since: datetime, verbose: bool = False) -> list[dict]:
    """Query Cloud Logging for failed DAG runs."""
    timestamp = since.strftime('%Y-%m-%dT%H:%M:%SZ')

    query = (
        'resource.type="k8s_container" AND '
        'resource.labels.namespace_name="telemetry-airflow-prod" AND '
        f'textPayload=~"DagRun Finished.*state=failed" AND '
        f'timestamp>="{timestamp}"'
    )

    cmd = [
        'gcloud', 'logging', 'read', query,
        '--project=moz-fx-dataservices-high-prod',
        '--limit=200',
        '--format=value(timestamp,textPayload)'
    ]

    if verbose:
        print(f"Running: gcloud logging read ...", file=sys.stderr)
        print(f"  Since: {timestamp}", file=sys.stderr)

    result = subprocess.run(cmd, capture_output=True, text=True)

    if result.returncode != 0:
        print(f"Error querying logs: {result.stderr}", file=sys.stderr)
        sys.exit(1)

    # Parse output
    failures = []
    for line in result.stdout.strip().split('\n'):
        if not line:
            continue

        # Extract fields from log line
        dag_match = re.search(r'dag_id=([^,]+)', line)
        run_match = re.search(r'run_id=([^,]+)', line)
        exec_match = re.search(r'execution_date=([^,]+)', line)

        if dag_match:
            failures.append({
                'dag_id': dag_match.group(1),
                'run_id': run_match.group(1) if run_match else 'unknown',
                'execution_date': exec_match.group(1) if exec_match else 'unknown',
            })

    return failures


def main():
    parser = argparse.ArgumentParser(
        description='List Airflow DAGs that failed within a time window.',
        formatter_class=argparse.RawDescriptionHelpFormatter,
        epilog=__doc__
    )
    parser.add_argument(
        '--since', '-s',
        default='24h',
        help='Time window: 24h, 3d, or YYYY-MM-DD (default: 24h)'
    )
    parser.add_argument(
        '--verbose', '-v',
        action='store_true',
        help='Show verbose output'
    )
    parser.add_argument(
        '--all', '-a',
        action='store_true',
        help='Show all failures with details (not just unique DAGs)'
    )

    args = parser.parse_args()

    try:
        since = parse_since(args.since)
    except ValueError as e:
        print(f"Error: {e}", file=sys.stderr)
        sys.exit(1)

    if args.verbose:
        print(f"Querying for failures since {since.isoformat()}", file=sys.stderr)

    failures = query_failed_dags(since, verbose=args.verbose)

    if not failures:
        print("No failed DAGs found in the specified time window.")
        return

    if args.all:
        # Show all failures with details
        print(f"{'DAG ID':<50} {'Execution Date':<30}")
        print("-" * 80)
        for f in failures:
            print(f"{f['dag_id']:<50} {f['execution_date']:<30}")
    else:
        # Show unique DAG IDs
        unique_dags = sorted(set(f['dag_id'] for f in failures))
        print(f"Failed DAGs ({len(unique_dags)} unique, {len(failures)} total failures):\n")
        for dag_id in unique_dags:
            count = sum(1 for f in failures if f['dag_id'] == dag_id)
            if count > 1:
                print(f"  {dag_id} ({count} failures)")
            else:
                print(f"  {dag_id}")


if __name__ == '__main__':
    main()
