#!/usr/bin/env python3
"""Batch crawler scheduler — V3 连续运行+自动重试
- 持续跑，不等30分钟cron间隔，一批接一批
- 失败脚本自动入重试队列，全部跑完一轮后立即重试
- 锁文件防止cron重复启动
- 状态持久化到文件，崩溃后cron可续跑
"""
import json, os, sys, subprocess, time, re
from datetime import datetime, date
from concurrent.futures import ThreadPoolExecutor, as_completed

DB_PATH = "/root/search.db"
STATE_FILE = "/tmp/crawl_batch_state.json"
RESULTS_FILE = "/tmp/crawl_results.json"
RUN_STATUS_FILE = "/root/gov_crawler/run_status.json"
CONFIG_FILE = "/root/gov_crawler/daily_crawl_config.json"
RETRY_FILE = "/tmp/crawl_retry_queue.json"
LOCK_FILE = "/tmp/crawl_batch.lock"
LOCK_STALE_SECS = 600  # 锁超过10分钟视为失效

BATCH_SIZE = 60
MAX_WORKERS = 3
PER_SCRIPT_TIMEOUT = 60
BATCH_TIMEOUT = 600      # 10分钟单批次安全阀
MAX_RETRY_ROUNDS = 2     # 失败脚本最多重试2轮

# 重试时跳过这些错误类型（肯定救不回来的）
SKIP_RETRY_TYPES = {"connection_refused", "dns_resolve", "unreachable", "read_timeout"}


# ── 锁 ──

def acquire_lock():
    if os.path.exists(LOCK_FILE):
        try:
            with open(LOCK_FILE) as f:
                lock_pid = f.read().strip()
            if lock_pid:
                try:
                    os.kill(int(lock_pid), 0)  # 发信号0检查进程存活
                    print(f"[{datetime.now().strftime('%H:%M:%S')}] 已有实例 (PID {lock_pid})，跳过")
                    return False
                except (OSError, ProcessLookupError):
                    print(f"[{datetime.now().strftime('%H:%M:%S')}] 锁文件残留 (PID {lock_pid}已死)，移除...")
            os.remove(LOCK_FILE)
        except Exception:
            pass
    with open(LOCK_FILE, "w") as f:
        f.write(str(os.getpid()))
    return True


def release_lock():
    if os.path.exists(LOCK_FILE):
        try:
            os.remove(LOCK_FILE)
        except Exception:
            pass


# ── 状态 ──

def load_state():
    today = str(date.today())
    if os.path.exists(STATE_FILE):
        try:
            with open(STATE_FILE) as f:
                state = json.load(f)
            if state.get("date") == today:
                return state
        except (json.JSONDecodeError, KeyError):
            pass
    return {
        "date": today,
        "batch_size": BATCH_SIZE,
        "current_index": 0,
        "total": 0,
        "results": {"success": [], "failed": []},
        "completed": False,
        "has_new_data": False,
        "retry_round": 0,   # 当前重试轮次
        "retry_done": False, # 重试是否全部完成
    }


def save_state(state):
    with open(STATE_FILE, "w") as f:
        json.dump(state, f, ensure_ascii=False, indent=2)


# ── 重试队列 ──

def load_retry_queue():
    """Load retry queue: list of config entry dicts that failed."""
    if os.path.exists(RETRY_FILE):
        try:
            with open(RETRY_FILE) as f:
                return json.load(f)
        except (json.JSONDecodeError, Exception):
            pass
    return []


def save_retry_queue(queue):
    with open(RETRY_FILE, "w") as f:
        json.dump(queue, f, ensure_ascii=False, indent=2)


# ── 配置 ──

def load_config():
    with open(CONFIG_FILE) as f:
        cfg = json.load(f)["crawlers"]
    return [c for c in cfg if c.get("enabled", False)]


def build_script_index(enabled):
    """Build {script_name: config_entry} lookup for retry matching."""
    return {c.get("script", ""): c for c in enabled if c.get("script")}


# ── 结果持久化 ──

