281 lines
9.8 KiB
C++
281 lines
9.8 KiB
C++
#include "net/ModelServerClient.h"
|
|
|
|
#include <QEventLoop>
|
|
#include <QJsonArray>
|
|
#include <QJsonDocument>
|
|
#include <QJsonObject>
|
|
#include <QNetworkAccessManager>
|
|
#include <QNetworkReply>
|
|
#include <QNetworkRequest>
|
|
#include <QTimer>
|
|
#include <QUrl>
|
|
|
|
#include <algorithm>
|
|
|
|
namespace core {
|
|
|
|
ModelServerClient::ModelServerClient(QObject* parent)
|
|
: QObject(parent)
|
|
, m_nam(new QNetworkAccessManager(this)) {
|
|
}
|
|
|
|
void ModelServerClient::setBaseUrl(const QUrl& baseUrl) {
|
|
m_baseUrl = baseUrl;
|
|
}
|
|
|
|
QUrl ModelServerClient::baseUrl() const {
|
|
return m_baseUrl;
|
|
}
|
|
|
|
QNetworkReply* ModelServerClient::computeDepthPng8Async(const QByteArray& imageBytes, QString* outImmediateError) {
|
|
if (outImmediateError) {
|
|
outImmediateError->clear();
|
|
}
|
|
if (!m_baseUrl.isValid() || m_baseUrl.isEmpty()) {
|
|
if (outImmediateError) *outImmediateError = QStringLiteral("后端地址无效。");
|
|
return nullptr;
|
|
}
|
|
if (imageBytes.isEmpty()) {
|
|
if (outImmediateError) *outImmediateError = QStringLiteral("输入图像为空。");
|
|
return nullptr;
|
|
}
|
|
|
|
const QUrl url = m_baseUrl.resolved(QUrl(QStringLiteral("/depth")));
|
|
QNetworkRequest req(url);
|
|
req.setHeader(QNetworkRequest::ContentTypeHeader, QStringLiteral("application/json"));
|
|
|
|
const QByteArray imageB64 = imageBytes.toBase64();
|
|
const QJsonObject payload{
|
|
{QStringLiteral("image_b64"), QString::fromLatin1(imageB64)},
|
|
};
|
|
const QByteArray body = QJsonDocument(payload).toJson(QJsonDocument::Compact);
|
|
return m_nam->post(req, body);
|
|
}
|
|
|
|
QNetworkReply* ModelServerClient::segmentSamPromptAsync(
|
|
const QByteArray& cropRgbPngBytes,
|
|
const QByteArray& overlayPngBytes,
|
|
const QJsonArray& pointCoords,
|
|
const QJsonArray& pointLabels,
|
|
const QJsonArray& boxXyxy,
|
|
QString* outImmediateError
|
|
) {
|
|
if (outImmediateError) {
|
|
outImmediateError->clear();
|
|
}
|
|
if (!m_baseUrl.isValid() || m_baseUrl.isEmpty()) {
|
|
if (outImmediateError) *outImmediateError = QStringLiteral("后端地址无效。");
|
|
return nullptr;
|
|
}
|
|
if (cropRgbPngBytes.isEmpty()) {
|
|
if (outImmediateError) *outImmediateError = QStringLiteral("裁剪图像为空。");
|
|
return nullptr;
|
|
}
|
|
|
|
const QUrl url = m_baseUrl.resolved(QUrl(QStringLiteral("/segment/sam_prompt")));
|
|
QNetworkRequest req(url);
|
|
req.setHeader(QNetworkRequest::ContentTypeHeader, QStringLiteral("application/json"));
|
|
|
|
QJsonObject payload;
|
|
payload.insert(QStringLiteral("image_b64"), QString::fromLatin1(cropRgbPngBytes.toBase64()));
|
|
if (!overlayPngBytes.isEmpty()) {
|
|
payload.insert(QStringLiteral("overlay_b64"), QString::fromLatin1(overlayPngBytes.toBase64()));
|
|
}
|
|
payload.insert(QStringLiteral("point_coords"), pointCoords);
|
|
payload.insert(QStringLiteral("point_labels"), pointLabels);
|
|
payload.insert(QStringLiteral("box_xyxy"), boxXyxy);
|
|
|
|
const QByteArray body = QJsonDocument(payload).toJson(QJsonDocument::Compact);
|
|
return m_nam->post(req, body);
|
|
}
|
|
|
|
QNetworkReply* ModelServerClient::inpaintAsync(
|
|
const QByteArray& cropRgbPngBytes,
|
|
const QByteArray& maskPngBytes,
|
|
const QString& modelName,
|
|
const QString& prompt,
|
|
const QString& negativePrompt,
|
|
double strength,
|
|
int maxSide,
|
|
QString* outImmediateError
|
|
) {
|
|
if (outImmediateError) {
|
|
outImmediateError->clear();
|
|
}
|
|
if (!m_baseUrl.isValid() || m_baseUrl.isEmpty()) {
|
|
if (outImmediateError) *outImmediateError = QStringLiteral("后端地址无效。");
|
|
return nullptr;
|
|
}
|
|
if (cropRgbPngBytes.isEmpty()) {
|
|
if (outImmediateError) *outImmediateError = QStringLiteral("裁剪图像为空。");
|
|
return nullptr;
|
|
}
|
|
if (maskPngBytes.isEmpty()) {
|
|
if (outImmediateError) *outImmediateError = QStringLiteral("Mask 为空。");
|
|
return nullptr;
|
|
}
|
|
|
|
const QUrl url = m_baseUrl.resolved(QUrl(QStringLiteral("/inpaint")));
|
|
QNetworkRequest req(url);
|
|
req.setHeader(QNetworkRequest::ContentTypeHeader, QStringLiteral("application/json"));
|
|
|
|
QJsonObject payload;
|
|
payload.insert(QStringLiteral("image_b64"), QString::fromLatin1(cropRgbPngBytes.toBase64()));
|
|
payload.insert(QStringLiteral("mask_b64"), QString::fromLatin1(maskPngBytes.toBase64()));
|
|
if (!modelName.trimmed().isEmpty()) {
|
|
payload.insert(QStringLiteral("model_name"), modelName.trimmed());
|
|
}
|
|
payload.insert(QStringLiteral("prompt"), prompt);
|
|
payload.insert(QStringLiteral("negative_prompt"), negativePrompt);
|
|
payload.insert(QStringLiteral("strength"), strength);
|
|
payload.insert(QStringLiteral("max_side"), maxSide);
|
|
|
|
const QByteArray body = QJsonDocument(payload).toJson(QJsonDocument::Compact);
|
|
return m_nam->post(req, body);
|
|
}
|
|
|
|
QNetworkReply* ModelServerClient::animateCharacterSequenceAsync(
|
|
const QByteArray& characterPngBytes,
|
|
const QString& prompt,
|
|
const QString& negativePrompt,
|
|
const QString& backgroundColor,
|
|
int backgroundTolerance,
|
|
int frameCount,
|
|
int numInferenceSteps,
|
|
double guidanceScale,
|
|
int maxSide,
|
|
int seed,
|
|
QString* outImmediateError
|
|
) {
|
|
if (outImmediateError) {
|
|
outImmediateError->clear();
|
|
}
|
|
if (!m_baseUrl.isValid() || m_baseUrl.isEmpty()) {
|
|
if (outImmediateError) *outImmediateError = QStringLiteral("后端地址无效。");
|
|
return nullptr;
|
|
}
|
|
if (characterPngBytes.isEmpty()) {
|
|
if (outImmediateError) *outImmediateError = QStringLiteral("角色 PNG 为空。");
|
|
return nullptr;
|
|
}
|
|
if (prompt.trimmed().isEmpty()) {
|
|
if (outImmediateError) *outImmediateError = QStringLiteral("动画提示词不能为空。");
|
|
return nullptr;
|
|
}
|
|
if (frameCount <= 0) {
|
|
if (outImmediateError) *outImmediateError = QStringLiteral("帧数必须大于 0。");
|
|
return nullptr;
|
|
}
|
|
|
|
const QUrl url = m_baseUrl.resolved(QUrl(QStringLiteral("/animate/character_sequence")));
|
|
QNetworkRequest req(url);
|
|
req.setHeader(QNetworkRequest::ContentTypeHeader, QStringLiteral("application/json"));
|
|
|
|
QJsonObject payload;
|
|
payload.insert(QStringLiteral("image_b64"), QString::fromLatin1(characterPngBytes.toBase64()));
|
|
payload.insert(QStringLiteral("model_name"), QStringLiteral("animatediff"));
|
|
payload.insert(QStringLiteral("prompt"), prompt.trimmed());
|
|
payload.insert(QStringLiteral("negative_prompt"), negativePrompt);
|
|
payload.insert(QStringLiteral("background_color"), backgroundColor.trimmed().isEmpty()
|
|
? QStringLiteral("#00FF00")
|
|
: backgroundColor.trimmed());
|
|
payload.insert(QStringLiteral("background_tolerance"), std::clamp(backgroundTolerance, 0, 255));
|
|
payload.insert(QStringLiteral("video_length"), frameCount);
|
|
payload.insert(QStringLiteral("num_inference_steps"), std::max(1, numInferenceSteps));
|
|
payload.insert(QStringLiteral("guidance_scale"), guidanceScale);
|
|
payload.insert(QStringLiteral("max_side"), std::clamp(maxSide, 128, 2048));
|
|
payload.insert(QStringLiteral("seed"), seed);
|
|
|
|
const QByteArray body = QJsonDocument(payload).toJson(QJsonDocument::Compact);
|
|
return m_nam->post(req, body);
|
|
}
|
|
|
|
bool ModelServerClient::computeDepthPng8(
|
|
const QByteArray& imageBytes,
|
|
QByteArray& outPngBytes,
|
|
QString& outError,
|
|
int timeoutMs
|
|
) {
|
|
outPngBytes.clear();
|
|
outError.clear();
|
|
|
|
if (!m_baseUrl.isValid() || m_baseUrl.isEmpty()) {
|
|
outError = QStringLiteral("后端地址无效。");
|
|
return false;
|
|
}
|
|
if (imageBytes.isEmpty()) {
|
|
outError = QStringLiteral("输入图像为空。");
|
|
return false;
|
|
}
|
|
|
|
const QUrl url = m_baseUrl.resolved(QUrl(QStringLiteral("/depth")));
|
|
|
|
QNetworkRequest req(url);
|
|
req.setHeader(QNetworkRequest::ContentTypeHeader, QStringLiteral("application/json"));
|
|
|
|
const QByteArray imageB64 = imageBytes.toBase64();
|
|
const QJsonObject payload{
|
|
{QStringLiteral("image_b64"), QString::fromLatin1(imageB64)},
|
|
};
|
|
const QByteArray body = QJsonDocument(payload).toJson(QJsonDocument::Compact);
|
|
|
|
QNetworkReply* reply = m_nam->post(req, body);
|
|
if (!reply) {
|
|
outError = QStringLiteral("创建网络请求失败。");
|
|
return false;
|
|
}
|
|
|
|
QEventLoop loop;
|
|
QTimer timer;
|
|
timer.setSingleShot(true);
|
|
const int t = (timeoutMs <= 0) ? 30000 : timeoutMs;
|
|
|
|
QObject::connect(reply, &QNetworkReply::finished, &loop, &QEventLoop::quit);
|
|
QObject::connect(&timer, &QTimer::timeout, &loop, &QEventLoop::quit);
|
|
timer.start(t);
|
|
loop.exec();
|
|
|
|
if (timer.isActive() == false && reply->isFinished() == false) {
|
|
reply->abort();
|
|
reply->deleteLater();
|
|
outError = QStringLiteral("请求超时(%1ms)。").arg(t);
|
|
return false;
|
|
}
|
|
|
|
const int httpStatus = reply->attribute(QNetworkRequest::HttpStatusCodeAttribute).toInt();
|
|
const QByteArray raw = reply->readAll();
|
|
const auto netErr = reply->error();
|
|
const QString netErrStr = reply->errorString();
|
|
reply->deleteLater();
|
|
|
|
if (netErr != QNetworkReply::NoError) {
|
|
outError = QStringLiteral("网络错误:%1").arg(netErrStr);
|
|
return false;
|
|
}
|
|
|
|
if (httpStatus != 200) {
|
|
// FastAPI HTTPException 默认返回 {"detail": "..."}
|
|
QString detail;
|
|
const QJsonDocument jd = QJsonDocument::fromJson(raw);
|
|
if (jd.isObject()) {
|
|
const auto obj = jd.object();
|
|
detail = obj.value(QStringLiteral("detail")).toString();
|
|
}
|
|
outError = detail.isEmpty()
|
|
? QStringLiteral("后端返回HTTP %1。").arg(httpStatus)
|
|
: QStringLiteral("后端错误(HTTP %1):%2").arg(httpStatus).arg(detail);
|
|
return false;
|
|
}
|
|
|
|
if (raw.isEmpty()) {
|
|
outError = QStringLiteral("后端返回空数据。");
|
|
return false;
|
|
}
|
|
|
|
outPngBytes = raw;
|
|
return true;
|
|
}
|
|
|
|
} // namespace core
|
|
|