refac
This commit is contained in:
@@ -72,20 +72,29 @@ log = logging.getLogger(__name__)
|
||||
##########################################
|
||||
|
||||
|
||||
async def send_get_request(url, key=None, user: UserModel = None):
|
||||
async def send_get_request(
|
||||
request: Request = None, url=None, key=None, user: UserModel = None, config=None,
|
||||
):
|
||||
timeout = aiohttp.ClientTimeout(total=AIOHTTP_CLIENT_TIMEOUT_MODEL_LIST)
|
||||
try:
|
||||
async with aiohttp.ClientSession(timeout=timeout, trust_env=True) as session:
|
||||
headers = {
|
||||
**({'Authorization': f'Bearer {key}'} if key else {}),
|
||||
}
|
||||
if request and config:
|
||||
headers, cookies = await get_headers_and_cookies(
|
||||
request, url, key, config, user=user
|
||||
)
|
||||
else:
|
||||
headers = {
|
||||
**({'Authorization': f'Bearer {key}'} if key else {}),
|
||||
}
|
||||
cookies = None
|
||||
|
||||
if ENABLE_FORWARD_USER_INFO_HEADERS and user:
|
||||
headers = include_user_info_headers(headers, user)
|
||||
if ENABLE_FORWARD_USER_INFO_HEADERS and user:
|
||||
headers = include_user_info_headers(headers, user)
|
||||
|
||||
async with session.get(
|
||||
url,
|
||||
headers=headers,
|
||||
cookies=cookies,
|
||||
ssl=AIOHTTP_CLIENT_SESSION_SSL,
|
||||
) as response:
|
||||
return await response.json()
|
||||
@@ -95,10 +104,12 @@ async def send_get_request(url, key=None, user: UserModel = None):
|
||||
return None
|
||||
|
||||
|
||||
async def get_models_request(url, key=None, user: UserModel = None):
|
||||
async def get_models_request(
|
||||
request: Request = None, url=None, key=None, user: UserModel = None, config=None,
|
||||
):
|
||||
if is_anthropic_url(url):
|
||||
return await get_anthropic_models(url, key, user=user)
|
||||
return await send_get_request(f'{url}/models', key, user=user)
|
||||
return await send_get_request(request, f'{url}/models', key, user=user, config=config)
|
||||
|
||||
|
||||
def openai_reasoning_model_handler(payload):
|
||||
@@ -360,7 +371,7 @@ async def get_all_models_responses(request: Request, user: UserModel) -> list:
|
||||
request_tasks = []
|
||||
for idx, url in enumerate(api_base_urls):
|
||||
if (str(idx) not in api_configs) and (url not in api_configs): # Legacy support
|
||||
request_tasks.append(get_models_request(url, api_keys[idx], user=user))
|
||||
request_tasks.append(get_models_request(request, url, api_keys[idx], user=user))
|
||||
else:
|
||||
api_config = api_configs.get(
|
||||
str(idx),
|
||||
@@ -372,7 +383,9 @@ async def get_all_models_responses(request: Request, user: UserModel) -> list:
|
||||
|
||||
if enable:
|
||||
if len(model_ids) == 0:
|
||||
request_tasks.append(get_models_request(url, api_keys[idx], user=user))
|
||||
request_tasks.append(
|
||||
get_models_request(request, url, api_keys[idx], user=user, config=api_config)
|
||||
)
|
||||
else:
|
||||
model_list = {
|
||||
'object': 'list',
|
||||
|
||||
Reference in New Issue
Block a user