"""
enrich.py — Fetch real plant species from Perenual, enrich agronomic traits
with Claude AI, then populate graft combinations.

Run from /backend:
    python -m app.seeds.enrich            # adds to existing data
    python -m app.seeds.enrich --reset    # clears plants + grafts first
"""
from __future__ import annotations

import argparse
import json
import os
import sys
import time
import uuid
import re
from datetime import datetime, timezone
from typing import Any, Dict, List, Optional

import requests

# Load .env before importing app modules
from dotenv import load_dotenv
load_dotenv()

import anthropic

from app.database import SessionLocal
from app.models.graft_combination import GraftCombination
from app.models.plant import Plant, RootDepthCategory, VarietyType
from app.models.scenario import Scenario, ScenarioResult
from app.models.soil_profile import DrainageRate, SoilProfile, SoilType, WaterRetention

# ---------------------------------------------------------------------------
# Config
# ---------------------------------------------------------------------------

PERENUAL_KEY = os.getenv("PERNUAL_API_KEY")          # note: typo preserved from .env
ANTHROPIC_KEY = os.getenv("ANTHROPIC_API_KEY")
PERENUAL_BASE = "https://perenual.com/api"
MODEL = "claude-sonnet-4-6"
NOW = datetime.now(timezone.utc)

_ai = anthropic.Anthropic(api_key=ANTHROPIC_KEY) if ANTHROPIC_KEY else None

# Cache file — stores raw Perenual results so the API is only hit once
CACHE_FILE = os.path.join(os.path.dirname(__file__), "perenual_cache.json")

# Search terms — Perenual returns edible species matching each query
SEARCH_TERMS = [
    # Solanaceae
    "tomato", "cherry tomato", "pepper", "chili pepper",
    "eggplant", "aubergine", "tomatillo",
    # Cucurbitaceae
    "cucumber", "melon", "watermelon", "squash",
    "zucchini", "pumpkin", "bitter melon", "luffa",
    # Fruit trees
    "apple", "pear", "citrus", "lemon", "orange",
    "peach", "plum", "cherry", "apricot", "mango",
    "avocado", "fig", "pomegranate", "guava",
    # Grapes & berries
    "grape", "strawberry", "blueberry",
    # Other vegetables
    "lettuce", "spinach", "kale", "broccoli",
    "carrot", "potato", "sweet potato", "onion",
    "garlic", "artichoke", "asparagus",
]
MAX_PER_TERM = 3   # Perenual results per search term


# ---------------------------------------------------------------------------
# Perenual helpers
# ---------------------------------------------------------------------------

def _perenual_get(path: str, params: Dict) -> Dict:
    params["key"] = PERENUAL_KEY
    resp = requests.get(f"{PERENUAL_BASE}{path}", params=params, timeout=12)
    resp.raise_for_status()
    return resp.json()


def _species_prompt(groups: str, start_id: int) -> str:
    return f"""Generate exactly 30 edible plant species for: {groups}

Return a JSON array of exactly 30 objects. Each object:
{{
  "id": <unique int starting from {start_id}>,
  "common_name": "<common name>",
  "scientific_name": ["<scientific name>"],
  "cycle": "<Annual|Biennial|Perennial>",
  "watering": "<Frequent|Average|Minimum>",
  "drought_tolerant": <true|false>,
  "salt_tolerant": <true|false>
}}"""


def generate_species_with_claude() -> List[Dict]:
    """Two batches of 30 to stay within token limits."""
    batch_a = _ask_claude(_species_prompt(
        "Solanaceae (tomato varieties, cherry tomato, beef tomato, roma tomato, pepper, chili, bell pepper, eggplant, tomatillo), "
        "Cucurbitaceae (cucumber, melon, watermelon, squash, zucchini, pumpkin, bitter melon, luffa), "
        "Citrus (lemon, orange, lime, grapefruit, mandarin, pomelo)",
        start_id=1,
    ), max_tokens=3000)

    batch_b = _ask_claude(_species_prompt(
        "Fruit trees (apple, pear, peach, plum, cherry, apricot, mango, avocado, fig, pomegranate, guava), "
        "Grapes and berries (grape, strawberry, blueberry, raspberry), "
        "Vegetables (lettuce, spinach, kale, broccoli, carrot, potato, sweet potato, onion, garlic, artichoke, asparagus)",
        start_id=31,
    ), max_tokens=3000)

    return batch_a + batch_b


