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/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
@@ -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