import argparse
import hashlib
import json
import mimetypes
import re
import urllib.error
import urllib.parse
import urllib.request
from concurrent.futures import ThreadPoolExecutor, as_completed
from pathlib import Path


ROOT = Path(__file__).resolve().parent
LIBRARY_FILE = ROOT / 'data' / 'prompt-library.json'
ASSET_DIR = ROOT / 'prompt-assets'
MAX_SOURCE_BYTES = 16 * 1024 * 1024
MAX_IMAGE_BYTES = 20 * 1024 * 1024

SOURCES = [
    ('nanobanana-trending', 'NanoBanana 热门提示词', 'https://duolapi.vip/prompt-sources/nanobanana-trending.json'),
    ('erickkkyt-gptimage2-prompts', 'Erickkkyt GPT Image 2', 'https://duolapi.vip/prompt-sources/erickkkyt-gptimage2-prompts.json'),
    ('bigpeng-gpt-image-lab', 'BigPeng / GPT ImageLab', 'https://duolapi.vip/prompt-sources/bigpeng-gpt-image-lab.json'),
    ('image-prompt-library-zh', '中文生图提示词库', 'https://duolapi.vip/prompt-sources/image-prompt-library-zh.json'),
    ('image-prompt-library-en', '英文生图提示词库', 'https://duolapi.vip/prompt-sources/image-prompt-library-en.json'),
    ('freestylefly-gpt-image-2', 'Freestylefly GPT Image 2', 'https://duolapi.vip/prompt-sources/freestylefly-gpt-image-2.json'),
    ('davidwu-gpt-image2-prompts', 'Davidwu GPT Image 2', 'https://duolapi.vip/prompt-sources/davidwu-gpt-image2-prompts.json'),
    ('banana-prompt-quicker', 'Nano Banana 快速提示词', 'https://duolapi.vip/prompt-sources/banana-prompt-quicker.json'),
    ('pyth0nb3st-gallery', 'pyth0nb3st Gallery', 'https://duolapi.vip/prompt-sources/pyth0nb3st-gallery.json'),
    ('hiapiai-gpt-image-2-prompts', 'HiAPI GPT Image 2', 'https://duolapi.vip/prompt-sources/hiapiai-gpt-image-2-prompts.json'),
    ('tiange-happycapy-gpt-image-2', 'TIANGE / Happycapy', 'https://duolapi.vip/prompt-sources/tiange-happycapy-gpt-image-2.json'),
    ('youmind-nano-banana-pro', 'YouMind Nano Banana Pro', 'https://duolapi.vip/prompt-sources/youmind-nano-banana-pro.json'),
    ('youmind-gpt-image-2', 'YouMind GPT Image 2', 'https://duolapi.vip/prompt-sources/youmind-gpt-image-2.json'),
    ('visual-prompt-cookbook', 'AI 视觉风格手册', 'https://duolapi.vip/prompt-sources/visual-prompt-cookbook.json'),
    ('awesome-gpt4o-image-prompts', 'GPT-4o 生图提示词', 'https://duolapi.vip/prompt-sources/awesome-gpt4o-image-prompts.json'),
    ('mrchen-gpt-image-2-prompts', 'MrChen GPT Image 2', 'https://duolapi.vip/prompt-sources/mrchen-gpt-image-2-prompts.json'),
    ('tosea-gpt-image-2-prompts', 'ToseaAI', 'https://duolapi.vip/prompt-sources/tosea-gpt-image-2-prompts.json'),
    ('awesome-gpt-image', 'GPT 生图综合合集', 'https://duolapi.vip/prompt-sources/awesome-gpt-image.json'),
    ('wuyoscar-gpt-image2-skill', 'Wuyoscar GPT Image2 Skill', 'https://duolapi.vip/prompt-sources/wuyoscar-gpt-image2-skill.json'),
]

CLASSIFICATION_RULES = [
    ('人物人像', ['portrait', 'person', 'people', 'face', 'model', 'character', 'fashion', '人像', '人物', '肖像', '模特', '角色']),
    ('商业产品', ['product', 'commercial', 'ecommerce', 'advertising', 'brand', 'packaging', '商品', '产品', '电商', '广告', '包装', '品牌']),
    ('美食静物', ['food', 'dish', 'cuisine', 'still life', '美食', '食物', '餐饮', '静物']),
    ('建筑空间', ['architecture', 'building', 'interior', 'room', 'architectural', '建筑', '空间', '室内', '房间']),
    ('旅行风景', ['travel', 'landscape', 'nature', 'cityscape', 'street', 'scenery', '旅行', '风景', '自然', '城市', '街景']),
    ('插画艺术', ['illustration', 'anime', 'cartoon', 'comic', 'concept art', '插画', '动漫', '漫画', '卡通', '艺术']),
    ('海报设计', ['poster', 'typography', 'graphic design', 'editorial', 'layout', '海报', '字体', '平面设计', '排版']),
    ('图片编辑', ['edit', 'editing', 'retouch', 'restyle', 'transform', '图生图', '改图', '编辑', '重绘', '变换']),
]


