""" PostgreSQL数据库连接和操作模块 提供数据库连接池、CRUD操作、数据清理等功能 """ import psycopg2 from psycopg2 import pool, extras from psycopg2.extras import RealDictCursor from contextlib import contextmanager from typing import List, Dict, Any, Optional, Generator from datetime import datetime, timedelta import threading from dataclasses import asdict from .models import Announcement, AnnouncementSource, AnnouncementType, CrawlResult, CrawlStatus from .config_manager import get_config from .logger import get_logger from .reliability import retry_on_exception, RetryConfig logger = get_logger(__name__) class DatabaseConnectionPool: """数据库连接池管理器""" _instance = None _pool = None _lock = threading.Lock() def __new__(cls): if cls._instance is None: with cls._lock: if cls._instance is None: cls._instance = super().__new__(cls) return cls._instance def __init__(self): if self._pool is None: self._pool = None self._config = None def init_pool(self, config): """ 初始化连接池 Args: config: 数据库配置 """ if self._pool is not None: return try: self._config = config self._pool = psycopg2.pool.SimpleConnectionPool( minconn=config.pool_size, maxconn=config.pool_size + config.max_overflow, host=config.host, port=config.port, database=config.name, user=config.user, password=config.password, connect_timeout=config.pool_timeout ) logger.info("数据库连接池初始化成功") except Exception as e: logger.error(f"数据库连接池初始化失败: {str(e)}") raise def get_connection(self): """ 获取数据库连接 Returns: 数据库连接对象 Raises: Exception: 获取连接失败 """ if self._pool is None: raise Exception("数据库连接池未初始化") try: conn = self._pool.getconn() # 设置自动提交为False,需要手动提交 conn.autocommit = False return conn except Exception as e: logger.error(f"获取数据库连接失败: {str(e)}") raise def return_connection(self, conn): """ 返回数据库连接到连接池 Args: conn: 数据库连接对象 """ if self._pool and conn: try: self._pool.putconn(conn) except Exception as e: logger.warning(f"返回数据库连接失败: {str(e)}") def close_all(self): """关闭所有连接""" if self._pool: try: self._pool.closeall() logger.info("数据库连接池已关闭") except Exception as e: logger.error(f"关闭数据库连接池失败: {str(e)}") # 全局连接池实例 _connection_pool = DatabaseConnectionPool() @contextmanager def get_db_connection(): """ 获取数据库连接的上下文管理器 Yields: 数据库连接对象 """ conn = None try: conn = _connection_pool.get_connection() yield conn except Exception as e: logger.error(f"数据库连接错误: {str(e)}") raise finally: if conn: _connection_pool.return_connection(conn) @contextmanager def get_db_cursor(commit: bool = True): """ 获取数据库游标的上下文管理器 Args: commit: 是否自动提交事务 Yields: 数据库游标对象 """ with get_db_connection() as conn: cursor = None try: cursor = conn.cursor(cursor_factory=RealDictCursor) yield cursor if commit: conn.commit() except Exception as e: conn.rollback() logger.error(f"数据库操作错误: {str(e)}") raise finally: if cursor: cursor.close() class DatabaseManager: """数据库管理器""" def __init__(self): self.config = get_config() if self.config.database.enabled: _connection_pool.init_pool(self.config.database) else: logger.warning("数据库功能已禁用") def init_database(self): """初始化数据库表结构""" if not self.config.database.enabled: return logger.info("开始初始化数据库表结构") # 创建表的SQL语句 create_tables_sql = """ -- 公告表 CREATE TABLE IF NOT EXISTS announcements ( id SERIAL PRIMARY KEY, title VARCHAR(500) NOT NULL, publish_date TIMESTAMP NOT NULL, purchase_name VARCHAR(200), content_url TEXT, source_code VARCHAR(50) NOT NULL, source_name VARCHAR(100) NOT NULL, announcement_type VARCHAR(50) NOT NULL, crawled_at TIMESTAMP, created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, content_hash VARCHAR(32) UNIQUE, keyword_matched BOOLEAN DEFAULT FALSE, date_filtered BOOLEAN DEFAULT TRUE, is_new BOOLEAN DEFAULT TRUE, is_today BOOLEAN DEFAULT FALSE ); -- 公告来源表 CREATE TABLE IF NOT EXISTS announcement_sources ( code VARCHAR(50) PRIMARY KEY, category_id INTEGER NOT NULL, name VARCHAR(100) NOT NULL, type VARCHAR(50) NOT NULL ); -- 爬取结果表 CREATE TABLE IF NOT EXISTS crawl_results ( id SERIAL PRIMARY KEY, source_code VARCHAR(50) NOT NULL, status VARCHAR(20) NOT NULL, total_count INTEGER DEFAULT 0, new_count INTEGER DEFAULT 0, error_message TEXT, crawled_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP, duration FLOAT DEFAULT 0.0, FOREIGN KEY (source_code) REFERENCES announcement_sources(code) ); -- 创建索引 CREATE INDEX IF NOT EXISTS idx_announcements_publish_date ON announcements(publish_date DESC); CREATE INDEX IF NOT EXISTS idx_announcements_source_code ON announcements(source_code); CREATE INDEX IF NOT EXISTS idx_announcements_content_hash ON announcements(content_hash); CREATE INDEX IF NOT EXISTS idx_announcements_created_at ON announcements(created_at DESC); CREATE INDEX IF NOT EXISTS idx_crawl_results_crawled_at ON crawl_results(crawled_at DESC); -- 创建更新时间触发器 CREATE OR REPLACE FUNCTION update_updated_at_column() RETURNS TRIGGER AS $$ BEGIN NEW.updated_at = CURRENT_TIMESTAMP; RETURN NEW; END; $$ language 'plpgsql'; DROP TRIGGER IF EXISTS update_announcements_updated_at ON announcements; CREATE TRIGGER update_announcements_updated_at BEFORE UPDATE ON announcements FOR EACH ROW EXECUTE FUNCTION update_updated_at_column(); """ with get_db_cursor() as cursor: cursor.execute(create_tables_sql) logger.info("数据库表结构初始化完成") @retry_on_exception(RetryConfig(max_retries=3)) def save_announcement(self, announcement: Announcement) -> bool: """ 保存公告到数据库 Args: announcement: 公告对象 Returns: bool: 保存是否成功 """ if not self.config.database.enabled: return False # 生成内容哈希(如果还没有) if not announcement.content_hash: announcement.generate_content_hash() sql = """ INSERT INTO announcements ( title, publish_date, purchase_name, content_url, source_code, source_name, announcement_type, crawled_at, content_hash, keyword_matched, date_filtered, is_new, is_today ) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) ON CONFLICT (content_hash) DO NOTHING """ values = ( announcement.title, announcement.publish_date, announcement.purchase_name, announcement.content_url, announcement.source_code, announcement.source_name, announcement.announcement_type.value, announcement.crawled_at, announcement.content_hash, announcement.keyword_matched, announcement.date_filtered, announcement.is_new, announcement.is_today ) try: with get_db_cursor() as cursor: cursor.execute(sql, values) affected_rows = cursor.rowcount if affected_rows > 0: logger.debug(f"成功保存公告: {announcement.title[:50]}...") return True else: logger.debug(f"公告已存在,跳过保存: {announcement.title[:50]}...") return False except Exception as e: logger.error(f"保存公告失败: {str(e)}") return False @retry_on_exception(RetryConfig(max_retries=3)) def save_announcements_batch(self, announcements: List[Announcement]) -> int: """ 批量保存公告 Args: announcements: 公告列表 Returns: int: 成功保存的数量 """ if not self.config.database.enabled: return 0 if not announcements: return 0 # 为没有哈希的公告生成哈希 for announcement in announcements: if not announcement.content_hash: announcement.generate_content_hash() sql = """ INSERT INTO announcements ( title, publish_date, purchase_name, content_url, source_code, source_name, announcement_type, crawled_at, content_hash, keyword_matched, date_filtered, is_new, is_today ) VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s) ON CONFLICT (content_hash) DO NOTHING """ values = [] for announcement in announcements: values.append(( announcement.title, announcement.publish_date, announcement.purchase_name, announcement.content_url, announcement.source_code, announcement.source_name, announcement.announcement_type.value, announcement.crawled_at, announcement.content_hash, announcement.keyword_matched, announcement.date_filtered, announcement.is_new, announcement.is_today )) try: with get_db_cursor() as cursor: extras.execute_batch(cursor, sql, values) affected_rows = cursor.rowcount logger.info(f"批量保存公告完成,成功保存 {affected_rows} 条") return affected_rows except Exception as e: logger.error(f"批量保存公告失败: {str(e)}") return 0 @retry_on_exception(RetryConfig(max_retries=3)) def get_announcements(self, source_code: Optional[str] = None, start_date: Optional[datetime] = None, end_date: Optional[datetime] = None, limit: int = 100, offset: int = 0) -> List[Announcement]: """ 查询公告 Args: source_code: 来源代码过滤 start_date: 开始日期 end_date: 结束日期 limit: 限制数量 offset: 偏移量 Returns: List[Announcement]: 公告列表 """ if not self.config.database.enabled: return [] sql = """ SELECT * FROM announcements WHERE 1=1 """ params = [] if source_code: sql += " AND source_code = %s" params.append(source_code) if start_date: sql += " AND publish_date >= %s" params.append(start_date) if end_date: sql += " AND publish_date <= %s" params.append(end_date) sql += " ORDER BY publish_date DESC LIMIT %s OFFSET %s" params.extend([limit, offset]) try: with get_db_cursor() as cursor: cursor.execute(sql, params) rows = cursor.fetchall() announcements = [] for row in rows: # 转换数据类型 row_dict = dict(row) row_dict['announcement_type'] = AnnouncementType(row_dict['announcement_type']) announcements.append(Announcement.from_dict(row_dict)) return announcements except Exception as e: logger.error(f"查询公告失败: {str(e)}") return [] @retry_on_exception(RetryConfig(max_retries=3)) def save_crawl_result(self, result: CrawlResult) -> bool: """ 保存爬取结果 Args: result: 爬取结果对象 Returns: bool: 保存是否成功 """ if not self.config.database.enabled: return False sql = """ INSERT INTO crawl_results ( source_code, status, total_count, new_count, error_message, crawled_at, duration ) VALUES (%s, %s, %s, %s, %s, %s, %s) """ values = ( result.source.code, result.status.value, result.total_count, result.new_count, result.error_message, result.crawled_at, result.duration ) try: with get_db_cursor() as cursor: cursor.execute(sql, values) logger.debug(f"保存爬取结果: {result.source.name}") return True except Exception as e: logger.error(f"保存爬取结果失败: {str(e)}") return False @retry_on_exception(RetryConfig(max_retries=3)) def cleanup_expired_data(self, days: int = 90) -> int: """ 清理过期数据 Args: days: 保留天数 Returns: int: 清理的记录数 """ if not self.config.database.enabled: return 0 cutoff_date = datetime.now() - timedelta(days=days) sql = "DELETE FROM announcements WHERE created_at < %s" try: with get_db_cursor() as cursor: cursor.execute(sql, (cutoff_date,)) deleted_count = cursor.rowcount logger.info(f"清理过期数据完成,删除 {deleted_count} 条记录") return deleted_count except Exception as e: logger.error(f"清理过期数据失败: {str(e)}") return 0 @retry_on_exception(RetryConfig(max_retries=3)) def get_statistics(self) -> Dict[str, Any]: """ 获取统计信息 Returns: Dict[str, Any]: 统计数据 """ if not self.config.database.enabled: return {} sql = """ SELECT COUNT(*) as total_announcements, COUNT(CASE WHEN is_today THEN 1 END) as today_announcements, COUNT(CASE WHEN is_new THEN 1 END) as new_announcements, COUNT(DISTINCT source_code) as sources_count, MAX(crawled_at) as last_crawl_time FROM announcements """ try: with get_db_cursor() as cursor: cursor.execute(sql) result = cursor.fetchone() return dict(result) if result else {} except Exception as e: logger.error(f"获取统计信息失败: {str(e)}") return {} def is_announcement_exists(self, content_hash: str) -> bool: """ 检查公告是否已存在 Args: content_hash: 内容哈希 Returns: bool: 是否存在 """ if not self.config.database.enabled: return False sql = "SELECT 1 FROM announcements WHERE content_hash = %s LIMIT 1" try: with get_db_cursor() as cursor: cursor.execute(sql, (content_hash,)) return cursor.fetchone() is not None except Exception as e: logger.error(f"检查公告存在性失败: {str(e)}") return False def get_recent_announcements(self, hours: int = 24) -> List[Announcement]: """ 获取最近的公告 Args: hours: 最近小时数 Returns: List[Announcement]: 公告列表 """ if not self.config.database.enabled: return [] cutoff_time = datetime.now() - timedelta(hours=hours) sql = """ SELECT * FROM announcements WHERE crawled_at >= %s ORDER BY crawled_at DESC """ try: with get_db_cursor() as cursor: cursor.execute(sql, (cutoff_time,)) rows = cursor.fetchall() announcements = [] for row in rows: row_dict = dict(row) row_dict['announcement_type'] = AnnouncementType(row_dict['announcement_type']) announcements.append(Announcement.from_dict(row_dict)) return announcements except Exception as e: logger.error(f"获取最近公告失败: {str(e)}") return [] # 全局数据库管理器实例 _db_manager = None def get_database_manager() -> DatabaseManager: """ 获取数据库管理器实例 Returns: DatabaseManager: 数据库管理器实例 """ global _db_manager if _db_manager is None: _db_manager = DatabaseManager() return _db_manager def init_database(): """初始化数据库""" manager = get_database_manager() manager.init_database() def cleanup_database(): """清理数据库连接""" _connection_pool.close_all()