mirror of
https://github.com/fawney19/Aether.git
synced 2026-09-01 17:00:21 +08:00
fix(model-resolution): 修复模型解析优先级和别名匹配范围
- 调整 ModelCacheService 解析优先级:优先直接匹配 GlobalModel.name, 避免被 provider_model_name 误导到错误的 GlobalModel - 为 check_model_allowed_with_aliases 新增 candidate_models 参数, 限制别名匹配只能落到 Provider 实际配置的模型名上 - 对 allowed_set 排序以确保别名匹配结果的确定性 - 调整 HTTP_READ_TIMEOUT 默认值从 300s 改为 60s - 优化 ModelMappingTab 组件布局,将正则表达式提取到 Key 级别展示 - 补充 .env.example 超时配置文档
This commit is contained in:
41
.env.example
41
.env.example
@@ -42,3 +42,44 @@ ADMIN_PASSWORD=admin123456
|
|||||||
# 示例: http://localhost:3000,https://example.com
|
# 示例: http://localhost:3000,https://example.com
|
||||||
# 默认: * (允许所有源)
|
# 默认: * (允许所有源)
|
||||||
# CORS_ORIGINS=*
|
# CORS_ORIGINS=*
|
||||||
|
|
||||||
|
# ==================== 超时配置 ====================
|
||||||
|
# 以下配置控制各种超时行为,影响故障转移和请求处理
|
||||||
|
|
||||||
|
# --- HTTP 连接层超时(httpx 底层) ---
|
||||||
|
# 这些是网络层的超时,控制 TCP 连接和数据传输
|
||||||
|
|
||||||
|
# TCP 连接建立超时(默认 10 秒)
|
||||||
|
# 无法在此时间内建立 TCP 连接则触发故障转移
|
||||||
|
# HTTP_CONNECT_TIMEOUT=10.0
|
||||||
|
|
||||||
|
# 读取数据超时(默认 60 秒)
|
||||||
|
# 两次数据包之间的最大等待时间,超时触发故障转移
|
||||||
|
# 注意:如果上游持续发送数据(哪怕很慢),此计时器会不断重置
|
||||||
|
# HTTP_READ_TIMEOUT=60.0
|
||||||
|
|
||||||
|
# 发送数据超时(默认 60 秒)
|
||||||
|
# 发送请求体到上游的超时时间
|
||||||
|
# HTTP_WRITE_TIMEOUT=60.0
|
||||||
|
|
||||||
|
# 连接池获取超时(默认 10 秒)
|
||||||
|
# 从连接池获取可用连接的超时时间
|
||||||
|
# HTTP_POOL_TIMEOUT=10.0
|
||||||
|
|
||||||
|
# --- 业务层超时 ---
|
||||||
|
# 这些是应用层的超时,控制请求处理流程
|
||||||
|
|
||||||
|
# 流式响应首字节超时(默认 30 秒,范围 10-120 秒)
|
||||||
|
# 从发起请求到收到第一个字节的最大等待时间
|
||||||
|
# 仅对流式请求生效,超时触发故障转移
|
||||||
|
# STREAM_FIRST_BYTE_TIMEOUT=30.0
|
||||||
|
|
||||||
|
# 请求体读取超时(默认 60 秒)
|
||||||
|
# 等待客户端发送完整请求体的超时时间
|
||||||
|
# 防止客户端发送不完整请求导致连接卡死
|
||||||
|
# REQUEST_BODY_TIMEOUT=60.0
|
||||||
|
|
||||||
|
# --- Provider 级别超时 ---
|
||||||
|
# 在管理面板中为每个 Provider 单独配置的 timeout 字段
|
||||||
|
# 控制"建立连接 + 获取首字节"的总时间(默认 300 秒)
|
||||||
|
# 建议根据 Provider 响应速度设置为 30-120 秒
|
||||||
|
|||||||
@@ -47,7 +47,7 @@
|
|||||||
class="w-4 h-4 text-muted-foreground shrink-0 transition-transform self-start mt-0.5"
|
class="w-4 h-4 text-muted-foreground shrink-0 transition-transform self-start mt-0.5"
|
||||||
:class="{ 'rotate-90': expandedItems.has(index) }"
|
:class="{ 'rotate-90': expandedItems.has(index) }"
|
||||||
/>
|
/>
|
||||||
<!-- 精确映射:两行显示 -->
|
<!-- 精确映射 -->
|
||||||
<template v-if="item.type === 'exact'">
|
<template v-if="item.type === 'exact'">
|
||||||
<div class="flex flex-col min-w-0">
|
<div class="flex flex-col min-w-0">
|
||||||
<span class="font-semibold text-sm truncate">
|
<span class="font-semibold text-sm truncate">
|
||||||
@@ -72,7 +72,7 @@
|
|||||||
| {{ item.mappings.length }} 个映射
|
| {{ item.mappings.length }} 个映射
|
||||||
</span>
|
</span>
|
||||||
</template>
|
</template>
|
||||||
<!-- 正则映射:两行显示 -->
|
<!-- 正则映射 -->
|
||||||
<template v-else>
|
<template v-else>
|
||||||
<div class="flex flex-col min-w-0">
|
<div class="flex flex-col min-w-0">
|
||||||
<span class="font-semibold text-sm truncate">
|
<span class="font-semibold text-sm truncate">
|
||||||
@@ -96,7 +96,7 @@
|
|||||||
<span class="text-xs text-muted-foreground shrink-0">
|
<span class="text-xs text-muted-foreground shrink-0">
|
||||||
| {{ item.mappings.length }} 个映射
|
| {{ item.mappings.length }} 个映射
|
||||||
</span>
|
</span>
|
||||||
<!-- 正则映射显示匹配的 Key 数量 -->
|
<!-- Key 数量 -->
|
||||||
<span
|
<span
|
||||||
v-if="item.matchedKeys && item.matchedKeys.length > 0"
|
v-if="item.matchedKeys && item.matchedKeys.length > 0"
|
||||||
class="text-xs text-muted-foreground shrink-0"
|
class="text-xs text-muted-foreground shrink-0"
|
||||||
@@ -180,31 +180,41 @@
|
|||||||
:key="keyItem.keyId"
|
:key="keyItem.keyId"
|
||||||
class="bg-background rounded-md border p-3"
|
class="bg-background rounded-md border p-3"
|
||||||
>
|
>
|
||||||
<!-- Key 信息 -->
|
<!-- Key 信息和正则表达式(两列布局) -->
|
||||||
<div class="flex items-center gap-2 text-sm mb-2">
|
<div class="flex items-center gap-3">
|
||||||
<Key class="w-3.5 h-3.5 text-muted-foreground shrink-0" />
|
<!-- 第一列:Key 名称 + sk -->
|
||||||
<span class="font-medium truncate">{{ keyItem.keyName || '未命名密钥' }}</span>
|
<div class="flex flex-col shrink-0">
|
||||||
<span class="text-xs text-muted-foreground font-mono ml-auto shrink-0">
|
<span class="font-medium text-sm">{{ keyItem.keyName || '未命名密钥' }}</span>
|
||||||
{{ keyItem.maskedKey }}
|
<span class="text-xs text-muted-foreground font-mono">
|
||||||
</span>
|
{{ keyItem.maskedKey }}
|
||||||
|
</span>
|
||||||
|
</div>
|
||||||
|
<!-- 分隔线 -->
|
||||||
|
<div
|
||||||
|
v-if="getKeyPatterns(keyItem).length > 0"
|
||||||
|
class="w-px h-8 bg-border shrink-0"
|
||||||
|
/>
|
||||||
|
<!-- 第二列:正则表达式(限制2行) -->
|
||||||
|
<div
|
||||||
|
v-if="getKeyPatterns(keyItem).length > 0"
|
||||||
|
class="flex-1 min-w-0"
|
||||||
|
>
|
||||||
|
<div
|
||||||
|
class="font-mono text-[11px] text-muted-foreground line-clamp-2"
|
||||||
|
:title="getKeyPatterns(keyItem).join(', ')"
|
||||||
|
>
|
||||||
|
{{ getKeyPatterns(keyItem).join(', ') }}
|
||||||
|
</div>
|
||||||
|
</div>
|
||||||
</div>
|
</div>
|
||||||
<!-- 匹配的模型列表 -->
|
<!-- 匹配的模型列表 -->
|
||||||
<div class="space-y-1">
|
<div class="mt-2 space-y-1">
|
||||||
<div
|
<div
|
||||||
v-for="match in keyItem.matches"
|
v-for="match in keyItem.matches"
|
||||||
:key="match.name"
|
:key="match.name"
|
||||||
class="flex items-center justify-between gap-2 py-1"
|
class="flex items-center justify-between gap-2 py-1"
|
||||||
>
|
>
|
||||||
<div class="flex items-center gap-2 flex-1 min-w-0">
|
<span class="font-mono text-sm truncate">{{ match.name }}</span>
|
||||||
<span class="font-mono text-sm truncate">{{ match.name }}</span>
|
|
||||||
<span
|
|
||||||
v-if="match.pattern"
|
|
||||||
class="text-xs text-muted-foreground truncate"
|
|
||||||
:title="match.pattern"
|
|
||||||
>
|
|
||||||
{{ match.pattern }}
|
|
||||||
</span>
|
|
||||||
</div>
|
|
||||||
<Button
|
<Button
|
||||||
variant="ghost"
|
variant="ghost"
|
||||||
size="icon"
|
size="icon"
|
||||||
@@ -270,7 +280,7 @@
|
|||||||
|
|
||||||
<script setup lang="ts">
|
<script setup lang="ts">
|
||||||
import { ref, computed, watch } from 'vue'
|
import { ref, computed, watch } from 'vue'
|
||||||
import { Tag, Plus, Edit, Trash2, ChevronRight, Loader2, Play, Key } from 'lucide-vue-next'
|
import { Tag, Plus, Edit, Trash2, ChevronRight, Loader2, Play } from 'lucide-vue-next'
|
||||||
import { Card, Button, Badge } from '@/components/ui'
|
import { Card, Button, Badge } from '@/components/ui'
|
||||||
import AlertDialog from '@/components/common/AlertDialog.vue'
|
import AlertDialog from '@/components/common/AlertDialog.vue'
|
||||||
import ModelMappingDialog, { type AliasGroup } from '../ModelMappingDialog.vue'
|
import ModelMappingDialog, { type AliasGroup } from '../ModelMappingDialog.vue'
|
||||||
@@ -491,6 +501,17 @@ function toggleExpand(index: number) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
// 获取单个 Key 的去重正则模式列表
|
||||||
|
function getKeyPatterns(keyItem: MatchedKeyInfo): string[] {
|
||||||
|
const patterns = new Set<string>()
|
||||||
|
for (const match of keyItem.matches) {
|
||||||
|
if (match.pattern) {
|
||||||
|
patterns.add(match.pattern)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
return Array.from(patterns)
|
||||||
|
}
|
||||||
|
|
||||||
// 打开添加对话框
|
// 打开添加对话框
|
||||||
function openAddDialog() {
|
function openAddDialog() {
|
||||||
editingGroup.value = null
|
editingGroup.value = null
|
||||||
|
|||||||
@@ -141,7 +141,7 @@ class Config:
|
|||||||
|
|
||||||
# HTTP 请求超时配置(秒)
|
# HTTP 请求超时配置(秒)
|
||||||
self.http_connect_timeout = float(os.getenv("HTTP_CONNECT_TIMEOUT", "10.0"))
|
self.http_connect_timeout = float(os.getenv("HTTP_CONNECT_TIMEOUT", "10.0"))
|
||||||
self.http_read_timeout = float(os.getenv("HTTP_READ_TIMEOUT", "300.0"))
|
self.http_read_timeout = float(os.getenv("HTTP_READ_TIMEOUT", "60.0"))
|
||||||
self.http_write_timeout = float(os.getenv("HTTP_WRITE_TIMEOUT", "60.0"))
|
self.http_write_timeout = float(os.getenv("HTTP_WRITE_TIMEOUT", "60.0"))
|
||||||
self.http_pool_timeout = float(os.getenv("HTTP_POOL_TIMEOUT", "10.0"))
|
self.http_pool_timeout = float(os.getenv("HTTP_POOL_TIMEOUT", "10.0"))
|
||||||
|
|
||||||
|
|||||||
@@ -16,7 +16,7 @@
|
|||||||
|
|
||||||
import re
|
import re
|
||||||
from functools import lru_cache
|
from functools import lru_cache
|
||||||
from typing import Dict, List, Optional, Set, Tuple, Union
|
from typing import Dict, List, Optional, Tuple, Union
|
||||||
|
|
||||||
import regex
|
import regex
|
||||||
|
|
||||||
@@ -35,7 +35,7 @@ AllowedModels = Optional[Union[List[str], Dict[str, List[str]]]]
|
|||||||
def normalize_allowed_models(
|
def normalize_allowed_models(
|
||||||
allowed_models: AllowedModels,
|
allowed_models: AllowedModels,
|
||||||
api_format: Optional[str] = None,
|
api_format: Optional[str] = None,
|
||||||
) -> Optional[Set[str]]:
|
) -> Optional[set[str]]:
|
||||||
"""
|
"""
|
||||||
将 allowed_models 规范化为模型名称集合
|
将 allowed_models 规范化为模型名称集合
|
||||||
|
|
||||||
@@ -45,7 +45,7 @@ def normalize_allowed_models(
|
|||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
- None: 不限制(允许所有模型)
|
- None: 不限制(允许所有模型)
|
||||||
- Set[str]: 允许的模型名称集合(可能为空集,表示拒绝所有)
|
- set[str]: 允许的模型名称集合(可能为空集,表示拒绝所有)
|
||||||
"""
|
"""
|
||||||
if allowed_models is None:
|
if allowed_models is None:
|
||||||
return None
|
return None
|
||||||
@@ -58,7 +58,7 @@ def normalize_allowed_models(
|
|||||||
if isinstance(allowed_models, dict):
|
if isinstance(allowed_models, dict):
|
||||||
if api_format is None:
|
if api_format is None:
|
||||||
# 没有指定格式,合并所有格式的模型
|
# 没有指定格式,合并所有格式的模型
|
||||||
all_models: Set[str] = set()
|
all_models: set[str] = set()
|
||||||
for models in allowed_models.values():
|
for models in allowed_models.values():
|
||||||
if isinstance(models, list):
|
if isinstance(models, list):
|
||||||
all_models.update(models)
|
all_models.update(models)
|
||||||
@@ -153,7 +153,7 @@ def merge_allowed_models(
|
|||||||
# 任一为字典模式:按 API 格式分别取交集,避免把 dict 合并成 list 导致权限过宽
|
# 任一为字典模式:按 API 格式分别取交集,避免把 dict 合并成 list 导致权限过宽
|
||||||
from src.core.enums import APIFormat
|
from src.core.enums import APIFormat
|
||||||
|
|
||||||
def merge_sets(a: Optional[Set[str]], b: Optional[Set[str]]) -> Optional[Set[str]]:
|
def merge_sets(a: Optional[set[str]], b: Optional[set[str]]) -> Optional[set[str]]:
|
||||||
# None 表示不限制:交集规则下等价于“只受另一方限制”
|
# None 表示不限制:交集规则下等价于“只受另一方限制”
|
||||||
if a is None:
|
if a is None:
|
||||||
return b
|
return b
|
||||||
@@ -163,7 +163,7 @@ def merge_allowed_models(
|
|||||||
|
|
||||||
known_formats = [fmt.value for fmt in APIFormat]
|
known_formats = [fmt.value for fmt in APIFormat]
|
||||||
|
|
||||||
per_format: Dict[str, Optional[Set[str]]] = {}
|
per_format: Dict[str, Optional[set[str]]] = {}
|
||||||
for fmt in known_formats:
|
for fmt in known_formats:
|
||||||
s1 = normalize_allowed_models(allowed_models_1, api_format=fmt)
|
s1 = normalize_allowed_models(allowed_models_1, api_format=fmt)
|
||||||
s2 = normalize_allowed_models(allowed_models_2, api_format=fmt)
|
s2 = normalize_allowed_models(allowed_models_2, api_format=fmt)
|
||||||
@@ -215,7 +215,7 @@ def get_allowed_models_preview(
|
|||||||
if allowed_models is None:
|
if allowed_models is None:
|
||||||
return "(不限制)"
|
return "(不限制)"
|
||||||
|
|
||||||
all_models: Set[str] = set()
|
all_models: set[str] = set()
|
||||||
|
|
||||||
if isinstance(allowed_models, list):
|
if isinstance(allowed_models, list):
|
||||||
all_models = set(allowed_models)
|
all_models = set(allowed_models)
|
||||||
@@ -295,7 +295,7 @@ def convert_to_simple_mode(allowed_models: AllowedModels) -> Optional[List[str]]
|
|||||||
return allowed_models
|
return allowed_models
|
||||||
|
|
||||||
if isinstance(allowed_models, dict):
|
if isinstance(allowed_models, dict):
|
||||||
all_models: Set[str] = set()
|
all_models: set[str] = set()
|
||||||
for models in allowed_models.values():
|
for models in allowed_models.values():
|
||||||
if isinstance(models, list):
|
if isinstance(models, list):
|
||||||
all_models.update(models)
|
all_models.update(models)
|
||||||
@@ -325,7 +325,7 @@ def parse_allowed_models_to_list(allowed_models: AllowedModels) -> List[str]:
|
|||||||
return allowed_models
|
return allowed_models
|
||||||
|
|
||||||
if isinstance(allowed_models, dict):
|
if isinstance(allowed_models, dict):
|
||||||
all_models: Set[str] = set()
|
all_models: set[str] = set()
|
||||||
for models in allowed_models.values():
|
for models in allowed_models.values():
|
||||||
if isinstance(models, list):
|
if isinstance(models, list):
|
||||||
all_models.update(models)
|
all_models.update(models)
|
||||||
@@ -533,6 +533,7 @@ def check_model_allowed_with_aliases(
|
|||||||
api_format: Optional[str] = None,
|
api_format: Optional[str] = None,
|
||||||
resolved_model_name: Optional[str] = None,
|
resolved_model_name: Optional[str] = None,
|
||||||
model_aliases: Optional[List[str]] = None,
|
model_aliases: Optional[List[str]] = None,
|
||||||
|
candidate_models: Optional[set[str]] = None,
|
||||||
) -> tuple[bool, Optional[str]]:
|
) -> tuple[bool, Optional[str]]:
|
||||||
"""
|
"""
|
||||||
检查模型是否被允许(支持别名通配符匹配)
|
检查模型是否被允许(支持别名通配符匹配)
|
||||||
@@ -554,6 +555,7 @@ def check_model_allowed_with_aliases(
|
|||||||
api_format: 当前请求的 API 格式
|
api_format: 当前请求的 API 格式
|
||||||
resolved_model_name: 解析后的 GlobalModel.name
|
resolved_model_name: 解析后的 GlobalModel.name
|
||||||
model_aliases: GlobalModel 的别名列表(来自 config.model_aliases)
|
model_aliases: GlobalModel 的别名列表(来自 config.model_aliases)
|
||||||
|
candidate_models: 可选的候选模型集合(用于限制别名匹配只能落到这些模型名上)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(is_allowed, matched_model_name):
|
(is_allowed, matched_model_name):
|
||||||
@@ -578,9 +580,16 @@ def check_model_allowed_with_aliases(
|
|||||||
# 空集合 = 拒绝所有
|
# 空集合 = 拒绝所有
|
||||||
return False, None
|
return False, None
|
||||||
|
|
||||||
|
# 如果提供了候选集合,只允许在候选集合中进行别名匹配
|
||||||
|
if candidate_models is not None:
|
||||||
|
allowed_set = allowed_set & candidate_models
|
||||||
|
if len(allowed_set) == 0:
|
||||||
|
return False, None
|
||||||
|
|
||||||
# 遍历 allowed_models 中的每个模型名,检查是否有别名能匹配
|
# 遍历 allowed_models 中的每个模型名,检查是否有别名能匹配
|
||||||
# 注意:返回第一个匹配的模型名,匹配顺序由 allowed_set 迭代顺序和 model_aliases 数组顺序决定
|
# 注意:为了避免 set 迭代顺序带来的非确定性,这里对 allowed_set 做排序
|
||||||
for allowed_model in allowed_set:
|
# 返回第一个匹配的模型名,匹配顺序由 allowed_models 排序结果和 model_aliases 数组顺序共同决定
|
||||||
|
for allowed_model in sorted(allowed_set):
|
||||||
for alias_pattern in model_aliases:
|
for alias_pattern in model_aliases:
|
||||||
if match_model_with_pattern(alias_pattern, allowed_model):
|
if match_model_with_pattern(alias_pattern, allowed_model):
|
||||||
# 返回匹配到的模型名,用于实际请求
|
# 返回匹配到的模型名,用于实际请求
|
||||||
|
|||||||
62
src/services/cache/aware_scheduler.py
vendored
62
src/services/cache/aware_scheduler.py
vendored
@@ -746,9 +746,10 @@ class CacheAwareScheduler:
|
|||||||
db: Session,
|
db: Session,
|
||||||
provider: Provider,
|
provider: Provider,
|
||||||
model_name: str,
|
model_name: str,
|
||||||
|
api_format: Optional[str] = None,
|
||||||
is_stream: bool = False,
|
is_stream: bool = False,
|
||||||
capability_requirements: Optional[Dict[str, bool]] = None,
|
capability_requirements: Optional[Dict[str, bool]] = None,
|
||||||
) -> Tuple[bool, Optional[str], Optional[List[str]]]:
|
) -> Tuple[bool, Optional[str], Optional[List[str]], Optional[set[str]]]:
|
||||||
"""
|
"""
|
||||||
检查 Provider 是否支持指定模型(可选检查流式支持和能力需求)
|
检查 Provider 是否支持指定模型(可选检查流式支持和能力需求)
|
||||||
|
|
||||||
@@ -768,20 +769,24 @@ class CacheAwareScheduler:
|
|||||||
capability_requirements: 能力需求(可选),用于检查模型是否支持所需能力
|
capability_requirements: 能力需求(可选),用于检查模型是否支持所需能力
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(is_supported, skip_reason, supported_capabilities) - 是否支持、跳过原因、模型支持的能力列表
|
(is_supported, skip_reason, supported_capabilities, provider_model_names)
|
||||||
|
- is_supported: 是否支持
|
||||||
|
- skip_reason: 跳过原因
|
||||||
|
- supported_capabilities: 模型支持的能力列表
|
||||||
|
- provider_model_names: Provider 侧可用的模型名称集合(主名称 + 映射名称,按 api_format 过滤)
|
||||||
"""
|
"""
|
||||||
# 使用 ModelCacheService 解析模型名称(支持映射名)
|
# 使用 ModelCacheService 解析模型名称(支持映射名)
|
||||||
global_model = await ModelCacheService.resolve_global_model_by_name_or_alias(db, model_name)
|
global_model = await ModelCacheService.resolve_global_model_by_name_or_alias(db, model_name)
|
||||||
|
|
||||||
if not global_model:
|
if not global_model:
|
||||||
# 完全未找到匹配
|
# 完全未找到匹配
|
||||||
return False, "模型不存在或 Provider 未配置此模型", None
|
return False, "模型不存在或 Provider 未配置此模型", None, None
|
||||||
|
|
||||||
# 找到 GlobalModel 后,检查当前 Provider 是否支持
|
# 找到 GlobalModel 后,检查当前 Provider 是否支持
|
||||||
is_supported, skip_reason, caps = await self._check_model_support_for_global_model(
|
is_supported, skip_reason, caps, provider_model_names = await self._check_model_support_for_global_model(
|
||||||
db, provider, global_model, model_name, is_stream, capability_requirements
|
db, provider, global_model, model_name, api_format, is_stream, capability_requirements
|
||||||
)
|
)
|
||||||
return is_supported, skip_reason, caps
|
return is_supported, skip_reason, caps, provider_model_names
|
||||||
|
|
||||||
async def _check_model_support_for_global_model(
|
async def _check_model_support_for_global_model(
|
||||||
self,
|
self,
|
||||||
@@ -789,9 +794,10 @@ class CacheAwareScheduler:
|
|||||||
provider: Provider,
|
provider: Provider,
|
||||||
global_model: "GlobalModel",
|
global_model: "GlobalModel",
|
||||||
model_name: str,
|
model_name: str,
|
||||||
|
api_format: Optional[str] = None,
|
||||||
is_stream: bool = False,
|
is_stream: bool = False,
|
||||||
capability_requirements: Optional[Dict[str, bool]] = None,
|
capability_requirements: Optional[Dict[str, bool]] = None,
|
||||||
) -> Tuple[bool, Optional[str], Optional[List[str]]]:
|
) -> Tuple[bool, Optional[str], Optional[List[str]], Optional[set[str]]]:
|
||||||
"""
|
"""
|
||||||
检查 Provider 是否支持指定的 GlobalModel
|
检查 Provider 是否支持指定的 GlobalModel
|
||||||
|
|
||||||
@@ -804,7 +810,7 @@ class CacheAwareScheduler:
|
|||||||
capability_requirements: 能力需求
|
capability_requirements: 能力需求
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(is_supported, skip_reason, supported_capabilities)
|
(is_supported, skip_reason, supported_capabilities, provider_model_names)
|
||||||
"""
|
"""
|
||||||
# 确保 global_model 附加到当前 Session
|
# 确保 global_model 附加到当前 Session
|
||||||
# 注意:从缓存重建的对象是 transient 状态,不能使用 load=False
|
# 注意:从缓存重建的对象是 transient 状态,不能使用 load=False
|
||||||
@@ -833,7 +839,7 @@ class CacheAwareScheduler:
|
|||||||
if is_stream:
|
if is_stream:
|
||||||
supports_streaming = model.get_effective_supports_streaming()
|
supports_streaming = model.get_effective_supports_streaming()
|
||||||
if not supports_streaming:
|
if not supports_streaming:
|
||||||
return False, f"模型 {model_name} 在此 Provider 不支持流式", None
|
return False, f"模型 {model_name} 在此 Provider 不支持流式", None, None
|
||||||
|
|
||||||
# 检查模型是否支持所需的能力(在 Provider 级别检查,而不是 Key 级别)
|
# 检查模型是否支持所需的能力(在 Provider 级别检查,而不是 Key 级别)
|
||||||
# 只有当 model_supported_capabilities 非空时才进行检查
|
# 只有当 model_supported_capabilities 非空时才进行检查
|
||||||
@@ -845,11 +851,32 @@ class CacheAwareScheduler:
|
|||||||
False,
|
False,
|
||||||
f"模型 {model_name} 不支持能力: {cap_name}",
|
f"模型 {model_name} 不支持能力: {cap_name}",
|
||||||
list(model_supported_capabilities),
|
list(model_supported_capabilities),
|
||||||
|
None,
|
||||||
)
|
)
|
||||||
|
|
||||||
return True, None, list(model_supported_capabilities)
|
provider_model_names: set[str] = {model.provider_model_name}
|
||||||
|
raw_mappings = model.provider_model_mappings
|
||||||
|
if isinstance(raw_mappings, list):
|
||||||
|
for raw in raw_mappings:
|
||||||
|
if not isinstance(raw, dict):
|
||||||
|
continue
|
||||||
|
name = raw.get("name")
|
||||||
|
if not isinstance(name, str) or not name.strip():
|
||||||
|
continue
|
||||||
|
|
||||||
return False, "Provider 未实现此模型", None
|
mapping_api_formats = raw.get("api_formats")
|
||||||
|
if api_format and mapping_api_formats:
|
||||||
|
if (
|
||||||
|
isinstance(mapping_api_formats, list)
|
||||||
|
and api_format not in mapping_api_formats
|
||||||
|
):
|
||||||
|
continue
|
||||||
|
|
||||||
|
provider_model_names.add(name.strip())
|
||||||
|
|
||||||
|
return True, None, list(model_supported_capabilities), provider_model_names
|
||||||
|
|
||||||
|
return False, "Provider 未实现此模型", None, None
|
||||||
|
|
||||||
def _check_key_availability(
|
def _check_key_availability(
|
||||||
self,
|
self,
|
||||||
@@ -859,6 +886,7 @@ class CacheAwareScheduler:
|
|||||||
capability_requirements: Optional[Dict[str, bool]] = None,
|
capability_requirements: Optional[Dict[str, bool]] = None,
|
||||||
resolved_model_name: Optional[str] = None,
|
resolved_model_name: Optional[str] = None,
|
||||||
model_aliases: Optional[List[str]] = None,
|
model_aliases: Optional[List[str]] = None,
|
||||||
|
candidate_models: Optional[set[str]] = None,
|
||||||
) -> Tuple[bool, Optional[str], Optional[str]]:
|
) -> Tuple[bool, Optional[str], Optional[str]]:
|
||||||
"""
|
"""
|
||||||
检查 API Key 的可用性
|
检查 API Key 的可用性
|
||||||
@@ -872,6 +900,7 @@ class CacheAwareScheduler:
|
|||||||
capability_requirements: 能力需求(可选)
|
capability_requirements: 能力需求(可选)
|
||||||
resolved_model_name: 解析后的 GlobalModel.name(可选)
|
resolved_model_name: 解析后的 GlobalModel.name(可选)
|
||||||
model_aliases: GlobalModel 的别名列表(用于通配符匹配)
|
model_aliases: GlobalModel 的别名列表(用于通配符匹配)
|
||||||
|
candidate_models: Provider 侧可用的模型名称集合(用于限制别名匹配范围)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
(is_available, skip_reason, alias_matched_model)
|
(is_available, skip_reason, alias_matched_model)
|
||||||
@@ -901,6 +930,7 @@ class CacheAwareScheduler:
|
|||||||
api_format=api_format,
|
api_format=api_format,
|
||||||
resolved_model_name=resolved_model_name,
|
resolved_model_name=resolved_model_name,
|
||||||
model_aliases=model_aliases,
|
model_aliases=model_aliases,
|
||||||
|
candidate_models=candidate_models,
|
||||||
)
|
)
|
||||||
except TimeoutError:
|
except TimeoutError:
|
||||||
# 正则匹配超时(可能是 ReDoS 攻击或复杂模式)
|
# 正则匹配超时(可能是 ReDoS 攻击或复杂模式)
|
||||||
@@ -970,8 +1000,13 @@ class CacheAwareScheduler:
|
|||||||
|
|
||||||
for provider in providers:
|
for provider in providers:
|
||||||
# 检查模型支持(同时检查流式支持和模型能力需求)
|
# 检查模型支持(同时检查流式支持和模型能力需求)
|
||||||
supports_model, skip_reason, _model_caps = await self._check_model_support(
|
supports_model, skip_reason, _model_caps, provider_model_names = await self._check_model_support(
|
||||||
db, provider, model_name, is_stream, capability_requirements
|
db,
|
||||||
|
provider,
|
||||||
|
model_name,
|
||||||
|
api_format=target_format_str,
|
||||||
|
is_stream=is_stream,
|
||||||
|
capability_requirements=capability_requirements,
|
||||||
)
|
)
|
||||||
if not supports_model:
|
if not supports_model:
|
||||||
logger.debug(f"Provider {provider.name} 不支持模型 {model_name}: {skip_reason}")
|
logger.debug(f"Provider {provider.name} 不支持模型 {model_name}: {skip_reason}")
|
||||||
@@ -1022,6 +1057,7 @@ class CacheAwareScheduler:
|
|||||||
capability_requirements,
|
capability_requirements,
|
||||||
resolved_model_name=resolved_model_name,
|
resolved_model_name=resolved_model_name,
|
||||||
model_aliases=model_aliases,
|
model_aliases=model_aliases,
|
||||||
|
candidate_models=provider_model_names,
|
||||||
)
|
)
|
||||||
|
|
||||||
candidate = ProviderCandidate(
|
candidate = ProviderCandidate(
|
||||||
|
|||||||
41
src/services/cache/model_cache.py
vendored
41
src/services/cache/model_cache.py
vendored
@@ -258,8 +258,8 @@ class ModelCacheService:
|
|||||||
|
|
||||||
查找顺序:
|
查找顺序:
|
||||||
1. 检查缓存
|
1. 检查缓存
|
||||||
2. 通过 provider_model_name 匹配(查询 Model 表)
|
2. 直接匹配 GlobalModel.name
|
||||||
3. 直接匹配 GlobalModel.name(兜底)
|
3. 通过 provider_model_name 匹配(查询 Model 表)
|
||||||
|
|
||||||
注意:此方法不使用 provider_model_mappings 进行全局解析。
|
注意:此方法不使用 provider_model_mappings 进行全局解析。
|
||||||
provider_model_mappings 是 Provider 级别的映射配置,只在特定 Provider 上下文中生效,
|
provider_model_mappings 是 Provider 级别的映射配置,只在特定 Provider 上下文中生效,
|
||||||
@@ -301,7 +301,25 @@ class ModelCacheService:
|
|||||||
logger.debug(f"GlobalModel 缓存命中(映射解析): {normalized_name}")
|
logger.debug(f"GlobalModel 缓存命中(映射解析): {normalized_name}")
|
||||||
return ModelCacheService._dict_to_global_model(cached_data)
|
return ModelCacheService._dict_to_global_model(cached_data)
|
||||||
|
|
||||||
# 2. 通过 provider_model_name 匹配(不考虑 provider_model_mappings)
|
# 2. 直接通过 GlobalModel.name 匹配(优先级最高)
|
||||||
|
# 说明:如果存在同名 GlobalModel,应优先解析为 GlobalModel 本身,
|
||||||
|
# 避免被某个 Provider 的 provider_model_name 误导导致解析到错误的 GlobalModel。
|
||||||
|
global_model = (
|
||||||
|
db.query(GlobalModel)
|
||||||
|
.filter(GlobalModel.name == normalized_name, GlobalModel.is_active == True)
|
||||||
|
.first()
|
||||||
|
)
|
||||||
|
|
||||||
|
if global_model:
|
||||||
|
resolution_method = "direct_match"
|
||||||
|
global_model_dict = ModelCacheService._global_model_to_dict(global_model)
|
||||||
|
await CacheService.set(
|
||||||
|
cache_key, global_model_dict, ttl_seconds=ModelCacheService.CACHE_TTL
|
||||||
|
)
|
||||||
|
logger.debug(f"GlobalModel 已缓存(映射解析-直接匹配): {normalized_name}")
|
||||||
|
return global_model
|
||||||
|
|
||||||
|
# 3. 通过 provider_model_name 匹配(不考虑 provider_model_mappings)
|
||||||
# 重要:provider_model_mappings 是 Provider 级别的映射配置,只在特定 Provider 上下文中生效
|
# 重要:provider_model_mappings 是 Provider 级别的映射配置,只在特定 Provider 上下文中生效
|
||||||
# 全局解析不应该受到某个 Provider 映射配置的影响
|
# 全局解析不应该受到某个 Provider 映射配置的影响
|
||||||
# 例如:Provider A 把 "haiku" 映射到 "sonnet",不应该影响 Provider B 的 "haiku" 解析
|
# 例如:Provider A 把 "haiku" 映射到 "sonnet",不应该影响 Provider B 的 "haiku" 解析
|
||||||
@@ -357,23 +375,6 @@ class ModelCacheService:
|
|||||||
)
|
)
|
||||||
return result_global_model
|
return result_global_model
|
||||||
|
|
||||||
# 3. 如果通过 provider 映射没找到,最后尝试直接通过 GlobalModel.name 查找
|
|
||||||
global_model = (
|
|
||||||
db.query(GlobalModel)
|
|
||||||
.filter(GlobalModel.name == normalized_name, GlobalModel.is_active == True)
|
|
||||||
.first()
|
|
||||||
)
|
|
||||||
|
|
||||||
if global_model:
|
|
||||||
resolution_method = "direct_match"
|
|
||||||
# 缓存结果
|
|
||||||
global_model_dict = ModelCacheService._global_model_to_dict(global_model)
|
|
||||||
await CacheService.set(
|
|
||||||
cache_key, global_model_dict, ttl_seconds=ModelCacheService.CACHE_TTL
|
|
||||||
)
|
|
||||||
logger.debug(f"GlobalModel 已缓存(映射解析-直接匹配): {normalized_name}")
|
|
||||||
return global_model
|
|
||||||
|
|
||||||
# 4. 完全未找到
|
# 4. 完全未找到
|
||||||
resolution_method = "not_found"
|
resolution_method = "not_found"
|
||||||
# 未找到匹配,缓存负结果
|
# 未找到匹配,缓存负结果
|
||||||
|
|||||||
50
tests/core/test_model_permissions.py
Normal file
50
tests/core/test_model_permissions.py
Normal file
@@ -0,0 +1,50 @@
|
|||||||
|
from src.core.model_permissions import check_model_allowed_with_aliases
|
||||||
|
|
||||||
|
|
||||||
|
class TestCheckModelAllowedWithAliases:
|
||||||
|
def test_exact_match_returns_allowed_without_mapping(self) -> None:
|
||||||
|
is_allowed, matched = check_model_allowed_with_aliases(
|
||||||
|
model_name="gpt-4o",
|
||||||
|
allowed_models=["gpt-4o"],
|
||||||
|
api_format="OPENAI",
|
||||||
|
resolved_model_name="gpt-4o",
|
||||||
|
model_aliases=[r"gpt-4o-.*"],
|
||||||
|
)
|
||||||
|
assert is_allowed is True
|
||||||
|
assert matched is None
|
||||||
|
|
||||||
|
def test_alias_match_is_deterministic(self) -> None:
|
||||||
|
is_allowed, matched = check_model_allowed_with_aliases(
|
||||||
|
model_name="target",
|
||||||
|
allowed_models=["b", "a"],
|
||||||
|
api_format="OPENAI",
|
||||||
|
resolved_model_name="target",
|
||||||
|
model_aliases=[r".*"],
|
||||||
|
)
|
||||||
|
assert is_allowed is True
|
||||||
|
assert matched == "a"
|
||||||
|
|
||||||
|
def test_alias_match_respects_candidate_models(self) -> None:
|
||||||
|
is_allowed, matched = check_model_allowed_with_aliases(
|
||||||
|
model_name="target",
|
||||||
|
allowed_models=["other-1", "allowed-1"],
|
||||||
|
api_format="OPENAI",
|
||||||
|
resolved_model_name="target",
|
||||||
|
model_aliases=[r".*-1"],
|
||||||
|
candidate_models={"allowed-1"},
|
||||||
|
)
|
||||||
|
assert is_allowed is True
|
||||||
|
assert matched == "allowed-1"
|
||||||
|
|
||||||
|
def test_alias_match_candidate_models_no_intersection(self) -> None:
|
||||||
|
is_allowed, matched = check_model_allowed_with_aliases(
|
||||||
|
model_name="target",
|
||||||
|
allowed_models=["allowed-1"],
|
||||||
|
api_format="OPENAI",
|
||||||
|
resolved_model_name="target",
|
||||||
|
model_aliases=[r".*-1"],
|
||||||
|
candidate_models={"not-present"},
|
||||||
|
)
|
||||||
|
assert is_allowed is False
|
||||||
|
assert matched is None
|
||||||
|
|
||||||
69
tests/services/test_model_cache_service.py
Normal file
69
tests/services/test_model_cache_service.py
Normal file
@@ -0,0 +1,69 @@
|
|||||||
|
import pytest
|
||||||
|
|
||||||
|
from src.core.cache_service import CacheService
|
||||||
|
from src.models.database import GlobalModel, Model
|
||||||
|
from src.services.cache.model_cache import ModelCacheService
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeQuery:
|
||||||
|
def __init__(self, *, first_result=None, all_result=None, on_all=None):
|
||||||
|
self._first_result = first_result
|
||||||
|
self._all_result = all_result if all_result is not None else []
|
||||||
|
self._on_all = on_all
|
||||||
|
|
||||||
|
def join(self, *_args, **_kwargs):
|
||||||
|
return self
|
||||||
|
|
||||||
|
def filter(self, *_args, **_kwargs):
|
||||||
|
return self
|
||||||
|
|
||||||
|
def first(self):
|
||||||
|
return self._first_result
|
||||||
|
|
||||||
|
def all(self):
|
||||||
|
if self._on_all:
|
||||||
|
self._on_all()
|
||||||
|
return self._all_result
|
||||||
|
|
||||||
|
|
||||||
|
class _FakeSession:
|
||||||
|
def __init__(self, *, direct_match: GlobalModel):
|
||||||
|
self._direct_match = direct_match
|
||||||
|
|
||||||
|
def query(self, *entities):
|
||||||
|
if entities == (GlobalModel,):
|
||||||
|
return _FakeQuery(first_result=self._direct_match)
|
||||||
|
|
||||||
|
# 如果 direct match 命中,不应再走 provider_model_name 分支
|
||||||
|
if entities == (Model, GlobalModel):
|
||||||
|
raise AssertionError("provider_model_name query should not run when direct match exists")
|
||||||
|
|
||||||
|
raise AssertionError(f"Unexpected query entities: {entities}")
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_resolve_global_model_prefers_direct_match(monkeypatch) -> None:
|
||||||
|
async def _fake_get(_key: str):
|
||||||
|
return None
|
||||||
|
|
||||||
|
async def _fake_set(_key: str, _value, ttl_seconds: int = 60): # noqa: ARG001
|
||||||
|
return True
|
||||||
|
|
||||||
|
monkeypatch.setattr(CacheService, "get", staticmethod(_fake_get))
|
||||||
|
monkeypatch.setattr(CacheService, "set", staticmethod(_fake_set))
|
||||||
|
|
||||||
|
global_model = GlobalModel(
|
||||||
|
id="gm-1",
|
||||||
|
name="claude-haiku-4-5-20251001",
|
||||||
|
display_name="Claude Haiku 4.5",
|
||||||
|
supported_capabilities=[],
|
||||||
|
config={},
|
||||||
|
default_tiered_pricing=None,
|
||||||
|
default_price_per_request=None,
|
||||||
|
is_active=True,
|
||||||
|
)
|
||||||
|
db = _FakeSession(direct_match=global_model)
|
||||||
|
|
||||||
|
resolved = await ModelCacheService.resolve_global_model_by_name_or_alias(db, global_model.name)
|
||||||
|
assert resolved is global_model
|
||||||
|
|
||||||
Reference in New Issue
Block a user