From a337cba3f73794c8b021c67bb611b11b63e5fb4f Mon Sep 17 00:00:00 2001 From: Wenjie Zhang Date: Mon, 22 Sep 2025 17:18:07 +0800 Subject: [PATCH] =?UTF-8?q?fix(migrate):=20=E4=BF=AE=E5=A4=8D=E7=94=A8?= =?UTF-8?q?=E6=88=B7=E5=AD=97=E6=AE=B5=E8=BF=81=E7=A7=BB=E8=84=9A=E6=9C=AC?= =?UTF-8?q?=EF=BC=8C=E4=BC=98=E5=8C=96=E5=BC=82=E5=B8=B8=E5=A4=84=E7=90=86?= =?UTF-8?q?=E5=92=8C=E5=AD=97=E6=AE=B5=E6=A3=80=E6=9F=A5=E9=80=BB=E8=BE=91?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - 更新迁移脚本以确保在数据库中添加缺失的字段(user_id、phone_number、avatar) - 优化异常处理,确保在迁移过程中发生错误时能够正确回滚 - 统一字符串引号格式,提升代码可读性 - 添加必要的导入语句,确保脚本正常运行 --- scripts/migrate_user_fields.py | 40 ++++++----- server/db_manager.py | 3 +- server/routers/auth_router.py | 88 +++++++---------------- server/routers/knowledge_router.py | 11 ++- server/utils/migrate.py | 28 ++++++-- src/knowledge/milvus_kb.py | 2 +- src/models/__init__.py | 4 +- src/models/{chat_model.py => chat.py} | 0 src/models/{embedding.py => embed.py} | 0 src/models/{rerank_model.py => rerank.py} | 0 10 files changed, 76 insertions(+), 100 deletions(-) rename src/models/{chat_model.py => chat.py} (100%) rename src/models/{embedding.py => embed.py} (100%) rename src/models/{rerank_model.py => rerank.py} (100%) diff --git a/scripts/migrate_user_fields.py b/scripts/migrate_user_fields.py index fb051f9f..5cd934ef 100644 --- a/scripts/migrate_user_fields.py +++ b/scripts/migrate_user_fields.py @@ -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() \ No newline at end of file + migrate_user_fields() diff --git a/server/db_manager.py b/server/db_manager.py index 9cc4fa38..dcb71d12 100644 --- a/server/db_manager.py +++ b/server/db_manager.py @@ -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() diff --git a/server/routers/auth_router.py b/server/routers/auth_router.py index 7895a50e..a1eaf41d 100644 --- a/server/routers/auth_router.py +++ b/server/routers/auth_router.py @@ -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)}") diff --git a/server/routers/knowledge_router.py b/server/routers/knowledge_router.py index 0a50a49f..cd01dcfb 100644 --- a/server/routers/knowledge_router.py +++ b/server/routers/knowledge_router.py @@ -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 diff --git a/server/utils/migrate.py b/server/utils/migrate.py index b692a117..e970ddc2 100644 --- a/server/utils/migrate.py +++ b/server/utils/migrate.py @@ -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() diff --git a/src/knowledge/milvus_kb.py b/src/knowledge/milvus_kb.py index 0cb685d8..0843cf86 100644 --- a/src/knowledge/milvus_kb.py +++ b/src/knowledge/milvus_kb.py @@ -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 diff --git a/src/models/__init__.py b/src/models/__init__.py index b01fbb2a..d530a04f 100644 --- a/src/models/__init__.py +++ b/src/models/__init__.py @@ -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"] diff --git a/src/models/chat_model.py b/src/models/chat.py similarity index 100% rename from src/models/chat_model.py rename to src/models/chat.py diff --git a/src/models/embedding.py b/src/models/embed.py similarity index 100% rename from src/models/embedding.py rename to src/models/embed.py diff --git a/src/models/rerank_model.py b/src/models/rerank.py similarity index 100% rename from src/models/rerank_model.py rename to src/models/rerank.py