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 {}