#!/usr/bin/env python3
"""
Step 1: Data Acquisition
Fetches adverse event data from FDA FAERS database and safety data from Open Targets.

Target drugs: semaglutide, tirzepatide, liraglutide
Target reactions: thyroid cancer, pancreatitis, gastroparesis
Target genes: GLP1R, GIPR
"""

import json
import time
import urllib.request
import urllib.parse
import urllib.error
from pathlib import Path
from datetime import datetime
from collections import defaultdict

# Configuration
BASE_DIR = Path("/app/sandbox/session_20260203_091322_5981a70f834a")
DATA_DIR = BASE_DIR / "workflow" / "data"
RESULTS_DIR = BASE_DIR / "results"

# OpenFDA API endpoint
OPENFDA_ENDPOINT = "https://api.fda.gov/drug/event.json"

# Drug names to query
DRUGS = ["semaglutide", "tirzepatide", "liraglutide"]

# Adverse events of interest (with variations)
ADVERSE_EVENTS = {
    "thyroid_cancer": [
        "thyroid cancer",
        "thyroid neoplasm",
        "thyroid carcinoma",
        "medullary thyroid carcinoma",
        "thyroid c-cell tumour",
        "thyroid neoplasm malignant",
    ],
    "pancreatitis": [
        "pancreatitis",
        "pancreatitis acute",
        "pancreatitis chronic",
        "acute pancreatitis",
        "chronic pancreatitis",
    ],
    "gastroparesis": [
        "gastroparesis",
        "delayed gastric emptying",
        "stomach paralysis",
    ],
}

# Open Targets GraphQL endpoint and target IDs
OT_GRAPHQL_ENDPOINT = "https://api.platform.opentargets.org/api/v4/graphql"
OT_TARGETS = {
    "GLP1R": "ENSG00000112164",
    "GIPR": "ENSG00000010310",
}


def make_request(url, data=None, headers=None, max_retries=3):
    """Make HTTP request with retry logic."""
    if headers is None:
        headers = {"Content-Type": "application/json"}

    for attempt in range(max_retries):
        try:
            if data:
                request = urllib.request.Request(
                    url,
                    data=json.dumps(data).encode("utf-8"),
                    headers=headers
                )
            else:
                request = urllib.request.Request(url, headers=headers)

            with urllib.request.urlopen(request, timeout=60) as response:
                return json.loads(response.read().decode("utf-8"))
        except urllib.error.HTTPError as e:
            if e.code == 429:  # Rate limit
                wait_time = (attempt + 1) * 10
                print(f"  Rate limited, waiting {wait_time}s...")
                time.sleep(wait_time)
            elif e.code == 404:
                return None  # No data found
            else:
                print(f"  HTTP Error {e.code}: {e.reason}")
                if attempt < max_retries - 1:
                    time.sleep(2)
        except Exception as e:
            print(f"  Error: {e}")
            if attempt < max_retries - 1:
                time.sleep(2)

    return None


def query_openfda_counts(drug_name, reaction_terms):
    """Query OpenFDA for adverse event counts."""
    # Build search query
    drug_search = f'patient.drug.medicinalproduct:"{drug_name}"'

    # Build reaction search with OR
    reaction_conditions = []
    for term in reaction_terms:
        reaction_conditions.append(f'patient.reaction.reactionmeddrapt:"{term}"')
    reaction_search = " OR ".join(reaction_conditions)

    search_query = f"({drug_search}) AND ({reaction_search})"

    # Encode URL
    params = urllib.parse.urlencode({
        "search": search_query,
        "count": "patient.reaction.reactionmeddrapt.exact"
    })
    url = f"{OPENFDA_ENDPOINT}?{params}"

    result = make_request(url)

    if result and "results" in result:
        return result["results"]
    return []


