Files
aiagents-stock/portfolio_db.py
T
2025-10-20 09:18:44 +08:00

626 lines
20 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""
持仓股票数据库管理模块
提供持仓股票和分析历史的数据库操作接口
"""
import sqlite3
from datetime import datetime
from typing import List, Dict, Optional, Tuple
import os
# 数据库文件路径
DB_PATH = "portfolio_stocks.db"
class PortfolioDB:
"""持仓股票数据库管理类"""
def __init__(self, db_path: str = DB_PATH):
"""
初始化数据库连接
Args:
db_path: 数据库文件路径
"""
self.db_path = db_path
self._init_database()
def _get_connection(self) -> sqlite3.Connection:
"""获取数据库连接"""
conn = sqlite3.connect(self.db_path)
conn.row_factory = sqlite3.Row # 使查询结果可以通过列名访问
return conn
def _init_database(self):
"""初始化数据库表结构"""
conn = self._get_connection()
cursor = conn.cursor()
try:
# 创建持仓股票表
cursor.execute('''
CREATE TABLE IF NOT EXISTS portfolio_stocks (
id INTEGER PRIMARY KEY AUTOINCREMENT,
code TEXT NOT NULL UNIQUE,
name TEXT NOT NULL,
cost_price REAL,
quantity INTEGER,
note TEXT,
auto_monitor BOOLEAN DEFAULT 1,
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
)
''')
# 创建持仓分析历史表
cursor.execute('''
CREATE TABLE IF NOT EXISTS portfolio_analysis_history (
id INTEGER PRIMARY KEY AUTOINCREMENT,
portfolio_stock_id INTEGER NOT NULL,
analysis_time TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
rating TEXT,
confidence REAL,
current_price REAL,
target_price REAL,
entry_min REAL,
entry_max REAL,
take_profit REAL,
stop_loss REAL,
summary TEXT,
FOREIGN KEY (portfolio_stock_id) REFERENCES portfolio_stocks(id) ON DELETE CASCADE
)
''')
# 创建索引以提升查询性能
cursor.execute('''
CREATE INDEX IF NOT EXISTS idx_portfolio_analysis_stock_id
ON portfolio_analysis_history(portfolio_stock_id)
''')
cursor.execute('''
CREATE INDEX IF NOT EXISTS idx_portfolio_analysis_time
ON portfolio_analysis_history(analysis_time DESC)
''')
conn.commit()
print(f"[OK] 数据库初始化成功: {self.db_path}")
except Exception as e:
print(f"[ERROR] 数据库初始化失败: {e}")
conn.rollback()
raise
finally:
conn.close()
# ==================== 持仓股票CRUD操作 ====================
def add_stock(self, code: str, name: str, cost_price: Optional[float] = None,
quantity: Optional[int] = None, note: str = "",
auto_monitor: bool = True) -> int:
"""
添加持仓股票
Args:
code: 股票代码
name: 股票名称
cost_price: 持仓成本价(可选)
quantity: 持仓数量(可选)
note: 备注信息
auto_monitor: 是否自动同步到监测列表
Returns:
新增股票的ID
Raises:
sqlite3.IntegrityError: 如果股票代码已存在
"""
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute('''
INSERT INTO portfolio_stocks
(code, name, cost_price, quantity, note, auto_monitor, created_at, updated_at)
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
''', (code, name, cost_price, quantity, note, auto_monitor,
datetime.now(), datetime.now()))
conn.commit()
stock_id = cursor.lastrowid
print(f"[OK] 添加持仓股票成功: {code} {name} (ID: {stock_id})")
return stock_id
except sqlite3.IntegrityError as e:
print(f"[ERROR] 股票代码已存在: {code}")
raise ValueError(f"股票代码 {code} 已存在") from e
except Exception as e:
print(f"[ERROR] 添加持仓股票失败: {e}")
conn.rollback()
raise
finally:
conn.close()
def update_stock(self, stock_id: int, **kwargs) -> bool:
"""
更新持仓股票信息
Args:
stock_id: 股票ID
**kwargs: 要更新的字段(code, name, cost_price, quantity, note, auto_monitor
Returns:
是否更新成功
"""
conn = self._get_connection()
cursor = conn.cursor()
# 允许更新的字段
allowed_fields = ['code', 'name', 'cost_price', 'quantity', 'note', 'auto_monitor']
update_fields = {k: v for k, v in kwargs.items() if k in allowed_fields}
if not update_fields:
print("[WARN] 没有需要更新的字段")
return False
# 添加更新时间
update_fields['updated_at'] = datetime.now()
# 构建SQL语句
set_clause = ', '.join([f"{field} = ?" for field in update_fields.keys()])
values = list(update_fields.values()) + [stock_id]
try:
cursor.execute(f'''
UPDATE portfolio_stocks
SET {set_clause}
WHERE id = ?
''', values)
conn.commit()
if cursor.rowcount > 0:
print(f"[OK] 更新持仓股票成功: ID {stock_id}")
return True
else:
print(f"[WARN] 未找到股票: ID {stock_id}")
return False
except Exception as e:
print(f"[ERROR] 更新持仓股票失败: {e}")
conn.rollback()
raise
finally:
conn.close()
def delete_stock(self, stock_id: int) -> bool:
"""
删除持仓股票(级联删除其所有分析历史)
Args:
stock_id: 股票ID
Returns:
是否删除成功
"""
conn = self._get_connection()
cursor = conn.cursor()
try:
# 由于设置了ON DELETE CASCADE,删除股票会自动删除其分析历史
cursor.execute('DELETE FROM portfolio_stocks WHERE id = ?', (stock_id,))
conn.commit()
if cursor.rowcount > 0:
print(f"[OK] 删除持仓股票成功: ID {stock_id}")
return True
else:
print(f"[WARN] 未找到股票: ID {stock_id}")
return False
except Exception as e:
print(f"[ERROR] 删除持仓股票失败: {e}")
conn.rollback()
raise
finally:
conn.close()
def get_stock(self, stock_id: int) -> Optional[Dict]:
"""
获取单只持仓股票信息
Args:
stock_id: 股票ID
Returns:
股票信息字典,不存在则返回None
"""
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute('SELECT * FROM portfolio_stocks WHERE id = ?', (stock_id,))
row = cursor.fetchone()
if row:
return dict(row)
return None
finally:
conn.close()
def get_stock_by_code(self, code: str) -> Optional[Dict]:
"""
根据股票代码获取持仓股票信息
Args:
code: 股票代码
Returns:
股票信息字典,不存在则返回None
"""
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute('SELECT * FROM portfolio_stocks WHERE code = ?', (code,))
row = cursor.fetchone()
if row:
return dict(row)
return None
finally:
conn.close()
def get_all_stocks(self, auto_monitor_only: bool = False) -> List[Dict]:
"""
获取所有持仓股票列表
Args:
auto_monitor_only: 是否只返回启用自动监测的股票
Returns:
股票信息字典列表
"""
conn = self._get_connection()
cursor = conn.cursor()
try:
if auto_monitor_only:
cursor.execute('''
SELECT * FROM portfolio_stocks
WHERE auto_monitor = 1
ORDER BY created_at DESC
''')
else:
cursor.execute('SELECT * FROM portfolio_stocks ORDER BY created_at DESC')
rows = cursor.fetchall()
return [dict(row) for row in rows]
finally:
conn.close()
def search_stocks(self, keyword: str) -> List[Dict]:
"""
搜索持仓股票(按代码或名称)
Args:
keyword: 搜索关键词
Returns:
匹配的股票信息字典列表
"""
conn = self._get_connection()
cursor = conn.cursor()
try:
keyword_pattern = f"%{keyword}%"
cursor.execute('''
SELECT * FROM portfolio_stocks
WHERE code LIKE ? OR name LIKE ?
ORDER BY created_at DESC
''', (keyword_pattern, keyword_pattern))
rows = cursor.fetchall()
return [dict(row) for row in rows]
finally:
conn.close()
def get_stock_count(self) -> int:
"""
获取持仓股票总数
Returns:
股票数量
"""
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute('SELECT COUNT(*) as count FROM portfolio_stocks')
result = cursor.fetchone()
return result['count']
finally:
conn.close()
# ==================== 分析历史记录操作 ====================
def save_analysis(self, stock_id: int, rating: str, confidence: float,
current_price: float, target_price: Optional[float] = None,
entry_min: Optional[float] = None, entry_max: Optional[float] = None,
take_profit: Optional[float] = None, stop_loss: Optional[float] = None,
summary: str = "") -> int:
"""
保存分析历史记录
Args:
stock_id: 持仓股票ID
rating: 投资评级(买入/持有/卖出)
confidence: 信心度(0-10
current_price: 当前价格
target_price: 目标价位
entry_min: 进场区间最小值
entry_max: 进场区间最大值
take_profit: 止盈位
stop_loss: 止损位
summary: 分析摘要
Returns:
新增分析记录的ID
"""
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute('''
INSERT INTO portfolio_analysis_history
(portfolio_stock_id, analysis_time, rating, confidence, current_price,
target_price, entry_min, entry_max, take_profit, stop_loss, summary)
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
''', (stock_id, datetime.now(), rating, confidence, current_price,
target_price, entry_min, entry_max, take_profit, stop_loss, summary))
conn.commit()
analysis_id = cursor.lastrowid
print(f"[OK] 保存分析历史成功: 股票ID {stock_id}, 评级 {rating}")
return analysis_id
except Exception as e:
print(f"[ERROR] 保存分析历史失败: {e}")
conn.rollback()
raise
finally:
conn.close()
def get_analysis_history(self, stock_id: int, limit: int = 10) -> List[Dict]:
"""
获取股票的分析历史记录
Args:
stock_id: 持仓股票ID
limit: 返回记录数量限制
Returns:
分析历史记录列表(按时间倒序)
"""
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute('''
SELECT * FROM portfolio_analysis_history
WHERE portfolio_stock_id = ?
ORDER BY analysis_time DESC
LIMIT ?
''', (stock_id, limit))
rows = cursor.fetchall()
return [dict(row) for row in rows]
finally:
conn.close()
def get_latest_analysis_history(self, stock_id: int, limit: int = 10) -> List[Dict]:
"""
获取股票的最新分析历史记录(按时间倒序)
这是 get_analysis_history 的别名方法,用于保持代码兼容性
Args:
stock_id: 持仓股票ID
limit: 返回记录数量限制
Returns:
分析历史记录列表(按时间倒序)
"""
return self.get_analysis_history(stock_id, limit)
def get_latest_analysis(self, stock_id: int) -> Optional[Dict]:
"""
获取股票的最新一次分析记录
Args:
stock_id: 持仓股票ID
Returns:
最新分析记录字典,不存在则返回None
"""
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute('''
SELECT * FROM portfolio_analysis_history
WHERE portfolio_stock_id = ?
ORDER BY analysis_time DESC
LIMIT 1
''', (stock_id,))
row = cursor.fetchone()
if row:
return dict(row)
return None
finally:
conn.close()
def get_rating_changes(self, stock_id: int, days: int = 30) -> List[Tuple[str, str, str]]:
"""
获取股票在指定天数内的评级变化
Args:
stock_id: 持仓股票ID
days: 查询天数
Returns:
评级变化列表 [(时间, 旧评级, 新评级), ...]
"""
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute('''
SELECT analysis_time, rating
FROM portfolio_analysis_history
WHERE portfolio_stock_id = ?
AND analysis_time >= datetime('now', '-' || ? || ' days')
ORDER BY analysis_time ASC
''', (stock_id, days))
rows = cursor.fetchall()
changes = []
for i in range(1, len(rows)):
prev_rating = rows[i-1]['rating']
curr_rating = rows[i]['rating']
if prev_rating != curr_rating:
changes.append((
rows[i]['analysis_time'],
prev_rating,
curr_rating
))
return changes
finally:
conn.close()
def delete_old_analysis(self, days: int = 90) -> int:
"""
删除超过指定天数的分析历史记录
Args:
days: 保留天数
Returns:
删除的记录数量
"""
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute('''
DELETE FROM portfolio_analysis_history
WHERE analysis_time < datetime('now', '-' || ? || ' days')
''', (days,))
conn.commit()
deleted_count = cursor.rowcount
print(f"[OK] 清理历史分析记录: 删除 {deleted_count} 条记录")
return deleted_count
except Exception as e:
print(f"[ERROR] 清理历史分析记录失败: {e}")
conn.rollback()
raise
finally:
conn.close()
def get_all_latest_analysis(self) -> List[Dict]:
"""
获取所有持仓股票的最新分析记录
Returns:
包含股票信息和最新分析的字典列表
"""
conn = self._get_connection()
cursor = conn.cursor()
try:
cursor.execute('''
SELECT
s.*,
h.rating, h.confidence, h.current_price, h.target_price,
h.entry_min, h.entry_max, h.take_profit, h.stop_loss,
h.analysis_time
FROM portfolio_stocks s
LEFT JOIN (
SELECT h1.*
FROM portfolio_analysis_history h1
INNER JOIN (
SELECT portfolio_stock_id, MAX(analysis_time) as max_time
FROM portfolio_analysis_history
GROUP BY portfolio_stock_id
) h2
ON h1.portfolio_stock_id = h2.portfolio_stock_id
AND h1.analysis_time = h2.max_time
) h ON s.id = h.portfolio_stock_id
ORDER BY s.created_at DESC
''')
rows = cursor.fetchall()
return [dict(row) for row in rows]
finally:
conn.close()
# 创建全局数据库实例
portfolio_db = PortfolioDB()
if __name__ == "__main__":
# 测试代码
print("=" * 50)
print("持仓股票数据库测试")
print("=" * 50)
# 初始化数据库
db = PortfolioDB("test_portfolio.db")
# 测试添加股票
try:
stock_id = db.add_stock("600519", "贵州茅台", 1650.5, 100, "长期持有")
print(f"\n添加股票ID: {stock_id}")
except ValueError as e:
print(f"\n{e}")
# 测试查询所有股票
print("\n所有持仓股票:")
stocks = db.get_all_stocks()
for stock in stocks:
print(f" {stock['code']} {stock['name']}")
# 测试保存分析历史
if stocks:
stock_id = stocks[0]['id']
analysis_id = db.save_analysis(
stock_id, "买入", 8.5, 1700.0, 1850.0,
1600.0, 1650.0, 1900.0, 1500.0,
"技术面和基本面均良好"
)
print(f"\n保存分析记录ID: {analysis_id}")
# 查询分析历史
print(f"\n股票 {stocks[0]['code']} 的分析历史:")
history = db.get_analysis_history(stock_id)
for h in history:
print(f" {h['analysis_time']}: {h['rating']} (信心度: {h['confidence']})")
print("\n[OK] 数据库测试完成")