fix: defer client disconnect cleanup

This commit is contained in:
开心 2026-07-10 22:40:20 +08:00 committed by Chris Rizzitello
parent 6934dbe62b
commit 6592dab84e
11 changed files with 483 additions and 16 deletions

View file

@ -41,6 +41,10 @@ enum class EventTypes : uint32_t
*/ */
ClientDisconnected, ClientDisconnected,
/** Internal client event used to defer disconnect cleanup until the current callback has returned.
*/
ClientDisconnectRequested,
/// A stream sends this event when \c read() will return with data. /// A stream sends this event when \c read() will return with data.
StreamInputReady, StreamInputReady,

View file

@ -37,6 +37,19 @@
// Client // Client
// //
Client::DisconnectRequest::DisconnectRequest(Kind kind, const char *message)
: m_kind(kind),
m_message(message != nullptr ? message : "")
{
}
Client::DisconnectRequest::DisconnectRequest(deskflow::core::ConnectionRefusal reason, const char *message)
: m_kind(Kind::Refuse),
m_refusalReason(reason),
m_message(message != nullptr ? message : "")
{
}
Client::Client( Client::Client(
IEventQueue *events, const std::string &name, const NetworkAddress &address, ISocketFactory *socketFactory, IEventQueue *events, const std::string &name, const NetworkAddress &address, ISocketFactory *socketFactory,
deskflow::Screen *screen deskflow::Screen *screen
@ -432,6 +445,9 @@ void Client::setupConnection()
{ {
assert(m_stream != nullptr); assert(m_stream != nullptr);
m_events->addHandler(EventTypes::ClientDisconnectRequested, m_stream->getEventTarget(), [this](const auto &e) {
handleDisconnectRequested(e);
});
m_events->addHandler(EventTypes::SocketDisconnected, m_stream->getEventTarget(), [this](const auto &) { m_events->addHandler(EventTypes::SocketDisconnected, m_stream->getEventTarget(), [this](const auto &) {
handleDisconnected(); handleDisconnected();
}); });
@ -483,6 +499,7 @@ void Client::cleanupConnecting()
{ {
if (m_stream != nullptr) { if (m_stream != nullptr) {
m_events->removeHandler(EventTypes::DataSocketConnected, m_stream->getEventTarget()); m_events->removeHandler(EventTypes::DataSocketConnected, m_stream->getEventTarget());
m_events->removeHandler(EventTypes::DataSocketSecureConnected, m_stream->getEventTarget());
m_events->removeHandler(EventTypes::DataSocketConnectionFailed, m_stream->getEventTarget()); m_events->removeHandler(EventTypes::DataSocketConnectionFailed, m_stream->getEventTarget());
} }
} }
@ -496,6 +513,7 @@ void Client::cleanupConnection()
m_events->removeHandler(StreamInputShutdown, m_stream->getEventTarget()); m_events->removeHandler(StreamInputShutdown, m_stream->getEventTarget());
m_events->removeHandler(StreamOutputShutdown, m_stream->getEventTarget()); m_events->removeHandler(StreamOutputShutdown, m_stream->getEventTarget());
m_events->removeHandler(SocketDisconnected, m_stream->getEventTarget()); m_events->removeHandler(SocketDisconnected, m_stream->getEventTarget());
m_events->removeHandler(ClientDisconnectRequested, m_stream->getEventTarget());
cleanupStream(); cleanupStream();
} }
} }
@ -583,6 +601,21 @@ void Client::handleDisconnected()
sendEvent(EventTypes::ClientDisconnected); sendEvent(EventTypes::ClientDisconnected);
} }
void Client::handleDisconnectRequested(const Event &event)
{
const auto *request = static_cast<const DisconnectRequest *>(event.getDataObject());
if (request == nullptr) {
disconnect(nullptr);
return;
}
if (request->kind() == DisconnectRequest::Kind::Refuse) {
refuseConnection(request->refusalReason(), request->message());
} else {
disconnect(request->message());
}
}
void Client::handleShapeChanged() void Client::handleShapeChanged()
{ {
LOG_DEBUG("resolution changed"); LOG_DEBUG("resolution changed");

View file

@ -10,12 +10,14 @@
#include "deskflow/IClient.h" #include "deskflow/IClient.h"
#include "base/Event.h"
#include "base/EventTypes.h" #include "base/EventTypes.h"
#include "common/Enums.h" #include "common/Enums.h"
#include "deskflow/IClipboard.h" #include "deskflow/IClipboard.h"
#include "net/NetworkAddress.h" #include "net/NetworkAddress.h"
#include <climits> #include <climits>
#include <string>
class Event; class Event;
class EventQueueTimer; class EventQueueTimer;
@ -39,6 +41,39 @@ This class implements the top-level client algorithms for deskflow.
class Client : public IClient class Client : public IClient
{ {
public: public:
class DisconnectRequest : public EventData
{
public:
enum class Kind
{
Disconnect,
Refuse
};
DisconnectRequest(Kind kind, const char *message);
DisconnectRequest(deskflow::core::ConnectionRefusal reason, const char *message);
Kind kind() const
{
return m_kind;
}
deskflow::core::ConnectionRefusal refusalReason() const
{
return m_refusalReason;
}
const char *message() const
{
return m_message.empty() ? nullptr : m_message.c_str();
}
private:
Kind m_kind = Kind::Disconnect;
deskflow::core::ConnectionRefusal m_refusalReason = deskflow::core::ConnectionRefusal::ProtocolError;
std::string m_message;
};
class FailInfo class FailInfo
{ {
public: public:
@ -175,6 +210,7 @@ private:
void handleConnectTimeout(); void handleConnectTimeout();
void handleOutputError(); void handleOutputError();
void handleDisconnected(); void handleDisconnected();
void handleDisconnectRequested(const Event &event);
void handleShapeChanged(); void handleShapeChanged();
void handleClipboardGrabbed(const Event &event); void handleClipboardGrabbed(const Event &event);
void handleHello(); void handleHello();

View file

@ -86,7 +86,7 @@ void ServerProxy::handleData()
// verify we got an entire code // verify we got an entire code
if (n != 4) { if (n != 4) {
LOG_ERR("incomplete message from server: %d bytes", n); LOG_ERR("incomplete message from server: %d bytes", n);
m_client->disconnect("incomplete message from server"); requestDisconnect("incomplete message from server");
return; return;
} }
@ -112,7 +112,7 @@ void ServerProxy::handleData()
} catch (const BadClientException &e) { } catch (const BadClientException &e) {
LOG_ERR("protocol error from server: %s", e.what()); LOG_ERR("protocol error from server: %s", e.what());
ProtocolUtil::writef(m_stream, kMsgEBad); ProtocolUtil::writef(m_stream, kMsgEBad);
m_client->disconnect("invalid message from server"); requestDisconnect("invalid message from server");
return; return;
} }
@ -167,7 +167,7 @@ ServerProxy::ConnectionResult ServerProxy::parseHandshakeMessage(const uint8_t *
else if (memcmp(code, kMsgCClose, 4) == 0) { else if (memcmp(code, kMsgCClose, 4) == 0) {
// server wants us to hangup // server wants us to hangup
LOG_VERBOSE("recv close"); LOG_VERBOSE("recv close");
m_client->disconnect(nullptr); requestDisconnect(nullptr);
return Disconnect; return Disconnect;
} }
@ -176,25 +176,25 @@ ServerProxy::ConnectionResult ServerProxy::parseHandshakeMessage(const uint8_t *
int32_t minor; int32_t minor;
ProtocolUtil::readf(m_stream, kMsgEIncompatible + 4, &major, &minor); ProtocolUtil::readf(m_stream, kMsgEIncompatible + 4, &major, &minor);
LOG_ERR("server has incompatible version %d.%d", major, minor); LOG_ERR("server has incompatible version %d.%d", major, minor);
m_client->refuseConnection(IncompatibleVersion, "server has incompatible version"); requestRefuseConnection(IncompatibleVersion, "server has incompatible version");
return Disconnect; return Disconnect;
} }
else if (memcmp(code, kMsgEBusy, 4) == 0) { else if (memcmp(code, kMsgEBusy, 4) == 0) {
LOG_ERR("server already has a connected client with name \"%s\"", m_client->getName().c_str()); LOG_ERR("server already has a connected client with name \"%s\"", m_client->getName().c_str());
m_client->refuseConnection(AlreadyConnected, "server already has a connected client with our name"); requestRefuseConnection(AlreadyConnected, "server already has a connected client with our name");
return Disconnect; return Disconnect;
} }
else if (memcmp(code, kMsgEUnknown, 4) == 0) { else if (memcmp(code, kMsgEUnknown, 4) == 0) {
LOG_ERR("server refused client with name \"%s\"", m_client->getName().c_str()); LOG_ERR("server refused client with name \"%s\"", m_client->getName().c_str());
m_client->refuseConnection(UnknownClient, "server refused client with our name"); requestRefuseConnection(UnknownClient, "server refused client with our name");
return Disconnect; return Disconnect;
} }
else if (memcmp(code, kMsgEBad, 4) == 0) { else if (memcmp(code, kMsgEBad, 4) == 0) {
LOG_ERR("server disconnected due to a protocol error"); LOG_ERR("server disconnected due to a protocol error");
m_client->refuseConnection(ProtocolError, "server reported a protocol error"); requestRefuseConnection(ProtocolError, "server reported a protocol error");
return Disconnect; return Disconnect;
} else if (memcmp(code, kMsgDLanguageSynchronisation, 4) == 0) { } else if (memcmp(code, kMsgDLanguageSynchronisation, 4) == 0) {
setServerLanguages(); setServerLanguages();
@ -312,11 +312,11 @@ ServerProxy::ConnectionResult ServerProxy::parseMessage(const uint8_t *code)
else if (memcmp(code, kMsgCClose, 4) == 0) { else if (memcmp(code, kMsgCClose, 4) == 0) {
// server wants us to hangup // server wants us to hangup
LOG_VERBOSE("recv close"); LOG_VERBOSE("recv close");
m_client->disconnect(nullptr); requestDisconnect(nullptr);
return Disconnect; return Disconnect;
} else if (memcmp(code, kMsgEBad, 4) == 0) { } else if (memcmp(code, kMsgEBad, 4) == 0) {
LOG_ERR("server disconnected due to a protocol error"); LOG_ERR("server disconnected due to a protocol error");
m_client->disconnect("server reported a protocol error"); requestDisconnect("server reported a protocol error");
return Disconnect; return Disconnect;
} else { } else {
return Unknown; return Unknown;
@ -337,7 +337,22 @@ ServerProxy::ConnectionResult ServerProxy::parseMessage(const uint8_t *code)
void ServerProxy::handleKeepAliveAlarm() void ServerProxy::handleKeepAliveAlarm()
{ {
LOG_INFO("server is dead"); LOG_INFO("server is dead");
m_client->disconnect("server is not responding"); requestDisconnect("server is not responding");
}
void ServerProxy::requestDisconnect(const char *message)
{
m_events->addEvent(Event(
EventTypes::ClientDisconnectRequested, m_stream->getEventTarget(),
new Client::DisconnectRequest(Client::DisconnectRequest::Kind::Disconnect, message)
));
}
void ServerProxy::requestRefuseConnection(deskflow::core::ConnectionRefusal reason, const char *message)
{
m_events->addEvent(Event(
EventTypes::ClientDisconnectRequested, m_stream->getEventTarget(), new Client::DisconnectRequest(reason, message)
));
} }
void ServerProxy::onInfoChanged() void ServerProxy::onInfoChanged()
@ -550,7 +565,7 @@ void ServerProxy::setClipboard()
LOG_INFO("clipboard was updated"); LOG_INFO("clipboard was updated");
} else if (r == TransferState::Error) { } else if (r == TransferState::Error) {
m_client->disconnect("invalid clipboard data from server"); requestDisconnect("invalid clipboard data from server");
} }
} }

View file

@ -8,6 +8,7 @@
#pragma once #pragma once
#include "common/Enums.h"
#include "deskflow/ClipboardChunk.h" #include "deskflow/ClipboardChunk.h"
#include "deskflow/ClipboardTypes.h" #include "deskflow/ClipboardTypes.h"
#include "deskflow/KeyTypes.h" #include "deskflow/KeyTypes.h"
@ -77,6 +78,8 @@ private:
// event handlers // event handlers
void handleData(); void handleData();
void handleKeepAliveAlarm(); void handleKeepAliveAlarm();
void requestDisconnect(const char *message);
void requestRefuseConnection(deskflow::core::ConnectionRefusal reason, const char *message);
// message handlers // message handlers
void enter(); void enter();

View file

@ -262,10 +262,11 @@ void ClientApp::closeClient(Client *client)
return; return;
} }
using enum EventTypes; using enum EventTypes;
getEvents()->removeHandler(ClientConnected, client); auto *target = client->getEventTarget();
getEvents()->removeHandler(ClientConnectionFailed, client); getEvents()->removeHandler(ClientConnected, target);
getEvents()->removeHandler(ClientConnectionRefused, client); getEvents()->removeHandler(ClientConnectionFailed, target);
getEvents()->removeHandler(ClientDisconnected, client); getEvents()->removeHandler(ClientConnectionRefused, target);
getEvents()->removeHandler(ClientDisconnected, target);
delete client; delete client;
} }