def query_openfda_records(drug_name, limit=100, skip=0):
    """Query OpenFDA for individual adverse event records."""
    drug_search = f'patient.drug.medicinalproduct:"{drug_name}"'

    # Search for any of our reactions of interest
    all_reactions = []
    for terms in ADVERSE_EVENTS.values():
        all_reactions.extend(terms)

    reaction_conditions = []
    for term in all_reactions:
        reaction_conditions.append(f'patient.reaction.reactionmeddrapt:"{term}"')
    reaction_search = " OR ".join(reaction_conditions)

    search_query = f"({drug_search}) AND ({reaction_search})"

    params = urllib.parse.urlencode({
        "search": search_query,
        "limit": limit,
        "skip": skip
    })
    url = f"{OPENFDA_ENDPOINT}?{params}"

    return make_request(url)


def fetch_faers_data():
    """Fetch FAERS data for all drugs and adverse events."""
    print("=" * 60)
    print("FETCHING FDA FAERS DATA")
    print("=" * 60)

    all_records = []
    summary = defaultdict(lambda: defaultdict(int))

    for drug in DRUGS:
        print(f"\nProcessing {drug}...")

        # First, get counts for each adverse event category
        for category, terms in ADVERSE_EVENTS.items():
            print(f"  Querying {category} counts...")
            counts = query_openfda_counts(drug, terms)

            if counts:
                # Sum up counts for all matching terms
                category_count = 0
                for item in counts:
                    term_lower = item["term"].lower()
                    for search_term in terms:
                        if search_term.lower() in term_lower or term_lower in search_term.lower():
                            category_count += item["count"]
                            break
                summary[drug][category] = category_count
                print(f"    Found {category_count} reports for {category}")
            else:
                print(f"    No data found for {category}")

            time.sleep(0.5)  # Rate limiting

        # Fetch actual records (limited sample)
        print(f"  Fetching detailed records for {drug}...")
        total_fetched = 0
        skip = 0

        while total_fetched < 500:  # Limit to 500 records per drug
            result = query_openfda_records(drug, limit=100, skip=skip)

            if result is None or "results" not in result:
                break

            records = result["results"]
            if not records:
                break

            # Process each record
            for record in records:
                processed = {
                    "drug": drug,
                    "receivedate": record.get("receivedate"),
                    "serious": record.get("serious"),
                    "seriousnessdeath": record.get("seriousnessdeath"),
                    "seriousnesshospitalization": record.get("seriousnesshospitalization"),
                    "patient_drugs": [],
                    "patient_reactions": [],
                    "matched_categories": [],
                }

                # Extract patient drugs
                if "patient" in record and "drug" in record["patient"]:
                    for d in record["patient"]["drug"]:
                        processed["patient_drugs"].append({
                            "medicinalproduct": d.get("medicinalproduct"),
                            "drugindication": d.get("drugindication"),
                            "drugcharacterization": d.get("drugcharacterization"),
                        })

                # Extract reactions and categorize
                if "patient" in record and "reaction" in record["patient"]:
                    for r in record["patient"]["reaction"]:
                        reaction = r.get("reactionmeddrapt", "")
                        processed["patient_reactions"].append({
                            "reactionmeddrapt": reaction,
                            "reactionoutcome": r.get("reactionoutcome"),
                        })

                        # Categorize reaction
                        for category, terms in ADVERSE_EVENTS.items():
                            for term in terms:
                                if term.lower() in reaction.lower() or reaction.lower() in term.lower():
                                    if category not in processed["matched_categories"]:
                                        processed["matched_categories"].append(category)
                                    break

                all_records.append(processed)
                total_fetched += 1

            skip += 100
            print(f"    Fetched {total_fetched} records so far...")
            time.sleep(0.5)  # Rate limiting

        print(f"  Total records fetched for {drug}: {total_fetched}")
        time.sleep(1)  # Rate limiting between drugs

    return all_records, dict(summary)


