48 [[nodiscard]] std::string
curlError(CURLcode code) {
return curl_easy_strerror(code); }
51 [[nodiscard]] std::string windowsError(DWORD error) {
52 LPWSTR message =
nullptr;
53 const DWORD
length = FormatMessageW(FORMAT_MESSAGE_ALLOCATE_BUFFER | FORMAT_MESSAGE_FROM_SYSTEM | FORMAT_MESSAGE_IGNORE_INSERTS,
nullptr, error, 0,
reinterpret_cast<LPWSTR
>(&message), 0,
nullptr);
55 if (
length != 0U && message !=
nullptr) {
56 const int utf8_length = WideCharToMultiByte(CP_UTF8, 0, message,
static_cast<int>(
length),
nullptr, 0,
nullptr,
nullptr);
57 if (utf8_length > 0) {
58 text.resize(
static_cast<std::size_t
>(utf8_length));
59 WideCharToMultiByte(CP_UTF8, 0, message,
static_cast<int>(
length), text.data(), utf8_length,
nullptr,
nullptr);
62 if (message !=
nullptr) {
65 while (!text.empty() && (text.back() ==
'\r' || text.back() ==
'\n')) {
68 return text.empty() ?
"Windows error " + std::to_string(error) : text;
71 [[nodiscard]] std::wstring utf8ToWide(std::string_view text) {
75 const int length = MultiByteToWideChar(CP_UTF8, MB_ERR_INVALID_CHARS, text.data(),
static_cast<int>(text.size()),
nullptr, 0);
77 throw std::runtime_error(
"unable to convert sd-server argument to UTF-16: " + windowsError(GetLastError()));
79 std::wstring output(
static_cast<std::size_t
>(
length), L
'\0');
80 if (MultiByteToWideChar(CP_UTF8, MB_ERR_INVALID_CHARS, text.data(),
static_cast<int>(text.size()), output.data(),
length) <= 0) {
81 throw std::runtime_error(
"unable to convert sd-server argument to UTF-16: " + windowsError(GetLastError()));
86 [[nodiscard]] std::wstring quoteWindowsArgument(std::wstring_view argument) {
87 if (!argument.empty() && argument.find_first_of(L
" \t\n\v\"") == std::wstring_view::npos) {
88 return std::wstring(argument);
90 std::wstring quoted(L
"\"");
91 std::size_t backslashes = 0;
92 for (
const wchar_t character : argument) {
93 if (character == L
'\\') {
97 if (character == L
'\"') {
98 quoted.append(backslashes * 2U + 1U, L
'\\');
99 quoted.push_back(L
'\"');
103 quoted.append(backslashes, L
'\\');
105 quoted.push_back(character);
107 quoted.append(backslashes * 2U, L
'\\');
108 quoted.push_back(L
'\"');
113 std::size_t
appendResponse(
char *data, std::size_t size, std::size_t count,
void *context) {
115 if (size != 0U && count > std::numeric_limits<std::size_t>::max() / size) {
116 buffer->overflow =
true;
119 const std::size_t bytes = size * count;
121 buffer->overflow =
true;
124 buffer->
value.append(data, bytes);
128 int checkCancelled(
void *context, curl_off_t, curl_off_t, curl_off_t, curl_off_t) {
129 const auto *cancelled =
static_cast<const std::function<
bool()
> *>(context);
130 return cancelled !=
nullptr && *cancelled && (*cancelled)() ? 1 : 0;
133 [[nodiscard]] std::string
encodeBase64(
const std::vector<std::uint8_t> &bytes) {
134 static constexpr std::string_view ALPHABET =
"ABCDEFGHIJKLMNOPQRSTUVWXYZabcdefghijklmnopqrstuvwxyz0123456789+/";
136 output.reserve(((bytes.size() + 2U) / 3U) * 4U);
137 for (std::size_t index = 0; index < bytes.size(); index += 3U) {
138 const std::uint32_t first = bytes[index];
139 const std::uint32_t second = index + 1U < bytes.size() ? bytes[index + 1U] : 0U;
140 const std::uint32_t third = index + 2U < bytes.size() ? bytes[index + 2U] : 0U;
141 const std::uint32_t value = (first << 16U) | (second << 8U) | third;
142 output.push_back(ALPHABET[(value >> 18U) & 0x3fU]);
143 output.push_back(ALPHABET[(value >> 12U) & 0x3fU]);
144 output.push_back(index + 1U < bytes.size() ? ALPHABET[(value >> 6U) & 0x3fU] :
'=');
145 output.push_back(index + 2U < bytes.size() ? ALPHABET[value & 0x3fU] :
'=');
151 if (character >=
'A' && character <=
'Z') {
152 return character -
'A';
154 if (character >=
'a' && character <=
'z') {
155 return character -
'a' + 26;
157 if (character >=
'0' && character <=
'9') {
158 return character -
'0' + 52;
160 if (character ==
'+') {
163 if (character ==
'/') {
169 [[nodiscard]] std::vector<std::uint8_t>
decodeBase64(std::string_view encoded) {
170 const std::size_t comma = encoded.find(
',');
171 if (encoded.starts_with(
"data:") && comma != std::string_view::npos) {
172 encoded.remove_prefix(comma + 1U);
174 std::vector<std::uint8_t> output;
175 output.reserve((encoded.size() / 4U) * 3U);
176 std::uint32_t accumulator = 0U;
178 for (
const unsigned char character : encoded) {
179 if (character ==
'=') {
184 if (character ==
' ' || character ==
'\n' || character ==
'\r' || character ==
'\t') {
187 throw std::runtime_error(
"sd-server returned invalid base64 image data");
189 accumulator = (accumulator << 6U) | static_cast<std::uint32_t>(value);
193 output.push_back(
static_cast<std::uint8_t
>((accumulator >>
static_cast<unsigned int>(bits)) & 0xffU));
199 [[nodiscard]] Json::Value
parseJson(std::string_view text, std::string_view context) {
200 Json::CharReaderBuilder builder;
201 builder[
"collectComments"] =
false;
202 std::unique_ptr<Json::CharReader> reader(builder.newCharReader());
205 if (!reader->parse(text.data(), text.data() + text.size(), &value, &error)) {
206 throw std::runtime_error(std::string(context) +
" returned invalid JSON: " + error);
211 [[nodiscard]] std::string
writeJson(
const Json::Value &value) {
212 Json::StreamWriterBuilder builder;
213 builder[
"indentation"] =
"";
214 return Json::writeString(builder, value);
217 [[nodiscard]]
ResponseBuffer request(std::string_view url,
const std::string *body,
long timeout_seconds,
long &status,
const std::function<
bool()> *cancelled) {
218 CURL *handle = curl_easy_init();
219 if (handle ==
nullptr) {
220 throw std::runtime_error(
"unable to initialize libcurl");
223 curl_slist *headers =
nullptr;
224 if (body !=
nullptr) {
225 headers = curl_slist_append(headers,
"Content-Type: application/json");
227 const std::string request_url(url);
228 curl_easy_setopt(handle, CURLOPT_URL, request_url.c_str());
229 curl_easy_setopt(handle, CURLOPT_NOPROXY,
"127.0.0.1,localhost");
231 curl_easy_setopt(handle, CURLOPT_WRITEDATA, &response);
232 curl_easy_setopt(handle, CURLOPT_CONNECTTIMEOUT, 2L);
233 curl_easy_setopt(handle, CURLOPT_TIMEOUT, timeout_seconds);
234 curl_easy_setopt(handle, CURLOPT_NOSIGNAL, 1L);
235 curl_easy_setopt(handle, CURLOPT_NOPROGRESS, 0L);
236 curl_easy_setopt(handle, CURLOPT_XFERINFOFUNCTION,
checkCancelled);
237 curl_easy_setopt(handle, CURLOPT_XFERINFODATA, cancelled);
238 if (body !=
nullptr) {
239 curl_easy_setopt(handle, CURLOPT_HTTPHEADER, headers);
240 curl_easy_setopt(handle, CURLOPT_POST, 1L);
241 curl_easy_setopt(handle, CURLOPT_POSTFIELDS, body->data());
242 curl_easy_setopt(handle, CURLOPT_POSTFIELDSIZE_LARGE,
static_cast<curl_off_t
>(body->size()));
244 const CURLcode result = curl_easy_perform(handle);
245 curl_easy_getinfo(handle, CURLINFO_RESPONSE_CODE, &status);
246 curl_slist_free_all(headers);
247 curl_easy_cleanup(handle);
248 if (result != CURLE_OK) {
249 if (result == CURLE_ABORTED_BY_CALLBACK) {
250 throw std::runtime_error(
"Stable Diffusion request cancelled");
253 throw std::runtime_error(
"sd-server response exceeded 64 MiB");
255 throw std::runtime_error(
"sd-server request failed: " +
curlError(result));
261 if (root.isMember(
"message") && root[
"message"].isString()) {
262 return root[
"message"].asString();
264 if (root.isMember(
"error") && root[
"error"].isString()) {
265 return root[
"error"].asString();
267 return "unknown server error";
276 constexpr double MAX_WORKING_PIXELS = 1280.0 * 720.0;
277 const double pixels =
static_cast<double>(width) * height;
278 if (pixels > MAX_WORKING_PIXELS) {
279 const double scale = std::sqrt(MAX_WORKING_PIXELS / pixels);
280 width = std::max(64,
static_cast<int>(std::floor(width * scale)));
281 height = std::max(64,
static_cast<int>(std::floor(height * scale)));
283 return {width, height};
287 if (settings.
loras.empty()) {
290 root[
"lora"] = Json::arrayValue;
291 for (
const Lora &lora : settings.
loras) {
293 entry[
"path"] = lora.
path.generic_string();
295 entry[
"is_high_noise"] =
false;
296 root[
"lora"].append(std::move(entry));
300 [[nodiscard]] std::string
submitAsyncJob(
const std::string &endpoint, std::string_view route,
const Json::Value &request_root,
const Settings &settings) {
301 const std::string body =
writeJson(request_root);
304 const Json::Value accepted =
parseJson(submission.
value,
"sd-server");
305 if (status != 202 || !accepted[
"id"].isString()) {
306 throw std::runtime_error(
"sd-server rejected frame: " +
serverError(accepted));
308 const std::string job_id = accepted[
"id"].asString();
309 const std::string job_url = endpoint +
"/sdcpp/v1/jobs/" + job_id;
310 const auto deadline = std::chrono::steady_clock::now() + std::chrono::hours(1);
311 while (std::chrono::steady_clock::now() < deadline) {
313 const std::string cancel_body =
"{}";
314 long cancel_status = 0;
316 static_cast<void>(
request(job_url +
"/cancel", &cancel_body, 3L, cancel_status,
nullptr));
317 }
catch (
const std::exception &) {
319 throw std::runtime_error(
"Stable Diffusion request cancelled");
321 long poll_status = 0;
324 if (poll_status != 200) {
325 throw std::runtime_error(
"sd-server job polling failed: " +
serverError(job));
327 const std::string state = job[
"status"].asString();
328 if (state ==
"completed") {
329 const Json::Value &images = job[
"result"][
"images"];
330 if (!images.isArray() || images.empty() || !images[0][
"b64_json"].isString()) {
331 throw std::runtime_error(
"sd-server job did not contain an output image");
333 return images[0][
"b64_json"].asString();
335 if (state ==
"failed" || state ==
"cancelled") {
336 throw std::runtime_error(
"sd-server frame job " + state +
": " +
serverError(job[
"error"]));
338 std::this_thread::sleep_for(std::chrono::milliseconds(50));
340 throw std::runtime_error(
"sd-server frame job timed out");
367 std::vector<std::string> argument_storage{executable,
"--listen-ip",
"127.0.0.1",
"--listen-port", port};
369 argument_storage.emplace_back(
"--upscale-model");
372 argument_storage.insert(argument_storage.end(), {
"--model", settings.model.string(),
"--type",
"f16",
"--mmap",
"--fa",
"--diffusion-conv-direct",
"--vae-conv-direct",
"--lora-model-dir", settings.lora_directory.string()});
375 argument_storage.emplace_back(
"--hires-upscalers-dir");
380 std::wstring command_line;
381 for (
const std::string &argument : argument_storage) {
382 if (!command_line.empty()) {
383 command_line.push_back(L
' ');
385 command_line += quoteWindowsArgument(utf8ToWide(argument));
388 STARTUPINFOW startup_info{};
389 startup_info.cb =
sizeof(startup_info);
390 PROCESS_INFORMATION process_info{};
391 HANDLE log_handle = INVALID_HANDLE_VALUE;
392 HANDLE input_handle = INVALID_HANDLE_VALUE;
394 std::vector<wchar_t> log_path(MAX_PATH, L
'\0');
395 if (GetTempFileNameW(std::filesystem::temp_directory_path().wstring().c_str(), L
"acm", 0, log_path.data()) == 0) {
396 throw std::runtime_error(
"unable to create sd-server diagnostic log: " + windowsError(GetLastError()));
399 log_handle = CreateFileW(
diagnostic_log_path.wstring().c_str(), FILE_APPEND_DATA, FILE_SHARE_READ | FILE_SHARE_WRITE,
nullptr, OPEN_EXISTING, FILE_ATTRIBUTE_NORMAL,
nullptr);
400 if (log_handle == INVALID_HANDLE_VALUE) {
401 throw std::runtime_error(
"unable to open sd-server diagnostic log: " + windowsError(GetLastError()));
403 input_handle = CreateFileW(L
"NUL", GENERIC_READ, FILE_SHARE_READ | FILE_SHARE_WRITE,
nullptr, OPEN_EXISTING, FILE_ATTRIBUTE_NORMAL,
nullptr);
404 if (input_handle == INVALID_HANDLE_VALUE) {
405 const DWORD error = GetLastError();
406 CloseHandle(log_handle);
407 throw std::runtime_error(
"unable to configure sd-server input: " + windowsError(error));
409 if (!SetHandleInformation(log_handle, HANDLE_FLAG_INHERIT, HANDLE_FLAG_INHERIT) || !SetHandleInformation(input_handle, HANDLE_FLAG_INHERIT, HANDLE_FLAG_INHERIT)) {
410 const DWORD error = GetLastError();
411 CloseHandle(input_handle);
412 CloseHandle(log_handle);
413 throw std::runtime_error(
"unable to configure sd-server output: " + windowsError(error));
415 startup_info.dwFlags = STARTF_USESTDHANDLES;
416 startup_info.hStdOutput = log_handle;
417 startup_info.hStdError = log_handle;
418 startup_info.hStdInput = input_handle;
420 std::vector<wchar_t> mutable_command(command_line.begin(), command_line.end());
421 mutable_command.push_back(L
'\0');
422 const BOOL launched = CreateProcessW(
nullptr, mutable_command.data(),
nullptr,
nullptr,
settings.
quiet ? TRUE : FALSE, CREATE_NO_WINDOW,
nullptr,
nullptr, &startup_info, &process_info);
423 const DWORD launch_error = launched ? ERROR_SUCCESS : GetLastError();
424 if (log_handle != INVALID_HANDLE_VALUE) {
425 CloseHandle(log_handle);
427 if (input_handle != INVALID_HANDLE_VALUE) {
428 CloseHandle(input_handle);
431 throw std::runtime_error(
"unable to launch sd-server: " + windowsError(launch_error));
433 CloseHandle(process_info.hThread);
434 process_handle =
reinterpret_cast<std::intptr_t
>(process_info.hProcess);
435 process_id =
static_cast<std::int64_t
>(process_info.dwProcessId);
437 std::vector<char *> arguments;
438 arguments.reserve(argument_storage.size() + 1U);
439 for (std::string &argument : argument_storage) {
440 arguments.push_back(argument.data());
442 arguments.push_back(
nullptr);
444 posix_spawn_file_actions_t file_actions;
445 posix_spawn_file_actions_t *actions =
nullptr;
447 std::string log_template = (std::filesystem::temp_directory_path() /
"acmxvk-sd-server-XXXXXX").
string();
448 std::vector<char> mutable_template(log_template.begin(), log_template.end());
449 mutable_template.push_back(
'\0');
450 const int log_fd = ::mkstemp(mutable_template.data());
452 throw std::runtime_error(
"unable to create sd-server diagnostic log: " + std::string(std::strerror(errno)));
457 const int init_result = posix_spawn_file_actions_init(&file_actions);
458 if (init_result != 0) {
459 throw std::runtime_error(
"unable to configure sd-server output: " + std::string(std::strerror(init_result)));
461 actions = &file_actions;
462 const int stdout_result = posix_spawn_file_actions_addopen(actions, STDOUT_FILENO,
diagnostic_log_path.c_str(), O_WRONLY | O_APPEND, 0600);
463 const int stderr_result = posix_spawn_file_actions_adddup2(actions, STDOUT_FILENO, STDERR_FILENO);
464 if (stdout_result != 0 || stderr_result != 0) {
465 posix_spawn_file_actions_destroy(actions);
466 const int error = stdout_result != 0 ? stdout_result : stderr_result;
467 throw std::runtime_error(
"unable to redirect sd-server output: " + std::string(std::strerror(error)));
470 const int result = posix_spawnp(&child, executable.c_str(), actions,
nullptr, arguments.data(),
environ);
471 if (actions !=
nullptr) {
472 posix_spawn_file_actions_destroy(actions);
475 throw std::runtime_error(
"unable to launch sd-server: " + std::string(std::strerror(result)));
479 std::cout <<
"acmxvk: launched sd-server process " <<
process_id <<
" on 127.0.0.1:" <<
settings.
port <<
'\n';
555 const std::string url =
endpoint + (require_capabilities ?
"/sdcpp/v1/capabilities" :
"/sdapi/v1/options");
556 const auto deadline = std::chrono::steady_clock::now() + std::chrono::minutes(5);
557 while (std::chrono::steady_clock::now() < deadline) {
559 throw std::runtime_error(
"Stable Diffusion startup cancelled");
561 bool server_exited =
false;
563 const HANDLE
process =
reinterpret_cast<HANDLE
>(process_handle);
564 const DWORD wait_result =
process ==
nullptr ? WAIT_OBJECT_0 : WaitForSingleObject(
process, 0U);
565 if (wait_result == WAIT_FAILED) {
566 throw std::runtime_error(
"unable to inspect sd-server process: " + windowsError(GetLastError()));
568 server_exited = wait_result == WAIT_OBJECT_0;
569 if (server_exited &&
process !=
nullptr) {
574 int child_status = 0;
575 const pid_t result = waitpid(
static_cast<pid_t
>(
process_id), &child_status, WNOHANG);
576 server_exited = result ==
static_cast<pid_t
>(
process_id);
580 throw std::runtime_error(
"sd-server exited before accepting requests" +
diagnosticLogDetails());
584 const ResponseBuffer response = request(url,
nullptr, 3L, status, &
settings.
cancelled);
586 if (!response.value.empty()) {
587 const Json::Value document = parseJson(response.value,
"sd-server");
591 for (
const Json::Value &entry : document[
"upscalers"]) {
592 if (entry[
"name"].asString() == upscaler_name) {
598 throw std::runtime_error(
"sd-server did not discover the requested "
605 for (
const Json::Value &entry : document[
"loras"]) {
606 if (entry[
"path"].asString() == lora.
path.generic_string()) {
612 throw std::runtime_error(
"sd-server did not discover the requested "
614 lora.
path.generic_string());
618 std::cout <<
"acmxvk: sd-server " << (
settings.
upscale_only ?
"upscaler" :
"model") <<
" ready";
627 std::cout <<
"; " <<
settings.
loras.size() <<
" LoRA model(s)";
633 if (status >= 400 && status < 500) {
634 throw std::runtime_error(
"sd-server readiness check returned HTTP " + std::to_string(status) +
diagnosticLogDetails());
636 }
catch (
const std::exception &) {
641 std::this_thread::sleep_for(std::chrono::milliseconds(250));
643 throw std::runtime_error(
"sd-server did not become ready within five minutes" +
diagnosticLogDetails());
647 if (rgba.empty() || rgba.type() != CV_8UC4) {
648 throw std::runtime_error(
"Stable Diffusion input must be a non-empty RGBA8 frame");
651 cv::cvtColor(rgba, bgr, cv::COLOR_RGBA2BGR);
652 std::vector<std::uint8_t> png;
653 if (!cv::imencode(
".png", bgr, png)) {
654 throw std::runtime_error(
"unable to encode Stable Diffusion input frame");
659 const std::string encoded_input = encodeBase64(png);
661 std::string encoded_output;
663 const cv::Size working_size = neuralUpscaleWorkingSize(
settings, rgba.size());
664 cv::Mat working_input = rgba;
665 if (working_input.size() != working_size) {
666 cv::resize(working_input, working_input, working_size, 0.0, 0.0, cv::INTER_LANCZOS4);
669 cv::cvtColor(working_input, working_bgr, cv::COLOR_RGBA2BGR);
670 std::vector<std::uint8_t> working_png;
671 if (!cv::imencode(
".png", working_bgr, working_png)) {
672 throw std::runtime_error(
"unable to encode Stable Diffusion ESRGAN input frame");
674 root[
"image"] = encodeBase64(working_png);
676 root[
"upscale_repeats"] = 1;
677 root[
"upscale_tile_size"] = 128;
678 root[
"output_compression"] = 100;
679 encoded_output = submitAsyncJob(
endpoint,
"/sdcpp/v1/upscale", root,
settings);
689 root[
"batch_size"] = 1;
693 root[
"init_images"] = Json::arrayValue;
694 root[
"init_images"].append(encoded_input);
695 const std::string body = writeJson(root);
698 const Json::Value document = parseJson(response.value,
"sd-server");
700 throw std::runtime_error(
"sd-server rejected frame: " + serverError(document));
702 if (!document[
"images"].isArray() || document[
"images"].empty() || !document[
"images"][0].isString()) {
703 throw std::runtime_error(
"sd-server response did not contain an output image");
705 encoded_output = document[
"images"][0].asString();
714 root[
"batch_count"] = 1;
715 root[
"init_image"] = encoded_input;
716 root[
"output_format"] =
"png";
717 root[
"output_compression"] = 100;
718 Json::Value &sample = root[
"sample_params"];
723 Json::Value &hires = root[
"hires"];
724 const cv::Size working_size = neuralUpscaleWorkingSize(
settings, rgba.size());
725 hires[
"enabled"] =
true;
727 hires[
"scale"] = 4.0;
728 hires[
"target_width"] = working_size.width;
729 hires[
"target_height"] = working_size.height;
732 hires[
"upscale_tile_size"] = 128;
733 encoded_output = submitAsyncJob(
endpoint,
"/sdcpp/v1/img_gen", root,
settings);
735 const std::vector<std::uint8_t> decoded = decodeBase64(encoded_output);
736 cv::Mat encoded(1,
static_cast<int>(decoded.size()), CV_8UC1,
const_cast<std::uint8_t *
>(decoded.data()));
737 cv::Mat output_bgr = cv::imdecode(encoded, cv::IMREAD_COLOR);
738 if (output_bgr.empty()) {
739 throw std::runtime_error(
"unable to decode image returned by sd-server");
742 cv::cvtColor(output_bgr, output_rgba, cv::COLOR_BGR2RGBA);
746 cv::resize(output_rgba, output_rgba, rgba.size(), 0.0, 0.0, cv::INTER_LINEAR);