Files
z-bar-qt/Plugins/ZShell/Llm/chat.cpp
T

495 lines
13 KiB
C++

#include "chat.hpp"
#include "config.hpp"
#include "llm.hpp"
#include <QDateTime>
#include <QDebug>
#include <QJsonArray>
#include <QJsonDocument>
#include <QJsonObject>
#include <QNetworkReply>
#include <QNetworkRequest>
#include <QSet>
#include <QUrl>
namespace ZShell {
QString Chat::titleFrom(const QString& content) {
const QString flat = content.simplified();
if (flat.isEmpty())
return QString();
if (flat.size() <= 48)
return flat;
return flat.left(47) + QStringLiteral("…");
}
QString Chat::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 Chat::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 Chat::setBusy(Chat* chat, bool value) {
if (chat->m_busy == value)
return;
chat->m_busy = value;
Q_EMIT chat->busyChanged();
}
void Chat::setStreamingChatId(Chat* chat, const QString& id) {
if (chat->m_streamingChatId == id)
return;
chat->m_streamingChatId = id;
Q_EMIT chat->streamingChatIdChanged();
}
Chat::Chat(QObject* parent) : QObject(parent), m_store(new ChatStore(this)) {
if (!config::Config::instance())
new config::Config();
const auto* llm = config::Config::instance()->llm();
m_endpoint = llm->endpoint();
m_temperature = llm->temperature();
m_model = llm->model();
connect(llm, &config::Llm::endpointChanged, this, [this, llm]() {
if (m_endpoint != llm->endpoint()) {
m_endpoint = llm->endpoint();
Q_EMIT endpointChanged();
probeContextSize();
}
if (m_model.isEmpty())
refreshModels();
});
connect(llm, &config::Llm::modelChanged, this, [this, llm]() {
if (m_model == llm->model())
return;
m_model = llm->model();
Q_EMIT modelChanged();
if (m_model.isEmpty())
refreshModels();
});
connect(
llm,
&config::Llm::temperatureChanged,
this,
[this, llm]() { m_temperature = llm->temperature(); });
if (m_model.isEmpty())
refreshModels();
probeContextSize();
}
Chat::~Chat() {
if (m_reply)
m_reply->abort();
}
Chat* Chat::s_instance = nullptr;
Chat* Chat::create(QQmlEngine*, QJSEngine*) {
if (!s_instance)
s_instance = new Chat();
return s_instance;
}
void Chat::send(const QString& chatId, const QString& content) {
const QString text = content.trimmed();
if (text.isEmpty() || m_busy)
return;
ChatSession* session = m_store->sessionById(chatId);
if (!session) {
qWarning() << "Chat: unknown chat id" << chatId;
return;
}
if (!m_lastError.isEmpty()) {
m_lastError.clear();
Q_EMIT lastErrorChanged();
}
session->ensureLoaded();
session->appendMessage(
ChatMessage::Role::User,
text,
QDateTime::currentMSecsSinceEpoch());
while (session->messageCount() > 200)
session->removeMessage(session->messages().first());
if (session->title().isEmpty())
session->setTitle(titleFrom(text));
if (m_contextSize > 0 && m_lastTokenCount > m_contextSize * 4 / 5)
trimHistory(session, m_contextSize);
m_active = session;
beginAssistant();
m_store->persist(session);
}
void Chat::beginAssistant() {
m_streaming = m_active->appendMessage(
ChatMessage::Role::Assistant,
QString(),
QDateTime::currentMSecsSinceEpoch());
m_streaming->setStreaming(true);
setBusy(this, true);
setStreamingChatId(this, m_active->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");
QJsonArray messages;
for (const auto* message : m_active->messages()) {
if (message == m_streaming)
continue;
QJsonObject messageObj;
messageObj[QStringLiteral("role")] =
message->role() == ChatMessage::Role::User
? QStringLiteral("user")
: QStringLiteral("assistant");
messageObj[QStringLiteral("content")] = message->content();
if (!message->reasoning().isEmpty())
messageObj[QStringLiteral("reasoning_content")] = message->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();
const QByteArray responseBody = reply->readAll();
reply->deleteLater();
drainBuffer();
if (!m_streaming)
return;
if (error == QNetworkReply::NoError ||
error == QNetworkReply::OperationCanceledError)
finalize();
else
fail(serverErrorMessage(responseBody, errorString));
});
}
void Chat::stop() {
if (!m_busy)
return;
if (m_reply)
m_reply->abort();
}
void Chat::dismissError() {
if (m_lastError.isEmpty())
return;
m_lastError.clear();
Q_EMIT lastErrorChanged();
}
void Chat::clearConversation(const QString& chatId) {
ChatSession* session = m_store->sessionById(chatId);
if (!session)
return;
if (m_busy && m_active == session) {
m_pendingClearChatId = chatId;
stop();
return;
}
session->ensureLoaded();
session->clearMessages();
m_store->persist(session);
}
void Chat::finalize() {
if (!m_streaming)
return;
ChatSession* session = m_active;
endStream();
if (session && m_pendingClearChatId == session->id()) {
session->clearMessages();
m_pendingClearChatId.clear();
}
if (session)
m_store->persist(session);
}
void Chat::endStream() {
if (!m_streaming)
return;
m_streaming->setStreaming(false);
if (m_streaming->content().isEmpty() && m_streaming->reasoning().isEmpty() && m_active)
m_active->removeMessage(m_streaming);
m_streaming = nullptr;
setBusy(this, false);
setStreamingChatId(this, QString());
}
void Chat::fail(const QString& message) {
qWarning() << "Chat:" << message;
const bool overflow = message.contains("overflow") ||
message.contains("exceed") ||
m_lastTokenCount > m_contextSize;
const QString shown = overflow
? QStringLiteral(
"%1 (using ~%2 of %3 tokens; the oldest messages were auto-removed "
"so the next one will fit)")
.arg(message, QString::number(m_lastTokenCount),
QString::number(m_contextSize))
: message;
m_lastError = shown;
Q_EMIT lastErrorChanged();
if (m_streaming) {
ChatSession* session = m_active;
endStream();
if (overflow)
trimHistory(session, m_contextSize);
if (session && m_pendingClearChatId == session->id()) {
session->clearMessages();
m_pendingClearChatId.clear();
}
if (session)
m_store->persist(session);
}
Q_EMIT errorOccurred(shown);
}
void Chat::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 Chat::handleLine(const QByteArray& line) {
if (line.isEmpty() || line.startsWith('#') || !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()) {
const QJsonObject delta = choiceValue.toObject()["delta"].toObject();
if (!m_streaming)
continue;
m_streaming->appendContent(delta["content"].toString());
QString reasoning = delta["reasoning_content"].toString();
if (reasoning.isEmpty())
reasoning = delta["reasoning"].toString();
m_streaming->appendReasoning(reasoning);
}
}
void Chat::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() << "Chat: 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 Chat::selectModel(const QString& id) {
if (id.isEmpty())
return;
m_model = id;
Q_EMIT modelChanged();
config::Config::instance()->llm()->set_model(id);
}
void Chat::setContextSize(int size) {
if (size <= 0 || m_contextSize == size)
return;
m_contextSize = size;
Q_EMIT contextSizeChanged();
}
void Chat::probeContextSize() {
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 Chat::updateTokenUsage(const QJsonObject& data) {
if (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_lastTokenCount = qMin(static_cast<int>(used), m_contextSize * 2);
}
void Chat::trimHistory(ChatSession* session, int contextSize) {
if (!session || contextSize <= 0)
return;
const int budget = contextSize * 4 / 5;
auto estimate = [&](const ChatMessage* message) -> qsizetype {
return (message->content().size() +
message->reasoning().size()) /
4;
};
qsizetype total = 0;
for (const auto* message : session->messages())
total += estimate(message);
while (total > budget && session->messageCount() >= 2) {
ChatMessage* first = session->messages().first();
const qsizetype used = estimate(first);
session->removeMessage(first);
total -= used;
if (session->messageCount() >= 2 &&
session->messages().first()->role() == ChatMessage::Role::Assistant) {
ChatMessage* second = session->messages().first();
const qsizetype usedSecond = estimate(second);
session->removeMessage(second);
total -= usedSecond;
}
}
}
} // namespace ZShell