Files
stock/stock-html/migrate_data.py
T
freedakgmail 9c7d7abdd4 Initial commit
2026-07-17 18:49:35 +08:00

243 lines
7.1 KiB
Python

"""
数据迁移脚本 - 将JSON数据迁移到PostgreSQL数据库
"""
import json
import os
import sys
import psycopg2
from psycopg2.extras import RealDictCursor
from werkzeug.security import generate_password_hash
from config import Config
def get_db():
"""获取数据库连接"""
return psycopg2.connect(
host=Config.DB_HOST,
port=Config.DB_PORT,
database=Config.DB_NAME,
user=Config.DB_USER,
password=Config.DB_PASSWORD
)
def create_default_user(conn, email=None, password=None):
"""创建用户"""
cur = conn.cursor(cursor_factory=RealDictCursor)
# 使用指定的或默认的邮箱密码
user_email = email or 'admin@admin.com'
user_password = password or 'admin123'
# 检查是否已存在用户
cur.execute("SELECT id FROM users WHERE email = %s OR username = %s", (user_email, user_email))
user = cur.fetchone()
if user:
print(f"用户 {user_email} 已存在,使用现有用户")
return user['id']
# 创建用户
password_hash = generate_password_hash(user_password)
cur.execute(
"INSERT INTO users (username, email, password_hash) VALUES (%s, %s, %s) RETURNING id",
(user_email, user_email, password_hash)
)
user_id = cur.fetchone()['id']
conn.commit()
print(f"创建用户: {user_email}")
return user_id
def migrate_trades(conn, user_id):
"""迁移交易记录"""
if not os.path.exists(Config.TRADES_FILE):
print("trades.json 不存在,跳过")
return 0
with open(Config.TRADES_FILE, 'r', encoding='utf-8') as f:
trades = json.load(f)
if not trades:
print("trades.json 为空,跳过")
return 0
cur = conn.cursor()
count = 0
for trade in trades:
try:
# 处理日期
trade_date = trade.get('trade_date')
if trade_date and len(trade_date) > 10:
trade_date = trade_date[:10]
# 处理数值
price = trade.get('price')
if price and price != '':
price = float(price)
else:
price = None
quantity = trade.get('quantity')
if quantity and quantity != '':
quantity = int(quantity)
else:
quantity = None
profit_amount = trade.get('profit_amount')
if profit_amount and profit_amount != '':
profit_amount = float(profit_amount)
else:
profit_amount = None
stop_loss_price = trade.get('stop_loss_price')
if stop_loss_price and stop_loss_price != '':
stop_loss_price = float(stop_loss_price)
else:
stop_loss_price = None
cur.execute("""
INSERT INTO trades (user_id, stock_code, stock_name, trade_type, price,
quantity, trade_date, reason, result, profit_amount,
stop_loss_price, notes)
VALUES (%s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s, %s)
""", (
user_id,
trade.get('stock_code'),
trade.get('stock_name'),
trade.get('trade_type'),
price,
quantity,
trade_date,
trade.get('reason'),
trade.get('result'),
profit_amount,
stop_loss_price,
trade.get('notes')
))
count += 1
except Exception as e:
print(f"迁移交易记录失败: {e}, 数据: {trade}")
conn.commit()
print(f"迁移交易记录: {count}")
return count
def migrate_watchlist(conn, user_id):
"""迁移关注列表"""
if not os.path.exists(Config.WATCHLIST_FILE):
print("watchlist.json 不存在,跳过")
return 0
with open(Config.WATCHLIST_FILE, 'r', encoding='utf-8') as f:
watchlist = json.load(f)
if not watchlist:
print("watchlist.json 为空,跳过")
return 0
cur = conn.cursor()
count = 0
for item in watchlist:
try:
cur.execute("""
INSERT INTO watchlist (user_id, stock_code, stock_name)
VALUES (%s, %s, %s)
ON CONFLICT (user_id, stock_code) DO NOTHING
""", (
user_id,
item.get('code'),
item.get('name')
))
count += 1
except Exception as e:
print(f"迁移关注列表失败: {e}, 数据: {item}")
conn.commit()
print(f"迁移关注列表: {count}")
return count
def migrate_alerts_cache(conn, user_id):
"""迁移分析缓存"""
if not os.path.exists(Config.ALERTS_CACHE_FILE):
print("alerts_cache.json 不存在,跳过")
return 0
with open(Config.ALERTS_CACHE_FILE, 'r', encoding='utf-8') as f:
cache = json.load(f)
alerts = cache.get('alerts', [])
if not alerts:
print("alerts_cache.json 为空,跳过")
return 0
cur = conn.cursor()
cur.execute("""
INSERT INTO alerts_cache (user_id, data, updated_at)
VALUES (%s, %s, NOW())
ON CONFLICT (user_id) DO UPDATE SET
data = EXCLUDED.data,
updated_at = NOW()
""", (user_id, json.dumps(alerts)))
conn.commit()
print(f"迁移分析缓存: {len(alerts)}")
return len(alerts)
def main():
import argparse
parser = argparse.ArgumentParser(description='数据迁移工具')
parser.add_argument('--email', default=None, help='用户邮箱')
parser.add_argument('--password', default=None, help='用户密码')
args = parser.parse_args()
print("=" * 60)
print("数据迁移 - JSON -> PostgreSQL")
print("=" * 60)
try:
conn = get_db()
print("数据库连接成功")
except Exception as e:
print(f"数据库连接失败: {e}")
print("\n请先执行: psql -U postgres -f init_db.sql")
sys.exit(1)
try:
# 创建用户
user_id = create_default_user(conn, args.email, args.password)
# 先清空现有数据
cur = conn.cursor()
cur.execute("DELETE FROM trades WHERE user_id = %s", (user_id,))
cur.execute("DELETE FROM watchlist WHERE user_id = %s", (user_id,))
cur.execute("DELETE FROM alerts_cache WHERE user_id = %s", (user_id,))
conn.commit()
print("清空现有数据")
# 迁移数据
migrate_trades(conn, user_id)
migrate_watchlist(conn, user_id)
migrate_alerts_cache(conn, user_id)
print("=" * 60)
print("迁移完成!")
print(f"登录邮箱: {args.email or 'admin@admin.com'}")
print(f"登录密码: {args.password or 'admin123'}")
print("=" * 60)
except Exception as e:
print(f"迁移失败: {e}")
conn.rollback()
raise
finally:
conn.close()
if __name__ == '__main__':
main()