手动模式
This commit is contained in:
@@ -0,0 +1,612 @@
|
||||
"""
|
||||
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()
|
||||
Reference in New Issue
Block a user