def load_run_status():
    if os.path.exists(RUN_STATUS_FILE):
        try:
            with open(RUN_STATUS_FILE) as f:
                return json.load(f)
        except (json.JSONDecodeError, Exception):
            pass
    return {}


def save_run_status(status_map):
    try:
        with open(RUN_STATUS_FILE, "w") as f:
            json.dump(status_map, f, ensure_ascii=False, indent=2)
    except Exception as e:
        print(f"  [WARN] 保存run_status失败: {e}")


def update_run_status(status_map, script_key, name_key, status_str, elapsed, error_type=""):
    ts = time.time()
    entry = {"status": status_str, "ts": ts, "elapsed": elapsed, "error_type": error_type}
    if script_key:
        status_map[script_key] = entry
    if name_key and name_key != script_key:
        status_map[name_key] = dict(entry)


def save_results(batch_results):
    try:
        today = str(date.today())
        existing = []
        if os.path.exists(RESULTS_FILE):
            try:
                with open(RESULTS_FILE) as f:
                    existing = json.load(f)
                if isinstance(existing, list):
                    existing = [r for r in existing if r.get("date") == today]
                else:
                    existing = []
            except (json.JSONDecodeError, Exception):
                existing = []
        existing.extend(batch_results)
        with open(RESULTS_FILE, "w") as f:
            json.dump(existing, f, ensure_ascii=False, indent=2)
    except Exception as e:
        print(f"  [WARN] 保存批次结果失败: {e}")


# ── 辅助 ──

def extract_new_count(stdout, stderr):
    """Extract new record count from script output.
    Only matches explicit patterns with 条 (records/tiao).
    Capped at 100000 to prevent regex false positives."""
    text = stdout + "\n" + stderr
    for line in text.split("\n"):
        # 只匹配明确带"条"的输出格式
        for pat in [
            r"新增\s*[:：]?\s*(\d+)\s*条",      # 新增: N条 / 新增 N 条
            r"本次新增[了]?\s*(\d+)\s*条",      # 本次新增N条
            r"成功入库\s*(\d+)\s*条",           # 成功入库 N 条
            r"\+(\d+)\s*条",                   # +N条
        ]:
            m = re.search(pat, line)
            if m:
                cnt = int(m.group(1))
                # 安全上限：单次运行不可能超过10万条
                if 0 < cnt < 100000:
                    return cnt
                return 0
    return 0


def classify_error(msg, stdout=""):
    full = (msg + "\n" + stdout).lower()
    if "timeout" in full:
        return "timeout"
    if any(x in full for x in ["connection refused", "connection reset", "connect failed",
                                "连接被拒绝"]):
        return "connection_refused"
    if any(x in full for x in ["name or service not known", "temporary failure in name resolution",
                                "无法解析", "dns", "resolve failed"]):
        return "dns_resolve"
    if any(x in full for x in ["no route to host", "network is unreachable", "无法访问",
                                "网络不可达", "连接超时"]):
        return "unreachable"
    if "traceback" in full:
        return "script_error"
    if "read timed out" in full:
        return "read_timeout"
    return "exit_error"


