#!/usr/bin/env python3
"""藤县人民政府 - 两个栏目爬虫（并发版）
1. 建设项目环境影响评价审批
2. 通知公告
"""

import re
import time
from datetime import datetime, date
from urllib.parse import urljoin
from concurrent.futures import ThreadPoolExecutor, as_completed
import requests
import sqlite3
import os

BASE_URL = "http://www.tengxian.gov.cn"
CUTOFF_DATE = date(2023, 6, 17)
DB_PATH = os.environ.get("DB_PATH", "/root/search.db")
MAX_WORKERS = 5  # 并发抓取详情页

COLUMNS = [
    {
        "name": "藤县-环评审批",
        "list_url": "http://www.tengxian.gov.cn/zfxxgk/xxgk/shgysy/hjbh/jsxmhp/index.shtml",
        "list_base": "http://www.tengxian.gov.cn/zfxxgk/xxgk/shgysy/hjbh/jsxmhp/",
        "max_pages": 25,
    },
    {
        "name": "藤县-通知公告",
        "list_url": "http://www.tengxian.gov.cn/txdt/tzgg/index.shtml",
        "list_base": "http://www.tengxian.gov.cn/txdt/tzgg/",
        "max_pages": 40,
    },
]

HEADERS = {"User-Agent": "Mozilla/5.0"}


def fetch(url, retries=3):
    for i in range(retries):
        try:
            resp = requests.get(url, timeout=30, headers=HEADERS)
            resp.encoding = "utf-8"
            if resp.status_code == 200:
                return resp.text
        except Exception as e:
            if i < retries - 1:
                time.sleep(2)
    return None


def extract_list_items(html, base_url):
    items = []
    if not html:
        return items
    pattern = r'<li><a href="([^"]+)"[^>]*title="([^"]*)"[^>]*>.*?</a><span>(\d{4}-\d{2}-\d{2})</span></li>'
    for m in re.finditer(pattern, html, re.DOTALL):
        href = m.group(1)
        title = m.group(2).strip()
        pub_date = m.group(3)
        detail_url = urljoin(base_url, href)
        items.append({"title": title, "url": detail_url, "pub_date": pub_date})
    return items


def get_page_count(html):
    m = re.search(r'createPageHTML\s*\(\s*(\d+)', html)
    return int(m.group(1)) if m else 1


def fetch_detail(detail_url):
    """获取单条详情，返回 (url, title, content, pub_date)"""
    html = fetch(detail_url)
    if not html:
        return (detail_url, None, None, None)

    title = ""
    content = ""
    pub_date = ""

    m = re.search(r'<meta\s+name="ArticleTitle"\s+content="([^"]*)"', html)
    if m:
        title = m.group(1).strip()
    m = re.search(r'<meta\s+name="PubDate"\s+content="([^"]*)"', html)
    if m:
        pub_date = m.group(1).strip()[:10]
    m = re.search(r'<div\s+class="trs_editor_view[^"]*"[^>]*>(.*?)</div>\s*', html, re.DOTALL)
    if m:
        content = m.group(1).strip()
    if not content:
        m2 = re.search(r'<div\s+class="article-con"[^>]*>(.*?)</div>\s*</div>', html, re.DOTALL)
        if m2:
            content = m2.group(1).strip()
    if not content:
        content = title

    return (detail_url, title, content, pub_date)


def crawl_column(col_config):
    site_name = col_config["name"]
    list_base = col_config["list_base"]
    max_pages = col_config["max_pages"]

    print("\n====== {} ======".format(site_name), flush=True)

    html = fetch(col_config["list_url"])
    if not html:
        print("  无法获取首页", flush=True)
        return 0, 0

    total_pages = get_page_count(html)
    if total_pages > max_pages:
        total_pages = max_pages
    print("  总页数: {}".format(total_pages), flush=True)

    # 收集所有列表页文章
    all_items = []
    for page in range(total_pages):
        page_url = col_config["list_url"] if page == 0 else "{}index_{}.shtml".format(list_base, page)
        print("  列表页 {}/{}".format(page + 1, total_pages), flush=True)
        page_html = html if page == 0 else fetch(page_url)
        if not page_html:
            continue
        items = extract_list_items(page_html, list_base)
        filtered = [it for it in items if datetime.strptime(it["pub_date"], "%Y-%m-%d").date() >= CUTOFF_DATE]
        all_items.extend(filtered)
        print("    本页{}条，近3年{}条".format(len(items), len(filtered)), flush=True)

        # 如果本页末尾已经超出范围，停止
        if items and not filtered and page > 0:
            print("    本页无近3年数据，停止翻页", flush=True)
            break

    print("  列表总计: {} 条（近3年）".format(len(all_items)), flush=True)
    if not all_items:
        return 0, 0

    # 并发获取详情页
    conn = sqlite3.connect(DB_PATH, timeout=60)
    cursor = conn.cursor()
    total_new = 0

    urls_to_fetch = [it["url"] for it in all_items]
    done = 0

    with ThreadPoolExecutor(max_workers=MAX_WORKERS) as executor:
        fut_map = {executor.submit(fetch_detail, url): url for url in urls_to_fetch}
        for fut in as_completed(fut_map):
            url, title, content, pub_date = fut.result()
            done += 1
            if done % 50 == 0:
                print("  详情 {}/{}...".format(done, len(all_items)), flush=True)

            if title is None:
                continue

            # 从列表项补充信息
            list_item = [it for it in all_items if it["url"] == url]
            if list_item:
                if not title:
                    title = list_item[0]["title"]
                if not pub_date:
                    pub_date = list_item[0]["pub_date"]

            if not content:
                content = title

            try:
                cursor.execute("""
                    INSERT OR IGNORE INTO gov_raw (site_name, title, page_url, content, publish_date, summary, date_rank)
                    VALUES (?, ?, ?, ?, ?, ?, ?)
                """, (site_name, title, url, content, pub_date, content,
                      int(datetime.strptime(pub_date, "%Y-%m-%d").timestamp())))
                if cursor.rowcount > 0:
                    total_new += 1
            except Exception as e:
                pass

            # 每100条提交一次
            if done % 100 == 0:
                conn.commit()

    conn.commit()
    conn.close()

    print("  结果: 新增{}".format(total_new), flush=True)
    return len(all_items), total_new


def main():
    print("=== 藤县人民政府 爬虫 ===", flush=True)
    print("数据库: {}".format(DB_PATH), flush=True)

    total_all = 0
    total_new_all = 0

    for col in COLUMNS:
        fetched, new = crawl_column(col)
        total_all += fetched
        total_new_all += new

    print("\n=== 全部完成 ===", flush=True)
    print("总抓取(近3年): {}".format(total_all), flush=True)
    print("总新增: {}".format(total_new_all), flush=True)


if __name__ == "__main__":
    main()