def fetch_species_list(query: str) -> List[Dict]:
    data = _perenual_get("/species-list", {"q": query, "edible": 1, "page": 1})
    return data.get("data", [])


def load_cache() -> List[Dict]:
    if os.path.exists(CACHE_FILE):
        with open(CACHE_FILE) as f:
            return json.load(f)
    return []


def save_cache(data: List[Dict]) -> None:
    with open(CACHE_FILE, "w") as f:
        json.dump(data, f, indent=2)
    print(f"  Cache saved → {CACHE_FILE}")


# ---------------------------------------------------------------------------
# Claude helpers — strip markdown fences, parse JSON
# ---------------------------------------------------------------------------

def _ask_claude(prompt: str, max_tokens: int = 4096) -> List[Dict]:
    if _ai is None:
        raise RuntimeError("ANTHROPIC_API_KEY is not configured")

    msg = _ai.messages.create(
        model=MODEL,
        max_tokens=max_tokens,
        system="You are an expert horticulturist and agronomist. Always respond with a valid JSON array only — no markdown fences, no commentary, no wrapping object.",
        messages=[{"role": "user", "content": prompt}],
    )
    raw = msg.content[0].text.strip()
    if raw.startswith("```"):
        raw = raw.split("```")[1]
        if raw.startswith("json"):
            raw = raw[4:]
        raw = raw.strip()
    result = json.loads(raw)
    # Normalise: if Claude wrapped the array in an object, extract the first list value
    if isinstance(result, dict):
        for v in result.values():
            if isinstance(v, list):
                return v
        return [result]
    return result


# ---------------------------------------------------------------------------
# Enrichment steps
# ---------------------------------------------------------------------------

def enrich_scions(raw_plants: List[Dict]) -> List[Dict]:
    """Send Perenual species data to Claude → structured ACI plant profiles."""
    summary = "\n".join(
        f"- {p.get('common_name','?')} ({', '.join(p.get('scientific_name', ['?']))})"
        f" | cycle: {p.get('cycle','?')}"
        f" | watering: {p.get('watering','?')}"
        f" | drought_tolerant: {p.get('drought_tolerant', False)}"
        f" | salt_tolerant: {p.get('salt_tolerant', False)}"
        for p in raw_plants
    )

    prompt = f"""I have these real plant species from a botanical database:
{summary}

For each plant, generate a complete ACI agronomic profile. Return a JSON array — one object per plant — with exactly these fields:

{{
  "name": "<common cultivar name>",
  "species": "<scientific name>",
  "variety_type": "scion",
  "drought_tolerance": <int 1-10>,
  "salinity_tolerance": <int 1-10>,
  "heat_tolerance": <int 1-10>,
  "cold_tolerance": <int 1-10>,
  "disease_resistance": {{"<disease_name>": <int 1-10>}},
  "optimal_ph_min": <float>,
  "optimal_ph_max": <float>,
  "nutrient_uptake_efficiency": {{"N": <float>, "P": <float>, "K": <float>}},
  "root_depth_category": "<shallow|medium|deep>",
  "growth_cycle_days": <int, days from transplant to first harvest>,
  "yield_potential_kg_per_plant": <float, realistic kg per plant>,
  "description": "<2-sentence agronomic description covering key traits and production notes>"
}}

Rules:
- disease_resistance: include 3-5 diseases relevant to that crop family with integer scores 1-10
- nutrient_uptake_efficiency: N/P/K floats around 1.0 (1.2 = 20% more efficient than baseline)
- Base all values on established agronomic science — be precise, not generic
- yield_potential_kg_per_plant: realistic figure for a single plant in ideal conditions"""

    return _ask_claude(prompt)


