import os from sqlalchemy import create_engine, event from sqlalchemy.orm import sessionmaker, declarative_base, Session BASE_DIR = os.path.dirname(os.path.abspath(__file__)) DB_FILE = os.path.join(BASE_DIR, "avatar.db") DATABASE_URL = os.getenv("DATABASE_URL", f"sqlite:///{DB_FILE}") IS_SQLITE = DATABASE_URL.startswith("sqlite:") engine = create_engine( DATABASE_URL, connect_args={"check_same_thread": False, "timeout": 30} if IS_SQLITE else {}, ) if IS_SQLITE: @event.listens_for(engine, "connect") def _configure_sqlite_connection(dbapi_connection, _connection_record): cursor = dbapi_connection.cursor() try: cursor.execute("PRAGMA synchronous=NORMAL") cursor.execute("PRAGMA busy_timeout=30000") finally: cursor.close() SessionLocal = sessionmaker(bind=engine, autoflush=False, expire_on_commit=False) Base = declarative_base() def get_db(): db = SessionLocal() try: yield db finally: db.close() def init_db(): import models if IS_SQLITE: with engine.connect() as conn: conn.exec_driver_sql("PRAGMA journal_mode=WAL") conn.commit() Base.metadata.create_all(bind=engine) # 轻量迁移:为已存在的表补充新列(SQLite 不支持自动 ALTER,逐列尝试) _try_add_columns( ("qa_pairs", "enabled", "BOOLEAN DEFAULT 1"), ("knowledge_docs", "vectorized", "BOOLEAN DEFAULT 0"), ("knowledge_docs", "embedding_model", "VARCHAR DEFAULT ''"), ("knowledge_docs", "chunk_count", "INTEGER DEFAULT 0"), ("knowledge_docs", "vectorized_at", "TIMESTAMP"), ("knowledge_docs", "error_message", "VARCHAR DEFAULT ''"), ("avatars", "owner_id", "VARCHAR DEFAULT ''"), ("authorizations", "takeover_enabled", "BOOLEAN DEFAULT 0"), ("authorizations", "takeover_mode", "VARCHAR DEFAULT 'immediate'"), ("authorizations", "takeover_delay_seconds", "INTEGER DEFAULT 180"), ("avatars", "share_token", "VARCHAR DEFAULT NULL"), ("token_account", "user_id", "VARCHAR DEFAULT ''"), ("token_account", "total_granted", "BIGINT DEFAULT 0"), ("token_account", "total_consumed", "BIGINT DEFAULT 0"), ("token_account", "created_at", "TIMESTAMP"), ("token_account", "updated_at", "TIMESTAMP"), ("takeover_messages", "attachment_id", "VARCHAR DEFAULT NULL"), ) _normalize_optional_unique_values() _normalize_takeover_delays() _create_token_indexes() def _try_add_columns(*cols): with engine.connect() as conn: for table, col, ddl in cols: try: conn.exec_driver_sql(f"ALTER TABLE {table} ADD COLUMN {col} {ddl}") conn.commit() except Exception: # 列已存在(或全新库由 create_all 建好)则忽略 pass def _normalize_optional_unique_values(): with engine.begin() as conn: conn.exec_driver_sql("UPDATE avatars SET share_token = NULL WHERE share_token = ''") def _normalize_takeover_delays(): with engine.begin() as conn: # The old 30-second column default was never wired into the scheduler. conn.exec_driver_sql( "UPDATE authorizations SET takeover_delay_seconds = 180 " "WHERE takeover_delay_seconds IS NULL OR takeover_delay_seconds = 30" ) def _create_token_indexes(): with engine.begin() as conn: conn.exec_driver_sql( "CREATE UNIQUE INDEX IF NOT EXISTS ux_token_account_user_id " "ON token_account(user_id) WHERE user_id <> ''" )