database.py

"""
QTrading SQLite 数据库操作模块
"""

import sqlite3
import pandas as pd
import os
from typing import Optional, List
from config.config import config
from utils.logger import get_logger

logger = get_logger(__name__)

class DatabaseManager:
    """SQLite 数据库管理器"""

    def __init__(self, db_path: str = None):
        """
        初始化数据库管理器

        Args:
            db_path: 数据库文件路径(可选,默认从配置文件读取)
        """
        self.db_path = db_path or config.database.db_path
        self._init_database()

    def _init_database(self) -> None:
        """初始化数据库"""
        db_dir = os.path.dirname(self.db_path)
        if not os.path.exists(db_dir):
            os.makedirs(db_dir)
            logger.info(f"创建数据库目录: {db_dir}")

        try:
            conn = self._get_connection()
            cursor = conn.cursor()
            # 创建日线数据表
            cursor.execute('''
                CREATE TABLE IF NOT EXISTS daily_kline (
                    ts_code TEXT NOT NULL,
                    trade_date TEXT NOT NULL,
                    open REAL,
                    high REAL,
                    low REAL,
                    close REAL,
                    vol REAL,
                    amount REAL,
                    PRIMARY KEY (ts_code, trade_date)
                )
            ''')
            # 创建复权因子表
            cursor.execute('''
                CREATE TABLE IF NOT EXISTS adj_factor (
                    ts_code TEXT NOT NULL,
                    trade_date TEXT NOT NULL,
                    adj_factor REAL,
                    PRIMARY KEY (ts_code, trade_date)
                )
            ''')
            conn.commit()
            logger.info(f"数据库初始化成功: {self.db_path}")
        except Exception as e:
            logger.error(f"数据库初始化失败: {e}")

    def _get_connection(self) -> sqlite3.Connection:
        """获取数据库连接"""
        return sqlite3.connect(self.db_path)

    def save_daily_kline(self, data: pd.DataFrame) -> bool:
        """
        保存日线数据到数据库

        Args:
            data: 日线数据 DataFrame

        Returns:
            bool: 是否保存成功
        """
        if data.empty:
            logger.warning("日线数据为空,无需保存")
            return False

        try:
            conn = self._get_connection()
            data.to_sql(
                'daily_kline',
                conn,
                if_exists='append',
                index=False,
                chunksize=1000
            )
            conn.commit()
            logger.info(f"成功保存 {len(data)} 条日线数据")
            return True
        except Exception as e:
            logger.error(f"保存日线数据失败: {e}")
            return False

    def save_adj_factor(self, data: pd.DataFrame) -> bool:
        """
        保存复权因子到数据库

        Args:
            data: 复权因子 DataFrame

        Returns:
            bool: 是否保存成功
        """
        if data.empty:
            logger.warning("复权因子数据为空,无需保存")
            return False

        try:
            conn = self._get_connection()
            data.to_sql(
                'adj_factor',
                conn,
                if_exists='append',
                index=False,
                chunksize=1000
            )
            conn.commit()
            logger.info(f"成功保存 {len(data)} 条复权因子数据")
            return True
        except Exception as e:
            logger.error(f"保存复权因子失败: {e}")
            return False

    def get_daily_kline(self, ts_code: str, start_date: str, end_date: str) -> Optional[pd.DataFrame]:
        """
        从数据库获取日线数据

        Args:
            ts_code: 股票代码(如 '000001.SZ')
            start_date: 开始日期(YYYYMMDD)
            end_date: 结束日期(YYYYMMDD)

        Returns:
            pd.DataFrame: 日线数据
        """
        try:
            conn = self._get_connection()
            query = '''
                SELECT * FROM daily_kline
                WHERE ts_code = ? AND trade_date BETWEEN ? AND ?
                ORDER BY trade_date
            '''
            data = pd.read_sql_query(
                query,
                conn,
                params=(ts_code, start_date, end_date)
            )
            logger.info(f"从数据库获取到 {len(data)} 条日线数据")
            return data
        except Exception as e:
            logger.error(f"从数据库获取日线数据失败: {e}")
            return None

    def get_adj_factor(self, ts_code: str, start_date: str, end_date: str) -> Optional[pd.DataFrame]:
        """
        从数据库获取复权因子

        Args:
            ts_code: 股票代码(如 '000001.SZ')
            start_date: 开始日期(YYYYMMDD)
            end_date: 结束日期(YYYYMMDD)

        Returns:
            pd.DataFrame: 复权因子数据
        """
        try:
            conn = self._get_connection()
            query = '''
                SELECT * FROM adj_factor
                WHERE ts_code = ? AND trade_date BETWEEN ? AND ?
                ORDER BY trade_date
            '''
            data = pd.read_sql_query(
                query,
                conn,
                params=(ts_code, start_date, end_date)
            )
            logger.info(f"从数据库获取到 {len(data)} 条复权因子数据")
            return data
        except Exception as e:
            logger.error(f"从数据库获取复权因子失败: {e}")
            return None

    def get_last_trade_date(self, ts_code: str) -> Optional[str]:
        """
        获取股票最后一条数据的交易日

        Args:
            ts_code: 股票代码(如 '000001.SZ')

        Returns:
            str: 最后交易日(YYYYMMDD)
        """
        try:
            conn = self._get_connection()
            query = '''
                SELECT MAX(trade_date) FROM daily_kline
                WHERE ts_code = ?
            '''
            cursor = conn.cursor()
            cursor.execute(query, (ts_code,))
            result = cursor.fetchone()
            return result[0] if result[0] else None
        except Exception as e:
            logger.error(f"获取最后交易日失败: {e}")
            return None

    def get_all_stock_codes(self) -> List[str]:
        """
        获取数据库中所有股票代码

        Returns:
            List[str]: 股票代码列表
        """
        try:
            conn = self._get_connection()
            query = '''
                SELECT DISTINCT ts_code FROM daily_kline
            '''
            data = pd.read_sql_query(query, conn)
            return data['ts_code'].tolist()
        except Exception as e:
            logger.error(f"获取股票代码列表失败: {e}")
            return []

    def delete_stock_data(self, ts_code: str) -> bool:
        """
        删除指定股票的所有数据

        Args:
            ts_code: 股票代码(如 '000001.SZ')

        Returns:
            bool: 是否删除成功
        """
        try:
            conn = self._get_connection()
            cursor = conn.cursor()
            cursor.execute('''
                DELETE FROM daily_kline WHERE ts_code = ?
            ''', (ts_code,))
            cursor.execute('''
                DELETE FROM adj_factor WHERE ts_code = ?
            ''', (ts_code,))
            conn.commit()
            logger.info(f"成功删除股票 {ts_code} 的所有数据")
            return True
        except Exception as e:
            logger.error(f"删除股票数据失败: {e}")
            return False

    def get_database_info(self) -> dict:
        """
        获取数据库信息

        Returns:
            dict: 数据库信息(包含表名、记录数)
        """
        info = {}
        try:
            conn = self._get_connection()
            cursor = conn.cursor()
            # 获取表名
            cursor.execute('''
                SELECT name FROM sqlite_master WHERE type='table'
            ''')
            tables = cursor.fetchall()
            for table in tables:
                table_name = table[0]
                cursor.execute(f'''
                    SELECT COUNT(*) FROM {table_name}
                ''')
                count = cursor.fetchone()[0]
                info[table_name] = count
            return info
        except Exception as e:
            logger.error(f"获取数据库信息失败: {e}")
            return {}