Skip to content

Commit 108b905

Browse files
committed
refactor(openai-embedding): apply code quality improvements, exception handling and empty validation from 1.md
1 parent 5ad2927 commit 108b905

1 file changed

Lines changed: 57 additions & 16 deletions

File tree

astrbot/core/provider/sources/openai_embedding_source.py

Lines changed: 57 additions & 16 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
from urllib.parse import urlparse
22

33
import httpx
4+
import openai
45
from openai import AsyncOpenAI
56

67
from astrbot import logger
@@ -20,12 +21,29 @@ def __init__(self, provider_config: dict, provider_settings: dict) -> None:
2021
super().__init__(provider_config, provider_settings)
2122
self.provider_config = provider_config
2223
self.provider_settings = provider_settings
23-
proxy = provider_config.get("proxy", "")
24+
2425
provider_id = provider_config.get("id", "unknown_id")
26+
27+
# 1. 强制校验 API Key (Fail-Fast)
28+
api_key: str = provider_config.get("embedding_api_key", "")
29+
if not api_key:
30+
raise ValueError(
31+
f"OpenAI embedding provider [{provider_id}] 配置错误: 缺少必需的 'embedding_api_key'"
32+
)
33+
34+
# 2. 安全获取并转换 timeout 避免空字符串导致 int() 崩溃
35+
raw_timeout = provider_config.get("timeout", 20)
36+
try:
37+
timeout_val = int(raw_timeout) if raw_timeout else 20
38+
except (ValueError, TypeError):
39+
timeout_val = 20
40+
41+
proxy = provider_config.get("proxy", "")
2542
http_client = None
2643
if proxy:
2744
logger.info(f"[OpenAI Embedding] {provider_id} Using proxy: {proxy}")
2845
http_client = httpx.AsyncClient(proxy=proxy)
46+
2947
api_base = (
3048
provider_config.get("embedding_api_base", "https://api.openai.com/v1")
3149
.strip()
@@ -35,34 +53,56 @@ def __init__(self, provider_config: dict, provider_settings: dict) -> None:
3553
if api_base and not api_base.endswith("/v1") and not api_base.endswith("/v4"):
3654
# /v4 see #5699
3755
api_base = api_base + "/v1"
56+
3857
logger.info(f"[OpenAI Embedding] {provider_id} Using API Base: {api_base}")
58+
3959
self.client = AsyncOpenAI(
40-
api_key=provider_config.get("embedding_api_key"),
60+
api_key=api_key,
4161
base_url=api_base,
42-
timeout=int(provider_config.get("timeout", 20)),
62+
timeout=timeout_val,
4363
http_client=http_client,
4464
)
4565
self.model = provider_config.get("embedding_model", "text-embedding-3-small")
4666

4767
async def get_embedding(self, text: str) -> list[float]:
4868
"""获取文本的嵌入"""
69+
# 3. 拦截空文本防 400 报错
70+
if not text or not text.strip():
71+
raise ValueError("输入文本不能为空")
72+
4973
kwargs = self._embedding_kwargs()
50-
embedding = await self.client.embeddings.create(
51-
input=text,
52-
model=self.model,
53-
**kwargs,
54-
)
55-
return embedding.data[0].embedding
74+
75+
try:
76+
embedding = await self.client.embeddings.create(
77+
input=text,
78+
model=self.model,
79+
**kwargs,
80+
)
81+
return embedding.data[0].embedding
82+
except openai.OpenAIError as e:
83+
# 4. 包装规范异常
84+
raise Exception(f"OpenAI Embedding API 请求失败: {str(e)}") from e
5685

5786
async def get_embeddings(self, text: list[str]) -> list[list[float]]:
5887
"""批量获取文本的嵌入"""
88+
# 5. 拦截空列表和内部脏数据
89+
if not text:
90+
return []
91+
92+
if any(not s or not s.strip() for s in text):
93+
raise ValueError("批量输入文本列表中不能包含空文本")
94+
5995
kwargs = self._embedding_kwargs()
60-
embeddings = await self.client.embeddings.create(
61-
input=text,
62-
model=self.model,
63-
**kwargs,
64-
)
65-
return [item.embedding for item in embeddings.data]
96+
97+
try:
98+
embeddings = await self.client.embeddings.create(
99+
input=text,
100+
model=self.model,
101+
**kwargs,
102+
)
103+
return [item.embedding for item in embeddings.data]
104+
except openai.OpenAIError as e:
105+
raise Exception(f"OpenAI Embedding API 批量请求失败: {str(e)}") from e
66106

67107
def _embedding_kwargs(self) -> dict:
68108
"""构建嵌入请求的可选参数"""
@@ -111,5 +151,6 @@ def get_dim(self) -> int:
111151
return 0
112152

113153
async def terminate(self):
114-
if self.client:
154+
"""释放资源"""
155+
if getattr(self, "client", None):
115156
await self.client.close()

0 commit comments

Comments
 (0)