def fetch_json(url):
    request = urllib.request.Request(url, headers={'User-Agent': 'GuangzhanPromptSync/1.0', 'Accept': 'application/json'})
    with urllib.request.urlopen(request, timeout=40) as response:
        payload = response.read(MAX_SOURCE_BYTES + 1)
    if len(payload) > MAX_SOURCE_BYTES:
        raise ValueError(f'{url} 超过 {MAX_SOURCE_BYTES // 1024 // 1024} MB')
    value = json.loads(payload.decode('utf-8-sig'))
    if not isinstance(value, list):
        raise ValueError(f'{url} 返回的不是数组')
    return value


def prompt_key(value):
    return re.sub(r'[^\w\u4e00-\u9fff]+', '', str(value or '').lower())


def classify(item):
    text = ' '.join(str(item.get(key) or '') for key in ('title', 'description', 'prompt', 'author'))
    text += ' ' + ' '.join(str(tag) for tag in item.get('tags', []) if tag)
    text = text.lower()
    for category, words in CLASSIFICATION_RULES:
        if any(word in text for word in words):
            return category
    return '综合灵感'


def image_url(value, source_url):
    raw = str(value or '').strip()
    if not raw:
        return ''
    return urllib.parse.urljoin(source_url, raw)


def normalize_sources():
    rows = []
    seen_prompts = set()
    errors = []
    for source_id, label, source_url in SOURCES:
        try:
            source_rows = fetch_json(source_url)
        except (OSError, ValueError, urllib.error.URLError) as error:
            errors.append({'source_id': source_id, 'label': label, 'error': str(error)})
            continue
        for item in source_rows:
            if not isinstance(item, dict) or not str(item.get('prompt') or '').strip():
                continue
            key = prompt_key(item.get('prompt'))
            if not key or key in seen_prompts:
                continue
            seen_prompts.add(key)
            cover = image_url(item.get('coverUrl') or item.get('preview') or (item.get('referenceImageUrls') or [''])[0], source_url)
            rows.append({
                'id': f'{source_id}:{item.get("id") or hashlib.sha1(key.encode("utf-8")).hexdigest()[:12]}',
                'source_id': source_id,
                'source_label': label,
                'source_url': item.get('sourceUrl') or source_url,
                'title': str(item.get('title') or str(item.get('prompt'))[:48])[:160],
                'prompt': str(item.get('prompt'))[:20000],
                'description': str(item.get('description') or (f'作者：{item.get("author")}' if item.get('author') else ''))[:400],
                'category': classify(item),
                'cover_url': cover,
                'author': str(item.get('author') or '')[:120],
                'tags': [str(tag)[:48] for tag in item.get('tags', [])[:12]] if isinstance(item.get('tags'), list) else [],
                'image_model': item.get('imageModel') or '',
                'image_size': item.get('imageSize') or '',
                'image_url': '',
            })
    return rows, errors


def asset_name(url):
    suffix = Path(urllib.parse.urlparse(url).path).suffix.lower()
    if suffix not in ('.jpg', '.jpeg', '.png', '.webp', '.gif', '.avif'):
        suffix = '.webp'
    return hashlib.sha256(url.encode('utf-8')).hexdigest() + suffix


def download_asset(row):
    url = row.get('cover_url')
    if not url.startswith('https://'):
        return row
    name = asset_name(url)
    target = ASSET_DIR / name
    if not target.exists():
        request = urllib.request.Request(url, headers={'User-Agent': 'GuangzhanPromptSync/1.0', 'Accept': 'image/*'})
        try:
            with urllib.request.urlopen(request, timeout=40) as response:
                content_type = response.headers.get('Content-Type', '').split(';', 1)[0].lower()
                if content_type and not content_type.startswith('image/'):
                    raise ValueError(f'响应类型不是图片：{content_type}')
                data = response.read(MAX_IMAGE_BYTES + 1)
            if len(data) > MAX_IMAGE_BYTES:
                raise ValueError('图片超过 20 MB')
            temp = target.with_suffix(target.suffix + '.part')
            temp.write_bytes(data)
            temp.replace(target)
        except (OSError, ValueError, urllib.error.URLError):
            return row
    row['image_url'] = f'/prompt-assets/{name}'
    return row


def sync(max_workers):
    ASSET_DIR.mkdir(parents=True, exist_ok=True)
    LIBRARY_FILE.parent.mkdir(parents=True, exist_ok=True)
    rows, errors = normalize_sources()
    completed = 0
    with ThreadPoolExecutor(max_workers=max_workers) as executor:
        futures = [executor.submit(download_asset, row) for row in rows]
        for future in as_completed(futures):
            future.result()
            completed += 1
            if completed % 100 == 0:
                print(f'downloaded {completed}/{len(rows)}')
    payload = {
        'updated_at': __import__('datetime').datetime.utcnow().isoformat() + 'Z',
        'total': len(rows),
        'items': rows,
        'source_errors': errors,
    }
    temp = LIBRARY_FILE.with_suffix('.json.tmp')
    temp.write_text(json.dumps(payload, ensure_ascii=False, indent=2) + '\n', encoding='utf-8')
    temp.replace(LIBRARY_FILE)
    print(json.dumps({'total': len(rows), 'images': sum(bool(row['image_url']) for row in rows), 'source_errors': errors}, ensure_ascii=False))


if __name__ == '__main__':
    parser = argparse.ArgumentParser(description='同步提示词和图片到服务器本地存储')
    parser.add_argument('--workers', type=int, default=8)
    args = parser.parse_args()
    sync(max(1, min(args.workers, 16)))
