11from urllib .parse import urlparse
22
33import httpx
4+ import openai
45from openai import AsyncOpenAI
56
67from 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