182 lines
5.9 KiB
Python
182 lines
5.9 KiB
Python
#!/usr/bin/env python3
|
||
"""
|
||
数据库初始化脚本
|
||
用于创建数据库表结构
|
||
"""
|
||
|
||
import asyncio
|
||
import sys
|
||
import os
|
||
|
||
# 添加项目根目录到Python路径
|
||
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
|
||
|
||
from app.database import init_db, engine
|
||
from app.config import settings
|
||
from sqlalchemy import text
|
||
|
||
|
||
async def create_database():
|
||
"""创建数据库(如果不存在)"""
|
||
try:
|
||
# 从DATABASE_URL中提取数据库名
|
||
db_url = settings.database_url
|
||
if "postgresql+asyncpg://" in db_url:
|
||
# 解析数据库URL获取数据库名
|
||
# 格式: postgresql+asyncpg://user:password@host:port/dbname
|
||
import re
|
||
match = re.match(r'.+/([^/]+)$', db_url)
|
||
if match:
|
||
db_name = match.group(1)
|
||
# 创建不包含数据库名的连接URL
|
||
base_url = db_url.rsplit('/', 1)[0]
|
||
|
||
# 连接到postgres默认数据库创建新数据库
|
||
from sqlalchemy.ext.asyncio import create_async_engine
|
||
temp_engine = create_async_engine(f"{base_url}/postgres")
|
||
|
||
async with temp_engine.begin() as conn:
|
||
# 检查数据库是否已存在
|
||
result = await conn.execute(
|
||
text("SELECT 1 FROM pg_database WHERE datname = :db_name"),
|
||
{"db_name": db_name}
|
||
)
|
||
exists = result.scalar()
|
||
|
||
if not exists:
|
||
await conn.execute(text(f"CREATE DATABASE {db_name}"))
|
||
print(f"✅ 数据库 '{db_name}' 创建成功")
|
||
else:
|
||
print(f"✅ 数据库 '{db_name}' 已存在")
|
||
|
||
await temp_engine.dispose()
|
||
|
||
except Exception as e:
|
||
print(f"❌ 创建数据库失败: {e}")
|
||
return False
|
||
|
||
return True
|
||
|
||
|
||
async def init_tables():
|
||
"""初始化数据库表结构"""
|
||
try:
|
||
print("🔄 开始初始化数据库表结构...")
|
||
|
||
# 初始化所有表
|
||
await init_db()
|
||
|
||
print("✅ 数据库表结构初始化成功")
|
||
|
||
# 验证表是否创建成功
|
||
async with engine.begin() as conn:
|
||
result = await conn.execute(text("""
|
||
SELECT table_name
|
||
FROM information_schema.tables
|
||
WHERE table_schema = 'public'
|
||
ORDER BY table_name
|
||
"""))
|
||
tables = [row[0] for row in result.fetchall()]
|
||
|
||
print(f"📋 已创建的表: {', '.join(tables)}")
|
||
|
||
# 检查必要的表
|
||
required_tables = ['callback_failure_logs', 'callback_failure_data', 'external_api_logs']
|
||
missing_tables = [table for table in required_tables if table not in tables]
|
||
|
||
if missing_tables:
|
||
print(f"❌ 缺少必要的表: {', '.join(missing_tables)}")
|
||
return False
|
||
else:
|
||
print("✅ 所有必要的表都已创建")
|
||
|
||
# 检查手机号索引
|
||
result = await conn.execute(text("""
|
||
SELECT indexname
|
||
FROM pg_indexes
|
||
WHERE tablename = 'callback_failure_data'
|
||
AND indexname = 'idx_phone_number'
|
||
"""))
|
||
phone_index = result.fetchone()
|
||
|
||
if phone_index:
|
||
print("✅ 手机号索引 'idx_phone_number' 已创建")
|
||
else:
|
||
print("⚠️ 手机号索引 'idx_phone_number' 未找到")
|
||
|
||
return True
|
||
|
||
except Exception as e:
|
||
print(f"❌ 初始化数据库表结构失败: {e}")
|
||
return False
|
||
|
||
|
||
async def verify_connection():
|
||
"""验证数据库连接"""
|
||
try:
|
||
print("🔄 验证数据库连接...")
|
||
|
||
async with engine.begin() as conn:
|
||
result = await conn.execute(text("SELECT version()"))
|
||
version = result.scalar()
|
||
print(f"✅ 数据库连接成功")
|
||
print(f"📊 PostgreSQL版本: {version}")
|
||
|
||
return True
|
||
|
||
except Exception as e:
|
||
print(f"❌ 数据库连接失败: {e}")
|
||
print("\n💡 请检查以下配置:")
|
||
print(f" - DATABASE_URL: {settings.database_url}")
|
||
print(" - PostgreSQL服务是否运行")
|
||
print(" - 用户名密码是否正确")
|
||
print(" - 网络连接是否正常")
|
||
return False
|
||
|
||
|
||
async def main():
|
||
"""主函数"""
|
||
print("🚀 开始数据库初始化...")
|
||
print(f"📋 配置信息:")
|
||
print(f" - 数据库URL: {settings.database_url}")
|
||
print(f" - Celery Broker URL: {settings.celery_broker_url}")
|
||
print(f" - 应用名称: {settings.app_name}")
|
||
print(f" - 外部API URL: {settings.external_api_url}")
|
||
print()
|
||
|
||
# 验证连接
|
||
if not await verify_connection():
|
||
sys.exit(1)
|
||
|
||
print()
|
||
|
||
# 创建数据库(如果需要)
|
||
if not await create_database():
|
||
sys.exit(1)
|
||
|
||
print()
|
||
|
||
# 初始化表结构
|
||
if not await init_tables():
|
||
sys.exit(1)
|
||
|
||
print()
|
||
print("🎉 数据库初始化完成!")
|
||
print("\n📝 下一步:")
|
||
print(" 1. 配置 .env 文件中的数据库连接信息")
|
||
print(" 2. 确保 Redis 服务正在运行(用于Celery)")
|
||
print(" 3. 启动 Celery Worker: python celery_worker.py")
|
||
print(" 4. 运行应用: python main.py")
|
||
print(f" 5. 访问API文档: http://localhost:8000/docs")
|
||
print(f" 6. 接口地址: POST /ai-talk/callback/{{siteId}}/failure")
|
||
|
||
|
||
if __name__ == "__main__":
|
||
try:
|
||
asyncio.run(main())
|
||
except KeyboardInterrupt:
|
||
print("\n⚠️ 用户中断操作")
|
||
sys.exit(1)
|
||
except Exception as e:
|
||
print(f"\n❌ 初始化过程中发生错误: {e}")
|
||
sys.exit(1) |