fix(migrate): 修复用户字段迁移脚本,优化异常处理和字段检查逻辑

- 更新迁移脚本以确保在数据库中添加缺失的字段(user_id、phone_number、avatar)
- 优化异常处理,确保在迁移过程中发生错误时能够正确回滚
- 统一字符串引号格式,提升代码可读性
- 添加必要的导入语句,确保脚本正常运行
This commit is contained in:
Wenjie Zhang 2025-09-22 17:18:07 +08:00
parent 767a3d4762
commit a337cba3f7
10 changed files with 76 additions and 100 deletions

View File

@ -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()

View File

@ -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()

View File

@ -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)}")

View File

@ -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

View File

@ -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()

View File

@ -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

View File

@ -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"]