def _rootstock_prompt(families: str, examples: str) -> str:
    return f"""Generate exactly 6 well-known grafting rootstock cultivars for: {families}.
Use these real cultivar names: {examples}

Return a JSON array of exactly 6 objects. Each object:
{{
  "name": "<cultivar name>",
  "species": "<scientific name>",
  "variety_type": "rootstock",
  "drought_tolerance": <int 1-10>,
  "salinity_tolerance": <int 1-10>,
  "heat_tolerance": <int 1-10>,
  "cold_tolerance": <int 1-10>,
  "disease_resistance": {{"<disease>": <int 1-10>, "<disease>": <int 1-10>, "<disease>": <int 1-10>}},
  "optimal_ph_min": <float>,
  "optimal_ph_max": <float>,
  "nutrient_uptake_efficiency": {{"N": <float>, "P": <float>, "K": <float>}},
  "root_depth_category": "<shallow|medium|deep>",
  "growth_cycle_days": <int>,
  "yield_potential_kg_per_plant": 0.0,
  "description": "<one sentence on key traits and use case>"
}}"""


def generate_rootstocks() -> List[Dict]:
    """Three batches of 6 to stay within token limits."""
    batch_a = _ask_claude(_rootstock_prompt(
        families="tomato/solanaceae vegetables",
        examples="Beaufort, Maxifort, Arnold, Brigeor, He-Man, Cobalt",
    ), max_tokens=3000)

    batch_b = _ask_claude(_rootstock_prompt(
        families="cucurbit vegetables (cucumber/melon/watermelon)",
        examples="RS841, Ferro, Shintoza, Strong Tosa, Macis, Tetsukabuto",
    ), max_tokens=3000)

    batch_c = _ask_claude(_rootstock_prompt(
        families="apple, pear, stone fruit (Prunus), grape (Vitis)",
        examples="M.9, MM.106, Quince A, Gisela 6, St. Julien A, SO4",
    ), max_tokens=3000)

    return batch_a + batch_b + batch_c


def _soil_prompt(regions: str) -> str:
    return f"""Generate exactly 6 soil profiles for these agricultural regions: {regions}

Return a JSON array of exactly 6 objects:
{{
  "name": "<descriptive name e.g. 'Nile Delta Silt'>",
  "soil_type": "<sandy|loam|clay|silt|peat|chalky>",
  "ph_level": <float>,
  "organic_matter_percent": <float>,
  "drainage_rate": "<poor|moderate|good|excessive>",
  "water_retention": "<low|medium|high>",
  "native_nitrogen": <float ppm>,
  "native_phosphorus": <float ppm>,
  "native_potassium": <float ppm>,
  "salinity_level": <float dS/m>,
  "temperature_range_min": <float Celsius>,
  "temperature_range_max": <float Celsius>,
  "description": "<one sentence on origin and crop suitability>",
  "climate_zone": "<arid|semi-arid|mediterranean|temperate|tropical|subtropical|continental>"
}}

Use realistic agronomic values."""


def generate_soils() -> List[Dict]:
    """Two batches of 6 soils to stay within token limits."""
    batch_a = _ask_claude(_soil_prompt(
        "UAE/Gulf desert (arid saline sandy), Nile Delta (fertile silt), "
        "Mediterranean basin (loam), Northern European lowland (heavy clay), "
        "Tropical volcanic Indonesia (andisol), California Central Valley (fertile loam)"
    ), max_tokens=2000)

    batch_b = _ask_claude(_soil_prompt(
        "Amazon basin (tropical acidic low-P), Sahel West Africa (semi-arid sandy), "
        "Rhine-Po alluvial plain (temperate silt), South African Karoo (semi-arid chalky), "
        "Southeast Asian paddy (waterlogged clay), Andean highland (cold peaty)"
    ), max_tokens=2000)

    return batch_a + batch_b