def run_single_script(c):
    """Run one crawler script. Thread-safe."""
    name = c.get("display_name", c.get("site_name", c.get("name", "Unknown")))
    script = c["script"]
    args_raw = c.get("args", "")
    if isinstance(args_raw, list):
        args = args_raw
    else:
        args = [a for a in args_raw.strip().split() if a]

    cmd = ["python3", script]
    if args:
        cmd.extend(args)

    script_timeout = c.get("timeout", PER_SCRIPT_TIMEOUT)

    result = {
        "script": script,
        "name": name,
        "ts": 0,
        "date": str(date.today()),
        "has_error": True,
        "rc": -1,
        "new_count": 0,
        "elapsed": 0,
        "error_type": "",
        "error": "",
        "timeout": False,
    }

    t0 = time.time()
    try:
        r = subprocess.run(
            cmd, capture_output=True, text=True,
            timeout=script_timeout, cwd="/root/gov_crawler",
        )
        elapsed = time.time() - t0
        result["ts"] = time.time()
        result["rc"] = r.returncode
        result["elapsed"] = round(elapsed, 1)

        if r.returncode == 0:
            new_count = extract_new_count(r.stdout, r.stderr)
            result["has_error"] = False
            result["new_count"] = new_count
            result["status"] = f"成功 +{new_count}条"
        else:
            err_text = (r.stderr or "")[:200] or f"exit code {r.returncode}"
            result["error"] = err_text
            error_type = classify_error(r.stderr, r.stdout)
            result["error_type"] = error_type
            if error_type == "timeout":
                result["status"] = "超时"
                result["timeout"] = True
            elif error_type in ("connection_refused", "dns_resolve", "unreachable", "read_timeout"):
                result["status"] = "站点不可达"
            elif error_type == "script_error":
                result["status"] = "脚本异常"
            else:
                result["status"] = f"失败 rc={r.returncode}"

    except subprocess.TimeoutExpired:
        elapsed = time.time() - t0
        result["ts"] = time.time()
        result["elapsed"] = round(elapsed, 1)
        result["error"] = "timeout"
        result["error_type"] = "timeout"
        result["timeout"] = True
        result["status"] = "超时"
    except Exception as e:
        elapsed = time.time() - t0
        result["ts"] = time.time()
        result["elapsed"] = round(elapsed, 1)
        result["error"] = str(e)[:200]
        result["error_type"] = "exception"
        result["status"] = f"失败 异常"

    return result


def process_batch_results(state, batch_results, batch_has_new, run_status_map):
    """Common result processing for both first-pass and retry batches."""
    s = state["results"]
    ok_cnt = len(s["success"])
    fail_cnt = len(s["failed"])
    new_total = sum(item.get("new", 0) for item in s["success"])

    save_results(batch_results)
    save_run_status(run_status_map)

    if batch_has_new:
        print("  检测到新数据, 重启 search-app...")
        restart_search_app()

    return ok_cnt, fail_cnt, new_total


# ── 批次执行（首轮：按序跑所有脚本）──

def run_batch(state, enabled):
    """Run one batch from current_index onward."""
    start = state["current_index"]
    page = start // state["batch_size"] + 1
    end = min(start + state["batch_size"], len(enabled))
    batch = enabled[start:end]

    if not batch:
        return

    t_start = time.time()
    batch_has_new = False
    batch_results = []
    run_status_map = load_run_status()

    print(f"\n{'='*60}")
    print(f"  首轮 批次 {page} | {start+1}-{end}/{len(enabled)} | {len(batch)} 个爬虫")
    print(f"  时间: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")
    print(f"{'='*60}")

    with ThreadPoolExecutor(max_workers=MAX_WORKERS) as executor:
        future_map = {executor.submit(run_single_script, c): c for c in batch}

        completed = 0
        for future in as_completed(future_map):
            c = future_map[future]
            try:
                result = future.result()
            except Exception as e:
                result = {
                    "script": c.get("script", ""),
                    "name": c.get("name", "Unknown"),
                    "status": f"失败 thread_error",
                    "has_error": True, "new_count": 0, "elapsed": 0,
                    "error_type": "exception", "error": str(e)[:200],
                    "timeout": False,
                }

            completed += 1
            seq = start + completed
            print(f"  [{seq}/{len(enabled)}] {result['name']}: {result['status']} ({result['elapsed']:.0f}s)")

            state["current_index"] = seq
            save_state(state)

            if not result["has_error"]:
                state["results"]["success"].append({"name": result["name"], "new": result["new_count"]})
                if result["new_count"] > 0:
                    batch_has_new = True
                    state["has_new_data"] = True
            else:
                state["results"]["failed"].append({
                    "name": result["name"],
                    "script": result["script"],
                    "error": result.get("error", "")[:200],
                    "error_type": result.get("error_type", ""),
                })

            batch_results.append(result)

            update_run_status(run_status_map,
                              result.get("script", ""), result.get("name", ""),
                              result["status"], result["elapsed"], result.get("error_type", ""))

            if completed % 10 == 0:
                save_run_status(run_status_map)

            if time.time() - t_start > BATCH_TIMEOUT:
                remaining = len(batch) - completed
                print(f"\n  [安全阀] 批次超时 {BATCH_TIMEOUT}s, {remaining} 个留待下批")
                for f in future_map:
                    if not f.done():
                        f.cancel()
                break

    ok_cnt, fail_cnt, new_total = process_batch_results(state, batch_results, batch_has_new, run_status_map)
    elapsed = time.time() - t_start
    print(f"  ─── 批次 {page} 小结: {elapsed:.0f}s | ✓{ok_cnt} ✗{fail_cnt} | 进度 {state['current_index']}/{len(enabled)}")

    if end >= len(enabled):
        state["completed"] = True
        print(f"\n{'='*60}")
        print(f"  ✅ 首轮全部完成！{len(enabled)} 个爬虫已处理")
        print(f"  新增总计: {new_total} 条")
        print(f"{'='*60}")

    save_state(state)


