#!/usr/bin/env python3 -u
"""
每日推送：2026年后含"项目"关键词的数据 → Notion数据库
"""

import os, re, sys, json, subprocess, ssl, time
from urllib.request import Request, urlopen
from urllib.error import HTTPError
from concurrent.futures import ThreadPoolExecutor, as_completed

DB_PATH = os.getenv("SEARCH_DB", "/root/search.db")
NOTION_API_KEY = "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"
NOTION_DB_ID = "3a6399b8-3c5b-803d-b92f-da0c63a31a6d"
DATA_SOURCE_ID = "3a6399b8-3c5b-8025-8989-000b93a513e1"
BATCH_SIZE = 50
TARGET_INDUSTRIES = ['cpi']  # 化工加工(CPI) —— 统一英文id写法; [] = 不限行业
KEYWORDS = ["项目"]  # 仅标题含"项目"
MAX_WORKERS = 2
_SSL_CTX = ssl.create_default_context()

# 加载爬虫配置映射表
CONFIG_PATH = os.path.join(os.path.dirname(os.path.abspath(__file__)), "daily_crawl_config.json")
SCRIPT_MAP = {}
if os.path.exists(CONFIG_PATH):
    with open(CONFIG_PATH, "r", encoding="utf-8") as f:
        for entry in json.load(f):
            name = entry.get("name", "")
            script = entry.get("script", "")
            if not name and script:
                name = entry.get("group", "")
            if name:
                if script:
                    SCRIPT_MAP[name] = script.replace('.py', '')
                else:
                    SCRIPT_MAP[name] = name


def notion_api(method, path, data=None, retries=5):
    url = "https://api.notion.com/v1" + path
    body = json.dumps(data).encode() if data else None

    for attempt in range(retries):
        try:
            req = Request(url, data=body, headers={
                "Authorization": "Bearer " + NOTION_API_KEY,
                "Notion-Version": "2022-06-28",
                "Content-Type": "application/json",
            }, method=method)
            resp = urlopen(req, timeout=20, context=_SSL_CTX)
            return json.loads(resp.read())
        except HTTPError as e:
            err = e.read().decode()
            # 429 rate limit: 指数退避重试
            if e.code == 429:
                if attempt < retries - 1:
                    wait = min(2 ** (attempt + 2), 30)
                    time.sleep(wait)
                    continue
                return {"_error": True, "status": e.code, "message": err[:600]}
            # 400类错误(非429)不重试
            if e.code >= 400 and e.code < 500:
                return {"_error": True, "status": e.code, "message": err[:600]}
            if attempt < retries - 1:
                time.sleep(2 ** attempt)
                continue
            return {"_error": True, "status": e.code, "message": err[:600]}
        except Exception as e:
            if attempt < retries - 1:
                time.sleep(2 ** attempt)
                continue
            return {"_error": True, "status": 0, "message": str(e)[:200]}


def get_script_name(site_name):
    """从daily_crawl_config.json匹配站点名称->爬虫标识名"""
    if not site_name:
        return ""
    # 精确匹配
    if site_name in SCRIPT_MAP:
        return SCRIPT_MAP[site_name]
    # 模糊匹配: config name应该包含在site_name中或反之
    best = ""
    best_len = 0
    for cname, label in SCRIPT_MAP.items():
        if cname in site_name or site_name in cname:
            match_len = min(len(cname), len(site_name))
            if match_len > best_len:
                best = label
                best_len = match_len
    # 如果还匹配不上，尝试用地名前缀匹配(如"沂源县"匹配"沂源县-环评审批信息")
    if not best:
        for prefix_len in range(4, 1, -1):
            if len(site_name) >= prefix_len:
                prefix = site_name[:prefix_len]
                for cname, label in SCRIPT_MAP.items():
                    if prefix in cname:
                        return label
    return best


def fetch_records(limit=None):
    """查询符合条件的记录（按行业+关键词）"""
    limit_clause = " LIMIT {}".format(limit) if limit else ""

    # 构建关键词LIKE条件
    kw_conditions = " OR ".join(["title LIKE '%{}%'".format(kw) for kw in KEYWORDS])

    ind_sql = ''
    if TARGET_INDUSTRIES:
        ind_sql = ' AND industry IN ({})'.format(','.join("'{}'".format(i) for i in TARGET_INDUSTRIES))
    sql = """SELECT COALESCE(NULLIF(page_url, ''), source_url) AS page_url,
       title, publish_date, site_name, group_name, industry
FROM gov_raw
WHERE publish_date >= '2026-01-01'
  {ind_sql}
  AND ({kw_conditions})
ORDER BY publish_date DESC""".format(ind_sql=ind_sql, kw_conditions=kw_conditions) + limit_clause

    r = subprocess.run(["sqlite3", "-json", DB_PATH],
                       input=sql, capture_output=True, text=True, timeout=120)
    if r.returncode != 0 or not r.stdout.strip():
        return []
    return json.loads(r.stdout)