def _graft_prompt(scion_lines: str, rootstock_lines: str, family_rule: str) -> str:
    return f"""Generate graft combinations. Family rule: {family_rule}

SCIONS:
{scion_lines}

ROOTSTOCKS:
{rootstock_lines}

Return a JSON array of the 5 most relevant/successful combinations. Each object:
{{
  "rootstock_id": "<exact UUID from rootstocks above>",
  "scion_id": "<exact UUID from scions above>",
  "compatibility_score": <float 0.65-1.0>,
  "vigor_boost": <float>,
  "yield_modifier": <float>,
  "stress_tolerance_override": {{"drought_tolerance": <int>, "salinity_tolerance": <int>, "heat_tolerance": <int>}},
  "notes": "<concise one-sentence note>",
  "source_reference": "<short citation>"
}}

Rule: Return exactly 5 objects. Use ONLY the UUIDs listed above. Be concise.
"""


def generate_graft_combinations(
    scions: List[Plant], rootstocks: List[Plant]
) -> List[Dict]:
    """Split by plant family to keep prompts small."""
    # Classify by species keywords
    def is_solanaceous(p: Plant) -> bool:
        s = (p.species + p.name).lower()
        return any(k in s for k in ["solanum", "capsicum", "lycopersicum", "melongena", "physalis", "tomato", "pepper", "eggplant", "aubergine", "tomatillo"])

    def is_cucurbit(p: Plant) -> bool:
        s = (p.species + p.name).lower()
        return any(k in s for k in ["cucumis", "cucurbita", "citrullus", "luffa", "momordica", "cucumber", "melon", "squash", "zucchini", "pumpkin", "gourd"])

    sol_scions = [p for p in scions if is_solanaceous(p)]
    cuc_scions  = [p for p in scions if is_cucurbit(p)]
    sol_roots   = [p for p in rootstocks if is_solanaceous(p)]
    cuc_roots   = [p for p in rootstocks if is_cucurbit(p)]

    results: List[Dict] = []

    if sol_scions and sol_roots:
        # Process in very small batches (3) to ensure Claude doesn't truncate the large response
        for i in range(0, len(sol_scions), 3):
            batch_s = sol_scions[i : i + 3]
            lines_s = "\n".join(f"  id={p.id} | {p.name} ({p.species})" for p in batch_s)
            lines_r = "\n".join(f"  id={p.id} | {p.name} ({p.species})" for p in sol_roots)
            try:
                results += _ask_claude(_graft_prompt(lines_s, lines_r, "solanaceous scions onto solanaceous rootstocks only"), max_tokens=4000)
            except Exception as e:
                print(f"  Warning: solanaceous graft batch {i//3 + 1} failed — {e}")

    if cuc_scions and cuc_roots:
        for i in range(0, len(cuc_scions), 3):
            batch_s = cuc_scions[i : i + 3]
            lines_s = "\n".join(f"  id={p.id} | {p.name} ({p.species})" for p in batch_s)
            lines_r = "\n".join(f"  id={p.id} | {p.name} ({p.species})" for p in cuc_roots)
            try:
                results += _ask_claude(_graft_prompt(lines_s, lines_r, "cucurbit scions onto cucurbit rootstocks only"), max_tokens=4000)
            except Exception as e:
                print(f"  Warning: cucurbit graft batch {i//3 + 1} failed — {e}")

    # Fruit tree / grape grafts
    other_scions = [p for p in scions if not is_solanaceous(p) and not is_cucurbit(p)]
    other_roots  = [p for p in rootstocks if not is_solanaceous(p) and not is_cucurbit(p)]
    if other_scions and other_roots:
        for i in range(0, len(other_scions), 3):
            batch_s = other_scions[i : i + 3]
            lines_s = "\n".join(f"  id={p.id} | {p.name} ({p.species})" for p in batch_s)
            lines_r = "\n".join(f"  id={p.id} | {p.name} ({p.species})" for p in other_roots)
            try:
                results += _ask_claude(_graft_prompt(lines_s, lines_r, "fruit/grape scions onto compatible rootstocks"), max_tokens=4000)
            except Exception as e:
                print(f"  Warning: fruit/grape graft batch {i//3 + 1} failed — {e}")

    return results