def query_open_targets_safety(target_id, target_name):
    """Query Open Targets for target safety information."""
    query = """
    query TargetSafety($ensemblId: String!) {
        target(ensemblId: $ensemblId) {
            id
            approvedSymbol
            approvedName
            safetyLiabilities {
                event
                eventId
                datasource
                url
                literature
                effects {
                    direction
                    dosing
                }
                biosamples {
                    cellFormat
                    cellLabel
                    tissue
                    tissueLabel
                }
                studies {
                    name
                    type
                    description
                }
            }
            geneticConstraint {
                constraintType
                score
                exp
                obs
                oe
                oeUpper
                upperBin
                upperRank
            }
        }
    }
    """

    variables = {"ensemblId": target_id}

    data = {
        "query": query,
        "variables": variables
    }

    result = make_request(OT_GRAPHQL_ENDPOINT, data=data)

    if result and "data" in result and result["data"]["target"]:
        return result["data"]["target"]
    return None


def fetch_open_targets_data():
    """Fetch Open Targets safety data for GLP1R and GIPR."""
    print("\n" + "=" * 60)
    print("FETCHING OPEN TARGETS SAFETY DATA")
    print("=" * 60)

    results = {}

    for gene_name, ensembl_id in OT_TARGETS.items():
        print(f"\nQuerying {gene_name} ({ensembl_id})...")

        target_data = query_open_targets_safety(ensembl_id, gene_name)

        if target_data:
            results[gene_name] = {
                "ensembl_id": target_data["id"],
                "symbol": target_data["approvedSymbol"],
                "name": target_data["approvedName"],
                "safety_liabilities": target_data.get("safetyLiabilities", []),
                "genetic_constraint": target_data.get("geneticConstraint", []),
            }

            n_safety = len(results[gene_name]["safety_liabilities"])
            n_constraint = len(results[gene_name]["genetic_constraint"])
            print(f"  Found {n_safety} safety liabilities")
            print(f"  Found {n_constraint} constraint scores")
        else:
            print(f"  No data found for {gene_name}")
            results[gene_name] = {"error": "No data found"}

        time.sleep(0.5)  # Rate limiting

    return results


def generate_summary(faers_records, faers_summary, ot_data):
    """Generate acquisition summary."""
    summary_lines = []
    summary_lines.append("=" * 70)
    summary_lines.append("DATA ACQUISITION SUMMARY")
    summary_lines.append(f"Generated: {datetime.now().isoformat()}")
    summary_lines.append("=" * 70)

    # FAERS Summary
    summary_lines.append("\n" + "-" * 50)
    summary_lines.append("FDA FAERS DATA")
    summary_lines.append("-" * 50)

    total_records = len(faers_records)
    summary_lines.append(f"\nTotal records retrieved: {total_records}")

    # Count by drug
    drug_counts = defaultdict(int)
    for record in faers_records:
        drug_counts[record["drug"]] += 1

    summary_lines.append("\nRecords per drug:")
    for drug in DRUGS:
        summary_lines.append(f"  - {drug}: {drug_counts[drug]} records")

    # Counts by adverse event category (from API counts)
    summary_lines.append("\nAdverse event counts (from FDA counts endpoint):")
    for drug in DRUGS:
        summary_lines.append(f"\n  {drug.upper()}:")
        for category in ADVERSE_EVENTS.keys():
            count = faers_summary.get(drug, {}).get(category, 0)
            summary_lines.append(f"    - {category}: {count} reports")

    # Category counts in fetched records
    summary_lines.append("\nMatched categories in fetched records:")
    category_counts = defaultdict(int)
    for record in faers_records:
        for cat in record.get("matched_categories", []):
            category_counts[cat] += 1

    for category in ADVERSE_EVENTS.keys():
        summary_lines.append(f"  - {category}: {category_counts[category]} records")

    # Serious events
    serious_count = sum(1 for r in faers_records if r.get("serious") == "1")
    death_count = sum(1 for r in faers_records if r.get("seriousnessdeath") == "1")
    hosp_count = sum(1 for r in faers_records if r.get("seriousnesshospitalization") == "1")

    summary_lines.append("\nSerious events in fetched records:")
    summary_lines.append(f"  - Serious: {serious_count}")
    summary_lines.append(f"  - Deaths: {death_count}")
    summary_lines.append(f"  - Hospitalizations: {hosp_count}")

    # Open Targets Summary
    summary_lines.append("\n" + "-" * 50)
    summary_lines.append("OPEN TARGETS SAFETY DATA")
    summary_lines.append("-" * 50)

    for gene_name, data in ot_data.items():
        summary_lines.append(f"\n{gene_name}:")
        if "error" in data:
            summary_lines.append(f"  Error: {data['error']}")
        else:
            summary_lines.append(f"  Ensembl ID: {data.get('ensembl_id', 'N/A')}")
            summary_lines.append(f"  Full Name: {data.get('name', 'N/A')}")

            safety = data.get("safety_liabilities", [])
            summary_lines.append(f"  Safety Liabilities: {len(safety)}")
            if safety:
                events = [s.get("event", "Unknown") for s in safety[:5]]
                summary_lines.append(f"    Top events: {', '.join(str(e) for e in events)}")

            constraint = data.get("genetic_constraint", [])
            summary_lines.append(f"  Genetic Constraint Scores: {len(constraint)}")
            for c in constraint:
                ct = c.get("constraintType", "Unknown")
                score = c.get("score", "N/A")
                oe = c.get("oe", "N/A")
                summary_lines.append(f"    - {ct}: score={score}, oe={oe}")

    # Success criteria check
    summary_lines.append("\n" + "-" * 50)
    summary_lines.append("SUCCESS CRITERIA CHECK")
    summary_lines.append("-" * 50)

    all_drugs_present = all(drug_counts[d] > 0 for d in DRUGS)
    ot_data_present = all("error" not in ot_data.get(g, {"error": True}) for g in OT_TARGETS.keys())

    summary_lines.append(f"\n[{'PASS' if all_drugs_present else 'FAIL'}] All drugs have FAERS records")
    summary_lines.append(f"[{'PASS' if ot_data_present else 'FAIL'}] Open Targets data retrieved for both targets")
    summary_lines.append(f"[{'PASS' if total_records > 0 else 'FAIL'}] FAERS records fetched: {total_records}")

    return "\n".join(summary_lines)


