mirror of
https://github.com/LostRuins/koboldcpp.git
synced 2026-09-20 01:31:42 +02:00
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.
This commit is contained in:
+73
-11
@@ -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.<br>You will have to connect via the main KoboldAI client, or <a href='https://lite.koboldai.net?local=1&port={self.port}'>use this URL</a> 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)
|
||||
|
||||
|
||||
Reference in New Issue
Block a user