mirror of
https://github.com/zhayujie/chatgpt-on-wechat.git
synced 2026-07-20 21:57:14 +08:00
compatible with openai bot
This commit is contained in:
@@ -2,6 +2,7 @@
|
|||||||
|
|
||||||
import requests
|
import requests
|
||||||
from bot.bot import Bot
|
from bot.bot import Bot
|
||||||
|
from bridge.reply import Reply, ReplyType
|
||||||
|
|
||||||
|
|
||||||
# Baidu Unit对话接口 (可用, 但能力较弱)
|
# Baidu Unit对话接口 (可用, 但能力较弱)
|
||||||
@@ -14,7 +15,8 @@ class BaiduUnitBot(Bot):
|
|||||||
headers = {'content-type': 'application/x-www-form-urlencoded'}
|
headers = {'content-type': 'application/x-www-form-urlencoded'}
|
||||||
response = requests.post(url, data=post_data.encode(), headers=headers)
|
response = requests.post(url, data=post_data.encode(), headers=headers)
|
||||||
if response:
|
if response:
|
||||||
return response.json()['result']['context']['SYS_PRESUMED_HIST'][1]
|
reply = Reply(ReplyType.TEXT, response.json()['result']['context']['SYS_PRESUMED_HIST'][1])
|
||||||
|
return reply
|
||||||
|
|
||||||
def get_token(self):
|
def get_token(self):
|
||||||
access_key = 'YOUR_ACCESS_KEY'
|
access_key = 'YOUR_ACCESS_KEY'
|
||||||
|
|||||||
@@ -1,6 +1,8 @@
|
|||||||
# encoding:utf-8
|
# encoding:utf-8
|
||||||
|
|
||||||
from bot.bot import Bot
|
from bot.bot import Bot
|
||||||
|
from bridge.context import ContextType
|
||||||
|
from bridge.reply import Reply, ReplyType
|
||||||
from config import conf
|
from config import conf
|
||||||
from common.log import logger
|
from common.log import logger
|
||||||
import openai
|
import openai
|
||||||
@@ -13,30 +15,31 @@ class OpenAIBot(Bot):
|
|||||||
def __init__(self):
|
def __init__(self):
|
||||||
openai.api_key = conf().get('open_ai_api_key')
|
openai.api_key = conf().get('open_ai_api_key')
|
||||||
|
|
||||||
|
|
||||||
def reply(self, query, context=None):
|
def reply(self, query, context=None):
|
||||||
# acquire reply content
|
# acquire reply content
|
||||||
if not context or not context.get('type') or context.get('type') == 'TEXT':
|
if context and context.type:
|
||||||
logger.info("[OPEN_AI] query={}".format(query))
|
if context.type == ContextType.TEXT:
|
||||||
from_user_id = context.get('from_user_id') or context.get('session_id')
|
logger.info("[OPEN_AI] query={}".format(query))
|
||||||
if query == '#清除记忆':
|
from_user_id = context['session_id']
|
||||||
Session.clear_session(from_user_id)
|
reply = None
|
||||||
return '记忆已清除'
|
if query == '#清除记忆':
|
||||||
elif query == '#清除所有':
|
Session.clear_session(from_user_id)
|
||||||
Session.clear_all_session()
|
reply = Reply(ReplyType.INFO, '记忆已清除')
|
||||||
return '所有人记忆已清除'
|
elif query == '#清除所有':
|
||||||
|
Session.clear_all_session()
|
||||||
|
reply = Reply(ReplyType.INFO, '所有人记忆已清除')
|
||||||
|
else:
|
||||||
|
new_query = Session.build_session_query(query, from_user_id)
|
||||||
|
logger.debug("[OPEN_AI] session query={}".format(new_query))
|
||||||
|
|
||||||
new_query = Session.build_session_query(query, from_user_id)
|
reply_content = self.reply_text(new_query, from_user_id, 0)
|
||||||
logger.debug("[OPEN_AI] session query={}".format(new_query))
|
logger.debug("[OPEN_AI] new_query={}, user={}, reply_cont={}".format(new_query, from_user_id, reply_content))
|
||||||
|
if reply_content and query:
|
||||||
reply_content = self.reply_text(new_query, from_user_id, 0)
|
Session.save_session(query, reply_content, from_user_id)
|
||||||
logger.debug("[OPEN_AI] new_query={}, user={}, reply_cont={}".format(new_query, from_user_id, reply_content))
|
reply = Reply(ReplyType.TEXT, reply_content)
|
||||||
if reply_content and query:
|
return reply
|
||||||
Session.save_session(query, reply_content, from_user_id)
|
elif context.type == ContextType.IMAGE_CREATE:
|
||||||
return reply_content
|
return self.create_img(query, 0)
|
||||||
|
|
||||||
elif context.get('type', None) == 'IMAGE_CREATE':
|
|
||||||
return self.create_img(query, 0)
|
|
||||||
|
|
||||||
def reply_text(self, query, user_id, retry_count=0):
|
def reply_text(self, query, user_id, retry_count=0):
|
||||||
try:
|
try:
|
||||||
|
|||||||
@@ -162,7 +162,7 @@ class Godcmd(Plugin):
|
|||||||
bot.sessions.clear_session(session_id)
|
bot.sessions.clear_session(session_id)
|
||||||
ok, result = True, "会话已重置"
|
ok, result = True, "会话已重置"
|
||||||
else:
|
else:
|
||||||
ok, result = False, "当前机器人不支持重置会话"
|
ok, result = False, "当前对话机器人不支持重置会话"
|
||||||
logger.debug("[Godcmd] command: %s by %s" % (cmd, user))
|
logger.debug("[Godcmd] command: %s by %s" % (cmd, user))
|
||||||
elif any(cmd in info['alias'] for info in ADMIN_COMMANDS.values()):
|
elif any(cmd in info['alias'] for info in ADMIN_COMMANDS.values()):
|
||||||
if isadmin:
|
if isadmin:
|
||||||
@@ -184,7 +184,7 @@ class Godcmd(Plugin):
|
|||||||
bot.sessions.clear_all_session()
|
bot.sessions.clear_all_session()
|
||||||
ok, result = True, "重置所有会话成功"
|
ok, result = True, "重置所有会话成功"
|
||||||
else:
|
else:
|
||||||
ok, result = False, "当前机器人不支持重置会话"
|
ok, result = False, "当前对话机器人不支持重置会话"
|
||||||
elif cmd == "debug":
|
elif cmd == "debug":
|
||||||
logger.setLevel('DEBUG')
|
logger.setLevel('DEBUG')
|
||||||
ok, result = True, "DEBUG模式已开启"
|
ok, result = True, "DEBUG模式已开启"
|
||||||
|
|||||||
Reference in New Issue
Block a user