View file

@ -57,6 +57,7 @@ enable_testing()
find_package(Qt6 ${REQUIRED_QT_VERSION} REQUIRED COMPONENTS Test) find_package(Qt6 ${REQUIRED_QT_VERSION} REQUIRED COMPONENTS Test)
add_subdirectory(base) add_subdirectory(base)
add_subdirectory(client)
add_subdirectory(common) add_subdirectory(common)
add_subdirectory(deskflow) add_subdirectory(deskflow)
add_subdirectory(gui) add_subdirectory(gui)

View file

@ -40,7 +40,7 @@ create_test(
create_test( create_test(
NAME EventQueueTests NAME EventQueueTests
DEPENDS base DEPENDS base
LIBS arch mt LIBS arch mt ${extra_libs}
SOURCE EventQueueTests.cpp SOURCE EventQueueTests.cpp
WORKING_DIRECTORY "${CMAKE_BINARY_DIR}/src/lib/base" WORKING_DIRECTORY "${CMAKE_BINARY_DIR}/src/lib/base"
) )

View file

@ -0,0 +1,14 @@
# SPDX-FileCopyrightText: (C) 2026 Deskflow Developers
# SPDX-License-Identifier: MIT
if(WIN32)
set(extra_libs version)
endif()
create_test(
NAME ServerProxyTests
DEPENDS client
LIBS arch base io mt net platform server app ${extra_libs}
SOURCE ServerProxyTests.cpp
WORKING_DIRECTORY "${CMAKE_BINARY_DIR}/src/lib/client"
)

View file

@ -0,0 +1,335 @@
/*
* Deskflow -- mouse and keyboard sharing utility
* SPDX-FileCopyrightText: (C) 2026 Deskflow Developers
* SPDX-License-Identifier: GPL-2.0-only WITH LicenseRef-OpenSSL-Exception
*/
#include "ServerProxyTests.h"
#include "base/Event.h"
#include "base/IEventQueue.h"
#include "client/Client.h"
#include "client/ServerProxy.h"
#include "deskflow/AppUtil.h"
#include "deskflow/ProtocolTypes.h"
#include "io/IStream.h"
#include <QTest>
#include <algorithm>
#include <cstdint>
#include <cstring>
#include <deque>
#include <functional>
#include <map>
#include <string>
#include <vector>
namespace {
class TestAppUtil : public AppUtil
{
public:
int run() override
{
return 0;
}
void startNode() override
{
}
std::vector<std::string> getKeyboardLayoutList() override
{
return {"en"};
}
std::string getCurrentLanguageCode() override
{
return "en";
}
};
class FakeStream : public deskflow::IStream
{
public:
void push(const std::string &bytes)
{
m_chunks.push_back(bytes);
}
void close() override
{
m_chunks.clear();
m_inputShutdown = true;
}
uint32_t read(void *buffer, uint32_t size) override
{
if (m_inputShutdown || m_chunks.empty() || size == 0) {
return 0;
}
auto &front = m_chunks.front();
const size_t bytesToRead = std::min(static_cast<size_t>(size), front.size());
if (buffer != nullptr) {
std::memcpy(buffer, front.data(), bytesToRead);
}
front.erase(0, bytesToRead);
if (front.empty()) {
m_chunks.pop_front();
}
return static_cast<uint32_t>(bytesToRead);
}
void write(const void *, uint32_t) override
{
}
void flush() override
{
}
void shutdownInput() override
{
close();
}
void shutdownOutput() override
{
}
void *getEventTarget() const override
{
return const_cast<FakeStream *>(this);
}
bool isReady() const override
{
return !m_inputShutdown && !m_chunks.empty();
}
uint32_t getSize() const override
{
size_t total = 0;
for (const auto &chunk : m_chunks) {
total += chunk.size();
}
return static_cast<uint32_t>(std::min<size_t>(total, UINT32_MAX));
}
private:
std::deque<std::string> m_chunks;
bool m_inputShutdown = false;
};
class RecordingEventQueue : public IEventQueue
{
public:
~RecordingEventQueue() override
{
for (const auto &event : m_addedEvents) {
Event::deleteData(event);
}
}
int loop() override
{
return 0;
}
void adoptBuffer(IEventQueueBuffer *) override
{
}
bool getEvent(Event &, double = -1.0) override
{
return false;
}
bool dispatchEvent(const Event &event) override
{
const auto handler = m_handlers.find(HandlerKey{event.getType(), event.getTarget()});
if (handler == m_handlers.end()) {
return false;
}
const auto callback = handler->second;
callback(event);
return true;
}
void addEvent(Event &&event) override
{
m_addedEvents.emplace_back(std::move(event));
}
EventQueueTimer *newTimer(double, void *) override
{
return timer();
}
EventQueueTimer *newOneShotTimer(double, void *) override
{
return timer();
}
void deleteTimer(EventQueueTimer *) override
{
}
void addHandler(EventTypes type, void *target, const EventHandler &handler) override
{
m_handlers[HandlerKey{type, target}] = handler;
}
void removeHandler(EventTypes type, void *target) override
{
m_handlers.erase(HandlerKey{type, target});
}
void removeHandlers(void *target) override
{
for (auto it = m_handlers.begin(); it != m_handlers.end();) {
if (it->first.target == target) {
it = m_handlers.erase(it);
} else {
++it;
}
}
}
void waitForReady() const override
{
}
void *getSystemTarget() override
{
return this;
}
EventQueueTimer *timer()
{
return reinterpret_cast<EventQueueTimer *>(&m_timerStorage);
}
const std::vector<Event> &addedEvents() const
{
return m_addedEvents;
}
private:
struct HandlerKey
{
EventTypes type;
void *target;
bool operator<(const HandlerKey &other) const
{
if (type != other.type) {
return static_cast<uint32_t>(type) < static_cast<uint32_t>(other.type);
}
return std::less<void *>{}(target, other.target);
}
};
int m_timerStorage = 0;
std::map<HandlerKey, EventHandler> m_handlers;
std::vector<Event> m_addedEvents;
};
class TestServerProxy : public ServerProxy
{
public:
using ServerProxy::ServerProxy;
bool parseHandshakeMessageReturnsDisconnect(const uint8_t *code)
{
return parseHandshakeMessage(code) == ConnectionResult::Disconnect;
}
};
Client *undereferenceableClient()
{
// These paths must queue cleanup without calling through to Client.
return reinterpret_cast<Client *>(0x1);
}
TestAppUtil &testAppUtil()
{
static TestAppUtil util;
return util;
}
const Client::DisconnectRequest *disconnectRequest(const RecordingEventQueue &events)
{
if (events.addedEvents().size() != 1) {
return nullptr;
}
return static_cast<const Client::DisconnectRequest *>(events.addedEvents().front().getDataObject());
}
} // namespace
void ServerProxyTests::initTestCase()
{
(void)testAppUtil();
m_log.setFilter(LogLevel::Level::Debug);
}
void ServerProxyTests::handleKeepAliveAlarm_timeout_queuesDisconnectRequest()
{
RecordingEventQueue events;
FakeStream stream;
ServerProxy proxy(undereferenceableClient(), &stream, &events);
QVERIFY(events.dispatchEvent(Event(EventTypes::Timer, events.timer())));
QCOMPARE(events.addedEvents().size(), static_cast<size_t>(1));
const auto &event = events.addedEvents().front();
QVERIFY(event.getType() == EventTypes::ClientDisconnectRequested);
QCOMPARE(event.getTarget(), stream.getEventTarget());
const auto *request = disconnectRequest(events);
QVERIFY(request != nullptr);
QVERIFY(request->kind() == Client::DisconnectRequest::Kind::Disconnect);
QCOMPARE(QString::fromUtf8(request->message()), QStringLiteral("server is not responding"));
}
void ServerProxyTests::handleData_incompleteMessage_queuesDisconnectRequest()
{
RecordingEventQueue events;
FakeStream stream;
stream.push("DF");
ServerProxy proxy(undereferenceableClient(), &stream, &events);
QVERIFY(events.dispatchEvent(Event(EventTypes::StreamInputReady, stream.getEventTarget())));
QCOMPARE(events.addedEvents().size(), static_cast<size_t>(1));
const auto &event = events.addedEvents().front();
QVERIFY(event.getType() == EventTypes::ClientDisconnectRequested);
QCOMPARE(event.getTarget(), stream.getEventTarget());
const auto *request = disconnectRequest(events);
QVERIFY(request != nullptr);
QVERIFY(request->kind() == Client::DisconnectRequest::Kind::Disconnect);
QCOMPARE(QString::fromUtf8(request->message()), QStringLiteral("incomplete message from server"));
}
void ServerProxyTests::parseHandshakeMessage_protocolError_queuesRefusalRequest()
{
RecordingEventQueue events;
FakeStream stream;
TestServerProxy proxy(undereferenceableClient(), &stream, &events);
QVERIFY(proxy.parseHandshakeMessageReturnsDisconnect(reinterpret_cast<const uint8_t *>(kMsgEBad)));
QCOMPARE(events.addedEvents().size(), static_cast<size_t>(1));
const auto *request = disconnectRequest(events);
QVERIFY(request != nullptr);
QVERIFY(request->kind() == Client::DisconnectRequest::Kind::Refuse);
QVERIFY(request->refusalReason() == deskflow::core::ConnectionRefusal::ProtocolError);
QCOMPARE(QString::fromUtf8(request->message()), QStringLiteral("server reported a protocol error"));
}
QTEST_MAIN(ServerProxyTests)

View file

@ -0,0 +1,25 @@
/*
* Deskflow -- mouse and keyboard sharing utility
* SPDX-FileCopyrightText: (C) 2026 Deskflow Developers
* SPDX-License-Identifier: GPL-2.0-only WITH LicenseRef-OpenSSL-Exception
*/
#pragma once
#include "base/Log.h"
#include <QObject>
class ServerProxyTests : public QObject
{
Q_OBJECT
private Q_SLOTS:
void initTestCase();
void handleKeepAliveAlarm_timeout_queuesDisconnectRequest();
void handleData_incompleteMessage_queuesDisconnectRequest();
void parseHandshakeMessage_protocolError_queuesRefusalRequest();
private:
Log m_log;
};