# ── 重试批次：只跑失败脚本 ──

def run_retry_batch(state, failed_scripts, retry_round):
    """Run a batch of retry scripts (failed ones from first pass)."""
    state["retry_round"] = retry_round
    save_state(state)

    if not failed_scripts:
        print("  无失败脚本需要重试")
        state["retry_done"] = True
        save_state(state)
        return 0, 0

    t_start = time.time()
    batch_has_new = False
    batch_results = []
    run_status_map = load_run_status()
    total = len(failed_scripts)

    print(f"\n{'='*60}")
    print(f"  重试第 {retry_round} 轮 | {total} 个失败脚本 | 并发 {MAX_WORKERS}")
    print(f"  时间: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")
    print(f"{'='*60}")

    with ThreadPoolExecutor(max_workers=MAX_WORKERS) as executor:
        future_map = {executor.submit(run_single_script, c): c for c in failed_scripts}

        completed = 0
        still_failed = []
        for future in as_completed(future_map):
            c = future_map[future]
            try:
                result = future.result()
            except Exception:
                result = {
                    "script": c.get("script", ""), "name": c.get("name", "Unknown"),
                    "status": "失败 thread_error", "has_error": True,
                    "new_count": 0, "elapsed": 0, "error_type": "exception",
                    "error": "thread error", "timeout": False,
                }

            completed += 1
            print(f"  [{completed}/{total}] {result['name']}: {result['status']} ({result['elapsed']:.0f}s)")

            if not result["has_error"]:
                state["results"]["success"].append({"name": result["name"], "new": result["new_count"]})
                if result["new_count"] > 0:
                    batch_has_new = True
                    state["has_new_data"] = True
            else:
                still_failed.append({
                    "name": result["name"],
                    "script": result["script"],
                    "error": result.get("error", "")[:200],
                })

            batch_results.append(result)

            update_run_status(run_status_map,
                              result.get("script", ""), result.get("name", ""),
                              result["status"], result["elapsed"], result.get("error_type", ""))

            if completed % 10 == 0:
                save_run_status(run_status_map)

    process_batch_results(state, batch_results, batch_has_new, run_status_map)
    elapsed = time.time() - t_start

    recovered = total - len(still_failed)
    print(f"\n  ─── 重试第 {retry_round} 轮小结: {elapsed:.0f}s | ✓{recovered}(恢复) ✗{len(still_failed)}(仍失败)")

    if still_failed:
        # 还有失败的，记录到重试队列
        save_retry_queue(still_failed)
        print(f"  {len(still_failed)} 个脚本仍需修复后手动重试:")
        for f in still_failed[:10]:
            print(f"    - {f['name']}: {f['error'][:80]}")
        if len(still_failed) > 10:
            print(f"    ... 还有 {len(still_failed)-10} 个")
    else:
        state["retry_done"] = True
        save_retry_queue([])
        print("  ✅ 全部重试成功！")

    return recovered, len(still_failed)


def restart_search_app():
    try:
        subprocess.run(
            ["systemctl", "restart", "search_app.service"],
            capture_output=True, timeout=10,
        )
        print("  [服务] search-app 已重启")
    except Exception as e:
        print(f"  [服务] 重启失败: {e}")


# ── 主循环 ──