def main():
    """Main execution function."""
    print(f"Starting data acquisition at {datetime.now().isoformat()}")
    print(f"Working directory: {BASE_DIR}")

    # Ensure directories exist
    DATA_DIR.mkdir(parents=True, exist_ok=True)
    RESULTS_DIR.mkdir(parents=True, exist_ok=True)

    # Fetch FAERS data
    faers_records, faers_summary = fetch_faers_data()

    # Save FAERS data
    faers_output_path = DATA_DIR / "faers_raw.json"
    with open(faers_output_path, "w") as f:
        json.dump({
            "metadata": {
                "fetch_date": datetime.now().isoformat(),
                "drugs": DRUGS,
                "adverse_events": ADVERSE_EVENTS,
                "total_records": len(faers_records),
            },
            "summary_counts": faers_summary,
            "records": faers_records,
        }, f, indent=2)
    print(f"\nSaved FAERS data to: {faers_output_path}")

    # Fetch Open Targets data
    ot_data = fetch_open_targets_data()

    # Save Open Targets data
    ot_output_path = DATA_DIR / "opentargets_safety.json"
    with open(ot_output_path, "w") as f:
        json.dump({
            "metadata": {
                "fetch_date": datetime.now().isoformat(),
                "targets": OT_TARGETS,
            },
            "data": ot_data,
        }, f, indent=2)
    print(f"\nSaved Open Targets data to: {ot_output_path}")

    # Generate and save summary
    summary_text = generate_summary(faers_records, faers_summary, ot_data)

    summary_path = RESULTS_DIR / "01_acquisition_summary.txt"
    with open(summary_path, "w") as f:
        f.write(summary_text)
    print(f"\nSaved summary to: {summary_path}")

    print("\n" + "=" * 60)
    print("DATA ACQUISITION COMPLETE")
    print("=" * 60)
    print(summary_text)

    return faers_records, ot_data


if __name__ == "__main__":
    main()
