CmvrControlStation/src/quictestserver.cpp

949 lines
34 KiB
C++
Raw Blame History

This file contains ambiguous Unicode characters

This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.

#include "quictestserver.h"
#include <QCoreApplication>
#include <QDateTime>
#include <QDir>
#include <QFile>
#include <QLibrary>
#include <QMetaObject>
#if CMVR_WITH_MSQUIC
#if defined(__MINGW32__)
#ifndef _In_
#define _In_
#endif
#ifndef _In_opt_
#define _In_opt_
#endif
#ifndef _IRQL_requires_max_
#define _IRQL_requires_max_(level)
#endif
#ifndef _Pre_defensive_
#define _Pre_defensive_
#endif
#ifndef _Out_writes_bytes_opt_
#define _Out_writes_bytes_opt_(size)
#endif
#ifndef _In_reads_or_z_opt_
#define _In_reads_or_z_opt_(size)
#endif
#ifndef _In_reads_
#define _In_reads_(size)
#endif
#ifndef _In_reads_bytes_
#define _In_reads_bytes_(size)
#endif
#ifndef _In_reads_bytes_opt_
#define _In_reads_bytes_opt_(size)
#endif
#ifndef _In_z_
#define _In_z_
#endif
#ifndef _Inout_
#define _Inout_
#endif
#ifndef _Out_
#define _Out_
#endif
#ifndef _Out_writes_
#define _Out_writes_(size)
#endif
#ifndef _Outptr_
#define _Outptr_
#endif
#ifndef _Success_
#define _Success_(condition)
#endif
#ifndef _In_range_
#define _In_range_(minimum, maximum)
#endif
#ifndef _Field_range_
#define _Field_range_(minimum, maximum)
#endif
#ifndef _Field_size_
#define _Field_size_(size)
#endif
#ifndef _Field_size_opt_
#define _Field_size_opt_(size)
#endif
#ifndef _Field_size_bytes_
#define _Field_size_bytes_(size)
#endif
#ifndef _Field_size_bytes_opt_
#define _Field_size_bytes_opt_(size)
#endif
#ifndef _Field_z_
#define _Field_z_
#endif
#ifndef _Function_class_
#define _Function_class_(name)
#endif
#ifndef _When_
#define _When_(condition, annotation)
#endif
#ifndef _At_
#define _At_(target, annotation)
#endif
#ifndef _At_buffer_
#define _At_buffer_(target, iterator, bound, annotation)
#endif
#ifndef _Reserved_
#define _Reserved_
#endif
#ifndef _Check_return_
#define _Check_return_
#endif
#ifndef __drv_allocatesMem
#define __drv_allocatesMem(kind)
#endif
#ifndef __drv_freesMem
#define __drv_freesMem(kind)
#endif
#ifndef PASSIVE_LEVEL
#define PASSIVE_LEVEL 0
#endif
#ifndef DISPATCH_LEVEL
#define DISPATCH_LEVEL 2
#endif
#endif
#include <msquic.h>
#include "cmvr/quic_edge/v1/quic_edge.pb.h"
#include <atomic>
#include <chrono>
#include <condition_variable>
#include <cstdint>
#include <cstring>
#include <limits>
#include <mutex>
#include <set>
#include <string>
#include <vector>
namespace {
const quint32 kProtocolVersion = 1;
const quint32 kMaximumControlFrameBytes = 1024U * 1024U;
const QUIC_UINT62 kShutdownCode = 0x100U;
const QUIC_UINT62 kProtocolErrorCode = 0x101U;
#ifdef Q_OS_WIN
QString tlsAlertName(quint8 alert)
{
switch (alert) {
case 42: return QStringLiteral("bad_certificate");
case 43: return QStringLiteral("unsupported_certificate");
case 44: return QStringLiteral("certificate_revoked");
case 45: return QStringLiteral("certificate_expired");
case 46: return QStringLiteral("certificate_unknown");
case 48: return QStringLiteral("unknown_ca");
default: return QStringLiteral("unknown");
}
}
#endif
QString statusText(const char *operation, QUIC_STATUS status)
{
QString result = QStringLiteral("%1 失败,QUIC_STATUS=0x%2")
.arg(QString::fromLatin1(operation))
.arg(static_cast<qulonglong>(status), 0, 16);
#ifdef Q_OS_WIN
const quint32 rawStatus = static_cast<quint32>(status);
if ((rawStatus & 0xffffff00U) == 0x80410100U) {
const quint8 alert = static_cast<quint8>(rawStatus & 0xffU);
result += QStringLiteral(",TLS alert %1 (%2)")
.arg(alert)
.arg(tlsAlertName(alert));
}
#endif
return result;
}
QString addressText(const QUIC_ADDR *address)
{
if (!address)
return QStringLiteral("--");
QUIC_ADDR_STR value = {};
if (!QuicAddrToString(address, &value))
return QStringLiteral("--");
return QString::fromLatin1(value.Address);
}
QString addressIpText(const QUIC_ADDR *address)
{
const QString endpoint = addressText(address);
if (endpoint.startsWith(QLatin1Char('['))) {
const int closingBracket = endpoint.indexOf(QLatin1Char(']'));
if (closingBracket > 1)
return endpoint.mid(1, closingBracket - 1);
}
if (endpoint.count(QLatin1Char(':')) == 1)
return endpoint.left(endpoint.indexOf(QLatin1Char(':')));
return endpoint;
}
}
class QuicTestServer::Impl
{
public:
explicit Impl(QuicTestServer *owner)
: m_owner(owner)
{
}
~Impl()
{
stop();
}
bool isRunning() const { return m_running.load(); }
quint16 boundPort() const { return m_boundPort.load(); }
void start(const QString &bindAddress, quint16 port, const QString &alpn,
const QString &pkcs12Path, const QString &password)
{
if (m_running.load()) {
postLog(QStringLiteral("WARN"), QStringLiteral("QUIC Server 已经在运行"));
return;
}
if (bindAddress.trimmed().isEmpty() || alpn.trimmed().isEmpty()) {
postRunning(false, 0, QStringLiteral("监听地址和 ALPN 不能为空"));
return;
}
m_registrations.store(0);
m_heartbeats.store(0);
m_datagrams.store(0);
m_datagramBytes.store(0);
postStatistics();
QFile certificateFile(pkcs12Path);
if (!certificateFile.open(QIODevice::ReadOnly)) {
postRunning(false, 0, QStringLiteral("无法读取测试证书:%1").arg(pkcs12Path));
return;
}
m_pkcs12 = certificateFile.readAll();
if (m_pkcs12.isEmpty()) {
postRunning(false, 0, QStringLiteral("测试证书内容为空"));
return;
}
// QLibrary adds the platform prefix and suffix. This resolves
// msquic.dll on Windows and libmsquic.so on Linux.
m_library.setFileName(QStringLiteral("msquic"));
if (!m_library.load()) {
postRunning(false, 0, QStringLiteral("无法加载 MsQuic:%1")
.arg(m_library.errorString()));
return;
}
m_open = reinterpret_cast<MsQuicOpenVersionFn>(m_library.resolve("MsQuicOpenVersion"));
m_close = reinterpret_cast<MsQuicCloseFn>(m_library.resolve("MsQuicClose"));
if (!m_open || !m_close) {
failStart(QStringLiteral("MsQuic 动态库缺少 MsQuicOpenVersion/MsQuicClose"));
return;
}
QUIC_STATUS status = m_open(QUIC_API_VERSION_2,
reinterpret_cast<const void **>(&m_api));
if (QUIC_FAILED(status)) {
failStart(statusText("MsQuicOpenVersion", status));
return;
}
const QUIC_REGISTRATION_CONFIG registrationConfig = {
"cmvr-qt-quic-test-server", QUIC_EXECUTION_PROFILE_LOW_LATENCY
};
status = m_api->RegistrationOpen(&registrationConfig, &m_registration);
if (QUIC_FAILED(status)) {
failStart(statusText("RegistrationOpen", status));
return;
}
m_alpn = alpn.trimmed().toUtf8();
QUIC_BUFFER alpnBuffer = {};
alpnBuffer.Buffer = reinterpret_cast<uint8_t *>(m_alpn.data());
alpnBuffer.Length = static_cast<uint32_t>(m_alpn.size());
QUIC_SETTINGS settings = {};
settings.IsSet.PeerBidiStreamCount = TRUE;
settings.PeerBidiStreamCount = 1;
settings.IsSet.PeerUnidiStreamCount = TRUE;
settings.PeerUnidiStreamCount = 0;
settings.IsSet.IdleTimeoutMs = TRUE;
settings.IdleTimeoutMs = 120000;
settings.IsSet.DatagramReceiveEnabled = TRUE;
settings.DatagramReceiveEnabled = TRUE;
status = m_api->ConfigurationOpen(m_registration, &alpnBuffer, 1, &settings,
sizeof(settings), nullptr, &m_configuration);
if (QUIC_FAILED(status)) {
failStart(statusText("ConfigurationOpen", status));
return;
}
const QByteArray passwordBytes = password.toUtf8();
QUIC_CERTIFICATE_PKCS12 certificate = {};
certificate.Asn1Blob = reinterpret_cast<uint8_t *>(m_pkcs12.data());
certificate.Asn1BlobLength = static_cast<uint32_t>(m_pkcs12.size());
certificate.PrivateKeyPassword = passwordBytes.constData();
QUIC_CREDENTIAL_CONFIG credential = {};
credential.Type = QUIC_CREDENTIAL_TYPE_CERTIFICATE_PKCS12;
credential.Flags = QUIC_CREDENTIAL_FLAG_USE_PORTABLE_CERTIFICATES;
credential.CertificatePkcs12 = &certificate;
status = m_api->ConfigurationLoadCredential(m_configuration, &credential);
if (QUIC_FAILED(status)) {
failStart(statusText("ConfigurationLoadCredential", status));
return;
}
status = m_api->ListenerOpen(m_registration, &Impl::listenerCallback, this, &m_listener);
if (QUIC_FAILED(status)) {
failStart(statusText("ListenerOpen", status));
return;
}
QUIC_ADDR address = {};
const QByteArray bindBytes = bindAddress.trimmed().toLatin1();
if (!QuicAddrFromString(bindBytes.constData(), port, &address)) {
failStart(QStringLiteral("QUIC 监听地址必须是数字 IPv4 或 IPv6 地址"));
return;
}
status = m_api->ListenerStart(m_listener, &alpnBuffer, 1, &address);
if (QUIC_FAILED(status)) {
failStart(statusText("ListenerStart", status));
return;
}
QUIC_ADDR boundAddress = {};
uint32_t boundAddressSize = sizeof(boundAddress);
status = m_api->GetParam(m_listener, QUIC_PARAM_LISTENER_LOCAL_ADDRESS,
&boundAddressSize, &boundAddress);
if (QUIC_FAILED(status)) {
failStart(statusText("GetParam(LISTENER_LOCAL_ADDRESS)", status));
return;
}
m_bindAddress = bindAddress.trimmed();
m_boundPort.store(QuicAddrGetPort(&boundAddress));
m_running.store(true);
m_stopping.store(false);
postRunning(true, m_boundPort.load(),
QStringLiteral("监听 %1,ALPN %2")
.arg(addressText(&boundAddress), QString::fromUtf8(m_alpn)));
postLog(QStringLiteral("INFO"),
QStringLiteral("真实 QUIC/TLS 服务端已启动:%1").arg(addressText(&boundAddress)));
}
void stop()
{
if (!m_api && !m_library.isLoaded())
return;
m_stopping.store(true);
m_running.store(false);
if (m_listener && m_api) {
m_api->ListenerStop(m_listener);
std::unique_lock<std::mutex> lock(m_stateMutex);
m_stateChanged.wait_for(lock, std::chrono::seconds(2), [this]() {
return m_listenerStopped;
});
lock.unlock();
m_api->ListenerClose(m_listener);
m_listener = nullptr;
}
std::vector<ConnectionContext *> connections;
{
std::lock_guard<std::mutex> lock(m_stateMutex);
connections.assign(m_connections.begin(), m_connections.end());
}
for (ConnectionContext *connection : connections) {
if (connection && connection->handle)
m_api->ConnectionShutdown(connection->handle,
QUIC_CONNECTION_SHUTDOWN_FLAG_NONE,
kShutdownCode);
}
{
std::unique_lock<std::mutex> lock(m_stateMutex);
m_stateChanged.wait_for(lock, std::chrono::seconds(3), [this]() {
return m_liveContexts == 0;
});
}
if (m_configuration && m_api) {
m_api->ConfigurationClose(m_configuration);
m_configuration = nullptr;
}
if (m_registration && m_api) {
m_api->RegistrationClose(m_registration);
m_registration = nullptr;
}
if (m_api && m_close) {
m_close(m_api);
m_api = nullptr;
}
m_open = nullptr;
m_close = nullptr;
m_library.unload();
m_pkcs12.clear();
m_boundPort.store(0);
m_listenerStopped = false;
if (m_stopping.exchange(false)) {
postRunning(false, 0, QStringLiteral("QUIC Server 已停止"));
postLog(QStringLiteral("INFO"), QStringLiteral("QUIC Server 已停止"));
}
}
private:
struct SendContext {
explicit SendContext(const QByteArray &value)
: bytes(value)
{
buffer.Buffer = reinterpret_cast<uint8_t *>(bytes.data());
buffer.Length = static_cast<uint32_t>(bytes.size());
}
QByteArray bytes;
QUIC_BUFFER buffer = {};
};
struct ConnectionContext {
Impl *server = nullptr;
HQUIC handle = nullptr;
HQUIC stream = nullptr;
QString peerAddress;
QString peerIp;
QByteArray receiveBuffer;
QString nodeId;
QString bootId;
QString sessionId;
quint64 lastInboundSequence = 0;
quint64 nextOutboundSequence = 0;
bool hasInboundSequence = false;
bool receivedFirstMessage = false;
bool registrationAttempted = false;
bool registered = false;
bool hasHeartbeatSequence = false;
quint64 lastHeartbeatSequence = 0;
std::atomic<int> references{1};
std::mutex mutex;
void addRef() { references.fetch_add(1); }
void release()
{
if (references.fetch_sub(1) == 1)
server->destroyContext(this);
}
};
void failStart(const QString &error)
{
postLog(QStringLiteral("ERROR"), error);
postRunning(false, 0, error);
stop();
}
void destroyContext(ConnectionContext *context)
{
{
std::lock_guard<std::mutex> lock(m_stateMutex);
if (m_liveContexts > 0)
--m_liveContexts;
}
m_stateChanged.notify_all();
delete context;
}
void addConnection(ConnectionContext *context)
{
int count = 0;
{
std::lock_guard<std::mutex> lock(m_stateMutex);
m_connections.insert(context);
++m_liveContexts;
count = static_cast<int>(m_connections.size());
}
postClientCount(count);
postLog(QStringLiteral("INFO"),
QStringLiteral("QUIC 客户端接入:%1").arg(context->peerAddress));
}
void removeConnection(ConnectionContext *context)
{
int count = 0;
{
std::lock_guard<std::mutex> lock(m_stateMutex);
m_connections.erase(context);
count = static_cast<int>(m_connections.size());
}
m_stateChanged.notify_all();
postClientCount(count);
postLog(QStringLiteral("INFO"),
QStringLiteral("QUIC 客户端断开:%1").arg(context->peerAddress));
}
bool processFrame(ConnectionContext *context, const QByteArray &payload)
{
cmvr::quic_edge::v1::EdgeControlEnvelope envelope;
if (!envelope.ParseFromArray(payload.constData(), payload.size())) {
protocolViolation(context, QStringLiteral("控制消息 protobuf 解析失败"), 0);
return false;
}
if (envelope.protocol_version() != kProtocolVersion) {
protocolViolation(context, QStringLiteral("协议版本不是 1"), envelope.message_sequence());
return false;
}
if (context->hasInboundSequence && envelope.message_sequence() <= context->lastInboundSequence) {
protocolViolation(context, QStringLiteral("message_sequence 未严格递增"),
envelope.message_sequence());
return false;
}
context->hasInboundSequence = true;
context->lastInboundSequence = envelope.message_sequence();
if (!context->receivedFirstMessage && !envelope.has_node_register_request()) {
protocolViolation(context, QStringLiteral("NodeRegisterRequest 必须是首条控制消息"),
envelope.message_sequence());
return false;
}
context->receivedFirstMessage = true;
if (envelope.has_node_register_request())
return handleRegistration(context, envelope.node_register_request());
if (envelope.has_node_heartbeat())
return handleHeartbeat(context, envelope.node_heartbeat());
if (envelope.has_media_session_open()) {
postLog(QStringLiteral("INFO"), QStringLiteral("媒体会话已打开:epoch=%1")
.arg(envelope.media_session_open().session_epoch()));
return true;
}
if (envelope.has_media_track_descriptor()) {
const auto &descriptor = envelope.media_track_descriptor();
postLog(QStringLiteral("INFO"), QStringLiteral("媒体轨道:ID=%1,设备=%2,编码=%3")
.arg(descriptor.track_id())
.arg(QString::fromStdString(descriptor.device_id()),
QString::fromStdString(descriptor.codec())));
return true;
}
postLog(QStringLiteral("WARN"), QStringLiteral("收到未处理的 QUIC 控制消息"));
return true;
}
bool handleRegistration(ConnectionContext *context,
const cmvr::quic_edge::v1::NodeRegisterRequest &request)
{
const auto &node = request.node();
if (context->registrationAttempted || context->registered ||
node.node_id().empty() || node.boot_id().empty()) {
protocolViolation(context, QStringLiteral("节点注册无效或重复"), context->lastInboundSequence);
return false;
}
context->registrationAttempted = true;
context->nodeId = QString::fromStdString(node.node_id());
context->bootId = QString::fromStdString(node.boot_id());
context->sessionId = QStringLiteral("qt-%1-%2")
.arg(QDateTime::currentMSecsSinceEpoch())
.arg(m_nextSession.fetch_add(1));
context->registered = true;
m_registrations.fetch_add(1);
cmvr::quic_edge::v1::EdgeControlEnvelope response;
response.set_protocol_version(kProtocolVersion);
response.set_message_sequence(context->nextOutboundSequence++);
auto *registration = response.mutable_node_register_response();
registration->set_accepted(true);
registration->set_session_id(context->sessionId.toStdString());
registration->set_message("accepted by CMVR Qt QUIC test server");
registration->set_heartbeat_interval_ms(1000);
registration->set_observed_source_ip(context->peerIp.toStdString());
if (!sendEnvelope(context, response))
return false;
const QString grpcEndpoint = QStringLiteral("%1:%2%3")
.arg(QString::fromStdString(node.grpc_endpoint().host()))
.arg(node.grpc_endpoint().port())
.arg(node.grpc_endpoint().tls() ? QStringLiteral(" TLS") : QString());
postNode(context->nodeId, context->bootId, grpcEndpoint, context->peerAddress);
postLog(QStringLiteral("OK"), QStringLiteral("节点注册成功:%1,session=%2")
.arg(context->nodeId, context->sessionId));
postStatistics();
return true;
}
bool handleHeartbeat(ConnectionContext *context,
const cmvr::quic_edge::v1::NodeHeartbeat &heartbeat)
{
if (!context->registered ||
QString::fromStdString(heartbeat.session_id()) != context->sessionId ||
QString::fromStdString(heartbeat.node_id()) != context->nodeId ||
QString::fromStdString(heartbeat.boot_id()) != context->bootId ||
(context->hasHeartbeatSequence &&
heartbeat.sequence() <= context->lastHeartbeatSequence)) {
protocolViolation(context, QStringLiteral("心跳 session 与注册会话不一致"),
context->lastInboundSequence);
return false;
}
context->hasHeartbeatSequence = true;
context->lastHeartbeatSequence = heartbeat.sequence();
m_heartbeats.fetch_add(1);
cmvr::quic_edge::v1::EdgeControlEnvelope response;
response.set_protocol_version(kProtocolVersion);
response.set_message_sequence(context->nextOutboundSequence++);
auto *ack = response.mutable_node_heartbeat_ack();
ack->set_accepted(true);
ack->set_acknowledged_sequence(heartbeat.sequence());
ack->set_message("ok");
ack->set_server_time_unix_ms(QDateTime::currentMSecsSinceEpoch());
ack->set_observed_source_ip(context->peerIp.toStdString());
ack->set_session_id(context->sessionId.toStdString());
if (!sendEnvelope(context, response))
return false;
postLog(QStringLiteral("INFO"), QStringLiteral("心跳 #%1,节点 %2,设备 %3")
.arg(heartbeat.sequence())
.arg(context->nodeId)
.arg(heartbeat.has_device_manager()
? heartbeat.device_manager().devices_size() : 0));
postStatistics();
return true;
}
bool sendEnvelope(ConnectionContext *context,
const cmvr::quic_edge::v1::EdgeControlEnvelope &envelope)
{
std::string payload;
if (!envelope.SerializeToString(&payload) || payload.size() > kMaximumControlFrameBytes)
return false;
QByteArray frame;
frame.resize(static_cast<int>(payload.size()) + 4);
const quint32 size = static_cast<quint32>(payload.size());
frame[0] = static_cast<char>((size >> 24) & 0xff);
frame[1] = static_cast<char>((size >> 16) & 0xff);
frame[2] = static_cast<char>((size >> 8) & 0xff);
frame[3] = static_cast<char>(size & 0xff);
if (!payload.empty())
std::memcpy(frame.data() + 4, payload.data(), payload.size());
SendContext *send = new SendContext(frame);
QUIC_STATUS status = QUIC_STATUS_INVALID_PARAMETER;
{
std::lock_guard<std::mutex> lock(context->mutex);
if (context->stream)
status = m_api->StreamSend(context->stream, &send->buffer, 1,
QUIC_SEND_FLAG_NONE, send);
}
if (QUIC_FAILED(status)) {
delete send;
postLog(QStringLiteral("ERROR"), statusText("StreamSend", status));
return false;
}
return true;
}
void protocolViolation(ConnectionContext *context, const QString &message,
quint64 relatedSequence)
{
postLog(QStringLiteral("ERROR"), QStringLiteral("QUIC 协议错误:%1").arg(message));
if (context->stream) {
cmvr::quic_edge::v1::EdgeControlEnvelope response;
response.set_protocol_version(kProtocolVersion);
response.set_message_sequence(context->nextOutboundSequence++);
auto *error = response.mutable_protocol_error();
error->set_code(1);
error->set_message(message.toStdString());
error->set_related_message_sequence(relatedSequence);
error->set_fatal(true);
sendEnvelope(context, response);
}
if (context->handle)
m_api->ConnectionShutdown(context->handle, QUIC_CONNECTION_SHUTDOWN_FLAG_NONE,
kProtocolErrorCode);
}
QUIC_STATUS onStreamReceive(ConnectionContext *context, QUIC_STREAM_EVENT *event)
{
quint64 total = 0;
std::vector<QByteArray> payloads;
bool frameTooLarge = false;
for (uint32_t index = 0; index < event->RECEIVE.BufferCount; ++index)
total += event->RECEIVE.Buffers[index].Length;
if (total > kMaximumControlFrameBytes + 4U)
return QUIC_STATUS_BUFFER_TOO_SMALL;
{
std::lock_guard<std::mutex> lock(context->mutex);
for (uint32_t index = 0; index < event->RECEIVE.BufferCount; ++index) {
const QUIC_BUFFER &buffer = event->RECEIVE.Buffers[index];
context->receiveBuffer.append(reinterpret_cast<const char *>(buffer.Buffer),
static_cast<int>(buffer.Length));
}
while (context->receiveBuffer.size() >= 4) {
const unsigned char *data = reinterpret_cast<const unsigned char *>(
context->receiveBuffer.constData());
const quint32 size = (static_cast<quint32>(data[0]) << 24) |
(static_cast<quint32>(data[1]) << 16) |
(static_cast<quint32>(data[2]) << 8) |
static_cast<quint32>(data[3]);
if (size > kMaximumControlFrameBytes) {
frameTooLarge = true;
break;
}
if (context->receiveBuffer.size() < static_cast<int>(size + 4U))
break;
payloads.push_back(context->receiveBuffer.mid(4, static_cast<int>(size)));
context->receiveBuffer.remove(0, static_cast<int>(size + 4U));
}
}
if (frameTooLarge) {
protocolViolation(context, QStringLiteral("控制帧超过 1 MiB"), 0);
return QUIC_STATUS_BUFFER_TOO_SMALL;
}
for (const QByteArray &payload : payloads) {
if (!processFrame(context, payload))
break;
}
return QUIC_STATUS_SUCCESS;
}
static QUIC_STATUS QUIC_API listenerCallback(HQUIC, void *context,
QUIC_LISTENER_EVENT *event)
{
Impl *server = static_cast<Impl *>(context);
if (!server || !event)
return QUIC_STATUS_INVALID_PARAMETER;
if (event->Type == QUIC_LISTENER_EVENT_STOP_COMPLETE) {
{
std::lock_guard<std::mutex> lock(server->m_stateMutex);
server->m_listenerStopped = true;
}
server->m_stateChanged.notify_all();
return QUIC_STATUS_SUCCESS;
}
if (event->Type != QUIC_LISTENER_EVENT_NEW_CONNECTION || server->m_stopping.load())
return QUIC_STATUS_SUCCESS;
ConnectionContext *connection = new ConnectionContext;
connection->server = server;
connection->handle = event->NEW_CONNECTION.Connection;
connection->peerAddress = event->NEW_CONNECTION.Info
? addressText(event->NEW_CONNECTION.Info->RemoteAddress) : QStringLiteral("--");
connection->peerIp = event->NEW_CONNECTION.Info
? addressIpText(event->NEW_CONNECTION.Info->RemoteAddress) : QStringLiteral("--");
server->m_api->SetCallbackHandler(connection->handle,
reinterpret_cast<void *>(&Impl::connectionCallback), connection);
server->addConnection(connection);
const QUIC_STATUS status = server->m_api->ConnectionSetConfiguration(
connection->handle, server->m_configuration);
if (QUIC_FAILED(status)) {
server->m_api->ConnectionClose(connection->handle);
connection->handle = nullptr;
server->removeConnection(connection);
connection->release();
}
return status;
}
static QUIC_STATUS QUIC_API connectionCallback(HQUIC handle, void *context,
QUIC_CONNECTION_EVENT *event)
{
ConnectionContext *connection = static_cast<ConnectionContext *>(context);
if (!connection || !event)
return QUIC_STATUS_INVALID_PARAMETER;
Impl *server = connection->server;
switch (event->Type) {
case QUIC_CONNECTION_EVENT_CONNECTED:
server->postLog(QStringLiteral("OK"),
QStringLiteral("QUIC/TLS 握手完成:%1").arg(connection->peerAddress));
break;
case QUIC_CONNECTION_EVENT_PEER_STREAM_STARTED:
if ((event->PEER_STREAM_STARTED.Flags & QUIC_STREAM_OPEN_FLAG_UNIDIRECTIONAL) ||
connection->stream) {
server->m_api->StreamClose(event->PEER_STREAM_STARTED.Stream);
server->m_api->ConnectionShutdown(handle, QUIC_CONNECTION_SHUTDOWN_FLAG_NONE,
kProtocolErrorCode);
break;
}
{
std::lock_guard<std::mutex> lock(connection->mutex);
connection->stream = event->PEER_STREAM_STARTED.Stream;
connection->addRef();
}
server->m_api->SetCallbackHandler(connection->stream,
reinterpret_cast<void *>(&Impl::streamCallback), connection);
server->postLog(QStringLiteral("INFO"), QStringLiteral("可靠控制流已建立"));
break;
case QUIC_CONNECTION_EVENT_DATAGRAM_RECEIVED:
if (event->DATAGRAM_RECEIVED.Buffer) {
server->m_datagrams.fetch_add(1);
server->m_datagramBytes.fetch_add(event->DATAGRAM_RECEIVED.Buffer->Length);
server->postStatistics();
}
break;
case QUIC_CONNECTION_EVENT_SHUTDOWN_INITIATED_BY_TRANSPORT:
server->postLog(QStringLiteral("WARN"),
statusText("QUIC transport", event->SHUTDOWN_INITIATED_BY_TRANSPORT.Status));
break;
case QUIC_CONNECTION_EVENT_SHUTDOWN_COMPLETE:
connection->handle = nullptr;
server->m_api->ConnectionClose(handle);
server->removeConnection(connection);
connection->release();
break;
default:
break;
}
return QUIC_STATUS_SUCCESS;
}
static QUIC_STATUS QUIC_API streamCallback(HQUIC stream, void *context,
QUIC_STREAM_EVENT *event)
{
ConnectionContext *connection = static_cast<ConnectionContext *>(context);
if (!connection || !event)
return QUIC_STATUS_INVALID_PARAMETER;
Impl *server = connection->server;
switch (event->Type) {
case QUIC_STREAM_EVENT_RECEIVE:
return server->onStreamReceive(connection, event);
case QUIC_STREAM_EVENT_SEND_COMPLETE:
delete static_cast<SendContext *>(event->SEND_COMPLETE.ClientContext);
break;
case QUIC_STREAM_EVENT_PEER_SEND_SHUTDOWN:
server->m_api->StreamShutdown(stream, QUIC_STREAM_SHUTDOWN_FLAG_GRACEFUL, 0);
break;
case QUIC_STREAM_EVENT_SHUTDOWN_COMPLETE:
{
std::lock_guard<std::mutex> lock(connection->mutex);
if (connection->stream == stream)
connection->stream = nullptr;
}
server->m_api->StreamClose(stream);
connection->release();
break;
default:
break;
}
return QUIC_STATUS_SUCCESS;
}
void postRunning(bool running, quint16 port, const QString &detail)
{
QMetaObject::invokeMethod(m_owner, [this, running, port, detail]() {
emit m_owner->runningChanged(running, port, detail);
}, Qt::QueuedConnection);
}
void postClientCount(int count)
{
QMetaObject::invokeMethod(m_owner, [this, count]() {
emit m_owner->clientCountChanged(count);
}, Qt::QueuedConnection);
}
void postStatistics()
{
const quint64 registrations = m_registrations.load();
const quint64 heartbeats = m_heartbeats.load();
const quint64 datagrams = m_datagrams.load();
const quint64 bytes = m_datagramBytes.load();
QMetaObject::invokeMethod(m_owner, [this, registrations, heartbeats, datagrams, bytes]() {
emit m_owner->statisticsChanged(registrations, heartbeats, datagrams, bytes);
}, Qt::QueuedConnection);
}
void postNode(const QString &nodeId, const QString &bootId,
const QString &grpcEndpoint, const QString &peerAddress)
{
QMetaObject::invokeMethod(m_owner,
[this, nodeId, bootId, grpcEndpoint, peerAddress]() {
emit m_owner->nodeRegistered(nodeId, bootId, grpcEndpoint, peerAddress);
}, Qt::QueuedConnection);
}
void postLog(const QString &level, const QString &message)
{
QMetaObject::invokeMethod(m_owner, [this, level, message]() {
emit m_owner->logMessage(level, message);
}, Qt::QueuedConnection);
}
QuicTestServer *m_owner = nullptr;
QLibrary m_library;
MsQuicOpenVersionFn m_open = nullptr;
MsQuicCloseFn m_close = nullptr;
const QUIC_API_TABLE *m_api = nullptr;
HQUIC m_registration = nullptr;
HQUIC m_configuration = nullptr;
HQUIC m_listener = nullptr;
QByteArray m_alpn;
QByteArray m_pkcs12;
QString m_bindAddress;
std::atomic<bool> m_running{false};
std::atomic<bool> m_stopping{false};
std::atomic<quint16> m_boundPort{0};
std::atomic<quint64> m_registrations{0};
std::atomic<quint64> m_heartbeats{0};
std::atomic<quint64> m_datagrams{0};
std::atomic<quint64> m_datagramBytes{0};
std::atomic<quint64> m_nextSession{1};
std::mutex m_stateMutex;
std::condition_variable m_stateChanged;
std::set<ConnectionContext *> m_connections;
int m_liveContexts = 0;
bool m_listenerStopped = false;
};
#else
class QuicTestServer::Impl
{
public:
explicit Impl(QuicTestServer *owner)
: m_owner(owner)
{
}
bool isRunning() const { return false; }
quint16 boundPort() const { return 0; }
void start(const QString &bindAddress, quint16 port, const QString &alpn,
const QString &pkcs12Path, const QString &password)
{
Q_UNUSED(bindAddress)
Q_UNUSED(port)
Q_UNUSED(alpn)
Q_UNUSED(pkcs12Path)
Q_UNUSED(password)
const QString detail = QStringLiteral(
"当前构建未启用 MsQuic;请使用 -DCMVR_WITH_MSQUIC=ON 重新配置");
QMetaObject::invokeMethod(m_owner, [this, detail]() {
emit m_owner->runningChanged(false, 0, detail);
emit m_owner->logMessage(QStringLiteral("WARN"), detail);
}, Qt::QueuedConnection);
}
void stop() {}
private:
QuicTestServer *m_owner = nullptr;
};
#endif
QuicTestServer::QuicTestServer(QObject *parent)
: QObject(parent), m_impl(new Impl(this))
{
}
QuicTestServer::~QuicTestServer() = default;
bool QuicTestServer::isRunning() const { return m_impl->isRunning(); }
quint16 QuicTestServer::boundPort() const { return m_impl->boundPort(); }
void QuicTestServer::startServer(const QString &bindAddress, quint16 port,
const QString &alpn, const QString &pkcs12Path,
const QString &password)
{
m_impl->start(bindAddress, port, alpn, pkcs12Path, password);
}
void QuicTestServer::stopServer()
{
m_impl->stop();
}