918bbed0fc
- 修复 deps.py 中 get_vn_service 的引用错误 (vn_service.vn_service -> vn_service) - 移除登录页面上的明文密码提示 - 改进前端错误处理,避免数据加载失败导致登录显示错误
174 lines
4.5 KiB
Python
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",
|
|
]
|