diff --git a/examples/server/routes_openai.cpp b/examples/server/routes_openai.cpp index af0cdfaa1..c20be3ba1 100644 --- a/examples/server/routes_openai.cpp +++ b/examples/server/routes_openai.cpp @@ -23,6 +23,7 @@ static std::string extract_and_remove_sd_cpp_extra_args(std::string& text) { static bool build_openai_generation_request(const httplib::Request& req, ServerRuntime& runtime, ImgGenJobRequest& request, + bool& stream, std::string& error_message) { if (req.body.empty()) { error_message = "empty body"; @@ -30,6 +31,7 @@ static bool build_openai_generation_request(const httplib::Request& req, } json j = json::parse(req.body); + stream = j.value("stream", false); std::string prompt = j.value("prompt", ""); int n = std::max(1, j.value("n", 1)); std::string size = j.value("size", ""); @@ -80,12 +82,21 @@ static bool build_openai_generation_request(const httplib::Request& req, static bool build_openai_edit_request(const httplib::Request& req, ServerRuntime& runtime, ImgGenJobRequest& request, + bool& stream, std::string& error_message) { if (!req.is_multipart_form_data()) { error_message = "Content-Type must be multipart/form-data"; return false; } + // Parse stream boolean fro m form field + stream = false; + if (req.form.has_field("stream")) { + std::string stream_val = req.form.get_field("stream"); + std::transform(stream_val.begin(), stream_val.end(), stream_val.begin(), ::tolower); + stream = (stream_val == "true" || stream_val == "1"); + } + std::string prompt = req.form.get_field("prompt"); if (prompt.empty()) { error_message = "prompt required"; @@ -247,6 +258,123 @@ static bool execute_sync_img_gen_request(ServerRuntime& runtime, return true; } +struct PreviewStreamContext { + httplib::DataSink* sink = nullptr; + const ImgGenJobRequest* request = nullptr; +}; + +static void preview_callback(int step, int frame_count, sd_image_t* frames, bool is_noisy, void* data) { + if (!data || frame_count <= 0 || !frames || !frames[0].data) { + return; + } + + auto* ctx = static_cast(data); + if (!ctx->sink || !ctx->sink->is_writable()) { + return; // Client disconnected + } + + auto image_bytes = encode_image_to_vector( + EncodedImageFormat::JPEG, + frames[0].data, + frames[0].width, + frames[0].height, + frames[0].channel, + "", + 80); + + if (image_bytes.empty()) { + return; + } + + json preview_chunk; + preview_chunk["object"] = "image.preview"; + preview_chunk["step"] = step; + preview_chunk["data"] = json::array({{{"b64_json", base64_encode(image_bytes)}}}); + + std::string sse_data = "data: " + preview_chunk.dump() + "\n\n"; + ctx->sink->write(sse_data.data(), sse_data.size()); +} + +static void execute_streaming_img_gen(ServerRuntime& runtime, + ImgGenJobRequest& request, + httplib::DataSink& sink) { + PreviewStreamContext preview_ctx{&sink, &request}; + sd_img_gen_params_t img_gen_params = request.to_sd_img_gen_params_t(); + SDImageVec results; + int num_results = 0; + + { + std::lock_guard lock(*runtime.sd_ctx_mutex); + + // TODO: support other preview methods? + sd_set_preview_callback(preview_callback, PREVIEW_PROJ, 1, true, false, &preview_ctx); + + try { + sd_image_t* raw_results = nullptr; + if (!generate_image(runtime.sd_ctx, &img_gen_params, &raw_results, &num_results)) { + raw_results = nullptr; + num_results = 0; + } + results.adopt(raw_results, num_results); + } catch (...) { + sd_set_preview_callback(nullptr, PREVIEW_NONE, 0, false, false, nullptr); + throw; + } + + sd_set_preview_callback(nullptr, PREVIEW_NONE, 0, false, false, nullptr); + } + + if (results.empty() || !sink.is_writable()) { + if (sink.is_writable()) { + json err = {{"error", "generate_image returned no results"}}; + std::string err_msg = "data: " + err.dump() + "\n\n"; + sink.write(err_msg.data(), err_msg.size()); + } + sink.done(); + return; + } + + json final_out; + final_out["created"] = static_cast(std::time(nullptr)); + final_out["data"] = json::array(); + final_out["output_format"] = request.output_format; + + int result_count = results.count(); + int images_per_batch = request.gen_params.batch_count > 0 + ? std::max(1, result_count / request.gen_params.batch_count) + : 1; + + for (int i = 0; i < result_count; ++i) { + if (results[i].data == nullptr) + continue; + + std::string params = request.gen_params.embed_image_metadata + ? get_image_params(*runtime.ctx_params, request.gen_params, + request.gen_params.seed + i / images_per_batch) + : ""; + auto image_bytes = encode_image_to_vector( + request.output_format == "jpeg" ? EncodedImageFormat::JPEG + : request.output_format == "webp" ? EncodedImageFormat::WEBP + : EncodedImageFormat::PNG, + results[i].data, results[i].width, results[i].height, + results[i].channel, params, request.output_compression); + + if (!image_bytes.empty()) { + final_out["data"].push_back({{"b64_json", base64_encode(image_bytes)}}); + } + } + + if (sink.is_writable()) { + std::string final_sse = "data: " + final_out.dump() + "\n\n"; + sink.write(final_sse.data(), final_sse.size()); + + const std::string done_msg = "data: [DONE]\n\n"; + sink.write(done_msg.data(), done_msg.size()); + } + + sink.done(); +} + void register_openai_api_endpoints(httplib::Server& svr, ServerRuntime& rt) { ServerRuntime* runtime = &rt; @@ -267,7 +395,8 @@ void register_openai_api_endpoints(httplib::Server& svr, ServerRuntime& rt) { ImgGenJobRequest request; std::string error_message; - if (!build_openai_generation_request(req, *runtime, request, error_message)) { + bool stream = false; + if (!build_openai_generation_request(req, *runtime, request, stream, error_message)) { res.status = 400; res.set_content(json({{"error", error_message}}).dump(), "application/json"); return; @@ -275,6 +404,20 @@ void register_openai_api_endpoints(httplib::Server& svr, ServerRuntime& rt) { LOG_VERBOSE("%s\n", request.gen_params.to_string().c_str()); + if (stream) { + res.set_header("Cache-Control", "no-cache"); + res.set_header("Connection", "keep-alive"); + res.set_chunked_content_provider( + "text/event-stream", + [runtime, request](size_t /*offset*/, httplib::DataSink& sink) mutable { + execute_streaming_img_gen(*runtime, request, sink); + return false; + } + ); + return; + } + + SDImageVec results; if (!execute_sync_img_gen_request(*runtime, request, results, error_message)) { res.status = 500; @@ -341,7 +484,8 @@ void register_openai_api_endpoints(httplib::Server& svr, ServerRuntime& rt) { ImgGenJobRequest request; std::string error_message; - if (!build_openai_edit_request(req, *runtime, request, error_message)) { + bool stream = false; + if (!build_openai_edit_request(req, *runtime, request, stream, error_message)) { res.status = 400; res.set_content(json({{"error", error_message}}).dump(), "application/json"); return; @@ -349,6 +493,22 @@ void register_openai_api_endpoints(httplib::Server& svr, ServerRuntime& rt) { LOG_VERBOSE("%s\n", request.gen_params.to_string().c_str()); + if (stream) { + res.set_header("Cache-Control", "no-cache"); + res.set_header("Connection", "keep-alive"); + + auto req_ptr = std::make_shared(std::move(request)); + res.set_chunked_content_provider( + "text/event-stream", + [runtime, req_ptr](size_t /*offset*/, httplib::DataSink& sink) { + execute_streaming_img_gen(*runtime, *req_ptr, sink); + return false; + } + ); + return; + } + + // Synchronous path SDImageVec results; if (!execute_sync_img_gen_request(*runtime, request, results, error_message)) { res.status = 500;