#!/usr/bin/env python3
"""
从 backup jsonl 恢复数据到服务器。
不覆盖已有数据（INSERT OR IGNORE by URL）。
"""
import json, sqlite3, os, sys, hashlib, subprocess, base64, glob
from pathlib import Path
from collections import OrderedDict

BACKUP_DIR = Path(__file__).parent / "output" / "backup"
SERVER_SSH = "root@1.94.217.116"
SERVER_CRAWLER_DB = "/root/crawler_results.db"
SERVER_SEARCH_DB = "/root/search.db"

def load_jsonl_files():
    """读取所有 success jsonl 文件，去重 by page_url"""
    seen = OrderedDict()
    files = sorted(glob.glob(str(BACKUP_DIR / "gov_list_success_*.jsonl")))
    total_lines = 0
    for fpath in files:
        fsize = os.path.getsize(fpath)
        if fsize == 0:
            continue
        print(f"  读取 {Path(fpath).name} ({fsize/1024/1024:.0f}MB)...")
        with open(fpath, 'r', encoding='utf-8') as f:
            for line in f:
                line = line.strip()
                if not line:
                    continue
                total_lines += 1
                try:
                    d = json.loads(line)
                except:
                    continue
                url = d.get("page_url") or d.get("source_url", "")
                if not url:
                    continue
                if url not in seen:
                    seen[url] = d
    print(f"  共读取 {total_lines} 行, 去重后 {len(seen)} 条唯一 URL")
    return list(seen.values())

def build_insert_sql(items, batch_size=100):
    """
    生成 INSERT OR IGNORE SQL 批处理语句。
    为每个唯一 domain/site_name 创建 site_id。
    """
    # 1. 收集所有站点
    sites = {}
    site_next_id = 100000  # 大 site_id 避免冲突
    for d in items:
        domain = (d.get("domain") or "").strip()
        sname = (d.get("site_name") or "").strip()
        if not domain and not sname:
            continue
        key = domain or sname
        if key not in sites:
            site_next_id += 1
            sites[key] = {
                "id": site_next_id,
                "name": sname or domain,
                "domain": domain,
                "url": d.get("source_url", ""),
            }

    # 2. 检查服务器上已有哪些站点（尝试匹配 domain）
    print(f"  📋 站点数: {len(sites)}")
    
    # 3. 生成全部 SQL
    all_sqls = []
    
    # 3a. 站点表 INSERT OR IGNORE
    for key, site in sites.items():
        name_esc = site["name"].replace("'", "''")
        domain_esc = site["domain"].replace("'", "''")
        url_esc = site["url"].replace("'", "''")
        all_sqls.append(
            f"INSERT OR IGNORE INTO crawl_sites (id, name, url, domain, crawl_method, status, total_crawled) "
            f"VALUES ({site['id']}, '{name_esc}', '{url_esc}', '{domain_esc}', 'scrapy', '已完成', 0);"
        )
    
    # 3b. 数据表 INSERT OR IGNORE（按 site_id 分批）
    site_items = {}
    for d in items:
        domain = (d.get("domain") or "").strip()
        sname = (d.get("site_name") or "").strip()
        key = domain or sname
        if key not in sites:
            continue
        sid = sites[key]["id"]
        if sid not in site_items:
            site_items[sid] = []
        site_items[sid].append(d)
    
    for sid, site_records in site_items.items():
        batch = []
        for d in site_records:
            title = (d.get("title") or "").replace("'", "''")
            url = (d.get("page_url") or d.get("source_url", "")).replace("'", "''")
            content = (d.get("content") or "").replace("'", "''")
            summary = (d.get("content_text") or content[:500] if content else "").replace("'", "''")
            date = (d.get("publish_date") or "")[:10]
            domain = (d.get("domain") or "").replace("'", "''")
            ch = hashlib.md5(content.encode()).hexdigest()
            
            batch.append(
                f"({sid}, '{title}', '{url}', '{content}', '{date}', '{summary}', '{domain}', 'gov', '{ch}')"
            )
            
            if len(batch) >= batch_size:
                sql = f"INSERT OR IGNORE INTO crawl_results (site_id, title, url, content, publish_date, summary, domain, category, content_hash) VALUES {','.join(batch)};"
                all_sqls.append(sql)
                batch = []
        
        if batch:
            sql = f"INSERT OR IGNORE INTO crawl_results (site_id, title, url, content, publish_date, summary, domain, category, content_hash) VALUES {','.join(batch)};"
            all_sqls.append(sql)
    
    return all_sqls, len(items)

