sdpy-kit 0.1.0__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- kit/__init__.py +68 -0
- kit/ai/__init__.py +47 -0
- kit/ai/base_service.py +94 -0
- kit/ai/chat_service.py +110 -0
- kit/ai/embedding_service.py +133 -0
- kit/ai/exceptions.py +57 -0
- kit/ai/failover.py +244 -0
- kit/ai/image_service.py +99 -0
- kit/ai/model_builder.py +119 -0
- kit/ai/model_client.py +214 -0
- kit/ai/pool.py +411 -0
- kit/ai/profiles.py +373 -0
- kit/ai/rerank_service.py +93 -0
- kit/ai/runner.py +580 -0
- kit/ai/types.py +28 -0
- kit/common_kit.py +221 -0
- kit/config/__init__.py +44 -0
- kit/config/kit_config.py +147 -0
- kit/config/log_config.py +134 -0
- kit/db/__init__.py +70 -0
- kit/db/base_db_kit.py +406 -0
- kit/db/crud.py +361 -0
- kit/db/engine.py +230 -0
- kit/db/exceptions.py +34 -0
- kit/db/models/__init__.py +33 -0
- kit/db/models/base.py +56 -0
- kit/db/mysql_kit.py +89 -0
- kit/db/pg_kit.py +87 -0
- kit/db/session.py +85 -0
- kit/db/sqlite_kit.py +284 -0
- kit/db/transaction.py +109 -0
- kit/doc_kit.py +1502 -0
- kit/file_kit.py +429 -0
- kit/mcp_kit.py +547 -0
- kit/mineru_kit.py +918 -0
- kit/redis_kit.py +392 -0
- kit/zvec_kit.py +671 -0
- sdpy_kit-0.1.0.dist-info/METADATA +128 -0
- sdpy_kit-0.1.0.dist-info/RECORD +40 -0
- sdpy_kit-0.1.0.dist-info/WHEEL +4 -0
kit/ai/profiles.py
ADDED
|
@@ -0,0 +1,373 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Copyright (c) 2021-2026 Clark Chang. All Rights Reserved.
|
|
3
|
+
|
|
4
|
+
Licensed under the Apache License, Version 2.0 (the "License");
|
|
5
|
+
you may not use this file except in compliance with the License.
|
|
6
|
+
You may obtain a copy of the License at
|
|
7
|
+
|
|
8
|
+
http://www.apache.org/licenses/LICENSE-2.0
|
|
9
|
+
|
|
10
|
+
Unless required by applicable law or agreed to in writing, software
|
|
11
|
+
distributed under the License is distributed on an "AS IS" BASIS,
|
|
12
|
+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
13
|
+
See the License for the specific language governing permissions and
|
|
14
|
+
limitations under the License.
|
|
15
|
+
|
|
16
|
+
Project: sdpy
|
|
17
|
+
Author: Clark Chang
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
# date: 2026-08-21
|
|
21
|
+
"""模型档案(ModelProfile)与注册表(ModelRegistry)
|
|
22
|
+
|
|
23
|
+
从 ``config/ai_models.json`` 读取四类型(chat / embedding / rerank / image)
|
|
24
|
+
多厂商模型档案,用 pydantic 模型校验结构,并在 ``get()`` 时解析档案引用
|
|
25
|
+
的环境变量密钥。
|
|
26
|
+
|
|
27
|
+
密钥来源优先级:真实环境变量 > 项目根 ``.env``(``load()`` 时经
|
|
28
|
+
``load_dotenv`` 注入,不覆盖已有环境变量)> 缺失报 ``AIConfigError``。
|
|
29
|
+
档案数量与密钥变量名完全由 JSON 决定,无需改 ``sys_config``。
|
|
30
|
+
"""
|
|
31
|
+
|
|
32
|
+
import enum
|
|
33
|
+
import os
|
|
34
|
+
from dataclasses import dataclass, replace
|
|
35
|
+
from functools import lru_cache
|
|
36
|
+
from pathlib import Path
|
|
37
|
+
from typing import Final
|
|
38
|
+
|
|
39
|
+
from dotenv import load_dotenv
|
|
40
|
+
from pydantic import BaseModel, Field, ValidationError
|
|
41
|
+
|
|
42
|
+
from kit.ai.exceptions import AIConfigError
|
|
43
|
+
from kit.config.kit_config import get_settings
|
|
44
|
+
from kit.config.log_config import get_default_logger
|
|
45
|
+
|
|
46
|
+
logger = get_default_logger(__name__)
|
|
47
|
+
|
|
48
|
+
# 支持的模型类型
|
|
49
|
+
_MODEL_TYPES: Final[tuple[str, ...]] = ("chat", "embedding", "rerank", "image")
|
|
50
|
+
|
|
51
|
+
# 文件缺失报错时附带的最小可用模板(api_key 为环境变量名,密钥配置在 .env)
|
|
52
|
+
_MINIMAL_MODELS_JSON: Final[str] = (
|
|
53
|
+
'{"chat": {"default_profile": "default", "profiles": {"default": '
|
|
54
|
+
'{"api_key": "SILICONFLOW_API_KEY", "base_url": "https://api.siliconflow.cn/v1", '
|
|
55
|
+
'"endpoint": "/chat/completions", "model": "Qwen/Qwen3-8B", "timeout": 120}}}}'
|
|
56
|
+
)
|
|
57
|
+
|
|
58
|
+
|
|
59
|
+
@dataclass(frozen=True)
|
|
60
|
+
class ModelProfile:
|
|
61
|
+
"""单个模型档案(运行时已解析密钥)。
|
|
62
|
+
|
|
63
|
+
Attributes:
|
|
64
|
+
name: 档案名(如 ``siliconflow``)
|
|
65
|
+
model_type: 模型类型(chat / embedding / rerank / image)
|
|
66
|
+
api_key_env: 密钥所在环境变量名
|
|
67
|
+
api_key: 已解析的密钥值
|
|
68
|
+
base_url: API 基础地址
|
|
69
|
+
endpoint: API 路径(以 ``/`` 开头)
|
|
70
|
+
model: 模型名称
|
|
71
|
+
timeout: 请求超时秒数
|
|
72
|
+
dimension: 向量维度(仅 embedding 有效)
|
|
73
|
+
batch_size: 单次请求最大文本数(仅 embedding 有效)
|
|
74
|
+
"""
|
|
75
|
+
|
|
76
|
+
name: str
|
|
77
|
+
model_type: str
|
|
78
|
+
api_key_env: str
|
|
79
|
+
api_key: str
|
|
80
|
+
base_url: str
|
|
81
|
+
endpoint: str
|
|
82
|
+
model: str
|
|
83
|
+
timeout: int = 60
|
|
84
|
+
dimension: int | None = None
|
|
85
|
+
batch_size: int = 32
|
|
86
|
+
|
|
87
|
+
@property
|
|
88
|
+
def url(self) -> str:
|
|
89
|
+
"""拼接完整请求 URL(base_url 去尾斜杠 + endpoint)"""
|
|
90
|
+
return self.base_url.rstrip("/") + self.endpoint
|
|
91
|
+
|
|
92
|
+
|
|
93
|
+
class ProfileState(enum.Enum):
|
|
94
|
+
"""档案运行时状态(候选链与状态快照使用)。"""
|
|
95
|
+
|
|
96
|
+
AVAILABLE = "available" # 密钥已配置且可用(探测通过或尚未验证)
|
|
97
|
+
COOLING = "cooling" # 调用失败冷却中(到期自动恢复)
|
|
98
|
+
INVALID_KEY = "invalid_key" # 探测判定密钥无效
|
|
99
|
+
UNREACHABLE = "unreachable" # 探测无法判定(超时 / 404 / 5xx)
|
|
100
|
+
UNCONFIGURED = "unconfigured" # 密钥未配置
|
|
101
|
+
EXCLUDED = "excluded" # 因约束被剔除(如 embedding 链内 model 不一致)
|
|
102
|
+
|
|
103
|
+
|
|
104
|
+
@dataclass(frozen=True)
|
|
105
|
+
class ProfileStatus:
|
|
106
|
+
"""档案状态快照行(``ModelRegistry.available`` / ``AI.available`` 返回)。
|
|
107
|
+
|
|
108
|
+
Attributes:
|
|
109
|
+
name: 档案名
|
|
110
|
+
model: 模型名称
|
|
111
|
+
state: 运行时状态
|
|
112
|
+
current: 是否为该类型当前使用项(链头)
|
|
113
|
+
last_error: 最近一次失败原因(无失败为 None)
|
|
114
|
+
last_ok_at: 最近一次成功/探测通过时刻(Unix 秒;从未成功为 None)
|
|
115
|
+
"""
|
|
116
|
+
|
|
117
|
+
name: str
|
|
118
|
+
model: str
|
|
119
|
+
state: ProfileState
|
|
120
|
+
current: bool = False
|
|
121
|
+
last_error: str | None = None
|
|
122
|
+
last_ok_at: float | None = None
|
|
123
|
+
|
|
124
|
+
|
|
125
|
+
class ProfileConfig(BaseModel):
|
|
126
|
+
"""档案配置(JSON 文件中单个 profile 的原始字段)。
|
|
127
|
+
|
|
128
|
+
注意:``api_key`` 存的是密钥所在环境变量名,
|
|
129
|
+
与 ``ModelProfile.api_key``(已解析的密钥值)同名不同义。
|
|
130
|
+
"""
|
|
131
|
+
|
|
132
|
+
api_key: str = Field(min_length=1)
|
|
133
|
+
base_url: str = Field(min_length=1)
|
|
134
|
+
endpoint: str = Field(min_length=1)
|
|
135
|
+
model: str = Field(min_length=1)
|
|
136
|
+
timeout: int = Field(default=60, gt=0)
|
|
137
|
+
dimension: int | None = Field(default=None, gt=0)
|
|
138
|
+
batch_size: int = Field(default=32, gt=0)
|
|
139
|
+
|
|
140
|
+
|
|
141
|
+
class TypeConfig(BaseModel):
|
|
142
|
+
"""某模型类型的档案集合"""
|
|
143
|
+
|
|
144
|
+
default_profile: str = Field(min_length=1)
|
|
145
|
+
profiles: dict[str, ProfileConfig]
|
|
146
|
+
|
|
147
|
+
|
|
148
|
+
class RegistryConfig(BaseModel):
|
|
149
|
+
"""模型档案文件整体结构(顶层对象)"""
|
|
150
|
+
|
|
151
|
+
chat: TypeConfig | None = None
|
|
152
|
+
embedding: TypeConfig | None = None
|
|
153
|
+
rerank: TypeConfig | None = None
|
|
154
|
+
image: TypeConfig | None = None
|
|
155
|
+
|
|
156
|
+
def type_config(self, model_type: str) -> TypeConfig:
|
|
157
|
+
config = getattr(self, model_type, None)
|
|
158
|
+
if not isinstance(config, TypeConfig):
|
|
159
|
+
available = ", ".join(_MODEL_TYPES)
|
|
160
|
+
raise AIConfigError(f"模型类型不存在: {model_type},可用类型: {available}")
|
|
161
|
+
return config
|
|
162
|
+
|
|
163
|
+
|
|
164
|
+
class ModelRegistry:
|
|
165
|
+
"""模型档案注册表。
|
|
166
|
+
|
|
167
|
+
用法::
|
|
168
|
+
|
|
169
|
+
registry = ModelRegistry()
|
|
170
|
+
registry.load("./config/ai_models.json")
|
|
171
|
+
profile = registry.get("chat", "siliconflow")
|
|
172
|
+
"""
|
|
173
|
+
|
|
174
|
+
def __init__(self) -> None:
|
|
175
|
+
self._config: RegistryConfig | None = None
|
|
176
|
+
|
|
177
|
+
def load(self, path: str | None = None) -> None:
|
|
178
|
+
"""加载并校验模型档案文件。
|
|
179
|
+
|
|
180
|
+
加载前将项目根 ``.env`` 注入 ``os.environ``(``override=False``:
|
|
181
|
+
真实环境变量优先,重复调用幂等),使密钥仅配置在 ``.env`` 时
|
|
182
|
+
``get()`` 也能解析——档案数量与密钥变量名由 JSON 决定,
|
|
183
|
+
无需改 ``sys_config``。
|
|
184
|
+
|
|
185
|
+
Args:
|
|
186
|
+
path: 档案文件路径,None 时取 ``Settings.AI_MODELS_CONFIG_PATH``
|
|
187
|
+
|
|
188
|
+
Raises:
|
|
189
|
+
AIConfigError: 文件缺失、JSON 非法、结构非法、档案必填字段缺失、
|
|
190
|
+
默认档案未配置。
|
|
191
|
+
"""
|
|
192
|
+
load_dotenv(Path.cwd() / ".env", override=False)
|
|
193
|
+
|
|
194
|
+
config_path = path or get_settings().AI_MODELS_CONFIG_PATH
|
|
195
|
+
file_path = Path(config_path)
|
|
196
|
+
|
|
197
|
+
if not file_path.is_file():
|
|
198
|
+
raise AIConfigError(
|
|
199
|
+
f"模型档案文件不存在: {config_path}。"
|
|
200
|
+
"请在项目根创建该文件(api_key 值为环境变量名,密钥配置在 .env),"
|
|
201
|
+
f"最小模板:\n{_MINIMAL_MODELS_JSON}\n"
|
|
202
|
+
"或调整 Settings.AI_MODELS_CONFIG_PATH 指向已有档案文件"
|
|
203
|
+
)
|
|
204
|
+
|
|
205
|
+
try:
|
|
206
|
+
self._config = RegistryConfig.model_validate_json(
|
|
207
|
+
file_path.read_text(encoding="utf-8")
|
|
208
|
+
)
|
|
209
|
+
except (OSError, ValidationError) as exc:
|
|
210
|
+
raise AIConfigError(f"模型档案文件解析失败: {config_path}: {exc}") from exc
|
|
211
|
+
|
|
212
|
+
loaded_types = [t for t in _MODEL_TYPES if getattr(self._config, t) is not None]
|
|
213
|
+
logger.info(
|
|
214
|
+
f"AI 模型档案已加载: {config_path},类型: {', '.join(loaded_types) or '无'}"
|
|
215
|
+
)
|
|
216
|
+
|
|
217
|
+
def _resolve_profile(
|
|
218
|
+
self, model_type: str, profile: str | None
|
|
219
|
+
) -> tuple[str, ProfileConfig]:
|
|
220
|
+
"""解析类型 + 档案名,返回 (档案名, 静态 ProfileConfig),不解析密钥。"""
|
|
221
|
+
if self._config is None:
|
|
222
|
+
raise AIConfigError("ModelRegistry 未加载,请先调用 load()")
|
|
223
|
+
|
|
224
|
+
type_config = self._config.type_config(model_type)
|
|
225
|
+
profile_name = profile or type_config.default_profile
|
|
226
|
+
profile_config = type_config.profiles.get(profile_name)
|
|
227
|
+
if profile_config is None:
|
|
228
|
+
raise AIConfigError(
|
|
229
|
+
f"档案不存在: {model_type}.{profile_name},"
|
|
230
|
+
f"可用档案: {', '.join(sorted(type_config.profiles))}"
|
|
231
|
+
)
|
|
232
|
+
return profile_name, profile_config
|
|
233
|
+
|
|
234
|
+
def get(self, model_type: str, profile: str | None = None) -> ModelProfile:
|
|
235
|
+
"""获取模型档案,profile 为 None 时回退到该类型默认档案。
|
|
236
|
+
|
|
237
|
+
Raises:
|
|
238
|
+
AIConfigError: 类型/档案不存在、密钥环境变量未配置。
|
|
239
|
+
"""
|
|
240
|
+
profile_name, profile_config = self._resolve_profile(model_type, profile)
|
|
241
|
+
|
|
242
|
+
api_key = os.environ.get(profile_config.api_key) or ""
|
|
243
|
+
if not api_key:
|
|
244
|
+
raise AIConfigError(
|
|
245
|
+
f"未配置环境变量 {profile_config.api_key},"
|
|
246
|
+
f"请检查 .env 中的 {profile_config.api_key}"
|
|
247
|
+
)
|
|
248
|
+
|
|
249
|
+
return self._build_profile(model_type, profile_name, profile_config, api_key)
|
|
250
|
+
|
|
251
|
+
@staticmethod
|
|
252
|
+
def _build_profile(
|
|
253
|
+
model_type: str, name: str, config: ProfileConfig, api_key: str
|
|
254
|
+
) -> ModelProfile:
|
|
255
|
+
"""由静态配置 + 已解析密钥构建运行时档案(get/chain 共用)。"""
|
|
256
|
+
return ModelProfile(
|
|
257
|
+
name=name,
|
|
258
|
+
model_type=model_type,
|
|
259
|
+
api_key_env=config.api_key,
|
|
260
|
+
api_key=api_key,
|
|
261
|
+
base_url=config.base_url,
|
|
262
|
+
endpoint=config.endpoint,
|
|
263
|
+
model=config.model,
|
|
264
|
+
timeout=config.timeout,
|
|
265
|
+
dimension=config.dimension,
|
|
266
|
+
batch_size=config.batch_size,
|
|
267
|
+
)
|
|
268
|
+
|
|
269
|
+
def get_profile(self, model_type: str, profile: str | None = None) -> ProfileConfig:
|
|
270
|
+
"""读取静态档案(不含密钥),用于维度/模型名等无需密钥的配置。"""
|
|
271
|
+
return self._resolve_profile(model_type, profile)[1]
|
|
272
|
+
|
|
273
|
+
def type_profiles(self, model_type: str) -> dict[str, ProfileConfig]:
|
|
274
|
+
"""返回该类型全部静态档案(含密钥未配置者),供状态快照展示。
|
|
275
|
+
|
|
276
|
+
Raises:
|
|
277
|
+
AIConfigError: 注册表未加载或类型不存在。
|
|
278
|
+
"""
|
|
279
|
+
if self._config is None:
|
|
280
|
+
raise AIConfigError("ModelRegistry 未加载,请先调用 load()")
|
|
281
|
+
return dict(self._config.type_config(model_type).profiles)
|
|
282
|
+
|
|
283
|
+
def chain(
|
|
284
|
+
self, model_type: str, *, env_profiles: str | None = None
|
|
285
|
+
) -> list[ModelProfile]:
|
|
286
|
+
"""解析该类型候选链:优先级链(env)> default 打头 + JSON 声明顺序。
|
|
287
|
+
|
|
288
|
+
规则:
|
|
289
|
+
1. ``env_profiles`` 参数或 ``Settings.AI_{TYPE}_PROFILES``(逗号
|
|
290
|
+
分隔档案名)设置时按其顺序;JSON 中不存在的档案名丢弃并告警;
|
|
291
|
+
2. 未设置时 ``default_profile`` 打头,其余按 JSON 声明顺序;
|
|
292
|
+
3. 密钥未配置的档案剔除出链(debug 日志)。
|
|
293
|
+
|
|
294
|
+
Args:
|
|
295
|
+
model_type: 模型类型(chat / embedding / rerank / image)
|
|
296
|
+
env_profiles: 显式优先级链(测试注入用);None 时读 Settings
|
|
297
|
+
|
|
298
|
+
Returns:
|
|
299
|
+
候选档案列表(密钥已解析;可能为空——全部未配置密钥)。
|
|
300
|
+
|
|
301
|
+
Raises:
|
|
302
|
+
AIConfigError: 注册表未加载或类型不存在。
|
|
303
|
+
"""
|
|
304
|
+
if self._config is None:
|
|
305
|
+
raise AIConfigError("ModelRegistry 未加载,请先调用 load()")
|
|
306
|
+
type_config = self._config.type_config(model_type)
|
|
307
|
+
|
|
308
|
+
env_value = env_profiles
|
|
309
|
+
if env_value is None:
|
|
310
|
+
settings = get_settings()
|
|
311
|
+
env_value = {
|
|
312
|
+
"chat": settings.AI_CHAT_PROFILES,
|
|
313
|
+
"embedding": settings.AI_EMBEDDING_PROFILES,
|
|
314
|
+
"rerank": settings.AI_RERANK_PROFILES,
|
|
315
|
+
"image": settings.AI_IMAGE_PROFILES,
|
|
316
|
+
}.get(model_type)
|
|
317
|
+
|
|
318
|
+
if env_value:
|
|
319
|
+
names = [name.strip() for name in env_value.split(",") if name.strip()]
|
|
320
|
+
unknown = [name for name in names if name not in type_config.profiles]
|
|
321
|
+
if unknown:
|
|
322
|
+
logger.warning(
|
|
323
|
+
f"{model_type} 候选链忽略 JSON 中不存在的档案: "
|
|
324
|
+
f"{', '.join(unknown)}"
|
|
325
|
+
)
|
|
326
|
+
order = [name for name in names if name in type_config.profiles]
|
|
327
|
+
else:
|
|
328
|
+
order = [
|
|
329
|
+
type_config.default_profile,
|
|
330
|
+
*(
|
|
331
|
+
name
|
|
332
|
+
for name in type_config.profiles
|
|
333
|
+
if name != type_config.default_profile
|
|
334
|
+
),
|
|
335
|
+
]
|
|
336
|
+
|
|
337
|
+
profiles: list[ModelProfile] = []
|
|
338
|
+
for name in order:
|
|
339
|
+
config = type_config.profiles[name]
|
|
340
|
+
api_key = os.environ.get(config.api_key) or ""
|
|
341
|
+
if not api_key:
|
|
342
|
+
logger.debug(
|
|
343
|
+
f"{model_type}.{name} 密钥未配置({config.api_key}),剔除出链"
|
|
344
|
+
)
|
|
345
|
+
continue
|
|
346
|
+
profiles.append(self._build_profile(model_type, name, config, api_key))
|
|
347
|
+
return profiles
|
|
348
|
+
|
|
349
|
+
|
|
350
|
+
def with_api_key(profile: ModelProfile, api_key: str) -> ModelProfile:
|
|
351
|
+
"""返回替换 api_key 后的新档案(ModelProfile 为不可变)。"""
|
|
352
|
+
return replace(profile, api_key=api_key)
|
|
353
|
+
|
|
354
|
+
|
|
355
|
+
@lru_cache(maxsize=1)
|
|
356
|
+
def get_model_registry() -> ModelRegistry:
|
|
357
|
+
"""获取全局单例 ModelRegistry(首次调用自动 load 默认档案文件)。"""
|
|
358
|
+
registry = ModelRegistry()
|
|
359
|
+
registry.load()
|
|
360
|
+
return registry
|
|
361
|
+
|
|
362
|
+
|
|
363
|
+
def embedding_dimension() -> int:
|
|
364
|
+
"""embedding 默认档案的向量维度(读静态字段,无需密钥)。"""
|
|
365
|
+
dimension = get_model_registry().get_profile("embedding").dimension
|
|
366
|
+
if dimension is None:
|
|
367
|
+
raise AIConfigError("ai_models.json 的 embedding 档案缺少 dimension")
|
|
368
|
+
return dimension
|
|
369
|
+
|
|
370
|
+
|
|
371
|
+
def embedding_model() -> str:
|
|
372
|
+
"""embedding 默认档案的模型名(读静态字段,无需密钥)。"""
|
|
373
|
+
return get_model_registry().get_profile("embedding").model
|
kit/ai/rerank_service.py
ADDED
|
@@ -0,0 +1,93 @@
|
|
|
1
|
+
"""
|
|
2
|
+
Copyright (c) 2021-2026 Clark Chang. All Rights Reserved.
|
|
3
|
+
|
|
4
|
+
Licensed under the Apache License, Version 2.0 (the "License");
|
|
5
|
+
you may not use this file except in compliance with the License.
|
|
6
|
+
You may obtain a copy of the License at
|
|
7
|
+
|
|
8
|
+
http://www.apache.org/licenses/LICENSE-2.0
|
|
9
|
+
|
|
10
|
+
Unless required by applicable law or agreed to in writing, software
|
|
11
|
+
distributed under the License is distributed on an "AS IS" BASIS,
|
|
12
|
+
WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
|
|
13
|
+
See the License for the specific language governing permissions and
|
|
14
|
+
limitations under the License.
|
|
15
|
+
|
|
16
|
+
Project: sdpy
|
|
17
|
+
Author: Clark Chang
|
|
18
|
+
"""
|
|
19
|
+
|
|
20
|
+
from dataclasses import dataclass
|
|
21
|
+
from typing import cast, override
|
|
22
|
+
|
|
23
|
+
from kit.ai.base_service import BaseAIService, safe_get
|
|
24
|
+
from kit.ai.exceptions import AIResponseError
|
|
25
|
+
from kit.ai.types import JSONObject
|
|
26
|
+
|
|
27
|
+
__all__ = ["RerankRequest", "RerankHit", "RerankService"]
|
|
28
|
+
|
|
29
|
+
|
|
30
|
+
@dataclass(frozen=True)
|
|
31
|
+
class RerankRequest:
|
|
32
|
+
"""重排序请求模型。"""
|
|
33
|
+
|
|
34
|
+
query: str
|
|
35
|
+
documents: list[str]
|
|
36
|
+
top_n: int | None = None
|
|
37
|
+
|
|
38
|
+
|
|
39
|
+
@dataclass(frozen=True)
|
|
40
|
+
class RerankHit:
|
|
41
|
+
"""重排序单条命中结果。"""
|
|
42
|
+
|
|
43
|
+
index: int
|
|
44
|
+
score: float
|
|
45
|
+
|
|
46
|
+
|
|
47
|
+
class RerankService(BaseAIService[RerankRequest, list[RerankHit]]):
|
|
48
|
+
"""重排序服务(全异步)。"""
|
|
49
|
+
|
|
50
|
+
@override
|
|
51
|
+
def _build_payload(self, request: RerankRequest) -> JSONObject:
|
|
52
|
+
payload: dict[str, object] = {
|
|
53
|
+
"model": self._profile.model,
|
|
54
|
+
"query": request.query,
|
|
55
|
+
"documents": request.documents,
|
|
56
|
+
}
|
|
57
|
+
if request.top_n is not None:
|
|
58
|
+
payload["top_n"] = request.top_n
|
|
59
|
+
return cast(JSONObject, payload)
|
|
60
|
+
|
|
61
|
+
@override
|
|
62
|
+
def _parse_response(self, data: JSONObject) -> list[RerankHit]:
|
|
63
|
+
items = safe_get(data, "results")
|
|
64
|
+
if not isinstance(items, list):
|
|
65
|
+
raise AIResponseError(f"重排序响应 results 必须是数组: {str(data)[:200]}")
|
|
66
|
+
hits: list[RerankHit] = []
|
|
67
|
+
for item in items:
|
|
68
|
+
if not isinstance(item, dict):
|
|
69
|
+
raise AIResponseError(f"重排序响应条目非法: {str(item)[:200]}")
|
|
70
|
+
index = item.get("index")
|
|
71
|
+
score = item.get("relevance_score")
|
|
72
|
+
if not isinstance(index, int) or not isinstance(score, (int, float)):
|
|
73
|
+
raise AIResponseError(f"重排序响应字段异常: {str(item)[:200]}")
|
|
74
|
+
hits.append(RerankHit(index=index, score=float(score)))
|
|
75
|
+
hits.sort(key=lambda h: h.score, reverse=True)
|
|
76
|
+
return hits
|
|
77
|
+
|
|
78
|
+
async def rerank(
|
|
79
|
+
self,
|
|
80
|
+
query: str,
|
|
81
|
+
documents: list[str],
|
|
82
|
+
*,
|
|
83
|
+
top_n: int | None = None,
|
|
84
|
+
) -> list[RerankHit]:
|
|
85
|
+
if not documents:
|
|
86
|
+
return []
|
|
87
|
+
hits = await self._call(RerankRequest(query, documents, top_n=top_n))
|
|
88
|
+
for hit in hits:
|
|
89
|
+
if hit.index < 0 or hit.index >= len(documents):
|
|
90
|
+
raise AIResponseError(
|
|
91
|
+
f"重排序 index 越界: {hit.index} (文档数 {len(documents)})"
|
|
92
|
+
)
|
|
93
|
+
return hits
|