import os
import sys
import json
import time
import glob
import logging
import threading
import random
from concurrent.futures import ThreadPoolExecutor, as_completed

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 pinscrape.v2 import Pinterest
from pinscrape.database import get_pending_combined, BatchWriter
from pinscrape.utils import load_proxies

logging.basicConfig(
    level=logging.INFO,
    format="%(asctime)s [%(levelname)s] %(message)s",
    datefmt="%H:%M:%S"
)
logger = logging.getLogger(__name__)

_worker_local = threading.local()


def load_config():
    config_path = os.path.join(ROOT_DIR, "config.json")
    defaults = {
        "image_result": 26,
        "min_snippet": 3,
        "page_size": 0,
        "concurrency": 5,
        "max_retries": 3,
        "retry_delay": 3,
        "warmup_request": True,
        "session_max_requests": 50,
        "db_batch_size": 100,
    }
    if os.path.exists(config_path):
        with open(config_path, "r") as f:
            cfg = json.load(f)
        defaults.update(cfg)
    return defaults


def resolve_page_size(config):
    page_size = int(config.get("page_size") or 0)
    if page_size > 0:
        return page_size
    return max(int(config.get("image_result", 26)), int(config.get("min_snippet", 3)))


def get_worker_client(proxy_list, config):
    p = getattr(_worker_local, "client", None)
    if p is None:
        p = Pinterest(
            proxy_list=proxy_list,
            sleep_time=(1.5, 4.0),
            max_retries=config.get("max_retries", 3),
            warmup=config.get("warmup_request", True),
            max_requests_per_session=config.get("session_max_requests", 50),
        )
        _worker_local.client = p
    return p


def scrape_keyword(db_path, item, writer, proxy_list, config):
    keyword_id = item["id"]
    keyword = item["keyword"]
    need_images = item["need_images"]
    need_snippet = item["need_snippet"]
    page_size = resolve_page_size(config)
    max_retries = config["max_retries"]

    p = get_worker_client(proxy_list, config)

    for attempt in range(1, max_retries + 1):
        try:
            out = p.search_combined(keyword, page_size=page_size)
            images = out["images"]
            descriptions = out["descriptions"]

            if need_images and images:
                writer.update_images(keyword_id, json.dumps(images, ensure_ascii=False))
            if need_snippet:
                writer.update_snippet_description(keyword_id, descriptions)

            error = None
            if need_images and not images:
                error = "no images found"
            return {
                "id": keyword_id,
                "keyword": keyword,
                "count": len(images),
                "desc_count": len(descriptions),
                "error": error,
            }
        except Exception as e:
            logger.warning(f"  Attempt {attempt}/{max_retries} failed for '{keyword}': {e}")
            if attempt < max_retries:
                wait = random.uniform(3, 8)
                logger.info(f"  Waiting {wait:.1f}s before retry...")
                time.sleep(wait)
                p.reset_session()

    return {
        "id": keyword_id,
        "keyword": keyword,
        "count": 0,
        "desc_count": 0,
        "error": "max retries exceeded",
    }


def scrape_database(db_path, proxy_list, config):
    db_name = os.path.basename(db_path)
    logger.info(f"\n{'='*50}")
    logger.info(f"Processing: {db_name}")
    logger.info(f"{'='*50}")

    pending = get_pending_combined(db_path)
    if not pending:
        logger.info(f"No pending keywords in {db_name}")
        return

    logger.info(f"Found {len(pending)} pending keywords")
    logger.info(f"Page size: {resolve_page_size(config)} (image_result={config.get('image_result')}, "
                f"min_snippet={config.get('min_snippet')})")
    concurrency = config["concurrency"]
    completed = 0
    lock = threading.Lock()
    writer = BatchWriter(db_path, batch_size=config.get("db_batch_size", 100)).start()

    def progress_callback(future):
        nonlocal completed
        result = future.result()
        with lock:
            completed += 1
            status = "OK" if not result["error"] else f"FAIL ({result['error']})"
            logger.info(f"  [{completed}/{len(pending)}] {result['keyword']} -> "
                        f"{result['count']} images, {result['desc_count']} desc [{status}]")

    try:
        with ThreadPoolExecutor(max_workers=concurrency) as executor:
            futures = {}
            for item in pending:
                future = executor.submit(
                    scrape_keyword,
                    db_path, item, writer,
                    proxy_list, config
                )
                futures[future] = item
                future.add_done_callback(progress_callback)

            for future in as_completed(futures):
                pass
    finally:
        writer.flush()
        writer.close()

    logger.info(f"Completed {db_name}: {completed}/{len(pending)} keywords processed")


def find_databases():
    db_dir = os.path.join(ROOT_DIR, "database")
    if not os.path.exists(db_dir):
        return []
    return sorted(glob.glob(os.path.join(db_dir, "*.sqlite")))


def main():
    config = load_config()
    logger.info(f"Config: page_size={resolve_page_size(config)}, image_result={config['image_result']}, "
                f"min_snippet={config.get('min_snippet')}, concurrency={config['concurrency']}, "
                f"max_retries={config['max_retries']}, warmup={config.get('warmup_request', True)}, "
                f"session_max_requests={config.get('session_max_requests', 50)}, "
                f"db_batch_size={config.get('db_batch_size', 100)}")

    proxy_list = load_proxies(os.path.join(ROOT_DIR, "proxies.txt"))
    logger.info(f"Loaded {len(proxy_list)} proxies")

    databases = find_databases()
    if not databases:
        logger.warning("No .sqlite files found in database/ folder")
        return

    logger.info(f"Found {len(databases)} database(s): {[os.path.basename(d) for d in databases]}")

    for db_path in databases:
        scrape_database(db_path, proxy_list, config)

    logger.info("\nAll databases processed!")


if __name__ == "__main__":
    main()
