Files
doc-forge/backend/database.py
T
2026-07-02 18:26:50 +08:00

48 lines
2.1 KiB
Python

from sqlalchemy.ext.asyncio import create_async_engine, async_sessionmaker, AsyncSession
from sqlalchemy import inspect, text
from sqlalchemy.orm import DeclarativeBase
from config import settings
engine = create_async_engine(settings.DATABASE_URL, echo=settings.DEBUG, pool_size=10, max_overflow=20)
async_session = async_sessionmaker(engine, class_=AsyncSession, expire_on_commit=False)
class Base(DeclarativeBase):
pass
async def get_db():
async with async_session() as session:
try:
yield session
finally:
await session.close()
async def init_db():
from models.template import Template
from models.paragraph import Paragraph
from models.ai_model import AiModel
from models.document import Document
from models.generation_log import GenerationLog
from models.reference_file import ReferenceFile
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.create_all)
dialect_name = conn.dialect.name
columns = await conn.run_sync(lambda sync_conn: [column["name"] for column in inspect(sync_conn).get_columns("ai_models")])
if "supports_streaming" not in columns:
if dialect_name == "sqlite":
await conn.execute(text("ALTER TABLE ai_models ADD COLUMN supports_streaming BOOLEAN DEFAULT 0"))
else:
await conn.execute(text("ALTER TABLE ai_models ADD COLUMN supports_streaming TINYINT(1) DEFAULT 0"))
if "enable_reasoning" not in columns:
if dialect_name == "sqlite":
await conn.execute(text("ALTER TABLE ai_models ADD COLUMN enable_reasoning BOOLEAN DEFAULT 0"))
else:
await conn.execute(text("ALTER TABLE ai_models ADD COLUMN enable_reasoning TINYINT(1) DEFAULT 0"))
document_columns = await conn.run_sync(
lambda sync_conn: [column["name"] for column in inspect(sync_conn).get_columns("documents")]
)
if "request_payload_json" not in document_columns:
await conn.execute(text("ALTER TABLE documents ADD COLUMN request_payload_json TEXT"))