Files
sanguo_vnpy_v2/sanguo_web/database.py
claude_dev 918bbed0fc fix: 修复登录500错误和移除明文密码提示
- 修复 deps.py 中 get_vn_service 的引用错误 (vn_service.vn_service -> vn_service)
- 移除登录页面上的明文密码提示
- 改进前端错误处理,避免数据加载失败导致登录显示错误
2026-07-02 12:23:55 +08:00

174 lines
4.5 KiB
Python

"""
数据库模型定义
使用 SQLAlchemy ORM 定义数据表结构
"""
from sqlalchemy import create_engine, Column, Integer, String, Float, DateTime, Boolean, Text
from sqlalchemy.ext.declarative import declarative_base
from sqlalchemy.orm import sessionmaker, Session
from datetime import datetime
import logging
logger = logging.getLogger(__name__)
# 声明基类
Base = declarative_base()
# ============================================
# 用户表
# ============================================
class User(Base):
"""用户表"""
__tablename__ = "users"
id = Column(Integer, primary_key=True, index=True)
username = Column(String(50), unique=True, index=True, nullable=False)
hashed_password = Column(String(200), nullable=False)
email = Column(String(100), unique=True, index=True)
full_name = Column(String(100))
is_active = Column(Boolean, default=True)
is_superuser = Column(Boolean, default=False)
created_at = Column(DateTime, default=datetime.utcnow)
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
# ============================================
# API Token 表
# ============================================
class APIToken(Base):
"""API Token 表"""
__tablename__ = "api_tokens"
id = Column(Integer, primary_key=True, index=True)
user_id = Column(Integer, nullable=False)
token = Column(String(200), unique=True, index=True, nullable=False)
name = Column(String(100))
is_active = Column(Boolean, default=True)
expires_at = Column(DateTime)
created_at = Column(DateTime, default=datetime.utcnow)
# ============================================
# 交易日志表
# ============================================
class TradeLog(Base):
"""交易日志表"""
__tablename__ = "trade_logs"
id = Column(Integer, primary_key=True, index=True)
order_id = Column(String(50), index=True)
symbol = Column(String(20), index=True)
exchange = Column(String(20))
direction = Column(String(10))
order_type = Column(String(20))
volume = Column(Float)
price = Column(Float)
traded_volume = Column(Float)
traded_price = Column(Float)
status = Column(String(20), index=True)
time = Column(DateTime, index=True)
created_at = Column(DateTime, default=datetime.utcnow)
# ============================================
# 策略日志表
# ============================================
class StrategyLog(Base):
"""策略日志表"""
__tablename__ = "strategy_logs"
id = Column(Integer, primary_key=True, index=True)
strategy_name = Column(String(50), index=True)
level = Column(String(20)) # INFO, WARNING, ERROR
message = Column(Text)
created_at = Column(DateTime, default=datetime.utcnow, index=True)
# ============================================
# 系统配置表
# ============================================
class SystemConfig(Base):
"""系统配置表"""
__tablename__ = "system_configs"
id = Column(Integer, primary_key=True, index=True)
key = Column(String(100), unique=True, index=True, nullable=False)
value = Column(Text)
description = Column(String(200))
updated_at = Column(DateTime, default=datetime.utcnow, onupdate=datetime.utcnow)
# ============================================
# 数据库初始化
# ============================================
# 全局数据库引擎和会话
_engine = None
_SessionLocal = None
def init_database(database_url: str = "sqlite:///sanguo_web.db") -> None:
"""
初始化数据库
- **database_url**: 数据库连接字符串
"""
global _engine, _SessionLocal
logger.info(f"Initializing database: {database_url}")
_engine = create_engine(
database_url,
connect_args={"check_same_thread": False} if database_url.startswith("sqlite") else {},
echo=False
)
_SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=_engine)
# 创建所有表
Base.metadata.create_all(bind=_engine)
logger.info("Database initialized successfully")
def get_db() -> Session:
"""
获取数据库会话
用于依赖注入
"""
global _SessionLocal
if _SessionLocal is None:
init_database()
db = _SessionLocal()
try:
yield db
finally:
db.close()
def get_engine():
"""获取数据库引擎"""
if _engine is None:
init_database()
return _engine
__all__ = [
"User",
"APIToken",
"TradeLog",
"StrategyLog",
"SystemConfig",
"init_database",
"get_db",
"get_engine",
]