# ---------------------------------------------------------------------------
# ORM conversion
# ---------------------------------------------------------------------------

_ST = {
    "sandy": SoilType.sandy, "loam": SoilType.loam, "clay": SoilType.clay,
    "silt": SoilType.silt, "peat": SoilType.peat, "chalky": SoilType.chalky,
}
_DR = {
    "poor": DrainageRate.poor, "moderate": DrainageRate.moderate,
    "good": DrainageRate.good, "excessive": DrainageRate.excessive,
}
_WR = {
    "low": WaterRetention.low, "medium": WaterRetention.medium, "high": WaterRetention.high,
}

_RDC = {
    "shallow": RootDepthCategory.shallow,
    "medium": RootDepthCategory.medium,
    "deep": RootDepthCategory.deep,
}
_VT = {
    "scion": VarietyType.scion,
    "rootstock": VarietyType.rootstock,
    "both": VarietyType.both,
}


def _normalize_text(value: Any) -> str:
    text = str(value or "").strip()
    return re.sub(r"\s+", " ", text)


def _first_nonempty(*values: Any) -> str:
    for value in values:
        text = _normalize_text(value)
        if text:
            return text
    return ""


def _perenual_image_url(raw: Dict[str, Any]) -> Optional[str]:
    image = raw.get("default_image") or {}
    if not isinstance(image, dict):
        return None

    for key in ["regular_url", "medium_url", "small_url", "thumbnail", "original_url"]:
        url = _normalize_text(image.get(key))
        if url:
            return url
    return None


def _watering_to_drought_tolerance(watering: str) -> int:
    watering_value = watering.lower()
    if "minimum" in watering_value or "low" in watering_value:
        return 8
    if "average" in watering_value or "moderate" in watering_value:
        return 5
    if "frequent" in watering_value or "high" in watering_value:
        return 3
    return 5


def _cycle_to_growth_days(cycle: str) -> int:
    cycle_value = cycle.lower()
    if "annual" in cycle_value:
        return 120
    if "biennial" in cycle_value:
        return 365
    if "perennial" in cycle_value:
        return 730
    return 180


def _infer_root_depth(name: str, species: str) -> RootDepthCategory:
    haystack = f"{name} {species}".lower()
    if any(keyword in haystack for keyword in ["tree", "apple", "pear", "citrus", "mango", "avocado", "grape", "fig", "pomegranate"]):
        return RootDepthCategory.deep
    if any(keyword in haystack for keyword in ["lettuce", "spinach", "onion", "garlic", "strawberry"]):
        return RootDepthCategory.shallow
    return RootDepthCategory.medium


def perenual_species_to_plant_dict(raw: Dict[str, Any]) -> Dict[str, Any]:
    common_name = _first_nonempty(raw.get("common_name"), raw.get("other_name", [None])[0] if raw.get("other_name") else None, "Unknown plant")
    scientific_names = raw.get("scientific_name") or []
    species = _first_nonempty(scientific_names[0] if scientific_names else None, common_name)
    cycle = _normalize_text(raw.get("cycle"))
    watering = _normalize_text(raw.get("watering"))

    drought_tolerance = _watering_to_drought_tolerance(watering)
    if raw.get("drought_tolerant") is True:
        drought_tolerance = min(10, drought_tolerance + 2)

    salinity_tolerance = 7 if raw.get("salt_tolerant") is True else 4
    description_bits = [
        f"Imported from Perenual as {common_name} ({species}).",
    ]
    if cycle:
        description_bits.append(f"Lifecycle: {cycle}.")
    if watering:
        description_bits.append(f"Watering guidance: {watering}.")

    return {
        "name": common_name[:255],
        "species": species[:255],
        "variety_type": "scion",
        "drought_tolerance": drought_tolerance,
        "salinity_tolerance": salinity_tolerance,
        "heat_tolerance": 6,
        "cold_tolerance": 5,
        "disease_resistance": {},
        "optimal_ph_min": 6.0,
        "optimal_ph_max": 7.0,
        "nutrient_uptake_efficiency": {"N": 1.0, "P": 1.0, "K": 1.0},
        "root_depth_category": _infer_root_depth(common_name, species).value,
        "growth_cycle_days": _cycle_to_growth_days(cycle),
        "yield_potential_kg_per_plant": 0.0,
        "description": " ".join(description_bits)[:1000],
        "image_url": _perenual_image_url(raw),
    }


