#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
海外中国投资项目分析简报 - RSS 采集脚本
用于从白名单媒体源采集相关新闻
"""

import argparse
import feedparser
import json
from datetime import datetime, timedelta
from pathlib import Path
import re
import time
import random


def load_feeds(feed_file: str) -> list:
    """加载媒体源白名单"""
    feeds = []
    with open(feed_file, 'r', encoding='utf-8') as f:
        for line in f:
            line = line.strip()
            if not line or line.startswith('#'):
                continue
            parts = line.split('|')
            if len(parts) >= 4:
                feeds.append({
                    'name': parts[0].strip(),
                    'url': parts[1].strip(),
                    'region': parts[2].strip(),
                    'trust_level': parts[3].strip()
                })
    return feeds


def load_blocklist(blocklist_file: str) -> set:
    """加载媒体黑名单"""
    blocklist = set()
    try:
        with open(blocklist_file, 'r', encoding='utf-8') as f:
            for line in f:
                line = line.strip()
                if not line or line.startswith('#'):
                    continue
                parts = line.split('|')
                if parts:
                    blocklist.add(parts[0].strip().lower())
    except FileNotFoundError:
        pass
    return blocklist


def is_blocked(title: str, blocklist: set) -> bool:
    """检查是否在黑名单中"""
    title_lower = title.lower()
    for blocked in blocklist:
        if blocked in title_lower:
            return True
    return False


def fetch_feed(feed_url: str, hours: int) -> list:
    """抓取单个 RSS 源"""
    entries = []
    try:
        # 随机 User-Agent 避免被屏蔽
        headers = {
            'User-Agent': 'Mozilla/5.0 (Windows NT 10.0; Win64; x64) AppleWebKit/537.36'
        }
        feed = feedparser.parse(feed_url, request_headers=headers)
        
        if not feed.entries:
            return entries
        
        cutoff = datetime.now() - timedelta(hours=hours)
        
        for entry in feed.entries[:30]:  # 限制每个源最多 30 条
            published = None
            if hasattr(entry, 'published_parsed') and entry.published_parsed:
                try:
                    published = datetime(*entry.published_parsed[:6])
                except:
                    pass
            elif hasattr(entry, 'updated_parsed') and entry.updated_parsed:
                try:
                    published = datetime(*entry.updated_parsed[:6])
                except:
                    pass
            
            # 如果没有时间信息，也保留（可能是最新文章）
            if published and published < cutoff:
                continue
            
            title = entry.get('title', '')
            summary = entry.get('summary', '') or entry.get('description', '')
            link = entry.get('link', '')
            
            # 清理 HTML 标签
            summary = re.sub(r'<[^>]+>', '', summary)[:500]
            
            entries.append({
                'title': title,
                'link': link,
                'published': published.isoformat() if published else None,
                'summary': summary,
                'source': feed.feed.get('title', 'Unknown')
            })
    except Exception as e:
        print(f"  ⚠ 抓取失败：{str(e)[:50]}")
    
    return entries


def filter_by_keywords(entries: list, keywords: list) -> list:
    """按关键词过滤"""
    if not keywords:
        return entries
    
    filtered = []
    for entry in entries:
        text = f"{entry['title']} {entry['summary']}"
        for keyword in keywords:
            if keyword and keyword in text:
                filtered.append(entry)
                break
    return filtered


def format_output(entries: list, hours: int) -> str:
    """格式化输出为 Markdown"""
    output = []
    output.append("# 原始采集素材\n")
    output.append(f"**采集时间：** {datetime.now().strftime('%Y-%m-%d %H:%M:%S')}")
    output.append(f"**数据时段：** 过去 {hours} 小时")
    output.append(f"**采集条数：** {len(entries)}\n")
    output.append("---\n")
    
    for i, entry in enumerate(entries, 1):
        output.append(f"## {i}. {entry['title']}\n")
        output.append(f"**来源：** {entry['source']}")
        if entry['published']:
            output.append(f" **发布时间：** {entry['published'][:10]}")
        output.append(f"\n**链接：** {entry['link']}\n")
        if entry['summary']:
            output.append(f"\n**摘要：** {entry['summary']}\n")
        output.append("\n---\n")
    
    return '\n'.join(output)


def main():
    parser = argparse.ArgumentParser(description='RSS 采集脚本')
    parser.add_argument('action', choices=['fetch'], help='操作类型')
    parser.add_argument('--feed-file', required=True, help='媒体源文件路径')
    parser.add_argument('--keywords', required=True, help='关键词列表（逗号分隔）')
    parser.add_argument('--hours', type=int, default=24, help='采集时间范围（小时）')
    parser.add_argument('--output', required=True, help='输出文件路径')
    
    args = parser.parse_args()
    
    # 加载配置
    feeds = load_feeds(args.feed_file)
    blocklist = load_blocklist(args.feed_file.replace('feeds.txt', 'blocklist.txt'))
    keywords = [k.strip() for k in args.keywords.split(',')]
    
    # 添加英文关键词
    en_keywords = ['China', 'Chinese', 'pipeline', 'oil', 'gas', 'mine', 'copper', 
                   'canal', 'investment', 'Myanmar', 'Cambodia', 'Niger', 'Benin',
                   'Belt and Road', ' BRI ']
    all_keywords = keywords + en_keywords
    
    print(f"加载 {len(feeds)} 个媒体源")
    print(f"中文关键词：{keywords}")
    print(f"英文关键词：{en_keywords}")
    print(f"时间范围：{args.hours} 小时\n")
    
    # 采集所有源
    all_entries = []
    success_count = 0
    
    for i, feed in enumerate(feeds):
        print(f"[{i+1}/{len(feeds)}] 采集：{feed['name']}...", end=' ')
        entries = fetch_feed(feed['url'], args.hours)
        
        # 过滤黑名单
        entries = [e for e in entries if not is_blocked(e['title'], blocklist)]
        
        if entries:
            print(f"✓ {len(entries)} 条")
            success_count += 1
        else:
            print("○ 无相关")
        
        all_entries.extend(entries)
        
        # 礼貌延迟，避免请求过快
        time.sleep(random.uniform(0.5, 1.5))
    
    print(f"\n原始采集：{len(all_entries)} 条")
    print(f"成功源数：{success_count}/{len(feeds)}")
    
    # 按关键词过滤
    if all_keywords:
        filtered = filter_by_keywords(all_entries, all_keywords)
        print(f"关键词过滤后：{len(filtered)} 条")
    else:
        filtered = all_entries
        print(f"过滤后：{len(filtered)} 条")
    
    # 去重
    seen = set()
    unique = []
    for entry in filtered:
        if entry['link'] not in seen:
            seen.add(entry['link'])
            unique.append(entry)
    
    print(f"去重后：{len(unique)} 条")
    
    # 输出
    output = format_output(unique, args.hours)
    Path(args.output).parent.mkdir(parents=True, exist_ok=True)
    with open(args.output, 'w', encoding='utf-8') as f:
        f.write(output)
    
    print(f"\n输出：{args.output}")


if __name__ == '__main__':
    main()
