#!/usr/bin/env python3
# -*- coding: utf-8 -*-
"""
信源采集模块
多源采集官方、媒体、智库信源，提供交叉验证基础
"""

import asyncio
import json
from datetime import datetime
from typing import List, Dict, Any, Optional
from dataclasses import dataclass, field
from urllib.parse import urljoin, urlparse

try:
    from scrapling.fetchers import Fetcher, FetcherSession
    from scrapling.parser import Selector
    SCRAPLING_AVAILABLE = True
except ImportError:
    SCRAPLING_AVAILABLE = False
    print("警告：Scrapling 未安装，部分功能将受限")


@dataclass
class SourceItem:
    """信源条目数据模型"""
    id: str
    url: str
    source_type: str  # official, media, think_tank
    source_name: str
    content: str = ""
    title: str = ""
    timestamp: str = ""
    metadata: Dict[str, Any] = field(default_factory=dict)
    
    def to_dict(self) -> Dict[str, Any]:
        return {
            'id': self.id,
            'url': self.url,
            'source_type': self.source_type,
            'source_name': self.source_name,
            'content': self.content,
            'title': self.title,
            'timestamp': self.timestamp,
            'metadata': self.metadata
        }


@dataclass
class VerificationResult:
    """验证结果"""
    item_id: str
    consistency_score: float  # 0-1
    verified_sources: int
    total_sources: int
    conflicts: List[str] = field(default_factory=list)
    confidence: str = ""  # low, medium, high