def merge_enriched_plant_data(raw_batch: List[Dict[str, Any]], enriched_batch: List[Dict[str, Any]]) -> List[Dict[str, Any]]:
    merged: List[Dict[str, Any]] = []

    for index, raw in enumerate(raw_batch):
        fallback = perenual_species_to_plant_dict(raw)
        ai_data = enriched_batch[index] if index < len(enriched_batch) and isinstance(enriched_batch[index], dict) else {}
        merged_item = {**fallback, **ai_data}
        merged_item["name"] = _first_nonempty(merged_item.get("name"), fallback["name"])[:255]
        merged_item["species"] = _first_nonempty(merged_item.get("species"), fallback["species"])[:255]
        merged_item["image_url"] = fallback.get("image_url")
        merged.append(merged_item)

    return merged


def upsert_plant(db, plant: Plant) -> Plant:
    existing = (
        db.query(Plant)
        .filter(
            Plant.name == plant.name,
            Plant.species == plant.species,
            Plant.variety_type == plant.variety_type,
        )
        .one_or_none()
    )

    if existing is None:
        db.add(plant)
        return plant

    existing.drought_tolerance = plant.drought_tolerance
    existing.salinity_tolerance = plant.salinity_tolerance
    existing.heat_tolerance = plant.heat_tolerance
    existing.cold_tolerance = plant.cold_tolerance
    existing.disease_resistance = plant.disease_resistance
    existing.optimal_ph_min = plant.optimal_ph_min
    existing.optimal_ph_max = plant.optimal_ph_max
    existing.nutrient_uptake_efficiency = plant.nutrient_uptake_efficiency
    existing.root_depth_category = plant.root_depth_category
    existing.growth_cycle_days = plant.growth_cycle_days
    existing.yield_potential_kg_per_plant = plant.yield_potential_kg_per_plant
    existing.description = plant.description
    existing.image_url = plant.image_url
    existing.updated_at = NOW
    return existing


def persist_raw_scions(db, raw_plants: List[Dict[str, Any]]) -> List[Plant]:
    persisted: List[Plant] = []

    for item in raw_plants:
        plant = dict_to_plant(perenual_species_to_plant_dict(item))
        persisted.append(upsert_plant(db, plant))

    db.commit()
    return persisted


def dict_to_soil(d: Dict) -> SoilProfile:
    return SoilProfile(
        id=uuid.uuid4(),
        name=str(d["name"])[:255],
        soil_type=_ST.get(d.get("soil_type", "loam"), SoilType.loam),
        ph_level=float(d.get("ph_level", 6.5)),
        organic_matter_percent=float(d.get("organic_matter_percent", 2.0)),
        drainage_rate=_DR.get(d.get("drainage_rate", "moderate"), DrainageRate.moderate),
        water_retention=_WR.get(d.get("water_retention", "medium"), WaterRetention.medium),
        native_nitrogen=float(d.get("native_nitrogen", 20.0)),
        native_phosphorus=float(d.get("native_phosphorus", 15.0)),
        native_potassium=float(d.get("native_potassium", 150.0)),
        salinity_level=float(d.get("salinity_level", 0.5)),
        temperature_range_min=float(d.get("temperature_range_min", 5.0)),
        temperature_range_max=float(d.get("temperature_range_max", 30.0)),
        description=str(d.get("description", "") or ""),
        climate_zone=str(d.get("climate_zone", "temperate"))[:100],
        created_at=NOW,
        updated_at=NOW,
    )