def parse_date(raw):
    """将各种中文日期格式转为ISO 8601 (YYYY-MM-DD)"""
    import re
    s = (raw or "").strip()
    m = re.search(r"(\d{4})-(\d{1,2})-(\d{1,2})", s)
    if m:
        return "{}-{:02d}-{:02d}".format(int(m.group(1)), int(m.group(2)), int(m.group(3)))
    m = re.search(r"(\d{4})[年/.]?(\d{1,2})[月/.]?(\d{1,2})", s)
    if m:
        return "{}-{:02d}-{:02d}".format(int(m.group(1)), int(m.group(2)), int(m.group(3)))
    m = re.search(r"(\d{4})", s)
    if m:
        return m.group(1) + "-01-01"
    return None


def push_to_notion(record):
    """推送单条到Notion"""
    title = (record.get("title") or "").strip()
    page_url = (record.get("page_url") or "").strip()
    pub_date = parse_date(record.get("publish_date"))
    site_name = (record.get("site_name") or "").strip()
    script_name = get_script_name(site_name)

    if not title or not page_url:
        return {"error": "missing title or url"}

    props = {
        "标题": {"title": [{"text": {"content": title[:2000]}}]},
        "URL": {"url": page_url},
    }
    if site_name:
        props["站点名称"] = {"rich_text": [{"text": {"content": site_name[:100]}}]}
    if script_name:
        props["脚本名称"] = {"rich_text": [{"text": {"content": script_name[:100]}}]}
    if pub_date and len(pub_date) == 10:
        props["发布日期"] = {"date": {"start": pub_date}}

    data = {
        "parent": {"database_id": NOTION_DB_ID},
        "properties": props,
    }

    result = notion_api("POST", "/pages", data)
    if result.get("_error"):
        return {"url": page_url, "title": title[:50], "error": result.get("message")}
    return {"url": page_url, "title": title[:50], "id": result.get("id")}


def get_pushed_urls():
    """获取已推送的URL列表（去重用）"""
    urls = set()
    cursor = None
    retries = 0
    while True:
        params = {"page_size": 100}
        if cursor:
            params["start_cursor"] = cursor
        result = notion_api("POST", "/databases/" + NOTION_DB_ID + "/query", params)
        if result.get("_error"):
            if "rate_limited" in str(result) and retries < 3:
                retries += 1
                time.sleep(10)
                continue
            retries = 0
            break
        retries = 0
        if not result.get("results"):
            break
        for item in result["results"]:
            props = item.get("properties", {})
            url_prop = props.get("URL", {})
            if url_prop.get("type") == "url":
                val = url_prop.get("url")
                if val:
                    urls.add(val)
        cursor = result.get("next_cursor")
        if not result.get("has_more"):
            break
        time.sleep(1.0)  # 限流间隔
    return urls


def main():
    import argparse
    parser = argparse.ArgumentParser(description="推送项目数据到Notion")
    parser.add_argument("--limit", type=int, default=0, help="限制推送条数(0=全部)")
    parser.add_argument("--skip-dedup", action="store_true", help="跳过去重(用于首次全量)")
    args = parser.parse_args()

    print("[项目推送] Starting...")

    # 去重
    if args.skip_dedup:
        pushed = set()
        print("  跳过去重")
    else:
        print("  查询Notion已有URL (去重)...")
        pushed = get_pushed_urls()
        print("  Notion中已有: {}条".format(len(pushed)))

    # 查询待推送记录
    all_records = fetch_records(limit=args.limit if args.limit else None)
    print("  数据库匹配: {}条".format(len(all_records)))

    # 过滤已推送的
    new_records = [r for r in all_records if r["page_url"] not in pushed]
    total = len(new_records)
    print("  待推送: {}条".format(total))

    if total == 0:
        print("[项目推送] Done: 无新记录")
        return

    # 推送
    print("  写入Notion ({} workers)...".format(MAX_WORKERS))
    ok = fail = 0
    t0 = time.time()

    with ThreadPoolExecutor(max_workers=MAX_WORKERS) as ex:
        fut_map = {ex.submit(push_to_notion, r): r for r in new_records}
        for i, fut in enumerate(as_completed(fut_map), 1):
            r = fut.result()
            if r.get("id"):
                ok += 1
            else:
                fail += 1
                if r.get("error"):
                    print("    FAIL: {} - {}".format(r.get("title", "?"), r["error"][:100]))
            if i % 200 == 0:
                print("    {}/{} ({:.0f}s)".format(i, total, time.time() - t0))

    t = time.time() - t0
    print("\n[项目推送] Done: {}成功, {}失败, {:.1f}s, {:.0f}条/分钟".format(
        ok, fail, t, total / (t / 60) if t > 0 else 0))


if __name__ == "__main__":
    main()
