fix: gemini doesn't receive system messages; change session to gpt method, add system messages as user messages to the gemini, and logging historical messages

This commit is contained in:
cmgzn
2024-09-20 09:10:21 +01:00
parent e4e1e2e944
commit f20d704390

View File

@@ -13,7 +13,7 @@ from bridge.context import ContextType, Context
from bridge.reply import Reply, ReplyType from bridge.reply import Reply, ReplyType
from common.log import logger from common.log import logger
from config import conf from config import conf
from bot.baidu.baidu_wenxin_session import BaiduWenxinSession from bot.chatgpt.chat_gpt_session import ChatGPTSession
# OpenAI对话模型API (可用) # OpenAI对话模型API (可用)
@@ -23,7 +23,7 @@ class GoogleGeminiBot(Bot):
super().__init__() super().__init__()
self.api_key = conf().get("gemini_api_key") self.api_key = conf().get("gemini_api_key")
# 复用文心的token计算方式 # 复用文心的token计算方式
self.sessions = SessionManager(BaiduWenxinSession, model=conf().get("model") or "gpt-3.5-turbo") self.sessions = SessionManager(ChatGPTSession, model=conf().get("model") or "gpt-3.5-turbo")
self.model = conf().get("model") or "gemini-pro" self.model = conf().get("model") or "gemini-pro"
if self.model == "gemini": if self.model == "gemini":
self.model = "gemini-pro" self.model = "gemini-pro"
@@ -36,6 +36,7 @@ class GoogleGeminiBot(Bot):
session_id = context["session_id"] session_id = context["session_id"]
session = self.sessions.session_query(query, session_id) session = self.sessions.session_query(query, session_id)
gemini_messages = self._convert_to_gemini_messages(self.filter_messages(session.messages)) gemini_messages = self._convert_to_gemini_messages(self.filter_messages(session.messages))
logger.info(f"[Gemini] messages={gemini_messages}")
genai.configure(api_key=self.api_key) genai.configure(api_key=self.api_key)
model = genai.GenerativeModel(self.model) model = genai.GenerativeModel(self.model)
response = model.generate_content(gemini_messages) response = model.generate_content(gemini_messages)
@@ -55,6 +56,8 @@ class GoogleGeminiBot(Bot):
role = "user" role = "user"
elif msg.get("role") == "assistant": elif msg.get("role") == "assistant":
role = "model" role = "model"
elif msg.get("role") == "system":
role = "user"
else: else:
continue continue
res.append({ res.append({
@@ -71,7 +74,11 @@ class GoogleGeminiBot(Bot):
return res return res
for i in range(len(messages) - 1, -1, -1): for i in range(len(messages) - 1, -1, -1):
message = messages[i] message = messages[i]
if message.get("role") != turn: role = message.get("role")
if role == "system":
res.insert(0, message)
continue
if role != turn:
continue continue
res.insert(0, message) res.insert(0, message)
if turn == "user": if turn == "user":