def main():
    if not acquire_lock():
        return

    try:
        enabled = load_config()
        script_index = build_script_index(enabled)
        state = load_state()
        state["total"] = len(enabled)

        # 如果状态显示今天已完成第一轮，检查是否有重试队列
        if state.get("completed") and not state.get("retry_done"):
            # 加载重试队列
            retry_entries = load_retry_queue()
            if not retry_entries:
                # 从state.results.failed重建
                failed = state["results"].get("failed", [])
                retry_entries = []
                for f in failed:
                    script = f.get("script", "")
                    entry = script_index.get(script)
                    if entry:
                        retry_entries.append(entry)
                save_retry_queue(retry_entries)

            if retry_entries:
                retry_round = state.get("retry_round", 0) + 1
                if retry_round <= MAX_RETRY_ROUNDS:
                    run_retry_batch(state, retry_entries, retry_round)
                else:
                    print(f"已达最大重试次数 {MAX_RETRY_ROUNDS}，{len(retry_entries)} 个脚本仍失败")
                    state["retry_done"] = True
                    save_state(state)
            else:
                state["retry_done"] = True
                save_state(state)
                print(f"[{datetime.now().strftime('%H:%M:%S')}] 今日全部完成，无重试需要")
            return

        if state.get("completed") and state.get("retry_done"):
            print(f"[{datetime.now().strftime('%H:%M:%S')}] 今日全部完成，跳过。")
            return

        # === Phase 1: 首轮跑完所有脚本 ===
        print(f"\n{'#'*60}")
        print(f"#  开始首轮爬取: 共 {len(enabled)} 个脚本")
        print(f"#  并发 {MAX_WORKERS} worker | 每批 {BATCH_SIZE} 个")
        print(f"{'#'*60}")

        while not state.get("completed"):
            run_batch(state, enabled)

            if state["current_index"] >= len(enabled):
                state["completed"] = True
                save_state(state)
                break

        # === Phase 2: 重试失败的脚本 ===
        failed = state["results"].get("failed", [])
        if not failed:
            print(f"\n[{datetime.now().strftime('%H:%M:%S')}] 全部脚本运行成功，无重试需要。")
            state["retry_done"] = True
            save_state(state)
            return

        # 重建重试队列（从config条目匹配，跳过不可达类）
        retry_entries = []
        skip_count = 0
        for f in failed:
            error_type = f.get("error_type", "")
            script = f.get("script", "")
            entry = script_index.get(script)
            if not entry:
                continue
            if error_type in SKIP_RETRY_TYPES:
                skip_count += 1
                continue
            retry_entries.append(entry)

        if skip_count > 0:
            print(f"  跳过 {skip_count} 个不可达脚本（不重试）")
        if not retry_entries:
            print("无可重试的有效条目，退出。")
            state["retry_done"] = True
            save_state(state)
            return

        save_retry_queue(retry_entries)

        for retry_round in range(1, MAX_RETRY_ROUNDS + 1):
            already_retried = run_retry_batch(state, retry_entries, retry_round)
            recovered, still_fail = already_retried  # tuple

            if still_fail == 0:
                state["retry_done"] = True
                save_state(state)
                break

            # 还有失败的，用最新重试队列继续
            retry_entries = load_retry_queue()
        else:
            print(f"\n[{datetime.now().strftime('%H:%M:%S')}] 重试 {MAX_RETRY_ROUNDS} 轮结束，仍有失败脚本。")
            state["retry_done"] = True
            save_state(state)

        # === Final summary ===
        r = load_run_status()
        total_status = len(r)
        ok_count = sum(1 for v in r.values() if v.get("status", "").startswith("成功"))
        fail_count = total_status - ok_count

        print(f"\n{'='*60}")
        print(f"  📊 今日爬取总结")
        print(f"  总脚本: {len(enabled)}")
        print(f"  有状态记录: {total_status}")
        print(f"  成功: {ok_count}")
        print(f"  失败/异常: {fail_count}")
        new_total = sum(item.get("new", 0) for item in state["results"]["success"])
        print(f"  新增数据: {new_total} 条")
        print(f"  时间: {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")
        print(f"{'='*60}")

    finally:
        release_lock()


if __name__ == "__main__":
    main()
