2024-07-07 01:58:23 +08:00
|
|
|
import os
|
2025-04-09 11:47:08 +08:00
|
|
|
import requests
|
2024-07-07 01:58:23 +08:00
|
|
|
from openai import OpenAI
|
2025-03-04 13:49:00 +08:00
|
|
|
from src.utils import logger, get_docker_safe_url
|
2025-03-24 19:07:51 +08:00
|
|
|
from langchain_openai import ChatOpenAI
|
2024-07-07 01:58:23 +08:00
|
|
|
|
2025-05-24 11:29:45 +08:00
|
|
|
class OpenAIBase:
|
2025-07-02 02:38:36 +08:00
|
|
|
def __init__(self, api_key, base_url, model_name, **kwargs):
|
2025-03-10 16:33:40 +08:00
|
|
|
self.api_key = api_key
|
|
|
|
|
self.base_url = base_url
|
2024-07-07 01:58:23 +08:00
|
|
|
self.client = OpenAI(api_key=api_key, base_url=base_url)
|
|
|
|
|
self.model_name = model_name
|
2025-04-09 11:47:08 +08:00
|
|
|
self.info = kwargs
|
2024-07-07 01:58:23 +08:00
|
|
|
|
|
|
|
|
def predict(self, message, stream=False):
|
|
|
|
|
if isinstance(message, str):
|
|
|
|
|
messages=[{"role": "user", "content": message}]
|
|
|
|
|
else:
|
|
|
|
|
messages = message
|
|
|
|
|
|
|
|
|
|
if stream:
|
|
|
|
|
return self._stream_response(messages)
|
|
|
|
|
else:
|
|
|
|
|
return self._get_response(messages)
|
|
|
|
|
|
|
|
|
|
def _stream_response(self, messages):
|
2025-05-09 14:48:32 +08:00
|
|
|
try:
|
|
|
|
|
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
|
|
|
|
|
|
|
|
|
|
except Exception as e:
|
|
|
|
|
err = f"Error streaming response: {e}, URL: {self.base_url}, API Key: {self.api_key[:5]}***, Model: {self.model_name}"
|
|
|
|
|
logger.error(err)
|
|
|
|
|
raise Exception(err)
|
2024-07-07 01:58:23 +08:00
|
|
|
|
|
|
|
|
def _get_response(self, messages):
|
|
|
|
|
response = self.client.chat.completions.create(
|
|
|
|
|
model=self.model_name,
|
|
|
|
|
messages=messages,
|
|
|
|
|
stream=False,
|
|
|
|
|
)
|
|
|
|
|
return response.choices[0].message
|
2025-03-24 19:07:51 +08:00
|
|
|
|
2025-03-10 16:33:40 +08:00
|
|
|
def get_models(self):
|
|
|
|
|
try:
|
2025-04-01 00:39:54 +08:00
|
|
|
return self.client.models.list(
|
|
|
|
|
extra_query={
|
|
|
|
|
"type": "text"
|
|
|
|
|
}
|
|
|
|
|
)
|
2025-03-10 16:33:40 +08:00
|
|
|
except Exception as e:
|
|
|
|
|
logger.error(f"Error getting models: {e}")
|
|
|
|
|
return []
|
2024-07-07 01:58:23 +08:00
|
|
|
|
|
|
|
|
|
2024-09-09 17:07:03 +08:00
|
|
|
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)
|
|
|
|
|
|
|
|
|
|
|
2025-02-28 14:39:03 +08:00
|
|
|
|
2024-10-08 22:16:17 +08:00
|
|
|
class CustomModel(OpenAIBase):
|
|
|
|
|
def __init__(self, model_info):
|
|
|
|
|
model_name = model_info["name"]
|
2025-03-29 19:19:43 +08:00
|
|
|
api_key = model_info.get("api_key") or "custom_model"
|
2025-02-28 14:39:03 +08:00
|
|
|
base_url = get_docker_safe_url(model_info["api_base"])
|
2025-03-11 14:18:48 +08:00
|
|
|
logger.info(f"> Custom model: {model_name}, base_url: {base_url}")
|
2025-02-28 14:39:03 +08:00
|
|
|
|
2024-10-08 22:16:17 +08:00
|
|
|
super().__init__(api_key=api_key, base_url=base_url, model_name=model_name)
|
|
|
|
|
|
2024-07-07 17:21:07 +08:00
|
|
|
|
2024-07-31 20:22:05 +08:00
|
|
|
class GeneralResponse:
|
2024-07-07 17:21:07 +08:00
|
|
|
def __init__(self, content):
|
|
|
|
|
self.content = content
|
2024-07-31 20:22:05 +08:00
|
|
|
self.is_full = False
|
2024-07-07 17:21:07 +08:00
|
|
|
|
|
|
|
|
|
2025-03-10 16:33:40 +08:00
|
|
|
class Qianfan(OpenAIBase):
|
|
|
|
|
"""弃用"""
|
2024-07-07 17:21:07 +08:00
|
|
|
|
2024-07-18 02:47:41 +08:00
|
|
|
def __init__(self, model_name="ernie_speed") -> None:
|
2024-08-07 17:24:46 +08:00
|
|
|
import qianfan
|
2024-07-07 17:21:07 +08:00
|
|
|
self.model_name = model_name
|
|
|
|
|
access_key = os.getenv("QIANFAN_ACCESS_KEY")
|
|
|
|
|
secret_key = os.getenv("QIANFAN_SECRET_KEY")
|
|
|
|
|
self.client = qianfan.ChatCompletion(ak=access_key, sk=secret_key)
|
|
|
|
|
|
|
|
|
|
def predict(self, message, stream=False):
|
|
|
|
|
if isinstance(message, str):
|
|
|
|
|
messages=[{"role": "user", "content": message}]
|
|
|
|
|
else:
|
|
|
|
|
messages = message
|
|
|
|
|
|
|
|
|
|
if stream:
|
|
|
|
|
return self._stream_response(messages)
|
|
|
|
|
else:
|
|
|
|
|
return self._get_response(messages)
|
|
|
|
|
|
|
|
|
|
def _stream_response(self, messages):
|
|
|
|
|
response = self.client.do(
|
|
|
|
|
model=self.model_name,
|
|
|
|
|
messages=messages,
|
|
|
|
|
stream=True,
|
|
|
|
|
)
|
|
|
|
|
for chunk in response:
|
2024-07-31 20:22:05 +08:00
|
|
|
yield GeneralResponse(chunk["body"]["result"])
|
2024-07-07 17:21:07 +08:00
|
|
|
|
|
|
|
|
def _get_response(self, messages):
|
|
|
|
|
response = self.client.do(
|
|
|
|
|
model=self.model_name,
|
|
|
|
|
messages=messages,
|
|
|
|
|
stream=False,
|
|
|
|
|
)
|
2024-07-31 20:22:05 +08:00
|
|
|
return GeneralResponse(response["body"]["result"])
|
2024-07-20 15:27:33 +08:00
|
|
|
|
|
|
|
|
|
2024-07-31 20:22:05 +08:00
|
|
|
if __name__ == "__main__":
|
2025-05-24 11:29:45 +08:00
|
|
|
pass
|