refactor(model-permissions): 简化 allowed_models 为纯列表格式

移除按 API 格式区分的字典模式(Dict[str, List[str]]),统一使用简单列表格式。

- 删除 normalize_allowed_models 的 api_format 参数
- 删除 check_model_allowed 的 api_format 参数
- 简化 merge_allowed_models 为直接列表交集
- 移除前端的字典模式兼容代码和警告 UI
- 删除 is_format_mode、convert_to_format_mode 等辅助函数
This commit is contained in:
fawney19
2026-01-14 17:57:21 +08:00
parent 133f3108f0
commit b272109055
9 changed files with 62 additions and 345 deletions

View File

@@ -1,10 +1,7 @@
"""
模型权限工具
支持两种 allowed_models 格式
1. 简单模式(列表): ["claude-sonnet-4", "gpt-4o"]
2. 按格式模式(字典): {"OPENAI": ["gpt-4o"], "CLAUDE": ["claude-sonnet-4"]}
allowed_models 格式: ["claude-sonnet-4", "gpt-4o"]
使用 None/null 表示不限制(允许所有模型)
支持模型别名匹配:
@@ -16,7 +13,7 @@
import re
from functools import lru_cache
from typing import Dict, List, Optional, Tuple, Union
from typing import List, Optional, Tuple
import regex
@@ -29,19 +26,15 @@ MAX_MODEL_NAME_LENGTH = 200 # 与 MAX_ALIAS_LENGTH 保持一致
REGEX_MATCH_TIMEOUT_MS = 100 # 正则匹配超时(毫秒)
# 类型别名
AllowedModels = Optional[Union[List[str], Dict[str, List[str]]]]
AllowedModels = Optional[List[str]]
def normalize_allowed_models(
allowed_models: AllowedModels,
api_format: Optional[str] = None,
) -> Optional[set[str]]:
def normalize_allowed_models(allowed_models: AllowedModels) -> Optional[set[str]]:
"""
将 allowed_models 规范化为模型名称集合
Args:
allowed_models: 允许的模型配置(列表或字典
api_format: 当前请求的 API 格式(用于字典模式)
allowed_models: 允许的模型配置(列表)
Returns:
- None: 不限制(允许所有模型)
@@ -50,41 +43,12 @@ def normalize_allowed_models(
if allowed_models is None:
return None
# 简单模式:直接是列表
if isinstance(allowed_models, list):
return set(allowed_models)
# 按格式模式:字典
if isinstance(allowed_models, dict):
if api_format is None:
# 没有指定格式,合并所有格式的模型
all_models: set[str] = set()
for models in allowed_models.values():
if isinstance(models, list):
all_models.update(models)
return all_models if all_models else None
# 查找指定格式的模型列表
api_format_upper = api_format.upper()
models = allowed_models.get(api_format_upper)
if models is None:
# 该格式未配置,检查是否有通配符 "*"
models = allowed_models.get("*")
if models is None:
# 字典模式下未配置的格式 = 不限制该格式
return None
return set(models) if isinstance(models, list) else None
# 未知类型,视为不限制
return None
return set(allowed_models)
def check_model_allowed(
model_name: str,
allowed_models: AllowedModels,
api_format: Optional[str] = None,
resolved_model_name: Optional[str] = None,
) -> bool:
"""
@@ -93,14 +57,13 @@ def check_model_allowed(
Args:
model_name: 请求的模型名称
allowed_models: 允许的模型配置
api_format: 当前请求的 API 格式
resolved_model_name: 解析后的 GlobalModel.name可选
Returns:
True: 允许使用该模型
False: 不允许使用该模型
"""
allowed_set = normalize_allowed_models(allowed_models, api_format)
allowed_set = normalize_allowed_models(allowed_models)
if allowed_set is None:
# 不限制
@@ -130,8 +93,6 @@ def merge_allowed_models(
规则:
- 如果任一为 None返回另一个
- 如果都有值,取交集
- 如果都是列表,取列表交集
- 如果有字典,按 API 格式分别取交集(保持字典语义,不丢失格式区分信息)
Args:
allowed_models_1: 第一个配置
@@ -145,57 +106,8 @@ def merge_allowed_models(
if allowed_models_2 is None:
return allowed_models_1
# 两个都是简单列表:直接取交集(返回确定性顺序)
if isinstance(allowed_models_1, list) and isinstance(allowed_models_2, list):
intersection = set(allowed_models_1) & set(allowed_models_2)
return sorted(intersection) if intersection else []
# 任一为字典模式:按 API 格式分别取交集,避免把 dict 合并成 list 导致权限过宽
from src.core.enums import APIFormat
def merge_sets(a: Optional[set[str]], b: Optional[set[str]]) -> Optional[set[str]]:
# None 表示不限制:交集规则下等价于“只受另一方限制”
if a is None:
return b
if b is None:
return a
return a & b
known_formats = [fmt.value for fmt in APIFormat]
per_format: Dict[str, Optional[set[str]]] = {}
for fmt in known_formats:
s1 = normalize_allowed_models(allowed_models_1, api_format=fmt)
s2 = normalize_allowed_models(allowed_models_2, api_format=fmt)
per_format[fmt] = merge_sets(s1, s2)
# 计算默认(未知格式)的交集,用 "*" 作为默认值以覆盖未枚举的格式
default_s1 = normalize_allowed_models(allowed_models_1, api_format="__DEFAULT__")
default_s2 = normalize_allowed_models(allowed_models_2, api_format="__DEFAULT__")
default_set = merge_sets(default_s1, default_s2)
# 如果 default_set 非 None 且不存在“某些格式不限制”的情况,可用 "*" 作为默认规则并按需覆盖
can_use_wildcard = default_set is not None and all(v is not None for v in per_format.values())
merged_dict: Dict[str, List[str]] = {}
if can_use_wildcard and default_set is not None:
merged_dict["*"] = sorted(default_set)
for fmt, s in per_format.items():
# can_use_wildcard 保证 s 非 None
if s is not None and s != default_set:
merged_dict[fmt] = sorted(s)
else:
for fmt, s in per_format.items():
if s is None:
continue
merged_dict[fmt] = sorted(s)
if not merged_dict:
# 全部不限制
return None
return merged_dict
intersection = set(allowed_models_1) & set(allowed_models_2)
return sorted(intersection) if intersection else []
def get_allowed_models_preview(
@@ -215,19 +127,10 @@ def get_allowed_models_preview(
if allowed_models is None:
return "(不限制)"
all_models: set[str] = set()
if isinstance(allowed_models, list):
all_models = set(allowed_models)
elif isinstance(allowed_models, dict):
for models in allowed_models.values():
if isinstance(models, list):
all_models.update(models)
if not all_models:
if not allowed_models:
return "(无)"
sorted_models = sorted(all_models)
sorted_models = sorted(allowed_models)
preview = ", ".join(sorted_models[:max_items])
if len(sorted_models) > max_items:
preview += f", ...共{len(sorted_models)}"
@@ -235,103 +138,20 @@ def get_allowed_models_preview(
return preview
def is_format_mode(allowed_models: AllowedModels) -> bool:
def parse_allowed_models_to_list(allowed_models: AllowedModels) -> List[str]:
"""
判断 allowed_models 是否为按格式模式
解析 allowed_models 为列表
Args:
allowed_models: 允许的模型配置
Returns:
True: 按格式模式(字典)
False: 简单模式(列表或 None
"""
return isinstance(allowed_models, dict)
def convert_to_format_mode(
allowed_models: AllowedModels,
api_formats: Optional[List[str]] = None,
) -> Dict[str, List[str]]:
"""
将 allowed_models 转换为按格式模式
Args:
allowed_models: 原始配置
api_formats: 要应用的 API 格式列表
Returns:
按格式模式的配置
"""
if allowed_models is None:
return {}
if isinstance(allowed_models, dict):
return allowed_models
# 简单列表模式 -> 按格式模式
if isinstance(allowed_models, list):
if not api_formats:
return {"*": allowed_models}
return {fmt.upper(): list(allowed_models) for fmt in api_formats}
return {}
def convert_to_simple_mode(allowed_models: AllowedModels) -> Optional[List[str]]:
"""
将 allowed_models 转换为简单列表模式
Args:
allowed_models: 原始配置
Returns:
简单列表或 None
"""
if allowed_models is None:
return None
if isinstance(allowed_models, list):
return allowed_models
if isinstance(allowed_models, dict):
all_models: set[str] = set()
for models in allowed_models.values():
if isinstance(models, list):
all_models.update(models)
return sorted(all_models) if all_models else None
return None
def parse_allowed_models_to_list(allowed_models: AllowedModels) -> List[str]:
"""
解析 allowed_models支持 list 和 dict 格式)为统一的列表
与 convert_to_simple_mode 的区别:
- 本函数返回空列表而非 None用于 UI 展示)
- convert_to_simple_mode 返回 None 表示不限制
Args:
allowed_models: 允许的模型配置(列表或字典)
Returns:
模型名称列表(可能为空)
"""
if allowed_models is None:
return []
if isinstance(allowed_models, list):
return allowed_models
if isinstance(allowed_models, dict):
all_models: set[str] = set()
for models in allowed_models.values():
if isinstance(models, list):
all_models.update(models)
return sorted(all_models)
return []
return list(allowed_models)
def validate_alias_pattern(pattern: str) -> Tuple[bool, Optional[str]]:
@@ -530,7 +350,6 @@ def match_model_with_pattern(pattern: str, model_name: str) -> bool:
def check_model_allowed_with_aliases(
model_name: str,
allowed_models: AllowedModels,
api_format: Optional[str] = None,
resolved_model_name: Optional[str] = None,
model_aliases: Optional[List[str]] = None,
candidate_models: Optional[set[str]] = None,
@@ -552,7 +371,6 @@ def check_model_allowed_with_aliases(
Args:
model_name: 请求的模型名称
allowed_models: 允许的模型配置(来自 Provider Key
api_format: 当前请求的 API 格式
resolved_model_name: 解析后的 GlobalModel.name
model_aliases: GlobalModel 的别名列表(来自 config.model_aliases
candidate_models: 可选的候选模型集合(用于限制别名匹配只能落到这些模型名上)
@@ -563,7 +381,7 @@ def check_model_allowed_with_aliases(
- matched_model_name: 通过别名匹配到的模型名(仅别名匹配时有值,精确匹配时为 None
"""
# 先尝试精确匹配(使用原有逻辑)
if check_model_allowed(model_name, allowed_models, api_format, resolved_model_name):
if check_model_allowed(model_name, allowed_models, resolved_model_name):
return True, None
# 如果精确匹配失败且有别名配置,尝试别名匹配
@@ -571,7 +389,7 @@ def check_model_allowed_with_aliases(
return False, None
# 获取 allowed_models 的集合
allowed_set = normalize_allowed_models(allowed_models, api_format)
allowed_set = normalize_allowed_models(allowed_models)
if allowed_set is None:
# 不限制,已在 check_model_allowed 中返回 True
return True, None