ForcePilot/src/models/chat.py
2025-10-11 02:34:42 +08:00

171 lines
5.3 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

import os
import traceback
from openai import OpenAI
from tenacity import before_sleep_log, retry, retry_if_exception_type, stop_after_attempt, wait_exponential
from src import config
from src.utils import logger
class OpenAIBase:
def __init__(self, api_key, base_url, model_name, **kwargs):
self.api_key = api_key
self.base_url = base_url
self.client = OpenAI(api_key=api_key, base_url=base_url)
self.model_name = model_name
self.info = kwargs
@retry(
stop=stop_after_attempt(3),
wait=wait_exponential(multiplier=1, min=1, max=10),
retry=retry_if_exception_type((Exception,)),
before_sleep=before_sleep_log(logger, log_level="WARNING"),
reraise=True,
)
def call(self, message, stream=False):
if isinstance(message, str):
messages = [{"role": "user", "content": message}]
else:
messages = message
try:
if stream:
response = self._stream_response(messages)
else:
response = self._get_response(messages)
except Exception as e:
err = (
f"Error streaming response: {e}, URL: {self.base_url}, "
f"API Key: {self.api_key[:5]}***, Model: {self.model_name}"
)
logger.error(err)
raise Exception(err)
return response
def _stream_response(self, messages):
response = self.client.chat.completions.create(
model=self.model_name,
messages=messages,
stream=True,
)
for chunk in response:
if len(chunk.choices) > 0:
yield chunk.choices[0].delta
def _get_response(self, messages):
response = self.client.chat.completions.create(
model=self.model_name,
messages=messages,
stream=False,
)
return response.choices[0].message
def get_models(self):
try:
return self.client.models.list(extra_query={"type": "text"})
except Exception as e:
logger.error(f"Error getting models: {e}")
return []
class OpenModel(OpenAIBase):
def __init__(self, model_name=None):
model_name = model_name or "gpt-4o-mini"
api_key = os.getenv("OPENAI_API_KEY")
base_url = os.getenv("OPENAI_API_BASE")
super().__init__(api_key=api_key, base_url=base_url, model_name=model_name)
class GeneralResponse:
def __init__(self, content):
self.content = content
self.is_full = False
def select_model(model_provider, model_name=None):
"""根据模型提供者选择模型"""
assert model_provider is not None, "Model provider not specified"
model_info = config.model_names.get(model_provider, {})
model_name = model_name or model_info.get("default", "")
logger.info(f"Selecting model from `{model_provider}` with `{model_name}`")
if model_provider == "openai":
return OpenModel(model_name)
# 其他模型默认使用OpenAIBase
try:
model = OpenAIBase(
api_key=os.getenv(model_info["env"]),
base_url=model_info["base_url"],
model_name=model_name,
)
return model
except Exception as e:
raise ValueError(f"Model provider {model_provider} load failed, {e} \n {traceback.format_exc()}")
async def test_chat_model_status(provider: str, model_name: str) -> dict:
"""
测试指定聊天模型的状态
Args:
provider: 模型提供商
model_name: 模型名称
Returns:
dict: 包含状态信息的字典
"""
try:
# 加载模型
logger.debug(f"Selecting chat model {provider}/{model_name}")
model = select_model(provider, model_name)
# 使用简单的测试消息
test_messages = [{"role": "user", "content": "Say 1"}]
# 发送测试请求
response = model.call(test_messages, stream=False)
logger.debug(f"Test chat model status response: {response}")
# 检查响应是否有效
if response and response.content:
return {"provider": provider, "model_name": model_name, "status": "available", "message": "连接正常"}
else:
return {"provider": provider, "model_name": model_name, "status": "unavailable", "message": "响应无效"}
except Exception as e:
logger.error(f"测试聊天模型状态失败 {provider}/{model_name}: {e}")
return {"provider": provider, "model_name": model_name, "status": "error", "message": str(e)}
async def test_all_chat_models_status() -> dict:
"""
测试所有支持的聊天模型状态
Returns:
dict: 包含所有模型状态的字典
"""
from src import config
results = {}
# 获取所有可用的模型
for provider, provider_info in config.model_names.items():
# 处理普通模型
for model_name in provider_info.models:
model_id = f"{provider}/{model_name}"
status = await test_chat_model_status(provider, model_name)
results[model_id] = status
available_count = len([m for m in results.values() if m["status"] == "available"])
return {"models": results, "total": len(results), "available": available_count}
if __name__ == "__main__":
pass