11from collections import defaultdict
22from typing import Optional
33from fastapi import APIRouter , File , Path , Query , UploadFile
4- from sqlmodel import SQLModel , or_ , select , delete as sqlmodel_delete
4+ from sqlmodel import SQLModel , case , or_ , select , delete as sqlmodel_delete
55from apps .system .crud .user import check_account_exists , check_email_exists , check_email_format , check_pwd_format , get_db_user , single_delete , user_ws_options
66from apps .system .crud .user_excel import batchUpload , downTemplate , download_error_file
77from apps .system .models .system_model import UserWsModel , WorkspaceModel
@@ -80,20 +80,36 @@ async def pager(
8080 if order_by and order_by != 'account' :
8181 select_columns .append (sort_field )
8282
83+ # 相似度排序:精确匹配 > 前缀匹配 > 包含匹配
84+ # 当有 keyword 时,将 similarity_score 加入 SELECT 列以满足 DISTINCT 约束
85+ similarity_score = None
86+ if keyword :
87+ similarity_score = case (
88+ (UserModel .account == keyword , 0 ),
89+ (UserModel .account .startswith (keyword ), 1 ),
90+ (UserModel .account .contains (keyword ), 2 ),
91+ else_ = 3
92+ )
93+ select_columns .append (similarity_score .label ('similarity_score' ))
94+
8395 origin_stmt = (
8496 select (* select_columns )
8597 .join (UserWsModel , UserModel .id == UserWsModel .uid , isouter = True )
8698 .where (UserModel .id != 1 )
8799 .distinct ()
88- .order_by (sort_clause )
89100 )
90-
101+ # 根据是否有 keyword 决定排序方式
102+ if keyword :
103+ origin_stmt = origin_stmt .order_by (similarity_score , sort_clause )
104+ else :
105+ origin_stmt = origin_stmt .order_by (sort_clause )
106+
91107 if oidlist :
92108 origin_stmt = origin_stmt .where (UserWsModel .oid .in_ (oidlist ))
93109 if origins :
94110 origin_stmt = origin_stmt .where (UserModel .origin .in_ (origins ))
95111 if status is not None :
96- origin_stmt = origin_stmt .where (UserModel .status == status )
112+ origin_stmt = origin_stmt .where (UserModel .status == status )
97113 if keyword :
98114 keyword_pattern = f"%{ keyword } %"
99115 origin_stmt = origin_stmt .where (
@@ -114,8 +130,18 @@ async def pager(
114130 select (UserModel , UserWsModel .oid .label ('ws_oid' ))
115131 .join (UserWsModel , UserModel .id == UserWsModel .uid , isouter = True )
116132 .where (UserModel .id .in_ (uid_list ))
117- .order_by (sort_clause )
118133 )
134+ # 第二次查询也需要应用相同的相似度排序
135+ if keyword :
136+ similarity_score = case (
137+ (UserModel .account == keyword , 0 ),
138+ (UserModel .account .startswith (keyword ), 1 ),
139+ (UserModel .account .contains (keyword ), 2 ),
140+ else_ = 3
141+ )
142+ stmt = stmt .order_by (similarity_score , sort_clause )
143+ else :
144+ stmt = stmt .order_by (sort_clause )
119145 user_workspaces = session .exec (stmt ).all ()
120146 merged = defaultdict (list )
121147 extra_attrs = {}
0 commit comments