fix(migrate): 修复用户字段迁移脚本,优化异常处理和字段检查逻辑
- 更新迁移脚本以确保在数据库中添加缺失的字段(user_id、phone_number、avatar) - 优化异常处理,确保在迁移过程中发生错误时能够正确回滚 - 统一字符串引号格式,提升代码可读性 - 添加必要的导入语句,确保脚本正常运行
This commit is contained in:
parent
767a3d4762
commit
a337cba3f7
@ -5,17 +5,19 @@
|
||||
将现有的 username 作为 user_id 的默认值
|
||||
"""
|
||||
|
||||
import os
|
||||
# ruff: noqa: E402
|
||||
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
from sqlalchemy import text
|
||||
|
||||
# 添加项目根目录到Python路径
|
||||
PROJECT_ROOT = Path(__file__).parent
|
||||
sys.path.insert(0, str(PROJECT_ROOT))
|
||||
|
||||
from sqlalchemy import text
|
||||
from server.db_manager import db_manager
|
||||
from server.models.user_model import User
|
||||
from server.models.user_model import User as User
|
||||
|
||||
|
||||
def migrate_user_fields():
|
||||
@ -41,15 +43,15 @@ def migrate_user_fields():
|
||||
print(f"现有字段: {existing_columns}")
|
||||
|
||||
# 添加缺失的字段
|
||||
if 'user_id' not in existing_columns:
|
||||
if "user_id" not in existing_columns:
|
||||
print("添加 user_id 字段...")
|
||||
db.execute(text("ALTER TABLE users ADD COLUMN user_id VARCHAR(255)"))
|
||||
|
||||
if 'phone_number' not in existing_columns:
|
||||
if "phone_number" not in existing_columns:
|
||||
print("添加 phone_number 字段...")
|
||||
db.execute(text("ALTER TABLE users ADD COLUMN phone_number VARCHAR(255)"))
|
||||
|
||||
if 'avatar' not in existing_columns:
|
||||
if "avatar" not in existing_columns:
|
||||
print("添加 avatar 字段...")
|
||||
db.execute(text("ALTER TABLE users ADD COLUMN avatar VARCHAR(500)"))
|
||||
|
||||
@ -82,10 +84,7 @@ def migrate_user_fields():
|
||||
for user_id, username in users_without_user_id:
|
||||
# 将 username 作为默认的 user_id
|
||||
print(f"为用户 {username} (ID: {user_id}) 设置 user_id: {username}")
|
||||
db.execute(
|
||||
text("UPDATE users SET user_id = :user_id WHERE id = :id"),
|
||||
{"user_id": username, "id": user_id}
|
||||
)
|
||||
db.execute(text("UPDATE users SET user_id = :user_id WHERE id = :id"), {"user_id": username, "id": user_id})
|
||||
|
||||
db.commit()
|
||||
|
||||
@ -96,13 +95,18 @@ def migrate_user_fields():
|
||||
try:
|
||||
db.execute(text("CREATE UNIQUE INDEX idx_users_user_id ON users(user_id)"))
|
||||
print("创建 user_id 唯一索引")
|
||||
except:
|
||||
except Exception:
|
||||
print("user_id 索引可能已存在")
|
||||
|
||||
try:
|
||||
db.execute(text("CREATE UNIQUE INDEX idx_users_phone_number ON users(phone_number) WHERE phone_number IS NOT NULL"))
|
||||
db.execute(
|
||||
text(
|
||||
"CREATE UNIQUE INDEX idx_users_phone_number ON users(phone_number) "
|
||||
"WHERE phone_number IS NOT NULL"
|
||||
)
|
||||
)
|
||||
print("创建 phone_number 唯一索引")
|
||||
except:
|
||||
except Exception:
|
||||
print("phone_number 索引可能已存在")
|
||||
|
||||
db.commit()
|
||||
@ -114,7 +118,9 @@ def migrate_user_fields():
|
||||
# 4. 验证迁移结果
|
||||
print("验证迁移结果...")
|
||||
total_users = db.execute(text("SELECT COUNT(*) FROM users")).scalar()
|
||||
users_with_user_id = db.execute(text("SELECT COUNT(*) FROM users WHERE user_id IS NOT NULL AND user_id != ''")).scalar()
|
||||
users_with_user_id = db.execute(
|
||||
text("SELECT COUNT(*) FROM users WHERE user_id IS NOT NULL AND user_id != ''")
|
||||
).scalar()
|
||||
|
||||
print(f"总用户数: {total_users}")
|
||||
print(f"已设置 user_id 的用户数: {users_with_user_id}")
|
||||
@ -126,13 +132,13 @@ def migrate_user_fields():
|
||||
|
||||
except Exception as e:
|
||||
print(f"迁移过程中发生错误: {e}")
|
||||
if 'db' in locals():
|
||||
if "db" in locals():
|
||||
db.rollback()
|
||||
raise
|
||||
finally:
|
||||
if 'db' in locals():
|
||||
if "db" in locals():
|
||||
db.close()
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
migrate_user_fields()
|
||||
migrate_user_fields()
|
||||
|
||||
@ -7,7 +7,7 @@ from sqlalchemy.orm import sessionmaker
|
||||
|
||||
from server.models import Base
|
||||
from server.models.user_model import User
|
||||
from server.utils.migrate import check_and_migrate, validate_database_schema
|
||||
from server.utils.migrate import validate_database_schema
|
||||
from src import config
|
||||
from src.utils import logger
|
||||
|
||||
@ -58,7 +58,6 @@ class DBManager:
|
||||
logger.warning("请运行以下 scripts/migrate_user_fields.py 来修复数据库结构:")
|
||||
logger.warning("=" * 60)
|
||||
|
||||
|
||||
def get_session(self):
|
||||
"""获取数据库会话"""
|
||||
return self.Session()
|
||||
|
||||
@ -185,7 +185,7 @@ async def initialize_admin(admin_data: InitializeAdmin, db: Session = Depends(ge
|
||||
hashed_password = AuthUtils.hash_password(admin_data.password)
|
||||
|
||||
# 验证用户ID格式(只支持字母数字和下划线)
|
||||
if not re.match(r'^[a-zA-Z0-9_]+$', admin_data.user_id):
|
||||
if not re.match(r"^[a-zA-Z0-9_]+$", admin_data.user_id):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="用户ID只能包含字母、数字和下划线",
|
||||
@ -199,10 +199,7 @@ async def initialize_admin(admin_data: InitializeAdmin, db: Session = Depends(ge
|
||||
|
||||
# 验证手机号格式(如果提供了)
|
||||
if admin_data.phone_number and not is_valid_phone_number(admin_data.phone_number):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="手机号格式不正确"
|
||||
)
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="手机号格式不正确")
|
||||
|
||||
# 由于是首次初始化,直接使用输入的user_id
|
||||
user_id = admin_data.user_id
|
||||
@ -214,7 +211,7 @@ async def initialize_admin(admin_data: InitializeAdmin, db: Session = Depends(ge
|
||||
avatar=None, # 初始化时头像为空
|
||||
password_hash=hashed_password,
|
||||
role="superadmin",
|
||||
last_login=datetime.now()
|
||||
last_login=datetime.now(),
|
||||
)
|
||||
|
||||
db.add(new_admin)
|
||||
@ -257,7 +254,7 @@ async def update_profile(
|
||||
profile_data: UserProfileUpdate,
|
||||
request: Request,
|
||||
current_user: User = Depends(get_required_user),
|
||||
db: Session = Depends(get_db)
|
||||
db: Session = Depends(get_db),
|
||||
):
|
||||
"""更新当前用户的个人资料"""
|
||||
update_details = []
|
||||
@ -266,22 +263,17 @@ async def update_profile(
|
||||
if profile_data.phone_number is not None:
|
||||
# 如果手机号不为空,验证格式
|
||||
if profile_data.phone_number and not is_valid_phone_number(profile_data.phone_number):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="手机号格式不正确"
|
||||
)
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="手机号格式不正确")
|
||||
|
||||
# 检查手机号是否已被其他用户使用
|
||||
if profile_data.phone_number:
|
||||
existing_phone = db.query(User).filter(
|
||||
User.phone_number == profile_data.phone_number,
|
||||
User.id != current_user.id
|
||||
).first()
|
||||
existing_phone = (
|
||||
db.query(User)
|
||||
.filter(User.phone_number == profile_data.phone_number, User.id != current_user.id)
|
||||
.first()
|
||||
)
|
||||
if existing_phone:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="手机号已被其他用户使用"
|
||||
)
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="手机号已被其他用户使用")
|
||||
|
||||
current_user.phone_number = profile_data.phone_number
|
||||
update_details.append(f"手机号: {profile_data.phone_number or '已清空'}")
|
||||
@ -290,13 +282,7 @@ async def update_profile(
|
||||
|
||||
# 记录操作
|
||||
if update_details:
|
||||
log_operation(
|
||||
db,
|
||||
current_user.id,
|
||||
"更新个人资料",
|
||||
f"更新个人资料: {', '.join(update_details)}",
|
||||
request
|
||||
)
|
||||
log_operation(db, current_user.id, "更新个人资料", f"更新个人资料: {', '.join(update_details)}", request)
|
||||
|
||||
return current_user.to_dict()
|
||||
|
||||
@ -363,7 +349,7 @@ async def create_user(
|
||||
user_id=user_id,
|
||||
phone_number=user_data.phone_number,
|
||||
password_hash=hashed_password,
|
||||
role=user_data.role
|
||||
role=user_data.role,
|
||||
)
|
||||
|
||||
db.add(new_user)
|
||||
@ -508,9 +494,7 @@ async def delete_user(
|
||||
# 路由:验证用户名并生成user_id
|
||||
@auth.post("/validate-username", response_model=UserIdGeneration)
|
||||
async def validate_username_and_generate_user_id(
|
||||
validation_data: UsernameValidation,
|
||||
current_user: User = Depends(get_admin_user),
|
||||
db: Session = Depends(get_db)
|
||||
validation_data: UsernameValidation, current_user: User = Depends(get_admin_user), db: Session = Depends(get_db)
|
||||
):
|
||||
"""验证用户名格式并生成可用的user_id"""
|
||||
# 验证用户名格式
|
||||
@ -533,42 +517,28 @@ async def validate_username_and_generate_user_id(
|
||||
existing_user_ids = [user.user_id for user in db.query(User.user_id).all()]
|
||||
user_id = generate_unique_user_id(validation_data.username, existing_user_ids)
|
||||
|
||||
return UserIdGeneration(
|
||||
username=validation_data.username,
|
||||
user_id=user_id,
|
||||
is_available=True
|
||||
)
|
||||
return UserIdGeneration(username=validation_data.username, user_id=user_id, is_available=True)
|
||||
|
||||
|
||||
# 路由:检查user_id是否可用
|
||||
@auth.get("/check-user-id/{user_id}")
|
||||
async def check_user_id_availability(
|
||||
user_id: str,
|
||||
current_user: User = Depends(get_admin_user),
|
||||
db: Session = Depends(get_db)
|
||||
user_id: str, current_user: User = Depends(get_admin_user), db: Session = Depends(get_db)
|
||||
):
|
||||
"""检查user_id是否可用"""
|
||||
existing_user = db.query(User).filter(User.user_id == user_id).first()
|
||||
return {
|
||||
"user_id": user_id,
|
||||
"is_available": existing_user is None
|
||||
}
|
||||
return {"user_id": user_id, "is_available": existing_user is None}
|
||||
|
||||
|
||||
# 路由:上传用户头像
|
||||
@auth.post("/upload-avatar")
|
||||
async def upload_user_avatar(
|
||||
file: UploadFile = File(...),
|
||||
current_user: User = Depends(get_required_user),
|
||||
db: Session = Depends(get_db)
|
||||
file: UploadFile = File(...), current_user: User = Depends(get_required_user), db: Session = Depends(get_db)
|
||||
):
|
||||
"""上传用户头像"""
|
||||
# 检查文件类型
|
||||
if not file.content_type or not file.content_type.startswith('image/'):
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="只能上传图片文件"
|
||||
)
|
||||
if not file.content_type or not file.content_type.startswith("image/"):
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="只能上传图片文件")
|
||||
|
||||
# 检查文件大小(5MB限制)
|
||||
file_size = 0
|
||||
@ -576,14 +546,11 @@ async def upload_user_avatar(
|
||||
file_size = len(file_content)
|
||||
|
||||
if file_size > 5 * 1024 * 1024: # 5MB
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_400_BAD_REQUEST,
|
||||
detail="文件大小不能超过5MB"
|
||||
)
|
||||
raise HTTPException(status_code=status.HTTP_400_BAD_REQUEST, detail="文件大小不能超过5MB")
|
||||
|
||||
try:
|
||||
# 获取文件扩展名
|
||||
file_extension = file.filename.split('.')[-1].lower() if file.filename and '.' in file.filename else 'jpg'
|
||||
file_extension = file.filename.split(".")[-1].lower() if file.filename and "." in file.filename else "jpg"
|
||||
|
||||
# 上传到MinIO
|
||||
avatar_url = upload_image_to_minio(file_content, file_extension)
|
||||
@ -595,14 +562,7 @@ async def upload_user_avatar(
|
||||
# 记录操作
|
||||
log_operation(db, current_user.id, "上传头像", f"更新头像: {avatar_url}")
|
||||
|
||||
return {
|
||||
"success": True,
|
||||
"avatar_url": avatar_url,
|
||||
"message": "头像上传成功"
|
||||
}
|
||||
return {"success": True, "avatar_url": avatar_url, "message": "头像上传成功"}
|
||||
|
||||
except Exception as e:
|
||||
raise HTTPException(
|
||||
status_code=status.HTTP_500_INTERNAL_SERVER_ERROR,
|
||||
detail=f"头像上传失败: {str(e)}"
|
||||
)
|
||||
raise HTTPException(status_code=status.HTTP_500_INTERNAL_SERVER_ERROR, detail=f"头像上传失败: {str(e)}")
|
||||
|
||||
@ -230,7 +230,7 @@ async def download_document(db_id: str, doc_id: str, request: Request, current_u
|
||||
|
||||
# 解码URL编码的文件名(如果有的话)
|
||||
try:
|
||||
decoded_filename = unquote(filename, encoding='utf-8')
|
||||
decoded_filename = unquote(filename, encoding="utf-8")
|
||||
logger.debug(f"Decoded filename: {decoded_filename}")
|
||||
except Exception as e:
|
||||
logger.debug(f"Failed to decode filename {filename}: {e}")
|
||||
@ -276,21 +276,18 @@ async def download_document(db_id: str, doc_id: str, request: Request, current_u
|
||||
media_type = media_types.get(ext.lower(), "application/octet-stream")
|
||||
|
||||
# 创建自定义FileResponse,避免文件名编码问题
|
||||
response = StarletteFileResponse(
|
||||
path=file_path,
|
||||
media_type=media_type
|
||||
)
|
||||
response = StarletteFileResponse(path=file_path, media_type=media_type)
|
||||
|
||||
# 正确处理中文文件名的HTTP头部设置
|
||||
# HTTP头部只能包含ASCII字符,所以需要对中文文件名进行编码
|
||||
try:
|
||||
# 尝试使用ASCII编码(适用于英文文件名)
|
||||
decoded_filename.encode('ascii')
|
||||
decoded_filename.encode("ascii")
|
||||
# 如果成功,直接使用简单格式
|
||||
response.headers["Content-Disposition"] = f'attachment; filename="{decoded_filename}"'
|
||||
except UnicodeEncodeError:
|
||||
# 如果包含非ASCII字符(如中文),使用RFC 2231格式
|
||||
encoded_filename = quote(decoded_filename.encode('utf-8'))
|
||||
encoded_filename = quote(decoded_filename.encode("utf-8"))
|
||||
response.headers["Content-Disposition"] = f"attachment; filename*=UTF-8''{encoded_filename}"
|
||||
|
||||
return response
|
||||
|
||||
@ -277,18 +277,32 @@ def validate_database_schema(db_path: str) -> tuple[bool, list[str]]:
|
||||
|
||||
# 检查users表必需字段
|
||||
required_fields = {
|
||||
'users': ['id', 'username', 'user_id', 'phone_number', 'avatar',
|
||||
'password_hash', 'role', 'created_at', 'last_login',
|
||||
'login_failed_count', 'last_failed_login', 'login_locked_until'],
|
||||
'operation_logs': ['id', 'user_id', 'operation', 'details', 'ip_address', 'timestamp']
|
||||
"users": [
|
||||
"id",
|
||||
"username",
|
||||
"user_id",
|
||||
"phone_number",
|
||||
"avatar",
|
||||
"password_hash",
|
||||
"role",
|
||||
"created_at",
|
||||
"last_login",
|
||||
"login_failed_count",
|
||||
"last_failed_login",
|
||||
"login_locked_until",
|
||||
],
|
||||
"operation_logs": ["id", "user_id", "operation", "details", "ip_address", "timestamp"],
|
||||
}
|
||||
|
||||
for table_name, fields in required_fields.items():
|
||||
# 检查表是否存在
|
||||
cursor.execute("""
|
||||
cursor.execute(
|
||||
"""
|
||||
SELECT name FROM sqlite_master
|
||||
WHERE type='table' AND name=?
|
||||
""", (table_name,))
|
||||
""",
|
||||
(table_name,),
|
||||
)
|
||||
|
||||
if not cursor.fetchone():
|
||||
missing_fields.append(f"表 {table_name} 不存在")
|
||||
@ -308,7 +322,7 @@ def validate_database_schema(db_path: str) -> tuple[bool, list[str]]:
|
||||
logger.error(f"验证数据库结构失败: {e}")
|
||||
return False, [f"验证失败: {str(e)}"]
|
||||
finally:
|
||||
if 'conn' in locals():
|
||||
if "conn" in locals():
|
||||
conn.close()
|
||||
|
||||
|
||||
|
||||
@ -14,7 +14,7 @@ from src.knowledge.kb_utils import (
|
||||
split_text_into_qa_chunks,
|
||||
)
|
||||
from src.knowledge.knowledge_base import KnowledgeBase
|
||||
from src.models.embedding import OtherEmbedding
|
||||
from src.models.embed import OtherEmbedding
|
||||
from src.utils import hashstr, logger
|
||||
|
||||
MILVUS_AVAILABLE = True
|
||||
|
||||
@ -1,4 +1,4 @@
|
||||
from src.models.chat_model import get_custom_model, select_model
|
||||
from src.models.embedding import select_embedding_model
|
||||
from src.models.chat import get_custom_model, select_model
|
||||
from src.models.embed import select_embedding_model
|
||||
|
||||
__all__ = ["select_model", "select_embedding_model", "get_custom_model"]
|
||||
|
||||
Loading…
Reference in New Issue
Block a user