from sqlalchemy import create_engine from sqlalchemy.ext.declarative import declarative_base from sqlalchemy.orm import sessionmaker from dotenv import load_dotenv import os import pymysql # 加载环境变量 load_dotenv() # 获取数据库连接 URL DATABASE_URL = os.getenv("DATABASE_URL") # 处理MySQL数据库自动创建 if DATABASE_URL.startswith("mysql"): # 解析数据库连接信息 import re match = re.match(r"mysql\+pymysql://([^:]+):([^@]+)@([^:]+):(\d+)/(.+)", DATABASE_URL) if match: user, password, host, port, db_name = match.groups() port = int(port) # 先连接到MySQL服务器 try: # 连接到MySQL服务器(不指定数据库) conn = pymysql.connect( host=host, port=port, user=user, password=password, charset='utf8mb4' ) cursor = conn.cursor() # 检查数据库是否存在 cursor.execute(f"SHOW DATABASES LIKE '{db_name}'") result = cursor.fetchone() # 如果数据库不存在,创建它 if not result: cursor.execute(f"CREATE DATABASE {db_name} CHARACTER SET utf8mb4 COLLATE utf8mb4_unicode_ci") print(f"数据库 {db_name} 已创建") cursor.close() conn.close() except Exception as e: print(f"创建数据库时出错: {e}") # 创建数据库引擎 # 对于SQLite,需要添加check_same_thread参数 if DATABASE_URL.startswith("sqlite"): engine = create_engine(DATABASE_URL, connect_args={"check_same_thread": False}) else: # 对于MySQL,添加charset参数确保utf8编码 if DATABASE_URL.startswith("mysql"): engine = create_engine(DATABASE_URL + "?charset=utf8mb4") else: engine = create_engine(DATABASE_URL) # 创建会话工厂 SessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine) # 创建基类 Base = declarative_base() # 检查表结构并自动更新 def check_and_update_tables(): """检查表结构是否与models一致,不一致则自动更新""" try: # 导入所有模型,确保它们被注册到 Base from models.product import Product from models.company import Company from models.user import User from models.category import Category from models.membership import Membership from models.operation_log import OperationLog from models.contact import Contact from models.template import Template from models.member_level import MemberLevel from models.customer_relationship import CustomerRelationship from models.product_certification import ProductCertification from models.product_media import ProductMedia from models.product_packaging import ProductPackaging from models.company_r2_config import CompanyR2Config from models.media import Media, MediaTag, MediaTagMapping, MediaDirectory # 检查数据库连接 with engine.connect() as conn: print("数据库连接成功") # 自动创建不存在的表 Base.metadata.create_all(bind=engine) print("表结构检查完成,不存在的表已创建") # 注意:SQLAlchemy的create_all()只会创建不存在的表,不会更新已存在表的结构 # 对于表结构的更新,建议使用Alembic迁移工具 # print("提示:已存在表的结构更新需要使用Alembic迁移工具") except Exception as e: print(f"检查表结构时出错: {e}") # 调用检查表结构的函数 check_and_update_tables() # 依赖项,用于获取数据库会话 def get_db(): db = SessionLocal() try: yield db finally: db.close()