52 lines
1.4 KiB
Python
52 lines
1.4 KiB
Python
"""初始化数据库结构 + 预置分类。
|
||
|
||
用法:python scripts/init_db.py
|
||
"""
|
||
from __future__ import annotations
|
||
|
||
import sys
|
||
from pathlib import Path
|
||
|
||
sys.path.insert(0, str(Path(__file__).resolve().parent.parent))
|
||
|
||
from app.constants import DEFAULT_CATEGORY_COLORS # noqa: E402
|
||
from app.db import SessionLocal, engine # noqa: E402
|
||
from app.models import Base, Category # noqa: E402
|
||
from app.schema_patch import ensure_columns # noqa: E402
|
||
|
||
DEFAULT_CATEGORIES = [
|
||
("平台产品", 0),
|
||
("AI 能力", 1),
|
||
("智能装备", 2),
|
||
("定制项目", 3),
|
||
("内部运营", 4),
|
||
]
|
||
|
||
|
||
def main():
|
||
Base.metadata.create_all(bind=engine)
|
||
added = ensure_columns(engine, Base.metadata)
|
||
if added:
|
||
print(f"[ok] 已补齐字段:{', '.join(added)}")
|
||
db = SessionLocal()
|
||
try:
|
||
for idx, (name, order) in enumerate(DEFAULT_CATEGORIES):
|
||
if db.query(Category).filter(Category.name == name).one_or_none():
|
||
continue
|
||
db.add(
|
||
Category(
|
||
name=name,
|
||
color=DEFAULT_CATEGORY_COLORS[idx % len(DEFAULT_CATEGORY_COLORS)],
|
||
sort_order=order,
|
||
)
|
||
)
|
||
db.commit()
|
||
print(f"[ok] 数据库已初始化:{engine.url}")
|
||
print(f"[ok] 预置分类:{', '.join(n for n, _ in DEFAULT_CATEGORIES)}")
|
||
finally:
|
||
db.close()
|
||
|
||
|
||
if __name__ == "__main__":
|
||
main()
|