#!/usr/bin/env python3
"""
eiafans 全量同步到服务器 + Playwright 正文补抓
============================================
步骤1: 把所有 204,578 条记录导入服务器 search.db
步骤2: 用 Playwright 逐条获取真实正文

用法:
  python3 backfill_eiafans.py --sync-db     # 仅同步所有记录到服务器
  python3 backfill_eiafans.py --crawl=100   # 抓取100条正文
  python3 backfill_eiafans.py --full        # 全量抓取所有204k条正文
  python3 backfill_eiafans.py --full --resume  # 续跑

策略：
  - 步骤2 从本地 quality_results 读取所有待补记录（204k条全量）
  - 先补最新的（id 降序），旧的慢慢补
  - 每50条自动同步到服务器
  - 中断后 --resume 自动跳过已补的
"""

import os, re, sys, time, sqlite3, subprocess, base64, asyncio

BASE_DIR = os.path.dirname(os.path.abspath(__file__))
QUALITY_DB = os.path.join(BASE_DIR, "quality_results.db")
SERVER_SSH = "root@1.94.217.116"
SERVER_SEARCH_DB = "/root/search.db"

BATCH_SIZE = 50
DELAY = 1.5
PLACEHOLDER_MARK = "该数据来自环评爱好者论坛"

stats = {"ok": 0, "skip": 0, "error": 0, "total": 0}


def load_quality_rows():
    """从本地 quality_results 加载所有 eiafans 记录"""
    conn = sqlite3.connect(QUALITY_DB)
    conn.row_factory = sqlite3.Row
    rows = conn.execute(
        "SELECT id, url, title, content, publish_date, summary "
        "FROM quality_results WHERE site_name='环评爱好者' AND domain='www.eiafans.com' "
        "ORDER BY id DESC"
    ).fetchall()
    conn.close()
    return [dict(r) for r in rows]