def dict_to_plant(d: Dict) -> Plant:
    return Plant(
        id=uuid.uuid4(),
        name=str(d["name"])[:255],
        species=str(d["species"])[:255],
        variety_type=_VT.get(d.get("variety_type", "scion"), VarietyType.scion),
        drought_tolerance=max(1, min(10, int(d.get("drought_tolerance", 5)))),
        salinity_tolerance=max(1, min(10, int(d.get("salinity_tolerance", 5)))),
        heat_tolerance=max(1, min(10, int(d.get("heat_tolerance", 5)))),
        cold_tolerance=max(1, min(10, int(d.get("cold_tolerance", 5)))),
        disease_resistance=d.get("disease_resistance") or {},
        optimal_ph_min=float(d.get("optimal_ph_min", 6.0)),
        optimal_ph_max=float(d.get("optimal_ph_max", 7.0)),
        nutrient_uptake_efficiency=d.get("nutrient_uptake_efficiency") or {"N": 1.0, "P": 1.0, "K": 1.0},
        root_depth_category=_RDC.get(d.get("root_depth_category", "medium"), RootDepthCategory.medium),
        growth_cycle_days=max(1, int(d.get("growth_cycle_days", 90))),
        yield_potential_kg_per_plant=max(0.0, float(d.get("yield_potential_kg_per_plant", 0.0))),
        description=str(d.get("description", "") or ""),
        image_url=d.get("image_url"),
        created_at=NOW,
        updated_at=NOW,
    )


# ---------------------------------------------------------------------------
# Main
# ---------------------------------------------------------------------------

