diff --git a/examples/server/routes_sdapi.cpp b/examples/server/routes_sdapi.cpp index eaf31db7c..9515faf67 100644 --- a/examples/server/routes_sdapi.cpp +++ b/examples/server/routes_sdapi.cpp @@ -301,6 +301,74 @@ static nlohmann::json prepare_info_field(const SDContextParams& ctx_params, return jsoninfo; } +namespace { + struct SdapiPreviewState { + std::mutex mutex; + std::string b64; + std::string id_task; + std::string job_type = "txt2img"; + int step = 0; + int total = 0; + bool active = false; + int job_count = 1; + int job_no = 0; + }; + + SdapiPreviewState g_preview; + +} + +static void preview_callback(int step, int frame_count, sd_image_t* frames, bool is_noisy, void* data) { + (void)is_noisy; + (void)data; + if (frame_count <= 0 || frames == nullptr || frames[0].data == nullptr) { + return; + } + + const sd_image_t& frame = frames[0]; + + auto image_bytes = encode_image_to_vector(EncodedImageFormat::JPEG, + frame.data, + frame.width, + frame.height, + frame.channel, + "", + 80); + if (image_bytes.empty()) { + return; + } + + std::lock_guard lock(g_preview.mutex); + g_preview.b64 = base64_encode(image_bytes); + + if (step < g_preview.step) { + if (g_preview.job_no < g_preview.job_count - 1) { + g_preview.job_no++; + } + } + g_preview.step = step; +} + +static void set_preview_callback(const std::string& id_task, int steps, int count, const std::string& job_type) { + std::lock_guard lock(g_preview.mutex); + g_preview.id_task = id_task; + g_preview.total = steps; + g_preview.job_type = job_type; + g_preview.b64.clear(); + g_preview.step = 0; + g_preview.active = true; + g_preview.job_count = count > 0 ? count : 1; + g_preview.job_no = 0; + sd_set_preview_callback(preview_callback, PREVIEW_PROJ, 1, true, false, nullptr); +} + +static void clear_preview_callback() { + sd_set_preview_callback(nullptr, PREVIEW_NONE, 1, false, false, nullptr); + std::lock_guard lock(g_preview.mutex); + g_preview.step = g_preview.total; + g_preview.active = false; +} + void register_sdapi_endpoints(httplib::Server& svr, ServerRuntime& rt) { ServerRuntime* runtime = &rt; @@ -317,7 +385,9 @@ void register_sdapi_endpoints(httplib::Server& svr, ServerRuntime& rt) { return; } - json j = json::parse(req.body); + std::string id_task = "sdapi"; + json j = json::parse(req.body); + id_task = j.value("id_task", id_task); ImgGenJobRequest request; std::string error_message; if (!build_sdapi_img_gen_request(j, *runtime, img2img, request, error_message)) { @@ -334,11 +404,16 @@ void register_sdapi_endpoints(httplib::Server& svr, ServerRuntime& rt) { { std::lock_guard lock(*runtime->sd_ctx_mutex); + set_preview_callback(id_task, + img_gen_params.sample_params.sample_steps, + img_gen_params.batch_count, + img2img ? "img2img" : "txt2img"); 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; } + clear_preview_callback(); results.adopt(raw_results, num_results); } @@ -528,4 +603,60 @@ void register_sdapi_endpoints(httplib::Server& svr, ServerRuntime& rt) { r["sd_model_checkpoint"] = model_path.stem(); res.set_content(r.dump(), "application/json"); }); + + svr.Get("/sdapi/v1/progress", [runtime](const httplib::Request& req, httplib::Response& res) { + bool skip_current_image = false; + bool has_id_param = req.has_param("id_task"); + std::string req_id_task = has_id_param ? req.get_param_value("id_task") : ""; + + if (req.has_param("skip_current_image")) { + std::string skip_str = req.get_param_value("skip_current_image"); + skip_current_image = (skip_str == "true" || skip_str == "1"); + } + + json r; + r["eta_relative"] = 0.0f; + json state; + state["skipped"] = false; + state["interrupted"] = false; + state["job_timestamp"] = "0"; + r["textinfo"] = nullptr; + + { + std::lock_guard lock(g_preview.mutex); + + bool is_matching_task = !has_id_param || (req_id_task == g_preview.id_task); + + int step = is_matching_task ? g_preview.step : 0; + int steps = is_matching_task ? g_preview.total : 0; + std::string b64 = is_matching_task ? g_preview.b64 : ""; + std::string job_type = is_matching_task ? g_preview.job_type : ""; + int job_count = is_matching_task ? g_preview.job_count : 1; + int job_no = is_matching_task ? g_preview.job_no : 0; + + float batch_progress = (steps > 0) ? (static_cast(step) / static_cast(steps)) : 0.0f; + if (batch_progress > 1.0f) + batch_progress = 1.0f; + + float overall_progress = (static_cast(job_no) + batch_progress) / static_cast(job_count); + if (overall_progress > 1.0f) + overall_progress = 1.0f; + + r["progress"] = overall_progress; + state["job"] = job_type; + state["sampling_step"] = step; + state["sampling_steps"] = steps; + state["job_count"] = job_count; + state["job_no"] = job_no; + r["state"] = state; + + if (!skip_current_image && !b64.empty()) { + r["current_image"] = b64; + } else { + r["current_image"] = nullptr; + } + } + + res.set_content(r.dump(), "application/json"); + }); }