mirror of
https://github.com/zhayujie/chatgpt-on-wechat.git
synced 2026-07-20 13:47:15 +08:00
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:
@@ -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":
|
||||||
|
|||||||
Reference in New Issue
Block a user