ClientConnection.cpp 12 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331
  1. // ClientConnection.cpp - 用户态连接管理实现
  2. #include <ntddk.h>
  3. #include <wdf.h>
  4. #include "ClientConnection.h"
  5. #include "EventRingBuffer.h"
  6. #include "SerialFilter.h" // g_sequence_counter
  7. #include "../common/CommKitIoctl.h"
  8. #include "../common/CommKitEvents.h"
  9. namespace commkit_driver {
  10. // 全局实例
  11. ClientConnection g_ClientConnection;
  12. void ClientConnection::Initialize() {
  13. ports_count_ = 0;
  14. control_device_ = nullptr;
  15. callback_registered_ = FALSE;
  16. KeInitializeSpinLock(&ports_lock_);
  17. RtlZeroMemory(ports_table_, sizeof(ports_table_));
  18. }
  19. void ClientConnection::Cleanup() {
  20. // 释放所有端口的环形缓冲
  21. KIRQL old_irql;
  22. KeAcquireSpinLock(&ports_lock_, &old_irql);
  23. for (ULONG i = 0; i < ports_count_; ++i) {
  24. PDEVICE_CONTEXT ctx = ports_table_[i].Context;
  25. if (ctx && ctx->RingBuffer) {
  26. EventRingBuffer* ring = (EventRingBuffer*)ctx->RingBuffer;
  27. ring->Cleanup();
  28. ExFreePoolWithTag(ring, 'RBCK');
  29. ctx->RingBuffer = nullptr;
  30. }
  31. }
  32. ports_count_ = 0;
  33. KeReleaseSpinLock(&ports_lock_, old_irql);
  34. }
  35. void ClientConnection::RegisterFilterDevice(WDFDEVICE device) {
  36. KIRQL old_irql;
  37. KeAcquireSpinLock(&ports_lock_, &old_irql);
  38. if (ports_count_ < 256) {
  39. PDEVICE_CONTEXT ctx = DeviceGetContext(device);
  40. ports_table_[ports_count_].ComNumber = 0; // 延迟到 UpdatePortComNumber
  41. ports_table_[ports_count_].Context = ctx;
  42. ports_count_++;
  43. ctx->ComNumber = 0;
  44. ctx->MonitoringEnabled = FALSE;
  45. ctx->WdfDevice = device;
  46. }
  47. KeReleaseSpinLock(&ports_lock_, old_irql);
  48. }
  49. void ClientConnection::UpdatePortComNumber(WDFDEVICE device, ULONG com_number) {
  50. KIRQL old_irql;
  51. KeAcquireSpinLock(&ports_lock_, &old_irql);
  52. for (ULONG i = 0; i < ports_count_; ++i) {
  53. if (ports_table_[i].Context &&
  54. ports_table_[i].Context->WdfDevice == device) {
  55. ports_table_[i].ComNumber = com_number;
  56. ports_table_[i].Context->ComNumber = com_number;
  57. break;
  58. }
  59. }
  60. KeReleaseSpinLock(&ports_lock_, old_irql);
  61. }
  62. void ClientConnection::UnregisterFilterDevice(ULONG com_number) {
  63. KIRQL old_irql;
  64. KeAcquireSpinLock(&ports_lock_, &old_irql);
  65. for (ULONG i = 0; i < ports_count_; ++i) {
  66. if (ports_table_[i].ComNumber == com_number) {
  67. // 移动最后一个元素到当前位置
  68. PDEVICE_CONTEXT ctx = ports_table_[i].Context;
  69. if (ctx && ctx->RingBuffer) {
  70. EventRingBuffer* ring = (EventRingBuffer*)ctx->RingBuffer;
  71. ring->Cleanup();
  72. ExFreePoolWithTag(ring, 'RBCK');
  73. ctx->RingBuffer = nullptr;
  74. }
  75. ports_table_[i] = ports_table_[ports_count_ - 1];
  76. ports_count_--;
  77. break;
  78. }
  79. }
  80. KeReleaseSpinLock(&ports_lock_, old_irql);
  81. }
  82. PDEVICE_CONTEXT ClientConnection::FindPortContext(ULONG com_number) {
  83. PDEVICE_CONTEXT result = nullptr;
  84. KIRQL old_irql;
  85. KeAcquireSpinLock(&ports_lock_, &old_irql);
  86. for (ULONG i = 0; i < ports_count_; ++i) {
  87. if (ports_table_[i].ComNumber == com_number) {
  88. result = ports_table_[i].Context;
  89. break;
  90. }
  91. }
  92. KeReleaseSpinLock(&ports_lock_, old_irql);
  93. return result;
  94. }
  95. NTSTATUS ClientConnection::HandleRegisterCallback() {
  96. callback_registered_ = TRUE;
  97. return STATUS_SUCCESS;
  98. }
  99. NTSTATUS ClientConnection::HandleAttachPort(ULONG com_number) {
  100. // 拒绝 com_number==0:EvtDevicePrepareHardware 解析失败时 ComNumber 保持 0,
  101. // 若允许 attach 会误匹配第一个未解析出编号的设备
  102. if (com_number == 0) {
  103. DbgPrint("[CommModifyKit] HandleAttachPort: com_number=0 rejected\n");
  104. return STATUS_INVALID_PARAMETER;
  105. }
  106. PDEVICE_CONTEXT ctx = FindPortContext(com_number);
  107. if (!ctx) {
  108. DbgPrint("[CommModifyKit] HandleAttachPort: COM%u NOT in ports_table (count=%u)\n",
  109. com_number, ports_count_);
  110. // 打印 ports_table_ 里所有已注册的 COM 编号,便于诊断
  111. KIRQL old_irql;
  112. KeAcquireSpinLock(&ports_lock_, &old_irql);
  113. for (ULONG i = 0; i < ports_count_; ++i) {
  114. DbgPrint("[CommModifyKit] ports_table[%u].ComNumber=%u\n",
  115. i, ports_table_[i].ComNumber);
  116. }
  117. KeReleaseSpinLock(&ports_lock_, old_irql);
  118. return STATUS_DEVICE_DOES_NOT_EXIST;
  119. }
  120. DbgPrint("[CommModifyKit] HandleAttachPort: COM%u found in ports_table\n", com_number);
  121. // 如果未分配环形缓冲,先分配
  122. if (!ctx->RingBuffer) {
  123. EventRingBuffer* ring = (EventRingBuffer*)ExAllocatePool2(
  124. POOL_FLAG_NON_PAGED, sizeof(EventRingBuffer), 'RBCK');
  125. if (!ring) {
  126. return STATUS_INSUFFICIENT_RESOURCES;
  127. }
  128. // placement new 等价:手动调用构造
  129. RtlZeroMemory(ring, sizeof(EventRingBuffer));
  130. NTSTATUS status = ring->Initialize();
  131. if (!NT_SUCCESS(status)) {
  132. ring->Cleanup();
  133. ExFreePoolWithTag(ring, 'RBCK');
  134. return status;
  135. }
  136. ctx->RingBuffer = ring;
  137. }
  138. // 推入 OP_OPEN 事件(使用全局共享序列号,保证所有事件类型 Sequence 唯一递增)
  139. EventRingBuffer* ring = (EventRingBuffer*)ctx->RingBuffer;
  140. LONG cur = InterlockedIncrement(&g_sequence_counter);
  141. UINT64 timestamp = KeQueryInterruptTime(); // 100ns 单位
  142. ring->Push((ULONG)cur, timestamp, com_number, COMMKIT_OP_OPEN, 0, nullptr);
  143. ctx->MonitoringEnabled = TRUE;
  144. return STATUS_SUCCESS;
  145. }
  146. NTSTATUS ClientConnection::HandleDetachPort(ULONG com_number) {
  147. PDEVICE_CONTEXT ctx = FindPortContext(com_number);
  148. if (!ctx) {
  149. return STATUS_DEVICE_DOES_NOT_EXIST;
  150. }
  151. ctx->MonitoringEnabled = FALSE;
  152. // 推入 OP_CLOSE 事件(使用全局共享序列号)
  153. if (ctx->RingBuffer) {
  154. EventRingBuffer* ring = (EventRingBuffer*)ctx->RingBuffer;
  155. LONG cur = InterlockedIncrement(&g_sequence_counter);
  156. UINT64 timestamp = KeQueryInterruptTime();
  157. ring->Push((ULONG)cur, timestamp, com_number, COMMKIT_OP_CLOSE, 0, nullptr);
  158. }
  159. return STATUS_SUCCESS;
  160. }
  161. NTSTATUS ClientConnection::HandleWritePort(ULONG com_number, PVOID data, ULONG len) {
  162. PDEVICE_CONTEXT ctx = FindPortContext(com_number);
  163. if (!ctx || !ctx->LowerDevice) {
  164. return STATUS_DEVICE_DOES_NOT_EXIST;
  165. }
  166. // 构造同步写 IRP 发送给下层串口设备
  167. KEVENT event;
  168. KeInitializeEvent(&event, NotificationEvent, FALSE);
  169. IO_STATUS_BLOCK io_status = {};
  170. PIRP irp = IoBuildSynchronousFsdRequest(
  171. IRP_MJ_WRITE, ctx->LowerDevice, data, len, nullptr, &event, &io_status);
  172. if (!irp) {
  173. return STATUS_INSUFFICIENT_RESOURCES;
  174. }
  175. NTSTATUS status = IoCallDriver(ctx->LowerDevice, irp);
  176. if (status == STATUS_PENDING) {
  177. KeWaitForSingleObject(&event, Executive, KernelMode, FALSE, nullptr);
  178. status = io_status.Status;
  179. }
  180. // 注:此 IRP 直接发给 LowerDevice,绕过 SerialFilter::DispatchWrite,
  181. // 不会触发 OnWriteComplete 捕获。需手动将写入数据 Push 为 OP_WRITE 事件,
  182. // 否则监控端看不到注入的写流量(监控数据不完整)
  183. if (NT_SUCCESS(status) && ctx->MonitoringEnabled && ctx->RingBuffer) {
  184. ULONG written = (ULONG)io_status.Information;
  185. if (written > 0 && data) {
  186. LONG seq = InterlockedIncrement(&g_sequence_counter);
  187. UINT64 timestamp = KeQueryInterruptTime();
  188. EventRingBuffer* ring = (EventRingBuffer*)ctx->RingBuffer;
  189. ring->Push((ULONG)seq, timestamp, com_number, COMMKIT_OP_WRITE,
  190. written, data);
  191. }
  192. }
  193. return status;
  194. }
  195. NTSTATUS ClientConnection::HandleReadPort(ULONG com_number, PVOID data, ULONG len) {
  196. PDEVICE_CONTEXT ctx = FindPortContext(com_number);
  197. if (!ctx || !ctx->LowerDevice) {
  198. return STATUS_DEVICE_DOES_NOT_EXIST;
  199. }
  200. // 构造同步读 IRP 发送给下层串口设备
  201. KEVENT event;
  202. KeInitializeEvent(&event, NotificationEvent, FALSE);
  203. IO_STATUS_BLOCK io_status = {};
  204. PIRP irp = IoBuildSynchronousFsdRequest(
  205. IRP_MJ_READ, ctx->LowerDevice, data, len, nullptr, &event, &io_status);
  206. if (!irp) {
  207. return STATUS_INSUFFICIENT_RESOURCES;
  208. }
  209. NTSTATUS status = IoCallDriver(ctx->LowerDevice, irp);
  210. if (status == STATUS_PENDING) {
  211. KeWaitForSingleObject(&event, Executive, KernelMode, FALSE, nullptr);
  212. status = io_status.Status;
  213. }
  214. // 注:此 IRP 直接发给 LowerDevice,绕过 SerialFilter::DispatchRead,
  215. // 不会触发 OnReadComplete 捕获。需手动将读取数据 Push 为 OP_READ 事件,
  216. // 否则读取的数据既不返回调用者(METHOD_BUFFERED + info=0)也不进 RingBuffer,完全丢失
  217. if (NT_SUCCESS(status) && ctx->MonitoringEnabled && ctx->RingBuffer) {
  218. ULONG read_bytes = (ULONG)io_status.Information;
  219. if (read_bytes > 0 && data) {
  220. LONG seq = InterlockedIncrement(&g_sequence_counter);
  221. UINT64 timestamp = KeQueryInterruptTime();
  222. EventRingBuffer* ring = (EventRingBuffer*)ctx->RingBuffer;
  223. ring->Push((ULONG)seq, timestamp, com_number, COMMKIT_OP_READ,
  224. read_bytes, data);
  225. }
  226. }
  227. return status;
  228. }
  229. NTSTATUS ClientConnection::HandleReadEvents(PVOID out_buf, ULONG out_size, PULONG returned) {
  230. *returned = 0;
  231. if (!callback_registered_) {
  232. return STATUS_DEVICE_NOT_READY;
  233. }
  234. PCOMMKIT_EVENT events = (PCOMMKIT_EVENT)out_buf;
  235. ULONG max_count = out_size / sizeof(COMMKIT_EVENT);
  236. if (max_count == 0) {
  237. return STATUS_BUFFER_TOO_SMALL;
  238. }
  239. ULONG total_popped = 0;
  240. KIRQL old_irql;
  241. KeAcquireSpinLock(&ports_lock_, &old_irql);
  242. // 按端口顺序弹出事件到输出缓冲
  243. for (ULONG i = 0; i < ports_count_ && total_popped < max_count; ++i) {
  244. PDEVICE_CONTEXT ctx = ports_table_[i].Context;
  245. if (!ctx || !ctx->RingBuffer || !ctx->MonitoringEnabled) {
  246. continue;
  247. }
  248. EventRingBuffer* ring = (EventRingBuffer*)ctx->RingBuffer;
  249. LONG popped = ring->PopBatch(events + total_popped, max_count - total_popped);
  250. total_popped += (ULONG)popped;
  251. }
  252. KeReleaseSpinLock(&ports_lock_, old_irql);
  253. // 弹出后按 Sequence 升序排序(锁外执行,不持自旋锁做 O(n²) 操作)
  254. // 各端口 RingBuffer 内部 Sequence 是递增的,但跨端口弹出后顺序被打乱
  255. // (如 port0 的 seq=100-109 可能排在 port1 的 seq=95-99 之前)
  256. // 使用插入排序:事件批量小(通常 ≤16),O(n²) 开销可忽略
  257. if (total_popped > 1) {
  258. for (ULONG i = 1; i < total_popped; ++i) {
  259. COMMKIT_EVENT key = events[i];
  260. LONG j = (LONG)i - 1;
  261. while (j >= 0 && events[j].Sequence > key.Sequence) {
  262. events[j + 1] = events[j];
  263. j--;
  264. }
  265. events[j + 1] = key;
  266. }
  267. }
  268. *returned = total_popped * sizeof(COMMKIT_EVENT);
  269. return STATUS_SUCCESS;
  270. }
  271. NTSTATUS ClientConnection::HandleFreeAll() {
  272. KIRQL old_irql;
  273. KeAcquireSpinLock(&ports_lock_, &old_irql);
  274. for (ULONG i = 0; i < ports_count_; ++i) {
  275. PDEVICE_CONTEXT ctx = ports_table_[i].Context;
  276. if (ctx) {
  277. ctx->MonitoringEnabled = FALSE;
  278. if (ctx->RingBuffer) {
  279. ((EventRingBuffer*)ctx->RingBuffer)->Clear();
  280. }
  281. }
  282. }
  283. KeReleaseSpinLock(&ports_lock_, old_irql);
  284. callback_registered_ = FALSE;
  285. return STATUS_SUCCESS;
  286. }
  287. } // namespace commkit_driver