mirror of
https://github.com/LostRuins/koboldcpp.git
synced 2026-09-19 01:05:09 +02:00
sd: cancel image generation on client disconnect (#2318)
This commit is contained in:
@@ -213,6 +213,10 @@ extern "C"
|
||||
{
|
||||
return sdtype_get_info();
|
||||
}
|
||||
void sd_abort_generation()
|
||||
{
|
||||
sdtype_abort_generation();
|
||||
}
|
||||
|
||||
bool whisper_load_model(const whisper_load_model_inputs inputs)
|
||||
{
|
||||
|
||||
+27
-4
@@ -992,6 +992,8 @@ def init_library():
|
||||
handle.sd_upscale.restype = sd_generation_outputs
|
||||
handle.sd_get_info.argtypes = []
|
||||
handle.sd_get_info.restype = sd_info_outputs
|
||||
handle.sd_abort_generation.argtypes = []
|
||||
handle.sd_abort_generation.restype = None
|
||||
handle.whisper_load_model.argtypes = [whisper_load_model_inputs]
|
||||
handle.whisper_load_model.restype = ctypes.c_bool
|
||||
handle.whisper_generate.argtypes = [whisper_generation_inputs]
|
||||
@@ -5859,7 +5861,7 @@ class KcppServerRequestHandler(http.server.SimpleHTTPRequestHandler):
|
||||
self.close_connection = True
|
||||
await asyncio.sleep(0.05)
|
||||
|
||||
async def monitor_connection(self): #Poll the socket to detect client disconnection during prompt processing
|
||||
async def monitor_connection(self, cancel_fn): #Poll the socket to detect client disconnection
|
||||
import select
|
||||
loop = asyncio.get_event_loop()
|
||||
def check_connection_closed():
|
||||
@@ -5882,7 +5884,7 @@ class KcppServerRequestHandler(http.server.SimpleHTTPRequestHandler):
|
||||
if disconnected:
|
||||
if args.debugmode:
|
||||
print("\nClient disconnected unexpectedly, aborting...")
|
||||
handle.abort_generate()
|
||||
cancel_fn()
|
||||
return
|
||||
except Exception:
|
||||
return
|
||||
@@ -5897,7 +5899,7 @@ class KcppServerRequestHandler(http.server.SimpleHTTPRequestHandler):
|
||||
generate_task = asyncio.create_task(self.generate_text(genparams, api_format, stream_flag))
|
||||
tasks.append(generate_task)
|
||||
if stream_flag:
|
||||
monitor_task = asyncio.create_task(self.monitor_connection())
|
||||
monitor_task = asyncio.create_task(self.monitor_connection(handle.abort_generate))
|
||||
await asyncio.gather(*tasks)
|
||||
generate_result = generate_task.result()
|
||||
return generate_result
|
||||
@@ -5916,6 +5918,27 @@ class KcppServerRequestHandler(http.server.SimpleHTTPRequestHandler):
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
async def handle_image_request(self, generate_fn, param, cancel_fn):
|
||||
monitor_task = None
|
||||
try:
|
||||
monitor_task = asyncio.create_task(self.monitor_connection(cancel_fn))
|
||||
result = await asyncio.to_thread(generate_fn, param)
|
||||
return result
|
||||
except (BrokenPipeError, ConnectionAbortedError) as cae: # attempt to abort if connection lost
|
||||
print("An ongoing connection was aborted or interrupted!")
|
||||
print(cae)
|
||||
cancel_fn()
|
||||
await asyncio.sleep(0.1) #short delay
|
||||
except Exception as e:
|
||||
print(e)
|
||||
finally:
|
||||
if monitor_task and not monitor_task.done():
|
||||
monitor_task.cancel()
|
||||
try:
|
||||
await monitor_task
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
def get_multiplayer_idle_state(self,userid):
|
||||
if modelbusy.locked() or batched_request_runner_count>0:
|
||||
return False
|
||||
@@ -7487,7 +7510,7 @@ Change Mode<br>
|
||||
if loras:
|
||||
genparams['prompt'] = prompt
|
||||
genparams['lora'] = lora_map_name_to_path(loras)
|
||||
gen = sd_generate(genparams)
|
||||
gen = asyncio.run(self.handle_image_request(sd_generate, genparams, handle.sd_abort_generation))
|
||||
gendat = gen["data"]
|
||||
genanim = gen["animated"]
|
||||
gendatextra = gen["data_extra"]
|
||||
|
||||
@@ -112,6 +112,7 @@ bool sdtype_load_model(const sd_load_model_inputs inputs);
|
||||
sd_generation_outputs sdtype_generate(const sd_generation_inputs inputs);
|
||||
sd_generation_outputs sdtype_upscale(const sd_upscale_inputs inputs);
|
||||
sd_info_outputs sdtype_get_info();
|
||||
void sdtype_abort_generation();
|
||||
|
||||
bool whispertype_load_model(const whisper_load_model_inputs inputs);
|
||||
whisper_generation_outputs whispertype_generate(const whisper_generation_inputs inputs);
|
||||
|
||||
@@ -972,6 +972,10 @@ static std::string upscale_image_to_png_base64(upscaler_ctx_t* upscaler_ctx, con
|
||||
return gen_data;
|
||||
}
|
||||
|
||||
void sdtype_abort_generation() {
|
||||
sd_cancel_generation(sd_ctx, SD_CANCEL_ALL);
|
||||
}
|
||||
|
||||
sd_generation_outputs sdtype_generate(const sd_generation_inputs inputs)
|
||||
{
|
||||
if(sd_ctx == nullptr || sd_params == nullptr)
|
||||
|
||||
Reference in New Issue
Block a user