569 lines
16 KiB
C++
569 lines
16 KiB
C++
#include "llmclient.hpp"
|
|
|
|
#include "message.hpp"
|
|
#include "messagemodel.hpp"
|
|
#include "session.hpp"
|
|
|
|
#include <QDebug>
|
|
#include <QJsonArray>
|
|
#include <QJsonDocument>
|
|
#include <QJsonObject>
|
|
#include <QNetworkReply>
|
|
#include <QNetworkRequest>
|
|
#include <QSet>
|
|
#include <QUrl>
|
|
|
|
namespace ZShell::llm {
|
|
|
|
QString LlmClient::completionsPath(
|
|
const QString& endpoint, const QString& subpath) {
|
|
QString base = endpoint.trimmed();
|
|
while (base.endsWith('/'))
|
|
base.chop(1);
|
|
if (!base.endsWith("/v1"))
|
|
base += "/v1";
|
|
return base + subpath;
|
|
}
|
|
|
|
QString LlmClient::serverErrorMessage(
|
|
const QByteArray& body, const QString& fallback) {
|
|
const QJsonDocument doc = QJsonDocument::fromJson(body);
|
|
if (doc.isObject()) {
|
|
const QJsonObject obj = doc.object();
|
|
if (obj.contains("error")) {
|
|
const QJsonValue errorValue = obj["error"];
|
|
if (errorValue.isObject()) {
|
|
const QString message =
|
|
errorValue.toObject()["message"].toString();
|
|
if (!message.isEmpty())
|
|
return message;
|
|
} else if (!errorValue.toString().isEmpty()) {
|
|
return errorValue.toString();
|
|
}
|
|
}
|
|
}
|
|
return fallback.isEmpty()
|
|
? QStringLiteral("Request to LLM server failed")
|
|
: fallback;
|
|
}
|
|
|
|
void LlmClient::setBusy(bool value) {
|
|
if (m_busy == value)
|
|
return;
|
|
m_busy = value;
|
|
Q_EMIT busyChanged();
|
|
}
|
|
|
|
void LlmClient::setStreamingChatId(const QString& id) {
|
|
if (m_streamingChatId == id)
|
|
return;
|
|
m_streamingChatId = id;
|
|
Q_EMIT streamingChatIdChanged();
|
|
}
|
|
|
|
LlmClient::LlmClient(QObject* parent) : QObject(parent) {
|
|
probeContextSize();
|
|
}
|
|
|
|
LlmClient::~LlmClient() {
|
|
if (m_reply)
|
|
m_reply->abort();
|
|
endStream();
|
|
}
|
|
|
|
void LlmClient::setEndpoint(const QString& value) {
|
|
if (m_endpoint == value)
|
|
return;
|
|
m_endpoint = value;
|
|
Q_EMIT endpointChanged();
|
|
probeContextSize();
|
|
if (m_model.isEmpty())
|
|
refreshModels();
|
|
}
|
|
|
|
void LlmClient::setModel(const QString& value) {
|
|
if (m_model == value)
|
|
return;
|
|
m_model = value;
|
|
Q_EMIT modelChanged();
|
|
if (m_model.isEmpty())
|
|
refreshModels();
|
|
}
|
|
|
|
void LlmClient::setTemperature(double value) {
|
|
m_temperature = value;
|
|
}
|
|
|
|
void LlmClient::startGeneration(
|
|
ChatSession* session, ChatGeneration* target) {
|
|
if (m_busy || !session || !target)
|
|
return;
|
|
m_active = session;
|
|
m_streaming = target;
|
|
m_streaming->setStreaming(true);
|
|
setBusy(true);
|
|
setStreamingChatId(session->id());
|
|
|
|
const QUrl url =
|
|
QUrl::fromUserInput(completionsPath(m_endpoint, "/chat/completions"));
|
|
if (!url.isValid() || url.host().isEmpty()) {
|
|
fail(QStringLiteral("Invalid LLM endpoint: %1").arg(m_endpoint));
|
|
return;
|
|
}
|
|
|
|
QNetworkRequest request(url);
|
|
request.setHeader(QNetworkRequest::ContentTypeHeader, "application/json");
|
|
request.setRawHeader("Accept", "text/event-stream");
|
|
|
|
const auto* model = session->messagesModel();
|
|
const int targetRow =
|
|
model->rowOf(qobject_cast<ChatMessage*>(target->parent()));
|
|
if (targetRow < 0) {
|
|
fail(QStringLiteral("Internal error: generation target is not in the "
|
|
"session"));
|
|
return;
|
|
}
|
|
|
|
// Context: every message older than the target, oldest first.
|
|
QJsonArray messages;
|
|
for (int row = model->rowCount() - 1; row > targetRow; --row) {
|
|
const auto* message = model->at(row);
|
|
const auto* generation = message->activeGeneration();
|
|
if (!generation)
|
|
continue;
|
|
QJsonObject messageObj;
|
|
messageObj[QStringLiteral("role")] =
|
|
message->role() == ChatMessage::Role::User
|
|
? QStringLiteral("user")
|
|
: QStringLiteral("assistant");
|
|
messageObj[QStringLiteral("content")] = generation->content();
|
|
if (!generation->reasoning().isEmpty())
|
|
messageObj[QStringLiteral("reasoning_content")] =
|
|
generation->reasoning();
|
|
messages.append(messageObj);
|
|
}
|
|
|
|
QJsonObject body;
|
|
QJsonObject streamOptions;
|
|
streamOptions[QStringLiteral("include_usage")] = true;
|
|
body[QStringLiteral("stream_options")] = streamOptions;
|
|
body[QStringLiteral("messages")] = messages;
|
|
body[QStringLiteral("stream")] = true;
|
|
body[QStringLiteral("temperature")] = m_temperature;
|
|
if (!m_model.isEmpty())
|
|
body[QStringLiteral("model")] = m_model;
|
|
|
|
m_buffer.clear();
|
|
m_reply = m_manager.post(request, QJsonDocument(body).toJson());
|
|
|
|
connect(m_reply, &QNetworkReply::readyRead, this, [this]() {
|
|
if (m_reply)
|
|
m_buffer.append(m_reply->readAll());
|
|
drainBuffer();
|
|
});
|
|
connect(m_reply, &QNetworkReply::finished, this, [this]() {
|
|
QNetworkReply* reply = m_reply;
|
|
if (!reply)
|
|
return;
|
|
m_reply = nullptr;
|
|
|
|
const QNetworkReply::NetworkError error = reply->error();
|
|
const QString errorString = reply->errorString();
|
|
// Data not already consumed by readyRead is only reachable here.
|
|
const QByteArray responseBody = reply->readAll();
|
|
m_buffer.append(responseBody);
|
|
reply->deleteLater();
|
|
|
|
drainBuffer();
|
|
if (!m_streaming)
|
|
return;
|
|
|
|
if (error == QNetworkReply::NoError ||
|
|
error == QNetworkReply::OperationCanceledError)
|
|
finalize();
|
|
else
|
|
fail(serverErrorMessage(responseBody, errorString));
|
|
});
|
|
}
|
|
|
|
void LlmClient::shortRequest(
|
|
const QString& tag,
|
|
const QString& systemPrompt,
|
|
const QString& userText,
|
|
std::function<void(QString result)> onResult) {
|
|
const QUrl url =
|
|
QUrl::fromUserInput(completionsPath(m_endpoint, "/chat/completions"));
|
|
if (!url.isValid() || url.host().isEmpty())
|
|
return;
|
|
|
|
qInfo() << "LlmClient:" << tag << "request POST" << url.toString()
|
|
<< "model=" << m_model;
|
|
|
|
QNetworkRequest request(url);
|
|
request.setHeader(QNetworkRequest::ContentTypeHeader, "application/json");
|
|
request.setRawHeader("Accept", "application/json");
|
|
|
|
QJsonArray messages;
|
|
QJsonObject system;
|
|
system[QStringLiteral("role")] = QStringLiteral("system");
|
|
system[QStringLiteral("content")] = systemPrompt;
|
|
messages.append(system);
|
|
QJsonObject user;
|
|
user[QStringLiteral("role")] = QStringLiteral("user");
|
|
user[QStringLiteral("content")] = userText.simplified().mid(0, 512);
|
|
messages.append(user);
|
|
|
|
QJsonObject body;
|
|
if (!m_model.isEmpty())
|
|
body[QStringLiteral("model")] = m_model;
|
|
body[QStringLiteral("stream")] = false;
|
|
body[QStringLiteral("temperature")] = 0.3;
|
|
body[QStringLiteral("max_tokens")] = 128;
|
|
QJsonObject templateKwargs;
|
|
templateKwargs[QStringLiteral("enable_thinking")] = false;
|
|
body[QStringLiteral("chat_template_kwargs")] = templateKwargs;
|
|
body[QStringLiteral("messages")] = messages;
|
|
|
|
auto* reply = m_manager.post(request, QJsonDocument(body).toJson());
|
|
connect(
|
|
reply,
|
|
&QNetworkReply::finished,
|
|
this,
|
|
[this, reply, tag, onResult = std::move(onResult)]() {
|
|
const QByteArray data = reply->readAll();
|
|
reply->deleteLater();
|
|
|
|
qInfo() << "LlmClient:" << tag << "request finished"
|
|
<< "error=" << reply->error() << reply->errorString()
|
|
<< "http=" << reply->attribute(
|
|
QNetworkRequest::HttpStatusCodeAttribute)
|
|
.toInt()
|
|
<< "response="
|
|
<< QString::fromUtf8(data.left(400)).simplified();
|
|
|
|
if (reply->error() != QNetworkReply::NoError)
|
|
return;
|
|
|
|
const QJsonDocument doc = QJsonDocument::fromJson(data);
|
|
const QJsonArray choices =
|
|
doc.object()[QStringLiteral("choices")].toArray();
|
|
if (choices.isEmpty())
|
|
return;
|
|
const QString result =
|
|
choices.at(0).toObject()[QStringLiteral("message")].toObject()
|
|
[QStringLiteral("content")].toString()
|
|
.trimmed();
|
|
qInfo() << "LlmClient:" << tag << "raw result" << result;
|
|
onResult(result);
|
|
});
|
|
}
|
|
|
|
void LlmClient::requestTitle(ChatSession* session, const QString& userText) {
|
|
shortRequest(
|
|
QStringLiteral("title"),
|
|
QStringLiteral(
|
|
"Write a short, concise title for a chat conversation starting "
|
|
"with the user's message below. At most six words, no quotation "
|
|
"marks, no trailing punctuation. Reply with the title only."),
|
|
userText,
|
|
[this, session = QPointer<ChatSession>(session)](QString title) {
|
|
if (!session) {
|
|
qWarning() << "LlmClient: title request: session gone";
|
|
return;
|
|
}
|
|
const auto isQuote = [](QChar c) {
|
|
return c == QLatin1Char('"') || c == QLatin1Char('\'') ||
|
|
c == QChar(u'\u201C') || c == QChar(u'\u201D') ||
|
|
c == QChar(u'\u2018') || c == QChar(u'\u2019');
|
|
};
|
|
while (title.size() >= 2 && isQuote(title.at(0)) &&
|
|
isQuote(title.at(title.size() - 1)))
|
|
title = title.mid(1, title.size() - 2).simplified();
|
|
while (!title.isEmpty() &&
|
|
(title.endsWith(QLatin1Char('.')) ||
|
|
title.endsWith(QLatin1Char('!')) ||
|
|
title.endsWith(QLatin1Char('?'))))
|
|
title.chop(1);
|
|
if (title.size() < 2) {
|
|
qWarning() << "LlmClient: title rejected (too short)"
|
|
<< title;
|
|
return;
|
|
}
|
|
if (title.size() > 48)
|
|
title = title.left(47) + QStringLiteral("…");
|
|
qInfo() << "LlmClient: suggesting title" << title;
|
|
Q_EMIT titleSuggested(session, title);
|
|
});
|
|
}
|
|
|
|
void LlmClient::requestIcon(ChatSession* session, const QString& userText) {
|
|
const QStringList icons = {
|
|
"chat", "lightbulb", "code",
|
|
"description", "article", "school",
|
|
"work", "build", "science",
|
|
"palette", "music_note", "sports_esports",
|
|
"takeout_dining", "flight", "photo_camera",
|
|
"psychology_alt", "favorite", "savings",
|
|
"gamepad", "auto_awesome",
|
|
};
|
|
const QString prompt =
|
|
QStringLiteral(
|
|
"Pick the single icon name from this list that best matches the "
|
|
"topic of the user's message below: %1. Reply with only the icon "
|
|
"name, exactly as written in the list, and nothing else.")
|
|
.arg(icons.join(QStringLiteral(", ")));
|
|
shortRequest(
|
|
QStringLiteral("icon"),
|
|
prompt,
|
|
userText,
|
|
[this, session = QPointer<ChatSession>(session), icons](
|
|
QString name) {
|
|
if (!session) {
|
|
qWarning() << "LlmClient: icon request: session gone";
|
|
return;
|
|
}
|
|
name = name.simplified().toLower();
|
|
const auto isQuote = [](QChar c) {
|
|
return c == QLatin1Char('"') || c == QLatin1Char('\'');
|
|
};
|
|
while (name.size() >= 2 && isQuote(name.at(0)) &&
|
|
isQuote(name.at(name.size() - 1)))
|
|
name = name.mid(1, name.size() - 2).simplified();
|
|
name.replace(QLatin1Char(' '), QLatin1Char('_'));
|
|
if (!icons.contains(name)) {
|
|
qWarning() << "LlmClient: icon not in list, using default"
|
|
<< name;
|
|
name = QStringLiteral("chat");
|
|
}
|
|
qInfo() << "LlmClient: suggesting icon" << name;
|
|
Q_EMIT iconSuggested(session, name);
|
|
});
|
|
}
|
|
|
|
void LlmClient::stop() {
|
|
if (!m_busy)
|
|
return;
|
|
if (m_reply)
|
|
m_reply->abort();
|
|
}
|
|
|
|
void LlmClient::endStream() {
|
|
if (!m_streaming)
|
|
return;
|
|
auto* generation = m_streaming;
|
|
auto* session = m_active;
|
|
m_streaming = nullptr;
|
|
m_active = nullptr;
|
|
generation->setStreaming(false);
|
|
if (generation->content().isEmpty() &&
|
|
generation->reasoning().isEmpty()) {
|
|
if (auto* message = qobject_cast<ChatMessage*>(generation->parent())) {
|
|
if (message->generationCount() <= 1) {
|
|
if (session)
|
|
session->removeMessage(message);
|
|
} else {
|
|
message->removeGeneration(generation);
|
|
}
|
|
}
|
|
}
|
|
setBusy(false);
|
|
setStreamingChatId(QString());
|
|
}
|
|
|
|
void LlmClient::finalize() {
|
|
if (!m_streaming)
|
|
return;
|
|
ChatSession* session = m_active;
|
|
endStream();
|
|
if (session && m_pendingClear == session) {
|
|
session->clearMessages();
|
|
m_pendingClear.clear();
|
|
}
|
|
if (session)
|
|
session->persist();
|
|
}
|
|
|
|
void LlmClient::clearOnFinish(ChatSession* session) {
|
|
m_pendingClear = session;
|
|
}
|
|
|
|
void LlmClient::sessionRemoved(ChatSession* session) {
|
|
if (m_pendingClear == session)
|
|
m_pendingClear.clear();
|
|
if (m_active == session) {
|
|
stop();
|
|
endStream();
|
|
}
|
|
}
|
|
|
|
void LlmClient::fail(const QString& message) {
|
|
qWarning() << "LlmClient:" << message;
|
|
|
|
if (m_streaming) {
|
|
ChatSession* session = m_active;
|
|
endStream();
|
|
if (session && m_pendingClear == session) {
|
|
session->clearMessages();
|
|
m_pendingClear.clear();
|
|
}
|
|
if (session)
|
|
session->persist();
|
|
}
|
|
Q_EMIT errorOccurred(message);
|
|
}
|
|
|
|
void LlmClient::drainBuffer() {
|
|
while (true) {
|
|
const qsizetype newline = m_buffer.indexOf('\n');
|
|
if (newline < 0)
|
|
break;
|
|
const QByteArray line = m_buffer.left(newline).trimmed();
|
|
m_buffer.remove(0, newline + 1);
|
|
handleLine(line);
|
|
}
|
|
}
|
|
|
|
void LlmClient::handleLine(const QByteArray& line) {
|
|
if (!m_streaming || line.isEmpty() || !line.startsWith("data:"))
|
|
return;
|
|
|
|
const QByteArray data = line.mid(5).trimmed();
|
|
if (data == "[DONE]") {
|
|
finalize();
|
|
return;
|
|
}
|
|
|
|
const QJsonDocument doc = QJsonDocument::fromJson(data);
|
|
if (!doc.isObject())
|
|
return;
|
|
const QJsonObject obj = doc.object();
|
|
updateTokenUsage(obj);
|
|
|
|
if (obj.contains("error")) {
|
|
const QJsonObject error = obj["error"].toObject();
|
|
const QString message = error["message"].toString();
|
|
fail(message.isEmpty()
|
|
? QStringLiteral("LLM server returned an error")
|
|
: message);
|
|
return;
|
|
}
|
|
|
|
for (const QJsonValue& choiceValue : obj["choices"].toArray()) {
|
|
if (!m_streaming)
|
|
continue;
|
|
const QJsonObject delta =
|
|
choiceValue.toObject()["delta"].toObject();
|
|
m_streaming->appendContent(delta["content"].toString());
|
|
QString reasoning = delta["reasoning_content"].toString();
|
|
if (reasoning.isEmpty())
|
|
reasoning = delta["reasoning"].toString();
|
|
m_streaming->appendReasoning(reasoning);
|
|
}
|
|
}
|
|
|
|
void LlmClient::refreshModels() {
|
|
const QUrl url =
|
|
QUrl::fromUserInput(completionsPath(m_endpoint, "/models"));
|
|
if (!url.isValid() || url.host().isEmpty())
|
|
return;
|
|
|
|
auto* reply = m_manager.get(QNetworkRequest(url));
|
|
connect(reply, &QNetworkReply::finished, this, [this, reply]() {
|
|
const QNetworkReply::NetworkError error = reply->error();
|
|
const QByteArray data = reply->readAll();
|
|
reply->deleteLater();
|
|
|
|
if (error != QNetworkReply::NoError) {
|
|
qWarning() << "LlmClient: failed to fetch models:" << error;
|
|
return;
|
|
}
|
|
const QJsonArray dataArr =
|
|
QJsonDocument::fromJson(data).object()["data"].toArray();
|
|
|
|
QStringList models;
|
|
QSet<QString> seen;
|
|
for (const QJsonValue& value : dataArr) {
|
|
const QString id = value.toObject()["id"].toString();
|
|
if (!id.isEmpty() && !seen.contains(id)) {
|
|
seen.insert(id);
|
|
models.append(id);
|
|
}
|
|
}
|
|
if (models.isEmpty())
|
|
return;
|
|
|
|
m_availableModels = models;
|
|
Q_EMIT availableModelsChanged();
|
|
|
|
if (m_model.isEmpty()) {
|
|
m_model = models.first();
|
|
Q_EMIT modelChanged();
|
|
}
|
|
});
|
|
}
|
|
|
|
void LlmClient::setContextSize(int size) {
|
|
if (size <= 0 || m_contextSize == size)
|
|
return;
|
|
m_contextSize = size;
|
|
Q_EMIT contextSizeChanged();
|
|
}
|
|
|
|
void LlmClient::probeContextSize() {
|
|
// llama.cpp-specific endpoint; other servers fall back to 4096.
|
|
QString base = m_endpoint.trimmed();
|
|
while (base.endsWith('/'))
|
|
base.chop(1);
|
|
const QUrl url = QUrl::fromUserInput(base + "/props");
|
|
if (!url.isValid() || url.host().isEmpty())
|
|
return;
|
|
|
|
auto* reply = m_manager.get(QNetworkRequest(url));
|
|
connect(
|
|
reply,
|
|
&QNetworkReply::finished,
|
|
this,
|
|
[this, reply]() {
|
|
const QNetworkReply::NetworkError error = reply->error();
|
|
const QByteArray data = reply->readAll();
|
|
reply->deleteLater();
|
|
|
|
int size = 0;
|
|
if (error == QNetworkReply::NoError) {
|
|
const QJsonDocument doc = QJsonDocument::fromJson(data);
|
|
if (doc.isArray()) {
|
|
for (const auto& value : doc.array()) {
|
|
const QJsonObject slot = value.toObject();
|
|
if (slot.contains("n_ctx")) {
|
|
size = slot["n_ctx"].toInt(0);
|
|
if (size > 0)
|
|
break;
|
|
}
|
|
}
|
|
} else if (doc.isObject()) {
|
|
const QJsonObject obj = doc.object();
|
|
size = obj["n_ctx"].toInt(0);
|
|
if (size <= 0)
|
|
size = obj["default_generation_settings"].toObject()
|
|
["n_ctx"].toInt(0);
|
|
}
|
|
}
|
|
setContextSize(size > 0 ? size : 4096);
|
|
});
|
|
}
|
|
|
|
void LlmClient::updateTokenUsage(const QJsonObject& data) {
|
|
if (!m_active || m_contextSize <= 0)
|
|
return;
|
|
const QJsonObject usage = data["usage"].toObject();
|
|
if (usage.isEmpty())
|
|
return;
|
|
const double used =
|
|
usage.value("prompt_tokens").toDouble() +
|
|
usage.value("completion_tokens").toDouble();
|
|
if (used > 0)
|
|
m_active->setLastTokenCount(static_cast<int>(used));
|
|
}
|
|
|
|
} // namespace ZShell::llm
|