def sync_all_to_server(rows):
    """全量同步到服务器 search.db"""
    print(f"\n📤 同步 {len(rows)} 条到服务器...")
    ok = 0
    for i in range(0, len(rows), 100):
        batch = rows[i:i+100]
        # 批量 INSERT OR IGNORE + UPDATE
        vals = []
        for r in batch:
            def e(v):
                if v is None: return "NULL"
                return "'" + str(v).replace("'", "''") + "'"
            title = (r["title"] or "").replace("'", "''")
            content = (r["content"] or "").replace("'", "''")
            summary = (r["summary"] or "").replace("'", "''")
            page_url = (r["url"] or "").replace("'", "''")
            date = (r.get("publish_date") or "").replace("'", "''")
            dr = f"CAST(strftime('%s', '{date}') AS INTEGER)" if date else "NULL"
            vals.append(
                f"('环评爱好者', '{page_url}', '{title}', '{content}', '{date}', '{summary}', {dr})"
            )
        
        sql = (
            "INSERT OR IGNORE INTO gov_raw (site_name, page_url, title, content, publish_date, summary, date_rank) "
            f"VALUES {','.join(vals)};"
        )
        b64 = base64.b64encode(sql.encode()).decode()
        r2 = subprocess.run(
            ["ssh", SERVER_SSH,
             f"echo {b64} | base64 -d | sqlite3 {SERVER_SEARCH_DB}"],
            capture_output=True, timeout=60,
        )
        if r2.returncode == 0:
            ok += len(batch)
        else:
            print(f"  ⚠️ 批次 {i//100+1} 出错: {r2.stderr.decode()[:100]}")
        
        if (i // 100) % 5 == 0:
            print(f"\r  ✅ {ok}/{len(rows)}", end="")
        time.sleep(0.3)
    
    print(f"\n  ✅ 同步完成: {ok} 条")
    
    # 重建 FTS
    subprocess.run(
        ["ssh", SERVER_SSH,
         f"sqlite3 {SERVER_SEARCH_DB} \"INSERT INTO gov_search(gov_search) VALUES('rebuild');\""],
        capture_output=True, timeout=30,
    )
    print(f"  ✅ FTS 索引重建")


def get_missing_locally():
    """获取本地所有还缺正文的 eiafans 记录"""
    conn = sqlite3.connect(QUALITY_DB)
    conn.row_factory = sqlite3.Row
    rows = conn.execute(
        "SELECT id, url, title, publish_date FROM quality_results "
        "WHERE site_name='环评爱好者' AND domain='www.eiafans.com' AND "
        "(content IS NULL OR content = '' OR content LIKE ?) "
        "ORDER BY id DESC",
        (f'%{PLACEHOLDER_MARK}%',)
    ).fetchall()
    conn.close()
    return [dict(r) for r in rows]


def find_local(url):
    """查找本地记录"""
    conn = sqlite3.connect(QUALITY_DB)
    r = conn.execute(
        "SELECT id, content FROM quality_results WHERE url=? AND site_name='环评爱好者'",
        (url,)
    ).fetchone()
    conn.close()
    return {"id": r[0], "content": r[1]} if r else None


def update_local(local_id, content, summary, attachments):
    """更新本地 + 同步到服务器"""
    attach_text = ""
    if attachments:
        attach_text = "\n" + "\n".join(f"📎 {a['name']}: {a['url']}" for a in attachments)
    full_summary = (summary or "")[:480] + attach_text
    full_summary = full_summary[:500]

    conn = sqlite3.connect(QUALITY_DB)
    conn.execute(
        "UPDATE quality_results SET content=?, summary=?, crawled_at=datetime('now','localtime') WHERE id=?",
        (content, full_summary, local_id)
    )
    conn.commit()
    conn.close()


def sync_to_server(url, content, summary):
    """实时同步单条到服务器"""
    esc_summary = summary.replace("'", "''")
    esc_content = content.replace("'", "''")
    esc_url = url.replace("'", "''")
    sql = (
        f"UPDATE gov_raw SET content='{esc_content}', summary='{esc_summary}' "
        f"WHERE page_url='{esc_url}' AND site_name='环评爱好者';"
    )
    b64 = base64.b64encode(sql.encode()).decode()
    subprocess.run(
        ["ssh", SERVER_SSH, f"echo {b64} | base64 -d | sqlite3 {SERVER_SEARCH_DB}"],
        capture_output=True, timeout=30,
    )


def clean_discuz_html(html):
    """清洗 Discuz! HTML 噪音"""
    if not html:
        return ""
    html = re.sub(r'<font[^>]*>\s*</font>', '', html)
    html = re.sub(r'<font[^>]*>(.*?)</font>', r'\1', html, flags=re.DOTALL)
    html = re.sub(r'<strong>\s*<font[^>]*>(.*?)</font>\s*</strong>', r'<strong>\1</strong>', html, flags=re.DOTALL)
    html = re.sub(r'<div[^>]*align="[^"]*"[^>]*>', '', html)
    html = re.sub(r'</div>', '\n', html)
    html = re.sub(r'<font[^>]*>&quot;</font>', '', html)
    html = re.sub(r'\s+style="[^"]*"', '', html)
    html = re.sub(r'\n{3,}', '\n\n', html)
    return html.strip()


async def extract_content(page):
    """从 Playwright 页面提取 Discuz! 正文"""
    content = ""
    attachments = []
    
    for sel in ["div[id^='postmessage_']", "td.t_f", ".t_fsz td", ".t_fsz"]:
        try:
            els = await page.query_selector_all(sel)
            if els:
                content = await els[0].inner_html()
                if content:
                    break
        except:
            pass

    try:
        for el in await page.query_selector_all("a[href*='forum.php?mod=attachment']"):
            href = await el.get_attribute("href")
            name = (await el.inner_text()).strip() or "附件"
            if href:
                if not href.startswith("http"):
                    href = f"http://www.eiafans.com/{href.lstrip('/')}"
                attachments.append({"name": name, "url": href})
    except:
        pass

    if content:
        content = re.sub(r'<div[^>]*class="sign"[^>]*>.*?</div>', '', content, flags=re.DOTALL)
        content = re.sub(r'<div[^>]*class="pstatus"[^>]*>.*?</div>', '', content, flags=re.DOTALL)
        content = content.strip()
    return content, attachments


async def crawl_one(page, url):
    """抓取一个详情页"""
    try:
        await page.goto(url, timeout=25000, wait_until="domcontentloaded")
        for _ in range(3):
            await page.wait_for_timeout(3000)
            content, attachments = await extract_content(page)
            if content and len(content) > 50:
                break
        
        summary = ""
        if content:
            text = re.sub(r'<[^>]+>', ' ', content)
            text = re.sub(r'\s+', ' ', text).strip()
            text = re.sub(r'本帖最后由.*?编辑', '', text)
            summary = text[:300]
        
        return {"ok": bool(content) and len(content) > 50, "content": content, "summary": summary, "attachments": attachments}
    except Exception as e:
        return {"ok": False, "error": str(e)[:60], "content": "", "summary": "", "attachments": []}


async def main():
    import argparse
    parser = argparse.ArgumentParser()
    parser.add_argument("--sync-db", action="store_true", help="仅全量同步到服务器")
    parser.add_argument("--full", action="store_true", help="全量抓取所有204k")
    parser.add_argument("--resume", action="store_true", help="续跑（跳过已补）")
    parser.add_argument("--crawl", type=int, default=0, help="抓取N条")
    parser.add_argument("--test", type=int, default=0, help="测试N条")
    args = parser.parse_args()

    # ── 步骤1: 全量同步到服务器 ──
    if args.sync_db:
        rows = load_quality_rows()
        sync_all_to_server(rows)
        return

    # ── 步骤2: 抓取正文 ──
    missing = get_missing_locally()
    total = len(missing)
    print(f"\n📊 本地待补: {total} 条")
    
    if args.test:
        limit = args.test
    elif args.crawl:
        limit = args.crawl
    elif args.full:
        limit = total
    else:
        limit = 500
    
    todo = missing[:limit]
    print(f"  ➡ 本次: {len(todo)} 条")
    
    from playwright.async_api import async_playwright
    
    async with async_playwright() as pw:
        browser = await pw.chromium.launch(
            headless=True,
            args=["--no-sandbox", "--disable-setuid-sandbox", "--disable-dev-shm-usage", "--disable-gpu"],
        )
        ctx = await browser.new_context(
            user_agent="Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36 Chrome/125.0.0.0 Safari/537.36",
            viewport={"width": 1280, "height": 1024}, locale="zh-CN",
        )
        page = await ctx.new_page()
        
        start = time.time()
        buffer = []
        
        for i, t in enumerate(todo):
            url = t["url"]
            tid = re.search(r'thread-(\d+)', url)
            tid = tid.group(1) if tid else "?"
            
            # 续跑跳过已补
            local = find_local(url)
            if local and local["content"] and "该数据来自环评" not in local["content"]:
                stats["skip"] += 1
                print(f"\r  [{i+1}/{len(todo)}] tid={tid} ⏭️ 已补", end="")
                continue
            
            print(f"\r  [{i+1}/{len(todo)}] tid={tid}... ", end="")
            sys.stdout.flush()
            
            result = await crawl_one(page, url)
            
            if result["ok"]:
                stats["ok"] += 1
                content = clean_discuz_html(result["content"])
                na = f" +{len(result['attachments'])}附件" if result["attachments"] else ""
                print(f"✅ {len(content)}B{na}", end="")
                
                if local:
                    update_local(local["id"], content, result["summary"], result["attachments"])
                else:
                    print("(本地无匹配)", end="")
                
                buffer.append({"url": url, "content": content, "summary": result["summary"]})
            else:
                stats["error"] += 1
                print(f"❌ {result.get('error','')[:40]}", end="")
            
            elapsed = time.time() - start
            rate = (i + 1) / max(elapsed, 1)
            eta = (len(todo) - i - 1) / max(rate, 0.001)
            print(f" | {(i+1)}/{len(todo)} | {rate:.1f}/s, ETA {eta/60:.0f}min", end="")
            
            if len(buffer) >= BATCH_SIZE:
                # 批量同步到服务器
                print(f"\n📤 同步 {len(buffer)} 条...")
                for r in buffer:
                    sync_to_server(r["url"], r["content"], r["summary"])
                # FTS
                subprocess.run(
                    ["ssh", SERVER_SSH,
                     f"sqlite3 {SERVER_SEARCH_DB} \"INSERT INTO gov_search(gov_search) VALUES('rebuild');\""],
                    capture_output=True, timeout=30,
                )
                buffer = []
            
            await asyncio.sleep(DELAY)
        
        # 最终同步
        if buffer:
            print(f"\n📤 同步 {len(buffer)} 条...")
            for r in buffer:
                sync_to_server(r["url"], r["content"], r["summary"])
            subprocess.run(
                ["ssh", SERVER_SSH,
                 f"sqlite3 {SERVER_SEARCH_DB} \"INSERT INTO gov_search(gov_search) VALUES('rebuild');\""],
                capture_output=True, timeout=30,
            )
        
        await browser.close()
    
    elapsed = time.time() - start
    done = stats["ok"] + stats["skip"] + stats["error"]
    print(f"\n\n{'='*50}")
    print(f"🏁 完成！")
    print(f"  ✅ 成功: {stats['ok']}")
    print(f"  ⏭️ 跳过: {stats['skip']}")
    print(f"  ❌ 失败: {stats['error']}")
    print(f"  ⏱ 耗时: {elapsed/60:.1f}分钟")
    print(f"  ⚡ 速率: {done/max(elapsed,1):.1f} 条/秒")


if __name__ == "__main__":
    asyncio.run(main())
