7 Sprints 全部完成: Sprint 0: 基础设施 (DataManager + MariaDB) Sprint 1: 因子引擎 (34因子/12分类) Sprint 2: VectorBT 回测 (5策略+截面) Sprint 3: Optuna 优化 (+Walk-Forward) Sprint 4: ML 模型 (LightGBM+CatBoost) Sprint 5: Qwen 情绪因子 (三源新闻+日期对齐) Sprint 6: Agent 系统 (4Agent+日报.md/.html) 生产加固 (15项): Tushare双源fallback, SSH自动恢复, pool_pre_ping, save_daily先删后插, load_dotenv绝对路径, 日报5d/20d修复, RiskAgent改上证指数, 昨日对比+数据截止, mac_report utf8mb4, CLAUDE-*.md 9条已知Bug, demo全参数化, djapi数据源归一化, indexDatas API修正 Co-Authored-By: Claude Opus 4.7 <noreply@anthropic.com>
95 lines
3.2 KiB
Python
95 lines
3.2 KiB
Python
import os
|
|
import mysql.connector
|
|
from mysql.connector import Error
|
|
|
|
|
|
class MySQLDB:
|
|
def __init__(self, host=None, port=None, username=None, password=None, database=None):
|
|
self.host = host or os.getenv('MYSQL_HOST', 'localhost')
|
|
self.port = port or int(os.getenv('MYSQL_PORT', '3306'))
|
|
self.username = username or os.getenv('MYSQL_USER', 'myquant')
|
|
self.password = password or os.getenv('MYSQL_PASSWORD', '')
|
|
self.database = database or os.getenv('MYSQL_DATABASE', 'myquant')
|
|
self.connection = None
|
|
self.connect()
|
|
|
|
def connect(self):
|
|
"""连接数据库"""
|
|
try:
|
|
self.connection = mysql.connector.connect(
|
|
host=self.host,
|
|
port=self.port,
|
|
user=self.username,
|
|
password=self.password,
|
|
database=self.database
|
|
)
|
|
if self.connection.is_connected():
|
|
print("成功连接到MySQL数据库")
|
|
except Error as e:
|
|
print(f"连接错误: {e}")
|
|
|
|
def insert_data(self, table, data):
|
|
"""插入数据"""
|
|
try:
|
|
cursor = self.connection.cursor()
|
|
columns = ', '.join(data.keys())
|
|
placeholders = ', '.join(['%s'] * len(data))
|
|
query = f"INSERT INTO {table} ({columns}) VALUES ({placeholders})"
|
|
|
|
cursor.execute(query, tuple(data.values()))
|
|
self.connection.commit()
|
|
print(f"成功插入数据,影响行数: {cursor.rowcount}")
|
|
return cursor.lastrowid
|
|
except Error as e:
|
|
print(f"插入错误: {e}")
|
|
return None
|
|
finally:
|
|
if cursor:
|
|
cursor.close()
|
|
|
|
def query_data(self, table, columns="*", where=None, params=None):
|
|
"""查询数据"""
|
|
try:
|
|
cursor = self.connection.cursor(dictionary=True)
|
|
|
|
query = f"SELECT {columns} FROM {table}"
|
|
if where:
|
|
query += f" WHERE {where}"
|
|
print(query)
|
|
cursor.execute(query, params or ())
|
|
result = cursor.fetchall()
|
|
return result
|
|
except Error as e:
|
|
print(f"查询错误: {e}")
|
|
return []
|
|
finally:
|
|
if cursor:
|
|
cursor.close()
|
|
|
|
def update_data(self, table, data, where, params=None):
|
|
"""更新数据"""
|
|
try:
|
|
cursor = self.connection.cursor()
|
|
|
|
set_clause = ', '.join([f"{key} = %s" for key in data.keys()])
|
|
query = f"UPDATE {table} SET {set_clause} WHERE {where}"
|
|
|
|
all_params = tuple(data.values()) + (params if params else ())
|
|
|
|
cursor.execute(query, all_params)
|
|
self.connection.commit()
|
|
print(f"成功更新数据,影响行数: {cursor.rowcount}")
|
|
return cursor.rowcount
|
|
except Error as e:
|
|
print(f"更新错误: {e}")
|
|
return 0
|
|
finally:
|
|
if cursor:
|
|
cursor.close()
|
|
|
|
def close(self):
|
|
"""关闭数据库连接"""
|
|
if self.connection and self.connection.is_connected():
|
|
self.connection.close()
|
|
print("数据库连接已关闭")
|