def main(reset: bool = False, refresh_cache: bool = False, skip_fetch: bool = False) -> None:
    db = SessionLocal()
    try:
        if reset:
            print("Clearing existing data...")
            db.query(ScenarioResult).delete()
            db.query(Scenario).delete()
            db.query(GraftCombination).delete()
            db.query(Plant).delete()
            db.query(SoilProfile).delete()
            db.commit()
            print("Cleared.\n")

        # 1. Fetch from Perenual (or load from cache, or skip entirely)
        cached = load_cache()
        if skip_fetch:
            print("[1/4] Skipping Perenual fetch — generating species list with Claude AI...")
            raw = generate_species_with_claude()
            save_cache(raw)
            print(f"  Generated {len(raw)} species\n")
        elif cached and not refresh_cache:
            raw = cached
            print(f"[1/4] Loaded {len(raw)} species from cache (use --refresh-cache to re-fetch)\n")
        else:
            print("[1/4] Fetching species from Perenual API...")
            raw = []
            seen: set = set()

            for term in SEARCH_TERMS:
                try:
                    results = fetch_species_list(term)
                    added = 0
                    for r in results:
                        if r["id"] not in seen and added < MAX_PER_TERM:
                            seen.add(r["id"])
                            raw.append(r)
                            added += 1
                    print(f"  '{term}' → {added} species")
                    time.sleep(2.0)
                except Exception as e:
                    print(f"  Warning: '{term}' fetch failed — {e}")
                    time.sleep(5.0)

            save_cache(raw)
            print(f"  Total: {len(raw)} species collected\n")

        print("[1.5/5] Saving fetched Perenual species to plants table...")
        try:
            persisted_scions = persist_raw_scions(db, raw)
            print(f"  Saved {len(persisted_scions)} raw scion plants ✓\n")
        except Exception as e:
            db.rollback()
            print(f"  Failed to save raw scion plants — {e}\n")
            raise

        # 2. Enrich + commit scions batch by batch
        print("[2/5] Enriching and saving scion plants...")
        scion_plants: List[Plant] = list(persisted_scions)
        BATCH = 8
        for i in range(0, len(raw), BATCH):
            batch = raw[i : i + BATCH]
            try:
                enriched = enrich_scions(batch)
                plant_dicts = merge_enriched_plant_data(batch, enriched)
            except Exception as e:
                print(f"  Warning: batch {i // BATCH + 1} AI enrichment failed — {e}")
                plant_dicts = [perenual_species_to_plant_dict(item) for item in batch]

            try:
                for d in plant_dicts:
                    p = upsert_plant(db, dict_to_plant(d))
                    if all(existing.id != p.id for existing in scion_plants):
                        scion_plants.append(p)
                db.commit()
                print(f"  Batch {i // BATCH + 1}: saved {len(plant_dicts)} plants ✓")
            except Exception as e:
                db.rollback()
                print(f"  Warning: batch {i // BATCH + 1} failed to persist — {e}")

        print(f"  Total scions saved: {len(scion_plants)}\n")

        # 3. Generate + commit rootstocks
        print("[3/5] Generating and saving rootstock cultivars...")
        rootstock_plants: List[Plant] = []
        try:
            rootstock_dicts = generate_rootstocks()
            for d in rootstock_dicts:
                p = dict_to_plant(d)
                db.add(p)
                rootstock_plants.append(p)
            db.commit()
            print(f"  Saved {len(rootstock_plants)} rootstocks ✓\n")
        except Exception as e:
            db.rollback()
            print(f"  Rootstock generation failed — {e}\n")

        # 4. Generate + commit soil profiles
        print("[4/5] Generating and saving soil profiles...")
        n_soils = 0
        try:
            soil_dicts = generate_soils()
            for d in soil_dicts:
                db.add(dict_to_soil(d))
            db.commit()
            n_soils = len(soil_dicts)
            print(f"  Saved {n_soils} soil profiles ✓\n")
        except Exception as e:
            db.rollback()
            print(f"  Soil generation failed — {e}\n")

        # 5. Generate + commit graft combinations
        print("[5/5] Generating and saving graft combinations...")
        inserted_grafts = 0
        try:
            plant_id_set = {str(p.id) for p in scion_plants + rootstock_plants}
            graft_dicts = generate_graft_combinations(scion_plants, rootstock_plants)
            for g in graft_dicts:
                rs_id = str(g.get("rootstock_id", ""))
                sc_id = str(g.get("scion_id", ""))
                if rs_id not in plant_id_set or sc_id not in plant_id_set:
                    continue
                db.add(GraftCombination(
                    id=uuid.uuid4(),
                    rootstock_id=uuid.UUID(rs_id),
                    scion_id=uuid.UUID(sc_id),
                    compatibility_score=float(g.get("compatibility_score", 0.8)),
                    vigor_boost=float(g.get("vigor_boost", 1.1)),
                    yield_modifier=float(g.get("yield_modifier", 1.1)),
                    stress_tolerance_override=g.get("stress_tolerance_override") or {},
                    notes=str(g.get("notes", "")),
                    source_reference=str(g.get("source_reference", "")),
                    created_at=NOW,
                    updated_at=NOW,
                ))
                inserted_grafts += 1
            db.commit()
            print(f"  Saved {inserted_grafts} graft combinations ✓\n")
        except Exception as e:
            db.rollback()
            print(f"  Graft generation failed — {e}\n")

        print("Done.")
        print(f"  {n_soils} soil profiles")
        print(f"  {len(scion_plants)} scion plants")
        print(f"  {len(rootstock_plants)} rootstock plants")
        print(f"  {inserted_grafts} graft combinations")

    except Exception as e:
        db.rollback()
        print(f"\nFailed: {e}", file=sys.stderr)
        raise
    finally:
        db.close()


if __name__ == "__main__":
    parser = argparse.ArgumentParser(description="Enrich ACI database with real plant data")
    parser.add_argument("--reset", action="store_true",
                        help="Clear all existing plants, soils, and grafts before enriching")
    parser.add_argument("--refresh-cache", action="store_true",
                        help="Re-fetch from Perenual even if a local cache exists")
    parser.add_argument("--skip-fetch", action="store_true",
                        help="Skip Perenual entirely — use Claude to generate the species list")
    args = parser.parse_args()
    main(reset=args.reset, refresh_cache=args.refresh_cache, skip_fetch=args.skip_fetch)