def run():
    print("=" * 50)
    print("🔄 从 backup 恢复数据到服务器")
    print("=" * 50)
    
    # 1. 读取 jsonl
    items = load_jsonl_files()
    if not items:
        print("❌ 没有可恢复的数据")
        return
    
    # 2. 生成 SQL
    print("\n📝 生成 INSERT SQL...")
    sqls, count = build_insert_sql(items)
    print(f"   SQL 批次数: {len(sqls)}, 记录数: {count}")
    
    # 3. 批量发送到服务器
    print(f"\n📤 发送到 {SERVER_SSH}...")
    total_sql = "\n".join(sqls)
    b64 = base64.b64encode(total_sql.encode()).decode()
    
    print(f"   SQL 大小: {len(total_sql)/1024:.0f}KB, Base64: {len(b64)/1024:.0f}KB")
    
    # 分块传输（避免单次命令过长）
    chunk_size = 500 * 1024  # 500KB chunks
    total_chunks = (len(b64) + chunk_size - 1) // chunk_size
    print(f"   分 {total_chunks} 块传输...")
    
    # 先把完整 SQL 写到本地临时文件，再 SCP 到服务器 sqlite3
    tmp_sql = "/tmp/restore_crawler.sql"
    with open(tmp_sql, 'w', encoding='utf-8') as f:
        f.write(total_sql)
    
    # SCP 到服务器
    try:
        subprocess.run([
            "scp", "-o", "ConnectTimeout=15", "-o", "StrictHostKeyChecking=no",
            tmp_sql, f"{SERVER_SSH}:/tmp/restore_crawler.sql"
        ], capture_output=True, timeout=60)
        print("   ✅ SQL 文件已上传到服务器")
    except Exception as e:
        print(f"   ❌ SCP 失败: {e}")
        return
    
    # 在服务器上执行
    cmd = f"sqlite3 {SERVER_CRAWLER_DB} '.read /tmp/restore_crawler.sql' && echo OK"
    result = subprocess.run([
        "ssh", "-o", "ConnectTimeout=15", "-o", "StrictHostKeyChecking=no",
        SERVER_SSH, cmd
    ], capture_output=True, text=True, timeout=300)
    
    if result.returncode == 0 and "OK" in result.stdout:
        print("   ✅ 数据已恢复到服务器 crawler_results.db")
    else:
        print(f"   ⚠️ 可能有问题: stdout={result.stdout[:200]}, stderr={result.stderr[:200]}")
    
    # 验证
    result = subprocess.run([
        "ssh", "-o", "ConnectTimeout=10", "-o", "StrictHostKeyChecking=no",
        SERVER_SSH, f"sqlite3 {SERVER_CRAWLER_DB} 'SELECT COUNT(*) FROM crawl_results;'"
    ], capture_output=True, text=True, timeout=15)
    print(f"\n📊 服务器 crawl_results 总数: {result.stdout.strip()}")
    
    # 清理临时文件
    os.remove(tmp_sql)
    subprocess.run([
        "ssh", "-o", "ConnectTimeout=10", "-o", "StrictHostKeyChecking=no",
        SERVER_SSH, "rm -f /tmp/restore_crawler.sql"
    ], capture_output=True, timeout=15)
    
    # 4. 同步到 search.db
    print("\n📤 同步到 search.db...")
    sync_sql = f"""
        ATTACH DATABASE '{SERVER_CRAWLER_DB}' AS crawl;
        INSERT OR IGNORE INTO main.gov_raw
            (site_name, source_url, page_url, title, publish_date, summary, category, visits)
        SELECT
            COALESCE((SELECT name FROM crawl.crawl_sites WHERE id=crawl.crawl_results.site_id), ''),
            url,
            url,
            title,
            publish_date,
            COALESCE(summary, content, ''),
            'gov',
            0
        FROM crawl.crawl_results
        WHERE url NOT IN (SELECT page_url FROM main.gov_raw);
        DETACH crawl;
    """
    b64 = base64.b64encode(sync_sql.encode()).decode()
    result = subprocess.run([
        "ssh", "-o", "ConnectTimeout=15", "-o", "StrictHostKeyChecking=no",
        SERVER_SSH, f"echo {b64} | base64 -d | sqlite3 {SERVER_SEARCH_DB}"
    ], capture_output=True, text=True, timeout=60)
    
    if result.returncode == 0:
        print("   ✅ 已同步到 search.db")
    else:
        print(f"   ⚠️ search.db 同步: {result.stderr[:200]}")
    
    # 验证 search.db
    result = subprocess.run([
        "ssh", "-o", "ConnectTimeout=10", "-o", "StrictHostKeyChecking=no",
        SERVER_SSH, f"sqlite3 {SERVER_SEARCH_DB} 'SELECT COUNT(*) FROM gov_raw;'"
    ], capture_output=True, text=True, timeout=15)
    print(f"📊 服务器 search.db(gov_raw) 总数: {result.stdout.strip()}")

if __name__ == "__main__":
    run()
