From 88ebdf1577266389813867b212ea3c2b24eeef9a Mon Sep 17 00:00:00 2001
From: teddybear082 <87204721+teddybear082@users.noreply.github.com>
Date: Tue, 3 Oct 2023 13:24:24 -0400
Subject: [PATCH] Implement basic chat/completions openai endpoint
-Basic support for openai chat/completions endpoint documented at: https://platform.openai.com/docs/api-reference/chat/create
-Tested with example code from openai for chat/completions and chat/completions with stream=True parameter found here: https://cookbook.openai.com/examples/how_to_stream_completions.
-Tested with Mantella, the skyrim mod that turns all the NPC's into AI chattable characters, which uses openai's acreate / async competions method: https://github.com/art-from-the-machine/Mantella/blob/main/src/output_manager.py
-Tested default koboldcpp api behavior with streaming and non-streaming generate endpoints and running GUI and seems to be fine.
-Still TODO / evaluate before merging:
(1) implement rest of openai chat/completion parameters to the extent possible, mapping to koboldcpp parameters
(2) determine if there is a way to use kobold's prompt formats for certain models when translating openai messages format into a prompt string. (Not sure if possible or where these are in the code)
(3) have chat/completions responses include the actual local model the user is using instead of just koboldcpp (Not sure if this is possible)
Note I am a python noob, so if there is a more elegant way of doing this at minimum hopefully I have done some of the grunt work for you to implement on your own.
---
koboldcpp.py | 84 +++++++++++++++++++++++++++++++++++++++++++++-------
1 file changed, 73 insertions(+), 11 deletions(-)
diff --git a/koboldcpp.py b/koboldcpp.py
index 7a733973b..96c04acd1 100755
--- a/koboldcpp.py
+++ b/koboldcpp.py
@@ -371,6 +371,7 @@ exitcounter = 0
totalgens = 0
currentusergenkey = "" #store a special key so polled streaming works even in multiuser
args = None #global args
+openaistreaming = False #store if using openai endpoint in streaming mode
class ServerRequestHandler(http.server.SimpleHTTPRequestHandler):
sys_version = ""
@@ -402,6 +403,26 @@ class ServerRequestHandler(http.server.SimpleHTTPRequestHandler):
scaled_rep_pen = genparams.get('presence_penalty', frqp) + 1
genparams["max_length"] = genparams.get('max_tokens', 50)
genparams["rep_pen"] = scaled_rep_pen
+ elif api_format==4:
+ # TODO: translate other openai unique chat completion parameters to kobold parameters
+ # translate openai chat completion messages format into one big string.
+ messages_array = genparams.get('messages', [])
+ messages_string = ""
+ for message in messages_array:
+ # TODO: use the template format of the model rather than generic Instruction/Response
+ if message['role'] == "system":
+ messages_string+="\n###Instruction: "
+ elif message['role'] == "user":
+ messages_string+="\n###Instruction: "
+ elif message['role'] == "assistant":
+ messages_string+="\n###Response: "
+ messages_string+=message['content']
+ messages_string+="\n###Response:"
+ genparams["prompt"] = messages_string
+ frqp = genparams.get('frequency_penalty', 0.1)
+ scaled_rep_pen = genparams.get('presence_penalty', frqp) + 1
+ genparams["max_length"] = genparams.get('max_tokens', 50)
+ genparams["rep_pen"] = scaled_rep_pen
return generate(
prompt=genparams.get('prompt', ""),
@@ -443,6 +464,9 @@ class ServerRequestHandler(http.server.SimpleHTTPRequestHandler):
elif api_format==3:
res = {"id": "cmpl-1", "object": "text_completion", "created": 1, "model": "koboldcpp",
"choices": [{"text": recvtxt, "index": 0, "finish_reason": "length"}]}
+ elif api_format==4:
+ res = {"id": "cmpl-1", "object": "chat.completion", "created": 1, "model": "koboldcpp",
+ "choices": [{"index": 0, "message":{"role": "assistant", "content": recvtxt,}, "finish_reason": "length"}]}
else:
res = {"results": [{"text": recvtxt}]}
@@ -453,8 +477,12 @@ class ServerRequestHandler(http.server.SimpleHTTPRequestHandler):
async def send_sse_event(self, event, data):
- self.wfile.write(f'event: {event}\n'.encode())
- self.wfile.write(f'data: {data}\n\n'.encode())
+ if openaistreaming == False:
+ self.wfile.write(f'event: {event}\n'.encode())
+ self.wfile.write(f'data: {data}\n\n'.encode())
+ else:
+ openai_sse_data = 'data: ' + data + '\r\n\r\n'
+ self.wfile.write(openai_sse_data.encode())
async def handle_sse_stream(self):
@@ -486,9 +514,22 @@ class ServerRequestHandler(http.server.SimpleHTTPRequestHandler):
if tokenStr!="":
event_data = {"token": tokenStr}
- event_str = json.dumps(event_data)
+ event_str = ""
+ # if openaistreaming endpoint, set format to expected openai streaming response
+ if openaistreaming == True:
+ event_data = {"id":"koboldcpp","object":"chat.completion.chunk","created":1,"model":"koboldcpp","choices":[{"index":0,"finish_reason":"length","delta":{'role':'assistant','content':tokenStr},}],}
+ event_str = json.dumps(event_data)
+ else:
+ event_str = json.dumps(event_data)
tokenStr = ""
+ print('\r\nSending event_str to SSE:')
+ print('\r\n' + event_str)
await self.send_sse_event("message", event_str)
+ if streamDone:
+ # if openai streaming chat endpoint, send last [DONE] message consistent with openai format
+ if openaistreaming == True:
+ print('\r\nevent stream done')
+ await self.send_sse_event("message", '[DONE]')
else:
await asyncio.sleep(0.02) #this should keep things responsive
@@ -525,6 +566,17 @@ class ServerRequestHandler(http.server.SimpleHTTPRequestHandler):
force_json = False
if self.path in ["", "/?"] or self.path.startswith(('/?','?')): #it's possible for the root url to have ?params without /
+ if args.stream and not "streaming=1" in self.path:
+ self.path = self.path.replace("streaming=0","")
+ if self.path.startswith(('/?','?')):
+ self.path += "&streaming=1"
+ else:
+ self.path = self.path + "?streaming=1"
+ self.send_response(302)
+ self.send_header("Location", self.path)
+ self.end_headers()
+ print("Force redirect to streaming mode, as --stream is set.")
+ return None
if self.embedded_kailite is None:
response_body = (f"Embedded Kobold Lite is not found.
You will have to connect via the main KoboldAI client, or use this URL to connect.").encode()
@@ -590,12 +642,11 @@ class ServerRequestHandler(http.server.SimpleHTTPRequestHandler):
return
def do_POST(self):
- global modelbusy, requestsinqueue, currentusergenkey, totalgens
+ global modelbusy, requestsinqueue, currentusergenkey, totalgens, openaistreaming
content_length = int(self.headers['Content-Length'])
body = self.rfile.read(content_length)
self.path = self.path.rstrip('/')
force_json = False
-
if self.path.endswith(('/api/extra/tokencount')):
try:
genparams = json.loads(body)
@@ -659,8 +710,9 @@ class ServerRequestHandler(http.server.SimpleHTTPRequestHandler):
try:
kai_sse_stream_flag = False
+ openaistreaming = False
- api_format = 0 #1=basic,2=kai,3=oai
+ api_format = 0 #1=basic,2=kai,3=oai,4=oai-chat
if self.path.endswith('/request'):
api_format = 1
@@ -672,10 +724,14 @@ class ServerRequestHandler(http.server.SimpleHTTPRequestHandler):
api_format = 2
kai_sse_stream_flag = True
- if self.path.endswith('/v1/completions') or self.path.endswith('/completions'):
+ if self.path.endswith('/v1/completions'):
api_format = 3
force_json = True
+ if self.path.endswith('/v1/chat/completions'):
+ api_format = 4
+ force_json = True
+
if api_format>0:
genparams = None
try:
@@ -689,7 +745,14 @@ class ServerRequestHandler(http.server.SimpleHTTPRequestHandler):
if args.foreground:
bring_terminal_to_foreground()
-
+ # Check if streaming chat completions, if so, set stream mode to true
+ if api_format == 4 and "stream" in genparams:
+ print(genparams["stream"])
+ if genparams["stream"] == True:
+ kai_sse_stream_flag = True
+ openaistreaming = True
+ print(openaistreaming)
+ print(kai_sse_stream_flag)
gen = asyncio.run(self.handle_request(genparams, api_format, kai_sse_stream_flag))
try:
@@ -721,8 +784,8 @@ class ServerRequestHandler(http.server.SimpleHTTPRequestHandler):
self.send_header('Access-Control-Allow-Origin', '*')
self.send_header('Access-Control-Allow-Methods', '*')
self.send_header('Access-Control-Allow-Headers', '*')
- if "/api" in self.path or force_json:
- if self.path.endswith("/stream"):
+ if "/api" in self.path or force_json or openaistreaming == True:
+ if self.path.endswith("/stream") or openaistreaming == True:
self.send_header('Content-type', 'text/event-stream')
self.send_header('Content-type', 'application/json')
else:
@@ -1098,7 +1161,6 @@ def show_new_gui():
for idx, name, in enumerate(token_boxes):
makecheckbox(tokens_tab, name, token_boxes[name], idx + 1)
-
# context size
makeslider(tokens_tab, "Context Size:",contextsize_text, context_var, 0, len(contextsize_text)-1, 20, set=2)