class SourceCollector:
    """信源采集器"""
    
    def __init__(self):
        self.sources: Dict[str, Dict] = {}
        self.collected_items: List[SourceItem] = []
        self.session = None
        
        if SCRAPLING_AVAILABLE:
            self.session = FetcherSession(impersonate='chrome')
    
    def load_sources_config(self, config_path: str = None):
        """加载信源配置"""
        if config_path is None:
            # 使用内置配置
            self._load_builtin_sources()
        else:
            with open(config_path, 'r', encoding='utf-8') as f:
                self.sources = json.load(f)
    
    def _load_builtin_sources(self):
        """加载内置信源配置"""
        self.sources = {
            'official': {
                'mod_gov_cn': {
                    'name': '中国国防部',
                    'url': 'http://www.mod.gov.cn/',
                    'type': 'official',
                    'priority': 1,
                    'css_selectors': {
                        'title': 'h1.title, h2.title, .news-title',
                        'content': '.news-content, article.content',
                        'timestamp': '.news-time, .time, .date'
                    }
                },
                'fmprc_gov_cn': {
                    'name': '中国外交部',
                    'url': 'https://www.fmprc.gov.cn/',
                    'type': 'official',
                    'priority': 1,
                    'css_selectors': {
                        'title': 'h1, h2, .title',
                        'content': '.content, article, .article-content',
                        'timestamp': '.time, .date, .publish-time'
                    }
                },
                'gwytb_gov_cn': {
                    'name': '国台办',
                    'url': 'http://www.gwytb.gov.cn/',
                    'type': 'official',
                    'priority': 1,
                    'css_selectors': {
                        'title': 'h1, h2, .title',
                        'content': '.content, article',
                        'timestamp': '.time, .date'
                    }
                }
            },
            'media': {
                'xinhua': {
                    'name': '新华社',
                    'url': 'http://www.xinhuanet.com/',
                    'type': 'media',
                    'priority': 2,
                    'css_selectors': {
                        'title': 'h1.title, h2.title',
                        'content': '.article-content, .content',
                        'timestamp': '.time, .pubtime'
                    }
                },
                'cctv': {
                    'name': '央视',
                    'url': 'https://news.cctv.com/',
                    'type': 'media',
                    'priority': 2,
                    'css_selectors': {
                        'title': 'h1, h2',
                        'content': '.cnt_bd, .article-content',
                        'timestamp': '.time, .date'
                    }
                }
            },
            'think_tank': {
                'cass_taiwan': {
                    'name': '中国社会科学院台湾研究所',
                    'url': 'http://www.taiwanstudies.cn/',
                    'type': 'think_tank',
                    'priority': 3,
                    'css_selectors': {
                        'title': 'h1, h2',
                        'content': '.content, article',
                        'timestamp': '.time, .date'
                    }
                }
            }
        }
    
    def collect_from_source(self, source_id: str, keywords: List[str] = None, 
                        max_items: int = 10) -> List[SourceItem]:
        """从指定信源采集数据"""
        if source_id not in sum([list(s.keys()) for s in self.sources.values()], []):
            print(f"错误：信源 {source_id} 不存在")
            return []
        
        # 查找信源配置
        source_config = None
        for category in self.sources.values():
            if source_id in category:
                source_config = category[source_id]
                break
        
        if not source_config:
            return []
        
        print(f'正在从 {source_config["name"]} 采集数据...')
        
        if not SCRAPLING_AVAILABLE:
            print("错误：Scrapling 不可用，无法采集")
            return []
        
        items = []
        try:
            # 使用 Fetcher 获取页面
            page = Fetcher.get(source_config['url'], stealthy_headers=True)
            
            # 提取内容
            title = page.css(source_config['css_selectors']['title'] + '::text').get('').strip()
            content = page.css(source_config['css_selectors']['content']).get('').strip()
            timestamp = page.css(source_config['css_selectors']['timestamp'] + '::text').get('').strip()
            
            if keywords:
                # 关键词过滤
                keyword_matches = []
                for keyword in keywords:
                    if keyword in title or keyword in content:
                        keyword_matches.append(keyword)
                
                if not keyword_matches:
                    return []
                
                print(f"找到关键词匹配: {', '.join(keyword_matches)}")
            
            item = SourceItem(
                id=f"{source_id}_{datetime.now().strftime('%Y%m%d%H%M%S')}",
                url=source_config['url'],
                source_type=source_config['type'],
                source_name=source_config['name'],
                content=content,
                title=title,
                timestamp=timestamp or datetime.now().isoformat(),
                metadata={
                    'collected_at': datetime.now().isoformat(),
                    'keywords': keywords or []
                }
            )
            
            items.append(item)
            print(f"采集成功: {title[:50]}...")
            
        except Exception as e:
            print(f"采集失败: {str(e)}")
        
        self.collected_items.extend(items)
        return items
    
    def collect_from_all(self, categories: List[str] = None, 
                      keywords: List[str] = None,
                      max_items: int = 10) -> List[SourceItem]:
        """从多个信源类别采集数据"""
        if categories is None:
            categories = ['official', 'media', 'think_tank']
        
        all_items = []
        for category in categories:
            if category not in self.sources:
                continue
            
            for source_id in self.sources[category]:
                items = self.collect_from_source(source_id, keywords, max_items)
                all_items.extend(items)
        
        print(f"总共采集到 {len(all_items)} 条信息")
        return all_items
    
    def verify_sources(self, items: List[SourceItem]) -> List[VerificationResult]:
        """多源交叉验证"""
        if len(items) < 2:
            return []
        
        print("正在进行多源交叉验证...")
        
        results = []
        # 简化验证逻辑：检查内容相似度
        for i, item in enumerate(items):
            similar_count = 0
            conflicts = []
            
            for j, other_item in enumerate(items):
                if i == j:
                    continue
                
                # 简单的内容重叠度计算
                if item.title and other_item.title:
                    title_overlap = len(set(item.title) & set(other_item.title))
                    if title_overlap > 5:  # 阈值
                        similar_count += 1
                else:
                    if item.source_type != other_item.source_type:
                        conflicts.append(other_item.source_name)
            
            consistency_score = min(1.0, similar_count / (len(items) - 1))
            
            # 置信度评级
            if consistency_score >= 0.7:
                confidence = "high"
            elif consistency_score >= 0.4:
                confidence = "medium"
            else:
                confidence = "low"
            
            result = VerificationResult(
                item_id=item.id,
                consistency_score=round(consistency_score, 3),
                verified_sources=similar_count + 1,
                total_sources=len(items),
                conflicts=conflicts,
                confidence=confidence
            )
            
            results.append(result)
        
        return results
    
    def export_to_json(self, filepath: str):
        """导出为JSON"""
        data = {
            'collected_at': datetime.now().isoformat(),
            'total_items': len(self.collected_items),
            'items': [item.to_dict() for item in self.collected_items]
        }
        
        with open(filepath, 'w', encoding='utf-8') as f:
            json.dump(data, f, ensure_ascii=False, indent=2)
        
        print(f"已导出到: {filepath}")
    
    def export_to_markdown(self, filepath: str):
        """导出为Markdown"""
        with open(filepath, 'w', encoding='utf-8') as f:
            f.write(f"# 信源采集报告\n\n")
            f.write(f"**采集时间**: {datetime.now().isoformat()}\n")
            f.write(f"**总条目数**: {len(self.collected_items)}\n\n")
            f.write("---\n\n")
            
            for item in self.collected_items:
                f.write(f"## {item.title}\n\n")
                f.write(f"**来源**: {item.source_name}\n")
                f.write(f"**类型**: {item.source_type}\n")
                f.write(f"**时间**: {item.timestamp}\n")
                f.write(f"**URL**: {item.url}\n\n")
                f.write(f"**内容摘要**:\n")
                f.write(f"{item.content[:500]}...\n\n")
                f.write("---\n\n")
        
        print(f"已导出到: {filepath}")


def main():
    """测试主函数"""
    collector = SourceCollector()
    collector.load_sources_config()
    
    # 采集测试
    print("=== 信源采集测试 ===\n")
    
    # 从单个信源采集
    items = collector.collect_from_source('mod_gov_cn', keywords=['演习', '台湾'])
    
    # 多源交叉验证
    if items:
        results = collector.verify_sources(items)
        for result in results:
            print(f"验证结果: {result.item_id} - 一致性: {result.consistency_score} - 置信度: {result.confidence}")
        
        collector.export_to_json('collected_data.json')
        collector.export_to_markdown('collected_data.md')


if __name__ == '__main__':
    main()
