mirror of
https://github.com/langbot-app/LangBot.git
synced 2026-07-24 21:36:06 +00:00
Merge pull request #219 from LINSTCL/modelmgr_optimization
优化模型接口底层的异常处理
This commit is contained in:
+24
-6
@@ -1,5 +1,6 @@
|
|||||||
# 提供与模型交互的抽象接口
|
# 提供与模型交互的抽象接口
|
||||||
import openai, logging, threading, asyncio
|
import openai, logging, threading, asyncio
|
||||||
|
import openai.error as aiE
|
||||||
|
|
||||||
COMPLETION_MODELS = {
|
COMPLETION_MODELS = {
|
||||||
'text-davinci-003',
|
'text-davinci-003',
|
||||||
@@ -29,22 +30,37 @@ class ModelRequest():
|
|||||||
"""GPT父类"""
|
"""GPT父类"""
|
||||||
can_chat = False
|
can_chat = False
|
||||||
runtime:threading.Thread = None
|
runtime:threading.Thread = None
|
||||||
ret = ""
|
ret = {}
|
||||||
proxy:str = None
|
proxy:str = None
|
||||||
|
request_ready = True
|
||||||
|
error_info:str = "若在没有任何错误的情况下看到这句话,请带着配置文件上报Issues"
|
||||||
|
|
||||||
def __init__(self, model_name, user_name, request_fun, http_proxy:str = None):
|
def __init__(self, model_name, user_name, request_fun, http_proxy:str = None, time_out = None):
|
||||||
self.model_name = model_name
|
self.model_name = model_name
|
||||||
self.user_name = user_name
|
self.user_name = user_name
|
||||||
self.request_fun = request_fun
|
self.request_fun = request_fun
|
||||||
|
self.time_out = time_out
|
||||||
if http_proxy != None:
|
if http_proxy != None:
|
||||||
self.proxy = http_proxy
|
self.proxy = http_proxy
|
||||||
openai.proxy = self.proxy
|
openai.proxy = self.proxy
|
||||||
|
self.request_ready = False
|
||||||
|
|
||||||
async def __a_request__(self, **kwargs):
|
async def __a_request__(self, **kwargs):
|
||||||
self.ret = await self.request_fun(**kwargs)
|
try:
|
||||||
|
self.ret:dict = await self.request_fun(**kwargs)
|
||||||
|
self.request_ready = True
|
||||||
|
except aiE.APIConnectionError as e:
|
||||||
|
self.error_info = "{}\n请检查网络连接或代理是否正常".format(e)
|
||||||
|
raise ConnectionError(self.error_info)
|
||||||
|
except ValueError as e:
|
||||||
|
self.error_info = "{}\n该错误可能是由于http_proxy格式设置错误引起的"
|
||||||
|
except Exception as e:
|
||||||
|
self.error_info = "{}\n由于请求异常产生的未知错误,请查看日志".format(e)
|
||||||
|
raise Exception(self.error_info)
|
||||||
|
|
||||||
def request(self, **kwargs):
|
def request(self, **kwargs):
|
||||||
if self.proxy != None: #异步请求
|
if self.proxy != None: #异步请求
|
||||||
|
self.request_ready = False
|
||||||
loop = asyncio.new_event_loop()
|
loop = asyncio.new_event_loop()
|
||||||
self.runtime = threading.Thread(
|
self.runtime = threading.Thread(
|
||||||
target=loop.run_until_complete,
|
target=loop.run_until_complete,
|
||||||
@@ -64,13 +80,15 @@ class ModelRequest():
|
|||||||
若重写该方法,应检查异步线程状态,或在需要检查处super该方法
|
若重写该方法,应检查异步线程状态,或在需要检查处super该方法
|
||||||
'''
|
'''
|
||||||
if self.runtime != None and isinstance(self.runtime, threading.Thread):
|
if self.runtime != None and isinstance(self.runtime, threading.Thread):
|
||||||
self.runtime.join()
|
self.runtime.join(self.time_out)
|
||||||
return
|
if self.request_ready:
|
||||||
|
return
|
||||||
|
raise Exception(self.error_info)
|
||||||
|
|
||||||
def get_total_tokens(self):
|
def get_total_tokens(self):
|
||||||
try:
|
try:
|
||||||
return self.ret['usage']['total_tokens']
|
return self.ret['usage']['total_tokens']
|
||||||
except Exception:
|
except:
|
||||||
return 0
|
return 0
|
||||||
|
|
||||||
def get_message(self):
|
def get_message(self):
|
||||||
|
|||||||
Reference in New Issue
Block a user