import asyncio
import json
import logging
import os
import sys
from pathlib import Path

logging.getLogger("httpx").setLevel(logging.WARNING)
logging.getLogger("httpcore").setLevel(logging.WARNING)

SCRIPTS_DIR = os.path.dirname(os.path.abspath(__file__))
ROOT_DIR = os.path.dirname(SCRIPTS_DIR)

if ROOT_DIR not in sys.path:
    sys.path.insert(0, ROOT_DIR)

from openai import AsyncOpenAI
from pinscrape.database import get_all_db_files, get_db_connection


class KeyPool:
    def __init__(self, keys):
        self.keys = keys
        self.index = 0
        self.lock = asyncio.Lock()

    async def next(self):
        async with self.lock:
            key = self.keys[self.index % len(self.keys)]
            self.index += 1
            return key

    def __len__(self):
        return len(self.keys)


def load_api_keys(ai_cfg, root_dir):
    key_file = ai_cfg.get("key_file")
    if key_file:
        key_path = Path(root_dir) / key_file
        if key_path.exists():
            keys = [
                line.strip() for line in key_path.read_text(encoding="utf-8").splitlines()
                if line.strip() and not line.strip().startswith("#")
            ]
            if keys:
                return keys

    single_key = ai_cfg.get("api_key")
    if single_key:
        return [single_key]

    return []


def build_model_chain(ai_cfg):
    primary = {
        "model": ai_cfg.get("model", "gpt-3.5-turbo"),
        "base_url": ai_cfg.get("base_url", "https://api.openai.com/v1"),
        "api_key": ai_cfg.get("api_key"),
    }
    chain = [primary]
    for fb in ai_cfg.get("fallback_models", []):
        if isinstance(fb, str):
            chain.append({
                "model": fb,
                "base_url": primary["base_url"],
                "api_key": primary["api_key"],
            })
        elif isinstance(fb, dict):
            chain.append({
                "model": fb.get("model", primary["model"]),
                "base_url": fb.get("base_url", primary["base_url"]),
                "api_key": fb.get("api_key", primary["api_key"]),
            })
    return chain


def clean_title(title):
    return title.replace('"', '').replace('\u201c', '').replace('\u201d', '').replace('\u2018', '').replace('\u2019', '').strip()


async def generate_content(model_chain, prompt_template, keyword, max_retries, retry_delay, key_pool=None):
    prompt = prompt_template.replace("{keyword}", keyword)
    last_error = None

    for model_info in model_chain:
        for attempt in range(1, max_retries + 1):
            api_key = await key_pool.next() if key_pool else model_info["api_key"]
            client = AsyncOpenAI(
                api_key=api_key,
                base_url=model_info["base_url"],
            )
            try:
                response = await client.chat.completions.create(
                    model=model_info["model"],
                    messages=[{"role": "user", "content": prompt}],
                )
                content = response.choices[0].message.content
                if content:
                    return content.strip()
                raise ValueError("Empty response content")
            except Exception as e:
                last_error = e
                err_str = str(e).lower()
                is_rate_limit = "429" in err_str or "rate" in err_str
                is_server_error = any(code in err_str for code in ["500", "502", "503", "504"])
                is_timeout = "timeout" in err_str or "timed out" in err_str

                if is_rate_limit or is_server_error or is_timeout:
                    wait = retry_delay * attempt
                    print(f"  [Retry {attempt}/{max_retries}] {model_info['model']} - {e} (waiting {wait}s)")
                    await asyncio.sleep(wait)
                else:
                    print(f"  [Skip model] {model_info['model']} - {e}")
                    break
        else:
            continue

    print(f"  All models failed for '{keyword}': {last_error}")
    return None


async def process_db_ai(semaphore, db_path, model_chain, prompt_template, mode, counter, total, lock, max_retries, retry_delay, key_pool=None):
    conn = get_db_connection(db_path)
    column = "ai_title" if mode == "title" else "ai_content"
    rows = conn.execute(
        f"SELECT id, keyword FROM posts WHERE {column} IS NULL "
        "AND images IS NOT NULL AND images != '' AND images != '[]'"
    ).fetchall()

    if not rows:
        conn.close()
        return

    print(f"Processing {len(rows)} items for {mode} in {os.path.basename(db_path)}...")

    async def task(row):
        async with semaphore:
            content = await generate_content(model_chain, prompt_template, row["keyword"], max_retries, retry_delay, key_pool)
            if content:
                if mode == "title":
                    content = clean_title(content)
                c = get_db_connection(db_path)
                c.execute(f"UPDATE posts SET {column} = ?, updated_at = datetime('now') WHERE id = ?", (content, row["id"]))
                c.commit()
                c.close()
                async with lock:
                    counter[0] += 1
                    current = counter[0]
                print(f"[{current}/{total}] Success:{row['keyword']} <{mode}>")

    tasks = [task(dict(row)) for row in rows]
    await asyncio.gather(*tasks)
    conn.close()


async def run_ai_generation(mode):
    config_path = os.path.join(ROOT_DIR, "config.json")
    with open(config_path, "r", encoding="utf-8") as f:
        config = json.load(f)

    ai_cfg = config.get("ai", {})
    api_keys = load_api_keys(ai_cfg, ROOT_DIR)
    if not api_keys:
        print("No API keys found. Set 'api_key' or 'key_file' in config.json")
        return

    key_pool = KeyPool(api_keys)
    model_chain = build_model_chain(ai_cfg)
    concurrency = ai_cfg.get("concurrency", 5)
    max_retries = ai_cfg.get("max_retries", 2)
    retry_delay = ai_cfg.get("retry_delay", 5)
    semaphore = asyncio.Semaphore(concurrency)

    print(f"Model chain: {' -> '.join(m['model'] for m in model_chain)}")
    print(f"API keys loaded: {len(key_pool)}")
    print(f"Max retries per model: {max_retries}, Retry delay: {retry_delay}s")

    prompt_path = Path(ROOT_DIR) / "prompts" / f"{mode}.txt"
    if not prompt_path.exists():
        print(f"Prompt file not found: {prompt_path}")
        return
    prompt_template = prompt_path.read_text(encoding="utf-8")

    db_files = get_all_db_files()

    total = 0
    for db_file in db_files:
        conn = get_db_connection(str(db_file))
        column = "ai_title" if mode == "title" else "ai_content"
        count = conn.execute(
            f"SELECT COUNT(*) FROM posts WHERE {column} IS NULL "
            "AND images IS NOT NULL AND images != '' AND images != '[]'"
        ).fetchone()[0]
        total += count
        conn.close()

    if total == 0:
        print(f"No items to process for {mode}.")
        return

    counter = [0]
    lock = asyncio.Lock()

    print(f"Total items to process: {total}")

    for db_file in db_files:
        await process_db_ai(semaphore, str(db_file), model_chain, prompt_template, mode, counter, total, lock, max_retries, retry_delay, key_pool)


if __name__ == "__main__":
    if len(sys.argv) < 2 or sys.argv[1] not in ("title", "article"):
        print("Usage: python scripts/ai_generator.py <mode>")
        print("  mode: title | article")
        sys.exit(1)
    asyncio.run(run_ai_generation(sys.argv[1]))
