@cryptotaxi247 / netdata / commits / 79d43657a

feat: replace cgroups-ebpf SHM transport with netipc IPC (#22221)

* feat: vendor netipc on latest master * plugin-ipc-integration-v2: Fix https://github.com/netdata/netdata/pull/22221#discussion_r3100547470 * plugin-ipc-integration-v2: Fix https://github.com/netdata/netdata/pull/22221#discussion_r3101183364 * plugin-ipc-integration-v2: Fix https://github.com/netdata/netdata/pull/22221#discussion_r3101183424 * plugin-ipc-integration-v2: Fix https://github.com/netdata/netdata/pull/22221#discussion_r3101183478 * plugin-ipc-integration-v2: Fix https://github.com/netdata/netdata/pull/22221#discussion_r3101183546 * plugin-ipc-integration-v2: Fix https://github.com/netdata/netdata/pull/22221#discussion_r3101183584 * plugin-ipc-integration-v2: Fix https://github.com/netdata/netdata/pull/22221#discussion_r3102848766 * vendor: sync rust windows pipe wait fix * vendor: sync go windows shm spin fix * vendor: sync windows pipe server wait fix --------- Co-authored-by: thiagoftsm <thiagoftsm@gmail.com> Co-authored-by: Stelios Fragkakis <52996999+stelfrag@users.noreply.github.com>

Costa Tsaousis committed Apr 21, 2026 at 16:55 UTC 79d43657a332daf9947c2f8dfa45a21b8656c5a4
110 files changed +60923 -438
CMakeLists.txt
+42
@@ -1172,6 +1172,13 @@ if(OS_WINDOWS)
1172 list(APPEND LIBNETDATA_FILES ${LIBNETDATA_WIN_FILES})
1173 endif()
1174
1175 +if(OS_LINUX OR OS_WINDOWS)
1176 + list(APPEND LIBNETDATA_FILES
1177 + src/libnetdata/netipc/netipc_netdata.c
1178 + src/libnetdata/netipc/netipc_netdata.h
1179 + )
1180 +endif()
1181 +
1182 if(NOT OS_FREEBSD)
1183 set(LIBNETDATA_STACKTRACE_FILES
1184 src/libnetdata/stacktrace/stacktrace.h
@@ -1954,6 +1961,8 @@ set(CGROUPS_PLUGIN_FILES
1961 src/collectors/cgroups.plugin/cgroup-discovery.c
1962 src/collectors/cgroups.plugin/cgroup-charts.c
1963 src/collectors/cgroups.plugin/cgroup-top.c
1964 + src/collectors/cgroups.plugin/cgroup-netipc.c
1965 + src/collectors/cgroups.plugin/cgroup-netipc.h
1966 )
1967
1968 set(DISKSPACE_PLUGIN_FILES
@@ -2257,6 +2266,35 @@ set_source_files_properties(src/libnetdata/libjudy/vendored/JudyL/j__udyLGet.c P
2266 set_source_files_properties(src/libnetdata/libjudy/vendored/JudyL/JudyLByCount.c PROPERTIES COMPILE_OPTIONS "-DNOSMARTJBB -DNOSMARTJBU -DNOSMARTJLB")
2267 set_source_files_properties(JudyLTables.c PROPERTIES COMPILE_OPTIONS "-I${CMAKE_SOURCE_DIR}/src/libnetdata/libjudy/src/JudyL")
2268
2269 +#
2270 +# build netipc (standalone IPC library, no Netdata deps)
2271 +#
2272 +
2273 +if(OS_LINUX)
2274 + set(NETIPC_FILES
2275 + src/libnetdata/netipc/src/protocol/netipc_protocol.c
2276 + src/libnetdata/netipc/src/transport/posix/netipc_uds.c
2277 + src/libnetdata/netipc/src/transport/posix/netipc_shm.c
2278 + src/libnetdata/netipc/src/service/netipc_service.c
2279 + )
2280 +
2281 + add_library(netipc STATIC ${NETIPC_FILES})
2282 + target_include_directories(netipc PUBLIC ${CMAKE_SOURCE_DIR}/src/libnetdata/netipc/include)
2283 + target_link_libraries(netipc PUBLIC pthread rt)
2284 + set_target_properties(netipc PROPERTIES C_STANDARD 11 C_STANDARD_REQUIRED ON)
2285 +elseif(OS_WINDOWS)
2286 + set(NETIPC_FILES
2287 + src/libnetdata/netipc/src/protocol/netipc_protocol.c
2288 + src/libnetdata/netipc/src/transport/windows/netipc_named_pipe.c
2289 + src/libnetdata/netipc/src/transport/windows/netipc_win_shm.c
2290 + src/libnetdata/netipc/src/service/netipc_service_win.c
2291 + )
2292 +
2293 + add_library(netipc STATIC ${NETIPC_FILES})
2294 + target_include_directories(netipc PUBLIC ${CMAKE_SOURCE_DIR}/src/libnetdata/netipc/include)
2295 + set_target_properties(netipc PROPERTIES C_STANDARD 11 C_STANDARD_REQUIRED ON)
2296 +endif()
2297 +
2298 #
2299 # build libnetdata
2300 #
@@ -2283,6 +2321,10 @@ target_link_libraries(libnetdata PUBLIC
2321 "$<$<BOOL:${LINK_LIBM}>:m>"
2322 "${SYSTEMD_LDFLAGS}")
2323
2324 +if(OS_LINUX OR OS_WINDOWS)
2325 + target_link_libraries(libnetdata PUBLIC netipc)
2326 +endif()
2327 +
2328 if(HAVE_LIBBACKTRACE)
2329 netdata_add_libbacktrace_to_target(libnetdata)
2330 endif()
src/collectors/cgroups.plugin/cgroup-discovery.c
+8 -157
@@ -1,6 +1,7 @@
1 // SPDX-License-Identifier: GPL-3.0-or-later
2
3 #include "cgroup-internals.h"
4 +#include "cgroup-netipc.h"
5
6 // discovery cgroup thread worker jobs
7 #define WORKER_DISCOVERY_INIT 0
@@ -12,11 +13,10 @@
13 #define WORKER_DISCOVERY_UPDATE 6
14 #define WORKER_DISCOVERY_CLEANUP 7
15 #define WORKER_DISCOVERY_COPY 8
15 -#define WORKER_DISCOVERY_SHARE 9
16 -#define WORKER_DISCOVERY_LOCK 10
16 +#define WORKER_DISCOVERY_LOCK 9
17
18 -#if WORKER_UTILIZATION_MAX_JOB_TYPES < 11
19 -#error WORKER_UTILIZATION_MAX_JOB_TYPES has to be at least 11
18 +#if WORKER_UTILIZATION_MAX_JOB_TYPES < 10
19 +#error WORKER_UTILIZATION_MAX_JOB_TYPES has to be at least 10
20 #endif
21
22 struct cgroup *discovered_cgroup_root = NULL;
@@ -25,10 +25,7 @@ char cgroup_chart_id_prefix[] = "cgroup_";
25 char services_chart_id_prefix[] = "systemd_";
26 const char *cgroups_rename_script = NULL;
27
28 -// Shared memory with information from detected cgroups
29 -netdata_ebpf_cgroup_shm_t shm_cgroup_ebpf = {NULL, NULL};
30 -int shm_fd_cgroup_ebpf = -1;
31 -sem_t *shm_mutex_cgroup_ebpf = SEM_FAILED;
28 +// (legacy SHM globals removed — replaced by netipc in cgroup-netipc.c)
29
30 // ----------------------------------------------------------------------------
31
@@ -280,28 +277,6 @@ static inline void discovery_rename_cgroup(struct cgroup *cg) {
277 cg->hash_chart_id = simple_hash(cg->chart_id);
278 }
279
283 -static void is_cgroup_procs_exist(netdata_ebpf_cgroup_shm_body_t *out, char *id) {
284 - struct stat buf;
285 -
286 - snprintfz(out->path, FILENAME_MAX, "%s%s/cgroup.procs", cgroup_cpuset_base, id);
287 - if (likely(stat(out->path, &buf) == 0)) {
288 - return;
289 - }
290 -
291 - snprintfz(out->path, FILENAME_MAX, "%s%s/cgroup.procs", cgroup_blkio_base, id);
292 - if (likely(stat(out->path, &buf) == 0)) {
293 - return;
294 - }
295 -
296 - snprintfz(out->path, FILENAME_MAX, "%s%s/cgroup.procs", cgroup_memory_base, id);
297 - if (likely(stat(out->path, &buf) == 0)) {
298 - return;
299 - }
300 -
301 - out->path[0] = '\0';
302 - out->enabled = 0;
303 -}
304 -
280 static inline void convert_cgroup_to_systemd_service(struct cgroup *cg) {
281 char buffer[CGROUP_CHARTID_LINE_MAX + 1];
282 cg->options |= CGROUP_OPTIONS_SYSTEM_SLICE_SERVICE;
@@ -813,47 +788,6 @@ static inline void discovery_copy_discovered_cgroups_to_reader() {
788 cgroup_root = discovered_cgroup_root;
789 }
790
816 -static inline void discovery_share_cgroups_with_ebpf() {
817 - struct cgroup *cg;
818 - int count;
819 - struct stat buf;
820 -
821 - if (shm_mutex_cgroup_ebpf == SEM_FAILED) {
822 - return;
823 - }
824 - sem_wait(shm_mutex_cgroup_ebpf);
825 -
826 - for (cg = cgroup_root, count = 0; cg && count < cgroup_root_max; cg = cg->next, count++) {
827 - netdata_ebpf_cgroup_shm_body_t *ptr = &shm_cgroup_ebpf.body[count];
828 - char *prefix = (is_cgroup_systemd_service(cg)) ? services_chart_id_prefix : cgroup_chart_id_prefix;
829 - snprintfz(ptr->name, CGROUP_EBPF_NAME_SHARED_LENGTH - 1, "%s%s", prefix, cg->chart_id);
830 - ptr->hash = simple_hash(ptr->name);
831 - ptr->options = cg->options;
832 - ptr->enabled = cg->enabled;
833 - if (cgroup_use_unified_cgroups) {
834 - snprintfz(ptr->path, FILENAME_MAX, "%s%s/cgroup.procs", cgroup_unified_base, cg->id);
835 - if (likely(stat(ptr->path, &buf) == -1)) {
836 - ptr->path[0] = '\0';
837 - ptr->enabled = 0;
838 - }
839 - } else {
840 - is_cgroup_procs_exist(ptr, cg->id);
841 - }
842 -
843 - netdata_log_debug(D_CGROUP, "cgroup shared: NAME=%s, ENABLED=%d", ptr->name, ptr->enabled);
844 - }
845 -
846 - if (unlikely(cg != NULL)) {
847 - nd_log_limit_static_global_var(erl, 3600, 0);
848 - nd_log_limit(&erl, NDLS_COLLECTORS, NDLP_WARNING,
849 - "CGROUP: shared memory buffer full (%d cgroups). Some cgroups were not shared with eBPF.",
850 - cgroup_root_max);
851 - }
852 -
853 - shm_cgroup_ebpf.header->cgroup_root_count = count;
854 - sem_post(shm_mutex_cgroup_ebpf);
855 -}
856 -
791 static inline void discovery_find_all_cgroups_v1() {
792 if (cgroup_enable_cpuacct) {
793 if (discovery_find_walkdir(cgroup_cpuacct_base, NULL) == -1) {
@@ -1039,87 +973,6 @@ static int discovery_is_cgroup_duplicate(struct cgroup *cg) {
973 return 0;
974 }
975
1042 -// ----------------------------------------------------------------------------
1043 -// ebpf shared memory
1044 -
1045 -static void netdata_cgroup_ebpf_set_values(size_t length)
1046 -{
1047 - sem_wait(shm_mutex_cgroup_ebpf);
1048 -
1049 - shm_cgroup_ebpf.header->cgroup_max = cgroup_root_max;
1050 - shm_cgroup_ebpf.header->systemd_enabled = CONFIG_BOOLEAN_YES;
1051 - shm_cgroup_ebpf.header->body_length = length;
1052 -
1053 - sem_post(shm_mutex_cgroup_ebpf);
1054 -}
1055 -
1056 -static void netdata_cgroup_ebpf_initialize_shm()
1057 -{
1058 - // Unlink any existing shared memory and semaphore to start fresh.
1059 - // This prevents truncating memory that another process might be using.
1060 - // Existing mappings in other processes remain valid until they unmap.
1061 - (void) shm_unlink(NETDATA_SHARED_MEMORY_EBPF_CGROUP_NAME);
1062 - (void) sem_unlink(NETDATA_NAMED_SEMAPHORE_EBPF_CGROUP_NAME);
1063 -
1064 - shm_fd_cgroup_ebpf = shm_open(NETDATA_SHARED_MEMORY_EBPF_CGROUP_NAME, O_CREAT | O_RDWR, 0660);
1065 - if (shm_fd_cgroup_ebpf < 0) {
1066 - collector_error("Cannot initialize shared memory used by cgroup and eBPF, integration won't happen.");
1067 - return;
1068 - }
1069 -
1070 - size_t length = sizeof(netdata_ebpf_cgroup_shm_header_t) + cgroup_root_max * sizeof(netdata_ebpf_cgroup_shm_body_t);
1071 - if (ftruncate(shm_fd_cgroup_ebpf, length)) {
1072 - collector_error("Cannot set size for shared memory.");
1073 - goto end_init_shm;
1074 - }
1075 -
1076 - shm_cgroup_ebpf.header = (netdata_ebpf_cgroup_shm_header_t *)
1077 - nd_mmap(NULL, length, PROT_READ | PROT_WRITE, MAP_SHARED, shm_fd_cgroup_ebpf, 0);
1078 -
1079 - if (unlikely(MAP_FAILED == shm_cgroup_ebpf.header)) {
1080 - shm_cgroup_ebpf.header = NULL;
1081 - collector_error("Cannot map shared memory used between cgroup and eBPF, integration won't happen");
1082 - goto end_init_shm;
1083 - }
1084 - shm_cgroup_ebpf.body = (netdata_ebpf_cgroup_shm_body_t *) ((char *)shm_cgroup_ebpf.header +
1085 - sizeof(netdata_ebpf_cgroup_shm_header_t));
1086 -
1087 - shm_mutex_cgroup_ebpf = sem_open(NETDATA_NAMED_SEMAPHORE_EBPF_CGROUP_NAME, O_CREAT,
1088 - S_IRUSR | S_IWUSR | S_IRGRP | S_IWGRP | S_IROTH | S_IWOTH, 1);
1089 -
1090 - if (shm_mutex_cgroup_ebpf != SEM_FAILED) {
1091 - netdata_cgroup_ebpf_set_values(length);
1092 - return;
1093 - }
1094 -
1095 - collector_error("Cannot create semaphore, integration between eBPF and cgroup won't happen");
1096 - nd_munmap(shm_cgroup_ebpf.header, length);
1097 - shm_cgroup_ebpf.header = NULL;
1098 -
1099 - end_init_shm:
1100 - close(shm_fd_cgroup_ebpf);
1101 - shm_fd_cgroup_ebpf = -1;
1102 - (void) shm_unlink(NETDATA_SHARED_MEMORY_EBPF_CGROUP_NAME);
1103 -}
1104 -
1105 -static void cgroup_cleanup_ebpf_integration()
1106 -{
1107 - if (shm_mutex_cgroup_ebpf != SEM_FAILED) {
1108 - sem_close(shm_mutex_cgroup_ebpf);
1109 - (void) sem_unlink(NETDATA_NAMED_SEMAPHORE_EBPF_CGROUP_NAME);
1110 - }
1111 -
1112 - if (shm_cgroup_ebpf.header) {
1113 - shm_cgroup_ebpf.header->cgroup_root_count = 0;
1114 - nd_munmap(shm_cgroup_ebpf.header, shm_cgroup_ebpf.header->body_length);
1115 - }
1116 -
1117 - if (shm_fd_cgroup_ebpf > 0) {
1118 - close(shm_fd_cgroup_ebpf);
1119 - }
1120 - (void) shm_unlink(NETDATA_SHARED_MEMORY_EBPF_CGROUP_NAME);
1121 -}
1122 -
976 // ----------------------------------------------------------------------------
977 // cgroup network interfaces
978
@@ -1296,8 +1149,7 @@ static inline void discovery_find_all_cgroups() {
1149
1150 netdata_mutex_unlock(&cgroup_root_mutex);
1151
1299 - worker_is_busy(WORKER_DISCOVERY_SHARE);
1300 - discovery_share_cgroups_with_ebpf();
1152 + // cgroup metadata is now served on-demand via netipc (cgroup-netipc.c)
1153
1154 netdata_log_debug(D_CGROUP, "done searching for cgroups");
1155 }
@@ -1317,7 +1169,6 @@ void cgroup_discovery_worker(void *ptr)
1169 worker_register_job_name(WORKER_DISCOVERY_UPDATE, "update");
1170 worker_register_job_name(WORKER_DISCOVERY_CLEANUP, "cleanup");
1171 worker_register_job_name(WORKER_DISCOVERY_COPY, "copy");
1320 - worker_register_job_name(WORKER_DISCOVERY_SHARE, "share");
1172 worker_register_job_name(WORKER_DISCOVERY_LOCK, "lock");
1173
1174 entrypoint_parent_process_comm = simple_pattern_create(
@@ -1328,7 +1179,7 @@ void cgroup_discovery_worker(void *ptr)
1179
1180 service_register(NULL, NULL, NULL);
1181
1331 - netdata_cgroup_ebpf_initialize_shm();
1182 + cgroup_netipc_init();
1183
1184 while (service_running(SERVICE_COLLECTORS)) {
1185 worker_is_idle();
@@ -1353,7 +1204,7 @@ void cgroup_discovery_worker(void *ptr)
1204 netdata_mutex_unlock(&cgroup_root_mutex);
1205
1206 collector_info("discovery thread stopped");
1356 - cgroup_cleanup_ebpf_integration();
1207 + cgroup_netipc_cleanup();
1208 worker_unregister();
1209 service_exits();
1210 __atomic_store_n(&discovery_thread.exited, 1, __ATOMIC_RELEASE);
src/collectors/cgroups.plugin/cgroup-netipc.c new
+189
@@ -0,0 +1,189 @@
1 +// SPDX-License-Identifier: GPL-3.0-or-later
2 +
3 +#include "cgroup-internals.h"
4 +#include "libnetdata/netipc/netipc_netdata.h"
5 +
6 +#ifdef OS_LINUX
7 +
8 +#define CGROUP_NETIPC_SERVICE_NAME "cgroups-snapshot"
9 +#define CGROUP_NETIPC_WORKER_COUNT 2
10 +
11 +static nipc_managed_server_t cgroup_netipc_server;
12 +static ND_THREAD *cgroup_netipc_thread = NULL;
13 +
14 +// find the cgroup.procs path for a given cgroup (v1 hierarchy)
15 +// writes into path_buf, returns true if found
16 +static bool cgroup_find_procs_path_v1(char *path_buf, size_t path_buf_size, const char *cg_id) {
17 + struct stat buf;
18 +
19 + snprintfz(path_buf, path_buf_size - 1, "%s%s/cgroup.procs", cgroup_cpuset_base, cg_id);
20 + if (stat(path_buf, &buf) == 0)
21 + return true;
22 +
23 + snprintfz(path_buf, path_buf_size - 1, "%s%s/cgroup.procs", cgroup_blkio_base, cg_id);
24 + if (stat(path_buf, &buf) == 0)
25 + return true;
26 +
27 + snprintfz(path_buf, path_buf_size - 1, "%s%s/cgroup.procs", cgroup_memory_base, cg_id);
28 + if (stat(path_buf, &buf) == 0)
29 + return true;
30 +
31 + path_buf[0] = '\0';
32 + return false;
33 +}
34 +
35 +// handler callback invoked by netipc worker threads when a client requests a snapshot
36 +static bool cgroups_snapshot_handler(void *user __maybe_unused,
37 + const nipc_cgroups_req_t *request __maybe_unused,
38 + nipc_cgroups_builder_t *builder) {
39 + static uint64_t generation = 0;
40 + static uint64_t last_logged_zero_generation = 0;
41 + static uint64_t last_logged_truncated_generation = 0;
42 + uint64_t snapshot_generation;
43 + char name_buf[256];
44 + char path_buf[FILENAME_MAX + 1];
45 +
46 + netdata_mutex_lock(&cgroup_root_mutex);
47 +
48 + // set snapshot header — systemd is always enabled in this codebase
49 + snapshot_generation = ++generation;
50 + nipc_cgroups_builder_set_header(builder, CONFIG_BOOLEAN_YES, snapshot_generation);
51 +
52 + struct cgroup *cg;
53 + int count;
54 + int enabled_count = 0;
55 + bool truncated = false;
56 + for (cg = cgroup_root, count = 0; cg && count < cgroup_root_max; cg = cg->next, count++) {
57 + const char *prefix = is_cgroup_systemd_service(cg)
58 + ? services_chart_id_prefix
59 + : cgroup_chart_id_prefix;
60 +
61 + snprintfz(name_buf, sizeof(name_buf) - 1, "%s%s", prefix, cg->chart_id);
62 + uint32_t hash = simple_hash(name_buf);
63 + uint32_t options = cg->options;
64 + uint32_t enabled = cg->enabled;
65 +
66 + // find the cgroup.procs path
67 + if (cgroup_use_unified_cgroups) {
68 + struct stat buf;
69 + snprintfz(path_buf, FILENAME_MAX, "%s%s/cgroup.procs", cgroup_unified_base, cg->id);
70 + if (stat(path_buf, &buf) == -1) {
71 + path_buf[0] = '\0';
72 + enabled = 0;
73 + }
74 + } else {
75 + if (!cgroup_find_procs_path_v1(path_buf, sizeof(path_buf), cg->id))
76 + enabled = 0;
77 + }
78 +
79 + if (enabled)
80 + enabled_count++;
81 +
82 + nipc_error_t err = nipc_cgroups_builder_add(
83 + builder, hash, options, enabled,
84 + name_buf, (uint32_t)strlen(name_buf),
85 + path_buf, (uint32_t)strlen(path_buf));
86 +
87 + if (err == NIPC_ERR_OVERFLOW) {
88 + truncated = true;
89 + break; // buffer full — send what we have
90 + }
91 + }
92 +
93 + netdata_mutex_unlock(&cgroup_root_mutex);
94 +
95 + bool log_zero_generation = false;
96 + bool log_truncated_generation = false;
97 +
98 + netdata_mutex_lock(&cgroup_root_mutex);
99 +
100 + if (count == 0 && last_logged_zero_generation != snapshot_generation) {
101 + last_logged_zero_generation = snapshot_generation;
102 + log_zero_generation = true;
103 + }
104 +
105 + if (truncated && last_logged_truncated_generation != snapshot_generation) {
106 + last_logged_truncated_generation = snapshot_generation;
107 + log_truncated_generation = true;
108 + }
109 +
110 + netdata_mutex_unlock(&cgroup_root_mutex);
111 +
112 + if (log_zero_generation) {
113 + collector_info("CGROUP: netipc snapshot generation=%llu returned zero items",
114 + (unsigned long long)snapshot_generation);
115 + }
116 +
117 + if (log_truncated_generation) {
118 + collector_error(
119 + "CGROUP: netipc snapshot generation=%llu truncated after %d items (%d enabled) due to response size limits",
120 + (unsigned long long)snapshot_generation,
121 + count,
122 + enabled_count);
123 + }
124 +
125 + return true;
126 +}
127 +
128 +// thread entry point for the netipc accept loop
129 +static void cgroup_netipc_server_thread(void *arg) {
130 + nipc_server_run((nipc_managed_server_t *)arg);
131 +}
132 +
133 +void cgroup_netipc_init(void) {
134 + uint64_t auth = netipc_auth_token();
135 +
136 + nipc_server_config_t config = {
137 + .supported_profiles = NIPC_PROFILE_BASELINE | NIPC_PROFILE_SHM_HYBRID | NIPC_PROFILE_SHM_FUTEX,
138 + .preferred_profiles = NIPC_PROFILE_SHM_FUTEX,
139 + .auth_token = auth,
140 + };
141 +
142 + nipc_cgroups_service_handler_t handler = {
143 + .handle = cgroups_snapshot_handler,
144 + .snapshot_max_items = 0, // auto-estimate from negotiated limits
145 + .user = NULL,
146 + };
147 +
148 + nipc_error_t err = nipc_server_init_typed(
149 + &cgroup_netipc_server,
150 + os_run_dir(true),
151 + CGROUP_NETIPC_SERVICE_NAME,
152 + &config,
153 + CGROUP_NETIPC_WORKER_COUNT,
154 + &handler);
155 +
156 + if (err != NIPC_OK) {
157 + collector_error("CGROUP: netipc server init failed (error %u), IPC sharing disabled", (unsigned int)err);
158 + return;
159 + }
160 +
161 + cgroup_netipc_thread = nd_thread_create(
162 + "P[cgroupsipc]", NETDATA_THREAD_OPTION_DONT_LOG_STARTUP,
163 + cgroup_netipc_server_thread, &cgroup_netipc_server);
164 +
165 + if (!cgroup_netipc_thread) {
166 + collector_error("CGROUP: failed to create netipc server thread");
167 + nipc_server_destroy(&cgroup_netipc_server);
168 + return;
169 + }
170 +
171 + collector_info("CGROUP: netipc server started on '%s/%s.sock'",
172 + os_run_dir(true), CGROUP_NETIPC_SERVICE_NAME);
173 +}
174 +
175 +void cgroup_netipc_cleanup(void) {
176 + if (!cgroup_netipc_thread)
177 + return;
178 +
179 + nipc_server_stop(&cgroup_netipc_server);
180 + nd_thread_join(cgroup_netipc_thread);
181 + cgroup_netipc_thread = NULL;
182 +
183 + nipc_server_drain(&cgroup_netipc_server, 5000);
184 + nipc_server_destroy(&cgroup_netipc_server);
185 +
186 + collector_info("CGROUP: netipc server stopped");
187 +}
188 +
189 +#endif // OS_LINUX
src/collectors/cgroups.plugin/cgroup-netipc.h new
+14
@@ -0,0 +1,14 @@
1 +// SPDX-License-Identifier: GPL-3.0-or-later
2 +
3 +#ifndef NETDATA_CGROUP_NETIPC_H
4 +#define NETDATA_CGROUP_NETIPC_H 1
5 +
6 +#ifdef OS_LINUX
7 +void cgroup_netipc_init(void);
8 +void cgroup_netipc_cleanup(void);
9 +#else
10 +static inline void cgroup_netipc_init(void) {}
11 +static inline void cgroup_netipc_cleanup(void) {}
12 +#endif
13 +
14 +#endif // NETDATA_CGROUP_NETIPC_H
src/collectors/cgroups.plugin/sys_fs_cgroup.h
+1 -27
@@ -14,33 +14,7 @@
14 #define CGROUP_OPTIONS_IS_UNIFIED (1 << 2)
15 #define CGROUP_OPTIONS_DISABLED_EXCLUDED (1 << 3)
16
17 -typedef struct netdata_ebpf_cgroup_shm_header {
18 - int cgroup_root_count;
19 - int cgroup_max;
20 - int systemd_enabled;
21 - int __pad;
22 - size_t body_length;
23 -} netdata_ebpf_cgroup_shm_header_t;
24 -
25 -#define CGROUP_EBPF_NAME_SHARED_LENGTH 256
26 -
27 -typedef struct netdata_ebpf_cgroup_shm_body {
28 - // Considering what is exposed in this link https://en.wikipedia.org/wiki/Comparison_of_file_systems#Limits
29 - // this length is enough to store what we want.
30 - char name[CGROUP_EBPF_NAME_SHARED_LENGTH];
31 - uint32_t hash;
32 - uint32_t options;
33 - int enabled;
34 - char path[FILENAME_MAX + 1];
35 -} netdata_ebpf_cgroup_shm_body_t;
36 -
37 -typedef struct netdata_ebpf_cgroup_shm {
38 - netdata_ebpf_cgroup_shm_header_t *header;
39 - netdata_ebpf_cgroup_shm_body_t *body;
40 -} netdata_ebpf_cgroup_shm_t;
41 -
42 -#define NETDATA_SHARED_MEMORY_EBPF_CGROUP_NAME "netdata_shm_cgroup_ebpf"
43 -#define NETDATA_NAMED_SEMAPHORE_EBPF_CGROUP_NAME "/netdata_sem_cgroup_ebpf"
17 +// legacy SHM structs removed — cgroup metadata is shared via netipc (cgroup-netipc.c)
18
19 #include "../proc.plugin/plugin_proc.h"
20
src/collectors/ebpf.plugin/ebpf.c
+4 -10
@@ -852,10 +852,9 @@ ebpf_sync_syscalls_t local_syscalls[] = {
852 #endif
853 .sync_maps = NULL}};
854
855 -// Link with cgroup.plugin
856 -netdata_ebpf_cgroup_shm_t shm_ebpf_cgroup = {NULL, NULL};
857 -int shm_fd_ebpf_cgroup = -1;
858 -sem_t *shm_sem_ebpf_cgroup = SEM_FAILED;
855 +// cgroup integration via netipc
856 +_Atomic int ebpf_cgroup_systemd_enabled = 0;
857 +_Atomic int ebpf_cgroup_integration_active = 0;
858 netdata_mutex_t mutex_cgroup_shm;
859
860 //Network viewer
@@ -1042,12 +1041,7 @@ static void ebpf_exit()
1041 ebpf_check_before2go();
1042 ebpf_pre_exit_check_done = true;
1043 }
1045 - netdata_mutex_lock(&mutex_cgroup_shm);
1046 - if (shm_ebpf_cgroup.header) {
1047 - ebpf_unmap_cgroup_shared_memory();
1048 - shm_unlink(NETDATA_SHARED_MEMORY_EBPF_CGROUP_NAME);
1049 - }
1050 - netdata_mutex_unlock(&mutex_cgroup_shm);
1044 + ebpf_cgroup_cache_cleanup();
1045 netdata_integration_cleanup_shm();
1046
1047 exit(0);
src/collectors/ebpf.plugin/ebpf.h
+35 -3
@@ -6,6 +6,7 @@
6 #ifndef __FreeBSD__
7 #include <linux/perf_event.h>
8 #endif
9 +#include <stdatomic.h>
10 #include <stdint.h>
11 #include <errno.h>
12 #include <signal.h>
@@ -23,6 +24,7 @@
24 #include "libbpf_api/ebpf.h"
25
26 #include "collectors/cgroups.plugin/sys_fs_cgroup.h"
27 +#include "libnetdata/netipc/netipc_netdata.h"
28
29 #include "ebpf_apps.h"
30 #include "ebpf_functions.h"
@@ -269,9 +271,9 @@ extern struct ebpf_pid_stat *ebpf_root_of_pids;
271 extern ebpf_cgroup_target_t *ebpf_cgroup_pids;
272 extern char *ebpf_algorithms[];
273 extern struct config collector_config;
272 -extern netdata_ebpf_cgroup_shm_t shm_ebpf_cgroup;
273 -extern int shm_fd_ebpf_cgroup;
274 -extern sem_t *shm_sem_ebpf_cgroup;
274 +extern _Atomic int ebpf_cgroup_systemd_enabled;
275 +extern _Atomic int ebpf_cgroup_integration_active;
276 +extern _Atomic int send_cgroup_chart;
277 extern netdata_mutex_t mutex_cgroup_shm;
278 extern size_t ebpf_all_pids_count;
279 extern ebpf_plugin_stats_t plugin_statistics;
@@ -328,6 +330,36 @@ extern volatile sig_atomic_t ebpf_stop_signal;
330 extern bool ebpf_plugin_exit;
331 extern uint64_t collect_pids;
332
333 +static inline void ebpf_cgroup_systemd_enabled_set(int value)
334 +{
335 + atomic_store_explicit(&ebpf_cgroup_systemd_enabled, value, memory_order_release);
336 +}
337 +
338 +static inline int ebpf_cgroup_systemd_enabled_get(void)
339 +{
340 + return atomic_load_explicit(&ebpf_cgroup_systemd_enabled, memory_order_acquire);
341 +}
342 +
343 +static inline void ebpf_cgroup_integration_active_set(int value)
344 +{
345 + atomic_store_explicit(&ebpf_cgroup_integration_active, value, memory_order_release);
346 +}
347 +
348 +static inline int ebpf_cgroup_integration_active_get(void)
349 +{
350 + return atomic_load_explicit(&ebpf_cgroup_integration_active, memory_order_acquire);
351 +}
352 +
353 +static inline void ebpf_send_cgroup_chart_set(int value)
354 +{
355 + atomic_store_explicit(&send_cgroup_chart, value, memory_order_release);
356 +}
357 +
358 +static inline int ebpf_send_cgroup_chart_get(void)
359 +{
360 + return atomic_load_explicit(&send_cgroup_chart, memory_order_acquire);
361 +}
362 +
363 static inline bool ebpf_plugin_stop(void)
364 {
365 return __atomic_load_n(&ebpf_plugin_exit, __ATOMIC_ACQUIRE) ||
src/collectors/ebpf.plugin/ebpf_cachestat.c
+4 -4
@@ -1055,7 +1055,7 @@ void ebpf_read_cachestat_thread(void *ptr)
1055 break;
1056 }
1057
1058 - if (cgroups && shm_ebpf_cgroup.header)
1058 + if (cgroups && ebpf_cgroup_integration_active_get())
1059 ebpf_update_cachestat_cgroup();
1060 if (sem_post(shm_mutex_ebpf_integration)) {
1061 netdata_log_error("CACHESTAT: Failed to post semaphore.");
@@ -1550,8 +1550,8 @@ void ebpf_cachestat_send_cgroup_data(int update_every)
1550 ebpf_cgroup_target_t *ect;
1551 ebpf_cachestat_calc_chart_values();
1552
1553 - if (shm_ebpf_cgroup.header->systemd_enabled) {
1554 - if (send_cgroup_chart) {
1553 + if (ebpf_cgroup_systemd_enabled_get()) {
1554 + if (ebpf_send_cgroup_chart_get()) {
1555 ebpf_create_systemd_cachestat_charts(update_every);
1556 }
1557
@@ -1629,7 +1629,7 @@ static void cachestat_collector(ebpf_module_t *em)
1629 break;
1630 }
1631
1632 - if (cgroups && shm_ebpf_cgroup.header)
1632 + if (cgroups && ebpf_cgroup_integration_active_get())
1633 ebpf_cachestat_send_cgroup_data(update_every);
1634
1635 netdata_mutex_unlock(&lock);
src/collectors/ebpf.plugin/ebpf_cgroup.c
+263 -136
@@ -7,112 +7,88 @@
7 #include "libbpf_api/ebpf_library.h"
8
9 ebpf_cgroup_target_t *ebpf_cgroup_pids = NULL;
10 -static void *ebpf_mapped_memory = NULL;
11 -int send_cgroup_chart = 0;
10 +_Atomic int send_cgroup_chart = 0;
11
13 -// --------------------------------------------------------------------------------------------------------------------
14 -// Map shared memory
12 +#ifdef OS_LINUX
13 +static nipc_cgroups_cache_t ebpf_cgroup_cache;
14 +static bool ebpf_cgroup_cache_initialized = false;
15
16 -/**
17 - * Map Shared Memory locally
18 - *
19 - * Map the shared memory for current process
20 - *
21 - * @param fd file descriptor returned after shm_open was called.
22 - * @param length length of the shared memory
23 - *
24 - * @return It returns a pointer to the region mapped on success and MAP_FAILED otherwise.
25 - */
26 -static inline void *ebpf_cgroup_map_shm_locally(int fd, size_t length)
16 +static const char *ebpf_netipc_client_state_name(nipc_client_state_t state)
17 {
28 - void *value = nd_mmap(NULL, length, PROT_READ | PROT_WRITE, MAP_SHARED, fd, 0);
29 - if (value == MAP_FAILED) {
30 - netdata_log_error(
31 - "Cannot map shared memory used between eBPF and cgroup, integration between processes won't happen");
32 - close(shm_fd_ebpf_cgroup);
33 - shm_fd_ebpf_cgroup = -1;
34 - shm_unlink(NETDATA_SHARED_MEMORY_EBPF_CGROUP_NAME);
18 + switch (state) {
19 + case NIPC_CLIENT_DISCONNECTED:
20 + return "disconnected";
21 + case NIPC_CLIENT_CONNECTING:
22 + return "connecting";
23 + case NIPC_CLIENT_READY:
24 + return "ready";
25 + case NIPC_CLIENT_NOT_FOUND:
26 + return "not-found";
27 + case NIPC_CLIENT_AUTH_FAILED:
28 + return "auth-failed";
29 + case NIPC_CLIENT_INCOMPATIBLE:
30 + return "incompatible";
31 + case NIPC_CLIENT_BROKEN:
32 + return "broken";
33 + default:
34 + return "unknown";
35 }
36 -
37 - return value;
36 }
37
40 -/**
41 - * Unmap Shared Memory
42 - *
43 - * Unmap shared memory used to integrate eBPF and cgroup plugin
44 - */
45 -void ebpf_unmap_cgroup_shared_memory()
38 +static const char *ebpf_netipc_profile_name(uint32_t profile)
39 {
47 - nd_munmap(ebpf_mapped_memory, shm_ebpf_cgroup.header->body_length);
40 + switch (profile) {
41 + case NIPC_PROFILE_BASELINE:
42 + return "baseline";
43 + case NIPC_PROFILE_SHM_HYBRID:
44 + return "shm-hybrid";
45 + case NIPC_PROFILE_SHM_FUTEX:
46 + return "shm-futex";
47 + default:
48 + return "none";
49 + }
50 }
51
50 -/**
51 - * Map cgroup shared memory
52 - *
53 - * Map cgroup shared memory from cgroup to plugin
54 - */
55 -void ebpf_map_cgroup_shared_memory()
52 +static void ebpf_log_cgroup_transport_state(uint64_t generation, uint32_t count, uint32_t enabled_count)
53 {
57 - static int limit_try = 0;
58 - static time_t next_try = 0;
59 -
60 - if (shm_ebpf_cgroup.header || limit_try > NETDATA_EBPF_CGROUP_MAX_TRIES)
61 - return;
62 -
63 - time_t curr_time = time(NULL);
64 - if (curr_time < next_try)
65 - return;
66 -
67 - limit_try++;
68 - next_try = curr_time + NETDATA_EBPF_CGROUP_NEXT_TRY_SEC;
69 -
70 - if (shm_fd_ebpf_cgroup < 0) {
71 - shm_fd_ebpf_cgroup = shm_open(NETDATA_SHARED_MEMORY_EBPF_CGROUP_NAME, O_RDWR, 0660);
72 - if (shm_fd_ebpf_cgroup < 0) {
73 - if (limit_try == NETDATA_EBPF_CGROUP_MAX_TRIES)
74 - netdata_log_error("Shared memory was not initialized, integration between processes won't happen.");
75 -
76 - return;
77 - }
78 - }
79 -
80 - // Map only header
81 - void *mapped = (netdata_ebpf_cgroup_shm_header_t *)ebpf_cgroup_map_shm_locally(
82 - shm_fd_ebpf_cgroup, sizeof(netdata_ebpf_cgroup_shm_header_t));
83 - if (unlikely(mapped == MAP_FAILED)) {
84 - return;
85 - }
86 - netdata_ebpf_cgroup_shm_header_t *header = mapped;
87 -
88 - size_t length = header->body_length;
89 -
90 - nd_munmap(header, sizeof(netdata_ebpf_cgroup_shm_header_t));
91 -
92 - if (length <= ((sizeof(netdata_ebpf_cgroup_shm_header_t) + sizeof(netdata_ebpf_cgroup_shm_body_t)))) {
54 + static nipc_client_state_t previous_state = NIPC_CLIENT_DISCONNECTED;
55 + static bool previous_session_valid = false;
56 + static uint32_t previous_profile = 0;
57 + static bool previous_using_shm = false;
58 + static uint64_t previous_session_id = 0;
59 +
60 + nipc_client_state_t state = ebpf_cgroup_cache.client.state;
61 + bool session_valid = ebpf_cgroup_cache.client.session_valid;
62 + uint32_t profile = session_valid ? ebpf_cgroup_cache.client.session.selected_profile : 0;
63 + bool using_shm = session_valid && ebpf_cgroup_cache.client.shm != NULL;
64 + uint64_t session_id = session_valid ? ebpf_cgroup_cache.client.session.session_id : 0;
65 +
66 + if (previous_state == state &&
67 + previous_session_valid == session_valid &&
68 + previous_profile == profile &&
69 + previous_using_shm == using_shm &&
70 + previous_session_id == session_id)
71 return;
94 - }
72
96 - ebpf_mapped_memory = (void *)ebpf_cgroup_map_shm_locally(shm_fd_ebpf_cgroup, length);
97 - if (unlikely(ebpf_mapped_memory == MAP_FAILED)) {
98 - return;
99 - }
100 - shm_ebpf_cgroup.header = ebpf_mapped_memory;
101 - shm_ebpf_cgroup.body = ebpf_mapped_memory + sizeof(netdata_ebpf_cgroup_shm_header_t);
102 -
103 - shm_sem_ebpf_cgroup = sem_open(NETDATA_NAMED_SEMAPHORE_EBPF_CGROUP_NAME, O_CREAT, 0660, 1);
104 -
105 - if (shm_sem_ebpf_cgroup == SEM_FAILED) {
106 - netdata_log_error("Cannot create semaphore, integration between eBPF and cgroup won't happen");
107 - limit_try = NETDATA_EBPF_CGROUP_MAX_TRIES + 1;
108 - nd_munmap(ebpf_mapped_memory, length);
109 - shm_ebpf_cgroup.header = NULL;
110 - shm_ebpf_cgroup.body = NULL;
111 - close(shm_fd_ebpf_cgroup);
112 - shm_fd_ebpf_cgroup = -1;
113 - shm_unlink(NETDATA_SHARED_MEMORY_EBPF_CGROUP_NAME);
114 - }
73 + collector_info(
74 + "EBPF CGROUP: netipc transport state=%s session_valid=%d session=%016llx "
75 + "selected_profile=%s data_plane=%s generation=%llu items=%u enabled=%u",
76 + ebpf_netipc_client_state_name(state),
77 + session_valid,
78 + (unsigned long long)session_id,
79 + ebpf_netipc_profile_name(profile),
80 + using_shm ? "shm" : (session_valid ? "baseline" : "none"),
81 + (unsigned long long)generation,
82 + count,
83 + enabled_count);
84 +
85 + previous_state = state;
86 + previous_session_valid = session_valid;
87 + previous_profile = profile;
88 + previous_using_shm = using_shm;
89 + previous_session_id = session_id;
90 }
91 +#endif
92
93 // --------------------------------------------------------------------------------------------------------------------
94 // Close and Cleanup
@@ -163,22 +139,44 @@ static void ebpf_remove_cgroup_target_update_list()
139 }
140 }
141
142 +static size_t ebpf_count_cgroup_targets_unsafe(void)
143 +{
144 + size_t count = 0;
145 +
146 + for (ebpf_cgroup_target_t *ect = ebpf_cgroup_pids; ect; ect = ect->next)
147 + count++;
148 +
149 + return count;
150 +}
151 +
152 +static size_t ebpf_count_cgroup_pids_unsafe(void)
153 +{
154 + size_t count = 0;
155 +
156 + for (ebpf_cgroup_target_t *ect = ebpf_cgroup_pids; ect; ect = ect->next) {
157 + for (struct pid_on_target2 *pt = ect->pids; pt; pt = pt->next)
158 + count++;
159 + }
160 +
161 + return count;
162 +}
163 +
164 // --------------------------------------------------------------------------------------------------------------------
165 // Fill variables
166
167 /**
168 * Set Target Data
169 *
172 - * Set local variable values according shared memory information.
170 + * Set local variable values from a netipc cache item.
171 *
174 - * @param out local output variable.
175 - * @param ptr input from shared memory.
172 + * @param out local output variable.
173 + * @param item netipc cache item.
174 */
177 -static inline void ebpf_cgroup_set_target_data(ebpf_cgroup_target_t *out, netdata_ebpf_cgroup_shm_body_t *ptr)
175 +static inline void ebpf_cgroup_set_target_data(ebpf_cgroup_target_t *out, const nipc_cgroups_cache_item_t *item)
176 {
179 - out->hash = ptr->hash;
180 - snprintfz(out->name, 255, "%s", ptr->name);
181 - out->systemd = ptr->options & CGROUP_OPTIONS_SYSTEM_SLICE_SERVICE;
177 + out->hash = item->hash;
178 + snprintfz(out->name, 255, "%s", item->name);
179 + out->systemd = item->options & CGROUP_OPTIONS_SYSTEM_SLICE_SERVICE;
180 out->updated = 1;
181 }
182
@@ -187,21 +185,21 @@ static inline void ebpf_cgroup_set_target_data(ebpf_cgroup_target_t *out, netdat
185 *
186 * Find the structure inside the link list or allocate and link when it is not present.
187 *
190 - * @param ptr Input from shared memory.
188 + * @param item netipc cache item.
189 *
190 * @return It returns a pointer for the structure associated with the input.
191 */
194 -static ebpf_cgroup_target_t *ebpf_cgroup_find_or_create(netdata_ebpf_cgroup_shm_body_t *ptr)
192 +static ebpf_cgroup_target_t *ebpf_cgroup_find_or_create(const nipc_cgroups_cache_item_t *item)
193 {
194 for (ebpf_cgroup_target_t *ect = ebpf_cgroup_pids; ect; ect = ect->next) {
197 - if (ect->hash == ptr->hash && !strcmp(ect->name, ptr->name)) {
195 + if (ect->hash == item->hash && !strcmp(ect->name, item->name)) {
196 ect->updated = 1;
197 return ect;
198 }
199 }
200
201 ebpf_cgroup_target_t *new_ect = callocz(1, sizeof(*new_ect));
204 - ebpf_cgroup_set_target_data(new_ect, ptr);
202 + ebpf_cgroup_set_target_data(new_ect, item);
203 new_ect->next = ebpf_cgroup_pids;
204 ebpf_cgroup_pids = new_ect;
205
@@ -216,7 +214,7 @@ static ebpf_cgroup_target_t *ebpf_cgroup_find_or_create(netdata_ebpf_cgroup_shm_
214 * @param ect cgroup structure where pids will be stored
215 * @param path file with PIDs associated to cgroup.
216 */
219 -static void ebpf_update_pid_link_list(ebpf_cgroup_target_t *ect, char *path)
217 +static void ebpf_update_pid_link_list(ebpf_cgroup_target_t *ect, const char *path)
218 {
219 procfile *ff = procfile_open_no_log(path, " \t:", PROCFILE_FLAG_DEFAULT);
220 if (!ff)
@@ -278,45 +276,174 @@ void ebpf_reset_updated_var()
276 }
277
278 /**
281 - * Parse cgroup shared memory
279 + * Initialize netipc cgroup cache
280 + *
281 + * Connect to the cgroups-snapshot service via netipc.
282 + */
283 +static void ebpf_cgroup_cache_init(void)
284 +{
285 +#ifdef OS_LINUX
286 + if (ebpf_cgroup_cache_initialized)
287 + return;
288 +
289 + uint64_t auth = netipc_auth_token();
290 +
291 + nipc_client_config_t config = {
292 + .supported_profiles = NIPC_PROFILE_BASELINE | NIPC_PROFILE_SHM_HYBRID | NIPC_PROFILE_SHM_FUTEX,
293 + .preferred_profiles = NIPC_PROFILE_SHM_FUTEX,
294 + .auth_token = auth,
295 + };
296 +
297 + nipc_cgroups_cache_init(&ebpf_cgroup_cache,
298 + os_run_dir(false),
299 + "cgroups-snapshot",
300 + &config);
301 +
302 + ebpf_cgroup_cache_initialized = true;
303 +#endif
304 +}
305 +
306 +/**
307 + * Close the netipc cgroup cache and release resources.
308 + */
309 +void ebpf_cgroup_cache_cleanup(void)
310 +{
311 +#ifdef OS_LINUX
312 + if (ebpf_cgroup_cache_initialized) {
313 + nipc_cgroups_cache_close(&ebpf_cgroup_cache);
314 + ebpf_cgroup_cache_initialized = false;
315 + }
316 +#endif
317 +}
318 +
319 +/**
320 + * Refresh cgroup data from netipc cache
321 *
283 - * This function is responsible to copy necessary data from shared memory to local memory.
322 + * Replaces the legacy SHM parse function.
323 */
285 -void ebpf_parse_cgroup_shm_data()
324 +static void ebpf_parse_cgroup_netipc_data(void)
325 {
287 - static int previous = 0;
288 - if (!shm_ebpf_cgroup.header || shm_sem_ebpf_cgroup == SEM_FAILED)
326 +#ifdef OS_LINUX
327 + static uint32_t previous_count = 0;
328 + static uint32_t previous_enabled_count = 0;
329 + static size_t previous_imported_targets = 0;
330 + static size_t previous_total_pids = 0;
331 + static int previous_integration_active = -1;
332 + static int previous_systemd_enabled = -1;
333 +
334 + if (!ebpf_cgroup_cache_initialized)
335 return;
336
291 - sem_wait(shm_sem_ebpf_cgroup);
292 - int i, end = shm_ebpf_cgroup.header->cgroup_root_count;
293 - if (end <= 0) {
294 - sem_post(shm_sem_ebpf_cgroup);
337 + static int refresh_fail_count = 0;
338 + if (!nipc_cgroups_cache_refresh(&ebpf_cgroup_cache)) {
339 + if (++refresh_fail_count % 10 == 1)
340 + collector_error("EBPF CGROUP: netipc refresh failed (%d consecutive failures)", refresh_fail_count);
341 + return;
342 + }
343 + refresh_fail_count = 0;
344 +
345 + uint32_t last_count = previous_count;
346 + uint32_t count = ebpf_cgroup_cache.item_count;
347 + uint64_t generation = ebpf_cgroup_cache.generation;
348 + uint32_t enabled_count = 0;
349 +
350 + // Publish the latest cgroup state before collectors read the lock-free flags.
351 + int systemd_enabled = (int)ebpf_cgroup_cache.systemd_enabled;
352 + int integration_active = (count > 0) ? 1 : 0;
353 + ebpf_cgroup_systemd_enabled_set(systemd_enabled);
354 + ebpf_cgroup_integration_active_set(integration_active);
355 +
356 + for (uint32_t i = 0; i < count; i++) {
357 + const nipc_cgroups_cache_item_t *item = &ebpf_cgroup_cache.items[i];
358 + if (item->enabled)
359 + enabled_count++;
360 + }
361 +
362 + ebpf_log_cgroup_transport_state(generation, count, enabled_count);
363 +
364 + // nothing to process; preserve existing targets rather than wiping them.
365 + // reset previous_count so the next non-zero snapshot triggers send_cgroup_chart=1,
366 + // ensuring systemd charts are recreated after a cgroup list becomes empty.
367 + if (count == 0) {
368 + size_t preserved_targets = 0;
369 + size_t preserved_pids = 0;
370 +
371 + netdata_mutex_lock(&mutex_cgroup_shm);
372 + preserved_targets = ebpf_count_cgroup_targets_unsafe();
373 + preserved_pids = ebpf_count_cgroup_pids_unsafe();
374 + netdata_mutex_unlock(&mutex_cgroup_shm);
375 +
376 + if (last_count != 0 ||
377 + previous_integration_active != integration_active ||
378 + previous_systemd_enabled != systemd_enabled) {
379 + collector_info(
380 + "EBPF CGROUP: empty netipc snapshot generation=%llu items=%u enabled=%u preserved_targets=%zu "
381 + "preserved_pids=%zu integration_active=%d systemd_enabled=%d refresh_failures=%d",
382 + (unsigned long long)generation,
383 + count,
384 + enabled_count,
385 + preserved_targets,
386 + preserved_pids,
387 + integration_active,
388 + systemd_enabled,
389 + refresh_fail_count);
390 + }
391 +
392 + previous_count = 0;
393 + previous_enabled_count = enabled_count;
394 + previous_imported_targets = preserved_targets;
395 + previous_total_pids = preserved_pids;
396 + previous_integration_active = integration_active;
397 + previous_systemd_enabled = systemd_enabled;
398 return;
399 }
400
401 netdata_mutex_lock(&mutex_cgroup_shm);
402 ebpf_remove_cgroup_target_update_list();
300 -
403 ebpf_reset_updated_var();
404
303 - for (i = 0; i < end; i++) {
304 - netdata_ebpf_cgroup_shm_body_t *ptr = &shm_ebpf_cgroup.body[i];
305 - if (ptr->enabled) {
306 - ebpf_cgroup_target_t *ect = ebpf_cgroup_find_or_create(ptr);
307 - ebpf_update_pid_link_list(ect, ptr->path);
405 + for (uint32_t i = 0; i < count; i++) {
406 + const nipc_cgroups_cache_item_t *item = &ebpf_cgroup_cache.items[i];
407 + if (item->enabled) {
408 + ebpf_cgroup_target_t *ect = ebpf_cgroup_find_or_create(item);
409 + ebpf_update_pid_link_list(ect, item->path);
410 }
411 }
310 - send_cgroup_chart = previous != shm_ebpf_cgroup.header->cgroup_root_count;
311 - previous = shm_ebpf_cgroup.header->cgroup_root_count;
312 - sem_post(shm_sem_ebpf_cgroup);
412 +
413 + size_t imported_targets = ebpf_count_cgroup_targets_unsafe();
414 + size_t total_pids = ebpf_count_cgroup_pids_unsafe();
415 +
416 + int chart_refresh_needed = previous_count != count;
417 + ebpf_send_cgroup_chart_set(chart_refresh_needed);
418 + previous_count = count;
419 netdata_mutex_unlock(&mutex_cgroup_shm);
314 -#ifdef NETDATA_DEV_MODE
315 - netdata_log_info(
316 - "Updating cgroup %d (Previous: %d, Current: %d)",
317 - send_cgroup_chart,
318 - previous,
319 - shm_ebpf_cgroup.header->cgroup_root_count);
420 +
421 + if (last_count != count ||
422 + previous_enabled_count != enabled_count ||
423 + previous_imported_targets != imported_targets ||
424 + previous_total_pids != total_pids ||
425 + previous_integration_active != integration_active ||
426 + previous_systemd_enabled != systemd_enabled ||
427 + enabled_count == 0 || imported_targets == 0 || total_pids == 0) {
428 + collector_info(
429 + "EBPF CGROUP: netipc snapshot generation=%llu items=%u enabled=%u imported_targets=%zu total_pids=%zu "
430 + "send_cgroup_chart=%d integration_active=%d systemd_enabled=%d refresh_failures=%d",
431 + (unsigned long long)generation,
432 + count,
433 + enabled_count,
434 + imported_targets,
435 + total_pids,
436 + chart_refresh_needed,
437 + integration_active,
438 + systemd_enabled,
439 + refresh_fail_count);
440 + }
441 +
442 + previous_enabled_count = enabled_count;
443 + previous_imported_targets = imported_targets;
444 + previous_total_pids = total_pids;
445 + previous_integration_active = integration_active;
446 + previous_systemd_enabled = systemd_enabled;
447 #endif
448 }
449
@@ -388,7 +515,7 @@ void ebpf_cgroup_integration(void *ptr __maybe_unused)
515 int counter = NETDATA_EBPF_CGROUP_UPDATE - 1;
516 heartbeat_t hb;
517 heartbeat_init(&hb, USEC_PER_SEC);
391 - //Plugin will be killed when it receives a signal
518 +
519 while (!ebpf_plugin_stop()) {
520 if (ebpf_plugin_stop())
521 break;
@@ -397,14 +524,14 @@ void ebpf_cgroup_integration(void *ptr __maybe_unused)
524
525 if (ebpf_plugin_stop())
526 break;
400 - // We are using a small heartbeat time to wake up thread,
401 - // but we should not update so frequently the shared memory data
527 +
528 + // refresh every NETDATA_EBPF_CGROUP_UPDATE seconds
529 if (++counter >= NETDATA_EBPF_CGROUP_UPDATE) {
530 counter = 0;
404 - if (!shm_ebpf_cgroup.header)
405 - ebpf_map_cgroup_shared_memory();
406 - else
407 - ebpf_parse_cgroup_shm_data();
531 + if (!ebpf_cgroup_cache_initialized)
532 + ebpf_cgroup_cache_init();
533 +
534 + ebpf_parse_cgroup_netipc_data();
535 }
536 }
537 }
src/collectors/ebpf.plugin/ebpf_cgroup.h
+2 -7
@@ -3,9 +3,6 @@
3 #ifndef NETDATA_EBPF_CGROUP_H
4 #define NETDATA_EBPF_CGROUP_H 1
5
6 -#define NETDATA_EBPF_CGROUP_MAX_TRIES 3
7 -#define NETDATA_EBPF_CGROUP_NEXT_TRY_SEC 30
8 -
6 #include "ebpf.h"
7 #include "ebpf_apps.h"
8
@@ -83,11 +80,9 @@ typedef struct ebpf_systemd_args {
80 char *dimension;
81 } ebpf_systemd_args_t;
82
86 -void ebpf_map_cgroup_shared_memory();
87 -void ebpf_parse_cgroup_shm_data();
83 void ebpf_create_charts_on_systemd(ebpf_systemd_args_t *chart);
84 void ebpf_cgroup_integration(void *ptr);
90 -void ebpf_unmap_cgroup_shared_memory();
91 -extern int send_cgroup_chart;
85 +void ebpf_cgroup_cache_cleanup(void);
86 +extern _Atomic int send_cgroup_chart;
87
88 #endif /* NETDATA_EBPF_CGROUP_H */
src/collectors/ebpf.plugin/ebpf_dcstat.c
+4 -4
@@ -750,7 +750,7 @@ void ebpf_read_dcstat_thread(void *ptr)
750 break;
751 }
752
753 - if (cgroups && shm_ebpf_cgroup.header)
753 + if (cgroups && ebpf_cgroup_integration_active_get())
754 ebpf_update_dc_cgroup();
755
756 if (sem_post(shm_mutex_ebpf_integration)) {
@@ -1354,8 +1354,8 @@ void ebpf_dc_send_cgroup_data(int update_every)
1354 ebpf_cgroup_target_t *ect;
1355 ebpf_dc_calc_chart_values();
1356
1357 - if (shm_ebpf_cgroup.header->systemd_enabled) {
1358 - if (send_cgroup_chart) {
1357 + if (ebpf_cgroup_systemd_enabled_get()) {
1358 + if (ebpf_send_cgroup_chart_get()) {
1359 ebpf_create_systemd_dc_charts(update_every);
1360 }
1361
@@ -1430,7 +1430,7 @@ static void dcstat_collector(ebpf_module_t *em)
1430 break;
1431 }
1432
1433 - if (cgroups && shm_ebpf_cgroup.header)
1433 + if (cgroups && ebpf_cgroup_integration_active_get())
1434 ebpf_dc_send_cgroup_data(update_every);
1435
1436 netdata_mutex_unlock(&lock);
src/collectors/ebpf.plugin/ebpf_fd.c
+4 -4
@@ -939,7 +939,7 @@ void ebpf_read_fd_thread(void *ptr)
939 break;
940 }
941
942 - if (cgroups && shm_ebpf_cgroup.header)
942 + if (cgroups && ebpf_cgroup_integration_active_get())
943 ebpf_update_fd_cgroup();
944
945 if (sem_post(shm_mutex_ebpf_integration)) {
@@ -1358,8 +1358,8 @@ static void ebpf_fd_send_cgroup_data(ebpf_module_t *em)
1358 return;
1359 }
1360
1361 - if (shm_ebpf_cgroup.header->systemd_enabled) {
1362 - if (send_cgroup_chart) {
1361 + if (ebpf_cgroup_systemd_enabled_get()) {
1362 + if (ebpf_send_cgroup_chart_get()) {
1363 ebpf_create_systemd_fd_charts(em);
1364 }
1365
@@ -1437,7 +1437,7 @@ static void fd_collector(ebpf_module_t *em)
1437 break;
1438 }
1439
1440 - if (cgroups && shm_ebpf_cgroup.header)
1440 + if (cgroups && ebpf_cgroup_integration_active_get())
1441 ebpf_fd_send_cgroup_data(em);
1442
1443 netdata_mutex_unlock(&lock);
src/collectors/ebpf.plugin/ebpf_oomkill.c
+4 -4
@@ -328,8 +328,8 @@ void ebpf_oomkill_send_cgroup_data(int update_every)
328 netdata_mutex_lock(&mutex_cgroup_shm);
329 ebpf_cgroup_target_t *ect;
330
331 - if (shm_ebpf_cgroup.header->systemd_enabled) {
332 - if (send_cgroup_chart) {
331 + if (ebpf_cgroup_systemd_enabled_get()) {
332 + if (ebpf_send_cgroup_chart_get()) {
333 ebpf_create_systemd_oomkill_charts(update_every);
334 }
335 ebpf_send_systemd_oomkill_charts();
@@ -479,7 +479,7 @@ static void oomkill_collector(ebpf_module_t *em)
479 stats[NETDATA_CONTROLLER_PID_TABLE_ADD] += (uint64_t)count;
480 stats[NETDATA_CONTROLLER_PID_TABLE_DEL] += (uint64_t)count;
481
482 - if (cgroups && shm_ebpf_cgroup.header)
482 + if (cgroups && ebpf_cgroup_integration_active_get())
483 ebpf_update_oomkill_cgroup(keys, count);
484
485 if (ebpf_plugin_stop())
@@ -488,7 +488,7 @@ static void oomkill_collector(ebpf_module_t *em)
488 netdata_apps_integration_flags_t apps = em->apps_charts;
489 netdata_mutex_lock(&lock);
490 // write everything from the ebpf map.
491 - if (cgroups && shm_ebpf_cgroup.header)
491 + if (cgroups && ebpf_cgroup_integration_active_get())
492 ebpf_oomkill_send_cgroup_data(update_every);
493
494 if (apps & NETDATA_EBPF_APPS_FLAG_CHART_CREATED)
src/collectors/ebpf.plugin/ebpf_process.c
+4 -4
@@ -1423,8 +1423,8 @@ static void ebpf_process_send_cgroup_data(ebpf_module_t *em)
1423 return;
1424 }
1425
1426 - if (shm_ebpf_cgroup.header->systemd_enabled) {
1427 - if (send_cgroup_chart) {
1426 + if (ebpf_cgroup_systemd_enabled_get()) {
1427 + if (ebpf_send_cgroup_chart_get()) {
1428 ebpf_create_systemd_process_charts(em);
1429 }
1430
@@ -1648,7 +1648,7 @@ static void process_collector(ebpf_module_t *em)
1648 netdata_mutex_lock(&collect_data_mutex);
1649 collect_data_for_all_processes(process_pid_fd, process_maps_per_core);
1650
1651 - if (cgroups && shm_ebpf_cgroup.header) {
1651 + if (cgroups && ebpf_cgroup_integration_active_get()) {
1652 ebpf_update_process_cgroup();
1653 }
1654 netdata_mutex_unlock(&collect_data_mutex);
@@ -1675,7 +1675,7 @@ static void process_collector(ebpf_module_t *em)
1675 ebpf_process_send_apps_data(apps_groups_root_target, em);
1676 }
1677
1678 - if (cgroups && shm_ebpf_cgroup.header) {
1678 + if (cgroups && ebpf_cgroup_integration_active_get()) {
1679 if (!ebpf_plugin_stop())
1680 ebpf_process_send_cgroup_data(em);
1681 }
src/collectors/ebpf.plugin/ebpf_shm.c
+4 -4
@@ -1058,8 +1058,8 @@ void ebpf_shm_send_cgroup_data(int update_every)
1058 return;
1059 }
1060
1061 - if (shm_ebpf_cgroup.header->systemd_enabled) {
1062 - if (send_cgroup_chart) {
1061 + if (ebpf_cgroup_systemd_enabled_get()) {
1062 + if (ebpf_send_cgroup_chart_get()) {
1063 ebpf_create_systemd_shm_charts(update_every);
1064 }
1065
@@ -1161,7 +1161,7 @@ void ebpf_read_shm_thread(void *ptr)
1161 break;
1162 }
1163
1164 - if (cgroups && shm_ebpf_cgroup.header)
1164 + if (cgroups && ebpf_cgroup_integration_active_get())
1165 ebpf_update_shm_cgroup();
1166
1167 if (sem_post(shm_mutex_ebpf_integration)) {
@@ -1224,7 +1224,7 @@ static void shm_collector(ebpf_module_t *em)
1224 break;
1225 }
1226
1227 - if (cgroups && shm_ebpf_cgroup.header) {
1227 + if (cgroups && ebpf_cgroup_integration_active_get()) {
1228 ebpf_shm_send_cgroup_data(update_every);
1229 }
1230
src/collectors/ebpf.plugin/ebpf_socket.c
+4 -4
@@ -2048,7 +2048,7 @@ void ebpf_read_socket_thread(void *ptr)
2048 break;
2049 }
2050
2051 - if (cgroups && shm_ebpf_cgroup.header)
2051 + if (cgroups && ebpf_cgroup_integration_active_get())
2052 ebpf_update_socket_cgroup();
2053
2054 if (sem_post(shm_mutex_ebpf_integration)) {
@@ -2803,8 +2803,8 @@ static void ebpf_socket_send_cgroup_data(int update_every)
2803 return;
2804 }
2805
2806 - if (shm_ebpf_cgroup.header->systemd_enabled) {
2807 - if (send_cgroup_chart) {
2806 + if (ebpf_cgroup_systemd_enabled_get()) {
2807 + if (ebpf_send_cgroup_chart_get()) {
2808 ebpf_create_systemd_socket_charts(update_every);
2809 }
2810 ebpf_send_systemd_socket_charts();
@@ -2893,7 +2893,7 @@ static void socket_collector(ebpf_module_t *em)
2893 break;
2894 }
2895
2896 - if (cgroups && shm_ebpf_cgroup.header)
2896 + if (cgroups && ebpf_cgroup_integration_active_get())
2897 ebpf_socket_send_cgroup_data(update_every);
2898
2899 fflush(stdout);
src/collectors/ebpf.plugin/ebpf_swap.c
+4 -4
@@ -702,7 +702,7 @@ void ebpf_read_swap_thread(void *ptr)
702 break;
703 }
704
705 - if (cgroups && shm_ebpf_cgroup.header)
705 + if (cgroups && ebpf_cgroup_integration_active_get())
706 ebpf_update_swap_cgroup();
707
708 if (sem_post(shm_mutex_ebpf_integration)) {
@@ -1025,8 +1025,8 @@ void ebpf_swap_send_cgroup_data(int update_every)
1025 return;
1026 }
1027
1028 - if (shm_ebpf_cgroup.header->systemd_enabled) {
1029 - if (send_cgroup_chart) {
1028 + if (ebpf_cgroup_systemd_enabled_get()) {
1029 + if (ebpf_send_cgroup_chart_get()) {
1030 ebpf_create_systemd_swap_charts(update_every);
1031 fflush(stdout);
1032 }
@@ -1101,7 +1101,7 @@ static void swap_collector(ebpf_module_t *em)
1101 break;
1102 }
1103
1104 - if (cgroup && shm_ebpf_cgroup.header)
1104 + if (cgroup && ebpf_cgroup_integration_active_get())
1105 ebpf_swap_send_cgroup_data(update_every);
1106
1107 netdata_mutex_unlock(&lock);
src/collectors/ebpf.plugin/ebpf_vfs.c
+4 -4
@@ -2258,8 +2258,8 @@ static void ebpf_vfs_send_cgroup_data(ebpf_module_t *em)
2258 return;
2259 }
2260
2261 - if (shm_ebpf_cgroup.header->systemd_enabled) {
2262 - if (send_cgroup_chart) {
2261 + if (ebpf_cgroup_systemd_enabled_get()) {
2262 + if (ebpf_send_cgroup_chart_get()) {
2263 ebpf_create_systemd_vfs_charts(em);
2264 }
2265 ebpf_send_systemd_vfs_charts(em);
@@ -2360,7 +2360,7 @@ void ebpf_read_vfs_thread(void *ptr)
2360 break;
2361 }
2362
2363 - if (cgroups && shm_ebpf_cgroup.header)
2363 + if (cgroups && ebpf_cgroup_integration_active_get())
2364 read_update_vfs_cgroup();
2365
2366 if (sem_post(shm_mutex_ebpf_integration)) {
@@ -2429,7 +2429,7 @@ static void vfs_collector(ebpf_module_t *em)
2429 break;
2430 }
2431
2432 - if (cgroups && shm_ebpf_cgroup.header)
2432 + if (cgroups && ebpf_cgroup_integration_active_get())
2433 ebpf_vfs_send_cgroup_data(em);
2434
2435 netdata_mutex_unlock(&lock);
src/crates/Cargo.lock
+54
@@ -311,6 +311,21 @@ dependencies = [
311 "serde",
312 ]
313
314 +[[package]]
315 +name = "bit-set"
316 +version = "0.8.0"
317 +source = "registry+https://github.com/rust-lang/crates.io-index"
318 +checksum = "08807e080ed7f9d5433fa9b275196cfc35414f66a0c79d864dc51a0d825231a3"
319 +dependencies = [
320 + "bit-vec",
321 +]
322 +
323 +[[package]]
324 +name = "bit-vec"
325 +version = "0.8.0"
326 +source = "registry+https://github.com/rust-lang/crates.io-index"
327 +checksum = "5e764a1d40d510daf35e07be9eb06e75770908c27d411ee6c92109c9840eaaf7"
328 +
329 [[package]]
330 name = "bitflags"
331 version = "1.3.2"
@@ -2031,6 +2046,14 @@ dependencies = [
2046 "serde_json",
2047 ]
2048
2049 +[[package]]
2050 +name = "netipc"
2051 +version = "0.1.0"
2052 +dependencies = [
2053 + "libc",
2054 + "proptest",
2055 +]
2056 +
2057 [[package]]
2058 name = "nix"
2059 version = "0.30.1"
@@ -2408,12 +2431,16 @@ version = "1.9.0"
2431 source = "registry+https://github.com/rust-lang/crates.io-index"
2432 checksum = "bee689443a2bd0a16ab0348b52ee43e3b2d1b1f931c8aa5c9f8de4c86fbe8c40"
2433 dependencies = [
2434 + "bit-set",
2435 + "bit-vec",
2436 "bitflags 2.10.0",
2437 "num-traits",
2438 "rand 0.9.2",
2439 "rand_chacha 0.9.0",
2440 "rand_xorshift",
2441 "regex-syntax",
2442 + "rusty-fork",
2443 + "tempfile",
2444 "unarray",
2445 ]
2446
@@ -2472,6 +2499,12 @@ dependencies = [
2499 "prost 0.13.5",
2500 ]
2501
2502 +[[package]]
2503 +name = "quick-error"
2504 +version = "1.2.3"
2505 +source = "registry+https://github.com/rust-lang/crates.io-index"
2506 +checksum = "a1d01941d82fa2ab50be1e79e6714289dd7cde78eba4c074bc5a4374f650dfe0"
2507 +
2508 [[package]]
2509 name = "quote"
2510 version = "1.0.42"
@@ -2793,6 +2826,18 @@ version = "1.0.22"
2826 source = "registry+https://github.com/rust-lang/crates.io-index"
2827 checksum = "b39cdef0fa800fc44525c84ccb54a029961a8215f9619753635a9c0d2538d46d"
2828
2829 +[[package]]
2830 +name = "rusty-fork"
2831 +version = "0.3.1"
2832 +source = "registry+https://github.com/rust-lang/crates.io-index"
2833 +checksum = "cc6bf79ff24e648f6da1f8d1f011e9cac26491b619e6b9280f2b47f1774e6ee2"
2834 +dependencies = [
2835 + "fnv",
2836 + "quick-error",
2837 + "tempfile",
2838 + "wait-timeout",
2839 +]
2840 +
2841 [[package]]
2842 name = "ruzstd"
2843 version = "0.8.2"
@@ -3596,6 +3641,15 @@ version = "0.9.5"
3641 source = "registry+https://github.com/rust-lang/crates.io-index"
3642 checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a"
3643
3644 +[[package]]
3645 +name = "wait-timeout"
3646 +version = "0.2.1"
3647 +source = "registry+https://github.com/rust-lang/crates.io-index"
3648 +checksum = "09ac3b126d3914f9849036f826e054cbabdc8519970b8998ddaf3b5bd3c65f11"
3649 +dependencies = [
3650 + "libc",
3651 +]
3652 +
3653 [[package]]
3654 name = "walkdir"
3655 version = "2.5.0"
src/crates/Cargo.toml
+3
@@ -27,6 +27,9 @@ members = [
27 # Netdata OTEL workspace members
28 "netdata-otel/otel-plugin",
29 "netdata-otel/flatten_otel",
30 +
31 + # Plugin IPC library
32 + "netipc",
33 ]
34
35 [workspace.package]
src/crates/netdata-otel/otel-plugin/src/plugin_config/env.rs
+20 -16
@@ -3,10 +3,10 @@ use bytesize::ByteSize;
3 use std::env;
4 use std::time::Duration;
5
6 +use super::PluginConfigOverride;
7 use super::endpoint::EndpointConfigOverride;
8 use super::logs::LogsConfigOverride;
9 use super::metrics::MetricsConfigOverride;
9 -use super::PluginConfigOverride;
10
11 /// Read an environment variable, returning `None` if not set and an error if not valid UTF-8.
12 fn read_env(name: &str) -> Result<Option<String>> {
@@ -77,9 +77,21 @@ impl PluginConfigOverride {
77 let logs = LogsConfigOverride::from_env()?;
78
79 Ok(Self {
80 - endpoint: if endpoint.has_overrides() { Some(endpoint) } else { None },
81 - metrics: if metrics.has_overrides() { Some(metrics) } else { None },
82 - logs: if logs.has_overrides() { Some(logs) } else { None },
80 + endpoint: if endpoint.has_overrides() {
81 + Some(endpoint)
82 + } else {
83 + None
84 + },
85 + metrics: if metrics.has_overrides() {
86 + Some(metrics)
87 + } else {
88 + None
89 + },
90 + logs: if logs.has_overrides() {
91 + Some(logs)
92 + } else {
93 + None
94 + },
95 })
96 }
97 }
@@ -128,18 +140,10 @@ impl LogsConfigOverride {
140 fn from_env() -> Result<Self> {
141 Ok(Self {
142 journal_dir: env_var("NETDATA_OTEL_LOGS_JOURNAL_DIR")?,
131 - size_of_journal_file: parse_env_bytesize(
132 - "NETDATA_OTEL_LOGS_SIZE_OF_JOURNAL_FILE",
133 - )?,
134 - entries_of_journal_file: parse_env_var(
135 - "NETDATA_OTEL_LOGS_ENTRIES_OF_JOURNAL_FILE",
136 - )?,
137 - number_of_journal_files: parse_env_var(
138 - "NETDATA_OTEL_LOGS_NUMBER_OF_JOURNAL_FILES",
139 - )?,
140 - size_of_journal_files: parse_env_bytesize(
141 - "NETDATA_OTEL_LOGS_SIZE_OF_JOURNAL_FILES",
142 - )?,
143 + size_of_journal_file: parse_env_bytesize("NETDATA_OTEL_LOGS_SIZE_OF_JOURNAL_FILE")?,
144 + entries_of_journal_file: parse_env_var("NETDATA_OTEL_LOGS_ENTRIES_OF_JOURNAL_FILE")?,
145 + number_of_journal_files: parse_env_var("NETDATA_OTEL_LOGS_NUMBER_OF_JOURNAL_FILES")?,
146 + size_of_journal_files: parse_env_bytesize("NETDATA_OTEL_LOGS_SIZE_OF_JOURNAL_FILES")?,
147 duration_of_journal_files: parse_env_duration(
148 "NETDATA_OTEL_LOGS_DURATION_OF_JOURNAL_FILES",
149 )?,
src/crates/netdata-otel/otel-plugin/src/plugin_config/mod.rs
+20 -46
@@ -94,9 +94,8 @@ impl PluginConfig {
94 .map(|p| p.join("otel.yaml"));
95
96 let mut config = match &stock_path {
97 - Some(path) => Self::from_yaml_file(path).with_context(|| {
98 - format!("loading stock config from {}", path.display())
99 - })?,
97 + Some(path) => Self::from_yaml_file(path)
98 + .with_context(|| format!("loading stock config from {}", path.display()))?,
99 None => anyhow::bail!("no stock configuration directory available"),
100 };
101
@@ -109,9 +108,7 @@ impl PluginConfig {
108 .map(|path| path.join("otel.yaml"))
109 {
110 if let Some(overrides) = Self::load_overrides(&user_path)
112 - .with_context(|| {
113 - format!("loading user config from {}", user_path.display())
114 - })?
111 + .with_context(|| format!("loading user config from {}", user_path.display()))?
112 {
113 config.apply_overrides(&overrides);
114 config.log_config(ConfigSource::User);
@@ -154,10 +151,7 @@ impl PluginConfig {
151 }
152
153 // Validate TLS configuration
157 - let tls_enabled = match (
158 - &self.endpoint.tls_cert_path,
159 - &self.endpoint.tls_key_path,
160 - ) {
154 + let tls_enabled = match (&self.endpoint.tls_cert_path, &self.endpoint.tls_key_path) {
155 (Some(cert_path), Some(key_path)) => {
156 if cert_path.is_empty() {
157 anyhow::bail!("TLS certificate path cannot be empty when provided");
@@ -207,9 +201,7 @@ impl PluginConfig {
201 Ok(contents) => contents,
202 Err(e) if e.kind() == std::io::ErrorKind::NotFound => return Ok(None),
203 Err(e) => {
210 - return Err(
211 - anyhow::Error::new(e).context(format!("reading {}", path.display()))
212 - );
204 + return Err(anyhow::Error::new(e).context(format!("reading {}", path.display())));
205 }
206 };
207 let overrides: PluginConfigOverride = serde_yaml::from_str(&contents)
@@ -496,13 +488,10 @@ logs:
488
489 #[test]
490 fn env_override_metrics_interval() {
499 - with_env_vars(
500 - &[("NETDATA_OTEL_METRICS_INTERVAL_SECS", "30")],
501 - || {
502 - let overrides = PluginConfigOverride::from_env().unwrap();
503 - assert_eq!(overrides.metrics.as_ref().unwrap().interval_secs, Some(30));
504 - },
505 - );
491 + with_env_vars(&[("NETDATA_OTEL_METRICS_INTERVAL_SECS", "30")], || {
492 + let overrides = PluginConfigOverride::from_env().unwrap();
493 + assert_eq!(overrides.metrics.as_ref().unwrap().interval_secs, Some(30));
494 + });
495 }
496
497 #[test]
@@ -535,30 +524,18 @@ logs:
524
525 #[test]
526 fn env_override_bool_values() {
538 - with_env_vars(
539 - &[("NETDATA_OTEL_LOGS_STORE_OTLP_JSON", "true")],
540 - || {
541 - let overrides = PluginConfigOverride::from_env().unwrap();
542 - assert_eq!(
543 - overrides.logs.as_ref().unwrap().store_otlp_json,
544 - Some(true)
545 - );
546 - },
547 - );
527 + with_env_vars(&[("NETDATA_OTEL_LOGS_STORE_OTLP_JSON", "true")], || {
528 + let overrides = PluginConfigOverride::from_env().unwrap();
529 + assert_eq!(overrides.logs.as_ref().unwrap().store_otlp_json, Some(true));
530 + });
531 }
532
533 #[test]
534 fn env_override_bool_accepts_yes_no() {
552 - with_env_vars(
553 - &[("NETDATA_OTEL_LOGS_STORE_OTLP_JSON", "yes")],
554 - || {
555 - let overrides = PluginConfigOverride::from_env().unwrap();
556 - assert_eq!(
557 - overrides.logs.as_ref().unwrap().store_otlp_json,
558 - Some(true)
559 - );
560 - },
561 - );
535 + with_env_vars(&[("NETDATA_OTEL_LOGS_STORE_OTLP_JSON", "yes")], || {
536 + let overrides = PluginConfigOverride::from_env().unwrap();
537 + assert_eq!(overrides.logs.as_ref().unwrap().store_otlp_json, Some(true));
538 + });
539 }
540
541 #[test]
@@ -573,12 +550,9 @@ logs:
550
551 #[test]
552 fn env_override_invalid_bool_is_rejected() {
576 - with_env_vars(
577 - &[("NETDATA_OTEL_LOGS_STORE_OTLP_JSON", "maybe")],
578 - || {
579 - assert!(PluginConfigOverride::from_env().is_err());
580 - },
581 - );
553 + with_env_vars(&[("NETDATA_OTEL_LOGS_STORE_OTLP_JSON", "maybe")], || {
554 + assert!(PluginConfigOverride::from_env().is_err());
555 + });
556 }
557
558 #[test]
src/crates/netipc/Cargo.toml new
+13
@@ -0,0 +1,13 @@
1 +[package]
2 +name = "netipc"
3 +version = "0.1.0"
4 +edition = "2021"
5 +
6 +[lib]
7 +path = "src/lib.rs"
8 +
9 +[dependencies]
10 +libc = "0.2"
11 +
12 +[dev-dependencies]
13 +proptest = "1"
src/crates/netipc/src/bin/interop_codec.rs new
+351
@@ -0,0 +1,351 @@
1 +//! Encode/decode test messages to/from files for cross-language interop testing.
2 +//!
3 +//! Usage:
4 +//! interop_codec encode <output_dir> - encode all test messages to files
5 +//! interop_codec decode <input_dir> - decode files and verify correctness
6 +//!
7 +//! Returns 0 on success, 1 on failure.
8 +
9 +use netipc::protocol::*;
10 +use std::fs;
11 +use std::path::Path;
12 +use std::process;
13 +
14 +fn write_file(dir: &str, name: &str, data: &[u8]) {
15 + let path = Path::new(dir).join(name);
16 + fs::write(&path, data).unwrap_or_else(|e| {
17 + eprintln!("ERROR: cannot write {}: {}", path.display(), e);
18 + process::exit(1);
19 + });
20 +}
21 +
22 +fn read_file(dir: &str, name: &str) -> Vec<u8> {
23 + let path = Path::new(dir).join(name);
24 + fs::read(&path).unwrap_or_else(|e| {
25 + eprintln!("ERROR: cannot read {}: {}", path.display(), e);
26 + process::exit(1);
27 + })
28 +}
29 +
30 +struct Checker {
31 + pass: u32,
32 + fail: u32,
33 +}
34 +
35 +impl Checker {
36 + fn new() -> Self {
37 + Checker { pass: 0, fail: 0 }
38 + }
39 +
40 + fn check(&mut self, cond: bool, name: &str) {
41 + if cond {
42 + self.pass += 1;
43 + } else {
44 + self.fail += 1;
45 + eprintln!("FAIL: {}", name);
46 + }
47 + }
48 +
49 + fn report(&self, label: &str) -> bool {
50 + println!("{}: {} passed, {} failed", label, self.pass, self.fail);
51 + self.fail == 0
52 + }
53 +}
54 +
55 +fn do_encode(dir: &str) {
56 + // 1. Outer message header
57 + {
58 + let h = Header {
59 + magic: MAGIC_MSG,
60 + version: VERSION,
61 + header_len: HEADER_LEN,
62 + kind: KIND_REQUEST,
63 + flags: FLAG_BATCH,
64 + code: METHOD_CGROUPS_SNAPSHOT,
65 + transport_status: STATUS_OK,
66 + payload_len: 12345,
67 + item_count: 42,
68 + message_id: 0xDEAD_BEEF_CAFE_BABE,
69 + };
70 + let mut buf = [0u8; 32];
71 + h.encode(&mut buf);
72 + write_file(dir, "header.bin", &buf);
73 + }
74 +
75 + // 2. Chunk continuation header
76 + {
77 + let c = ChunkHeader {
78 + magic: MAGIC_CHUNK,
79 + version: VERSION,
80 + flags: 0,
81 + message_id: 0x1234_5678_90AB_CDEF,
82 + total_message_len: 100000,
83 + chunk_index: 3,
84 + chunk_count: 10,
85 + chunk_payload_len: 8192,
86 + };
87 + let mut buf = [0u8; 32];
88 + c.encode(&mut buf);
89 + write_file(dir, "chunk_header.bin", &buf);
90 + }
91 +
92 + // 3. Hello payload
93 + {
94 + let h = Hello {
95 + layout_version: 1,
96 + flags: 0,
97 + supported_profiles: PROFILE_BASELINE | PROFILE_SHM_FUTEX,
98 + preferred_profiles: PROFILE_SHM_FUTEX,
99 + max_request_payload_bytes: 4096,
100 + max_request_batch_items: 100,
101 + max_response_payload_bytes: 1048576,
102 + max_response_batch_items: 1,
103 + auth_token: 0xAABB_CCDD_EEFF_0011,
104 + packet_size: 65536,
105 + };
106 + let mut buf = [0u8; 44];
107 + h.encode(&mut buf);
108 + write_file(dir, "hello.bin", &buf);
109 + }
110 +
111 + // 4. Hello-ack payload
112 + {
113 + let h = HelloAck {
114 + layout_version: 1,
115 + flags: 0,
116 + server_supported_profiles: 0x07,
117 + intersection_profiles: 0x05,
118 + selected_profile: PROFILE_SHM_FUTEX,
119 + agreed_max_request_payload_bytes: 2048,
120 + agreed_max_request_batch_items: 50,
121 + agreed_max_response_payload_bytes: 65536,
122 + agreed_max_response_batch_items: 1,
123 + agreed_packet_size: 32768,
124 + session_id: 1,
125 + };
126 + let mut buf = [0u8; 48];
127 + h.encode(&mut buf);
128 + write_file(dir, "hello_ack.bin", &buf);
129 + }
130 +
131 + // 5. Cgroups request
132 + {
133 + let r = CgroupsRequest {
134 + layout_version: 1,
135 + flags: 0,
136 + };
137 + let mut buf = [0u8; 4];
138 + r.encode(&mut buf);
139 + write_file(dir, "cgroups_req.bin", &buf);
140 + }
141 +
142 + // 6. Cgroups snapshot response (multi-item)
143 + {
144 + let mut buf = [0u8; 8192];
145 + let mut b = CgroupsBuilder::new(&mut buf, 3, 1, 999);
146 +
147 + b.add(100, 0, 1, b"init.scope", b"/sys/fs/cgroup/init.scope")
148 + .unwrap();
149 + b.add(
150 + 200,
151 + 0x02,
152 + 0,
153 + b"system.slice/docker-abc.scope",
154 + b"/sys/fs/cgroup/system.slice/docker-abc.scope",
155 + )
156 + .unwrap();
157 + b.add(300, 0, 1, b"", b"").unwrap();
158 +
159 + let total = b.finish();
160 + write_file(dir, "cgroups_resp.bin", &buf[..total]);
161 + }
162 +
163 + // 7. Empty cgroups snapshot
164 + {
165 + let mut buf = [0u8; 8192];
166 + let b = CgroupsBuilder::new(&mut buf, 0, 0, 42);
167 + let total = b.finish();
168 + write_file(dir, "cgroups_resp_empty.bin", &buf[..total]);
169 + }
170 +}
171 +
172 +fn do_decode(dir: &str) -> bool {
173 + let mut c = Checker::new();
174 +
175 + // 1. Outer message header
176 + {
177 + let data = read_file(dir, "header.bin");
178 + let hdr = Header::decode(&data);
179 + c.check(hdr.is_ok(), "decode header");
180 + if let Ok(h) = hdr {
181 + c.check(h.magic == MAGIC_MSG, "header magic");
182 + c.check(h.version == VERSION, "header version");
183 + c.check(h.header_len == HEADER_LEN, "header header_len");
184 + c.check(h.kind == KIND_REQUEST, "header kind");
185 + c.check(h.flags == FLAG_BATCH, "header flags");
186 + c.check(h.code == METHOD_CGROUPS_SNAPSHOT, "header code");
187 + c.check(h.transport_status == STATUS_OK, "header transport_status");
188 + c.check(h.payload_len == 12345, "header payload_len");
189 + c.check(h.item_count == 42, "header item_count");
190 + c.check(h.message_id == 0xDEAD_BEEF_CAFE_BABE, "header message_id");
191 + }
192 + }
193 +
194 + // 2. Chunk continuation header
195 + {
196 + let data = read_file(dir, "chunk_header.bin");
197 + let chk = ChunkHeader::decode(&data);
198 + c.check(chk.is_ok(), "decode chunk");
199 + if let Ok(ch) = chk {
200 + c.check(ch.magic == MAGIC_CHUNK, "chunk magic");
201 + c.check(ch.message_id == 0x1234_5678_90AB_CDEF, "chunk message_id");
202 + c.check(ch.total_message_len == 100000, "chunk total_message_len");
203 + c.check(ch.chunk_index == 3, "chunk chunk_index");
204 + c.check(ch.chunk_count == 10, "chunk chunk_count");
205 + c.check(ch.chunk_payload_len == 8192, "chunk chunk_payload_len");
206 + }
207 + }
208 +
209 + // 3. Hello payload
210 + {
211 + let data = read_file(dir, "hello.bin");
212 + let hello = Hello::decode(&data);
213 + c.check(hello.is_ok(), "decode hello");
214 + if let Ok(h) = hello {
215 + c.check(
216 + h.supported_profiles == (PROFILE_BASELINE | PROFILE_SHM_FUTEX),
217 + "hello supported",
218 + );
219 + c.check(h.preferred_profiles == PROFILE_SHM_FUTEX, "hello preferred");
220 + c.check(h.max_request_payload_bytes == 4096, "hello max_req_payload");
221 + c.check(h.max_request_batch_items == 100, "hello max_req_batch");
222 + c.check(
223 + h.max_response_payload_bytes == 1048576,
224 + "hello max_resp_payload",
225 + );
226 + c.check(h.max_response_batch_items == 1, "hello max_resp_batch");
227 + c.check(h.auth_token == 0xAABB_CCDD_EEFF_0011, "hello auth_token");
228 + c.check(h.packet_size == 65536, "hello packet_size");
229 + }
230 + }
231 +
232 + // 4. Hello-ack payload
233 + {
234 + let data = read_file(dir, "hello_ack.bin");
235 + let ack = HelloAck::decode(&data);
236 + c.check(ack.is_ok(), "decode hello_ack");
237 + if let Ok(h) = ack {
238 + c.check(
239 + h.server_supported_profiles == 0x07,
240 + "hello_ack server_supported",
241 + );
242 + c.check(h.intersection_profiles == 0x05, "hello_ack intersection");
243 + c.check(
244 + h.selected_profile == PROFILE_SHM_FUTEX,
245 + "hello_ack selected",
246 + );
247 + c.check(
248 + h.agreed_max_request_payload_bytes == 2048,
249 + "hello_ack req_payload",
250 + );
251 + c.check(
252 + h.agreed_max_request_batch_items == 50,
253 + "hello_ack req_batch",
254 + );
255 + c.check(
256 + h.agreed_max_response_payload_bytes == 65536,
257 + "hello_ack resp_payload",
258 + );
259 + c.check(
260 + h.agreed_max_response_batch_items == 1,
261 + "hello_ack resp_batch",
262 + );
263 + c.check(h.agreed_packet_size == 32768, "hello_ack pkt_size");
264 + }
265 + }
266 +
267 + // 5. Cgroups request
268 + {
269 + let data = read_file(dir, "cgroups_req.bin");
270 + let req = CgroupsRequest::decode(&data);
271 + c.check(req.is_ok(), "decode cgroups_req");
272 + if let Ok(r) = req {
273 + c.check(r.layout_version == 1, "cgroups_req layout_version");
274 + c.check(r.flags == 0, "cgroups_req flags");
275 + }
276 + }
277 +
278 + // 6. Cgroups snapshot response (multi-item)
279 + {
280 + let data = read_file(dir, "cgroups_resp.bin");
281 + let view = CgroupsResponseView::decode(&data);
282 + c.check(view.is_ok(), "decode cgroups_resp");
283 + if let Ok(v) = view {
284 + c.check(v.item_count == 3, "cgroups_resp item_count");
285 + c.check(v.systemd_enabled == 1, "cgroups_resp systemd_enabled");
286 + c.check(v.generation == 999, "cgroups_resp generation");
287 +
288 + if let Ok(item) = v.item(0) {
289 + c.check(item.hash == 100, "item 0 hash");
290 + c.check(item.options == 0, "item 0 options");
291 + c.check(item.enabled == 1, "item 0 enabled");
292 + c.check(item.name.as_bytes() == b"init.scope", "item 0 name");
293 + c.check(
294 + item.path.as_bytes() == b"/sys/fs/cgroup/init.scope",
295 + "item 0 path",
296 + );
297 + }
298 +
299 + if let Ok(item) = v.item(1) {
300 + c.check(item.hash == 200, "item 1 hash");
301 + c.check(item.options == 0x02, "item 1 options");
302 + c.check(item.enabled == 0, "item 1 enabled");
303 + c.check(
304 + item.name.as_bytes() == b"system.slice/docker-abc.scope",
305 + "item 1 name",
306 + );
307 + }
308 +
309 + if let Ok(item) = v.item(2) {
310 + c.check(item.hash == 300, "item 2 hash");
311 + c.check(item.name.len == 0, "item 2 name empty");
312 + c.check(item.path.len == 0, "item 2 path empty");
313 + }
314 + }
315 + }
316 +
317 + // 7. Empty cgroups snapshot
318 + {
319 + let data = read_file(dir, "cgroups_resp_empty.bin");
320 + let view = CgroupsResponseView::decode(&data);
321 + c.check(view.is_ok(), "decode cgroups_resp_empty");
322 + if let Ok(v) = view {
323 + c.check(v.item_count == 0, "empty item_count");
324 + c.check(v.systemd_enabled == 0, "empty systemd_enabled");
325 + c.check(v.generation == 42, "empty generation");
326 + }
327 + }
328 +
329 + c.report("Rust decode")
330 +}
331 +
332 +fn main() {
333 + let args: Vec<String> = std::env::args().collect();
334 + if args.len() != 3 {
335 + eprintln!("Usage: {} <encode|decode> <dir>", args[0]);
336 + process::exit(1);
337 + }
338 +
339 + match args[1].as_str() {
340 + "encode" => do_encode(&args[2]),
341 + "decode" => {
342 + if !do_decode(&args[2]) {
343 + process::exit(1);
344 + }
345 + }
346 + other => {
347 + eprintln!("Unknown command: {}", other);
348 + process::exit(1);
349 + }
350 + }
351 +}
src/crates/netipc/src/lib.rs new
+5
@@ -0,0 +1,5 @@
1 +// netipc - Netdata Plugin IPC Library (Rust)
2 +
3 +pub mod protocol;
4 +pub mod service;
5 +pub mod transport;
src/crates/netipc/src/protocol/cgroups.rs new
+560
@@ -0,0 +1,560 @@
1 +//! Cgroups snapshot codec -- request, response view, builder, dispatch.
2 +
3 +use super::{align8, NipcError, ALIGNMENT};
4 +
5 +const CGROUPS_REQ_SIZE: usize = 4;
6 +const CGROUPS_RESP_HDR_SIZE: usize = 24;
7 +const CGROUPS_DIR_ENTRY_SIZE: usize = 8;
8 +const CGROUPS_ITEM_HDR_SIZE: usize = 32;
9 +
10 +// ---------------------------------------------------------------------------
11 +// Cgroups snapshot request (4 bytes)
12 +// ---------------------------------------------------------------------------
13 +
14 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
15 +pub struct CgroupsRequest {
16 + pub layout_version: u16,
17 + pub flags: u16,
18 +}
19 +
20 +impl CgroupsRequest {
21 + /// Encode into `buf`. Returns 4 on success, 0 if buf is too small.
22 + pub fn encode(&self, buf: &mut [u8]) -> usize {
23 + if buf.len() < CGROUPS_REQ_SIZE {
24 + return 0;
25 + }
26 + buf[0..2].copy_from_slice(&self.layout_version.to_ne_bytes());
27 + buf[2..4].copy_from_slice(&self.flags.to_ne_bytes());
28 + CGROUPS_REQ_SIZE
29 + }
30 +
31 + /// Decode from `buf`. Validates layout_version.
32 + pub fn decode(buf: &[u8]) -> Result<Self, NipcError> {
33 + if buf.len() < CGROUPS_REQ_SIZE {
34 + return Err(NipcError::Truncated);
35 + }
36 + let r = CgroupsRequest {
37 + layout_version: u16::from_ne_bytes(buf[0..2].try_into().unwrap()),
38 + flags: u16::from_ne_bytes(buf[2..4].try_into().unwrap()),
39 + };
40 + if r.layout_version != 1 {
41 + return Err(NipcError::BadLayout);
42 + }
43 + // flags must be zero (reserved for future use)
44 + if r.flags != 0 {
45 + return Err(NipcError::BadLayout);
46 + }
47 + Ok(r)
48 + }
49 +}
50 +
51 +// ---------------------------------------------------------------------------
52 +// Cgroups snapshot response
53 +// ---------------------------------------------------------------------------
54 +
55 +/// Borrowed string view into the payload buffer.
56 +/// Valid only while the underlying payload buffer is alive.
57 +#[derive(Debug, Clone, Copy, PartialEq)]
58 +pub struct StrView<'a> {
59 + /// Slice into payload, NUL-terminated.
60 + pub bytes: &'a [u8],
61 + /// Length excluding the NUL.
62 + pub len: u32,
63 +}
64 +
65 +impl<'a> StrView<'a> {
66 + /// Return the string content as a `&str`, or a UTF-8 error.
67 + pub fn as_str(&self) -> Result<&'a str, core::str::Utf8Error> {
68 + core::str::from_utf8(&self.bytes[..self.len as usize])
69 + }
70 +
71 + /// Return the string content as a byte slice (without the NUL).
72 + pub fn as_bytes(&self) -> &'a [u8] {
73 + &self.bytes[..self.len as usize]
74 + }
75 +}
76 +
77 +/// Per-item view -- ephemeral, borrows the payload buffer.
78 +/// Valid only while the payload buffer is alive.
79 +#[derive(Debug, Clone, Copy, PartialEq)]
80 +pub struct CgroupsItemView<'a> {
81 + pub layout_version: u16,
82 + pub flags: u16,
83 + pub hash: u32,
84 + pub options: u32,
85 + pub enabled: u32,
86 + pub name: StrView<'a>,
87 + pub path: StrView<'a>,
88 +}
89 +
90 +/// Full snapshot view -- ephemeral, borrows the payload buffer.
91 +/// Valid only during the current library call or callback.
92 +/// Copy immediately if the data is needed later.
93 +#[derive(Debug)]
94 +pub struct CgroupsResponseView<'a> {
95 + pub layout_version: u16,
96 + pub flags: u16,
97 + pub item_count: u32,
98 + pub systemd_enabled: u32,
99 + pub generation: u64,
100 + payload: &'a [u8],
101 +}
102 +
103 +impl<'a> CgroupsResponseView<'a> {
104 + /// Decode the snapshot response header and validate the item directory.
105 + /// On success, use `item()` to access individual items.
106 + pub fn decode(buf: &'a [u8]) -> Result<Self, NipcError> {
107 + if buf.len() < CGROUPS_RESP_HDR_SIZE {
108 + return Err(NipcError::Truncated);
109 + }
110 +
111 + let layout_version = u16::from_ne_bytes(buf[0..2].try_into().unwrap());
112 + let flags = u16::from_ne_bytes(buf[2..4].try_into().unwrap());
113 + let item_count = u32::from_ne_bytes(buf[4..8].try_into().unwrap());
114 + let systemd_enabled = u32::from_ne_bytes(buf[8..12].try_into().unwrap());
115 + // buf[12..16] reserved, must be zero
116 + let reserved = u32::from_ne_bytes(buf[12..16].try_into().unwrap());
117 + let generation = u64::from_ne_bytes(buf[16..24].try_into().unwrap());
118 +
119 + if layout_version != 1 {
120 + return Err(NipcError::BadLayout);
121 + }
122 +
123 + // flags must be zero
124 + if flags != 0 {
125 + return Err(NipcError::BadLayout);
126 + }
127 +
128 + // reserved field must be zero
129 + if reserved != 0 {
130 + return Err(NipcError::BadLayout);
131 + }
132 +
133 + // Validate directory fits (checked arithmetic for 32-bit safety)
134 + let dir_size = (item_count as usize)
135 + .checked_mul(CGROUPS_DIR_ENTRY_SIZE)
136 + .ok_or(NipcError::BadLayout)?;
137 + let dir_end = CGROUPS_RESP_HDR_SIZE
138 + .checked_add(dir_size)
139 + .ok_or(NipcError::BadLayout)?;
140 + if dir_end > buf.len() {
141 + return Err(NipcError::Truncated);
142 + }
143 +
144 + let packed_area_len = buf.len() - dir_end;
145 +
146 + // Validate each directory entry
147 + for i in 0..item_count as usize {
148 + let base = CGROUPS_RESP_HDR_SIZE + i * 8;
149 + let off = u32::from_ne_bytes(buf[base..base + 4].try_into().unwrap());
150 + let len = u32::from_ne_bytes(buf[base + 4..base + 8].try_into().unwrap());
151 +
152 + if (off as usize) % ALIGNMENT != 0 {
153 + return Err(NipcError::BadAlignment);
154 + }
155 + if (off as u64) + (len as u64) > packed_area_len as u64 {
156 + return Err(NipcError::OutOfBounds);
157 + }
158 + if (len as usize) < CGROUPS_ITEM_HDR_SIZE {
159 + return Err(NipcError::Truncated);
160 + }
161 + }
162 +
163 + Ok(CgroupsResponseView {
164 + layout_version,
165 + flags,
166 + item_count,
167 + systemd_enabled,
168 + generation,
169 + payload: buf,
170 + })
171 + }
172 +
173 + /// Access item at `index`. Returns an ephemeral item view.
174 + pub fn item(&self, index: u32) -> Result<CgroupsItemView<'a>, NipcError> {
175 + if index >= self.item_count {
176 + return Err(NipcError::OutOfBounds);
177 + }
178 +
179 + let dir_start = CGROUPS_RESP_HDR_SIZE;
180 + let dir_size = self.item_count as usize * CGROUPS_DIR_ENTRY_SIZE;
181 + let packed_area_start = dir_start + dir_size;
182 +
183 + let dir_base = dir_start + index as usize * 8;
184 + let item_off = u32::from_ne_bytes(self.payload[dir_base..dir_base + 4].try_into().unwrap());
185 + let item_len =
186 + u32::from_ne_bytes(self.payload[dir_base + 4..dir_base + 8].try_into().unwrap());
187 +
188 + let item_start = packed_area_start + item_off as usize;
189 + let item = &self.payload[item_start..item_start + item_len as usize];
190 +
191 + let layout_version = u16::from_ne_bytes(item[0..2].try_into().unwrap());
192 + let flags = u16::from_ne_bytes(item[2..4].try_into().unwrap());
193 + let hash = u32::from_ne_bytes(item[4..8].try_into().unwrap());
194 + let options = u32::from_ne_bytes(item[8..12].try_into().unwrap());
195 + let enabled = u32::from_ne_bytes(item[12..16].try_into().unwrap());
196 +
197 + let name_off = u32::from_ne_bytes(item[16..20].try_into().unwrap()) as usize;
198 + let name_len = u32::from_ne_bytes(item[20..24].try_into().unwrap());
199 + let path_off = u32::from_ne_bytes(item[24..28].try_into().unwrap()) as usize;
200 + let path_len = u32::from_ne_bytes(item[28..32].try_into().unwrap());
201 +
202 + if layout_version != 1 {
203 + return Err(NipcError::BadLayout);
204 + }
205 +
206 + // item flags must be zero
207 + if flags != 0 {
208 + return Err(NipcError::BadLayout);
209 + }
210 +
211 + // Validate name string
212 + if name_off < CGROUPS_ITEM_HDR_SIZE {
213 + return Err(NipcError::OutOfBounds);
214 + }
215 + if (name_off as u64) + (name_len as u64) + 1 > item_len as u64 {
216 + return Err(NipcError::OutOfBounds);
217 + }
218 + if item[name_off + name_len as usize] != 0 {
219 + return Err(NipcError::MissingNul);
220 + }
221 +
222 + // Validate path string
223 + if path_off < CGROUPS_ITEM_HDR_SIZE {
224 + return Err(NipcError::OutOfBounds);
225 + }
226 + if (path_off as u64) + (path_len as u64) + 1 > item_len as u64 {
227 + return Err(NipcError::OutOfBounds);
228 + }
229 + if item[path_off + path_len as usize] != 0 {
230 + return Err(NipcError::MissingNul);
231 + }
232 +
233 + // Reject overlapping name and path regions (including NUL)
234 + {
235 + let name_start = name_off as u64;
236 + let name_end = name_start + name_len as u64 + 1;
237 + let path_start = path_off as u64;
238 + let path_end = path_start + path_len as u64 + 1;
239 + if name_start < path_end && path_start < name_end {
240 + return Err(NipcError::BadLayout);
241 + }
242 + }
243 +
244 + let name = StrView {
245 + bytes: &item[name_off..name_off + name_len as usize + 1],
246 + len: name_len,
247 + };
248 + let path = StrView {
249 + bytes: &item[path_off..path_off + path_len as usize + 1],
250 + len: path_len,
251 + };
252 +
253 + Ok(CgroupsItemView {
254 + layout_version,
255 + flags,
256 + hash,
257 + options,
258 + enabled,
259 + name,
260 + path,
261 + })
262 + }
263 +}
264 +
265 +// ---------------------------------------------------------------------------
266 +// Cgroups snapshot response builder
267 +// ---------------------------------------------------------------------------
268 +
269 +/// Builds a cgroups snapshot response payload.
270 +///
271 +/// Layout during building (max_items directory slots reserved):
272 +/// [24-byte header space] [max_items*8 directory] [packed items]
273 +///
274 +/// Layout after finish (compacted to actual item_count):
275 +/// [24-byte header] [item_count*8 directory] [packed items]
276 +pub struct CgroupsBuilder<'a> {
277 + buf: &'a mut [u8],
278 + systemd_enabled: u32,
279 + generation: u64,
280 + item_count: u32,
281 + max_items: u32,
282 + data_offset: usize, // current write position (absolute in buf)
283 +}
284 +
285 +impl<'a> CgroupsBuilder<'a> {
286 + /// Initialize the builder. `buf` must be caller-owned and large enough
287 + /// for the expected snapshot.
288 + /// Initialize the builder. `buf` must be at least `CGROUPS_RESP_HDR_SIZE`
289 + /// (24) bytes plus `max_items * 8` directory bytes.
290 + ///
291 + /// # Panics
292 + /// Panics if `buf` is too small for the response header and item directory.
293 + pub fn new(buf: &'a mut [u8], max_items: u32, systemd_enabled: u32, generation: u64) -> Self {
294 + let dir_size = (max_items as usize).saturating_mul(CGROUPS_DIR_ENTRY_SIZE);
295 + let min_required = CGROUPS_RESP_HDR_SIZE.saturating_add(dir_size);
296 + assert!(
297 + buf.len() >= min_required,
298 + "CgroupsBuilder buffer too small: need at least {} bytes (hdr {} + dir {}), got {}",
299 + min_required,
300 + CGROUPS_RESP_HDR_SIZE,
301 + dir_size,
302 + buf.len()
303 + );
304 + let data_offset = min_required;
305 + CgroupsBuilder {
306 + buf,
307 + systemd_enabled,
308 + generation,
309 + item_count: 0,
310 + max_items,
311 + data_offset,
312 + }
313 + }
314 +
315 + /// Update the response header fields written by `finish()`.
316 + pub fn set_header(&mut self, systemd_enabled: u32, generation: u64) {
317 + self.systemd_enabled = systemd_enabled;
318 + self.generation = generation;
319 + }
320 +
321 + /// Add one cgroup item. Handles offset bookkeeping, NUL termination,
322 + /// and alignment.
323 + pub fn add(
324 + &mut self,
325 + hash: u32,
326 + options: u32,
327 + enabled: u32,
328 + name: &[u8],
329 + path: &[u8],
330 + ) -> Result<(), NipcError> {
331 + if self.item_count >= self.max_items {
332 + return Err(NipcError::Overflow);
333 + }
334 +
335 + // Align item start to 8 bytes
336 + let item_start = align8(self.data_offset);
337 +
338 + // Item payload: 32-byte header + name + NUL + path + NUL
339 + let item_size = CGROUPS_ITEM_HDR_SIZE
340 + .checked_add(name.len())
341 + .and_then(|v| v.checked_add(1))
342 + .and_then(|v| v.checked_add(path.len()))
343 + .and_then(|v| v.checked_add(1))
344 + .ok_or(NipcError::Overflow)?;
345 +
346 + let item_end = item_start
347 + .checked_add(item_size)
348 + .ok_or(NipcError::Overflow)?;
349 + if item_end > self.buf.len() {
350 + return Err(NipcError::Overflow);
351 + }
352 +
353 + let item_start = checked_u32(item_start)?;
354 + let item_size = checked_u32(item_size)?;
355 + let name_offset = checked_u32(CGROUPS_ITEM_HDR_SIZE)?;
356 + let name_len = checked_u32(name.len())?;
357 + let path_offset = CGROUPS_ITEM_HDR_SIZE
358 + .checked_add(name.len())
359 + .and_then(|v| v.checked_add(1))
360 + .ok_or(NipcError::Overflow)?;
361 + let path_offset = checked_u32(path_offset)?;
362 + let path_len = checked_u32(path.len())?;
363 +
364 + // Zero alignment padding
365 + if item_start as usize > self.data_offset {
366 + self.buf[self.data_offset..item_start as usize].fill(0);
367 + }
368 +
369 + // Write item header
370 + let p = item_start as usize;
371 + self.buf[p..p + 2].copy_from_slice(&1u16.to_ne_bytes()); // layout_version
372 + self.buf[p + 2..p + 4].copy_from_slice(&0u16.to_ne_bytes()); // flags
373 + self.buf[p + 4..p + 8].copy_from_slice(&hash.to_ne_bytes());
374 + self.buf[p + 8..p + 12].copy_from_slice(&options.to_ne_bytes());
375 + self.buf[p + 12..p + 16].copy_from_slice(&enabled.to_ne_bytes());
376 + self.buf[p + 16..p + 20].copy_from_slice(&name_offset.to_ne_bytes());
377 + self.buf[p + 20..p + 24].copy_from_slice(&name_len.to_ne_bytes());
378 + self.buf[p + 24..p + 28].copy_from_slice(&path_offset.to_ne_bytes());
379 + self.buf[p + 28..p + 32].copy_from_slice(&path_len.to_ne_bytes());
380 +
381 + // Write strings with NUL terminators
382 + let name_start = p + name_offset as usize;
383 + self.buf[name_start..name_start + name.len()].copy_from_slice(name);
384 + self.buf[name_start + name.len()] = 0;
385 +
386 + let path_start = p + path_offset as usize;
387 + self.buf[path_start..path_start + path.len()].copy_from_slice(path);
388 + self.buf[path_start + path.len()] = 0;
389 +
390 + // Write directory entry (absolute offset stored temporarily)
391 + let dir_entry = CGROUPS_RESP_HDR_SIZE + self.item_count as usize * CGROUPS_DIR_ENTRY_SIZE;
392 + self.buf[dir_entry..dir_entry + 4].copy_from_slice(&item_start.to_ne_bytes());
393 + self.buf[dir_entry + 4..dir_entry + 8].copy_from_slice(&item_size.to_ne_bytes());
394 +
395 + self.data_offset = p + item_size as usize;
396 + self.item_count += 1;
397 + Ok(())
398 + }
399 +
400 + /// Finalize the builder. Returns the total payload size. The buffer now
401 + /// contains a complete, decodable cgroups snapshot response payload.
402 + pub fn finish(self) -> usize {
403 + let p = &mut *{ self.buf };
404 +
405 + if self.item_count == 0 {
406 + p[0..2].copy_from_slice(&1u16.to_ne_bytes());
407 + p[2..4].copy_from_slice(&0u16.to_ne_bytes());
408 + p[4..8].copy_from_slice(&0u32.to_ne_bytes());
409 + p[8..12].copy_from_slice(&self.systemd_enabled.to_ne_bytes());
410 + p[12..16].copy_from_slice(&0u32.to_ne_bytes());
411 + p[16..24].copy_from_slice(&self.generation.to_ne_bytes());
412 + return CGROUPS_RESP_HDR_SIZE;
413 + }
414 +
415 + // Where the decoder expects packed data to start
416 + let final_packed_start =
417 + CGROUPS_RESP_HDR_SIZE + self.item_count as usize * CGROUPS_DIR_ENTRY_SIZE;
418 +
419 + // Read the first directory entry to find where packed data begins
420 + let first_item_abs = u32::from_ne_bytes(
421 + p[CGROUPS_RESP_HDR_SIZE..CGROUPS_RESP_HDR_SIZE + 4]
422 + .try_into()
423 + .unwrap(),
424 + ) as usize;
425 +
426 + let packed_data_len = self.data_offset - first_item_abs;
427 +
428 + if final_packed_start < first_item_abs {
429 + // Shift packed data left
430 + p.copy_within(
431 + first_item_abs..first_item_abs + packed_data_len,
432 + final_packed_start,
433 + );
434 + }
435 +
436 + // Convert directory entries from absolute to relative offsets
437 + let dir_base = CGROUPS_RESP_HDR_SIZE;
438 + for i in 0..self.item_count as usize {
439 + let entry = dir_base + i * CGROUPS_DIR_ENTRY_SIZE;
440 + let abs_off = u32::from_ne_bytes(p[entry..entry + 4].try_into().unwrap());
441 + let rel_off = abs_off - first_item_abs as u32;
442 + p[entry..entry + 4].copy_from_slice(&rel_off.to_ne_bytes());
443 + // length stays the same
444 + }
445 +
446 + // Write snapshot header
447 + p[0..2].copy_from_slice(&1u16.to_ne_bytes());
448 + p[2..4].copy_from_slice(&0u16.to_ne_bytes());
449 + p[4..8].copy_from_slice(&self.item_count.to_ne_bytes());
450 + p[8..12].copy_from_slice(&self.systemd_enabled.to_ne_bytes());
451 + p[12..16].copy_from_slice(&0u32.to_ne_bytes());
452 + p[16..24].copy_from_slice(&self.generation.to_ne_bytes());
453 +
454 + final_packed_start + packed_data_len
455 + }
456 +}
457 +
458 +/// Estimate a safe upper bound for the number of cgroup items that can fit in
459 +/// `buf_size`. This is a reservation hint for the builder, not a guarantee for
460 +/// arbitrary string lengths.
461 +pub fn estimate_cgroups_max_items(buf_size: usize) -> u32 {
462 + if buf_size <= CGROUPS_RESP_HDR_SIZE {
463 + return 0;
464 + }
465 +
466 + let min_aligned_item = align8(CGROUPS_ITEM_HDR_SIZE + 2);
467 + ((buf_size - CGROUPS_RESP_HDR_SIZE) / (CGROUPS_DIR_ENTRY_SIZE + min_aligned_item)) as u32
468 +}
469 +
470 +fn checked_u32(value: usize) -> Result<u32, NipcError> {
471 + u32::try_from(value).map_err(|_| NipcError::Overflow)
472 +}
473 +
474 +/// CGROUPS_SNAPSHOT dispatch: decode request, build response via handler.
475 +pub fn dispatch_cgroups_snapshot<F>(
476 + req: &[u8],
477 + resp: &mut [u8],
478 + max_items: u32,
479 + handler: F,
480 +) -> Option<usize>
481 +where
482 + F: FnOnce(&CgroupsRequest, &mut CgroupsBuilder) -> bool,
483 +{
484 + let request = CgroupsRequest::decode(req).ok()?;
485 + let min_required = (max_items as usize)
486 + .checked_mul(CGROUPS_DIR_ENTRY_SIZE)
487 + .and_then(|d| d.checked_add(CGROUPS_RESP_HDR_SIZE));
488 + let min_required = match min_required {
489 + Some(v) => v,
490 + None => return None,
491 + };
492 + if resp.len() < min_required {
493 + return None;
494 + }
495 + let mut builder = CgroupsBuilder::new(resp, max_items, 0, 0);
496 + if !handler(&request, &mut builder) {
497 + return None;
498 + }
499 + Some(builder.finish())
500 +}
501 +
502 +// ---------------------------------------------------------------------------
503 +// Tests
504 +// ---------------------------------------------------------------------------
505 +
506 +#[cfg(test)]
507 +mod tests {
508 + use super::*;
509 +
510 + #[test]
511 + fn checked_u32_rejects_oversized_values() {
512 + assert_eq!(checked_u32(u32::MAX as usize).unwrap(), u32::MAX);
513 + if usize::BITS > 32 {
514 + assert_eq!(
515 + checked_u32((u32::MAX as usize) + 1).unwrap_err(),
516 + NipcError::Overflow
517 + );
518 + }
519 + }
520 +
521 + #[test]
522 + fn cgroups_item_overlapping_regions() {
523 + let mut buf = [0u8; 4096];
524 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 1);
525 + b.add(1, 0, 1, b"test", b"/test").unwrap();
526 + let total = b.finish();
527 +
528 + let dir_end = CGROUPS_RESP_HDR_SIZE + CGROUPS_DIR_ENTRY_SIZE;
529 + let item_off = u32::from_ne_bytes(
530 + buf[CGROUPS_RESP_HDR_SIZE..CGROUPS_RESP_HDR_SIZE + 4]
531 + .try_into()
532 + .unwrap(),
533 + ) as usize;
534 + let item_start = dir_end + item_off;
535 +
536 + let name_off =
537 + u32::from_ne_bytes(buf[item_start + 16..item_start + 20].try_into().unwrap());
538 + buf[item_start + 24..item_start + 28].copy_from_slice(&name_off.to_ne_bytes());
539 + let name_len =
540 + u32::from_ne_bytes(buf[item_start + 20..item_start + 24].try_into().unwrap());
541 + buf[item_start + 28..item_start + 32].copy_from_slice(&name_len.to_ne_bytes());
542 +
543 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
544 + assert_eq!(view.item(0).unwrap_err(), NipcError::BadLayout);
545 + }
546 +
547 + #[test]
548 + fn dispatch_cgroups_empty_finish_returns_some_header_only() {
549 + let req = CgroupsRequest {
550 + layout_version: 1,
551 + flags: 0,
552 + };
553 + let mut req_buf = [0u8; 4];
554 + req.encode(&mut req_buf);
555 + let mut resp = [0u8; 4096];
556 + let result = dispatch_cgroups_snapshot(&req_buf, &mut resp, 1, |_req, _builder| true);
557 + assert!(result.is_some());
558 + assert_eq!(result.unwrap(), CGROUPS_RESP_HDR_SIZE);
559 + }
560 +}
src/crates/netipc/src/protocol/increment.rs new
+92
@@ -0,0 +1,92 @@
1 +//! INCREMENT codec (method 1) -- 8-byte payload: { u64 value }
2 +
3 +use super::NipcError;
4 +
5 +pub const INCREMENT_PAYLOAD_SIZE: usize = 8;
6 +
7 +pub fn increment_encode(value: u64, buf: &mut [u8]) -> usize {
8 + if buf.len() < INCREMENT_PAYLOAD_SIZE {
9 + return 0;
10 + }
11 + buf[..8].copy_from_slice(&value.to_ne_bytes());
12 + INCREMENT_PAYLOAD_SIZE
13 +}
14 +
15 +pub fn increment_decode(buf: &[u8]) -> Result<u64, NipcError> {
16 + if buf.len() < INCREMENT_PAYLOAD_SIZE {
17 + return Err(NipcError::Truncated);
18 + }
19 + Ok(u64::from_ne_bytes(buf[..8].try_into().unwrap()))
20 +}
21 +
22 +/// INCREMENT dispatch: decode -> handler -> encode.
23 +pub fn dispatch_increment<F>(req: &[u8], resp: &mut [u8], handler: F) -> Option<usize>
24 +where
25 + F: FnOnce(u64) -> Option<u64>,
26 +{
27 + let value = increment_decode(req).ok()?;
28 + let result = handler(value)?;
29 + let n = increment_encode(result, resp);
30 + if n == 0 {
31 + return None;
32 + }
33 + Some(n)
34 +}
35 +
36 +#[cfg(test)]
37 +mod tests {
38 + use super::*;
39 +
40 + #[test]
41 + fn encode_too_small() {
42 + let mut buf = [0u8; 4];
43 + assert_eq!(increment_encode(42, &mut buf), 0);
44 + }
45 +
46 + #[test]
47 + fn decode_too_small() {
48 + let buf = [0u8; 7];
49 + assert_eq!(increment_decode(&buf), Err(NipcError::Truncated));
50 + }
51 +
52 + #[test]
53 + fn roundtrip() {
54 + let mut buf = [0u8; 8];
55 + let n = increment_encode(0xDEAD_BEEF, &mut buf);
56 + assert_eq!(n, 8);
57 + let val = increment_decode(&buf).unwrap();
58 + assert_eq!(val, 0xDEAD_BEEF);
59 + }
60 +
61 + #[test]
62 + fn dispatch_resp_too_small() {
63 + let mut req = [0u8; 8];
64 + increment_encode(42, &mut req);
65 + let mut resp = [0u8; 4];
66 + assert!(dispatch_increment(&req, &mut resp, |v| Some(v + 1)).is_none());
67 + }
68 +
69 + #[test]
70 + fn dispatch_bad_request() {
71 + let mut resp = [0u8; 8];
72 + assert!(dispatch_increment(&[0u8; 4], &mut resp, |v| Some(v)).is_none());
73 + }
74 +
75 + #[test]
76 + fn dispatch_handler_returns_none() {
77 + let mut req = [0u8; 8];
78 + increment_encode(42, &mut req);
79 + let mut resp = [0u8; 8];
80 + assert!(dispatch_increment(&req, &mut resp, |_| None).is_none());
81 + }
82 +
83 + #[test]
84 + fn dispatch_success() {
85 + let mut req = [0u8; 8];
86 + increment_encode(99, &mut req);
87 + let mut resp = [0u8; 8];
88 + let n = dispatch_increment(&req, &mut resp, |v| Some(v + 1)).unwrap();
89 + assert_eq!(n, 8);
90 + assert_eq!(increment_decode(&resp).unwrap(), 100);
91 + }
92 +}
src/crates/netipc/src/protocol/mod.rs new
+2395
@@ -0,0 +1,2395 @@
1 +//! Wire envelope and codec for the netipc protocol.
2 +//!
3 +//! Pure byte-layout encode/decode. No I/O, no transport, no allocation on
4 +//! decode. Localhost-only IPC -- all multi-byte fields use host byte order.
5 +//!
6 +//! Decoded `View` types borrow the underlying buffer and are valid only while
7 +//! that buffer lives. Copy immediately if the data is needed later.
8 +
9 +mod cgroups;
10 +mod increment;
11 +mod string_reverse;
12 +
13 +// Re-export all public symbols from submodules.
14 +pub use cgroups::*;
15 +pub use increment::*;
16 +pub use string_reverse::*;
17 +
18 +// ---------------------------------------------------------------------------
19 +// Constants
20 +// ---------------------------------------------------------------------------
21 +
22 +pub const MAGIC_MSG: u32 = 0x4e495043; // "NIPC"
23 +pub const MAGIC_CHUNK: u32 = 0x4e43484b; // "NCHK"
24 +pub const VERSION: u16 = 1;
25 +pub const HEADER_LEN: u16 = 32;
26 +pub const HEADER_SIZE: usize = 32;
27 +
28 +// Message kinds
29 +pub const KIND_REQUEST: u16 = 1;
30 +pub const KIND_RESPONSE: u16 = 2;
31 +pub const KIND_CONTROL: u16 = 3;
32 +
33 +// Flags
34 +pub const FLAG_BATCH: u16 = 0x0001;
35 +
36 +// Transport status
37 +pub const STATUS_OK: u16 = 0;
38 +pub const STATUS_BAD_ENVELOPE: u16 = 1;
39 +pub const STATUS_AUTH_FAILED: u16 = 2;
40 +pub const STATUS_INCOMPATIBLE: u16 = 3;
41 +pub const STATUS_UNSUPPORTED: u16 = 4;
42 +pub const STATUS_LIMIT_EXCEEDED: u16 = 5;
43 +pub const STATUS_INTERNAL_ERROR: u16 = 6;
44 +
45 +// Control opcodes
46 +pub const CODE_HELLO: u16 = 1;
47 +pub const CODE_HELLO_ACK: u16 = 2;
48 +
49 +// Method codes
50 +pub const METHOD_INCREMENT: u16 = 1;
51 +pub const METHOD_CGROUPS_SNAPSHOT: u16 = 2;
52 +pub const METHOD_STRING_REVERSE: u16 = 3;
53 +
54 +// Profile bits
55 +pub const PROFILE_BASELINE: u32 = 0x01;
56 +pub const PROFILE_SHM_HYBRID: u32 = 0x02;
57 +pub const PROFILE_SHM_FUTEX: u32 = 0x04;
58 +pub const PROFILE_SHM_WAITADDR: u32 = 0x08;
59 +
60 +// Defaults
61 +pub const MAX_PAYLOAD_DEFAULT: u32 = 1024;
62 +
63 +/// Hard cap on negotiated request payload sizes (1 MiB) — prevents a
64 +/// compromised peer from forcing excessive memory allocation.
65 +pub const MAX_PAYLOAD_CAP: u32 = 1024 * 1024;
66 +
67 +// Alignment
68 +pub const ALIGNMENT: usize = 8;
69 +
70 +// Payload sizes
71 +const HELLO_SIZE: usize = 44;
72 +const HELLO_ACK_SIZE: usize = 48;
73 +
74 +// ---------------------------------------------------------------------------
75 +// Errors
76 +// ---------------------------------------------------------------------------
77 +
78 +#[derive(Debug, Clone, Copy, PartialEq, Eq)]
79 +pub enum NipcError {
80 + /// Buffer too short for the expected structure.
81 + Truncated,
82 + /// Magic value mismatch.
83 + BadMagic,
84 + /// Unsupported version.
85 + BadVersion,
86 + /// header_len != 32.
87 + BadHeaderLen,
88 + /// Unknown message kind.
89 + BadKind,
90 + /// Unknown layout_version in a payload.
91 + BadLayout,
92 + /// Offset+length exceeds available data.
93 + OutOfBounds,
94 + /// String not NUL-terminated.
95 + MissingNul,
96 + /// Item not 8-byte aligned.
97 + BadAlignment,
98 + /// Directory inconsistent with payload size.
99 + BadItemCount,
100 + /// Builder ran out of space.
101 + Overflow,
102 +}
103 +
104 +impl core::fmt::Display for NipcError {
105 + fn fmt(&self, f: &mut core::fmt::Formatter<'_>) -> core::fmt::Result {
106 + match self {
107 + NipcError::Truncated => write!(f, "buffer too short"),
108 + NipcError::BadMagic => write!(f, "magic value mismatch"),
109 + NipcError::BadVersion => write!(f, "unsupported version"),
110 + NipcError::BadHeaderLen => write!(f, "header_len != 32"),
111 + NipcError::BadKind => write!(f, "unknown message kind"),
112 + NipcError::BadLayout => write!(f, "unknown layout_version"),
113 + NipcError::OutOfBounds => write!(f, "offset+length exceeds data"),
114 + NipcError::MissingNul => write!(f, "string not NUL-terminated"),
115 + NipcError::BadAlignment => write!(f, "item not 8-byte aligned"),
116 + NipcError::BadItemCount => write!(f, "item count inconsistent"),
117 + NipcError::Overflow => write!(f, "builder out of space"),
118 + }
119 + }
120 +}
121 +
122 +impl std::error::Error for NipcError {}
123 +
124 +// ---------------------------------------------------------------------------
125 +// Utility
126 +// ---------------------------------------------------------------------------
127 +
128 +/// Round `v` up to the next multiple of 8.
129 +#[inline]
130 +pub fn align8(v: usize) -> usize {
131 + (v + 7) & !7
132 +}
133 +
134 +// ---------------------------------------------------------------------------
135 +// Outer message header (32 bytes)
136 +// ---------------------------------------------------------------------------
137 +
138 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
139 +pub struct Header {
140 + pub magic: u32,
141 + pub version: u16,
142 + pub header_len: u16,
143 + pub kind: u16,
144 + pub flags: u16,
145 + pub code: u16,
146 + pub transport_status: u16,
147 + pub payload_len: u32,
148 + pub item_count: u32,
149 + pub message_id: u64,
150 +}
151 +
152 +impl Header {
153 + /// Encode into `buf`. Returns 32 on success, 0 if buf is too small.
154 + pub fn encode(&self, buf: &mut [u8]) -> usize {
155 + if buf.len() < HEADER_SIZE {
156 + return 0;
157 + }
158 + buf[0..4].copy_from_slice(&self.magic.to_ne_bytes());
159 + buf[4..6].copy_from_slice(&self.version.to_ne_bytes());
160 + buf[6..8].copy_from_slice(&self.header_len.to_ne_bytes());
161 + buf[8..10].copy_from_slice(&self.kind.to_ne_bytes());
162 + buf[10..12].copy_from_slice(&self.flags.to_ne_bytes());
163 + buf[12..14].copy_from_slice(&self.code.to_ne_bytes());
164 + buf[14..16].copy_from_slice(&self.transport_status.to_ne_bytes());
165 + buf[16..20].copy_from_slice(&self.payload_len.to_ne_bytes());
166 + buf[20..24].copy_from_slice(&self.item_count.to_ne_bytes());
167 + buf[24..32].copy_from_slice(&self.message_id.to_ne_bytes());
168 + HEADER_SIZE
169 + }
170 +
171 + /// Decode from `buf`. Validates magic, version, header_len, kind.
172 + pub fn decode(buf: &[u8]) -> Result<Self, NipcError> {
173 + if buf.len() < HEADER_SIZE {
174 + return Err(NipcError::Truncated);
175 + }
176 + let hdr = Header {
177 + magic: u32::from_ne_bytes(buf[0..4].try_into().unwrap()),
178 + version: u16::from_ne_bytes(buf[4..6].try_into().unwrap()),
179 + header_len: u16::from_ne_bytes(buf[6..8].try_into().unwrap()),
180 + kind: u16::from_ne_bytes(buf[8..10].try_into().unwrap()),
181 + flags: u16::from_ne_bytes(buf[10..12].try_into().unwrap()),
182 + code: u16::from_ne_bytes(buf[12..14].try_into().unwrap()),
183 + transport_status: u16::from_ne_bytes(buf[14..16].try_into().unwrap()),
184 + payload_len: u32::from_ne_bytes(buf[16..20].try_into().unwrap()),
185 + item_count: u32::from_ne_bytes(buf[20..24].try_into().unwrap()),
186 + message_id: u64::from_ne_bytes(buf[24..32].try_into().unwrap()),
187 + };
188 +
189 + if hdr.magic != MAGIC_MSG {
190 + return Err(NipcError::BadMagic);
191 + }
192 + if hdr.version != VERSION {
193 + return Err(NipcError::BadVersion);
194 + }
195 + if hdr.header_len != HEADER_LEN {
196 + return Err(NipcError::BadHeaderLen);
197 + }
198 + if hdr.kind < KIND_REQUEST || hdr.kind > KIND_CONTROL {
199 + return Err(NipcError::BadKind);
200 + }
201 + Ok(hdr)
202 + }
203 +}
204 +
205 +// ---------------------------------------------------------------------------
206 +// Chunk continuation header (32 bytes)
207 +// ---------------------------------------------------------------------------
208 +
209 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
210 +pub struct ChunkHeader {
211 + pub magic: u32,
212 + pub version: u16,
213 + pub flags: u16,
214 + pub message_id: u64,
215 + pub total_message_len: u32,
216 + pub chunk_index: u32,
217 + pub chunk_count: u32,
218 + pub chunk_payload_len: u32,
219 +}
220 +
221 +impl ChunkHeader {
222 + /// Encode into `buf`. Returns 32 on success, 0 if buf is too small.
223 + pub fn encode(&self, buf: &mut [u8]) -> usize {
224 + if buf.len() < HEADER_SIZE {
225 + return 0;
226 + }
227 + buf[0..4].copy_from_slice(&self.magic.to_ne_bytes());
228 + buf[4..6].copy_from_slice(&self.version.to_ne_bytes());
229 + buf[6..8].copy_from_slice(&self.flags.to_ne_bytes());
230 + buf[8..16].copy_from_slice(&self.message_id.to_ne_bytes());
231 + buf[16..20].copy_from_slice(&self.total_message_len.to_ne_bytes());
232 + buf[20..24].copy_from_slice(&self.chunk_index.to_ne_bytes());
233 + buf[24..28].copy_from_slice(&self.chunk_count.to_ne_bytes());
234 + buf[28..32].copy_from_slice(&self.chunk_payload_len.to_ne_bytes());
235 + HEADER_SIZE
236 + }
237 +
238 + /// Decode from `buf`. Validates magic and version.
239 + pub fn decode(buf: &[u8]) -> Result<Self, NipcError> {
240 + if buf.len() < HEADER_SIZE {
241 + return Err(NipcError::Truncated);
242 + }
243 + let chk = ChunkHeader {
244 + magic: u32::from_ne_bytes(buf[0..4].try_into().unwrap()),
245 + version: u16::from_ne_bytes(buf[4..6].try_into().unwrap()),
246 + flags: u16::from_ne_bytes(buf[6..8].try_into().unwrap()),
247 + message_id: u64::from_ne_bytes(buf[8..16].try_into().unwrap()),
248 + total_message_len: u32::from_ne_bytes(buf[16..20].try_into().unwrap()),
249 + chunk_index: u32::from_ne_bytes(buf[20..24].try_into().unwrap()),
250 + chunk_count: u32::from_ne_bytes(buf[24..28].try_into().unwrap()),
251 + chunk_payload_len: u32::from_ne_bytes(buf[28..32].try_into().unwrap()),
252 + };
253 +
254 + if chk.magic != MAGIC_CHUNK {
255 + return Err(NipcError::BadMagic);
256 + }
257 + if chk.version != VERSION {
258 + return Err(NipcError::BadVersion);
259 + }
260 + if chk.flags != 0 {
261 + return Err(NipcError::BadLayout);
262 + }
263 + if chk.chunk_payload_len == 0 {
264 + return Err(NipcError::BadLayout);
265 + }
266 + Ok(chk)
267 + }
268 +}
269 +
270 +// ---------------------------------------------------------------------------
271 +// Batch item directory
272 +// ---------------------------------------------------------------------------
273 +
274 +/// One entry in a batch item directory (8 bytes on wire).
275 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
276 +pub struct BatchEntry {
277 + pub offset: u32,
278 + pub length: u32,
279 +}
280 +
281 +/// Encode `entries` into `buf`. Returns total bytes written (entries.len() * 8),
282 +/// or 0 if buf is too small.
283 +pub fn batch_dir_encode(entries: &[BatchEntry], buf: &mut [u8]) -> usize {
284 + let need = entries.len() * 8;
285 + if buf.len() < need {
286 + return 0;
287 + }
288 + for (i, e) in entries.iter().enumerate() {
289 + let base = i * 8;
290 + buf[base..base + 4].copy_from_slice(&e.offset.to_ne_bytes());
291 + buf[base + 4..base + 8].copy_from_slice(&e.length.to_ne_bytes());
292 + }
293 + need
294 +}
295 +
296 +/// Decode `item_count` directory entries from `buf`. Validates alignment and
297 +/// that each entry falls within `packed_area_len`.
298 +pub fn batch_dir_decode(
299 + buf: &[u8],
300 + item_count: u32,
301 + packed_area_len: u32,
302 +) -> Result<Vec<BatchEntry>, NipcError> {
303 + let count = item_count as usize;
304 + let dir_size = count.checked_mul(8).ok_or(NipcError::BadLayout)?;
305 + if buf.len() < dir_size {
306 + return Err(NipcError::Truncated);
307 + }
308 +
309 + let mut out = Vec::with_capacity(count);
310 + for i in 0..count {
311 + let base = i * 8;
312 + let offset = u32::from_ne_bytes(buf[base..base + 4].try_into().unwrap());
313 + let length = u32::from_ne_bytes(buf[base + 4..base + 8].try_into().unwrap());
314 +
315 + if (offset as usize) % ALIGNMENT != 0 {
316 + return Err(NipcError::BadAlignment);
317 + }
318 + if (offset as u64) + (length as u64) > packed_area_len as u64 {
319 + return Err(NipcError::OutOfBounds);
320 + }
321 + out.push(BatchEntry { offset, length });
322 + }
323 + Ok(out)
324 +}
325 +
326 +/// Validate a batch directory without allocating. Checks alignment and
327 +/// that each entry falls within `packed_area_len`. Mirrors C's
328 +/// `nipc_batch_dir_validate`.
329 +pub fn batch_dir_validate(
330 + buf: &[u8],
331 + item_count: u32,
332 + packed_area_len: u32,
333 +) -> Result<(), NipcError> {
334 + let count = item_count as usize;
335 + let dir_size = count.checked_mul(8).ok_or(NipcError::BadLayout)?;
336 + if buf.len() < dir_size {
337 + return Err(NipcError::Truncated);
338 + }
339 + for i in 0..count {
340 + let base = i * 8;
341 + let offset = u32::from_ne_bytes(buf[base..base + 4].try_into().unwrap());
342 + let length = u32::from_ne_bytes(buf[base + 4..base + 8].try_into().unwrap());
343 + if (offset as usize) % ALIGNMENT != 0 {
344 + return Err(NipcError::BadAlignment);
345 + }
346 + if (offset as u64) + (length as u64) > packed_area_len as u64 {
347 + return Err(NipcError::OutOfBounds);
348 + }
349 + }
350 + Ok(())
351 +}
352 +
353 +/// Extract a single batch item by index from a complete batch payload.
354 +/// Returns (item_slice, item_len) on success.
355 +pub fn batch_item_get(
356 + payload: &[u8],
357 + item_count: u32,
358 + index: u32,
359 +) -> Result<(&[u8], u32), NipcError> {
360 + if index >= item_count {
361 + return Err(NipcError::OutOfBounds);
362 + }
363 +
364 + let dir_size = (item_count as usize)
365 + .checked_mul(8)
366 + .ok_or(NipcError::BadLayout)?;
367 + let dir_aligned = align8(dir_size);
368 +
369 + if payload.len() < dir_aligned {
370 + return Err(NipcError::Truncated);
371 + }
372 +
373 + let idx = index as usize;
374 + let base = idx * 8;
375 + let off = u32::from_ne_bytes(payload[base..base + 4].try_into().unwrap());
376 + let len = u32::from_ne_bytes(payload[base + 4..base + 8].try_into().unwrap());
377 +
378 + let packed_area_start = dir_aligned;
379 + let packed_area_len = payload.len() - packed_area_start;
380 +
381 + if (off as usize) % ALIGNMENT != 0 {
382 + return Err(NipcError::BadAlignment);
383 + }
384 + if (off as u64) + (len as u64) > packed_area_len as u64 {
385 + return Err(NipcError::OutOfBounds);
386 + }
387 +
388 + let start = packed_area_start + off as usize;
389 + let end = start + len as usize;
390 + Ok((&payload[start..end], len))
391 +}
392 +
393 +// ---------------------------------------------------------------------------
394 +// Batch builder
395 +// ---------------------------------------------------------------------------
396 +
397 +/// Builds a batch payload: [directory] [align-pad] [packed items].
398 +pub struct BatchBuilder<'a> {
399 + buf: &'a mut [u8],
400 + item_count: u32,
401 + max_items: u32,
402 + dir_end: usize, // byte offset where directory reservation ends
403 + data_offset: usize, // current offset within the packed data area (relative)
404 +}
405 +
406 +impl<'a> BatchBuilder<'a> {
407 + /// Create a new batch builder. `buf` must be large enough for
408 + /// `max_items * 8` (directory) + packed data.
409 + pub fn new(buf: &'a mut [u8], max_items: u32) -> Self {
410 + let dir_end = align8(max_items as usize * 8);
411 + BatchBuilder {
412 + buf,
413 + item_count: 0,
414 + max_items,
415 + dir_end,
416 + data_offset: 0,
417 + }
418 + }
419 +
420 + /// Add an item payload. Handles alignment padding.
421 + pub fn add(&mut self, item: &[u8]) -> Result<(), NipcError> {
422 + if self.item_count >= self.max_items {
423 + return Err(NipcError::Overflow);
424 + }
425 +
426 + let aligned_off = align8(self.data_offset);
427 + let abs_pos = self.dir_end + aligned_off;
428 +
429 + if abs_pos + item.len() > self.buf.len() {
430 + return Err(NipcError::Overflow);
431 + }
432 +
433 + // Zero alignment padding
434 + if aligned_off > self.data_offset {
435 + let pad_start = self.dir_end + self.data_offset;
436 + let pad_end = self.dir_end + aligned_off;
437 + self.buf[pad_start..pad_end].fill(0);
438 + }
439 +
440 + self.buf[abs_pos..abs_pos + item.len()].copy_from_slice(item);
441 +
442 + // Write directory entry
443 + let idx = self.item_count as usize;
444 + let dir_base = idx * 8;
445 + self.buf[dir_base..dir_base + 4].copy_from_slice(&(aligned_off as u32).to_ne_bytes());
446 + self.buf[dir_base + 4..dir_base + 8].copy_from_slice(&(item.len() as u32).to_ne_bytes());
447 +
448 + self.data_offset = aligned_off + item.len();
449 + self.item_count += 1;
450 + Ok(())
451 + }
452 +
453 + /// Finalize the batch. Returns (total_payload_size, item_count).
454 + /// Compacts if fewer items were added than max_items.
455 + pub fn finish(self) -> (usize, u32) {
456 + let count = self.item_count;
457 + let final_dir_aligned = align8(count as usize * 8);
458 +
459 + if final_dir_aligned < self.dir_end && self.data_offset > 0 {
460 + // Shift packed data left
461 + self.buf.copy_within(
462 + self.dir_end..self.dir_end + self.data_offset,
463 + final_dir_aligned,
464 + );
465 + }
466 +
467 + let total = final_dir_aligned + align8(self.data_offset);
468 + // Zero trailing alignment padding to avoid leaking stale buffer data
469 + if total < self.buf.len() {
470 + let pad_start = final_dir_aligned + self.data_offset;
471 + if pad_start < total {
472 + self.buf[pad_start..total].fill(0);
473 + }
474 + }
475 + (total, count)
476 + }
477 +}
478 +
479 +// ---------------------------------------------------------------------------
480 +// Hello payload (44 bytes)
481 +// ---------------------------------------------------------------------------
482 +
483 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
484 +pub struct Hello {
485 + pub layout_version: u16,
486 + pub flags: u16,
487 + pub supported_profiles: u32,
488 + pub preferred_profiles: u32,
489 + pub max_request_payload_bytes: u32,
490 + pub max_request_batch_items: u32,
491 + pub max_response_payload_bytes: u32,
492 + pub max_response_batch_items: u32,
493 + pub auth_token: u64,
494 + pub packet_size: u32,
495 +}
496 +
497 +impl Hello {
498 + /// Encode into `buf`. Returns 44 on success, 0 if buf is too small.
499 + pub fn encode(&self, buf: &mut [u8]) -> usize {
500 + if buf.len() < HELLO_SIZE {
501 + return 0;
502 + }
503 + buf[0..2].copy_from_slice(&self.layout_version.to_ne_bytes());
504 + buf[2..4].copy_from_slice(&self.flags.to_ne_bytes());
505 + buf[4..8].copy_from_slice(&self.supported_profiles.to_ne_bytes());
506 + buf[8..12].copy_from_slice(&self.preferred_profiles.to_ne_bytes());
507 + buf[12..16].copy_from_slice(&self.max_request_payload_bytes.to_ne_bytes());
508 + buf[16..20].copy_from_slice(&self.max_request_batch_items.to_ne_bytes());
509 + buf[20..24].copy_from_slice(&self.max_response_payload_bytes.to_ne_bytes());
510 + buf[24..28].copy_from_slice(&self.max_response_batch_items.to_ne_bytes());
511 + buf[28..32].copy_from_slice(&0u32.to_ne_bytes()); // padding
512 + buf[32..40].copy_from_slice(&self.auth_token.to_ne_bytes());
513 + buf[40..44].copy_from_slice(&self.packet_size.to_ne_bytes());
514 + HELLO_SIZE
515 + }
516 +
517 + /// Decode from `buf`. Validates layout_version.
518 + pub fn decode(buf: &[u8]) -> Result<Self, NipcError> {
519 + if buf.len() < HELLO_SIZE {
520 + return Err(NipcError::Truncated);
521 + }
522 + let h = Hello {
523 + layout_version: u16::from_ne_bytes(buf[0..2].try_into().unwrap()),
524 + flags: u16::from_ne_bytes(buf[2..4].try_into().unwrap()),
525 + supported_profiles: u32::from_ne_bytes(buf[4..8].try_into().unwrap()),
526 + preferred_profiles: u32::from_ne_bytes(buf[8..12].try_into().unwrap()),
527 + max_request_payload_bytes: u32::from_ne_bytes(buf[12..16].try_into().unwrap()),
528 + max_request_batch_items: u32::from_ne_bytes(buf[16..20].try_into().unwrap()),
529 + max_response_payload_bytes: u32::from_ne_bytes(buf[20..24].try_into().unwrap()),
530 + max_response_batch_items: u32::from_ne_bytes(buf[24..28].try_into().unwrap()),
531 + // buf[28..32] is reserved padding, must be zero
532 + auth_token: u64::from_ne_bytes(buf[32..40].try_into().unwrap()),
533 + packet_size: u32::from_ne_bytes(buf[40..44].try_into().unwrap()),
534 + };
535 +
536 + if h.layout_version != 1 {
537 + return Err(NipcError::BadLayout);
538 + }
539 +
540 + // Validate padding bytes 28..32 are zero
541 + if u32::from_ne_bytes(buf[28..32].try_into().unwrap()) != 0 {
542 + return Err(NipcError::BadLayout);
543 + }
544 +
545 + Ok(h)
546 + }
547 +}
548 +
549 +// ---------------------------------------------------------------------------
550 +// Hello-ack payload (48 bytes)
551 +// ---------------------------------------------------------------------------
552 +
553 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)]
554 +pub struct HelloAck {
555 + pub layout_version: u16,
556 + pub flags: u16,
557 + pub server_supported_profiles: u32,
558 + pub intersection_profiles: u32,
559 + pub selected_profile: u32,
560 + pub agreed_max_request_payload_bytes: u32,
561 + pub agreed_max_request_batch_items: u32,
562 + pub agreed_max_response_payload_bytes: u32,
563 + pub agreed_max_response_batch_items: u32,
564 + pub agreed_packet_size: u32,
565 + pub session_id: u64,
566 +}
567 +
568 +impl HelloAck {
569 + /// Encode into `buf`. Returns 48 on success, 0 if buf is too small.
570 + pub fn encode(&self, buf: &mut [u8]) -> usize {
571 + if buf.len() < HELLO_ACK_SIZE {
572 + return 0;
573 + }
574 + buf[0..2].copy_from_slice(&self.layout_version.to_ne_bytes());
575 + buf[2..4].copy_from_slice(&self.flags.to_ne_bytes());
576 + buf[4..8].copy_from_slice(&self.server_supported_profiles.to_ne_bytes());
577 + buf[8..12].copy_from_slice(&self.intersection_profiles.to_ne_bytes());
578 + buf[12..16].copy_from_slice(&self.selected_profile.to_ne_bytes());
579 + buf[16..20].copy_from_slice(&self.agreed_max_request_payload_bytes.to_ne_bytes());
580 + buf[20..24].copy_from_slice(&self.agreed_max_request_batch_items.to_ne_bytes());
581 + buf[24..28].copy_from_slice(&self.agreed_max_response_payload_bytes.to_ne_bytes());
582 + buf[28..32].copy_from_slice(&self.agreed_max_response_batch_items.to_ne_bytes());
583 + buf[32..36].copy_from_slice(&self.agreed_packet_size.to_ne_bytes());
584 + buf[36..40].copy_from_slice(&0u32.to_ne_bytes()); // padding
585 + buf[40..48].copy_from_slice(&self.session_id.to_ne_bytes());
586 + HELLO_ACK_SIZE
587 + }
588 +
589 + /// Decode from `buf`. Validates layout_version.
590 + pub fn decode(buf: &[u8]) -> Result<Self, NipcError> {
591 + if buf.len() < HELLO_ACK_SIZE {
592 + return Err(NipcError::Truncated);
593 + }
594 + let h = HelloAck {
595 + layout_version: u16::from_ne_bytes(buf[0..2].try_into().unwrap()),
596 + flags: u16::from_ne_bytes(buf[2..4].try_into().unwrap()),
597 + server_supported_profiles: u32::from_ne_bytes(buf[4..8].try_into().unwrap()),
598 + intersection_profiles: u32::from_ne_bytes(buf[8..12].try_into().unwrap()),
599 + selected_profile: u32::from_ne_bytes(buf[12..16].try_into().unwrap()),
600 + agreed_max_request_payload_bytes: u32::from_ne_bytes(buf[16..20].try_into().unwrap()),
601 + agreed_max_request_batch_items: u32::from_ne_bytes(buf[20..24].try_into().unwrap()),
602 + agreed_max_response_payload_bytes: u32::from_ne_bytes(buf[24..28].try_into().unwrap()),
603 + agreed_max_response_batch_items: u32::from_ne_bytes(buf[28..32].try_into().unwrap()),
604 + agreed_packet_size: u32::from_ne_bytes(buf[32..36].try_into().unwrap()),
605 + // skip padding at 36..40
606 + session_id: u64::from_ne_bytes(buf[40..48].try_into().unwrap()),
607 + };
608 +
609 + if h.layout_version != 1 {
610 + return Err(NipcError::BadLayout);
611 + }
612 + if h.flags != 0 {
613 + return Err(NipcError::BadLayout);
614 + }
615 + Ok(h)
616 + }
617 +}
618 +
619 +// ===========================================================================
620 +// Tests
621 +// ===========================================================================
622 +
623 +#[cfg(test)]
624 +mod tests {
625 + use super::*;
626 +
627 + // -----------------------------------------------------------------------
628 + // Outer message header tests
629 + // -----------------------------------------------------------------------
630 +
631 + #[test]
632 + fn header_roundtrip() {
633 + let h = Header {
634 + magic: MAGIC_MSG,
635 + version: VERSION,
636 + header_len: HEADER_LEN,
637 + kind: KIND_REQUEST,
638 + flags: FLAG_BATCH,
639 + code: METHOD_CGROUPS_SNAPSHOT,
640 + transport_status: STATUS_OK,
641 + payload_len: 12345,
642 + item_count: 42,
643 + message_id: 0xDEAD_BEEF_CAFE_BABE,
644 + };
645 +
646 + let mut buf = [0u8; 64];
647 + let n = h.encode(&mut buf);
648 + assert_eq!(n, 32);
649 +
650 + let out = Header::decode(&buf[..n]).unwrap();
651 + assert_eq!(out, h);
652 + }
653 +
654 + #[test]
655 + fn header_encode_too_small() {
656 + let h = Header::default();
657 + let mut buf = [0u8; 16];
658 + assert_eq!(h.encode(&mut buf), 0);
659 + }
660 +
661 + #[test]
662 + fn header_decode_truncated() {
663 + let buf = [0u8; 31];
664 + assert_eq!(Header::decode(&buf), Err(NipcError::Truncated));
665 + }
666 +
667 + #[test]
668 + fn header_decode_bad_magic() {
669 + let h = Header {
670 + magic: 0x12345678,
671 + version: VERSION,
672 + header_len: HEADER_LEN,
673 + kind: KIND_REQUEST,
674 + ..Default::default()
675 + };
676 + let mut buf = [0u8; 32];
677 + h.encode(&mut buf);
678 + assert_eq!(Header::decode(&buf), Err(NipcError::BadMagic));
679 + }
680 +
681 + #[test]
682 + fn header_decode_bad_version() {
683 + let h = Header {
684 + magic: MAGIC_MSG,
685 + version: 99,
686 + header_len: HEADER_LEN,
687 + kind: KIND_REQUEST,
688 + ..Default::default()
689 + };
690 + let mut buf = [0u8; 32];
691 + h.encode(&mut buf);
692 + assert_eq!(Header::decode(&buf), Err(NipcError::BadVersion));
693 + }
694 +
695 + #[test]
696 + fn header_decode_bad_header_len() {
697 + let h = Header {
698 + magic: MAGIC_MSG,
699 + version: VERSION,
700 + header_len: 64,
701 + kind: KIND_REQUEST,
702 + ..Default::default()
703 + };
704 + let mut buf = [0u8; 32];
705 + h.encode(&mut buf);
706 + assert_eq!(Header::decode(&buf), Err(NipcError::BadHeaderLen));
707 + }
708 +
709 + #[test]
710 + fn header_decode_bad_kind() {
711 + // kind = 0
712 + let h = Header {
713 + magic: MAGIC_MSG,
714 + version: VERSION,
715 + header_len: HEADER_LEN,
716 + kind: 0,
717 + ..Default::default()
718 + };
719 + let mut buf = [0u8; 32];
720 + h.encode(&mut buf);
721 + assert_eq!(Header::decode(&buf), Err(NipcError::BadKind));
722 +
723 + // kind = 4
724 + let h2 = Header { kind: 4, ..h };
725 + h2.encode(&mut buf);
726 + assert_eq!(Header::decode(&buf), Err(NipcError::BadKind));
727 + }
728 +
729 + #[test]
730 + fn header_all_kinds() {
731 + for k in KIND_REQUEST..=KIND_CONTROL {
732 + let h = Header {
733 + magic: MAGIC_MSG,
734 + version: VERSION,
735 + header_len: HEADER_LEN,
736 + kind: k,
737 + ..Default::default()
738 + };
739 + let mut buf = [0u8; 32];
740 + h.encode(&mut buf);
741 + let out = Header::decode(&buf).unwrap();
742 + assert_eq!(out.kind, k);
743 + }
744 + }
745 +
746 + #[test]
747 + fn header_wire_bytes() {
748 + let h = Header {
749 + magic: MAGIC_MSG,
750 + version: VERSION,
751 + header_len: HEADER_LEN,
752 + kind: KIND_REQUEST,
753 + flags: 0,
754 + code: METHOD_CGROUPS_SNAPSHOT,
755 + transport_status: STATUS_OK,
756 + payload_len: 4,
757 + item_count: 1,
758 + message_id: 1,
759 + };
760 +
761 + let mut buf = [0u8; 32];
762 + h.encode(&mut buf);
763 +
764 + // magic = 0x4e495043 LE: 43 50 49 4e
765 + assert_eq!(&buf[0..4], &[0x43, 0x50, 0x49, 0x4e]);
766 + // version = 1 LE: 01 00
767 + assert_eq!(&buf[4..6], &[0x01, 0x00]);
768 + // header_len = 32 LE: 20 00
769 + assert_eq!(&buf[6..8], &[0x20, 0x00]);
770 + // kind = 1 LE: 01 00
771 + assert_eq!(&buf[8..10], &[0x01, 0x00]);
772 + // code = 2 LE: 02 00
773 + assert_eq!(&buf[12..14], &[0x02, 0x00]);
774 + }
775 +
776 + // -----------------------------------------------------------------------
777 + // Chunk continuation header tests
778 + // -----------------------------------------------------------------------
779 +
780 + #[test]
781 + fn chunk_header_roundtrip() {
782 + let c = ChunkHeader {
783 + magic: MAGIC_CHUNK,
784 + version: VERSION,
785 + flags: 0,
786 + message_id: 0x1234_5678_90AB_CDEF,
787 + total_message_len: 100000,
788 + chunk_index: 3,
789 + chunk_count: 10,
790 + chunk_payload_len: 8192,
791 + };
792 +
793 + let mut buf = [0u8; 64];
794 + let n = c.encode(&mut buf);
795 + assert_eq!(n, 32);
796 +
797 + let out = ChunkHeader::decode(&buf[..n]).unwrap();
798 + assert_eq!(out, c);
799 + }
800 +
801 + #[test]
802 + fn chunk_decode_truncated() {
803 + let buf = [0u8; 31];
804 + assert_eq!(ChunkHeader::decode(&buf), Err(NipcError::Truncated));
805 + }
806 +
807 + #[test]
808 + fn chunk_decode_bad_magic() {
809 + let c = ChunkHeader {
810 + magic: MAGIC_MSG, // wrong magic for chunk
811 + version: VERSION,
812 + ..Default::default()
813 + };
814 + let mut buf = [0u8; 32];
815 + c.encode(&mut buf);
816 + assert_eq!(ChunkHeader::decode(&buf), Err(NipcError::BadMagic));
817 + }
818 +
819 + #[test]
820 + fn chunk_decode_bad_version() {
821 + let c = ChunkHeader {
822 + magic: MAGIC_CHUNK,
823 + version: 2,
824 + ..Default::default()
825 + };
826 + let mut buf = [0u8; 32];
827 + c.encode(&mut buf);
828 + assert_eq!(ChunkHeader::decode(&buf), Err(NipcError::BadVersion));
829 + }
830 +
831 + #[test]
832 + fn chunk_encode_too_small() {
833 + let c = ChunkHeader::default();
834 + let mut buf = [0u8; 16];
835 + assert_eq!(c.encode(&mut buf), 0);
836 + }
837 +
838 + #[test]
839 + fn chunk_wire_bytes() {
840 + let c = ChunkHeader {
841 + magic: MAGIC_CHUNK,
842 + version: VERSION,
843 + flags: 0,
844 + message_id: 1,
845 + total_message_len: 256,
846 + chunk_index: 1,
847 + chunk_count: 3,
848 + chunk_payload_len: 100,
849 + };
850 +
851 + let mut buf = [0u8; 32];
852 + c.encode(&mut buf);
853 +
854 + // magic = 0x4e43484b LE: 4b 48 43 4e
855 + assert_eq!(&buf[0..4], &[0x4b, 0x48, 0x43, 0x4e]);
856 + }
857 +
858 + // -----------------------------------------------------------------------
859 + // Batch item directory tests
860 + // -----------------------------------------------------------------------
861 +
862 + #[test]
863 + fn batch_dir_roundtrip() {
864 + let entries = [
865 + BatchEntry {
866 + offset: 0,
867 + length: 100,
868 + },
869 + BatchEntry {
870 + offset: 104,
871 + length: 200,
872 + },
873 + BatchEntry {
874 + offset: 304,
875 + length: 50,
876 + },
877 + ];
878 +
879 + let mut buf = [0u8; 64];
880 + let n = batch_dir_encode(&entries, &mut buf);
881 + assert_eq!(n, 24);
882 +
883 + let out = batch_dir_decode(&buf[..n], 3, 400).unwrap();
884 + assert_eq!(out[0], entries[0]);
885 + assert_eq!(out[1], entries[1]);
886 + assert_eq!(out[2], entries[2]);
887 + }
888 +
889 + #[test]
890 + fn batch_dir_decode_truncated() {
891 + let buf = [0u8; 12];
892 + assert_eq!(batch_dir_decode(&buf, 2, 1000), Err(NipcError::Truncated));
893 + }
894 +
895 + #[test]
896 + fn batch_dir_decode_oob() {
897 + let e = BatchEntry {
898 + offset: 0,
899 + length: 200,
900 + };
901 + let mut buf = [0u8; 8];
902 + batch_dir_encode(&[e], &mut buf);
903 + assert_eq!(batch_dir_decode(&buf, 1, 100), Err(NipcError::OutOfBounds));
904 + }
905 +
906 + #[test]
907 + fn batch_dir_decode_bad_alignment() {
908 + let mut buf = [0u8; 8];
909 + // Manually write unaligned offset
910 + buf[0..4].copy_from_slice(&3u32.to_ne_bytes());
911 + buf[4..8].copy_from_slice(&10u32.to_ne_bytes());
912 + assert_eq!(batch_dir_decode(&buf, 1, 100), Err(NipcError::BadAlignment));
913 + }
914 +
915 + // -----------------------------------------------------------------------
916 + // Batch builder + extraction tests
917 + // -----------------------------------------------------------------------
918 +
919 + #[test]
920 + fn batch_builder_roundtrip() {
921 + let mut buf = [0u8; 1024];
922 + let mut b = BatchBuilder::new(&mut buf, 4);
923 +
924 + let item1 = [1u8, 2, 3, 4, 5];
925 + let item2 = [10u8, 20, 30];
926 + let item3 = [0xAAu8, 0xBB];
927 +
928 + b.add(&item1).unwrap();
929 + b.add(&item2).unwrap();
930 + b.add(&item3).unwrap();
931 +
932 + let (total, count) = b.finish();
933 + assert_eq!(count, 3);
934 + assert!(total > 0);
935 +
936 + // Extract items
937 + let (data, len) = batch_item_get(&buf[..total], 3, 0).unwrap();
938 + assert_eq!(len as usize, item1.len());
939 + assert_eq!(data, &item1);
940 +
941 + let (data, len) = batch_item_get(&buf[..total], 3, 1).unwrap();
942 + assert_eq!(len as usize, item2.len());
943 + assert_eq!(data, &item2);
944 +
945 + let (data, len) = batch_item_get(&buf[..total], 3, 2).unwrap();
946 + assert_eq!(len as usize, item3.len());
947 + assert_eq!(data, &item3);
948 + }
949 +
950 + #[test]
951 + fn batch_builder_overflow() {
952 + let mut buf = [0u8; 32];
953 + let mut b = BatchBuilder::new(&mut buf, 1);
954 + let item = [1u8];
955 + b.add(&item).unwrap();
956 + assert_eq!(b.add(&item), Err(NipcError::Overflow));
957 + }
958 +
959 + #[test]
960 + fn batch_builder_buf_overflow() {
961 + let mut buf = [0u8; 24];
962 + let mut b = BatchBuilder::new(&mut buf, 1);
963 + let big = [0u8; 100];
964 + assert_eq!(b.add(&big), Err(NipcError::Overflow));
965 + }
966 +
967 + #[test]
968 + fn batch_item_get_oob_index() {
969 + let mut buf = [0u8; 64];
970 + let mut b = BatchBuilder::new(&mut buf, 2);
971 + b.add(&[1u8]).unwrap();
972 + let (total, count) = b.finish();
973 + assert_eq!(
974 + batch_item_get(&buf[..total], count, 5),
975 + Err(NipcError::OutOfBounds)
976 + );
977 + }
978 +
979 + #[test]
980 + fn batch_empty() {
981 + let mut buf = [0u8; 64];
982 + let b = BatchBuilder::new(&mut buf, 4);
983 + let (total, count) = b.finish();
984 + assert_eq!(count, 0);
985 + assert_eq!(total, 0);
986 + }
987 +
988 + // -----------------------------------------------------------------------
989 + // Hello payload tests
990 + // -----------------------------------------------------------------------
991 +
992 + #[test]
993 + fn hello_roundtrip() {
994 + let h = Hello {
995 + layout_version: 1,
996 + flags: 0,
997 + supported_profiles: PROFILE_BASELINE | PROFILE_SHM_FUTEX,
998 + preferred_profiles: PROFILE_SHM_FUTEX,
999 + max_request_payload_bytes: 4096,
1000 + max_request_batch_items: 100,
1001 + max_response_payload_bytes: 1048576,
1002 + max_response_batch_items: 1,
1003 + auth_token: 0xAABB_CCDD_EEFF_0011,
1004 + packet_size: 65536,
1005 + };
1006 +
1007 + let mut buf = [0u8; 64];
1008 + let n = h.encode(&mut buf);
1009 + assert_eq!(n, 44);
1010 +
1011 + let out = Hello::decode(&buf[..n]).unwrap();
1012 + assert_eq!(out, h);
1013 + }
1014 +
1015 + #[test]
1016 + fn hello_decode_truncated() {
1017 + let buf = [0u8; 43];
1018 + assert_eq!(Hello::decode(&buf), Err(NipcError::Truncated));
1019 + }
1020 +
1021 + #[test]
1022 + fn hello_decode_bad_layout() {
1023 + let h = Hello {
1024 + layout_version: 99,
1025 + ..Default::default()
1026 + };
1027 + let mut buf = [0u8; 44];
1028 + h.encode(&mut buf);
1029 + assert_eq!(Hello::decode(&buf), Err(NipcError::BadLayout));
1030 + }
1031 +
1032 + #[test]
1033 + fn hello_encode_too_small() {
1034 + let h = Hello::default();
1035 + let mut buf = [0u8; 10];
1036 + assert_eq!(h.encode(&mut buf), 0);
1037 + }
1038 +
1039 + // -----------------------------------------------------------------------
1040 + // Hello-ack payload tests
1041 + // -----------------------------------------------------------------------
1042 +
1043 + #[test]
1044 + fn hello_ack_roundtrip() {
1045 + let h = HelloAck {
1046 + layout_version: 1,
1047 + flags: 0,
1048 + server_supported_profiles: 0x07,
1049 + intersection_profiles: 0x05,
1050 + selected_profile: PROFILE_SHM_FUTEX,
1051 + agreed_max_request_payload_bytes: 2048,
1052 + agreed_max_request_batch_items: 50,
1053 + agreed_max_response_payload_bytes: 65536,
1054 + agreed_max_response_batch_items: 1,
1055 + agreed_packet_size: 32768,
1056 + session_id: 42,
1057 + };
1058 +
1059 + let mut buf = [0u8; 64];
1060 + let n = h.encode(&mut buf);
1061 + assert_eq!(n, 48);
1062 +
1063 + let out = HelloAck::decode(&buf[..n]).unwrap();
1064 + assert_eq!(out, h);
1065 + }
1066 +
1067 + #[test]
1068 + fn hello_ack_decode_truncated() {
1069 + let buf = [0u8; 47];
1070 + assert_eq!(HelloAck::decode(&buf), Err(NipcError::Truncated));
1071 + }
1072 +
1073 + #[test]
1074 + fn hello_ack_decode_bad_layout() {
1075 + let h = HelloAck {
1076 + layout_version: 0,
1077 + ..Default::default()
1078 + };
1079 + let mut buf = [0u8; 48];
1080 + h.encode(&mut buf);
1081 + assert_eq!(HelloAck::decode(&buf), Err(NipcError::BadLayout));
1082 + }
1083 +
1084 + #[test]
1085 + fn hello_ack_encode_too_small() {
1086 + let h = HelloAck::default();
1087 + let mut buf = [0u8; 10];
1088 + assert_eq!(h.encode(&mut buf), 0);
1089 + }
1090 +
1091 + // -----------------------------------------------------------------------
1092 + // Cgroups snapshot request tests
1093 + // -----------------------------------------------------------------------
1094 +
1095 + #[test]
1096 + fn cgroups_req_roundtrip() {
1097 + let r = CgroupsRequest {
1098 + layout_version: 1,
1099 + flags: 0,
1100 + };
1101 +
1102 + let mut buf = [0u8; 16];
1103 + let n = r.encode(&mut buf);
1104 + assert_eq!(n, 4);
1105 +
1106 + let out = CgroupsRequest::decode(&buf[..n]).unwrap();
1107 + assert_eq!(out, r);
1108 + }
1109 +
1110 + #[test]
1111 + fn cgroups_req_decode_truncated() {
1112 + let buf = [0u8; 3];
1113 + assert_eq!(CgroupsRequest::decode(&buf), Err(NipcError::Truncated));
1114 + }
1115 +
1116 + #[test]
1117 + fn cgroups_req_decode_bad_layout() {
1118 + let r = CgroupsRequest {
1119 + layout_version: 5,
1120 + flags: 0,
1121 + };
1122 + let mut buf = [0u8; 4];
1123 + r.encode(&mut buf);
1124 + assert_eq!(CgroupsRequest::decode(&buf), Err(NipcError::BadLayout));
1125 + }
1126 +
1127 + #[test]
1128 + fn cgroups_req_encode_too_small() {
1129 + let r = CgroupsRequest::default();
1130 + let mut buf = [0u8; 2];
1131 + assert_eq!(r.encode(&mut buf), 0);
1132 + }
1133 +
1134 + // -----------------------------------------------------------------------
1135 + // Cgroups snapshot response tests
1136 + // -----------------------------------------------------------------------
1137 +
1138 + // Private constants needed by tests -- mirror the values from cgroups.rs
1139 + const CGROUPS_RESP_HDR_SIZE: usize = 24;
1140 + const CGROUPS_DIR_ENTRY_SIZE: usize = 8;
1141 +
1142 + #[test]
1143 + fn cgroups_resp_empty() {
1144 + let mut buf = [0u8; 4096];
1145 + let b = CgroupsBuilder::new(&mut buf, 0, 1, 42);
1146 + let total = b.finish();
1147 + assert_eq!(total, 24);
1148 +
1149 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
1150 + assert_eq!(view.item_count, 0);
1151 + assert_eq!(view.systemd_enabled, 1);
1152 + assert_eq!(view.generation, 42);
1153 + }
1154 +
1155 + #[test]
1156 + fn cgroups_resp_single_item() {
1157 + let mut buf = [0u8; 4096];
1158 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 100);
1159 +
1160 + let name = b"docker-abc123";
1161 + let path = b"/sys/fs/cgroup/docker/abc123";
1162 + b.add(12345, 0x01, 1, name, path).unwrap();
1163 +
1164 + let total = b.finish();
1165 + assert!(total > 24);
1166 +
1167 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
1168 + assert_eq!(view.item_count, 1);
1169 + assert_eq!(view.systemd_enabled, 0);
1170 + assert_eq!(view.generation, 100);
1171 +
1172 + let item = view.item(0).unwrap();
1173 + assert_eq!(item.hash, 12345);
1174 + assert_eq!(item.options, 0x01);
1175 + assert_eq!(item.enabled, 1);
1176 + assert_eq!(item.name.len as usize, name.len());
1177 + assert_eq!(item.name.as_bytes(), name);
1178 + assert_eq!(item.name.bytes[name.len()], 0); // NUL
1179 + assert_eq!(item.path.len as usize, path.len());
1180 + assert_eq!(item.path.as_bytes(), path);
1181 + assert_eq!(item.path.bytes[path.len()], 0); // NUL
1182 + }
1183 +
1184 + #[test]
1185 + fn cgroups_resp_multiple_items() {
1186 + let mut buf = [0u8; 8192];
1187 + let mut b = CgroupsBuilder::new(&mut buf, 5, 1, 999);
1188 +
1189 + // Item 0
1190 + let n0 = b"init.scope";
1191 + let p0 = b"/sys/fs/cgroup/init.scope";
1192 + b.add(100, 0, 1, n0, p0).unwrap();
1193 +
1194 + // Item 1
1195 + let n1 = b"system.slice/docker-abc.scope";
1196 + let p1 = b"/sys/fs/cgroup/system.slice/docker-abc.scope";
1197 + b.add(200, 0x02, 0, n1, p1).unwrap();
1198 +
1199 + // Item 2 - empty strings
1200 + b.add(300, 0, 1, b"", b"").unwrap();
1201 +
1202 + let total = b.finish();
1203 +
1204 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
1205 + assert_eq!(view.item_count, 3);
1206 + assert_eq!(view.systemd_enabled, 1);
1207 + assert_eq!(view.generation, 999);
1208 +
1209 + // Verify item 0
1210 + let item = view.item(0).unwrap();
1211 + assert_eq!(item.hash, 100);
1212 + assert_eq!(item.name.len as usize, n0.len());
1213 + assert_eq!(item.name.as_bytes(), n0);
1214 + assert_eq!(item.path.len as usize, p0.len());
1215 + assert_eq!(item.path.as_bytes(), p0);
1216 +
1217 + // Verify item 1
1218 + let item = view.item(1).unwrap();
1219 + assert_eq!(item.hash, 200);
1220 + assert_eq!(item.options, 0x02);
1221 + assert_eq!(item.enabled, 0);
1222 + assert_eq!(item.name.len as usize, n1.len());
1223 + assert_eq!(item.name.as_bytes(), n1);
1224 +
1225 + // Verify item 2 (empty strings)
1226 + let item = view.item(2).unwrap();
1227 + assert_eq!(item.hash, 300);
1228 + assert_eq!(item.name.len, 0);
1229 + assert_eq!(item.name.bytes[0], 0); // NUL
1230 + assert_eq!(item.path.len, 0);
1231 + assert_eq!(item.path.bytes[0], 0); // NUL
1232 +
1233 + // Out-of-bounds index
1234 + assert_eq!(view.item(3), Err(NipcError::OutOfBounds));
1235 + }
1236 +
1237 + #[test]
1238 + fn cgroups_resp_decode_truncated_header() {
1239 + let buf = [0u8; 23];
1240 + assert_eq!(
1241 + CgroupsResponseView::decode(&buf).unwrap_err(),
1242 + NipcError::Truncated
1243 + );
1244 + }
1245 +
1246 + #[test]
1247 + fn cgroups_resp_decode_bad_layout() {
1248 + let mut buf = [0u8; 24];
1249 + buf[0..2].copy_from_slice(&99u16.to_ne_bytes());
1250 + assert_eq!(
1251 + CgroupsResponseView::decode(&buf).unwrap_err(),
1252 + NipcError::BadLayout
1253 + );
1254 + }
1255 +
1256 + #[test]
1257 + fn cgroups_resp_decode_truncated_dir() {
1258 + // Header says item_count=2 but payload is only 24 bytes
1259 + let mut buf = [0u8; 24];
1260 + buf[0..2].copy_from_slice(&1u16.to_ne_bytes());
1261 + buf[4..8].copy_from_slice(&2u32.to_ne_bytes());
1262 + assert_eq!(
1263 + CgroupsResponseView::decode(&buf).unwrap_err(),
1264 + NipcError::Truncated
1265 + );
1266 + }
1267 +
1268 + #[test]
1269 + fn cgroups_resp_decode_oob_dir() {
1270 + // Header + 1 dir entry pointing beyond payload
1271 + let mut buf = [0u8; 64];
1272 + buf[0..2].copy_from_slice(&1u16.to_ne_bytes());
1273 + buf[4..8].copy_from_slice(&1u32.to_ne_bytes());
1274 + // Dir entry at offset 24: offset=0, length=9999
1275 + buf[24..28].copy_from_slice(&0u32.to_ne_bytes());
1276 + buf[28..32].copy_from_slice(&9999u32.to_ne_bytes());
1277 + assert_eq!(
1278 + CgroupsResponseView::decode(&buf).unwrap_err(),
1279 + NipcError::OutOfBounds
1280 + );
1281 + }
1282 +
1283 + #[test]
1284 + fn cgroups_resp_decode_item_too_small() {
1285 + // Dir entry with length < 32
1286 + let mut buf = [0u8; 64];
1287 + buf[0..2].copy_from_slice(&1u16.to_ne_bytes());
1288 + buf[4..8].copy_from_slice(&1u32.to_ne_bytes());
1289 + buf[24..28].copy_from_slice(&0u32.to_ne_bytes());
1290 + buf[28..32].copy_from_slice(&16u32.to_ne_bytes());
1291 + assert_eq!(
1292 + CgroupsResponseView::decode(&buf).unwrap_err(),
1293 + NipcError::Truncated
1294 + );
1295 + }
1296 +
1297 + #[test]
1298 + fn cgroups_resp_item_missing_nul() {
1299 + // Build valid snapshot then corrupt the NUL terminator
1300 + let mut buf = [0u8; 4096];
1301 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 1);
1302 + b.add(1, 0, 1, b"test", b"/test").unwrap();
1303 + let total = b.finish();
1304 +
1305 + // Find item data and corrupt the name's NUL terminator
1306 + let dir_end = CGROUPS_RESP_HDR_SIZE + 1 * CGROUPS_DIR_ENTRY_SIZE;
1307 + let item_off = u32::from_ne_bytes(
1308 + buf[CGROUPS_RESP_HDR_SIZE..CGROUPS_RESP_HDR_SIZE + 4]
1309 + .try_into()
1310 + .unwrap(),
1311 + ) as usize;
1312 + let item_start = dir_end + item_off;
1313 +
1314 + let noff =
1315 + u32::from_ne_bytes(buf[item_start + 16..item_start + 20].try_into().unwrap()) as usize;
1316 + let nlen =
1317 + u32::from_ne_bytes(buf[item_start + 20..item_start + 24].try_into().unwrap()) as usize;
1318 +
1319 + buf[item_start + noff + nlen] = b'X'; // corrupt NUL
1320 +
1321 + // Re-decode after corruption -- header/dir still valid
1322 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
1323 + assert_eq!(view.item(0).unwrap_err(), NipcError::MissingNul);
1324 + }
1325 +
1326 + #[test]
1327 + fn cgroups_resp_item_string_oob() {
1328 + // Build valid snapshot then corrupt string length to be huge
1329 + let mut buf = [0u8; 4096];
1330 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 1);
1331 + b.add(1, 0, 1, b"test", b"/test").unwrap();
1332 + let total = b.finish();
1333 +
1334 + // Corrupt name_length to huge value
1335 + let dir_end = CGROUPS_RESP_HDR_SIZE + 1 * CGROUPS_DIR_ENTRY_SIZE;
1336 + let item_off = u32::from_ne_bytes(
1337 + buf[CGROUPS_RESP_HDR_SIZE..CGROUPS_RESP_HDR_SIZE + 4]
1338 + .try_into()
1339 + .unwrap(),
1340 + ) as usize;
1341 + let item_start = dir_end + item_off;
1342 +
1343 + buf[item_start + 20..item_start + 24].copy_from_slice(&99999u32.to_ne_bytes());
1344 +
1345 + // Re-decode after corruption
1346 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
1347 + assert_eq!(view.item(0).unwrap_err(), NipcError::OutOfBounds);
1348 + }
1349 +
1350 + #[test]
1351 + fn cgroups_builder_overflow() {
1352 + let mut buf = [0u8; 64]; // too small for any real item
1353 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 0);
1354 + let long_name = [b'A'; 200];
1355 + assert_eq!(b.add(1, 0, 1, &long_name, b""), Err(NipcError::Overflow));
1356 + }
1357 +
1358 + #[test]
1359 + fn cgroups_builder_max_items_exceeded() {
1360 + let mut buf = [0u8; 4096];
1361 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 0);
1362 + b.add(1, 0, 1, b"a", b"b").unwrap();
1363 + assert_eq!(b.add(2, 0, 1, b"c", b"d"), Err(NipcError::Overflow));
1364 + }
1365 +
1366 + #[test]
1367 + fn cgroups_builder_compaction() {
1368 + let mut buf = [0u8; 4096];
1369 + // Reserve 10 directory slots but only add 2 items
1370 + let mut b = CgroupsBuilder::new(&mut buf, 10, 1, 77);
1371 +
1372 + b.add(10, 0, 1, b"slice-a", b"/cgroup/slice-a").unwrap();
1373 + b.add(20, 0, 0, b"slice-b", b"/cgroup/slice-b").unwrap();
1374 +
1375 + let total = b.finish();
1376 +
1377 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
1378 + assert_eq!(view.item_count, 2);
1379 + assert_eq!(view.generation, 77);
1380 +
1381 + let item = view.item(0).unwrap();
1382 + assert_eq!(item.hash, 10);
1383 + assert_eq!(item.name.as_bytes(), b"slice-a");
1384 +
1385 + let item = view.item(1).unwrap();
1386 + assert_eq!(item.hash, 20);
1387 + assert_eq!(item.name.as_bytes(), b"slice-b");
1388 + }
1389 +
1390 + // -----------------------------------------------------------------------
1391 + // Alignment utility test
1392 + // -----------------------------------------------------------------------
1393 +
1394 + #[test]
1395 + fn test_align8() {
1396 + assert_eq!(align8(0), 0);
1397 + assert_eq!(align8(1), 8);
1398 + assert_eq!(align8(7), 8);
1399 + assert_eq!(align8(8), 8);
1400 + assert_eq!(align8(9), 16);
1401 + assert_eq!(align8(16), 16);
1402 + assert_eq!(align8(17), 24);
1403 + }
1404 +
1405 + // -----------------------------------------------------------------------
1406 + // Cross-language wire compatibility: C-Rust byte identity
1407 + //
1408 + // These tests encode in Rust and verify the exact bytes match what the
1409 + // C implementation produces for the same inputs. This ensures identical
1410 + // wire output across languages.
1411 + // -----------------------------------------------------------------------
1412 +
1413 + #[test]
1414 + fn c_rust_header_bytes_identical() {
1415 + // Encode in Rust
1416 + let h = Header {
1417 + magic: MAGIC_MSG,
1418 + version: VERSION,
1419 + header_len: HEADER_LEN,
1420 + kind: KIND_REQUEST,
1421 + flags: FLAG_BATCH,
1422 + code: METHOD_CGROUPS_SNAPSHOT,
1423 + transport_status: STATUS_OK,
1424 + payload_len: 12345,
1425 + item_count: 42,
1426 + message_id: 0xDEAD_BEEF_CAFE_BABE,
1427 + };
1428 + let mut rust_buf = [0u8; 32];
1429 + h.encode(&mut rust_buf);
1430 +
1431 + // Known LE bytes for this header
1432 + let expected: [u8; 32] = [
1433 + 0x43, 0x50, 0x49, 0x4e, // magic
1434 + 0x01, 0x00, // version
1435 + 0x20, 0x00, // header_len
1436 + 0x01, 0x00, // kind
1437 + 0x01, 0x00, // flags
1438 + 0x02, 0x00, // code
1439 + 0x00, 0x00, // transport_status
1440 + 0x39, 0x30, 0x00, 0x00, // payload_len = 12345
1441 + 0x2a, 0x00, 0x00, 0x00, // item_count = 42
1442 + 0xbe, 0xba, 0xfe, 0xca, 0xef, 0xbe, 0xad, 0xde, // message_id
1443 + ];
1444 + assert_eq!(rust_buf, expected);
1445 + }
1446 +
1447 + #[test]
1448 + fn c_rust_chunk_bytes_identical() {
1449 + let c = ChunkHeader {
1450 + magic: MAGIC_CHUNK,
1451 + version: VERSION,
1452 + flags: 0,
1453 + message_id: 1,
1454 + total_message_len: 256,
1455 + chunk_index: 1,
1456 + chunk_count: 3,
1457 + chunk_payload_len: 100,
1458 + };
1459 + let mut rust_buf = [0u8; 32];
1460 + c.encode(&mut rust_buf);
1461 +
1462 + let expected: [u8; 32] = [
1463 + 0x4b, 0x48, 0x43, 0x4e, // magic
1464 + 0x01, 0x00, // version
1465 + 0x00, 0x00, // flags
1466 + 0x01, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, // message_id
1467 + 0x00, 0x01, 0x00, 0x00, // total_message_len = 256
1468 + 0x01, 0x00, 0x00, 0x00, // chunk_index
1469 + 0x03, 0x00, 0x00, 0x00, // chunk_count
1470 + 0x64, 0x00, 0x00, 0x00, // chunk_payload_len = 100
1471 + ];
1472 + assert_eq!(rust_buf, expected);
1473 + }
1474 +
1475 + #[test]
1476 + fn c_rust_hello_bytes_identical() {
1477 + let h = Hello {
1478 + layout_version: 1,
1479 + flags: 0,
1480 + supported_profiles: PROFILE_BASELINE | PROFILE_SHM_FUTEX,
1481 + preferred_profiles: PROFILE_SHM_FUTEX,
1482 + max_request_payload_bytes: 4096,
1483 + max_request_batch_items: 100,
1484 + max_response_payload_bytes: 1048576,
1485 + max_response_batch_items: 1,
1486 + auth_token: 0xAABB_CCDD_EEFF_0011,
1487 + packet_size: 65536,
1488 + };
1489 +
1490 + let mut rust_buf = [0u8; 44];
1491 + h.encode(&mut rust_buf);
1492 +
1493 + // Verify key byte positions
1494 + assert_eq!(&rust_buf[0..2], &[0x01, 0x00]); // layout_version
1495 + assert_eq!(&rust_buf[2..4], &[0x00, 0x00]); // flags
1496 + assert_eq!(&rust_buf[4..8], &[0x05, 0x00, 0x00, 0x00]); // supported = 0x05
1497 + assert_eq!(&rust_buf[8..12], &[0x04, 0x00, 0x00, 0x00]); // preferred = 0x04
1498 + assert_eq!(&rust_buf[28..32], &[0x00, 0x00, 0x00, 0x00]); // padding = 0
1499 + assert_eq!(
1500 + &rust_buf[32..40],
1501 + &[0x11, 0x00, 0xFF, 0xEE, 0xDD, 0xCC, 0xBB, 0xAA]
1502 + ); // auth_token
1503 +
1504 + // Round-trip
1505 + let out = Hello::decode(&rust_buf).unwrap();
1506 + assert_eq!(out, h);
1507 + }
1508 +
1509 + #[test]
1510 + fn c_rust_hello_ack_bytes_identical() {
1511 + let h = HelloAck {
1512 + layout_version: 1,
1513 + flags: 0,
1514 + server_supported_profiles: 0x07,
1515 + intersection_profiles: 0x05,
1516 + selected_profile: PROFILE_SHM_FUTEX,
1517 + agreed_max_request_payload_bytes: 2048,
1518 + agreed_max_request_batch_items: 50,
1519 + agreed_max_response_payload_bytes: 65536,
1520 + agreed_max_response_batch_items: 1,
1521 + agreed_packet_size: 32768,
1522 + session_id: 0x0000_0001_0000_0007,
1523 + };
1524 + let mut rust_buf = [0u8; 48];
1525 + h.encode(&mut rust_buf);
1526 +
1527 + assert_eq!(&rust_buf[0..2], &[0x01, 0x00]);
1528 + assert_eq!(&rust_buf[4..8], &[0x07, 0x00, 0x00, 0x00]); // server_supported
1529 + assert_eq!(&rust_buf[12..16], &[0x04, 0x00, 0x00, 0x00]); // selected = SHM_FUTEX
1530 + assert_eq!(&rust_buf[36..40], &[0x00, 0x00, 0x00, 0x00]); // padding
1531 + assert_eq!(
1532 + &rust_buf[40..48],
1533 + &[0x07, 0x00, 0x00, 0x00, 0x01, 0x00, 0x00, 0x00]
1534 + ); // session_id LE
1535 +
1536 + let out = HelloAck::decode(&rust_buf).unwrap();
1537 + assert_eq!(out, h);
1538 + }
1539 +
1540 + #[test]
1541 + fn c_rust_cgroups_req_bytes_identical() {
1542 + let r = CgroupsRequest {
1543 + layout_version: 1,
1544 + flags: 0,
1545 + };
1546 + let mut rust_buf = [0u8; 4];
1547 + r.encode(&mut rust_buf);
1548 +
1549 + assert_eq!(rust_buf, [0x01, 0x00, 0x00, 0x00]);
1550 +
1551 + let out = CgroupsRequest::decode(&rust_buf).unwrap();
1552 + assert_eq!(out, r);
1553 + }
1554 +
1555 + #[test]
1556 + fn c_rust_cgroups_snapshot_bytes_identical() {
1557 + // Build a snapshot with the exact same inputs as the C test
1558 + let mut buf = [0u8; 4096];
1559 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 100);
1560 + b.add(
1561 + 12345,
1562 + 0x01,
1563 + 1,
1564 + b"docker-abc123",
1565 + b"/sys/fs/cgroup/docker/abc123",
1566 + )
1567 + .unwrap();
1568 + let total = b.finish();
1569 +
1570 + // Verify the snapshot header bytes
1571 + assert_eq!(&buf[0..2], &[0x01, 0x00]); // layout_version
1572 + assert_eq!(&buf[2..4], &[0x00, 0x00]); // flags
1573 + assert_eq!(&buf[4..8], &[0x01, 0x00, 0x00, 0x00]); // item_count
1574 + assert_eq!(&buf[8..12], &[0x00, 0x00, 0x00, 0x00]); // systemd_enabled
1575 + assert_eq!(&buf[12..16], &[0x00, 0x00, 0x00, 0x00]); // reserved
1576 + assert_eq!(
1577 + &buf[16..24],
1578 + &[0x64, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00, 0x00]
1579 + ); // generation
1580 +
1581 + // Verify it decodes correctly
1582 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
1583 + assert_eq!(view.item_count, 1);
1584 + assert_eq!(view.generation, 100);
1585 +
1586 + let item = view.item(0).unwrap();
1587 + assert_eq!(item.hash, 12345);
1588 + assert_eq!(item.name.as_bytes(), b"docker-abc123");
1589 + assert_eq!(item.path.as_bytes(), b"/sys/fs/cgroup/docker/abc123");
1590 + }
1591 +
1592 + #[test]
1593 + fn cgroups_resp_dir_bad_alignment() {
1594 + // Dir entry with unaligned offset
1595 + let mut buf = [0u8; 128];
1596 + buf[0..2].copy_from_slice(&1u16.to_ne_bytes());
1597 + buf[4..8].copy_from_slice(&1u32.to_ne_bytes());
1598 + // offset=3 (not 8-byte aligned), length=32
1599 + buf[24..28].copy_from_slice(&3u32.to_ne_bytes());
1600 + buf[28..32].copy_from_slice(&32u32.to_ne_bytes());
1601 + assert_eq!(
1602 + CgroupsResponseView::decode(&buf).unwrap_err(),
1603 + NipcError::BadAlignment
1604 + );
1605 + }
1606 +
1607 + #[test]
1608 + fn cgroups_resp_item_bad_layout_version() {
1609 + // Build valid snapshot, then corrupt the item's layout_version
1610 + let mut buf = [0u8; 4096];
1611 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 1);
1612 + b.add(1, 0, 1, b"test", b"/test").unwrap();
1613 + let total = b.finish();
1614 +
1615 + // Corrupt item layout_version
1616 + let dir_end = CGROUPS_RESP_HDR_SIZE + 1 * CGROUPS_DIR_ENTRY_SIZE;
1617 + let item_off = u32::from_ne_bytes(
1618 + buf[CGROUPS_RESP_HDR_SIZE..CGROUPS_RESP_HDR_SIZE + 4]
1619 + .try_into()
1620 + .unwrap(),
1621 + ) as usize;
1622 + let item_start = dir_end + item_off;
1623 + buf[item_start..item_start + 2].copy_from_slice(&99u16.to_ne_bytes());
1624 +
1625 + // Re-decode after corruption
1626 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
1627 + assert_eq!(view.item(0).unwrap_err(), NipcError::BadLayout);
1628 + }
1629 +
1630 + #[test]
1631 + fn cgroups_resp_item_name_off_below_header() {
1632 + // Build valid snapshot, then set name_offset < 32
1633 + let mut buf = [0u8; 4096];
1634 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 1);
1635 + b.add(1, 0, 1, b"test", b"/test").unwrap();
1636 + let total = b.finish();
1637 +
1638 + let dir_end = CGROUPS_RESP_HDR_SIZE + 1 * CGROUPS_DIR_ENTRY_SIZE;
1639 + let item_off = u32::from_ne_bytes(
1640 + buf[CGROUPS_RESP_HDR_SIZE..CGROUPS_RESP_HDR_SIZE + 4]
1641 + .try_into()
1642 + .unwrap(),
1643 + ) as usize;
1644 + let item_start = dir_end + item_off;
1645 + // Set name_offset to 0 (below header)
1646 + buf[item_start + 16..item_start + 20].copy_from_slice(&0u32.to_ne_bytes());
1647 +
1648 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
1649 + assert_eq!(view.item(0).unwrap_err(), NipcError::OutOfBounds);
1650 + }
1651 +
1652 + #[test]
1653 + fn cgroups_resp_item_path_off_below_header() {
1654 + let mut buf = [0u8; 4096];
1655 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 1);
1656 + b.add(1, 0, 1, b"test", b"/test").unwrap();
1657 + let total = b.finish();
1658 +
1659 + let dir_end = CGROUPS_RESP_HDR_SIZE + 1 * CGROUPS_DIR_ENTRY_SIZE;
1660 + let item_off = u32::from_ne_bytes(
1661 + buf[CGROUPS_RESP_HDR_SIZE..CGROUPS_RESP_HDR_SIZE + 4]
1662 + .try_into()
1663 + .unwrap(),
1664 + ) as usize;
1665 + let item_start = dir_end + item_off;
1666 + // Set path_offset to 16 (below header)
1667 + buf[item_start + 24..item_start + 28].copy_from_slice(&16u32.to_ne_bytes());
1668 +
1669 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
1670 + assert_eq!(view.item(0).unwrap_err(), NipcError::OutOfBounds);
1671 + }
1672 +
1673 + #[test]
1674 + fn cgroups_resp_item_path_missing_nul() {
1675 + let mut buf = [0u8; 4096];
1676 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 1);
1677 + b.add(1, 0, 1, b"test", b"/test").unwrap();
1678 + let total = b.finish();
1679 +
1680 + // Corrupt path NUL
1681 + let dir_end = CGROUPS_RESP_HDR_SIZE + 1 * CGROUPS_DIR_ENTRY_SIZE;
1682 + let item_off = u32::from_ne_bytes(
1683 + buf[CGROUPS_RESP_HDR_SIZE..CGROUPS_RESP_HDR_SIZE + 4]
1684 + .try_into()
1685 + .unwrap(),
1686 + ) as usize;
1687 + let item_start = dir_end + item_off;
1688 + let poff =
1689 + u32::from_ne_bytes(buf[item_start + 24..item_start + 28].try_into().unwrap()) as usize;
1690 + let plen =
1691 + u32::from_ne_bytes(buf[item_start + 28..item_start + 32].try_into().unwrap()) as usize;
1692 + buf[item_start + poff + plen] = b'X';
1693 +
1694 + // Re-decode after corruption
1695 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
1696 + assert_eq!(view.item(0).unwrap_err(), NipcError::MissingNul);
1697 + }
1698 +
1699 + #[test]
1700 + fn cgroups_resp_item_overlap_rejected() {
1701 + // Build a valid item, then manually set path_offset to overlap with name
1702 + let mut buf = [0u8; 4096];
1703 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 1);
1704 + b.add(1, 0, 1, b"hello", b"/path").unwrap();
1705 + let total = b.finish();
1706 +
1707 + let dir_end = CGROUPS_RESP_HDR_SIZE + 1 * CGROUPS_DIR_ENTRY_SIZE;
1708 + let item_off = u32::from_ne_bytes(
1709 + buf[CGROUPS_RESP_HDR_SIZE..CGROUPS_RESP_HDR_SIZE + 4]
1710 + .try_into()
1711 + .unwrap(),
1712 + ) as usize;
1713 + let item_start = dir_end + item_off;
1714 +
1715 + // name_off=32, name_len=5, so name region is [32..38)
1716 + // Set path_off=34 (inside name region), path_len=1
1717 + buf[item_start + 24..item_start + 28].copy_from_slice(&34u32.to_ne_bytes());
1718 + buf[item_start + 28..item_start + 32].copy_from_slice(&1u32.to_ne_bytes());
1719 + // Ensure NUL at item[34+1]=item[35]
1720 + buf[item_start + 35] = 0;
1721 +
1722 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
1723 + assert_eq!(view.item(0).unwrap_err(), NipcError::BadLayout);
1724 + }
1725 +
1726 + // -------------------------------------------------------------------
1727 + // Proptest: fuzz / property-based tests for all decode paths
1728 + // -------------------------------------------------------------------
1729 +
1730 + mod proptests {
1731 + use super::*;
1732 + use proptest::prelude::*;
1733 +
1734 + // Arbitrary bytes -- no decode path may panic on any input.
1735 +
1736 + proptest! {
1737 + #[test]
1738 + fn decode_header_never_panics(data: Vec<u8>) {
1739 + let _ = Header::decode(&data);
1740 + }
1741 +
1742 + #[test]
1743 + fn decode_chunk_header_never_panics(data: Vec<u8>) {
1744 + let _ = ChunkHeader::decode(&data);
1745 + }
1746 +
1747 + #[test]
1748 + fn decode_hello_never_panics(data: Vec<u8>) {
1749 + let _ = Hello::decode(&data);
1750 + }
1751 +
1752 + #[test]
1753 + fn decode_hello_ack_never_panics(data: Vec<u8>) {
1754 + let _ = HelloAck::decode(&data);
1755 + }
1756 +
1757 + #[test]
1758 + fn decode_cgroups_request_never_panics(data: Vec<u8>) {
1759 + let _ = CgroupsRequest::decode(&data);
1760 + }
1761 +
1762 + #[test]
1763 + fn decode_cgroups_response_never_panics(data: Vec<u8>) {
1764 + let result = CgroupsResponseView::decode(&data);
1765 + if let Ok(view) = result {
1766 + // Exercise item access on valid decodes.
1767 + let limit = view.item_count.min(64);
1768 + for i in 0..limit {
1769 + let _ = view.item(i);
1770 + }
1771 + // Out-of-bounds must not panic.
1772 + let _ = view.item(view.item_count);
1773 + }
1774 + }
1775 +
1776 + #[test]
1777 + fn batch_dir_decode_never_panics(
1778 + data: Vec<u8>,
1779 + item_count in 0u32..128,
1780 + packed_area_len in 0u32..65536,
1781 + ) {
1782 + let _ = batch_dir_decode(&data, item_count, packed_area_len);
1783 + }
1784 +
1785 + #[test]
1786 + fn batch_item_get_never_panics(
1787 + data: Vec<u8>,
1788 + item_count in 0u32..128,
1789 + index in 0u32..128,
1790 + ) {
1791 + let _ = batch_item_get(&data, item_count, index);
1792 + }
1793 + }
1794 +
1795 + // Roundtrip tests: encode random valid values, decode, verify match.
1796 +
1797 + proptest! {
1798 + #[test]
1799 + fn encode_decode_header_roundtrip(
1800 + kind in 1u16..=3,
1801 + flags in any::<u16>(),
1802 + code in any::<u16>(),
1803 + transport_status in any::<u16>(),
1804 + payload_len in any::<u32>(),
1805 + item_count in any::<u32>(),
1806 + message_id in any::<u64>(),
1807 + ) {
1808 + let h = Header {
1809 + magic: MAGIC_MSG,
1810 + version: VERSION,
1811 + header_len: HEADER_LEN,
1812 + kind,
1813 + flags,
1814 + code,
1815 + transport_status,
1816 + payload_len,
1817 + item_count,
1818 + message_id,
1819 + };
1820 + let mut buf = [0u8; 64];
1821 + let n = h.encode(&mut buf);
1822 + prop_assert_eq!(n, HEADER_SIZE);
1823 + let decoded = Header::decode(&buf[..n]).unwrap();
1824 + prop_assert_eq!(decoded, h);
1825 + }
1826 +
1827 + #[test]
1828 + fn encode_decode_hello_roundtrip(
1829 + supported in any::<u32>(),
1830 + preferred in any::<u32>(),
1831 + max_req_payload in any::<u32>(),
1832 + max_req_batch in any::<u32>(),
1833 + max_resp_payload in any::<u32>(),
1834 + max_resp_batch in any::<u32>(),
1835 + auth_token in any::<u64>(),
1836 + packet_size in any::<u32>(),
1837 + ) {
1838 + let h = Hello {
1839 + layout_version: 1,
1840 + flags: 0,
1841 + supported_profiles: supported,
1842 + preferred_profiles: preferred,
1843 + max_request_payload_bytes: max_req_payload,
1844 + max_request_batch_items: max_req_batch,
1845 + max_response_payload_bytes: max_resp_payload,
1846 + max_response_batch_items: max_resp_batch,
1847 + auth_token,
1848 + packet_size,
1849 + };
1850 + let mut buf = [0u8; 64];
1851 + let n = h.encode(&mut buf);
1852 + prop_assert_eq!(n, HELLO_SIZE);
1853 + let decoded = Hello::decode(&buf[..n]).unwrap();
1854 + prop_assert_eq!(decoded, h);
1855 + }
1856 + }
1857 + }
1858 +
1859 + // -----------------------------------------------------------------------
1860 + // NipcError Display coverage
1861 + // -----------------------------------------------------------------------
1862 +
1863 + #[test]
1864 + fn nipc_error_display_all_variants() {
1865 + // Exercise the Display impl for every NipcError variant (lines 101-113)
1866 + let cases: Vec<(NipcError, &str)> = vec![
1867 + (NipcError::Truncated, "buffer too short"),
1868 + (NipcError::BadMagic, "magic value mismatch"),
1869 + (NipcError::BadVersion, "unsupported version"),
1870 + (NipcError::BadHeaderLen, "header_len != 32"),
1871 + (NipcError::BadKind, "unknown message kind"),
1872 + (NipcError::BadLayout, "unknown layout_version"),
1873 + (NipcError::OutOfBounds, "offset+length exceeds data"),
1874 + (NipcError::MissingNul, "string not NUL-terminated"),
1875 + (NipcError::BadAlignment, "item not 8-byte aligned"),
1876 + (NipcError::BadItemCount, "item count inconsistent"),
1877 + (NipcError::Overflow, "builder out of space"),
1878 + ];
1879 + for (err, expected) in cases {
1880 + let msg = format!("{}", err);
1881 + assert_eq!(msg, expected, "Display for {:?}", err);
1882 + }
1883 + // Also verify std::error::Error is implemented
1884 + let err: &dyn std::error::Error = &NipcError::Truncated;
1885 + let _ = format!("{err}");
1886 + }
1887 +
1888 + // -----------------------------------------------------------------------
1889 + // ChunkHeader decode: flags != 0 and chunk_payload_len == 0
1890 + // -----------------------------------------------------------------------
1891 +
1892 + #[test]
1893 + fn chunk_decode_bad_flags() {
1894 + // Line 257: flags != 0 -> BadLayout
1895 + let c = ChunkHeader {
1896 + magic: MAGIC_CHUNK,
1897 + version: VERSION,
1898 + flags: 0x01, // non-zero flags
1899 + message_id: 1,
1900 + total_message_len: 100,
1901 + chunk_index: 0,
1902 + chunk_count: 1,
1903 + chunk_payload_len: 50,
1904 + };
1905 + let mut buf = [0u8; 32];
1906 + c.encode(&mut buf);
1907 + assert_eq!(ChunkHeader::decode(&buf), Err(NipcError::BadLayout));
1908 + }
1909 +
1910 + #[test]
1911 + fn chunk_decode_zero_payload_len() {
1912 + // Line 260: chunk_payload_len == 0 -> BadLayout
1913 + let c = ChunkHeader {
1914 + magic: MAGIC_CHUNK,
1915 + version: VERSION,
1916 + flags: 0,
1917 + message_id: 1,
1918 + total_message_len: 100,
1919 + chunk_index: 0,
1920 + chunk_count: 1,
1921 + chunk_payload_len: 0,
1922 + };
1923 + let mut buf = [0u8; 32];
1924 + c.encode(&mut buf);
1925 + assert_eq!(ChunkHeader::decode(&buf), Err(NipcError::BadLayout));
1926 + }
1927 +
1928 + // -----------------------------------------------------------------------
1929 + // batch_dir_encode: buffer too small (line 282)
1930 + // -----------------------------------------------------------------------
1931 +
1932 + #[test]
1933 + fn batch_dir_encode_too_small() {
1934 + let entries = [
1935 + BatchEntry {
1936 + offset: 0,
1937 + length: 8,
1938 + },
1939 + BatchEntry {
1940 + offset: 8,
1941 + length: 8,
1942 + },
1943 + ];
1944 + let mut buf = [0u8; 12]; // needs 16, only 12
1945 + assert_eq!(batch_dir_encode(&entries, &mut buf), 0);
1946 + }
1947 +
1948 + // -----------------------------------------------------------------------
1949 + // batch_dir_validate error paths (lines 329, 336, 339)
1950 + // -----------------------------------------------------------------------
1951 +
1952 + #[test]
1953 + fn batch_dir_validate_truncated() {
1954 + let buf = [0u8; 4]; // too short for 1 entry (needs 8)
1955 + assert_eq!(batch_dir_validate(&buf, 1, 100), Err(NipcError::Truncated));
1956 + }
1957 +
1958 + #[test]
1959 + fn batch_dir_validate_bad_alignment() {
1960 + let mut buf = [0u8; 8];
1961 + buf[0..4].copy_from_slice(&3u32.to_ne_bytes()); // unaligned offset
1962 + buf[4..8].copy_from_slice(&8u32.to_ne_bytes());
1963 + assert_eq!(
1964 + batch_dir_validate(&buf, 1, 100),
1965 + Err(NipcError::BadAlignment)
1966 + );
1967 + }
1968 +
1969 + #[test]
1970 + fn batch_dir_validate_out_of_bounds() {
1971 + let mut buf = [0u8; 8];
1972 + buf[0..4].copy_from_slice(&0u32.to_ne_bytes());
1973 + buf[4..8].copy_from_slice(&200u32.to_ne_bytes()); // exceeds packed_area_len
1974 + assert_eq!(
1975 + batch_dir_validate(&buf, 1, 100),
1976 + Err(NipcError::OutOfBounds)
1977 + );
1978 + }
1979 +
1980 + #[test]
1981 + fn batch_dir_validate_ok() {
1982 + let mut buf = [0u8; 16];
1983 + buf[0..4].copy_from_slice(&0u32.to_ne_bytes());
1984 + buf[4..8].copy_from_slice(&8u32.to_ne_bytes());
1985 + buf[8..12].copy_from_slice(&8u32.to_ne_bytes());
1986 + buf[12..16].copy_from_slice(&8u32.to_ne_bytes());
1987 + assert!(batch_dir_validate(&buf, 2, 100).is_ok());
1988 + }
1989 +
1990 + // -----------------------------------------------------------------------
1991 + // batch_item_get: alignment check (line 375)
1992 + // -----------------------------------------------------------------------
1993 +
1994 + #[test]
1995 + fn batch_item_get_bad_alignment() {
1996 + // Manually craft a batch payload with unaligned offset
1997 + let mut buf = [0u8; 64];
1998 + // Directory: 1 entry at offset 0 of buf
1999 + buf[0..4].copy_from_slice(&3u32.to_ne_bytes()); // unaligned offset
2000 + buf[4..8].copy_from_slice(&4u32.to_ne_bytes());
2001 + assert_eq!(batch_item_get(&buf, 1, 0), Err(NipcError::BadAlignment));
2002 + }
2003 +
2004 + #[test]
2005 + fn batch_item_get_truncated_dir() {
2006 + // Payload too small to hold the directory
2007 + let buf = [0u8; 4]; // needs at least 8 for 1 item directory
2008 + assert_eq!(batch_item_get(&buf, 1, 0), Err(NipcError::Truncated));
2009 + }
2010 +
2011 + // -----------------------------------------------------------------------
2012 + // BatchBuilder::finish compaction (lines 451-456)
2013 + // -----------------------------------------------------------------------
2014 +
2015 + #[test]
2016 + fn batch_builder_compaction() {
2017 + // Reserve space for 8 items but add only 2 -- triggers compaction
2018 + let mut buf = [0u8; 1024];
2019 + let mut b = BatchBuilder::new(&mut buf, 8);
2020 +
2021 + let item1 = [1u8, 2, 3, 4, 5, 6, 7, 8];
2022 + let item2 = [10u8, 20, 30, 40];
2023 +
2024 + b.add(&item1).unwrap();
2025 + b.add(&item2).unwrap();
2026 +
2027 + // dir_end for 8 items = align8(8*8) = 64
2028 + // final_dir_aligned for 2 items = align8(2*8) = 16
2029 + // This triggers the copy_within compaction branch (line 453-456)
2030 + let (total, count) = b.finish();
2031 + assert_eq!(count, 2);
2032 + assert!(total > 0);
2033 +
2034 + // Verify the items can still be extracted correctly
2035 + let (data, len) = batch_item_get(&buf[..total], 2, 0).unwrap();
2036 + assert_eq!(len as usize, item1.len());
2037 + assert_eq!(data, &item1);
2038 +
2039 + let (data, len) = batch_item_get(&buf[..total], 2, 1).unwrap();
2040 + assert_eq!(len as usize, item2.len());
2041 + assert_eq!(data, &item2);
2042 + }
2043 +
2044 + // -----------------------------------------------------------------------
2045 + // HelloAck decode: flags != 0 (line 606)
2046 + // -----------------------------------------------------------------------
2047 +
2048 + #[test]
2049 + fn hello_ack_decode_bad_flags() {
2050 + let h = HelloAck {
2051 + layout_version: 1,
2052 + flags: 1, // non-zero flags
2053 + ..Default::default()
2054 + };
2055 + let mut buf = [0u8; 48];
2056 + h.encode(&mut buf);
2057 + assert_eq!(HelloAck::decode(&buf), Err(NipcError::BadLayout));
2058 + }
2059 +
2060 + // -----------------------------------------------------------------------
2061 + // Hello decode: non-zero padding (line 527)
2062 + // -----------------------------------------------------------------------
2063 +
2064 + #[test]
2065 + fn hello_decode_bad_padding() {
2066 + let h = Hello {
2067 + layout_version: 1,
2068 + flags: 0,
2069 + ..Default::default()
2070 + };
2071 + let mut buf = [0u8; 44];
2072 + h.encode(&mut buf);
2073 + // Corrupt the padding bytes at 28..32
2074 + buf[28..32].copy_from_slice(&1u32.to_ne_bytes());
2075 + assert_eq!(Hello::decode(&buf), Err(NipcError::BadLayout));
2076 + }
2077 +
2078 + // -----------------------------------------------------------------------
2079 + // Cgroups request: non-zero flags (line 45)
2080 + // -----------------------------------------------------------------------
2081 +
2082 + #[test]
2083 + fn cgroups_req_decode_bad_flags() {
2084 + let r = CgroupsRequest {
2085 + layout_version: 1,
2086 + flags: 1, // non-zero flags -> BadLayout
2087 + };
2088 + let mut buf = [0u8; 4];
2089 + r.encode(&mut buf);
2090 + assert_eq!(CgroupsRequest::decode(&buf), Err(NipcError::BadLayout));
2091 + }
2092 +
2093 + // -----------------------------------------------------------------------
2094 + // CgroupsResponseView: non-zero flags and reserved (lines 125, 130)
2095 + // -----------------------------------------------------------------------
2096 +
2097 + #[test]
2098 + fn cgroups_resp_decode_bad_flags() {
2099 + // Line 125: flags != 0 -> BadLayout
2100 + let mut buf = [0u8; 24];
2101 + buf[0..2].copy_from_slice(&1u16.to_ne_bytes()); // layout_version = 1
2102 + buf[2..4].copy_from_slice(&1u16.to_ne_bytes()); // flags = 1 (non-zero)
2103 + assert_eq!(
2104 + CgroupsResponseView::decode(&buf).unwrap_err(),
2105 + NipcError::BadLayout
2106 + );
2107 + }
2108 +
2109 + #[test]
2110 + fn cgroups_resp_decode_bad_reserved() {
2111 + // Line 130: reserved != 0 -> BadLayout
2112 + let mut buf = [0u8; 24];
2113 + buf[0..2].copy_from_slice(&1u16.to_ne_bytes()); // layout_version = 1
2114 + buf[2..4].copy_from_slice(&0u16.to_ne_bytes()); // flags = 0
2115 + buf[12..16].copy_from_slice(&1u32.to_ne_bytes()); // reserved = 1
2116 + assert_eq!(
2117 + CgroupsResponseView::decode(&buf).unwrap_err(),
2118 + NipcError::BadLayout
2119 + );
2120 + }
2121 +
2122 + // -----------------------------------------------------------------------
2123 + // CgroupsResponseView: bad alignment in directory (line 149)
2124 + // -----------------------------------------------------------------------
2125 +
2126 + #[test]
2127 + fn cgroups_resp_decode_bad_dir_alignment() {
2128 + // dir entry with offset not aligned to 8
2129 + let mut buf = [0u8; 128];
2130 + buf[0..2].copy_from_slice(&1u16.to_ne_bytes()); // layout_version
2131 + buf[4..8].copy_from_slice(&1u32.to_ne_bytes()); // item_count = 1
2132 + // Dir entry at offset 24: offset=3 (unaligned), length=32
2133 + buf[24..28].copy_from_slice(&3u32.to_ne_bytes());
2134 + buf[28..32].copy_from_slice(&32u32.to_ne_bytes());
2135 + assert_eq!(
2136 + CgroupsResponseView::decode(&buf).unwrap_err(),
2137 + NipcError::BadAlignment
2138 + );
2139 + }
2140 +
2141 + // -----------------------------------------------------------------------
2142 + // CgroupsItemView: bad layout_version, bad flags (lines 200-206)
2143 + // -----------------------------------------------------------------------
2144 +
2145 + #[test]
2146 + fn cgroups_item_bad_layout_version() {
2147 + // Build valid snapshot then corrupt item layout_version
2148 + let mut buf = [0u8; 4096];
2149 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 1);
2150 + b.add(1, 0, 1, b"test", b"/test").unwrap();
2151 + let total = b.finish();
2152 +
2153 + // Find item start
2154 + let dir_end = CGROUPS_RESP_HDR_SIZE + 1 * CGROUPS_DIR_ENTRY_SIZE;
2155 + let item_off = u32::from_ne_bytes(
2156 + buf[CGROUPS_RESP_HDR_SIZE..CGROUPS_RESP_HDR_SIZE + 4]
2157 + .try_into()
2158 + .unwrap(),
2159 + ) as usize;
2160 + let item_start = dir_end + item_off;
2161 +
2162 + // Corrupt layout_version to 99
2163 + buf[item_start..item_start + 2].copy_from_slice(&99u16.to_ne_bytes());
2164 +
2165 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
2166 + assert_eq!(view.item(0).unwrap_err(), NipcError::BadLayout);
2167 + }
2168 +
2169 + #[test]
2170 + fn cgroups_item_bad_flags() {
2171 + // Build valid snapshot then corrupt item flags
2172 + let mut buf = [0u8; 4096];
2173 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 1);
2174 + b.add(1, 0, 1, b"test", b"/test").unwrap();
2175 + let total = b.finish();
2176 +
2177 + let dir_end = CGROUPS_RESP_HDR_SIZE + 1 * CGROUPS_DIR_ENTRY_SIZE;
2178 + let item_off = u32::from_ne_bytes(
2179 + buf[CGROUPS_RESP_HDR_SIZE..CGROUPS_RESP_HDR_SIZE + 4]
2180 + .try_into()
2181 + .unwrap(),
2182 + ) as usize;
2183 + let item_start = dir_end + item_off;
2184 +
2185 + // Corrupt flags to non-zero
2186 + buf[item_start + 2..item_start + 4].copy_from_slice(&1u16.to_ne_bytes());
2187 +
2188 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
2189 + assert_eq!(view.item(0).unwrap_err(), NipcError::BadLayout);
2190 + }
2191 +
2192 + // -----------------------------------------------------------------------
2193 + // CgroupsItemView: name_off < ITEM_HDR_SIZE (line 211)
2194 + // -----------------------------------------------------------------------
2195 +
2196 + #[test]
2197 + fn cgroups_item_name_off_too_small() {
2198 + let mut buf = [0u8; 4096];
2199 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 1);
2200 + b.add(1, 0, 1, b"test", b"/test").unwrap();
2201 + let total = b.finish();
2202 +
2203 + let dir_end = CGROUPS_RESP_HDR_SIZE + 1 * CGROUPS_DIR_ENTRY_SIZE;
2204 + let item_off = u32::from_ne_bytes(
2205 + buf[CGROUPS_RESP_HDR_SIZE..CGROUPS_RESP_HDR_SIZE + 4]
2206 + .try_into()
2207 + .unwrap(),
2208 + ) as usize;
2209 + let item_start = dir_end + item_off;
2210 +
2211 + // Set name_offset to 0 (< 32 = CGROUPS_ITEM_HDR_SIZE)
2212 + buf[item_start + 16..item_start + 20].copy_from_slice(&0u32.to_ne_bytes());
2213 +
2214 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
2215 + assert_eq!(view.item(0).unwrap_err(), NipcError::OutOfBounds);
2216 + }
2217 +
2218 + // -----------------------------------------------------------------------
2219 + // CgroupsItemView: path NUL missing (line 228)
2220 + // -----------------------------------------------------------------------
2221 +
2222 + #[test]
2223 + fn cgroups_item_path_missing_nul() {
2224 + let mut buf = [0u8; 4096];
2225 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 1);
2226 + b.add(1, 0, 1, b"test", b"/test").unwrap();
2227 + let total = b.finish();
2228 +
2229 + let dir_end = CGROUPS_RESP_HDR_SIZE + 1 * CGROUPS_DIR_ENTRY_SIZE;
2230 + let item_off = u32::from_ne_bytes(
2231 + buf[CGROUPS_RESP_HDR_SIZE..CGROUPS_RESP_HDR_SIZE + 4]
2232 + .try_into()
2233 + .unwrap(),
2234 + ) as usize;
2235 + let item_start = dir_end + item_off;
2236 +
2237 + // Find the path NUL terminator and corrupt it
2238 + let path_off =
2239 + u32::from_ne_bytes(buf[item_start + 24..item_start + 28].try_into().unwrap()) as usize;
2240 + let path_len =
2241 + u32::from_ne_bytes(buf[item_start + 28..item_start + 32].try_into().unwrap()) as usize;
2242 + buf[item_start + path_off + path_len] = b'X'; // corrupt path NUL
2243 +
2244 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
2245 + assert_eq!(view.item(0).unwrap_err(), NipcError::MissingNul);
2246 + }
2247 +
2248 + // -----------------------------------------------------------------------
2249 + // CgroupsItemView: path_off < ITEM_HDR_SIZE (line 222)
2250 + // -----------------------------------------------------------------------
2251 +
2252 + #[test]
2253 + fn cgroups_item_path_off_too_small() {
2254 + let mut buf = [0u8; 4096];
2255 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 1);
2256 + b.add(1, 0, 1, b"test", b"/test").unwrap();
2257 + let total = b.finish();
2258 +
2259 + let dir_end = CGROUPS_RESP_HDR_SIZE + 1 * CGROUPS_DIR_ENTRY_SIZE;
2260 + let item_off = u32::from_ne_bytes(
2261 + buf[CGROUPS_RESP_HDR_SIZE..CGROUPS_RESP_HDR_SIZE + 4]
2262 + .try_into()
2263 + .unwrap(),
2264 + ) as usize;
2265 + let item_start = dir_end + item_off;
2266 +
2267 + // Set path_offset to 0 (< 32 = CGROUPS_ITEM_HDR_SIZE)
2268 + buf[item_start + 24..item_start + 28].copy_from_slice(&0u32.to_ne_bytes());
2269 +
2270 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
2271 + assert_eq!(view.item(0).unwrap_err(), NipcError::OutOfBounds);
2272 + }
2273 +
2274 + // -----------------------------------------------------------------------
2275 + // CgroupsItemView: path string OOB (line 225)
2276 + // -----------------------------------------------------------------------
2277 +
2278 + #[test]
2279 + fn cgroups_item_path_string_oob() {
2280 + let mut buf = [0u8; 4096];
2281 + let mut b = CgroupsBuilder::new(&mut buf, 1, 0, 1);
2282 + b.add(1, 0, 1, b"test", b"/test").unwrap();
2283 + let total = b.finish();
2284 +
2285 + let dir_end = CGROUPS_RESP_HDR_SIZE + 1 * CGROUPS_DIR_ENTRY_SIZE;
2286 + let item_off = u32::from_ne_bytes(
2287 + buf[CGROUPS_RESP_HDR_SIZE..CGROUPS_RESP_HDR_SIZE + 4]
2288 + .try_into()
2289 + .unwrap(),
2290 + ) as usize;
2291 + let item_start = dir_end + item_off;
2292 +
2293 + // Corrupt path_length to huge value
2294 + buf[item_start + 28..item_start + 32].copy_from_slice(&99999u32.to_ne_bytes());
2295 +
2296 + let view = CgroupsResponseView::decode(&buf[..total]).unwrap();
2297 + assert_eq!(view.item(0).unwrap_err(), NipcError::OutOfBounds);
2298 + }
2299 +
2300 + // -----------------------------------------------------------------------
2301 + // Cgroups dispatch (lines 438-447)
2302 + // -----------------------------------------------------------------------
2303 +
2304 + #[test]
2305 + fn dispatch_cgroups_snapshot_bad_request() {
2306 + // Bad request (too short) -> dispatch returns None (line 438)
2307 + let mut resp = [0u8; 4096];
2308 + let result =
2309 + crate::protocol::dispatch_cgroups_snapshot(&[], &mut resp, 1, |_req, _builder| true);
2310 + assert!(result.is_none());
2311 + }
2312 +
2313 + #[test]
2314 + fn dispatch_cgroups_snapshot_handler_returns_false() {
2315 + // Handler returns false -> dispatch returns None (lines 440-441)
2316 + let req = CgroupsRequest {
2317 + layout_version: 1,
2318 + flags: 0,
2319 + };
2320 + let mut req_buf = [0u8; 4];
2321 + req.encode(&mut req_buf);
2322 + let mut resp = [0u8; 4096];
2323 + let result =
2324 + crate::protocol::dispatch_cgroups_snapshot(&req_buf, &mut resp, 1, |_req, _builder| {
2325 + false
2326 + });
2327 + assert!(result.is_none());
2328 + }
2329 +
2330 + #[test]
2331 + fn dispatch_cgroups_snapshot_success() {
2332 + let req = CgroupsRequest {
2333 + layout_version: 1,
2334 + flags: 0,
2335 + };
2336 + let mut req_buf = [0u8; 4];
2337 + req.encode(&mut req_buf);
2338 + let mut resp = [0u8; 4096];
2339 + let result =
2340 + crate::protocol::dispatch_cgroups_snapshot(&req_buf, &mut resp, 2, |_req, builder| {
2341 + builder.add(1, 0, 1, b"cg1", b"/test").unwrap();
2342 + true
2343 + });
2344 + assert!(result.is_some());
2345 + let n = result.unwrap();
2346 + let view = CgroupsResponseView::decode(&resp[..n]).unwrap();
2347 + assert_eq!(view.item_count, 1);
2348 + }
2349 +
2350 + // -----------------------------------------------------------------------
2351 + // Dispatch increment / string_reverse buf overflow (lines 27, 58)
2352 + // -----------------------------------------------------------------------
2353 +
2354 + #[test]
2355 + fn dispatch_increment_resp_too_small() {
2356 + // Response buffer too small -> encode returns 0 -> dispatch returns None
2357 + let req_val = 42u64;
2358 + let mut req_buf = [0u8; 8];
2359 + crate::protocol::increment_encode(req_val, &mut req_buf);
2360 + let mut resp = [0u8; 4]; // too small for 8-byte response
2361 + let result = crate::protocol::dispatch_increment(&req_buf, &mut resp, |v| Some(v + 1));
2362 + assert!(result.is_none());
2363 + }
2364 +
2365 + #[test]
2366 + fn dispatch_increment_handler_none() {
2367 + let mut req_buf = [0u8; 8];
2368 + crate::protocol::increment_encode(42, &mut req_buf);
2369 + let mut resp = [0u8; 8];
2370 + let result = crate::protocol::dispatch_increment(&req_buf, &mut resp, |_| None);
2371 + assert!(result.is_none());
2372 + }
2373 +
2374 + #[test]
2375 + fn dispatch_string_reverse_resp_too_small() {
2376 + let s = b"hello";
2377 + let mut req_buf = [0u8; 64];
2378 + crate::protocol::string_reverse_encode(s, &mut req_buf);
2379 + let mut resp = [0u8; 4]; // too small
2380 + let result = crate::protocol::dispatch_string_reverse(&req_buf, &mut resp, |data| {
2381 + Some(data.iter().rev().copied().collect())
2382 + });
2383 + assert!(result.is_none());
2384 + }
2385 +
2386 + #[test]
2387 + fn dispatch_string_reverse_handler_none() {
2388 + let s = b"hello";
2389 + let mut req_buf = [0u8; 64];
2390 + crate::protocol::string_reverse_encode(s, &mut req_buf);
2391 + let mut resp = [0u8; 64];
2392 + let result = crate::protocol::dispatch_string_reverse(&req_buf, &mut resp, |_| None);
2393 + assert!(result.is_none());
2394 + }
2395 +}
src/crates/netipc/src/protocol/string_reverse.rs new
+195
@@ -0,0 +1,195 @@
1 +//! STRING_REVERSE codec (method 3) -- variable-length payload:
2 +//! [0:4] u32 str_offset (from payload start, always 8)
3 +//! [4:8] u32 str_length (excluding NUL)
4 +//! [8:N+1] string data + NUL
5 +
6 +use super::NipcError;
7 +
8 +pub const STRING_REVERSE_HDR_SIZE: usize = 8;
9 +
10 +/// Ephemeral view into a decoded STRING_REVERSE payload.
11 +#[derive(Debug, Clone)]
12 +pub struct StringReverseView<'a> {
13 + pub str_data: &'a [u8], // slice into payload, excludes NUL
14 + pub str_len: u32,
15 +}
16 +
17 +impl<'a> StringReverseView<'a> {
18 + pub fn as_str(&self) -> &'a str {
19 + std::str::from_utf8(self.str_data).unwrap_or("")
20 + }
21 +}
22 +
23 +pub fn string_reverse_encode(s: &[u8], buf: &mut [u8]) -> usize {
24 + if s.len() > u32::MAX as usize {
25 + return 0;
26 + }
27 + let total = STRING_REVERSE_HDR_SIZE + s.len() + 1;
28 + if buf.len() < total {
29 + return 0;
30 + }
31 + let offset: u32 = STRING_REVERSE_HDR_SIZE as u32;
32 + let length: u32 = s.len() as u32;
33 + buf[0..4].copy_from_slice(&offset.to_ne_bytes());
34 + buf[4..8].copy_from_slice(&length.to_ne_bytes());
35 + if !s.is_empty() {
36 + buf[8..8 + s.len()].copy_from_slice(s);
37 + }
38 + buf[8 + s.len()] = 0; // NUL
39 + total
40 +}
41 +
42 +pub fn string_reverse_decode(buf: &[u8]) -> Result<StringReverseView<'_>, NipcError> {
43 + if buf.len() < STRING_REVERSE_HDR_SIZE {
44 + return Err(NipcError::Truncated);
45 + }
46 + let str_offset = u32::from_ne_bytes(buf[0..4].try_into().unwrap()) as usize;
47 + if str_offset < STRING_REVERSE_HDR_SIZE {
48 + return Err(NipcError::BadLayout);
49 + }
50 + let str_length = u32::from_ne_bytes(buf[4..8].try_into().unwrap()) as usize;
51 + let end = str_offset
52 + .checked_add(str_length)
53 + .and_then(|v| v.checked_add(1))
54 + .ok_or(NipcError::OutOfBounds)?;
55 + if end > buf.len() {
56 + return Err(NipcError::OutOfBounds);
57 + }
58 + if buf[str_offset + str_length] != 0 {
59 + return Err(NipcError::MissingNul);
60 + }
61 + Ok(StringReverseView {
62 + str_data: &buf[str_offset..str_offset + str_length],
63 + str_len: str_length as u32,
64 + })
65 +}
66 +
67 +/// STRING_REVERSE dispatch: decode -> handler -> encode.
68 +pub fn dispatch_string_reverse<F>(req: &[u8], resp: &mut [u8], handler: F) -> Option<usize>
69 +where
70 + F: FnOnce(&[u8]) -> Option<Vec<u8>>,
71 +{
72 + let view = string_reverse_decode(req).ok()?;
73 + let result = handler(view.str_data)?;
74 + let n = string_reverse_encode(&result, resp);
75 + if n == 0 {
76 + return None;
77 + }
78 + Some(n)
79 +}
80 +
81 +#[cfg(test)]
82 +mod tests {
83 + use super::*;
84 +
85 + #[test]
86 + fn encode_decode_roundtrip() {
87 + let s = b"hello world";
88 + let mut buf = [0u8; 64];
89 + let n = string_reverse_encode(s, &mut buf);
90 + assert_eq!(n, STRING_REVERSE_HDR_SIZE + s.len() + 1);
91 +
92 + let view = string_reverse_decode(&buf[..n]).unwrap();
93 + assert_eq!(view.str_data, s);
94 + assert_eq!(view.str_len, s.len() as u32);
95 + assert_eq!(view.as_str(), "hello world");
96 + }
97 +
98 + #[test]
99 + fn encode_empty() {
100 + let mut buf = [0u8; 64];
101 + let n = string_reverse_encode(b"", &mut buf);
102 + assert_eq!(n, STRING_REVERSE_HDR_SIZE + 1);
103 + let view = string_reverse_decode(&buf[..n]).unwrap();
104 + assert_eq!(view.str_data, b"");
105 + assert_eq!(view.str_len, 0);
106 + }
107 +
108 + #[test]
109 + fn encode_too_small() {
110 + let mut buf = [0u8; 4];
111 + assert_eq!(string_reverse_encode(b"hello", &mut buf), 0);
112 + }
113 +
114 + #[test]
115 + fn decode_truncated() {
116 + assert!(matches!(
117 + string_reverse_decode(&[0u8; 4]),
118 + Err(NipcError::Truncated)
119 + ));
120 + }
121 +
122 + #[test]
123 + fn decode_oob() {
124 + let mut buf = [0u8; 16];
125 + buf[0..4].copy_from_slice(&8u32.to_ne_bytes());
126 + buf[4..8].copy_from_slice(&99u32.to_ne_bytes());
127 + assert!(matches!(
128 + string_reverse_decode(&buf),
129 + Err(NipcError::OutOfBounds)
130 + ));
131 + }
132 +
133 + #[test]
134 + fn decode_missing_nul() {
135 + let mut buf = [0u8; 16];
136 + buf[0..4].copy_from_slice(&8u32.to_ne_bytes());
137 + buf[4..8].copy_from_slice(&3u32.to_ne_bytes());
138 + buf[8] = b'a';
139 + buf[9] = b'b';
140 + buf[10] = b'c';
141 + buf[11] = b'X';
142 + assert!(matches!(
143 + string_reverse_decode(&buf),
144 + Err(NipcError::MissingNul)
145 + ));
146 + }
147 +
148 + #[test]
149 + fn dispatch_resp_too_small() {
150 + let s = b"hello";
151 + let mut req_buf = [0u8; 64];
152 + string_reverse_encode(s, &mut req_buf);
153 + let mut resp = [0u8; 4];
154 + let result = dispatch_string_reverse(&req_buf, &mut resp, |data| {
155 + Some(data.iter().rev().copied().collect())
156 + });
157 + assert!(result.is_none());
158 + }
159 +
160 + #[test]
161 + fn dispatch_handler_returns_none() {
162 + let mut req_buf = [0u8; 64];
163 + string_reverse_encode(b"test", &mut req_buf);
164 + let mut resp = [0u8; 64];
165 + assert!(dispatch_string_reverse(&req_buf, &mut resp, |_| None).is_none());
166 + }
167 +
168 + #[test]
169 + fn dispatch_bad_request() {
170 + let mut resp = [0u8; 64];
171 + assert!(dispatch_string_reverse(&[0u8; 4], &mut resp, |_| None).is_none());
172 + }
173 +
174 + #[test]
175 + fn dispatch_success() {
176 + let mut req_buf = [0u8; 64];
177 + let n = string_reverse_encode(b"abc", &mut req_buf);
178 + let mut resp = [0u8; 64];
179 + let rn = dispatch_string_reverse(&req_buf[..n], &mut resp, |data| {
180 + Some(data.iter().rev().copied().collect())
181 + })
182 + .unwrap();
183 + let view = string_reverse_decode(&resp[..rn]).unwrap();
184 + assert_eq!(view.str_data, b"cba");
185 + }
186 +
187 + #[test]
188 + fn as_str_non_utf8() {
189 + let view = StringReverseView {
190 + str_data: &[0xFF, 0xFE],
191 + str_len: 2,
192 + };
193 + assert_eq!(view.as_str(), "");
194 + }
195 +}
src/crates/netipc/src/service/cgroups.rs new
+264
@@ -0,0 +1,264 @@
1 +//! L2 cgroups-snapshot service facade.
2 +//!
3 +//! The public service surface is service-kind specific: one endpoint, one
4 +//! request kind. The request code remains in the outer envelope only for
5 +//! validation, not for public multi-method dispatch.
6 +
7 +use super::raw;
8 +use crate::protocol::{CgroupsResponseView, NipcError, METHOD_CGROUPS_SNAPSHOT, PROFILE_BASELINE};
9 +
10 +#[cfg(unix)]
11 +use crate::transport::posix::{
12 + ClientConfig as TransportClientConfig, ServerConfig as TransportServerConfig,
13 +};
14 +
15 +#[cfg(windows)]
16 +use crate::transport::windows::{
17 + ClientConfig as TransportClientConfig, ServerConfig as TransportServerConfig,
18 +};
19 +
20 +use std::sync::atomic::AtomicBool;
21 +use std::sync::Arc;
22 +
23 +pub use raw::{CgroupsCacheItem, CgroupsCacheStatus, ClientState, ClientStatus, SnapshotHandler};
24 +
25 +/// Public L2/L3 client configuration for the cgroups-snapshot service.
26 +///
27 +/// This service-level configuration is shared across supported operating
28 +/// systems. Transport-only tuning stays below the public typed API.
29 +#[derive(Debug, Clone)]
30 +pub struct ClientConfig {
31 + pub supported_profiles: u32,
32 + pub preferred_profiles: u32,
33 + pub max_request_batch_items: u32,
34 + pub max_response_payload_bytes: u32,
35 + pub auth_token: u64,
36 +}
37 +
38 +impl Default for ClientConfig {
39 + fn default() -> Self {
40 + Self {
41 + supported_profiles: PROFILE_BASELINE,
42 + preferred_profiles: 0,
43 + max_request_batch_items: 0,
44 + max_response_payload_bytes: 0,
45 + auth_token: 0,
46 + }
47 + }
48 +}
49 +
50 +impl ClientConfig {
51 + fn into_transport(self) -> TransportClientConfig {
52 + let mut transport = TransportClientConfig::default();
53 + transport.supported_profiles = self.supported_profiles;
54 + transport.preferred_profiles = self.preferred_profiles;
55 + transport.max_request_batch_items = self.max_request_batch_items;
56 + transport.max_response_payload_bytes = self.max_response_payload_bytes;
57 + transport.max_response_batch_items = self.max_request_batch_items;
58 + transport.auth_token = self.auth_token;
59 + transport
60 + }
61 +}
62 +
63 +/// Public typed-server configuration for the cgroups-snapshot service.
64 +///
65 +/// This configuration is intentionally transport-agnostic. Transport-only
66 +/// knobs such as socket backlog or packet sizing stay below the service layer.
67 +#[derive(Debug, Clone)]
68 +pub struct ServerConfig {
69 + pub supported_profiles: u32,
70 + pub preferred_profiles: u32,
71 + pub max_request_batch_items: u32,
72 + pub max_response_payload_bytes: u32,
73 + pub auth_token: u64,
74 +}
75 +
76 +impl Default for ServerConfig {
77 + fn default() -> Self {
78 + Self {
79 + supported_profiles: PROFILE_BASELINE,
80 + preferred_profiles: 0,
81 + max_request_batch_items: 0,
82 + max_response_payload_bytes: 0,
83 + auth_token: 0,
84 + }
85 + }
86 +}
87 +
88 +impl ServerConfig {
89 + fn into_transport(self) -> TransportServerConfig {
90 + let mut transport = TransportServerConfig::default();
91 + transport.supported_profiles = self.supported_profiles;
92 + transport.preferred_profiles = self.preferred_profiles;
93 + transport.max_request_batch_items = self.max_request_batch_items;
94 + transport.max_response_payload_bytes = self.max_response_payload_bytes;
95 + transport.max_response_batch_items = self.max_request_batch_items;
96 + transport.auth_token = self.auth_token;
97 + transport
98 + }
99 +}
100 +
101 +/// L2 client context for the cgroups-snapshot service.
102 +pub struct CgroupsClient {
103 + inner: raw::RawClient,
104 +}
105 +
106 +impl CgroupsClient {
107 + /// Create a new client context. Does NOT connect. Does NOT require the
108 + /// server to be running.
109 + pub fn new(run_dir: &str, service_name: &str, config: ClientConfig) -> Self {
110 + Self {
111 + inner: raw::RawClient::new_snapshot(run_dir, service_name, config.into_transport()),
112 + }
113 + }
114 +
115 + /// Attempt connect if DISCONNECTED/NOT_FOUND, reconnect if BROKEN.
116 + /// Returns true if the state changed.
117 + pub fn refresh(&mut self) -> bool {
118 + self.inner.refresh()
119 + }
120 +
121 + /// Cheap cached boolean. No I/O, no syscalls.
122 + #[inline]
123 + pub fn ready(&self) -> bool {
124 + self.inner.ready()
125 + }
126 +
127 + /// Detailed status snapshot for diagnostics.
128 + pub fn status(&self) -> ClientStatus {
129 + self.inner.status()
130 + }
131 +
132 + /// Blocking typed call for the cgroups-snapshot service.
133 + pub fn call_snapshot(&mut self) -> Result<CgroupsResponseView<'_>, NipcError> {
134 + self.inner.call_snapshot()
135 + }
136 +
137 + /// Tear down connection and release resources.
138 + pub fn close(&mut self) {
139 + self.inner.close();
140 + }
141 +}
142 +
143 +impl Drop for CgroupsClient {
144 + fn drop(&mut self) {
145 + self.close();
146 + }
147 +}
148 +
149 +/// Typed server handler surface for the cgroups-snapshot service.
150 +#[derive(Clone, Default)]
151 +pub struct Handler {
152 + pub handle: Option<SnapshotHandler>,
153 + pub snapshot_max_items: u32,
154 +}
155 +
156 +/// Managed server for the cgroups-snapshot service kind.
157 +pub struct ManagedServer {
158 + inner: raw::ManagedServer,
159 +}
160 +
161 +impl ManagedServer {
162 + /// Create a new managed server. Does NOT start listening yet.
163 + pub fn new(run_dir: &str, service_name: &str, config: ServerConfig, handler: Handler) -> Self {
164 + Self::with_workers(run_dir, service_name, config, handler, 8)
165 + }
166 +
167 + /// Create a server with an explicit worker count limit.
168 + pub fn with_workers(
169 + run_dir: &str,
170 + service_name: &str,
171 + config: ServerConfig,
172 + handler: Handler,
173 + worker_count: usize,
174 + ) -> Self {
175 + let raw_handler = handler
176 + .handle
177 + .map(|handle| raw::snapshot_dispatch(handle, handler.snapshot_max_items));
178 +
179 + Self {
180 + inner: raw::ManagedServer::with_workers(
181 + run_dir,
182 + service_name,
183 + config.into_transport(),
184 + METHOD_CGROUPS_SNAPSHOT,
185 + raw_handler,
186 + worker_count,
187 + ),
188 + }
189 + }
190 +
191 + /// Run the acceptor loop. Blocking. Returns when `stop()` is called or on
192 + /// fatal error.
193 + pub fn run(&mut self) -> Result<(), NipcError> {
194 + self.inner.run()
195 + }
196 +
197 + /// Signal shutdown.
198 + pub fn stop(&self) {
199 + self.inner.stop();
200 + }
201 +
202 + /// Clone of the internal running flag for diagnostics and test helpers.
203 + ///
204 + /// For reliable shutdown, call `stop()`. On Windows, changing this flag
205 + /// alone does not wake a blocking listener accept.
206 + pub fn running_flag(&self) -> Arc<AtomicBool> {
207 + self.inner.running_flag()
208 + }
209 +}
210 +
211 +/// L3 client-side cgroups snapshot cache.
212 +pub struct CgroupsCache {
213 + inner: raw::CgroupsCache,
214 +}
215 +
216 +impl CgroupsCache {
217 + /// Create a new L3 cache. Creates the underlying L2 client context.
218 + /// Does NOT connect. Does NOT require the server to be running.
219 + pub fn new(run_dir: &str, service_name: &str, config: ClientConfig) -> Self {
220 + Self {
221 + inner: raw::CgroupsCache::new(run_dir, service_name, config.into_transport()),
222 + }
223 + }
224 +
225 + /// Refresh the cache. Returns true if the cache was updated.
226 + pub fn refresh(&mut self) -> bool {
227 + self.inner.refresh()
228 + }
229 +
230 + /// Returns true if at least one successful refresh has occurred.
231 + #[inline]
232 + pub fn ready(&self) -> bool {
233 + self.inner.ready()
234 + }
235 +
236 + /// Look up a cached item by hash + name. O(1), no I/O.
237 + pub fn lookup(&self, hash: u32, name: &str) -> Option<&CgroupsCacheItem> {
238 + self.inner.lookup(hash, name)
239 + }
240 +
241 + /// Fill a status snapshot for diagnostics.
242 + pub fn status(&self) -> CgroupsCacheStatus {
243 + self.inner.status()
244 + }
245 +
246 + /// Close the cache and underlying L2 client.
247 + pub fn close(&mut self) {
248 + self.inner.close();
249 + }
250 +}
251 +
252 +impl Drop for CgroupsCache {
253 + fn drop(&mut self) {
254 + self.close();
255 + }
256 +}
257 +
258 +#[cfg(all(test, unix))]
259 +#[path = "cgroups_unix_tests.rs"]
260 +mod tests;
261 +
262 +#[cfg(all(test, windows))]
263 +#[path = "cgroups_windows_tests.rs"]
264 +mod windows_tests;
src/crates/netipc/src/service/cgroups_unix_tests.rs new
+266
@@ -0,0 +1,266 @@
1 +use super::*;
2 +#[cfg(target_os = "linux")]
3 +use crate::protocol::PROFILE_SHM_FUTEX;
4 +use crate::protocol::{CgroupsBuilder, NipcError, PROFILE_BASELINE};
5 +use std::path::PathBuf;
6 +use std::sync::atomic::{AtomicU64, Ordering};
7 +use std::sync::Arc;
8 +use std::thread;
9 +use std::time::{Duration, SystemTime, UNIX_EPOCH};
10 +
11 +const TEST_RUN_DIR: &str = "/tmp/nipc_cgroups_rust_test";
12 +const AUTH_TOKEN: u64 = 0xDEADBEEFCAFEBABE;
13 +const RESPONSE_BUF_SIZE: usize = 65536;
14 +static SERVICE_COUNTER: AtomicU64 = AtomicU64::new(0);
15 +
16 +fn ensure_run_dir() {
17 + let _ = std::fs::create_dir_all(TEST_RUN_DIR);
18 +}
19 +
20 +fn unique_service(prefix: &str) -> String {
21 + let stamp = SystemTime::now()
22 + .duration_since(UNIX_EPOCH)
23 + .unwrap_or_default()
24 + .as_nanos();
25 + format!(
26 + "{}_{}_{}_{}",
27 + prefix,
28 + std::process::id(),
29 + SERVICE_COUNTER.fetch_add(1, Ordering::Relaxed) + 1,
30 + stamp
31 + )
32 +}
33 +
34 +fn cleanup_all(service: &str) {
35 + let _ = std::fs::remove_file(format!("{TEST_RUN_DIR}/{service}.sock"));
36 +}
37 +
38 +fn socket_path(service: &str) -> PathBuf {
39 + PathBuf::from(format!("{TEST_RUN_DIR}/{service}.sock"))
40 +}
41 +
42 +fn wait_for_listener_bind(service: &str) {
43 + let sock = socket_path(service);
44 + for _ in 0..2000 {
45 + if sock.exists() {
46 + return;
47 + }
48 + thread::sleep(Duration::from_micros(500));
49 + }
50 +
51 + panic!("listener did not bind for service {service}");
52 +}
53 +
54 +fn server_config() -> ServerConfig {
55 + ServerConfig {
56 + supported_profiles: PROFILE_BASELINE,
57 + preferred_profiles: PROFILE_BASELINE,
58 + max_request_batch_items: 1,
59 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
60 + auth_token: AUTH_TOKEN,
61 + ..ServerConfig::default()
62 + }
63 +}
64 +
65 +fn client_config() -> ClientConfig {
66 + ClientConfig {
67 + supported_profiles: PROFILE_BASELINE,
68 + preferred_profiles: PROFILE_BASELINE,
69 + max_request_batch_items: 1,
70 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
71 + auth_token: AUTH_TOKEN,
72 + ..ClientConfig::default()
73 + }
74 +}
75 +
76 +#[cfg(target_os = "linux")]
77 +fn shm_server_config() -> ServerConfig {
78 + ServerConfig {
79 + supported_profiles: PROFILE_BASELINE | PROFILE_SHM_FUTEX,
80 + preferred_profiles: PROFILE_SHM_FUTEX,
81 + max_request_batch_items: 1,
82 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
83 + auth_token: AUTH_TOKEN,
84 + ..ServerConfig::default()
85 + }
86 +}
87 +
88 +#[cfg(target_os = "linux")]
89 +fn shm_client_config() -> ClientConfig {
90 + ClientConfig {
91 + supported_profiles: PROFILE_BASELINE | PROFILE_SHM_FUTEX,
92 + preferred_profiles: PROFILE_SHM_FUTEX,
93 + max_request_batch_items: 1,
94 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
95 + auth_token: AUTH_TOKEN,
96 + ..ClientConfig::default()
97 + }
98 +}
99 +
100 +fn fill_snapshot(builder: &mut CgroupsBuilder<'_>) -> bool {
101 + let items = [
102 + (
103 + 1001u32,
104 + 0u32,
105 + 1u32,
106 + b"docker-abc123" as &[u8],
107 + b"/sys/fs/cgroup/docker/abc123" as &[u8],
108 + ),
109 + (2002, 0, 1, b"k8s-pod-xyz", b"/sys/fs/cgroup/kubepods/xyz"),
110 + (
111 + 3003,
112 + 0,
113 + 0,
114 + b"systemd-user",
115 + b"/sys/fs/cgroup/user.slice/user-1000",
116 + ),
117 + ];
118 +
119 + for (hash, options, enabled, name, path) in &items {
120 + if builder.add(*hash, *options, *enabled, name, path).is_err() {
121 + return false;
122 + }
123 + }
124 +
125 + true
126 +}
127 +
128 +fn snapshot_handler() -> Handler {
129 + Handler {
130 + handle: Some(Arc::new(|req, builder| {
131 + if req.layout_version != 1 || req.flags != 0 {
132 + return false;
133 + }
134 + builder.set_header(1, 42);
135 + fill_snapshot(builder)
136 + })),
137 + snapshot_max_items: 3,
138 + }
139 +}
140 +
141 +fn connect_ready(client: &mut CgroupsClient) {
142 + for _ in 0..200 {
143 + client.refresh();
144 + if client.ready() {
145 + return;
146 + }
147 + thread::sleep(Duration::from_millis(10));
148 + }
149 +
150 + panic!("client did not reach READY state");
151 +}
152 +
153 +struct TestServer {
154 + stop_flag: Arc<std::sync::atomic::AtomicBool>,
155 + thread: Option<thread::JoinHandle<()>>,
156 +}
157 +
158 +impl TestServer {
159 + fn start(service: &str, config: ServerConfig) -> Self {
160 + ensure_run_dir();
161 + cleanup_all(service);
162 +
163 + let ready_flag = Arc::new(std::sync::atomic::AtomicBool::new(false));
164 + let ready_clone = ready_flag.clone();
165 + let mut server = ManagedServer::new(TEST_RUN_DIR, service, config, snapshot_handler());
166 + let stop_flag = server.running_flag();
167 + let thread = thread::spawn(move || {
168 + ready_clone.store(true, std::sync::atomic::Ordering::Release);
169 + let _ = server.run();
170 + });
171 +
172 + for _ in 0..2000 {
173 + if ready_flag.load(std::sync::atomic::Ordering::Acquire) {
174 + break;
175 + }
176 + thread::sleep(Duration::from_micros(500));
177 + }
178 + wait_for_listener_bind(service);
179 +
180 + Self {
181 + stop_flag,
182 + thread: Some(thread),
183 + }
184 + }
185 +}
186 +
187 +impl Drop for TestServer {
188 + fn drop(&mut self) {
189 + self.stop_flag
190 + .store(false, std::sync::atomic::Ordering::Release);
191 + if let Some(thread) = self.thread.take() {
192 + let _ = thread.join();
193 + }
194 + }
195 +}
196 +
197 +#[test]
198 +fn test_snapshot_round_trip_unix() {
199 + let service = unique_service("snapshot");
200 + let _server = TestServer::start(&service, server_config());
201 +
202 + let mut client = CgroupsClient::new(TEST_RUN_DIR, &service, client_config());
203 + connect_ready(&mut client);
204 +
205 + let view = client.call_snapshot().expect("snapshot");
206 + assert_eq!(view.item_count, 3);
207 + assert_eq!(view.systemd_enabled, 1);
208 + assert_eq!(view.generation, 42);
209 +
210 + let item0 = view.item(0).expect("item 0");
211 + assert_eq!(item0.hash, 1001);
212 + assert_eq!(item0.name.as_bytes(), b"docker-abc123");
213 + assert_eq!(item0.path.as_bytes(), b"/sys/fs/cgroup/docker/abc123");
214 +}
215 +
216 +#[cfg(target_os = "linux")]
217 +#[test]
218 +fn test_snapshot_round_trip_shm_unix() {
219 + let service = unique_service("snapshot_shm");
220 + let _server = TestServer::start(&service, shm_server_config());
221 +
222 + let mut client = CgroupsClient::new(TEST_RUN_DIR, &service, shm_client_config());
223 + connect_ready(&mut client);
224 +
225 + let view = client.call_snapshot().expect("snapshot");
226 + assert_eq!(view.item_count, 3);
227 + assert_eq!(view.generation, 42);
228 +}
229 +
230 +#[test]
231 +fn test_cache_round_trip_unix() {
232 + let service = unique_service("cache");
233 + let _server = TestServer::start(&service, server_config());
234 +
235 + let mut cache = CgroupsCache::new(TEST_RUN_DIR, &service, client_config());
236 + let mut updated = false;
237 + for _ in 0..200 {
238 + if cache.refresh() {
239 + updated = true;
240 + break;
241 + }
242 + thread::sleep(Duration::from_millis(10));
243 + }
244 + assert!(updated);
245 + assert!(cache.ready());
246 +
247 + let item = cache.lookup(1001, "docker-abc123").expect("lookup");
248 + assert_eq!(item.path, "/sys/fs/cgroup/docker/abc123");
249 +
250 + let status = cache.status();
251 + assert!(status.populated);
252 + assert_eq!(status.item_count, 3);
253 + assert_eq!(status.generation, 42);
254 +}
255 +
256 +#[test]
257 +fn test_client_not_ready_returns_error_unix() {
258 + let service = unique_service("not_ready");
259 + cleanup_all(&service);
260 +
261 + let mut client = CgroupsClient::new(TEST_RUN_DIR, &service, client_config());
262 + match client.call_snapshot() {
263 + Err(NipcError::BadLayout) => {}
264 + other => panic!("unexpected result: {other:?}"),
265 + }
266 +}
src/crates/netipc/src/service/cgroups_windows_tests.rs new
+224
@@ -0,0 +1,224 @@
1 +use super::*;
2 +use crate::protocol::{CgroupsBuilder, NipcError, PROFILE_BASELINE, PROFILE_SHM_HYBRID};
3 +use std::sync::atomic::{AtomicU64, Ordering};
4 +use std::sync::Arc;
5 +use std::thread;
6 +use std::time::Duration;
7 +
8 +const TEST_RUN_DIR: &str = r"C:\Temp\nipc_cgroups_rust_test";
9 +const AUTH_TOKEN: u64 = 0xDEADBEEFCAFEBABE;
10 +const RESPONSE_BUF_SIZE: usize = 65536;
11 +static SERVICE_COUNTER: AtomicU64 = AtomicU64::new(0);
12 +
13 +fn ensure_run_dir() {
14 + let _ = std::fs::create_dir_all(TEST_RUN_DIR);
15 +}
16 +
17 +fn unique_service(prefix: &str) -> String {
18 + format!(
19 + "{}_{}_{}",
20 + prefix,
21 + std::process::id(),
22 + SERVICE_COUNTER.fetch_add(1, Ordering::Relaxed) + 1
23 + )
24 +}
25 +
26 +fn server_config() -> ServerConfig {
27 + ServerConfig {
28 + supported_profiles: PROFILE_BASELINE,
29 + max_request_batch_items: 1,
30 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
31 + auth_token: AUTH_TOKEN,
32 + ..ServerConfig::default()
33 + }
34 +}
35 +
36 +fn client_config() -> ClientConfig {
37 + ClientConfig {
38 + supported_profiles: PROFILE_BASELINE,
39 + max_request_batch_items: 1,
40 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
41 + auth_token: AUTH_TOKEN,
42 + ..ClientConfig::default()
43 + }
44 +}
45 +
46 +fn shm_server_config() -> ServerConfig {
47 + ServerConfig {
48 + supported_profiles: PROFILE_SHM_HYBRID | PROFILE_BASELINE,
49 + max_request_batch_items: 1,
50 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
51 + auth_token: AUTH_TOKEN,
52 + ..ServerConfig::default()
53 + }
54 +}
55 +
56 +fn shm_client_config() -> ClientConfig {
57 + ClientConfig {
58 + supported_profiles: PROFILE_SHM_HYBRID | PROFILE_BASELINE,
59 + max_request_batch_items: 1,
60 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
61 + auth_token: AUTH_TOKEN,
62 + ..ClientConfig::default()
63 + }
64 +}
65 +
66 +fn fill_snapshot(builder: &mut CgroupsBuilder<'_>) -> bool {
67 + let items = [
68 + (
69 + 1001u32,
70 + 0u32,
71 + 1u32,
72 + b"docker-abc123" as &[u8],
73 + b"/sys/fs/cgroup/docker/abc123" as &[u8],
74 + ),
75 + (2002, 0, 1, b"k8s-pod-xyz", b"/sys/fs/cgroup/kubepods/xyz"),
76 + (
77 + 3003,
78 + 0,
79 + 0,
80 + b"systemd-user",
81 + b"/sys/fs/cgroup/user.slice/user-1000",
82 + ),
83 + ];
84 +
85 + for (hash, options, enabled, name, path) in &items {
86 + if builder.add(*hash, *options, *enabled, name, path).is_err() {
87 + return false;
88 + }
89 + }
90 +
91 + true
92 +}
93 +
94 +fn snapshot_handler() -> Handler {
95 + Handler {
96 + handle: Some(Arc::new(|req, builder| {
97 + if req.layout_version != 1 || req.flags != 0 {
98 + return false;
99 + }
100 + builder.set_header(1, 42);
101 + fill_snapshot(builder)
102 + })),
103 + snapshot_max_items: 3,
104 + }
105 +}
106 +
107 +fn connect_ready(client: &mut CgroupsClient) {
108 + for _ in 0..200 {
109 + client.refresh();
110 + if client.ready() {
111 + return;
112 + }
113 + thread::sleep(Duration::from_millis(10));
114 + }
115 +
116 + panic!("client did not reach READY state");
117 +}
118 +
119 +struct TestServer {
120 + service: String,
121 + wake_config: ClientConfig,
122 + stop_flag: Arc<std::sync::atomic::AtomicBool>,
123 + thread: Option<thread::JoinHandle<()>>,
124 +}
125 +
126 +impl TestServer {
127 + fn start(service: &str, config: ServerConfig) -> Self {
128 + ensure_run_dir();
129 +
130 + let svc = service.to_string();
131 + let wake_config = ClientConfig {
132 + supported_profiles: config.supported_profiles,
133 + preferred_profiles: config.preferred_profiles,
134 + max_request_batch_items: config.max_request_batch_items,
135 + max_response_payload_bytes: config.max_response_payload_bytes,
136 + auth_token: config.auth_token,
137 + ..ClientConfig::default()
138 + };
139 + let mut server = ManagedServer::new(TEST_RUN_DIR, service, config, snapshot_handler());
140 + let stop_flag = server.running_flag();
141 + let thread = thread::spawn(move || {
142 + let _ = server.run();
143 + });
144 +
145 + thread::sleep(Duration::from_millis(200));
146 +
147 + Self {
148 + service: svc,
149 + wake_config,
150 + stop_flag,
151 + thread: Some(thread),
152 + }
153 + }
154 +}
155 +
156 +impl Drop for TestServer {
157 + fn drop(&mut self) {
158 + self.stop_flag
159 + .store(false, std::sync::atomic::Ordering::Release);
160 + let mut wake = CgroupsClient::new(TEST_RUN_DIR, &self.service, self.wake_config.clone());
161 + let _ = wake.refresh();
162 + if let Some(thread) = self.thread.take() {
163 + let _ = thread.join();
164 + }
165 + }
166 +}
167 +
168 +#[test]
169 +fn test_snapshot_round_trip_windows() {
170 + let service = unique_service("snapshot");
171 + let _server = TestServer::start(&service, server_config());
172 +
173 + let mut client = CgroupsClient::new(TEST_RUN_DIR, &service, client_config());
174 + connect_ready(&mut client);
175 +
176 + let view = client.call_snapshot().expect("snapshot");
177 + assert_eq!(view.item_count, 3);
178 + assert_eq!(view.systemd_enabled, 1);
179 + assert_eq!(view.generation, 42);
180 +}
181 +
182 +#[test]
183 +fn test_snapshot_round_trip_win_shm() {
184 + let service = unique_service("snapshot_shm");
185 + let _server = TestServer::start(&service, shm_server_config());
186 +
187 + let mut client = CgroupsClient::new(TEST_RUN_DIR, &service, shm_client_config());
188 + connect_ready(&mut client);
189 +
190 + let view = client.call_snapshot().expect("snapshot");
191 + assert_eq!(view.item_count, 3);
192 + assert_eq!(view.generation, 42);
193 +}
194 +
195 +#[test]
196 +fn test_cache_round_trip_windows() {
197 + let service = unique_service("cache");
198 + let _server = TestServer::start(&service, server_config());
199 +
200 + let mut cache = CgroupsCache::new(TEST_RUN_DIR, &service, client_config());
201 + let mut updated = false;
202 + for _ in 0..200 {
203 + if cache.refresh() {
204 + updated = true;
205 + break;
206 + }
207 + thread::sleep(Duration::from_millis(10));
208 + }
209 + assert!(updated);
210 + assert!(cache.ready());
211 +
212 + let item = cache.lookup(1001, "docker-abc123").expect("lookup");
213 + assert_eq!(item.path, "/sys/fs/cgroup/docker/abc123");
214 +}
215 +
216 +#[test]
217 +fn test_client_not_ready_returns_error_windows() {
218 + let service = unique_service("not_ready");
219 + let mut client = CgroupsClient::new(TEST_RUN_DIR, &service, client_config());
220 + match client.call_snapshot() {
221 + Err(NipcError::BadLayout) => {}
222 + other => panic!("unexpected result: {other:?}"),
223 + }
224 +}
src/crates/netipc/src/service/mod.rs new
+9
@@ -0,0 +1,9 @@
1 +//! L2 orchestration: service-specific client contexts and managed servers.
2 +//!
3 +//! Production-facing modules are service-kind specific. Internal helpers may
4 +//! remain generic for tests and benchmarks, but every running endpoint still
5 +//! serves exactly one request kind.
6 +
7 +pub mod cgroups;
8 +#[doc(hidden)]
9 +pub mod raw;
src/crates/netipc/src/service/raw.rs new
+2643
@@ -0,0 +1,2643 @@
1 +//! Internal single-kind L2 service helpers for tests and benchmarks.
2 +//!
3 +//! This module is not the production service contract. It exists to preserve
4 +//! internal benchmark/stress coverage while the public service modules expose
5 +//! one service kind per endpoint.
6 +
7 +use crate::protocol::{
8 + self, batch_item_get, increment_decode, increment_encode, string_reverse_decode,
9 + string_reverse_encode, BatchBuilder, CgroupsRequest, CgroupsResponseView, Header, NipcError,
10 + FLAG_BATCH, HEADER_SIZE, INCREMENT_PAYLOAD_SIZE, KIND_REQUEST, KIND_RESPONSE, MAGIC_MSG,
11 + MAX_PAYLOAD_CAP, MAX_PAYLOAD_DEFAULT, METHOD_CGROUPS_SNAPSHOT, METHOD_INCREMENT,
12 + METHOD_STRING_REVERSE, STATUS_BAD_ENVELOPE, STATUS_INTERNAL_ERROR, STATUS_LIMIT_EXCEEDED,
13 + STATUS_OK, STRING_REVERSE_HDR_SIZE, VERSION,
14 +};
15 +
16 +#[cfg(unix)]
17 +use crate::protocol::{PROFILE_SHM_FUTEX, PROFILE_SHM_HYBRID};
18 +
19 +#[cfg(unix)]
20 +use crate::transport::posix::{ClientConfig, ServerConfig, UdsListener, UdsSession};
21 +
22 +#[cfg(target_os = "linux")]
23 +use crate::transport::shm::ShmContext;
24 +
25 +#[cfg(windows)]
26 +use crate::transport::windows::{ClientConfig, NpError, NpListener, NpSession, ServerConfig};
27 +
28 +#[cfg(windows)]
29 +use crate::transport::win_shm::{
30 + WinShmContext, PROFILE_BUSYWAIT as WIN_SHM_PROFILE_BUSYWAIT,
31 + PROFILE_HYBRID as WIN_SHM_PROFILE_HYBRID,
32 +};
33 +
34 +use std::sync::atomic::{AtomicBool, AtomicU32, Ordering};
35 +use std::sync::Arc;
36 +
37 +/// Poll/receive timeout for server loops (ms). Controls shutdown detection latency.
38 +const SERVER_POLL_TIMEOUT_MS: u32 = 100;
39 +const CLIENT_SHM_ATTACH_RETRY_INTERVAL_MS: u64 = 5;
40 +const CLIENT_SHM_ATTACH_RETRY_TIMEOUT_MS: u64 = 5_000;
41 +
42 +fn next_power_of_2_u32(n: u32) -> u32 {
43 + if n < 16 {
44 + return 16;
45 + }
46 +
47 + // Cap at 2^31 — the largest power of 2 that fits in u32
48 + if n > (1u32 << 31) {
49 + return 1u32 << 31;
50 + }
51 +
52 + let mut value = n - 1;
53 + value |= value >> 1;
54 + value |= value >> 2;
55 + value |= value >> 4;
56 + value |= value >> 8;
57 + value |= value >> 16;
58 + value + 1
59 +}
60 +
61 +// ---------------------------------------------------------------------------
62 +// Client state
63 +// ---------------------------------------------------------------------------
64 +
65 +/// Client connection state machine.
66 +#[derive(Debug, Clone, Copy, PartialEq, Eq)]
67 +pub enum ClientState {
68 + Disconnected,
69 + Connecting,
70 + Ready,
71 + NotFound,
72 + AuthFailed,
73 + Incompatible,
74 + Broken,
75 +}
76 +
77 +/// Diagnostic counters snapshot.
78 +#[derive(Debug, Clone, Copy, PartialEq, Eq)]
79 +pub struct ClientStatus {
80 + pub state: ClientState,
81 + pub connect_count: u32,
82 + pub reconnect_count: u32,
83 + pub call_count: u32,
84 + pub error_count: u32,
85 +}
86 +
87 +// ---------------------------------------------------------------------------
88 +// Client context
89 +// ---------------------------------------------------------------------------
90 +
91 +/// L2 client context bound to one service kind.
92 +///
93 +/// Manages connection lifecycle and provides typed blocking calls with
94 +/// at-least-once retry semantics. The outer request code remains only for
95 +/// validation; each client instance is bound to one expected request kind.
96 +pub struct RawClient {
97 + state: ClientState,
98 + run_dir: String,
99 + service_name: String,
100 + expected_method_code: u16,
101 + transport_config: ClientConfig,
102 +
103 + // Connection (managed internally)
104 + #[cfg(unix)]
105 + session: Option<UdsSession>,
106 + #[cfg(target_os = "linux")]
107 + shm: Option<ShmContext>,
108 +
109 + #[cfg(windows)]
110 + session: Option<NpSession>,
111 + #[cfg(windows)]
112 + shm: Option<WinShmContext>,
113 +
114 + // Reusable scratch buffers owned by the client for hot request paths.
115 + request_buf: Vec<u8>,
116 + send_buf: Vec<u8>,
117 + transport_buf: Vec<u8>,
118 +
119 + // Stats
120 + connect_count: u32,
121 + reconnect_count: u32,
122 + call_count: u32,
123 + error_count: u32,
124 +}
125 +
126 +impl RawClient {
127 + fn new_bound(
128 + run_dir: &str,
129 + service_name: &str,
130 + expected_method_code: u16,
131 + config: ClientConfig,
132 + ) -> Self {
133 + RawClient {
134 + state: ClientState::Disconnected,
135 + run_dir: run_dir.to_string(),
136 + service_name: service_name.to_string(),
137 + expected_method_code,
138 + transport_config: config,
139 + session: None,
140 + #[cfg(target_os = "linux")]
141 + shm: None,
142 + #[cfg(windows)]
143 + shm: None,
144 + request_buf: Vec::new(),
145 + send_buf: Vec::new(),
146 + transport_buf: Vec::new(),
147 + connect_count: 0,
148 + reconnect_count: 0,
149 + call_count: 0,
150 + error_count: 0,
151 + }
152 + }
153 +
154 + /// Create a new client context bound to the cgroups-snapshot service kind.
155 + /// Does NOT connect. Does NOT require the server to be running.
156 + pub fn new_snapshot(run_dir: &str, service_name: &str, config: ClientConfig) -> Self {
157 + Self::new_bound(run_dir, service_name, METHOD_CGROUPS_SNAPSHOT, config)
158 + }
159 +
160 + /// Create a new client context bound to the increment service kind.
161 + /// Does NOT connect. Does NOT require the server to be running.
162 + pub fn new_increment(run_dir: &str, service_name: &str, config: ClientConfig) -> Self {
163 + Self::new_bound(run_dir, service_name, METHOD_INCREMENT, config)
164 + }
165 +
166 + /// Create a new client context bound to the string-reverse service kind.
167 + /// Does NOT connect. Does NOT require the server to be running.
168 + pub fn new_string_reverse(run_dir: &str, service_name: &str, config: ClientConfig) -> Self {
169 + Self::new_bound(run_dir, service_name, METHOD_STRING_REVERSE, config)
170 + }
171 +
172 + /// Attempt connect if DISCONNECTED/NOT_FOUND, reconnect if BROKEN.
173 + /// Returns true if the state changed.
174 + pub fn refresh(&mut self) -> bool {
175 + let old_state = self.state;
176 +
177 + match self.state {
178 + ClientState::Disconnected | ClientState::NotFound => {
179 + self.state = ClientState::Connecting;
180 + self.state = self.try_connect();
181 + if self.state == ClientState::Ready {
182 + self.connect_count += 1;
183 + }
184 + }
185 + ClientState::Broken => {
186 + self.disconnect();
187 + self.state = ClientState::Connecting;
188 + self.state = self.try_connect();
189 + if self.state == ClientState::Ready {
190 + self.reconnect_count += 1;
191 + }
192 + }
193 + ClientState::Ready
194 + | ClientState::Connecting
195 + | ClientState::AuthFailed
196 + | ClientState::Incompatible => {}
197 + }
198 +
199 + self.state != old_state
200 + }
201 +
202 + /// Cheap cached boolean. No I/O, no syscalls.
203 + #[inline]
204 + pub fn ready(&self) -> bool {
205 + self.state == ClientState::Ready
206 + }
207 +
208 + /// Detailed status snapshot for diagnostics.
209 + pub fn status(&self) -> ClientStatus {
210 + ClientStatus {
211 + state: self.state,
212 + connect_count: self.connect_count,
213 + reconnect_count: self.reconnect_count,
214 + call_count: self.call_count,
215 + error_count: self.error_count,
216 + }
217 + }
218 +
219 + fn session_max_request_payload_bytes(&self) -> u32 {
220 + #[cfg(unix)]
221 + if let Some(ref session) = self.session {
222 + return session.max_request_payload_bytes;
223 + }
224 +
225 + #[cfg(windows)]
226 + if let Some(ref session) = self.session {
227 + return session.max_request_payload_bytes;
228 + }
229 +
230 + self.transport_config.max_request_payload_bytes
231 + }
232 +
233 + fn session_max_response_payload_bytes(&self) -> u32 {
234 + #[cfg(unix)]
235 + if let Some(ref session) = self.session {
236 + return session.max_response_payload_bytes;
237 + }
238 +
239 + #[cfg(windows)]
240 + if let Some(ref session) = self.session {
241 + return session.max_response_payload_bytes;
242 + }
243 +
244 + self.transport_config.max_response_payload_bytes
245 + }
246 +
247 + fn client_note_request_capacity(&mut self, payload_len: u32) {
248 + let grown = next_power_of_2_u32(payload_len).min(MAX_PAYLOAD_CAP);
249 + if grown > self.transport_config.max_request_payload_bytes {
250 + self.transport_config.max_request_payload_bytes = grown;
251 + }
252 + }
253 +
254 + fn client_note_response_capacity(&mut self, payload_len: u32) {
255 + let grown = next_power_of_2_u32(payload_len).min(MAX_PAYLOAD_CAP);
256 + if grown > self.transport_config.max_response_payload_bytes {
257 + self.transport_config.max_response_payload_bytes = grown;
258 + }
259 + }
260 +
261 + fn validate_method(&self, method_code: u16) -> Result<(), NipcError> {
262 + if self.expected_method_code == method_code {
263 + Ok(())
264 + } else {
265 + Err(NipcError::BadLayout)
266 + }
267 + }
268 +
269 + /// Blocking typed call: encode request, send, receive, check
270 + /// transport_status, decode response.
271 + ///
272 + /// The returned view is valid until the next typed call on this client.
273 + ///
274 + /// Retry policy (per spec): if the call fails and the context was
275 + /// previously READY, disconnect, reconnect (full handshake), and retry.
276 + /// Ordinary failures retry once. Overflow-driven resize recovery may
277 + /// reconnect more than once while negotiated capacities grow.
278 + pub fn call_snapshot(&mut self) -> Result<CgroupsResponseView<'_>, NipcError> {
279 + self.validate_method(METHOD_CGROUPS_SNAPSHOT)?;
280 + let req = CgroupsRequest {
281 + layout_version: 1,
282 + flags: 0,
283 + };
284 + let mut req_buf = [0u8; 4];
285 + let req_len = req.encode(&mut req_buf);
286 + if req_len == 0 {
287 + return Err(NipcError::Truncated);
288 + }
289 +
290 + let response = self.raw_call_with_retry(METHOD_CGROUPS_SNAPSHOT, &req_buf[..req_len])?;
291 + CgroupsResponseView::decode(self.response_payload(response)?)
292 + }
293 +
294 + /// Blocking typed call: INCREMENT method.
295 + /// Sends a u64 value, receives the incremented u64 back.
296 + pub fn call_increment(&mut self, value: u64) -> Result<u64, NipcError> {
297 + self.validate_method(METHOD_INCREMENT)?;
298 + let mut req_buf = [0u8; INCREMENT_PAYLOAD_SIZE];
299 + let req_len = increment_encode(value, &mut req_buf);
300 + if req_len == 0 {
301 + return Err(NipcError::Truncated);
302 + }
303 +
304 + let response = self.raw_call_with_retry(METHOD_INCREMENT, &req_buf[..req_len])?;
305 + increment_decode(self.response_payload(response)?)
306 + }
307 +
308 + /// Blocking typed call: STRING_REVERSE method.
309 + /// Sends a string, receives the reversed string back.
310 + ///
311 + /// The returned view is valid until the next typed call on this client.
312 + pub fn call_string_reverse(
313 + &mut self,
314 + s: &str,
315 + ) -> Result<protocol::StringReverseView<'_>, NipcError> {
316 + self.validate_method(METHOD_STRING_REVERSE)?;
317 + let req_size = STRING_REVERSE_HDR_SIZE + s.len() + 1;
318 + let req_buf = ensure_client_scratch(&mut self.request_buf, req_size);
319 + let req_len = string_reverse_encode(s.as_bytes(), req_buf);
320 + if req_len == 0 {
321 + return Err(NipcError::Truncated);
322 + }
323 +
324 + let response = self.raw_call_with_retry_request_buf(METHOD_STRING_REVERSE, req_len)?;
325 + string_reverse_decode(self.response_payload(response)?)
326 + }
327 +
328 + /// Blocking typed batch call: INCREMENT method.
329 + /// Sends multiple u64 values, receives the incremented u64s back.
330 + pub fn call_increment_batch(&mut self, values: &[u64]) -> Result<Vec<u64>, NipcError> {
331 + self.validate_method(METHOD_INCREMENT)?;
332 + if values.is_empty() {
333 + return Ok(Vec::new());
334 + }
335 +
336 + // Single value: use the non-batch path
337 + if values.len() == 1 {
338 + let r = self.call_increment(values[0])?;
339 + return Ok(vec![r]);
340 + }
341 +
342 + let count = values.len() as u32;
343 +
344 + let req_buf_size = protocol::align8(count as usize * 8)
345 + + count as usize * protocol::align8(INCREMENT_PAYLOAD_SIZE)
346 + + 64;
347 + let req_buf = ensure_client_scratch(&mut self.request_buf, req_buf_size);
348 + let req_len = {
349 + let mut bb = BatchBuilder::new(req_buf, count);
350 + for &v in values {
351 + let mut item_buf = [0u8; INCREMENT_PAYLOAD_SIZE];
352 + if increment_encode(v, &mut item_buf) == 0 {
353 + return Err(NipcError::Truncated);
354 + }
355 + bb.add(&item_buf).map_err(|_| NipcError::Overflow)?;
356 + }
357 + let (req_len, _out_count) = bb.finish();
358 + req_len
359 + };
360 +
361 + let response =
362 + self.raw_batch_call_with_retry_request_buf(METHOD_INCREMENT, req_len, count)?;
363 + let resp_payload = self.response_payload(response)?;
364 + let mut results = Vec::with_capacity(values.len());
365 + for i in 0..count {
366 + let (item_data, _item_len) = batch_item_get(resp_payload, count, i)?;
367 + let val = increment_decode(item_data)?;
368 + results.push(val);
369 + }
370 +
371 + Ok(results)
372 + }
373 +
374 + /// Tear down connection and release resources.
375 + pub fn close(&mut self) {
376 + self.disconnect();
377 + self.state = ClientState::Disconnected;
378 + }
379 +
380 + // ------------------------------------------------------------------
381 + // Internal helpers
382 + // ------------------------------------------------------------------
383 +
384 + /// Tear down the current connection.
385 + fn disconnect(&mut self) {
386 + #[cfg(target_os = "linux")]
387 + {
388 + if let Some(mut shm) = self.shm.take() {
389 + shm.close();
390 + }
391 + }
392 +
393 + #[cfg(windows)]
394 + {
395 + if let Some(mut shm) = self.shm.take() {
396 + shm.close();
397 + }
398 + }
399 +
400 + // Drop the session (closes handle/fd via Drop impl)
401 + self.session.take();
402 + }
403 +
404 + /// Attempt a full connection: transport connect + handshake, then SHM
405 + /// upgrade if negotiated.
406 + #[cfg(unix)]
407 + fn try_connect(&mut self) -> ClientState {
408 + match UdsSession::connect(&self.run_dir, &self.service_name, &self.transport_config) {
409 + Ok(session) => {
410 + #[cfg(target_os = "linux")]
411 + let selected_profile = session.selected_profile;
412 + #[cfg(target_os = "linux")]
413 + let session_id = session.session_id;
414 +
415 + // SHM upgrade if negotiated
416 + #[cfg(target_os = "linux")]
417 + {
418 + if selected_profile == PROFILE_SHM_HYBRID
419 + || selected_profile == PROFILE_SHM_FUTEX
420 + {
421 + // Retry attach: server creates the SHM region after
422 + // the UDS handshake, so it may not exist yet.
423 + let mut shm_ok = false;
424 + let deadline = std::time::Instant::now()
425 + + std::time::Duration::from_millis(CLIENT_SHM_ATTACH_RETRY_TIMEOUT_MS);
426 + loop {
427 + match ShmContext::client_attach(
428 + &self.run_dir,
429 + &self.service_name,
430 + session_id,
431 + ) {
432 + Ok(ctx) => {
433 + self.shm = Some(ctx);
434 + shm_ok = true;
435 + break;
436 + }
437 + Err(_) => {
438 + if std::time::Instant::now() >= deadline {
439 + break;
440 + }
441 + std::thread::sleep(std::time::Duration::from_millis(
442 + CLIENT_SHM_ATTACH_RETRY_INTERVAL_MS,
443 + ));
444 + }
445 + }
446 + }
447 + if !shm_ok {
448 + // SHM attach failed after negotiation. Close that session,
449 + // blacklist SHM for this client context, and retry baseline.
450 + drop(session);
451 + self.transport_config.supported_profiles &=
452 + !(PROFILE_SHM_HYBRID | PROFILE_SHM_FUTEX);
453 + self.transport_config.preferred_profiles &=
454 + !(PROFILE_SHM_HYBRID | PROFILE_SHM_FUTEX);
455 + if self.transport_config.supported_profiles == 0 {
456 + return ClientState::Disconnected;
457 + }
458 + return self.try_connect();
459 + }
460 + }
461 + }
462 +
463 + self.session = Some(session);
464 + ClientState::Ready
465 + }
466 + Err(e) => {
467 + use crate::transport::posix::UdsError;
468 + match e {
469 + UdsError::Connect(_) => ClientState::NotFound,
470 + UdsError::AuthFailed => ClientState::AuthFailed,
471 + UdsError::NoProfile => ClientState::Incompatible,
472 + UdsError::Incompatible(_) => ClientState::Incompatible,
473 + _ => ClientState::Disconnected,
474 + }
475 + }
476 + }
477 + }
478 +
479 + /// Windows: attempt a full Named Pipe connection + Win SHM upgrade.
480 + #[cfg(windows)]
481 + fn try_connect(&mut self) -> ClientState {
482 + match NpSession::connect(&self.run_dir, &self.service_name, &self.transport_config) {
483 + Ok(session) => {
484 + let selected_profile = session.selected_profile;
485 +
486 + // Win SHM upgrade if negotiated
487 + if selected_profile == WIN_SHM_PROFILE_HYBRID
488 + || selected_profile == WIN_SHM_PROFILE_BUSYWAIT
489 + {
490 + let mut shm_ok = false;
491 + let deadline = std::time::Instant::now()
492 + + std::time::Duration::from_millis(CLIENT_SHM_ATTACH_RETRY_TIMEOUT_MS);
493 + loop {
494 + match WinShmContext::client_attach(
495 + &self.run_dir,
496 + &self.service_name,
497 + self.transport_config.auth_token,
498 + session.session_id,
499 + selected_profile,
500 + ) {
501 + Ok(ctx) => {
502 + self.shm = Some(ctx);
503 + shm_ok = true;
504 + break;
505 + }
506 + Err(_) => {
507 + if std::time::Instant::now() >= deadline {
508 + break;
509 + }
510 + std::thread::sleep(std::time::Duration::from_millis(
511 + CLIENT_SHM_ATTACH_RETRY_INTERVAL_MS,
512 + ));
513 + }
514 + }
515 + }
516 + if !shm_ok {
517 + // WinSHM attach failed after negotiation. Close that
518 + // session, blacklist WinSHM for this client context,
519 + // and retry baseline.
520 + drop(session);
521 + self.transport_config.supported_profiles &=
522 + !(WIN_SHM_PROFILE_HYBRID | WIN_SHM_PROFILE_BUSYWAIT);
523 + self.transport_config.preferred_profiles &=
524 + !(WIN_SHM_PROFILE_HYBRID | WIN_SHM_PROFILE_BUSYWAIT);
525 + if self.transport_config.supported_profiles == 0 {
526 + return ClientState::Disconnected;
527 + }
528 + return self.try_connect();
529 + }
530 + }
531 +
532 + self.session = Some(session);
533 + ClientState::Ready
534 + }
535 + Err(e) => match e {
536 + NpError::Connect(_) => ClientState::NotFound,
537 + NpError::AuthFailed => ClientState::AuthFailed,
538 + NpError::NoProfile => ClientState::Incompatible,
539 + NpError::Incompatible(_) => ClientState::Incompatible,
540 + _ => ClientState::Disconnected,
541 + },
542 + }
543 + }
544 +
545 + /// Reconnect-driven recovery for a single-item raw call.
546 + /// Ordinary failures retry once. Overflow-driven resize recovery may
547 + /// reconnect more than once while negotiated capacities grow.
548 + fn raw_call_with_retry<'a>(
549 + &mut self,
550 + method_code: u16,
551 + request_payload: &[u8],
552 + ) -> Result<ClientResponseRef, NipcError> {
553 + if self.state != ClientState::Ready {
554 + self.error_count += 1;
555 + return Err(NipcError::BadLayout);
556 + }
557 +
558 + // Cap overflow-driven retries: payloads grow by powers of 2, so 8
559 + // retries allows ~256x growth from the initial negotiated size.
560 + let mut overflow_retries = 0u32;
561 + loop {
562 + let prev_req = self.session_max_request_payload_bytes();
563 + let prev_resp = self.session_max_response_payload_bytes();
564 + let prev_cfg_req = self.transport_config.max_request_payload_bytes;
565 + let prev_cfg_resp = self.transport_config.max_response_payload_bytes;
566 +
567 + match self.do_raw_call(method_code, request_payload) {
568 + Ok(payload) => {
569 + self.call_count += 1;
570 + return Ok(payload);
571 + }
572 + Err(first_err) => {
573 + if first_err != NipcError::Overflow {
574 + self.disconnect();
575 + self.state = ClientState::Broken;
576 + self.state = self.try_connect();
577 + if self.state != ClientState::Ready {
578 + self.error_count += 1;
579 + return Err(first_err);
580 + }
581 + self.reconnect_count += 1;
582 +
583 + match self.do_raw_call(method_code, request_payload) {
584 + Ok(payload) => {
585 + self.call_count += 1;
586 + return Ok(payload);
587 + }
588 + Err(retry_err) => {
589 + self.disconnect();
590 + self.state = ClientState::Broken;
591 + self.error_count += 1;
592 + return Err(retry_err);
593 + }
594 + }
595 + }
596 +
597 + self.disconnect();
598 + self.state = ClientState::Broken;
599 + self.state = self.try_connect();
600 + if self.state != ClientState::Ready {
601 + self.error_count += 1;
602 + return Err(first_err);
603 + }
604 + self.reconnect_count += 1;
605 +
606 + if self.session_max_request_payload_bytes() <= prev_req
607 + && self.session_max_response_payload_bytes() <= prev_resp
608 + && self.transport_config.max_request_payload_bytes <= prev_cfg_req
609 + && self.transport_config.max_response_payload_bytes <= prev_cfg_resp
610 + {
611 + self.disconnect();
612 + self.state = ClientState::Broken;
613 + self.error_count += 1;
614 + return Err(first_err);
615 + }
616 +
617 + overflow_retries += 1;
618 + if overflow_retries >= 8 {
619 + self.disconnect();
620 + self.state = ClientState::Broken;
621 + self.error_count += 1;
622 + return Err(first_err);
623 + }
624 + }
625 + }
626 + }
627 + }
628 +
629 + fn raw_call_with_retry_request_buf<'a>(
630 + &mut self,
631 + method_code: u16,
632 + req_len: usize,
633 + ) -> Result<ClientResponseRef, NipcError> {
634 + if self.state != ClientState::Ready {
635 + self.error_count += 1;
636 + return Err(NipcError::BadLayout);
637 + }
638 +
639 + let mut overflow_retries = 0u32;
640 + loop {
641 + let prev_req = self.session_max_request_payload_bytes();
642 + let prev_resp = self.session_max_response_payload_bytes();
643 + let prev_cfg_req = self.transport_config.max_request_payload_bytes;
644 + let prev_cfg_resp = self.transport_config.max_response_payload_bytes;
645 +
646 + match self.do_raw_call_from_request_buf(method_code, req_len) {
647 + Ok(payload) => {
648 + self.call_count += 1;
649 + return Ok(payload);
650 + }
651 + Err(first_err) => {
652 + if first_err != NipcError::Overflow {
653 + self.disconnect();
654 + self.state = ClientState::Broken;
655 + self.state = self.try_connect();
656 + if self.state != ClientState::Ready {
657 + self.error_count += 1;
658 + return Err(first_err);
659 + }
660 + self.reconnect_count += 1;
661 +
662 + match self.do_raw_call_from_request_buf(method_code, req_len) {
663 + Ok(payload) => {
664 + self.call_count += 1;
665 + return Ok(payload);
666 + }
667 + Err(retry_err) => {
668 + self.disconnect();
669 + self.state = ClientState::Broken;
670 + self.error_count += 1;
671 + return Err(retry_err);
672 + }
673 + }
674 + }
675 +
676 + self.disconnect();
677 + self.state = ClientState::Broken;
678 + self.state = self.try_connect();
679 + if self.state != ClientState::Ready {
680 + self.error_count += 1;
681 + return Err(first_err);
682 + }
683 + self.reconnect_count += 1;
684 +
685 + if self.session_max_request_payload_bytes() <= prev_req
686 + && self.session_max_response_payload_bytes() <= prev_resp
687 + && self.transport_config.max_request_payload_bytes <= prev_cfg_req
688 + && self.transport_config.max_response_payload_bytes <= prev_cfg_resp
689 + {
690 + self.disconnect();
691 + self.state = ClientState::Broken;
692 + self.error_count += 1;
693 + return Err(first_err);
694 + }
695 +
696 + overflow_retries += 1;
697 + if overflow_retries >= 8 {
698 + self.disconnect();
699 + self.state = ClientState::Broken;
700 + self.error_count += 1;
701 + return Err(first_err);
702 + }
703 + }
704 + }
705 + }
706 + }
707 +
708 + fn raw_batch_call_with_retry_request_buf<'a>(
709 + &mut self,
710 + method_code: u16,
711 + req_len: usize,
712 + item_count: u32,
713 + ) -> Result<ClientResponseRef, NipcError> {
714 + if self.state != ClientState::Ready {
715 + self.error_count += 1;
716 + return Err(NipcError::BadLayout);
717 + }
718 +
719 + let mut overflow_retries = 0u32;
720 + loop {
721 + let prev_req = self.session_max_request_payload_bytes();
722 + let prev_resp = self.session_max_response_payload_bytes();
723 + let prev_cfg_req = self.transport_config.max_request_payload_bytes;
724 + let prev_cfg_resp = self.transport_config.max_response_payload_bytes;
725 +
726 + match self.do_raw_batch_call_from_request_buf(method_code, req_len, item_count) {
727 + Ok(payload) => {
728 + self.call_count += 1;
729 + return Ok(payload);
730 + }
731 + Err(first_err) => {
732 + if first_err != NipcError::Overflow {
733 + self.disconnect();
734 + self.state = ClientState::Broken;
735 + self.state = self.try_connect();
736 + if self.state != ClientState::Ready {
737 + self.error_count += 1;
738 + return Err(first_err);
739 + }
740 + self.reconnect_count += 1;
741 +
742 + match self.do_raw_batch_call_from_request_buf(
743 + method_code,
744 + req_len,
745 + item_count,
746 + ) {
747 + Ok(payload) => {
748 + self.call_count += 1;
749 + return Ok(payload);
750 + }
751 + Err(retry_err) => {
752 + self.disconnect();
753 + self.state = ClientState::Broken;
754 + self.error_count += 1;
755 + return Err(retry_err);
756 + }
757 + }
758 + }
759 +
760 + self.disconnect();
761 + self.state = ClientState::Broken;
762 + self.state = self.try_connect();
763 + if self.state != ClientState::Ready {
764 + self.error_count += 1;
765 + return Err(first_err);
766 + }
767 + self.reconnect_count += 1;
768 +
769 + if self.session_max_request_payload_bytes() <= prev_req
770 + && self.session_max_response_payload_bytes() <= prev_resp
771 + && self.transport_config.max_request_payload_bytes <= prev_cfg_req
772 + && self.transport_config.max_response_payload_bytes <= prev_cfg_resp
773 + {
774 + self.disconnect();
775 + self.state = ClientState::Broken;
776 + self.error_count += 1;
777 + return Err(first_err);
778 + }
779 +
780 + overflow_retries += 1;
781 + if overflow_retries >= 8 {
782 + self.disconnect();
783 + self.state = ClientState::Broken;
784 + self.error_count += 1;
785 + return Err(first_err);
786 + }
787 + }
788 + }
789 + }
790 + }
791 +
792 + /// Single attempt at a raw call for any method.
793 + fn do_raw_call(
794 + &mut self,
795 + method_code: u16,
796 + request_payload: &[u8],
797 + ) -> Result<ClientResponseRef, NipcError> {
798 + // 1. Build outer header
799 + let mut hdr = Header {
800 + kind: KIND_REQUEST,
801 + code: method_code,
802 + flags: 0,
803 + item_count: 1,
804 + message_id: (self.call_count as u64) + 1,
805 + transport_status: STATUS_OK,
806 + ..Header::default()
807 + };
808 +
809 + // 2. Send via L1 (SHM or UDS)
810 + self.transport_send(&mut hdr, request_payload)?;
811 +
812 + // 3. Receive via L1
813 + let (resp_hdr, response) = self.transport_receive()?;
814 +
815 + // 4. Verify response envelope fields before decode
816 + if resp_hdr.kind != KIND_RESPONSE {
817 + return Err(NipcError::BadKind);
818 + }
819 + if resp_hdr.code != method_code {
820 + return Err(NipcError::BadLayout);
821 + }
822 + if resp_hdr.message_id != hdr.message_id {
823 + return Err(NipcError::BadLayout);
824 + }
825 +
826 + // 5. Check transport_status BEFORE decode (spec requirement)
827 + match resp_hdr.transport_status {
828 + STATUS_OK => {}
829 + STATUS_LIMIT_EXCEEDED => {
830 + let current = self.session_max_response_payload_bytes();
831 + if current > 0 {
832 + self.client_note_response_capacity(current.saturating_mul(2));
833 + }
834 + return Err(NipcError::Overflow);
835 + }
836 + _ => return Err(NipcError::BadLayout),
837 + }
838 + Ok(response)
839 + }
840 +
841 + fn do_raw_call_from_request_buf(
842 + &mut self,
843 + method_code: u16,
844 + req_len: usize,
845 + ) -> Result<ClientResponseRef, NipcError> {
846 + let mut hdr = Header {
847 + kind: KIND_REQUEST,
848 + code: method_code,
849 + flags: 0,
850 + item_count: 1,
851 + message_id: (self.call_count as u64) + 1,
852 + transport_status: STATUS_OK,
853 + ..Header::default()
854 + };
855 +
856 + self.transport_send_request_buf(&mut hdr, req_len)?;
857 + let (resp_hdr, response) = self.transport_receive()?;
858 +
859 + if resp_hdr.kind != KIND_RESPONSE {
860 + return Err(NipcError::BadKind);
861 + }
862 + if resp_hdr.code != method_code {
863 + return Err(NipcError::BadLayout);
864 + }
865 + if resp_hdr.message_id != hdr.message_id {
866 + return Err(NipcError::BadLayout);
867 + }
868 + match resp_hdr.transport_status {
869 + STATUS_OK => {}
870 + STATUS_LIMIT_EXCEEDED => {
871 + let current = self.session_max_response_payload_bytes();
872 + if current > 0 {
873 + self.client_note_response_capacity(current.saturating_mul(2));
874 + }
875 + return Err(NipcError::Overflow);
876 + }
877 + _ => return Err(NipcError::BadLayout),
878 + }
879 + Ok(response)
880 + }
881 +
882 + /// Single attempt at a raw batch call. Like `do_raw_call` but sets
883 + /// FLAG_BATCH and item_count, and validates the response matches.
884 + fn do_raw_batch_call_from_request_buf(
885 + &mut self,
886 + method_code: u16,
887 + req_len: usize,
888 + item_count: u32,
889 + ) -> Result<ClientResponseRef, NipcError> {
890 + let mut hdr = Header {
891 + kind: KIND_REQUEST,
892 + code: method_code,
893 + flags: FLAG_BATCH,
894 + item_count,
895 + message_id: (self.call_count as u64) + 1,
896 + transport_status: STATUS_OK,
897 + ..Header::default()
898 + };
899 +
900 + self.transport_send_request_buf(&mut hdr, req_len)?;
901 +
902 + let (resp_hdr, response) = self.transport_receive()?;
903 +
904 + if resp_hdr.kind != KIND_RESPONSE {
905 + return Err(NipcError::BadKind);
906 + }
907 + if resp_hdr.code != method_code {
908 + return Err(NipcError::BadLayout);
909 + }
910 + if resp_hdr.message_id != hdr.message_id {
911 + return Err(NipcError::BadLayout);
912 + }
913 + match resp_hdr.transport_status {
914 + STATUS_OK => {}
915 + STATUS_LIMIT_EXCEEDED => {
916 + let current = self.session_max_response_payload_bytes();
917 + if current > 0 {
918 + self.client_note_response_capacity(current.saturating_mul(2));
919 + }
920 + return Err(NipcError::Overflow);
921 + }
922 + _ => return Err(NipcError::BadLayout),
923 + }
924 + if resp_hdr.item_count != item_count {
925 + return Err(NipcError::BadItemCount);
926 + }
927 + Ok(response)
928 + }
929 +
930 + /// Send via the active transport (SHM if available, baseline otherwise).
931 + fn transport_send(&mut self, hdr: &mut Header, payload: &[u8]) -> Result<(), NipcError> {
932 + let max_request_payload_bytes = self.session_max_request_payload_bytes();
933 +
934 + // SHM path (POSIX or Windows)
935 + #[cfg(target_os = "linux")]
936 + {
937 + if let Some(ref mut shm) = self.shm {
938 + if payload.len() > max_request_payload_bytes as usize {
939 + self.client_note_request_capacity(payload.len() as u32);
940 + return Err(NipcError::Overflow);
941 + }
942 +
943 + let msg_len = HEADER_SIZE + payload.len();
944 + let msg = ensure_client_scratch(&mut self.send_buf, msg_len);
945 +
946 + hdr.magic = MAGIC_MSG;
947 + hdr.version = VERSION;
948 + hdr.header_len = protocol::HEADER_LEN;
949 + hdr.payload_len = payload.len() as u32;
950 +
951 + hdr.encode(&mut msg[..HEADER_SIZE]);
952 + if !payload.is_empty() {
953 + msg[HEADER_SIZE..].copy_from_slice(payload);
954 + }
955 +
956 + let send_result = shm.send(&msg);
957 + return match send_result {
958 + Ok(()) => Ok(()),
959 + Err(crate::transport::shm::ShmError::MsgTooLarge) => {
960 + self.client_note_request_capacity(payload.len() as u32);
961 + Err(NipcError::Overflow)
962 + }
963 + Err(_) => Err(NipcError::Truncated),
964 + };
965 + }
966 + }
967 +
968 + #[cfg(windows)]
969 + {
970 + if let Some(ref mut shm) = self.shm {
971 + if payload.len() > max_request_payload_bytes as usize {
972 + self.client_note_request_capacity(payload.len() as u32);
973 + return Err(NipcError::Overflow);
974 + }
975 +
976 + let msg_len = HEADER_SIZE + payload.len();
977 + let msg = ensure_client_scratch(&mut self.send_buf, msg_len);
978 +
979 + hdr.magic = MAGIC_MSG;
980 + hdr.version = VERSION;
981 + hdr.header_len = protocol::HEADER_LEN;
982 + hdr.payload_len = payload.len() as u32;
983 +
984 + hdr.encode(&mut msg[..HEADER_SIZE]);
985 + if !payload.is_empty() {
986 + msg[HEADER_SIZE..].copy_from_slice(payload);
987 + }
988 +
989 + let send_result = shm.send(&msg);
990 + return match send_result {
991 + Ok(()) => Ok(()),
992 + Err(crate::transport::win_shm::WinShmError::MsgTooLarge) => {
993 + self.client_note_request_capacity(payload.len() as u32);
994 + Err(NipcError::Overflow)
995 + }
996 + Err(_) => Err(NipcError::Truncated),
997 + };
998 + }
999 + }
1000 +
1001 + // Baseline transport path
1002 + let send_result = {
1003 + let session = self.session.as_mut().ok_or(NipcError::Truncated)?;
1004 + session.send(hdr, payload)
1005 + };
1006 + match send_result {
1007 + Ok(()) => Ok(()),
1008 + #[cfg(unix)]
1009 + Err(crate::transport::posix::UdsError::LimitExceeded) => {
1010 + self.client_note_request_capacity(payload.len() as u32);
1011 + Err(NipcError::Overflow)
1012 + }
1013 + #[cfg(windows)]
1014 + Err(crate::transport::windows::NpError::LimitExceeded) => {
1015 + self.client_note_request_capacity(payload.len() as u32);
1016 + Err(NipcError::Overflow)
1017 + }
1018 + Err(_) => Err(NipcError::Truncated),
1019 + }
1020 + }
1021 +
1022 + fn transport_send_request_buf(
1023 + &mut self,
1024 + hdr: &mut Header,
1025 + req_len: usize,
1026 + ) -> Result<(), NipcError> {
1027 + let max_request_payload_bytes = self.session_max_request_payload_bytes();
1028 +
1029 + #[cfg(target_os = "linux")]
1030 + {
1031 + if let Some(ref mut shm) = self.shm {
1032 + if req_len > max_request_payload_bytes as usize {
1033 + self.client_note_request_capacity(req_len as u32);
1034 + return Err(NipcError::Overflow);
1035 + }
1036 +
1037 + let msg_len = HEADER_SIZE + req_len;
1038 + let msg = ensure_client_scratch(&mut self.send_buf, msg_len);
1039 +
1040 + hdr.magic = MAGIC_MSG;
1041 + hdr.version = VERSION;
1042 + hdr.header_len = protocol::HEADER_LEN;
1043 + hdr.payload_len = req_len as u32;
1044 +
1045 + hdr.encode(&mut msg[..HEADER_SIZE]);
1046 + if req_len > 0 {
1047 + msg[HEADER_SIZE..HEADER_SIZE + req_len]
1048 + .copy_from_slice(&self.request_buf[..req_len]);
1049 + }
1050 +
1051 + let send_result = shm.send(&msg[..msg_len]);
1052 + return match send_result {
1053 + Ok(()) => Ok(()),
1054 + Err(crate::transport::shm::ShmError::MsgTooLarge) => {
1055 + self.client_note_request_capacity(req_len as u32);
1056 + Err(NipcError::Overflow)
1057 + }
1058 + Err(_) => Err(NipcError::Truncated),
1059 + };
1060 + }
1061 + }
1062 +
1063 + #[cfg(windows)]
1064 + {
1065 + if let Some(ref mut shm) = self.shm {
1066 + if req_len > max_request_payload_bytes as usize {
1067 + self.client_note_request_capacity(req_len as u32);
1068 + return Err(NipcError::Overflow);
1069 + }
1070 +
1071 + let msg_len = HEADER_SIZE + req_len;
1072 + let msg = ensure_client_scratch(&mut self.send_buf, msg_len);
1073 +
1074 + hdr.magic = MAGIC_MSG;
1075 + hdr.version = VERSION;
1076 + hdr.header_len = protocol::HEADER_LEN;
1077 + hdr.payload_len = req_len as u32;
1078 +
1079 + hdr.encode(&mut msg[..HEADER_SIZE]);
1080 + if req_len > 0 {
1081 + msg[HEADER_SIZE..HEADER_SIZE + req_len]
1082 + .copy_from_slice(&self.request_buf[..req_len]);
1083 + }
1084 +
1085 + let send_result = shm.send(&msg[..msg_len]);
1086 + return match send_result {
1087 + Ok(()) => Ok(()),
1088 + Err(crate::transport::win_shm::WinShmError::MsgTooLarge) => {
1089 + self.client_note_request_capacity(req_len as u32);
1090 + Err(NipcError::Overflow)
1091 + }
1092 + Err(_) => Err(NipcError::Truncated),
1093 + };
1094 + }
1095 + }
1096 +
1097 + let send_result = {
1098 + let session = self.session.as_mut().ok_or(NipcError::Truncated)?;
1099 + session.send(hdr, &self.request_buf[..req_len])
1100 + };
1101 + match send_result {
1102 + Ok(()) => Ok(()),
1103 + #[cfg(unix)]
1104 + Err(crate::transport::posix::UdsError::LimitExceeded) => {
1105 + self.client_note_request_capacity(req_len as u32);
1106 + Err(NipcError::Overflow)
1107 + }
1108 + #[cfg(windows)]
1109 + Err(crate::transport::windows::NpError::LimitExceeded) => {
1110 + self.client_note_request_capacity(req_len as u32);
1111 + Err(NipcError::Overflow)
1112 + }
1113 + Err(_) => Err(NipcError::Truncated),
1114 + }
1115 + }
1116 +
1117 + /// Receive via the active transport. Returns (header, payload_view).
1118 + fn transport_receive(&mut self) -> Result<(Header, ClientResponseRef), NipcError> {
1119 + let needed = self.max_receive_message_bytes();
1120 + let scratch = ensure_client_scratch(&mut self.transport_buf, needed);
1121 +
1122 + // SHM path (POSIX or Windows)
1123 + #[cfg(target_os = "linux")]
1124 + {
1125 + if let Some(ref mut shm) = self.shm {
1126 + let mlen = shm
1127 + .receive(scratch, 30000)
1128 + .map_err(|_| NipcError::Truncated)?;
1129 +
1130 + if mlen < HEADER_SIZE {
1131 + return Err(NipcError::Truncated);
1132 + }
1133 +
1134 + let hdr = Header::decode(&scratch[..mlen])?;
1135 + return Ok((
1136 + hdr,
1137 + ClientResponseRef {
1138 + source: ClientResponseSource::TransportBuf,
1139 + len: mlen - HEADER_SIZE,
1140 + },
1141 + ));
1142 + }
1143 + }
1144 +
1145 + #[cfg(windows)]
1146 + {
1147 + if let Some(ref mut shm) = self.shm {
1148 + let mlen = shm
1149 + .receive(scratch, 30000)
1150 + .map_err(|_| NipcError::Truncated)?;
1151 +
1152 + if mlen < HEADER_SIZE {
1153 + return Err(NipcError::Truncated);
1154 + }
1155 +
1156 + let hdr = Header::decode(&scratch[..mlen])?;
1157 + return Ok((
1158 + hdr,
1159 + ClientResponseRef {
1160 + source: ClientResponseSource::TransportBuf,
1161 + len: mlen - HEADER_SIZE,
1162 + },
1163 + ));
1164 + }
1165 + }
1166 +
1167 + // Baseline transport: UDS on POSIX, Named Pipe on Windows
1168 + let session = self.session.as_mut().ok_or(NipcError::Truncated)?;
1169 +
1170 + #[cfg(unix)]
1171 + {
1172 + let scratch_payload_ptr = unsafe { scratch.as_ptr().add(HEADER_SIZE) };
1173 + let (hdr, payload) = session.receive(scratch).map_err(|_| NipcError::Truncated)?;
1174 + let source = if payload.as_ptr() == scratch_payload_ptr {
1175 + ClientResponseSource::TransportBuf
1176 + } else {
1177 + ClientResponseSource::SessionBuf
1178 + };
1179 + Ok((
1180 + hdr,
1181 + ClientResponseRef {
1182 + source,
1183 + len: payload.len(),
1184 + },
1185 + ))
1186 + }
1187 +
1188 + #[cfg(windows)]
1189 + {
1190 + let scratch_payload_ptr = unsafe { scratch.as_ptr().add(HEADER_SIZE) };
1191 + let (hdr, payload) = session.receive(scratch).map_err(|_| NipcError::Truncated)?;
1192 + let source = if payload.as_ptr() == scratch_payload_ptr {
1193 + ClientResponseSource::TransportBuf
1194 + } else {
1195 + ClientResponseSource::SessionBuf
1196 + };
1197 + Ok((
1198 + hdr,
1199 + ClientResponseRef {
1200 + source,
1201 + len: payload.len(),
1202 + },
1203 + ))
1204 + }
1205 + }
1206 +
1207 + fn response_payload(&self, response: ClientResponseRef) -> Result<&[u8], NipcError> {
1208 + match response.source {
1209 + ClientResponseSource::TransportBuf => {
1210 + let start = HEADER_SIZE;
1211 + let end = HEADER_SIZE + response.len;
1212 + if end > self.transport_buf.len() {
1213 + return Err(NipcError::Truncated);
1214 + }
1215 + Ok(&self.transport_buf[start..end])
1216 + }
1217 + ClientResponseSource::SessionBuf => {
1218 + #[cfg(unix)]
1219 + {
1220 + let session = self.session.as_ref().ok_or(NipcError::Truncated)?;
1221 + return Ok(session.received_payload(response.len));
1222 + }
1223 + #[cfg(windows)]
1224 + {
1225 + let session = self.session.as_ref().ok_or(NipcError::Truncated)?;
1226 + return Ok(session.received_payload(response.len));
1227 + }
1228 + #[allow(unreachable_code)]
1229 + Err(NipcError::Truncated)
1230 + }
1231 + }
1232 + }
1233 +
1234 + fn max_receive_message_bytes(&self) -> usize {
1235 + let mut max_payload = self.transport_config.max_response_payload_bytes as usize;
1236 + #[cfg(unix)]
1237 + if let Some(ref session) = self.session {
1238 + if session.max_response_payload_bytes > 0 {
1239 + max_payload = session.max_response_payload_bytes as usize;
1240 + }
1241 + }
1242 + #[cfg(windows)]
1243 + if let Some(ref session) = self.session {
1244 + if session.max_response_payload_bytes > 0 {
1245 + max_payload = session.max_response_payload_bytes as usize;
1246 + }
1247 + }
1248 + if max_payload == 0 {
1249 + max_payload = CACHE_RESPONSE_BUF_SIZE;
1250 + }
1251 + HEADER_SIZE + max_payload
1252 + }
1253 +}
1254 +
1255 +impl Drop for RawClient {
1256 + fn drop(&mut self) {
1257 + self.close();
1258 + }
1259 +}
1260 +
1261 +fn ensure_client_scratch(buf: &mut Vec<u8>, needed: usize) -> &mut [u8] {
1262 + if buf.len() < needed {
1263 + buf.resize(needed, 0);
1264 + }
1265 + &mut buf[..needed]
1266 +}
1267 +
1268 +#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1269 +enum ClientResponseSource {
1270 + TransportBuf,
1271 + SessionBuf,
1272 +}
1273 +
1274 +#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1275 +struct ClientResponseRef {
1276 + source: ClientResponseSource,
1277 + len: usize,
1278 +}
1279 +
1280 +fn dispatch_single_internal(
1281 + expected_method_code: u16,
1282 + handler: Option<&DispatchHandler>,
1283 + method_code: u16,
1284 + request: &[u8],
1285 + response_buf: &mut [u8],
1286 +) -> Result<usize, DispatchError> {
1287 + if method_code != expected_method_code {
1288 + return Err(DispatchError::HandlerFailed);
1289 + }
1290 +
1291 + match handler {
1292 + Some(dispatch) => match dispatch(request, response_buf) {
1293 + Ok(n) if n <= response_buf.len() => Ok(n),
1294 + Ok(_) => Err(DispatchError::Overflow),
1295 + Err(err) => Err(err),
1296 + },
1297 + None => Err(DispatchError::HandlerFailed),
1298 + }
1299 +}
1300 +
1301 +#[cfg(test)]
1302 +#[allow(dead_code)]
1303 +fn dispatch_single(
1304 + expected_method_code: u16,
1305 + handler: Option<&DispatchHandler>,
1306 + method_code: u16,
1307 + request: &[u8],
1308 + response_buf: &mut [u8],
1309 +) -> Result<usize, DispatchError> {
1310 + dispatch_single_internal(
1311 + expected_method_code,
1312 + handler,
1313 + method_code,
1314 + request,
1315 + response_buf,
1316 + )
1317 +}
1318 +
1319 +fn method_supported_internal(
1320 + expected_method_code: u16,
1321 + handler: Option<&DispatchHandler>,
1322 + method_code: u16,
1323 +) -> bool {
1324 + handler.is_some() && method_code == expected_method_code
1325 +}
1326 +
1327 +fn server_note_payload_capacity(target: &AtomicU32, payload_len: u32) {
1328 + let grown = next_power_of_2_u32(payload_len);
1329 + let mut current = target.load(Ordering::Relaxed);
1330 + while grown > current {
1331 + match target.compare_exchange_weak(current, grown, Ordering::Release, Ordering::Relaxed) {
1332 + Ok(_) => break,
1333 + Err(observed) => current = observed,
1334 + }
1335 + }
1336 +}
1337 +
1338 +// ---------------------------------------------------------------------------
1339 +// Managed server
1340 +// ---------------------------------------------------------------------------
1341 +
1342 +pub type IncrementHandler = Arc<dyn Fn(u64) -> Option<u64> + Send + Sync>;
1343 +pub type StringReverseHandler = Arc<dyn Fn(&str) -> Option<String> + Send + Sync>;
1344 +pub type SnapshotHandler =
1345 + Arc<dyn for<'a> Fn(&CgroupsRequest, &mut protocol::CgroupsBuilder<'a>) -> bool + Send + Sync>;
1346 +
1347 +#[derive(Debug, Clone, Copy, PartialEq, Eq)]
1348 +pub enum DispatchError {
1349 + BadEnvelope,
1350 + Overflow,
1351 + HandlerFailed,
1352 +}
1353 +
1354 +pub type DispatchHandler =
1355 + Arc<dyn Fn(&[u8], &mut [u8]) -> Result<usize, DispatchError> + Send + Sync>;
1356 +
1357 +pub fn increment_dispatch(handler: IncrementHandler) -> DispatchHandler {
1358 + Arc::new(move |request, response_buf| {
1359 + let value = increment_decode(request).map_err(|_| DispatchError::BadEnvelope)?;
1360 + let result = handler(value).ok_or(DispatchError::HandlerFailed)?;
1361 + let n = increment_encode(result, response_buf);
1362 + if n == 0 {
1363 + return Err(DispatchError::Overflow);
1364 + }
1365 + Ok(n)
1366 + })
1367 +}
1368 +
1369 +pub fn string_reverse_dispatch(handler: StringReverseHandler) -> DispatchHandler {
1370 + Arc::new(move |request, response_buf| {
1371 + let view = string_reverse_decode(request).map_err(|_| DispatchError::BadEnvelope)?;
1372 + let result = handler(view.as_str()).ok_or(DispatchError::HandlerFailed)?;
1373 + let n = string_reverse_encode(result.as_bytes(), response_buf);
1374 + if n == 0 {
1375 + return Err(DispatchError::Overflow);
1376 + }
1377 + Ok(n)
1378 + })
1379 +}
1380 +
1381 +pub fn snapshot_max_items(response_buf_size: usize, override_max_items: u32) -> u32 {
1382 + if override_max_items != 0 {
1383 + return override_max_items;
1384 + }
1385 + protocol::estimate_cgroups_max_items(response_buf_size)
1386 +}
1387 +
1388 +pub fn snapshot_dispatch(handler: SnapshotHandler, max_items: u32) -> DispatchHandler {
1389 + Arc::new(move |request, response_buf| {
1390 + let request = CgroupsRequest::decode(request).map_err(|_| DispatchError::BadEnvelope)?;
1391 + let item_budget = snapshot_max_items(response_buf.len(), max_items);
1392 + if item_budget == 0 {
1393 + return Err(DispatchError::Overflow);
1394 + }
1395 + let mut builder = protocol::CgroupsBuilder::new(response_buf, item_budget, 0, 0);
1396 + if !handler(&request, &mut builder) {
1397 + return Err(DispatchError::HandlerFailed);
1398 + }
1399 + let n = builder.finish();
1400 + if n == 0 {
1401 + return Err(DispatchError::Overflow);
1402 + }
1403 + Ok(n)
1404 + })
1405 +}
1406 +
1407 +/// L2 managed server. Typed request/response dispatcher.
1408 +///
1409 +/// Handles accept, spawns a thread per session (up to worker_count),
1410 +/// reads requests, dispatches to handler, sends responses.
1411 +pub struct ManagedServer {
1412 + run_dir: String,
1413 + service_name: String,
1414 + server_config: ServerConfig,
1415 + expected_method_code: u16,
1416 + handler: Option<DispatchHandler>,
1417 + running: Arc<AtomicBool>,
1418 + learned_request_payload_bytes: Arc<AtomicU32>,
1419 + learned_response_payload_bytes: Arc<AtomicU32>,
1420 + next_session_id: u64,
1421 + worker_count: usize,
1422 + /// Windows: stored listener handle so stop() can close it to unblock Accept.
1423 + #[cfg(windows)]
1424 + listener_handle: Arc<std::sync::Mutex<Option<usize>>>,
1425 +}
1426 +
1427 +impl ManagedServer {
1428 + /// Create a new managed server for a single service kind.
1429 + pub fn new(
1430 + run_dir: &str,
1431 + service_name: &str,
1432 + config: ServerConfig,
1433 + expected_method_code: u16,
1434 + handler: Option<DispatchHandler>,
1435 + ) -> Self {
1436 + Self::with_workers(
1437 + run_dir,
1438 + service_name,
1439 + config,
1440 + expected_method_code,
1441 + handler,
1442 + 8,
1443 + )
1444 + }
1445 +
1446 + /// Create a managed server with an explicit worker count.
1447 + pub fn with_workers(
1448 + run_dir: &str,
1449 + service_name: &str,
1450 + config: ServerConfig,
1451 + expected_method_code: u16,
1452 + handler: Option<DispatchHandler>,
1453 + worker_count: usize,
1454 + ) -> Self {
1455 + let learned_request = if config.max_request_payload_bytes != 0 {
1456 + config.max_request_payload_bytes
1457 + } else {
1458 + MAX_PAYLOAD_DEFAULT
1459 + };
1460 + let learned_response = if config.max_response_payload_bytes != 0 {
1461 + config.max_response_payload_bytes
1462 + } else {
1463 + MAX_PAYLOAD_DEFAULT
1464 + };
1465 +
1466 + ManagedServer {
1467 + run_dir: run_dir.to_string(),
1468 + service_name: service_name.to_string(),
1469 + server_config: config,
1470 + expected_method_code,
1471 + handler,
1472 + running: Arc::new(AtomicBool::new(false)),
1473 + learned_request_payload_bytes: Arc::new(AtomicU32::new(learned_request)),
1474 + learned_response_payload_bytes: Arc::new(AtomicU32::new(learned_response)),
1475 + next_session_id: 1,
1476 + worker_count: if worker_count < 1 { 1 } else { worker_count },
1477 + #[cfg(windows)]
1478 + listener_handle: Arc::new(std::sync::Mutex::new(None)),
1479 + }
1480 + }
1481 +
1482 + /// Run the acceptor loop. Blocking. Accepts clients, spawns a
1483 + /// thread per session (up to worker_count concurrent sessions).
1484 + ///
1485 + /// Returns when `stop()` is called or on fatal error.
1486 + #[cfg(unix)]
1487 + pub fn run(&mut self) -> Result<(), NipcError> {
1488 + #[cfg(target_os = "linux")]
1489 + crate::transport::shm::cleanup_stale(&self.run_dir, &self.service_name);
1490 +
1491 + let listener = UdsListener::bind(
1492 + &self.run_dir,
1493 + &self.service_name,
1494 + self.server_config.clone(),
1495 + )
1496 + .map_err(|_| NipcError::BadLayout)?;
1497 +
1498 + self.running.store(true, Ordering::Release);
1499 +
1500 + let mut session_threads: Vec<std::thread::JoinHandle<()>> = Vec::new();
1501 +
1502 + while self.running.load(Ordering::Acquire) {
1503 + let ready = poll_fd(listener.fd(), SERVER_POLL_TIMEOUT_MS as i32);
1504 + if ready < 0 {
1505 + break;
1506 + }
1507 + if ready == 0 {
1508 + // Reap finished threads periodically
1509 + session_threads.retain(|t| !t.is_finished());
1510 + continue;
1511 + }
1512 +
1513 + let (session_id, accept_cfg, precreated_shm, ready) = self.prepare_unix_accept();
1514 + if !ready {
1515 + std::thread::sleep(std::time::Duration::from_millis(10));
1516 + continue;
1517 + }
1518 +
1519 + let session = match listener.accept_with_config(session_id, accept_cfg) {
1520 + Ok(s) => s,
1521 + Err(_) => {
1522 + #[cfg(target_os = "linux")]
1523 + if let Some(mut shm) = precreated_shm {
1524 + shm.destroy();
1525 + }
1526 + if !self.running.load(Ordering::Acquire) {
1527 + break;
1528 + }
1529 + std::thread::sleep(std::time::Duration::from_millis(10));
1530 + continue;
1531 + }
1532 + };
1533 +
1534 + // Check worker count limit (non-blocking)
1535 + // Reap finished threads first
1536 + session_threads.retain(|t| !t.is_finished());
1537 + if session_threads.len() >= self.worker_count {
1538 + // At capacity: reject client
1539 + #[cfg(target_os = "linux")]
1540 + if let Some(mut shm) = precreated_shm {
1541 + shm.destroy();
1542 + }
1543 + drop(session);
1544 + continue;
1545 + }
1546 +
1547 + #[cfg(target_os = "linux")]
1548 + let shm = match self.finalize_unix_shm(&session, precreated_shm) {
1549 + Some(shm) => Some(shm),
1550 + None if session.selected_profile == PROFILE_SHM_HYBRID
1551 + || session.selected_profile == PROFILE_SHM_FUTEX =>
1552 + {
1553 + drop(session);
1554 + continue;
1555 + }
1556 + None => None,
1557 + };
1558 + #[cfg(not(target_os = "linux"))]
1559 + let shm: Option<()> = None;
1560 +
1561 + // Spawn a handler thread for this session
1562 + let expected_method_code = self.expected_method_code;
1563 + let handler = self.handler.clone();
1564 + let running = self.running.clone();
1565 + let learned_request_payload_bytes = self.learned_request_payload_bytes.clone();
1566 + let learned_response_payload_bytes = self.learned_response_payload_bytes.clone();
1567 +
1568 + let t = std::thread::spawn(move || {
1569 + handle_session_threaded(
1570 + session,
1571 + #[cfg(target_os = "linux")]
1572 + shm,
1573 + #[cfg(not(target_os = "linux"))]
1574 + shm,
1575 + expected_method_code,
1576 + handler,
1577 + running,
1578 + learned_request_payload_bytes,
1579 + learned_response_payload_bytes,
1580 + );
1581 + });
1582 + session_threads.push(t);
1583 + }
1584 +
1585 + // Wait for all active session threads
1586 + for t in session_threads {
1587 + let _ = t.join();
1588 + }
1589 +
1590 + Ok(())
1591 + }
1592 +
1593 + /// Windows: run the acceptor loop over Named Pipes.
1594 + #[cfg(windows)]
1595 + pub fn run(&mut self) -> Result<(), NipcError> {
1596 + // Win SHM cleanup is a no-op: kernel objects auto-clean on handle close.
1597 +
1598 + let mut listener = NpListener::bind(
1599 + &self.run_dir,
1600 + &self.service_name,
1601 + self.server_config.clone(),
1602 + )
1603 + .map_err(|_| NipcError::BadLayout)?;
1604 +
1605 + // Store listener handle so stop() can close it to unblock Accept
1606 + *self.listener_handle.lock().unwrap() = Some(listener.handle() as usize);
1607 +
1608 + self.running.store(true, Ordering::Release);
1609 +
1610 + let mut session_threads: Vec<std::thread::JoinHandle<()>> = Vec::new();
1611 +
1612 + while self.running.load(Ordering::Acquire) {
1613 + let (session_id, accept_cfg, prepared_shm, ready) = self.prepare_windows_accept();
1614 + if !ready {
1615 + std::thread::sleep(std::time::Duration::from_millis(10));
1616 + continue;
1617 + }
1618 +
1619 + let session = match listener.accept_with_config(session_id, accept_cfg) {
1620 + Ok(s) => s,
1621 + Err(_) => {
1622 + if let Some(mut prepared) = prepared_shm {
1623 + prepared.destroy_all();
1624 + }
1625 + if !self.running.load(Ordering::Acquire) {
1626 + break;
1627 + }
1628 + std::thread::sleep(std::time::Duration::from_millis(10));
1629 + continue;
1630 + }
1631 + };
1632 +
1633 + // Reap finished threads
1634 + session_threads.retain(|t| !t.is_finished());
1635 + if session_threads.len() >= self.worker_count {
1636 + if let Some(mut prepared) = prepared_shm {
1637 + prepared.destroy_all();
1638 + }
1639 + drop(session);
1640 + continue;
1641 + }
1642 +
1643 + let shm = match self.finalize_windows_shm(&session, prepared_shm) {
1644 + Some(shm) => Some(shm),
1645 + None if session.selected_profile == WIN_SHM_PROFILE_HYBRID
1646 + || session.selected_profile == WIN_SHM_PROFILE_BUSYWAIT =>
1647 + {
1648 + drop(session);
1649 + continue;
1650 + }
1651 + None => None,
1652 + };
1653 +
1654 + if shm.is_none()
1655 + && (session.selected_profile == WIN_SHM_PROFILE_HYBRID
1656 + || session.selected_profile == WIN_SHM_PROFILE_BUSYWAIT)
1657 + {
1658 + drop(session);
1659 + continue;
1660 + }
1661 +
1662 + let expected_method_code = self.expected_method_code;
1663 + let handler = self.handler.clone();
1664 + let running = self.running.clone();
1665 + let learned_request_payload_bytes = self.learned_request_payload_bytes.clone();
1666 + let learned_response_payload_bytes = self.learned_response_payload_bytes.clone();
1667 + let t = std::thread::spawn(move || {
1668 + handle_session_win_threaded(
1669 + session,
1670 + shm,
1671 + expected_method_code,
1672 + handler,
1673 + running,
1674 + learned_request_payload_bytes,
1675 + learned_response_payload_bytes,
1676 + );
1677 + });
1678 + session_threads.push(t);
1679 + }
1680 +
1681 + for t in session_threads {
1682 + let _ = t.join();
1683 + }
1684 +
1685 + Ok(())
1686 + }
1687 +
1688 + /// Signal shutdown. On Windows, also closes the listener pipe to
1689 + /// unblock ConnectNamedPipe in the accept loop.
1690 + pub fn stop(&self) {
1691 + self.running.store(false, Ordering::Release);
1692 +
1693 + #[cfg(windows)]
1694 + {
1695 + let mut guard = self.listener_handle.lock().unwrap();
1696 + if let Some(h) = guard.take() {
1697 + // Close the listener pipe to unblock ConnectNamedPipe
1698 + extern "system" {
1699 + fn CloseHandle(h: isize) -> i32;
1700 + }
1701 + unsafe {
1702 + CloseHandle(h as isize);
1703 + }
1704 + }
1705 + }
1706 + }
1707 +
1708 + /// Returns the internal running flag for diagnostics and test helpers.
1709 + ///
1710 + /// For reliable shutdown, call `stop()`. On Windows, flipping this flag
1711 + /// alone does not wake a blocking listener accept.
1712 + pub fn running_flag(&self) -> Arc<AtomicBool> {
1713 + self.running.clone()
1714 + }
1715 +
1716 + // ------------------------------------------------------------------
1717 + // Internal helpers
1718 + // ------------------------------------------------------------------
1719 +
1720 + #[cfg(target_os = "linux")]
1721 + fn prepare_unix_accept(&mut self) -> (u64, ServerConfig, Option<ShmContext>, bool) {
1722 + let session_id = self.next_session_id;
1723 + self.next_session_id += 1;
1724 +
1725 + let mut cfg = self.server_config.clone();
1726 + cfg.max_request_payload_bytes = self.learned_request_payload_bytes.load(Ordering::Acquire);
1727 + cfg.max_response_payload_bytes =
1728 + self.learned_response_payload_bytes.load(Ordering::Acquire);
1729 +
1730 + let shm_profiles = cfg.supported_profiles & (PROFILE_SHM_HYBRID | PROFILE_SHM_FUTEX);
1731 + if shm_profiles == 0 {
1732 + return (session_id, cfg, None, true);
1733 + }
1734 +
1735 + match ShmContext::server_create(
1736 + &self.run_dir,
1737 + &self.service_name,
1738 + session_id,
1739 + cfg.max_request_payload_bytes + HEADER_SIZE as u32,
1740 + cfg.max_response_payload_bytes + HEADER_SIZE as u32,
1741 + ) {
1742 + Ok(ctx) => (session_id, cfg, Some(ctx), true),
1743 + Err(_) => {
1744 + cfg.supported_profiles &= !(PROFILE_SHM_HYBRID | PROFILE_SHM_FUTEX);
1745 + cfg.preferred_profiles &= !(PROFILE_SHM_HYBRID | PROFILE_SHM_FUTEX);
1746 + (session_id, cfg.clone(), None, cfg.supported_profiles != 0)
1747 + }
1748 + }
1749 + }
1750 +
1751 + #[cfg(target_os = "linux")]
1752 + fn finalize_unix_shm(
1753 + &self,
1754 + session: &UdsSession,
1755 + mut shm: Option<ShmContext>,
1756 + ) -> Option<ShmContext> {
1757 + let profile = session.selected_profile;
1758 + if profile != PROFILE_SHM_HYBRID && profile != PROFILE_SHM_FUTEX {
1759 + if let Some(ref mut ctx) = shm {
1760 + ctx.destroy();
1761 + }
1762 + return None;
1763 + }
1764 + shm
1765 + }
1766 +
1767 + #[cfg(windows)]
1768 + fn prepare_windows_accept(&mut self) -> (u64, ServerConfig, Option<PreparedWinShm>, bool) {
1769 + let session_id = self.next_session_id;
1770 + self.next_session_id += 1;
1771 +
1772 + let mut cfg = self.server_config.clone();
1773 + cfg.max_request_payload_bytes = self.learned_request_payload_bytes.load(Ordering::Acquire);
1774 + cfg.max_response_payload_bytes =
1775 + self.learned_response_payload_bytes.load(Ordering::Acquire);
1776 +
1777 + let shm_profiles =
1778 + cfg.supported_profiles & (WIN_SHM_PROFILE_HYBRID | WIN_SHM_PROFILE_BUSYWAIT);
1779 + if shm_profiles == 0 {
1780 + return (session_id, cfg, None, true);
1781 + }
1782 +
1783 + let mut prepared = PreparedWinShm::default();
1784 + for profile in [WIN_SHM_PROFILE_HYBRID, WIN_SHM_PROFILE_BUSYWAIT] {
1785 + if cfg.supported_profiles & profile == 0 {
1786 + continue;
1787 + }
1788 +
1789 + match WinShmContext::server_create(
1790 + &self.run_dir,
1791 + &self.service_name,
1792 + self.server_config.auth_token,
1793 + session_id,
1794 + profile,
1795 + cfg.max_request_payload_bytes + HEADER_SIZE as u32,
1796 + cfg.max_response_payload_bytes + HEADER_SIZE as u32,
1797 + ) {
1798 + Ok(ctx) => prepared.insert(profile, ctx),
1799 + Err(_) => {
1800 + cfg.supported_profiles &= !profile;
1801 + cfg.preferred_profiles &= !profile;
1802 + }
1803 + }
1804 + }
1805 +
1806 + if cfg.supported_profiles == 0 {
1807 + prepared.destroy_all();
1808 + return (session_id, cfg, None, false);
1809 + }
1810 +
1811 + if prepared.is_empty() {
1812 + return (session_id, cfg, None, true);
1813 + }
1814 +
1815 + (session_id, cfg, Some(prepared), true)
1816 + }
1817 +
1818 + #[cfg(windows)]
1819 + fn finalize_windows_shm(
1820 + &self,
1821 + session: &NpSession,
1822 + mut prepared: Option<PreparedWinShm>,
1823 + ) -> Option<WinShmContext> {
1824 + let profile = session.selected_profile;
1825 + if profile != WIN_SHM_PROFILE_HYBRID && profile != WIN_SHM_PROFILE_BUSYWAIT {
1826 + if let Some(ref mut prepared) = prepared {
1827 + prepared.destroy_all();
1828 + }
1829 + return None;
1830 + }
1831 + let mut prepared = prepared?;
1832 + let selected = prepared.take(profile);
1833 + prepared.destroy_all();
1834 + selected
1835 + }
1836 +}
1837 +
1838 +#[cfg(windows)]
1839 +#[derive(Default)]
1840 +struct PreparedWinShm {
1841 + hybrid: Option<WinShmContext>,
1842 + busywait: Option<WinShmContext>,
1843 +}
1844 +
1845 +#[cfg(windows)]
1846 +impl PreparedWinShm {
1847 + fn insert(&mut self, profile: u32, ctx: WinShmContext) {
1848 + if profile == WIN_SHM_PROFILE_HYBRID {
1849 + self.hybrid = Some(ctx);
1850 + } else if profile == WIN_SHM_PROFILE_BUSYWAIT {
1851 + self.busywait = Some(ctx);
1852 + }
1853 + }
1854 +
1855 + fn take(&mut self, profile: u32) -> Option<WinShmContext> {
1856 + if profile == WIN_SHM_PROFILE_HYBRID {
1857 + self.hybrid.take()
1858 + } else if profile == WIN_SHM_PROFILE_BUSYWAIT {
1859 + self.busywait.take()
1860 + } else {
1861 + None
1862 + }
1863 + }
1864 +
1865 + fn destroy_all(&mut self) {
1866 + if let Some(mut ctx) = self.hybrid.take() {
1867 + ctx.destroy();
1868 + }
1869 + if let Some(mut ctx) = self.busywait.take() {
1870 + ctx.destroy();
1871 + }
1872 + }
1873 +
1874 + fn is_empty(&self) -> bool {
1875 + self.hybrid.is_none() && self.busywait.is_none()
1876 + }
1877 +}
1878 +
1879 +/// Windows: handle one client session over Named Pipe + optional Win SHM.
1880 +/// Standalone function for use in per-session threads.
1881 +#[cfg(windows)]
1882 +fn handle_session_win_threaded(
1883 + mut session: NpSession,
1884 + mut shm: Option<WinShmContext>,
1885 + expected_method_code: u16,
1886 + handler: Option<DispatchHandler>,
1887 + running: Arc<AtomicBool>,
1888 + learned_request_payload_bytes: Arc<AtomicU32>,
1889 + learned_response_payload_bytes: Arc<AtomicU32>,
1890 +) {
1891 + let mut recv_buf = vec![0u8; HEADER_SIZE + session.max_request_payload_bytes as usize];
1892 + let mut resp_buf = vec![0u8; session.max_response_payload_bytes as usize];
1893 + let mut item_resp_buf = vec![0u8; session.max_response_payload_bytes as usize];
1894 + let mut msg_buf = vec![0u8; HEADER_SIZE + session.max_response_payload_bytes as usize];
1895 +
1896 + while running.load(Ordering::Acquire) {
1897 + let (hdr, payload) = {
1898 + if let Some(ref mut shm_ctx) = shm {
1899 + match shm_ctx.receive(&mut recv_buf, SERVER_POLL_TIMEOUT_MS) {
1900 + Ok(mlen) => {
1901 + if mlen < HEADER_SIZE {
1902 + break;
1903 + }
1904 + let hdr = match Header::decode(&recv_buf[..mlen]) {
1905 + Ok(h) => h,
1906 + Err(_) => break,
1907 + };
1908 + let payload = &recv_buf[HEADER_SIZE..mlen];
1909 + (hdr, payload)
1910 + }
1911 + Err(crate::transport::win_shm::WinShmError::Timeout) => continue,
1912 + Err(_) => break,
1913 + }
1914 + } else {
1915 + // Named Pipe path
1916 + match session.wait_readable(SERVER_POLL_TIMEOUT_MS) {
1917 + Ok(true) => {}
1918 + Ok(false) => continue,
1919 + Err(_) => break,
1920 + }
1921 + match session.receive(&mut recv_buf) {
1922 + Ok((hdr, payload)) => (hdr, payload),
1923 + Err(_) => break,
1924 + }
1925 + }
1926 + };
1927 +
1928 + // Protocol violation: unexpected message kind terminates session
1929 + if hdr.kind != KIND_REQUEST {
1930 + break;
1931 + }
1932 +
1933 + if payload.len() <= u32::MAX as usize {
1934 + server_note_payload_capacity(&learned_request_payload_bytes, payload.len() as u32);
1935 + }
1936 +
1937 + if !method_supported_internal(expected_method_code, handler.as_ref(), hdr.code) {
1938 + let mut resp_hdr = Header {
1939 + kind: KIND_RESPONSE,
1940 + code: hdr.code,
1941 + message_id: hdr.message_id,
1942 + transport_status: protocol::STATUS_UNSUPPORTED,
1943 + item_count: 1,
1944 + ..Header::default()
1945 + };
1946 +
1947 + if let Some(ref mut shm_ctx) = shm {
1948 + let msg = ensure_client_scratch(&mut msg_buf, HEADER_SIZE);
1949 + resp_hdr.magic = MAGIC_MSG;
1950 + resp_hdr.version = VERSION;
1951 + resp_hdr.header_len = protocol::HEADER_LEN;
1952 + resp_hdr.payload_len = 0;
1953 + resp_hdr.encode(&mut msg[..HEADER_SIZE]);
1954 + if shm_ctx.send(&msg[..HEADER_SIZE]).is_err() {
1955 + break;
1956 + }
1957 + } else if session.send(&mut resp_hdr, &[]).is_err() {
1958 + break;
1959 + }
1960 + continue;
1961 + }
1962 +
1963 + // Dispatch: single-item or batch
1964 + let is_batch = (hdr.flags & FLAG_BATCH) != 0 && hdr.item_count >= 1;
1965 + let response_len;
1966 + let dispatch_result = if !is_batch {
1967 + dispatch_single_internal(
1968 + expected_method_code,
1969 + handler.as_ref(),
1970 + hdr.code,
1971 + payload,
1972 + &mut resp_buf,
1973 + )
1974 + } else {
1975 + let mut bb = BatchBuilder::new(&mut resp_buf, hdr.item_count);
1976 + let mut batch_result = Ok(0usize);
1977 +
1978 + for i in 0..hdr.item_count {
1979 + let (item_data, _item_len) = match batch_item_get(payload, hdr.item_count, i) {
1980 + Ok(v) => v,
1981 + Err(_) => {
1982 + batch_result = Err(DispatchError::BadEnvelope);
1983 + break;
1984 + }
1985 + };
1986 + let item_len = match dispatch_single_internal(
1987 + expected_method_code,
1988 + handler.as_ref(),
1989 + hdr.code,
1990 + item_data,
1991 + &mut item_resp_buf,
1992 + ) {
1993 + Ok(n) => n,
1994 + Err(err) => {
1995 + batch_result = Err(err);
1996 + break;
1997 + }
1998 + };
1999 + if bb.add(&item_resp_buf[..item_len]).is_err() {
2000 + batch_result = Err(DispatchError::Overflow);
2001 + break;
2002 + }
2003 + }
2004 + if batch_result.is_ok() {
2005 + let (n, _) = bb.finish();
2006 + batch_result = Ok(n);
2007 + }
2008 + batch_result
2009 + };
2010 +
2011 + let mut resp_hdr = Header {
2012 + kind: KIND_RESPONSE,
2013 + code: hdr.code,
2014 + message_id: hdr.message_id,
2015 + ..Header::default()
2016 + };
2017 +
2018 + match dispatch_result {
2019 + Ok(n) => {
2020 + response_len = n;
2021 + if response_len <= u32::MAX as usize {
2022 + server_note_payload_capacity(
2023 + &learned_response_payload_bytes,
2024 + response_len as u32,
2025 + );
2026 + }
2027 + resp_hdr.transport_status = STATUS_OK;
2028 + if is_batch {
2029 + resp_hdr.flags = FLAG_BATCH;
2030 + resp_hdr.item_count = hdr.item_count;
2031 + } else {
2032 + resp_hdr.flags = 0;
2033 + resp_hdr.item_count = 1;
2034 + }
2035 + }
2036 + Err(DispatchError::Overflow) => {
2037 + let current = session.max_response_payload_bytes;
2038 + if current >= u32::MAX / 2 {
2039 + server_note_payload_capacity(&learned_response_payload_bytes, u32::MAX);
2040 + } else {
2041 + server_note_payload_capacity(&learned_response_payload_bytes, current * 2);
2042 + }
2043 + resp_hdr.transport_status = STATUS_LIMIT_EXCEEDED;
2044 + resp_hdr.item_count = 1;
2045 + resp_hdr.flags = 0;
2046 + response_len = 0;
2047 + }
2048 + Err(DispatchError::BadEnvelope) => {
2049 + resp_hdr.transport_status = STATUS_BAD_ENVELOPE;
2050 + resp_hdr.item_count = 1;
2051 + resp_hdr.flags = 0;
2052 + response_len = 0;
2053 + }
2054 + Err(DispatchError::HandlerFailed) => {
2055 + resp_hdr.transport_status = STATUS_INTERNAL_ERROR;
2056 + resp_hdr.item_count = 1;
2057 + resp_hdr.flags = 0;
2058 + response_len = 0;
2059 + }
2060 + }
2061 +
2062 + if let Some(ref mut shm_ctx) = shm {
2063 + let msg_len = HEADER_SIZE + response_len;
2064 + let msg = ensure_client_scratch(&mut msg_buf, msg_len);
2065 +
2066 + resp_hdr.magic = MAGIC_MSG;
2067 + resp_hdr.version = VERSION;
2068 + resp_hdr.header_len = protocol::HEADER_LEN;
2069 + resp_hdr.payload_len = response_len as u32;
2070 +
2071 + resp_hdr.encode(&mut msg[..HEADER_SIZE]);
2072 + if response_len > 0 {
2073 + msg[HEADER_SIZE..].copy_from_slice(&resp_buf[..response_len]);
2074 + }
2075 +
2076 + if shm_ctx.send(msg).is_err() {
2077 + break;
2078 + }
2079 + if resp_hdr.transport_status == STATUS_LIMIT_EXCEEDED {
2080 + break;
2081 + }
2082 + continue;
2083 + }
2084 +
2085 + if session
2086 + .send(&mut resp_hdr, &resp_buf[..response_len])
2087 + .is_err()
2088 + {
2089 + break;
2090 + }
2091 + if resp_hdr.transport_status == STATUS_LIMIT_EXCEEDED {
2092 + break;
2093 + }
2094 + }
2095 +
2096 + if let Some(mut shm_ctx) = shm {
2097 + shm_ctx.destroy();
2098 + }
2099 + session.close();
2100 +}
2101 +
2102 +/// POSIX: Handle one client session in its own thread.
2103 +#[cfg(unix)]
2104 +fn handle_session_threaded(
2105 + mut session: UdsSession,
2106 + #[cfg(target_os = "linux")] mut shm: Option<ShmContext>,
2107 + #[cfg(not(target_os = "linux"))] _shm: Option<()>,
2108 + expected_method_code: u16,
2109 + handler: Option<DispatchHandler>,
2110 + running: Arc<AtomicBool>,
2111 + learned_request_payload_bytes: Arc<AtomicU32>,
2112 + learned_response_payload_bytes: Arc<AtomicU32>,
2113 +) {
2114 + let mut recv_buf = vec![0u8; HEADER_SIZE + session.max_request_payload_bytes as usize];
2115 + let mut resp_buf = vec![0u8; session.max_response_payload_bytes as usize];
2116 + let mut item_resp_buf = vec![0u8; session.max_response_payload_bytes as usize];
2117 + let mut msg_buf = vec![0u8; HEADER_SIZE + session.max_response_payload_bytes as usize];
2118 +
2119 + while running.load(Ordering::Acquire) {
2120 + // Receive request via the active transport
2121 + let (hdr, payload) = {
2122 + #[cfg(target_os = "linux")]
2123 + {
2124 + if let Some(ref mut shm_ctx) = shm {
2125 + match shm_ctx.receive(&mut recv_buf, SERVER_POLL_TIMEOUT_MS) {
2126 + Ok(mlen) => {
2127 + if mlen < HEADER_SIZE {
2128 + break;
2129 + }
2130 + let hdr = match Header::decode(&recv_buf[..mlen]) {
2131 + Ok(h) => h,
2132 + Err(_) => break,
2133 + };
2134 + let payload = &recv_buf[HEADER_SIZE..mlen];
2135 + (hdr, payload)
2136 + }
2137 + Err(crate::transport::shm::ShmError::Timeout) => continue,
2138 + Err(_) => break,
2139 + }
2140 + } else {
2141 + // UDS path with poll
2142 + let ready = poll_fd(session.fd(), SERVER_POLL_TIMEOUT_MS as i32);
2143 + if ready < 0 {
2144 + break;
2145 + }
2146 + if ready == 0 {
2147 + continue;
2148 + }
2149 +
2150 + match session.receive(&mut recv_buf) {
2151 + Ok((hdr, payload)) => (hdr, payload),
2152 + Err(_) => break,
2153 + }
2154 + }
2155 + }
2156 +
2157 + #[cfg(not(target_os = "linux"))]
2158 + {
2159 + let ready = poll_fd(session.fd(), SERVER_POLL_TIMEOUT_MS as i32);
2160 + if ready < 0 {
2161 + break;
2162 + }
2163 + if ready == 0 {
2164 + continue;
2165 + }
2166 +
2167 + match session.receive(&mut recv_buf) {
2168 + Ok((hdr, payload)) => (hdr, payload),
2169 + Err(_) => break,
2170 + }
2171 + }
2172 + };
2173 +
2174 + // Protocol violation: unexpected message kind terminates session
2175 + if hdr.kind != KIND_REQUEST {
2176 + break;
2177 + }
2178 +
2179 + if payload.len() <= u32::MAX as usize {
2180 + server_note_payload_capacity(&learned_request_payload_bytes, payload.len() as u32);
2181 + }
2182 +
2183 + if !method_supported_internal(expected_method_code, handler.as_ref(), hdr.code) {
2184 + let mut resp_hdr = Header {
2185 + kind: KIND_RESPONSE,
2186 + code: hdr.code,
2187 + message_id: hdr.message_id,
2188 + transport_status: protocol::STATUS_UNSUPPORTED,
2189 + item_count: 1,
2190 + ..Header::default()
2191 + };
2192 +
2193 + #[cfg(target_os = "linux")]
2194 + {
2195 + if let Some(ref mut shm_ctx) = shm {
2196 + let msg = ensure_client_scratch(&mut msg_buf, HEADER_SIZE);
2197 + resp_hdr.magic = MAGIC_MSG;
2198 + resp_hdr.version = VERSION;
2199 + resp_hdr.header_len = protocol::HEADER_LEN;
2200 + resp_hdr.payload_len = 0;
2201 + resp_hdr.encode(&mut msg[..HEADER_SIZE]);
2202 + if shm_ctx.send(&msg[..HEADER_SIZE]).is_err() {
2203 + break;
2204 + }
2205 + continue;
2206 + }
2207 + }
2208 +
2209 + if session.send(&mut resp_hdr, &[]).is_err() {
2210 + break;
2211 + }
2212 + continue;
2213 + }
2214 +
2215 + // Dispatch: single-item or batch
2216 + let is_batch = (hdr.flags & FLAG_BATCH) != 0 && hdr.item_count >= 1;
2217 + let response_len;
2218 + let dispatch_result = if !is_batch {
2219 + dispatch_single_internal(
2220 + expected_method_code,
2221 + handler.as_ref(),
2222 + hdr.code,
2223 + payload,
2224 + &mut resp_buf,
2225 + )
2226 + } else {
2227 + let mut bb = BatchBuilder::new(&mut resp_buf, hdr.item_count);
2228 + let mut batch_result = Ok(0usize);
2229 +
2230 + for i in 0..hdr.item_count {
2231 + let (item_data, _item_len) = match batch_item_get(payload, hdr.item_count, i) {
2232 + Ok(v) => v,
2233 + Err(_) => {
2234 + batch_result = Err(DispatchError::BadEnvelope);
2235 + break;
2236 + }
2237 + };
2238 + let item_len = match dispatch_single_internal(
2239 + expected_method_code,
2240 + handler.as_ref(),
2241 + hdr.code,
2242 + item_data,
2243 + &mut item_resp_buf,
2244 + ) {
2245 + Ok(n) => n,
2246 + Err(err) => {
2247 + batch_result = Err(err);
2248 + break;
2249 + }
2250 + };
2251 + if bb.add(&item_resp_buf[..item_len]).is_err() {
2252 + batch_result = Err(DispatchError::Overflow);
2253 + break;
2254 + }
2255 + }
2256 + if batch_result.is_ok() {
2257 + let (n, _) = bb.finish();
2258 + batch_result = Ok(n);
2259 + }
2260 + batch_result
2261 + };
2262 +
2263 + // Build response header
2264 + let mut resp_hdr = Header {
2265 + kind: KIND_RESPONSE,
2266 + code: hdr.code,
2267 + message_id: hdr.message_id,
2268 + ..Header::default()
2269 + };
2270 +
2271 + match dispatch_result {
2272 + Ok(n) => {
2273 + response_len = n;
2274 + if response_len <= u32::MAX as usize {
2275 + server_note_payload_capacity(
2276 + &learned_response_payload_bytes,
2277 + response_len as u32,
2278 + );
2279 + }
2280 + resp_hdr.transport_status = STATUS_OK;
2281 + if is_batch {
2282 + resp_hdr.flags = FLAG_BATCH;
2283 + resp_hdr.item_count = hdr.item_count;
2284 + } else {
2285 + resp_hdr.flags = 0;
2286 + resp_hdr.item_count = 1;
2287 + }
2288 + }
2289 + Err(DispatchError::Overflow) => {
2290 + let current = session.max_response_payload_bytes;
2291 + if current >= u32::MAX / 2 {
2292 + server_note_payload_capacity(&learned_response_payload_bytes, u32::MAX);
2293 + } else {
2294 + server_note_payload_capacity(&learned_response_payload_bytes, current * 2);
2295 + }
2296 + resp_hdr.transport_status = STATUS_LIMIT_EXCEEDED;
2297 + resp_hdr.item_count = 1;
2298 + resp_hdr.flags = 0;
2299 + response_len = 0;
2300 + }
2301 + Err(DispatchError::BadEnvelope) => {
2302 + resp_hdr.transport_status = STATUS_BAD_ENVELOPE;
2303 + resp_hdr.item_count = 1;
2304 + resp_hdr.flags = 0;
2305 + response_len = 0;
2306 + }
2307 + Err(DispatchError::HandlerFailed) => {
2308 + resp_hdr.transport_status = STATUS_INTERNAL_ERROR;
2309 + resp_hdr.item_count = 1;
2310 + resp_hdr.flags = 0;
2311 + response_len = 0;
2312 + }
2313 + }
2314 +
2315 + // Send response via the active transport
2316 + #[cfg(target_os = "linux")]
2317 + {
2318 + if let Some(ref mut shm_ctx) = shm {
2319 + let msg_len = HEADER_SIZE + response_len;
2320 + let msg = ensure_client_scratch(&mut msg_buf, msg_len);
2321 +
2322 + resp_hdr.magic = MAGIC_MSG;
2323 + resp_hdr.version = VERSION;
2324 + resp_hdr.header_len = protocol::HEADER_LEN;
2325 + resp_hdr.payload_len = response_len as u32;
2326 +
2327 + resp_hdr.encode(&mut msg[..HEADER_SIZE]);
2328 + if response_len > 0 {
2329 + msg[HEADER_SIZE..].copy_from_slice(&resp_buf[..response_len]);
2330 + }
2331 +
2332 + if shm_ctx.send(msg).is_err() {
2333 + break;
2334 + }
2335 + if resp_hdr.transport_status == STATUS_LIMIT_EXCEEDED {
2336 + break;
2337 + }
2338 + continue;
2339 + }
2340 + }
2341 +
2342 + // UDS path
2343 + if session
2344 + .send(&mut resp_hdr, &resp_buf[..response_len])
2345 + .is_err()
2346 + {
2347 + break;
2348 + }
2349 + if resp_hdr.transport_status == STATUS_LIMIT_EXCEEDED {
2350 + break;
2351 + }
2352 + }
2353 +
2354 + // Cleanup
2355 + #[cfg(target_os = "linux")]
2356 + {
2357 + if let Some(mut shm_ctx) = shm {
2358 + shm_ctx.destroy();
2359 + }
2360 + }
2361 + drop(session);
2362 +}
2363 +
2364 +// ---------------------------------------------------------------------------
2365 +// Internal: poll helper
2366 +// ---------------------------------------------------------------------------
2367 +
2368 +/// Poll a file descriptor for readability with a timeout in milliseconds.
2369 +/// Returns: 1 = data ready, 0 = timeout, -1 = error/hangup.
2370 +#[cfg(unix)]
2371 +fn poll_fd(fd: i32, timeout_ms: i32) -> i32 {
2372 + let mut pfd = libc::pollfd {
2373 + fd,
2374 + events: libc::POLLIN,
2375 + revents: 0,
2376 + };
2377 +
2378 + let ret = unsafe { libc::poll(&mut pfd, 1, timeout_ms) };
2379 +
2380 + if ret < 0 {
2381 + let errno = unsafe { *libc::__errno_location() };
2382 + if errno == libc::EINTR {
2383 + return 0;
2384 + }
2385 + return -1;
2386 + }
2387 +
2388 + if ret == 0 {
2389 + return 0;
2390 + }
2391 +
2392 + if pfd.revents & (libc::POLLERR | libc::POLLHUP | libc::POLLNVAL) != 0 {
2393 + return -1;
2394 + }
2395 +
2396 + if pfd.revents & libc::POLLIN != 0 {
2397 + return 1;
2398 + }
2399 +
2400 + 0
2401 +}
2402 +
2403 +// ---------------------------------------------------------------------------
2404 +// L3: Client-side cgroups snapshot cache
2405 +// ---------------------------------------------------------------------------
2406 +
2407 +/// Cached copy of a single cgroup item. Owns its strings.
2408 +/// Built from ephemeral L2 views during cache construction.
2409 +#[derive(Debug, Clone)]
2410 +pub struct CgroupsCacheItem {
2411 + pub hash: u32,
2412 + pub options: u32,
2413 + pub enabled: u32,
2414 + pub name: String,
2415 + pub path: String,
2416 +}
2417 +
2418 +/// L3 cache status snapshot (for diagnostics, not hot path).
2419 +#[derive(Debug, Clone, Copy, PartialEq, Eq)]
2420 +pub struct CgroupsCacheStatus {
2421 + pub populated: bool,
2422 + pub item_count: u32,
2423 + pub systemd_enabled: u32,
2424 + pub generation: u64,
2425 + pub refresh_success_count: u32,
2426 + pub refresh_failure_count: u32,
2427 + pub connection_state: ClientState,
2428 + /// Monotonic milliseconds of last successful refresh (0 if never).
2429 + pub last_refresh_ts: u64,
2430 +}
2431 +
2432 +/// Default response buffer size for L3 cache refresh.
2433 +const CACHE_RESPONSE_BUF_SIZE: usize = 65536;
2434 +
2435 +#[derive(Debug, Clone, Copy, Default)]
2436 +struct CgroupsHashBucket {
2437 + index: u32,
2438 + used: bool,
2439 +}
2440 +
2441 +fn cache_hash_name(name: &str) -> u32 {
2442 + let mut h: u32 = 5381;
2443 + for b in name.as_bytes() {
2444 + h = ((h << 5).wrapping_add(h)).wrapping_add(*b as u32);
2445 + }
2446 + h
2447 +}
2448 +
2449 +/// L3 client-side cgroups snapshot cache.
2450 +///
2451 +/// Wraps an L2 client and maintains a local owned copy of the most
2452 +/// recent successful snapshot. Lookup by hash+name is O(1) via HashMap.
2453 +///
2454 +/// On refresh failure, the previous cache is preserved. The cache
2455 +/// is empty only if no successful refresh has ever occurred.
2456 +pub struct CgroupsCache {
2457 + client: RawClient,
2458 + items: Vec<CgroupsCacheItem>,
2459 + /// Open-addressing hash table: (hash ^ djb2(name)) -> index into items vec
2460 + buckets: Vec<CgroupsHashBucket>,
2461 + systemd_enabled: u32,
2462 + generation: u64,
2463 + populated: bool,
2464 + refresh_success_count: u32,
2465 + refresh_failure_count: u32,
2466 + /// Monotonic reference point for timestamp calculation
2467 + epoch: std::time::Instant,
2468 + /// Monotonic ms of last successful refresh (0 if never)
2469 + pub last_refresh_ts: u64,
2470 +}
2471 +
2472 +impl CgroupsCache {
2473 + /// Create a new L3 cache. Creates the underlying L2 client context.
2474 + /// Does NOT connect. Does NOT require the server to be running.
2475 + /// Cache starts empty (populated == false).
2476 + pub fn new(run_dir: &str, service_name: &str, config: ClientConfig) -> Self {
2477 + CgroupsCache {
2478 + client: RawClient::new_snapshot(run_dir, service_name, config),
2479 + items: Vec::new(),
2480 + buckets: Vec::new(),
2481 + systemd_enabled: 0,
2482 + generation: 0,
2483 + populated: false,
2484 + refresh_success_count: 0,
2485 + refresh_failure_count: 0,
2486 + epoch: std::time::Instant::now(),
2487 + last_refresh_ts: 0,
2488 + }
2489 + }
2490 +
2491 + /// Refresh the cache. Drives the L2 client (connect/reconnect as
2492 + /// needed) and requests a fresh snapshot. On success, rebuilds the
2493 + /// local cache. On failure, preserves the previous cache.
2494 + ///
2495 + /// Returns true if the cache was updated.
2496 + pub fn refresh(&mut self) -> bool {
2497 + // Drive L2 connection lifecycle
2498 + self.client.refresh();
2499 +
2500 + // Attempt snapshot call
2501 + match self.client.call_snapshot() {
2502 + Ok(view) => {
2503 + // Build new cache from snapshot view
2504 + let mut new_items = Vec::with_capacity(view.item_count as usize);
2505 + for i in 0..view.item_count {
2506 + match view.item(i) {
2507 + Ok(iv) => {
2508 + let name = match iv.name.as_str() {
2509 + Ok(s) => s.to_string(),
2510 + Err(_) => {
2511 + // Non-UTF8 name: use lossy conversion
2512 + String::from_utf8_lossy(iv.name.as_bytes()).into_owned()
2513 + }
2514 + };
2515 + let path = match iv.path.as_str() {
2516 + Ok(s) => s.to_string(),
2517 + Err(_) => String::from_utf8_lossy(iv.path.as_bytes()).into_owned(),
2518 + };
2519 + new_items.push(CgroupsCacheItem {
2520 + hash: iv.hash,
2521 + options: iv.options,
2522 + enabled: iv.enabled,
2523 + name,
2524 + path,
2525 + });
2526 + }
2527 + Err(_) => {
2528 + // Malformed item: abort, preserve old cache
2529 + self.refresh_failure_count += 1;
2530 + return false;
2531 + }
2532 + }
2533 + }
2534 +
2535 + // Rebuild open-addressing lookup table.
2536 + let mut buckets = Vec::new();
2537 + if !new_items.is_empty() {
2538 + let bcount = next_power_of_2_u32((new_items.len() as u32) * 2) as usize;
2539 + buckets.resize(bcount, CgroupsHashBucket::default());
2540 + let mask = (bcount - 1) as u32;
2541 + for (i, item) in new_items.iter().enumerate() {
2542 + let mut slot = (item.hash ^ cache_hash_name(&item.name)) & mask;
2543 + while buckets[slot as usize].used {
2544 + slot = (slot + 1) & mask;
2545 + }
2546 + buckets[slot as usize] = CgroupsHashBucket {
2547 + index: i as u32,
2548 + used: true,
2549 + };
2550 + }
2551 + }
2552 +
2553 + // Replace old cache
2554 + self.items = new_items;
2555 + self.buckets = buckets;
2556 + self.systemd_enabled = view.systemd_enabled;
2557 + self.generation = view.generation;
2558 + self.populated = true;
2559 + self.refresh_success_count += 1;
2560 + self.last_refresh_ts = self.epoch.elapsed().as_millis() as u64;
2561 + true
2562 + }
2563 + Err(_) => {
2564 + // Refresh failed: preserve previous cache
2565 + self.refresh_failure_count += 1;
2566 + false
2567 + }
2568 + }
2569 + }
2570 +
2571 + /// Returns true if at least one successful refresh has occurred.
2572 + /// Cheap cached boolean. No I/O, no syscalls.
2573 + ///
2574 + /// Note: ready means "has cached data", not "is connected."
2575 + #[inline]
2576 + pub fn ready(&self) -> bool {
2577 + self.populated
2578 + }
2579 +
2580 + /// Look up a cached item by hash + name. O(1) via open-addressing hash
2581 + /// table. No I/O.
2582 + pub fn lookup(&self, hash: u32, name: &str) -> Option<&CgroupsCacheItem> {
2583 + if !self.populated {
2584 + return None;
2585 + }
2586 + if !self.buckets.is_empty() {
2587 + let mask = (self.buckets.len() - 1) as u32;
2588 + let mut slot = (hash ^ cache_hash_name(name)) & mask;
2589 + while self.buckets[slot as usize].used {
2590 + let item = &self.items[self.buckets[slot as usize].index as usize];
2591 + if item.hash == hash && item.name == name {
2592 + return Some(item);
2593 + }
2594 + slot = (slot + 1) & mask;
2595 + }
2596 + return None;
2597 + }
2598 +
2599 + self.items
2600 + .iter()
2601 + .find(|item| item.hash == hash && item.name == name)
2602 + }
2603 +
2604 + /// Fill a status snapshot for diagnostics.
2605 + pub fn status(&self) -> CgroupsCacheStatus {
2606 + CgroupsCacheStatus {
2607 + populated: self.populated,
2608 + item_count: self.items.len() as u32,
2609 + systemd_enabled: self.systemd_enabled,
2610 + generation: self.generation,
2611 + refresh_success_count: self.refresh_success_count,
2612 + refresh_failure_count: self.refresh_failure_count,
2613 + connection_state: self.client.state,
2614 + last_refresh_ts: self.last_refresh_ts,
2615 + }
2616 + }
2617 +
2618 + /// Close the cache: free all cached items, close the L2 client.
2619 + pub fn close(&mut self) {
2620 + self.items.clear();
2621 + self.buckets.clear();
2622 + self.populated = false;
2623 + self.client.close();
2624 + }
2625 +}
2626 +
2627 +impl Drop for CgroupsCache {
2628 + fn drop(&mut self) {
2629 + self.close();
2630 + }
2631 +}
2632 +
2633 +// ---------------------------------------------------------------------------
2634 +// Tests
2635 +// ---------------------------------------------------------------------------
2636 +
2637 +#[cfg(all(test, unix))]
2638 +#[path = "raw_unix_tests.rs"]
2639 +mod tests;
2640 +
2641 +#[cfg(all(test, windows))]
2642 +#[path = "raw_windows_tests.rs"]
2643 +mod windows_tests;
src/crates/netipc/src/service/raw_unix_tests.rs new
+4165
@@ -0,0 +1,4165 @@
1 +use super::*;
2 +#[cfg(target_os = "linux")]
3 +use crate::protocol::PROFILE_SHM_FUTEX;
4 +use crate::protocol::{increment_encode, BatchBuilder, CgroupsBuilder, PROFILE_BASELINE};
5 +use std::os::fd::RawFd;
6 +use std::os::unix::ffi::OsStrExt;
7 +use std::path::PathBuf;
8 +use std::thread;
9 +use std::time::Duration;
10 +
11 +const TEST_RUN_DIR: &str = "/tmp/nipc_svc_rust_test";
12 +const AUTH_TOKEN: u64 = 0xDEADBEEFCAFEBABE;
13 +const RESPONSE_BUF_SIZE: usize = 65536;
14 +static RAW_SERVICE_COUNTER: std::sync::atomic::AtomicU64 = std::sync::atomic::AtomicU64::new(0);
15 +
16 +fn ensure_run_dir() {
17 + let _ = std::fs::create_dir_all(TEST_RUN_DIR);
18 +}
19 +
20 +fn cleanup_all(service: &str) {
21 + let _ = std::fs::remove_file(format!("{TEST_RUN_DIR}/{service}.sock"));
22 + #[cfg(target_os = "linux")]
23 + crate::transport::shm::cleanup_stale(TEST_RUN_DIR, service);
24 +}
25 +
26 +fn socket_path(service: &str) -> PathBuf {
27 + PathBuf::from(format!("{TEST_RUN_DIR}/{service}.sock"))
28 +}
29 +
30 +fn unique_service(prefix: &str) -> String {
31 + format!(
32 + "{}_{}_{}",
33 + prefix,
34 + std::process::id(),
35 + RAW_SERVICE_COUNTER.fetch_add(1, std::sync::atomic::Ordering::Relaxed) + 1
36 + )
37 +}
38 +
39 +fn wait_for_listener_bind(service: &str) {
40 + let sock = socket_path(service);
41 + for _ in 0..2000 {
42 + if sock.exists() {
43 + return;
44 + }
45 + thread::sleep(Duration::from_micros(500));
46 + }
47 +
48 + panic!("listener did not bind for service {service}");
49 +}
50 +
51 +fn server_config() -> ServerConfig {
52 + ServerConfig {
53 + supported_profiles: PROFILE_BASELINE,
54 + max_request_payload_bytes: 4096,
55 + max_request_batch_items: 1,
56 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
57 + max_response_batch_items: 1,
58 + auth_token: AUTH_TOKEN,
59 + backlog: 4,
60 + ..ServerConfig::default()
61 + }
62 +}
63 +
64 +fn client_config() -> ClientConfig {
65 + ClientConfig {
66 + supported_profiles: PROFILE_BASELINE,
67 + max_request_payload_bytes: 4096,
68 + max_request_batch_items: 1,
69 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
70 + max_response_batch_items: 1,
71 + auth_token: AUTH_TOKEN,
72 + ..ClientConfig::default()
73 + }
74 +}
75 +
76 +#[cfg(target_os = "linux")]
77 +fn shm_server_config() -> ServerConfig {
78 + ServerConfig {
79 + supported_profiles: PROFILE_BASELINE | PROFILE_SHM_FUTEX,
80 + preferred_profiles: PROFILE_SHM_FUTEX,
81 + max_request_payload_bytes: 4096,
82 + max_request_batch_items: 16,
83 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
84 + max_response_batch_items: 16,
85 + auth_token: AUTH_TOKEN,
86 + backlog: 4,
87 + ..ServerConfig::default()
88 + }
89 +}
90 +
91 +#[cfg(target_os = "linux")]
92 +fn shm_client_config() -> ClientConfig {
93 + ClientConfig {
94 + supported_profiles: PROFILE_BASELINE | PROFILE_SHM_FUTEX,
95 + preferred_profiles: PROFILE_SHM_FUTEX,
96 + max_request_payload_bytes: 4096,
97 + max_request_batch_items: 16,
98 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
99 + max_response_batch_items: 16,
100 + auth_token: AUTH_TOKEN,
101 + ..ClientConfig::default()
102 + }
103 +}
104 +
105 +fn batch_server_config() -> ServerConfig {
106 + ServerConfig {
107 + supported_profiles: PROFILE_BASELINE,
108 + max_request_payload_bytes: 4096,
109 + max_request_batch_items: 16,
110 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
111 + max_response_batch_items: 16,
112 + auth_token: AUTH_TOKEN,
113 + backlog: 4,
114 + ..ServerConfig::default()
115 + }
116 +}
117 +
118 +fn batch_client_config() -> ClientConfig {
119 + ClientConfig {
120 + supported_profiles: PROFILE_BASELINE,
121 + max_request_payload_bytes: 4096,
122 + max_request_batch_items: 16,
123 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
124 + max_response_batch_items: 16,
125 + auth_token: AUTH_TOKEN,
126 + ..ClientConfig::default()
127 + }
128 +}
129 +
130 +fn snapshot_client(service: &str, config: ClientConfig) -> RawClient {
131 + RawClient::new_snapshot(TEST_RUN_DIR, service, config)
132 +}
133 +
134 +fn increment_client(service: &str, config: ClientConfig) -> RawClient {
135 + RawClient::new_increment(TEST_RUN_DIR, service, config)
136 +}
137 +
138 +fn string_reverse_client(service: &str, config: ClientConfig) -> RawClient {
139 + RawClient::new_string_reverse(TEST_RUN_DIR, service, config)
140 +}
141 +
142 +fn connect_ready(client: &mut RawClient) {
143 + for _ in 0..200 {
144 + client.refresh();
145 + if client.ready() {
146 + return;
147 + }
148 + thread::sleep(Duration::from_millis(10));
149 + }
150 +
151 + panic!("client did not reach READY state");
152 +}
153 +
154 +fn fill_test_cgroups_snapshot(builder: &mut CgroupsBuilder<'_>) -> bool {
155 + let items = [
156 + (
157 + 1001u32,
158 + 0u32,
159 + 1u32,
160 + b"docker-abc123" as &[u8],
161 + b"/sys/fs/cgroup/docker/abc123" as &[u8],
162 + ),
163 + (2002, 0, 1, b"k8s-pod-xyz", b"/sys/fs/cgroup/kubepods/xyz"),
164 + (
165 + 3003,
166 + 0,
167 + 0,
168 + b"systemd-user",
169 + b"/sys/fs/cgroup/user.slice/user-1000",
170 + ),
171 + ];
172 +
173 + for (hash, options, enabled, name, path) in &items {
174 + if builder.add(*hash, *options, *enabled, name, path).is_err() {
175 + return false;
176 + }
177 + }
178 +
179 + true
180 +}
181 +
182 +fn panic_payload_to_string(payload: Box<dyn std::any::Any + Send>) -> String {
183 + match payload.downcast::<String>() {
184 + Ok(msg) => *msg,
185 + Err(payload) => match payload.downcast::<&'static str>() {
186 + Ok(msg) => (*msg).to_string(),
187 + Err(_) => "<non-string panic>".to_string(),
188 + },
189 + }
190 +}
191 +
192 +fn send_raw_packet(fd: i32, data: &[u8]) {
193 + let sent = unsafe { libc::send(fd, data.as_ptr() as *const libc::c_void, data.len(), 0) };
194 + assert_eq!(
195 + sent,
196 + data.len() as isize,
197 + "raw send failed: {:?}",
198 + std::io::Error::last_os_error()
199 + );
200 +}
201 +
202 +fn build_increment_request_message(message_id: u64, value: u64) -> Vec<u8> {
203 + let mut payload = [0u8; INCREMENT_PAYLOAD_SIZE];
204 + let payload_len = increment_encode(value, &mut payload);
205 + assert_eq!(
206 + payload_len, INCREMENT_PAYLOAD_SIZE,
207 + "increment_encode should fit the fixed-size request buffer"
208 + );
209 +
210 + let hdr = Header {
211 + magic: MAGIC_MSG,
212 + version: VERSION,
213 + header_len: protocol::HEADER_LEN,
214 + kind: KIND_REQUEST,
215 + code: METHOD_INCREMENT,
216 + flags: 0,
217 + payload_len: payload_len as u32,
218 + item_count: 1,
219 + message_id,
220 + transport_status: STATUS_OK,
221 + };
222 +
223 + let mut msg = vec![0u8; HEADER_SIZE + payload_len];
224 + hdr.encode(&mut msg[..HEADER_SIZE]);
225 + msg[HEADER_SIZE..].copy_from_slice(&payload[..payload_len]);
226 + msg
227 +}
228 +
229 +fn verify_increment_service_ok(service: &str, config: ClientConfig) {
230 + let mut verify = increment_client(service, config);
231 + connect_ready(&mut verify);
232 + assert_eq!(
233 + verify.call_increment(1).expect("verification call"),
234 + 2,
235 + "server should remain healthy after rejecting the malformed session"
236 + );
237 + verify.close();
238 +}
239 +
240 +fn test_cgroups_snapshot_handler() -> SnapshotHandler {
241 + Arc::new(|req, builder| {
242 + if req.layout_version != 1 || req.flags != 0 {
243 + return false;
244 + }
245 + builder.set_header(1, 42);
246 + fill_test_cgroups_snapshot(builder)
247 + })
248 +}
249 +
250 +fn test_cgroups_dispatch() -> DispatchHandler {
251 + snapshot_dispatch(test_cgroups_snapshot_handler(), 3)
252 +}
253 +
254 +fn increment_handler() -> IncrementHandler {
255 + Arc::new(|value| Some(value + 1))
256 +}
257 +
258 +fn increment_dispatch_handler() -> DispatchHandler {
259 + increment_dispatch(increment_handler())
260 +}
261 +
262 +fn string_reverse_handler() -> StringReverseHandler {
263 + Arc::new(|s| Some(s.chars().rev().collect()))
264 +}
265 +
266 +fn string_reverse_dispatch_handler() -> DispatchHandler {
267 + string_reverse_dispatch(string_reverse_handler())
268 +}
269 +
270 +struct TestServer {
271 + stop_flag: Arc<AtomicBool>,
272 + thread: Option<thread::JoinHandle<()>>,
273 +}
274 +
275 +impl TestServer {
276 + fn start(service: &str, expected_method_code: u16, handler: Option<DispatchHandler>) -> Self {
277 + Self::start_with(service, server_config(), expected_method_code, handler, 8)
278 + }
279 +
280 + #[cfg(target_os = "linux")]
281 + fn start_shm(
282 + service: &str,
283 + expected_method_code: u16,
284 + handler: Option<DispatchHandler>,
285 + ) -> Self {
286 + Self::start_with(
287 + service,
288 + shm_server_config(),
289 + expected_method_code,
290 + handler,
291 + 8,
292 + )
293 + }
294 +
295 + fn start_with_workers(
296 + service: &str,
297 + expected_method_code: u16,
298 + handler: Option<DispatchHandler>,
299 + worker_count: usize,
300 + ) -> Self {
301 + Self::start_with(
302 + service,
303 + server_config(),
304 + expected_method_code,
305 + handler,
306 + worker_count,
307 + )
308 + }
309 +
310 + fn start_with(
311 + service: &str,
312 + config: ServerConfig,
313 + expected_method_code: u16,
314 + handler: Option<DispatchHandler>,
315 + worker_count: usize,
316 + ) -> Self {
317 + ensure_run_dir();
318 + cleanup_all(service);
319 +
320 + let svc = service.to_string();
321 + let ready_flag = Arc::new(AtomicBool::new(false));
322 + let ready_clone = ready_flag.clone();
323 +
324 + let mut server = ManagedServer::with_workers(
325 + TEST_RUN_DIR,
326 + &svc,
327 + config,
328 + expected_method_code,
329 + handler,
330 + worker_count,
331 + );
332 + let stop_flag = server.running_flag();
333 +
334 + let thread = thread::spawn(move || {
335 + // We need to signal readiness after bind but before accept loop.
336 + // The run() method binds internally, so we signal immediately
337 + // after it starts (it blocks on accept).
338 + ready_clone.store(true, Ordering::Release);
339 + let _ = server.run();
340 + });
341 +
342 + // Wait for server to be ready
343 + for _ in 0..2000 {
344 + if ready_flag.load(Ordering::Acquire) {
345 + break;
346 + }
347 + thread::sleep(Duration::from_micros(500));
348 + }
349 + wait_for_listener_bind(service);
350 +
351 + TestServer {
352 + stop_flag,
353 + thread: Some(thread),
354 + }
355 + }
356 +
357 + fn start_with_resp_size(
358 + service: &str,
359 + expected_method_code: u16,
360 + handler: Option<DispatchHandler>,
361 + resp_buf_size: usize,
362 + ) -> Self {
363 + ensure_run_dir();
364 + cleanup_all(service);
365 +
366 + let svc = service.to_string();
367 + let ready_flag = Arc::new(AtomicBool::new(false));
368 + let ready_clone = ready_flag.clone();
369 +
370 + let mut scfg = server_config();
371 + scfg.max_response_payload_bytes = resp_buf_size as u32;
372 +
373 + let mut server =
374 + ManagedServer::new(TEST_RUN_DIR, &svc, scfg, expected_method_code, handler);
375 + let stop_flag = server.running_flag();
376 +
377 + let thread = thread::spawn(move || {
378 + ready_clone.store(true, Ordering::Release);
379 + let _ = server.run();
380 + });
381 +
382 + for _ in 0..2000 {
383 + if ready_flag.load(Ordering::Acquire) {
384 + break;
385 + }
386 + thread::sleep(Duration::from_micros(500));
387 + }
388 + wait_for_listener_bind(service);
389 +
390 + TestServer {
391 + stop_flag,
392 + thread: Some(thread),
393 + }
394 + }
395 +
396 + fn stop(&mut self) {
397 + self.stop_flag.store(false, Ordering::Release);
398 + if let Some(t) = self.thread.take() {
399 + let _ = t.join();
400 + }
401 + }
402 +}
403 +
404 +impl Drop for TestServer {
405 + fn drop(&mut self) {
406 + self.stop();
407 + }
408 +}
409 +
410 +struct RawSessionServer {
411 + thread: Option<thread::JoinHandle<Result<(), String>>>,
412 +}
413 +
414 +struct RawHelloAckServer {
415 + thread: Option<thread::JoinHandle<Result<(), String>>>,
416 +}
417 +
418 +fn raw_listener_fd_for_service(service: &str) -> RawFd {
419 + let path = socket_path(service);
420 + let path_bytes = path.as_os_str().as_bytes();
421 + assert!(
422 + path_bytes.len() < std::mem::size_of::<libc::sockaddr_un>() - 2,
423 + "socket path too long"
424 + );
425 +
426 + let fd = unsafe { libc::socket(libc::AF_UNIX, libc::SOCK_SEQPACKET, 0) };
427 + assert!(
428 + fd >= 0,
429 + "socket failed: {}",
430 + std::io::Error::last_os_error()
431 + );
432 +
433 + let mut addr: libc::sockaddr_un = unsafe { std::mem::zeroed() };
434 + addr.sun_family = libc::AF_UNIX as libc::sa_family_t;
435 + for (idx, byte) in path_bytes.iter().enumerate() {
436 + addr.sun_path[idx] = *byte as libc::c_char;
437 + }
438 +
439 + let rc = unsafe {
440 + libc::bind(
441 + fd,
442 + &addr as *const libc::sockaddr_un as *const libc::sockaddr,
443 + std::mem::size_of::<libc::sockaddr_un>() as libc::socklen_t,
444 + )
445 + };
446 + assert_eq!(rc, 0, "bind failed: {}", std::io::Error::last_os_error());
447 +
448 + let rc = unsafe { libc::listen(fd, 4) };
449 + assert_eq!(rc, 0, "listen failed: {}", std::io::Error::last_os_error());
450 + fd
451 +}
452 +
453 +fn hello_ack_packet_with_version(version: u16, status: u16, layout_version: u16) -> Vec<u8> {
454 + let ack = crate::protocol::HelloAck {
455 + layout_version,
456 + flags: 0,
457 + server_supported_profiles: crate::protocol::PROFILE_BASELINE,
458 + intersection_profiles: crate::protocol::PROFILE_BASELINE,
459 + selected_profile: crate::protocol::PROFILE_BASELINE,
460 + agreed_max_request_payload_bytes: crate::protocol::MAX_PAYLOAD_DEFAULT,
461 + agreed_max_request_batch_items: 1,
462 + agreed_max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
463 + agreed_max_response_batch_items: 1,
464 + agreed_packet_size: 0,
465 + session_id: 77,
466 + };
467 +
468 + let mut payload = vec![0u8; 48];
469 + let payload_len = ack.encode(&mut payload);
470 + payload.truncate(payload_len);
471 +
472 + let hdr = crate::protocol::Header {
473 + magic: crate::protocol::MAGIC_MSG,
474 + version,
475 + header_len: crate::protocol::HEADER_LEN,
476 + kind: crate::protocol::KIND_CONTROL,
477 + flags: 0,
478 + code: crate::protocol::CODE_HELLO_ACK,
479 + transport_status: status,
480 + payload_len: payload.len() as u32,
481 + item_count: 1,
482 + message_id: 0,
483 + };
484 +
485 + let mut pkt = vec![0u8; crate::protocol::HEADER_SIZE + payload.len()];
486 + hdr.encode(&mut pkt[..crate::protocol::HEADER_SIZE]);
487 + pkt[crate::protocol::HEADER_SIZE..].copy_from_slice(&payload);
488 + pkt
489 +}
490 +
491 +fn start_raw_hello_ack_server(service: &str, packet: Vec<u8>) -> RawHelloAckServer {
492 + ensure_run_dir();
493 + cleanup_all(service);
494 +
495 + let svc = service.to_string();
496 + let thread = thread::spawn(move || {
497 + let fd = raw_listener_fd_for_service(&svc);
498 + let client_fd = unsafe { libc::accept(fd, std::ptr::null_mut(), std::ptr::null_mut()) };
499 + if client_fd < 0 {
500 + unsafe { libc::close(fd) };
501 + return Err(format!("accept: {}", std::io::Error::last_os_error()));
502 + }
503 +
504 + let mut hello_buf = [0u8; crate::protocol::HEADER_SIZE + 128];
505 + let n = unsafe {
506 + libc::recv(
507 + client_fd,
508 + hello_buf.as_mut_ptr() as *mut libc::c_void,
509 + hello_buf.len(),
510 + 0,
511 + )
512 + };
513 + if n < 0 {
514 + unsafe {
515 + libc::close(client_fd);
516 + libc::close(fd);
517 + }
518 + return Err(format!("recv: {}", std::io::Error::last_os_error()));
519 + }
520 +
521 + let wrote = unsafe {
522 + libc::send(
523 + client_fd,
524 + packet.as_ptr() as *const libc::c_void,
525 + packet.len(),
526 + 0,
527 + )
528 + };
529 + unsafe {
530 + libc::close(client_fd);
531 + libc::close(fd);
532 + }
533 + if wrote != packet.len() as isize {
534 + return Err(format!(
535 + "send short write: wrote {wrote}, want {}",
536 + packet.len()
537 + ));
538 + }
539 +
540 + Ok(())
541 + });
542 +
543 + thread::sleep(Duration::from_millis(50));
544 + RawHelloAckServer {
545 + thread: Some(thread),
546 + }
547 +}
548 +
549 +impl RawHelloAckServer {
550 + fn wait(&mut self) {
551 + if let Some(thread) = self.thread.take() {
552 + match thread.join() {
553 + Ok(Ok(())) => {}
554 + Ok(Err(err)) => panic!("raw hello-ack server failed: {err}"),
555 + Err(_) => panic!("raw hello-ack server panicked"),
556 + }
557 + }
558 + }
559 +}
560 +
561 +fn start_raw_session_server<F>(service: &str, cfg: ServerConfig, handler: F) -> RawSessionServer
562 +where
563 + F: FnOnce(&mut UdsSession, Header, &[u8]) -> Result<(), String> + Send + 'static,
564 +{
565 + ensure_run_dir();
566 + cleanup_all(service);
567 +
568 + let svc = service.to_string();
569 + let ready = Arc::new(AtomicBool::new(false));
570 + let ready_clone = ready.clone();
571 +
572 + let thread = thread::spawn(move || {
573 + let listener =
574 + UdsListener::bind(TEST_RUN_DIR, &svc, cfg).map_err(|e| format!("bind: {e}"))?;
575 + ready_clone.store(true, Ordering::Release);
576 + let mut session = listener.accept().map_err(|e| format!("accept: {e}"))?;
577 +
578 + let (hdr, payload) = {
579 + let mut recv_buf = vec![0u8; RESPONSE_BUF_SIZE];
580 + let (hdr, payload) = session
581 + .receive(&mut recv_buf)
582 + .map_err(|e| format!("receive: {e}"))?;
583 + (hdr, payload.to_vec())
584 + };
585 +
586 + handler(&mut session, hdr, &payload)
587 + });
588 +
589 + for _ in 0..2000 {
590 + if ready.load(Ordering::Acquire) {
591 + break;
592 + }
593 + thread::sleep(Duration::from_micros(500));
594 + }
595 + thread::sleep(Duration::from_millis(50));
596 +
597 + RawSessionServer {
598 + thread: Some(thread),
599 + }
600 +}
601 +
602 +impl RawSessionServer {
603 + fn wait(&mut self) {
604 + if let Some(thread) = self.thread.take() {
605 + match thread.join() {
606 + Ok(Ok(())) => {}
607 + Ok(Err(err)) => panic!("raw unix session server failed: {err}"),
608 + Err(_) => panic!("raw unix session server panicked"),
609 + }
610 + }
611 + }
612 +}
613 +
614 +#[cfg(target_os = "linux")]
615 +struct RawShmSessionServer {
616 + thread: Option<thread::JoinHandle<Result<(), String>>>,
617 +}
618 +
619 +#[cfg(target_os = "linux")]
620 +fn start_raw_shm_session_server<F>(
621 + service: &str,
622 + cfg: ServerConfig,
623 + handler: F,
624 +) -> RawShmSessionServer
625 +where
626 + F: FnOnce(&mut ShmContext, Header, &[u8]) -> Result<(), String> + Send + 'static,
627 +{
628 + ensure_run_dir();
629 + cleanup_all(service);
630 +
631 + let svc = service.to_string();
632 + let ready = Arc::new(AtomicBool::new(false));
633 + let ready_clone = ready.clone();
634 +
635 + let thread = thread::spawn(move || {
636 + let listener =
637 + UdsListener::bind(TEST_RUN_DIR, &svc, cfg).map_err(|e| format!("bind: {e}"))?;
638 + ready_clone.store(true, Ordering::Release);
639 + let session = listener.accept().map_err(|e| format!("accept: {e}"))?;
640 + if session.selected_profile != PROFILE_SHM_FUTEX {
641 + return Err(format!("unexpected profile {}", session.selected_profile));
642 + }
643 +
644 + let mut shm = ShmContext::server_create(
645 + TEST_RUN_DIR,
646 + &svc,
647 + session.session_id,
648 + session.max_request_payload_bytes + HEADER_SIZE as u32,
649 + session.max_response_payload_bytes + HEADER_SIZE as u32,
650 + )
651 + .map_err(|e| format!("server_create: {e}"))?;
652 +
653 + let mut recv_buf = vec![0u8; session.max_request_payload_bytes as usize + HEADER_SIZE];
654 + let mlen = shm
655 + .receive(&mut recv_buf, 5000)
656 + .map_err(|e| format!("shm receive: {e}"))?;
657 + if mlen < HEADER_SIZE {
658 + return Err(format!("request too short: {mlen}"));
659 + }
660 +
661 + let hdr = Header::decode(&recv_buf[..mlen]).map_err(|e| format!("decode: {e:?}"))?;
662 + let payload = recv_buf[HEADER_SIZE..mlen].to_vec();
663 + handler(&mut shm, hdr, &payload)
664 + });
665 +
666 + for _ in 0..2000 {
667 + if ready.load(Ordering::Acquire) {
668 + break;
669 + }
670 + thread::sleep(Duration::from_micros(500));
671 + }
672 + thread::sleep(Duration::from_millis(50));
673 +
674 + RawShmSessionServer {
675 + thread: Some(thread),
676 + }
677 +}
678 +
679 +#[cfg(target_os = "linux")]
680 +impl RawShmSessionServer {
681 + fn wait(&mut self) {
682 + if let Some(thread) = self.thread.take() {
683 + match thread.join() {
684 + Ok(Ok(())) => {}
685 + Ok(Err(err)) => panic!("raw shm session server failed: {err}"),
686 + Err(_) => panic!("raw shm session server panicked"),
687 + }
688 + }
689 + }
690 +}
691 +
692 +#[cfg(target_os = "linux")]
693 +fn encode_raw_message(hdr: &Header, payload: &[u8]) -> Vec<u8> {
694 + let mut msg = vec![0u8; HEADER_SIZE + payload.len()];
695 + hdr.encode(&mut msg[..HEADER_SIZE]);
696 + if !payload.is_empty() {
697 + msg[HEADER_SIZE..].copy_from_slice(payload);
698 + }
699 + msg
700 +}
701 +
702 +#[test]
703 +fn test_client_lifecycle() {
704 + let svc = "rs_svc_lifecycle";
705 + ensure_run_dir();
706 + cleanup_all(svc);
707 +
708 + // Init without server running
709 + let mut client = snapshot_client(svc, client_config());
710 + assert_eq!(client.state, ClientState::Disconnected);
711 + assert!(!client.ready());
712 +
713 + // Refresh without server -> NOT_FOUND
714 + let changed = client.refresh();
715 + assert!(changed);
716 + assert_eq!(client.state, ClientState::NotFound);
717 +
718 + // Start server
719 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
720 +
721 + // Refresh -> READY
722 + let changed = client.refresh();
723 + assert!(changed);
724 + assert_eq!(client.state, ClientState::Ready);
725 + assert!(client.ready());
726 +
727 + // Status reporting
728 + let status = client.status();
729 + assert_eq!(status.connect_count, 1);
730 + assert_eq!(status.reconnect_count, 0);
731 +
732 + // Close
733 + client.close();
734 + assert_eq!(client.state, ClientState::Disconnected);
735 + assert!(!client.ready());
736 +
737 + server.stop();
738 + cleanup_all(svc);
739 +}
740 +
741 +#[test]
742 +fn test_fill_test_cgroups_snapshot_small_builder_returns_false() {
743 + let mut buf = [0u8; 64];
744 + let mut builder = CgroupsBuilder::new(&mut buf, 3, 0, 0);
745 + assert!(
746 + !fill_test_cgroups_snapshot(&mut builder),
747 + "small builder should reject the synthetic snapshot"
748 + );
749 +}
750 +
751 +#[test]
752 +fn test_snapshot_handler_rejects_bad_request_metadata() {
753 + let on_snapshot = test_cgroups_snapshot_handler();
754 + let req = CgroupsRequest {
755 + layout_version: 2,
756 + flags: 0,
757 + };
758 + let mut response_payload = [0u8; RESPONSE_BUF_SIZE];
759 + let mut builder = CgroupsBuilder::new(&mut response_payload, 4, 1, 99);
760 + assert!(
761 + !on_snapshot(&req, &mut builder),
762 + "snapshot handler should reject bad request metadata"
763 + );
764 +}
765 +
766 +#[test]
767 +fn test_cgroups_call() {
768 + let svc = "rs_svc_cgroups";
769 + ensure_run_dir();
770 + cleanup_all(svc);
771 +
772 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
773 +
774 + let mut client = snapshot_client(svc, client_config());
775 + client.refresh();
776 + assert!(client.ready());
777 +
778 + let view = client.call_snapshot().expect("call should succeed");
779 +
780 + assert_eq!(view.item_count, 3);
781 + assert_eq!(view.systemd_enabled, 1);
782 + assert_eq!(view.generation, 42);
783 +
784 + // Verify first item
785 + let item0 = view.item(0).expect("item 0");
786 + assert_eq!(item0.hash, 1001);
787 + assert_eq!(item0.enabled, 1);
788 + assert_eq!(item0.name.as_bytes(), b"docker-abc123");
789 + assert_eq!(item0.path.as_bytes(), b"/sys/fs/cgroup/docker/abc123");
790 +
791 + // Verify third item
792 + let item2 = view.item(2).expect("item 2");
793 + assert_eq!(item2.hash, 3003);
794 + assert_eq!(item2.enabled, 0);
795 + assert_eq!(item2.name.as_bytes(), b"systemd-user");
796 +
797 + // Verify stats
798 + let status = client.status();
799 + assert_eq!(status.call_count, 1);
800 + assert_eq!(status.error_count, 0);
801 +
802 + client.close();
803 + server.stop();
804 + cleanup_all(svc);
805 +}
806 +
807 +#[cfg(target_os = "linux")]
808 +#[test]
809 +fn test_cgroups_call_shm() {
810 + let svc = "rs_svc_cgroups_shm";
811 + ensure_run_dir();
812 + cleanup_all(svc);
813 +
814 + let mut server =
815 + TestServer::start_shm(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
816 +
817 + let mut client = snapshot_client(svc, shm_client_config());
818 + connect_ready(&mut client);
819 + assert!(client.shm.is_some(), "expected SHM to be negotiated");
820 + assert_eq!(
821 + client.session.as_ref().map(|s| s.selected_profile),
822 + Some(PROFILE_SHM_FUTEX)
823 + );
824 +
825 + let view = client.call_snapshot().expect("snapshot over SHM");
826 + assert_eq!(view.item_count, 3);
827 + assert_eq!(view.generation, 42);
828 + assert_eq!(view.item(0).expect("item 0").hash, 1001);
829 +
830 + client.close();
831 + server.stop();
832 + cleanup_all(svc);
833 +}
834 +
835 +#[cfg(target_os = "linux")]
836 +#[test]
837 +fn test_client_call_string_reverse_shm_success() {
838 + let svc = "rs_svc_strrev_shm";
839 + ensure_run_dir();
840 + cleanup_all(svc);
841 +
842 + let mut server = TestServer::start_shm(
843 + svc,
844 + METHOD_STRING_REVERSE,
845 + Some(string_reverse_dispatch_handler()),
846 + );
847 +
848 + let mut client = string_reverse_client(svc, shm_client_config());
849 + connect_ready(&mut client);
850 + assert!(client.shm.is_some(), "expected SHM to be negotiated");
851 +
852 + let result = client
853 + .call_string_reverse("hello")
854 + .expect("string reverse over SHM");
855 + assert_eq!(result.as_str(), "olleh");
856 +
857 + client.close();
858 + server.stop();
859 + cleanup_all(svc);
860 +}
861 +
862 +#[cfg(target_os = "linux")]
863 +#[test]
864 +fn test_increment_batch_shm() {
865 + let svc = "rs_pp_batch_shm";
866 + ensure_run_dir();
867 + cleanup_all(svc);
868 +
869 + let mut server =
870 + TestServer::start_shm(svc, METHOD_INCREMENT, Some(increment_dispatch_handler()));
871 +
872 + let mut client = increment_client(svc, shm_client_config());
873 + connect_ready(&mut client);
874 + assert!(client.shm.is_some(), "expected SHM to be negotiated");
875 +
876 + let values = vec![10u64, 20, 30, 40];
877 + let results = client
878 + .call_increment_batch(&values)
879 + .expect("batch over SHM");
880 + assert_eq!(results, vec![11, 21, 31, 41]);
881 +
882 + client.close();
883 + server.stop();
884 + cleanup_all(svc);
885 +}
886 +
887 +#[cfg(target_os = "linux")]
888 +#[test]
889 +fn test_refresh_shm_attach_failure_falls_back_to_baseline() {
890 + let svc = "rs_svc_shm_attach_fail";
891 + ensure_run_dir();
892 + cleanup_all(svc);
893 +
894 + let ready = Arc::new(AtomicBool::new(false));
895 + let ready_clone = ready.clone();
896 + let svc_clone = svc.to_string();
897 +
898 + let server_thread = thread::spawn(move || {
899 + let listener = UdsListener::bind(TEST_RUN_DIR, &svc_clone, shm_server_config())
900 + .expect("bind shm listener");
901 + ready_clone.store(true, Ordering::Release);
902 + let mut first = listener.accept().expect("accept negotiated shm session");
903 + assert_eq!(first.selected_profile, PROFILE_SHM_FUTEX);
904 +
905 + let mut recv_buf = vec![0u8; RESPONSE_BUF_SIZE];
906 + let first_result = first.receive(&mut recv_buf);
907 + assert!(
908 + first_result.is_err(),
909 + "first SHM-selected session should disconnect after attach failure"
910 + );
911 +
912 + let mut second = listener.accept().expect("accept baseline fallback session");
913 + assert_eq!(second.selected_profile, PROFILE_BASELINE);
914 +
915 + let second_result = second.receive(&mut recv_buf);
916 + assert!(
917 + second_result.is_err(),
918 + "second baseline session should close cleanly when client closes"
919 + );
920 + });
921 +
922 + while !ready.load(Ordering::Acquire) {
923 + thread::sleep(Duration::from_millis(1));
924 + }
925 +
926 + let mut client = snapshot_client(svc, shm_client_config());
927 + assert!(
928 + client.refresh(),
929 + "refresh should transition to READY via baseline fallback"
930 + );
931 + assert_eq!(client.state, ClientState::Ready);
932 + assert!(client.ready());
933 + assert!(
934 + client.session.is_some(),
935 + "attach fallback should end with a live baseline session"
936 + );
937 + assert!(
938 + client.shm.is_none(),
939 + "fallback session must not retain SHM state"
940 + );
941 + assert_eq!(
942 + client.session.as_ref().map(|s| s.selected_profile),
943 + Some(PROFILE_BASELINE)
944 + );
945 + assert_eq!(client.transport_config.supported_profiles, PROFILE_BASELINE);
946 + assert_eq!(client.transport_config.preferred_profiles, 0);
947 +
948 + client.close();
949 + server_thread.join().expect("server join");
950 + cleanup_all(svc);
951 +}
952 +
953 +#[cfg(target_os = "linux")]
954 +#[test]
955 +fn test_server_falls_back_to_baseline_when_linux_shm_prepare_fails() {
956 + let svc = "rs_svc_shm_upgrade_fail";
957 + ensure_run_dir();
958 + cleanup_all(svc);
959 +
960 + let shm_path = format!("{TEST_RUN_DIR}/{svc}-{:016x}.ipcshm", 1u64);
961 + let _ = std::fs::remove_dir_all(&shm_path);
962 + std::fs::create_dir(&shm_path).expect("create SHM obstruction directory");
963 +
964 + let mut server = TestServer::start_with(
965 + svc,
966 + shm_server_config(),
967 + METHOD_INCREMENT,
968 + Some(increment_dispatch_handler()),
969 + 8,
970 + );
971 + let mut client = increment_client(svc, shm_client_config());
972 +
973 + assert!(client.refresh(), "client should transition to READY");
974 + assert!(client.ready(), "client should remain usable over baseline");
975 + assert_eq!(client.state, ClientState::Ready);
976 + assert_eq!(
977 + client.session.as_ref().map(|s| s.selected_profile),
978 + Some(PROFILE_BASELINE)
979 + );
980 + assert!(client.shm.is_none(), "fallback session must not attach SHM");
981 +
982 + client.close();
983 + server.stop();
984 + let _ = std::fs::remove_dir_all(&shm_path);
985 + cleanup_all(svc);
986 +}
987 +
988 +#[cfg(target_os = "linux")]
989 +#[test]
990 +fn test_call_increment_shm_short_response_truncated() {
991 + let svc = "rs_svc_shm_short_resp";
992 + ensure_run_dir();
993 + cleanup_all(svc);
994 +
995 + let ready = Arc::new(AtomicBool::new(false));
996 + let ready_clone = ready.clone();
997 + let svc_clone = svc.to_string();
998 +
999 + let server_thread = thread::spawn(move || -> Result<(), String> {
1000 + let listener = UdsListener::bind(TEST_RUN_DIR, &svc_clone, shm_server_config())
1001 + .map_err(|e| format!("bind: {e}"))?;
1002 + ready_clone.store(true, Ordering::Release);
1003 + let session = listener.accept().map_err(|e| format!("accept: {e}"))?;
1004 + if session.selected_profile != PROFILE_SHM_FUTEX {
1005 + return Err(format!("unexpected profile {}", session.selected_profile));
1006 + }
1007 +
1008 + let mut shm = ShmContext::server_create(
1009 + TEST_RUN_DIR,
1010 + &svc_clone,
1011 + session.session_id,
1012 + session.max_request_payload_bytes + HEADER_SIZE as u32,
1013 + session.max_response_payload_bytes + HEADER_SIZE as u32,
1014 + )
1015 + .map_err(|e| format!("server_create: {e}"))?;
1016 +
1017 + let mut req_buf = vec![0u8; session.max_request_payload_bytes as usize + HEADER_SIZE];
1018 + let mlen = shm
1019 + .receive(&mut req_buf, 5000)
1020 + .map_err(|e| format!("shm receive: {e}"))?;
1021 + if mlen < HEADER_SIZE {
1022 + return Err(format!("request too short: {mlen}"));
1023 + }
1024 +
1025 + shm.send(&[0xAB])
1026 + .map_err(|e| format!("shm send short response: {e}"))?;
1027 + Ok(())
1028 + });
1029 +
1030 + while !ready.load(Ordering::Acquire) {
1031 + thread::sleep(Duration::from_millis(1));
1032 + }
1033 +
1034 + let mut client = increment_client(svc, shm_client_config());
1035 + connect_ready(&mut client);
1036 + assert!(client.shm.is_some(), "expected SHM to be negotiated");
1037 +
1038 + let err = client
1039 + .call_increment(41)
1040 + .expect_err("short SHM response should be truncated");
1041 + assert_eq!(err, NipcError::Truncated);
1042 +
1043 + client.close();
1044 + match server_thread.join() {
1045 + Ok(Ok(())) => {}
1046 + Ok(Err(err)) => panic!("raw shm server failed: {err}"),
1047 + Err(_) => panic!("raw shm server panicked"),
1048 + }
1049 + cleanup_all(svc);
1050 +}
1051 +
1052 +#[cfg(target_os = "linux")]
1053 +#[test]
1054 +fn test_raw_shm_session_server_wait_panics_on_baseline_profile() {
1055 + let svc = "rs_svc_shm_helper_bad_profile";
1056 + ensure_run_dir();
1057 + cleanup_all(svc);
1058 +
1059 + let mut server = start_raw_shm_session_server(svc, server_config(), move |_, _, _| Ok(()));
1060 +
1061 + let mut client = increment_client(svc, client_config());
1062 + connect_ready(&mut client);
1063 + client.close();
1064 +
1065 + let panic = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| server.wait()))
1066 + .expect_err("baseline profile should panic in raw SHM helper");
1067 + let msg = panic_payload_to_string(panic);
1068 + assert!(
1069 + msg.contains("unexpected profile"),
1070 + "unexpected panic message: {msg}"
1071 + );
1072 +
1073 + cleanup_all(svc);
1074 +}
1075 +
1076 +#[cfg(target_os = "linux")]
1077 +#[test]
1078 +fn test_raw_shm_session_server_wait_panics_on_short_request() {
1079 + let svc = "rs_svc_shm_helper_short_req";
1080 + ensure_run_dir();
1081 + cleanup_all(svc);
1082 +
1083 + let mut server = start_raw_shm_session_server(svc, shm_server_config(), move |_, _, _| Ok(()));
1084 +
1085 + let mut client = increment_client(svc, shm_client_config());
1086 + connect_ready(&mut client);
1087 + assert!(
1088 + client.shm.is_some(),
1089 + "expected SHM transport to be negotiated"
1090 + );
1091 +
1092 + client
1093 + .shm
1094 + .as_mut()
1095 + .expect("shm")
1096 + .send(&[0xAB])
1097 + .expect("send short SHM request");
1098 + client.close();
1099 +
1100 + let panic = std::panic::catch_unwind(std::panic::AssertUnwindSafe(|| server.wait()))
1101 + .expect_err("short SHM request should panic in raw SHM helper");
1102 + let msg = panic_payload_to_string(panic);
1103 + assert!(
1104 + msg.contains("request too short"),
1105 + "unexpected panic message: {msg}"
1106 + );
1107 +
1108 + cleanup_all(svc);
1109 +}
1110 +
1111 +#[cfg(target_os = "linux")]
1112 +#[test]
1113 +fn test_call_increment_shm_server_thread_rejects_baseline_profile() {
1114 + let svc = "rs_svc_shm_short_resp_bad_profile";
1115 + ensure_run_dir();
1116 + cleanup_all(svc);
1117 +
1118 + let ready = Arc::new(AtomicBool::new(false));
1119 + let ready_clone = ready.clone();
1120 + let svc_clone = svc.to_string();
1121 +
1122 + let server_thread = thread::spawn(move || -> Result<(), String> {
1123 + let listener = UdsListener::bind(TEST_RUN_DIR, &svc_clone, shm_server_config())
1124 + .map_err(|e| format!("bind: {e}"))?;
1125 + ready_clone.store(true, Ordering::Release);
1126 + let session = listener.accept().map_err(|e| format!("accept: {e}"))?;
1127 + if session.selected_profile != PROFILE_SHM_FUTEX {
1128 + return Err(format!("unexpected profile {}", session.selected_profile));
1129 + }
1130 + Ok(())
1131 + });
1132 +
1133 + while !ready.load(Ordering::Acquire) {
1134 + thread::sleep(Duration::from_millis(1));
1135 + }
1136 +
1137 + let mut client = increment_client(svc, client_config());
1138 + connect_ready(&mut client);
1139 + client.close();
1140 +
1141 + match server_thread.join() {
1142 + Ok(Err(err)) => assert!(
1143 + err.contains("unexpected profile"),
1144 + "unexpected error: {err}"
1145 + ),
1146 + Ok(Ok(())) => panic!("baseline profile should not satisfy the SHM-only helper"),
1147 + Err(_) => panic!("raw shm server panicked"),
1148 + }
1149 +
1150 + cleanup_all(svc);
1151 +}
1152 +
1153 +#[cfg(target_os = "linux")]
1154 +#[test]
1155 +fn test_call_increment_shm_server_thread_rejects_short_request() {
1156 + let svc = "rs_svc_shm_short_resp_short_req";
1157 + ensure_run_dir();
1158 + cleanup_all(svc);
1159 +
1160 + let ready = Arc::new(AtomicBool::new(false));
1161 + let ready_clone = ready.clone();
1162 + let svc_clone = svc.to_string();
1163 +
1164 + let server_thread = thread::spawn(move || -> Result<(), String> {
1165 + let listener = UdsListener::bind(TEST_RUN_DIR, &svc_clone, shm_server_config())
1166 + .map_err(|e| format!("bind: {e}"))?;
1167 + ready_clone.store(true, Ordering::Release);
1168 + let session = listener.accept().map_err(|e| format!("accept: {e}"))?;
1169 + if session.selected_profile != PROFILE_SHM_FUTEX {
1170 + return Err(format!("unexpected profile {}", session.selected_profile));
1171 + }
1172 +
1173 + let mut shm = ShmContext::server_create(
1174 + TEST_RUN_DIR,
1175 + &svc_clone,
1176 + session.session_id,
1177 + session.max_request_payload_bytes + HEADER_SIZE as u32,
1178 + session.max_response_payload_bytes + HEADER_SIZE as u32,
1179 + )
1180 + .map_err(|e| format!("server_create: {e}"))?;
1181 +
1182 + let mut req_buf = vec![0u8; session.max_request_payload_bytes as usize + HEADER_SIZE];
1183 + let mlen = shm
1184 + .receive(&mut req_buf, 5000)
1185 + .map_err(|e| format!("shm receive: {e}"))?;
1186 + if mlen < HEADER_SIZE {
1187 + return Err(format!("request too short: {mlen}"));
1188 + }
1189 +
1190 + Ok(())
1191 + });
1192 +
1193 + while !ready.load(Ordering::Acquire) {
1194 + thread::sleep(Duration::from_millis(1));
1195 + }
1196 +
1197 + let mut client = increment_client(svc, shm_client_config());
1198 + connect_ready(&mut client);
1199 + assert!(client.shm.is_some(), "expected SHM to be negotiated");
1200 + client
1201 + .shm
1202 + .as_mut()
1203 + .expect("shm")
1204 + .send(&[0xAB])
1205 + .expect("send short SHM request");
1206 + client.close();
1207 +
1208 + match server_thread.join() {
1209 + Ok(Err(err)) => assert!(err.contains("request too short"), "unexpected error: {err}"),
1210 + Ok(Ok(())) => panic!("short SHM request should not be accepted"),
1211 + Err(_) => panic!("raw shm server panicked"),
1212 + }
1213 +
1214 + cleanup_all(svc);
1215 +}
1216 +
1217 +#[cfg(target_os = "linux")]
1218 +#[test]
1219 +fn test_call_increment_shm_rejects_bad_message_id() {
1220 + let svc = "rs_svc_shm_inc_bad_mid";
1221 + ensure_run_dir();
1222 + cleanup_all(svc);
1223 +
1224 + let mut server =
1225 + start_raw_shm_session_server(svc, shm_server_config(), move |shm, req_hdr, _| {
1226 + let mut payload = [0u8; INCREMENT_PAYLOAD_SIZE];
1227 + let n = increment_encode(43, &mut payload);
1228 + if n != INCREMENT_PAYLOAD_SIZE {
1229 + return Err(format!("increment_encode returned {n}"));
1230 + }
1231 +
1232 + let resp_hdr = Header {
1233 + magic: MAGIC_MSG,
1234 + version: VERSION,
1235 + header_len: protocol::HEADER_LEN,
1236 + kind: KIND_RESPONSE,
1237 + code: METHOD_INCREMENT,
1238 + flags: 0,
1239 + payload_len: n as u32,
1240 + item_count: 1,
1241 + message_id: req_hdr.message_id + 1,
1242 + transport_status: STATUS_OK,
1243 + };
1244 + let msg = encode_raw_message(&resp_hdr, &payload[..n]);
1245 + shm.send(&msg).map_err(|e| format!("shm send: {e}"))
1246 + });
1247 +
1248 + let mut client = increment_client(svc, shm_client_config());
1249 + connect_ready(&mut client);
1250 + assert!(client.shm.is_some(), "expected SHM to be negotiated");
1251 +
1252 + let err = client
1253 + .call_increment(42)
1254 + .expect_err("bad SHM response message_id");
1255 + assert_eq!(err, NipcError::BadLayout);
1256 +
1257 + client.close();
1258 + server.wait();
1259 + cleanup_all(svc);
1260 +}
1261 +
1262 +#[cfg(target_os = "linux")]
1263 +#[test]
1264 +fn test_call_string_reverse_shm_rejects_bad_message_id() {
1265 + let svc = "rs_svc_shm_str_bad_mid";
1266 + ensure_run_dir();
1267 + cleanup_all(svc);
1268 +
1269 + let mut server =
1270 + start_raw_shm_session_server(svc, shm_server_config(), move |shm, req_hdr, _| {
1271 + let mut payload = [0u8; 128];
1272 + let n = string_reverse_encode(b"olleh", &mut payload);
1273 + if n == 0 {
1274 + return Err("string_reverse_encode returned 0".into());
1275 + }
1276 +
1277 + let resp_hdr = Header {
1278 + magic: MAGIC_MSG,
1279 + version: VERSION,
1280 + header_len: protocol::HEADER_LEN,
1281 + kind: KIND_RESPONSE,
1282 + code: METHOD_STRING_REVERSE,
1283 + flags: 0,
1284 + payload_len: n as u32,
1285 + item_count: 1,
1286 + message_id: req_hdr.message_id + 1,
1287 + transport_status: STATUS_OK,
1288 + };
1289 + let msg = encode_raw_message(&resp_hdr, &payload[..n]);
1290 + shm.send(&msg).map_err(|e| format!("shm send: {e}"))
1291 + });
1292 +
1293 + let mut client = string_reverse_client(svc, shm_client_config());
1294 + connect_ready(&mut client);
1295 + assert!(client.shm.is_some(), "expected SHM to be negotiated");
1296 +
1297 + let err = client
1298 + .call_string_reverse("hello")
1299 + .expect_err("bad SHM response message_id");
1300 + assert_eq!(err, NipcError::BadLayout);
1301 +
1302 + client.close();
1303 + server.wait();
1304 + cleanup_all(svc);
1305 +}
1306 +
1307 +#[cfg(target_os = "linux")]
1308 +#[test]
1309 +fn test_call_increment_batch_shm_rejects_bad_message_id() {
1310 + let svc = "rs_svc_shm_batch_bad_mid";
1311 + ensure_run_dir();
1312 + cleanup_all(svc);
1313 +
1314 + let mut server =
1315 + start_raw_shm_session_server(svc, shm_server_config(), move |shm, req_hdr, _| {
1316 + let mut encoded = [0u8; INCREMENT_PAYLOAD_SIZE];
1317 + let n = increment_encode(11, &mut encoded);
1318 + if n != INCREMENT_PAYLOAD_SIZE {
1319 + return Err(format!("increment_encode returned {n}"));
1320 + }
1321 +
1322 + let mut response_buf = vec![0u8; 128];
1323 + let resp_len = {
1324 + let mut batch = BatchBuilder::new(&mut response_buf, 2);
1325 + batch
1326 + .add(&encoded)
1327 + .map_err(|e| format!("batch add 1: {e:?}"))?;
1328 + batch
1329 + .add(&encoded)
1330 + .map_err(|e| format!("batch add 2: {e:?}"))?;
1331 + let (len, _count) = batch.finish();
1332 + len
1333 + };
1334 +
1335 + let resp_hdr = Header {
1336 + magic: MAGIC_MSG,
1337 + version: VERSION,
1338 + header_len: protocol::HEADER_LEN,
1339 + kind: KIND_RESPONSE,
1340 + code: METHOD_INCREMENT,
1341 + flags: FLAG_BATCH,
1342 + payload_len: resp_len as u32,
1343 + item_count: 2,
1344 + message_id: req_hdr.message_id + 1,
1345 + transport_status: STATUS_OK,
1346 + };
1347 + let msg = encode_raw_message(&resp_hdr, &response_buf[..resp_len]);
1348 + shm.send(&msg).map_err(|e| format!("shm send: {e}"))
1349 + });
1350 +
1351 + let mut client = increment_client(svc, shm_client_config());
1352 + connect_ready(&mut client);
1353 + assert!(client.shm.is_some(), "expected SHM to be negotiated");
1354 +
1355 + let err = client
1356 + .call_increment_batch(&[10, 20])
1357 + .expect_err("bad SHM batch response message_id");
1358 + assert_eq!(err, NipcError::BadLayout);
1359 +
1360 + client.close();
1361 + server.wait();
1362 + cleanup_all(svc);
1363 +}
1364 +
1365 +#[test]
1366 +fn test_retry_on_failure() {
1367 + let svc = "rs_svc_retry";
1368 + ensure_run_dir();
1369 + cleanup_all(svc);
1370 +
1371 + let mut server1 =
1372 + TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
1373 +
1374 + let mut client = snapshot_client(svc, client_config());
1375 + client.refresh();
1376 + assert!(client.ready());
1377 +
1378 + // First call succeeds
1379 + let view = client.call_snapshot().expect("first call");
1380 + assert_eq!(view.item_count, 3);
1381 +
1382 + // Kill server
1383 + server1.stop();
1384 + cleanup_all(svc);
1385 + thread::sleep(Duration::from_millis(50));
1386 +
1387 + // Restart server
1388 + let mut server2 =
1389 + TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
1390 +
1391 + // Next call triggers reconnect + retry
1392 + let view2 = client.call_snapshot().expect("retry call");
1393 + assert_eq!(view2.item_count, 3);
1394 +
1395 + // Verify reconnect happened
1396 + let status = client.status();
1397 + assert!(status.reconnect_count >= 1);
1398 +
1399 + client.close();
1400 + server2.stop();
1401 + cleanup_all(svc);
1402 +}
1403 +
1404 +#[test]
1405 +fn test_string_reverse_retry_on_failure() {
1406 + let svc = "rs_svc_retry_str";
1407 + ensure_run_dir();
1408 + cleanup_all(svc);
1409 +
1410 + let mut server1 = TestServer::start(
1411 + svc,
1412 + METHOD_STRING_REVERSE,
1413 + Some(string_reverse_dispatch_handler()),
1414 + );
1415 +
1416 + let mut client = string_reverse_client(svc, client_config());
1417 + client.refresh();
1418 + assert!(client.ready());
1419 + assert_eq!(
1420 + client
1421 + .call_string_reverse("hello")
1422 + .expect("first reverse")
1423 + .as_str(),
1424 + "olleh"
1425 + );
1426 +
1427 + server1.stop();
1428 + cleanup_all(svc);
1429 + thread::sleep(Duration::from_millis(50));
1430 +
1431 + let mut server2 = TestServer::start(
1432 + svc,
1433 + METHOD_STRING_REVERSE,
1434 + Some(string_reverse_dispatch_handler()),
1435 + );
1436 +
1437 + let result = client.call_string_reverse("hello").expect("retry reverse");
1438 + assert_eq!(result.as_str(), "olleh");
1439 +
1440 + let status = client.status();
1441 + assert!(status.reconnect_count >= 1);
1442 +
1443 + client.close();
1444 + server2.stop();
1445 + cleanup_all(svc);
1446 +}
1447 +
1448 +#[test]
1449 +fn test_increment_batch_retry_on_failure() {
1450 + let svc = "rs_svc_retry_batch";
1451 + ensure_run_dir();
1452 + cleanup_all(svc);
1453 +
1454 + let mut server1 = TestServer::start_with(
1455 + svc,
1456 + batch_server_config(),
1457 + METHOD_INCREMENT,
1458 + Some(increment_dispatch_handler()),
1459 + 8,
1460 + );
1461 +
1462 + let mut client = increment_client(svc, batch_client_config());
1463 + client.refresh();
1464 + assert!(client.ready());
1465 + assert_eq!(
1466 + client.call_increment_batch(&[10, 20]).expect("first batch"),
1467 + vec![11, 21]
1468 + );
1469 +
1470 + server1.stop();
1471 + cleanup_all(svc);
1472 + thread::sleep(Duration::from_millis(50));
1473 +
1474 + let mut server2 = TestServer::start_with(
1475 + svc,
1476 + batch_server_config(),
1477 + METHOD_INCREMENT,
1478 + Some(increment_dispatch_handler()),
1479 + 8,
1480 + );
1481 +
1482 + let result = client.call_increment_batch(&[10, 20]).expect("retry batch");
1483 + assert_eq!(result, vec![11, 21]);
1484 +
1485 + let status = client.status();
1486 + assert!(status.reconnect_count >= 1);
1487 +
1488 + client.close();
1489 + server2.stop();
1490 + cleanup_all(svc);
1491 +}
1492 +
1493 +#[test]
1494 +fn test_string_reverse_retry_second_failure() {
1495 + let svc = "rs_svc_retry_str_second_fail";
1496 + ensure_run_dir();
1497 + cleanup_all(svc);
1498 +
1499 + let mut server1 = TestServer::start_with(
1500 + svc,
1501 + batch_server_config(),
1502 + METHOD_STRING_REVERSE,
1503 + Some(string_reverse_dispatch_handler()),
1504 + 8,
1505 + );
1506 +
1507 + let mut client = string_reverse_client(svc, client_config());
1508 + client.refresh();
1509 + assert!(client.ready());
1510 + assert_eq!(
1511 + client
1512 + .call_string_reverse("hello")
1513 + .expect("initial reverse")
1514 + .as_str(),
1515 + "olleh"
1516 + );
1517 +
1518 + server1.stop();
1519 + cleanup_all(svc);
1520 + thread::sleep(Duration::from_millis(50));
1521 +
1522 + let mut server2 =
1523 + start_raw_session_server(svc, server_config(), move |session, req_hdr, payload| {
1524 + let decoded =
1525 + string_reverse_decode(payload).map_err(|e| format!("decode request: {e:?}"))?;
1526 + if decoded.as_str() != "hello" {
1527 + return Err("unexpected request payload".into());
1528 + }
1529 +
1530 + let mut resp_hdr = Header {
1531 + kind: KIND_REQUEST,
1532 + code: METHOD_STRING_REVERSE,
1533 + flags: 0,
1534 + item_count: 1,
1535 + message_id: req_hdr.message_id,
1536 + transport_status: STATUS_OK,
1537 + ..Header::default()
1538 + };
1539 + session
1540 + .send(&mut resp_hdr, payload)
1541 + .map_err(|e| format!("send: {e}"))
1542 + });
1543 +
1544 + let err = client
1545 + .call_string_reverse("hello")
1546 + .expect_err("retry should fail on malformed second response");
1547 + assert_eq!(err, NipcError::BadKind);
1548 +
1549 + let status = client.status();
1550 + assert_eq!(status.state, ClientState::Broken);
1551 + assert!(status.reconnect_count >= 1);
1552 + assert!(status.error_count >= 1);
1553 +
1554 + client.close();
1555 + server2.wait();
1556 + cleanup_all(svc);
1557 +}
1558 +
1559 +#[test]
1560 +fn test_increment_batch_retry_second_failure() {
1561 + let svc = "rs_svc_retry_batch_second_fail";
1562 + ensure_run_dir();
1563 + cleanup_all(svc);
1564 +
1565 + let mut server1 = TestServer::start_with(
1566 + svc,
1567 + batch_server_config(),
1568 + METHOD_INCREMENT,
1569 + Some(increment_dispatch_handler()),
1570 + 8,
1571 + );
1572 +
1573 + let mut client = increment_client(svc, batch_client_config());
1574 + client.refresh();
1575 + assert!(client.ready());
1576 + assert_eq!(
1577 + client
1578 + .call_increment_batch(&[10, 20])
1579 + .expect("initial batch"),
1580 + vec![11, 21]
1581 + );
1582 +
1583 + server1.stop();
1584 + cleanup_all(svc);
1585 + thread::sleep(Duration::from_millis(50));
1586 +
1587 + let mut server2 = start_raw_session_server(
1588 + svc,
1589 + batch_server_config(),
1590 + move |session, req_hdr, payload| {
1591 + let (item0, _) =
1592 + batch_item_get(payload, 2, 0).map_err(|e| format!("decode batch item 0: {e:?}"))?;
1593 + let (item1, _) =
1594 + batch_item_get(payload, 2, 1).map_err(|e| format!("decode batch item 1: {e:?}"))?;
1595 + let v0 = increment_decode(item0).map_err(|e| format!("decode inc0: {e:?}"))?;
1596 + let v1 = increment_decode(item1).map_err(|e| format!("decode inc1: {e:?}"))?;
1597 + if v0 != 10 || v1 != 20 {
1598 + return Err("unexpected batch payload".into());
1599 + }
1600 +
1601 + let mut resp_hdr = Header {
1602 + kind: KIND_REQUEST,
1603 + code: METHOD_INCREMENT,
1604 + flags: 0,
1605 + item_count: 1,
1606 + message_id: req_hdr.message_id,
1607 + transport_status: STATUS_OK,
1608 + ..Header::default()
1609 + };
1610 + session
1611 + .send(&mut resp_hdr, &[])
1612 + .map_err(|e| format!("send: {e}"))
1613 + },
1614 + );
1615 +
1616 + let err = client
1617 + .call_increment_batch(&[10, 20])
1618 + .expect_err("retry should fail on malformed second batch response");
1619 + assert_eq!(err, NipcError::BadKind);
1620 +
1621 + let status = client.status();
1622 + assert_eq!(status.state, ClientState::Broken);
1623 + assert!(status.reconnect_count >= 1);
1624 + assert!(status.error_count >= 1);
1625 +
1626 + client.close();
1627 + server2.wait();
1628 + cleanup_all(svc);
1629 +}
1630 +
1631 +#[test]
1632 +fn test_multiple_clients() {
1633 + let svc = "rs_svc_multi";
1634 + ensure_run_dir();
1635 + cleanup_all(svc);
1636 +
1637 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
1638 +
1639 + // Create and connect client 1
1640 + let mut client1 = snapshot_client(svc, client_config());
1641 + client1.refresh();
1642 + assert!(client1.ready());
1643 +
1644 + let view1 = client1.call_snapshot().expect("client 1 call");
1645 + assert_eq!(view1.item_count, 3);
1646 +
1647 + // Now multi-client: keep client 1 open, connect client 2
1648 + let mut client2 = snapshot_client(svc, client_config());
1649 + client2.refresh();
1650 + assert!(client2.ready());
1651 +
1652 + let view2 = client2.call_snapshot().expect("client 2 call");
1653 + assert_eq!(view2.item_count, 3);
1654 +
1655 + client1.close();
1656 + client2.close();
1657 + server.stop();
1658 + cleanup_all(svc);
1659 +}
1660 +
1661 +#[test]
1662 +fn test_server_rejects_session_at_worker_capacity() {
1663 + let svc = "rs_svc_worker_capacity";
1664 + ensure_run_dir();
1665 + cleanup_all(svc);
1666 +
1667 + let handler_calls = Arc::new(std::sync::atomic::AtomicU32::new(0));
1668 + let release = Arc::new(AtomicBool::new(false));
1669 +
1670 + let handler = {
1671 + let handler_calls = handler_calls.clone();
1672 + let release = release.clone();
1673 + increment_dispatch(Arc::new(move |value| {
1674 + handler_calls.fetch_add(1, Ordering::AcqRel);
1675 + while !release.load(Ordering::Acquire) {
1676 + thread::sleep(Duration::from_millis(1));
1677 + }
1678 + Some(value + 1)
1679 + }))
1680 + };
1681 +
1682 + let mut server = TestServer::start_with_workers(svc, METHOD_INCREMENT, Some(handler), 1);
1683 +
1684 + let (call_tx, call_rx) = std::sync::mpsc::channel();
1685 + let svc_name = svc.to_string();
1686 + let caller = thread::spawn(move || {
1687 + let mut client = increment_client(&svc_name, client_config());
1688 + connect_ready(&mut client);
1689 + let result = client.call_increment(41);
1690 + client.close();
1691 + call_tx.send(result).expect("send first call result");
1692 + });
1693 +
1694 + for _ in 0..500 {
1695 + if handler_calls.load(Ordering::Acquire) >= 1 {
1696 + break;
1697 + }
1698 + thread::sleep(Duration::from_millis(1));
1699 + }
1700 + assert_eq!(
1701 + handler_calls.load(Ordering::Acquire),
1702 + 1,
1703 + "first client should occupy the only worker slot"
1704 + );
1705 +
1706 + let mut session2 =
1707 + UdsSession::connect(TEST_RUN_DIR, svc, &client_config()).expect("second connect");
1708 + let mut req_hdr = Header {
1709 + kind: KIND_REQUEST,
1710 + code: METHOD_INCREMENT,
1711 + item_count: 1,
1712 + message_id: 1,
1713 + ..Header::default()
1714 + };
1715 + let mut req_payload = [0u8; INCREMENT_PAYLOAD_SIZE];
1716 + assert_eq!(
1717 + increment_encode(9, &mut req_payload),
1718 + INCREMENT_PAYLOAD_SIZE,
1719 + "IncrementEncode should fill the request buffer"
1720 + );
1721 +
1722 + let send_err = session2.send(&mut req_hdr, &req_payload).err();
1723 + let recv_ok = if send_err.is_none() {
1724 + let mut recv_buf = [0u8; HEADER_SIZE + 64];
1725 + session2.receive(&mut recv_buf).is_ok()
1726 + } else {
1727 + false
1728 + };
1729 + assert!(
1730 + !(send_err.is_none() && recv_ok),
1731 + "second session should be rejected while the server is at worker capacity"
1732 + );
1733 + assert_eq!(
1734 + handler_calls.load(Ordering::Acquire),
1735 + 1,
1736 + "second session should not enter the handler while the first session is active"
1737 + );
1738 + drop(session2);
1739 +
1740 + release.store(true, Ordering::Release);
1741 + let first_result = call_rx.recv().expect("first call result");
1742 + assert_eq!(first_result.expect("first call should succeed"), 42);
1743 + caller.join().expect("caller join");
1744 +
1745 + thread::sleep(Duration::from_millis(100));
1746 +
1747 + let mut verify = increment_client(svc, client_config());
1748 + connect_ready(&mut verify);
1749 + assert_eq!(
1750 + verify.call_increment(1).expect("verification call"),
1751 + 2,
1752 + "server should remain healthy after rejecting a session at capacity"
1753 + );
1754 + verify.close();
1755 +
1756 + server.stop();
1757 + cleanup_all(svc);
1758 +}
1759 +
1760 +#[test]
1761 +fn test_poll_fd_invalid_fd_returns_error() {
1762 + let mut fds = [0; 2];
1763 + let rc = unsafe { libc::pipe(fds.as_mut_ptr()) };
1764 + assert_eq!(rc, 0, "pipe failed");
1765 +
1766 + let read_fd = fds[0];
1767 + let write_fd = fds[1];
1768 + unsafe {
1769 + libc::close(read_fd);
1770 + }
1771 +
1772 + let rc = poll_fd(read_fd, 0);
1773 + unsafe {
1774 + libc::close(write_fd);
1775 + }
1776 +
1777 + assert_eq!(rc, -1);
1778 +}
1779 +
1780 +#[test]
1781 +fn test_poll_fd_readable_returns_ready() {
1782 + let mut fds = [0; 2];
1783 + let rc = unsafe { libc::pipe(fds.as_mut_ptr()) };
1784 + assert_eq!(rc, 0, "pipe failed");
1785 +
1786 + let read_fd = fds[0];
1787 + let write_fd = fds[1];
1788 + let wrote = unsafe { libc::write(write_fd, b"x".as_ptr() as *const libc::c_void, 1) };
1789 + assert_eq!(wrote, 1, "write failed");
1790 +
1791 + let rc = poll_fd(read_fd, 0);
1792 +
1793 + unsafe {
1794 + libc::close(read_fd);
1795 + libc::close(write_fd);
1796 + }
1797 +
1798 + assert_eq!(rc, 1);
1799 +}
1800 +
1801 +#[cfg(target_os = "linux")]
1802 +#[test]
1803 +fn test_poll_fd_eintr_returns_timeout() {
1804 + unsafe extern "C" fn noop_signal_handler(_: libc::c_int) {}
1805 +
1806 + let mut fds = [0; 2];
1807 + let rc = unsafe { libc::pipe(fds.as_mut_ptr()) };
1808 + assert_eq!(rc, 0, "pipe failed");
1809 +
1810 + let read_fd = fds[0];
1811 + let write_fd = fds[1];
1812 + let mut action: libc::sigaction = unsafe { std::mem::zeroed() };
1813 + let mut old_action: libc::sigaction = unsafe { std::mem::zeroed() };
1814 + action.sa_flags = 0;
1815 + action.sa_sigaction = noop_signal_handler as usize;
1816 + unsafe { libc::sigemptyset(&mut action.sa_mask) };
1817 + assert_eq!(
1818 + unsafe { libc::sigaction(libc::SIGUSR1, &action, &mut old_action) },
1819 + 0
1820 + );
1821 +
1822 + let tid = unsafe { libc::pthread_self() };
1823 + let signaler = thread::spawn(move || {
1824 + thread::sleep(Duration::from_millis(50));
1825 + assert_eq!(unsafe { libc::pthread_kill(tid, libc::SIGUSR1) }, 0);
1826 + });
1827 +
1828 + let rc = poll_fd(read_fd, 5000);
1829 + signaler.join().expect("signaler join");
1830 + unsafe {
1831 + libc::sigaction(libc::SIGUSR1, &old_action, std::ptr::null_mut());
1832 + libc::close(read_fd);
1833 + libc::close(write_fd);
1834 + }
1835 +
1836 + assert_eq!(rc, 0);
1837 +}
1838 +
1839 +#[test]
1840 +fn test_managed_server_recovers_after_short_uds_request() {
1841 + let svc = "rs_svc_short_uds_req";
1842 + ensure_run_dir();
1843 + cleanup_all(svc);
1844 +
1845 + let mut server = TestServer::start(svc, METHOD_INCREMENT, Some(increment_dispatch_handler()));
1846 + let session = UdsSession::connect(TEST_RUN_DIR, svc, &client_config()).expect("connect");
1847 + send_raw_packet(session.fd(), &[0xAB]);
1848 + drop(session);
1849 + thread::sleep(Duration::from_millis(50));
1850 +
1851 + verify_increment_service_ok(svc, client_config());
1852 +
1853 + server.stop();
1854 + cleanup_all(svc);
1855 +}
1856 +
1857 +#[test]
1858 +fn test_managed_server_recovers_after_bad_uds_header() {
1859 + let svc = "rs_svc_bad_uds_hdr";
1860 + ensure_run_dir();
1861 + cleanup_all(svc);
1862 +
1863 + let mut server = TestServer::start(svc, METHOD_INCREMENT, Some(increment_dispatch_handler()));
1864 + let session = UdsSession::connect(TEST_RUN_DIR, svc, &client_config()).expect("connect");
1865 + let bad_header = [0u8; HEADER_SIZE];
1866 + send_raw_packet(session.fd(), &bad_header);
1867 + drop(session);
1868 + thread::sleep(Duration::from_millis(50));
1869 +
1870 + verify_increment_service_ok(svc, client_config());
1871 +
1872 + server.stop();
1873 + cleanup_all(svc);
1874 +}
1875 +
1876 +#[test]
1877 +fn test_managed_server_recovers_after_uds_peer_closes_before_response() {
1878 + let svc = "rs_svc_uds_send_break";
1879 + ensure_run_dir();
1880 + cleanup_all(svc);
1881 +
1882 + let handler = increment_dispatch(Arc::new(|value| {
1883 + thread::sleep(Duration::from_millis(50));
1884 + Some(value + 1)
1885 + }));
1886 + let mut server = TestServer::start(svc, METHOD_INCREMENT, Some(handler));
1887 + let session = UdsSession::connect(TEST_RUN_DIR, svc, &client_config()).expect("connect");
1888 + let request = build_increment_request_message(77, 41);
1889 + send_raw_packet(session.fd(), &request);
1890 + assert_eq!(unsafe { libc::shutdown(session.fd(), libc::SHUT_RDWR) }, 0);
1891 + drop(session);
1892 + thread::sleep(Duration::from_millis(100));
1893 +
1894 + verify_increment_service_ok(svc, client_config());
1895 +
1896 + server.stop();
1897 + cleanup_all(svc);
1898 +}
1899 +
1900 +#[cfg(target_os = "linux")]
1901 +#[test]
1902 +fn test_managed_server_recovers_after_short_shm_request() {
1903 + let svc = "rs_svc_short_shm_req";
1904 + ensure_run_dir();
1905 + cleanup_all(svc);
1906 +
1907 + let mut server =
1908 + TestServer::start_shm(svc, METHOD_INCREMENT, Some(increment_dispatch_handler()));
1909 + let mut client = increment_client(svc, shm_client_config());
1910 + connect_ready(&mut client);
1911 + assert!(
1912 + client.shm.is_some(),
1913 + "expected SHM transport to be negotiated"
1914 + );
1915 +
1916 + client
1917 + .shm
1918 + .as_mut()
1919 + .expect("shm")
1920 + .send(&[0xAB])
1921 + .expect("send short SHM request");
1922 + client.close();
1923 + thread::sleep(Duration::from_millis(50));
1924 +
1925 + verify_increment_service_ok(svc, shm_client_config());
1926 +
1927 + server.stop();
1928 + cleanup_all(svc);
1929 +}
1930 +
1931 +#[cfg(target_os = "linux")]
1932 +#[test]
1933 +fn test_managed_server_recovers_after_bad_shm_header() {
1934 + let svc = "rs_svc_bad_shm_hdr";
1935 + ensure_run_dir();
1936 + cleanup_all(svc);
1937 +
1938 + let mut server =
1939 + TestServer::start_shm(svc, METHOD_INCREMENT, Some(increment_dispatch_handler()));
1940 + let mut client = increment_client(svc, shm_client_config());
1941 + connect_ready(&mut client);
1942 + assert!(
1943 + client.shm.is_some(),
1944 + "expected SHM transport to be negotiated"
1945 + );
1946 +
1947 + let bad_header = [0u8; HEADER_SIZE];
1948 + client
1949 + .shm
1950 + .as_mut()
1951 + .expect("shm")
1952 + .send(&bad_header)
1953 + .expect("send bad SHM header");
1954 + client.close();
1955 + thread::sleep(Duration::from_millis(50));
1956 +
1957 + verify_increment_service_ok(svc, shm_client_config());
1958 +
1959 + server.stop();
1960 + cleanup_all(svc);
1961 +}
1962 +
1963 +#[test]
1964 +fn test_concurrent_clients() {
1965 + let svc = "rs_svc_concurrent";
1966 + ensure_run_dir();
1967 + cleanup_all(svc);
1968 +
1969 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
1970 +
1971 + const NUM_CLIENTS: usize = 5;
1972 + const REQUESTS_PER: usize = 10;
1973 +
1974 + let mut handles = Vec::new();
1975 +
1976 + for _ in 0..NUM_CLIENTS {
1977 + let svc_name = svc.to_string();
1978 + let handle = thread::spawn(move || {
1979 + let mut client = snapshot_client(&svc_name, client_config());
1980 +
1981 + // Connect with retry
1982 + for _ in 0..100 {
1983 + client.refresh();
1984 + if client.ready() {
1985 + break;
1986 + }
1987 + thread::sleep(Duration::from_millis(10));
1988 + }
1989 +
1990 + assert!(client.ready(), "client must be ready");
1991 +
1992 + let mut successes = 0usize;
1993 + for _ in 0..REQUESTS_PER {
1994 + match client.call_snapshot() {
1995 + Ok(view) => {
1996 + assert_eq!(view.item_count, 3);
1997 + assert_eq!(view.generation, 42);
1998 +
1999 + // Verify first item content
2000 + let item0 = view.item(0).expect("item 0");
2001 + assert_eq!(item0.hash, 1001);
2002 + assert_eq!(
2003 + std::str::from_utf8(item0.name.as_bytes()).unwrap(),
2004 + "docker-abc123"
2005 + );
2006 +
2007 + successes += 1;
2008 + }
2009 + Err(e) => panic!("call failed: {:?}", e),
2010 + }
2011 + }
2012 + client.close();
2013 + successes
2014 + });
2015 + handles.push(handle);
2016 + }
2017 +
2018 + let mut total = 0usize;
2019 + for h in handles {
2020 + total += h.join().expect("client thread panicked");
2021 + }
2022 +
2023 + assert_eq!(total, NUM_CLIENTS * REQUESTS_PER);
2024 +
2025 + server.stop();
2026 + cleanup_all(svc);
2027 +}
2028 +
2029 +#[test]
2030 +fn test_handler_failure() {
2031 + let svc = "rs_svc_hfail";
2032 + ensure_run_dir();
2033 + cleanup_all(svc);
2034 +
2035 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, None);
2036 +
2037 + let mut client = snapshot_client(svc, client_config());
2038 + client.refresh();
2039 + assert!(client.ready());
2040 +
2041 + // Call should fail (handler returns None -> INTERNAL_ERROR)
2042 + let err = client.call_snapshot();
2043 + assert!(err.is_err());
2044 +
2045 + let status = client.status();
2046 + assert!(status.error_count >= 1);
2047 +
2048 + client.close();
2049 + server.stop();
2050 + cleanup_all(svc);
2051 +}
2052 +
2053 +#[test]
2054 +fn test_status_reporting() {
2055 + let svc = "rs_svc_status";
2056 + ensure_run_dir();
2057 + cleanup_all(svc);
2058 +
2059 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
2060 +
2061 + let mut client = snapshot_client(svc, client_config());
2062 + client.refresh();
2063 + assert!(client.ready());
2064 +
2065 + // Initial counters
2066 + let s0 = client.status();
2067 + assert_eq!(s0.connect_count, 1);
2068 + assert_eq!(s0.call_count, 0);
2069 + assert_eq!(s0.error_count, 0);
2070 +
2071 + // Make 3 successful calls
2072 + for _ in 0..3 {
2073 + client.call_snapshot().expect("call ok");
2074 + }
2075 +
2076 + let s1 = client.status();
2077 + assert_eq!(s1.call_count, 3);
2078 + assert_eq!(s1.error_count, 0);
2079 +
2080 + // Call on disconnected client
2081 + client.close();
2082 + let err = client.call_snapshot();
2083 + assert!(err.is_err());
2084 +
2085 + let s2 = client.status();
2086 + assert_eq!(s2.error_count, 1);
2087 +
2088 + server.stop();
2089 + cleanup_all(svc);
2090 +}
2091 +
2092 +#[test]
2093 +fn test_non_request_terminates_session() {
2094 + // Send a RESPONSE message to a server; the server must terminate
2095 + // the session (protocol violation), so subsequent requests fail.
2096 + let svc = "rs_svc_nonreq";
2097 + ensure_run_dir();
2098 + cleanup_all(svc);
2099 +
2100 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
2101 +
2102 + // Connect via raw UDS session
2103 + let mut session = UdsSession::connect(TEST_RUN_DIR, svc, &client_config()).expect("connect");
2104 +
2105 + // Send a RESPONSE (not REQUEST) - protocol violation
2106 + let mut hdr = Header {
2107 + kind: KIND_RESPONSE,
2108 + code: METHOD_CGROUPS_SNAPSHOT,
2109 + flags: 0,
2110 + item_count: 0,
2111 + message_id: 1,
2112 + transport_status: STATUS_OK,
2113 + ..Header::default()
2114 + };
2115 + let send_result = session.send(&mut hdr, &[]);
2116 + // Send may succeed (the bytes go out)
2117 + if send_result.is_ok() {
2118 + // But subsequent communication should fail because the
2119 + // server terminated the session
2120 + thread::sleep(Duration::from_millis(100));
2121 + let mut recv_buf = vec![0u8; 4096];
2122 + // Try to send a valid request and receive - should fail
2123 + let mut hdr2 = Header {
2124 + kind: KIND_REQUEST,
2125 + code: METHOD_CGROUPS_SNAPSHOT,
2126 + flags: 0,
2127 + item_count: 1,
2128 + message_id: 2,
2129 + transport_status: STATUS_OK,
2130 + ..Header::default()
2131 + };
2132 + let req = CgroupsRequest {
2133 + layout_version: 1,
2134 + flags: 0,
2135 + };
2136 + let mut req_buf = [0u8; 4];
2137 + req.encode(&mut req_buf);
2138 + let _ = session.send(&mut hdr2, &req_buf);
2139 + let recv = session.receive(&mut recv_buf);
2140 + assert!(
2141 + recv.is_err(),
2142 + "server should have terminated session after non-request message"
2143 + );
2144 + }
2145 +
2146 + drop(session);
2147 +
2148 + // Verify server is still alive: connect a new client and do a normal call
2149 + let mut verify_client = snapshot_client(svc, client_config());
2150 + verify_client.refresh();
2151 + assert!(
2152 + verify_client.ready(),
2153 + "server should still be alive after bad client"
2154 + );
2155 +
2156 + let view = verify_client
2157 + .call_snapshot()
2158 + .expect("normal call should succeed after bad client");
2159 + assert_eq!(
2160 + view.item_count, 3,
2161 + "response should be correct after bad client"
2162 + );
2163 +
2164 + verify_client.close();
2165 + server.stop();
2166 + cleanup_all(svc);
2167 +}
2168 +
2169 +// ---------------------------------------------------------------
2170 +// L3 Cache tests
2171 +// ---------------------------------------------------------------
2172 +
2173 +#[test]
2174 +fn test_cache_full_round_trip() {
2175 + let svc = "rs_cache_rt";
2176 + ensure_run_dir();
2177 + cleanup_all(svc);
2178 +
2179 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
2180 +
2181 + let mut cache = CgroupsCache::new(TEST_RUN_DIR, svc, client_config());
2182 + assert!(!cache.ready());
2183 +
2184 + // Ensure monotonic epoch advances past 0ms before refresh
2185 + thread::sleep(Duration::from_millis(2));
2186 +
2187 + // Refresh populates the cache
2188 + let updated = cache.refresh();
2189 + assert!(updated);
2190 + assert!(cache.ready());
2191 +
2192 + // Lookup by hash + name
2193 + let item = cache.lookup(1001, "docker-abc123");
2194 + assert!(item.is_some());
2195 + let item = item.unwrap();
2196 + assert_eq!(item.hash, 1001);
2197 + assert_eq!(item.options, 0);
2198 + assert_eq!(item.enabled, 1);
2199 + assert_eq!(item.name, "docker-abc123");
2200 + assert_eq!(item.path, "/sys/fs/cgroup/docker/abc123");
2201 +
2202 + let item2 = cache.lookup(3003, "systemd-user");
2203 + assert!(item2.is_some());
2204 + assert_eq!(item2.unwrap().enabled, 0);
2205 +
2206 + // Status
2207 + let status = cache.status();
2208 + assert!(status.populated);
2209 + assert_eq!(status.item_count, 3);
2210 + assert_eq!(status.systemd_enabled, 1);
2211 + assert_eq!(status.generation, 42);
2212 + assert_eq!(status.refresh_success_count, 1);
2213 + assert_eq!(status.refresh_failure_count, 0);
2214 + assert_eq!(status.connection_state, ClientState::Ready);
2215 + assert!(status.last_refresh_ts > 0);
2216 +
2217 + cache.close();
2218 + server.stop();
2219 + cleanup_all(svc);
2220 +}
2221 +
2222 +#[test]
2223 +fn test_cache_refresh_failure_preserves() {
2224 + let svc = "rs_cache_preserve";
2225 + ensure_run_dir();
2226 + cleanup_all(svc);
2227 +
2228 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
2229 +
2230 + let mut cache = CgroupsCache::new(TEST_RUN_DIR, svc, client_config());
2231 +
2232 + // First refresh populates cache
2233 + assert!(cache.refresh());
2234 + assert!(cache.ready());
2235 + assert!(cache.lookup(1001, "docker-abc123").is_some());
2236 +
2237 + // Kill server
2238 + server.stop();
2239 + cleanup_all(svc);
2240 + thread::sleep(Duration::from_millis(50));
2241 +
2242 + // Refresh fails, but old cache is preserved
2243 + let updated = cache.refresh();
2244 + assert!(!updated);
2245 + assert!(cache.ready()); // still has cached data
2246 + assert!(cache.lookup(1001, "docker-abc123").is_some());
2247 +
2248 + let status = cache.status();
2249 + assert_eq!(status.refresh_success_count, 1);
2250 + assert!(status.refresh_failure_count >= 1);
2251 +
2252 + cache.close();
2253 + cleanup_all(svc);
2254 +}
2255 +
2256 +#[test]
2257 +fn test_cache_reconnect_rebuilds() {
2258 + let svc = "rs_cache_reconn";
2259 + ensure_run_dir();
2260 + cleanup_all(svc);
2261 +
2262 + let mut server1 =
2263 + TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
2264 +
2265 + let mut cache = CgroupsCache::new(TEST_RUN_DIR, svc, client_config());
2266 + assert!(cache.refresh());
2267 + assert_eq!(cache.status().item_count, 3);
2268 +
2269 + // Kill and restart server
2270 + server1.stop();
2271 + cleanup_all(svc);
2272 + thread::sleep(Duration::from_millis(50));
2273 +
2274 + let mut server2 =
2275 + TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
2276 +
2277 + // Refresh should reconnect and rebuild cache
2278 + let updated = cache.refresh();
2279 + assert!(updated);
2280 + assert!(cache.ready());
2281 + assert_eq!(cache.status().item_count, 3);
2282 + assert_eq!(cache.status().refresh_success_count, 2);
2283 +
2284 + cache.close();
2285 + server2.stop();
2286 + cleanup_all(svc);
2287 +}
2288 +
2289 +#[test]
2290 +fn test_cache_lookup_not_found() {
2291 + let svc = "rs_cache_notfound";
2292 + ensure_run_dir();
2293 + cleanup_all(svc);
2294 +
2295 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
2296 +
2297 + let mut cache = CgroupsCache::new(TEST_RUN_DIR, svc, client_config());
2298 + assert!(cache.refresh());
2299 +
2300 + // Non-existent hash
2301 + assert!(cache.lookup(9999, "nonexistent").is_none());
2302 +
2303 + // Correct hash, wrong name
2304 + assert!(cache.lookup(1001, "wrong-name").is_none());
2305 +
2306 + // Correct name, wrong hash
2307 + assert!(cache.lookup(9999, "docker-abc123").is_none());
2308 +
2309 + cache.close();
2310 + server.stop();
2311 + cleanup_all(svc);
2312 +}
2313 +
2314 +#[test]
2315 +fn test_cache_empty() {
2316 + let svc = "rs_cache_empty";
2317 + ensure_run_dir();
2318 + cleanup_all(svc);
2319 +
2320 + let cache = CgroupsCache::new(TEST_RUN_DIR, svc, client_config());
2321 +
2322 + // Not ready before any refresh
2323 + assert!(!cache.ready());
2324 +
2325 + // Lookup on empty cache returns None
2326 + assert!(cache.lookup(1001, "docker-abc123").is_none());
2327 +
2328 + let status = cache.status();
2329 + assert!(!status.populated);
2330 + assert_eq!(status.item_count, 0);
2331 + assert_eq!(status.refresh_success_count, 0);
2332 + assert_eq!(status.refresh_failure_count, 0);
2333 +
2334 + cleanup_all(svc);
2335 +}
2336 +
2337 +#[test]
2338 +fn test_cache_large_dataset() {
2339 + let svc = "rs_cache_large";
2340 + ensure_run_dir();
2341 + cleanup_all(svc);
2342 +
2343 + const N: u32 = 1000;
2344 +
2345 + // Handler that builds N items
2346 + fn large_snapshot_dispatch() -> DispatchHandler {
2347 + snapshot_dispatch(
2348 + Arc::new(|req, builder| {
2349 + if req.layout_version != 1 || req.flags != 0 {
2350 + return false;
2351 + }
2352 + builder.set_header(1, 100);
2353 +
2354 + for i in 0..N {
2355 + let name = format!("cgroup-{i}");
2356 + let path = format!("/sys/fs/cgroup/test/{i}");
2357 + if builder
2358 + .add(
2359 + i + 1000,
2360 + 0,
2361 + if i % 3 == 0 { 0 } else { 1 },
2362 + name.as_bytes(),
2363 + path.as_bytes(),
2364 + )
2365 + .is_err()
2366 + {
2367 + return false;
2368 + }
2369 + }
2370 +
2371 + true
2372 + }),
2373 + N,
2374 + )
2375 + }
2376 +
2377 + // Use a larger response buf size
2378 + let mut cfg = client_config();
2379 + cfg.max_response_payload_bytes = 256 * N;
2380 +
2381 + let mut server = TestServer::start_with_resp_size(
2382 + svc,
2383 + METHOD_CGROUPS_SNAPSHOT,
2384 + Some(large_snapshot_dispatch()),
2385 + 256 * N as usize,
2386 + );
2387 +
2388 + let mut cache = CgroupsCache::new(TEST_RUN_DIR, svc, cfg);
2389 +
2390 + assert!(cache.refresh());
2391 + assert_eq!(cache.status().item_count, N);
2392 +
2393 + // Verify all lookups
2394 + for i in 0..N {
2395 + let name = format!("cgroup-{i}");
2396 + let item = cache.lookup(i + 1000, &name);
2397 + assert!(item.is_some(), "item {i} not found");
2398 + let item = item.unwrap();
2399 + assert_eq!(item.hash, i + 1000);
2400 + let expected_path = format!("/sys/fs/cgroup/test/{i}");
2401 + assert_eq!(item.path, expected_path);
2402 + }
2403 +
2404 + cache.close();
2405 + server.stop();
2406 + cleanup_all(svc);
2407 +}
2408 +
2409 +#[test]
2410 +fn test_cache_refresh_lossy_utf8() {
2411 + let svc = "rs_cache_lossy";
2412 + ensure_run_dir();
2413 + cleanup_all(svc);
2414 +
2415 + let handler = snapshot_dispatch(
2416 + Arc::new(|_, builder| {
2417 + builder.set_header(1, 7);
2418 + builder
2419 + .add(1001, 0, 1, b"bad-\xFF-name", b"/bad/\xFF/path")
2420 + .is_ok()
2421 + }),
2422 + 1,
2423 + );
2424 +
2425 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(handler));
2426 + let mut cache = CgroupsCache::new(TEST_RUN_DIR, svc, client_config());
2427 +
2428 + assert!(cache.refresh(), "cache refresh should succeed");
2429 + let item = cache
2430 + .lookup(1001, "bad-\u{FFFD}-name")
2431 + .expect("lossy lookup");
2432 + assert_eq!(item.name, "bad-\u{FFFD}-name");
2433 + assert_eq!(item.path, "/bad/\u{FFFD}/path");
2434 +
2435 + cache.close();
2436 + server.stop();
2437 + cleanup_all(svc);
2438 +}
2439 +
2440 +#[test]
2441 +fn test_cache_refresh_preserves_old_cache_on_malformed_snapshot_item() {
2442 + let svc = "rs_cache_preserve_bad_item";
2443 + ensure_run_dir();
2444 + cleanup_all(svc);
2445 +
2446 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
2447 + let mut cache = CgroupsCache::new(TEST_RUN_DIR, svc, client_config());
2448 +
2449 + assert!(cache.refresh(), "initial refresh should succeed");
2450 + let old_status = cache.status();
2451 + let old_item = cache
2452 + .lookup(1001, "docker-abc123")
2453 + .expect("existing cache item")
2454 + .clone();
2455 +
2456 + server.stop();
2457 + cleanup_all(svc);
2458 +
2459 + let mut raw_server =
2460 + start_raw_session_server(svc, server_config(), move |session, req_hdr, payload| {
2461 + let req = crate::protocol::CgroupsRequest::decode(payload)
2462 + .map_err(|e| format!("decode request: {e:?}"))?;
2463 + if req.layout_version != 1 || req.flags != 0 {
2464 + return Err("unexpected snapshot request".into());
2465 + }
2466 +
2467 + let mut response_payload = [0u8; 512];
2468 + let response_len = {
2469 + let mut builder = CgroupsBuilder::new(&mut response_payload, 1, 1, 99);
2470 + builder
2471 + .add(9999, 0, 1, b"new-item", b"/new/path")
2472 + .map_err(|e| format!("builder add: {e:?}"))?;
2473 + builder.finish()
2474 + };
2475 +
2476 + let dir_end = 24 + 8;
2477 + let item_off =
2478 + u32::from_ne_bytes(response_payload[24..28].try_into().unwrap()) as usize;
2479 + let item_start = dir_end + item_off;
2480 + response_payload[item_start..item_start + 2].copy_from_slice(&99u16.to_ne_bytes());
2481 +
2482 + let mut resp_hdr = Header {
2483 + kind: KIND_RESPONSE,
2484 + code: METHOD_CGROUPS_SNAPSHOT,
2485 + flags: 0,
2486 + item_count: 1,
2487 + message_id: req_hdr.message_id,
2488 + transport_status: STATUS_OK,
2489 + ..Header::default()
2490 + };
2491 + session
2492 + .send(&mut resp_hdr, &response_payload[..response_len])
2493 + .map_err(|e| format!("send: {e}"))
2494 + });
2495 +
2496 + assert!(
2497 + !cache.refresh(),
2498 + "malformed snapshot item should preserve the old cache"
2499 + );
2500 +
2501 + let status = cache.status();
2502 + assert!(cache.ready(), "cache should stay populated");
2503 + assert_eq!(status.item_count, old_status.item_count);
2504 + assert_eq!(status.generation, old_status.generation);
2505 + assert_eq!(
2506 + status.refresh_success_count,
2507 + old_status.refresh_success_count
2508 + );
2509 + assert_eq!(
2510 + status.refresh_failure_count,
2511 + old_status.refresh_failure_count + 1
2512 + );
2513 + let preserved = cache
2514 + .lookup(1001, "docker-abc123")
2515 + .expect("old cache item should remain");
2516 + assert_eq!(preserved.hash, old_item.hash);
2517 + assert_eq!(preserved.path, old_item.path);
2518 + assert!(
2519 + cache.lookup(9999, "new-item").is_none(),
2520 + "bad refresh must not replace the old cache"
2521 + );
2522 +
2523 + cache.close();
2524 + raw_server.wait();
2525 + cleanup_all(svc);
2526 +}
2527 +
2528 +// ---------------------------------------------------------------
2529 +// Stress tests (Phase H4)
2530 +// ---------------------------------------------------------------
2531 +
2532 +/// djb2 hash matching the C implementation
2533 +fn simple_hash(s: &str) -> u32 {
2534 + let mut hash: u32 = 5381;
2535 + for c in s.bytes() {
2536 + hash = hash
2537 + .wrapping_shl(5)
2538 + .wrapping_add(hash)
2539 + .wrapping_add(c as u32);
2540 + }
2541 + hash
2542 +}
2543 +
2544 +struct StressTestServer {
2545 + stop_flag: Arc<AtomicBool>,
2546 + thread: Option<thread::JoinHandle<()>>,
2547 +}
2548 +
2549 +impl StressTestServer {
2550 + fn start(service: &str, n: u32, resp_buf_size: usize) -> Self {
2551 + ensure_run_dir();
2552 + cleanup_all(service);
2553 +
2554 + let svc = service.to_string();
2555 + let ready_flag = Arc::new(AtomicBool::new(false));
2556 + let ready_clone = ready_flag.clone();
2557 +
2558 + let mut scfg = server_config();
2559 + scfg.max_response_payload_bytes = resp_buf_size as u32;
2560 + scfg.packet_size = 65536; // force smaller packets for chunked transport
2561 +
2562 + let handler = snapshot_dispatch(
2563 + Arc::new(move |req, builder| {
2564 + if req.layout_version != 1 || req.flags != 0 {
2565 + return false;
2566 + }
2567 + builder.set_header(1, 42);
2568 +
2569 + for i in 0..n {
2570 + let name = format!("container-{i:04}");
2571 + let path = format!("/sys/fs/cgroup/docker/{i:04}");
2572 + let hash = simple_hash(&name);
2573 + let enabled = if i % 5 == 0 { 0 } else { 1 };
2574 + if builder
2575 + .add(hash, 0x10, enabled, name.as_bytes(), path.as_bytes())
2576 + .is_err()
2577 + {
2578 + return false;
2579 + }
2580 + }
2581 +
2582 + true
2583 + }),
2584 + n,
2585 + );
2586 +
2587 + let mut server = ManagedServer::new(
2588 + TEST_RUN_DIR,
2589 + &svc,
2590 + scfg,
2591 + METHOD_CGROUPS_SNAPSHOT,
2592 + Some(handler),
2593 + );
2594 + let stop_flag = server.running_flag();
2595 +
2596 + let thread = thread::spawn(move || {
2597 + ready_clone.store(true, Ordering::Release);
2598 + let _ = server.run();
2599 + });
2600 +
2601 + for _ in 0..2000 {
2602 + if ready_flag.load(Ordering::Acquire) {
2603 + break;
2604 + }
2605 + thread::sleep(Duration::from_micros(500));
2606 + }
2607 + thread::sleep(Duration::from_millis(50));
2608 +
2609 + StressTestServer {
2610 + stop_flag,
2611 + thread: Some(thread),
2612 + }
2613 + }
2614 +
2615 + fn stop(&mut self) {
2616 + self.stop_flag.store(false, Ordering::Release);
2617 + if let Some(t) = self.thread.take() {
2618 + let _ = t.join();
2619 + }
2620 + }
2621 +}
2622 +
2623 +impl Drop for StressTestServer {
2624 + fn drop(&mut self) {
2625 + self.stop();
2626 + }
2627 +}
2628 +
2629 +#[test]
2630 +fn test_stress_1000_items() {
2631 + let svc = "rs_stress_1k";
2632 +
2633 + const N: u32 = 1000;
2634 + const BUF_SIZE: usize = 300 * N as usize;
2635 +
2636 + let mut server = StressTestServer::start(svc, N, BUF_SIZE);
2637 +
2638 + let mut cfg = client_config();
2639 + cfg.max_response_payload_bytes = BUF_SIZE as u32;
2640 + cfg.packet_size = 65536;
2641 +
2642 + let mut client = snapshot_client(svc, cfg);
2643 + client.refresh();
2644 + assert!(client.ready(), "client not ready");
2645 +
2646 + let start = std::time::Instant::now();
2647 + let view = client.call_snapshot().expect("call should succeed");
2648 + let elapsed = start.elapsed();
2649 +
2650 + eprintln!(" 1000 items: {:?}", elapsed);
2651 +
2652 + assert_eq!(view.item_count, N);
2653 + assert_eq!(view.systemd_enabled, 1);
2654 + assert_eq!(view.generation, 42);
2655 +
2656 + // Verify ALL items
2657 + for i in 0..N {
2658 + let item = view
2659 + .item(i)
2660 + .unwrap_or_else(|_| panic!("item {i} decode failed"));
2661 + let expected_name = format!("container-{i:04}");
2662 + let expected_path = format!("/sys/fs/cgroup/docker/{i:04}");
2663 + let expected_hash = simple_hash(&expected_name);
2664 + let expected_enabled = if i % 5 == 0 { 0 } else { 1 };
2665 +
2666 + assert_eq!(item.hash, expected_hash, "item {i} hash mismatch");
2667 + assert_eq!(
2668 + std::str::from_utf8(item.name.as_bytes()).unwrap(),
2669 + expected_name,
2670 + "item {i} name mismatch"
2671 + );
2672 + assert_eq!(
2673 + std::str::from_utf8(item.path.as_bytes()).unwrap(),
2674 + expected_path,
2675 + "item {i} path mismatch"
2676 + );
2677 + assert_eq!(item.enabled, expected_enabled, "item {i} enabled mismatch");
2678 + assert_eq!(item.options, 0x10, "item {i} options mismatch");
2679 + }
2680 +
2681 + client.close();
2682 + server.stop();
2683 + cleanup_all(svc);
2684 +}
2685 +
2686 +#[test]
2687 +fn test_stress_5000_items() {
2688 + let svc = "rs_stress_5k";
2689 +
2690 + const N: u32 = 5000;
2691 + const BUF_SIZE: usize = 300 * N as usize;
2692 +
2693 + let mut server = StressTestServer::start(svc, N, BUF_SIZE);
2694 +
2695 + let mut cfg = client_config();
2696 + cfg.max_response_payload_bytes = BUF_SIZE as u32;
2697 + cfg.packet_size = 65536;
2698 +
2699 + let mut client = snapshot_client(svc, cfg);
2700 + client.refresh();
2701 + assert!(client.ready(), "client not ready");
2702 +
2703 + let start = std::time::Instant::now();
2704 + let view = client.call_snapshot().expect("call should succeed");
2705 + let elapsed = start.elapsed();
2706 +
2707 + eprintln!(" 5000 items: {:?}", elapsed);
2708 +
2709 + assert_eq!(view.item_count, N);
2710 +
2711 + // Spot-check first, middle, last
2712 + for idx in [0, N / 2, N - 1] {
2713 + let item = view.item(idx).unwrap();
2714 + let expected_name = format!("container-{idx:04}");
2715 + let expected_hash = simple_hash(&expected_name);
2716 + assert_eq!(item.hash, expected_hash);
2717 + assert_eq!(
2718 + std::str::from_utf8(item.name.as_bytes()).unwrap(),
2719 + expected_name
2720 + );
2721 + }
2722 +
2723 + client.close();
2724 + server.stop();
2725 + cleanup_all(svc);
2726 +}
2727 +
2728 +#[test]
2729 +fn test_stress_concurrent_clients() {
2730 + let svc = "rs_stress_concurrent";
2731 + ensure_run_dir();
2732 + cleanup_all(svc);
2733 +
2734 + let mut server = TestServer::start_with_workers(
2735 + svc,
2736 + METHOD_CGROUPS_SNAPSHOT,
2737 + Some(test_cgroups_dispatch()),
2738 + 64,
2739 + );
2740 +
2741 + const NUM_CLIENTS: usize = 50;
2742 + const REQUESTS_PER: usize = 10;
2743 +
2744 + let start = std::time::Instant::now();
2745 +
2746 + let mut handles = Vec::new();
2747 + for client_id in 0..NUM_CLIENTS {
2748 + let svc_name = svc.to_string();
2749 + let handle = thread::spawn(move || {
2750 + let mut client = snapshot_client(&svc_name, client_config());
2751 +
2752 + for _ in 0..200 {
2753 + client.refresh();
2754 + if client.ready() {
2755 + break;
2756 + }
2757 + thread::sleep(Duration::from_millis(5));
2758 + }
2759 +
2760 + assert!(client.ready(), "client {client_id} not ready");
2761 +
2762 + let mut successes = 0usize;
2763 + for _ in 0..REQUESTS_PER {
2764 + match client.call_snapshot() {
2765 + Ok(view) => {
2766 + assert_eq!(view.item_count, 3);
2767 + assert_eq!(view.generation, 42);
2768 + let item0 = view.item(0).expect("item 0");
2769 + assert_eq!(item0.hash, 1001);
2770 + assert_eq!(
2771 + std::str::from_utf8(item0.name.as_bytes()).unwrap(),
2772 + "docker-abc123"
2773 + );
2774 + let item2 = view.item(2).expect("item 2");
2775 + assert_eq!(item2.hash, 3003);
2776 + successes += 1;
2777 + }
2778 + Err(e) => panic!("client {client_id} call failed: {:?}", e),
2779 + }
2780 + }
2781 + client.close();
2782 + successes
2783 + });
2784 + handles.push(handle);
2785 + }
2786 +
2787 + let mut total = 0usize;
2788 + for h in handles {
2789 + total += h.join().expect("client thread panicked");
2790 + }
2791 +
2792 + let elapsed = start.elapsed();
2793 + eprintln!(
2794 + " {NUM_CLIENTS} clients x {REQUESTS_PER} req: {total}/{} in {:?}",
2795 + NUM_CLIENTS * REQUESTS_PER,
2796 + elapsed
2797 + );
2798 +
2799 + assert_eq!(total, NUM_CLIENTS * REQUESTS_PER);
2800 +
2801 + server.stop();
2802 + cleanup_all(svc);
2803 +}
2804 +
2805 +#[test]
2806 +fn test_stress_rapid_connect_disconnect() {
2807 + let svc = "rs_stress_rapid";
2808 + ensure_run_dir();
2809 + cleanup_all(svc);
2810 +
2811 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
2812 +
2813 + const CYCLES: usize = 1000;
2814 + let mut successes = 0usize;
2815 + let mut failures = 0usize;
2816 +
2817 + let start = std::time::Instant::now();
2818 +
2819 + for _ in 0..CYCLES {
2820 + let mut client = snapshot_client(svc, client_config());
2821 +
2822 + let mut connected = false;
2823 + /* Under full ctest -j load, a freshly spawned server may need a
2824 + * slightly longer connect window than the hot-path unit tests. */
2825 + for _ in 0..200 {
2826 + client.refresh();
2827 + if client.ready() {
2828 + connected = true;
2829 + break;
2830 + }
2831 + thread::sleep(Duration::from_millis(2));
2832 + }
2833 +
2834 + if !connected {
2835 + failures += 1;
2836 + client.close();
2837 + continue;
2838 + }
2839 +
2840 + match client.call_snapshot() {
2841 + Ok(view) => {
2842 + if view.item_count == 3 && view.generation == 42 {
2843 + successes += 1;
2844 + } else {
2845 + failures += 1;
2846 + }
2847 + }
2848 + Err(_) => failures += 1,
2849 + }
2850 +
2851 + client.close();
2852 + }
2853 +
2854 + let elapsed = start.elapsed();
2855 + eprintln!(
2856 + " {CYCLES} rapid cycles: {successes} ok, {failures} fail, {:?}",
2857 + elapsed
2858 + );
2859 +
2860 + assert_eq!(successes, CYCLES, "all cycles should succeed");
2861 + assert_eq!(failures, 0, "no failures expected");
2862 +
2863 + server.stop();
2864 + cleanup_all(svc);
2865 +}
2866 +
2867 +#[test]
2868 +fn test_stress_cache_concurrent() {
2869 + let svc = "rs_stress_cache";
2870 + ensure_run_dir();
2871 + cleanup_all(svc);
2872 +
2873 + let mut server = TestServer::start_with_workers(
2874 + svc,
2875 + METHOD_CGROUPS_SNAPSHOT,
2876 + Some(test_cgroups_dispatch()),
2877 + 16,
2878 + );
2879 +
2880 + const NUM_CLIENTS: usize = 10;
2881 + const REQUESTS_PER: usize = 100;
2882 +
2883 + let start = std::time::Instant::now();
2884 +
2885 + let mut handles = Vec::new();
2886 + for _ in 0..NUM_CLIENTS {
2887 + let svc_name = svc.to_string();
2888 + let handle = thread::spawn(move || {
2889 + let mut cache = CgroupsCache::new(TEST_RUN_DIR, &svc_name, client_config());
2890 + let mut successes = 0usize;
2891 +
2892 + for _ in 0..REQUESTS_PER {
2893 + let updated = cache.refresh();
2894 + if updated || cache.ready() {
2895 + let status = cache.status();
2896 + if status.item_count != 3 {
2897 + continue;
2898 + }
2899 + let item = cache.lookup(1001, "docker-abc123");
2900 + if item.is_some() && item.unwrap().hash == 1001 {
2901 + successes += 1;
2902 + }
2903 + }
2904 + }
2905 + cache.close();
2906 + successes
2907 + });
2908 + handles.push(handle);
2909 + }
2910 +
2911 + let mut total = 0usize;
2912 + for h in handles {
2913 + total += h.join().expect("cache thread panicked");
2914 + }
2915 +
2916 + let elapsed = start.elapsed();
2917 + eprintln!(
2918 + " {NUM_CLIENTS} cache clients x {REQUESTS_PER} req: {total}/{} in {:?}",
2919 + NUM_CLIENTS * REQUESTS_PER,
2920 + elapsed
2921 + );
2922 +
2923 + assert_eq!(total, NUM_CLIENTS * REQUESTS_PER);
2924 +
2925 + server.stop();
2926 + cleanup_all(svc);
2927 +}
2928 +
2929 +#[test]
2930 +fn test_stress_long_running() {
2931 + let svc = "rs_stress_long";
2932 + ensure_run_dir();
2933 + cleanup_all(svc);
2934 +
2935 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
2936 +
2937 + const NUM_CLIENTS: usize = 5;
2938 + let run_duration = Duration::from_secs(30);
2939 +
2940 + let running = Arc::new(AtomicBool::new(true));
2941 + let mut handles = Vec::new();
2942 +
2943 + for _ in 0..NUM_CLIENTS {
2944 + let svc_name = svc.to_string();
2945 + let r = running.clone();
2946 + let handle = thread::spawn(move || {
2947 + let mut cache = CgroupsCache::new(TEST_RUN_DIR, &svc_name, client_config());
2948 + let mut refreshes = 0u64;
2949 + let mut errors = 0u64;
2950 +
2951 + while r.load(Ordering::Acquire) {
2952 + let updated = cache.refresh();
2953 + if updated || cache.ready() {
2954 + let status = cache.status();
2955 + if status.item_count == 3 {
2956 + refreshes += 1;
2957 + } else {
2958 + errors += 1;
2959 + }
2960 + } else {
2961 + errors += 1;
2962 + }
2963 + thread::sleep(Duration::from_millis(1));
2964 + }
2965 +
2966 + cache.close();
2967 + (refreshes, errors)
2968 + });
2969 + handles.push(handle);
2970 + }
2971 +
2972 + thread::sleep(run_duration);
2973 + running.store(false, Ordering::Release);
2974 +
2975 + let mut total_refreshes = 0u64;
2976 + let mut total_errors = 0u64;
2977 + for h in handles {
2978 + let (r, e) = h.join().expect("client thread panicked");
2979 + total_refreshes += r;
2980 + total_errors += e;
2981 + }
2982 +
2983 + eprintln!(" 30s run: {total_refreshes} refreshes, {total_errors} errors");
2984 +
2985 + assert!(total_refreshes > 0, "expected some refreshes");
2986 + assert_eq!(total_errors, 0, "expected zero errors in 60s run");
2987 +
2988 + server.stop();
2989 + cleanup_all(svc);
2990 +}
2991 +
2992 +// ---------------------------------------------------------------
2993 +// Ping-pong tests per service kind
2994 +// ---------------------------------------------------------------
2995 +
2996 +#[test]
2997 +fn test_increment_ping_pong() {
2998 + let svc = "rs_pp_incr";
2999 + ensure_run_dir();
3000 + cleanup_all(svc);
3001 +
3002 + let mut server = TestServer::start(svc, METHOD_INCREMENT, Some(increment_dispatch_handler()));
3003 +
3004 + let mut client = increment_client(svc, client_config());
3005 + client.refresh();
3006 + assert!(client.ready(), "client not ready");
3007 +
3008 + // Ping-pong: send 0 -> get 1 -> send 1 -> get 2 -> ... -> 10
3009 + let mut value = 0u64;
3010 + let mut responses_received = 0u64;
3011 + for round in 0..10 {
3012 + let sent = value;
3013 + let result = client
3014 + .call_increment(sent)
3015 + .unwrap_or_else(|e| panic!("round {round}: call_increment({sent}) failed: {e:?}"));
3016 + assert_eq!(
3017 + result,
3018 + sent + 1,
3019 + "round {round}: expected {} got {result}",
3020 + sent + 1
3021 + );
3022 + responses_received += 1;
3023 + value = result;
3024 + }
3025 + assert_eq!(
3026 + responses_received, 10,
3027 + "expected 10 responses, got {responses_received}"
3028 + );
3029 + assert_eq!(value, 10, "final value after 10 rounds");
3030 +
3031 + client.close();
3032 + server.stop();
3033 + cleanup_all(svc);
3034 +}
3035 +
3036 +#[test]
3037 +fn test_string_reverse_ping_pong() {
3038 + let svc = "rs_pp_strrev";
3039 + ensure_run_dir();
3040 + cleanup_all(svc);
3041 +
3042 + let mut server = TestServer::start(
3043 + svc,
3044 + METHOD_STRING_REVERSE,
3045 + Some(string_reverse_dispatch_handler()),
3046 + );
3047 +
3048 + let mut client = string_reverse_client(svc, client_config());
3049 + client.refresh();
3050 + assert!(client.ready(), "client not ready");
3051 +
3052 + let original = "abcdefghijklmnopqrstuvwxyz";
3053 + let mut current = original.to_string();
3054 + let mut responses_received = 0u64;
3055 +
3056 + // 6 rounds: feed each response back as next request
3057 + for round in 0..6 {
3058 + let sent = current.clone();
3059 + let expected: String = sent.chars().rev().collect();
3060 + let result = client.call_string_reverse(&sent).unwrap_or_else(|e| {
3061 + panic!("round {round}: call_string_reverse({sent:?}) failed: {e:?}")
3062 + });
3063 + assert_eq!(
3064 + result.as_str(),
3065 + expected,
3066 + "round {round}: reverse of {sent:?} should be {expected:?}, got {result:?}"
3067 + );
3068 + responses_received += 1;
3069 + current = result.as_str().to_string();
3070 + }
3071 + assert_eq!(
3072 + responses_received, 6,
3073 + "expected 6 responses, got {responses_received}"
3074 + );
3075 + // even number of reversals = identity
3076 + assert_eq!(
3077 + current, original,
3078 + "6 reversals should restore original string"
3079 + );
3080 +
3081 + client.close();
3082 + server.stop();
3083 + cleanup_all(svc);
3084 +}
3085 +
3086 +#[test]
3087 +fn test_increment_batch() {
3088 + let svc = "rs_pp_batch";
3089 + ensure_run_dir();
3090 + cleanup_all(svc);
3091 +
3092 + // Need batch items > 1 for both client and server configs
3093 + fn batch_server_config() -> ServerConfig {
3094 + ServerConfig {
3095 + supported_profiles: PROFILE_BASELINE,
3096 + max_request_payload_bytes: 4096,
3097 + max_request_batch_items: 16,
3098 + max_response_payload_bytes: 4096,
3099 + max_response_batch_items: 16,
3100 + auth_token: AUTH_TOKEN,
3101 + backlog: 4,
3102 + ..ServerConfig::default()
3103 + }
3104 + }
3105 +
3106 + fn batch_client_config() -> ClientConfig {
3107 + ClientConfig {
3108 + supported_profiles: PROFILE_BASELINE,
3109 + max_request_payload_bytes: 4096,
3110 + max_request_batch_items: 16,
3111 + max_response_payload_bytes: 4096,
3112 + max_response_batch_items: 16,
3113 + auth_token: AUTH_TOKEN,
3114 + ..ClientConfig::default()
3115 + }
3116 + }
3117 +
3118 + // Start server with batch-capable config
3119 + let svc_name = svc.to_string();
3120 + let ready_flag = Arc::new(AtomicBool::new(false));
3121 + let ready_clone = ready_flag.clone();
3122 +
3123 + let mut server_obj = ManagedServer::with_workers(
3124 + TEST_RUN_DIR,
3125 + &svc_name,
3126 + batch_server_config(),
3127 + METHOD_INCREMENT,
3128 + Some(increment_dispatch_handler()),
3129 + 8,
3130 + );
3131 + let stop_flag = server_obj.running_flag();
3132 +
3133 + let thread_handle = thread::spawn(move || {
3134 + ready_clone.store(true, Ordering::Release);
3135 + let _ = server_obj.run();
3136 + });
3137 +
3138 + for _ in 0..2000 {
3139 + if ready_flag.load(Ordering::Acquire) {
3140 + break;
3141 + }
3142 + thread::sleep(Duration::from_micros(500));
3143 + }
3144 + thread::sleep(Duration::from_millis(50));
3145 +
3146 + let mut client = increment_client(svc, batch_client_config());
3147 + client.refresh();
3148 + assert!(client.ready(), "client not ready");
3149 +
3150 + // Send batch of [10, 20, 30, 40, 50]
3151 + let values = vec![10u64, 20, 30, 40, 50];
3152 + let results = client.call_increment_batch(&values).expect("batch call");
3153 +
3154 + assert_eq!(results.len(), 5);
3155 + for (i, (&input, &output)) in values.iter().zip(results.iter()).enumerate() {
3156 + assert_eq!(
3157 + output,
3158 + input + 1,
3159 + "batch item {i}: expected {}, got {output}",
3160 + input + 1
3161 + );
3162 + }
3163 +
3164 + // Single item batch
3165 + let single = client
3166 + .call_increment_batch(&[99])
3167 + .expect("single-item batch");
3168 + assert_eq!(single, vec![100]);
3169 +
3170 + // Empty batch
3171 + let empty = client.call_increment_batch(&[]).expect("empty batch");
3172 + assert!(empty.is_empty());
3173 +
3174 + client.close();
3175 + stop_flag.store(false, Ordering::Release);
3176 + let _ = thread_handle.join();
3177 + cleanup_all(svc);
3178 +}
3179 +
3180 +// ---------------------------------------------------------------
3181 +// Client state machine: auth failure (lines 438-440)
3182 +// ---------------------------------------------------------------
3183 +
3184 +#[test]
3185 +fn test_client_auth_failure() {
3186 + let svc = "rs_svc_authfail";
3187 + ensure_run_dir();
3188 + cleanup_all(svc);
3189 +
3190 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
3191 +
3192 + // Client with wrong auth token
3193 + let mut bad_cfg = client_config();
3194 + bad_cfg.auth_token = 0xBAD_BAD_BAD;
3195 +
3196 + let mut client = snapshot_client(svc, bad_cfg);
3197 + client.refresh();
3198 + assert_eq!(client.state, ClientState::AuthFailed);
3199 + assert!(!client.ready());
3200 +
3201 + // Subsequent refresh stays stuck in AuthFailed
3202 + client.refresh();
3203 + assert_eq!(client.state, ClientState::AuthFailed);
3204 +
3205 + client.close();
3206 + server.stop();
3207 + cleanup_all(svc);
3208 +}
3209 +
3210 +// ---------------------------------------------------------------
3211 +// Client state machine: incompatible (lines 439-440)
3212 +// ---------------------------------------------------------------
3213 +
3214 +#[test]
3215 +fn test_client_incompatible() {
3216 + let svc = "rs_svc_incompat";
3217 + ensure_run_dir();
3218 + cleanup_all(svc);
3219 +
3220 + // Server supports only PROFILE_BASELINE, but start it first
3221 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
3222 +
3223 + // Client requires SHM_FUTEX only (no baseline)
3224 + let mut bad_cfg = client_config();
3225 + #[cfg(target_os = "linux")]
3226 + {
3227 + bad_cfg.supported_profiles = crate::protocol::PROFILE_SHM_FUTEX;
3228 + }
3229 + #[cfg(not(target_os = "linux"))]
3230 + {
3231 + // On non-Linux, use a profile bit that won't match
3232 + bad_cfg.supported_profiles = 0x80000000;
3233 + }
3234 +
3235 + let mut client = snapshot_client(svc, bad_cfg);
3236 + client.refresh();
3237 + assert_eq!(client.state, ClientState::Incompatible);
3238 + assert!(!client.ready());
3239 +
3240 + // Stays stuck
3241 + client.refresh();
3242 + assert_eq!(client.state, ClientState::Incompatible);
3243 +
3244 + client.close();
3245 + server.stop();
3246 + cleanup_all(svc);
3247 +}
3248 +
3249 +#[test]
3250 +fn test_client_protocol_version_incompatible() {
3251 + let svc = unique_service("rs_svc_proto_incompat");
3252 + ensure_run_dir();
3253 + cleanup_all(&svc);
3254 +
3255 + let packet =
3256 + hello_ack_packet_with_version(crate::protocol::VERSION + 1, crate::protocol::STATUS_OK, 1);
3257 + let mut server = start_raw_hello_ack_server(&svc, packet);
3258 +
3259 + let mut client = snapshot_client(&svc, client_config());
3260 + let changed = client.refresh();
3261 + assert!(changed, "refresh should move client into INCOMPATIBLE");
3262 + assert_eq!(client.state, ClientState::Incompatible);
3263 + assert!(!client.ready());
3264 +
3265 + server.wait();
3266 +
3267 + let changed = client.refresh();
3268 + assert!(!changed, "refresh from incompatible should be a no-op");
3269 + assert_eq!(client.state, ClientState::Incompatible);
3270 +
3271 + client.close();
3272 + cleanup_all(&svc);
3273 +}
3274 +
3275 +// ---------------------------------------------------------------
3276 +// Client: call_snapshot when not ready (line 324-326)
3277 +// ---------------------------------------------------------------
3278 +
3279 +#[test]
3280 +fn test_call_when_not_ready() {
3281 + let svc = "rs_svc_noready";
3282 + ensure_run_dir();
3283 + cleanup_all(svc);
3284 +
3285 + let mut snapshot = snapshot_client(svc, client_config());
3286 + assert_eq!(snapshot.state, ClientState::Disconnected);
3287 + assert!(snapshot.call_snapshot().is_err());
3288 + assert_eq!(snapshot.status().error_count, 1);
3289 + snapshot.close();
3290 +
3291 + let mut increment = increment_client(svc, client_config());
3292 + assert_eq!(increment.state, ClientState::Disconnected);
3293 + assert!(increment.call_increment(42).is_err());
3294 + assert!(increment.call_increment_batch(&[1, 2]).is_err());
3295 + assert_eq!(increment.status().error_count, 2);
3296 + increment.close();
3297 +
3298 + let mut string_reverse = string_reverse_client(svc, client_config());
3299 + assert_eq!(string_reverse.state, ClientState::Disconnected);
3300 + assert!(string_reverse.call_string_reverse("test").is_err());
3301 + assert_eq!(string_reverse.status().error_count, 1);
3302 + string_reverse.close();
3303 +
3304 + cleanup_all(svc);
3305 +}
3306 +
3307 +#[test]
3308 +fn test_client_invalid_service_name_maps_to_disconnected() {
3309 + let bad_service = "x".repeat(400);
3310 + let mut client = snapshot_client(&bad_service, client_config());
3311 +
3312 + client.refresh();
3313 +
3314 + assert_eq!(client.state, ClientState::Disconnected);
3315 + assert!(!client.ready());
3316 +}
3317 +
3318 +// ---------------------------------------------------------------
3319 +// Client: broken -> reconnect cycle (lines 136-141)
3320 +// ---------------------------------------------------------------
3321 +
3322 +#[test]
3323 +fn test_broken_reconnect() {
3324 + let svc = "rs_svc_broken";
3325 + ensure_run_dir();
3326 + cleanup_all(svc);
3327 +
3328 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
3329 +
3330 + let mut client = snapshot_client(svc, client_config());
3331 + client.refresh();
3332 + assert_eq!(client.state, ClientState::Ready);
3333 +
3334 + // Force broken state
3335 + client.state = ClientState::Broken;
3336 +
3337 + // refresh from Broken should disconnect, reconnect
3338 + let changed = client.refresh();
3339 + assert!(changed);
3340 + assert_eq!(client.state, ClientState::Ready);
3341 + assert!(client.status().reconnect_count >= 1);
3342 +
3343 + client.close();
3344 + server.stop();
3345 + cleanup_all(svc);
3346 +}
3347 +
3348 +// ---------------------------------------------------------------
3349 +// Cache: close resets everything (line 1561-1566)
3350 +// ---------------------------------------------------------------
3351 +
3352 +#[test]
3353 +fn test_cache_close_resets() {
3354 + let svc = "rs_cache_close";
3355 + ensure_run_dir();
3356 + cleanup_all(svc);
3357 +
3358 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, Some(test_cgroups_dispatch()));
3359 +
3360 + let mut cache = CgroupsCache::new(TEST_RUN_DIR, svc, client_config());
3361 + assert!(cache.refresh());
3362 + assert!(cache.ready());
3363 + assert_eq!(cache.status().item_count, 3);
3364 +
3365 + cache.close();
3366 + assert!(!cache.ready());
3367 + assert!(cache.lookup(1001, "docker-abc123").is_none());
3368 + assert_eq!(cache.status().item_count, 0);
3369 +
3370 + server.stop();
3371 + cleanup_all(svc);
3372 +}
3373 +
3374 +// ---------------------------------------------------------------
3375 +// Cache: with max_response_payload_bytes = 0 (line 1437-1440)
3376 +// ---------------------------------------------------------------
3377 +
3378 +#[test]
3379 +fn test_cache_default_buf_size() {
3380 + let svc = "rs_cache_defbuf";
3381 + ensure_run_dir();
3382 + cleanup_all(svc);
3383 +
3384 + // Config with max_response_payload_bytes = 0 triggers default buf
3385 + let mut cfg = client_config();
3386 + cfg.max_response_payload_bytes = 0;
3387 +
3388 + let cache = CgroupsCache::new(TEST_RUN_DIR, svc, cfg);
3389 + assert_eq!(
3390 + cache.client.max_receive_message_bytes(),
3391 + HEADER_SIZE + CACHE_RESPONSE_BUF_SIZE
3392 + );
3393 +
3394 + cleanup_all(svc);
3395 +}
3396 +
3397 +// ---------------------------------------------------------------
3398 +// ManagedServer: worker_count = 0 -> clamped to 1 (line 778)
3399 +// ---------------------------------------------------------------
3400 +
3401 +#[test]
3402 +fn test_server_worker_count_clamped() {
3403 + let svc = "rs_svc_w0";
3404 + ensure_run_dir();
3405 + cleanup_all(svc);
3406 +
3407 + let server = ManagedServer::with_workers(
3408 + TEST_RUN_DIR,
3409 + svc,
3410 + server_config(),
3411 + METHOD_INCREMENT,
3412 + None,
3413 + 0,
3414 + );
3415 + assert_eq!(server.worker_count, 1);
3416 +
3417 + cleanup_all(svc);
3418 +}
3419 +
3420 +// ---------------------------------------------------------------
3421 +// ManagedServer: stop flag (line 946)
3422 +// ---------------------------------------------------------------
3423 +
3424 +#[test]
3425 +fn test_server_stop_flag() {
3426 + let svc = "rs_svc_stopflag";
3427 + ensure_run_dir();
3428 + cleanup_all(svc);
3429 +
3430 + let server = ManagedServer::new(TEST_RUN_DIR, svc, server_config(), METHOD_INCREMENT, None);
3431 + let flag = server.running_flag();
3432 + assert!(!flag.load(Ordering::Acquire));
3433 +
3434 + // stop sets running to false
3435 + server.stop();
3436 + assert!(!flag.load(Ordering::Acquire));
3437 +
3438 + cleanup_all(svc);
3439 +}
3440 +
3441 +// ---------------------------------------------------------------
3442 +// ClientStatus / CgroupsCacheStatus fields
3443 +// ---------------------------------------------------------------
3444 +
3445 +#[test]
3446 +fn test_client_status_fields() {
3447 + let svc = "rs_svc_csf";
3448 + ensure_run_dir();
3449 + cleanup_all(svc);
3450 +
3451 + let client = snapshot_client(svc, client_config());
3452 + let status = client.status();
3453 + assert_eq!(status.state, ClientState::Disconnected);
3454 + assert_eq!(status.connect_count, 0);
3455 + assert_eq!(status.reconnect_count, 0);
3456 + assert_eq!(status.call_count, 0);
3457 + assert_eq!(status.error_count, 0);
3458 + cleanup_all(svc);
3459 +}
3460 +
3461 +// ---------------------------------------------------------------
3462 +// call_increment and call_string_reverse success paths
3463 +// ---------------------------------------------------------------
3464 +
3465 +#[test]
3466 +fn test_client_call_increment_success() {
3467 + let svc = "rs_svc_incr_ok";
3468 + ensure_run_dir();
3469 + cleanup_all(svc);
3470 +
3471 + let mut server = TestServer::start(svc, METHOD_INCREMENT, Some(increment_dispatch_handler()));
3472 +
3473 + let mut client = increment_client(svc, client_config());
3474 + client.refresh();
3475 + assert!(client.ready());
3476 +
3477 + let result = client.call_increment(99).expect("increment");
3478 + assert_eq!(result, 100);
3479 +
3480 + client.close();
3481 + server.stop();
3482 + cleanup_all(svc);
3483 +}
3484 +
3485 +#[test]
3486 +fn test_client_call_string_reverse_success() {
3487 + let svc = "rs_svc_strrev_ok";
3488 + ensure_run_dir();
3489 + cleanup_all(svc);
3490 +
3491 + let mut server = TestServer::start(
3492 + svc,
3493 + METHOD_STRING_REVERSE,
3494 + Some(string_reverse_dispatch_handler()),
3495 + );
3496 +
3497 + let mut client = string_reverse_client(svc, client_config());
3498 + client.refresh();
3499 + assert!(client.ready());
3500 +
3501 + let result = client.call_string_reverse("hello").expect("reverse");
3502 + assert_eq!(result.as_str(), "olleh");
3503 +
3504 + client.close();
3505 + server.stop();
3506 + cleanup_all(svc);
3507 +}
3508 +
3509 +#[test]
3510 +fn test_dispatch_single_helper_paths() {
3511 + let mut response_buf = [0u8; 128];
3512 +
3513 + assert!(matches!(
3514 + dispatch_single(
3515 + METHOD_INCREMENT,
3516 + None,
3517 + METHOD_INCREMENT,
3518 + &[0; 8],
3519 + &mut response_buf,
3520 + ),
3521 + Err(_)
3522 + ));
3523 +
3524 + assert!(matches!(
3525 + dispatch_single(
3526 + METHOD_STRING_REVERSE,
3527 + None,
3528 + METHOD_STRING_REVERSE,
3529 + &[0; 8],
3530 + &mut response_buf,
3531 + ),
3532 + Err(_)
3533 + ));
3534 +
3535 + assert!(matches!(
3536 + dispatch_single(
3537 + METHOD_CGROUPS_SNAPSHOT,
3538 + None,
3539 + METHOD_CGROUPS_SNAPSHOT,
3540 + &[1, 0, 0, 0],
3541 + &mut response_buf,
3542 + ),
3543 + Err(_)
3544 + ));
3545 +
3546 + let snapshot_handler = snapshot_dispatch(Arc::new(|_, _| true), 0);
3547 + assert!(matches!(
3548 + dispatch_single(
3549 + METHOD_CGROUPS_SNAPSHOT,
3550 + Some(&snapshot_handler),
3551 + METHOD_CGROUPS_SNAPSHOT,
3552 + &[1, 0, 0, 0],
3553 + &mut [],
3554 + ),
3555 + Err(_)
3556 + ));
3557 +
3558 + let reverse_handler = string_reverse_dispatch_handler();
3559 + let mut invalid_utf8_req = [0u8; 16];
3560 + let invalid_len = string_reverse_encode(&[0xff], &mut invalid_utf8_req);
3561 + let n = dispatch_single(
3562 + METHOD_STRING_REVERSE,
3563 + Some(&reverse_handler),
3564 + METHOD_STRING_REVERSE,
3565 + &invalid_utf8_req[..invalid_len],
3566 + &mut response_buf,
3567 + )
3568 + .expect("non-UTF8 string_reverse input should decode as an empty string");
3569 + let view =
3570 + string_reverse_decode(&response_buf[..n]).expect("decode empty-string reverse response");
3571 + assert_eq!(view.as_str(), "");
3572 +
3573 + assert!(matches!(
3574 + dispatch_single(
3575 + METHOD_STRING_REVERSE,
3576 + Some(&reverse_handler),
3577 + METHOD_INCREMENT,
3578 + &invalid_utf8_req[..invalid_len],
3579 + &mut response_buf,
3580 + ),
3581 + Err(_)
3582 + ));
3583 +
3584 + let snapshot_fail_handler = snapshot_dispatch(Arc::new(|_, _| false), 1);
3585 + assert!(matches!(
3586 + dispatch_single(
3587 + METHOD_CGROUPS_SNAPSHOT,
3588 + Some(&snapshot_fail_handler),
3589 + METHOD_CGROUPS_SNAPSHOT,
3590 + &[1, 0, 0, 0],
3591 + &mut response_buf,
3592 + ),
3593 + Err(_)
3594 + ));
3595 +
3596 + assert!(matches!(
3597 + dispatch_single(0xFFFF, None, 0xFFFF, &[], &mut response_buf),
3598 + Err(_)
3599 + ));
3600 +
3601 + assert!(
3602 + snapshot_max_items(4096, 0) > 0,
3603 + "default snapshot item estimate should be positive for a non-empty buffer"
3604 + );
3605 +}
3606 +
3607 +#[test]
3608 +fn test_response_payload_transport_buf_bounds() {
3609 + let mut client = snapshot_client("rs_payload_bounds", client_config());
3610 + client.transport_buf.resize(HEADER_SIZE + 8, 0);
3611 +
3612 + let response = ClientResponseRef {
3613 + source: ClientResponseSource::TransportBuf,
3614 + len: 16,
3615 + };
3616 +
3617 + assert_eq!(client.response_payload(response), Err(NipcError::Truncated));
3618 +}
3619 +
3620 +#[test]
3621 +fn test_call_increment_rejects_malformed_response_envelope_unix() {
3622 + struct Case {
3623 + name: &'static str,
3624 + kind: u16,
3625 + code: u16,
3626 + status: u16,
3627 + message_id_delta: u64,
3628 + want: NipcError,
3629 + }
3630 +
3631 + let cases = [
3632 + Case {
3633 + name: "bad kind",
3634 + kind: KIND_REQUEST,
3635 + code: METHOD_INCREMENT,
3636 + status: STATUS_OK,
3637 + message_id_delta: 0,
3638 + want: NipcError::BadKind,
3639 + },
3640 + Case {
3641 + name: "bad code",
3642 + kind: KIND_RESPONSE,
3643 + code: METHOD_STRING_REVERSE,
3644 + status: STATUS_OK,
3645 + message_id_delta: 0,
3646 + want: NipcError::BadLayout,
3647 + },
3648 + Case {
3649 + name: "bad status",
3650 + kind: KIND_RESPONSE,
3651 + code: METHOD_INCREMENT,
3652 + status: STATUS_INTERNAL_ERROR,
3653 + message_id_delta: 0,
3654 + want: NipcError::BadLayout,
3655 + },
3656 + Case {
3657 + name: "bad message id",
3658 + kind: KIND_RESPONSE,
3659 + code: METHOD_INCREMENT,
3660 + status: STATUS_OK,
3661 + message_id_delta: 1,
3662 + want: NipcError::Truncated,
3663 + },
3664 + ];
3665 +
3666 + for tc in cases {
3667 + let svc = format!("rs_unix_inc_env_{}", tc.name.replace(' ', "_"));
3668 + let mut server =
3669 + start_raw_session_server(&svc, server_config(), move |session, req_hdr, _| {
3670 + let mut payload = [0u8; INCREMENT_PAYLOAD_SIZE];
3671 + let n = increment_encode(43, &mut payload);
3672 + if n != INCREMENT_PAYLOAD_SIZE {
3673 + return Err(format!("increment_encode returned {n}"));
3674 + }
3675 +
3676 + let mut resp_hdr = Header {
3677 + kind: tc.kind,
3678 + code: tc.code,
3679 + flags: 0,
3680 + item_count: 1,
3681 + message_id: req_hdr.message_id + tc.message_id_delta,
3682 + transport_status: tc.status,
3683 + ..Header::default()
3684 + };
3685 + session
3686 + .send(&mut resp_hdr, &payload)
3687 + .map_err(|e| format!("send: {e}"))
3688 + });
3689 +
3690 + let mut client = increment_client(&svc, client_config());
3691 + connect_ready(&mut client);
3692 +
3693 + let err = client.call_increment(42).expect_err(tc.name);
3694 + assert_eq!(err, tc.want, "{}", tc.name);
3695 +
3696 + client.close();
3697 + server.wait();
3698 + cleanup_all(&svc);
3699 + }
3700 +}
3701 +
3702 +#[test]
3703 +fn test_call_string_reverse_rejects_malformed_response_envelope_unix() {
3704 + struct Case {
3705 + name: &'static str,
3706 + kind: u16,
3707 + code: u16,
3708 + status: u16,
3709 + message_id_delta: u64,
3710 + want: NipcError,
3711 + }
3712 +
3713 + let cases = [
3714 + Case {
3715 + name: "bad kind",
3716 + kind: KIND_REQUEST,
3717 + code: METHOD_STRING_REVERSE,
3718 + status: STATUS_OK,
3719 + message_id_delta: 0,
3720 + want: NipcError::BadKind,
3721 + },
3722 + Case {
3723 + name: "bad code",
3724 + kind: KIND_RESPONSE,
3725 + code: METHOD_INCREMENT,
3726 + status: STATUS_OK,
3727 + message_id_delta: 0,
3728 + want: NipcError::BadLayout,
3729 + },
3730 + Case {
3731 + name: "bad status",
3732 + kind: KIND_RESPONSE,
3733 + code: METHOD_STRING_REVERSE,
3734 + status: STATUS_INTERNAL_ERROR,
3735 + message_id_delta: 0,
3736 + want: NipcError::BadLayout,
3737 + },
3738 + Case {
3739 + name: "bad message id",
3740 + kind: KIND_RESPONSE,
3741 + code: METHOD_STRING_REVERSE,
3742 + status: STATUS_OK,
3743 + message_id_delta: 1,
3744 + want: NipcError::Truncated,
3745 + },
3746 + ];
3747 +
3748 + for tc in cases {
3749 + let svc = format!("rs_unix_str_env_{}", tc.name.replace(' ', "_"));
3750 + let mut server =
3751 + start_raw_session_server(&svc, server_config(), move |session, req_hdr, _| {
3752 + let mut payload = [0u8; 128];
3753 + let n = string_reverse_encode(b"olleh", &mut payload);
3754 + if n == 0 {
3755 + return Err("string_reverse_encode returned 0".into());
3756 + }
3757 +
3758 + let mut resp_hdr = Header {
3759 + kind: tc.kind,
3760 + code: tc.code,
3761 + flags: 0,
3762 + item_count: 1,
3763 + message_id: req_hdr.message_id + tc.message_id_delta,
3764 + transport_status: tc.status,
3765 + ..Header::default()
3766 + };
3767 + session
3768 + .send(&mut resp_hdr, &payload[..n])
3769 + .map_err(|e| format!("send: {e}"))
3770 + });
3771 +
3772 + let mut client = string_reverse_client(&svc, client_config());
3773 + connect_ready(&mut client);
3774 +
3775 + let err = client.call_string_reverse("hello").expect_err(tc.name);
3776 + assert_eq!(err, tc.want, "{}", tc.name);
3777 +
3778 + client.close();
3779 + server.wait();
3780 + cleanup_all(&svc);
3781 + }
3782 +}
3783 +
3784 +#[test]
3785 +fn test_call_increment_batch_rejects_wrong_item_count_unix() {
3786 + let svc = "rs_unix_batch_count";
3787 + ensure_run_dir();
3788 + cleanup_all(svc);
3789 +
3790 + let mut server =
3791 + start_raw_session_server(svc, batch_server_config(), move |session, req_hdr, _| {
3792 + let mut encoded = [0u8; INCREMENT_PAYLOAD_SIZE];
3793 + let n = increment_encode(11, &mut encoded);
3794 + if n != INCREMENT_PAYLOAD_SIZE {
3795 + return Err(format!("increment_encode returned {n}"));
3796 + }
3797 +
3798 + let mut response_buf = vec![0u8; 128];
3799 + let resp_len = {
3800 + let mut batch = BatchBuilder::new(&mut response_buf, 2);
3801 + batch
3802 + .add(&encoded)
3803 + .map_err(|e| format!("batch add 1: {e:?}"))?;
3804 + batch
3805 + .add(&encoded)
3806 + .map_err(|e| format!("batch add 2: {e:?}"))?;
3807 + let (len, _count) = batch.finish();
3808 + len
3809 + };
3810 +
3811 + let mut resp_hdr = Header {
3812 + kind: KIND_RESPONSE,
3813 + code: METHOD_INCREMENT,
3814 + flags: FLAG_BATCH,
3815 + item_count: 1,
3816 + message_id: req_hdr.message_id,
3817 + transport_status: STATUS_OK,
3818 + ..Header::default()
3819 + };
3820 + session
3821 + .send(&mut resp_hdr, &response_buf[..resp_len])
3822 + .map_err(|e| format!("send: {e}"))
3823 + });
3824 +
3825 + let mut client = increment_client(svc, batch_client_config());
3826 + connect_ready(&mut client);
3827 +
3828 + let err = client
3829 + .call_increment_batch(&[10, 20])
3830 + .expect_err("wrong batch item_count");
3831 + assert_eq!(err, NipcError::BadItemCount);
3832 +
3833 + client.close();
3834 + server.wait();
3835 + cleanup_all(svc);
3836 +}
3837 +
3838 +#[test]
3839 +fn test_call_increment_batch_rejects_malformed_response_envelope_unix() {
3840 + struct Case {
3841 + name: &'static str,
3842 + kind: u16,
3843 + code: u16,
3844 + status: u16,
3845 + message_id_delta: u64,
3846 + want: NipcError,
3847 + }
3848 +
3849 + let cases = [
3850 + Case {
3851 + name: "bad kind",
3852 + kind: KIND_REQUEST,
3853 + code: METHOD_INCREMENT,
3854 + status: STATUS_OK,
3855 + message_id_delta: 0,
3856 + want: NipcError::BadKind,
3857 + },
3858 + Case {
3859 + name: "bad code",
3860 + kind: KIND_RESPONSE,
3861 + code: METHOD_STRING_REVERSE,
3862 + status: STATUS_OK,
3863 + message_id_delta: 0,
3864 + want: NipcError::BadLayout,
3865 + },
3866 + Case {
3867 + name: "bad status",
3868 + kind: KIND_RESPONSE,
3869 + code: METHOD_INCREMENT,
3870 + status: STATUS_INTERNAL_ERROR,
3871 + message_id_delta: 0,
3872 + want: NipcError::BadLayout,
3873 + },
3874 + Case {
3875 + name: "bad message id",
3876 + kind: KIND_RESPONSE,
3877 + code: METHOD_INCREMENT,
3878 + status: STATUS_OK,
3879 + message_id_delta: 1,
3880 + want: NipcError::Truncated,
3881 + },
3882 + ];
3883 +
3884 + for tc in cases {
3885 + let svc = format!("rs_unix_batch_env_{}", tc.name.replace(' ', "_"));
3886 + let mut server =
3887 + start_raw_session_server(&svc, batch_server_config(), move |session, req_hdr, _| {
3888 + let mut encoded = [0u8; INCREMENT_PAYLOAD_SIZE];
3889 + let n = increment_encode(11, &mut encoded);
3890 + if n != INCREMENT_PAYLOAD_SIZE {
3891 + return Err(format!("increment_encode returned {n}"));
3892 + }
3893 +
3894 + let mut response_buf = vec![0u8; 128];
3895 + let resp_len = {
3896 + let mut batch = BatchBuilder::new(&mut response_buf, 2);
3897 + batch
3898 + .add(&encoded)
3899 + .map_err(|e| format!("batch add 1: {e:?}"))?;
3900 + batch
3901 + .add(&encoded)
3902 + .map_err(|e| format!("batch add 2: {e:?}"))?;
3903 + let (len, _count) = batch.finish();
3904 + len
3905 + };
3906 +
3907 + let mut resp_hdr = Header {
3908 + kind: tc.kind,
3909 + code: tc.code,
3910 + flags: FLAG_BATCH,
3911 + item_count: 2,
3912 + message_id: req_hdr.message_id + tc.message_id_delta,
3913 + transport_status: tc.status,
3914 + ..Header::default()
3915 + };
3916 + session
3917 + .send(&mut resp_hdr, &response_buf[..resp_len])
3918 + .map_err(|e| format!("send: {e}"))
3919 + });
3920 +
3921 + let mut client = increment_client(&svc, batch_client_config());
3922 + connect_ready(&mut client);
3923 +
3924 + let err = client.call_increment_batch(&[10, 20]).expect_err(tc.name);
3925 + assert_eq!(err, tc.want, "{}", tc.name);
3926 +
3927 + client.close();
3928 + server.wait();
3929 + cleanup_all(&svc);
3930 + }
3931 +}
3932 +
3933 +#[test]
3934 +fn test_call_string_reverse_chunked_response_unix() {
3935 + let svc = "rs_unix_chunked_reverse";
3936 + ensure_run_dir();
3937 + cleanup_all(svc);
3938 +
3939 + let long_input = "abcdefghi".repeat(16);
3940 + let expected: String = long_input.chars().rev().collect();
3941 + let scfg = ServerConfig {
3942 + packet_size: 64,
3943 + max_response_payload_bytes: 4096,
3944 + ..server_config()
3945 + };
3946 + let ccfg = ClientConfig {
3947 + packet_size: 64,
3948 + max_response_payload_bytes: 4096,
3949 + ..client_config()
3950 + };
3951 +
3952 + let mut server = TestServer::start_with(
3953 + svc,
3954 + scfg,
3955 + METHOD_STRING_REVERSE,
3956 + Some(string_reverse_dispatch_handler()),
3957 + 8,
3958 + );
3959 + let mut client = string_reverse_client(svc, ccfg);
3960 + connect_ready(&mut client);
3961 +
3962 + let result = client
3963 + .call_string_reverse(&long_input)
3964 + .expect("chunked reverse");
3965 + assert_eq!(result.as_str(), expected);
3966 +
3967 + client.close();
3968 + server.stop();
3969 + cleanup_all(svc);
3970 +}
3971 +
3972 +// ---------------------------------------------------------------
3973 +// Batch dispatch: handler failure returns INTERNAL_ERROR (lines 1244-1248)
3974 +// ---------------------------------------------------------------
3975 +
3976 +#[test]
3977 +fn test_batch_dispatch_handler_failure() {
3978 + let svc = "rs_svc_batchfail";
3979 + ensure_run_dir();
3980 + cleanup_all(svc);
3981 +
3982 + // Handler that fails on the 2nd item
3983 + fn fail_second_increment_handler() -> IncrementHandler {
3984 + static CALL_COUNT: std::sync::atomic::AtomicU32 = std::sync::atomic::AtomicU32::new(0);
3985 + Arc::new(move |value| {
3986 + let n = CALL_COUNT.fetch_add(1, std::sync::atomic::Ordering::Relaxed);
3987 + if n % 3 == 1 {
3988 + return None;
3989 + }
3990 + Some(value + 1)
3991 + })
3992 + }
3993 +
3994 + fn batch_server_config() -> ServerConfig {
3995 + ServerConfig {
3996 + supported_profiles: PROFILE_BASELINE,
3997 + max_request_payload_bytes: 4096,
3998 + max_request_batch_items: 16,
3999 + max_response_payload_bytes: 4096,
4000 + max_response_batch_items: 16,
4001 + auth_token: AUTH_TOKEN,
4002 + backlog: 4,
4003 + ..ServerConfig::default()
4004 + }
4005 + }
4006 +
4007 + fn batch_client_config() -> ClientConfig {
4008 + ClientConfig {
4009 + supported_profiles: PROFILE_BASELINE,
4010 + max_request_payload_bytes: 4096,
4011 + max_request_batch_items: 16,
4012 + max_response_payload_bytes: 4096,
4013 + max_response_batch_items: 16,
4014 + auth_token: AUTH_TOKEN,
4015 + ..ClientConfig::default()
4016 + }
4017 + }
4018 +
4019 + let svc_name = svc.to_string();
4020 + let ready_flag = Arc::new(AtomicBool::new(false));
4021 + let ready_clone = ready_flag.clone();
4022 +
4023 + let mut server_obj = ManagedServer::with_workers(
4024 + TEST_RUN_DIR,
4025 + &svc_name,
4026 + batch_server_config(),
4027 + METHOD_INCREMENT,
4028 + Some(increment_dispatch(fail_second_increment_handler())),
4029 + 8,
4030 + );
4031 + let stop_flag = server_obj.running_flag();
4032 +
4033 + let thread_handle = thread::spawn(move || {
4034 + ready_clone.store(true, Ordering::Release);
4035 + let _ = server_obj.run();
4036 + });
4037 +
4038 + for _ in 0..2000 {
4039 + if ready_flag.load(Ordering::Acquire) {
4040 + break;
4041 + }
4042 + thread::sleep(Duration::from_micros(500));
4043 + }
4044 + thread::sleep(Duration::from_millis(50));
4045 +
4046 + let mut client = increment_client(svc, batch_client_config());
4047 + client.refresh();
4048 + assert!(client.ready());
4049 +
4050 + // Batch of 3: handler fails on the 2nd -> server returns INTERNAL_ERROR
4051 + let values = vec![10u64, 20, 30];
4052 + let result = client.call_increment_batch(&values);
4053 + // The batch should fail because the handler returned None for item 2
4054 + assert!(result.is_err());
4055 +
4056 + client.close();
4057 + stop_flag.store(false, Ordering::Release);
4058 + let _ = thread_handle.join();
4059 + cleanup_all(svc);
4060 +}
4061 +
4062 +#[test]
4063 +fn test_batch_dispatch_builder_overflow_retries_and_recovers() {
4064 + let svc = "rs_svc_batch_overflow";
4065 + ensure_run_dir();
4066 + cleanup_all(svc);
4067 +
4068 + let mut scfg = batch_server_config();
4069 + scfg.max_response_payload_bytes = 8;
4070 +
4071 + let mut server = TestServer::start_with(
4072 + svc,
4073 + scfg,
4074 + METHOD_INCREMENT,
4075 + Some(increment_dispatch_handler()),
4076 + 8,
4077 + );
4078 + let mut client = increment_client(svc, batch_client_config());
4079 + connect_ready(&mut client);
4080 +
4081 + let values = client
4082 + .call_increment_batch(&[10, 20])
4083 + .expect("batch builder overflow should transparently reconnect and retry");
4084 + assert_eq!(values, vec![11, 21]);
4085 + assert!(
4086 + client.ready(),
4087 + "client should stay READY after overflow recovery"
4088 + );
4089 + assert!(
4090 + client.status().reconnect_count >= 1,
4091 + "overflow recovery should reconnect at least once"
4092 + );
4093 +
4094 + client.close();
4095 + server.stop();
4096 + cleanup_all(svc);
4097 +}
4098 +
4099 +#[cfg(target_os = "linux")]
4100 +#[test]
4101 +fn test_shm_batch_request_item_decode_failure_returns_bad_envelope() {
4102 + let svc = "rs_svc_shm_batch_bad_item";
4103 + ensure_run_dir();
4104 + cleanup_all(svc);
4105 +
4106 + let mut server = TestServer::start_with(
4107 + svc,
4108 + shm_server_config(),
4109 + METHOD_INCREMENT,
4110 + Some(increment_dispatch_handler()),
4111 + 8,
4112 + );
4113 + let mut client = increment_client(svc, shm_client_config());
4114 + connect_ready(&mut client);
4115 + assert!(
4116 + client.shm.is_some(),
4117 + "expected SHM transport to be negotiated"
4118 + );
4119 +
4120 + let mut bad_payload = [0u8; 16];
4121 + bad_payload[0..4].copy_from_slice(&0u32.to_ne_bytes());
4122 + bad_payload[4..8].copy_from_slice(&32u32.to_ne_bytes());
4123 + bad_payload[8..12].copy_from_slice(&0u32.to_ne_bytes());
4124 + bad_payload[12..16].copy_from_slice(&4u32.to_ne_bytes());
4125 +
4126 + let req_hdr = Header {
4127 + magic: MAGIC_MSG,
4128 + version: VERSION,
4129 + header_len: protocol::HEADER_LEN,
4130 + kind: KIND_REQUEST,
4131 + code: METHOD_INCREMENT,
4132 + flags: FLAG_BATCH,
4133 + payload_len: bad_payload.len() as u32,
4134 + item_count: 2,
4135 + message_id: 7,
4136 + transport_status: STATUS_OK,
4137 + };
4138 + let mut msg = [0u8; HEADER_SIZE + 16];
4139 + req_hdr.encode(&mut msg[..HEADER_SIZE]);
4140 + msg[HEADER_SIZE..].copy_from_slice(&bad_payload);
4141 + client
4142 + .shm
4143 + .as_mut()
4144 + .expect("shm")
4145 + .send(&msg)
4146 + .expect("send malformed batch request");
4147 +
4148 + let (resp_hdr, response) = client.transport_receive().expect("receive response");
4149 + assert_eq!(resp_hdr.kind, KIND_RESPONSE);
4150 + assert_eq!(resp_hdr.code, METHOD_INCREMENT);
4151 + assert_eq!(resp_hdr.transport_status, STATUS_BAD_ENVELOPE);
4152 + assert_eq!(resp_hdr.flags, 0);
4153 + assert_eq!(resp_hdr.item_count, 1);
4154 + assert!(
4155 + client
4156 + .response_payload(response)
4157 + .expect("response payload view")
4158 + .is_empty(),
4159 + "bad-envelope response should have no payload"
4160 + );
4161 +
4162 + client.close();
4163 + server.stop();
4164 + cleanup_all(svc);
4165 +}
src/crates/netipc/src/service/raw_windows_tests.rs new
+1145
@@ -0,0 +1,1145 @@
1 +use super::*;
2 +use crate::protocol::{
3 + increment_encode, BatchBuilder, CgroupsBuilder, CgroupsRequest, Header, HelloAck, NipcError,
4 + CODE_HELLO_ACK, FLAG_BATCH, HEADER_SIZE, INCREMENT_PAYLOAD_SIZE, KIND_CONTROL, KIND_REQUEST,
5 + KIND_RESPONSE, METHOD_CGROUPS_SNAPSHOT, METHOD_INCREMENT, METHOD_STRING_REVERSE,
6 + PROFILE_BASELINE, PROFILE_SHM_HYBRID, STATUS_INTERNAL_ERROR, STATUS_OK, VERSION,
7 +};
8 +use crate::transport::windows::build_pipe_name;
9 +use std::ptr;
10 +use std::sync::atomic::{AtomicBool, AtomicU64, Ordering};
11 +use std::sync::Arc;
12 +use std::thread;
13 +use std::time::Duration;
14 +
15 +const TEST_RUN_DIR: &str = r"C:\Temp\nipc_svc_rust_test";
16 +const AUTH_TOKEN: u64 = 0xDEADBEEFCAFEBABE;
17 +const RESPONSE_BUF_SIZE: usize = 65536;
18 +static WIN_SERVICE_COUNTER: AtomicU64 = AtomicU64::new(0);
19 +
20 +fn ensure_run_dir() {
21 + let _ = std::fs::create_dir_all(TEST_RUN_DIR);
22 +}
23 +
24 +fn cleanup_all(_service: &str) {}
25 +
26 +fn unique_service(prefix: &str) -> String {
27 + format!(
28 + "{}_{}_{}",
29 + prefix,
30 + std::process::id(),
31 + WIN_SERVICE_COUNTER.fetch_add(1, Ordering::Relaxed) + 1
32 + )
33 +}
34 +
35 +fn server_config() -> ServerConfig {
36 + ServerConfig {
37 + supported_profiles: PROFILE_BASELINE,
38 + max_request_payload_bytes: 4096,
39 + max_request_batch_items: 1,
40 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
41 + max_response_batch_items: 1,
42 + auth_token: AUTH_TOKEN,
43 + ..ServerConfig::default()
44 + }
45 +}
46 +
47 +fn client_config() -> ClientConfig {
48 + ClientConfig {
49 + supported_profiles: PROFILE_BASELINE,
50 + max_request_payload_bytes: 4096,
51 + max_request_batch_items: 1,
52 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
53 + max_response_batch_items: 1,
54 + auth_token: AUTH_TOKEN,
55 + ..ClientConfig::default()
56 + }
57 +}
58 +
59 +fn shm_server_config() -> ServerConfig {
60 + ServerConfig {
61 + supported_profiles: PROFILE_SHM_HYBRID | PROFILE_BASELINE,
62 + max_request_payload_bytes: 4096,
63 + max_request_batch_items: 1,
64 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
65 + max_response_batch_items: 1,
66 + auth_token: AUTH_TOKEN,
67 + ..ServerConfig::default()
68 + }
69 +}
70 +
71 +fn shm_client_config() -> ClientConfig {
72 + ClientConfig {
73 + supported_profiles: PROFILE_SHM_HYBRID | PROFILE_BASELINE,
74 + max_request_payload_bytes: 4096,
75 + max_request_batch_items: 1,
76 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
77 + max_response_batch_items: 1,
78 + auth_token: AUTH_TOKEN,
79 + ..ClientConfig::default()
80 + }
81 +}
82 +
83 +fn batch_server_config() -> ServerConfig {
84 + ServerConfig {
85 + supported_profiles: PROFILE_BASELINE,
86 + max_request_payload_bytes: 4096,
87 + max_request_batch_items: 16,
88 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
89 + max_response_batch_items: 16,
90 + auth_token: AUTH_TOKEN,
91 + ..ServerConfig::default()
92 + }
93 +}
94 +
95 +fn batch_client_config() -> ClientConfig {
96 + ClientConfig {
97 + supported_profiles: PROFILE_BASELINE,
98 + max_request_payload_bytes: 4096,
99 + max_request_batch_items: 16,
100 + max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
101 + max_response_batch_items: 16,
102 + auth_token: AUTH_TOKEN,
103 + ..ClientConfig::default()
104 + }
105 +}
106 +
107 +fn snapshot_client(service: &str, config: ClientConfig) -> RawClient {
108 + RawClient::new_snapshot(TEST_RUN_DIR, service, config)
109 +}
110 +
111 +fn increment_client(service: &str, config: ClientConfig) -> RawClient {
112 + RawClient::new_increment(TEST_RUN_DIR, service, config)
113 +}
114 +
115 +fn string_reverse_client(service: &str, config: ClientConfig) -> RawClient {
116 + RawClient::new_string_reverse(TEST_RUN_DIR, service, config)
117 +}
118 +
119 +fn fill_test_cgroups_snapshot(builder: &mut CgroupsBuilder<'_>) -> bool {
120 + let items = [
121 + (
122 + 1001u32,
123 + 0u32,
124 + 1u32,
125 + b"docker-abc123" as &[u8],
126 + b"/sys/fs/cgroup/docker/abc123" as &[u8],
127 + ),
128 + (2002, 0, 1, b"k8s-pod-xyz", b"/sys/fs/cgroup/kubepods/xyz"),
129 + (
130 + 3003,
131 + 0,
132 + 0,
133 + b"systemd-user",
134 + b"/sys/fs/cgroup/user.slice/user-1000",
135 + ),
136 + ];
137 +
138 + for (hash, options, enabled, name, path) in &items {
139 + if builder.add(*hash, *options, *enabled, name, path).is_err() {
140 + return false;
141 + }
142 + }
143 +
144 + true
145 +}
146 +
147 +fn test_cgroups_dispatch() -> DispatchHandler {
148 + snapshot_dispatch(
149 + Arc::new(|req, builder| {
150 + if req.layout_version != 1 || req.flags != 0 {
151 + return false;
152 + }
153 + builder.set_header(1, 42);
154 + fill_test_cgroups_snapshot(builder)
155 + }),
156 + 3,
157 + )
158 +}
159 +
160 +fn increment_dispatch_handler() -> DispatchHandler {
161 + increment_dispatch(Arc::new(|value| Some(value + 1)))
162 +}
163 +
164 +fn connect_ready(client: &mut RawClient) {
165 + for _ in 0..200 {
166 + client.refresh();
167 + if client.ready() {
168 + return;
169 + }
170 + thread::sleep(Duration::from_millis(10));
171 + }
172 +
173 + panic!("client did not reach READY state");
174 +}
175 +
176 +fn wait_for_state(client: &mut RawClient, want: ClientState) {
177 + for _ in 0..200 {
178 + client.refresh();
179 + if client.state == want {
180 + return;
181 + }
182 + thread::sleep(Duration::from_millis(10));
183 + }
184 +
185 + panic!("client did not reach state {:?}", want);
186 +}
187 +
188 +struct TestServer {
189 + service: String,
190 + wake_config: ClientConfig,
191 + stop_flag: Arc<AtomicBool>,
192 + thread: Option<thread::JoinHandle<()>>,
193 +}
194 +
195 +impl TestServer {
196 + fn start(service: &str, expected_method_code: u16, handler: DispatchHandler) -> Self {
197 + Self::start_with(
198 + service,
199 + server_config(),
200 + client_config(),
201 + expected_method_code,
202 + handler,
203 + 8,
204 + )
205 + }
206 +
207 + fn start_shm(service: &str, expected_method_code: u16, handler: DispatchHandler) -> Self {
208 + Self::start_with(
209 + service,
210 + shm_server_config(),
211 + shm_client_config(),
212 + expected_method_code,
213 + handler,
214 + 8,
215 + )
216 + }
217 +
218 + fn start_batch(service: &str, expected_method_code: u16, handler: DispatchHandler) -> Self {
219 + Self::start_with(
220 + service,
221 + batch_server_config(),
222 + batch_client_config(),
223 + expected_method_code,
224 + handler,
225 + 8,
226 + )
227 + }
228 +
229 + fn start_with(
230 + service: &str,
231 + config: ServerConfig,
232 + wake_config: ClientConfig,
233 + expected_method_code: u16,
234 + handler: DispatchHandler,
235 + worker_count: usize,
236 + ) -> Self {
237 + ensure_run_dir();
238 + cleanup_all(service);
239 +
240 + let svc = service.to_string();
241 + let stop_cfg = wake_config.clone();
242 + let mut server = ManagedServer::with_workers(
243 + TEST_RUN_DIR,
244 + &svc,
245 + config,
246 + expected_method_code,
247 + Some(handler),
248 + worker_count,
249 + );
250 + let stop_flag = server.running_flag();
251 +
252 + let thread = thread::spawn(move || {
253 + let _ = server.run();
254 + });
255 + thread::sleep(Duration::from_millis(50));
256 +
257 + TestServer {
258 + service: svc,
259 + wake_config: stop_cfg,
260 + stop_flag,
261 + thread: Some(thread),
262 + }
263 + }
264 +
265 + fn stop(&mut self) {
266 + self.stop_flag.store(false, Ordering::Release);
267 +
268 + // Wake a blocking ConnectNamedPipe() so the accept loop can observe
269 + // the stop flag and exit.
270 + let _ = NpSession::connect(TEST_RUN_DIR, &self.service, &self.wake_config);
271 +
272 + if let Some(handle) = self.thread.take() {
273 + let _ = handle.join();
274 + }
275 +
276 + cleanup_all(&self.service);
277 + }
278 +}
279 +
280 +impl Drop for TestServer {
281 + fn drop(&mut self) {
282 + self.stop();
283 + }
284 +}
285 +
286 +struct RawSessionServer {
287 + thread: Option<thread::JoinHandle<Result<(), String>>>,
288 +}
289 +
290 +type RawHandle = isize;
291 +type Dword = u32;
292 +type Bool = i32;
293 +
294 +const INVALID_HANDLE_VALUE: RawHandle = -1;
295 +const PIPE_ACCESS_DUPLEX: Dword = 0x00000003;
296 +const PIPE_TYPE_MESSAGE: Dword = 0x00000004;
297 +const PIPE_READMODE_MESSAGE: Dword = 0x00000002;
298 +const PIPE_WAIT: Dword = 0x00000000;
299 +const ERROR_PIPE_CONNECTED: Dword = 535;
300 +const RAW_PIPE_PACKET_SIZE: Dword = 65536;
301 +
302 +unsafe extern "system" {
303 + fn CreateNamedPipeW(
304 + lp_name: *const u16,
305 + dw_open_mode: Dword,
306 + dw_pipe_mode: Dword,
307 + n_max_instances: Dword,
308 + n_out_buffer_size: Dword,
309 + n_in_buffer_size: Dword,
310 + n_default_time_out: Dword,
311 + lp_security_attributes: *const core::ffi::c_void,
312 + ) -> RawHandle;
313 +
314 + fn ConnectNamedPipe(handle: RawHandle, overlapped: *mut core::ffi::c_void) -> Bool;
315 + fn ReadFile(
316 + handle: RawHandle,
317 + buffer: *mut core::ffi::c_void,
318 + bytes_to_read: Dword,
319 + bytes_read: *mut Dword,
320 + overlapped: *mut core::ffi::c_void,
321 + ) -> Bool;
322 + fn WriteFile(
323 + handle: RawHandle,
324 + buffer: *const core::ffi::c_void,
325 + bytes_to_write: Dword,
326 + bytes_written: *mut Dword,
327 + overlapped: *mut core::ffi::c_void,
328 + ) -> Bool;
329 + fn CloseHandle(handle: RawHandle) -> Bool;
330 + fn GetLastError() -> Dword;
331 +}
332 +
333 +struct RawHelloAckServer {
334 + accepted: Arc<AtomicBool>,
335 + thread: Option<thread::JoinHandle<Result<(), String>>>,
336 +}
337 +
338 +fn encode_hello_ack_packet_with_version(version: u16, status: u16, layout_version: u16) -> Vec<u8> {
339 + let ack = HelloAck {
340 + layout_version,
341 + flags: 0,
342 + server_supported_profiles: PROFILE_BASELINE,
343 + intersection_profiles: PROFILE_BASELINE,
344 + selected_profile: PROFILE_BASELINE,
345 + agreed_max_request_payload_bytes: crate::protocol::MAX_PAYLOAD_DEFAULT,
346 + agreed_max_request_batch_items: 1,
347 + agreed_max_response_payload_bytes: RESPONSE_BUF_SIZE as u32,
348 + agreed_max_response_batch_items: 1,
349 + agreed_packet_size: 0,
350 + session_id: 77,
351 + };
352 +
353 + let mut payload = [0u8; 48];
354 + let payload_len = ack.encode(&mut payload);
355 +
356 + let hdr = Header {
357 + magic: crate::protocol::MAGIC_MSG,
358 + version,
359 + header_len: HEADER_SIZE as u16,
360 + kind: KIND_CONTROL,
361 + code: CODE_HELLO_ACK,
362 + flags: 0,
363 + item_count: 1,
364 + message_id: 0,
365 + payload_len: payload_len as u32,
366 + transport_status: status,
367 + };
368 +
369 + let mut packet = vec![0u8; HEADER_SIZE + payload_len];
370 + hdr.encode(&mut packet[..HEADER_SIZE]);
371 + packet[HEADER_SIZE..].copy_from_slice(&payload[..payload_len]);
372 + packet
373 +}
374 +
375 +fn start_raw_hello_ack_version_server(service: &str, version: u16) -> RawHelloAckServer {
376 + ensure_run_dir();
377 + cleanup_all(service);
378 +
379 + let accepted = Arc::new(AtomicBool::new(false));
380 + let accepted_flag = Arc::clone(&accepted);
381 + let svc = service.to_string();
382 + let packet = encode_hello_ack_packet_with_version(version, STATUS_OK, 1);
383 +
384 + let thread = thread::spawn(move || {
385 + let pipe_name =
386 + build_pipe_name(TEST_RUN_DIR, &svc).map_err(|e| format!("pipe name: {e}"))?;
387 + let pipe = unsafe {
388 + CreateNamedPipeW(
389 + pipe_name.as_ptr(),
390 + PIPE_ACCESS_DUPLEX,
391 + PIPE_TYPE_MESSAGE | PIPE_READMODE_MESSAGE | PIPE_WAIT,
392 + 1,
393 + RAW_PIPE_PACKET_SIZE,
394 + RAW_PIPE_PACKET_SIZE,
395 + 0,
396 + ptr::null(),
397 + )
398 + };
399 + if pipe == INVALID_HANDLE_VALUE {
400 + return Err(format!("CreateNamedPipeW failed: {}", unsafe {
401 + GetLastError()
402 + }));
403 + }
404 +
405 + let connect_ok = unsafe { ConnectNamedPipe(pipe, ptr::null_mut()) };
406 + if connect_ok == 0 {
407 + let err = unsafe { GetLastError() };
408 + if err != ERROR_PIPE_CONNECTED {
409 + unsafe {
410 + CloseHandle(pipe);
411 + }
412 + return Err(format!("ConnectNamedPipe failed: {err}"));
413 + }
414 + }
415 +
416 + accepted_flag.store(true, Ordering::Release);
417 +
418 + let mut hello_buf = [0u8; 256];
419 + let mut hello_n = 0u32;
420 + let read_ok = unsafe {
421 + ReadFile(
422 + pipe,
423 + hello_buf.as_mut_ptr().cast(),
424 + hello_buf.len() as u32,
425 + &mut hello_n,
426 + ptr::null_mut(),
427 + )
428 + };
429 + if read_ok == 0 || hello_n == 0 {
430 + unsafe {
431 + CloseHandle(pipe);
432 + }
433 + return Err(format!("ReadFile failed: {}", unsafe { GetLastError() }));
434 + }
435 +
436 + let mut written = 0u32;
437 + let write_ok = unsafe {
438 + WriteFile(
439 + pipe,
440 + packet.as_ptr().cast(),
441 + packet.len() as u32,
442 + &mut written,
443 + ptr::null_mut(),
444 + )
445 + };
446 + let write_err = unsafe { GetLastError() };
447 + unsafe {
448 + CloseHandle(pipe);
449 + }
450 + if write_ok == 0 || written as usize != packet.len() {
451 + return Err(format!("WriteFile failed: {write_err}"));
452 + }
453 +
454 + Ok(())
455 + });
456 +
457 + thread::sleep(Duration::from_millis(200));
458 + RawHelloAckServer {
459 + accepted,
460 + thread: Some(thread),
461 + }
462 +}
463 +
464 +impl RawHelloAckServer {
465 + fn wait(&mut self) {
466 + if let Some(thread) = self.thread.take() {
467 + match thread.join() {
468 + Ok(Ok(())) => {}
469 + Ok(Err(err)) => panic!("raw windows hello_ack server failed: {err}"),
470 + Err(_) => panic!("raw windows hello_ack server panicked"),
471 + }
472 + }
473 +
474 + assert!(
475 + self.accepted.load(Ordering::Acquire),
476 + "raw windows hello_ack server accepted no clients"
477 + );
478 + }
479 +}
480 +
481 +fn start_raw_session_server<F>(service: &str, cfg: ServerConfig, handler: F) -> RawSessionServer
482 +where
483 + F: FnOnce(&mut NpSession, Header, &[u8]) -> Result<(), String> + Send + 'static,
484 +{
485 + ensure_run_dir();
486 + cleanup_all(service);
487 +
488 + let svc = service.to_string();
489 + let thread = thread::spawn(move || {
490 + let mut listener =
491 + NpListener::bind(TEST_RUN_DIR, &svc, cfg).map_err(|e| format!("bind: {e}"))?;
492 + let mut session = listener.accept().map_err(|e| format!("accept: {e}"))?;
493 +
494 + let (hdr, payload) = {
495 + let mut recv_buf = vec![0u8; RESPONSE_BUF_SIZE];
496 + let (hdr, payload) = session
497 + .receive(&mut recv_buf)
498 + .map_err(|e| format!("receive: {e}"))?;
499 + (hdr, payload.to_vec())
500 + };
501 +
502 + handler(&mut session, hdr, &payload)
503 + });
504 +
505 + thread::sleep(Duration::from_millis(200));
506 + RawSessionServer {
507 + thread: Some(thread),
508 + }
509 +}
510 +
511 +impl RawSessionServer {
512 + fn wait(&mut self) {
513 + if let Some(thread) = self.thread.take() {
514 + match thread.join() {
515 + Ok(Ok(())) => {}
516 + Ok(Err(err)) => panic!("raw windows session server failed: {err}"),
517 + Err(_) => panic!("raw windows session server panicked"),
518 + }
519 + }
520 + }
521 +}
522 +
523 +#[test]
524 +fn test_client_lifecycle_windows() {
525 + let svc = "rs_win_svc_lifecycle";
526 + ensure_run_dir();
527 + cleanup_all(svc);
528 +
529 + let mut client = snapshot_client(svc, client_config());
530 + assert_eq!(client.state, ClientState::Disconnected);
531 + assert!(!client.ready());
532 +
533 + let changed = client.refresh();
534 + assert!(changed);
535 + assert_eq!(client.state, ClientState::NotFound);
536 +
537 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, test_cgroups_dispatch());
538 +
539 + connect_ready(&mut client);
540 + assert_eq!(client.state, ClientState::Ready);
541 + assert!(client.ready());
542 + assert_eq!(client.status().connect_count, 1);
543 +
544 + client.close();
545 + assert_eq!(client.state, ClientState::Disconnected);
546 +
547 + server.stop();
548 +}
549 +
550 +#[test]
551 +fn test_cgroups_call_windows_baseline() {
552 + let svc = "rs_win_svc_cgroups";
553 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, test_cgroups_dispatch());
554 +
555 + let mut client = snapshot_client(svc, client_config());
556 + connect_ready(&mut client);
557 +
558 + let view = client.call_snapshot().expect("snapshot");
559 + assert_eq!(view.item_count, 3);
560 + assert_eq!(view.systemd_enabled, 1);
561 + assert_eq!(view.generation, 42);
562 +
563 + let item0 = view.item(0).expect("item 0");
564 + assert_eq!(item0.hash, 1001);
565 + assert_eq!(item0.name.as_bytes(), b"docker-abc123");
566 +
567 + client.close();
568 + server.stop();
569 +}
570 +
571 +#[test]
572 +fn test_cgroups_call_windows_shm() {
573 + let svc = "rs_win_svc_cgroups_shm";
574 + let mut server = TestServer::start_shm(svc, METHOD_CGROUPS_SNAPSHOT, test_cgroups_dispatch());
575 +
576 + let mut client = snapshot_client(svc, shm_client_config());
577 + connect_ready(&mut client);
578 +
579 + assert!(client.shm.is_some(), "expected Win SHM to be negotiated");
580 + assert_eq!(
581 + client.session.as_ref().map(|s| s.selected_profile),
582 + Some(PROFILE_SHM_HYBRID)
583 + );
584 +
585 + let view = client.call_snapshot().expect("snapshot");
586 + assert_eq!(view.item_count, 3);
587 + assert_eq!(view.generation, 42);
588 +
589 + client.close();
590 + server.stop();
591 +}
592 +
593 +#[test]
594 +fn test_retry_on_failure_windows() {
595 + let svc = unique_service("rs_win_svc_retry");
596 + let mut server1 = TestServer::start(&svc, METHOD_CGROUPS_SNAPSHOT, test_cgroups_dispatch());
597 +
598 + let mut client = snapshot_client(&svc, client_config());
599 + connect_ready(&mut client);
600 +
601 + let view = client.call_snapshot().expect("first call");
602 + assert_eq!(view.item_count, 3);
603 +
604 + let (stop_tx, stop_rx) = std::sync::mpsc::channel();
605 + thread::spawn(move || {
606 + server1.stop();
607 + let _ = stop_tx.send(());
608 + });
609 + assert!(
610 + stop_rx.recv_timeout(Duration::from_secs(2)).is_ok(),
611 + "server stop should not hang with an active client session"
612 + );
613 + thread::sleep(Duration::from_millis(100));
614 +
615 + let mut server2 = TestServer::start(&svc, METHOD_CGROUPS_SNAPSHOT, test_cgroups_dispatch());
616 +
617 + let view2 = client.call_snapshot().expect("retry call");
618 + assert_eq!(view2.item_count, 3);
619 + assert!(client.status().reconnect_count >= 1);
620 +
621 + client.close();
622 + server2.stop();
623 +}
624 +
625 +#[test]
626 +fn test_non_request_terminates_session_windows() {
627 + let svc = "rs_win_svc_nonreq";
628 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, test_cgroups_dispatch());
629 +
630 + let mut session = NpSession::connect(TEST_RUN_DIR, svc, &client_config()).expect("connect");
631 +
632 + let mut hdr = Header {
633 + kind: KIND_RESPONSE,
634 + code: METHOD_CGROUPS_SNAPSHOT,
635 + flags: 0,
636 + item_count: 0,
637 + message_id: 1,
638 + transport_status: STATUS_OK,
639 + ..Header::default()
640 + };
641 + let send_result = session.send(&mut hdr, &[]);
642 + if send_result.is_ok() {
643 + thread::sleep(Duration::from_millis(100));
644 + let mut recv_buf = vec![0u8; 4096];
645 + let mut hdr2 = Header {
646 + kind: KIND_REQUEST,
647 + code: METHOD_CGROUPS_SNAPSHOT,
648 + flags: 0,
649 + item_count: 1,
650 + message_id: 2,
651 + transport_status: STATUS_OK,
652 + ..Header::default()
653 + };
654 + let req = CgroupsRequest {
655 + layout_version: 1,
656 + flags: 0,
657 + };
658 + let mut req_buf = [0u8; 4];
659 + req.encode(&mut req_buf);
660 + let _ = session.send(&mut hdr2, &req_buf);
661 + let recv = session.receive(&mut recv_buf);
662 + assert!(
663 + recv.is_err(),
664 + "server should terminate the offending session"
665 + );
666 + }
667 +
668 + drop(session);
669 +
670 + let mut verify = snapshot_client(svc, client_config());
671 + connect_ready(&mut verify);
672 + let view = verify.call_snapshot().expect("normal call");
673 + assert_eq!(view.item_count, 3);
674 +
675 + verify.close();
676 + server.stop();
677 +}
678 +
679 +#[test]
680 +fn test_cache_full_round_trip_windows() {
681 + let svc = "rs_win_cache_roundtrip";
682 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, test_cgroups_dispatch());
683 +
684 + let mut cache = CgroupsCache::new(TEST_RUN_DIR, svc, client_config());
685 + assert!(!cache.ready());
686 +
687 + thread::sleep(Duration::from_millis(2));
688 +
689 + let updated = cache.refresh();
690 + assert!(updated);
691 + assert!(cache.ready());
692 +
693 + let item = cache.lookup(1001, "docker-abc123").expect("lookup");
694 + assert_eq!(item.hash, 1001);
695 + assert_eq!(item.path, "/sys/fs/cgroup/docker/abc123");
696 +
697 + let status = cache.status();
698 + assert!(status.populated);
699 + assert_eq!(status.item_count, 3);
700 + assert_eq!(status.generation, 42);
701 + assert_eq!(status.connection_state, ClientState::Ready);
702 +
703 + cache.close();
704 + server.stop();
705 +}
706 +
707 +#[test]
708 +fn test_increment_ping_pong_windows() {
709 + let svc = "rs_win_pp_increment";
710 + let mut server = TestServer::start(svc, METHOD_INCREMENT, increment_dispatch_handler());
711 +
712 + let mut client = increment_client(svc, client_config());
713 + connect_ready(&mut client);
714 +
715 + let mut value = 0u64;
716 + for _ in 0..10 {
717 + value = client.call_increment(value).expect("increment");
718 + }
719 + assert_eq!(value, 10);
720 +
721 + client.close();
722 + server.stop();
723 +}
724 +
725 +#[test]
726 +fn test_increment_batch_windows() {
727 + let svc = "rs_win_pp_batch";
728 + let mut server = TestServer::start_batch(svc, METHOD_INCREMENT, increment_dispatch_handler());
729 +
730 + let mut client = increment_client(svc, batch_client_config());
731 + connect_ready(&mut client);
732 +
733 + let values = vec![10u64, 20, 30, 40];
734 + let results = client.call_increment_batch(&values).expect("batch call");
735 + assert_eq!(results, vec![11, 21, 31, 41]);
736 +
737 + let single = client
738 + .call_increment_batch(&[99])
739 + .expect("single-item batch");
740 + assert_eq!(single, vec![100]);
741 +
742 + client.close();
743 + server.stop();
744 +}
745 +
746 +#[test]
747 +fn test_client_auth_failure_windows() {
748 + let svc = "rs_win_svc_authfail";
749 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, test_cgroups_dispatch());
750 +
751 + let mut bad_cfg = client_config();
752 + bad_cfg.auth_token = 0xBAD_BAD_BAD;
753 +
754 + let mut client = snapshot_client(svc, bad_cfg);
755 + wait_for_state(&mut client, ClientState::AuthFailed);
756 + assert_eq!(client.state, ClientState::AuthFailed);
757 + assert!(!client.ready());
758 +
759 + client.refresh();
760 + assert_eq!(client.state, ClientState::AuthFailed);
761 +
762 + client.close();
763 + server.stop();
764 +}
765 +
766 +#[test]
767 +fn test_refresh_winshm_attach_failure_falls_back_to_baseline() {
768 + let svc = unique_service("rs_win_svc_shm_attach_fail");
769 + ensure_run_dir();
770 + cleanup_all(&svc);
771 +
772 + let ready = Arc::new(AtomicBool::new(false));
773 + let ready_clone = ready.clone();
774 + let server_svc = svc.clone();
775 + let server_thread = thread::spawn(move || -> Result<(), String> {
776 + let mut listener = NpListener::bind(TEST_RUN_DIR, &server_svc, shm_server_config())
777 + .map_err(|e| format!("bind: {e}"))?;
778 + ready_clone.store(true, Ordering::Release);
779 +
780 + let mut first = listener
781 + .accept()
782 + .map_err(|e| format!("accept first: {e}"))?;
783 + if first.selected_profile != PROFILE_SHM_HYBRID {
784 + return Err(format!(
785 + "first selected profile = {}, want {}",
786 + first.selected_profile, PROFILE_SHM_HYBRID
787 + ));
788 + }
789 +
790 + let mut recv_buf = vec![0u8; RESPONSE_BUF_SIZE];
791 + if first.receive(&mut recv_buf).is_ok() {
792 + return Err("first SHM-selected session should disconnect after attach failure".into());
793 + }
794 +
795 + let mut second = listener
796 + .accept()
797 + .map_err(|e| format!("accept second: {e}"))?;
798 + if second.selected_profile != PROFILE_BASELINE {
799 + return Err(format!(
800 + "second selected profile = {}, want {}",
801 + second.selected_profile, PROFILE_BASELINE
802 + ));
803 + }
804 +
805 + if second.receive(&mut recv_buf).is_ok() {
806 + return Err("second baseline session should close cleanly when client closes".into());
807 + }
808 +
809 + Ok(())
810 + });
811 +
812 + while !ready.load(Ordering::Acquire) {
813 + thread::sleep(Duration::from_millis(1));
814 + }
815 +
816 + let mut client = increment_client(&svc, shm_client_config());
817 + assert!(
818 + client.refresh(),
819 + "refresh should transition to READY via baseline fallback"
820 + );
821 + assert_eq!(client.state, ClientState::Ready);
822 + assert!(client.ready());
823 + assert!(
824 + client.shm.is_none(),
825 + "fallback session must not attach WinSHM"
826 + );
827 + assert_eq!(
828 + client.session.as_ref().map(|s| s.selected_profile),
829 + Some(PROFILE_BASELINE)
830 + );
831 + assert_eq!(client.transport_config.supported_profiles, PROFILE_BASELINE);
832 + assert_eq!(client.transport_config.preferred_profiles, 0);
833 +
834 + client.close();
835 + match server_thread.join() {
836 + Ok(Ok(())) => {}
837 + Ok(Err(err)) => panic!("raw win attach-failure server failed: {err}"),
838 + Err(_) => panic!("raw win attach-failure server panicked"),
839 + }
840 +}
841 +
842 +#[test]
843 +fn test_client_incompatible_windows() {
844 + let svc = "rs_win_svc_incompat";
845 + let mut server = TestServer::start(svc, METHOD_CGROUPS_SNAPSHOT, test_cgroups_dispatch());
846 +
847 + let mut bad_cfg = client_config();
848 + bad_cfg.supported_profiles = 0x80000000;
849 +
850 + let mut client = snapshot_client(svc, bad_cfg);
851 + wait_for_state(&mut client, ClientState::Incompatible);
852 + assert_eq!(client.state, ClientState::Incompatible);
853 + assert!(!client.ready());
854 +
855 + client.refresh();
856 + assert_eq!(client.state, ClientState::Incompatible);
857 +
858 + client.close();
859 + server.stop();
860 +}
861 +
862 +#[test]
863 +fn test_client_protocol_version_incompatible_windows() {
864 + let svc = unique_service("rs_win_svc_proto_incompat");
865 + let mut server = start_raw_hello_ack_version_server(&svc, VERSION + 1);
866 +
867 + let mut client = snapshot_client(&svc, client_config());
868 + let changed = client.refresh();
869 + assert!(changed, "refresh should move client into INCOMPATIBLE");
870 + assert_eq!(client.state, ClientState::Incompatible);
871 + assert!(!client.ready());
872 +
873 + server.wait();
874 +
875 + let changed = client.refresh();
876 + assert!(!changed, "refresh from incompatible should be a no-op");
877 + assert_eq!(client.state, ClientState::Incompatible);
878 +
879 + client.close();
880 + cleanup_all(&svc);
881 +}
882 +
883 +#[test]
884 +fn test_call_when_not_ready_windows() {
885 + let svc = "rs_win_svc_noready";
886 + let mut snapshot = snapshot_client(svc, client_config());
887 + assert_eq!(snapshot.state, ClientState::Disconnected);
888 + assert!(snapshot.call_snapshot().is_err());
889 + assert_eq!(snapshot.status().error_count, 1);
890 + snapshot.close();
891 +
892 + let mut increment = increment_client(svc, client_config());
893 + assert_eq!(increment.state, ClientState::Disconnected);
894 + assert!(increment.call_increment(42).is_err());
895 + assert_eq!(increment.status().error_count, 1);
896 + increment.close();
897 +
898 + let mut string_reverse = string_reverse_client(svc, client_config());
899 + assert_eq!(string_reverse.state, ClientState::Disconnected);
900 + assert!(string_reverse.call_string_reverse("test").is_err());
901 + assert_eq!(string_reverse.status().error_count, 1);
902 + string_reverse.close();
903 +}
904 +
905 +#[test]
906 +fn test_server_worker_count_clamped_windows() {
907 + let svc = "rs_win_svc_w0";
908 + ensure_run_dir();
909 + cleanup_all(svc);
910 +
911 + let server = ManagedServer::with_workers(
912 + TEST_RUN_DIR,
913 + svc,
914 + server_config(),
915 + METHOD_INCREMENT,
916 + None,
917 + 0,
918 + );
919 + assert_eq!(server.worker_count, 1);
920 +}
921 +
922 +#[test]
923 +fn test_server_stop_flag_windows() {
924 + let svc = "rs_win_svc_stopflag";
925 + ensure_run_dir();
926 + cleanup_all(svc);
927 +
928 + let server = ManagedServer::new(TEST_RUN_DIR, svc, server_config(), METHOD_INCREMENT, None);
929 + let flag = server.running_flag();
930 + assert!(!flag.load(Ordering::Acquire));
931 +
932 + server.stop();
933 + assert!(!flag.load(Ordering::Acquire));
934 +}
935 +
936 +#[test]
937 +fn test_transport_without_session_windows() {
938 + let svc = unique_service("rs_win_transport");
939 + let mut client = increment_client(&svc, client_config());
940 +
941 + let mut hdr = Header {
942 + kind: KIND_REQUEST,
943 + code: METHOD_INCREMENT,
944 + flags: 0,
945 + item_count: 1,
946 + message_id: 1,
947 + transport_status: STATUS_OK,
948 + ..Header::default()
949 + };
950 +
951 + assert_eq!(
952 + client.transport_send(&mut hdr, &[]),
953 + Err(NipcError::Truncated)
954 + );
955 + assert!(matches!(
956 + client.transport_receive(),
957 + Err(NipcError::Truncated)
958 + ));
959 + client.close();
960 +}
961 +
962 +#[test]
963 +fn test_call_increment_rejects_malformed_response_envelope_windows() {
964 + struct Case {
965 + name: &'static str,
966 + kind: u16,
967 + code: u16,
968 + status: u16,
969 + want: NipcError,
970 + }
971 +
972 + let cases = [
973 + Case {
974 + name: "bad kind",
975 + kind: KIND_REQUEST,
976 + code: METHOD_INCREMENT,
977 + status: STATUS_OK,
978 + want: NipcError::BadKind,
979 + },
980 + Case {
981 + name: "bad code",
982 + kind: KIND_RESPONSE,
983 + code: METHOD_STRING_REVERSE,
984 + status: STATUS_OK,
985 + want: NipcError::BadLayout,
986 + },
987 + Case {
988 + name: "bad status",
989 + kind: KIND_RESPONSE,
990 + code: METHOD_INCREMENT,
991 + status: STATUS_INTERNAL_ERROR,
992 + want: NipcError::BadLayout,
993 + },
994 + ];
995 +
996 + for tc in cases {
997 + let svc = unique_service("rs_win_inc_env");
998 + let mut server =
999 + start_raw_session_server(&svc, server_config(), move |session, req_hdr, _| {
1000 + let mut payload = [0u8; INCREMENT_PAYLOAD_SIZE];
1001 + let n = increment_encode(43, &mut payload);
1002 + if n != INCREMENT_PAYLOAD_SIZE {
1003 + return Err(format!("increment_encode returned {n}"));
1004 + }
1005 +
1006 + let mut resp_hdr = Header {
1007 + kind: tc.kind,
1008 + code: tc.code,
1009 + flags: 0,
1010 + item_count: 1,
1011 + message_id: req_hdr.message_id,
1012 + transport_status: tc.status,
1013 + ..Header::default()
1014 + };
1015 + session
1016 + .send(&mut resp_hdr, &payload)
1017 + .map_err(|e| format!("send: {e}"))
1018 + });
1019 +
1020 + let mut client = increment_client(&svc, client_config());
1021 + connect_ready(&mut client);
1022 +
1023 + let err = client.call_increment(42).expect_err(tc.name);
1024 + assert_eq!(err, tc.want, "{}", tc.name);
1025 +
1026 + client.close();
1027 + server.wait();
1028 + }
1029 +}
1030 +
1031 +#[test]
1032 +fn test_call_increment_rejects_malformed_payload_windows() {
1033 + let svc = unique_service("rs_win_inc_payload");
1034 + let mut server = start_raw_session_server(&svc, server_config(), move |session, req_hdr, _| {
1035 + let payload = [1u8, 2, 3, 4];
1036 + let mut resp_hdr = Header {
1037 + kind: KIND_RESPONSE,
1038 + code: METHOD_INCREMENT,
1039 + flags: 0,
1040 + item_count: 1,
1041 + message_id: req_hdr.message_id,
1042 + transport_status: STATUS_OK,
1043 + ..Header::default()
1044 + };
1045 + session
1046 + .send(&mut resp_hdr, &payload)
1047 + .map_err(|e| format!("send: {e}"))
1048 + });
1049 +
1050 + let mut client = increment_client(&svc, client_config());
1051 + connect_ready(&mut client);
1052 +
1053 + let err = client
1054 + .call_increment(42)
1055 + .expect_err("malformed increment response");
1056 + assert_eq!(err, NipcError::Truncated);
1057 +
1058 + client.close();
1059 + server.wait();
1060 +}
1061 +
1062 +#[test]
1063 +fn test_call_string_reverse_rejects_missing_nul_windows() {
1064 + let svc = unique_service("rs_win_str_payload");
1065 + let mut server = start_raw_session_server(&svc, server_config(), move |session, req_hdr, _| {
1066 + let payload = [
1067 + 8u8, 0, 0, 0, // str_offset = 8
1068 + 2, 0, 0, 0, // str_length = 2
1069 + b'o', b'k', b'!', // missing trailing NUL
1070 + ];
1071 + let mut resp_hdr = Header {
1072 + kind: KIND_RESPONSE,
1073 + code: METHOD_STRING_REVERSE,
1074 + flags: 0,
1075 + item_count: 1,
1076 + message_id: req_hdr.message_id,
1077 + transport_status: STATUS_OK,
1078 + ..Header::default()
1079 + };
1080 + session
1081 + .send(&mut resp_hdr, &payload)
1082 + .map_err(|e| format!("send: {e}"))
1083 + });
1084 +
1085 + let mut client = string_reverse_client(&svc, client_config());
1086 + connect_ready(&mut client);
1087 +
1088 + let err = client
1089 + .call_string_reverse("ok")
1090 + .expect_err("malformed string response");
1091 + assert_eq!(err, NipcError::MissingNul);
1092 +
1093 + client.close();
1094 + server.wait();
1095 +}
1096 +
1097 +#[test]
1098 +fn test_call_increment_batch_rejects_wrong_item_count_windows() {
1099 + let svc = unique_service("rs_win_batch_count");
1100 + let mut server =
1101 + start_raw_session_server(&svc, batch_server_config(), move |session, req_hdr, _| {
1102 + let mut encoded = [0u8; INCREMENT_PAYLOAD_SIZE];
1103 + let n = increment_encode(11, &mut encoded);
1104 + if n != INCREMENT_PAYLOAD_SIZE {
1105 + return Err(format!("increment_encode returned {n}"));
1106 + }
1107 +
1108 + let mut response_buf = vec![0u8; 128];
1109 + let resp_len = {
1110 + let mut batch = BatchBuilder::new(&mut response_buf, 2);
1111 + batch
1112 + .add(&encoded)
1113 + .map_err(|e| format!("batch add 1: {e:?}"))?;
1114 + batch
1115 + .add(&encoded)
1116 + .map_err(|e| format!("batch add 2: {e:?}"))?;
1117 + let (len, _count) = batch.finish();
1118 + len
1119 + };
1120 +
1121 + let mut resp_hdr = Header {
1122 + kind: KIND_RESPONSE,
1123 + code: METHOD_INCREMENT,
1124 + flags: FLAG_BATCH,
1125 + item_count: 1,
1126 + message_id: req_hdr.message_id,
1127 + transport_status: STATUS_OK,
1128 + ..Header::default()
1129 + };
1130 + session
1131 + .send(&mut resp_hdr, &response_buf[..resp_len])
1132 + .map_err(|e| format!("send: {e}"))
1133 + });
1134 +
1135 + let mut client = increment_client(&svc, batch_client_config());
1136 + connect_ready(&mut client);
1137 +
1138 + let err = client
1139 + .call_increment_batch(&[10, 20])
1140 + .expect_err("wrong batch item_count");
1141 + assert_eq!(err, NipcError::BadItemCount);
1142 +
1143 + client.close();
1144 + server.wait();
1145 +}
src/crates/netipc/src/transport/mod.rs new
+16
@@ -0,0 +1,16 @@
1 +//! L1 transport backends.
2 +//!
3 +//! Each transport provides connection lifecycle, handshake, and
4 +//! send/receive with transparent chunking.
5 +
6 +#[cfg(unix)]
7 +pub mod posix;
8 +
9 +#[cfg(target_os = "linux")]
10 +pub mod shm;
11 +
12 +#[cfg(windows)]
13 +pub mod windows;
14 +
15 +#[cfg(windows)]
16 +pub mod win_shm;
src/crates/netipc/src/transport/posix.rs new
+1282
@@ -0,0 +1,1282 @@
1 +//! L1 POSIX UDS SEQPACKET transport.
2 +//!
3 +//! Connection lifecycle, handshake with profile/limit negotiation,
4 +//! and send/receive with transparent chunking over AF_UNIX SEQPACKET sockets.
5 +//! Wire-compatible with the C implementation in netipc_uds.c.
6 +
7 +use crate::protocol::{
8 + self, align8, ChunkHeader, Header, Hello, HelloAck, FLAG_BATCH, HEADER_SIZE, KIND_REQUEST,
9 + KIND_RESPONSE, MAGIC_CHUNK, MAGIC_MSG, MAX_PAYLOAD_CAP, MAX_PAYLOAD_DEFAULT, PROFILE_BASELINE,
10 + VERSION,
11 +};
12 +use std::collections::HashSet;
13 +use std::ffi::CString;
14 +use std::io;
15 +use std::os::unix::io::RawFd;
16 +use std::path::{Path, PathBuf};
17 +use std::sync::atomic::{AtomicU64, Ordering};
18 +
19 +// ---------------------------------------------------------------------------
20 +// Constants
21 +// ---------------------------------------------------------------------------
22 +
23 +const DEFAULT_BACKLOG: i32 = 16;
24 +const DEFAULT_BATCH_ITEMS: u32 = 1;
25 +const DEFAULT_PACKET_SIZE_FALLBACK: u32 = 65536;
26 +const HELLO_PAYLOAD_SIZE: usize = 44;
27 +const HELLO_ACK_PAYLOAD_SIZE: usize = 48;
28 +
29 +// ---------------------------------------------------------------------------
30 +// Errors
31 +// ---------------------------------------------------------------------------
32 +
33 +/// Transport-level errors.
34 +#[derive(Debug, Clone, PartialEq, Eq)]
35 +pub enum UdsError {
36 + /// Socket path exceeds sun_path limit.
37 + PathTooLong,
38 + /// socket()/bind()/listen() syscall failed.
39 + Socket(i32),
40 + /// connect() failed.
41 + Connect(i32),
42 + /// accept() failed.
43 + Accept(i32),
44 + /// send()/sendmsg() failed.
45 + Send(i32),
46 + /// recv() failed or peer disconnected.
47 + Recv(i32),
48 + /// Handshake protocol error.
49 + Handshake(String),
50 + /// Authentication token rejected.
51 + AuthFailed,
52 + /// No common profile between client and server.
53 + NoProfile,
54 + /// Protocol or layout version mismatch.
55 + Incompatible(String),
56 + /// Wire protocol violation.
57 + Protocol(String),
58 + /// A live server already owns this socket path.
59 + AddrInUse,
60 + /// Chunk header mismatch during reassembly.
61 + Chunk(String),
62 + /// Memory allocation failed.
63 + Alloc,
64 + /// Payload or batch count exceeds negotiated limit.
65 + LimitExceeded,
66 + /// Invalid argument.
67 + BadParam(String),
68 + /// Duplicate message_id on send.
69 + DuplicateMsgId(u64),
70 + /// Unknown message_id on receive.
71 + UnknownMsgId(u64),
72 +}
73 +
74 +impl std::fmt::Display for UdsError {
75 + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
76 + match self {
77 + UdsError::PathTooLong => write!(f, "socket path exceeds sun_path limit"),
78 + UdsError::Socket(e) => write!(f, "socket syscall failed: errno {e}"),
79 + UdsError::Connect(e) => write!(f, "connect failed: errno {e}"),
80 + UdsError::Accept(e) => write!(f, "accept failed: errno {e}"),
81 + UdsError::Send(e) => write!(f, "send failed: errno {e}"),
82 + UdsError::Recv(e) => write!(f, "recv failed: errno {e}"),
83 + UdsError::Handshake(s) => write!(f, "handshake error: {s}"),
84 + UdsError::AuthFailed => write!(f, "authentication token rejected"),
85 + UdsError::NoProfile => write!(f, "no common transport profile"),
86 + UdsError::Incompatible(s) => write!(f, "incompatible protocol: {s}"),
87 + UdsError::Protocol(s) => write!(f, "protocol violation: {s}"),
88 + UdsError::AddrInUse => write!(f, "address already in use by live server"),
89 + UdsError::Chunk(s) => write!(f, "chunk error: {s}"),
90 + UdsError::Alloc => write!(f, "memory allocation failed"),
91 + UdsError::LimitExceeded => write!(f, "negotiated limit exceeded"),
92 + UdsError::BadParam(s) => write!(f, "bad parameter: {s}"),
93 + UdsError::DuplicateMsgId(id) => write!(f, "duplicate message_id: {id}"),
94 + UdsError::UnknownMsgId(id) => write!(f, "unknown response message_id: {id}"),
95 + }
96 + }
97 +}
98 +
99 +fn header_version_incompatible(buf: &[u8], expected_code: u16) -> bool {
100 + if buf.len() < HEADER_SIZE {
101 + return false;
102 + }
103 +
104 + let magic = u32::from_ne_bytes(buf[0..4].try_into().unwrap());
105 + let version = u16::from_ne_bytes(buf[4..6].try_into().unwrap());
106 + let header_len = u16::from_ne_bytes(buf[6..8].try_into().unwrap());
107 + let kind = u16::from_ne_bytes(buf[8..10].try_into().unwrap());
108 + let code = u16::from_ne_bytes(buf[12..14].try_into().unwrap());
109 +
110 + magic == MAGIC_MSG
111 + && version != VERSION
112 + && header_len == protocol::HEADER_LEN
113 + && kind == protocol::KIND_CONTROL
114 + && code == expected_code
115 +}
116 +
117 +fn hello_layout_incompatible(buf: &[u8]) -> bool {
118 + buf.len() >= 2 && u16::from_ne_bytes(buf[0..2].try_into().unwrap()) != 1
119 +}
120 +
121 +fn hello_ack_layout_incompatible(buf: &[u8]) -> bool {
122 + buf.len() >= 2 && u16::from_ne_bytes(buf[0..2].try_into().unwrap()) != 1
123 +}
124 +
125 +impl std::error::Error for UdsError {}
126 +
127 +// ---------------------------------------------------------------------------
128 +// Role
129 +// ---------------------------------------------------------------------------
130 +
131 +#[derive(Debug, Clone, Copy, PartialEq, Eq)]
132 +pub enum Role {
133 + Client = 1,
134 + Server = 2,
135 +}
136 +
137 +// ---------------------------------------------------------------------------
138 +// Configuration
139 +// ---------------------------------------------------------------------------
140 +
141 +/// Client connection configuration.
142 +#[derive(Debug, Clone)]
143 +pub struct ClientConfig {
144 + pub supported_profiles: u32,
145 + pub preferred_profiles: u32,
146 + pub max_request_payload_bytes: u32,
147 + pub max_request_batch_items: u32,
148 + pub max_response_payload_bytes: u32,
149 + pub max_response_batch_items: u32,
150 + pub auth_token: u64,
151 + /// 0 = auto-detect from SO_SNDBUF.
152 + pub packet_size: u32,
153 +}
154 +
155 +impl Default for ClientConfig {
156 + fn default() -> Self {
157 + Self {
158 + supported_profiles: PROFILE_BASELINE,
159 + preferred_profiles: 0,
160 + max_request_payload_bytes: 0,
161 + max_request_batch_items: 0,
162 + max_response_payload_bytes: 0,
163 + max_response_batch_items: 0,
164 + auth_token: 0,
165 + packet_size: 0,
166 + }
167 + }
168 +}
169 +
170 +/// Server configuration for listen + accept.
171 +#[derive(Debug, Clone)]
172 +pub struct ServerConfig {
173 + pub supported_profiles: u32,
174 + pub preferred_profiles: u32,
175 + pub max_request_payload_bytes: u32,
176 + pub max_request_batch_items: u32,
177 + pub max_response_payload_bytes: u32,
178 + pub max_response_batch_items: u32,
179 + pub auth_token: u64,
180 + /// 0 = auto-detect from SO_SNDBUF.
181 + pub packet_size: u32,
182 + /// listen() backlog, 0 = default (16).
183 + pub backlog: i32,
184 +}
185 +
186 +impl Default for ServerConfig {
187 + fn default() -> Self {
188 + Self {
189 + supported_profiles: PROFILE_BASELINE,
190 + preferred_profiles: 0,
191 + max_request_payload_bytes: 0,
192 + max_request_batch_items: 0,
193 + max_response_payload_bytes: 0,
194 + max_response_batch_items: 0,
195 + auth_token: 0,
196 + packet_size: 0,
197 + backlog: 0,
198 + }
199 + }
200 +}
201 +
202 +// ---------------------------------------------------------------------------
203 +// Session
204 +// ---------------------------------------------------------------------------
205 +
206 +/// A connected UDS SEQPACKET session (client or server side).
207 +pub struct UdsSession {
208 + fd: RawFd,
209 + role: Role,
210 +
211 + // Negotiated limits
212 + pub max_request_payload_bytes: u32,
213 + pub max_request_batch_items: u32,
214 + pub max_response_payload_bytes: u32,
215 + pub max_response_batch_items: u32,
216 + pub packet_size: u32,
217 + pub selected_profile: u32,
218 + pub session_id: u64,
219 +
220 + // Internal receive buffer for chunked reassembly
221 + recv_buf: Vec<u8>,
222 + pkt_buf: Vec<u8>,
223 +
224 + // In-flight message_id set (client-side only)
225 + inflight_ids: HashSet<u64>,
226 +}
227 +
228 +impl UdsSession {
229 + fn fail_all_inflight(&mut self) {
230 + if self.role == Role::Client {
231 + self.inflight_ids.clear();
232 + }
233 + }
234 +
235 + /// Get the raw fd for poll/epoll integration.
236 + pub fn fd(&self) -> RawFd {
237 + self.fd
238 + }
239 +
240 + /// Get the session role.
241 + pub fn role(&self) -> Role {
242 + self.role
243 + }
244 +
245 + /// Return the most recently reassembled payload stored in the internal
246 + /// receive buffer.
247 + pub fn received_payload(&self, len: usize) -> &[u8] {
248 + &self.recv_buf[..len]
249 + }
250 +
251 + /// Connect to a server at `{run_dir}/{service_name}.sock`.
252 + /// Performs the full handshake. Blocks until connected + handshake done.
253 + pub fn connect(
254 + run_dir: &str,
255 + service_name: &str,
256 + config: &ClientConfig,
257 + ) -> Result<Self, UdsError> {
258 + let path = build_socket_path(run_dir, service_name)?;
259 +
260 + let fd = unsafe { libc::socket(libc::AF_UNIX, libc::SOCK_SEQPACKET, 0) };
261 + if fd < 0 {
262 + return Err(UdsError::Socket(errno()));
263 + }
264 +
265 + let result = connect_and_handshake(fd, &path, config);
266 + if result.is_err() {
267 + unsafe {
268 + libc::close(fd);
269 + }
270 + }
271 + result
272 + }
273 +
274 + /// Send one logical message. `hdr` is the 32-byte outer header (caller
275 + /// fills kind, code, flags, item_count, message_id; this function sets
276 + /// magic/version/header_len/payload_len).
277 + ///
278 + /// If the total message (32 + payload_len) exceeds packet_size, the
279 + /// message is chunked transparently.
280 + pub fn send(&mut self, hdr: &mut Header, payload: &[u8]) -> Result<(), UdsError> {
281 + if self.fd < 0 {
282 + return Err(UdsError::BadParam("session closed".into()));
283 + }
284 +
285 + // Validate payload against negotiated directional limits before transmitting.
286 + // The u32 cast below requires payload.len() <= u32::MAX.
287 + let (max_payload, max_items) = if self.role == Role::Client {
288 + (self.max_request_payload_bytes, self.max_request_batch_items)
289 + } else {
290 + (
291 + self.max_response_payload_bytes,
292 + self.max_response_batch_items,
293 + )
294 + };
295 + if payload.len() > max_payload as usize || payload.len() > u32::MAX as usize {
296 + return Err(UdsError::LimitExceeded);
297 + }
298 + if hdr.item_count > max_items {
299 + return Err(UdsError::LimitExceeded);
300 + }
301 +
302 + // Client-side: track in-flight message_ids for requests
303 + if self.role == Role::Client && hdr.kind == KIND_REQUEST {
304 + if !self.inflight_ids.insert(hdr.message_id) {
305 + return Err(UdsError::DuplicateMsgId(hdr.message_id));
306 + }
307 + }
308 +
309 + // Fill envelope fields
310 + hdr.magic = MAGIC_MSG;
311 + hdr.version = VERSION;
312 + hdr.header_len = protocol::HEADER_LEN;
313 + hdr.payload_len = payload.len() as u32;
314 +
315 + let tracked = self.role == Role::Client && hdr.kind == KIND_REQUEST;
316 + let msg_id = hdr.message_id;
317 +
318 + let result = self.send_inner(hdr, payload);
319 +
320 + if let Err(err) = &result {
321 + if tracked {
322 + match err {
323 + UdsError::Send(_) => self.fail_all_inflight(),
324 + _ => {
325 + self.inflight_ids.remove(&msg_id);
326 + }
327 + }
328 + }
329 + }
330 +
331 + result
332 + }
333 +
334 + /// Inner send logic, separated so the caller can rollback on failure.
335 + fn send_inner(&mut self, hdr: &mut Header, payload: &[u8]) -> Result<(), UdsError> {
336 + let total_msg = HEADER_SIZE + payload.len();
337 +
338 + // Single packet?
339 + if total_msg <= self.packet_size as usize {
340 + let mut hdr_buf = [0u8; HEADER_SIZE];
341 + hdr.encode(&mut hdr_buf);
342 + return raw_send_iov(self.fd, &hdr_buf, payload);
343 + }
344 +
345 + // Chunked send
346 + let chunk_payload_budget = self.packet_size as usize - HEADER_SIZE;
347 + if chunk_payload_budget == 0 {
348 + return Err(UdsError::BadParam("packet_size too small".into()));
349 + }
350 +
351 + let first_chunk_payload = payload.len().min(chunk_payload_budget);
352 + let remaining_after_first = payload.len() - first_chunk_payload;
353 +
354 + let continuation_chunks = if remaining_after_first > 0 {
355 + (remaining_after_first + chunk_payload_budget - 1) / chunk_payload_budget
356 + } else {
357 + 0
358 + };
359 + let chunk_count = 1 + continuation_chunks as u32;
360 +
361 + // Send first chunk: outer header + first part of payload
362 + let mut hdr_buf = [0u8; HEADER_SIZE];
363 + hdr.encode(&mut hdr_buf);
364 + raw_send_iov(self.fd, &hdr_buf, &payload[..first_chunk_payload])?;
365 +
366 + // Send continuation chunks
367 + let mut offset = first_chunk_payload;
368 + for ci in 1..chunk_count {
369 + let remaining = payload.len() - offset;
370 + let this_chunk = remaining.min(chunk_payload_budget);
371 +
372 + let chk = ChunkHeader {
373 + magic: MAGIC_CHUNK,
374 + version: VERSION,
375 + flags: 0,
376 + message_id: hdr.message_id,
377 + total_message_len: total_msg as u32,
378 + chunk_index: ci,
379 + chunk_count,
380 + chunk_payload_len: this_chunk as u32,
381 + };
382 +
383 + let mut chk_buf = [0u8; HEADER_SIZE];
384 + chk.encode(&mut chk_buf);
385 + raw_send_iov(self.fd, &chk_buf, &payload[offset..offset + this_chunk])?;
386 +
387 + offset += this_chunk;
388 + }
389 +
390 + Ok(())
391 + }
392 +
393 + /// Receive one logical message. Blocks until a complete message arrives.
394 + ///
395 + /// On success, returns (header, payload_view). The payload view is valid
396 + /// until the next receive call on this session.
397 + pub fn receive<'a>(&'a mut self, buf: &'a mut [u8]) -> Result<(Header, &'a [u8]), UdsError> {
398 + if self.fd < 0 {
399 + return Err(UdsError::BadParam("session closed".into()));
400 + }
401 +
402 + // Read first packet
403 + let n = match raw_recv(self.fd, buf) {
404 + Ok(n) => n,
405 + Err(err) => {
406 + self.fail_all_inflight();
407 + return Err(err);
408 + }
409 + };
410 +
411 + if n < HEADER_SIZE {
412 + return Err(UdsError::Protocol("packet too short for header".into()));
413 + }
414 +
415 + let hdr = Header::decode(&buf[..n])
416 + .map_err(|e| UdsError::Protocol(format!("header decode: {e}")))?;
417 +
418 + // Validate payload_len against negotiated directional limit.
419 + // Server receives requests; client receives responses.
420 + let max_payload = if self.role == Role::Server {
421 + self.max_request_payload_bytes
422 + } else {
423 + self.max_response_payload_bytes
424 + };
425 + if hdr.payload_len > max_payload {
426 + return Err(UdsError::LimitExceeded);
427 + }
428 +
429 + // Validate item_count against negotiated directional batch limit.
430 + let max_batch = if self.role == Role::Server {
431 + self.max_request_batch_items
432 + } else {
433 + self.max_response_batch_items
434 + };
435 + if hdr.item_count > max_batch {
436 + return Err(UdsError::LimitExceeded);
437 + }
438 +
439 + // Client-side: validate response message_id is in-flight
440 + if self.role == Role::Client && hdr.kind == KIND_RESPONSE {
441 + if !self.inflight_ids.remove(&hdr.message_id) {
442 + return Err(UdsError::UnknownMsgId(hdr.message_id));
443 + }
444 + }
445 +
446 + let total_msg = HEADER_SIZE + hdr.payload_len as usize;
447 +
448 + // Non-chunked: entire message in one packet
449 + if n >= total_msg {
450 + let payload = &buf[HEADER_SIZE..HEADER_SIZE + hdr.payload_len as usize];
451 +
452 + // Validate batch directory
453 + if hdr.flags & FLAG_BATCH != 0 && hdr.item_count > 1 {
454 + let dir_bytes = hdr.item_count as usize * 8;
455 + let dir_aligned = align8(dir_bytes);
456 + if payload.len() < dir_aligned {
457 + return Err(UdsError::Protocol("batch directory exceeds payload".into()));
458 + }
459 + let packed_area_len = (payload.len() - dir_aligned) as u32;
460 + protocol::batch_dir_validate(
461 + &payload[..dir_bytes],
462 + hdr.item_count,
463 + packed_area_len,
464 + )
465 + .map_err(|e| UdsError::Protocol(format!("batch directory: {e:?}")))?;
466 + }
467 +
468 + return Ok((hdr, payload));
469 + }
470 +
471 + // Chunked: first packet has partial payload
472 + let first_payload_bytes = n - HEADER_SIZE;
473 +
474 + // Grow recv_buf to hold full payload
475 + let needed = hdr.payload_len as usize;
476 + if self.recv_buf.len() < needed {
477 + self.recv_buf.resize(needed, 0);
478 + }
479 +
480 + // Copy first chunk's payload
481 + self.recv_buf[..first_payload_bytes]
482 + .copy_from_slice(&buf[HEADER_SIZE..HEADER_SIZE + first_payload_bytes]);
483 +
484 + let mut assembled = first_payload_bytes;
485 + let chunk_payload_budget = self.packet_size as usize - HEADER_SIZE;
486 +
487 + // Expected chunk count
488 + let remaining_after_first = hdr.payload_len as usize - first_payload_bytes;
489 + let expected_continuations = if remaining_after_first > 0 && chunk_payload_budget > 0 {
490 + (remaining_after_first + chunk_payload_budget - 1) / chunk_payload_budget
491 + } else {
492 + 0
493 + };
494 + let expected_chunk_count = 1 + expected_continuations as u32;
495 +
496 + if self.pkt_buf.len() < self.packet_size as usize {
497 + self.pkt_buf.resize(self.packet_size as usize, 0);
498 + }
499 +
500 + let mut ci = 1u32;
501 + while assembled < hdr.payload_len as usize {
502 + let cn = match raw_recv(self.fd, &mut self.pkt_buf) {
503 + Ok(n) => n,
504 + Err(err) => {
505 + self.fail_all_inflight();
506 + return Err(err);
507 + }
508 + };
509 +
510 + if cn < HEADER_SIZE {
511 + return Err(UdsError::Chunk("continuation too short".into()));
512 + }
513 +
514 + let chk = ChunkHeader::decode(&self.pkt_buf[..cn])
515 + .map_err(|e| UdsError::Chunk(format!("chunk header: {e}")))?;
516 +
517 + // Validate chunk header
518 + if chk.message_id != hdr.message_id {
519 + return Err(UdsError::Chunk("message_id mismatch".into()));
520 + }
521 + if chk.chunk_index != ci {
522 + return Err(UdsError::Chunk(format!(
523 + "chunk_index mismatch: expected {ci}, got {}",
524 + chk.chunk_index
525 + )));
526 + }
527 + if chk.chunk_count != expected_chunk_count {
528 + return Err(UdsError::Chunk("chunk_count mismatch".into()));
529 + }
530 + if chk.total_message_len != total_msg as u32 {
531 + return Err(UdsError::Chunk("total_message_len mismatch".into()));
532 + }
533 +
534 + let chunk_data = cn - HEADER_SIZE;
535 + if chunk_data != chk.chunk_payload_len as usize {
536 + return Err(UdsError::Chunk("chunk_payload_len mismatch".into()));
537 + }
538 + if assembled + chunk_data > hdr.payload_len as usize {
539 + return Err(UdsError::Chunk("chunk exceeds payload_len".into()));
540 + }
541 +
542 + self.recv_buf[assembled..assembled + chunk_data]
543 + .copy_from_slice(&self.pkt_buf[HEADER_SIZE..HEADER_SIZE + chunk_data]);
544 + assembled += chunk_data;
545 + ci += 1;
546 + }
547 +
548 + let payload = &self.recv_buf[..hdr.payload_len as usize];
549 +
550 + // Validate batch directory
551 + if hdr.flags & FLAG_BATCH != 0 && hdr.item_count > 1 {
552 + let dir_bytes = hdr.item_count as usize * 8;
553 + let dir_aligned = align8(dir_bytes);
554 + if payload.len() < dir_aligned {
555 + return Err(UdsError::Protocol("batch directory exceeds payload".into()));
556 + }
557 + let packed_area_len = (payload.len() - dir_aligned) as u32;
558 + protocol::batch_dir_validate(&payload[..dir_bytes], hdr.item_count, packed_area_len)
559 + .map_err(|e| UdsError::Protocol(format!("batch directory: {e:?}")))?;
560 + }
561 +
562 + Ok((hdr, payload))
563 + }
564 +}
565 +
566 +impl Drop for UdsSession {
567 + fn drop(&mut self) {
568 + if self.fd >= 0 {
569 + unsafe {
570 + libc::close(self.fd);
571 + }
572 + self.fd = -1;
573 + }
574 + }
575 +}
576 +
577 +// ---------------------------------------------------------------------------
578 +// Listener
579 +// ---------------------------------------------------------------------------
580 +
581 +/// A listening UDS SEQPACKET endpoint.
582 +pub struct UdsListener {
583 + fd: RawFd,
584 + config: ServerConfig,
585 + path: PathBuf,
586 + next_session_id: AtomicU64,
587 +}
588 +
589 +impl UdsListener {
590 + /// Create a listener on `{run_dir}/{service_name}.sock`.
591 + /// Performs stale endpoint recovery.
592 + pub fn bind(run_dir: &str, service_name: &str, config: ServerConfig) -> Result<Self, UdsError> {
593 + let path = build_socket_path(run_dir, service_name)?;
594 +
595 + // Stale recovery
596 + match check_and_recover_stale(&path) {
597 + StaleResult::LiveServer => return Err(UdsError::AddrInUse),
598 + StaleResult::Stale | StaleResult::NotExist => { /* proceed */ }
599 + }
600 +
601 + let fd = unsafe { libc::socket(libc::AF_UNIX, libc::SOCK_SEQPACKET, 0) };
602 + if fd < 0 {
603 + return Err(UdsError::Socket(errno()));
604 + }
605 +
606 + // Bind
607 + if let Err(e) = bind_unix(fd, &path) {
608 + unsafe {
609 + libc::close(fd);
610 + }
611 + return Err(e);
612 + }
613 +
614 + let backlog = if config.backlog > 0 {
615 + config.backlog
616 + } else {
617 + DEFAULT_BACKLOG
618 + };
619 +
620 + if unsafe { libc::listen(fd, backlog) } < 0 {
621 + let e = errno();
622 + unsafe {
623 + libc::close(fd);
624 + }
625 + let _ = std::fs::remove_file(&path);
626 + return Err(UdsError::Socket(e));
627 + }
628 +
629 + Ok(UdsListener {
630 + fd,
631 + config,
632 + path: PathBuf::from(&path),
633 + next_session_id: AtomicU64::new(1),
634 + })
635 + }
636 +
637 + /// Get the raw fd for poll/epoll integration.
638 + pub fn fd(&self) -> RawFd {
639 + self.fd
640 + }
641 +
642 + /// Update the payload limits used for future handshakes.
643 + pub fn set_payload_limits(
644 + &mut self,
645 + max_request_payload_bytes: u32,
646 + max_response_payload_bytes: u32,
647 + ) {
648 + self.config.max_request_payload_bytes = max_request_payload_bytes;
649 + self.config.max_response_payload_bytes = max_response_payload_bytes;
650 + }
651 +
652 + /// Accept one client. Performs the full handshake.
653 + /// Blocks until a client connects and the handshake completes.
654 + pub fn accept(&self) -> Result<UdsSession, UdsError> {
655 + let session_id = self.next_session_id.fetch_add(1, Ordering::Relaxed);
656 + self.accept_with_config(session_id, self.config.clone())
657 + }
658 +
659 + /// Accept one client using a caller-provided per-session server config
660 + /// and session ID.
661 + pub fn accept_with_config(
662 + &self,
663 + session_id: u64,
664 + config: ServerConfig,
665 + ) -> Result<UdsSession, UdsError> {
666 + let client_fd =
667 + unsafe { libc::accept(self.fd, std::ptr::null_mut(), std::ptr::null_mut()) };
668 + if client_fd < 0 {
669 + return Err(UdsError::Accept(errno()));
670 + }
671 +
672 + match server_handshake(client_fd, &config, session_id) {
673 + Ok(session) => Ok(session),
674 + Err(e) => {
675 + unsafe {
676 + libc::close(client_fd);
677 + }
678 + Err(e)
679 + }
680 + }
681 + }
682 +}
683 +
684 +impl Drop for UdsListener {
685 + fn drop(&mut self) {
686 + if self.fd >= 0 {
687 + unsafe {
688 + libc::close(self.fd);
689 + }
690 + self.fd = -1;
691 + }
692 + if self.path.exists() {
693 + let _ = std::fs::remove_file(&self.path);
694 + }
695 + }
696 +}
697 +
698 +// ---------------------------------------------------------------------------
699 +// Internal helpers
700 +// ---------------------------------------------------------------------------
701 +
702 +fn errno() -> i32 {
703 + io::Error::last_os_error().raw_os_error().unwrap_or(0)
704 +}
705 +
706 +/// Validate service_name: only [a-zA-Z0-9._-], non-empty, not "." or "..".
707 +fn validate_service_name(name: &str) -> Result<(), UdsError> {
708 + if name.is_empty() {
709 + return Err(UdsError::BadParam("empty service name".into()));
710 + }
711 + if name == "." || name == ".." {
712 + return Err(UdsError::BadParam(
713 + "service name cannot be '.' or '..'".into(),
714 + ));
715 + }
716 + for c in name.bytes() {
717 + match c {
718 + b'a'..=b'z' | b'A'..=b'Z' | b'0'..=b'9' | b'.' | b'_' | b'-' => {}
719 + _ => {
720 + return Err(UdsError::BadParam(format!(
721 + "service name contains invalid character: {:?}",
722 + c as char
723 + )))
724 + }
725 + }
726 + }
727 + Ok(())
728 +}
729 +
730 +/// Build `{run_dir}/{service_name}.sock` and validate length.
731 +fn build_socket_path(run_dir: &str, service_name: &str) -> Result<String, UdsError> {
732 + validate_service_name(service_name)?;
733 +
734 + let path = format!("{run_dir}/{service_name}.sock");
735 +
736 + // sun_path is typically 108 bytes on Linux, 104 on macOS
737 + let max_sun_path =
738 + std::mem::size_of::<libc::sockaddr_un>() - std::mem::size_of::<libc::sa_family_t>();
739 +
740 + if path.len() >= max_sun_path {
741 + return Err(UdsError::PathTooLong);
742 + }
743 + Ok(path)
744 +}
745 +
746 +/// Get SO_SNDBUF as packet size.
747 +fn detect_packet_size(fd: RawFd) -> u32 {
748 + let mut val: libc::c_int = 0;
749 + let mut len: libc::socklen_t = std::mem::size_of::<libc::c_int>() as libc::socklen_t;
750 +
751 + let rc = unsafe {
752 + libc::getsockopt(
753 + fd,
754 + libc::SOL_SOCKET,
755 + libc::SO_SNDBUF,
756 + &mut val as *mut _ as *mut libc::c_void,
757 + &mut len,
758 + )
759 + };
760 +
761 + if rc < 0 || val <= 0 {
762 + DEFAULT_PACKET_SIZE_FALLBACK
763 + } else {
764 + val as u32
765 + }
766 +}
767 +
768 +/// Highest set bit in a bitmask (0 if empty).
769 +fn highest_bit(mask: u32) -> u32 {
770 + if mask == 0 {
771 + return 0;
772 + }
773 + 1u32 << (31 - mask.leading_zeros() as u32)
774 +}
775 +
776 +fn apply_default(val: u32, def: u32) -> u32 {
777 + if val == 0 {
778 + def
779 + } else {
780 + val
781 + }
782 +}
783 +
784 +// ---------------------------------------------------------------------------
785 +// Low-level I/O
786 +// ---------------------------------------------------------------------------
787 +
788 +/// Send header + payload as one SEQPACKET message using sendmsg.
789 +fn raw_send_iov(fd: RawFd, hdr: &[u8], payload: &[u8]) -> Result<(), UdsError> {
790 + let total = hdr.len() + payload.len();
791 +
792 + let mut iov = [
793 + libc::iovec {
794 + iov_base: hdr.as_ptr() as *mut libc::c_void,
795 + iov_len: hdr.len(),
796 + },
797 + libc::iovec {
798 + iov_base: payload.as_ptr() as *mut libc::c_void,
799 + iov_len: payload.len(),
800 + },
801 + ];
802 +
803 + let iovcnt = if payload.is_empty() { 1 } else { 2 };
804 +
805 + let msg = libc::msghdr {
806 + msg_name: std::ptr::null_mut(),
807 + msg_namelen: 0,
808 + msg_iov: iov.as_mut_ptr(),
809 + msg_iovlen: iovcnt,
810 + msg_control: std::ptr::null_mut(),
811 + msg_controllen: 0,
812 + msg_flags: 0,
813 + };
814 +
815 + let n = unsafe { libc::sendmsg(fd, &msg, libc::MSG_NOSIGNAL) };
816 + if n < 0 || n as usize != total {
817 + return Err(UdsError::Send(errno()));
818 + }
819 + Ok(())
820 +}
821 +
822 +/// Send a contiguous buffer as one SEQPACKET message.
823 +fn raw_send(fd: RawFd, data: &[u8]) -> Result<(), UdsError> {
824 + let n = unsafe {
825 + libc::send(
826 + fd,
827 + data.as_ptr() as *const libc::c_void,
828 + data.len(),
829 + libc::MSG_NOSIGNAL,
830 + )
831 + };
832 + if n < 0 || n as usize != data.len() {
833 + return Err(UdsError::Send(errno()));
834 + }
835 + Ok(())
836 +}
837 +
838 +/// Receive one SEQPACKET message. Returns bytes received.
839 +fn raw_recv(fd: RawFd, buf: &mut [u8]) -> Result<usize, UdsError> {
840 + let n = unsafe { libc::recv(fd, buf.as_mut_ptr() as *mut libc::c_void, buf.len(), 0) };
841 + if n <= 0 {
842 + return Err(UdsError::Recv(if n == 0 { 0 } else { errno() }));
843 + }
844 + Ok(n as usize)
845 +}
846 +
847 +// ---------------------------------------------------------------------------
848 +// Socket helpers
849 +// ---------------------------------------------------------------------------
850 +
851 +/// Bind a Unix socket to the given path.
852 +fn bind_unix(fd: RawFd, path: &str) -> Result<(), UdsError> {
853 + let c_path = CString::new(path).map_err(|_| UdsError::PathTooLong)?;
854 +
855 + let mut addr: libc::sockaddr_un = unsafe { std::mem::zeroed() };
856 + addr.sun_family = libc::AF_UNIX as libc::sa_family_t;
857 +
858 + let path_bytes = c_path.as_bytes_with_nul();
859 + let sun_path_ptr = addr.sun_path.as_mut_ptr() as *mut u8;
860 + unsafe {
861 + std::ptr::copy_nonoverlapping(path_bytes.as_ptr(), sun_path_ptr, path_bytes.len());
862 + }
863 +
864 + let rc = unsafe {
865 + libc::bind(
866 + fd,
867 + &addr as *const libc::sockaddr_un as *const libc::sockaddr,
868 + std::mem::size_of::<libc::sockaddr_un>() as libc::socklen_t,
869 + )
870 + };
871 + if rc < 0 {
872 + return Err(UdsError::Socket(errno()));
873 + }
874 + Ok(())
875 +}
876 +
877 +/// Connect to a Unix socket at the given path.
878 +fn connect_unix(fd: RawFd, path: &str) -> Result<(), UdsError> {
879 + let c_path = CString::new(path).map_err(|_| UdsError::PathTooLong)?;
880 +
881 + let mut addr: libc::sockaddr_un = unsafe { std::mem::zeroed() };
882 + addr.sun_family = libc::AF_UNIX as libc::sa_family_t;
883 +
884 + let path_bytes = c_path.as_bytes_with_nul();
885 + let sun_path_ptr = addr.sun_path.as_mut_ptr() as *mut u8;
886 + unsafe {
887 + std::ptr::copy_nonoverlapping(path_bytes.as_ptr(), sun_path_ptr, path_bytes.len());
888 + }
889 +
890 + let rc = unsafe {
891 + libc::connect(
892 + fd,
893 + &addr as *const libc::sockaddr_un as *const libc::sockaddr,
894 + std::mem::size_of::<libc::sockaddr_un>() as libc::socklen_t,
895 + )
896 + };
897 + if rc < 0 {
898 + return Err(UdsError::Connect(errno()));
899 + }
900 + Ok(())
901 +}
902 +
903 +// ---------------------------------------------------------------------------
904 +// Stale endpoint recovery
905 +// ---------------------------------------------------------------------------
906 +
907 +enum StaleResult {
908 + NotExist,
909 + Stale,
910 + LiveServer,
911 +}
912 +
913 +fn check_and_recover_stale(path: &str) -> StaleResult {
914 + if !Path::new(path).exists() {
915 + return StaleResult::NotExist;
916 + }
917 +
918 + // Try connecting to see if a live server is there
919 + let probe = unsafe { libc::socket(libc::AF_UNIX, libc::SOCK_SEQPACKET, 0) };
920 + if probe < 0 {
921 + return StaleResult::NotExist;
922 + }
923 +
924 + let result = match connect_unix(probe, path) {
925 + Ok(()) => {
926 + // Connected => live server
927 + StaleResult::LiveServer
928 + }
929 + Err(UdsError::Connect(e)) if e == libc::ECONNREFUSED || e == libc::ENOENT => {
930 + // Connection refused or no such socket => stale, unlink
931 + let _ = std::fs::remove_file(path);
932 + StaleResult::Stale
933 + }
934 + Err(_) => {
935 + // Other errors (EACCES, etc.) — can't determine ownership,
936 + // treat as live to prevent overwriting
937 + StaleResult::LiveServer
938 + }
939 + };
940 +
941 + unsafe {
942 + libc::close(probe);
943 + }
944 + result
945 +}
946 +
947 +// ---------------------------------------------------------------------------
948 +// Handshake: client side
949 +// ---------------------------------------------------------------------------
950 +
951 +fn connect_and_handshake(
952 + fd: RawFd,
953 + path: &str,
954 + config: &ClientConfig,
955 +) -> Result<UdsSession, UdsError> {
956 + connect_unix(fd, path)?;
957 +
958 + let pkt_size = if config.packet_size == 0 {
959 + detect_packet_size(fd)
960 + } else {
961 + config.packet_size
962 + };
963 +
964 + let supported = if config.supported_profiles == 0 {
965 + PROFILE_BASELINE
966 + } else {
967 + config.supported_profiles
968 + };
969 +
970 + // Build HELLO
971 + let hello = Hello {
972 + layout_version: 1,
973 + flags: 0,
974 + supported_profiles: supported,
975 + preferred_profiles: config.preferred_profiles,
976 + max_request_payload_bytes: apply_default(
977 + config.max_request_payload_bytes,
978 + MAX_PAYLOAD_DEFAULT,
979 + ),
980 + max_request_batch_items: apply_default(config.max_request_batch_items, DEFAULT_BATCH_ITEMS),
981 + max_response_payload_bytes: apply_default(
982 + config.max_response_payload_bytes,
983 + MAX_PAYLOAD_DEFAULT,
984 + ),
985 + max_response_batch_items: apply_default(
986 + config.max_response_batch_items,
987 + DEFAULT_BATCH_ITEMS,
988 + ),
989 + auth_token: config.auth_token,
990 + packet_size: pkt_size,
991 + };
992 +
993 + let mut hello_buf = [0u8; HELLO_PAYLOAD_SIZE];
994 + hello.encode(&mut hello_buf);
995 +
996 + // Build outer CONTROL header
997 + let hdr = Header {
998 + magic: MAGIC_MSG,
999 + version: VERSION,
1000 + header_len: protocol::HEADER_LEN,
1001 + kind: protocol::KIND_CONTROL,
1002 + flags: 0,
1003 + code: protocol::CODE_HELLO,
1004 + transport_status: protocol::STATUS_OK,
1005 + payload_len: HELLO_PAYLOAD_SIZE as u32,
1006 + item_count: 1,
1007 + message_id: 0,
1008 + };
1009 +
1010 + let mut pkt = [0u8; HEADER_SIZE + HELLO_PAYLOAD_SIZE];
1011 + hdr.encode(&mut pkt[..HEADER_SIZE]);
1012 + pkt[HEADER_SIZE..].copy_from_slice(&hello_buf);
1013 +
1014 + // Send HELLO
1015 + raw_send(fd, &pkt)?;
1016 +
1017 + // Receive HELLO_ACK
1018 + let mut buf = [0u8; 128];
1019 + let n = raw_recv(fd, &mut buf)?;
1020 +
1021 + // Decode outer header
1022 + let ack_hdr = match Header::decode(&buf[..n]) {
1023 + Ok(hdr) => hdr,
1024 + Err(crate::protocol::NipcError::BadVersion) => {
1025 + return Err(UdsError::Incompatible("ack header version mismatch".into()))
1026 + }
1027 + Err(e) => return Err(UdsError::Protocol(format!("ack header: {e}"))),
1028 + };
1029 +
1030 + if ack_hdr.kind != protocol::KIND_CONTROL || ack_hdr.code != protocol::CODE_HELLO_ACK {
1031 + return Err(UdsError::Protocol("expected HELLO_ACK".into()));
1032 + }
1033 +
1034 + // Check transport_status for rejection
1035 + if ack_hdr.transport_status == protocol::STATUS_AUTH_FAILED {
1036 + return Err(UdsError::AuthFailed);
1037 + }
1038 + if ack_hdr.transport_status == protocol::STATUS_UNSUPPORTED {
1039 + return Err(UdsError::NoProfile);
1040 + }
1041 + if ack_hdr.transport_status == protocol::STATUS_INCOMPATIBLE {
1042 + return Err(UdsError::Incompatible(
1043 + "ack transport_status incompatible".into(),
1044 + ));
1045 + }
1046 + if ack_hdr.transport_status == protocol::STATUS_LIMIT_EXCEEDED {
1047 + return Err(UdsError::LimitExceeded);
1048 + }
1049 + if ack_hdr.transport_status != protocol::STATUS_OK {
1050 + return Err(UdsError::Handshake(format!(
1051 + "transport_status={}",
1052 + ack_hdr.transport_status
1053 + )));
1054 + }
1055 +
1056 + // Decode hello-ack payload
1057 + if n < HEADER_SIZE + HELLO_ACK_PAYLOAD_SIZE {
1058 + return Err(UdsError::Protocol("ack payload truncated".into()));
1059 + }
1060 + let ack = match HelloAck::decode(&buf[HEADER_SIZE..n]) {
1061 + Ok(ack) => ack,
1062 + Err(crate::protocol::NipcError::BadLayout)
1063 + if hello_ack_layout_incompatible(&buf[HEADER_SIZE..n]) =>
1064 + {
1065 + return Err(UdsError::Incompatible(
1066 + "ack payload layout version mismatch".into(),
1067 + ))
1068 + }
1069 + Err(e) => return Err(UdsError::Protocol(format!("ack payload: {e}"))),
1070 + };
1071 +
1072 + // Sanity: reject a packet_size too small for chunking arithmetic
1073 + if ack.agreed_packet_size <= HEADER_SIZE as u32 {
1074 + return Err(UdsError::Protocol("agreed packet_size too small".into()));
1075 + }
1076 +
1077 + Ok(UdsSession {
1078 + fd,
1079 + role: Role::Client,
1080 + max_request_payload_bytes: ack.agreed_max_request_payload_bytes,
1081 + max_request_batch_items: ack.agreed_max_request_batch_items,
1082 + max_response_payload_bytes: ack.agreed_max_response_payload_bytes,
1083 + max_response_batch_items: ack.agreed_max_response_batch_items,
1084 + packet_size: ack.agreed_packet_size,
1085 + selected_profile: ack.selected_profile,
1086 + session_id: ack.session_id,
1087 + recv_buf: Vec::new(),
1088 + pkt_buf: Vec::new(),
1089 + inflight_ids: HashSet::new(),
1090 + })
1091 +}
1092 +
1093 +// ---------------------------------------------------------------------------
1094 +// Handshake: server side
1095 +// ---------------------------------------------------------------------------
1096 +
1097 +fn server_handshake(
1098 + fd: RawFd,
1099 + config: &ServerConfig,
1100 + session_id: u64,
1101 +) -> Result<UdsSession, UdsError> {
1102 + let server_pkt_size = if config.packet_size == 0 {
1103 + detect_packet_size(fd)
1104 + } else {
1105 + config.packet_size
1106 + };
1107 +
1108 + let s_resp_pay = apply_default(config.max_response_payload_bytes, MAX_PAYLOAD_DEFAULT);
1109 + let s_profiles = if config.supported_profiles == 0 {
1110 + PROFILE_BASELINE
1111 + } else {
1112 + config.supported_profiles
1113 + };
1114 + let s_preferred = config.preferred_profiles;
1115 +
1116 + let send_rejection = |status: u16| -> Result<(), UdsError> {
1117 + let ack = HelloAck {
1118 + layout_version: 1,
1119 + ..HelloAck::default()
1120 + };
1121 + let mut ack_buf = [0u8; HELLO_ACK_PAYLOAD_SIZE];
1122 + ack.encode(&mut ack_buf);
1123 +
1124 + let ack_hdr = Header {
1125 + magic: MAGIC_MSG,
1126 + version: VERSION,
1127 + header_len: protocol::HEADER_LEN,
1128 + kind: protocol::KIND_CONTROL,
1129 + flags: 0,
1130 + code: protocol::CODE_HELLO_ACK,
1131 + transport_status: status,
1132 + payload_len: HELLO_ACK_PAYLOAD_SIZE as u32,
1133 + item_count: 1,
1134 + message_id: 0,
1135 + };
1136 +
1137 + let mut pkt = [0u8; HEADER_SIZE + HELLO_ACK_PAYLOAD_SIZE];
1138 + ack_hdr.encode(&mut pkt[..HEADER_SIZE]);
1139 + pkt[HEADER_SIZE..].copy_from_slice(&ack_buf);
1140 + let _ = raw_send(fd, &pkt);
1141 + Ok(())
1142 + };
1143 +
1144 + // Receive HELLO
1145 + let mut buf = [0u8; 128];
1146 + let n = raw_recv(fd, &mut buf)?;
1147 +
1148 + let hdr = match Header::decode(&buf[..n]) {
1149 + Ok(hdr) => hdr,
1150 + Err(crate::protocol::NipcError::BadVersion)
1151 + if header_version_incompatible(&buf[..n], protocol::CODE_HELLO) =>
1152 + {
1153 + send_rejection(protocol::STATUS_INCOMPATIBLE)?;
1154 + return Err(UdsError::Incompatible(
1155 + "hello header version mismatch".into(),
1156 + ));
1157 + }
1158 + Err(e) => return Err(UdsError::Protocol(format!("hello header: {e}"))),
1159 + };
1160 +
1161 + if hdr.kind != protocol::KIND_CONTROL || hdr.code != protocol::CODE_HELLO {
1162 + return Err(UdsError::Protocol("expected HELLO".into()));
1163 + }
1164 +
1165 + let hello = match Hello::decode(&buf[HEADER_SIZE..n]) {
1166 + Ok(hello) => hello,
1167 + Err(crate::protocol::NipcError::BadLayout)
1168 + if hello_layout_incompatible(&buf[HEADER_SIZE..n]) =>
1169 + {
1170 + send_rejection(protocol::STATUS_INCOMPATIBLE)?;
1171 + return Err(UdsError::Incompatible(
1172 + "hello payload layout version mismatch".into(),
1173 + ));
1174 + }
1175 + Err(e) => return Err(UdsError::Protocol(format!("hello payload: {e}"))),
1176 + };
1177 +
1178 + // Compute intersection
1179 + let intersection = hello.supported_profiles & s_profiles;
1180 +
1181 + // Check intersection
1182 + if intersection == 0 {
1183 + send_rejection(protocol::STATUS_UNSUPPORTED)?;
1184 + return Err(UdsError::NoProfile);
1185 + }
1186 +
1187 + // Check auth
1188 + if hello.auth_token != config.auth_token {
1189 + send_rejection(protocol::STATUS_AUTH_FAILED)?;
1190 + return Err(UdsError::AuthFailed);
1191 + }
1192 +
1193 + // Select profile
1194 + let preferred_intersection = intersection & hello.preferred_profiles & s_preferred;
1195 + let selected = if preferred_intersection != 0 {
1196 + highest_bit(preferred_intersection)
1197 + } else {
1198 + highest_bit(intersection)
1199 + };
1200 +
1201 + if hello.max_request_payload_bytes > MAX_PAYLOAD_CAP {
1202 + send_rejection(protocol::STATUS_LIMIT_EXCEEDED)?;
1203 + return Err(UdsError::LimitExceeded);
1204 + }
1205 +
1206 + // Negotiate limits:
1207 + // - request payload and batch size are client-proposed and echoed
1208 + // - response payload is server-authoritative
1209 + // - response batch size is symmetric with request batch size
1210 + let agreed_req_pay = hello.max_request_payload_bytes;
1211 + let agreed_req_bat = hello.max_request_batch_items;
1212 + let agreed_resp_pay = s_resp_pay;
1213 + let agreed_resp_bat = agreed_req_bat;
1214 + let agreed_pkt = hello.packet_size.min(server_pkt_size);
1215 +
1216 + // packet_size must be large enough for a usable message packet
1217 + if agreed_pkt <= HEADER_SIZE as u32 {
1218 + send_rejection(protocol::STATUS_INCOMPATIBLE)?;
1219 + return Err(UdsError::Incompatible(
1220 + "packet size too small for negotiated session".into(),
1221 + ));
1222 + }
1223 +
1224 + // Send HELLO_ACK (success)
1225 + let ack = HelloAck {
1226 + layout_version: 1,
1227 + flags: 0,
1228 + server_supported_profiles: s_profiles,
1229 + intersection_profiles: intersection,
1230 + selected_profile: selected,
1231 + agreed_max_request_payload_bytes: agreed_req_pay,
1232 + agreed_max_request_batch_items: agreed_req_bat,
1233 + agreed_max_response_payload_bytes: agreed_resp_pay,
1234 + agreed_max_response_batch_items: agreed_resp_bat,
1235 + agreed_packet_size: agreed_pkt,
1236 + session_id,
1237 + };
1238 +
1239 + let mut ack_buf = [0u8; HELLO_ACK_PAYLOAD_SIZE];
1240 + ack.encode(&mut ack_buf);
1241 +
1242 + let ack_hdr = Header {
1243 + magic: MAGIC_MSG,
1244 + version: VERSION,
1245 + header_len: protocol::HEADER_LEN,
1246 + kind: protocol::KIND_CONTROL,
1247 + flags: 0,
1248 + code: protocol::CODE_HELLO_ACK,
1249 + transport_status: protocol::STATUS_OK,
1250 + payload_len: HELLO_ACK_PAYLOAD_SIZE as u32,
1251 + item_count: 1,
1252 + message_id: 0,
1253 + };
1254 +
1255 + let mut pkt = [0u8; HEADER_SIZE + HELLO_ACK_PAYLOAD_SIZE];
1256 + ack_hdr.encode(&mut pkt[..HEADER_SIZE]);
1257 + pkt[HEADER_SIZE..].copy_from_slice(&ack_buf);
1258 + raw_send(fd, &pkt)?;
1259 +
1260 + Ok(UdsSession {
1261 + fd,
1262 + role: Role::Server,
1263 + max_request_payload_bytes: agreed_req_pay,
1264 + max_request_batch_items: agreed_req_bat,
1265 + max_response_payload_bytes: agreed_resp_pay,
1266 + max_response_batch_items: agreed_resp_bat,
1267 + packet_size: agreed_pkt,
1268 + selected_profile: selected,
1269 + session_id,
1270 + recv_buf: Vec::new(),
1271 + pkt_buf: Vec::new(),
1272 + inflight_ids: HashSet::new(),
1273 + })
1274 +}
1275 +
1276 +// ---------------------------------------------------------------------------
1277 +// Tests
1278 +// ---------------------------------------------------------------------------
1279 +
1280 +#[cfg(test)]
1281 +#[path = "posix_tests.rs"]
1282 +mod tests;
src/crates/netipc/src/transport/posix_tests.rs new
+2725
@@ -0,0 +1,2725 @@
1 +use super::*;
2 +use crate::protocol;
3 +use std::sync::atomic::{AtomicBool, AtomicU32, Ordering};
4 +use std::sync::Arc;
5 +use std::thread;
6 +use std::time::Duration;
7 +
8 +const TEST_RUN_DIR: &str = "/tmp/nipc_rust_test";
9 +const AUTH_TOKEN: u64 = 0xDEADBEEFCAFEBABE;
10 +
11 +fn ensure_run_dir() {
12 + let _ = std::fs::create_dir_all(TEST_RUN_DIR);
13 +}
14 +
15 +fn cleanup_socket(service: &str) {
16 + let path = format!("{TEST_RUN_DIR}/{service}.sock");
17 + let _ = std::fs::remove_file(&path);
18 +}
19 +
20 +fn default_server_config() -> ServerConfig {
21 + ServerConfig {
22 + supported_profiles: PROFILE_BASELINE,
23 + preferred_profiles: 0,
24 + max_request_payload_bytes: 4096,
25 + max_request_batch_items: 16,
26 + max_response_payload_bytes: 4096,
27 + max_response_batch_items: 16,
28 + auth_token: AUTH_TOKEN,
29 + packet_size: 0,
30 + backlog: 4,
31 + }
32 +}
33 +
34 +fn default_client_config() -> ClientConfig {
35 + ClientConfig {
36 + supported_profiles: PROFILE_BASELINE,
37 + preferred_profiles: 0,
38 + max_request_payload_bytes: 4096,
39 + max_request_batch_items: 16,
40 + max_response_payload_bytes: 4096,
41 + max_response_batch_items: 16,
42 + auth_token: AUTH_TOKEN,
43 + packet_size: 0,
44 + }
45 +}
46 +
47 +/// Unique service name to avoid parallel test collisions.
48 +static TEST_COUNTER: AtomicU32 = AtomicU32::new(0);
49 +fn unique_service(prefix: &str) -> String {
50 + let n = TEST_COUNTER.fetch_add(1, Ordering::Relaxed);
51 + format!("{prefix}_{n}_{}", std::process::id())
52 +}
53 +
54 +fn socketpair_seqpacket() -> (RawFd, RawFd) {
55 + let mut fds = [-1; 2];
56 + let rc = unsafe { libc::socketpair(libc::AF_UNIX, libc::SOCK_SEQPACKET, 0, fds.as_mut_ptr()) };
57 + assert_eq!(rc, 0, "socketpair failed: {}", errno());
58 + (fds[0], fds[1])
59 +}
60 +
61 +fn test_session(fd: RawFd, role: Role, packet_size: u32) -> UdsSession {
62 + UdsSession {
63 + fd,
64 + role,
65 + max_request_payload_bytes: 4096,
66 + max_request_batch_items: 16,
67 + max_response_payload_bytes: 4096,
68 + max_response_batch_items: 16,
69 + packet_size,
70 + selected_profile: PROFILE_BASELINE,
71 + session_id: 1,
72 + recv_buf: Vec::new(),
73 + pkt_buf: Vec::new(),
74 + inflight_ids: HashSet::new(),
75 + }
76 +}
77 +
78 +fn raw_listener_fd(service: &str) -> RawFd {
79 + let path = build_socket_path(TEST_RUN_DIR, service).expect("socket path");
80 + let fd = unsafe { libc::socket(libc::AF_UNIX, libc::SOCK_SEQPACKET, 0) };
81 + assert!(fd >= 0, "socket failed: {}", errno());
82 + bind_unix(fd, &path).expect("bind raw listener");
83 + let rc = unsafe { libc::listen(fd, DEFAULT_BACKLOG) };
84 + assert_eq!(rc, 0, "listen failed: {}", errno());
85 + fd
86 +}
87 +
88 +// -----------------------------------------------------------------------
89 +// Test 1: Single client ping-pong
90 +// -----------------------------------------------------------------------
91 +
92 +#[test]
93 +fn test_ping_pong() {
94 + ensure_run_dir();
95 + let svc = unique_service("rs_ping");
96 + cleanup_socket(&svc);
97 +
98 + let svc_clone = svc.clone();
99 + let ready = Arc::new(AtomicBool::new(false));
100 + let ready_clone = ready.clone();
101 +
102 + let server_thread = thread::spawn(move || {
103 + let listener =
104 + UdsListener::bind(TEST_RUN_DIR, &svc_clone, default_server_config()).expect("listen");
105 + ready_clone.store(true, Ordering::Release);
106 +
107 + let mut session = listener.accept().expect("accept");
108 +
109 + let mut buf = [0u8; 8192];
110 + let (hdr, payload) = session.receive(&mut buf).expect("recv");
111 + let payload = payload.to_vec();
112 +
113 + // Echo back as response
114 + let mut resp = hdr;
115 + resp.kind = protocol::KIND_RESPONSE;
116 + resp.transport_status = protocol::STATUS_OK;
117 + session.send(&mut resp, &payload).expect("send");
118 + });
119 +
120 + // Wait for server
121 + while !ready.load(Ordering::Acquire) {
122 + thread::sleep(Duration::from_millis(1));
123 + }
124 +
125 + let mut session =
126 + UdsSession::connect(TEST_RUN_DIR, &svc, &default_client_config()).expect("connect");
127 +
128 + assert_eq!(session.selected_profile, PROFILE_BASELINE);
129 +
130 + let payload = [0x01u8, 0x02, 0x03, 0x04];
131 + let mut hdr = Header {
132 + kind: protocol::KIND_REQUEST,
133 + code: protocol::METHOD_INCREMENT,
134 + flags: 0,
135 + item_count: 1,
136 + message_id: 42,
137 + ..Header::default()
138 + };
139 +
140 + session.send(&mut hdr, &payload).expect("send");
141 +
142 + let mut rbuf = [0u8; 4096];
143 + let (rhdr, rpayload) = session.receive(&mut rbuf).expect("recv");
144 + assert_eq!(rhdr.kind, protocol::KIND_RESPONSE);
145 + assert_eq!(rhdr.message_id, 42);
146 + assert_eq!(rpayload, payload);
147 +
148 + drop(session);
149 + server_thread.join().expect("server join");
150 + cleanup_socket(&svc);
151 +}
152 +
153 +// -----------------------------------------------------------------------
154 +// Test 2: Multi-client (2 clients, 1 listener)
155 +// -----------------------------------------------------------------------
156 +
157 +#[test]
158 +fn test_multi_client() {
159 + ensure_run_dir();
160 + let svc = unique_service("rs_multi");
161 + cleanup_socket(&svc);
162 +
163 + let svc_clone = svc.clone();
164 + let ready = Arc::new(AtomicBool::new(false));
165 + let ready_clone = ready.clone();
166 +
167 + let server_thread = thread::spawn(move || {
168 + let listener =
169 + UdsListener::bind(TEST_RUN_DIR, &svc_clone, default_server_config()).expect("listen");
170 + ready_clone.store(true, Ordering::Release);
171 +
172 + for _ in 0..2 {
173 + let mut session = listener.accept().expect("accept");
174 + let mut buf = [0u8; 8192];
175 + let (hdr, payload) = session.receive(&mut buf).expect("recv");
176 + let payload = payload.to_vec();
177 + let mut resp = hdr;
178 + resp.kind = protocol::KIND_RESPONSE;
179 + resp.transport_status = protocol::STATUS_OK;
180 + session.send(&mut resp, &payload).expect("send");
181 + }
182 + });
183 +
184 + while !ready.load(Ordering::Acquire) {
185 + thread::sleep(Duration::from_millis(1));
186 + }
187 +
188 + let results: Vec<_> = (0..2)
189 + .map(|i| {
190 + let svc = svc.clone();
191 + thread::spawn(move || {
192 + let mut session = UdsSession::connect(TEST_RUN_DIR, &svc, &default_client_config())
193 + .expect("connect");
194 +
195 + let payload = [0xAA + i as u8];
196 + let msg_id = 100 + i as u64;
197 + let mut hdr = Header {
198 + kind: protocol::KIND_REQUEST,
199 + code: 1,
200 + item_count: 1,
201 + message_id: msg_id,
202 + ..Header::default()
203 + };
204 + session.send(&mut hdr, &payload).expect("send");
205 +
206 + let mut rbuf = [0u8; 4096];
207 + let (rhdr, rpayload) = session.receive(&mut rbuf).expect("recv");
208 + assert_eq!(rhdr.message_id, msg_id);
209 + assert_eq!(rpayload, payload);
210 + })
211 + })
212 + .collect();
213 +
214 + for t in results {
215 + t.join().expect("client join");
216 + }
217 +
218 + server_thread.join().expect("server join");
219 + cleanup_socket(&svc);
220 +}
221 +
222 +// -----------------------------------------------------------------------
223 +// Test 3: Pipelining (send N, receive N)
224 +// -----------------------------------------------------------------------
225 +
226 +#[test]
227 +fn test_pipelining() {
228 + ensure_run_dir();
229 + let svc = unique_service("rs_pipe");
230 + cleanup_socket(&svc);
231 +
232 + let svc_clone = svc.clone();
233 + let ready = Arc::new(AtomicBool::new(false));
234 + let ready_clone = ready.clone();
235 +
236 + let server_thread = thread::spawn(move || {
237 + let listener =
238 + UdsListener::bind(TEST_RUN_DIR, &svc_clone, default_server_config()).expect("listen");
239 + ready_clone.store(true, Ordering::Release);
240 +
241 + let mut session = listener.accept().expect("accept");
242 +
243 + for _ in 0..3 {
244 + let mut buf = [0u8; 8192];
245 + let (hdr, payload) = session.receive(&mut buf).expect("recv");
246 + let payload = payload.to_vec();
247 + let mut resp = hdr;
248 + resp.kind = protocol::KIND_RESPONSE;
249 + resp.transport_status = protocol::STATUS_OK;
250 + session.send(&mut resp, &payload).expect("send");
251 + }
252 + });
253 +
254 + while !ready.load(Ordering::Acquire) {
255 + thread::sleep(Duration::from_millis(1));
256 + }
257 +
258 + let mut session =
259 + UdsSession::connect(TEST_RUN_DIR, &svc, &default_client_config()).expect("connect");
260 +
261 + // Send 3 requests
262 + for i in 1u64..=3 {
263 + let payload = [i as u8];
264 + let mut hdr = Header {
265 + kind: protocol::KIND_REQUEST,
266 + code: 1,
267 + item_count: 1,
268 + message_id: i,
269 + ..Header::default()
270 + };
271 + session.send(&mut hdr, &payload).expect("send");
272 + }
273 +
274 + // Receive 3 responses
275 + for i in 1u64..=3 {
276 + let mut rbuf = [0u8; 4096];
277 + let (rhdr, rpayload) = session.receive(&mut rbuf).expect("recv");
278 + assert_eq!(rhdr.message_id, i);
279 + assert_eq!(rpayload, [i as u8]);
280 + }
281 +
282 + drop(session);
283 + server_thread.join().expect("server join");
284 + cleanup_socket(&svc);
285 +}
286 +
287 +// -----------------------------------------------------------------------
288 +// Test 4: Chunking (large message with small packet_size)
289 +// -----------------------------------------------------------------------
290 +
291 +#[test]
292 +fn test_chunking() {
293 + ensure_run_dir();
294 + let svc = unique_service("rs_chunk");
295 + cleanup_socket(&svc);
296 +
297 + let svc_clone = svc.clone();
298 + let ready = Arc::new(AtomicBool::new(false));
299 + let ready_clone = ready.clone();
300 +
301 + let server_thread = thread::spawn(move || {
302 + let scfg = ServerConfig {
303 + packet_size: 128,
304 + max_request_payload_bytes: 65536,
305 + max_response_payload_bytes: 65536,
306 + ..default_server_config()
307 + };
308 + let listener = UdsListener::bind(TEST_RUN_DIR, &svc_clone, scfg).expect("listen");
309 + ready_clone.store(true, Ordering::Release);
310 +
311 + let mut session = listener.accept().expect("accept");
312 +
313 + let mut buf = [0u8; 256]; // small buf, forces recv_buf usage
314 + let (hdr, payload) = session.receive(&mut buf).expect("recv");
315 + let payload = payload.to_vec();
316 +
317 + let mut resp = hdr;
318 + resp.kind = protocol::KIND_RESPONSE;
319 + session.send(&mut resp, &payload).expect("send");
320 + });
321 +
322 + while !ready.load(Ordering::Acquire) {
323 + thread::sleep(Duration::from_millis(1));
324 + }
325 +
326 + let ccfg = ClientConfig {
327 + packet_size: 128,
328 + max_request_payload_bytes: 65536,
329 + max_response_payload_bytes: 65536,
330 + ..default_client_config()
331 + };
332 + let mut session = UdsSession::connect(TEST_RUN_DIR, &svc, &ccfg).expect("connect");
333 +
334 + assert_eq!(session.packet_size, 128);
335 +
336 + // Build a payload larger than 128 - 32 = 96 bytes
337 + let big_len = 500;
338 + let big: Vec<u8> = (0..big_len).map(|i| (i & 0xFF) as u8).collect();
339 +
340 + let mut hdr = Header {
341 + kind: protocol::KIND_REQUEST,
342 + code: 1,
343 + item_count: 1,
344 + message_id: 7,
345 + ..Header::default()
346 + };
347 +
348 + session.send(&mut hdr, &big).expect("send chunked");
349 +
350 + let mut rbuf = [0u8; 256];
351 + let (rhdr, rpayload) = session.receive(&mut rbuf).expect("recv chunked");
352 + assert_eq!(rhdr.message_id, 7);
353 + assert_eq!(rpayload.len(), big_len);
354 + assert_eq!(rpayload, big);
355 +
356 + drop(session);
357 + server_thread.join().expect("server join");
358 + cleanup_socket(&svc);
359 +}
360 +
361 +// -----------------------------------------------------------------------
362 +// Test 5: Handshake failure - bad auth
363 +// -----------------------------------------------------------------------
364 +
365 +#[test]
366 +fn test_bad_auth() {
367 + ensure_run_dir();
368 + let svc = unique_service("rs_badauth");
369 + cleanup_socket(&svc);
370 +
371 + let svc_clone = svc.clone();
372 + let ready = Arc::new(AtomicBool::new(false));
373 + let ready_clone = ready.clone();
374 +
375 + let server_thread = thread::spawn(move || {
376 + let listener =
377 + UdsListener::bind(TEST_RUN_DIR, &svc_clone, default_server_config()).expect("listen");
378 + ready_clone.store(true, Ordering::Release);
379 + // accept will fail due to auth mismatch
380 + let _ = listener.accept();
381 + });
382 +
383 + while !ready.load(Ordering::Acquire) {
384 + thread::sleep(Duration::from_millis(1));
385 + }
386 +
387 + let ccfg = ClientConfig {
388 + auth_token: 0xBAD,
389 + ..default_client_config()
390 + };
391 +
392 + let result = UdsSession::connect(TEST_RUN_DIR, &svc, &ccfg);
393 + assert!(matches!(result, Err(UdsError::AuthFailed)));
394 +
395 + server_thread.join().expect("server join");
396 + cleanup_socket(&svc);
397 +}
398 +
399 +// -----------------------------------------------------------------------
400 +// Test 6: Handshake failure - profile mismatch
401 +// -----------------------------------------------------------------------
402 +
403 +#[test]
404 +fn test_profile_mismatch() {
405 + ensure_run_dir();
406 + let svc = unique_service("rs_badprof");
407 + cleanup_socket(&svc);
408 +
409 + let svc_clone = svc.clone();
410 + let ready = Arc::new(AtomicBool::new(false));
411 + let ready_clone = ready.clone();
412 +
413 + let server_thread = thread::spawn(move || {
414 + let scfg = ServerConfig {
415 + supported_profiles: protocol::PROFILE_SHM_FUTEX,
416 + ..default_server_config()
417 + };
418 + let listener = UdsListener::bind(TEST_RUN_DIR, &svc_clone, scfg).expect("listen");
419 + ready_clone.store(true, Ordering::Release);
420 + let _ = listener.accept();
421 + });
422 +
423 + while !ready.load(Ordering::Acquire) {
424 + thread::sleep(Duration::from_millis(1));
425 + }
426 +
427 + let ccfg = ClientConfig {
428 + supported_profiles: PROFILE_BASELINE,
429 + ..default_client_config()
430 + };
431 +
432 + let result = UdsSession::connect(TEST_RUN_DIR, &svc, &ccfg);
433 + assert!(matches!(result, Err(UdsError::NoProfile)));
434 +
435 + server_thread.join().expect("server join");
436 + cleanup_socket(&svc);
437 +}
438 +
439 +#[test]
440 +fn test_request_payload_over_cap() {
441 + ensure_run_dir();
442 + let svc = unique_service("rs_reqcap");
443 + cleanup_socket(&svc);
444 +
445 + let svc_clone = svc.clone();
446 + let ready = Arc::new(AtomicBool::new(false));
447 + let ready_clone = ready.clone();
448 +
449 + let server_thread = thread::spawn(move || {
450 + let listener =
451 + UdsListener::bind(TEST_RUN_DIR, &svc_clone, default_server_config()).expect("listen");
452 + ready_clone.store(true, Ordering::Release);
453 + let result = listener.accept();
454 + assert!(matches!(result, Err(UdsError::LimitExceeded)));
455 + });
456 +
457 + while !ready.load(Ordering::Acquire) {
458 + thread::sleep(Duration::from_millis(1));
459 + }
460 +
461 + let ccfg = ClientConfig {
462 + max_request_payload_bytes: protocol::MAX_PAYLOAD_CAP + 1,
463 + ..default_client_config()
464 + };
465 + let result = UdsSession::connect(TEST_RUN_DIR, &svc, &ccfg);
466 + assert!(matches!(result, Err(UdsError::LimitExceeded)));
467 +
468 + server_thread.join().expect("server join");
469 + cleanup_socket(&svc);
470 +}
471 +
472 +// -----------------------------------------------------------------------
473 +// Test 7: Stale socket recovery
474 +// -----------------------------------------------------------------------
475 +
476 +#[test]
477 +fn test_stale_recovery() {
478 + ensure_run_dir();
479 + let svc = unique_service("rs_stale");
480 + cleanup_socket(&svc);
481 +
482 + let path = format!("{TEST_RUN_DIR}/{svc}.sock");
483 +
484 + // Create a stale socket file (bound but not listening)
485 + unsafe {
486 + let sock = libc::socket(libc::AF_UNIX, libc::SOCK_SEQPACKET, 0);
487 + assert!(sock >= 0);
488 +
489 + let c_path = CString::new(path.as_str()).unwrap();
490 + let mut addr: libc::sockaddr_un = std::mem::zeroed();
491 + addr.sun_family = libc::AF_UNIX as libc::sa_family_t;
492 + let path_bytes = c_path.as_bytes_with_nul();
493 + let sun_path_ptr = addr.sun_path.as_mut_ptr() as *mut u8;
494 + std::ptr::copy_nonoverlapping(path_bytes.as_ptr(), sun_path_ptr, path_bytes.len());
495 +
496 + libc::bind(
497 + sock,
498 + &addr as *const libc::sockaddr_un as *const libc::sockaddr,
499 + std::mem::size_of::<libc::sockaddr_un>() as libc::socklen_t,
500 + );
501 + // Close without unlink => stale
502 + libc::close(sock);
503 + }
504 +
505 + assert!(Path::new(&path).exists(), "stale socket should exist");
506 +
507 + // listen should recover it
508 + let listener = UdsListener::bind(TEST_RUN_DIR, &svc, default_server_config())
509 + .expect("listen should recover stale socket");
510 + drop(listener);
511 + cleanup_socket(&svc);
512 +}
513 +
514 +// -----------------------------------------------------------------------
515 +// Test 8: Disconnect detection
516 +// -----------------------------------------------------------------------
517 +
518 +#[test]
519 +fn test_disconnect_detection() {
520 + ensure_run_dir();
521 + let svc = unique_service("rs_disc");
522 + cleanup_socket(&svc);
523 +
524 + let svc_clone = svc.clone();
525 + let ready = Arc::new(AtomicBool::new(false));
526 + let ready_clone = ready.clone();
527 +
528 + let server_thread = thread::spawn(move || {
529 + let listener =
530 + UdsListener::bind(TEST_RUN_DIR, &svc_clone, default_server_config()).expect("listen");
531 + ready_clone.store(true, Ordering::Release);
532 +
533 + let mut session = listener.accept().expect("accept");
534 + // Read request then close without responding
535 + let mut buf = [0u8; 4096];
536 + let _ = session.receive(&mut buf);
537 + drop(session); // close socket
538 + });
539 +
540 + while !ready.load(Ordering::Acquire) {
541 + thread::sleep(Duration::from_millis(1));
542 + }
543 +
544 + let mut session =
545 + UdsSession::connect(TEST_RUN_DIR, &svc, &default_client_config()).expect("connect");
546 +
547 + let mut hdr = Header {
548 + kind: protocol::KIND_REQUEST,
549 + code: 1,
550 + item_count: 1,
551 + message_id: 99,
552 + ..Header::default()
553 + };
554 + session.send(&mut hdr, &[0xFF]).expect("send");
555 + session.inflight_ids.insert(100);
556 +
557 + // Receive should fail because server disconnected
558 + let mut rbuf = [0u8; 4096];
559 + let result = session.receive(&mut rbuf);
560 + assert!(result.is_err());
561 + assert!(
562 + session.inflight_ids.is_empty(),
563 + "disconnect must fail every in-flight request on the session"
564 + );
565 +
566 + drop(session);
567 + server_thread.join().expect("server join");
568 + cleanup_socket(&svc);
569 +}
570 +
571 +// -----------------------------------------------------------------------
572 +// Test 9: Batch send/receive
573 +// -----------------------------------------------------------------------
574 +
575 +#[test]
576 +fn test_batch() {
577 + ensure_run_dir();
578 + let svc = unique_service("rs_batch");
579 + cleanup_socket(&svc);
580 +
581 + let svc_clone = svc.clone();
582 + let ready = Arc::new(AtomicBool::new(false));
583 + let ready_clone = ready.clone();
584 +
585 + let server_thread = thread::spawn(move || {
586 + let listener =
587 + UdsListener::bind(TEST_RUN_DIR, &svc_clone, default_server_config()).expect("listen");
588 + ready_clone.store(true, Ordering::Release);
589 +
590 + let mut session = listener.accept().expect("accept");
591 + let mut buf = [0u8; 8192];
592 + let (hdr, payload) = session.receive(&mut buf).expect("recv");
593 + let payload = payload.to_vec();
594 + let mut resp = hdr;
595 + resp.kind = protocol::KIND_RESPONSE;
596 + resp.transport_status = protocol::STATUS_OK;
597 + session.send(&mut resp, &payload).expect("send");
598 + });
599 +
600 + while !ready.load(Ordering::Acquire) {
601 + thread::sleep(Duration::from_millis(1));
602 + }
603 +
604 + let mut session =
605 + UdsSession::connect(TEST_RUN_DIR, &svc, &default_client_config()).expect("connect");
606 +
607 + // Build batch using protocol layer
608 + let mut batch_buf = [0u8; 2048];
609 + let mut builder = protocol::BatchBuilder::new(&mut batch_buf, 3);
610 +
611 + let item0 = [0x10u8, 0x20];
612 + let item1 = [0x30u8, 0x40, 0x50];
613 + let item2 = [0x60u8];
614 +
615 + builder.add(&item0).expect("add item0");
616 + builder.add(&item1).expect("add item1");
617 + builder.add(&item2).expect("add item2");
618 +
619 + let (batch_len, batch_count) = builder.finish();
620 + assert_eq!(batch_count, 3);
621 +
622 + let mut hdr = Header {
623 + kind: protocol::KIND_REQUEST,
624 + code: protocol::METHOD_INCREMENT,
625 + flags: protocol::FLAG_BATCH,
626 + item_count: batch_count,
627 + message_id: 55,
628 + ..Header::default()
629 + };
630 +
631 + session
632 + .send(&mut hdr, &batch_buf[..batch_len])
633 + .expect("send batch");
634 +
635 + let mut rbuf = [0u8; 4096];
636 + let (rhdr, rpayload) = session.receive(&mut rbuf).expect("recv batch");
637 + assert_eq!(rhdr.message_id, 55);
638 + assert!(rhdr.flags & protocol::FLAG_BATCH != 0);
639 + assert_eq!(rhdr.item_count, 3);
640 +
641 + // Verify items
642 + let (ip0, len0) = protocol::batch_item_get(&rpayload, 3, 0).expect("item0");
643 + assert_eq!(len0, 2);
644 + assert_eq!(ip0, &item0);
645 +
646 + let (ip1, len1) = protocol::batch_item_get(&rpayload, 3, 1).expect("item1");
647 + assert_eq!(len1, 3);
648 + assert_eq!(ip1, &item1);
649 +
650 + let (ip2, len2) = protocol::batch_item_get(&rpayload, 3, 2).expect("item2");
651 + assert_eq!(len2, 1);
652 + assert_eq!(ip2, &item2);
653 +
654 + drop(session);
655 + server_thread.join().expect("server join");
656 + cleanup_socket(&svc);
657 +}
658 +
659 +// -----------------------------------------------------------------------
660 +// Test 10: Large chunked pipelining
661 +// -----------------------------------------------------------------------
662 +
663 +#[test]
664 +fn test_chunked_pipelining() {
665 + ensure_run_dir();
666 + let svc = unique_service("rs_chkpipe");
667 + cleanup_socket(&svc);
668 +
669 + let svc_clone = svc.clone();
670 + let ready = Arc::new(AtomicBool::new(false));
671 + let ready_clone = ready.clone();
672 +
673 + let server_thread = thread::spawn(move || {
674 + let scfg = ServerConfig {
675 + packet_size: 128,
676 + max_request_payload_bytes: 65536,
677 + max_response_payload_bytes: 65536,
678 + ..default_server_config()
679 + };
680 + let listener = UdsListener::bind(TEST_RUN_DIR, &svc_clone, scfg).expect("listen");
681 + ready_clone.store(true, Ordering::Release);
682 +
683 + let mut session = listener.accept().expect("accept");
684 +
685 + // Echo 3 messages
686 + for _ in 0..3 {
687 + let mut buf = [0u8; 256];
688 + let (hdr, payload) = session.receive(&mut buf).expect("recv");
689 + let payload = payload.to_vec();
690 + let mut resp = hdr;
691 + resp.kind = protocol::KIND_RESPONSE;
692 + session.send(&mut resp, &payload).expect("send");
693 + }
694 + });
695 +
696 + while !ready.load(Ordering::Acquire) {
697 + thread::sleep(Duration::from_millis(1));
698 + }
699 +
700 + let ccfg = ClientConfig {
701 + packet_size: 128,
702 + max_request_payload_bytes: 65536,
703 + max_response_payload_bytes: 65536,
704 + ..default_client_config()
705 + };
706 + let mut session = UdsSession::connect(TEST_RUN_DIR, &svc, &ccfg).expect("connect");
707 +
708 + // Send 3 chunked messages in pipeline
709 + let sizes = [200usize, 500, 300];
710 + for (i, &sz) in sizes.iter().enumerate() {
711 + let payload: Vec<u8> = (0..sz).map(|j| ((i + j) & 0xFF) as u8).collect();
712 + let mut hdr = Header {
713 + kind: protocol::KIND_REQUEST,
714 + code: 1,
715 + item_count: 1,
716 + message_id: (i + 1) as u64,
717 + ..Header::default()
718 + };
719 + session.send(&mut hdr, &payload).expect("send");
720 + }
721 +
722 + // Receive 3 responses
723 + for (i, &sz) in sizes.iter().enumerate() {
724 + let mut rbuf = [0u8; 256];
725 + let (rhdr, rpayload) = session.receive(&mut rbuf).expect("recv");
726 + assert_eq!(rhdr.message_id, (i + 1) as u64);
727 + let expected: Vec<u8> = (0..sz).map(|j| ((i + j) & 0xFF) as u8).collect();
728 + assert_eq!(rpayload, expected);
729 + }
730 +
731 + drop(session);
732 + server_thread.join().expect("server join");
733 + cleanup_socket(&svc);
734 +}
735 +
736 +// -----------------------------------------------------------------------
737 +// Test: Pipeline 10 requests, verify all matched by message_id
738 +// -----------------------------------------------------------------------
739 +
740 +#[test]
741 +fn test_pipeline_10() {
742 + ensure_run_dir();
743 + let svc = unique_service("rs_pipe10");
744 + cleanup_socket(&svc);
745 +
746 + let svc_clone = svc.clone();
747 + let ready = Arc::new(AtomicBool::new(false));
748 + let ready_clone = ready.clone();
749 +
750 + let server_thread = thread::spawn(move || {
751 + let listener =
752 + UdsListener::bind(TEST_RUN_DIR, &svc_clone, default_server_config()).expect("listen");
753 + ready_clone.store(true, Ordering::Release);
754 +
755 + let mut session = listener.accept().expect("accept");
756 +
757 + for _ in 0..10 {
758 + let mut buf = [0u8; 8192];
759 + let (hdr, payload) = session.receive(&mut buf).expect("recv");
760 + let payload = payload.to_vec();
761 + let mut resp = hdr;
762 + resp.kind = protocol::KIND_RESPONSE;
763 + resp.transport_status = protocol::STATUS_OK;
764 + session.send(&mut resp, &payload).expect("send");
765 + }
766 + });
767 +
768 + while !ready.load(Ordering::Acquire) {
769 + thread::sleep(Duration::from_millis(1));
770 + }
771 +
772 + let mut session =
773 + UdsSession::connect(TEST_RUN_DIR, &svc, &default_client_config()).expect("connect");
774 +
775 + // Send 10 requests before reading any response
776 + for i in 1u64..=10 {
777 + let payload = i.to_ne_bytes();
778 + let mut hdr = Header {
779 + kind: protocol::KIND_REQUEST,
780 + code: 1,
781 + item_count: 1,
782 + message_id: i,
783 + ..Header::default()
784 + };
785 + session.send(&mut hdr, &payload).expect("send");
786 + }
787 +
788 + // Receive 10 responses, verify message_id and payload
789 + for i in 1u64..=10 {
790 + let mut rbuf = [0u8; 4096];
791 + let (rhdr, rpayload) = session.receive(&mut rbuf).expect("recv");
792 + assert_eq!(rhdr.message_id, i, "message_id mismatch at {i}");
793 + assert_eq!(rpayload.len(), 8, "payload len at {i}");
794 + let val = u64::from_ne_bytes(rpayload.try_into().unwrap());
795 + assert_eq!(val, i, "payload value at {i}");
796 + }
797 +
798 + drop(session);
799 + server_thread.join().expect("server join");
800 + cleanup_socket(&svc);
801 +}
802 +
803 +// -----------------------------------------------------------------------
804 +// Test: Pipeline 100 requests (stress pipelining)
805 +// -----------------------------------------------------------------------
806 +
807 +#[test]
808 +fn test_pipeline_100() {
809 + ensure_run_dir();
810 + let svc = unique_service("rs_pipe100");
811 + cleanup_socket(&svc);
812 +
813 + let svc_clone = svc.clone();
814 + let ready = Arc::new(AtomicBool::new(false));
815 + let ready_clone = ready.clone();
816 +
817 + let server_thread = thread::spawn(move || {
818 + let scfg = ServerConfig {
819 + max_request_payload_bytes: 65536,
820 + max_response_payload_bytes: 65536,
821 + ..default_server_config()
822 + };
823 + let listener = UdsListener::bind(TEST_RUN_DIR, &svc_clone, scfg).expect("listen");
824 + ready_clone.store(true, Ordering::Release);
825 +
826 + let mut session = listener.accept().expect("accept");
827 +
828 + for _ in 0..100 {
829 + let mut buf = [0u8; 8192];
830 + let (hdr, payload) = session.receive(&mut buf).expect("recv");
831 + let payload = payload.to_vec();
832 + let mut resp = hdr;
833 + resp.kind = protocol::KIND_RESPONSE;
834 + resp.transport_status = protocol::STATUS_OK;
835 + session.send(&mut resp, &payload).expect("send");
836 + }
837 + });
838 +
839 + while !ready.load(Ordering::Acquire) {
840 + thread::sleep(Duration::from_millis(1));
841 + }
842 +
843 + let ccfg = ClientConfig {
844 + max_request_payload_bytes: 65536,
845 + max_response_payload_bytes: 65536,
846 + ..default_client_config()
847 + };
848 + let mut session = UdsSession::connect(TEST_RUN_DIR, &svc, &ccfg).expect("connect");
849 +
850 + // Send 100 requests
851 + for i in 1u64..=100 {
852 + let payload = i.to_ne_bytes();
853 + let mut hdr = Header {
854 + kind: protocol::KIND_REQUEST,
855 + code: 1,
856 + item_count: 1,
857 + message_id: i,
858 + ..Header::default()
859 + };
860 + session.send(&mut hdr, &payload).expect("send");
861 + }
862 +
863 + // Receive 100 responses
864 + for i in 1u64..=100 {
865 + let mut rbuf = [0u8; 4096];
866 + let (rhdr, rpayload) = session.receive(&mut rbuf).expect("recv");
867 + assert_eq!(rhdr.message_id, i);
868 + let val = u64::from_ne_bytes(rpayload.try_into().unwrap());
869 + assert_eq!(val, i);
870 + }
871 +
872 + drop(session);
873 + server_thread.join().expect("server join");
874 + cleanup_socket(&svc);
875 +}
876 +
877 +// -----------------------------------------------------------------------
878 +// Test: Pipeline with mixed message sizes
879 +// -----------------------------------------------------------------------
880 +
881 +#[test]
882 +fn test_pipeline_mixed_sizes() {
883 + ensure_run_dir();
884 + let svc = unique_service("rs_pipemix");
885 + cleanup_socket(&svc);
886 +
887 + let svc_clone = svc.clone();
888 + let ready = Arc::new(AtomicBool::new(false));
889 + let ready_clone = ready.clone();
890 +
891 + let sizes = [8usize, 256, 1024, 8, 256, 1024, 8, 256, 1024];
892 + let count = sizes.len();
893 +
894 + let server_thread = thread::spawn(move || {
895 + let scfg = ServerConfig {
896 + max_request_payload_bytes: 65536,
897 + max_response_payload_bytes: 65536,
898 + ..default_server_config()
899 + };
900 + let listener = UdsListener::bind(TEST_RUN_DIR, &svc_clone, scfg).expect("listen");
901 + ready_clone.store(true, Ordering::Release);
902 +
903 + let mut session = listener.accept().expect("accept");
904 +
905 + for _ in 0..count {
906 + let mut buf = [0u8; 8192];
907 + let (hdr, payload) = session.receive(&mut buf).expect("recv");
908 + let payload = payload.to_vec();
909 + let mut resp = hdr;
910 + resp.kind = protocol::KIND_RESPONSE;
911 + resp.transport_status = protocol::STATUS_OK;
912 + session.send(&mut resp, &payload).expect("send");
913 + }
914 + });
915 +
916 + while !ready.load(Ordering::Acquire) {
917 + thread::sleep(Duration::from_millis(1));
918 + }
919 +
920 + let ccfg = ClientConfig {
921 + max_request_payload_bytes: 65536,
922 + max_response_payload_bytes: 65536,
923 + ..default_client_config()
924 + };
925 + let mut session = UdsSession::connect(TEST_RUN_DIR, &svc, &ccfg).expect("connect");
926 +
927 + // Send all messages with varying sizes
928 + for (i, &sz) in sizes.iter().enumerate() {
929 + let payload: Vec<u8> = (0..sz).map(|j| ((i * 37 + j) & 0xFF) as u8).collect();
930 + let mut hdr = Header {
931 + kind: protocol::KIND_REQUEST,
932 + code: 1,
933 + item_count: 1,
934 + message_id: (i + 1) as u64,
935 + ..Header::default()
936 + };
937 + session.send(&mut hdr, &payload).expect("send");
938 + }
939 +
940 + // Receive all responses
941 + for (i, &sz) in sizes.iter().enumerate() {
942 + let mut rbuf = [0u8; 4096];
943 + let (rhdr, rpayload) = session.receive(&mut rbuf).expect("recv");
944 + assert_eq!(rhdr.message_id, (i + 1) as u64, "message_id at {i}");
945 + assert_eq!(rpayload.len(), sz, "payload len at {i}");
946 + let expected: Vec<u8> = (0..sz).map(|j| ((i * 37 + j) & 0xFF) as u8).collect();
947 + assert_eq!(rpayload, expected, "payload data at {i}");
948 + }
949 +
950 + drop(session);
951 + server_thread.join().expect("server join");
952 + cleanup_socket(&svc);
953 +}
954 +
955 +// -----------------------------------------------------------------------
956 +// Test: Pipeline with chunked messages (> packet_size)
957 +// -----------------------------------------------------------------------
958 +
959 +#[test]
960 +fn test_pipeline_chunked_multi() {
961 + ensure_run_dir();
962 + let svc = unique_service("rs_pipechk2");
963 + cleanup_socket(&svc);
964 +
965 + let svc_clone = svc.clone();
966 + let ready = Arc::new(AtomicBool::new(false));
967 + let ready_clone = ready.clone();
968 +
969 + let sizes = [200usize, 500, 300, 800, 150];
970 + let count = sizes.len();
971 +
972 + let server_thread = thread::spawn(move || {
973 + let scfg = ServerConfig {
974 + packet_size: 128,
975 + max_request_payload_bytes: 65536,
976 + max_response_payload_bytes: 65536,
977 + ..default_server_config()
978 + };
979 + let listener = UdsListener::bind(TEST_RUN_DIR, &svc_clone, scfg).expect("listen");
980 + ready_clone.store(true, Ordering::Release);
981 +
982 + let mut session = listener.accept().expect("accept");
983 +
984 + for _ in 0..count {
985 + let mut buf = [0u8; 256];
986 + let (hdr, payload) = session.receive(&mut buf).expect("recv");
987 + let payload = payload.to_vec();
988 + let mut resp = hdr;
989 + resp.kind = protocol::KIND_RESPONSE;
990 + session.send(&mut resp, &payload).expect("send");
991 + }
992 + });
993 +
994 + while !ready.load(Ordering::Acquire) {
995 + thread::sleep(Duration::from_millis(1));
996 + }
997 +
998 + let ccfg = ClientConfig {
999 + packet_size: 128,
1000 + max_request_payload_bytes: 65536,
1001 + max_response_payload_bytes: 65536,
1002 + ..default_client_config()
1003 + };
1004 + let mut session = UdsSession::connect(TEST_RUN_DIR, &svc, &ccfg).expect("connect");
1005 +
1006 + // Send all chunked messages
1007 + for (i, &sz) in sizes.iter().enumerate() {
1008 + let payload: Vec<u8> = (0..sz).map(|j| ((i + j) & 0xFF) as u8).collect();
1009 + let mut hdr = Header {
1010 + kind: protocol::KIND_REQUEST,
1011 + code: 1,
1012 + item_count: 1,
1013 + message_id: (i + 1) as u64,
1014 + ..Header::default()
1015 + };
1016 + session.send(&mut hdr, &payload).expect("send");
1017 + }
1018 +
1019 + // Receive all responses
1020 + for (i, &sz) in sizes.iter().enumerate() {
1021 + let mut rbuf = [0u8; 256];
1022 + let (rhdr, rpayload) = session.receive(&mut rbuf).expect("recv");
1023 + assert_eq!(rhdr.message_id, (i + 1) as u64);
1024 + let expected: Vec<u8> = (0..sz).map(|j| ((i + j) & 0xFF) as u8).collect();
1025 + assert_eq!(rpayload, expected);
1026 + }
1027 +
1028 + drop(session);
1029 + server_thread.join().expect("server join");
1030 + cleanup_socket(&svc);
1031 +}
1032 +
1033 +#[test]
1034 +fn test_invalid_service_name() {
1035 + let bad_names = &["", ".", "..", "foo/bar", "../etc", "name space", "a@b"];
1036 + for name in bad_names {
1037 + let result = validate_service_name(name);
1038 + assert!(result.is_err(), "should reject {:?}", name);
1039 + }
1040 +
1041 + let good_names = &[
1042 + "valid-name",
1043 + "valid_name",
1044 + "valid.name",
1045 + "ValidName123",
1046 + "a",
1047 + ];
1048 + for name in good_names {
1049 + validate_service_name(name).unwrap_or_else(|e| panic!("{:?} should be valid: {e}", name));
1050 + }
1051 +}
1052 +
1053 +#[test]
1054 +fn test_hello_decode_nonzero_padding() {
1055 + let h = Hello {
1056 + layout_version: 1,
1057 + supported_profiles: PROFILE_BASELINE,
1058 + max_request_payload_bytes: 1024,
1059 + max_request_batch_items: 1,
1060 + max_response_payload_bytes: 1024,
1061 + max_response_batch_items: 1,
1062 + packet_size: 65536,
1063 + ..Default::default()
1064 + };
1065 +
1066 + let mut buf = [0u8; 44];
1067 + h.encode(&mut buf);
1068 + Hello::decode(&buf).expect("valid hello should decode");
1069 +
1070 + // Corrupt padding bytes 28..32
1071 + buf[28] = 0xFF;
1072 + assert_eq!(Hello::decode(&buf), Err(protocol::NipcError::BadLayout));
1073 +}
1074 +
1075 +// -----------------------------------------------------------------------
1076 +// UdsError Display coverage (lines 73-91)
1077 +// -----------------------------------------------------------------------
1078 +
1079 +#[test]
1080 +fn uds_error_display_all_variants() {
1081 + let cases: Vec<(UdsError, &str)> = vec![
1082 + (UdsError::PathTooLong, "socket path exceeds sun_path limit"),
1083 + (UdsError::Socket(22), "socket syscall failed: errno 22"),
1084 + (UdsError::Connect(111), "connect failed: errno 111"),
1085 + (UdsError::Accept(24), "accept failed: errno 24"),
1086 + (UdsError::Send(32), "send failed: errno 32"),
1087 + (UdsError::Recv(0), "recv failed: errno 0"),
1088 + (UdsError::Handshake("test".into()), "handshake error: test"),
1089 + (UdsError::AuthFailed, "authentication token rejected"),
1090 + (UdsError::NoProfile, "no common transport profile"),
1091 + (
1092 + UdsError::Incompatible("version mismatch".into()),
1093 + "incompatible protocol: version mismatch",
1094 + ),
1095 + (UdsError::Protocol("bad".into()), "protocol violation: bad"),
1096 + (UdsError::AddrInUse, "address already in use by live server"),
1097 + (UdsError::Chunk("mismatch".into()), "chunk error: mismatch"),
1098 + (UdsError::Alloc, "memory allocation failed"),
1099 + (UdsError::LimitExceeded, "negotiated limit exceeded"),
1100 + (UdsError::BadParam("foo".into()), "bad parameter: foo"),
1101 + (UdsError::DuplicateMsgId(42), "duplicate message_id: 42"),
1102 + (
1103 + UdsError::UnknownMsgId(99),
1104 + "unknown response message_id: 99",
1105 + ),
1106 + ];
1107 + for (err, expected) in cases {
1108 + assert_eq!(format!("{}", err), expected);
1109 + }
1110 + // Verify Error trait
1111 + let e: &dyn std::error::Error = &UdsError::PathTooLong;
1112 + let _ = format!("{e}");
1113 +}
1114 +
1115 +// -----------------------------------------------------------------------
1116 +// UdsSession::role() coverage (lines 205-206)
1117 +// -----------------------------------------------------------------------
1118 +
1119 +#[test]
1120 +fn test_session_role() {
1121 + ensure_run_dir();
1122 + let svc = unique_service("rs_role");
1123 + cleanup_socket(&svc);
1124 +
1125 + let svc_clone = svc.clone();
1126 + let ready = Arc::new(AtomicBool::new(false));
1127 + let ready_clone = ready.clone();
1128 +
1129 + let server_thread = thread::spawn(move || {
1130 + let listener =
1131 + UdsListener::bind(TEST_RUN_DIR, &svc_clone, default_server_config()).expect("listen");
1132 + ready_clone.store(true, Ordering::Release);
1133 + let session = listener.accept().expect("accept");
1134 + assert_eq!(session.role(), Role::Server);
1135 + });
1136 +
1137 + while !ready.load(Ordering::Acquire) {
1138 + thread::sleep(Duration::from_millis(1));
1139 + }
1140 +
1141 + let session =
1142 + UdsSession::connect(TEST_RUN_DIR, &svc, &default_client_config()).expect("connect");
1143 + assert_eq!(session.role(), Role::Client);
1144 +
1145 + drop(session);
1146 + server_thread.join().expect("server join");
1147 + cleanup_socket(&svc);
1148 +}
1149 +
1150 +// -----------------------------------------------------------------------
1151 +// Send on closed session (line 238)
1152 +// -----------------------------------------------------------------------
1153 +
1154 +#[test]
1155 +fn test_send_on_closed_session() {
1156 + ensure_run_dir();
1157 + let svc = unique_service("rs_sendclosed");
1158 + cleanup_socket(&svc);
1159 +
1160 + let svc_clone = svc.clone();
1161 + let ready = Arc::new(AtomicBool::new(false));
1162 + let ready_clone = ready.clone();
1163 +
1164 + let server_thread = thread::spawn(move || {
1165 + let listener =
1166 + UdsListener::bind(TEST_RUN_DIR, &svc_clone, default_server_config()).expect("listen");
1167 + ready_clone.store(true, Ordering::Release);
1168 + let _session = listener.accept().expect("accept");
1169 + });
1170 +
1171 + while !ready.load(Ordering::Acquire) {
1172 + thread::sleep(Duration::from_millis(1));
1173 + }
1174 +
1175 + let mut session =
1176 + UdsSession::connect(TEST_RUN_DIR, &svc, &default_client_config()).expect("connect");
1177 +
1178 + // Simulate closed session by closing the fd via Drop-like behavior
1179 + // We can't call close() (no such method), so we close fd directly
1180 + unsafe {
1181 + libc::close(session.fd);
1182 + }
1183 + session.fd = -1;
1184 +
1185 + let mut hdr = Header {
1186 + kind: protocol::KIND_REQUEST,
1187 + code: 1,
1188 + item_count: 1,
1189 + message_id: 1,
1190 + ..Header::default()
1191 + };
1192 + let result = session.send(&mut hdr, &[1, 2, 3]);
1193 + assert!(matches!(result, Err(UdsError::BadParam(_))));
1194 +
1195 + // Receive on closed session (line 333)
1196 + let mut rbuf = [0u8; 4096];
1197 + let result = session.receive(&mut rbuf);
1198 + assert!(matches!(result, Err(UdsError::BadParam(_))));
1199 +
1200 + server_thread.join().expect("server join");
1201 + cleanup_socket(&svc);
1202 +}
1203 +
1204 +#[test]
1205 +fn test_accept_on_closed_listener_returns_accept_error() {
1206 + ensure_run_dir();
1207 + let svc = unique_service("rs_accept_closed");
1208 + cleanup_socket(&svc);
1209 +
1210 + let mut listener =
1211 + UdsListener::bind(TEST_RUN_DIR, &svc, default_server_config()).expect("listen");
1212 + let fd = listener.fd;
1213 + assert!(fd >= 0, "listener fd should be valid");
1214 + assert_eq!(unsafe { libc::close(fd) }, 0, "close listener fd");
1215 +
1216 + let err = match listener.accept() {
1217 + Ok(_) => panic!("accept on closed listener should fail"),
1218 + Err(err) => err,
1219 + };
1220 + assert!(matches!(err, UdsError::Accept(_)));
1221 +
1222 + listener.fd = -1;
1223 + drop(listener);
1224 + cleanup_socket(&svc);
1225 +}
1226 +
1227 +// -----------------------------------------------------------------------
1228 +// Duplicate message_id (line 244)
1229 +// -----------------------------------------------------------------------
1230 +
1231 +#[test]
1232 +fn test_duplicate_message_id() {
1233 + ensure_run_dir();
1234 + let svc = unique_service("rs_dupmsg");
1235 + cleanup_socket(&svc);
1236 +
1237 + let svc_clone = svc.clone();
1238 + let ready = Arc::new(AtomicBool::new(false));
1239 + let ready_clone = ready.clone();
1240 +
1241 + let server_thread = thread::spawn(move || {
1242 + let listener =
1243 + UdsListener::bind(TEST_RUN_DIR, &svc_clone, default_server_config()).expect("listen");
1244 + ready_clone.store(true, Ordering::Release);
1245 + let _session = listener.accept().expect("accept");
1246 + // Hold connection open while client tests
1247 + thread::sleep(Duration::from_millis(500));
1248 + });
1249 +
1250 + while !ready.load(Ordering::Acquire) {
1251 + thread::sleep(Duration::from_millis(1));
1252 + }
1253 +
1254 + let mut session =
1255 + UdsSession::connect(TEST_RUN_DIR, &svc, &default_client_config()).expect("connect");
1256 +
1257 + // First send with message_id=42 should succeed
1258 + let mut hdr1 = Header {
1259 + kind: protocol::KIND_REQUEST,
1260 + code: 1,
1261 + item_count: 1,
1262 + message_id: 42,
1263 + ..Header::default()
1264 + };
1265 + session.send(&mut hdr1, &[1]).expect("first send ok");
1266 +
1267 + // Second send with same message_id=42 should fail
1268 + let mut hdr2 = Header {
1269 + kind: protocol::KIND_REQUEST,
1270 + code: 1,
1271 + item_count: 1,
1272 + message_id: 42,
1273 + ..Header::default()
1274 + };
1275 + let result = session.send(&mut hdr2, &[2]);
1276 + assert!(matches!(result, Err(UdsError::DuplicateMsgId(42))));
1277 +
1278 + drop(session);
1279 + server_thread.join().expect("server join");
1280 + cleanup_socket(&svc);
1281 +}
1282 +
1283 +// -----------------------------------------------------------------------
1284 +// Receive: payload exceeds limit (line 354)
1285 +// -----------------------------------------------------------------------
1286 +
1287 +#[test]
1288 +fn test_receive_payload_exceeds_limit() {
1289 + let (fd0, fd1) = socketpair_seqpacket();
1290 + let mut session = test_session(fd0, Role::Client, 4096);
1291 + session.max_response_payload_bytes = 16;
1292 + session.inflight_ids.insert(99);
1293 +
1294 + let payload = [0xAB; 32];
1295 + let hdr = Header {
1296 + magic: MAGIC_MSG,
1297 + version: VERSION,
1298 + header_len: protocol::HEADER_LEN,
1299 + kind: KIND_RESPONSE,
1300 + code: 1,
1301 + flags: 0,
1302 + transport_status: protocol::STATUS_OK,
1303 + payload_len: payload.len() as u32,
1304 + item_count: 1,
1305 + message_id: 99,
1306 + };
1307 + let mut pkt = [0u8; HEADER_SIZE + 32];
1308 + hdr.encode(&mut pkt[..HEADER_SIZE]);
1309 + pkt[HEADER_SIZE..].copy_from_slice(&payload);
1310 + raw_send(fd1, &pkt).expect("send oversized payload");
1311 +
1312 + let mut buf = [0u8; 64];
1313 + let err = session
1314 + .receive(&mut buf)
1315 + .expect_err("payload exceeds limit");
1316 + assert!(matches!(err, UdsError::LimitExceeded));
1317 +
1318 + unsafe { libc::close(fd1) };
1319 +}
1320 +
1321 +// -----------------------------------------------------------------------
1322 +// Receive: batch item_count exceeds limit (line 364)
1323 +// -----------------------------------------------------------------------
1324 +
1325 +#[test]
1326 +fn test_receive_batch_count_exceeds_limit() {
1327 + let (fd0, fd1) = socketpair_seqpacket();
1328 + let mut session = test_session(fd0, Role::Client, 4096);
1329 + session.max_response_batch_items = 16;
1330 + session.inflight_ids.insert(1);
1331 +
1332 + let hdr = Header {
1333 + magic: MAGIC_MSG,
1334 + version: VERSION,
1335 + header_len: protocol::HEADER_LEN,
1336 + kind: KIND_RESPONSE,
1337 + code: 1,
1338 + flags: 0,
1339 + transport_status: protocol::STATUS_OK,
1340 + payload_len: 1,
1341 + item_count: 17,
1342 + message_id: 1,
1343 + };
1344 + let mut pkt = [0u8; HEADER_SIZE + 1];
1345 + hdr.encode(&mut pkt[..HEADER_SIZE]);
1346 + pkt[HEADER_SIZE] = 0xAD;
1347 + raw_send(fd1, &pkt).expect("send oversized batch-count response");
1348 +
1349 + let mut buf = [0u8; 64];
1350 + let err = session
1351 + .receive(&mut buf)
1352 + .expect_err("batch count exceeds limit");
1353 + assert!(matches!(err, UdsError::LimitExceeded));
1354 +
1355 + unsafe { libc::close(fd1) };
1356 +}
1357 +
1358 +#[test]
1359 +fn test_directional_limit_negotiation() {
1360 + ensure_run_dir();
1361 + let svc = unique_service("rs_dir_limits");
1362 + cleanup_socket(&svc);
1363 +
1364 + let svc_clone = svc.clone();
1365 + let ready = Arc::new(AtomicBool::new(false));
1366 + let ready_clone = ready.clone();
1367 +
1368 + let server_thread = thread::spawn(move || {
1369 + let scfg = ServerConfig {
1370 + max_request_payload_bytes: 2048,
1371 + max_request_batch_items: 8,
1372 + max_response_payload_bytes: 8192,
1373 + max_response_batch_items: 32,
1374 + ..default_server_config()
1375 + };
1376 + let listener = UdsListener::bind(TEST_RUN_DIR, &svc_clone, scfg).expect("listen");
1377 + ready_clone.store(true, Ordering::Release);
1378 +
1379 + let session = listener.accept().expect("accept");
1380 + assert_eq!(session.max_request_payload_bytes, 4096);
1381 + assert_eq!(session.max_request_batch_items, 16);
1382 + assert_eq!(session.max_response_payload_bytes, 8192);
1383 + assert_eq!(session.max_response_batch_items, 16);
1384 + assert_ne!(session.session_id, 0);
1385 + });
1386 +
1387 + while !ready.load(Ordering::Acquire) {
1388 + thread::sleep(Duration::from_millis(1));
1389 + }
1390 +
1391 + let ccfg = ClientConfig {
1392 + max_request_payload_bytes: 4096,
1393 + max_request_batch_items: 16,
1394 + max_response_payload_bytes: 4096,
1395 + max_response_batch_items: 16,
1396 + ..default_client_config()
1397 + };
1398 + let session = UdsSession::connect(TEST_RUN_DIR, &svc, &ccfg).expect("connect");
1399 +
1400 + assert_eq!(session.max_request_payload_bytes, 4096);
1401 + assert_eq!(session.max_request_batch_items, 16);
1402 + assert_eq!(session.max_response_payload_bytes, 8192);
1403 + assert_eq!(session.max_response_batch_items, 16);
1404 + assert_ne!(session.session_id, 0);
1405 +
1406 + drop(session);
1407 + server_thread.join().expect("server join");
1408 + cleanup_socket(&svc);
1409 +}
1410 +
1411 +// -----------------------------------------------------------------------
1412 +// Connect to nonexistent socket (line 220)
1413 +// -----------------------------------------------------------------------
1414 +
1415 +#[test]
1416 +fn test_connect_nonexistent() {
1417 + ensure_run_dir();
1418 + let svc = unique_service("rs_noexist");
1419 + cleanup_socket(&svc);
1420 +
1421 + let result = UdsSession::connect(TEST_RUN_DIR, &svc, &default_client_config());
1422 + assert!(matches!(result, Err(UdsError::Connect(_))));
1423 +}
1424 +
1425 +#[test]
1426 +fn test_send_packet_size_too_small_rolls_back_inflight() {
1427 + let (fd0, fd1) = socketpair_seqpacket();
1428 + let mut session = test_session(fd0, Role::Client, HEADER_SIZE as u32);
1429 +
1430 + let mut hdr = Header {
1431 + kind: KIND_REQUEST,
1432 + code: 1,
1433 + item_count: 1,
1434 + message_id: 77,
1435 + ..Header::default()
1436 + };
1437 +
1438 + let err = session
1439 + .send(&mut hdr, &[0xAA])
1440 + .expect_err("packet_size too small");
1441 + assert!(matches!(err, UdsError::BadParam(_)));
1442 + assert!(
1443 + !session.inflight_ids.contains(&77),
1444 + "failed send should roll back the tracked message_id"
1445 + );
1446 +
1447 + unsafe { libc::close(fd1) };
1448 +}
1449 +
1450 +#[test]
1451 +fn test_receive_packet_too_short_for_header() {
1452 + let (fd0, fd1) = socketpair_seqpacket();
1453 + let mut session = test_session(fd0, Role::Server, 4096);
1454 +
1455 + raw_send(fd1, &[0x01]).expect("send short packet");
1456 +
1457 + let mut buf = [0u8; 64];
1458 + let err = session.receive(&mut buf).expect_err("packet too short");
1459 + assert!(matches!(err, UdsError::Protocol(_)));
1460 +
1461 + unsafe { libc::close(fd1) };
1462 +}
1463 +
1464 +#[test]
1465 +fn test_receive_batch_directory_too_short_nonchunked() {
1466 + let (fd0, fd1) = socketpair_seqpacket();
1467 + let mut session = test_session(fd0, Role::Server, 4096);
1468 +
1469 + let payload = [0u8; 8];
1470 + let mut pkt = [0u8; HEADER_SIZE + 8];
1471 + let hdr = Header {
1472 + magic: MAGIC_MSG,
1473 + version: VERSION,
1474 + header_len: protocol::HEADER_LEN,
1475 + kind: KIND_REQUEST,
1476 + code: 1,
1477 + flags: FLAG_BATCH,
1478 + transport_status: protocol::STATUS_OK,
1479 + payload_len: payload.len() as u32,
1480 + item_count: 2,
1481 + message_id: 1,
1482 + };
1483 + hdr.encode(&mut pkt[..HEADER_SIZE]);
1484 + pkt[HEADER_SIZE..].copy_from_slice(&payload);
1485 +
1486 + raw_send(fd1, &pkt).expect("send malformed batch packet");
1487 +
1488 + let mut buf = [0u8; 128];
1489 + let err = session
1490 + .receive(&mut buf)
1491 + .expect_err("batch directory exceeds payload");
1492 + assert!(matches!(err, UdsError::Protocol(_)));
1493 +
1494 + unsafe { libc::close(fd1) };
1495 +}
1496 +
1497 +#[test]
1498 +fn test_receive_batch_directory_invalid_nonchunked() {
1499 + let (fd0, fd1) = socketpair_seqpacket();
1500 + let mut session = test_session(fd0, Role::Server, 4096);
1501 +
1502 + let mut payload = [0u8; 16];
1503 + payload[0..4].copy_from_slice(&0u32.to_ne_bytes());
1504 + payload[4..8].copy_from_slice(&32u32.to_ne_bytes());
1505 + payload[8..12].copy_from_slice(&0u32.to_ne_bytes());
1506 + payload[12..16].copy_from_slice(&32u32.to_ne_bytes());
1507 +
1508 + let mut pkt = [0u8; HEADER_SIZE + 16];
1509 + let hdr = Header {
1510 + magic: MAGIC_MSG,
1511 + version: VERSION,
1512 + header_len: protocol::HEADER_LEN,
1513 + kind: KIND_REQUEST,
1514 + code: 1,
1515 + flags: FLAG_BATCH,
1516 + transport_status: protocol::STATUS_OK,
1517 + payload_len: payload.len() as u32,
1518 + item_count: 2,
1519 + message_id: 2,
1520 + };
1521 + hdr.encode(&mut pkt[..HEADER_SIZE]);
1522 + pkt[HEADER_SIZE..].copy_from_slice(&payload);
1523 +
1524 + raw_send(fd1, &pkt).expect("send invalid batch packet");
1525 +
1526 + let mut buf = [0u8; 128];
1527 + let err = session
1528 + .receive(&mut buf)
1529 + .expect_err("batch directory validation should fail");
1530 + assert!(matches!(err, UdsError::Protocol(_)));
1531 +
1532 + unsafe { libc::close(fd1) };
1533 +}
1534 +
1535 +#[test]
1536 +fn test_receive_chunk_message_id_mismatch() {
1537 + let (fd0, fd1) = socketpair_seqpacket();
1538 + let mut session = test_session(fd0, Role::Server, (HEADER_SIZE + 10) as u32);
1539 +
1540 + let first_payload = [1u8; 10];
1541 + let hdr = Header {
1542 + magic: MAGIC_MSG,
1543 + version: VERSION,
1544 + header_len: protocol::HEADER_LEN,
1545 + kind: KIND_REQUEST,
1546 + code: 1,
1547 + flags: 0,
1548 + transport_status: protocol::STATUS_OK,
1549 + payload_len: 20,
1550 + item_count: 1,
1551 + message_id: 11,
1552 + };
1553 + let mut first_pkt = [0u8; HEADER_SIZE + 10];
1554 + hdr.encode(&mut first_pkt[..HEADER_SIZE]);
1555 + first_pkt[HEADER_SIZE..].copy_from_slice(&first_payload);
1556 + raw_send(fd1, &first_pkt).expect("send first chunk");
1557 +
1558 + let second_payload = [2u8; 10];
1559 + let chk = ChunkHeader {
1560 + magic: MAGIC_CHUNK,
1561 + version: VERSION,
1562 + flags: 0,
1563 + message_id: 12,
1564 + total_message_len: (HEADER_SIZE + 20) as u32,
1565 + chunk_index: 1,
1566 + chunk_count: 2,
1567 + chunk_payload_len: second_payload.len() as u32,
1568 + };
1569 + let mut second_pkt = [0u8; HEADER_SIZE + 10];
1570 + chk.encode(&mut second_pkt[..HEADER_SIZE]);
1571 + second_pkt[HEADER_SIZE..].copy_from_slice(&second_payload);
1572 + raw_send(fd1, &second_pkt).expect("send mismatched continuation");
1573 +
1574 + let mut buf = [0u8; 64];
1575 + let err = session.receive(&mut buf).expect_err("message_id mismatch");
1576 + assert!(matches!(err, UdsError::Chunk(_)));
1577 +
1578 + unsafe { libc::close(fd1) };
1579 +}
1580 +
1581 +#[test]
1582 +fn test_receive_chunk_index_mismatch() {
1583 + let (fd0, fd1) = socketpair_seqpacket();
1584 + let mut session = test_session(fd0, Role::Server, (HEADER_SIZE + 10) as u32);
1585 +
1586 + let first_payload = [1u8; 10];
1587 + let hdr = Header {
1588 + magic: MAGIC_MSG,
1589 + version: VERSION,
1590 + header_len: protocol::HEADER_LEN,
1591 + kind: KIND_REQUEST,
1592 + code: 1,
1593 + flags: 0,
1594 + transport_status: protocol::STATUS_OK,
1595 + payload_len: 20,
1596 + item_count: 1,
1597 + message_id: 11,
1598 + };
1599 + let mut first_pkt = [0u8; HEADER_SIZE + 10];
1600 + hdr.encode(&mut first_pkt[..HEADER_SIZE]);
1601 + first_pkt[HEADER_SIZE..].copy_from_slice(&first_payload);
1602 + raw_send(fd1, &first_pkt).expect("send first chunk");
1603 +
1604 + let second_payload = [2u8; 10];
1605 + let chk = ChunkHeader {
1606 + magic: MAGIC_CHUNK,
1607 + version: VERSION,
1608 + flags: 0,
1609 + message_id: 11,
1610 + total_message_len: (HEADER_SIZE + 20) as u32,
1611 + chunk_index: 2,
1612 + chunk_count: 2,
1613 + chunk_payload_len: second_payload.len() as u32,
1614 + };
1615 + let mut second_pkt = [0u8; HEADER_SIZE + 10];
1616 + chk.encode(&mut second_pkt[..HEADER_SIZE]);
1617 + second_pkt[HEADER_SIZE..].copy_from_slice(&second_payload);
1618 + raw_send(fd1, &second_pkt).expect("send mismatched continuation");
1619 +
1620 + let mut buf = [0u8; 64];
1621 + let err = session.receive(&mut buf).expect_err("chunk index mismatch");
1622 + assert!(matches!(
1623 + err,
1624 + UdsError::Chunk(ref msg) if msg == "chunk_index mismatch: expected 1, got 2"
1625 + ));
1626 +
1627 + unsafe { libc::close(fd1) };
1628 +}
1629 +
1630 +#[test]
1631 +fn test_receive_chunked_batch_directory_too_short() {
1632 + let (fd0, fd1) = socketpair_seqpacket();
1633 + let mut session = test_session(fd0, Role::Server, (HEADER_SIZE + 16) as u32);
1634 +
1635 + let first_payload = [0u8; 8];
1636 + let hdr = Header {
1637 + magic: MAGIC_MSG,
1638 + version: VERSION,
1639 + header_len: protocol::HEADER_LEN,
1640 + kind: KIND_REQUEST,
1641 + code: 1,
1642 + flags: FLAG_BATCH,
1643 + transport_status: protocol::STATUS_OK,
1644 + payload_len: 12,
1645 + item_count: 2,
1646 + message_id: 21,
1647 + };
1648 + let mut first_pkt = [0u8; HEADER_SIZE + 8];
1649 + hdr.encode(&mut first_pkt[..HEADER_SIZE]);
1650 + first_pkt[HEADER_SIZE..].copy_from_slice(&first_payload);
1651 + raw_send(fd1, &first_pkt).expect("send first chunk");
1652 +
1653 + let second_payload = [0u8; 4];
1654 + let chk = ChunkHeader {
1655 + magic: MAGIC_CHUNK,
1656 + version: VERSION,
1657 + flags: 0,
1658 + message_id: 21,
1659 + total_message_len: (HEADER_SIZE + 12) as u32,
1660 + chunk_index: 1,
1661 + chunk_count: 2,
1662 + chunk_payload_len: second_payload.len() as u32,
1663 + };
1664 + let mut second_pkt = [0u8; HEADER_SIZE + 4];
1665 + chk.encode(&mut second_pkt[..HEADER_SIZE]);
1666 + second_pkt[HEADER_SIZE..].copy_from_slice(&second_payload);
1667 + raw_send(fd1, &second_pkt).expect("send continuation");
1668 +
1669 + let mut buf = [0u8; 64];
1670 + let err = session
1671 + .receive(&mut buf)
1672 + .expect_err("chunked batch directory too short");
1673 + assert!(matches!(err, UdsError::Protocol(_)));
1674 +
1675 + unsafe { libc::close(fd1) };
1676 +}
1677 +
1678 +#[test]
1679 +fn test_receive_chunked_batch_directory_invalid() {
1680 + let (fd0, fd1) = socketpair_seqpacket();
1681 + let mut session = test_session(fd0, Role::Server, (HEADER_SIZE + 8) as u32);
1682 +
1683 + let mut full_payload = [0u8; 24];
1684 + full_payload[0..4].copy_from_slice(&0u32.to_ne_bytes());
1685 + full_payload[4..8].copy_from_slice(&16u32.to_ne_bytes());
1686 + full_payload[8..12].copy_from_slice(&0u32.to_ne_bytes());
1687 + full_payload[12..16].copy_from_slice(&4u32.to_ne_bytes());
1688 + full_payload[16..24].copy_from_slice(b"payload!");
1689 +
1690 + let hdr = Header {
1691 + magic: MAGIC_MSG,
1692 + version: VERSION,
1693 + header_len: protocol::HEADER_LEN,
1694 + kind: KIND_REQUEST,
1695 + code: 1,
1696 + flags: FLAG_BATCH,
1697 + transport_status: protocol::STATUS_OK,
1698 + payload_len: full_payload.len() as u32,
1699 + item_count: 2,
1700 + message_id: 46,
1701 + };
1702 + let mut first_pkt = [0u8; HEADER_SIZE + 16];
1703 + hdr.encode(&mut first_pkt[..HEADER_SIZE]);
1704 + first_pkt[HEADER_SIZE..].copy_from_slice(&full_payload[..16]);
1705 + raw_send(fd1, &first_pkt).expect("send first chunk");
1706 +
1707 + let second_payload = &full_payload[16..];
1708 + let chk = ChunkHeader {
1709 + magic: MAGIC_CHUNK,
1710 + version: VERSION,
1711 + flags: 0,
1712 + message_id: 46,
1713 + total_message_len: (HEADER_SIZE + full_payload.len()) as u32,
1714 + chunk_index: 1,
1715 + chunk_count: 2,
1716 + chunk_payload_len: second_payload.len() as u32,
1717 + };
1718 + let mut second_pkt = [0u8; HEADER_SIZE + 8];
1719 + chk.encode(&mut second_pkt[..HEADER_SIZE]);
1720 + second_pkt[HEADER_SIZE..HEADER_SIZE + second_payload.len()].copy_from_slice(second_payload);
1721 + raw_send(fd1, &second_pkt).expect("send invalid continuation");
1722 +
1723 + let mut buf = [0u8; 64];
1724 + let err = session
1725 + .receive(&mut buf)
1726 + .expect_err("chunked batch directory validation should fail");
1727 + assert!(matches!(err, UdsError::Protocol(_)));
1728 +
1729 + unsafe { libc::close(fd1) };
1730 +}
1731 +
1732 +#[test]
1733 +fn test_receive_chunk_total_message_len_mismatch() {
1734 + let (fd0, fd1) = socketpair_seqpacket();
1735 + let mut session = test_session(fd0, Role::Server, (HEADER_SIZE + 10) as u32);
1736 +
1737 + let first_payload = [1u8; 10];
1738 + let hdr = Header {
1739 + magic: MAGIC_MSG,
1740 + version: VERSION,
1741 + header_len: protocol::HEADER_LEN,
1742 + kind: KIND_REQUEST,
1743 + code: 1,
1744 + flags: 0,
1745 + transport_status: protocol::STATUS_OK,
1746 + payload_len: 20,
1747 + item_count: 1,
1748 + message_id: 41,
1749 + };
1750 + let mut first_pkt = [0u8; HEADER_SIZE + 10];
1751 + hdr.encode(&mut first_pkt[..HEADER_SIZE]);
1752 + first_pkt[HEADER_SIZE..].copy_from_slice(&first_payload);
1753 + raw_send(fd1, &first_pkt).expect("send first chunk");
1754 +
1755 + let second_payload = [2u8; 10];
1756 + let chk = ChunkHeader {
1757 + magic: MAGIC_CHUNK,
1758 + version: VERSION,
1759 + flags: 0,
1760 + message_id: 41,
1761 + total_message_len: (HEADER_SIZE + 21) as u32,
1762 + chunk_index: 1,
1763 + chunk_count: 2,
1764 + chunk_payload_len: second_payload.len() as u32,
1765 + };
1766 + let mut second_pkt = [0u8; HEADER_SIZE + 10];
1767 + chk.encode(&mut second_pkt[..HEADER_SIZE]);
1768 + second_pkt[HEADER_SIZE..].copy_from_slice(&second_payload);
1769 + raw_send(fd1, &second_pkt).expect("send bad total_message_len");
1770 +
1771 + let mut buf = [0u8; 64];
1772 + let err = session
1773 + .receive(&mut buf)
1774 + .expect_err("total_message_len mismatch");
1775 + assert!(matches!(err, UdsError::Chunk(_)));
1776 +
1777 + unsafe { libc::close(fd1) };
1778 +}
1779 +
1780 +#[test]
1781 +fn test_receive_chunk_payload_len_mismatch() {
1782 + let (fd0, fd1) = socketpair_seqpacket();
1783 + let mut session = test_session(fd0, Role::Server, (HEADER_SIZE + 10) as u32);
1784 +
1785 + let first_payload = [1u8; 10];
1786 + let hdr = Header {
1787 + magic: MAGIC_MSG,
1788 + version: VERSION,
1789 + header_len: protocol::HEADER_LEN,
1790 + kind: KIND_REQUEST,
1791 + code: 1,
1792 + flags: 0,
1793 + transport_status: protocol::STATUS_OK,
1794 + payload_len: 20,
1795 + item_count: 1,
1796 + message_id: 42,
1797 + };
1798 + let mut first_pkt = [0u8; HEADER_SIZE + 10];
1799 + hdr.encode(&mut first_pkt[..HEADER_SIZE]);
1800 + first_pkt[HEADER_SIZE..].copy_from_slice(&first_payload);
1801 + raw_send(fd1, &first_pkt).expect("send first chunk");
1802 +
1803 + let second_payload = [2u8; 10];
1804 + let chk = ChunkHeader {
1805 + magic: MAGIC_CHUNK,
1806 + version: VERSION,
1807 + flags: 0,
1808 + message_id: 42,
1809 + total_message_len: (HEADER_SIZE + 20) as u32,
1810 + chunk_index: 1,
1811 + chunk_count: 2,
1812 + chunk_payload_len: 9,
1813 + };
1814 + let mut second_pkt = [0u8; HEADER_SIZE + 10];
1815 + chk.encode(&mut second_pkt[..HEADER_SIZE]);
1816 + second_pkt[HEADER_SIZE..].copy_from_slice(&second_payload);
1817 + raw_send(fd1, &second_pkt).expect("send bad chunk length");
1818 +
1819 + let mut buf = [0u8; 64];
1820 + let err = session
1821 + .receive(&mut buf)
1822 + .expect_err("chunk_payload_len mismatch");
1823 + assert!(matches!(err, UdsError::Chunk(_)));
1824 +
1825 + unsafe { libc::close(fd1) };
1826 +}
1827 +
1828 +#[test]
1829 +fn test_receive_continuation_too_short() {
1830 + let (fd0, fd1) = socketpair_seqpacket();
1831 + let mut session = test_session(fd0, Role::Server, (HEADER_SIZE + 10) as u32);
1832 +
1833 + let first_payload = [1u8; 10];
1834 + let hdr = Header {
1835 + magic: MAGIC_MSG,
1836 + version: VERSION,
1837 + header_len: protocol::HEADER_LEN,
1838 + kind: KIND_REQUEST,
1839 + code: 1,
1840 + flags: 0,
1841 + transport_status: protocol::STATUS_OK,
1842 + payload_len: 20,
1843 + item_count: 1,
1844 + message_id: 43,
1845 + };
1846 + let mut first_pkt = [0u8; HEADER_SIZE + 10];
1847 + hdr.encode(&mut first_pkt[..HEADER_SIZE]);
1848 + first_pkt[HEADER_SIZE..].copy_from_slice(&first_payload);
1849 + raw_send(fd1, &first_pkt).expect("send first chunk");
1850 + raw_send(fd1, &[0x01]).expect("send truncated continuation");
1851 +
1852 + let mut buf = [0u8; 64];
1853 + let err = session
1854 + .receive(&mut buf)
1855 + .expect_err("continuation too short");
1856 + assert!(matches!(err, UdsError::Chunk(_)));
1857 +
1858 + unsafe { libc::close(fd1) };
1859 +}
1860 +
1861 +#[test]
1862 +fn test_receive_chunk_count_mismatch() {
1863 + let (fd0, fd1) = socketpair_seqpacket();
1864 + let mut session = test_session(fd0, Role::Server, (HEADER_SIZE + 10) as u32);
1865 +
1866 + let first_payload = [1u8; 10];
1867 + let hdr = Header {
1868 + magic: MAGIC_MSG,
1869 + version: VERSION,
1870 + header_len: protocol::HEADER_LEN,
1871 + kind: KIND_REQUEST,
1872 + code: 1,
1873 + flags: 0,
1874 + transport_status: protocol::STATUS_OK,
1875 + payload_len: 20,
1876 + item_count: 1,
1877 + message_id: 44,
1878 + };
1879 + let mut first_pkt = [0u8; HEADER_SIZE + 10];
1880 + hdr.encode(&mut first_pkt[..HEADER_SIZE]);
1881 + first_pkt[HEADER_SIZE..].copy_from_slice(&first_payload);
1882 + raw_send(fd1, &first_pkt).expect("send first chunk");
1883 +
1884 + let second_payload = [2u8; 10];
1885 + let chk = ChunkHeader {
1886 + magic: MAGIC_CHUNK,
1887 + version: VERSION,
1888 + flags: 0,
1889 + message_id: 44,
1890 + total_message_len: (HEADER_SIZE + 20) as u32,
1891 + chunk_index: 1,
1892 + chunk_count: 3,
1893 + chunk_payload_len: second_payload.len() as u32,
1894 + };
1895 + let mut second_pkt = [0u8; HEADER_SIZE + 10];
1896 + chk.encode(&mut second_pkt[..HEADER_SIZE]);
1897 + second_pkt[HEADER_SIZE..].copy_from_slice(&second_payload);
1898 + raw_send(fd1, &second_pkt).expect("send bad chunk count");
1899 +
1900 + let mut buf = [0u8; 64];
1901 + let err = session.receive(&mut buf).expect_err("chunk_count mismatch");
1902 + assert!(matches!(err, UdsError::Chunk(_)));
1903 +
1904 + unsafe { libc::close(fd1) };
1905 +}
1906 +
1907 +#[test]
1908 +fn test_receive_chunk_exceeds_payload_len() {
1909 + let (fd0, fd1) = socketpair_seqpacket();
1910 + let mut session = test_session(fd0, Role::Server, (HEADER_SIZE + 10) as u32);
1911 +
1912 + let first_payload = [1u8; 10];
1913 + let hdr = Header {
1914 + magic: MAGIC_MSG,
1915 + version: VERSION,
1916 + header_len: protocol::HEADER_LEN,
1917 + kind: KIND_REQUEST,
1918 + code: 1,
1919 + flags: 0,
1920 + transport_status: protocol::STATUS_OK,
1921 + payload_len: 15,
1922 + item_count: 1,
1923 + message_id: 45,
1924 + };
1925 + let mut first_pkt = [0u8; HEADER_SIZE + 10];
1926 + hdr.encode(&mut first_pkt[..HEADER_SIZE]);
1927 + first_pkt[HEADER_SIZE..].copy_from_slice(&first_payload);
1928 + raw_send(fd1, &first_pkt).expect("send first chunk");
1929 +
1930 + let second_payload = [2u8; 10];
1931 + let chk = ChunkHeader {
1932 + magic: MAGIC_CHUNK,
1933 + version: VERSION,
1934 + flags: 0,
1935 + message_id: 45,
1936 + total_message_len: (HEADER_SIZE + 15) as u32,
1937 + chunk_index: 1,
1938 + chunk_count: 2,
1939 + chunk_payload_len: second_payload.len() as u32,
1940 + };
1941 + let mut second_pkt = [0u8; HEADER_SIZE + 10];
1942 + chk.encode(&mut second_pkt[..HEADER_SIZE]);
1943 + second_pkt[HEADER_SIZE..].copy_from_slice(&second_payload);
1944 + raw_send(fd1, &second_pkt).expect("send oversized continuation");
1945 +
1946 + let mut buf = [0u8; 64];
1947 + let err = session
1948 + .receive(&mut buf)
1949 + .expect_err("chunk exceeds payload_len");
1950 + assert!(matches!(err, UdsError::Chunk(_)));
1951 +
1952 + unsafe { libc::close(fd1) };
1953 +}
1954 +
1955 +#[test]
1956 +fn test_bind_rejects_live_server_addr_in_use() {
1957 + ensure_run_dir();
1958 + let svc = unique_service("rs_live_bind");
1959 + cleanup_socket(&svc);
1960 +
1961 + let listener =
1962 + UdsListener::bind(TEST_RUN_DIR, &svc, default_server_config()).expect("first bind");
1963 + let result = UdsListener::bind(TEST_RUN_DIR, &svc, default_server_config());
1964 + assert!(matches!(result, Err(UdsError::AddrInUse)));
1965 +
1966 + drop(listener);
1967 + cleanup_socket(&svc);
1968 +}
1969 +
1970 +#[test]
1971 +fn test_bind_missing_run_dir_fails() {
1972 + let svc = unique_service("rs_missing_bind");
1973 + let bad_run_dir = "/tmp/nipc_rust_missing_parent/does/not/exist";
1974 + let err = match UdsListener::bind(bad_run_dir, &svc, default_server_config()) {
1975 + Ok(_) => panic!("missing parent bind should fail"),
1976 + Err(err) => err,
1977 + };
1978 + assert!(matches!(err, UdsError::Socket(_)));
1979 +}
1980 +
1981 +#[test]
1982 +fn test_profile_preference_selects_shm_futex() {
1983 + ensure_run_dir();
1984 + let svc = unique_service("rs_pref_prof");
1985 + cleanup_socket(&svc);
1986 +
1987 + let svc_clone = svc.clone();
1988 + let ready = Arc::new(AtomicBool::new(false));
1989 + let ready_clone = ready.clone();
1990 +
1991 + let server_thread = thread::spawn(move || {
1992 + let scfg = ServerConfig {
1993 + supported_profiles: PROFILE_BASELINE | protocol::PROFILE_SHM_FUTEX,
1994 + preferred_profiles: protocol::PROFILE_SHM_FUTEX,
1995 + ..default_server_config()
1996 + };
1997 + let listener = UdsListener::bind(TEST_RUN_DIR, &svc_clone, scfg).expect("listen");
1998 + ready_clone.store(true, Ordering::Release);
1999 + let session = listener.accept().expect("accept");
2000 + assert_eq!(session.selected_profile, protocol::PROFILE_SHM_FUTEX);
2001 + });
2002 +
2003 + while !ready.load(Ordering::Acquire) {
2004 + thread::sleep(Duration::from_millis(1));
2005 + }
2006 +
2007 + let ccfg = ClientConfig {
2008 + supported_profiles: PROFILE_BASELINE | protocol::PROFILE_SHM_FUTEX,
2009 + preferred_profiles: protocol::PROFILE_SHM_FUTEX,
2010 + ..default_client_config()
2011 + };
2012 + let session = UdsSession::connect(TEST_RUN_DIR, &svc, &ccfg).expect("connect");
2013 + assert_eq!(session.selected_profile, protocol::PROFILE_SHM_FUTEX);
2014 +
2015 + drop(session);
2016 + server_thread.join().expect("server join");
2017 + cleanup_socket(&svc);
2018 +}
2019 +
2020 +#[test]
2021 +fn test_zero_supported_profiles_defaults_to_baseline() {
2022 + ensure_run_dir();
2023 + let svc = unique_service("rs_default_profiles");
2024 + cleanup_socket(&svc);
2025 +
2026 + let svc_clone = svc.clone();
2027 + let ready = Arc::new(AtomicBool::new(false));
2028 + let ready_clone = ready.clone();
2029 +
2030 + let server_thread = thread::spawn(move || {
2031 + let scfg = ServerConfig {
2032 + supported_profiles: 0,
2033 + preferred_profiles: 0,
2034 + ..default_server_config()
2035 + };
2036 + let listener = UdsListener::bind(TEST_RUN_DIR, &svc_clone, scfg).expect("listen");
2037 + ready_clone.store(true, Ordering::Release);
2038 + let session = listener.accept().expect("accept");
2039 + assert_eq!(session.selected_profile, PROFILE_BASELINE);
2040 + });
2041 +
2042 + while !ready.load(Ordering::Acquire) {
2043 + thread::sleep(Duration::from_millis(1));
2044 + }
2045 +
2046 + let ccfg = ClientConfig {
2047 + supported_profiles: 0,
2048 + preferred_profiles: 0,
2049 + ..default_client_config()
2050 + };
2051 + let session = UdsSession::connect(TEST_RUN_DIR, &svc, &ccfg).expect("connect");
2052 + assert_eq!(session.selected_profile, PROFILE_BASELINE);
2053 +
2054 + drop(session);
2055 + server_thread.join().expect("server join");
2056 + cleanup_socket(&svc);
2057 +}
2058 +
2059 +#[test]
2060 +fn test_bind_zero_backlog_uses_default() {
2061 + ensure_run_dir();
2062 + let svc = unique_service("rs_default_backlog");
2063 + cleanup_socket(&svc);
2064 +
2065 + let listener = UdsListener::bind(
2066 + TEST_RUN_DIR,
2067 + &svc,
2068 + ServerConfig {
2069 + backlog: 0,
2070 + ..default_server_config()
2071 + },
2072 + )
2073 + .expect("bind with default backlog");
2074 +
2075 + drop(listener);
2076 + cleanup_socket(&svc);
2077 +}
2078 +
2079 +#[test]
2080 +fn test_connect_rejects_bad_hello_ack_kind() {
2081 + ensure_run_dir();
2082 + let svc = unique_service("rs_bad_ack_kind");
2083 + cleanup_socket(&svc);
2084 +
2085 + let ready = Arc::new(AtomicBool::new(false));
2086 + let ready_clone = ready.clone();
2087 + let svc_clone = svc.clone();
2088 + let server = thread::spawn(move || {
2089 + let fd = raw_listener_fd(&svc_clone);
2090 + ready_clone.store(true, Ordering::Release);
2091 + let client_fd = unsafe { libc::accept(fd, std::ptr::null_mut(), std::ptr::null_mut()) };
2092 + assert!(client_fd >= 0, "accept failed: {}", errno());
2093 +
2094 + let mut hello = [0u8; 128];
2095 + let n = raw_recv(client_fd, &mut hello).expect("recv hello");
2096 + assert!(n >= HEADER_SIZE + HELLO_PAYLOAD_SIZE);
2097 +
2098 + let ack = HelloAck {
2099 + layout_version: 1,
2100 + ..HelloAck::default()
2101 + };
2102 + let mut ack_buf = [0u8; HELLO_ACK_PAYLOAD_SIZE];
2103 + ack.encode(&mut ack_buf);
2104 + let hdr = Header {
2105 + magic: MAGIC_MSG,
2106 + version: VERSION,
2107 + header_len: protocol::HEADER_LEN,
2108 + kind: KIND_RESPONSE,
2109 + flags: 0,
2110 + code: protocol::CODE_HELLO_ACK,
2111 + transport_status: protocol::STATUS_OK,
2112 + payload_len: HELLO_ACK_PAYLOAD_SIZE as u32,
2113 + item_count: 1,
2114 + message_id: 0,
2115 + };
2116 + let mut pkt = [0u8; HEADER_SIZE + HELLO_ACK_PAYLOAD_SIZE];
2117 + hdr.encode(&mut pkt[..HEADER_SIZE]);
2118 + pkt[HEADER_SIZE..].copy_from_slice(&ack_buf);
2119 + raw_send(client_fd, &pkt).expect("send malformed ack");
2120 +
2121 + unsafe {
2122 + libc::close(client_fd);
2123 + libc::close(fd);
2124 + }
2125 + });
2126 +
2127 + while !ready.load(Ordering::Acquire) {
2128 + thread::sleep(Duration::from_millis(1));
2129 + }
2130 +
2131 + let err = match UdsSession::connect(TEST_RUN_DIR, &svc, &default_client_config()) {
2132 + Ok(_) => panic!("bad hello ack kind should fail"),
2133 + Err(err) => err,
2134 + };
2135 + assert!(matches!(err, UdsError::Protocol(_)));
2136 +
2137 + server.join().expect("server join");
2138 + cleanup_socket(&svc);
2139 +}
2140 +
2141 +#[test]
2142 +fn test_connect_rejects_bad_hello_ack_status() {
2143 + ensure_run_dir();
2144 + let svc = unique_service("rs_bad_ack_status");
2145 + cleanup_socket(&svc);
2146 +
2147 + let ready = Arc::new(AtomicBool::new(false));
2148 + let ready_clone = ready.clone();
2149 + let svc_clone = svc.clone();
2150 + let server = thread::spawn(move || {
2151 + let fd = raw_listener_fd(&svc_clone);
2152 + ready_clone.store(true, Ordering::Release);
2153 + let client_fd = unsafe { libc::accept(fd, std::ptr::null_mut(), std::ptr::null_mut()) };
2154 + assert!(client_fd >= 0, "accept failed: {}", errno());
2155 +
2156 + let mut hello = [0u8; 128];
2157 + let _ = raw_recv(client_fd, &mut hello).expect("recv hello");
2158 +
2159 + let ack = HelloAck {
2160 + layout_version: 1,
2161 + ..HelloAck::default()
2162 + };
2163 + let mut ack_buf = [0u8; HELLO_ACK_PAYLOAD_SIZE];
2164 + ack.encode(&mut ack_buf);
2165 + let hdr = Header {
2166 + magic: MAGIC_MSG,
2167 + version: VERSION,
2168 + header_len: protocol::HEADER_LEN,
2169 + kind: protocol::KIND_CONTROL,
2170 + flags: 0,
2171 + code: protocol::CODE_HELLO_ACK,
2172 + transport_status: 0x7777,
2173 + payload_len: HELLO_ACK_PAYLOAD_SIZE as u32,
2174 + item_count: 1,
2175 + message_id: 0,
2176 + };
2177 + let mut pkt = [0u8; HEADER_SIZE + HELLO_ACK_PAYLOAD_SIZE];
2178 + hdr.encode(&mut pkt[..HEADER_SIZE]);
2179 + pkt[HEADER_SIZE..].copy_from_slice(&ack_buf);
2180 + raw_send(client_fd, &pkt).expect("send malformed ack");
2181 +
2182 + unsafe {
2183 + libc::close(client_fd);
2184 + libc::close(fd);
2185 + }
2186 + });
2187 +
2188 + while !ready.load(Ordering::Acquire) {
2189 + thread::sleep(Duration::from_millis(1));
2190 + }
2191 +
2192 + let err = match UdsSession::connect(TEST_RUN_DIR, &svc, &default_client_config()) {
2193 + Ok(_) => panic!("bad hello ack status should fail"),
2194 + Err(err) => err,
2195 + };
2196 + assert!(matches!(err, UdsError::Handshake(_)));
2197 +
2198 + server.join().expect("server join");
2199 + cleanup_socket(&svc);
2200 +}
2201 +
2202 +#[test]
2203 +fn test_connect_rejects_bad_hello_ack_version_as_incompatible() {
2204 + ensure_run_dir();
2205 + let svc = unique_service("rs_bad_ack_version");
2206 + cleanup_socket(&svc);
2207 +
2208 + let ready = Arc::new(AtomicBool::new(false));
2209 + let ready_clone = ready.clone();
2210 + let svc_clone = svc.clone();
2211 + let server = thread::spawn(move || {
2212 + let fd = raw_listener_fd(&svc_clone);
2213 + ready_clone.store(true, Ordering::Release);
2214 + let client_fd = unsafe { libc::accept(fd, std::ptr::null_mut(), std::ptr::null_mut()) };
2215 + assert!(client_fd >= 0, "accept failed: {}", errno());
2216 +
2217 + let mut hello = [0u8; 128];
2218 + let _ = raw_recv(client_fd, &mut hello).expect("recv hello");
2219 +
2220 + let ack = HelloAck {
2221 + layout_version: 1,
2222 + ..HelloAck::default()
2223 + };
2224 + let mut ack_buf = [0u8; HELLO_ACK_PAYLOAD_SIZE];
2225 + ack.encode(&mut ack_buf);
2226 + let hdr = Header {
2227 + magic: MAGIC_MSG,
2228 + version: VERSION + 1,
2229 + header_len: protocol::HEADER_LEN,
2230 + kind: protocol::KIND_CONTROL,
2231 + flags: 0,
2232 + code: protocol::CODE_HELLO_ACK,
2233 + transport_status: protocol::STATUS_OK,
2234 + payload_len: HELLO_ACK_PAYLOAD_SIZE as u32,
2235 + item_count: 1,
2236 + message_id: 0,
2237 + };
2238 + let mut pkt = [0u8; HEADER_SIZE + HELLO_ACK_PAYLOAD_SIZE];
2239 + hdr.encode(&mut pkt[..HEADER_SIZE]);
2240 + pkt[HEADER_SIZE..].copy_from_slice(&ack_buf);
2241 + raw_send(client_fd, &pkt).expect("send bad-version ack");
2242 +
2243 + unsafe {
2244 + libc::close(client_fd);
2245 + libc::close(fd);
2246 + }
2247 + });
2248 +
2249 + while !ready.load(Ordering::Acquire) {
2250 + thread::sleep(Duration::from_millis(1));
2251 + }
2252 +
2253 + let err = match UdsSession::connect(TEST_RUN_DIR, &svc, &default_client_config()) {
2254 + Ok(_) => panic!("bad hello ack version should fail"),
2255 + Err(err) => err,
2256 + };
2257 + assert!(matches!(err, UdsError::Incompatible(_)));
2258 +
2259 + server.join().expect("server join");
2260 + cleanup_socket(&svc);
2261 +}
2262 +
2263 +#[test]
2264 +fn test_connect_rejects_incompatible_hello_ack_status() {
2265 + ensure_run_dir();
2266 + let svc = unique_service("rs_ack_incompat_status");
2267 + cleanup_socket(&svc);
2268 +
2269 + let ready = Arc::new(AtomicBool::new(false));
2270 + let ready_clone = ready.clone();
2271 + let svc_clone = svc.clone();
2272 + let server = thread::spawn(move || {
2273 + let fd = raw_listener_fd(&svc_clone);
2274 + ready_clone.store(true, Ordering::Release);
2275 + let client_fd = unsafe { libc::accept(fd, std::ptr::null_mut(), std::ptr::null_mut()) };
2276 + assert!(client_fd >= 0, "accept failed: {}", errno());
2277 +
2278 + let mut hello = [0u8; 128];
2279 + let _ = raw_recv(client_fd, &mut hello).expect("recv hello");
2280 +
2281 + let ack = HelloAck {
2282 + layout_version: 1,
2283 + ..HelloAck::default()
2284 + };
2285 + let mut ack_buf = [0u8; HELLO_ACK_PAYLOAD_SIZE];
2286 + ack.encode(&mut ack_buf);
2287 + let hdr = Header {
2288 + magic: MAGIC_MSG,
2289 + version: VERSION,
2290 + header_len: protocol::HEADER_LEN,
2291 + kind: protocol::KIND_CONTROL,
2292 + flags: 0,
2293 + code: protocol::CODE_HELLO_ACK,
2294 + transport_status: protocol::STATUS_INCOMPATIBLE,
2295 + payload_len: HELLO_ACK_PAYLOAD_SIZE as u32,
2296 + item_count: 1,
2297 + message_id: 0,
2298 + };
2299 + let mut pkt = [0u8; HEADER_SIZE + HELLO_ACK_PAYLOAD_SIZE];
2300 + hdr.encode(&mut pkt[..HEADER_SIZE]);
2301 + pkt[HEADER_SIZE..].copy_from_slice(&ack_buf);
2302 + raw_send(client_fd, &pkt).expect("send incompatible-status ack");
2303 +
2304 + unsafe {
2305 + libc::close(client_fd);
2306 + libc::close(fd);
2307 + }
2308 + });
2309 +
2310 + while !ready.load(Ordering::Acquire) {
2311 + thread::sleep(Duration::from_millis(1));
2312 + }
2313 +
2314 + let err = match UdsSession::connect(TEST_RUN_DIR, &svc, &default_client_config()) {
2315 + Ok(_) => panic!("incompatible hello ack status should fail"),
2316 + Err(err) => err,
2317 + };
2318 + assert!(matches!(err, UdsError::Incompatible(_)));
2319 +
2320 + server.join().expect("server join");
2321 + cleanup_socket(&svc);
2322 +}
2323 +
2324 +#[test]
2325 +fn test_connect_rejects_bad_hello_ack_layout_as_incompatible() {
2326 + ensure_run_dir();
2327 + let svc = unique_service("rs_bad_ack_layout");
2328 + cleanup_socket(&svc);
2329 +
2330 + let ready = Arc::new(AtomicBool::new(false));
2331 + let ready_clone = ready.clone();
2332 + let svc_clone = svc.clone();
2333 + let server = thread::spawn(move || {
2334 + let fd = raw_listener_fd(&svc_clone);
2335 + ready_clone.store(true, Ordering::Release);
2336 + let client_fd = unsafe { libc::accept(fd, std::ptr::null_mut(), std::ptr::null_mut()) };
2337 + assert!(client_fd >= 0, "accept failed: {}", errno());
2338 +
2339 + let mut hello = [0u8; 128];
2340 + let _ = raw_recv(client_fd, &mut hello).expect("recv hello");
2341 +
2342 + let ack = HelloAck {
2343 + layout_version: 2,
2344 + ..HelloAck::default()
2345 + };
2346 + let mut ack_buf = [0u8; HELLO_ACK_PAYLOAD_SIZE];
2347 + ack.encode(&mut ack_buf);
2348 + let hdr = Header {
2349 + magic: MAGIC_MSG,
2350 + version: VERSION,
2351 + header_len: protocol::HEADER_LEN,
2352 + kind: protocol::KIND_CONTROL,
2353 + flags: 0,
2354 + code: protocol::CODE_HELLO_ACK,
2355 + transport_status: protocol::STATUS_OK,
2356 + payload_len: HELLO_ACK_PAYLOAD_SIZE as u32,
2357 + item_count: 1,
2358 + message_id: 0,
2359 + };
2360 + let mut pkt = [0u8; HEADER_SIZE + HELLO_ACK_PAYLOAD_SIZE];
2361 + hdr.encode(&mut pkt[..HEADER_SIZE]);
2362 + pkt[HEADER_SIZE..].copy_from_slice(&ack_buf);
2363 + raw_send(client_fd, &pkt).expect("send bad-layout ack");
2364 +
2365 + unsafe {
2366 + libc::close(client_fd);
2367 + libc::close(fd);
2368 + }
2369 + });
2370 +
2371 + while !ready.load(Ordering::Acquire) {
2372 + thread::sleep(Duration::from_millis(1));
2373 + }
2374 +
2375 + let err = match UdsSession::connect(TEST_RUN_DIR, &svc, &default_client_config()) {
2376 + Ok(_) => panic!("bad hello ack layout should fail"),
2377 + Err(err) => err,
2378 + };
2379 + assert!(matches!(err, UdsError::Incompatible(_)));
2380 +
2381 + server.join().expect("server join");
2382 + cleanup_socket(&svc);
2383 +}
2384 +
2385 +#[test]
2386 +fn test_connect_rejects_truncated_hello_ack() {
2387 + ensure_run_dir();
2388 + let svc = unique_service("rs_bad_ack_trunc");
2389 + cleanup_socket(&svc);
2390 +
2391 + let ready = Arc::new(AtomicBool::new(false));
2392 + let ready_clone = ready.clone();
2393 + let svc_clone = svc.clone();
2394 + let server = thread::spawn(move || {
2395 + let fd = raw_listener_fd(&svc_clone);
2396 + ready_clone.store(true, Ordering::Release);
2397 + let client_fd = unsafe { libc::accept(fd, std::ptr::null_mut(), std::ptr::null_mut()) };
2398 + assert!(client_fd >= 0, "accept failed: {}", errno());
2399 +
2400 + let mut hello = [0u8; 128];
2401 + let _ = raw_recv(client_fd, &mut hello).expect("recv hello");
2402 +
2403 + let hdr = Header {
2404 + magic: MAGIC_MSG,
2405 + version: VERSION,
2406 + header_len: protocol::HEADER_LEN,
2407 + kind: protocol::KIND_CONTROL,
2408 + flags: 0,
2409 + code: protocol::CODE_HELLO_ACK,
2410 + transport_status: protocol::STATUS_OK,
2411 + payload_len: HELLO_ACK_PAYLOAD_SIZE as u32,
2412 + item_count: 1,
2413 + message_id: 0,
2414 + };
2415 + let mut pkt = [0u8; HEADER_SIZE];
2416 + hdr.encode(&mut pkt);
2417 + raw_send(client_fd, &pkt).expect("send truncated ack");
2418 +
2419 + unsafe {
2420 + libc::close(client_fd);
2421 + libc::close(fd);
2422 + }
2423 + });
2424 +
2425 + while !ready.load(Ordering::Acquire) {
2426 + thread::sleep(Duration::from_millis(1));
2427 + }
2428 +
2429 + let err = match UdsSession::connect(TEST_RUN_DIR, &svc, &default_client_config()) {
2430 + Ok(_) => panic!("truncated hello ack should fail"),
2431 + Err(err) => err,
2432 + };
2433 + assert!(matches!(err, UdsError::Protocol(_)));
2434 +
2435 + server.join().expect("server join");
2436 + cleanup_socket(&svc);
2437 +}
2438 +
2439 +#[test]
2440 +fn test_accept_rejects_bad_hello_kind() {
2441 + ensure_run_dir();
2442 + let svc = unique_service("rs_bad_hello_kind");
2443 + cleanup_socket(&svc);
2444 +
2445 + let listener =
2446 + UdsListener::bind(TEST_RUN_DIR, &svc, default_server_config()).expect("bind listener");
2447 +
2448 + let svc_clone = svc.clone();
2449 + let client = thread::spawn(move || {
2450 + let path = build_socket_path(TEST_RUN_DIR, &svc_clone).expect("socket path");
2451 + let fd = unsafe { libc::socket(libc::AF_UNIX, libc::SOCK_SEQPACKET, 0) };
2452 + assert!(fd >= 0, "socket failed: {}", errno());
2453 + connect_unix(fd, &path).expect("connect");
2454 +
2455 + let hdr = Header {
2456 + magic: MAGIC_MSG,
2457 + version: VERSION,
2458 + header_len: protocol::HEADER_LEN,
2459 + kind: KIND_REQUEST,
2460 + flags: 0,
2461 + code: protocol::CODE_HELLO,
2462 + transport_status: protocol::STATUS_OK,
2463 + payload_len: HELLO_PAYLOAD_SIZE as u32,
2464 + item_count: 1,
2465 + message_id: 0,
2466 + };
2467 + let mut payload = [0u8; HELLO_PAYLOAD_SIZE];
2468 + let hello = Hello {
2469 + layout_version: 1,
2470 + supported_profiles: PROFILE_BASELINE,
2471 + max_request_payload_bytes: 4096,
2472 + max_request_batch_items: 16,
2473 + max_response_payload_bytes: 4096,
2474 + max_response_batch_items: 16,
2475 + auth_token: AUTH_TOKEN,
2476 + packet_size: DEFAULT_PACKET_SIZE_FALLBACK,
2477 + ..Hello::default()
2478 + };
2479 + hello.encode(&mut payload);
2480 + let mut pkt = [0u8; HEADER_SIZE + HELLO_PAYLOAD_SIZE];
2481 + hdr.encode(&mut pkt[..HEADER_SIZE]);
2482 + pkt[HEADER_SIZE..].copy_from_slice(&payload);
2483 + raw_send(fd, &pkt).expect("send malformed hello");
2484 + unsafe { libc::close(fd) };
2485 + });
2486 +
2487 + let err = match listener.accept() {
2488 + Ok(_) => panic!("bad hello kind should fail"),
2489 + Err(err) => err,
2490 + };
2491 + assert!(matches!(err, UdsError::Protocol(_)));
2492 +
2493 + client.join().expect("client join");
2494 + drop(listener);
2495 + cleanup_socket(&svc);
2496 +}
2497 +
2498 +#[test]
2499 +fn test_accept_rejects_bad_hello_version_as_incompatible() {
2500 + ensure_run_dir();
2501 + let svc = unique_service("rs_bad_hello_version");
2502 + cleanup_socket(&svc);
2503 +
2504 + let listener =
2505 + UdsListener::bind(TEST_RUN_DIR, &svc, default_server_config()).expect("bind listener");
2506 +
2507 + let svc_clone = svc.clone();
2508 + let client = thread::spawn(move || {
2509 + let path = build_socket_path(TEST_RUN_DIR, &svc_clone).expect("socket path");
2510 + let fd = unsafe { libc::socket(libc::AF_UNIX, libc::SOCK_SEQPACKET, 0) };
2511 + assert!(fd >= 0, "socket failed: {}", errno());
2512 + connect_unix(fd, &path).expect("connect");
2513 +
2514 + let hdr = Header {
2515 + magic: MAGIC_MSG,
2516 + version: VERSION + 1,
2517 + header_len: protocol::HEADER_LEN,
2518 + kind: protocol::KIND_CONTROL,
2519 + flags: 0,
2520 + code: protocol::CODE_HELLO,
2521 + transport_status: protocol::STATUS_OK,
2522 + payload_len: HELLO_PAYLOAD_SIZE as u32,
2523 + item_count: 1,
2524 + message_id: 0,
2525 + };
2526 + let mut payload = [0u8; HELLO_PAYLOAD_SIZE];
2527 + let hello = Hello {
2528 + layout_version: 1,
2529 + supported_profiles: PROFILE_BASELINE,
2530 + max_request_payload_bytes: 4096,
2531 + max_request_batch_items: 16,
2532 + max_response_payload_bytes: 4096,
2533 + max_response_batch_items: 16,
2534 + auth_token: AUTH_TOKEN,
2535 + packet_size: DEFAULT_PACKET_SIZE_FALLBACK,
2536 + ..Hello::default()
2537 + };
2538 + hello.encode(&mut payload);
2539 + let mut pkt = [0u8; HEADER_SIZE + HELLO_PAYLOAD_SIZE];
2540 + hdr.encode(&mut pkt[..HEADER_SIZE]);
2541 + pkt[HEADER_SIZE..].copy_from_slice(&payload);
2542 + raw_send(fd, &pkt).expect("send incompatible hello");
2543 + unsafe { libc::close(fd) };
2544 + });
2545 +
2546 + let err = match listener.accept() {
2547 + Ok(_) => panic!("bad hello version should fail"),
2548 + Err(err) => err,
2549 + };
2550 + assert!(matches!(err, UdsError::Incompatible(_)));
2551 +
2552 + client.join().expect("client join");
2553 + drop(listener);
2554 + cleanup_socket(&svc);
2555 +}
2556 +
2557 +#[test]
2558 +fn test_accept_rejects_bad_hello_layout_as_incompatible() {
2559 + ensure_run_dir();
2560 + let svc = unique_service("rs_bad_hello_layout");
2561 + cleanup_socket(&svc);
2562 +
2563 + let listener =
2564 + UdsListener::bind(TEST_RUN_DIR, &svc, default_server_config()).expect("bind listener");
2565 +
2566 + let svc_clone = svc.clone();
2567 + let client = thread::spawn(move || {
2568 + let path = build_socket_path(TEST_RUN_DIR, &svc_clone).expect("socket path");
2569 + let fd = unsafe { libc::socket(libc::AF_UNIX, libc::SOCK_SEQPACKET, 0) };
2570 + assert!(fd >= 0, "socket failed: {}", errno());
2571 + connect_unix(fd, &path).expect("connect");
2572 +
2573 + let hdr = Header {
2574 + magic: MAGIC_MSG,
2575 + version: VERSION,
2576 + header_len: protocol::HEADER_LEN,
2577 + kind: protocol::KIND_CONTROL,
2578 + flags: 0,
2579 + code: protocol::CODE_HELLO,
2580 + transport_status: protocol::STATUS_OK,
2581 + payload_len: HELLO_PAYLOAD_SIZE as u32,
2582 + item_count: 1,
2583 + message_id: 0,
2584 + };
2585 + let mut payload = [0u8; HELLO_PAYLOAD_SIZE];
2586 + let hello = Hello {
2587 + layout_version: 2,
2588 + supported_profiles: PROFILE_BASELINE,
2589 + max_request_payload_bytes: 4096,
2590 + max_request_batch_items: 16,
2591 + max_response_payload_bytes: 4096,
2592 + max_response_batch_items: 16,
2593 + auth_token: AUTH_TOKEN,
2594 + packet_size: DEFAULT_PACKET_SIZE_FALLBACK,
2595 + ..Hello::default()
2596 + };
2597 + hello.encode(&mut payload);
2598 + let mut pkt = [0u8; HEADER_SIZE + HELLO_PAYLOAD_SIZE];
2599 + hdr.encode(&mut pkt[..HEADER_SIZE]);
2600 + pkt[HEADER_SIZE..].copy_from_slice(&payload);
2601 + raw_send(fd, &pkt).expect("send incompatible-layout hello");
2602 + unsafe { libc::close(fd) };
2603 + });
2604 +
2605 + let err = match listener.accept() {
2606 + Ok(_) => panic!("bad hello layout should fail"),
2607 + Err(err) => err,
2608 + };
2609 + assert!(matches!(err, UdsError::Incompatible(_)));
2610 +
2611 + client.join().expect("client join");
2612 + drop(listener);
2613 + cleanup_socket(&svc);
2614 +}
2615 +
2616 +// -----------------------------------------------------------------------
2617 +// Receive: unknown response message_id (line 370)
2618 +// -----------------------------------------------------------------------
2619 +
2620 +#[test]
2621 +fn test_unknown_response_msg_id() {
2622 + ensure_run_dir();
2623 + let svc = unique_service("rs_unkmsg");
2624 + cleanup_socket(&svc);
2625 +
2626 + let svc_clone = svc.clone();
2627 + let ready = Arc::new(AtomicBool::new(false));
2628 + let ready_clone = ready.clone();
2629 +
2630 + let server_thread = thread::spawn(move || {
2631 + let listener =
2632 + UdsListener::bind(TEST_RUN_DIR, &svc_clone, default_server_config()).expect("listen");
2633 + ready_clone.store(true, Ordering::Release);
2634 +
2635 + let mut session = listener.accept().expect("accept");
2636 + let mut buf = [0u8; 8192];
2637 + let (hdr, payload) = session.receive(&mut buf).expect("recv");
2638 + let payload = payload.to_vec();
2639 +
2640 + // Respond with a different message_id
2641 + let mut resp = Header {
2642 + kind: protocol::KIND_RESPONSE,
2643 + code: hdr.code,
2644 + message_id: hdr.message_id + 999, // wrong message_id
2645 + item_count: 1,
2646 + transport_status: protocol::STATUS_OK,
2647 + ..Header::default()
2648 + };
2649 + session.send(&mut resp, &payload).expect("send");
2650 + });
2651 +
2652 + while !ready.load(Ordering::Acquire) {
2653 + thread::sleep(Duration::from_millis(1));
2654 + }
2655 +
2656 + let mut session =
2657 + UdsSession::connect(TEST_RUN_DIR, &svc, &default_client_config()).expect("connect");
2658 +
2659 + let mut hdr = Header {
2660 + kind: protocol::KIND_REQUEST,
2661 + code: 1,
2662 + item_count: 1,
2663 + message_id: 50,
2664 + ..Header::default()
2665 + };
2666 + session.send(&mut hdr, &[1]).expect("send");
2667 +
2668 + let mut rbuf = [0u8; 8192];
2669 + let result = session.receive(&mut rbuf);
2670 + assert!(matches!(result, Err(UdsError::UnknownMsgId(_))));
2671 +
2672 + drop(session);
2673 + server_thread.join().expect("server join");
2674 + cleanup_socket(&svc);
2675 +}
2676 +
2677 +// -----------------------------------------------------------------------
2678 +// Path length validation
2679 +// -----------------------------------------------------------------------
2680 +
2681 +#[test]
2682 +fn test_path_too_long() {
2683 + let long_name = "a".repeat(200);
2684 + let result = build_socket_path("/tmp", &long_name);
2685 + assert!(matches!(result, Err(UdsError::PathTooLong)));
2686 +}
2687 +
2688 +// -----------------------------------------------------------------------
2689 +// Helpers: highest_bit, apply_default
2690 +// -----------------------------------------------------------------------
2691 +
2692 +#[test]
2693 +fn test_highest_bit() {
2694 + assert_eq!(highest_bit(0), 0);
2695 + assert_eq!(highest_bit(1), 1);
2696 + assert_eq!(highest_bit(0b0101), 4);
2697 + assert_eq!(highest_bit(0b1000), 8);
2698 + assert_eq!(highest_bit(0xFF), 128);
2699 +}
2700 +
2701 +#[test]
2702 +fn test_apply_default() {
2703 + assert_eq!(apply_default(0, 42), 42);
2704 + assert_eq!(apply_default(10, 42), 10);
2705 +}
2706 +
2707 +#[test]
2708 +fn test_detect_packet_size_invalid_fd_returns_fallback() {
2709 + assert_eq!(detect_packet_size(-1), DEFAULT_PACKET_SIZE_FALLBACK);
2710 +}
2711 +
2712 +#[test]
2713 +fn test_raw_send_closed_peer_returns_send_error() {
2714 + let (fd0, fd1) = socketpair_seqpacket();
2715 + unsafe {
2716 + libc::close(fd1);
2717 + }
2718 +
2719 + let err = raw_send(fd0, &[1, 2, 3]).expect_err("closed peer should fail send");
2720 + assert!(matches!(err, UdsError::Send(_)));
2721 +
2722 + unsafe {
2723 + libc::close(fd0);
2724 + }
2725 +}
src/crates/netipc/src/transport/shm.rs new
+1049
@@ -0,0 +1,1049 @@
1 +//! L1 POSIX SHM transport (Linux only).
2 +//!
3 +//! Shared memory data plane with spin+futex synchronization.
4 +//! The SHM region carries the same outer protocol envelope as the UDS
5 +//! transport. Higher levels see no difference.
6 +//!
7 +//! Wire-compatible with the C implementation in netipc_shm.c.
8 +
9 +use std::path::{Path, PathBuf};
10 +use std::ptr;
11 +
12 +// ---------------------------------------------------------------------------
13 +// Constants
14 +// ---------------------------------------------------------------------------
15 +
16 +/// Magic value: "NSHM" as u32 LE.
17 +pub const REGION_MAGIC: u32 = 0x4e53484d;
18 +pub const REGION_VERSION: u16 = 3;
19 +pub const REGION_ALIGNMENT: u32 = 64;
20 +pub const HEADER_LEN: u16 = 64;
21 +pub const DEFAULT_SPIN_TRIES: u32 = 128;
22 +
23 +// Byte offsets of atomic fields in the region header.
24 +const OFF_REQ_SEQ: usize = 32;
25 +const OFF_RESP_SEQ: usize = 40;
26 +const OFF_REQ_LEN: usize = 48;
27 +const OFF_RESP_LEN: usize = 52;
28 +const OFF_REQ_SIGNAL: usize = 56;
29 +const OFF_RESP_SIGNAL: usize = 60;
30 +
31 +// futex operations
32 +const FUTEX_WAIT: i32 = 0;
33 +const FUTEX_WAKE: i32 = 1;
34 +
35 +// ---------------------------------------------------------------------------
36 +// Errors
37 +// ---------------------------------------------------------------------------
38 +
39 +/// SHM transport errors.
40 +#[derive(Debug, Clone, PartialEq, Eq)]
41 +pub enum ShmError {
42 + /// SHM path exceeds limit.
43 + PathTooLong,
44 + /// open/shm_open failed.
45 + Open(i32),
46 + /// ftruncate failed.
47 + Truncate(i32),
48 + /// mmap failed.
49 + Mmap(i32),
50 + /// Header magic mismatch.
51 + BadMagic,
52 + /// Header version mismatch.
53 + BadVersion,
54 + /// header_len mismatch or corrupt.
55 + BadHeader,
56 + /// File too small / capacity mismatch.
57 + BadSize,
58 + /// Live server owns the region.
59 + AddrInUse,
60 + /// Server hasn't finished setup (retry).
61 + NotReady,
62 + /// Message exceeds area capacity.
63 + MsgTooLarge,
64 + /// Futex wait timed out.
65 + Timeout,
66 + /// Invalid argument.
67 + BadParam(String),
68 + /// Owner process has exited.
69 + PeerDead,
70 +}
71 +
72 +impl std::fmt::Display for ShmError {
73 + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
74 + match self {
75 + ShmError::PathTooLong => write!(f, "SHM path exceeds limit"),
76 + ShmError::Open(e) => write!(f, "open failed: errno {e}"),
77 + ShmError::Truncate(e) => write!(f, "ftruncate failed: errno {e}"),
78 + ShmError::Mmap(e) => write!(f, "mmap failed: errno {e}"),
79 + ShmError::BadMagic => write!(f, "SHM header magic mismatch"),
80 + ShmError::BadVersion => write!(f, "SHM header version mismatch"),
81 + ShmError::BadHeader => write!(f, "SHM header_len mismatch"),
82 + ShmError::BadSize => write!(f, "SHM file too small for declared areas"),
83 + ShmError::AddrInUse => write!(f, "SHM region owned by live server"),
84 + ShmError::NotReady => write!(f, "SHM server not ready"),
85 + ShmError::MsgTooLarge => write!(f, "message exceeds SHM area capacity"),
86 + ShmError::Timeout => write!(f, "SHM futex wait timed out"),
87 + ShmError::BadParam(s) => write!(f, "bad parameter: {s}"),
88 + ShmError::PeerDead => write!(f, "SHM owner process has exited"),
89 + }
90 + }
91 +}
92 +
93 +impl std::error::Error for ShmError {}
94 +
95 +// ---------------------------------------------------------------------------
96 +// Role
97 +// ---------------------------------------------------------------------------
98 +
99 +#[derive(Debug, Clone, Copy, PartialEq, Eq)]
100 +pub enum ShmRole {
101 + Server = 1,
102 + Client = 2,
103 +}
104 +
105 +// ---------------------------------------------------------------------------
106 +// Region header (64 bytes at offset 0)
107 +// ---------------------------------------------------------------------------
108 +
109 +/// On-disk layout. Not used directly for atomic accesses; we use
110 +/// raw pointer arithmetic like the C implementation.
111 +#[repr(C)]
112 +struct RegionHeader {
113 + magic: u32, // 0
114 + version: u16, // 4
115 + header_len: u16, // 6
116 + owner_pid: i32, // 8
117 + owner_generation: u32, // 12
118 + request_offset: u32, // 16
119 + request_capacity: u32, // 20
120 + response_offset: u32, // 24
121 + response_capacity: u32, // 28
122 + req_seq: u64, // 32
123 + resp_seq: u64, // 40
124 + req_len: u32, // 48
125 + resp_len: u32, // 52
126 + req_signal: u32, // 56
127 + resp_signal: u32, // 60
128 +}
129 +
130 +const _: () = assert!(std::mem::size_of::<RegionHeader>() == 64);
131 +
132 +// ---------------------------------------------------------------------------
133 +// SHM context
134 +// ---------------------------------------------------------------------------
135 +
136 +/// A handle to a shared memory region (server or client side).
137 +pub struct ShmContext {
138 + role: ShmRole,
139 + fd: i32,
140 + base: *mut u8,
141 + region_size: usize,
142 +
143 + // Cached from header
144 + request_offset: u32,
145 + request_capacity: u32,
146 + response_offset: u32,
147 + response_capacity: u32,
148 +
149 + // Sequence tracking
150 + local_req_seq: u64,
151 + local_resp_seq: u64,
152 +
153 + spin_tries: u32,
154 + owner_generation: u32,
155 + path: PathBuf,
156 +}
157 +
158 +// ShmContext is Send: the mmap pointer is process-global shared memory,
159 +// and the context is used by one thread at a time (single in-flight).
160 +unsafe impl Send for ShmContext {}
161 +
162 +impl ShmContext {
163 + /// Returns the role.
164 + pub fn role(&self) -> ShmRole {
165 + self.role
166 + }
167 +
168 + /// Returns the raw file descriptor.
169 + pub fn fd(&self) -> i32 {
170 + self.fd
171 + }
172 +
173 + /// Check if the region's owner process is still alive.
174 + pub fn owner_alive(&self) -> bool {
175 + if self.base.is_null() || self.region_size < HEADER_LEN as usize {
176 + return false;
177 + }
178 + // Read the PID and generation fields by computing their byte offsets
179 + // directly from self.base, avoiding any pointer dereference through
180 + // a reference type. This is safer than casting to *const RegionHeader
181 + // and dereferencing because we never materialize a borrow of the
182 + // header struct — we just read raw bytes at known offsets.
183 + //
184 + // SAFETY: self.base is a non-null mmap'd region of at least HEADER_LEN
185 + // bytes (checked above). The field offsets below are within HEADER_LEN.
186 + // owner_pid is at offset 8, owner_generation at offset 12
187 + let pid: i32 = unsafe {
188 + let p = self.base.add(8) as *const i32;
189 + ptr::read_volatile(p)
190 + };
191 + if !pid_alive(pid) {
192 + return false;
193 + }
194 + if self.owner_generation != 0 {
195 + let cur_gen: u32 = unsafe {
196 + let p = self.base.add(12) as *const u32;
197 + ptr::read_volatile(p)
198 + };
199 + if cur_gen != self.owner_generation {
200 + return false;
201 + }
202 + }
203 + true
204 + }
205 +
206 + /// Create a SHM region (server side).
207 + ///
208 + /// Creates `{run_dir}/{service_name}-{session_id:016x}.ipcshm` with O_EXCL.
209 + /// Pre-checks for stale files and unlinks them before creating.
210 + pub fn server_create(
211 + run_dir: &str,
212 + service_name: &str,
213 + session_id: u64,
214 + req_capacity: u32,
215 + resp_capacity: u32,
216 + ) -> Result<Self, ShmError> {
217 + let path = build_shm_path(run_dir, service_name, session_id)?;
218 +
219 + // Round capacities to alignment (fails if rounding would overflow u32)
220 + let req_cap = align64(req_capacity)
221 + .ok_or_else(|| ShmError::BadParam("request capacity overflow".into()))?;
222 + let resp_cap = align64(resp_capacity)
223 + .ok_or_else(|| ShmError::BadParam("response capacity overflow".into()))?;
224 +
225 + let req_off = align64(HEADER_LEN as u32)
226 + .ok_or_else(|| ShmError::BadParam("header offset overflow".into()))?;
227 + let resp_off = req_off
228 + .checked_add(req_cap)
229 + .and_then(align64)
230 + .ok_or_else(|| ShmError::BadParam("region offset overflow".into()))?;
231 + let region_size = resp_off
232 + .checked_add(resp_cap)
233 + .ok_or_else(|| ShmError::BadParam("region size overflow".into()))?
234 + as usize;
235 +
236 + // Try O_EXCL create first (fast path, no stale check needed).
237 + let c_path = path_to_cstring(&path)?;
238 + let mut fd = unsafe {
239 + libc::open(
240 + c_path.as_ptr(),
241 + libc::O_RDWR | libc::O_CREAT | libc::O_EXCL,
242 + 0o600,
243 + )
244 + };
245 +
246 + // If O_EXCL failed (file exists), do stale recovery and retry.
247 + if fd < 0 && unsafe { *libc::__errno_location() } == libc::EEXIST {
248 + let stale = check_shm_stale(&path);
249 + if stale == StaleResult::LiveServer {
250 + return Err(ShmError::AddrInUse);
251 + }
252 + // Stale file was (hopefully) unlinked, retry create
253 + fd = unsafe {
254 + libc::open(
255 + c_path.as_ptr(),
256 + libc::O_RDWR | libc::O_CREAT | libc::O_EXCL,
257 + 0o600,
258 + )
259 + };
260 + // If retry still fails with EEXIST, the stale check couldn't
261 + // unlink (e.g., EACCES) — treat as address-in-use rather than
262 + // leaking EEXIST up the stack.
263 + if fd < 0 && unsafe { *libc::__errno_location() } == libc::EEXIST {
264 + return Err(ShmError::AddrInUse);
265 + }
266 + }
267 + if fd < 0 {
268 + return Err(ShmError::Open(errno()));
269 + }
270 +
271 + if unsafe { libc::ftruncate(fd, region_size as libc::off_t) } < 0 {
272 + let e = errno();
273 + unsafe {
274 + libc::close(fd);
275 + libc::unlink(c_path.as_ptr());
276 + }
277 + return Err(ShmError::Truncate(e));
278 + }
279 +
280 + let base = unsafe {
281 + libc::mmap(
282 + ptr::null_mut(),
283 + region_size,
284 + libc::PROT_READ | libc::PROT_WRITE,
285 + libc::MAP_SHARED,
286 + fd,
287 + 0,
288 + )
289 + };
290 + if base == libc::MAP_FAILED {
291 + let e = errno();
292 + unsafe {
293 + libc::close(fd);
294 + libc::unlink(c_path.as_ptr());
295 + }
296 + return Err(ShmError::Mmap(e));
297 + }
298 +
299 + let base = base as *mut u8;
300 +
301 + // Zero the region
302 + unsafe { ptr::write_bytes(base, 0, region_size) };
303 +
304 + // Use a time-based generation to detect PID reuse across restarts.
305 + let generation = {
306 + let mut ts: libc::timespec = unsafe { std::mem::zeroed() };
307 + unsafe { libc::clock_gettime(libc::CLOCK_MONOTONIC, &mut ts) };
308 + (ts.tv_sec as u32) ^ ((ts.tv_nsec >> 10) as u32)
309 + };
310 +
311 + // Write header
312 + let hdr = base as *mut RegionHeader;
313 + unsafe {
314 + (*hdr).magic = REGION_MAGIC;
315 + (*hdr).version = REGION_VERSION;
316 + (*hdr).header_len = HEADER_LEN;
317 + (*hdr).owner_pid = libc::getpid();
318 + (*hdr).owner_generation = generation;
319 + (*hdr).request_offset = req_off;
320 + (*hdr).request_capacity = req_cap;
321 + (*hdr).response_offset = resp_off;
322 + (*hdr).response_capacity = resp_cap;
323 + }
324 +
325 + // Release fence so clients see header writes
326 + std::sync::atomic::fence(std::sync::atomic::Ordering::Release);
327 +
328 + Ok(ShmContext {
329 + role: ShmRole::Server,
330 + fd,
331 + base,
332 + region_size,
333 + request_offset: req_off,
334 + request_capacity: req_cap,
335 + response_offset: resp_off,
336 + response_capacity: resp_cap,
337 + local_req_seq: 0,
338 + local_resp_seq: 0,
339 + spin_tries: DEFAULT_SPIN_TRIES,
340 + owner_generation: generation,
341 + path,
342 + })
343 + }
344 +
345 + /// Attach to an existing SHM region (client side).
346 + pub fn client_attach(
347 + run_dir: &str,
348 + service_name: &str,
349 + session_id: u64,
350 + ) -> Result<Self, ShmError> {
351 + let path = build_shm_path(run_dir, service_name, session_id)?;
352 + let c_path = path_to_cstring(&path)?;
353 +
354 + let fd = unsafe { libc::open(c_path.as_ptr(), libc::O_RDWR) };
355 + if fd < 0 {
356 + return Err(ShmError::Open(errno()));
357 + }
358 +
359 + // Check file size
360 + let mut st: libc::stat = unsafe { std::mem::zeroed() };
361 + if unsafe { libc::fstat(fd, &mut st) } < 0 {
362 + unsafe { libc::close(fd) };
363 + return Err(ShmError::Open(errno()));
364 + }
365 +
366 + let file_size = st.st_size as usize;
367 + if file_size < HEADER_LEN as usize {
368 + unsafe { libc::close(fd) };
369 + return Err(ShmError::NotReady);
370 + }
371 +
372 + let base = unsafe {
373 + libc::mmap(
374 + ptr::null_mut(),
375 + file_size,
376 + libc::PROT_READ | libc::PROT_WRITE,
377 + libc::MAP_SHARED,
378 + fd,
379 + 0,
380 + )
381 + };
382 + if base == libc::MAP_FAILED {
383 + unsafe { libc::close(fd) };
384 + return Err(ShmError::Mmap(errno()));
385 + }
386 +
387 + let base = base as *mut u8;
388 +
389 + // Acquire fence to see server's header writes
390 + std::sync::atomic::fence(std::sync::atomic::Ordering::Acquire);
391 +
392 + let hdr = base as *const RegionHeader;
393 + let (magic, version, hdr_len) =
394 + unsafe { ((*hdr).magic, (*hdr).version, (*hdr).header_len) };
395 +
396 + if magic != REGION_MAGIC {
397 + unsafe {
398 + libc::munmap(base as *mut libc::c_void, file_size);
399 + libc::close(fd);
400 + }
401 + return Err(ShmError::BadMagic);
402 + }
403 + if version != REGION_VERSION {
404 + unsafe {
405 + libc::munmap(base as *mut libc::c_void, file_size);
406 + libc::close(fd);
407 + }
408 + return Err(ShmError::BadVersion);
409 + }
410 + if hdr_len != HEADER_LEN {
411 + unsafe {
412 + libc::munmap(base as *mut libc::c_void, file_size);
413 + libc::close(fd);
414 + }
415 + return Err(ShmError::BadHeader);
416 + }
417 +
418 + let (req_off, req_cap, resp_off, resp_cap) = unsafe {
419 + (
420 + (*hdr).request_offset,
421 + (*hdr).request_capacity,
422 + (*hdr).response_offset,
423 + (*hdr).response_capacity,
424 + )
425 + };
426 +
427 + let header_end = align64(HEADER_LEN as u32).unwrap_or(HEADER_LEN as u32);
428 + if req_off < header_end || req_cap == 0 || resp_off < header_end || resp_cap == 0 {
429 + unsafe {
430 + libc::munmap(base as *mut libc::c_void, file_size);
431 + libc::close(fd);
432 + }
433 + return Err(ShmError::NotReady);
434 + }
435 +
436 + if req_off % REGION_ALIGNMENT != 0
437 + || req_cap % REGION_ALIGNMENT != 0
438 + || resp_off % REGION_ALIGNMENT != 0
439 + || resp_cap % REGION_ALIGNMENT != 0
440 + || req_off
441 + .checked_add(req_cap)
442 + .and_then(align64)
443 + .map_or(true, |min| resp_off < min)
444 + {
445 + unsafe {
446 + libc::munmap(base as *mut libc::c_void, file_size);
447 + libc::close(fd);
448 + }
449 + return Err(ShmError::BadSize);
450 + }
451 +
452 + // Validate region size
453 + let req_end = req_off as usize + req_cap as usize;
454 + let resp_end = resp_off as usize + resp_cap as usize;
455 + let needed = req_end.max(resp_end);
456 + if file_size < needed {
457 + unsafe {
458 + libc::munmap(base as *mut libc::c_void, file_size);
459 + libc::close(fd);
460 + }
461 + return Err(ShmError::BadSize);
462 + }
463 +
464 + // Read current sequence numbers
465 + let cur_req_seq = atomic_load_u64(base, OFF_REQ_SEQ);
466 + let cur_resp_seq = atomic_load_u64(base, OFF_RESP_SEQ);
467 + let generation = unsafe { (*hdr).owner_generation };
468 +
469 + Ok(ShmContext {
470 + role: ShmRole::Client,
471 + fd,
472 + base,
473 + region_size: file_size,
474 + request_offset: req_off,
475 + request_capacity: req_cap,
476 + response_offset: resp_off,
477 + response_capacity: resp_cap,
478 + local_req_seq: cur_req_seq,
479 + local_resp_seq: cur_resp_seq,
480 + spin_tries: DEFAULT_SPIN_TRIES,
481 + owner_generation: generation,
482 + path,
483 + })
484 + }
485 +
486 + /// Publish a message (client sends request, server sends response).
487 + ///
488 + /// The message must include the 32-byte outer header + payload,
489 + /// exactly as sent over UDS.
490 + pub fn send(&mut self, msg: &[u8]) -> Result<(), ShmError> {
491 + if self.base.is_null() || msg.is_empty() {
492 + return Err(ShmError::BadParam("null context or empty message".into()));
493 + }
494 +
495 + let (area_offset, area_capacity, seq_off, len_off, sig_off) = match self.role {
496 + ShmRole::Client => (
497 + self.request_offset,
498 + self.request_capacity,
499 + OFF_REQ_SEQ,
500 + OFF_REQ_LEN,
501 + OFF_REQ_SIGNAL,
502 + ),
503 + ShmRole::Server => (
504 + self.response_offset,
505 + self.response_capacity,
506 + OFF_RESP_SEQ,
507 + OFF_RESP_LEN,
508 + OFF_RESP_SIGNAL,
509 + ),
510 + };
511 +
512 + if msg.len() > area_capacity as usize {
513 + return Err(ShmError::MsgTooLarge);
514 + }
515 +
516 + // 1. Write message data into the area
517 + unsafe {
518 + ptr::copy_nonoverlapping(msg.as_ptr(), self.base.add(area_offset as usize), msg.len());
519 + }
520 +
521 + // 2. Store message length (release)
522 + atomic_store_u32(self.base, len_off, msg.len() as u32);
523 +
524 + // 3. Increment sequence number (release) to publish
525 + atomic_add_u64(self.base, seq_off, 1);
526 +
527 + // 4. Wake the peer via futex
528 + atomic_add_u32(self.base, sig_off, 1);
529 + futex_wake(unsafe { self.base.add(sig_off) as *mut u32 }, 1);
530 +
531 + // Track locally
532 + match self.role {
533 + ShmRole::Client => self.local_req_seq += 1,
534 + ShmRole::Server => self.local_resp_seq += 1,
535 + }
536 +
537 + Ok(())
538 + }
539 +
540 + /// Receive a message into the caller-provided buffer.
541 + ///
542 + /// On success, returns the number of bytes written to `buf`.
543 + /// Returns `MsgTooLarge` if the message exceeds `buf.len()`.
544 + pub fn receive(&mut self, buf: &mut [u8], timeout_ms: u32) -> Result<usize, ShmError> {
545 + if self.base.is_null() {
546 + return Err(ShmError::BadParam("null context".into()));
547 + }
548 + if buf.is_empty() {
549 + return Err(ShmError::BadParam("empty buffer".into()));
550 + }
551 +
552 + let (area_offset, area_capacity, seq_off, len_off, sig_off, expected_seq) = match self.role
553 + {
554 + ShmRole::Server => (
555 + self.request_offset,
556 + self.request_capacity,
557 + OFF_REQ_SEQ,
558 + OFF_REQ_LEN,
559 + OFF_REQ_SIGNAL,
560 + self.local_req_seq + 1,
561 + ),
562 + ShmRole::Client => (
563 + self.response_offset,
564 + self.response_capacity,
565 + OFF_RESP_SEQ,
566 + OFF_RESP_LEN,
567 + OFF_RESP_SIGNAL,
568 + self.local_resp_seq + 1,
569 + ),
570 + };
571 +
572 + let max_copy = buf.len().min(area_capacity as usize);
573 +
574 + // Phase 1: spin. Copy immediately on observing the advance.
575 + let mut observed = false;
576 + let mut mlen = 0usize;
577 + for _ in 0..self.spin_tries {
578 + let cur = atomic_load_u64(self.base, seq_off);
579 + if cur >= expected_seq {
580 + mlen = atomic_load_u32(self.base, len_off) as usize;
581 + if mlen > 0 && mlen <= max_copy {
582 + unsafe {
583 + ptr::copy_nonoverlapping(
584 + self.base.add(area_offset as usize),
585 + buf.as_mut_ptr(),
586 + mlen,
587 + );
588 + }
589 + }
590 + observed = true;
591 + break;
592 + }
593 + cpu_relax();
594 + }
595 +
596 + // Phase 2: futex wait with deadline-based retry loop.
597 + //
598 + // Handles spurious wakeups (EAGAIN when signal word changed
599 + // between read and syscall, or EINTR from signal delivery).
600 + // Computes a wall-clock deadline so total wait never exceeds
601 + // timeout_ms regardless of retries.
602 + if !observed {
603 + let deadline_ns: u64 = if timeout_ms > 0 {
604 + let mut ts = libc::timespec {
605 + tv_sec: 0,
606 + tv_nsec: 0,
607 + };
608 + unsafe { libc::clock_gettime(libc::CLOCK_MONOTONIC, &mut ts) };
609 + ts.tv_sec as u64 * 1_000_000_000 + ts.tv_nsec as u64 + timeout_ms as u64 * 1_000_000
610 + } else {
611 + 0
612 + };
613 +
614 + loop {
615 + let sig_val = atomic_load_u32(self.base, sig_off);
616 +
617 + let cur = atomic_load_u64(self.base, seq_off);
618 + if cur >= expected_seq {
619 + break; // response arrived
620 + }
621 +
622 + // Compute remaining timeout for this futex_wait call
623 + let timeout = if deadline_ns > 0 {
624 + let mut now_ts = libc::timespec {
625 + tv_sec: 0,
626 + tv_nsec: 0,
627 + };
628 + unsafe { libc::clock_gettime(libc::CLOCK_MONOTONIC, &mut now_ts) };
629 + let now_val = now_ts.tv_sec as u64 * 1_000_000_000 + now_ts.tv_nsec as u64;
630 + if now_val >= deadline_ns {
631 + return Err(ShmError::Timeout);
632 + }
633 + let remain = deadline_ns - now_val;
634 + Some(libc::timespec {
635 + tv_sec: (remain / 1_000_000_000) as libc::time_t,
636 + tv_nsec: (remain % 1_000_000_000) as libc::c_long,
637 + })
638 + } else {
639 + None
640 + };
641 +
642 + let ret = futex_wait(
643 + unsafe { self.base.add(sig_off) as *mut u32 },
644 + sig_val,
645 + timeout.as_ref(),
646 + );
647 +
648 + if ret < 0 && errno() == libc::ETIMEDOUT {
649 + return Err(ShmError::Timeout);
650 + }
651 +
652 + // EAGAIN (value changed) or EINTR (signal): re-check seq
653 + }
654 +
655 + // Copy immediately after observing the sequence advance
656 + mlen = atomic_load_u32(self.base, len_off) as usize;
657 + if mlen > 0 && mlen <= max_copy {
658 + unsafe {
659 + ptr::copy_nonoverlapping(
660 + self.base.add(area_offset as usize),
661 + buf.as_mut_ptr(),
662 + mlen,
663 + );
664 + }
665 + }
666 + }
667 +
668 + // Advance local tracking (message is consumed from SHM perspective)
669 + match self.role {
670 + ShmRole::Server => self.local_req_seq = expected_seq,
671 + ShmRole::Client => self.local_resp_seq = expected_seq,
672 + }
673 +
674 + // mlen==0 after sequence advance indicates SHM corruption (send rejects 0-length)
675 + if mlen == 0 {
676 + return Err(ShmError::BadHeader);
677 + }
678 +
679 + // Message larger than caller buffer or area capacity
680 + if mlen > max_copy {
681 + return Err(ShmError::MsgTooLarge);
682 + }
683 +
684 + Ok(mlen)
685 + }
686 +
687 + /// Close client (munmap, close fd, no unlink).
688 + pub fn close(&mut self) {
689 + if !self.base.is_null() {
690 + unsafe { libc::munmap(self.base as *mut libc::c_void, self.region_size) };
691 + self.base = ptr::null_mut();
692 + }
693 + if self.fd >= 0 {
694 + unsafe { libc::close(self.fd) };
695 + self.fd = -1;
696 + }
697 + self.region_size = 0;
698 + }
699 +
700 + /// Destroy server (munmap, close, unlink).
701 + pub fn destroy(&mut self) {
702 + if !self.base.is_null() {
703 + unsafe { libc::munmap(self.base as *mut libc::c_void, self.region_size) };
704 + self.base = ptr::null_mut();
705 + }
706 + if self.fd >= 0 {
707 + unsafe { libc::close(self.fd) };
708 + self.fd = -1;
709 + }
710 + if !self.path.as_os_str().is_empty() {
711 + if let Ok(c) = std::ffi::CString::new(self.path.to_string_lossy().as_bytes()) {
712 + unsafe { libc::unlink(c.as_ptr()) };
713 + }
714 + self.path = PathBuf::new();
715 + }
716 + self.region_size = 0;
717 + }
718 +}
719 +
720 +impl Drop for ShmContext {
721 + fn drop(&mut self) {
722 + match self.role {
723 + ShmRole::Server => self.destroy(),
724 + ShmRole::Client => self.close(),
725 + }
726 + }
727 +}
728 +
729 +// ---------------------------------------------------------------------------
730 +// Stale session cleanup
731 +// ---------------------------------------------------------------------------
732 +
733 +/// Scan `run_dir` for files matching `{service_name}-*.ipcshm`, check
734 +/// owner_pid liveness for each, and unlink stale ones.
735 +pub fn cleanup_stale(run_dir: &str, service_name: &str) {
736 + let prefix = format!("{service_name}-");
737 + let suffix = ".ipcshm";
738 +
739 + let entries = match std::fs::read_dir(run_dir) {
740 + Ok(e) => e,
741 + Err(_) => return,
742 + };
743 +
744 + for entry in entries.flatten() {
745 + let name = match entry.file_name().into_string() {
746 + Ok(n) => n,
747 + Err(_) => continue,
748 + };
749 +
750 + if !name.starts_with(&prefix) || !name.ends_with(suffix) {
751 + continue;
752 + }
753 +
754 + let path = entry.path();
755 + let c_path = match path_to_cstring(&path) {
756 + Ok(c) => c,
757 + Err(_) => continue,
758 + };
759 +
760 + // Open read-only to inspect the header
761 + let fd = unsafe { libc::open(c_path.as_ptr(), libc::O_RDONLY) };
762 + if fd < 0 {
763 + // The entry vanished after readdir() (for example, a dangling symlink
764 + // target disappeared) — remove the stale directory entry. Any other
765 + // open failure is ambiguous, so leave the entry alone.
766 + if should_unlink_cleanup_open_failure(errno()) {
767 + unsafe { libc::unlink(c_path.as_ptr()) };
768 + }
769 + continue;
770 + }
771 +
772 + let mut st: libc::stat = unsafe { std::mem::zeroed() };
773 + if unsafe { libc::fstat(fd, &mut st) } != 0 || (st.st_size as usize) < HEADER_LEN as usize {
774 + unsafe {
775 + libc::close(fd);
776 + libc::unlink(c_path.as_ptr());
777 + }
778 + continue;
779 + }
780 +
781 + let map = unsafe {
782 + libc::mmap(
783 + ptr::null_mut(),
784 + HEADER_LEN as usize,
785 + libc::PROT_READ,
786 + libc::MAP_SHARED,
787 + fd,
788 + 0,
789 + )
790 + };
791 + unsafe { libc::close(fd) };
792 +
793 + if map == libc::MAP_FAILED {
794 + unsafe { libc::unlink(c_path.as_ptr()) };
795 + continue;
796 + }
797 +
798 + let hdr = map as *const RegionHeader;
799 + let magic = unsafe { (*hdr).magic };
800 + if magic != REGION_MAGIC {
801 + unsafe {
802 + libc::munmap(map, HEADER_LEN as usize);
803 + libc::unlink(c_path.as_ptr());
804 + }
805 + continue;
806 + }
807 +
808 + let owner = unsafe { (*hdr).owner_pid };
809 + let gen = unsafe { (*hdr).owner_generation };
810 + unsafe { libc::munmap(map, HEADER_LEN as usize) };
811 +
812 + // If owner is dead (or generation is zero / legacy), unlink
813 + if !pid_alive(owner) || gen == 0 {
814 + unsafe { libc::unlink(c_path.as_ptr()) };
815 + }
816 + }
817 +}
818 +
819 +// ---------------------------------------------------------------------------
820 +// Internal helpers
821 +// ---------------------------------------------------------------------------
822 +
823 +/// Round v up to REGION_ALIGNMENT. Returns None if the rounded value would
824 +/// overflow u32 — callers must reject such inputs rather than silently wrap.
825 +fn align64(v: u32) -> Option<u32> {
826 + v.checked_add(REGION_ALIGNMENT - 1)
827 + .map(|x| x & !(REGION_ALIGNMENT - 1))
828 +}
829 +
830 +/// Validate service_name: only [a-zA-Z0-9._-], non-empty, not "." or "..".
831 +fn validate_service_name(name: &str) -> Result<(), ShmError> {
832 + if name.is_empty() {
833 + return Err(ShmError::BadParam("empty service name".into()));
834 + }
835 + if name == "." || name == ".." {
836 + return Err(ShmError::BadParam(
837 + "service name cannot be '.' or '..'".into(),
838 + ));
839 + }
840 + for c in name.bytes() {
841 + match c {
842 + b'a'..=b'z' | b'A'..=b'Z' | b'0'..=b'9' | b'.' | b'_' | b'-' => {}
843 + _ => {
844 + return Err(ShmError::BadParam(format!(
845 + "service name contains invalid character: {:?}",
846 + c as char
847 + )))
848 + }
849 + }
850 + }
851 + Ok(())
852 +}
853 +
854 +fn build_shm_path(run_dir: &str, service_name: &str, session_id: u64) -> Result<PathBuf, ShmError> {
855 + validate_service_name(service_name)?;
856 + let path = Path::new(run_dir).join(format!("{service_name}-{session_id:016x}.ipcshm"));
857 + if path.to_string_lossy().len() >= 256 {
858 + return Err(ShmError::PathTooLong);
859 + }
860 + Ok(path)
861 +}
862 +
863 +fn path_to_cstring(path: &Path) -> Result<std::ffi::CString, ShmError> {
864 + std::ffi::CString::new(path.to_string_lossy().as_bytes())
865 + .map_err(|_| ShmError::BadParam("path contains null byte".into()))
866 +}
867 +
868 +fn errno() -> i32 {
869 + unsafe { *libc::__errno_location() }
870 +}
871 +
872 +fn pid_alive(pid: i32) -> bool {
873 + if pid <= 0 {
874 + return false;
875 + }
876 + let ret = unsafe { libc::kill(pid, 0) };
877 + ret == 0 || errno() == libc::EPERM
878 +}
879 +
880 +#[inline]
881 +fn cpu_relax() {
882 + #[cfg(any(target_arch = "x86_64", target_arch = "x86"))]
883 + unsafe {
884 + std::arch::asm!("pause", options(nomem, nostack));
885 + }
886 + #[cfg(target_arch = "aarch64")]
887 + unsafe {
888 + std::arch::asm!("yield", options(nomem, nostack));
889 + }
890 + #[cfg(not(any(target_arch = "x86_64", target_arch = "x86", target_arch = "aarch64")))]
891 + {
892 + std::sync::atomic::fence(std::sync::atomic::Ordering::SeqCst);
893 + }
894 +}
895 +
896 +// Atomic helpers: operate on raw pointers into the mmap'd region.
897 +// These match the C implementation's __atomic builtins.
898 +
899 +fn atomic_load_u64(base: *mut u8, offset: usize) -> u64 {
900 + let ptr = unsafe { base.add(offset) as *const std::sync::atomic::AtomicU64 };
901 + unsafe { (*ptr).load(std::sync::atomic::Ordering::Acquire) }
902 +}
903 +
904 +fn atomic_load_u32(base: *mut u8, offset: usize) -> u32 {
905 + let ptr = unsafe { base.add(offset) as *const std::sync::atomic::AtomicU32 };
906 + unsafe { (*ptr).load(std::sync::atomic::Ordering::Acquire) }
907 +}
908 +
909 +fn atomic_store_u32(base: *mut u8, offset: usize, val: u32) {
910 + let ptr = unsafe { base.add(offset) as *const std::sync::atomic::AtomicU32 };
911 + unsafe { (*ptr).store(val, std::sync::atomic::Ordering::Release) };
912 +}
913 +
914 +fn atomic_add_u64(base: *mut u8, offset: usize, val: u64) {
915 + let ptr = unsafe { base.add(offset) as *const std::sync::atomic::AtomicU64 };
916 + unsafe { (*ptr).fetch_add(val, std::sync::atomic::Ordering::Release) };
917 +}
918 +
919 +fn atomic_add_u32(base: *mut u8, offset: usize, val: u32) {
920 + let ptr = unsafe { base.add(offset) as *const std::sync::atomic::AtomicU32 };
921 + unsafe { (*ptr).fetch_add(val, std::sync::atomic::Ordering::Release) };
922 +}
923 +
924 +fn futex_wake(addr: *mut u32, count: i32) -> i32 {
925 + unsafe {
926 + libc::syscall(
927 + libc::SYS_futex,
928 + addr,
929 + FUTEX_WAKE,
930 + count,
931 + ptr::null::<libc::timespec>(),
932 + ptr::null::<u32>(),
933 + 0i32,
934 + ) as i32
935 + }
936 +}
937 +
938 +fn futex_wait(addr: *mut u32, expected: u32, timeout: Option<&libc::timespec>) -> i32 {
939 + let tsp = match timeout {
940 + Some(ts) => ts as *const libc::timespec,
941 + None => ptr::null(),
942 + };
943 + unsafe {
944 + libc::syscall(
945 + libc::SYS_futex,
946 + addr,
947 + FUTEX_WAIT,
948 + expected,
949 + tsp,
950 + ptr::null::<u32>(),
951 + 0i32,
952 + ) as i32
953 + }
954 +}
955 +
956 +// ---------------------------------------------------------------------------
957 +// Stale region recovery
958 +// ---------------------------------------------------------------------------
959 +
960 +#[derive(PartialEq, Eq)]
961 +#[allow(dead_code)]
962 +enum StaleResult {
963 + NotExist,
964 + Recovered,
965 + LiveServer,
966 + Invalid,
967 +}
968 +
969 +fn should_unlink_cleanup_open_failure(err: i32) -> bool {
970 + err == libc::ENOENT
971 +}
972 +
973 +fn classify_stale_open_failure(err: i32) -> StaleResult {
974 + if err == libc::ENOENT {
975 + StaleResult::NotExist
976 + } else {
977 + StaleResult::Invalid
978 + }
979 +}
980 +
981 +#[allow(dead_code)]
982 +fn check_shm_stale(path: &Path) -> StaleResult {
983 + let c_path = match path_to_cstring(path) {
984 + Ok(c) => c,
985 + Err(_) => return StaleResult::NotExist,
986 + };
987 +
988 + let mut st: libc::stat = unsafe { std::mem::zeroed() };
989 + if unsafe { libc::stat(c_path.as_ptr(), &mut st) } != 0 {
990 + return StaleResult::NotExist;
991 + }
992 +
993 + if (st.st_size as usize) < HEADER_LEN as usize {
994 + unsafe { libc::unlink(c_path.as_ptr()) };
995 + return StaleResult::Invalid;
996 + }
997 +
998 + let fd = unsafe { libc::open(c_path.as_ptr(), libc::O_RDONLY) };
999 + if fd < 0 {
1000 + return classify_stale_open_failure(errno());
1001 + }
1002 +
1003 + let map = unsafe {
1004 + libc::mmap(
1005 + ptr::null_mut(),
1006 + HEADER_LEN as usize,
1007 + libc::PROT_READ,
1008 + libc::MAP_SHARED,
1009 + fd,
1010 + 0,
1011 + )
1012 + };
1013 + unsafe { libc::close(fd) };
1014 +
1015 + if map == libc::MAP_FAILED {
1016 + unsafe { libc::unlink(c_path.as_ptr()) };
1017 + return StaleResult::Invalid;
1018 + }
1019 +
1020 + let hdr = map as *const RegionHeader;
1021 + let magic = unsafe { (*hdr).magic };
1022 + if magic != REGION_MAGIC {
1023 + unsafe {
1024 + libc::munmap(map, HEADER_LEN as usize);
1025 + libc::unlink(c_path.as_ptr());
1026 + }
1027 + return StaleResult::Invalid;
1028 + }
1029 +
1030 + let owner = unsafe { (*hdr).owner_pid };
1031 + let gen = unsafe { (*hdr).owner_generation };
1032 + unsafe { libc::munmap(map, HEADER_LEN as usize) };
1033 +
1034 + if pid_alive(owner) && gen != 0 {
1035 + return StaleResult::LiveServer;
1036 + }
1037 +
1038 + // Dead owner or zero generation (PID reuse / legacy) — stale
1039 + unsafe { libc::unlink(c_path.as_ptr()) };
1040 + StaleResult::Recovered
1041 +}
1042 +
1043 +// ---------------------------------------------------------------------------
1044 +// Tests
1045 +// ---------------------------------------------------------------------------
1046 +
1047 +#[cfg(test)]
1048 +#[path = "shm_tests.rs"]
1049 +mod tests;
src/crates/netipc/src/transport/shm_tests.rs new
+1446
@@ -0,0 +1,1446 @@
1 +use super::*;
2 +use crate::protocol;
3 +use std::ffi::OsString;
4 +use std::os::unix::ffi::OsStringExt;
5 +use std::path::PathBuf;
6 +use std::sync::mpsc;
7 +use std::thread;
8 +use std::time::Duration;
9 +
10 +const TEST_RUN_DIR: &str = "/tmp/nipc_shm_rust_test";
11 +const SERVER_READY_TIMEOUT: Duration = Duration::from_secs(2);
12 +
13 +fn ensure_run_dir() {
14 + let _ = std::fs::create_dir_all(TEST_RUN_DIR);
15 +}
16 +
17 +fn cleanup_shm(service: &str, session_id: u64) {
18 + let path = format!("{TEST_RUN_DIR}/{service}-{session_id:016x}.ipcshm");
19 + let _ = std::fs::remove_file(&path);
20 +}
21 +
22 +fn wait_for_server_ready(rx: &mpsc::Receiver<u64>, service: &str) -> u64 {
23 + rx.recv_timeout(SERVER_READY_TIMEOUT).unwrap_or_else(|err| {
24 + panic!("timed out waiting for server readiness for service={service}: {err}")
25 + })
26 +}
27 +
28 +/// Build a complete message (outer header + payload) for SHM.
29 +fn build_message(kind: u16, code: u16, message_id: u64, payload: &[u8]) -> Vec<u8> {
30 + let hdr = protocol::Header {
31 + magic: protocol::MAGIC_MSG,
32 + version: protocol::VERSION,
33 + header_len: protocol::HEADER_LEN,
34 + kind,
35 + code,
36 + flags: 0,
37 + transport_status: protocol::STATUS_OK,
38 + payload_len: payload.len() as u32,
39 + item_count: 1,
40 + message_id,
41 + };
42 + let mut buf = vec![0u8; protocol::HEADER_SIZE + payload.len()];
43 + hdr.encode(&mut buf[..protocol::HEADER_SIZE]);
44 + buf[protocol::HEADER_SIZE..].copy_from_slice(payload);
45 + buf
46 +}
47 +
48 +#[test]
49 +fn test_direct_roundtrip() {
50 + ensure_run_dir();
51 + let svc = "rs_shm_rt";
52 + let sid: u64 = 1;
53 + cleanup_shm(svc, sid);
54 +
55 + let (ready_tx, ready_rx) = mpsc::channel();
56 + let svc_clone = svc.to_string();
57 + let server_thread = thread::spawn(move || {
58 + let mut ctx = ShmContext::server_create(TEST_RUN_DIR, &svc_clone, sid, 4096, 4096)
59 + .expect("server create");
60 + ready_tx.send(sid).expect("server ready signal");
61 +
62 + // Receive request
63 + let mut buf = vec![0u8; 65536];
64 + let mlen = ctx.receive(&mut buf, 5000).expect("server receive");
65 + let msg = &buf[..mlen];
66 + assert!(msg.len() >= protocol::HEADER_SIZE);
67 +
68 + // Parse and echo as response
69 + let hdr = protocol::Header::decode(msg).expect("decode");
70 + let payload = msg[protocol::HEADER_SIZE..].to_vec();
71 + let resp = build_message(protocol::KIND_RESPONSE, hdr.code, hdr.message_id, &payload);
72 + ctx.send(&resp).expect("server send");
73 + ctx.destroy();
74 + });
75 +
76 + let ready_sid = wait_for_server_ready(&ready_rx, svc);
77 + assert_eq!(ready_sid, sid, "unexpected server ready sid");
78 + let mut client = ShmContext::client_attach(TEST_RUN_DIR, svc, sid).expect("client attach");
79 +
80 + let payload = vec![0xCA, 0xFE, 0xBA, 0xBE];
81 + let msg = build_message(
82 + protocol::KIND_REQUEST,
83 + protocol::METHOD_INCREMENT,
84 + 42,
85 + &payload,
86 + );
87 + client.send(&msg).expect("client send");
88 +
89 + let mut resp_buf = vec![0u8; 65536];
90 + let rlen = client.receive(&mut resp_buf, 5000).expect("client receive");
91 + let resp = &resp_buf[..rlen];
92 + assert_eq!(resp.len(), protocol::HEADER_SIZE + payload.len());
93 +
94 + let rhdr = protocol::Header::decode(resp).expect("decode response");
95 + assert_eq!(rhdr.kind, protocol::KIND_RESPONSE);
96 + assert_eq!(rhdr.message_id, 42);
97 + assert_eq!(&resp[protocol::HEADER_SIZE..], &payload[..]);
98 +
99 + client.close();
100 + server_thread.join().unwrap();
101 + cleanup_shm(svc, sid);
102 +}
103 +
104 +#[test]
105 +fn test_multiple_roundtrips() {
106 + ensure_run_dir();
107 + let svc = "rs_shm_multi";
108 + let sid: u64 = 2;
109 + cleanup_shm(svc, sid);
110 +
111 + let (ready_tx, ready_rx) = mpsc::channel();
112 + let svc_clone = svc.to_string();
113 + let server_thread = thread::spawn(move || {
114 + let mut ctx = ShmContext::server_create(TEST_RUN_DIR, &svc_clone, sid, 4096, 4096)
115 + .expect("server create");
116 + ready_tx.send(sid).expect("server ready signal");
117 +
118 + let mut buf = vec![0u8; 65536];
119 + for _ in 0..10 {
120 + let mlen = ctx.receive(&mut buf, 5000).expect("server receive");
121 + let msg = &buf[..mlen];
122 + let hdr = protocol::Header::decode(msg).expect("decode");
123 + let payload = msg[protocol::HEADER_SIZE..].to_vec();
124 + let resp = build_message(protocol::KIND_RESPONSE, hdr.code, hdr.message_id, &payload);
125 + ctx.send(&resp).expect("server send");
126 + }
127 + ctx.destroy();
128 + });
129 +
130 + let ready_sid = wait_for_server_ready(&ready_rx, svc);
131 + assert_eq!(ready_sid, sid, "unexpected server ready sid");
132 + let mut client = ShmContext::client_attach(TEST_RUN_DIR, svc, sid).expect("client attach");
133 +
134 + let mut resp_buf = vec![0u8; 65536];
135 + for i in 0u64..10 {
136 + let payload = vec![i as u8];
137 + let msg = build_message(protocol::KIND_REQUEST, 1, i + 1, &payload);
138 + client.send(&msg).expect("client send");
139 + let rlen = client.receive(&mut resp_buf, 5000).expect("client receive");
140 + let resp = &resp_buf[..rlen];
141 + let rhdr = protocol::Header::decode(resp).expect("decode");
142 + assert_eq!(rhdr.kind, protocol::KIND_RESPONSE);
143 + assert_eq!(rhdr.message_id, i + 1);
144 + assert_eq!(resp[protocol::HEADER_SIZE], i as u8);
145 + }
146 +
147 + client.close();
148 + server_thread.join().unwrap();
149 + cleanup_shm(svc, sid);
150 +}
151 +
152 +#[test]
153 +fn test_stale_recovery() {
154 + ensure_run_dir();
155 + let svc = "rs_shm_stale";
156 + let sid: u64 = 3;
157 + cleanup_shm(svc, sid);
158 +
159 + // Create a region, corrupt owner_pid to simulate dead process
160 + let mut first =
161 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("first create");
162 + let hdr = first.base as *mut RegionHeader;
163 + unsafe { (*hdr).owner_pid = 99999 }; // very unlikely alive
164 + first.close(); // close without unlink
165 +
166 + // cleanup_stale should remove the stale file
167 + cleanup_stale(TEST_RUN_DIR, svc);
168 +
169 + // Now O_EXCL create should succeed
170 + let mut second = ShmContext::server_create(TEST_RUN_DIR, svc, sid, 2048, 2048)
171 + .expect("stale recovery create");
172 + assert!(second.request_capacity >= 2048);
173 + second.destroy();
174 + cleanup_shm(svc, sid);
175 +}
176 +
177 +#[test]
178 +fn test_large_message() {
179 + ensure_run_dir();
180 + let svc = "rs_shm_large";
181 + let sid: u64 = 4;
182 + cleanup_shm(svc, sid);
183 +
184 + let (ready_tx, ready_rx) = mpsc::channel();
185 + let svc_clone = svc.to_string();
186 + let server_thread = thread::spawn(move || {
187 + let mut ctx = ShmContext::server_create(TEST_RUN_DIR, &svc_clone, sid, 65536, 65536)
188 + .expect("server create");
189 + ready_tx.send(sid).expect("server ready signal");
190 + let mut buf = vec![0u8; 65536];
191 + let mlen = ctx.receive(&mut buf, 5000).expect("server receive");
192 + let msg = &buf[..mlen];
193 + let hdr = protocol::Header::decode(msg).expect("decode");
194 + let payload = msg[protocol::HEADER_SIZE..].to_vec();
195 + let resp = build_message(protocol::KIND_RESPONSE, hdr.code, hdr.message_id, &payload);
196 + ctx.send(&resp).expect("server send");
197 + ctx.destroy();
198 + });
199 +
200 + let ready_sid = wait_for_server_ready(&ready_rx, svc);
201 + assert_eq!(ready_sid, sid, "unexpected server ready sid");
202 + let mut client = ShmContext::client_attach(TEST_RUN_DIR, svc, sid).expect("client attach");
203 +
204 + // 60000 bytes of payload
205 + let payload: Vec<u8> = (0..60000).map(|i| (i & 0xFF) as u8).collect();
206 + let msg = build_message(protocol::KIND_REQUEST, 1, 999, &payload);
207 + client.send(&msg).expect("send large");
208 +
209 + let mut resp_buf = vec![0u8; 65536];
210 + let rlen = client.receive(&mut resp_buf, 5000).expect("receive large");
211 + let resp = &resp_buf[..rlen];
212 + assert_eq!(resp.len(), protocol::HEADER_SIZE + payload.len());
213 + assert_eq!(&resp[protocol::HEADER_SIZE..], &payload[..]);
214 +
215 + client.close();
216 + server_thread.join().unwrap();
217 + cleanup_shm(svc, sid);
218 +}
219 +
220 +#[test]
221 +fn test_shm_chaos_forged_length() {
222 + // Verify that forged/malicious req_len and resp_len values in the SHM
223 + // header are handled safely: no panic, no out-of-bounds read.
224 + ensure_run_dir();
225 + let svc = "rs_shm_forge";
226 + let sid: u64 = 100;
227 + cleanup_shm(svc, sid);
228 +
229 + let mut server_ctx =
230 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
231 +
232 + let mut client_ctx = ShmContext::client_attach(TEST_RUN_DIR, svc, sid).expect("client attach");
233 +
234 + let base = server_ctx.base;
235 + let req_cap = server_ctx.request_capacity as usize;
236 + let resp_cap = server_ctx.response_capacity as usize;
237 +
238 + // --- Test forged req_len (server-side receive) ---
239 +
240 + // Forged lengths: 0 → BadHeader (corruption), within capacity → Ok,
241 + // over capacity → MsgTooLarge.
242 + let forged_req_lengths: &[(u32, bool)] = &[
243 + // (0, ...) tested separately below as BadHeader
244 + (1, false), // tiny valid
245 + (req_cap as u32 - 1, false), // just under capacity
246 + (req_cap as u32, false), // exactly at capacity
247 + (req_cap as u32 + 1, true), // one over capacity
248 + (0x0001_0000, true), // moderately large
249 + (0x7FFF_FFFF, true), // large positive
250 + (0xFFFF_FFFF, true), // u32::MAX
251 + ];
252 +
253 + let mut recv_buf = vec![0u8; 65536];
254 +
255 + for &(forged_len, expect_too_large) in forged_req_lengths {
256 + // Write garbage into the request area
257 + unsafe {
258 + let req_area = base.add(server_ctx.request_offset as usize);
259 + std::ptr::write_bytes(req_area, 0xAB, req_cap);
260 + }
261 +
262 + // Store forged req_len at offset 48
263 + atomic_store_u32(base, OFF_REQ_LEN, forged_len);
264 +
265 + // Increment req_seq at offset 32 to signal a new "message"
266 + atomic_add_u64(base, OFF_REQ_SEQ, 1);
267 +
268 + // Bump req_signal at offset 56 to wake futex
269 + atomic_add_u32(base, OFF_REQ_SIGNAL, 1);
270 +
271 + let result = server_ctx.receive(&mut recv_buf, 100);
272 +
273 + if expect_too_large {
274 + assert_eq!(
275 + result.unwrap_err(),
276 + ShmError::MsgTooLarge,
277 + "forged req_len={forged_len:#x} should return MsgTooLarge"
278 + );
279 + } else {
280 + let n = result.unwrap_or_else(|e| {
281 + panic!("forged req_len={forged_len:#x} should succeed, got {e:?}")
282 + });
283 + assert_eq!(
284 + n, forged_len as usize,
285 + "returned length should match forged req_len={forged_len:#x}"
286 + );
287 + }
288 + }
289 +
290 + // --- Test forged resp_len (client-side receive) ---
291 +
292 + let forged_resp_lengths: &[(u32, bool)] = &[
293 + // (0, ...) tested separately below as BadHeader
294 + (1, false),
295 + (resp_cap as u32 - 1, false),
296 + (resp_cap as u32, false),
297 + (resp_cap as u32 + 1, true),
298 + (0x0001_0000, true),
299 + (0x7FFF_FFFF, true),
300 + (0xFFFF_FFFF, true),
301 + ];
302 +
303 + for &(forged_len, expect_too_large) in forged_resp_lengths {
304 + // Write garbage into the response area
305 + unsafe {
306 + let resp_area = base.add(server_ctx.response_offset as usize);
307 + std::ptr::write_bytes(resp_area, 0xCD, resp_cap);
308 + }
309 +
310 + // Store forged resp_len at offset 52
311 + atomic_store_u32(base, OFF_RESP_LEN, forged_len);
312 +
313 + // Increment resp_seq at offset 40
314 + atomic_add_u64(base, OFF_RESP_SEQ, 1);
315 +
316 + // Bump resp_signal at offset 60
317 + atomic_add_u32(base, OFF_RESP_SIGNAL, 1);
318 +
319 + let result = client_ctx.receive(&mut recv_buf, 100);
320 +
321 + if expect_too_large {
322 + assert_eq!(
323 + result.unwrap_err(),
324 + ShmError::MsgTooLarge,
325 + "forged resp_len={forged_len:#x} should return MsgTooLarge"
326 + );
327 + } else {
328 + let n = result.unwrap_or_else(|e| {
329 + panic!("forged resp_len={forged_len:#x} should succeed, got {e:?}")
330 + });
331 + assert_eq!(
332 + n, forged_len as usize,
333 + "returned length should match forged resp_len={forged_len:#x}"
334 + );
335 + }
336 + }
337 +
338 + // --- mlen==0 → BadHeader (corruption indicator) ---
339 + // Server-side: forge req_len=0
340 + atomic_store_u32(base, OFF_REQ_LEN, 0);
341 + atomic_add_u64(base, OFF_REQ_SEQ, 1);
342 + atomic_add_u32(base, OFF_REQ_SIGNAL, 1);
343 + assert_eq!(
344 + server_ctx.receive(&mut recv_buf, 100).unwrap_err(),
345 + ShmError::BadHeader,
346 + "forged req_len=0 should return BadHeader"
347 + );
348 +
349 + // Client-side: forge resp_len=0
350 + atomic_store_u32(base, OFF_RESP_LEN, 0);
351 + atomic_add_u64(base, OFF_RESP_SEQ, 1);
352 + atomic_add_u32(base, OFF_RESP_SIGNAL, 1);
353 + assert_eq!(
354 + client_ctx.receive(&mut recv_buf, 100).unwrap_err(),
355 + ShmError::BadHeader,
356 + "forged resp_len=0 should return BadHeader"
357 + );
358 +
359 + client_ctx.close();
360 + server_ctx.destroy();
361 + cleanup_shm(svc, sid);
362 +}
363 +
364 +#[test]
365 +fn test_shm_multi_client() {
366 + // Verify multiple clients with independent SHM regions have no
367 + // cross-contamination: each client/server pair exchanges unique data.
368 + ensure_run_dir();
369 + let svc = "rs_shm_mc";
370 + let session_ids: [u64; 3] = [1, 2, 3];
371 +
372 + for &sid in &session_ids {
373 + cleanup_shm(svc, sid);
374 + }
375 +
376 + let svc_name = svc.to_string();
377 + let (ready_tx, ready_rx) = mpsc::channel();
378 +
379 + // Spawn 3 server threads, one per session
380 + let server_handles: Vec<_> = session_ids
381 + .iter()
382 + .map(|&sid| {
383 + let svc_clone = svc_name.clone();
384 + let ready_tx = ready_tx.clone();
385 + thread::spawn(move || {
386 + let mut ctx = ShmContext::server_create(TEST_RUN_DIR, &svc_clone, sid, 4096, 4096)
387 + .expect(&format!("server create sid={sid}"));
388 + ready_tx.send(sid).expect("server ready signal");
389 +
390 + // Receive request
391 + let mut buf = vec![0u8; 8192];
392 + let mlen = ctx
393 + .receive(&mut buf, 5000)
394 + .expect(&format!("server receive sid={sid}"));
395 + let msg = &buf[..mlen];
396 + let hdr = protocol::Header::decode(msg).expect("decode req");
397 + let req_payload = msg[protocol::HEADER_SIZE..].to_vec();
398 +
399 + // Echo back as response with same message_id
400 + let resp = build_message(
401 + protocol::KIND_RESPONSE,
402 + hdr.code,
403 + hdr.message_id,
404 + &req_payload,
405 + );
406 + ctx.send(&resp).expect(&format!("server send sid={sid}"));
407 + ctx.destroy();
408 + })
409 + })
410 + .collect();
411 +
412 + drop(ready_tx);
413 + for _ in &session_ids {
414 + let ready_sid = wait_for_server_ready(&ready_rx, svc);
415 + assert!(
416 + session_ids.contains(&ready_sid),
417 + "unexpected server ready sid={ready_sid}"
418 + );
419 + }
420 +
421 + // Attach 3 clients, send unique payloads, verify isolation
422 + let client_handles: Vec<_> = session_ids
423 + .iter()
424 + .map(|&sid| {
425 + let svc_clone = svc_name.clone();
426 + thread::spawn(move || {
427 + let mut client = ShmContext::client_attach(TEST_RUN_DIR, &svc_clone, sid)
428 + .expect(&format!("client attach sid={sid}"));
429 +
430 + // Unique payload: session_id repeated to fill a pattern
431 + let unique_byte = sid as u8;
432 + let payload: Vec<u8> = vec![unique_byte; 64];
433 + let msg_id = 1000 + sid;
434 +
435 + let msg = build_message(
436 + protocol::KIND_REQUEST,
437 + protocol::METHOD_INCREMENT,
438 + msg_id,
439 + &payload,
440 + );
441 + client.send(&msg).expect(&format!("client send sid={sid}"));
442 +
443 + // Receive response
444 + let mut resp_buf = vec![0u8; 8192];
445 + let rlen = client
446 + .receive(&mut resp_buf, 5000)
447 + .expect(&format!("client receive sid={sid}"));
448 + let resp = &resp_buf[..rlen];
449 +
450 + let rhdr = protocol::Header::decode(resp).expect("decode resp");
451 + assert_eq!(rhdr.kind, protocol::KIND_RESPONSE);
452 + assert_eq!(
453 + rhdr.message_id, msg_id,
454 + "sid={sid}: message_id mismatch (cross-contamination?)"
455 + );
456 +
457 + // Verify the payload matches what we sent
458 + let resp_payload = &resp[protocol::HEADER_SIZE..];
459 + assert_eq!(
460 + resp_payload,
461 + &payload[..],
462 + "sid={sid}: payload mismatch (cross-contamination?)"
463 + );
464 +
465 + // Verify every byte is our unique marker
466 + for (i, &b) in resp_payload.iter().enumerate() {
467 + assert_eq!(
468 + b, unique_byte,
469 + "sid={sid}: byte {i} is {b:#x}, expected {unique_byte:#x}"
470 + );
471 + }
472 +
473 + client.close();
474 + })
475 + })
476 + .collect();
477 +
478 + for h in client_handles {
479 + h.join().unwrap();
480 + }
481 + for h in server_handles {
482 + h.join().unwrap();
483 + }
484 +
485 + for &sid in &session_ids {
486 + cleanup_shm(svc, sid);
487 + }
488 +}
489 +
490 +// -----------------------------------------------------------------------
491 +// ShmError Display coverage (lines 73-88)
492 +// -----------------------------------------------------------------------
493 +
494 +#[test]
495 +fn shm_error_display_all_variants() {
496 + let cases: Vec<(ShmError, &str)> = vec![
497 + (ShmError::PathTooLong, "SHM path exceeds limit"),
498 + (ShmError::Open(2), "open failed: errno 2"),
499 + (ShmError::Truncate(28), "ftruncate failed: errno 28"),
500 + (ShmError::Mmap(12), "mmap failed: errno 12"),
501 + (ShmError::BadMagic, "SHM header magic mismatch"),
502 + (ShmError::BadVersion, "SHM header version mismatch"),
503 + (ShmError::BadHeader, "SHM header_len mismatch"),
504 + (ShmError::BadSize, "SHM file too small for declared areas"),
505 + (ShmError::AddrInUse, "SHM region owned by live server"),
506 + (ShmError::NotReady, "SHM server not ready"),
507 + (ShmError::MsgTooLarge, "message exceeds SHM area capacity"),
508 + (ShmError::Timeout, "SHM futex wait timed out"),
509 + (ShmError::BadParam("test".into()), "bad parameter: test"),
510 + (ShmError::PeerDead, "SHM owner process has exited"),
511 + ];
512 + for (err, expected) in cases {
513 + assert_eq!(format!("{}", err), expected);
514 + }
515 + let e: &dyn std::error::Error = &ShmError::PathTooLong;
516 + let _ = format!("{e}");
517 +}
518 +
519 +// -----------------------------------------------------------------------
520 +// ShmContext accessors: role(), fd() (lines 164-165, 169-170)
521 +// -----------------------------------------------------------------------
522 +
523 +#[test]
524 +fn test_shm_role_and_fd() {
525 + ensure_run_dir();
526 + let svc = "rs_shm_role";
527 + let sid: u64 = 50;
528 + cleanup_shm(svc, sid);
529 +
530 + let server =
531 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
532 + assert_eq!(server.role(), ShmRole::Server);
533 + assert!(server.fd() >= 0);
534 +
535 + let client = ShmContext::client_attach(TEST_RUN_DIR, svc, sid).expect("client attach");
536 + assert_eq!(client.role(), ShmRole::Client);
537 + assert!(client.fd() >= 0);
538 +
539 + // Cleanup
540 + let mut c = client;
541 + let mut s = server;
542 + c.close();
543 + s.destroy();
544 + cleanup_shm(svc, sid);
545 +}
546 +
547 +// -----------------------------------------------------------------------
548 +// ShmContext::owner_alive() (lines 174-191)
549 +// -----------------------------------------------------------------------
550 +
551 +#[test]
552 +fn test_owner_alive() {
553 + ensure_run_dir();
554 + let svc = "rs_shm_alive";
555 + let sid: u64 = 51;
556 + cleanup_shm(svc, sid);
557 +
558 + let server =
559 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
560 +
561 + // Owner is the current process, so should be alive
562 + assert!(server.owner_alive());
563 +
564 + // Client should also report alive (same process owns it)
565 + let client = ShmContext::client_attach(TEST_RUN_DIR, svc, sid).expect("client attach");
566 + assert!(client.owner_alive());
567 +
568 + let mut c = client;
569 + let mut s = server;
570 + c.close();
571 + s.destroy();
572 + cleanup_shm(svc, sid);
573 +}
574 +
575 +#[test]
576 +fn test_owner_alive_dead_pid() {
577 + ensure_run_dir();
578 + let svc = "rs_shm_dead";
579 + let sid: u64 = 52;
580 + cleanup_shm(svc, sid);
581 +
582 + let mut server =
583 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
584 +
585 + // Forge a dead PID
586 + let hdr = server.base as *mut RegionHeader;
587 + unsafe { (*hdr).owner_pid = 99999 };
588 +
589 + assert!(!server.owner_alive());
590 +
591 + server.destroy();
592 + cleanup_shm(svc, sid);
593 +}
594 +
595 +#[test]
596 +fn test_owner_alive_null_base() {
597 + ensure_run_dir();
598 + let svc = "rs_shm_null";
599 + let sid: u64 = 53;
600 + cleanup_shm(svc, sid);
601 +
602 + let mut server =
603 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
604 +
605 + // Null base -> not alive
606 + let saved_base = server.base;
607 + server.base = std::ptr::null_mut();
608 + assert!(!server.owner_alive());
609 +
610 + // Restore for cleanup
611 + server.base = saved_base;
612 + server.destroy();
613 + cleanup_shm(svc, sid);
614 +}
615 +
616 +#[test]
617 +fn test_owner_alive_generation_mismatch() {
618 + ensure_run_dir();
619 + let svc = "rs_shm_gen";
620 + let sid: u64 = 54;
621 + cleanup_shm(svc, sid);
622 +
623 + let mut server =
624 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
625 +
626 + // Forge a different generation in the header
627 + let hdr = server.base as *mut RegionHeader;
628 + unsafe { (*hdr).owner_generation = server.owner_generation + 1 };
629 +
630 + // Generation mismatch -> not alive
631 + assert!(!server.owner_alive());
632 +
633 + server.destroy();
634 + cleanup_shm(svc, sid);
635 +}
636 +
637 +#[test]
638 +fn test_owner_alive_legacy_generation_skips_generation_check() {
639 + ensure_run_dir();
640 + let svc = "rs_shm_gen_legacy";
641 + let sid: u64 = 55;
642 + cleanup_shm(svc, sid);
643 +
644 + let mut server =
645 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
646 +
647 + let hdr = server.base as *mut RegionHeader;
648 + unsafe { (*hdr).owner_generation = server.owner_generation.wrapping_add(123) };
649 + server.owner_generation = 0;
650 +
651 + assert!(
652 + server.owner_alive(),
653 + "cached generation 0 should skip owner-generation mismatch checks"
654 + );
655 +
656 + server.destroy();
657 + cleanup_shm(svc, sid);
658 +}
659 +
660 +// -----------------------------------------------------------------------
661 +// ShmContext::send() error paths (lines 436-437, 458)
662 +// -----------------------------------------------------------------------
663 +
664 +#[test]
665 +fn test_send_null_context() {
666 + ensure_run_dir();
667 + let svc = "rs_shm_sendnull";
668 + let sid: u64 = 55;
669 + cleanup_shm(svc, sid);
670 +
671 + let mut server =
672 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
673 +
674 + // Empty message -> error
675 + assert!(matches!(server.send(&[]), Err(ShmError::BadParam(_))));
676 +
677 + // Null base -> error
678 + let saved_base = server.base;
679 + server.base = std::ptr::null_mut();
680 + assert!(matches!(
681 + server.send(&[1, 2, 3]),
682 + Err(ShmError::BadParam(_))
683 + ));
684 + server.base = saved_base;
685 +
686 + server.destroy();
687 + cleanup_shm(svc, sid);
688 +}
689 +
690 +#[test]
691 +fn test_send_msg_too_large() {
692 + ensure_run_dir();
693 + let svc = "rs_shm_sendlrg";
694 + let sid: u64 = 56;
695 + cleanup_shm(svc, sid);
696 +
697 + // Small capacity
698 + let mut server =
699 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 64, 64).expect("server create");
700 +
701 + // Create a message bigger than response capacity
702 + let big = vec![0u8; 1024];
703 + assert_eq!(server.send(&big).unwrap_err(), ShmError::MsgTooLarge);
704 +
705 + server.destroy();
706 + cleanup_shm(svc, sid);
707 +}
708 +
709 +// -----------------------------------------------------------------------
710 +// ShmContext::receive() error paths (lines 494-498)
711 +// -----------------------------------------------------------------------
712 +
713 +#[test]
714 +fn test_receive_null_context() {
715 + ensure_run_dir();
716 + let svc = "rs_shm_recvnull";
717 + let sid: u64 = 57;
718 + cleanup_shm(svc, sid);
719 +
720 + let mut server =
721 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
722 +
723 + // Empty buffer -> error
724 + assert!(matches!(
725 + server.receive(&mut [], 100),
726 + Err(ShmError::BadParam(_))
727 + ));
728 +
729 + // Null base -> error
730 + let saved_base = server.base;
731 + server.base = std::ptr::null_mut();
732 + let mut buf = [0u8; 64];
733 + assert!(matches!(
734 + server.receive(&mut buf, 100),
735 + Err(ShmError::BadParam(_))
736 + ));
737 + server.base = saved_base;
738 +
739 + server.destroy();
740 + cleanup_shm(svc, sid);
741 +}
742 +
743 +// -----------------------------------------------------------------------
744 +// Timeout (line 575, 593)
745 +// -----------------------------------------------------------------------
746 +
747 +#[test]
748 +fn test_receive_timeout() {
749 + ensure_run_dir();
750 + let svc = "rs_shm_timeout";
751 + let sid: u64 = 58;
752 + cleanup_shm(svc, sid);
753 +
754 + let mut server =
755 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
756 +
757 + // No client sends anything, so receive must timeout
758 + let mut buf = [0u8; 1024];
759 + let result = server.receive(&mut buf, 50); // 50ms timeout
760 + assert_eq!(result.unwrap_err(), ShmError::Timeout);
761 +
762 + server.destroy();
763 + cleanup_shm(svc, sid);
764 +}
765 +
766 +#[test]
767 +fn test_receive_without_deadline_waits_for_message() {
768 + ensure_run_dir();
769 + let svc = "rs_shm_wait_forever";
770 + let sid: u64 = 580;
771 + cleanup_shm(svc, sid);
772 +
773 + let mut server =
774 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
775 + let mut client = ShmContext::client_attach(TEST_RUN_DIR, svc, sid).expect("client attach");
776 +
777 + let sender = thread::spawn({
778 + let svc_name = svc.to_string();
779 + move || {
780 + thread::sleep(std::time::Duration::from_millis(50));
781 + let mut ctx =
782 + ShmContext::client_attach(TEST_RUN_DIR, &svc_name, sid).expect("attach sender");
783 + let msg = build_message(protocol::KIND_REQUEST, 1, 9, &[0xAA, 0xBB]);
784 + ctx.send(&msg).expect("send");
785 + ctx.close();
786 + }
787 + });
788 +
789 + let mut buf = [0u8; 128];
790 + let n = server.receive(&mut buf, 0).expect("receive with timeout 0");
791 + assert_eq!(n, protocol::HEADER_SIZE + 2);
792 +
793 + sender.join().unwrap();
794 + client.close();
795 + server.destroy();
796 + cleanup_shm(svc, sid);
797 +}
798 +
799 +#[test]
800 +fn test_receive_with_timeout_wakes_for_message() {
801 + ensure_run_dir();
802 + let svc = "rs_shm_wait_budget";
803 + let sid: u64 = 581;
804 + cleanup_shm(svc, sid);
805 +
806 + let mut server =
807 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
808 + let mut client = ShmContext::client_attach(TEST_RUN_DIR, svc, sid).expect("client attach");
809 +
810 + let sender = thread::spawn({
811 + let svc_name = svc.to_string();
812 + move || {
813 + thread::sleep(std::time::Duration::from_millis(50));
814 + let mut ctx =
815 + ShmContext::client_attach(TEST_RUN_DIR, &svc_name, sid).expect("attach sender");
816 + let msg = build_message(protocol::KIND_REQUEST, 1, 10, &[0xCC, 0xDD, 0xEE]);
817 + ctx.send(&msg).expect("send");
818 + ctx.close();
819 + }
820 + });
821 +
822 + let mut buf = [0u8; 128];
823 + let n = server
824 + .receive(&mut buf, 5000)
825 + .expect("receive with finite timeout");
826 + assert_eq!(n, protocol::HEADER_SIZE + 3);
827 +
828 + let hdr = protocol::Header::decode(&buf[..protocol::HEADER_SIZE]).expect("decode");
829 + assert_eq!(hdr.kind, protocol::KIND_REQUEST);
830 + assert_eq!(hdr.message_id, 10);
831 + assert_eq!(&buf[protocol::HEADER_SIZE..n], &[0xCC, 0xDD, 0xEE]);
832 +
833 + sender.join().unwrap();
834 + client.close();
835 + server.destroy();
836 + cleanup_shm(svc, sid);
837 +}
838 +
839 +// -----------------------------------------------------------------------
840 +// Client attach: bad magic, bad version, bad header_len (lines 366-385)
841 +// -----------------------------------------------------------------------
842 +
843 +#[test]
844 +fn test_client_attach_bad_magic() {
845 + ensure_run_dir();
846 + let svc = "rs_shm_badmag";
847 + let sid: u64 = 59;
848 + cleanup_shm(svc, sid);
849 +
850 + let mut server =
851 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
852 +
853 + // Corrupt magic
854 + let hdr = server.base as *mut RegionHeader;
855 + unsafe { (*hdr).magic = 0xDEADBEEF };
856 +
857 + let result = ShmContext::client_attach(TEST_RUN_DIR, svc, sid);
858 + assert!(matches!(result, Err(ShmError::BadMagic)));
859 +
860 + // Restore for cleanup
861 + unsafe { (*hdr).magic = REGION_MAGIC };
862 + server.destroy();
863 + cleanup_shm(svc, sid);
864 +}
865 +
866 +#[test]
867 +fn test_client_attach_bad_version() {
868 + ensure_run_dir();
869 + let svc = "rs_shm_badver";
870 + let sid: u64 = 60;
871 + cleanup_shm(svc, sid);
872 +
873 + let mut server =
874 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
875 +
876 + let hdr = server.base as *mut RegionHeader;
877 + unsafe { (*hdr).version = 99 };
878 +
879 + let result = ShmContext::client_attach(TEST_RUN_DIR, svc, sid);
880 + assert!(matches!(result, Err(ShmError::BadVersion)));
881 +
882 + unsafe { (*hdr).version = REGION_VERSION };
883 + server.destroy();
884 + cleanup_shm(svc, sid);
885 +}
886 +
887 +#[test]
888 +fn test_client_attach_bad_header_len() {
889 + ensure_run_dir();
890 + let svc = "rs_shm_badhdr";
891 + let sid: u64 = 61;
892 + cleanup_shm(svc, sid);
893 +
894 + let mut server =
895 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
896 +
897 + let hdr = server.base as *mut RegionHeader;
898 + unsafe { (*hdr).header_len = 128 };
899 +
900 + let result = ShmContext::client_attach(TEST_RUN_DIR, svc, sid);
901 + assert!(matches!(result, Err(ShmError::BadHeader)));
902 +
903 + unsafe { (*hdr).header_len = HEADER_LEN };
904 + server.destroy();
905 + cleanup_shm(svc, sid);
906 +}
907 +
908 +#[test]
909 +fn test_client_attach_partial_header_not_ready() {
910 + ensure_run_dir();
911 + let svc = "rs_shm_partial";
912 + let sid: u64 = 63;
913 + cleanup_shm(svc, sid);
914 +
915 + let mut server =
916 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
917 +
918 + let hdr = server.base as *mut RegionHeader;
919 + let saved = unsafe {
920 + (
921 + (*hdr).request_offset,
922 + (*hdr).request_capacity,
923 + (*hdr).response_offset,
924 + (*hdr).response_capacity,
925 + )
926 + };
927 +
928 + unsafe {
929 + (*hdr).request_offset = 0;
930 + (*hdr).request_capacity = 0;
931 + (*hdr).response_offset = 0;
932 + (*hdr).response_capacity = 0;
933 + }
934 +
935 + let result = ShmContext::client_attach(TEST_RUN_DIR, svc, sid);
936 + assert!(matches!(result, Err(ShmError::NotReady)));
937 +
938 + unsafe {
939 + (*hdr).request_offset = saved.0;
940 + (*hdr).request_capacity = saved.1;
941 + (*hdr).response_offset = saved.2;
942 + (*hdr).response_capacity = saved.3;
943 + }
944 + server.destroy();
945 + cleanup_shm(svc, sid);
946 +}
947 +
948 +#[test]
949 +fn test_server_create_rejects_live_region() {
950 + ensure_run_dir();
951 + let svc = "rs_shm_live_region";
952 + let sid: u64 = 631;
953 + cleanup_shm(svc, sid);
954 +
955 + let mut server =
956 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
957 + let err = match ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024) {
958 + Ok(_) => panic!("live region should reject second create"),
959 + Err(err) => err,
960 + };
961 + assert_eq!(err, ShmError::AddrInUse);
962 +
963 + server.destroy();
964 + cleanup_shm(svc, sid);
965 +}
966 +
967 +#[test]
968 +fn test_server_create_recovers_invalid_stale_file() {
969 + ensure_run_dir();
970 + let svc = "rs_shm_invalid_stale_retry";
971 + let sid: u64 = 6311;
972 + cleanup_shm(svc, sid);
973 +
974 + let path = build_shm_path(TEST_RUN_DIR, svc, sid).expect("path");
975 + std::fs::write(&path, [0u8; 8]).expect("write stale short file");
976 +
977 + let mut server = ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024)
978 + .expect("server create should recover invalid stale file");
979 + assert_eq!(server.role, ShmRole::Server);
980 +
981 + server.destroy();
982 + cleanup_shm(svc, sid);
983 +}
984 +
985 +#[test]
986 +fn test_client_attach_short_file_not_ready() {
987 + ensure_run_dir();
988 + let svc = "rs_shm_short";
989 + let sid: u64 = 632;
990 + cleanup_shm(svc, sid);
991 +
992 + let path = build_shm_path(TEST_RUN_DIR, svc, sid).expect("path");
993 + let c_path = path_to_cstring(&path).expect("cstring");
994 + let fd = unsafe {
995 + libc::open(
996 + c_path.as_ptr(),
997 + libc::O_RDWR | libc::O_CREAT | libc::O_TRUNC,
998 + 0o600,
999 + )
1000 + };
1001 + assert!(fd >= 0, "open failed: {}", errno());
1002 + assert_eq!(unsafe { libc::ftruncate(fd, 8) }, 0, "ftruncate failed");
1003 + unsafe { libc::close(fd) };
1004 +
1005 + let err = match ShmContext::client_attach(TEST_RUN_DIR, svc, sid) {
1006 + Ok(_) => panic!("short file should not attach"),
1007 + Err(err) => err,
1008 + };
1009 + assert_eq!(err, ShmError::NotReady);
1010 +
1011 + let _ = std::fs::remove_file(path);
1012 +}
1013 +
1014 +#[test]
1015 +fn test_client_attach_region_smaller_than_declared_capacity() {
1016 + ensure_run_dir();
1017 + let svc = "rs_shm_truncated_region";
1018 + let sid: u64 = 633;
1019 + cleanup_shm(svc, sid);
1020 +
1021 + let mut server =
1022 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
1023 + assert_eq!(
1024 + unsafe { libc::ftruncate(server.fd, HEADER_LEN as libc::off_t) },
1025 + 0,
1026 + "truncate region"
1027 + );
1028 +
1029 + let err = match ShmContext::client_attach(TEST_RUN_DIR, svc, sid) {
1030 + Ok(_) => panic!("truncated region should not attach"),
1031 + Err(err) => err,
1032 + };
1033 + assert_eq!(err, ShmError::BadSize);
1034 +
1035 + server.destroy();
1036 + cleanup_shm(svc, sid);
1037 +}
1038 +
1039 +#[test]
1040 +fn test_client_attach_bad_size() {
1041 + ensure_run_dir();
1042 + let svc = "rs_shm_badsz";
1043 + let sid: u64 = 64;
1044 + cleanup_shm(svc, sid);
1045 +
1046 + let mut server =
1047 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
1048 +
1049 + // Corrupt response capacity to be huge, so region_size < needed
1050 + let hdr = server.base as *mut RegionHeader;
1051 + let saved_cap = unsafe { (*hdr).response_capacity };
1052 + unsafe { (*hdr).response_capacity = 0xFFFF_FFFF };
1053 +
1054 + let result = ShmContext::client_attach(TEST_RUN_DIR, svc, sid);
1055 + assert!(matches!(result, Err(ShmError::BadSize)));
1056 +
1057 + unsafe { (*hdr).response_capacity = saved_cap };
1058 + server.destroy();
1059 + cleanup_shm(svc, sid);
1060 +}
1061 +
1062 +// -----------------------------------------------------------------------
1063 +// validate_service_name (lines 764-778)
1064 +// -----------------------------------------------------------------------
1065 +
1066 +#[test]
1067 +fn test_validate_service_name() {
1068 + assert!(validate_service_name("").is_err());
1069 + assert!(validate_service_name(".").is_err());
1070 + assert!(validate_service_name("..").is_err());
1071 + assert!(validate_service_name("a/b").is_err());
1072 + assert!(validate_service_name("a b").is_err());
1073 + assert!(validate_service_name("a@b").is_err());
1074 +
1075 + assert!(validate_service_name("valid").is_ok());
1076 + assert!(validate_service_name("valid-name").is_ok());
1077 + assert!(validate_service_name("valid_name").is_ok());
1078 + assert!(validate_service_name("valid.name.123").is_ok());
1079 +}
1080 +
1081 +// -----------------------------------------------------------------------
1082 +// build_shm_path: path too long (line 785)
1083 +// -----------------------------------------------------------------------
1084 +
1085 +#[test]
1086 +fn test_build_shm_path_too_long() {
1087 + // Path = "{run_dir}/{name}-{session_id:016x}.ipcshm" must be >= 256
1088 + // /tmp/aaa...aaa-0000000000000001.ipcshm
1089 + // 5 + name_len + 1 + 16 + 7 = 29 + name_len >= 256 -> name_len >= 227
1090 + let long_name = "a".repeat(230);
1091 + let result = build_shm_path("/tmp", &long_name, 1);
1092 + assert!(matches!(result, Err(ShmError::PathTooLong)));
1093 +}
1094 +
1095 +// -----------------------------------------------------------------------
1096 +// pid_alive edge cases (lines 800-804)
1097 +// -----------------------------------------------------------------------
1098 +
1099 +#[test]
1100 +fn test_pid_alive() {
1101 + assert!(!pid_alive(0));
1102 + assert!(!pid_alive(-1));
1103 + // Current process should be alive
1104 + assert!(pid_alive(unsafe { libc::getpid() }));
1105 + // Very unlikely PID
1106 + assert!(!pid_alive(99999));
1107 +}
1108 +
1109 +// -----------------------------------------------------------------------
1110 +// Client attach: nonexistent file (line 325)
1111 +// -----------------------------------------------------------------------
1112 +
1113 +#[test]
1114 +fn test_client_attach_nonexistent() {
1115 + ensure_run_dir();
1116 + let result = ShmContext::client_attach(TEST_RUN_DIR, "rs_shm_nofile", 99999);
1117 + assert!(matches!(result, Err(ShmError::Open(_))));
1118 +}
1119 +
1120 +#[test]
1121 +fn test_check_shm_stale_nonexistent_returns_not_exist() {
1122 + ensure_run_dir();
1123 + let svc = "rs_shm_stale_missing";
1124 + let sid: u64 = 700;
1125 + cleanup_shm(svc, sid);
1126 +
1127 + let path = build_shm_path(TEST_RUN_DIR, svc, sid).expect("path");
1128 + assert!(matches!(check_shm_stale(&path), StaleResult::NotExist));
1129 +}
1130 +
1131 +#[test]
1132 +fn test_check_shm_stale_invalid_cstring_returns_not_exist() {
1133 + let bad_path = PathBuf::from(OsString::from_vec(vec![
1134 + b'/', b't', b'm', b'p', b'/', b'n', b'i', b'p', b'c', 0, b'b',
1135 + ]));
1136 + assert!(matches!(check_shm_stale(&bad_path), StaleResult::NotExist));
1137 +}
1138 +
1139 +#[test]
1140 +fn test_stale_open_failure_policies_are_conservative() {
1141 + assert!(should_unlink_cleanup_open_failure(libc::ENOENT));
1142 + assert!(!should_unlink_cleanup_open_failure(libc::EACCES));
1143 + assert!(!should_unlink_cleanup_open_failure(libc::EPERM));
1144 + assert!(!should_unlink_cleanup_open_failure(libc::EMFILE));
1145 + assert!(!should_unlink_cleanup_open_failure(libc::ENFILE));
1146 + assert!(!should_unlink_cleanup_open_failure(libc::ELOOP));
1147 +
1148 + assert!(matches!(
1149 + classify_stale_open_failure(libc::ENOENT),
1150 + StaleResult::NotExist
1151 + ));
1152 + assert!(matches!(
1153 + classify_stale_open_failure(libc::EACCES),
1154 + StaleResult::Invalid
1155 + ));
1156 + assert!(matches!(
1157 + classify_stale_open_failure(libc::EMFILE),
1158 + StaleResult::Invalid
1159 + ));
1160 + assert!(matches!(
1161 + classify_stale_open_failure(libc::ELOOP),
1162 + StaleResult::Invalid
1163 + ));
1164 +}
1165 +
1166 +#[test]
1167 +fn test_cleanup_stale_missing_run_dir_is_noop() {
1168 + cleanup_stale("/tmp/nipc_shm_rust_missing_dir", "rs_shm_missing");
1169 +}
1170 +
1171 +#[test]
1172 +fn test_cleanup_stale_ignores_unrelated_and_non_utf8_entries() {
1173 + ensure_run_dir();
1174 + let svc = "rs_shm_cleanup_skip";
1175 + let unrelated = PathBuf::from(format!("{TEST_RUN_DIR}/not-a-shm-entry.txt"));
1176 + let invalid_name = OsString::from_vec(vec![
1177 + b'r', b's', b'_', b's', b'h', b'm', b'_', b'c', b'l', b'e', b'a', b'n', b'u', b'p', b'_',
1178 + b's', b'k', b'i', b'p', b'-', 0xff, b'.', b'i', b'p', b'c', b's', b'h', b'm',
1179 + ]);
1180 + let invalid_path = PathBuf::from(TEST_RUN_DIR).join(invalid_name);
1181 +
1182 + let _ = std::fs::remove_file(&unrelated);
1183 + let _ = std::fs::remove_file(&invalid_path);
1184 + std::fs::write(&unrelated, b"skip").expect("write unrelated file");
1185 + std::fs::write(&invalid_path, b"skip").expect("write invalid utf8 file");
1186 +
1187 + cleanup_stale(TEST_RUN_DIR, svc);
1188 +
1189 + assert!(unrelated.exists(), "unrelated entries should be ignored");
1190 + assert!(invalid_path.exists(), "non-UTF8 entries should be ignored");
1191 +
1192 + let _ = std::fs::remove_file(&unrelated);
1193 + let _ = std::fs::remove_file(&invalid_path);
1194 +}
1195 +
1196 +#[test]
1197 +fn test_cleanup_stale_invalid_entries() {
1198 + ensure_run_dir();
1199 + let svc = "rs_shm_cleanup_invalid";
1200 + let short_sid: u64 = 701;
1201 + let magic_sid: u64 = 702;
1202 + let unreadable_sid: u64 = 703;
1203 +
1204 + for sid in [short_sid, magic_sid, unreadable_sid] {
1205 + cleanup_shm(svc, sid);
1206 + }
1207 +
1208 + let short_path = build_shm_path(TEST_RUN_DIR, svc, short_sid).expect("short path");
1209 + std::fs::write(&short_path, [0u8; 8]).expect("write short file");
1210 +
1211 + let mut bad_magic =
1212 + ShmContext::server_create(TEST_RUN_DIR, svc, magic_sid, 1024, 1024).expect("server create");
1213 + let hdr = bad_magic.base as *mut RegionHeader;
1214 + unsafe { (*hdr).magic = 0xDEADBEEF };
1215 + bad_magic.close();
1216 +
1217 + let unreadable_path = build_shm_path(TEST_RUN_DIR, svc, unreadable_sid).expect("path");
1218 + let unreadable_c = path_to_cstring(&unreadable_path).expect("cstring");
1219 + let unreadable_fd = unsafe {
1220 + libc::open(
1221 + unreadable_c.as_ptr(),
1222 + libc::O_RDWR | libc::O_CREAT | libc::O_TRUNC,
1223 + 0o000,
1224 + )
1225 + };
1226 + assert!(unreadable_fd >= 0, "open unreadable");
1227 + unsafe { libc::close(unreadable_fd) };
1228 +
1229 + cleanup_stale(TEST_RUN_DIR, svc);
1230 +
1231 + assert!(
1232 + !short_path.exists(),
1233 + "short invalid entry should be removed"
1234 + );
1235 + assert!(!build_shm_path(TEST_RUN_DIR, svc, magic_sid)
1236 + .unwrap()
1237 + .exists());
1238 + // Under non-root: unreadable file is preserved (EACCES skip).
1239 + // Under root: chmod 000 has no effect, so the file gets opened,
1240 + // inspected, and removed normally.
1241 + if unsafe { libc::geteuid() } != 0 {
1242 + assert!(
1243 + unreadable_path.exists(),
1244 + "unreadable entry should be preserved (EACCES skip)"
1245 + );
1246 + }
1247 + // Clean up
1248 + unsafe { libc::chmod(unreadable_c.as_ptr(), 0o600) };
1249 + std::fs::remove_file(&unreadable_path).ok();
1250 +}
1251 +
1252 +#[test]
1253 +fn test_check_shm_stale_short_file_invalid() {
1254 + ensure_run_dir();
1255 + let svc = "rs_shm_stale_short_direct";
1256 + let sid: u64 = 704;
1257 + cleanup_shm(svc, sid);
1258 +
1259 + let path = build_shm_path(TEST_RUN_DIR, svc, sid).expect("path");
1260 + std::fs::write(&path, [0u8; 8]).expect("write short file");
1261 +
1262 + assert!(matches!(check_shm_stale(&path), StaleResult::Invalid));
1263 + assert!(!path.exists(), "short stale file should be removed");
1264 +}
1265 +
1266 +#[test]
1267 +fn test_check_shm_stale_bad_magic_invalid() {
1268 + ensure_run_dir();
1269 + let svc = "rs_shm_stale_magic_direct";
1270 + let sid: u64 = 705;
1271 + cleanup_shm(svc, sid);
1272 +
1273 + let mut server =
1274 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
1275 + let hdr = server.base as *mut RegionHeader;
1276 + unsafe { (*hdr).magic = 0xDEADBEEF };
1277 + server.close();
1278 +
1279 + let path = build_shm_path(TEST_RUN_DIR, svc, sid).expect("path");
1280 + assert!(matches!(check_shm_stale(&path), StaleResult::Invalid));
1281 + assert!(!path.exists(), "bad magic stale file should be removed");
1282 +}
1283 +
1284 +#[test]
1285 +fn test_check_shm_stale_zero_generation_recovers() {
1286 + ensure_run_dir();
1287 + let svc = "rs_shm_stale_zero_gen";
1288 + let sid: u64 = 7051;
1289 + cleanup_shm(svc, sid);
1290 +
1291 + let mut server =
1292 + ShmContext::server_create(TEST_RUN_DIR, svc, sid, 1024, 1024).expect("server create");
1293 + let hdr = server.base as *mut RegionHeader;
1294 + unsafe { (*hdr).owner_generation = 0 };
1295 + server.close();
1296 +
1297 + let path = build_shm_path(TEST_RUN_DIR, svc, sid).expect("path");
1298 + assert!(matches!(check_shm_stale(&path), StaleResult::Recovered));
1299 + assert!(
1300 + !path.exists(),
1301 + "zero-generation stale file should be removed"
1302 + );
1303 +}
1304 +
1305 +#[test]
1306 +fn test_check_shm_stale_open_failure_invalid() {
1307 + ensure_run_dir();
1308 + let svc = "rs_shm_stale_open_fail";
1309 + let sid: u64 = 7052;
1310 + cleanup_shm(svc, sid);
1311 +
1312 + let path = build_shm_path(TEST_RUN_DIR, svc, sid).expect("path");
1313 + let c_path = path_to_cstring(&path).expect("cstring");
1314 + let fd = unsafe {
1315 + libc::open(
1316 + c_path.as_ptr(),
1317 + libc::O_CREAT | libc::O_TRUNC | libc::O_RDWR,
1318 + 0o600,
1319 + )
1320 + };
1321 + assert!(fd >= 0, "open failed: {}", errno());
1322 + assert_eq!(
1323 + unsafe { libc::ftruncate(fd, HEADER_LEN as libc::off_t) },
1324 + 0,
1325 + "ftruncate failed: {}",
1326 + errno()
1327 + );
1328 + unsafe { libc::close(fd) };
1329 + assert_eq!(
1330 + unsafe { libc::chmod(c_path.as_ptr(), 0) },
1331 + 0,
1332 + "chmod failed: {}",
1333 + errno()
1334 + );
1335 +
1336 + assert!(matches!(check_shm_stale(&path), StaleResult::Invalid));
1337 + // Under non-root: file preserved (EACCES). Under root: chmod 000
1338 + // has no effect, so the file is opened, inspected, and removed.
1339 + if unsafe { libc::geteuid() } != 0 {
1340 + assert!(
1341 + path.exists(),
1342 + "unopenable file should be preserved (EACCES skip)"
1343 + );
1344 + }
1345 + unsafe { libc::chmod(c_path.as_ptr(), 0o600) };
1346 + std::fs::remove_file(&path).ok();
1347 +}
1348 +
1349 +#[test]
1350 +fn test_check_shm_stale_directory_symlink_invalid() {
1351 + ensure_run_dir();
1352 + let svc = "rs_shm_stale_dir_symlink";
1353 + let sid: u64 = 7053;
1354 + cleanup_shm(svc, sid);
1355 +
1356 + let path = build_shm_path(TEST_RUN_DIR, svc, sid).expect("path");
1357 + let target = PathBuf::from(format!("{TEST_RUN_DIR}/{svc}-target-dir"));
1358 + let _ = std::fs::remove_dir_all(&target);
1359 + std::fs::create_dir_all(&target).expect("create target dir");
1360 + std::os::unix::fs::symlink(&target, &path).expect("create symlink");
1361 +
1362 + assert!(matches!(check_shm_stale(&path), StaleResult::Invalid));
1363 + assert!(
1364 + !path.exists(),
1365 + "directory symlink stale entry should be removed"
1366 + );
1367 + assert!(target.is_dir(), "target directory should remain");
1368 +
1369 + std::fs::remove_dir_all(&target).expect("remove target dir");
1370 +}
1371 +
1372 +#[test]
1373 +fn test_cleanup_stale_unlinks_dangling_symlink() {
1374 + ensure_run_dir();
1375 + let svc = "rs_shm_cleanup_symlink";
1376 + let sid: u64 = 706;
1377 + cleanup_shm(svc, sid);
1378 +
1379 + let path = build_shm_path(TEST_RUN_DIR, svc, sid).expect("path");
1380 + let target = path.with_extension("missing-target");
1381 + let _ = std::fs::remove_file(&target);
1382 + std::os::unix::fs::symlink(&target, &path).expect("create dangling symlink");
1383 + assert!(
1384 + std::fs::symlink_metadata(&path)
1385 + .map(|meta| meta.file_type().is_symlink())
1386 + .unwrap_or(false),
1387 + "dangling symlink test entry should exist before cleanup"
1388 + );
1389 +
1390 + cleanup_stale(TEST_RUN_DIR, svc);
1391 +
1392 + assert!(
1393 + std::fs::symlink_metadata(&path).is_err(),
1394 + "dangling symlink entry should be removed"
1395 + );
1396 +}
1397 +
1398 +#[test]
1399 +fn test_cleanup_stale_preserves_self_referential_symlink_open_failure() {
1400 + ensure_run_dir();
1401 + let svc = "rs_shm_cleanup_self_symlink";
1402 + let sid: u64 = 7061;
1403 + cleanup_shm(svc, sid);
1404 +
1405 + let path = build_shm_path(TEST_RUN_DIR, svc, sid).expect("path");
1406 + let _ = std::fs::remove_file(&path);
1407 + std::os::unix::fs::symlink(&path, &path).expect("create self symlink");
1408 + assert!(
1409 + std::fs::symlink_metadata(&path)
1410 + .map(|meta| meta.file_type().is_symlink())
1411 + .unwrap_or(false),
1412 + "self-referential symlink test entry should exist before cleanup"
1413 + );
1414 +
1415 + cleanup_stale(TEST_RUN_DIR, svc);
1416 +
1417 + assert!(
1418 + std::fs::symlink_metadata(&path)
1419 + .map(|meta| meta.file_type().is_symlink())
1420 + .unwrap_or(false),
1421 + "ambiguous open failures must not delete the entry"
1422 + );
1423 +
1424 + std::fs::remove_file(&path).expect("remove self symlink");
1425 +}
1426 +
1427 +#[test]
1428 +fn test_cleanup_stale_unlinks_directory_symlink_when_mmap_fails() {
1429 + ensure_run_dir();
1430 + let svc = "rs_shm_cleanup_dir_symlink";
1431 + let sid: u64 = 707;
1432 + cleanup_shm(svc, sid);
1433 +
1434 + let path = build_shm_path(TEST_RUN_DIR, svc, sid).expect("path");
1435 + let target = PathBuf::from(format!("{TEST_RUN_DIR}/{svc}-target-dir"));
1436 + let _ = std::fs::remove_dir_all(&target);
1437 + std::fs::create_dir_all(&target).expect("create target dir");
1438 + std::os::unix::fs::symlink(&target, &path).expect("create symlink");
1439 +
1440 + cleanup_stale(TEST_RUN_DIR, svc);
1441 +
1442 + assert!(!path.exists(), "directory symlink entry should be removed");
1443 + assert!(target.is_dir(), "target directory should remain");
1444 +
1445 + std::fs::remove_dir_all(&target).expect("remove target dir");
1446 +}
src/crates/netipc/src/transport/win_shm.rs new
+1984
@@ -0,0 +1,1984 @@
1 +//! L1 Windows SHM transport.
2 +//!
3 +//! Shared memory data plane with spin + kernel event synchronization.
4 +//! Uses CreateFileMappingW/MapViewOfFile for the region and auto-reset
5 +//! kernel events for signaling (SHM_HYBRID profile).
6 +//!
7 +//! Wire-compatible with the C and Go implementations.
8 +
9 +#[cfg(windows)]
10 +mod ffi {
11 + #![allow(non_snake_case, non_camel_case_types, dead_code)]
12 +
13 + pub type HANDLE = isize;
14 + pub type DWORD = u32;
15 + pub type BOOL = i32;
16 + pub type LPCWSTR = *const u16;
17 + pub type LONG = i32;
18 + pub type LONG64 = i64;
19 +
20 + pub const INVALID_HANDLE_VALUE: HANDLE = -1;
21 + pub const PAGE_READWRITE: DWORD = 0x04;
22 + pub const FILE_MAP_ALL_ACCESS: DWORD = 0x000F001F;
23 + pub const EVENT_MODIFY_STATE: DWORD = 0x0002;
24 + pub const SYNCHRONIZE: DWORD = 0x00100000;
25 + pub const INFINITE: DWORD = 0xFFFFFFFF;
26 + pub const WAIT_OBJECT_0: DWORD = 0x00000000;
27 + pub const WAIT_TIMEOUT: DWORD = 0x00000102;
28 +
29 + extern "system" {
30 + pub fn CreateFileMappingW(
31 + hFile: HANDLE,
32 + lpFileMappingAttributes: *const core::ffi::c_void,
33 + flProtect: DWORD,
34 + dwMaximumSizeHigh: DWORD,
35 + dwMaximumSizeLow: DWORD,
36 + lpName: LPCWSTR,
37 + ) -> HANDLE;
38 +
39 + pub fn OpenFileMappingW(
40 + dwDesiredAccess: DWORD,
41 + bInheritHandle: BOOL,
42 + lpName: LPCWSTR,
43 + ) -> HANDLE;
44 +
45 + pub fn MapViewOfFile(
46 + hFileMappingObject: HANDLE,
47 + dwDesiredAccess: DWORD,
48 + dwFileOffsetHigh: DWORD,
49 + dwFileOffsetLow: DWORD,
50 + dwNumberOfBytesToMap: usize,
51 + ) -> *mut core::ffi::c_void;
52 +
53 + pub fn UnmapViewOfFile(lpBaseAddress: *const core::ffi::c_void) -> BOOL;
54 +
55 + pub fn CreateEventW(
56 + lpEventAttributes: *const core::ffi::c_void,
57 + bManualReset: BOOL,
58 + bInitialState: BOOL,
59 + lpName: LPCWSTR,
60 + ) -> HANDLE;
61 +
62 + pub fn OpenEventW(dwDesiredAccess: DWORD, bInheritHandle: BOOL, lpName: LPCWSTR) -> HANDLE;
63 +
64 + pub fn SetEvent(hEvent: HANDLE) -> BOOL;
65 +
66 + pub fn WaitForSingleObject(hHandle: HANDLE, dwMilliseconds: DWORD) -> DWORD;
67 +
68 + pub fn CloseHandle(hObject: HANDLE) -> BOOL;
69 +
70 + pub fn GetLastError() -> DWORD;
71 + pub fn SetLastError(dwErrCode: DWORD);
72 +
73 + pub fn GetTickCount() -> DWORD;
74 + pub fn GetTickCount64() -> u64;
75 +
76 + // Note: InterlockedXxx are MSVC compiler intrinsics, not linkable
77 + // symbols. We use Rust's std::sync::atomic instead (see helpers below).
78 + }
79 +}
80 +
81 +// ---------------------------------------------------------------------------
82 +// Constants
83 +// ---------------------------------------------------------------------------
84 +
85 +/// Magic value: "NSWH" as u32 LE.
86 +pub const REGION_MAGIC: u32 = 0x4e535748;
87 +pub const REGION_VERSION: u32 = 3;
88 +pub const HEADER_LEN: u32 = 128;
89 +pub const CACHELINE_SIZE: u32 = 64;
90 +pub const DEFAULT_SPIN_TRIES: u32 = 1024;
91 +pub const BUSYWAIT_POLL_MASK: u32 = 1023;
92 +
93 +pub const PROFILE_HYBRID: u32 = 0x02;
94 +pub const PROFILE_BUSYWAIT: u32 = 0x04;
95 +const ERROR_ALREADY_EXISTS: u32 = 183;
96 +
97 +// Header field byte offsets
98 +const OFF_MAGIC: usize = 0;
99 +const OFF_VERSION: usize = 4;
100 +const OFF_HEADER_LEN: usize = 8;
101 +const OFF_PROFILE: usize = 12;
102 +const OFF_REQ_OFFSET: usize = 16;
103 +const OFF_REQ_CAPACITY: usize = 20;
104 +const OFF_RESP_OFFSET: usize = 24;
105 +const OFF_RESP_CAPACITY: usize = 28;
106 +const OFF_SPIN_TRIES: usize = 32;
107 +const OFF_REQ_LEN: usize = 36;
108 +const OFF_RESP_LEN: usize = 40;
109 +const OFF_REQ_CLIENT_CLOSED: usize = 44;
110 +const OFF_REQ_SERVER_WAITING: usize = 48;
111 +const OFF_RESP_SERVER_CLOSED: usize = 52;
112 +const OFF_RESP_CLIENT_WAITING: usize = 56;
113 +const OFF_REQ_SEQ: usize = 64;
114 +const OFF_RESP_SEQ: usize = 72;
115 +
116 +// FNV-1a 64-bit constants
117 +const FNV1A_OFFSET_BASIS: u64 = 0xcbf29ce484222325;
118 +const FNV1A_PRIME: u64 = 0x00000100000001B3;
119 +
120 +#[cfg(all(test, windows))]
121 +#[derive(Debug, Clone, Copy, PartialEq, Eq)]
122 +enum WinShmFaultSite {
123 + CreateFileMapping,
124 + OpenFileMapping,
125 + MapViewOfFile,
126 + CreateEvent,
127 + OpenEvent,
128 +}
129 +
130 +#[cfg(all(test, windows))]
131 +#[derive(Debug, Clone, Copy)]
132 +struct WinShmFaultHook {
133 + site: WinShmFaultSite,
134 + error: u32,
135 + skip_matches: u32,
136 +}
137 +
138 +#[cfg(all(test, windows))]
139 +thread_local! {
140 + static WIN_SHM_FAULT_HOOK: std::cell::RefCell<Option<WinShmFaultHook>> =
141 + const { std::cell::RefCell::new(None) };
142 +}
143 +
144 +// ---------------------------------------------------------------------------
145 +// Errors
146 +// ---------------------------------------------------------------------------
147 +
148 +/// Windows SHM transport errors.
149 +#[derive(Debug, Clone, PartialEq, Eq)]
150 +pub enum WinShmError {
151 + /// Invalid argument.
152 + BadParam(String),
153 + /// CreateFileMappingW failed.
154 + CreateMapping(u32),
155 + /// OpenFileMappingW failed.
156 + OpenMapping(u32),
157 + /// MapViewOfFile failed.
158 + MapView(u32),
159 + /// CreateEventW failed.
160 + CreateEvent(u32),
161 + /// OpenEventW failed.
162 + OpenEvent(u32),
163 + /// Named mapping/event already exists.
164 + AddrInUse,
165 + /// Header magic mismatch.
166 + BadMagic,
167 + /// Header version mismatch.
168 + BadVersion,
169 + /// header_len mismatch.
170 + BadHeader,
171 + /// Profile mismatch.
172 + BadProfile,
173 + /// Message exceeds area capacity.
174 + MsgTooLarge,
175 + /// Wait timed out.
176 + Timeout,
177 + /// Peer closed.
178 + Disconnected,
179 +}
180 +
181 +impl std::fmt::Display for WinShmError {
182 + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
183 + match self {
184 + WinShmError::BadParam(s) => write!(f, "bad parameter: {s}"),
185 + WinShmError::CreateMapping(e) => write!(f, "CreateFileMappingW failed: {e}"),
186 + WinShmError::OpenMapping(e) => write!(f, "OpenFileMappingW failed: {e}"),
187 + WinShmError::MapView(e) => write!(f, "MapViewOfFile failed: {e}"),
188 + WinShmError::CreateEvent(e) => write!(f, "CreateEventW failed: {e}"),
189 + WinShmError::OpenEvent(e) => write!(f, "OpenEventW failed: {e}"),
190 + WinShmError::AddrInUse => {
191 + write!(f, "Windows SHM object name already in use by live server")
192 + }
193 + WinShmError::BadMagic => write!(f, "SHM header magic mismatch"),
194 + WinShmError::BadVersion => write!(f, "SHM header version mismatch"),
195 + WinShmError::BadHeader => write!(f, "SHM header_len mismatch"),
196 + WinShmError::BadProfile => write!(f, "SHM profile mismatch"),
197 + WinShmError::MsgTooLarge => write!(f, "message exceeds SHM area capacity"),
198 + WinShmError::Timeout => write!(f, "SHM wait timed out"),
199 + WinShmError::Disconnected => write!(f, "peer closed SHM session"),
200 + }
201 + }
202 +}
203 +
204 +impl std::error::Error for WinShmError {}
205 +
206 +// ---------------------------------------------------------------------------
207 +// Role
208 +// ---------------------------------------------------------------------------
209 +
210 +#[derive(Debug, Clone, Copy, PartialEq, Eq)]
211 +pub enum WinShmRole {
212 + Server = 1,
213 + Client = 2,
214 +}
215 +
216 +// ---------------------------------------------------------------------------
217 +// Compile-time header size assertion
218 +// ---------------------------------------------------------------------------
219 +
220 +// We don't map a Rust struct onto the header; we use raw pointer arithmetic.
221 +// Assert that our offset constants are consistent.
222 +const _: () = assert!(OFF_REQ_SEQ == 64);
223 +const _: () = assert!(OFF_RESP_SEQ == 72);
224 +const _: () = assert!(HEADER_LEN == 128);
225 +
226 +#[cfg(all(test, windows))]
227 +fn install_fault_hook(site: WinShmFaultSite, error: u32, skip_matches: u32) {
228 + WIN_SHM_FAULT_HOOK.with(|slot| {
229 + let mut slot = slot.borrow_mut();
230 + assert!(slot.is_none(), "win_shm fault hook already installed");
231 + *slot = Some(WinShmFaultHook {
232 + site,
233 + error,
234 + skip_matches,
235 + });
236 + });
237 +}
238 +
239 +#[cfg(all(test, windows))]
240 +fn clear_fault_hook() {
241 + WIN_SHM_FAULT_HOOK.with(|slot| {
242 + *slot.borrow_mut() = None;
243 + });
244 +}
245 +
246 +#[cfg(all(test, windows))]
247 +fn take_fault_hook(site: WinShmFaultSite) -> Option<u32> {
248 + WIN_SHM_FAULT_HOOK.with(|slot| {
249 + let mut slot = slot.borrow_mut();
250 + match *slot {
251 + Some(mut hook) if hook.site == site => {
252 + if hook.skip_matches > 0 {
253 + hook.skip_matches -= 1;
254 + *slot = Some(hook);
255 + None
256 + } else {
257 + *slot = None;
258 + Some(hook.error)
259 + }
260 + }
261 + _ => None,
262 + }
263 + })
264 +}
265 +
266 +#[cfg(windows)]
267 +unsafe fn create_file_mapping(
268 + h_file: ffi::HANDLE,
269 + attrs: *const core::ffi::c_void,
270 + protect: u32,
271 + size_high: u32,
272 + size_low: u32,
273 + name: *const u16,
274 +) -> ffi::HANDLE {
275 + #[cfg(test)]
276 + if let Some(error) = take_fault_hook(WinShmFaultSite::CreateFileMapping) {
277 + ffi::SetLastError(error);
278 + return 0;
279 + }
280 +
281 + ffi::CreateFileMappingW(h_file, attrs, protect, size_high, size_low, name)
282 +}
283 +
284 +#[cfg(windows)]
285 +unsafe fn open_file_mapping(access: u32, inherit: i32, name: *const u16) -> ffi::HANDLE {
286 + #[cfg(test)]
287 + if let Some(error) = take_fault_hook(WinShmFaultSite::OpenFileMapping) {
288 + ffi::SetLastError(error);
289 + return 0;
290 + }
291 +
292 + ffi::OpenFileMappingW(access, inherit, name)
293 +}
294 +
295 +#[cfg(windows)]
296 +unsafe fn map_view_of_file(
297 + mapping: ffi::HANDLE,
298 + access: u32,
299 + offset_high: u32,
300 + offset_low: u32,
301 + bytes: usize,
302 +) -> *mut core::ffi::c_void {
303 + #[cfg(test)]
304 + if let Some(error) = take_fault_hook(WinShmFaultSite::MapViewOfFile) {
305 + ffi::SetLastError(error);
306 + return std::ptr::null_mut();
307 + }
308 +
309 + ffi::MapViewOfFile(mapping, access, offset_high, offset_low, bytes)
310 +}
311 +
312 +#[cfg(windows)]
313 +unsafe fn create_event(
314 + attrs: *const core::ffi::c_void,
315 + manual_reset: i32,
316 + initial_state: i32,
317 + name: *const u16,
318 +) -> ffi::HANDLE {
319 + #[cfg(test)]
320 + if let Some(error) = take_fault_hook(WinShmFaultSite::CreateEvent) {
321 + ffi::SetLastError(error);
322 + return 0;
323 + }
324 +
325 + ffi::CreateEventW(attrs, manual_reset, initial_state, name)
326 +}
327 +
328 +#[cfg(windows)]
329 +unsafe fn open_event(access: u32, inherit: i32, name: *const u16) -> ffi::HANDLE {
330 + #[cfg(test)]
331 + if let Some(error) = take_fault_hook(WinShmFaultSite::OpenEvent) {
332 + ffi::SetLastError(error);
333 + return 0;
334 + }
335 +
336 + ffi::OpenEventW(access, inherit, name)
337 +}
338 +
339 +// ---------------------------------------------------------------------------
340 +// Context
341 +// ---------------------------------------------------------------------------
342 +
343 +/// A handle to a Windows SHM region.
344 +#[cfg(windows)]
345 +pub struct WinShmContext {
346 + role: WinShmRole,
347 + mapping: ffi::HANDLE,
348 + base: *mut u8,
349 + region_size: usize,
350 +
351 + req_event: ffi::HANDLE,
352 + resp_event: ffi::HANDLE,
353 +
354 + profile: u32,
355 + request_offset: u32,
356 + request_capacity: u32,
357 + response_offset: u32,
358 + response_capacity: u32,
359 + spin_tries: u32,
360 +
361 + local_req_seq: i64,
362 + local_resp_seq: i64,
363 +}
364 +
365 +#[cfg(windows)]
366 +unsafe impl Send for WinShmContext {}
367 +
368 +#[cfg(windows)]
369 +impl WinShmContext {
370 + pub fn role(&self) -> WinShmRole {
371 + self.role
372 + }
373 +
374 + /// Create a per-session Windows SHM region (server side).
375 + pub fn server_create(
376 + run_dir: &str,
377 + service_name: &str,
378 + auth_token: u64,
379 + session_id: u64,
380 + profile: u32,
381 + req_capacity: u32,
382 + resp_capacity: u32,
383 + ) -> Result<Self, WinShmError> {
384 + validate_service_name(service_name)?;
385 + validate_profile(profile)?;
386 +
387 + let hash = compute_shm_hash(run_dir, service_name, auth_token);
388 + let mapping_name = build_object_name(hash, service_name, profile, session_id, "mapping")?;
389 +
390 + let req_cap = align_cacheline(req_capacity);
391 + let resp_cap = align_cacheline(resp_capacity);
392 + let req_off = align_cacheline(HEADER_LEN);
393 + let resp_off = req_off
394 + .checked_add(req_cap)
395 + .map(align_cacheline)
396 + .ok_or_else(|| WinShmError::BadParam("region offset overflow".into()))?;
397 + let region_size = resp_off
398 + .checked_add(resp_cap)
399 + .ok_or_else(|| WinShmError::BadParam("region size overflow".into()))?
400 + as usize;
401 +
402 + // Create page-file backed mapping
403 + unsafe { ffi::SetLastError(0) };
404 + let mapping = unsafe {
405 + create_file_mapping(
406 + ffi::INVALID_HANDLE_VALUE,
407 + std::ptr::null(),
408 + ffi::PAGE_READWRITE,
409 + (region_size >> 32) as u32,
410 + (region_size & 0xFFFFFFFF) as u32,
411 + mapping_name.as_ptr(),
412 + )
413 + };
414 + if mapping == 0 {
415 + return Err(WinShmError::CreateMapping(last_error()));
416 + }
417 + let mapping_err = last_error();
418 + if mapping_err == ERROR_ALREADY_EXISTS {
419 + unsafe { ffi::CloseHandle(mapping) };
420 + return Err(WinShmError::AddrInUse);
421 + }
422 +
423 + let base =
424 + unsafe { map_view_of_file(mapping, ffi::FILE_MAP_ALL_ACCESS, 0, 0, region_size) };
425 + if base.is_null() {
426 + let e = last_error();
427 + unsafe { ffi::CloseHandle(mapping) };
428 + return Err(WinShmError::MapView(e));
429 + }
430 + let base = base as *mut u8;
431 +
432 + // Zero region
433 + unsafe { std::ptr::write_bytes(base, 0, region_size) };
434 +
435 + // Write header
436 + write_u32(base, OFF_MAGIC, REGION_MAGIC);
437 + write_u32(base, OFF_VERSION, REGION_VERSION);
438 + write_u32(base, OFF_HEADER_LEN, HEADER_LEN);
439 + write_u32(base, OFF_PROFILE, profile);
440 + write_u32(base, OFF_REQ_OFFSET, req_off);
441 + write_u32(base, OFF_REQ_CAPACITY, req_cap);
442 + write_u32(base, OFF_RESP_OFFSET, resp_off);
443 + write_u32(base, OFF_RESP_CAPACITY, resp_cap);
444 + write_u32(base, OFF_SPIN_TRIES, DEFAULT_SPIN_TRIES);
445 +
446 + // Memory barrier
447 + std::sync::atomic::fence(std::sync::atomic::Ordering::Release);
448 +
449 + // Create events for HYBRID
450 + let (req_event, resp_event) = if profile == PROFILE_HYBRID {
451 + let re_name = build_object_name(hash, service_name, profile, session_id, "req_event")?;
452 + unsafe { ffi::SetLastError(0) };
453 + let re = unsafe { create_event(std::ptr::null(), 0, 0, re_name.as_ptr()) };
454 + if re == 0 {
455 + let e = last_error();
456 + unsafe {
457 + ffi::UnmapViewOfFile(base as *const _);
458 + ffi::CloseHandle(mapping);
459 + }
460 + return Err(WinShmError::CreateEvent(e));
461 + }
462 + if last_error() == ERROR_ALREADY_EXISTS {
463 + unsafe {
464 + ffi::CloseHandle(re);
465 + ffi::UnmapViewOfFile(base as *const _);
466 + ffi::CloseHandle(mapping);
467 + }
468 + return Err(WinShmError::AddrInUse);
469 + }
470 +
471 + let rsp_name =
472 + build_object_name(hash, service_name, profile, session_id, "resp_event")?;
473 + unsafe { ffi::SetLastError(0) };
474 + let rsp = unsafe { create_event(std::ptr::null(), 0, 0, rsp_name.as_ptr()) };
475 + if rsp == 0 {
476 + let e = last_error();
477 + unsafe {
478 + ffi::CloseHandle(re);
479 + ffi::UnmapViewOfFile(base as *const _);
480 + ffi::CloseHandle(mapping);
481 + }
482 + return Err(WinShmError::CreateEvent(e));
483 + }
484 + if last_error() == ERROR_ALREADY_EXISTS {
485 + unsafe {
486 + ffi::CloseHandle(rsp);
487 + ffi::CloseHandle(re);
488 + ffi::UnmapViewOfFile(base as *const _);
489 + ffi::CloseHandle(mapping);
490 + }
491 + return Err(WinShmError::AddrInUse);
492 + }
493 + (re, rsp)
494 + } else {
495 + (ffi::INVALID_HANDLE_VALUE, ffi::INVALID_HANDLE_VALUE)
496 + };
497 +
498 + Ok(WinShmContext {
499 + role: WinShmRole::Server,
500 + mapping,
501 + base,
502 + region_size,
503 + req_event,
504 + resp_event,
505 + profile,
506 + request_offset: req_off,
507 + request_capacity: req_cap,
508 + response_offset: resp_off,
509 + response_capacity: resp_cap,
510 + spin_tries: DEFAULT_SPIN_TRIES,
511 + local_req_seq: 0,
512 + local_resp_seq: 0,
513 + })
514 + }
515 +
516 + /// Attach to an existing per-session Windows SHM region (client side).
517 + pub fn client_attach(
518 + run_dir: &str,
519 + service_name: &str,
520 + auth_token: u64,
521 + session_id: u64,
522 + profile: u32,
523 + ) -> Result<Self, WinShmError> {
524 + validate_service_name(service_name)?;
525 + validate_profile(profile)?;
526 +
527 + let hash = compute_shm_hash(run_dir, service_name, auth_token);
528 + let mapping_name = build_object_name(hash, service_name, profile, session_id, "mapping")?;
529 +
530 + let mapping =
531 + unsafe { open_file_mapping(ffi::FILE_MAP_ALL_ACCESS, 0, mapping_name.as_ptr()) };
532 + if mapping == 0 {
533 + return Err(WinShmError::OpenMapping(last_error()));
534 + }
535 +
536 + let base = unsafe { map_view_of_file(mapping, ffi::FILE_MAP_ALL_ACCESS, 0, 0, 0) };
537 + if base.is_null() {
538 + let e = last_error();
539 + unsafe { ffi::CloseHandle(mapping) };
540 + return Err(WinShmError::MapView(e));
541 + }
542 + let base = base as *mut u8;
543 +
544 + // Acquire barrier
545 + std::sync::atomic::fence(std::sync::atomic::Ordering::Acquire);
546 +
547 + // Validate header
548 + let magic = read_u32(base, OFF_MAGIC);
549 + if magic != REGION_MAGIC {
550 + unsafe {
551 + ffi::UnmapViewOfFile(base as *const _);
552 + ffi::CloseHandle(mapping);
553 + }
554 + return Err(WinShmError::BadMagic);
555 + }
556 +
557 + let version = read_u32(base, OFF_VERSION);
558 + if version != REGION_VERSION {
559 + unsafe {
560 + ffi::UnmapViewOfFile(base as *const _);
561 + ffi::CloseHandle(mapping);
562 + }
563 + return Err(WinShmError::BadVersion);
564 + }
565 +
566 + let hdr_len = read_u32(base, OFF_HEADER_LEN);
567 + if hdr_len != HEADER_LEN {
568 + unsafe {
569 + ffi::UnmapViewOfFile(base as *const _);
570 + ffi::CloseHandle(mapping);
571 + }
572 + return Err(WinShmError::BadHeader);
573 + }
574 +
575 + let hdr_profile = read_u32(base, OFF_PROFILE);
576 + if hdr_profile != profile {
577 + unsafe {
578 + ffi::UnmapViewOfFile(base as *const _);
579 + ffi::CloseHandle(mapping);
580 + }
581 + return Err(WinShmError::BadProfile);
582 + }
583 +
584 + let req_off = read_u32(base, OFF_REQ_OFFSET);
585 + let req_cap = read_u32(base, OFF_REQ_CAPACITY);
586 + let resp_off = read_u32(base, OFF_RESP_OFFSET);
587 + let resp_cap = read_u32(base, OFF_RESP_CAPACITY);
588 + let spin = read_u32(base, OFF_SPIN_TRIES);
589 +
590 + // Validate header fields from shared memory
591 + if req_off == 0
592 + || req_cap == 0
593 + || resp_off == 0
594 + || resp_cap == 0
595 + || req_off % 64 != 0
596 + || req_cap % 64 != 0
597 + || resp_off % 64 != 0
598 + || resp_cap % 64 != 0
599 + || req_off
600 + .checked_add(req_cap)
601 + .map_or(true, |end| resp_off < end)
602 + {
603 + unsafe {
604 + ffi::UnmapViewOfFile(base as *const _);
605 + ffi::CloseHandle(mapping);
606 + }
607 + return Err(WinShmError::BadHeader);
608 + }
609 +
610 + let region_size = (resp_off as usize) + (resp_cap as usize);
611 +
612 + // Read current sequence numbers via interlocked
613 + let cur_req_seq = interlocked_read_i64(base, OFF_REQ_SEQ);
614 + let cur_resp_seq = interlocked_read_i64(base, OFF_RESP_SEQ);
615 +
616 + // Open events for HYBRID
617 + let (req_event, resp_event) = if profile == PROFILE_HYBRID {
618 + let re_name = build_object_name(hash, service_name, profile, session_id, "req_event")?;
619 + let re = unsafe {
620 + open_event(
621 + ffi::EVENT_MODIFY_STATE | ffi::SYNCHRONIZE,
622 + 0,
623 + re_name.as_ptr(),
624 + )
625 + };
626 + if re == 0 {
627 + let e = last_error();
628 + unsafe {
629 + ffi::UnmapViewOfFile(base as *const _);
630 + ffi::CloseHandle(mapping);
631 + }
632 + return Err(WinShmError::OpenEvent(e));
633 + }
634 +
635 + let rsp_name =
636 + build_object_name(hash, service_name, profile, session_id, "resp_event")?;
637 + let rsp = unsafe {
638 + open_event(
639 + ffi::EVENT_MODIFY_STATE | ffi::SYNCHRONIZE,
640 + 0,
641 + rsp_name.as_ptr(),
642 + )
643 + };
644 + if rsp == 0 {
645 + let e = last_error();
646 + unsafe {
647 + ffi::CloseHandle(re);
648 + ffi::UnmapViewOfFile(base as *const _);
649 + ffi::CloseHandle(mapping);
650 + }
651 + return Err(WinShmError::OpenEvent(e));
652 + }
653 + (re, rsp)
654 + } else {
655 + (ffi::INVALID_HANDLE_VALUE, ffi::INVALID_HANDLE_VALUE)
656 + };
657 +
658 + Ok(WinShmContext {
659 + role: WinShmRole::Client,
660 + mapping,
661 + base,
662 + region_size,
663 + req_event,
664 + resp_event,
665 + profile,
666 + request_offset: req_off,
667 + request_capacity: req_cap,
668 + response_offset: resp_off,
669 + response_capacity: resp_cap,
670 + spin_tries: spin,
671 + local_req_seq: cur_req_seq,
672 + local_resp_seq: cur_resp_seq,
673 + })
674 + }
675 +
676 + /// Publish a message. Client sends to request area; server to response.
677 + pub fn send(&mut self, msg: &[u8]) -> Result<(), WinShmError> {
678 + if self.base.is_null() || msg.is_empty() {
679 + return Err(WinShmError::BadParam(
680 + "null context or empty message".into(),
681 + ));
682 + }
683 +
684 + let (area_off, area_cap, len_off, seq_off, peer_waiting_off, peer_event) = match self.role {
685 + WinShmRole::Client => (
686 + self.request_offset,
687 + self.request_capacity,
688 + OFF_REQ_LEN,
689 + OFF_REQ_SEQ,
690 + OFF_REQ_SERVER_WAITING,
691 + self.req_event,
692 + ),
693 + WinShmRole::Server => (
694 + self.response_offset,
695 + self.response_capacity,
696 + OFF_RESP_LEN,
697 + OFF_RESP_SEQ,
698 + OFF_RESP_CLIENT_WAITING,
699 + self.resp_event,
700 + ),
701 + };
702 +
703 + if msg.len() > area_cap as usize || msg.len() > i32::MAX as usize {
704 + return Err(WinShmError::MsgTooLarge);
705 + }
706 +
707 + // 1. Write message data
708 + unsafe {
709 + std::ptr::copy_nonoverlapping(
710 + msg.as_ptr(),
711 + self.base.add(area_off as usize),
712 + msg.len(),
713 + );
714 + }
715 +
716 + // 2. Store message length (interlocked exchange)
717 + interlocked_exchange_i32(self.base, len_off, msg.len() as i32);
718 +
719 + // 3. Increment sequence number
720 + interlocked_increment_i64(self.base, seq_off);
721 +
722 + // 4. If HYBRID and peer waiting, signal event
723 + if self.profile == PROFILE_HYBRID {
724 + let waiting = interlocked_read_i32(self.base, peer_waiting_off);
725 + if waiting != 0 {
726 + unsafe { ffi::SetEvent(peer_event) };
727 + }
728 + }
729 +
730 + match self.role {
731 + WinShmRole::Client => self.local_req_seq += 1,
732 + WinShmRole::Server => self.local_resp_seq += 1,
733 + }
734 +
735 + Ok(())
736 + }
737 +
738 + /// Receive a message into the caller-provided buffer.
739 + pub fn receive(&mut self, buf: &mut [u8], timeout_ms: u32) -> Result<usize, WinShmError> {
740 + if self.base.is_null() {
741 + return Err(WinShmError::BadParam("null context".into()));
742 + }
743 + if buf.is_empty() {
744 + return Err(WinShmError::BadParam("empty buffer".into()));
745 + }
746 +
747 + let (
748 + area_off,
749 + area_cap,
750 + len_off,
751 + seq_off,
752 + self_waiting_off,
753 + peer_closed_off,
754 + wait_event,
755 + expected_seq,
756 + ) = match self.role {
757 + WinShmRole::Server => (
758 + self.request_offset,
759 + self.request_capacity,
760 + OFF_REQ_LEN,
761 + OFF_REQ_SEQ,
762 + OFF_REQ_SERVER_WAITING,
763 + OFF_REQ_CLIENT_CLOSED,
764 + self.req_event,
765 + self.local_req_seq + 1,
766 + ),
767 + WinShmRole::Client => (
768 + self.response_offset,
769 + self.response_capacity,
770 + OFF_RESP_LEN,
771 + OFF_RESP_SEQ,
772 + OFF_RESP_CLIENT_WAITING,
773 + OFF_RESP_SERVER_CLOSED,
774 + self.resp_event,
775 + self.local_resp_seq + 1,
776 + ),
777 + };
778 +
779 + // The copy ceiling is the smaller of the caller buffer and the
780 + // SHM area capacity. Prevents out-of-bounds reads on forged lengths.
781 + let max_copy = std::cmp::min(buf.len(), area_cap as usize);
782 +
783 + // Phase 1: spin
784 + let mut observed = false;
785 + let mut mlen: i32 = 0;
786 + for _ in 0..self.spin_tries {
787 + let cur = interlocked_read_i64(self.base, seq_off);
788 + if cur >= expected_seq {
789 + mlen = interlocked_read_i32(self.base, len_off);
790 + if mlen > 0 && (mlen as usize) <= max_copy {
791 + unsafe {
792 + std::ptr::copy_nonoverlapping(
793 + self.base.add(area_off as usize),
794 + buf.as_mut_ptr(),
795 + mlen as usize,
796 + );
797 + }
798 + }
799 + observed = true;
800 + break;
801 + }
802 + cpu_relax();
803 + }
804 +
805 + // Phase 2: kernel wait or busy-wait (deadline-based retry for
806 + // spurious wakes — same pattern as POSIX SHM Phase H6 fix).
807 + if !observed {
808 + if self.profile == PROFILE_HYBRID {
809 + let deadline_ms = if timeout_ms == 0 {
810 + ffi::INFINITE
811 + } else {
812 + timeout_ms
813 + };
814 + let start_tick = unsafe { ffi::GetTickCount64() };
815 +
816 + loop {
817 + interlocked_exchange_i32(self.base, self_waiting_off, 1);
818 + std::sync::atomic::fence(std::sync::atomic::Ordering::SeqCst);
819 +
820 + let cur = interlocked_read_i64(self.base, seq_off);
821 + if cur >= expected_seq {
822 + interlocked_exchange_i32(self.base, self_waiting_off, 0);
823 + break; // data available
824 + }
825 +
826 + // Compute remaining wait time
827 + let wait_ms = if deadline_ms == ffi::INFINITE {
828 + ffi::INFINITE
829 + } else {
830 + let elapsed = (unsafe { ffi::GetTickCount64() } - start_tick) as u32;
831 + if elapsed >= deadline_ms {
832 + interlocked_exchange_i32(self.base, self_waiting_off, 0);
833 + return Err(WinShmError::Timeout);
834 + }
835 + deadline_ms - elapsed
836 + };
837 +
838 + let ret = unsafe { ffi::WaitForSingleObject(wait_event, wait_ms) };
839 + interlocked_exchange_i32(self.base, self_waiting_off, 0);
840 +
841 + // Check sequence — data may have arrived
842 + let cur = interlocked_read_i64(self.base, seq_off);
843 + if cur >= expected_seq {
844 + break; // data available
845 + }
846 +
847 + // No data — check peer close
848 + if interlocked_read_i32(self.base, peer_closed_off) != 0 {
849 + let cur = interlocked_read_i64(self.base, seq_off);
850 + if cur >= expected_seq {
851 + break;
852 + }
853 + self.advance_seq(expected_seq);
854 + return Err(WinShmError::Disconnected);
855 + }
856 +
857 + // Actual timeout (not spurious)
858 + if ret == ffi::WAIT_TIMEOUT {
859 + return Err(WinShmError::Timeout);
860 + }
861 +
862 + // Spurious wake — retry with remaining deadline
863 + }
864 +
865 + // Copy after waking
866 + mlen = interlocked_read_i32(self.base, len_off);
867 + if mlen > 0 && (mlen as usize) <= max_copy {
868 + unsafe {
869 + std::ptr::copy_nonoverlapping(
870 + self.base.add(area_off as usize),
871 + buf.as_mut_ptr(),
872 + mlen as usize,
873 + );
874 + }
875 + }
876 + } else {
877 + // BUSYWAIT
878 + let start = unsafe { ffi::GetTickCount() };
879 + loop {
880 + let cur = interlocked_read_i64(self.base, seq_off);
881 + if cur >= expected_seq {
882 + mlen = interlocked_read_i32(self.base, len_off);
883 + if mlen > 0 && (mlen as usize) <= max_copy {
884 + unsafe {
885 + std::ptr::copy_nonoverlapping(
886 + self.base.add(area_off as usize),
887 + buf.as_mut_ptr(),
888 + mlen as usize,
889 + );
890 + }
891 + }
892 + break;
893 + }
894 +
895 + if timeout_ms > 0 {
896 + let elapsed = unsafe { ffi::GetTickCount() }.wrapping_sub(start);
897 + if elapsed >= timeout_ms {
898 + return Err(WinShmError::Timeout);
899 + }
900 + }
901 +
902 + if interlocked_read_i32(self.base, peer_closed_off) != 0 {
903 + let cur = interlocked_read_i64(self.base, seq_off);
904 + if cur >= expected_seq {
905 + mlen = interlocked_read_i32(self.base, len_off);
906 + if mlen > 0 && (mlen as usize) <= max_copy {
907 + unsafe {
908 + std::ptr::copy_nonoverlapping(
909 + self.base.add(area_off as usize),
910 + buf.as_mut_ptr(),
911 + mlen as usize,
912 + );
913 + }
914 + }
915 + break;
916 + }
917 + self.advance_seq(expected_seq);
918 + return Err(WinShmError::Disconnected);
919 + }
920 +
921 + cpu_relax();
922 + }
923 + }
924 + }
925 +
926 + self.advance_seq(expected_seq);
927 +
928 + // mlen==0 after sequence advance indicates SHM corruption (send rejects 0-length)
929 + if mlen == 0 {
930 + return Err(WinShmError::BadHeader);
931 + }
932 +
933 + if (mlen as usize) > max_copy {
934 + return Err(WinShmError::MsgTooLarge);
935 + }
936 +
937 + Ok(mlen as usize)
938 + }
939 +
940 + fn advance_seq(&mut self, expected_seq: i64) {
941 + match self.role {
942 + WinShmRole::Server => self.local_req_seq = expected_seq,
943 + WinShmRole::Client => self.local_resp_seq = expected_seq,
944 + }
945 + }
946 +
947 + /// Close client (unmap, close handles, set close flag).
948 + pub fn close(&mut self) {
949 + if !self.base.is_null() {
950 + // Set close flag
951 + if self.role == WinShmRole::Client {
952 + interlocked_exchange_i32(self.base, OFF_REQ_CLIENT_CLOSED, 1);
953 + }
954 + std::sync::atomic::fence(std::sync::atomic::Ordering::SeqCst);
955 + }
956 +
957 + if self.profile == PROFILE_HYBRID && self.req_event != ffi::INVALID_HANDLE_VALUE {
958 + if self.role == WinShmRole::Client {
959 + unsafe { ffi::SetEvent(self.req_event) };
960 + }
961 + }
962 +
963 + self.cleanup_handles();
964 + }
965 +
966 + /// Destroy server (set close flag, signal, unmap, close handles).
967 + pub fn destroy(&mut self) {
968 + if !self.base.is_null() {
969 + interlocked_exchange_i32(self.base, OFF_RESP_SERVER_CLOSED, 1);
970 + std::sync::atomic::fence(std::sync::atomic::Ordering::SeqCst);
971 + }
972 +
973 + if self.profile == PROFILE_HYBRID && self.resp_event != ffi::INVALID_HANDLE_VALUE {
974 + unsafe { ffi::SetEvent(self.resp_event) };
975 + }
976 +
977 + self.cleanup_handles();
978 + }
979 +
980 + fn cleanup_handles(&mut self) {
981 + if !self.base.is_null() {
982 + unsafe { ffi::UnmapViewOfFile(self.base as *const _) };
983 + self.base = std::ptr::null_mut();
984 + }
985 + if self.mapping != 0 {
986 + unsafe { ffi::CloseHandle(self.mapping) };
987 + self.mapping = 0;
988 + }
989 + if self.req_event != ffi::INVALID_HANDLE_VALUE {
990 + unsafe { ffi::CloseHandle(self.req_event) };
991 + self.req_event = ffi::INVALID_HANDLE_VALUE;
992 + }
993 + if self.resp_event != ffi::INVALID_HANDLE_VALUE {
994 + unsafe { ffi::CloseHandle(self.resp_event) };
995 + self.resp_event = ffi::INVALID_HANDLE_VALUE;
996 + }
997 + self.region_size = 0;
998 + }
999 +}
1000 +
1001 +#[cfg(windows)]
1002 +impl Drop for WinShmContext {
1003 + fn drop(&mut self) {
1004 + match self.role {
1005 + WinShmRole::Server => self.destroy(),
1006 + WinShmRole::Client => self.close(),
1007 + }
1008 + }
1009 +}
1010 +
1011 +// ---------------------------------------------------------------------------
1012 +// Internal helpers
1013 +// ---------------------------------------------------------------------------
1014 +
1015 +fn align_cacheline(v: u32) -> u32 {
1016 + (v + (CACHELINE_SIZE - 1)) & !(CACHELINE_SIZE - 1)
1017 +}
1018 +
1019 +fn validate_service_name(name: &str) -> Result<(), WinShmError> {
1020 + if name.is_empty() {
1021 + return Err(WinShmError::BadParam("empty service name".into()));
1022 + }
1023 + if name == "." || name == ".." {
1024 + return Err(WinShmError::BadParam(
1025 + "service name cannot be '.' or '..'".into(),
1026 + ));
1027 + }
1028 + for c in name.bytes() {
1029 + match c {
1030 + b'a'..=b'z' | b'A'..=b'Z' | b'0'..=b'9' | b'.' | b'_' | b'-' => {}
1031 + _ => {
1032 + return Err(WinShmError::BadParam(format!(
1033 + "service name contains invalid character: {:?}",
1034 + c as char,
1035 + )))
1036 + }
1037 + }
1038 + }
1039 + Ok(())
1040 +}
1041 +
1042 +fn validate_profile(profile: u32) -> Result<(), WinShmError> {
1043 + if profile != PROFILE_HYBRID && profile != PROFILE_BUSYWAIT {
1044 + return Err(WinShmError::BadParam(format!("invalid profile: {profile}")));
1045 + }
1046 + Ok(())
1047 +}
1048 +
1049 +/// FNV-1a 64-bit hash.
1050 +pub fn fnv1a_64(data: &[u8]) -> u64 {
1051 + let mut hash = FNV1A_OFFSET_BASIS;
1052 + for &b in data {
1053 + hash ^= b as u64;
1054 + hash = hash.wrapping_mul(FNV1A_PRIME);
1055 + }
1056 + hash
1057 +}
1058 +
1059 +fn compute_shm_hash(run_dir: &str, service_name: &str, auth_token: u64) -> u64 {
1060 + let input = format!("{}\n{}\n{}", run_dir, service_name, auth_token);
1061 + fnv1a_64(input.as_bytes())
1062 +}
1063 +
1064 +fn build_object_name(
1065 + hash: u64,
1066 + service_name: &str,
1067 + profile: u32,
1068 + session_id: u64,
1069 + suffix: &str,
1070 +) -> Result<Vec<u16>, WinShmError> {
1071 + let narrow = format!(
1072 + "Local\\netipc-{:016x}-{}-p{}-s{:016x}-{}",
1073 + hash, service_name, profile, session_id, suffix
1074 + );
1075 + if narrow.len() >= 256 {
1076 + return Err(WinShmError::BadParam("object name too long".into()));
1077 + }
1078 + let mut wide: Vec<u16> = narrow.encode_utf16().collect();
1079 + wide.push(0);
1080 + Ok(wide)
1081 +}
1082 +
1083 +#[cfg(windows)]
1084 +fn last_error() -> u32 {
1085 + unsafe { ffi::GetLastError() }
1086 +}
1087 +
1088 +fn read_u32(base: *mut u8, offset: usize) -> u32 {
1089 + unsafe { std::ptr::read_unaligned(base.add(offset) as *const u32) }
1090 +}
1091 +
1092 +fn write_u32(base: *mut u8, offset: usize, val: u32) {
1093 + unsafe { std::ptr::write_unaligned(base.add(offset) as *mut u32, val) };
1094 +}
1095 +
1096 +#[cfg(windows)]
1097 +fn interlocked_read_i32(base: *mut u8, offset: usize) -> i32 {
1098 + use std::sync::atomic::{AtomicI32, Ordering};
1099 + unsafe {
1100 + let ptr = base.add(offset) as *const AtomicI32;
1101 + (*ptr).load(Ordering::Acquire)
1102 + }
1103 +}
1104 +
1105 +#[cfg(windows)]
1106 +fn interlocked_exchange_i32(base: *mut u8, offset: usize, val: i32) {
1107 + use std::sync::atomic::{AtomicI32, Ordering};
1108 + unsafe {
1109 + let ptr = base.add(offset) as *const AtomicI32;
1110 + (*ptr).store(val, Ordering::Release);
1111 + }
1112 +}
1113 +
1114 +#[cfg(windows)]
1115 +fn interlocked_read_i64(base: *mut u8, offset: usize) -> i64 {
1116 + use std::sync::atomic::{AtomicI64, Ordering};
1117 + unsafe {
1118 + let ptr = base.add(offset) as *const AtomicI64;
1119 + (*ptr).load(Ordering::Acquire)
1120 + }
1121 +}
1122 +
1123 +#[cfg(windows)]
1124 +fn interlocked_increment_i64(base: *mut u8, offset: usize) {
1125 + use std::sync::atomic::{AtomicI64, Ordering};
1126 + unsafe {
1127 + let ptr = base.add(offset) as *const AtomicI64;
1128 + (*ptr).fetch_add(1, Ordering::Release);
1129 + }
1130 +}
1131 +
1132 +#[cfg(windows)]
1133 +#[inline]
1134 +fn cpu_relax() {
1135 + #[cfg(any(target_arch = "x86_64", target_arch = "x86"))]
1136 + unsafe {
1137 + std::arch::asm!("pause", options(nomem, nostack));
1138 + }
1139 + #[cfg(target_arch = "aarch64")]
1140 + unsafe {
1141 + std::arch::asm!("yield", options(nomem, nostack));
1142 + }
1143 + #[cfg(not(any(target_arch = "x86_64", target_arch = "x86", target_arch = "aarch64")))]
1144 + {
1145 + std::sync::atomic::fence(std::sync::atomic::Ordering::SeqCst);
1146 + }
1147 +}
1148 +
1149 +// ---------------------------------------------------------------------------
1150 +// Tests (non-Windows: compile-check only)
1151 +// ---------------------------------------------------------------------------
1152 +
1153 +#[cfg(test)]
1154 +mod tests {
1155 + use super::*;
1156 +
1157 + #[cfg(windows)]
1158 + use std::sync::atomic::{AtomicU64, Ordering};
1159 + #[cfg(windows)]
1160 + use std::thread;
1161 + #[cfg(windows)]
1162 + use std::time::Duration;
1163 +
1164 + #[cfg(windows)]
1165 + static WIN_SHM_TEST_COUNTER: AtomicU64 = AtomicU64::new(0);
1166 +
1167 + #[cfg(windows)]
1168 + fn test_run_dir() -> String {
1169 + let dir = std::env::temp_dir().join("nipc_win_shm_rust_test");
1170 + let _ = std::fs::create_dir_all(&dir);
1171 + dir.display().to_string()
1172 + }
1173 +
1174 + #[cfg(windows)]
1175 + fn unique_service(prefix: &str) -> String {
1176 + format!(
1177 + "{}_{}_{}",
1178 + prefix,
1179 + std::process::id(),
1180 + WIN_SHM_TEST_COUNTER.fetch_add(1, Ordering::Relaxed) + 1
1181 + )
1182 + }
1183 +
1184 + #[cfg(windows)]
1185 + struct FaultHookGuard;
1186 +
1187 + #[cfg(windows)]
1188 + impl FaultHookGuard {
1189 + fn install(site: WinShmFaultSite, error: u32) -> Self {
1190 + install_fault_hook(site, error, 0);
1191 + Self
1192 + }
1193 +
1194 + fn install_after(site: WinShmFaultSite, error: u32, skip_matches: u32) -> Self {
1195 + install_fault_hook(site, error, skip_matches);
1196 + Self
1197 + }
1198 + }
1199 +
1200 + #[cfg(windows)]
1201 + impl Drop for FaultHookGuard {
1202 + fn drop(&mut self) {
1203 + clear_fault_hook();
1204 + }
1205 + }
1206 +
1207 + #[test]
1208 + fn test_fnv1a_64_deterministic() {
1209 + let h1 = fnv1a_64(b"C:\\Temp\\netdata\ncgroups-snapshot\n12345");
1210 + let h2 = fnv1a_64(b"C:\\Temp\\netdata\ncgroups-snapshot\n12345");
1211 + assert_eq!(h1, h2);
1212 + assert_ne!(h1, 0);
1213 + }
1214 +
1215 + #[test]
1216 + fn test_fnv1a_64_different_tokens() {
1217 + let h1 = fnv1a_64(b"dir\nsvc\n100");
1218 + let h2 = fnv1a_64(b"dir\nsvc\n200");
1219 + assert_ne!(h1, h2);
1220 + }
1221 +
1222 + #[test]
1223 + fn test_validate_service_name() {
1224 + assert!(validate_service_name("cgroups-snapshot").is_ok());
1225 + assert!(validate_service_name("test.v1").is_ok());
1226 + assert!(validate_service_name("").is_err());
1227 + assert!(validate_service_name(".").is_err());
1228 + assert!(validate_service_name("..").is_err());
1229 + assert!(validate_service_name("has/slash").is_err());
1230 + assert!(validate_service_name("has space").is_err());
1231 + }
1232 +
1233 + #[test]
1234 + fn test_align_cacheline() {
1235 + assert_eq!(align_cacheline(0), 0);
1236 + assert_eq!(align_cacheline(1), 64);
1237 + assert_eq!(align_cacheline(64), 64);
1238 + assert_eq!(align_cacheline(65), 128);
1239 + assert_eq!(align_cacheline(128), 128);
1240 + }
1241 +
1242 + #[test]
1243 + fn test_object_name_format() {
1244 + let name = build_object_name(0xDEADBEEF, "test-svc", 2, 0x42, "mapping").unwrap();
1245 + // Check it's NUL-terminated
1246 + assert_eq!(*name.last().unwrap(), 0);
1247 + let narrow: String = name[..name.len() - 1]
1248 + .iter()
1249 + .map(|&c| c as u8 as char)
1250 + .collect();
1251 + assert!(narrow.starts_with("Local\\netipc-"));
1252 + assert!(narrow.contains("-test-svc-p2-s0000000000000042-mapping"));
1253 + }
1254 +
1255 + #[test]
1256 + fn test_error_display_variants() {
1257 + let cases = [
1258 + (WinShmError::BadParam("oops".into()), "bad parameter: oops"),
1259 + (
1260 + WinShmError::CreateMapping(5),
1261 + "CreateFileMappingW failed: 5",
1262 + ),
1263 + (WinShmError::OpenMapping(6), "OpenFileMappingW failed: 6"),
1264 + (WinShmError::MapView(7), "MapViewOfFile failed: 7"),
1265 + (WinShmError::CreateEvent(8), "CreateEventW failed: 8"),
1266 + (WinShmError::OpenEvent(9), "OpenEventW failed: 9"),
1267 + (
1268 + WinShmError::AddrInUse,
1269 + "Windows SHM object name already in use by live server",
1270 + ),
1271 + (WinShmError::BadMagic, "SHM header magic mismatch"),
1272 + (WinShmError::BadVersion, "SHM header version mismatch"),
1273 + (WinShmError::BadHeader, "SHM header_len mismatch"),
1274 + (WinShmError::BadProfile, "SHM profile mismatch"),
1275 + (
1276 + WinShmError::MsgTooLarge,
1277 + "message exceeds SHM area capacity",
1278 + ),
1279 + (WinShmError::Timeout, "SHM wait timed out"),
1280 + (WinShmError::Disconnected, "peer closed SHM session"),
1281 + ];
1282 +
1283 + for (err, expected) in cases {
1284 + assert_eq!(err.to_string(), expected);
1285 + }
1286 + }
1287 +
1288 + #[test]
1289 + fn test_build_object_name_too_long() {
1290 + let long_service = "a".repeat(240);
1291 + let err = build_object_name(0xDEADBEEF, &long_service, PROFILE_HYBRID, 1, "mapping")
1292 + .expect_err("object name should be rejected");
1293 + assert_eq!(err, WinShmError::BadParam("object name too long".into()));
1294 + }
1295 +
1296 + #[test]
1297 + fn test_validate_profile_rejects_invalid() {
1298 + let err = validate_profile(0).expect_err("invalid profile should fail");
1299 + assert_eq!(err, WinShmError::BadParam("invalid profile: 0".into()));
1300 + }
1301 +
1302 + #[cfg(windows)]
1303 + #[test]
1304 + fn test_client_attach_bad_magic_windows() {
1305 + let run_dir = test_run_dir();
1306 + let service = unique_service("rs_win_shm_bad_magic");
1307 + let auth_token: u64 = 0x123456;
1308 + let session_id: u64 = 15;
1309 +
1310 + let mut server = WinShmContext::server_create(
1311 + &run_dir,
1312 + &service,
1313 + auth_token,
1314 + session_id,
1315 + PROFILE_HYBRID,
1316 + 4096,
1317 + 4096,
1318 + )
1319 + .expect("server_create");
1320 +
1321 + write_u32(server.base, OFF_MAGIC, 0);
1322 + let err = match WinShmContext::client_attach(
1323 + &run_dir,
1324 + &service,
1325 + auth_token,
1326 + session_id,
1327 + PROFILE_HYBRID,
1328 + ) {
1329 + Ok(_) => panic!("bad magic must fail"),
1330 + Err(err) => err,
1331 + };
1332 + assert_eq!(err, WinShmError::BadMagic);
1333 +
1334 + server.destroy();
1335 + }
1336 +
1337 + #[cfg(windows)]
1338 + #[test]
1339 + fn test_server_create_rejects_existing_objects_windows() {
1340 + let run_dir = test_run_dir();
1341 + let service = unique_service("rs_win_shm_addr_in_use");
1342 + let auth_token: u64 = 0x424242;
1343 + let session_id: u64 = 22;
1344 +
1345 + let mut server = WinShmContext::server_create(
1346 + &run_dir,
1347 + &service,
1348 + auth_token,
1349 + session_id,
1350 + PROFILE_HYBRID,
1351 + 4096,
1352 + 4096,
1353 + )
1354 + .expect("server_create");
1355 +
1356 + let err = match WinShmContext::server_create(
1357 + &run_dir,
1358 + &service,
1359 + auth_token,
1360 + session_id,
1361 + PROFILE_HYBRID,
1362 + 4096,
1363 + 4096,
1364 + ) {
1365 + Ok(_) => panic!("duplicate server_create must fail"),
1366 + Err(err) => err,
1367 + };
1368 + assert_eq!(err, WinShmError::AddrInUse);
1369 +
1370 + server.destroy();
1371 + }
1372 +
1373 + #[cfg(windows)]
1374 + #[test]
1375 + fn test_role_send_and_receive_guards_windows() {
1376 + let run_dir = test_run_dir();
1377 + let service = unique_service("rs_win_shm_guards");
1378 + let auth_token: u64 = 0xabcdef;
1379 + let session_id: u64 = 31;
1380 +
1381 + let mut server = WinShmContext::server_create(
1382 + &run_dir,
1383 + &service,
1384 + auth_token,
1385 + session_id,
1386 + PROFILE_HYBRID,
1387 + 128,
1388 + 128,
1389 + )
1390 + .expect("server_create");
1391 + assert_eq!(server.role(), WinShmRole::Server);
1392 +
1393 + let mut client = WinShmContext::client_attach(
1394 + &run_dir,
1395 + &service,
1396 + auth_token,
1397 + session_id,
1398 + PROFILE_HYBRID,
1399 + )
1400 + .expect("client_attach");
1401 + assert_eq!(client.role(), WinShmRole::Client);
1402 +
1403 + let empty_send = client.send(&[]).expect_err("empty send must fail");
1404 + assert_eq!(
1405 + empty_send,
1406 + WinShmError::BadParam("null context or empty message".into())
1407 + );
1408 +
1409 + let oversize = vec![0u8; 256];
1410 + let oversize_err = client.send(&oversize).expect_err("oversize send must fail");
1411 + assert_eq!(oversize_err, WinShmError::MsgTooLarge);
1412 +
1413 + let empty_buf_err = server
1414 + .receive(&mut [], 10)
1415 + .expect_err("empty receive buffer must fail");
1416 + assert_eq!(empty_buf_err, WinShmError::BadParam("empty buffer".into()));
1417 +
1418 + client.close();
1419 +
1420 + let mut buf = [0u8; 16];
1421 + let null_ctx_err = client
1422 + .receive(&mut buf, 10)
1423 + .expect_err("closed client must reject receive");
1424 + assert_eq!(null_ctx_err, WinShmError::BadParam("null context".into()));
1425 +
1426 + server.destroy();
1427 + }
1428 +
1429 + #[cfg(windows)]
1430 + #[test]
1431 + fn test_client_attach_bad_version_windows() {
1432 + let run_dir = test_run_dir();
1433 + let service = unique_service("rs_win_shm_bad_version");
1434 + let auth_token: u64 = 0x123457;
1435 + let session_id: u64 = 16;
1436 +
1437 + let mut server = WinShmContext::server_create(
1438 + &run_dir,
1439 + &service,
1440 + auth_token,
1441 + session_id,
1442 + PROFILE_HYBRID,
1443 + 4096,
1444 + 4096,
1445 + )
1446 + .expect("server_create");
1447 +
1448 + write_u32(server.base, OFF_VERSION, REGION_VERSION + 1);
1449 + let err = match WinShmContext::client_attach(
1450 + &run_dir,
1451 + &service,
1452 + auth_token,
1453 + session_id,
1454 + PROFILE_HYBRID,
1455 + ) {
1456 + Ok(_) => panic!("bad version must fail"),
1457 + Err(err) => err,
1458 + };
1459 + assert_eq!(err, WinShmError::BadVersion);
1460 +
1461 + server.destroy();
1462 + }
1463 +
1464 + #[cfg(windows)]
1465 + #[test]
1466 + fn test_client_attach_bad_header_windows() {
1467 + let run_dir = test_run_dir();
1468 + let service = unique_service("rs_win_shm_bad_header");
1469 + let auth_token: u64 = 0x123458;
1470 + let session_id: u64 = 17;
1471 +
1472 + let mut server = WinShmContext::server_create(
1473 + &run_dir,
1474 + &service,
1475 + auth_token,
1476 + session_id,
1477 + PROFILE_HYBRID,
1478 + 4096,
1479 + 4096,
1480 + )
1481 + .expect("server_create");
1482 +
1483 + write_u32(server.base, OFF_HEADER_LEN, HEADER_LEN + 64);
1484 + let err = match WinShmContext::client_attach(
1485 + &run_dir,
1486 + &service,
1487 + auth_token,
1488 + session_id,
1489 + PROFILE_HYBRID,
1490 + ) {
1491 + Ok(_) => panic!("bad header_len must fail"),
1492 + Err(err) => err,
1493 + };
1494 + assert_eq!(err, WinShmError::BadHeader);
1495 +
1496 + server.destroy();
1497 + }
1498 +
1499 + #[cfg(windows)]
1500 + #[test]
1501 + fn test_client_attach_bad_profile_windows() {
1502 + let run_dir = test_run_dir();
1503 + let service = unique_service("rs_win_shm_bad_profile");
1504 + let auth_token: u64 = 0x123459;
1505 + let session_id: u64 = 18;
1506 +
1507 + let mut server = WinShmContext::server_create(
1508 + &run_dir,
1509 + &service,
1510 + auth_token,
1511 + session_id,
1512 + PROFILE_HYBRID,
1513 + 4096,
1514 + 4096,
1515 + )
1516 + .expect("server_create");
1517 +
1518 + write_u32(server.base, OFF_PROFILE, PROFILE_BUSYWAIT);
1519 + let err = match WinShmContext::client_attach(
1520 + &run_dir,
1521 + &service,
1522 + auth_token,
1523 + session_id,
1524 + PROFILE_HYBRID,
1525 + ) {
1526 + Ok(_) => panic!("bad profile must fail"),
1527 + Err(err) => err,
1528 + };
1529 + assert_eq!(err, WinShmError::BadProfile);
1530 +
1531 + server.destroy();
1532 + }
1533 +
1534 + #[cfg(windows)]
1535 + #[test]
1536 + fn test_receive_timeout_hybrid_windows() {
1537 + let run_dir = test_run_dir();
1538 + let service = unique_service("rs_win_shm_timeout_h");
1539 + let auth_token: u64 = 0x8912;
1540 + let session_id: u64 = 19;
1541 +
1542 + let mut server = WinShmContext::server_create(
1543 + &run_dir,
1544 + &service,
1545 + auth_token,
1546 + session_id,
1547 + PROFILE_HYBRID,
1548 + 4096,
1549 + 4096,
1550 + )
1551 + .expect("server_create");
1552 + let mut client = WinShmContext::client_attach(
1553 + &run_dir,
1554 + &service,
1555 + auth_token,
1556 + session_id,
1557 + PROFILE_HYBRID,
1558 + )
1559 + .expect("client_attach");
1560 +
1561 + let mut buf = [0u8; 128];
1562 + assert_eq!(server.receive(&mut buf, 10), Err(WinShmError::Timeout));
1563 +
1564 + client.close();
1565 + server.destroy();
1566 + }
1567 +
1568 + #[cfg(windows)]
1569 + #[test]
1570 + fn test_receive_timeout_busywait_windows() {
1571 + let run_dir = test_run_dir();
1572 + let service = unique_service("rs_win_shm_timeout_b");
1573 + let auth_token: u64 = 0x8913;
1574 + let session_id: u64 = 20;
1575 +
1576 + let mut server = WinShmContext::server_create(
1577 + &run_dir,
1578 + &service,
1579 + auth_token,
1580 + session_id,
1581 + PROFILE_BUSYWAIT,
1582 + 4096,
1583 + 4096,
1584 + )
1585 + .expect("server_create");
1586 + let mut client = WinShmContext::client_attach(
1587 + &run_dir,
1588 + &service,
1589 + auth_token,
1590 + session_id,
1591 + PROFILE_BUSYWAIT,
1592 + )
1593 + .expect("client_attach");
1594 +
1595 + let mut buf = [0u8; 128];
1596 + assert_eq!(server.receive(&mut buf, 10), Err(WinShmError::Timeout));
1597 +
1598 + client.close();
1599 + server.destroy();
1600 + }
1601 +
1602 + #[cfg(windows)]
1603 + #[test]
1604 + fn test_receive_detects_peer_closed_windows() {
1605 + let run_dir = test_run_dir();
1606 + let service = unique_service("rs_win_shm_peer_closed");
1607 + let auth_token: u64 = 0x789a;
1608 + let session_id: u64 = 21;
1609 +
1610 + let server = WinShmContext::server_create(
1611 + &run_dir,
1612 + &service,
1613 + auth_token,
1614 + session_id,
1615 + PROFILE_HYBRID,
1616 + 4096,
1617 + 4096,
1618 + )
1619 + .expect("server_create");
1620 + let mut client = WinShmContext::client_attach(
1621 + &run_dir,
1622 + &service,
1623 + auth_token,
1624 + session_id,
1625 + PROFILE_HYBRID,
1626 + )
1627 + .expect("client_attach");
1628 +
1629 + let (sender, receiver) = std::sync::mpsc::channel();
1630 +
1631 + let server_ptr = std::sync::Arc::new(std::sync::Mutex::new(server));
1632 + let recv_server = server_ptr.clone();
1633 + let handle = thread::spawn(move || {
1634 + let mut buf = [0u8; 128];
1635 + let mut guard = recv_server.lock().unwrap();
1636 + let recv = guard.receive(&mut buf, 1000);
1637 + sender.send(recv).unwrap();
1638 + });
1639 +
1640 + thread::sleep(Duration::from_millis(20));
1641 + client.close();
1642 +
1643 + let recv = receiver
1644 + .recv_timeout(Duration::from_secs(2))
1645 + .expect("receive result");
1646 + assert_eq!(recv, Err(WinShmError::Disconnected));
1647 +
1648 + handle.join().expect("receiver thread");
1649 + let mutex = match std::sync::Arc::try_unwrap(server_ptr) {
1650 + Ok(mutex) => mutex,
1651 + Err(_) => panic!("receiver thread still holds the SHM server"),
1652 + };
1653 + let mut server = mutex.into_inner().expect("mutex");
1654 + server.destroy();
1655 + }
1656 +
1657 + #[cfg(windows)]
1658 + #[test]
1659 + fn test_server_create_fault_create_mapping_windows() {
1660 + let run_dir = test_run_dir();
1661 + let service = unique_service("rs_win_shm_fault_create_mapping");
1662 + let auth_token: u64 = 0x5511;
1663 + let session_id: u64 = 41;
1664 +
1665 + let _fault = FaultHookGuard::install(WinShmFaultSite::CreateFileMapping, 5);
1666 + let err = match WinShmContext::server_create(
1667 + &run_dir,
1668 + &service,
1669 + auth_token,
1670 + session_id,
1671 + PROFILE_HYBRID,
1672 + 4096,
1673 + 4096,
1674 + ) {
1675 + Ok(_) => panic!("CreateFileMappingW fault must fail"),
1676 + Err(err) => err,
1677 + };
1678 + assert_eq!(err, WinShmError::CreateMapping(5));
1679 + drop(_fault);
1680 +
1681 + let mut server = WinShmContext::server_create(
1682 + &run_dir,
1683 + &service,
1684 + auth_token,
1685 + session_id,
1686 + PROFILE_HYBRID,
1687 + 4096,
1688 + 4096,
1689 + )
1690 + .expect("server_create after CreateFileMappingW fault");
1691 + server.destroy();
1692 + }
1693 +
1694 + #[cfg(windows)]
1695 + #[test]
1696 + fn test_server_create_fault_map_view_releases_mapping_windows() {
1697 + let run_dir = test_run_dir();
1698 + let service = unique_service("rs_win_shm_fault_map_view");
1699 + let auth_token: u64 = 0x5512;
1700 + let session_id: u64 = 42;
1701 +
1702 + let _fault = FaultHookGuard::install(WinShmFaultSite::MapViewOfFile, 6);
1703 + let err = match WinShmContext::server_create(
1704 + &run_dir,
1705 + &service,
1706 + auth_token,
1707 + session_id,
1708 + PROFILE_HYBRID,
1709 + 4096,
1710 + 4096,
1711 + ) {
1712 + Ok(_) => panic!("MapViewOfFile fault must fail"),
1713 + Err(err) => err,
1714 + };
1715 + assert_eq!(err, WinShmError::MapView(6));
1716 + drop(_fault);
1717 +
1718 + let mut server = WinShmContext::server_create(
1719 + &run_dir,
1720 + &service,
1721 + auth_token,
1722 + session_id,
1723 + PROFILE_HYBRID,
1724 + 4096,
1725 + 4096,
1726 + )
1727 + .expect("server_create after MapViewOfFile fault");
1728 + server.destroy();
1729 + }
1730 +
1731 + #[cfg(windows)]
1732 + #[test]
1733 + fn test_server_create_fault_req_event_releases_mapping_windows() {
1734 + let run_dir = test_run_dir();
1735 + let service = unique_service("rs_win_shm_fault_req_event");
1736 + let auth_token: u64 = 0x5513;
1737 + let session_id: u64 = 43;
1738 +
1739 + let _fault = FaultHookGuard::install(WinShmFaultSite::CreateEvent, 7);
1740 + let err = match WinShmContext::server_create(
1741 + &run_dir,
1742 + &service,
1743 + auth_token,
1744 + session_id,
1745 + PROFILE_HYBRID,
1746 + 4096,
1747 + 4096,
1748 + ) {
1749 + Ok(_) => panic!("req_event CreateEventW fault must fail"),
1750 + Err(err) => err,
1751 + };
1752 + assert_eq!(err, WinShmError::CreateEvent(7));
1753 + drop(_fault);
1754 +
1755 + let mut server = WinShmContext::server_create(
1756 + &run_dir,
1757 + &service,
1758 + auth_token,
1759 + session_id,
1760 + PROFILE_HYBRID,
1761 + 4096,
1762 + 4096,
1763 + )
1764 + .expect("server_create after req_event fault");
1765 + server.destroy();
1766 + }
1767 +
1768 + #[cfg(windows)]
1769 + #[test]
1770 + fn test_server_create_fault_resp_event_releases_partial_objects_windows() {
1771 + let run_dir = test_run_dir();
1772 + let service = unique_service("rs_win_shm_fault_resp_event");
1773 + let auth_token: u64 = 0x5514;
1774 + let session_id: u64 = 44;
1775 +
1776 + let _fault = FaultHookGuard::install_after(WinShmFaultSite::CreateEvent, 8, 1);
1777 + let err = match WinShmContext::server_create(
1778 + &run_dir,
1779 + &service,
1780 + auth_token,
1781 + session_id,
1782 + PROFILE_HYBRID,
1783 + 4096,
1784 + 4096,
1785 + ) {
1786 + Ok(_) => panic!("second CreateEventW fault must fail"),
1787 + Err(err) => err,
1788 + };
1789 + assert_eq!(err, WinShmError::CreateEvent(8));
1790 + drop(_fault);
1791 +
1792 + let mut server = WinShmContext::server_create(
1793 + &run_dir,
1794 + &service,
1795 + auth_token,
1796 + session_id,
1797 + PROFILE_HYBRID,
1798 + 4096,
1799 + 4096,
1800 + )
1801 + .expect("server_create after resp_event fault");
1802 + server.destroy();
1803 + }
1804 +
1805 + #[cfg(windows)]
1806 + #[test]
1807 + fn test_client_attach_fault_open_mapping_windows() {
1808 + let run_dir = test_run_dir();
1809 + let service = unique_service("rs_win_shm_fault_open_mapping");
1810 + let auth_token: u64 = 0x6611;
1811 + let session_id: u64 = 51;
1812 +
1813 + let mut server = WinShmContext::server_create(
1814 + &run_dir,
1815 + &service,
1816 + auth_token,
1817 + session_id,
1818 + PROFILE_HYBRID,
1819 + 4096,
1820 + 4096,
1821 + )
1822 + .expect("server_create");
1823 +
1824 + let _fault = FaultHookGuard::install(WinShmFaultSite::OpenFileMapping, 9);
1825 + let err = match WinShmContext::client_attach(
1826 + &run_dir,
1827 + &service,
1828 + auth_token,
1829 + session_id,
1830 + PROFILE_HYBRID,
1831 + ) {
1832 + Ok(_) => panic!("OpenFileMappingW fault must fail"),
1833 + Err(err) => err,
1834 + };
1835 + assert_eq!(err, WinShmError::OpenMapping(9));
1836 + drop(_fault);
1837 +
1838 + let mut client = WinShmContext::client_attach(
1839 + &run_dir,
1840 + &service,
1841 + auth_token,
1842 + session_id,
1843 + PROFILE_HYBRID,
1844 + )
1845 + .expect("client_attach after OpenFileMappingW fault");
1846 + client.close();
1847 + server.destroy();
1848 + }
1849 +
1850 + #[cfg(windows)]
1851 + #[test]
1852 + fn test_client_attach_fault_map_view_windows() {
1853 + let run_dir = test_run_dir();
1854 + let service = unique_service("rs_win_shm_fault_attach_map_view");
1855 + let auth_token: u64 = 0x6612;
1856 + let session_id: u64 = 52;
1857 +
1858 + let mut server = WinShmContext::server_create(
1859 + &run_dir,
1860 + &service,
1861 + auth_token,
1862 + session_id,
1863 + PROFILE_HYBRID,
1864 + 4096,
1865 + 4096,
1866 + )
1867 + .expect("server_create");
1868 +
1869 + let _fault = FaultHookGuard::install(WinShmFaultSite::MapViewOfFile, 10);
1870 + let err = match WinShmContext::client_attach(
1871 + &run_dir,
1872 + &service,
1873 + auth_token,
1874 + session_id,
1875 + PROFILE_HYBRID,
1876 + ) {
1877 + Ok(_) => panic!("MapViewOfFile attach fault must fail"),
1878 + Err(err) => err,
1879 + };
1880 + assert_eq!(err, WinShmError::MapView(10));
1881 + drop(_fault);
1882 +
1883 + let mut client = WinShmContext::client_attach(
1884 + &run_dir,
1885 + &service,
1886 + auth_token,
1887 + session_id,
1888 + PROFILE_HYBRID,
1889 + )
1890 + .expect("client_attach after MapViewOfFile fault");
1891 + client.close();
1892 + server.destroy();
1893 + }
1894 +
1895 + #[cfg(windows)]
1896 + #[test]
1897 + fn test_client_attach_fault_req_event_windows() {
1898 + let run_dir = test_run_dir();
1899 + let service = unique_service("rs_win_shm_fault_req_open_event");
1900 + let auth_token: u64 = 0x6613;
1901 + let session_id: u64 = 53;
1902 +
1903 + let mut server = WinShmContext::server_create(
1904 + &run_dir,
1905 + &service,
1906 + auth_token,
1907 + session_id,
1908 + PROFILE_HYBRID,
1909 + 4096,
1910 + 4096,
1911 + )
1912 + .expect("server_create");
1913 +
1914 + let _fault = FaultHookGuard::install(WinShmFaultSite::OpenEvent, 11);
1915 + let err = match WinShmContext::client_attach(
1916 + &run_dir,
1917 + &service,
1918 + auth_token,
1919 + session_id,
1920 + PROFILE_HYBRID,
1921 + ) {
1922 + Ok(_) => panic!("req_event OpenEventW fault must fail"),
1923 + Err(err) => err,
1924 + };
1925 + assert_eq!(err, WinShmError::OpenEvent(11));
1926 + drop(_fault);
1927 +
1928 + let mut client = WinShmContext::client_attach(
1929 + &run_dir,
1930 + &service,
1931 + auth_token,
1932 + session_id,
1933 + PROFILE_HYBRID,
1934 + )
1935 + .expect("client_attach after req_event OpenEventW fault");
1936 + client.close();
1937 + server.destroy();
1938 + }
1939 +
1940 + #[cfg(windows)]
1941 + #[test]
1942 + fn test_client_attach_fault_resp_event_windows() {
1943 + let run_dir = test_run_dir();
1944 + let service = unique_service("rs_win_shm_fault_resp_open_event");
1945 + let auth_token: u64 = 0x6614;
1946 + let session_id: u64 = 54;
1947 +
1948 + let mut server = WinShmContext::server_create(
1949 + &run_dir,
1950 + &service,
1951 + auth_token,
1952 + session_id,
1953 + PROFILE_HYBRID,
1954 + 4096,
1955 + 4096,
1956 + )
1957 + .expect("server_create");
1958 +
1959 + let _fault = FaultHookGuard::install_after(WinShmFaultSite::OpenEvent, 12, 1);
1960 + let err = match WinShmContext::client_attach(
1961 + &run_dir,
1962 + &service,
1963 + auth_token,
1964 + session_id,
1965 + PROFILE_HYBRID,
1966 + ) {
1967 + Ok(_) => panic!("resp_event OpenEventW fault must fail"),
1968 + Err(err) => err,
1969 + };
1970 + assert_eq!(err, WinShmError::OpenEvent(12));
1971 + drop(_fault);
1972 +
1973 + let mut client = WinShmContext::client_attach(
1974 + &run_dir,
1975 + &service,
1976 + auth_token,
1977 + session_id,
1978 + PROFILE_HYBRID,
1979 + )
1980 + .expect("client_attach after resp_event OpenEventW fault");
1981 + client.close();
1982 + server.destroy();
1983 + }
1984 +}
src/crates/netipc/src/transport/windows.rs new
+2146
@@ -0,0 +1,2146 @@
1 +//! L1 Windows Named Pipe transport.
2 +//!
3 +//! Connection lifecycle, handshake with profile/limit negotiation,
4 +//! and send/receive with transparent chunking over Win32 Named Pipes
5 +//! in message mode. Wire-compatible with the C and Go implementations.
6 +
7 +use crate::protocol::{
8 + self, align8, ChunkHeader, Header, Hello, HelloAck, FLAG_BATCH, HEADER_SIZE, KIND_REQUEST,
9 + KIND_RESPONSE, MAGIC_CHUNK, MAGIC_MSG, MAX_PAYLOAD_CAP, MAX_PAYLOAD_DEFAULT, PROFILE_BASELINE,
10 + VERSION,
11 +};
12 +use std::collections::HashSet;
13 +use std::ptr;
14 +use std::sync::atomic::{AtomicU64, Ordering};
15 +
16 +// ---------------------------------------------------------------------------
17 +// Win32 FFI — using windows-sys when available, raw bindings as fallback
18 +// ---------------------------------------------------------------------------
19 +
20 +#[cfg(windows)]
21 +mod ffi {
22 + #![allow(non_snake_case, non_camel_case_types, dead_code)]
23 +
24 + pub type HANDLE = isize;
25 + pub type DWORD = u32;
26 + pub type BOOL = i32;
27 + pub type LPCWSTR = *const u16;
28 + pub type LPVOID = *mut core::ffi::c_void;
29 + pub type LPCVOID = *const core::ffi::c_void;
30 + pub type LPDWORD = *mut DWORD;
31 +
32 + pub const INVALID_HANDLE_VALUE: HANDLE = -1;
33 + pub const PIPE_ACCESS_DUPLEX: DWORD = 0x00000003;
34 + pub const FILE_FLAG_FIRST_PIPE_INSTANCE: DWORD = 0x00080000;
35 + pub const PIPE_TYPE_MESSAGE: DWORD = 0x00000004;
36 + pub const PIPE_READMODE_MESSAGE: DWORD = 0x00000002;
37 + pub const PIPE_WAIT: DWORD = 0x00000000;
38 + pub const PIPE_UNLIMITED_INSTANCES: DWORD = 255;
39 + pub const GENERIC_READ: DWORD = 0x80000000;
40 + pub const GENERIC_WRITE: DWORD = 0x40000000;
41 + pub const OPEN_EXISTING: DWORD = 3;
42 +
43 + pub const ERROR_PIPE_CONNECTED: DWORD = 535;
44 + pub const ERROR_BROKEN_PIPE: DWORD = 109;
45 + pub const ERROR_NO_DATA: DWORD = 232;
46 + pub const ERROR_PIPE_NOT_CONNECTED: DWORD = 233;
47 + pub const ERROR_ACCESS_DENIED: DWORD = 5;
48 + pub const ERROR_PIPE_BUSY: DWORD = 231;
49 +
50 + extern "system" {
51 + pub fn CreateNamedPipeW(
52 + lpName: LPCWSTR,
53 + dwOpenMode: DWORD,
54 + dwPipeMode: DWORD,
55 + nMaxInstances: DWORD,
56 + nOutBufferSize: DWORD,
57 + nInBufferSize: DWORD,
58 + nDefaultTimeOut: DWORD,
59 + lpSecurityAttributes: *const core::ffi::c_void,
60 + ) -> HANDLE;
61 +
62 + pub fn ConnectNamedPipe(hNamedPipe: HANDLE, lpOverlapped: *mut core::ffi::c_void) -> BOOL;
63 +
64 + pub fn DisconnectNamedPipe(hNamedPipe: HANDLE) -> BOOL;
65 + pub fn FlushFileBuffers(hFile: HANDLE) -> BOOL;
66 +
67 + pub fn CreateFileW(
68 + lpFileName: LPCWSTR,
69 + dwDesiredAccess: DWORD,
70 + dwShareMode: DWORD,
71 + lpSecurityAttributes: *const core::ffi::c_void,
72 + dwCreationDisposition: DWORD,
73 + dwFlagsAndAttributes: DWORD,
74 + hTemplateFile: HANDLE,
75 + ) -> HANDLE;
76 +
77 + pub fn ReadFile(
78 + hFile: HANDLE,
79 + lpBuffer: LPVOID,
80 + nNumberOfBytesToRead: DWORD,
81 + lpNumberOfBytesRead: LPDWORD,
82 + lpOverlapped: *mut core::ffi::c_void,
83 + ) -> BOOL;
84 +
85 + pub fn WriteFile(
86 + hFile: HANDLE,
87 + lpBuffer: LPCVOID,
88 + nNumberOfBytesToWrite: DWORD,
89 + lpNumberOfBytesWritten: LPDWORD,
90 + lpOverlapped: *mut core::ffi::c_void,
91 + ) -> BOOL;
92 +
93 + pub fn CloseHandle(hObject: HANDLE) -> BOOL;
94 +
95 + pub fn GetLastError() -> DWORD;
96 + pub fn GetTickCount64() -> u64;
97 + pub fn PeekNamedPipe(
98 + hNamedPipe: HANDLE,
99 + lpBuffer: LPVOID,
100 + nBufferSize: DWORD,
101 + lpBytesRead: LPDWORD,
102 + lpTotalBytesAvail: LPDWORD,
103 + lpBytesLeftThisMessage: LPDWORD,
104 + ) -> BOOL;
105 +
106 + pub fn SwitchToThread() -> BOOL;
107 +
108 + pub fn SetNamedPipeHandleState(
109 + hNamedPipe: HANDLE,
110 + lpMode: *const DWORD,
111 + lpMaxCollectionCount: *const DWORD,
112 + lpCollectDataTimeout: *const DWORD,
113 + ) -> BOOL;
114 + }
115 +}
116 +
117 +// ---------------------------------------------------------------------------
118 +// Constants
119 +// ---------------------------------------------------------------------------
120 +
121 +const DEFAULT_BATCH_ITEMS: u32 = 1;
122 +const DEFAULT_PACKET_SIZE: u32 = 65536;
123 +const DEFAULT_PIPE_BUF_SIZE: u32 = 65536;
124 +const HELLO_PAYLOAD_SIZE: usize = 44;
125 +const HELLO_ACK_PAYLOAD_SIZE: usize = 48;
126 +const MAX_PIPE_NAME_CHARS: usize = 256;
127 +
128 +// FNV-1a 64-bit constants
129 +const FNV1A_OFFSET_BASIS: u64 = 0xcbf29ce484222325;
130 +const FNV1A_PRIME: u64 = 0x00000100000001B3;
131 +
132 +// ---------------------------------------------------------------------------
133 +// Errors
134 +// ---------------------------------------------------------------------------
135 +
136 +/// Transport-level errors for Named Pipe transport.
137 +#[derive(Debug, Clone, PartialEq, Eq)]
138 +pub enum NpError {
139 + /// Pipe name derivation failed.
140 + PipeName(String),
141 + /// CreateNamedPipeW failed.
142 + CreatePipe(u32),
143 + /// CreateFileW / connection failed.
144 + Connect(u32),
145 + /// ConnectNamedPipe (accept) failed.
146 + Accept(u32),
147 + /// WriteFile failed.
148 + Send(u32),
149 + /// ReadFile failed or peer disconnected.
150 + Recv(u32),
151 + /// Handshake protocol error.
152 + Handshake(String),
153 + /// Authentication token rejected.
154 + AuthFailed,
155 + /// No common profile.
156 + NoProfile,
157 + /// Protocol or layout version mismatch.
158 + Incompatible(String),
159 + /// Wire protocol violation.
160 + Protocol(String),
161 + /// Pipe name already in use by live server.
162 + AddrInUse,
163 + /// Chunk header mismatch.
164 + Chunk(String),
165 + /// Memory allocation failed.
166 + Alloc,
167 + /// Payload or batch count exceeds negotiated limit.
168 + LimitExceeded,
169 + /// Invalid argument.
170 + BadParam(String),
171 + /// Duplicate message_id on send.
172 + DuplicateMsgId(u64),
173 + /// Unknown message_id on receive.
174 + UnknownMsgId(u64),
175 + /// Peer disconnected (graceful).
176 + Disconnected,
177 +}
178 +
179 +impl std::fmt::Display for NpError {
180 + fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
181 + match self {
182 + NpError::PipeName(s) => write!(f, "pipe name error: {s}"),
183 + NpError::CreatePipe(e) => write!(f, "CreateNamedPipeW failed: {e}"),
184 + NpError::Connect(e) => write!(f, "connect failed: {e}"),
185 + NpError::Accept(e) => write!(f, "accept failed: {e}"),
186 + NpError::Send(e) => write!(f, "send failed: {e}"),
187 + NpError::Recv(e) => write!(f, "recv failed: {e}"),
188 + NpError::Handshake(s) => write!(f, "handshake error: {s}"),
189 + NpError::AuthFailed => write!(f, "authentication token rejected"),
190 + NpError::NoProfile => write!(f, "no common transport profile"),
191 + NpError::Incompatible(s) => write!(f, "incompatible protocol: {s}"),
192 + NpError::Protocol(s) => write!(f, "protocol violation: {s}"),
193 + NpError::AddrInUse => write!(f, "pipe name already in use by live server"),
194 + NpError::Chunk(s) => write!(f, "chunk error: {s}"),
195 + NpError::Alloc => write!(f, "memory allocation failed"),
196 + NpError::LimitExceeded => write!(f, "negotiated limit exceeded"),
197 + NpError::BadParam(s) => write!(f, "bad parameter: {s}"),
198 + NpError::DuplicateMsgId(id) => write!(f, "duplicate message_id: {id}"),
199 + NpError::UnknownMsgId(id) => write!(f, "unknown response message_id: {id}"),
200 + NpError::Disconnected => write!(f, "peer disconnected"),
201 + }
202 + }
203 +}
204 +
205 +fn header_version_incompatible(buf: &[u8], expected_code: u16) -> bool {
206 + if buf.len() < HEADER_SIZE {
207 + return false;
208 + }
209 +
210 + let magic = u32::from_ne_bytes(buf[0..4].try_into().unwrap());
211 + let version = u16::from_ne_bytes(buf[4..6].try_into().unwrap());
212 + let header_len = u16::from_ne_bytes(buf[6..8].try_into().unwrap());
213 + let kind = u16::from_ne_bytes(buf[8..10].try_into().unwrap());
214 + let code = u16::from_ne_bytes(buf[12..14].try_into().unwrap());
215 +
216 + magic == MAGIC_MSG
217 + && version != VERSION
218 + && header_len == protocol::HEADER_LEN
219 + && kind == protocol::KIND_CONTROL
220 + && code == expected_code
221 +}
222 +
223 +fn hello_layout_incompatible(buf: &[u8]) -> bool {
224 + buf.len() >= 2 && u16::from_ne_bytes(buf[0..2].try_into().unwrap()) != 1
225 +}
226 +
227 +fn hello_ack_layout_incompatible(buf: &[u8]) -> bool {
228 + buf.len() >= 2 && u16::from_ne_bytes(buf[0..2].try_into().unwrap()) != 1
229 +}
230 +
231 +impl std::error::Error for NpError {}
232 +
233 +// ---------------------------------------------------------------------------
234 +// Role
235 +// ---------------------------------------------------------------------------
236 +
237 +#[derive(Debug, Clone, Copy, PartialEq, Eq)]
238 +pub enum Role {
239 + Client = 1,
240 + Server = 2,
241 +}
242 +
243 +// ---------------------------------------------------------------------------
244 +// Configuration
245 +// ---------------------------------------------------------------------------
246 +
247 +/// Client connection configuration.
248 +#[derive(Debug, Clone)]
249 +pub struct ClientConfig {
250 + pub supported_profiles: u32,
251 + pub preferred_profiles: u32,
252 + pub max_request_payload_bytes: u32,
253 + pub max_request_batch_items: u32,
254 + pub max_response_payload_bytes: u32,
255 + pub max_response_batch_items: u32,
256 + pub auth_token: u64,
257 + /// 0 = use default (65536).
258 + pub packet_size: u32,
259 +}
260 +
261 +impl Default for ClientConfig {
262 + fn default() -> Self {
263 + Self {
264 + supported_profiles: PROFILE_BASELINE,
265 + preferred_profiles: 0,
266 + max_request_payload_bytes: 0,
267 + max_request_batch_items: 0,
268 + max_response_payload_bytes: 0,
269 + max_response_batch_items: 0,
270 + auth_token: 0,
271 + packet_size: 0,
272 + }
273 + }
274 +}
275 +
276 +/// Server configuration for listen + accept.
277 +#[derive(Debug, Clone)]
278 +pub struct ServerConfig {
279 + pub supported_profiles: u32,
280 + pub preferred_profiles: u32,
281 + pub max_request_payload_bytes: u32,
282 + pub max_request_batch_items: u32,
283 + pub max_response_payload_bytes: u32,
284 + pub max_response_batch_items: u32,
285 + pub auth_token: u64,
286 + /// 0 = use default (65536).
287 + pub packet_size: u32,
288 +}
289 +
290 +impl Default for ServerConfig {
291 + fn default() -> Self {
292 + Self {
293 + supported_profiles: PROFILE_BASELINE,
294 + preferred_profiles: 0,
295 + max_request_payload_bytes: 0,
296 + max_request_batch_items: 0,
297 + max_response_payload_bytes: 0,
298 + max_response_batch_items: 0,
299 + auth_token: 0,
300 + packet_size: 0,
301 + }
302 + }
303 +}
304 +
305 +// ---------------------------------------------------------------------------
306 +// FNV-1a 64-bit hash
307 +// ---------------------------------------------------------------------------
308 +
309 +/// Compute FNV-1a 64-bit hash of data.
310 +pub fn fnv1a_64(data: &[u8]) -> u64 {
311 + let mut hash = FNV1A_OFFSET_BASIS;
312 + for &byte in data {
313 + hash ^= byte as u64;
314 + hash = hash.wrapping_mul(FNV1A_PRIME);
315 + }
316 + hash
317 +}
318 +
319 +// ---------------------------------------------------------------------------
320 +// Service name validation
321 +// ---------------------------------------------------------------------------
322 +
323 +fn validate_service_name(name: &str) -> Result<(), NpError> {
324 + if name.is_empty() {
325 + return Err(NpError::BadParam("empty service name".into()));
326 + }
327 + if name == "." || name == ".." {
328 + return Err(NpError::BadParam(
329 + "service name cannot be '.' or '..'".into(),
330 + ));
331 + }
332 + for &b in name.as_bytes() {
333 + if (b >= b'a' && b <= b'z')
334 + || (b >= b'A' && b <= b'Z')
335 + || (b >= b'0' && b <= b'9')
336 + || b == b'.'
337 + || b == b'_'
338 + || b == b'-'
339 + {
340 + continue;
341 + }
342 + return Err(NpError::BadParam(format!(
343 + "service name contains invalid character: {:?}",
344 + b as char
345 + )));
346 + }
347 + Ok(())
348 +}
349 +
350 +// ---------------------------------------------------------------------------
351 +// Pipe name derivation
352 +// ---------------------------------------------------------------------------
353 +
354 +/// Build pipe name from run_dir and service_name.
355 +/// Returns the pipe name as a NUL-terminated wide string vector.
356 +pub fn build_pipe_name(run_dir: &str, service_name: &str) -> Result<Vec<u16>, NpError> {
357 + validate_service_name(service_name)?;
358 +
359 + let hash = fnv1a_64(run_dir.as_bytes());
360 + let narrow = format!("\\\\.\\pipe\\netipc-{:016x}-{}", hash, service_name);
361 +
362 + if narrow.len() >= MAX_PIPE_NAME_CHARS {
363 + return Err(NpError::PipeName("pipe name too long".into()));
364 + }
365 +
366 + // Convert to UTF-16 with NUL terminator
367 + let mut wide: Vec<u16> = narrow.encode_utf16().collect();
368 + wide.push(0);
369 + Ok(wide)
370 +}
371 +
372 +// ---------------------------------------------------------------------------
373 +// Internal helpers
374 +// ---------------------------------------------------------------------------
375 +
376 +fn apply_default(val: u32, def: u32) -> u32 {
377 + if val == 0 {
378 + def
379 + } else {
380 + val
381 + }
382 +}
383 +
384 +fn min_u32(a: u32, b: u32) -> u32 {
385 + a.min(b)
386 +}
387 +
388 +fn max_u32(a: u32, b: u32) -> u32 {
389 + a.max(b)
390 +}
391 +
392 +#[cfg(windows)]
393 +fn pipe_buffer_size(packet_size: u32) -> u32 {
394 + // The protocol packet size controls logical framing and chunk size. The
395 + // underlying pipe quota must stay large enough for full-duplex pipelining
396 + // even when tests force a tiny protocol packet size.
397 + max_u32(
398 + apply_default(packet_size, DEFAULT_PIPE_BUF_SIZE),
399 + DEFAULT_PIPE_BUF_SIZE,
400 + )
401 +}
402 +
403 +fn highest_bit(mask: u32) -> u32 {
404 + if mask == 0 {
405 + return 0;
406 + }
407 + let mut bit: u32 = 1 << 31;
408 + while bit & mask == 0 {
409 + bit >>= 1;
410 + }
411 + bit
412 +}
413 +
414 +// ---------------------------------------------------------------------------
415 +// Low-level I/O (Windows-only)
416 +// ---------------------------------------------------------------------------
417 +
418 +#[cfg(windows)]
419 +fn is_disconnect_error(err: u32) -> bool {
420 + err == ffi::ERROR_BROKEN_PIPE
421 + || err == ffi::ERROR_NO_DATA
422 + || err == ffi::ERROR_PIPE_NOT_CONNECTED
423 +}
424 +
425 +#[cfg(windows)]
426 +fn last_error() -> u32 {
427 + unsafe { ffi::GetLastError() }
428 +}
429 +
430 +#[cfg(windows)]
431 +fn raw_write(handle: ffi::HANDLE, data: &[u8]) -> Result<(), NpError> {
432 + let mut written: u32 = 0;
433 + let ok = unsafe {
434 + ffi::WriteFile(
435 + handle,
436 + data.as_ptr() as ffi::LPCVOID,
437 + data.len() as u32,
438 + &mut written,
439 + ptr::null_mut(),
440 + )
441 + };
442 + if ok == 0 {
443 + let err = last_error();
444 + if is_disconnect_error(err) {
445 + return Err(NpError::Disconnected);
446 + }
447 + return Err(NpError::Send(err));
448 + }
449 + if written != data.len() as u32 {
450 + return Err(NpError::Send(0));
451 + }
452 + Ok(())
453 +}
454 +
455 +/// Send header + payload as one pipe message.
456 +#[cfg(windows)]
457 +fn raw_send_msg(
458 + handle: ffi::HANDLE,
459 + scratch: &mut Vec<u8>,
460 + hdr: &[u8],
461 + payload: &[u8],
462 +) -> Result<(), NpError> {
463 + let total = hdr.len() + payload.len();
464 + if scratch.len() < total {
465 + scratch.resize(total, 0);
466 + }
467 + scratch[..hdr.len()].copy_from_slice(hdr);
468 + scratch[hdr.len()..total].copy_from_slice(payload);
469 + raw_write(handle, &scratch[..total])
470 +}
471 +
472 +/// Read one pipe message. Returns number of bytes read.
473 +#[cfg(windows)]
474 +fn raw_recv(handle: ffi::HANDLE, buf: &mut [u8]) -> Result<usize, NpError> {
475 + let mut read: u32 = 0;
476 + let ok = unsafe {
477 + ffi::ReadFile(
478 + handle,
479 + buf.as_mut_ptr() as ffi::LPVOID,
480 + buf.len() as u32,
481 + &mut read,
482 + ptr::null_mut(),
483 + )
484 + };
485 + if ok == 0 {
486 + let err = last_error();
487 + if is_disconnect_error(err) {
488 + return Err(NpError::Disconnected);
489 + }
490 + return Err(NpError::Recv(err));
491 + }
492 + if read == 0 {
493 + return Err(NpError::Disconnected);
494 + }
495 + Ok(read as usize)
496 +}
497 +
498 +#[cfg(windows)]
499 +fn close_handle(handle: ffi::HANDLE) {
500 + if handle != ffi::INVALID_HANDLE_VALUE && handle != 0 {
501 + unsafe {
502 + ffi::CloseHandle(handle);
503 + }
504 + }
505 +}
506 +
507 +// ---------------------------------------------------------------------------
508 +// Session
509 +// ---------------------------------------------------------------------------
510 +
511 +/// A connected Named Pipe session (client or server side).
512 +#[cfg(windows)]
513 +pub struct NpSession {
514 + handle: ffi::HANDLE,
515 + role: Role,
516 +
517 + // Negotiated limits
518 + pub max_request_payload_bytes: u32,
519 + pub max_request_batch_items: u32,
520 + pub max_response_payload_bytes: u32,
521 + pub max_response_batch_items: u32,
522 + pub packet_size: u32,
523 + pub selected_profile: u32,
524 + pub session_id: u64,
525 +
526 + // Internal receive buffer for chunked reassembly
527 + recv_buf: Vec<u8>,
528 + pkt_buf: Vec<u8>,
529 + send_buf: Vec<u8>,
530 +
531 + // In-flight message_id set (client-side only)
532 + inflight_ids: HashSet<u64>,
533 +}
534 +
535 +#[cfg(windows)]
536 +impl NpSession {
537 + fn fail_all_inflight(&mut self) {
538 + if self.role == Role::Client {
539 + self.inflight_ids.clear();
540 + }
541 + }
542 +
543 + /// Get the raw HANDLE for WaitForSingleObject integration.
544 + pub fn handle(&self) -> ffi::HANDLE {
545 + self.handle
546 + }
547 +
548 + /// Get the session role.
549 + pub fn role(&self) -> Role {
550 + self.role
551 + }
552 +
553 + /// Wait until the pipe becomes readable or the timeout expires.
554 + pub fn wait_readable(&self, timeout_ms: u32) -> Result<bool, NpError> {
555 + if self.handle == ffi::INVALID_HANDLE_VALUE || self.handle == 0 {
556 + return Err(NpError::BadParam("session closed".into()));
557 + }
558 +
559 + let deadline = unsafe { ffi::GetTickCount64() }.saturating_add(timeout_ms as u64);
560 + let mut yielded = false;
561 + loop {
562 + let mut available: u32 = 0;
563 + let ok = unsafe {
564 + ffi::PeekNamedPipe(
565 + self.handle,
566 + ptr::null_mut(),
567 + 0,
568 + ptr::null_mut(),
569 + &mut available,
570 + ptr::null_mut(),
571 + )
572 + };
573 + if ok == 0 {
574 + let err = last_error();
575 + if is_disconnect_error(err) {
576 + return Err(NpError::Disconnected);
577 + }
578 + return Err(NpError::Recv(err));
579 + }
580 + if available > 0 {
581 + return Ok(true);
582 + }
583 +
584 + if unsafe { ffi::GetTickCount64() } >= deadline {
585 + return Ok(false);
586 + }
587 +
588 + if !yielded {
589 + yielded = true;
590 + for _ in 0..256 {
591 + unsafe {
592 + ffi::SwitchToThread();
593 + }
594 +
595 + let mut yielded_available: u32 = 0;
596 + let yielded_ok = unsafe {
597 + ffi::PeekNamedPipe(
598 + self.handle,
599 + ptr::null_mut(),
600 + 0,
601 + ptr::null_mut(),
602 + &mut yielded_available,
603 + ptr::null_mut(),
604 + )
605 + };
606 + if yielded_ok == 0 {
607 + let err = last_error();
608 + if is_disconnect_error(err) {
609 + return Err(NpError::Disconnected);
610 + }
611 + return Err(NpError::Recv(err));
612 + }
613 + if yielded_available > 0 {
614 + return Ok(true);
615 + }
616 +
617 + if unsafe { ffi::GetTickCount64() } >= deadline {
618 + return Ok(false);
619 + }
620 + }
621 + continue;
622 + }
623 +
624 + std::thread::sleep(std::time::Duration::from_millis(1));
625 + }
626 + }
627 +
628 + /// Return the most recently reassembled payload stored in the internal
629 + /// receive buffer.
630 + pub fn received_payload(&self, len: usize) -> &[u8] {
631 + &self.recv_buf[..len]
632 + }
633 +
634 + /// Connect to a server pipe derived from run_dir + service_name.
635 + pub fn connect(
636 + run_dir: &str,
637 + service_name: &str,
638 + config: &ClientConfig,
639 + ) -> Result<Self, NpError> {
640 + let pipe_name = build_pipe_name(run_dir, service_name)?;
641 +
642 + let handle = unsafe {
643 + ffi::CreateFileW(
644 + pipe_name.as_ptr(),
645 + ffi::GENERIC_READ | ffi::GENERIC_WRITE,
646 + 0,
647 + ptr::null(),
648 + ffi::OPEN_EXISTING,
649 + 0,
650 + 0,
651 + )
652 + };
653 +
654 + if handle == ffi::INVALID_HANDLE_VALUE {
655 + return Err(NpError::Connect(last_error()));
656 + }
657 +
658 + // Set read mode to message mode
659 + let mode: u32 = ffi::PIPE_READMODE_MESSAGE;
660 + let ok = unsafe { ffi::SetNamedPipeHandleState(handle, &mode, ptr::null(), ptr::null()) };
661 + if ok == 0 {
662 + let err = last_error();
663 + close_handle(handle);
664 + return Err(NpError::Connect(err));
665 + }
666 +
667 + match client_handshake(handle, config) {
668 + Ok(session) => Ok(session),
669 + Err(e) => {
670 + close_handle(handle);
671 + Err(e)
672 + }
673 + }
674 + }
675 +
676 + /// Send one logical message. Fills magic/version/header_len/payload_len.
677 + /// Chunked transparently if message exceeds packet_size.
678 + pub fn send(&mut self, hdr: &mut Header, payload: &[u8]) -> Result<(), NpError> {
679 + if self.handle == ffi::INVALID_HANDLE_VALUE {
680 + return Err(NpError::BadParam("session closed".into()));
681 + }
682 +
683 + // Validate payload against negotiated directional limits before transmitting.
684 + let (max_payload, max_items) = if self.role == Role::Client {
685 + (self.max_request_payload_bytes, self.max_request_batch_items)
686 + } else {
687 + (
688 + self.max_response_payload_bytes,
689 + self.max_response_batch_items,
690 + )
691 + };
692 + if payload.len() > max_payload as usize || payload.len() > u32::MAX as usize {
693 + return Err(NpError::LimitExceeded);
694 + }
695 + if hdr.item_count > max_items {
696 + return Err(NpError::LimitExceeded);
697 + }
698 +
699 + // Client-side: track in-flight message_ids for requests
700 + if self.role == Role::Client && hdr.kind == KIND_REQUEST {
701 + if !self.inflight_ids.insert(hdr.message_id) {
702 + return Err(NpError::DuplicateMsgId(hdr.message_id));
703 + }
704 + }
705 +
706 + // Fill envelope fields
707 + hdr.magic = MAGIC_MSG;
708 + hdr.version = VERSION;
709 + hdr.header_len = protocol::HEADER_LEN;
710 + hdr.payload_len = payload.len() as u32;
711 +
712 + let tracked = self.role == Role::Client && hdr.kind == KIND_REQUEST;
713 + let msg_id = hdr.message_id;
714 +
715 + let result = self.send_inner(hdr, payload);
716 +
717 + if let Err(err) = &result {
718 + if tracked {
719 + match err {
720 + NpError::Send(_) | NpError::Disconnected => self.fail_all_inflight(),
721 + _ => {
722 + self.inflight_ids.remove(&msg_id);
723 + }
724 + }
725 + }
726 + }
727 +
728 + result
729 + }
730 +
731 + fn send_inner(&mut self, hdr: &mut Header, payload: &[u8]) -> Result<(), NpError> {
732 + let total_msg = HEADER_SIZE + payload.len();
733 +
734 + // Single packet?
735 + if total_msg <= self.packet_size as usize {
736 + let mut hdr_buf = [0u8; HEADER_SIZE];
737 + hdr.encode(&mut hdr_buf);
738 + return raw_send_msg(self.handle, &mut self.send_buf, &hdr_buf, payload);
739 + }
740 +
741 + // Chunked send
742 + let chunk_payload_budget = self.packet_size as usize - HEADER_SIZE;
743 + if chunk_payload_budget == 0 {
744 + return Err(NpError::BadParam("packet_size too small".into()));
745 + }
746 +
747 + let first_chunk_payload = payload.len().min(chunk_payload_budget);
748 + let remaining_after_first = payload.len() - first_chunk_payload;
749 +
750 + let continuation_chunks = if remaining_after_first > 0 {
751 + (remaining_after_first + chunk_payload_budget - 1) / chunk_payload_budget
752 + } else {
753 + 0
754 + };
755 + let chunk_count = 1 + continuation_chunks as u32;
756 +
757 + // First chunk
758 + let mut hdr_buf = [0u8; HEADER_SIZE];
759 + hdr.encode(&mut hdr_buf);
760 + raw_send_msg(
761 + self.handle,
762 + &mut self.send_buf,
763 + &hdr_buf,
764 + &payload[..first_chunk_payload],
765 + )?;
766 +
767 + // Continuation chunks
768 + let mut offset = first_chunk_payload;
769 + for ci in 1..chunk_count {
770 + let remaining = payload.len() - offset;
771 + let this_chunk = remaining.min(chunk_payload_budget);
772 +
773 + let chk = ChunkHeader {
774 + magic: MAGIC_CHUNK,
775 + version: VERSION,
776 + flags: 0,
777 + message_id: hdr.message_id,
778 + total_message_len: total_msg as u32,
779 + chunk_index: ci,
780 + chunk_count,
781 + chunk_payload_len: this_chunk as u32,
782 + };
783 +
784 + let mut chk_buf = [0u8; HEADER_SIZE];
785 + chk.encode(&mut chk_buf);
786 + raw_send_msg(
787 + self.handle,
788 + &mut self.send_buf,
789 + &chk_buf,
790 + &payload[offset..offset + this_chunk],
791 + )?;
792 +
793 + offset += this_chunk;
794 + }
795 +
796 + Ok(())
797 + }
798 +
799 + /// Receive one logical message. buf is a scratch buffer for the first read.
800 + /// The payload view is valid until the next receive call on this session.
801 + pub fn receive<'a>(&'a mut self, buf: &'a mut [u8]) -> Result<(Header, &'a [u8]), NpError> {
802 + if self.handle == ffi::INVALID_HANDLE_VALUE {
803 + return Err(NpError::BadParam("session closed".into()));
804 + }
805 +
806 + let n = match raw_recv(self.handle, buf) {
807 + Ok(n) => n,
808 + Err(err) => {
809 + self.fail_all_inflight();
810 + return Err(err);
811 + }
812 + };
813 +
814 + if n < HEADER_SIZE {
815 + return Err(NpError::Protocol("packet too short for header".into()));
816 + }
817 +
818 + let hdr = Header::decode(&buf[..n])
819 + .map_err(|e| NpError::Protocol(format!("header decode: {e}")))?;
820 +
821 + // Validate payload_len against negotiated limit
822 + let max_payload = if self.role == Role::Server {
823 + self.max_request_payload_bytes
824 + } else {
825 + self.max_response_payload_bytes
826 + };
827 + if hdr.payload_len > max_payload {
828 + return Err(NpError::LimitExceeded);
829 + }
830 +
831 + // Validate item_count
832 + let max_batch = if self.role == Role::Server {
833 + self.max_request_batch_items
834 + } else {
835 + self.max_response_batch_items
836 + };
837 + if hdr.item_count > max_batch {
838 + return Err(NpError::LimitExceeded);
839 + }
840 +
841 + // Client-side: validate response message_id
842 + if self.role == Role::Client && hdr.kind == KIND_RESPONSE {
843 + if !self.inflight_ids.remove(&hdr.message_id) {
844 + return Err(NpError::UnknownMsgId(hdr.message_id));
845 + }
846 + }
847 +
848 + let total_msg = HEADER_SIZE + hdr.payload_len as usize;
849 +
850 + // Non-chunked
851 + if n >= total_msg {
852 + let payload = &buf[HEADER_SIZE..HEADER_SIZE + hdr.payload_len as usize];
853 +
854 + // Validate batch directory
855 + if hdr.flags & FLAG_BATCH != 0 && hdr.item_count > 1 {
856 + let dir_bytes = hdr.item_count as usize * 8;
857 + let dir_aligned = align8(dir_bytes);
858 + if payload.len() < dir_aligned {
859 + return Err(NpError::Protocol("batch directory exceeds payload".into()));
860 + }
861 + let packed_area_len = (payload.len() - dir_aligned) as u32;
862 + protocol::batch_dir_validate(
863 + &payload[..dir_bytes],
864 + hdr.item_count,
865 + packed_area_len,
866 + )
867 + .map_err(|e| NpError::Protocol(format!("batch directory: {e:?}")))?;
868 + }
869 +
870 + return Ok((hdr, payload));
871 + }
872 +
873 + // Chunked
874 + let first_payload_bytes = n - HEADER_SIZE;
875 + let needed = hdr.payload_len as usize;
876 + if self.recv_buf.len() < needed {
877 + self.recv_buf.resize(needed, 0);
878 + }
879 +
880 + self.recv_buf[..first_payload_bytes]
881 + .copy_from_slice(&buf[HEADER_SIZE..HEADER_SIZE + first_payload_bytes]);
882 +
883 + let mut assembled = first_payload_bytes;
884 + let chunk_payload_budget = self.packet_size as usize - HEADER_SIZE;
885 +
886 + let remaining_after_first = hdr.payload_len as usize - first_payload_bytes;
887 + let expected_continuations = if remaining_after_first > 0 && chunk_payload_budget > 0 {
888 + (remaining_after_first + chunk_payload_budget - 1) / chunk_payload_budget
889 + } else {
890 + 0
891 + };
892 + let expected_chunk_count = 1 + expected_continuations as u32;
893 +
894 + if self.pkt_buf.len() < self.packet_size as usize {
895 + self.pkt_buf.resize(self.packet_size as usize, 0);
896 + }
897 +
898 + let mut ci: u32 = 1;
899 + while assembled < hdr.payload_len as usize {
900 + let cn = match raw_recv(self.handle, &mut self.pkt_buf) {
901 + Ok(n) => n,
902 + Err(err) => {
903 + self.fail_all_inflight();
904 + return Err(err);
905 + }
906 + };
907 +
908 + if cn < HEADER_SIZE {
909 + return Err(NpError::Chunk("continuation too short".into()));
910 + }
911 +
912 + let chk = ChunkHeader::decode(&self.pkt_buf[..cn])
913 + .map_err(|e| NpError::Chunk(format!("chunk header: {e}")))?;
914 +
915 + if chk.message_id != hdr.message_id {
916 + return Err(NpError::Chunk("message_id mismatch".into()));
917 + }
918 + if chk.chunk_index != ci {
919 + return Err(NpError::Chunk(format!(
920 + "chunk_index mismatch: expected {ci}, got {}",
921 + chk.chunk_index
922 + )));
923 + }
924 + if chk.chunk_count != expected_chunk_count {
925 + return Err(NpError::Chunk("chunk_count mismatch".into()));
926 + }
927 + if chk.total_message_len != total_msg as u32 {
928 + return Err(NpError::Chunk("total_message_len mismatch".into()));
929 + }
930 +
931 + let chunk_data = cn - HEADER_SIZE;
932 + if chunk_data != chk.chunk_payload_len as usize {
933 + return Err(NpError::Chunk("chunk_payload_len mismatch".into()));
934 + }
935 + if assembled + chunk_data > hdr.payload_len as usize {
936 + return Err(NpError::Chunk("chunk exceeds payload_len".into()));
937 + }
938 +
939 + self.recv_buf[assembled..assembled + chunk_data]
940 + .copy_from_slice(&self.pkt_buf[HEADER_SIZE..HEADER_SIZE + chunk_data]);
941 + assembled += chunk_data;
942 + ci += 1;
943 + }
944 +
945 + let payload = &self.recv_buf[..hdr.payload_len as usize];
946 +
947 + // Validate batch directory
948 + if hdr.flags & FLAG_BATCH != 0 && hdr.item_count > 1 {
949 + let dir_bytes = hdr.item_count as usize * 8;
950 + let dir_aligned = align8(dir_bytes);
951 + if payload.len() < dir_aligned {
952 + return Err(NpError::Protocol("batch directory exceeds payload".into()));
953 + }
954 + let packed_area_len = (payload.len() - dir_aligned) as u32;
955 + protocol::batch_dir_validate(&payload[..dir_bytes], hdr.item_count, packed_area_len)
956 + .map_err(|e| NpError::Protocol(format!("batch directory: {e:?}")))?;
957 + }
958 +
959 + Ok((hdr, payload))
960 + }
961 +
962 + /// Close the session.
963 + pub fn close(&mut self) {
964 + if self.handle != ffi::INVALID_HANDLE_VALUE && self.handle != 0 {
965 + // Flush pending writes so the peer reads all data
966 + unsafe {
967 + ffi::FlushFileBuffers(self.handle);
968 + }
969 + if self.role == Role::Server {
970 + unsafe {
971 + ffi::DisconnectNamedPipe(self.handle);
972 + }
973 + }
974 + close_handle(self.handle);
975 + self.handle = ffi::INVALID_HANDLE_VALUE;
976 + }
977 + self.fail_all_inflight();
978 + self.recv_buf.clear();
979 + }
980 +}
981 +
982 +#[cfg(windows)]
983 +impl Drop for NpSession {
984 + fn drop(&mut self) {
985 + self.close();
986 + }
987 +}
988 +
989 +// ---------------------------------------------------------------------------
990 +// Listener
991 +// ---------------------------------------------------------------------------
992 +
993 +/// A Named Pipe listener that accepts client connections.
994 +#[cfg(windows)]
995 +pub struct NpListener {
996 + handle: ffi::HANDLE,
997 + config: ServerConfig,
998 + pipe_name: Vec<u16>,
999 + next_session_id: AtomicU64,
1000 +}
1001 +
1002 +#[cfg(windows)]
1003 +impl NpListener {
1004 + /// Create a listener on a Named Pipe derived from run_dir + service_name.
1005 + pub fn bind(run_dir: &str, service_name: &str, config: ServerConfig) -> Result<Self, NpError> {
1006 + let pipe_name = build_pipe_name(run_dir, service_name)?;
1007 + let buf_size = pipe_buffer_size(config.packet_size);
1008 +
1009 + // Create first instance with FILE_FLAG_FIRST_PIPE_INSTANCE
1010 + let handle = create_pipe_instance(&pipe_name, buf_size, true)?;
1011 +
1012 + Ok(Self {
1013 + handle,
1014 + config,
1015 + pipe_name,
1016 + next_session_id: AtomicU64::new(1),
1017 + })
1018 + }
1019 +
1020 + /// Get the raw HANDLE.
1021 + pub fn handle(&self) -> ffi::HANDLE {
1022 + self.handle
1023 + }
1024 +
1025 + /// Update the payload limits used for future handshakes.
1026 + pub fn set_payload_limits(
1027 + &mut self,
1028 + max_request_payload_bytes: u32,
1029 + max_response_payload_bytes: u32,
1030 + ) {
1031 + self.config.max_request_payload_bytes = max_request_payload_bytes;
1032 + self.config.max_response_payload_bytes = max_response_payload_bytes;
1033 + }
1034 +
1035 + /// Accept one client connection. Performs the full handshake.
1036 + pub fn accept(&mut self) -> Result<NpSession, NpError> {
1037 + let session_id = self.next_session_id.fetch_add(1, Ordering::Relaxed);
1038 + self.accept_with_config(session_id, self.config.clone())
1039 + }
1040 +
1041 + /// Accept one client using a caller-provided per-session server config
1042 + /// and session ID.
1043 + pub fn accept_with_config(
1044 + &mut self,
1045 + session_id: u64,
1046 + config: ServerConfig,
1047 + ) -> Result<NpSession, NpError> {
1048 + // Wait for client
1049 + let connected = unsafe { ffi::ConnectNamedPipe(self.handle, ptr::null_mut()) };
1050 + if connected == 0 {
1051 + let err = last_error();
1052 + if err != ffi::ERROR_PIPE_CONNECTED {
1053 + return Err(NpError::Accept(err));
1054 + }
1055 + }
1056 +
1057 + let session_handle = self.handle;
1058 +
1059 + // Create new pipe instance for next client
1060 + let buf_size = pipe_buffer_size(self.config.packet_size);
1061 + let next = match create_pipe_instance(&self.pipe_name, buf_size, false) {
1062 + Ok(h) => h,
1063 + Err(e) => {
1064 + // Failed to create replacement instance; disconnect and close
1065 + // the accepted client to avoid an orphaned connection.
1066 + unsafe {
1067 + ffi::DisconnectNamedPipe(session_handle);
1068 + }
1069 + close_handle(session_handle);
1070 + self.handle = ffi::INVALID_HANDLE_VALUE;
1071 + return Err(e);
1072 + }
1073 + };
1074 + self.handle = next;
1075 +
1076 + // Perform handshake
1077 + match server_handshake(session_handle, &config, session_id) {
1078 + Ok(session) => Ok(session),
1079 + Err(e) => {
1080 + unsafe {
1081 + ffi::DisconnectNamedPipe(session_handle);
1082 + }
1083 + close_handle(session_handle);
1084 + Err(e)
1085 + }
1086 + }
1087 + }
1088 +
1089 + /// Close the listener.
1090 + pub fn close(&mut self) {
1091 + close_handle(self.handle);
1092 + self.handle = ffi::INVALID_HANDLE_VALUE;
1093 + }
1094 +}
1095 +
1096 +#[cfg(windows)]
1097 +impl Drop for NpListener {
1098 + fn drop(&mut self) {
1099 + self.close();
1100 + }
1101 +}
1102 +
1103 +// ---------------------------------------------------------------------------
1104 +// Pipe instance creation
1105 +// ---------------------------------------------------------------------------
1106 +
1107 +#[cfg(windows)]
1108 +fn create_pipe_instance(
1109 + pipe_name: &[u16],
1110 + buf_size: u32,
1111 + first_instance: bool,
1112 +) -> Result<ffi::HANDLE, NpError> {
1113 + let mut open_mode = ffi::PIPE_ACCESS_DUPLEX;
1114 + if first_instance {
1115 + open_mode |= ffi::FILE_FLAG_FIRST_PIPE_INSTANCE;
1116 + }
1117 +
1118 + let handle = unsafe {
1119 + ffi::CreateNamedPipeW(
1120 + pipe_name.as_ptr(),
1121 + open_mode,
1122 + ffi::PIPE_TYPE_MESSAGE | ffi::PIPE_READMODE_MESSAGE | ffi::PIPE_WAIT,
1123 + ffi::PIPE_UNLIMITED_INSTANCES,
1124 + buf_size,
1125 + buf_size,
1126 + 0,
1127 + ptr::null(),
1128 + )
1129 + };
1130 +
1131 + if handle == ffi::INVALID_HANDLE_VALUE {
1132 + let err = last_error();
1133 + if err == ffi::ERROR_ACCESS_DENIED || err == ffi::ERROR_PIPE_BUSY {
1134 + return Err(NpError::AddrInUse);
1135 + }
1136 + return Err(NpError::CreatePipe(err));
1137 + }
1138 +
1139 + Ok(handle)
1140 +}
1141 +
1142 +// ---------------------------------------------------------------------------
1143 +// Client handshake
1144 +// ---------------------------------------------------------------------------
1145 +
1146 +#[cfg(windows)]
1147 +fn client_handshake(handle: ffi::HANDLE, config: &ClientConfig) -> Result<NpSession, NpError> {
1148 + let pkt_size = apply_default(config.packet_size, DEFAULT_PACKET_SIZE);
1149 +
1150 + let supported = if config.supported_profiles != 0 {
1151 + config.supported_profiles
1152 + } else {
1153 + PROFILE_BASELINE
1154 + };
1155 +
1156 + let hello = Hello {
1157 + layout_version: 1,
1158 + flags: 0,
1159 + supported_profiles: supported,
1160 + preferred_profiles: config.preferred_profiles,
1161 + max_request_payload_bytes: apply_default(
1162 + config.max_request_payload_bytes,
1163 + MAX_PAYLOAD_DEFAULT,
1164 + ),
1165 + max_request_batch_items: apply_default(config.max_request_batch_items, DEFAULT_BATCH_ITEMS),
1166 + max_response_payload_bytes: apply_default(
1167 + config.max_response_payload_bytes,
1168 + MAX_PAYLOAD_DEFAULT,
1169 + ),
1170 + max_response_batch_items: apply_default(
1171 + config.max_response_batch_items,
1172 + DEFAULT_BATCH_ITEMS,
1173 + ),
1174 + auth_token: config.auth_token,
1175 + packet_size: pkt_size,
1176 + };
1177 +
1178 + let mut hello_buf = [0u8; HELLO_PAYLOAD_SIZE];
1179 + hello.encode(&mut hello_buf);
1180 +
1181 + let hdr = Header {
1182 + magic: MAGIC_MSG,
1183 + version: VERSION,
1184 + header_len: protocol::HEADER_LEN,
1185 + kind: protocol::KIND_CONTROL,
1186 + flags: 0,
1187 + code: protocol::CODE_HELLO,
1188 + transport_status: protocol::STATUS_OK,
1189 + payload_len: HELLO_PAYLOAD_SIZE as u32,
1190 + item_count: 1,
1191 + message_id: 0,
1192 + };
1193 +
1194 + let mut pkt = [0u8; HEADER_SIZE + HELLO_PAYLOAD_SIZE];
1195 + hdr.encode(&mut pkt[..HEADER_SIZE]);
1196 + pkt[HEADER_SIZE..].copy_from_slice(&hello_buf);
1197 +
1198 + raw_write(handle, &pkt)?;
1199 +
1200 + // Receive HELLO_ACK
1201 + let mut ack_buf = [0u8; 128];
1202 + let n = raw_recv(handle, &mut ack_buf)?;
1203 +
1204 + let ack_hdr = match Header::decode(&ack_buf[..n]) {
1205 + Ok(hdr) => hdr,
1206 + Err(crate::protocol::NipcError::BadVersion) => {
1207 + return Err(NpError::Incompatible("ack header version mismatch".into()))
1208 + }
1209 + Err(e) => return Err(NpError::Protocol(format!("ack header: {e}"))),
1210 + };
1211 +
1212 + if ack_hdr.kind != protocol::KIND_CONTROL || ack_hdr.code != protocol::CODE_HELLO_ACK {
1213 + return Err(NpError::Protocol("expected HELLO_ACK".into()));
1214 + }
1215 +
1216 + if ack_hdr.transport_status == protocol::STATUS_AUTH_FAILED {
1217 + return Err(NpError::AuthFailed);
1218 + }
1219 + if ack_hdr.transport_status == protocol::STATUS_UNSUPPORTED {
1220 + return Err(NpError::NoProfile);
1221 + }
1222 + if ack_hdr.transport_status == protocol::STATUS_INCOMPATIBLE {
1223 + return Err(NpError::Incompatible(
1224 + "ack transport_status incompatible".into(),
1225 + ));
1226 + }
1227 + if ack_hdr.transport_status == protocol::STATUS_LIMIT_EXCEEDED {
1228 + return Err(NpError::LimitExceeded);
1229 + }
1230 + if ack_hdr.transport_status != protocol::STATUS_OK {
1231 + return Err(NpError::Handshake(format!(
1232 + "transport_status={}",
1233 + ack_hdr.transport_status
1234 + )));
1235 + }
1236 +
1237 + if n < HEADER_SIZE + HELLO_ACK_PAYLOAD_SIZE {
1238 + return Err(NpError::Protocol("ack payload truncated".into()));
1239 + }
1240 +
1241 + let ack = match HelloAck::decode(&ack_buf[HEADER_SIZE..n]) {
1242 + Ok(ack) => ack,
1243 + Err(crate::protocol::NipcError::BadLayout)
1244 + if hello_ack_layout_incompatible(&ack_buf[HEADER_SIZE..n]) =>
1245 + {
1246 + return Err(NpError::Incompatible(
1247 + "ack payload layout version mismatch".into(),
1248 + ))
1249 + }
1250 + Err(e) => return Err(NpError::Protocol(format!("ack payload: {e}"))),
1251 + };
1252 +
1253 + // Reject packet sizes too small for a header
1254 + if ack.agreed_packet_size <= HEADER_SIZE as u32 {
1255 + return Err(NpError::Handshake(
1256 + "agreed packet_size <= HEADER_SIZE".into(),
1257 + ));
1258 + }
1259 +
1260 + Ok(NpSession {
1261 + handle,
1262 + role: Role::Client,
1263 + max_request_payload_bytes: ack.agreed_max_request_payload_bytes,
1264 + max_request_batch_items: ack.agreed_max_request_batch_items,
1265 + max_response_payload_bytes: ack.agreed_max_response_payload_bytes,
1266 + max_response_batch_items: ack.agreed_max_response_batch_items,
1267 + packet_size: ack.agreed_packet_size,
1268 + selected_profile: ack.selected_profile,
1269 + session_id: ack.session_id,
1270 + recv_buf: Vec::new(),
1271 + pkt_buf: Vec::new(),
1272 + send_buf: Vec::new(),
1273 + inflight_ids: HashSet::new(),
1274 + })
1275 +}
1276 +
1277 +// ---------------------------------------------------------------------------
1278 +// Server handshake
1279 +// ---------------------------------------------------------------------------
1280 +
1281 +#[cfg(windows)]
1282 +fn server_handshake(
1283 + handle: ffi::HANDLE,
1284 + config: &ServerConfig,
1285 + session_id: u64,
1286 +) -> Result<NpSession, NpError> {
1287 + let server_pkt_size = apply_default(config.packet_size, DEFAULT_PACKET_SIZE);
1288 + let s_resp_pay = apply_default(config.max_response_payload_bytes, MAX_PAYLOAD_DEFAULT);
1289 + let s_profiles = if config.supported_profiles != 0 {
1290 + config.supported_profiles
1291 + } else {
1292 + PROFILE_BASELINE
1293 + };
1294 + let s_preferred = config.preferred_profiles;
1295 +
1296 + let send_rejection = |status: u16| {
1297 + let ack = HelloAck {
1298 + layout_version: 1,
1299 + ..HelloAck::default()
1300 + };
1301 + let mut ack_buf = [0u8; HELLO_ACK_PAYLOAD_SIZE];
1302 + ack.encode(&mut ack_buf);
1303 +
1304 + let ack_hdr = Header {
1305 + magic: MAGIC_MSG,
1306 + version: VERSION,
1307 + header_len: protocol::HEADER_LEN,
1308 + kind: protocol::KIND_CONTROL,
1309 + flags: 0,
1310 + code: protocol::CODE_HELLO_ACK,
1311 + transport_status: status,
1312 + payload_len: HELLO_ACK_PAYLOAD_SIZE as u32,
1313 + item_count: 1,
1314 + message_id: 0,
1315 + };
1316 +
1317 + let mut pkt = [0u8; HEADER_SIZE + HELLO_ACK_PAYLOAD_SIZE];
1318 + ack_hdr.encode(&mut pkt[..HEADER_SIZE]);
1319 + pkt[HEADER_SIZE..].copy_from_slice(&ack_buf);
1320 + let _ = raw_write(handle, &pkt);
1321 + };
1322 +
1323 + // Receive HELLO
1324 + let mut buf = [0u8; 128];
1325 + let n = raw_recv(handle, &mut buf)?;
1326 +
1327 + let hdr = match Header::decode(&buf[..n]) {
1328 + Ok(hdr) => hdr,
1329 + Err(crate::protocol::NipcError::BadVersion)
1330 + if header_version_incompatible(&buf[..n], protocol::CODE_HELLO) =>
1331 + {
1332 + send_rejection(protocol::STATUS_INCOMPATIBLE);
1333 + return Err(NpError::Incompatible(
1334 + "hello header version mismatch".into(),
1335 + ));
1336 + }
1337 + Err(e) => return Err(NpError::Protocol(format!("hello header: {e}"))),
1338 + };
1339 +
1340 + if hdr.kind != protocol::KIND_CONTROL || hdr.code != protocol::CODE_HELLO {
1341 + return Err(NpError::Protocol("expected HELLO".into()));
1342 + }
1343 +
1344 + let hello = match Hello::decode(&buf[HEADER_SIZE..n]) {
1345 + Ok(hello) => hello,
1346 + Err(crate::protocol::NipcError::BadLayout)
1347 + if hello_layout_incompatible(&buf[HEADER_SIZE..n]) =>
1348 + {
1349 + send_rejection(protocol::STATUS_INCOMPATIBLE);
1350 + return Err(NpError::Incompatible(
1351 + "hello payload layout version mismatch".into(),
1352 + ));
1353 + }
1354 + Err(e) => return Err(NpError::Protocol(format!("hello payload: {e}"))),
1355 + };
1356 +
1357 + let intersection = hello.supported_profiles & s_profiles;
1358 +
1359 + if intersection == 0 {
1360 + send_rejection(protocol::STATUS_UNSUPPORTED);
1361 + return Err(NpError::NoProfile);
1362 + }
1363 +
1364 + if hello.auth_token != config.auth_token {
1365 + send_rejection(protocol::STATUS_AUTH_FAILED);
1366 + return Err(NpError::AuthFailed);
1367 + }
1368 +
1369 + // Select profile
1370 + let preferred_intersection = intersection & hello.preferred_profiles & s_preferred;
1371 + let selected = if preferred_intersection != 0 {
1372 + highest_bit(preferred_intersection)
1373 + } else {
1374 + highest_bit(intersection)
1375 + };
1376 +
1377 + if hello.max_request_payload_bytes > MAX_PAYLOAD_CAP {
1378 + send_rejection(protocol::STATUS_LIMIT_EXCEEDED);
1379 + return Err(NpError::LimitExceeded);
1380 + }
1381 +
1382 + // Negotiate limits:
1383 + // - request payload and batch size are client-proposed and echoed
1384 + // - response payload is server-authoritative
1385 + // - response batch size is symmetric with request batch size
1386 + let agreed_req_pay = hello.max_request_payload_bytes;
1387 + let agreed_req_bat = hello.max_request_batch_items;
1388 + let agreed_resp_pay = s_resp_pay;
1389 + let agreed_resp_bat = agreed_req_bat;
1390 + let agreed_pkt = min_u32(hello.packet_size, server_pkt_size);
1391 +
1392 + // Reject packet sizes too small for a usable message packet
1393 + if agreed_pkt <= HEADER_SIZE as u32 {
1394 + send_rejection(protocol::STATUS_INCOMPATIBLE);
1395 + return Err(NpError::Incompatible(
1396 + "packet size too small for negotiated session".into(),
1397 + ));
1398 + }
1399 +
1400 + // Send HELLO_ACK
1401 + let ack = HelloAck {
1402 + layout_version: 1,
1403 + flags: 0,
1404 + server_supported_profiles: s_profiles,
1405 + intersection_profiles: intersection,
1406 + selected_profile: selected,
1407 + agreed_max_request_payload_bytes: agreed_req_pay,
1408 + agreed_max_request_batch_items: agreed_req_bat,
1409 + agreed_max_response_payload_bytes: agreed_resp_pay,
1410 + agreed_max_response_batch_items: agreed_resp_bat,
1411 + agreed_packet_size: agreed_pkt,
1412 + session_id,
1413 + };
1414 +
1415 + let mut ack_buf = [0u8; HELLO_ACK_PAYLOAD_SIZE];
1416 + ack.encode(&mut ack_buf);
1417 +
1418 + let ack_hdr = Header {
1419 + magic: MAGIC_MSG,
1420 + version: VERSION,
1421 + header_len: protocol::HEADER_LEN,
1422 + kind: protocol::KIND_CONTROL,
1423 + flags: 0,
1424 + code: protocol::CODE_HELLO_ACK,
1425 + transport_status: protocol::STATUS_OK,
1426 + payload_len: HELLO_ACK_PAYLOAD_SIZE as u32,
1427 + item_count: 1,
1428 + message_id: 0,
1429 + };
1430 +
1431 + let mut pkt = [0u8; HEADER_SIZE + HELLO_ACK_PAYLOAD_SIZE];
1432 + ack_hdr.encode(&mut pkt[..HEADER_SIZE]);
1433 + pkt[HEADER_SIZE..].copy_from_slice(&ack_buf);
1434 +
1435 + raw_write(handle, &pkt)?;
1436 +
1437 + Ok(NpSession {
1438 + handle,
1439 + role: Role::Server,
1440 + max_request_payload_bytes: agreed_req_pay,
1441 + max_request_batch_items: agreed_req_bat,
1442 + max_response_payload_bytes: agreed_resp_pay,
1443 + max_response_batch_items: agreed_resp_bat,
1444 + packet_size: agreed_pkt,
1445 + selected_profile: selected,
1446 + session_id,
1447 + recv_buf: Vec::new(),
1448 + pkt_buf: Vec::new(),
1449 + send_buf: Vec::new(),
1450 + inflight_ids: HashSet::new(),
1451 + })
1452 +}
1453 +
1454 +// ---------------------------------------------------------------------------
1455 +// Tests (cross-platform unit tests for non-Win32 logic)
1456 +// ---------------------------------------------------------------------------
1457 +
1458 +#[cfg(test)]
1459 +mod tests {
1460 + use super::*;
1461 + #[cfg(windows)]
1462 + use std::sync::atomic::{AtomicU64, Ordering};
1463 + #[cfg(windows)]
1464 + use std::thread;
1465 +
1466 + #[cfg(windows)]
1467 + const TEST_RUN_DIR: &str = r"C:\Temp\nipc_transport_rust_test";
1468 + #[cfg(windows)]
1469 + static TEST_COUNTER: AtomicU64 = AtomicU64::new(0);
1470 +
1471 + #[cfg(windows)]
1472 + fn ensure_run_dir() {
1473 + let _ = std::fs::create_dir_all(TEST_RUN_DIR);
1474 + }
1475 +
1476 + #[cfg(windows)]
1477 + fn unique_service(prefix: &str) -> String {
1478 + format!(
1479 + "{}_{}_{}",
1480 + prefix,
1481 + std::process::id(),
1482 + TEST_COUNTER.fetch_add(1, Ordering::Relaxed) + 1
1483 + )
1484 + }
1485 +
1486 + #[cfg(windows)]
1487 + fn default_client_config() -> ClientConfig {
1488 + ClientConfig {
1489 + supported_profiles: PROFILE_BASELINE,
1490 + preferred_profiles: 0,
1491 + max_request_payload_bytes: 4096,
1492 + max_request_batch_items: 16,
1493 + max_response_payload_bytes: 4096,
1494 + max_response_batch_items: 16,
1495 + auth_token: 0xDEADBEEFCAFEBABE,
1496 + packet_size: 0,
1497 + }
1498 + }
1499 +
1500 + #[cfg(windows)]
1501 + fn default_server_config() -> ServerConfig {
1502 + ServerConfig {
1503 + supported_profiles: PROFILE_BASELINE,
1504 + preferred_profiles: 0,
1505 + max_request_payload_bytes: 4096,
1506 + max_request_batch_items: 16,
1507 + max_response_payload_bytes: 4096,
1508 + max_response_batch_items: 16,
1509 + auth_token: 0xDEADBEEFCAFEBABE,
1510 + packet_size: 0,
1511 + }
1512 + }
1513 +
1514 + #[test]
1515 + fn test_fnv1a_64_empty() {
1516 + assert_eq!(fnv1a_64(b""), FNV1A_OFFSET_BASIS);
1517 + }
1518 +
1519 + #[test]
1520 + fn test_fnv1a_64_known() {
1521 + // FNV-1a of "foobar" — verified against reference implementation
1522 + let hash = fnv1a_64(b"foobar");
1523 + assert_ne!(hash, 0);
1524 + assert_ne!(hash, FNV1A_OFFSET_BASIS);
1525 + }
1526 +
1527 + #[test]
1528 + fn test_fnv1a_64_deterministic() {
1529 + let h1 = fnv1a_64(b"/var/run/netdata");
1530 + let h2 = fnv1a_64(b"/var/run/netdata");
1531 + assert_eq!(h1, h2);
1532 + }
1533 +
1534 + #[test]
1535 + fn test_fnv1a_64_different_inputs() {
1536 + let h1 = fnv1a_64(b"/var/run/netdata");
1537 + let h2 = fnv1a_64(b"/tmp/netdata");
1538 + assert_ne!(h1, h2);
1539 + }
1540 +
1541 + #[test]
1542 + fn test_validate_service_name_valid() {
1543 + assert!(validate_service_name("cgroups-snapshot").is_ok());
1544 + assert!(validate_service_name("test_service.v1").is_ok());
1545 + assert!(validate_service_name("A-Z_09").is_ok());
1546 + }
1547 +
1548 + #[test]
1549 + fn test_validate_service_name_invalid() {
1550 + assert!(validate_service_name("").is_err());
1551 + assert!(validate_service_name(".").is_err());
1552 + assert!(validate_service_name("..").is_err());
1553 + assert!(validate_service_name("has space").is_err());
1554 + assert!(validate_service_name("has/slash").is_err());
1555 + assert!(validate_service_name("has\\backslash").is_err());
1556 + }
1557 +
1558 + #[test]
1559 + fn test_build_pipe_name() {
1560 + let name = build_pipe_name("/var/run/netdata", "cgroups-snapshot").unwrap();
1561 + // Should produce a valid wide string ending in NUL
1562 + assert!(*name.last().unwrap() == 0);
1563 + // Convert back to narrow for checking prefix
1564 + let narrow: String = name[..name.len() - 1]
1565 + .iter()
1566 + .map(|&c| c as u8 as char)
1567 + .collect();
1568 + assert!(narrow.starts_with("\\\\.\\pipe\\netipc-"));
1569 + assert!(narrow.ends_with("-cgroups-snapshot"));
1570 + // Hash should be 16 hex chars
1571 + let parts: Vec<&str> = narrow.split('-').collect();
1572 + // \\.\pipe\netipc - {hash} - cgroups - snapshot
1573 + // parts: ["\\\\.", "\\pipe\\netipc", "{hash}", "cgroups", "snapshot"]
1574 + // Actually the split is on '-' so:
1575 + // "\\\\.\\pipe\\netipc" - "hash" - "cgroups" - "snapshot"
1576 + assert!(parts.len() >= 3);
1577 + // The pipe name is \\.\pipe\netipc-{hash}-{service}, so parts[1]
1578 + // is the 16-character hash component.
1579 + assert_eq!(parts[1].len(), 16, "hash should be 16 hex chars");
1580 + }
1581 +
1582 + #[test]
1583 + fn test_build_pipe_name_invalid_service() {
1584 + assert!(build_pipe_name("/var/run", "").is_err());
1585 + assert!(build_pipe_name("/var/run", "bad/name").is_err());
1586 + assert!(build_pipe_name("/var/run", ".").is_err());
1587 + }
1588 +
1589 + #[test]
1590 + fn test_pipe_name_deterministic() {
1591 + let n1 = build_pipe_name("/var/run/netdata", "test-svc").unwrap();
1592 + let n2 = build_pipe_name("/var/run/netdata", "test-svc").unwrap();
1593 + assert_eq!(n1, n2);
1594 + }
1595 +
1596 + #[test]
1597 + fn test_pipe_name_different_run_dir() {
1598 + let n1 = build_pipe_name("/var/run/netdata", "svc").unwrap();
1599 + let n2 = build_pipe_name("/tmp/netdata", "svc").unwrap();
1600 + assert_ne!(n1, n2);
1601 + }
1602 +
1603 + #[test]
1604 + fn test_np_error_display() {
1605 + let cases = [
1606 + (NpError::PipeName("bad".into()), "pipe name error: bad"),
1607 + (NpError::CreatePipe(5), "CreateNamedPipeW failed: 5"),
1608 + (NpError::Connect(231), "connect failed: 231"),
1609 + (NpError::Accept(24), "accept failed: 24"),
1610 + (NpError::Send(32), "send failed: 32"),
1611 + (NpError::Recv(0), "recv failed: 0"),
1612 + (NpError::Handshake("test".into()), "handshake error: test"),
1613 + (NpError::AuthFailed, "authentication token rejected"),
1614 + (NpError::NoProfile, "no common transport profile"),
1615 + (
1616 + NpError::Incompatible("version mismatch".into()),
1617 + "incompatible protocol: version mismatch",
1618 + ),
1619 + (NpError::Protocol("bad".into()), "protocol violation: bad"),
1620 + (
1621 + NpError::AddrInUse,
1622 + "pipe name already in use by live server",
1623 + ),
1624 + (NpError::Chunk("mismatch".into()), "chunk error: mismatch"),
1625 + (NpError::Alloc, "memory allocation failed"),
1626 + (NpError::LimitExceeded, "negotiated limit exceeded"),
1627 + (NpError::BadParam("foo".into()), "bad parameter: foo"),
1628 + (NpError::DuplicateMsgId(42), "duplicate message_id: 42"),
1629 + (NpError::UnknownMsgId(99), "unknown response message_id: 99"),
1630 + (NpError::Disconnected, "peer disconnected"),
1631 + ];
1632 +
1633 + for (err, expected) in cases {
1634 + assert_eq!(format!("{err}"), expected);
1635 + }
1636 +
1637 + let e: &dyn std::error::Error = &NpError::Disconnected;
1638 + let _ = format!("{e}");
1639 + }
1640 +
1641 + #[test]
1642 + fn test_incompatible_classifiers() {
1643 + let hdr = Header {
1644 + magic: MAGIC_MSG,
1645 + version: VERSION + 1,
1646 + header_len: protocol::HEADER_LEN,
1647 + kind: protocol::KIND_CONTROL,
1648 + flags: 0,
1649 + code: protocol::CODE_HELLO,
1650 + transport_status: protocol::STATUS_OK,
1651 + payload_len: HELLO_PAYLOAD_SIZE as u32,
1652 + item_count: 1,
1653 + message_id: 0,
1654 + };
1655 + let mut hdr_buf = [0u8; HEADER_SIZE];
1656 + hdr.encode(&mut hdr_buf);
1657 + assert!(header_version_incompatible(&hdr_buf, protocol::CODE_HELLO));
1658 + assert!(!header_version_incompatible(
1659 + &hdr_buf,
1660 + protocol::CODE_HELLO_ACK
1661 + ));
1662 +
1663 + let hello = Hello {
1664 + layout_version: 2,
1665 + ..Hello::default()
1666 + };
1667 + let mut hello_buf = [0u8; HELLO_PAYLOAD_SIZE];
1668 + hello.encode(&mut hello_buf);
1669 + assert!(hello_layout_incompatible(&hello_buf));
1670 +
1671 + let ack = HelloAck {
1672 + layout_version: 2,
1673 + ..HelloAck::default()
1674 + };
1675 + let mut ack_buf = [0u8; HELLO_ACK_PAYLOAD_SIZE];
1676 + ack.encode(&mut ack_buf);
1677 + assert!(hello_ack_layout_incompatible(&ack_buf));
1678 + }
1679 +
1680 + #[test]
1681 + fn test_highest_bit() {
1682 + assert_eq!(highest_bit(0), 0);
1683 + assert_eq!(highest_bit(1), 1);
1684 + assert_eq!(highest_bit(0b0101), 4);
1685 + assert_eq!(highest_bit(0b1000), 8);
1686 + assert_eq!(highest_bit(0xFF), 128);
1687 + }
1688 +
1689 + #[test]
1690 + fn test_build_pipe_name_too_long() {
1691 + let long_name = "a".repeat(300);
1692 + let result = build_pipe_name("/var/run", &long_name);
1693 + assert!(matches!(result, Err(NpError::PipeName(_))));
1694 + }
1695 +
1696 + #[cfg(windows)]
1697 + #[test]
1698 + fn test_connect_nonexistent() {
1699 + ensure_run_dir();
1700 + let svc = unique_service("rs_noexist");
1701 + let result = NpSession::connect(TEST_RUN_DIR, &svc, &default_client_config());
1702 + assert!(matches!(result, Err(NpError::Connect(_))));
1703 + }
1704 +
1705 + #[cfg(windows)]
1706 + #[test]
1707 + fn test_send_receive_on_closed_session() {
1708 + let mut session = NpSession {
1709 + handle: ffi::INVALID_HANDLE_VALUE,
1710 + role: Role::Client,
1711 + max_request_payload_bytes: 4096,
1712 + max_request_batch_items: 16,
1713 + max_response_payload_bytes: 4096,
1714 + max_response_batch_items: 16,
1715 + packet_size: DEFAULT_PACKET_SIZE,
1716 + selected_profile: PROFILE_BASELINE,
1717 + session_id: 1,
1718 + recv_buf: Vec::new(),
1719 + pkt_buf: Vec::new(),
1720 + send_buf: Vec::new(),
1721 + inflight_ids: HashSet::new(),
1722 + };
1723 +
1724 + let mut hdr = Header {
1725 + kind: protocol::KIND_REQUEST,
1726 + code: protocol::METHOD_INCREMENT,
1727 + item_count: 1,
1728 + message_id: 1,
1729 + ..Header::default()
1730 + };
1731 + assert!(matches!(
1732 + session.send(&mut hdr, &[1, 2, 3]),
1733 + Err(NpError::BadParam(_))
1734 + ));
1735 +
1736 + let mut buf = [0u8; 64];
1737 + assert!(matches!(
1738 + session.receive(&mut buf),
1739 + Err(NpError::BadParam(_))
1740 + ));
1741 + }
1742 +
1743 + #[cfg(windows)]
1744 + #[test]
1745 + fn test_roundtrip_and_accessors() {
1746 + ensure_run_dir();
1747 + let svc = unique_service("rs_roundtrip");
1748 +
1749 + let mut listener =
1750 + NpListener::bind(TEST_RUN_DIR, &svc, default_server_config()).expect("bind");
1751 + assert_ne!(listener.handle(), ffi::INVALID_HANDLE_VALUE);
1752 + assert_ne!(listener.handle(), 0);
1753 +
1754 + let server = thread::spawn(move || {
1755 + let mut session = listener.accept().expect("accept");
1756 + assert_eq!(session.role(), Role::Server);
1757 + assert_ne!(session.handle(), ffi::INVALID_HANDLE_VALUE);
1758 + let mut buf = [0u8; 256];
1759 + let (hdr, payload) = session.receive(&mut buf).expect("recv");
1760 + let payload = payload.to_vec();
1761 + let mut resp = hdr;
1762 + resp.kind = KIND_RESPONSE;
1763 + resp.transport_status = protocol::STATUS_OK;
1764 + session.send(&mut resp, &payload).expect("send");
1765 + session.close();
1766 + });
1767 +
1768 + let mut session =
1769 + NpSession::connect(TEST_RUN_DIR, &svc, &default_client_config()).expect("connect");
1770 + assert_eq!(session.role(), Role::Client);
1771 + assert_ne!(session.handle(), ffi::INVALID_HANDLE_VALUE);
1772 + assert_eq!(session.selected_profile, PROFILE_BASELINE);
1773 +
1774 + let payload = [1u8, 2, 3, 4];
1775 + let mut hdr = Header {
1776 + kind: KIND_REQUEST,
1777 + code: protocol::METHOD_INCREMENT,
1778 + flags: 0,
1779 + item_count: 1,
1780 + message_id: 42,
1781 + ..Header::default()
1782 + };
1783 + session.send(&mut hdr, &payload).expect("send");
1784 +
1785 + let mut rbuf = [0u8; 256];
1786 + let (rhdr, rpayload) = session.receive(&mut rbuf).expect("recv");
1787 + assert_eq!(rhdr.kind, KIND_RESPONSE);
1788 + assert_eq!(rhdr.message_id, 42);
1789 + assert_eq!(rpayload, payload);
1790 +
1791 + session.close();
1792 + server.join().expect("server join");
1793 + }
1794 +
1795 + #[cfg(windows)]
1796 + #[test]
1797 + fn test_chunking_and_received_payload() {
1798 + ensure_run_dir();
1799 + let svc = unique_service("rs_chunk");
1800 +
1801 + let scfg = ServerConfig {
1802 + packet_size: 128,
1803 + max_request_payload_bytes: 65536,
1804 + max_response_payload_bytes: 65536,
1805 + ..default_server_config()
1806 + };
1807 + let mut listener = NpListener::bind(TEST_RUN_DIR, &svc, scfg).expect("bind");
1808 +
1809 + let server = thread::spawn(move || {
1810 + let mut session = listener.accept().expect("accept");
1811 + let mut buf = [0u8; 256];
1812 + let (hdr, payload) = session.receive(&mut buf).expect("recv");
1813 + let payload = payload.to_vec();
1814 + let mut resp = hdr;
1815 + resp.kind = KIND_RESPONSE;
1816 + resp.transport_status = protocol::STATUS_OK;
1817 + session.send(&mut resp, &payload).expect("send");
1818 + session.close();
1819 + });
1820 +
1821 + let ccfg = ClientConfig {
1822 + packet_size: 128,
1823 + max_request_payload_bytes: 65536,
1824 + max_response_payload_bytes: 65536,
1825 + ..default_client_config()
1826 + };
1827 + let mut session = NpSession::connect(TEST_RUN_DIR, &svc, &ccfg).expect("connect");
1828 + assert_eq!(session.packet_size, 128);
1829 +
1830 + let big: Vec<u8> = (0..500).map(|i| (i & 0xff) as u8).collect();
1831 + let mut hdr = Header {
1832 + kind: KIND_REQUEST,
1833 + code: protocol::METHOD_INCREMENT,
1834 + item_count: 1,
1835 + message_id: 7,
1836 + ..Header::default()
1837 + };
1838 + session.send(&mut hdr, &big).expect("send chunked");
1839 +
1840 + let mut rbuf = [0u8; 256];
1841 + let (rhdr, received_len) = {
1842 + let (rhdr, rpayload) = session.receive(&mut rbuf).expect("recv chunked");
1843 + assert_eq!(rpayload, big);
1844 + (rhdr, rpayload.len())
1845 + };
1846 + assert_eq!(rhdr.message_id, 7);
1847 + assert_eq!(session.received_payload(received_len), big.as_slice());
1848 +
1849 + session.close();
1850 + server.join().expect("server join");
1851 + }
1852 +
1853 + #[cfg(windows)]
1854 + #[test]
1855 + fn test_pipeline_chunked() {
1856 + ensure_run_dir();
1857 + let svc = unique_service("rs_pipe_chunked");
1858 + let forced_packet_size = 128u32;
1859 +
1860 + let scfg = ServerConfig {
1861 + packet_size: forced_packet_size,
1862 + max_request_payload_bytes: 65536,
1863 + max_response_payload_bytes: 65536,
1864 + ..default_server_config()
1865 + };
1866 + let mut listener = NpListener::bind(TEST_RUN_DIR, &svc, scfg).expect("bind");
1867 +
1868 + let server = thread::spawn(move || {
1869 + let mut session = listener.accept().expect("accept");
1870 + let mut buf = vec![0u8; forced_packet_size as usize];
1871 + for _ in 0..5 {
1872 + let (hdr, payload) = session.receive(&mut buf).expect("recv");
1873 + let payload = payload.to_vec();
1874 + let mut resp = hdr;
1875 + resp.kind = KIND_RESPONSE;
1876 + resp.transport_status = protocol::STATUS_OK;
1877 + session.send(&mut resp, &payload).expect("send");
1878 + }
1879 + session.close();
1880 + });
1881 +
1882 + let ccfg = ClientConfig {
1883 + packet_size: forced_packet_size,
1884 + max_request_payload_bytes: 65536,
1885 + max_response_payload_bytes: 65536,
1886 + ..default_client_config()
1887 + };
1888 + let mut session = NpSession::connect(TEST_RUN_DIR, &svc, &ccfg).expect("connect");
1889 +
1890 + let sizes = [200usize, 500, 300, 800, 150];
1891 + for (i, &sz) in sizes.iter().enumerate() {
1892 + let payload: Vec<u8> = (0..sz).map(|j| ((i + j) & 0xFF) as u8).collect();
1893 + let mut hdr = Header {
1894 + kind: KIND_REQUEST,
1895 + code: protocol::METHOD_INCREMENT,
1896 + item_count: 1,
1897 + message_id: (i + 1) as u64,
1898 + ..Header::default()
1899 + };
1900 + session.send(&mut hdr, &payload).expect("send");
1901 + }
1902 +
1903 + let mut rbuf = vec![0u8; forced_packet_size as usize];
1904 + for (i, &sz) in sizes.iter().enumerate() {
1905 + let (rhdr, rpayload) = session.receive(&mut rbuf).expect("recv");
1906 + assert_eq!(rhdr.message_id, (i + 1) as u64, "message_id at {i}");
1907 + assert_eq!(rpayload.len(), sz, "payload len at {i}");
1908 + let expected: Vec<u8> = (0..sz).map(|j| ((i + j) & 0xFF) as u8).collect();
1909 + assert_eq!(rpayload, expected, "payload data at {i}");
1910 + }
1911 +
1912 + session.close();
1913 + server.join().expect("server join");
1914 + }
1915 +
1916 + #[cfg(windows)]
1917 + #[test]
1918 + fn test_duplicate_message_id() {
1919 + ensure_run_dir();
1920 + let svc = unique_service("rs_dupmsg");
1921 + let mut listener =
1922 + NpListener::bind(TEST_RUN_DIR, &svc, default_server_config()).expect("bind");
1923 +
1924 + let server = thread::spawn(move || {
1925 + let mut session = listener.accept().expect("accept");
1926 + let mut buf = [0u8; 256];
1927 + let _ = session.receive(&mut buf).expect("recv");
1928 + session.close();
1929 + });
1930 +
1931 + let mut session =
1932 + NpSession::connect(TEST_RUN_DIR, &svc, &default_client_config()).expect("connect");
1933 +
1934 + let mut hdr1 = Header {
1935 + kind: KIND_REQUEST,
1936 + code: protocol::METHOD_INCREMENT,
1937 + item_count: 1,
1938 + message_id: 42,
1939 + ..Header::default()
1940 + };
1941 + session.send(&mut hdr1, &[1]).expect("first send");
1942 +
1943 + let mut hdr2 = Header {
1944 + kind: KIND_REQUEST,
1945 + code: protocol::METHOD_INCREMENT,
1946 + item_count: 1,
1947 + message_id: 42,
1948 + ..Header::default()
1949 + };
1950 + assert!(matches!(
1951 + session.send(&mut hdr2, &[2]),
1952 + Err(NpError::DuplicateMsgId(42))
1953 + ));
1954 +
1955 + session.close();
1956 + server.join().expect("server join");
1957 + }
1958 +
1959 + #[cfg(windows)]
1960 + #[test]
1961 + fn test_directional_limit_negotiation() {
1962 + ensure_run_dir();
1963 + let svc = unique_service("rs_dir_limits");
1964 +
1965 + let scfg = ServerConfig {
1966 + max_request_payload_bytes: 2048,
1967 + max_request_batch_items: 8,
1968 + max_response_payload_bytes: 8192,
1969 + max_response_batch_items: 32,
1970 + ..default_server_config()
1971 + };
1972 + let mut listener = NpListener::bind(TEST_RUN_DIR, &svc, scfg).expect("bind");
1973 +
1974 + let server = thread::spawn(move || {
1975 + let mut session = listener.accept().expect("accept");
1976 + assert_eq!(session.max_request_payload_bytes, 4096);
1977 + assert_eq!(session.max_request_batch_items, 16);
1978 + assert_eq!(session.max_response_payload_bytes, 8192);
1979 + assert_eq!(session.max_response_batch_items, 16);
1980 + assert_ne!(session.session_id, 0);
1981 + session.close();
1982 + });
1983 +
1984 + let ccfg = ClientConfig {
1985 + max_request_payload_bytes: 4096,
1986 + max_request_batch_items: 16,
1987 + max_response_payload_bytes: 4096,
1988 + max_response_batch_items: 16,
1989 + ..default_client_config()
1990 + };
1991 + let mut session = NpSession::connect(TEST_RUN_DIR, &svc, &ccfg).expect("connect");
1992 + assert_eq!(session.max_request_payload_bytes, 4096);
1993 + assert_eq!(session.max_request_batch_items, 16);
1994 + assert_eq!(session.max_response_payload_bytes, 8192);
1995 + assert_eq!(session.max_response_batch_items, 16);
1996 + assert_ne!(session.session_id, 0);
1997 +
1998 + session.close();
1999 + server.join().expect("server join");
2000 + }
2001 +
2002 + #[cfg(windows)]
2003 + #[test]
2004 + fn test_request_payload_over_cap() {
2005 + ensure_run_dir();
2006 + let svc = unique_service("rs_reqcap");
2007 +
2008 + let mut listener =
2009 + NpListener::bind(TEST_RUN_DIR, &svc, default_server_config()).expect("bind");
2010 +
2011 + let server = thread::spawn(move || {
2012 + let result = listener.accept();
2013 + assert!(matches!(result, Err(NpError::LimitExceeded)));
2014 + });
2015 +
2016 + let ccfg = ClientConfig {
2017 + max_request_payload_bytes: protocol::MAX_PAYLOAD_CAP + 1,
2018 + ..default_client_config()
2019 + };
2020 + let result = NpSession::connect(TEST_RUN_DIR, &svc, &ccfg);
2021 + assert!(matches!(result, Err(NpError::LimitExceeded)));
2022 +
2023 + server.join().expect("server join");
2024 + }
2025 +
2026 + #[cfg(windows)]
2027 + #[test]
2028 + fn test_disconnect_clears_all_inflight() {
2029 + ensure_run_dir();
2030 + let svc = unique_service("rs_disc");
2031 + let mut listener =
2032 + NpListener::bind(TEST_RUN_DIR, &svc, default_server_config()).expect("bind");
2033 +
2034 + let server = thread::spawn(move || {
2035 + let mut session = listener.accept().expect("accept");
2036 + let mut buf = [0u8; 256];
2037 + let _ = session.receive(&mut buf).expect("recv request");
2038 + session.close();
2039 + });
2040 +
2041 + let mut session =
2042 + NpSession::connect(TEST_RUN_DIR, &svc, &default_client_config()).expect("connect");
2043 +
2044 + let mut hdr = Header {
2045 + kind: KIND_REQUEST,
2046 + code: protocol::METHOD_INCREMENT,
2047 + item_count: 1,
2048 + message_id: 99,
2049 + ..Header::default()
2050 + };
2051 + session.send(&mut hdr, &[0xFF]).expect("send");
2052 + session.inflight_ids.insert(100);
2053 +
2054 + let mut rbuf = [0u8; 256];
2055 + assert!(matches!(
2056 + session.receive(&mut rbuf),
2057 + Err(NpError::Disconnected)
2058 + ));
2059 + assert!(
2060 + session.inflight_ids.is_empty(),
2061 + "disconnect must fail every in-flight request on the session"
2062 + );
2063 +
2064 + session.close();
2065 + server.join().expect("server join");
2066 + }
2067 +
2068 + #[cfg(windows)]
2069 + #[test]
2070 + fn test_unknown_response_msg_id() {
2071 + ensure_run_dir();
2072 + let svc = unique_service("rs_unkmsg");
2073 + let mut listener =
2074 + NpListener::bind(TEST_RUN_DIR, &svc, default_server_config()).expect("bind");
2075 +
2076 + let server = thread::spawn(move || {
2077 + let mut session = listener.accept().expect("accept");
2078 + let mut buf = [0u8; 256];
2079 + let (hdr, payload) = session.receive(&mut buf).expect("recv");
2080 + let payload = payload.to_vec();
2081 +
2082 + let mut resp = Header {
2083 + kind: KIND_RESPONSE,
2084 + code: hdr.code,
2085 + message_id: hdr.message_id + 999,
2086 + item_count: 1,
2087 + transport_status: protocol::STATUS_OK,
2088 + ..Header::default()
2089 + };
2090 + session.send(&mut resp, &payload).expect("send");
2091 + session.close();
2092 + });
2093 +
2094 + let mut session =
2095 + NpSession::connect(TEST_RUN_DIR, &svc, &default_client_config()).expect("connect");
2096 +
2097 + let mut hdr = Header {
2098 + kind: KIND_REQUEST,
2099 + code: protocol::METHOD_INCREMENT,
2100 + item_count: 1,
2101 + message_id: 50,
2102 + ..Header::default()
2103 + };
2104 + session.send(&mut hdr, &[1]).expect("send");
2105 +
2106 + let mut rbuf = [0u8; 256];
2107 + assert!(matches!(
2108 + session.receive(&mut rbuf),
2109 + Err(NpError::UnknownMsgId(_))
2110 + ));
2111 +
2112 + session.close();
2113 + server.join().expect("server join");
2114 + }
2115 +
2116 + #[cfg(windows)]
2117 + #[test]
2118 + fn test_preferred_profile_selected() {
2119 + ensure_run_dir();
2120 + let svc = unique_service("rs_pref");
2121 + let preferred = crate::protocol::PROFILE_SHM_HYBRID;
2122 +
2123 + let scfg = ServerConfig {
2124 + supported_profiles: PROFILE_BASELINE | preferred,
2125 + preferred_profiles: preferred,
2126 + ..default_server_config()
2127 + };
2128 + let mut listener = NpListener::bind(TEST_RUN_DIR, &svc, scfg).expect("bind");
2129 +
2130 + let server = thread::spawn(move || {
2131 + let mut session = listener.accept().expect("accept");
2132 + assert_eq!(session.selected_profile, preferred);
2133 + session.close();
2134 + });
2135 +
2136 + let ccfg = ClientConfig {
2137 + supported_profiles: PROFILE_BASELINE | preferred,
2138 + preferred_profiles: preferred,
2139 + ..default_client_config()
2140 + };
2141 + let mut session = NpSession::connect(TEST_RUN_DIR, &svc, &ccfg).expect("connect");
2142 + assert_eq!(session.selected_profile, preferred);
2143 + session.close();
2144 + server.join().expect("server join");
2145 + }
2146 +}
src/go/pkg/netipc/protocol/cgroups.go new
+460
@@ -0,0 +1,460 @@
1 +// Cgroups snapshot codec -- request, response view, builder, dispatch.
2 +
3 +package protocol
4 +
5 +import (
6 + "fmt"
7 +)
8 +
9 +// ---------------------------------------------------------------------------
10 +// Cgroups snapshot request (4 bytes)
11 +// ---------------------------------------------------------------------------
12 +
13 +// CgroupsRequest is the cgroups snapshot request payload (4 bytes).
14 +type CgroupsRequest struct {
15 + LayoutVersion uint16
16 + Flags uint16
17 +}
18 +
19 +// Encode writes the request into buf. Returns 4 on success, 0 if buf is
20 +// too small.
21 +func (r *CgroupsRequest) Encode(buf []byte) int {
22 + if len(buf) < cgroupsReqSize {
23 + return 0
24 + }
25 + ne.PutUint16(buf[0:2], r.LayoutVersion)
26 + ne.PutUint16(buf[2:4], r.Flags)
27 + return cgroupsReqSize
28 +}
29 +
30 +// DecodeCgroupsRequest decodes a cgroups request from buf. Validates
31 +// layout_version.
32 +func DecodeCgroupsRequest(buf []byte) (CgroupsRequest, error) {
33 + if len(buf) < cgroupsReqSize {
34 + return CgroupsRequest{}, ErrTruncated
35 + }
36 + r := CgroupsRequest{
37 + LayoutVersion: ne.Uint16(buf[0:2]),
38 + Flags: ne.Uint16(buf[2:4]),
39 + }
40 + if r.LayoutVersion != 1 {
41 + return CgroupsRequest{}, ErrBadLayout
42 + }
43 + // flags must be zero (reserved for future use)
44 + if r.Flags != 0 {
45 + return CgroupsRequest{}, ErrBadLayout
46 + }
47 + return r, nil
48 +}
49 +
50 +// ---------------------------------------------------------------------------
51 +// CStringView - borrowed string view into payload buffer
52 +// ---------------------------------------------------------------------------
53 +
54 +// CStringView is a borrowed, zero-copy string view into the payload buffer.
55 +// It wraps a byte slice that includes the NUL terminator. The view is
56 +// ephemeral and valid only while the underlying payload buffer lives.
57 +// Copy immediately via String() if the data is needed later.
58 +type CStringView struct {
59 + data []byte // includes trailing NUL
60 + len uint32 // length excluding NUL
61 +}
62 +
63 +// NewCStringView creates a CStringView from a slice that includes the NUL
64 +// terminator and the length excluding the NUL.
65 +func NewCStringView(data []byte, length uint32) CStringView {
66 + return CStringView{data: data, len: length}
67 +}
68 +
69 +// Bytes returns the string content as a byte slice (without the NUL).
70 +func (v CStringView) Bytes() []byte {
71 + return v.data[:v.len]
72 +}
73 +
74 +// Len returns the string length excluding the NUL terminator.
75 +func (v CStringView) Len() uint32 {
76 + return v.len
77 +}
78 +
79 +// String returns a copy of the string content. This allocates.
80 +func (v CStringView) String() string {
81 + return string(v.data[:v.len])
82 +}
83 +
84 +// GoString implements fmt.GoStringer for debug output.
85 +func (v CStringView) GoString() string {
86 + return fmt.Sprintf("CStringView(%q)", v.data[:v.len])
87 +}
88 +
89 +// ---------------------------------------------------------------------------
90 +// Cgroups snapshot response
91 +// ---------------------------------------------------------------------------
92 +
93 +// CgroupsItemView is a per-item view -- ephemeral, borrows the payload
94 +// buffer. Valid only while the payload buffer is alive.
95 +type CgroupsItemView struct {
96 + LayoutVersion uint16
97 + Flags uint16
98 + Hash uint32
99 + Options uint32
100 + Enabled uint32
101 + Name CStringView
102 + Path CStringView
103 +}
104 +
105 +// CgroupsResponseView is a full snapshot view -- ephemeral, borrows the
106 +// payload buffer. Valid only during the current library call or callback.
107 +// Copy immediately if the data is needed later.
108 +type CgroupsResponseView struct {
109 + LayoutVersion uint16
110 + Flags uint16
111 + ItemCount uint32
112 + SystemdEnabled uint32
113 + Generation uint64
114 + payload []byte // full payload for item access
115 +}
116 +
117 +// DecodeCgroupsResponse decodes the snapshot response header and validates
118 +// the item directory. On success, use Item() to access individual items.
119 +func DecodeCgroupsResponse(buf []byte) (CgroupsResponseView, error) {
120 + if len(buf) < cgroupsRespHdr {
121 + return CgroupsResponseView{}, ErrTruncated
122 + }
123 +
124 + layoutVersion := ne.Uint16(buf[0:2])
125 + flags := ne.Uint16(buf[2:4])
126 + itemCount := ne.Uint32(buf[4:8])
127 + systemdEnabled := ne.Uint32(buf[8:12])
128 + reserved := ne.Uint32(buf[12:16])
129 + generation := ne.Uint64(buf[16:24])
130 +
131 + if layoutVersion != 1 {
132 + return CgroupsResponseView{}, ErrBadLayout
133 + }
134 +
135 + // flags must be zero
136 + if flags != 0 {
137 + return CgroupsResponseView{}, ErrBadLayout
138 + }
139 +
140 + // reserved field must be zero
141 + if reserved != 0 {
142 + return CgroupsResponseView{}, ErrBadLayout
143 + }
144 +
145 + // Validate directory fits (use uint64 to prevent int overflow on 32-bit).
146 + dirSize64 := uint64(itemCount) * uint64(cgroupsDirEntry)
147 + dirEnd64 := uint64(cgroupsRespHdr) + dirSize64
148 + if dirEnd64 > uint64(len(buf)) {
149 + return CgroupsResponseView{}, ErrTruncated
150 + }
151 + dirEnd := int(dirEnd64)
152 +
153 + packedAreaLen := len(buf) - dirEnd
154 +
155 + // Validate each directory entry.
156 + dirSize := int(dirSize64)
157 + for i := 0; i < dirSize; i += cgroupsDirEntry {
158 + base := cgroupsRespHdr + i
159 + off := ne.Uint32(buf[base : base+4])
160 + length := ne.Uint32(buf[base+4 : base+8])
161 +
162 + if int(off)%Alignment != 0 {
163 + return CgroupsResponseView{}, ErrBadAlignment
164 + }
165 + if uint64(off)+uint64(length) > uint64(packedAreaLen) {
166 + return CgroupsResponseView{}, ErrOutOfBounds
167 + }
168 + if int(length) < cgroupsItemHdr {
169 + return CgroupsResponseView{}, ErrTruncated
170 + }
171 + }
172 +
173 + return CgroupsResponseView{
174 + LayoutVersion: layoutVersion,
175 + Flags: flags,
176 + ItemCount: itemCount,
177 + SystemdEnabled: systemdEnabled,
178 + Generation: generation,
179 + payload: buf,
180 + }, nil
181 +}
182 +
183 +// Item accesses the item at index from a decoded snapshot view. Returns an
184 +// ephemeral item view.
185 +func (v *CgroupsResponseView) Item(index uint32) (CgroupsItemView, error) {
186 + if index >= v.ItemCount {
187 + return CgroupsItemView{}, ErrOutOfBounds
188 + }
189 +
190 + dirStart := cgroupsRespHdr
191 + dirSize := int(uint64(v.ItemCount) * uint64(cgroupsDirEntry))
192 + packedAreaStart := dirStart + dirSize
193 +
194 + dirBase := dirStart + int(index)*cgroupsDirEntry
195 + itemOff := int(ne.Uint32(v.payload[dirBase : dirBase+4]))
196 + itemLen := int(ne.Uint32(v.payload[dirBase+4 : dirBase+8]))
197 +
198 + itemStart := packedAreaStart + itemOff
199 + item := v.payload[itemStart : itemStart+itemLen]
200 +
201 + layoutVersion := ne.Uint16(item[0:2])
202 + flags := ne.Uint16(item[2:4])
203 + hash := ne.Uint32(item[4:8])
204 + options := ne.Uint32(item[8:12])
205 + enabled := ne.Uint32(item[12:16])
206 +
207 + nameOff := int(ne.Uint32(item[16:20]))
208 + nameLen := ne.Uint32(item[20:24])
209 + pathOff := int(ne.Uint32(item[24:28]))
210 + pathLen := ne.Uint32(item[28:32])
211 +
212 + if layoutVersion != 1 {
213 + return CgroupsItemView{}, ErrBadLayout
214 + }
215 +
216 + // item flags must be zero
217 + if flags != 0 {
218 + return CgroupsItemView{}, ErrBadLayout
219 + }
220 +
221 + // Validate name string.
222 + if nameOff < cgroupsItemHdr {
223 + return CgroupsItemView{}, ErrOutOfBounds
224 + }
225 + if uint64(nameOff)+uint64(nameLen)+1 > uint64(itemLen) {
226 + return CgroupsItemView{}, ErrOutOfBounds
227 + }
228 + if item[nameOff+int(nameLen)] != 0 {
229 + return CgroupsItemView{}, ErrMissingNul
230 + }
231 +
232 + // Validate path string.
233 + if pathOff < cgroupsItemHdr {
234 + return CgroupsItemView{}, ErrOutOfBounds
235 + }
236 + if uint64(pathOff)+uint64(pathLen)+1 > uint64(itemLen) {
237 + return CgroupsItemView{}, ErrOutOfBounds
238 + }
239 + if item[pathOff+int(pathLen)] != 0 {
240 + return CgroupsItemView{}, ErrMissingNul
241 + }
242 +
243 + // Reject overlapping name and path regions (including NUL)
244 + {
245 + nameStart := uint64(nameOff)
246 + nameEnd := nameStart + uint64(nameLen) + 1
247 + pathStart := uint64(pathOff)
248 + pathEnd := pathStart + uint64(pathLen) + 1
249 + if nameStart < pathEnd && pathStart < nameEnd {
250 + return CgroupsItemView{}, ErrBadLayout
251 + }
252 + }
253 +
254 + name := NewCStringView(item[nameOff:nameOff+int(nameLen)+1], nameLen)
255 + path := NewCStringView(item[pathOff:pathOff+int(pathLen)+1], pathLen)
256 +
257 + return CgroupsItemView{
258 + LayoutVersion: layoutVersion,
259 + Flags: flags,
260 + Hash: hash,
261 + Options: options,
262 + Enabled: enabled,
263 + Name: name,
264 + Path: path,
265 + }, nil
266 +}
267 +
268 +// ---------------------------------------------------------------------------
269 +// Cgroups snapshot response builder
270 +// ---------------------------------------------------------------------------
271 +
272 +// CgroupsBuilder builds a cgroups snapshot response payload.
273 +//
274 +// Layout during building (maxItems directory slots reserved):
275 +//
276 +// [24-byte header space] [maxItems*8 directory] [packed items]
277 +//
278 +// Layout after Finish (compacted to actual itemCount):
279 +//
280 +// [24-byte header] [itemCount*8 directory] [packed items]
281 +type CgroupsBuilder struct {
282 + buf []byte
283 + systemdEnabled uint32
284 + generation uint64
285 + itemCount uint32
286 + maxItems uint32
287 + dataOffset int // current write position (absolute in buf)
288 +}
289 +
290 +// NewCgroupsBuilder initializes a cgroups response builder. buf must be
291 +// caller-owned and large enough for the expected snapshot.
292 +func NewCgroupsBuilder(buf []byte, maxItems uint32, systemdEnabled uint32, generation uint64) *CgroupsBuilder {
293 + minRequired, ok := CgroupsBuilderMinBytes(maxItems)
294 + if !ok || len(buf) < minRequired {
295 + panic(fmt.Sprintf("CgroupsBuilder buffer too small: need at least %d bytes, got %d",
296 + minRequired, len(buf)))
297 + }
298 + dataOffset := minRequired
299 + return &CgroupsBuilder{
300 + buf: buf,
301 + systemdEnabled: systemdEnabled,
302 + generation: generation,
303 + maxItems: maxItems,
304 + dataOffset: dataOffset,
305 + }
306 +}
307 +
308 +// CgroupsBuilderMinBytes returns the minimum response buffer required to
309 +// reserve directory slots for maxItems before packed item data is appended.
310 +func CgroupsBuilderMinBytes(maxItems uint32) (int, bool) {
311 + minRequired := uint64(cgroupsRespHdr) + uint64(maxItems)*uint64(cgroupsDirEntry)
312 + maxInt := uint64(int(^uint(0) >> 1))
313 + if minRequired > maxInt {
314 + return 0, false
315 + }
316 + return int(minRequired), true
317 +}
318 +
319 +// SetHeader updates the response header fields written by Finish().
320 +func (b *CgroupsBuilder) SetHeader(systemdEnabled uint32, generation uint64) {
321 + b.systemdEnabled = systemdEnabled
322 + b.generation = generation
323 +}
324 +
325 +// EstimateCgroupsMaxItems returns a safe upper bound for the number of
326 +// cgroup items that can fit in a response buffer of size bufSize.
327 +//
328 +// This is an upper bound for builder reservation, not a promise that all of
329 +// those items will fit with arbitrary string lengths.
330 +func EstimateCgroupsMaxItems(bufSize int) uint32 {
331 + if bufSize <= cgroupsRespHdr {
332 + return 0
333 + }
334 +
335 + minAlignedItem := Align8(cgroupsItemHdr + 2)
336 + return uint32((bufSize - cgroupsRespHdr) / (cgroupsDirEntry + minAlignedItem))
337 +}
338 +
339 +// Add adds one cgroup item. Handles offset bookkeeping, NUL termination,
340 +// and alignment.
341 +func (b *CgroupsBuilder) Add(hash, options, enabled uint32, name, path []byte) error {
342 + if b.itemCount >= b.maxItems {
343 + return ErrOverflow
344 + }
345 +
346 + // Align item start to 8 bytes.
347 + itemStart := Align8(b.dataOffset)
348 +
349 + // Item payload: 32-byte header + name + NUL + path + NUL.
350 + itemSize := cgroupsItemHdr + len(name) + 1 + len(path) + 1
351 +
352 + if itemStart+itemSize > len(b.buf) {
353 + return ErrOverflow
354 + }
355 +
356 + // Zero alignment padding.
357 + if itemStart > b.dataOffset {
358 + clear(b.buf[b.dataOffset:itemStart])
359 + }
360 +
361 + nameOffset := uint32(cgroupsItemHdr)
362 + pathOffset := uint32(cgroupsItemHdr) + uint32(len(name)) + 1
363 +
364 + // Write item header.
365 + p := itemStart
366 + ne.PutUint16(b.buf[p:p+2], 1) // layout_version
367 + ne.PutUint16(b.buf[p+2:p+4], 0) // flags
368 + ne.PutUint32(b.buf[p+4:p+8], hash)
369 + ne.PutUint32(b.buf[p+8:p+12], options)
370 + ne.PutUint32(b.buf[p+12:p+16], enabled)
371 + ne.PutUint32(b.buf[p+16:p+20], nameOffset)
372 + ne.PutUint32(b.buf[p+20:p+24], uint32(len(name)))
373 + ne.PutUint32(b.buf[p+24:p+28], pathOffset)
374 + ne.PutUint32(b.buf[p+28:p+32], uint32(len(path)))
375 +
376 + // Write strings with NUL terminators.
377 + ns := p + int(nameOffset)
378 + copy(b.buf[ns:], name)
379 + b.buf[ns+len(name)] = 0
380 +
381 + ps := p + int(pathOffset)
382 + copy(b.buf[ps:], path)
383 + b.buf[ps+len(path)] = 0
384 +
385 + // Write directory entry (absolute offset stored temporarily).
386 + dirEntry := cgroupsRespHdr + int(b.itemCount)*cgroupsDirEntry
387 + ne.PutUint32(b.buf[dirEntry:dirEntry+4], uint32(itemStart))
388 + ne.PutUint32(b.buf[dirEntry+4:dirEntry+8], uint32(itemSize))
389 +
390 + b.dataOffset = itemStart + itemSize
391 + b.itemCount++
392 + return nil
393 +}
394 +
395 +// Finish finalizes the builder. Returns the total payload size. The buffer
396 +// now contains a complete, decodable cgroups snapshot response payload.
397 +func (b *CgroupsBuilder) Finish() int {
398 + p := b.buf
399 +
400 + if b.itemCount == 0 {
401 + ne.PutUint16(p[0:2], 1)
402 + ne.PutUint16(p[2:4], 0)
403 + ne.PutUint32(p[4:8], 0)
404 + ne.PutUint32(p[8:12], b.systemdEnabled)
405 + ne.PutUint32(p[12:16], 0)
406 + ne.PutUint64(p[16:24], b.generation)
407 + return cgroupsRespHdr
408 + }
409 +
410 + // Where the decoder expects packed data to start.
411 + finalPackedStart := cgroupsRespHdr + int(b.itemCount)*cgroupsDirEntry
412 +
413 + // Read the first directory entry to find where packed data begins.
414 + firstItemAbs := int(ne.Uint32(p[cgroupsRespHdr : cgroupsRespHdr+4]))
415 +
416 + packedDataLen := b.dataOffset - firstItemAbs
417 +
418 + if finalPackedStart < firstItemAbs {
419 + // Shift packed data left.
420 + copy(p[finalPackedStart:], p[firstItemAbs:firstItemAbs+packedDataLen])
421 + }
422 +
423 + // Convert directory entries from absolute to relative offsets.
424 + dirBase := cgroupsRespHdr
425 + for i := 0; i < int(b.itemCount); i++ {
426 + entry := dirBase + i*cgroupsDirEntry
427 + absOff := ne.Uint32(p[entry : entry+4])
428 + relOff := absOff - uint32(firstItemAbs)
429 + ne.PutUint32(p[entry:entry+4], relOff)
430 + // length stays the same.
431 + }
432 +
433 + // Write snapshot header.
434 + ne.PutUint16(p[0:2], 1)
435 + ne.PutUint16(p[2:4], 0)
436 + ne.PutUint32(p[4:8], b.itemCount)
437 + ne.PutUint32(p[8:12], b.systemdEnabled)
438 + ne.PutUint32(p[12:16], 0)
439 + ne.PutUint64(p[16:24], b.generation)
440 +
441 + return finalPackedStart + packedDataLen
442 +}
443 +
444 +// DispatchCgroupsSnapshot decodes request, builds response via handler.
445 +func DispatchCgroupsSnapshot(req []byte, resp []byte, maxItems uint32,
446 + handler func(*CgroupsRequest, *CgroupsBuilder) bool) (int, bool) {
447 + request, err := DecodeCgroupsRequest(req)
448 + if err != nil {
449 + return 0, false
450 + }
451 + minRequired, ok := CgroupsBuilderMinBytes(maxItems)
452 + if !ok || len(resp) < minRequired {
453 + return 0, false
454 + }
455 + builder := NewCgroupsBuilder(resp, maxItems, 0, 0)
456 + if !handler(&request, builder) {
457 + return 0, false
458 + }
459 + return builder.Finish(), true
460 +}
src/go/pkg/netipc/protocol/codec_edge_test.go new
+920
@@ -0,0 +1,920 @@
1 +package protocol
2 +
3 +import (
4 + "testing"
5 +)
6 +
7 +// ---------------------------------------------------------------------------
8 +// DecodeChunkHeader error paths
9 +// ---------------------------------------------------------------------------
10 +
11 +func TestDecodeChunkHeaderBadFlags(t *testing.T) {
12 + // Valid chunk header but with non-zero flags (must be 0)
13 + c := ChunkHeader{
14 + Magic: MagicChunk,
15 + Version: Version,
16 + Flags: 1, // invalid
17 + MessageID: 1,
18 + TotalMessageLen: 256,
19 + ChunkIndex: 0,
20 + ChunkCount: 3,
21 + ChunkPayloadLen: 100,
22 + }
23 + var buf [HeaderSize]byte
24 + c.Encode(buf[:])
25 + _, err := DecodeChunkHeader(buf[:])
26 + if err != ErrBadLayout {
27 + t.Fatalf("expected ErrBadLayout for non-zero flags, got %v", err)
28 + }
29 +}
30 +
31 +func TestDecodeChunkHeaderZeroPayloadLen(t *testing.T) {
32 + // Valid chunk header but with zero chunk_payload_len
33 + c := ChunkHeader{
34 + Magic: MagicChunk,
35 + Version: Version,
36 + Flags: 0,
37 + MessageID: 1,
38 + TotalMessageLen: 256,
39 + ChunkIndex: 0,
40 + ChunkCount: 3,
41 + ChunkPayloadLen: 0, // invalid
42 + }
43 + var buf [HeaderSize]byte
44 + c.Encode(buf[:])
45 + _, err := DecodeChunkHeader(buf[:])
46 + if err != ErrBadLayout {
47 + t.Fatalf("expected ErrBadLayout for zero chunk_payload_len, got %v", err)
48 + }
49 +}
50 +
51 +func TestDecodeChunkHeaderBadVersion(t *testing.T) {
52 + c := ChunkHeader{
53 + Magic: MagicChunk,
54 + Version: 99,
55 + Flags: 0,
56 + MessageID: 1,
57 + TotalMessageLen: 256,
58 + ChunkIndex: 0,
59 + ChunkCount: 3,
60 + ChunkPayloadLen: 100,
61 + }
62 + var buf [HeaderSize]byte
63 + c.Encode(buf[:])
64 + _, err := DecodeChunkHeader(buf[:])
65 + if err != ErrBadVersion {
66 + t.Fatalf("expected ErrBadVersion, got %v", err)
67 + }
68 +}
69 +
70 +// ---------------------------------------------------------------------------
71 +// DecodeHelloAck error paths
72 +// ---------------------------------------------------------------------------
73 +
74 +func TestDecodeHelloAckBadLayout(t *testing.T) {
75 + h := HelloAck{
76 + LayoutVersion: 99, // invalid
77 + }
78 + var buf [64]byte
79 + h.Encode(buf[:])
80 + _, err := DecodeHelloAck(buf[:])
81 + if err != ErrBadLayout {
82 + t.Fatalf("expected ErrBadLayout for bad layout_version, got %v", err)
83 + }
84 +}
85 +
86 +func TestDecodeHelloAckBadFlags(t *testing.T) {
87 + h := HelloAck{
88 + LayoutVersion: 1,
89 + Flags: 0x0001, // non-zero flags
90 + }
91 + var buf [64]byte
92 + h.Encode(buf[:])
93 + _, err := DecodeHelloAck(buf[:])
94 + if err != ErrBadLayout {
95 + t.Fatalf("expected ErrBadLayout for non-zero flags, got %v", err)
96 + }
97 +}
98 +
99 +// ---------------------------------------------------------------------------
100 +// BatchDirValidate
101 +// ---------------------------------------------------------------------------
102 +
103 +func TestBatchDirValidateSuccess(t *testing.T) {
104 + // Build a valid directory with 2 entries
105 + var buf [16]byte
106 + ne.PutUint32(buf[0:4], 0) // offset=0, aligned
107 + ne.PutUint32(buf[4:8], 10) // length=10
108 + ne.PutUint32(buf[8:12], 16) // offset=16, aligned
109 + ne.PutUint32(buf[12:16], 5) // length=5
110 +
111 + err := BatchDirValidate(buf[:], 2, 100)
112 + if err != nil {
113 + t.Fatalf("expected success, got %v", err)
114 + }
115 +}
116 +
117 +func TestBatchDirValidateTruncated(t *testing.T) {
118 + err := BatchDirValidate(make([]byte, 4), 2, 100)
119 + if err != ErrTruncated {
120 + t.Fatalf("expected ErrTruncated, got %v", err)
121 + }
122 +}
123 +
124 +func TestBatchDirValidateBadAlignment(t *testing.T) {
125 + var buf [8]byte
126 + ne.PutUint32(buf[0:4], 3) // offset=3, not 8-byte aligned
127 + ne.PutUint32(buf[4:8], 5)
128 +
129 + err := BatchDirValidate(buf[:], 1, 100)
130 + if err != ErrBadAlignment {
131 + t.Fatalf("expected ErrBadAlignment, got %v", err)
132 + }
133 +}
134 +
135 +func TestBatchDirValidateOutOfBounds(t *testing.T) {
136 + var buf [8]byte
137 + ne.PutUint32(buf[0:4], 0)
138 + ne.PutUint32(buf[4:8], 200) // length=200 exceeds packedAreaLen=100
139 +
140 + err := BatchDirValidate(buf[:], 1, 100)
141 + if err != ErrOutOfBounds {
142 + t.Fatalf("expected ErrOutOfBounds, got %v", err)
143 + }
144 +}
145 +
146 +func TestBatchDirValidateZeroItems(t *testing.T) {
147 + err := BatchDirValidate(nil, 0, 0)
148 + if err != nil {
149 + t.Fatalf("expected success for zero items, got %v", err)
150 + }
151 +}
152 +
153 +// ---------------------------------------------------------------------------
154 +// BatchItemGet error paths
155 +// ---------------------------------------------------------------------------
156 +
157 +func TestBatchItemGetBadAlignment(t *testing.T) {
158 + // Build a batch payload with 1 item, but set misaligned offset
159 + var buf [64]byte
160 + ne.PutUint32(buf[0:4], 3) // offset=3, not aligned
161 + ne.PutUint32(buf[4:8], 5) // length=5
162 +
163 + _, err := BatchItemGet(buf[:], 1, 0)
164 + if err != ErrBadAlignment {
165 + t.Fatalf("expected ErrBadAlignment, got %v", err)
166 + }
167 +}
168 +
169 +func TestBatchItemGetDirTruncated(t *testing.T) {
170 + // 2 items = 16 bytes dir, aligned to 16. But give only 8 bytes total.
171 + _, err := BatchItemGet(make([]byte, 8), 2, 0)
172 + if err != ErrTruncated {
173 + t.Fatalf("expected ErrTruncated, got %v", err)
174 + }
175 +}
176 +
177 +// ---------------------------------------------------------------------------
178 +// DecodeCgroupsRequest error paths
179 +// ---------------------------------------------------------------------------
180 +
181 +func TestDecodeCgroupsRequestBadFlags(t *testing.T) {
182 + // Valid layout_version but non-zero flags
183 + var buf [4]byte
184 + ne.PutUint16(buf[0:2], 1) // layout_version = 1
185 + ne.PutUint16(buf[2:4], 0x01) // flags = 1 (invalid)
186 +
187 + _, err := DecodeCgroupsRequest(buf[:])
188 + if err != ErrBadLayout {
189 + t.Fatalf("expected ErrBadLayout for non-zero flags, got %v", err)
190 + }
191 +}
192 +
193 +// ---------------------------------------------------------------------------
194 +// DecodeCgroupsResponse error paths
195 +// ---------------------------------------------------------------------------
196 +
197 +func TestDecodeCgroupsResponseBadFlags(t *testing.T) {
198 + var buf [24]byte
199 + ne.PutUint16(buf[0:2], 1) // layout_version = 1
200 + ne.PutUint16(buf[2:4], 0x01) // flags = 1 (invalid)
201 + ne.PutUint32(buf[4:8], 0) // item_count = 0
202 + ne.PutUint32(buf[8:12], 0) // systemd_enabled
203 + ne.PutUint32(buf[12:16], 0) // reserved
204 + ne.PutUint64(buf[16:24], 0) // generation
205 +
206 + _, err := DecodeCgroupsResponse(buf[:])
207 + if err != ErrBadLayout {
208 + t.Fatalf("expected ErrBadLayout for non-zero flags, got %v", err)
209 + }
210 +}
211 +
212 +func TestDecodeCgroupsResponseBadReserved(t *testing.T) {
213 + var buf [24]byte
214 + ne.PutUint16(buf[0:2], 1) // layout_version = 1
215 + ne.PutUint16(buf[2:4], 0) // flags = 0
216 + ne.PutUint32(buf[4:8], 0) // item_count = 0
217 + ne.PutUint32(buf[8:12], 0) // systemd_enabled
218 + ne.PutUint32(buf[12:16], 99) // reserved non-zero (invalid)
219 + ne.PutUint64(buf[16:24], 0) // generation
220 +
221 + _, err := DecodeCgroupsResponse(buf[:])
222 + if err != ErrBadLayout {
223 + t.Fatalf("expected ErrBadLayout for non-zero reserved field, got %v", err)
224 + }
225 +}
226 +
227 +func TestDecodeCgroupsResponseDirTruncated(t *testing.T) {
228 + // Declare 1 item but no space for the directory entry
229 + var buf [24]byte
230 + ne.PutUint16(buf[0:2], 1) // layout_version = 1
231 + ne.PutUint16(buf[2:4], 0) // flags = 0
232 + ne.PutUint32(buf[4:8], 1) // item_count = 1
233 + ne.PutUint32(buf[8:12], 0) // systemd_enabled
234 + ne.PutUint32(buf[12:16], 0)
235 + ne.PutUint64(buf[16:24], 0)
236 +
237 + _, err := DecodeCgroupsResponse(buf[:])
238 + if err != ErrTruncated {
239 + t.Fatalf("expected ErrTruncated, got %v", err)
240 + }
241 +}
242 +
243 +func TestDecodeCgroupsResponseDirBadAlignment(t *testing.T) {
244 + // 1 item: dir at offset 24, packed area starts at 32
245 + // Set item offset to 3 (misaligned)
246 + var buf [128]byte
247 + ne.PutUint16(buf[0:2], 1) // layout_version
248 + ne.PutUint16(buf[2:4], 0) // flags
249 + ne.PutUint32(buf[4:8], 1) // item_count = 1
250 + ne.PutUint32(buf[8:12], 0) // systemd_enabled
251 + ne.PutUint32(buf[12:16], 0)
252 + ne.PutUint64(buf[16:24], 0)
253 + // Directory entry at offset 24
254 + ne.PutUint32(buf[24:28], 3) // offset = 3 (misaligned)
255 + ne.PutUint32(buf[28:32], 40) // length = 40
256 +
257 + _, err := DecodeCgroupsResponse(buf[:])
258 + if err != ErrBadAlignment {
259 + t.Fatalf("expected ErrBadAlignment, got %v", err)
260 + }
261 +}
262 +
263 +func TestDecodeCgroupsResponseDirItemTooSmall(t *testing.T) {
264 + // 1 item with length < cgroupsItemHdr (32)
265 + var buf [128]byte
266 + ne.PutUint16(buf[0:2], 1) // layout_version
267 + ne.PutUint16(buf[2:4], 0) // flags
268 + ne.PutUint32(buf[4:8], 1) // item_count = 1
269 + ne.PutUint32(buf[8:12], 0) // systemd_enabled
270 + ne.PutUint32(buf[12:16], 0)
271 + ne.PutUint64(buf[16:24], 0)
272 + // Directory entry at offset 24
273 + ne.PutUint32(buf[24:28], 0) // offset = 0
274 + ne.PutUint32(buf[28:32], 16) // length = 16 (< cgroupsItemHdr=32)
275 +
276 + _, err := DecodeCgroupsResponse(buf[:])
277 + if err != ErrTruncated {
278 + t.Fatalf("expected ErrTruncated for item too small, got %v", err)
279 + }
280 +}
281 +
282 +// ---------------------------------------------------------------------------
283 +// CgroupsResponseView.Item error paths
284 +// ---------------------------------------------------------------------------
285 +
286 +func TestCgroupsItemBadLayoutVersion(t *testing.T) {
287 + // Build a valid snapshot, then corrupt the item's layout_version
288 + var buf [4096]byte
289 + b := NewCgroupsBuilder(buf[:], 1, 0, 1)
290 + if err := b.Add(42, 0, 1, []byte("test"), []byte("/path")); err != nil {
291 + t.Fatalf("Add: %v", err)
292 + }
293 + total := b.Finish()
294 + payload := buf[:total]
295 +
296 + // Decode successfully first
297 + view, err := DecodeCgroupsResponse(payload)
298 + if err != nil {
299 + t.Fatalf("DecodeCgroupsResponse: %v", err)
300 + }
301 +
302 + // Find the item in the buffer and corrupt layout_version
303 + dirBase := cgroupsRespHdr
304 + itemOff := int(ne.Uint32(payload[dirBase : dirBase+4]))
305 + packedStart := cgroupsRespHdr + int(view.ItemCount)*cgroupsDirEntry
306 + itemAbsStart := packedStart + itemOff
307 + ne.PutUint16(payload[itemAbsStart:itemAbsStart+2], 99) // bad layout_version
308 +
309 + _, err = view.Item(0)
310 + if err != ErrBadLayout {
311 + t.Fatalf("expected ErrBadLayout for bad item layout_version, got %v", err)
312 + }
313 +}
314 +
315 +func TestCgroupsItemBadFlags(t *testing.T) {
316 + var buf [4096]byte
317 + b := NewCgroupsBuilder(buf[:], 1, 0, 1)
318 + if err := b.Add(42, 0, 1, []byte("test"), []byte("/path")); err != nil {
319 + t.Fatalf("Add: %v", err)
320 + }
321 + total := b.Finish()
322 + payload := buf[:total]
323 +
324 + view, err := DecodeCgroupsResponse(payload)
325 + if err != nil {
326 + t.Fatalf("DecodeCgroupsResponse: %v", err)
327 + }
328 +
329 + // Corrupt item flags
330 + dirBase := cgroupsRespHdr
331 + itemOff := int(ne.Uint32(payload[dirBase : dirBase+4]))
332 + packedStart := cgroupsRespHdr + int(view.ItemCount)*cgroupsDirEntry
333 + itemAbsStart := packedStart + itemOff
334 + ne.PutUint16(payload[itemAbsStart+2:itemAbsStart+4], 0x01) // bad flags
335 +
336 + _, err = view.Item(0)
337 + if err != ErrBadLayout {
338 + t.Fatalf("expected ErrBadLayout for bad item flags, got %v", err)
339 + }
340 +}
341 +
342 +// ---------------------------------------------------------------------------
343 +// CStringView.GoString
344 +// ---------------------------------------------------------------------------
345 +
346 +func TestCStringViewGoString(t *testing.T) {
347 + data := []byte("hello\x00")
348 + v := NewCStringView(data, 5)
349 + gs := v.GoString()
350 + if gs != `CStringView("hello")` {
351 + t.Fatalf("GoString = %q, want CStringView(\"hello\")", gs)
352 + }
353 +}
354 +
355 +// ---------------------------------------------------------------------------
356 +// DispatchCgroupsSnapshot
357 +// ---------------------------------------------------------------------------
358 +
359 +func TestDispatchCgroupsSnapshotSuccess(t *testing.T) {
360 + // Valid request
361 + var reqBuf [4]byte
362 + req := CgroupsRequest{LayoutVersion: 1, Flags: 0}
363 + req.Encode(reqBuf[:])
364 +
365 + resp := make([]byte, 4096)
366 + n, ok := DispatchCgroupsSnapshot(reqBuf[:], resp, 1,
367 + func(r *CgroupsRequest, b *CgroupsBuilder) bool {
368 + if err := b.Add(1, 0, 1, []byte("cg"), []byte("/path")); err != nil {
369 + return false
370 + }
371 + return true
372 + })
373 + if !ok {
374 + t.Fatal("expected success")
375 + }
376 + if n == 0 {
377 + t.Fatal("expected non-zero payload size")
378 + }
379 +
380 + // Verify the result decodes
381 + view, err := DecodeCgroupsResponse(resp[:n])
382 + if err != nil {
383 + t.Fatalf("DecodeCgroupsResponse: %v", err)
384 + }
385 + if view.ItemCount != 1 {
386 + t.Fatalf("expected 1 item, got %d", view.ItemCount)
387 + }
388 +}
389 +
390 +func TestDispatchCgroupsSnapshotBadRequest(t *testing.T) {
391 + // Truncated request
392 + _, ok := DispatchCgroupsSnapshot([]byte{0}, make([]byte, 4096), 1,
393 + func(r *CgroupsRequest, b *CgroupsBuilder) bool {
394 + return true
395 + })
396 + if ok {
397 + t.Fatal("expected failure for bad request")
398 + }
399 +}
400 +
401 +func TestDispatchCgroupsSnapshotHandlerFails(t *testing.T) {
402 + var reqBuf [4]byte
403 + req := CgroupsRequest{LayoutVersion: 1, Flags: 0}
404 + req.Encode(reqBuf[:])
405 +
406 + _, ok := DispatchCgroupsSnapshot(reqBuf[:], make([]byte, 4096), 1,
407 + func(r *CgroupsRequest, b *CgroupsBuilder) bool {
408 + return false
409 + })
410 + if ok {
411 + t.Fatal("expected failure when handler returns false")
412 + }
413 +}
414 +
415 +func TestDispatchCgroupsSnapshotEmptyResult(t *testing.T) {
416 + var reqBuf [4]byte
417 + req := CgroupsRequest{LayoutVersion: 1, Flags: 0}
418 + req.Encode(reqBuf[:])
419 +
420 + // Handler succeeds but adds no items - Finish returns cgroupsRespHdr (24) which is > 0
421 + n, ok := DispatchCgroupsSnapshot(reqBuf[:], make([]byte, 4096), 0,
422 + func(r *CgroupsRequest, b *CgroupsBuilder) bool {
423 + return true
424 + })
425 + if !ok {
426 + t.Fatal("expected success for empty snapshot")
427 + }
428 + if n != cgroupsRespHdr {
429 + t.Fatalf("expected %d bytes, got %d", cgroupsRespHdr, n)
430 + }
431 +}
432 +
433 +// ---------------------------------------------------------------------------
434 +// IncrementEncode / IncrementDecode
435 +// ---------------------------------------------------------------------------
436 +
437 +func TestIncrementEncodeSuccess(t *testing.T) {
438 + var buf [8]byte
439 + n := IncrementEncode(0xDEADBEEFCAFEBABE, buf[:])
440 + if n != 8 {
441 + t.Fatalf("expected 8, got %d", n)
442 + }
443 + val, err := IncrementDecode(buf[:])
444 + if err != nil {
445 + t.Fatalf("decode: %v", err)
446 + }
447 + if val != 0xDEADBEEFCAFEBABE {
448 + t.Fatalf("expected 0xDEADBEEFCAFEBABE, got 0x%x", val)
449 + }
450 +}
451 +
452 +func TestIncrementEncodeTooSmall(t *testing.T) {
453 + var buf [4]byte
454 + n := IncrementEncode(42, buf[:])
455 + if n != 0 {
456 + t.Fatalf("expected 0 for too-small buffer, got %d", n)
457 + }
458 +}
459 +
460 +func TestIncrementDecodeTruncated(t *testing.T) {
461 + _, err := IncrementDecode(make([]byte, 3))
462 + if err != ErrTruncated {
463 + t.Fatalf("expected ErrTruncated, got %v", err)
464 + }
465 +}
466 +
467 +// ---------------------------------------------------------------------------
468 +// DispatchIncrement
469 +// ---------------------------------------------------------------------------
470 +
471 +func TestDispatchIncrementSuccess(t *testing.T) {
472 + var reqBuf [8]byte
473 + IncrementEncode(100, reqBuf[:])
474 +
475 + var respBuf [8]byte
476 + n, ok := DispatchIncrement(reqBuf[:], respBuf[:], func(v uint64) (uint64, bool) {
477 + return v + 1, true
478 + })
479 + if !ok {
480 + t.Fatal("expected success")
481 + }
482 + if n != 8 {
483 + t.Fatalf("expected 8, got %d", n)
484 + }
485 + val, err := IncrementDecode(respBuf[:])
486 + if err != nil {
487 + t.Fatalf("decode: %v", err)
488 + }
489 + if val != 101 {
490 + t.Fatalf("expected 101, got %d", val)
491 + }
492 +}
493 +
494 +func TestDispatchIncrementBadRequest(t *testing.T) {
495 + _, ok := DispatchIncrement([]byte{0, 1}, make([]byte, 8), func(v uint64) (uint64, bool) {
496 + return v, true
497 + })
498 + if ok {
499 + t.Fatal("expected failure for truncated request")
500 + }
501 +}
502 +
503 +func TestDispatchIncrementHandlerFails(t *testing.T) {
504 + var reqBuf [8]byte
505 + IncrementEncode(42, reqBuf[:])
506 +
507 + _, ok := DispatchIncrement(reqBuf[:], make([]byte, 8), func(v uint64) (uint64, bool) {
508 + return 0, false
509 + })
510 + if ok {
511 + t.Fatal("expected failure when handler returns false")
512 + }
513 +}
514 +
515 +func TestDispatchIncrementRespTooSmall(t *testing.T) {
516 + var reqBuf [8]byte
517 + IncrementEncode(42, reqBuf[:])
518 +
519 + n, ok := DispatchIncrement(reqBuf[:], make([]byte, 2), func(v uint64) (uint64, bool) {
520 + return 42, true
521 + })
522 + // IncrementEncode returns 0 for too-small buffer, so n=0, ok = (0>0) = false
523 + if ok {
524 + t.Fatal("expected failure for too-small response buffer")
525 + }
526 + if n != 0 {
527 + t.Fatalf("expected n=0, got %d", n)
528 + }
529 +}
530 +
531 +// ---------------------------------------------------------------------------
532 +// StringReverseEncode / StringReverseDecode
533 +// ---------------------------------------------------------------------------
534 +
535 +func TestStringReverseRoundtrip(t *testing.T) {
536 + s := "hello world"
537 + buf := make([]byte, StringReverseHdrSize+len(s)+1)
538 + n := StringReverseEncode(s, buf)
539 + if n != len(buf) {
540 + t.Fatalf("expected %d, got %d", len(buf), n)
541 + }
542 +
543 + view, err := StringReverseDecode(buf)
544 + if err != nil {
545 + t.Fatalf("decode: %v", err)
546 + }
547 + if view.Str != s {
548 + t.Fatalf("expected %q, got %q", s, view.Str)
549 + }
550 + if view.StrLen != uint32(len(s)) {
551 + t.Fatalf("expected len=%d, got %d", len(s), view.StrLen)
552 + }
553 +}
554 +
555 +func TestStringReverseEncodeEmpty(t *testing.T) {
556 + buf := make([]byte, StringReverseHdrSize+1) // 8 + 0 + 1 NUL
557 + n := StringReverseEncode("", buf)
558 + if n != StringReverseHdrSize+1 {
559 + t.Fatalf("expected %d, got %d", StringReverseHdrSize+1, n)
560 + }
561 +
562 + view, err := StringReverseDecode(buf[:n])
563 + if err != nil {
564 + t.Fatalf("decode: %v", err)
565 + }
566 + if view.Str != "" {
567 + t.Fatalf("expected empty, got %q", view.Str)
568 + }
569 +}
570 +
571 +func TestStringReverseEncodeTooSmall(t *testing.T) {
572 + n := StringReverseEncode("hello", make([]byte, 5))
573 + if n != 0 {
574 + t.Fatalf("expected 0 for too-small buffer, got %d", n)
575 + }
576 +}
577 +
578 +func TestStringReverseDecodeTruncated(t *testing.T) {
579 + _, err := StringReverseDecode(make([]byte, 4))
580 + if err != ErrTruncated {
581 + t.Fatalf("expected ErrTruncated, got %v", err)
582 + }
583 +}
584 +
585 +func TestStringReverseDecodeOutOfBounds(t *testing.T) {
586 + var buf [8]byte
587 + ne.PutUint32(buf[0:4], 8) // str_offset
588 + ne.PutUint32(buf[4:8], 99) // str_length (overflows)
589 +
590 + _, err := StringReverseDecode(buf[:])
591 + if err != ErrOutOfBounds {
592 + t.Fatalf("expected ErrOutOfBounds, got %v", err)
593 + }
594 +}
595 +
596 +func TestStringReverseDecodeMissingNul(t *testing.T) {
597 + // Build payload where NUL terminator is wrong
598 + buf := make([]byte, StringReverseHdrSize+6)
599 + ne.PutUint32(buf[0:4], uint32(StringReverseHdrSize)) // str_offset
600 + ne.PutUint32(buf[4:8], 5) // str_length
601 + copy(buf[8:13], "hello")
602 + buf[13] = 'X' // should be 0
603 +
604 + _, err := StringReverseDecode(buf)
605 + if err != ErrMissingNul {
606 + t.Fatalf("expected ErrMissingNul, got %v", err)
607 + }
608 +}
609 +
610 +// ---------------------------------------------------------------------------
611 +// DispatchStringReverse
612 +// ---------------------------------------------------------------------------
613 +
614 +func TestDispatchStringReverseSuccess(t *testing.T) {
615 + input := "hello"
616 + reqBuf := make([]byte, StringReverseHdrSize+len(input)+1)
617 + StringReverseEncode(input, reqBuf)
618 +
619 + respBuf := make([]byte, 128)
620 + n, ok := DispatchStringReverse(reqBuf, respBuf, func(s string) (string, bool) {
621 + // Reverse the string
622 + b := []byte(s)
623 + for i, j := 0, len(b)-1; i < j; i, j = i+1, j-1 {
624 + b[i], b[j] = b[j], b[i]
625 + }
626 + return string(b), true
627 + })
628 + if !ok {
629 + t.Fatal("expected success")
630 + }
631 + if n == 0 {
632 + t.Fatal("expected non-zero response size")
633 + }
634 +
635 + view, err := StringReverseDecode(respBuf[:n])
636 + if err != nil {
637 + t.Fatalf("decode: %v", err)
638 + }
639 + if view.Str != "olleh" {
640 + t.Fatalf("expected 'olleh', got %q", view.Str)
641 + }
642 +}
643 +
644 +func TestDispatchStringReverseBadRequest(t *testing.T) {
645 + _, ok := DispatchStringReverse([]byte{0}, make([]byte, 128), func(s string) (string, bool) {
646 + return s, true
647 + })
648 + if ok {
649 + t.Fatal("expected failure for bad request")
650 + }
651 +}
652 +
653 +func TestDispatchStringReverseHandlerFails(t *testing.T) {
654 + input := "test"
655 + reqBuf := make([]byte, StringReverseHdrSize+len(input)+1)
656 + StringReverseEncode(input, reqBuf)
657 +
658 + _, ok := DispatchStringReverse(reqBuf, make([]byte, 128), func(s string) (string, bool) {
659 + return "", false
660 + })
661 + if ok {
662 + t.Fatal("expected failure when handler returns false")
663 + }
664 +}
665 +
666 +func TestDispatchStringReverseRespTooSmall(t *testing.T) {
667 + input := "hello"
668 + reqBuf := make([]byte, StringReverseHdrSize+len(input)+1)
669 + StringReverseEncode(input, reqBuf)
670 +
671 + // Response buffer too small for result
672 + n, ok := DispatchStringReverse(reqBuf, make([]byte, 2), func(s string) (string, bool) {
673 + return "very long response string", true
674 + })
675 + if ok {
676 + t.Fatal("expected failure for too-small response buffer")
677 + }
678 + if n != 0 {
679 + t.Fatalf("expected n=0, got %d", n)
680 + }
681 +}
682 +
683 +// ---------------------------------------------------------------------------
684 +// CgroupsBuilder edge cases
685 +// ---------------------------------------------------------------------------
686 +
687 +func TestCgroupsBuilderOverflow(t *testing.T) {
688 + // Buffer too small for the item
689 + var buf [64]byte
690 + b := NewCgroupsBuilder(buf[:], 1, 0, 1)
691 + // name + path would overflow the small buffer
692 + err := b.Add(1, 0, 1,
693 + []byte("this-name-is-quite-long-to-overflow"),
694 + []byte("/this/path/is/also/long/enough/to/overflow"))
695 + if err != ErrOverflow {
696 + t.Fatalf("expected ErrOverflow, got %v", err)
697 + }
698 +}
699 +
700 +func TestCgroupsBuilderMaxItemsExceeded(t *testing.T) {
701 + var buf [4096]byte
702 + b := NewCgroupsBuilder(buf[:], 1, 0, 1)
703 + // First item succeeds
704 + if err := b.Add(1, 0, 1, []byte("a"), []byte("/b")); err != nil {
705 + t.Fatalf("first Add: %v", err)
706 + }
707 + // Second item exceeds maxItems
708 + err := b.Add(2, 0, 1, []byte("c"), []byte("/d"))
709 + if err != ErrOverflow {
710 + t.Fatalf("expected ErrOverflow, got %v", err)
711 + }
712 +}
713 +
714 +func TestCgroupsBuilderFinishCompaction(t *testing.T) {
715 + // Reserve space for 4 items but only add 2 — tests compaction path
716 + var buf [4096]byte
717 + b := NewCgroupsBuilder(buf[:], 4, 1, 42)
718 + if err := b.Add(1, 0, 1, []byte("name1"), []byte("/path1")); err != nil {
719 + t.Fatalf("Add 1: %v", err)
720 + }
721 + if err := b.Add(2, 0, 1, []byte("name2"), []byte("/path2")); err != nil {
722 + t.Fatalf("Add 2: %v", err)
723 + }
724 + total := b.Finish()
725 +
726 + view, err := DecodeCgroupsResponse(buf[:total])
727 + if err != nil {
728 + t.Fatalf("decode: %v", err)
729 + }
730 + if view.ItemCount != 2 {
731 + t.Fatalf("expected 2 items, got %d", view.ItemCount)
732 + }
733 + if view.SystemdEnabled != 1 {
734 + t.Fatalf("expected systemd_enabled=1, got %d", view.SystemdEnabled)
735 + }
736 + if view.Generation != 42 {
737 + t.Fatalf("expected generation=42, got %d", view.Generation)
738 + }
739 +
740 + // Verify both items decode correctly
741 + item0, err := view.Item(0)
742 + if err != nil {
743 + t.Fatalf("item 0: %v", err)
744 + }
745 + if item0.Name.String() != "name1" || item0.Path.String() != "/path1" {
746 + t.Fatalf("item 0 mismatch: name=%q path=%q", item0.Name.String(), item0.Path.String())
747 + }
748 +
749 + item1, err := view.Item(1)
750 + if err != nil {
751 + t.Fatalf("item 1: %v", err)
752 + }
753 + if item1.Name.String() != "name2" || item1.Path.String() != "/path2" {
754 + t.Fatalf("item 1 mismatch: name=%q path=%q", item1.Name.String(), item1.Path.String())
755 + }
756 +}
757 +
758 +// ---------------------------------------------------------------------------
759 +// BatchBuilder edge cases
760 +// ---------------------------------------------------------------------------
761 +
762 +func TestBatchBuilderCompactionWide(t *testing.T) {
763 + // Reserve 8 slots, use only 2 (wider gap than existing test)
764 + var buf [4096]byte
765 + bb := NewBatchBuilder(buf[:], 8)
766 + if err := bb.Add([]byte{1, 2, 3}); err != nil {
767 + t.Fatalf("Add 1: %v", err)
768 + }
769 + if err := bb.Add([]byte{4, 5}); err != nil {
770 + t.Fatalf("Add 2: %v", err)
771 + }
772 + total, count := bb.Finish()
773 + if count != 2 {
774 + t.Fatalf("expected count=2, got %d", count)
775 + }
776 +
777 + // Verify items are accessible
778 + item0, err := BatchItemGet(buf[:total], count, 0)
779 + if err != nil {
780 + t.Fatalf("get item 0: %v", err)
781 + }
782 + if len(item0) != 3 || item0[0] != 1 {
783 + t.Fatalf("item 0 mismatch: %v", item0)
784 + }
785 +
786 + item1, err := BatchItemGet(buf[:total], count, 1)
787 + if err != nil {
788 + t.Fatalf("get item 1: %v", err)
789 + }
790 + if len(item1) != 2 || item1[0] != 4 {
791 + t.Fatalf("item 1 mismatch: %v", item1)
792 + }
793 +}
794 +
795 +func TestBatchBuilderOverflowMaxItemsSingle(t *testing.T) {
796 + var buf [4096]byte
797 + bb := NewBatchBuilder(buf[:], 1)
798 + if err := bb.Add([]byte{1}); err != nil {
799 + t.Fatalf("first Add: %v", err)
800 + }
801 + err := bb.Add([]byte{2})
802 + if err != ErrOverflow {
803 + t.Fatalf("expected ErrOverflow, got %v", err)
804 + }
805 +}
806 +
807 +func TestBatchBuilderOverflowBufferFull(t *testing.T) {
808 + // Tiny buffer: dir for 1 item = 8 bytes aligned, so 8 bytes.
809 + // Total buf = 16 bytes, leaving 8 for data.
810 + var buf [16]byte
811 + bb := NewBatchBuilder(buf[:], 1)
812 + err := bb.Add(make([]byte, 100)) // too large
813 + if err != ErrOverflow {
814 + t.Fatalf("expected ErrOverflow, got %v", err)
815 + }
816 +}
817 +
818 +func TestBatchBuilderFinishNoCompaction(t *testing.T) {
819 + // Use exactly the reserved number of slots - no compaction needed
820 + var buf [4096]byte
821 + bb := NewBatchBuilder(buf[:], 2)
822 + if err := bb.Add([]byte{1, 2, 3, 4, 5, 6, 7, 8}); err != nil {
823 + t.Fatalf("Add 1: %v", err)
824 + }
825 + if err := bb.Add([]byte{10, 20}); err != nil {
826 + t.Fatalf("Add 2: %v", err)
827 + }
828 + total, count := bb.Finish()
829 + if count != 2 {
830 + t.Fatalf("expected count=2, got %d", count)
831 + }
832 +
833 + // Verify access
834 + for i := uint32(0); i < count; i++ {
835 + _, err := BatchItemGet(buf[:total], count, i)
836 + if err != nil {
837 + t.Fatalf("item %d: %v", i, err)
838 + }
839 + }
840 +}
841 +
842 +// ---------------------------------------------------------------------------
843 +// DecodeHello padding validation
844 +// ---------------------------------------------------------------------------
845 +
846 +func TestDecodeHelloBadPadding(t *testing.T) {
847 + h := Hello{
848 + LayoutVersion: 1,
849 + SupportedProfiles: ProfileBaseline,
850 + MaxRequestPayloadBytes: 1024,
851 + MaxRequestBatchItems: 1,
852 + MaxResponsePayloadBytes: 1024,
853 + MaxResponseBatchItems: 1,
854 + AuthToken: 42,
855 + PacketSize: 65536,
856 + }
857 + var buf [64]byte
858 + h.Encode(buf[:])
859 +
860 + // Corrupt padding bytes at offset 28..32
861 + ne.PutUint32(buf[28:32], 0xFFFF)
862 +
863 + _, err := DecodeHello(buf[:44])
864 + if err != ErrBadLayout {
865 + t.Fatalf("expected ErrBadLayout for bad padding, got %v", err)
866 + }
867 +}
868 +
869 +// ---------------------------------------------------------------------------
870 +// CgroupsBuilder utility functions
871 +// ---------------------------------------------------------------------------
872 +
873 +func TestCgroupsBuilderSetHeader(t *testing.T) {
874 + var buf [4096]byte
875 + b := NewCgroupsBuilder(buf[:], 10, 0, 0)
876 +
877 + // Initial values should be zero
878 + if b.systemdEnabled != 0 {
879 + t.Fatalf("initial systemdEnabled should be 0, got %d", b.systemdEnabled)
880 + }
881 + if b.generation != 0 {
882 + t.Fatalf("initial generation should be 0, got %d", b.generation)
883 + }
884 +
885 + // Set header values
886 + b.SetHeader(1, 12345)
887 +
888 + if b.systemdEnabled != 1 {
889 + t.Fatalf("SetHeader systemdEnabled should be 1, got %d", b.systemdEnabled)
890 + }
891 + if b.generation != 12345 {
892 + t.Fatalf("SetHeader generation should be 12345, got %d", b.generation)
893 + }
894 +}
895 +
896 +func TestEstimateCgroupsMaxItems(t *testing.T) {
897 + // Buffer too small - should return 0
898 + maxItems := EstimateCgroupsMaxItems(cgroupsRespHdr)
899 + if maxItems != 0 {
900 + t.Fatalf("EstimateCgroupsMaxItems(%d) should be 0, got %d", cgroupsRespHdr, maxItems)
901 + }
902 +
903 + // Small buffer - should return 0
904 + maxItems = EstimateCgroupsMaxItems(cgroupsRespHdr + 10)
905 + if maxItems != 0 {
906 + t.Fatalf("EstimateCgroupsMaxItems(small) should be 0, got %d", maxItems)
907 + }
908 +
909 + // Reasonable buffer - should return positive value
910 + maxItems = EstimateCgroupsMaxItems(4096)
911 + if maxItems == 0 {
912 + t.Fatalf("EstimateCgroupsMaxItems(4096) should be > 0, got 0")
913 + }
914 +
915 + // Larger buffer - should return larger value
916 + maxItemsLarge := EstimateCgroupsMaxItems(65536)
917 + if maxItemsLarge <= maxItems {
918 + t.Fatalf("EstimateCgroupsMaxItems(65536)=%d should be > EstimateCgroupsMaxItems(4096)=%d", maxItemsLarge, maxItems)
919 + }
920 +}
src/go/pkg/netipc/protocol/frame.go new
+559
@@ -0,0 +1,559 @@
1 +// Package protocol implements the wire envelope and codec for the netipc
2 +// protocol. Pure byte-layout encode/decode. No I/O, no transport, no
3 +// allocation on decode. Localhost-only IPC — all multi-byte fields use
4 +// host byte order.
5 +//
6 +// Decoded "View" types borrow the underlying buffer and are valid only while
7 +// that buffer lives. Copy immediately if the data is needed later.
8 +package protocol
9 +
10 +import (
11 + "encoding/binary"
12 + "errors"
13 +)
14 +
15 +// ---------------------------------------------------------------------------
16 +// Constants
17 +// ---------------------------------------------------------------------------
18 +
19 +const (
20 + MagicMsg uint32 = 0x4e495043 // "NIPC"
21 + MagicChunk uint32 = 0x4e43484b // "NCHK"
22 + Version uint16 = 1
23 + HeaderLen uint16 = 32
24 + HeaderSize = 32
25 +
26 + // Message kinds.
27 + KindRequest uint16 = 1
28 + KindResponse uint16 = 2
29 + KindControl uint16 = 3
30 +
31 + // Flags.
32 + FlagBatch uint16 = 0x0001
33 +
34 + // Transport status.
35 + StatusOK uint16 = 0
36 + StatusBadEnvelope uint16 = 1
37 + StatusAuthFailed uint16 = 2
38 + StatusIncompatible uint16 = 3
39 + StatusUnsupported uint16 = 4
40 + StatusLimitExceeded uint16 = 5
41 + StatusInternalError uint16 = 6
42 +
43 + // Control opcodes.
44 + CodeHello uint16 = 1
45 + CodeHelloAck uint16 = 2
46 +
47 + // Method codes.
48 + MethodIncrement uint16 = 1
49 + MethodCgroupsSnapshot uint16 = 2
50 + MethodStringReverse uint16 = 3
51 +
52 + // Profile bits.
53 + ProfileBaseline uint32 = 0x01
54 + ProfileSHMHybrid uint32 = 0x02
55 + ProfileSHMFutex uint32 = 0x04
56 + ProfileSHMWaitAddr uint32 = 0x08
57 +
58 + // Defaults.
59 + MaxPayloadDefault uint32 = 1024
60 +
61 + // MaxPayloadCap is the hard cap on negotiated request payload sizes
62 + // (1 MiB) to prevent excessive memory allocation from a compromised peer.
63 + MaxPayloadCap uint32 = 1024 * 1024
64 +
65 + // Alignment for batch items and cgroups items.
66 + Alignment = 8
67 +
68 + // Payload sizes.
69 + helloSize = 44
70 + helloAckSize = 48
71 + cgroupsReqSize = 4
72 + cgroupsRespHdr = 24
73 + cgroupsDirEntry = 8
74 + cgroupsItemHdr = 32
75 +)
76 +
77 +var ne = binary.NativeEndian
78 +
79 +// ---------------------------------------------------------------------------
80 +// Errors
81 +// ---------------------------------------------------------------------------
82 +
83 +var (
84 + ErrTruncated = errors.New("buffer too short")
85 + ErrBadMagic = errors.New("magic value mismatch")
86 + ErrBadVersion = errors.New("unsupported version")
87 + ErrBadHeaderLen = errors.New("header_len != 32")
88 + ErrBadKind = errors.New("unknown message kind")
89 + ErrBadLayout = errors.New("unknown layout_version")
90 + ErrOutOfBounds = errors.New("offset+length exceeds data")
91 + ErrMissingNul = errors.New("string not NUL-terminated")
92 + ErrBadAlignment = errors.New("item not 8-byte aligned")
93 + ErrBadItemCount = errors.New("item count inconsistent")
94 + ErrOverflow = errors.New("builder out of space")
95 +)
96 +
97 +// ---------------------------------------------------------------------------
98 +// Utility
99 +// ---------------------------------------------------------------------------
100 +
101 +// Align8 rounds v up to the next multiple of 8.
102 +func Align8(v int) int {
103 + return (v + 7) &^ 7
104 +}
105 +
106 +// ---------------------------------------------------------------------------
107 +// Outer message header (32 bytes)
108 +// ---------------------------------------------------------------------------
109 +
110 +// Header is the outer message header (32 bytes on the wire).
111 +type Header struct {
112 + Magic uint32
113 + Version uint16
114 + HeaderLen uint16
115 + Kind uint16
116 + Flags uint16
117 + Code uint16
118 + TransportStatus uint16
119 + PayloadLen uint32
120 + ItemCount uint32
121 + MessageID uint64
122 +}
123 +
124 +// Encode writes the header into buf. Returns 32 on success, 0 if buf is
125 +// too small.
126 +func (h *Header) Encode(buf []byte) int {
127 + if len(buf) < HeaderSize {
128 + return 0
129 + }
130 + ne.PutUint32(buf[0:4], h.Magic)
131 + ne.PutUint16(buf[4:6], h.Version)
132 + ne.PutUint16(buf[6:8], h.HeaderLen)
133 + ne.PutUint16(buf[8:10], h.Kind)
134 + ne.PutUint16(buf[10:12], h.Flags)
135 + ne.PutUint16(buf[12:14], h.Code)
136 + ne.PutUint16(buf[14:16], h.TransportStatus)
137 + ne.PutUint32(buf[16:20], h.PayloadLen)
138 + ne.PutUint32(buf[20:24], h.ItemCount)
139 + ne.PutUint64(buf[24:32], h.MessageID)
140 + return HeaderSize
141 +}
142 +
143 +// DecodeHeader decodes an outer message header from buf. Validates magic,
144 +// version, header_len, and kind.
145 +func DecodeHeader(buf []byte) (Header, error) {
146 + if len(buf) < HeaderSize {
147 + return Header{}, ErrTruncated
148 + }
149 + h := Header{
150 + Magic: ne.Uint32(buf[0:4]),
151 + Version: ne.Uint16(buf[4:6]),
152 + HeaderLen: ne.Uint16(buf[6:8]),
153 + Kind: ne.Uint16(buf[8:10]),
154 + Flags: ne.Uint16(buf[10:12]),
155 + Code: ne.Uint16(buf[12:14]),
156 + TransportStatus: ne.Uint16(buf[14:16]),
157 + PayloadLen: ne.Uint32(buf[16:20]),
158 + ItemCount: ne.Uint32(buf[20:24]),
159 + MessageID: ne.Uint64(buf[24:32]),
160 + }
161 + if h.Magic != MagicMsg {
162 + return Header{}, ErrBadMagic
163 + }
164 + if h.Version != Version {
165 + return Header{}, ErrBadVersion
166 + }
167 + if h.HeaderLen != HeaderLen {
168 + return Header{}, ErrBadHeaderLen
169 + }
170 + if h.Kind < KindRequest || h.Kind > KindControl {
171 + return Header{}, ErrBadKind
172 + }
173 + return h, nil
174 +}
175 +
176 +// ---------------------------------------------------------------------------
177 +// Chunk continuation header (32 bytes)
178 +// ---------------------------------------------------------------------------
179 +
180 +// ChunkHeader is a chunk continuation header (32 bytes on the wire).
181 +type ChunkHeader struct {
182 + Magic uint32
183 + Version uint16
184 + Flags uint16
185 + MessageID uint64
186 + TotalMessageLen uint32
187 + ChunkIndex uint32
188 + ChunkCount uint32
189 + ChunkPayloadLen uint32
190 +}
191 +
192 +// Encode writes the chunk header into buf. Returns 32 on success, 0 if
193 +// buf is too small.
194 +func (c *ChunkHeader) Encode(buf []byte) int {
195 + if len(buf) < HeaderSize {
196 + return 0
197 + }
198 + ne.PutUint32(buf[0:4], c.Magic)
199 + ne.PutUint16(buf[4:6], c.Version)
200 + ne.PutUint16(buf[6:8], c.Flags)
201 + ne.PutUint64(buf[8:16], c.MessageID)
202 + ne.PutUint32(buf[16:20], c.TotalMessageLen)
203 + ne.PutUint32(buf[20:24], c.ChunkIndex)
204 + ne.PutUint32(buf[24:28], c.ChunkCount)
205 + ne.PutUint32(buf[28:32], c.ChunkPayloadLen)
206 + return HeaderSize
207 +}
208 +
209 +// DecodeChunkHeader decodes a chunk continuation header from buf.
210 +// Validates magic and version.
211 +func DecodeChunkHeader(buf []byte) (ChunkHeader, error) {
212 + if len(buf) < HeaderSize {
213 + return ChunkHeader{}, ErrTruncated
214 + }
215 + c := ChunkHeader{
216 + Magic: ne.Uint32(buf[0:4]),
217 + Version: ne.Uint16(buf[4:6]),
218 + Flags: ne.Uint16(buf[6:8]),
219 + MessageID: ne.Uint64(buf[8:16]),
220 + TotalMessageLen: ne.Uint32(buf[16:20]),
221 + ChunkIndex: ne.Uint32(buf[20:24]),
222 + ChunkCount: ne.Uint32(buf[24:28]),
223 + ChunkPayloadLen: ne.Uint32(buf[28:32]),
224 + }
225 + if c.Magic != MagicChunk {
226 + return ChunkHeader{}, ErrBadMagic
227 + }
228 + if c.Version != Version {
229 + return ChunkHeader{}, ErrBadVersion
230 + }
231 + if c.Flags != 0 {
232 + return ChunkHeader{}, ErrBadLayout
233 + }
234 + if c.ChunkPayloadLen == 0 {
235 + return ChunkHeader{}, ErrBadLayout
236 + }
237 + return c, nil
238 +}
239 +
240 +// ---------------------------------------------------------------------------
241 +// Batch item directory
242 +// ---------------------------------------------------------------------------
243 +
244 +// BatchEntry is one entry in a batch item directory (8 bytes on wire).
245 +type BatchEntry struct {
246 + Offset uint32
247 + Length uint32
248 +}
249 +
250 +// BatchDirEncode encodes entries into buf. Returns total bytes written
251 +// (len(entries) * 8), or 0 if buf is too small.
252 +func BatchDirEncode(entries []BatchEntry, buf []byte) int {
253 + need := len(entries) * 8
254 + if len(buf) < need {
255 + return 0
256 + }
257 + for i, e := range entries {
258 + base := i * 8
259 + ne.PutUint32(buf[base:base+4], e.Offset)
260 + ne.PutUint32(buf[base+4:base+8], e.Length)
261 + }
262 + return need
263 +}
264 +
265 +// BatchDirDecode decodes itemCount directory entries from buf. Validates
266 +// alignment and that each entry falls within packedAreaLen.
267 +func BatchDirDecode(buf []byte, itemCount uint32, packedAreaLen uint32) ([]BatchEntry, error) {
268 + count := int(itemCount)
269 + dirSize := count * 8
270 + if len(buf) < dirSize {
271 + return nil, ErrTruncated
272 + }
273 +
274 + out := make([]BatchEntry, count)
275 + for i := 0; i < count; i++ {
276 + base := i * 8
277 + off := ne.Uint32(buf[base : base+4])
278 + length := ne.Uint32(buf[base+4 : base+8])
279 +
280 + if int(off)%Alignment != 0 {
281 + return nil, ErrBadAlignment
282 + }
283 + if uint64(off)+uint64(length) > uint64(packedAreaLen) {
284 + return nil, ErrOutOfBounds
285 + }
286 + out[i] = BatchEntry{Offset: off, Length: length}
287 + }
288 + return out, nil
289 +}
290 +
291 +// BatchDirValidate validates the batch directory without allocating.
292 +// Checks alignment and that each entry falls within packedAreaLen.
293 +func BatchDirValidate(buf []byte, itemCount uint32, packedAreaLen uint32) error {
294 + count := int(itemCount)
295 + dirSize := count * 8
296 + if len(buf) < dirSize {
297 + return ErrTruncated
298 + }
299 + for i := 0; i < count; i++ {
300 + base := i * 8
301 + off := ne.Uint32(buf[base : base+4])
302 + length := ne.Uint32(buf[base+4 : base+8])
303 + if int(off)%Alignment != 0 {
304 + return ErrBadAlignment
305 + }
306 + if uint64(off)+uint64(length) > uint64(packedAreaLen) {
307 + return ErrOutOfBounds
308 + }
309 + }
310 + return nil
311 +}
312 +
313 +// BatchItemGet extracts a single batch item by index from a complete batch
314 +// payload. Returns the item slice on success.
315 +func BatchItemGet(payload []byte, itemCount uint32, index uint32) ([]byte, error) {
316 + if index >= itemCount {
317 + return nil, ErrOutOfBounds
318 + }
319 +
320 + dirSize := int(itemCount) * 8
321 + dirAligned := Align8(dirSize)
322 +
323 + if len(payload) < dirAligned {
324 + return nil, ErrTruncated
325 + }
326 +
327 + idx := int(index)
328 + base := idx * 8
329 + off := ne.Uint32(payload[base : base+4])
330 + length := ne.Uint32(payload[base+4 : base+8])
331 +
332 + packedAreaStart := dirAligned
333 + packedAreaLen := len(payload) - packedAreaStart
334 +
335 + if int(off)%Alignment != 0 {
336 + return nil, ErrBadAlignment
337 + }
338 + if uint64(off)+uint64(length) > uint64(packedAreaLen) {
339 + return nil, ErrOutOfBounds
340 + }
341 +
342 + start := packedAreaStart + int(off)
343 + end := start + int(length)
344 + return payload[start:end], nil
345 +}
346 +
347 +// ---------------------------------------------------------------------------
348 +// Batch builder
349 +// ---------------------------------------------------------------------------
350 +
351 +// BatchBuilder builds a batch payload: [directory] [align-pad] [packed items].
352 +type BatchBuilder struct {
353 + buf []byte
354 + itemCount uint32
355 + maxItems uint32
356 + dirEnd int // byte offset where directory reservation ends
357 + dataOffset int // current offset within the packed data area (relative)
358 +}
359 +
360 +// Reset reinitializes a batch builder against a caller-provided buffer.
361 +// This lets hot paths reuse stack-allocated builders instead of allocating
362 +// a fresh helper object for every request.
363 +func (b *BatchBuilder) Reset(buf []byte, maxItems uint32) {
364 + b.buf = buf
365 + b.itemCount = 0
366 + b.maxItems = maxItems
367 + b.dirEnd = Align8(int(maxItems) * 8)
368 + b.dataOffset = 0
369 +}
370 +
371 +// NewBatchBuilder creates a new batch builder. buf must be large enough for
372 +// maxItems*8 (directory) + packed data.
373 +func NewBatchBuilder(buf []byte, maxItems uint32) *BatchBuilder {
374 + b := &BatchBuilder{}
375 + b.Reset(buf, maxItems)
376 + return b
377 +}
378 +
379 +// Add appends an item payload. Handles alignment padding.
380 +func (b *BatchBuilder) Add(item []byte) error {
381 + if b.itemCount >= b.maxItems {
382 + return ErrOverflow
383 + }
384 +
385 + alignedOff := Align8(b.dataOffset)
386 + absPos := b.dirEnd + alignedOff
387 +
388 + if absPos+len(item) > len(b.buf) {
389 + return ErrOverflow
390 + }
391 +
392 + // Zero alignment padding.
393 + if alignedOff > b.dataOffset {
394 + padStart := b.dirEnd + b.dataOffset
395 + padEnd := b.dirEnd + alignedOff
396 + clear(b.buf[padStart:padEnd])
397 + }
398 +
399 + copy(b.buf[absPos:], item)
400 +
401 + // Write directory entry.
402 + idx := int(b.itemCount) * 8
403 + ne.PutUint32(b.buf[idx:idx+4], uint32(alignedOff))
404 + ne.PutUint32(b.buf[idx+4:idx+8], uint32(len(item)))
405 +
406 + b.dataOffset = alignedOff + len(item)
407 + b.itemCount++
408 + return nil
409 +}
410 +
411 +// Finish finalizes the batch. Returns (totalPayloadSize, itemCount).
412 +// Compacts if fewer items were added than maxItems.
413 +func (b *BatchBuilder) Finish() (int, uint32) {
414 + count := b.itemCount
415 + finalDirAligned := Align8(int(count) * 8)
416 +
417 + if finalDirAligned < b.dirEnd && b.dataOffset > 0 {
418 + // Shift packed data left.
419 + copy(b.buf[finalDirAligned:], b.buf[b.dirEnd:b.dirEnd+b.dataOffset])
420 + }
421 +
422 + total := finalDirAligned + Align8(b.dataOffset)
423 + return total, count
424 +}
425 +
426 +// ---------------------------------------------------------------------------
427 +// Hello payload (44 bytes)
428 +// ---------------------------------------------------------------------------
429 +
430 +// Hello is the client handshake payload (44 bytes on the wire).
431 +type Hello struct {
432 + LayoutVersion uint16
433 + Flags uint16
434 + SupportedProfiles uint32
435 + PreferredProfiles uint32
436 + MaxRequestPayloadBytes uint32
437 + MaxRequestBatchItems uint32
438 + MaxResponsePayloadBytes uint32
439 + MaxResponseBatchItems uint32
440 + AuthToken uint64
441 + PacketSize uint32
442 +}
443 +
444 +// Encode writes the Hello payload into buf. Returns 44 on success, 0 if
445 +// buf is too small.
446 +func (h *Hello) Encode(buf []byte) int {
447 + if len(buf) < helloSize {
448 + return 0
449 + }
450 + ne.PutUint16(buf[0:2], h.LayoutVersion)
451 + ne.PutUint16(buf[2:4], h.Flags)
452 + ne.PutUint32(buf[4:8], h.SupportedProfiles)
453 + ne.PutUint32(buf[8:12], h.PreferredProfiles)
454 + ne.PutUint32(buf[12:16], h.MaxRequestPayloadBytes)
455 + ne.PutUint32(buf[16:20], h.MaxRequestBatchItems)
456 + ne.PutUint32(buf[20:24], h.MaxResponsePayloadBytes)
457 + ne.PutUint32(buf[24:28], h.MaxResponseBatchItems)
458 + ne.PutUint32(buf[28:32], 0) // padding
459 + ne.PutUint64(buf[32:40], h.AuthToken)
460 + ne.PutUint32(buf[40:44], h.PacketSize)
461 + return helloSize
462 +}
463 +
464 +// DecodeHello decodes a Hello payload from buf. Validates layout_version.
465 +func DecodeHello(buf []byte) (Hello, error) {
466 + if len(buf) < helloSize {
467 + return Hello{}, ErrTruncated
468 + }
469 + h := Hello{
470 + LayoutVersion: ne.Uint16(buf[0:2]),
471 + Flags: ne.Uint16(buf[2:4]),
472 + SupportedProfiles: ne.Uint32(buf[4:8]),
473 + PreferredProfiles: ne.Uint32(buf[8:12]),
474 + MaxRequestPayloadBytes: ne.Uint32(buf[12:16]),
475 + MaxRequestBatchItems: ne.Uint32(buf[16:20]),
476 + MaxResponsePayloadBytes: ne.Uint32(buf[20:24]),
477 + MaxResponseBatchItems: ne.Uint32(buf[24:28]),
478 + // buf[28:32] is reserved padding, must be zero.
479 + AuthToken: ne.Uint64(buf[32:40]),
480 + PacketSize: ne.Uint32(buf[40:44]),
481 + }
482 + if h.LayoutVersion != 1 {
483 + return Hello{}, ErrBadLayout
484 + }
485 + // Validate padding bytes 28..32 are zero
486 + if ne.Uint32(buf[28:32]) != 0 {
487 + return Hello{}, ErrBadLayout
488 + }
489 + return h, nil
490 +}
491 +
492 +// ---------------------------------------------------------------------------
493 +// Hello-ack payload (44 bytes)
494 +// ---------------------------------------------------------------------------
495 +
496 +// HelloAck is the server handshake response payload (48 bytes on the wire).
497 +type HelloAck struct {
498 + LayoutVersion uint16
499 + Flags uint16
500 + ServerSupportedProfiles uint32
501 + IntersectionProfiles uint32
502 + SelectedProfile uint32
503 + AgreedMaxRequestPayloadBytes uint32
504 + AgreedMaxRequestBatchItems uint32
505 + AgreedMaxResponsePayloadBytes uint32
506 + AgreedMaxResponseBatchItems uint32
507 + AgreedPacketSize uint32
508 + SessionID uint64
509 +}
510 +
511 +// Encode writes the HelloAck payload into buf. Returns 48 on success, 0
512 +// if buf is too small.
513 +func (h *HelloAck) Encode(buf []byte) int {
514 + if len(buf) < helloAckSize {
515 + return 0
516 + }
517 + ne.PutUint16(buf[0:2], h.LayoutVersion)
518 + ne.PutUint16(buf[2:4], h.Flags)
519 + ne.PutUint32(buf[4:8], h.ServerSupportedProfiles)
520 + ne.PutUint32(buf[8:12], h.IntersectionProfiles)
521 + ne.PutUint32(buf[12:16], h.SelectedProfile)
522 + ne.PutUint32(buf[16:20], h.AgreedMaxRequestPayloadBytes)
523 + ne.PutUint32(buf[20:24], h.AgreedMaxRequestBatchItems)
524 + ne.PutUint32(buf[24:28], h.AgreedMaxResponsePayloadBytes)
525 + ne.PutUint32(buf[28:32], h.AgreedMaxResponseBatchItems)
526 + ne.PutUint32(buf[32:36], h.AgreedPacketSize)
527 + ne.PutUint32(buf[36:40], 0) // padding
528 + ne.PutUint64(buf[40:48], h.SessionID)
529 + return helloAckSize
530 +}
531 +
532 +// DecodeHelloAck decodes a HelloAck payload from buf. Validates
533 +// layout_version.
534 +func DecodeHelloAck(buf []byte) (HelloAck, error) {
535 + if len(buf) < helloAckSize {
536 + return HelloAck{}, ErrTruncated
537 + }
538 + h := HelloAck{
539 + LayoutVersion: ne.Uint16(buf[0:2]),
540 + Flags: ne.Uint16(buf[2:4]),
541 + ServerSupportedProfiles: ne.Uint32(buf[4:8]),
542 + IntersectionProfiles: ne.Uint32(buf[8:12]),
543 + SelectedProfile: ne.Uint32(buf[12:16]),
544 + AgreedMaxRequestPayloadBytes: ne.Uint32(buf[16:20]),
545 + AgreedMaxRequestBatchItems: ne.Uint32(buf[20:24]),
546 + AgreedMaxResponsePayloadBytes: ne.Uint32(buf[24:28]),
547 + AgreedMaxResponseBatchItems: ne.Uint32(buf[28:32]),
548 + AgreedPacketSize: ne.Uint32(buf[32:36]),
549 + // skip padding at 36:40
550 + SessionID: ne.Uint64(buf[40:48]),
551 + }
552 + if h.LayoutVersion != 1 {
553 + return HelloAck{}, ErrBadLayout
554 + }
555 + if h.Flags != 0 {
556 + return HelloAck{}, ErrBadLayout
557 + }
558 + return h, nil
559 +}
src/go/pkg/netipc/protocol/frame_test.go new
+1685
@@ -0,0 +1,1685 @@
1 +package protocol
2 +
3 +import (
4 + "bytes"
5 + "testing"
6 +)
7 +
8 +// ---------------------------------------------------------------------------
9 +// Utility tests
10 +// ---------------------------------------------------------------------------
11 +
12 +func TestAlign8(t *testing.T) {
13 + cases := []struct {
14 + in, want int
15 + }{
16 + {0, 0},
17 + {1, 8},
18 + {7, 8},
19 + {8, 8},
20 + {9, 16},
21 + {15, 16},
22 + {16, 16},
23 + {17, 24},
24 + }
25 + for _, tc := range cases {
26 + if got := Align8(tc.in); got != tc.want {
27 + t.Errorf("Align8(%d) = %d, want %d", tc.in, got, tc.want)
28 + }
29 + }
30 +}
31 +
32 +// ---------------------------------------------------------------------------
33 +// Outer message header tests
34 +// ---------------------------------------------------------------------------
35 +
36 +func TestHeaderRoundtrip(t *testing.T) {
37 + h := Header{
38 + Magic: MagicMsg,
39 + Version: Version,
40 + HeaderLen: HeaderLen,
41 + Kind: KindRequest,
42 + Flags: FlagBatch,
43 + Code: MethodCgroupsSnapshot,
44 + TransportStatus: StatusOK,
45 + PayloadLen: 12345,
46 + ItemCount: 42,
47 + MessageID: 0xDEADBEEFCAFEBABE,
48 + }
49 +
50 + var buf [64]byte
51 + n := h.Encode(buf[:])
52 + if n != 32 {
53 + t.Fatalf("Encode returned %d, want 32", n)
54 + }
55 +
56 + out, err := DecodeHeader(buf[:n])
57 + if err != nil {
58 + t.Fatalf("Decode error: %v", err)
59 + }
60 + if out != h {
61 + t.Fatalf("roundtrip mismatch:\ngot: %+v\nwant: %+v", out, h)
62 + }
63 +}
64 +
65 +func TestHeaderEncodeTooSmall(t *testing.T) {
66 + h := Header{}
67 + var buf [16]byte
68 + if n := h.Encode(buf[:]); n != 0 {
69 + t.Fatalf("Encode returned %d, want 0", n)
70 + }
71 +}
72 +
73 +func TestHeaderDecodeTruncated(t *testing.T) {
74 + buf := make([]byte, 31)
75 + _, err := DecodeHeader(buf)
76 + if err != ErrTruncated {
77 + t.Fatalf("got %v, want ErrTruncated", err)
78 + }
79 +}
80 +
81 +func TestHeaderDecodeBadMagic(t *testing.T) {
82 + h := Header{
83 + Magic: 0x12345678,
84 + Version: Version,
85 + HeaderLen: HeaderLen,
86 + Kind: KindRequest,
87 + }
88 + var buf [32]byte
89 + h.Encode(buf[:])
90 + _, err := DecodeHeader(buf[:])
91 + if err != ErrBadMagic {
92 + t.Fatalf("got %v, want ErrBadMagic", err)
93 + }
94 +}
95 +
96 +func TestHeaderDecodeBadVersion(t *testing.T) {
97 + h := Header{
98 + Magic: MagicMsg,
99 + Version: 99,
100 + HeaderLen: HeaderLen,
101 + Kind: KindRequest,
102 + }
103 + var buf [32]byte
104 + h.Encode(buf[:])
105 + _, err := DecodeHeader(buf[:])
106 + if err != ErrBadVersion {
107 + t.Fatalf("got %v, want ErrBadVersion", err)
108 + }
109 +}
110 +
111 +func TestHeaderDecodeBadHeaderLen(t *testing.T) {
112 + h := Header{
113 + Magic: MagicMsg,
114 + Version: Version,
115 + HeaderLen: 64,
116 + Kind: KindRequest,
117 + }
118 + var buf [32]byte
119 + h.Encode(buf[:])
120 + _, err := DecodeHeader(buf[:])
121 + if err != ErrBadHeaderLen {
122 + t.Fatalf("got %v, want ErrBadHeaderLen", err)
123 + }
124 +}
125 +
126 +func TestHeaderDecodeBadKind(t *testing.T) {
127 + // kind = 0
128 + h := Header{
129 + Magic: MagicMsg,
130 + Version: Version,
131 + HeaderLen: HeaderLen,
132 + Kind: 0,
133 + }
134 + var buf [32]byte
135 + h.Encode(buf[:])
136 + _, err := DecodeHeader(buf[:])
137 + if err != ErrBadKind {
138 + t.Fatalf("kind=0: got %v, want ErrBadKind", err)
139 + }
140 +
141 + // kind = 4
142 + h.Kind = 4
143 + h.Encode(buf[:])
144 + _, err = DecodeHeader(buf[:])
145 + if err != ErrBadKind {
146 + t.Fatalf("kind=4: got %v, want ErrBadKind", err)
147 + }
148 +}
149 +
150 +func TestHeaderAllKinds(t *testing.T) {
151 + for k := KindRequest; k <= KindControl; k++ {
152 + h := Header{
153 + Magic: MagicMsg,
154 + Version: Version,
155 + HeaderLen: HeaderLen,
156 + Kind: k,
157 + }
158 + var buf [32]byte
159 + h.Encode(buf[:])
160 + out, err := DecodeHeader(buf[:])
161 + if err != nil {
162 + t.Fatalf("kind=%d: %v", k, err)
163 + }
164 + if out.Kind != k {
165 + t.Fatalf("kind=%d: got %d", k, out.Kind)
166 + }
167 + }
168 +}
169 +
170 +func TestHeaderWireBytes(t *testing.T) {
171 + h := Header{
172 + Magic: MagicMsg,
173 + Version: Version,
174 + HeaderLen: HeaderLen,
175 + Kind: KindRequest,
176 + Flags: 0,
177 + Code: MethodCgroupsSnapshot,
178 + TransportStatus: StatusOK,
179 + PayloadLen: 4,
180 + ItemCount: 1,
181 + MessageID: 1,
182 + }
183 +
184 + var buf [32]byte
185 + h.Encode(buf[:])
186 +
187 + // magic = 0x4e495043 LE: 43 50 49 4e
188 + if !bytes.Equal(buf[0:4], []byte{0x43, 0x50, 0x49, 0x4e}) {
189 + t.Errorf("magic bytes: %x", buf[0:4])
190 + }
191 + // version = 1 LE: 01 00
192 + if !bytes.Equal(buf[4:6], []byte{0x01, 0x00}) {
193 + t.Errorf("version bytes: %x", buf[4:6])
194 + }
195 + // header_len = 32 LE: 20 00
196 + if !bytes.Equal(buf[6:8], []byte{0x20, 0x00}) {
197 + t.Errorf("header_len bytes: %x", buf[6:8])
198 + }
199 + // kind = 1 LE: 01 00
200 + if !bytes.Equal(buf[8:10], []byte{0x01, 0x00}) {
201 + t.Errorf("kind bytes: %x", buf[8:10])
202 + }
203 + // code = 2 LE: 02 00
204 + if !bytes.Equal(buf[12:14], []byte{0x02, 0x00}) {
205 + t.Errorf("code bytes: %x", buf[12:14])
206 + }
207 +}
208 +
209 +// ---------------------------------------------------------------------------
210 +// Chunk continuation header tests
211 +// ---------------------------------------------------------------------------
212 +
213 +func TestChunkHeaderRoundtrip(t *testing.T) {
214 + c := ChunkHeader{
215 + Magic: MagicChunk,
216 + Version: Version,
217 + Flags: 0,
218 + MessageID: 0x1234567890ABCDEF,
219 + TotalMessageLen: 100000,
220 + ChunkIndex: 3,
221 + ChunkCount: 10,
222 + ChunkPayloadLen: 8192,
223 + }
224 +
225 + var buf [64]byte
226 + n := c.Encode(buf[:])
227 + if n != 32 {
228 + t.Fatalf("Encode returned %d, want 32", n)
229 + }
230 +
231 + out, err := DecodeChunkHeader(buf[:n])
232 + if err != nil {
233 + t.Fatalf("Decode error: %v", err)
234 + }
235 + if out != c {
236 + t.Fatalf("roundtrip mismatch:\ngot: %+v\nwant: %+v", out, c)
237 + }
238 +}
239 +
240 +func TestChunkDecodeTruncated(t *testing.T) {
241 + buf := make([]byte, 31)
242 + _, err := DecodeChunkHeader(buf)
243 + if err != ErrTruncated {
244 + t.Fatalf("got %v, want ErrTruncated", err)
245 + }
246 +}
247 +
248 +func TestChunkDecodeBadMagic(t *testing.T) {
249 + c := ChunkHeader{
250 + Magic: MagicMsg, // wrong magic for chunk
251 + Version: Version,
252 + }
253 + var buf [32]byte
254 + c.Encode(buf[:])
255 + _, err := DecodeChunkHeader(buf[:])
256 + if err != ErrBadMagic {
257 + t.Fatalf("got %v, want ErrBadMagic", err)
258 + }
259 +}
260 +
261 +func TestChunkDecodeBadVersion(t *testing.T) {
262 + c := ChunkHeader{
263 + Magic: MagicChunk,
264 + Version: 2,
265 + }
266 + var buf [32]byte
267 + c.Encode(buf[:])
268 + _, err := DecodeChunkHeader(buf[:])
269 + if err != ErrBadVersion {
270 + t.Fatalf("got %v, want ErrBadVersion", err)
271 + }
272 +}
273 +
274 +func TestChunkEncodeTooSmall(t *testing.T) {
275 + c := ChunkHeader{}
276 + var buf [16]byte
277 + if n := c.Encode(buf[:]); n != 0 {
278 + t.Fatalf("Encode returned %d, want 0", n)
279 + }
280 +}
281 +
282 +func TestChunkWireBytes(t *testing.T) {
283 + c := ChunkHeader{
284 + Magic: MagicChunk,
285 + Version: Version,
286 + Flags: 0,
287 + MessageID: 1,
288 + TotalMessageLen: 256,
289 + ChunkIndex: 1,
290 + ChunkCount: 3,
291 + ChunkPayloadLen: 100,
292 + }
293 +
294 + var buf [32]byte
295 + c.Encode(buf[:])
296 +
297 + // magic = 0x4e43484b LE: 4b 48 43 4e
298 + if !bytes.Equal(buf[0:4], []byte{0x4b, 0x48, 0x43, 0x4e}) {
299 + t.Errorf("magic bytes: %x", buf[0:4])
300 + }
301 +}
302 +
303 +// ---------------------------------------------------------------------------
304 +// Batch item directory tests
305 +// ---------------------------------------------------------------------------
306 +
307 +func TestBatchDirRoundtrip(t *testing.T) {
308 + entries := []BatchEntry{
309 + {Offset: 0, Length: 100},
310 + {Offset: 104, Length: 200},
311 + {Offset: 304, Length: 50},
312 + }
313 +
314 + buf := make([]byte, 24)
315 + n := BatchDirEncode(entries, buf)
316 + if n != 24 {
317 + t.Fatalf("BatchDirEncode returned %d, want 24", n)
318 + }
319 +
320 + out, err := BatchDirDecode(buf, 3, 1000)
321 + if err != nil {
322 + t.Fatalf("BatchDirDecode error: %v", err)
323 + }
324 + for i, e := range entries {
325 + if out[i] != e {
326 + t.Errorf("entry[%d]: got %+v, want %+v", i, out[i], e)
327 + }
328 + }
329 +}
330 +
331 +func TestBatchDirEncodeTooSmall(t *testing.T) {
332 + entries := []BatchEntry{{Offset: 0, Length: 10}}
333 + buf := make([]byte, 4)
334 + if n := BatchDirEncode(entries, buf); n != 0 {
335 + t.Fatalf("got %d, want 0", n)
336 + }
337 +}
338 +
339 +func TestBatchDirDecodeTruncated(t *testing.T) {
340 + buf := make([]byte, 12)
341 + _, err := BatchDirDecode(buf, 2, 1000)
342 + if err != ErrTruncated {
343 + t.Fatalf("got %v, want ErrTruncated", err)
344 + }
345 +}
346 +
347 +func TestBatchDirDecodeBadAlignment(t *testing.T) {
348 + buf := make([]byte, 8)
349 + ne.PutUint32(buf[0:4], 3) // offset not aligned to 8
350 + ne.PutUint32(buf[4:8], 10)
351 + _, err := BatchDirDecode(buf, 1, 100)
352 + if err != ErrBadAlignment {
353 + t.Fatalf("got %v, want ErrBadAlignment", err)
354 + }
355 +}
356 +
357 +func TestBatchDirDecodeOutOfBounds(t *testing.T) {
358 + buf := make([]byte, 8)
359 + ne.PutUint32(buf[0:4], 0)
360 + ne.PutUint32(buf[4:8], 200) // exceeds packed area
361 + _, err := BatchDirDecode(buf, 1, 100)
362 + if err != ErrOutOfBounds {
363 + t.Fatalf("got %v, want ErrOutOfBounds", err)
364 + }
365 +}
366 +
367 +func TestBatchItemGetBasic(t *testing.T) {
368 + // Build a batch with 2 items using the builder.
369 + buf := make([]byte, 1024)
370 + b := NewBatchBuilder(buf, 2)
371 +
372 + item0 := []byte("hello")
373 + item1 := []byte("world!!!")
374 +
375 + if err := b.Add(item0); err != nil {
376 + t.Fatal(err)
377 + }
378 + if err := b.Add(item1); err != nil {
379 + t.Fatal(err)
380 + }
381 +
382 + total, count := b.Finish()
383 + if count != 2 {
384 + t.Fatalf("count = %d, want 2", count)
385 + }
386 +
387 + // Extract items.
388 + got0, err := BatchItemGet(buf[:total], count, 0)
389 + if err != nil {
390 + t.Fatalf("item 0: %v", err)
391 + }
392 + if !bytes.Equal(got0, item0) {
393 + t.Fatalf("item 0: got %q, want %q", got0, item0)
394 + }
395 +
396 + got1, err := BatchItemGet(buf[:total], count, 1)
397 + if err != nil {
398 + t.Fatalf("item 1: %v", err)
399 + }
400 + if !bytes.Equal(got1, item1) {
401 + t.Fatalf("item 1: got %q, want %q", got1, item1)
402 + }
403 +}
404 +
405 +func TestBatchItemGetOutOfBounds(t *testing.T) {
406 + buf := make([]byte, 16)
407 + _, err := BatchItemGet(buf, 1, 1) // index >= count
408 + if err != ErrOutOfBounds {
409 + t.Fatalf("got %v, want ErrOutOfBounds", err)
410 + }
411 +}
412 +
413 +func TestBatchItemGetTruncated(t *testing.T) {
414 + buf := make([]byte, 4) // too small for even 1 directory entry aligned
415 + _, err := BatchItemGet(buf, 1, 0)
416 + if err != ErrTruncated {
417 + t.Fatalf("got %v, want ErrTruncated", err)
418 + }
419 +}
420 +
421 +// ---------------------------------------------------------------------------
422 +// Batch builder tests
423 +// ---------------------------------------------------------------------------
424 +
425 +func TestBatchBuilderOverflowMaxItems(t *testing.T) {
426 + buf := make([]byte, 1024)
427 + b := NewBatchBuilder(buf, 1)
428 + if err := b.Add([]byte("one")); err != nil {
429 + t.Fatal(err)
430 + }
431 + if err := b.Add([]byte("two")); err != ErrOverflow {
432 + t.Fatalf("got %v, want ErrOverflow", err)
433 + }
434 +}
435 +
436 +func TestBatchBuilderOverflowBuffer(t *testing.T) {
437 + buf := make([]byte, 16) // dir = 8 bytes, 8 bytes data
438 + b := NewBatchBuilder(buf, 1)
439 + if err := b.Add(make([]byte, 100)); err != ErrOverflow {
440 + t.Fatalf("got %v, want ErrOverflow", err)
441 + }
442 +}
443 +
444 +func TestBatchBuilderEmpty(t *testing.T) {
445 + buf := make([]byte, 64)
446 + b := NewBatchBuilder(buf, 5)
447 + total, count := b.Finish()
448 + if count != 0 {
449 + t.Fatalf("count = %d, want 0", count)
450 + }
451 + if total != 0 {
452 + t.Fatalf("total = %d, want 0", total)
453 + }
454 +}
455 +
456 +func TestBatchBuilderCompaction(t *testing.T) {
457 + // Reserve space for 10 items, only add 1.
458 + buf := make([]byte, 1024)
459 + b := NewBatchBuilder(buf, 10)
460 + if err := b.Add([]byte("compact")); err != nil {
461 + t.Fatal(err)
462 + }
463 + total, count := b.Finish()
464 + if count != 1 {
465 + t.Fatalf("count = %d, want 1", count)
466 + }
467 +
468 + // Verify the item can be extracted.
469 + got, err := BatchItemGet(buf[:total], count, 0)
470 + if err != nil {
471 + t.Fatal(err)
472 + }
473 + if !bytes.Equal(got, []byte("compact")) {
474 + t.Fatalf("got %q, want %q", got, "compact")
475 + }
476 +}
477 +
478 +// ---------------------------------------------------------------------------
479 +// Hello payload tests
480 +// ---------------------------------------------------------------------------
481 +
482 +func TestHelloRoundtrip(t *testing.T) {
483 + h := Hello{
484 + LayoutVersion: 1,
485 + Flags: 0,
486 + SupportedProfiles: ProfileBaseline | ProfileSHMFutex,
487 + PreferredProfiles: ProfileSHMFutex,
488 + MaxRequestPayloadBytes: 4096,
489 + MaxRequestBatchItems: 100,
490 + MaxResponsePayloadBytes: 1048576,
491 + MaxResponseBatchItems: 1,
492 + AuthToken: 0xAABBCCDDEEFF0011,
493 + PacketSize: 65536,
494 + }
495 +
496 + var buf [64]byte
497 + n := h.Encode(buf[:])
498 + if n != 44 {
499 + t.Fatalf("Encode returned %d, want 44", n)
500 + }
501 +
502 + out, err := DecodeHello(buf[:n])
503 + if err != nil {
504 + t.Fatalf("Decode error: %v", err)
505 + }
506 + if out != h {
507 + t.Fatalf("roundtrip mismatch:\ngot: %+v\nwant: %+v", out, h)
508 + }
509 +}
510 +
511 +func TestHelloEncodeTooSmall(t *testing.T) {
512 + h := Hello{LayoutVersion: 1}
513 + var buf [20]byte
514 + if n := h.Encode(buf[:]); n != 0 {
515 + t.Fatalf("Encode returned %d, want 0", n)
516 + }
517 +}
518 +
519 +func TestHelloDecodeTruncated(t *testing.T) {
520 + buf := make([]byte, 43)
521 + _, err := DecodeHello(buf)
522 + if err != ErrTruncated {
523 + t.Fatalf("got %v, want ErrTruncated", err)
524 + }
525 +}
526 +
527 +func TestHelloDecodeBadLayout(t *testing.T) {
528 + h := Hello{LayoutVersion: 2}
529 + var buf [44]byte
530 + h.Encode(buf[:])
531 + _, err := DecodeHello(buf[:])
532 + if err != ErrBadLayout {
533 + t.Fatalf("got %v, want ErrBadLayout", err)
534 + }
535 +}
536 +
537 +func TestHelloPaddingIsZero(t *testing.T) {
538 + h := Hello{
539 + LayoutVersion: 1,
540 + MaxResponseBatchItems: 0xFFFFFFFF,
541 + AuthToken: 0xAAAAAAAAAAAAAAAA,
542 + }
543 + var buf [44]byte
544 + h.Encode(buf[:])
545 +
546 + // Padding at offset 28..32 must be zero.
547 + if !bytes.Equal(buf[28:32], []byte{0, 0, 0, 0}) {
548 + t.Errorf("padding not zero: %x", buf[28:32])
549 + }
550 +}
551 +
552 +func TestHelloDecodeNonzeroPadding(t *testing.T) {
553 + h := Hello{
554 + LayoutVersion: 1,
555 + SupportedProfiles: ProfileBaseline,
556 + MaxRequestPayloadBytes: 1024,
557 + MaxRequestBatchItems: 1,
558 + MaxResponsePayloadBytes: 1024,
559 + MaxResponseBatchItems: 1,
560 + PacketSize: 65536,
561 + }
562 + var buf [44]byte
563 + h.Encode(buf[:])
564 +
565 + // Valid first
566 + if _, err := DecodeHello(buf[:]); err != nil {
567 + t.Fatalf("valid hello failed: %v", err)
568 + }
569 +
570 + // Corrupt padding
571 + buf[28] = 0xFF
572 + if _, err := DecodeHello(buf[:]); err != ErrBadLayout {
573 + t.Errorf("nonzero padding: got %v, want ErrBadLayout", err)
574 + }
575 +}
576 +
577 +func TestHelloWireBytes(t *testing.T) {
578 + h := Hello{
579 + LayoutVersion: 1,
580 + Flags: 0,
581 + SupportedProfiles: ProfileBaseline | ProfileSHMFutex,
582 + PreferredProfiles: ProfileSHMFutex,
583 + MaxRequestPayloadBytes: 4096,
584 + MaxRequestBatchItems: 100,
585 + MaxResponsePayloadBytes: 1048576,
586 + MaxResponseBatchItems: 1,
587 + AuthToken: 0xAABBCCDDEEFF0011,
588 + PacketSize: 65536,
589 + }
590 + var buf [44]byte
591 + h.Encode(buf[:])
592 +
593 + // supported_profiles = 0x05 LE at offset 4
594 + if !bytes.Equal(buf[4:8], []byte{0x05, 0x00, 0x00, 0x00}) {
595 + t.Errorf("supported_profiles bytes: %x", buf[4:8])
596 + }
597 + // auth_token at offset 32
598 + if !bytes.Equal(buf[32:40], []byte{0x11, 0x00, 0xFF, 0xEE, 0xDD, 0xCC, 0xBB, 0xAA}) {
599 + t.Errorf("auth_token bytes: %x", buf[32:40])
600 + }
601 +}
602 +
603 +// ---------------------------------------------------------------------------
604 +// Hello-ack payload tests
605 +// ---------------------------------------------------------------------------
606 +
607 +func TestHelloAckRoundtrip(t *testing.T) {
608 + h := HelloAck{
609 + LayoutVersion: 1,
610 + Flags: 0,
611 + ServerSupportedProfiles: 0x07,
612 + IntersectionProfiles: 0x05,
613 + SelectedProfile: ProfileSHMFutex,
614 + AgreedMaxRequestPayloadBytes: 2048,
615 + AgreedMaxRequestBatchItems: 50,
616 + AgreedMaxResponsePayloadBytes: 65536,
617 + AgreedMaxResponseBatchItems: 1,
618 + AgreedPacketSize: 32768,
619 + }
620 +
621 + var buf [64]byte
622 + n := h.Encode(buf[:])
623 + if n != 48 {
624 + t.Fatalf("Encode returned %d, want 48", n)
625 + }
626 +
627 + out, err := DecodeHelloAck(buf[:n])
628 + if err != nil {
629 + t.Fatalf("Decode error: %v", err)
630 + }
631 + if out != h {
632 + t.Fatalf("roundtrip mismatch:\ngot: %+v\nwant: %+v", out, h)
633 + }
634 +}
635 +
636 +func TestHelloAckEncodeTooSmall(t *testing.T) {
637 + h := HelloAck{LayoutVersion: 1}
638 + var buf [20]byte
639 + if n := h.Encode(buf[:]); n != 0 {
640 + t.Fatalf("Encode returned %d, want 0", n)
641 + }
642 +}
643 +
644 +func TestHelloAckDecodeTruncated(t *testing.T) {
645 + buf := make([]byte, 47)
646 + _, err := DecodeHelloAck(buf)
647 + if err != ErrTruncated {
648 + t.Fatalf("got %v, want ErrTruncated", err)
649 + }
650 +}
651 +
652 +func TestHelloAckDecodeBadLayout(t *testing.T) {
653 + h := HelloAck{LayoutVersion: 99}
654 + var buf [48]byte
655 + h.Encode(buf[:])
656 + _, err := DecodeHelloAck(buf[:])
657 + if err != ErrBadLayout {
658 + t.Fatalf("got %v, want ErrBadLayout", err)
659 + }
660 +}
661 +
662 +// ---------------------------------------------------------------------------
663 +// Cgroups request tests
664 +// ---------------------------------------------------------------------------
665 +
666 +func TestCgroupsRequestRoundtrip(t *testing.T) {
667 + r := CgroupsRequest{LayoutVersion: 1, Flags: 0}
668 +
669 + var buf [16]byte
670 + n := r.Encode(buf[:])
671 + if n != 4 {
672 + t.Fatalf("Encode returned %d, want 4", n)
673 + }
674 +
675 + out, err := DecodeCgroupsRequest(buf[:n])
676 + if err != nil {
677 + t.Fatalf("Decode error: %v", err)
678 + }
679 + if out != r {
680 + t.Fatalf("roundtrip mismatch:\ngot: %+v\nwant: %+v", out, r)
681 + }
682 +}
683 +
684 +func TestCgroupsRequestEncodeTooSmall(t *testing.T) {
685 + r := CgroupsRequest{LayoutVersion: 1}
686 + var buf [2]byte
687 + if n := r.Encode(buf[:]); n != 0 {
688 + t.Fatalf("Encode returned %d, want 0", n)
689 + }
690 +}
691 +
692 +func TestCgroupsRequestDecodeTruncated(t *testing.T) {
693 + buf := make([]byte, 3)
694 + _, err := DecodeCgroupsRequest(buf)
695 + if err != ErrTruncated {
696 + t.Fatalf("got %v, want ErrTruncated", err)
697 + }
698 +}
699 +
700 +func TestCgroupsRequestDecodeBadLayout(t *testing.T) {
701 + r := CgroupsRequest{LayoutVersion: 2}
702 + var buf [4]byte
703 + r.Encode(buf[:])
704 + _, err := DecodeCgroupsRequest(buf[:])
705 + if err != ErrBadLayout {
706 + t.Fatalf("got %v, want ErrBadLayout", err)
707 + }
708 +}
709 +
710 +// ---------------------------------------------------------------------------
711 +// CStringView tests
712 +// ---------------------------------------------------------------------------
713 +
714 +func TestCStringViewBasic(t *testing.T) {
715 + data := []byte("hello\x00")
716 + v := NewCStringView(data, 5)
717 +
718 + if v.Len() != 5 {
719 + t.Fatalf("Len() = %d, want 5", v.Len())
720 + }
721 + if !bytes.Equal(v.Bytes(), []byte("hello")) {
722 + t.Fatalf("Bytes() = %q, want %q", v.Bytes(), "hello")
723 + }
724 + if v.String() != "hello" {
725 + t.Fatalf("String() = %q, want %q", v.String(), "hello")
726 + }
727 +}
728 +
729 +func TestCStringViewEmpty(t *testing.T) {
730 + data := []byte{0}
731 + v := NewCStringView(data, 0)
732 +
733 + if v.Len() != 0 {
734 + t.Fatalf("Len() = %d, want 0", v.Len())
735 + }
736 + if len(v.Bytes()) != 0 {
737 + t.Fatalf("Bytes() len = %d, want 0", len(v.Bytes()))
738 + }
739 + if v.String() != "" {
740 + t.Fatalf("String() = %q, want empty", v.String())
741 + }
742 +}
743 +
744 +// ---------------------------------------------------------------------------
745 +// Cgroups snapshot response tests
746 +// ---------------------------------------------------------------------------
747 +
748 +func TestCgroupsResponseEmptyRoundtrip(t *testing.T) {
749 + buf := make([]byte, 8192)
750 + b := NewCgroupsBuilder(buf, 0, 0, 42)
751 + total := b.Finish()
752 + if total != 24 {
753 + t.Fatalf("Finish returned %d, want 24", total)
754 + }
755 +
756 + view, err := DecodeCgroupsResponse(buf[:total])
757 + if err != nil {
758 + t.Fatalf("Decode error: %v", err)
759 + }
760 + if view.ItemCount != 0 {
761 + t.Fatalf("ItemCount = %d, want 0", view.ItemCount)
762 + }
763 + if view.SystemdEnabled != 0 {
764 + t.Fatalf("SystemdEnabled = %d, want 0", view.SystemdEnabled)
765 + }
766 + if view.Generation != 42 {
767 + t.Fatalf("Generation = %d, want 42", view.Generation)
768 + }
769 +}
770 +
771 +func TestCgroupsResponseSingleItemRoundtrip(t *testing.T) {
772 + buf := make([]byte, 8192)
773 + b := NewCgroupsBuilder(buf, 1, 1, 100)
774 +
775 + name := []byte("init.scope")
776 + path := []byte("/sys/fs/cgroup/init.scope")
777 + if err := b.Add(42, 0x01, 1, name, path); err != nil {
778 + t.Fatal(err)
779 + }
780 +
781 + total := b.Finish()
782 +
783 + view, err := DecodeCgroupsResponse(buf[:total])
784 + if err != nil {
785 + t.Fatalf("Decode error: %v", err)
786 + }
787 + if view.ItemCount != 1 {
788 + t.Fatalf("ItemCount = %d, want 1", view.ItemCount)
789 + }
790 + if view.SystemdEnabled != 1 {
791 + t.Fatalf("SystemdEnabled = %d, want 1", view.SystemdEnabled)
792 + }
793 + if view.Generation != 100 {
794 + t.Fatalf("Generation = %d, want 100", view.Generation)
795 + }
796 +
797 + item, err := view.Item(0)
798 + if err != nil {
799 + t.Fatalf("Item(0) error: %v", err)
800 + }
801 + if item.Hash != 42 {
802 + t.Fatalf("Hash = %d, want 42", item.Hash)
803 + }
804 + if item.Options != 0x01 {
805 + t.Fatalf("Options = %d, want 1", item.Options)
806 + }
807 + if item.Enabled != 1 {
808 + t.Fatalf("Enabled = %d, want 1", item.Enabled)
809 + }
810 + if item.Name.String() != "init.scope" {
811 + t.Fatalf("Name = %q, want %q", item.Name.String(), "init.scope")
812 + }
813 + if item.Path.String() != "/sys/fs/cgroup/init.scope" {
814 + t.Fatalf("Path = %q, want %q", item.Path.String(), "/sys/fs/cgroup/init.scope")
815 + }
816 +}
817 +
818 +func TestCgroupsResponseMultiItemRoundtrip(t *testing.T) {
819 + buf := make([]byte, 8192)
820 + b := NewCgroupsBuilder(buf, 3, 1, 999)
821 +
822 + if err := b.Add(100, 0, 1,
823 + []byte("init.scope"),
824 + []byte("/sys/fs/cgroup/init.scope")); err != nil {
825 + t.Fatal(err)
826 + }
827 + if err := b.Add(200, 0x02, 0,
828 + []byte("system.slice/docker-abc.scope"),
829 + []byte("/sys/fs/cgroup/system.slice/docker-abc.scope")); err != nil {
830 + t.Fatal(err)
831 + }
832 + if err := b.Add(300, 0, 1, []byte(""), []byte("")); err != nil {
833 + t.Fatal(err)
834 + }
835 +
836 + total := b.Finish()
837 +
838 + view, err := DecodeCgroupsResponse(buf[:total])
839 + if err != nil {
840 + t.Fatalf("Decode error: %v", err)
841 + }
842 + if view.ItemCount != 3 {
843 + t.Fatalf("ItemCount = %d, want 3", view.ItemCount)
844 + }
845 + if view.SystemdEnabled != 1 {
846 + t.Fatalf("SystemdEnabled = %d, want 1", view.SystemdEnabled)
847 + }
848 + if view.Generation != 999 {
849 + t.Fatalf("Generation = %d, want 999", view.Generation)
850 + }
851 +
852 + // Item 0
853 + item, err := view.Item(0)
854 + if err != nil {
855 + t.Fatal(err)
856 + }
857 + if item.Hash != 100 {
858 + t.Errorf("item 0 Hash = %d, want 100", item.Hash)
859 + }
860 + if item.Options != 0 {
861 + t.Errorf("item 0 Options = %d, want 0", item.Options)
862 + }
863 + if item.Enabled != 1 {
864 + t.Errorf("item 0 Enabled = %d, want 1", item.Enabled)
865 + }
866 + if item.Name.String() != "init.scope" {
867 + t.Errorf("item 0 Name = %q", item.Name.String())
868 + }
869 + if item.Path.String() != "/sys/fs/cgroup/init.scope" {
870 + t.Errorf("item 0 Path = %q", item.Path.String())
871 + }
872 +
873 + // Item 1
874 + item, err = view.Item(1)
875 + if err != nil {
876 + t.Fatal(err)
877 + }
878 + if item.Hash != 200 {
879 + t.Errorf("item 1 Hash = %d, want 200", item.Hash)
880 + }
881 + if item.Options != 0x02 {
882 + t.Errorf("item 1 Options = %d, want 2", item.Options)
883 + }
884 + if item.Enabled != 0 {
885 + t.Errorf("item 1 Enabled = %d, want 0", item.Enabled)
886 + }
887 + if item.Name.String() != "system.slice/docker-abc.scope" {
888 + t.Errorf("item 1 Name = %q", item.Name.String())
889 + }
890 +
891 + // Item 2 (empty strings)
892 + item, err = view.Item(2)
893 + if err != nil {
894 + t.Fatal(err)
895 + }
896 + if item.Hash != 300 {
897 + t.Errorf("item 2 Hash = %d, want 300", item.Hash)
898 + }
899 + if item.Name.Len() != 0 {
900 + t.Errorf("item 2 Name.Len() = %d, want 0", item.Name.Len())
901 + }
902 + if item.Path.Len() != 0 {
903 + t.Errorf("item 2 Path.Len() = %d, want 0", item.Path.Len())
904 + }
905 +}
906 +
907 +func TestCgroupsResponseCompaction(t *testing.T) {
908 + // Reserve space for 100 items, add only 2.
909 + buf := make([]byte, 8192)
910 + b := NewCgroupsBuilder(buf, 100, 0, 7)
911 +
912 + if err := b.Add(1, 0, 1, []byte("a"), []byte("b")); err != nil {
913 + t.Fatal(err)
914 + }
915 + if err := b.Add(2, 0, 1, []byte("c"), []byte("d")); err != nil {
916 + t.Fatal(err)
917 + }
918 +
919 + total := b.Finish()
920 +
921 + view, err := DecodeCgroupsResponse(buf[:total])
922 + if err != nil {
923 + t.Fatalf("Decode error: %v", err)
924 + }
925 + if view.ItemCount != 2 {
926 + t.Fatalf("ItemCount = %d, want 2", view.ItemCount)
927 + }
928 +
929 + item, err := view.Item(0)
930 + if err != nil {
931 + t.Fatal(err)
932 + }
933 + if item.Name.String() != "a" {
934 + t.Fatalf("item 0 Name = %q, want %q", item.Name.String(), "a")
935 + }
936 +
937 + item, err = view.Item(1)
938 + if err != nil {
939 + t.Fatal(err)
940 + }
941 + if item.Name.String() != "c" {
942 + t.Fatalf("item 1 Name = %q, want %q", item.Name.String(), "c")
943 + }
944 +}
945 +
946 +func TestCgroupsResponseItemOutOfBounds(t *testing.T) {
947 + buf := make([]byte, 8192)
948 + b := NewCgroupsBuilder(buf, 1, 0, 1)
949 + if err := b.Add(1, 0, 1, []byte("x"), []byte("y")); err != nil {
950 + t.Fatal(err)
951 + }
952 + total := b.Finish()
953 +
954 + view, err := DecodeCgroupsResponse(buf[:total])
955 + if err != nil {
956 + t.Fatal(err)
957 + }
958 +
959 + _, err = view.Item(1) // index out of bounds
960 + if err != ErrOutOfBounds {
961 + t.Fatalf("got %v, want ErrOutOfBounds", err)
962 + }
963 +}
964 +
965 +func TestCgroupsResponseDecodeTruncated(t *testing.T) {
966 + buf := make([]byte, 23) // less than 24-byte header
967 + _, err := DecodeCgroupsResponse(buf)
968 + if err != ErrTruncated {
969 + t.Fatalf("got %v, want ErrTruncated", err)
970 + }
971 +}
972 +
973 +func TestCgroupsResponseDecodeBadLayout(t *testing.T) {
974 + buf := make([]byte, 24)
975 + ne.PutUint16(buf[0:2], 99) // bad layout_version
976 + _, err := DecodeCgroupsResponse(buf)
977 + if err != ErrBadLayout {
978 + t.Fatalf("got %v, want ErrBadLayout", err)
979 + }
980 +}
981 +
982 +func TestCgroupsResponseDecodeDirectoryTruncated(t *testing.T) {
983 + buf := make([]byte, 28) // 24 header + 4 bytes, not enough for 1 dir entry (8)
984 + ne.PutUint16(buf[0:2], 1) // layout_version
985 + ne.PutUint16(buf[2:4], 0) // flags
986 + ne.PutUint32(buf[4:8], 1) // item_count = 1
987 + ne.PutUint32(buf[8:12], 0) // systemd_enabled
988 + ne.PutUint32(buf[12:16], 0) // reserved
989 + ne.PutUint64(buf[16:24], 0) // generation
990 +
991 + _, err := DecodeCgroupsResponse(buf)
992 + if err != ErrTruncated {
993 + t.Fatalf("got %v, want ErrTruncated", err)
994 + }
995 +}
996 +
997 +func TestCgroupsResponseDecodeItemBadAlignment(t *testing.T) {
998 + // Create a payload with a directory entry that has a non-aligned offset.
999 + buf := make([]byte, 128)
1000 + ne.PutUint16(buf[0:2], 1) // layout_version
1001 + ne.PutUint16(buf[2:4], 0) // flags
1002 + ne.PutUint32(buf[4:8], 1) // item_count = 1
1003 + ne.PutUint32(buf[8:12], 0) // systemd_enabled
1004 + ne.PutUint32(buf[12:16], 0) // reserved
1005 + ne.PutUint64(buf[16:24], 0) // generation
1006 +
1007 + // Directory entry at offset 24: offset=3 (not aligned), length=32
1008 + ne.PutUint32(buf[24:28], 3) // bad alignment
1009 + ne.PutUint32(buf[28:32], 32)
1010 +
1011 + _, err := DecodeCgroupsResponse(buf)
1012 + if err != ErrBadAlignment {
1013 + t.Fatalf("got %v, want ErrBadAlignment", err)
1014 + }
1015 +}
1016 +
1017 +func TestCgroupsResponseDecodeItemOutOfBounds(t *testing.T) {
1018 + buf := make([]byte, 64)
1019 + ne.PutUint16(buf[0:2], 1)
1020 + ne.PutUint16(buf[2:4], 0)
1021 + ne.PutUint32(buf[4:8], 1) // item_count = 1
1022 + ne.PutUint32(buf[8:12], 0)
1023 + ne.PutUint32(buf[12:16], 0)
1024 + ne.PutUint64(buf[16:24], 0)
1025 +
1026 + // Dir entry: offset=0, length=1000 (exceeds buffer)
1027 + ne.PutUint32(buf[24:28], 0)
1028 + ne.PutUint32(buf[28:32], 1000)
1029 +
1030 + _, err := DecodeCgroupsResponse(buf)
1031 + if err != ErrOutOfBounds {
1032 + t.Fatalf("got %v, want ErrOutOfBounds", err)
1033 + }
1034 +}
1035 +
1036 +func TestCgroupsResponseDecodeItemTooSmall(t *testing.T) {
1037 + // Item is present but smaller than 32-byte item header.
1038 + buf := make([]byte, 64)
1039 + ne.PutUint16(buf[0:2], 1)
1040 + ne.PutUint16(buf[2:4], 0)
1041 + ne.PutUint32(buf[4:8], 1) // item_count = 1
1042 + ne.PutUint32(buf[8:12], 0)
1043 + ne.PutUint32(buf[12:16], 0)
1044 + ne.PutUint64(buf[16:24], 0)
1045 +
1046 + // Dir entry: offset=0, length=16 (< 32 item header)
1047 + ne.PutUint32(buf[24:28], 0)
1048 + ne.PutUint32(buf[28:32], 16)
1049 +
1050 + _, err := DecodeCgroupsResponse(buf)
1051 + if err != ErrTruncated {
1052 + t.Fatalf("got %v, want ErrTruncated", err)
1053 + }
1054 +}
1055 +
1056 +func TestCgroupsResponseItemBadLayout(t *testing.T) {
1057 + // Build a valid response, then corrupt the item's layout_version.
1058 + buf := make([]byte, 8192)
1059 + b := NewCgroupsBuilder(buf, 1, 0, 1)
1060 + if err := b.Add(1, 0, 1, []byte("x"), []byte("y")); err != nil {
1061 + t.Fatal(err)
1062 + }
1063 + total := b.Finish()
1064 +
1065 + // Find the item start and corrupt layout_version.
1066 + dirBase := cgroupsRespHdr
1067 + itemOff := int(ne.Uint32(buf[dirBase : dirBase+4]))
1068 + packedStart := cgroupsRespHdr + 1*cgroupsDirEntry
1069 + itemStart := packedStart + itemOff
1070 + ne.PutUint16(buf[itemStart:itemStart+2], 99) // corrupt layout_version
1071 +
1072 + view, err := DecodeCgroupsResponse(buf[:total])
1073 + if err != nil {
1074 + t.Fatal(err)
1075 + }
1076 +
1077 + _, err = view.Item(0)
1078 + if err != ErrBadLayout {
1079 + t.Fatalf("got %v, want ErrBadLayout", err)
1080 + }
1081 +}
1082 +
1083 +func TestCgroupsResponseItemMissingNul(t *testing.T) {
1084 + // Build a valid response, then overwrite the name's NUL terminator.
1085 + buf := make([]byte, 8192)
1086 + b := NewCgroupsBuilder(buf, 1, 0, 1)
1087 + if err := b.Add(1, 0, 1, []byte("test"), []byte("path")); err != nil {
1088 + t.Fatal(err)
1089 + }
1090 + total := b.Finish()
1091 +
1092 + // Find item and overwrite the NUL after "test".
1093 + dirBase := cgroupsRespHdr
1094 + itemOff := int(ne.Uint32(buf[dirBase : dirBase+4]))
1095 + packedStart := cgroupsRespHdr + 1*cgroupsDirEntry
1096 + itemStart := packedStart + itemOff
1097 + // name is at offset 32 within the item, length 4, NUL at 36.
1098 + buf[itemStart+32+4] = 'X' // overwrite NUL
1099 +
1100 + view, err := DecodeCgroupsResponse(buf[:total])
1101 + if err != nil {
1102 + t.Fatal(err)
1103 + }
1104 + _, err = view.Item(0)
1105 + if err != ErrMissingNul {
1106 + t.Fatalf("got %v, want ErrMissingNul", err)
1107 + }
1108 +}
1109 +
1110 +func TestCgroupsResponseItemNameOutOfBounds(t *testing.T) {
1111 + // Build valid, then corrupt name_offset to point beyond item bounds.
1112 + buf := make([]byte, 8192)
1113 + b := NewCgroupsBuilder(buf, 1, 0, 1)
1114 + if err := b.Add(1, 0, 1, []byte("x"), []byte("y")); err != nil {
1115 + t.Fatal(err)
1116 + }
1117 + total := b.Finish()
1118 +
1119 + dirBase := cgroupsRespHdr
1120 + itemOff := int(ne.Uint32(buf[dirBase : dirBase+4]))
1121 + packedStart := cgroupsRespHdr + 1*cgroupsDirEntry
1122 + itemStart := packedStart + itemOff
1123 + // Corrupt name_offset to a huge value.
1124 + ne.PutUint32(buf[itemStart+16:itemStart+20], 9999)
1125 +
1126 + view, err := DecodeCgroupsResponse(buf[:total])
1127 + if err != nil {
1128 + t.Fatal(err)
1129 + }
1130 + _, err = view.Item(0)
1131 + if err != ErrOutOfBounds {
1132 + t.Fatalf("got %v, want ErrOutOfBounds", err)
1133 + }
1134 +}
1135 +
1136 +func TestCgroupsResponseItemNameOffsetBelowHeader(t *testing.T) {
1137 + // name_offset < 32 (item header size).
1138 + buf := make([]byte, 8192)
1139 + b := NewCgroupsBuilder(buf, 1, 0, 1)
1140 + if err := b.Add(1, 0, 1, []byte("x"), []byte("y")); err != nil {
1141 + t.Fatal(err)
1142 + }
1143 + total := b.Finish()
1144 +
1145 + dirBase := cgroupsRespHdr
1146 + itemOff := int(ne.Uint32(buf[dirBase : dirBase+4]))
1147 + packedStart := cgroupsRespHdr + 1*cgroupsDirEntry
1148 + itemStart := packedStart + itemOff
1149 + ne.PutUint32(buf[itemStart+16:itemStart+20], 4) // below 32
1150 +
1151 + view, err := DecodeCgroupsResponse(buf[:total])
1152 + if err != nil {
1153 + t.Fatal(err)
1154 + }
1155 + _, err = view.Item(0)
1156 + if err != ErrOutOfBounds {
1157 + t.Fatalf("got %v, want ErrOutOfBounds", err)
1158 + }
1159 +}
1160 +
1161 +func TestCgroupsBuilderOverflowMaxItems(t *testing.T) {
1162 + buf := make([]byte, 8192)
1163 + b := NewCgroupsBuilder(buf, 1, 0, 1)
1164 + if err := b.Add(1, 0, 1, []byte("a"), []byte("b")); err != nil {
1165 + t.Fatal(err)
1166 + }
1167 + if err := b.Add(2, 0, 1, []byte("c"), []byte("d")); err != ErrOverflow {
1168 + t.Fatalf("got %v, want ErrOverflow", err)
1169 + }
1170 +}
1171 +
1172 +func TestCgroupsBuilderOverflowBuffer(t *testing.T) {
1173 + buf := make([]byte, 40) // too small for header + dir + item
1174 + b := NewCgroupsBuilder(buf, 1, 0, 1)
1175 + err := b.Add(1, 0, 1, []byte("long-name-that-wont-fit"), []byte("long-path-too"))
1176 + if err != ErrOverflow {
1177 + t.Fatalf("got %v, want ErrOverflow", err)
1178 + }
1179 +}
1180 +
1181 +func TestCgroupsResponseEmptyStrings(t *testing.T) {
1182 + buf := make([]byte, 8192)
1183 + b := NewCgroupsBuilder(buf, 1, 0, 1)
1184 + if err := b.Add(42, 0, 1, []byte(""), []byte("")); err != nil {
1185 + t.Fatal(err)
1186 + }
1187 + total := b.Finish()
1188 +
1189 + view, err := DecodeCgroupsResponse(buf[:total])
1190 + if err != nil {
1191 + t.Fatal(err)
1192 + }
1193 +
1194 + item, err := view.Item(0)
1195 + if err != nil {
1196 + t.Fatal(err)
1197 + }
1198 + if item.Name.Len() != 0 {
1199 + t.Fatalf("Name.Len() = %d, want 0", item.Name.Len())
1200 + }
1201 + if item.Path.Len() != 0 {
1202 + t.Fatalf("Path.Len() = %d, want 0", item.Path.Len())
1203 + }
1204 + if item.Name.String() != "" {
1205 + t.Fatalf("Name.String() = %q, want empty", item.Name.String())
1206 + }
1207 + if item.Path.String() != "" {
1208 + t.Fatalf("Path.String() = %q, want empty", item.Path.String())
1209 + }
1210 +}
1211 +
1212 +// ---------------------------------------------------------------------------
1213 +// Interop reference: test values matching C and Rust interop binaries
1214 +// ---------------------------------------------------------------------------
1215 +
1216 +func TestInteropHeaderValues(t *testing.T) {
1217 + h := Header{
1218 + Magic: MagicMsg,
1219 + Version: Version,
1220 + HeaderLen: HeaderLen,
1221 + Kind: KindRequest,
1222 + Flags: FlagBatch,
1223 + Code: MethodCgroupsSnapshot,
1224 + TransportStatus: StatusOK,
1225 + PayloadLen: 12345,
1226 + ItemCount: 42,
1227 + MessageID: 0xDEADBEEFCAFEBABE,
1228 + }
1229 +
1230 + var buf [32]byte
1231 + h.Encode(buf[:])
1232 +
1233 + // Decode and verify all fields match.
1234 + out, err := DecodeHeader(buf[:])
1235 + if err != nil {
1236 + t.Fatal(err)
1237 + }
1238 + if out.Magic != MagicMsg {
1239 + t.Errorf("magic = %x", out.Magic)
1240 + }
1241 + if out.PayloadLen != 12345 {
1242 + t.Errorf("payload_len = %d", out.PayloadLen)
1243 + }
1244 + if out.ItemCount != 42 {
1245 + t.Errorf("item_count = %d", out.ItemCount)
1246 + }
1247 + if out.MessageID != 0xDEADBEEFCAFEBABE {
1248 + t.Errorf("message_id = %x", out.MessageID)
1249 + }
1250 +}
1251 +
1252 +func TestInteropChunkValues(t *testing.T) {
1253 + c := ChunkHeader{
1254 + Magic: MagicChunk,
1255 + Version: Version,
1256 + Flags: 0,
1257 + MessageID: 0x1234567890ABCDEF,
1258 + TotalMessageLen: 100000,
1259 + ChunkIndex: 3,
1260 + ChunkCount: 10,
1261 + ChunkPayloadLen: 8192,
1262 + }
1263 +
1264 + var buf [32]byte
1265 + c.Encode(buf[:])
1266 +
1267 + out, err := DecodeChunkHeader(buf[:])
1268 + if err != nil {
1269 + t.Fatal(err)
1270 + }
1271 + if out.MessageID != 0x1234567890ABCDEF {
1272 + t.Errorf("message_id = %x", out.MessageID)
1273 + }
1274 + if out.TotalMessageLen != 100000 {
1275 + t.Errorf("total_message_len = %d", out.TotalMessageLen)
1276 + }
1277 + if out.ChunkIndex != 3 {
1278 + t.Errorf("chunk_index = %d", out.ChunkIndex)
1279 + }
1280 + if out.ChunkCount != 10 {
1281 + t.Errorf("chunk_count = %d", out.ChunkCount)
1282 + }
1283 + if out.ChunkPayloadLen != 8192 {
1284 + t.Errorf("chunk_payload_len = %d", out.ChunkPayloadLen)
1285 + }
1286 +}
1287 +
1288 +func TestInteropHelloValues(t *testing.T) {
1289 + h := Hello{
1290 + LayoutVersion: 1,
1291 + Flags: 0,
1292 + SupportedProfiles: ProfileBaseline | ProfileSHMFutex,
1293 + PreferredProfiles: ProfileSHMFutex,
1294 + MaxRequestPayloadBytes: 4096,
1295 + MaxRequestBatchItems: 100,
1296 + MaxResponsePayloadBytes: 1048576,
1297 + MaxResponseBatchItems: 1,
1298 + AuthToken: 0xAABBCCDDEEFF0011,
1299 + PacketSize: 65536,
1300 + }
1301 +
1302 + var buf [44]byte
1303 + h.Encode(buf[:])
1304 +
1305 + out, err := DecodeHello(buf[:])
1306 + if err != nil {
1307 + t.Fatal(err)
1308 + }
1309 + if out.SupportedProfiles != 0x05 {
1310 + t.Errorf("supported = %x", out.SupportedProfiles)
1311 + }
1312 + if out.AuthToken != 0xAABBCCDDEEFF0011 {
1313 + t.Errorf("auth_token = %x", out.AuthToken)
1314 + }
1315 + if out.PacketSize != 65536 {
1316 + t.Errorf("packet_size = %d", out.PacketSize)
1317 + }
1318 +}
1319 +
1320 +func TestInteropHelloAckValues(t *testing.T) {
1321 + h := HelloAck{
1322 + LayoutVersion: 1,
1323 + Flags: 0,
1324 + ServerSupportedProfiles: 0x07,
1325 + IntersectionProfiles: 0x05,
1326 + SelectedProfile: ProfileSHMFutex,
1327 + AgreedMaxRequestPayloadBytes: 2048,
1328 + AgreedMaxRequestBatchItems: 50,
1329 + AgreedMaxResponsePayloadBytes: 65536,
1330 + AgreedMaxResponseBatchItems: 1,
1331 + AgreedPacketSize: 32768,
1332 + }
1333 +
1334 + var buf [48]byte
1335 + h.Encode(buf[:])
1336 +
1337 + out, err := DecodeHelloAck(buf[:])
1338 + if err != nil {
1339 + t.Fatal(err)
1340 + }
1341 + if out.ServerSupportedProfiles != 0x07 {
1342 + t.Errorf("server_supported = %x", out.ServerSupportedProfiles)
1343 + }
1344 + if out.AgreedPacketSize != 32768 {
1345 + t.Errorf("agreed_packet_size = %d", out.AgreedPacketSize)
1346 + }
1347 +}
1348 +
1349 +func TestInteropCgroupsResponseValues(t *testing.T) {
1350 + buf := make([]byte, 8192)
1351 + b := NewCgroupsBuilder(buf, 3, 1, 999)
1352 +
1353 + if err := b.Add(100, 0, 1,
1354 + []byte("init.scope"),
1355 + []byte("/sys/fs/cgroup/init.scope")); err != nil {
1356 + t.Fatal(err)
1357 + }
1358 + if err := b.Add(200, 0x02, 0,
1359 + []byte("system.slice/docker-abc.scope"),
1360 + []byte("/sys/fs/cgroup/system.slice/docker-abc.scope")); err != nil {
1361 + t.Fatal(err)
1362 + }
1363 + if err := b.Add(300, 0, 1, []byte(""), []byte("")); err != nil {
1364 + t.Fatal(err)
1365 + }
1366 +
1367 + total := b.Finish()
1368 +
1369 + view, err := DecodeCgroupsResponse(buf[:total])
1370 + if err != nil {
1371 + t.Fatal(err)
1372 + }
1373 + if view.ItemCount != 3 {
1374 + t.Fatalf("ItemCount = %d", view.ItemCount)
1375 + }
1376 + if view.SystemdEnabled != 1 {
1377 + t.Fatalf("SystemdEnabled = %d", view.SystemdEnabled)
1378 + }
1379 + if view.Generation != 999 {
1380 + t.Fatalf("Generation = %d", view.Generation)
1381 + }
1382 +
1383 + item0, _ := view.Item(0)
1384 + if item0.Hash != 100 || item0.Name.String() != "init.scope" {
1385 + t.Errorf("item 0: hash=%d name=%q", item0.Hash, item0.Name.String())
1386 + }
1387 +
1388 + item1, _ := view.Item(1)
1389 + if item1.Hash != 200 || item1.Options != 0x02 || item1.Enabled != 0 {
1390 + t.Errorf("item 1: hash=%d options=%d enabled=%d",
1391 + item1.Hash, item1.Options, item1.Enabled)
1392 + }
1393 + if item1.Name.String() != "system.slice/docker-abc.scope" {
1394 + t.Errorf("item 1 name: %q", item1.Name.String())
1395 + }
1396 +
1397 + item2, _ := view.Item(2)
1398 + if item2.Hash != 300 || item2.Name.Len() != 0 || item2.Path.Len() != 0 {
1399 + t.Errorf("item 2: hash=%d namelen=%d pathlen=%d",
1400 + item2.Hash, item2.Name.Len(), item2.Path.Len())
1401 + }
1402 +}
1403 +
1404 +func TestInteropCgroupsResponseEmptyValues(t *testing.T) {
1405 + buf := make([]byte, 8192)
1406 + b := NewCgroupsBuilder(buf, 0, 0, 42)
1407 + total := b.Finish()
1408 +
1409 + view, err := DecodeCgroupsResponse(buf[:total])
1410 + if err != nil {
1411 + t.Fatal(err)
1412 + }
1413 + if view.ItemCount != 0 {
1414 + t.Errorf("ItemCount = %d", view.ItemCount)
1415 + }
1416 + if view.SystemdEnabled != 0 {
1417 + t.Errorf("SystemdEnabled = %d", view.SystemdEnabled)
1418 + }
1419 + if view.Generation != 42 {
1420 + t.Errorf("Generation = %d", view.Generation)
1421 + }
1422 +}
1423 +
1424 +// ---------------------------------------------------------------------------
1425 +// Large snapshot test
1426 +// ---------------------------------------------------------------------------
1427 +
1428 +func TestCgroupsResponseLargeSnapshot(t *testing.T) {
1429 + const n = 200
1430 + buf := make([]byte, 1024*1024)
1431 + b := NewCgroupsBuilder(buf, n, 1, 12345)
1432 +
1433 + for i := 0; i < n; i++ {
1434 + name := []byte("cgroup-" + string(rune('A'+i%26)))
1435 + path := []byte("/sys/fs/cgroup/system.slice/cgroup-" + string(rune('A'+i%26)))
1436 + if err := b.Add(uint32(i), uint32(i%4), uint32(i%2), name, path); err != nil {
1437 + t.Fatalf("item %d: %v", i, err)
1438 + }
1439 + }
1440 +
1441 + total := b.Finish()
1442 +
1443 + view, err := DecodeCgroupsResponse(buf[:total])
1444 + if err != nil {
1445 + t.Fatal(err)
1446 + }
1447 + if view.ItemCount != n {
1448 + t.Fatalf("ItemCount = %d, want %d", view.ItemCount, n)
1449 + }
1450 +
1451 + for i := uint32(0); i < n; i++ {
1452 + item, err := view.Item(i)
1453 + if err != nil {
1454 + t.Fatalf("item %d: %v", i, err)
1455 + }
1456 + if item.Hash != i {
1457 + t.Errorf("item %d: hash = %d", i, item.Hash)
1458 + }
1459 + }
1460 +}
1461 +
1462 +// ---------------------------------------------------------------------------
1463 +// Error string tests (ensure errors have useful messages)
1464 +// ---------------------------------------------------------------------------
1465 +
1466 +func TestErrorStrings(t *testing.T) {
1467 + errs := []error{
1468 + ErrTruncated, ErrBadMagic, ErrBadVersion, ErrBadHeaderLen,
1469 + ErrBadKind, ErrBadLayout, ErrOutOfBounds, ErrMissingNul,
1470 + ErrBadAlignment, ErrBadItemCount, ErrOverflow,
1471 + }
1472 + for _, e := range errs {
1473 + if e.Error() == "" {
1474 + t.Errorf("error has empty string: %v", e)
1475 + }
1476 + }
1477 +}
1478 +
1479 +// ---------------------------------------------------------------------------
1480 +// Decode from zero/garbage bytes (robustness)
1481 +// ---------------------------------------------------------------------------
1482 +
1483 +func TestDecodeZeroBytes(t *testing.T) {
1484 + buf := make([]byte, 128)
1485 +
1486 + _, err := DecodeHeader(buf)
1487 + if err == nil {
1488 + t.Error("expected error for zero header")
1489 + }
1490 +
1491 + _, err = DecodeChunkHeader(buf)
1492 + if err == nil {
1493 + t.Error("expected error for zero chunk header")
1494 + }
1495 +
1496 + _, err = DecodeHello(buf)
1497 + if err == nil {
1498 + t.Error("expected error for zero hello")
1499 + }
1500 +
1501 + _, err = DecodeHelloAck(buf)
1502 + if err == nil {
1503 + t.Error("expected error for zero hello_ack")
1504 + }
1505 +
1506 + _, err = DecodeCgroupsRequest(buf)
1507 + if err == nil {
1508 + t.Error("expected error for zero cgroups_req")
1509 + }
1510 +
1511 + _, err = DecodeCgroupsResponse(buf)
1512 + if err == nil {
1513 + t.Error("expected error for zero cgroups_resp")
1514 + }
1515 +}
1516 +
1517 +func TestDecodeGarbage(t *testing.T) {
1518 + buf := make([]byte, 128)
1519 + for i := range buf {
1520 + buf[i] = 0xFF
1521 + }
1522 +
1523 + _, err := DecodeHeader(buf)
1524 + if err == nil {
1525 + t.Error("expected error for garbage header")
1526 + }
1527 +
1528 + _, err = DecodeChunkHeader(buf)
1529 + if err == nil {
1530 + t.Error("expected error for garbage chunk header")
1531 + }
1532 +
1533 + _, err = DecodeHello(buf)
1534 + if err == nil {
1535 + t.Error("expected error for garbage hello")
1536 + }
1537 +
1538 + _, err = DecodeHelloAck(buf)
1539 + if err == nil {
1540 + t.Error("expected error for garbage hello_ack")
1541 + }
1542 +
1543 + _, err = DecodeCgroupsRequest(buf)
1544 + if err == nil {
1545 + t.Error("expected error for garbage cgroups_req")
1546 + }
1547 +
1548 + _, err = DecodeCgroupsResponse(buf)
1549 + if err == nil {
1550 + t.Error("expected error for garbage cgroups_resp")
1551 + }
1552 +}
1553 +
1554 +func TestDecodeEmptyBuf(t *testing.T) {
1555 + empty := []byte{}
1556 + if _, err := DecodeHeader(empty); err != ErrTruncated {
1557 + t.Errorf("header: got %v", err)
1558 + }
1559 + if _, err := DecodeChunkHeader(empty); err != ErrTruncated {
1560 + t.Errorf("chunk: got %v", err)
1561 + }
1562 + if _, err := DecodeHello(empty); err != ErrTruncated {
1563 + t.Errorf("hello: got %v", err)
1564 + }
1565 + if _, err := DecodeHelloAck(empty); err != ErrTruncated {
1566 + t.Errorf("hello_ack: got %v", err)
1567 + }
1568 + if _, err := DecodeCgroupsRequest(empty); err != ErrTruncated {
1569 + t.Errorf("cgroups_req: got %v", err)
1570 + }
1571 + if _, err := DecodeCgroupsResponse(empty); err != ErrTruncated {
1572 + t.Errorf("cgroups_resp: got %v", err)
1573 + }
1574 +}
1575 +
1576 +// ---------------------------------------------------------------------------
1577 +// Path-level string validation in cgroups item
1578 +// ---------------------------------------------------------------------------
1579 +
1580 +func TestCgroupsResponseItemPathMissingNul(t *testing.T) {
1581 + buf := make([]byte, 8192)
1582 + b := NewCgroupsBuilder(buf, 1, 0, 1)
1583 + if err := b.Add(1, 0, 1, []byte("ok"), []byte("bad")); err != nil {
1584 + t.Fatal(err)
1585 + }
1586 + total := b.Finish()
1587 +
1588 + // Overwrite the NUL after "bad" (the path string).
1589 + dirBase := cgroupsRespHdr
1590 + itemOff := int(ne.Uint32(buf[dirBase : dirBase+4]))
1591 + packedStart := cgroupsRespHdr + 1*cgroupsDirEntry
1592 + itemStart := packedStart + itemOff
1593 +
1594 + // path_offset is at item+24, path_len at item+28.
1595 + pathOff := int(ne.Uint32(buf[itemStart+24 : itemStart+28]))
1596 + pathLen := int(ne.Uint32(buf[itemStart+28 : itemStart+32]))
1597 + buf[itemStart+pathOff+pathLen] = 'X' // corrupt NUL
1598 +
1599 + view, err := DecodeCgroupsResponse(buf[:total])
1600 + if err != nil {
1601 + t.Fatal(err)
1602 + }
1603 + _, err = view.Item(0)
1604 + if err != ErrMissingNul {
1605 + t.Fatalf("got %v, want ErrMissingNul", err)
1606 + }
1607 +}
1608 +
1609 +func TestCgroupsResponseItemPathOutOfBounds(t *testing.T) {
1610 + buf := make([]byte, 8192)
1611 + b := NewCgroupsBuilder(buf, 1, 0, 1)
1612 + if err := b.Add(1, 0, 1, []byte("x"), []byte("y")); err != nil {
1613 + t.Fatal(err)
1614 + }
1615 + total := b.Finish()
1616 +
1617 + dirBase := cgroupsRespHdr
1618 + itemOff := int(ne.Uint32(buf[dirBase : dirBase+4]))
1619 + packedStart := cgroupsRespHdr + 1*cgroupsDirEntry
1620 + itemStart := packedStart + itemOff
1621 + ne.PutUint32(buf[itemStart+24:itemStart+28], 9999) // corrupt path_offset
1622 +
1623 + view, err := DecodeCgroupsResponse(buf[:total])
1624 + if err != nil {
1625 + t.Fatal(err)
1626 + }
1627 + _, err = view.Item(0)
1628 + if err != ErrOutOfBounds {
1629 + t.Fatalf("got %v, want ErrOutOfBounds", err)
1630 + }
1631 +}
1632 +
1633 +func TestCgroupsResponseItemPathOffsetBelowHeader(t *testing.T) {
1634 + buf := make([]byte, 8192)
1635 + b := NewCgroupsBuilder(buf, 1, 0, 1)
1636 + if err := b.Add(1, 0, 1, []byte("x"), []byte("y")); err != nil {
1637 + t.Fatal(err)
1638 + }
1639 + total := b.Finish()
1640 +
1641 + dirBase := cgroupsRespHdr
1642 + itemOff := int(ne.Uint32(buf[dirBase : dirBase+4]))
1643 + packedStart := cgroupsRespHdr + 1*cgroupsDirEntry
1644 + itemStart := packedStart + itemOff
1645 + ne.PutUint32(buf[itemStart+24:itemStart+28], 4) // below 32
1646 +
1647 + view, err := DecodeCgroupsResponse(buf[:total])
1648 + if err != nil {
1649 + t.Fatal(err)
1650 + }
1651 + _, err = view.Item(0)
1652 + if err != ErrOutOfBounds {
1653 + t.Fatalf("got %v, want ErrOutOfBounds", err)
1654 + }
1655 +}
1656 +
1657 +func TestCgroupsResponseItemOverlapRejected(t *testing.T) {
1658 + buf := make([]byte, 8192)
1659 + b := NewCgroupsBuilder(buf, 1, 0, 1)
1660 + if err := b.Add(1, 0, 1, []byte("hello"), []byte("/path")); err != nil {
1661 + t.Fatal(err)
1662 + }
1663 + total := b.Finish()
1664 +
1665 + dirBase := cgroupsRespHdr
1666 + itemOff := int(ne.Uint32(buf[dirBase : dirBase+4]))
1667 + packedStart := cgroupsRespHdr + 1*cgroupsDirEntry
1668 + itemStart := packedStart + itemOff
1669 +
1670 + // name region is [32..38) (name_off=32, name_len=5, +1 for NUL)
1671 + // Set path_off=34 (inside name), path_len=1
1672 + ne.PutUint32(buf[itemStart+24:itemStart+28], 34)
1673 + ne.PutUint32(buf[itemStart+28:itemStart+32], 1)
1674 + // Ensure NUL at item[35]
1675 + buf[itemStart+35] = 0
1676 +
1677 + view, err := DecodeCgroupsResponse(buf[:total])
1678 + if err != nil {
1679 + t.Fatal(err)
1680 + }
1681 + _, err = view.Item(0)
1682 + if err != ErrBadLayout {
1683 + t.Fatalf("got %v, want ErrBadLayout (overlapping fields)", err)
1684 + }
1685 +}
src/go/pkg/netipc/protocol/fuzz_test.go new
+273
@@ -0,0 +1,273 @@
1 +package protocol
2 +
3 +import (
4 + "testing"
5 +)
6 +
7 +// ---------------------------------------------------------------------------
8 +// Fuzz targets for all decode paths.
9 +//
10 +// Each target feeds arbitrary bytes to a decode function and, if the decode
11 +// succeeds, exercises the result. The invariant: no input may cause a panic.
12 +// ---------------------------------------------------------------------------
13 +
14 +func FuzzDecodeHeader(f *testing.F) {
15 + // Seed with a valid header.
16 + var seed [HeaderSize]byte
17 + h := Header{
18 + Magic: MagicMsg, Version: Version, HeaderLen: HeaderLen,
19 + Kind: KindRequest, Code: MethodIncrement, PayloadLen: 0,
20 + ItemCount: 1, MessageID: 1,
21 + }
22 + h.Encode(seed[:])
23 + f.Add(seed[:])
24 +
25 + // Seed with truncated and zeroed inputs.
26 + f.Add([]byte{})
27 + f.Add(make([]byte, 31))
28 + f.Add(make([]byte, 32))
29 +
30 + f.Fuzz(func(t *testing.T, data []byte) {
31 + hdr, err := DecodeHeader(data)
32 + if err != nil {
33 + return
34 + }
35 + // Exercise the decoded result.
36 + _ = hdr.Magic
37 + _ = hdr.Version
38 + _ = hdr.Kind
39 + _ = hdr.Flags
40 + _ = hdr.Code
41 + _ = hdr.TransportStatus
42 + _ = hdr.PayloadLen
43 + _ = hdr.ItemCount
44 + _ = hdr.MessageID
45 + })
46 +}
47 +
48 +func FuzzDecodeChunkHeader(f *testing.F) {
49 + var seed [HeaderSize]byte
50 + c := ChunkHeader{
51 + Magic: MagicChunk, Version: Version, Flags: 0,
52 + MessageID: 1, TotalMessageLen: 256,
53 + ChunkIndex: 0, ChunkCount: 3, ChunkPayloadLen: 100,
54 + }
55 + c.Encode(seed[:])
56 + f.Add(seed[:])
57 +
58 + f.Add([]byte{})
59 + f.Add(make([]byte, 31))
60 + f.Add(make([]byte, 32))
61 +
62 + f.Fuzz(func(t *testing.T, data []byte) {
63 + chk, err := DecodeChunkHeader(data)
64 + if err != nil {
65 + return
66 + }
67 + _ = chk.Magic
68 + _ = chk.Version
69 + _ = chk.Flags
70 + _ = chk.MessageID
71 + _ = chk.TotalMessageLen
72 + _ = chk.ChunkIndex
73 + _ = chk.ChunkCount
74 + _ = chk.ChunkPayloadLen
75 + })
76 +}
77 +
78 +func FuzzDecodeHello(f *testing.F) {
79 + var seed [64]byte
80 + h := Hello{
81 + LayoutVersion: 1, SupportedProfiles: ProfileBaseline,
82 + PreferredProfiles: ProfileBaseline,
83 + MaxRequestPayloadBytes: 1024, MaxRequestBatchItems: 1,
84 + MaxResponsePayloadBytes: 1024, MaxResponseBatchItems: 1,
85 + AuthToken: 0xABCD, PacketSize: 65536,
86 + }
87 + h.Encode(seed[:])
88 + f.Add(seed[:44])
89 +
90 + f.Add([]byte{})
91 + f.Add(make([]byte, 43))
92 + f.Add(make([]byte, 44))
93 +
94 + f.Fuzz(func(t *testing.T, data []byte) {
95 + hello, err := DecodeHello(data)
96 + if err != nil {
97 + return
98 + }
99 + _ = hello.LayoutVersion
100 + _ = hello.Flags
101 + _ = hello.SupportedProfiles
102 + _ = hello.PreferredProfiles
103 + _ = hello.MaxRequestPayloadBytes
104 + _ = hello.MaxRequestBatchItems
105 + _ = hello.MaxResponsePayloadBytes
106 + _ = hello.MaxResponseBatchItems
107 + _ = hello.AuthToken
108 + _ = hello.PacketSize
109 + })
110 +}
111 +
112 +func FuzzDecodeHelloAck(f *testing.F) {
113 + var seed [64]byte
114 + h := HelloAck{
115 + LayoutVersion: 1, ServerSupportedProfiles: 0x07,
116 + IntersectionProfiles: 0x05, SelectedProfile: ProfileSHMFutex,
117 + AgreedMaxRequestPayloadBytes: 2048, AgreedMaxRequestBatchItems: 50,
118 + AgreedMaxResponsePayloadBytes: 65536, AgreedMaxResponseBatchItems: 1,
119 + AgreedPacketSize: 32768,
120 + }
121 + h.Encode(seed[:])
122 + f.Add(seed[:48])
123 +
124 + f.Add([]byte{})
125 + f.Add(make([]byte, 47))
126 + f.Add(make([]byte, 48))
127 +
128 + f.Fuzz(func(t *testing.T, data []byte) {
129 + ack, err := DecodeHelloAck(data)
130 + if err != nil {
131 + return
132 + }
133 + _ = ack.LayoutVersion
134 + _ = ack.Flags
135 + _ = ack.ServerSupportedProfiles
136 + _ = ack.IntersectionProfiles
137 + _ = ack.SelectedProfile
138 + _ = ack.AgreedMaxRequestPayloadBytes
139 + _ = ack.AgreedMaxRequestBatchItems
140 + _ = ack.AgreedMaxResponsePayloadBytes
141 + _ = ack.AgreedMaxResponseBatchItems
142 + _ = ack.AgreedPacketSize
143 + })
144 +}
145 +
146 +func FuzzDecodeCgroupsRequest(f *testing.F) {
147 + var seed [4]byte
148 + r := CgroupsRequest{LayoutVersion: 1, Flags: 0}
149 + r.Encode(seed[:])
150 + f.Add(seed[:])
151 +
152 + f.Add([]byte{})
153 + f.Add(make([]byte, 3))
154 + f.Add(make([]byte, 4))
155 +
156 + f.Fuzz(func(t *testing.T, data []byte) {
157 + req, err := DecodeCgroupsRequest(data)
158 + if err != nil {
159 + return
160 + }
161 + _ = req.LayoutVersion
162 + _ = req.Flags
163 + })
164 +}
165 +
166 +func FuzzDecodeCgroupsResponse(f *testing.F) {
167 + // Seed: empty snapshot (24 bytes).
168 + var emptyBuf [4096]byte
169 + eb := NewCgroupsBuilder(emptyBuf[:], 0, 1, 42)
170 + emptyTotal := eb.Finish()
171 + f.Add(emptyBuf[:emptyTotal])
172 +
173 + // Seed: single-item snapshot.
174 + var singleBuf [4096]byte
175 + sb := NewCgroupsBuilder(singleBuf[:], 1, 0, 100)
176 + sb.Add(12345, 0x01, 1, []byte("docker-abc123"), []byte("/sys/fs/cgroup/docker/abc123"))
177 + singleTotal := sb.Finish()
178 + f.Add(singleBuf[:singleTotal])
179 +
180 + // Seed: garbage inputs.
181 + f.Add([]byte{})
182 + f.Add(make([]byte, 23))
183 + f.Add(make([]byte, 64))
184 +
185 + f.Fuzz(func(t *testing.T, data []byte) {
186 + view, err := DecodeCgroupsResponse(data)
187 + if err != nil {
188 + return
189 + }
190 + // Exercise all items.
191 + for i := uint32(0); i < view.ItemCount; i++ {
192 + item, ierr := view.Item(i)
193 + if ierr != nil {
194 + continue
195 + }
196 + _ = item.LayoutVersion
197 + _ = item.Flags
198 + _ = item.Hash
199 + _ = item.Options
200 + _ = item.Enabled
201 + _ = item.Name.Bytes()
202 + _ = item.Name.Len()
203 + _ = item.Name.String()
204 + _ = item.Path.Bytes()
205 + _ = item.Path.Len()
206 + _ = item.Path.String()
207 + }
208 + // Out-of-bounds item access should not panic.
209 + _, _ = view.Item(view.ItemCount)
210 + if view.ItemCount > 0 {
211 + _, _ = view.Item(view.ItemCount + 1)
212 + }
213 + })
214 +}
215 +
216 +func FuzzBatchDirDecode(f *testing.F) {
217 + // Seed: valid 2-entry directory.
218 + var seed [16]byte
219 + ne.PutUint32(seed[0:4], 0) // offset=0, aligned
220 + ne.PutUint32(seed[4:8], 10) // length=10
221 + ne.PutUint32(seed[8:12], 16) // offset=16, aligned
222 + ne.PutUint32(seed[12:16], 5) // length=5
223 + f.Add(seed[:], uint32(2), uint32(100))
224 +
225 + // Seed: empty.
226 + f.Add([]byte{}, uint32(0), uint32(0))
227 + f.Add(make([]byte, 8), uint32(1), uint32(50))
228 + f.Add(make([]byte, 4), uint32(1), uint32(50)) // truncated
229 +
230 + f.Fuzz(func(t *testing.T, data []byte, itemCount uint32, packedAreaLen uint32) {
231 + // Clamp item_count to prevent huge allocations that slow the fuzzer.
232 + if itemCount > 1024 {
233 + return
234 + }
235 + entries, err := BatchDirDecode(data, itemCount, packedAreaLen)
236 + if err != nil {
237 + return
238 + }
239 + for _, e := range entries {
240 + _ = e.Offset
241 + _ = e.Length
242 + }
243 + })
244 +}
245 +
246 +func FuzzBatchItemGet(f *testing.F) {
247 + // Seed: valid single-item batch payload built via BatchBuilder.
248 + var batchBuf [256]byte
249 + bb := NewBatchBuilder(batchBuf[:], 2)
250 + bb.Add([]byte{1, 2, 3, 4, 5})
251 + bb.Add([]byte{10, 20, 30})
252 + total, count := bb.Finish()
253 + f.Add(batchBuf[:total], count, uint32(0))
254 + f.Add(batchBuf[:total], count, uint32(1))
255 +
256 + f.Add([]byte{}, uint32(0), uint32(0))
257 + f.Add(make([]byte, 8), uint32(1), uint32(0))
258 +
259 + f.Fuzz(func(t *testing.T, payload []byte, itemCount uint32, index uint32) {
260 + // Clamp to prevent huge allocations from itemCount.
261 + if itemCount > 1024 {
262 + return
263 + }
264 + item, err := BatchItemGet(payload, itemCount, index)
265 + if err != nil {
266 + return
267 + }
268 + _ = len(item)
269 + if len(item) > 0 {
270 + _ = item[0]
271 + }
272 + })
273 +}
src/go/pkg/netipc/protocol/increment.go new
+37
@@ -0,0 +1,37 @@
1 +// INCREMENT codec (method 1) -- 8-byte payload: { u64 value }
2 +
3 +package protocol
4 +
5 +const IncrementPayloadSize = 8
6 +
7 +// IncrementEncode writes a u64 value into buf. Returns 8 on success, 0 if
8 +// buf is too small.
9 +func IncrementEncode(value uint64, buf []byte) int {
10 + if len(buf) < IncrementPayloadSize {
11 + return 0
12 + }
13 + ne.PutUint64(buf[:8], value)
14 + return IncrementPayloadSize
15 +}
16 +
17 +// IncrementDecode reads a u64 value from buf.
18 +func IncrementDecode(buf []byte) (uint64, error) {
19 + if len(buf) < IncrementPayloadSize {
20 + return 0, ErrTruncated
21 + }
22 + return ne.Uint64(buf[:8]), nil
23 +}
24 +
25 +// DispatchIncrement decodes request, calls handler, encodes response.
26 +func DispatchIncrement(req []byte, resp []byte, handler func(uint64) (uint64, bool)) (int, bool) {
27 + value, err := IncrementDecode(req)
28 + if err != nil {
29 + return 0, false
30 + }
31 + result, ok := handler(value)
32 + if !ok {
33 + return 0, false
34 + }
35 + n := IncrementEncode(result, resp)
36 + return n, n > 0
37 +}
src/go/pkg/netipc/protocol/string_reverse.go new
+64
@@ -0,0 +1,64 @@
1 +// STRING_REVERSE codec (method 3) -- variable-length payload:
2 +//
3 +// [0:4] u32 str_offset (from payload start, always 8)
4 +// [4:8] u32 str_length (excluding NUL)
5 +// [8:N+1] string data + NUL
6 +
7 +package protocol
8 +
9 +const StringReverseHdrSize = 8
10 +
11 +// StringReverseView is the decoded result of a STRING_REVERSE payload.
12 +type StringReverseView struct {
13 + Str string
14 + StrLen uint32
15 +}
16 +
17 +// StringReverseEncode writes a STRING_REVERSE payload into buf.
18 +// Returns total bytes written, or 0 if buf is too small.
19 +func StringReverseEncode(s string, buf []byte) int {
20 + total := StringReverseHdrSize + len(s) + 1
21 + if len(buf) < total {
22 + return 0
23 + }
24 + ne.PutUint32(buf[0:4], uint32(StringReverseHdrSize)) // str_offset
25 + ne.PutUint32(buf[4:8], uint32(len(s))) // str_length
26 + if len(s) > 0 {
27 + copy(buf[8:8+len(s)], s)
28 + }
29 + buf[8+len(s)] = 0 // NUL terminator
30 + return total
31 +}
32 +
33 +// StringReverseDecode decodes a STRING_REVERSE payload from buf.
34 +func StringReverseDecode(buf []byte) (StringReverseView, error) {
35 + if len(buf) < StringReverseHdrSize {
36 + return StringReverseView{}, ErrTruncated
37 + }
38 + strOffset := int(ne.Uint32(buf[0:4]))
39 + strLength := int(ne.Uint32(buf[4:8]))
40 + if strOffset+strLength+1 > len(buf) {
41 + return StringReverseView{}, ErrOutOfBounds
42 + }
43 + if buf[strOffset+strLength] != 0 {
44 + return StringReverseView{}, ErrMissingNul
45 + }
46 + return StringReverseView{
47 + Str: string(buf[strOffset : strOffset+strLength]),
48 + StrLen: uint32(strLength),
49 + }, nil
50 +}
51 +
52 +// DispatchStringReverse decodes request, calls handler, encodes response.
53 +func DispatchStringReverse(req []byte, resp []byte, handler func(string) (string, bool)) (int, bool) {
54 + view, err := StringReverseDecode(req)
55 + if err != nil {
56 + return 0, false
57 + }
58 + result, ok := handler(view.Str)
59 + if !ok {
60 + return 0, false
61 + }
62 + n := StringReverseEncode(result, resp)
63 + return n, n > 0
64 +}
src/go/pkg/netipc/service/cgroups/cache.go new
+42
@@ -0,0 +1,42 @@
1 +//go:build unix
2 +
3 +package cgroups
4 +
5 +import (
6 + raw "github.com/netdata/netdata/go/plugins/pkg/netipc/service/raw"
7 +)
8 +
9 +// Cache is the public L3 client-side cgroups snapshot cache.
10 +type Cache struct {
11 + inner *raw.Cache
12 +}
13 +
14 +// NewCache creates a new L3 cache. Does NOT connect.
15 +func NewCache(runDir, serviceName string, config ClientConfig) *Cache {
16 + return &Cache{inner: raw.NewCache(runDir, serviceName, clientConfigToTransport(config))}
17 +}
18 +
19 +// Refresh drives the L2 client and requests a fresh snapshot.
20 +func (c *Cache) Refresh() bool {
21 + return c.inner.Refresh()
22 +}
23 +
24 +// Ready returns true if at least one successful refresh has occurred.
25 +func (c *Cache) Ready() bool {
26 + return c.inner.Ready()
27 +}
28 +
29 +// Lookup finds a cached item by hash + name. O(1), no I/O.
30 +func (c *Cache) Lookup(hash uint32, name string) (CacheItem, bool) {
31 + return c.inner.Lookup(hash, name)
32 +}
33 +
34 +// Status returns a diagnostic snapshot for the L3 cache.
35 +func (c *Cache) Status() CacheStatus {
36 + return c.inner.Status()
37 +}
38 +
39 +// Close frees all cached items and closes the L2 client.
40 +func (c *Cache) Close() {
41 + c.inner.Close()
42 +}
src/go/pkg/netipc/service/cgroups/cache_windows.go new
+42
@@ -0,0 +1,42 @@
1 +//go:build windows
2 +
3 +package cgroups
4 +
5 +import (
6 + raw "github.com/netdata/netdata/go/plugins/pkg/netipc/service/raw"
7 +)
8 +
9 +// Cache is the public L3 client-side cgroups snapshot cache.
10 +type Cache struct {
11 + inner *raw.Cache
12 +}
13 +
14 +// NewCache creates a new L3 cache. Does NOT connect.
15 +func NewCache(runDir, serviceName string, config ClientConfig) *Cache {
16 + return &Cache{inner: raw.NewCache(runDir, serviceName, clientConfigToTransport(config))}
17 +}
18 +
19 +// Refresh drives the L2 client and requests a fresh snapshot.
20 +func (c *Cache) Refresh() bool {
21 + return c.inner.Refresh()
22 +}
23 +
24 +// Ready returns true if at least one successful refresh has occurred.
25 +func (c *Cache) Ready() bool {
26 + return c.inner.Ready()
27 +}
28 +
29 +// Lookup finds a cached item by hash + name. O(1), no I/O.
30 +func (c *Cache) Lookup(hash uint32, name string) (CacheItem, bool) {
31 + return c.inner.Lookup(hash, name)
32 +}
33 +
34 +// Status returns a diagnostic snapshot for the L3 cache.
35 +func (c *Cache) Status() CacheStatus {
36 + return c.inner.Status()
37 +}
38 +
39 +// Close frees all cached items and closes the L2 client.
40 +func (c *Cache) Close() {
41 + c.inner.Close()
42 +}
src/go/pkg/netipc/service/cgroups/cgroups_unix_test.go new
+203
@@ -0,0 +1,203 @@
1 +//go:build unix
2 +
3 +package cgroups
4 +
5 +import (
6 + "fmt"
7 + "os"
8 + "sync/atomic"
9 + "testing"
10 + "time"
11 +
12 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
13 +)
14 +
15 +const (
16 + testRunDirUnix = "/tmp/nipc_go_cgroups_public"
17 + testAuthToken = uint64(0xDEADBEEFCAFEBABE)
18 + testResponseSize = 65536
19 +)
20 +
21 +var unixServiceCounter atomic.Uint64
22 +
23 +func uniqueUnixService(prefix string) string {
24 + return fmt.Sprintf("%s_%d_%d_%d", prefix, os.Getpid(), unixServiceCounter.Add(1), time.Now().UnixNano())
25 +}
26 +
27 +func ensureUnixRunDir(t *testing.T) {
28 + t.Helper()
29 + if err := os.MkdirAll(testRunDirUnix, 0o700); err != nil {
30 + t.Fatalf("mkdir: %v", err)
31 + }
32 +}
33 +
34 +func cleanupUnix(service string) {
35 + _ = os.Remove(fmt.Sprintf("%s/%s.sock", testRunDirUnix, service))
36 +}
37 +
38 +func testUnixServerConfig() ServerConfig {
39 + return ServerConfig{
40 + SupportedProfiles: protocol.ProfileBaseline,
41 + PreferredProfiles: protocol.ProfileBaseline,
42 + MaxRequestBatchItems: 1,
43 + MaxResponsePayloadBytes: testResponseSize,
44 + AuthToken: testAuthToken,
45 + }
46 +}
47 +
48 +func testUnixClientConfig() ClientConfig {
49 + return ClientConfig{
50 + SupportedProfiles: protocol.ProfileBaseline,
51 + PreferredProfiles: protocol.ProfileBaseline,
52 + MaxRequestBatchItems: 1,
53 + MaxResponsePayloadBytes: testResponseSize,
54 + AuthToken: testAuthToken,
55 + }
56 +}
57 +
58 +func testUnixHandler() Handler {
59 + return Handler{
60 + Handle: func(req *protocol.CgroupsRequest, builder *protocol.CgroupsBuilder) bool {
61 + if req.LayoutVersion != 1 || req.Flags != 0 {
62 + return false
63 + }
64 + builder.SetHeader(1, 42)
65 + items := []struct {
66 + hash, options, enabled uint32
67 + name, path []byte
68 + }{
69 + {1001, 0, 1, []byte("docker-abc123"), []byte("/sys/fs/cgroup/docker/abc123")},
70 + {2002, 0, 1, []byte("k8s-pod-xyz"), []byte("/sys/fs/cgroup/kubepods/xyz")},
71 + {3003, 0, 0, []byte("systemd-user"), []byte("/sys/fs/cgroup/user.slice/user-1000")},
72 + }
73 + for _, item := range items {
74 + if err := builder.Add(item.hash, item.options, item.enabled, item.name, item.path); err != nil {
75 + return false
76 + }
77 + }
78 + return true
79 + },
80 + SnapshotMaxItems: 3,
81 + }
82 +}
83 +
84 +type unixTestServer struct {
85 + server *Server
86 + done chan struct{}
87 +}
88 +
89 +func startUnixTestServer(t *testing.T, service string) *unixTestServer {
90 + t.Helper()
91 + ensureUnixRunDir(t)
92 + cleanupUnix(service)
93 +
94 + s := NewServer(testRunDirUnix, service, testUnixServerConfig(), testUnixHandler())
95 + done := make(chan struct{})
96 + go func() {
97 + defer close(done)
98 + _ = s.Run()
99 + }()
100 +
101 + time.Sleep(50 * time.Millisecond)
102 + return &unixTestServer{server: s, done: done}
103 +}
104 +
105 +func (ts *unixTestServer) stop() {
106 + ts.server.Stop()
107 + select {
108 + case <-ts.done:
109 + case <-time.After(2 * time.Second):
110 + }
111 +}
112 +
113 +func connectReadyUnix(t *testing.T, client *Client) {
114 + t.Helper()
115 + for i := 0; i < 200; i++ {
116 + client.Refresh()
117 + if client.Ready() {
118 + return
119 + }
120 + time.Sleep(10 * time.Millisecond)
121 + }
122 + t.Fatal("client did not reach READY state")
123 +}
124 +
125 +func TestSnapshotRoundTripUnix(t *testing.T) {
126 + service := uniqueUnixService("snapshot")
127 + ts := startUnixTestServer(t, service)
128 + defer ts.stop()
129 +
130 + client := NewClient(testRunDirUnix, service, testUnixClientConfig())
131 + defer client.Close()
132 + connectReadyUnix(t, client)
133 +
134 + view, err := client.CallSnapshot()
135 + if err != nil {
136 + t.Fatalf("CallSnapshot failed: %v", err)
137 + }
138 + if view.ItemCount != 3 {
139 + t.Fatalf("ItemCount = %d, want 3", view.ItemCount)
140 + }
141 + if view.SystemdEnabled != 1 {
142 + t.Fatalf("SystemdEnabled = %d, want 1", view.SystemdEnabled)
143 + }
144 + if view.Generation != 42 {
145 + t.Fatalf("Generation = %d, want 42", view.Generation)
146 + }
147 + item0, err := view.Item(0)
148 + if err != nil {
149 + t.Fatalf("Item(0): %v", err)
150 + }
151 + if item0.Hash != 1001 || item0.Name.String() != "docker-abc123" {
152 + t.Fatalf("unexpected item0: %+v", item0)
153 + }
154 +}
155 +
156 +func TestCacheRoundTripUnix(t *testing.T) {
157 + service := uniqueUnixService("cache")
158 + ts := startUnixTestServer(t, service)
159 + defer ts.stop()
160 +
161 + cache := NewCache(testRunDirUnix, service, testUnixClientConfig())
162 + defer cache.Close()
163 +
164 + var updated bool
165 + for i := 0; i < 200; i++ {
166 + if cache.Refresh() {
167 + updated = true
168 + break
169 + }
170 + time.Sleep(10 * time.Millisecond)
171 + }
172 + if !updated {
173 + t.Fatal("Refresh never succeeded")
174 + }
175 + if !cache.Ready() {
176 + t.Fatal("cache not ready after refresh")
177 + }
178 +
179 + item, ok := cache.Lookup(1001, "docker-abc123")
180 + if !ok {
181 + t.Fatal("lookup failed")
182 + }
183 + if item.Path != "/sys/fs/cgroup/docker/abc123" {
184 + t.Fatalf("unexpected path: %q", item.Path)
185 + }
186 +
187 + status := cache.Status()
188 + if !status.Populated || status.ItemCount != 3 || status.Generation != 42 {
189 + t.Fatalf("unexpected status: %+v", status)
190 + }
191 +}
192 +
193 +func TestClientNotReadyReturnsErrorUnix(t *testing.T) {
194 + service := uniqueUnixService("not_ready")
195 + cleanupUnix(service)
196 +
197 + client := NewClient(testRunDirUnix, service, testUnixClientConfig())
198 + defer client.Close()
199 +
200 + if _, err := client.CallSnapshot(); err != protocol.ErrBadLayout {
201 + t.Fatalf("CallSnapshot err = %v, want %v", err, protocol.ErrBadLayout)
202 + }
203 +}
src/go/pkg/netipc/service/cgroups/cgroups_windows_test.go new
+227
@@ -0,0 +1,227 @@
1 +//go:build windows
2 +
3 +package cgroups
4 +
5 +import (
6 + "fmt"
7 + "os"
8 + "sync/atomic"
9 + "testing"
10 + "time"
11 +
12 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
13 +)
14 +
15 +const testWinRunDir = `C:\Temp\nipc_go_cgroups_public`
16 +
17 +const (
18 + testWinAuthToken = uint64(0xDEADBEEFCAFEBABE)
19 + testWinResponseSize = 65536
20 +)
21 +
22 +var winServiceCounter atomic.Uint64
23 +
24 +func uniqueWinService(prefix string) string {
25 + return fmt.Sprintf("%s_%d_%d", prefix, os.Getpid(), winServiceCounter.Add(1))
26 +}
27 +
28 +func ensureWinRunDir(t *testing.T) {
29 + t.Helper()
30 + if err := os.MkdirAll(testWinRunDir, 0o700); err != nil {
31 + t.Fatalf("mkdir: %v", err)
32 + }
33 +}
34 +
35 +func testWinServerConfig() ServerConfig {
36 + return ServerConfig{
37 + SupportedProfiles: protocol.ProfileBaseline,
38 + PreferredProfiles: protocol.ProfileBaseline,
39 + MaxRequestBatchItems: 1,
40 + MaxResponsePayloadBytes: testWinResponseSize,
41 + AuthToken: testWinAuthToken,
42 + }
43 +}
44 +
45 +func testWinClientConfig() ClientConfig {
46 + return ClientConfig{
47 + SupportedProfiles: protocol.ProfileBaseline,
48 + PreferredProfiles: protocol.ProfileBaseline,
49 + MaxRequestBatchItems: 1,
50 + MaxResponsePayloadBytes: testWinResponseSize,
51 + AuthToken: testWinAuthToken,
52 + }
53 +}
54 +
55 +func testWinHandler() Handler {
56 + return Handler{
57 + Handle: func(req *protocol.CgroupsRequest, builder *protocol.CgroupsBuilder) bool {
58 + if req.LayoutVersion != 1 || req.Flags != 0 {
59 + return false
60 + }
61 + builder.SetHeader(1, 42)
62 + items := []struct {
63 + hash, options, enabled uint32
64 + name, path []byte
65 + }{
66 + {1001, 0, 1, []byte("docker-abc123"), []byte("/sys/fs/cgroup/docker/abc123")},
67 + {2002, 0, 1, []byte("k8s-pod-xyz"), []byte("/sys/fs/cgroup/kubepods/xyz")},
68 + {3003, 0, 0, []byte("systemd-user"), []byte("/sys/fs/cgroup/user.slice/user-1000")},
69 + }
70 + for _, item := range items {
71 + if err := builder.Add(item.hash, item.options, item.enabled, item.name, item.path); err != nil {
72 + return false
73 + }
74 + }
75 + return true
76 + },
77 + SnapshotMaxItems: 3,
78 + }
79 +}
80 +
81 +type winTestServer struct {
82 + server *Server
83 + done chan struct{}
84 +}
85 +
86 +func startWinTestServer(t *testing.T, service string) *winTestServer {
87 + t.Helper()
88 + ensureWinRunDir(t)
89 + s := NewServer(testWinRunDir, service, testWinServerConfig(), testWinHandler())
90 + done := make(chan struct{})
91 + go func() {
92 + defer close(done)
93 + _ = s.Run()
94 + }()
95 + time.Sleep(200 * time.Millisecond)
96 + return &winTestServer{server: s, done: done}
97 +}
98 +
99 +func (ts *winTestServer) stop() {
100 + ts.server.Stop()
101 + select {
102 + case <-ts.done:
103 + case <-time.After(2 * time.Second):
104 + }
105 +}
106 +
107 +func connectReadyWin(t *testing.T, client *Client) {
108 + t.Helper()
109 + for i := 0; i < 200; i++ {
110 + client.Refresh()
111 + if client.Ready() {
112 + return
113 + }
114 + time.Sleep(10 * time.Millisecond)
115 + }
116 + t.Fatal("client did not reach READY state")
117 +}
118 +
119 +func TestSnapshotRoundTripWindows(t *testing.T) {
120 + service := uniqueWinService("snapshot")
121 + ts := startWinTestServer(t, service)
122 + defer ts.stop()
123 +
124 + client := NewClient(testWinRunDir, service, testWinClientConfig())
125 + defer client.Close()
126 + connectReadyWin(t, client)
127 +
128 + view, err := client.CallSnapshot()
129 + if err != nil {
130 + t.Fatalf("CallSnapshot failed: %v", err)
131 + }
132 + if view.ItemCount != 3 || view.SystemdEnabled != 1 || view.Generation != 42 {
133 + t.Fatalf("unexpected snapshot header: %+v", view)
134 + }
135 +}
136 +
137 +func TestCacheRoundTripWindows(t *testing.T) {
138 + service := uniqueWinService("cache")
139 + ts := startWinTestServer(t, service)
140 + defer ts.stop()
141 +
142 + cache := NewCache(testWinRunDir, service, testWinClientConfig())
143 + defer cache.Close()
144 +
145 + if cache.Ready() {
146 + t.Fatal("cache unexpectedly ready before refresh")
147 + }
148 + status0 := cache.Status()
149 + if status0.Populated || status0.ItemCount != 0 || status0.RefreshSuccessCount != 0 || status0.RefreshFailureCount != 0 || status0.ConnectionState != StateDisconnected || status0.LastRefreshTs != 0 {
150 + t.Fatalf("unexpected initial cache status: %+v", status0)
151 + }
152 +
153 + var updated bool
154 + for i := 0; i < 200; i++ {
155 + if cache.Refresh() {
156 + updated = true
157 + break
158 + }
159 + time.Sleep(10 * time.Millisecond)
160 + }
161 + if !updated {
162 + t.Fatal("Refresh never succeeded")
163 + }
164 + if !cache.Ready() {
165 + t.Fatal("cache not ready after refresh")
166 + }
167 + item, ok := cache.Lookup(1001, "docker-abc123")
168 + if !ok || item.Path != "/sys/fs/cgroup/docker/abc123" {
169 + t.Fatalf("unexpected cache item: %+v ok=%v", item, ok)
170 + }
171 + status1 := cache.Status()
172 + if !status1.Populated || status1.ItemCount != 3 || status1.SystemdEnabled != 1 || status1.Generation != 42 || status1.RefreshSuccessCount != 1 || status1.RefreshFailureCount != 0 || status1.ConnectionState != StateReady || status1.LastRefreshTs < 0 {
173 + t.Fatalf("unexpected refreshed cache status: %+v", status1)
174 + }
175 +}
176 +
177 +func TestClientNotReadyReturnsErrorWindows(t *testing.T) {
178 + service := uniqueWinService("not_ready")
179 + client := NewClient(testWinRunDir, service, testWinClientConfig())
180 + defer client.Close()
181 +
182 + status := client.Status()
183 + if status.State != StateDisconnected || status.ConnectCount != 0 || status.ReconnectCount != 0 || status.CallCount != 0 || status.ErrorCount != 0 {
184 + t.Fatalf("unexpected initial client status: %+v", status)
185 + }
186 +
187 + if _, err := client.CallSnapshot(); err != protocol.ErrBadLayout {
188 + t.Fatalf("CallSnapshot err = %v, want %v", err, protocol.ErrBadLayout)
189 + }
190 +
191 + status = client.Status()
192 + if status.State != StateDisconnected || status.ErrorCount != 1 {
193 + t.Fatalf("unexpected error-path client status: %+v", status)
194 + }
195 +}
196 +
197 +func TestClientStatusWindows(t *testing.T) {
198 + service := uniqueWinService("status")
199 + ts := startWinTestServer(t, service)
200 + defer ts.stop()
201 +
202 + client := NewClient(testWinRunDir, service, testWinClientConfig())
203 + defer client.Close()
204 + connectReadyWin(t, client)
205 +
206 + status0 := client.Status()
207 + if status0.State != StateReady || status0.ConnectCount != 1 || status0.ReconnectCount != 0 || status0.CallCount != 0 || status0.ErrorCount != 0 {
208 + t.Fatalf("unexpected ready client status: %+v", status0)
209 + }
210 +
211 + if _, err := client.CallSnapshot(); err != nil {
212 + t.Fatalf("CallSnapshot failed: %v", err)
213 + }
214 +
215 + status1 := client.Status()
216 + if status1.State != StateReady || status1.ConnectCount != 1 || status1.ReconnectCount != 0 || status1.CallCount != 1 || status1.ErrorCount != 0 {
217 + t.Fatalf("unexpected post-call client status: %+v", status1)
218 + }
219 +}
220 +
221 +func TestNewServerWithWorkersWindows(t *testing.T) {
222 + service := uniqueWinService("workers")
223 + server := NewServerWithWorkers(testWinRunDir, service, testWinServerConfig(), testWinHandler(), 3)
224 + if server == nil || server.inner == nil {
225 + t.Fatal("NewServerWithWorkers returned nil")
226 + }
227 +}
src/go/pkg/netipc/service/cgroups/client.go new
+113
@@ -0,0 +1,113 @@
1 +//go:build unix
2 +
3 +package cgroups
4 +
5 +import (
6 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
7 + raw "github.com/netdata/netdata/go/plugins/pkg/netipc/service/raw"
8 + "github.com/netdata/netdata/go/plugins/pkg/netipc/transport/posix"
9 +)
10 +
11 +func snapshotDispatch(handler Handler) raw.DispatchHandler {
12 + return raw.SnapshotDispatch(handler.Handle, handler.SnapshotMaxItems)
13 +}
14 +
15 +func clientConfigToTransport(config ClientConfig) posix.ClientConfig {
16 + return posix.ClientConfig{
17 + SupportedProfiles: config.SupportedProfiles,
18 + PreferredProfiles: config.PreferredProfiles,
19 + MaxRequestBatchItems: config.MaxRequestBatchItems,
20 + MaxResponsePayloadBytes: config.MaxResponsePayloadBytes,
21 + MaxResponseBatchItems: config.MaxRequestBatchItems,
22 + AuthToken: config.AuthToken,
23 + }
24 +}
25 +
26 +func serverConfigToTransport(config ServerConfig) posix.ServerConfig {
27 + return posix.ServerConfig{
28 + SupportedProfiles: config.SupportedProfiles,
29 + PreferredProfiles: config.PreferredProfiles,
30 + MaxRequestBatchItems: config.MaxRequestBatchItems,
31 + MaxResponsePayloadBytes: config.MaxResponsePayloadBytes,
32 + MaxResponseBatchItems: config.MaxRequestBatchItems,
33 + AuthToken: config.AuthToken,
34 + }
35 +}
36 +
37 +// Client is the public L2 client context for the cgroups-snapshot service.
38 +type Client struct {
39 + inner *raw.Client
40 +}
41 +
42 +// NewClient creates a new client context. Does NOT connect.
43 +func NewClient(runDir, serviceName string, config ClientConfig) *Client {
44 + return &Client{inner: raw.NewSnapshotClient(runDir, serviceName, clientConfigToTransport(config))}
45 +}
46 +
47 +// Refresh attempts connect if DISCONNECTED/NOT_FOUND, reconnect if BROKEN.
48 +func (c *Client) Refresh() bool {
49 + return c.inner.Refresh()
50 +}
51 +
52 +// Ready returns true only if the client is in the READY state.
53 +func (c *Client) Ready() bool {
54 + return c.inner.Ready()
55 +}
56 +
57 +// Status returns a diagnostic counters snapshot.
58 +func (c *Client) Status() ClientStatus {
59 + return c.inner.Status()
60 +}
61 +
62 +// CallSnapshot performs a blocking typed cgroups snapshot call.
63 +func (c *Client) CallSnapshot() (*protocol.CgroupsResponseView, error) {
64 + return c.inner.CallSnapshot()
65 +}
66 +
67 +// Close tears down the connection and releases resources.
68 +func (c *Client) Close() {
69 + c.inner.Close()
70 +}
71 +
72 +// Server is the public managed server for the cgroups-snapshot service kind.
73 +type Server struct {
74 + inner *raw.Server
75 +}
76 +
77 +// NewServer creates a new managed server.
78 +func NewServer(runDir, serviceName string, config ServerConfig, handler Handler) *Server {
79 + return &Server{
80 + inner: raw.NewServer(
81 + runDir,
82 + serviceName,
83 + serverConfigToTransport(config),
84 + protocol.MethodCgroupsSnapshot,
85 + snapshotDispatch(handler),
86 + ),
87 + }
88 +}
89 +
90 +// NewServerWithWorkers creates a server with an explicit worker count limit.
91 +func NewServerWithWorkers(runDir, serviceName string, config ServerConfig,
92 + handler Handler, workerCount int) *Server {
93 + return &Server{
94 + inner: raw.NewServerWithWorkers(
95 + runDir,
96 + serviceName,
97 + serverConfigToTransport(config),
98 + protocol.MethodCgroupsSnapshot,
99 + snapshotDispatch(handler),
100 + workerCount,
101 + ),
102 + }
103 +}
104 +
105 +// Run starts the acceptor loop. Blocking.
106 +func (s *Server) Run() error {
107 + return s.inner.Run()
108 +}
109 +
110 +// Stop signals the server to stop.
111 +func (s *Server) Stop() {
112 + s.inner.Stop()
113 +}
src/go/pkg/netipc/service/cgroups/client_windows.go new
+113
@@ -0,0 +1,113 @@
1 +//go:build windows
2 +
3 +package cgroups
4 +
5 +import (
6 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
7 + raw "github.com/netdata/netdata/go/plugins/pkg/netipc/service/raw"
8 + windows "github.com/netdata/netdata/go/plugins/pkg/netipc/transport/windows"
9 +)
10 +
11 +func snapshotDispatch(handler Handler) raw.DispatchHandler {
12 + return raw.SnapshotDispatch(handler.Handle, handler.SnapshotMaxItems)
13 +}
14 +
15 +func clientConfigToTransport(config ClientConfig) windows.ClientConfig {
16 + return windows.ClientConfig{
17 + SupportedProfiles: config.SupportedProfiles,
18 + PreferredProfiles: config.PreferredProfiles,
19 + MaxRequestBatchItems: config.MaxRequestBatchItems,
20 + MaxResponsePayloadBytes: config.MaxResponsePayloadBytes,
21 + MaxResponseBatchItems: config.MaxRequestBatchItems,
22 + AuthToken: config.AuthToken,
23 + }
24 +}
25 +
26 +func serverConfigToTransport(config ServerConfig) windows.ServerConfig {
27 + return windows.ServerConfig{
28 + SupportedProfiles: config.SupportedProfiles,
29 + PreferredProfiles: config.PreferredProfiles,
30 + MaxRequestBatchItems: config.MaxRequestBatchItems,
31 + MaxResponsePayloadBytes: config.MaxResponsePayloadBytes,
32 + MaxResponseBatchItems: config.MaxRequestBatchItems,
33 + AuthToken: config.AuthToken,
34 + }
35 +}
36 +
37 +// Client is the public L2 client context for the cgroups-snapshot service.
38 +type Client struct {
39 + inner *raw.Client
40 +}
41 +
42 +// NewClient creates a new client context. Does NOT connect.
43 +func NewClient(runDir, serviceName string, config ClientConfig) *Client {
44 + return &Client{inner: raw.NewSnapshotClient(runDir, serviceName, clientConfigToTransport(config))}
45 +}
46 +
47 +// Refresh attempts connect if DISCONNECTED/NOT_FOUND, reconnect if BROKEN.
48 +func (c *Client) Refresh() bool {
49 + return c.inner.Refresh()
50 +}
51 +
52 +// Ready returns true only if the client is in the READY state.
53 +func (c *Client) Ready() bool {
54 + return c.inner.Ready()
55 +}
56 +
57 +// Status returns a diagnostic counters snapshot.
58 +func (c *Client) Status() ClientStatus {
59 + return c.inner.Status()
60 +}
61 +
62 +// CallSnapshot performs a blocking typed cgroups snapshot call.
63 +func (c *Client) CallSnapshot() (*protocol.CgroupsResponseView, error) {
64 + return c.inner.CallSnapshot()
65 +}
66 +
67 +// Close tears down the connection and releases resources.
68 +func (c *Client) Close() {
69 + c.inner.Close()
70 +}
71 +
72 +// Server is the public managed server for the cgroups-snapshot service kind.
73 +type Server struct {
74 + inner *raw.Server
75 +}
76 +
77 +// NewServer creates a new managed server.
78 +func NewServer(runDir, serviceName string, config ServerConfig, handler Handler) *Server {
79 + return &Server{
80 + inner: raw.NewServer(
81 + runDir,
82 + serviceName,
83 + serverConfigToTransport(config),
84 + protocol.MethodCgroupsSnapshot,
85 + snapshotDispatch(handler),
86 + ),
87 + }
88 +}
89 +
90 +// NewServerWithWorkers creates a server with an explicit worker count limit.
91 +func NewServerWithWorkers(runDir, serviceName string, config ServerConfig,
92 + handler Handler, workerCount int) *Server {
93 + return &Server{
94 + inner: raw.NewServerWithWorkers(
95 + runDir,
96 + serviceName,
97 + serverConfigToTransport(config),
98 + protocol.MethodCgroupsSnapshot,
99 + snapshotDispatch(handler),
100 + workerCount,
101 + ),
102 + }
103 +}
104 +
105 +// Run starts the acceptor loop. Blocking.
106 +func (s *Server) Run() error {
107 + return s.inner.Run()
108 +}
109 +
110 +// Stop signals the server to stop.
111 +func (s *Server) Stop() {
112 + s.inner.Stop()
113 +}
src/go/pkg/netipc/service/cgroups/types.go new
+72
@@ -0,0 +1,72 @@
1 +// Package cgroups provides the public single-kind L2/L3 surface for the
2 +// cgroups-snapshot service.
3 +//
4 +// Clients connect to a service kind, not to a plugin identity. One service
5 +// endpoint serves one request kind only. The outer request code remains part
6 +// of the envelope for validation, not public multi-method dispatch.
7 +package cgroups
8 +
9 +import (
10 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
11 + raw "github.com/netdata/netdata/go/plugins/pkg/netipc/service/raw"
12 +)
13 +
14 +// ClientState represents the connection state machine.
15 +type ClientState = raw.ClientState
16 +
17 +const (
18 + StateDisconnected = raw.StateDisconnected
19 + StateConnecting = raw.StateConnecting
20 + StateReady = raw.StateReady
21 + StateNotFound = raw.StateNotFound
22 + StateAuthFailed = raw.StateAuthFailed
23 + StateIncompatible = raw.StateIncompatible
24 + StateBroken = raw.StateBroken
25 +)
26 +
27 +// ClientStatus is a diagnostic counters snapshot.
28 +type ClientStatus = raw.ClientStatus
29 +
30 +// ClientConfig is the public L2/L3 client configuration for the
31 +// cgroups-snapshot service.
32 +//
33 +// Transport-only tuning stays below the public typed API.
34 +type ClientConfig struct {
35 + SupportedProfiles uint32
36 + PreferredProfiles uint32
37 + MaxRequestBatchItems uint32
38 + MaxResponsePayloadBytes uint32
39 + AuthToken uint64
40 +}
41 +
42 +// ServerConfig is the public typed-server configuration for the
43 +// cgroups-snapshot service.
44 +//
45 +// Transport-only tuning stays below the public typed API.
46 +type ServerConfig struct {
47 + SupportedProfiles uint32
48 + PreferredProfiles uint32
49 + MaxRequestBatchItems uint32
50 + MaxResponsePayloadBytes uint32
51 + AuthToken uint64
52 +}
53 +
54 +// SnapshotHandler is the typed callback used by the cgroups-snapshot service.
55 +type SnapshotHandler = func(*protocol.CgroupsRequest, *protocol.CgroupsBuilder) bool
56 +
57 +// Handler defines the public typed callback surface for the cgroups-snapshot
58 +// service. A nil handler means the service is unavailable.
59 +type Handler struct {
60 + Handle SnapshotHandler
61 +
62 + // SnapshotMaxItems optionally caps the number of snapshot items the
63 + // internal builder reserves directory space for. When zero, the library
64 + // derives a safe upper bound from the negotiated response buffer size.
65 + SnapshotMaxItems uint32
66 +}
67 +
68 +// CacheItem is an owned copy of a single cgroup item.
69 +type CacheItem = raw.CacheItem
70 +
71 +// CacheStatus is a diagnostic snapshot for the L3 cache.
72 +type CacheStatus = raw.CacheStatus
src/go/pkg/netipc/service/raw/cache.go new
+163
@@ -0,0 +1,163 @@
1 +//go:build unix
2 +
3 +// L3: Client-side cgroups snapshot cache (POSIX).
4 +
5 +package raw
6 +
7 +import (
8 + "time"
9 +
10 + "github.com/netdata/netdata/go/plugins/pkg/netipc/transport/posix"
11 +)
12 +
13 +// cacheBucket is one open-addressing bucket for hash+name lookup.
14 +type cacheBucket struct {
15 + index int
16 + used bool
17 +}
18 +
19 +func cacheHashName(name string) uint32 {
20 + h := uint32(5381)
21 + for i := 0; i < len(name); i++ {
22 + h = ((h << 5) + h) + uint32(name[i])
23 + }
24 + return h
25 +}
26 +
27 +// Cache is an L3 client-side cgroups snapshot cache.
28 +type Cache struct {
29 + client *Client
30 + items []CacheItem
31 + // Open-addressing hash table: (hash ^ djb2(name)) -> index into items slice.
32 + buckets []cacheBucket
33 +
34 + systemdEnabled uint32
35 + generation uint64
36 + populated bool
37 + refreshSuccessCount uint32
38 + refreshFailureCount uint32
39 + epoch time.Time // monotonic reference point
40 + lastRefreshTs int64 // elapsed ms since epoch
41 +}
42 +
43 +// NewCache creates a new L3 cache. Creates the underlying L2 client
44 +// context. Does NOT connect. Does NOT require the server to be running.
45 +// Cache starts empty (populated == false).
46 +func NewCache(runDir, serviceName string, config posix.ClientConfig) *Cache {
47 + return &Cache{
48 + client: NewSnapshotClient(runDir, serviceName, config),
49 + epoch: time.Now(),
50 + }
51 +}
52 +
53 +// Refresh drives the L2 client (connect/reconnect as needed) and
54 +// requests a fresh snapshot. On success, rebuilds the local cache.
55 +// On failure, preserves the previous cache.
56 +//
57 +// Returns true if the cache was updated.
58 +func (c *Cache) Refresh() bool {
59 + c.client.Refresh()
60 +
61 + view, err := c.client.CallSnapshot()
62 + if err != nil {
63 + c.refreshFailureCount++
64 + return false
65 + }
66 +
67 + newItems := make([]CacheItem, 0, view.ItemCount)
68 + for i := uint32(0); i < view.ItemCount; i++ {
69 + iv, ierr := view.Item(i)
70 + if ierr != nil {
71 + c.refreshFailureCount++
72 + return false
73 + }
74 + newItems = append(newItems, CacheItem{
75 + Hash: iv.Hash,
76 + Options: iv.Options,
77 + Enabled: iv.Enabled,
78 + Name: iv.Name.String(),
79 + Path: iv.Path.String(),
80 + })
81 + }
82 +
83 + // Rebuild open-addressing lookup table.
84 + var buckets []cacheBucket
85 + if len(newItems) > 0 {
86 + bcount := nextPowerOf2U32(uint32(len(newItems)) * 2)
87 + buckets = make([]cacheBucket, bcount)
88 + mask := bcount - 1
89 + for i := range newItems {
90 + slot := (newItems[i].Hash ^ cacheHashName(newItems[i].Name)) & mask
91 + for buckets[slot].used {
92 + slot = (slot + 1) & mask
93 + }
94 + buckets[slot].index = i
95 + buckets[slot].used = true
96 + }
97 + }
98 +
99 + c.items = newItems
100 + c.buckets = buckets
101 + c.systemdEnabled = view.SystemdEnabled
102 + c.generation = view.Generation
103 + c.populated = true
104 + c.refreshSuccessCount++
105 + c.lastRefreshTs = time.Since(c.epoch).Milliseconds()
106 +
107 + return true
108 +}
109 +
110 +// Ready returns true if at least one successful refresh has occurred.
111 +func (c *Cache) Ready() bool {
112 + return c.populated
113 +}
114 +
115 +// Lookup finds a cached item by hash + name. O(1) via open-addressing hash
116 +// table. No I/O.
117 +func (c *Cache) Lookup(hash uint32, name string) (CacheItem, bool) {
118 + if !c.populated {
119 + return CacheItem{}, false
120 + }
121 +
122 + if len(c.buckets) > 0 {
123 + mask := uint32(len(c.buckets) - 1)
124 + slot := (hash ^ cacheHashName(name)) & mask
125 + for c.buckets[slot].used {
126 + item := &c.items[c.buckets[slot].index]
127 + if item.Hash == hash && item.Name == name {
128 + return *item, true
129 + }
130 + slot = (slot + 1) & mask
131 + }
132 + return CacheItem{}, false
133 + }
134 +
135 + for i := range c.items {
136 + if c.items[i].Hash == hash && c.items[i].Name == name {
137 + return c.items[i], true
138 + }
139 + }
140 + return CacheItem{}, false
141 +}
142 +
143 +// Status returns a diagnostic snapshot for the L3 cache.
144 +func (c *Cache) Status() CacheStatus {
145 + return CacheStatus{
146 + Populated: c.populated,
147 + ItemCount: uint32(len(c.items)),
148 + SystemdEnabled: c.systemdEnabled,
149 + Generation: c.generation,
150 + RefreshSuccessCount: c.refreshSuccessCount,
151 + RefreshFailureCount: c.refreshFailureCount,
152 + ConnectionState: c.client.state,
153 + LastRefreshTs: c.lastRefreshTs,
154 + }
155 +}
156 +
157 +// Close frees all cached items and closes the L2 client.
158 +func (c *Cache) Close() {
159 + c.items = nil
160 + c.buckets = nil
161 + c.populated = false
162 + c.client.Close()
163 +}
src/go/pkg/netipc/service/raw/cache_test.go new
+340
@@ -0,0 +1,340 @@
1 +//go:build unix
2 +
3 +package raw
4 +
5 +import (
6 + "fmt"
7 + "testing"
8 + "time"
9 +
10 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
11 +)
12 +
13 +// --- L3 Cache Tests ---
14 +
15 +func TestCacheFullRoundTrip(t *testing.T) {
16 + svc := "go_cache_rt"
17 + ensureRunDir()
18 + cleanupAll(svc)
19 +
20 + ts := startTestServer(svc, testSnapshotDispatch())
21 + defer ts.stop()
22 +
23 + cache := NewCache(testRunDir, svc, testClientConfig())
24 + if cache.Ready() {
25 + t.Fatal("should not be ready before refresh")
26 + }
27 +
28 + // Ensure monotonic epoch advances past 0ms before refresh
29 + time.Sleep(2 * time.Millisecond)
30 +
31 + // Refresh populates the cache
32 + updated := cache.Refresh()
33 + if !updated {
34 + t.Fatal("refresh should update the cache")
35 + }
36 + if !cache.Ready() {
37 + t.Fatal("should be ready after successful refresh")
38 + }
39 +
40 + // Lookup by hash + name
41 + item, found := cache.Lookup(1001, "docker-abc123")
42 + if !found {
43 + t.Fatal("item should be found")
44 + }
45 + if item.Hash != 1001 {
46 + t.Fatalf("expected hash=1001, got %d", item.Hash)
47 + }
48 + if item.Options != 0 {
49 + t.Fatalf("expected options=0, got %d", item.Options)
50 + }
51 + if item.Enabled != 1 {
52 + t.Fatalf("expected enabled=1, got %d", item.Enabled)
53 + }
54 + if item.Name != "docker-abc123" {
55 + t.Fatalf("expected name=docker-abc123, got %q", item.Name)
56 + }
57 + if item.Path != "/sys/fs/cgroup/docker/abc123" {
58 + t.Fatalf("expected path=/sys/fs/cgroup/docker/abc123, got %q", item.Path)
59 + }
60 +
61 + item2, found2 := cache.Lookup(3003, "systemd-user")
62 + if !found2 {
63 + t.Fatal("item 2 should be found")
64 + }
65 + if item2.Enabled != 0 {
66 + t.Fatalf("expected enabled=0, got %d", item2.Enabled)
67 + }
68 +
69 + // Status
70 + status := cache.Status()
71 + if !status.Populated {
72 + t.Fatal("status should be populated")
73 + }
74 + if status.ItemCount != 3 {
75 + t.Fatalf("expected item_count=3, got %d", status.ItemCount)
76 + }
77 + if status.SystemdEnabled != 1 {
78 + t.Fatalf("expected systemd_enabled=1, got %d", status.SystemdEnabled)
79 + }
80 + if status.Generation != 42 {
81 + t.Fatalf("expected generation=42, got %d", status.Generation)
82 + }
83 + if status.RefreshSuccessCount != 1 {
84 + t.Fatalf("expected success_count=1, got %d", status.RefreshSuccessCount)
85 + }
86 + if status.RefreshFailureCount != 0 {
87 + t.Fatalf("expected failure_count=0, got %d", status.RefreshFailureCount)
88 + }
89 + if status.ConnectionState != StateReady {
90 + t.Fatalf("expected connection_state=StateReady, got %d", status.ConnectionState)
91 + }
92 + if status.LastRefreshTs <= 0 {
93 + t.Fatalf("expected positive last_refresh_ts, got %d", status.LastRefreshTs)
94 + }
95 +
96 + cache.Close()
97 + cleanupAll(svc)
98 +}
99 +
100 +func TestCacheRefreshFailurePreserves(t *testing.T) {
101 + svc := "go_cache_preserve"
102 + ensureRunDir()
103 + cleanupAll(svc)
104 +
105 + ts := startTestServer(svc, testSnapshotDispatch())
106 +
107 + cache := NewCache(testRunDir, svc, testClientConfig())
108 +
109 + // First refresh populates cache
110 + if !cache.Refresh() {
111 + t.Fatal("first refresh should succeed")
112 + }
113 + if !cache.Ready() {
114 + t.Fatal("should be ready")
115 + }
116 + _, found := cache.Lookup(1001, "docker-abc123")
117 + if !found {
118 + t.Fatal("item should be found after first refresh")
119 + }
120 +
121 + // Kill server
122 + ts.stop()
123 + cleanupAll(svc)
124 + time.Sleep(50 * time.Millisecond)
125 +
126 + // Refresh fails, old cache preserved
127 + updated := cache.Refresh()
128 + if updated {
129 + t.Fatal("refresh should fail with no server")
130 + }
131 + if !cache.Ready() {
132 + t.Fatal("should still be ready (old cache preserved)")
133 + }
134 + _, found = cache.Lookup(1001, "docker-abc123")
135 + if !found {
136 + t.Fatal("item should still be found (old cache preserved)")
137 + }
138 +
139 + status := cache.Status()
140 + if status.RefreshSuccessCount != 1 {
141 + t.Fatalf("expected success_count=1, got %d", status.RefreshSuccessCount)
142 + }
143 + if status.RefreshFailureCount < 1 {
144 + t.Fatalf("expected failure_count >= 1, got %d", status.RefreshFailureCount)
145 + }
146 +
147 + cache.Close()
148 + cleanupAll(svc)
149 +}
150 +
151 +func TestCacheReconnectRebuilds(t *testing.T) {
152 + svc := "go_cache_reconn"
153 + ensureRunDir()
154 + cleanupAll(svc)
155 +
156 + ts1 := startTestServer(svc, testSnapshotDispatch())
157 +
158 + cache := NewCache(testRunDir, svc, testClientConfig())
159 + if !cache.Refresh() {
160 + t.Fatal("first refresh should succeed")
161 + }
162 + if cache.Status().ItemCount != 3 {
163 + t.Fatal("expected 3 items")
164 + }
165 +
166 + // Kill and restart server
167 + ts1.stop()
168 + cleanupAll(svc)
169 + time.Sleep(50 * time.Millisecond)
170 +
171 + ts2 := startTestServer(svc, testSnapshotDispatch())
172 + defer ts2.stop()
173 +
174 + // Refresh should reconnect and rebuild cache
175 + updated := cache.Refresh()
176 + if !updated {
177 + t.Fatal("refresh after reconnect should succeed")
178 + }
179 + if !cache.Ready() {
180 + t.Fatal("should be ready after reconnect")
181 + }
182 + if cache.Status().ItemCount != 3 {
183 + t.Fatal("expected 3 items after reconnect")
184 + }
185 + if cache.Status().RefreshSuccessCount != 2 {
186 + t.Fatalf("expected success_count=2, got %d", cache.Status().RefreshSuccessCount)
187 + }
188 +
189 + cache.Close()
190 + cleanupAll(svc)
191 +}
192 +
193 +func TestCacheLookupNotFound(t *testing.T) {
194 + svc := "go_cache_notfound"
195 + ensureRunDir()
196 + cleanupAll(svc)
197 +
198 + ts := startTestServer(svc, testSnapshotDispatch())
199 + defer ts.stop()
200 +
201 + cache := NewCache(testRunDir, svc, testClientConfig())
202 + if !cache.Refresh() {
203 + t.Fatal("refresh should succeed")
204 + }
205 +
206 + // Non-existent hash
207 + _, found := cache.Lookup(9999, "nonexistent")
208 + if found {
209 + t.Fatal("should not find nonexistent item")
210 + }
211 +
212 + // Correct hash, wrong name
213 + _, found = cache.Lookup(1001, "wrong-name")
214 + if found {
215 + t.Fatal("should not find with wrong name")
216 + }
217 +
218 + // Correct name, wrong hash
219 + _, found = cache.Lookup(9999, "docker-abc123")
220 + if found {
221 + t.Fatal("should not find with wrong hash")
222 + }
223 +
224 + cache.Close()
225 + cleanupAll(svc)
226 +}
227 +
228 +func TestCacheEmpty(t *testing.T) {
229 + svc := "go_cache_empty"
230 + ensureRunDir()
231 + cleanupAll(svc)
232 +
233 + cache := NewCache(testRunDir, svc, testClientConfig())
234 +
235 + // Not ready before any refresh
236 + if cache.Ready() {
237 + t.Fatal("should not be ready")
238 + }
239 +
240 + // Lookup on empty cache returns not-found
241 + _, found := cache.Lookup(1001, "docker-abc123")
242 + if found {
243 + t.Fatal("should not find in empty cache")
244 + }
245 +
246 + status := cache.Status()
247 + if status.Populated {
248 + t.Fatal("should not be populated")
249 + }
250 + if status.ItemCount != 0 {
251 + t.Fatalf("expected item_count=0, got %d", status.ItemCount)
252 + }
253 + if status.RefreshSuccessCount != 0 {
254 + t.Fatalf("expected success_count=0, got %d", status.RefreshSuccessCount)
255 + }
256 + if status.RefreshFailureCount != 0 {
257 + t.Fatalf("expected failure_count=0, got %d", status.RefreshFailureCount)
258 + }
259 +
260 + cleanupAll(svc)
261 +}
262 +
263 +func TestCacheLargeDataset(t *testing.T) {
264 + svc := "go_cache_large"
265 + ensureRunDir()
266 + cleanupAll(svc)
267 +
268 + const N = 1000
269 +
270 + // Handler that builds N items
271 + largeHandler := SnapshotDispatch(
272 + func(request *protocol.CgroupsRequest, builder *protocol.CgroupsBuilder) bool {
273 + if request.LayoutVersion != 1 || request.Flags != 0 {
274 + return false
275 + }
276 + builder.SetHeader(1, 100)
277 +
278 + for i := uint32(0); i < N; i++ {
279 + name := fmt.Sprintf("cgroup-%d", i)
280 + path := fmt.Sprintf("/sys/fs/cgroup/test/%d", i)
281 + enabled := uint32(1)
282 + if i%3 == 0 {
283 + enabled = 0
284 + }
285 + if err := builder.Add(i+1000, 0, enabled,
286 + []byte(name), []byte(path)); err != nil {
287 + return false
288 + }
289 + }
290 + return true
291 + },
292 + N,
293 + )
294 +
295 + cfg := testServerConfig()
296 + cfg.MaxResponsePayloadBytes = 256 * N
297 +
298 + s := NewServer(testRunDir, svc, cfg, protocol.MethodCgroupsSnapshot, largeHandler)
299 + doneCh := make(chan struct{})
300 + go func() {
301 + defer close(doneCh)
302 + s.Run()
303 + }()
304 + time.Sleep(100 * time.Millisecond)
305 +
306 + defer func() {
307 + s.Stop()
308 + <-doneCh
309 + }()
310 +
311 + ccfg := testClientConfig()
312 + ccfg.MaxResponsePayloadBytes = 256 * N
313 +
314 + cache := NewCache(testRunDir, svc, ccfg)
315 + if !cache.Refresh() {
316 + t.Fatal("refresh should succeed")
317 + }
318 + if cache.Status().ItemCount != N {
319 + t.Fatalf("expected %d items, got %d", N, cache.Status().ItemCount)
320 + }
321 +
322 + // Verify all lookups
323 + for i := uint32(0); i < N; i++ {
324 + name := fmt.Sprintf("cgroup-%d", i)
325 + item, found := cache.Lookup(i+1000, name)
326 + if !found {
327 + t.Fatalf("item %d not found", i)
328 + }
329 + if item.Hash != i+1000 {
330 + t.Fatalf("item %d: expected hash=%d, got %d", i, i+1000, item.Hash)
331 + }
332 + expectedPath := fmt.Sprintf("/sys/fs/cgroup/test/%d", i)
333 + if item.Path != expectedPath {
334 + t.Fatalf("item %d: expected path=%q, got %q", i, expectedPath, item.Path)
335 + }
336 + }
337 +
338 + cache.Close()
339 + cleanupAll(svc)
340 +}
src/go/pkg/netipc/service/raw/cache_windows.go new
+159
@@ -0,0 +1,159 @@
1 +//go:build windows
2 +
3 +// L3: Client-side cgroups snapshot cache (Windows).
4 +//
5 +// Identical cache logic as the POSIX version. Uses Windows Client.
6 +//
7 +// Pure Go — no cgo. Works with CGO_ENABLED=0.
8 +
9 +package raw
10 +
11 +import (
12 + "time"
13 +
14 + windows "github.com/netdata/netdata/go/plugins/pkg/netipc/transport/windows"
15 +)
16 +
17 +// cacheBucket is one open-addressing bucket for hash+name lookup.
18 +type cacheBucket struct {
19 + index int
20 + used bool
21 +}
22 +
23 +func cacheHashName(name string) uint32 {
24 + h := uint32(5381)
25 + for i := 0; i < len(name); i++ {
26 + h = ((h << 5) + h) + uint32(name[i])
27 + }
28 + return h
29 +}
30 +
31 +// Cache is an L3 client-side cgroups snapshot cache.
32 +type Cache struct {
33 + client *Client
34 + items []CacheItem
35 + // Open-addressing hash table: (hash ^ djb2(name)) -> index into items slice.
36 + buckets []cacheBucket
37 +
38 + systemdEnabled uint32
39 + generation uint64
40 + populated bool
41 + refreshSuccessCount uint32
42 + refreshFailureCount uint32
43 + epoch time.Time
44 + lastRefreshTs int64
45 +}
46 +
47 +// NewCache creates a new L3 cache.
48 +func NewCache(runDir, serviceName string, config windows.ClientConfig) *Cache {
49 + return &Cache{
50 + client: NewSnapshotClient(runDir, serviceName, config),
51 + epoch: time.Now(),
52 + }
53 +}
54 +
55 +// Refresh drives the L2 client and requests a fresh snapshot.
56 +func (c *Cache) Refresh() bool {
57 + c.client.Refresh()
58 +
59 + view, err := c.client.CallSnapshot()
60 + if err != nil {
61 + c.refreshFailureCount++
62 + return false
63 + }
64 +
65 + newItems := make([]CacheItem, 0, view.ItemCount)
66 + for i := uint32(0); i < view.ItemCount; i++ {
67 + iv, ierr := view.Item(i)
68 + if ierr != nil {
69 + c.refreshFailureCount++
70 + return false
71 + }
72 + newItems = append(newItems, CacheItem{
73 + Hash: iv.Hash,
74 + Options: iv.Options,
75 + Enabled: iv.Enabled,
76 + Name: iv.Name.String(),
77 + Path: iv.Path.String(),
78 + })
79 + }
80 +
81 + // Rebuild open-addressing lookup table.
82 + var buckets []cacheBucket
83 + if len(newItems) > 0 {
84 + bcount := nextPowerOf2U32(uint32(len(newItems)) * 2)
85 + buckets = make([]cacheBucket, bcount)
86 + mask := bcount - 1
87 + for i := range newItems {
88 + slot := (newItems[i].Hash ^ cacheHashName(newItems[i].Name)) & mask
89 + for buckets[slot].used {
90 + slot = (slot + 1) & mask
91 + }
92 + buckets[slot].index = i
93 + buckets[slot].used = true
94 + }
95 + }
96 +
97 + c.items = newItems
98 + c.buckets = buckets
99 + c.systemdEnabled = view.SystemdEnabled
100 + c.generation = view.Generation
101 + c.populated = true
102 + c.refreshSuccessCount++
103 + c.lastRefreshTs = time.Since(c.epoch).Milliseconds()
104 +
105 + return true
106 +}
107 +
108 +// Ready returns true if at least one successful refresh has occurred.
109 +func (c *Cache) Ready() bool {
110 + return c.populated
111 +}
112 +
113 +// Lookup finds a cached item by hash + name. O(1) via open-addressing hash
114 +// table. No I/O.
115 +func (c *Cache) Lookup(hash uint32, name string) (CacheItem, bool) {
116 + if !c.populated {
117 + return CacheItem{}, false
118 + }
119 + if len(c.buckets) > 0 {
120 + mask := uint32(len(c.buckets) - 1)
121 + slot := (hash ^ cacheHashName(name)) & mask
122 + for c.buckets[slot].used {
123 + item := c.items[c.buckets[slot].index]
124 + if item.Hash == hash && item.Name == name {
125 + return item, true
126 + }
127 + slot = (slot + 1) & mask
128 + }
129 + return CacheItem{}, false
130 + }
131 + for i := range c.items {
132 + if c.items[i].Hash == hash && c.items[i].Name == name {
133 + return c.items[i], true
134 + }
135 + }
136 + return CacheItem{}, false
137 +}
138 +
139 +// Status returns a diagnostic snapshot for the L3 cache.
140 +func (c *Cache) Status() CacheStatus {
141 + return CacheStatus{
142 + Populated: c.populated,
143 + ItemCount: uint32(len(c.items)),
144 + SystemdEnabled: c.systemdEnabled,
145 + Generation: c.generation,
146 + RefreshSuccessCount: c.refreshSuccessCount,
147 + RefreshFailureCount: c.refreshFailureCount,
148 + ConnectionState: c.client.state,
149 + LastRefreshTs: c.lastRefreshTs,
150 + }
151 +}
152 +
153 +// Close frees all cached items and closes the L2 client.
154 +func (c *Cache) Close() {
155 + c.items = nil
156 + c.buckets = nil
157 + c.populated = false
158 + c.client.Close()
159 +}
src/go/pkg/netipc/service/raw/client.go new
+1149
@@ -0,0 +1,1149 @@
1 +//go:build unix
2 +
3 +package raw
4 +
5 +import (
6 + "encoding/binary"
7 + "errors"
8 + "sync"
9 + "sync/atomic"
10 + "syscall"
11 + "time"
12 + "unsafe"
13 +
14 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
15 + "github.com/netdata/netdata/go/plugins/pkg/netipc/transport/posix"
16 +)
17 +
18 +const clientShmAttachRetryInterval = 5 * time.Millisecond
19 +const clientShmAttachRetryTimeout = 5 * time.Second
20 +
21 +// ---------------------------------------------------------------------------
22 +// Client context
23 +// ---------------------------------------------------------------------------
24 +
25 +// Client is an internal L2 client context bound to one service kind.
26 +// It manages connection lifecycle and provides typed blocking calls with
27 +// at-least-once retry semantics.
28 +type Client struct {
29 + state ClientState
30 + runDir string
31 + serviceName string
32 + expectedMethodCode uint16
33 + config posix.ClientConfig
34 +
35 + // Connection (managed internally)
36 + session *posix.Session
37 + shm *posix.ShmContext
38 +
39 + // Reusable scratch buffers owned by the client for hot request paths.
40 + requestBuf []byte
41 + sendBuf []byte
42 + transportBuf []byte
43 +
44 + // Stats
45 + connectCount uint32
46 + reconnectCount uint32
47 + callCount uint32
48 + errorCount uint32
49 +}
50 +
51 +func newClient(runDir, serviceName string, config posix.ClientConfig, expectedMethodCode uint16) *Client {
52 + return &Client{
53 + state: StateDisconnected,
54 + runDir: runDir,
55 + serviceName: serviceName,
56 + expectedMethodCode: expectedMethodCode,
57 + config: config,
58 + }
59 +}
60 +
61 +// NewSnapshotClient creates a raw client bound to the cgroups-snapshot service kind.
62 +func NewSnapshotClient(runDir, serviceName string, config posix.ClientConfig) *Client {
63 + return newClient(runDir, serviceName, config, protocol.MethodCgroupsSnapshot)
64 +}
65 +
66 +// NewIncrementClient creates a raw client bound to the increment service kind.
67 +func NewIncrementClient(runDir, serviceName string, config posix.ClientConfig) *Client {
68 + return newClient(runDir, serviceName, config, protocol.MethodIncrement)
69 +}
70 +
71 +// NewStringReverseClient creates a raw client bound to the string-reverse service kind.
72 +func NewStringReverseClient(runDir, serviceName string, config posix.ClientConfig) *Client {
73 + return newClient(runDir, serviceName, config, protocol.MethodStringReverse)
74 +}
75 +
76 +func nextPowerOf2U32(n uint32) uint32 {
77 + if n < 16 {
78 + return 16
79 + }
80 + n--
81 + n |= n >> 1
82 + n |= n >> 2
83 + n |= n >> 4
84 + n |= n >> 8
85 + n |= n >> 16
86 + return n + 1
87 +}
88 +
89 +func (c *Client) validateMethod(methodCode uint16) error {
90 + if c.expectedMethodCode != methodCode {
91 + return protocol.ErrBadLayout
92 + }
93 + return nil
94 +}
95 +
96 +// Refresh attempts connect if DISCONNECTED/NOT_FOUND, reconnect if BROKEN.
97 +// Returns true if the state changed.
98 +func (c *Client) Refresh() bool {
99 + oldState := c.state
100 +
101 + switch c.state {
102 + case StateDisconnected, StateNotFound:
103 + c.state = StateConnecting
104 + c.state = c.tryConnect()
105 + if c.state == StateReady {
106 + c.connectCount++
107 + }
108 +
109 + case StateBroken:
110 + c.disconnect()
111 + c.state = StateConnecting
112 + c.state = c.tryConnect()
113 + if c.state == StateReady {
114 + c.reconnectCount++
115 + }
116 +
117 + case StateReady, StateConnecting, StateAuthFailed, StateIncompatible:
118 + // No action needed
119 + }
120 +
121 + return c.state != oldState
122 +}
123 +
124 +// Ready returns true only if the client is in the READY state.
125 +// Cheap cached boolean, no I/O.
126 +func (c *Client) Ready() bool {
127 + return c.state == StateReady
128 +}
129 +
130 +// Status returns a diagnostic counters snapshot.
131 +func (c *Client) Status() ClientStatus {
132 + return ClientStatus{
133 + State: c.state,
134 + ConnectCount: c.connectCount,
135 + ReconnectCount: c.reconnectCount,
136 + CallCount: c.callCount,
137 + ErrorCount: c.errorCount,
138 + }
139 +}
140 +
141 +func (c *Client) sessionMaxRequestPayloadBytes() uint32 {
142 + if c.session != nil {
143 + return c.session.MaxRequestPayloadBytes
144 + }
145 + return c.config.MaxRequestPayloadBytes
146 +}
147 +
148 +func (c *Client) sessionMaxResponsePayloadBytes() uint32 {
149 + if c.session != nil {
150 + return c.session.MaxResponsePayloadBytes
151 + }
152 + return c.config.MaxResponsePayloadBytes
153 +}
154 +
155 +func (c *Client) noteRequestCapacity(payloadLen uint32) {
156 + grown := nextPowerOf2U32(payloadLen)
157 + if grown > protocol.MaxPayloadCap {
158 + grown = protocol.MaxPayloadCap
159 + }
160 + if grown > c.config.MaxRequestPayloadBytes {
161 + c.config.MaxRequestPayloadBytes = grown
162 + }
163 +}
164 +
165 +func (c *Client) noteResponseCapacity(payloadLen uint32) {
166 + grown := nextPowerOf2U32(payloadLen)
167 + if grown > protocol.MaxPayloadCap {
168 + grown = protocol.MaxPayloadCap
169 + }
170 + if grown > c.config.MaxResponsePayloadBytes {
171 + c.config.MaxResponsePayloadBytes = grown
172 + }
173 +}
174 +
175 +// callWithRetry manages reconnect-driven recovery for a typed call.
176 +// Ordinary failures retry once. Overflow-driven resize recovery may
177 +// reconnect more than once while negotiated capacities grow.
178 +func (c *Client) callWithRetry(attempt func() error) error {
179 + // Fail fast if not READY
180 + if c.state != StateReady {
181 + c.errorCount++
182 + return protocol.ErrBadLayout
183 + }
184 +
185 + // Cap overflow-driven retries: payloads grow by powers of 2, so 8
186 + // retries allows ~256x growth from the initial negotiated size.
187 + overflowRetries := 0
188 + for {
189 + prevReq := c.sessionMaxRequestPayloadBytes()
190 + prevResp := c.sessionMaxResponsePayloadBytes()
191 + prevCfgReq := c.config.MaxRequestPayloadBytes
192 + prevCfgResp := c.config.MaxResponsePayloadBytes
193 +
194 + firstErr := attempt()
195 + if firstErr == nil {
196 + c.callCount++
197 + return nil
198 + }
199 +
200 + if !errors.Is(firstErr, protocol.ErrOverflow) {
201 + c.disconnect()
202 + c.state = StateBroken
203 +
204 + c.state = c.tryConnect()
205 + if c.state != StateReady {
206 + c.errorCount++
207 + return firstErr
208 + }
209 + c.reconnectCount++
210 +
211 + retryErr := attempt()
212 + if retryErr == nil {
213 + c.callCount++
214 + return nil
215 + }
216 +
217 + c.disconnect()
218 + c.state = StateBroken
219 + c.errorCount++
220 + return retryErr
221 + }
222 +
223 + c.disconnect()
224 + c.state = StateBroken
225 +
226 + c.state = c.tryConnect()
227 + if c.state != StateReady {
228 + c.errorCount++
229 + return firstErr
230 + }
231 + c.reconnectCount++
232 +
233 + if c.sessionMaxRequestPayloadBytes() <= prevReq &&
234 + c.sessionMaxResponsePayloadBytes() <= prevResp &&
235 + c.config.MaxRequestPayloadBytes <= prevCfgReq &&
236 + c.config.MaxResponsePayloadBytes <= prevCfgResp {
237 + c.disconnect()
238 + c.state = StateBroken
239 + c.errorCount++
240 + return firstErr
241 + }
242 +
243 + overflowRetries++
244 + if overflowRetries >= 8 {
245 + c.disconnect()
246 + c.state = StateBroken
247 + c.errorCount++
248 + return firstErr
249 + }
250 + }
251 +}
252 +
253 +// doRawCall sends a request and receives/validates the response envelope.
254 +// Returns the validated response header and a borrowed payload view.
255 +func (c *Client) doRawCall(methodCode uint16, reqPayload []byte) (protocol.Header, []byte, error) {
256 + hdr := protocol.Header{
257 + Kind: protocol.KindRequest,
258 + Code: methodCode,
259 + Flags: 0,
260 + ItemCount: 1,
261 + MessageID: uint64(c.callCount) + 1,
262 + TransportStatus: protocol.StatusOK,
263 + }
264 +
265 + if err := c.transportSend(&hdr, reqPayload); err != nil {
266 + return protocol.Header{}, nil, err
267 + }
268 +
269 + respHdr, payload, err := c.transportReceive()
270 + if err != nil {
271 + return protocol.Header{}, nil, err
272 + }
273 +
274 + if respHdr.Kind != protocol.KindResponse {
275 + return protocol.Header{}, nil, protocol.ErrBadKind
276 + }
277 + if respHdr.Code != methodCode {
278 + return protocol.Header{}, nil, protocol.ErrBadLayout
279 + }
280 + if respHdr.MessageID != hdr.MessageID {
281 + return protocol.Header{}, nil, protocol.ErrBadLayout
282 + }
283 + switch respHdr.TransportStatus {
284 + case protocol.StatusOK:
285 + case protocol.StatusLimitExceeded:
286 + if current := c.sessionMaxResponsePayloadBytes(); current > 0 {
287 + if current >= ^uint32(0)/2 {
288 + c.noteResponseCapacity(^uint32(0))
289 + } else {
290 + c.noteResponseCapacity(current * 2)
291 + }
292 + }
293 + return protocol.Header{}, nil, protocol.ErrOverflow
294 + default:
295 + return protocol.Header{}, nil, protocol.ErrBadLayout
296 + }
297 +
298 + return respHdr, payload, nil
299 +}
300 +
301 +// CallSnapshot performs a blocking typed cgroups snapshot call.
302 +// The returned view is valid until the next typed call on this client.
303 +func (c *Client) CallSnapshot() (*protocol.CgroupsResponseView, error) {
304 + if err := c.validateMethod(protocol.MethodCgroupsSnapshot); err != nil {
305 + return nil, err
306 + }
307 +
308 + var result *protocol.CgroupsResponseView
309 +
310 + err := c.callWithRetry(func() error {
311 + req := protocol.CgroupsRequest{LayoutVersion: 1, Flags: 0}
312 + var reqBuf [4]byte
313 + if req.Encode(reqBuf[:]) == 0 {
314 + return protocol.ErrTruncated
315 + }
316 +
317 + _, payload, rerr := c.doRawCall(protocol.MethodCgroupsSnapshot, reqBuf[:])
318 + if rerr != nil {
319 + return rerr
320 + }
321 +
322 + view, derr := protocol.DecodeCgroupsResponse(payload)
323 + if derr != nil {
324 + return derr
325 + }
326 + result = &view
327 + return nil
328 + })
329 + if err != nil {
330 + return nil, err
331 + }
332 + return result, nil
333 +}
334 +
335 +// CallIncrement performs a blocking INCREMENT call.
336 +// Sends requestValue, returns the server's response value.
337 +func (c *Client) CallIncrement(requestValue uint64) (uint64, error) {
338 + if err := c.validateMethod(protocol.MethodIncrement); err != nil {
339 + return 0, err
340 + }
341 +
342 + var result uint64
343 +
344 + err := c.callWithRetry(func() error {
345 + var reqBuf [protocol.IncrementPayloadSize]byte
346 + if protocol.IncrementEncode(requestValue, reqBuf[:]) == 0 {
347 + return protocol.ErrTruncated
348 + }
349 +
350 + _, payload, rerr := c.doRawCall(protocol.MethodIncrement, reqBuf[:])
351 + if rerr != nil {
352 + return rerr
353 + }
354 +
355 + val, derr := protocol.IncrementDecode(payload)
356 + if derr != nil {
357 + return derr
358 + }
359 + result = val
360 + return nil
361 + })
362 + return result, err
363 +}
364 +
365 +// CallStringReverse performs a blocking STRING_REVERSE call.
366 +// The returned view is valid until the next typed call on this client.
367 +func (c *Client) CallStringReverse(requestStr string) (*protocol.StringReverseView, error) {
368 + if err := c.validateMethod(protocol.MethodStringReverse); err != nil {
369 + return nil, err
370 + }
371 +
372 + var result *protocol.StringReverseView
373 +
374 + err := c.callWithRetry(func() error {
375 + reqBuf := ensureClientScratch(&c.requestBuf, protocol.StringReverseHdrSize+len(requestStr)+1)
376 + if protocol.StringReverseEncode(requestStr, reqBuf) == 0 {
377 + return protocol.ErrTruncated
378 + }
379 +
380 + _, payload, rerr := c.doRawCall(protocol.MethodStringReverse, reqBuf)
381 + if rerr != nil {
382 + return rerr
383 + }
384 +
385 + view, derr := protocol.StringReverseDecode(payload)
386 + if derr != nil {
387 + return derr
388 + }
389 + result = &view
390 + return nil
391 + })
392 + if err != nil {
393 + return nil, err
394 + }
395 + return result, nil
396 +}
397 +
398 +// CallIncrementBatch performs a blocking batch INCREMENT call.
399 +// Sends multiple values, returns the server's response values.
400 +func (c *Client) CallIncrementBatch(values []uint64) ([]uint64, error) {
401 + if err := c.validateMethod(protocol.MethodIncrement); err != nil {
402 + return nil, err
403 + }
404 +
405 + if len(values) == 0 {
406 + return nil, nil
407 + }
408 +
409 + var results []uint64
410 + itemCount := uint32(len(values))
411 +
412 + err := c.callWithRetry(func() error {
413 + // Build batch request payload
414 + batchBufSize := protocol.Align8(int(itemCount)*8) + int(itemCount)*protocol.IncrementPayloadSize + int(itemCount)*protocol.Alignment
415 + batchBuf := ensureClientScratch(&c.requestBuf, batchBufSize)
416 + bb := protocol.NewBatchBuilder(batchBuf, itemCount)
417 +
418 + for _, v := range values {
419 + var item [protocol.IncrementPayloadSize]byte
420 + if protocol.IncrementEncode(v, item[:]) == 0 {
421 + return protocol.ErrTruncated
422 + }
423 + if err := bb.Add(item[:]); err != nil {
424 + return err
425 + }
426 + }
427 +
428 + totalPayloadLen, _ := bb.Finish()
429 + reqPayload := batchBuf[:totalPayloadLen]
430 +
431 + // Build and send header with batch flags
432 + hdr := protocol.Header{
433 + Kind: protocol.KindRequest,
434 + Code: protocol.MethodIncrement,
435 + Flags: protocol.FlagBatch,
436 + ItemCount: itemCount,
437 + MessageID: uint64(c.callCount) + 1,
438 + TransportStatus: protocol.StatusOK,
439 + }
440 +
441 + if err := c.transportSend(&hdr, reqPayload); err != nil {
442 + return err
443 + }
444 +
445 + // Receive response
446 + respHdr, respPayload, err := c.transportReceive()
447 + if err != nil {
448 + return err
449 + }
450 +
451 + if respHdr.Kind != protocol.KindResponse {
452 + return protocol.ErrBadKind
453 + }
454 + if respHdr.Code != protocol.MethodIncrement {
455 + return protocol.ErrBadLayout
456 + }
457 + if respHdr.MessageID != hdr.MessageID {
458 + return protocol.ErrBadLayout
459 + }
460 + switch respHdr.TransportStatus {
461 + case protocol.StatusOK:
462 + case protocol.StatusLimitExceeded:
463 + if current := c.sessionMaxResponsePayloadBytes(); current > 0 {
464 + if current >= ^uint32(0)/2 {
465 + c.noteResponseCapacity(^uint32(0))
466 + } else {
467 + c.noteResponseCapacity(current * 2)
468 + }
469 + }
470 + return protocol.ErrOverflow
471 + default:
472 + return protocol.ErrBadLayout
473 + }
474 + if respHdr.Flags&protocol.FlagBatch == 0 || respHdr.ItemCount != itemCount {
475 + return protocol.ErrBadItemCount
476 + }
477 +
478 + // Extract each response item
479 + out := make([]uint64, itemCount)
480 + for i := uint32(0); i < itemCount; i++ {
481 + itemData, gerr := protocol.BatchItemGet(respPayload, itemCount, i)
482 + if gerr != nil {
483 + return gerr
484 + }
485 + val, derr := protocol.IncrementDecode(itemData)
486 + if derr != nil {
487 + return derr
488 + }
489 + out[i] = val
490 + }
491 + results = out
492 + return nil
493 + })
494 + return results, err
495 +}
496 +
497 +// Close tears down the connection and releases resources.
498 +func (c *Client) Close() {
499 + c.disconnect()
500 + c.state = StateDisconnected
501 +}
502 +
503 +// ------------------------------------------------------------------
504 +// Internal helpers
505 +// ------------------------------------------------------------------
506 +
507 +func (c *Client) disconnect() {
508 + if c.shm != nil {
509 + c.shm.ShmClose()
510 + c.shm = nil
511 + }
512 + if c.session != nil {
513 + c.session.Close()
514 + c.session = nil
515 + }
516 +}
517 +
518 +func (c *Client) tryConnect() ClientState {
519 + session, err := posix.Connect(c.runDir, c.serviceName, &c.config)
520 + if err != nil {
521 + switch {
522 + case isConnectError(err):
523 + return StateNotFound
524 + case isAuthError(err):
525 + return StateAuthFailed
526 + case isIncompatibleError(err):
527 + return StateIncompatible
528 + default:
529 + return StateDisconnected
530 + }
531 + }
532 +
533 + // SHM upgrade if negotiated
534 + if session.SelectedProfile == protocol.ProfileSHMHybrid ||
535 + session.SelectedProfile == protocol.ProfileSHMFutex {
536 + // Retry attach: the server prepared SHM before handshake, but the
537 + // client may still race slightly with the peer exposing the region.
538 + deadline := time.Now().Add(clientShmAttachRetryTimeout)
539 + for {
540 + shm, serr := posix.ShmClientAttach(c.runDir, c.serviceName, session.SessionID)
541 + if serr == nil {
542 + c.shm = shm
543 + break
544 + }
545 + if !time.Now().Before(deadline) {
546 + break
547 + }
548 + time.Sleep(clientShmAttachRetryInterval)
549 + }
550 + if c.shm == nil {
551 + // SHM attach failed after negotiation. Close that session,
552 + // blacklist SHM for this client context, and retry baseline.
553 + session.Close()
554 + c.config.SupportedProfiles &^= posixShmProfiles
555 + c.config.PreferredProfiles &^= posixShmProfiles
556 + if c.config.SupportedProfiles == 0 {
557 + return StateDisconnected
558 + }
559 + return c.tryConnect()
560 + }
561 + }
562 +
563 + c.session = session
564 + return StateReady
565 +}
566 +
567 +func (c *Client) transportSend(hdr *protocol.Header, payload []byte) error {
568 + if c.shm != nil {
569 + if len(payload) > int(c.sessionMaxRequestPayloadBytes()) {
570 + c.noteRequestCapacity(uint32(len(payload)))
571 + return protocol.ErrOverflow
572 + }
573 +
574 + msgLen := protocol.HeaderSize + len(payload)
575 + msg := ensureClientScratch(&c.sendBuf, msgLen)
576 +
577 + hdr.Magic = protocol.MagicMsg
578 + hdr.Version = protocol.Version
579 + hdr.HeaderLen = protocol.HeaderLen
580 + hdr.PayloadLen = uint32(len(payload))
581 +
582 + hdr.Encode(msg[:protocol.HeaderSize])
583 + if len(payload) > 0 {
584 + copy(msg[protocol.HeaderSize:], payload)
585 + }
586 +
587 + if err := c.shm.ShmSend(msg[:msgLen]); err != nil {
588 + if errors.Is(err, posix.ErrShmMsgTooLarge) {
589 + c.noteRequestCapacity(uint32(len(payload)))
590 + return protocol.ErrOverflow
591 + }
592 + return protocol.ErrTruncated
593 + }
594 + return nil
595 + }
596 +
597 + // UDS path
598 + if c.session == nil {
599 + return protocol.ErrTruncated
600 + }
601 + if err := c.session.Send(hdr, payload); err != nil {
602 + if errors.Is(err, posix.ErrLimitExceeded) {
603 + c.noteRequestCapacity(uint32(len(payload)))
604 + return protocol.ErrOverflow
605 + }
606 + return protocol.ErrTruncated
607 + }
608 + return nil
609 +}
610 +
611 +func (c *Client) transportReceive() (protocol.Header, []byte, error) {
612 + scratch := ensureClientScratch(&c.transportBuf, c.maxReceiveMessageBytes())
613 +
614 + if c.shm != nil {
615 + mlen, err := c.shm.ShmReceive(scratch, 30000)
616 + if err != nil {
617 + return protocol.Header{}, nil, protocol.ErrTruncated
618 + }
619 + if mlen < protocol.HeaderSize {
620 + return protocol.Header{}, nil, protocol.ErrTruncated
621 + }
622 +
623 + hdr, err := protocol.DecodeHeader(scratch[:mlen])
624 + if err != nil {
625 + return protocol.Header{}, nil, err
626 + }
627 + return hdr, scratch[protocol.HeaderSize:mlen], nil
628 + }
629 +
630 + // UDS path: receive returns (Header, payload, error)
631 + if c.session == nil {
632 + return protocol.Header{}, nil, protocol.ErrTruncated
633 + }
634 +
635 + hdr, payload, err := c.session.Receive(scratch)
636 + if err != nil {
637 + return protocol.Header{}, nil, protocol.ErrTruncated
638 + }
639 + return hdr, payload, nil
640 +}
641 +
642 +func (c *Client) maxReceiveMessageBytes() int {
643 + maxPayload := c.config.MaxResponsePayloadBytes
644 + if c.session != nil && c.session.MaxResponsePayloadBytes > 0 {
645 + maxPayload = c.session.MaxResponsePayloadBytes
646 + }
647 + if maxPayload == 0 {
648 + maxPayload = cacheResponseBufSize
649 + }
650 + return protocol.HeaderSize + int(maxPayload)
651 +}
652 +
653 +// Error classification helpers
654 +func isConnectError(err error) bool {
655 + return errors.Is(err, posix.ErrConnect) || errors.Is(err, posix.ErrSocket)
656 +}
657 +
658 +func isAuthError(err error) bool {
659 + return errors.Is(err, posix.ErrAuthFailed)
660 +}
661 +
662 +func isIncompatibleError(err error) bool {
663 + return errors.Is(err, posix.ErrNoProfile) || errors.Is(err, posix.ErrIncompatible)
664 +}
665 +
666 +// ---------------------------------------------------------------------------
667 +// Managed server
668 +// ---------------------------------------------------------------------------
669 +
670 +// Server is an internal managed server bound to one expected request kind.
671 +// Supports multiple concurrent client sessions up to workerCount.
672 +type Server struct {
673 + runDir string
674 + serviceName string
675 + config posix.ServerConfig
676 + expectedMethodCode uint16
677 + handler DispatchHandler
678 + running atomic.Bool
679 + learnedRequestPayloadBytes atomic.Uint32
680 + learnedResponsePayloadBytes atomic.Uint32
681 + nextSessionID atomic.Uint64
682 + workerCount int
683 + wg sync.WaitGroup
684 +}
685 +
686 +// NewServer creates a new managed server. workerCount limits the
687 +// maximum number of concurrent client sessions (default 1 if <= 0).
688 +func NewServer(
689 + runDir, serviceName string,
690 + config posix.ServerConfig,
691 + expectedMethodCode uint16,
692 + handler DispatchHandler,
693 +) *Server {
694 + return NewServerWithWorkers(runDir, serviceName, config, expectedMethodCode, handler, 8)
695 +}
696 +
697 +// NewServerWithWorkers creates a server with an explicit worker count limit.
698 +func NewServerWithWorkers(
699 + runDir, serviceName string,
700 + config posix.ServerConfig,
701 + expectedMethodCode uint16,
702 + handler DispatchHandler,
703 + workerCount int,
704 +) *Server {
705 + if workerCount < 1 {
706 + workerCount = 1
707 + }
708 + learnedRequest := config.MaxRequestPayloadBytes
709 + if learnedRequest == 0 {
710 + learnedRequest = protocol.MaxPayloadDefault
711 + }
712 + learnedResponse := config.MaxResponsePayloadBytes
713 + if learnedResponse == 0 {
714 + learnedResponse = protocol.MaxPayloadDefault
715 + }
716 + s := &Server{
717 + runDir: runDir,
718 + serviceName: serviceName,
719 + config: config,
720 + expectedMethodCode: expectedMethodCode,
721 + handler: handler,
722 + workerCount: workerCount,
723 + }
724 + s.learnedRequestPayloadBytes.Store(learnedRequest)
725 + s.learnedResponsePayloadBytes.Store(learnedResponse)
726 + // Session ids are 1-based; prepareAcceptConfig() allocates with Add(1).
727 + s.nextSessionID.Store(0)
728 + return s
729 +}
730 +
731 +func (s *Server) dispatchSingle(methodCode uint16, request []byte, responseBuf []byte) (int, error) {
732 + if methodCode != s.expectedMethodCode || s.handler == nil {
733 + return 0, errHandlerFailed
734 + }
735 +
736 + return s.handler(request, responseBuf)
737 +}
738 +
739 +func (s *Server) methodSupported(methodCode uint16) bool {
740 + return s.handler != nil && methodCode == s.expectedMethodCode
741 +}
742 +
743 +func serverNotePayloadCapacity(target *atomic.Uint32, payloadLen uint32) {
744 + grown := nextPowerOf2U32(payloadLen)
745 + for {
746 + current := target.Load()
747 + if grown <= current {
748 + return
749 + }
750 + if target.CompareAndSwap(current, grown) {
751 + return
752 + }
753 + }
754 +}
755 +
756 +const posixShmProfiles = protocol.ProfileSHMHybrid | protocol.ProfileSHMFutex
757 +
758 +func (s *Server) prepareAcceptConfig() (uint64, posix.ServerConfig, *posix.ShmContext, bool) {
759 + sessionID := s.nextSessionID.Add(1)
760 + cfg := s.config
761 + cfg.MaxRequestPayloadBytes = s.learnedRequestPayloadBytes.Load()
762 + cfg.MaxResponsePayloadBytes = s.learnedResponsePayloadBytes.Load()
763 +
764 + if cfg.SupportedProfiles&posixShmProfiles == 0 {
765 + return sessionID, cfg, nil, true
766 + }
767 +
768 + shm, err := posix.ShmServerCreate(
769 + s.runDir, s.serviceName, sessionID,
770 + cfg.MaxRequestPayloadBytes+uint32(protocol.HeaderSize),
771 + cfg.MaxResponsePayloadBytes+uint32(protocol.HeaderSize),
772 + )
773 + if err == nil {
774 + return sessionID, cfg, shm, true
775 + }
776 +
777 + cfg.SupportedProfiles &^= posixShmProfiles
778 + cfg.PreferredProfiles &^= posixShmProfiles
779 + if cfg.SupportedProfiles == 0 {
780 + return sessionID, cfg, nil, false
781 + }
782 +
783 + return sessionID, cfg, nil, true
784 +}
785 +
786 +// Run starts the acceptor loop. Blocking. Accepts clients, spawns a
787 +// goroutine per session (up to workerCount concurrently).
788 +// Returns when Stop() is called or on fatal error.
789 +func (s *Server) Run() error {
790 + posix.ShmCleanupStale(s.runDir, s.serviceName)
791 +
792 + listener, err := posix.Listen(s.runDir, s.serviceName, s.config)
793 + if err != nil {
794 + return err
795 + }
796 + defer listener.Close()
797 +
798 + s.running.Store(true)
799 +
800 + /* Semaphore channel limits concurrent sessions */
801 + sem := make(chan struct{}, s.workerCount)
802 +
803 + for s.running.Load() {
804 + // Poll the listener fd before blocking on accept
805 + ready := pollFd(listener.Fd(), serverPollTimeoutMs)
806 + if ready < 0 {
807 + break
808 + }
809 + if ready == 0 {
810 + continue
811 + }
812 +
813 + sessionID, acceptCfg, precreatedShm, ok := s.prepareAcceptConfig()
814 + if !ok {
815 + time.Sleep(10 * time.Millisecond)
816 + continue
817 + }
818 +
819 + session, err := listener.AcceptWithConfig(sessionID, acceptCfg)
820 + if err != nil {
821 + if precreatedShm != nil {
822 + precreatedShm.ShmDestroy()
823 + }
824 + if !s.running.Load() {
825 + break
826 + }
827 + time.Sleep(10 * time.Millisecond)
828 + continue
829 + }
830 +
831 + // Try to acquire a worker slot (non-blocking check)
832 + select {
833 + case sem <- struct{}{}:
834 + // Got a slot
835 + default:
836 + // At capacity: reject client
837 + if precreatedShm != nil {
838 + precreatedShm.ShmDestroy()
839 + }
840 + session.Close()
841 + continue
842 + }
843 +
844 + var shm *posix.ShmContext
845 + if session.SelectedProfile == protocol.ProfileSHMHybrid ||
846 + session.SelectedProfile == protocol.ProfileSHMFutex {
847 + if precreatedShm == nil {
848 + session.Close()
849 + <-sem
850 + continue
851 + }
852 + shm = precreatedShm
853 + } else if precreatedShm != nil {
854 + precreatedShm.ShmDestroy()
855 + }
856 +
857 + // Handle this session in a goroutine
858 + s.wg.Add(1)
859 + go func(sess *posix.Session, shmCtx *posix.ShmContext) {
860 + defer func() {
861 + if r := recover(); r != nil {
862 + // Session handler panicked; log but don't crash the server
863 + }
864 + <-sem // release worker slot
865 + s.wg.Done()
866 + }()
867 + s.handleSession(sess, shmCtx)
868 + }(session, shm)
869 + }
870 +
871 + // Wait for all active session goroutines to finish
872 + s.wg.Wait()
873 +
874 + return nil
875 +}
876 +
877 +// Stop signals the server to stop.
878 +func (s *Server) Stop() {
879 + s.running.Store(false)
880 +}
881 +
882 +func (s *Server) handleSession(session *posix.Session, shm *posix.ShmContext) {
883 + recvBuf := make([]byte, protocol.HeaderSize+int(session.MaxRequestPayloadBytes))
884 + respBuf := make([]byte, int(session.MaxResponsePayloadBytes))
885 + itemRespBuf := make([]byte, int(session.MaxResponsePayloadBytes))
886 + msgBuf := make([]byte, int(session.MaxResponsePayloadBytes)+protocol.HeaderSize)
887 +
888 + defer func() {
889 + if shm != nil {
890 + shm.ShmDestroy()
891 + }
892 + session.Close()
893 + }()
894 +
895 + for s.running.Load() {
896 + var hdr protocol.Header
897 + var payload []byte
898 +
899 + if shm != nil {
900 + mlen, err := shm.ShmReceive(recvBuf, serverPollTimeoutMs)
901 + if err != nil {
902 + if err == posix.ErrShmTimeout {
903 + continue
904 + }
905 + return
906 + }
907 + if mlen < protocol.HeaderSize {
908 + return
909 + }
910 + h, err := protocol.DecodeHeader(recvBuf[:mlen])
911 + if err != nil {
912 + return
913 + }
914 + hdr = h
915 + payload = recvBuf[protocol.HeaderSize:mlen]
916 + } else {
917 + // Poll the session fd before blocking on receive
918 + ready := pollFd(session.Fd(), serverPollTimeoutMs)
919 + if ready < 0 {
920 + return
921 + }
922 + if ready == 0 {
923 + continue
924 + }
925 +
926 + h, p, err := session.Receive(recvBuf)
927 + if err != nil {
928 + return
929 + }
930 + hdr = h
931 + payload = p
932 + }
933 +
934 + // Protocol violation: unexpected message kind terminates session
935 + if hdr.Kind != protocol.KindRequest {
936 + return
937 + }
938 +
939 + if len(payload) <= int(^uint32(0)) {
940 + serverNotePayloadCapacity(&s.learnedRequestPayloadBytes, uint32(len(payload)))
941 + }
942 +
943 + if !s.methodSupported(hdr.Code) {
944 + respHdr := protocol.Header{
945 + Kind: protocol.KindResponse,
946 + Code: hdr.Code,
947 + MessageID: hdr.MessageID,
948 + TransportStatus: protocol.StatusUnsupported,
949 + ItemCount: 1,
950 + }
951 +
952 + if shm != nil {
953 + if len(msgBuf) < protocol.HeaderSize {
954 + msgBuf = make([]byte, protocol.HeaderSize)
955 + }
956 + msg := msgBuf[:protocol.HeaderSize]
957 + respHdr.Magic = protocol.MagicMsg
958 + respHdr.Version = protocol.Version
959 + respHdr.HeaderLen = protocol.HeaderLen
960 + respHdr.PayloadLen = 0
961 + respHdr.Encode(msg[:protocol.HeaderSize])
962 + if err := shm.ShmSend(msg); err != nil {
963 + return
964 + }
965 + } else if err := session.Send(&respHdr, nil); err != nil {
966 + return
967 + }
968 + continue
969 + }
970 +
971 + // Dispatch: single-item or batch
972 + responseLen := 0
973 + isBatch := (hdr.Flags&protocol.FlagBatch != 0) && hdr.ItemCount >= 1
974 + var dispatchErr error
975 +
976 + if !isBatch {
977 + var derr error
978 + responseLen, derr = s.dispatchSingle(hdr.Code, payload, respBuf)
979 + if derr != nil {
980 + dispatchErr = derr
981 + responseLen = 0
982 + } else if responseLen < 0 || responseLen > len(respBuf) {
983 + dispatchErr = protocol.ErrOverflow
984 + responseLen = 0
985 + }
986 + } else {
987 + var bb protocol.BatchBuilder
988 + bb.Reset(respBuf, hdr.ItemCount)
989 +
990 + for i := uint32(0); i < hdr.ItemCount && dispatchErr == nil; i++ {
991 + itemData, gerr := protocol.BatchItemGet(payload, hdr.ItemCount, i)
992 + if gerr != nil {
993 + dispatchErr = gerr
994 + break
995 + }
996 +
997 + itemResultLen, derr := s.dispatchSingle(hdr.Code, itemData, itemRespBuf)
998 + if derr != nil {
999 + dispatchErr = derr
1000 + break
1001 + }
1002 + if itemResultLen < 0 || itemResultLen > len(itemRespBuf) {
1003 + dispatchErr = protocol.ErrOverflow
1004 + break
1005 + }
1006 +
1007 + if aerr := bb.Add(itemRespBuf[:itemResultLen]); aerr != nil {
1008 + dispatchErr = aerr
1009 + break
1010 + }
1011 + }
1012 +
1013 + if dispatchErr == nil {
1014 + responseLen, _ = bb.Finish()
1015 + }
1016 + }
1017 +
1018 + // Build response header
1019 + respHdr := protocol.Header{
1020 + Kind: protocol.KindResponse,
1021 + Code: hdr.Code,
1022 + MessageID: hdr.MessageID,
1023 + }
1024 +
1025 + if dispatchErr == nil {
1026 + if responseLen <= int(^uint32(0)) {
1027 + serverNotePayloadCapacity(&s.learnedResponsePayloadBytes, uint32(responseLen))
1028 + }
1029 + respHdr.TransportStatus = protocol.StatusOK
1030 + if isBatch {
1031 + respHdr.Flags = protocol.FlagBatch
1032 + respHdr.ItemCount = hdr.ItemCount
1033 + } else {
1034 + respHdr.ItemCount = 1
1035 + }
1036 + } else if errors.Is(dispatchErr, protocol.ErrOverflow) {
1037 + current := session.MaxResponsePayloadBytes
1038 + if current >= ^uint32(0)/2 {
1039 + serverNotePayloadCapacity(&s.learnedResponsePayloadBytes, ^uint32(0))
1040 + } else {
1041 + serverNotePayloadCapacity(&s.learnedResponsePayloadBytes, current*2)
1042 + }
1043 + respHdr.TransportStatus = protocol.StatusLimitExceeded
1044 + respHdr.ItemCount = 1
1045 + responseLen = 0
1046 + } else if errors.Is(dispatchErr, errHandlerFailed) {
1047 + respHdr.TransportStatus = protocol.StatusInternalError
1048 + respHdr.ItemCount = 1
1049 + responseLen = 0
1050 + } else {
1051 + respHdr.TransportStatus = protocol.StatusBadEnvelope
1052 + respHdr.ItemCount = 1
1053 + responseLen = 0
1054 + }
1055 +
1056 + // Send response via the active transport
1057 + if shm != nil {
1058 + msgLen := protocol.HeaderSize + responseLen
1059 + if len(msgBuf) < msgLen {
1060 + msgBuf = make([]byte, msgLen)
1061 + }
1062 + msg := msgBuf[:msgLen]
1063 +
1064 + respHdr.Magic = protocol.MagicMsg
1065 + respHdr.Version = protocol.Version
1066 + respHdr.HeaderLen = protocol.HeaderLen
1067 + respHdr.PayloadLen = uint32(responseLen)
1068 +
1069 + respHdr.Encode(msg[:protocol.HeaderSize])
1070 + if responseLen > 0 {
1071 + copy(msg[protocol.HeaderSize:], respBuf[:responseLen])
1072 + }
1073 +
1074 + if err := shm.ShmSend(msg); err != nil {
1075 + return
1076 + }
1077 + if respHdr.TransportStatus == protocol.StatusLimitExceeded {
1078 + return
1079 + }
1080 + } else {
1081 + if err := session.Send(&respHdr, respBuf[:responseLen]); err != nil {
1082 + return
1083 + }
1084 + if respHdr.TransportStatus == protocol.StatusLimitExceeded {
1085 + return
1086 + }
1087 + }
1088 + }
1089 +}
1090 +
1091 +// ---------------------------------------------------------------------------
1092 +// Internal: poll helper (raw syscall, pure Go, no cgo)
1093 +// ---------------------------------------------------------------------------
1094 +
1095 +// poll constants (not exported by Go's syscall package)
1096 +const (
1097 + _POLLIN = 0x0001
1098 + _POLLERR = 0x0008
1099 + _POLLHUP = 0x0010
1100 + _POLLNVAL = 0x0020
1101 +)
1102 +
1103 +// pollfd matches struct pollfd from <poll.h>.
1104 +type pollfd struct {
1105 + fd int32
1106 + events int16
1107 + revents int16
1108 +}
1109 +
1110 +// pollFd polls a file descriptor for readability with a timeout in ms.
1111 +// Returns: 1 = data ready, 0 = timeout, -1 = error/hangup.
1112 +func pollFd(fd int, timeoutMs int) int {
1113 + pfd := pollfd{
1114 + fd: int32(fd),
1115 + events: _POLLIN,
1116 + }
1117 +
1118 + r, _, errno := syscall.Syscall(
1119 + syscall.SYS_POLL,
1120 + uintptr(unsafe.Pointer(&pfd)),
1121 + 1,
1122 + uintptr(timeoutMs),
1123 + )
1124 +
1125 + n := int(r)
1126 + if n < 0 {
1127 + if errno == syscall.EINTR {
1128 + return 0
1129 + }
1130 + return -1
1131 + }
1132 +
1133 + if n == 0 {
1134 + return 0
1135 + }
1136 +
1137 + if pfd.revents&(_POLLERR|_POLLHUP|_POLLNVAL) != 0 {
1138 + return -1
1139 + }
1140 +
1141 + if pfd.revents&_POLLIN != 0 {
1142 + return 1
1143 + }
1144 +
1145 + return 0
1146 +}
1147 +
1148 +// Suppress unused import warnings.
1149 +var _ = binary.NativeEndian
src/go/pkg/netipc/service/raw/client_test.go new
+622
@@ -0,0 +1,622 @@
1 +//go:build unix
2 +
3 +package raw
4 +
5 +import (
6 + "os"
7 + "testing"
8 + "time"
9 +
10 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
11 + "github.com/netdata/netdata/go/plugins/pkg/netipc/transport/posix"
12 +)
13 +
14 +const (
15 + testRunDir = "/tmp/nipc_svc_go_test"
16 + authToken = uint64(0xDEADBEEFCAFEBABE)
17 + responseBufSize = 65536
18 +)
19 +
20 +func ensureRunDir() {
21 + os.MkdirAll(testRunDir, 0700)
22 +}
23 +
24 +func cleanupAll(service string) {
25 + os.Remove(testRunDir + "/" + service + ".sock")
26 + posix.ShmCleanupStale(testRunDir, service)
27 +}
28 +
29 +func waitUnixSocketReady(service string) {
30 + path := testRunDir + "/" + service + ".sock"
31 + deadline := time.Now().Add(2 * time.Second)
32 + for time.Now().Before(deadline) {
33 + if _, err := os.Stat(path); err == nil {
34 + return
35 + }
36 + time.Sleep(10 * time.Millisecond)
37 + }
38 +}
39 +
40 +func waitUnixServerReady(service string) {
41 + waitUnixSocketReady(service)
42 +
43 + cfg := testClientConfig()
44 + deadline := time.Now().Add(2 * time.Second)
45 + for time.Now().Before(deadline) {
46 + session, err := posix.Connect(testRunDir, service, &cfg)
47 + if err == nil {
48 + session.Close()
49 + return
50 + }
51 + time.Sleep(10 * time.Millisecond)
52 + }
53 +}
54 +
55 +func waitUnixRawSessionConnect(service string, cfg posix.ClientConfig) (*posix.Session, error) {
56 + deadline := time.Now().Add(2 * time.Second)
57 + var lastErr error
58 + for time.Now().Before(deadline) {
59 + session, err := posix.Connect(testRunDir, service, &cfg)
60 + if err == nil {
61 + return session, nil
62 + }
63 + lastErr = err
64 + time.Sleep(10 * time.Millisecond)
65 + }
66 + return nil, lastErr
67 +}
68 +
69 +func testServerConfig() posix.ServerConfig {
70 + return posix.ServerConfig{
71 + SupportedProfiles: protocol.ProfileBaseline,
72 + MaxRequestPayloadBytes: 4096,
73 + MaxRequestBatchItems: 1,
74 + MaxResponsePayloadBytes: responseBufSize,
75 + MaxResponseBatchItems: 1,
76 + AuthToken: authToken,
77 + Backlog: 4,
78 + }
79 +}
80 +
81 +func testClientConfig() posix.ClientConfig {
82 + return posix.ClientConfig{
83 + SupportedProfiles: protocol.ProfileBaseline,
84 + MaxRequestPayloadBytes: 4096,
85 + MaxRequestBatchItems: 1,
86 + MaxResponsePayloadBytes: responseBufSize,
87 + MaxResponseBatchItems: 1,
88 + AuthToken: authToken,
89 + }
90 +}
91 +
92 +// testCgroupsHandler builds a snapshot with 3 test items.
93 +func testCgroupsHandler(request *protocol.CgroupsRequest, builder *protocol.CgroupsBuilder) bool {
94 + if request.LayoutVersion != 1 || request.Flags != 0 {
95 + return false
96 + }
97 +
98 + builder.SetHeader(1, 42)
99 +
100 + items := []struct {
101 + hash, options, enabled uint32
102 + name, path []byte
103 + }{
104 + {1001, 0, 1, []byte("docker-abc123"), []byte("/sys/fs/cgroup/docker/abc123")},
105 + {2002, 0, 1, []byte("k8s-pod-xyz"), []byte("/sys/fs/cgroup/kubepods/xyz")},
106 + {3003, 0, 0, []byte("systemd-user"), []byte("/sys/fs/cgroup/user.slice/user-1000")},
107 + }
108 +
109 + for _, item := range items {
110 + if err := builder.Add(item.hash, item.options, item.enabled, item.name, item.path); err != nil {
111 + return false
112 + }
113 + }
114 +
115 + return true
116 +}
117 +
118 +// failingHandler always fails.
119 +func failingHandler(*protocol.CgroupsRequest, *protocol.CgroupsBuilder) bool {
120 + return false
121 +}
122 +
123 +func testSnapshotDispatch() DispatchHandler {
124 + return SnapshotDispatch(testCgroupsHandler, 3)
125 +}
126 +
127 +func failingSnapshotDispatch() DispatchHandler {
128 + return SnapshotDispatch(failingHandler, 3)
129 +}
130 +
131 +type testServer struct {
132 + server *Server
133 + doneCh chan struct{}
134 +}
135 +
136 +func startTestServer(service string, handler DispatchHandler) *testServer {
137 + ensureRunDir()
138 + cleanupAll(service)
139 +
140 + s := NewServer(testRunDir, service, testServerConfig(), protocol.MethodCgroupsSnapshot, handler)
141 + doneCh := make(chan struct{})
142 +
143 + go func() {
144 + defer close(doneCh)
145 + s.Run()
146 + }()
147 +
148 + // Wait until the server is actually accepting sessions.
149 + waitUnixServerReady(service)
150 +
151 + return &testServer{server: s, doneCh: doneCh}
152 +}
153 +
154 +func (ts *testServer) stop() {
155 + ts.server.Stop()
156 + <-ts.doneCh
157 +}
158 +
159 +func TestClientLifecycle(t *testing.T) {
160 + svc := "go_svc_lifecycle"
161 + ensureRunDir()
162 + cleanupAll(svc)
163 +
164 + // Init without server running
165 + client := NewSnapshotClient(testRunDir, svc, testClientConfig())
166 + if client.state != StateDisconnected {
167 + t.Fatal("expected DISCONNECTED")
168 + }
169 + if client.Ready() {
170 + t.Fatal("should not be ready")
171 + }
172 +
173 + // Refresh without server -> NOT_FOUND
174 + changed := client.Refresh()
175 + if !changed {
176 + t.Fatal("state should have changed")
177 + }
178 + if client.state != StateNotFound {
179 + t.Fatalf("expected NOT_FOUND, got %d", client.state)
180 + }
181 +
182 + // Start server
183 + ts := startTestServer(svc, testSnapshotDispatch())
184 + defer ts.stop()
185 +
186 + // Refresh -> READY
187 + changed = client.Refresh()
188 + if !changed {
189 + t.Fatal("state should have changed")
190 + }
191 + if client.state != StateReady {
192 + t.Fatalf("expected READY, got %d", client.state)
193 + }
194 + if !client.Ready() {
195 + t.Fatal("should be ready")
196 + }
197 +
198 + // Status reporting
199 + status := client.Status()
200 + if status.ConnectCount != 1 {
201 + t.Fatalf("expected connect_count=1, got %d", status.ConnectCount)
202 + }
203 + if status.ReconnectCount != 0 {
204 + t.Fatalf("expected reconnect_count=0, got %d", status.ReconnectCount)
205 + }
206 +
207 + // Close
208 + client.Close()
209 + if client.state != StateDisconnected {
210 + t.Fatal("expected DISCONNECTED after close")
211 + }
212 + if client.Ready() {
213 + t.Fatal("should not be ready after close")
214 + }
215 +
216 + cleanupAll(svc)
217 +}
218 +
219 +func TestCgroupsCall(t *testing.T) {
220 + svc := "go_svc_cgroups"
221 + ensureRunDir()
222 + cleanupAll(svc)
223 +
224 + ts := startTestServer(svc, testSnapshotDispatch())
225 + defer ts.stop()
226 +
227 + client := NewSnapshotClient(testRunDir, svc, testClientConfig())
228 + client.Refresh()
229 + if !client.Ready() {
230 + t.Fatal("client not ready")
231 + }
232 +
233 + view, err := client.CallSnapshot()
234 + if err != nil {
235 + t.Fatalf("call failed: %v", err)
236 + }
237 +
238 + if view.ItemCount != 3 {
239 + t.Fatalf("expected 3 items, got %d", view.ItemCount)
240 + }
241 + if view.SystemdEnabled != 1 {
242 + t.Fatalf("expected systemd_enabled=1, got %d", view.SystemdEnabled)
243 + }
244 + if view.Generation != 42 {
245 + t.Fatalf("expected generation=42, got %d", view.Generation)
246 + }
247 +
248 + // Verify first item
249 + item0, err := view.Item(0)
250 + if err != nil {
251 + t.Fatalf("item 0 error: %v", err)
252 + }
253 + if item0.Hash != 1001 {
254 + t.Fatalf("item 0 hash: got %d", item0.Hash)
255 + }
256 + if item0.Enabled != 1 {
257 + t.Fatalf("item 0 enabled: got %d", item0.Enabled)
258 + }
259 + if item0.Name.String() != "docker-abc123" {
260 + t.Fatalf("item 0 name: got %q", item0.Name.String())
261 + }
262 + if item0.Path.String() != "/sys/fs/cgroup/docker/abc123" {
263 + t.Fatalf("item 0 path: got %q", item0.Path.String())
264 + }
265 +
266 + // Verify third item
267 + item2, err := view.Item(2)
268 + if err != nil {
269 + t.Fatalf("item 2 error: %v", err)
270 + }
271 + if item2.Hash != 3003 {
272 + t.Fatalf("item 2 hash: got %d", item2.Hash)
273 + }
274 + if item2.Enabled != 0 {
275 + t.Fatalf("item 2 enabled: got %d", item2.Enabled)
276 + }
277 + if item2.Name.String() != "systemd-user" {
278 + t.Fatalf("item 2 name: got %q", item2.Name.String())
279 + }
280 +
281 + // Verify stats
282 + status := client.Status()
283 + if status.CallCount != 1 {
284 + t.Fatalf("expected call_count=1, got %d", status.CallCount)
285 + }
286 + if status.ErrorCount != 0 {
287 + t.Fatalf("expected error_count=0, got %d", status.ErrorCount)
288 + }
289 +
290 + client.Close()
291 + cleanupAll(svc)
292 +}
293 +
294 +func TestRetryOnFailure(t *testing.T) {
295 + svc := "go_svc_retry"
296 + ensureRunDir()
297 + cleanupAll(svc)
298 +
299 + ts1 := startTestServer(svc, testSnapshotDispatch())
300 +
301 + client := NewSnapshotClient(testRunDir, svc, testClientConfig())
302 + client.Refresh()
303 + if !client.Ready() {
304 + t.Fatal("client not ready")
305 + }
306 +
307 + // First call succeeds
308 + view, err := client.CallSnapshot()
309 + if err != nil {
310 + t.Fatalf("first call failed: %v", err)
311 + }
312 + if view.ItemCount != 3 {
313 + t.Fatalf("expected 3 items, got %d", view.ItemCount)
314 + }
315 +
316 + // Kill server
317 + ts1.stop()
318 + cleanupAll(svc)
319 + time.Sleep(50 * time.Millisecond)
320 +
321 + // Restart server
322 + ts2 := startTestServer(svc, testSnapshotDispatch())
323 + defer ts2.stop()
324 +
325 + // Next call triggers reconnect + retry
326 + view2, err := client.CallSnapshot()
327 + if err != nil {
328 + t.Fatalf("retry call failed: %v", err)
329 + }
330 + if view2.ItemCount != 3 {
331 + t.Fatalf("expected 3 items after retry, got %d", view2.ItemCount)
332 + }
333 +
334 + // Verify reconnect happened
335 + status := client.Status()
336 + if status.ReconnectCount < 1 {
337 + t.Fatalf("expected reconnect_count >= 1, got %d", status.ReconnectCount)
338 + }
339 +
340 + client.Close()
341 + cleanupAll(svc)
342 +}
343 +
344 +func TestMultipleClients(t *testing.T) {
345 + svc := "go_svc_multi"
346 + ensureRunDir()
347 + cleanupAll(svc)
348 +
349 + ts := startTestServer(svc, testSnapshotDispatch())
350 + defer ts.stop()
351 +
352 + // Client 1
353 + client1 := NewSnapshotClient(testRunDir, svc, testClientConfig())
354 + client1.Refresh()
355 + if !client1.Ready() {
356 + t.Fatal("client 1 not ready")
357 + }
358 +
359 + view1, err := client1.CallSnapshot()
360 + if err != nil {
361 + t.Fatalf("client 1 call failed: %v", err)
362 + }
363 + if view1.ItemCount != 3 {
364 + t.Fatalf("client 1: expected 3 items, got %d", view1.ItemCount)
365 + }
366 +
367 + // Now multi-client: keep client 1 open, connect client 2
368 + client2 := NewSnapshotClient(testRunDir, svc, testClientConfig())
369 + client2.Refresh()
370 + if !client2.Ready() {
371 + t.Fatal("client 2 not ready")
372 + }
373 +
374 + view2, err := client2.CallSnapshot()
375 + if err != nil {
376 + t.Fatalf("client 2 call failed: %v", err)
377 + }
378 + if view2.ItemCount != 3 {
379 + t.Fatalf("client 2: expected 3 items, got %d", view2.ItemCount)
380 + }
381 +
382 + client1.Close()
383 + client2.Close()
384 + cleanupAll(svc)
385 +}
386 +
387 +func TestConcurrentClients(t *testing.T) {
388 + svc := "go_svc_concurrent"
389 + ensureRunDir()
390 + cleanupAll(svc)
391 +
392 + ts := startTestServer(svc, testSnapshotDispatch())
393 + defer ts.stop()
394 +
395 + const numClients = 5
396 + const requestsPerClient = 10
397 +
398 + type result struct {
399 + successes int
400 + failures int
401 + }
402 +
403 + results := make(chan result, numClients)
404 +
405 + for i := 0; i < numClients; i++ {
406 + go func() {
407 + r := result{}
408 + client := NewSnapshotClient(testRunDir, svc, testClientConfig())
409 + defer client.Close()
410 +
411 + for retry := 0; retry < 100; retry++ {
412 + client.Refresh()
413 + if client.Ready() {
414 + break
415 + }
416 + time.Sleep(10 * time.Millisecond)
417 + }
418 +
419 + if !client.Ready() {
420 + r.failures = requestsPerClient
421 + results <- r
422 + return
423 + }
424 +
425 + for j := 0; j < requestsPerClient; j++ {
426 + view, err := client.CallSnapshot()
427 + if err != nil || view.ItemCount != 3 {
428 + r.failures++
429 + continue
430 + }
431 + // Verify content
432 + item0, err := view.Item(0)
433 + if err != nil || item0.Hash != 1001 || item0.Name.String() != "docker-abc123" {
434 + r.failures++
435 + continue
436 + }
437 + r.successes++
438 + }
439 + results <- r
440 + }()
441 + }
442 +
443 + totalSuccess := 0
444 + totalFailure := 0
445 + for i := 0; i < numClients; i++ {
446 + r := <-results
447 + totalSuccess += r.successes
448 + totalFailure += r.failures
449 + }
450 +
451 + expected := numClients * requestsPerClient
452 + if totalSuccess != expected {
453 + t.Fatalf("expected %d successes, got %d (failures: %d)", expected, totalSuccess, totalFailure)
454 + }
455 + if totalFailure != 0 {
456 + t.Fatalf("expected 0 failures, got %d", totalFailure)
457 + }
458 +
459 + cleanupAll(svc)
460 +}
461 +
462 +func TestHandlerFailure(t *testing.T) {
463 + svc := "go_svc_hfail"
464 + ensureRunDir()
465 + cleanupAll(svc)
466 +
467 + ts := startTestServer(svc, failingSnapshotDispatch())
468 + defer ts.stop()
469 +
470 + client := NewSnapshotClient(testRunDir, svc, testClientConfig())
471 + client.Refresh()
472 + if !client.Ready() {
473 + t.Fatal("client not ready")
474 + }
475 +
476 + _, err := client.CallSnapshot()
477 + if err == nil {
478 + t.Fatal("expected error when handler fails")
479 + }
480 +
481 + status := client.Status()
482 + if status.ErrorCount < 1 {
483 + t.Fatalf("expected error_count >= 1, got %d", status.ErrorCount)
484 + }
485 +
486 + client.Close()
487 + cleanupAll(svc)
488 +}
489 +
490 +func TestStatusReporting(t *testing.T) {
491 + svc := "go_svc_status"
492 + ensureRunDir()
493 + cleanupAll(svc)
494 +
495 + ts := startTestServer(svc, testSnapshotDispatch())
496 + defer ts.stop()
497 +
498 + client := NewSnapshotClient(testRunDir, svc, testClientConfig())
499 + client.Refresh()
500 + if !client.Ready() {
501 + t.Fatal("client not ready")
502 + }
503 +
504 + // Initial counters
505 + s0 := client.Status()
506 + if s0.ConnectCount != 1 {
507 + t.Fatalf("expected connect_count=1, got %d", s0.ConnectCount)
508 + }
509 + if s0.CallCount != 0 {
510 + t.Fatalf("expected call_count=0, got %d", s0.CallCount)
511 + }
512 + if s0.ErrorCount != 0 {
513 + t.Fatalf("expected error_count=0, got %d", s0.ErrorCount)
514 + }
515 +
516 + // Make 3 successful calls
517 + for i := 0; i < 3; i++ {
518 + _, err := client.CallSnapshot()
519 + if err != nil {
520 + t.Fatalf("call %d failed: %v", i, err)
521 + }
522 + }
523 +
524 + s1 := client.Status()
525 + if s1.CallCount != 3 {
526 + t.Fatalf("expected call_count=3, got %d", s1.CallCount)
527 + }
528 + if s1.ErrorCount != 0 {
529 + t.Fatalf("expected error_count=0, got %d", s1.ErrorCount)
530 + }
531 +
532 + // Call on disconnected client
533 + client.Close()
534 + _, err := client.CallSnapshot()
535 + if err == nil {
536 + t.Fatal("expected error on disconnected client")
537 + }
538 +
539 + s2 := client.Status()
540 + if s2.ErrorCount != 1 {
541 + t.Fatalf("expected error_count=1, got %d", s2.ErrorCount)
542 + }
543 +
544 + cleanupAll(svc)
545 +}
546 +
547 +func TestNonRequestTerminatesSession(t *testing.T) {
548 + svc := "go_svc_nonreq"
549 + ensureRunDir()
550 + cleanupAll(svc)
551 +
552 + ts := startTestServer(svc, testSnapshotDispatch())
553 + defer ts.stop()
554 +
555 + // Connect via raw UDS session (transport level)
556 + session, err := waitUnixRawSessionConnect(svc, posix.ClientConfig{
557 + SupportedProfiles: protocol.ProfileBaseline,
558 + MaxRequestPayloadBytes: 4096,
559 + MaxRequestBatchItems: 1,
560 + MaxResponsePayloadBytes: responseBufSize,
561 + MaxResponseBatchItems: 1,
562 + AuthToken: authToken,
563 + })
564 + if err != nil {
565 + t.Fatalf("raw connect failed: %v", err)
566 + }
567 +
568 + // Send a RESPONSE message (not REQUEST) - protocol violation
569 + badHdr := &protocol.Header{
570 + Kind: protocol.KindResponse, // wrong kind!
571 + Code: protocol.MethodCgroupsSnapshot,
572 + Flags: 0,
573 + ItemCount: 0,
574 + MessageID: 1,
575 + TransportStatus: protocol.StatusOK,
576 + }
577 + sendErr := session.Send(badHdr, nil)
578 + if sendErr != nil {
579 + // Send might fail immediately; that's acceptable
580 + t.Logf("send non-request failed immediately: %v", sendErr)
581 + } else {
582 + // Wait for server to process and terminate the session
583 + time.Sleep(200 * time.Millisecond)
584 +
585 + // Try to send a valid request and receive - should fail
586 + reqHdr := &protocol.Header{
587 + Kind: protocol.KindRequest,
588 + Code: protocol.MethodCgroupsSnapshot,
589 + Flags: 0,
590 + ItemCount: 1,
591 + MessageID: 2,
592 + TransportStatus: protocol.StatusOK,
593 + }
594 + var reqBuf [4]byte
595 + req := protocol.CgroupsRequest{LayoutVersion: 1, Flags: 0}
596 + req.Encode(reqBuf[:])
597 +
598 + _ = session.Send(reqHdr, reqBuf[:])
599 +
600 + recvBuf := make([]byte, 4096)
601 + _, _, recvErr := session.Receive(recvBuf)
602 + if recvErr == nil {
603 + t.Fatal("recv after non-request should fail (server terminated session)")
604 + }
605 + }
606 + session.Close()
607 +
608 + // Verify server is still alive: connect a new client and do a normal call
609 + verifyClient := NewSnapshotClient(testRunDir, svc, testClientConfig())
610 + refreshUnixClientReady(t, verifyClient)
611 +
612 + view, err := verifyClient.CallSnapshot()
613 + if err != nil {
614 + t.Fatalf("normal call should succeed after bad client: %v", err)
615 + }
616 + if view.ItemCount != 3 {
617 + t.Fatalf("expected 3 items, got %d", view.ItemCount)
618 + }
619 +
620 + verifyClient.Close()
621 + cleanupAll(svc)
622 +}
src/go/pkg/netipc/service/raw/client_windows.go new
+1121
@@ -0,0 +1,1121 @@
1 +//go:build windows
2 +
3 +// Internal typed service client for Windows.
4 +//
5 +// Identical state machine and retry logic as the POSIX client.
6 +// Uses Named Pipe + Win SHM transports instead of UDS + POSIX SHM.
7 +//
8 +// Pure Go — no cgo. Works with CGO_ENABLED=0.
9 +
10 +package raw
11 +
12 +import (
13 + "encoding/binary"
14 + "errors"
15 + "sync"
16 + "sync/atomic"
17 + "time"
18 +
19 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
20 + windows "github.com/netdata/netdata/go/plugins/pkg/netipc/transport/windows"
21 +)
22 +
23 +const clientShmAttachRetryInterval = 5 * time.Millisecond
24 +const clientShmAttachRetryTimeout = 5 * time.Second
25 +
26 +// ---------------------------------------------------------------------------
27 +// Client context
28 +// ---------------------------------------------------------------------------
29 +
30 +// Client is an internal L2 client context bound to one service kind.
31 +type Client struct {
32 + state ClientState
33 + runDir string
34 + serviceName string
35 + expectedMethodCode uint16
36 + config windows.ClientConfig
37 +
38 + session *windows.Session
39 + shm *windows.WinShmContext
40 +
41 + requestBuf []byte
42 + sendBuf []byte
43 + transportBuf []byte
44 +
45 + connectCount uint32
46 + reconnectCount uint32
47 + callCount uint32
48 + errorCount uint32
49 +}
50 +
51 +func newClient(runDir, serviceName string, config windows.ClientConfig, expectedMethodCode uint16) *Client {
52 + return &Client{
53 + state: StateDisconnected,
54 + runDir: runDir,
55 + serviceName: serviceName,
56 + expectedMethodCode: expectedMethodCode,
57 + config: config,
58 + }
59 +}
60 +
61 +// NewSnapshotClient creates a raw client bound to the cgroups-snapshot service kind.
62 +func NewSnapshotClient(runDir, serviceName string, config windows.ClientConfig) *Client {
63 + return newClient(runDir, serviceName, config, protocol.MethodCgroupsSnapshot)
64 +}
65 +
66 +// NewIncrementClient creates a raw client bound to the increment service kind.
67 +func NewIncrementClient(runDir, serviceName string, config windows.ClientConfig) *Client {
68 + return newClient(runDir, serviceName, config, protocol.MethodIncrement)
69 +}
70 +
71 +// NewStringReverseClient creates a raw client bound to the string-reverse service kind.
72 +func NewStringReverseClient(runDir, serviceName string, config windows.ClientConfig) *Client {
73 + return newClient(runDir, serviceName, config, protocol.MethodStringReverse)
74 +}
75 +
76 +func nextPowerOf2U32(n uint32) uint32 {
77 + if n < 16 {
78 + return 16
79 + }
80 + n--
81 + n |= n >> 1
82 + n |= n >> 2
83 + n |= n >> 4
84 + n |= n >> 8
85 + n |= n >> 16
86 + return n + 1
87 +}
88 +
89 +func (c *Client) validateMethod(methodCode uint16) error {
90 + if c.expectedMethodCode != methodCode {
91 + return protocol.ErrBadLayout
92 + }
93 + return nil
94 +}
95 +
96 +// Refresh attempts connect if DISCONNECTED/NOT_FOUND, reconnect if BROKEN.
97 +func (c *Client) Refresh() bool {
98 + oldState := c.state
99 +
100 + switch c.state {
101 + case StateDisconnected, StateNotFound:
102 + c.state = StateConnecting
103 + c.state = c.tryConnect()
104 + if c.state == StateReady {
105 + c.connectCount++
106 + }
107 +
108 + case StateBroken:
109 + c.disconnect()
110 + c.state = StateConnecting
111 + c.state = c.tryConnect()
112 + if c.state == StateReady {
113 + c.reconnectCount++
114 + }
115 +
116 + case StateReady, StateConnecting, StateAuthFailed, StateIncompatible:
117 + // No action needed
118 + }
119 +
120 + return c.state != oldState
121 +}
122 +
123 +// Ready returns true only if the client is in the READY state.
124 +func (c *Client) Ready() bool {
125 + return c.state == StateReady
126 +}
127 +
128 +// Status returns a diagnostic counters snapshot.
129 +func (c *Client) Status() ClientStatus {
130 + return ClientStatus{
131 + State: c.state,
132 + ConnectCount: c.connectCount,
133 + ReconnectCount: c.reconnectCount,
134 + CallCount: c.callCount,
135 + ErrorCount: c.errorCount,
136 + }
137 +}
138 +
139 +func (c *Client) sessionMaxRequestPayloadBytes() uint32 {
140 + if c.session != nil {
141 + return c.session.MaxRequestPayloadBytes
142 + }
143 + return c.config.MaxRequestPayloadBytes
144 +}
145 +
146 +func (c *Client) sessionMaxResponsePayloadBytes() uint32 {
147 + if c.session != nil {
148 + return c.session.MaxResponsePayloadBytes
149 + }
150 + return c.config.MaxResponsePayloadBytes
151 +}
152 +
153 +func (c *Client) noteRequestCapacity(payloadLen uint32) {
154 + grown := nextPowerOf2U32(payloadLen)
155 + if grown > protocol.MaxPayloadCap {
156 + grown = protocol.MaxPayloadCap
157 + }
158 + if grown > c.config.MaxRequestPayloadBytes {
159 + c.config.MaxRequestPayloadBytes = grown
160 + }
161 +}
162 +
163 +func (c *Client) noteResponseCapacity(payloadLen uint32) {
164 + grown := nextPowerOf2U32(payloadLen)
165 + if grown > protocol.MaxPayloadCap {
166 + grown = protocol.MaxPayloadCap
167 + }
168 + if grown > c.config.MaxResponsePayloadBytes {
169 + c.config.MaxResponsePayloadBytes = grown
170 + }
171 +}
172 +
173 +// callWithRetry manages reconnect-driven recovery for a typed call.
174 +// Ordinary failures retry once. Overflow-driven resize recovery may
175 +// reconnect more than once while negotiated capacities grow.
176 +func (c *Client) callWithRetry(attempt func() error) error {
177 + if c.state != StateReady {
178 + c.errorCount++
179 + return protocol.ErrBadLayout
180 + }
181 +
182 + for {
183 + prevReq := c.sessionMaxRequestPayloadBytes()
184 + prevResp := c.sessionMaxResponsePayloadBytes()
185 + prevCfgReq := c.config.MaxRequestPayloadBytes
186 + prevCfgResp := c.config.MaxResponsePayloadBytes
187 +
188 + firstErr := attempt()
189 + if firstErr == nil {
190 + c.callCount++
191 + return nil
192 + }
193 +
194 + if !errors.Is(firstErr, protocol.ErrOverflow) {
195 + c.disconnect()
196 + c.state = StateBroken
197 +
198 + c.state = c.tryConnect()
199 + if c.state != StateReady {
200 + c.errorCount++
201 + return firstErr
202 + }
203 + c.reconnectCount++
204 +
205 + retryErr := attempt()
206 + if retryErr == nil {
207 + c.callCount++
208 + return nil
209 + }
210 +
211 + c.disconnect()
212 + c.state = StateBroken
213 + c.errorCount++
214 + return retryErr
215 + }
216 +
217 + c.disconnect()
218 + c.state = StateBroken
219 +
220 + c.state = c.tryConnect()
221 + if c.state != StateReady {
222 + c.errorCount++
223 + return firstErr
224 + }
225 + c.reconnectCount++
226 +
227 + if c.sessionMaxRequestPayloadBytes() <= prevReq &&
228 + c.sessionMaxResponsePayloadBytes() <= prevResp &&
229 + c.config.MaxRequestPayloadBytes <= prevCfgReq &&
230 + c.config.MaxResponsePayloadBytes <= prevCfgResp {
231 + c.disconnect()
232 + c.state = StateBroken
233 + c.errorCount++
234 + return firstErr
235 + }
236 + }
237 +}
238 +
239 +// doRawCall sends a request and receives/validates the response envelope.
240 +// Returns the validated response header and a borrowed payload view.
241 +func (c *Client) doRawCall(methodCode uint16, reqPayload []byte) (protocol.Header, []byte, error) {
242 + hdr := protocol.Header{
243 + Kind: protocol.KindRequest,
244 + Code: methodCode,
245 + Flags: 0,
246 + ItemCount: 1,
247 + MessageID: uint64(c.callCount) + 1,
248 + TransportStatus: protocol.StatusOK,
249 + }
250 +
251 + if err := c.transportSend(&hdr, reqPayload); err != nil {
252 + return protocol.Header{}, nil, err
253 + }
254 +
255 + respHdr, payload, err := c.transportReceive()
256 + if err != nil {
257 + return protocol.Header{}, nil, err
258 + }
259 +
260 + if respHdr.Kind != protocol.KindResponse {
261 + return protocol.Header{}, nil, protocol.ErrBadKind
262 + }
263 + if respHdr.Code != methodCode {
264 + return protocol.Header{}, nil, protocol.ErrBadLayout
265 + }
266 + if respHdr.MessageID != hdr.MessageID {
267 + return protocol.Header{}, nil, protocol.ErrBadLayout
268 + }
269 + switch respHdr.TransportStatus {
270 + case protocol.StatusOK:
271 + case protocol.StatusLimitExceeded:
272 + if current := c.sessionMaxResponsePayloadBytes(); current > 0 {
273 + if current >= ^uint32(0)/2 {
274 + c.noteResponseCapacity(^uint32(0))
275 + } else {
276 + c.noteResponseCapacity(current * 2)
277 + }
278 + }
279 + return protocol.Header{}, nil, protocol.ErrOverflow
280 + default:
281 + return protocol.Header{}, nil, protocol.ErrBadLayout
282 + }
283 +
284 + return respHdr, payload, nil
285 +}
286 +
287 +// CallSnapshot performs a blocking typed cgroups snapshot call.
288 +func (c *Client) CallSnapshot() (*protocol.CgroupsResponseView, error) {
289 + if err := c.validateMethod(protocol.MethodCgroupsSnapshot); err != nil {
290 + return nil, err
291 + }
292 +
293 + var result *protocol.CgroupsResponseView
294 +
295 + err := c.callWithRetry(func() error {
296 + req := protocol.CgroupsRequest{LayoutVersion: 1, Flags: 0}
297 + var reqBuf [4]byte
298 + if req.Encode(reqBuf[:]) == 0 {
299 + return protocol.ErrTruncated
300 + }
301 +
302 + _, payload, rerr := c.doRawCall(protocol.MethodCgroupsSnapshot, reqBuf[:])
303 + if rerr != nil {
304 + return rerr
305 + }
306 +
307 + view, derr := protocol.DecodeCgroupsResponse(payload)
308 + if derr != nil {
309 + return derr
310 + }
311 + result = &view
312 + return nil
313 + })
314 + if err != nil {
315 + return nil, err
316 + }
317 + return result, nil
318 +}
319 +
320 +// CallIncrement performs a blocking INCREMENT call.
321 +// Sends requestValue, returns the server's response value.
322 +func (c *Client) CallIncrement(requestValue uint64) (uint64, error) {
323 + if err := c.validateMethod(protocol.MethodIncrement); err != nil {
324 + return 0, err
325 + }
326 +
327 + var result uint64
328 +
329 + err := c.callWithRetry(func() error {
330 + var reqBuf [protocol.IncrementPayloadSize]byte
331 + if protocol.IncrementEncode(requestValue, reqBuf[:]) == 0 {
332 + return protocol.ErrTruncated
333 + }
334 +
335 + _, payload, rerr := c.doRawCall(protocol.MethodIncrement, reqBuf[:])
336 + if rerr != nil {
337 + return rerr
338 + }
339 +
340 + val, derr := protocol.IncrementDecode(payload)
341 + if derr != nil {
342 + return derr
343 + }
344 + result = val
345 + return nil
346 + })
347 + return result, err
348 +}
349 +
350 +// CallStringReverse performs a blocking STRING_REVERSE call.
351 +// Sends requestStr, returns the server's reversed string view.
352 +func (c *Client) CallStringReverse(requestStr string) (*protocol.StringReverseView, error) {
353 + if err := c.validateMethod(protocol.MethodStringReverse); err != nil {
354 + return nil, err
355 + }
356 +
357 + var result *protocol.StringReverseView
358 +
359 + err := c.callWithRetry(func() error {
360 + reqBuf := ensureClientScratch(&c.requestBuf, protocol.StringReverseHdrSize+len(requestStr)+1)
361 + if protocol.StringReverseEncode(requestStr, reqBuf) == 0 {
362 + return protocol.ErrTruncated
363 + }
364 +
365 + _, payload, rerr := c.doRawCall(protocol.MethodStringReverse, reqBuf)
366 + if rerr != nil {
367 + return rerr
368 + }
369 +
370 + view, derr := protocol.StringReverseDecode(payload)
371 + if derr != nil {
372 + return derr
373 + }
374 + result = &view
375 + return nil
376 + })
377 + if err != nil {
378 + return nil, err
379 + }
380 + return result, nil
381 +}
382 +
383 +// CallIncrementBatch performs a blocking batch INCREMENT call.
384 +// Sends multiple values, returns the server's response values.
385 +func (c *Client) CallIncrementBatch(values []uint64) ([]uint64, error) {
386 + if err := c.validateMethod(protocol.MethodIncrement); err != nil {
387 + return nil, err
388 + }
389 +
390 + if len(values) == 0 {
391 + return nil, nil
392 + }
393 +
394 + var results []uint64
395 + itemCount := uint32(len(values))
396 +
397 + err := c.callWithRetry(func() error {
398 + // Build batch request payload
399 + batchBufSize := protocol.Align8(int(itemCount)*8) + int(itemCount)*protocol.IncrementPayloadSize + int(itemCount)*protocol.Alignment
400 + batchBuf := ensureClientScratch(&c.requestBuf, batchBufSize)
401 + bb := protocol.NewBatchBuilder(batchBuf, itemCount)
402 +
403 + for _, v := range values {
404 + var item [protocol.IncrementPayloadSize]byte
405 + if protocol.IncrementEncode(v, item[:]) == 0 {
406 + return protocol.ErrTruncated
407 + }
408 + if err := bb.Add(item[:]); err != nil {
409 + return err
410 + }
411 + }
412 +
413 + totalPayloadLen, _ := bb.Finish()
414 + reqPayload := batchBuf[:totalPayloadLen]
415 +
416 + // Build and send header with batch flags
417 + hdr := protocol.Header{
418 + Kind: protocol.KindRequest,
419 + Code: protocol.MethodIncrement,
420 + Flags: protocol.FlagBatch,
421 + ItemCount: itemCount,
422 + MessageID: uint64(c.callCount) + 1,
423 + TransportStatus: protocol.StatusOK,
424 + }
425 +
426 + if err := c.transportSend(&hdr, reqPayload); err != nil {
427 + return err
428 + }
429 +
430 + // Receive response
431 + respHdr, respPayload, err := c.transportReceive()
432 + if err != nil {
433 + return err
434 + }
435 +
436 + if respHdr.Kind != protocol.KindResponse {
437 + return protocol.ErrBadKind
438 + }
439 + if respHdr.Code != protocol.MethodIncrement {
440 + return protocol.ErrBadLayout
441 + }
442 + if respHdr.MessageID != hdr.MessageID {
443 + return protocol.ErrBadLayout
444 + }
445 + switch respHdr.TransportStatus {
446 + case protocol.StatusOK:
447 + case protocol.StatusLimitExceeded:
448 + if current := c.sessionMaxResponsePayloadBytes(); current > 0 {
449 + if current >= ^uint32(0)/2 {
450 + c.noteResponseCapacity(^uint32(0))
451 + } else {
452 + c.noteResponseCapacity(current * 2)
453 + }
454 + }
455 + return protocol.ErrOverflow
456 + default:
457 + return protocol.ErrBadLayout
458 + }
459 + if respHdr.Flags&protocol.FlagBatch == 0 || respHdr.ItemCount != itemCount {
460 + return protocol.ErrBadItemCount
461 + }
462 +
463 + // Extract each response item
464 + out := make([]uint64, itemCount)
465 + for i := uint32(0); i < itemCount; i++ {
466 + itemData, gerr := protocol.BatchItemGet(respPayload, itemCount, i)
467 + if gerr != nil {
468 + return gerr
469 + }
470 + val, derr := protocol.IncrementDecode(itemData)
471 + if derr != nil {
472 + return derr
473 + }
474 + out[i] = val
475 + }
476 + results = out
477 + return nil
478 + })
479 + return results, err
480 +}
481 +
482 +// Close tears down the connection and releases resources.
483 +func (c *Client) Close() {
484 + c.disconnect()
485 + c.state = StateDisconnected
486 +}
487 +
488 +// ------------------------------------------------------------------
489 +// Internal helpers
490 +// ------------------------------------------------------------------
491 +
492 +func (c *Client) disconnect() {
493 + if c.shm != nil {
494 + c.shm.WinShmClose()
495 + c.shm = nil
496 + }
497 + if c.session != nil {
498 + c.session.Close()
499 + c.session = nil
500 + }
501 +}
502 +
503 +func (c *Client) tryConnect() ClientState {
504 + session, err := windows.Connect(c.runDir, c.serviceName, &c.config)
505 + if err != nil {
506 + switch {
507 + case isConnectError(err):
508 + return StateNotFound
509 + case isAuthError(err):
510 + return StateAuthFailed
511 + case isIncompatibleError(err):
512 + return StateIncompatible
513 + default:
514 + return StateDisconnected
515 + }
516 + }
517 +
518 + // Win SHM upgrade if negotiated
519 + if session.SelectedProfile == windows.WinShmProfileHybrid ||
520 + session.SelectedProfile == windows.WinShmProfileBusywait {
521 + deadline := time.Now().Add(clientShmAttachRetryTimeout)
522 + for {
523 + shm, serr := windows.WinShmClientAttach(
524 + c.runDir, c.serviceName,
525 + c.config.AuthToken,
526 + session.SessionID,
527 + session.SelectedProfile,
528 + )
529 + if serr == nil {
530 + c.shm = shm
531 + break
532 + }
533 + if !time.Now().Before(deadline) {
534 + break
535 + }
536 + time.Sleep(clientShmAttachRetryInterval)
537 + }
538 + if c.shm == nil {
539 + // WinSHM attach failed after negotiation. Close that session,
540 + // blacklist WinSHM for this client context, and retry baseline.
541 + session.Close()
542 + c.config.SupportedProfiles &^= winShmProfiles
543 + c.config.PreferredProfiles &^= winShmProfiles
544 + if c.config.SupportedProfiles == 0 {
545 + return StateDisconnected
546 + }
547 + return c.tryConnect()
548 + }
549 + }
550 +
551 + c.session = session
552 + return StateReady
553 +}
554 +
555 +func (c *Client) transportSend(hdr *protocol.Header, payload []byte) error {
556 + if c.shm != nil {
557 + if len(payload) > int(c.sessionMaxRequestPayloadBytes()) {
558 + c.noteRequestCapacity(uint32(len(payload)))
559 + return protocol.ErrOverflow
560 + }
561 +
562 + msgLen := protocol.HeaderSize + len(payload)
563 + msg := ensureClientScratch(&c.sendBuf, msgLen)
564 +
565 + hdr.Magic = protocol.MagicMsg
566 + hdr.Version = protocol.Version
567 + hdr.HeaderLen = protocol.HeaderLen
568 + hdr.PayloadLen = uint32(len(payload))
569 +
570 + hdr.Encode(msg[:protocol.HeaderSize])
571 + if len(payload) > 0 {
572 + copy(msg[protocol.HeaderSize:], payload)
573 + }
574 +
575 + if err := c.shm.WinShmSend(msg[:msgLen]); err != nil {
576 + if errors.Is(err, windows.ErrWinShmMsgTooLarge) {
577 + c.noteRequestCapacity(uint32(len(payload)))
578 + return protocol.ErrOverflow
579 + }
580 + return protocol.ErrTruncated
581 + }
582 + return nil
583 + }
584 +
585 + if c.session == nil {
586 + return protocol.ErrTruncated
587 + }
588 + if err := c.session.Send(hdr, payload); err != nil {
589 + if errors.Is(err, windows.ErrLimitExceeded) {
590 + c.noteRequestCapacity(uint32(len(payload)))
591 + return protocol.ErrOverflow
592 + }
593 + return protocol.ErrTruncated
594 + }
595 + return nil
596 +}
597 +
598 +func (c *Client) transportReceive() (protocol.Header, []byte, error) {
599 + scratch := ensureClientScratch(&c.transportBuf, c.maxReceiveMessageBytes())
600 +
601 + if c.shm != nil {
602 + mlen, err := c.shm.WinShmReceive(scratch, 30000)
603 + if err != nil {
604 + return protocol.Header{}, nil, protocol.ErrTruncated
605 + }
606 + if mlen < protocol.HeaderSize {
607 + return protocol.Header{}, nil, protocol.ErrTruncated
608 + }
609 +
610 + hdr, err := protocol.DecodeHeader(scratch[:mlen])
611 + if err != nil {
612 + return protocol.Header{}, nil, err
613 + }
614 + return hdr, scratch[protocol.HeaderSize:mlen], nil
615 + }
616 +
617 + if c.session == nil {
618 + return protocol.Header{}, nil, protocol.ErrTruncated
619 + }
620 +
621 + hdr, payload, err := c.session.Receive(scratch)
622 + if err != nil {
623 + return protocol.Header{}, nil, protocol.ErrTruncated
624 + }
625 + return hdr, payload, nil
626 +}
627 +
628 +func (c *Client) maxReceiveMessageBytes() int {
629 + maxPayload := c.config.MaxResponsePayloadBytes
630 + if c.session != nil && c.session.MaxResponsePayloadBytes > 0 {
631 + maxPayload = c.session.MaxResponsePayloadBytes
632 + }
633 + if maxPayload == 0 {
634 + maxPayload = cacheResponseBufSize
635 + }
636 + return protocol.HeaderSize + int(maxPayload)
637 +}
638 +
639 +// Error classification helpers
640 +func isConnectError(err error) bool {
641 + return errors.Is(err, windows.ErrConnect) || errors.Is(err, windows.ErrCreatePipe)
642 +}
643 +
644 +func isAuthError(err error) bool {
645 + return errors.Is(err, windows.ErrAuthFailed)
646 +}
647 +
648 +func isIncompatibleError(err error) bool {
649 + return errors.Is(err, windows.ErrNoProfile) || errors.Is(err, windows.ErrIncompatible)
650 +}
651 +
652 +// ---------------------------------------------------------------------------
653 +// Managed server
654 +// ---------------------------------------------------------------------------
655 +
656 +// Server is an internal managed server bound to one expected request kind.
657 +type Server struct {
658 + runDir string
659 + serviceName string
660 + config windows.ServerConfig
661 + expectedMethodCode uint16
662 + handler DispatchHandler
663 + running atomic.Bool
664 + learnedRequestPayloadBytes atomic.Uint32
665 + learnedResponsePayloadBytes atomic.Uint32
666 + nextSessionID atomic.Uint64
667 + workerCount int
668 + wg sync.WaitGroup
669 + listener *windows.Listener // stored so Stop() can close it
670 +}
671 +
672 +// NewServer creates a new managed server.
673 +func NewServer(
674 + runDir, serviceName string,
675 + config windows.ServerConfig,
676 + expectedMethodCode uint16,
677 + handler DispatchHandler,
678 +) *Server {
679 + return NewServerWithWorkers(runDir, serviceName, config, expectedMethodCode, handler, 8)
680 +}
681 +
682 +// NewServerWithWorkers creates a server with an explicit worker count limit.
683 +func NewServerWithWorkers(
684 + runDir, serviceName string,
685 + config windows.ServerConfig,
686 + expectedMethodCode uint16,
687 + handler DispatchHandler,
688 + workerCount int,
689 +) *Server {
690 + if workerCount < 1 {
691 + workerCount = 1
692 + }
693 + learnedRequest := config.MaxRequestPayloadBytes
694 + if learnedRequest == 0 {
695 + learnedRequest = protocol.MaxPayloadDefault
696 + }
697 + learnedResponse := config.MaxResponsePayloadBytes
698 + if learnedResponse == 0 {
699 + learnedResponse = protocol.MaxPayloadDefault
700 + }
701 + s := &Server{
702 + runDir: runDir,
703 + serviceName: serviceName,
704 + config: config,
705 + expectedMethodCode: expectedMethodCode,
706 + handler: handler,
707 + workerCount: workerCount,
708 + }
709 + s.learnedRequestPayloadBytes.Store(learnedRequest)
710 + s.learnedResponsePayloadBytes.Store(learnedResponse)
711 + // Session ids are 1-based; prepareAcceptConfig() allocates with Add(1).
712 + s.nextSessionID.Store(0)
713 + return s
714 +}
715 +
716 +func (s *Server) dispatchSingle(methodCode uint16, request []byte, responseBuf []byte) (int, error) {
717 + if methodCode != s.expectedMethodCode || s.handler == nil {
718 + return 0, errHandlerFailed
719 + }
720 +
721 + return s.handler(request, responseBuf)
722 +}
723 +
724 +func (s *Server) methodSupported(methodCode uint16) bool {
725 + return s.handler != nil && methodCode == s.expectedMethodCode
726 +}
727 +
728 +func serverNotePayloadCapacity(target *atomic.Uint32, payloadLen uint32) {
729 + grown := nextPowerOf2U32(payloadLen)
730 + for {
731 + current := target.Load()
732 + if grown <= current {
733 + return
734 + }
735 + if target.CompareAndSwap(current, grown) {
736 + return
737 + }
738 + }
739 +}
740 +
741 +type preparedWinShm struct {
742 + hybrid *windows.WinShmContext
743 + busywait *windows.WinShmContext
744 +}
745 +
746 +func (p *preparedWinShm) take(profile uint32) *windows.WinShmContext {
747 + if p == nil {
748 + return nil
749 + }
750 + switch profile {
751 + case windows.WinShmProfileHybrid:
752 + ctx := p.hybrid
753 + p.hybrid = nil
754 + return ctx
755 + case windows.WinShmProfileBusywait:
756 + ctx := p.busywait
757 + p.busywait = nil
758 + return ctx
759 + default:
760 + return nil
761 + }
762 +}
763 +
764 +func (p *preparedWinShm) destroyAll() {
765 + if p == nil {
766 + return
767 + }
768 + if p.hybrid != nil {
769 + p.hybrid.WinShmDestroy()
770 + p.hybrid = nil
771 + }
772 + if p.busywait != nil {
773 + p.busywait.WinShmDestroy()
774 + p.busywait = nil
775 + }
776 +}
777 +
778 +const winShmProfiles = windows.WinShmProfileHybrid | windows.WinShmProfileBusywait
779 +
780 +func (s *Server) prepareAcceptConfig() (uint64, windows.ServerConfig, *preparedWinShm, bool) {
781 + sessionID := s.nextSessionID.Add(1)
782 + cfg := s.config
783 + cfg.MaxRequestPayloadBytes = s.learnedRequestPayloadBytes.Load()
784 + cfg.MaxResponsePayloadBytes = s.learnedResponsePayloadBytes.Load()
785 +
786 + if cfg.SupportedProfiles&winShmProfiles == 0 {
787 + return sessionID, cfg, nil, true
788 + }
789 +
790 + prepared := &preparedWinShm{}
791 + for _, profile := range []uint32{windows.WinShmProfileHybrid, windows.WinShmProfileBusywait} {
792 + if cfg.SupportedProfiles&profile == 0 {
793 + continue
794 + }
795 + shm, err := windows.WinShmServerCreate(
796 + s.runDir, s.serviceName,
797 + cfg.AuthToken,
798 + sessionID,
799 + profile,
800 + cfg.MaxRequestPayloadBytes+uint32(protocol.HeaderSize),
801 + cfg.MaxResponsePayloadBytes+uint32(protocol.HeaderSize),
802 + )
803 + if err != nil {
804 + cfg.SupportedProfiles &^= profile
805 + cfg.PreferredProfiles &^= profile
806 + continue
807 + }
808 + if profile == windows.WinShmProfileHybrid {
809 + prepared.hybrid = shm
810 + } else {
811 + prepared.busywait = shm
812 + }
813 + }
814 +
815 + if cfg.SupportedProfiles == 0 {
816 + prepared.destroyAll()
817 + return sessionID, cfg, nil, false
818 + }
819 +
820 + if prepared.hybrid == nil && prepared.busywait == nil {
821 + return sessionID, cfg, nil, true
822 + }
823 +
824 + return sessionID, cfg, prepared, true
825 +}
826 +
827 +// Run starts the acceptor loop. Blocking.
828 +func (s *Server) Run() error {
829 + listener, err := windows.Listen(s.runDir, s.serviceName, s.config)
830 + if err != nil {
831 + return err
832 + }
833 + s.listener = listener
834 + defer func() {
835 + listener.Close()
836 + s.listener = nil
837 + }()
838 +
839 + s.running.Store(true)
840 + sem := make(chan struct{}, s.workerCount)
841 +
842 + for s.running.Load() {
843 + sessionID, acceptCfg, preparedShm, ok := s.prepareAcceptConfig()
844 + if !ok {
845 + time.Sleep(10 * time.Millisecond)
846 + continue
847 + }
848 +
849 + session, err := listener.AcceptWithConfig(sessionID, acceptCfg)
850 + if err != nil {
851 + if preparedShm != nil {
852 + preparedShm.destroyAll()
853 + }
854 + if !s.running.Load() {
855 + break
856 + }
857 + time.Sleep(10 * time.Millisecond)
858 + continue
859 + }
860 +
861 + select {
862 + case sem <- struct{}{}:
863 + default:
864 + if preparedShm != nil {
865 + preparedShm.destroyAll()
866 + }
867 + session.Close()
868 + continue
869 + }
870 +
871 + var shm *windows.WinShmContext
872 + if session.SelectedProfile == windows.WinShmProfileHybrid ||
873 + session.SelectedProfile == windows.WinShmProfileBusywait {
874 + shm = preparedShm.take(session.SelectedProfile)
875 + if shm == nil {
876 + if preparedShm != nil {
877 + preparedShm.destroyAll()
878 + }
879 + session.Close()
880 + <-sem
881 + continue
882 + }
883 + }
884 + if preparedShm != nil {
885 + preparedShm.destroyAll()
886 + }
887 +
888 + s.wg.Add(1)
889 + go func(sess *windows.Session, shmCtx *windows.WinShmContext) {
890 + defer func() {
891 + if r := recover(); r != nil {
892 + // Session handler panicked; log but don't crash the server
893 + }
894 + <-sem
895 + s.wg.Done()
896 + }()
897 + s.handleSession(sess, shmCtx)
898 + }(session, shm)
899 + }
900 +
901 + s.wg.Wait()
902 + return nil
903 +}
904 +
905 +// Stop signals the server to stop and unblocks Accept by closing the listener.
906 +func (s *Server) Stop() {
907 + s.running.Store(false)
908 + if s.listener != nil {
909 + s.listener.Close()
910 + }
911 +}
912 +
913 +func (s *Server) handleSession(session *windows.Session, shm *windows.WinShmContext) {
914 + recvBuf := make([]byte, protocol.HeaderSize+int(session.MaxRequestPayloadBytes))
915 + respBuf := make([]byte, int(session.MaxResponsePayloadBytes))
916 + itemRespBuf := make([]byte, int(session.MaxResponsePayloadBytes))
917 + msgBuf := make([]byte, int(session.MaxResponsePayloadBytes)+protocol.HeaderSize)
918 +
919 + defer func() {
920 + if shm != nil {
921 + shm.WinShmDestroy()
922 + }
923 + session.Close()
924 + }()
925 +
926 + for s.running.Load() {
927 + var hdr protocol.Header
928 + var payload []byte
929 +
930 + if shm != nil {
931 + mlen, err := shm.WinShmReceive(recvBuf, serverPollTimeoutMs)
932 + if err != nil {
933 + if err == windows.ErrWinShmTimeout {
934 + continue
935 + }
936 + return
937 + }
938 + if mlen < protocol.HeaderSize {
939 + return
940 + }
941 + h, err := protocol.DecodeHeader(recvBuf[:mlen])
942 + if err != nil {
943 + return
944 + }
945 + hdr = h
946 + payload = recvBuf[protocol.HeaderSize:mlen]
947 + } else {
948 + // Named Pipe path
949 + ready, waitErr := session.WaitReadable(serverPollTimeoutMs)
950 + if waitErr != nil {
951 + return
952 + }
953 + if !ready {
954 + continue
955 + }
956 + h, p, err := session.Receive(recvBuf)
957 + if err != nil {
958 + return
959 + }
960 + hdr = h
961 + payload = p
962 + }
963 +
964 + // Protocol violation: unexpected message kind terminates session
965 + if hdr.Kind != protocol.KindRequest {
966 + return
967 + }
968 +
969 + if len(payload) <= int(^uint32(0)) {
970 + serverNotePayloadCapacity(&s.learnedRequestPayloadBytes, uint32(len(payload)))
971 + }
972 +
973 + if !s.methodSupported(hdr.Code) {
974 + respHdr := protocol.Header{
975 + Kind: protocol.KindResponse,
976 + Code: hdr.Code,
977 + MessageID: hdr.MessageID,
978 + TransportStatus: protocol.StatusUnsupported,
979 + ItemCount: 1,
980 + }
981 +
982 + if shm != nil {
983 + if len(msgBuf) < protocol.HeaderSize {
984 + msgBuf = make([]byte, protocol.HeaderSize)
985 + }
986 + msg := msgBuf[:protocol.HeaderSize]
987 + respHdr.Magic = protocol.MagicMsg
988 + respHdr.Version = protocol.Version
989 + respHdr.HeaderLen = protocol.HeaderLen
990 + respHdr.PayloadLen = 0
991 + respHdr.Encode(msg[:protocol.HeaderSize])
992 + if err := shm.WinShmSend(msg); err != nil {
993 + return
994 + }
995 + } else if err := session.Send(&respHdr, nil); err != nil {
996 + return
997 + }
998 + continue
999 + }
1000 +
1001 + // Dispatch: single-item or batch
1002 + responseLen := 0
1003 + isBatch := (hdr.Flags&protocol.FlagBatch != 0) && hdr.ItemCount >= 1
1004 + var dispatchErr error
1005 +
1006 + if !isBatch {
1007 + var derr error
1008 + responseLen, derr = s.dispatchSingle(hdr.Code, payload, respBuf)
1009 + if derr != nil {
1010 + dispatchErr = derr
1011 + responseLen = 0
1012 + } else if responseLen < 0 || responseLen > len(respBuf) {
1013 + dispatchErr = protocol.ErrOverflow
1014 + responseLen = 0
1015 + }
1016 + } else {
1017 + var bb protocol.BatchBuilder
1018 + bb.Reset(respBuf, hdr.ItemCount)
1019 +
1020 + for i := uint32(0); i < hdr.ItemCount && dispatchErr == nil; i++ {
1021 + itemData, gerr := protocol.BatchItemGet(payload, hdr.ItemCount, i)
1022 + if gerr != nil {
1023 + dispatchErr = gerr
1024 + break
1025 + }
1026 +
1027 + itemResultLen, derr := s.dispatchSingle(hdr.Code, itemData, itemRespBuf)
1028 + if derr != nil {
1029 + dispatchErr = derr
1030 + break
1031 + }
1032 + if itemResultLen < 0 || itemResultLen > len(itemRespBuf) {
1033 + dispatchErr = protocol.ErrOverflow
1034 + break
1035 + }
1036 +
1037 + if aerr := bb.Add(itemRespBuf[:itemResultLen]); aerr != nil {
1038 + dispatchErr = aerr
1039 + break
1040 + }
1041 + }
1042 +
1043 + if dispatchErr == nil {
1044 + responseLen, _ = bb.Finish()
1045 + }
1046 + }
1047 +
1048 + // Build response header
1049 + respHdr := protocol.Header{
1050 + Kind: protocol.KindResponse,
1051 + Code: hdr.Code,
1052 + MessageID: hdr.MessageID,
1053 + }
1054 +
1055 + if dispatchErr == nil {
1056 + if responseLen <= int(^uint32(0)) {
1057 + serverNotePayloadCapacity(&s.learnedResponsePayloadBytes, uint32(responseLen))
1058 + }
1059 + respHdr.TransportStatus = protocol.StatusOK
1060 + if isBatch {
1061 + respHdr.Flags = protocol.FlagBatch
1062 + respHdr.ItemCount = hdr.ItemCount
1063 + } else {
1064 + respHdr.ItemCount = 1
1065 + }
1066 + } else if errors.Is(dispatchErr, protocol.ErrOverflow) {
1067 + current := session.MaxResponsePayloadBytes
1068 + if current >= ^uint32(0)/2 {
1069 + serverNotePayloadCapacity(&s.learnedResponsePayloadBytes, ^uint32(0))
1070 + } else {
1071 + serverNotePayloadCapacity(&s.learnedResponsePayloadBytes, current*2)
1072 + }
1073 + respHdr.TransportStatus = protocol.StatusLimitExceeded
1074 + respHdr.ItemCount = 1
1075 + responseLen = 0
1076 + } else if errors.Is(dispatchErr, errHandlerFailed) {
1077 + respHdr.TransportStatus = protocol.StatusInternalError
1078 + respHdr.ItemCount = 1
1079 + responseLen = 0
1080 + } else {
1081 + respHdr.TransportStatus = protocol.StatusBadEnvelope
1082 + respHdr.ItemCount = 1
1083 + responseLen = 0
1084 + }
1085 +
1086 + if shm != nil {
1087 + msgLen := protocol.HeaderSize + responseLen
1088 + if len(msgBuf) < msgLen {
1089 + msgBuf = make([]byte, msgLen)
1090 + }
1091 + msg := msgBuf[:msgLen]
1092 +
1093 + respHdr.Magic = protocol.MagicMsg
1094 + respHdr.Version = protocol.Version
1095 + respHdr.HeaderLen = protocol.HeaderLen
1096 + respHdr.PayloadLen = uint32(responseLen)
1097 +
1098 + respHdr.Encode(msg[:protocol.HeaderSize])
1099 + if responseLen > 0 {
1100 + copy(msg[protocol.HeaderSize:], respBuf[:responseLen])
1101 + }
1102 +
1103 + if err := shm.WinShmSend(msg); err != nil {
1104 + return
1105 + }
1106 + if respHdr.TransportStatus == protocol.StatusLimitExceeded {
1107 + return
1108 + }
1109 + } else {
1110 + if err := session.Send(&respHdr, respBuf[:responseLen]); err != nil {
1111 + return
1112 + }
1113 + if respHdr.TransportStatus == protocol.StatusLimitExceeded {
1114 + return
1115 + }
1116 + }
1117 + }
1118 +}
1119 +
1120 +// Suppress unused import warnings.
1121 +var _ = binary.NativeEndian
src/go/pkg/netipc/service/raw/edge_test.go new
+546
@@ -0,0 +1,546 @@
1 +//go:build unix
2 +
3 +package raw
4 +
5 +import (
6 + "testing"
7 + "time"
8 +
9 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
10 + "github.com/netdata/netdata/go/plugins/pkg/netipc/transport/posix"
11 +)
12 +
13 +// ---------------------------------------------------------------------------
14 +// Client.Refresh: state transitions
15 +// ---------------------------------------------------------------------------
16 +
17 +func TestClientRefreshFromReady(t *testing.T) {
18 + svc := "go_edge_ready"
19 + ensureRunDir()
20 + cleanupAll(svc)
21 +
22 + ts := startTestServer(svc, testSnapshotDispatch())
23 + defer ts.stop()
24 +
25 + client := NewSnapshotClient(testRunDir, svc, testClientConfig())
26 + defer client.Close()
27 +
28 + client.Refresh() // DISCONNECTED -> READY
29 +
30 + // Refresh from READY should be a no-op
31 + changed := client.Refresh()
32 + if changed {
33 + t.Fatal("Refresh from READY should not change state")
34 + }
35 + if client.state != StateReady {
36 + t.Fatalf("expected READY, got %d", client.state)
37 + }
38 +
39 + cleanupAll(svc)
40 +}
41 +
42 +func TestClientRefreshFromAuthFailed(t *testing.T) {
43 + svc := "go_edge_auth"
44 + ensureRunDir()
45 + cleanupAll(svc)
46 +
47 + // Server with token=1
48 + sCfg := testServerConfig()
49 + sCfg.AuthToken = 0x1111
50 + s := NewServer(
51 + testRunDir,
52 + svc,
53 + sCfg,
54 + protocol.MethodCgroupsSnapshot,
55 + testSnapshotDispatch(),
56 + )
57 + doneCh := make(chan struct{})
58 + go func() {
59 + defer close(doneCh)
60 + s.Run()
61 + }()
62 + time.Sleep(100 * time.Millisecond)
63 + defer func() {
64 + s.Stop()
65 + <-doneCh
66 + }()
67 +
68 + // Client with wrong token
69 + cCfg := testClientConfig()
70 + cCfg.AuthToken = 0x2222
71 + client := NewSnapshotClient(testRunDir, svc, cCfg)
72 + defer client.Close()
73 +
74 + client.Refresh()
75 + if client.state != StateAuthFailed {
76 + t.Fatalf("expected StateAuthFailed, got %d", client.state)
77 + }
78 +
79 + // Refresh from AuthFailed should be a no-op
80 + changed := client.Refresh()
81 + if changed {
82 + t.Fatal("Refresh from StateAuthFailed should not change state")
83 + }
84 +
85 + cleanupAll(svc)
86 +}
87 +
88 +func TestClientRefreshFromIncompatible(t *testing.T) {
89 + svc := "go_edge_incompat"
90 + ensureRunDir()
91 + cleanupAll(svc)
92 +
93 + // Server supports only SHMFutex
94 + sCfg := testServerConfig()
95 + sCfg.SupportedProfiles = protocol.ProfileSHMFutex
96 + s := NewServer(
97 + testRunDir,
98 + svc,
99 + sCfg,
100 + protocol.MethodCgroupsSnapshot,
101 + testSnapshotDispatch(),
102 + )
103 + doneCh := make(chan struct{})
104 + go func() {
105 + defer close(doneCh)
106 + s.Run()
107 + }()
108 + time.Sleep(100 * time.Millisecond)
109 + defer func() {
110 + s.Stop()
111 + <-doneCh
112 + }()
113 +
114 + // Client supports only Baseline
115 + cCfg := testClientConfig()
116 + cCfg.SupportedProfiles = protocol.ProfileBaseline
117 + client := NewSnapshotClient(testRunDir, svc, cCfg)
118 + defer client.Close()
119 +
120 + client.Refresh()
121 + if client.state != StateIncompatible {
122 + t.Fatalf("expected StateIncompatible, got %d", client.state)
123 + }
124 +
125 + // Refresh from Incompatible should be a no-op
126 + changed := client.Refresh()
127 + if changed {
128 + t.Fatal("Refresh from StateIncompatible should not change state")
129 + }
130 +
131 + cleanupAll(svc)
132 +}
133 +
134 +func TestClientRefreshFromProtocolVersionMismatch(t *testing.T) {
135 + svc := uniqueUnixService("go_edge_proto_incompat")
136 + packet := encodeHelloAckPacketWithVersion(protocol.Version+1, protocol.StatusOK, 1)
137 + srv := startRawPosixHelloAckServer(t, svc, packet)
138 +
139 + client := NewSnapshotClient(testRunDir, svc, testClientConfig())
140 + defer client.Close()
141 +
142 + changed := client.Refresh()
143 + if !changed {
144 + t.Fatal("Refresh should move client into StateIncompatible")
145 + }
146 + if client.state != StateIncompatible {
147 + t.Fatalf("expected StateIncompatible, got %d", client.state)
148 + }
149 + if client.Ready() {
150 + t.Fatal("client should not be ready after protocol version mismatch")
151 + }
152 +
153 + srv.wait(t)
154 + if srv.accepted.Load() != 1 {
155 + t.Fatalf("expected exactly one raw handshake attempt, got %d", srv.accepted.Load())
156 + }
157 +
158 + changed = client.Refresh()
159 + if changed {
160 + t.Fatal("Refresh from StateIncompatible should be a no-op after protocol mismatch")
161 + }
162 + if client.state != StateIncompatible {
163 + t.Fatalf("expected StateIncompatible after second refresh, got %d", client.state)
164 + }
165 +
166 + cleanupAll(svc)
167 +}
168 +
169 +func TestClientRefreshFromBroken(t *testing.T) {
170 + svc := "go_edge_broken"
171 + ensureRunDir()
172 + cleanupAll(svc)
173 +
174 + ts := startTestServer(svc, testSnapshotDispatch())
175 +
176 + client := NewSnapshotClient(testRunDir, svc, testClientConfig())
177 + defer client.Close()
178 +
179 + client.Refresh()
180 + if !client.Ready() {
181 + t.Fatal("expected READY")
182 + }
183 +
184 + // Kill server to make connection broken
185 + ts.stop()
186 + cleanupAll(svc)
187 + time.Sleep(50 * time.Millisecond)
188 +
189 + // Force a call to break the connection
190 + _, _ = client.CallSnapshot()
191 +
192 + // At this point, state should be BROKEN
193 + if client.state != StateBroken {
194 + t.Logf("state is %d (may not be broken if retry succeeded; skipping)", client.state)
195 + return
196 + }
197 +
198 + // Start a new server
199 + ts2 := startTestServer(svc, testSnapshotDispatch())
200 + defer ts2.stop()
201 +
202 + // Refresh from BROKEN should reconnect
203 + changed := client.Refresh()
204 + if !changed {
205 + t.Fatal("Refresh from BROKEN should change state")
206 + }
207 + if client.state != StateReady {
208 + t.Fatalf("expected READY after reconnect, got %d", client.state)
209 + }
210 + if client.reconnectCount < 1 {
211 + t.Fatalf("expected reconnectCount >= 1, got %d", client.reconnectCount)
212 + }
213 +
214 + cleanupAll(svc)
215 +}
216 +
217 +// ---------------------------------------------------------------------------
218 +// Client.callWithRetry: not-ready fast-fail
219 +// ---------------------------------------------------------------------------
220 +
221 +func TestCallWithRetryNotReady(t *testing.T) {
222 + svc := "go_edge_notready"
223 + ensureRunDir()
224 +
225 + client := NewSnapshotClient(testRunDir, svc, testClientConfig())
226 + defer client.Close()
227 +
228 + // Don't call Refresh - client is DISCONNECTED
229 + _, err := client.CallSnapshot()
230 + if err == nil {
231 + t.Fatal("expected error when client is not ready")
232 + }
233 + if client.errorCount != 1 {
234 + t.Fatalf("expected errorCount=1, got %d", client.errorCount)
235 + }
236 +}
237 +
238 +// ---------------------------------------------------------------------------
239 +// NewServerWithWorkers: workerCount < 1 defaults to 1
240 +// ---------------------------------------------------------------------------
241 +
242 +func TestNewServerWithWorkersMinimum(t *testing.T) {
243 + s := NewServerWithWorkers(
244 + "/tmp",
245 + "test",
246 + posix.ServerConfig{},
247 + protocol.MethodCgroupsSnapshot,
248 + testSnapshotDispatch(),
249 + 0,
250 + )
251 + if s.workerCount != 1 {
252 + t.Fatalf("expected workerCount=1 for input 0, got %d", s.workerCount)
253 + }
254 + s2 := NewServerWithWorkers(
255 + "/tmp",
256 + "test",
257 + posix.ServerConfig{},
258 + protocol.MethodCgroupsSnapshot,
259 + testSnapshotDispatch(),
260 + -5,
261 + )
262 + if s2.workerCount != 1 {
263 + t.Fatalf("expected workerCount=1 for input -5, got %d", s2.workerCount)
264 + }
265 +}
266 +
267 +// ---------------------------------------------------------------------------
268 +// Cache: Lookup on empty (unpopulated) cache
269 +// ---------------------------------------------------------------------------
270 +
271 +func TestCacheLookupBeforeRefresh(t *testing.T) {
272 + cache := NewCache(testRunDir, "nonexistent", testClientConfig())
273 + defer cache.Close()
274 +
275 + if cache.Ready() {
276 + t.Fatal("should not be ready")
277 + }
278 +
279 + _, found := cache.Lookup(123, "anything")
280 + if found {
281 + t.Fatal("should not find items in empty cache")
282 + }
283 +}
284 +
285 +// ---------------------------------------------------------------------------
286 +// Cache: Refresh with no server running
287 +// ---------------------------------------------------------------------------
288 +
289 +func TestCacheRefreshNoServer(t *testing.T) {
290 + svc := "go_edge_cache_nosrv"
291 + ensureRunDir()
292 + cleanupAll(svc)
293 +
294 + cache := NewCache(testRunDir, svc, testClientConfig())
295 + defer cache.Close()
296 +
297 + updated := cache.Refresh()
298 + if updated {
299 + t.Fatal("refresh should fail with no server")
300 + }
301 + if cache.Ready() {
302 + t.Fatal("should not be ready after failed refresh")
303 + }
304 +
305 + status := cache.Status()
306 + if status.RefreshFailureCount != 1 {
307 + t.Fatalf("expected failure_count=1, got %d", status.RefreshFailureCount)
308 + }
309 +
310 + cleanupAll(svc)
311 +}
312 +
313 +// ---------------------------------------------------------------------------
314 +// Cache: Status on fresh cache
315 +// ---------------------------------------------------------------------------
316 +
317 +func TestCacheStatusFresh(t *testing.T) {
318 + cache := NewCache(testRunDir, "fresh", testClientConfig())
319 + defer cache.Close()
320 +
321 + status := cache.Status()
322 + if status.Populated {
323 + t.Fatal("should not be populated")
324 + }
325 + if status.ItemCount != 0 {
326 + t.Fatalf("expected item_count=0, got %d", status.ItemCount)
327 + }
328 + if status.RefreshSuccessCount != 0 {
329 + t.Fatalf("expected success_count=0, got %d", status.RefreshSuccessCount)
330 + }
331 + if status.RefreshFailureCount != 0 {
332 + t.Fatalf("expected failure_count=0, got %d", status.RefreshFailureCount)
333 + }
334 + if status.LastRefreshTs != 0 {
335 + t.Fatalf("expected last_refresh_ts=0, got %d", status.LastRefreshTs)
336 + }
337 +}
338 +
339 +// ---------------------------------------------------------------------------
340 +// Cache: Close resets state
341 +// ---------------------------------------------------------------------------
342 +
343 +func TestCacheCloseResetsState(t *testing.T) {
344 + svc := "go_edge_cache_close"
345 + ensureRunDir()
346 + cleanupAll(svc)
347 +
348 + ts := startTestServer(svc, testSnapshotDispatch())
349 + defer ts.stop()
350 +
351 + cache := NewCache(testRunDir, svc, testClientConfig())
352 + if !cache.Refresh() {
353 + t.Fatal("refresh should succeed")
354 + }
355 + if !cache.Ready() {
356 + t.Fatal("should be ready")
357 + }
358 +
359 + cache.Close()
360 +
361 + if cache.Ready() {
362 + t.Fatal("should not be ready after close")
363 + }
364 + _, found := cache.Lookup(1001, "docker-abc123")
365 + if found {
366 + t.Fatal("should not find items after close")
367 + }
368 +
369 + cleanupAll(svc)
370 +}
371 +
372 +// ---------------------------------------------------------------------------
373 +// Cache: NewCache with custom MaxResponsePayloadBytes
374 +// ---------------------------------------------------------------------------
375 +
376 +func TestCacheCustomBufferSize(t *testing.T) {
377 + cfg := testClientConfig()
378 + cfg.MaxResponsePayloadBytes = 1024
379 +
380 + cache := NewCache(testRunDir, "custom", cfg)
381 + defer cache.Close()
382 +
383 + expectedBufSize := protocol.HeaderSize + 1024
384 + if got := cache.client.maxReceiveMessageBytes(); got != expectedBufSize {
385 + t.Fatalf("expected response buf size=%d, got %d", expectedBufSize, got)
386 + }
387 +}
388 +
389 +func TestCacheDefaultBufferSize(t *testing.T) {
390 + cfg := testClientConfig()
391 + cfg.MaxResponsePayloadBytes = 0
392 +
393 + cache := NewCache(testRunDir, "default", cfg)
394 + defer cache.Close()
395 +
396 + expectedBufSize := protocol.HeaderSize + cacheResponseBufSize
397 + if got := cache.client.maxReceiveMessageBytes(); got != expectedBufSize {
398 + t.Fatalf("expected response buf size=%d, got %d", expectedBufSize, got)
399 + }
400 +}
401 +
402 +// ---------------------------------------------------------------------------
403 +// Cache: Lookup miss with correct hash but wrong name, and vice versa
404 +// ---------------------------------------------------------------------------
405 +
406 +func TestCacheLookupHashNameMismatch(t *testing.T) {
407 + svc := "go_edge_cache_mismatch"
408 + ensureRunDir()
409 + cleanupAll(svc)
410 +
411 + ts := startTestServer(svc, testSnapshotDispatch())
412 + defer ts.stop()
413 +
414 + cache := NewCache(testRunDir, svc, testClientConfig())
415 + defer cache.Close()
416 +
417 + if !cache.Refresh() {
418 + t.Fatal("refresh should succeed")
419 + }
420 +
421 + // Correct hash (1001), wrong name
422 + _, found := cache.Lookup(1001, "wrong-name")
423 + if found {
424 + t.Fatal("should not find with wrong name")
425 + }
426 +
427 + // Wrong hash, correct name
428 + _, found = cache.Lookup(9999, "docker-abc123")
429 + if found {
430 + t.Fatal("should not find with wrong hash")
431 + }
432 +
433 + // Both correct
434 + item, found := cache.Lookup(1001, "docker-abc123")
435 + if !found {
436 + t.Fatal("should find with correct hash+name")
437 + }
438 + if item.Hash != 1001 {
439 + t.Fatalf("expected hash=1001, got %d", item.Hash)
440 + }
441 +
442 + cleanupAll(svc)
443 +}
444 +
445 +// ---------------------------------------------------------------------------
446 +// CallIncrementBatch: empty slice
447 +// ---------------------------------------------------------------------------
448 +
449 +func TestCallIncrementBatchEmpty(t *testing.T) {
450 + svc := "go_edge_batch_empty"
451 + ensureRunDir()
452 + cleanupAll(svc)
453 +
454 + client := NewIncrementClient(testRunDir, svc, testClientConfig())
455 + defer client.Close()
456 +
457 + results, err := client.CallIncrementBatch(nil)
458 + if err != nil {
459 + t.Fatalf("expected nil error for empty batch, got %v", err)
460 + }
461 + if results != nil {
462 + t.Fatalf("expected nil results for empty batch, got %v", results)
463 + }
464 +
465 + results2, err := client.CallIncrementBatch([]uint64{})
466 + if err != nil {
467 + t.Fatalf("expected nil error for empty slice, got %v", err)
468 + }
469 + if results2 != nil {
470 + t.Fatalf("expected nil results for empty slice, got %v", results2)
471 + }
472 +
473 + cleanupAll(svc)
474 +}
475 +
476 +// ---------------------------------------------------------------------------
477 +// Client error classification helpers
478 +// ---------------------------------------------------------------------------
479 +
480 +func TestIsConnectError(t *testing.T) {
481 + if !isConnectError(posix.ErrConnect) {
482 + t.Error("should match ErrConnect")
483 + }
484 + if !isConnectError(posix.ErrSocket) {
485 + t.Error("should match ErrSocket")
486 + }
487 + if isConnectError(posix.ErrAuthFailed) {
488 + t.Error("should not match ErrAuthFailed")
489 + }
490 + if isConnectError(posix.ErrNoProfile) {
491 + t.Error("should not match ErrNoProfile")
492 + }
493 + if isConnectError(nil) {
494 + t.Error("should not match nil")
495 + }
496 +}
497 +
498 +func TestIsAuthError(t *testing.T) {
499 + if !isAuthError(posix.ErrAuthFailed) {
500 + t.Error("should match ErrAuthFailed")
501 + }
502 + if isAuthError(posix.ErrConnect) {
503 + t.Error("should not match ErrConnect")
504 + }
505 + if isAuthError(nil) {
506 + t.Error("should not match nil")
507 + }
508 +}
509 +
510 +func TestIsIncompatibleError(t *testing.T) {
511 + if !isIncompatibleError(posix.ErrNoProfile) {
512 + t.Error("should match ErrNoProfile")
513 + }
514 + if !isIncompatibleError(posix.ErrIncompatible) {
515 + t.Error("should match ErrIncompatible")
516 + }
517 + if isIncompatibleError(posix.ErrAuthFailed) {
518 + t.Error("should not match ErrAuthFailed")
519 + }
520 + if isIncompatibleError(nil) {
521 + t.Error("should not match nil")
522 + }
523 +}
524 +
525 +// ---------------------------------------------------------------------------
526 +// Client.Status on fresh client
527 +// ---------------------------------------------------------------------------
528 +
529 +func TestClientStatusFresh(t *testing.T) {
530 + client := NewSnapshotClient(testRunDir, "nosvc", testClientConfig())
531 + defer client.Close()
532 +
533 + status := client.Status()
534 + if status.State != StateDisconnected {
535 + t.Fatalf("expected DISCONNECTED, got %d", status.State)
536 + }
537 + if status.ConnectCount != 0 {
538 + t.Fatalf("expected connect_count=0, got %d", status.ConnectCount)
539 + }
540 + if status.CallCount != 0 {
541 + t.Fatalf("expected call_count=0, got %d", status.CallCount)
542 + }
543 + if status.ErrorCount != 0 {
544 + t.Fatalf("expected error_count=0, got %d", status.ErrorCount)
545 + }
546 +}
src/go/pkg/netipc/service/raw/helpers_windows_test.go new
+439
@@ -0,0 +1,439 @@
1 +//go:build windows
2 +
3 +package raw
4 +
5 +import (
6 + "fmt"
7 + "os"
8 + "sync/atomic"
9 + "syscall"
10 + "testing"
11 + "time"
12 + "unsafe"
13 +
14 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
15 + windows "github.com/netdata/netdata/go/plugins/pkg/netipc/transport/windows"
16 +)
17 +
18 +const (
19 + rawPipeAccessDuplex = 0x00000003
20 + rawPipeTypeMessage = 0x00000004
21 + rawPipeReadModeMsg = 0x00000002
22 + rawPipeWait = 0x00000000
23 + rawErrorPipeConnected = 535
24 + rawPipeBufSize = 65536
25 +)
26 +
27 +var (
28 + rawKernel32 = syscall.NewLazyDLL("kernel32.dll")
29 + rawCreateNamedPipeW = rawKernel32.NewProc("CreateNamedPipeW")
30 + rawConnectNamedPipe = rawKernel32.NewProc("ConnectNamedPipe")
31 +)
32 +
33 +var winServiceCounter atomic.Uint64
34 +
35 +func uniqueWinService(prefix string) string {
36 + return fmt.Sprintf("%s_%d_%d", prefix, os.Getpid(), winServiceCounter.Add(1))
37 +}
38 +
39 +func testWinShmServerConfig() windows.ServerConfig {
40 + cfg := testWinServerConfig()
41 + cfg.SupportedProfiles = protocol.ProfileBaseline | protocol.ProfileSHMHybrid
42 + cfg.PreferredProfiles = protocol.ProfileSHMHybrid
43 + return cfg
44 +}
45 +
46 +func testWinShmClientConfig() windows.ClientConfig {
47 + cfg := testWinClientConfig()
48 + cfg.SupportedProfiles = protocol.ProfileBaseline | protocol.ProfileSHMHybrid
49 + cfg.PreferredProfiles = protocol.ProfileSHMHybrid
50 + return cfg
51 +}
52 +
53 +func waitWinClientReady(t *testing.T, client *Client) {
54 + t.Helper()
55 +
56 + for i := 0; i < 100; i++ {
57 + client.Refresh()
58 + if client.Ready() {
59 + return
60 + }
61 + time.Sleep(10 * time.Millisecond)
62 + }
63 +
64 + t.Fatalf("client not ready, final state=%d", client.state)
65 +}
66 +
67 +func startTestServerWinWithConfig(
68 + service string,
69 + cfg windows.ServerConfig,
70 + expectedMethodCode uint16,
71 + handler DispatchHandler,
72 +) *winTestServer {
73 + s := NewServer(winTestRunDir, service, cfg, expectedMethodCode, handler)
74 + doneCh := make(chan struct{})
75 +
76 + go func() {
77 + defer close(doneCh)
78 + _ = s.Run()
79 + }()
80 +
81 + time.Sleep(200 * time.Millisecond)
82 +
83 + return &winTestServer{server: s, doneCh: doneCh}
84 +}
85 +
86 +func startTestIncrementServerWinWithConfig(service string, cfg windows.ServerConfig) *winTestServer {
87 + return startTestServerWinWithConfig(service, cfg, protocol.MethodIncrement, winIncrementDispatchHandler())
88 +}
89 +
90 +func startTestStringReverseServerWinWithConfig(service string, cfg windows.ServerConfig) *winTestServer {
91 + return startTestServerWinWithConfig(service, cfg, protocol.MethodStringReverse, winStringReverseDispatchHandler())
92 +}
93 +
94 +func startTestSnapshotServerWinWithConfig(service string, cfg windows.ServerConfig) *winTestServer {
95 + return startTestServerWinWithConfig(service, cfg, protocol.MethodCgroupsSnapshot, winSnapshotDispatchHandler())
96 +}
97 +
98 +func newRawWinShmClient(t *testing.T, profile uint32) (*Client, *windows.WinShmContext) {
99 + t.Helper()
100 +
101 + runDir := t.TempDir()
102 + service := uniqueWinService("go_win_raw_shm")
103 + sessionID := winServiceCounter.Add(1)
104 + reqCap := uint32(protocol.HeaderSize + 4096)
105 + respCap := uint32(protocol.HeaderSize + winResponseBufSize)
106 +
107 + server, err := windows.WinShmServerCreate(runDir, service, winAuthToken, sessionID, profile, reqCap, respCap)
108 + if err != nil {
109 + t.Fatalf("WinShmServerCreate failed: %v", err)
110 + }
111 +
112 + clientShm, err := windows.WinShmClientAttach(runDir, service, winAuthToken, sessionID, profile)
113 + if err != nil {
114 + server.WinShmDestroy()
115 + t.Fatalf("WinShmClientAttach failed: %v", err)
116 + }
117 +
118 + cfg := testWinShmClientConfig()
119 + client := NewIncrementClient(runDir, service, cfg)
120 + client.state = StateReady
121 + client.shm = clientShm
122 +
123 + t.Cleanup(func() {
124 + client.Close()
125 + server.WinShmDestroy()
126 + })
127 +
128 + return client, server
129 +}
130 +
131 +func encodeRawWinMessage(hdr protocol.Header, payload []byte) []byte {
132 + hdr.Magic = protocol.MagicMsg
133 + hdr.Version = protocol.Version
134 + hdr.HeaderLen = protocol.HeaderLen
135 + hdr.PayloadLen = uint32(len(payload))
136 +
137 + msg := make([]byte, protocol.HeaderSize+len(payload))
138 + hdr.Encode(msg[:protocol.HeaderSize])
139 + copy(msg[protocol.HeaderSize:], payload)
140 + return msg
141 +}
142 +
143 +func encodeWinIncrementBatchPayload(t *testing.T, values ...uint64) []byte {
144 + t.Helper()
145 +
146 + itemCount := uint32(len(values))
147 + if itemCount == 0 {
148 + return nil
149 + }
150 +
151 + bufSize := protocol.Align8(int(itemCount)*8) +
152 + int(itemCount)*protocol.IncrementPayloadSize +
153 + int(itemCount)*protocol.Alignment
154 + buf := make([]byte, bufSize)
155 + bb := protocol.NewBatchBuilder(buf, itemCount)
156 +
157 + for _, v := range values {
158 + var item [protocol.IncrementPayloadSize]byte
159 + if protocol.IncrementEncode(v, item[:]) == 0 {
160 + t.Fatal("IncrementEncode failed")
161 + }
162 + if err := bb.Add(item[:]); err != nil {
163 + t.Fatalf("batch add failed: %v", err)
164 + }
165 + }
166 +
167 + n, _ := bb.Finish()
168 + return buf[:n]
169 +}
170 +
171 +type winRawSessionServer struct {
172 + listener *windows.Listener
173 + doneCh chan error
174 +}
175 +
176 +type winHelloAckServer struct {
177 + doneCh chan error
178 + accepted atomic.Uint32
179 +}
180 +
181 +func encodeHelloAckPacketWithVersion(version uint16, status uint16, layoutVersion uint16) []byte {
182 + ack := protocol.HelloAck{
183 + LayoutVersion: layoutVersion,
184 + Flags: 0,
185 + ServerSupportedProfiles: protocol.ProfileBaseline,
186 + IntersectionProfiles: protocol.ProfileBaseline,
187 + SelectedProfile: protocol.ProfileBaseline,
188 + AgreedMaxRequestPayloadBytes: protocol.MaxPayloadDefault,
189 + AgreedMaxRequestBatchItems: 1,
190 + AgreedMaxResponsePayloadBytes: winResponseBufSize,
191 + AgreedMaxResponseBatchItems: 1,
192 + AgreedPacketSize: 0,
193 + SessionID: 77,
194 + }
195 +
196 + payload := make([]byte, 48)
197 + n := ack.Encode(payload)
198 + payload = payload[:n]
199 +
200 + hdr := protocol.Header{
201 + Magic: protocol.MagicMsg,
202 + Version: version,
203 + HeaderLen: protocol.HeaderSize,
204 + Kind: protocol.KindControl,
205 + Code: protocol.CodeHelloAck,
206 + TransportStatus: status,
207 + PayloadLen: uint32(len(payload)),
208 + ItemCount: 1,
209 + MessageID: 0,
210 + }
211 +
212 + pkt := make([]byte, protocol.HeaderSize+len(payload))
213 + hdr.Encode(pkt[:protocol.HeaderSize])
214 + copy(pkt[protocol.HeaderSize:], payload)
215 + return pkt
216 +}
217 +
218 +func startRawWinHelloAckServer(t *testing.T, service string, packet []byte) *winHelloAckServer {
219 + t.Helper()
220 +
221 + if err := os.MkdirAll(winTestRunDir, 0o700); err != nil {
222 + t.Fatalf("MkdirAll failed: %v", err)
223 + }
224 +
225 + pipeName, err := windows.BuildPipeName(winTestRunDir, service)
226 + if err != nil {
227 + t.Fatalf("BuildPipeName failed: %v", err)
228 + }
229 +
230 + pipe, err := createRawNamedPipe(
231 + &pipeName[0],
232 + rawPipeAccessDuplex,
233 + rawPipeTypeMessage|rawPipeReadModeMsg|rawPipeWait,
234 + 1,
235 + rawPipeBufSize,
236 + rawPipeBufSize,
237 + 0,
238 + )
239 + if err != nil {
240 + t.Fatalf("CreateNamedPipe failed: %v", err)
241 + }
242 +
243 + srv := &winHelloAckServer{
244 + doneCh: make(chan error, 1),
245 + }
246 +
247 + go func() {
248 + defer close(srv.doneCh)
249 + defer syscall.CloseHandle(pipe)
250 +
251 + err := connectRawNamedPipe(pipe)
252 + if err != nil && err != syscall.Errno(rawErrorPipeConnected) {
253 + srv.doneCh <- fmt.Errorf("ConnectNamedPipe failed: %w", err)
254 + return
255 + }
256 +
257 + srv.accepted.Store(1)
258 +
259 + helloBuf := make([]byte, 256)
260 + var helloN uint32
261 + if err := syscall.ReadFile(pipe, helloBuf, &helloN, nil); err != nil {
262 + srv.doneCh <- fmt.Errorf("ReadFile failed: %w", err)
263 + return
264 + }
265 + if helloN == 0 {
266 + srv.doneCh <- fmt.Errorf("ReadFile returned zero bytes")
267 + return
268 + }
269 +
270 + var written uint32
271 + if err := syscall.WriteFile(pipe, packet, &written, nil); err != nil {
272 + srv.doneCh <- fmt.Errorf("WriteFile failed: %w", err)
273 + return
274 + }
275 + if int(written) != len(packet) {
276 + srv.doneCh <- fmt.Errorf("short write: %d/%d", written, len(packet))
277 + return
278 + }
279 +
280 + srv.doneCh <- nil
281 + }()
282 +
283 + time.Sleep(200 * time.Millisecond)
284 + return srv
285 +}
286 +
287 +func createRawNamedPipe(name *uint16, openMode uint32, pipeMode uint32, maxInstances uint32,
288 + outBufferSize uint32, inBufferSize uint32, defaultTimeout uint32) (syscall.Handle, error) {
289 + r0, _, e1 := rawCreateNamedPipeW.Call(
290 + uintptr(unsafe.Pointer(name)),
291 + uintptr(openMode),
292 + uintptr(pipeMode),
293 + uintptr(maxInstances),
294 + uintptr(outBufferSize),
295 + uintptr(inBufferSize),
296 + uintptr(defaultTimeout),
297 + 0,
298 + )
299 + handle := syscall.Handle(r0)
300 + if handle == syscall.InvalidHandle {
301 + if e1 != syscall.Errno(0) {
302 + return syscall.InvalidHandle, e1
303 + }
304 + return syscall.InvalidHandle, syscall.EINVAL
305 + }
306 + return handle, nil
307 +}
308 +
309 +func connectRawNamedPipe(handle syscall.Handle) error {
310 + r1, _, e1 := rawConnectNamedPipe.Call(uintptr(handle), 0)
311 + if r1 != 0 {
312 + return nil
313 + }
314 + if e1 != syscall.Errno(0) {
315 + return e1
316 + }
317 + return syscall.EINVAL
318 +}
319 +
320 +func (s *winHelloAckServer) wait(t *testing.T) {
321 + t.Helper()
322 +
323 + if err := <-s.doneCh; err != nil {
324 + t.Fatalf("raw windows hello_ack server failed: %v", err)
325 + }
326 + if got := s.accepted.Load(); got != 1 {
327 + t.Fatalf("expected exactly one raw handshake accept, got %d", got)
328 + }
329 +}
330 +
331 +func startRawWinSessionServer(t *testing.T, service string, cfg windows.ServerConfig,
332 + handler func(*windows.Session, protocol.Header, []byte) error) *winRawSessionServer {
333 + return startRawWinSessionServerN(t, service, cfg, 1, handler)
334 +}
335 +
336 +func startRawWinSessionServerN(t *testing.T, service string, cfg windows.ServerConfig, accepts int,
337 + handler func(*windows.Session, protocol.Header, []byte) error) *winRawSessionServer {
338 + t.Helper()
339 +
340 + listener, err := windows.Listen(winTestRunDir, service, cfg)
341 + if err != nil {
342 + t.Fatalf("windows.Listen failed: %v", err)
343 + }
344 +
345 + srv := &winRawSessionServer{
346 + listener: listener,
347 + doneCh: make(chan error, 1),
348 + }
349 +
350 + go func() {
351 + defer close(srv.doneCh)
352 + defer listener.Close()
353 +
354 + for i := 0; i < accepts; i++ {
355 + session, err := listener.Accept()
356 + if err != nil {
357 + srv.doneCh <- err
358 + return
359 + }
360 +
361 + recvBuf := make([]byte, protocol.HeaderSize+int(cfg.MaxRequestPayloadBytes))
362 + hdr, payload, err := session.Receive(recvBuf)
363 + if err != nil {
364 + session.Close()
365 + srv.doneCh <- err
366 + return
367 + }
368 +
369 + if err := handler(session, hdr, payload); err != nil {
370 + session.Close()
371 + srv.doneCh <- err
372 + return
373 + }
374 +
375 + session.Close()
376 + }
377 +
378 + srv.doneCh <- nil
379 + }()
380 +
381 + time.Sleep(200 * time.Millisecond)
382 + return srv
383 +}
384 +
385 +func (s *winRawSessionServer) wait(t *testing.T) {
386 + t.Helper()
387 + if err := <-s.doneCh; err != nil {
388 + t.Fatalf("raw windows session server failed: %v", err)
389 + }
390 +}
391 +
392 +func winTestCgroupsHandler(request *protocol.CgroupsRequest, builder *protocol.CgroupsBuilder) bool {
393 + if request.LayoutVersion != 1 || request.Flags != 0 {
394 + return false
395 + }
396 +
397 + builder.SetHeader(1, 42)
398 +
399 + items := []struct {
400 + hash, options, enabled uint32
401 + name, path []byte
402 + }{
403 + {1001, 0, 1, []byte("docker-abc123"), []byte("/sys/fs/cgroup/docker/abc123")},
404 + {2002, 0, 1, []byte("k8s-pod-xyz"), []byte("/sys/fs/cgroup/kubepods/xyz")},
405 + {3003, 0, 0, []byte("systemd-user"), []byte("/sys/fs/cgroup/user.slice/user-1000")},
406 + }
407 +
408 + for _, item := range items {
409 + if err := builder.Add(item.hash, item.options, item.enabled, item.name, item.path); err != nil {
410 + return false
411 + }
412 + }
413 +
414 + return true
415 +}
416 +
417 +func winFailingCgroupsHandler(*protocol.CgroupsRequest, *protocol.CgroupsBuilder) bool {
418 + return false
419 +}
420 +
421 +func winIncrementDispatchHandler() DispatchHandler {
422 + return IncrementDispatch(func(v uint64) (uint64, bool) {
423 + return v + 1, true
424 + })
425 +}
426 +
427 +func winStringReverseDispatchHandler() DispatchHandler {
428 + return StringReverseDispatch(func(s string) (string, bool) {
429 + return winReverseString(s), true
430 + })
431 +}
432 +
433 +func winSnapshotDispatchHandler() DispatchHandler {
434 + return SnapshotDispatch(winTestCgroupsHandler, 3)
435 +}
436 +
437 +func winFailingSnapshotDispatchHandler() DispatchHandler {
438 + return SnapshotDispatch(winFailingCgroupsHandler, 3)
439 +}
src/go/pkg/netipc/service/raw/more_unix_test.go new
+1470
@@ -0,0 +1,1470 @@
1 +//go:build unix
2 +
3 +package raw
4 +
5 +import (
6 + "encoding/binary"
7 + "errors"
8 + "fmt"
9 + "os"
10 + "path/filepath"
11 + "sync/atomic"
12 + "syscall"
13 + "testing"
14 + "time"
15 +
16 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
17 + "github.com/netdata/netdata/go/plugins/pkg/netipc/transport/posix"
18 +)
19 +
20 +var unixServiceCounter atomic.Uint64
21 +
22 +func uniqueUnixService(prefix string) string {
23 + return fmt.Sprintf("%s_%d", prefix, unixServiceCounter.Add(1))
24 +}
25 +
26 +type unixRawSessionServer struct {
27 + doneCh chan error
28 +}
29 +
30 +type unixHelloAckServer struct {
31 + doneCh chan error
32 + accepted atomic.Uint32
33 +}
34 +
35 +func encodeHelloAckPacketWithVersion(version uint16, status uint16, layoutVersion uint16) []byte {
36 + ack := protocol.HelloAck{
37 + LayoutVersion: layoutVersion,
38 + Flags: 0,
39 + ServerSupportedProfiles: protocol.ProfileBaseline,
40 + IntersectionProfiles: protocol.ProfileBaseline,
41 + SelectedProfile: protocol.ProfileBaseline,
42 + AgreedMaxRequestPayloadBytes: protocol.MaxPayloadDefault,
43 + AgreedMaxRequestBatchItems: 1,
44 + AgreedMaxResponsePayloadBytes: responseBufSize,
45 + AgreedMaxResponseBatchItems: 1,
46 + AgreedPacketSize: 0,
47 + SessionID: 77,
48 + }
49 +
50 + payload := make([]byte, 48)
51 + n := ack.Encode(payload)
52 + payload = payload[:n]
53 +
54 + hdr := protocol.Header{
55 + Magic: protocol.MagicMsg,
56 + Version: version,
57 + HeaderLen: protocol.HeaderSize,
58 + Kind: protocol.KindControl,
59 + Code: protocol.CodeHelloAck,
60 + TransportStatus: status,
61 + PayloadLen: uint32(len(payload)),
62 + ItemCount: 1,
63 + MessageID: 0,
64 + }
65 +
66 + pkt := make([]byte, protocol.HeaderSize+len(payload))
67 + hdr.Encode(pkt[:protocol.HeaderSize])
68 + copy(pkt[protocol.HeaderSize:], payload)
69 + return pkt
70 +}
71 +
72 +func startRawPosixHelloAckServer(t *testing.T, service string, packet []byte) *unixHelloAckServer {
73 + t.Helper()
74 +
75 + ensureRunDir()
76 + cleanupAll(service)
77 +
78 + path := filepath.Join(testRunDir, service+".sock")
79 + _ = os.Remove(path)
80 +
81 + fd, err := syscall.Socket(syscall.AF_UNIX, syscall.SOCK_SEQPACKET, 0)
82 + if err != nil {
83 + t.Fatalf("socket failed: %v", err)
84 + }
85 +
86 + addr := &syscall.SockaddrUnix{Name: path}
87 + if err := syscall.Bind(fd, addr); err != nil {
88 + _ = syscall.Close(fd)
89 + t.Fatalf("bind failed: %v", err)
90 + }
91 + if err := syscall.Listen(fd, 4); err != nil {
92 + _ = syscall.Close(fd)
93 + t.Fatalf("listen failed: %v", err)
94 + }
95 +
96 + srv := &unixHelloAckServer{doneCh: make(chan error, 1)}
97 + go func() {
98 + defer close(srv.doneCh)
99 + defer syscall.Close(fd)
100 +
101 + connFD, _, err := syscall.Accept(fd)
102 + if err != nil {
103 + srv.doneCh <- err
104 + return
105 + }
106 + srv.accepted.Store(1)
107 + defer syscall.Close(connFD)
108 +
109 + buf := make([]byte, protocol.HeaderSize+128)
110 + if _, err := syscall.Read(connFD, buf); err != nil {
111 + srv.doneCh <- err
112 + return
113 + }
114 + if _, err := syscall.Write(connFD, packet); err != nil {
115 + srv.doneCh <- err
116 + return
117 + }
118 +
119 + srv.doneCh <- nil
120 + }()
121 +
122 + time.Sleep(100 * time.Millisecond)
123 + return srv
124 +}
125 +
126 +func (s *unixHelloAckServer) wait(t *testing.T) {
127 + t.Helper()
128 + if err := <-s.doneCh; err != nil {
129 + t.Fatalf("raw unix hello-ack server failed: %v", err)
130 + }
131 +}
132 +
133 +func startRawPosixSessionServerN(
134 + t *testing.T,
135 + service string,
136 + cfg posix.ServerConfig,
137 + accepts int,
138 + handler func(*posix.Session, protocol.Header, []byte) error,
139 +) *unixRawSessionServer {
140 + t.Helper()
141 + ensureRunDir()
142 + cleanupAll(service)
143 +
144 + listener, err := posix.Listen(testRunDir, service, cfg)
145 + if err != nil {
146 + t.Fatalf("Listen failed: %v", err)
147 + }
148 +
149 + srv := &unixRawSessionServer{doneCh: make(chan error, 1)}
150 + go func() {
151 + defer listener.Close()
152 +
153 + for i := 0; i < accepts; i++ {
154 + session, err := listener.Accept()
155 + if err != nil {
156 + srv.doneCh <- err
157 + return
158 + }
159 +
160 + recvBuf := make([]byte, protocol.HeaderSize+int(cfg.MaxRequestPayloadBytes))
161 + hdr, payload, err := session.Receive(recvBuf)
162 + if err != nil {
163 + session.Close()
164 + srv.doneCh <- err
165 + return
166 + }
167 +
168 + if err := handler(session, hdr, payload); err != nil {
169 + session.Close()
170 + srv.doneCh <- err
171 + return
172 + }
173 +
174 + session.Close()
175 + }
176 +
177 + srv.doneCh <- nil
178 + }()
179 +
180 + time.Sleep(200 * time.Millisecond)
181 + return srv
182 +}
183 +
184 +func startRawPosixSessionServer(
185 + t *testing.T,
186 + service string,
187 + cfg posix.ServerConfig,
188 + handler func(*posix.Session, protocol.Header, []byte) error,
189 +) *unixRawSessionServer {
190 + t.Helper()
191 + return startRawPosixSessionServerN(t, service, cfg, 1, handler)
192 +}
193 +
194 +func (s *unixRawSessionServer) wait(t *testing.T) {
195 + t.Helper()
196 + if err := <-s.doneCh; err != nil {
197 + t.Fatalf("raw unix session server failed: %v", err)
198 + }
199 +}
200 +
201 +func refreshUnixClientReady(t *testing.T, client *Client) {
202 + t.Helper()
203 + deadline := time.Now().Add(2 * time.Second)
204 + for time.Now().Before(deadline) {
205 + client.Refresh()
206 + if client.Ready() {
207 + return
208 + }
209 + time.Sleep(10 * time.Millisecond)
210 + }
211 + t.Fatalf("client not ready, state=%d", client.state)
212 +}
213 +
214 +func newRawPosixShmClient(t *testing.T) (*Client, *posix.ShmContext) {
215 + t.Helper()
216 +
217 + runDir := t.TempDir()
218 + service := uniqueUnixService("go_unix_raw_shm")
219 + sessionID := unixServiceCounter.Add(1)
220 + reqCap := uint32(protocol.HeaderSize + 4096)
221 + respCap := uint32(protocol.HeaderSize + responseBufSize)
222 +
223 + server, err := posix.ShmServerCreate(runDir, service, sessionID, reqCap, respCap)
224 + if err != nil {
225 + t.Fatalf("ShmServerCreate failed: %v", err)
226 + }
227 +
228 + clientShm, err := posix.ShmClientAttach(runDir, service, sessionID)
229 + if err != nil {
230 + server.ShmDestroy()
231 + t.Fatalf("ShmClientAttach failed: %v", err)
232 + }
233 +
234 + client := NewIncrementClient(runDir, service, testClientConfig())
235 + client.state = StateReady
236 + client.shm = clientShm
237 +
238 + t.Cleanup(func() {
239 + client.Close()
240 + server.ShmDestroy()
241 + })
242 +
243 + return client, server
244 +}
245 +
246 +func startServerWithWorkersSocketReady(
247 + service string,
248 + expectedMethodCode uint16,
249 + handler DispatchHandler,
250 + workers int,
251 +) *testServer {
252 + ensureRunDir()
253 + cleanupAll(service)
254 +
255 + s := NewServerWithWorkers(
256 + testRunDir,
257 + service,
258 + testServerConfig(),
259 + expectedMethodCode,
260 + handler,
261 + workers,
262 + )
263 + doneCh := make(chan struct{})
264 +
265 + go func() {
266 + defer close(doneCh)
267 + s.Run()
268 + }()
269 +
270 + waitUnixSocketReady(service)
271 + return &testServer{server: s, doneCh: doneCh}
272 +}
273 +
274 +func startServerWithConfigSocketReady(
275 + service string,
276 + cfg posix.ServerConfig,
277 + expectedMethodCode uint16,
278 + handler DispatchHandler,
279 +) *testServer {
280 + ensureRunDir()
281 + cleanupAll(service)
282 +
283 + s := NewServer(testRunDir, service, cfg, expectedMethodCode, handler)
284 + doneCh := make(chan struct{})
285 +
286 + go func() {
287 + defer close(doneCh)
288 + s.Run()
289 + }()
290 +
291 + waitUnixSocketReady(service)
292 + return &testServer{server: s, doneCh: doneCh}
293 +}
294 +
295 +func unixShmSessionPath(service string, sessionID uint64) string {
296 + return filepath.Join(testRunDir, fmt.Sprintf("%s-%016x.ipcshm", service, sessionID))
297 +}
298 +
299 +func TestUnixServerStopWhileIdle(t *testing.T) {
300 + svc := uniqueUnixService("go_unix_stop_idle")
301 + cleanupAll(svc)
302 + server := NewServer(
303 + testRunDir,
304 + svc,
305 + testServerConfig(),
306 + protocol.MethodCgroupsSnapshot,
307 + testSnapshotDispatch(),
308 + )
309 +
310 + done := make(chan error, 1)
311 + go func() {
312 + done <- server.Run()
313 + }()
314 +
315 + time.Sleep(100 * time.Millisecond)
316 + server.Stop()
317 +
318 + select {
319 + case err := <-done:
320 + if err != nil {
321 + t.Fatalf("Run after Stop = %v, want nil", err)
322 + }
323 + case <-time.After(2 * time.Second):
324 + t.Fatal("Run did not exit after Stop")
325 + }
326 +
327 + cleanupAll(svc)
328 +}
329 +
330 +func TestUnixTryConnectInvalidServiceReturnsDisconnected(t *testing.T) {
331 + client := NewSnapshotClient(testRunDir, "bad/service", testClientConfig())
332 + defer client.Close()
333 +
334 + if state := client.tryConnect(); state != StateDisconnected {
335 + t.Fatalf("tryConnect invalid service = %d, want %d", state, StateDisconnected)
336 + }
337 +}
338 +
339 +func TestUnixPollFdRejectsInvalidFdAndHangup(t *testing.T) {
340 + fds := []int{0, 0}
341 + if err := syscall.Pipe(fds); err != nil {
342 + t.Fatalf("Pipe failed: %v", err)
343 + }
344 + defer syscall.Close(fds[0])
345 +
346 + if err := syscall.Close(fds[1]); err != nil {
347 + t.Fatalf("close writer failed: %v", err)
348 + }
349 +
350 + if got := pollFd(fds[0], 50); got != -1 {
351 + t.Fatalf("pollFd(read end with closed writer) = %d, want -1", got)
352 + }
353 +
354 + closedFDs := []int{0, 0}
355 + if err := syscall.Pipe(closedFDs); err != nil {
356 + t.Fatalf("Pipe for POLLNVAL failed: %v", err)
357 + }
358 + if err := syscall.Close(closedFDs[0]); err != nil {
359 + t.Fatalf("close reader failed: %v", err)
360 + }
361 + defer syscall.Close(closedFDs[1])
362 +
363 + if got := pollFd(closedFDs[0], 50); got != -1 {
364 + t.Fatalf("pollFd(closed reader) = %d, want -1", got)
365 + }
366 +}
367 +
368 +func TestUnixRefreshFromBrokenReconnects(t *testing.T) {
369 + svc := uniqueUnixService("go_unix_refresh_broken")
370 + ts := startTestServer(svc, testSnapshotDispatch())
371 + defer ts.stop()
372 +
373 + client := NewSnapshotClient(testRunDir, svc, testClientConfig())
374 + defer client.Close()
375 +
376 + refreshUnixClientReady(t, client)
377 +
378 + client.session.Close()
379 + client.state = StateBroken
380 +
381 + changed := client.Refresh()
382 + if !changed {
383 + t.Fatal("refresh from BROKEN should change state")
384 + }
385 + if client.state != StateReady {
386 + t.Fatalf("expected READY after reconnect, got %d", client.state)
387 + }
388 + if client.reconnectCount < 1 {
389 + t.Fatalf("expected reconnect_count >= 1, got %d", client.reconnectCount)
390 + }
391 +}
392 +
393 +func TestUnixSingleResponseOverflowRetriesAndRecovers(t *testing.T) {
394 + svc := uniqueUnixService("go_unix_single_overflow")
395 +
396 + scfg := testServerConfig()
397 + scfg.MaxResponsePayloadBytes = 64
398 + var calls atomic.Int32
399 + ts := startTestServerUnixWithConfig(
400 + svc,
401 + scfg,
402 + protocol.MethodCgroupsSnapshot,
403 + SnapshotDispatch(func(request *protocol.CgroupsRequest, builder *protocol.CgroupsBuilder) bool {
404 + if request.LayoutVersion != 1 || request.Flags != 0 {
405 + return false
406 + }
407 + builder.SetHeader(1, uint64(calls.Add(1)))
408 + if calls.Load() == 1 {
409 + return builder.Add(
410 + 1001,
411 + 0,
412 + 1,
413 + []byte("oversized-snapshot"),
414 + []byte("/sys/fs/cgroup/this-path-is-intentionally-too-large-for-the-first-response"),
415 + ) == nil
416 + }
417 + return true
418 + }, 1),
419 + )
420 + defer ts.stop()
421 +
422 + client := NewSnapshotClient(testRunDir, svc, testClientConfig())
423 + defer client.Close()
424 + refreshUnixClientReady(t, client)
425 +
426 + view, err := client.CallSnapshot()
427 + if err != nil {
428 + t.Fatalf("CallSnapshot after single-response overflow failed: %v", err)
429 + }
430 + if view.ItemCount != 0 || view.Generation != 2 {
431 + t.Fatalf("CallSnapshot after single-response overflow returned item_count=%d generation=%d, want 0 and 2", view.ItemCount, view.Generation)
432 + }
433 + if client.state != StateReady {
434 + t.Fatalf("client state after single-response overflow = %d, want READY", client.state)
435 + }
436 + if client.reconnectCount < 1 {
437 + t.Fatalf("expected reconnect_count >= 1 after single-response overflow, got %d", client.reconnectCount)
438 + }
439 +}
440 +
441 +func TestUnixShmPrepareFailureFallsBackToBaseline(t *testing.T) {
442 + svc := uniqueUnixService("go_unix_shm_create_fail")
443 + obstruction := unixShmSessionPath(svc, 1)
444 +
445 + scfg := testUnixShmServerConfig()
446 + ts := startServerWithConfigSocketReady(
447 + svc,
448 + scfg,
449 + protocol.MethodIncrement,
450 + IncrementDispatch(func(v uint64) (uint64, bool) {
451 + return v + 1, true
452 + }),
453 + )
454 + defer func() {
455 + ts.stop()
456 + _ = os.RemoveAll(obstruction)
457 + }()
458 +
459 + if err := os.MkdirAll(filepath.Join(obstruction, "keep"), 0700); err != nil {
460 + t.Fatalf("mkdir SHM obstruction failed: %v", err)
461 + }
462 +
463 + badClient := NewIncrementClient(testRunDir, svc, testUnixShmClientConfig())
464 + defer badClient.Close()
465 +
466 + if changed := badClient.Refresh(); !changed {
467 + t.Fatal("Refresh should transition to READY when baseline stays viable")
468 + }
469 + if !badClient.Ready() {
470 + t.Fatal("client should stay usable over baseline after server SHM prepare failure")
471 + }
472 + if badClient.state != StateReady {
473 + t.Fatalf("client state after SHM prepare failure = %d, want READY", badClient.state)
474 + }
475 + if badClient.shm != nil {
476 + t.Fatalf("fallback baseline session must not attach SHM: selected=%#x session=%+v", badClient.session.SelectedProfile, badClient.session)
477 + }
478 + if badClient.session == nil || badClient.session.SelectedProfile != protocol.ProfileBaseline {
479 + t.Fatalf("selected profile after SHM prepare failure = %#v, want baseline", badClient.session)
480 + }
481 +
482 + if _, err := os.Stat(obstruction); err != nil {
483 + t.Fatalf("SHM obstruction should still exist after failed create, stat err = %v", err)
484 + }
485 + if err := os.RemoveAll(obstruction); err != nil {
486 + t.Fatalf("remove SHM obstruction failed: %v", err)
487 + }
488 +
489 + goodClient := NewIncrementClient(testRunDir, svc, testUnixShmClientConfig())
490 + defer goodClient.Close()
491 + refreshUnixClientReady(t, goodClient)
492 +
493 + if goodClient.shm == nil {
494 + t.Fatal("expected SHM attachment after obstruction removal")
495 + }
496 + got, err := goodClient.CallIncrement(4)
497 + if err != nil {
498 + t.Fatalf("CallIncrement after SHM create recovery failed: %v", err)
499 + }
500 + if got != 5 {
501 + t.Fatalf("CallIncrement after SHM create recovery = %d, want 5", got)
502 + }
503 +}
504 +
505 +func TestUnixServerRunRejectsInvalidServiceName(t *testing.T) {
506 + server := NewServer(
507 + testRunDir,
508 + "bad/service",
509 + testServerConfig(),
510 + protocol.MethodCgroupsSnapshot,
511 + testSnapshotDispatch(),
512 + )
513 + if err := server.Run(); !errors.Is(err, posix.ErrBadParam) {
514 + t.Fatalf("Run() invalid service name = %v, want %v", err, posix.ErrBadParam)
515 + }
516 +}
517 +
518 +func TestUnixServerRejectsSessionAtWorkerCapacity(t *testing.T) {
519 + svc := uniqueUnixService("go_unix_capacity")
520 + entered := make(chan struct{}, 1)
521 + release := make(chan struct{})
522 + released := false
523 + releaseOnce := func() {
524 + if !released {
525 + close(release)
526 + released = true
527 + }
528 + }
529 + defer releaseOnce()
530 +
531 + var handlerCalls atomic.Int32
532 + ts := startServerWithWorkersSocketReady(
533 + svc,
534 + protocol.MethodIncrement,
535 + IncrementDispatch(func(v uint64) (uint64, bool) {
536 + handlerCalls.Add(1)
537 + if v == 41 {
538 + select {
539 + case entered <- struct{}{}:
540 + default:
541 + }
542 + <-release
543 + }
544 + return v + 1, true
545 + }),
546 + 1,
547 + )
548 + defer ts.stop()
549 +
550 + client1 := NewIncrementClient(testRunDir, svc, testClientConfig())
551 + defer client1.Close()
552 + refreshUnixClientReady(t, client1)
553 +
554 + type callResult struct {
555 + got uint64
556 + err error
557 + }
558 + callDone := make(chan callResult, 1)
559 + go func() {
560 + got, err := client1.CallIncrement(41)
561 + callDone <- callResult{got: got, err: err}
562 + }()
563 +
564 + select {
565 + case <-entered:
566 + case <-time.After(2 * time.Second):
567 + t.Fatal("first client did not occupy the only worker slot")
568 + }
569 +
570 + ccfg := testClientConfig()
571 + session2, err := posix.Connect(testRunDir, svc, &ccfg)
572 + if err != nil {
573 + t.Fatalf("second raw connect failed: %v", err)
574 + }
575 + defer session2.Close()
576 +
577 + time.Sleep(100 * time.Millisecond)
578 +
579 + reqHdr := &protocol.Header{
580 + Kind: protocol.KindRequest,
581 + Code: protocol.MethodIncrement,
582 + ItemCount: 1,
583 + MessageID: 1,
584 + TransportStatus: protocol.StatusOK,
585 + }
586 + var reqPayload [protocol.IncrementPayloadSize]byte
587 + if protocol.IncrementEncode(9, reqPayload[:]) == 0 {
588 + t.Fatal("IncrementEncode failed")
589 + }
590 +
591 + sendErr := session2.Send(reqHdr, reqPayload[:])
592 + var recvErr error
593 + if sendErr == nil {
594 + recvBuf := make([]byte, protocol.HeaderSize+64)
595 + _, _, recvErr = session2.Receive(recvBuf)
596 + }
597 + if sendErr == nil && recvErr == nil {
598 + t.Fatal("second session should be rejected while the server is at worker capacity")
599 + }
600 + if handlerCalls.Load() != 1 {
601 + t.Fatalf("handler entered %d times, want 1 while second session is rejected", handlerCalls.Load())
602 + }
603 + session2.Close()
604 +
605 + releaseOnce()
606 + res := <-callDone
607 + if res.err != nil {
608 + t.Fatalf("first client call failed: %v", res.err)
609 + }
610 + if res.got != 42 {
611 + t.Fatalf("first client result = %d, want 42", res.got)
612 + }
613 + client1.Close()
614 +
615 + time.Sleep(100 * time.Millisecond)
616 +
617 + verify := NewIncrementClient(testRunDir, svc, testClientConfig())
618 + defer verify.Close()
619 + refreshUnixClientReady(t, verify)
620 +
621 + got, err := verify.CallIncrement(1)
622 + if err != nil {
623 + t.Fatalf("verification call after worker-capacity reject failed: %v", err)
624 + }
625 + if got != 2 {
626 + t.Fatalf("verification increment = %d, want 2", got)
627 + }
628 +}
629 +
630 +func TestUnixIdlePeerDisconnectKeepsServerHealthy(t *testing.T) {
631 + svc := uniqueUnixService("go_unix_idle_disconnect")
632 + ts := startTestServer(svc, testSnapshotDispatch())
633 + defer ts.stop()
634 +
635 + ccfg := testClientConfig()
636 + session, err := posix.Connect(testRunDir, svc, &ccfg)
637 + if err != nil {
638 + t.Fatalf("raw connect failed: %v", err)
639 + }
640 + session.Close()
641 +
642 + time.Sleep(200 * time.Millisecond)
643 +
644 + verify := NewSnapshotClient(testRunDir, svc, testClientConfig())
645 + defer verify.Close()
646 + refreshUnixClientReady(t, verify)
647 +
648 + view, err := verify.CallSnapshot()
649 + if err != nil {
650 + t.Fatalf("normal call after idle disconnect failed: %v", err)
651 + }
652 + if view.ItemCount != 3 {
653 + t.Fatalf("expected 3 items after idle disconnect, got %d", view.ItemCount)
654 + }
655 +}
656 +
657 +func TestUnixClientTransportWithoutSession(t *testing.T) {
658 + client := NewIncrementClient(testRunDir, uniqueUnixService("go_unix_transport"), testClientConfig())
659 + defer client.Close()
660 +
661 + hdr := protocol.Header{
662 + Kind: protocol.KindRequest,
663 + Code: protocol.MethodIncrement,
664 + ItemCount: 1,
665 + MessageID: 1,
666 + TransportStatus: protocol.StatusOK,
667 + }
668 +
669 + if err := client.transportSend(&hdr, nil); !errors.Is(err, protocol.ErrTruncated) {
670 + t.Fatalf("transportSend without session = %v, want ErrTruncated", err)
671 + }
672 + if _, _, err := client.transportReceive(); !errors.Is(err, protocol.ErrTruncated) {
673 + t.Fatalf("transportReceive without session = %v, want ErrTruncated", err)
674 + }
675 +}
676 +
677 +func TestUnixClientTransportReceiveShmError(t *testing.T) {
678 + client := NewIncrementClient(testRunDir, uniqueUnixService("go_unix_transport_shm_err"), testClientConfig())
679 + defer client.Close()
680 +
681 + client.state = StateReady
682 + client.shm = &posix.ShmContext{}
683 +
684 + if _, _, err := client.transportReceive(); !errors.Is(err, protocol.ErrTruncated) {
685 + t.Fatalf("transportReceive with invalid SHM context = %v, want %v", err, protocol.ErrTruncated)
686 + }
687 +}
688 +
689 +func TestUnixClientTransportReceiveShmRejectsShortMessage(t *testing.T) {
690 + client, serverShm := newRawPosixShmClient(t)
691 +
692 + if err := serverShm.ShmSend([]byte{1, 2, 3, 4}); err != nil {
693 + t.Fatalf("ShmSend failed: %v", err)
694 + }
695 +
696 + if _, _, err := client.transportReceive(); !errors.Is(err, protocol.ErrTruncated) {
697 + t.Fatalf("transportReceive short SHM message = %v, want %v", err, protocol.ErrTruncated)
698 + }
699 +}
700 +
701 +func TestUnixClientTransportReceiveShmRejectsBadHeader(t *testing.T) {
702 + client, serverShm := newRawPosixShmClient(t)
703 +
704 + msg := make([]byte, protocol.HeaderSize)
705 + if err := serverShm.ShmSend(msg); err != nil {
706 + t.Fatalf("ShmSend failed: %v", err)
707 + }
708 +
709 + if _, _, err := client.transportReceive(); !errors.Is(err, protocol.ErrBadMagic) {
710 + t.Fatalf("transportReceive bad SHM header = %v, want %v", err, protocol.ErrBadMagic)
711 + }
712 +}
713 +
714 +func TestUnixDispatchSingleUnsupportedMethods(t *testing.T) {
715 + server := NewServer(
716 + testRunDir,
717 + uniqueUnixService("go_unix_dispatch"),
718 + testServerConfig(),
719 + protocol.MethodCgroupsSnapshot,
720 + nil,
721 + )
722 + responseBuf := make([]byte, 128)
723 +
724 + for _, methodCode := range []uint16{
725 + protocol.MethodIncrement,
726 + protocol.MethodStringReverse,
727 + protocol.MethodCgroupsSnapshot,
728 + 0xffff,
729 + } {
730 + if n, err := server.dispatchSingle(methodCode, nil, responseBuf); err == nil || n != 0 {
731 + t.Fatalf("dispatchSingle(%d) = (%d, %v), want (0, error)", methodCode, n, err)
732 + }
733 + }
734 +
735 + server = NewServer(
736 + testRunDir,
737 + uniqueUnixService("go_unix_snapshot_dispatch"),
738 + testServerConfig(),
739 + protocol.MethodCgroupsSnapshot,
740 + testSnapshotDispatch(),
741 + )
742 + if n, err := server.dispatchSingle(protocol.MethodCgroupsSnapshot, nil, nil); err == nil || n != 0 {
743 + t.Fatalf("dispatchSingle snapshot with zero response buffer = (%d, %v), want (0, error)", n, err)
744 + }
745 +
746 + server = NewServer(
747 + testRunDir,
748 + uniqueUnixService("go_unix_snapshot_dispatch_zero"),
749 + testServerConfig(),
750 + protocol.MethodCgroupsSnapshot,
751 + SnapshotDispatch(testCgroupsHandler, 0),
752 + )
753 + if n, err := server.dispatchSingle(protocol.MethodCgroupsSnapshot, nil, nil); err == nil || n != 0 {
754 + t.Fatalf("dispatchSingle snapshot with derived zero max items = (%d, %v), want (0, error)", n, err)
755 + }
756 +}
757 +
758 +func TestUnixNonRequestTerminatesSession(t *testing.T) {
759 + svc := uniqueUnixService("go_unix_nonreq")
760 + ts := startTestServer(svc, testSnapshotDispatch())
761 + defer ts.stop()
762 +
763 + ccfg := testClientConfig()
764 + session, err := posix.Connect(testRunDir, svc, &ccfg)
765 + if err != nil {
766 + t.Fatalf("raw connect failed: %v", err)
767 + }
768 +
769 + badHdr := &protocol.Header{
770 + Kind: protocol.KindResponse,
771 + Code: protocol.MethodCgroupsSnapshot,
772 + ItemCount: 0,
773 + MessageID: 1,
774 + TransportStatus: protocol.StatusOK,
775 + }
776 + if err := session.Send(badHdr, nil); err != nil {
777 + t.Fatalf("send non-request failed: %v", err)
778 + }
779 +
780 + time.Sleep(200 * time.Millisecond)
781 +
782 + req := protocol.CgroupsRequest{LayoutVersion: 1, Flags: 0}
783 + var reqBuf [4]byte
784 + if req.Encode(reqBuf[:]) == 0 {
785 + t.Fatal("request encode failed")
786 + }
787 + goodHdr := &protocol.Header{
788 + Kind: protocol.KindRequest,
789 + Code: protocol.MethodCgroupsSnapshot,
790 + ItemCount: 1,
791 + MessageID: 2,
792 + TransportStatus: protocol.StatusOK,
793 + }
794 + _ = session.Send(goodHdr, reqBuf[:])
795 +
796 + recvBuf := make([]byte, 4096)
797 + if _, _, err := session.Receive(recvBuf); err == nil {
798 + t.Fatal("receive after non-request should fail")
799 + }
800 + session.Close()
801 +
802 + verify := NewSnapshotClient(testRunDir, svc, testClientConfig())
803 + defer verify.Close()
804 + refreshUnixClientReady(t, verify)
805 +
806 + view, err := verify.CallSnapshot()
807 + if err != nil {
808 + t.Fatalf("normal call after bad session failed: %v", err)
809 + }
810 + if view.ItemCount != 3 {
811 + t.Fatalf("expected 3 items after recovery, got %d", view.ItemCount)
812 + }
813 +}
814 +
815 +func TestUnixTruncatedRawRequestKeepsServerHealthy(t *testing.T) {
816 + svc := uniqueUnixService("go_unix_short_req")
817 + ts := startTestServer(svc, testSnapshotDispatch())
818 + defer ts.stop()
819 +
820 + ccfg := testClientConfig()
821 + session, err := posix.Connect(testRunDir, svc, &ccfg)
822 + if err != nil {
823 + t.Fatalf("raw connect failed: %v", err)
824 + }
825 + if _, err := syscall.SendmsgN(session.Fd(), []byte{1, 2, 3, 4}, nil, nil, 0); err != nil {
826 + session.Close()
827 + t.Fatalf("send truncated raw request failed: %v", err)
828 + }
829 + time.Sleep(100 * time.Millisecond)
830 + session.Close()
831 +
832 + time.Sleep(200 * time.Millisecond)
833 +
834 + verify := NewSnapshotClient(testRunDir, svc, testClientConfig())
835 + defer verify.Close()
836 + refreshUnixClientReady(t, verify)
837 +
838 + view, err := verify.CallSnapshot()
839 + if err != nil {
840 + t.Fatalf("normal call after truncated raw request failed: %v", err)
841 + }
842 + if view.ItemCount != 3 {
843 + t.Fatalf("expected 3 items after truncated raw request, got %d", view.ItemCount)
844 + }
845 +}
846 +
847 +func TestUnixPeerDisconnectDuringResponseKeepsServerHealthy(t *testing.T) {
848 + svc := uniqueUnixService("go_unix_send_fail")
849 + ts := startTestServerUnixWithConfig(
850 + svc,
851 + testServerConfig(),
852 + protocol.MethodIncrement,
853 + IncrementDispatch(func(v uint64) (uint64, bool) {
854 + time.Sleep(100 * time.Millisecond)
855 + return v + 1, true
856 + }),
857 + )
858 + defer ts.stop()
859 +
860 + ccfg := testClientConfig()
861 + session, err := posix.Connect(testRunDir, svc, &ccfg)
862 + if err != nil {
863 + t.Fatalf("raw connect failed: %v", err)
864 + }
865 +
866 + reqHdr := &protocol.Header{
867 + Kind: protocol.KindRequest,
868 + Code: protocol.MethodIncrement,
869 + ItemCount: 1,
870 + MessageID: 1,
871 + TransportStatus: protocol.StatusOK,
872 + }
873 + var reqPayload [protocol.IncrementPayloadSize]byte
874 + if protocol.IncrementEncode(41, reqPayload[:]) == 0 {
875 + t.Fatal("IncrementEncode failed")
876 + }
877 + if err := session.Send(reqHdr, reqPayload[:]); err != nil {
878 + t.Fatalf("send increment request failed: %v", err)
879 + }
880 + session.Close()
881 +
882 + time.Sleep(250 * time.Millisecond)
883 +
884 + verify := NewIncrementClient(testRunDir, svc, testClientConfig())
885 + defer verify.Close()
886 + refreshUnixClientReady(t, verify)
887 +
888 + got, err := verify.CallIncrement(9)
889 + if err != nil {
890 + t.Fatalf("normal call after peer disconnect failed: %v", err)
891 + }
892 + if got != 10 {
893 + t.Fatalf("increment after peer disconnect = %d, want 10", got)
894 + }
895 +}
896 +
897 +func TestUnixCallSnapshotWithMalformedTransportState(t *testing.T) {
898 + svc := uniqueUnixService("go_unix_malformed_state")
899 + ts := startTestServer(svc, testSnapshotDispatch())
900 + defer ts.stop()
901 +
902 + client := NewSnapshotClient(testRunDir, svc, testClientConfig())
903 + defer client.Close()
904 +
905 + refreshUnixClientReady(t, client)
906 +
907 + client.session.Close()
908 + client.session = nil
909 +
910 + view, err := client.CallSnapshot()
911 + if err != nil {
912 + t.Fatalf("expected reconnect to recover from nil session, got %v", err)
913 + }
914 + if view.ItemCount != 3 {
915 + t.Fatalf("expected recovered snapshot, got %d items", view.ItemCount)
916 + }
917 +}
918 +
919 +func TestUnixCallIncrementBatchWithMalformedTransportState(t *testing.T) {
920 + svc := uniqueUnixService("go_unix_batch_malformed_state")
921 +
922 + scfg := testServerConfig()
923 + scfg.MaxRequestBatchItems = 16
924 + scfg.MaxResponseBatchItems = 16
925 + s := NewServer(
926 + testRunDir,
927 + svc,
928 + scfg,
929 + protocol.MethodIncrement,
930 + pingPongIncrementDispatch(),
931 + )
932 + doneCh := make(chan struct{})
933 + go func() {
934 + defer close(doneCh)
935 + s.Run()
936 + }()
937 + time.Sleep(100 * time.Millisecond)
938 + defer func() {
939 + s.Stop()
940 + <-doneCh
941 + cleanupAll(svc)
942 + }()
943 +
944 + ccfg := testClientConfig()
945 + ccfg.MaxRequestBatchItems = 16
946 + ccfg.MaxResponseBatchItems = 16
947 + client := NewIncrementClient(testRunDir, svc, ccfg)
948 + defer client.Close()
949 +
950 + refreshUnixClientReady(t, client)
951 +
952 + client.session.Close()
953 + client.session = nil
954 +
955 + got, err := client.CallIncrementBatch([]uint64{1, 41})
956 + if err != nil {
957 + t.Fatalf("expected batch reconnect to recover from nil session, got %v", err)
958 + }
959 +
960 + want := []uint64{2, 42}
961 + if len(got) != len(want) {
962 + t.Fatalf("batch result len = %d, want %d", len(got), len(want))
963 + }
964 + for i, v := range want {
965 + if got[i] != v {
966 + t.Fatalf("batch[%d] = %d, want %d", i, got[i], v)
967 + }
968 + }
969 +}
970 +
971 +func TestUnixCallIncrementRejectsMalformedResponseEnvelope(t *testing.T) {
972 + cases := []struct {
973 + name string
974 + want error
975 + mutate func(*protocol.Header)
976 + }{
977 + {
978 + name: "bad kind",
979 + want: protocol.ErrBadKind,
980 + mutate: func(h *protocol.Header) {
981 + h.Kind = protocol.KindRequest
982 + },
983 + },
984 + {
985 + name: "bad code",
986 + want: protocol.ErrBadLayout,
987 + mutate: func(h *protocol.Header) {
988 + h.Code = protocol.MethodStringReverse
989 + },
990 + },
991 + {
992 + name: "bad status",
993 + want: protocol.ErrBadLayout,
994 + mutate: func(h *protocol.Header) {
995 + h.TransportStatus = protocol.StatusInternalError
996 + },
997 + },
998 + {
999 + name: "bad message_id",
1000 + want: protocol.ErrTruncated,
1001 + mutate: func(h *protocol.Header) {
1002 + h.MessageID++
1003 + },
1004 + },
1005 + }
1006 +
1007 + for _, tc := range cases {
1008 + t.Run(tc.name, func(t *testing.T) {
1009 + svc := uniqueUnixService("go_unix_bad_incr_resp")
1010 + srv := startRawPosixSessionServer(t, svc, testServerConfig(),
1011 + func(session *posix.Session, hdr protocol.Header, payload []byte) error {
1012 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodIncrement {
1013 + return fmt.Errorf("unexpected request header: %+v", hdr)
1014 + }
1015 +
1016 + var respPayload [protocol.IncrementPayloadSize]byte
1017 + protocol.IncrementEncode(42, respPayload[:])
1018 +
1019 + respHdr := protocol.Header{
1020 + Kind: protocol.KindResponse,
1021 + Code: protocol.MethodIncrement,
1022 + ItemCount: 1,
1023 + MessageID: hdr.MessageID,
1024 + TransportStatus: protocol.StatusOK,
1025 + }
1026 + tc.mutate(&respHdr)
1027 + return session.Send(&respHdr, respPayload[:])
1028 + })
1029 +
1030 + client := NewIncrementClient(testRunDir, svc, testClientConfig())
1031 + defer client.Close()
1032 +
1033 + refreshUnixClientReady(t, client)
1034 +
1035 + _, err := client.CallIncrement(41)
1036 + if !errors.Is(err, tc.want) {
1037 + t.Fatalf("CallIncrement error = %v, want %v", err, tc.want)
1038 + }
1039 +
1040 + srv.wait(t)
1041 + cleanupAll(svc)
1042 + })
1043 + }
1044 +}
1045 +
1046 +func TestUnixCallIncrementRejectsMalformedPayload(t *testing.T) {
1047 + svc := uniqueUnixService("go_unix_bad_incr_payload")
1048 + srv := startRawPosixSessionServer(t, svc, testServerConfig(),
1049 + func(session *posix.Session, hdr protocol.Header, payload []byte) error {
1050 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodIncrement {
1051 + return fmt.Errorf("unexpected request header: %+v", hdr)
1052 + }
1053 +
1054 + respHdr := protocol.Header{
1055 + Kind: protocol.KindResponse,
1056 + Code: protocol.MethodIncrement,
1057 + ItemCount: 1,
1058 + MessageID: hdr.MessageID,
1059 + TransportStatus: protocol.StatusOK,
1060 + }
1061 + return session.Send(&respHdr, []byte{1, 2, 3, 4})
1062 + })
1063 +
1064 + client := NewIncrementClient(testRunDir, svc, testClientConfig())
1065 + defer client.Close()
1066 +
1067 + refreshUnixClientReady(t, client)
1068 +
1069 + _, err := client.CallIncrement(41)
1070 + if !errors.Is(err, protocol.ErrTruncated) {
1071 + t.Fatalf("CallIncrement error = %v, want %v", err, protocol.ErrTruncated)
1072 + }
1073 +
1074 + srv.wait(t)
1075 + cleanupAll(svc)
1076 +}
1077 +
1078 +func TestUnixCallStringReverseRejectsMalformedResponseEnvelope(t *testing.T) {
1079 + cases := []struct {
1080 + name string
1081 + want error
1082 + mutate func(*protocol.Header)
1083 + }{
1084 + {
1085 + name: "bad kind",
1086 + want: protocol.ErrBadKind,
1087 + mutate: func(h *protocol.Header) {
1088 + h.Kind = protocol.KindRequest
1089 + },
1090 + },
1091 + {
1092 + name: "bad code",
1093 + want: protocol.ErrBadLayout,
1094 + mutate: func(h *protocol.Header) {
1095 + h.Code = protocol.MethodIncrement
1096 + },
1097 + },
1098 + {
1099 + name: "bad status",
1100 + want: protocol.ErrBadLayout,
1101 + mutate: func(h *protocol.Header) {
1102 + h.TransportStatus = protocol.StatusInternalError
1103 + },
1104 + },
1105 + {
1106 + name: "bad message_id",
1107 + want: protocol.ErrTruncated,
1108 + mutate: func(h *protocol.Header) {
1109 + h.MessageID++
1110 + },
1111 + },
1112 + }
1113 +
1114 + for _, tc := range cases {
1115 + t.Run(tc.name, func(t *testing.T) {
1116 + svc := uniqueUnixService("go_unix_bad_str_resp")
1117 + srv := startRawPosixSessionServer(t, svc, testServerConfig(),
1118 + func(session *posix.Session, hdr protocol.Header, payload []byte) error {
1119 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodStringReverse {
1120 + return fmt.Errorf("unexpected request header: %+v", hdr)
1121 + }
1122 +
1123 + respPayload := make([]byte, protocol.StringReverseHdrSize+len("olleh")+1)
1124 + if protocol.StringReverseEncode("olleh", respPayload) == 0 {
1125 + return fmt.Errorf("StringReverseEncode failed")
1126 + }
1127 +
1128 + respHdr := protocol.Header{
1129 + Kind: protocol.KindResponse,
1130 + Code: protocol.MethodStringReverse,
1131 + ItemCount: 1,
1132 + MessageID: hdr.MessageID,
1133 + TransportStatus: protocol.StatusOK,
1134 + }
1135 + tc.mutate(&respHdr)
1136 + return session.Send(&respHdr, respPayload)
1137 + })
1138 +
1139 + client := NewStringReverseClient(testRunDir, svc, testClientConfig())
1140 + defer client.Close()
1141 +
1142 + refreshUnixClientReady(t, client)
1143 +
1144 + _, err := client.CallStringReverse("hello")
1145 + if !errors.Is(err, tc.want) {
1146 + t.Fatalf("CallStringReverse error = %v, want %v", err, tc.want)
1147 + }
1148 +
1149 + srv.wait(t)
1150 + cleanupAll(svc)
1151 + })
1152 + }
1153 +}
1154 +
1155 +func TestUnixCallSnapshotRejectsMalformedPayload(t *testing.T) {
1156 + svc := uniqueUnixService("go_unix_bad_snapshot_payload")
1157 + srv := startRawPosixSessionServer(t, svc, testServerConfig(),
1158 + func(session *posix.Session, hdr protocol.Header, payload []byte) error {
1159 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodCgroupsSnapshot {
1160 + return fmt.Errorf("unexpected request header: %+v", hdr)
1161 + }
1162 +
1163 + respHdr := protocol.Header{
1164 + Kind: protocol.KindResponse,
1165 + Code: protocol.MethodCgroupsSnapshot,
1166 + ItemCount: 1,
1167 + MessageID: hdr.MessageID,
1168 + TransportStatus: protocol.StatusOK,
1169 + }
1170 + return session.Send(&respHdr, []byte{1, 2, 3, 4})
1171 + })
1172 +
1173 + client := NewSnapshotClient(testRunDir, svc, testClientConfig())
1174 + defer client.Close()
1175 +
1176 + refreshUnixClientReady(t, client)
1177 +
1178 + _, err := client.CallSnapshot()
1179 + if !errors.Is(err, protocol.ErrTruncated) {
1180 + t.Fatalf("CallSnapshot error = %v, want %v", err, protocol.ErrTruncated)
1181 + }
1182 +
1183 + srv.wait(t)
1184 + cleanupAll(svc)
1185 +}
1186 +
1187 +func TestUnixCallStringReverseRejectsMalformedPayload(t *testing.T) {
1188 + svc := uniqueUnixService("go_unix_bad_str_payload")
1189 + srv := startRawPosixSessionServer(t, svc, testServerConfig(),
1190 + func(session *posix.Session, hdr protocol.Header, payload []byte) error {
1191 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodStringReverse {
1192 + return fmt.Errorf("unexpected request header: %+v", hdr)
1193 + }
1194 +
1195 + respHdr := protocol.Header{
1196 + Kind: protocol.KindResponse,
1197 + Code: protocol.MethodStringReverse,
1198 + ItemCount: 1,
1199 + MessageID: hdr.MessageID,
1200 + TransportStatus: protocol.StatusOK,
1201 + }
1202 +
1203 + respPayload := make([]byte, protocol.StringReverseHdrSize+3)
1204 + binary.NativeEndian.PutUint32(respPayload[0:4], uint32(protocol.StringReverseHdrSize))
1205 + binary.NativeEndian.PutUint32(respPayload[4:8], 2)
1206 + copy(respPayload[8:], []byte{'o', 'k', 'x'})
1207 + return session.Send(&respHdr, respPayload)
1208 + })
1209 +
1210 + client := NewStringReverseClient(testRunDir, svc, testClientConfig())
1211 + defer client.Close()
1212 +
1213 + refreshUnixClientReady(t, client)
1214 +
1215 + _, err := client.CallStringReverse("hello")
1216 + if !errors.Is(err, protocol.ErrMissingNul) {
1217 + t.Fatalf("CallStringReverse error = %v, want %v", err, protocol.ErrMissingNul)
1218 + }
1219 +
1220 + srv.wait(t)
1221 + cleanupAll(svc)
1222 +}
1223 +
1224 +func TestUnixCallIncrementBatchRejectsMalformedResponseEnvelope(t *testing.T) {
1225 + cases := []struct {
1226 + name string
1227 + want error
1228 + mutate func(*protocol.Header)
1229 + }{
1230 + {
1231 + name: "bad kind",
1232 + want: protocol.ErrBadKind,
1233 + mutate: func(h *protocol.Header) {
1234 + h.Kind = protocol.KindRequest
1235 + },
1236 + },
1237 + {
1238 + name: "bad code",
1239 + want: protocol.ErrBadLayout,
1240 + mutate: func(h *protocol.Header) {
1241 + h.Code = protocol.MethodStringReverse
1242 + },
1243 + },
1244 + {
1245 + name: "bad status",
1246 + want: protocol.ErrBadLayout,
1247 + mutate: func(h *protocol.Header) {
1248 + h.TransportStatus = protocol.StatusInternalError
1249 + },
1250 + },
1251 + {
1252 + name: "bad message_id",
1253 + want: protocol.ErrTruncated,
1254 + mutate: func(h *protocol.Header) {
1255 + h.MessageID++
1256 + },
1257 + },
1258 + {
1259 + name: "missing batch flag",
1260 + want: protocol.ErrBadItemCount,
1261 + mutate: func(h *protocol.Header) {
1262 + h.Flags = 0
1263 + },
1264 + },
1265 + {
1266 + name: "wrong item count",
1267 + want: protocol.ErrBadItemCount,
1268 + mutate: func(h *protocol.Header) {
1269 + h.Flags = protocol.FlagBatch
1270 + h.ItemCount = 1
1271 + },
1272 + },
1273 + }
1274 +
1275 + for _, tc := range cases {
1276 + t.Run(tc.name, func(t *testing.T) {
1277 + cfg := testServerConfig()
1278 + cfg.MaxRequestBatchItems = 16
1279 + cfg.MaxResponseBatchItems = 16
1280 +
1281 + svc := uniqueUnixService("go_unix_bad_batch_resp")
1282 + srv := startRawPosixSessionServer(t, svc, cfg,
1283 + func(session *posix.Session, hdr protocol.Header, payload []byte) error {
1284 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodIncrement {
1285 + return fmt.Errorf("unexpected request header: %+v", hdr)
1286 + }
1287 +
1288 + respHdr := protocol.Header{
1289 + Kind: protocol.KindResponse,
1290 + Code: protocol.MethodIncrement,
1291 + Flags: protocol.FlagBatch,
1292 + ItemCount: hdr.ItemCount,
1293 + MessageID: hdr.MessageID,
1294 + TransportStatus: protocol.StatusOK,
1295 + }
1296 + tc.mutate(&respHdr)
1297 +
1298 + var itemA [protocol.IncrementPayloadSize]byte
1299 + var itemB [protocol.IncrementPayloadSize]byte
1300 + if protocol.IncrementEncode(2, itemA[:]) == 0 || protocol.IncrementEncode(3, itemB[:]) == 0 {
1301 + return fmt.Errorf("IncrementEncode failed")
1302 + }
1303 +
1304 + respBuf := make([]byte, 64)
1305 + bb := protocol.NewBatchBuilder(respBuf, 2)
1306 + if err := bb.Add(itemA[:]); err != nil {
1307 + return err
1308 + }
1309 + if err := bb.Add(itemB[:]); err != nil {
1310 + return err
1311 + }
1312 + n, _ := bb.Finish()
1313 + return session.Send(&respHdr, respBuf[:n])
1314 + })
1315 +
1316 + ccfg := testClientConfig()
1317 + ccfg.MaxRequestBatchItems = 16
1318 + ccfg.MaxResponseBatchItems = 16
1319 +
1320 + client := NewIncrementClient(testRunDir, svc, ccfg)
1321 + defer client.Close()
1322 +
1323 + refreshUnixClientReady(t, client)
1324 +
1325 + _, err := client.CallIncrementBatch([]uint64{1, 2})
1326 + if !errors.Is(err, tc.want) {
1327 + t.Fatalf("CallIncrementBatch error = %v, want %v", err, tc.want)
1328 + }
1329 +
1330 + srv.wait(t)
1331 + cleanupAll(svc)
1332 + })
1333 + }
1334 +}
1335 +
1336 +func TestUnixCallIncrementBatchRejectsMalformedPayload(t *testing.T) {
1337 + t.Run("bad batch directory", func(t *testing.T) {
1338 + cfg := testServerConfig()
1339 + cfg.MaxRequestBatchItems = 16
1340 + cfg.MaxResponseBatchItems = 16
1341 +
1342 + svc := uniqueUnixService("go_unix_bad_batch_dir")
1343 + srv := startRawPosixSessionServerN(t, svc, cfg, 2,
1344 + func(session *posix.Session, hdr protocol.Header, payload []byte) error {
1345 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodIncrement {
1346 + return fmt.Errorf("unexpected request header: %+v", hdr)
1347 + }
1348 +
1349 + respPayload := make([]byte, 24)
1350 + binary.NativeEndian.PutUint32(respPayload[0:4], 1)
1351 + binary.NativeEndian.PutUint32(respPayload[4:8], 4)
1352 + binary.NativeEndian.PutUint32(respPayload[8:12], 0)
1353 + binary.NativeEndian.PutUint32(respPayload[12:16], 4)
1354 + copy(respPayload[16:], []byte("payload!!"))
1355 +
1356 + respHdr := protocol.Header{
1357 + Kind: protocol.KindResponse,
1358 + Code: protocol.MethodIncrement,
1359 + Flags: protocol.FlagBatch,
1360 + ItemCount: 2,
1361 + MessageID: hdr.MessageID,
1362 + TransportStatus: protocol.StatusOK,
1363 + }
1364 + return session.Send(&respHdr, respPayload)
1365 + })
1366 +
1367 + ccfg := testClientConfig()
1368 + ccfg.MaxRequestBatchItems = 16
1369 + ccfg.MaxResponseBatchItems = 16
1370 + client := NewIncrementClient(testRunDir, svc, ccfg)
1371 + defer client.Close()
1372 +
1373 + refreshUnixClientReady(t, client)
1374 +
1375 + _, err := client.CallIncrementBatch([]uint64{1, 2})
1376 + if !errors.Is(err, protocol.ErrTruncated) {
1377 + t.Fatalf("CallIncrementBatch error = %v, want %v", err, protocol.ErrTruncated)
1378 + }
1379 +
1380 + srv.wait(t)
1381 + cleanupAll(svc)
1382 + })
1383 +
1384 + t.Run("truncated batch item", func(t *testing.T) {
1385 + cfg := testServerConfig()
1386 + cfg.MaxRequestBatchItems = 16
1387 + cfg.MaxResponseBatchItems = 16
1388 +
1389 + svc := uniqueUnixService("go_unix_bad_batch_payload")
1390 + srv := startRawPosixSessionServerN(t, svc, cfg, 1,
1391 + func(session *posix.Session, hdr protocol.Header, payload []byte) error {
1392 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodIncrement {
1393 + return fmt.Errorf("unexpected request header: %+v", hdr)
1394 + }
1395 +
1396 + respBuf := make([]byte, 64)
1397 + bb := protocol.NewBatchBuilder(respBuf, 1)
1398 + if err := bb.Add([]byte{1, 2, 3, 4}); err != nil {
1399 + return err
1400 + }
1401 + n, _ := bb.Finish()
1402 +
1403 + respHdr := protocol.Header{
1404 + Kind: protocol.KindResponse,
1405 + Code: protocol.MethodIncrement,
1406 + Flags: protocol.FlagBatch,
1407 + ItemCount: 1,
1408 + MessageID: hdr.MessageID,
1409 + TransportStatus: protocol.StatusOK,
1410 + }
1411 + return session.Send(&respHdr, respBuf[:n])
1412 + })
1413 +
1414 + ccfg := testClientConfig()
1415 + ccfg.MaxRequestBatchItems = 16
1416 + ccfg.MaxResponseBatchItems = 16
1417 + client := NewIncrementClient(testRunDir, svc, ccfg)
1418 + defer client.Close()
1419 +
1420 + refreshUnixClientReady(t, client)
1421 +
1422 + _, err := client.CallIncrementBatch([]uint64{1})
1423 + if !errors.Is(err, protocol.ErrTruncated) {
1424 + t.Fatalf("CallIncrementBatch error = %v, want %v", err, protocol.ErrTruncated)
1425 + }
1426 +
1427 + srv.wait(t)
1428 + cleanupAll(svc)
1429 + })
1430 +
1431 + t.Run("missing batch body", func(t *testing.T) {
1432 + cfg := testServerConfig()
1433 + cfg.MaxRequestBatchItems = 16
1434 + cfg.MaxResponseBatchItems = 16
1435 +
1436 + svc := uniqueUnixService("go_unix_bad_batch_body")
1437 + srv := startRawPosixSessionServerN(t, svc, cfg, 1,
1438 + func(session *posix.Session, hdr protocol.Header, payload []byte) error {
1439 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodIncrement {
1440 + return fmt.Errorf("unexpected request header: %+v", hdr)
1441 + }
1442 +
1443 + respHdr := protocol.Header{
1444 + Kind: protocol.KindResponse,
1445 + Code: protocol.MethodIncrement,
1446 + Flags: protocol.FlagBatch,
1447 + ItemCount: hdr.ItemCount,
1448 + MessageID: hdr.MessageID,
1449 + TransportStatus: protocol.StatusOK,
1450 + }
1451 + return session.Send(&respHdr, nil)
1452 + })
1453 +
1454 + ccfg := testClientConfig()
1455 + ccfg.MaxRequestBatchItems = 16
1456 + ccfg.MaxResponseBatchItems = 16
1457 + client := NewIncrementClient(testRunDir, svc, ccfg)
1458 + defer client.Close()
1459 +
1460 + refreshUnixClientReady(t, client)
1461 +
1462 + _, err := client.CallIncrementBatch([]uint64{1, 2})
1463 + if !errors.Is(err, protocol.ErrTruncated) {
1464 + t.Fatalf("CallIncrementBatch error = %v, want %v", err, protocol.ErrTruncated)
1465 + }
1466 +
1467 + srv.wait(t)
1468 + cleanupAll(svc)
1469 + })
1470 +}
src/go/pkg/netipc/service/raw/more_windows_test.go new
+1709
@@ -0,0 +1,1709 @@
1 +//go:build windows
2 +
3 +package raw
4 +
5 +import (
6 + "encoding/binary"
7 + "errors"
8 + "fmt"
9 + "testing"
10 + "time"
11 +
12 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
13 + windows "github.com/netdata/netdata/go/plugins/pkg/netipc/transport/windows"
14 +)
15 +
16 +func TestWinClientStatusFresh(t *testing.T) {
17 + client := NewSnapshotClient(winTestRunDir, uniqueWinService("go_win_fresh"), testWinClientConfig())
18 + defer client.Close()
19 +
20 + status := client.Status()
21 + if status.State != StateDisconnected {
22 + t.Fatalf("expected DISCONNECTED, got %d", status.State)
23 + }
24 + if status.ConnectCount != 0 || status.ReconnectCount != 0 || status.CallCount != 0 || status.ErrorCount != 0 {
25 + t.Fatalf("unexpected fresh counters: %+v", status)
26 + }
27 +}
28 +
29 +func TestWinClientLifecycle(t *testing.T) {
30 + svc := uniqueWinService("go_win_lifecycle")
31 +
32 + client := NewIncrementClient(winTestRunDir, svc, testWinClientConfig())
33 + defer client.Close()
34 +
35 + if client.state != StateDisconnected {
36 + t.Fatalf("expected DISCONNECTED, got %d", client.state)
37 + }
38 + if client.Ready() {
39 + t.Fatal("client should not be ready before refresh")
40 + }
41 +
42 + changed := client.Refresh()
43 + if !changed {
44 + t.Fatal("refresh without server should change state")
45 + }
46 + if client.state != StateNotFound {
47 + t.Fatalf("expected NOT_FOUND, got %d", client.state)
48 + }
49 +
50 + ts := startTestIncrementServerWinWithConfig(svc, testWinServerConfig())
51 +
52 + changed = client.Refresh()
53 + if !changed {
54 + t.Fatal("refresh with server should change state")
55 + }
56 + if client.state != StateReady {
57 + t.Fatalf("expected READY, got %d", client.state)
58 + }
59 + if !client.Ready() {
60 + t.Fatal("client should be ready")
61 + }
62 +
63 + status := client.Status()
64 + if status.ConnectCount != 1 {
65 + t.Fatalf("expected connect_count=1, got %d", status.ConnectCount)
66 + }
67 + if status.ReconnectCount != 0 {
68 + t.Fatalf("expected reconnect_count=0, got %d", status.ReconnectCount)
69 + }
70 +
71 + client.Close()
72 + if client.state != StateDisconnected {
73 + t.Fatalf("expected DISCONNECTED after close, got %d", client.state)
74 + }
75 +
76 + ts.stop()
77 +}
78 +
79 +func TestWinClientRefreshFromReadyNoop(t *testing.T) {
80 + svc := uniqueWinService("go_win_ready")
81 + ts := startTestIncrementServerWinWithConfig(svc, testWinServerConfig())
82 + defer ts.stop()
83 +
84 + client := NewIncrementClient(winTestRunDir, svc, testWinClientConfig())
85 + defer client.Close()
86 +
87 + client.Refresh()
88 + if !client.Ready() {
89 + t.Fatal("client not ready")
90 + }
91 +
92 + changed := client.Refresh()
93 + if changed {
94 + t.Fatal("refresh from READY should be a no-op")
95 + }
96 + if client.state != StateReady {
97 + t.Fatalf("expected READY, got %d", client.state)
98 + }
99 +}
100 +
101 +func TestWinServerStopWhileIdle(t *testing.T) {
102 + svc := uniqueWinService("go_win_stop_idle")
103 + server := NewServer(
104 + winTestRunDir,
105 + svc,
106 + testWinServerConfig(),
107 + protocol.MethodIncrement,
108 + winIncrementDispatchHandler(),
109 + )
110 +
111 + done := make(chan error, 1)
112 + go func() {
113 + done <- server.Run()
114 + }()
115 +
116 + deadline := time.Now().Add(2 * time.Second)
117 + for server.listener == nil && time.Now().Before(deadline) {
118 + time.Sleep(10 * time.Millisecond)
119 + }
120 + if server.listener == nil {
121 + t.Fatal("server listener did not start")
122 + }
123 +
124 + server.Stop()
125 +
126 + select {
127 + case err := <-done:
128 + if err != nil {
129 + t.Fatalf("Run after Stop = %v, want nil", err)
130 + }
131 + case <-time.After(2 * time.Second):
132 + t.Fatal("Run did not exit after Stop")
133 + }
134 +}
135 +
136 +func TestWinServerStopWithActiveClientAndRestart(t *testing.T) {
137 + svc := uniqueWinService("go_win_stop_active")
138 + ts1 := startTestSnapshotServerWinWithConfig(svc, testWinServerConfig())
139 +
140 + client := NewSnapshotClient(winTestRunDir, svc, testWinClientConfig())
141 + defer client.Close()
142 + waitWinClientReady(t, client)
143 +
144 + if _, err := client.CallSnapshot(); err != nil {
145 + t.Fatalf("first snapshot failed: %v", err)
146 + }
147 +
148 + stopDone := make(chan struct{})
149 + go func() {
150 + ts1.stop()
151 + close(stopDone)
152 + }()
153 +
154 + select {
155 + case <-stopDone:
156 + case <-time.After(2 * time.Second):
157 + t.Fatal("active server stop did not exit")
158 + }
159 +
160 + ts2 := startTestSnapshotServerWinWithConfig(svc, testWinServerConfig())
161 + defer ts2.stop()
162 +
163 + view, err := client.CallSnapshot()
164 + if err != nil {
165 + t.Fatalf("snapshot after restart failed: %v", err)
166 + }
167 + if view.ItemCount != 3 {
168 + t.Fatalf("expected 3 items after restart, got %d", view.ItemCount)
169 + }
170 +
171 + status := client.Status()
172 + if status.ReconnectCount < 1 {
173 + t.Fatalf("expected reconnect_count >= 1, got %d", status.ReconnectCount)
174 + }
175 +}
176 +
177 +func TestWinClientRefreshFromAuthFailed(t *testing.T) {
178 + svc := uniqueWinService("go_win_auth")
179 +
180 + scfg := testWinServerConfig()
181 + scfg.AuthToken = 0x1111
182 + ts := startTestIncrementServerWinWithConfig(svc, scfg)
183 + defer ts.stop()
184 +
185 + ccfg := testWinClientConfig()
186 + ccfg.AuthToken = 0x2222
187 + client := NewIncrementClient(winTestRunDir, svc, ccfg)
188 + defer client.Close()
189 +
190 + client.Refresh()
191 + if client.state != StateAuthFailed {
192 + t.Fatalf("expected AUTH_FAILED, got %d", client.state)
193 + }
194 +
195 + changed := client.Refresh()
196 + if changed {
197 + t.Fatal("refresh from AUTH_FAILED should be a no-op")
198 + }
199 +}
200 +
201 +func TestWinClientRefreshFromIncompatible(t *testing.T) {
202 + svc := uniqueWinService("go_win_incompat")
203 +
204 + scfg := testWinServerConfig()
205 + scfg.SupportedProfiles = protocol.ProfileSHMFutex
206 + ts := startTestIncrementServerWinWithConfig(svc, scfg)
207 + defer ts.stop()
208 +
209 + ccfg := testWinClientConfig()
210 + ccfg.SupportedProfiles = protocol.ProfileBaseline
211 + client := NewIncrementClient(winTestRunDir, svc, ccfg)
212 + defer client.Close()
213 +
214 + client.Refresh()
215 + if client.state != StateIncompatible {
216 + t.Fatalf("expected INCOMPATIBLE, got %d", client.state)
217 + }
218 +
219 + changed := client.Refresh()
220 + if changed {
221 + t.Fatal("refresh from INCOMPATIBLE should be a no-op")
222 + }
223 +}
224 +
225 +func TestWinClientRefreshFromProtocolVersionMismatch(t *testing.T) {
226 + svc := uniqueWinService("go_win_proto_incompat")
227 + packet := encodeHelloAckPacketWithVersion(protocol.Version+1, protocol.StatusOK, 1)
228 + srv := startRawWinHelloAckServer(t, svc, packet)
229 +
230 + client := NewSnapshotClient(winTestRunDir, svc, testWinClientConfig())
231 + defer client.Close()
232 +
233 + changed := client.Refresh()
234 + if !changed {
235 + t.Fatal("Refresh should move client into StateIncompatible")
236 + }
237 + if client.state != StateIncompatible {
238 + t.Fatalf("expected StateIncompatible, got %d", client.state)
239 + }
240 + if client.Ready() {
241 + t.Fatal("client should not be ready after protocol version mismatch")
242 + }
243 +
244 + srv.wait(t)
245 +
246 + changed = client.Refresh()
247 + if changed {
248 + t.Fatal("Refresh from StateIncompatible should be a no-op after protocol mismatch")
249 + }
250 + if client.state != StateIncompatible {
251 + t.Fatalf("expected StateIncompatible after second refresh, got %d", client.state)
252 + }
253 +}
254 +
255 +func TestWinClientRefreshFromBroken(t *testing.T) {
256 + svc := uniqueWinService("go_win_broken")
257 + ts := startTestIncrementServerWinWithConfig(svc, testWinServerConfig())
258 + defer ts.stop()
259 +
260 + client := NewIncrementClient(winTestRunDir, svc, testWinClientConfig())
261 + defer client.Close()
262 +
263 + client.Refresh()
264 + if !client.Ready() {
265 + t.Fatal("client not ready")
266 + }
267 +
268 + client.session.Close()
269 + client.state = StateBroken
270 +
271 + changed := client.Refresh()
272 + if !changed {
273 + t.Fatal("refresh from BROKEN should change state")
274 + }
275 + if client.state != StateReady {
276 + t.Fatalf("expected READY after reconnect, got %d", client.state)
277 + }
278 + if client.reconnectCount < 1 {
279 + t.Fatalf("expected reconnect_count >= 1, got %d", client.reconnectCount)
280 + }
281 +}
282 +
283 +func TestWinCallWithRetryNotReady(t *testing.T) {
284 + client := NewSnapshotClient(winTestRunDir, uniqueWinService("go_win_notready"), testWinClientConfig())
285 + defer client.Close()
286 +
287 + if _, err := client.CallSnapshot(); err == nil {
288 + t.Fatal("expected error when client is not ready")
289 + }
290 + if client.errorCount != 1 {
291 + t.Fatalf("expected error_count=1, got %d", client.errorCount)
292 + }
293 +}
294 +
295 +func TestWinClientTransportWithoutSession(t *testing.T) {
296 + client := NewIncrementClient(winTestRunDir, uniqueWinService("go_win_transport"), testWinClientConfig())
297 + defer client.Close()
298 +
299 + hdr := protocol.Header{
300 + Kind: protocol.KindRequest,
301 + Code: protocol.MethodIncrement,
302 + ItemCount: 1,
303 + MessageID: 1,
304 + TransportStatus: protocol.StatusOK,
305 + }
306 +
307 + if err := client.transportSend(&hdr, nil); !errors.Is(err, protocol.ErrTruncated) {
308 + t.Fatalf("transportSend without session = %v, want ErrTruncated", err)
309 + }
310 +
311 + if _, _, err := client.transportReceive(); !errors.Is(err, protocol.ErrTruncated) {
312 + t.Fatalf("transportReceive without session = %v, want ErrTruncated", err)
313 + }
314 +}
315 +
316 +func TestWinClientTransportReceiveWinShmError(t *testing.T) {
317 + client := NewIncrementClient(winTestRunDir, uniqueWinService("go_win_transport_shm_err"), testWinShmClientConfig())
318 + defer client.Close()
319 +
320 + client.state = StateReady
321 + client.shm = &windows.WinShmContext{}
322 +
323 + if _, _, err := client.transportReceive(); !errors.Is(err, protocol.ErrTruncated) {
324 + t.Fatalf("transportReceive with invalid WinSHM context = %v, want %v", err, protocol.ErrTruncated)
325 + }
326 +}
327 +
328 +func TestWinClientTransportReceiveWinShmRejectsShortMessage(t *testing.T) {
329 + client, serverShm := newRawWinShmClient(t, windows.WinShmProfileHybrid)
330 +
331 + if err := serverShm.WinShmSend([]byte{1, 2, 3, 4}); err != nil {
332 + t.Fatalf("WinShmSend failed: %v", err)
333 + }
334 +
335 + if _, _, err := client.transportReceive(); !errors.Is(err, protocol.ErrTruncated) {
336 + t.Fatalf("transportReceive short WinSHM message = %v, want %v", err, protocol.ErrTruncated)
337 + }
338 +}
339 +
340 +func TestWinClientTransportReceiveWinShmRejectsBadHeader(t *testing.T) {
341 + client, serverShm := newRawWinShmClient(t, windows.WinShmProfileHybrid)
342 +
343 + msg := make([]byte, protocol.HeaderSize)
344 + if err := serverShm.WinShmSend(msg); err != nil {
345 + t.Fatalf("WinShmSend failed: %v", err)
346 + }
347 +
348 + if _, _, err := client.transportReceive(); !errors.Is(err, protocol.ErrBadMagic) {
349 + t.Fatalf("transportReceive bad WinSHM header = %v, want %v", err, protocol.ErrBadMagic)
350 + }
351 +}
352 +
353 +func TestWinDoRawCallWinShmRejectsBadMessageID(t *testing.T) {
354 + client, serverShm := newRawWinShmClient(t, windows.WinShmProfileHybrid)
355 +
356 + serverDone := make(chan error, 1)
357 + go func() {
358 + reqBuf := make([]byte, protocol.HeaderSize+32)
359 + n, err := serverShm.WinShmReceive(reqBuf, 1000)
360 + if err != nil {
361 + serverDone <- err
362 + return
363 + }
364 +
365 + hdr, err := protocol.DecodeHeader(reqBuf[:n])
366 + if err != nil {
367 + serverDone <- err
368 + return
369 + }
370 +
371 + var respPayload [protocol.IncrementPayloadSize]byte
372 + if protocol.IncrementEncode(42, respPayload[:]) == 0 {
373 + serverDone <- fmt.Errorf("IncrementEncode failed")
374 + return
375 + }
376 +
377 + respHdr := protocol.Header{
378 + Kind: protocol.KindResponse,
379 + Code: protocol.MethodIncrement,
380 + ItemCount: 1,
381 + MessageID: hdr.MessageID + 1,
382 + TransportStatus: protocol.StatusOK,
383 + }
384 + serverDone <- serverShm.WinShmSend(encodeRawWinMessage(respHdr, respPayload[:]))
385 + }()
386 +
387 + var reqPayload [protocol.IncrementPayloadSize]byte
388 + if protocol.IncrementEncode(41, reqPayload[:]) == 0 {
389 + t.Fatal("IncrementEncode failed")
390 + }
391 +
392 + _, _, err := client.doRawCall(protocol.MethodIncrement, reqPayload[:])
393 + if !errors.Is(err, protocol.ErrBadLayout) {
394 + t.Fatalf("doRawCall bad WinSHM message_id = %v, want %v", err, protocol.ErrBadLayout)
395 + }
396 +
397 + if err := <-serverDone; err != nil {
398 + t.Fatalf("raw WinSHM server failed: %v", err)
399 + }
400 +}
401 +
402 +func TestWinCallIncrementBatchWinShmRejectsBadMessageID(t *testing.T) {
403 + client, serverShm := newRawWinShmClient(t, windows.WinShmProfileHybrid)
404 +
405 + serverDone := make(chan error, 1)
406 + go func() {
407 + reqBuf := make([]byte, protocol.HeaderSize+256)
408 + n, err := serverShm.WinShmReceive(reqBuf, 1000)
409 + if err != nil {
410 + serverDone <- err
411 + return
412 + }
413 +
414 + hdr, err := protocol.DecodeHeader(reqBuf[:n])
415 + if err != nil {
416 + serverDone <- err
417 + return
418 + }
419 +
420 + var itemA [protocol.IncrementPayloadSize]byte
421 + var itemB [protocol.IncrementPayloadSize]byte
422 + if protocol.IncrementEncode(2, itemA[:]) == 0 || protocol.IncrementEncode(3, itemB[:]) == 0 {
423 + serverDone <- fmt.Errorf("IncrementEncode failed")
424 + return
425 + }
426 +
427 + respBuf := make([]byte, 64)
428 + bb := protocol.NewBatchBuilder(respBuf, 2)
429 + if err := bb.Add(itemA[:]); err != nil {
430 + serverDone <- err
431 + return
432 + }
433 + if err := bb.Add(itemB[:]); err != nil {
434 + serverDone <- err
435 + return
436 + }
437 + respLen, _ := bb.Finish()
438 +
439 + respHdr := protocol.Header{
440 + Kind: protocol.KindResponse,
441 + Code: protocol.MethodIncrement,
442 + Flags: protocol.FlagBatch,
443 + ItemCount: 2,
444 + MessageID: hdr.MessageID + 1,
445 + TransportStatus: protocol.StatusOK,
446 + }
447 + serverDone <- serverShm.WinShmSend(encodeRawWinMessage(respHdr, respBuf[:respLen]))
448 + }()
449 +
450 + _, err := client.CallIncrementBatch([]uint64{1, 2})
451 + if !errors.Is(err, protocol.ErrBadLayout) {
452 + t.Fatalf("CallIncrementBatch bad WinSHM message_id = %v, want %v", err, protocol.ErrBadLayout)
453 + }
454 +
455 + if err := <-serverDone; err != nil {
456 + t.Fatalf("raw WinSHM server failed: %v", err)
457 + }
458 +}
459 +
460 +func TestWinCallIncrementBatchWinShmRejectsMalformedPayload(t *testing.T) {
461 + client, serverShm := newRawWinShmClient(t, windows.WinShmProfileHybrid)
462 +
463 + serverDone := make(chan error, 1)
464 + go func() {
465 + reqBuf := make([]byte, protocol.HeaderSize+256)
466 + n, err := serverShm.WinShmReceive(reqBuf, 1000)
467 + if err != nil {
468 + serverDone <- err
469 + return
470 + }
471 +
472 + hdr, err := protocol.DecodeHeader(reqBuf[:n])
473 + if err != nil {
474 + serverDone <- err
475 + return
476 + }
477 +
478 + respHdr := protocol.Header{
479 + Kind: protocol.KindResponse,
480 + Code: protocol.MethodIncrement,
481 + Flags: protocol.FlagBatch,
482 + ItemCount: 2,
483 + MessageID: hdr.MessageID,
484 + TransportStatus: protocol.StatusOK,
485 + }
486 + // ItemCount=2 requires a 16-byte aligned directory. An 8-byte payload
487 + // makes the client-side BatchItemGet() fail while decoding the response.
488 + serverDone <- serverShm.WinShmSend(encodeRawWinMessage(respHdr, make([]byte, 8)))
489 + }()
490 +
491 + _, err := client.CallIncrementBatch([]uint64{1, 2})
492 + if !errors.Is(err, protocol.ErrTruncated) {
493 + t.Fatalf("CallIncrementBatch malformed WinSHM payload = %v, want %v", err, protocol.ErrTruncated)
494 + }
495 +
496 + if err := <-serverDone; err != nil {
497 + t.Fatalf("raw WinSHM server failed: %v", err)
498 + }
499 +}
500 +
501 +func TestWinClientMaxReceiveMessageBytes(t *testing.T) {
502 + client := NewSnapshotClient(winTestRunDir, uniqueWinService("go_win_recvmax"), testWinClientConfig())
503 + defer client.Close()
504 +
505 + if got := client.maxReceiveMessageBytes(); got != protocol.HeaderSize+cacheResponseBufSize {
506 + t.Fatalf("default maxReceiveMessageBytes = %d, want %d", got, protocol.HeaderSize+cacheResponseBufSize)
507 + }
508 +
509 + client.config.MaxResponsePayloadBytes = 1234
510 + if got := client.maxReceiveMessageBytes(); got != protocol.HeaderSize+1234 {
511 + t.Fatalf("config maxReceiveMessageBytes = %d, want %d", got, protocol.HeaderSize+1234)
512 + }
513 +
514 + client.session = &windows.Session{MaxResponsePayloadBytes: 4321}
515 + if got := client.maxReceiveMessageBytes(); got != protocol.HeaderSize+4321 {
516 + t.Fatalf("session maxReceiveMessageBytes = %d, want %d", got, protocol.HeaderSize+4321)
517 + }
518 +}
519 +
520 +func TestWinServerDispatchSingleMissingHandler(t *testing.T) {
521 + server := &Server{expectedMethodCode: protocol.MethodIncrement}
522 + responseBuf := make([]byte, 128)
523 +
524 + for _, methodCode := range []uint16{
525 + protocol.MethodIncrement,
526 + protocol.MethodStringReverse,
527 + protocol.MethodCgroupsSnapshot,
528 + 0xFFFF,
529 + } {
530 + if n, err := server.dispatchSingle(methodCode, nil, responseBuf); !errors.Is(err, errHandlerFailed) || n != 0 {
531 + t.Fatalf("dispatchSingle(%d) = (%d, %v), want (0, %v)", methodCode, n, err, errHandlerFailed)
532 + }
533 + }
534 +}
535 +
536 +func TestWinServerDispatchSingleSnapshotZeroCapacity(t *testing.T) {
537 + server := &Server{
538 + expectedMethodCode: protocol.MethodCgroupsSnapshot,
539 + handler: winSnapshotDispatchHandler(),
540 + }
541 +
542 + req := protocol.CgroupsRequest{LayoutVersion: 1, Flags: 0}
543 + var reqBuf [4]byte
544 + if req.Encode(reqBuf[:]) == 0 {
545 + t.Fatal("request encode failed")
546 + }
547 + if n, err := server.dispatchSingle(protocol.MethodCgroupsSnapshot, reqBuf[:], nil); err == nil || n != 0 {
548 + t.Fatalf("dispatchSingle snapshot with zero response buffer = (%d, %v), want (0, error)", n, err)
549 + }
550 +}
551 +
552 +func TestWinCgroupsCall(t *testing.T) {
553 + svc := uniqueWinService("go_win_cgroups")
554 + ts := startTestSnapshotServerWinWithConfig(svc, testWinServerConfig())
555 + defer ts.stop()
556 +
557 + client := NewSnapshotClient(winTestRunDir, svc, testWinClientConfig())
558 + defer client.Close()
559 +
560 + client.Refresh()
561 + if !client.Ready() {
562 + t.Fatal("client not ready")
563 + }
564 +
565 + view, err := client.CallSnapshot()
566 + if err != nil {
567 + t.Fatalf("snapshot call failed: %v", err)
568 + }
569 + if view.ItemCount != 3 {
570 + t.Fatalf("expected 3 items, got %d", view.ItemCount)
571 + }
572 + if view.SystemdEnabled != 1 {
573 + t.Fatalf("expected systemd_enabled=1, got %d", view.SystemdEnabled)
574 + }
575 + if view.Generation != 42 {
576 + t.Fatalf("expected generation=42, got %d", view.Generation)
577 + }
578 +
579 + item0, err := view.Item(0)
580 + if err != nil {
581 + t.Fatalf("item 0 error: %v", err)
582 + }
583 + if item0.Hash != 1001 || item0.Enabled != 1 {
584 + t.Fatalf("unexpected item0: %+v", item0)
585 + }
586 + if item0.Name.String() != "docker-abc123" {
587 + t.Fatalf("unexpected item0 name: %q", item0.Name.String())
588 + }
589 + if item0.Path.String() != "/sys/fs/cgroup/docker/abc123" {
590 + t.Fatalf("unexpected item0 path: %q", item0.Path.String())
591 + }
592 +
593 + status := client.Status()
594 + if status.CallCount != 1 || status.ErrorCount != 0 {
595 + t.Fatalf("unexpected status after snapshot: %+v", status)
596 + }
597 +}
598 +
599 +func TestWinCallIncrementBatch(t *testing.T) {
600 + svc := uniqueWinService("go_win_batch")
601 + cfg := testWinServerConfig()
602 + cfg.MaxRequestBatchItems = 16
603 + cfg.MaxResponseBatchItems = 16
604 + ts := startTestServerWinWithConfig(svc, cfg, protocol.MethodIncrement, winIncrementDispatchHandler())
605 + defer ts.stop()
606 +
607 + ccfg := testWinClientConfig()
608 + ccfg.MaxRequestBatchItems = 16
609 + ccfg.MaxResponseBatchItems = 16
610 + client := NewIncrementClient(winTestRunDir, svc, ccfg)
611 + defer client.Close()
612 +
613 + client.Refresh()
614 + if !client.Ready() {
615 + t.Fatal("client not ready")
616 + }
617 +
618 + got, err := client.CallIncrementBatch([]uint64{1, 41, 99, 1000})
619 + if err != nil {
620 + t.Fatalf("CallIncrementBatch failed: %v", err)
621 + }
622 +
623 + want := []uint64{2, 42, 100, 1001}
624 + if len(got) != len(want) {
625 + t.Fatalf("expected %d results, got %d", len(want), len(got))
626 + }
627 + for i := range want {
628 + if got[i] != want[i] {
629 + t.Fatalf("result[%d] = %d, want %d", i, got[i], want[i])
630 + }
631 + }
632 +}
633 +
634 +func TestWinCallIncrementBatchEmpty(t *testing.T) {
635 + client := NewIncrementClient(winTestRunDir, uniqueWinService("go_win_empty_batch"), testWinClientConfig())
636 + defer client.Close()
637 +
638 + results, err := client.CallIncrementBatch(nil)
639 + if err != nil {
640 + t.Fatalf("expected nil error for nil batch, got %v", err)
641 + }
642 + if results != nil {
643 + t.Fatalf("expected nil results, got %v", results)
644 + }
645 +
646 + results, err = client.CallIncrementBatch([]uint64{})
647 + if err != nil {
648 + t.Fatalf("expected nil error for empty slice, got %v", err)
649 + }
650 + if results != nil {
651 + t.Fatalf("expected nil results, got %v", results)
652 + }
653 +}
654 +
655 +func TestWinRetryOnClosedSession(t *testing.T) {
656 + svc := uniqueWinService("go_win_retry")
657 + ts := startTestSnapshotServerWinWithConfig(svc, testWinServerConfig())
658 + defer ts.stop()
659 +
660 + client := NewSnapshotClient(winTestRunDir, svc, testWinClientConfig())
661 + defer client.Close()
662 +
663 + client.Refresh()
664 + if !client.Ready() {
665 + t.Fatal("client not ready")
666 + }
667 +
668 + if _, err := client.CallSnapshot(); err != nil {
669 + t.Fatalf("first call failed: %v", err)
670 + }
671 +
672 + client.session.Close()
673 +
674 + view, err := client.CallSnapshot()
675 + if err != nil {
676 + t.Fatalf("retry call failed: %v", err)
677 + }
678 + if view.ItemCount != 3 {
679 + t.Fatalf("expected 3 items after retry, got %d", view.ItemCount)
680 + }
681 +
682 + status := client.Status()
683 + if status.ReconnectCount < 1 {
684 + t.Fatalf("expected reconnect_count >= 1, got %d", status.ReconnectCount)
685 + }
686 +}
687 +
688 +func TestWinHandlerFailure(t *testing.T) {
689 + svc := uniqueWinService("go_win_handler_fail")
690 + ts := startTestServerWinWithConfig(svc, testWinServerConfig(), protocol.MethodCgroupsSnapshot, winFailingSnapshotDispatchHandler())
691 + defer ts.stop()
692 +
693 + client := NewSnapshotClient(winTestRunDir, svc, testWinClientConfig())
694 + defer client.Close()
695 +
696 + client.Refresh()
697 + if !client.Ready() {
698 + t.Fatal("client not ready")
699 + }
700 +
701 + if _, err := client.CallSnapshot(); err == nil {
702 + t.Fatal("expected error when handler fails")
703 + }
704 +
705 + status := client.Status()
706 + if status.ErrorCount < 1 {
707 + t.Fatalf("expected error_count >= 1, got %d", status.ErrorCount)
708 + }
709 +}
710 +
711 +func TestWinStatusReporting(t *testing.T) {
712 + svc := uniqueWinService("go_win_status")
713 + ts := startTestSnapshotServerWinWithConfig(svc, testWinServerConfig())
714 +
715 + client := NewSnapshotClient(winTestRunDir, svc, testWinClientConfig())
716 + client.Refresh()
717 + if !client.Ready() {
718 + t.Fatal("client not ready")
719 + }
720 +
721 + s0 := client.Status()
722 + if s0.ConnectCount != 1 || s0.CallCount != 0 || s0.ErrorCount != 0 {
723 + t.Fatalf("unexpected initial status: %+v", s0)
724 + }
725 +
726 + for i := 0; i < 3; i++ {
727 + if _, err := client.CallSnapshot(); err != nil {
728 + t.Fatalf("call %d failed: %v", i, err)
729 + }
730 + }
731 +
732 + s1 := client.Status()
733 + if s1.CallCount != 3 || s1.ErrorCount != 0 {
734 + t.Fatalf("unexpected status after calls: %+v", s1)
735 + }
736 +
737 + client.Close()
738 + if _, err := client.CallSnapshot(); err == nil {
739 + t.Fatal("expected error on disconnected client")
740 + }
741 +
742 + s2 := client.Status()
743 + if s2.ErrorCount != 1 {
744 + t.Fatalf("expected error_count=1, got %+v", s2)
745 + }
746 +
747 + ts.stop()
748 +}
749 +
750 +func TestWinCacheLookupBeforeRefresh(t *testing.T) {
751 + cache := NewCache(winTestRunDir, uniqueWinService("go_win_cache_empty"), testWinClientConfig())
752 + defer cache.Close()
753 +
754 + if cache.Ready() {
755 + t.Fatal("cache should not be ready")
756 + }
757 + if _, found := cache.Lookup(123, "anything"); found {
758 + t.Fatal("lookup before refresh should miss")
759 + }
760 +}
761 +
762 +func TestWinCacheStatusFresh(t *testing.T) {
763 + cache := NewCache(winTestRunDir, uniqueWinService("go_win_cache_status"), testWinClientConfig())
764 + defer cache.Close()
765 +
766 + status := cache.Status()
767 + if status.Populated || status.ItemCount != 0 || status.RefreshSuccessCount != 0 || status.RefreshFailureCount != 0 || status.LastRefreshTs != 0 {
768 + t.Fatalf("unexpected fresh cache status: %+v", status)
769 + }
770 +}
771 +
772 +func TestWinCacheFullRoundTrip(t *testing.T) {
773 + svc := uniqueWinService("go_win_cache_rt")
774 + ts := startTestSnapshotServerWinWithConfig(svc, testWinServerConfig())
775 +
776 + cache := NewCache(winTestRunDir, svc, testWinClientConfig())
777 + if cache.Ready() {
778 + t.Fatal("cache should not be ready before refresh")
779 + }
780 +
781 + time.Sleep(2 * time.Millisecond)
782 +
783 + if !cache.Refresh() {
784 + t.Fatal("refresh should succeed")
785 + }
786 + if !cache.Ready() {
787 + t.Fatal("cache should be ready after refresh")
788 + }
789 +
790 + item, found := cache.Lookup(1001, "docker-abc123")
791 + if !found {
792 + t.Fatal("expected cached item")
793 + }
794 + if item.Path != "/sys/fs/cgroup/docker/abc123" {
795 + t.Fatalf("unexpected cached path: %q", item.Path)
796 + }
797 +
798 + status := cache.Status()
799 + if !status.Populated || status.ItemCount != 3 || status.SystemdEnabled != 1 || status.Generation != 42 || status.RefreshSuccessCount != 1 || status.RefreshFailureCount != 0 || status.ConnectionState != StateReady || status.LastRefreshTs <= 0 {
800 + t.Fatalf("unexpected cache status: %+v", status)
801 + }
802 +
803 + cache.Close()
804 + ts.stop()
805 +}
806 +
807 +func TestWinCacheRefreshNoServer(t *testing.T) {
808 + cache := NewCache(winTestRunDir, uniqueWinService("go_win_cache_noserver"), testWinClientConfig())
809 + defer cache.Close()
810 +
811 + if cache.Refresh() {
812 + t.Fatal("refresh should fail with no server")
813 + }
814 + if cache.Ready() {
815 + t.Fatal("cache should not be ready")
816 + }
817 + if cache.Status().RefreshFailureCount != 1 {
818 + t.Fatalf("expected failure_count=1, got %d", cache.Status().RefreshFailureCount)
819 + }
820 +}
821 +
822 +func TestWinCacheRefreshFailurePreserves(t *testing.T) {
823 + svc := uniqueWinService("go_win_cache_preserve")
824 + ts := startTestSnapshotServerWinWithConfig(svc, testWinServerConfig())
825 +
826 + cache := NewCache(winTestRunDir, svc, testWinClientConfig())
827 + if !cache.Refresh() {
828 + t.Fatal("first refresh should succeed")
829 + }
830 + if !cache.Ready() {
831 + t.Fatal("cache should be ready")
832 + }
833 + if _, found := cache.Lookup(1001, "docker-abc123"); !found {
834 + t.Fatal("expected first cached item")
835 + }
836 +
837 + cache.client.Close()
838 + ts.stop()
839 +
840 + if cache.Refresh() {
841 + t.Fatal("refresh should fail without server")
842 + }
843 + if !cache.Ready() {
844 + t.Fatal("cache should preserve readiness")
845 + }
846 + if _, found := cache.Lookup(1001, "docker-abc123"); !found {
847 + t.Fatal("old cache data should be preserved")
848 + }
849 +
850 + status := cache.Status()
851 + if status.RefreshSuccessCount != 1 || status.RefreshFailureCount < 1 {
852 + t.Fatalf("unexpected preserved cache status: %+v", status)
853 + }
854 +
855 + cache.Close()
856 +}
857 +
858 +func TestWinCacheReconnectRebuilds(t *testing.T) {
859 + svc := uniqueWinService("go_win_cache_reconnect")
860 + ts := startTestSnapshotServerWinWithConfig(svc, testWinServerConfig())
861 + defer ts.stop()
862 +
863 + cache := NewCache(winTestRunDir, svc, testWinClientConfig())
864 + defer cache.Close()
865 +
866 + if !cache.Refresh() {
867 + t.Fatal("first refresh should succeed")
868 + }
869 +
870 + cache.client.session.Close()
871 +
872 + if !cache.Refresh() {
873 + t.Fatal("refresh after reconnect should succeed")
874 + }
875 + if cache.Status().RefreshSuccessCount != 2 {
876 + t.Fatalf("expected success_count=2, got %d", cache.Status().RefreshSuccessCount)
877 + }
878 +}
879 +
880 +func TestWinCacheLookupHashNameMismatch(t *testing.T) {
881 + svc := uniqueWinService("go_win_cache_lookup")
882 + ts := startTestSnapshotServerWinWithConfig(svc, testWinServerConfig())
883 + defer ts.stop()
884 +
885 + cache := NewCache(winTestRunDir, svc, testWinClientConfig())
886 + defer cache.Close()
887 +
888 + if !cache.Refresh() {
889 + t.Fatal("refresh should succeed")
890 + }
891 +
892 + if _, found := cache.Lookup(1001, "wrong-name"); found {
893 + t.Fatal("lookup with wrong name should miss")
894 + }
895 + if _, found := cache.Lookup(9999, "docker-abc123"); found {
896 + t.Fatal("lookup with wrong hash should miss")
897 + }
898 + if item, found := cache.Lookup(1001, "docker-abc123"); !found || item.Hash != 1001 {
899 + t.Fatalf("expected exact hash+name match, got found=%v item=%+v", found, item)
900 + }
901 +}
902 +
903 +func TestWinCacheCloseResetsState(t *testing.T) {
904 + svc := uniqueWinService("go_win_cache_close")
905 + ts := startTestSnapshotServerWinWithConfig(svc, testWinServerConfig())
906 +
907 + cache := NewCache(winTestRunDir, svc, testWinClientConfig())
908 + if !cache.Refresh() {
909 + t.Fatal("refresh should succeed")
910 + }
911 + cache.Close()
912 +
913 + if cache.Ready() {
914 + t.Fatal("cache should not be ready after close")
915 + }
916 + if _, found := cache.Lookup(1001, "docker-abc123"); found {
917 + t.Fatal("lookup after close should miss")
918 + }
919 +
920 + ts.stop()
921 +}
922 +
923 +func TestWinCacheCustomAndDefaultBufferSize(t *testing.T) {
924 + customCfg := testWinClientConfig()
925 + customCfg.MaxResponsePayloadBytes = 1024
926 + custom := NewCache(winTestRunDir, uniqueWinService("go_win_custom_buf"), customCfg)
927 + defer custom.Close()
928 +
929 + if got, want := custom.client.maxReceiveMessageBytes(), protocol.HeaderSize+1024; got != want {
930 + t.Fatalf("custom buffer size = %d, want %d", got, want)
931 + }
932 +
933 + defaultCfg := testWinClientConfig()
934 + defaultCfg.MaxResponsePayloadBytes = 0
935 + def := NewCache(winTestRunDir, uniqueWinService("go_win_default_buf"), defaultCfg)
936 + defer def.Close()
937 +
938 + if got, want := def.client.maxReceiveMessageBytes(), protocol.HeaderSize+cacheResponseBufSize; got != want {
939 + t.Fatalf("default buffer size = %d, want %d", got, want)
940 + }
941 +}
942 +
943 +func TestWinCacheLargeDataset(t *testing.T) {
944 + svc := uniqueWinService("go_win_cache_large")
945 + const itemCount = 512
946 +
947 + cfg := testWinServerConfig()
948 + cfg.MaxResponsePayloadBytes = 256 * itemCount
949 + ts := startTestServerWinWithConfig(svc, cfg, protocol.MethodCgroupsSnapshot, SnapshotDispatch(func(request *protocol.CgroupsRequest, builder *protocol.CgroupsBuilder) bool {
950 + if request.LayoutVersion != 1 || request.Flags != 0 {
951 + return false
952 + }
953 + builder.SetHeader(1, 100)
954 +
955 + for i := uint32(0); i < itemCount; i++ {
956 + name := fmt.Sprintf("cgroup-%d", i)
957 + path := fmt.Sprintf("/sys/fs/cgroup/test/%d", i)
958 + enabled := uint32(1)
959 + if i%3 == 0 {
960 + enabled = 0
961 + }
962 + if err := builder.Add(i+1000, 0, enabled, []byte(name), []byte(path)); err != nil {
963 + return false
964 + }
965 + }
966 + return true
967 + }, itemCount))
968 + defer ts.stop()
969 +
970 + ccfg := testWinClientConfig()
971 + ccfg.MaxResponsePayloadBytes = 256 * itemCount
972 +
973 + cache := NewCache(winTestRunDir, svc, ccfg)
974 + defer cache.Close()
975 +
976 + if !cache.Refresh() {
977 + t.Fatal("refresh should succeed")
978 + }
979 + if got := cache.Status().ItemCount; got != itemCount {
980 + t.Fatalf("expected %d items, got %d", itemCount, got)
981 + }
982 +
983 + for i := uint32(0); i < itemCount; i++ {
984 + name := fmt.Sprintf("cgroup-%d", i)
985 + item, found := cache.Lookup(i+1000, name)
986 + if !found {
987 + t.Fatalf("item %d not found", i)
988 + }
989 + wantPath := fmt.Sprintf("/sys/fs/cgroup/test/%d", i)
990 + if item.Path != wantPath {
991 + t.Fatalf("item %d path = %q, want %q", i, item.Path, wantPath)
992 + }
993 + }
994 +}
995 +
996 +func TestWinIsErrorHelpers(t *testing.T) {
997 + if !isConnectError(windows.ErrConnect) || !isConnectError(windows.ErrCreatePipe) {
998 + t.Fatal("connect errors should match")
999 + }
1000 + if isConnectError(windows.ErrAuthFailed) || isConnectError(windows.ErrNoProfile) || isConnectError(nil) {
1001 + t.Fatal("non-connect errors should not match connect classification")
1002 + }
1003 +
1004 + if !isAuthError(windows.ErrAuthFailed) || isAuthError(windows.ErrConnect) || isAuthError(nil) {
1005 + t.Fatal("auth error classification mismatch")
1006 + }
1007 +
1008 + if !isIncompatibleError(windows.ErrNoProfile) || !isIncompatibleError(windows.ErrIncompatible) ||
1009 + isIncompatibleError(windows.ErrAuthFailed) || isIncompatibleError(nil) {
1010 + t.Fatal("incompatible error classification mismatch")
1011 + }
1012 +}
1013 +
1014 +func TestWinNonRequestTerminatesSession(t *testing.T) {
1015 + svc := uniqueWinService("go_win_nonreq")
1016 + ts := startTestSnapshotServerWinWithConfig(svc, testWinServerConfig())
1017 + defer ts.stop()
1018 +
1019 + ccfg := testWinClientConfig()
1020 + session, err := windows.Connect(winTestRunDir, svc, &ccfg)
1021 + if err != nil {
1022 + t.Fatalf("raw connect failed: %v", err)
1023 + }
1024 +
1025 + badHdr := &protocol.Header{
1026 + Kind: protocol.KindResponse,
1027 + Code: protocol.MethodCgroupsSnapshot,
1028 + ItemCount: 0,
1029 + MessageID: 1,
1030 + TransportStatus: protocol.StatusOK,
1031 + }
1032 + if err := session.Send(badHdr, nil); err != nil {
1033 + t.Fatalf("send non-request failed: %v", err)
1034 + }
1035 +
1036 + time.Sleep(200 * time.Millisecond)
1037 +
1038 + req := protocol.CgroupsRequest{LayoutVersion: 1, Flags: 0}
1039 + var reqBuf [4]byte
1040 + if req.Encode(reqBuf[:]) == 0 {
1041 + t.Fatal("request encode failed")
1042 + }
1043 + goodHdr := &protocol.Header{
1044 + Kind: protocol.KindRequest,
1045 + Code: protocol.MethodCgroupsSnapshot,
1046 + ItemCount: 1,
1047 + MessageID: 2,
1048 + TransportStatus: protocol.StatusOK,
1049 + }
1050 + _ = session.Send(goodHdr, reqBuf[:])
1051 +
1052 + recvBuf := make([]byte, 4096)
1053 + if _, _, err := session.Receive(recvBuf); err == nil {
1054 + t.Fatal("receive after non-request should fail")
1055 + }
1056 + session.Close()
1057 +
1058 + verify := NewSnapshotClient(winTestRunDir, svc, testWinClientConfig())
1059 + defer verify.Close()
1060 +
1061 + verify.Refresh()
1062 + if !verify.Ready() {
1063 + t.Fatal("server should still accept new clients after bad session")
1064 + }
1065 +
1066 + view, err := verify.CallSnapshot()
1067 + if err != nil {
1068 + t.Fatalf("normal call after bad session failed: %v", err)
1069 + }
1070 + if view.ItemCount != 3 {
1071 + t.Fatalf("expected 3 items after recovery, got %d", view.ItemCount)
1072 + }
1073 +}
1074 +
1075 +func TestWinPeerDisconnectDuringResponseKeepsServerHealthy(t *testing.T) {
1076 + svc := uniqueWinService("go_win_send_fail")
1077 + ts := startTestServerWinWithConfig(svc, testWinServerConfig(), protocol.MethodIncrement, IncrementDispatch(func(v uint64) (uint64, bool) {
1078 + time.Sleep(100 * time.Millisecond)
1079 + return v + 1, true
1080 + }))
1081 + defer ts.stop()
1082 +
1083 + ccfg := testWinClientConfig()
1084 + session, err := windows.Connect(winTestRunDir, svc, &ccfg)
1085 + if err != nil {
1086 + t.Fatalf("raw connect failed: %v", err)
1087 + }
1088 +
1089 + reqHdr := &protocol.Header{
1090 + Kind: protocol.KindRequest,
1091 + Code: protocol.MethodIncrement,
1092 + ItemCount: 1,
1093 + MessageID: 1,
1094 + TransportStatus: protocol.StatusOK,
1095 + }
1096 + var reqPayload [protocol.IncrementPayloadSize]byte
1097 + if protocol.IncrementEncode(41, reqPayload[:]) == 0 {
1098 + t.Fatal("IncrementEncode failed")
1099 + }
1100 + if err := session.Send(reqHdr, reqPayload[:]); err != nil {
1101 + t.Fatalf("send increment request failed: %v", err)
1102 + }
1103 + session.Close()
1104 +
1105 + time.Sleep(250 * time.Millisecond)
1106 +
1107 + verify := NewIncrementClient(winTestRunDir, svc, testWinClientConfig())
1108 + defer verify.Close()
1109 +
1110 + verify.Refresh()
1111 + if !verify.Ready() {
1112 + t.Fatal("server should still accept new clients after peer disconnect during response")
1113 + }
1114 +
1115 + got, err := verify.CallIncrement(9)
1116 + if err != nil {
1117 + t.Fatalf("normal call after peer disconnect failed: %v", err)
1118 + }
1119 + if got != 10 {
1120 + t.Fatalf("increment after peer disconnect = %d, want 10", got)
1121 + }
1122 +}
1123 +
1124 +func TestWinCallSnapshotWithMalformedTransportState(t *testing.T) {
1125 + svc := uniqueWinService("go_win_malformed_state")
1126 + ts := startTestSnapshotServerWinWithConfig(svc, testWinServerConfig())
1127 + defer ts.stop()
1128 +
1129 + client := NewSnapshotClient(winTestRunDir, svc, testWinClientConfig())
1130 + defer client.Close()
1131 +
1132 + client.Refresh()
1133 + if !client.Ready() {
1134 + t.Fatal("client not ready")
1135 + }
1136 +
1137 + client.session.Close()
1138 + client.session = nil
1139 + view, err := client.CallSnapshot()
1140 + if err != nil {
1141 + t.Fatalf("expected reconnect to recover from nil session, got %v", err)
1142 + }
1143 + if view.ItemCount != 3 {
1144 + t.Fatalf("expected recovered snapshot, got %d items", view.ItemCount)
1145 + }
1146 +}
1147 +
1148 +func TestWinCallIncrementBatchWithMalformedTransportState(t *testing.T) {
1149 + svc := uniqueWinService("go_win_batch_malformed_state")
1150 +
1151 + scfg := testWinServerConfig()
1152 + scfg.MaxRequestBatchItems = 16
1153 + scfg.MaxResponseBatchItems = 16
1154 + ccfg := testWinClientConfig()
1155 + ccfg.MaxRequestBatchItems = 16
1156 + ccfg.MaxResponseBatchItems = 16
1157 +
1158 + ts := startTestServerWinWithConfig(svc, scfg, protocol.MethodIncrement, winIncrementDispatchHandler())
1159 + defer ts.stop()
1160 +
1161 + client := NewIncrementClient(winTestRunDir, svc, ccfg)
1162 + defer client.Close()
1163 +
1164 + client.Refresh()
1165 + if !client.Ready() {
1166 + t.Fatal("client not ready")
1167 + }
1168 +
1169 + client.session.Close()
1170 + client.session = nil
1171 +
1172 + got, err := client.CallIncrementBatch([]uint64{1, 41})
1173 + if err != nil {
1174 + t.Fatalf("expected batch reconnect to recover from nil session, got %v", err)
1175 + }
1176 +
1177 + want := []uint64{2, 42}
1178 + if len(got) != len(want) {
1179 + t.Fatalf("batch result len = %d, want %d", len(got), len(want))
1180 + }
1181 + for i, v := range want {
1182 + if got[i] != v {
1183 + t.Fatalf("batch[%d] = %d, want %d", i, got[i], v)
1184 + }
1185 + }
1186 +}
1187 +
1188 +func TestWinCallIncrementRejectsMalformedResponseEnvelope(t *testing.T) {
1189 + cases := []struct {
1190 + name string
1191 + want error
1192 + mutate func(*protocol.Header)
1193 + }{
1194 + {
1195 + name: "bad kind",
1196 + want: protocol.ErrBadKind,
1197 + mutate: func(h *protocol.Header) {
1198 + h.Kind = protocol.KindRequest
1199 + },
1200 + },
1201 + {
1202 + name: "bad code",
1203 + want: protocol.ErrBadLayout,
1204 + mutate: func(h *protocol.Header) {
1205 + h.Code = protocol.MethodStringReverse
1206 + },
1207 + },
1208 + {
1209 + name: "bad status",
1210 + want: protocol.ErrBadLayout,
1211 + mutate: func(h *protocol.Header) {
1212 + h.TransportStatus = protocol.StatusInternalError
1213 + },
1214 + },
1215 + {
1216 + name: "bad message_id",
1217 + want: protocol.ErrTruncated,
1218 + mutate: func(h *protocol.Header) {
1219 + h.MessageID++
1220 + },
1221 + },
1222 + }
1223 +
1224 + for _, tc := range cases {
1225 + t.Run(tc.name, func(t *testing.T) {
1226 + svc := uniqueWinService("go_win_bad_resp")
1227 + srv := startRawWinSessionServer(t, svc, testWinServerConfig(),
1228 + func(session *windows.Session, hdr protocol.Header, payload []byte) error {
1229 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodIncrement {
1230 + return fmt.Errorf("unexpected request header: %+v", hdr)
1231 + }
1232 +
1233 + var respPayload [protocol.IncrementPayloadSize]byte
1234 + protocol.IncrementEncode(42, respPayload[:])
1235 +
1236 + respHdr := protocol.Header{
1237 + Kind: protocol.KindResponse,
1238 + Code: protocol.MethodIncrement,
1239 + ItemCount: 1,
1240 + MessageID: hdr.MessageID,
1241 + TransportStatus: protocol.StatusOK,
1242 + }
1243 + tc.mutate(&respHdr)
1244 + return session.Send(&respHdr, respPayload[:])
1245 + })
1246 +
1247 + client := NewIncrementClient(winTestRunDir, svc, testWinClientConfig())
1248 + defer client.Close()
1249 +
1250 + client.Refresh()
1251 + if !client.Ready() {
1252 + t.Fatal("client not ready")
1253 + }
1254 +
1255 + _, err := client.CallIncrement(41)
1256 + if !errors.Is(err, tc.want) {
1257 + t.Fatalf("CallIncrement error = %v, want %v", err, tc.want)
1258 + }
1259 +
1260 + srv.wait(t)
1261 + })
1262 + }
1263 +}
1264 +
1265 +func TestWinCallIncrementRejectsMalformedPayload(t *testing.T) {
1266 + svc := uniqueWinService("go_win_bad_incr_payload")
1267 + srv := startRawWinSessionServer(t, svc, testWinServerConfig(),
1268 + func(session *windows.Session, hdr protocol.Header, payload []byte) error {
1269 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodIncrement {
1270 + return fmt.Errorf("unexpected request header: %+v", hdr)
1271 + }
1272 +
1273 + respHdr := protocol.Header{
1274 + Kind: protocol.KindResponse,
1275 + Code: protocol.MethodIncrement,
1276 + ItemCount: 1,
1277 + MessageID: hdr.MessageID,
1278 + TransportStatus: protocol.StatusOK,
1279 + }
1280 + return session.Send(&respHdr, []byte{1, 2, 3, 4})
1281 + })
1282 +
1283 + client := NewIncrementClient(winTestRunDir, svc, testWinClientConfig())
1284 + defer client.Close()
1285 +
1286 + client.Refresh()
1287 + if !client.Ready() {
1288 + t.Fatal("client not ready")
1289 + }
1290 +
1291 + _, err := client.CallIncrement(41)
1292 + if !errors.Is(err, protocol.ErrTruncated) {
1293 + t.Fatalf("CallIncrement error = %v, want %v", err, protocol.ErrTruncated)
1294 + }
1295 +
1296 + srv.wait(t)
1297 +}
1298 +
1299 +func TestWinCallStringReverseRejectsMalformedResponseEnvelope(t *testing.T) {
1300 + cases := []struct {
1301 + name string
1302 + want error
1303 + mutate func(*protocol.Header)
1304 + }{
1305 + {
1306 + name: "bad kind",
1307 + want: protocol.ErrBadKind,
1308 + mutate: func(h *protocol.Header) {
1309 + h.Kind = protocol.KindRequest
1310 + },
1311 + },
1312 + {
1313 + name: "bad code",
1314 + want: protocol.ErrBadLayout,
1315 + mutate: func(h *protocol.Header) {
1316 + h.Code = protocol.MethodIncrement
1317 + },
1318 + },
1319 + {
1320 + name: "bad status",
1321 + want: protocol.ErrBadLayout,
1322 + mutate: func(h *protocol.Header) {
1323 + h.TransportStatus = protocol.StatusInternalError
1324 + },
1325 + },
1326 + {
1327 + name: "bad message_id",
1328 + want: protocol.ErrTruncated,
1329 + mutate: func(h *protocol.Header) {
1330 + h.MessageID++
1331 + },
1332 + },
1333 + }
1334 +
1335 + for _, tc := range cases {
1336 + t.Run(tc.name, func(t *testing.T) {
1337 + svc := uniqueWinService("go_win_bad_str_resp")
1338 + srv := startRawWinSessionServer(t, svc, testWinServerConfig(),
1339 + func(session *windows.Session, hdr protocol.Header, payload []byte) error {
1340 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodStringReverse {
1341 + return fmt.Errorf("unexpected request header: %+v", hdr)
1342 + }
1343 +
1344 + respPayload := make([]byte, protocol.StringReverseHdrSize+len("olleh")+1)
1345 + if protocol.StringReverseEncode("olleh", respPayload) == 0 {
1346 + return fmt.Errorf("StringReverseEncode failed")
1347 + }
1348 +
1349 + respHdr := protocol.Header{
1350 + Kind: protocol.KindResponse,
1351 + Code: protocol.MethodStringReverse,
1352 + ItemCount: 1,
1353 + MessageID: hdr.MessageID,
1354 + TransportStatus: protocol.StatusOK,
1355 + }
1356 + tc.mutate(&respHdr)
1357 + return session.Send(&respHdr, respPayload)
1358 + })
1359 +
1360 + client := NewStringReverseClient(winTestRunDir, svc, testWinClientConfig())
1361 + defer client.Close()
1362 +
1363 + client.Refresh()
1364 + if !client.Ready() {
1365 + t.Fatal("client not ready")
1366 + }
1367 +
1368 + _, err := client.CallStringReverse("hello")
1369 + if !errors.Is(err, tc.want) {
1370 + t.Fatalf("CallStringReverse error = %v, want %v", err, tc.want)
1371 + }
1372 +
1373 + srv.wait(t)
1374 + })
1375 + }
1376 +}
1377 +
1378 +func TestWinCallSnapshotRejectsMalformedPayload(t *testing.T) {
1379 + svc := uniqueWinService("go_win_bad_snapshot_payload")
1380 + srv := startRawWinSessionServer(t, svc, testWinServerConfig(),
1381 + func(session *windows.Session, hdr protocol.Header, payload []byte) error {
1382 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodCgroupsSnapshot {
1383 + return fmt.Errorf("unexpected request header: %+v", hdr)
1384 + }
1385 +
1386 + respHdr := protocol.Header{
1387 + Kind: protocol.KindResponse,
1388 + Code: protocol.MethodCgroupsSnapshot,
1389 + ItemCount: 1,
1390 + MessageID: hdr.MessageID,
1391 + TransportStatus: protocol.StatusOK,
1392 + }
1393 + return session.Send(&respHdr, []byte{1, 2, 3, 4})
1394 + })
1395 +
1396 + client := NewSnapshotClient(winTestRunDir, svc, testWinClientConfig())
1397 + defer client.Close()
1398 +
1399 + client.Refresh()
1400 + if !client.Ready() {
1401 + t.Fatal("client not ready")
1402 + }
1403 +
1404 + _, err := client.CallSnapshot()
1405 + if !errors.Is(err, protocol.ErrTruncated) {
1406 + t.Fatalf("CallSnapshot error = %v, want %v", err, protocol.ErrTruncated)
1407 + }
1408 +
1409 + srv.wait(t)
1410 +}
1411 +
1412 +func TestWinCallStringReverseRejectsMalformedPayload(t *testing.T) {
1413 + svc := uniqueWinService("go_win_bad_str_payload")
1414 + srv := startRawWinSessionServer(t, svc, testWinServerConfig(),
1415 + func(session *windows.Session, hdr protocol.Header, payload []byte) error {
1416 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodStringReverse {
1417 + return fmt.Errorf("unexpected request header: %+v", hdr)
1418 + }
1419 +
1420 + respHdr := protocol.Header{
1421 + Kind: protocol.KindResponse,
1422 + Code: protocol.MethodStringReverse,
1423 + ItemCount: 1,
1424 + MessageID: hdr.MessageID,
1425 + TransportStatus: protocol.StatusOK,
1426 + }
1427 +
1428 + // Valid offsets/lengths, but missing the required NUL terminator.
1429 + respPayload := make([]byte, protocol.StringReverseHdrSize+3)
1430 + binary.NativeEndian.PutUint32(respPayload[0:4], uint32(protocol.StringReverseHdrSize))
1431 + binary.NativeEndian.PutUint32(respPayload[4:8], 2)
1432 + copy(respPayload[8:], []byte{'o', 'k', 'x'})
1433 + return session.Send(&respHdr, respPayload)
1434 + })
1435 +
1436 + client := NewStringReverseClient(winTestRunDir, svc, testWinClientConfig())
1437 + defer client.Close()
1438 +
1439 + client.Refresh()
1440 + if !client.Ready() {
1441 + t.Fatal("client not ready")
1442 + }
1443 +
1444 + _, err := client.CallStringReverse("hello")
1445 + if !errors.Is(err, protocol.ErrMissingNul) {
1446 + t.Fatalf("CallStringReverse error = %v, want %v", err, protocol.ErrMissingNul)
1447 + }
1448 +
1449 + srv.wait(t)
1450 +}
1451 +
1452 +func TestWinCallIncrementBatchRejectsMalformedResponseEnvelope(t *testing.T) {
1453 + cases := []struct {
1454 + name string
1455 + want error
1456 + mutate func(*protocol.Header)
1457 + }{
1458 + {
1459 + name: "bad kind",
1460 + want: protocol.ErrBadKind,
1461 + mutate: func(h *protocol.Header) {
1462 + h.Kind = protocol.KindRequest
1463 + },
1464 + },
1465 + {
1466 + name: "bad code",
1467 + want: protocol.ErrBadLayout,
1468 + mutate: func(h *protocol.Header) {
1469 + h.Code = protocol.MethodStringReverse
1470 + },
1471 + },
1472 + {
1473 + name: "bad status",
1474 + want: protocol.ErrBadLayout,
1475 + mutate: func(h *protocol.Header) {
1476 + h.TransportStatus = protocol.StatusInternalError
1477 + },
1478 + },
1479 + {
1480 + name: "bad message_id",
1481 + want: protocol.ErrTruncated,
1482 + mutate: func(h *protocol.Header) {
1483 + h.MessageID++
1484 + },
1485 + },
1486 + {
1487 + name: "missing batch flag",
1488 + want: protocol.ErrBadItemCount,
1489 + mutate: func(h *protocol.Header) {
1490 + h.Flags = 0
1491 + },
1492 + },
1493 + {
1494 + name: "wrong item count",
1495 + want: protocol.ErrBadItemCount,
1496 + mutate: func(h *protocol.Header) {
1497 + h.Flags = protocol.FlagBatch
1498 + h.ItemCount = 1
1499 + },
1500 + },
1501 + }
1502 +
1503 + for _, tc := range cases {
1504 + t.Run(tc.name, func(t *testing.T) {
1505 + cfg := testWinServerConfig()
1506 + cfg.MaxRequestBatchItems = 16
1507 + cfg.MaxResponseBatchItems = 16
1508 +
1509 + svc := uniqueWinService("go_win_bad_batch_resp")
1510 + srv := startRawWinSessionServer(t, svc, cfg,
1511 + func(session *windows.Session, hdr protocol.Header, payload []byte) error {
1512 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodIncrement {
1513 + return fmt.Errorf("unexpected request header: %+v", hdr)
1514 + }
1515 +
1516 + respHdr := protocol.Header{
1517 + Kind: protocol.KindResponse,
1518 + Code: protocol.MethodIncrement,
1519 + Flags: protocol.FlagBatch,
1520 + ItemCount: hdr.ItemCount,
1521 + MessageID: hdr.MessageID,
1522 + TransportStatus: protocol.StatusOK,
1523 + }
1524 + tc.mutate(&respHdr)
1525 +
1526 + var itemA [protocol.IncrementPayloadSize]byte
1527 + var itemB [protocol.IncrementPayloadSize]byte
1528 + if protocol.IncrementEncode(2, itemA[:]) == 0 || protocol.IncrementEncode(3, itemB[:]) == 0 {
1529 + return fmt.Errorf("IncrementEncode failed")
1530 + }
1531 +
1532 + respBuf := make([]byte, 64)
1533 + bb := protocol.NewBatchBuilder(respBuf, 2)
1534 + if err := bb.Add(itemA[:]); err != nil {
1535 + return err
1536 + }
1537 + if err := bb.Add(itemB[:]); err != nil {
1538 + return err
1539 + }
1540 + n, _ := bb.Finish()
1541 + return session.Send(&respHdr, respBuf[:n])
1542 + })
1543 +
1544 + ccfg := testWinClientConfig()
1545 + ccfg.MaxRequestBatchItems = 16
1546 + ccfg.MaxResponseBatchItems = 16
1547 +
1548 + client := NewIncrementClient(winTestRunDir, svc, ccfg)
1549 + defer client.Close()
1550 +
1551 + client.Refresh()
1552 + if !client.Ready() {
1553 + t.Fatal("client not ready")
1554 + }
1555 +
1556 + _, err := client.CallIncrementBatch([]uint64{1, 2})
1557 + if !errors.Is(err, tc.want) {
1558 + t.Fatalf("CallIncrementBatch error = %v, want %v", err, tc.want)
1559 + }
1560 +
1561 + srv.wait(t)
1562 + })
1563 + }
1564 +}
1565 +
1566 +func TestWinCallIncrementBatchRejectsMalformedPayload(t *testing.T) {
1567 + t.Run("bad batch directory", func(t *testing.T) {
1568 + cfg := testWinServerConfig()
1569 + cfg.MaxRequestBatchItems = 16
1570 + cfg.MaxResponseBatchItems = 16
1571 +
1572 + svc := uniqueWinService("go_win_bad_batch_dir")
1573 + srv := startRawWinSessionServerN(t, svc, cfg, 2,
1574 + func(session *windows.Session, hdr protocol.Header, payload []byte) error {
1575 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodIncrement {
1576 + return fmt.Errorf("unexpected request header: %+v", hdr)
1577 + }
1578 +
1579 + respPayload := make([]byte, 24)
1580 + binary.NativeEndian.PutUint32(respPayload[0:4], 1)
1581 + binary.NativeEndian.PutUint32(respPayload[4:8], 4)
1582 + binary.NativeEndian.PutUint32(respPayload[8:12], 0)
1583 + binary.NativeEndian.PutUint32(respPayload[12:16], 4)
1584 + copy(respPayload[16:], []byte("payload!!"))
1585 +
1586 + respHdr := protocol.Header{
1587 + Kind: protocol.KindResponse,
1588 + Code: protocol.MethodIncrement,
1589 + Flags: protocol.FlagBatch,
1590 + ItemCount: 2,
1591 + MessageID: hdr.MessageID,
1592 + TransportStatus: protocol.StatusOK,
1593 + }
1594 + return session.Send(&respHdr, respPayload)
1595 + })
1596 +
1597 + ccfg := testWinClientConfig()
1598 + ccfg.MaxRequestBatchItems = 16
1599 + ccfg.MaxResponseBatchItems = 16
1600 +
1601 + client := NewIncrementClient(winTestRunDir, svc, ccfg)
1602 + defer client.Close()
1603 +
1604 + client.Refresh()
1605 + if !client.Ready() {
1606 + t.Fatal("client not ready")
1607 + }
1608 +
1609 + _, err := client.CallIncrementBatch([]uint64{1, 2})
1610 + if !errors.Is(err, protocol.ErrTruncated) {
1611 + t.Fatalf("CallIncrementBatch error = %v, want %v", err, protocol.ErrTruncated)
1612 + }
1613 +
1614 + srv.wait(t)
1615 + })
1616 +
1617 + t.Run("truncated batch item", func(t *testing.T) {
1618 + cfg := testWinServerConfig()
1619 + cfg.MaxRequestBatchItems = 16
1620 + cfg.MaxResponseBatchItems = 16
1621 +
1622 + svc := uniqueWinService("go_win_bad_batch_payload")
1623 + srv := startRawWinSessionServerN(t, svc, cfg, 2,
1624 + func(session *windows.Session, hdr protocol.Header, payload []byte) error {
1625 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodIncrement {
1626 + return fmt.Errorf("unexpected request header: %+v", hdr)
1627 + }
1628 +
1629 + respBuf := make([]byte, 64)
1630 + bb := protocol.NewBatchBuilder(respBuf, 1)
1631 + if err := bb.Add([]byte{1, 2, 3, 4}); err != nil {
1632 + return err
1633 + }
1634 + n, _ := bb.Finish()
1635 +
1636 + respHdr := protocol.Header{
1637 + Kind: protocol.KindResponse,
1638 + Code: protocol.MethodIncrement,
1639 + Flags: protocol.FlagBatch,
1640 + ItemCount: 1,
1641 + MessageID: hdr.MessageID,
1642 + TransportStatus: protocol.StatusOK,
1643 + }
1644 + return session.Send(&respHdr, respBuf[:n])
1645 + })
1646 +
1647 + ccfg := testWinClientConfig()
1648 + ccfg.MaxRequestBatchItems = 16
1649 + ccfg.MaxResponseBatchItems = 16
1650 +
1651 + client := NewIncrementClient(winTestRunDir, svc, ccfg)
1652 + defer client.Close()
1653 +
1654 + client.Refresh()
1655 + if !client.Ready() {
1656 + t.Fatal("client not ready")
1657 + }
1658 +
1659 + _, err := client.CallIncrementBatch([]uint64{1})
1660 + if !errors.Is(err, protocol.ErrTruncated) {
1661 + t.Fatalf("CallIncrementBatch error = %v, want %v", err, protocol.ErrTruncated)
1662 + }
1663 +
1664 + srv.wait(t)
1665 + })
1666 +
1667 + t.Run("missing batch body", func(t *testing.T) {
1668 + cfg := testWinServerConfig()
1669 + cfg.MaxRequestBatchItems = 16
1670 + cfg.MaxResponseBatchItems = 16
1671 +
1672 + svc := uniqueWinService("go_win_bad_batch_body")
1673 + srv := startRawWinSessionServerN(t, svc, cfg, 2,
1674 + func(session *windows.Session, hdr protocol.Header, payload []byte) error {
1675 + if hdr.Kind != protocol.KindRequest || hdr.Code != protocol.MethodIncrement {
1676 + return fmt.Errorf("unexpected request header: %+v", hdr)
1677 + }
1678 +
1679 + respHdr := protocol.Header{
1680 + Kind: protocol.KindResponse,
1681 + Code: protocol.MethodIncrement,
1682 + Flags: protocol.FlagBatch,
1683 + ItemCount: hdr.ItemCount,
1684 + MessageID: hdr.MessageID,
1685 + TransportStatus: protocol.StatusOK,
1686 + }
1687 + return session.Send(&respHdr, nil)
1688 + })
1689 +
1690 + ccfg := testWinClientConfig()
1691 + ccfg.MaxRequestBatchItems = 16
1692 + ccfg.MaxResponseBatchItems = 16
1693 +
1694 + client := NewIncrementClient(winTestRunDir, svc, ccfg)
1695 + defer client.Close()
1696 +
1697 + client.Refresh()
1698 + if !client.Ready() {
1699 + t.Fatal("client not ready")
1700 + }
1701 +
1702 + _, err := client.CallIncrementBatch([]uint64{1, 2})
1703 + if !errors.Is(err, protocol.ErrTruncated) {
1704 + t.Fatalf("CallIncrementBatch error = %v, want %v", err, protocol.ErrTruncated)
1705 + }
1706 +
1707 + srv.wait(t)
1708 + })
1709 +}
src/go/pkg/netipc/service/raw/ping_pong_test.go new
+219
@@ -0,0 +1,219 @@
1 +//go:build unix
2 +
3 +package raw
4 +
5 +import (
6 + "testing"
7 + "time"
8 +
9 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
10 +)
11 +
12 +func pingPongIncrementDispatch() DispatchHandler {
13 + return IncrementDispatch(func(v uint64) (uint64, bool) {
14 + return v + 1, true
15 + })
16 +}
17 +
18 +func pingPongStringReverseDispatch() DispatchHandler {
19 + return StringReverseDispatch(func(s string) (string, bool) {
20 + return reverseString(s), true
21 + })
22 +}
23 +
24 +// reverseString reverses a string byte-by-byte.
25 +func reverseString(s string) string {
26 + b := []byte(s)
27 + for i, j := 0, len(b)-1; i < j; i, j = i+1, j-1 {
28 + b[i], b[j] = b[j], b[i]
29 + }
30 + return string(b)
31 +}
32 +
33 +func TestIncrementPingPong(t *testing.T) {
34 + svc := "go_pp_incr"
35 + ensureRunDir()
36 + cleanupAll(svc)
37 +
38 + ts := startTestServerUnixWithConfig(
39 + svc,
40 + testServerConfig(),
41 + protocol.MethodIncrement,
42 + pingPongIncrementDispatch(),
43 + )
44 + defer ts.stop()
45 +
46 + client := NewIncrementClient(testRunDir, svc, testClientConfig())
47 + client.Refresh()
48 + if !client.Ready() {
49 + t.Fatal("client not ready")
50 + }
51 +
52 + // 10 rounds: send 0 -> get 1 -> send 1 -> get 2 -> ... -> value == 10
53 + var val uint64
54 + responsesReceived := 0
55 + for i := 0; i < 10; i++ {
56 + got, err := client.CallIncrement(val)
57 + if err != nil {
58 + t.Fatalf("round %d: CallIncrement(%d) failed: %v", i, val, err)
59 + }
60 + responsesReceived++
61 + expected := val + 1
62 + if got != expected {
63 + t.Fatalf("round %d: expected %d, got %d", i, expected, got)
64 + }
65 + val = got
66 + }
67 +
68 + if responsesReceived != 10 {
69 + t.Fatalf("expected 10 responses received, got %d", responsesReceived)
70 + }
71 + if val != 10 {
72 + t.Fatalf("expected final value 10, got %d", val)
73 + }
74 +
75 + status := client.Status()
76 + if status.CallCount != 10 {
77 + t.Fatalf("expected call_count=10, got %d", status.CallCount)
78 + }
79 + if status.ErrorCount != 0 {
80 + t.Fatalf("expected error_count=0, got %d", status.ErrorCount)
81 + }
82 +
83 + client.Close()
84 + cleanupAll(svc)
85 +}
86 +
87 +func TestStringReversePingPong(t *testing.T) {
88 + svc := "go_pp_strrev"
89 + ensureRunDir()
90 + cleanupAll(svc)
91 +
92 + ts := startTestServerUnixWithConfig(
93 + svc,
94 + testServerConfig(),
95 + protocol.MethodStringReverse,
96 + pingPongStringReverseDispatch(),
97 + )
98 + defer ts.stop()
99 +
100 + client := NewStringReverseClient(testRunDir, svc, testClientConfig())
101 + client.Refresh()
102 + if !client.Ready() {
103 + t.Fatal("client not ready")
104 + }
105 +
106 + original := "abcdefghijklmnopqrstuvwxyz"
107 +
108 + // 6 rounds: feed each response back as the next request
109 + responsesReceived := 0
110 + current := original
111 + for i := 0; i < 6; i++ {
112 + view, err := client.CallStringReverse(current)
113 + if err != nil {
114 + t.Fatalf("round %d: CallStringReverse(%q) failed: %v", i+1, current, err)
115 + }
116 + responsesReceived++
117 +
118 + // verify response is the character-by-character reverse of the sent string
119 + expectedReversed := reverseString(current)
120 + if view.Str != expectedReversed {
121 + t.Fatalf("round %d: sent %q, expected reverse %q, got %q", i+1, current, expectedReversed, view.Str)
122 + }
123 +
124 + current = view.Str
125 + }
126 +
127 + if responsesReceived != 6 {
128 + t.Fatalf("expected 6 responses received, got %d", responsesReceived)
129 + }
130 +
131 + // even number of reversals = identity
132 + if current != original {
133 + t.Fatalf("after 6 reversals expected original %q, got %q", original, current)
134 + }
135 +
136 + status := client.Status()
137 + if status.CallCount != 6 {
138 + t.Fatalf("expected call_count=6, got %d", status.CallCount)
139 + }
140 + if status.ErrorCount != 0 {
141 + t.Fatalf("expected error_count=0, got %d", status.ErrorCount)
142 + }
143 +
144 + client.Close()
145 + cleanupAll(svc)
146 +}
147 +
148 +func TestIncrementBatch(t *testing.T) {
149 + svc := "go_pp_batch"
150 + ensureRunDir()
151 + cleanupAll(svc)
152 +
153 + // Server config with batch support
154 + sCfg := testServerConfig()
155 + sCfg.MaxRequestBatchItems = 16
156 + sCfg.MaxResponseBatchItems = 16
157 + sCfg.MaxRequestPayloadBytes = 65536
158 +
159 + s := NewServer(
160 + testRunDir,
161 + svc,
162 + sCfg,
163 + protocol.MethodIncrement,
164 + pingPongIncrementDispatch(),
165 + )
166 + doneCh := make(chan struct{})
167 + go func() {
168 + defer close(doneCh)
169 + s.Run()
170 + }()
171 + defer func() {
172 + s.Stop()
173 + <-doneCh
174 + }()
175 +
176 + // Wait for server
177 + time.Sleep(100 * time.Millisecond)
178 +
179 + // Client config with batch support
180 + cCfg := testClientConfig()
181 + cCfg.MaxRequestBatchItems = 16
182 + cCfg.MaxResponseBatchItems = 16
183 + cCfg.MaxRequestPayloadBytes = 65536
184 +
185 + client := NewIncrementClient(testRunDir, svc, cCfg)
186 + client.Refresh()
187 + if !client.Ready() {
188 + t.Fatal("client not ready")
189 + }
190 +
191 + input := []uint64{10, 20, 30, 40, 50}
192 +
193 + results, err := client.CallIncrementBatch(input)
194 + if err != nil {
195 + t.Fatalf("CallIncrementBatch failed: %v", err)
196 + }
197 +
198 + if len(results) != len(input) {
199 + t.Fatalf("expected %d results, got %d", len(input), len(results))
200 + }
201 +
202 + for i, v := range input {
203 + expected := v + 1
204 + if results[i] != expected {
205 + t.Fatalf("item %d: expected %d, got %d", i, expected, results[i])
206 + }
207 + }
208 +
209 + status := client.Status()
210 + if status.CallCount != 1 {
211 + t.Fatalf("expected call_count=1, got %d", status.CallCount)
212 + }
213 + if status.ErrorCount != 0 {
214 + t.Fatalf("expected error_count=0, got %d", status.ErrorCount)
215 + }
216 +
217 + client.Close()
218 + cleanupAll(svc)
219 +}
src/go/pkg/netipc/service/raw/ping_pong_windows_test.go new
+182
@@ -0,0 +1,182 @@
1 +//go:build windows
2 +
3 +package raw
4 +
5 +import (
6 + "testing"
7 + "time"
8 +
9 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
10 + windows "github.com/netdata/netdata/go/plugins/pkg/netipc/transport/windows"
11 +)
12 +
13 +const (
14 + winTestRunDir = `C:\ProgramData\netipc_test`
15 + winAuthToken = uint64(0xDEADBEEFCAFEBABE)
16 + winResponseBufSize = 65536
17 +)
18 +
19 +// testWinServerConfig returns a baseline Windows server config for tests.
20 +func testWinServerConfig() windows.ServerConfig {
21 + return windows.ServerConfig{
22 + SupportedProfiles: protocol.ProfileBaseline,
23 + MaxRequestPayloadBytes: 4096,
24 + MaxRequestBatchItems: 1,
25 + MaxResponsePayloadBytes: winResponseBufSize,
26 + MaxResponseBatchItems: 1,
27 + AuthToken: winAuthToken,
28 + }
29 +}
30 +
31 +// testWinClientConfig returns a baseline Windows client config for tests.
32 +func testWinClientConfig() windows.ClientConfig {
33 + return windows.ClientConfig{
34 + SupportedProfiles: protocol.ProfileBaseline,
35 + MaxRequestPayloadBytes: 4096,
36 + MaxRequestBatchItems: 1,
37 + MaxResponsePayloadBytes: winResponseBufSize,
38 + MaxResponseBatchItems: 1,
39 + AuthToken: winAuthToken,
40 + }
41 +}
42 +
43 +// winTestServer wraps a Server with a done channel for clean shutdown.
44 +type winTestServer struct {
45 + server *Server
46 + doneCh chan struct{}
47 +}
48 +
49 +// startTestServerWin creates a Windows Server (Named Pipe transport)
50 +// and starts it in a background goroutine. Returns a handle for cleanup.
51 +func startTestServerWin(service string, expectedMethodCode uint16, handler DispatchHandler) *winTestServer {
52 + s := NewServer(winTestRunDir, service, testWinServerConfig(), expectedMethodCode, handler)
53 + doneCh := make(chan struct{})
54 +
55 + go func() {
56 + defer close(doneCh)
57 + s.Run()
58 + }()
59 +
60 + // Allow the Named Pipe listener to start
61 + time.Sleep(200 * time.Millisecond)
62 +
63 + return &winTestServer{server: s, doneCh: doneCh}
64 +}
65 +
66 +func (ts *winTestServer) stop() {
67 + ts.server.Stop()
68 + <-ts.doneCh
69 +}
70 +
71 +func winPingPongIncrementDispatch() DispatchHandler {
72 + return IncrementDispatch(func(v uint64) (uint64, bool) {
73 + return v + 1, true
74 + })
75 +}
76 +
77 +func winPingPongStringReverseDispatch() DispatchHandler {
78 + return StringReverseDispatch(func(s string) (string, bool) {
79 + return winReverseString(s), true
80 + })
81 +}
82 +
83 +// winReverseString reverses a string byte-by-byte.
84 +func winReverseString(s string) string {
85 + b := []byte(s)
86 + for i, j := 0, len(b)-1; i < j; i, j = i+1, j-1 {
87 + b[i], b[j] = b[j], b[i]
88 + }
89 + return string(b)
90 +}
91 +
92 +func TestWinIncrementPingPong(t *testing.T) {
93 + svc := "go_win_pp_incr"
94 +
95 + ts := startTestServerWin(svc, protocol.MethodIncrement, winPingPongIncrementDispatch())
96 + defer ts.stop()
97 +
98 + client := NewIncrementClient(winTestRunDir, svc, testWinClientConfig())
99 + waitWinClientReady(t, client)
100 +
101 + // 10 rounds: send 0 -> get 1 -> send 1 -> get 2 -> ... -> value == 10
102 + var val uint64
103 + responsesReceived := 0
104 + for i := 0; i < 10; i++ {
105 + got, err := client.CallIncrement(val)
106 + if err != nil {
107 + t.Fatalf("round %d: CallIncrement(%d) failed: %v", i, val, err)
108 + }
109 + responsesReceived++
110 + expected := val + 1
111 + if got != expected {
112 + t.Fatalf("round %d: expected %d, got %d", i, expected, got)
113 + }
114 + val = got
115 + }
116 +
117 + if responsesReceived != 10 {
118 + t.Fatalf("expected 10 responses received, got %d", responsesReceived)
119 + }
120 + if val != 10 {
121 + t.Fatalf("expected final value 10, got %d", val)
122 + }
123 +
124 + status := client.Status()
125 + if status.CallCount != 10 {
126 + t.Fatalf("expected call_count=10, got %d", status.CallCount)
127 + }
128 + if status.ErrorCount != 0 {
129 + t.Fatalf("expected error_count=0, got %d", status.ErrorCount)
130 + }
131 +
132 + client.Close()
133 +}
134 +
135 +func TestWinStringReversePingPong(t *testing.T) {
136 + svc := "go_win_pp_strrev"
137 +
138 + ts := startTestServerWin(svc, protocol.MethodStringReverse, winPingPongStringReverseDispatch())
139 + defer ts.stop()
140 +
141 + client := NewStringReverseClient(winTestRunDir, svc, testWinClientConfig())
142 + waitWinClientReady(t, client)
143 +
144 + original := "abcdefghijklmnopqrstuvwxyz"
145 +
146 + // 6 rounds: feed each response back as the next request
147 + responsesReceived := 0
148 + current := original
149 + for i := 0; i < 6; i++ {
150 + view, err := client.CallStringReverse(current)
151 + if err != nil {
152 + t.Fatalf("round %d: CallStringReverse(%q) failed: %v", i+1, current, err)
153 + }
154 + responsesReceived++
155 +
156 + expectedReversed := winReverseString(current)
157 + if view.Str != expectedReversed {
158 + t.Fatalf("round %d: sent %q, expected reverse %q, got %q", i+1, current, expectedReversed, view.Str)
159 + }
160 +
161 + current = view.Str
162 + }
163 +
164 + if responsesReceived != 6 {
165 + t.Fatalf("expected 6 responses received, got %d", responsesReceived)
166 + }
167 +
168 + // even number of reversals = identity
169 + if current != original {
170 + t.Fatalf("after 6 reversals expected original %q, got %q", original, current)
171 + }
172 +
173 + status := client.Status()
174 + if status.CallCount != 6 {
175 + t.Fatalf("expected call_count=6, got %d", status.CallCount)
176 + }
177 + if status.ErrorCount != 0 {
178 + t.Fatalf("expected error_count=0, got %d", status.ErrorCount)
179 + }
180 +
181 + client.Close()
182 +}
src/go/pkg/netipc/service/raw/scratch.go new
+8
@@ -0,0 +1,8 @@
1 +package raw
2 +
3 +func ensureClientScratch(buf *[]byte, needed int) []byte {
4 + if len(*buf) < needed {
5 + *buf = make([]byte, needed)
6 + }
7 + return (*buf)[:needed]
8 +}
src/go/pkg/netipc/service/raw/shm_unix_test.go new
+577
@@ -0,0 +1,577 @@
1 +//go:build unix
2 +
3 +package raw
4 +
5 +import (
6 + "errors"
7 + "testing"
8 + "time"
9 +
10 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
11 + "github.com/netdata/netdata/go/plugins/pkg/netipc/transport/posix"
12 +)
13 +
14 +func testUnixShmServerConfig() posix.ServerConfig {
15 + cfg := testServerConfig()
16 + cfg.SupportedProfiles = protocol.ProfileBaseline | protocol.ProfileSHMHybrid
17 + cfg.PreferredProfiles = protocol.ProfileSHMHybrid
18 + return cfg
19 +}
20 +
21 +func testUnixShmClientConfig() posix.ClientConfig {
22 + cfg := testClientConfig()
23 + cfg.SupportedProfiles = protocol.ProfileBaseline | protocol.ProfileSHMHybrid
24 + cfg.PreferredProfiles = protocol.ProfileSHMHybrid
25 + return cfg
26 +}
27 +
28 +func startTestServerUnixWithConfig(
29 + service string,
30 + cfg posix.ServerConfig,
31 + expectedMethodCode uint16,
32 + handler DispatchHandler,
33 +) *testServer {
34 + ensureRunDir()
35 + cleanupAll(service)
36 +
37 + s := NewServer(testRunDir, service, cfg, expectedMethodCode, handler)
38 + doneCh := make(chan struct{})
39 +
40 + go func() {
41 + defer close(doneCh)
42 + _ = s.Run()
43 + }()
44 +
45 + waitUnixServerReady(service)
46 + return &testServer{server: s, doneCh: doneCh}
47 +}
48 +
49 +func encodeRawUnixMessage(hdr protocol.Header, payload []byte) []byte {
50 + hdr.Magic = protocol.MagicMsg
51 + hdr.Version = protocol.Version
52 + hdr.HeaderLen = protocol.HeaderLen
53 + hdr.PayloadLen = uint32(len(payload))
54 +
55 + msg := make([]byte, protocol.HeaderSize+len(payload))
56 + hdr.Encode(msg[:protocol.HeaderSize])
57 + copy(msg[protocol.HeaderSize:], payload)
58 + return msg
59 +}
60 +
61 +func encodeUnixIncrementBatchPayload(t *testing.T, values ...uint64) []byte {
62 + t.Helper()
63 +
64 + itemCount := uint32(len(values))
65 + if itemCount == 0 {
66 + return nil
67 + }
68 +
69 + bufSize := protocol.Align8(int(itemCount)*8) +
70 + int(itemCount)*protocol.IncrementPayloadSize +
71 + int(itemCount)*protocol.Alignment
72 + buf := make([]byte, bufSize)
73 + entries := make([]protocol.BatchEntry, len(values))
74 +
75 + offset := protocol.Align8(int(itemCount) * 8)
76 + for i, v := range values {
77 + itemOff := offset
78 + itemLen := protocol.IncrementPayloadSize
79 + entries[i] = protocol.BatchEntry{Offset: uint32(itemOff), Length: uint32(itemLen)}
80 + if protocol.IncrementEncode(v, buf[itemOff:itemOff+itemLen]) == 0 {
81 + t.Fatalf("IncrementEncode(%d) failed", v)
82 + }
83 + offset += protocol.Align8(itemLen)
84 + }
85 +
86 + if n := protocol.BatchDirEncode(entries, buf[:int(itemCount)*8]); n != int(itemCount)*8 {
87 + t.Fatalf("BatchDirEncode returned %d, want %d", n, int(itemCount)*8)
88 + }
89 +
90 + return buf[:offset]
91 +}
92 +
93 +func TestUnixShmRoundTrip(t *testing.T) {
94 + svc := uniqueUnixService("go_unix_shm_roundtrip")
95 + ts := startTestServerUnixWithConfig(
96 + svc,
97 + testUnixShmServerConfig(),
98 + protocol.MethodIncrement,
99 + pingPongIncrementDispatch(),
100 + )
101 + defer ts.stop()
102 +
103 + client := NewIncrementClient(testRunDir, svc, testUnixShmClientConfig())
104 + defer client.Close()
105 +
106 + refreshUnixClientReady(t, client)
107 +
108 + if client.session == nil {
109 + t.Fatal("expected negotiated session")
110 + }
111 + if client.shm == nil {
112 + t.Fatal("expected SHM attachment")
113 + }
114 + if client.session.SelectedProfile != protocol.ProfileSHMHybrid {
115 + t.Fatalf("selected profile = %d, want %d", client.session.SelectedProfile, protocol.ProfileSHMHybrid)
116 + }
117 +
118 + got, err := client.CallIncrement(41)
119 + if err != nil {
120 + t.Fatalf("CallIncrement failed: %v", err)
121 + }
122 + if got != 42 {
123 + t.Fatalf("CallIncrement = %d, want 42", got)
124 + }
125 +}
126 +
127 +func TestUnixShmAttachFailureFallsBackToBaseline(t *testing.T) {
128 + svc := uniqueUnixService("go_unix_shm_attach_fail")
129 + cfg := testUnixShmServerConfig()
130 + ensureRunDir()
131 + cleanupAll(svc)
132 +
133 + listener, err := posix.Listen(testRunDir, svc, cfg)
134 + if err != nil {
135 + t.Fatalf("posix.Listen failed: %v", err)
136 + }
137 + defer listener.Close()
138 +
139 + type attachFailureResult struct {
140 + firstSelected uint32
141 + secondSelected uint32
142 + err error
143 + }
144 + doneCh := make(chan attachFailureResult, 1)
145 + go func() {
146 + first, err := listener.Accept()
147 + if err != nil {
148 + doneCh <- attachFailureResult{err: err}
149 + return
150 + }
151 + defer first.Close()
152 + firstSelected := first.SelectedProfile
153 +
154 + recvBuf := make([]byte, protocol.HeaderSize+4096)
155 + _, _, err = first.Receive(recvBuf)
156 + if err == nil {
157 + doneCh <- attachFailureResult{err: errors.New("expected first receive to fail after SHM attach fallback disconnect")}
158 + return
159 + }
160 +
161 + second, err := listener.Accept()
162 + if err != nil {
163 + doneCh <- attachFailureResult{err: err}
164 + return
165 + }
166 + defer second.Close()
167 + secondSelected := second.SelectedProfile
168 +
169 + _, _, err = second.Receive(recvBuf)
170 + if err == nil {
171 + doneCh <- attachFailureResult{err: errors.New("expected second receive to fail after client close")}
172 + return
173 + }
174 +
175 + doneCh <- attachFailureResult{
176 + firstSelected: firstSelected,
177 + secondSelected: secondSelected,
178 + }
179 + }()
180 +
181 + client := NewIncrementClient(testRunDir, svc, testUnixShmClientConfig())
182 +
183 + if changed := client.Refresh(); !changed {
184 + t.Fatal("refresh should transition to READY via baseline fallback after SHM attach failure")
185 + }
186 + if !client.Ready() {
187 + t.Fatal("client should be ready after baseline fallback")
188 + }
189 + if client.state != StateReady {
190 + t.Fatalf("client state = %d, want READY", client.state)
191 + }
192 + if client.shm != nil {
193 + t.Fatal("expected no SHM attachment after fallback")
194 + }
195 + if client.session == nil {
196 + t.Fatal("expected live baseline session after fallback")
197 + }
198 + if client.session.SelectedProfile != protocol.ProfileBaseline {
199 + t.Fatalf("selected profile after fallback = %#x, want baseline", client.session.SelectedProfile)
200 + }
201 + if client.config.SupportedProfiles&protocol.ProfileSHMHybrid != 0 {
202 + t.Fatalf("supported profiles should drop SHM after attach failure, got %#x", client.config.SupportedProfiles)
203 + }
204 + if client.config.PreferredProfiles&protocol.ProfileSHMHybrid != 0 {
205 + t.Fatalf("preferred profiles should drop SHM after attach failure, got %#x", client.config.PreferredProfiles)
206 + }
207 +
208 + client.Close()
209 +
210 + result := <-doneCh
211 + if result.err != nil {
212 + t.Fatalf("attach-failure fallback server failed: %v", result.err)
213 + }
214 + if result.firstSelected != protocol.ProfileSHMHybrid {
215 + t.Fatalf("first selected profile = %#x, want SHM hybrid", result.firstSelected)
216 + }
217 + if result.secondSelected != protocol.ProfileBaseline {
218 + t.Fatalf("second selected profile = %#x, want baseline", result.secondSelected)
219 + }
220 +}
221 +
222 +func TestUnixDoRawCallShmRejectsBadMessageID(t *testing.T) {
223 + client, serverShm := newRawPosixShmClient(t)
224 +
225 + serverDone := make(chan error, 1)
226 + go func() {
227 + reqBuf := make([]byte, protocol.HeaderSize+32)
228 + n, err := serverShm.ShmReceive(reqBuf, 1000)
229 + if err != nil {
230 + serverDone <- err
231 + return
232 + }
233 +
234 + hdr, err := protocol.DecodeHeader(reqBuf[:n])
235 + if err != nil {
236 + serverDone <- err
237 + return
238 + }
239 +
240 + var respPayload [protocol.IncrementPayloadSize]byte
241 + if protocol.IncrementEncode(42, respPayload[:]) == 0 {
242 + serverDone <- errors.New("IncrementEncode failed")
243 + return
244 + }
245 +
246 + respHdr := protocol.Header{
247 + Kind: protocol.KindResponse,
248 + Code: protocol.MethodIncrement,
249 + ItemCount: 1,
250 + MessageID: hdr.MessageID + 1,
251 + TransportStatus: protocol.StatusOK,
252 + }
253 + serverDone <- serverShm.ShmSend(encodeRawUnixMessage(respHdr, respPayload[:]))
254 + }()
255 +
256 + var reqPayload [protocol.IncrementPayloadSize]byte
257 + if protocol.IncrementEncode(41, reqPayload[:]) == 0 {
258 + t.Fatal("IncrementEncode failed")
259 + }
260 +
261 + _, _, err := client.doRawCall(protocol.MethodIncrement, reqPayload[:])
262 + if !errors.Is(err, protocol.ErrBadLayout) {
263 + t.Fatalf("doRawCall bad SHM message_id = %v, want %v", err, protocol.ErrBadLayout)
264 + }
265 +
266 + if err := <-serverDone; err != nil {
267 + t.Fatalf("raw POSIX SHM server failed: %v", err)
268 + }
269 +}
270 +
271 +func TestUnixCallIncrementBatchShmRejectsBadMessageID(t *testing.T) {
272 + client, serverShm := newRawPosixShmClient(t)
273 +
274 + serverDone := make(chan error, 1)
275 + go func() {
276 + reqBuf := make([]byte, protocol.HeaderSize+256)
277 + n, err := serverShm.ShmReceive(reqBuf, 1000)
278 + if err != nil {
279 + serverDone <- err
280 + return
281 + }
282 +
283 + hdr, err := protocol.DecodeHeader(reqBuf[:n])
284 + if err != nil {
285 + serverDone <- err
286 + return
287 + }
288 +
289 + var itemA [protocol.IncrementPayloadSize]byte
290 + var itemB [protocol.IncrementPayloadSize]byte
291 + if protocol.IncrementEncode(2, itemA[:]) == 0 || protocol.IncrementEncode(3, itemB[:]) == 0 {
292 + serverDone <- errors.New("IncrementEncode failed")
293 + return
294 + }
295 +
296 + respBuf := make([]byte, 64)
297 + bb := protocol.NewBatchBuilder(respBuf, 2)
298 + if err := bb.Add(itemA[:]); err != nil {
299 + serverDone <- err
300 + return
301 + }
302 + if err := bb.Add(itemB[:]); err != nil {
303 + serverDone <- err
304 + return
305 + }
306 + respLen, _ := bb.Finish()
307 +
308 + respHdr := protocol.Header{
309 + Kind: protocol.KindResponse,
310 + Code: protocol.MethodIncrement,
311 + Flags: protocol.FlagBatch,
312 + ItemCount: 2,
313 + MessageID: hdr.MessageID + 1,
314 + TransportStatus: protocol.StatusOK,
315 + }
316 + serverDone <- serverShm.ShmSend(encodeRawUnixMessage(respHdr, respBuf[:respLen]))
317 + }()
318 +
319 + _, err := client.CallIncrementBatch([]uint64{1, 2})
320 + if !errors.Is(err, protocol.ErrBadLayout) {
321 + t.Fatalf("CallIncrementBatch bad SHM message_id = %v, want %v", err, protocol.ErrBadLayout)
322 + }
323 +
324 + if err := <-serverDone; err != nil {
325 + t.Fatalf("raw POSIX SHM server failed: %v", err)
326 + }
327 +}
328 +
329 +func TestUnixCallIncrementBatchShmRejectsMalformedPayload(t *testing.T) {
330 + client, serverShm := newRawPosixShmClient(t)
331 +
332 + serverDone := make(chan error, 1)
333 + go func() {
334 + reqBuf := make([]byte, protocol.HeaderSize+256)
335 + n, err := serverShm.ShmReceive(reqBuf, 1000)
336 + if err != nil {
337 + serverDone <- err
338 + return
339 + }
340 +
341 + hdr, err := protocol.DecodeHeader(reqBuf[:n])
342 + if err != nil {
343 + serverDone <- err
344 + return
345 + }
346 +
347 + respHdr := protocol.Header{
348 + Kind: protocol.KindResponse,
349 + Code: protocol.MethodIncrement,
350 + Flags: protocol.FlagBatch,
351 + ItemCount: 2,
352 + MessageID: hdr.MessageID,
353 + TransportStatus: protocol.StatusOK,
354 + }
355 + serverDone <- serverShm.ShmSend(encodeRawUnixMessage(respHdr, make([]byte, 8)))
356 + }()
357 +
358 + _, err := client.CallIncrementBatch([]uint64{1, 2})
359 + if !errors.Is(err, protocol.ErrTruncated) {
360 + t.Fatalf("CallIncrementBatch malformed SHM payload = %v, want %v", err, protocol.ErrTruncated)
361 + }
362 +
363 + if err := <-serverDone; err != nil {
364 + t.Fatalf("raw POSIX SHM server failed: %v", err)
365 + }
366 +}
367 +
368 +func TestUnixShmMalformedBatchRequestRecovers(t *testing.T) {
369 + svc := uniqueUnixService("go_unix_shm_bad_batch")
370 + ts := startTestServerUnixWithConfig(
371 + svc,
372 + testUnixShmServerConfig(),
373 + protocol.MethodIncrement,
374 + pingPongIncrementDispatch(),
375 + )
376 + defer ts.stop()
377 +
378 + client := NewIncrementClient(testRunDir, svc, testUnixShmClientConfig())
379 + defer client.Close()
380 +
381 + refreshUnixClientReady(t, client)
382 +
383 + if client.shm == nil {
384 + t.Fatal("expected SHM attachment")
385 + }
386 +
387 + reqHdr := protocol.Header{
388 + Kind: protocol.KindRequest,
389 + Code: protocol.MethodIncrement,
390 + Flags: protocol.FlagBatch,
391 + ItemCount: 2,
392 + MessageID: 1,
393 + TransportStatus: protocol.StatusOK,
394 + }
395 + badPayload := encodeUnixIncrementBatchPayload(t, 1, 2)[:8]
396 + if err := client.shm.ShmSend(encodeRawUnixMessage(reqHdr, badPayload)); err != nil {
397 + t.Fatalf("ShmSend malformed batch request failed: %v", err)
398 + }
399 +
400 + time.Sleep(50 * time.Millisecond)
401 +
402 + got, err := client.CallIncrement(21)
403 + if err != nil {
404 + t.Fatalf("CallIncrement after malformed batch SHM request failed: %v", err)
405 + }
406 + if got != 22 {
407 + t.Fatalf("CallIncrement after malformed batch SHM request = %d, want 22", got)
408 + }
409 + if !client.Ready() {
410 + t.Fatalf("client should stay usable after malformed batch SHM request, got status %+v", client.Status())
411 + }
412 +}
413 +
414 +func TestUnixShmShortRequestKeepsServerHealthy(t *testing.T) {
415 + svc := uniqueUnixService("go_unix_shm_short_req")
416 + ts := startTestServerUnixWithConfig(
417 + svc,
418 + testUnixShmServerConfig(),
419 + protocol.MethodIncrement,
420 + pingPongIncrementDispatch(),
421 + )
422 + defer ts.stop()
423 +
424 + client := NewIncrementClient(testRunDir, svc, testUnixShmClientConfig())
425 + defer client.Close()
426 + refreshUnixClientReady(t, client)
427 +
428 + if client.shm == nil {
429 + t.Fatal("expected SHM attachment")
430 + }
431 + if err := client.shm.ShmSend([]byte{1, 2, 3, 4}); err != nil {
432 + t.Fatalf("ShmSend short request failed: %v", err)
433 + }
434 +
435 + time.Sleep(50 * time.Millisecond)
436 +
437 + verify := NewIncrementClient(testRunDir, svc, testUnixShmClientConfig())
438 + defer verify.Close()
439 + refreshUnixClientReady(t, verify)
440 +
441 + got, err := verify.CallIncrement(21)
442 + if err != nil {
443 + t.Fatalf("CallIncrement after short SHM request failed: %v", err)
444 + }
445 + if got != 22 {
446 + t.Fatalf("CallIncrement after short SHM request = %d, want 22", got)
447 + }
448 +}
449 +
450 +func TestUnixShmBadHeaderKeepsServerHealthy(t *testing.T) {
451 + svc := uniqueUnixService("go_unix_shm_bad_header")
452 + ts := startTestServerUnixWithConfig(
453 + svc,
454 + testUnixShmServerConfig(),
455 + protocol.MethodIncrement,
456 + pingPongIncrementDispatch(),
457 + )
458 + defer ts.stop()
459 +
460 + client := NewIncrementClient(testRunDir, svc, testUnixShmClientConfig())
461 + defer client.Close()
462 + refreshUnixClientReady(t, client)
463 +
464 + if client.shm == nil {
465 + t.Fatal("expected SHM attachment")
466 + }
467 + badMsg := make([]byte, protocol.HeaderSize)
468 + if err := client.shm.ShmSend(badMsg); err != nil {
469 + t.Fatalf("ShmSend bad header request failed: %v", err)
470 + }
471 +
472 + time.Sleep(50 * time.Millisecond)
473 +
474 + verify := NewIncrementClient(testRunDir, svc, testUnixShmClientConfig())
475 + defer verify.Close()
476 + refreshUnixClientReady(t, verify)
477 +
478 + got, err := verify.CallIncrement(9)
479 + if err != nil {
480 + t.Fatalf("CallIncrement after bad SHM header failed: %v", err)
481 + }
482 + if got != 10 {
483 + t.Fatalf("CallIncrement after bad SHM header = %d, want 10", got)
484 + }
485 +}
486 +
487 +func TestUnixShmBatchHandlerFailureNeedsRefresh(t *testing.T) {
488 + svc := uniqueUnixService("go_unix_shm_batch_fail")
489 + handler := IncrementDispatch(func(v uint64) (uint64, bool) {
490 + if v == 99 {
491 + return 0, false
492 + }
493 + return v + 1, true
494 + })
495 + ts := startTestServerUnixWithConfig(
496 + svc,
497 + testUnixShmServerConfig(),
498 + protocol.MethodIncrement,
499 + handler,
500 + )
501 + defer ts.stop()
502 +
503 + client := NewIncrementClient(testRunDir, svc, testUnixShmClientConfig())
504 + defer client.Close()
505 +
506 + refreshUnixClientReady(t, client)
507 +
508 + _, err := client.CallIncrementBatch([]uint64{99})
509 + if !errors.Is(err, protocol.ErrBadLayout) {
510 + t.Fatalf("CallIncrementBatch handler failure = %v, want %v", err, protocol.ErrBadLayout)
511 + }
512 + if client.state != StateBroken {
513 + t.Fatalf("client state after batch handler failure = %d, want BROKEN", client.state)
514 + }
515 +
516 + changed := client.Refresh()
517 + if !changed || !client.Ready() {
518 + t.Fatalf("Refresh after batch handler failure = %v, state=%d, want READY", changed, client.state)
519 + }
520 +
521 + got, err := client.CallIncrement(5)
522 + if err != nil {
523 + t.Fatalf("CallIncrement after Refresh failed: %v", err)
524 + }
525 + if got != 6 {
526 + t.Fatalf("CallIncrement after Refresh = %d, want 6", got)
527 + }
528 +}
529 +
530 +func TestUnixShmBatchResponseOverflowRetriesAndRecovers(t *testing.T) {
531 + svc := uniqueUnixService("go_unix_shm_batch_overflow")
532 +
533 + scfg := testUnixShmServerConfig()
534 + scfg.MaxResponsePayloadBytes = 24
535 + scfg.MaxResponseBatchItems = 2
536 + ccfg := testUnixShmClientConfig()
537 + ccfg.MaxResponsePayloadBytes = 24
538 + ccfg.MaxResponseBatchItems = 2
539 + ccfg.MaxRequestBatchItems = 2
540 +
541 + ts := startTestServerUnixWithConfig(
542 + svc,
543 + scfg,
544 + protocol.MethodIncrement,
545 + IncrementDispatch(func(v uint64) (uint64, bool) {
546 + return v + 1, true
547 + }),
548 + )
549 + defer ts.stop()
550 +
551 + client := NewIncrementClient(testRunDir, svc, ccfg)
552 + defer client.Close()
553 +
554 + refreshUnixClientReady(t, client)
555 +
556 + gotBatch, err := client.CallIncrementBatch([]uint64{1, 2})
557 + if err != nil {
558 + t.Fatalf("CallIncrementBatch after SHM overflow failed: %v", err)
559 + }
560 + if len(gotBatch) != 2 || gotBatch[0] != 2 || gotBatch[1] != 3 {
561 + t.Fatalf("CallIncrementBatch after SHM overflow = %v, want [2 3]", gotBatch)
562 + }
563 + if client.state != StateReady {
564 + t.Fatalf("client state after batch overflow recovery = %d, want READY", client.state)
565 + }
566 + if client.reconnectCount < 1 {
567 + t.Fatalf("expected reconnect_count >= 1 after SHM batch overflow, got %d", client.reconnectCount)
568 + }
569 +
570 + got, err := client.CallIncrement(8)
571 + if err != nil {
572 + t.Fatalf("CallIncrement after SHM overflow recovery failed: %v", err)
573 + }
574 + if got != 9 {
575 + t.Fatalf("CallIncrement after SHM overflow recovery = %d, want 9", got)
576 + }
577 +}
src/go/pkg/netipc/service/raw/shm_windows_test.go new
+446
@@ -0,0 +1,446 @@
1 +//go:build windows
2 +
3 +package raw
4 +
5 +import (
6 + "errors"
7 + "testing"
8 + "time"
9 +
10 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
11 + windows "github.com/netdata/netdata/go/plugins/pkg/netipc/transport/windows"
12 +)
13 +
14 +func TestWinShmRoundTrip(t *testing.T) {
15 + t.Run("snapshot", func(t *testing.T) {
16 + svc := uniqueWinService("go_win_shm_roundtrip_snapshot")
17 + ts := startTestSnapshotServerWinWithConfig(svc, testWinShmServerConfig())
18 + defer ts.stop()
19 +
20 + client := NewSnapshotClient(winTestRunDir, svc, testWinShmClientConfig())
21 + defer client.Close()
22 +
23 + waitWinClientReady(t, client)
24 +
25 + if client.session == nil {
26 + t.Fatal("expected negotiated session")
27 + }
28 + if client.shm == nil {
29 + t.Fatal("expected WinSHM attachment")
30 + }
31 + if client.session.SelectedProfile != windows.WinShmProfileHybrid {
32 + t.Fatalf("selected profile = %d, want %d", client.session.SelectedProfile, windows.WinShmProfileHybrid)
33 + }
34 +
35 + view, err := client.CallSnapshot()
36 + if err != nil {
37 + t.Fatalf("CallSnapshot failed: %v", err)
38 + }
39 + if view.ItemCount != 3 {
40 + t.Fatalf("snapshot item count = %d, want 3", view.ItemCount)
41 + }
42 + if status := client.Status(); status.CallCount != 1 || status.ErrorCount != 0 {
43 + t.Fatalf("unexpected client status: %+v", status)
44 + }
45 + })
46 +
47 + t.Run("increment", func(t *testing.T) {
48 + svc := uniqueWinService("go_win_shm_roundtrip_increment")
49 + ts := startTestIncrementServerWinWithConfig(svc, testWinShmServerConfig())
50 + defer ts.stop()
51 +
52 + client := NewIncrementClient(winTestRunDir, svc, testWinShmClientConfig())
53 + defer client.Close()
54 + waitWinClientReady(t, client)
55 +
56 + got, err := client.CallIncrement(41)
57 + if err != nil {
58 + t.Fatalf("CallIncrement failed: %v", err)
59 + }
60 + if got != 42 {
61 + t.Fatalf("increment result = %d, want 42", got)
62 + }
63 +
64 + batch, err := client.CallIncrementBatch([]uint64{1, 41, 99})
65 + if err != nil {
66 + t.Fatalf("CallIncrementBatch failed: %v", err)
67 + }
68 + wantBatch := []uint64{2, 42, 100}
69 + if len(batch) != len(wantBatch) {
70 + t.Fatalf("batch result len = %d, want %d", len(batch), len(wantBatch))
71 + }
72 + for i, want := range wantBatch {
73 + if batch[i] != want {
74 + t.Fatalf("batch[%d] = %d, want %d", i, batch[i], want)
75 + }
76 + }
77 + })
78 +
79 + t.Run("string-reverse", func(t *testing.T) {
80 + svc := uniqueWinService("go_win_shm_roundtrip_reverse")
81 + ts := startTestStringReverseServerWinWithConfig(svc, testWinShmServerConfig())
82 + defer ts.stop()
83 +
84 + client := NewStringReverseClient(winTestRunDir, svc, testWinShmClientConfig())
85 + defer client.Close()
86 + waitWinClientReady(t, client)
87 +
88 + reversed, err := client.CallStringReverse("hello")
89 + if err != nil {
90 + t.Fatalf("CallStringReverse failed: %v", err)
91 + }
92 + if reversed.Str != "olleh" {
93 + t.Fatalf("string reverse result = %q, want %q", reversed.Str, "olleh")
94 + }
95 + })
96 +}
97 +
98 +func TestWinShmIdleTimeoutKeepsSessionAlive(t *testing.T) {
99 + svc := uniqueWinService("go_win_shm_idle")
100 + ts := startTestIncrementServerWinWithConfig(svc, testWinShmServerConfig())
101 + defer ts.stop()
102 +
103 + client := NewIncrementClient(winTestRunDir, svc, testWinShmClientConfig())
104 + defer client.Close()
105 +
106 + waitWinClientReady(t, client)
107 +
108 + time.Sleep(350 * time.Millisecond)
109 +
110 + got, err := client.CallIncrement(9)
111 + if err != nil {
112 + t.Fatalf("CallIncrement after idle timeout failed: %v", err)
113 + }
114 + if got != 10 {
115 + t.Fatalf("CallIncrement after idle timeout = %d, want 10", got)
116 + }
117 +}
118 +
119 +func TestWinServerRunInvalidServiceName(t *testing.T) {
120 + server := NewServer(winTestRunDir, "bad/name", testWinServerConfig(), protocol.MethodIncrement, winIncrementDispatchHandler())
121 + if err := server.Run(); err == nil {
122 + t.Fatal("expected Run() to fail for invalid service name")
123 + }
124 +}
125 +
126 +func TestWinShmAttachFailureFallsBackToBaseline(t *testing.T) {
127 + svc := uniqueWinService("go_win_shm_attach_fail")
128 + cfg := testWinShmServerConfig()
129 +
130 + listener, err := windows.Listen(winTestRunDir, svc, cfg)
131 + if err != nil {
132 + t.Fatalf("windows.Listen failed: %v", err)
133 + }
134 + defer listener.Close()
135 +
136 + type attachFailureResult struct {
137 + firstSelected uint32
138 + secondSelected uint32
139 + err error
140 + }
141 + doneCh := make(chan attachFailureResult, 1)
142 + go func() {
143 + first, err := listener.Accept()
144 + if err != nil {
145 + doneCh <- attachFailureResult{err: err}
146 + return
147 + }
148 + defer first.Close()
149 + firstSelected := first.SelectedProfile
150 +
151 + recvBuf := make([]byte, protocol.HeaderSize+4096)
152 + _, _, err = first.Receive(recvBuf)
153 + if err == nil {
154 + doneCh <- attachFailureResult{err: errors.New("expected first receive to fail after WinSHM attach fallback disconnect")}
155 + return
156 + }
157 +
158 + second, err := listener.Accept()
159 + if err != nil {
160 + doneCh <- attachFailureResult{err: err}
161 + return
162 + }
163 + defer second.Close()
164 + secondSelected := second.SelectedProfile
165 +
166 + _, _, err = second.Receive(recvBuf)
167 + if err == nil {
168 + doneCh <- attachFailureResult{err: errors.New("expected second receive to fail after client close")}
169 + return
170 + }
171 +
172 + doneCh <- attachFailureResult{
173 + firstSelected: firstSelected,
174 + secondSelected: secondSelected,
175 + }
176 + }()
177 +
178 + client := NewIncrementClient(winTestRunDir, svc, testWinShmClientConfig())
179 +
180 + if changed := client.Refresh(); !changed {
181 + t.Fatal("refresh should transition to READY via baseline fallback after WinSHM attach failure")
182 + }
183 + if !client.Ready() {
184 + t.Fatal("client should be ready after baseline fallback")
185 + }
186 + if client.state != StateReady {
187 + t.Fatalf("client state = %d, want READY", client.state)
188 + }
189 + if client.shm != nil {
190 + t.Fatal("expected no WinSHM attachment after fallback")
191 + }
192 + if client.session == nil {
193 + t.Fatal("expected live baseline session after fallback")
194 + }
195 + if client.session.SelectedProfile != protocol.ProfileBaseline {
196 + t.Fatalf("selected profile after fallback = %#x, want baseline", client.session.SelectedProfile)
197 + }
198 + if client.config.SupportedProfiles&windows.WinShmProfileHybrid != 0 {
199 + t.Fatalf("supported profiles should drop WinSHM after attach failure, got %#x", client.config.SupportedProfiles)
200 + }
201 + if client.config.PreferredProfiles&windows.WinShmProfileHybrid != 0 {
202 + t.Fatalf("preferred profiles should drop WinSHM after attach failure, got %#x", client.config.PreferredProfiles)
203 + }
204 +
205 + client.Close()
206 +
207 + result := <-doneCh
208 + if result.err != nil {
209 + t.Fatalf("attach-failure fallback server failed: %v", result.err)
210 + }
211 + if result.firstSelected != windows.WinShmProfileHybrid {
212 + t.Fatalf("first selected profile = %#x, want WinSHM hybrid", result.firstSelected)
213 + }
214 + if result.secondSelected != protocol.ProfileBaseline {
215 + t.Fatalf("second selected profile = %#x, want baseline", result.secondSelected)
216 + }
217 +}
218 +
219 +func TestWinShmMalformedShortRequestRecovers(t *testing.T) {
220 + svc := uniqueWinService("go_win_shm_short_req")
221 + ts := startTestIncrementServerWinWithConfig(svc, testWinShmServerConfig())
222 + defer ts.stop()
223 +
224 + client := NewIncrementClient(winTestRunDir, svc, testWinShmClientConfig())
225 + defer client.Close()
226 +
227 + waitWinClientReady(t, client)
228 +
229 + if client.shm == nil {
230 + t.Fatal("expected WinSHM attachment")
231 + }
232 + if err := client.shm.WinShmSend([]byte{1, 2, 3, 4}); err != nil {
233 + t.Fatalf("WinShmSend malformed short request failed: %v", err)
234 + }
235 +
236 + time.Sleep(50 * time.Millisecond)
237 +
238 + got, err := client.CallIncrement(9)
239 + if err != nil {
240 + t.Fatalf("CallIncrement after malformed short WinSHM request failed: %v", err)
241 + }
242 + if got != 10 {
243 + t.Fatalf("CallIncrement after malformed short WinSHM request = %d, want 10", got)
244 + }
245 + if client.Status().ReconnectCount < 1 {
246 + t.Fatalf("expected reconnect after malformed short WinSHM request, got status %+v", client.Status())
247 + }
248 +}
249 +
250 +func TestWinShmMalformedHeaderRequestRecovers(t *testing.T) {
251 + svc := uniqueWinService("go_win_shm_bad_hdr")
252 + ts := startTestIncrementServerWinWithConfig(svc, testWinShmServerConfig())
253 + defer ts.stop()
254 +
255 + client := NewIncrementClient(winTestRunDir, svc, testWinShmClientConfig())
256 + defer client.Close()
257 +
258 + waitWinClientReady(t, client)
259 +
260 + if client.shm == nil {
261 + t.Fatal("expected WinSHM attachment")
262 + }
263 + msg := make([]byte, protocol.HeaderSize)
264 + if err := client.shm.WinShmSend(msg); err != nil {
265 + t.Fatalf("WinShmSend malformed header request failed: %v", err)
266 + }
267 +
268 + time.Sleep(50 * time.Millisecond)
269 +
270 + got, err := client.CallIncrement(11)
271 + if err != nil {
272 + t.Fatalf("CallIncrement after malformed header WinSHM request failed: %v", err)
273 + }
274 + if got != 12 {
275 + t.Fatalf("CallIncrement after malformed header WinSHM request = %d, want 12", got)
276 + }
277 + if client.Status().ReconnectCount < 1 {
278 + t.Fatalf("expected reconnect after malformed header WinSHM request, got status %+v", client.Status())
279 + }
280 +}
281 +
282 +func TestWinShmUnexpectedMessageKindRecovers(t *testing.T) {
283 + svc := uniqueWinService("go_win_shm_bad_kind")
284 + ts := startTestIncrementServerWinWithConfig(svc, testWinShmServerConfig())
285 + defer ts.stop()
286 +
287 + client := NewIncrementClient(winTestRunDir, svc, testWinShmClientConfig())
288 + defer client.Close()
289 +
290 + waitWinClientReady(t, client)
291 +
292 + if client.shm == nil {
293 + t.Fatal("expected WinSHM attachment")
294 + }
295 +
296 + reqHdr := protocol.Header{
297 + Kind: protocol.KindResponse,
298 + Code: protocol.MethodIncrement,
299 + ItemCount: 1,
300 + MessageID: 1,
301 + TransportStatus: protocol.StatusOK,
302 + }
303 + var reqPayload [protocol.IncrementPayloadSize]byte
304 + if protocol.IncrementEncode(9, reqPayload[:]) == 0 {
305 + t.Fatal("IncrementEncode failed")
306 + }
307 +
308 + if err := client.shm.WinShmSend(encodeRawWinMessage(reqHdr, reqPayload[:])); err != nil {
309 + t.Fatalf("WinShmSend unexpected-kind request failed: %v", err)
310 + }
311 +
312 + time.Sleep(50 * time.Millisecond)
313 +
314 + got, err := client.CallIncrement(13)
315 + if err != nil {
316 + t.Fatalf("CallIncrement after unexpected-kind WinSHM request failed: %v", err)
317 + }
318 + if got != 14 {
319 + t.Fatalf("CallIncrement after unexpected-kind WinSHM request = %d, want 14", got)
320 + }
321 + if client.Status().ReconnectCount < 1 {
322 + t.Fatalf("expected reconnect after unexpected-kind WinSHM request, got status %+v", client.Status())
323 + }
324 +}
325 +
326 +func TestWinShmMalformedBatchRequestRecovers(t *testing.T) {
327 + svc := uniqueWinService("go_win_shm_bad_batch")
328 + ts := startTestIncrementServerWinWithConfig(svc, testWinShmServerConfig())
329 + defer ts.stop()
330 +
331 + client := NewIncrementClient(winTestRunDir, svc, testWinShmClientConfig())
332 + defer client.Close()
333 +
334 + waitWinClientReady(t, client)
335 +
336 + if client.shm == nil {
337 + t.Fatal("expected WinSHM attachment")
338 + }
339 +
340 + reqHdr := protocol.Header{
341 + Kind: protocol.KindRequest,
342 + Code: protocol.MethodIncrement,
343 + Flags: protocol.FlagBatch,
344 + ItemCount: 2,
345 + MessageID: 1,
346 + TransportStatus: protocol.StatusOK,
347 + }
348 + // ItemCount=2 requires a 16-byte aligned directory. An 8-byte payload
349 + // forces BatchItemGet() to fail in the server batch loop.
350 + badPayload := encodeWinIncrementBatchPayload(t, 1, 2)[:8]
351 + if err := client.shm.WinShmSend(encodeRawWinMessage(reqHdr, badPayload)); err != nil {
352 + t.Fatalf("WinShmSend malformed batch request failed: %v", err)
353 + }
354 +
355 + time.Sleep(50 * time.Millisecond)
356 +
357 + got, err := client.CallIncrement(21)
358 + if err != nil {
359 + t.Fatalf("CallIncrement after malformed batch WinSHM request failed: %v", err)
360 + }
361 + if got != 22 {
362 + t.Fatalf("CallIncrement after malformed batch WinSHM request = %d, want 22", got)
363 + }
364 + if !client.Ready() {
365 + t.Fatalf("client should stay usable after malformed batch WinSHM request, got status %+v", client.Status())
366 + }
367 +}
368 +
369 +func TestWinShmBatchHandlerFailureNeedsRefresh(t *testing.T) {
370 + svc := uniqueWinService("go_win_shm_batch_fail")
371 + ts := startTestServerWinWithConfig(svc, testWinShmServerConfig(), protocol.MethodIncrement, IncrementDispatch(func(v uint64) (uint64, bool) {
372 + if v == 99 {
373 + return 0, false
374 + }
375 + return v + 1, true
376 + }))
377 + defer ts.stop()
378 +
379 + client := NewIncrementClient(winTestRunDir, svc, testWinShmClientConfig())
380 + defer client.Close()
381 +
382 + waitWinClientReady(t, client)
383 +
384 + _, err := client.CallIncrementBatch([]uint64{99})
385 + if !errors.Is(err, protocol.ErrBadLayout) {
386 + t.Fatalf("CallIncrementBatch handler failure = %v, want %v", err, protocol.ErrBadLayout)
387 + }
388 + if client.state != StateBroken {
389 + t.Fatalf("client state after batch handler failure = %d, want BROKEN", client.state)
390 + }
391 +
392 + changed := client.Refresh()
393 + if !changed || !client.Ready() {
394 + t.Fatalf("Refresh after batch handler failure = %v, state=%d, want READY", changed, client.state)
395 + }
396 +
397 + got, err := client.CallIncrement(5)
398 + if err != nil {
399 + t.Fatalf("CallIncrement after Refresh failed: %v", err)
400 + }
401 + if got != 6 {
402 + t.Fatalf("CallIncrement after Refresh = %d, want 6", got)
403 + }
404 +}
405 +
406 +func TestWinShmBatchResponseOverflowRetriesAndRecovers(t *testing.T) {
407 + svc := uniqueWinService("go_win_shm_batch_overflow")
408 +
409 + scfg := testWinShmServerConfig()
410 + scfg.MaxResponsePayloadBytes = 24
411 + scfg.MaxResponseBatchItems = 2
412 + ccfg := testWinShmClientConfig()
413 + ccfg.MaxResponsePayloadBytes = 24
414 + ccfg.MaxResponseBatchItems = 2
415 + ccfg.MaxRequestBatchItems = 2
416 +
417 + ts := startTestServerWinWithConfig(svc, scfg, protocol.MethodIncrement, winIncrementDispatchHandler())
418 + defer ts.stop()
419 +
420 + client := NewIncrementClient(winTestRunDir, svc, ccfg)
421 + defer client.Close()
422 +
423 + waitWinClientReady(t, client)
424 +
425 + gotBatch, err := client.CallIncrementBatch([]uint64{1, 2})
426 + if err != nil {
427 + t.Fatalf("CallIncrementBatch after WinSHM overflow failed: %v", err)
428 + }
429 + if len(gotBatch) != 2 || gotBatch[0] != 2 || gotBatch[1] != 3 {
430 + t.Fatalf("CallIncrementBatch after WinSHM overflow = %v, want [2 3]", gotBatch)
431 + }
432 + if client.state != StateReady {
433 + t.Fatalf("client state after batch overflow recovery = %d, want READY", client.state)
434 + }
435 + if client.reconnectCount < 1 {
436 + t.Fatalf("expected reconnect_count >= 1 after WinSHM batch overflow, got %d", client.reconnectCount)
437 + }
438 +
439 + got, err := client.CallIncrement(8)
440 + if err != nil {
441 + t.Fatalf("CallIncrement after WinSHM overflow recovery failed: %v", err)
442 + }
443 + if got != 9 {
444 + t.Fatalf("CallIncrement after WinSHM overflow recovery = %d, want 9", got)
445 + }
446 +}
src/go/pkg/netipc/service/raw/stress_test.go new
+691
@@ -0,0 +1,691 @@
1 +//go:build unix
2 +
3 +package raw
4 +
5 +import (
6 + "fmt"
7 + "sync"
8 + "sync/atomic"
9 + "testing"
10 + "time"
11 +
12 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
13 + "github.com/netdata/netdata/go/plugins/pkg/netipc/transport/posix"
14 +)
15 +
16 +// simpleHash: djb2 matching the C implementation
17 +func simpleHash(s string) uint32 {
18 + var hash uint32 = 5381
19 + for _, c := range []byte(s) {
20 + hash = ((hash << 5) + hash) + uint32(c)
21 + }
22 + return hash
23 +}
24 +
25 +// largeHandler builds a snapshot with N items (realistic container names).
26 +func largeHandler(n int) DispatchHandler {
27 + return SnapshotDispatch(func(request *protocol.CgroupsRequest, builder *protocol.CgroupsBuilder) bool {
28 + if request.LayoutVersion != 1 || request.Flags != 0 {
29 + return false
30 + }
31 + builder.SetHeader(1, 42)
32 +
33 + for i := 0; i < n; i++ {
34 + name := fmt.Sprintf("container-%04d", i)
35 + path := fmt.Sprintf("/sys/fs/cgroup/docker/%04d", i)
36 + hash := simpleHash(name)
37 + enabled := uint32(1)
38 + if i%5 == 0 {
39 + enabled = 0
40 + }
41 + if err := builder.Add(hash, 0x10, enabled,
42 + []byte(name), []byte(path)); err != nil {
43 + return false
44 + }
45 + }
46 + return true
47 + }, uint32(n))
48 +}
49 +
50 +// verifySnapshot checks all items in a snapshot view match expected data.
51 +func verifySnapshot(t *testing.T, view *protocol.CgroupsResponseView, n int) bool {
52 + t.Helper()
53 + if int(view.ItemCount) != n {
54 + t.Errorf("expected %d items, got %d", n, view.ItemCount)
55 + return false
56 + }
57 + if view.SystemdEnabled != 1 {
58 + t.Errorf("expected systemd_enabled=1, got %d", view.SystemdEnabled)
59 + return false
60 + }
61 + if view.Generation != 42 {
62 + t.Errorf("expected generation=42, got %d", view.Generation)
63 + return false
64 + }
65 +
66 + // Spot-check first, middle, last
67 + indices := []int{0, n / 2, n - 1}
68 + for _, idx := range indices {
69 + item, err := view.Item(uint32(idx))
70 + if err != nil {
71 + t.Errorf("item %d decode error: %v", idx, err)
72 + return false
73 + }
74 + expectedName := fmt.Sprintf("container-%04d", idx)
75 + expectedPath := fmt.Sprintf("/sys/fs/cgroup/docker/%04d", idx)
76 + expectedHash := simpleHash(expectedName)
77 + expectedEnabled := uint32(1)
78 + if idx%5 == 0 {
79 + expectedEnabled = 0
80 + }
81 +
82 + if item.Hash != expectedHash {
83 + t.Errorf("item %d: hash=%d, expected=%d", idx, item.Hash, expectedHash)
84 + return false
85 + }
86 + if item.Name.String() != expectedName {
87 + t.Errorf("item %d: name=%q, expected=%q", idx, item.Name.String(), expectedName)
88 + return false
89 + }
90 + if item.Path.String() != expectedPath {
91 + t.Errorf("item %d: path=%q, expected=%q", idx, item.Path.String(), expectedPath)
92 + return false
93 + }
94 + if item.Enabled != expectedEnabled {
95 + t.Errorf("item %d: enabled=%d, expected=%d", idx, item.Enabled, expectedEnabled)
96 + return false
97 + }
98 + if item.Options != 0x10 {
99 + t.Errorf("item %d: options=0x%x, expected=0x10", idx, item.Options)
100 + return false
101 + }
102 + }
103 + return true
104 +}
105 +
106 +// startLargeServer creates a server with a large response buffer.
107 +// startLargeServer creates a server with large response buffer and explicit
108 +// packet_size to enable chunked transport for payloads > 64KB.
109 +func startLargeServer(service string, handler DispatchHandler, maxResp uint32) *testServer {
110 + ensureRunDir()
111 + cleanupAll(service)
112 +
113 + cfg := testServerConfig()
114 + cfg.MaxResponsePayloadBytes = maxResp
115 + cfg.PacketSize = 65536 // force smaller packets for reliable chunking
116 +
117 + s := NewServer(testRunDir, service, cfg, protocol.MethodCgroupsSnapshot, handler)
118 + doneCh := make(chan struct{})
119 +
120 + go func() {
121 + defer close(doneCh)
122 + s.Run()
123 + }()
124 +
125 + waitUnixServerReady(service)
126 + return &testServer{server: s, doneCh: doneCh}
127 +}
128 +
129 +// startServerWithWorkers creates a server with explicit worker count.
130 +func startServerWithWorkers(service string, handler DispatchHandler, workers int) *testServer {
131 + ensureRunDir()
132 + cleanupAll(service)
133 +
134 + s := NewServerWithWorkers(
135 + testRunDir,
136 + service,
137 + testServerConfig(),
138 + protocol.MethodCgroupsSnapshot,
139 + handler,
140 + workers,
141 + )
142 + doneCh := make(chan struct{})
143 +
144 + go func() {
145 + defer close(doneCh)
146 + s.Run()
147 + }()
148 +
149 + waitUnixServerReady(service)
150 + return &testServer{server: s, doneCh: doneCh}
151 +}
152 +
153 +func TestStress1000Items(t *testing.T) {
154 + const N = 1000
155 + const bufSize = 300 * N
156 +
157 + svc := "go_stress_1k"
158 + ts := startLargeServer(svc, largeHandler(N), bufSize)
159 + defer ts.stop()
160 +
161 + ccfg := testClientConfig()
162 + ccfg.MaxResponsePayloadBytes = bufSize
163 + ccfg.PacketSize = 65536
164 +
165 + client := NewSnapshotClient(testRunDir, svc, ccfg)
166 + client.Refresh()
167 + if !client.Ready() {
168 + t.Fatal("client not ready")
169 + }
170 +
171 + start := time.Now()
172 + view, err := client.CallSnapshot()
173 + elapsed := time.Since(start)
174 + if err != nil {
175 + t.Fatalf("call failed: %v", err)
176 + }
177 +
178 + t.Logf("1000 items: %v", elapsed)
179 +
180 + // Verify ALL items
181 + if int(view.ItemCount) != N {
182 + t.Fatalf("expected %d items, got %d", N, view.ItemCount)
183 + }
184 + for i := 0; i < N; i++ {
185 + item, ierr := view.Item(uint32(i))
186 + if ierr != nil {
187 + t.Fatalf("item %d decode error: %v", i, ierr)
188 + }
189 + expectedName := fmt.Sprintf("container-%04d", i)
190 + expectedPath := fmt.Sprintf("/sys/fs/cgroup/docker/%04d", i)
191 + expectedHash := simpleHash(expectedName)
192 + if item.Hash != expectedHash {
193 + t.Fatalf("item %d: hash mismatch", i)
194 + }
195 + if item.Name.String() != expectedName {
196 + t.Fatalf("item %d: name=%q expected=%q", i, item.Name.String(), expectedName)
197 + }
198 + if item.Path.String() != expectedPath {
199 + t.Fatalf("item %d: path=%q expected=%q", i, item.Path.String(), expectedPath)
200 + }
201 + }
202 +
203 + client.Close()
204 + cleanupAll(svc)
205 +}
206 +
207 +func TestStress5000Items(t *testing.T) {
208 + const N = 5000
209 + const bufSize = 300 * N
210 +
211 + svc := "go_stress_5k"
212 + ts := startLargeServer(svc, largeHandler(N), bufSize)
213 + defer ts.stop()
214 +
215 + ccfg := testClientConfig()
216 + ccfg.MaxResponsePayloadBytes = bufSize
217 + ccfg.PacketSize = 65536
218 +
219 + client := NewSnapshotClient(testRunDir, svc, ccfg)
220 + client.Refresh()
221 + if !client.Ready() {
222 + t.Fatal("client not ready")
223 + }
224 +
225 + start := time.Now()
226 + view, err := client.CallSnapshot()
227 + elapsed := time.Since(start)
228 + if err != nil {
229 + t.Fatalf("call failed: %v", err)
230 + }
231 +
232 + t.Logf("5000 items: %v", elapsed)
233 +
234 + if int(view.ItemCount) != N {
235 + t.Fatalf("expected %d items, got %d", N, view.ItemCount)
236 + }
237 +
238 + // Verify spot checks
239 + if !verifySnapshot(t, view, N) {
240 + t.Fatal("snapshot verification failed")
241 + }
242 +
243 + client.Close()
244 + cleanupAll(svc)
245 +}
246 +
247 +func TestStress50Clients(t *testing.T) {
248 + svc := "go_stress_mc50"
249 + ensureRunDir()
250 + cleanupAll(svc)
251 +
252 + ts := startServerWithWorkers(svc, testSnapshotDispatch(), 64)
253 + defer ts.stop()
254 +
255 + const numClients = 50
256 + const requestsPerClient = 10
257 +
258 + type result struct {
259 + clientID int
260 + successes int
261 + failures int
262 + }
263 +
264 + results := make(chan result, numClients)
265 +
266 + start := time.Now()
267 +
268 + for i := 0; i < numClients; i++ {
269 + go func(id int) {
270 + r := result{clientID: id}
271 + client := NewSnapshotClient(testRunDir, svc, testClientConfig())
272 + defer client.Close()
273 +
274 + for retry := 0; retry < 200; retry++ {
275 + client.Refresh()
276 + if client.Ready() {
277 + break
278 + }
279 + time.Sleep(5 * time.Millisecond)
280 + }
281 +
282 + if !client.Ready() {
283 + r.failures = requestsPerClient
284 + results <- r
285 + return
286 + }
287 +
288 + for j := 0; j < requestsPerClient; j++ {
289 + view, err := client.CallSnapshot()
290 + if err != nil || view.ItemCount != 3 {
291 + r.failures++
292 + continue
293 + }
294 + // Verify content correctness
295 + item0, ierr := view.Item(0)
296 + if ierr != nil || item0.Hash != 1001 ||
297 + item0.Name.String() != "docker-abc123" {
298 + r.failures++
299 + continue
300 + }
301 + item2, ierr := view.Item(2)
302 + if ierr != nil || item2.Hash != 3003 ||
303 + item2.Name.String() != "systemd-user" {
304 + r.failures++
305 + continue
306 + }
307 + r.successes++
308 + }
309 + results <- r
310 + }(i)
311 + }
312 +
313 + totalSuccess := 0
314 + totalFailure := 0
315 + for i := 0; i < numClients; i++ {
316 + r := <-results
317 + totalSuccess += r.successes
318 + totalFailure += r.failures
319 + }
320 +
321 + elapsed := time.Since(start)
322 + expected := numClients * requestsPerClient
323 + t.Logf("50 clients x 10 req: %d/%d succeeded, %d failures, %v",
324 + totalSuccess, expected, totalFailure, elapsed)
325 +
326 + if totalSuccess != expected {
327 + t.Fatalf("expected %d successes, got %d (failures: %d)",
328 + expected, totalSuccess, totalFailure)
329 + }
330 + if totalFailure != 0 {
331 + t.Fatalf("expected 0 failures, got %d", totalFailure)
332 + }
333 +
334 + cleanupAll(svc)
335 +}
336 +
337 +func TestStressConcurrentCacheClients(t *testing.T) {
338 + svc := "go_stress_cache10"
339 + ensureRunDir()
340 + cleanupAll(svc)
341 +
342 + ts := startServerWithWorkers(svc, testSnapshotDispatch(), 16)
343 + defer ts.stop()
344 +
345 + const numClients = 10
346 + const requestsPerClient = 100
347 +
348 + type result struct {
349 + successes int
350 + failures int
351 + }
352 +
353 + results := make(chan result, numClients)
354 +
355 + start := time.Now()
356 +
357 + for i := 0; i < numClients; i++ {
358 + go func() {
359 + r := result{}
360 + cache := NewCache(testRunDir, svc, testClientConfig())
361 + defer cache.Close()
362 +
363 + for j := 0; j < requestsPerClient; j++ {
364 + updated := cache.Refresh()
365 + if updated || cache.Ready() {
366 + status := cache.Status()
367 + if status.ItemCount != 3 {
368 + r.failures++
369 + continue
370 + }
371 + // Verify a lookup
372 + item, found := cache.Lookup(1001, "docker-abc123")
373 + if !found || item.Hash != 1001 ||
374 + item.Path != "/sys/fs/cgroup/docker/abc123" {
375 + r.failures++
376 + continue
377 + }
378 + r.successes++
379 + } else {
380 + r.failures++
381 + }
382 + }
383 + results <- r
384 + }()
385 + }
386 +
387 + totalSuccess := 0
388 + totalFailure := 0
389 + for i := 0; i < numClients; i++ {
390 + r := <-results
391 + totalSuccess += r.successes
392 + totalFailure += r.failures
393 + }
394 +
395 + elapsed := time.Since(start)
396 + expected := numClients * requestsPerClient
397 + t.Logf("10 cache clients x 100 req: %d/%d succeeded, %d failures, %v",
398 + totalSuccess, expected, totalFailure, elapsed)
399 +
400 + if totalSuccess != expected {
401 + t.Fatalf("expected %d successes, got %d", expected, totalSuccess)
402 + }
403 + if totalFailure != 0 {
404 + t.Fatalf("expected 0 failures, got %d", totalFailure)
405 + }
406 +
407 + cleanupAll(svc)
408 +}
409 +
410 +func TestStressRapidConnectDisconnect(t *testing.T) {
411 + svc := "go_stress_rapid"
412 + ensureRunDir()
413 + cleanupAll(svc)
414 +
415 + ts := startServerWithWorkers(svc, testSnapshotDispatch(), 16)
416 + defer ts.stop()
417 +
418 + const cycles = 1000
419 + successes := 0
420 + failures := 0
421 +
422 + start := time.Now()
423 +
424 + for i := 0; i < cycles; i++ {
425 + client := NewSnapshotClient(testRunDir, svc, testClientConfig())
426 +
427 + for r := 0; r < 50; r++ {
428 + client.Refresh()
429 + if client.Ready() {
430 + break
431 + }
432 + time.Sleep(2 * time.Millisecond)
433 + }
434 +
435 + if !client.Ready() {
436 + failures++
437 + client.Close()
438 + continue
439 + }
440 +
441 + view, err := client.CallSnapshot()
442 + if err != nil || view.ItemCount != 3 {
443 + failures++
444 + } else {
445 + item0, ierr := view.Item(0)
446 + if ierr != nil || item0.Hash != 1001 {
447 + failures++
448 + } else {
449 + successes++
450 + }
451 + }
452 + client.Close()
453 + }
454 +
455 + elapsed := time.Since(start)
456 + t.Logf("1000 rapid cycles: %d ok, %d fail, %v", successes, failures, elapsed)
457 +
458 + if successes != cycles {
459 + t.Fatalf("expected %d successes, got %d (failures: %d)",
460 + cycles, successes, failures)
461 + }
462 +
463 + cleanupAll(svc)
464 +}
465 +
466 +func TestStressLongRunning60s(t *testing.T) {
467 + if testing.Short() {
468 + t.Skip("skipping 60s test in short mode")
469 + }
470 +
471 + svc := "go_stress_long"
472 + ensureRunDir()
473 + cleanupAll(svc)
474 +
475 + ts := startServerWithWorkers(svc, testSnapshotDispatch(), 8)
476 + defer ts.stop()
477 +
478 + const numClients = 5
479 + const duration = 60 * time.Second
480 +
481 + var totalRefreshes int64
482 + var totalErrors int64
483 + var wg sync.WaitGroup
484 + stop := make(chan struct{})
485 +
486 + for i := 0; i < numClients; i++ {
487 + wg.Add(1)
488 + go func() {
489 + defer wg.Done()
490 + cache := NewCache(testRunDir, svc, testClientConfig())
491 + defer cache.Close()
492 +
493 + for {
494 + select {
495 + case <-stop:
496 + return
497 + default:
498 + }
499 +
500 + updated := cache.Refresh()
501 + if updated || cache.Ready() {
502 + status := cache.Status()
503 + if status.ItemCount != 3 {
504 + atomic.AddInt64(&totalErrors, 1)
505 + } else {
506 + atomic.AddInt64(&totalRefreshes, 1)
507 + }
508 + } else {
509 + atomic.AddInt64(&totalErrors, 1)
510 + }
511 +
512 + time.Sleep(time.Millisecond)
513 + }
514 + }()
515 + }
516 +
517 + time.Sleep(duration)
518 + close(stop)
519 + wg.Wait()
520 +
521 + refreshes := atomic.LoadInt64(&totalRefreshes)
522 + errors := atomic.LoadInt64(&totalErrors)
523 +
524 + t.Logf("60s run: %d refreshes, %d errors", refreshes, errors)
525 +
526 + if refreshes == 0 {
527 + t.Fatal("expected some refreshes to succeed")
528 + }
529 + if errors != 0 {
530 + t.Fatalf("expected 0 errors, got %d", errors)
531 + }
532 +
533 + cleanupAll(svc)
534 +}
535 +
536 +func TestStressMixedTransport(t *testing.T) {
537 + svc := "go_stress_mixed"
538 + ensureRunDir()
539 + cleanupAll(svc)
540 +
541 + // Server supports both baseline and SHM
542 + cfg := posix.ServerConfig{
543 + SupportedProfiles: protocol.ProfileBaseline | protocol.ProfileSHMHybrid,
544 + PreferredProfiles: protocol.ProfileSHMHybrid,
545 + MaxRequestPayloadBytes: 4096,
546 + MaxRequestBatchItems: 1,
547 + MaxResponsePayloadBytes: responseBufSize,
548 + MaxResponseBatchItems: 1,
549 + AuthToken: authToken,
550 + Backlog: 16,
551 + }
552 +
553 + s := NewServer(
554 + testRunDir,
555 + svc,
556 + cfg,
557 + protocol.MethodCgroupsSnapshot,
558 + testSnapshotDispatch(),
559 + )
560 + doneCh := make(chan struct{})
561 + go func() {
562 + defer close(doneCh)
563 + s.Run()
564 + }()
565 + waitUnixServerReady(svc)
566 +
567 + defer func() {
568 + s.Stop()
569 + <-doneCh
570 + }()
571 +
572 + type result struct {
573 + clientID int
574 + profile string
575 + success int
576 + failure int
577 + }
578 +
579 + results := make(chan result, 3)
580 +
581 + // Client 0: SHM-capable
582 + go func() {
583 + ccfg := testClientConfig()
584 + ccfg.SupportedProfiles = protocol.ProfileBaseline | protocol.ProfileSHMHybrid
585 + ccfg.PreferredProfiles = protocol.ProfileSHMHybrid
586 +
587 + r := result{clientID: 0, profile: "SHM"}
588 + client := NewSnapshotClient(testRunDir, svc, ccfg)
589 + defer client.Close()
590 +
591 + for retry := 0; retry < 200; retry++ {
592 + client.Refresh()
593 + if client.Ready() {
594 + break
595 + }
596 + time.Sleep(5 * time.Millisecond)
597 + }
598 +
599 + for i := 0; i < 10; i++ {
600 + view, err := client.CallSnapshot()
601 + if err == nil && view.ItemCount == 3 {
602 + item0, ierr := view.Item(0)
603 + if ierr == nil && item0.Hash == 1001 {
604 + r.success++
605 + continue
606 + }
607 + }
608 + r.failure++
609 + }
610 + results <- r
611 + }()
612 +
613 + // Client 1: SHM-capable
614 + go func() {
615 + ccfg := testClientConfig()
616 + ccfg.SupportedProfiles = protocol.ProfileBaseline | protocol.ProfileSHMHybrid
617 + ccfg.PreferredProfiles = protocol.ProfileSHMHybrid
618 +
619 + r := result{clientID: 1, profile: "SHM"}
620 + client := NewSnapshotClient(testRunDir, svc, ccfg)
621 + defer client.Close()
622 +
623 + for retry := 0; retry < 200; retry++ {
624 + client.Refresh()
625 + if client.Ready() {
626 + break
627 + }
628 + time.Sleep(5 * time.Millisecond)
629 + }
630 +
631 + for i := 0; i < 10; i++ {
632 + view, err := client.CallSnapshot()
633 + if err == nil && view.ItemCount == 3 {
634 + r.success++
635 + } else {
636 + r.failure++
637 + }
638 + }
639 + results <- r
640 + }()
641 +
642 + // Client 2: UDS-only (baseline)
643 + go func() {
644 + ccfg := testClientConfig()
645 + ccfg.SupportedProfiles = protocol.ProfileBaseline
646 + ccfg.PreferredProfiles = 0
647 +
648 + r := result{clientID: 2, profile: "UDS"}
649 + client := NewSnapshotClient(testRunDir, svc, ccfg)
650 + defer client.Close()
651 +
652 + for retry := 0; retry < 200; retry++ {
653 + client.Refresh()
654 + if client.Ready() {
655 + break
656 + }
657 + time.Sleep(5 * time.Millisecond)
658 + }
659 +
660 + for i := 0; i < 10; i++ {
661 + view, err := client.CallSnapshot()
662 + if err == nil && view.ItemCount == 3 {
663 + item0, ierr := view.Item(0)
664 + if ierr == nil && item0.Hash == 1001 {
665 + r.success++
666 + continue
667 + }
668 + }
669 + r.failure++
670 + }
671 + results <- r
672 + }()
673 +
674 + totalSuccess := 0
675 + totalFailure := 0
676 + for i := 0; i < 3; i++ {
677 + r := <-results
678 + t.Logf("client %d (%s): %d ok, %d fail", r.clientID, r.profile, r.success, r.failure)
679 + totalSuccess += r.success
680 + totalFailure += r.failure
681 + }
682 +
683 + if totalSuccess != 30 {
684 + t.Fatalf("expected 30 successes, got %d (failures: %d)", totalSuccess, totalFailure)
685 + }
686 + if totalFailure != 0 {
687 + t.Fatalf("expected 0 failures, got %d", totalFailure)
688 + }
689 +
690 + cleanupAll(svc)
691 +}
src/go/pkg/netipc/service/raw/types.go new
+170
@@ -0,0 +1,170 @@
1 +// Package raw provides internal L2/L3 helpers used by the protocol, fixtures,
2 +// and benchmarks.
3 +//
4 +// Pure composition of L1 transport + Codec. No direct socket/pipe calls.
5 +// Client manages connection lifecycle with at-least-once retry.
6 +// Server handles accept, read, dispatch, respond.
7 +//
8 +// Pure Go — no cgo. Works with CGO_ENABLED=0.
9 +package raw
10 +
11 +import (
12 + "errors"
13 +
14 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
15 +)
16 +
17 +// Poll/receive timeout for server loops (ms). Controls shutdown detection latency.
18 +const serverPollTimeoutMs = 100
19 +
20 +// ---------------------------------------------------------------------------
21 +// Client state (shared across platforms)
22 +// ---------------------------------------------------------------------------
23 +
24 +// ClientState represents the connection state machine.
25 +type ClientState int
26 +
27 +const (
28 + StateDisconnected ClientState = iota
29 + StateConnecting
30 + StateReady
31 + StateNotFound
32 + StateAuthFailed
33 + StateIncompatible
34 + StateBroken
35 +)
36 +
37 +// ClientStatus is a diagnostic counters snapshot.
38 +type ClientStatus struct {
39 + State ClientState
40 + ConnectCount uint32
41 + ReconnectCount uint32
42 + CallCount uint32
43 + ErrorCount uint32
44 +}
45 +
46 +// IncrementHandler serves a single INCREMENT service kind.
47 +type IncrementHandler func(uint64) (uint64, bool)
48 +
49 +// StringReverseHandler serves a single STRING_REVERSE service kind.
50 +type StringReverseHandler func(string) (string, bool)
51 +
52 +// SnapshotHandler serves a single CGROUPS_SNAPSHOT service kind.
53 +type SnapshotHandler func(*protocol.CgroupsRequest, *protocol.CgroupsBuilder) bool
54 +
55 +var errHandlerFailed = errors.New("dispatch handler failed")
56 +
57 +// DispatchHandler validates/decodes a single service kind request and writes
58 +// the matching response into responseBuf.
59 +type DispatchHandler func(request []byte, responseBuf []byte) (int, error)
60 +
61 +// IncrementDispatch adapts a typed increment handler to the raw dispatch shape.
62 +func IncrementDispatch(handle IncrementHandler) DispatchHandler {
63 + if handle == nil {
64 + return nil
65 + }
66 + return func(request []byte, responseBuf []byte) (int, error) {
67 + value, err := protocol.IncrementDecode(request)
68 + if err != nil {
69 + return 0, err
70 + }
71 + result, ok := handle(value)
72 + if !ok {
73 + return 0, errHandlerFailed
74 + }
75 + n := protocol.IncrementEncode(result, responseBuf)
76 + if n == 0 {
77 + return 0, protocol.ErrOverflow
78 + }
79 + return n, nil
80 + }
81 +}
82 +
83 +// StringReverseDispatch adapts a typed string-reverse handler to the raw dispatch shape.
84 +func StringReverseDispatch(handle StringReverseHandler) DispatchHandler {
85 + if handle == nil {
86 + return nil
87 + }
88 + return func(request []byte, responseBuf []byte) (int, error) {
89 + view, err := protocol.StringReverseDecode(request)
90 + if err != nil {
91 + return 0, err
92 + }
93 + result, ok := handle(view.Str)
94 + if !ok {
95 + return 0, errHandlerFailed
96 + }
97 + n := protocol.StringReverseEncode(result, responseBuf)
98 + if n == 0 {
99 + return 0, protocol.ErrOverflow
100 + }
101 + return n, nil
102 + }
103 +}
104 +
105 +// SnapshotMaxItems returns the item budget for a single snapshot service kind.
106 +func SnapshotMaxItems(responseBufSize int, override uint32) uint32 {
107 + if override != 0 {
108 + return override
109 + }
110 + return protocol.EstimateCgroupsMaxItems(responseBufSize)
111 +}
112 +
113 +// SnapshotDispatch adapts a typed snapshot handler to the raw dispatch shape.
114 +func SnapshotDispatch(handle SnapshotHandler, maxItems uint32) DispatchHandler {
115 + if handle == nil {
116 + return nil
117 + }
118 + return func(request []byte, responseBuf []byte) (int, error) {
119 + req, err := protocol.DecodeCgroupsRequest(request)
120 + if err != nil {
121 + return 0, err
122 + }
123 + itemBudget := SnapshotMaxItems(len(responseBuf), maxItems)
124 + if itemBudget == 0 {
125 + return 0, protocol.ErrOverflow
126 + }
127 + minRequired, ok := protocol.CgroupsBuilderMinBytes(itemBudget)
128 + if !ok || len(responseBuf) < minRequired {
129 + return 0, protocol.ErrOverflow
130 + }
131 + builder := protocol.NewCgroupsBuilder(responseBuf, itemBudget, 0, 0)
132 + if !handle(&req, builder) {
133 + return 0, errHandlerFailed
134 + }
135 + n := builder.Finish()
136 + if n == 0 {
137 + return 0, protocol.ErrOverflow
138 + }
139 + return n, nil
140 + }
141 +}
142 +
143 +// ---------------------------------------------------------------------------
144 +// L3 cache types (shared across platforms)
145 +// ---------------------------------------------------------------------------
146 +
147 +// Default response buffer size for L3 cache refresh.
148 +const cacheResponseBufSize = 65536
149 +
150 +// CacheItem is an owned copy of a single cgroup item.
151 +// Built from ephemeral L2 views during cache construction.
152 +type CacheItem struct {
153 + Hash uint32
154 + Options uint32
155 + Enabled uint32
156 + Name string // owned copy
157 + Path string // owned copy
158 +}
159 +
160 +// CacheStatus is a diagnostic snapshot for the L3 cache.
161 +type CacheStatus struct {
162 + Populated bool
163 + ItemCount uint32
164 + SystemdEnabled uint32
165 + Generation uint64
166 + RefreshSuccessCount uint32
167 + RefreshFailureCount uint32
168 + ConnectionState ClientState // underlying L2 client state
169 + LastRefreshTs int64 // monotonic timestamp (ms) of last successful refresh, 0 if never
170 +}
src/go/pkg/netipc/service/raw/types_more_test.go new
+19
@@ -0,0 +1,19 @@
1 +package raw
2 +
3 +import (
4 + "testing"
5 +
6 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
7 +)
8 +
9 +func TestSnapshotMaxItems(t *testing.T) {
10 + got := SnapshotMaxItems(4096, 0)
11 + want := protocol.EstimateCgroupsMaxItems(4096)
12 + if got != want {
13 + t.Fatalf("SnapshotMaxItems default = %d, want %d", got, want)
14 + }
15 +
16 + if got := SnapshotMaxItems(4096, 7); got != 7 {
17 + t.Fatalf("SnapshotMaxItems override = %d, want 7", got)
18 + }
19 +}
src/go/pkg/netipc/transport/posix/shm_edge_test.go new
+296
@@ -0,0 +1,296 @@
1 +//go:build linux
2 +
3 +package posix
4 +
5 +import (
6 + "fmt"
7 + "os"
8 + "testing"
9 +)
10 +
11 +// ---------------------------------------------------------------------------
12 +// ShmContext accessor methods
13 +// ---------------------------------------------------------------------------
14 +
15 +func TestShmContextRole(t *testing.T) {
16 + ensureShmRunDir(t)
17 + svc := "go_shm_role"
18 + cleanupShmFiles(t, svc)
19 + defer cleanupShmFiles(t, svc)
20 +
21 + ctx, err := ShmServerCreate(testShmRunDir, svc, 1, 1024, 1024)
22 + if err != nil {
23 + t.Fatalf("server create: %v", err)
24 + }
25 + defer ctx.ShmDestroy()
26 +
27 + if ctx.Role() != ShmRoleServer {
28 + t.Errorf("expected ShmRoleServer, got %d", ctx.Role())
29 + }
30 + if ctx.Fd() < 0 {
31 + t.Errorf("expected fd >= 0, got %d", ctx.Fd())
32 + }
33 +}
34 +
35 +func TestShmContextClientRole(t *testing.T) {
36 + ensureShmRunDir(t)
37 + svc := "go_shm_crole"
38 + cleanupShmFiles(t, svc)
39 + defer cleanupShmFiles(t, svc)
40 +
41 + srv, err := ShmServerCreate(testShmRunDir, svc, 10, 1024, 1024)
42 + if err != nil {
43 + t.Fatalf("server create: %v", err)
44 + }
45 + defer srv.ShmDestroy()
46 +
47 + client, err := ShmClientAttach(testShmRunDir, svc, 10)
48 + if err != nil {
49 + t.Fatalf("client attach: %v", err)
50 + }
51 + defer client.ShmClose()
52 +
53 + if client.Role() != ShmRoleClient {
54 + t.Errorf("expected ShmRoleClient, got %d", client.Role())
55 + }
56 + if client.Fd() < 0 {
57 + t.Errorf("expected fd >= 0, got %d", client.Fd())
58 + }
59 +}
60 +
61 +// ---------------------------------------------------------------------------
62 +// OwnerAlive
63 +// ---------------------------------------------------------------------------
64 +
65 +func TestShmOwnerAlive(t *testing.T) {
66 + ensureShmRunDir(t)
67 + svc := "go_shm_alive"
68 + cleanupShmFiles(t, svc)
69 + defer cleanupShmFiles(t, svc)
70 +
71 + ctx, err := ShmServerCreate(testShmRunDir, svc, 1, 1024, 1024)
72 + if err != nil {
73 + t.Fatalf("server create: %v", err)
74 + }
75 + defer ctx.ShmDestroy()
76 +
77 + // Owner is current process, should be alive
78 + if !ctx.OwnerAlive() {
79 + t.Error("expected OwnerAlive() = true for current process")
80 + }
81 +}
82 +
83 +// ---------------------------------------------------------------------------
84 +// validateShmServiceName edge cases
85 +// ---------------------------------------------------------------------------
86 +
87 +func TestValidateShmServiceNameEdgeCases(t *testing.T) {
88 + valid := []string{"a", "test-svc", "test.svc", "test_svc", "Test123"}
89 + for _, name := range valid {
90 + if err := validateShmServiceName(name); err != nil {
91 + t.Errorf("validateShmServiceName(%q) = %v, want nil", name, err)
92 + }
93 + }
94 +
95 + invalid := []string{"", ".", "..", "a/b", "a b", "a@b", "a#b"}
96 + for _, name := range invalid {
97 + if err := validateShmServiceName(name); err == nil {
98 + t.Errorf("validateShmServiceName(%q) = nil, want error", name)
99 + }
100 + }
101 +}
102 +
103 +// ---------------------------------------------------------------------------
104 +// buildShmPath edge cases
105 +// ---------------------------------------------------------------------------
106 +
107 +func TestBuildShmPathEdgeCases(t *testing.T) {
108 + // Valid
109 + path, err := buildShmPath("/tmp", "test", 1)
110 + if err != nil {
111 + t.Fatalf("buildShmPath: %v", err)
112 + }
113 + expected := fmt.Sprintf("/tmp/test-%016x.ipcshm", uint64(1))
114 + if path != expected {
115 + t.Fatalf("path = %q, want %q", path, expected)
116 + }
117 +
118 + // Invalid service name
119 + _, err = buildShmPath("/tmp", "", 1)
120 + if err == nil {
121 + t.Fatal("expected error for empty service name")
122 + }
123 +
124 + // Path too long
125 + longDir := "/tmp/" + string(make([]byte, 300))
126 + _, err = buildShmPath(longDir, "test", 1)
127 + if err != ErrShmPathTooLong {
128 + t.Fatalf("expected ErrShmPathTooLong, got %v", err)
129 + }
130 +}
131 +
132 +// ---------------------------------------------------------------------------
133 +// pidAlive edge cases
134 +// ---------------------------------------------------------------------------
135 +
136 +func TestPidAliveEdgeCases(t *testing.T) {
137 + // PID 0 or negative should return false
138 + if pidAlive(0) {
139 + t.Error("pidAlive(0) should be false")
140 + }
141 + if pidAlive(-1) {
142 + t.Error("pidAlive(-1) should be false")
143 + }
144 +
145 + // Current process should be alive
146 + if !pidAlive(os.Getpid()) {
147 + t.Error("current process should be alive")
148 + }
149 +}
150 +
151 +// ---------------------------------------------------------------------------
152 +// Atomic operation bounds checking
153 +// ---------------------------------------------------------------------------
154 +
155 +func TestAtomicLoadU64OutOfBounds(t *testing.T) {
156 + data := make([]byte, 4) // too small for u64
157 + _, err := atomicLoadU64(data, 0)
158 + if err != errShmOutOfBounds {
159 + t.Fatalf("expected errShmOutOfBounds, got %v", err)
160 + }
161 +
162 + _, err = atomicLoadU64(data, -1)
163 + if err != errShmOutOfBounds {
164 + t.Fatalf("expected errShmOutOfBounds for negative offset, got %v", err)
165 + }
166 +}
167 +
168 +func TestAtomicLoadU32OutOfBounds(t *testing.T) {
169 + data := make([]byte, 2) // too small for u32
170 + _, err := atomicLoadU32(data, 0)
171 + if err != errShmOutOfBounds {
172 + t.Fatalf("expected errShmOutOfBounds, got %v", err)
173 + }
174 +}
175 +
176 +func TestAtomicStoreU32OutOfBounds(t *testing.T) {
177 + data := make([]byte, 2)
178 + err := atomicStoreU32(data, 0, 42)
179 + if err != errShmOutOfBounds {
180 + t.Fatalf("expected errShmOutOfBounds, got %v", err)
181 + }
182 +}
183 +
184 +func TestAtomicAddU64OutOfBounds(t *testing.T) {
185 + data := make([]byte, 4)
186 + err := atomicAddU64(data, 0, 1)
187 + if err != errShmOutOfBounds {
188 + t.Fatalf("expected errShmOutOfBounds, got %v", err)
189 + }
190 +}
191 +
192 +func TestAtomicAddU32OutOfBounds(t *testing.T) {
193 + data := make([]byte, 2)
194 + err := atomicAddU32(data, 0, 1)
195 + if err != errShmOutOfBounds {
196 + t.Fatalf("expected errShmOutOfBounds, got %v", err)
197 + }
198 +}
199 +
200 +func TestFutexWakeCallOutOfBounds(t *testing.T) {
201 + data := make([]byte, 2)
202 + ret := futexWakeCall(data, 0, 1)
203 + if ret != -1 {
204 + t.Fatalf("expected -1 for out-of-bounds, got %d", ret)
205 + }
206 +}
207 +
208 +func TestFutexWaitCallOutOfBounds(t *testing.T) {
209 + data := make([]byte, 2)
210 + ret := futexWaitCall(data, 0, 0, nil)
211 + if ret != -1 {
212 + t.Fatalf("expected -1 for out-of-bounds, got %d", ret)
213 + }
214 +}
215 +
216 +// ---------------------------------------------------------------------------
217 +// SHM client attach: file does not exist
218 +// ---------------------------------------------------------------------------
219 +
220 +func TestShmClientAttachMissing(t *testing.T) {
221 + ensureShmRunDir(t)
222 + _, err := ShmClientAttach(testShmRunDir, "nonexistent", 99999)
223 + if err == nil {
224 + t.Fatal("expected error for non-existent SHM file")
225 + }
226 +}
227 +
228 +// ---------------------------------------------------------------------------
229 +// SHM close/destroy double-call safety
230 +// ---------------------------------------------------------------------------
231 +
232 +func TestShmCloseDoubleCall(t *testing.T) {
233 + ensureShmRunDir(t)
234 + svc := "go_shm_dclose"
235 + cleanupShmFiles(t, svc)
236 + defer cleanupShmFiles(t, svc)
237 +
238 + ctx, err := ShmServerCreate(testShmRunDir, svc, 1, 1024, 1024)
239 + if err != nil {
240 + t.Fatalf("server create: %v", err)
241 + }
242 +
243 + // Double close should not panic
244 + ctx.ShmClose()
245 + ctx.ShmClose()
246 +}
247 +
248 +func TestShmDestroyDoubleCall(t *testing.T) {
249 + ensureShmRunDir(t)
250 + svc := "go_shm_ddestroy"
251 + cleanupShmFiles(t, svc)
252 + defer cleanupShmFiles(t, svc)
253 +
254 + ctx, err := ShmServerCreate(testShmRunDir, svc, 1, 1024, 1024)
255 + if err != nil {
256 + t.Fatalf("server create: %v", err)
257 + }
258 +
259 + // Double destroy should not panic
260 + ctx.ShmDestroy()
261 + ctx.ShmDestroy()
262 +}
263 +
264 +// ---------------------------------------------------------------------------
265 +// SHM send: message exceeds capacity
266 +// ---------------------------------------------------------------------------
267 +
268 +func TestShmSendTooLarge(t *testing.T) {
269 + ensureShmRunDir(t)
270 + svc := "go_shm_toolarge"
271 + cleanupShmFiles(t, svc)
272 + defer cleanupShmFiles(t, svc)
273 +
274 + ctx, err := ShmServerCreate(testShmRunDir, svc, 1, 64, 64)
275 + if err != nil {
276 + t.Fatalf("server create: %v", err)
277 + }
278 + defer ctx.ShmDestroy()
279 +
280 + // Try to send a message larger than the response capacity
281 + bigMsg := make([]byte, 200) // > 64 bytes
282 + err = ctx.ShmSend(bigMsg)
283 + if err == nil {
284 + t.Fatal("expected error for oversized message")
285 + }
286 +}
287 +
288 +// ---------------------------------------------------------------------------
289 +// SHM cleanup stale: no files to clean
290 +// ---------------------------------------------------------------------------
291 +
292 +func TestShmCleanupStaleNoFiles(t *testing.T) {
293 + ensureShmRunDir(t)
294 + // Should not panic even with no matching files
295 + ShmCleanupStale(testShmRunDir, "nonexistent_svc_xyz")
296 +}
src/go/pkg/netipc/transport/posix/shm_linux.go new
+826
@@ -0,0 +1,826 @@
1 +//go:build linux
2 +
3 +// SHM transport for Linux — shared memory data plane with spin+futex
4 +// synchronization. Wire-compatible with the C and Rust implementations.
5 +//
6 +// Pure Go — no cgo. Works with CGO_ENABLED=0.
7 +
8 +package posix
9 +
10 +import (
11 + "encoding/binary"
12 + "errors"
13 + "fmt"
14 + "os"
15 + "path/filepath"
16 + "sync/atomic"
17 + "syscall"
18 + "time"
19 + "unsafe"
20 +)
21 +
22 +// ---------------------------------------------------------------------------
23 +// Constants
24 +// ---------------------------------------------------------------------------
25 +
26 +const (
27 + shmRegionMagic uint32 = 0x4e53484d // "NSHM"
28 + shmRegionVersion uint16 = 3
29 + shmRegionAlignment uint32 = 64
30 + shmHeaderLen uint16 = 64
31 + shmDefaultSpin uint32 = 128
32 +
33 + // Byte offsets of all fields in the 64-byte region header.
34 + shmHeaderMagicOff = 0
35 + shmHeaderVersionOff = 4
36 + shmHeaderHeaderLenOff = 6
37 + shmHeaderOwnerPidOff = 8
38 + shmHeaderOwnerGenOff = 12
39 + shmHeaderReqOffOff = 16
40 + shmHeaderReqCapOff = 20
41 + shmHeaderRespOffOff = 24
42 + shmHeaderRespCapOff = 28
43 + shmHeaderReqSeqOff = 32
44 + shmHeaderRespSeqOff = 40
45 + shmHeaderReqLenOff = 48
46 + shmHeaderRespLenOff = 52
47 + shmHeaderReqSignalOff = 56
48 + shmHeaderRespSignalOff = 60
49 +
50 + // futex operations
51 + futexWait = 0
52 + futexWake = 1
53 +
54 + shmMaxPath = 256
55 +)
56 +
57 +// ---------------------------------------------------------------------------
58 +// Errors
59 +// ---------------------------------------------------------------------------
60 +
61 +var (
62 + ErrShmPathTooLong = errors.New("SHM path exceeds limit")
63 + ErrShmOpen = errors.New("SHM open failed")
64 + ErrShmTruncate = errors.New("SHM ftruncate failed")
65 + ErrShmMmap = errors.New("SHM mmap failed")
66 + ErrShmBadMagic = errors.New("SHM header magic mismatch")
67 + ErrShmBadVersion = errors.New("SHM header version mismatch")
68 + ErrShmBadHeader = errors.New("SHM header_len mismatch")
69 + ErrShmBadSize = errors.New("SHM file too small for declared areas")
70 + ErrShmAddrInUse = errors.New("SHM region owned by live server")
71 + ErrShmNotReady = errors.New("SHM server not ready")
72 + ErrShmMsgTooLarge = errors.New("message exceeds SHM area capacity")
73 + ErrShmTimeout = errors.New("SHM futex wait timed out")
74 + ErrShmBadParam = errors.New("invalid SHM argument")
75 + ErrShmPeerDead = errors.New("SHM owner process has exited")
76 +)
77 +
78 +// ---------------------------------------------------------------------------
79 +// Role
80 +// ---------------------------------------------------------------------------
81 +
82 +// ShmRole distinguishes server vs client SHM contexts.
83 +type ShmRole int
84 +
85 +const (
86 + ShmRoleServer ShmRole = 1
87 + ShmRoleClient ShmRole = 2
88 +)
89 +
90 +// ---------------------------------------------------------------------------
91 +// SHM context
92 +// ---------------------------------------------------------------------------
93 +
94 +// ShmContext is a handle to a shared memory region.
95 +type ShmContext struct {
96 + role ShmRole
97 + fd int
98 + data []byte // mmap'd region (via syscall.Mmap)
99 +
100 + requestOffset uint32
101 + requestCapacity uint32
102 + responseOffset uint32
103 + responseCapacity uint32
104 +
105 + localReqSeq uint64
106 + localRespSeq uint64
107 +
108 + SpinTries uint32
109 + ownerGeneration uint32 // cached for PID reuse detection
110 + path string
111 +}
112 +
113 +// Role returns the context role.
114 +func (c *ShmContext) Role() ShmRole { return c.role }
115 +
116 +// Fd returns the file descriptor.
117 +func (c *ShmContext) Fd() int { return c.fd }
118 +
119 +// OwnerAlive checks if the region's owner process is still alive.
120 +func (c *ShmContext) OwnerAlive() bool {
121 + if len(c.data) < int(shmHeaderLen) {
122 + return false
123 + }
124 + pid := int32(binary.NativeEndian.Uint32(c.data[shmHeaderOwnerPidOff : shmHeaderOwnerPidOff+4]))
125 + if !pidAlive(int(pid)) {
126 + return false
127 + }
128 + // Verify generation matches to detect PID reuse.
129 + // Skip check if cached generation is 0 (legacy region).
130 + if c.ownerGeneration != 0 {
131 + curGen := binary.NativeEndian.Uint32(c.data[shmHeaderOwnerGenOff : shmHeaderOwnerGenOff+4])
132 + if curGen != c.ownerGeneration {
133 + return false
134 + }
135 + }
136 + return true
137 +}
138 +
139 +// ---------------------------------------------------------------------------
140 +// Server API
141 +// ---------------------------------------------------------------------------
142 +
143 +// ShmServerCreate creates a SHM region at {runDir}/{serviceName}-{sessionID}.ipcshm.
144 +func ShmServerCreate(runDir, serviceName string, sessionID uint64, reqCapacity, respCapacity uint32) (*ShmContext, error) {
145 + path, err := buildShmPath(runDir, serviceName, sessionID)
146 + if err != nil {
147 + return nil, err
148 + }
149 +
150 + // Round capacities
151 + reqCap := shmAlign64(reqCapacity)
152 + respCap := shmAlign64(respCapacity)
153 +
154 + reqOff := shmAlign64(uint32(shmHeaderLen))
155 + respOff := shmAlign64(reqOff + reqCap)
156 + regionSize := int(respOff + respCap)
157 +
158 + // Try O_EXCL create first (fast path, no stale check needed).
159 + f, err := os.OpenFile(path, os.O_RDWR|os.O_CREATE|os.O_EXCL, 0600)
160 + if err != nil && os.IsExist(err) {
161 + // File exists — do stale recovery and retry.
162 + stale := checkShmStale(path)
163 + if stale == shmStaleLive {
164 + return nil, fmt.Errorf("%w: live server owns SHM region", ErrShmOpen)
165 + }
166 + // Stale file was unlinked, retry create
167 + f, err = os.OpenFile(path, os.O_RDWR|os.O_CREATE|os.O_EXCL, 0600)
168 + }
169 + if err != nil {
170 + return nil, fmt.Errorf("%w: %v", ErrShmOpen, err)
171 + }
172 + fd := int(f.Fd())
173 +
174 + if err := syscall.Ftruncate(fd, int64(regionSize)); err != nil {
175 + f.Close()
176 + os.Remove(path)
177 + return nil, fmt.Errorf("%w: %v", ErrShmTruncate, err)
178 + }
179 +
180 + data, err := syscall.Mmap(fd, 0, regionSize,
181 + syscall.PROT_READ|syscall.PROT_WRITE, syscall.MAP_SHARED)
182 + if err != nil {
183 + f.Close()
184 + os.Remove(path)
185 + return nil, fmt.Errorf("%w: %v", ErrShmMmap, err)
186 + }
187 +
188 + // Zero region (Mmap may return zero-filled from kernel, but be explicit)
189 + for i := range data {
190 + data[i] = 0
191 + }
192 +
193 + // Use a time-based generation to detect PID reuse across restarts.
194 + now := time.Now()
195 + generation := uint32(now.Unix()) ^ uint32(now.Nanosecond()>>10)
196 +
197 + // Write header fields (host byte order)
198 + binary.NativeEndian.PutUint32(data[shmHeaderMagicOff:shmHeaderMagicOff+4], shmRegionMagic)
199 + binary.NativeEndian.PutUint16(data[shmHeaderVersionOff:shmHeaderVersionOff+2], shmRegionVersion)
200 + binary.NativeEndian.PutUint16(data[shmHeaderHeaderLenOff:shmHeaderHeaderLenOff+2], uint16(shmHeaderLen))
201 + binary.NativeEndian.PutUint32(data[shmHeaderOwnerPidOff:shmHeaderOwnerPidOff+4], uint32(int32(os.Getpid())))
202 + binary.NativeEndian.PutUint32(data[shmHeaderOwnerGenOff:shmHeaderOwnerGenOff+4], generation)
203 + binary.NativeEndian.PutUint32(data[shmHeaderReqOffOff:shmHeaderReqOffOff+4], reqOff)
204 + binary.NativeEndian.PutUint32(data[shmHeaderReqCapOff:shmHeaderReqCapOff+4], reqCap)
205 + binary.NativeEndian.PutUint32(data[shmHeaderRespOffOff:shmHeaderRespOffOff+4], respOff)
206 + binary.NativeEndian.PutUint32(data[shmHeaderRespCapOff:shmHeaderRespCapOff+4], respCap)
207 +
208 + // Release fence: ensure header writes are visible before clients
209 + atomic.StoreUint32((*uint32)(unsafe.Pointer(&data[shmHeaderReqSignalOff])), 0)
210 +
211 + // Close the os.File but keep the fd open (Mmap holds a reference).
212 + // Actually, we need to keep the fd ourselves for the context.
213 + // os.File.Close() would close the fd, so we dup first or just
214 + // not use os.File.Close(). We already have the fd from f.Fd().
215 + // Trick: prevent Go's finalizer from closing fd.
216 + // The safe way: dup the fd, then close the file.
217 + newFd, err := syscall.Dup(fd)
218 + if err != nil {
219 + syscall.Munmap(data)
220 + f.Close()
221 + os.Remove(path)
222 + return nil, fmt.Errorf("%w: dup: %v", ErrShmOpen, err)
223 + }
224 + f.Close() // closes original fd
225 +
226 + return &ShmContext{
227 + role: ShmRoleServer,
228 + fd: newFd,
229 + data: data,
230 + requestOffset: reqOff,
231 + requestCapacity: reqCap,
232 + responseOffset: respOff,
233 + responseCapacity: respCap,
234 + localReqSeq: 0,
235 + localRespSeq: 0,
236 + SpinTries: shmDefaultSpin,
237 + ownerGeneration: generation,
238 + path: path,
239 + }, nil
240 +}
241 +
242 +// ShmDestroy destroys a server SHM region (munmap, close, unlink).
243 +func (c *ShmContext) ShmDestroy() {
244 + if c.data != nil {
245 + syscall.Munmap(c.data)
246 + c.data = nil
247 + }
248 + if c.fd >= 0 {
249 + syscall.Close(c.fd)
250 + c.fd = -1
251 + }
252 + if c.path != "" {
253 + os.Remove(c.path)
254 + c.path = ""
255 + }
256 +}
257 +
258 +// ---------------------------------------------------------------------------
259 +// Client API
260 +// ---------------------------------------------------------------------------
261 +
262 +// ShmClientAttach attaches to an existing SHM region.
263 +func ShmClientAttach(runDir, serviceName string, sessionID uint64) (*ShmContext, error) {
264 + path, err := buildShmPath(runDir, serviceName, sessionID)
265 + if err != nil {
266 + return nil, err
267 + }
268 +
269 + f, err := os.OpenFile(path, os.O_RDWR, 0)
270 + if err != nil {
271 + return nil, fmt.Errorf("%w: %v", ErrShmOpen, err)
272 + }
273 + fd := int(f.Fd())
274 +
275 + info, err := f.Stat()
276 + if err != nil {
277 + f.Close()
278 + return nil, fmt.Errorf("%w: stat: %v", ErrShmOpen, err)
279 + }
280 +
281 + fileSize := int(info.Size())
282 + if fileSize < int(shmHeaderLen) {
283 + f.Close()
284 + return nil, ErrShmNotReady
285 + }
286 +
287 + data, err := syscall.Mmap(fd, 0, fileSize,
288 + syscall.PROT_READ|syscall.PROT_WRITE, syscall.MAP_SHARED)
289 + if err != nil {
290 + f.Close()
291 + return nil, fmt.Errorf("%w: %v", ErrShmMmap, err)
292 + }
293 +
294 + // Acquire fence
295 + atomic.LoadUint32((*uint32)(unsafe.Pointer(&data[shmHeaderReqSignalOff])))
296 +
297 + // Validate header
298 + magic := binary.NativeEndian.Uint32(data[shmHeaderMagicOff : shmHeaderMagicOff+4])
299 + if magic != shmRegionMagic {
300 + syscall.Munmap(data)
301 + f.Close()
302 + return nil, ErrShmBadMagic
303 + }
304 +
305 + version := binary.NativeEndian.Uint16(data[shmHeaderVersionOff : shmHeaderVersionOff+2])
306 + if version != shmRegionVersion {
307 + syscall.Munmap(data)
308 + f.Close()
309 + return nil, ErrShmBadVersion
310 + }
311 +
312 + hdrLen := binary.NativeEndian.Uint16(data[shmHeaderHeaderLenOff : shmHeaderHeaderLenOff+2])
313 + if hdrLen != uint16(shmHeaderLen) {
314 + syscall.Munmap(data)
315 + f.Close()
316 + return nil, ErrShmBadHeader
317 + }
318 +
319 + reqOff := binary.NativeEndian.Uint32(data[shmHeaderReqOffOff : shmHeaderReqOffOff+4])
320 + reqCap := binary.NativeEndian.Uint32(data[shmHeaderReqCapOff : shmHeaderReqCapOff+4])
321 + respOff := binary.NativeEndian.Uint32(data[shmHeaderRespOffOff : shmHeaderRespOffOff+4])
322 + respCap := binary.NativeEndian.Uint32(data[shmHeaderRespCapOff : shmHeaderRespCapOff+4])
323 +
324 + headerEnd := shmAlign64(uint32(shmHeaderLen))
325 + if reqOff < headerEnd || reqCap == 0 || respOff < headerEnd || respCap == 0 {
326 + syscall.Munmap(data)
327 + f.Close()
328 + return nil, ErrShmNotReady
329 + }
330 +
331 + if reqOff%shmRegionAlignment != 0 ||
332 + reqCap%shmRegionAlignment != 0 ||
333 + respOff%shmRegionAlignment != 0 ||
334 + respCap%shmRegionAlignment != 0 ||
335 + respOff < shmAlign64(reqOff+reqCap) {
336 + syscall.Munmap(data)
337 + f.Close()
338 + return nil, ErrShmBadSize
339 + }
340 +
341 + // Validate region size
342 + reqEnd := int(reqOff) + int(reqCap)
343 + respEnd := int(respOff) + int(respCap)
344 + needed := reqEnd
345 + if respEnd > needed {
346 + needed = respEnd
347 + }
348 + if fileSize < needed {
349 + syscall.Munmap(data)
350 + f.Close()
351 + return nil, ErrShmBadSize
352 + }
353 +
354 + // Read current sequence numbers
355 + curReqSeq, err := atomicLoadU64(data, shmHeaderReqSeqOff)
356 + if err != nil {
357 + syscall.Munmap(data)
358 + f.Close()
359 + return nil, fmt.Errorf("%w: load req_seq: %v", ErrShmBadParam, err)
360 + }
361 + curRespSeq, err := atomicLoadU64(data, shmHeaderRespSeqOff)
362 + if err != nil {
363 + syscall.Munmap(data)
364 + f.Close()
365 + return nil, fmt.Errorf("%w: load resp_seq: %v", ErrShmBadParam, err)
366 + }
367 + ownerGen := binary.NativeEndian.Uint32(data[shmHeaderOwnerGenOff : shmHeaderOwnerGenOff+4])
368 +
369 + // Dup fd and close file
370 + newFd, err := syscall.Dup(fd)
371 + if err != nil {
372 + syscall.Munmap(data)
373 + f.Close()
374 + return nil, fmt.Errorf("%w: dup: %v", ErrShmOpen, err)
375 + }
376 + f.Close()
377 +
378 + return &ShmContext{
379 + role: ShmRoleClient,
380 + fd: newFd,
381 + data: data,
382 + requestOffset: reqOff,
383 + requestCapacity: reqCap,
384 + responseOffset: respOff,
385 + responseCapacity: respCap,
386 + localReqSeq: curReqSeq,
387 + localRespSeq: curRespSeq,
388 + SpinTries: shmDefaultSpin,
389 + ownerGeneration: ownerGen,
390 + path: path,
391 + }, nil
392 +}
393 +
394 +// ShmClose closes a client SHM context (no unlink).
395 +func (c *ShmContext) ShmClose() {
396 + if c.data != nil {
397 + syscall.Munmap(c.data)
398 + c.data = nil
399 + }
400 + if c.fd >= 0 {
401 + syscall.Close(c.fd)
402 + c.fd = -1
403 + }
404 +}
405 +
406 +// ---------------------------------------------------------------------------
407 +// Data plane
408 +// ---------------------------------------------------------------------------
409 +
410 +// ShmSend publishes a message. The message must include the 32-byte
411 +// outer header + payload, exactly as sent over UDS.
412 +func (c *ShmContext) ShmSend(msg []byte) error {
413 + if c.data == nil || len(msg) == 0 {
414 + return fmt.Errorf("%w: null context or empty message", ErrShmBadParam)
415 + }
416 +
417 + var areaOff, areaCap uint32
418 + var seqOff, lenOff, sigOff int
419 +
420 + if c.role == ShmRoleClient {
421 + areaOff = c.requestOffset
422 + areaCap = c.requestCapacity
423 + seqOff = shmHeaderReqSeqOff
424 + lenOff = shmHeaderReqLenOff
425 + sigOff = shmHeaderReqSignalOff
426 + } else {
427 + areaOff = c.responseOffset
428 + areaCap = c.responseCapacity
429 + seqOff = shmHeaderRespSeqOff
430 + lenOff = shmHeaderRespLenOff
431 + sigOff = shmHeaderRespSignalOff
432 + }
433 +
434 + if uint32(len(msg)) > areaCap {
435 + return ErrShmMsgTooLarge
436 + }
437 +
438 + // 1. Write message data
439 + copy(c.data[areaOff:], msg)
440 +
441 + // 2. Store message length (release)
442 + if err := atomicStoreU32(c.data, lenOff, uint32(len(msg))); err != nil {
443 + return fmt.Errorf("%w: store msg_len: %v", ErrShmBadParam, err)
444 + }
445 +
446 + // 3. Increment sequence number (release)
447 + if err := atomicAddU64(c.data, seqOff, 1); err != nil {
448 + return fmt.Errorf("%w: add seq: %v", ErrShmBadParam, err)
449 + }
450 +
451 + // 4. Wake peer via futex
452 + if err := atomicAddU32(c.data, sigOff, 1); err != nil {
453 + return fmt.Errorf("%w: add signal: %v", ErrShmBadParam, err)
454 + }
455 + futexWakeCall(c.data, sigOff, 1)
456 +
457 + // Track locally
458 + if c.role == ShmRoleClient {
459 + c.localReqSeq++
460 + } else {
461 + c.localRespSeq++
462 + }
463 +
464 + return nil
465 +}
466 +
467 +// ShmReceive receives a message into the caller-provided buffer.
468 +// On success, returns the number of bytes written to buf.
469 +// Returns ErrShmMsgTooLarge if the message exceeds len(buf).
470 +func (c *ShmContext) ShmReceive(buf []byte, timeoutMs uint32) (int, error) {
471 + if c.data == nil {
472 + return 0, fmt.Errorf("%w: null context", ErrShmBadParam)
473 + }
474 + if len(buf) == 0 {
475 + return 0, fmt.Errorf("%w: empty buffer", ErrShmBadParam)
476 + }
477 +
478 + var areaOff, areaCap uint32
479 + var seqOff, lenOff, sigOff int
480 + var expectedSeq uint64
481 +
482 + if c.role == ShmRoleServer {
483 + areaOff = c.requestOffset
484 + areaCap = c.requestCapacity
485 + seqOff = shmHeaderReqSeqOff
486 + lenOff = shmHeaderReqLenOff
487 + sigOff = shmHeaderReqSignalOff
488 + expectedSeq = c.localReqSeq + 1
489 + } else {
490 + areaOff = c.responseOffset
491 + areaCap = c.responseCapacity
492 + seqOff = shmHeaderRespSeqOff
493 + lenOff = shmHeaderRespLenOff
494 + sigOff = shmHeaderRespSignalOff
495 + expectedSeq = c.localRespSeq + 1
496 + }
497 +
498 + // Limit copy to the smaller of caller buffer and SHM area capacity
499 + maxCopy := len(buf)
500 + if int(areaCap) < maxCopy {
501 + maxCopy = int(areaCap)
502 + }
503 +
504 + // Phase 1: spin. Copy immediately on observing the advance.
505 + observed := false
506 + var mlen uint32
507 + for i := uint32(0); i < c.SpinTries; i++ {
508 + cur, err := atomicLoadU64(c.data, seqOff)
509 + if err != nil {
510 + return 0, fmt.Errorf("%w: load seq: %v", ErrShmBadParam, err)
511 + }
512 + if cur >= expectedSeq {
513 + mlen, err = atomicLoadU32(c.data, lenOff)
514 + if err != nil {
515 + return 0, fmt.Errorf("%w: load msg_len: %v", ErrShmBadParam, err)
516 + }
517 + if mlen > 0 && int(mlen) <= maxCopy {
518 + copy(buf[:mlen], c.data[areaOff:areaOff+mlen])
519 + }
520 + observed = true
521 + break
522 + }
523 + spinPause()
524 + }
525 +
526 + // Phase 2: futex wait with deadline-based retry loop.
527 + //
528 + // Handles spurious wakeups (EAGAIN when signal word changed
529 + // between read and syscall, or EINTR from signal delivery).
530 + // Computes a wall-clock deadline so total wait never exceeds
531 + // timeoutMs regardless of retries.
532 + if !observed {
533 + var deadlineNs uint64
534 + if timeoutMs > 0 {
535 + var nowTs syscall.Timespec
536 + syscall.Syscall(syscall.SYS_CLOCK_GETTIME, 1 /* CLOCK_MONOTONIC */, uintptr(unsafe.Pointer(&nowTs)), 0)
537 + deadlineNs = uint64(nowTs.Sec)*1_000_000_000 + uint64(nowTs.Nsec) +
538 + uint64(timeoutMs)*1_000_000
539 + }
540 +
541 + for {
542 + sigVal, serr := atomicLoadU32(c.data, sigOff)
543 + if serr != nil {
544 + return 0, fmt.Errorf("%w: load signal: %v", ErrShmBadParam, serr)
545 + }
546 +
547 + cur, serr := atomicLoadU64(c.data, seqOff)
548 + if serr != nil {
549 + return 0, fmt.Errorf("%w: load seq: %v", ErrShmBadParam, serr)
550 + }
551 + if cur >= expectedSeq {
552 + break // response arrived
553 + }
554 +
555 + // Compute remaining timeout for this futex_wait call
556 + var ts *syscall.Timespec
557 + if deadlineNs > 0 {
558 + var nowTs syscall.Timespec
559 + syscall.Syscall(syscall.SYS_CLOCK_GETTIME, 1 /* CLOCK_MONOTONIC */, uintptr(unsafe.Pointer(&nowTs)), 0)
560 + nowVal := uint64(nowTs.Sec)*1_000_000_000 + uint64(nowTs.Nsec)
561 + if nowVal >= deadlineNs {
562 + return 0, ErrShmTimeout
563 + }
564 + remain := deadlineNs - nowVal
565 + ts = &syscall.Timespec{
566 + Sec: int64(remain / 1_000_000_000),
567 + Nsec: int64(remain % 1_000_000_000),
568 + }
569 + }
570 +
571 + ret := futexWaitCall(c.data, sigOff, sigVal, ts)
572 + if ret < 0 {
573 + errno := syscall.Errno(-ret)
574 + if errno == syscall.ETIMEDOUT {
575 + return 0, ErrShmTimeout
576 + }
577 + }
578 +
579 + // EAGAIN (value changed) or EINTR (signal): re-check seq
580 + }
581 +
582 + // Copy immediately after observing the sequence advance
583 + var lerr error
584 + mlen, lerr = atomicLoadU32(c.data, lenOff)
585 + if lerr != nil {
586 + return 0, fmt.Errorf("%w: load msg_len: %v", ErrShmBadParam, lerr)
587 + }
588 + if mlen > 0 && int(mlen) <= maxCopy {
589 + copy(buf[:mlen], c.data[areaOff:areaOff+mlen])
590 + }
591 + }
592 +
593 + // Advance local tracking (message is consumed from SHM perspective)
594 + if c.role == ShmRoleServer {
595 + c.localReqSeq = expectedSeq
596 + } else {
597 + c.localRespSeq = expectedSeq
598 + }
599 +
600 + // Message larger than safe copy limit
601 + if int(mlen) > maxCopy {
602 + return int(mlen), ErrShmMsgTooLarge
603 + }
604 +
605 + return int(mlen), nil
606 +}
607 +
608 +// ---------------------------------------------------------------------------
609 +// Internal helpers
610 +// ---------------------------------------------------------------------------
611 +
612 +func shmAlign64(v uint32) uint32 {
613 + return (v + (shmRegionAlignment - 1)) & ^(shmRegionAlignment - 1)
614 +}
615 +
616 +// validateShmServiceName checks that name contains only [a-zA-Z0-9._-],
617 +// is non-empty, and is not "." or "..".
618 +func validateShmServiceName(name string) error {
619 + if name == "" {
620 + return fmt.Errorf("%w: empty service name", ErrShmBadParam)
621 + }
622 + if name == "." || name == ".." {
623 + return fmt.Errorf("%w: service name cannot be '.' or '..'", ErrShmBadParam)
624 + }
625 + for i := 0; i < len(name); i++ {
626 + c := name[i]
627 + if (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') ||
628 + (c >= '0' && c <= '9') || c == '.' || c == '_' || c == '-' {
629 + continue
630 + }
631 + return fmt.Errorf("%w: service name contains invalid character: %q", ErrShmBadParam, c)
632 + }
633 + return nil
634 +}
635 +
636 +func buildShmPath(runDir, serviceName string, sessionID uint64) (string, error) {
637 + if err := validateShmServiceName(serviceName); err != nil {
638 + return "", err
639 + }
640 + path := filepath.Join(runDir, fmt.Sprintf("%s-%016x.ipcshm", serviceName, sessionID))
641 + if len(path) >= shmMaxPath {
642 + return "", ErrShmPathTooLong
643 + }
644 + return path, nil
645 +}
646 +
647 +func pidAlive(pid int) bool {
648 + if pid <= 0 {
649 + return false
650 + }
651 + err := syscall.Kill(pid, 0)
652 + return err == nil || err == syscall.EPERM
653 +}
654 +
655 +// Atomic operations on the mmap'd region with bounds checking.
656 +
657 +var errShmOutOfBounds = errors.New("SHM atomic: offset out of bounds")
658 +
659 +func atomicLoadU64(data []byte, off int) (uint64, error) {
660 + if off < 0 || off+8 > len(data) {
661 + return 0, errShmOutOfBounds
662 + }
663 + ptr := (*uint64)(unsafe.Pointer(&data[off]))
664 + return atomic.LoadUint64(ptr), nil
665 +}
666 +
667 +func atomicLoadU32(data []byte, off int) (uint32, error) {
668 + if off < 0 || off+4 > len(data) {
669 + return 0, errShmOutOfBounds
670 + }
671 + ptr := (*uint32)(unsafe.Pointer(&data[off]))
672 + return atomic.LoadUint32(ptr), nil
673 +}
674 +
675 +func atomicStoreU32(data []byte, off int, val uint32) error {
676 + if off < 0 || off+4 > len(data) {
677 + return errShmOutOfBounds
678 + }
679 + ptr := (*uint32)(unsafe.Pointer(&data[off]))
680 + atomic.StoreUint32(ptr, val)
681 + return nil
682 +}
683 +
684 +func atomicAddU64(data []byte, off int, val uint64) error {
685 + if off < 0 || off+8 > len(data) {
686 + return errShmOutOfBounds
687 + }
688 + ptr := (*uint64)(unsafe.Pointer(&data[off]))
689 + atomic.AddUint64(ptr, val)
690 + return nil
691 +}
692 +
693 +func atomicAddU32(data []byte, off int, val uint32) error {
694 + if off < 0 || off+4 > len(data) {
695 + return errShmOutOfBounds
696 + }
697 + ptr := (*uint32)(unsafe.Pointer(&data[off]))
698 + atomic.AddUint32(ptr, val)
699 + return nil
700 +}
701 +
702 +func futexWakeCall(data []byte, off int, count int) int {
703 + if off < 0 || off+4 > len(data) {
704 + return -1
705 + }
706 + addr := unsafe.Pointer(&data[off])
707 + r1, _, _ := syscall.Syscall6(
708 + syscall.SYS_FUTEX,
709 + uintptr(addr),
710 + uintptr(futexWake),
711 + uintptr(count),
712 + 0, 0, 0,
713 + )
714 + return int(r1)
715 +}
716 +
717 +func futexWaitCall(data []byte, off int, expected uint32, ts *syscall.Timespec) int {
718 + if off < 0 || off+4 > len(data) {
719 + return -1
720 + }
721 + addr := unsafe.Pointer(&data[off])
722 + var tsPtr uintptr
723 + if ts != nil {
724 + tsPtr = uintptr(unsafe.Pointer(ts))
725 + }
726 + r1, _, errno := syscall.Syscall6(
727 + syscall.SYS_FUTEX,
728 + uintptr(addr),
729 + uintptr(futexWait),
730 + uintptr(expected),
731 + tsPtr,
732 + 0, 0,
733 + )
734 + if errno != 0 {
735 + return -int(errno)
736 + }
737 + return int(r1)
738 +}
739 +
740 +// ---------------------------------------------------------------------------
741 +// Stale cleanup
742 +// ---------------------------------------------------------------------------
743 +
744 +// ShmCleanupStale scans runDir for SHM files matching {serviceName}-*.ipcshm,
745 +// checks if the owner PID is alive for each, and unlinks stale ones.
746 +func ShmCleanupStale(runDir, serviceName string) {
747 + entries, err := os.ReadDir(runDir)
748 + if err != nil {
749 + return
750 + }
751 + prefix := serviceName + "-"
752 + suffix := ".ipcshm"
753 + for _, e := range entries {
754 + name := e.Name()
755 + if !e.Type().IsRegular() {
756 + continue
757 + }
758 + if len(name) < len(prefix)+len(suffix) {
759 + continue
760 + }
761 + if name[:len(prefix)] != prefix || name[len(name)-len(suffix):] != suffix {
762 + continue
763 + }
764 + path := filepath.Join(runDir, name)
765 + result := checkShmStale(path)
766 + _ = result // checkShmStale already unlinks stale/invalid files
767 + }
768 +}
769 +
770 +// ---------------------------------------------------------------------------
771 +// Stale region recovery
772 +// ---------------------------------------------------------------------------
773 +
774 +type shmStaleResult int
775 +
776 +const (
777 + shmStaleNotExist shmStaleResult = iota
778 + shmStaleRecovered
779 + shmStaleLive
780 + shmStaleInvalid
781 +)
782 +
783 +func checkShmStale(path string) shmStaleResult {
784 + info, err := os.Stat(path)
785 + if err != nil {
786 + return shmStaleNotExist
787 + }
788 +
789 + if info.Size() < int64(shmHeaderLen) {
790 + os.Remove(path)
791 + return shmStaleInvalid
792 + }
793 +
794 + f, err := os.Open(path)
795 + if err != nil {
796 + os.Remove(path)
797 + return shmStaleInvalid
798 + }
799 +
800 + data, err := syscall.Mmap(int(f.Fd()), 0, int(shmHeaderLen),
801 + syscall.PROT_READ, syscall.MAP_SHARED)
802 + f.Close()
803 + if err != nil {
804 + os.Remove(path)
805 + return shmStaleInvalid
806 + }
807 +
808 + magic := binary.NativeEndian.Uint32(data[shmHeaderMagicOff : shmHeaderMagicOff+4])
809 + if magic != shmRegionMagic {
810 + syscall.Munmap(data)
811 + os.Remove(path)
812 + return shmStaleInvalid
813 + }
814 +
815 + ownerPid := int(int32(binary.NativeEndian.Uint32(data[shmHeaderOwnerPidOff : shmHeaderOwnerPidOff+4])))
816 + ownerGen := binary.NativeEndian.Uint32(data[shmHeaderOwnerGenOff : shmHeaderOwnerGenOff+4])
817 + syscall.Munmap(data)
818 +
819 + if pidAlive(ownerPid) && ownerGen != 0 {
820 + return shmStaleLive
821 + }
822 +
823 + // Dead owner or zero generation (PID reuse / legacy) — stale
824 + os.Remove(path)
825 + return shmStaleRecovered
826 +}
src/go/pkg/netipc/transport/posix/shm_linux_test.go new
+626
@@ -0,0 +1,626 @@
1 +//go:build linux
2 +
3 +package posix
4 +
5 +import (
6 + "bytes"
7 + "encoding/binary"
8 + "errors"
9 + "fmt"
10 + "os"
11 + "sync"
12 + "testing"
13 + "time"
14 +
15 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
16 +)
17 +
18 +const testShmRunDir = "/tmp/nipc_shm_go_test"
19 +
20 +func ensureShmRunDir(t *testing.T) {
21 + t.Helper()
22 + if err := os.MkdirAll(testShmRunDir, 0700); err != nil {
23 + t.Fatalf("cannot create SHM run dir: %v", err)
24 + }
25 +}
26 +
27 +func cleanupShmFiles(t *testing.T, service string) {
28 + t.Helper()
29 + // Clean up SHM files with session ID pattern
30 + entries, _ := os.ReadDir(testShmRunDir)
31 + for _, e := range entries {
32 + name := e.Name()
33 + if len(name) > len(service)+1 && name[:len(service)+1] == service+"-" {
34 + os.Remove(fmt.Sprintf("%s/%s", testShmRunDir, name))
35 + }
36 + }
37 +}
38 +
39 +func uniqueShmService(t *testing.T, prefix string) string {
40 + t.Helper()
41 + return fmt.Sprintf("%s_%s", prefix, uniqueService(t))
42 +}
43 +
44 +func waitShmClientAttach(t *testing.T, runDir, service string, sessionID uint64) *ShmContext {
45 + t.Helper()
46 +
47 + deadline := time.Now().Add(5 * time.Second)
48 + var lastErr error
49 + for time.Now().Before(deadline) {
50 + ctx, err := ShmClientAttach(runDir, service, sessionID)
51 + if err == nil {
52 + return ctx
53 + }
54 + lastErr = err
55 +
56 + if !errors.Is(err, ErrShmOpen) &&
57 + !errors.Is(err, ErrShmNotReady) &&
58 + !errors.Is(err, ErrShmBadMagic) &&
59 + !errors.Is(err, ErrShmBadVersion) &&
60 + !errors.Is(err, ErrShmBadHeader) &&
61 + !errors.Is(err, ErrShmBadSize) {
62 + t.Fatalf("attach should not fail with non-transient error while waiting for readiness: %v", err)
63 + }
64 +
65 + time.Sleep(10 * time.Millisecond)
66 + }
67 +
68 + t.Fatalf("timed out waiting for SHM region readiness: %v", lastErr)
69 + return nil
70 +}
71 +
72 +// buildShmMessage creates a complete wire message (32-byte header + payload).
73 +func buildShmMessage(kind, code uint16, messageID uint64, payload []byte) []byte {
74 + hdr := protocol.Header{
75 + Magic: protocol.MagicMsg,
76 + Version: protocol.Version,
77 + HeaderLen: protocol.HeaderLen,
78 + Kind: kind,
79 + Code: code,
80 + ItemCount: 1,
81 + MessageID: messageID,
82 + PayloadLen: uint32(len(payload)),
83 + }
84 + buf := make([]byte, protocol.HeaderSize+len(payload))
85 + hdr.Encode(buf[:protocol.HeaderSize])
86 + copy(buf[protocol.HeaderSize:], payload)
87 + return buf
88 +}
89 +
90 +func TestShmDirectRoundtrip(t *testing.T) {
91 + ensureShmRunDir(t)
92 + svc := uniqueShmService(t, "go_shm_rt")
93 + cleanupShmFiles(t, svc)
94 + defer cleanupShmFiles(t, svc)
95 +
96 + var wg sync.WaitGroup
97 + var serverErr error
98 +
99 + wg.Add(1)
100 + go func() {
101 + defer wg.Done()
102 + ctx, err := ShmServerCreate(testShmRunDir, svc, 1, 4096, 4096)
103 + if err != nil {
104 + serverErr = fmt.Errorf("server create: %w", err)
105 + return
106 + }
107 + defer ctx.ShmDestroy()
108 +
109 + buf := make([]byte, 65536)
110 + mlen, err := ctx.ShmReceive(buf, 5000)
111 + if err != nil {
112 + serverErr = fmt.Errorf("server receive: %w", err)
113 + return
114 + }
115 +
116 + if mlen < protocol.HeaderSize {
117 + serverErr = fmt.Errorf("message too short: %d", mlen)
118 + return
119 + }
120 +
121 + // Parse header, echo as response
122 + hdr, err := protocol.DecodeHeader(buf[:mlen])
123 + if err != nil {
124 + serverErr = fmt.Errorf("decode header: %w", err)
125 + return
126 + }
127 + payload := make([]byte, mlen-protocol.HeaderSize)
128 + copy(payload, buf[protocol.HeaderSize:mlen])
129 + resp := buildShmMessage(protocol.KindResponse, hdr.Code, hdr.MessageID, payload)
130 + if err := ctx.ShmSend(resp); err != nil {
131 + serverErr = fmt.Errorf("server send: %w", err)
132 + }
133 + }()
134 +
135 + client := waitShmClientAttach(t, testShmRunDir, svc, 1)
136 + defer client.ShmClose()
137 +
138 + payload := []byte{0xCA, 0xFE, 0xBA, 0xBE}
139 + msg := buildShmMessage(protocol.KindRequest, protocol.MethodIncrement, 42, payload)
140 + if err := client.ShmSend(msg); err != nil {
141 + t.Fatalf("client send: %v", err)
142 + }
143 +
144 + respBuf := make([]byte, 65536)
145 + rlen, err := client.ShmReceive(respBuf, 5000)
146 + if err != nil {
147 + t.Fatalf("client receive: %v", err)
148 + }
149 +
150 + if rlen != protocol.HeaderSize+len(payload) {
151 + t.Fatalf("response length: got %d, want %d", rlen, protocol.HeaderSize+len(payload))
152 + }
153 +
154 + rhdr, err := protocol.DecodeHeader(respBuf[:rlen])
155 + if err != nil {
156 + t.Fatalf("decode response header: %v", err)
157 + }
158 + if rhdr.Kind != protocol.KindResponse {
159 + t.Errorf("response kind: got %d, want %d", rhdr.Kind, protocol.KindResponse)
160 + }
161 + if rhdr.MessageID != 42 {
162 + t.Errorf("response message_id: got %d, want 42", rhdr.MessageID)
163 + }
164 + respPayload := respBuf[protocol.HeaderSize:rlen]
165 + if !bytes.Equal(respPayload, payload) {
166 + t.Errorf("response payload mismatch")
167 + }
168 +
169 + wg.Wait()
170 + if serverErr != nil {
171 + t.Fatalf("server error: %v", serverErr)
172 + }
173 +}
174 +
175 +func TestShmMultipleRoundtrips(t *testing.T) {
176 + ensureShmRunDir(t)
177 + svc := uniqueShmService(t, "go_shm_multi")
178 + cleanupShmFiles(t, svc)
179 + defer cleanupShmFiles(t, svc)
180 +
181 + var wg sync.WaitGroup
182 + var serverErr error
183 +
184 + wg.Add(1)
185 + go func() {
186 + defer wg.Done()
187 + ctx, err := ShmServerCreate(testShmRunDir, svc, 2, 4096, 4096)
188 + if err != nil {
189 + serverErr = fmt.Errorf("server create: %w", err)
190 + return
191 + }
192 + defer ctx.ShmDestroy()
193 +
194 + buf := make([]byte, 65536)
195 + for i := 0; i < 10; i++ {
196 + mlen, err := ctx.ShmReceive(buf, 5000)
197 + if err != nil {
198 + serverErr = fmt.Errorf("server receive %d: %w", i, err)
199 + return
200 + }
201 + hdr, err := protocol.DecodeHeader(buf[:mlen])
202 + if err != nil {
203 + serverErr = fmt.Errorf("decode header %d: %w", i, err)
204 + return
205 + }
206 + payload := make([]byte, mlen-protocol.HeaderSize)
207 + copy(payload, buf[protocol.HeaderSize:mlen])
208 + resp := buildShmMessage(protocol.KindResponse, hdr.Code, hdr.MessageID, payload)
209 + if err := ctx.ShmSend(resp); err != nil {
210 + serverErr = fmt.Errorf("server send %d: %w", i, err)
211 + return
212 + }
213 + }
214 + }()
215 +
216 + client := waitShmClientAttach(t, testShmRunDir, svc, 2)
217 + defer client.ShmClose()
218 +
219 + respBuf := make([]byte, 65536)
220 + for i := uint64(0); i < 10; i++ {
221 + payload := []byte{byte(i)}
222 + msg := buildShmMessage(protocol.KindRequest, 1, i+1, payload)
223 + if err := client.ShmSend(msg); err != nil {
224 + t.Fatalf("client send %d: %v", i, err)
225 + }
226 +
227 + rlen, err := client.ShmReceive(respBuf, 5000)
228 + if err != nil {
229 + t.Fatalf("client receive %d: %v", i, err)
230 + }
231 +
232 + rhdr, err := protocol.DecodeHeader(respBuf[:rlen])
233 + if err != nil {
234 + t.Fatalf("decode response %d: %v", i, err)
235 + }
236 + if rhdr.Kind != protocol.KindResponse {
237 + t.Errorf("round %d: kind=%d, want %d", i, rhdr.Kind, protocol.KindResponse)
238 + }
239 + if rhdr.MessageID != i+1 {
240 + t.Errorf("round %d: message_id=%d, want %d", i, rhdr.MessageID, i+1)
241 + }
242 + if respBuf[protocol.HeaderSize] != byte(i) {
243 + t.Errorf("round %d: payload byte=%d, want %d", i, respBuf[protocol.HeaderSize], i)
244 + }
245 + }
246 +
247 + wg.Wait()
248 + if serverErr != nil {
249 + t.Fatalf("server error: %v", serverErr)
250 + }
251 +}
252 +
253 +func TestShmStaleRecovery(t *testing.T) {
254 + ensureShmRunDir(t)
255 + svc := uniqueShmService(t, "go_shm_stale")
256 + cleanupShmFiles(t, svc)
257 + defer cleanupShmFiles(t, svc)
258 +
259 + // Create a region, then corrupt owner_pid to simulate dead process
260 + first, err := ShmServerCreate(testShmRunDir, svc, 3, 1024, 1024)
261 + if err != nil {
262 + t.Fatalf("first create: %v", err)
263 + }
264 +
265 + // Write a dead PID into the header
266 + binary.NativeEndian.PutUint32(first.data[8:12], 99999) // very unlikely alive
267 + first.ShmClose() // close without unlink
268 +
269 + // Clean up stale regions (as production server would)
270 + ShmCleanupStale(testShmRunDir, svc)
271 +
272 + // Should succeed after stale recovery
273 + second, err := ShmServerCreate(testShmRunDir, svc, 3, 2048, 2048)
274 + if err != nil {
275 + t.Fatalf("stale recovery create: %v", err)
276 + }
277 + if second.requestCapacity < 2048 {
278 + t.Errorf("new region capacity: %d, want >= 2048", second.requestCapacity)
279 + }
280 + second.ShmDestroy()
281 +}
282 +
283 +func TestShmClientAttachPartialHeaderNotReady(t *testing.T) {
284 + ensureShmRunDir(t)
285 + svc := uniqueShmService(t, "go_shm_partial")
286 + cleanupShmFiles(t, svc)
287 + defer cleanupShmFiles(t, svc)
288 +
289 + server, err := ShmServerCreate(testShmRunDir, svc, 5, 1024, 1024)
290 + if err != nil {
291 + t.Fatalf("server create: %v", err)
292 + }
293 + defer server.ShmDestroy()
294 +
295 + reqOff := binary.NativeEndian.Uint32(server.data[shmHeaderReqOffOff : shmHeaderReqOffOff+4])
296 + reqCap := binary.NativeEndian.Uint32(server.data[shmHeaderReqCapOff : shmHeaderReqCapOff+4])
297 + respOff := binary.NativeEndian.Uint32(server.data[shmHeaderRespOffOff : shmHeaderRespOffOff+4])
298 + respCap := binary.NativeEndian.Uint32(server.data[shmHeaderRespCapOff : shmHeaderRespCapOff+4])
299 +
300 + binary.NativeEndian.PutUint32(server.data[shmHeaderReqOffOff:shmHeaderReqOffOff+4], 0)
301 + binary.NativeEndian.PutUint32(server.data[shmHeaderReqCapOff:shmHeaderReqCapOff+4], 0)
302 + binary.NativeEndian.PutUint32(server.data[shmHeaderRespOffOff:shmHeaderRespOffOff+4], 0)
303 + binary.NativeEndian.PutUint32(server.data[shmHeaderRespCapOff:shmHeaderRespCapOff+4], 0)
304 +
305 + _, err = ShmClientAttach(testShmRunDir, svc, 5)
306 + if !errors.Is(err, ErrShmNotReady) {
307 + t.Fatalf("client attach error = %v, want %v", err, ErrShmNotReady)
308 + }
309 +
310 + binary.NativeEndian.PutUint32(server.data[shmHeaderReqOffOff:shmHeaderReqOffOff+4], reqOff)
311 + binary.NativeEndian.PutUint32(server.data[shmHeaderReqCapOff:shmHeaderReqCapOff+4], reqCap)
312 + binary.NativeEndian.PutUint32(server.data[shmHeaderRespOffOff:shmHeaderRespOffOff+4], respOff)
313 + binary.NativeEndian.PutUint32(server.data[shmHeaderRespCapOff:shmHeaderRespCapOff+4], respCap)
314 +}
315 +
316 +func TestShmLargeMessage(t *testing.T) {
317 + ensureShmRunDir(t)
318 + svc := uniqueShmService(t, "go_shm_large")
319 + cleanupShmFiles(t, svc)
320 + defer cleanupShmFiles(t, svc)
321 +
322 + var wg sync.WaitGroup
323 + var serverErr error
324 +
325 + wg.Add(1)
326 + go func() {
327 + defer wg.Done()
328 + ctx, err := ShmServerCreate(testShmRunDir, svc, 4, 65536, 65536)
329 + if err != nil {
330 + serverErr = fmt.Errorf("server create: %w", err)
331 + return
332 + }
333 + defer ctx.ShmDestroy()
334 +
335 + buf := make([]byte, 65536)
336 + mlen, err := ctx.ShmReceive(buf, 5000)
337 + if err != nil {
338 + serverErr = fmt.Errorf("server receive: %w", err)
339 + return
340 + }
341 +
342 + hdr, err := protocol.DecodeHeader(buf[:mlen])
343 + if err != nil {
344 + serverErr = fmt.Errorf("decode: %w", err)
345 + return
346 + }
347 + payload := make([]byte, mlen-protocol.HeaderSize)
348 + copy(payload, buf[protocol.HeaderSize:mlen])
349 + resp := buildShmMessage(protocol.KindResponse, hdr.Code, hdr.MessageID, payload)
350 + if err := ctx.ShmSend(resp); err != nil {
351 + serverErr = fmt.Errorf("server send: %w", err)
352 + }
353 + }()
354 +
355 + client := waitShmClientAttach(t, testShmRunDir, svc, 4)
356 + defer client.ShmClose()
357 +
358 + // 60000 bytes of payload
359 + payload := make([]byte, 60000)
360 + for i := range payload {
361 + payload[i] = byte(i & 0xFF)
362 + }
363 + msg := buildShmMessage(protocol.KindRequest, 1, 999, payload)
364 + if err := client.ShmSend(msg); err != nil {
365 + t.Fatalf("client send: %v", err)
366 + }
367 +
368 + respBuf := make([]byte, 65536)
369 + rlen, err := client.ShmReceive(respBuf, 5000)
370 + if err != nil {
371 + t.Fatalf("client receive: %v", err)
372 + }
373 +
374 + if rlen != protocol.HeaderSize+len(payload) {
375 + t.Fatalf("response length: got %d, want %d", rlen, protocol.HeaderSize+len(payload))
376 + }
377 +
378 + respPayload := respBuf[protocol.HeaderSize:rlen]
379 + if !bytes.Equal(respPayload, payload) {
380 + t.Errorf("response payload pattern mismatch")
381 + }
382 +
383 + wg.Wait()
384 + if serverErr != nil {
385 + t.Fatalf("server error: %v", serverErr)
386 + }
387 +}
388 +
389 +func TestShmChaosForgedLength(t *testing.T) {
390 + ensureShmRunDir(t)
391 + svc := uniqueShmService(t, "go_shm_forged")
392 + cleanupShmFiles(t, svc)
393 + defer cleanupShmFiles(t, svc)
394 +
395 + const reqCap uint32 = 1024
396 + const respCap uint32 = 1024
397 +
398 + // --- Server-side receive with forged req_len ---
399 +
400 + srv, err := ShmServerCreate(testShmRunDir, svc, 100, reqCap, respCap)
401 + if err != nil {
402 + t.Fatalf("server create: %v", err)
403 + }
404 + defer srv.ShmDestroy()
405 +
406 + // Attach a client so we can also test client-side forged resp_len
407 + client, err := ShmClientAttach(testShmRunDir, svc, 100)
408 + if err != nil {
409 + t.Fatalf("client attach: %v", err)
410 + }
411 + defer client.ShmClose()
412 +
413 + buf := make([]byte, 65536)
414 +
415 + // Test forged req_len values on the server side.
416 + // We directly manipulate the mapped region to simulate a malicious client.
417 + forgedLengths := []uint32{0, reqCap + 1, 0xFFFFFFFF}
418 +
419 + for _, forgedLen := range forgedLengths {
420 + // Write garbage into the request area
421 + for i := srv.requestOffset; i < srv.requestOffset+reqCap; i++ {
422 + srv.data[i] = 0xAA
423 + }
424 +
425 + // Store forged req_len
426 + if err := atomicStoreU32(srv.data, shmHeaderReqLenOff, forgedLen); err != nil {
427 + t.Fatalf("store forged req_len=%d: %v", forgedLen, err)
428 + }
429 +
430 + // Bump req_seq to signal a "message" arrived
431 + if err := atomicAddU64(srv.data, shmHeaderReqSeqOff, 1); err != nil {
432 + t.Fatalf("add req_seq for forged=%d: %v", forgedLen, err)
433 + }
434 +
435 + // Bump req_signal to wake futex
436 + if err := atomicAddU32(srv.data, shmHeaderReqSignalOff, 1); err != nil {
437 + t.Fatalf("add req_signal for forged=%d: %v", forgedLen, err)
438 + }
439 +
440 + mlen, recvErr := srv.ShmReceive(buf, 1000)
441 +
442 + if forgedLen == 0 {
443 + // Zero-length: no copy, no error, returns 0 bytes
444 + if recvErr != nil {
445 + t.Errorf("forged req_len=0: unexpected error: %v", recvErr)
446 + }
447 + if mlen != 0 {
448 + t.Errorf("forged req_len=0: got mlen=%d, want 0", mlen)
449 + }
450 + } else {
451 + // Oversized: must return ErrShmMsgTooLarge, no panic
452 + if !errors.Is(recvErr, ErrShmMsgTooLarge) {
453 + t.Errorf("forged req_len=%d: got err=%v, want ErrShmMsgTooLarge", forgedLen, recvErr)
454 + }
455 + }
456 + }
457 +
458 + // --- Client-side receive with forged resp_len ---
459 +
460 + forgedRespLengths := []uint32{0, respCap + 1, 0xFFFFFFFF}
461 +
462 + for _, forgedLen := range forgedRespLengths {
463 + // Write garbage into the response area
464 + for i := srv.responseOffset; i < srv.responseOffset+respCap; i++ {
465 + srv.data[i] = 0xBB
466 + }
467 +
468 + // Store forged resp_len (server and client share the same mapped region)
469 + if err := atomicStoreU32(client.data, shmHeaderRespLenOff, forgedLen); err != nil {
470 + t.Fatalf("store forged resp_len=%d: %v", forgedLen, err)
471 + }
472 +
473 + // Bump resp_seq
474 + if err := atomicAddU64(client.data, shmHeaderRespSeqOff, 1); err != nil {
475 + t.Fatalf("add resp_seq for forged=%d: %v", forgedLen, err)
476 + }
477 +
478 + // Bump resp_signal
479 + if err := atomicAddU32(client.data, shmHeaderRespSignalOff, 1); err != nil {
480 + t.Fatalf("add resp_signal for forged=%d: %v", forgedLen, err)
481 + }
482 +
483 + mlen, recvErr := client.ShmReceive(buf, 1000)
484 +
485 + if forgedLen == 0 {
486 + if recvErr != nil {
487 + t.Errorf("forged resp_len=0: unexpected error: %v", recvErr)
488 + }
489 + if mlen != 0 {
490 + t.Errorf("forged resp_len=0: got mlen=%d, want 0", mlen)
491 + }
492 + } else {
493 + if !errors.Is(recvErr, ErrShmMsgTooLarge) {
494 + t.Errorf("forged resp_len=%d: got err=%v, want ErrShmMsgTooLarge", forgedLen, recvErr)
495 + }
496 + }
497 + }
498 +}
499 +
500 +func TestShmMultiClient(t *testing.T) {
501 + ensureShmRunDir(t)
502 + svc := uniqueShmService(t, "go_shm_mcli")
503 + cleanupShmFiles(t, svc)
504 + defer cleanupShmFiles(t, svc)
505 +
506 + const numClients = 3
507 +
508 + type serverSlot struct {
509 + ctx *ShmContext
510 + err error
511 + got []byte // received payload
512 + }
513 +
514 + var wg sync.WaitGroup
515 + slots := make([]serverSlot, numClients)
516 +
517 + // Create server regions and start goroutines that receive + echo
518 + for i := 0; i < numClients; i++ {
519 + sessionID := uint64(i + 1)
520 + ctx, err := ShmServerCreate(testShmRunDir, svc, sessionID, 4096, 4096)
521 + if err != nil {
522 + // Clean up already-created regions
523 + for j := 0; j < i; j++ {
524 + slots[j].ctx.ShmDestroy()
525 + }
526 + t.Fatalf("server create session %d: %v", sessionID, err)
527 + }
528 + slots[i].ctx = ctx
529 +
530 + idx := i
531 + wg.Add(1)
532 + go func() {
533 + defer wg.Done()
534 + buf := make([]byte, 65536)
535 + mlen, err := slots[idx].ctx.ShmReceive(buf, 5000)
536 + if err != nil {
537 + slots[idx].err = fmt.Errorf("receive: %w", err)
538 + return
539 + }
540 + // Save received payload for verification
541 + slots[idx].got = make([]byte, mlen)
542 + copy(slots[idx].got, buf[:mlen])
543 +
544 + // Parse and echo back as response
545 + hdr, err := protocol.DecodeHeader(buf[:mlen])
546 + if err != nil {
547 + slots[idx].err = fmt.Errorf("decode: %w", err)
548 + return
549 + }
550 + payload := make([]byte, mlen-protocol.HeaderSize)
551 + copy(payload, buf[protocol.HeaderSize:mlen])
552 + resp := buildShmMessage(protocol.KindResponse, hdr.Code, hdr.MessageID, payload)
553 + if err := slots[idx].ctx.ShmSend(resp); err != nil {
554 + slots[idx].err = fmt.Errorf("send: %w", err)
555 + }
556 + }()
557 + }
558 +
559 + // Attach clients and send unique messages
560 + clients := make([]*ShmContext, numClients)
561 + payloads := make([][]byte, numClients)
562 + for i := 0; i < numClients; i++ {
563 + sessionID := uint64(i + 1)
564 + c := waitShmClientAttach(t, testShmRunDir, svc, sessionID)
565 + defer c.ShmClose()
566 + clients[i] = c
567 +
568 + // Each client sends a unique payload: [0xC0+i, session_id_byte, 0xDE, 0xAD]
569 + payloads[i] = []byte{byte(0xC0 + i), byte(sessionID), 0xDE, 0xAD}
570 + msg := buildShmMessage(protocol.KindRequest, protocol.MethodIncrement, uint64(100+i), payloads[i])
571 + if err := c.ShmSend(msg); err != nil {
572 + t.Fatalf("client %d send: %v", i, err)
573 + }
574 + }
575 +
576 + // Each client receives its own response
577 + for i := 0; i < numClients; i++ {
578 + respBuf := make([]byte, 65536)
579 + rlen, err := clients[i].ShmReceive(respBuf, 5000)
580 + if err != nil {
581 + t.Fatalf("client %d receive: %v", i, err)
582 + }
583 +
584 + rhdr, err := protocol.DecodeHeader(respBuf[:rlen])
585 + if err != nil {
586 + t.Fatalf("client %d decode response: %v", i, err)
587 + }
588 +
589 + if rhdr.Kind != protocol.KindResponse {
590 + t.Errorf("client %d: kind=%d, want %d", i, rhdr.Kind, protocol.KindResponse)
591 + }
592 + if rhdr.MessageID != uint64(100+i) {
593 + t.Errorf("client %d: message_id=%d, want %d", i, rhdr.MessageID, 100+i)
594 + }
595 +
596 + respPayload := respBuf[protocol.HeaderSize:rlen]
597 + if !bytes.Equal(respPayload, payloads[i]) {
598 + t.Errorf("client %d: payload mismatch: got %x, want %x", i, respPayload, payloads[i])
599 + }
600 + }
601 +
602 + // Wait for server goroutines to finish
603 + wg.Wait()
604 +
605 + // Check for server errors and verify no cross-contamination
606 + for i := 0; i < numClients; i++ {
607 + if slots[i].err != nil {
608 + t.Errorf("server %d error: %v", i, slots[i].err)
609 + continue
610 + }
611 + // Verify each server got the right client's message (check payload)
612 + if len(slots[i].got) < protocol.HeaderSize+len(payloads[i]) {
613 + t.Errorf("server %d: received too few bytes: %d", i, len(slots[i].got))
614 + continue
615 + }
616 + srvPayload := slots[i].got[protocol.HeaderSize:]
617 + if !bytes.Equal(srvPayload, payloads[i]) {
618 + t.Errorf("server %d: cross-contamination: got %x, want %x", i, srvPayload, payloads[i])
619 + }
620 + }
621 +
622 + // Cleanup all server regions
623 + for i := 0; i < numClients; i++ {
624 + slots[i].ctx.ShmDestroy()
625 + }
626 +}
src/go/pkg/netipc/transport/posix/shm_more_edge_test.go new
+493
@@ -0,0 +1,493 @@
1 +//go:build linux
2 +
3 +package posix
4 +
5 +import (
6 + "encoding/binary"
7 + "errors"
8 + "os"
9 + "path/filepath"
10 + "testing"
11 +)
12 +
13 +func defaultShmLayout() (uint32, uint32, uint32, uint32) {
14 + reqCap := shmAlign64(128)
15 + respCap := shmAlign64(128)
16 + reqOff := shmAlign64(uint32(shmHeaderLen))
17 + respOff := shmAlign64(reqOff + reqCap)
18 + return reqOff, reqCap, respOff, respCap
19 +}
20 +
21 +func fillShmHeader(data []byte, ownerPID int32, ownerGen uint32, reqOff, reqCap, respOff, respCap uint32) {
22 + if len(data) < int(shmHeaderLen) {
23 + return
24 + }
25 + binary.NativeEndian.PutUint32(data[shmHeaderMagicOff:shmHeaderMagicOff+4], shmRegionMagic)
26 + binary.NativeEndian.PutUint16(data[shmHeaderVersionOff:shmHeaderVersionOff+2], shmRegionVersion)
27 + binary.NativeEndian.PutUint16(data[shmHeaderHeaderLenOff:shmHeaderHeaderLenOff+2], uint16(shmHeaderLen))
28 + binary.NativeEndian.PutUint32(data[shmHeaderOwnerPidOff:shmHeaderOwnerPidOff+4], uint32(ownerPID))
29 + binary.NativeEndian.PutUint32(data[shmHeaderOwnerGenOff:shmHeaderOwnerGenOff+4], ownerGen)
30 + binary.NativeEndian.PutUint32(data[shmHeaderReqOffOff:shmHeaderReqOffOff+4], reqOff)
31 + binary.NativeEndian.PutUint32(data[shmHeaderReqCapOff:shmHeaderReqCapOff+4], reqCap)
32 + binary.NativeEndian.PutUint32(data[shmHeaderRespOffOff:shmHeaderRespOffOff+4], respOff)
33 + binary.NativeEndian.PutUint32(data[shmHeaderRespCapOff:shmHeaderRespCapOff+4], respCap)
34 +}
35 +
36 +func writeRawShmRegionFile(t *testing.T, runDir, service string, sessionID uint64, size int, fill func([]byte)) string {
37 + t.Helper()
38 + if err := os.MkdirAll(runDir, 0700); err != nil {
39 + t.Fatalf("mkdir %s: %v", runDir, err)
40 + }
41 +
42 + path, err := buildShmPath(runDir, service, sessionID)
43 + if err != nil {
44 + t.Fatalf("build path: %v", err)
45 + }
46 +
47 + f, err := os.OpenFile(path, os.O_CREATE|os.O_TRUNC|os.O_RDWR, 0600)
48 + if err != nil {
49 + t.Fatalf("open %s: %v", path, err)
50 + }
51 + defer f.Close()
52 +
53 + if err := f.Truncate(int64(size)); err != nil {
54 + t.Fatalf("truncate %s: %v", path, err)
55 + }
56 + if size == 0 {
57 + return path
58 + }
59 +
60 + data := make([]byte, size)
61 + if fill != nil {
62 + fill(data)
63 + }
64 + if _, err := f.WriteAt(data, 0); err != nil {
65 + t.Fatalf("write %s: %v", path, err)
66 + }
67 + return path
68 +}
69 +
70 +func TestShmOwnerAliveFalsePaths(t *testing.T) {
71 + if (&ShmContext{data: make([]byte, 8)}).OwnerAlive() {
72 + t.Fatal("short header should not be treated as alive")
73 + }
74 +
75 + runDir := t.TempDir()
76 + svc := "go_shm_owner_false"
77 + ctx, err := ShmServerCreate(runDir, svc, 1, 1024, 1024)
78 + if err != nil {
79 + t.Fatalf("server create: %v", err)
80 + }
81 + defer ctx.ShmDestroy()
82 +
83 + binary.NativeEndian.PutUint32(ctx.data[shmHeaderOwnerPidOff:shmHeaderOwnerPidOff+4], 0)
84 + if ctx.OwnerAlive() {
85 + t.Fatal("owner pid 0 should not be treated as alive")
86 + }
87 +
88 + binary.NativeEndian.PutUint32(ctx.data[shmHeaderOwnerPidOff:shmHeaderOwnerPidOff+4], uint32(int32(os.Getpid())))
89 + binary.NativeEndian.PutUint32(ctx.data[shmHeaderOwnerGenOff:shmHeaderOwnerGenOff+4], ctx.ownerGeneration+1)
90 + if ctx.OwnerAlive() {
91 + t.Fatal("generation mismatch should not be treated as alive")
92 + }
93 +
94 + ctx.ownerGeneration = 0
95 + if !ctx.OwnerAlive() {
96 + t.Fatal("legacy zero cached generation should skip generation check")
97 + }
98 +}
99 +
100 +func TestShmServerCreateRejectsLiveRegion(t *testing.T) {
101 + runDir := t.TempDir()
102 + svc := "go_shm_live_region"
103 +
104 + first, err := ShmServerCreate(runDir, svc, 7, 1024, 1024)
105 + if err != nil {
106 + t.Fatalf("first create: %v", err)
107 + }
108 + defer first.ShmDestroy()
109 +
110 + _, err = ShmServerCreate(runDir, svc, 7, 1024, 1024)
111 + if !errors.Is(err, ErrShmOpen) {
112 + t.Fatalf("second create error = %v, want %v", err, ErrShmOpen)
113 + }
114 +}
115 +
116 +func TestShmServerCreateRecoversInvalidAndLegacyFiles(t *testing.T) {
117 + runDir := t.TempDir()
118 +
119 + t.Run("tiny-invalid-file", func(t *testing.T) {
120 + svc := "go_shm_recover_tiny"
121 + writeRawShmRegionFile(t, runDir, svc, 1, 8, nil)
122 +
123 + ctx, err := ShmServerCreate(runDir, svc, 1, 1024, 1024)
124 + if err != nil {
125 + t.Fatalf("create after tiny stale file: %v", err)
126 + }
127 + defer ctx.ShmDestroy()
128 + })
129 +
130 + t.Run("zero-generation-legacy-file", func(t *testing.T) {
131 + svc := "go_shm_recover_legacy"
132 + reqOff, reqCap, respOff, respCap := defaultShmLayout()
133 + writeRawShmRegionFile(t, runDir, svc, 2, int(respOff+respCap), func(data []byte) {
134 + fillShmHeader(data, int32(os.Getpid()), 0, reqOff, reqCap, respOff, respCap)
135 + })
136 +
137 + ctx, err := ShmServerCreate(runDir, svc, 2, 1024, 1024)
138 + if err != nil {
139 + t.Fatalf("create after zero-generation stale file: %v", err)
140 + }
141 + defer ctx.ShmDestroy()
142 + })
143 +}
144 +
145 +func TestShmClientAttachRejectsMalformedHeaders(t *testing.T) {
146 + runDir := t.TempDir()
147 + reqOff, reqCap, respOff, respCap := defaultShmLayout()
148 + validSize := int(respOff + respCap)
149 +
150 + tests := []struct {
151 + name string
152 + size int
153 + fill func([]byte)
154 + want error
155 + }{
156 + {
157 + name: "file-too-small",
158 + size: 8,
159 + fill: nil,
160 + want: ErrShmNotReady,
161 + },
162 + {
163 + name: "bad-magic",
164 + size: validSize,
165 + fill: func(data []byte) {
166 + fillShmHeader(data, int32(os.Getpid()), 1, reqOff, reqCap, respOff, respCap)
167 + binary.NativeEndian.PutUint32(data[shmHeaderMagicOff:shmHeaderMagicOff+4], 0)
168 + },
169 + want: ErrShmBadMagic,
170 + },
171 + {
172 + name: "bad-version",
173 + size: validSize,
174 + fill: func(data []byte) {
175 + fillShmHeader(data, int32(os.Getpid()), 1, reqOff, reqCap, respOff, respCap)
176 + binary.NativeEndian.PutUint16(data[shmHeaderVersionOff:shmHeaderVersionOff+2], shmRegionVersion+1)
177 + },
178 + want: ErrShmBadVersion,
179 + },
180 + {
181 + name: "bad-header-length",
182 + size: validSize,
183 + fill: func(data []byte) {
184 + fillShmHeader(data, int32(os.Getpid()), 1, reqOff, reqCap, respOff, respCap)
185 + binary.NativeEndian.PutUint16(data[shmHeaderHeaderLenOff:shmHeaderHeaderLenOff+2], 32)
186 + },
187 + want: ErrShmBadHeader,
188 + },
189 + {
190 + name: "bad-layout-alignment",
191 + size: validSize,
192 + fill: func(data []byte) {
193 + fillShmHeader(data, int32(os.Getpid()), 1, reqOff+1, reqCap, respOff, respCap)
194 + },
195 + want: ErrShmBadSize,
196 + },
197 + {
198 + name: "declared-region-beyond-file-size",
199 + size: int(shmHeaderLen) + 64,
200 + fill: func(data []byte) {
201 + fillShmHeader(data, int32(os.Getpid()), 1, reqOff, reqCap, respOff, respCap)
202 + },
203 + want: ErrShmBadSize,
204 + },
205 + }
206 +
207 + for i, tc := range tests {
208 + t.Run(tc.name, func(t *testing.T) {
209 + svc := "go_shm_attach_bad"
210 + writeRawShmRegionFile(t, runDir, svc, uint64(i+1), tc.size, tc.fill)
211 +
212 + _, err := ShmClientAttach(runDir, svc, uint64(i+1))
213 + if !errors.Is(err, tc.want) {
214 + t.Fatalf("attach error = %v, want %v", err, tc.want)
215 + }
216 + })
217 + }
218 +}
219 +
220 +func TestShmCreateAndAttachRejectInvalidServiceName(t *testing.T) {
221 + runDir := t.TempDir()
222 +
223 + if _, err := ShmServerCreate(runDir, "bad/name", 1, 1024, 1024); !errors.Is(err, ErrShmBadParam) {
224 + t.Fatalf("ShmServerCreate invalid service name = %v, want %v", err, ErrShmBadParam)
225 + }
226 + if _, err := ShmClientAttach(runDir, "bad/name", 1); !errors.Is(err, ErrShmBadParam) {
227 + t.Fatalf("ShmClientAttach invalid service name = %v, want %v", err, ErrShmBadParam)
228 + }
229 +}
230 +
231 +func TestShmSendBadParamGuards(t *testing.T) {
232 + if err := (&ShmContext{}).ShmSend([]byte{1}); !errors.Is(err, ErrShmBadParam) {
233 + t.Fatalf("ShmSend nil context = %v, want %v", err, ErrShmBadParam)
234 + }
235 +
236 + storeLenCtx := &ShmContext{
237 + role: ShmRoleClient,
238 + data: make([]byte, 39),
239 + requestCapacity: 64,
240 + }
241 + if err := storeLenCtx.ShmSend([]byte{1}); !errors.Is(err, ErrShmBadParam) {
242 + t.Fatalf("ShmSend short backing slice for len store = %v, want %v", err, ErrShmBadParam)
243 + }
244 +
245 + addSeqCtx := &ShmContext{
246 + role: ShmRoleClient,
247 + data: make([]byte, 48),
248 + requestCapacity: 64,
249 + }
250 + if err := addSeqCtx.ShmSend([]byte{1}); !errors.Is(err, ErrShmBadParam) {
251 + t.Fatalf("ShmSend short backing slice for seq add = %v, want %v", err, ErrShmBadParam)
252 + }
253 +
254 + addSignalCtx := &ShmContext{
255 + role: ShmRoleClient,
256 + data: make([]byte, 56),
257 + requestCapacity: 64,
258 + }
259 + if err := addSignalCtx.ShmSend([]byte{1}); !errors.Is(err, ErrShmBadParam) {
260 + t.Fatalf("ShmSend short backing slice for signal add = %v, want %v", err, ErrShmBadParam)
261 + }
262 +}
263 +
264 +func TestShmReceiveBadParamAndTimeoutPaths(t *testing.T) {
265 + if _, err := (&ShmContext{}).ShmReceive(make([]byte, 8), 1); !errors.Is(err, ErrShmBadParam) {
266 + t.Fatalf("ShmReceive nil context = %v, want %v", err, ErrShmBadParam)
267 + }
268 +
269 + if _, err := (&ShmContext{data: make([]byte, 8)}).ShmReceive(nil, 1); !errors.Is(err, ErrShmBadParam) {
270 + t.Fatalf("ShmReceive empty buffer = %v, want %v", err, ErrShmBadParam)
271 + }
272 +
273 + spinLoadSeqCtx := &ShmContext{
274 + role: ShmRoleClient,
275 + data: make([]byte, 8),
276 + responseCapacity: 64,
277 + SpinTries: 1,
278 + }
279 + if _, err := spinLoadSeqCtx.ShmReceive(make([]byte, 8), 1); !errors.Is(err, ErrShmBadParam) {
280 + t.Fatalf("ShmReceive spin-phase seq load = %v, want %v", err, ErrShmBadParam)
281 + }
282 +
283 + spinLoadLenCtx := &ShmContext{
284 + role: ShmRoleClient,
285 + data: make([]byte, 48),
286 + responseCapacity: 64,
287 + SpinTries: 1,
288 + }
289 + binary.NativeEndian.PutUint64(
290 + spinLoadLenCtx.data[shmHeaderRespSeqOff:shmHeaderRespSeqOff+8],
291 + 1,
292 + )
293 + if _, err := spinLoadLenCtx.ShmReceive(make([]byte, 8), 1); !errors.Is(err, ErrShmBadParam) {
294 + t.Fatalf("ShmReceive spin-phase msg_len load = %v, want %v", err, ErrShmBadParam)
295 + }
296 +
297 + futexLoadSignalCtx := &ShmContext{
298 + role: ShmRoleClient,
299 + data: make([]byte, 8),
300 + responseCapacity: 64,
301 + SpinTries: 0,
302 + }
303 + if _, err := futexLoadSignalCtx.ShmReceive(make([]byte, 8), 1); !errors.Is(err, ErrShmBadParam) {
304 + t.Fatalf("ShmReceive futex-phase signal load = %v, want %v", err, ErrShmBadParam)
305 + }
306 +
307 + futexLoadSeqCtx := &ShmContext{
308 + role: ShmRoleClient,
309 + data: make([]byte, 48),
310 + responseCapacity: 64,
311 + SpinTries: 0,
312 + }
313 + if _, err := futexLoadSeqCtx.ShmReceive(make([]byte, 8), 1); !errors.Is(err, ErrShmBadParam) {
314 + t.Fatalf("ShmReceive futex-phase seq load = %v, want %v", err, ErrShmBadParam)
315 + }
316 +
317 + runDir := t.TempDir()
318 + svc := "go_shm_timeout"
319 + srv, err := ShmServerCreate(runDir, svc, 11, 1024, 1024)
320 + if err != nil {
321 + t.Fatalf("ShmServerCreate timeout fixture failed: %v", err)
322 + }
323 + defer srv.ShmDestroy()
324 +
325 + client, err := ShmClientAttach(runDir, svc, 11)
326 + if err != nil {
327 + t.Fatalf("ShmClientAttach timeout fixture failed: %v", err)
328 + }
329 + defer client.ShmClose()
330 +
331 + buf := make([]byte, 64)
332 + if _, err := srv.ShmReceive(buf, 1); !errors.Is(err, ErrShmTimeout) {
333 + t.Fatalf("server ShmReceive timeout = %v, want %v", err, ErrShmTimeout)
334 + }
335 + if _, err := client.ShmReceive(buf, 1); !errors.Is(err, ErrShmTimeout) {
336 + t.Fatalf("client ShmReceive timeout = %v, want %v", err, ErrShmTimeout)
337 + }
338 +}
339 +
340 +func TestCheckShmStaleVariants(t *testing.T) {
341 + runDir := t.TempDir()
342 + reqOff, reqCap, respOff, respCap := defaultShmLayout()
343 + validSize := int(respOff + respCap)
344 +
345 + missingPath := filepath.Join(runDir, "missing.ipcshm")
346 + if got := checkShmStale(missingPath); got != shmStaleNotExist {
347 + t.Fatalf("missing path result = %v, want %v", got, shmStaleNotExist)
348 + }
349 +
350 + tinyPath := writeRawShmRegionFile(t, runDir, "go_shm_check_tiny", 1, 8, nil)
351 + if got := checkShmStale(tinyPath); got != shmStaleInvalid {
352 + t.Fatalf("tiny file result = %v, want %v", got, shmStaleInvalid)
353 + }
354 + if _, err := os.Stat(tinyPath); !errors.Is(err, os.ErrNotExist) {
355 + t.Fatalf("tiny file should be removed, stat err = %v", err)
356 + }
357 +
358 + badMagicPath := writeRawShmRegionFile(t, runDir, "go_shm_check_magic", 2, validSize, func(data []byte) {
359 + fillShmHeader(data, int32(os.Getpid()), 1, reqOff, reqCap, respOff, respCap)
360 + binary.NativeEndian.PutUint32(data[shmHeaderMagicOff:shmHeaderMagicOff+4], 0)
361 + })
362 + if got := checkShmStale(badMagicPath); got != shmStaleInvalid {
363 + t.Fatalf("bad magic result = %v, want %v", got, shmStaleInvalid)
364 + }
365 +
366 + livePath := writeRawShmRegionFile(t, runDir, "go_shm_check_live", 3, validSize, func(data []byte) {
367 + fillShmHeader(data, int32(os.Getpid()), 7, reqOff, reqCap, respOff, respCap)
368 + })
369 + if got := checkShmStale(livePath); got != shmStaleLive {
370 + t.Fatalf("live path result = %v, want %v", got, shmStaleLive)
371 + }
372 + if _, err := os.Stat(livePath); err != nil {
373 + t.Fatalf("live file should remain, stat err = %v", err)
374 + }
375 +
376 + legacyPath := writeRawShmRegionFile(t, runDir, "go_shm_check_legacy", 4, validSize, func(data []byte) {
377 + fillShmHeader(data, int32(os.Getpid()), 0, reqOff, reqCap, respOff, respCap)
378 + })
379 + if got := checkShmStale(legacyPath); got != shmStaleRecovered {
380 + t.Fatalf("legacy path result = %v, want %v", got, shmStaleRecovered)
381 + }
382 + if _, err := os.Stat(legacyPath); !errors.Is(err, os.ErrNotExist) {
383 + t.Fatalf("legacy stale file should be removed, stat err = %v", err)
384 + }
385 +
386 + unreadablePath := writeRawShmRegionFile(t, runDir, "go_shm_check_unreadable", 5, validSize, func(data []byte) {
387 + fillShmHeader(data, int32(os.Getpid()), 1, reqOff, reqCap, respOff, respCap)
388 + })
389 + if err := os.Chmod(unreadablePath, 0); err != nil {
390 + t.Fatalf("chmod unreadable stale file: %v", err)
391 + }
392 + if got := checkShmStale(unreadablePath); got != shmStaleInvalid {
393 + t.Fatalf("unreadable path result = %v, want %v", got, shmStaleInvalid)
394 + }
395 + if _, err := os.Stat(unreadablePath); !errors.Is(err, os.ErrNotExist) {
396 + t.Fatalf("unreadable stale file should be removed, stat err = %v", err)
397 + }
398 +
399 + dirPath, err := buildShmPath(runDir, "go_shm_check_dir", 6)
400 + if err != nil {
401 + t.Fatalf("build dir path: %v", err)
402 + }
403 + if err := os.MkdirAll(filepath.Join(dirPath, "keep"), 0700); err != nil {
404 + t.Fatalf("mkdir stale dir: %v", err)
405 + }
406 + if got := checkShmStale(dirPath); got != shmStaleInvalid {
407 + t.Fatalf("directory path result = %v, want %v", got, shmStaleInvalid)
408 + }
409 + if info, err := os.Stat(dirPath); err != nil || !info.IsDir() {
410 + t.Fatalf("non-empty stale directory should remain, stat err = %v", err)
411 + }
412 +}
413 +
414 +func TestShmServerCreateFailsWhenObstructionSurvivesRecovery(t *testing.T) {
415 + runDir := t.TempDir()
416 + svc := "go_shm_create_blocked"
417 + path, err := buildShmPath(runDir, svc, 7)
418 + if err != nil {
419 + t.Fatalf("build path: %v", err)
420 + }
421 + if err := os.MkdirAll(filepath.Join(path, "keep"), 0700); err != nil {
422 + t.Fatalf("mkdir obstructing directory: %v", err)
423 + }
424 +
425 + _, err = ShmServerCreate(runDir, svc, 7, 1024, 1024)
426 + if !errors.Is(err, ErrShmOpen) {
427 + t.Fatalf("ShmServerCreate blocked path error = %v, want %v", err, ErrShmOpen)
428 + }
429 + if info, statErr := os.Stat(path); statErr != nil || !info.IsDir() {
430 + t.Fatalf("obstructing directory should remain after failed recovery, stat err = %v", statErr)
431 + }
432 +}
433 +
434 +func TestShmCleanupStaleMixedEntries(t *testing.T) {
435 + runDir := t.TempDir()
436 + reqOff, reqCap, respOff, respCap := defaultShmLayout()
437 + validSize := int(respOff + respCap)
438 + svc := "go_shm_cleanup"
439 +
440 + live, err := ShmServerCreate(runDir, svc, 1, 1024, 1024)
441 + if err != nil {
442 + t.Fatalf("live server create: %v", err)
443 + }
444 + defer live.ShmDestroy()
445 +
446 + tinyPath := writeRawShmRegionFile(t, runDir, svc, 2, 8, nil)
447 + legacyPath := writeRawShmRegionFile(t, runDir, svc, 3, validSize, func(data []byte) {
448 + fillShmHeader(data, int32(os.Getpid()), 0, reqOff, reqCap, respOff, respCap)
449 + })
450 + unrelatedPath := filepath.Join(runDir, "other-file.ipcshm")
451 + if err := os.WriteFile(unrelatedPath, []byte("keep"), 0600); err != nil {
452 + t.Fatalf("write unrelated file: %v", err)
453 + }
454 + matchingDir := filepath.Join(runDir, svc+"-dir.ipcshm")
455 + if err := os.Mkdir(matchingDir, 0700); err != nil {
456 + t.Fatalf("mkdir matching dir: %v", err)
457 + }
458 +
459 + ShmCleanupStale(runDir, svc)
460 +
461 + if _, err := os.Stat(live.path); err != nil {
462 + t.Fatalf("live region should remain, stat err = %v", err)
463 + }
464 + if _, err := os.Stat(tinyPath); !errors.Is(err, os.ErrNotExist) {
465 + t.Fatalf("tiny stale file should be removed, stat err = %v", err)
466 + }
467 + if _, err := os.Stat(legacyPath); !errors.Is(err, os.ErrNotExist) {
468 + t.Fatalf("legacy stale file should be removed, stat err = %v", err)
469 + }
470 + if _, err := os.Stat(unrelatedPath); err != nil {
471 + t.Fatalf("unrelated file should remain, stat err = %v", err)
472 + }
473 + if info, err := os.Stat(matchingDir); err != nil || !info.IsDir() {
474 + t.Fatalf("matching directory should remain, stat err = %v", err)
475 + }
476 +}
477 +
478 +func TestShmCleanupStaleMissingDirAndUnrelatedFile(t *testing.T) {
479 + missingDir := filepath.Join(t.TempDir(), "missing")
480 + ShmCleanupStale(missingDir, "go_shm_missing")
481 +
482 + runDir := t.TempDir()
483 + unrelatedPath := filepath.Join(runDir, "go_shm_other-note.txt")
484 + if err := os.WriteFile(unrelatedPath, []byte("x"), 0600); err != nil {
485 + t.Fatalf("write unrelated file: %v", err)
486 + }
487 +
488 + ShmCleanupStale(runDir, "go_shm_other")
489 +
490 + if _, err := os.Stat(unrelatedPath); err != nil {
491 + t.Fatalf("unrelated file should remain after cleanup, stat err = %v", err)
492 + }
493 +}
src/go/pkg/netipc/transport/posix/shm_pause_amd64.go new
+6
@@ -0,0 +1,6 @@
1 +//go:build linux && amd64
2 +
3 +package posix
4 +
5 +// spinPause emits a PAUSE instruction (defined in shm_pause_amd64.s).
6 +func spinPause()
src/go/pkg/netipc/transport/posix/shm_pause_amd64.s new
+9
@@ -0,0 +1,9 @@
1 +//go:build linux && amd64
2 +
3 +#include "textflag.h"
4 +
5 +// func spinPause()
6 +// CPU PAUSE hint for spin loops — reduces pipeline contention.
7 +TEXT ·spinPause(SB),NOSPLIT,$0-0
8 + PAUSE
9 + RET
src/go/pkg/netipc/transport/posix/shm_pause_other.go new
+12
@@ -0,0 +1,12 @@
1 +//go:build linux && !amd64
2 +
3 +package posix
4 +
5 +import "runtime"
6 +
7 +// spinPause is a no-op fallback on non-amd64 architectures.
8 +// Gosched yields to the scheduler briefly, which is the closest
9 +// portable equivalent to a CPU pause hint.
10 +func spinPause() {
11 + runtime.Gosched()
12 +}
src/go/pkg/netipc/transport/posix/uds.go new
+1079
@@ -0,0 +1,1079 @@
1 +//go:build unix
2 +
3 +// Package posix implements the L1 POSIX UDS SEQPACKET transport.
4 +//
5 +// Connection lifecycle, handshake with profile/limit negotiation,
6 +// and send/receive with transparent chunking over AF_UNIX SEQPACKET sockets.
7 +// Wire-compatible with the C and Rust implementations.
8 +//
9 +// Pure Go — no cgo. Works with CGO_ENABLED=0.
10 +package posix
11 +
12 +import (
13 + "encoding/binary"
14 + "errors"
15 + "fmt"
16 + "net"
17 + "os"
18 + "path/filepath"
19 + "sync/atomic"
20 + "syscall"
21 + "unsafe"
22 +
23 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
24 +)
25 +
26 +// ---------------------------------------------------------------------------
27 +// Constants
28 +// ---------------------------------------------------------------------------
29 +
30 +const (
31 + defaultBacklog = 16
32 + defaultBatchItems = 1
33 + defaultPacketSizeFallback uint32 = 65536
34 + helloPayloadSize = 44
35 + helloAckPayloadSize = 48
36 +
37 + // sun_path max — 108 on Linux, 104 on macOS/FreeBSD.
38 + // We use a conservative limit.
39 + maxSunPath = 104
40 +)
41 +
42 +// ---------------------------------------------------------------------------
43 +// Errors
44 +// ---------------------------------------------------------------------------
45 +
46 +var (
47 + ErrPathTooLong = errors.New("socket path exceeds sun_path limit")
48 + ErrSocket = errors.New("socket syscall failed")
49 + ErrConnect = errors.New("connect failed")
50 + ErrAccept = errors.New("accept failed")
51 + ErrSend = errors.New("send failed")
52 + ErrRecv = errors.New("recv failed or peer disconnected")
53 + ErrHandshake = errors.New("handshake protocol error")
54 + ErrAuthFailed = errors.New("authentication token rejected")
55 + ErrNoProfile = errors.New("no common transport profile")
56 + ErrIncompatible = errors.New("protocol or layout version mismatch")
57 + ErrProtocol = errors.New("wire protocol violation")
58 + ErrAddrInUse = errors.New("address already in use by live server")
59 + ErrChunk = errors.New("chunk header mismatch")
60 + ErrLimitExceeded = errors.New("negotiated limit exceeded")
61 + ErrBadParam = errors.New("invalid argument")
62 + ErrDuplicateMsgID = errors.New("duplicate message_id")
63 + ErrUnknownMsgID = errors.New("unknown response message_id")
64 +)
65 +
66 +// wrapErr creates a descriptive error wrapping a sentinel.
67 +func wrapErr(sentinel error, detail string) error {
68 + return fmt.Errorf("%w: %s", sentinel, detail)
69 +}
70 +
71 +func headerVersionIncompatible(buf []byte, expectedCode uint16) bool {
72 + if len(buf) < protocol.HeaderSize {
73 + return false
74 + }
75 +
76 + return binary.NativeEndian.Uint32(buf[0:4]) == protocol.MagicMsg &&
77 + binary.NativeEndian.Uint16(buf[4:6]) != protocol.Version &&
78 + binary.NativeEndian.Uint16(buf[6:8]) == protocol.HeaderLen &&
79 + binary.NativeEndian.Uint16(buf[8:10]) == protocol.KindControl &&
80 + binary.NativeEndian.Uint16(buf[12:14]) == expectedCode
81 +}
82 +
83 +func helloLayoutIncompatible(buf []byte) bool {
84 + return len(buf) >= 2 && binary.NativeEndian.Uint16(buf[0:2]) != 1
85 +}
86 +
87 +func helloAckLayoutIncompatible(buf []byte) bool {
88 + return len(buf) >= 2 && binary.NativeEndian.Uint16(buf[0:2]) != 1
89 +}
90 +
91 +// ---------------------------------------------------------------------------
92 +// Role
93 +// ---------------------------------------------------------------------------
94 +
95 +// Role distinguishes client vs server sessions.
96 +type Role int
97 +
98 +const (
99 + RoleClient Role = 1
100 + RoleServer Role = 2
101 +)
102 +
103 +// ---------------------------------------------------------------------------
104 +// Configuration
105 +// ---------------------------------------------------------------------------
106 +
107 +// ClientConfig configures a client connection.
108 +type ClientConfig struct {
109 + SupportedProfiles uint32
110 + PreferredProfiles uint32
111 + MaxRequestPayloadBytes uint32 // 0 = use default (1024)
112 + MaxRequestBatchItems uint32 // 0 = use default (1)
113 + MaxResponsePayloadBytes uint32
114 + MaxResponseBatchItems uint32
115 + AuthToken uint64
116 + PacketSize uint32 // 0 = auto-detect from SO_SNDBUF
117 +}
118 +
119 +// ServerConfig configures a listener and its accepted sessions.
120 +type ServerConfig struct {
121 + SupportedProfiles uint32
122 + PreferredProfiles uint32
123 + MaxRequestPayloadBytes uint32
124 + MaxRequestBatchItems uint32
125 + MaxResponsePayloadBytes uint32
126 + MaxResponseBatchItems uint32
127 + AuthToken uint64
128 + PacketSize uint32 // 0 = auto-detect from SO_SNDBUF
129 + Backlog int // 0 = default (16)
130 +}
131 +
132 +// ---------------------------------------------------------------------------
133 +// Session
134 +// ---------------------------------------------------------------------------
135 +
136 +// Session is a connected UDS SEQPACKET session (client or server side).
137 +type Session struct {
138 + fd int
139 + role Role
140 +
141 + // Negotiated limits
142 + MaxRequestPayloadBytes uint32
143 + MaxRequestBatchItems uint32
144 + MaxResponsePayloadBytes uint32
145 + MaxResponseBatchItems uint32
146 + PacketSize uint32
147 + SelectedProfile uint32
148 + SessionID uint64
149 +
150 + // Internal receive buffer for chunked reassembly
151 + recvBuf []byte
152 +
153 + // Reusable packet scratch buffer for receive chunk assembly.
154 + pktBuf []byte
155 +
156 + // In-flight message_id set (client-side only)
157 + inflightIDs map[uint64]struct{}
158 +}
159 +
160 +func (s *Session) failAllInflight() {
161 + if s.role != RoleClient || len(s.inflightIDs) == 0 {
162 + return
163 + }
164 + clear(s.inflightIDs)
165 +}
166 +
167 +// Fd returns the raw file descriptor for poll/epoll integration.
168 +func (s *Session) Fd() int {
169 + return s.fd
170 +}
171 +
172 +// Role returns the session role.
173 +func (s *Session) Role() Role {
174 + return s.role
175 +}
176 +
177 +// Close closes the session and releases resources.
178 +func (s *Session) Close() {
179 + if s.fd >= 0 {
180 + syscall.Close(s.fd)
181 + s.fd = -1
182 + }
183 + s.recvBuf = nil
184 + s.pktBuf = nil
185 + s.failAllInflight()
186 +}
187 +
188 +// Connect establishes a session to a server at {runDir}/{serviceName}.sock.
189 +// Performs the full handshake. Blocks until connected + handshake done.
190 +func Connect(runDir, serviceName string, config *ClientConfig) (*Session, error) {
191 + path, err := buildSocketPath(runDir, serviceName)
192 + if err != nil {
193 + return nil, err
194 + }
195 +
196 + fd, err := syscall.Socket(syscall.AF_UNIX, syscall.SOCK_SEQPACKET, 0)
197 + if err != nil {
198 + return nil, wrapErr(ErrSocket, err.Error())
199 + }
200 +
201 + session, herr := connectAndHandshake(fd, path, config)
202 + if herr != nil {
203 + syscall.Close(fd)
204 + return nil, herr
205 + }
206 + return session, nil
207 +}
208 +
209 +// Send sends one logical message. The caller fills Kind, Code, Flags,
210 +// ItemCount, MessageID in hdr; this function sets Magic/Version/
211 +// HeaderLen/PayloadLen. If the total message exceeds PacketSize,
212 +// it is chunked transparently.
213 +func (s *Session) Send(hdr *protocol.Header, payload []byte) error {
214 + if s.fd < 0 {
215 + return wrapErr(ErrBadParam, "session closed")
216 + }
217 +
218 + // Client-side: track in-flight message_ids for requests
219 + if s.role == RoleClient && hdr.Kind == protocol.KindRequest {
220 + if s.inflightIDs == nil {
221 + s.inflightIDs = make(map[uint64]struct{})
222 + }
223 + if _, exists := s.inflightIDs[hdr.MessageID]; exists {
224 + return wrapErr(ErrDuplicateMsgID, fmt.Sprintf("message_id %d", hdr.MessageID))
225 + }
226 + s.inflightIDs[hdr.MessageID] = struct{}{}
227 + }
228 +
229 + // Fill envelope fields
230 + hdr.Magic = protocol.MagicMsg
231 + hdr.Version = protocol.Version
232 + hdr.HeaderLen = protocol.HeaderLen
233 + hdr.PayloadLen = uint32(len(payload))
234 +
235 + tracked := s.role == RoleClient && hdr.Kind == protocol.KindRequest
236 +
237 + sendErr := s.sendInner(hdr, payload)
238 +
239 + // Rollback: remove message_id from in-flight set on send failure
240 + if sendErr != nil && tracked {
241 + if errors.Is(sendErr, ErrSend) {
242 + s.failAllInflight()
243 + } else {
244 + delete(s.inflightIDs, hdr.MessageID)
245 + }
246 + }
247 +
248 + return sendErr
249 +}
250 +
251 +// sendInner performs the actual send logic, separated so the caller can
252 +// rollback the in-flight set on failure.
253 +func (s *Session) sendInner(hdr *protocol.Header, payload []byte) error {
254 + totalMsg := protocol.HeaderSize + len(payload)
255 +
256 + // Single packet?
257 + if totalMsg <= int(s.PacketSize) {
258 + var hdrBuf [protocol.HeaderSize]byte
259 + hdr.Encode(hdrBuf[:])
260 + return rawSendIov(s.fd, hdrBuf[:], payload)
261 + }
262 +
263 + // Chunked send
264 + chunkPayloadBudget := int(s.PacketSize) - protocol.HeaderSize
265 + if chunkPayloadBudget <= 0 {
266 + return wrapErr(ErrBadParam, "packet_size too small")
267 + }
268 +
269 + firstChunkPayload := len(payload)
270 + if firstChunkPayload > chunkPayloadBudget {
271 + firstChunkPayload = chunkPayloadBudget
272 + }
273 +
274 + remainingAfterFirst := len(payload) - firstChunkPayload
275 + continuationChunks := uint32(0)
276 + if remainingAfterFirst > 0 {
277 + continuationChunks = uint32((remainingAfterFirst + chunkPayloadBudget - 1) / chunkPayloadBudget)
278 + }
279 + chunkCount := 1 + continuationChunks
280 +
281 + // Send first chunk: outer header + first part of payload
282 + var hdrBuf [protocol.HeaderSize]byte
283 + hdr.Encode(hdrBuf[:])
284 + if err := rawSendIov(s.fd, hdrBuf[:], payload[:firstChunkPayload]); err != nil {
285 + return err
286 + }
287 +
288 + // Send continuation chunks
289 + offset := firstChunkPayload
290 + for ci := uint32(1); ci < chunkCount; ci++ {
291 + remaining := len(payload) - offset
292 + thisChunk := remaining
293 + if thisChunk > chunkPayloadBudget {
294 + thisChunk = chunkPayloadBudget
295 + }
296 +
297 + chk := protocol.ChunkHeader{
298 + Magic: protocol.MagicChunk,
299 + Version: protocol.Version,
300 + Flags: 0,
301 + MessageID: hdr.MessageID,
302 + TotalMessageLen: uint32(totalMsg),
303 + ChunkIndex: ci,
304 + ChunkCount: chunkCount,
305 + ChunkPayloadLen: uint32(thisChunk),
306 + }
307 +
308 + var chkBuf [protocol.HeaderSize]byte
309 + chk.Encode(chkBuf[:])
310 + if err := rawSendIov(s.fd, chkBuf[:], payload[offset:offset+thisChunk]); err != nil {
311 + return err
312 + }
313 +
314 + offset += thisChunk
315 + }
316 +
317 + return nil
318 +}
319 +
320 +// Receive reads one logical message. Blocks until a complete message
321 +// arrives. buf is a caller-provided scratch buffer for the first packet.
322 +// On success, returns the header and a payload view valid until the next
323 +// Receive call on this session.
324 +func (s *Session) Receive(buf []byte) (protocol.Header, []byte, error) {
325 + if s.fd < 0 {
326 + return protocol.Header{}, nil, wrapErr(ErrBadParam, "session closed")
327 + }
328 +
329 + // Read first packet
330 + n, err := rawRecv(s.fd, buf)
331 + if err != nil {
332 + if errors.Is(err, ErrRecv) {
333 + s.failAllInflight()
334 + }
335 + return protocol.Header{}, nil, err
336 + }
337 +
338 + if n < protocol.HeaderSize {
339 + return protocol.Header{}, nil, wrapErr(ErrProtocol, "packet too short for header")
340 + }
341 +
342 + hdr, err := protocol.DecodeHeader(buf[:n])
343 + if err != nil {
344 + return protocol.Header{}, nil, wrapErr(ErrProtocol, "header decode: "+err.Error())
345 + }
346 +
347 + // Validate payload_len against negotiated directional limit.
348 + // Server receives requests; client receives responses.
349 + var maxPayload uint32
350 + if s.role == RoleServer {
351 + maxPayload = s.MaxRequestPayloadBytes
352 + } else {
353 + maxPayload = s.MaxResponsePayloadBytes
354 + }
355 + if hdr.PayloadLen > maxPayload {
356 + return protocol.Header{}, nil, wrapErr(ErrLimitExceeded,
357 + fmt.Sprintf("payload_len %d exceeds negotiated max %d", hdr.PayloadLen, maxPayload))
358 + }
359 +
360 + // Validate item_count against negotiated directional batch limit.
361 + var maxBatch uint32
362 + if s.role == RoleServer {
363 + maxBatch = s.MaxRequestBatchItems
364 + } else {
365 + maxBatch = s.MaxResponseBatchItems
366 + }
367 + if hdr.ItemCount > maxBatch {
368 + return protocol.Header{}, nil, wrapErr(ErrLimitExceeded,
369 + fmt.Sprintf("item_count %d exceeds negotiated max %d", hdr.ItemCount, maxBatch))
370 + }
371 +
372 + // Client-side: validate response message_id is in-flight
373 + if s.role == RoleClient && hdr.Kind == protocol.KindResponse {
374 + if s.inflightIDs == nil {
375 + return protocol.Header{}, nil, wrapErr(ErrUnknownMsgID,
376 + fmt.Sprintf("message_id %d", hdr.MessageID))
377 + }
378 + if _, exists := s.inflightIDs[hdr.MessageID]; !exists {
379 + return protocol.Header{}, nil, wrapErr(ErrUnknownMsgID,
380 + fmt.Sprintf("message_id %d", hdr.MessageID))
381 + }
382 + delete(s.inflightIDs, hdr.MessageID)
383 + }
384 +
385 + totalMsg := protocol.HeaderSize + int(hdr.PayloadLen)
386 +
387 + // Non-chunked: entire message in one packet
388 + if n >= totalMsg {
389 + payload := buf[protocol.HeaderSize : protocol.HeaderSize+int(hdr.PayloadLen)]
390 +
391 + // Validate batch directory
392 + if hdr.Flags&protocol.FlagBatch != 0 && hdr.ItemCount > 1 {
393 + dirBytes := int(hdr.ItemCount) * 8
394 + dirAligned := protocol.Align8(dirBytes)
395 + if len(payload) < dirAligned {
396 + return protocol.Header{}, nil, wrapErr(ErrProtocol, "batch dir exceeds payload")
397 + }
398 + packedAreaLen := uint32(len(payload) - dirAligned)
399 + if err := protocol.BatchDirValidate(payload[:dirBytes], hdr.ItemCount, packedAreaLen); err != nil {
400 + return protocol.Header{}, nil, wrapErr(ErrProtocol, "batch dir: "+err.Error())
401 + }
402 + }
403 +
404 + return hdr, payload, nil
405 + }
406 +
407 + // Chunked: first packet has partial payload
408 + firstPayloadBytes := n - protocol.HeaderSize
409 +
410 + // Ensure recv buffer is large enough
411 + needed := int(hdr.PayloadLen)
412 + if len(s.recvBuf) < needed {
413 + s.recvBuf = make([]byte, needed)
414 + }
415 +
416 + // Copy first chunk's payload
417 + copy(s.recvBuf[:firstPayloadBytes], buf[protocol.HeaderSize:protocol.HeaderSize+firstPayloadBytes])
418 +
419 + assembled := firstPayloadBytes
420 + chunkPayloadBudget := int(s.PacketSize) - protocol.HeaderSize
421 +
422 + // Expected chunk count
423 + remainingAfterFirst := int(hdr.PayloadLen) - firstPayloadBytes
424 + expectedContinuations := uint32(0)
425 + if remainingAfterFirst > 0 && chunkPayloadBudget > 0 {
426 + expectedContinuations = uint32((remainingAfterFirst + chunkPayloadBudget - 1) / chunkPayloadBudget)
427 + }
428 + expectedChunkCount := 1 + expectedContinuations
429 +
430 + // Temporary buffer for continuation packets
431 + pktBuf := ensureScratchBuf(&s.pktBuf, int(s.PacketSize))
432 +
433 + ci := uint32(1)
434 + for assembled < int(hdr.PayloadLen) {
435 + cn, err := rawRecv(s.fd, pktBuf)
436 + if err != nil {
437 + if errors.Is(err, ErrRecv) {
438 + s.failAllInflight()
439 + }
440 + return protocol.Header{}, nil, wrapErr(ErrRecv, "continuation recv: "+err.Error())
441 + }
442 +
443 + if cn < protocol.HeaderSize {
444 + return protocol.Header{}, nil, wrapErr(ErrChunk, "continuation too short")
445 + }
446 +
447 + chk, err := protocol.DecodeChunkHeader(pktBuf[:cn])
448 + if err != nil {
449 + return protocol.Header{}, nil, wrapErr(ErrChunk, "chunk header: "+err.Error())
450 + }
451 +
452 + // Validate chunk header
453 + if chk.MessageID != hdr.MessageID {
454 + return protocol.Header{}, nil, wrapErr(ErrChunk, "message_id mismatch")
455 + }
456 + if chk.ChunkIndex != ci {
457 + return protocol.Header{}, nil, wrapErr(ErrChunk, fmt.Sprintf(
458 + "chunk_index mismatch: expected %d, got %d", ci, chk.ChunkIndex))
459 + }
460 + if chk.ChunkCount != expectedChunkCount {
461 + return protocol.Header{}, nil, wrapErr(ErrChunk, "chunk_count mismatch")
462 + }
463 + if chk.TotalMessageLen != uint32(totalMsg) {
464 + return protocol.Header{}, nil, wrapErr(ErrChunk, "total_message_len mismatch")
465 + }
466 +
467 + chunkData := cn - protocol.HeaderSize
468 + if chunkData != int(chk.ChunkPayloadLen) {
469 + return protocol.Header{}, nil, wrapErr(ErrChunk, "chunk_payload_len mismatch")
470 + }
471 + if assembled+chunkData > int(hdr.PayloadLen) {
472 + return protocol.Header{}, nil, wrapErr(ErrChunk, "chunk exceeds payload_len")
473 + }
474 +
475 + copy(s.recvBuf[assembled:assembled+chunkData], pktBuf[protocol.HeaderSize:protocol.HeaderSize+chunkData])
476 + assembled += chunkData
477 + ci++
478 + }
479 +
480 + payload := s.recvBuf[:hdr.PayloadLen]
481 +
482 + // Validate batch directory
483 + if hdr.Flags&protocol.FlagBatch != 0 && hdr.ItemCount > 1 {
484 + dirBytes := int(hdr.ItemCount) * 8
485 + dirAligned := protocol.Align8(dirBytes)
486 + if len(payload) < dirAligned {
487 + return protocol.Header{}, nil, wrapErr(ErrProtocol, "batch dir exceeds payload")
488 + }
489 + packedAreaLen := uint32(len(payload) - dirAligned)
490 + if err := protocol.BatchDirValidate(payload[:dirBytes], hdr.ItemCount, packedAreaLen); err != nil {
491 + return protocol.Header{}, nil, wrapErr(ErrProtocol, "batch dir: "+err.Error())
492 + }
493 + }
494 +
495 + return hdr, payload, nil
496 +}
497 +
498 +// ---------------------------------------------------------------------------
499 +// Listener
500 +// ---------------------------------------------------------------------------
501 +
502 +// Listener is a listening UDS SEQPACKET endpoint.
503 +type Listener struct {
504 + fd int
505 + config ServerConfig
506 + path string
507 + nextSessionID atomic.Uint64
508 +}
509 +
510 +// Listen creates a listener on {runDir}/{serviceName}.sock.
511 +// Performs stale endpoint recovery.
512 +func Listen(runDir, serviceName string, config ServerConfig) (*Listener, error) {
513 + path, err := buildSocketPath(runDir, serviceName)
514 + if err != nil {
515 + return nil, err
516 + }
517 +
518 + // Stale recovery
519 + stale := checkAndRecoverStale(path)
520 + if stale == staleLiveServer {
521 + return nil, ErrAddrInUse
522 + }
523 +
524 + fd, err := syscall.Socket(syscall.AF_UNIX, syscall.SOCK_SEQPACKET, 0)
525 + if err != nil {
526 + return nil, wrapErr(ErrSocket, err.Error())
527 + }
528 +
529 + // Bind
530 + sa := &syscall.SockaddrUnix{Name: path}
531 + if err := syscall.Bind(fd, sa); err != nil {
532 + syscall.Close(fd)
533 + return nil, wrapErr(ErrSocket, "bind: "+err.Error())
534 + }
535 +
536 + backlog := config.Backlog
537 + if backlog <= 0 {
538 + backlog = defaultBacklog
539 + }
540 +
541 + if err := syscall.Listen(fd, backlog); err != nil {
542 + syscall.Close(fd)
543 + os.Remove(path)
544 + return nil, wrapErr(ErrSocket, "listen: "+err.Error())
545 + }
546 +
547 + return &Listener{
548 + fd: fd,
549 + config: config,
550 + path: path,
551 + }, nil
552 +}
553 +
554 +// Fd returns the raw file descriptor for poll/epoll integration.
555 +func (l *Listener) Fd() int {
556 + return l.fd
557 +}
558 +
559 +// SetPayloadLimits updates the payload limits used for future handshakes.
560 +func (l *Listener) SetPayloadLimits(maxRequestPayloadBytes, maxResponsePayloadBytes uint32) {
561 + l.config.MaxRequestPayloadBytes = maxRequestPayloadBytes
562 + l.config.MaxResponsePayloadBytes = maxResponsePayloadBytes
563 +}
564 +
565 +// Accept accepts one client connection. Performs the full handshake.
566 +// Blocks until a client connects and the handshake completes.
567 +func (l *Listener) Accept() (*Session, error) {
568 + sessionID := l.nextSessionID.Add(1)
569 + return l.AcceptWithConfig(sessionID, l.config)
570 +}
571 +
572 +// AcceptWithConfig accepts one client connection using a caller-provided
573 +// per-session server config and session ID.
574 +func (l *Listener) AcceptWithConfig(sessionID uint64, config ServerConfig) (*Session, error) {
575 + nfd, _, err := syscall.Accept(l.fd)
576 + if err != nil {
577 + return nil, wrapErr(ErrAccept, err.Error())
578 + }
579 +
580 + session, herr := serverHandshake(nfd, &config, sessionID)
581 + if herr != nil {
582 + syscall.Close(nfd)
583 + return nil, herr
584 + }
585 + return session, nil
586 +}
587 +
588 +// Close closes the listener, stops accepting, and unlinks the socket file.
589 +func (l *Listener) Close() {
590 + if l.fd >= 0 {
591 + syscall.Close(l.fd)
592 + l.fd = -1
593 + }
594 + if l.path != "" {
595 + os.Remove(l.path)
596 + l.path = ""
597 + }
598 +}
599 +
600 +// ---------------------------------------------------------------------------
601 +// Internal helpers
602 +// ---------------------------------------------------------------------------
603 +
604 +// validateServiceName checks that name contains only [a-zA-Z0-9._-],
605 +// is non-empty, and is not "." or "..".
606 +func validateServiceName(name string) error {
607 + if name == "" {
608 + return wrapErr(ErrBadParam, "empty service name")
609 + }
610 + if name == "." || name == ".." {
611 + return wrapErr(ErrBadParam, "service name cannot be '.' or '..'")
612 + }
613 + for i := 0; i < len(name); i++ {
614 + c := name[i]
615 + if (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') ||
616 + (c >= '0' && c <= '9') || c == '.' || c == '_' || c == '-' {
617 + continue
618 + }
619 + return wrapErr(ErrBadParam, fmt.Sprintf("service name contains invalid character: %q", c))
620 + }
621 + return nil
622 +}
623 +
624 +// buildSocketPath constructs {runDir}/{serviceName}.sock and validates length.
625 +func buildSocketPath(runDir, serviceName string) (string, error) {
626 + if err := validateServiceName(serviceName); err != nil {
627 + return "", err
628 + }
629 + path := filepath.Join(runDir, serviceName+".sock")
630 + // sun_path limit check. On Linux it's 108, macOS/FreeBSD 104.
631 + // We use the smaller value for portability.
632 + if len(path) >= maxSunPath {
633 + return "", ErrPathTooLong
634 + }
635 + return path, nil
636 +}
637 +
638 +// detectPacketSize reads SO_SNDBUF from the socket.
639 +func detectPacketSize(fd int) uint32 {
640 + val, err := syscall.GetsockoptInt(fd, syscall.SOL_SOCKET, syscall.SO_SNDBUF)
641 + if err != nil || val <= 0 {
642 + return defaultPacketSizeFallback
643 + }
644 + return uint32(val)
645 +}
646 +
647 +// highestBit returns the highest set bit in a bitmask (0 if empty).
648 +func highestBit(mask uint32) uint32 {
649 + if mask == 0 {
650 + return 0
651 + }
652 + bit := uint32(1) << 31
653 + for bit&mask == 0 {
654 + bit >>= 1
655 + }
656 + return bit
657 +}
658 +
659 +func applyDefault(val, def uint32) uint32 {
660 + if val == 0 {
661 + return def
662 + }
663 + return val
664 +}
665 +
666 +func minU32(a, b uint32) uint32 {
667 + if a < b {
668 + return a
669 + }
670 + return b
671 +}
672 +
673 +func maxU32(a, b uint32) uint32 {
674 + if a > b {
675 + return a
676 + }
677 + return b
678 +}
679 +
680 +// ---------------------------------------------------------------------------
681 +// Low-level I/O
682 +// ---------------------------------------------------------------------------
683 +
684 +// rawSendIov sends header + payload as one SEQPACKET message using sendmsg.
685 +func rawSendIov(fd int, hdr []byte, payload []byte) error {
686 + var iov [2]syscall.Iovec
687 + iov[0].Base = unsafe.SliceData(hdr)
688 + iov[0].SetLen(len(hdr))
689 +
690 + iovlen := uint64(1)
691 + if len(payload) > 0 {
692 + iov[1].Base = unsafe.SliceData(payload)
693 + iov[1].SetLen(len(payload))
694 + iovlen = 2
695 + }
696 +
697 + msg := syscall.Msghdr{
698 + Iov: &iov[0],
699 + Iovlen: iovlen,
700 + }
701 +
702 + n, _, errno := syscall.Syscall(
703 + syscall.SYS_SENDMSG,
704 + uintptr(fd),
705 + uintptr(unsafe.Pointer(&msg)),
706 + uintptr(syscall.MSG_NOSIGNAL),
707 + )
708 + if errno != 0 {
709 + return wrapErr(ErrSend, errno.Error())
710 + }
711 +
712 + expected := len(hdr) + len(payload)
713 + if int(n) != expected {
714 + return wrapErr(ErrSend, fmt.Sprintf("short write: %d/%d", n, expected))
715 + }
716 + return nil
717 +}
718 +
719 +func ensureScratchBuf(buf *[]byte, needed int) []byte {
720 + if len(*buf) < needed {
721 + *buf = make([]byte, needed)
722 + }
723 + return (*buf)[:needed]
724 +}
725 +
726 +// rawRecv receives one SEQPACKET message. Returns bytes received.
727 +func rawRecv(fd int, buf []byte) (int, error) {
728 + n, _, _, _, err := syscall.Recvmsg(fd, buf, nil, 0)
729 + if err != nil {
730 + return 0, wrapErr(ErrRecv, err.Error())
731 + }
732 + if n == 0 {
733 + return 0, wrapErr(ErrRecv, "peer disconnected")
734 + }
735 + return n, nil
736 +}
737 +
738 +// ---------------------------------------------------------------------------
739 +// Stale endpoint recovery
740 +// ---------------------------------------------------------------------------
741 +
742 +type staleResult int
743 +
744 +const (
745 + staleNotExist staleResult = 0
746 + staleRecovered staleResult = 1
747 + staleLiveServer staleResult = 2
748 +)
749 +
750 +func checkAndRecoverStale(path string) staleResult {
751 + _, err := os.Stat(path)
752 + if err != nil {
753 + return staleNotExist
754 + }
755 +
756 + // Try connecting to see if a live server is there.
757 + // We use net.Dial instead of raw syscalls for the probe — it handles
758 + // all the sockaddr setup and is fine for a one-shot connectivity test.
759 + conn, err := net.Dial("unixpacket", path)
760 + if err == nil {
761 + // Connected => live server
762 + conn.Close()
763 + return staleLiveServer
764 + }
765 +
766 + // Only unlink on connection-refused (stale socket).
767 + // Other errors (EACCES, etc.) should not remove the file.
768 + if errors.Is(err, syscall.ECONNREFUSED) || errors.Is(err, syscall.ENOENT) {
769 + os.Remove(path)
770 + return staleRecovered
771 + }
772 + // Can't determine ownership — treat as live to prevent overwriting
773 + return staleLiveServer
774 +}
775 +
776 +// ---------------------------------------------------------------------------
777 +// Handshake: client side
778 +// ---------------------------------------------------------------------------
779 +
780 +func connectAndHandshake(fd int, path string, config *ClientConfig) (*Session, error) {
781 + // Connect
782 + sa := &syscall.SockaddrUnix{Name: path}
783 + if err := syscall.Connect(fd, sa); err != nil {
784 + return nil, wrapErr(ErrConnect, err.Error())
785 + }
786 +
787 + pktSize := config.PacketSize
788 + if pktSize == 0 {
789 + pktSize = detectPacketSize(fd)
790 + }
791 +
792 + supported := config.SupportedProfiles
793 + if supported == 0 {
794 + supported = protocol.ProfileBaseline
795 + }
796 +
797 + // Build HELLO
798 + hello := protocol.Hello{
799 + LayoutVersion: 1,
800 + Flags: 0,
801 + SupportedProfiles: supported,
802 + PreferredProfiles: config.PreferredProfiles,
803 + MaxRequestPayloadBytes: applyDefault(config.MaxRequestPayloadBytes, protocol.MaxPayloadDefault),
804 + MaxRequestBatchItems: applyDefault(config.MaxRequestBatchItems, defaultBatchItems),
805 + MaxResponsePayloadBytes: applyDefault(config.MaxResponsePayloadBytes, protocol.MaxPayloadDefault),
806 + MaxResponseBatchItems: applyDefault(config.MaxResponseBatchItems, defaultBatchItems),
807 + AuthToken: config.AuthToken,
808 + PacketSize: pktSize,
809 + }
810 +
811 + var helloBuf [helloPayloadSize]byte
812 + hello.Encode(helloBuf[:])
813 +
814 + // Build outer CONTROL header
815 + hdr := protocol.Header{
816 + Magic: protocol.MagicMsg,
817 + Version: protocol.Version,
818 + HeaderLen: protocol.HeaderLen,
819 + Kind: protocol.KindControl,
820 + Flags: 0,
821 + Code: protocol.CodeHello,
822 + TransportStatus: protocol.StatusOK,
823 + PayloadLen: helloPayloadSize,
824 + ItemCount: 1,
825 + MessageID: 0,
826 + }
827 +
828 + var pkt [protocol.HeaderSize + helloPayloadSize]byte
829 + hdr.Encode(pkt[:protocol.HeaderSize])
830 + copy(pkt[protocol.HeaderSize:], helloBuf[:])
831 +
832 + // Send HELLO
833 + n, err := syscall.SendmsgN(fd, pkt[:], nil, nil, 0)
834 + if err != nil {
835 + return nil, wrapErr(ErrSend, "hello send: "+err.Error())
836 + }
837 + if n != len(pkt) {
838 + return nil, wrapErr(ErrSend, "hello short write")
839 + }
840 +
841 + // Receive HELLO_ACK
842 + var ackBuf [128]byte
843 + an, _, _, _, err := syscall.Recvmsg(fd, ackBuf[:], nil, 0)
844 + if err != nil {
845 + return nil, wrapErr(ErrRecv, "hello_ack recv: "+err.Error())
846 + }
847 + if an == 0 {
848 + return nil, wrapErr(ErrRecv, "peer disconnected during handshake")
849 + }
850 +
851 + // Decode outer header
852 + ackHdr, err := protocol.DecodeHeader(ackBuf[:an])
853 + if err != nil {
854 + if errors.Is(err, protocol.ErrBadVersion) {
855 + return nil, wrapErr(ErrIncompatible, "ack header version mismatch")
856 + }
857 + return nil, wrapErr(ErrProtocol, "ack header: "+err.Error())
858 + }
859 +
860 + if ackHdr.Kind != protocol.KindControl || ackHdr.Code != protocol.CodeHelloAck {
861 + return nil, wrapErr(ErrProtocol, "expected HELLO_ACK")
862 + }
863 +
864 + // Check transport_status for rejection
865 + if ackHdr.TransportStatus == protocol.StatusAuthFailed {
866 + return nil, ErrAuthFailed
867 + }
868 + if ackHdr.TransportStatus == protocol.StatusUnsupported {
869 + return nil, ErrNoProfile
870 + }
871 + if ackHdr.TransportStatus == protocol.StatusIncompatible {
872 + return nil, ErrIncompatible
873 + }
874 + if ackHdr.TransportStatus == protocol.StatusLimitExceeded {
875 + return nil, ErrLimitExceeded
876 + }
877 + if ackHdr.TransportStatus != protocol.StatusOK {
878 + return nil, wrapErr(ErrHandshake, fmt.Sprintf("transport_status=%d", ackHdr.TransportStatus))
879 + }
880 +
881 + // Decode hello-ack payload
882 + if an < protocol.HeaderSize+helloAckPayloadSize {
883 + return nil, wrapErr(ErrProtocol, "ack payload truncated")
884 + }
885 + ack, err := protocol.DecodeHelloAck(ackBuf[protocol.HeaderSize:an])
886 + if err != nil {
887 + if errors.Is(err, protocol.ErrBadLayout) &&
888 + helloAckLayoutIncompatible(ackBuf[protocol.HeaderSize:an]) {
889 + return nil, wrapErr(ErrIncompatible, "ack payload layout version mismatch")
890 + }
891 + return nil, wrapErr(ErrProtocol, "ack payload: "+err.Error())
892 + }
893 +
894 + return &Session{
895 + fd: fd,
896 + role: RoleClient,
897 + MaxRequestPayloadBytes: ack.AgreedMaxRequestPayloadBytes,
898 + MaxRequestBatchItems: ack.AgreedMaxRequestBatchItems,
899 + MaxResponsePayloadBytes: ack.AgreedMaxResponsePayloadBytes,
900 + MaxResponseBatchItems: ack.AgreedMaxResponseBatchItems,
901 + PacketSize: ack.AgreedPacketSize,
902 + SelectedProfile: ack.SelectedProfile,
903 + SessionID: ack.SessionID,
904 + inflightIDs: make(map[uint64]struct{}),
905 + }, nil
906 +}
907 +
908 +// ---------------------------------------------------------------------------
909 +// Handshake: server side
910 +// ---------------------------------------------------------------------------
911 +
912 +func serverHandshake(fd int, config *ServerConfig, sessionID uint64) (*Session, error) {
913 + serverPktSize := config.PacketSize
914 + if serverPktSize == 0 {
915 + serverPktSize = detectPacketSize(fd)
916 + }
917 +
918 + sRespPay := applyDefault(config.MaxResponsePayloadBytes, protocol.MaxPayloadDefault)
919 + sProfiles := config.SupportedProfiles
920 + if sProfiles == 0 {
921 + sProfiles = protocol.ProfileBaseline
922 + }
923 + sPreferred := config.PreferredProfiles
924 +
925 + // Helper: send rejection ACK
926 + sendRejection := func(status uint16) {
927 + ack := protocol.HelloAck{LayoutVersion: 1}
928 + var ackPayBuf [helloAckPayloadSize]byte
929 + ack.Encode(ackPayBuf[:])
930 +
931 + ackHdr := protocol.Header{
932 + Magic: protocol.MagicMsg,
933 + Version: protocol.Version,
934 + HeaderLen: protocol.HeaderLen,
935 + Kind: protocol.KindControl,
936 + Code: protocol.CodeHelloAck,
937 + TransportStatus: status,
938 + PayloadLen: helloAckPayloadSize,
939 + ItemCount: 1,
940 + }
941 +
942 + var pkt [protocol.HeaderSize + helloAckPayloadSize]byte
943 + ackHdr.Encode(pkt[:protocol.HeaderSize])
944 + copy(pkt[protocol.HeaderSize:], ackPayBuf[:])
945 + // Best effort send
946 + syscall.SendmsgN(fd, pkt[:], nil, nil, 0) //nolint:errcheck
947 + }
948 +
949 + // Receive HELLO
950 + var buf [128]byte
951 + n, _, _, _, err := syscall.Recvmsg(fd, buf[:], nil, 0)
952 + if err != nil {
953 + return nil, wrapErr(ErrRecv, "hello recv: "+err.Error())
954 + }
955 + if n == 0 {
956 + return nil, wrapErr(ErrRecv, "peer disconnected during handshake")
957 + }
958 +
959 + hdr, err := protocol.DecodeHeader(buf[:n])
960 + if err != nil {
961 + if errors.Is(err, protocol.ErrBadVersion) &&
962 + headerVersionIncompatible(buf[:n], protocol.CodeHello) {
963 + sendRejection(protocol.StatusIncompatible)
964 + return nil, ErrIncompatible
965 + }
966 + return nil, wrapErr(ErrProtocol, "hello header: "+err.Error())
967 + }
968 +
969 + if hdr.Kind != protocol.KindControl || hdr.Code != protocol.CodeHello {
970 + return nil, wrapErr(ErrProtocol, "expected HELLO")
971 + }
972 +
973 + hello, err := protocol.DecodeHello(buf[protocol.HeaderSize:n])
974 + if err != nil {
975 + if errors.Is(err, protocol.ErrBadLayout) &&
976 + helloLayoutIncompatible(buf[protocol.HeaderSize:n]) {
977 + sendRejection(protocol.StatusIncompatible)
978 + return nil, ErrIncompatible
979 + }
980 + return nil, wrapErr(ErrProtocol, "hello payload: "+err.Error())
981 + }
982 +
983 + // Compute intersection
984 + intersection := hello.SupportedProfiles & sProfiles
985 +
986 + // Check intersection
987 + if intersection == 0 {
988 + sendRejection(protocol.StatusUnsupported)
989 + return nil, ErrNoProfile
990 + }
991 +
992 + // Check auth
993 + if hello.AuthToken != config.AuthToken {
994 + sendRejection(protocol.StatusAuthFailed)
995 + return nil, ErrAuthFailed
996 + }
997 +
998 + // Select profile: prefer preferred_intersection, then intersection
999 + preferredIntersection := intersection & hello.PreferredProfiles & sPreferred
1000 + var selected uint32
1001 + if preferredIntersection != 0 {
1002 + selected = highestBit(preferredIntersection)
1003 + } else {
1004 + selected = highestBit(intersection)
1005 + }
1006 +
1007 + if hello.MaxRequestPayloadBytes > protocol.MaxPayloadCap {
1008 + sendRejection(protocol.StatusLimitExceeded)
1009 + return nil, ErrLimitExceeded
1010 + }
1011 +
1012 + // Negotiate limits:
1013 + // - request payload and batch size are client-proposed and echoed
1014 + // - response payload is server-authoritative
1015 + // - response batch size is symmetric with request batch size
1016 + agreedReqPay := hello.MaxRequestPayloadBytes
1017 + agreedReqBat := hello.MaxRequestBatchItems
1018 + agreedRespPay := sRespPay
1019 + agreedRespBat := agreedReqBat
1020 + agreedPkt := minU32(hello.PacketSize, serverPktSize)
1021 + if agreedPkt <= protocol.HeaderSize {
1022 + sendRejection(protocol.StatusIncompatible)
1023 + return nil, ErrIncompatible
1024 + }
1025 +
1026 + // Send HELLO_ACK (success)
1027 + ack := protocol.HelloAck{
1028 + LayoutVersion: 1,
1029 + Flags: 0,
1030 + ServerSupportedProfiles: sProfiles,
1031 + IntersectionProfiles: intersection,
1032 + SelectedProfile: selected,
1033 + AgreedMaxRequestPayloadBytes: agreedReqPay,
1034 + AgreedMaxRequestBatchItems: agreedReqBat,
1035 + AgreedMaxResponsePayloadBytes: agreedRespPay,
1036 + AgreedMaxResponseBatchItems: agreedRespBat,
1037 + AgreedPacketSize: agreedPkt,
1038 + SessionID: sessionID,
1039 + }
1040 +
1041 + var ackPayBuf [helloAckPayloadSize]byte
1042 + ack.Encode(ackPayBuf[:])
1043 +
1044 + ackHdr := protocol.Header{
1045 + Magic: protocol.MagicMsg,
1046 + Version: protocol.Version,
1047 + HeaderLen: protocol.HeaderLen,
1048 + Kind: protocol.KindControl,
1049 + Code: protocol.CodeHelloAck,
1050 + TransportStatus: protocol.StatusOK,
1051 + PayloadLen: helloAckPayloadSize,
1052 + ItemCount: 1,
1053 + }
1054 +
1055 + var pkt [protocol.HeaderSize + helloAckPayloadSize]byte
1056 + ackHdr.Encode(pkt[:protocol.HeaderSize])
1057 + copy(pkt[protocol.HeaderSize:], ackPayBuf[:])
1058 +
1059 + sn, err := syscall.SendmsgN(fd, pkt[:], nil, nil, 0)
1060 + if err != nil {
1061 + return nil, wrapErr(ErrSend, "hello_ack send: "+err.Error())
1062 + }
1063 + if sn != len(pkt) {
1064 + return nil, wrapErr(ErrSend, "hello_ack short write")
1065 + }
1066 +
1067 + return &Session{
1068 + fd: fd,
1069 + role: RoleServer,
1070 + MaxRequestPayloadBytes: agreedReqPay,
1071 + MaxRequestBatchItems: agreedReqBat,
1072 + MaxResponsePayloadBytes: agreedRespPay,
1073 + MaxResponseBatchItems: agreedRespBat,
1074 + PacketSize: agreedPkt,
1075 + SelectedProfile: selected,
1076 + SessionID: sessionID,
1077 + inflightIDs: make(map[uint64]struct{}),
1078 + }, nil
1079 +}
src/go/pkg/netipc/transport/posix/uds_edge_test.go new
+305
@@ -0,0 +1,305 @@
1 +//go:build unix
2 +
3 +package posix
4 +
5 +import (
6 + "errors"
7 + "os"
8 + "path/filepath"
9 + "testing"
10 +
11 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
12 +)
13 +
14 +// ---------------------------------------------------------------------------
15 +// Receive: payload_len exceeds negotiated limit
16 +// ---------------------------------------------------------------------------
17 +
18 +func TestReceivePayloadExceedsLimit(t *testing.T) {
19 + runDir := testRunDir(t)
20 + service := uniqueService(t)
21 + defer os.Remove(filepath.Join(runDir, service+".sock"))
22 +
23 + // Server allows up to 128 bytes request payload
24 + sCfg := defaultServerConfig()
25 + sCfg.MaxRequestPayloadBytes = 128
26 + sCfg.MaxResponsePayloadBytes = 128
27 + listener := startListener(t, runDir, service, sCfg)
28 + defer listener.Close()
29 +
30 + acceptCh := acceptAsync(listener)
31 +
32 + // Client also advertises 128 so negotiation settles on 128
33 + cCfg := defaultClientConfig()
34 + cCfg.MaxRequestPayloadBytes = 128
35 + cCfg.MaxResponsePayloadBytes = 128
36 + client, err := Connect(runDir, service, &cCfg)
37 + if err != nil {
38 + t.Fatalf("Connect: %v", err)
39 + }
40 + defer client.Close()
41 +
42 + sr := <-acceptCh
43 + if sr.err != nil {
44 + t.Fatalf("Accept: %v", sr.err)
45 + }
46 + server := sr.session
47 + defer server.Close()
48 +
49 + // Client sends request with message_id=1, server receives OK
50 + hdr := protocol.Header{
51 + Kind: protocol.KindRequest,
52 + Code: protocol.MethodIncrement,
53 + ItemCount: 1,
54 + MessageID: 1,
55 + }
56 + if err := client.Send(&hdr, []byte("ok")); err != nil {
57 + t.Fatalf("Send: %v", err)
58 + }
59 +
60 + buf := make([]byte, 4096)
61 + _, _, err = server.Receive(buf)
62 + if err != nil {
63 + t.Fatalf("server Receive: %v", err)
64 + }
65 +
66 + // Server sends response with artificially large payload that exceeds
67 + // the client's negotiated limit. We need to forge this via raw send
68 + // since session.Send doesn't validate outgoing size.
69 + resp := protocol.Header{
70 + Kind: protocol.KindResponse,
71 + Code: protocol.MethodIncrement,
72 + ItemCount: 1,
73 + MessageID: 1,
74 + }
75 + // Forge a payload > 128 bytes
76 + bigPayload := make([]byte, 200)
77 + if err := server.Send(&resp, bigPayload); err != nil {
78 + t.Fatalf("server Send: %v", err)
79 + }
80 +
81 + // Client receive should fail with limit exceeded
82 + _, _, err = client.Receive(buf)
83 + if err == nil {
84 + t.Fatal("expected limit exceeded error")
85 + }
86 + if !errors.Is(err, ErrLimitExceeded) {
87 + t.Errorf("error = %v, want ErrLimitExceeded", err)
88 + }
89 +}
90 +
91 +// ---------------------------------------------------------------------------
92 +// Receive: item_count exceeds negotiated batch limit
93 +// ---------------------------------------------------------------------------
94 +
95 +func TestReceiveBatchExceedsLimit(t *testing.T) {
96 + runDir := testRunDir(t)
97 + service := uniqueService(t)
98 + defer os.Remove(filepath.Join(runDir, service+".sock"))
99 +
100 + // Server allows batch items = 2
101 + sCfg := defaultServerConfig()
102 + sCfg.MaxRequestBatchItems = 2
103 + sCfg.MaxResponseBatchItems = 2
104 + listener := startListener(t, runDir, service, sCfg)
105 + defer listener.Close()
106 +
107 + acceptCh := acceptAsync(listener)
108 +
109 + cCfg := defaultClientConfig()
110 + cCfg.MaxRequestBatchItems = 2
111 + cCfg.MaxResponseBatchItems = 2
112 + client, err := Connect(runDir, service, &cCfg)
113 + if err != nil {
114 + t.Fatalf("Connect: %v", err)
115 + }
116 + defer client.Close()
117 +
118 + sr := <-acceptCh
119 + if sr.err != nil {
120 + t.Fatalf("Accept: %v", sr.err)
121 + }
122 + server := sr.session
123 + defer server.Close()
124 +
125 + // Client sends a valid request
126 + hdr := protocol.Header{
127 + Kind: protocol.KindRequest,
128 + Code: protocol.MethodIncrement,
129 + ItemCount: 1,
130 + MessageID: 1,
131 + }
132 + if err := client.Send(&hdr, []byte("x")); err != nil {
133 + t.Fatalf("Send: %v", err)
134 + }
135 +
136 + buf := make([]byte, 4096)
137 + _, _, err = server.Receive(buf)
138 + if err != nil {
139 + t.Fatalf("server Receive: %v", err)
140 + }
141 +
142 + // Server responds with item_count=10 (exceeds negotiated 2)
143 + resp := protocol.Header{
144 + Kind: protocol.KindResponse,
145 + Code: protocol.MethodIncrement,
146 + ItemCount: 10,
147 + MessageID: 1,
148 + }
149 + if err := server.Send(&resp, []byte("y")); err != nil {
150 + t.Fatalf("server Send: %v", err)
151 + }
152 +
153 + // Client receive should fail
154 + _, _, err = client.Receive(buf)
155 + if err == nil {
156 + t.Fatal("expected limit exceeded error for batch items")
157 + }
158 + if !errors.Is(err, ErrLimitExceeded) {
159 + t.Errorf("error = %v, want ErrLimitExceeded", err)
160 + }
161 +}
162 +
163 +// ---------------------------------------------------------------------------
164 +// Handshake: client gets non-OK transport_status (generic)
165 +// ---------------------------------------------------------------------------
166 +
167 +func TestWrapErr(t *testing.T) {
168 + // Verify wrapErr produces a properly wrapped error
169 + err := wrapErr(ErrConnect, "some detail")
170 + if !errors.Is(err, ErrConnect) {
171 + t.Error("wrapErr result should wrap ErrConnect")
172 + }
173 + if err.Error() != "connect failed: some detail" {
174 + t.Errorf("unexpected error string: %q", err.Error())
175 + }
176 +
177 + err2 := wrapErr(ErrAuthFailed, "bad token")
178 + if !errors.Is(err2, ErrAuthFailed) {
179 + t.Error("wrapErr result should wrap ErrAuthFailed")
180 + }
181 +}
182 +
183 +// ---------------------------------------------------------------------------
184 +// validateServiceName edge cases
185 +// ---------------------------------------------------------------------------
186 +
187 +func TestValidateServiceNameEdgeCases(t *testing.T) {
188 + // Valid names
189 + valid := []string{"a", "A", "0", "test-service", "test.service", "test_service", "Test123"}
190 + for _, name := range valid {
191 + if err := validateServiceName(name); err != nil {
192 + t.Errorf("validateServiceName(%q) = %v, want nil", name, err)
193 + }
194 + }
195 +
196 + // Invalid names
197 + invalid := []string{"", ".", "..", "a/b", "a b", "a\tb", "a@b", "a#b", "a$b", "a!b"}
198 + for _, name := range invalid {
199 + if err := validateServiceName(name); err == nil {
200 + t.Errorf("validateServiceName(%q) = nil, want error", name)
201 + }
202 + }
203 +}
204 +
205 +// ---------------------------------------------------------------------------
206 +// buildSocketPath edge cases
207 +// ---------------------------------------------------------------------------
208 +
209 +func TestBuildSocketPathEdgeCases(t *testing.T) {
210 + // Valid case
211 + path, err := buildSocketPath("/tmp", "test")
212 + if err != nil {
213 + t.Fatalf("buildSocketPath: %v", err)
214 + }
215 + expected := filepath.Join("/tmp", "test.sock")
216 + if path != expected {
217 + t.Fatalf("path = %q, want %q", path, expected)
218 + }
219 +
220 + // Invalid service name
221 + _, err = buildSocketPath("/tmp", "")
222 + if err == nil {
223 + t.Fatal("expected error for empty service name")
224 + }
225 +
226 + // Path too long
227 + longDir := "/tmp/" + string(make([]byte, 200))
228 + _, err = buildSocketPath(longDir, "test")
229 + if err != ErrPathTooLong {
230 + t.Fatalf("expected ErrPathTooLong, got %v", err)
231 + }
232 +}
233 +
234 +// ---------------------------------------------------------------------------
235 +// applyDefault and minU32 helpers
236 +// ---------------------------------------------------------------------------
237 +
238 +func TestApplyDefault(t *testing.T) {
239 + if got := applyDefault(0, 42); got != 42 {
240 + t.Errorf("applyDefault(0, 42) = %d, want 42", got)
241 + }
242 + if got := applyDefault(10, 42); got != 10 {
243 + t.Errorf("applyDefault(10, 42) = %d, want 10", got)
244 + }
245 +}
246 +
247 +func TestMinU32(t *testing.T) {
248 + if got := minU32(5, 10); got != 5 {
249 + t.Errorf("minU32(5, 10) = %d, want 5", got)
250 + }
251 + if got := minU32(10, 5); got != 5 {
252 + t.Errorf("minU32(10, 5) = %d, want 5", got)
253 + }
254 + if got := minU32(7, 7); got != 7 {
255 + t.Errorf("minU32(7, 7) = %d, want 7", got)
256 + }
257 +}
258 +
259 +// ---------------------------------------------------------------------------
260 +// Session.Close double-close safety
261 +// ---------------------------------------------------------------------------
262 +
263 +func TestSessionDoubleClose(t *testing.T) {
264 + runDir := testRunDir(t)
265 + service := uniqueService(t)
266 + defer os.Remove(filepath.Join(runDir, service+".sock"))
267 +
268 + sCfg := defaultServerConfig()
269 + listener := startListener(t, runDir, service, sCfg)
270 + defer listener.Close()
271 +
272 + acceptCh := acceptAsync(listener)
273 +
274 + cCfg := defaultClientConfig()
275 + client, err := Connect(runDir, service, &cCfg)
276 + if err != nil {
277 + t.Fatalf("Connect: %v", err)
278 + }
279 +
280 + sr := <-acceptCh
281 + if sr.err != nil {
282 + t.Fatalf("Accept: %v", sr.err)
283 + }
284 + sr.session.Close()
285 +
286 + // Double close should not panic
287 + client.Close()
288 + client.Close()
289 +}
290 +
291 +// ---------------------------------------------------------------------------
292 +// Listener.Close double-close safety
293 +// ---------------------------------------------------------------------------
294 +
295 +func TestListenerDoubleClose(t *testing.T) {
296 + runDir := testRunDir(t)
297 + service := uniqueService(t)
298 +
299 + sCfg := defaultServerConfig()
300 + listener := startListener(t, runDir, service, sCfg)
301 +
302 + // Double close should not panic
303 + listener.Close()
304 + listener.Close()
305 +}
src/go/pkg/netipc/transport/posix/uds_more_edge_test.go new
+829
@@ -0,0 +1,829 @@
1 +//go:build unix
2 +
3 +package posix
4 +
5 +import (
6 + "encoding/binary"
7 + "errors"
8 + "fmt"
9 + "os"
10 + "path/filepath"
11 + "strings"
12 + "syscall"
13 + "testing"
14 +
15 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
16 +)
17 +
18 +func udsSessionPair(t *testing.T, sCfg ServerConfig, cCfg ClientConfig) (*Session, *Session) {
19 + t.Helper()
20 +
21 + runDir := testRunDir(t)
22 + service := uniqueService(t)
23 + socketPath := filepath.Join(runDir, service+".sock")
24 + t.Cleanup(func() { _ = os.Remove(socketPath) })
25 +
26 + listener := startListener(t, runDir, service, sCfg)
27 + t.Cleanup(listener.Close)
28 +
29 + acceptCh := acceptAsync(listener)
30 +
31 + client, err := Connect(runDir, service, &cCfg)
32 + if err != nil {
33 + t.Fatalf("Connect failed: %v", err)
34 + }
35 + t.Cleanup(client.Close)
36 +
37 + sr := <-acceptCh
38 + if sr.err != nil {
39 + t.Fatalf("Accept failed: %v", sr.err)
40 + }
41 + server := sr.session
42 + t.Cleanup(server.Close)
43 +
44 + return client, server
45 +}
46 +
47 +func rawSendPacket(fd int, pkt []byte) error {
48 + n, err := syscall.SendmsgN(fd, pkt, nil, nil, 0)
49 + if err != nil {
50 + return err
51 + }
52 + if n != len(pkt) {
53 + return fmt.Errorf("short write: %d/%d", n, len(pkt))
54 + }
55 + return nil
56 +}
57 +
58 +func encodePacket(hdr protocol.Header, payload []byte) []byte {
59 + hdr.Magic = protocol.MagicMsg
60 + hdr.Version = protocol.Version
61 + hdr.HeaderLen = protocol.HeaderLen
62 + hdr.PayloadLen = uint32(len(payload))
63 +
64 + pkt := make([]byte, protocol.HeaderSize+len(payload))
65 + hdr.Encode(pkt[:protocol.HeaderSize])
66 + copy(pkt[protocol.HeaderSize:], payload)
67 + return pkt
68 +}
69 +
70 +func encodePartialPacket(hdr protocol.Header, advertisedPayloadLen uint32, payload []byte) []byte {
71 + hdr.Magic = protocol.MagicMsg
72 + hdr.Version = protocol.Version
73 + hdr.HeaderLen = protocol.HeaderLen
74 + hdr.PayloadLen = advertisedPayloadLen
75 +
76 + pkt := make([]byte, protocol.HeaderSize+len(payload))
77 + hdr.Encode(pkt[:protocol.HeaderSize])
78 + copy(pkt[protocol.HeaderSize:], payload)
79 + return pkt
80 +}
81 +
82 +func encodeChunkPacket(chk protocol.ChunkHeader, payload []byte) []byte {
83 + pkt := make([]byte, protocol.HeaderSize+len(payload))
84 + chk.Encode(pkt[:protocol.HeaderSize])
85 + copy(pkt[protocol.HeaderSize:], payload)
86 + return pkt
87 +}
88 +
89 +func validHelloAckPacket(status uint16) []byte {
90 + ack := protocol.HelloAck{
91 + LayoutVersion: 1,
92 + Flags: 0,
93 + ServerSupportedProfiles: protocol.ProfileBaseline,
94 + IntersectionProfiles: protocol.ProfileBaseline,
95 + SelectedProfile: protocol.ProfileBaseline,
96 + AgreedMaxRequestPayloadBytes: protocol.MaxPayloadDefault,
97 + AgreedMaxRequestBatchItems: 1,
98 + AgreedMaxResponsePayloadBytes: protocol.MaxPayloadDefault,
99 + AgreedMaxResponseBatchItems: 1,
100 + AgreedPacketSize: defaultPacketSizeFallback,
101 + SessionID: 77,
102 + }
103 +
104 + payload := make([]byte, helloAckPayloadSize)
105 + ack.Encode(payload)
106 + return encodePacket(protocol.Header{
107 + Kind: protocol.KindControl,
108 + Code: protocol.CodeHelloAck,
109 + TransportStatus: status,
110 + ItemCount: 1,
111 + }, payload)
112 +}
113 +
114 +func validHelloPacket() []byte {
115 + hello := protocol.Hello{
116 + LayoutVersion: 1,
117 + Flags: 0,
118 + SupportedProfiles: protocol.ProfileBaseline,
119 + PreferredProfiles: protocol.ProfileBaseline,
120 + MaxRequestPayloadBytes: protocol.MaxPayloadDefault,
121 + MaxRequestBatchItems: 1,
122 + MaxResponsePayloadBytes: protocol.MaxPayloadDefault,
123 + MaxResponseBatchItems: 1,
124 + AuthToken: testAuthToken,
125 + PacketSize: defaultPacketSizeFallback,
126 + }
127 +
128 + payload := make([]byte, helloPayloadSize)
129 + hello.Encode(payload)
130 + return encodePacket(protocol.Header{
131 + Kind: protocol.KindControl,
132 + Code: protocol.CodeHello,
133 + ItemCount: 1,
134 + }, payload)
135 +}
136 +
137 +func rawSeqpacketListener(t *testing.T, runDir, service string) (int, string) {
138 + t.Helper()
139 +
140 + path, err := buildSocketPath(runDir, service)
141 + if err != nil {
142 + t.Fatalf("buildSocketPath failed: %v", err)
143 + }
144 + _ = os.Remove(path)
145 +
146 + fd, err := syscall.Socket(syscall.AF_UNIX, syscall.SOCK_SEQPACKET, 0)
147 + if err != nil {
148 + t.Fatalf("socket failed: %v", err)
149 + }
150 + t.Cleanup(func() {
151 + _ = syscall.Close(fd)
152 + _ = os.Remove(path)
153 + })
154 +
155 + if err := syscall.Bind(fd, &syscall.SockaddrUnix{Name: path}); err != nil {
156 + t.Fatalf("bind failed: %v", err)
157 + }
158 + if err := syscall.Listen(fd, defaultBacklog); err != nil {
159 + t.Fatalf("listen failed: %v", err)
160 + }
161 + return fd, path
162 +}
163 +
164 +func rawSeqpacketConnect(t *testing.T, path string) int {
165 + t.Helper()
166 +
167 + fd, err := syscall.Socket(syscall.AF_UNIX, syscall.SOCK_SEQPACKET, 0)
168 + if err != nil {
169 + t.Fatalf("socket failed: %v", err)
170 + }
171 + t.Cleanup(func() { _ = syscall.Close(fd) })
172 +
173 + if err := syscall.Connect(fd, &syscall.SockaddrUnix{Name: path}); err != nil {
174 + t.Fatalf("connect failed: %v", err)
175 + }
176 + return fd
177 +}
178 +
179 +func TestUdsAcceptOnClosedListener(t *testing.T) {
180 + runDir := testRunDir(t)
181 + service := uniqueService(t)
182 + listener := startListener(t, runDir, service, defaultServerConfig())
183 + listener.Close()
184 +
185 + session, err := listener.Accept()
186 + if session != nil {
187 + session.Close()
188 + }
189 + if !errors.Is(err, ErrAccept) {
190 + t.Fatalf("Accept after close = %v, want ErrAccept", err)
191 + }
192 +}
193 +
194 +func TestUdsListenFailsWhenRunDirMissing(t *testing.T) {
195 + runDir := filepath.Join(t.TempDir(), "missing")
196 + _, err := Listen(runDir, uniqueService(t), defaultServerConfig())
197 + if !errors.Is(err, ErrSocket) {
198 + t.Fatalf("Listen(missing runDir) = %v, want ErrSocket", err)
199 + }
200 + if !containsErrText(err, "bind:") {
201 + t.Fatalf("Listen(missing runDir) = %v, want bind failure detail", err)
202 + }
203 +}
204 +
205 +func TestUdsRawGenericErrors(t *testing.T) {
206 + if err := rawSendIov(-1, []byte("x"), nil); !errors.Is(err, ErrSend) {
207 + t.Fatalf("rawSendIov(invalid fd) = %v, want ErrSend", err)
208 + }
209 +
210 + buf := make([]byte, 16)
211 + if _, err := rawRecv(-1, buf); !errors.Is(err, ErrRecv) {
212 + t.Fatalf("rawRecv(invalid fd) = %v, want ErrRecv", err)
213 + }
214 +}
215 +
216 +func TestUdsSessionSendRejectsTooSmallPacketSize(t *testing.T) {
217 + client, _ := udsSessionPair(t, defaultServerConfig(), defaultClientConfig())
218 +
219 + client.PacketSize = uint32(protocol.HeaderSize)
220 + hdr := protocol.Header{
221 + Kind: protocol.KindRequest,
222 + Code: protocol.MethodIncrement,
223 + ItemCount: 1,
224 + MessageID: 1,
225 + }
226 + if err := client.Send(&hdr, []byte("x")); !errors.Is(err, ErrBadParam) {
227 + t.Fatalf("Send(packet_size too small) = %v, want ErrBadParam", err)
228 + }
229 +}
230 +
231 +func TestUdsSendInitializesInflightSetOnFirstRequest(t *testing.T) {
232 + client, server := udsSessionPair(t, defaultServerConfig(), defaultClientConfig())
233 +
234 + client.inflightIDs = nil
235 + payload := []byte("ping")
236 + hdr := protocol.Header{
237 + Kind: protocol.KindRequest,
238 + Code: protocol.MethodIncrement,
239 + ItemCount: 1,
240 + MessageID: 55,
241 + }
242 +
243 + if err := client.Send(&hdr, payload); err != nil {
244 + t.Fatalf("Send failed: %v", err)
245 + }
246 + if client.inflightIDs == nil {
247 + t.Fatal("Send should initialize inflightIDs for the first client request")
248 + }
249 + if _, ok := client.inflightIDs[55]; !ok {
250 + t.Fatal("Send should track the in-flight message_id")
251 + }
252 +
253 + buf := make([]byte, 4096)
254 + rHdr, rPayload, err := server.Receive(buf)
255 + if err != nil {
256 + t.Fatalf("server Receive failed: %v", err)
257 + }
258 + if rHdr.MessageID != 55 {
259 + t.Fatalf("server received message_id=%d, want 55", rHdr.MessageID)
260 + }
261 + if string(rPayload) != string(payload) {
262 + t.Fatalf("server received payload=%q, want %q", rPayload, payload)
263 + }
264 +}
265 +
266 +func TestUdsReceiveRejectsMalformedFirstPacket(t *testing.T) {
267 + cases := []struct {
268 + name string
269 + pkt []byte
270 + want error
271 + }{
272 + {
273 + name: "short packet",
274 + pkt: []byte{1, 2, 3},
275 + want: ErrProtocol,
276 + },
277 + {
278 + name: "bad header",
279 + pkt: make([]byte, protocol.HeaderSize),
280 + want: ErrProtocol,
281 + },
282 + }
283 +
284 + for _, tc := range cases {
285 + t.Run(tc.name, func(t *testing.T) {
286 + client, server := udsSessionPair(t, defaultServerConfig(), defaultClientConfig())
287 +
288 + client.inflightIDs[42] = struct{}{}
289 + if err := rawSendPacket(server.fd, tc.pkt); err != nil {
290 + t.Fatalf("rawSendPacket failed: %v", err)
291 + }
292 +
293 + buf := make([]byte, 4096)
294 + _, _, err := client.Receive(buf)
295 + if !errors.Is(err, tc.want) {
296 + t.Fatalf("Receive error = %v, want %v", err, tc.want)
297 + }
298 + })
299 + }
300 +}
301 +
302 +func TestUdsReceiveRejectsUnknownMessageIDWithNilInflightSet(t *testing.T) {
303 + client, server := udsSessionPair(t, defaultServerConfig(), defaultClientConfig())
304 +
305 + client.inflightIDs = nil
306 + respPkt := encodePacket(protocol.Header{
307 + Kind: protocol.KindResponse,
308 + Code: protocol.MethodIncrement,
309 + ItemCount: 1,
310 + MessageID: 42,
311 + }, []byte("ok"))
312 + if err := rawSendPacket(server.fd, respPkt); err != nil {
313 + t.Fatalf("rawSendPacket failed: %v", err)
314 + }
315 +
316 + buf := make([]byte, 4096)
317 + _, _, err := client.Receive(buf)
318 + if !errors.Is(err, ErrUnknownMsgID) {
319 + t.Fatalf("Receive error = %v, want ErrUnknownMsgID", err)
320 + }
321 +}
322 +
323 +func TestUdsReceiveRejectsMalformedBatchPayload(t *testing.T) {
324 + client, server := udsSessionPair(t, defaultServerConfig(), defaultClientConfig())
325 +
326 + client.inflightIDs[42] = struct{}{}
327 + payload := make([]byte, 24)
328 + binary.NativeEndian.PutUint32(payload[0:4], 1)
329 + binary.NativeEndian.PutUint32(payload[4:8], 4)
330 + binary.NativeEndian.PutUint32(payload[8:12], 0)
331 + binary.NativeEndian.PutUint32(payload[12:16], 4)
332 + copy(payload[16:], []byte("payload!!"))
333 +
334 + respPkt := encodePacket(protocol.Header{
335 + Kind: protocol.KindResponse,
336 + Code: protocol.MethodIncrement,
337 + Flags: protocol.FlagBatch,
338 + ItemCount: 2,
339 + MessageID: 42,
340 + }, payload)
341 + if err := rawSendPacket(server.fd, respPkt); err != nil {
342 + t.Fatalf("rawSendPacket failed: %v", err)
343 + }
344 +
345 + buf := make([]byte, 4096)
346 + _, _, err := client.Receive(buf)
347 + if !errors.Is(err, ErrProtocol) {
348 + t.Fatalf("Receive error = %v, want ErrProtocol", err)
349 + }
350 +}
351 +
352 +func TestUdsReceiveRejectsMalformedChunkedBatchPayload(t *testing.T) {
353 + cases := []struct {
354 + name string
355 + payload []byte
356 + first int
357 + }{
358 + {
359 + name: "batch dir exceeds payload",
360 + payload: make([]byte, 8),
361 + first: 4,
362 + },
363 + {
364 + name: "batch dir invalid after reassembly",
365 + payload: func() []byte {
366 + payload := make([]byte, 24)
367 + binary.NativeEndian.PutUint32(payload[0:4], 1)
368 + binary.NativeEndian.PutUint32(payload[4:8], 4)
369 + binary.NativeEndian.PutUint32(payload[8:12], 0)
370 + binary.NativeEndian.PutUint32(payload[12:16], 4)
371 + copy(payload[16:], []byte("payload!!"))
372 + return payload
373 + }(),
374 + first: 8,
375 + },
376 + }
377 +
378 + for _, tc := range cases {
379 + t.Run(tc.name, func(t *testing.T) {
380 + client, server := udsSessionPair(t, defaultServerConfig(), defaultClientConfig())
381 +
382 + client.inflightIDs[42] = struct{}{}
383 + firstPkt := encodePartialPacket(protocol.Header{
384 + Kind: protocol.KindResponse,
385 + Code: protocol.MethodIncrement,
386 + Flags: protocol.FlagBatch,
387 + ItemCount: 2,
388 + MessageID: 42,
389 + }, uint32(len(tc.payload)), tc.payload[:tc.first])
390 + if err := rawSendPacket(server.fd, firstPkt); err != nil {
391 + t.Fatalf("rawSendPacket(first) failed: %v", err)
392 + }
393 +
394 + second := tc.payload[tc.first:]
395 + if len(second) > 0 {
396 + secondPkt := encodeChunkPacket(protocol.ChunkHeader{
397 + Magic: protocol.MagicChunk,
398 + Version: protocol.Version,
399 + MessageID: 42,
400 + TotalMessageLen: uint32(protocol.HeaderSize + len(tc.payload)),
401 + ChunkIndex: 1,
402 + ChunkCount: 2,
403 + ChunkPayloadLen: uint32(len(second)),
404 + }, second)
405 + if err := rawSendPacket(server.fd, secondPkt); err != nil {
406 + t.Fatalf("rawSendPacket(second) failed: %v", err)
407 + }
408 + }
409 +
410 + buf := make([]byte, 4096)
411 + _, _, err := client.Receive(buf)
412 + if !errors.Is(err, ErrProtocol) {
413 + t.Fatalf("Receive error = %v, want ErrProtocol", err)
414 + }
415 + })
416 + }
417 +}
418 +
419 +func TestUdsReceiveRejectsMalformedChunks(t *testing.T) {
420 + cases := []struct {
421 + name string
422 + second []byte
423 + want error
424 + wantMsg string
425 + }{
426 + {
427 + name: "continuation too short",
428 + second: []byte{1, 2, 3},
429 + want: ErrChunk,
430 + },
431 + {
432 + name: "bad chunk header",
433 + second: encodeChunkPacket(protocol.ChunkHeader{
434 + Magic: protocol.MagicChunk,
435 + Version: protocol.Version + 1,
436 + MessageID: 42,
437 + TotalMessageLen: uint32(protocol.HeaderSize + 8),
438 + ChunkIndex: 1,
439 + ChunkCount: 2,
440 + ChunkPayloadLen: 4,
441 + }, []byte("tail")),
442 + want: ErrChunk,
443 + },
444 + {
445 + name: "message id mismatch",
446 + second: encodeChunkPacket(protocol.ChunkHeader{
447 + Magic: protocol.MagicChunk,
448 + Version: protocol.Version,
449 + MessageID: 99,
450 + TotalMessageLen: uint32(protocol.HeaderSize + 8),
451 + ChunkIndex: 1,
452 + ChunkCount: 2,
453 + ChunkPayloadLen: 4,
454 + }, []byte("tail")),
455 + want: ErrChunk,
456 + },
457 + {
458 + name: "chunk index mismatch",
459 + second: encodeChunkPacket(protocol.ChunkHeader{
460 + Magic: protocol.MagicChunk,
461 + Version: protocol.Version,
462 + MessageID: 42,
463 + TotalMessageLen: uint32(protocol.HeaderSize + 8),
464 + ChunkIndex: 2,
465 + ChunkCount: 2,
466 + ChunkPayloadLen: 4,
467 + }, []byte("tail")),
468 + want: ErrChunk,
469 + },
470 + {
471 + name: "chunk count mismatch",
472 + second: encodeChunkPacket(protocol.ChunkHeader{
473 + Magic: protocol.MagicChunk,
474 + Version: protocol.Version,
475 + MessageID: 42,
476 + TotalMessageLen: uint32(protocol.HeaderSize + 8),
477 + ChunkIndex: 1,
478 + ChunkCount: 3,
479 + ChunkPayloadLen: 4,
480 + }, []byte("tail")),
481 + want: ErrChunk,
482 + },
483 + {
484 + name: "total len mismatch",
485 + second: encodeChunkPacket(protocol.ChunkHeader{
486 + Magic: protocol.MagicChunk,
487 + Version: protocol.Version,
488 + MessageID: 42,
489 + TotalMessageLen: uint32(protocol.HeaderSize + 9),
490 + ChunkIndex: 1,
491 + ChunkCount: 2,
492 + ChunkPayloadLen: 4,
493 + }, []byte("tail")),
494 + want: ErrChunk,
495 + },
496 + {
497 + name: "chunk payload len mismatch",
498 + second: encodeChunkPacket(protocol.ChunkHeader{
499 + Magic: protocol.MagicChunk,
500 + Version: protocol.Version,
501 + MessageID: 42,
502 + TotalMessageLen: uint32(protocol.HeaderSize + 8),
503 + ChunkIndex: 1,
504 + ChunkCount: 2,
505 + ChunkPayloadLen: 3,
506 + }, []byte("tail")),
507 + want: ErrChunk,
508 + },
509 + {
510 + name: "chunk exceeds payload",
511 + second: encodeChunkPacket(protocol.ChunkHeader{
512 + Magic: protocol.MagicChunk,
513 + Version: protocol.Version,
514 + MessageID: 42,
515 + TotalMessageLen: uint32(protocol.HeaderSize + 8),
516 + ChunkIndex: 1,
517 + ChunkCount: 2,
518 + ChunkPayloadLen: 5,
519 + }, []byte("tails")),
520 + want: ErrChunk,
521 + },
522 + }
523 +
524 + for _, tc := range cases {
525 + t.Run(tc.name, func(t *testing.T) {
526 + client, server := udsSessionPair(t, defaultServerConfig(), defaultClientConfig())
527 +
528 + client.inflightIDs[42] = struct{}{}
529 + firstPkt := encodePartialPacket(protocol.Header{
530 + Kind: protocol.KindResponse,
531 + Code: protocol.MethodIncrement,
532 + ItemCount: 1,
533 + MessageID: 42,
534 + }, 8, []byte("head"))
535 + if err := rawSendPacket(server.fd, firstPkt); err != nil {
536 + t.Fatalf("rawSendPacket(first) failed: %v", err)
537 + }
538 + if err := rawSendPacket(server.fd, tc.second); err != nil {
539 + t.Fatalf("rawSendPacket(second) failed: %v", err)
540 + }
541 +
542 + buf := make([]byte, 4096)
543 + _, _, err := client.Receive(buf)
544 + if !errors.Is(err, tc.want) {
545 + t.Fatalf("Receive error = %v, want %v", err, tc.want)
546 + }
547 + })
548 + }
549 +}
550 +
551 +func TestUdsClientHandshakeRejectsMalformedAck(t *testing.T) {
552 + cases := []struct {
553 + name string
554 + packetFn func() []byte
555 + want error
556 + wantText string
557 + }{
558 + {
559 + name: "bad ack header",
560 + packetFn: func() []byte {
561 + ack := validHelloAckPacket(protocol.StatusOK)
562 + ack[0] = 0
563 + return ack
564 + },
565 + want: ErrProtocol,
566 + wantText: "ack header",
567 + },
568 + {
569 + name: "wrong kind",
570 + packetFn: func() []byte {
571 + payload := make([]byte, helloAckPayloadSize)
572 + return encodePacket(protocol.Header{
573 + Kind: protocol.KindRequest,
574 + Code: protocol.CodeHelloAck,
575 + TransportStatus: protocol.StatusOK,
576 + ItemCount: 1,
577 + }, payload)
578 + },
579 + want: ErrProtocol,
580 + wantText: "expected HELLO_ACK",
581 + },
582 + {
583 + name: "unexpected status",
584 + packetFn: func() []byte {
585 + return validHelloAckPacket(999)
586 + },
587 + want: ErrHandshake,
588 + wantText: "transport_status=999",
589 + },
590 + {
591 + name: "truncated payload",
592 + packetFn: func() []byte {
593 + ack := validHelloAckPacket(protocol.StatusOK)
594 + return ack[:protocol.HeaderSize+8]
595 + },
596 + want: ErrProtocol,
597 + wantText: "ack payload truncated",
598 + },
599 + {
600 + name: "bad ack payload layout",
601 + packetFn: func() []byte {
602 + ack := validHelloAckPacket(protocol.StatusOK)
603 + binary.NativeEndian.PutUint16(ack[protocol.HeaderSize:protocol.HeaderSize+2], 2)
604 + return ack
605 + },
606 + want: ErrIncompatible,
607 + wantText: "ack payload",
608 + },
609 + }
610 +
611 + for _, tc := range cases {
612 + t.Run(tc.name, func(t *testing.T) {
613 + runDir := testRunDir(t)
614 + service := uniqueService(t)
615 + lfd, path := rawSeqpacketListener(t, runDir, service)
616 +
617 + serverDone := make(chan error, 1)
618 + go func() {
619 + nfd, _, err := syscall.Accept(lfd)
620 + if err != nil {
621 + serverDone <- err
622 + return
623 + }
624 + defer syscall.Close(nfd)
625 +
626 + var helloBuf [128]byte
627 + if _, err := rawRecv(nfd, helloBuf[:]); err != nil {
628 + serverDone <- err
629 + return
630 + }
631 + serverDone <- rawSendPacket(nfd, tc.packetFn())
632 + }()
633 +
634 + client, err := Connect(runDir, service, defaultClientConfigPtr())
635 + if client != nil {
636 + client.Close()
637 + }
638 + if err == nil {
639 + t.Fatal("Connect should fail on malformed HELLO_ACK")
640 + }
641 + if !errors.Is(err, tc.want) {
642 + t.Fatalf("Connect error = %v, want %v", err, tc.want)
643 + }
644 + if tc.wantText != "" && !containsErrText(err, tc.wantText) {
645 + t.Fatalf("Connect error = %v, want text %q", err, tc.wantText)
646 + }
647 + if serr := <-serverDone; serr != nil {
648 + t.Fatalf("raw server failed: %v (path %s)", serr, path)
649 + }
650 + })
651 + }
652 +}
653 +
654 +func TestUdsClientHandshakePeerDisconnectBeforeAck(t *testing.T) {
655 + runDir := testRunDir(t)
656 + service := uniqueService(t)
657 + lfd, path := rawSeqpacketListener(t, runDir, service)
658 +
659 + serverDone := make(chan error, 1)
660 + go func() {
661 + nfd, _, err := syscall.Accept(lfd)
662 + if err != nil {
663 + serverDone <- err
664 + return
665 + }
666 + var helloBuf [128]byte
667 + _, recvErr := rawRecv(nfd, helloBuf[:])
668 + _ = syscall.Close(nfd)
669 + serverDone <- recvErr
670 + }()
671 +
672 + client, err := Connect(runDir, service, defaultClientConfigPtr())
673 + if client != nil {
674 + client.Close()
675 + }
676 + if err == nil {
677 + t.Fatal("Connect should fail when peer disconnects before HELLO_ACK")
678 + }
679 + if !errors.Is(err, ErrRecv) {
680 + t.Fatalf("Connect error = %v, want ErrRecv", err)
681 + }
682 + if serr := <-serverDone; serr != nil && !errors.Is(serr, ErrRecv) {
683 + t.Fatalf("raw server failed: %v (path %s)", serr, path)
684 + }
685 +}
686 +
687 +func TestUdsServerHandshakeRejectsMalformedHello(t *testing.T) {
688 + cases := []struct {
689 + name string
690 + packetFn func() []byte
691 + want error
692 + wantText string
693 + }{
694 + {
695 + name: "bad hello header",
696 + packetFn: func() []byte {
697 + hello := validHelloPacket()
698 + hello[0] = 0
699 + return hello
700 + },
701 + want: ErrProtocol,
702 + wantText: "hello header",
703 + },
704 + {
705 + name: "wrong kind",
706 + packetFn: func() []byte {
707 + payload := make([]byte, helloPayloadSize)
708 + return encodePacket(protocol.Header{
709 + Kind: protocol.KindRequest,
710 + Code: protocol.CodeHello,
711 + ItemCount: 1,
712 + }, payload)
713 + },
714 + want: ErrProtocol,
715 + wantText: "expected HELLO",
716 + },
717 + {
718 + name: "truncated payload",
719 + packetFn: func() []byte {
720 + hello := validHelloPacket()
721 + return hello[:protocol.HeaderSize+8]
722 + },
723 + want: ErrProtocol,
724 + wantText: "hello payload",
725 + },
726 + {
727 + name: "bad hello payload layout",
728 + packetFn: func() []byte {
729 + hello := validHelloPacket()
730 + binary.NativeEndian.PutUint16(hello[protocol.HeaderSize:protocol.HeaderSize+2], 2)
731 + return hello
732 + },
733 + want: ErrIncompatible,
734 + wantText: "protocol or layout version mismatch",
735 + },
736 + }
737 +
738 + for _, tc := range cases {
739 + t.Run(tc.name, func(t *testing.T) {
740 + runDir := testRunDir(t)
741 + service := uniqueService(t)
742 + lfd, path := rawSeqpacketListener(t, runDir, service)
743 +
744 + handshakeDone := make(chan error, 1)
745 + go func() {
746 + nfd, _, err := syscall.Accept(lfd)
747 + if err != nil {
748 + handshakeDone <- err
749 + return
750 + }
751 + defer syscall.Close(nfd)
752 +
753 + _, err = serverHandshake(nfd, defaultServerConfigPtr(), 1)
754 + handshakeDone <- err
755 + }()
756 +
757 + cfd := rawSeqpacketConnect(t, path)
758 + if err := rawSendPacket(cfd, tc.packetFn()); err != nil {
759 + t.Fatalf("rawSendPacket failed: %v", err)
760 + }
761 +
762 + err := <-handshakeDone
763 + if err == nil {
764 + t.Fatal("serverHandshake should fail on malformed HELLO")
765 + }
766 + if !errors.Is(err, tc.want) {
767 + t.Fatalf("serverHandshake error = %v, want %v", err, tc.want)
768 + }
769 + if tc.wantText != "" && !containsErrText(err, tc.wantText) {
770 + t.Fatalf("serverHandshake error = %v, want text %q", err, tc.wantText)
771 + }
772 + })
773 + }
774 +}
775 +
776 +func TestUdsServerHandshakePeerDisconnectBeforeHello(t *testing.T) {
777 + runDir := testRunDir(t)
778 + service := uniqueService(t)
779 + lfd, path := rawSeqpacketListener(t, runDir, service)
780 +
781 + handshakeDone := make(chan error, 1)
782 + go func() {
783 + nfd, _, err := syscall.Accept(lfd)
784 + if err != nil {
785 + handshakeDone <- err
786 + return
787 + }
788 + defer syscall.Close(nfd)
789 +
790 + _, err = serverHandshake(nfd, defaultServerConfigPtr(), 1)
791 + handshakeDone <- err
792 + }()
793 +
794 + cfd := rawSeqpacketConnect(t, path)
795 + _ = syscall.Close(cfd)
796 +
797 + err := <-handshakeDone
798 + if err == nil {
799 + t.Fatal("serverHandshake should fail when peer disconnects before HELLO")
800 + }
801 + if !errors.Is(err, ErrRecv) {
802 + t.Fatalf("serverHandshake error = %v, want ErrRecv", err)
803 + }
804 +}
805 +
806 +func defaultClientConfigPtr() *ClientConfig {
807 + cfg := defaultClientConfig()
808 + return &cfg
809 +}
810 +
811 +func defaultServerConfigPtr() *ServerConfig {
812 + cfg := defaultServerConfig()
813 + return &cfg
814 +}
815 +
816 +func containsErrText(err error, want string) bool {
817 + return err != nil && want != "" && strings.Contains(err.Error(), want)
818 +}
819 +
820 +func TestDetectPacketSizeFallbackAndSuccess(t *testing.T) {
821 + if got := detectPacketSize(-1); got != defaultPacketSizeFallback {
822 + t.Fatalf("detectPacketSize(-1) = %d, want fallback %d", got, defaultPacketSizeFallback)
823 + }
824 +
825 + client, _ := udsSessionPair(t, defaultServerConfig(), defaultClientConfig())
826 + if got := detectPacketSize(client.fd); got == 0 {
827 + t.Fatal("detectPacketSize(valid fd) returned 0")
828 + }
829 +}
src/go/pkg/netipc/transport/posix/uds_test.go new
+1859
@@ -0,0 +1,1859 @@
1 +//go:build unix
2 +
3 +package posix
4 +
5 +import (
6 + "bytes"
7 + "encoding/binary"
8 + "errors"
9 + "fmt"
10 + "os"
11 + "path/filepath"
12 + "sync"
13 + "sync/atomic"
14 + "testing"
15 + "time"
16 +
17 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
18 +)
19 +
20 +const (
21 + testAuthToken uint64 = 0xDEADBEEFCAFEBABE
22 +)
23 +
24 +var testServiceCounter atomic.Uint64
25 +
26 +// uniqueService returns a unique service name per test to avoid socket conflicts.
27 +func uniqueService(t *testing.T) string {
28 + t.Helper()
29 + n := testServiceCounter.Add(1)
30 + return fmt.Sprintf("gotest_%d_%d", os.Getpid(), n)
31 +}
32 +
33 +func testRunDir(t *testing.T) string {
34 + t.Helper()
35 + dir := filepath.Join(os.TempDir(), "nipc_go_test")
36 + if err := os.MkdirAll(dir, 0700); err != nil {
37 + t.Fatalf("cannot create run dir: %v", err)
38 + }
39 + return dir
40 +}
41 +
42 +func defaultServerConfig() ServerConfig {
43 + return ServerConfig{
44 + SupportedProfiles: protocol.ProfileBaseline,
45 + MaxRequestPayloadBytes: 4096,
46 + MaxRequestBatchItems: 16,
47 + MaxResponsePayloadBytes: 4096,
48 + MaxResponseBatchItems: 16,
49 + AuthToken: testAuthToken,
50 + Backlog: 4,
51 + }
52 +}
53 +
54 +func defaultClientConfig() ClientConfig {
55 + return ClientConfig{
56 + SupportedProfiles: protocol.ProfileBaseline,
57 + MaxRequestPayloadBytes: 4096,
58 + MaxRequestBatchItems: 16,
59 + MaxResponsePayloadBytes: 4096,
60 + MaxResponseBatchItems: 16,
61 + AuthToken: testAuthToken,
62 + }
63 +}
64 +
65 +// serverResult holds the result of an Accept call.
66 +type serverResult struct {
67 + session *Session
68 + err error
69 +}
70 +
71 +// startListener creates a listener and returns it. The caller must close it.
72 +func startListener(t *testing.T, runDir, service string, cfg ServerConfig) *Listener {
73 + t.Helper()
74 + listener, err := Listen(runDir, service, cfg)
75 + if err != nil {
76 + t.Fatalf("Listen failed: %v", err)
77 + }
78 + return listener
79 +}
80 +
81 +// acceptAsync starts accepting in a goroutine and returns a channel.
82 +func acceptAsync(listener *Listener) <-chan serverResult {
83 + ch := make(chan serverResult, 1)
84 + go func() {
85 + session, err := listener.Accept()
86 + ch <- serverResult{session, err}
87 + }()
88 + return ch
89 +}
90 +
91 +// ---------------------------------------------------------------------------
92 +// Test: Single client ping-pong
93 +// ---------------------------------------------------------------------------
94 +
95 +func TestSingleClientPingPong(t *testing.T) {
96 + runDir := testRunDir(t)
97 + service := uniqueService(t)
98 + defer os.Remove(filepath.Join(runDir, service+".sock"))
99 +
100 + sCfg := defaultServerConfig()
101 + listener := startListener(t, runDir, service, sCfg)
102 + defer listener.Close()
103 +
104 + acceptCh := acceptAsync(listener)
105 +
106 + // Client connects
107 + cCfg := defaultClientConfig()
108 + client, err := Connect(runDir, service, &cCfg)
109 + if err != nil {
110 + t.Fatalf("Connect failed: %v", err)
111 + }
112 + defer client.Close()
113 +
114 + // Wait for server accept
115 + sr := <-acceptCh
116 + if sr.err != nil {
117 + t.Fatalf("Accept failed: %v", sr.err)
118 + }
119 + server := sr.session
120 + defer server.Close()
121 +
122 + // Verify negotiated values
123 + if client.SelectedProfile != protocol.ProfileBaseline {
124 + t.Errorf("client profile = 0x%x, want 0x%x", client.SelectedProfile, protocol.ProfileBaseline)
125 + }
126 + if server.SelectedProfile != protocol.ProfileBaseline {
127 + t.Errorf("server profile = 0x%x, want 0x%x", server.SelectedProfile, protocol.ProfileBaseline)
128 + }
129 +
130 + // Client sends request
131 + payload := []byte("hello from client")
132 + hdr := protocol.Header{
133 + Kind: protocol.KindRequest,
134 + Code: protocol.MethodIncrement,
135 + ItemCount: 1,
136 + MessageID: 42,
137 + }
138 + if err := client.Send(&hdr, payload); err != nil {
139 + t.Fatalf("client Send: %v", err)
140 + }
141 +
142 + // Server receives
143 + recvBuf := make([]byte, 4096)
144 + rHdr, rPayload, err := server.Receive(recvBuf)
145 + if err != nil {
146 + t.Fatalf("server Receive: %v", err)
147 + }
148 +
149 + if rHdr.Kind != protocol.KindRequest {
150 + t.Errorf("server received kind=%d, want %d", rHdr.Kind, protocol.KindRequest)
151 + }
152 + if rHdr.MessageID != 42 {
153 + t.Errorf("server received message_id=%d, want 42", rHdr.MessageID)
154 + }
155 + if !bytes.Equal(rPayload, payload) {
156 + t.Errorf("server received payload mismatch")
157 + }
158 +
159 + // Server sends response
160 + respHdr := protocol.Header{
161 + Kind: protocol.KindResponse,
162 + Code: protocol.MethodIncrement,
163 + ItemCount: 1,
164 + MessageID: 42,
165 + }
166 + respPayload := []byte("response from server")
167 + if err := server.Send(&respHdr, respPayload); err != nil {
168 + t.Fatalf("server Send: %v", err)
169 + }
170 +
171 + // Client receives response
172 + rHdr, rPayload, err = client.Receive(recvBuf)
173 + if err != nil {
174 + t.Fatalf("client Receive: %v", err)
175 + }
176 + if rHdr.Kind != protocol.KindResponse {
177 + t.Errorf("client received kind=%d, want %d", rHdr.Kind, protocol.KindResponse)
178 + }
179 + if rHdr.MessageID != 42 {
180 + t.Errorf("client received message_id=%d, want 42", rHdr.MessageID)
181 + }
182 + if !bytes.Equal(rPayload, respPayload) {
183 + t.Errorf("client received payload mismatch")
184 + }
185 +}
186 +
187 +// ---------------------------------------------------------------------------
188 +// Test: Multi-client (2 clients)
189 +// ---------------------------------------------------------------------------
190 +
191 +func TestMultiClient(t *testing.T) {
192 + runDir := testRunDir(t)
193 + service := uniqueService(t)
194 + defer os.Remove(filepath.Join(runDir, service+".sock"))
195 +
196 + sCfg := defaultServerConfig()
197 + listener := startListener(t, runDir, service, sCfg)
198 + defer listener.Close()
199 +
200 + // Connect two clients
201 + const numClients = 2
202 + clients := make([]*Session, numClients)
203 + servers := make([]*Session, numClients)
204 +
205 + for i := 0; i < numClients; i++ {
206 + acceptCh := acceptAsync(listener)
207 +
208 + cCfg := defaultClientConfig()
209 + c, err := Connect(runDir, service, &cCfg)
210 + if err != nil {
211 + t.Fatalf("Connect[%d] failed: %v", i, err)
212 + }
213 + clients[i] = c
214 +
215 + sr := <-acceptCh
216 + if sr.err != nil {
217 + t.Fatalf("Accept[%d] failed: %v", i, sr.err)
218 + }
219 + servers[i] = sr.session
220 + }
221 +
222 + defer func() {
223 + for i := 0; i < numClients; i++ {
224 + clients[i].Close()
225 + servers[i].Close()
226 + }
227 + }()
228 +
229 + // Each client sends a unique message
230 + for i := 0; i < numClients; i++ {
231 + payload := []byte(fmt.Sprintf("client_%d", i))
232 + hdr := protocol.Header{
233 + Kind: protocol.KindRequest,
234 + Code: protocol.MethodIncrement,
235 + ItemCount: 1,
236 + MessageID: uint64(100 + i),
237 + }
238 + if err := clients[i].Send(&hdr, payload); err != nil {
239 + t.Fatalf("client[%d] Send: %v", i, err)
240 + }
241 + }
242 +
243 + // Each server receives and echoes
244 + buf := make([]byte, 4096)
245 + for i := 0; i < numClients; i++ {
246 + rHdr, rPayload, err := servers[i].Receive(buf)
247 + if err != nil {
248 + t.Fatalf("server[%d] Receive: %v", i, err)
249 + }
250 +
251 + expected := []byte(fmt.Sprintf("client_%d", i))
252 + if !bytes.Equal(rPayload, expected) {
253 + t.Errorf("server[%d] payload = %q, want %q", i, rPayload, expected)
254 + }
255 +
256 + resp := protocol.Header{
257 + Kind: protocol.KindResponse,
258 + Code: rHdr.Code,
259 + ItemCount: 1,
260 + MessageID: rHdr.MessageID,
261 + }
262 + if err := servers[i].Send(&resp, rPayload); err != nil {
263 + t.Fatalf("server[%d] Send: %v", i, err)
264 + }
265 + }
266 +
267 + // Each client receives its echo
268 + for i := 0; i < numClients; i++ {
269 + rHdr, rPayload, err := clients[i].Receive(buf)
270 + if err != nil {
271 + t.Fatalf("client[%d] Receive: %v", i, err)
272 + }
273 + if rHdr.MessageID != uint64(100+i) {
274 + t.Errorf("client[%d] message_id = %d, want %d", i, rHdr.MessageID, 100+i)
275 + }
276 + expected := []byte(fmt.Sprintf("client_%d", i))
277 + if !bytes.Equal(rPayload, expected) {
278 + t.Errorf("client[%d] response payload = %q, want %q", i, rPayload, expected)
279 + }
280 + }
281 +}
282 +
283 +// ---------------------------------------------------------------------------
284 +// Test: Pipelining (3 requests, match by message_id)
285 +// ---------------------------------------------------------------------------
286 +
287 +func TestPipelining(t *testing.T) {
288 + runDir := testRunDir(t)
289 + service := uniqueService(t)
290 + defer os.Remove(filepath.Join(runDir, service+".sock"))
291 +
292 + sCfg := defaultServerConfig()
293 + listener := startListener(t, runDir, service, sCfg)
294 + defer listener.Close()
295 +
296 + acceptCh := acceptAsync(listener)
297 +
298 + cCfg := defaultClientConfig()
299 + client, err := Connect(runDir, service, &cCfg)
300 + if err != nil {
301 + t.Fatalf("Connect failed: %v", err)
302 + }
303 + defer client.Close()
304 +
305 + sr := <-acceptCh
306 + if sr.err != nil {
307 + t.Fatalf("Accept failed: %v", sr.err)
308 + }
309 + server := sr.session
310 + defer server.Close()
311 +
312 + // Client sends 3 requests without waiting
313 + messageIDs := []uint64{10, 20, 30}
314 + for _, mid := range messageIDs {
315 + hdr := protocol.Header{
316 + Kind: protocol.KindRequest,
317 + Code: protocol.MethodIncrement,
318 + ItemCount: 1,
319 + MessageID: mid,
320 + }
321 + payload := []byte(fmt.Sprintf("req_%d", mid))
322 + if err := client.Send(&hdr, payload); err != nil {
323 + t.Fatalf("client Send(%d): %v", mid, err)
324 + }
325 + }
326 +
327 + // Server receives all 3 and responds in reverse order
328 + buf := make([]byte, 4096)
329 + type reqInfo struct {
330 + hdr protocol.Header
331 + payload []byte
332 + }
333 + reqs := make([]reqInfo, 0, 3)
334 +
335 + for i := 0; i < 3; i++ {
336 + rHdr, rPayload, err := server.Receive(buf)
337 + if err != nil {
338 + t.Fatalf("server Receive[%d]: %v", i, err)
339 + }
340 + reqs = append(reqs, reqInfo{rHdr, append([]byte(nil), rPayload...)})
341 + }
342 +
343 + // Respond in reverse order
344 + for i := len(reqs) - 1; i >= 0; i-- {
345 + resp := protocol.Header{
346 + Kind: protocol.KindResponse,
347 + Code: reqs[i].hdr.Code,
348 + ItemCount: 1,
349 + MessageID: reqs[i].hdr.MessageID,
350 + }
351 + respPayload := append([]byte("resp_"), reqs[i].payload...)
352 + if err := server.Send(&resp, respPayload); err != nil {
353 + t.Fatalf("server Send[%d]: %v", i, err)
354 + }
355 + }
356 +
357 + // Client receives 3 responses (should arrive in reverse order)
358 + received := make(map[uint64][]byte)
359 + for i := 0; i < 3; i++ {
360 + rHdr, rPayload, err := client.Receive(buf)
361 + if err != nil {
362 + t.Fatalf("client Receive[%d]: %v", i, err)
363 + }
364 + if rHdr.Kind != protocol.KindResponse {
365 + t.Errorf("received kind=%d, want RESPONSE", rHdr.Kind)
366 + }
367 + received[rHdr.MessageID] = append([]byte(nil), rPayload...)
368 + }
369 +
370 + // Verify all messages received with correct data
371 + for _, mid := range messageIDs {
372 + payload, ok := received[mid]
373 + if !ok {
374 + t.Errorf("missing response for message_id %d", mid)
375 + continue
376 + }
377 + expected := []byte(fmt.Sprintf("resp_req_%d", mid))
378 + if !bytes.Equal(payload, expected) {
379 + t.Errorf("message_id %d: payload = %q, want %q", mid, payload, expected)
380 + }
381 + }
382 +}
383 +
384 +// ---------------------------------------------------------------------------
385 +// Test: Chunking (large message, small packet_size)
386 +// ---------------------------------------------------------------------------
387 +
388 +func TestChunking(t *testing.T) {
389 + runDir := testRunDir(t)
390 + service := uniqueService(t)
391 + defer os.Remove(filepath.Join(runDir, service+".sock"))
392 +
393 + // Use a small packet_size to force chunking
394 + const forcedPacketSize = 128
395 +
396 + sCfg := defaultServerConfig()
397 + sCfg.PacketSize = forcedPacketSize
398 + sCfg.MaxRequestPayloadBytes = 65536
399 + sCfg.MaxResponsePayloadBytes = 65536
400 + listener := startListener(t, runDir, service, sCfg)
401 + defer listener.Close()
402 +
403 + acceptCh := acceptAsync(listener)
404 +
405 + cCfg := defaultClientConfig()
406 + cCfg.PacketSize = forcedPacketSize
407 + cCfg.MaxRequestPayloadBytes = 65536
408 + cCfg.MaxResponsePayloadBytes = 65536
409 + client, err := Connect(runDir, service, &cCfg)
410 + if err != nil {
411 + t.Fatalf("Connect failed: %v", err)
412 + }
413 + defer client.Close()
414 +
415 + sr := <-acceptCh
416 + if sr.err != nil {
417 + t.Fatalf("Accept failed: %v", sr.err)
418 + }
419 + server := sr.session
420 + defer server.Close()
421 +
422 + // Verify small packet size was negotiated
423 + if client.PacketSize > forcedPacketSize {
424 + t.Errorf("client packet_size = %d, want <= %d", client.PacketSize, forcedPacketSize)
425 + }
426 +
427 + // Build a large payload (much larger than packet_size)
428 + largePayload := make([]byte, 2000)
429 + for i := range largePayload {
430 + largePayload[i] = byte(i & 0xFF)
431 + }
432 +
433 + // Client sends large message (will be chunked)
434 + hdr := protocol.Header{
435 + Kind: protocol.KindRequest,
436 + Code: protocol.MethodIncrement,
437 + ItemCount: 1,
438 + MessageID: 99,
439 + }
440 + if err := client.Send(&hdr, largePayload); err != nil {
441 + t.Fatalf("client Send (chunked): %v", err)
442 + }
443 +
444 + // Server receives (reassembles chunks)
445 + recvBuf := make([]byte, forcedPacketSize)
446 + rHdr, rPayload, err := server.Receive(recvBuf)
447 + if err != nil {
448 + t.Fatalf("server Receive (chunked): %v", err)
449 + }
450 +
451 + if rHdr.MessageID != 99 {
452 + t.Errorf("message_id = %d, want 99", rHdr.MessageID)
453 + }
454 + if !bytes.Equal(rPayload, largePayload) {
455 + t.Errorf("chunked payload mismatch: got %d bytes, want %d", len(rPayload), len(largePayload))
456 + }
457 +
458 + // Server echoes back (also chunked)
459 + resp := protocol.Header{
460 + Kind: protocol.KindResponse,
461 + Code: protocol.MethodIncrement,
462 + ItemCount: 1,
463 + MessageID: 99,
464 + }
465 + if err := server.Send(&resp, rPayload); err != nil {
466 + t.Fatalf("server Send (chunked echo): %v", err)
467 + }
468 +
469 + // Client receives
470 + rHdr, rPayload, err = client.Receive(recvBuf)
471 + if err != nil {
472 + t.Fatalf("client Receive (chunked): %v", err)
473 + }
474 + if !bytes.Equal(rPayload, largePayload) {
475 + t.Errorf("client chunked payload mismatch: got %d bytes, want %d", len(rPayload), len(largePayload))
476 + }
477 +}
478 +
479 +// ---------------------------------------------------------------------------
480 +// Test: Handshake failures - bad auth
481 +// ---------------------------------------------------------------------------
482 +
483 +func TestHandshakeBadAuth(t *testing.T) {
484 + runDir := testRunDir(t)
485 + service := uniqueService(t)
486 + defer os.Remove(filepath.Join(runDir, service+".sock"))
487 +
488 + sCfg := defaultServerConfig()
489 + sCfg.AuthToken = 0x1111111111111111
490 + listener := startListener(t, runDir, service, sCfg)
491 + defer listener.Close()
492 +
493 + acceptCh := acceptAsync(listener)
494 +
495 + cCfg := defaultClientConfig()
496 + cCfg.AuthToken = 0x2222222222222222 // wrong token
497 + _, err := Connect(runDir, service, &cCfg)
498 + if err == nil {
499 + t.Fatal("expected auth failure, got nil")
500 + }
501 +
502 + if err != ErrAuthFailed {
503 + t.Errorf("error = %v, want ErrAuthFailed", err)
504 + }
505 +
506 + // Server side should also fail
507 + sr := <-acceptCh
508 + if sr.err == nil {
509 + sr.session.Close()
510 + t.Fatal("server expected auth failure")
511 + }
512 + if sr.err != ErrAuthFailed {
513 + t.Errorf("server error = %v, want ErrAuthFailed", sr.err)
514 + }
515 +}
516 +
517 +// ---------------------------------------------------------------------------
518 +// Test: Handshake failures - profile mismatch
519 +// ---------------------------------------------------------------------------
520 +
521 +func TestHandshakeProfileMismatch(t *testing.T) {
522 + runDir := testRunDir(t)
523 + service := uniqueService(t)
524 + defer os.Remove(filepath.Join(runDir, service+".sock"))
525 +
526 + sCfg := defaultServerConfig()
527 + sCfg.SupportedProfiles = protocol.ProfileSHMFutex // only SHM
528 + listener := startListener(t, runDir, service, sCfg)
529 + defer listener.Close()
530 +
531 + acceptCh := acceptAsync(listener)
532 +
533 + cCfg := defaultClientConfig()
534 + cCfg.SupportedProfiles = protocol.ProfileBaseline // only baseline
535 + _, err := Connect(runDir, service, &cCfg)
536 + if err == nil {
537 + t.Fatal("expected profile mismatch, got nil")
538 + }
539 +
540 + if err != ErrNoProfile {
541 + t.Errorf("error = %v, want ErrNoProfile", err)
542 + }
543 +
544 + sr := <-acceptCh
545 + if sr.err == nil {
546 + sr.session.Close()
547 + t.Fatal("server expected profile mismatch")
548 + }
549 +}
550 +
551 +func TestHandshakeRequestPayloadOverCap(t *testing.T) {
552 + runDir := testRunDir(t)
553 + service := uniqueService(t)
554 + defer os.Remove(filepath.Join(runDir, service+".sock"))
555 +
556 + listener := startListener(t, runDir, service, defaultServerConfig())
557 + defer listener.Close()
558 +
559 + acceptCh := acceptAsync(listener)
560 +
561 + cCfg := defaultClientConfig()
562 + cCfg.MaxRequestPayloadBytes = protocol.MaxPayloadCap + 1
563 + if _, err := Connect(runDir, service, &cCfg); !errors.Is(err, ErrLimitExceeded) {
564 + t.Fatalf("expected ErrLimitExceeded, got %v", err)
565 + }
566 +
567 + sr := <-acceptCh
568 + if !errors.Is(sr.err, ErrLimitExceeded) {
569 + t.Fatalf("server expected ErrLimitExceeded, got %v", sr.err)
570 + }
571 +}
572 +
573 +func TestHandshakeIncompatibleClassifierHelpers(t *testing.T) {
574 + hdr := protocol.Header{
575 + Magic: protocol.MagicMsg,
576 + Version: protocol.Version + 1,
577 + HeaderLen: protocol.HeaderLen,
578 + Kind: protocol.KindControl,
579 + Code: protocol.CodeHello,
580 + }
581 + hdrBuf := make([]byte, protocol.HeaderSize)
582 + hdr.Encode(hdrBuf)
583 + if !headerVersionIncompatible(hdrBuf, protocol.CodeHello) {
584 + t.Fatal("headerVersionIncompatible should detect bad HELLO version")
585 + }
586 + if headerVersionIncompatible(hdrBuf, protocol.CodeHelloAck) {
587 + t.Fatal("headerVersionIncompatible should respect expected code")
588 + }
589 +
590 + hello := protocol.Hello{LayoutVersion: 2}
591 + helloBuf := make([]byte, 44)
592 + hello.Encode(helloBuf)
593 + if !helloLayoutIncompatible(helloBuf) {
594 + t.Fatal("helloLayoutIncompatible should detect bad layout_version")
595 + }
596 +
597 + ack := protocol.HelloAck{LayoutVersion: 2}
598 + ackBuf := make([]byte, 48)
599 + ack.Encode(ackBuf)
600 + if !helloAckLayoutIncompatible(ackBuf) {
601 + t.Fatal("helloAckLayoutIncompatible should detect bad layout_version")
602 + }
603 +}
604 +
605 +// ---------------------------------------------------------------------------
606 +// Test: Stale socket recovery
607 +// ---------------------------------------------------------------------------
608 +
609 +func TestStaleRecovery(t *testing.T) {
610 + runDir := testRunDir(t)
611 + service := uniqueService(t)
612 + sockPath := filepath.Join(runDir, service+".sock")
613 + defer os.Remove(sockPath)
614 +
615 + // Create a stale socket file (not a real listener)
616 + f, err := os.Create(sockPath)
617 + if err != nil {
618 + t.Fatalf("cannot create stale file: %v", err)
619 + }
620 + f.Close()
621 +
622 + // Listen should recover and succeed
623 + sCfg := defaultServerConfig()
624 + listener, err := Listen(runDir, service, sCfg)
625 + if err != nil {
626 + t.Fatalf("Listen (stale recovery) failed: %v", err)
627 + }
628 + defer listener.Close()
629 +
630 + // Now try to listen again — should get AddrInUse because
631 + // the first listener is alive
632 + _, err = Listen(runDir, service, sCfg)
633 + if err != ErrAddrInUse {
634 + t.Errorf("second Listen: error = %v, want ErrAddrInUse", err)
635 + }
636 +}
637 +
638 +// ---------------------------------------------------------------------------
639 +// Test: Disconnect detection
640 +// ---------------------------------------------------------------------------
641 +
642 +func TestDisconnectDetection(t *testing.T) {
643 + runDir := testRunDir(t)
644 + service := uniqueService(t)
645 + defer os.Remove(filepath.Join(runDir, service+".sock"))
646 +
647 + sCfg := defaultServerConfig()
648 + listener := startListener(t, runDir, service, sCfg)
649 + defer listener.Close()
650 +
651 + acceptCh := acceptAsync(listener)
652 +
653 + cCfg := defaultClientConfig()
654 + client, err := Connect(runDir, service, &cCfg)
655 + if err != nil {
656 + t.Fatalf("Connect failed: %v", err)
657 + }
658 +
659 + sr := <-acceptCh
660 + if sr.err != nil {
661 + t.Fatalf("Accept failed: %v", sr.err)
662 + }
663 + server := sr.session
664 + defer server.Close()
665 +
666 + // Close client
667 + client.Close()
668 +
669 + // Server should get an error on Receive
670 + buf := make([]byte, 4096)
671 + _, _, err = server.Receive(buf)
672 + if err == nil {
673 + t.Fatal("expected error on Receive after client disconnect, got nil")
674 + }
675 +}
676 +
677 +// ---------------------------------------------------------------------------
678 +// Test: Listener Fd
679 +// ---------------------------------------------------------------------------
680 +
681 +func TestListenerFd(t *testing.T) {
682 + runDir := testRunDir(t)
683 + service := uniqueService(t)
684 + defer os.Remove(filepath.Join(runDir, service+".sock"))
685 +
686 + sCfg := defaultServerConfig()
687 + listener := startListener(t, runDir, service, sCfg)
688 + defer listener.Close()
689 +
690 + if listener.Fd() < 0 {
691 + t.Errorf("listener fd = %d, want >= 0", listener.Fd())
692 + }
693 +}
694 +
695 +// ---------------------------------------------------------------------------
696 +// Test: Session Fd and Role
697 +// ---------------------------------------------------------------------------
698 +
699 +func TestSessionFdAndRole(t *testing.T) {
700 + runDir := testRunDir(t)
701 + service := uniqueService(t)
702 + defer os.Remove(filepath.Join(runDir, service+".sock"))
703 +
704 + sCfg := defaultServerConfig()
705 + listener := startListener(t, runDir, service, sCfg)
706 + defer listener.Close()
707 +
708 + acceptCh := acceptAsync(listener)
709 +
710 + cCfg := defaultClientConfig()
711 + client, err := Connect(runDir, service, &cCfg)
712 + if err != nil {
713 + t.Fatalf("Connect failed: %v", err)
714 + }
715 + defer client.Close()
716 +
717 + sr := <-acceptCh
718 + if sr.err != nil {
719 + t.Fatalf("Accept failed: %v", sr.err)
720 + }
721 + server := sr.session
722 + defer server.Close()
723 +
724 + if client.Fd() < 0 {
725 + t.Errorf("client fd = %d, want >= 0", client.Fd())
726 + }
727 + if server.Fd() < 0 {
728 + t.Errorf("server fd = %d, want >= 0", server.Fd())
729 + }
730 + if client.Role() != RoleClient {
731 + t.Errorf("client role = %d, want RoleClient", client.Role())
732 + }
733 + if server.Role() != RoleServer {
734 + t.Errorf("server role = %d, want RoleServer", server.Role())
735 + }
736 +}
737 +
738 +// ---------------------------------------------------------------------------
739 +// Test: Directional limit negotiation
740 +// ---------------------------------------------------------------------------
741 +
742 +func TestDirectionalLimitNegotiation(t *testing.T) {
743 + runDir := testRunDir(t)
744 + service := uniqueService(t)
745 + defer os.Remove(filepath.Join(runDir, service+".sock"))
746 +
747 + sCfg := defaultServerConfig()
748 + sCfg.MaxRequestPayloadBytes = 2048
749 + sCfg.MaxRequestBatchItems = 8
750 + sCfg.MaxResponsePayloadBytes = 8192
751 + sCfg.MaxResponseBatchItems = 32
752 + listener := startListener(t, runDir, service, sCfg)
753 + defer listener.Close()
754 +
755 + acceptCh := acceptAsync(listener)
756 +
757 + cCfg := defaultClientConfig()
758 + cCfg.MaxRequestPayloadBytes = 4096
759 + cCfg.MaxRequestBatchItems = 16
760 + cCfg.MaxResponsePayloadBytes = 4096
761 + cCfg.MaxResponseBatchItems = 16
762 + client, err := Connect(runDir, service, &cCfg)
763 + if err != nil {
764 + t.Fatalf("Connect failed: %v", err)
765 + }
766 + defer client.Close()
767 +
768 + sr := <-acceptCh
769 + if sr.err != nil {
770 + t.Fatalf("Accept failed: %v", sr.err)
771 + }
772 + server := sr.session
773 + defer server.Close()
774 +
775 + // Request-side values are echoed back to the caller, while response payload
776 + // size remains server-owned. Response batch size stays symmetric with the
777 + // negotiated request batch size.
778 + if client.MaxRequestPayloadBytes != 4096 {
779 + t.Errorf("request payload = %d, want 4096", client.MaxRequestPayloadBytes)
780 + }
781 + if client.MaxRequestBatchItems != 16 {
782 + t.Errorf("request batch = %d, want 16", client.MaxRequestBatchItems)
783 + }
784 + if client.MaxResponsePayloadBytes != 8192 {
785 + t.Errorf("response payload = %d, want 8192", client.MaxResponsePayloadBytes)
786 + }
787 + if client.MaxResponseBatchItems != 16 {
788 + t.Errorf("response batch = %d, want 16", client.MaxResponseBatchItems)
789 + }
790 + if client.SessionID == 0 {
791 + t.Fatal("session id must be non-zero")
792 + }
793 +
794 + // Server should have the same negotiated values
795 + if server.MaxRequestPayloadBytes != client.MaxRequestPayloadBytes {
796 + t.Errorf("server req_payload = %d, want %d", server.MaxRequestPayloadBytes, client.MaxRequestPayloadBytes)
797 + }
798 + if server.MaxRequestBatchItems != client.MaxRequestBatchItems {
799 + t.Errorf("server req_batch = %d, want %d", server.MaxRequestBatchItems, client.MaxRequestBatchItems)
800 + }
801 + if server.MaxResponsePayloadBytes != client.MaxResponsePayloadBytes {
802 + t.Errorf("server resp_payload = %d, want %d", server.MaxResponsePayloadBytes, client.MaxResponsePayloadBytes)
803 + }
804 + if server.MaxResponseBatchItems != client.MaxResponseBatchItems {
805 + t.Errorf("server resp_batch = %d, want %d", server.MaxResponseBatchItems, client.MaxResponseBatchItems)
806 + }
807 + if server.SessionID != client.SessionID {
808 + t.Errorf("server session_id = %d, want %d", server.SessionID, client.SessionID)
809 + }
810 +}
811 +
812 +// ---------------------------------------------------------------------------
813 +// Test: Profile selection (highest bit in preferred_intersection)
814 +// ---------------------------------------------------------------------------
815 +
816 +func TestProfileSelection(t *testing.T) {
817 + runDir := testRunDir(t)
818 + service := uniqueService(t)
819 + defer os.Remove(filepath.Join(runDir, service+".sock"))
820 +
821 + sCfg := defaultServerConfig()
822 + sCfg.SupportedProfiles = protocol.ProfileBaseline | protocol.ProfileSHMHybrid | protocol.ProfileSHMFutex
823 + sCfg.PreferredProfiles = protocol.ProfileSHMFutex
824 + listener := startListener(t, runDir, service, sCfg)
825 + defer listener.Close()
826 +
827 + acceptCh := acceptAsync(listener)
828 +
829 + cCfg := defaultClientConfig()
830 + cCfg.SupportedProfiles = protocol.ProfileBaseline | protocol.ProfileSHMHybrid | protocol.ProfileSHMFutex
831 + cCfg.PreferredProfiles = protocol.ProfileSHMFutex | protocol.ProfileSHMHybrid
832 + client, err := Connect(runDir, service, &cCfg)
833 + if err != nil {
834 + t.Fatalf("Connect failed: %v", err)
835 + }
836 + defer client.Close()
837 +
838 + sr := <-acceptCh
839 + if sr.err != nil {
840 + t.Fatalf("Accept failed: %v", sr.err)
841 + }
842 + server := sr.session
843 + defer server.Close()
844 +
845 + // Preferred intersection = SHMFutex (0x04), highest bit = 0x04
846 + if client.SelectedProfile != protocol.ProfileSHMFutex {
847 + t.Errorf("selected = 0x%x, want 0x%x (SHMFutex)", client.SelectedProfile, protocol.ProfileSHMFutex)
848 + }
849 +}
850 +
851 +// ---------------------------------------------------------------------------
852 +// Test: Empty payload
853 +// ---------------------------------------------------------------------------
854 +
855 +func TestEmptyPayload(t *testing.T) {
856 + runDir := testRunDir(t)
857 + service := uniqueService(t)
858 + defer os.Remove(filepath.Join(runDir, service+".sock"))
859 +
860 + sCfg := defaultServerConfig()
861 + listener := startListener(t, runDir, service, sCfg)
862 + defer listener.Close()
863 +
864 + acceptCh := acceptAsync(listener)
865 +
866 + cCfg := defaultClientConfig()
867 + client, err := Connect(runDir, service, &cCfg)
868 + if err != nil {
869 + t.Fatalf("Connect failed: %v", err)
870 + }
871 + defer client.Close()
872 +
873 + sr := <-acceptCh
874 + if sr.err != nil {
875 + t.Fatalf("Accept failed: %v", sr.err)
876 + }
877 + server := sr.session
878 + defer server.Close()
879 +
880 + // Send empty payload
881 + hdr := protocol.Header{
882 + Kind: protocol.KindRequest,
883 + Code: protocol.MethodIncrement,
884 + ItemCount: 1,
885 + MessageID: 1,
886 + }
887 + if err := client.Send(&hdr, nil); err != nil {
888 + t.Fatalf("Send empty: %v", err)
889 + }
890 +
891 + buf := make([]byte, 4096)
892 + rHdr, rPayload, err := server.Receive(buf)
893 + if err != nil {
894 + t.Fatalf("Receive empty: %v", err)
895 + }
896 + if rHdr.MessageID != 1 {
897 + t.Errorf("message_id = %d, want 1", rHdr.MessageID)
898 + }
899 + if len(rPayload) != 0 {
900 + t.Errorf("payload len = %d, want 0", len(rPayload))
901 + }
902 +}
903 +
904 +// ---------------------------------------------------------------------------
905 +// Test: Concurrent send/receive (stress test)
906 +// ---------------------------------------------------------------------------
907 +
908 +func TestConcurrentSendReceive(t *testing.T) {
909 + runDir := testRunDir(t)
910 + service := uniqueService(t)
911 + defer os.Remove(filepath.Join(runDir, service+".sock"))
912 +
913 + sCfg := defaultServerConfig()
914 + sCfg.MaxRequestPayloadBytes = 65536
915 + sCfg.MaxResponsePayloadBytes = 65536
916 + listener := startListener(t, runDir, service, sCfg)
917 + defer listener.Close()
918 +
919 + acceptCh := acceptAsync(listener)
920 +
921 + cCfg := defaultClientConfig()
922 + cCfg.MaxRequestPayloadBytes = 65536
923 + cCfg.MaxResponsePayloadBytes = 65536
924 + client, err := Connect(runDir, service, &cCfg)
925 + if err != nil {
926 + t.Fatalf("Connect failed: %v", err)
927 + }
928 + defer client.Close()
929 +
930 + sr := <-acceptCh
931 + if sr.err != nil {
932 + t.Fatalf("Accept failed: %v", sr.err)
933 + }
934 + server := sr.session
935 + defer server.Close()
936 +
937 + const numMessages = 20
938 + var wg sync.WaitGroup
939 +
940 + // Server goroutine: receive and echo
941 + wg.Add(1)
942 + go func() {
943 + defer wg.Done()
944 + buf := make([]byte, 65600)
945 + for i := 0; i < numMessages; i++ {
946 + rHdr, rPayload, err := server.Receive(buf)
947 + if err != nil {
948 + t.Errorf("server Receive[%d]: %v", i, err)
949 + return
950 + }
951 + resp := protocol.Header{
952 + Kind: protocol.KindResponse,
953 + Code: rHdr.Code,
954 + ItemCount: 1,
955 + MessageID: rHdr.MessageID,
956 + }
957 + if err := server.Send(&resp, rPayload); err != nil {
958 + t.Errorf("server Send[%d]: %v", i, err)
959 + return
960 + }
961 + }
962 + }()
963 +
964 + // Client: send all, then receive all
965 + for i := 0; i < numMessages; i++ {
966 + payload := []byte(fmt.Sprintf("message_%d", i))
967 + hdr := protocol.Header{
968 + Kind: protocol.KindRequest,
969 + Code: protocol.MethodIncrement,
970 + ItemCount: 1,
971 + MessageID: uint64(i),
972 + }
973 + if err := client.Send(&hdr, payload); err != nil {
974 + t.Fatalf("client Send[%d]: %v", i, err)
975 + }
976 + }
977 +
978 + received := make(map[uint64]bool)
979 + buf := make([]byte, 65600)
980 + for i := 0; i < numMessages; i++ {
981 + rHdr, _, err := client.Receive(buf)
982 + if err != nil {
983 + t.Fatalf("client Receive[%d]: %v", i, err)
984 + }
985 + received[rHdr.MessageID] = true
986 + }
987 +
988 + wg.Wait()
989 +
990 + for i := 0; i < numMessages; i++ {
991 + if !received[uint64(i)] {
992 + t.Errorf("missing response for message_id %d", i)
993 + }
994 + }
995 +}
996 +
997 +// ---------------------------------------------------------------------------
998 +// Test: Close then Send/Receive returns error
999 +// ---------------------------------------------------------------------------
1000 +
1001 +func TestClosedSessionErrors(t *testing.T) {
1002 + runDir := testRunDir(t)
1003 + service := uniqueService(t)
1004 + defer os.Remove(filepath.Join(runDir, service+".sock"))
1005 +
1006 + sCfg := defaultServerConfig()
1007 + listener := startListener(t, runDir, service, sCfg)
1008 + defer listener.Close()
1009 +
1010 + acceptCh := acceptAsync(listener)
1011 +
1012 + cCfg := defaultClientConfig()
1013 + client, err := Connect(runDir, service, &cCfg)
1014 + if err != nil {
1015 + t.Fatalf("Connect failed: %v", err)
1016 + }
1017 +
1018 + sr := <-acceptCh
1019 + if sr.err != nil {
1020 + t.Fatalf("Accept failed: %v", sr.err)
1021 + }
1022 + sr.session.Close()
1023 +
1024 + client.Close()
1025 +
1026 + // Send on closed session
1027 + hdr := protocol.Header{
1028 + Kind: protocol.KindRequest,
1029 + Code: protocol.MethodIncrement,
1030 + ItemCount: 1,
1031 + MessageID: 1,
1032 + }
1033 + if err := client.Send(&hdr, []byte("x")); err == nil {
1034 + t.Error("Send on closed session should fail")
1035 + }
1036 +
1037 + // Receive on closed session
1038 + buf := make([]byte, 4096)
1039 + _, _, err = client.Receive(buf)
1040 + if err == nil {
1041 + t.Error("Receive on closed session should fail")
1042 + }
1043 +}
1044 +
1045 +// ---------------------------------------------------------------------------
1046 +// Test: Path too long
1047 +// ---------------------------------------------------------------------------
1048 +
1049 +func TestPathTooLong(t *testing.T) {
1050 + longDir := "/tmp/" + string(make([]byte, 200))
1051 + _, err := Connect(longDir, "test", &ClientConfig{})
1052 + if err != ErrPathTooLong {
1053 + t.Errorf("expected ErrPathTooLong, got %v", err)
1054 + }
1055 +
1056 + _, err = Listen(longDir, "test", ServerConfig{})
1057 + if err != ErrPathTooLong {
1058 + t.Errorf("expected ErrPathTooLong on Listen, got %v", err)
1059 + }
1060 +}
1061 +
1062 +// ---------------------------------------------------------------------------
1063 +// Test: Listener close prevents new accepts
1064 +// ---------------------------------------------------------------------------
1065 +
1066 +func TestListenerClose(t *testing.T) {
1067 + runDir := testRunDir(t)
1068 + service := uniqueService(t)
1069 + defer os.Remove(filepath.Join(runDir, service+".sock"))
1070 +
1071 + sCfg := defaultServerConfig()
1072 + listener := startListener(t, runDir, service, sCfg)
1073 + listener.Close()
1074 +
1075 + // Connect should fail after listener close
1076 + cCfg := defaultClientConfig()
1077 + _, err := Connect(runDir, service, &cCfg)
1078 + if err == nil {
1079 + t.Error("Connect after listener close should fail")
1080 + }
1081 +}
1082 +
1083 +// ---------------------------------------------------------------------------
1084 +// Test: Defaults applied when config values are 0
1085 +// ---------------------------------------------------------------------------
1086 +
1087 +func TestDefaultsApplied(t *testing.T) {
1088 + runDir := testRunDir(t)
1089 + service := uniqueService(t)
1090 + defer os.Remove(filepath.Join(runDir, service+".sock"))
1091 +
1092 + // Zero config = all defaults
1093 + sCfg := ServerConfig{
1094 + AuthToken: testAuthToken,
1095 + }
1096 + listener := startListener(t, runDir, service, sCfg)
1097 + defer listener.Close()
1098 +
1099 + acceptCh := acceptAsync(listener)
1100 +
1101 + cCfg := ClientConfig{
1102 + AuthToken: testAuthToken,
1103 + }
1104 + client, err := Connect(runDir, service, &cCfg)
1105 + if err != nil {
1106 + t.Fatalf("Connect failed: %v", err)
1107 + }
1108 + defer client.Close()
1109 +
1110 + sr := <-acceptCh
1111 + if sr.err != nil {
1112 + t.Fatalf("Accept failed: %v", sr.err)
1113 + }
1114 + server := sr.session
1115 + defer server.Close()
1116 +
1117 + // Defaults: MaxPayloadDefault=1024, batch=1
1118 + if client.MaxRequestPayloadBytes != protocol.MaxPayloadDefault {
1119 + t.Errorf("req_payload = %d, want %d", client.MaxRequestPayloadBytes, protocol.MaxPayloadDefault)
1120 + }
1121 + if client.MaxRequestBatchItems != 1 {
1122 + t.Errorf("req_batch = %d, want 1", client.MaxRequestBatchItems)
1123 + }
1124 + if client.MaxResponsePayloadBytes != protocol.MaxPayloadDefault {
1125 + t.Errorf("resp_payload = %d, want %d", client.MaxResponsePayloadBytes, protocol.MaxPayloadDefault)
1126 + }
1127 + if client.MaxResponseBatchItems != 1 {
1128 + t.Errorf("resp_batch = %d, want 1", client.MaxResponseBatchItems)
1129 + }
1130 + // PacketSize should be > 0 (auto-detected from SO_SNDBUF)
1131 + if client.PacketSize == 0 {
1132 + t.Error("packet_size should be auto-detected, got 0")
1133 + }
1134 + // Profile should be baseline
1135 + if client.SelectedProfile != protocol.ProfileBaseline {
1136 + t.Errorf("selected_profile = 0x%x, want 0x%x", client.SelectedProfile, protocol.ProfileBaseline)
1137 + }
1138 +}
1139 +
1140 +// ---------------------------------------------------------------------------
1141 +// Test: highestBit helper
1142 +// ---------------------------------------------------------------------------
1143 +
1144 +func TestHighestBit(t *testing.T) {
1145 + tests := []struct {
1146 + in uint32
1147 + want uint32
1148 + }{
1149 + {0, 0},
1150 + {1, 1},
1151 + {0x03, 0x02},
1152 + {0x07, 0x04},
1153 + {0x80000000, 0x80000000},
1154 + {0xFF, 0x80},
1155 + }
1156 + for _, tc := range tests {
1157 + got := highestBit(tc.in)
1158 + if got != tc.want {
1159 + t.Errorf("highestBit(0x%x) = 0x%x, want 0x%x", tc.in, got, tc.want)
1160 + }
1161 + }
1162 +}
1163 +
1164 +// ---------------------------------------------------------------------------
1165 +// Test: Multiple sequential chunk messages on same session
1166 +// ---------------------------------------------------------------------------
1167 +
1168 +func TestMultipleChunkedMessages(t *testing.T) {
1169 + runDir := testRunDir(t)
1170 + service := uniqueService(t)
1171 + defer os.Remove(filepath.Join(runDir, service+".sock"))
1172 +
1173 + const forcedPacketSize = 96
1174 +
1175 + sCfg := defaultServerConfig()
1176 + sCfg.PacketSize = forcedPacketSize
1177 + sCfg.MaxRequestPayloadBytes = 65536
1178 + sCfg.MaxResponsePayloadBytes = 65536
1179 + listener := startListener(t, runDir, service, sCfg)
1180 + defer listener.Close()
1181 +
1182 + acceptCh := acceptAsync(listener)
1183 +
1184 + cCfg := defaultClientConfig()
1185 + cCfg.PacketSize = forcedPacketSize
1186 + cCfg.MaxRequestPayloadBytes = 65536
1187 + cCfg.MaxResponsePayloadBytes = 65536
1188 + client, err := Connect(runDir, service, &cCfg)
1189 + if err != nil {
1190 + t.Fatalf("Connect failed: %v", err)
1191 + }
1192 + defer client.Close()
1193 +
1194 + sr := <-acceptCh
1195 + if sr.err != nil {
1196 + t.Fatalf("Accept failed: %v", sr.err)
1197 + }
1198 + server := sr.session
1199 + defer server.Close()
1200 +
1201 + // Send 3 chunked messages sequentially
1202 + for i := 0; i < 3; i++ {
1203 + size := 500 + i*200
1204 + payload := make([]byte, size)
1205 + for j := range payload {
1206 + payload[j] = byte((i*31 + j) & 0xFF)
1207 + }
1208 +
1209 + hdr := protocol.Header{
1210 + Kind: protocol.KindRequest,
1211 + Code: protocol.MethodIncrement,
1212 + ItemCount: 1,
1213 + MessageID: uint64(i + 1),
1214 + }
1215 + if err := client.Send(&hdr, payload); err != nil {
1216 + t.Fatalf("Send[%d]: %v", i, err)
1217 + }
1218 +
1219 + buf := make([]byte, forcedPacketSize)
1220 + rHdr, rPayload, err := server.Receive(buf)
1221 + if err != nil {
1222 + t.Fatalf("Receive[%d]: %v", i, err)
1223 + }
1224 + if rHdr.MessageID != uint64(i+1) {
1225 + t.Errorf("msg[%d] message_id = %d, want %d", i, rHdr.MessageID, i+1)
1226 + }
1227 + if !bytes.Equal(rPayload, payload) {
1228 + t.Errorf("msg[%d] payload mismatch (%d vs %d bytes)", i, len(rPayload), len(payload))
1229 + }
1230 + }
1231 +}
1232 +
1233 +// ---------------------------------------------------------------------------
1234 +// Test: Connect to non-existent socket
1235 +// ---------------------------------------------------------------------------
1236 +
1237 +func TestConnectNoServer(t *testing.T) {
1238 + runDir := testRunDir(t)
1239 + service := uniqueService(t)
1240 +
1241 + cCfg := defaultClientConfig()
1242 + _, err := Connect(runDir, service, &cCfg)
1243 + if err == nil {
1244 + t.Fatal("expected error connecting to non-existent socket")
1245 + }
1246 +}
1247 +
1248 +// ---------------------------------------------------------------------------
1249 +// Test: Timeout guard for Accept (server should not block forever in test)
1250 +// ---------------------------------------------------------------------------
1251 +
1252 +func TestAcceptWithTimeout(t *testing.T) {
1253 + runDir := testRunDir(t)
1254 + service := uniqueService(t)
1255 + defer os.Remove(filepath.Join(runDir, service+".sock"))
1256 +
1257 + sCfg := defaultServerConfig()
1258 + listener := startListener(t, runDir, service, sCfg)
1259 + defer listener.Close()
1260 +
1261 + // Start accept in background, connect after a short delay
1262 + done := make(chan bool, 1)
1263 + go func() {
1264 + acceptCh := acceptAsync(listener)
1265 + select {
1266 + case sr := <-acceptCh:
1267 + if sr.err != nil {
1268 + done <- false
1269 + } else {
1270 + sr.session.Close()
1271 + done <- true
1272 + }
1273 + case <-time.After(5 * time.Second):
1274 + done <- false
1275 + }
1276 + }()
1277 +
1278 + // Short delay then connect
1279 + time.Sleep(50 * time.Millisecond)
1280 + cCfg := defaultClientConfig()
1281 + client, err := Connect(runDir, service, &cCfg)
1282 + if err != nil {
1283 + t.Fatalf("Connect: %v", err)
1284 + }
1285 + client.Close()
1286 +
1287 + if ok := <-done; !ok {
1288 + t.Error("accept did not complete successfully")
1289 + }
1290 +}
1291 +
1292 +// ---------------------------------------------------------------------------
1293 +// Test: Invalid service names rejected
1294 +// ---------------------------------------------------------------------------
1295 +
1296 +func TestInvalidServiceName(t *testing.T) {
1297 + runDir := testRunDir(t)
1298 +
1299 + badNames := []string{
1300 + "",
1301 + ".",
1302 + "..",
1303 + "foo/bar",
1304 + "../etc",
1305 + "name with spaces",
1306 + "name\ttab",
1307 + "hello@world",
1308 + }
1309 + goodNames := []string{
1310 + "valid-name",
1311 + "valid_name",
1312 + "valid.name",
1313 + "ValidName123",
1314 + "a",
1315 + }
1316 +
1317 + for _, name := range badNames {
1318 + _, err := Listen(runDir, name, defaultServerConfig())
1319 + if err == nil {
1320 + t.Errorf("Listen(%q) should fail", name)
1321 + }
1322 + _, err = Connect(runDir, name, &ClientConfig{AuthToken: testAuthToken})
1323 + if err == nil {
1324 + t.Errorf("Connect(%q) should fail", name)
1325 + }
1326 + }
1327 +
1328 + for _, name := range goodNames {
1329 + // These should not fail due to validation (may fail because no
1330 + // server is listening, but that's a different error).
1331 + if err := validateServiceName(name); err != nil {
1332 + t.Errorf("validateServiceName(%q) = %v, want nil", name, err)
1333 + }
1334 + }
1335 +}
1336 +
1337 +// ---------------------------------------------------------------------------
1338 +// Test: Pipeline 10 requests, verify all matched by message_id
1339 +// ---------------------------------------------------------------------------
1340 +
1341 +func TestPipeline10(t *testing.T) {
1342 + runDir := testRunDir(t)
1343 + service := uniqueService(t)
1344 + defer os.Remove(filepath.Join(runDir, service+".sock"))
1345 +
1346 + sCfg := defaultServerConfig()
1347 + listener := startListener(t, runDir, service, sCfg)
1348 + defer listener.Close()
1349 +
1350 + acceptCh := acceptAsync(listener)
1351 +
1352 + cCfg := defaultClientConfig()
1353 + client, err := Connect(runDir, service, &cCfg)
1354 + if err != nil {
1355 + t.Fatalf("Connect failed: %v", err)
1356 + }
1357 + defer client.Close()
1358 +
1359 + sr := <-acceptCh
1360 + if sr.err != nil {
1361 + t.Fatalf("Accept failed: %v", sr.err)
1362 + }
1363 + server := sr.session
1364 + defer server.Close()
1365 +
1366 + const count = 10
1367 +
1368 + // Server goroutine: receive and echo
1369 + var wg sync.WaitGroup
1370 + wg.Add(1)
1371 + go func() {
1372 + defer wg.Done()
1373 + buf := make([]byte, 4096)
1374 + for i := 0; i < count; i++ {
1375 + rHdr, rPayload, err := server.Receive(buf)
1376 + if err != nil {
1377 + t.Errorf("server Receive[%d]: %v", i, err)
1378 + return
1379 + }
1380 + resp := protocol.Header{
1381 + Kind: protocol.KindResponse,
1382 + Code: rHdr.Code,
1383 + ItemCount: 1,
1384 + MessageID: rHdr.MessageID,
1385 + }
1386 + if err := server.Send(&resp, rPayload); err != nil {
1387 + t.Errorf("server Send[%d]: %v", i, err)
1388 + return
1389 + }
1390 + }
1391 + }()
1392 +
1393 + // Client sends 10 requests before reading any
1394 + for i := uint64(1); i <= count; i++ {
1395 + payload := make([]byte, 8)
1396 + binary.NativeEndian.PutUint64(payload, i)
1397 + hdr := protocol.Header{
1398 + Kind: protocol.KindRequest,
1399 + Code: protocol.MethodIncrement,
1400 + ItemCount: 1,
1401 + MessageID: i,
1402 + }
1403 + if err := client.Send(&hdr, payload); err != nil {
1404 + t.Fatalf("client Send(%d): %v", i, err)
1405 + }
1406 + }
1407 +
1408 + // Receive 10 responses
1409 + buf := make([]byte, 4096)
1410 + for i := uint64(1); i <= count; i++ {
1411 + rHdr, rPayload, err := client.Receive(buf)
1412 + if err != nil {
1413 + t.Fatalf("client Receive(%d): %v", i, err)
1414 + }
1415 + if rHdr.MessageID != i {
1416 + t.Errorf("message_id = %d, want %d", rHdr.MessageID, i)
1417 + }
1418 + if len(rPayload) != 8 {
1419 + t.Errorf("payload len = %d, want 8", len(rPayload))
1420 + continue
1421 + }
1422 + val := binary.NativeEndian.Uint64(rPayload)
1423 + if val != i {
1424 + t.Errorf("payload value = %d, want %d", val, i)
1425 + }
1426 + }
1427 +
1428 + wg.Wait()
1429 +}
1430 +
1431 +// ---------------------------------------------------------------------------
1432 +// Test: Pipeline 100 requests (stress pipelining)
1433 +// ---------------------------------------------------------------------------
1434 +
1435 +func TestPipeline100(t *testing.T) {
1436 + runDir := testRunDir(t)
1437 + service := uniqueService(t)
1438 + defer os.Remove(filepath.Join(runDir, service+".sock"))
1439 +
1440 + sCfg := defaultServerConfig()
1441 + sCfg.MaxRequestPayloadBytes = 65536
1442 + sCfg.MaxResponsePayloadBytes = 65536
1443 + listener := startListener(t, runDir, service, sCfg)
1444 + defer listener.Close()
1445 +
1446 + acceptCh := acceptAsync(listener)
1447 +
1448 + cCfg := defaultClientConfig()
1449 + cCfg.MaxRequestPayloadBytes = 65536
1450 + cCfg.MaxResponsePayloadBytes = 65536
1451 + client, err := Connect(runDir, service, &cCfg)
1452 + if err != nil {
1453 + t.Fatalf("Connect failed: %v", err)
1454 + }
1455 + defer client.Close()
1456 +
1457 + sr := <-acceptCh
1458 + if sr.err != nil {
1459 + t.Fatalf("Accept failed: %v", sr.err)
1460 + }
1461 + server := sr.session
1462 + defer server.Close()
1463 +
1464 + const count = 100
1465 +
1466 + // Server goroutine
1467 + var wg sync.WaitGroup
1468 + wg.Add(1)
1469 + go func() {
1470 + defer wg.Done()
1471 + buf := make([]byte, 4096)
1472 + for i := 0; i < count; i++ {
1473 + rHdr, rPayload, err := server.Receive(buf)
1474 + if err != nil {
1475 + t.Errorf("server Receive[%d]: %v", i, err)
1476 + return
1477 + }
1478 + resp := protocol.Header{
1479 + Kind: protocol.KindResponse,
1480 + Code: rHdr.Code,
1481 + ItemCount: 1,
1482 + MessageID: rHdr.MessageID,
1483 + }
1484 + if err := server.Send(&resp, rPayload); err != nil {
1485 + t.Errorf("server Send[%d]: %v", i, err)
1486 + return
1487 + }
1488 + }
1489 + }()
1490 +
1491 + // Client sends 100 requests
1492 + for i := uint64(1); i <= count; i++ {
1493 + payload := make([]byte, 8)
1494 + binary.NativeEndian.PutUint64(payload, i)
1495 + hdr := protocol.Header{
1496 + Kind: protocol.KindRequest,
1497 + Code: protocol.MethodIncrement,
1498 + ItemCount: 1,
1499 + MessageID: i,
1500 + }
1501 + if err := client.Send(&hdr, payload); err != nil {
1502 + t.Fatalf("client Send(%d): %v", i, err)
1503 + }
1504 + }
1505 +
1506 + // Receive 100 responses
1507 + buf := make([]byte, 4096)
1508 + for i := uint64(1); i <= count; i++ {
1509 + rHdr, rPayload, err := client.Receive(buf)
1510 + if err != nil {
1511 + t.Fatalf("client Receive(%d): %v", i, err)
1512 + }
1513 + if rHdr.MessageID != i {
1514 + t.Errorf("message_id = %d, want %d", rHdr.MessageID, i)
1515 + }
1516 + val := binary.NativeEndian.Uint64(rPayload)
1517 + if val != i {
1518 + t.Errorf("payload value = %d, want %d", val, i)
1519 + }
1520 + }
1521 +
1522 + wg.Wait()
1523 +}
1524 +
1525 +// ---------------------------------------------------------------------------
1526 +// Test: Pipeline with mixed message sizes
1527 +// ---------------------------------------------------------------------------
1528 +
1529 +func TestPipelineMixedSizes(t *testing.T) {
1530 + runDir := testRunDir(t)
1531 + service := uniqueService(t)
1532 + defer os.Remove(filepath.Join(runDir, service+".sock"))
1533 +
1534 + sCfg := defaultServerConfig()
1535 + sCfg.MaxRequestPayloadBytes = 65536
1536 + sCfg.MaxResponsePayloadBytes = 65536
1537 + listener := startListener(t, runDir, service, sCfg)
1538 + defer listener.Close()
1539 +
1540 + acceptCh := acceptAsync(listener)
1541 +
1542 + cCfg := defaultClientConfig()
1543 + cCfg.MaxRequestPayloadBytes = 65536
1544 + cCfg.MaxResponsePayloadBytes = 65536
1545 + client, err := Connect(runDir, service, &cCfg)
1546 + if err != nil {
1547 + t.Fatalf("Connect failed: %v", err)
1548 + }
1549 + defer client.Close()
1550 +
1551 + sr := <-acceptCh
1552 + if sr.err != nil {
1553 + t.Fatalf("Accept failed: %v", sr.err)
1554 + }
1555 + server := sr.session
1556 + defer server.Close()
1557 +
1558 + sizes := []int{8, 256, 1024, 8, 256, 1024, 8, 256, 1024}
1559 + count := len(sizes)
1560 +
1561 + // Server goroutine
1562 + var wg sync.WaitGroup
1563 + wg.Add(1)
1564 + go func() {
1565 + defer wg.Done()
1566 + buf := make([]byte, 8192)
1567 + for i := 0; i < count; i++ {
1568 + rHdr, rPayload, err := server.Receive(buf)
1569 + if err != nil {
1570 + t.Errorf("server Receive[%d]: %v", i, err)
1571 + return
1572 + }
1573 + resp := protocol.Header{
1574 + Kind: protocol.KindResponse,
1575 + Code: rHdr.Code,
1576 + ItemCount: 1,
1577 + MessageID: rHdr.MessageID,
1578 + }
1579 + if err := server.Send(&resp, rPayload); err != nil {
1580 + t.Errorf("server Send[%d]: %v", i, err)
1581 + return
1582 + }
1583 + }
1584 + }()
1585 +
1586 + // Client sends all messages
1587 + for i, sz := range sizes {
1588 + payload := make([]byte, sz)
1589 + for j := range payload {
1590 + payload[j] = byte((i*37 + j) & 0xFF)
1591 + }
1592 + hdr := protocol.Header{
1593 + Kind: protocol.KindRequest,
1594 + Code: protocol.MethodIncrement,
1595 + ItemCount: 1,
1596 + MessageID: uint64(i + 1),
1597 + }
1598 + if err := client.Send(&hdr, payload); err != nil {
1599 + t.Fatalf("client Send[%d]: %v", i, err)
1600 + }
1601 + }
1602 +
1603 + // Receive all responses
1604 + buf := make([]byte, 8192)
1605 + for i, sz := range sizes {
1606 + rHdr, rPayload, err := client.Receive(buf)
1607 + if err != nil {
1608 + t.Fatalf("client Receive[%d]: %v", i, err)
1609 + }
1610 + if rHdr.MessageID != uint64(i+1) {
1611 + t.Errorf("[%d] message_id = %d, want %d", i, rHdr.MessageID, i+1)
1612 + }
1613 + if len(rPayload) != sz {
1614 + t.Errorf("[%d] payload len = %d, want %d", i, len(rPayload), sz)
1615 + continue
1616 + }
1617 + expected := make([]byte, sz)
1618 + for j := range expected {
1619 + expected[j] = byte((i*37 + j) & 0xFF)
1620 + }
1621 + if !bytes.Equal(rPayload, expected) {
1622 + t.Errorf("[%d] payload data mismatch", i)
1623 + }
1624 + }
1625 +
1626 + wg.Wait()
1627 +}
1628 +
1629 +// ---------------------------------------------------------------------------
1630 +// Test: Pipeline with chunked messages (> packet_size)
1631 +// ---------------------------------------------------------------------------
1632 +
1633 +func TestPipelineChunked(t *testing.T) {
1634 + runDir := testRunDir(t)
1635 + service := uniqueService(t)
1636 + defer os.Remove(filepath.Join(runDir, service+".sock"))
1637 +
1638 + const forcedPacketSize = 128
1639 +
1640 + sCfg := defaultServerConfig()
1641 + sCfg.PacketSize = forcedPacketSize
1642 + sCfg.MaxRequestPayloadBytes = 65536
1643 + sCfg.MaxResponsePayloadBytes = 65536
1644 + listener := startListener(t, runDir, service, sCfg)
1645 + defer listener.Close()
1646 +
1647 + acceptCh := acceptAsync(listener)
1648 +
1649 + cCfg := defaultClientConfig()
1650 + cCfg.PacketSize = forcedPacketSize
1651 + cCfg.MaxRequestPayloadBytes = 65536
1652 + cCfg.MaxResponsePayloadBytes = 65536
1653 + client, err := Connect(runDir, service, &cCfg)
1654 + if err != nil {
1655 + t.Fatalf("Connect failed: %v", err)
1656 + }
1657 + defer client.Close()
1658 +
1659 + sr := <-acceptCh
1660 + if sr.err != nil {
1661 + t.Fatalf("Accept failed: %v", sr.err)
1662 + }
1663 + server := sr.session
1664 + defer server.Close()
1665 +
1666 + sizes := []int{200, 500, 300, 800, 150}
1667 + count := len(sizes)
1668 +
1669 + // Server goroutine
1670 + var wg sync.WaitGroup
1671 + wg.Add(1)
1672 + go func() {
1673 + defer wg.Done()
1674 + buf := make([]byte, forcedPacketSize)
1675 + for i := 0; i < count; i++ {
1676 + rHdr, rPayload, err := server.Receive(buf)
1677 + if err != nil {
1678 + t.Errorf("server Receive[%d]: %v", i, err)
1679 + return
1680 + }
1681 + resp := protocol.Header{
1682 + Kind: protocol.KindResponse,
1683 + Code: rHdr.Code,
1684 + ItemCount: 1,
1685 + MessageID: rHdr.MessageID,
1686 + }
1687 + if err := server.Send(&resp, rPayload); err != nil {
1688 + t.Errorf("server Send[%d]: %v", i, err)
1689 + return
1690 + }
1691 + }
1692 + }()
1693 +
1694 + // Client sends all chunked messages
1695 + for i, sz := range sizes {
1696 + payload := make([]byte, sz)
1697 + for j := range payload {
1698 + payload[j] = byte((i + j) & 0xFF)
1699 + }
1700 + hdr := protocol.Header{
1701 + Kind: protocol.KindRequest,
1702 + Code: protocol.MethodIncrement,
1703 + ItemCount: 1,
1704 + MessageID: uint64(i + 1),
1705 + }
1706 + if err := client.Send(&hdr, payload); err != nil {
1707 + t.Fatalf("client Send[%d]: %v", i, err)
1708 + }
1709 + }
1710 +
1711 + // Receive all responses
1712 + buf := make([]byte, forcedPacketSize)
1713 + for i, sz := range sizes {
1714 + rHdr, rPayload, err := client.Receive(buf)
1715 + if err != nil {
1716 + t.Fatalf("client Receive[%d]: %v", i, err)
1717 + }
1718 + if rHdr.MessageID != uint64(i+1) {
1719 + t.Errorf("[%d] message_id = %d, want %d", i, rHdr.MessageID, i+1)
1720 + }
1721 + if len(rPayload) != sz {
1722 + t.Errorf("[%d] payload len = %d, want %d", i, len(rPayload), sz)
1723 + continue
1724 + }
1725 + expected := make([]byte, sz)
1726 + for j := range expected {
1727 + expected[j] = byte((i + j) & 0xFF)
1728 + }
1729 + if !bytes.Equal(rPayload, expected) {
1730 + t.Errorf("[%d] chunked payload data mismatch", i)
1731 + }
1732 + }
1733 +
1734 + wg.Wait()
1735 +}
1736 +
1737 +// ---------------------------------------------------------------------------
1738 +// Test: Duplicate message_id rejected on send
1739 +// ---------------------------------------------------------------------------
1740 +
1741 +func TestDuplicateMessageID(t *testing.T) {
1742 + runDir := testRunDir(t)
1743 + service := uniqueService(t)
1744 + defer os.Remove(filepath.Join(runDir, service+".sock"))
1745 +
1746 + sCfg := defaultServerConfig()
1747 + listener := startListener(t, runDir, service, sCfg)
1748 + defer listener.Close()
1749 +
1750 + acceptCh := acceptAsync(listener)
1751 +
1752 + cCfg := defaultClientConfig()
1753 + client, err := Connect(runDir, service, &cCfg)
1754 + if err != nil {
1755 + t.Fatalf("Connect: %v", err)
1756 + }
1757 + defer client.Close()
1758 +
1759 + sr := <-acceptCh
1760 + if sr.err != nil {
1761 + t.Fatalf("Accept: %v", sr.err)
1762 + }
1763 + defer sr.session.Close()
1764 +
1765 + hdr := protocol.Header{
1766 + Kind: protocol.KindRequest,
1767 + Code: protocol.MethodIncrement,
1768 + ItemCount: 1,
1769 + MessageID: 42,
1770 + }
1771 +
1772 + // First send should succeed
1773 + if err := client.Send(&hdr, []byte("hello")); err != nil {
1774 + t.Fatalf("first Send: %v", err)
1775 + }
1776 +
1777 + // Second send with same message_id should fail
1778 + hdr2 := protocol.Header{
1779 + Kind: protocol.KindRequest,
1780 + Code: protocol.MethodIncrement,
1781 + ItemCount: 1,
1782 + MessageID: 42,
1783 + }
1784 + err = client.Send(&hdr2, []byte("hello"))
1785 + if err == nil {
1786 + t.Fatal("duplicate message_id should fail")
1787 + }
1788 + if !errors.Is(err, ErrDuplicateMsgID) {
1789 + t.Errorf("error = %v, want ErrDuplicateMsgID", err)
1790 + }
1791 +}
1792 +
1793 +// ---------------------------------------------------------------------------
1794 +// Test: Unknown message_id rejected on receive
1795 +// ---------------------------------------------------------------------------
1796 +
1797 +func TestUnknownMessageID(t *testing.T) {
1798 + runDir := testRunDir(t)
1799 + service := uniqueService(t)
1800 + defer os.Remove(filepath.Join(runDir, service+".sock"))
1801 +
1802 + sCfg := defaultServerConfig()
1803 + listener := startListener(t, runDir, service, sCfg)
1804 + defer listener.Close()
1805 +
1806 + acceptCh := acceptAsync(listener)
1807 +
1808 + cCfg := defaultClientConfig()
1809 + client, err := Connect(runDir, service, &cCfg)
1810 + if err != nil {
1811 + t.Fatalf("Connect: %v", err)
1812 + }
1813 + defer client.Close()
1814 +
1815 + sr := <-acceptCh
1816 + if sr.err != nil {
1817 + t.Fatalf("Accept: %v", sr.err)
1818 + }
1819 + server := sr.session
1820 + defer server.Close()
1821 +
1822 + // Client sends request with message_id=10
1823 + hdr := protocol.Header{
1824 + Kind: protocol.KindRequest,
1825 + Code: protocol.MethodIncrement,
1826 + ItemCount: 1,
1827 + MessageID: 10,
1828 + }
1829 + if err := client.Send(&hdr, []byte("x")); err != nil {
1830 + t.Fatalf("Send: %v", err)
1831 + }
1832 +
1833 + // Server receives
1834 + buf := make([]byte, 4096)
1835 + rHdr, rPayload, err := server.Receive(buf)
1836 + if err != nil {
1837 + t.Fatalf("server Receive: %v", err)
1838 + }
1839 +
1840 + // Server responds with WRONG message_id (999)
1841 + resp := protocol.Header{
1842 + Kind: protocol.KindResponse,
1843 + Code: rHdr.Code,
1844 + ItemCount: 1,
1845 + MessageID: 999, // not in-flight
1846 + }
1847 + if err := server.Send(&resp, rPayload); err != nil {
1848 + t.Fatalf("server Send: %v", err)
1849 + }
1850 +
1851 + // Client receive should fail with unknown message_id
1852 + _, _, err = client.Receive(buf)
1853 + if err == nil {
1854 + t.Fatal("expected unknown message_id error")
1855 + }
1856 + if !errors.Is(err, ErrUnknownMsgID) {
1857 + t.Errorf("error = %v, want ErrUnknownMsgID", err)
1858 + }
1859 +}
src/go/pkg/netipc/transport/windows/pipe.go new
+1267
@@ -0,0 +1,1267 @@
1 +//go:build windows
2 +
3 +// Package windows implements the L1 Windows Named Pipe transport.
4 +//
5 +// Connection lifecycle, handshake with profile/limit negotiation,
6 +// and send/receive with transparent chunking over Win32 Named Pipes
7 +// in message mode. Wire-compatible with the C and Rust implementations.
8 +//
9 +// Pure Go — no cgo. Works with CGO_ENABLED=0.
10 +package windows
11 +
12 +import (
13 + "encoding/binary"
14 + "errors"
15 + "fmt"
16 + "sync"
17 + "sync/atomic"
18 + "syscall"
19 + "time"
20 + "unicode/utf16"
21 + "unsafe"
22 +
23 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
24 +)
25 +
26 +// ---------------------------------------------------------------------------
27 +// Constants
28 +// ---------------------------------------------------------------------------
29 +
30 +const (
31 + defaultBatchItems uint32 = 1
32 + defaultPacketSize uint32 = 65536
33 + defaultPipeBufSize uint32 = 65536
34 + helloPayloadSize = 44
35 + helloAckPayloadSize = 48
36 + maxPipeNameChars = 256
37 +
38 + // FNV-1a 64-bit constants
39 + fnv1aOffsetBasis uint64 = 0xcbf29ce484222325
40 + fnv1aPrime uint64 = 0x00000100000001B3
41 +)
42 +
43 +// Win32 constants
44 +const (
45 + _PIPE_ACCESS_DUPLEX = 0x00000003
46 + _FILE_FLAG_FIRST_PIPE_INSTANCE = 0x00080000
47 + _PIPE_TYPE_MESSAGE = 0x00000004
48 + _PIPE_READMODE_MESSAGE = 0x00000002
49 + _PIPE_WAIT = 0x00000000
50 + _PIPE_UNLIMITED_INSTANCES = 255
51 + _GENERIC_READ = 0x80000000
52 + _GENERIC_WRITE = 0x40000000
53 + _OPEN_EXISTING = 3
54 +
55 + _ERROR_PIPE_CONNECTED = 535
56 + _ERROR_BROKEN_PIPE = 109
57 + _ERROR_NO_DATA = 232
58 + _ERROR_PIPE_NOT_CONNECTED = 233
59 + _ERROR_ACCESS_DENIED = 5
60 + _ERROR_PIPE_BUSY = 231
61 +)
62 +
63 +// ---------------------------------------------------------------------------
64 +// Errors
65 +// ---------------------------------------------------------------------------
66 +
67 +var (
68 + ErrPipeName = errors.New("pipe name derivation failed")
69 + ErrCreatePipe = errors.New("CreateNamedPipe failed")
70 + ErrConnect = errors.New("connect failed")
71 + ErrAccept = errors.New("accept failed")
72 + ErrSend = errors.New("send failed")
73 + ErrRecv = errors.New("recv failed or peer disconnected")
74 + ErrHandshake = errors.New("handshake protocol error")
75 + ErrAuthFailed = errors.New("authentication token rejected")
76 + ErrNoProfile = errors.New("no common transport profile")
77 + ErrIncompatible = errors.New("protocol or layout version mismatch")
78 + ErrProtocol = errors.New("wire protocol violation")
79 + ErrAddrInUse = errors.New("pipe name already in use by live server")
80 + ErrChunk = errors.New("chunk header mismatch")
81 + ErrLimitExceeded = errors.New("negotiated limit exceeded")
82 + ErrBadParam = errors.New("invalid argument")
83 + ErrDuplicateMsgID = errors.New("duplicate message_id")
84 + ErrUnknownMsgID = errors.New("unknown response message_id")
85 + ErrDisconnected = errors.New("peer disconnected")
86 +)
87 +
88 +func wrapErr(sentinel error, detail string) error {
89 + return fmt.Errorf("%w: %s", sentinel, detail)
90 +}
91 +
92 +func headerVersionIncompatible(buf []byte, expectedCode uint16) bool {
93 + if len(buf) < protocol.HeaderSize {
94 + return false
95 + }
96 +
97 + return binary.NativeEndian.Uint32(buf[0:4]) == protocol.MagicMsg &&
98 + binary.NativeEndian.Uint16(buf[4:6]) != protocol.Version &&
99 + binary.NativeEndian.Uint16(buf[6:8]) == protocol.HeaderLen &&
100 + binary.NativeEndian.Uint16(buf[8:10]) == protocol.KindControl &&
101 + binary.NativeEndian.Uint16(buf[12:14]) == expectedCode
102 +}
103 +
104 +func helloLayoutIncompatible(buf []byte) bool {
105 + return len(buf) >= 2 && binary.NativeEndian.Uint16(buf[0:2]) != 1
106 +}
107 +
108 +func helloAckLayoutIncompatible(buf []byte) bool {
109 + return len(buf) >= 2 && binary.NativeEndian.Uint16(buf[0:2]) != 1
110 +}
111 +
112 +// ---------------------------------------------------------------------------
113 +// Win32 syscall imports (pure Go, no cgo)
114 +// ---------------------------------------------------------------------------
115 +
116 +var (
117 + modkernel32 = syscall.NewLazyDLL("kernel32.dll")
118 +
119 + procCreateNamedPipeW = modkernel32.NewProc("CreateNamedPipeW")
120 + procConnectNamedPipe = modkernel32.NewProc("ConnectNamedPipe")
121 + procDisconnectNamedPipe = modkernel32.NewProc("DisconnectNamedPipe")
122 + procFlushFileBuffers = modkernel32.NewProc("FlushFileBuffers")
123 + procPeekNamedPipe = modkernel32.NewProc("PeekNamedPipe")
124 + procSetNamedPipeHandleState = modkernel32.NewProc("SetNamedPipeHandleState")
125 + procSwitchToThread = modkernel32.NewProc("SwitchToThread")
126 +)
127 +
128 +func createNamedPipe(name *uint16, openMode, pipeMode, maxInstances, outBufSize, inBufSize, defaultTimeout uint32) (syscall.Handle, error) {
129 + r, _, err := procCreateNamedPipeW.Call(
130 + uintptr(unsafe.Pointer(name)),
131 + uintptr(openMode),
132 + uintptr(pipeMode),
133 + uintptr(maxInstances),
134 + uintptr(outBufSize),
135 + uintptr(inBufSize),
136 + uintptr(defaultTimeout),
137 + 0, // NULL security attributes
138 + )
139 + handle := syscall.Handle(r)
140 + if handle == syscall.InvalidHandle {
141 + return handle, err
142 + }
143 + return handle, nil
144 +}
145 +
146 +func connectNamedPipe(handle syscall.Handle) error {
147 + r, _, err := procConnectNamedPipe.Call(uintptr(handle), 0)
148 + if r == 0 {
149 + return err
150 + }
151 + return nil
152 +}
153 +
154 +func disconnectNamedPipe(handle syscall.Handle) {
155 + procDisconnectNamedPipe.Call(uintptr(handle))
156 +}
157 +
158 +func flushFileBuffers(handle syscall.Handle) {
159 + procFlushFileBuffers.Call(uintptr(handle))
160 +}
161 +
162 +func peekNamedPipeAvailable(handle syscall.Handle) (uint32, error) {
163 + var available uint32
164 + r, _, err := procPeekNamedPipe.Call(
165 + uintptr(handle),
166 + 0,
167 + 0,
168 + 0,
169 + uintptr(unsafe.Pointer(&available)),
170 + 0,
171 + )
172 + if r == 0 {
173 + return 0, err
174 + }
175 + return available, nil
176 +}
177 +
178 +func setNamedPipeHandleState(handle syscall.Handle, mode *uint32) error {
179 + r, _, err := procSetNamedPipeHandleState.Call(
180 + uintptr(handle),
181 + uintptr(unsafe.Pointer(mode)),
182 + 0, 0,
183 + )
184 + if r == 0 {
185 + return err
186 + }
187 + return nil
188 +}
189 +
190 +// ---------------------------------------------------------------------------
191 +// FNV-1a 64-bit hash
192 +// ---------------------------------------------------------------------------
193 +
194 +// FNV1a64 computes the FNV-1a 64-bit hash of data.
195 +func FNV1a64(data []byte) uint64 {
196 + hash := fnv1aOffsetBasis
197 + for _, b := range data {
198 + hash ^= uint64(b)
199 + hash *= fnv1aPrime
200 + }
201 + return hash
202 +}
203 +
204 +// ---------------------------------------------------------------------------
205 +// Service name validation
206 +// ---------------------------------------------------------------------------
207 +
208 +func validateServiceName(name string) error {
209 + if name == "" {
210 + return wrapErr(ErrBadParam, "empty service name")
211 + }
212 + if name == "." || name == ".." {
213 + return wrapErr(ErrBadParam, "service name cannot be '.' or '..'")
214 + }
215 + for i := 0; i < len(name); i++ {
216 + c := name[i]
217 + if (c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') ||
218 + (c >= '0' && c <= '9') || c == '.' || c == '_' || c == '-' {
219 + continue
220 + }
221 + return wrapErr(ErrBadParam, fmt.Sprintf("service name contains invalid character: %q", c))
222 + }
223 + return nil
224 +}
225 +
226 +// ---------------------------------------------------------------------------
227 +// Pipe name derivation
228 +// ---------------------------------------------------------------------------
229 +
230 +// BuildPipeName constructs the Named Pipe path from run_dir and service_name.
231 +// Returns the pipe name as a NUL-terminated UTF-16 slice.
232 +func BuildPipeName(runDir, serviceName string) ([]uint16, error) {
233 + if err := validateServiceName(serviceName); err != nil {
234 + return nil, err
235 + }
236 +
237 + hash := FNV1a64([]byte(runDir))
238 + narrow := fmt.Sprintf(`\\.\pipe\netipc-%016x-%s`, hash, serviceName)
239 +
240 + if len(narrow) >= maxPipeNameChars {
241 + return nil, wrapErr(ErrPipeName, "pipe name too long")
242 + }
243 +
244 + return utf16.Encode(append([]rune(narrow), 0)), nil
245 +}
246 +
247 +// ---------------------------------------------------------------------------
248 +// Internal helpers
249 +// ---------------------------------------------------------------------------
250 +
251 +func applyDefault(val, def uint32) uint32 {
252 + if val == 0 {
253 + return def
254 + }
255 + return val
256 +}
257 +
258 +func minU32(a, b uint32) uint32 {
259 + if a < b {
260 + return a
261 + }
262 + return b
263 +}
264 +
265 +func maxU32(a, b uint32) uint32 {
266 + if a > b {
267 + return a
268 + }
269 + return b
270 +}
271 +
272 +func pipeBufferSize(packetSize uint32) uint32 {
273 + // The protocol packet size controls logical framing and chunk size. The
274 + // underlying pipe quota must stay large enough for full-duplex pipelining
275 + // even when tests force a tiny protocol packet size.
276 + return maxU32(applyDefault(packetSize, defaultPipeBufSize), defaultPipeBufSize)
277 +}
278 +
279 +func highestBit(mask uint32) uint32 {
280 + if mask == 0 {
281 + return 0
282 + }
283 + bit := uint32(1) << 31
284 + for bit&mask == 0 {
285 + bit >>= 1
286 + }
287 + return bit
288 +}
289 +
290 +func isDisconnectError(err error) bool {
291 + errno, ok := err.(syscall.Errno)
292 + if !ok {
293 + return false
294 + }
295 + return errno == _ERROR_BROKEN_PIPE ||
296 + errno == _ERROR_NO_DATA ||
297 + errno == _ERROR_PIPE_NOT_CONNECTED
298 +}
299 +
300 +// ---------------------------------------------------------------------------
301 +// Role
302 +// ---------------------------------------------------------------------------
303 +
304 +// Role distinguishes client vs server sessions.
305 +type Role int
306 +
307 +const (
308 + RoleClient Role = 1
309 + RoleServer Role = 2
310 +)
311 +
312 +// ---------------------------------------------------------------------------
313 +// Configuration
314 +// ---------------------------------------------------------------------------
315 +
316 +// ClientConfig configures a client connection.
317 +type ClientConfig struct {
318 + SupportedProfiles uint32
319 + PreferredProfiles uint32
320 + MaxRequestPayloadBytes uint32 // 0 = use default (1024)
321 + MaxRequestBatchItems uint32 // 0 = use default (1)
322 + MaxResponsePayloadBytes uint32
323 + MaxResponseBatchItems uint32
324 + AuthToken uint64
325 + PacketSize uint32 // 0 = use default (65536)
326 +}
327 +
328 +// ServerConfig configures a listener and its accepted sessions.
329 +type ServerConfig struct {
330 + SupportedProfiles uint32
331 + PreferredProfiles uint32
332 + MaxRequestPayloadBytes uint32
333 + MaxRequestBatchItems uint32
334 + MaxResponsePayloadBytes uint32
335 + MaxResponseBatchItems uint32
336 + AuthToken uint64
337 + PacketSize uint32 // 0 = use default (65536)
338 +}
339 +
340 +// ---------------------------------------------------------------------------
341 +// Low-level I/O
342 +// ---------------------------------------------------------------------------
343 +
344 +func rawWrite(handle syscall.Handle, data []byte) error {
345 + var written uint32
346 + err := syscall.WriteFile(handle, data, &written, nil)
347 + if err != nil {
348 + if isDisconnectError(err) {
349 + return ErrDisconnected
350 + }
351 + return wrapErr(ErrSend, err.Error())
352 + }
353 + if written != uint32(len(data)) {
354 + return wrapErr(ErrSend, fmt.Sprintf("short write: %d/%d", written, len(data)))
355 + }
356 + return nil
357 +}
358 +
359 +func rawSendMsg(handle syscall.Handle, msg []byte) error {
360 + return rawWrite(handle, msg)
361 +}
362 +
363 +func rawRecv(handle syscall.Handle, buf []byte) (int, error) {
364 + var read uint32
365 + err := syscall.ReadFile(handle, buf, &read, nil)
366 + if err != nil {
367 + // ERROR_MORE_DATA (234): message mode pipe message is larger
368 + // than the buffer. The data read so far is valid; the
369 + // remaining data can be read with another ReadFile call.
370 + // For our protocol this should not happen if the buffer is
371 + // sized correctly, but treat it as a successful partial read
372 + // rather than a fatal error.
373 + if err == syscall.Errno(234) {
374 + if read > 0 {
375 + return int(read), nil
376 + }
377 + }
378 + if isDisconnectError(err) {
379 + return 0, ErrDisconnected
380 + }
381 + return 0, wrapErr(ErrRecv, err.Error())
382 + }
383 + if read == 0 {
384 + return 0, ErrDisconnected
385 + }
386 + return int(read), nil
387 +}
388 +
389 +// ---------------------------------------------------------------------------
390 +// Session
391 +// ---------------------------------------------------------------------------
392 +
393 +// Session is a connected Named Pipe session (client or server side).
394 +type Session struct {
395 + handle syscall.Handle
396 + role Role
397 +
398 + // Negotiated limits
399 + MaxRequestPayloadBytes uint32
400 + MaxRequestBatchItems uint32
401 + MaxResponsePayloadBytes uint32
402 + MaxResponseBatchItems uint32
403 + PacketSize uint32
404 + SelectedProfile uint32
405 + SessionID uint64
406 +
407 + // Internal receive buffer for chunked reassembly
408 + recvBuf []byte
409 +
410 + // Reusable packet scratch buffers for send/receive chunk assembly.
411 + sendBuf []byte
412 + pktBuf []byte
413 +
414 + // In-flight message_id set (client-side only)
415 + inflightIDs map[uint64]struct{}
416 +}
417 +
418 +func (s *Session) failAllInflight() {
419 + if s.role != RoleClient || len(s.inflightIDs) == 0 {
420 + return
421 + }
422 + clear(s.inflightIDs)
423 +}
424 +
425 +// Handle returns the raw HANDLE for WaitForSingleObject integration.
426 +func (s *Session) Handle() syscall.Handle {
427 + return s.handle
428 +}
429 +
430 +// Role returns the session role.
431 +func (s *Session) GetRole() Role {
432 + return s.role
433 +}
434 +
435 +// WaitReadable waits until bytes are available to read or the timeout expires.
436 +func (s *Session) WaitReadable(timeoutMs uint32) (bool, error) {
437 + if s.handle == syscall.InvalidHandle {
438 + return false, wrapErr(ErrBadParam, "session closed")
439 + }
440 +
441 + deadline := time.Now().Add(time.Duration(timeoutMs) * time.Millisecond)
442 + yielded := false
443 + for {
444 + available, err := peekNamedPipeAvailable(s.handle)
445 + if err != nil {
446 + if isDisconnectError(err) {
447 + s.failAllInflight()
448 + return false, ErrDisconnected
449 + }
450 + return false, wrapErr(ErrRecv, err.Error())
451 + }
452 + if available > 0 {
453 + return true, nil
454 + }
455 + if !time.Now().Before(deadline) {
456 + return false, nil
457 + }
458 + if !yielded {
459 + yielded = true
460 + for i := 0; i < 256; i++ {
461 + procSwitchToThread.Call()
462 + available, err = peekNamedPipeAvailable(s.handle)
463 + if err != nil {
464 + if isDisconnectError(err) {
465 + s.failAllInflight()
466 + return false, ErrDisconnected
467 + }
468 + return false, wrapErr(ErrRecv, err.Error())
469 + }
470 + if available > 0 {
471 + return true, nil
472 + }
473 + if !time.Now().Before(deadline) {
474 + return false, nil
475 + }
476 + }
477 + continue
478 + }
479 + time.Sleep(time.Millisecond)
480 + }
481 +}
482 +
483 +// Close closes the session and releases resources.
484 +func (s *Session) Close() {
485 + if s.handle != syscall.InvalidHandle {
486 + if s.role == RoleServer {
487 + // Flush before server-side disconnect so the client can
488 + // consume any final response bytes already written.
489 + flushFileBuffers(s.handle)
490 + disconnectNamedPipe(s.handle)
491 + }
492 + syscall.CloseHandle(s.handle)
493 + s.handle = syscall.InvalidHandle
494 + }
495 + s.recvBuf = nil
496 + s.sendBuf = nil
497 + s.pktBuf = nil
498 + s.failAllInflight()
499 +}
500 +
501 +// Connect establishes a session to a server pipe derived from runDir + serviceName.
502 +func Connect(runDir, serviceName string, config *ClientConfig) (*Session, error) {
503 + pipeName, err := BuildPipeName(runDir, serviceName)
504 + if err != nil {
505 + return nil, err
506 + }
507 +
508 + handle, err := syscall.CreateFile(
509 + &pipeName[0],
510 + _GENERIC_READ|_GENERIC_WRITE,
511 + 0,
512 + nil,
513 + _OPEN_EXISTING,
514 + 0,
515 + 0,
516 + )
517 + if err != nil {
518 + return nil, wrapErr(ErrConnect, err.Error())
519 + }
520 +
521 + // Set read mode to message mode
522 + mode := uint32(_PIPE_READMODE_MESSAGE)
523 + if err := setNamedPipeHandleState(handle, &mode); err != nil {
524 + syscall.CloseHandle(handle)
525 + return nil, wrapErr(ErrConnect, "SetNamedPipeHandleState: "+err.Error())
526 + }
527 +
528 + session, herr := clientHandshake(handle, config)
529 + if herr != nil {
530 + syscall.CloseHandle(handle)
531 + return nil, herr
532 + }
533 + return session, nil
534 +}
535 +
536 +// Send sends one logical message. Fills magic/version/header_len/payload_len.
537 +func (s *Session) Send(hdr *protocol.Header, payload []byte) error {
538 + if s.handle == syscall.InvalidHandle {
539 + return wrapErr(ErrBadParam, "session closed")
540 + }
541 +
542 + // Client-side: track in-flight message_ids
543 + if s.role == RoleClient && hdr.Kind == protocol.KindRequest {
544 + if s.inflightIDs == nil {
545 + s.inflightIDs = make(map[uint64]struct{})
546 + }
547 + if _, exists := s.inflightIDs[hdr.MessageID]; exists {
548 + return wrapErr(ErrDuplicateMsgID, fmt.Sprintf("message_id %d", hdr.MessageID))
549 + }
550 + s.inflightIDs[hdr.MessageID] = struct{}{}
551 + }
552 +
553 + // Fill envelope
554 + hdr.Magic = protocol.MagicMsg
555 + hdr.Version = protocol.Version
556 + hdr.HeaderLen = protocol.HeaderLen
557 + hdr.PayloadLen = uint32(len(payload))
558 +
559 + tracked := s.role == RoleClient && hdr.Kind == protocol.KindRequest
560 +
561 + sendErr := s.sendInner(hdr, payload)
562 +
563 + if sendErr != nil && tracked {
564 + if errors.Is(sendErr, ErrDisconnected) {
565 + s.failAllInflight()
566 + } else {
567 + delete(s.inflightIDs, hdr.MessageID)
568 + }
569 + }
570 +
571 + return sendErr
572 +}
573 +
574 +func (s *Session) sendInner(hdr *protocol.Header, payload []byte) error {
575 + totalMsg := protocol.HeaderSize + len(payload)
576 +
577 + // Single packet?
578 + if totalMsg <= int(s.PacketSize) {
579 + msg := ensurePipeScratchBuf(&s.sendBuf, totalMsg)
580 + hdr.Encode(msg[:protocol.HeaderSize])
581 + copy(msg[protocol.HeaderSize:], payload)
582 + return rawSendMsg(s.handle, msg[:totalMsg])
583 + }
584 +
585 + // Chunked send
586 + chunkPayloadBudget := int(s.PacketSize) - protocol.HeaderSize
587 + if chunkPayloadBudget <= 0 {
588 + return wrapErr(ErrBadParam, "packet_size too small")
589 + }
590 +
591 + firstChunkPayload := len(payload)
592 + if firstChunkPayload > chunkPayloadBudget {
593 + firstChunkPayload = chunkPayloadBudget
594 + }
595 +
596 + remainingAfterFirst := len(payload) - firstChunkPayload
597 + continuationChunks := uint32(0)
598 + if remainingAfterFirst > 0 {
599 + continuationChunks = uint32((remainingAfterFirst + chunkPayloadBudget - 1) / chunkPayloadBudget)
600 + }
601 + chunkCount := 1 + continuationChunks
602 +
603 + // First chunk
604 + pktBuf := ensurePipeScratchBuf(&s.sendBuf, int(s.PacketSize))
605 + hdr.Encode(pktBuf[:protocol.HeaderSize])
606 + copy(pktBuf[protocol.HeaderSize:], payload[:firstChunkPayload])
607 + if err := rawSendMsg(s.handle, pktBuf[:protocol.HeaderSize+firstChunkPayload]); err != nil {
608 + return err
609 + }
610 +
611 + // Continuation chunks
612 + offset := firstChunkPayload
613 + for ci := uint32(1); ci < chunkCount; ci++ {
614 + remaining := len(payload) - offset
615 + thisChunk := remaining
616 + if thisChunk > chunkPayloadBudget {
617 + thisChunk = chunkPayloadBudget
618 + }
619 +
620 + chk := protocol.ChunkHeader{
621 + Magic: protocol.MagicChunk,
622 + Version: protocol.Version,
623 + Flags: 0,
624 + MessageID: hdr.MessageID,
625 + TotalMessageLen: uint32(totalMsg),
626 + ChunkIndex: ci,
627 + ChunkCount: chunkCount,
628 + ChunkPayloadLen: uint32(thisChunk),
629 + }
630 +
631 + chk.Encode(pktBuf[:protocol.HeaderSize])
632 + copy(pktBuf[protocol.HeaderSize:], payload[offset:offset+thisChunk])
633 + if err := rawSendMsg(s.handle, pktBuf[:protocol.HeaderSize+thisChunk]); err != nil {
634 + return err
635 + }
636 +
637 + offset += thisChunk
638 + }
639 +
640 + return nil
641 +}
642 +
643 +// Receive reads one logical message. buf is a scratch buffer.
644 +func (s *Session) Receive(buf []byte) (protocol.Header, []byte, error) {
645 + if s.handle == syscall.InvalidHandle {
646 + return protocol.Header{}, nil, wrapErr(ErrBadParam, "session closed")
647 + }
648 +
649 + n, err := rawRecv(s.handle, buf)
650 + if err != nil {
651 + if errors.Is(err, ErrDisconnected) {
652 + s.failAllInflight()
653 + }
654 + return protocol.Header{}, nil, err
655 + }
656 +
657 + if n < protocol.HeaderSize {
658 + return protocol.Header{}, nil, wrapErr(ErrProtocol, "packet too short for header")
659 + }
660 +
661 + hdr, err := protocol.DecodeHeader(buf[:n])
662 + if err != nil {
663 + return protocol.Header{}, nil, wrapErr(ErrProtocol, "header decode: "+err.Error())
664 + }
665 +
666 + // Validate payload_len
667 + var maxPayload uint32
668 + if s.role == RoleServer {
669 + maxPayload = s.MaxRequestPayloadBytes
670 + } else {
671 + maxPayload = s.MaxResponsePayloadBytes
672 + }
673 + if hdr.PayloadLen > maxPayload {
674 + return protocol.Header{}, nil, wrapErr(ErrLimitExceeded,
675 + fmt.Sprintf("payload_len %d exceeds negotiated max %d", hdr.PayloadLen, maxPayload))
676 + }
677 +
678 + // Validate item_count
679 + var maxBatch uint32
680 + if s.role == RoleServer {
681 + maxBatch = s.MaxRequestBatchItems
682 + } else {
683 + maxBatch = s.MaxResponseBatchItems
684 + }
685 + if hdr.ItemCount > maxBatch {
686 + return protocol.Header{}, nil, wrapErr(ErrLimitExceeded,
687 + fmt.Sprintf("item_count %d exceeds negotiated max %d", hdr.ItemCount, maxBatch))
688 + }
689 +
690 + // Client-side: validate response message_id
691 + if s.role == RoleClient && hdr.Kind == protocol.KindResponse {
692 + if s.inflightIDs == nil {
693 + return protocol.Header{}, nil, wrapErr(ErrUnknownMsgID,
694 + fmt.Sprintf("message_id %d", hdr.MessageID))
695 + }
696 + if _, exists := s.inflightIDs[hdr.MessageID]; !exists {
697 + return protocol.Header{}, nil, wrapErr(ErrUnknownMsgID,
698 + fmt.Sprintf("message_id %d", hdr.MessageID))
699 + }
700 + delete(s.inflightIDs, hdr.MessageID)
701 + }
702 +
703 + totalMsg := protocol.HeaderSize + int(hdr.PayloadLen)
704 +
705 + // Non-chunked
706 + if n >= totalMsg {
707 + payload := buf[protocol.HeaderSize : protocol.HeaderSize+int(hdr.PayloadLen)]
708 +
709 + // Validate batch directory
710 + if hdr.Flags&protocol.FlagBatch != 0 && hdr.ItemCount > 1 {
711 + dirBytes := int(hdr.ItemCount) * 8
712 + dirAligned := protocol.Align8(dirBytes)
713 + if len(payload) < dirAligned {
714 + return protocol.Header{}, nil, wrapErr(ErrProtocol, "batch dir exceeds payload")
715 + }
716 + packedAreaLen := uint32(len(payload) - dirAligned)
717 + if err := protocol.BatchDirValidate(payload[:dirBytes], hdr.ItemCount, packedAreaLen); err != nil {
718 + return protocol.Header{}, nil, wrapErr(ErrProtocol, "batch dir: "+err.Error())
719 + }
720 + }
721 +
722 + return hdr, payload, nil
723 + }
724 +
725 + // Chunked
726 + firstPayloadBytes := n - protocol.HeaderSize
727 + needed := int(hdr.PayloadLen)
728 + if len(s.recvBuf) < needed {
729 + s.recvBuf = make([]byte, needed)
730 + }
731 +
732 + copy(s.recvBuf[:firstPayloadBytes], buf[protocol.HeaderSize:protocol.HeaderSize+firstPayloadBytes])
733 +
734 + assembled := firstPayloadBytes
735 + chunkPayloadBudget := int(s.PacketSize) - protocol.HeaderSize
736 +
737 + remainingAfterFirst := int(hdr.PayloadLen) - firstPayloadBytes
738 + expectedContinuations := uint32(0)
739 + if remainingAfterFirst > 0 && chunkPayloadBudget > 0 {
740 + expectedContinuations = uint32((remainingAfterFirst + chunkPayloadBudget - 1) / chunkPayloadBudget)
741 + }
742 + expectedChunkCount := 1 + expectedContinuations
743 +
744 + pktBuf := ensurePipeScratchBuf(&s.pktBuf, int(s.PacketSize))
745 +
746 + ci := uint32(1)
747 + for assembled < int(hdr.PayloadLen) {
748 + cn, err := rawRecv(s.handle, pktBuf)
749 + if err != nil {
750 + if errors.Is(err, ErrDisconnected) {
751 + s.failAllInflight()
752 + }
753 + return protocol.Header{}, nil, wrapErr(ErrRecv, "continuation recv: "+err.Error())
754 + }
755 +
756 + if cn < protocol.HeaderSize {
757 + return protocol.Header{}, nil, wrapErr(ErrChunk, "continuation too short")
758 + }
759 +
760 + chk, err := protocol.DecodeChunkHeader(pktBuf[:cn])
761 + if err != nil {
762 + return protocol.Header{}, nil, wrapErr(ErrChunk, "chunk header: "+err.Error())
763 + }
764 +
765 + if chk.MessageID != hdr.MessageID {
766 + return protocol.Header{}, nil, wrapErr(ErrChunk, "message_id mismatch")
767 + }
768 + if chk.ChunkIndex != ci {
769 + return protocol.Header{}, nil, wrapErr(ErrChunk, fmt.Sprintf(
770 + "chunk_index mismatch: expected %d, got %d", ci, chk.ChunkIndex))
771 + }
772 + if chk.ChunkCount != expectedChunkCount {
773 + return protocol.Header{}, nil, wrapErr(ErrChunk, "chunk_count mismatch")
774 + }
775 + if chk.TotalMessageLen != uint32(totalMsg) {
776 + return protocol.Header{}, nil, wrapErr(ErrChunk, "total_message_len mismatch")
777 + }
778 +
779 + chunkData := cn - protocol.HeaderSize
780 + if chunkData != int(chk.ChunkPayloadLen) {
781 + return protocol.Header{}, nil, wrapErr(ErrChunk, "chunk_payload_len mismatch")
782 + }
783 + if assembled+chunkData > int(hdr.PayloadLen) {
784 + return protocol.Header{}, nil, wrapErr(ErrChunk, "chunk exceeds payload_len")
785 + }
786 +
787 + copy(s.recvBuf[assembled:assembled+chunkData], pktBuf[protocol.HeaderSize:protocol.HeaderSize+chunkData])
788 + assembled += chunkData
789 + ci++
790 + }
791 +
792 + payload := s.recvBuf[:hdr.PayloadLen]
793 +
794 + // Validate batch directory
795 + if hdr.Flags&protocol.FlagBatch != 0 && hdr.ItemCount > 1 {
796 + dirBytes := int(hdr.ItemCount) * 8
797 + dirAligned := protocol.Align8(dirBytes)
798 + if len(payload) < dirAligned {
799 + return protocol.Header{}, nil, wrapErr(ErrProtocol, "batch dir exceeds payload")
800 + }
801 + packedAreaLen := uint32(len(payload) - dirAligned)
802 + if err := protocol.BatchDirValidate(payload[:dirBytes], hdr.ItemCount, packedAreaLen); err != nil {
803 + return protocol.Header{}, nil, wrapErr(ErrProtocol, "batch dir: "+err.Error())
804 + }
805 + }
806 +
807 + return hdr, payload, nil
808 +}
809 +
810 +func ensurePipeScratchBuf(buf *[]byte, needed int) []byte {
811 + if len(*buf) < needed {
812 + *buf = make([]byte, needed)
813 + }
814 + return (*buf)[:needed]
815 +}
816 +
817 +// ---------------------------------------------------------------------------
818 +// Listener
819 +// ---------------------------------------------------------------------------
820 +
821 +// Listener is a listening Named Pipe endpoint.
822 +type Listener struct {
823 + mu sync.Mutex
824 + handle syscall.Handle
825 + config ServerConfig
826 + pipeName []uint16
827 + nextSessionID atomic.Uint64
828 + closing bool
829 + accepting bool
830 +}
831 +
832 +// Listen creates a listener on a Named Pipe derived from runDir + serviceName.
833 +func Listen(runDir, serviceName string, config ServerConfig) (*Listener, error) {
834 + pipeName, err := BuildPipeName(runDir, serviceName)
835 + if err != nil {
836 + return nil, err
837 + }
838 +
839 + bufSize := pipeBufferSize(config.PacketSize)
840 +
841 + // Create first instance with FILE_FLAG_FIRST_PIPE_INSTANCE
842 + handle, err := createPipeInstance(pipeName, bufSize, true)
843 + if err != nil {
844 + return nil, err
845 + }
846 +
847 + return &Listener{
848 + handle: handle,
849 + config: config,
850 + pipeName: pipeName,
851 + }, nil
852 +}
853 +
854 +// Handle returns the raw HANDLE.
855 +func (l *Listener) Handle() syscall.Handle {
856 + l.mu.Lock()
857 + defer l.mu.Unlock()
858 + return l.handle
859 +}
860 +
861 +// SetPayloadLimits updates the payload limits used for future handshakes.
862 +func (l *Listener) SetPayloadLimits(maxRequestPayloadBytes, maxResponsePayloadBytes uint32) {
863 + l.mu.Lock()
864 + defer l.mu.Unlock()
865 + l.config.MaxRequestPayloadBytes = maxRequestPayloadBytes
866 + l.config.MaxResponsePayloadBytes = maxResponsePayloadBytes
867 +}
868 +
869 +// Accept accepts one client connection. Performs the full handshake.
870 +func (l *Listener) Accept() (*Session, error) {
871 + sessionID := l.nextSessionID.Add(1)
872 + return l.AcceptWithConfig(sessionID, l.config)
873 +}
874 +
875 +// AcceptWithConfig accepts one client connection using a caller-provided
876 +// per-session server config and session ID.
877 +func (l *Listener) AcceptWithConfig(sessionID uint64, config ServerConfig) (*Session, error) {
878 + l.mu.Lock()
879 + if l.handle == syscall.InvalidHandle {
880 + l.mu.Unlock()
881 + return nil, wrapErr(ErrAccept, "listener closed")
882 + }
883 + sessionHandle := l.handle
884 + l.accepting = true
885 + l.mu.Unlock()
886 +
887 + err := connectNamedPipe(sessionHandle)
888 + if err != nil {
889 + // ERROR_PIPE_CONNECTED is fine — client connected between
890 + // CreateNamedPipe and ConnectNamedPipe
891 + if errno, ok := err.(syscall.Errno); !ok || errno != _ERROR_PIPE_CONNECTED {
892 + l.mu.Lock()
893 + l.accepting = false
894 + l.mu.Unlock()
895 + return nil, wrapErr(ErrAccept, err.Error())
896 + }
897 + }
898 +
899 + l.mu.Lock()
900 + l.accepting = false
901 + if l.closing {
902 + if l.handle == sessionHandle {
903 + l.handle = syscall.InvalidHandle
904 + }
905 + l.mu.Unlock()
906 + disconnectNamedPipe(sessionHandle)
907 + syscall.CloseHandle(sessionHandle)
908 + return nil, wrapErr(ErrAccept, "listener closed")
909 + }
910 +
911 + // Create new pipe instance for next client
912 + bufSize := pipeBufferSize(l.config.PacketSize)
913 + next, perr := createPipeInstance(l.pipeName, bufSize, false)
914 + if perr != nil {
915 + if l.handle == sessionHandle {
916 + l.handle = syscall.InvalidHandle
917 + }
918 + l.mu.Unlock()
919 + disconnectNamedPipe(sessionHandle)
920 + syscall.CloseHandle(sessionHandle)
921 + return nil, perr
922 + }
923 + l.handle = next
924 + l.mu.Unlock()
925 +
926 + // Handshake
927 + session, herr := serverHandshake(sessionHandle, &config, sessionID)
928 + if herr != nil {
929 + disconnectNamedPipe(sessionHandle)
930 + syscall.CloseHandle(sessionHandle)
931 + return nil, herr
932 + }
933 + return session, nil
934 +}
935 +
936 +// Close closes the listener.
937 +func (l *Listener) Close() {
938 + l.mu.Lock()
939 + handle := l.handle
940 + if handle == syscall.InvalidHandle {
941 + l.mu.Unlock()
942 + return
943 + }
944 + l.closing = true
945 + accepting := l.accepting
946 + if !accepting {
947 + l.handle = syscall.InvalidHandle
948 + }
949 + pipeName := l.pipeName
950 + l.mu.Unlock()
951 +
952 + if accepting && len(pipeName) > 0 && pipeName[0] != 0 {
953 + // A loopback connect reliably wakes a blocking ConnectNamedPipe()
954 + // so Accept() can observe shutdown and close the live listener handle
955 + // from the owning goroutine.
956 + wake, err := syscall.CreateFile(
957 + &pipeName[0],
958 + _GENERIC_READ|_GENERIC_WRITE,
959 + 0,
960 + nil,
961 + _OPEN_EXISTING,
962 + 0,
963 + 0,
964 + )
965 + if err == nil && wake != syscall.InvalidHandle && wake != 0 {
966 + syscall.CloseHandle(wake)
967 + }
968 + return
969 + }
970 +
971 + syscall.CloseHandle(handle)
972 +}
973 +
974 +// ---------------------------------------------------------------------------
975 +// Pipe instance creation
976 +// ---------------------------------------------------------------------------
977 +
978 +func createPipeInstance(pipeName []uint16, bufSize uint32, firstInstance bool) (syscall.Handle, error) {
979 + openMode := uint32(_PIPE_ACCESS_DUPLEX)
980 + if firstInstance {
981 + openMode |= _FILE_FLAG_FIRST_PIPE_INSTANCE
982 + }
983 +
984 + handle, err := createNamedPipe(
985 + &pipeName[0],
986 + openMode,
987 + _PIPE_TYPE_MESSAGE|_PIPE_READMODE_MESSAGE|_PIPE_WAIT,
988 + _PIPE_UNLIMITED_INSTANCES,
989 + bufSize,
990 + bufSize,
991 + 0,
992 + )
993 + if err != nil {
994 + errno, ok := err.(syscall.Errno)
995 + if ok && (errno == _ERROR_ACCESS_DENIED || errno == _ERROR_PIPE_BUSY) {
996 + return syscall.InvalidHandle, ErrAddrInUse
997 + }
998 + return syscall.InvalidHandle, wrapErr(ErrCreatePipe, err.Error())
999 + }
1000 + return handle, nil
1001 +}
1002 +
1003 +// ---------------------------------------------------------------------------
1004 +// Client handshake
1005 +// ---------------------------------------------------------------------------
1006 +
1007 +func clientHandshake(handle syscall.Handle, config *ClientConfig) (*Session, error) {
1008 + pktSize := applyDefault(config.PacketSize, defaultPacketSize)
1009 +
1010 + supported := config.SupportedProfiles
1011 + if supported == 0 {
1012 + supported = protocol.ProfileBaseline
1013 + }
1014 +
1015 + hello := protocol.Hello{
1016 + LayoutVersion: 1,
1017 + Flags: 0,
1018 + SupportedProfiles: supported,
1019 + PreferredProfiles: config.PreferredProfiles,
1020 + MaxRequestPayloadBytes: applyDefault(config.MaxRequestPayloadBytes, protocol.MaxPayloadDefault),
1021 + MaxRequestBatchItems: applyDefault(config.MaxRequestBatchItems, defaultBatchItems),
1022 + MaxResponsePayloadBytes: applyDefault(config.MaxResponsePayloadBytes, protocol.MaxPayloadDefault),
1023 + MaxResponseBatchItems: applyDefault(config.MaxResponseBatchItems, defaultBatchItems),
1024 + AuthToken: config.AuthToken,
1025 + PacketSize: pktSize,
1026 + }
1027 +
1028 + var helloBuf [helloPayloadSize]byte
1029 + hello.Encode(helloBuf[:])
1030 +
1031 + hdr := protocol.Header{
1032 + Magic: protocol.MagicMsg,
1033 + Version: protocol.Version,
1034 + HeaderLen: protocol.HeaderLen,
1035 + Kind: protocol.KindControl,
1036 + Flags: 0,
1037 + Code: protocol.CodeHello,
1038 + TransportStatus: protocol.StatusOK,
1039 + PayloadLen: helloPayloadSize,
1040 + ItemCount: 1,
1041 + MessageID: 0,
1042 + }
1043 +
1044 + var pkt [protocol.HeaderSize + helloPayloadSize]byte
1045 + hdr.Encode(pkt[:protocol.HeaderSize])
1046 + copy(pkt[protocol.HeaderSize:], helloBuf[:])
1047 +
1048 + // Send HELLO
1049 + if err := rawWrite(handle, pkt[:]); err != nil {
1050 + return nil, wrapErr(ErrSend, "hello send: "+err.Error())
1051 + }
1052 +
1053 + // Receive HELLO_ACK
1054 + var ackBuf [128]byte
1055 + an, err := rawRecv(handle, ackBuf[:])
1056 + if err != nil {
1057 + return nil, wrapErr(ErrRecv, "hello_ack recv: "+err.Error())
1058 + }
1059 +
1060 + ackHdr, err := protocol.DecodeHeader(ackBuf[:an])
1061 + if err != nil {
1062 + if errors.Is(err, protocol.ErrBadVersion) {
1063 + return nil, wrapErr(ErrIncompatible, "ack header version mismatch")
1064 + }
1065 + return nil, wrapErr(ErrProtocol, "ack header: "+err.Error())
1066 + }
1067 +
1068 + if ackHdr.Kind != protocol.KindControl || ackHdr.Code != protocol.CodeHelloAck {
1069 + return nil, wrapErr(ErrProtocol, "expected HELLO_ACK")
1070 + }
1071 +
1072 + if ackHdr.TransportStatus == protocol.StatusAuthFailed {
1073 + return nil, ErrAuthFailed
1074 + }
1075 + if ackHdr.TransportStatus == protocol.StatusUnsupported {
1076 + return nil, ErrNoProfile
1077 + }
1078 + if ackHdr.TransportStatus == protocol.StatusIncompatible {
1079 + return nil, ErrIncompatible
1080 + }
1081 + if ackHdr.TransportStatus == protocol.StatusLimitExceeded {
1082 + return nil, ErrLimitExceeded
1083 + }
1084 + if ackHdr.TransportStatus != protocol.StatusOK {
1085 + return nil, wrapErr(ErrHandshake, fmt.Sprintf("transport_status=%d", ackHdr.TransportStatus))
1086 + }
1087 +
1088 + if an < protocol.HeaderSize+helloAckPayloadSize {
1089 + return nil, wrapErr(ErrProtocol, "ack payload truncated")
1090 + }
1091 + ack, err := protocol.DecodeHelloAck(ackBuf[protocol.HeaderSize:an])
1092 + if err != nil {
1093 + if errors.Is(err, protocol.ErrBadLayout) &&
1094 + helloAckLayoutIncompatible(ackBuf[protocol.HeaderSize:an]) {
1095 + return nil, wrapErr(ErrIncompatible, "ack payload layout version mismatch")
1096 + }
1097 + return nil, wrapErr(ErrProtocol, "ack payload: "+err.Error())
1098 + }
1099 +
1100 + return &Session{
1101 + handle: handle,
1102 + role: RoleClient,
1103 + MaxRequestPayloadBytes: ack.AgreedMaxRequestPayloadBytes,
1104 + MaxRequestBatchItems: ack.AgreedMaxRequestBatchItems,
1105 + MaxResponsePayloadBytes: ack.AgreedMaxResponsePayloadBytes,
1106 + MaxResponseBatchItems: ack.AgreedMaxResponseBatchItems,
1107 + PacketSize: ack.AgreedPacketSize,
1108 + SelectedProfile: ack.SelectedProfile,
1109 + SessionID: ack.SessionID,
1110 + inflightIDs: make(map[uint64]struct{}),
1111 + }, nil
1112 +}
1113 +
1114 +// ---------------------------------------------------------------------------
1115 +// Server handshake
1116 +// ---------------------------------------------------------------------------
1117 +
1118 +func serverHandshake(handle syscall.Handle, config *ServerConfig, sessionID uint64) (*Session, error) {
1119 + serverPktSize := applyDefault(config.PacketSize, defaultPacketSize)
1120 + sRespPay := applyDefault(config.MaxResponsePayloadBytes, protocol.MaxPayloadDefault)
1121 + sProfiles := config.SupportedProfiles
1122 + if sProfiles == 0 {
1123 + sProfiles = protocol.ProfileBaseline
1124 + }
1125 + sPreferred := config.PreferredProfiles
1126 +
1127 + // Helper: send rejection
1128 + sendRejection := func(status uint16) {
1129 + ack := protocol.HelloAck{LayoutVersion: 1}
1130 + var ackPayBuf [helloAckPayloadSize]byte
1131 + ack.Encode(ackPayBuf[:])
1132 +
1133 + ackHdr := protocol.Header{
1134 + Magic: protocol.MagicMsg,
1135 + Version: protocol.Version,
1136 + HeaderLen: protocol.HeaderLen,
1137 + Kind: protocol.KindControl,
1138 + Code: protocol.CodeHelloAck,
1139 + TransportStatus: status,
1140 + PayloadLen: helloAckPayloadSize,
1141 + ItemCount: 1,
1142 + }
1143 +
1144 + var pkt [protocol.HeaderSize + helloAckPayloadSize]byte
1145 + ackHdr.Encode(pkt[:protocol.HeaderSize])
1146 + copy(pkt[protocol.HeaderSize:], ackPayBuf[:])
1147 + rawWrite(handle, pkt[:]) //nolint:errcheck
1148 + }
1149 +
1150 + // Receive HELLO
1151 + var buf [128]byte
1152 + n, err := rawRecv(handle, buf[:])
1153 + if err != nil {
1154 + return nil, wrapErr(ErrRecv, "hello recv: "+err.Error())
1155 + }
1156 +
1157 + hdr, err := protocol.DecodeHeader(buf[:n])
1158 + if err != nil {
1159 + if errors.Is(err, protocol.ErrBadVersion) &&
1160 + headerVersionIncompatible(buf[:n], protocol.CodeHello) {
1161 + sendRejection(protocol.StatusIncompatible)
1162 + return nil, ErrIncompatible
1163 + }
1164 + return nil, wrapErr(ErrProtocol, "hello header: "+err.Error())
1165 + }
1166 +
1167 + if hdr.Kind != protocol.KindControl || hdr.Code != protocol.CodeHello {
1168 + return nil, wrapErr(ErrProtocol, "expected HELLO")
1169 + }
1170 +
1171 + hello, err := protocol.DecodeHello(buf[protocol.HeaderSize:n])
1172 + if err != nil {
1173 + if errors.Is(err, protocol.ErrBadLayout) &&
1174 + helloLayoutIncompatible(buf[protocol.HeaderSize:n]) {
1175 + sendRejection(protocol.StatusIncompatible)
1176 + return nil, ErrIncompatible
1177 + }
1178 + return nil, wrapErr(ErrProtocol, "hello payload: "+err.Error())
1179 + }
1180 +
1181 + intersection := hello.SupportedProfiles & sProfiles
1182 +
1183 + if intersection == 0 {
1184 + sendRejection(protocol.StatusUnsupported)
1185 + return nil, ErrNoProfile
1186 + }
1187 +
1188 + if hello.AuthToken != config.AuthToken {
1189 + sendRejection(protocol.StatusAuthFailed)
1190 + return nil, ErrAuthFailed
1191 + }
1192 +
1193 + // Select profile
1194 + preferredIntersection := intersection & hello.PreferredProfiles & sPreferred
1195 + var selected uint32
1196 + if preferredIntersection != 0 {
1197 + selected = highestBit(preferredIntersection)
1198 + } else {
1199 + selected = highestBit(intersection)
1200 + }
1201 +
1202 + if hello.MaxRequestPayloadBytes > protocol.MaxPayloadCap {
1203 + sendRejection(protocol.StatusLimitExceeded)
1204 + return nil, ErrLimitExceeded
1205 + }
1206 +
1207 + // Negotiate limits
1208 + agreedReqPay := hello.MaxRequestPayloadBytes
1209 + agreedReqBat := hello.MaxRequestBatchItems
1210 + agreedRespPay := sRespPay
1211 + agreedRespBat := agreedReqBat
1212 + agreedPkt := minU32(hello.PacketSize, serverPktSize)
1213 + if agreedPkt <= protocol.HeaderSize {
1214 + sendRejection(protocol.StatusIncompatible)
1215 + return nil, ErrIncompatible
1216 + }
1217 +
1218 + // Send HELLO_ACK
1219 + ack := protocol.HelloAck{
1220 + LayoutVersion: 1,
1221 + Flags: 0,
1222 + ServerSupportedProfiles: sProfiles,
1223 + IntersectionProfiles: intersection,
1224 + SelectedProfile: selected,
1225 + AgreedMaxRequestPayloadBytes: agreedReqPay,
1226 + AgreedMaxRequestBatchItems: agreedReqBat,
1227 + AgreedMaxResponsePayloadBytes: agreedRespPay,
1228 + AgreedMaxResponseBatchItems: agreedRespBat,
1229 + AgreedPacketSize: agreedPkt,
1230 + SessionID: sessionID,
1231 + }
1232 +
1233 + var ackPayBuf [helloAckPayloadSize]byte
1234 + ack.Encode(ackPayBuf[:])
1235 +
1236 + ackHdr := protocol.Header{
1237 + Magic: protocol.MagicMsg,
1238 + Version: protocol.Version,
1239 + HeaderLen: protocol.HeaderLen,
1240 + Kind: protocol.KindControl,
1241 + Code: protocol.CodeHelloAck,
1242 + TransportStatus: protocol.StatusOK,
1243 + PayloadLen: helloAckPayloadSize,
1244 + ItemCount: 1,
1245 + }
1246 +
1247 + var pkt [protocol.HeaderSize + helloAckPayloadSize]byte
1248 + ackHdr.Encode(pkt[:protocol.HeaderSize])
1249 + copy(pkt[protocol.HeaderSize:], ackPayBuf[:])
1250 +
1251 + if err := rawWrite(handle, pkt[:]); err != nil {
1252 + return nil, wrapErr(ErrSend, "hello_ack send: "+err.Error())
1253 + }
1254 +
1255 + return &Session{
1256 + handle: handle,
1257 + role: RoleServer,
1258 + MaxRequestPayloadBytes: agreedReqPay,
1259 + MaxRequestBatchItems: agreedReqBat,
1260 + MaxResponsePayloadBytes: agreedRespPay,
1261 + MaxResponseBatchItems: agreedRespBat,
1262 + PacketSize: agreedPkt,
1263 + SelectedProfile: selected,
1264 + SessionID: sessionID,
1265 + inflightIDs: make(map[uint64]struct{}),
1266 + }, nil
1267 +}
src/go/pkg/netipc/transport/windows/pipe_edge_test.go new
+1135
@@ -0,0 +1,1135 @@
1 +//go:build windows
2 +
3 +package windows
4 +
5 +import (
6 + "bytes"
7 + "encoding/binary"
8 + "errors"
9 + "strings"
10 + "syscall"
11 + "testing"
12 + "time"
13 +
14 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
15 +)
16 +
17 +func rawPipePair(t *testing.T) (syscall.Handle, syscall.Handle) {
18 + t.Helper()
19 +
20 + service := uniquePipeService(t)
21 + pipeName, err := BuildPipeName(testPipeRunDir, service)
22 + if err != nil {
23 + t.Fatalf("BuildPipeName failed: %v", err)
24 + }
25 +
26 + serverHandle, err := createPipeInstance(pipeName, defaultPipeBufSize, true)
27 + if err != nil {
28 + t.Fatalf("createPipeInstance failed: %v", err)
29 + }
30 + t.Cleanup(func() {
31 + if serverHandle != syscall.InvalidHandle && serverHandle != 0 {
32 + disconnectNamedPipe(serverHandle)
33 + syscall.CloseHandle(serverHandle)
34 + }
35 + })
36 +
37 + acceptCh := make(chan error, 1)
38 + go func() {
39 + acceptCh <- connectNamedPipe(serverHandle)
40 + }()
41 +
42 + clientHandle, err := syscall.CreateFile(
43 + &pipeName[0],
44 + _GENERIC_READ|_GENERIC_WRITE,
45 + 0,
46 + nil,
47 + _OPEN_EXISTING,
48 + 0,
49 + 0,
50 + )
51 + if err != nil {
52 + t.Fatalf("CreateFile failed: %v", err)
53 + }
54 + t.Cleanup(func() {
55 + if clientHandle != syscall.InvalidHandle && clientHandle != 0 {
56 + syscall.CloseHandle(clientHandle)
57 + }
58 + })
59 +
60 + mode := uint32(_PIPE_READMODE_MESSAGE)
61 + if err := setNamedPipeHandleState(clientHandle, &mode); err != nil {
62 + t.Fatalf("setNamedPipeHandleState failed: %v", err)
63 + }
64 +
65 + if err := <-acceptCh; err != nil && !errors.Is(err, syscall.Errno(_ERROR_PIPE_CONNECTED)) {
66 + t.Fatalf("connectNamedPipe failed: %v", err)
67 + }
68 +
69 + return clientHandle, serverHandle
70 +}
71 +
72 +func sessionPair(t *testing.T, sCfg ServerConfig, cCfg ClientConfig) (*Session, *Session) {
73 + t.Helper()
74 +
75 + service := uniquePipeService(t)
76 + listener := startListener(t, testPipeRunDir, service, sCfg)
77 + t.Cleanup(listener.Close)
78 +
79 + acceptCh := acceptAsync(listener)
80 +
81 + client, err := Connect(testPipeRunDir, service, &cCfg)
82 + if err != nil {
83 + t.Fatalf("Connect failed: %v", err)
84 + }
85 + t.Cleanup(client.Close)
86 +
87 + sr := <-acceptCh
88 + if sr.err != nil {
89 + t.Fatalf("Accept failed: %v", sr.err)
90 + }
91 + server := sr.session
92 + t.Cleanup(server.Close)
93 +
94 + return client, server
95 +}
96 +
97 +func encodePacket(hdr protocol.Header, payload []byte) []byte {
98 + hdr.Magic = protocol.MagicMsg
99 + hdr.Version = protocol.Version
100 + hdr.HeaderLen = protocol.HeaderLen
101 + hdr.PayloadLen = uint32(len(payload))
102 +
103 + pkt := make([]byte, protocol.HeaderSize+len(payload))
104 + hdr.Encode(pkt[:protocol.HeaderSize])
105 + copy(pkt[protocol.HeaderSize:], payload)
106 + return pkt
107 +}
108 +
109 +func encodeChunkPacket(chk protocol.ChunkHeader, payload []byte) []byte {
110 + pkt := make([]byte, protocol.HeaderSize+len(payload))
111 + chk.Encode(pkt[:protocol.HeaderSize])
112 + copy(pkt[protocol.HeaderSize:], payload)
113 + return pkt
114 +}
115 +
116 +func validHelloAckPacket(status uint16) []byte {
117 + ack := protocol.HelloAck{
118 + LayoutVersion: 1,
119 + Flags: 0,
120 + ServerSupportedProfiles: protocol.ProfileBaseline,
121 + IntersectionProfiles: protocol.ProfileBaseline,
122 + SelectedProfile: protocol.ProfileBaseline,
123 + AgreedMaxRequestPayloadBytes: protocol.MaxPayloadDefault,
124 + AgreedMaxRequestBatchItems: 1,
125 + AgreedMaxResponsePayloadBytes: protocol.MaxPayloadDefault,
126 + AgreedMaxResponseBatchItems: 1,
127 + AgreedPacketSize: defaultPacketSize,
128 + SessionID: 77,
129 + }
130 +
131 + payload := make([]byte, helloAckPayloadSize)
132 + ack.Encode(payload)
133 + return encodePacket(protocol.Header{
134 + Kind: protocol.KindControl,
135 + Code: protocol.CodeHelloAck,
136 + TransportStatus: status,
137 + ItemCount: 1,
138 + }, payload)
139 +}
140 +
141 +func validHelloPacket() []byte {
142 + hello := protocol.Hello{
143 + LayoutVersion: 1,
144 + Flags: 0,
145 + SupportedProfiles: protocol.ProfileBaseline,
146 + PreferredProfiles: protocol.ProfileBaseline,
147 + MaxRequestPayloadBytes: protocol.MaxPayloadDefault,
148 + MaxRequestBatchItems: 1,
149 + MaxResponsePayloadBytes: protocol.MaxPayloadDefault,
150 + MaxResponseBatchItems: 1,
151 + AuthToken: testAuthToken,
152 + PacketSize: defaultPacketSize,
153 + }
154 +
155 + payload := make([]byte, helloPayloadSize)
156 + hello.Encode(payload)
157 + return encodePacket(protocol.Header{
158 + Kind: protocol.KindControl,
159 + Code: protocol.CodeHello,
160 + ItemCount: 1,
161 + }, payload)
162 +}
163 +
164 +func TestPipeConnectRejectsInvalidServiceName(t *testing.T) {
165 + if _, err := Connect(testPipeRunDir, "bad/name", defaultClientConfigPtr()); !errors.Is(err, ErrBadParam) {
166 + t.Fatalf("Connect(invalid service) = %v, want ErrBadParam", err)
167 + }
168 +}
169 +
170 +func TestPipeConnectRejectsPipeNameTooLong(t *testing.T) {
171 + service := strings.Repeat("a", maxPipeNameChars)
172 + if _, err := Connect(testPipeRunDir, service, defaultClientConfigPtr()); !errors.Is(err, ErrPipeName) {
173 + t.Fatalf("Connect(long service) = %v, want ErrPipeName", err)
174 + }
175 +}
176 +
177 +func TestPipeAcceptUnblocksWhenListenerClosed(t *testing.T) {
178 + service := uniquePipeService(t)
179 + listener := startListener(t, testPipeRunDir, service, defaultServerConfig())
180 +
181 + acceptCh := acceptAsync(listener)
182 + time.Sleep(50 * time.Millisecond)
183 + listener.Close()
184 +
185 + select {
186 + case sr := <-acceptCh:
187 + if sr.err == nil {
188 + if sr.session != nil {
189 + sr.session.Close()
190 + }
191 + t.Fatal("Accept after listener close should fail")
192 + }
193 + case <-time.After(2 * time.Second):
194 + t.Fatal("Accept did not unblock after listener close")
195 + }
196 +}
197 +
198 +func TestRawRecvMoreData(t *testing.T) {
199 + clientHandle, serverHandle := rawPipePair(t)
200 +
201 + msg := make([]byte, 96)
202 + for i := range msg {
203 + msg[i] = byte(i)
204 + }
205 + if err := rawWrite(clientHandle, msg); err != nil {
206 + t.Fatalf("rawWrite failed: %v", err)
207 + }
208 +
209 + small := make([]byte, 16)
210 + n, err := rawRecv(serverHandle, small)
211 + if err != nil {
212 + t.Fatalf("rawRecv(small) failed: %v", err)
213 + }
214 + if n != len(small) {
215 + t.Fatalf("rawRecv(small) len=%d, want %d", n, len(small))
216 + }
217 +
218 + rest := make([]byte, 96)
219 + n, err = rawRecv(serverHandle, rest)
220 + if err != nil {
221 + t.Fatalf("rawRecv(rest) failed: %v", err)
222 + }
223 + if n != len(msg)-len(small) {
224 + t.Fatalf("rawRecv(rest) len=%d, want %d", n, len(msg)-len(small))
225 + }
226 +}
227 +
228 +func TestRawPipeDisconnectErrors(t *testing.T) {
229 + clientHandle, serverHandle := rawPipePair(t)
230 +
231 + syscall.CloseHandle(serverHandle)
232 + serverHandle = syscall.InvalidHandle
233 +
234 + if err := rawWrite(clientHandle, []byte("x")); !errors.Is(err, ErrDisconnected) {
235 + t.Fatalf("rawWrite after peer close = %v, want ErrDisconnected", err)
236 + }
237 +
238 + clientHandle2, serverHandle2 := rawPipePair(t)
239 + syscall.CloseHandle(clientHandle2)
240 + clientHandle2 = syscall.InvalidHandle
241 +
242 + buf := make([]byte, 32)
243 + if _, err := rawRecv(serverHandle2, buf); !errors.Is(err, ErrDisconnected) {
244 + t.Fatalf("rawRecv after peer close = %v, want ErrDisconnected", err)
245 + }
246 +}
247 +
248 +func TestRawPipeGenericErrors(t *testing.T) {
249 + if err := rawWrite(syscall.InvalidHandle, []byte("x")); !errors.Is(err, ErrSend) {
250 + t.Fatalf("rawWrite(invalid handle) = %v, want ErrSend", err)
251 + }
252 +
253 + buf := make([]byte, 16)
254 + if _, err := rawRecv(syscall.InvalidHandle, buf); !errors.Is(err, ErrRecv) {
255 + t.Fatalf("rawRecv(invalid handle) = %v, want ErrRecv", err)
256 + }
257 +}
258 +
259 +func TestClientHandshakeRejectsMalformedAck(t *testing.T) {
260 + cases := []struct {
261 + name string
262 + packetFn func() []byte
263 + want error
264 + wantText string
265 + }{
266 + {
267 + name: "bad ack header",
268 + packetFn: func() []byte {
269 + ack := validHelloAckPacket(protocol.StatusOK)
270 + ack[0] = 0
271 + return ack
272 + },
273 + want: ErrProtocol,
274 + wantText: "ack header",
275 + },
276 + {
277 + name: "wrong kind",
278 + packetFn: func() []byte {
279 + payload := make([]byte, helloAckPayloadSize)
280 + return encodePacket(protocol.Header{
281 + Kind: protocol.KindRequest,
282 + Code: protocol.CodeHelloAck,
283 + TransportStatus: protocol.StatusOK,
284 + ItemCount: 1,
285 + }, payload)
286 + },
287 + want: ErrProtocol,
288 + wantText: "expected HELLO_ACK",
289 + },
290 + {
291 + name: "unexpected status",
292 + packetFn: func() []byte {
293 + return validHelloAckPacket(999)
294 + },
295 + want: ErrHandshake,
296 + wantText: "transport_status=999",
297 + },
298 + {
299 + name: "truncated payload",
300 + packetFn: func() []byte {
301 + ack := validHelloAckPacket(protocol.StatusOK)
302 + return ack[:protocol.HeaderSize+8]
303 + },
304 + want: ErrProtocol,
305 + wantText: "ack payload truncated",
306 + },
307 + {
308 + name: "bad ack payload layout",
309 + packetFn: func() []byte {
310 + ack := validHelloAckPacket(protocol.StatusOK)
311 + binary.NativeEndian.PutUint16(ack[protocol.HeaderSize:protocol.HeaderSize+2], 2)
312 + return ack
313 + },
314 + want: ErrIncompatible,
315 + wantText: "ack payload",
316 + },
317 + }
318 +
319 + for _, tc := range cases {
320 + t.Run(tc.name, func(t *testing.T) {
321 + service := uniquePipeService(t)
322 + pipeName, err := BuildPipeName(testPipeRunDir, service)
323 + if err != nil {
324 + t.Fatalf("BuildPipeName failed: %v", err)
325 + }
326 +
327 + serverHandle, err := createPipeInstance(pipeName, defaultPipeBufSize, true)
328 + if err != nil {
329 + t.Fatalf("createPipeInstance failed: %v", err)
330 + }
331 + t.Cleanup(func() {
332 + if serverHandle != syscall.InvalidHandle && serverHandle != 0 {
333 + disconnectNamedPipe(serverHandle)
334 + syscall.CloseHandle(serverHandle)
335 + }
336 + })
337 +
338 + serverDone := make(chan error, 1)
339 + go func() {
340 + err := connectNamedPipe(serverHandle)
341 + if err != nil && !errors.Is(err, syscall.Errno(_ERROR_PIPE_CONNECTED)) {
342 + serverDone <- err
343 + return
344 + }
345 +
346 + var helloBuf [128]byte
347 + if _, err := rawRecv(serverHandle, helloBuf[:]); err != nil {
348 + serverDone <- err
349 + return
350 + }
351 +
352 + serverDone <- rawWrite(serverHandle, tc.packetFn())
353 + }()
354 +
355 + clientHandle, err := syscall.CreateFile(
356 + &pipeName[0],
357 + _GENERIC_READ|_GENERIC_WRITE,
358 + 0,
359 + nil,
360 + _OPEN_EXISTING,
361 + 0,
362 + 0,
363 + )
364 + if err != nil {
365 + t.Fatalf("CreateFile failed: %v", err)
366 + }
367 + t.Cleanup(func() {
368 + if clientHandle != syscall.InvalidHandle && clientHandle != 0 {
369 + syscall.CloseHandle(clientHandle)
370 + }
371 + })
372 +
373 + mode := uint32(_PIPE_READMODE_MESSAGE)
374 + if err := setNamedPipeHandleState(clientHandle, &mode); err != nil {
375 + t.Fatalf("setNamedPipeHandleState failed: %v", err)
376 + }
377 +
378 + session, err := clientHandshake(clientHandle, defaultClientConfigPtr())
379 + if err == nil {
380 + session.Close()
381 + t.Fatal("clientHandshake should fail")
382 + }
383 + if !errors.Is(err, tc.want) {
384 + t.Fatalf("clientHandshake error = %v, want %v", err, tc.want)
385 + }
386 + if !strings.Contains(err.Error(), tc.wantText) {
387 + t.Fatalf("clientHandshake error = %v, want text %q", err, tc.wantText)
388 + }
389 +
390 + if err := <-serverDone; err != nil {
391 + t.Fatalf("raw server failed: %v", err)
392 + }
393 + })
394 + }
395 +}
396 +
397 +func TestClientHandshakeTransportErrors(t *testing.T) {
398 + t.Run("hello send disconnected", func(t *testing.T) {
399 + clientHandle, serverHandle := rawPipePair(t)
400 +
401 + if err := syscall.CloseHandle(serverHandle); err != nil {
402 + t.Fatalf("CloseHandle(server) failed: %v", err)
403 + }
404 +
405 + _, err := clientHandshake(clientHandle, defaultClientConfigPtr())
406 + if err == nil {
407 + t.Fatal("clientHandshake should fail")
408 + }
409 + if !errors.Is(err, ErrSend) {
410 + t.Fatalf("clientHandshake error = %v, want ErrSend", err)
411 + }
412 + if !strings.Contains(err.Error(), "hello send") {
413 + t.Fatalf("clientHandshake error = %v, want text %q", err, "hello send")
414 + }
415 + })
416 +
417 + t.Run("hello ack recv disconnected", func(t *testing.T) {
418 + clientHandle, serverHandle := rawPipePair(t)
419 +
420 + serverDone := make(chan error, 1)
421 + go func() {
422 + var helloBuf [128]byte
423 + _, err := rawRecv(serverHandle, helloBuf[:])
424 + if err == nil {
425 + err = syscall.CloseHandle(serverHandle)
426 + }
427 + serverDone <- err
428 + }()
429 +
430 + _, err := clientHandshake(clientHandle, defaultClientConfigPtr())
431 + if err == nil {
432 + t.Fatal("clientHandshake should fail")
433 + }
434 + if !errors.Is(err, ErrRecv) {
435 + t.Fatalf("clientHandshake error = %v, want ErrRecv", err)
436 + }
437 + if !strings.Contains(err.Error(), "hello_ack recv") {
438 + t.Fatalf("clientHandshake error = %v, want text %q", err, "hello_ack recv")
439 + }
440 +
441 + if err := <-serverDone; err != nil {
442 + t.Fatalf("raw server failed: %v", err)
443 + }
444 + })
445 +}
446 +
447 +func TestServerHandshakeRejectsMalformedHello(t *testing.T) {
448 + cases := []struct {
449 + name string
450 + packetFn func() []byte
451 + want error
452 + wantText string
453 + }{
454 + {
455 + name: "bad hello header",
456 + packetFn: func() []byte {
457 + hello := validHelloPacket()
458 + hello[0] = 0
459 + return hello
460 + },
461 + want: ErrProtocol,
462 + wantText: "hello header",
463 + },
464 + {
465 + name: "wrong hello kind",
466 + packetFn: func() []byte {
467 + payload := make([]byte, helloPayloadSize)
468 + return encodePacket(protocol.Header{
469 + Kind: protocol.KindRequest,
470 + Code: protocol.CodeHello,
471 + ItemCount: 1,
472 + }, payload)
473 + },
474 + want: ErrProtocol,
475 + wantText: "expected HELLO",
476 + },
477 + {
478 + name: "bad hello payload layout",
479 + packetFn: func() []byte {
480 + hello := validHelloPacket()
481 + binary.NativeEndian.PutUint16(hello[protocol.HeaderSize:protocol.HeaderSize+2], 2)
482 + return hello
483 + },
484 + want: ErrIncompatible,
485 + wantText: "protocol or layout version mismatch",
486 + },
487 + }
488 +
489 + for _, tc := range cases {
490 + t.Run(tc.name, func(t *testing.T) {
491 + clientHandle, serverHandle := rawPipePair(t)
492 +
493 + if err := rawWrite(clientHandle, tc.packetFn()); err != nil {
494 + t.Fatalf("rawWrite failed: %v", err)
495 + }
496 +
497 + cfg := defaultServerConfig()
498 + _, err := serverHandshake(serverHandle, &cfg, 1)
499 + if err == nil {
500 + t.Fatal("serverHandshake should fail")
501 + }
502 + if !errors.Is(err, tc.want) {
503 + t.Fatalf("serverHandshake error = %v, want %v", err, tc.want)
504 + }
505 + if !strings.Contains(err.Error(), tc.wantText) {
506 + t.Fatalf("serverHandshake error = %v, want text %q", err, tc.wantText)
507 + }
508 + })
509 + }
510 +}
511 +
512 +func TestServerHandshakeHelloAckSendFailure(t *testing.T) {
513 + clientHandle, serverHandle := rawPipePair(t)
514 +
515 + if err := rawWrite(clientHandle, validHelloPacket()); err != nil {
516 + t.Fatalf("rawWrite(hello) failed: %v", err)
517 + }
518 + if err := syscall.CloseHandle(clientHandle); err != nil {
519 + t.Fatalf("CloseHandle(client) failed: %v", err)
520 + }
521 +
522 + cfg := defaultServerConfig()
523 + _, err := serverHandshake(serverHandle, &cfg, 7)
524 + if err == nil {
525 + t.Fatal("serverHandshake should fail")
526 + }
527 + if !errors.Is(err, ErrSend) {
528 + t.Fatalf("serverHandshake error = %v, want ErrSend", err)
529 + }
530 + if !strings.Contains(err.Error(), "hello_ack send") {
531 + t.Fatalf("serverHandshake error = %v, want text %q", err, "hello_ack send")
532 + }
533 +}
534 +
535 +func TestPipeListenRejectsBadServiceNames(t *testing.T) {
536 + if _, err := Listen(testPipeRunDir, "bad/name", defaultServerConfig()); !errors.Is(err, ErrBadParam) {
537 + t.Fatalf("Listen(invalid service) = %v, want ErrBadParam", err)
538 + }
539 +
540 + service := strings.Repeat("a", maxPipeNameChars)
541 + if _, err := Listen(testPipeRunDir, service, defaultServerConfig()); !errors.Is(err, ErrPipeName) {
542 + t.Fatalf("Listen(long service) = %v, want ErrPipeName", err)
543 + }
544 +}
545 +
546 +func TestSessionReceiveRejectsMalformedMessages(t *testing.T) {
547 + cases := []struct {
548 + name string
549 + sendFn func(t *testing.T, sender syscall.Handle)
550 + want error
551 + wantText string
552 + }{
553 + {
554 + name: "packet too short",
555 + sendFn: func(t *testing.T, sender syscall.Handle) {
556 + t.Helper()
557 + if err := rawWrite(sender, []byte{1, 2, 3, 4}); err != nil {
558 + t.Fatalf("rawWrite failed: %v", err)
559 + }
560 + },
561 + want: ErrProtocol,
562 + wantText: "packet too short for header",
563 + },
564 + {
565 + name: "bad header",
566 + sendFn: func(t *testing.T, sender syscall.Handle) {
567 + t.Helper()
568 + hdr := protocol.Header{
569 + Magic: 0,
570 + Version: protocol.Version,
571 + HeaderLen: protocol.HeaderLen,
572 + Kind: protocol.KindRequest,
573 + Code: protocol.MethodIncrement,
574 + ItemCount: 1,
575 + }
576 + pkt := make([]byte, protocol.HeaderSize)
577 + hdr.Encode(pkt)
578 + if err := rawWrite(sender, pkt); err != nil {
579 + t.Fatalf("rawWrite failed: %v", err)
580 + }
581 + },
582 + want: ErrProtocol,
583 + wantText: "header decode",
584 + },
585 + {
586 + name: "payload exceeds negotiated max",
587 + sendFn: func(t *testing.T, sender syscall.Handle) {
588 + t.Helper()
589 + pkt := encodePacket(protocol.Header{
590 + Kind: protocol.KindRequest,
591 + Code: protocol.MethodIncrement,
592 + ItemCount: 1,
593 + PayloadLen: 0, // overwritten by encodePacket
594 + }, nil)
595 + hdr, err := protocol.DecodeHeader(pkt[:protocol.HeaderSize])
596 + if err != nil {
597 + t.Fatalf("DecodeHeader failed: %v", err)
598 + }
599 + hdr.PayloadLen = 5000
600 + hdr.Encode(pkt[:protocol.HeaderSize])
601 + if err := rawWrite(sender, pkt[:protocol.HeaderSize]); err != nil {
602 + t.Fatalf("rawWrite failed: %v", err)
603 + }
604 + },
605 + want: ErrLimitExceeded,
606 + wantText: "payload_len",
607 + },
608 + {
609 + name: "item count exceeds negotiated max",
610 + sendFn: func(t *testing.T, sender syscall.Handle) {
611 + t.Helper()
612 + pkt := encodePacket(protocol.Header{
613 + Kind: protocol.KindRequest,
614 + Code: protocol.MethodIncrement,
615 + ItemCount: 99,
616 + }, nil)
617 + if err := rawWrite(sender, pkt); err != nil {
618 + t.Fatalf("rawWrite failed: %v", err)
619 + }
620 + },
621 + want: ErrLimitExceeded,
622 + wantText: "item_count",
623 + },
624 + {
625 + name: "batch dir exceeds payload",
626 + sendFn: func(t *testing.T, sender syscall.Handle) {
627 + t.Helper()
628 + payload := make([]byte, 8)
629 + pkt := encodePacket(protocol.Header{
630 + Kind: protocol.KindRequest,
631 + Code: protocol.MethodIncrement,
632 + Flags: protocol.FlagBatch,
633 + ItemCount: 2,
634 + }, payload)
635 + if err := rawWrite(sender, pkt); err != nil {
636 + t.Fatalf("rawWrite failed: %v", err)
637 + }
638 + },
639 + want: ErrProtocol,
640 + wantText: "batch dir exceeds payload",
641 + },
642 + {
643 + name: "invalid batch dir",
644 + sendFn: func(t *testing.T, sender syscall.Handle) {
645 + t.Helper()
646 + payload := make([]byte, 24)
647 + binary.NativeEndian.PutUint32(payload[0:4], 1)
648 + binary.NativeEndian.PutUint32(payload[4:8], 4)
649 + binary.NativeEndian.PutUint32(payload[8:12], 0)
650 + binary.NativeEndian.PutUint32(payload[12:16], 4)
651 + copy(payload[16:], []byte("payload!!"))
652 +
653 + pkt := encodePacket(protocol.Header{
654 + Kind: protocol.KindRequest,
655 + Code: protocol.MethodIncrement,
656 + Flags: protocol.FlagBatch,
657 + ItemCount: 2,
658 + }, payload)
659 + if err := rawWrite(sender, pkt); err != nil {
660 + t.Fatalf("rawWrite failed: %v", err)
661 + }
662 + },
663 + want: ErrProtocol,
664 + wantText: "batch dir",
665 + },
666 + }
667 +
668 + for _, tc := range cases {
669 + t.Run(tc.name, func(t *testing.T) {
670 + client, server := sessionPair(t, defaultServerConfig(), defaultClientConfig())
671 + tc.sendFn(t, client.handle)
672 +
673 + buf := make([]byte, 4096)
674 + _, _, err := server.Receive(buf)
675 + if err == nil {
676 + t.Fatal("Receive should fail")
677 + }
678 + if !errors.Is(err, tc.want) {
679 + t.Fatalf("Receive error = %v, want %v", err, tc.want)
680 + }
681 + if !strings.Contains(err.Error(), tc.wantText) {
682 + t.Fatalf("Receive error = %v, want text %q", err, tc.wantText)
683 + }
684 + })
685 + }
686 +}
687 +
688 +func TestSessionReceiveRejectsMalformedChunks(t *testing.T) {
689 + type chunkSpec struct {
690 + totalPayloadLen uint32
691 + firstPayload []byte
692 + chunkHeader protocol.ChunkHeader
693 + chunkPayload []byte
694 + closeSender bool
695 + }
696 +
697 + cases := []struct {
698 + name string
699 + spec chunkSpec
700 + want error
701 + wantText string
702 + }{
703 + {
704 + name: "bad chunk header",
705 + spec: chunkSpec{
706 + totalPayloadLen: 20,
707 + firstPayload: []byte("0123456789"),
708 + chunkHeader: protocol.ChunkHeader{
709 + Magic: protocol.MagicChunk,
710 + Version: protocol.Version + 1,
711 + MessageID: 7,
712 + TotalMessageLen: uint32(protocol.HeaderSize + 20),
713 + ChunkIndex: 1,
714 + ChunkCount: 2,
715 + ChunkPayloadLen: 10,
716 + },
717 + chunkPayload: []byte("abcdefghij"),
718 + },
719 + want: ErrChunk,
720 + wantText: "chunk header",
721 + },
722 + {
723 + name: "continuation recv disconnect",
724 + spec: chunkSpec{
725 + totalPayloadLen: 20,
726 + firstPayload: []byte("0123456789"),
727 + closeSender: true,
728 + },
729 + want: ErrRecv,
730 + wantText: "continuation recv",
731 + },
732 + {
733 + name: "continuation too short",
734 + spec: chunkSpec{
735 + totalPayloadLen: 20,
736 + firstPayload: []byte("0123456789"),
737 + chunkPayload: []byte{1, 2, 3},
738 + },
739 + want: ErrChunk,
740 + wantText: "continuation too short",
741 + },
742 + {
743 + name: "message id mismatch",
744 + spec: chunkSpec{
745 + totalPayloadLen: 20,
746 + firstPayload: []byte("0123456789"),
747 + chunkHeader: protocol.ChunkHeader{
748 + Magic: protocol.MagicChunk,
749 + Version: protocol.Version,
750 + MessageID: 999,
751 + TotalMessageLen: uint32(protocol.HeaderSize + 20),
752 + ChunkIndex: 1,
753 + ChunkCount: 2,
754 + ChunkPayloadLen: 10,
755 + },
756 + chunkPayload: []byte("abcdefghij"),
757 + },
758 + want: ErrChunk,
759 + wantText: "message_id mismatch",
760 + },
761 + {
762 + name: "chunk index mismatch",
763 + spec: chunkSpec{
764 + totalPayloadLen: 20,
765 + firstPayload: []byte("0123456789"),
766 + chunkHeader: protocol.ChunkHeader{
767 + Magic: protocol.MagicChunk,
768 + Version: protocol.Version,
769 + MessageID: 7,
770 + TotalMessageLen: uint32(protocol.HeaderSize + 20),
771 + ChunkIndex: 2,
772 + ChunkCount: 2,
773 + ChunkPayloadLen: 10,
774 + },
775 + chunkPayload: []byte("abcdefghij"),
776 + },
777 + want: ErrChunk,
778 + wantText: "chunk_index mismatch",
779 + },
780 + {
781 + name: "chunk count mismatch",
782 + spec: chunkSpec{
783 + totalPayloadLen: 20,
784 + firstPayload: []byte("0123456789"),
785 + chunkHeader: protocol.ChunkHeader{
786 + Magic: protocol.MagicChunk,
787 + Version: protocol.Version,
788 + MessageID: 7,
789 + TotalMessageLen: uint32(protocol.HeaderSize + 20),
790 + ChunkIndex: 1,
791 + ChunkCount: 3,
792 + ChunkPayloadLen: 10,
793 + },
794 + chunkPayload: []byte("abcdefghij"),
795 + },
796 + want: ErrChunk,
797 + wantText: "chunk_count mismatch",
798 + },
799 + {
800 + name: "total message len mismatch",
801 + spec: chunkSpec{
802 + totalPayloadLen: 20,
803 + firstPayload: []byte("0123456789"),
804 + chunkHeader: protocol.ChunkHeader{
805 + Magic: protocol.MagicChunk,
806 + Version: protocol.Version,
807 + MessageID: 7,
808 + TotalMessageLen: uint32(protocol.HeaderSize + 21),
809 + ChunkIndex: 1,
810 + ChunkCount: 2,
811 + ChunkPayloadLen: 10,
812 + },
813 + chunkPayload: []byte("abcdefghij"),
814 + },
815 + want: ErrChunk,
816 + wantText: "total_message_len mismatch",
817 + },
818 + {
819 + name: "chunk payload len mismatch",
820 + spec: chunkSpec{
821 + totalPayloadLen: 20,
822 + firstPayload: []byte("0123456789"),
823 + chunkHeader: protocol.ChunkHeader{
824 + Magic: protocol.MagicChunk,
825 + Version: protocol.Version,
826 + MessageID: 7,
827 + TotalMessageLen: uint32(protocol.HeaderSize + 20),
828 + ChunkIndex: 1,
829 + ChunkCount: 2,
830 + ChunkPayloadLen: 9,
831 + },
832 + chunkPayload: []byte("abcdefghij"),
833 + },
834 + want: ErrChunk,
835 + wantText: "chunk_payload_len mismatch",
836 + },
837 + {
838 + name: "chunk exceeds payload len",
839 + spec: chunkSpec{
840 + totalPayloadLen: 15,
841 + firstPayload: []byte("0123456789"),
842 + chunkHeader: protocol.ChunkHeader{
843 + Magic: protocol.MagicChunk,
844 + Version: protocol.Version,
845 + MessageID: 7,
846 + TotalMessageLen: uint32(protocol.HeaderSize + 15),
847 + ChunkIndex: 1,
848 + ChunkCount: 2,
849 + ChunkPayloadLen: 10,
850 + },
851 + chunkPayload: []byte("abcdefghij"),
852 + },
853 + want: ErrChunk,
854 + wantText: "chunk exceeds payload_len",
855 + },
856 + }
857 +
858 + for _, tc := range cases {
859 + t.Run(tc.name, func(t *testing.T) {
860 + client, server := sessionPair(t, defaultServerConfig(), defaultClientConfig())
861 +
862 + first := encodePacket(protocol.Header{
863 + Kind: protocol.KindRequest,
864 + Code: protocol.MethodIncrement,
865 + ItemCount: 1,
866 + MessageID: 7,
867 + }, tc.spec.firstPayload)
868 + hdr, err := protocol.DecodeHeader(first[:protocol.HeaderSize])
869 + if err != nil {
870 + t.Fatalf("DecodeHeader failed: %v", err)
871 + }
872 + hdr.PayloadLen = tc.spec.totalPayloadLen
873 + hdr.Encode(first[:protocol.HeaderSize])
874 +
875 + if err := rawWrite(client.handle, first); err != nil {
876 + t.Fatalf("rawWrite(first) failed: %v", err)
877 + }
878 +
879 + if tc.spec.closeSender {
880 + syscall.CloseHandle(client.handle)
881 + client.handle = syscall.InvalidHandle
882 + } else {
883 + var pkt []byte
884 + if len(tc.spec.chunkPayload) < protocol.HeaderSize && tc.spec.chunkHeader.Magic == 0 {
885 + pkt = tc.spec.chunkPayload
886 + } else {
887 + pkt = encodeChunkPacket(tc.spec.chunkHeader, tc.spec.chunkPayload)
888 + }
889 + if err := rawWrite(client.handle, pkt); err != nil {
890 + t.Fatalf("rawWrite(chunk) failed: %v", err)
891 + }
892 + }
893 +
894 + buf := make([]byte, 4096)
895 + _, _, err = server.Receive(buf)
896 + if err == nil {
897 + t.Fatal("Receive should fail")
898 + }
899 + if !errors.Is(err, tc.want) {
900 + t.Fatalf("Receive error = %v, want %v", err, tc.want)
901 + }
902 + if !strings.Contains(err.Error(), tc.wantText) {
903 + t.Fatalf("Receive error = %v, want text %q", err, tc.wantText)
904 + }
905 + })
906 + }
907 +}
908 +
909 +func TestSessionReceiveRejectsMalformedChunkedBatchMessages(t *testing.T) {
910 + type batchCase struct {
911 + name string
912 + totalPayload uint32
913 + firstPayload []byte
914 + restPayload []byte
915 + wantText string
916 + }
917 +
918 + cases := []batchCase{
919 + {
920 + name: "batch dir exceeds payload after chunking",
921 + totalPayload: 12,
922 + firstPayload: []byte("abcd"),
923 + restPayload: []byte("efghijkl"),
924 + wantText: "batch dir exceeds payload",
925 + },
926 + {
927 + name: "invalid batch dir after chunking",
928 + totalPayload: 24,
929 + firstPayload: func() []byte {
930 + payload := make([]byte, 24)
931 + binary.NativeEndian.PutUint32(payload[0:4], 1)
932 + binary.NativeEndian.PutUint32(payload[4:8], 4)
933 + binary.NativeEndian.PutUint32(payload[8:12], 0)
934 + binary.NativeEndian.PutUint32(payload[12:16], 4)
935 + copy(payload[16:], []byte("payload!!"))
936 + return payload[:8]
937 + }(),
938 + restPayload: func() []byte {
939 + payload := make([]byte, 24)
940 + binary.NativeEndian.PutUint32(payload[0:4], 1)
941 + binary.NativeEndian.PutUint32(payload[4:8], 4)
942 + binary.NativeEndian.PutUint32(payload[8:12], 0)
943 + binary.NativeEndian.PutUint32(payload[12:16], 4)
944 + copy(payload[16:], []byte("payload!!"))
945 + return payload[8:]
946 + }(),
947 + wantText: "batch dir",
948 + },
949 + }
950 +
951 + for _, tc := range cases {
952 + t.Run(tc.name, func(t *testing.T) {
953 + client, server := sessionPair(t, defaultServerConfig(), defaultClientConfig())
954 +
955 + first := encodePacket(protocol.Header{
956 + Kind: protocol.KindRequest,
957 + Code: protocol.MethodIncrement,
958 + Flags: protocol.FlagBatch,
959 + ItemCount: 2,
960 + MessageID: 7,
961 + }, tc.firstPayload)
962 + hdr, err := protocol.DecodeHeader(first[:protocol.HeaderSize])
963 + if err != nil {
964 + t.Fatalf("DecodeHeader failed: %v", err)
965 + }
966 + hdr.PayloadLen = tc.totalPayload
967 + hdr.Encode(first[:protocol.HeaderSize])
968 +
969 + chunk := encodeChunkPacket(protocol.ChunkHeader{
970 + Magic: protocol.MagicChunk,
971 + Version: protocol.Version,
972 + MessageID: 7,
973 + TotalMessageLen: uint32(protocol.HeaderSize) + tc.totalPayload,
974 + ChunkIndex: 1,
975 + ChunkCount: 2,
976 + ChunkPayloadLen: uint32(len(tc.restPayload)),
977 + }, tc.restPayload)
978 +
979 + if err := rawWrite(client.handle, first); err != nil {
980 + t.Fatalf("rawWrite(first) failed: %v", err)
981 + }
982 + if err := rawWrite(client.handle, chunk); err != nil {
983 + t.Fatalf("rawWrite(chunk) failed: %v", err)
984 + }
985 +
986 + buf := make([]byte, 4096)
987 + _, _, err = server.Receive(buf)
988 + if err == nil {
989 + t.Fatal("Receive should fail")
990 + }
991 + if !errors.Is(err, ErrProtocol) {
992 + t.Fatalf("Receive error = %v, want ErrProtocol", err)
993 + }
994 + if !strings.Contains(err.Error(), tc.wantText) {
995 + t.Fatalf("Receive error = %v, want text %q", err, tc.wantText)
996 + }
997 + })
998 + }
999 +}
1000 +
1001 +func TestSessionSendClearsInflightOnDisconnect(t *testing.T) {
1002 + client, server := sessionPair(t, defaultServerConfig(), defaultClientConfig())
1003 + server.Close()
1004 +
1005 + hdr := protocol.Header{
1006 + Kind: protocol.KindRequest,
1007 + Code: protocol.MethodIncrement,
1008 + ItemCount: 1,
1009 + MessageID: 77,
1010 + }
1011 + err := client.Send(&hdr, []byte("payload"))
1012 + if err == nil {
1013 + t.Fatal("Send should fail after peer disconnect")
1014 + }
1015 + if !errors.Is(err, ErrDisconnected) {
1016 + t.Fatalf("Send error = %v, want ErrDisconnected", err)
1017 + }
1018 + if _, exists := client.inflightIDs[77]; exists {
1019 + t.Fatalf("message_id 77 should be removed from inflightIDs after send failure")
1020 + }
1021 +}
1022 +
1023 +func TestSessionReceiveReportsPeerDisconnect(t *testing.T) {
1024 + client, server := sessionPair(t, defaultServerConfig(), defaultClientConfig())
1025 + server.Close()
1026 +
1027 + buf := make([]byte, 4096)
1028 + _, _, err := client.Receive(buf)
1029 + if err == nil {
1030 + t.Fatal("Receive should fail after peer disconnect")
1031 + }
1032 + if !errors.Is(err, ErrDisconnected) {
1033 + t.Fatalf("Receive error = %v, want ErrDisconnected", err)
1034 + }
1035 +}
1036 +
1037 +func TestSessionChunkedSendDisconnectPaths(t *testing.T) {
1038 + t.Run("first chunk send failure", func(t *testing.T) {
1039 + sCfg := defaultServerConfig()
1040 + cCfg := defaultClientConfig()
1041 + sCfg.PacketSize = 96
1042 + cCfg.PacketSize = 96
1043 +
1044 + client, server := sessionPair(t, sCfg, cCfg)
1045 + if err := syscall.CloseHandle(server.handle); err != nil {
1046 + t.Fatalf("CloseHandle(server) failed: %v", err)
1047 + }
1048 + server.handle = syscall.InvalidHandle
1049 +
1050 + hdr := protocol.Header{
1051 + Kind: protocol.KindRequest,
1052 + Code: protocol.MethodIncrement,
1053 + ItemCount: 1,
1054 + MessageID: 88,
1055 + }
1056 + payload := bytes.Repeat([]byte("x"), 256)
1057 + err := client.Send(&hdr, payload)
1058 + if err == nil {
1059 + t.Fatal("Send should fail after peer disconnect")
1060 + }
1061 + if !errors.Is(err, ErrDisconnected) {
1062 + t.Fatalf("Send error = %v, want ErrDisconnected", err)
1063 + }
1064 + if _, exists := client.inflightIDs[88]; exists {
1065 + t.Fatalf("message_id 88 should be removed from inflightIDs after chunked send failure")
1066 + }
1067 + })
1068 +
1069 + t.Run("continuation chunk send failure", func(t *testing.T) {
1070 + sCfg := defaultServerConfig()
1071 + cCfg := defaultClientConfig()
1072 + sCfg.PacketSize = 96
1073 + cCfg.PacketSize = 96
1074 +
1075 + client, server := sessionPair(t, sCfg, cCfg)
1076 +
1077 + closeDone := make(chan error, 1)
1078 + go func() {
1079 + buf := make([]byte, 96)
1080 + _, err := rawRecv(server.handle, buf)
1081 + if err == nil {
1082 + err = syscall.CloseHandle(server.handle)
1083 + server.handle = syscall.InvalidHandle
1084 + }
1085 + closeDone <- err
1086 + }()
1087 +
1088 + hdr := protocol.Header{
1089 + Kind: protocol.KindRequest,
1090 + Code: protocol.MethodIncrement,
1091 + ItemCount: 1,
1092 + MessageID: 89,
1093 + }
1094 + payload := bytes.Repeat([]byte("y"), 4096)
1095 + err := client.Send(&hdr, payload)
1096 + if err == nil {
1097 + buf := make([]byte, 96)
1098 + _, _, err = client.Receive(buf)
1099 + if err == nil {
1100 + t.Fatal("session should observe disconnect after peer close")
1101 + }
1102 + }
1103 + if !errors.Is(err, ErrDisconnected) && !errors.Is(err, ErrRecv) {
1104 + t.Fatalf("error = %v, want ErrDisconnected or ErrRecv", err)
1105 + }
1106 + if _, exists := client.inflightIDs[89]; exists {
1107 + t.Fatalf("message_id 89 should be removed from inflightIDs after session disconnect")
1108 + }
1109 + if err := <-closeDone; err != nil {
1110 + t.Fatalf("server close goroutine failed: %v", err)
1111 + }
1112 + })
1113 +}
1114 +
1115 +func TestSessionSendRejectsTooSmallPacketSize(t *testing.T) {
1116 + sCfg := defaultServerConfig()
1117 + cCfg := defaultClientConfig()
1118 + sCfg.PacketSize = 16
1119 + cCfg.PacketSize = 16
1120 +
1121 + service := uniquePipeService(t)
1122 + listener := startListener(t, testPipeRunDir, service, sCfg)
1123 + defer listener.Close()
1124 +
1125 + acceptCh := acceptAsync(listener)
1126 +
1127 + if _, err := Connect(testPipeRunDir, service, &cCfg); !errors.Is(err, ErrIncompatible) {
1128 + t.Fatalf("expected ErrIncompatible, got %v", err)
1129 + }
1130 +
1131 + sr := <-acceptCh
1132 + if !errors.Is(sr.err, ErrIncompatible) {
1133 + t.Fatalf("server expected ErrIncompatible, got %v", sr.err)
1134 + }
1135 +}
src/go/pkg/netipc/transport/windows/pipe_integration_test.go new
+1326
@@ -0,0 +1,1326 @@
1 +//go:build windows
2 +
3 +package windows
4 +
5 +import (
6 + "bytes"
7 + "encoding/binary"
8 + "errors"
9 + "fmt"
10 + "os"
11 + "sync"
12 + "sync/atomic"
13 + "syscall"
14 + "testing"
15 + "time"
16 +
17 + "github.com/netdata/netdata/go/plugins/pkg/netipc/protocol"
18 +)
19 +
20 +const (
21 + testPipeRunDir = `C:\ProgramData\nipc_go_test`
22 + testAuthToken = uint64(0xDEADBEEFCAFEBABE)
23 +)
24 +
25 +var testPipeCounter atomic.Uint64
26 +
27 +func uniquePipeService(t *testing.T) string {
28 + t.Helper()
29 + return fmt.Sprintf("go_win_pipe_%d_%d", os.Getpid(), testPipeCounter.Add(1))
30 +}
31 +
32 +func defaultServerConfig() ServerConfig {
33 + return ServerConfig{
34 + SupportedProfiles: protocol.ProfileBaseline,
35 + MaxRequestPayloadBytes: 4096,
36 + MaxRequestBatchItems: 16,
37 + MaxResponsePayloadBytes: 4096,
38 + MaxResponseBatchItems: 16,
39 + AuthToken: testAuthToken,
40 + }
41 +}
42 +
43 +func defaultClientConfig() ClientConfig {
44 + return ClientConfig{
45 + SupportedProfiles: protocol.ProfileBaseline,
46 + MaxRequestPayloadBytes: 4096,
47 + MaxRequestBatchItems: 16,
48 + MaxResponsePayloadBytes: 4096,
49 + MaxResponseBatchItems: 16,
50 + AuthToken: testAuthToken,
51 + }
52 +}
53 +
54 +func defaultClientConfigPtr() *ClientConfig {
55 + cfg := defaultClientConfig()
56 + return &cfg
57 +}
58 +
59 +type serverResult struct {
60 + session *Session
61 + err error
62 +}
63 +
64 +func startListener(t *testing.T, runDir, service string, cfg ServerConfig) *Listener {
65 + t.Helper()
66 + listener, err := Listen(runDir, service, cfg)
67 + if err != nil {
68 + t.Fatalf("Listen failed: %v", err)
69 + }
70 + return listener
71 +}
72 +
73 +func acceptAsync(listener *Listener) <-chan serverResult {
74 + ch := make(chan serverResult, 1)
75 + go func() {
76 + session, err := listener.Accept()
77 + ch <- serverResult{session: session, err: err}
78 + }()
79 + return ch
80 +}
81 +
82 +func TestPipeSingleClientPingPong(t *testing.T) {
83 + service := uniquePipeService(t)
84 + listener := startListener(t, testPipeRunDir, service, defaultServerConfig())
85 + defer listener.Close()
86 +
87 + acceptCh := acceptAsync(listener)
88 +
89 + client, err := Connect(testPipeRunDir, service, defaultClientConfigPtr())
90 + if err != nil {
91 + t.Fatalf("Connect failed: %v", err)
92 + }
93 + defer client.Close()
94 +
95 + sr := <-acceptCh
96 + if sr.err != nil {
97 + t.Fatalf("Accept failed: %v", sr.err)
98 + }
99 + server := sr.session
100 + defer server.Close()
101 +
102 + payload := []byte("hello from client")
103 + hdr := protocol.Header{
104 + Kind: protocol.KindRequest,
105 + Code: protocol.MethodIncrement,
106 + ItemCount: 1,
107 + MessageID: 42,
108 + }
109 + if err := client.Send(&hdr, payload); err != nil {
110 + t.Fatalf("client Send: %v", err)
111 + }
112 +
113 + recvBuf := make([]byte, 4096)
114 + rHdr, rPayload, err := server.Receive(recvBuf)
115 + if err != nil {
116 + t.Fatalf("server Receive: %v", err)
117 + }
118 + if rHdr.Kind != protocol.KindRequest || rHdr.MessageID != 42 || !bytes.Equal(rPayload, payload) {
119 + t.Fatalf("unexpected server receive: hdr=%+v payload=%q", rHdr, rPayload)
120 + }
121 +
122 + respHdr := protocol.Header{
123 + Kind: protocol.KindResponse,
124 + Code: protocol.MethodIncrement,
125 + ItemCount: 1,
126 + MessageID: 42,
127 + }
128 + respPayload := []byte("response from server")
129 + if err := server.Send(&respHdr, respPayload); err != nil {
130 + t.Fatalf("server Send: %v", err)
131 + }
132 +
133 + rHdr, rPayload, err = client.Receive(recvBuf)
134 + if err != nil {
135 + t.Fatalf("client Receive: %v", err)
136 + }
137 + if rHdr.Kind != protocol.KindResponse || rHdr.MessageID != 42 || !bytes.Equal(rPayload, respPayload) {
138 + t.Fatalf("unexpected client receive: hdr=%+v payload=%q", rHdr, rPayload)
139 + }
140 +}
141 +
142 +func TestPipeMultiClient(t *testing.T) {
143 + service := uniquePipeService(t)
144 + listener := startListener(t, testPipeRunDir, service, defaultServerConfig())
145 + defer listener.Close()
146 +
147 + const numClients = 2
148 + clients := make([]*Session, numClients)
149 + servers := make([]*Session, numClients)
150 +
151 + for i := 0; i < numClients; i++ {
152 + acceptCh := acceptAsync(listener)
153 + c, err := Connect(testPipeRunDir, service, defaultClientConfigPtr())
154 + if err != nil {
155 + t.Fatalf("Connect[%d] failed: %v", i, err)
156 + }
157 + clients[i] = c
158 +
159 + sr := <-acceptCh
160 + if sr.err != nil {
161 + t.Fatalf("Accept[%d] failed: %v", i, sr.err)
162 + }
163 + servers[i] = sr.session
164 + }
165 +
166 + defer func() {
167 + for i := 0; i < numClients; i++ {
168 + clients[i].Close()
169 + servers[i].Close()
170 + }
171 + }()
172 +
173 + for i := 0; i < numClients; i++ {
174 + payload := []byte(fmt.Sprintf("client_%d", i))
175 + hdr := protocol.Header{
176 + Kind: protocol.KindRequest,
177 + Code: protocol.MethodIncrement,
178 + ItemCount: 1,
179 + MessageID: uint64(100 + i),
180 + }
181 + if err := clients[i].Send(&hdr, payload); err != nil {
182 + t.Fatalf("client[%d] Send: %v", i, err)
183 + }
184 + }
185 +
186 + buf := make([]byte, 4096)
187 + for i := 0; i < numClients; i++ {
188 + rHdr, rPayload, err := servers[i].Receive(buf)
189 + if err != nil {
190 + t.Fatalf("server[%d] Receive: %v", i, err)
191 + }
192 + expected := []byte(fmt.Sprintf("client_%d", i))
193 + if rHdr.MessageID != uint64(100+i) || !bytes.Equal(rPayload, expected) {
194 + t.Fatalf("server[%d] unexpected receive: hdr=%+v payload=%q", i, rHdr, rPayload)
195 + }
196 +
197 + resp := protocol.Header{
198 + Kind: protocol.KindResponse,
199 + Code: rHdr.Code,
200 + ItemCount: 1,
201 + MessageID: rHdr.MessageID,
202 + }
203 + if err := servers[i].Send(&resp, rPayload); err != nil {
204 + t.Fatalf("server[%d] Send: %v", i, err)
205 + }
206 + }
207 +
208 + for i := 0; i < numClients; i++ {
209 + rHdr, rPayload, err := clients[i].Receive(buf)
210 + if err != nil {
211 + t.Fatalf("client[%d] Receive: %v", i, err)
212 + }
213 + expected := []byte(fmt.Sprintf("client_%d", i))
214 + if rHdr.MessageID != uint64(100+i) || !bytes.Equal(rPayload, expected) {
215 + t.Fatalf("client[%d] unexpected receive: hdr=%+v payload=%q", i, rHdr, rPayload)
216 + }
217 + }
218 +}
219 +
220 +func TestPipePipelining(t *testing.T) {
221 + service := uniquePipeService(t)
222 + listener := startListener(t, testPipeRunDir, service, defaultServerConfig())
223 + defer listener.Close()
224 +
225 + acceptCh := acceptAsync(listener)
226 +
227 + client, err := Connect(testPipeRunDir, service, defaultClientConfigPtr())
228 + if err != nil {
229 + t.Fatalf("Connect failed: %v", err)
230 + }
231 + defer client.Close()
232 +
233 + sr := <-acceptCh
234 + if sr.err != nil {
235 + t.Fatalf("Accept failed: %v", sr.err)
236 + }
237 + server := sr.session
238 + defer server.Close()
239 +
240 + messageIDs := []uint64{10, 20, 30}
241 + for _, mid := range messageIDs {
242 + hdr := protocol.Header{
243 + Kind: protocol.KindRequest,
244 + Code: protocol.MethodIncrement,
245 + ItemCount: 1,
246 + MessageID: mid,
247 + }
248 + payload := []byte(fmt.Sprintf("req_%d", mid))
249 + if err := client.Send(&hdr, payload); err != nil {
250 + t.Fatalf("client Send(%d): %v", mid, err)
251 + }
252 + }
253 +
254 + buf := make([]byte, 4096)
255 + type reqInfo struct {
256 + hdr protocol.Header
257 + payload []byte
258 + }
259 + reqs := make([]reqInfo, 0, len(messageIDs))
260 + for i := 0; i < len(messageIDs); i++ {
261 + rHdr, rPayload, err := server.Receive(buf)
262 + if err != nil {
263 + t.Fatalf("server Receive[%d]: %v", i, err)
264 + }
265 + reqs = append(reqs, reqInfo{hdr: rHdr, payload: append([]byte(nil), rPayload...)})
266 + }
267 +
268 + for i := len(reqs) - 1; i >= 0; i-- {
269 + resp := protocol.Header{
270 + Kind: protocol.KindResponse,
271 + Code: reqs[i].hdr.Code,
272 + ItemCount: 1,
273 + MessageID: reqs[i].hdr.MessageID,
274 + }
275 + respPayload := append([]byte("resp_"), reqs[i].payload...)
276 + if err := server.Send(&resp, respPayload); err != nil {
277 + t.Fatalf("server Send[%d]: %v", i, err)
278 + }
279 + }
280 +
281 + received := make(map[uint64][]byte)
282 + for i := 0; i < len(messageIDs); i++ {
283 + rHdr, rPayload, err := client.Receive(buf)
284 + if err != nil {
285 + t.Fatalf("client Receive[%d]: %v", i, err)
286 + }
287 + received[rHdr.MessageID] = append([]byte(nil), rPayload...)
288 + }
289 +
290 + for _, mid := range messageIDs {
291 + want := []byte(fmt.Sprintf("resp_req_%d", mid))
292 + if got, ok := received[mid]; !ok || !bytes.Equal(got, want) {
293 + t.Fatalf("missing or mismatched pipeline response for %d: %q", mid, got)
294 + }
295 + }
296 +}
297 +
298 +func TestPipeChunking(t *testing.T) {
299 + service := uniquePipeService(t)
300 + const forcedPacketSize = 128
301 +
302 + sCfg := defaultServerConfig()
303 + sCfg.PacketSize = forcedPacketSize
304 + sCfg.MaxRequestPayloadBytes = 65536
305 + sCfg.MaxResponsePayloadBytes = 65536
306 + listener := startListener(t, testPipeRunDir, service, sCfg)
307 + defer listener.Close()
308 +
309 + acceptCh := acceptAsync(listener)
310 +
311 + cCfg := defaultClientConfig()
312 + cCfg.PacketSize = forcedPacketSize
313 + cCfg.MaxRequestPayloadBytes = 65536
314 + cCfg.MaxResponsePayloadBytes = 65536
315 + client, err := Connect(testPipeRunDir, service, &cCfg)
316 + if err != nil {
317 + t.Fatalf("Connect failed: %v", err)
318 + }
319 + defer client.Close()
320 +
321 + sr := <-acceptCh
322 + if sr.err != nil {
323 + t.Fatalf("Accept failed: %v", sr.err)
324 + }
325 + server := sr.session
326 + defer server.Close()
327 +
328 + if client.PacketSize > forcedPacketSize {
329 + t.Fatalf("client packet_size=%d, want <= %d", client.PacketSize, forcedPacketSize)
330 + }
331 +
332 + largePayload := make([]byte, 2000)
333 + for i := range largePayload {
334 + largePayload[i] = byte(i & 0xFF)
335 + }
336 +
337 + hdr := protocol.Header{
338 + Kind: protocol.KindRequest,
339 + Code: protocol.MethodIncrement,
340 + ItemCount: 1,
341 + MessageID: 99,
342 + }
343 + recvBuf := make([]byte, forcedPacketSize)
344 + type receiveResult struct {
345 + hdr protocol.Header
346 + payload []byte
347 + err error
348 + }
349 + serverRecvCh := make(chan receiveResult, 1)
350 + go func() {
351 + rHdr, rPayload, err := server.Receive(recvBuf)
352 + if err == nil {
353 + rPayload = append([]byte(nil), rPayload...)
354 + }
355 + serverRecvCh <- receiveResult{hdr: rHdr, payload: rPayload, err: err}
356 + }()
357 +
358 + if err := client.Send(&hdr, largePayload); err != nil {
359 + t.Fatalf("client Send (chunked): %v", err)
360 + }
361 +
362 + serverRecv := <-serverRecvCh
363 + if serverRecv.err != nil {
364 + t.Fatalf("server Receive (chunked): %v", serverRecv.err)
365 + }
366 + if serverRecv.hdr.MessageID != 99 || !bytes.Equal(serverRecv.payload, largePayload) {
367 + t.Fatalf("unexpected chunked receive: hdr=%+v payload-len=%d", serverRecv.hdr, len(serverRecv.payload))
368 + }
369 +
370 + resp := protocol.Header{
371 + Kind: protocol.KindResponse,
372 + Code: protocol.MethodIncrement,
373 + ItemCount: 1,
374 + MessageID: 99,
375 + }
376 + clientRecvCh := make(chan receiveResult, 1)
377 + go func() {
378 + rHdr, rPayload, err := client.Receive(recvBuf)
379 + if err == nil {
380 + rPayload = append([]byte(nil), rPayload...)
381 + }
382 + clientRecvCh <- receiveResult{hdr: rHdr, payload: rPayload, err: err}
383 + }()
384 +
385 + if err := server.Send(&resp, serverRecv.payload); err != nil {
386 + t.Fatalf("server Send (chunked echo): %v", err)
387 + }
388 +
389 + clientRecv := <-clientRecvCh
390 + if clientRecv.err != nil {
391 + t.Fatalf("client Receive (chunked): %v", clientRecv.err)
392 + }
393 + if clientRecv.hdr.MessageID != 99 || !bytes.Equal(clientRecv.payload, largePayload) {
394 + t.Fatalf("unexpected client chunked receive: hdr=%+v payload-len=%d", clientRecv.hdr, len(clientRecv.payload))
395 + }
396 +}
397 +
398 +func TestPipeHandshakeBadAuth(t *testing.T) {
399 + service := uniquePipeService(t)
400 +
401 + sCfg := defaultServerConfig()
402 + sCfg.AuthToken = 0x1111111111111111
403 + listener := startListener(t, testPipeRunDir, service, sCfg)
404 + defer listener.Close()
405 +
406 + acceptCh := acceptAsync(listener)
407 +
408 + cCfg := defaultClientConfig()
409 + cCfg.AuthToken = 0x2222222222222222
410 + if _, err := Connect(testPipeRunDir, service, &cCfg); !errors.Is(err, ErrAuthFailed) {
411 + t.Fatalf("expected ErrAuthFailed, got %v", err)
412 + }
413 +
414 + sr := <-acceptCh
415 + if !errors.Is(sr.err, ErrAuthFailed) {
416 + t.Fatalf("server expected ErrAuthFailed, got %v", sr.err)
417 + }
418 +}
419 +
420 +func TestPipeHandshakeProfileMismatch(t *testing.T) {
421 + service := uniquePipeService(t)
422 +
423 + sCfg := defaultServerConfig()
424 + sCfg.SupportedProfiles = protocol.ProfileSHMFutex
425 + listener := startListener(t, testPipeRunDir, service, sCfg)
426 + defer listener.Close()
427 +
428 + acceptCh := acceptAsync(listener)
429 +
430 + cCfg := defaultClientConfig()
431 + cCfg.SupportedProfiles = protocol.ProfileBaseline
432 + if _, err := Connect(testPipeRunDir, service, &cCfg); !errors.Is(err, ErrNoProfile) {
433 + t.Fatalf("expected ErrNoProfile, got %v", err)
434 + }
435 +
436 + sr := <-acceptCh
437 + if !errors.Is(sr.err, ErrNoProfile) {
438 + t.Fatalf("server expected ErrNoProfile, got %v", sr.err)
439 + }
440 +}
441 +
442 +func TestPipeHandshakeRequestPayloadOverCap(t *testing.T) {
443 + service := uniquePipeService(t)
444 +
445 + listener := startListener(t, testPipeRunDir, service, defaultServerConfig())
446 + defer listener.Close()
447 +
448 + acceptCh := acceptAsync(listener)
449 +
450 + cCfg := defaultClientConfig()
451 + cCfg.MaxRequestPayloadBytes = protocol.MaxPayloadCap + 1
452 + if _, err := Connect(testPipeRunDir, service, &cCfg); !errors.Is(err, ErrLimitExceeded) {
453 + t.Fatalf("expected ErrLimitExceeded, got %v", err)
454 + }
455 +
456 + sr := <-acceptCh
457 + if !errors.Is(sr.err, ErrLimitExceeded) {
458 + t.Fatalf("server expected ErrLimitExceeded, got %v", sr.err)
459 + }
460 +}
461 +
462 +func TestPipeHandshakeIncompatibleClassifierHelpers(t *testing.T) {
463 + hdr := protocol.Header{
464 + Magic: protocol.MagicMsg,
465 + Version: protocol.Version + 1,
466 + HeaderLen: protocol.HeaderLen,
467 + Kind: protocol.KindControl,
468 + Code: protocol.CodeHello,
469 + }
470 + hdrBuf := make([]byte, protocol.HeaderSize)
471 + hdr.Encode(hdrBuf)
472 + if !headerVersionIncompatible(hdrBuf, protocol.CodeHello) {
473 + t.Fatal("headerVersionIncompatible should detect bad HELLO version")
474 + }
475 + if headerVersionIncompatible(hdrBuf, protocol.CodeHelloAck) {
476 + t.Fatal("headerVersionIncompatible should respect expected code")
477 + }
478 +
479 + hello := protocol.Hello{LayoutVersion: 2}
480 + helloBuf := make([]byte, 44)
481 + hello.Encode(helloBuf)
482 + if !helloLayoutIncompatible(helloBuf) {
483 + t.Fatal("helloLayoutIncompatible should detect bad layout_version")
484 + }
485 +
486 + ack := protocol.HelloAck{LayoutVersion: 2}
487 + ackBuf := make([]byte, 48)
488 + ack.Encode(ackBuf)
489 + if !helloAckLayoutIncompatible(ackBuf) {
490 + t.Fatal("helloAckLayoutIncompatible should detect bad layout_version")
491 + }
492 +}
493 +
494 +func TestPipeAddrInUse(t *testing.T) {
495 + service := uniquePipeService(t)
496 + listener := startListener(t, testPipeRunDir, service, defaultServerConfig())
497 + defer listener.Close()
498 +
499 + if _, err := Listen(testPipeRunDir, service, defaultServerConfig()); !errors.Is(err, ErrAddrInUse) {
500 + t.Fatalf("expected ErrAddrInUse, got %v", err)
501 + }
502 +}
503 +
504 +func TestPipeDirectionalLimitNegotiation(t *testing.T) {
505 + service := uniquePipeService(t)
506 +
507 + sCfg := defaultServerConfig()
508 + sCfg.MaxRequestPayloadBytes = 2048
509 + sCfg.MaxRequestBatchItems = 8
510 + sCfg.MaxResponsePayloadBytes = 8192
511 + sCfg.MaxResponseBatchItems = 32
512 + listener := startListener(t, testPipeRunDir, service, sCfg)
513 + defer listener.Close()
514 +
515 + acceptCh := acceptAsync(listener)
516 +
517 + cCfg := defaultClientConfig()
518 + cCfg.MaxRequestPayloadBytes = 4096
519 + cCfg.MaxRequestBatchItems = 16
520 + cCfg.MaxResponsePayloadBytes = 4096
521 + cCfg.MaxResponseBatchItems = 16
522 + client, err := Connect(testPipeRunDir, service, &cCfg)
523 + if err != nil {
524 + t.Fatalf("Connect failed: %v", err)
525 + }
526 + defer client.Close()
527 +
528 + sr := <-acceptCh
529 + if sr.err != nil {
530 + t.Fatalf("Accept failed: %v", sr.err)
531 + }
532 + server := sr.session
533 + defer server.Close()
534 +
535 + if client.MaxRequestPayloadBytes != 4096 || client.MaxRequestBatchItems != 16 || client.MaxResponsePayloadBytes != 8192 || client.MaxResponseBatchItems != 16 {
536 + t.Fatalf("unexpected negotiated client limits: %+v", client)
537 + }
538 + if client.SessionID == 0 {
539 + t.Fatal("session id must be non-zero")
540 + }
541 + if server.MaxRequestPayloadBytes != client.MaxRequestPayloadBytes ||
542 + server.MaxRequestBatchItems != client.MaxRequestBatchItems ||
543 + server.MaxResponsePayloadBytes != client.MaxResponsePayloadBytes ||
544 + server.MaxResponseBatchItems != client.MaxResponseBatchItems ||
545 + server.SessionID != client.SessionID {
546 + t.Fatalf("server/client negotiation mismatch: server=%+v client=%+v", server, client)
547 + }
548 +}
549 +
550 +func TestPipeProfileSelection(t *testing.T) {
551 + service := uniquePipeService(t)
552 +
553 + sCfg := defaultServerConfig()
554 + sCfg.SupportedProfiles = protocol.ProfileBaseline | protocol.ProfileSHMHybrid | protocol.ProfileSHMFutex
555 + sCfg.PreferredProfiles = protocol.ProfileSHMFutex
556 + listener := startListener(t, testPipeRunDir, service, sCfg)
557 + defer listener.Close()
558 +
559 + acceptCh := acceptAsync(listener)
560 +
561 + cCfg := defaultClientConfig()
562 + cCfg.SupportedProfiles = protocol.ProfileBaseline | protocol.ProfileSHMHybrid | protocol.ProfileSHMFutex
563 + cCfg.PreferredProfiles = protocol.ProfileSHMFutex | protocol.ProfileSHMHybrid
564 + client, err := Connect(testPipeRunDir, service, &cCfg)
565 + if err != nil {
566 + t.Fatalf("Connect failed: %v", err)
567 + }
568 + defer client.Close()
569 +
570 + sr := <-acceptCh
571 + if sr.err != nil {
572 + t.Fatalf("Accept failed: %v", sr.err)
573 + }
574 + defer sr.session.Close()
575 +
576 + if client.SelectedProfile != protocol.ProfileSHMFutex {
577 + t.Fatalf("selected profile = 0x%x, want SHMFutex", client.SelectedProfile)
578 + }
579 +}
580 +
581 +func TestPipeEmptyPayload(t *testing.T) {
582 + service := uniquePipeService(t)
583 + listener := startListener(t, testPipeRunDir, service, defaultServerConfig())
584 + defer listener.Close()
585 +
586 + acceptCh := acceptAsync(listener)
587 + client, err := Connect(testPipeRunDir, service, defaultClientConfigPtr())
588 + if err != nil {
589 + t.Fatalf("Connect failed: %v", err)
590 + }
591 + defer client.Close()
592 +
593 + sr := <-acceptCh
594 + if sr.err != nil {
595 + t.Fatalf("Accept failed: %v", sr.err)
596 + }
597 + server := sr.session
598 + defer server.Close()
599 +
600 + hdr := protocol.Header{
601 + Kind: protocol.KindRequest,
602 + Code: protocol.MethodIncrement,
603 + ItemCount: 1,
604 + MessageID: 1,
605 + }
606 + if err := client.Send(&hdr, nil); err != nil {
607 + t.Fatalf("Send empty payload: %v", err)
608 + }
609 +
610 + buf := make([]byte, 4096)
611 + rHdr, rPayload, err := server.Receive(buf)
612 + if err != nil {
613 + t.Fatalf("Receive empty payload: %v", err)
614 + }
615 + if rHdr.MessageID != 1 || len(rPayload) != 0 {
616 + t.Fatalf("unexpected empty payload receive: hdr=%+v len=%d", rHdr, len(rPayload))
617 + }
618 +}
619 +
620 +func TestPipeConcurrentSendReceive(t *testing.T) {
621 + service := uniquePipeService(t)
622 +
623 + sCfg := defaultServerConfig()
624 + sCfg.MaxRequestPayloadBytes = 65536
625 + sCfg.MaxResponsePayloadBytes = 65536
626 + listener := startListener(t, testPipeRunDir, service, sCfg)
627 + defer listener.Close()
628 +
629 + acceptCh := acceptAsync(listener)
630 + cCfg := defaultClientConfig()
631 + cCfg.MaxRequestPayloadBytes = 65536
632 + cCfg.MaxResponsePayloadBytes = 65536
633 + client, err := Connect(testPipeRunDir, service, &cCfg)
634 + if err != nil {
635 + t.Fatalf("Connect failed: %v", err)
636 + }
637 + defer client.Close()
638 +
639 + sr := <-acceptCh
640 + if sr.err != nil {
641 + t.Fatalf("Accept failed: %v", sr.err)
642 + }
643 + server := sr.session
644 + defer server.Close()
645 +
646 + const numMessages = 20
647 + var wg sync.WaitGroup
648 +
649 + wg.Add(1)
650 + go func() {
651 + defer wg.Done()
652 + buf := make([]byte, 65600)
653 + for i := 0; i < numMessages; i++ {
654 + rHdr, rPayload, err := server.Receive(buf)
655 + if err != nil {
656 + t.Errorf("server Receive[%d]: %v", i, err)
657 + return
658 + }
659 + resp := protocol.Header{
660 + Kind: protocol.KindResponse,
661 + Code: rHdr.Code,
662 + ItemCount: 1,
663 + MessageID: rHdr.MessageID,
664 + }
665 + if err := server.Send(&resp, rPayload); err != nil {
666 + t.Errorf("server Send[%d]: %v", i, err)
667 + return
668 + }
669 + }
670 + }()
671 +
672 + for i := 0; i < numMessages; i++ {
673 + payload := []byte(fmt.Sprintf("message_%d", i))
674 + hdr := protocol.Header{
675 + Kind: protocol.KindRequest,
676 + Code: protocol.MethodIncrement,
677 + ItemCount: 1,
678 + MessageID: uint64(i),
679 + }
680 + if err := client.Send(&hdr, payload); err != nil {
681 + t.Fatalf("client Send[%d]: %v", i, err)
682 + }
683 + }
684 +
685 + received := make(map[uint64]bool)
686 + buf := make([]byte, 65600)
687 + for i := 0; i < numMessages; i++ {
688 + rHdr, _, err := client.Receive(buf)
689 + if err != nil {
690 + t.Fatalf("client Receive[%d]: %v", i, err)
691 + }
692 + received[rHdr.MessageID] = true
693 + }
694 +
695 + wg.Wait()
696 +
697 + for i := 0; i < numMessages; i++ {
698 + if !received[uint64(i)] {
699 + t.Fatalf("missing response for message_id %d", i)
700 + }
701 + }
702 +}
703 +
704 +func TestPipeClosedSessionErrors(t *testing.T) {
705 + service := uniquePipeService(t)
706 + listener := startListener(t, testPipeRunDir, service, defaultServerConfig())
707 + defer listener.Close()
708 +
709 + acceptCh := acceptAsync(listener)
710 + client, err := Connect(testPipeRunDir, service, defaultClientConfigPtr())
711 + if err != nil {
712 + t.Fatalf("Connect failed: %v", err)
713 + }
714 +
715 + sr := <-acceptCh
716 + if sr.err != nil {
717 + t.Fatalf("Accept failed: %v", sr.err)
718 + }
719 + sr.session.Close()
720 + client.Close()
721 +
722 + hdr := protocol.Header{
723 + Kind: protocol.KindRequest,
724 + Code: protocol.MethodIncrement,
725 + ItemCount: 1,
726 + MessageID: 1,
727 + }
728 + if err := client.Send(&hdr, []byte("x")); err == nil {
729 + t.Fatal("Send on closed session should fail")
730 + }
731 +
732 + buf := make([]byte, 4096)
733 + if _, _, err := client.Receive(buf); err == nil {
734 + t.Fatal("Receive on closed session should fail")
735 + }
736 +}
737 +
738 +func TestPipeListenerClose(t *testing.T) {
739 + service := uniquePipeService(t)
740 + listener := startListener(t, testPipeRunDir, service, defaultServerConfig())
741 + listener.Close()
742 +
743 + if _, err := Connect(testPipeRunDir, service, defaultClientConfigPtr()); err == nil {
744 + t.Fatal("Connect after listener close should fail")
745 + }
746 +}
747 +
748 +func TestPipeDefaultsApplied(t *testing.T) {
749 + service := uniquePipeService(t)
750 +
751 + sCfg := ServerConfig{AuthToken: testAuthToken}
752 + listener := startListener(t, testPipeRunDir, service, sCfg)
753 + defer listener.Close()
754 +
755 + acceptCh := acceptAsync(listener)
756 +
757 + cCfg := ClientConfig{AuthToken: testAuthToken}
758 + client, err := Connect(testPipeRunDir, service, &cCfg)
759 + if err != nil {
760 + t.Fatalf("Connect failed: %v", err)
761 + }
762 + defer client.Close()
763 +
764 + sr := <-acceptCh
765 + if sr.err != nil {
766 + t.Fatalf("Accept failed: %v", sr.err)
767 + }
768 + defer sr.session.Close()
769 +
770 + if client.MaxRequestPayloadBytes != protocol.MaxPayloadDefault || client.MaxRequestBatchItems != 1 || client.MaxResponsePayloadBytes != protocol.MaxPayloadDefault || client.MaxResponseBatchItems != 1 {
771 + t.Fatalf("unexpected defaults: %+v", client)
772 + }
773 + if client.PacketSize != defaultPacketSize {
774 + t.Fatalf("packet_size=%d, want %d", client.PacketSize, defaultPacketSize)
775 + }
776 + if client.SelectedProfile != protocol.ProfileBaseline {
777 + t.Fatalf("selected profile = 0x%x, want baseline", client.SelectedProfile)
778 + }
779 +}
780 +
781 +func TestPipeHighestBit(t *testing.T) {
782 + tests := []struct {
783 + in uint32
784 + want uint32
785 + }{
786 + {0, 0},
787 + {1, 1},
788 + {0x03, 0x02},
789 + {0x07, 0x04},
790 + {0x80000000, 0x80000000},
791 + {0xFF, 0x80},
792 + }
793 + for _, tc := range tests {
794 + if got := highestBit(tc.in); got != tc.want {
795 + t.Fatalf("highestBit(0x%x) = 0x%x, want 0x%x", tc.in, got, tc.want)
796 + }
797 + }
798 +}
799 +
800 +func TestPipeMultipleChunkedMessages(t *testing.T) {
801 + service := uniquePipeService(t)
802 + const forcedPacketSize = 96
803 +
804 + sCfg := defaultServerConfig()
805 + sCfg.PacketSize = forcedPacketSize
806 + sCfg.MaxRequestPayloadBytes = 65536
807 + sCfg.MaxResponsePayloadBytes = 65536
808 + listener := startListener(t, testPipeRunDir, service, sCfg)
809 + defer listener.Close()
810 +
811 + acceptCh := acceptAsync(listener)
812 +
813 + cCfg := defaultClientConfig()
814 + cCfg.PacketSize = forcedPacketSize
815 + cCfg.MaxRequestPayloadBytes = 65536
816 + cCfg.MaxResponsePayloadBytes = 65536
817 + client, err := Connect(testPipeRunDir, service, &cCfg)
818 + if err != nil {
819 + t.Fatalf("Connect failed: %v", err)
820 + }
821 + defer client.Close()
822 +
823 + sr := <-acceptCh
824 + if sr.err != nil {
825 + t.Fatalf("Accept failed: %v", sr.err)
826 + }
827 + server := sr.session
828 + defer server.Close()
829 +
830 + type receiveResult struct {
831 + hdr protocol.Header
832 + payload []byte
833 + err error
834 + }
835 + serverRecvCh := make(chan receiveResult, 3)
836 + go func() {
837 + buf := make([]byte, forcedPacketSize)
838 + for i := 0; i < 3; i++ {
839 + rHdr, rPayload, err := server.Receive(buf)
840 + if err == nil {
841 + rPayload = append([]byte(nil), rPayload...)
842 + }
843 + serverRecvCh <- receiveResult{hdr: rHdr, payload: rPayload, err: err}
844 + if err != nil {
845 + return
846 + }
847 + }
848 + }()
849 +
850 + for i := 0; i < 3; i++ {
851 + size := 500 + i*200
852 + payload := make([]byte, size)
853 + for j := range payload {
854 + payload[j] = byte((i*31 + j) & 0xFF)
855 + }
856 +
857 + hdr := protocol.Header{
858 + Kind: protocol.KindRequest,
859 + Code: protocol.MethodIncrement,
860 + ItemCount: 1,
861 + MessageID: uint64(i + 1),
862 + }
863 + if err := client.Send(&hdr, payload); err != nil {
864 + t.Fatalf("Send[%d]: %v", i, err)
865 + }
866 +
867 + serverRecv := <-serverRecvCh
868 + if serverRecv.err != nil {
869 + t.Fatalf("Receive[%d]: %v", i, serverRecv.err)
870 + }
871 + if serverRecv.hdr.MessageID != uint64(i+1) || !bytes.Equal(serverRecv.payload, payload) {
872 + t.Fatalf("unexpected chunked message[%d]: hdr=%+v len=%d", i, serverRecv.hdr, len(serverRecv.payload))
873 + }
874 + }
875 +}
876 +
877 +func TestPipeConnectNoServer(t *testing.T) {
878 + service := uniquePipeService(t)
879 + if _, err := Connect(testPipeRunDir, service, defaultClientConfigPtr()); err == nil {
880 + t.Fatal("expected error connecting to missing server")
881 + }
882 +}
883 +
884 +func TestPipeAcceptWithDelayedConnect(t *testing.T) {
885 + service := uniquePipeService(t)
886 + listener := startListener(t, testPipeRunDir, service, defaultServerConfig())
887 + defer listener.Close()
888 +
889 + done := make(chan bool, 1)
890 + go func() {
891 + acceptCh := acceptAsync(listener)
892 + select {
893 + case sr := <-acceptCh:
894 + if sr.err != nil {
895 + done <- false
896 + return
897 + }
898 + sr.session.Close()
899 + done <- true
900 + case <-time.After(5 * time.Second):
901 + done <- false
902 + }
903 + }()
904 +
905 + time.Sleep(50 * time.Millisecond)
906 + client, err := Connect(testPipeRunDir, service, defaultClientConfigPtr())
907 + if err != nil {
908 + t.Fatalf("Connect failed: %v", err)
909 + }
910 + client.Close()
911 +
912 + if ok := <-done; !ok {
913 + t.Fatal("accept did not complete successfully")
914 + }
915 +}
916 +
917 +func TestPipeDuplicateMessageID(t *testing.T) {
918 + service := uniquePipeService(t)
919 + listener := startListener(t, testPipeRunDir, service, defaultServerConfig())
920 + defer listener.Close()
921 +
922 + acceptCh := acceptAsync(listener)
923 + client, err := Connect(testPipeRunDir, service, defaultClientConfigPtr())
924 + if err != nil {
925 + t.Fatalf("Connect failed: %v", err)
926 + }
927 + defer client.Close()
928 +
929 + sr := <-acceptCh
930 + if sr.err != nil {
931 + t.Fatalf("Accept failed: %v", sr.err)
932 + }
933 + defer sr.session.Close()
934 +
935 + hdr := protocol.Header{
936 + Kind: protocol.KindRequest,
937 + Code: protocol.MethodIncrement,
938 + ItemCount: 1,
939 + MessageID: 42,
940 + }
941 + if err := client.Send(&hdr, []byte("hello")); err != nil {
942 + t.Fatalf("first Send: %v", err)
943 + }
944 +
945 + if err := client.Send(&hdr, []byte("hello")); !errors.Is(err, ErrDuplicateMsgID) {
946 + t.Fatalf("expected ErrDuplicateMsgID, got %v", err)
947 + }
948 +}
949 +
950 +func TestPipeUnknownMessageID(t *testing.T) {
951 + service := uniquePipeService(t)
952 + listener := startListener(t, testPipeRunDir, service, defaultServerConfig())
953 + defer listener.Close()
954 +
955 + acceptCh := acceptAsync(listener)
956 + client, err := Connect(testPipeRunDir, service, defaultClientConfigPtr())
957 + if err != nil {
958 + t.Fatalf("Connect failed: %v", err)
959 + }
960 + defer client.Close()
961 +
962 + sr := <-acceptCh
963 + if sr.err != nil {
964 + t.Fatalf("Accept failed: %v", sr.err)
965 + }
966 + server := sr.session
967 + defer server.Close()
968 +
969 + hdr := protocol.Header{
970 + Kind: protocol.KindRequest,
971 + Code: protocol.MethodIncrement,
972 + ItemCount: 1,
973 + MessageID: 10,
974 + }
975 + if err := client.Send(&hdr, []byte("x")); err != nil {
976 + t.Fatalf("Send failed: %v", err)
977 + }
978 +
979 + buf := make([]byte, 4096)
980 + rHdr, rPayload, err := server.Receive(buf)
981 + if err != nil {
982 + t.Fatalf("server Receive failed: %v", err)
983 + }
984 +
985 + resp := protocol.Header{
986 + Kind: protocol.KindResponse,
987 + Code: rHdr.Code,
988 + ItemCount: 1,
989 + MessageID: 999,
990 + }
991 + if err := server.Send(&resp, rPayload); err != nil {
992 + t.Fatalf("server Send failed: %v", err)
993 + }
994 +
995 + if _, _, err := client.Receive(buf); !errors.Is(err, ErrUnknownMsgID) {
996 + t.Fatalf("expected ErrUnknownMsgID, got %v", err)
997 + }
998 +}
999 +
1000 +func TestPipeHandleAndRole(t *testing.T) {
1001 + service := uniquePipeService(t)
1002 + listener := startListener(t, testPipeRunDir, service, defaultServerConfig())
1003 + defer listener.Close()
1004 +
1005 + acceptCh := acceptAsync(listener)
1006 + client, err := Connect(testPipeRunDir, service, defaultClientConfigPtr())
1007 + if err != nil {
1008 + t.Fatalf("Connect failed: %v", err)
1009 + }
1010 + defer client.Close()
1011 +
1012 + sr := <-acceptCh
1013 + if sr.err != nil {
1014 + t.Fatalf("Accept failed: %v", sr.err)
1015 + }
1016 + server := sr.session
1017 + defer server.Close()
1018 +
1019 + if listener.Handle() == syscall.InvalidHandle {
1020 + t.Fatal("listener handle should be valid")
1021 + }
1022 + if client.Handle() == syscall.InvalidHandle || server.Handle() == syscall.InvalidHandle {
1023 + t.Fatal("session handles should be valid")
1024 + }
1025 + if client.GetRole() != RoleClient || server.GetRole() != RoleServer {
1026 + t.Fatalf("unexpected roles client=%d server=%d", client.GetRole(), server.GetRole())
1027 + }
1028 +}
1029 +
1030 +func TestPipePipelineChunked(t *testing.T) {
1031 + service := uniquePipeService(t)
1032 + const forcedPacketSize = 128
1033 +
1034 + sCfg := defaultServerConfig()
1035 + sCfg.PacketSize = forcedPacketSize
1036 + sCfg.MaxRequestPayloadBytes = 65536
1037 + sCfg.MaxResponsePayloadBytes = 65536
1038 + listener := startListener(t, testPipeRunDir, service, sCfg)
1039 + defer listener.Close()
1040 +
1041 + acceptCh := acceptAsync(listener)
1042 +
1043 + cCfg := defaultClientConfig()
1044 + cCfg.PacketSize = forcedPacketSize
1045 + cCfg.MaxRequestPayloadBytes = 65536
1046 + cCfg.MaxResponsePayloadBytes = 65536
1047 + client, err := Connect(testPipeRunDir, service, &cCfg)
1048 + if err != nil {
1049 + t.Fatalf("Connect failed: %v", err)
1050 + }
1051 + defer client.Close()
1052 +
1053 + sr := <-acceptCh
1054 + if sr.err != nil {
1055 + t.Fatalf("Accept failed: %v", sr.err)
1056 + }
1057 + server := sr.session
1058 + defer server.Close()
1059 +
1060 + sizes := []int{200, 500, 300, 800, 150}
1061 + count := len(sizes)
1062 +
1063 + var wg sync.WaitGroup
1064 + wg.Add(1)
1065 + go func() {
1066 + defer wg.Done()
1067 + buf := make([]byte, forcedPacketSize)
1068 + for i := 0; i < count; i++ {
1069 + rHdr, rPayload, err := server.Receive(buf)
1070 + if err != nil {
1071 + t.Errorf("server Receive[%d]: %v", i, err)
1072 + return
1073 + }
1074 + resp := protocol.Header{
1075 + Kind: protocol.KindResponse,
1076 + Code: rHdr.Code,
1077 + ItemCount: 1,
1078 + MessageID: rHdr.MessageID,
1079 + }
1080 + if err := server.Send(&resp, rPayload); err != nil {
1081 + t.Errorf("server Send[%d]: %v", i, err)
1082 + return
1083 + }
1084 + }
1085 + }()
1086 +
1087 + for i, sz := range sizes {
1088 + payload := make([]byte, sz)
1089 + for j := range payload {
1090 + payload[j] = byte((i + j) & 0xFF)
1091 + }
1092 + hdr := protocol.Header{
1093 + Kind: protocol.KindRequest,
1094 + Code: protocol.MethodIncrement,
1095 + ItemCount: 1,
1096 + MessageID: uint64(i + 1),
1097 + }
1098 + if err := client.Send(&hdr, payload); err != nil {
1099 + t.Fatalf("client Send[%d]: %v", i, err)
1100 + }
1101 + }
1102 +
1103 + buf := make([]byte, forcedPacketSize)
1104 + for i, sz := range sizes {
1105 + rHdr, rPayload, err := client.Receive(buf)
1106 + if err != nil {
1107 + t.Fatalf("client Receive[%d]: %v", i, err)
1108 + }
1109 + if rHdr.MessageID != uint64(i+1) {
1110 + t.Errorf("[%d] message_id = %d, want %d", i, rHdr.MessageID, i+1)
1111 + }
1112 + if len(rPayload) != sz {
1113 + t.Errorf("[%d] payload len = %d, want %d", i, len(rPayload), sz)
1114 + continue
1115 + }
1116 + expected := make([]byte, sz)
1117 + for j := range expected {
1118 + expected[j] = byte((i + j) & 0xFF)
1119 + }
1120 + if !bytes.Equal(rPayload, expected) {
1121 + t.Errorf("[%d] chunked payload data mismatch", i)
1122 + }
1123 + }
1124 +
1125 + wg.Wait()
1126 +}
1127 +
1128 +func TestPipeBatchRoundTrip(t *testing.T) {
1129 + service := uniquePipeService(t)
1130 + listener := startListener(t, testPipeRunDir, service, defaultServerConfig())
1131 + defer listener.Close()
1132 +
1133 + acceptCh := acceptAsync(listener)
1134 + client, err := Connect(testPipeRunDir, service, defaultClientConfigPtr())
1135 + if err != nil {
1136 + t.Fatalf("Connect failed: %v", err)
1137 + }
1138 + defer client.Close()
1139 +
1140 + sr := <-acceptCh
1141 + if sr.err != nil {
1142 + t.Fatalf("Accept failed: %v", sr.err)
1143 + }
1144 + server := sr.session
1145 + defer server.Close()
1146 +
1147 + payloadBuf := make([]byte, 256)
1148 + builder := protocol.NewBatchBuilder(payloadBuf, 3)
1149 + for _, v := range []uint64{1, 41, 99} {
1150 + var item [protocol.IncrementPayloadSize]byte
1151 + if protocol.IncrementEncode(v, item[:]) == 0 {
1152 + t.Fatal("increment encode failed")
1153 + }
1154 + if err := builder.Add(item[:]); err != nil {
1155 + t.Fatalf("batch add failed: %v", err)
1156 + }
1157 + }
1158 + totalLen, _ := builder.Finish()
1159 + reqPayload := payloadBuf[:totalLen]
1160 +
1161 + hdr := protocol.Header{
1162 + Kind: protocol.KindRequest,
1163 + Code: protocol.MethodIncrement,
1164 + Flags: protocol.FlagBatch,
1165 + ItemCount: 3,
1166 + MessageID: 7,
1167 + }
1168 + if err := client.Send(&hdr, reqPayload); err != nil {
1169 + t.Fatalf("client Send batch failed: %v", err)
1170 + }
1171 +
1172 + buf := make([]byte, 4096)
1173 + rHdr, rPayload, err := server.Receive(buf)
1174 + if err != nil {
1175 + t.Fatalf("server Receive batch failed: %v", err)
1176 + }
1177 + if rHdr.Flags&protocol.FlagBatch == 0 || rHdr.ItemCount != 3 {
1178 + t.Fatalf("unexpected batch header: %+v", rHdr)
1179 + }
1180 + for i, want := range []uint64{1, 41, 99} {
1181 + item, err := protocol.BatchItemGet(rPayload, 3, uint32(i))
1182 + if err != nil {
1183 + t.Fatalf("BatchItemGet[%d] failed: %v", i, err)
1184 + }
1185 + got, err := protocol.IncrementDecode(item)
1186 + if err != nil {
1187 + t.Fatalf("IncrementDecode[%d] failed: %v", i, err)
1188 + }
1189 + if got != want {
1190 + t.Fatalf("item[%d]=%d, want %d", i, got, want)
1191 + }
1192 + }
1193 +}
1194 +
1195 +func TestPipeLargePayloadsMixedSizes(t *testing.T) {
1196 + service := uniquePipeService(t)
1197 +
1198 + sCfg := defaultServerConfig()
1199 + sCfg.MaxRequestPayloadBytes = 65536
1200 + sCfg.MaxResponsePayloadBytes = 65536
1201 + listener := startListener(t, testPipeRunDir, service, sCfg)
1202 + defer listener.Close()
1203 +
1204 + acceptCh := acceptAsync(listener)
1205 +
1206 + cCfg := defaultClientConfig()
1207 + cCfg.MaxRequestPayloadBytes = 65536
1208 + cCfg.MaxResponsePayloadBytes = 65536
1209 + client, err := Connect(testPipeRunDir, service, &cCfg)
1210 + if err != nil {
1211 + t.Fatalf("Connect failed: %v", err)
1212 + }
1213 + defer client.Close()
1214 +
1215 + sr := <-acceptCh
1216 + if sr.err != nil {
1217 + t.Fatalf("Accept failed: %v", sr.err)
1218 + }
1219 + server := sr.session
1220 + defer server.Close()
1221 +
1222 + sizes := []int{8, 256, 1024, 8, 256, 1024}
1223 + var wg sync.WaitGroup
1224 + wg.Add(1)
1225 + go func() {
1226 + defer wg.Done()
1227 + buf := make([]byte, 8192)
1228 + for i := 0; i < len(sizes); i++ {
1229 + rHdr, rPayload, err := server.Receive(buf)
1230 + if err != nil {
1231 + t.Errorf("server Receive[%d]: %v", i, err)
1232 + return
1233 + }
1234 + resp := protocol.Header{
1235 + Kind: protocol.KindResponse,
1236 + Code: rHdr.Code,
1237 + ItemCount: 1,
1238 + MessageID: rHdr.MessageID,
1239 + }
1240 + if err := server.Send(&resp, rPayload); err != nil {
1241 + t.Errorf("server Send[%d]: %v", i, err)
1242 + return
1243 + }
1244 + }
1245 + }()
1246 +
1247 + for i, sz := range sizes {
1248 + payload := make([]byte, sz)
1249 + for j := range payload {
1250 + payload[j] = byte((i*37 + j) & 0xFF)
1251 + }
1252 + hdr := protocol.Header{
1253 + Kind: protocol.KindRequest,
1254 + Code: protocol.MethodIncrement,
1255 + ItemCount: 1,
1256 + MessageID: uint64(i + 1),
1257 + }
1258 + if err := client.Send(&hdr, payload); err != nil {
1259 + t.Fatalf("client Send[%d]: %v", i, err)
1260 + }
1261 + }
1262 +
1263 + buf := make([]byte, 8192)
1264 + for i, sz := range sizes {
1265 + rHdr, rPayload, err := client.Receive(buf)
1266 + if err != nil {
1267 + t.Fatalf("client Receive[%d]: %v", i, err)
1268 + }
1269 + if rHdr.MessageID != uint64(i+1) || len(rPayload) != sz {
1270 + t.Fatalf("unexpected mixed-size response[%d]: hdr=%+v len=%d", i, rHdr, len(rPayload))
1271 + }
1272 + expected := make([]byte, sz)
1273 + for j := range expected {
1274 + expected[j] = byte((i*37 + j) & 0xFF)
1275 + }
1276 + if !bytes.Equal(rPayload, expected) {
1277 + t.Fatalf("payload mismatch for mixed-size response[%d]", i)
1278 + }
1279 + }
1280 +
1281 + wg.Wait()
1282 +}
1283 +
1284 +func TestPipeBinaryPayloadRoundTrip(t *testing.T) {
1285 + service := uniquePipeService(t)
1286 + listener := startListener(t, testPipeRunDir, service, defaultServerConfig())
1287 + defer listener.Close()
1288 +
1289 + acceptCh := acceptAsync(listener)
1290 + client, err := Connect(testPipeRunDir, service, defaultClientConfigPtr())
1291 + if err != nil {
1292 + t.Fatalf("Connect failed: %v", err)
1293 + }
1294 + defer client.Close()
1295 +
1296 + sr := <-acceptCh
1297 + if sr.err != nil {
1298 + t.Fatalf("Accept failed: %v", sr.err)
1299 + }
1300 + server := sr.session
1301 + defer server.Close()
1302 +
1303 + payload := make([]byte, 8)
1304 + binary.NativeEndian.PutUint64(payload, 0x1122334455667788)
1305 + hdr := protocol.Header{
1306 + Kind: protocol.KindRequest,
1307 + Code: protocol.MethodIncrement,
1308 + ItemCount: 1,
1309 + MessageID: 55,
1310 + }
1311 + if err := client.Send(&hdr, payload); err != nil {
1312 + t.Fatalf("client Send failed: %v", err)
1313 + }
1314 +
1315 + buf := make([]byte, 4096)
1316 + rHdr, rPayload, err := server.Receive(buf)
1317 + if err != nil {
1318 + t.Fatalf("server Receive failed: %v", err)
1319 + }
1320 + if rHdr.MessageID != 55 {
1321 + t.Fatalf("unexpected message_id %d", rHdr.MessageID)
1322 + }
1323 + if got := binary.NativeEndian.Uint64(rPayload); got != 0x1122334455667788 {
1324 + t.Fatalf("unexpected payload value 0x%x", got)
1325 + }
1326 +}
src/go/pkg/netipc/transport/windows/pipe_test.go new
+220
@@ -0,0 +1,220 @@
1 +//go:build windows
2 +
3 +package windows
4 +
5 +import (
6 + "errors"
7 + "strings"
8 + "syscall"
9 + "testing"
10 +)
11 +
12 +func TestApplyDefault(t *testing.T) {
13 + if got := applyDefault(0, 42); got != 42 {
14 + t.Fatalf("applyDefault(0, 42) = %d, want 42", got)
15 + }
16 + if got := applyDefault(7, 42); got != 7 {
17 + t.Fatalf("applyDefault(7, 42) = %d, want 7", got)
18 + }
19 +}
20 +
21 +func TestMinU32(t *testing.T) {
22 + if got := minU32(1, 9); got != 1 {
23 + t.Fatalf("minU32(1, 9) = %d, want 1", got)
24 + }
25 + if got := minU32(9, 1); got != 1 {
26 + t.Fatalf("minU32(9, 1) = %d, want 1", got)
27 + }
28 + if got := minU32(5, 5); got != 5 {
29 + t.Fatalf("minU32(5, 5) = %d, want 5", got)
30 + }
31 +}
32 +
33 +func TestHighestBit(t *testing.T) {
34 + cases := []struct {
35 + mask uint32
36 + want uint32
37 + }{
38 + {0, 0},
39 + {1, 1},
40 + {2, 2},
41 + {3, 2},
42 + {0x10, 0x10},
43 + {0x101, 0x100},
44 + {0x80000001, 0x80000000},
45 + }
46 +
47 + for _, tc := range cases {
48 + if got := highestBit(tc.mask); got != tc.want {
49 + t.Fatalf("highestBit(0x%x) = 0x%x, want 0x%x", tc.mask, got, tc.want)
50 + }
51 + }
52 +}
53 +
54 +func TestIsDisconnectError(t *testing.T) {
55 + for _, errno := range []syscall.Errno{
56 + _ERROR_BROKEN_PIPE,
57 + _ERROR_NO_DATA,
58 + _ERROR_PIPE_NOT_CONNECTED,
59 + } {
60 + if !isDisconnectError(errno) {
61 + t.Fatalf("isDisconnectError(%v) = false, want true", errno)
62 + }
63 + }
64 +
65 + for _, err := range []error{
66 + nil,
67 + errors.New("plain"),
68 + syscall.Errno(_ERROR_ACCESS_DENIED),
69 + } {
70 + if isDisconnectError(err) {
71 + t.Fatalf("isDisconnectError(%v) = true, want false", err)
72 + }
73 + }
74 +}
75 +
76 +func TestFNV1a64Empty(t *testing.T) {
77 + got := FNV1a64([]byte{})
78 + if got != fnv1aOffsetBasis {
79 + t.Errorf("FNV1a64(empty) = %x, want %x", got, fnv1aOffsetBasis)
80 + }
81 +}
82 +
83 +func TestFNV1a64Deterministic(t *testing.T) {
84 + data := []byte("/var/run/netdata")
85 + h1 := FNV1a64(data)
86 + h2 := FNV1a64(data)
87 + if h1 != h2 {
88 + t.Errorf("FNV1a64 not deterministic: %x != %x", h1, h2)
89 + }
90 +}
91 +
92 +func TestFNV1a64DifferentInputs(t *testing.T) {
93 + h1 := FNV1a64([]byte("/var/run/netdata"))
94 + h2 := FNV1a64([]byte("/tmp/netdata"))
95 + if h1 == h2 {
96 + t.Error("FNV1a64 should produce different hashes for different inputs")
97 + }
98 +}
99 +
100 +func TestValidateServiceName(t *testing.T) {
101 + valid := []string{
102 + "cgroups-snapshot",
103 + "test_service.v1",
104 + "A-Z_09",
105 + "a",
106 + "abc123",
107 + }
108 + for _, name := range valid {
109 + if err := validateServiceName(name); err != nil {
110 + t.Errorf("validateServiceName(%q) = %v, want nil", name, err)
111 + }
112 + }
113 +
114 + invalid := []string{
115 + "",
116 + ".",
117 + "..",
118 + "has space",
119 + "has/slash",
120 + "has\\backslash",
121 + "has@at",
122 + "has:colon",
123 + }
124 + for _, name := range invalid {
125 + if err := validateServiceName(name); err == nil {
126 + t.Errorf("validateServiceName(%q) = nil, want error", name)
127 + }
128 + }
129 +}
130 +
131 +func TestBuildPipeName(t *testing.T) {
132 + name, err := BuildPipeName("/var/run/netdata", "cgroups-snapshot")
133 + if err != nil {
134 + t.Fatalf("BuildPipeName failed: %v", err)
135 + }
136 +
137 + // Should end with NUL
138 + if name[len(name)-1] != 0 {
139 + t.Error("pipe name should end with NUL")
140 + }
141 +
142 + // Convert to narrow string for checking
143 + narrow := make([]byte, len(name)-1)
144 + for i, c := range name[:len(name)-1] {
145 + narrow[i] = byte(c)
146 + }
147 + s := string(narrow)
148 +
149 + if s[:len(`\\.\pipe\netipc-`)] != `\\.\pipe\netipc-` {
150 + t.Errorf("unexpected prefix: %s", s)
151 + }
152 +
153 + // Should end with service name
154 + suffix := "-cgroups-snapshot"
155 + if s[len(s)-len(suffix):] != suffix {
156 + t.Errorf("unexpected suffix in: %s", s)
157 + }
158 +}
159 +
160 +func TestBuildPipeNameDeterministic(t *testing.T) {
161 + n1, _ := BuildPipeName("/var/run", "svc")
162 + n2, _ := BuildPipeName("/var/run", "svc")
163 + if len(n1) != len(n2) {
164 + t.Fatal("different lengths")
165 + }
166 + for i := range n1 {
167 + if n1[i] != n2[i] {
168 + t.Fatalf("mismatch at %d", i)
169 + }
170 + }
171 +}
172 +
173 +func TestBuildPipeNameDifferentRunDir(t *testing.T) {
174 + n1, _ := BuildPipeName("/var/run/netdata", "svc")
175 + n2, _ := BuildPipeName("/tmp/netdata", "svc")
176 + // They should differ (different hash)
177 + if len(n1) == len(n2) {
178 + same := true
179 + for i := range n1 {
180 + if n1[i] != n2[i] {
181 + same = false
182 + break
183 + }
184 + }
185 + if same {
186 + t.Error("different run_dir should produce different pipe names")
187 + }
188 + }
189 +}
190 +
191 +func TestBuildPipeNameInvalidService(t *testing.T) {
192 + _, err := BuildPipeName("/var/run", "")
193 + if err == nil {
194 + t.Error("expected error for empty service name")
195 + }
196 +
197 + _, err = BuildPipeName("/var/run", "bad/name")
198 + if err == nil {
199 + t.Error("expected error for service with /")
200 + }
201 +
202 + _, err = BuildPipeName("/var/run", ".")
203 + if err == nil {
204 + t.Error("expected error for '.'")
205 + }
206 +}
207 +
208 +func TestBuildPipeNameTooLong(t *testing.T) {
209 + service := strings.Repeat("a", maxPipeNameChars)
210 + if _, err := BuildPipeName("/var/run", service); !errors.Is(err, ErrPipeName) {
211 + t.Fatalf("BuildPipeName(long) = %v, want ErrPipeName", err)
212 + }
213 +}
214 +
215 +func TestSetNamedPipeHandleStateInvalidHandle(t *testing.T) {
216 + mode := uint32(_PIPE_READMODE_MESSAGE)
217 + if err := setNamedPipeHandleState(syscall.InvalidHandle, &mode); err == nil {
218 + t.Fatal("setNamedPipeHandleState on invalid handle should fail")
219 + }
220 +}
src/go/pkg/netipc/transport/windows/shm.go new
+784
@@ -0,0 +1,784 @@
1 +//go:build windows
2 +
3 +// Windows SHM transport — shared memory data plane with spin + kernel event
4 +// synchronization. Wire-compatible with the C and Rust implementations.
5 +//
6 +// Pure Go — no cgo. Works with CGO_ENABLED=0.
7 +
8 +package windows
9 +
10 +import (
11 + "encoding/binary"
12 + "errors"
13 + "fmt"
14 + "sync/atomic"
15 + "syscall"
16 + "unsafe"
17 +)
18 +
19 +// ---------------------------------------------------------------------------
20 +// Constants
21 +// ---------------------------------------------------------------------------
22 +
23 +const (
24 + winShmMagic uint32 = 0x4e535748 // "NSWH"
25 + winShmVersion uint32 = 3
26 + winShmHeaderLen uint32 = 128
27 + winShmCachelineSize uint32 = 64
28 + winShmDefaultSpin uint32 = 1024
29 +
30 + WinShmProfileHybrid uint32 = 0x02
31 + WinShmProfileBusywait uint32 = 0x04
32 +
33 + // Header field offsets
34 + wshOFFMagic = 0
35 + wshOFFVersion = 4
36 + wshOFFHeaderLen = 8
37 + wshOFFProfile = 12
38 + wshOFFReqOffset = 16
39 + wshOFFReqCapacity = 20
40 + wshOFFRespOffset = 24
41 + wshOFFRespCapacity = 28
42 + wshOFFSpinTries = 32
43 + wshOFFReqLen = 36
44 + wshOFFRespLen = 40
45 + wshOFFReqClientClosed = 44
46 + wshOFFReqServerWaiting = 48
47 + wshOFFRespServerClosed = 52
48 + wshOFFRespClientWaiting = 56
49 + wshOFFReqSeq = 64
50 + wshOFFRespSeq = 72
51 +)
52 +
53 +// ---------------------------------------------------------------------------
54 +// Errors
55 +// ---------------------------------------------------------------------------
56 +
57 +var (
58 + ErrWinShmBadParam = errors.New("invalid Windows SHM argument")
59 + ErrWinShmCreateMapping = errors.New("CreateFileMappingW failed")
60 + ErrWinShmOpenMapping = errors.New("OpenFileMappingW failed")
61 + ErrWinShmMapView = errors.New("MapViewOfFile failed")
62 + ErrWinShmCreateEvent = errors.New("CreateEventW failed")
63 + ErrWinShmOpenEvent = errors.New("OpenEventW failed")
64 + ErrWinShmAddrInUse = errors.New("Windows SHM object name already in use by live server")
65 + ErrWinShmBadMagic = errors.New("Windows SHM header magic mismatch")
66 + ErrWinShmBadVersion = errors.New("Windows SHM header version mismatch")
67 + ErrWinShmBadHeader = errors.New("Windows SHM header_len mismatch")
68 + ErrWinShmBadProfile = errors.New("Windows SHM profile mismatch")
69 + ErrWinShmMsgTooLarge = errors.New("message exceeds Windows SHM area capacity")
70 + ErrWinShmTimeout = errors.New("Windows SHM wait timed out")
71 + ErrWinShmDisconnected = errors.New("Windows SHM peer closed")
72 +)
73 +
74 +// ---------------------------------------------------------------------------
75 +// Win32 syscall imports
76 +// ---------------------------------------------------------------------------
77 +
78 +var (
79 + procCreateFileMappingW = modkernel32.NewProc("CreateFileMappingW")
80 + procOpenFileMappingW = modkernel32.NewProc("OpenFileMappingW")
81 + procMapViewOfFile = modkernel32.NewProc("MapViewOfFile")
82 + procUnmapViewOfFile = modkernel32.NewProc("UnmapViewOfFile")
83 + procCreateEventW = modkernel32.NewProc("CreateEventW")
84 + procOpenEventW = modkernel32.NewProc("OpenEventW")
85 + procSetEvent = modkernel32.NewProc("SetEvent")
86 + procWaitForSingleObj = modkernel32.NewProc("WaitForSingleObject")
87 + procGetTickCount64 = modkernel32.NewProc("GetTickCount64")
88 +)
89 +
90 +type winShmProcCall func(a ...uintptr) (uintptr, uintptr, error)
91 +
92 +func callCreateFileMappingW(a ...uintptr) (uintptr, uintptr, error) {
93 + return procCreateFileMappingW.Call(a...)
94 +}
95 +func callOpenFileMappingW(a ...uintptr) (uintptr, uintptr, error) {
96 + return procOpenFileMappingW.Call(a...)
97 +}
98 +func callMapViewOfFile(a ...uintptr) (uintptr, uintptr, error) { return procMapViewOfFile.Call(a...) }
99 +func callCreateEventW(a ...uintptr) (uintptr, uintptr, error) { return procCreateEventW.Call(a...) }
100 +func callOpenEventW(a ...uintptr) (uintptr, uintptr, error) { return procOpenEventW.Call(a...) }
101 +
102 +var (
103 + winShmCreateFileMappingW winShmProcCall = callCreateFileMappingW
104 + winShmOpenFileMappingW winShmProcCall = callOpenFileMappingW
105 + winShmMapViewOfFile winShmProcCall = callMapViewOfFile
106 + winShmCreateEventW winShmProcCall = callCreateEventW
107 + winShmOpenEventW winShmProcCall = callOpenEventW
108 +)
109 +
110 +const (
111 + _PAGE_READWRITE = 0x04
112 + _FILE_MAP_ALL_ACCESS = 0x000F001F
113 + _EVENT_MODIFY_STATE = 0x0002
114 + _SYNCHRONIZE = 0x00100000
115 + _INFINITE = 0xFFFFFFFF
116 + _WAIT_TIMEOUT = 0x00000102
117 + _ERROR_ALREADY_EXISTS = 183
118 +)
119 +
120 +func isWindowsErrno(err error, want syscall.Errno) bool {
121 + errno, ok := err.(syscall.Errno)
122 + return ok && errno == want
123 +}
124 +
125 +// ---------------------------------------------------------------------------
126 +// Role
127 +// ---------------------------------------------------------------------------
128 +
129 +// WinShmRole distinguishes server vs client.
130 +type WinShmRole int
131 +
132 +const (
133 + WinShmRoleServer WinShmRole = 1
134 + WinShmRoleClient WinShmRole = 2
135 +)
136 +
137 +// ---------------------------------------------------------------------------
138 +// Context
139 +// ---------------------------------------------------------------------------
140 +
141 +// WinShmContext is a handle to a Windows SHM region.
142 +type WinShmContext struct {
143 + role WinShmRole
144 + mapping syscall.Handle
145 + base uintptr
146 + size uintptr
147 +
148 + reqEvent syscall.Handle
149 + respEvent syscall.Handle
150 +
151 + profile uint32
152 + requestOffset uint32
153 + requestCapacity uint32
154 + responseOffset uint32
155 + responseCapacity uint32
156 + SpinTries uint32
157 +
158 + localReqSeq int64
159 + localRespSeq int64
160 +}
161 +
162 +// Role returns the context role.
163 +func (c *WinShmContext) GetRole() WinShmRole { return c.role }
164 +
165 +// ---------------------------------------------------------------------------
166 +// Server API
167 +// ---------------------------------------------------------------------------
168 +
169 +// WinShmServerCreate creates a per-session Windows SHM region.
170 +func WinShmServerCreate(runDir, serviceName string, authToken, sessionID uint64,
171 + profile, reqCapacity, respCapacity uint32) (*WinShmContext, error) {
172 +
173 + if err := validateServiceName(serviceName); err != nil {
174 + return nil, err
175 + }
176 + if err := validateWinShmProfile(profile); err != nil {
177 + return nil, err
178 + }
179 +
180 + hash := computeShmHash(runDir, serviceName, authToken)
181 + mappingName, err := buildWinShmObjectName(hash, serviceName, profile, sessionID, "mapping")
182 + if err != nil {
183 + return nil, err
184 + }
185 +
186 + reqCap := winShmAlignCacheline(reqCapacity)
187 + respCap := winShmAlignCacheline(respCapacity)
188 + reqOff := winShmAlignCacheline(winShmHeaderLen)
189 + respOff := winShmAlignCacheline(reqOff + reqCap)
190 + regionSize := uintptr(respOff + respCap)
191 +
192 + // Create page-file backed mapping
193 + r, _, callErr := winShmCreateFileMappingW(
194 + uintptr(syscall.InvalidHandle), // page file
195 + 0, // NULL security
196 + uintptr(_PAGE_READWRITE),
197 + uintptr(regionSize>>32),
198 + uintptr(regionSize&0xFFFFFFFF),
199 + uintptr(unsafe.Pointer(&mappingName[0])),
200 + )
201 + mapping := syscall.Handle(r)
202 + if mapping == 0 {
203 + return nil, fmt.Errorf("%w: %v", ErrWinShmCreateMapping, callErr)
204 + }
205 + if isWindowsErrno(callErr, syscall.Errno(_ERROR_ALREADY_EXISTS)) {
206 + syscall.CloseHandle(mapping)
207 + return nil, ErrWinShmAddrInUse
208 + }
209 +
210 + // Map view
211 + base, _, callErr := winShmMapViewOfFile(
212 + uintptr(mapping),
213 + uintptr(_FILE_MAP_ALL_ACCESS),
214 + 0, 0,
215 + regionSize,
216 + )
217 + if base == 0 {
218 + syscall.CloseHandle(mapping)
219 + return nil, fmt.Errorf("%w: %v", ErrWinShmMapView, callErr)
220 + }
221 +
222 + // Zero region
223 + data := unsafe.Slice((*byte)(unsafe.Pointer(base)), regionSize)
224 + for i := range data {
225 + data[i] = 0
226 + }
227 +
228 + // Write header
229 + binary.NativeEndian.PutUint32(data[wshOFFMagic:], winShmMagic)
230 + binary.NativeEndian.PutUint32(data[wshOFFVersion:], winShmVersion)
231 + binary.NativeEndian.PutUint32(data[wshOFFHeaderLen:], winShmHeaderLen)
232 + binary.NativeEndian.PutUint32(data[wshOFFProfile:], profile)
233 + binary.NativeEndian.PutUint32(data[wshOFFReqOffset:], reqOff)
234 + binary.NativeEndian.PutUint32(data[wshOFFReqCapacity:], reqCap)
235 + binary.NativeEndian.PutUint32(data[wshOFFRespOffset:], respOff)
236 + binary.NativeEndian.PutUint32(data[wshOFFRespCapacity:], respCap)
237 + binary.NativeEndian.PutUint32(data[wshOFFSpinTries:], winShmDefaultSpin)
238 +
239 + // Release fence
240 + atomic.StoreInt32((*int32)(unsafe.Pointer(&data[wshOFFReqLen])), 0)
241 +
242 + // Create events for HYBRID
243 + var reqEvent, respEvent syscall.Handle
244 + reqEvent = syscall.InvalidHandle
245 + respEvent = syscall.InvalidHandle
246 +
247 + if profile == WinShmProfileHybrid {
248 + reqEventName, err := buildWinShmObjectName(hash, serviceName, profile, sessionID, "req_event")
249 + if err != nil {
250 + procUnmapViewOfFile.Call(base)
251 + syscall.CloseHandle(mapping)
252 + return nil, err
253 + }
254 +
255 + r, _, callErr := winShmCreateEventW(0, 0, 0,
256 + uintptr(unsafe.Pointer(&reqEventName[0])))
257 + if r == 0 {
258 + procUnmapViewOfFile.Call(base)
259 + syscall.CloseHandle(mapping)
260 + return nil, fmt.Errorf("%w: req_event: %v", ErrWinShmCreateEvent, callErr)
261 + }
262 + reqEvent = syscall.Handle(r)
263 + if isWindowsErrno(callErr, syscall.Errno(_ERROR_ALREADY_EXISTS)) {
264 + syscall.CloseHandle(reqEvent)
265 + procUnmapViewOfFile.Call(base)
266 + syscall.CloseHandle(mapping)
267 + return nil, ErrWinShmAddrInUse
268 + }
269 +
270 + respEventName, err := buildWinShmObjectName(hash, serviceName, profile, sessionID, "resp_event")
271 + if err != nil {
272 + syscall.CloseHandle(reqEvent)
273 + procUnmapViewOfFile.Call(base)
274 + syscall.CloseHandle(mapping)
275 + return nil, err
276 + }
277 +
278 + r, _, callErr = winShmCreateEventW(0, 0, 0,
279 + uintptr(unsafe.Pointer(&respEventName[0])))
280 + if r == 0 {
281 + syscall.CloseHandle(reqEvent)
282 + procUnmapViewOfFile.Call(base)
283 + syscall.CloseHandle(mapping)
284 + return nil, fmt.Errorf("%w: resp_event: %v", ErrWinShmCreateEvent, callErr)
285 + }
286 + respEvent = syscall.Handle(r)
287 + if isWindowsErrno(callErr, syscall.Errno(_ERROR_ALREADY_EXISTS)) {
288 + syscall.CloseHandle(respEvent)
289 + syscall.CloseHandle(reqEvent)
290 + procUnmapViewOfFile.Call(base)
291 + syscall.CloseHandle(mapping)
292 + return nil, ErrWinShmAddrInUse
293 + }
294 + }
295 +
296 + return &WinShmContext{
297 + role: WinShmRoleServer,
298 + mapping: mapping,
299 + base: base,
300 + size: regionSize,
301 + reqEvent: reqEvent,
302 + respEvent: respEvent,
303 + profile: profile,
304 + requestOffset: reqOff,
305 + requestCapacity: reqCap,
306 + responseOffset: respOff,
307 + responseCapacity: respCap,
308 + SpinTries: winShmDefaultSpin,
309 + localReqSeq: 0,
310 + localRespSeq: 0,
311 + }, nil
312 +}
313 +
314 +// WinShmDestroy destroys a server SHM region.
315 +func (c *WinShmContext) WinShmDestroy() {
316 + if c.base != 0 {
317 + data := unsafe.Slice((*byte)(unsafe.Pointer(c.base)), c.size)
318 + atomic.StoreInt32((*int32)(unsafe.Pointer(&data[wshOFFRespServerClosed])), 1)
319 + }
320 +
321 + if c.profile == WinShmProfileHybrid && c.respEvent != syscall.InvalidHandle {
322 + procSetEvent.Call(uintptr(c.respEvent))
323 + }
324 +
325 + c.cleanupHandles()
326 +}
327 +
328 +// ---------------------------------------------------------------------------
329 +// Client API
330 +// ---------------------------------------------------------------------------
331 +
332 +// WinShmClientAttach attaches to an existing per-session Windows SHM region.
333 +func WinShmClientAttach(runDir, serviceName string, authToken, sessionID uint64,
334 + profile uint32) (*WinShmContext, error) {
335 +
336 + if err := validateServiceName(serviceName); err != nil {
337 + return nil, err
338 + }
339 + if err := validateWinShmProfile(profile); err != nil {
340 + return nil, err
341 + }
342 +
343 + hash := computeShmHash(runDir, serviceName, authToken)
344 + mappingName, err := buildWinShmObjectName(hash, serviceName, profile, sessionID, "mapping")
345 + if err != nil {
346 + return nil, err
347 + }
348 +
349 + r, _, callErr := winShmOpenFileMappingW(
350 + uintptr(_FILE_MAP_ALL_ACCESS),
351 + 0,
352 + uintptr(unsafe.Pointer(&mappingName[0])),
353 + )
354 + mapping := syscall.Handle(r)
355 + if mapping == 0 {
356 + return nil, fmt.Errorf("%w: %v", ErrWinShmOpenMapping, callErr)
357 + }
358 +
359 + base, _, callErr := winShmMapViewOfFile(
360 + uintptr(mapping),
361 + uintptr(_FILE_MAP_ALL_ACCESS),
362 + 0, 0, 0,
363 + )
364 + if base == 0 {
365 + syscall.CloseHandle(mapping)
366 + return nil, fmt.Errorf("%w: %v", ErrWinShmMapView, callErr)
367 + }
368 +
369 + // We need at least header_len to validate
370 + data := unsafe.Slice((*byte)(unsafe.Pointer(base)), winShmHeaderLen)
371 +
372 + // Acquire fence
373 + atomic.LoadInt32((*int32)(unsafe.Pointer(&data[wshOFFReqLen])))
374 +
375 + // Validate header
376 + magic := binary.NativeEndian.Uint32(data[wshOFFMagic:])
377 + if magic != winShmMagic {
378 + procUnmapViewOfFile.Call(base)
379 + syscall.CloseHandle(mapping)
380 + return nil, ErrWinShmBadMagic
381 + }
382 + version := binary.NativeEndian.Uint32(data[wshOFFVersion:])
383 + if version != winShmVersion {
384 + procUnmapViewOfFile.Call(base)
385 + syscall.CloseHandle(mapping)
386 + return nil, ErrWinShmBadVersion
387 + }
388 + hdrLen := binary.NativeEndian.Uint32(data[wshOFFHeaderLen:])
389 + if hdrLen != winShmHeaderLen {
390 + procUnmapViewOfFile.Call(base)
391 + syscall.CloseHandle(mapping)
392 + return nil, ErrWinShmBadHeader
393 + }
394 + hdrProfile := binary.NativeEndian.Uint32(data[wshOFFProfile:])
395 + if hdrProfile != profile {
396 + procUnmapViewOfFile.Call(base)
397 + syscall.CloseHandle(mapping)
398 + return nil, ErrWinShmBadProfile
399 + }
400 +
401 + reqOff := binary.NativeEndian.Uint32(data[wshOFFReqOffset:])
402 + reqCap := binary.NativeEndian.Uint32(data[wshOFFReqCapacity:])
403 + respOff := binary.NativeEndian.Uint32(data[wshOFFRespOffset:])
404 + respCap := binary.NativeEndian.Uint32(data[wshOFFRespCapacity:])
405 + spin := binary.NativeEndian.Uint32(data[wshOFFSpinTries:])
406 + regionSize := uintptr(respOff + respCap)
407 +
408 + // Now reslice to full region
409 + fullData := unsafe.Slice((*byte)(unsafe.Pointer(base)), regionSize)
410 +
411 + // Read current sequence numbers via atomic
412 + curReqSeq := atomic.LoadInt64((*int64)(unsafe.Pointer(&fullData[wshOFFReqSeq])))
413 + curRespSeq := atomic.LoadInt64((*int64)(unsafe.Pointer(&fullData[wshOFFRespSeq])))
414 +
415 + // Open events for HYBRID
416 + var reqEvent, respEvent syscall.Handle
417 + reqEvent = syscall.InvalidHandle
418 + respEvent = syscall.InvalidHandle
419 +
420 + if profile == WinShmProfileHybrid {
421 + reqEventName, err := buildWinShmObjectName(hash, serviceName, profile, sessionID, "req_event")
422 + if err != nil {
423 + procUnmapViewOfFile.Call(base)
424 + syscall.CloseHandle(mapping)
425 + return nil, err
426 + }
427 +
428 + r, _, callErr := winShmOpenEventW(
429 + uintptr(_EVENT_MODIFY_STATE|_SYNCHRONIZE),
430 + 0,
431 + uintptr(unsafe.Pointer(&reqEventName[0])),
432 + )
433 + if r == 0 {
434 + procUnmapViewOfFile.Call(base)
435 + syscall.CloseHandle(mapping)
436 + return nil, fmt.Errorf("%w: req_event: %v", ErrWinShmOpenEvent, callErr)
437 + }
438 + reqEvent = syscall.Handle(r)
439 +
440 + respEventName, err := buildWinShmObjectName(hash, serviceName, profile, sessionID, "resp_event")
441 + if err != nil {
442 + syscall.CloseHandle(reqEvent)
443 + procUnmapViewOfFile.Call(base)
444 + syscall.CloseHandle(mapping)
445 + return nil, err
446 + }
447 +
448 + r, _, callErr = winShmOpenEventW(
449 + uintptr(_EVENT_MODIFY_STATE|_SYNCHRONIZE),
450 + 0,
451 + uintptr(unsafe.Pointer(&respEventName[0])),
452 + )
453 + if r == 0 {
454 + syscall.CloseHandle(reqEvent)
455 + procUnmapViewOfFile.Call(base)
456 + syscall.CloseHandle(mapping)
457 + return nil, fmt.Errorf("%w: resp_event: %v", ErrWinShmOpenEvent, callErr)
458 + }
459 + respEvent = syscall.Handle(r)
460 + }
461 +
462 + return &WinShmContext{
463 + role: WinShmRoleClient,
464 + mapping: mapping,
465 + base: base,
466 + size: regionSize,
467 + reqEvent: reqEvent,
468 + respEvent: respEvent,
469 + profile: profile,
470 + requestOffset: reqOff,
471 + requestCapacity: reqCap,
472 + responseOffset: respOff,
473 + responseCapacity: respCap,
474 + SpinTries: spin,
475 + localReqSeq: curReqSeq,
476 + localRespSeq: curRespSeq,
477 + }, nil
478 +}
479 +
480 +// WinShmClose closes a client SHM context.
481 +func (c *WinShmContext) WinShmClose() {
482 + if c.base != 0 {
483 + data := unsafe.Slice((*byte)(unsafe.Pointer(c.base)), c.size)
484 + atomic.StoreInt32((*int32)(unsafe.Pointer(&data[wshOFFReqClientClosed])), 1)
485 + }
486 +
487 + if c.profile == WinShmProfileHybrid && c.reqEvent != syscall.InvalidHandle {
488 + procSetEvent.Call(uintptr(c.reqEvent))
489 + }
490 +
491 + c.cleanupHandles()
492 +}
493 +
494 +func (c *WinShmContext) cleanupHandles() {
495 + if c.base != 0 {
496 + procUnmapViewOfFile.Call(c.base)
497 + c.base = 0
498 + }
499 + if c.mapping != syscall.InvalidHandle && c.mapping != 0 {
500 + syscall.CloseHandle(c.mapping)
501 + c.mapping = syscall.InvalidHandle
502 + }
503 + if c.reqEvent != syscall.InvalidHandle {
504 + syscall.CloseHandle(c.reqEvent)
505 + c.reqEvent = syscall.InvalidHandle
506 + }
507 + if c.respEvent != syscall.InvalidHandle {
508 + syscall.CloseHandle(c.respEvent)
509 + c.respEvent = syscall.InvalidHandle
510 + }
511 + c.size = 0
512 +}
513 +
514 +// ---------------------------------------------------------------------------
515 +// Data plane
516 +// ---------------------------------------------------------------------------
517 +
518 +// WinShmSend publishes a message. The message must include the 32-byte
519 +// outer header + payload, exactly as sent over Named Pipe.
520 +func (c *WinShmContext) WinShmSend(msg []byte) error {
521 + if c.base == 0 || len(msg) == 0 {
522 + return fmt.Errorf("%w: null context or empty message", ErrWinShmBadParam)
523 + }
524 +
525 + var areaOff, areaCap uint32
526 + var lenOff, seqOff, peerWaitingOff int
527 + var peerEvent syscall.Handle
528 +
529 + if c.role == WinShmRoleClient {
530 + areaOff = c.requestOffset
531 + areaCap = c.requestCapacity
532 + lenOff = wshOFFReqLen
533 + seqOff = wshOFFReqSeq
534 + peerWaitingOff = wshOFFReqServerWaiting
535 + peerEvent = c.reqEvent
536 + } else {
537 + areaOff = c.responseOffset
538 + areaCap = c.responseCapacity
539 + lenOff = wshOFFRespLen
540 + seqOff = wshOFFRespSeq
541 + peerWaitingOff = wshOFFRespClientWaiting
542 + peerEvent = c.respEvent
543 + }
544 +
545 + if uint32(len(msg)) > areaCap {
546 + return ErrWinShmMsgTooLarge
547 + }
548 +
549 + data := unsafe.Slice((*byte)(unsafe.Pointer(c.base)), c.size)
550 +
551 + // 1. Write message data
552 + copy(data[areaOff:], msg)
553 +
554 + // 2. Store message length (atomic)
555 + atomic.StoreInt32((*int32)(unsafe.Pointer(&data[lenOff])), int32(len(msg)))
556 +
557 + // 3. Increment sequence number (atomic)
558 + atomic.AddInt64((*int64)(unsafe.Pointer(&data[seqOff])), 1)
559 +
560 + // 4. If HYBRID and peer waiting, signal event
561 + if c.profile == WinShmProfileHybrid {
562 + if atomic.LoadInt32((*int32)(unsafe.Pointer(&data[peerWaitingOff]))) != 0 {
563 + procSetEvent.Call(uintptr(peerEvent))
564 + }
565 + }
566 +
567 + if c.role == WinShmRoleClient {
568 + c.localReqSeq++
569 + } else {
570 + c.localRespSeq++
571 + }
572 +
573 + return nil
574 +}
575 +
576 +// WinShmReceive receives a message into the caller-provided buffer.
577 +func (c *WinShmContext) WinShmReceive(buf []byte, timeoutMs uint32) (int, error) {
578 + if c.base == 0 {
579 + return 0, fmt.Errorf("%w: null context", ErrWinShmBadParam)
580 + }
581 + if len(buf) == 0 {
582 + return 0, fmt.Errorf("%w: empty buffer", ErrWinShmBadParam)
583 + }
584 +
585 + var areaOff, areaCap uint32
586 + var lenOff, seqOff, selfWaitingOff, peerClosedOff int
587 + var waitEvent syscall.Handle
588 + var expectedSeq int64
589 +
590 + if c.role == WinShmRoleServer {
591 + areaOff = c.requestOffset
592 + areaCap = c.requestCapacity
593 + lenOff = wshOFFReqLen
594 + seqOff = wshOFFReqSeq
595 + selfWaitingOff = wshOFFReqServerWaiting
596 + peerClosedOff = wshOFFReqClientClosed
597 + waitEvent = c.reqEvent
598 + expectedSeq = c.localReqSeq + 1
599 + } else {
600 + areaOff = c.responseOffset
601 + areaCap = c.responseCapacity
602 + lenOff = wshOFFRespLen
603 + seqOff = wshOFFRespSeq
604 + selfWaitingOff = wshOFFRespClientWaiting
605 + peerClosedOff = wshOFFRespServerClosed
606 + waitEvent = c.respEvent
607 + expectedSeq = c.localRespSeq + 1
608 + }
609 +
610 + // The copy ceiling is the smaller of the caller buffer and the
611 + // SHM area capacity. Prevents out-of-bounds reads on forged lengths.
612 + maxCopy := len(buf)
613 + if int(areaCap) < maxCopy {
614 + maxCopy = int(areaCap)
615 + }
616 +
617 + data := unsafe.Slice((*byte)(unsafe.Pointer(c.base)), c.size)
618 +
619 + // Phase 1: spin
620 + observed := false
621 + var mlen int32
622 + for i := uint32(0); i < c.SpinTries; i++ {
623 + cur := atomic.LoadInt64((*int64)(unsafe.Pointer(&data[seqOff])))
624 + if cur >= expectedSeq {
625 + mlen = atomic.LoadInt32((*int32)(unsafe.Pointer(&data[lenOff])))
626 + if mlen > 0 && int(mlen) <= maxCopy {
627 + copy(buf[:mlen], data[areaOff:areaOff+uint32(mlen)])
628 + }
629 + observed = true
630 + break
631 + }
632 + spinPause()
633 + }
634 +
635 + // Phase 2: kernel wait or busy-wait
636 + if !observed {
637 + if c.profile == WinShmProfileHybrid {
638 + deadlineMs := uint32(_INFINITE)
639 + if timeoutMs > 0 {
640 + deadlineMs = timeoutMs
641 + }
642 + start, _, _ := procGetTickCount64.Call()
643 +
644 + for {
645 + atomic.StoreInt32((*int32)(unsafe.Pointer(&data[selfWaitingOff])), 1)
646 + atomic.LoadInt32((*int32)(unsafe.Pointer(&data[selfWaitingOff])))
647 +
648 + cur := atomic.LoadInt64((*int64)(unsafe.Pointer(&data[seqOff])))
649 + if cur >= expectedSeq {
650 + atomic.StoreInt32((*int32)(unsafe.Pointer(&data[selfWaitingOff])), 0)
651 + break
652 + }
653 +
654 + waitMs := uintptr(_INFINITE)
655 + if deadlineMs != _INFINITE {
656 + now, _, _ := procGetTickCount64.Call()
657 + elapsed := uint32(now - start)
658 + if elapsed >= deadlineMs {
659 + atomic.StoreInt32((*int32)(unsafe.Pointer(&data[selfWaitingOff])), 0)
660 + return 0, ErrWinShmTimeout
661 + }
662 + waitMs = uintptr(deadlineMs - elapsed)
663 + }
664 +
665 + ret, _, _ := procWaitForSingleObj.Call(uintptr(waitEvent), waitMs)
666 + atomic.StoreInt32((*int32)(unsafe.Pointer(&data[selfWaitingOff])), 0)
667 +
668 + cur = atomic.LoadInt64((*int64)(unsafe.Pointer(&data[seqOff])))
669 + if cur >= expectedSeq {
670 + break
671 + }
672 +
673 + if atomic.LoadInt32((*int32)(unsafe.Pointer(&data[peerClosedOff]))) != 0 {
674 + cur = atomic.LoadInt64((*int64)(unsafe.Pointer(&data[seqOff])))
675 + if cur >= expectedSeq {
676 + break
677 + }
678 + c.advanceSeq(expectedSeq)
679 + return 0, ErrWinShmDisconnected
680 + }
681 +
682 + if ret == _WAIT_TIMEOUT {
683 + return 0, ErrWinShmTimeout
684 + }
685 + }
686 +
687 + mlen = atomic.LoadInt32((*int32)(unsafe.Pointer(&data[lenOff])))
688 + if mlen > 0 && int(mlen) <= maxCopy {
689 + copy(buf[:mlen], data[areaOff:areaOff+uint32(mlen)])
690 + }
691 + } else {
692 + // BUSYWAIT
693 + start, _, _ := procGetTickCount64.Call()
694 + for {
695 + cur := atomic.LoadInt64((*int64)(unsafe.Pointer(&data[seqOff])))
696 + if cur >= expectedSeq {
697 + mlen = atomic.LoadInt32((*int32)(unsafe.Pointer(&data[lenOff])))
698 + if mlen > 0 && int(mlen) <= maxCopy {
699 + copy(buf[:mlen], data[areaOff:areaOff+uint32(mlen)])
700 + }
701 + break
702 + }
703 +
704 + if timeoutMs > 0 {
705 + now, _, _ := procGetTickCount64.Call()
706 + elapsed := uint64(now - start)
707 + if elapsed >= uint64(timeoutMs) {
708 + return 0, ErrWinShmTimeout
709 + }
710 + }
711 +
712 + if atomic.LoadInt32((*int32)(unsafe.Pointer(&data[peerClosedOff]))) != 0 {
713 + cur := atomic.LoadInt64((*int64)(unsafe.Pointer(&data[seqOff])))
714 + if cur >= expectedSeq {
715 + mlen = atomic.LoadInt32((*int32)(unsafe.Pointer(&data[lenOff])))
716 + if mlen > 0 && int(mlen) <= maxCopy {
717 + copy(buf[:mlen], data[areaOff:areaOff+uint32(mlen)])
718 + }
719 + break
720 + }
721 + c.advanceSeq(expectedSeq)
722 + return 0, ErrWinShmDisconnected
723 + }
724 +
725 + spinPause()
726 + }
727 + }
728 + }
729 +
730 + c.advanceSeq(expectedSeq)
731 +
732 + if int(mlen) > maxCopy {
733 + return int(mlen), ErrWinShmMsgTooLarge
734 + }
735 +
736 + return int(mlen), nil
737 +}
738 +
739 +func (c *WinShmContext) advanceSeq(expectedSeq int64) {
740 + if c.role == WinShmRoleServer {
741 + c.localReqSeq = expectedSeq
742 + } else {
743 + c.localRespSeq = expectedSeq
744 + }
745 +}
746 +
747 +// ---------------------------------------------------------------------------
748 +// Internal helpers
749 +// ---------------------------------------------------------------------------
750 +
751 +func winShmAlignCacheline(v uint32) uint32 {
752 + return (v + (winShmCachelineSize - 1)) & ^(winShmCachelineSize - 1)
753 +}
754 +
755 +func validateWinShmProfile(profile uint32) error {
756 + if profile != WinShmProfileHybrid && profile != WinShmProfileBusywait {
757 + return fmt.Errorf("%w: invalid profile %d", ErrWinShmBadParam, profile)
758 + }
759 + return nil
760 +}
761 +
762 +func computeShmHash(runDir, serviceName string, authToken uint64) uint64 {
763 + input := fmt.Sprintf("%s\n%s\n%d", runDir, serviceName, authToken)
764 + return FNV1a64([]byte(input))
765 +}
766 +
767 +func buildWinShmObjectName(hash uint64, serviceName string,
768 + profile uint32, sessionID uint64, suffix string) ([]uint16, error) {
769 +
770 + narrow := fmt.Sprintf(`Local\netipc-%016x-%s-p%d-s%016x-%s`,
771 + hash, serviceName, profile, sessionID, suffix)
772 + if len(narrow) >= 256 {
773 + return nil, fmt.Errorf("%w: object name too long", ErrWinShmBadParam)
774 + }
775 +
776 + // Convert to NUL-terminated UTF-16
777 + runes := []rune(narrow)
778 + runes = append(runes, 0)
779 + utf16 := make([]uint16, len(runes))
780 + for i, r := range runes {
781 + utf16[i] = uint16(r)
782 + }
783 + return utf16, nil
784 +}
src/go/pkg/netipc/transport/windows/shm_pause.go new
+9
@@ -0,0 +1,9 @@
1 +//go:build windows && !amd64
2 +
3 +package windows
4 +
5 +// spinPause yields the current Windows thread on architectures where we
6 +// do not provide a CPU pause instruction helper.
7 +func spinPause() {
8 + procSwitchToThread.Call()
9 +}
src/go/pkg/netipc/transport/windows/shm_pause_amd64.go new
+6
@@ -0,0 +1,6 @@
1 +//go:build windows && amd64
2 +
3 +package windows
4 +
5 +// spinPause emits a PAUSE instruction (defined in shm_pause_amd64.s).
6 +func spinPause()
src/go/pkg/netipc/transport/windows/shm_pause_amd64.s new
+9
@@ -0,0 +1,9 @@
1 +//go:build windows && amd64
2 +
3 +#include "textflag.h"
4 +
5 +// func spinPause()
6 +// CPU PAUSE hint for SHM spin loops.
7 +TEXT ·spinPause(SB),NOSPLIT,$0-0
8 + PAUSE
9 + RET
src/go/pkg/netipc/transport/windows/shm_test.go new
+770
@@ -0,0 +1,770 @@
1 +//go:build windows
2 +
3 +package windows
4 +
5 +import (
6 + "encoding/binary"
7 + "errors"
8 + "strings"
9 + "sync/atomic"
10 + "syscall"
11 + "testing"
12 + "time"
13 + "unsafe"
14 +)
15 +
16 +type shmReceiveResult struct {
17 + payload []byte
18 + err error
19 +}
20 +
21 +func installWinShmOneShotFault(t *testing.T, target *winShmProcCall, failOnCall int, errno syscall.Errno) {
22 + t.Helper()
23 +
24 + orig := *target
25 + callCount := 0
26 + *target = func(a ...uintptr) (uintptr, uintptr, error) {
27 + callCount++
28 + if callCount == failOnCall {
29 + return 0, 0, errno
30 + }
31 + return orig(a...)
32 + }
33 + t.Cleanup(func() {
34 + *target = orig
35 + })
36 +}
37 +
38 +func TestWinShmHelpers(t *testing.T) {
39 + if got := winShmAlignCacheline(1); got != 64 {
40 + t.Fatalf("winShmAlignCacheline(1) = %d, want 64", got)
41 + }
42 + if got := winShmAlignCacheline(64); got != 64 {
43 + t.Fatalf("winShmAlignCacheline(64) = %d, want 64", got)
44 + }
45 + if got := winShmAlignCacheline(65); got != 128 {
46 + t.Fatalf("winShmAlignCacheline(65) = %d, want 128", got)
47 + }
48 +
49 + if err := validateWinShmProfile(0); !errors.Is(err, ErrWinShmBadParam) {
50 + t.Fatalf("validateWinShmProfile(0) = %v, want ErrWinShmBadParam", err)
51 + }
52 + if err := validateWinShmProfile(WinShmProfileHybrid); err != nil {
53 + t.Fatalf("validateWinShmProfile(hybrid) failed: %v", err)
54 + }
55 + if err := validateWinShmProfile(WinShmProfileBusywait); err != nil {
56 + t.Fatalf("validateWinShmProfile(busywait) failed: %v", err)
57 + }
58 +
59 + h1 := computeShmHash("run", "svc", 123)
60 + h2 := computeShmHash("run", "svc", 123)
61 + h3 := computeShmHash("run", "svc", 124)
62 + if h1 != h2 {
63 + t.Fatalf("computeShmHash not deterministic: %d != %d", h1, h2)
64 + }
65 + if h1 == h3 {
66 + t.Fatal("computeShmHash should change when auth token changes")
67 + }
68 +
69 + name, err := buildWinShmObjectName(h1, "svc", WinShmProfileHybrid, 7, "mapping")
70 + if err != nil {
71 + t.Fatalf("buildWinShmObjectName failed: %v", err)
72 + }
73 + if got := syscall.UTF16ToString(name); !strings.Contains(got, "svc") {
74 + t.Fatalf("object name %q does not contain service name", got)
75 + }
76 +
77 + tooLong := strings.Repeat("a", 260)
78 + if _, err := buildWinShmObjectName(h1, tooLong, WinShmProfileHybrid, 7, "mapping"); !errors.Is(err, ErrWinShmBadParam) {
79 + t.Fatalf("buildWinShmObjectName long name = %v, want ErrWinShmBadParam", err)
80 + }
81 +}
82 +
83 +func TestWinShmCreateAttachAndCloseValidation(t *testing.T) {
84 + runDir := t.TempDir()
85 + service := "validation"
86 + const authToken uint64 = 0x5678
87 + const sessionID uint64 = 9
88 +
89 + if _, err := WinShmServerCreate(runDir, "bad/name", authToken, sessionID, WinShmProfileHybrid, 4096, 4096); err == nil {
90 + t.Fatal("WinShmServerCreate with invalid service name should fail")
91 + }
92 + if _, err := WinShmServerCreate(runDir, service, authToken, sessionID, 0, 4096, 4096); !errors.Is(err, ErrWinShmBadParam) {
93 + t.Fatalf("WinShmServerCreate invalid profile = %v, want ErrWinShmBadParam", err)
94 + }
95 + if _, err := WinShmClientAttach(runDir, service, authToken, sessionID, 0); !errors.Is(err, ErrWinShmBadParam) {
96 + t.Fatalf("WinShmClientAttach invalid profile = %v, want ErrWinShmBadParam", err)
97 + }
98 + if _, err := WinShmClientAttach(runDir, service, authToken, sessionID, WinShmProfileHybrid); err == nil {
99 + t.Fatal("WinShmClientAttach without mapping should fail")
100 + }
101 +
102 + server, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileHybrid, 4096, 4096)
103 + if err != nil {
104 + t.Fatalf("WinShmServerCreate failed: %v", err)
105 + }
106 + defer server.WinShmDestroy()
107 +
108 + client, err := WinShmClientAttach(runDir, service, authToken, sessionID, WinShmProfileHybrid)
109 + if err != nil {
110 + t.Fatalf("WinShmClientAttach failed: %v", err)
111 + }
112 + defer client.WinShmClose()
113 +
114 + if server.GetRole() != WinShmRoleServer {
115 + t.Fatalf("server role = %d, want %d", server.GetRole(), WinShmRoleServer)
116 + }
117 + if client.GetRole() != WinShmRoleClient {
118 + t.Fatalf("client role = %d, want %d", client.GetRole(), WinShmRoleClient)
119 + }
120 +}
121 +
122 +func TestWinShmCreateAttachRejectsLongObjectName(t *testing.T) {
123 + runDir := t.TempDir()
124 + service := strings.Repeat("a", 220)
125 + const authToken uint64 = 0x9988
126 + const sessionID uint64 = 23
127 +
128 + if _, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileHybrid, 4096, 4096); !errors.Is(err, ErrWinShmBadParam) {
129 + t.Fatalf("WinShmServerCreate(long name) = %v, want ErrWinShmBadParam", err)
130 + }
131 + if _, err := WinShmClientAttach(runDir, service, authToken, sessionID, WinShmProfileHybrid); !errors.Is(err, ErrWinShmBadParam) {
132 + t.Fatalf("WinShmClientAttach(long name) = %v, want ErrWinShmBadParam", err)
133 + }
134 +}
135 +
136 +func TestWinShmServerCreateRejectsExistingObjects(t *testing.T) {
137 + runDir := t.TempDir()
138 + service := "addr-in-use"
139 + const authToken uint64 = 0x445566
140 + const sessionID uint64 = 25
141 +
142 + server, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileHybrid, 4096, 4096)
143 + if err != nil {
144 + t.Fatalf("WinShmServerCreate failed: %v", err)
145 + }
146 + defer server.WinShmDestroy()
147 +
148 + if _, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileHybrid, 4096, 4096); !errors.Is(err, ErrWinShmAddrInUse) {
149 + t.Fatalf("second WinShmServerCreate = %v, want ErrWinShmAddrInUse", err)
150 + }
151 +}
152 +
153 +func TestWinShmServerCreateRejectsLateEventNameOverflow(t *testing.T) {
154 + runDir := t.TempDir()
155 + const authToken uint64 = 0x9989
156 + const sessionID uint64 = 24
157 +
158 + t.Run("req_event name overflow", func(t *testing.T) {
159 + // 195 characters keeps the mapping name below the Win32 object-name
160 + // limit but pushes req_event over it.
161 + service := strings.Repeat("b", 195)
162 + if _, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileHybrid, 4096, 4096); !errors.Is(err, ErrWinShmBadParam) {
163 + t.Fatalf("WinShmServerCreate(req_event overflow) = %v, want ErrWinShmBadParam", err)
164 + }
165 + })
166 +
167 + t.Run("resp_event name overflow", func(t *testing.T) {
168 + // 194 characters still allows req_event creation, but resp_event
169 + // crosses the object-name limit and exercises the cleanup path.
170 + service := strings.Repeat("c", 194)
171 + if _, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileHybrid, 4096, 4096); !errors.Is(err, ErrWinShmBadParam) {
172 + t.Fatalf("WinShmServerCreate(resp_event overflow) = %v, want ErrWinShmBadParam", err)
173 + }
174 + })
175 +}
176 +
177 +func TestWinShmClientAttachRejectsCorruptHeader(t *testing.T) {
178 + cases := []struct {
179 + name string
180 + want error
181 + mutate func([]byte)
182 + }{
183 + {
184 + name: "bad magic",
185 + want: ErrWinShmBadMagic,
186 + mutate: func(data []byte) {
187 + binary.NativeEndian.PutUint32(data[wshOFFMagic:], 0)
188 + },
189 + },
190 + {
191 + name: "bad version",
192 + want: ErrWinShmBadVersion,
193 + mutate: func(data []byte) {
194 + binary.NativeEndian.PutUint32(data[wshOFFVersion:], winShmVersion+1)
195 + },
196 + },
197 + {
198 + name: "bad header len",
199 + want: ErrWinShmBadHeader,
200 + mutate: func(data []byte) {
201 + binary.NativeEndian.PutUint32(data[wshOFFHeaderLen:], winShmHeaderLen+64)
202 + },
203 + },
204 + {
205 + name: "bad profile",
206 + want: ErrWinShmBadProfile,
207 + mutate: func(data []byte) {
208 + binary.NativeEndian.PutUint32(data[wshOFFProfile:], WinShmProfileBusywait)
209 + },
210 + },
211 + }
212 +
213 + for _, tc := range cases {
214 + t.Run(tc.name, func(t *testing.T) {
215 + runDir := t.TempDir()
216 + service := "corrupt-header"
217 + const authToken uint64 = 0x123456
218 + const sessionID uint64 = 15
219 +
220 + server, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileHybrid, 4096, 4096)
221 + if err != nil {
222 + t.Fatalf("WinShmServerCreate failed: %v", err)
223 + }
224 + defer server.WinShmDestroy()
225 +
226 + data := unsafe.Slice((*byte)(unsafe.Pointer(server.base)), server.size)
227 + tc.mutate(data)
228 +
229 + _, err = WinShmClientAttach(runDir, service, authToken, sessionID, WinShmProfileHybrid)
230 + if !errors.Is(err, tc.want) {
231 + t.Fatalf("WinShmClientAttach error = %v, want %v", err, tc.want)
232 + }
233 + })
234 + }
235 +}
236 +
237 +func TestWinShmClientAttachFailsWhenEventsAreMissing(t *testing.T) {
238 + runDir := t.TempDir()
239 + service := "missing-events"
240 + const authToken uint64 = 0x556677
241 +
242 + t.Run("req_event missing", func(t *testing.T) {
243 + const sessionID uint64 = 31
244 +
245 + server, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileHybrid, 4096, 4096)
246 + if err != nil {
247 + t.Fatalf("WinShmServerCreate failed: %v", err)
248 + }
249 + defer server.WinShmDestroy()
250 +
251 + if err := syscall.CloseHandle(server.reqEvent); err != nil {
252 + t.Fatalf("CloseHandle(reqEvent) failed: %v", err)
253 + }
254 + server.reqEvent = syscall.InvalidHandle
255 +
256 + _, err = WinShmClientAttach(runDir, service, authToken, sessionID, WinShmProfileHybrid)
257 + if err == nil || !errors.Is(err, ErrWinShmOpenEvent) {
258 + t.Fatalf("WinShmClientAttach missing req_event = %v, want ErrWinShmOpenEvent", err)
259 + }
260 + })
261 +
262 + t.Run("resp_event missing", func(t *testing.T) {
263 + const sessionID uint64 = 33
264 +
265 + server, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileHybrid, 4096, 4096)
266 + if err != nil {
267 + t.Fatalf("WinShmServerCreate failed: %v", err)
268 + }
269 + defer server.WinShmDestroy()
270 +
271 + if err := syscall.CloseHandle(server.respEvent); err != nil {
272 + t.Fatalf("CloseHandle(respEvent) failed: %v", err)
273 + }
274 + server.respEvent = syscall.InvalidHandle
275 +
276 + _, err = WinShmClientAttach(runDir, service, authToken, sessionID, WinShmProfileHybrid)
277 + if err == nil || !errors.Is(err, ErrWinShmOpenEvent) {
278 + t.Fatalf("WinShmClientAttach missing resp_event = %v, want ErrWinShmOpenEvent", err)
279 + }
280 + })
281 +}
282 +
283 +func TestWinShmServerCreateWin32Failures(t *testing.T) {
284 + cases := []struct {
285 + name string
286 + target *winShmProcCall
287 + failOnCall int
288 + errno syscall.Errno
289 + wantErr error
290 + wantText string
291 + }{
292 + {
293 + name: "CreateFileMappingW",
294 + target: &winShmCreateFileMappingW,
295 + failOnCall: 1,
296 + errno: syscall.Errno(5),
297 + wantErr: ErrWinShmCreateMapping,
298 + },
299 + {
300 + name: "MapViewOfFile",
301 + target: &winShmMapViewOfFile,
302 + failOnCall: 1,
303 + errno: syscall.Errno(6),
304 + wantErr: ErrWinShmMapView,
305 + },
306 + {
307 + name: "CreateEventW req_event",
308 + target: &winShmCreateEventW,
309 + failOnCall: 1,
310 + errno: syscall.Errno(7),
311 + wantErr: ErrWinShmCreateEvent,
312 + wantText: "req_event",
313 + },
314 + {
315 + name: "CreateEventW resp_event",
316 + target: &winShmCreateEventW,
317 + failOnCall: 2,
318 + errno: syscall.Errno(8),
319 + wantErr: ErrWinShmCreateEvent,
320 + wantText: "resp_event",
321 + },
322 + }
323 +
324 + for _, tc := range cases {
325 + t.Run(tc.name, func(t *testing.T) {
326 + runDir := t.TempDir()
327 + service := "fault-create"
328 + const authToken uint64 = 0x7711
329 + const sessionID uint64 = 61
330 +
331 + installWinShmOneShotFault(t, tc.target, tc.failOnCall, tc.errno)
332 +
333 + _, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileHybrid, 4096, 4096)
334 + if !errors.Is(err, tc.wantErr) {
335 + t.Fatalf("WinShmServerCreate injected fault = %v, want %v", err, tc.wantErr)
336 + }
337 + if tc.wantText != "" && !strings.Contains(err.Error(), tc.wantText) {
338 + t.Fatalf("WinShmServerCreate error %q does not mention %q", err, tc.wantText)
339 + }
340 +
341 + server, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileHybrid, 4096, 4096)
342 + if err != nil {
343 + t.Fatalf("WinShmServerCreate recovery failed: %v", err)
344 + }
345 + server.WinShmDestroy()
346 + })
347 + }
348 +}
349 +
350 +func TestWinShmClientAttachWin32Failures(t *testing.T) {
351 + cases := []struct {
352 + name string
353 + target *winShmProcCall
354 + failOnCall int
355 + errno syscall.Errno
356 + wantErr error
357 + wantText string
358 + }{
359 + {
360 + name: "OpenFileMappingW",
361 + target: &winShmOpenFileMappingW,
362 + failOnCall: 1,
363 + errno: syscall.Errno(9),
364 + wantErr: ErrWinShmOpenMapping,
365 + },
366 + {
367 + name: "MapViewOfFile",
368 + target: &winShmMapViewOfFile,
369 + failOnCall: 1,
370 + errno: syscall.Errno(10),
371 + wantErr: ErrWinShmMapView,
372 + },
373 + {
374 + name: "OpenEventW req_event",
375 + target: &winShmOpenEventW,
376 + failOnCall: 1,
377 + errno: syscall.Errno(11),
378 + wantErr: ErrWinShmOpenEvent,
379 + wantText: "req_event",
380 + },
381 + {
382 + name: "OpenEventW resp_event",
383 + target: &winShmOpenEventW,
384 + failOnCall: 2,
385 + errno: syscall.Errno(12),
386 + wantErr: ErrWinShmOpenEvent,
387 + wantText: "resp_event",
388 + },
389 + }
390 +
391 + for _, tc := range cases {
392 + t.Run(tc.name, func(t *testing.T) {
393 + runDir := t.TempDir()
394 + service := "fault-attach"
395 + const authToken uint64 = 0x7712
396 + const sessionID uint64 = 62
397 +
398 + server, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileHybrid, 4096, 4096)
399 + if err != nil {
400 + t.Fatalf("WinShmServerCreate failed: %v", err)
401 + }
402 + defer server.WinShmDestroy()
403 +
404 + installWinShmOneShotFault(t, tc.target, tc.failOnCall, tc.errno)
405 +
406 + _, err = WinShmClientAttach(runDir, service, authToken, sessionID, WinShmProfileHybrid)
407 + if !errors.Is(err, tc.wantErr) {
408 + t.Fatalf("WinShmClientAttach injected fault = %v, want %v", err, tc.wantErr)
409 + }
410 + if tc.wantText != "" && !strings.Contains(err.Error(), tc.wantText) {
411 + t.Fatalf("WinShmClientAttach error %q does not mention %q", err, tc.wantText)
412 + }
413 +
414 + client, err := WinShmClientAttach(runDir, service, authToken, sessionID, WinShmProfileHybrid)
415 + if err != nil {
416 + t.Fatalf("WinShmClientAttach recovery failed: %v", err)
417 + }
418 + client.WinShmClose()
419 + })
420 + }
421 +}
422 +
423 +func TestWinShmSendReceiveValidation(t *testing.T) {
424 + runDir := t.TempDir()
425 + service := "validation-io"
426 + const authToken uint64 = 0x6789
427 + const sessionID uint64 = 11
428 +
429 + server, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileHybrid, 4096, 4096)
430 + if err != nil {
431 + t.Fatalf("WinShmServerCreate failed: %v", err)
432 + }
433 + defer server.WinShmDestroy()
434 +
435 + client, err := WinShmClientAttach(runDir, service, authToken, sessionID, WinShmProfileHybrid)
436 + if err != nil {
437 + t.Fatalf("WinShmClientAttach failed: %v", err)
438 + }
439 + defer client.WinShmClose()
440 +
441 + if err := server.WinShmSend(nil); !errors.Is(err, ErrWinShmBadParam) {
442 + t.Fatalf("WinShmSend(nil) = %v, want ErrWinShmBadParam", err)
443 + }
444 + if _, err := client.WinShmReceive(nil, 10); !errors.Is(err, ErrWinShmBadParam) {
445 + t.Fatalf("WinShmReceive(nil) = %v, want ErrWinShmBadParam", err)
446 + }
447 +
448 + tooLarge := make([]byte, server.responseCapacity+1)
449 + if err := server.WinShmSend(tooLarge); !errors.Is(err, ErrWinShmMsgTooLarge) {
450 + t.Fatalf("WinShmSend(tooLarge) = %v, want ErrWinShmMsgTooLarge", err)
451 + }
452 +
453 + msg := []byte("0123456789")
454 + if err := server.WinShmSend(msg); err != nil {
455 + t.Fatalf("WinShmSend failed: %v", err)
456 + }
457 +
458 + smallBuf := make([]byte, 4)
459 + n, err := client.WinShmReceive(smallBuf, 1000)
460 + if !errors.Is(err, ErrWinShmMsgTooLarge) {
461 + t.Fatalf("WinShmReceive(smallBuf) = %v, want ErrWinShmMsgTooLarge", err)
462 + }
463 + if n != len(msg) {
464 + t.Fatalf("WinShmReceive reported len %d, want %d", n, len(msg))
465 + }
466 +}
467 +
468 +func TestWinShmReceiveDetectsPeerClosed(t *testing.T) {
469 + runDir := t.TempDir()
470 + service := "peer-closed"
471 + const authToken uint64 = 0x789a
472 + const sessionID uint64 = 13
473 +
474 + server, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileHybrid, 4096, 4096)
475 + if err != nil {
476 + t.Fatalf("WinShmServerCreate failed: %v", err)
477 + }
478 + defer server.WinShmDestroy()
479 +
480 + client, err := WinShmClientAttach(runDir, service, authToken, sessionID, WinShmProfileHybrid)
481 + if err != nil {
482 + t.Fatalf("WinShmClientAttach failed: %v", err)
483 + }
484 +
485 + results := make(chan error, 1)
486 + go func() {
487 + buf := make([]byte, 128)
488 + _, err := server.WinShmReceive(buf, 1000)
489 + results <- err
490 + }()
491 +
492 + time.Sleep(20 * time.Millisecond)
493 + client.WinShmClose()
494 +
495 + select {
496 + case err := <-results:
497 + if !errors.Is(err, ErrWinShmDisconnected) {
498 + t.Fatalf("WinShmReceive after peer close = %v, want ErrWinShmDisconnected", err)
499 + }
500 + case <-time.After(2 * time.Second):
501 + t.Fatal("WinShmReceive did not observe peer close")
502 + }
503 +}
504 +
505 +func TestWinShmReceiveTimeoutHybrid(t *testing.T) {
506 + runDir := t.TempDir()
507 + service := "timeout-hybrid"
508 + const authToken uint64 = 0x8912
509 + const sessionID uint64 = 17
510 +
511 + server, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileHybrid, 4096, 4096)
512 + if err != nil {
513 + t.Fatalf("WinShmServerCreate failed: %v", err)
514 + }
515 + defer server.WinShmDestroy()
516 +
517 + client, err := WinShmClientAttach(runDir, service, authToken, sessionID, WinShmProfileHybrid)
518 + if err != nil {
519 + t.Fatalf("WinShmClientAttach failed: %v", err)
520 + }
521 + defer client.WinShmClose()
522 +
523 + buf := make([]byte, 128)
524 + if _, err := server.WinShmReceive(buf, 10); !errors.Is(err, ErrWinShmTimeout) {
525 + t.Fatalf("WinShmReceive hybrid timeout = %v, want %v", err, ErrWinShmTimeout)
526 + }
527 +}
528 +
529 +func TestWinShmReceiveTimeoutBusywait(t *testing.T) {
530 + runDir := t.TempDir()
531 + service := "timeout-busywait"
532 + const authToken uint64 = 0x8913
533 + const sessionID uint64 = 19
534 +
535 + server, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileBusywait, 4096, 4096)
536 + if err != nil {
537 + t.Fatalf("WinShmServerCreate failed: %v", err)
538 + }
539 + defer server.WinShmDestroy()
540 +
541 + client, err := WinShmClientAttach(runDir, service, authToken, sessionID, WinShmProfileBusywait)
542 + if err != nil {
543 + t.Fatalf("WinShmClientAttach failed: %v", err)
544 + }
545 + defer client.WinShmClose()
546 +
547 + buf := make([]byte, 128)
548 + if _, err := server.WinShmReceive(buf, 10); !errors.Is(err, ErrWinShmTimeout) {
549 + t.Fatalf("WinShmReceive busywait timeout = %v, want %v", err, ErrWinShmTimeout)
550 + }
551 +}
552 +
553 +func TestWinShmReceiveHybridImmediateReadyAfterSpinSkipped(t *testing.T) {
554 + runDir := t.TempDir()
555 + service := "hybrid-immediate-ready"
556 + const authToken uint64 = 0x8915
557 + const sessionID uint64 = 25
558 +
559 + server, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileHybrid, 4096, 4096)
560 + if err != nil {
561 + t.Fatalf("WinShmServerCreate failed: %v", err)
562 + }
563 + defer server.WinShmDestroy()
564 +
565 + client, err := WinShmClientAttach(runDir, service, authToken, sessionID, WinShmProfileHybrid)
566 + if err != nil {
567 + t.Fatalf("WinShmClientAttach failed: %v", err)
568 + }
569 + defer client.WinShmClose()
570 +
571 + client.SpinTries = 0
572 +
573 + msg := []byte("hybrid-ready")
574 + if err := server.WinShmSend(msg); err != nil {
575 + t.Fatalf("WinShmSend failed: %v", err)
576 + }
577 +
578 + buf := make([]byte, 8192)
579 + n, err := client.WinShmReceive(buf, 1000)
580 + if err != nil {
581 + t.Fatalf("WinShmReceive failed: %v", err)
582 + }
583 + if got := string(buf[:n]); got != string(msg) {
584 + t.Fatalf("payload = %q, want %q", got, string(msg))
585 + }
586 +}
587 +
588 +func TestWinShmReceiveBusywaitImmediateReadyAfterSpinSkipped(t *testing.T) {
589 + runDir := t.TempDir()
590 + service := "busywait-immediate-ready"
591 + const authToken uint64 = 0x8916
592 + const sessionID uint64 = 27
593 +
594 + server, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileBusywait, 4096, 4096)
595 + if err != nil {
596 + t.Fatalf("WinShmServerCreate failed: %v", err)
597 + }
598 + defer server.WinShmDestroy()
599 +
600 + client, err := WinShmClientAttach(runDir, service, authToken, sessionID, WinShmProfileBusywait)
601 + if err != nil {
602 + t.Fatalf("WinShmClientAttach failed: %v", err)
603 + }
604 + defer client.WinShmClose()
605 +
606 + client.SpinTries = 0
607 +
608 + msg := []byte("busywait-ready")
609 + if err := server.WinShmSend(msg); err != nil {
610 + t.Fatalf("WinShmSend failed: %v", err)
611 + }
612 +
613 + buf := make([]byte, 8192)
614 + n, err := client.WinShmReceive(buf, 1000)
615 + if err != nil {
616 + t.Fatalf("WinShmReceive failed: %v", err)
617 + }
618 + if got := string(buf[:n]); got != string(msg) {
619 + t.Fatalf("payload = %q, want %q", got, string(msg))
620 + }
621 +}
622 +
623 +func TestWinShmReceiveBusywaitDetectsPeerClosed(t *testing.T) {
624 + runDir := t.TempDir()
625 + service := "busywait-peer-closed"
626 + const authToken uint64 = 0x8914
627 + const sessionID uint64 = 21
628 +
629 + server, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileBusywait, 4096, 4096)
630 + if err != nil {
631 + t.Fatalf("WinShmServerCreate failed: %v", err)
632 + }
633 + defer server.WinShmDestroy()
634 +
635 + client, err := WinShmClientAttach(runDir, service, authToken, sessionID, WinShmProfileBusywait)
636 + if err != nil {
637 + t.Fatalf("WinShmClientAttach failed: %v", err)
638 + }
639 +
640 + results := make(chan error, 1)
641 + go func() {
642 + buf := make([]byte, 128)
643 + _, err := server.WinShmReceive(buf, 1000)
644 + results <- err
645 + }()
646 +
647 + time.Sleep(20 * time.Millisecond)
648 + client.WinShmClose()
649 +
650 + select {
651 + case err := <-results:
652 + if !errors.Is(err, ErrWinShmDisconnected) {
653 + t.Fatalf("WinShmReceive busywait after peer close = %v, want ErrWinShmDisconnected", err)
654 + }
655 + case <-time.After(2 * time.Second):
656 + t.Fatal("WinShmReceive busywait did not observe peer close")
657 + }
658 +}
659 +
660 +func TestWinShmReceiveNullContext(t *testing.T) {
661 + var ctx WinShmContext
662 + buf := make([]byte, 16)
663 + if _, err := ctx.WinShmReceive(buf, 1); !errors.Is(err, ErrWinShmBadParam) {
664 + t.Fatalf("WinShmReceive on null context = %v, want ErrWinShmBadParam", err)
665 + }
666 +}
667 +
668 +func TestWinShmReceiveIgnoresSpuriousWakeClient(t *testing.T) {
669 + testWinShmReceiveIgnoresSpuriousWake(t, false)
670 +}
671 +
672 +func TestWinShmReceiveIgnoresSpuriousWakeServer(t *testing.T) {
673 + testWinShmReceiveIgnoresSpuriousWake(t, true)
674 +}
675 +
676 +func testWinShmReceiveIgnoresSpuriousWake(t *testing.T, serverReceives bool) {
677 + t.Helper()
678 +
679 + runDir := t.TempDir()
680 + service := "spurious-wake"
681 + const authToken uint64 = 0x1234
682 + const sessionID uint64 = 7
683 +
684 + server, err := WinShmServerCreate(runDir, service, authToken, sessionID, WinShmProfileHybrid, 4096, 4096)
685 + if err != nil {
686 + t.Fatalf("WinShmServerCreate failed: %v", err)
687 + }
688 + defer server.WinShmDestroy()
689 +
690 + client, err := WinShmClientAttach(runDir, service, authToken, sessionID, WinShmProfileHybrid)
691 + if err != nil {
692 + t.Fatalf("WinShmClientAttach failed: %v", err)
693 + }
694 + defer client.WinShmClose()
695 +
696 + first := []byte("first-message")
697 + second := []byte("second-message")
698 +
699 + var sender *WinShmContext
700 + var receiver *WinShmContext
701 + var waitingOff int
702 + var waitEvent syscall.Handle
703 +
704 + if serverReceives {
705 + sender = client
706 + receiver = server
707 + waitingOff = wshOFFReqServerWaiting
708 + waitEvent = server.reqEvent
709 + } else {
710 + sender = server
711 + receiver = client
712 + waitingOff = wshOFFRespClientWaiting
713 + waitEvent = client.respEvent
714 + }
715 +
716 + if err := sender.WinShmSend(first); err != nil {
717 + t.Fatalf("first WinShmSend failed: %v", err)
718 + }
719 +
720 + firstBuf := make([]byte, 128)
721 + firstLen, err := receiver.WinShmReceive(firstBuf, 1000)
722 + if err != nil {
723 + t.Fatalf("first WinShmReceive failed: %v", err)
724 + }
725 + if got := string(firstBuf[:firstLen]); got != string(first) {
726 + t.Fatalf("first payload = %q, want %q", got, string(first))
727 + }
728 +
729 + results := make(chan shmReceiveResult, 1)
730 + go func() {
731 + buf := make([]byte, 128)
732 + n, err := receiver.WinShmReceive(buf, 1000)
733 + if err != nil {
734 + results <- shmReceiveResult{err: err}
735 + return
736 + }
737 + results <- shmReceiveResult{payload: append([]byte(nil), buf[:n]...)}
738 + }()
739 +
740 + data := unsafe.Slice((*byte)(unsafe.Pointer(receiver.base)), receiver.size)
741 + deadline := time.Now().Add(time.Second)
742 + for atomic.LoadInt32((*int32)(unsafe.Pointer(&data[waitingOff]))) == 0 {
743 + if time.Now().After(deadline) {
744 + t.Fatal("receiver never entered the wait state")
745 + }
746 + time.Sleep(time.Millisecond)
747 + }
748 +
749 + if ret, _, _ := procSetEvent.Call(uintptr(waitEvent)); ret == 0 {
750 + t.Fatal("SetEvent failed for spurious wake probe")
751 + }
752 +
753 + time.Sleep(10 * time.Millisecond)
754 +
755 + if err := sender.WinShmSend(second); err != nil {
756 + t.Fatalf("second WinShmSend failed: %v", err)
757 + }
758 +
759 + select {
760 + case res := <-results:
761 + if res.err != nil {
762 + t.Fatalf("second WinShmReceive failed: %v", res.err)
763 + }
764 + if got := string(res.payload); got != string(second) {
765 + t.Fatalf("second payload = %q, want %q", got, string(second))
766 + }
767 + case <-time.After(2 * time.Second):
768 + t.Fatal("second WinShmReceive timed out")
769 + }
770 +}
src/libnetdata/netipc/include/netipc/netipc_named_pipe.h new
+280
@@ -0,0 +1,280 @@
1 +/*
2 + * netipc_named_pipe.h - L1 Windows Named Pipe transport.
3 + *
4 + * Connection lifecycle, handshake, send/receive with transparent chunking
5 + * over Win32 Named Pipes in message mode. Uses the wire envelope from
6 + * netipc_protocol.h for all framing.
7 + *
8 + * Pipe name derivation:
9 + * \\.\pipe\netipc-{FNV1a64(run_dir):016llx}-{service_name}
10 + */
11 +
12 +#ifndef NETIPC_NAMED_PIPE_H
13 +#define NETIPC_NAMED_PIPE_H
14 +
15 +#if defined(_WIN32) || defined(__MSYS__)
16 +
17 +#include "netipc_protocol.h"
18 +#include <stdbool.h>
19 +#include <stddef.h>
20 +#include <stdint.h>
21 +#include <windows.h>
22 +
23 +#ifdef __cplusplus
24 +extern "C" {
25 +#endif
26 +
27 +/* ------------------------------------------------------------------ */
28 +/* Constants */
29 +/* ------------------------------------------------------------------ */
30 +
31 +/* Maximum pipe name length (\\.\pipe\ prefix + hash + service) */
32 +#define NIPC_NP_MAX_PIPE_NAME 256
33 +
34 +/* Default pipe buffer size and packet size */
35 +#define NIPC_NP_DEFAULT_PIPE_BUF_SIZE 65536
36 +#define NIPC_NP_DEFAULT_PACKET_SIZE 65536
37 +#define NIPC_NP_DEFAULT_BATCH_ITEMS 1
38 +
39 +/* Max concurrent pipe instances */
40 +#define NIPC_NP_MAX_INSTANCES PIPE_UNLIMITED_INSTANCES
41 +
42 +/* FNV-1a 64-bit constants */
43 +#define NIPC_FNV1A_OFFSET_BASIS 0xcbf29ce484222325ull
44 +#define NIPC_FNV1A_PRIME 0x00000100000001B3ull
45 +
46 +/* ------------------------------------------------------------------ */
47 +/* Error codes (transport-level) */
48 +/* ------------------------------------------------------------------ */
49 +
50 +typedef enum {
51 + NIPC_NP_OK = 0,
52 + NIPC_NP_ERR_PIPE_NAME, /* pipe name derivation failed */
53 + NIPC_NP_ERR_CREATE_PIPE, /* CreateNamedPipeW failed */
54 + NIPC_NP_ERR_CONNECT, /* ConnectNamedPipe / CreateFileW failed */
55 + NIPC_NP_ERR_ACCEPT, /* ConnectNamedPipe failed waiting for client */
56 + NIPC_NP_ERR_SEND, /* WriteFile failed */
57 + NIPC_NP_ERR_RECV, /* ReadFile failed / peer disconnected */
58 + NIPC_NP_ERR_HANDSHAKE, /* handshake protocol error */
59 + NIPC_NP_ERR_AUTH_FAILED, /* auth token rejected */
60 + NIPC_NP_ERR_NO_PROFILE, /* no common profile */
61 + NIPC_NP_ERR_INCOMPATIBLE, /* protocol/layout version mismatch */
62 + NIPC_NP_ERR_PROTOCOL, /* wire protocol violation */
63 + NIPC_NP_ERR_ADDR_IN_USE, /* pipe name already in use */
64 + NIPC_NP_ERR_CHUNK, /* chunk header mismatch */
65 + NIPC_NP_ERR_ALLOC, /* memory allocation failed */
66 + NIPC_NP_ERR_LIMIT_EXCEEDED, /* payload/batch exceeds negotiated */
67 + NIPC_NP_ERR_BAD_PARAM, /* invalid argument */
68 + NIPC_NP_ERR_DUPLICATE_MSG_ID, /* message_id already in-flight */
69 + NIPC_NP_ERR_UNKNOWN_MSG_ID, /* response message_id not in-flight */
70 + NIPC_NP_ERR_DISCONNECTED, /* peer disconnected (graceful) */
71 +} nipc_np_error_t;
72 +
73 +/* ------------------------------------------------------------------ */
74 +/* Role */
75 +/* ------------------------------------------------------------------ */
76 +
77 +typedef enum {
78 + NIPC_NP_ROLE_CLIENT = 1,
79 + NIPC_NP_ROLE_SERVER = 2,
80 +} nipc_np_role_t;
81 +
82 +/* ------------------------------------------------------------------ */
83 +/* Client connect configuration */
84 +/* ------------------------------------------------------------------ */
85 +
86 +typedef struct {
87 + uint32_t supported_profiles;
88 + uint32_t preferred_profiles;
89 + uint32_t max_request_payload_bytes; /* 0 = use default */
90 + uint32_t max_request_batch_items; /* 0 = use default (1) */
91 + uint32_t max_response_payload_bytes; /* 0 = use default */
92 + uint32_t max_response_batch_items; /* 0 = use default (1) */
93 + uint64_t auth_token;
94 + uint32_t packet_size; /* 0 = use default (65536) */
95 +} nipc_np_client_config_t;
96 +
97 +/* ------------------------------------------------------------------ */
98 +/* Server configuration (for listen + accept) */
99 +/* ------------------------------------------------------------------ */
100 +
101 +typedef struct {
102 + uint32_t supported_profiles;
103 + uint32_t preferred_profiles;
104 + uint32_t max_request_payload_bytes; /* 0 = use default */
105 + uint32_t max_request_batch_items; /* 0 = use default (1) */
106 + uint32_t max_response_payload_bytes; /* 0 = use default */
107 + uint32_t max_response_batch_items; /* 0 = use default (1) */
108 + uint64_t auth_token;
109 + uint32_t packet_size; /* 0 = use default (65536) */
110 +} nipc_np_server_config_t;
111 +
112 +/* ------------------------------------------------------------------ */
113 +/* Session */
114 +/* ------------------------------------------------------------------ */
115 +
116 +typedef struct {
117 + HANDLE pipe; /* connected pipe handle (native wait object) */
118 + nipc_np_role_t role;
119 +
120 + /* Negotiated limits */
121 + uint32_t max_request_payload_bytes;
122 + uint32_t max_request_batch_items;
123 + uint32_t max_response_payload_bytes;
124 + uint32_t max_response_batch_items;
125 + uint32_t packet_size;
126 + uint32_t selected_profile;
127 + uint64_t session_id; /* server-assigned, for per-session SHM */
128 +
129 + /* Internal receive buffer for chunked reassembly */
130 + uint8_t *recv_buf;
131 + size_t recv_buf_size;
132 +
133 + /* In-flight message_id set (client-side only, dynamically grown) */
134 + uint64_t *inflight_ids;
135 + uint32_t inflight_count;
136 + uint32_t inflight_capacity;
137 +} nipc_np_session_t;
138 +
139 +/* ------------------------------------------------------------------ */
140 +/* Listener */
141 +/* ------------------------------------------------------------------ */
142 +
143 +typedef struct {
144 + HANDLE pipe; /* current listening pipe instance */
145 + nipc_np_server_config_t config;
146 + wchar_t pipe_name[NIPC_NP_MAX_PIPE_NAME]; /* stored for new instances */
147 +} nipc_np_listener_t;
148 +
149 +/* ------------------------------------------------------------------ */
150 +/* FNV-1a 64-bit hash */
151 +/* ------------------------------------------------------------------ */
152 +
153 +/* Compute FNV-1a 64-bit hash of data[0..len). */
154 +uint64_t nipc_fnv1a_64(const void *data, size_t len);
155 +
156 +/* ------------------------------------------------------------------ */
157 +/* Pipe name derivation */
158 +/* ------------------------------------------------------------------ */
159 +
160 +/*
161 + * Build pipe name into dst (wide string).
162 + * Format: \\.\pipe\netipc-{FNV1a64(run_dir):016llx}-{service_name}
163 + * Returns 0 on success, -1 on error.
164 + */
165 +int nipc_np_build_pipe_name(wchar_t *dst, size_t dst_chars,
166 + const char *run_dir,
167 + const char *service_name);
168 +
169 +/* ------------------------------------------------------------------ */
170 +/* Connection lifecycle */
171 +/* ------------------------------------------------------------------ */
172 +
173 +/*
174 + * Create a listener on a Named Pipe derived from run_dir + service_name.
175 + * Creates the first pipe instance and prepares for client connections.
176 + */
177 +nipc_np_error_t nipc_np_listen(const char *run_dir,
178 + const char *service_name,
179 + const nipc_np_server_config_t *config,
180 + nipc_np_listener_t *out);
181 +
182 +/*
183 + * Accept one client on a listener. Performs the full handshake.
184 + * session_id is placed into the hello-ack so the client can attach
185 + * to the correct per-session SHM region.
186 + * Blocks until a client connects and the handshake completes.
187 + * After accepting, creates a new pipe instance for subsequent clients.
188 + */
189 +nipc_np_error_t nipc_np_accept(nipc_np_listener_t *listener,
190 + uint64_t session_id,
191 + nipc_np_session_t *out);
192 +
193 +/*
194 + * Connect to a server pipe derived from run_dir + service_name.
195 + * Performs the full handshake. Blocks until connected + handshake done.
196 + */
197 +nipc_np_error_t nipc_np_connect(const char *run_dir,
198 + const char *service_name,
199 + const nipc_np_client_config_t *config,
200 + nipc_np_session_t *out);
201 +
202 +/*
203 + * Close a session. Releases pipe handle and internal buffers.
204 + * Safe to call on a zero-initialized session (no-op).
205 + */
206 +void nipc_np_close_session(nipc_np_session_t *session);
207 +
208 +/*
209 + * Close a listener. Closes the pipe handle.
210 + */
211 +void nipc_np_close_listener(nipc_np_listener_t *listener);
212 +
213 +/* ------------------------------------------------------------------ */
214 +/* Message send / receive */
215 +/* ------------------------------------------------------------------ */
216 +
217 +/*
218 + * Send one logical message. hdr is the 32-byte outer header (caller fills
219 + * kind, code, flags, payload_len, item_count, message_id; this function
220 + * sets magic/version/header_len). payload is the opaque payload bytes.
221 + *
222 + * If the total message (32 + payload_len) exceeds packet_size, the
223 + * message is chunked transparently.
224 + */
225 +nipc_np_error_t nipc_np_send(nipc_np_session_t *session,
226 + nipc_header_t *hdr,
227 + const void *payload,
228 + size_t payload_len);
229 +
230 +/*
231 + * Receive one logical message. Blocks until a complete message arrives.
232 + *
233 + * On success:
234 + * - hdr_out is filled with the decoded outer header.
235 + * - *payload_out points to the payload bytes (inside session->recv_buf
236 + * or buf). Valid until the next receive call.
237 + * - *payload_len_out is the payload length.
238 + *
239 + * buf/buf_size: caller-provided buffer for the first packet.
240 + */
241 +nipc_np_error_t nipc_np_receive(nipc_np_session_t *session,
242 + void *buf, size_t buf_size,
243 + nipc_header_t *hdr_out,
244 + const void **payload_out,
245 + size_t *payload_len_out);
246 +
247 +/*
248 + * Poll until a session becomes readable or the timeout expires.
249 + *
250 + * On success:
251 + * - returns NIPC_NP_OK
252 + * - sets *readable_out to true if bytes are available
253 + * - sets *readable_out to false on timeout with no pending bytes
254 + *
255 + * Returns NIPC_NP_ERR_DISCONNECTED if the peer has gone away while waiting.
256 + */
257 +nipc_np_error_t nipc_np_wait_readable(nipc_np_session_t *session,
258 + uint32_t timeout_ms,
259 + bool *readable_out);
260 +
261 +/* ------------------------------------------------------------------ */
262 +/* Utility */
263 +/* ------------------------------------------------------------------ */
264 +
265 +/* Get the HANDLE for WaitForSingleObject / WaitForMultipleObjects. */
266 +static inline HANDLE nipc_np_session_handle(const nipc_np_session_t *s) {
267 + return s->pipe;
268 +}
269 +
270 +static inline HANDLE nipc_np_listener_handle(const nipc_np_listener_t *l) {
271 + return l->pipe;
272 +}
273 +
274 +#ifdef __cplusplus
275 +}
276 +#endif
277 +
278 +#endif /* _WIN32 || __MSYS__ */
279 +
280 +#endif /* NETIPC_NAMED_PIPE_H */
src/libnetdata/netipc/include/netipc/netipc_protocol.h new
+517
@@ -0,0 +1,517 @@
1 +/*
2 + * netipc_protocol.h - Wire envelope and codec for the netipc protocol.
3 + *
4 + * Pure byte-layout encode/decode. No I/O, no transport, no allocation on
5 + * decode. Localhost-only IPC — all multi-byte fields use host byte order.
6 + * Struct layouts match wire format exactly; encode/decode is a single memcpy
7 + * plus validation.
8 + *
9 + * Decoded "View" types borrow the underlying buffer and are valid only while
10 + * that buffer lives. Copy immediately if the data is needed later.
11 + */
12 +
13 +#ifndef NETIPC_PROTOCOL_H
14 +#define NETIPC_PROTOCOL_H
15 +
16 +#include <stdbool.h>
17 +#include <stddef.h>
18 +#include <stdint.h>
19 +#include <string.h>
20 +
21 +#ifdef __cplusplus
22 +extern "C" {
23 +#endif
24 +
25 +/* ------------------------------------------------------------------ */
26 +/* Constants */
27 +/* ------------------------------------------------------------------ */
28 +
29 +#define NIPC_MAGIC_MSG 0x4e495043u /* "NIPC" */
30 +#define NIPC_MAGIC_CHUNK 0x4e43484bu /* "NCHK" */
31 +#define NIPC_VERSION 1u
32 +#define NIPC_HEADER_LEN 32u
33 +
34 +/* Message kinds */
35 +#define NIPC_KIND_REQUEST 1u
36 +#define NIPC_KIND_RESPONSE 2u
37 +#define NIPC_KIND_CONTROL 3u
38 +
39 +/* Flags */
40 +#define NIPC_FLAG_BATCH 0x0001u
41 +
42 +/* Transport status */
43 +#define NIPC_STATUS_OK 0u
44 +#define NIPC_STATUS_BAD_ENVELOPE 1u
45 +#define NIPC_STATUS_AUTH_FAILED 2u
46 +#define NIPC_STATUS_INCOMPATIBLE 3u
47 +#define NIPC_STATUS_UNSUPPORTED 4u
48 +#define NIPC_STATUS_LIMIT_EXCEEDED 5u
49 +#define NIPC_STATUS_INTERNAL_ERROR 6u
50 +
51 +/* Control opcodes */
52 +#define NIPC_CODE_HELLO 1u
53 +#define NIPC_CODE_HELLO_ACK 2u
54 +
55 +/* Method codes */
56 +#define NIPC_METHOD_INCREMENT 1u
57 +#define NIPC_METHOD_CGROUPS_SNAPSHOT 2u
58 +#define NIPC_METHOD_STRING_REVERSE 3u
59 +
60 +/* Profile bits */
61 +#define NIPC_PROFILE_BASELINE 0x01u
62 +#define NIPC_PROFILE_SHM_HYBRID 0x02u
63 +#define NIPC_PROFILE_SHM_FUTEX 0x04u
64 +#define NIPC_PROFILE_SHM_WAITADDR 0x08u
65 +
66 +/* Defaults */
67 +#define NIPC_MAX_PAYLOAD_DEFAULT 1024u
68 +
69 +/* Hard cap on negotiated request payload sizes (1 MiB) — prevents a
70 + * compromised peer from forcing excessive memory allocation. */
71 +#define NIPC_MAX_PAYLOAD_CAP (1024u * 1024u)
72 +
73 +/* Alignment for batch items and cgroups items */
74 +#define NIPC_ALIGNMENT 8u
75 +
76 +/* ------------------------------------------------------------------ */
77 +/* Error codes */
78 +/* ------------------------------------------------------------------ */
79 +
80 +typedef enum {
81 + NIPC_OK = 0,
82 + NIPC_ERR_TRUNCATED, /* buffer too short for the expected structure */
83 + NIPC_ERR_BAD_MAGIC, /* magic value mismatch */
84 + NIPC_ERR_BAD_VERSION, /* unsupported version */
85 + NIPC_ERR_BAD_HEADER_LEN, /* header_len != 32 */
86 + NIPC_ERR_BAD_KIND, /* unknown message kind */
87 + NIPC_ERR_BAD_LAYOUT, /* unknown layout_version in a payload */
88 + NIPC_ERR_OUT_OF_BOUNDS, /* offset+length exceeds available data */
89 + NIPC_ERR_MISSING_NUL, /* string not NUL-terminated */
90 + NIPC_ERR_BAD_ALIGNMENT, /* item not 8-byte aligned */
91 + NIPC_ERR_BAD_ITEM_COUNT, /* directory inconsistent with payload size */
92 + NIPC_ERR_OVERFLOW, /* builder ran out of space */
93 + NIPC_ERR_HANDLER_FAILED, /* typed handler rejected an otherwise valid request */
94 + NIPC_ERR_NOT_READY, /* client not connected / service unavailable */
95 +} nipc_error_t;
96 +
97 +/* ------------------------------------------------------------------ */
98 +/* Outer message header (32 bytes) */
99 +/* ------------------------------------------------------------------ */
100 +
101 +typedef struct {
102 + uint32_t magic;
103 + uint16_t version;
104 + uint16_t header_len;
105 + uint16_t kind;
106 + uint16_t flags;
107 + uint16_t code;
108 + uint16_t transport_status;
109 + uint32_t payload_len;
110 + uint32_t item_count;
111 + uint64_t message_id;
112 +} nipc_header_t;
113 +
114 +/* Encode header into buf (must be >= 32 bytes). Returns 32. */
115 +size_t nipc_header_encode(const nipc_header_t *hdr, void *buf, size_t buf_len);
116 +
117 +/* Decode header from buf. Returns NIPC_OK or an error. */
118 +nipc_error_t nipc_header_decode(const void *buf, size_t buf_len,
119 + nipc_header_t *out);
120 +
121 +/* ------------------------------------------------------------------ */
122 +/* Chunk continuation header (32 bytes) */
123 +/* ------------------------------------------------------------------ */
124 +
125 +typedef struct {
126 + uint32_t magic;
127 + uint16_t version;
128 + uint16_t flags;
129 + uint64_t message_id;
130 + uint32_t total_message_len;
131 + uint32_t chunk_index;
132 + uint32_t chunk_count;
133 + uint32_t chunk_payload_len;
134 +} nipc_chunk_header_t;
135 +
136 +size_t nipc_chunk_header_encode(const nipc_chunk_header_t *chk,
137 + void *buf, size_t buf_len);
138 +
139 +nipc_error_t nipc_chunk_header_decode(const void *buf, size_t buf_len,
140 + nipc_chunk_header_t *out);
141 +
142 +/* ------------------------------------------------------------------ */
143 +/* Batch item directory */
144 +/* ------------------------------------------------------------------ */
145 +
146 +typedef struct {
147 + uint32_t offset;
148 + uint32_t length;
149 +} nipc_batch_entry_t;
150 +
151 +/*
152 + * Encode item_count directory entries into buf.
153 + * Returns total bytes written (item_count * 8).
154 + */
155 +size_t nipc_batch_dir_encode(const nipc_batch_entry_t *entries,
156 + uint32_t item_count,
157 + void *buf, size_t buf_len);
158 +
159 +/*
160 + * Decode and validate item_count directory entries from buf.
161 + * packed_area_len is the size of the packed item area that follows
162 + * the directory. Each entry's offset+length must fall within it.
163 + * Returns NIPC_OK or an error.
164 + */
165 +nipc_error_t nipc_batch_dir_decode(const void *buf, size_t buf_len,
166 + uint32_t item_count,
167 + uint32_t packed_area_len,
168 + nipc_batch_entry_t *out);
169 +
170 +/*
171 + * Validate a batch directory without allocating an output array.
172 + * Checks alignment and bounds for each entry. For use in L1 receive
173 + * paths where allocation is undesirable.
174 + * buf/buf_len: the directory bytes (item_count * 8).
175 + * packed_area_len: size of the packed item area after the directory.
176 + */
177 +nipc_error_t nipc_batch_dir_validate(const void *buf, size_t buf_len,
178 + uint32_t item_count,
179 + uint32_t packed_area_len);
180 +
181 +/*
182 + * Extract a single batch item by index from a complete batch payload.
183 + * payload points to the first byte after the outer header.
184 + * On success, *item_ptr and *item_len are set.
185 + */
186 +nipc_error_t nipc_batch_item_get(const void *payload, size_t payload_len,
187 + uint32_t item_count, uint32_t index,
188 + const void **item_ptr, uint32_t *item_len);
189 +
190 +/* ------------------------------------------------------------------ */
191 +/* Batch builder */
192 +/* ------------------------------------------------------------------ */
193 +
194 +typedef struct {
195 + uint8_t *buf; /* caller-owned output buffer */
196 + size_t buf_len; /* total buffer capacity */
197 + uint32_t item_count; /* items added so far */
198 + uint32_t max_items; /* max items (from directory capacity) */
199 + size_t dir_end; /* byte offset where directory ends */
200 + size_t data_offset; /* current write position in packed area */
201 +} nipc_batch_builder_t;
202 +
203 +/*
204 + * Initialize a batch builder.
205 + * buf must be large enough for: max_items*8 (directory) + packed data.
206 + * The directory is written at the front; packed data grows after it.
207 + */
208 +void nipc_batch_builder_init(nipc_batch_builder_t *b,
209 + void *buf, size_t buf_len,
210 + uint32_t max_items);
211 +
212 +/* Add an item payload. Returns NIPC_OK or NIPC_ERR_OVERFLOW. */
213 +nipc_error_t nipc_batch_builder_add(nipc_batch_builder_t *b,
214 + const void *item, size_t item_len);
215 +
216 +/*
217 + * Finalize: writes the directory entries and returns the total payload
218 + * size (directory + packed items). Sets *item_count_out.
219 + */
220 +size_t nipc_batch_builder_finish(nipc_batch_builder_t *b,
221 + uint32_t *item_count_out);
222 +
223 +/* ------------------------------------------------------------------ */
224 +/* Hello payload (44 bytes) */
225 +/* ------------------------------------------------------------------ */
226 +
227 +typedef struct {
228 + uint16_t layout_version;
229 + uint16_t flags;
230 + uint32_t supported_profiles;
231 + uint32_t preferred_profiles;
232 + uint32_t max_request_payload_bytes;
233 + uint32_t max_request_batch_items;
234 + uint32_t max_response_payload_bytes;
235 + uint32_t max_response_batch_items;
236 + uint32_t _reserved; /* wire offset 28, must be 0 */
237 + uint64_t auth_token;
238 + uint32_t packet_size;
239 +} nipc_hello_t;
240 +
241 +/* Wire size (44) differs from sizeof (48) due to trailing alignment. */
242 +#define NIPC_HELLO_WIRE_SIZE 44u
243 +
244 +size_t nipc_hello_encode(const nipc_hello_t *h, void *buf, size_t buf_len);
245 +
246 +nipc_error_t nipc_hello_decode(const void *buf, size_t buf_len,
247 + nipc_hello_t *out);
248 +
249 +/* ------------------------------------------------------------------ */
250 +/* Hello-ack payload (48 bytes) */
251 +/* ------------------------------------------------------------------ */
252 +
253 +typedef struct {
254 + uint16_t layout_version;
255 + uint16_t flags;
256 + uint32_t server_supported_profiles;
257 + uint32_t intersection_profiles;
258 + uint32_t selected_profile;
259 + uint32_t agreed_max_request_payload_bytes;
260 + uint32_t agreed_max_request_batch_items;
261 + uint32_t agreed_max_response_payload_bytes;
262 + uint32_t agreed_max_response_batch_items;
263 + uint32_t agreed_packet_size;
264 + uint32_t _reserved; /* wire offset 36, must be 0 */
265 + uint64_t session_id; /* server-assigned, for per-session SHM path */
266 +} nipc_hello_ack_t;
267 +
268 +size_t nipc_hello_ack_encode(const nipc_hello_ack_t *h,
269 + void *buf, size_t buf_len);
270 +
271 +nipc_error_t nipc_hello_ack_decode(const void *buf, size_t buf_len,
272 + nipc_hello_ack_t *out);
273 +
274 +/* ------------------------------------------------------------------ */
275 +/* Cgroups snapshot request (4 bytes) */
276 +/* ------------------------------------------------------------------ */
277 +
278 +typedef struct {
279 + uint16_t layout_version;
280 + uint16_t flags;
281 +} nipc_cgroups_req_t;
282 +
283 +size_t nipc_cgroups_req_encode(const nipc_cgroups_req_t *r,
284 + void *buf, size_t buf_len);
285 +
286 +nipc_error_t nipc_cgroups_req_decode(const void *buf, size_t buf_len,
287 + nipc_cgroups_req_t *out);
288 +
289 +/* ------------------------------------------------------------------ */
290 +/* Cgroups snapshot response */
291 +/* ------------------------------------------------------------------ */
292 +
293 +/* Snapshot-level header (24 bytes) */
294 +typedef struct {
295 + uint16_t layout_version;
296 + uint16_t flags;
297 + uint32_t item_count;
298 + uint32_t systemd_enabled;
299 + uint32_t reserved;
300 + uint64_t generation;
301 +} nipc_cgroups_resp_header_t;
302 +
303 +/* Borrowed string view into the payload buffer. */
304 +typedef struct {
305 + const char *ptr; /* points into payload, NUL-terminated */
306 + uint32_t len; /* length excluding the NUL */
307 +} nipc_str_view_t;
308 +
309 +/*
310 + * Per-item view -- ephemeral, borrows the payload buffer.
311 + * Valid only while the payload buffer is alive.
312 + */
313 +typedef struct {
314 + uint16_t layout_version;
315 + uint16_t flags;
316 + uint32_t hash;
317 + uint32_t options;
318 + uint32_t enabled;
319 + nipc_str_view_t name;
320 + nipc_str_view_t path;
321 +} nipc_cgroups_item_view_t;
322 +
323 +/* Full snapshot view -- ephemeral. */
324 +typedef struct {
325 + uint16_t layout_version;
326 + uint16_t flags;
327 + uint32_t item_count;
328 + uint32_t systemd_enabled;
329 + uint64_t generation;
330 +
331 + /* Internal: payload pointer and size for item access */
332 + const uint8_t *_payload;
333 + size_t _payload_len;
334 +} nipc_cgroups_resp_view_t;
335 +
336 +/*
337 + * Decode the snapshot response header and validate the item directory.
338 + * On success, use nipc_cgroups_resp_item() to access individual items.
339 + */
340 +nipc_error_t nipc_cgroups_resp_decode(const void *buf, size_t buf_len,
341 + nipc_cgroups_resp_view_t *out);
342 +
343 +/*
344 + * Access item at index from a decoded snapshot view.
345 + * index must be < view->item_count.
346 + */
347 +nipc_error_t nipc_cgroups_resp_item(const nipc_cgroups_resp_view_t *view,
348 + uint32_t index,
349 + nipc_cgroups_item_view_t *out);
350 +
351 +/* ------------------------------------------------------------------ */
352 +/* Cgroups snapshot response builder */
353 +/* ------------------------------------------------------------------ */
354 +
355 +#define NIPC_CGROUPS_ITEM_HDR_SIZE 32u
356 +#define NIPC_CGROUPS_RESP_HDR_SIZE 24u
357 +#define NIPC_CGROUPS_DIR_ENTRY_SIZE 8u
358 +
359 +typedef struct {
360 + uint8_t *buf;
361 + size_t buf_len;
362 + uint32_t systemd_enabled;
363 + uint64_t generation;
364 + uint32_t item_count;
365 + uint32_t max_items; /* directory slots reserved at init */
366 + nipc_error_t error; /* sticky builder failure for dispatch */
367 +
368 + /* Current write position for packed item data (absolute). */
369 + size_t data_offset;
370 +} nipc_cgroups_builder_t;
371 +
372 +/*
373 + * Initialize the builder. buf must be caller-owned and large enough
374 + * for the expected snapshot. max_items is a hint for directory space
375 + * reservation.
376 + */
377 +void nipc_cgroups_builder_init(nipc_cgroups_builder_t *b,
378 + void *buf, size_t buf_len,
379 + uint32_t max_items,
380 + uint32_t systemd_enabled,
381 + uint64_t generation);
382 +
383 +/* Update the snapshot header fields that finish() writes. */
384 +void nipc_cgroups_builder_set_header(nipc_cgroups_builder_t *b,
385 + uint32_t systemd_enabled,
386 + uint64_t generation);
387 +
388 +/* Return a safe upper bound for the number of snapshot items that can
389 + * fit in a response buffer of size buf_len. This is for directory
390 + * reservation only, not a promise for arbitrary string sizes. */
391 +uint32_t nipc_cgroups_builder_estimate_max_items(size_t buf_len);
392 +
393 +/*
394 + * Add one cgroup item. The builder handles offset bookkeeping,
395 + * NUL termination, and alignment.
396 + */
397 +nipc_error_t nipc_cgroups_builder_add(nipc_cgroups_builder_t *b,
398 + uint32_t hash,
399 + uint32_t options,
400 + uint32_t enabled,
401 + const char *name, uint32_t name_len,
402 + const char *path, uint32_t path_len);
403 +
404 +/*
405 + * Finalize the builder. Writes the snapshot header and returns the
406 + * total payload size. The buffer now contains a complete, decodable
407 + * cgroups snapshot response payload.
408 + */
409 +size_t nipc_cgroups_builder_finish(nipc_cgroups_builder_t *b);
410 +
411 +/* ------------------------------------------------------------------ */
412 +/* INCREMENT codec (8 bytes) */
413 +/* ------------------------------------------------------------------ */
414 +
415 +/*
416 + * INCREMENT payload: { uint64_t value }.
417 + * Same layout for both request and response.
418 + */
419 +#define NIPC_INCREMENT_PAYLOAD_SIZE 8u
420 +
421 +/* Encode value into buf. Returns 8 on success, 0 if buf too small. */
422 +size_t nipc_increment_encode(uint64_t value, void *buf, size_t buf_len);
423 +
424 +/* Decode value from buf. Returns NIPC_OK or NIPC_ERR_TRUNCATED. */
425 +nipc_error_t nipc_increment_decode(const void *buf, size_t buf_len,
426 + uint64_t *value_out);
427 +
428 +/* ------------------------------------------------------------------ */
429 +/* STRING_REVERSE codec (variable length) */
430 +/* ------------------------------------------------------------------ */
431 +
432 +/*
433 + * STRING_REVERSE payload wire layout:
434 + * | 0 | 4 | u32 | str_offset (from payload start, always 8) |
435 + * | 4 | 4 | u32 | str_length (excluding NUL) |
436 + * | 8 | N+1 | bytes | string data + NUL |
437 + *
438 + * Same layout for both request and response.
439 + * Total size = 8 + str_length + 1.
440 + */
441 +#define NIPC_STRING_REVERSE_HDR_SIZE 8u
442 +
443 +/* Ephemeral view into a decoded STRING_REVERSE payload.
444 + * Borrows the payload buffer — valid only during the current call. */
445 +typedef struct {
446 + const char *str; /* pointer into payload, NUL-terminated */
447 + uint32_t str_len; /* length excluding NUL */
448 +} nipc_string_reverse_view_t;
449 +
450 +/* Encode str (str_len bytes) into buf. Returns total bytes written,
451 + * or 0 if buf too small. Appends trailing NUL. */
452 +size_t nipc_string_reverse_encode(const char *str, uint32_t str_len,
453 + void *buf, size_t buf_len);
454 +
455 +/* Decode payload into an ephemeral view. Validates bounds and NUL.
456 + * Returns NIPC_OK or error. */
457 +nipc_error_t nipc_string_reverse_decode(const void *buf, size_t buf_len,
458 + nipc_string_reverse_view_t *view_out);
459 +
460 +/* ------------------------------------------------------------------ */
461 +/* Server-side typed dispatch helpers */
462 +/* ------------------------------------------------------------------ */
463 +
464 +/*
465 + * Per-method dispatch: decode request → call typed handler → encode
466 + * response. Handlers never touch wire format — pure business logic.
467 + *
468 + * Each helper takes (raw_request, raw_response_buf, handler_fn, user).
469 + * Returns true on success (response written), false on failure.
470 + */
471 +
472 +/* INCREMENT: handler receives decoded u64, returns u64. */
473 +typedef bool (*nipc_increment_handler_fn)(
474 + void *user, uint64_t request, uint64_t *response);
475 +
476 +bool nipc_dispatch_increment(
477 + const uint8_t *req, size_t req_len,
478 + uint8_t *resp, size_t resp_size, size_t *resp_len,
479 + nipc_increment_handler_fn handler, void *user);
480 +
481 +/* STRING_REVERSE: handler receives decoded string, writes response string. */
482 +typedef bool (*nipc_string_reverse_handler_fn)(
483 + void *user,
484 + const char *request_str, uint32_t request_str_len,
485 + char *response_str, uint32_t response_capacity,
486 + uint32_t *response_str_len);
487 +
488 +bool nipc_dispatch_string_reverse(
489 + const uint8_t *req, size_t req_len,
490 + uint8_t *resp, size_t resp_size, size_t *resp_len,
491 + nipc_string_reverse_handler_fn handler, void *user);
492 +
493 +/* CGROUPS_SNAPSHOT: handler receives decoded request, fills builder. */
494 +typedef bool (*nipc_cgroups_handler_fn)(
495 + void *user,
496 + const nipc_cgroups_req_t *request,
497 + nipc_cgroups_builder_t *builder);
498 +
499 +nipc_error_t nipc_dispatch_cgroups_snapshot(
500 + const uint8_t *req, size_t req_len,
501 + uint8_t *resp, size_t resp_size, size_t *resp_len,
502 + uint32_t max_items,
503 + nipc_cgroups_handler_fn handler, void *user);
504 +
505 +/* ------------------------------------------------------------------ */
506 +/* Utility: 8-byte alignment */
507 +/* ------------------------------------------------------------------ */
508 +
509 +static inline size_t nipc_align8(size_t v) {
510 + return (v + 7u) & ~(size_t)7u;
511 +}
512 +
513 +#ifdef __cplusplus
514 +}
515 +#endif
516 +
517 +#endif /* NETIPC_PROTOCOL_H */
src/libnetdata/netipc/include/netipc/netipc_service.h new
+547
@@ -0,0 +1,547 @@
1 +/*
2 + * netipc_service.h - L2 orchestration: client context and managed server.
3 + *
4 + * Pure convenience layer. Uses L1 transport + Codec exclusively.
5 + * Adds zero wire behavior. Provides lifecycle management, typed
6 + * cgroups-snapshot calls, and managed multi-client worker dispatch for
7 + * one service kind per endpoint.
8 + *
9 + * The public L2 contract is service-oriented:
10 + * - clients connect to a service kind, not a plugin identity
11 + * - one server endpoint serves one request kind only
12 + * - outer request codes remain part of the envelope for validation,
13 + * not public multi-method dispatch
14 + *
15 + * L2 callers never see transports, handshakes, or chunking.
16 + *
17 + * Platform-specific transport types are selected at compile time:
18 + * POSIX: UDS + SHM (netipc_uds.h, netipc_shm.h)
19 + * Windows: Named Pipe + Win SHM (netipc_named_pipe.h, netipc_win_shm.h)
20 + */
21 +
22 +#ifndef NETIPC_SERVICE_H
23 +#define NETIPC_SERVICE_H
24 +
25 +#include "netipc_protocol.h"
26 +
27 +#if defined(_WIN32) || defined(__MSYS__)
28 +#include "netipc_named_pipe.h"
29 +#include "netipc_win_shm.h"
30 +#include <windows.h>
31 +#else
32 +#include "netipc_uds.h"
33 +#include "netipc_shm.h"
34 +#include <pthread.h>
35 +#endif
36 +
37 +#include <stdbool.h>
38 +#include <stddef.h>
39 +#include <stdint.h>
40 +
41 +#ifdef __cplusplus
42 +extern "C" {
43 +#endif
44 +
45 +/* ------------------------------------------------------------------ */
46 +/* Client context state */
47 +/* ------------------------------------------------------------------ */
48 +
49 +typedef enum {
50 + NIPC_CLIENT_DISCONNECTED = 0,
51 + NIPC_CLIENT_CONNECTING,
52 + NIPC_CLIENT_READY,
53 + NIPC_CLIENT_NOT_FOUND,
54 + NIPC_CLIENT_AUTH_FAILED,
55 + NIPC_CLIENT_INCOMPATIBLE,
56 + NIPC_CLIENT_BROKEN,
57 +} nipc_client_state_t;
58 +
59 +/* ------------------------------------------------------------------ */
60 +/* Client status snapshot (for diagnostics, not hot path) */
61 +/* ------------------------------------------------------------------ */
62 +
63 +typedef struct {
64 + nipc_client_state_t state;
65 + uint32_t connect_count;
66 + uint32_t reconnect_count;
67 + uint32_t call_count;
68 + uint32_t error_count;
69 +} nipc_client_status_t;
70 +
71 +/* ------------------------------------------------------------------ */
72 +/* Public L2/L3 service configuration */
73 +/* ------------------------------------------------------------------ */
74 +
75 +/*
76 + * Public service-level client configuration shared by the L2 and L3 APIs.
77 + *
78 + * This configuration is intentionally transport-agnostic. Transport-only
79 + * tuning such as packet sizing stays below the public L2/L3 boundary.
80 + *
81 + * Zero-valued fields map to the normal library defaults.
82 + */
83 +typedef struct {
84 + uint32_t supported_profiles;
85 + uint32_t preferred_profiles;
86 + uint32_t max_request_batch_items;
87 + uint32_t max_response_payload_bytes;
88 + uint64_t auth_token;
89 +} nipc_client_config_t;
90 +
91 +/*
92 + * Public service-level server configuration shared by the typed L2 API.
93 + *
94 + * Transport-only tuning such as socket backlog or packet sizing stays in the
95 + * transport layer and is not part of the public L2/L3 contract.
96 + *
97 + * Zero-valued fields map to the normal library defaults.
98 + */
99 +typedef struct {
100 + uint32_t supported_profiles;
101 + uint32_t preferred_profiles;
102 + uint32_t max_request_batch_items;
103 + uint32_t max_response_payload_bytes;
104 + uint64_t auth_token;
105 +} nipc_server_config_t;
106 +
107 +/* ------------------------------------------------------------------ */
108 +/* Client context */
109 +/* ------------------------------------------------------------------ */
110 +
111 +typedef struct {
112 + /* State */
113 + nipc_client_state_t state;
114 +
115 + /* Configuration (mutable learned sizing, updated across reconnects) */
116 + char run_dir[256];
117 + char service_name[128];
118 +
119 +#if defined(_WIN32) || defined(__MSYS__)
120 + nipc_np_client_config_t transport_config;
121 +
122 + /* Connection (managed internally) */
123 + nipc_np_session_t session;
124 + bool session_valid;
125 + nipc_win_shm_ctx_t *shm; /* non-NULL if SHM profile negotiated */
126 +#else
127 + nipc_uds_client_config_t transport_config;
128 +
129 + /* Connection (managed internally) */
130 + nipc_uds_session_t session;
131 + bool session_valid;
132 + nipc_shm_ctx_t *shm; /* non-NULL if SHM profile negotiated */
133 +#endif
134 +
135 + /* Stats */
136 + uint32_t connect_count;
137 + uint32_t reconnect_count;
138 + uint32_t call_count;
139 + uint32_t error_count;
140 +
141 + /* Internal reusable Level 2 scratch, allocated lazily from session sizes. */
142 + uint8_t *response_buf;
143 + size_t response_buf_size;
144 + uint8_t *send_buf;
145 + size_t send_buf_size;
146 +} nipc_client_ctx_t;
147 +
148 +/* ------------------------------------------------------------------ */
149 +/* Client API */
150 +/* ------------------------------------------------------------------ */
151 +
152 +/*
153 + * Initialize a client context. Does NOT connect. Does NOT require the
154 + * server to be running. State starts as DISCONNECTED.
155 + */
156 +void nipc_client_init(nipc_client_ctx_t *ctx,
157 + const char *run_dir,
158 + const char *service_name,
159 + const nipc_client_config_t *config);
160 +
161 +/*
162 + * Attempt connect if DISCONNECTED, reconnect if BROKEN.
163 + * Returns true if the state changed (e.g. DISCONNECTED -> READY),
164 + * false if unchanged.
165 + *
166 + * No hidden threads. Call from your own loop at your own cadence.
167 + */
168 +bool nipc_client_refresh(nipc_client_ctx_t *ctx);
169 +
170 +/*
171 + * Cheap cached boolean. No I/O, no syscalls.
172 + * Returns true only if state == READY.
173 + */
174 +static inline bool nipc_client_ready(const nipc_client_ctx_t *ctx) {
175 + return ctx->state == NIPC_CLIENT_READY;
176 +}
177 +
178 +/*
179 + * Detailed status snapshot. For diagnostics and logging, not hot path.
180 + */
181 +void nipc_client_status(const nipc_client_ctx_t *ctx,
182 + nipc_client_status_t *out);
183 +
184 +/*
185 + * Tear down connection and release resources. Safe on a zero-init ctx.
186 + */
187 +void nipc_client_close(nipc_client_ctx_t *ctx);
188 +
189 +/* ------------------------------------------------------------------ */
190 +/* Typed cgroups snapshot call */
191 +/* ------------------------------------------------------------------ */
192 +
193 +/*
194 + * Blocking typed call: encode request, send, receive, check
195 + * transport_status, decode response.
196 + *
197 + * view_out: on success, filled with the ephemeral snapshot view.
198 + * Valid only until the next call on this context.
199 + *
200 + * Retry policy (per spec):
201 + * If the call fails and the context was previously READY, the client
202 + * disconnects, reconnects (full handshake), and retries.
203 + * - Ordinary transport / peer failures retry once.
204 + * - Overflow-driven resize recovery may reconnect more than once
205 + * while negotiated capacities grow.
206 + * If not previously READY, fails immediately.
207 + *
208 + * Returns NIPC_OK on success, or an error code.
209 + */
210 +nipc_error_t nipc_client_call_cgroups_snapshot(
211 + nipc_client_ctx_t *ctx,
212 + nipc_cgroups_resp_view_t *view_out);
213 +
214 +/* ------------------------------------------------------------------ */
215 +/* Managed server */
216 +/* ------------------------------------------------------------------ */
217 +
218 +/* Internal dispatch callback shape used by the managed-server loop. */
219 +typedef nipc_error_t (*nipc_server_handler_fn)(
220 + void *user,
221 + const nipc_header_t *request_hdr,
222 + const uint8_t *request_payload, size_t request_len,
223 + uint8_t *response_buf, size_t response_buf_size,
224 + size_t *response_len_out);
225 +
226 +/* Typed managed-server callback surface for the cgroups-snapshot service. */
227 +typedef struct {
228 + nipc_cgroups_handler_fn handle;
229 +
230 + /* Optional explicit reservation hint for snapshot builders. When 0,
231 + * the library derives a safe upper bound from negotiated response
232 + * limits. */
233 + uint32_t snapshot_max_items;
234 +
235 + void *user;
236 +} nipc_cgroups_service_handler_t;
237 +
238 +typedef struct nipc_managed_server nipc_managed_server_t;
239 +
240 +/* Per-session context for multi-client server */
241 +typedef struct nipc_session_ctx {
242 + nipc_managed_server_t *server; /* back-pointer */
243 +#if defined(_WIN32) || defined(__MSYS__)
244 + nipc_np_session_t session;
245 + nipc_win_shm_ctx_t *shm; /* non-NULL if SHM negotiated */
246 + HANDLE thread;
247 +#else
248 + nipc_uds_session_t session;
249 + nipc_shm_ctx_t *shm; /* non-NULL if SHM negotiated */
250 + pthread_t thread;
251 +#endif
252 + uint64_t id;
253 + bool active; /* use __atomic builtins for cross-thread access */
254 +} nipc_session_ctx_t;
255 +
256 +struct nipc_managed_server {
257 + /* Listener */
258 +#if defined(_WIN32) || defined(__MSYS__)
259 + nipc_np_listener_t listener;
260 + nipc_np_server_config_t base_config;
261 +#else
262 + nipc_uds_listener_t listener;
263 + nipc_uds_server_config_t base_config;
264 +#endif
265 +
266 + /* Concurrency control */
267 + int worker_count; /* max concurrent sessions */
268 +
269 + /* Callback */
270 + nipc_server_handler_fn handler;
271 + void *handler_user;
272 + nipc_cgroups_service_handler_t service_handler;
273 + uint16_t expected_method_code;
274 + uint32_t learned_request_payload_bytes;
275 + uint32_t learned_response_payload_bytes;
276 +
277 + /* State — use __atomic builtins (POSIX) or Interlocked* (Windows) */
278 +#if defined(_WIN32) || defined(__MSYS__)
279 + volatile LONG running;
280 + volatile LONG accept_loop_active;
281 +#else
282 + bool running;
283 +#endif
284 +
285 + /* Session tracking */
286 + nipc_session_ctx_t **sessions; /* dynamic array of active sessions */
287 + int session_count; /* current active session count */
288 + int session_capacity; /* allocated slots */
289 + uint64_t next_session_id; /* monotonic session ID counter */
290 +#if defined(_WIN32) || defined(__MSYS__)
291 + CRITICAL_SECTION sessions_lock; /* protects session array + count */
292 +#else
293 + pthread_mutex_t sessions_lock; /* protects session array + count */
294 + pthread_t acceptor_thread;
295 + bool acceptor_started;
296 +#endif
297 +
298 + /* Configuration */
299 + char run_dir[256];
300 + char service_name[128];
301 +
302 +#if defined(_WIN32) || defined(__MSYS__)
303 + /* Auth token needed for SHM kernel object naming */
304 + uint64_t auth_token;
305 +#endif
306 +};
307 +
308 +/*
309 + * Initialize a managed server for the typed cgroups-snapshot service kind.
310 + * One endpoint serves one request kind only.
311 + * Does NOT start workers. Call nipc_server_run() to start the
312 + * acceptor+worker loop.
313 + */
314 +nipc_error_t nipc_server_init_typed(nipc_managed_server_t *server,
315 + const char *run_dir,
316 + const char *service_name,
317 + const nipc_server_config_t *config,
318 + int worker_count,
319 + const nipc_cgroups_service_handler_t *service_handler);
320 +
321 +#ifdef NIPC_INTERNAL_TESTING
322 +/*
323 + * Internal compatibility entrypoint for repo tests and benchmarks that
324 + * intentionally exercise raw dispatch or malformed response paths. Not part
325 + * of the public Level 2 contract.
326 + */
327 +nipc_error_t nipc_server_init_raw_for_tests(
328 + nipc_managed_server_t *server,
329 + const char *run_dir,
330 + const char *service_name,
331 +#if defined(_WIN32) || defined(__MSYS__)
332 + const nipc_np_server_config_t *config,
333 +#else
334 + const nipc_uds_server_config_t *config,
335 +#endif
336 + int worker_count,
337 + uint16_t expected_method_code,
338 + nipc_server_handler_fn handler,
339 + void *user);
340 +
341 +#if defined(_WIN32) || defined(__MSYS__)
342 +typedef enum {
343 + NIPC_WIN_SERVICE_TEST_FAULT_CLIENT_RESPONSE_BUF_REALLOC = 1,
344 + NIPC_WIN_SERVICE_TEST_FAULT_CLIENT_SEND_BUF_REALLOC,
345 + NIPC_WIN_SERVICE_TEST_FAULT_CLIENT_SHM_CTX_CALLOC,
346 + NIPC_WIN_SERVICE_TEST_FAULT_SERVER_SHM_CTX_CALLOC,
347 + NIPC_WIN_SERVICE_TEST_FAULT_SERVER_RECV_BUF_MALLOC,
348 + NIPC_WIN_SERVICE_TEST_FAULT_SERVER_RESP_BUF_MALLOC,
349 + NIPC_WIN_SERVICE_TEST_FAULT_SERVER_SESSIONS_CALLOC,
350 + NIPC_WIN_SERVICE_TEST_FAULT_SERVER_SESSION_CTX_CALLOC,
351 + NIPC_WIN_SERVICE_TEST_FAULT_SERVER_THREAD_CREATE,
352 + NIPC_WIN_SERVICE_TEST_FAULT_CACHE_BUCKETS_CALLOC,
353 + NIPC_WIN_SERVICE_TEST_FAULT_CACHE_ITEMS_CALLOC,
354 + NIPC_WIN_SERVICE_TEST_FAULT_CACHE_ITEM_NAME_MALLOC,
355 + NIPC_WIN_SERVICE_TEST_FAULT_CACHE_ITEM_PATH_MALLOC,
356 +} nipc_win_service_test_fault_site_t;
357 +
358 +void nipc_win_service_test_fault_set(int site, uint32_t skip_matches);
359 +void nipc_win_service_test_fault_clear(void);
360 +#else
361 +typedef enum {
362 + NIPC_POSIX_SERVICE_TEST_FAULT_CLIENT_RESPONSE_BUF_REALLOC = 1,
363 + NIPC_POSIX_SERVICE_TEST_FAULT_CLIENT_SEND_BUF_REALLOC,
364 + NIPC_POSIX_SERVICE_TEST_FAULT_CLIENT_SHM_CTX_CALLOC,
365 + NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_SHM_CTX_CALLOC,
366 + NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_RECV_BUF_MALLOC,
367 + NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_RESP_BUF_MALLOC,
368 + NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_SESSIONS_CALLOC,
369 + NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_SESSION_CTX_CALLOC,
370 + NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_THREAD_CREATE,
371 + NIPC_POSIX_SERVICE_TEST_FAULT_CACHE_BUCKETS_CALLOC,
372 + NIPC_POSIX_SERVICE_TEST_FAULT_CACHE_ITEMS_CALLOC,
373 + NIPC_POSIX_SERVICE_TEST_FAULT_CACHE_ITEM_NAME_MALLOC,
374 + NIPC_POSIX_SERVICE_TEST_FAULT_CACHE_ITEM_PATH_MALLOC,
375 +} nipc_posix_service_test_fault_site_t;
376 +
377 +void nipc_posix_service_test_fault_set(int site, uint32_t skip_matches);
378 +void nipc_posix_service_test_fault_clear(void);
379 +#endif
380 +
381 +#define nipc_server_init nipc_server_init_raw_for_tests
382 +#endif
383 +
384 +/*
385 + * Run the acceptor loop. Blocking. Accepts clients, reads requests,
386 + * dispatches to the handler, sends responses.
387 + *
388 + * Returns when nipc_server_stop() is called or on fatal error.
389 + */
390 +void nipc_server_run(nipc_managed_server_t *server);
391 +
392 +/*
393 + * Signal shutdown. The acceptor loop will exit after current work.
394 + */
395 +void nipc_server_stop(nipc_managed_server_t *server);
396 +
397 +/*
398 + * Graceful drain: stop accepting new clients, wait for in-flight
399 + * sessions to complete (up to timeout_ms), then close everything.
400 + *
401 + * Combines stop + wait + destroy in one call. If sessions don't
402 + * finish within timeout_ms, they are forcibly closed.
403 + *
404 + * Returns true if all sessions completed within the timeout,
405 + * false if the timeout expired and sessions were forcibly closed.
406 + */
407 +bool nipc_server_drain(nipc_managed_server_t *server, uint32_t timeout_ms);
408 +
409 +/*
410 + * Cleanup: close listener, free workers. Safe after stop.
411 + */
412 +void nipc_server_destroy(nipc_managed_server_t *server);
413 +
414 +/* ------------------------------------------------------------------ */
415 +/* L3: Client-side cgroups snapshot cache */
416 +/* ------------------------------------------------------------------ */
417 +
418 +/* Default response buffer size for L3 cache refresh (when config is 0) */
419 +#define NIPC_CGROUPS_CACHE_BUF_SIZE_DEFAULT 65536
420 +
421 +/*
422 + * Cached copy of a single cgroup item. Owns its strings.
423 + * Built from ephemeral L2 views during cache construction.
424 + */
425 +typedef struct {
426 + uint32_t hash;
427 + uint32_t options;
428 + uint32_t enabled;
429 + char *name; /* owned NUL-terminated copy */
430 + char *path; /* owned NUL-terminated copy */
431 +} nipc_cgroups_cache_item_t;
432 +
433 +/*
434 + * L3 cache status snapshot (for diagnostics, not hot path).
435 + */
436 +typedef struct {
437 + bool populated;
438 + uint32_t item_count;
439 + uint32_t systemd_enabled;
440 + uint64_t generation;
441 + uint32_t refresh_success_count;
442 + uint32_t refresh_failure_count;
443 + nipc_client_state_t connection_state; /* current L2 connection state */
444 + uint64_t last_refresh_ts; /* monotonic timestamp of last successful refresh (ms), 0 if never */
445 +} nipc_cgroups_cache_status_t;
446 +
447 +/*
448 + * L3 client-side cgroups snapshot cache.
449 + *
450 + * Wraps an L2 client context and maintains a local owned copy of the
451 + * most recent successful snapshot. Lookup by hash+name is O(1) via
452 + * an open-addressing hash table, with no I/O.
453 + *
454 + * On refresh failure, the previous cache is preserved. The cache
455 + * becomes empty only if no successful refresh has ever occurred.
456 + */
457 +
458 +/* Hash table bucket for O(1) cache lookup */
459 +typedef struct {
460 + uint32_t index; /* index into items[] array */
461 + bool used; /* true if this bucket is occupied */
462 +} nipc_cgroups_hash_bucket_t;
463 +
464 +typedef struct {
465 + nipc_client_ctx_t client;
466 +
467 + /* Cache data (owned) */
468 + nipc_cgroups_cache_item_t *items;
469 + uint32_t item_count;
470 + uint32_t systemd_enabled;
471 + uint64_t generation;
472 + bool populated;
473 +
474 + /* Hash table for O(1) lookup (open addressing, rebuilt on refresh) */
475 + nipc_cgroups_hash_bucket_t *buckets;
476 + uint32_t bucket_count; /* always a power of 2 */
477 +
478 + /* Counters */
479 + uint32_t refresh_success_count;
480 + uint32_t refresh_failure_count;
481 + uint64_t last_refresh_ts; /* monotonic ms of last successful refresh */
482 +
483 + /* Internal: response buffer for L2 calls */
484 + uint8_t *response_buf;
485 + size_t response_buf_size;
486 +} nipc_cgroups_cache_t;
487 +
488 +/*
489 + * Initialize an L3 cache. Creates the underlying L2 client context.
490 + * Does NOT connect. Does NOT require the server to be running.
491 + * Cache starts empty (populated == false).
492 + */
493 +void nipc_cgroups_cache_init(nipc_cgroups_cache_t *cache,
494 + const char *run_dir,
495 + const char *service_name,
496 + const nipc_client_config_t *config);
497 +
498 +/*
499 + * Refresh the cache. Drives the L2 client (connect/reconnect as
500 + * needed) and requests a fresh snapshot. On success, rebuilds the
501 + * local cache from the response (copies strings). On failure,
502 + * preserves the previous cache.
503 + *
504 + * Caller-driven: call from your own loop at your own cadence.
505 + * Returns true if the cache was updated, false otherwise.
506 + */
507 +bool nipc_cgroups_cache_refresh(nipc_cgroups_cache_t *cache);
508 +
509 +/*
510 + * Returns true if at least one successful refresh has occurred.
511 + * Cheap cached boolean. No I/O, no syscalls.
512 + *
513 + * Note: ready means "has cached data", not "is connected."
514 + */
515 +static inline bool nipc_cgroups_cache_ready(const nipc_cgroups_cache_t *cache) {
516 + return cache->populated;
517 +}
518 +
519 +/*
520 + * Look up a cached item by hash + name. Pure in-memory, no I/O.
521 + *
522 + * Returns a pointer to the cached item, or NULL if not found or
523 + * cache is empty. The returned pointer is valid until the next
524 + * successful refresh.
525 + */
526 +const nipc_cgroups_cache_item_t *nipc_cgroups_cache_lookup(
527 + const nipc_cgroups_cache_t *cache,
528 + uint32_t hash,
529 + const char *name);
530 +
531 +/*
532 + * Fill a status snapshot for diagnostics.
533 + */
534 +void nipc_cgroups_cache_status(const nipc_cgroups_cache_t *cache,
535 + nipc_cgroups_cache_status_t *out);
536 +
537 +/*
538 + * Close the cache: free all cached items, close the L2 client.
539 + * Safe on a zero-initialized cache.
540 + */
541 +void nipc_cgroups_cache_close(nipc_cgroups_cache_t *cache);
542 +
543 +#ifdef __cplusplus
544 +}
545 +#endif
546 +
547 +#endif /* NETIPC_SERVICE_H */
src/libnetdata/netipc/include/netipc/netipc_shm.h new
+229
@@ -0,0 +1,229 @@
1 +/*
2 + * netipc_shm.h - L1 POSIX SHM transport (Linux only).
3 + *
4 + * Shared memory data plane with spin+futex synchronization.
5 + * The SHM region carries the same outer protocol envelope as the UDS
6 + * transport. Higher levels see no difference.
7 + *
8 + * Lifecycle:
9 + * 1. UDS handshake negotiates PROFILE_SHM_HYBRID.
10 + * 2. Server creates the SHM region; client attaches.
11 + * 3. Data plane switches to SHM; UDS socket stays open.
12 + */
13 +
14 +#ifndef NETIPC_SHM_H
15 +#define NETIPC_SHM_H
16 +
17 +#include "netipc_protocol.h"
18 +#include <stdbool.h>
19 +#include <stddef.h>
20 +#include <stdint.h>
21 +
22 +#ifdef __cplusplus
23 +extern "C" {
24 +#endif
25 +
26 +/* ------------------------------------------------------------------ */
27 +/* Constants */
28 +/* ------------------------------------------------------------------ */
29 +
30 +#define NIPC_SHM_REGION_MAGIC 0x4e53484du /* "NSHM" */
31 +#define NIPC_SHM_REGION_VERSION 3u
32 +#define NIPC_SHM_REGION_ALIGNMENT 64u
33 +#define NIPC_SHM_HEADER_LEN 64u
34 +#define NIPC_SHM_DEFAULT_SPIN 128u
35 +
36 +/* ------------------------------------------------------------------ */
37 +/* Error codes */
38 +/* ------------------------------------------------------------------ */
39 +
40 +typedef enum {
41 + NIPC_SHM_OK = 0,
42 + NIPC_SHM_ERR_PATH_TOO_LONG, /* SHM path exceeds limit */
43 + NIPC_SHM_ERR_OPEN, /* open/shm_open failed */
44 + NIPC_SHM_ERR_TRUNCATE, /* ftruncate failed */
45 + NIPC_SHM_ERR_MMAP, /* mmap failed */
46 + NIPC_SHM_ERR_BAD_MAGIC, /* header magic mismatch */
47 + NIPC_SHM_ERR_BAD_VERSION, /* header version mismatch */
48 + NIPC_SHM_ERR_BAD_HEADER, /* header_len mismatch or corrupt */
49 + NIPC_SHM_ERR_BAD_SIZE, /* file too small / capacity mismatch */
50 + NIPC_SHM_ERR_ADDR_IN_USE, /* live server owns the region */
51 + NIPC_SHM_ERR_NOT_READY, /* server hasn't finished setup (retry) */
52 + NIPC_SHM_ERR_MSG_TOO_LARGE, /* message exceeds area capacity */
53 + NIPC_SHM_ERR_TIMEOUT, /* futex wait timed out */
54 + NIPC_SHM_ERR_BAD_PARAM, /* invalid argument */
55 + NIPC_SHM_ERR_PEER_DEAD, /* owner process has exited */
56 +} nipc_shm_error_t;
57 +
58 +/* ------------------------------------------------------------------ */
59 +/* Role */
60 +/* ------------------------------------------------------------------ */
61 +
62 +typedef enum {
63 + NIPC_SHM_ROLE_SERVER = 1,
64 + NIPC_SHM_ROLE_CLIENT = 2,
65 +} nipc_shm_role_t;
66 +
67 +/* ------------------------------------------------------------------ */
68 +/* Region header (64 bytes, mapped at offset 0) */
69 +/* ------------------------------------------------------------------ */
70 +
71 +/*
72 + * This struct is the on-disk/on-memory layout. Atomic fields are
73 + * accessed through __atomic builtins, not through direct struct reads.
74 + * The struct is declared solely for offset verification via
75 + * _Static_assert and for header initialization.
76 + */
77 +typedef struct {
78 + uint32_t magic; /* 0: NIPC_SHM_REGION_MAGIC */
79 + uint16_t version; /* 4: NIPC_SHM_REGION_VERSION */
80 + uint16_t header_len; /* 6: 64 */
81 + int32_t owner_pid; /* 8: server PID */
82 + uint32_t owner_generation; /* 12: generation for PID reuse detection */
83 + uint32_t request_offset; /* 16: byte offset to request area */
84 + uint32_t request_capacity; /* 20: request area size */
85 + uint32_t response_offset; /* 24: byte offset to response area */
86 + uint32_t response_capacity; /* 28: response area size */
87 + uint64_t req_seq; /* 32: request sequence (atomic) */
88 + uint64_t resp_seq; /* 40: response sequence (atomic) */
89 + uint32_t req_len; /* 48: current request msg length (atomic) */
90 + uint32_t resp_len; /* 52: current response msg length (atomic) */
91 + uint32_t req_signal; /* 56: request futex word (atomic) */
92 + uint32_t resp_signal; /* 60: response futex word (atomic) */
93 +} nipc_shm_region_header_t;
94 +
95 +_Static_assert(sizeof(nipc_shm_region_header_t) == 64,
96 + "SHM region header must be exactly 64 bytes");
97 +
98 +/* ------------------------------------------------------------------ */
99 +/* SHM context */
100 +/* ------------------------------------------------------------------ */
101 +
102 +typedef struct {
103 + nipc_shm_role_t role;
104 +
105 + int fd; /* file descriptor for the SHM region */
106 + void *base; /* mmap base pointer */
107 + size_t region_size; /* total mapped size */
108 +
109 + /* Cached from header (avoid repeated volatile reads) */
110 + uint32_t request_offset;
111 + uint32_t request_capacity;
112 + uint32_t response_offset;
113 + uint32_t response_capacity;
114 +
115 + /* Sequence tracking for send/receive */
116 + uint64_t local_req_seq; /* last known req_seq */
117 + uint64_t local_resp_seq; /* last known resp_seq */
118 +
119 + uint32_t spin_tries; /* spin count before futex wait */
120 + uint32_t owner_generation; /* cached for PID reuse detection */
121 +
122 + char path[256]; /* stored for unlink on destroy */
123 +
124 +} nipc_shm_ctx_t;
125 +
126 +/* ------------------------------------------------------------------ */
127 +/* Server API */
128 +/* ------------------------------------------------------------------ */
129 +
130 +/*
131 + * Create a per-session SHM region at
132 + * {run_dir}/{service_name}-{session_id:016x}.ipcshm
133 + *
134 + * session_id is the server-assigned session identifier (from hello-ack).
135 + * req_capacity / resp_capacity are the data area sizes in bytes.
136 + * They will be rounded up to NIPC_SHM_REGION_ALIGNMENT.
137 + */
138 +nipc_shm_error_t nipc_shm_server_create(const char *run_dir,
139 + const char *service_name,
140 + uint64_t session_id,
141 + uint32_t req_capacity,
142 + uint32_t resp_capacity,
143 + nipc_shm_ctx_t *out);
144 +
145 +/*
146 + * Destroy a server SHM region: munmap, close, unlink.
147 + */
148 +void nipc_shm_destroy(nipc_shm_ctx_t *ctx);
149 +
150 +/* ------------------------------------------------------------------ */
151 +/* Client API */
152 +/* ------------------------------------------------------------------ */
153 +
154 +/*
155 + * Attach to a per-session SHM region at
156 + * {run_dir}/{service_name}-{session_id:016x}.ipcshm
157 + *
158 + * session_id is from the hello-ack received during handshake.
159 + * Validates the header (magic, version, sizes). If the file is
160 + * undersized (server not ready), returns NIPC_SHM_ERR_NOT_READY.
161 + */
162 +nipc_shm_error_t nipc_shm_client_attach(const char *run_dir,
163 + const char *service_name,
164 + uint64_t session_id,
165 + nipc_shm_ctx_t *out);
166 +
167 +/*
168 + * Detach a client from the SHM region: munmap, close (no unlink).
169 + */
170 +void nipc_shm_close(nipc_shm_ctx_t *ctx);
171 +
172 +/* ------------------------------------------------------------------ */
173 +/* Data plane */
174 +/* ------------------------------------------------------------------ */
175 +
176 +/*
177 + * Publish a message into the SHM region (client sends request,
178 + * server sends response). The role determines which area is written.
179 + *
180 + * The message must include the 32-byte outer header + payload, exactly
181 + * as would be sent over UDS.
182 + *
183 + * Returns NIPC_SHM_ERR_MSG_TOO_LARGE if msg_len exceeds the area capacity.
184 + */
185 +nipc_shm_error_t nipc_shm_send(nipc_shm_ctx_t *ctx,
186 + const void *msg, size_t msg_len);
187 +
188 +/*
189 + * Receive a message from the SHM region (server receives request,
190 + * client receives response).
191 + *
192 + * Spins for ctx->spin_tries iterations, then falls back to futex wait
193 + * with timeout_ms. Pass 0 for no timeout (infinite wait).
194 + *
195 + * The message is copied into the caller-provided buffer (buf, buf_size).
196 + * On success, *msg_len_out is the message length. Returns
197 + * NIPC_SHM_ERR_MSG_TOO_LARGE if the message exceeds buf_size.
198 + */
199 +nipc_shm_error_t nipc_shm_receive(nipc_shm_ctx_t *ctx,
200 + void *buf,
201 + size_t buf_size,
202 + size_t *msg_len_out,
203 + uint32_t timeout_ms);
204 +
205 +/* ------------------------------------------------------------------ */
206 +/* Utility */
207 +/* ------------------------------------------------------------------ */
208 +
209 +/* Check if the region's owner process is still alive. */
210 +bool nipc_shm_owner_alive(const nipc_shm_ctx_t *ctx);
211 +
212 +/*
213 + * Scan for and unlink stale per-session SHM files left by a crashed
214 + * server. Call once at server startup before accepting connections.
215 + * Files matching {run_dir}/{service_name}-*.ipcshm whose owner_pid
216 + * is dead (or whose generation mismatches) are unlinked.
217 + */
218 +void nipc_shm_cleanup_stale(const char *run_dir, const char *service_name);
219 +
220 +/* Get the file descriptor for external event integration. */
221 +static inline int nipc_shm_fd(const nipc_shm_ctx_t *ctx) {
222 + return ctx->fd;
223 +}
224 +
225 +#ifdef __cplusplus
226 +}
227 +#endif
228 +
229 +#endif /* NETIPC_SHM_H */
src/libnetdata/netipc/include/netipc/netipc_uds.h new
+225
@@ -0,0 +1,225 @@
1 +/*
2 + * netipc_uds.h - L1 POSIX UDS SEQPACKET transport.
3 + *
4 + * Connection lifecycle, handshake, send/receive with transparent chunking.
5 + * Uses the wire envelope from netipc_protocol.h for all framing.
6 + */
7 +
8 +#ifndef NETIPC_UDS_H
9 +#define NETIPC_UDS_H
10 +
11 +#include "netipc_protocol.h"
12 +#include <stdbool.h>
13 +#include <stddef.h>
14 +#include <stdint.h>
15 +
16 +#ifdef __cplusplus
17 +extern "C" {
18 +#endif
19 +
20 +/* ------------------------------------------------------------------ */
21 +/* Error codes (transport-level) */
22 +/* ------------------------------------------------------------------ */
23 +
24 +typedef enum {
25 + NIPC_UDS_OK = 0,
26 + NIPC_UDS_ERR_PATH_TOO_LONG, /* socket path exceeds sun_path */
27 + NIPC_UDS_ERR_SOCKET, /* socket()/bind()/listen() failed */
28 + NIPC_UDS_ERR_CONNECT, /* connect() failed */
29 + NIPC_UDS_ERR_ACCEPT, /* accept() failed */
30 + NIPC_UDS_ERR_SEND, /* send() failed */
31 + NIPC_UDS_ERR_RECV, /* recv() failed / peer disconnected */
32 + NIPC_UDS_ERR_HANDSHAKE, /* handshake protocol error */
33 + NIPC_UDS_ERR_AUTH_FAILED, /* auth token rejected */
34 + NIPC_UDS_ERR_NO_PROFILE, /* no common profile */
35 + NIPC_UDS_ERR_INCOMPATIBLE, /* protocol/layout version mismatch */
36 + NIPC_UDS_ERR_PROTOCOL, /* wire protocol violation */
37 + NIPC_UDS_ERR_ADDR_IN_USE, /* live server on the socket path */
38 + NIPC_UDS_ERR_CHUNK, /* chunk header mismatch */
39 + NIPC_UDS_ERR_ALLOC, /* memory allocation failed */
40 + NIPC_UDS_ERR_LIMIT_EXCEEDED, /* payload/batch exceeds negotiated */
41 + NIPC_UDS_ERR_BAD_PARAM, /* invalid argument */
42 + NIPC_UDS_ERR_DUPLICATE_MSG_ID, /* message_id already in-flight */
43 + NIPC_UDS_ERR_UNKNOWN_MSG_ID, /* response message_id not in-flight */
44 +} nipc_uds_error_t;
45 +
46 +/* ------------------------------------------------------------------ */
47 +/* Role */
48 +/* ------------------------------------------------------------------ */
49 +
50 +typedef enum {
51 + NIPC_UDS_ROLE_CLIENT = 1,
52 + NIPC_UDS_ROLE_SERVER = 2,
53 +} nipc_uds_role_t;
54 +
55 +/* ------------------------------------------------------------------ */
56 +/* Client connect configuration */
57 +/* ------------------------------------------------------------------ */
58 +
59 +typedef struct {
60 + uint32_t supported_profiles; /* bitmask */
61 + uint32_t preferred_profiles; /* bitmask */
62 + uint32_t max_request_payload_bytes; /* 0 = use default */
63 + uint32_t max_request_batch_items; /* 0 = use default (1) */
64 + uint32_t max_response_payload_bytes;/* 0 = use default */
65 + uint32_t max_response_batch_items; /* 0 = use default (1) */
66 + uint64_t auth_token;
67 + uint32_t packet_size; /* 0 = auto-detect from SO_SNDBUF */
68 +} nipc_uds_client_config_t;
69 +
70 +/* ------------------------------------------------------------------ */
71 +/* Server configuration (for listen + accept) */
72 +/* ------------------------------------------------------------------ */
73 +
74 +typedef struct {
75 + uint32_t supported_profiles;
76 + uint32_t preferred_profiles;
77 + uint32_t max_request_payload_bytes; /* 0 = use default */
78 + uint32_t max_request_batch_items; /* 0 = use default (1) */
79 + uint32_t max_response_payload_bytes; /* 0 = use default */
80 + uint32_t max_response_batch_items; /* 0 = use default (1) */
81 + uint64_t auth_token; /* expected token from clients */
82 + uint32_t packet_size; /* 0 = auto-detect from SO_SNDBUF */
83 + int backlog; /* listen backlog, 0 = default (16) */
84 +} nipc_uds_server_config_t;
85 +
86 +/* ------------------------------------------------------------------ */
87 +/* Session */
88 +/* ------------------------------------------------------------------ */
89 +
90 +typedef struct {
91 + int fd; /* connected socket (native wait object) */
92 + nipc_uds_role_t role;
93 +
94 + /* Negotiated limits */
95 + uint32_t max_request_payload_bytes;
96 + uint32_t max_request_batch_items;
97 + uint32_t max_response_payload_bytes;
98 + uint32_t max_response_batch_items;
99 + uint32_t packet_size;
100 + uint32_t selected_profile;
101 +
102 + /* Server-assigned session ID (from hello-ack) */
103 + uint64_t session_id;
104 +
105 + /* Internal receive buffer for chunked reassembly */
106 + uint8_t *recv_buf;
107 + size_t recv_buf_size;
108 +
109 + /* In-flight message_id set (client-side only, dynamically grown) */
110 + uint64_t *inflight_ids;
111 + uint32_t inflight_count;
112 + uint32_t inflight_capacity;
113 +} nipc_uds_session_t;
114 +
115 +/* ------------------------------------------------------------------ */
116 +/* Listener */
117 +/* ------------------------------------------------------------------ */
118 +
119 +typedef struct {
120 + int fd; /* listening socket */
121 + nipc_uds_server_config_t config;
122 + char path[108]; /* stored for cleanup */
123 +} nipc_uds_listener_t;
124 +
125 +/* ------------------------------------------------------------------ */
126 +/* Connection lifecycle */
127 +/* ------------------------------------------------------------------ */
128 +
129 +/*
130 + * Create a listener on {run_dir}/{service_name}.sock.
131 + * Performs stale endpoint recovery: if the socket file exists and no
132 + * live server is connected, unlinks and recreates it.
133 + */
134 +nipc_uds_error_t nipc_uds_listen(const char *run_dir,
135 + const char *service_name,
136 + const nipc_uds_server_config_t *config,
137 + nipc_uds_listener_t *out);
138 +
139 +/*
140 + * Accept one client on a listener. Performs the full handshake.
141 + * session_id is placed into the hello-ack so the client can attach
142 + * to the correct per-session SHM region.
143 + * Blocks until a client connects and the handshake completes (or fails).
144 + */
145 +nipc_uds_error_t nipc_uds_accept(nipc_uds_listener_t *listener,
146 + uint64_t session_id,
147 + nipc_uds_session_t *out);
148 +
149 +/*
150 + * Connect to a server at {run_dir}/{service_name}.sock.
151 + * Performs the full handshake. Blocks until connected + handshake done.
152 + */
153 +nipc_uds_error_t nipc_uds_connect(const char *run_dir,
154 + const char *service_name,
155 + const nipc_uds_client_config_t *config,
156 + nipc_uds_session_t *out);
157 +
158 +/*
159 + * Close a session. Releases socket and internal buffers.
160 + * Safe to call on a zero-initialized session (no-op).
161 + */
162 +void nipc_uds_close_session(nipc_uds_session_t *session);
163 +
164 +/*
165 + * Close a listener. Stops accepting, closes the socket, unlinks
166 + * the socket file.
167 + */
168 +void nipc_uds_close_listener(nipc_uds_listener_t *listener);
169 +
170 +/* ------------------------------------------------------------------ */
171 +/* Message send / receive */
172 +/* ------------------------------------------------------------------ */
173 +
174 +/*
175 + * Send one logical message. hdr is the 32-byte outer header (caller fills
176 + * kind, code, flags, payload_len, item_count, message_id; this function
177 + * sets magic/version/header_len). payload is the opaque payload bytes.
178 + *
179 + * If the total message (32 + payload_len) exceeds packet_size, the
180 + * message is chunked transparently.
181 + *
182 + * Blocks until all chunks are sent or an error occurs.
183 + */
184 +nipc_uds_error_t nipc_uds_send(nipc_uds_session_t *session,
185 + nipc_header_t *hdr,
186 + const void *payload,
187 + size_t payload_len);
188 +
189 +/*
190 + * Receive one logical message. Blocks until a complete message arrives.
191 + *
192 + * On success:
193 + * - hdr_out is filled with the decoded outer header.
194 + * - *payload_out points to the payload bytes (inside session->recv_buf
195 + * or buf). Valid until the next nipc_uds_receive call.
196 + * - *payload_len_out is the payload length.
197 + *
198 + * buf/buf_size: caller-provided buffer. If the message fits, it is
199 + * placed here. If the message is chunked and exceeds buf_size, the
200 + * session's internal recv_buf is grown to fit (allocated via realloc).
201 + */
202 +nipc_uds_error_t nipc_uds_receive(nipc_uds_session_t *session,
203 + void *buf, size_t buf_size,
204 + nipc_header_t *hdr_out,
205 + const void **payload_out,
206 + size_t *payload_len_out);
207 +
208 +/* ------------------------------------------------------------------ */
209 +/* Utility */
210 +/* ------------------------------------------------------------------ */
211 +
212 +/* Get the fd for poll/epoll integration. */
213 +static inline int nipc_uds_session_fd(const nipc_uds_session_t *s) {
214 + return s->fd;
215 +}
216 +
217 +static inline int nipc_uds_listener_fd(const nipc_uds_listener_t *l) {
218 + return l->fd;
219 +}
220 +
221 +#ifdef __cplusplus
222 +}
223 +#endif
224 +
225 +#endif /* NETIPC_UDS_H */
src/libnetdata/netipc/include/netipc/netipc_win_shm.h new
+277
@@ -0,0 +1,277 @@
1 +/*
2 + * netipc_win_shm.h - L1 Windows SHM transport.
3 + *
4 + * Shared memory data plane with spin + kernel event synchronization.
5 + * Uses CreateFileMappingW/MapViewOfFile for the region and
6 + * auto-reset kernel events for synchronization.
7 + *
8 + * Kernel object name derivation:
9 + * Local\netipc-{FNV1a64(run_dir+"\n"+service_name+"\n"+auth_token):016llx}-{service}-p{profile}-s{session_id:016llx}-mapping
10 + * Local\netipc-{hash}-{service}-p{profile}-s{session_id:016llx}-req_event
11 + * Local\netipc-{hash}-{service}-p{profile}-s{session_id:016llx}-resp_event
12 + */
13 +
14 +#ifndef NETIPC_WIN_SHM_H
15 +#define NETIPC_WIN_SHM_H
16 +
17 +#if defined(_WIN32) || defined(__MSYS__)
18 +
19 +#include "netipc_protocol.h"
20 +#include <stdbool.h>
21 +#include <stddef.h>
22 +#include <stdint.h>
23 +#include <windows.h>
24 +
25 +#ifdef __cplusplus
26 +extern "C" {
27 +#endif
28 +
29 +/* ------------------------------------------------------------------ */
30 +/* Constants */
31 +/* ------------------------------------------------------------------ */
32 +
33 +/* Magic: "NSWH" as u32 LE */
34 +#define NIPC_WIN_SHM_MAGIC 0x4e535748u
35 +#define NIPC_WIN_SHM_VERSION 3u
36 +#define NIPC_WIN_SHM_HEADER_LEN 128u
37 +#define NIPC_WIN_SHM_CACHELINE 64u
38 +
39 +/* Profile bits */
40 +#define NIPC_WIN_SHM_PROFILE_HYBRID 0x02u
41 +#define NIPC_WIN_SHM_PROFILE_BUSYWAIT 0x04u
42 +
43 +/* Default spin count (higher than POSIX due to Windows kernel overhead) */
44 +#define NIPC_WIN_SHM_DEFAULT_SPIN 1024u
45 +
46 +/* Busy-wait deadline poll mask */
47 +#define NIPC_WIN_SHM_BUSYWAIT_POLL_MASK 1023u
48 +
49 +/* Max kernel object name length */
50 +#define NIPC_WIN_SHM_MAX_NAME 256
51 +
52 +/* ------------------------------------------------------------------ */
53 +/* Error codes */
54 +/* ------------------------------------------------------------------ */
55 +
56 +typedef enum {
57 + NIPC_WIN_SHM_OK = 0,
58 + NIPC_WIN_SHM_ERR_BAD_PARAM, /* invalid argument */
59 + NIPC_WIN_SHM_ERR_CREATE_MAPPING, /* CreateFileMappingW failed */
60 + NIPC_WIN_SHM_ERR_OPEN_MAPPING, /* OpenFileMappingW failed */
61 + NIPC_WIN_SHM_ERR_MAP_VIEW, /* MapViewOfFile failed */
62 + NIPC_WIN_SHM_ERR_CREATE_EVENT, /* CreateEventW failed */
63 + NIPC_WIN_SHM_ERR_OPEN_EVENT, /* OpenEventW failed */
64 + NIPC_WIN_SHM_ERR_ADDR_IN_USE, /* named mapping/event already exists */
65 + NIPC_WIN_SHM_ERR_BAD_MAGIC, /* header magic mismatch */
66 + NIPC_WIN_SHM_ERR_BAD_VERSION, /* header version mismatch */
67 + NIPC_WIN_SHM_ERR_BAD_HEADER, /* header_len mismatch */
68 + NIPC_WIN_SHM_ERR_BAD_PROFILE, /* profile mismatch */
69 + NIPC_WIN_SHM_ERR_MSG_TOO_LARGE, /* message exceeds area capacity */
70 + NIPC_WIN_SHM_ERR_TIMEOUT, /* wait timed out */
71 + NIPC_WIN_SHM_ERR_DISCONNECTED, /* peer closed */
72 +} nipc_win_shm_error_t;
73 +
74 +/* ------------------------------------------------------------------ */
75 +/* Role */
76 +/* ------------------------------------------------------------------ */
77 +
78 +typedef enum {
79 + NIPC_WIN_SHM_ROLE_SERVER = 1,
80 + NIPC_WIN_SHM_ROLE_CLIENT = 2,
81 +} nipc_win_shm_role_t;
82 +
83 +/* ------------------------------------------------------------------ */
84 +/* Region header (128 bytes, mapped at offset 0) */
85 +/* ------------------------------------------------------------------ */
86 +
87 +/*
88 + * On-disk/on-memory layout. Volatile fields are accessed via
89 + * InterlockedExchange / InterlockedCompareExchange, not direct reads.
90 + * Declared for offset verification and header initialization.
91 + */
92 +#pragma pack(push, 1)
93 +typedef struct {
94 + uint32_t magic; /* 0: NIPC_WIN_SHM_MAGIC */
95 + uint32_t version; /* 4: NIPC_WIN_SHM_VERSION */
96 + uint32_t header_len; /* 8: 128 */
97 + uint32_t profile; /* 12: selected profile */
98 + uint32_t request_offset; /* 16: byte offset to request area */
99 + uint32_t request_capacity; /* 20: request area size */
100 + uint32_t response_offset; /* 24: byte offset to response area */
101 + uint32_t response_capacity; /* 28: response area size */
102 + uint32_t spin_tries; /* 32: spin iterations before kernel wait */
103 + volatile LONG req_len; /* 36: current request message length */
104 + volatile LONG resp_len; /* 40: current response message length */
105 + volatile LONG req_client_closed; /* 44: client-side close flag */
106 + volatile LONG req_server_waiting; /* 48: server waiting for request */
107 + volatile LONG resp_server_closed; /* 52: server-side close flag */
108 + volatile LONG resp_client_waiting;/* 56: client waiting for response */
109 + uint32_t _padding; /* 60: reserved */
110 + volatile LONG64 req_seq; /* 64: request sequence number */
111 + volatile LONG64 resp_seq; /* 72: response sequence number */
112 + uint8_t _reserved[48]; /* 80: reserved for future use */
113 +} nipc_win_shm_region_header_t;
114 +#pragma pack(pop)
115 +
116 +/* Compile-time assertions for header layout */
117 +_Static_assert(sizeof(nipc_win_shm_region_header_t) == 128,
118 + "Windows SHM region header must be exactly 128 bytes");
119 +_Static_assert(offsetof(nipc_win_shm_region_header_t, spin_tries) == 32,
120 + "spin_tries must be at offset 32");
121 +_Static_assert(offsetof(nipc_win_shm_region_header_t, req_len) == 36,
122 + "req_len must be at offset 36");
123 +_Static_assert(offsetof(nipc_win_shm_region_header_t, resp_len) == 40,
124 + "resp_len must be at offset 40");
125 +_Static_assert(offsetof(nipc_win_shm_region_header_t, req_client_closed) == 44,
126 + "req_client_closed must be at offset 44");
127 +_Static_assert(offsetof(nipc_win_shm_region_header_t, req_server_waiting) == 48,
128 + "req_server_waiting must be at offset 48");
129 +_Static_assert(offsetof(nipc_win_shm_region_header_t, resp_server_closed) == 52,
130 + "resp_server_closed must be at offset 52");
131 +_Static_assert(offsetof(nipc_win_shm_region_header_t, resp_client_waiting) == 56,
132 + "resp_client_waiting must be at offset 56");
133 +_Static_assert(offsetof(nipc_win_shm_region_header_t, req_seq) == 64,
134 + "req_seq must be at offset 64");
135 +_Static_assert(offsetof(nipc_win_shm_region_header_t, resp_seq) == 72,
136 + "resp_seq must be at offset 72");
137 +
138 +/* ------------------------------------------------------------------ */
139 +/* SHM context */
140 +/* ------------------------------------------------------------------ */
141 +
142 +typedef struct {
143 + nipc_win_shm_role_t role;
144 +
145 + HANDLE mapping; /* file mapping handle */
146 + void *base; /* MapViewOfFile base pointer */
147 + size_t region_size; /* total mapped size */
148 +
149 + /* Kernel events (SHM_HYBRID only; INVALID_HANDLE_VALUE for BUSYWAIT) */
150 + HANDLE req_event;
151 + HANDLE resp_event;
152 +
153 + /* Cached from header */
154 + uint32_t profile;
155 + uint32_t request_offset;
156 + uint32_t request_capacity;
157 + uint32_t response_offset;
158 + uint32_t response_capacity;
159 + uint32_t spin_tries;
160 +
161 + /* Sequence tracking */
162 + LONG64 local_req_seq;
163 + LONG64 local_resp_seq;
164 +
165 +} nipc_win_shm_ctx_t;
166 +
167 +/* ------------------------------------------------------------------ */
168 +/* Server API */
169 +/* ------------------------------------------------------------------ */
170 +
171 +/*
172 + * Create a per-session Windows SHM region.
173 + *
174 + * The kernel objects are named using the FNV-1a hash of
175 + * run_dir + "\n" + service_name + "\n" + auth_token_decimal,
176 + * plus the session_id for per-session isolation.
177 + *
178 + * session_id: server-assigned session identifier (from hello-ack).
179 + * profile: NIPC_WIN_SHM_PROFILE_HYBRID or NIPC_WIN_SHM_PROFILE_BUSYWAIT.
180 + * req_capacity / resp_capacity: data area sizes in bytes.
181 + */
182 +nipc_win_shm_error_t nipc_win_shm_server_create(
183 + const char *run_dir,
184 + const char *service_name,
185 + uint64_t auth_token,
186 + uint64_t session_id,
187 + uint32_t profile,
188 + uint32_t req_capacity,
189 + uint32_t resp_capacity,
190 + nipc_win_shm_ctx_t *ctx);
191 +
192 +/*
193 + * Destroy a server SHM region.
194 + * Sets close flags, signals events, unmaps, closes handles.
195 + */
196 +void nipc_win_shm_destroy(nipc_win_shm_ctx_t *ctx);
197 +
198 +/* ------------------------------------------------------------------ */
199 +/* Client API */
200 +/* ------------------------------------------------------------------ */
201 +
202 +/*
203 + * Attach to an existing per-session Windows SHM region.
204 + * session_id is from the hello-ack received during handshake.
205 + * Validates header and opens event handles.
206 + */
207 +nipc_win_shm_error_t nipc_win_shm_client_attach(
208 + const char *run_dir,
209 + const char *service_name,
210 + uint64_t auth_token,
211 + uint64_t session_id,
212 + uint32_t profile,
213 + nipc_win_shm_ctx_t *ctx);
214 +
215 +/*
216 + * Cleanup stale Windows SHM kernel objects.
217 + * On Windows, kernel objects are reference-counted and auto-cleaned
218 + * when all handles close, so this is a no-op. Provided for API
219 + * symmetry with the POSIX SHM transport.
220 + */
221 +void nipc_win_shm_cleanup_stale(const char *run_dir, const char *service_name);
222 +
223 +/*
224 + * Close a client SHM context.
225 + * Sets close flags, signals events, unmaps, closes handles.
226 + */
227 +void nipc_win_shm_close(nipc_win_shm_ctx_t *ctx);
228 +
229 +/* ------------------------------------------------------------------ */
230 +/* Data plane */
231 +/* ------------------------------------------------------------------ */
232 +
233 +/*
234 + * Publish a message into the SHM region.
235 + * Client sends to request area; server sends to response area.
236 + * msg must include the 32-byte outer header + payload.
237 + */
238 +nipc_win_shm_error_t nipc_win_shm_send(
239 + nipc_win_shm_ctx_t *ctx,
240 + const void *msg,
241 + size_t msg_len);
242 +
243 +/*
244 + * Receive a message from the SHM region.
245 + * Server reads from request area; client reads from response area.
246 + * Copies into caller-provided buffer. Returns message length in *msg_len_out.
247 + */
248 +nipc_win_shm_error_t nipc_win_shm_receive(
249 + nipc_win_shm_ctx_t *ctx,
250 + void *buf,
251 + size_t buf_size,
252 + size_t *msg_len_out,
253 + uint32_t timeout_ms);
254 +
255 +#ifdef NIPC_INTERNAL_TESTING
256 +typedef enum {
257 + NIPC_WIN_SHM_TEST_FAULT_NONE = 0,
258 + NIPC_WIN_SHM_TEST_FAULT_CREATE_MAPPING,
259 + NIPC_WIN_SHM_TEST_FAULT_OPEN_MAPPING,
260 + NIPC_WIN_SHM_TEST_FAULT_MAP_VIEW,
261 + NIPC_WIN_SHM_TEST_FAULT_CREATE_EVENT,
262 + NIPC_WIN_SHM_TEST_FAULT_OPEN_EVENT,
263 +} nipc_win_shm_test_fault_site_t;
264 +
265 +void nipc_win_shm_test_fault_set(nipc_win_shm_test_fault_site_t site,
266 + DWORD error_code,
267 + uint32_t skip_matches);
268 +void nipc_win_shm_test_fault_clear(void);
269 +#endif
270 +
271 +#ifdef __cplusplus
272 +}
273 +#endif
274 +
275 +#endif /* _WIN32 || __MSYS__ */
276 +
277 +#endif /* NETIPC_WIN_SHM_H */
src/libnetdata/netipc/netipc_netdata.c new
+14
@@ -0,0 +1,14 @@
1 +// SPDX-License-Identifier: GPL-3.0-or-later
2 +
3 +#include "netipc_netdata.h"
4 +
5 +#if defined(OS_LINUX) || defined(OS_WINDOWS)
6 +
7 +#include "libnetdata/xxHash/xxhash.h"
8 +
9 +uint64_t netipc_auth_token(void) {
10 + ND_UUID id = nd_log_get_invocation_id();
11 + return XXH3_64bits(id.uuid, sizeof(id.uuid));
12 +}
13 +
14 +#endif // OS_LINUX || OS_WINDOWS
src/libnetdata/netipc/netipc_netdata.h new
+17
@@ -0,0 +1,17 @@
1 +// SPDX-License-Identifier: GPL-3.0-or-later
2 +
3 +#ifndef NETDATA_NETIPC_NETDATA_H
4 +#define NETDATA_NETIPC_NETDATA_H 1
5 +
6 +#include "libnetdata/libnetdata.h"
7 +
8 +#if defined(OS_LINUX) || defined(OS_WINDOWS)
9 +#include "netipc/netipc_service.h"
10 +#include "netipc/netipc_protocol.h"
11 +
12 +// derive a netipc auth token from the current NETDATA_INVOCATION_ID
13 +uint64_t netipc_auth_token(void);
14 +
15 +#endif // OS_LINUX || OS_WINDOWS
16 +
17 +#endif // NETDATA_NETIPC_NETDATA_H
src/libnetdata/netipc/src/protocol/netipc_protocol.c new
+866
@@ -0,0 +1,866 @@
1 +/*
2 + * netipc_protocol.c - Wire envelope and codec implementation.
3 + *
4 + * Localhost-only IPC — struct layouts match wire format exactly.
5 + * Encode/decode uses direct memcpy (single copy per struct).
6 + * No endianness conversion — both peers share host byte order.
7 + */
8 +
9 +#include "netipc/netipc_protocol.h"
10 +#include <stddef.h>
11 +#include <string.h>
12 +
13 +/*
14 + * Safe multiplication check: returns true if count * entry_size would
15 + * overflow size_t. Portable across 32-bit and 64-bit without triggering
16 + * -Wtype-limits.
17 + */
18 +static inline bool mul_would_overflow(size_t count, size_t entry_size) {
19 + return entry_size != 0 && count > SIZE_MAX / entry_size;
20 +}
21 +
22 +/* ------------------------------------------------------------------ */
23 +/* Compile-time layout assertions */
24 +/* */
25 +/* These guarantee struct layouts match wire format exactly, so */
26 +/* encode/decode can be a single memcpy. */
27 +/* ------------------------------------------------------------------ */
28 +
29 +/* Outer message header (32 bytes) */
30 +_Static_assert(sizeof(nipc_header_t) == 32,
31 + "nipc_header_t must be 32 bytes");
32 +_Static_assert(offsetof(nipc_header_t, magic) == 0, "");
33 +_Static_assert(offsetof(nipc_header_t, version) == 4, "");
34 +_Static_assert(offsetof(nipc_header_t, header_len) == 6, "");
35 +_Static_assert(offsetof(nipc_header_t, kind) == 8, "");
36 +_Static_assert(offsetof(nipc_header_t, flags) == 10, "");
37 +_Static_assert(offsetof(nipc_header_t, code) == 12, "");
38 +_Static_assert(offsetof(nipc_header_t, transport_status) == 14, "");
39 +_Static_assert(offsetof(nipc_header_t, payload_len) == 16, "");
40 +_Static_assert(offsetof(nipc_header_t, item_count) == 20, "");
41 +_Static_assert(offsetof(nipc_header_t, message_id) == 24, "");
42 +
43 +/* Chunk continuation header (32 bytes) */
44 +_Static_assert(sizeof(nipc_chunk_header_t) == 32,
45 + "nipc_chunk_header_t must be 32 bytes");
46 +_Static_assert(offsetof(nipc_chunk_header_t, magic) == 0, "");
47 +_Static_assert(offsetof(nipc_chunk_header_t, version) == 4, "");
48 +_Static_assert(offsetof(nipc_chunk_header_t, flags) == 6, "");
49 +_Static_assert(offsetof(nipc_chunk_header_t, message_id) == 8, "");
50 +_Static_assert(offsetof(nipc_chunk_header_t, total_message_len) == 16, "");
51 +_Static_assert(offsetof(nipc_chunk_header_t, chunk_index) == 20, "");
52 +_Static_assert(offsetof(nipc_chunk_header_t, chunk_count) == 24, "");
53 +_Static_assert(offsetof(nipc_chunk_header_t, chunk_payload_len) == 28, "");
54 +
55 +/* Batch entry (8 bytes) */
56 +_Static_assert(sizeof(nipc_batch_entry_t) == 8,
57 + "nipc_batch_entry_t must be 8 bytes");
58 +_Static_assert(offsetof(nipc_batch_entry_t, offset) == 0, "");
59 +_Static_assert(offsetof(nipc_batch_entry_t, length) == 4, "");
60 +
61 +/* Hello (wire = 44, sizeof may be 48 due to trailing alignment) */
62 +_Static_assert(offsetof(nipc_hello_t, layout_version) == 0, "");
63 +_Static_assert(offsetof(nipc_hello_t, flags) == 2, "");
64 +_Static_assert(offsetof(nipc_hello_t, supported_profiles) == 4, "");
65 +_Static_assert(offsetof(nipc_hello_t, preferred_profiles) == 8, "");
66 +_Static_assert(offsetof(nipc_hello_t, max_request_payload_bytes) == 12, "");
67 +_Static_assert(offsetof(nipc_hello_t, max_request_batch_items) == 16, "");
68 +_Static_assert(offsetof(nipc_hello_t, max_response_payload_bytes) == 20, "");
69 +_Static_assert(offsetof(nipc_hello_t, max_response_batch_items) == 24, "");
70 +_Static_assert(offsetof(nipc_hello_t, _reserved) == 28, "");
71 +_Static_assert(offsetof(nipc_hello_t, auth_token) == 32, "");
72 +_Static_assert(offsetof(nipc_hello_t, packet_size) == 40, "");
73 +
74 +/* Hello-ack (48 bytes) */
75 +_Static_assert(sizeof(nipc_hello_ack_t) == 48,
76 + "nipc_hello_ack_t must be 48 bytes");
77 +_Static_assert(offsetof(nipc_hello_ack_t, layout_version) == 0, "");
78 +_Static_assert(offsetof(nipc_hello_ack_t, flags) == 2, "");
79 +_Static_assert(offsetof(nipc_hello_ack_t, server_supported_profiles) == 4, "");
80 +_Static_assert(offsetof(nipc_hello_ack_t, intersection_profiles) == 8, "");
81 +_Static_assert(offsetof(nipc_hello_ack_t, selected_profile) == 12, "");
82 +_Static_assert(offsetof(nipc_hello_ack_t, agreed_max_request_payload_bytes) == 16, "");
83 +_Static_assert(offsetof(nipc_hello_ack_t, agreed_max_request_batch_items) == 20, "");
84 +_Static_assert(offsetof(nipc_hello_ack_t, agreed_max_response_payload_bytes) == 24, "");
85 +_Static_assert(offsetof(nipc_hello_ack_t, agreed_max_response_batch_items) == 28, "");
86 +_Static_assert(offsetof(nipc_hello_ack_t, agreed_packet_size) == 32, "");
87 +_Static_assert(offsetof(nipc_hello_ack_t, _reserved) == 36, "");
88 +_Static_assert(offsetof(nipc_hello_ack_t, session_id) == 40, "");
89 +
90 +/* Cgroups snapshot response header (24 bytes) */
91 +_Static_assert(sizeof(nipc_cgroups_resp_header_t) == 24,
92 + "nipc_cgroups_resp_header_t must be 24 bytes");
93 +_Static_assert(offsetof(nipc_cgroups_resp_header_t, layout_version) == 0, "");
94 +_Static_assert(offsetof(nipc_cgroups_resp_header_t, flags) == 2, "");
95 +_Static_assert(offsetof(nipc_cgroups_resp_header_t, item_count) == 4, "");
96 +_Static_assert(offsetof(nipc_cgroups_resp_header_t, systemd_enabled) == 8, "");
97 +_Static_assert(offsetof(nipc_cgroups_resp_header_t, reserved) == 12, "");
98 +_Static_assert(offsetof(nipc_cgroups_resp_header_t, generation) == 16, "");
99 +
100 +/* Cgroups item wire header (internal, 32 bytes) */
101 +typedef struct {
102 + uint16_t layout_version;
103 + uint16_t flags;
104 + uint32_t hash;
105 + uint32_t options;
106 + uint32_t enabled;
107 + uint32_t name_offset;
108 + uint32_t name_length;
109 + uint32_t path_offset;
110 + uint32_t path_length;
111 +} nipc_cgroups_item_wire_t;
112 +
113 +_Static_assert(sizeof(nipc_cgroups_item_wire_t) == 32,
114 + "nipc_cgroups_item_wire_t must be 32 bytes");
115 +
116 +/* Cgroups request (4 bytes) */
117 +_Static_assert(sizeof(nipc_cgroups_req_t) == 4,
118 + "nipc_cgroups_req_t must be 4 bytes");
119 +
120 +/* ------------------------------------------------------------------ */
121 +/* Outer message header (32 bytes) */
122 +/* ------------------------------------------------------------------ */
123 +
124 +size_t nipc_header_encode(const nipc_header_t *hdr, void *buf, size_t buf_len) {
125 + if (buf_len < NIPC_HEADER_LEN)
126 + return 0;
127 +
128 + memcpy(buf, hdr, NIPC_HEADER_LEN);
129 + return NIPC_HEADER_LEN;
130 +}
131 +
132 +nipc_error_t nipc_header_decode(const void *buf, size_t buf_len,
133 + nipc_header_t *out) {
134 + if (buf_len < NIPC_HEADER_LEN)
135 + return NIPC_ERR_TRUNCATED;
136 +
137 + memcpy(out, buf, NIPC_HEADER_LEN);
138 +
139 + if (out->magic != NIPC_MAGIC_MSG)
140 + return NIPC_ERR_BAD_MAGIC;
141 + if (out->version != NIPC_VERSION)
142 + return NIPC_ERR_BAD_VERSION;
143 + if (out->header_len != NIPC_HEADER_LEN)
144 + return NIPC_ERR_BAD_HEADER_LEN;
145 + if (out->kind < NIPC_KIND_REQUEST || out->kind > NIPC_KIND_CONTROL)
146 + return NIPC_ERR_BAD_KIND;
147 +
148 + return NIPC_OK;
149 +}
150 +
151 +/* ------------------------------------------------------------------ */
152 +/* Chunk continuation header (32 bytes) */
153 +/* ------------------------------------------------------------------ */
154 +
155 +size_t nipc_chunk_header_encode(const nipc_chunk_header_t *chk,
156 + void *buf, size_t buf_len) {
157 + if (buf_len < NIPC_HEADER_LEN)
158 + return 0;
159 +
160 + memcpy(buf, chk, NIPC_HEADER_LEN);
161 + return NIPC_HEADER_LEN;
162 +}
163 +
164 +nipc_error_t nipc_chunk_header_decode(const void *buf, size_t buf_len,
165 + nipc_chunk_header_t *out) {
166 + if (buf_len < NIPC_HEADER_LEN)
167 + return NIPC_ERR_TRUNCATED;
168 +
169 + memcpy(out, buf, NIPC_HEADER_LEN);
170 +
171 + if (out->magic != NIPC_MAGIC_CHUNK)
172 + return NIPC_ERR_BAD_MAGIC;
173 + if (out->version != NIPC_VERSION)
174 + return NIPC_ERR_BAD_VERSION;
175 + if (out->flags != 0)
176 + return NIPC_ERR_BAD_LAYOUT;
177 + if (out->chunk_payload_len == 0)
178 + return NIPC_ERR_BAD_LAYOUT;
179 +
180 + return NIPC_OK;
181 +}
182 +
183 +/* ------------------------------------------------------------------ */
184 +/* Batch item directory */
185 +/* ------------------------------------------------------------------ */
186 +
187 +size_t nipc_batch_dir_encode(const nipc_batch_entry_t *entries,
188 + uint32_t item_count,
189 + void *buf, size_t buf_len) {
190 + size_t need = (size_t)item_count * sizeof(nipc_batch_entry_t);
191 + if (buf_len < need)
192 + return 0;
193 +
194 + memcpy(buf, entries, need);
195 + return need;
196 +}
197 +
198 +nipc_error_t nipc_batch_dir_decode(const void *buf, size_t buf_len,
199 + uint32_t item_count,
200 + uint32_t packed_area_len,
201 + nipc_batch_entry_t *out) {
202 + if (mul_would_overflow((size_t)item_count, sizeof(nipc_batch_entry_t)))
203 + return NIPC_ERR_BAD_ITEM_COUNT;
204 + size_t dir_size = (size_t)item_count * sizeof(nipc_batch_entry_t);
205 + if (buf_len < dir_size)
206 + return NIPC_ERR_TRUNCATED;
207 +
208 + memcpy(out, buf, dir_size);
209 +
210 + for (uint32_t i = 0; i < item_count; i++) {
211 + if (out[i].offset % NIPC_ALIGNMENT != 0)
212 + return NIPC_ERR_BAD_ALIGNMENT;
213 + if ((uint64_t)out[i].offset + out[i].length > packed_area_len)
214 + return NIPC_ERR_OUT_OF_BOUNDS;
215 + }
216 + return NIPC_OK;
217 +}
218 +
219 +nipc_error_t nipc_batch_dir_validate(const void *buf, size_t buf_len,
220 + uint32_t item_count,
221 + uint32_t packed_area_len) {
222 + if (mul_would_overflow((size_t)item_count, sizeof(nipc_batch_entry_t)))
223 + return NIPC_ERR_BAD_ITEM_COUNT;
224 + size_t dir_size = (size_t)item_count * sizeof(nipc_batch_entry_t);
225 + if (buf_len < dir_size)
226 + return NIPC_ERR_TRUNCATED;
227 +
228 + const uint8_t *p = (const uint8_t *)buf;
229 + for (uint32_t i = 0; i < item_count; i++) {
230 + uint32_t off, len;
231 + memcpy(&off, p + i * 8, 4);
232 + memcpy(&len, p + i * 8 + 4, 4);
233 + if (off % NIPC_ALIGNMENT != 0)
234 + return NIPC_ERR_BAD_ALIGNMENT;
235 + if ((uint64_t)off + len > packed_area_len)
236 + return NIPC_ERR_OUT_OF_BOUNDS;
237 + }
238 + return NIPC_OK;
239 +}
240 +
241 +nipc_error_t nipc_batch_item_get(const void *payload, size_t payload_len,
242 + uint32_t item_count, uint32_t index,
243 + const void **item_ptr, uint32_t *item_len) {
244 + if (index >= item_count)
245 + return NIPC_ERR_OUT_OF_BOUNDS;
246 +
247 + if (mul_would_overflow((size_t)item_count, sizeof(nipc_batch_entry_t)))
248 + return NIPC_ERR_BAD_ITEM_COUNT;
249 + size_t dir_size = (size_t)item_count * sizeof(nipc_batch_entry_t);
250 + size_t dir_aligned = nipc_align8(dir_size);
251 +
252 + if (payload_len < dir_aligned)
253 + return NIPC_ERR_TRUNCATED;
254 +
255 + nipc_batch_entry_t entry;
256 + memcpy(&entry, (const uint8_t *)payload + index * sizeof(entry), sizeof(entry));
257 +
258 + size_t packed_area_start = dir_aligned;
259 + size_t packed_area_len = payload_len - packed_area_start;
260 +
261 + if (entry.offset % NIPC_ALIGNMENT != 0)
262 + return NIPC_ERR_BAD_ALIGNMENT;
263 + if ((uint64_t)entry.offset + entry.length > packed_area_len)
264 + return NIPC_ERR_OUT_OF_BOUNDS;
265 +
266 + *item_ptr = (const uint8_t *)payload + packed_area_start + entry.offset;
267 + *item_len = entry.length;
268 + return NIPC_OK;
269 +}
270 +
271 +/* ------------------------------------------------------------------ */
272 +/* Batch builder */
273 +/* ------------------------------------------------------------------ */
274 +
275 +void nipc_batch_builder_init(nipc_batch_builder_t *b,
276 + void *buf, size_t buf_len,
277 + uint32_t max_items) {
278 + b->buf = (uint8_t *)buf;
279 + b->buf_len = buf_len;
280 + b->item_count = 0;
281 + b->max_items = max_items;
282 + b->dir_end = nipc_align8((size_t)max_items * sizeof(nipc_batch_entry_t));
283 + b->data_offset = 0;
284 +}
285 +
286 +nipc_error_t nipc_batch_builder_add(nipc_batch_builder_t *b,
287 + const void *item, size_t item_len) {
288 + if (b->item_count >= b->max_items)
289 + return NIPC_ERR_OVERFLOW;
290 +
291 + size_t aligned_off = nipc_align8(b->data_offset);
292 + size_t abs_pos = b->dir_end + aligned_off;
293 +
294 + if (abs_pos + item_len > b->buf_len)
295 + return NIPC_ERR_OVERFLOW;
296 +
297 + /* Zero alignment padding */
298 + if (aligned_off > b->data_offset)
299 + memset(b->buf + b->dir_end + b->data_offset, 0,
300 + aligned_off - b->data_offset);
301 +
302 + memcpy(b->buf + abs_pos, item, item_len);
303 +
304 + /* Write directory entry */
305 + nipc_batch_entry_t entry = {
306 + .offset = (uint32_t)aligned_off,
307 + .length = (uint32_t)item_len,
308 + };
309 + memcpy(b->buf + b->item_count * sizeof(entry), &entry, sizeof(entry));
310 +
311 + b->data_offset = aligned_off + item_len;
312 + b->item_count++;
313 + return NIPC_OK;
314 +}
315 +
316 +size_t nipc_batch_builder_finish(nipc_batch_builder_t *b,
317 + uint32_t *item_count_out) {
318 + if (item_count_out)
319 + *item_count_out = b->item_count;
320 +
321 + /* The decoder expects: [dir: item_count*8] [align pad] [packed items].
322 + * During building we placed packed data after dir_end = align8(max_items*8).
323 + * If item_count < max_items, compact by shifting packed data left. */
324 + size_t final_dir_aligned = nipc_align8(
325 + (size_t)b->item_count * sizeof(nipc_batch_entry_t));
326 +
327 + if (final_dir_aligned < b->dir_end && b->data_offset > 0) {
328 + memmove(b->buf + final_dir_aligned,
329 + b->buf + b->dir_end,
330 + b->data_offset);
331 + }
332 +
333 + /* Zero trailing alignment padding for deterministic output,
334 + * but only within the caller's buffer. */
335 + size_t data_aligned = nipc_align8(b->data_offset);
336 + if (data_aligned > b->data_offset) {
337 + size_t total_with_pad = final_dir_aligned + data_aligned;
338 + if (total_with_pad <= b->buf_len) {
339 + memset(b->buf + final_dir_aligned + b->data_offset,
340 + 0, data_aligned - b->data_offset);
341 + } else if (final_dir_aligned + b->data_offset < b->buf_len) {
342 + /* Partial padding: zero only within bounds */
343 + memset(b->buf + final_dir_aligned + b->data_offset,
344 + 0, b->buf_len - (final_dir_aligned + b->data_offset));
345 + }
346 + }
347 +
348 + size_t total = final_dir_aligned + data_aligned;
349 + return total <= b->buf_len ? total : final_dir_aligned + b->data_offset;
350 +}
351 +
352 +/* ------------------------------------------------------------------ */
353 +/* Hello payload (44 bytes on wire) */
354 +/* ------------------------------------------------------------------ */
355 +
356 +size_t nipc_hello_encode(const nipc_hello_t *h, void *buf, size_t buf_len) {
357 + if (buf_len < NIPC_HELLO_WIRE_SIZE)
358 + return 0;
359 +
360 + memcpy(buf, h, NIPC_HELLO_WIRE_SIZE);
361 + return NIPC_HELLO_WIRE_SIZE;
362 +}
363 +
364 +nipc_error_t nipc_hello_decode(const void *buf, size_t buf_len,
365 + nipc_hello_t *out) {
366 + if (buf_len < NIPC_HELLO_WIRE_SIZE)
367 + return NIPC_ERR_TRUNCATED;
368 +
369 + memcpy(out, buf, NIPC_HELLO_WIRE_SIZE);
370 +
371 + if (out->layout_version != 1)
372 + return NIPC_ERR_BAD_LAYOUT;
373 + if (out->_reserved != 0)
374 + return NIPC_ERR_BAD_LAYOUT;
375 +
376 + return NIPC_OK;
377 +}
378 +
379 +/* ------------------------------------------------------------------ */
380 +/* Hello-ack payload (48 bytes) */
381 +/* ------------------------------------------------------------------ */
382 +
383 +size_t nipc_hello_ack_encode(const nipc_hello_ack_t *h,
384 + void *buf, size_t buf_len) {
385 + if (buf_len < sizeof(nipc_hello_ack_t))
386 + return 0;
387 +
388 + memcpy(buf, h, sizeof(nipc_hello_ack_t));
389 + return sizeof(nipc_hello_ack_t);
390 +}
391 +
392 +nipc_error_t nipc_hello_ack_decode(const void *buf, size_t buf_len,
393 + nipc_hello_ack_t *out) {
394 + if (buf_len < sizeof(nipc_hello_ack_t))
395 + return NIPC_ERR_TRUNCATED;
396 +
397 + memcpy(out, buf, sizeof(nipc_hello_ack_t));
398 +
399 + if (out->layout_version != 1)
400 + return NIPC_ERR_BAD_LAYOUT;
401 + if (out->flags != 0)
402 + return NIPC_ERR_BAD_LAYOUT;
403 +
404 + return NIPC_OK;
405 +}
406 +
407 +/* ------------------------------------------------------------------ */
408 +/* Cgroups snapshot request (4 bytes) */
409 +/* ------------------------------------------------------------------ */
410 +
411 +size_t nipc_cgroups_req_encode(const nipc_cgroups_req_t *r,
412 + void *buf, size_t buf_len) {
413 + if (buf_len < sizeof(nipc_cgroups_req_t))
414 + return 0;
415 +
416 + memcpy(buf, r, sizeof(nipc_cgroups_req_t));
417 + return sizeof(nipc_cgroups_req_t);
418 +}
419 +
420 +nipc_error_t nipc_cgroups_req_decode(const void *buf, size_t buf_len,
421 + nipc_cgroups_req_t *out) {
422 + if (buf_len < sizeof(nipc_cgroups_req_t))
423 + return NIPC_ERR_TRUNCATED;
424 +
425 + memcpy(out, buf, sizeof(nipc_cgroups_req_t));
426 +
427 + if (out->layout_version != 1)
428 + return NIPC_ERR_BAD_LAYOUT;
429 + if (out->flags != 0)
430 + return NIPC_ERR_BAD_LAYOUT;
431 +
432 + return NIPC_OK;
433 +}
434 +
435 +/* ------------------------------------------------------------------ */
436 +/* Cgroups snapshot response decode */
437 +/* ------------------------------------------------------------------ */
438 +
439 +nipc_error_t nipc_cgroups_resp_decode(const void *buf, size_t buf_len,
440 + nipc_cgroups_resp_view_t *out) {
441 + if (buf_len < NIPC_CGROUPS_RESP_HDR_SIZE)
442 + return NIPC_ERR_TRUNCATED;
443 +
444 + nipc_cgroups_resp_header_t hdr;
445 + memcpy(&hdr, buf, sizeof(hdr));
446 +
447 + if (hdr.layout_version != 1)
448 + return NIPC_ERR_BAD_LAYOUT;
449 + if (hdr.flags != 0)
450 + return NIPC_ERR_BAD_LAYOUT;
451 + if (hdr.reserved != 0)
452 + return NIPC_ERR_BAD_LAYOUT;
453 +
454 + out->layout_version = hdr.layout_version;
455 + out->flags = hdr.flags;
456 + out->item_count = hdr.item_count;
457 + out->systemd_enabled = hdr.systemd_enabled;
458 + out->generation = hdr.generation;
459 +
460 + /* Validate directory fits (with overflow check) */
461 + if (mul_would_overflow((size_t)out->item_count, NIPC_CGROUPS_DIR_ENTRY_SIZE))
462 + return NIPC_ERR_BAD_ITEM_COUNT;
463 + size_t dir_size = (size_t)out->item_count * NIPC_CGROUPS_DIR_ENTRY_SIZE;
464 + size_t dir_end = NIPC_CGROUPS_RESP_HDR_SIZE + dir_size;
465 + if (dir_end > buf_len)
466 + return NIPC_ERR_TRUNCATED;
467 +
468 + size_t packed_area_len = buf_len - dir_end;
469 +
470 + /* Validate each directory entry */
471 + const uint8_t *dir = (const uint8_t *)buf + NIPC_CGROUPS_RESP_HDR_SIZE;
472 + for (uint32_t i = 0; i < out->item_count; i++) {
473 + nipc_batch_entry_t entry;
474 + memcpy(&entry, dir + i * sizeof(entry), sizeof(entry));
475 +
476 + if (entry.offset % NIPC_ALIGNMENT != 0)
477 + return NIPC_ERR_BAD_ALIGNMENT;
478 + if ((uint64_t)entry.offset + entry.length > packed_area_len)
479 + return NIPC_ERR_OUT_OF_BOUNDS;
480 + if (entry.length < NIPC_CGROUPS_ITEM_HDR_SIZE)
481 + return NIPC_ERR_TRUNCATED;
482 + }
483 +
484 + out->_payload = (const uint8_t *)buf;
485 + out->_payload_len = buf_len;
486 + return NIPC_OK;
487 +}
488 +
489 +nipc_error_t nipc_cgroups_resp_item(const nipc_cgroups_resp_view_t *view,
490 + uint32_t index,
491 + nipc_cgroups_item_view_t *out) {
492 + if (index >= view->item_count)
493 + return NIPC_ERR_OUT_OF_BOUNDS;
494 +
495 + /* Overflow already checked in nipc_cgroups_resp_decode, but
496 + * guard defensively since this is a public API. */
497 + if (mul_would_overflow((size_t)view->item_count, NIPC_CGROUPS_DIR_ENTRY_SIZE))
498 + return NIPC_ERR_BAD_ITEM_COUNT;
499 +
500 + size_t dir_start = NIPC_CGROUPS_RESP_HDR_SIZE;
501 + size_t dir_size = (size_t)view->item_count * NIPC_CGROUPS_DIR_ENTRY_SIZE;
502 + size_t packed_area_start = dir_start + dir_size;
503 +
504 + /* Read directory entry */
505 + nipc_batch_entry_t dir_entry;
506 + memcpy(&dir_entry,
507 + view->_payload + dir_start + index * sizeof(dir_entry),
508 + sizeof(dir_entry));
509 +
510 + const uint8_t *item = view->_payload + packed_area_start + dir_entry.offset;
511 + uint32_t item_len = dir_entry.length;
512 +
513 + /* Read the 32-byte item wire header in one copy */
514 + nipc_cgroups_item_wire_t wire;
515 + memcpy(&wire, item, sizeof(wire));
516 +
517 + if (wire.layout_version != 1)
518 + return NIPC_ERR_BAD_LAYOUT;
519 + if (wire.flags != 0)
520 + return NIPC_ERR_BAD_LAYOUT;
521 +
522 + /* Validate name string */
523 + if (wire.name_offset < NIPC_CGROUPS_ITEM_HDR_SIZE)
524 + return NIPC_ERR_OUT_OF_BOUNDS;
525 + if ((uint64_t)wire.name_offset + wire.name_length + 1 > item_len)
526 + return NIPC_ERR_OUT_OF_BOUNDS;
527 + if (item[wire.name_offset + wire.name_length] != '\0')
528 + return NIPC_ERR_MISSING_NUL;
529 +
530 + /* Validate path string */
531 + if (wire.path_offset < NIPC_CGROUPS_ITEM_HDR_SIZE)
532 + return NIPC_ERR_OUT_OF_BOUNDS;
533 + if ((uint64_t)wire.path_offset + wire.path_length + 1 > item_len)
534 + return NIPC_ERR_OUT_OF_BOUNDS;
535 + if (item[wire.path_offset + wire.path_length] != '\0')
536 + return NIPC_ERR_MISSING_NUL;
537 +
538 + /* Reject overlapping name and path regions (including NUL) */
539 + {
540 + uint64_t name_start = wire.name_offset;
541 + uint64_t name_end = name_start + wire.name_length + 1;
542 + uint64_t path_start = wire.path_offset;
543 + uint64_t path_end = path_start + wire.path_length + 1;
544 + if (name_start < path_end && path_start < name_end)
545 + return NIPC_ERR_BAD_LAYOUT;
546 + }
547 +
548 + out->layout_version = wire.layout_version;
549 + out->flags = wire.flags;
550 + out->hash = wire.hash;
551 + out->options = wire.options;
552 + out->enabled = wire.enabled;
553 + out->name.ptr = (const char *)(item + wire.name_offset);
554 + out->name.len = wire.name_length;
555 + out->path.ptr = (const char *)(item + wire.path_offset);
556 + out->path.len = wire.path_length;
557 +
558 + return NIPC_OK;
559 +}
560 +
561 +/* ------------------------------------------------------------------ */
562 +/* Cgroups snapshot response builder */
563 +/* */
564 +/* Layout during building (max_items directory slots reserved): */
565 +/* [24-byte header space] [max_items*8 directory] [packed items] */
566 +/* */
567 +/* Layout after finish (compacted to actual item_count): */
568 +/* [24-byte header] [item_count*8 directory] [packed items] */
569 +/* */
570 +/* If item_count < max_items, finish() shifts packed data left and */
571 +/* adjusts directory offsets accordingly. */
572 +/* ------------------------------------------------------------------ */
573 +
574 +void nipc_cgroups_builder_init(nipc_cgroups_builder_t *b,
575 + void *buf, size_t buf_len,
576 + uint32_t max_items,
577 + uint32_t systemd_enabled,
578 + uint64_t generation) {
579 + b->buf = (uint8_t *)buf;
580 + b->buf_len = buf_len;
581 + b->systemd_enabled = systemd_enabled;
582 + b->generation = generation;
583 + b->item_count = 0;
584 + b->max_items = max_items;
585 + b->error = NIPC_OK;
586 +
587 + /* Packed item data starts after reserved directory */
588 + b->data_offset = NIPC_CGROUPS_RESP_HDR_SIZE +
589 + (size_t)max_items * NIPC_CGROUPS_DIR_ENTRY_SIZE;
590 +}
591 +
592 +void nipc_cgroups_builder_set_header(nipc_cgroups_builder_t *b,
593 + uint32_t systemd_enabled,
594 + uint64_t generation) {
595 + b->systemd_enabled = systemd_enabled;
596 + b->generation = generation;
597 +}
598 +
599 +uint32_t nipc_cgroups_builder_estimate_max_items(size_t buf_len) {
600 + if (buf_len <= NIPC_CGROUPS_RESP_HDR_SIZE)
601 + return 0;
602 +
603 + size_t min_aligned_item = nipc_align8(NIPC_CGROUPS_ITEM_HDR_SIZE + 2u);
604 + return (uint32_t)((buf_len - NIPC_CGROUPS_RESP_HDR_SIZE) /
605 + (NIPC_CGROUPS_DIR_ENTRY_SIZE + min_aligned_item));
606 +}
607 +
608 +nipc_error_t nipc_cgroups_builder_add(nipc_cgroups_builder_t *b,
609 + uint32_t hash,
610 + uint32_t options,
611 + uint32_t enabled,
612 + const char *name, uint32_t name_len,
613 + const char *path, uint32_t path_len) {
614 + if (b->item_count >= b->max_items) {
615 + b->error = NIPC_ERR_OVERFLOW;
616 + return NIPC_ERR_OVERFLOW;
617 + }
618 +
619 + /* Align item start to 8 bytes */
620 + size_t item_start = nipc_align8(b->data_offset);
621 +
622 + /* Item payload: 32-byte header + name + NUL + path + NUL */
623 + size_t item_size = NIPC_CGROUPS_ITEM_HDR_SIZE +
624 + (size_t)name_len + 1 +
625 + (size_t)path_len + 1;
626 +
627 + if (item_start + item_size > b->buf_len) {
628 + b->error = NIPC_ERR_OVERFLOW;
629 + return NIPC_ERR_OVERFLOW;
630 + }
631 +
632 + /* Zero alignment padding */
633 + if (item_start > b->data_offset)
634 + memset(b->buf + b->data_offset, 0, item_start - b->data_offset);
635 +
636 + uint8_t *item = b->buf + item_start;
637 +
638 + /* Write item header as a single struct copy */
639 + nipc_cgroups_item_wire_t wire = {
640 + .layout_version = 1,
641 + .flags = 0,
642 + .hash = hash,
643 + .options = options,
644 + .enabled = enabled,
645 + .name_offset = NIPC_CGROUPS_ITEM_HDR_SIZE,
646 + .name_length = name_len,
647 + .path_offset = NIPC_CGROUPS_ITEM_HDR_SIZE + name_len + 1,
648 + .path_length = path_len,
649 + };
650 + memcpy(item, &wire, sizeof(wire));
651 +
652 + /* Write strings with NUL terminators */
653 + memcpy(item + wire.name_offset, name, name_len);
654 + item[wire.name_offset + name_len] = '\0';
655 + memcpy(item + wire.path_offset, path, path_len);
656 + item[wire.path_offset + path_len] = '\0';
657 +
658 + /* Write directory entry (absolute offset stored temporarily) */
659 + nipc_batch_entry_t dir_entry = {
660 + .offset = (uint32_t)item_start,
661 + .length = (uint32_t)item_size,
662 + };
663 + size_t dir_pos = NIPC_CGROUPS_RESP_HDR_SIZE +
664 + (size_t)b->item_count * NIPC_CGROUPS_DIR_ENTRY_SIZE;
665 + memcpy(b->buf + dir_pos, &dir_entry, sizeof(dir_entry));
666 +
667 + b->data_offset = item_start + item_size;
668 + b->item_count++;
669 + return NIPC_OK;
670 +}
671 +
672 +size_t nipc_cgroups_builder_finish(nipc_cgroups_builder_t *b) {
673 + uint8_t *p = b->buf;
674 +
675 + nipc_cgroups_resp_header_t hdr = {
676 + .layout_version = 1,
677 + .flags = 0,
678 + .item_count = b->item_count,
679 + .systemd_enabled = b->systemd_enabled,
680 + .reserved = 0,
681 + .generation = b->generation,
682 + };
683 +
684 + if (b->item_count == 0) {
685 + memcpy(p, &hdr, sizeof(hdr));
686 + return NIPC_CGROUPS_RESP_HDR_SIZE;
687 + }
688 +
689 + /* Where the decoder expects packed data to start */
690 + size_t final_packed_start = NIPC_CGROUPS_RESP_HDR_SIZE +
691 + (size_t)b->item_count * NIPC_CGROUPS_DIR_ENTRY_SIZE;
692 +
693 + /* Read the first directory entry to find where packed data actually begins */
694 + nipc_batch_entry_t first_entry;
695 + memcpy(&first_entry, p + NIPC_CGROUPS_RESP_HDR_SIZE, sizeof(first_entry));
696 + uint32_t first_item_abs = first_entry.offset;
697 +
698 + /* Guard against underflow if builder state is inconsistent */
699 + if (b->data_offset < first_item_abs) {
700 + hdr.item_count = 0;
701 + memcpy(p, &hdr, sizeof(hdr));
702 + return NIPC_CGROUPS_RESP_HDR_SIZE;
703 + }
704 +
705 + size_t packed_data_len = b->data_offset - first_item_abs;
706 +
707 + if (final_packed_start < first_item_abs) {
708 + memmove(p + final_packed_start, p + first_item_abs, packed_data_len);
709 + }
710 +
711 + /* Convert directory entries from absolute offsets to relative offsets */
712 + size_t dir_base = NIPC_CGROUPS_RESP_HDR_SIZE;
713 + for (uint32_t i = 0; i < b->item_count; i++) {
714 + size_t entry_pos = dir_base + (size_t)i * NIPC_CGROUPS_DIR_ENTRY_SIZE;
715 + nipc_batch_entry_t entry;
716 + memcpy(&entry, p + entry_pos, sizeof(entry));
717 + if (entry.offset < first_item_abs)
718 + continue; /* skip corrupted entry */
719 + entry.offset -= first_item_abs;
720 + memcpy(p + entry_pos, &entry, sizeof(entry));
721 + }
722 +
723 + /* Write snapshot header */
724 + memcpy(p, &hdr, sizeof(hdr));
725 +
726 + return final_packed_start + packed_data_len;
727 +}
728 +
729 +/* ------------------------------------------------------------------ */
730 +/* INCREMENT codec */
731 +/* ------------------------------------------------------------------ */
732 +
733 +size_t nipc_increment_encode(uint64_t value, void *buf, size_t buf_len) {
734 + if (buf_len < NIPC_INCREMENT_PAYLOAD_SIZE)
735 + return 0;
736 + memcpy(buf, &value, 8);
737 + return NIPC_INCREMENT_PAYLOAD_SIZE;
738 +}
739 +
740 +nipc_error_t nipc_increment_decode(const void *buf, size_t buf_len,
741 + uint64_t *value_out) {
742 + if (buf_len < NIPC_INCREMENT_PAYLOAD_SIZE)
743 + return NIPC_ERR_TRUNCATED;
744 + memcpy(value_out, buf, 8);
745 + return NIPC_OK;
746 +}
747 +
748 +/* ------------------------------------------------------------------ */
749 +/* STRING_REVERSE codec */
750 +/* ------------------------------------------------------------------ */
751 +
752 +size_t nipc_string_reverse_encode(const char *str, uint32_t str_len,
753 + void *buf, size_t buf_len) {
754 + /* Guard against size_t overflow on 32-bit platforms */
755 + if (str_len > SIZE_MAX - NIPC_STRING_REVERSE_HDR_SIZE - 1)
756 + return 0;
757 +
758 + size_t total = NIPC_STRING_REVERSE_HDR_SIZE + str_len + 1;
759 + if (buf_len < total)
760 + return 0;
761 +
762 + uint8_t *p = (uint8_t *)buf;
763 + uint32_t offset = NIPC_STRING_REVERSE_HDR_SIZE;
764 + memcpy(p + 0, &offset, 4);
765 + memcpy(p + 4, &str_len, 4);
766 + if (str_len > 0)
767 + memcpy(p + offset, str, str_len);
768 + p[offset + str_len] = '\0';
769 + return total;
770 +}
771 +
772 +nipc_error_t nipc_string_reverse_decode(const void *buf, size_t buf_len,
773 + nipc_string_reverse_view_t *view_out) {
774 + if (buf_len < NIPC_STRING_REVERSE_HDR_SIZE)
775 + return NIPC_ERR_TRUNCATED;
776 +
777 + const uint8_t *p = (const uint8_t *)buf;
778 + uint32_t str_offset, str_length;
779 + memcpy(&str_offset, p + 0, 4);
780 + memcpy(&str_length, p + 4, 4);
781 +
782 + if ((uint64_t)str_offset + str_length + 1 > buf_len)
783 + return NIPC_ERR_OUT_OF_BOUNDS;
784 +
785 + if (p[str_offset + str_length] != '\0')
786 + return NIPC_ERR_MISSING_NUL;
787 +
788 + view_out->str = (const char *)(p + str_offset);
789 + view_out->str_len = str_length;
790 + return NIPC_OK;
791 +}
792 +
793 +/* ------------------------------------------------------------------ */
794 +/* Server-side typed dispatch helpers */
795 +/* ------------------------------------------------------------------ */
796 +
797 +bool nipc_dispatch_increment(
798 + const uint8_t *req, size_t req_len,
799 + uint8_t *resp, size_t resp_size, size_t *resp_len,
800 + nipc_increment_handler_fn handler, void *user)
801 +{
802 + uint64_t value;
803 + if (nipc_increment_decode(req, req_len, &value) != NIPC_OK)
804 + return false;
805 +
806 + uint64_t result;
807 + if (!handler(user, value, &result))
808 + return false;
809 +
810 + *resp_len = nipc_increment_encode(result, resp, resp_size);
811 + return *resp_len > 0;
812 +}
813 +
814 +bool nipc_dispatch_string_reverse(
815 + const uint8_t *req, size_t req_len,
816 + uint8_t *resp, size_t resp_size, size_t *resp_len,
817 + nipc_string_reverse_handler_fn handler, void *user)
818 +{
819 + nipc_string_reverse_view_t view;
820 + if (nipc_string_reverse_decode(req, req_len, &view) != NIPC_OK)
821 + return false;
822 +
823 + /* Provide a scratch buffer for the handler's response string.
824 + * The handler writes the response string into it; we encode after. */
825 + uint32_t capacity = (resp_size > NIPC_STRING_REVERSE_HDR_SIZE + 1)
826 + ? (uint32_t)(resp_size - NIPC_STRING_REVERSE_HDR_SIZE - 1)
827 + : 0;
828 + char *scratch = (char *)(resp + NIPC_STRING_REVERSE_HDR_SIZE);
829 +
830 + uint32_t response_str_len = 0;
831 + if (!handler(user, view.str, view.str_len,
832 + scratch, capacity, &response_str_len))
833 + return false;
834 +
835 + /* Encode from the scratch area (already at the right offset) */
836 + *resp_len = nipc_string_reverse_encode(scratch, response_str_len,
837 + resp, resp_size);
838 + return *resp_len > 0;
839 +}
840 +
841 +nipc_error_t nipc_dispatch_cgroups_snapshot(
842 + const uint8_t *req, size_t req_len,
843 + uint8_t *resp, size_t resp_size, size_t *resp_len,
844 + uint32_t max_items,
845 + nipc_cgroups_handler_fn handler, void *user)
846 +{
847 + nipc_cgroups_req_t request;
848 + nipc_error_t err = nipc_cgroups_req_decode(req, req_len, &request);
849 + if (err != NIPC_OK)
850 + return err;
851 +
852 + nipc_cgroups_builder_t builder;
853 + nipc_cgroups_builder_init(&builder, resp, resp_size, max_items, 0, 0);
854 +
855 + if (!handler(user, &request, &builder)) {
856 + if (builder.error != NIPC_OK)
857 + return builder.error;
858 + return NIPC_ERR_HANDLER_FAILED;
859 + }
860 +
861 + if (builder.error != NIPC_OK)
862 + return builder.error;
863 +
864 + *resp_len = nipc_cgroups_builder_finish(&builder);
865 + return (*resp_len > 0) ? NIPC_OK : NIPC_ERR_OVERFLOW;
866 +}
src/libnetdata/netipc/src/service/netipc_service.c new
+1747
@@ -0,0 +1,1747 @@
1 +/*
2 + * netipc_service.c - L2 orchestration implementation.
3 + *
4 + * Pure composition of L1 (UDS/SHM) + Codec. No direct socket/mmap calls
5 + * for data framing. Uses poll() for timeout-based shutdown detection
6 + * in the managed server.
7 + *
8 + * Client context manages connection lifecycle with at-least-once retry.
9 + * Managed server handles accept, read, dispatch, respond.
10 + */
11 +
12 +#include "netipc/netipc_service.h"
13 +#include "netipc/netipc_protocol.h"
14 +#include "netipc/netipc_uds.h"
15 +#include "netipc/netipc_shm.h"
16 +
17 +#include <errno.h>
18 +#include <poll.h>
19 +#include <stdlib.h>
20 +#include <string.h>
21 +#include <sys/socket.h>
22 +#include <time.h>
23 +#include <unistd.h>
24 +
25 +/* Poll timeout for server loops: 100ms between shutdown checks */
26 +#define SERVER_POLL_TIMEOUT_MS 100
27 +#define NIPC_CLIENT_BUF_DEFAULT 65536u
28 +#define CLIENT_SHM_ATTACH_RETRY_INTERVAL_MS 5u
29 +#define CLIENT_SHM_ATTACH_RETRY_TIMEOUT_MS 5000u
30 +
31 +enum {
32 + NIPC_POSIX_SERVICE_TEST_FAULT_CLIENT_RESPONSE_BUF_REALLOC_INTERNAL = 1,
33 + NIPC_POSIX_SERVICE_TEST_FAULT_CLIENT_SEND_BUF_REALLOC_INTERNAL,
34 + NIPC_POSIX_SERVICE_TEST_FAULT_CLIENT_SHM_CTX_CALLOC_INTERNAL,
35 + NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_SHM_CTX_CALLOC_INTERNAL,
36 + NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_RECV_BUF_MALLOC_INTERNAL,
37 + NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_RESP_BUF_MALLOC_INTERNAL,
38 + NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_SESSIONS_CALLOC_INTERNAL,
39 + NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_SESSION_CTX_CALLOC_INTERNAL,
40 + NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_THREAD_CREATE_INTERNAL,
41 + NIPC_POSIX_SERVICE_TEST_FAULT_CACHE_BUCKETS_CALLOC_INTERNAL,
42 + NIPC_POSIX_SERVICE_TEST_FAULT_CACHE_ITEMS_CALLOC_INTERNAL,
43 + NIPC_POSIX_SERVICE_TEST_FAULT_CACHE_ITEM_NAME_MALLOC_INTERNAL,
44 + NIPC_POSIX_SERVICE_TEST_FAULT_CACHE_ITEM_PATH_MALLOC_INTERNAL,
45 +};
46 +
47 +static uint64_t g_posix_service_test_fault_state = 0;
48 +
49 +static uint64_t posix_service_fault_state_make(int site, uint32_t skip_matches)
50 +{
51 + return ((uint64_t)skip_matches << 32) | (uint32_t)site;
52 +}
53 +
54 +void nipc_posix_service_test_fault_set(int site, uint32_t skip_matches)
55 +{
56 + __atomic_store_n(&g_posix_service_test_fault_state,
57 + posix_service_fault_state_make(site, skip_matches),
58 + __ATOMIC_RELEASE);
59 +}
60 +
61 +void nipc_posix_service_test_fault_clear(void)
62 +{
63 + __atomic_store_n(&g_posix_service_test_fault_state, 0, __ATOMIC_RELEASE);
64 +}
65 +
66 +static bool service_test_should_fail(int site)
67 +{
68 + for (;;) {
69 + uint64_t state = __atomic_load_n(&g_posix_service_test_fault_state,
70 + __ATOMIC_ACQUIRE);
71 + uint32_t current_site = (uint32_t)state;
72 + uint32_t current_skip = (uint32_t)(state >> 32);
73 + uint64_t next_state;
74 +
75 + if (current_site != (uint32_t)site)
76 + return false;
77 +
78 + if (current_skip > 0)
79 + next_state = posix_service_fault_state_make(site, current_skip - 1u);
80 + else
81 + next_state = 0;
82 +
83 + if (__atomic_compare_exchange_n(&g_posix_service_test_fault_state,
84 + &state, next_state, false,
85 + __ATOMIC_ACQ_REL, __ATOMIC_ACQUIRE))
86 + return current_skip == 0;
87 + }
88 +}
89 +
90 +static void *service_malloc(size_t size, int fault_site)
91 +{
92 + if (service_test_should_fail(fault_site))
93 + return NULL;
94 + return malloc(size);
95 +}
96 +
97 +static void *service_calloc(size_t count, size_t size, int fault_site)
98 +{
99 + if (service_test_should_fail(fault_site))
100 + return NULL;
101 + return calloc(count, size);
102 +}
103 +
104 +static void *service_realloc(void *ptr, size_t size, int fault_site)
105 +{
106 + if (service_test_should_fail(fault_site))
107 + return NULL;
108 + return realloc(ptr, size);
109 +}
110 +
111 +static int service_pthread_create(pthread_t *thread,
112 + const pthread_attr_t *attr,
113 + void *(*start_routine)(void *),
114 + void *arg)
115 +{
116 + if (service_test_should_fail(NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_THREAD_CREATE_INTERNAL))
117 + return EAGAIN;
118 + return pthread_create(thread, attr, start_routine, arg);
119 +}
120 +
121 +static uint64_t monotonic_time_ms(void)
122 +{
123 + struct timespec ts;
124 + clock_gettime(CLOCK_MONOTONIC, &ts);
125 + return (uint64_t)ts.tv_sec * 1000u + (uint64_t)ts.tv_nsec / 1000000u;
126 +}
127 +
128 +static uint32_t next_power_of_2_u32(uint32_t n)
129 +{
130 + if (n < 16)
131 + return 16;
132 + n--;
133 + n |= n >> 1;
134 + n |= n >> 2;
135 + n |= n >> 4;
136 + n |= n >> 8;
137 + n |= n >> 16;
138 + return n + 1;
139 +}
140 +
141 +static bool ensure_buffer(uint8_t **buf, size_t *buf_size, size_t need, int fault_site)
142 +{
143 + if (*buf && *buf_size >= need)
144 + return true;
145 +
146 + uint8_t *new_buf = service_realloc(*buf, need, fault_site);
147 + if (!new_buf)
148 + return false;
149 +
150 + *buf = new_buf;
151 + *buf_size = need;
152 + return true;
153 +}
154 +
155 +static void client_note_request_capacity(nipc_client_ctx_t *ctx, uint32_t payload_len)
156 +{
157 + uint32_t grown = next_power_of_2_u32(payload_len);
158 + if (grown > NIPC_MAX_PAYLOAD_CAP)
159 + grown = NIPC_MAX_PAYLOAD_CAP;
160 + if (grown > ctx->transport_config.max_request_payload_bytes)
161 + ctx->transport_config.max_request_payload_bytes = grown;
162 +}
163 +
164 +static void client_note_response_capacity(nipc_client_ctx_t *ctx, uint32_t payload_len)
165 +{
166 + uint32_t grown = next_power_of_2_u32(payload_len);
167 + if (grown > NIPC_MAX_PAYLOAD_CAP)
168 + grown = NIPC_MAX_PAYLOAD_CAP;
169 + if (grown > ctx->transport_config.max_response_payload_bytes)
170 + ctx->transport_config.max_response_payload_bytes = grown;
171 +}
172 +
173 +static bool client_prepare_session_buffers(nipc_client_ctx_t *ctx)
174 +{
175 + size_t response_need = (size_t)ctx->session.max_response_payload_bytes + NIPC_HEADER_LEN;
176 + if (response_need < NIPC_HEADER_LEN + 1024u)
177 + response_need = NIPC_HEADER_LEN + 1024u;
178 +
179 + if (!ensure_buffer(&ctx->response_buf, &ctx->response_buf_size, response_need,
180 + NIPC_POSIX_SERVICE_TEST_FAULT_CLIENT_RESPONSE_BUF_REALLOC_INTERNAL))
181 + return false;
182 +
183 + if (ctx->session.selected_profile == NIPC_PROFILE_SHM_HYBRID ||
184 + ctx->session.selected_profile == NIPC_PROFILE_SHM_FUTEX) {
185 + size_t send_need = (size_t)ctx->session.max_request_payload_bytes + NIPC_HEADER_LEN;
186 + if (!ensure_buffer(&ctx->send_buf, &ctx->send_buf_size, send_need,
187 + NIPC_POSIX_SERVICE_TEST_FAULT_CLIENT_SEND_BUF_REALLOC_INTERNAL))
188 + return false;
189 + }
190 +
191 + return true;
192 +}
193 +
194 +static uint32_t cgroups_request_payload_default(void)
195 +{
196 + return 16u;
197 +}
198 +
199 +static uint32_t cgroups_response_payload_default(void)
200 +{
201 + return NIPC_CLIENT_BUF_DEFAULT;
202 +}
203 +
204 +static nipc_uds_client_config_t service_client_config_to_transport(
205 + const nipc_client_config_t *config)
206 +{
207 + nipc_uds_client_config_t transport = {0};
208 +
209 + if (!config)
210 + return transport;
211 +
212 + transport.supported_profiles = config->supported_profiles;
213 + transport.preferred_profiles = config->preferred_profiles;
214 + transport.max_request_batch_items = config->max_request_batch_items;
215 + transport.max_response_payload_bytes = config->max_response_payload_bytes;
216 + transport.max_response_batch_items = config->max_request_batch_items;
217 + transport.auth_token = config->auth_token;
218 +
219 + return transport;
220 +}
221 +
222 +static nipc_uds_server_config_t service_server_config_to_transport(
223 + const nipc_server_config_t *config)
224 +{
225 + nipc_uds_server_config_t transport = {0};
226 +
227 + if (!config)
228 + return transport;
229 +
230 + transport.supported_profiles = config->supported_profiles;
231 + transport.preferred_profiles = config->preferred_profiles;
232 + transport.max_request_batch_items = config->max_request_batch_items;
233 + transport.max_response_payload_bytes = config->max_response_payload_bytes;
234 + transport.max_response_batch_items = config->max_request_batch_items;
235 + transport.auth_token = config->auth_token;
236 +
237 + return transport;
238 +}
239 +
240 +/* ------------------------------------------------------------------ */
241 +/* Internal: client connection helpers */
242 +/* ------------------------------------------------------------------ */
243 +
244 +/* Tear down the current connection (UDS session + SHM if any). */
245 +static void client_disconnect(nipc_client_ctx_t *ctx)
246 +{
247 + if (ctx->shm) {
248 + nipc_shm_close(ctx->shm);
249 + free(ctx->shm);
250 + ctx->shm = NULL;
251 + }
252 +
253 + if (ctx->session_valid) {
254 + nipc_uds_close_session(&ctx->session);
255 + ctx->session_valid = false;
256 + }
257 +}
258 +
259 +static void client_disable_shm_profiles(nipc_client_ctx_t *ctx)
260 +{
261 + ctx->transport_config.supported_profiles &=
262 + ~(NIPC_PROFILE_SHM_HYBRID | NIPC_PROFILE_SHM_FUTEX);
263 + ctx->transport_config.preferred_profiles &=
264 + ~(NIPC_PROFILE_SHM_HYBRID | NIPC_PROFILE_SHM_FUTEX);
265 +}
266 +
267 +/* Attempt a full connection: UDS connect + handshake, then SHM upgrade
268 + * if negotiated. Returns the new state. */
269 +static nipc_client_state_t client_try_connect(nipc_client_ctx_t *ctx)
270 +{
271 + nipc_uds_session_t session;
272 + memset(&session, 0, sizeof(session));
273 + session.fd = -1;
274 +
275 + nipc_uds_error_t err = nipc_uds_connect(
276 + ctx->run_dir, ctx->service_name,
277 + &ctx->transport_config, &session);
278 +
279 + switch (err) {
280 + case NIPC_UDS_OK:
281 + break;
282 + case NIPC_UDS_ERR_CONNECT:
283 + return NIPC_CLIENT_NOT_FOUND;
284 + case NIPC_UDS_ERR_AUTH_FAILED:
285 + return NIPC_CLIENT_AUTH_FAILED;
286 + case NIPC_UDS_ERR_NO_PROFILE:
287 + case NIPC_UDS_ERR_INCOMPATIBLE:
288 + return NIPC_CLIENT_INCOMPATIBLE;
289 + default:
290 + return NIPC_CLIENT_DISCONNECTED;
291 + }
292 +
293 + ctx->session = session;
294 + ctx->session_valid = true;
295 +
296 + if (!client_prepare_session_buffers(ctx)) {
297 + nipc_uds_close_session(&ctx->session);
298 + ctx->session_valid = false;
299 + return NIPC_CLIENT_DISCONNECTED;
300 + }
301 +
302 + /* SHM upgrade if negotiated */
303 + if (session.selected_profile == NIPC_PROFILE_SHM_HYBRID ||
304 + session.selected_profile == NIPC_PROFILE_SHM_FUTEX) {
305 +
306 + nipc_shm_ctx_t *shm = service_calloc(
307 + 1, sizeof(nipc_shm_ctx_t),
308 + NIPC_POSIX_SERVICE_TEST_FAULT_CLIENT_SHM_CTX_CALLOC_INTERNAL);
309 + if (!shm) {
310 + nipc_uds_close_session(&ctx->session);
311 + ctx->session_valid = false;
312 + return NIPC_CLIENT_DISCONNECTED;
313 + }
314 + {
315 + /* Retry attach: server creates the SHM region after
316 + * the UDS handshake, so it may not exist yet. */
317 + nipc_shm_error_t serr = NIPC_SHM_ERR_NOT_READY;
318 + uint64_t deadline_ms = monotonic_time_ms() + CLIENT_SHM_ATTACH_RETRY_TIMEOUT_MS;
319 + for (;;) {
320 + serr = nipc_shm_client_attach(
321 + ctx->run_dir, ctx->service_name,
322 + session.session_id, shm);
323 + if (serr == NIPC_SHM_OK)
324 + break;
325 + if (serr != NIPC_SHM_ERR_NOT_READY &&
326 + serr != NIPC_SHM_ERR_OPEN &&
327 + serr != NIPC_SHM_ERR_BAD_MAGIC)
328 + break;
329 + if (monotonic_time_ms() >= deadline_ms)
330 + break;
331 + usleep(CLIENT_SHM_ATTACH_RETRY_INTERVAL_MS * 1000u);
332 + }
333 +
334 + if (serr == NIPC_SHM_OK) {
335 + ctx->shm = shm;
336 + } else {
337 + /* SHM attach failed after negotiation. Close that session,
338 + * blacklist SHM for this client context, and retry
339 + * baseline via a new handshake. */
340 + free(shm);
341 + nipc_uds_close_session(&ctx->session);
342 + ctx->session_valid = false;
343 + client_disable_shm_profiles(ctx);
344 + if (ctx->transport_config.supported_profiles == 0)
345 + return NIPC_CLIENT_DISCONNECTED;
346 + return client_try_connect(ctx);
347 + }
348 + }
349 + }
350 +
351 + return NIPC_CLIENT_READY;
352 +}
353 +
354 +/* ------------------------------------------------------------------ */
355 +/* Internal: send/receive via the active transport */
356 +/* ------------------------------------------------------------------ */
357 +
358 +/*
359 + * Send a complete message (header + payload) using whichever transport
360 + * is active: SHM if negotiated, UDS otherwise.
361 + */
362 +static nipc_error_t transport_send(nipc_client_ctx_t *ctx,
363 + nipc_header_t *hdr,
364 + const void *payload,
365 + size_t payload_len)
366 +{
367 + if (payload_len > UINT32_MAX)
368 + return NIPC_ERR_OVERFLOW;
369 +
370 + if (ctx->shm) {
371 + if (payload_len > ctx->session.max_request_payload_bytes) {
372 + client_note_request_capacity(ctx, (uint32_t)payload_len);
373 + return NIPC_ERR_OVERFLOW;
374 + }
375 +
376 + size_t msg_len = NIPC_HEADER_LEN + payload_len;
377 + uint8_t *msg = ctx->send_buf;
378 + if (!msg || msg_len > ctx->send_buf_size)
379 + return NIPC_ERR_OVERFLOW;
380 +
381 + hdr->magic = NIPC_MAGIC_MSG;
382 + hdr->version = NIPC_VERSION;
383 + hdr->header_len = NIPC_HEADER_LEN;
384 + hdr->payload_len = (uint32_t)payload_len;
385 +
386 + nipc_header_encode(hdr, msg, NIPC_HEADER_LEN);
387 + if (payload_len > 0)
388 + memcpy(msg + NIPC_HEADER_LEN, payload, payload_len);
389 +
390 + nipc_shm_error_t serr = nipc_shm_send(ctx->shm, msg, msg_len);
391 + if (serr == NIPC_SHM_ERR_MSG_TOO_LARGE) {
392 + client_note_request_capacity(ctx, (uint32_t)payload_len);
393 + return NIPC_ERR_OVERFLOW;
394 + }
395 + return (serr == NIPC_SHM_OK) ? NIPC_OK : NIPC_ERR_NOT_READY;
396 + }
397 +
398 + /* UDS path */
399 + nipc_uds_error_t uerr = nipc_uds_send(&ctx->session, hdr,
400 + payload, payload_len);
401 + if (uerr == NIPC_UDS_ERR_LIMIT_EXCEEDED) {
402 + client_note_request_capacity(ctx, (uint32_t)payload_len);
403 + return NIPC_ERR_OVERFLOW;
404 + }
405 + return (uerr == NIPC_UDS_OK) ? NIPC_OK : NIPC_ERR_NOT_READY;
406 +}
407 +
408 +/*
409 + * Receive a complete message. For SHM, reads from the SHM region.
410 + * For UDS, reads from the socket into the caller's buffer.
411 + *
412 + * On success, hdr_out is filled, and payload_out + payload_len_out
413 + * point to the payload bytes (valid until next receive).
414 + */
415 +static nipc_error_t transport_receive(nipc_client_ctx_t *ctx,
416 + void *buf, size_t buf_size,
417 + nipc_header_t *hdr_out,
418 + const void **payload_out,
419 + size_t *payload_len_out)
420 +{
421 + if (ctx->shm) {
422 + size_t msg_len;
423 + nipc_shm_error_t serr = nipc_shm_receive(ctx->shm, buf, buf_size,
424 + &msg_len, 30000);
425 + if (serr != NIPC_SHM_OK)
426 + return NIPC_ERR_TRUNCATED;
427 +
428 + if (msg_len < NIPC_HEADER_LEN)
429 + return NIPC_ERR_TRUNCATED;
430 +
431 + nipc_error_t perr = nipc_header_decode(buf, msg_len, hdr_out);
432 + if (perr != NIPC_OK)
433 + return perr;
434 +
435 + *payload_out = (const uint8_t *)buf + NIPC_HEADER_LEN;
436 + *payload_len_out = msg_len - NIPC_HEADER_LEN;
437 + return NIPC_OK;
438 + }
439 +
440 + /* UDS path */
441 + nipc_uds_error_t uerr = nipc_uds_receive(&ctx->session, buf, buf_size,
442 + hdr_out, payload_out,
443 + payload_len_out);
444 + return (uerr == NIPC_UDS_OK) ? NIPC_OK : NIPC_ERR_TRUNCATED;
445 +}
446 +
447 +/* ------------------------------------------------------------------ */
448 +/* Internal: generic raw call (send request, receive response) */
449 +/* ------------------------------------------------------------------ */
450 +
451 +/*
452 + * Single-attempt raw call: build envelope, send, receive, validate
453 + * envelope. The caller handles encode before and decode after.
454 + *
455 + * On success, response_payload_out and response_len_out point into the
456 + * internal client response buffer (valid until next call on this context).
457 + */
458 +static nipc_error_t do_raw_call(nipc_client_ctx_t *ctx,
459 + uint16_t method_code,
460 + const void *request_payload,
461 + size_t request_len,
462 + const void **response_payload_out,
463 + size_t *response_len_out)
464 +{
465 + nipc_header_t hdr = {0};
466 + hdr.kind = NIPC_KIND_REQUEST;
467 + hdr.code = method_code;
468 + hdr.flags = 0;
469 + hdr.item_count = 1;
470 + hdr.message_id = (uint64_t)(ctx->call_count + 1);
471 + hdr.transport_status = NIPC_STATUS_OK;
472 +
473 + nipc_error_t err = transport_send(ctx, &hdr, request_payload, request_len);
474 + if (err != NIPC_OK) {
475 + return err;
476 + }
477 +
478 + nipc_header_t resp_hdr;
479 + err = transport_receive(ctx, ctx->response_buf, ctx->response_buf_size,
480 + &resp_hdr, response_payload_out, response_len_out);
481 + if (err != NIPC_OK) {
482 + return err;
483 + }
484 +
485 + if (resp_hdr.kind != NIPC_KIND_RESPONSE)
486 + return NIPC_ERR_BAD_KIND;
487 + if (resp_hdr.code != method_code)
488 + return NIPC_ERR_BAD_LAYOUT;
489 + if (resp_hdr.message_id != hdr.message_id)
490 + return NIPC_ERR_BAD_LAYOUT;
491 +
492 + switch (resp_hdr.transport_status) {
493 + case NIPC_STATUS_OK:
494 + break;
495 + case NIPC_STATUS_LIMIT_EXCEEDED:
496 + if (ctx->session.max_response_payload_bytes > 0) {
497 + uint32_t current = ctx->session.max_response_payload_bytes;
498 + client_note_response_capacity(
499 + ctx, current >= UINT32_MAX / 2u ? UINT32_MAX : current * 2u);
500 + }
501 + return NIPC_ERR_OVERFLOW;
502 + case NIPC_STATUS_UNSUPPORTED:
503 + return NIPC_ERR_BAD_LAYOUT;
504 + case NIPC_STATUS_BAD_ENVELOPE:
505 + case NIPC_STATUS_INTERNAL_ERROR:
506 + default:
507 + return NIPC_ERR_BAD_LAYOUT;
508 + }
509 +
510 + return NIPC_OK;
511 +}
512 +
513 +/*
514 + * Generic call-with-retry:
515 + * - ordinary failures reconnect and retry once
516 + * - overflow-driven resize recovery may reconnect repeatedly until
517 + * negotiated capacities grow or recovery fails
518 + * The caller provides a function pointer for the single-attempt logic.
519 + */
520 +typedef nipc_error_t (*nipc_attempt_fn)(nipc_client_ctx_t *ctx, void *state);
521 +
522 +static nipc_error_t call_with_retry(nipc_client_ctx_t *ctx,
523 + nipc_attempt_fn attempt,
524 + void *state)
525 +{
526 + if (ctx->state != NIPC_CLIENT_READY) {
527 + ctx->error_count++;
528 + return NIPC_ERR_NOT_READY;
529 + }
530 +
531 + /* Cap overflow-driven retries: payloads grow by powers of 2, so 8
532 + * retries allows ~256x growth from the initial negotiated size. */
533 + int overflow_retries = 0;
534 + for (;;) {
535 + uint32_t prev_req = ctx->session.max_request_payload_bytes;
536 + uint32_t prev_resp = ctx->session.max_response_payload_bytes;
537 + uint32_t prev_cfg_req = ctx->transport_config.max_request_payload_bytes;
538 + uint32_t prev_cfg_resp = ctx->transport_config.max_response_payload_bytes;
539 +
540 + nipc_error_t err = attempt(ctx, state);
541 + if (err == NIPC_OK) {
542 + ctx->call_count++;
543 + return NIPC_OK;
544 + }
545 +
546 + if (err != NIPC_ERR_OVERFLOW) {
547 + client_disconnect(ctx);
548 + ctx->state = NIPC_CLIENT_BROKEN;
549 + ctx->state = client_try_connect(ctx);
550 + if (ctx->state != NIPC_CLIENT_READY) {
551 + ctx->error_count++;
552 + return err;
553 + }
554 +
555 + ctx->reconnect_count++;
556 + err = attempt(ctx, state);
557 + if (err == NIPC_OK) {
558 + ctx->call_count++;
559 + return NIPC_OK;
560 + }
561 +
562 + client_disconnect(ctx);
563 + ctx->state = NIPC_CLIENT_BROKEN;
564 + ctx->error_count++;
565 + return err;
566 + }
567 +
568 + client_disconnect(ctx);
569 + ctx->state = NIPC_CLIENT_BROKEN;
570 + ctx->state = client_try_connect(ctx);
571 + if (ctx->state != NIPC_CLIENT_READY) {
572 + ctx->error_count++;
573 + return err;
574 + }
575 + ctx->reconnect_count++;
576 +
577 + if (ctx->session.max_request_payload_bytes <= prev_req &&
578 + ctx->session.max_response_payload_bytes <= prev_resp &&
579 + ctx->transport_config.max_request_payload_bytes <= prev_cfg_req &&
580 + ctx->transport_config.max_response_payload_bytes <= prev_cfg_resp) {
581 + client_disconnect(ctx);
582 + ctx->state = NIPC_CLIENT_BROKEN;
583 + ctx->error_count++;
584 + return err;
585 + }
586 +
587 + if (++overflow_retries >= 8) {
588 + client_disconnect(ctx);
589 + ctx->state = NIPC_CLIENT_BROKEN;
590 + ctx->error_count++;
591 + return err;
592 + }
593 + }
594 +}
595 +
596 +/* ------------------------------------------------------------------ */
597 +/* Internal: single attempt at a cgroups snapshot call */
598 +/* ------------------------------------------------------------------ */
599 +
600 +typedef struct {
601 + nipc_cgroups_resp_view_t *view_out;
602 +} cgroups_call_state_t;
603 +
604 +static nipc_error_t do_cgroups_attempt(nipc_client_ctx_t *ctx, void *state)
605 +{
606 + cgroups_call_state_t *s = (cgroups_call_state_t *)state;
607 +
608 + nipc_cgroups_req_t req = { .layout_version = 1, .flags = 0 };
609 + uint8_t req_buf[4];
610 + size_t req_len = nipc_cgroups_req_encode(&req, req_buf, sizeof(req_buf));
611 + if (req_len == 0)
612 + return NIPC_ERR_TRUNCATED;
613 +
614 + const void *payload;
615 + size_t payload_len;
616 + nipc_error_t err = do_raw_call(ctx, NIPC_METHOD_CGROUPS_SNAPSHOT,
617 + req_buf, req_len,
618 + &payload, &payload_len);
619 + if (err != NIPC_OK)
620 + return err;
621 +
622 + return nipc_cgroups_resp_decode(payload, payload_len, s->view_out);
623 +}
624 +
625 +/* ------------------------------------------------------------------ */
626 +/* Public API: client lifecycle */
627 +/* ------------------------------------------------------------------ */
628 +
629 +void nipc_client_init(nipc_client_ctx_t *ctx,
630 + const char *run_dir,
631 + const char *service_name,
632 + const nipc_client_config_t *config)
633 +{
634 + memset(ctx, 0, sizeof(*ctx));
635 + ctx->state = NIPC_CLIENT_DISCONNECTED;
636 + ctx->session.fd = -1;
637 + ctx->session_valid = false;
638 + ctx->shm = NULL;
639 +
640 + if (run_dir) {
641 + size_t len = strlen(run_dir);
642 + if (len >= sizeof(ctx->run_dir))
643 + len = sizeof(ctx->run_dir) - 1;
644 + memcpy(ctx->run_dir, run_dir, len);
645 + ctx->run_dir[len] = '\0';
646 + }
647 +
648 + if (service_name) {
649 + size_t len = strlen(service_name);
650 + if (len >= sizeof(ctx->service_name))
651 + len = sizeof(ctx->service_name) - 1;
652 + memcpy(ctx->service_name, service_name, len);
653 + ctx->service_name[len] = '\0';
654 + }
655 +
656 + ctx->transport_config = service_client_config_to_transport(config);
657 + if (ctx->transport_config.max_request_payload_bytes == 0)
658 + ctx->transport_config.max_request_payload_bytes = cgroups_request_payload_default();
659 + if (ctx->transport_config.max_response_payload_bytes == 0)
660 + ctx->transport_config.max_response_payload_bytes = cgroups_response_payload_default();
661 +}
662 +
663 +bool nipc_client_refresh(nipc_client_ctx_t *ctx)
664 +{
665 + nipc_client_state_t old_state = ctx->state;
666 +
667 + switch (ctx->state) {
668 + case NIPC_CLIENT_DISCONNECTED:
669 + case NIPC_CLIENT_NOT_FOUND:
670 + /* Attempt to connect */
671 + ctx->state = NIPC_CLIENT_CONNECTING;
672 + ctx->state = client_try_connect(ctx);
673 + if (ctx->state == NIPC_CLIENT_READY)
674 + ctx->connect_count++;
675 + break;
676 +
677 + case NIPC_CLIENT_BROKEN:
678 + /* Reconnect: tear down old connection first */
679 + client_disconnect(ctx);
680 + ctx->state = NIPC_CLIENT_CONNECTING;
681 + ctx->state = client_try_connect(ctx);
682 + if (ctx->state == NIPC_CLIENT_READY)
683 + ctx->reconnect_count++;
684 + break;
685 +
686 + case NIPC_CLIENT_READY:
687 + case NIPC_CLIENT_CONNECTING:
688 + case NIPC_CLIENT_AUTH_FAILED:
689 + case NIPC_CLIENT_INCOMPATIBLE:
690 + /* No action needed */
691 + break;
692 + }
693 +
694 + return ctx->state != old_state;
695 +}
696 +
697 +void nipc_client_status(const nipc_client_ctx_t *ctx,
698 + nipc_client_status_t *out)
699 +{
700 + out->state = ctx->state;
701 + out->connect_count = ctx->connect_count;
702 + out->reconnect_count = ctx->reconnect_count;
703 + out->call_count = ctx->call_count;
704 + out->error_count = ctx->error_count;
705 +}
706 +
707 +void nipc_client_close(nipc_client_ctx_t *ctx)
708 +{
709 + client_disconnect(ctx);
710 + free(ctx->response_buf);
711 + free(ctx->send_buf);
712 + ctx->response_buf = NULL;
713 + ctx->send_buf = NULL;
714 + ctx->response_buf_size = 0;
715 + ctx->send_buf_size = 0;
716 + ctx->state = NIPC_CLIENT_DISCONNECTED;
717 +}
718 +
719 +/* ------------------------------------------------------------------ */
720 +/* Public API: typed cgroups snapshot call */
721 +/* ------------------------------------------------------------------ */
722 +
723 +nipc_error_t nipc_client_call_cgroups_snapshot(
724 + nipc_client_ctx_t *ctx,
725 + nipc_cgroups_resp_view_t *view_out)
726 +{
727 + cgroups_call_state_t state = {
728 + .view_out = view_out,
729 + };
730 + return call_with_retry(ctx, do_cgroups_attempt, &state);
731 +}
732 +
733 +/* ------------------------------------------------------------------ */
734 +/* Internal: managed server session handler */
735 +/* ------------------------------------------------------------------ */
736 +
737 +/*
738 + * Wait for data on a file descriptor with periodic shutdown checks.
739 + * Returns: 1 = data ready, 0 = server stopping, -1 = error/hangup.
740 + */
741 +static int poll_with_shutdown(int fd, bool *running)
742 +{
743 + while (__atomic_load_n(running, __ATOMIC_RELAXED)) {
744 + struct pollfd pfd = { .fd = fd, .events = POLLIN };
745 + int ret = poll(&pfd, 1, SERVER_POLL_TIMEOUT_MS);
746 +
747 + if (ret < 0) {
748 + if (errno == EINTR)
749 + continue;
750 + return -1;
751 + }
752 +
753 + if (ret == 0)
754 + continue; /* timeout, check running flag */
755 +
756 + if (pfd.revents & (POLLERR | POLLHUP | POLLNVAL))
757 + return -1;
758 +
759 + if (pfd.revents & POLLIN)
760 + return 1;
761 + }
762 + return 0;
763 +}
764 +
765 +static uint32_t server_snapshot_max_items(size_t response_buf_size,
766 + const nipc_cgroups_service_handler_t *service_handler)
767 +{
768 + if (service_handler->snapshot_max_items != 0)
769 + return service_handler->snapshot_max_items;
770 + return nipc_cgroups_builder_estimate_max_items(response_buf_size);
771 +}
772 +
773 +static void server_note_request_capacity(nipc_managed_server_t *server,
774 + uint32_t payload_len)
775 +{
776 + uint32_t grown = next_power_of_2_u32(payload_len);
777 + uint32_t current = __atomic_load_n(&server->learned_request_payload_bytes,
778 + __ATOMIC_RELAXED);
779 + while (grown > current &&
780 + !__atomic_compare_exchange_n(&server->learned_request_payload_bytes,
781 + &current, grown, false,
782 + __ATOMIC_RELEASE, __ATOMIC_RELAXED)) {
783 + }
784 +}
785 +
786 +static void server_note_response_capacity(nipc_managed_server_t *server,
787 + uint32_t payload_len)
788 +{
789 + uint32_t grown = next_power_of_2_u32(payload_len);
790 + uint32_t current = __atomic_load_n(&server->learned_response_payload_bytes,
791 + __ATOMIC_RELAXED);
792 + while (grown > current &&
793 + !__atomic_compare_exchange_n(&server->learned_response_payload_bytes,
794 + &current, grown, false,
795 + __ATOMIC_RELEASE, __ATOMIC_RELAXED)) {
796 + }
797 +}
798 +
799 +static nipc_error_t server_typed_dispatch(void *user,
800 + const nipc_header_t *request_hdr,
801 + const uint8_t *request_payload,
802 + size_t request_len,
803 + uint8_t *response_buf,
804 + size_t response_buf_size,
805 + size_t *response_len_out)
806 +{
807 + nipc_managed_server_t *server = (nipc_managed_server_t *)user;
808 + nipc_cgroups_service_handler_t *service_handler = &server->service_handler;
809 + (void)request_hdr;
810 +
811 + if (!service_handler->handle)
812 + return NIPC_ERR_HANDLER_FAILED;
813 +
814 + return nipc_dispatch_cgroups_snapshot(
815 + request_payload, request_len,
816 + response_buf, response_buf_size, response_len_out,
817 + server_snapshot_max_items(response_buf_size, service_handler),
818 + service_handler->handle, service_handler->user);
819 +}
820 +
821 +/*
822 + * Handle one client session: read requests, dispatch to handler,
823 + * send responses. Each session gets its own response buffer.
824 + * Runs until the client disconnects or server stops.
825 + */
826 +static void server_handle_session(nipc_managed_server_t *server,
827 + nipc_uds_session_t *session,
828 + nipc_shm_ctx_t *shm,
829 + uint8_t *resp_buf,
830 + size_t resp_buf_size)
831 +{
832 + /* Allocate recv buffer based on negotiated max request size */
833 + size_t recv_size = NIPC_HEADER_LEN + session->max_request_payload_bytes;
834 + if (recv_size < NIPC_HEADER_LEN + 1024u)
835 + recv_size = NIPC_HEADER_LEN + 1024u;
836 + uint8_t *recv_buf = service_malloc(
837 + recv_size, NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_RECV_BUF_MALLOC_INTERNAL);
838 + if (!recv_buf)
839 + return;
840 +
841 + while (__atomic_load_n(&server->running, __ATOMIC_RELAXED)) {
842 + nipc_header_t hdr;
843 + const void *payload;
844 + size_t payload_len;
845 +
846 + /* Receive request via the active transport */
847 + if (shm) {
848 + size_t msg_len;
849 + nipc_shm_error_t serr = nipc_shm_receive(shm, recv_buf, recv_size,
850 + &msg_len, SERVER_POLL_TIMEOUT_MS);
851 + if (serr == NIPC_SHM_ERR_TIMEOUT)
852 + continue; /* check running flag */
853 + if (serr != NIPC_SHM_OK)
854 + break;
855 + if (msg_len < NIPC_HEADER_LEN)
856 + break;
857 +
858 + nipc_error_t perr = nipc_header_decode(recv_buf, msg_len, &hdr);
859 + if (perr != NIPC_OK)
860 + break;
861 +
862 + payload = recv_buf + NIPC_HEADER_LEN;
863 + payload_len = msg_len - NIPC_HEADER_LEN;
864 + } else {
865 + /* Poll the session fd before blocking on receive */
866 + int pr = poll_with_shutdown(session->fd, &server->running);
867 + if (pr <= 0)
868 + break; /* shutdown or error */
869 +
870 + nipc_uds_error_t uerr = nipc_uds_receive(
871 + session, recv_buf, recv_size,
872 + &hdr, &payload, &payload_len);
873 + if (uerr == NIPC_UDS_ERR_LIMIT_EXCEEDED) {
874 + if (hdr.kind == NIPC_KIND_REQUEST) {
875 + if (hdr.payload_len > 0)
876 + server_note_request_capacity(server, hdr.payload_len);
877 +
878 + nipc_header_t resp_hdr = {0};
879 + resp_hdr.kind = NIPC_KIND_RESPONSE;
880 + resp_hdr.code = hdr.code;
881 + resp_hdr.message_id = hdr.message_id;
882 + resp_hdr.transport_status = NIPC_STATUS_LIMIT_EXCEEDED;
883 + resp_hdr.item_count = 1;
884 + resp_hdr.flags = 0;
885 +
886 + if (nipc_uds_send(session, &resp_hdr, NULL, 0) != NIPC_UDS_OK)
887 + break;
888 + }
889 + break;
890 + }
891 + if (uerr != NIPC_UDS_OK)
892 + break;
893 + }
894 +
895 + /* Protocol violation: unexpected message kind terminates session */
896 + if (hdr.kind != NIPC_KIND_REQUEST)
897 + break;
898 +
899 + if (hdr.code != server->expected_method_code) {
900 + nipc_header_t resp_hdr = {0};
901 + resp_hdr.kind = NIPC_KIND_RESPONSE;
902 + resp_hdr.code = hdr.code;
903 + resp_hdr.message_id = hdr.message_id;
904 + resp_hdr.transport_status = NIPC_STATUS_UNSUPPORTED;
905 + resp_hdr.item_count = 1;
906 + resp_hdr.flags = 0;
907 +
908 + if (shm) {
909 + uint8_t msg[NIPC_HEADER_LEN];
910 + resp_hdr.magic = NIPC_MAGIC_MSG;
911 + resp_hdr.version = NIPC_VERSION;
912 + resp_hdr.header_len = NIPC_HEADER_LEN;
913 + resp_hdr.payload_len = 0;
914 + nipc_header_encode(&resp_hdr, msg, sizeof(msg));
915 + if (nipc_shm_send(shm, msg, sizeof(msg)) != NIPC_SHM_OK)
916 + break;
917 + } else {
918 + if (nipc_uds_send(session, &resp_hdr, NULL, 0) != NIPC_UDS_OK)
919 + break;
920 + }
921 + continue;
922 + }
923 +
924 + if (payload_len <= UINT32_MAX)
925 + server_note_request_capacity(server, (uint32_t)payload_len);
926 +
927 + /* Dispatch: one request kind per service endpoint. */
928 + size_t response_len = 0;
929 + nipc_error_t dispatch_err = server->handler(
930 + server->handler_user,
931 + &hdr,
932 + (const uint8_t *)payload, payload_len,
933 + resp_buf, resp_buf_size,
934 + &response_len);
935 +
936 + /* Build response header */
937 + nipc_header_t resp_hdr = {0};
938 + resp_hdr.kind = NIPC_KIND_RESPONSE;
939 + resp_hdr.code = hdr.code;
940 + resp_hdr.message_id = hdr.message_id;
941 + if ((hdr.flags & NIPC_FLAG_BATCH) && hdr.item_count >= 1) {
942 + resp_hdr.item_count = hdr.item_count;
943 + resp_hdr.flags = NIPC_FLAG_BATCH;
944 + } else {
945 + resp_hdr.item_count = 1;
946 + resp_hdr.flags = 0;
947 + }
948 +
949 + switch (dispatch_err) {
950 + case NIPC_OK:
951 + if (response_len > session->max_response_payload_bytes) {
952 + server_note_response_capacity(
953 + server,
954 + response_len >= UINT32_MAX ? UINT32_MAX : (uint32_t)response_len);
955 + resp_hdr.transport_status = NIPC_STATUS_LIMIT_EXCEEDED;
956 + response_len = 0;
957 + } else {
958 + if (response_len <= UINT32_MAX)
959 + server_note_response_capacity(server, (uint32_t)response_len);
960 + resp_hdr.transport_status = NIPC_STATUS_OK;
961 + }
962 + break;
963 + case NIPC_ERR_OVERFLOW:
964 + if (session->max_response_payload_bytes >= UINT32_MAX / 2u)
965 + server_note_response_capacity(server, UINT32_MAX);
966 + else
967 + server_note_response_capacity(server,
968 + session->max_response_payload_bytes * 2u);
969 + resp_hdr.transport_status = NIPC_STATUS_LIMIT_EXCEEDED;
970 + response_len = 0;
971 + break;
972 + case NIPC_ERR_TRUNCATED:
973 + case NIPC_ERR_BAD_LAYOUT:
974 + case NIPC_ERR_OUT_OF_BOUNDS:
975 + case NIPC_ERR_MISSING_NUL:
976 + case NIPC_ERR_BAD_ALIGNMENT:
977 + case NIPC_ERR_BAD_ITEM_COUNT:
978 + resp_hdr.transport_status = NIPC_STATUS_BAD_ENVELOPE;
979 + response_len = 0;
980 + break;
981 + case NIPC_ERR_HANDLER_FAILED:
982 + default:
983 + resp_hdr.transport_status = NIPC_STATUS_INTERNAL_ERROR;
984 + response_len = 0;
985 + break;
986 + }
987 +
988 + /* Send response via the active transport */
989 + if (shm) {
990 + size_t msg_len = NIPC_HEADER_LEN + response_len;
991 +
992 + resp_hdr.magic = NIPC_MAGIC_MSG;
993 + resp_hdr.version = NIPC_VERSION;
994 + resp_hdr.header_len = NIPC_HEADER_LEN;
995 + resp_hdr.payload_len = (uint32_t)response_len;
996 +
997 + /* Use a stack buffer for small responses, heap for large ones */
998 + uint8_t stack_msg[4096];
999 + uint8_t *msg = (msg_len <= sizeof(stack_msg)) ? stack_msg :
1000 + service_malloc(msg_len, NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_RESP_BUF_MALLOC_INTERNAL);
1001 + if (!msg)
1002 + break;
1003 +
1004 + nipc_header_encode(&resp_hdr, msg, NIPC_HEADER_LEN);
1005 + if (response_len > 0)
1006 + memcpy(msg + NIPC_HEADER_LEN, resp_buf, response_len);
1007 +
1008 + nipc_shm_error_t serr = nipc_shm_send(shm, msg, msg_len);
1009 + if (msg != stack_msg)
1010 + free(msg);
1011 + if (serr != NIPC_SHM_OK)
1012 + break;
1013 + } else {
1014 + nipc_uds_error_t uerr = nipc_uds_send(
1015 + session, &resp_hdr, resp_buf, response_len);
1016 + if (uerr != NIPC_UDS_OK)
1017 + break;
1018 + }
1019 +
1020 + if (dispatch_err == NIPC_ERR_OVERFLOW)
1021 + break;
1022 + }
1023 +
1024 + free(recv_buf);
1025 +}
1026 +
1027 +/* ------------------------------------------------------------------ */
1028 +/* Internal: per-session handler thread */
1029 +/* ------------------------------------------------------------------ */
1030 +
1031 +/* Thread function: handles one client session from accept to disconnect. */
1032 +static void *session_handler_thread(void *arg)
1033 +{
1034 + nipc_session_ctx_t *sctx = (nipc_session_ctx_t *)arg;
1035 + nipc_managed_server_t *server = sctx->server;
1036 +
1037 + /* Allocate a per-session response buffer */
1038 + size_t resp_size = (size_t)sctx->session.max_response_payload_bytes;
1039 + if (resp_size < 1024u)
1040 + resp_size = 1024u;
1041 + uint8_t *resp_buf = service_malloc(
1042 + resp_size, NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_RESP_BUF_MALLOC_INTERNAL);
1043 + if (resp_buf) {
1044 + server_handle_session(server, &sctx->session, sctx->shm,
1045 + resp_buf, resp_size);
1046 + free(resp_buf);
1047 + }
1048 +
1049 + /* Cleanup SHM and session */
1050 + if (sctx->shm) {
1051 + nipc_shm_destroy(sctx->shm);
1052 + free(sctx->shm);
1053 + }
1054 + nipc_uds_close_session(&sctx->session);
1055 +
1056 + /* Mark inactive so the acceptor's reap loop (or server destroy)
1057 + * can join this thread and free sctx. Do NOT remove from the
1058 + * tracking array here — the reap/destroy path owns that. */
1059 + __atomic_store_n(&sctx->active, false, __ATOMIC_RELEASE);
1060 + return NULL;
1061 +}
1062 +
1063 +/* ------------------------------------------------------------------ */
1064 +/* Internal: reap finished session threads */
1065 +/* ------------------------------------------------------------------ */
1066 +
1067 +/* Reap all finished (inactive) session threads. Called with lock held. */
1068 +static void server_reap_sessions_locked(nipc_managed_server_t *server)
1069 +{
1070 + int i = 0;
1071 + while (i < server->session_count) {
1072 + nipc_session_ctx_t *s = server->sessions[i];
1073 + if (!__atomic_load_n(&s->active, __ATOMIC_ACQUIRE)) {
1074 + pthread_join(s->thread, NULL);
1075 + /* Swap with last, free */
1076 + server->sessions[i] = server->sessions[server->session_count - 1];
1077 + server->session_count--;
1078 + free(s);
1079 + } else {
1080 + i++;
1081 + }
1082 + }
1083 +
1084 +}
1085 +
1086 +static void server_destroy_precreated_shm(nipc_shm_ctx_t **shm)
1087 +{
1088 + if (!shm || !*shm)
1089 + return;
1090 + nipc_shm_destroy(*shm);
1091 + free(*shm);
1092 + *shm = NULL;
1093 +}
1094 +
1095 +static bool server_prepare_accept_config(nipc_managed_server_t *server,
1096 + uint64_t sid,
1097 + nipc_uds_server_config_t *cfg_out,
1098 + nipc_shm_ctx_t **shm_out)
1099 +{
1100 + *cfg_out = server->base_config;
1101 + cfg_out->max_request_payload_bytes =
1102 + __atomic_load_n(&server->learned_request_payload_bytes, __ATOMIC_ACQUIRE);
1103 + cfg_out->max_response_payload_bytes =
1104 + __atomic_load_n(&server->learned_response_payload_bytes, __ATOMIC_ACQUIRE);
1105 + *shm_out = NULL;
1106 +
1107 + uint32_t shm_profiles = cfg_out->supported_profiles &
1108 + (NIPC_PROFILE_SHM_HYBRID | NIPC_PROFILE_SHM_FUTEX);
1109 + if (shm_profiles == 0)
1110 + return true;
1111 +
1112 + nipc_shm_ctx_t *shm = service_calloc(
1113 + 1, sizeof(nipc_shm_ctx_t),
1114 + NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_SHM_CTX_CALLOC_INTERNAL);
1115 + if (!shm)
1116 + return false;
1117 +
1118 + /* HELLO has not been read yet, so the request segment must cover any
1119 + * client proposal the handshake may legally echo back. */
1120 + nipc_shm_error_t serr = nipc_shm_server_create(
1121 + server->run_dir, server->service_name,
1122 + sid,
1123 + NIPC_MAX_PAYLOAD_CAP + NIPC_HEADER_LEN,
1124 + cfg_out->max_response_payload_bytes + NIPC_HEADER_LEN,
1125 + shm);
1126 + if (serr == NIPC_SHM_OK) {
1127 + *shm_out = shm;
1128 + return true;
1129 + }
1130 +
1131 + free(shm);
1132 + cfg_out->supported_profiles &= ~(NIPC_PROFILE_SHM_HYBRID | NIPC_PROFILE_SHM_FUTEX);
1133 + cfg_out->preferred_profiles &= ~(NIPC_PROFILE_SHM_HYBRID | NIPC_PROFILE_SHM_FUTEX);
1134 + return cfg_out->supported_profiles != 0;
1135 +}
1136 +
1137 +/* ------------------------------------------------------------------ */
1138 +/* Public API: managed server */
1139 +/* ------------------------------------------------------------------ */
1140 +
1141 +static nipc_error_t server_init_raw(nipc_managed_server_t *server,
1142 + const char *run_dir,
1143 + const char *service_name,
1144 + const nipc_uds_server_config_t *config,
1145 + int worker_count,
1146 + uint16_t expected_method_code,
1147 + nipc_server_handler_fn handler,
1148 + void *user)
1149 +{
1150 + memset(server, 0, sizeof(*server));
1151 + server->listener.fd = -1;
1152 + __atomic_store_n(&server->running, false, __ATOMIC_RELAXED);
1153 + server->acceptor_started = false;
1154 +
1155 + if (!run_dir || !service_name || !handler)
1156 + return NIPC_ERR_BAD_LAYOUT;
1157 +
1158 + if (worker_count < 1)
1159 + worker_count = 1;
1160 +
1161 + /* Store config */
1162 + {
1163 + size_t len = strlen(run_dir);
1164 + if (len >= sizeof(server->run_dir))
1165 + len = sizeof(server->run_dir) - 1;
1166 + memcpy(server->run_dir, run_dir, len);
1167 + server->run_dir[len] = '\0';
1168 + }
1169 + {
1170 + size_t len = strlen(service_name);
1171 + if (len >= sizeof(server->service_name))
1172 + len = sizeof(server->service_name) - 1;
1173 + memcpy(server->service_name, service_name, len);
1174 + server->service_name[len] = '\0';
1175 + }
1176 +
1177 + server->handler = handler;
1178 + server->handler_user = user;
1179 + server->worker_count = worker_count;
1180 + server->expected_method_code = expected_method_code;
1181 + server->base_config = *config;
1182 + server->learned_request_payload_bytes =
1183 + (config && config->max_request_payload_bytes > 0)
1184 + ? config->max_request_payload_bytes
1185 + : NIPC_MAX_PAYLOAD_DEFAULT;
1186 + server->learned_response_payload_bytes =
1187 + (config && config->max_response_payload_bytes > 0)
1188 + ? config->max_response_payload_bytes
1189 + : NIPC_MAX_PAYLOAD_DEFAULT;
1190 +
1191 + /* Initialize session tracking */
1192 + server->session_capacity = worker_count * 2; /* room for slots being reaped */
1193 + if (server->session_capacity < 16)
1194 + server->session_capacity = 16;
1195 + server->sessions = service_calloc((size_t)server->session_capacity,
1196 + sizeof(nipc_session_ctx_t *),
1197 + NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_SESSIONS_CALLOC_INTERNAL);
1198 + if (!server->sessions)
1199 + return NIPC_ERR_OVERFLOW;
1200 + server->session_count = 0;
1201 + server->next_session_id = 1; /* spec: monotonic counter starting at 1 */
1202 + pthread_mutex_init(&server->sessions_lock, NULL);
1203 +
1204 + /* Clean up stale SHM regions from previous crashes (spec requirement:
1205 + * runs once at server startup, before the listener begins accepting). */
1206 + nipc_shm_cleanup_stale(run_dir, service_name);
1207 +
1208 + /* Start listening via L1 */
1209 + nipc_uds_error_t uerr = nipc_uds_listen(
1210 + run_dir, service_name, config, &server->listener);
1211 + if (uerr != NIPC_UDS_OK) {
1212 + free(server->sessions);
1213 + server->sessions = NULL;
1214 + pthread_mutex_destroy(&server->sessions_lock);
1215 + return NIPC_ERR_BAD_LAYOUT;
1216 + }
1217 +
1218 + return NIPC_OK;
1219 +}
1220 +
1221 +nipc_error_t nipc_server_init_typed(nipc_managed_server_t *server,
1222 + const char *run_dir,
1223 + const char *service_name,
1224 + const nipc_server_config_t *config,
1225 + int worker_count,
1226 + const nipc_cgroups_service_handler_t *service_handler)
1227 +{
1228 + if (!service_handler)
1229 + return NIPC_ERR_BAD_LAYOUT;
1230 +
1231 + nipc_uds_server_config_t typed_cfg = service_server_config_to_transport(config);
1232 + if (typed_cfg.max_request_payload_bytes == 0)
1233 + typed_cfg.max_request_payload_bytes = cgroups_request_payload_default();
1234 + if (typed_cfg.max_response_payload_bytes == 0)
1235 + typed_cfg.max_response_payload_bytes = cgroups_response_payload_default();
1236 +
1237 + nipc_error_t err = server_init_raw(server, run_dir, service_name,
1238 + &typed_cfg, worker_count,
1239 + NIPC_METHOD_CGROUPS_SNAPSHOT,
1240 + server_typed_dispatch, server);
1241 + if (err != NIPC_OK)
1242 + return err;
1243 +
1244 + server->service_handler = *service_handler;
1245 + return NIPC_OK;
1246 +}
1247 +
1248 +nipc_error_t nipc_server_init_raw_for_tests(nipc_managed_server_t *server,
1249 + const char *run_dir,
1250 + const char *service_name,
1251 + const nipc_uds_server_config_t *config,
1252 + int worker_count,
1253 + uint16_t expected_method_code,
1254 + nipc_server_handler_fn handler,
1255 + void *user)
1256 +{
1257 + return server_init_raw(server, run_dir, service_name, config,
1258 + worker_count, expected_method_code, handler, user);
1259 +}
1260 +
1261 +void nipc_server_run(nipc_managed_server_t *server)
1262 +{
1263 + __atomic_store_n(&server->running, true, __ATOMIC_RELEASE);
1264 +
1265 + while (__atomic_load_n(&server->running, __ATOMIC_RELAXED)) {
1266 + /* Poll the listener fd before blocking on accept */
1267 + int pr = poll_with_shutdown(server->listener.fd, &server->running);
1268 + if (pr <= 0)
1269 + break; /* shutdown or error */
1270 +
1271 + /* Accept one client via L1 */
1272 + nipc_uds_session_t session;
1273 + memset(&session, 0, sizeof(session));
1274 + session.fd = -1;
1275 +
1276 + uint64_t sid = server->next_session_id++;
1277 + nipc_uds_server_config_t accept_cfg;
1278 + nipc_shm_ctx_t *prepared_shm = NULL;
1279 + if (!server_prepare_accept_config(server, sid, &accept_cfg, &prepared_shm)) {
1280 + usleep(10000);
1281 + continue;
1282 + }
1283 +
1284 + server->listener.config = accept_cfg;
1285 + nipc_uds_error_t uerr = nipc_uds_accept(
1286 + &server->listener, sid, &session);
1287 + if (uerr != NIPC_UDS_OK) {
1288 + server_destroy_precreated_shm(&prepared_shm);
1289 + if (!__atomic_load_n(&server->running, __ATOMIC_RELAXED))
1290 + break;
1291 + usleep(10000);
1292 + continue;
1293 + }
1294 +
1295 + server_note_request_capacity(server, session.max_request_payload_bytes);
1296 + server_note_response_capacity(server, session.max_response_payload_bytes);
1297 +
1298 + /* Enforce worker_count limit: reap finished sessions, check count */
1299 + pthread_mutex_lock(&server->sessions_lock);
1300 + server_reap_sessions_locked(server);
1301 +
1302 + if (server->session_count >= server->worker_count) {
1303 + /* At capacity: reject this client by closing the session */
1304 + pthread_mutex_unlock(&server->sessions_lock);
1305 + server_destroy_precreated_shm(&prepared_shm);
1306 + nipc_uds_close_session(&session);
1307 + continue;
1308 + }
1309 +
1310 + /* SHM profile guarantee: only negotiate SHM for sessions that already
1311 + * have a prepared per-session SHM region. */
1312 + nipc_shm_ctx_t *shm = prepared_shm;
1313 + if (session.selected_profile == NIPC_PROFILE_SHM_HYBRID ||
1314 + session.selected_profile == NIPC_PROFILE_SHM_FUTEX) {
1315 + if (!shm) {
1316 + pthread_mutex_unlock(&server->sessions_lock);
1317 + nipc_uds_close_session(&session);
1318 + continue;
1319 + }
1320 + } else {
1321 + server_destroy_precreated_shm(&prepared_shm);
1322 + shm = NULL;
1323 + }
1324 +
1325 + /* Create session context */
1326 + nipc_session_ctx_t *sctx = service_calloc(
1327 + 1, sizeof(nipc_session_ctx_t),
1328 + NIPC_POSIX_SERVICE_TEST_FAULT_SERVER_SESSION_CTX_CALLOC_INTERNAL);
1329 + if (!sctx) {
1330 + if (shm) { nipc_shm_destroy(shm); free(shm); }
1331 + pthread_mutex_unlock(&server->sessions_lock);
1332 + nipc_uds_close_session(&session);
1333 + continue;
1334 + }
1335 +
1336 + sctx->server = server;
1337 + sctx->session = session;
1338 + sctx->shm = shm;
1339 + sctx->id = sid;
1340 + __atomic_store_n(&sctx->active, true, __ATOMIC_RELAXED);
1341 +
1342 + server->sessions[server->session_count++] = sctx;
1343 + pthread_mutex_unlock(&server->sessions_lock);
1344 +
1345 + /* Spawn handler thread for this session */
1346 + int rc = service_pthread_create(&sctx->thread, NULL,
1347 + session_handler_thread, sctx);
1348 + if (rc != 0) {
1349 + /* Thread creation failed: clean up */
1350 + pthread_mutex_lock(&server->sessions_lock);
1351 + /* Remove the sctx we just added */
1352 + for (int i = 0; i < server->session_count; i++) {
1353 + if (server->sessions[i] == sctx) {
1354 + server->sessions[i] = server->sessions[server->session_count - 1];
1355 + server->session_count--;
1356 + break;
1357 + }
1358 + }
1359 + pthread_mutex_unlock(&server->sessions_lock);
1360 +
1361 + if (shm) { nipc_shm_destroy(shm); free(shm); }
1362 + nipc_uds_close_session(&session);
1363 + free(sctx);
1364 + }
1365 + }
1366 +}
1367 +
1368 +void nipc_server_stop(nipc_managed_server_t *server)
1369 +{
1370 + __atomic_store_n(&server->running, false, __ATOMIC_RELEASE);
1371 +}
1372 +
1373 +bool nipc_server_drain(nipc_managed_server_t *server, uint32_t timeout_ms)
1374 +{
1375 + /* 1. Stop accepting new clients.
1376 + * Do NOT close the listener here — the run loop may still be
1377 + * polling on listener.fd. Setting the flag is enough; the run
1378 + * loop will exit on its next poll timeout (100ms). The listener
1379 + * is closed later by nipc_server_destroy(). */
1380 + __atomic_store_n(&server->running, false, __ATOMIC_RELEASE);
1381 +
1382 + /* 2. Wait for in-flight sessions to complete */
1383 + bool all_drained = true;
1384 + if (server->sessions) {
1385 + struct timespec deadline;
1386 + clock_gettime(CLOCK_MONOTONIC, &deadline);
1387 + deadline.tv_sec += timeout_ms / 1000;
1388 + deadline.tv_nsec += (timeout_ms % 1000) * 1000000L;
1389 + if (deadline.tv_nsec >= 1000000000L) {
1390 + deadline.tv_sec++;
1391 + deadline.tv_nsec -= 1000000000L;
1392 + }
1393 +
1394 + /* Poll until all sessions are inactive or timeout */
1395 + while (1) {
1396 + pthread_mutex_lock(&server->sessions_lock);
1397 + int active_count = 0;
1398 + for (int i = 0; i < server->session_count; i++) {
1399 + if (__atomic_load_n(&server->sessions[i]->active,
1400 + __ATOMIC_ACQUIRE))
1401 + active_count++;
1402 + }
1403 + pthread_mutex_unlock(&server->sessions_lock);
1404 +
1405 + if (active_count == 0)
1406 + break;
1407 +
1408 + struct timespec now;
1409 + clock_gettime(CLOCK_MONOTONIC, &now);
1410 + if (now.tv_sec > deadline.tv_sec ||
1411 + (now.tv_sec == deadline.tv_sec &&
1412 + now.tv_nsec >= deadline.tv_nsec)) {
1413 + /* Timeout: force-close session fds to unblock recv.
1414 + * Closing the fd causes poll/recv to return error,
1415 + * which terminates the session handler loop. */
1416 + pthread_mutex_lock(&server->sessions_lock);
1417 + for (int i = 0; i < server->session_count; i++) {
1418 + nipc_session_ctx_t *s = server->sessions[i];
1419 + if (__atomic_load_n(&s->active, __ATOMIC_ACQUIRE)) {
1420 + if (s->session.fd >= 0) {
1421 + shutdown(s->session.fd, SHUT_RDWR);
1422 + }
1423 + }
1424 + }
1425 + pthread_mutex_unlock(&server->sessions_lock);
1426 + all_drained = false;
1427 + break;
1428 + }
1429 +
1430 + usleep(5000); /* 5ms poll interval */
1431 + }
1432 +
1433 + /* 3. Join all session threads (finished or not) */
1434 + pthread_mutex_lock(&server->sessions_lock);
1435 + for (int i = 0; i < server->session_count; i++) {
1436 + nipc_session_ctx_t *s = server->sessions[i];
1437 + pthread_mutex_unlock(&server->sessions_lock);
1438 + pthread_join(s->thread, NULL);
1439 + free(s);
1440 + pthread_mutex_lock(&server->sessions_lock);
1441 + }
1442 + server->session_count = 0;
1443 + pthread_mutex_unlock(&server->sessions_lock);
1444 +
1445 + free(server->sessions);
1446 + server->sessions = NULL;
1447 + server->session_capacity = 0;
1448 + pthread_mutex_destroy(&server->sessions_lock);
1449 + }
1450 +
1451 + server->worker_count = 0;
1452 + return all_drained;
1453 +}
1454 +
1455 +void nipc_server_destroy(nipc_managed_server_t *server)
1456 +{
1457 + __atomic_store_n(&server->running, false, __ATOMIC_RELEASE);
1458 + nipc_uds_close_listener(&server->listener);
1459 +
1460 + /* Join all active session threads */
1461 + if (server->sessions) {
1462 + pthread_mutex_lock(&server->sessions_lock);
1463 + for (int i = 0; i < server->session_count; i++) {
1464 + nipc_session_ctx_t *s = server->sessions[i];
1465 + pthread_mutex_unlock(&server->sessions_lock);
1466 + pthread_join(s->thread, NULL);
1467 + free(s);
1468 + pthread_mutex_lock(&server->sessions_lock);
1469 + }
1470 + server->session_count = 0;
1471 + pthread_mutex_unlock(&server->sessions_lock);
1472 +
1473 + free(server->sessions);
1474 + server->sessions = NULL;
1475 + server->session_capacity = 0;
1476 + pthread_mutex_destroy(&server->sessions_lock);
1477 + }
1478 +
1479 + server->worker_count = 0;
1480 +}
1481 +
1482 +/* ------------------------------------------------------------------ */
1483 +/* L3: Client-side cgroups snapshot cache */
1484 +/* ------------------------------------------------------------------ */
1485 +
1486 +/* Free all owned strings in cache items and the items array itself. */
1487 +static void cache_free_items(nipc_cgroups_cache_item_t *items, uint32_t count)
1488 +{
1489 + if (!items)
1490 + return;
1491 +
1492 + for (uint32_t i = 0; i < count; i++) {
1493 + free(items[i].name);
1494 + free(items[i].path);
1495 + }
1496 + free(items);
1497 +}
1498 +
1499 +/* Hash a name string (djb2). Combined with item hash for bucket index. */
1500 +static uint32_t cache_hash_name(const char *name)
1501 +{
1502 + uint32_t h = 5381;
1503 + for (const unsigned char *p = (const unsigned char *)name; *p; p++)
1504 + h = ((h << 5) + h) + *p;
1505 + return h;
1506 +}
1507 +
1508 +/*
1509 + * Build the open-addressing hash table from the items array.
1510 + * Uses (item.hash ^ name_hash) as the probe key.
1511 + * Load factor <= 0.5 (bucket_count >= 2 * item_count).
1512 + */
1513 +static bool cache_build_hashtable(nipc_cgroups_cache_t *cache)
1514 +{
1515 + free(cache->buckets);
1516 + cache->buckets = NULL;
1517 + cache->bucket_count = 0;
1518 +
1519 + if (cache->item_count == 0)
1520 + return true;
1521 +
1522 + uint32_t bcount = next_power_of_2_u32(cache->item_count * 2);
1523 + nipc_cgroups_hash_bucket_t *buckets = service_calloc(bcount,
1524 + sizeof(nipc_cgroups_hash_bucket_t),
1525 + NIPC_POSIX_SERVICE_TEST_FAULT_CACHE_BUCKETS_CALLOC_INTERNAL);
1526 + if (!buckets)
1527 + return false;
1528 +
1529 + uint32_t mask = bcount - 1;
1530 + for (uint32_t i = 0; i < cache->item_count; i++) {
1531 + uint32_t key = cache->items[i].hash ^ cache_hash_name(cache->items[i].name);
1532 + uint32_t slot = key & mask;
1533 +
1534 + /* Linear probe for an empty bucket */
1535 + while (buckets[slot].used)
1536 + slot = (slot + 1) & mask;
1537 +
1538 + buckets[slot].index = i;
1539 + buckets[slot].used = true;
1540 + }
1541 +
1542 + cache->buckets = buckets;
1543 + cache->bucket_count = bcount;
1544 + return true;
1545 +}
1546 +
1547 +/*
1548 + * Build a new cache from a decoded snapshot view. Copies all strings
1549 + * from the ephemeral view into owned heap allocations.
1550 + *
1551 + * Returns the new items array and sets *count_out. Returns NULL on
1552 + * allocation failure.
1553 + */
1554 +static nipc_cgroups_cache_item_t *cache_build_items(
1555 + const nipc_cgroups_resp_view_t *view,
1556 + uint32_t *count_out)
1557 +{
1558 + uint32_t n = view->item_count;
1559 + *count_out = 0;
1560 +
1561 + if (n == 0)
1562 + return NULL; /* empty snapshot is valid */
1563 +
1564 + nipc_cgroups_cache_item_t *items = service_calloc(
1565 + n, sizeof(nipc_cgroups_cache_item_t),
1566 + NIPC_POSIX_SERVICE_TEST_FAULT_CACHE_ITEMS_CALLOC_INTERNAL);
1567 + if (!items)
1568 + return NULL;
1569 +
1570 + for (uint32_t i = 0; i < n; i++) {
1571 + nipc_cgroups_item_view_t iv;
1572 + nipc_error_t err = nipc_cgroups_resp_item(view, i, &iv);
1573 + if (err != NIPC_OK) {
1574 + /* Malformed item: abort build, free partial */
1575 + cache_free_items(items, i);
1576 + return NULL;
1577 + }
1578 +
1579 + items[i].hash = iv.hash;
1580 + items[i].options = iv.options;
1581 + items[i].enabled = iv.enabled;
1582 +
1583 + /* Copy name (add NUL terminator) */
1584 + items[i].name = service_malloc(
1585 + iv.name.len + 1, NIPC_POSIX_SERVICE_TEST_FAULT_CACHE_ITEM_NAME_MALLOC_INTERNAL);
1586 + if (!items[i].name) {
1587 + cache_free_items(items, i);
1588 + return NULL;
1589 + }
1590 + if (iv.name.len > 0)
1591 + memcpy(items[i].name, iv.name.ptr, iv.name.len);
1592 + items[i].name[iv.name.len] = '\0';
1593 +
1594 + /* Copy path (add NUL terminator) */
1595 + items[i].path = service_malloc(
1596 + iv.path.len + 1, NIPC_POSIX_SERVICE_TEST_FAULT_CACHE_ITEM_PATH_MALLOC_INTERNAL);
1597 + if (!items[i].path) {
1598 + free(items[i].name);
1599 + cache_free_items(items, i);
1600 + return NULL;
1601 + }
1602 + if (iv.path.len > 0)
1603 + memcpy(items[i].path, iv.path.ptr, iv.path.len);
1604 + items[i].path[iv.path.len] = '\0';
1605 + }
1606 +
1607 + *count_out = n;
1608 + return items;
1609 +}
1610 +
1611 +void nipc_cgroups_cache_init(nipc_cgroups_cache_t *cache,
1612 + const char *run_dir,
1613 + const char *service_name,
1614 + const nipc_client_config_t *config)
1615 +{
1616 + memset(cache, 0, sizeof(*cache));
1617 +
1618 + nipc_client_init(&cache->client, run_dir, service_name, config);
1619 +
1620 + cache->items = NULL;
1621 + cache->item_count = 0;
1622 + cache->systemd_enabled = 0;
1623 + cache->generation = 0;
1624 + cache->populated = false;
1625 + cache->buckets = NULL;
1626 + cache->bucket_count = 0;
1627 + cache->refresh_success_count = 0;
1628 + cache->refresh_failure_count = 0;
1629 +
1630 + cache->response_buf = NULL;
1631 + cache->response_buf_size = 0;
1632 +}
1633 +
1634 +bool nipc_cgroups_cache_refresh(nipc_cgroups_cache_t *cache)
1635 +{
1636 + /* Drive L2 connection lifecycle */
1637 + nipc_client_refresh(&cache->client);
1638 +
1639 + /* Attempt snapshot call */
1640 + nipc_cgroups_resp_view_t view;
1641 + nipc_error_t err = nipc_client_call_cgroups_snapshot(&cache->client, &view);
1642 +
1643 + if (err != NIPC_OK) {
1644 + /* Refresh failed -- preserve previous cache */
1645 + cache->refresh_failure_count++;
1646 + return false;
1647 + }
1648 +
1649 + /* Build new cache from the snapshot view */
1650 + uint32_t new_count = 0;
1651 + nipc_cgroups_cache_item_t *new_items = NULL;
1652 +
1653 + if (view.item_count > 0) {
1654 + new_items = cache_build_items(&view, &new_count);
1655 + if (!new_items && view.item_count > 0) {
1656 + /* Build failed (allocation error) -- preserve old cache */
1657 + cache->refresh_failure_count++;
1658 + return false;
1659 + }
1660 + }
1661 +
1662 + /* Replace old cache with new one */
1663 + cache_free_items(cache->items, cache->item_count);
1664 + cache->items = new_items;
1665 + cache->item_count = new_count;
1666 + cache->systemd_enabled = view.systemd_enabled;
1667 + cache->generation = view.generation;
1668 + cache->populated = true;
1669 + cache->refresh_success_count++;
1670 +
1671 + /* Record monotonic timestamp */
1672 + struct timespec ts;
1673 + clock_gettime(CLOCK_MONOTONIC, &ts);
1674 + cache->last_refresh_ts = (uint64_t)ts.tv_sec * 1000 + (uint64_t)ts.tv_nsec / 1000000;
1675 +
1676 + /* Rebuild hash table for O(1) lookup */
1677 + cache_build_hashtable(cache);
1678 +
1679 + return true;
1680 +}
1681 +
1682 +const nipc_cgroups_cache_item_t *nipc_cgroups_cache_lookup(
1683 + const nipc_cgroups_cache_t *cache,
1684 + uint32_t hash,
1685 + const char *name)
1686 +{
1687 + if (!cache->populated || !cache->items || !name)
1688 + return NULL;
1689 +
1690 + /* Use hash table if available, fall back to linear scan */
1691 + if (cache->buckets && cache->bucket_count > 0) {
1692 + uint32_t key = hash ^ cache_hash_name(name);
1693 + uint32_t mask = cache->bucket_count - 1;
1694 + uint32_t slot = key & mask;
1695 +
1696 + while (cache->buckets[slot].used) {
1697 + uint32_t idx = cache->buckets[slot].index;
1698 + if (cache->items[idx].hash == hash &&
1699 + strcmp(cache->items[idx].name, name) == 0) {
1700 + return &cache->items[idx];
1701 + }
1702 + slot = (slot + 1) & mask;
1703 + }
1704 + return NULL;
1705 + }
1706 +
1707 + /* Fallback linear scan (hash table allocation failed) */
1708 + for (uint32_t i = 0; i < cache->item_count; i++) {
1709 + if (cache->items[i].hash == hash &&
1710 + strcmp(cache->items[i].name, name) == 0) {
1711 + return &cache->items[i];
1712 + }
1713 + }
1714 +
1715 + return NULL;
1716 +}
1717 +
1718 +void nipc_cgroups_cache_status(const nipc_cgroups_cache_t *cache,
1719 + nipc_cgroups_cache_status_t *out)
1720 +{
1721 + out->populated = cache->populated;
1722 + out->item_count = cache->item_count;
1723 + out->systemd_enabled = cache->systemd_enabled;
1724 + out->generation = cache->generation;
1725 + out->refresh_success_count = cache->refresh_success_count;
1726 + out->refresh_failure_count = cache->refresh_failure_count;
1727 + out->connection_state = cache->client.state;
1728 + out->last_refresh_ts = cache->last_refresh_ts;
1729 +}
1730 +
1731 +void nipc_cgroups_cache_close(nipc_cgroups_cache_t *cache)
1732 +{
1733 + cache_free_items(cache->items, cache->item_count);
1734 + cache->items = NULL;
1735 + cache->item_count = 0;
1736 + cache->populated = false;
1737 +
1738 + free(cache->buckets);
1739 + cache->buckets = NULL;
1740 + cache->bucket_count = 0;
1741 +
1742 + free(cache->response_buf);
1743 + cache->response_buf = NULL;
1744 + cache->response_buf_size = 0;
1745 +
1746 + nipc_client_close(&cache->client);
1747 +}
src/libnetdata/netipc/src/service/netipc_service_win.c new
+1793
@@ -0,0 +1,1793 @@
1 +/*
2 + * netipc_service_win.c - L2 orchestration for Windows.
3 + *
4 + * Pure composition of L1 (Named Pipe / Win SHM) + Codec.
5 + * Identical state machine and retry logic as the POSIX implementation,
6 + * using Windows transport calls instead of UDS/POSIX SHM.
7 + *
8 + * Client context manages connection lifecycle with at-least-once retry.
9 + * Managed server handles accept, read, dispatch, respond.
10 + */
11 +
12 +#if defined(_WIN32) || defined(__MSYS__)
13 +
14 +#include "netipc/netipc_service.h"
15 +#include "netipc/netipc_protocol.h"
16 +#include "netipc/netipc_named_pipe.h"
17 +#include "netipc/netipc_win_shm.h"
18 +
19 +#include <stdlib.h>
20 +#include <string.h>
21 +#include <process.h>
22 +#include <windows.h>
23 +
24 +/* WaitForSingleObject timeout for server poll loops (ms) */
25 +#define SERVER_POLL_TIMEOUT_MS 100
26 +#define NIPC_CLIENT_BUF_DEFAULT 65536u
27 +#define CLIENT_SHM_ATTACH_RETRY_INTERVAL_MS 5u
28 +#define CLIENT_SHM_ATTACH_RETRY_TIMEOUT_MS 5000u
29 +
30 +enum {
31 + NIPC_WIN_SERVICE_TEST_FAULT_CLIENT_RESPONSE_BUF_REALLOC_INTERNAL = 1,
32 + NIPC_WIN_SERVICE_TEST_FAULT_CLIENT_SEND_BUF_REALLOC_INTERNAL,
33 + NIPC_WIN_SERVICE_TEST_FAULT_CLIENT_SHM_CTX_CALLOC_INTERNAL,
34 + NIPC_WIN_SERVICE_TEST_FAULT_SERVER_SHM_CTX_CALLOC_INTERNAL,
35 + NIPC_WIN_SERVICE_TEST_FAULT_SERVER_RECV_BUF_MALLOC_INTERNAL,
36 + NIPC_WIN_SERVICE_TEST_FAULT_SERVER_RESP_BUF_MALLOC_INTERNAL,
37 + NIPC_WIN_SERVICE_TEST_FAULT_SERVER_SESSIONS_CALLOC_INTERNAL,
38 + NIPC_WIN_SERVICE_TEST_FAULT_SERVER_SESSION_CTX_CALLOC_INTERNAL,
39 + NIPC_WIN_SERVICE_TEST_FAULT_SERVER_THREAD_CREATE_INTERNAL,
40 + NIPC_WIN_SERVICE_TEST_FAULT_CACHE_BUCKETS_CALLOC_INTERNAL,
41 + NIPC_WIN_SERVICE_TEST_FAULT_CACHE_ITEMS_CALLOC_INTERNAL,
42 + NIPC_WIN_SERVICE_TEST_FAULT_CACHE_ITEM_NAME_MALLOC_INTERNAL,
43 + NIPC_WIN_SERVICE_TEST_FAULT_CACHE_ITEM_PATH_MALLOC_INTERNAL,
44 +};
45 +
46 +static uint64_t g_win_service_test_fault_state = 0;
47 +
48 +static uint64_t win_service_fault_state_make(int site, uint32_t skip_matches)
49 +{
50 + return ((uint64_t)skip_matches << 32) | (uint32_t)site;
51 +}
52 +
53 +void nipc_win_service_test_fault_set(int site, uint32_t skip_matches)
54 +{
55 + __atomic_store_n(&g_win_service_test_fault_state,
56 + win_service_fault_state_make(site, skip_matches),
57 + __ATOMIC_RELEASE);
58 +}
59 +
60 +void nipc_win_service_test_fault_clear(void)
61 +{
62 + __atomic_store_n(&g_win_service_test_fault_state, 0, __ATOMIC_RELEASE);
63 +}
64 +
65 +static bool service_test_should_fail(int site)
66 +{
67 + for (;;) {
68 + uint64_t state = __atomic_load_n(&g_win_service_test_fault_state,
69 + __ATOMIC_ACQUIRE);
70 + uint32_t current_site = (uint32_t)state;
71 + uint32_t current_skip = (uint32_t)(state >> 32);
72 + uint64_t next_state;
73 +
74 + if (current_site != (uint32_t)site)
75 + return false;
76 +
77 + if (current_skip > 0)
78 + next_state = win_service_fault_state_make(site, current_skip - 1u);
79 + else
80 + next_state = 0;
81 +
82 + if (__atomic_compare_exchange_n(&g_win_service_test_fault_state,
83 + &state, next_state, false,
84 + __ATOMIC_ACQ_REL, __ATOMIC_ACQUIRE))
85 + return current_skip == 0;
86 + }
87 +}
88 +
89 +static void *service_malloc(size_t size, int fault_site)
90 +{
91 + if (service_test_should_fail(fault_site))
92 + return NULL;
93 + return malloc(size);
94 +}
95 +
96 +static void *service_calloc(size_t count, size_t size, int fault_site)
97 +{
98 + if (service_test_should_fail(fault_site))
99 + return NULL;
100 + return calloc(count, size);
101 +}
102 +
103 +static void *service_realloc(void *ptr, size_t size, int fault_site)
104 +{
105 + if (service_test_should_fail(fault_site))
106 + return NULL;
107 + return realloc(ptr, size);
108 +}
109 +
110 +static uintptr_t service_beginthreadex(void *security,
111 + unsigned stack_size,
112 + unsigned (__stdcall *start_address)(void *),
113 + void *arglist,
114 + unsigned initflag,
115 + unsigned *thrdaddr)
116 +{
117 + if (service_test_should_fail(NIPC_WIN_SERVICE_TEST_FAULT_SERVER_THREAD_CREATE_INTERNAL))
118 + return 0;
119 +#ifdef __MSYS__
120 + {
121 + DWORD thread_id = 0;
122 + HANDLE thread = CreateThread((LPSECURITY_ATTRIBUTES)security,
123 + (SIZE_T)stack_size,
124 + (LPTHREAD_START_ROUTINE)start_address,
125 + arglist,
126 + (DWORD)initflag,
127 + &thread_id);
128 + if (thrdaddr)
129 + *thrdaddr = (unsigned)thread_id;
130 + return (uintptr_t)thread;
131 + }
132 +#else
133 + return _beginthreadex(security, stack_size, start_address,
134 + arglist, initflag, thrdaddr);
135 +#endif
136 +}
137 +
138 +static uint32_t next_power_of_2_u32(uint32_t n)
139 +{
140 + if (n < 16)
141 + return 16;
142 + n--;
143 + n |= n >> 1;
144 + n |= n >> 2;
145 + n |= n >> 4;
146 + n |= n >> 8;
147 + n |= n >> 16;
148 + return n + 1;
149 +}
150 +
151 +static bool ensure_buffer(uint8_t **buf, size_t *buf_size, size_t need, int fault_site)
152 +{
153 + if (*buf && *buf_size >= need)
154 + return true;
155 +
156 + uint8_t *new_buf = service_realloc(*buf, need, fault_site);
157 + if (!new_buf)
158 + return false;
159 +
160 + *buf = new_buf;
161 + *buf_size = need;
162 + return true;
163 +}
164 +
165 +static void client_note_request_capacity(nipc_client_ctx_t *ctx, uint32_t payload_len)
166 +{
167 + uint32_t grown = next_power_of_2_u32(payload_len);
168 + if (grown > NIPC_MAX_PAYLOAD_CAP)
169 + grown = NIPC_MAX_PAYLOAD_CAP;
170 + if (grown > ctx->transport_config.max_request_payload_bytes)
171 + ctx->transport_config.max_request_payload_bytes = grown;
172 +}
173 +
174 +static void client_note_response_capacity(nipc_client_ctx_t *ctx, uint32_t payload_len)
175 +{
176 + uint32_t grown = next_power_of_2_u32(payload_len);
177 + if (grown > NIPC_MAX_PAYLOAD_CAP)
178 + grown = NIPC_MAX_PAYLOAD_CAP;
179 + if (grown > ctx->transport_config.max_response_payload_bytes)
180 + ctx->transport_config.max_response_payload_bytes = grown;
181 +}
182 +
183 +static bool client_prepare_session_buffers(nipc_client_ctx_t *ctx)
184 +{
185 + size_t response_need = (size_t)ctx->session.max_response_payload_bytes + NIPC_HEADER_LEN;
186 + if (response_need < NIPC_HEADER_LEN + 1024u)
187 + response_need = NIPC_HEADER_LEN + 1024u;
188 +
189 + if (!ensure_buffer(&ctx->response_buf, &ctx->response_buf_size, response_need,
190 + NIPC_WIN_SERVICE_TEST_FAULT_CLIENT_RESPONSE_BUF_REALLOC_INTERNAL))
191 + return false;
192 +
193 + if (ctx->session.selected_profile == NIPC_WIN_SHM_PROFILE_HYBRID ||
194 + ctx->session.selected_profile == NIPC_WIN_SHM_PROFILE_BUSYWAIT) {
195 + size_t send_need = (size_t)ctx->session.max_request_payload_bytes + NIPC_HEADER_LEN;
196 + if (!ensure_buffer(&ctx->send_buf, &ctx->send_buf_size, send_need,
197 + NIPC_WIN_SERVICE_TEST_FAULT_CLIENT_SEND_BUF_REALLOC_INTERNAL))
198 + return false;
199 + }
200 +
201 + return true;
202 +}
203 +
204 +static uint32_t cgroups_request_payload_default(void)
205 +{
206 + return 16u;
207 +}
208 +
209 +static uint32_t cgroups_response_payload_default(void)
210 +{
211 + return NIPC_CLIENT_BUF_DEFAULT;
212 +}
213 +
214 +static nipc_np_client_config_t service_client_config_to_transport(
215 + const nipc_client_config_t *config)
216 +{
217 + nipc_np_client_config_t transport = {0};
218 +
219 + if (!config)
220 + return transport;
221 +
222 + transport.supported_profiles = config->supported_profiles;
223 + transport.preferred_profiles = config->preferred_profiles;
224 + transport.max_request_batch_items = config->max_request_batch_items;
225 + transport.max_response_payload_bytes = config->max_response_payload_bytes;
226 + transport.max_response_batch_items = config->max_request_batch_items;
227 + transport.auth_token = config->auth_token;
228 +
229 + return transport;
230 +}
231 +
232 +static nipc_np_server_config_t service_server_config_to_transport(
233 + const nipc_server_config_t *config)
234 +{
235 + nipc_np_server_config_t transport = {0};
236 +
237 + if (!config)
238 + return transport;
239 +
240 + transport.supported_profiles = config->supported_profiles;
241 + transport.preferred_profiles = config->preferred_profiles;
242 + transport.max_request_batch_items = config->max_request_batch_items;
243 + transport.max_response_payload_bytes = config->max_response_payload_bytes;
244 + transport.max_response_batch_items = config->max_request_batch_items;
245 + transport.auth_token = config->auth_token;
246 +
247 + return transport;
248 +}
249 +
250 +static void server_note_request_capacity(nipc_managed_server_t *server,
251 + uint32_t payload_len);
252 +static void server_note_response_capacity(nipc_managed_server_t *server,
253 + uint32_t payload_len);
254 +
255 +/* ------------------------------------------------------------------ */
256 +/* Internal: client connection helpers */
257 +/* ------------------------------------------------------------------ */
258 +
259 +/* Tear down the current connection (Named Pipe session + Win SHM). */
260 +static void client_disconnect(nipc_client_ctx_t *ctx)
261 +{
262 + if (ctx->shm) {
263 + nipc_win_shm_close(ctx->shm);
264 + free(ctx->shm);
265 + ctx->shm = NULL;
266 + }
267 +
268 + if (ctx->session_valid) {
269 + nipc_np_close_session(&ctx->session);
270 + ctx->session_valid = false;
271 + }
272 +}
273 +
274 +static void client_disable_shm_profiles(nipc_client_ctx_t *ctx)
275 +{
276 + ctx->transport_config.supported_profiles &=
277 + ~(NIPC_WIN_SHM_PROFILE_HYBRID | NIPC_WIN_SHM_PROFILE_BUSYWAIT);
278 + ctx->transport_config.preferred_profiles &=
279 + ~(NIPC_WIN_SHM_PROFILE_HYBRID | NIPC_WIN_SHM_PROFILE_BUSYWAIT);
280 +}
281 +
282 +/* Attempt a full connection: Named Pipe connect + handshake, then
283 + * Win SHM upgrade if negotiated. Returns the new state. */
284 +static nipc_client_state_t client_try_connect(nipc_client_ctx_t *ctx)
285 +{
286 + nipc_np_session_t session;
287 + memset(&session, 0, sizeof(session));
288 + session.pipe = INVALID_HANDLE_VALUE;
289 +
290 + nipc_np_error_t err = nipc_np_connect(
291 + ctx->run_dir, ctx->service_name,
292 + &ctx->transport_config, &session);
293 +
294 + switch (err) {
295 + case NIPC_NP_OK:
296 + break;
297 + case NIPC_NP_ERR_CONNECT:
298 + return NIPC_CLIENT_NOT_FOUND;
299 + case NIPC_NP_ERR_AUTH_FAILED:
300 + return NIPC_CLIENT_AUTH_FAILED;
301 + case NIPC_NP_ERR_NO_PROFILE:
302 + case NIPC_NP_ERR_INCOMPATIBLE:
303 + return NIPC_CLIENT_INCOMPATIBLE;
304 + default:
305 + return NIPC_CLIENT_DISCONNECTED;
306 + }
307 +
308 + ctx->session = session;
309 + ctx->session_valid = true;
310 +
311 + if (!client_prepare_session_buffers(ctx)) {
312 + nipc_np_close_session(&ctx->session);
313 + ctx->session_valid = false;
314 + return NIPC_CLIENT_DISCONNECTED;
315 + }
316 +
317 + /* Win SHM upgrade if negotiated */
318 + if (session.selected_profile == NIPC_WIN_SHM_PROFILE_HYBRID ||
319 + session.selected_profile == NIPC_WIN_SHM_PROFILE_BUSYWAIT) {
320 +
321 + nipc_win_shm_ctx_t *shm = service_calloc(
322 + 1, sizeof(nipc_win_shm_ctx_t),
323 + NIPC_WIN_SERVICE_TEST_FAULT_CLIENT_SHM_CTX_CALLOC_INTERNAL);
324 + if (!shm) {
325 + nipc_np_close_session(&ctx->session);
326 + ctx->session_valid = false;
327 + return NIPC_CLIENT_DISCONNECTED;
328 + }
329 + {
330 + /* Retry attach: the server prepares SHM before accepting the
331 + * handshake, but filesystem/object visibility may lag briefly. */
332 + nipc_win_shm_error_t serr = NIPC_WIN_SHM_ERR_OPEN_MAPPING;
333 + ULONGLONG deadline_ms = GetTickCount64() + CLIENT_SHM_ATTACH_RETRY_TIMEOUT_MS;
334 + for (;;) {
335 + serr = nipc_win_shm_client_attach(
336 + ctx->run_dir, ctx->service_name,
337 + ctx->transport_config.auth_token,
338 + session.session_id,
339 + session.selected_profile,
340 + shm);
341 + if (serr == NIPC_WIN_SHM_OK)
342 + break;
343 + if (serr != NIPC_WIN_SHM_ERR_OPEN_MAPPING &&
344 + serr != NIPC_WIN_SHM_ERR_OPEN_EVENT &&
345 + serr != NIPC_WIN_SHM_ERR_BAD_MAGIC)
346 + break;
347 + if (GetTickCount64() >= deadline_ms)
348 + break;
349 + Sleep(CLIENT_SHM_ATTACH_RETRY_INTERVAL_MS);
350 + }
351 +
352 + if (serr == NIPC_WIN_SHM_OK) {
353 + ctx->shm = shm;
354 + } else {
355 + /* WinSHM attach failed after negotiation. Close that
356 + * session, blacklist WinSHM for this client context, and
357 + * retry baseline via a new handshake. */
358 + free(shm);
359 + nipc_np_close_session(&ctx->session);
360 + ctx->session_valid = false;
361 + client_disable_shm_profiles(ctx);
362 + if (ctx->transport_config.supported_profiles == 0)
363 + return NIPC_CLIENT_DISCONNECTED;
364 + return client_try_connect(ctx);
365 + }
366 + }
367 + }
368 +
369 + return NIPC_CLIENT_READY;
370 +}
371 +
372 +/* ------------------------------------------------------------------ */
373 +/* Internal: send/receive via the active transport */
374 +/* ------------------------------------------------------------------ */
375 +
376 +static nipc_error_t transport_send(nipc_client_ctx_t *ctx,
377 + nipc_header_t *hdr,
378 + const void *payload,
379 + size_t payload_len)
380 +{
381 + if (payload_len > UINT32_MAX)
382 + return NIPC_ERR_OVERFLOW;
383 +
384 + if (ctx->shm) {
385 + if (payload_len > ctx->session.max_request_payload_bytes) {
386 + client_note_request_capacity(ctx, (uint32_t)payload_len);
387 + return NIPC_ERR_OVERFLOW;
388 + }
389 +
390 + size_t msg_len = NIPC_HEADER_LEN + payload_len;
391 + uint8_t *msg = ctx->send_buf;
392 + if (!msg || msg_len > ctx->send_buf_size)
393 + return NIPC_ERR_OVERFLOW;
394 +
395 + hdr->magic = NIPC_MAGIC_MSG;
396 + hdr->version = NIPC_VERSION;
397 + hdr->header_len = NIPC_HEADER_LEN;
398 + hdr->payload_len = (uint32_t)payload_len;
399 +
400 + nipc_header_encode(hdr, msg, NIPC_HEADER_LEN);
401 + if (payload_len > 0)
402 + memcpy(msg + NIPC_HEADER_LEN, payload, payload_len);
403 +
404 + nipc_win_shm_error_t serr = nipc_win_shm_send(ctx->shm, msg, msg_len);
405 + if (serr == NIPC_WIN_SHM_ERR_MSG_TOO_LARGE) {
406 + client_note_request_capacity(ctx, (uint32_t)payload_len);
407 + return NIPC_ERR_OVERFLOW;
408 + }
409 + return (serr == NIPC_WIN_SHM_OK) ? NIPC_OK : NIPC_ERR_NOT_READY;
410 + }
411 +
412 + /* Named Pipe path */
413 + nipc_np_error_t uerr = nipc_np_send(&ctx->session, hdr,
414 + payload, payload_len);
415 + if (uerr == NIPC_NP_ERR_LIMIT_EXCEEDED) {
416 + client_note_request_capacity(ctx, (uint32_t)payload_len);
417 + return NIPC_ERR_OVERFLOW;
418 + }
419 + return (uerr == NIPC_NP_OK) ? NIPC_OK : NIPC_ERR_NOT_READY;
420 +}
421 +
422 +static nipc_error_t transport_receive(nipc_client_ctx_t *ctx,
423 + void *buf, size_t buf_size,
424 + nipc_header_t *hdr_out,
425 + const void **payload_out,
426 + size_t *payload_len_out)
427 +{
428 + if (ctx->shm) {
429 + size_t msg_len;
430 + nipc_win_shm_error_t serr = nipc_win_shm_receive(ctx->shm, buf, buf_size,
431 + &msg_len, 30000);
432 + if (serr != NIPC_WIN_SHM_OK)
433 + return NIPC_ERR_TRUNCATED;
434 +
435 + if (msg_len < NIPC_HEADER_LEN)
436 + return NIPC_ERR_TRUNCATED;
437 +
438 + nipc_error_t perr = nipc_header_decode(buf, msg_len, hdr_out);
439 + if (perr != NIPC_OK)
440 + return perr;
441 +
442 + *payload_out = (const uint8_t *)buf + NIPC_HEADER_LEN;
443 + *payload_len_out = msg_len - NIPC_HEADER_LEN;
444 + return NIPC_OK;
445 + }
446 +
447 + /* Named Pipe path */
448 + nipc_np_error_t uerr = nipc_np_receive(&ctx->session, buf, buf_size,
449 + hdr_out, payload_out,
450 + payload_len_out);
451 + return (uerr == NIPC_NP_OK) ? NIPC_OK : NIPC_ERR_TRUNCATED;
452 +}
453 +
454 +/* ------------------------------------------------------------------ */
455 +/* Internal: generic raw call (send request, receive response) */
456 +/* ------------------------------------------------------------------ */
457 +
458 +/*
459 + * Single-attempt raw call: build envelope, send, receive, validate
460 + * envelope. The caller handles encode before and decode after.
461 + *
462 + * On success, response_payload_out and response_len_out point into the
463 + * internal client response buffer (valid until next call on this context).
464 + */
465 +static nipc_error_t do_raw_call(nipc_client_ctx_t *ctx,
466 + uint16_t method_code,
467 + const void *request_payload,
468 + size_t request_len,
469 + const void **response_payload_out,
470 + size_t *response_len_out)
471 +{
472 + nipc_header_t hdr = {0};
473 + hdr.kind = NIPC_KIND_REQUEST;
474 + hdr.code = method_code;
475 + hdr.flags = 0;
476 + hdr.item_count = 1;
477 + hdr.message_id = (uint64_t)(ctx->call_count + 1);
478 + hdr.transport_status = NIPC_STATUS_OK;
479 +
480 + nipc_error_t err = transport_send(ctx, &hdr, request_payload, request_len);
481 + if (err != NIPC_OK)
482 + return err;
483 +
484 + nipc_header_t resp_hdr;
485 + err = transport_receive(ctx, ctx->response_buf, ctx->response_buf_size,
486 + &resp_hdr, response_payload_out, response_len_out);
487 + if (err != NIPC_OK)
488 + return err;
489 +
490 + if (resp_hdr.kind != NIPC_KIND_RESPONSE)
491 + return NIPC_ERR_BAD_KIND;
492 + if (resp_hdr.code != method_code)
493 + return NIPC_ERR_BAD_LAYOUT;
494 + if (resp_hdr.message_id != hdr.message_id)
495 + return NIPC_ERR_BAD_LAYOUT;
496 + switch (resp_hdr.transport_status) {
497 + case NIPC_STATUS_OK:
498 + break;
499 + case NIPC_STATUS_LIMIT_EXCEEDED:
500 + if (ctx->session.max_response_payload_bytes > 0) {
501 + uint32_t current = ctx->session.max_response_payload_bytes;
502 + client_note_response_capacity(
503 + ctx, current >= UINT32_MAX / 2u ? UINT32_MAX : current * 2u);
504 + }
505 + return NIPC_ERR_OVERFLOW;
506 + case NIPC_STATUS_UNSUPPORTED:
507 + return NIPC_ERR_BAD_LAYOUT;
508 + case NIPC_STATUS_BAD_ENVELOPE:
509 + case NIPC_STATUS_INTERNAL_ERROR:
510 + default:
511 + return NIPC_ERR_BAD_LAYOUT;
512 + }
513 +
514 + return NIPC_OK;
515 +}
516 +
517 +/*
518 + * Generic call-with-retry:
519 + * - ordinary failures reconnect and retry once
520 + * - overflow-driven resize recovery may reconnect repeatedly until
521 + * negotiated capacities grow or recovery fails
522 + * The caller provides a function pointer for the single-attempt logic.
523 + */
524 +typedef nipc_error_t (*nipc_attempt_fn)(nipc_client_ctx_t *ctx, void *state);
525 +
526 +static nipc_error_t call_with_retry(nipc_client_ctx_t *ctx,
527 + nipc_attempt_fn attempt,
528 + void *state)
529 +{
530 + if (ctx->state != NIPC_CLIENT_READY) {
531 + ctx->error_count++;
532 + return NIPC_ERR_NOT_READY;
533 + }
534 +
535 + /* Cap overflow-driven retries: payloads grow by powers of 2, so 8
536 + * retries allows ~256x growth from the initial negotiated size. */
537 + int overflow_retries = 0;
538 + for (;;) {
539 + uint32_t prev_req = ctx->session.max_request_payload_bytes;
540 + uint32_t prev_resp = ctx->session.max_response_payload_bytes;
541 + uint32_t prev_cfg_req = ctx->transport_config.max_request_payload_bytes;
542 + uint32_t prev_cfg_resp = ctx->transport_config.max_response_payload_bytes;
543 +
544 + nipc_error_t err = attempt(ctx, state);
545 + if (err == NIPC_OK) {
546 + ctx->call_count++;
547 + return NIPC_OK;
548 + }
549 +
550 + if (err != NIPC_ERR_OVERFLOW) {
551 + client_disconnect(ctx);
552 + ctx->state = NIPC_CLIENT_BROKEN;
553 + ctx->state = client_try_connect(ctx);
554 + if (ctx->state != NIPC_CLIENT_READY) {
555 + ctx->error_count++;
556 + return err;
557 + }
558 +
559 + ctx->reconnect_count++;
560 + err = attempt(ctx, state);
561 + if (err == NIPC_OK) {
562 + ctx->call_count++;
563 + return NIPC_OK;
564 + }
565 +
566 + client_disconnect(ctx);
567 + ctx->state = NIPC_CLIENT_BROKEN;
568 + ctx->error_count++;
569 + return err;
570 + }
571 +
572 + client_disconnect(ctx);
573 + ctx->state = NIPC_CLIENT_BROKEN;
574 + ctx->state = client_try_connect(ctx);
575 + if (ctx->state != NIPC_CLIENT_READY) {
576 + ctx->error_count++;
577 + return err;
578 + }
579 + ctx->reconnect_count++;
580 +
581 + if (ctx->session.max_request_payload_bytes <= prev_req &&
582 + ctx->session.max_response_payload_bytes <= prev_resp &&
583 + ctx->transport_config.max_request_payload_bytes <= prev_cfg_req &&
584 + ctx->transport_config.max_response_payload_bytes <= prev_cfg_resp) {
585 + client_disconnect(ctx);
586 + ctx->state = NIPC_CLIENT_BROKEN;
587 + ctx->error_count++;
588 + return err;
589 + }
590 +
591 + if (++overflow_retries >= 8) {
592 + client_disconnect(ctx);
593 + ctx->state = NIPC_CLIENT_BROKEN;
594 + ctx->error_count++;
595 + return err;
596 + }
597 + }
598 +}
599 +
600 +/* ------------------------------------------------------------------ */
601 +/* Internal: single attempt at a cgroups snapshot call */
602 +/* ------------------------------------------------------------------ */
603 +
604 +typedef struct {
605 + nipc_cgroups_resp_view_t *view_out;
606 +} cgroups_call_state_t;
607 +
608 +static nipc_error_t do_cgroups_attempt(nipc_client_ctx_t *ctx, void *state)
609 +{
610 + cgroups_call_state_t *s = (cgroups_call_state_t *)state;
611 +
612 + nipc_cgroups_req_t req = { .layout_version = 1, .flags = 0 };
613 + uint8_t req_buf[4];
614 + size_t req_len = nipc_cgroups_req_encode(&req, req_buf, sizeof(req_buf));
615 + if (req_len == 0)
616 + return NIPC_ERR_TRUNCATED;
617 +
618 + const void *payload;
619 + size_t payload_len;
620 + nipc_error_t err = do_raw_call(ctx, NIPC_METHOD_CGROUPS_SNAPSHOT,
621 + req_buf, req_len,
622 + &payload, &payload_len);
623 + if (err != NIPC_OK)
624 + return err;
625 +
626 + return nipc_cgroups_resp_decode(payload, payload_len, s->view_out);
627 +}
628 +
629 +/* ------------------------------------------------------------------ */
630 +/* Public API: client lifecycle */
631 +/* ------------------------------------------------------------------ */
632 +
633 +void nipc_client_init(nipc_client_ctx_t *ctx,
634 + const char *run_dir,
635 + const char *service_name,
636 + const nipc_client_config_t *config)
637 +{
638 + memset(ctx, 0, sizeof(*ctx));
639 + ctx->state = NIPC_CLIENT_DISCONNECTED;
640 + ctx->session.pipe = INVALID_HANDLE_VALUE;
641 + ctx->session_valid = false;
642 + ctx->shm = NULL;
643 +
644 + if (run_dir) {
645 + size_t len = strlen(run_dir);
646 + if (len >= sizeof(ctx->run_dir))
647 + len = sizeof(ctx->run_dir) - 1;
648 + memcpy(ctx->run_dir, run_dir, len);
649 + ctx->run_dir[len] = '\0';
650 + }
651 +
652 + if (service_name) {
653 + size_t len = strlen(service_name);
654 + if (len >= sizeof(ctx->service_name))
655 + len = sizeof(ctx->service_name) - 1;
656 + memcpy(ctx->service_name, service_name, len);
657 + ctx->service_name[len] = '\0';
658 + }
659 +
660 + ctx->transport_config = service_client_config_to_transport(config);
661 + if (ctx->transport_config.max_request_payload_bytes == 0)
662 + ctx->transport_config.max_request_payload_bytes = cgroups_request_payload_default();
663 + if (ctx->transport_config.max_response_payload_bytes == 0)
664 + ctx->transport_config.max_response_payload_bytes = cgroups_response_payload_default();
665 +}
666 +
667 +bool nipc_client_refresh(nipc_client_ctx_t *ctx)
668 +{
669 + nipc_client_state_t old_state = ctx->state;
670 +
671 + switch (ctx->state) {
672 + case NIPC_CLIENT_DISCONNECTED:
673 + case NIPC_CLIENT_NOT_FOUND:
674 + ctx->state = NIPC_CLIENT_CONNECTING;
675 + ctx->state = client_try_connect(ctx);
676 + if (ctx->state == NIPC_CLIENT_READY)
677 + ctx->connect_count++;
678 + break;
679 +
680 + case NIPC_CLIENT_BROKEN:
681 + client_disconnect(ctx);
682 + ctx->state = NIPC_CLIENT_CONNECTING;
683 + ctx->state = client_try_connect(ctx);
684 + if (ctx->state == NIPC_CLIENT_READY)
685 + ctx->reconnect_count++;
686 + break;
687 +
688 + case NIPC_CLIENT_READY:
689 + case NIPC_CLIENT_CONNECTING:
690 + case NIPC_CLIENT_AUTH_FAILED:
691 + case NIPC_CLIENT_INCOMPATIBLE:
692 + break;
693 + }
694 +
695 + return ctx->state != old_state;
696 +}
697 +
698 +void nipc_client_status(const nipc_client_ctx_t *ctx,
699 + nipc_client_status_t *out)
700 +{
701 + out->state = ctx->state;
702 + out->connect_count = ctx->connect_count;
703 + out->reconnect_count = ctx->reconnect_count;
704 + out->call_count = ctx->call_count;
705 + out->error_count = ctx->error_count;
706 +}
707 +
708 +void nipc_client_close(nipc_client_ctx_t *ctx)
709 +{
710 + client_disconnect(ctx);
711 + free(ctx->response_buf);
712 + free(ctx->send_buf);
713 + ctx->response_buf = NULL;
714 + ctx->send_buf = NULL;
715 + ctx->response_buf_size = 0;
716 + ctx->send_buf_size = 0;
717 + ctx->state = NIPC_CLIENT_DISCONNECTED;
718 +}
719 +
720 +/* ------------------------------------------------------------------ */
721 +/* Public API: typed cgroups snapshot call */
722 +/* ------------------------------------------------------------------ */
723 +
724 +nipc_error_t nipc_client_call_cgroups_snapshot(
725 + nipc_client_ctx_t *ctx,
726 + nipc_cgroups_resp_view_t *view_out)
727 +{
728 + cgroups_call_state_t state = {
729 + .view_out = view_out,
730 + };
731 + return call_with_retry(ctx, do_cgroups_attempt, &state);
732 +}
733 +
734 +/* ------------------------------------------------------------------ */
735 +/* Internal: managed server session handler */
736 +/* ------------------------------------------------------------------ */
737 +
738 +/*
739 + * Handle one client session: read requests, dispatch to handler,
740 + * send responses. Each session gets its own response buffer.
741 + * Runs until the client disconnects or server stops.
742 + */
743 +static void server_handle_session(nipc_managed_server_t *server,
744 + nipc_np_session_t *session,
745 + nipc_win_shm_ctx_t *shm,
746 + uint8_t *resp_buf,
747 + size_t resp_buf_size)
748 +{
749 + /* Dynamically allocate recv buffer based on negotiated max */
750 + size_t recv_size = NIPC_HEADER_LEN + session->max_request_payload_bytes;
751 + if (recv_size < NIPC_HEADER_LEN + 1024u)
752 + recv_size = NIPC_HEADER_LEN + 1024u;
753 + uint8_t *recv_buf = service_malloc(
754 + recv_size, NIPC_WIN_SERVICE_TEST_FAULT_SERVER_RECV_BUF_MALLOC_INTERNAL);
755 + if (!recv_buf)
756 + return;
757 +
758 + while (InterlockedCompareExchange(&server->running, 0, 0)) {
759 + nipc_header_t hdr;
760 + const void *payload;
761 + size_t payload_len;
762 +
763 + /* Receive request via the active transport */
764 + if (shm) {
765 + size_t msg_len;
766 + nipc_win_shm_error_t serr = nipc_win_shm_receive(shm, recv_buf, recv_size,
767 + &msg_len, SERVER_POLL_TIMEOUT_MS);
768 + if (serr == NIPC_WIN_SHM_ERR_TIMEOUT)
769 + continue;
770 + if (serr != NIPC_WIN_SHM_OK)
771 + break;
772 + if (msg_len < NIPC_HEADER_LEN)
773 + break;
774 +
775 + nipc_error_t perr = nipc_header_decode(recv_buf, msg_len, &hdr);
776 + if (perr != NIPC_OK)
777 + break;
778 +
779 + payload = recv_buf + NIPC_HEADER_LEN;
780 + payload_len = msg_len - NIPC_HEADER_LEN;
781 + } else {
782 + /* Named Pipe path: wait for readability first, then receive.
783 + * This mirrors the Go/Rust Windows server loops and avoids
784 + * relying on a blocking ReadFile wake-up for each ping-pong
785 + * request. */
786 + bool readable = false;
787 + nipc_np_error_t werr = nipc_np_wait_readable(
788 + session, SERVER_POLL_TIMEOUT_MS, &readable);
789 + if (werr == NIPC_NP_ERR_DISCONNECTED)
790 + break;
791 + if (werr != NIPC_NP_OK)
792 + break;
793 + if (!readable)
794 + continue;
795 +
796 + nipc_np_error_t uerr = nipc_np_receive(
797 + session, recv_buf, recv_size,
798 + &hdr, &payload, &payload_len);
799 + if (uerr == NIPC_NP_ERR_LIMIT_EXCEEDED) {
800 + if (hdr.kind == NIPC_KIND_REQUEST) {
801 + if (hdr.payload_len > 0)
802 + server_note_request_capacity(server, hdr.payload_len);
803 +
804 + nipc_header_t resp_hdr = {0};
805 + resp_hdr.kind = NIPC_KIND_RESPONSE;
806 + resp_hdr.code = hdr.code;
807 + resp_hdr.message_id = hdr.message_id;
808 + resp_hdr.transport_status = NIPC_STATUS_LIMIT_EXCEEDED;
809 + resp_hdr.item_count = 1;
810 + resp_hdr.flags = 0;
811 +
812 + if (nipc_np_send(session, &resp_hdr, NULL, 0) != NIPC_NP_OK)
813 + break;
814 + }
815 + break;
816 + }
817 + if (uerr != NIPC_NP_OK)
818 + break;
819 + }
820 +
821 + /* Protocol violation: unexpected message kind terminates session */
822 + if (hdr.kind != NIPC_KIND_REQUEST)
823 + break;
824 +
825 + if (hdr.code != server->expected_method_code) {
826 + nipc_header_t resp_hdr = {0};
827 + resp_hdr.kind = NIPC_KIND_RESPONSE;
828 + resp_hdr.code = hdr.code;
829 + resp_hdr.message_id = hdr.message_id;
830 + resp_hdr.transport_status = NIPC_STATUS_UNSUPPORTED;
831 + resp_hdr.item_count = 1;
832 + resp_hdr.flags = 0;
833 +
834 + if (shm) {
835 + uint8_t msg[NIPC_HEADER_LEN];
836 + resp_hdr.magic = NIPC_MAGIC_MSG;
837 + resp_hdr.version = NIPC_VERSION;
838 + resp_hdr.header_len = NIPC_HEADER_LEN;
839 + resp_hdr.payload_len = 0;
840 + nipc_header_encode(&resp_hdr, msg, sizeof(msg));
841 + if (nipc_win_shm_send(shm, msg, sizeof(msg)) != NIPC_WIN_SHM_OK)
842 + break;
843 + } else {
844 + if (nipc_np_send(session, &resp_hdr, NULL, 0) != NIPC_NP_OK)
845 + break;
846 + }
847 + continue;
848 + }
849 +
850 + if (payload_len <= UINT32_MAX)
851 + server_note_request_capacity(server, (uint32_t)payload_len);
852 +
853 + /* Dispatch: one request kind per service endpoint. */
854 + size_t response_len = 0;
855 + nipc_error_t dispatch_err = server->handler(
856 + server->handler_user,
857 + &hdr,
858 + (const uint8_t *)payload, payload_len,
859 + resp_buf, resp_buf_size,
860 + &response_len);
861 +
862 + /* Build response header */
863 + nipc_header_t resp_hdr = {0};
864 + resp_hdr.kind = NIPC_KIND_RESPONSE;
865 + resp_hdr.code = hdr.code;
866 + resp_hdr.message_id = hdr.message_id;
867 + if ((hdr.flags & NIPC_FLAG_BATCH) && hdr.item_count >= 1) {
868 + resp_hdr.item_count = hdr.item_count;
869 + resp_hdr.flags = NIPC_FLAG_BATCH;
870 + } else {
871 + resp_hdr.item_count = 1;
872 + resp_hdr.flags = 0;
873 + }
874 +
875 + switch (dispatch_err) {
876 + case NIPC_OK:
877 + if (response_len > session->max_response_payload_bytes) {
878 + server_note_response_capacity(
879 + server,
880 + response_len >= UINT32_MAX ? UINT32_MAX : (uint32_t)response_len);
881 + resp_hdr.transport_status = NIPC_STATUS_LIMIT_EXCEEDED;
882 + response_len = 0;
883 + } else {
884 + if (response_len <= UINT32_MAX)
885 + server_note_response_capacity(server, (uint32_t)response_len);
886 + resp_hdr.transport_status = NIPC_STATUS_OK;
887 + }
888 + break;
889 + case NIPC_ERR_OVERFLOW:
890 + if (session->max_response_payload_bytes >= UINT32_MAX / 2u)
891 + server_note_response_capacity(server, UINT32_MAX);
892 + else
893 + server_note_response_capacity(server,
894 + session->max_response_payload_bytes * 2u);
895 + resp_hdr.transport_status = NIPC_STATUS_LIMIT_EXCEEDED;
896 + response_len = 0;
897 + break;
898 + case NIPC_ERR_TRUNCATED:
899 + case NIPC_ERR_BAD_LAYOUT:
900 + case NIPC_ERR_OUT_OF_BOUNDS:
901 + case NIPC_ERR_MISSING_NUL:
902 + case NIPC_ERR_BAD_ALIGNMENT:
903 + case NIPC_ERR_BAD_ITEM_COUNT:
904 + resp_hdr.transport_status = NIPC_STATUS_BAD_ENVELOPE;
905 + response_len = 0;
906 + break;
907 + case NIPC_ERR_HANDLER_FAILED:
908 + default:
909 + resp_hdr.transport_status = NIPC_STATUS_INTERNAL_ERROR;
910 + response_len = 0;
911 + break;
912 + }
913 +
914 + /* Send response via the active transport */
915 + if (shm) {
916 + size_t msg_len = NIPC_HEADER_LEN + response_len;
917 +
918 + resp_hdr.magic = NIPC_MAGIC_MSG;
919 + resp_hdr.version = NIPC_VERSION;
920 + resp_hdr.header_len = NIPC_HEADER_LEN;
921 + resp_hdr.payload_len = (uint32_t)response_len;
922 +
923 + uint8_t stack_msg[4096];
924 + uint8_t *msg = (msg_len <= sizeof(stack_msg)) ? stack_msg : malloc(msg_len);
925 + if (!msg)
926 + break;
927 +
928 + nipc_header_encode(&resp_hdr, msg, NIPC_HEADER_LEN);
929 + if (response_len > 0)
930 + memcpy(msg + NIPC_HEADER_LEN, resp_buf, response_len);
931 +
932 + nipc_win_shm_error_t serr = nipc_win_shm_send(shm, msg, msg_len);
933 + if (msg != stack_msg)
934 + free(msg);
935 + if (serr != NIPC_WIN_SHM_OK)
936 + break;
937 + } else {
938 + nipc_np_error_t uerr = nipc_np_send(
939 + session, &resp_hdr, resp_buf, response_len);
940 + if (uerr != NIPC_NP_OK)
941 + break;
942 + }
943 +
944 + if (dispatch_err == NIPC_ERR_OVERFLOW)
945 + break;
946 + }
947 +
948 + free(recv_buf);
949 +}
950 +
951 +static uint32_t server_snapshot_max_items(size_t response_buf_size,
952 + const nipc_cgroups_service_handler_t *service_handler)
953 +{
954 + if (service_handler->snapshot_max_items != 0)
955 + return service_handler->snapshot_max_items;
956 + return nipc_cgroups_builder_estimate_max_items(response_buf_size);
957 +}
958 +
959 +static void server_note_request_capacity(nipc_managed_server_t *server,
960 + uint32_t payload_len)
961 +{
962 + uint32_t grown = next_power_of_2_u32(payload_len);
963 + uint32_t current = server->learned_request_payload_bytes;
964 + while (grown > current) {
965 + uint32_t previous = (uint32_t)InterlockedCompareExchange(
966 + (volatile LONG *)&server->learned_request_payload_bytes,
967 + (LONG)grown, (LONG)current);
968 + if (previous == current)
969 + break;
970 + current = previous;
971 + }
972 +}
973 +
974 +static void server_note_response_capacity(nipc_managed_server_t *server,
975 + uint32_t payload_len)
976 +{
977 + uint32_t grown = next_power_of_2_u32(payload_len);
978 + uint32_t current = server->learned_response_payload_bytes;
979 + while (grown > current) {
980 + uint32_t previous = (uint32_t)InterlockedCompareExchange(
981 + (volatile LONG *)&server->learned_response_payload_bytes,
982 + (LONG)grown, (LONG)current);
983 + if (previous == current)
984 + break;
985 + current = previous;
986 + }
987 +}
988 +
989 +static nipc_error_t server_typed_dispatch(void *user,
990 + const nipc_header_t *request_hdr,
991 + const uint8_t *request_payload,
992 + size_t request_len,
993 + uint8_t *response_buf,
994 + size_t response_buf_size,
995 + size_t *response_len_out)
996 +{
997 + nipc_managed_server_t *server = (nipc_managed_server_t *)user;
998 + nipc_cgroups_service_handler_t *service_handler = &server->service_handler;
999 + (void)request_hdr;
1000 +
1001 + if (!service_handler->handle)
1002 + return NIPC_ERR_HANDLER_FAILED;
1003 +
1004 + return nipc_dispatch_cgroups_snapshot(
1005 + request_payload, request_len,
1006 + response_buf, response_buf_size, response_len_out,
1007 + server_snapshot_max_items(response_buf_size, service_handler),
1008 + service_handler->handle, service_handler->user);
1009 +}
1010 +
1011 +/* ------------------------------------------------------------------ */
1012 +/* Internal: per-session handler thread */
1013 +/* ------------------------------------------------------------------ */
1014 +
1015 +/* Thread function: handles one client session from accept to disconnect. */
1016 +static unsigned __stdcall session_handler_thread(void *arg)
1017 +{
1018 + nipc_session_ctx_t *sctx = (nipc_session_ctx_t *)arg;
1019 + nipc_managed_server_t *server = sctx->server;
1020 + /* Allocate a per-session response buffer */
1021 + size_t resp_size = (size_t)sctx->session.max_response_payload_bytes;
1022 + if (resp_size < 1024u)
1023 + resp_size = 1024u;
1024 + uint8_t *resp_buf = service_malloc(
1025 + resp_size, NIPC_WIN_SERVICE_TEST_FAULT_SERVER_RESP_BUF_MALLOC_INTERNAL);
1026 + if (resp_buf) {
1027 + server_handle_session(server, &sctx->session, sctx->shm,
1028 + resp_buf, resp_size);
1029 + free(resp_buf);
1030 + }
1031 +
1032 + /* Cleanup SHM and session */
1033 + if (sctx->shm) {
1034 + nipc_win_shm_destroy(sctx->shm);
1035 + free(sctx->shm);
1036 + }
1037 + nipc_np_close_session(&sctx->session);
1038 +
1039 + /* Mark inactive; the reap/destroy path owns removal from the array */
1040 + InterlockedExchange((volatile LONG *)&sctx->active, 0);
1041 + return 0;
1042 +}
1043 +
1044 +/* ------------------------------------------------------------------ */
1045 +/* Internal: reap finished session threads */
1046 +/* ------------------------------------------------------------------ */
1047 +
1048 +/* Reap all finished (inactive) session threads. Called with lock held. */
1049 +static void server_reap_sessions_locked(nipc_managed_server_t *server)
1050 +{
1051 + int i = 0;
1052 + while (i < server->session_count) {
1053 + nipc_session_ctx_t *s = server->sessions[i];
1054 + if (!InterlockedCompareExchange((volatile LONG *)&s->active, 0, 0)) {
1055 + WaitForSingleObject(s->thread, INFINITE);
1056 + CloseHandle(s->thread);
1057 + /* Swap with last, free */
1058 + server->sessions[i] = server->sessions[server->session_count - 1];
1059 + server->session_count--;
1060 + free(s);
1061 + } else {
1062 + i++;
1063 + }
1064 + }
1065 +
1066 +}
1067 +
1068 +typedef struct {
1069 + nipc_win_shm_ctx_t *hybrid;
1070 + nipc_win_shm_ctx_t *busywait;
1071 +} prepared_win_shm_t;
1072 +
1073 +static void server_destroy_prepared_win_shm(prepared_win_shm_t *prepared)
1074 +{
1075 + if (!prepared)
1076 + return;
1077 + if (prepared->hybrid) {
1078 + nipc_win_shm_destroy(prepared->hybrid);
1079 + free(prepared->hybrid);
1080 + prepared->hybrid = NULL;
1081 + }
1082 + if (prepared->busywait) {
1083 + nipc_win_shm_destroy(prepared->busywait);
1084 + free(prepared->busywait);
1085 + prepared->busywait = NULL;
1086 + }
1087 +}
1088 +
1089 +static nipc_win_shm_ctx_t *server_take_prepared_win_shm(prepared_win_shm_t *prepared,
1090 + uint32_t profile)
1091 +{
1092 + if (!prepared)
1093 + return NULL;
1094 + if (profile == NIPC_WIN_SHM_PROFILE_HYBRID) {
1095 + nipc_win_shm_ctx_t *ctx = prepared->hybrid;
1096 + prepared->hybrid = NULL;
1097 + return ctx;
1098 + }
1099 + if (profile == NIPC_WIN_SHM_PROFILE_BUSYWAIT) {
1100 + nipc_win_shm_ctx_t *ctx = prepared->busywait;
1101 + prepared->busywait = NULL;
1102 + return ctx;
1103 + }
1104 + return NULL;
1105 +}
1106 +
1107 +static bool server_prepare_accept_config(nipc_managed_server_t *server,
1108 + uint64_t sid,
1109 + nipc_np_server_config_t *cfg_out,
1110 + prepared_win_shm_t *prepared)
1111 +{
1112 + *cfg_out = server->base_config;
1113 + cfg_out->max_request_payload_bytes = server->learned_request_payload_bytes;
1114 + cfg_out->max_response_payload_bytes = server->learned_response_payload_bytes;
1115 + memset(prepared, 0, sizeof(*prepared));
1116 +
1117 + uint32_t shm_profiles = cfg_out->supported_profiles &
1118 + (NIPC_WIN_SHM_PROFILE_HYBRID |
1119 + NIPC_WIN_SHM_PROFILE_BUSYWAIT);
1120 + if (shm_profiles == 0)
1121 + return true;
1122 +
1123 + const uint32_t profiles[] = {
1124 + NIPC_WIN_SHM_PROFILE_HYBRID,
1125 + NIPC_WIN_SHM_PROFILE_BUSYWAIT,
1126 + };
1127 + for (size_t i = 0; i < sizeof(profiles) / sizeof(profiles[0]); i++) {
1128 + uint32_t profile = profiles[i];
1129 + if (!(cfg_out->supported_profiles & profile))
1130 + continue;
1131 +
1132 + nipc_win_shm_ctx_t *ctx = service_calloc(
1133 + 1, sizeof(nipc_win_shm_ctx_t),
1134 + NIPC_WIN_SERVICE_TEST_FAULT_SERVER_SHM_CTX_CALLOC_INTERNAL);
1135 + if (!ctx)
1136 + continue;
1137 +
1138 + /* HELLO has not been read yet, so the request segment must cover any
1139 + * client proposal the handshake may legally echo back. */
1140 + nipc_win_shm_error_t serr = nipc_win_shm_server_create(
1141 + server->run_dir, server->service_name,
1142 + server->auth_token,
1143 + sid,
1144 + profile,
1145 + NIPC_MAX_PAYLOAD_CAP + NIPC_HEADER_LEN,
1146 + cfg_out->max_response_payload_bytes + NIPC_HEADER_LEN,
1147 + ctx);
1148 + if (serr == NIPC_WIN_SHM_OK) {
1149 + if (profile == NIPC_WIN_SHM_PROFILE_HYBRID)
1150 + prepared->hybrid = ctx;
1151 + else
1152 + prepared->busywait = ctx;
1153 + continue;
1154 + }
1155 +
1156 + free(ctx);
1157 + cfg_out->supported_profiles &= ~profile;
1158 + cfg_out->preferred_profiles &= ~profile;
1159 + }
1160 +
1161 + return cfg_out->supported_profiles != 0;
1162 +}
1163 +
1164 +static void server_wake_listener(nipc_managed_server_t *server)
1165 +{
1166 + if (server->listener.pipe == INVALID_HANDLE_VALUE ||
1167 + server->listener.pipe == NULL ||
1168 + server->listener.pipe_name[0] == L'\0')
1169 + return;
1170 +
1171 + HANDLE wake = CreateFileW(
1172 + server->listener.pipe_name,
1173 + GENERIC_READ | GENERIC_WRITE,
1174 + 0,
1175 + NULL,
1176 + OPEN_EXISTING,
1177 + 0,
1178 + NULL);
1179 + if (wake != INVALID_HANDLE_VALUE)
1180 + CloseHandle(wake);
1181 +}
1182 +
1183 +/* Ask active session threads to leave synchronous ReadFile/WriteFile waits.
1184 + * This is safer than targeting pipe handles from another thread because the
1185 + * thread handle stays valid until the owner joins it. */
1186 +static void server_cancel_active_session_io_locked(nipc_managed_server_t *server)
1187 +{
1188 + for (int i = 0; i < server->session_count; i++) {
1189 + nipc_session_ctx_t *s = server->sessions[i];
1190 + if (!s)
1191 + continue;
1192 + if (!InterlockedCompareExchange((volatile LONG *)&s->active, 0, 0))
1193 + continue;
1194 + if (s->thread == NULL || s->thread == INVALID_HANDLE_VALUE)
1195 + continue;
1196 + CancelSynchronousIo(s->thread);
1197 + }
1198 +}
1199 +
1200 +/* ------------------------------------------------------------------ */
1201 +/* Public API: managed server */
1202 +/* ------------------------------------------------------------------ */
1203 +
1204 +static nipc_error_t server_init_raw(nipc_managed_server_t *server,
1205 + const char *run_dir,
1206 + const char *service_name,
1207 + const nipc_np_server_config_t *config,
1208 + int worker_count,
1209 + uint16_t expected_method_code,
1210 + nipc_server_handler_fn handler,
1211 + void *user)
1212 +{
1213 + memset(server, 0, sizeof(*server));
1214 + server->listener.pipe = INVALID_HANDLE_VALUE;
1215 + InterlockedExchange(&server->running, 0);
1216 +
1217 + if (!run_dir || !service_name || !handler)
1218 + return NIPC_ERR_BAD_LAYOUT;
1219 +
1220 + if (worker_count < 1)
1221 + worker_count = 1;
1222 +
1223 + /* Store config */
1224 + {
1225 + size_t len = strlen(run_dir);
1226 + if (len >= sizeof(server->run_dir))
1227 + len = sizeof(server->run_dir) - 1;
1228 + memcpy(server->run_dir, run_dir, len);
1229 + server->run_dir[len] = '\0';
1230 + }
1231 + {
1232 + size_t len = strlen(service_name);
1233 + if (len >= sizeof(server->service_name))
1234 + len = sizeof(server->service_name) - 1;
1235 + memcpy(server->service_name, service_name, len);
1236 + server->service_name[len] = '\0';
1237 + }
1238 +
1239 + server->handler = handler;
1240 + server->handler_user = user;
1241 + server->worker_count = worker_count;
1242 + server->expected_method_code = expected_method_code;
1243 + server->base_config = *config;
1244 + server->learned_request_payload_bytes =
1245 + (config && config->max_request_payload_bytes > 0)
1246 + ? config->max_request_payload_bytes
1247 + : NIPC_MAX_PAYLOAD_DEFAULT;
1248 + server->learned_response_payload_bytes =
1249 + (config && config->max_response_payload_bytes > 0)
1250 + ? config->max_response_payload_bytes
1251 + : NIPC_MAX_PAYLOAD_DEFAULT;
1252 + server->auth_token = config ? config->auth_token : 0;
1253 +
1254 +
1255 + /* Initialize session tracking */
1256 + server->session_capacity = worker_count * 2;
1257 + if (server->session_capacity < 16)
1258 + server->session_capacity = 16;
1259 + server->sessions = service_calloc(
1260 + (size_t)server->session_capacity, sizeof(nipc_session_ctx_t *),
1261 + NIPC_WIN_SERVICE_TEST_FAULT_SERVER_SESSIONS_CALLOC_INTERNAL);
1262 + if (!server->sessions)
1263 + return NIPC_ERR_OVERFLOW;
1264 + server->session_count = 0;
1265 + server->next_session_id = 1; /* spec: monotonic counter starting at 1 */
1266 + InitializeCriticalSection(&server->sessions_lock);
1267 +
1268 + /* Clean up stale SHM kernel objects from previous crashes (no-op on
1269 + * Windows but maintains API symmetry with the POSIX transport). */
1270 + nipc_win_shm_cleanup_stale(run_dir, service_name);
1271 +
1272 + /* Start listening via L1 */
1273 + nipc_np_error_t uerr = nipc_np_listen(
1274 + run_dir, service_name, config, &server->listener);
1275 + if (uerr != NIPC_NP_OK) {
1276 + free(server->sessions);
1277 + server->sessions = NULL;
1278 + DeleteCriticalSection(&server->sessions_lock);
1279 + return NIPC_ERR_BAD_LAYOUT;
1280 + }
1281 +
1282 + return NIPC_OK;
1283 +}
1284 +
1285 +nipc_error_t nipc_server_init_typed(nipc_managed_server_t *server,
1286 + const char *run_dir,
1287 + const char *service_name,
1288 + const nipc_server_config_t *config,
1289 + int worker_count,
1290 + const nipc_cgroups_service_handler_t *service_handler)
1291 +{
1292 + if (!service_handler)
1293 + return NIPC_ERR_BAD_LAYOUT;
1294 +
1295 + nipc_np_server_config_t typed_cfg = service_server_config_to_transport(config);
1296 + if (typed_cfg.max_request_payload_bytes == 0)
1297 + typed_cfg.max_request_payload_bytes = cgroups_request_payload_default();
1298 + if (typed_cfg.max_response_payload_bytes == 0)
1299 + typed_cfg.max_response_payload_bytes = cgroups_response_payload_default();
1300 +
1301 + nipc_error_t err = server_init_raw(server, run_dir, service_name,
1302 + &typed_cfg, worker_count,
1303 + NIPC_METHOD_CGROUPS_SNAPSHOT,
1304 + server_typed_dispatch, server);
1305 + if (err != NIPC_OK)
1306 + return err;
1307 +
1308 + server->service_handler = *service_handler;
1309 + return NIPC_OK;
1310 +}
1311 +
1312 +nipc_error_t nipc_server_init_raw_for_tests(nipc_managed_server_t *server,
1313 + const char *run_dir,
1314 + const char *service_name,
1315 + const nipc_np_server_config_t *config,
1316 + int worker_count,
1317 + uint16_t expected_method_code,
1318 + nipc_server_handler_fn handler,
1319 + void *user)
1320 +{
1321 + return server_init_raw(server, run_dir, service_name, config,
1322 + worker_count, expected_method_code, handler, user);
1323 +}
1324 +
1325 +void nipc_server_run(nipc_managed_server_t *server)
1326 +{
1327 + InterlockedExchange(&server->accept_loop_active, 1);
1328 + InterlockedExchange(&server->running, 1);
1329 +
1330 + while (InterlockedCompareExchange(&server->running, 0, 0)) {
1331 + /* Accept one client via L1 (blocking with internal timeout) */
1332 + nipc_np_session_t session;
1333 + memset(&session, 0, sizeof(session));
1334 + session.pipe = INVALID_HANDLE_VALUE;
1335 +
1336 + uint64_t sid = server->next_session_id++;
1337 + nipc_np_server_config_t accept_cfg;
1338 + prepared_win_shm_t prepared_shm;
1339 + if (!server_prepare_accept_config(server, sid, &accept_cfg, &prepared_shm)) {
1340 + Sleep(10);
1341 + continue;
1342 + }
1343 +
1344 + server->listener.config = accept_cfg;
1345 + nipc_np_error_t uerr = nipc_np_accept(&server->listener, sid, &session);
1346 + if (uerr != NIPC_NP_OK) {
1347 + server_destroy_prepared_win_shm(&prepared_shm);
1348 + if (!InterlockedCompareExchange(&server->running, 0, 0))
1349 + break;
1350 + Sleep(10);
1351 + continue;
1352 + }
1353 +
1354 + server_note_request_capacity(server, session.max_request_payload_bytes);
1355 + server_note_response_capacity(server, session.max_response_payload_bytes);
1356 +
1357 + /* Enforce worker_count limit: reap finished sessions, check count */
1358 + EnterCriticalSection(&server->sessions_lock);
1359 + server_reap_sessions_locked(server);
1360 +
1361 + if (server->session_count >= server->worker_count) {
1362 + /* At capacity: reject this client by closing the session */
1363 + LeaveCriticalSection(&server->sessions_lock);
1364 + server_destroy_prepared_win_shm(&prepared_shm);
1365 + nipc_np_close_session(&session);
1366 + continue;
1367 + }
1368 +
1369 + /* SHM profile guarantee: only negotiate SHM for sessions backed by
1370 + * prepared per-session kernel objects for the selected profile. */
1371 + nipc_win_shm_ctx_t *shm = NULL;
1372 + if (session.selected_profile == NIPC_WIN_SHM_PROFILE_HYBRID ||
1373 + session.selected_profile == NIPC_WIN_SHM_PROFILE_BUSYWAIT) {
1374 + shm = server_take_prepared_win_shm(&prepared_shm, session.selected_profile);
1375 + if (!shm) {
1376 + server_destroy_prepared_win_shm(&prepared_shm);
1377 + LeaveCriticalSection(&server->sessions_lock);
1378 + nipc_np_close_session(&session);
1379 + continue;
1380 + }
1381 + server_destroy_prepared_win_shm(&prepared_shm);
1382 + } else {
1383 + server_destroy_prepared_win_shm(&prepared_shm);
1384 + }
1385 +
1386 + /* Create session context */
1387 + nipc_session_ctx_t *sctx = service_calloc(
1388 + 1, sizeof(nipc_session_ctx_t),
1389 + NIPC_WIN_SERVICE_TEST_FAULT_SERVER_SESSION_CTX_CALLOC_INTERNAL);
1390 + if (!sctx) {
1391 + LeaveCriticalSection(&server->sessions_lock);
1392 + if (shm) { nipc_win_shm_destroy(shm); free(shm); }
1393 + nipc_np_close_session(&session);
1394 + continue;
1395 + }
1396 +
1397 + sctx->server = server;
1398 + sctx->session = session;
1399 + sctx->shm = shm;
1400 + sctx->id = sid;
1401 + InterlockedExchange((volatile LONG *)&sctx->active, 1);
1402 +
1403 + server->sessions[server->session_count++] = sctx;
1404 + LeaveCriticalSection(&server->sessions_lock);
1405 +
1406 + /* Spawn handler thread for this session */
1407 + unsigned tid_unused;
1408 + sctx->thread = (HANDLE)service_beginthreadex(
1409 + NULL, 0, session_handler_thread, sctx, 0, &tid_unused);
1410 + if (sctx->thread == 0) {
1411 + /* Thread creation failed: clean up */
1412 + EnterCriticalSection(&server->sessions_lock);
1413 + for (int i = 0; i < server->session_count; i++) {
1414 + if (server->sessions[i] == sctx) {
1415 + server->sessions[i] = server->sessions[server->session_count - 1];
1416 + server->session_count--;
1417 + break;
1418 + }
1419 + }
1420 + LeaveCriticalSection(&server->sessions_lock);
1421 +
1422 + if (shm) { nipc_win_shm_destroy(shm); free(shm); }
1423 + nipc_np_close_session(&session);
1424 + free(sctx);
1425 + }
1426 + }
1427 +
1428 + InterlockedExchange(&server->accept_loop_active, 0);
1429 + nipc_np_close_listener(&server->listener);
1430 +}
1431 +
1432 +void nipc_server_stop(nipc_managed_server_t *server)
1433 +{
1434 + InterlockedExchange(&server->running, 0);
1435 + if (InterlockedCompareExchange(&server->accept_loop_active, 0, 0))
1436 + server_wake_listener(server);
1437 + else
1438 + nipc_np_close_listener(&server->listener);
1439 +
1440 + if (server->sessions) {
1441 + EnterCriticalSection(&server->sessions_lock);
1442 + server_cancel_active_session_io_locked(server);
1443 + LeaveCriticalSection(&server->sessions_lock);
1444 + }
1445 +}
1446 +
1447 +bool nipc_server_drain(nipc_managed_server_t *server, uint32_t timeout_ms)
1448 +{
1449 + /* 1. Stop accepting new clients */
1450 + InterlockedExchange(&server->running, 0);
1451 + if (InterlockedCompareExchange(&server->accept_loop_active, 0, 0))
1452 + server_wake_listener(server);
1453 + else
1454 + nipc_np_close_listener(&server->listener);
1455 +
1456 + /* 2. Wait for in-flight sessions to complete */
1457 + bool all_drained = true;
1458 + if (server->sessions) {
1459 + ULONGLONG deadline = GetTickCount64() + timeout_ms;
1460 +
1461 + /* Poll until all sessions are inactive or timeout */
1462 + while (1) {
1463 + EnterCriticalSection(&server->sessions_lock);
1464 + int active_count = 0;
1465 + for (int i = 0; i < server->session_count; i++) {
1466 + if (InterlockedCompareExchange(
1467 + (volatile LONG *)&server->sessions[i]->active, 0, 0))
1468 + active_count++;
1469 + }
1470 + LeaveCriticalSection(&server->sessions_lock);
1471 +
1472 + if (active_count == 0)
1473 + break;
1474 +
1475 + if (GetTickCount64() >= deadline) {
1476 + /* Timeout: cancel synchronous session I/O to unblock threads. */
1477 + EnterCriticalSection(&server->sessions_lock);
1478 + server_cancel_active_session_io_locked(server);
1479 + LeaveCriticalSection(&server->sessions_lock);
1480 + all_drained = false;
1481 + break;
1482 + }
1483 +
1484 + Sleep(5); /* 5ms poll interval */
1485 + }
1486 +
1487 + /* 3. Join all session threads */
1488 + EnterCriticalSection(&server->sessions_lock);
1489 + for (int i = 0; i < server->session_count; i++) {
1490 + nipc_session_ctx_t *s = server->sessions[i];
1491 + LeaveCriticalSection(&server->sessions_lock);
1492 + WaitForSingleObject(s->thread, INFINITE);
1493 + CloseHandle(s->thread);
1494 + free(s);
1495 + EnterCriticalSection(&server->sessions_lock);
1496 + }
1497 + server->session_count = 0;
1498 + LeaveCriticalSection(&server->sessions_lock);
1499 +
1500 + free(server->sessions);
1501 + server->sessions = NULL;
1502 + server->session_capacity = 0;
1503 + DeleteCriticalSection(&server->sessions_lock);
1504 + }
1505 +
1506 + server->worker_count = 0;
1507 +
1508 + return all_drained;
1509 +}
1510 +
1511 +void nipc_server_destroy(nipc_managed_server_t *server)
1512 +{
1513 + InterlockedExchange(&server->running, 0);
1514 + if (InterlockedCompareExchange(&server->accept_loop_active, 0, 0))
1515 + server_wake_listener(server);
1516 + else
1517 + nipc_np_close_listener(&server->listener);
1518 +
1519 + /* Join all active session threads */
1520 + if (server->sessions) {
1521 + EnterCriticalSection(&server->sessions_lock);
1522 + server_cancel_active_session_io_locked(server);
1523 + for (int i = 0; i < server->session_count; i++) {
1524 + nipc_session_ctx_t *s = server->sessions[i];
1525 + LeaveCriticalSection(&server->sessions_lock);
1526 + WaitForSingleObject(s->thread, INFINITE);
1527 + CloseHandle(s->thread);
1528 + free(s);
1529 + EnterCriticalSection(&server->sessions_lock);
1530 + }
1531 + server->session_count = 0;
1532 + LeaveCriticalSection(&server->sessions_lock);
1533 +
1534 + free(server->sessions);
1535 + server->sessions = NULL;
1536 + server->session_capacity = 0;
1537 + DeleteCriticalSection(&server->sessions_lock);
1538 + }
1539 +
1540 + server->worker_count = 0;
1541 +
1542 +}
1543 +
1544 +/* ------------------------------------------------------------------ */
1545 +/* L3: Client-side cgroups snapshot cache */
1546 +/* ------------------------------------------------------------------ */
1547 +
1548 +/* Free all owned strings in cache items and the items array itself. */
1549 +static void cache_free_items(nipc_cgroups_cache_item_t *items, uint32_t count)
1550 +{
1551 + if (!items)
1552 + return;
1553 +
1554 + for (uint32_t i = 0; i < count; i++) {
1555 + free(items[i].name);
1556 + free(items[i].path);
1557 + }
1558 + free(items);
1559 +}
1560 +
1561 +/* Hash a name string (djb2). Combined with item hash for bucket index. */
1562 +static uint32_t cache_hash_name(const char *name)
1563 +{
1564 + uint32_t h = 5381;
1565 + for (const unsigned char *p = (const unsigned char *)name; *p; p++)
1566 + h = ((h << 5) + h) + *p;
1567 + return h;
1568 +}
1569 +
1570 +/*
1571 + * Build the open-addressing hash table from the items array.
1572 + * Uses (item.hash ^ name_hash) as the probe key.
1573 + * Load factor <= 0.5 (bucket_count >= 2 * item_count).
1574 + */
1575 +static bool cache_build_hashtable(nipc_cgroups_cache_t *cache)
1576 +{
1577 + free(cache->buckets);
1578 + cache->buckets = NULL;
1579 + cache->bucket_count = 0;
1580 +
1581 + if (cache->item_count == 0)
1582 + return true;
1583 +
1584 + uint32_t bcount = next_power_of_2_u32(cache->item_count * 2);
1585 + nipc_cgroups_hash_bucket_t *buckets = service_calloc(
1586 + bcount, sizeof(nipc_cgroups_hash_bucket_t),
1587 + NIPC_WIN_SERVICE_TEST_FAULT_CACHE_BUCKETS_CALLOC_INTERNAL);
1588 + if (!buckets)
1589 + return false;
1590 +
1591 + uint32_t mask = bcount - 1;
1592 + for (uint32_t i = 0; i < cache->item_count; i++) {
1593 + uint32_t key = cache->items[i].hash ^ cache_hash_name(cache->items[i].name);
1594 + uint32_t slot = key & mask;
1595 +
1596 + /* Linear probe for an empty bucket */
1597 + while (buckets[slot].used)
1598 + slot = (slot + 1) & mask;
1599 +
1600 + buckets[slot].index = i;
1601 + buckets[slot].used = true;
1602 + }
1603 +
1604 + cache->buckets = buckets;
1605 + cache->bucket_count = bcount;
1606 + return true;
1607 +}
1608 +
1609 +static nipc_cgroups_cache_item_t *cache_build_items(
1610 + const nipc_cgroups_resp_view_t *view,
1611 + uint32_t *count_out)
1612 +{
1613 + uint32_t n = view->item_count;
1614 + *count_out = 0;
1615 +
1616 + if (n == 0)
1617 + return NULL;
1618 +
1619 + nipc_cgroups_cache_item_t *items = service_calloc(
1620 + n, sizeof(nipc_cgroups_cache_item_t),
1621 + NIPC_WIN_SERVICE_TEST_FAULT_CACHE_ITEMS_CALLOC_INTERNAL);
1622 + if (!items)
1623 + return NULL;
1624 +
1625 + for (uint32_t i = 0; i < n; i++) {
1626 + nipc_cgroups_item_view_t iv;
1627 + nipc_error_t err = nipc_cgroups_resp_item(view, i, &iv);
1628 + if (err != NIPC_OK) {
1629 + cache_free_items(items, i);
1630 + return NULL;
1631 + }
1632 +
1633 + items[i].hash = iv.hash;
1634 + items[i].options = iv.options;
1635 + items[i].enabled = iv.enabled;
1636 +
1637 + items[i].name = service_malloc(
1638 + iv.name.len + 1, NIPC_WIN_SERVICE_TEST_FAULT_CACHE_ITEM_NAME_MALLOC_INTERNAL);
1639 + if (!items[i].name) {
1640 + cache_free_items(items, i);
1641 + return NULL;
1642 + }
1643 + if (iv.name.len > 0)
1644 + memcpy(items[i].name, iv.name.ptr, iv.name.len);
1645 + items[i].name[iv.name.len] = '\0';
1646 +
1647 + items[i].path = service_malloc(
1648 + iv.path.len + 1, NIPC_WIN_SERVICE_TEST_FAULT_CACHE_ITEM_PATH_MALLOC_INTERNAL);
1649 + if (!items[i].path) {
1650 + free(items[i].name);
1651 + cache_free_items(items, i);
1652 + return NULL;
1653 + }
1654 + if (iv.path.len > 0)
1655 + memcpy(items[i].path, iv.path.ptr, iv.path.len);
1656 + items[i].path[iv.path.len] = '\0';
1657 + }
1658 +
1659 + *count_out = n;
1660 + return items;
1661 +}
1662 +
1663 +void nipc_cgroups_cache_init(nipc_cgroups_cache_t *cache,
1664 + const char *run_dir,
1665 + const char *service_name,
1666 + const nipc_client_config_t *config)
1667 +{
1668 + memset(cache, 0, sizeof(*cache));
1669 +
1670 + nipc_client_init(&cache->client, run_dir, service_name, config);
1671 +
1672 + cache->items = NULL;
1673 + cache->item_count = 0;
1674 + cache->systemd_enabled = 0;
1675 + cache->generation = 0;
1676 + cache->populated = false;
1677 + cache->buckets = NULL;
1678 + cache->bucket_count = 0;
1679 + cache->refresh_success_count = 0;
1680 + cache->refresh_failure_count = 0;
1681 +
1682 + cache->response_buf = NULL;
1683 + cache->response_buf_size = 0;
1684 +}
1685 +
1686 +bool nipc_cgroups_cache_refresh(nipc_cgroups_cache_t *cache)
1687 +{
1688 + nipc_client_refresh(&cache->client);
1689 +
1690 + nipc_cgroups_resp_view_t view;
1691 + nipc_error_t err = nipc_client_call_cgroups_snapshot(&cache->client, &view);
1692 +
1693 + if (err != NIPC_OK) {
1694 + cache->refresh_failure_count++;
1695 + return false;
1696 + }
1697 +
1698 + uint32_t new_count = 0;
1699 + nipc_cgroups_cache_item_t *new_items = NULL;
1700 +
1701 + if (view.item_count > 0) {
1702 + new_items = cache_build_items(&view, &new_count);
1703 + if (!new_items && view.item_count > 0) {
1704 + cache->refresh_failure_count++;
1705 + return false;
1706 + }
1707 + }
1708 +
1709 + cache_free_items(cache->items, cache->item_count);
1710 + cache->items = new_items;
1711 + cache->item_count = new_count;
1712 + cache->systemd_enabled = view.systemd_enabled;
1713 + cache->generation = view.generation;
1714 + cache->populated = true;
1715 + cache->refresh_success_count++;
1716 +
1717 + /* Record monotonic timestamp (GetTickCount64 is always available on Windows) */
1718 + cache->last_refresh_ts = GetTickCount64();
1719 +
1720 + /* Rebuild hash table for O(1) lookup */
1721 + cache_build_hashtable(cache);
1722 +
1723 + return true;
1724 +}
1725 +
1726 +const nipc_cgroups_cache_item_t *nipc_cgroups_cache_lookup(
1727 + const nipc_cgroups_cache_t *cache,
1728 + uint32_t hash,
1729 + const char *name)
1730 +{
1731 + if (!cache->populated || !cache->items || !name)
1732 + return NULL;
1733 +
1734 + /* Use hash table if available, fall back to linear scan */
1735 + if (cache->buckets && cache->bucket_count > 0) {
1736 + uint32_t key = hash ^ cache_hash_name(name);
1737 + uint32_t mask = cache->bucket_count - 1;
1738 + uint32_t slot = key & mask;
1739 +
1740 + while (cache->buckets[slot].used) {
1741 + uint32_t idx = cache->buckets[slot].index;
1742 + if (cache->items[idx].hash == hash &&
1743 + strcmp(cache->items[idx].name, name) == 0) {
1744 + return &cache->items[idx];
1745 + }
1746 + slot = (slot + 1) & mask;
1747 + }
1748 + return NULL;
1749 + }
1750 +
1751 + /* Fallback linear scan (hash table allocation failed) */
1752 + for (uint32_t i = 0; i < cache->item_count; i++) {
1753 + if (cache->items[i].hash == hash &&
1754 + strcmp(cache->items[i].name, name) == 0) {
1755 + return &cache->items[i];
1756 + }
1757 + }
1758 +
1759 + return NULL;
1760 +}
1761 +
1762 +void nipc_cgroups_cache_status(const nipc_cgroups_cache_t *cache,
1763 + nipc_cgroups_cache_status_t *out)
1764 +{
1765 + out->populated = cache->populated;
1766 + out->item_count = cache->item_count;
1767 + out->systemd_enabled = cache->systemd_enabled;
1768 + out->generation = cache->generation;
1769 + out->refresh_success_count = cache->refresh_success_count;
1770 + out->refresh_failure_count = cache->refresh_failure_count;
1771 + out->connection_state = cache->client.state;
1772 + out->last_refresh_ts = cache->last_refresh_ts;
1773 +}
1774 +
1775 +void nipc_cgroups_cache_close(nipc_cgroups_cache_t *cache)
1776 +{
1777 + cache_free_items(cache->items, cache->item_count);
1778 + cache->items = NULL;
1779 + cache->item_count = 0;
1780 + cache->populated = false;
1781 +
1782 + free(cache->buckets);
1783 + cache->buckets = NULL;
1784 + cache->bucket_count = 0;
1785 +
1786 + free(cache->response_buf);
1787 + cache->response_buf = NULL;
1788 + cache->response_buf_size = 0;
1789 +
1790 + nipc_client_close(&cache->client);
1791 +}
1792 +
1793 +#endif /* _WIN32 || __MSYS__ */
src/libnetdata/netipc/src/transport/posix/netipc_shm.c new
+786
@@ -0,0 +1,786 @@
1 +/*
2 + * netipc_shm.c - L1 POSIX SHM transport (Linux only).
3 + *
4 + * Shared memory data plane with spin+futex synchronization.
5 + * Uses the same wire envelope as UDS -- higher levels are unaware
6 + * of the underlying transport.
7 + */
8 +
9 +#include "netipc/netipc_shm.h"
10 +
11 +#include <dirent.h>
12 +#include <errno.h>
13 +#include <fcntl.h>
14 +#include <inttypes.h>
15 +#include <signal.h>
16 +#include <stdio.h>
17 +#include <stdlib.h>
18 +#include <string.h>
19 +#include <time.h>
20 +#include <unistd.h>
21 +
22 +#include <sys/mman.h>
23 +#include <sys/stat.h>
24 +#include <sys/syscall.h>
25 +
26 +#include <linux/futex.h>
27 +
28 +/* ------------------------------------------------------------------ */
29 +/* Internal helpers */
30 +/* ------------------------------------------------------------------ */
31 +
32 +/* Round up to 64-byte alignment. */
33 +static inline uint32_t align64(uint32_t v)
34 +{
35 + return (v + (NIPC_SHM_REGION_ALIGNMENT - 1)) & ~(uint32_t)(NIPC_SHM_REGION_ALIGNMENT - 1);
36 +}
37 +
38 +/* Validate service_name: only [a-zA-Z0-9._-], non-empty, not "." or "..". */
39 +static int validate_service_name(const char *name)
40 +{
41 + if (!name || name[0] == '\0')
42 + return -1;
43 +
44 + /* Reject "." and ".." */
45 + if (name[0] == '.' && (name[1] == '\0' || (name[1] == '.' && name[2] == '\0')))
46 + return -1;
47 +
48 + for (const char *p = name; *p; p++) {
49 + char c = *p;
50 + if ((c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') ||
51 + (c >= '0' && c <= '9') || c == '.' || c == '_' || c == '-')
52 + continue;
53 + return -1;
54 + }
55 + return 0;
56 +}
57 +
58 +/* Build per-session SHM file path: {run_dir}/{service_name}-{session_id:016x}.ipcshm */
59 +static int build_shm_path(char *dst, size_t dst_len,
60 + const char *run_dir, const char *service_name,
61 + uint64_t session_id)
62 +{
63 + if (validate_service_name(service_name) < 0)
64 + return -2; /* invalid service name */
65 +
66 + int n = snprintf(dst, dst_len, "%s/%s-%016" PRIx64 ".ipcshm",
67 + run_dir, service_name, session_id);
68 + if (n < 0 || (size_t)n >= dst_len)
69 + return -1;
70 + return 0;
71 +}
72 +
73 +/* Thin wrapper around the futex syscall. */
74 +static int futex_wake(uint32_t *addr, int count)
75 +{
76 + return (int)syscall(SYS_futex, addr, FUTEX_WAKE, count, NULL, NULL, 0);
77 +}
78 +
79 +static int futex_wait(uint32_t *addr, uint32_t expected,
80 + const struct timespec *timeout)
81 +{
82 + return (int)syscall(SYS_futex, addr, FUTEX_WAIT, expected, timeout, NULL, 0);
83 +}
84 +
85 +/* CPU pause hint for spin loops. */
86 +static inline void cpu_relax(void)
87 +{
88 +#if defined(__x86_64__) || defined(__i386__)
89 + __builtin_ia32_pause();
90 +#elif defined(__aarch64__)
91 + __asm__ volatile("yield" ::: "memory");
92 +#else
93 + /* generic: compiler barrier */
94 + __atomic_signal_fence(__ATOMIC_SEQ_CST);
95 +#endif
96 +}
97 +
98 +/* Check if a PID is alive (without sending a signal). */
99 +static bool pid_alive(pid_t pid)
100 +{
101 + if (pid <= 0)
102 + return false;
103 + return kill(pid, 0) == 0 || errno == EPERM;
104 +}
105 +
106 +/* Pointer into the mapped region at a byte offset. */
107 +static inline void *region_ptr(const nipc_shm_ctx_t *ctx, uint32_t offset)
108 +{
109 + return (uint8_t *)ctx->base + offset;
110 +}
111 +
112 +/* Pointer to the header (always at offset 0). Used for non-atomic fields. */
113 +static inline nipc_shm_region_header_t *region_hdr(const nipc_shm_ctx_t *ctx)
114 +{
115 + return (nipc_shm_region_header_t *)ctx->base;
116 +}
117 +
118 +/*
119 + * Byte-offset accessors for atomic fields. These avoid taking the
120 + * address of a packed struct member, which GCC warns about.
121 + */
122 +#define SHM_OFF_REQ_SEQ 32
123 +#define SHM_OFF_RESP_SEQ 40
124 +#define SHM_OFF_REQ_LEN 48
125 +#define SHM_OFF_RESP_LEN 52
126 +#define SHM_OFF_REQ_SIGNAL 56
127 +#define SHM_OFF_RESP_SIGNAL 60
128 +
129 +static inline uint64_t *shm_seq_ptr(void *base, int offset)
130 +{
131 + return (uint64_t *)((uint8_t *)base + offset);
132 +}
133 +
134 +static inline uint32_t *shm_u32_ptr(void *base, int offset)
135 +{
136 + return (uint32_t *)((uint8_t *)base + offset);
137 +}
138 +
139 +/* ------------------------------------------------------------------ */
140 +/* Stale region recovery */
141 +/* ------------------------------------------------------------------ */
142 +
143 +/*
144 + * Returns:
145 + * 0 = stale (unlinked)
146 + * +1 = live server
147 + * -1 = doesn't exist
148 + * -2 = exists but undersized / invalid (treated as stale, unlinked)
149 + */
150 +static int check_shm_stale(const char *path)
151 +{
152 + struct stat st;
153 + if (stat(path, &st) != 0)
154 + return -1;
155 +
156 + /* Must be at least header-sized to inspect. */
157 + if ((size_t)st.st_size < NIPC_SHM_HEADER_LEN) {
158 + unlink(path);
159 + return -2;
160 + }
161 +
162 + int fd = open(path, O_RDONLY);
163 + if (fd < 0) {
164 + /* Don't unlink on permission errors — the file may be owned by
165 + * another user/process and we just can't read it. */
166 + if (errno != EACCES && errno != EPERM)
167 + unlink(path);
168 + return -2;
169 + }
170 +
171 + void *map = mmap(NULL, NIPC_SHM_HEADER_LEN, PROT_READ, MAP_SHARED, fd, 0);
172 + close(fd);
173 + if (map == MAP_FAILED) {
174 + unlink(path);
175 + return -2;
176 + }
177 +
178 + const nipc_shm_region_header_t *hdr = (const nipc_shm_region_header_t *)map;
179 +
180 + /* Validate magic first. */
181 + if (hdr->magic != NIPC_SHM_REGION_MAGIC) {
182 + munmap(map, NIPC_SHM_HEADER_LEN);
183 + unlink(path);
184 + return -2;
185 + }
186 +
187 + int32_t owner = hdr->owner_pid;
188 + uint32_t gen = hdr->owner_generation;
189 + munmap(map, NIPC_SHM_HEADER_LEN);
190 +
191 + if (pid_alive((pid_t)owner) && gen != 0) {
192 + return 1; /* live: PID is alive and generation is valid */
193 + }
194 +
195 + /* Dead owner or zero generation (uninitialized/legacy) — stale */
196 + unlink(path);
197 + return 0;
198 +}
199 +
200 +/* ------------------------------------------------------------------ */
201 +/* Server: create */
202 +/* ------------------------------------------------------------------ */
203 +
204 +nipc_shm_error_t nipc_shm_server_create(const char *run_dir,
205 + const char *service_name,
206 + uint64_t session_id,
207 + uint32_t req_capacity,
208 + uint32_t resp_capacity,
209 + nipc_shm_ctx_t *out)
210 +{
211 + if (!run_dir || !service_name || !out)
212 + return NIPC_SHM_ERR_BAD_PARAM;
213 +
214 + memset(out, 0, sizeof(*out));
215 + out->fd = -1;
216 +
217 + /* Build per-session path (validates service_name) */
218 + char path[256];
219 + int path_rc = build_shm_path(path, sizeof(path), run_dir, service_name,
220 + session_id);
221 + if (path_rc == -2)
222 + return NIPC_SHM_ERR_BAD_PARAM;
223 + if (path_rc < 0)
224 + return NIPC_SHM_ERR_PATH_TOO_LONG;
225 +
226 + /* Round capacities up to alignment. */
227 + req_capacity = align64(req_capacity);
228 + resp_capacity = align64(resp_capacity);
229 +
230 + uint32_t req_off = align64(NIPC_SHM_HEADER_LEN);
231 + /* Guard against uint32 overflow before computing resp_off */
232 + if (req_capacity > UINT32_MAX - req_off)
233 + return NIPC_SHM_ERR_BAD_PARAM;
234 + uint32_t resp_off = align64(req_off + req_capacity);
235 + size_t region_size = (size_t)resp_off + resp_capacity;
236 +
237 + /* Try O_EXCL create first (fast path, no stale check needed). */
238 + int fd = open(path, O_RDWR | O_CREAT | O_EXCL, 0600);
239 +
240 + /* If O_EXCL failed because file exists, do stale recovery and retry. */
241 + if (fd < 0 && errno == EEXIST) {
242 + int stale = check_shm_stale(path);
243 + if (stale == 1)
244 + return NIPC_SHM_ERR_ADDR_IN_USE;
245 + /* Stale file was unlinked, retry create */
246 + fd = open(path, O_RDWR | O_CREAT | O_EXCL, 0600);
247 + }
248 + if (fd < 0)
249 + return NIPC_SHM_ERR_OPEN;
250 +
251 + if (ftruncate(fd, (off_t)region_size) < 0) {
252 + close(fd);
253 + unlink(path);
254 + return NIPC_SHM_ERR_TRUNCATE;
255 + }
256 +
257 + void *map = mmap(NULL, region_size, PROT_READ | PROT_WRITE,
258 + MAP_SHARED, fd, 0);
259 + if (map == MAP_FAILED) {
260 + close(fd);
261 + unlink(path);
262 + return NIPC_SHM_ERR_MMAP;
263 + }
264 +
265 + /* Zero the region first to init all atomics to 0. */
266 + memset(map, 0, region_size);
267 +
268 + /* Write header. */
269 + nipc_shm_region_header_t *hdr = (nipc_shm_region_header_t *)map;
270 + hdr->magic = NIPC_SHM_REGION_MAGIC;
271 + hdr->version = NIPC_SHM_REGION_VERSION;
272 + hdr->header_len = NIPC_SHM_HEADER_LEN;
273 + hdr->owner_pid = (int32_t)getpid();
274 +
275 + /* Use a time-based generation to detect PID reuse across restarts. */
276 + {
277 + struct timespec ts;
278 + clock_gettime(CLOCK_MONOTONIC, &ts);
279 + hdr->owner_generation = (uint32_t)(ts.tv_sec ^ (ts.tv_nsec >> 10));
280 + }
281 + hdr->request_offset = req_off;
282 + hdr->request_capacity = req_capacity;
283 + hdr->response_offset = resp_off;
284 + hdr->response_capacity = resp_capacity;
285 +
286 + /* Ensure header writes are visible before clients read. */
287 + __atomic_thread_fence(__ATOMIC_RELEASE);
288 +
289 + /* Fill context. */
290 + out->role = NIPC_SHM_ROLE_SERVER;
291 + out->fd = fd;
292 + out->base = map;
293 + out->region_size = region_size;
294 + out->request_offset = req_off;
295 + out->request_capacity = req_capacity;
296 + out->response_offset = resp_off;
297 + out->response_capacity = resp_capacity;
298 + out->local_req_seq = 0;
299 + out->local_resp_seq = 0;
300 + out->spin_tries = NIPC_SHM_DEFAULT_SPIN;
301 + out->owner_generation = hdr->owner_generation;
302 + strncpy(out->path, path, sizeof(out->path) - 1);
303 + out->path[sizeof(out->path) - 1] = '\0';
304 +
305 + return NIPC_SHM_OK;
306 +}
307 +
308 +/* ------------------------------------------------------------------ */
309 +/* Server: destroy */
310 +/* ------------------------------------------------------------------ */
311 +
312 +void nipc_shm_destroy(nipc_shm_ctx_t *ctx)
313 +{
314 + if (!ctx)
315 + return;
316 +
317 + if (ctx->base && ctx->base != MAP_FAILED) {
318 + munmap(ctx->base, ctx->region_size);
319 + ctx->base = NULL;
320 + }
321 +
322 + if (ctx->fd >= 0) {
323 + close(ctx->fd);
324 + ctx->fd = -1;
325 + }
326 +
327 + if (ctx->path[0]) {
328 + unlink(ctx->path);
329 + ctx->path[0] = '\0';
330 + }
331 +
332 + ctx->region_size = 0;
333 +}
334 +
335 +/* ------------------------------------------------------------------ */
336 +/* Client: attach */
337 +/* ------------------------------------------------------------------ */
338 +
339 +nipc_shm_error_t nipc_shm_client_attach(const char *run_dir,
340 + const char *service_name,
341 + uint64_t session_id,
342 + nipc_shm_ctx_t *out)
343 +{
344 + if (!run_dir || !service_name || !out)
345 + return NIPC_SHM_ERR_BAD_PARAM;
346 +
347 + memset(out, 0, sizeof(*out));
348 + out->fd = -1;
349 +
350 + char path[256];
351 + int path_rc = build_shm_path(path, sizeof(path), run_dir, service_name,
352 + session_id);
353 + if (path_rc == -2)
354 + return NIPC_SHM_ERR_BAD_PARAM;
355 + if (path_rc < 0)
356 + return NIPC_SHM_ERR_PATH_TOO_LONG;
357 +
358 + /* Open the file. */
359 + int fd = open(path, O_RDWR);
360 + if (fd < 0)
361 + return NIPC_SHM_ERR_OPEN;
362 +
363 + /* Check file size. */
364 + struct stat st;
365 + if (fstat(fd, &st) < 0) {
366 + close(fd);
367 + return NIPC_SHM_ERR_OPEN;
368 + }
369 +
370 + if ((size_t)st.st_size < NIPC_SHM_HEADER_LEN) {
371 + close(fd);
372 + return NIPC_SHM_ERR_NOT_READY;
373 + }
374 +
375 + /* Map the region. */
376 + size_t file_size = (size_t)st.st_size;
377 + void *map = mmap(NULL, file_size, PROT_READ | PROT_WRITE,
378 + MAP_SHARED, fd, 0);
379 + if (map == MAP_FAILED) {
380 + close(fd);
381 + return NIPC_SHM_ERR_MMAP;
382 + }
383 +
384 + /* Acquire fence so we see the server's header writes. */
385 + __atomic_thread_fence(__ATOMIC_ACQUIRE);
386 +
387 + const nipc_shm_region_header_t *hdr =
388 + (const nipc_shm_region_header_t *)map;
389 +
390 + /* Validate header. */
391 + if (hdr->magic != NIPC_SHM_REGION_MAGIC) {
392 + munmap(map, file_size);
393 + close(fd);
394 + return NIPC_SHM_ERR_BAD_MAGIC;
395 + }
396 +
397 + if (hdr->version != NIPC_SHM_REGION_VERSION) {
398 + munmap(map, file_size);
399 + close(fd);
400 + return NIPC_SHM_ERR_BAD_VERSION;
401 + }
402 +
403 + if (hdr->header_len != NIPC_SHM_HEADER_LEN) {
404 + munmap(map, file_size);
405 + close(fd);
406 + return NIPC_SHM_ERR_BAD_HEADER;
407 + }
408 +
409 + size_t header_end = align64((uint32_t)NIPC_SHM_HEADER_LEN);
410 + if (hdr->request_offset < header_end ||
411 + hdr->request_capacity == 0 ||
412 + hdr->response_offset < header_end ||
413 + hdr->response_capacity == 0) {
414 + munmap(map, file_size);
415 + close(fd);
416 + return NIPC_SHM_ERR_NOT_READY;
417 + }
418 +
419 + /* Guard against uint32 overflow in request_offset + request_capacity */
420 + if (hdr->request_offset > UINT32_MAX - hdr->request_capacity) {
421 + munmap(map, file_size);
422 + close(fd);
423 + return NIPC_SHM_ERR_BAD_SIZE;
424 + }
425 +
426 + if ((hdr->request_offset % NIPC_SHM_REGION_ALIGNMENT) != 0 ||
427 + (hdr->request_capacity % NIPC_SHM_REGION_ALIGNMENT) != 0 ||
428 + (hdr->response_offset % NIPC_SHM_REGION_ALIGNMENT) != 0 ||
429 + (hdr->response_capacity % NIPC_SHM_REGION_ALIGNMENT) != 0 ||
430 + hdr->response_offset < align64(hdr->request_offset + hdr->request_capacity)) {
431 + munmap(map, file_size);
432 + close(fd);
433 + return NIPC_SHM_ERR_BAD_SIZE;
434 + }
435 +
436 + /* Validate region is large enough for the declared areas. */
437 + size_t needed = 0;
438 + size_t req_end = (size_t)hdr->request_offset + hdr->request_capacity;
439 + size_t resp_end = (size_t)hdr->response_offset + hdr->response_capacity;
440 + needed = req_end > resp_end ? req_end : resp_end;
441 +
442 + if (file_size < needed) {
443 + munmap(map, file_size);
444 + close(fd);
445 + return NIPC_SHM_ERR_BAD_SIZE;
446 + }
447 +
448 + /* Read the current sequence numbers so we don't see stale data. */
449 + uint64_t cur_req_seq = __atomic_load_n(
450 + (uint64_t *)((uint8_t *)map + offsetof(nipc_shm_region_header_t, req_seq)),
451 + __ATOMIC_ACQUIRE);
452 + uint64_t cur_resp_seq = __atomic_load_n(
453 + (uint64_t *)((uint8_t *)map + offsetof(nipc_shm_region_header_t, resp_seq)),
454 + __ATOMIC_ACQUIRE);
455 +
456 + /* Fill context. */
457 + out->role = NIPC_SHM_ROLE_CLIENT;
458 + out->fd = fd;
459 + out->base = map;
460 + out->region_size = file_size;
461 + out->request_offset = hdr->request_offset;
462 + out->request_capacity = hdr->request_capacity;
463 + out->response_offset = hdr->response_offset;
464 + out->response_capacity = hdr->response_capacity;
465 + out->local_req_seq = cur_req_seq;
466 + out->local_resp_seq = cur_resp_seq;
467 + out->spin_tries = NIPC_SHM_DEFAULT_SPIN;
468 + out->owner_generation = hdr->owner_generation;
469 + strncpy(out->path, path, sizeof(out->path) - 1);
470 + out->path[sizeof(out->path) - 1] = '\0';
471 +
472 + return NIPC_SHM_OK;
473 +}
474 +
475 +/* ------------------------------------------------------------------ */
476 +/* Client: close (no unlink) */
477 +/* ------------------------------------------------------------------ */
478 +
479 +void nipc_shm_close(nipc_shm_ctx_t *ctx)
480 +{
481 + if (!ctx)
482 + return;
483 +
484 + if (ctx->base && ctx->base != MAP_FAILED) {
485 + munmap(ctx->base, ctx->region_size);
486 + ctx->base = NULL;
487 + }
488 +
489 + if (ctx->fd >= 0) {
490 + close(ctx->fd);
491 + ctx->fd = -1;
492 + }
493 +
494 + ctx->region_size = 0;
495 +}
496 +
497 +/* ------------------------------------------------------------------ */
498 +/* Data plane: send */
499 +/* ------------------------------------------------------------------ */
500 +
501 +nipc_shm_error_t nipc_shm_send(nipc_shm_ctx_t *ctx,
502 + const void *msg, size_t msg_len)
503 +{
504 + if (!ctx || !ctx->base || !msg || msg_len == 0)
505 + return NIPC_SHM_ERR_BAD_PARAM;
506 +
507 + /*
508 + * Client writes to the request area; server writes to the response
509 + * area. The direction is determined by role.
510 + */
511 + uint32_t area_offset;
512 + uint32_t area_capacity;
513 + int seq_off, len_off, sig_off;
514 +
515 + if (ctx->role == NIPC_SHM_ROLE_CLIENT) {
516 + area_offset = ctx->request_offset;
517 + area_capacity = ctx->request_capacity;
518 + seq_off = SHM_OFF_REQ_SEQ;
519 + len_off = SHM_OFF_REQ_LEN;
520 + sig_off = SHM_OFF_REQ_SIGNAL;
521 + } else {
522 + area_offset = ctx->response_offset;
523 + area_capacity = ctx->response_capacity;
524 + seq_off = SHM_OFF_RESP_SEQ;
525 + len_off = SHM_OFF_RESP_LEN;
526 + sig_off = SHM_OFF_RESP_SIGNAL;
527 + }
528 +
529 + if (msg_len > area_capacity)
530 + return NIPC_SHM_ERR_MSG_TOO_LARGE;
531 +
532 + uint64_t *seq_ptr = shm_seq_ptr(ctx->base, seq_off);
533 + uint32_t *len_ptr = shm_u32_ptr(ctx->base, len_off);
534 + uint32_t *signal_ptr = shm_u32_ptr(ctx->base, sig_off);
535 +
536 + /* 1. Write message data into the area. */
537 + memcpy(region_ptr(ctx, area_offset), msg, msg_len);
538 +
539 + /* 2. Store message length (release). */
540 + __atomic_store_n(len_ptr, (uint32_t)msg_len, __ATOMIC_RELEASE);
541 +
542 + /* 3. Increment sequence number (release) to publish. */
543 + __atomic_add_fetch(seq_ptr, 1, __ATOMIC_RELEASE);
544 +
545 + /* 4. Wake the peer via futex. */
546 + __atomic_add_fetch(signal_ptr, 1, __ATOMIC_RELEASE);
547 + futex_wake(signal_ptr, 1);
548 +
549 + /* Track locally. */
550 + if (ctx->role == NIPC_SHM_ROLE_CLIENT)
551 + ctx->local_req_seq++;
552 + else
553 + ctx->local_resp_seq++;
554 +
555 + return NIPC_SHM_OK;
556 +}
557 +
558 +/* ------------------------------------------------------------------ */
559 +/* Data plane: receive */
560 +/* ------------------------------------------------------------------ */
561 +
562 +nipc_shm_error_t nipc_shm_receive(nipc_shm_ctx_t *ctx,
563 + void *buf,
564 + size_t buf_size,
565 + size_t *msg_len_out,
566 + uint32_t timeout_ms)
567 +{
568 + if (!ctx || !ctx->base || !buf || !msg_len_out || buf_size == 0)
569 + return NIPC_SHM_ERR_BAD_PARAM;
570 +
571 + /*
572 + * Server reads from the request area; client reads from the
573 + * response area.
574 + */
575 + uint32_t area_offset;
576 + uint32_t area_capacity;
577 + int seq_off, len_off, sig_off;
578 + uint64_t expected_seq;
579 +
580 + if (ctx->role == NIPC_SHM_ROLE_SERVER) {
581 + area_offset = ctx->request_offset;
582 + area_capacity = ctx->request_capacity;
583 + seq_off = SHM_OFF_REQ_SEQ;
584 + len_off = SHM_OFF_REQ_LEN;
585 + sig_off = SHM_OFF_REQ_SIGNAL;
586 + expected_seq = ctx->local_req_seq + 1;
587 + } else {
588 + area_offset = ctx->response_offset;
589 + area_capacity = ctx->response_capacity;
590 + seq_off = SHM_OFF_RESP_SEQ;
591 + len_off = SHM_OFF_RESP_LEN;
592 + sig_off = SHM_OFF_RESP_SIGNAL;
593 + expected_seq = ctx->local_resp_seq + 1;
594 + }
595 +
596 + /* The copy ceiling is the smaller of the caller buffer and the
597 + * SHM area capacity. This prevents out-of-bounds reads even if
598 + * the peer writes a forged length value. */
599 + uint32_t max_copy = (buf_size < area_capacity) ? (uint32_t)buf_size : area_capacity;
600 +
601 + uint64_t *seq_ptr = shm_seq_ptr(ctx->base, seq_off);
602 + uint32_t *len_ptr = shm_u32_ptr(ctx->base, len_off);
603 + uint32_t *signal_ptr = shm_u32_ptr(ctx->base, sig_off);
604 +
605 + void *data_ptr = region_ptr(ctx, area_offset);
606 +
607 + /*
608 + * Phase 1: spin. On detecting the sequence advance, immediately
609 + * copy the message into the caller's buffer while still in the
610 + * same iteration. The shared area can be overwritten by the peer
611 + * within nanoseconds of the sequence advancing.
612 + */
613 + bool observed = false;
614 + uint32_t mlen = 0;
615 + for (uint32_t i = 0; i < ctx->spin_tries; i++) {
616 + uint64_t cur = __atomic_load_n(seq_ptr, __ATOMIC_ACQUIRE);
617 + if (cur >= expected_seq) {
618 + mlen = __atomic_load_n(len_ptr, __ATOMIC_ACQUIRE);
619 + if (mlen > 0 && mlen <= max_copy)
620 + memcpy(buf, data_ptr, mlen);
621 + observed = true;
622 + break;
623 + }
624 + cpu_relax();
625 + }
626 +
627 + /* Phase 2: futex wait (if spinning didn't observe the advance).
628 + *
629 + * The futex_wait loop handles spurious wakeups (EAGAIN when the
630 + * signal word changed between our read and the syscall, or EINTR
631 + * from signal delivery). We compute a wall-clock deadline so the
632 + * total wait never exceeds timeout_ms regardless of retries. */
633 + if (!observed) {
634 + uint64_t deadline_ns = 0; /* 0 = no timeout */
635 + if (timeout_ms > 0) {
636 + struct timespec now_ts;
637 + clock_gettime(CLOCK_MONOTONIC, &now_ts);
638 + deadline_ns = (uint64_t)now_ts.tv_sec * 1000000000ull
639 + + (uint64_t)now_ts.tv_nsec
640 + + (uint64_t)timeout_ms * 1000000ull;
641 + }
642 +
643 + for (;;) {
644 + uint32_t sig_val = __atomic_load_n(signal_ptr, __ATOMIC_ACQUIRE);
645 +
646 + uint64_t cur = __atomic_load_n(seq_ptr, __ATOMIC_ACQUIRE);
647 + if (cur >= expected_seq)
648 + break; /* response arrived */
649 +
650 + /* Compute remaining timeout for this futex_wait call. */
651 + struct timespec ts;
652 + struct timespec *tsp = NULL;
653 + if (deadline_ns > 0) {
654 + struct timespec now_ts;
655 + clock_gettime(CLOCK_MONOTONIC, &now_ts);
656 + uint64_t now_val = (uint64_t)now_ts.tv_sec * 1000000000ull
657 + + (uint64_t)now_ts.tv_nsec;
658 + if (now_val >= deadline_ns)
659 + return NIPC_SHM_ERR_TIMEOUT;
660 +
661 + uint64_t remain = deadline_ns - now_val;
662 + ts.tv_sec = (time_t)(remain / 1000000000ull);
663 + ts.tv_nsec = (long)(remain % 1000000000ull);
664 + tsp = &ts;
665 + }
666 +
667 + int ret = futex_wait(signal_ptr, sig_val, tsp);
668 + if (ret < 0 && errno == ETIMEDOUT)
669 + return NIPC_SHM_ERR_TIMEOUT;
670 +
671 + /* EAGAIN (value changed) or EINTR (signal): re-check seq. */
672 + }
673 +
674 + /* Copy immediately after observing the sequence advance. */
675 + mlen = __atomic_load_n(len_ptr, __ATOMIC_ACQUIRE);
676 + if (mlen > 0 && mlen <= max_copy)
677 + memcpy(buf, data_ptr, mlen);
678 + }
679 +
680 + /* Message larger than caller buffer or area capacity */
681 + if (mlen > max_copy) {
682 + *msg_len_out = mlen;
683 + /* Still advance tracking -- message is consumed from SHM perspective */
684 + if (ctx->role == NIPC_SHM_ROLE_SERVER)
685 + ctx->local_req_seq = expected_seq;
686 + else
687 + ctx->local_resp_seq = expected_seq;
688 + return NIPC_SHM_ERR_MSG_TOO_LARGE;
689 + }
690 +
691 + /* mlen==0 after a sequence advance indicates corruption (send rejects 0-length) */
692 + if (mlen == 0) {
693 + if (ctx->role == NIPC_SHM_ROLE_SERVER)
694 + ctx->local_req_seq = expected_seq;
695 + else
696 + ctx->local_resp_seq = expected_seq;
697 + *msg_len_out = 0;
698 + return NIPC_SHM_ERR_BAD_HEADER;
699 + }
700 +
701 + *msg_len_out = mlen;
702 +
703 + /* Advance local tracking. */
704 + if (ctx->role == NIPC_SHM_ROLE_SERVER)
705 + ctx->local_req_seq = expected_seq;
706 + else
707 + ctx->local_resp_seq = expected_seq;
708 +
709 + return NIPC_SHM_OK;
710 +}
711 +
712 +/* ------------------------------------------------------------------ */
713 +/* Utility */
714 +/* ------------------------------------------------------------------ */
715 +
716 +bool nipc_shm_owner_alive(const nipc_shm_ctx_t *ctx)
717 +{
718 + if (!ctx || !ctx->base)
719 + return false;
720 +
721 + const nipc_shm_region_header_t *hdr =
722 + (const nipc_shm_region_header_t *)ctx->base;
723 +
724 + if (!pid_alive((pid_t)hdr->owner_pid))
725 + return false;
726 +
727 + /* PID is alive; verify generation matches to detect PID reuse.
728 + * If owner_generation is 0, skip the check (legacy region). */
729 + if (ctx->owner_generation != 0 &&
730 + hdr->owner_generation != ctx->owner_generation)
731 + return false;
732 +
733 + return true;
734 +}
735 +
736 +/* ------------------------------------------------------------------ */
737 +/* Stale cleanup on server startup */
738 +/* ------------------------------------------------------------------ */
739 +
740 +void nipc_shm_cleanup_stale(const char *run_dir, const char *service_name)
741 +{
742 + if (!run_dir || !service_name)
743 + return;
744 +
745 + if (validate_service_name(service_name) < 0)
746 + return;
747 +
748 + /* Build the prefix to match: "{service_name}-" */
749 + char prefix[256];
750 + int pn = snprintf(prefix, sizeof(prefix), "%s-", service_name);
751 + if (pn < 0 || (size_t)pn >= sizeof(prefix))
752 + return;
753 +
754 + size_t prefix_len = (size_t)pn;
755 + const char *suffix = ".ipcshm";
756 + size_t suffix_len = strlen(suffix);
757 +
758 + DIR *dir = opendir(run_dir);
759 + if (!dir)
760 + return;
761 +
762 + struct dirent *ent;
763 + while ((ent = readdir(dir)) != NULL) {
764 + size_t nlen = strlen(ent->d_name);
765 +
766 + /* Must start with "{service_name}-" and end with ".ipcshm" */
767 + if (nlen <= prefix_len + suffix_len)
768 + continue;
769 + if (strncmp(ent->d_name, prefix, prefix_len) != 0)
770 + continue;
771 + if (strcmp(ent->d_name + nlen - suffix_len, suffix) != 0)
772 + continue;
773 +
774 + /* Build full path and check if stale */
775 + char path[512];
776 + int n = snprintf(path, sizeof(path), "%s/%s", run_dir, ent->d_name);
777 + if (n < 0 || (size_t)n >= sizeof(path))
778 + continue;
779 +
780 + /* check_shm_stale unlinks stale files and returns:
781 + * 0 = stale (unlinked), +1 = live, -1 = gone, -2 = invalid (unlinked) */
782 + check_shm_stale(path);
783 + }
784 +
785 + closedir(dir);
786 +}
src/libnetdata/netipc/src/transport/posix/netipc_uds.c new
+1083
@@ -0,0 +1,1083 @@
1 +/*
2 + * netipc_uds.c - L1 POSIX UDS SEQPACKET transport.
3 + *
4 + * Implements connection lifecycle, handshake with profile/limit negotiation,
5 + * and send/receive with transparent chunking over AF_UNIX SEQPACKET sockets.
6 + */
7 +
8 +#include "netipc/netipc_uds.h"
9 +#include "netipc/netipc_protocol.h"
10 +
11 +#include <errno.h>
12 +#include <fcntl.h>
13 +#include <stdlib.h>
14 +#include <stdio.h>
15 +#include <string.h>
16 +#include <unistd.h>
17 +
18 +#include <sys/socket.h>
19 +#include <sys/stat.h>
20 +#include <sys/un.h>
21 +
22 +/* ------------------------------------------------------------------ */
23 +/* Internal constants */
24 +/* ------------------------------------------------------------------ */
25 +
26 +#define UDS_DEFAULT_BACKLOG 16
27 +#define UDS_DEFAULT_BATCH_ITEMS 1
28 +#define UDS_INITIAL_RECV_BUF 4096
29 +
30 +/* ------------------------------------------------------------------ */
31 +/* Internal helpers */
32 +/* ------------------------------------------------------------------ */
33 +
34 +/* Validate service_name: only [a-zA-Z0-9._-], non-empty, not "." or "..". */
35 +static int validate_service_name(const char *name)
36 +{
37 + if (!name || name[0] == '\0')
38 + return -1;
39 +
40 + /* Reject "." and ".." */
41 + if (name[0] == '.' && (name[1] == '\0' || (name[1] == '.' && name[2] == '\0')))
42 + return -1;
43 +
44 + for (const char *p = name; *p; p++) {
45 + char c = *p;
46 + if ((c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') ||
47 + (c >= '0' && c <= '9') || c == '.' || c == '_' || c == '-')
48 + continue;
49 + return -1;
50 + }
51 + return 0;
52 +}
53 +
54 +/* Build socket path into dst, return 0 on success, -1 if too long. */
55 +static int build_socket_path(char *dst, size_t dst_len,
56 + const char *run_dir, const char *service_name)
57 +{
58 + if (validate_service_name(service_name) < 0)
59 + return -2; /* invalid service name */
60 +
61 + int n = snprintf(dst, dst_len, "%s/%s.sock", run_dir, service_name);
62 + if (n < 0 || (size_t)n >= dst_len)
63 + return -1; /* path too long */
64 + return 0;
65 +}
66 +
67 +/* Get the socket's send buffer size as the packet size. */
68 +static uint32_t detect_packet_size(int fd)
69 +{
70 + int val = 0;
71 + socklen_t len = sizeof(val);
72 + if (getsockopt(fd, SOL_SOCKET, SO_SNDBUF, &val, &len) < 0)
73 + return 65536; /* safe default */
74 +
75 + /* Linux doubles SO_SNDBUF internally; use the value as-is, it
76 + * represents the actual kernel buffer available. Clamp to u32. */
77 + if (val <= 0)
78 + return 65536;
79 + return (uint32_t)val;
80 +}
81 +
82 +/* Highest set bit in a bitmask (0 if empty). */
83 +static uint32_t highest_bit(uint32_t mask)
84 +{
85 + if (mask == 0)
86 + return 0;
87 +
88 + uint32_t bit = 1u << 31;
89 + while (!(mask & bit))
90 + bit >>= 1;
91 + return bit;
92 +}
93 +
94 +static inline uint32_t min_u32(uint32_t a, uint32_t b)
95 +{
96 + return a < b ? a : b;
97 +}
98 +
99 +static inline uint32_t max_u32(uint32_t a, uint32_t b)
100 +{
101 + return a > b ? a : b;
102 +}
103 +
104 +static inline uint32_t apply_default(uint32_t val, uint32_t def)
105 +{
106 + return val == 0 ? def : val;
107 +}
108 +
109 +static nipc_uds_error_t raw_send(int fd, const void *data, size_t len);
110 +
111 +static bool header_version_incompatible(const void *buf, size_t buf_len,
112 + uint16_t expected_code)
113 +{
114 + if (buf_len < NIPC_HEADER_LEN)
115 + return false;
116 +
117 + nipc_header_t hdr;
118 + memcpy(&hdr, buf, sizeof(hdr));
119 + return hdr.magic == NIPC_MAGIC_MSG &&
120 + hdr.version != NIPC_VERSION &&
121 + hdr.header_len == NIPC_HEADER_LEN &&
122 + hdr.kind == NIPC_KIND_CONTROL &&
123 + hdr.code == expected_code;
124 +}
125 +
126 +static bool hello_layout_incompatible(const void *buf, size_t buf_len)
127 +{
128 + if (buf_len < sizeof(uint16_t))
129 + return false;
130 +
131 + nipc_hello_t hello;
132 + memset(&hello, 0, sizeof(hello));
133 + memcpy(&hello, buf, sizeof(uint16_t));
134 + return hello.layout_version != 1;
135 +}
136 +
137 +static bool hello_ack_layout_incompatible(const void *buf, size_t buf_len)
138 +{
139 + if (buf_len < sizeof(uint16_t))
140 + return false;
141 +
142 + nipc_hello_ack_t ack;
143 + memset(&ack, 0, sizeof(ack));
144 + memcpy(&ack, buf, sizeof(uint16_t));
145 + return ack.layout_version != 1;
146 +}
147 +
148 +static void send_rejection_ack(int fd, uint16_t status)
149 +{
150 + nipc_hello_ack_t ack = { .layout_version = 1 };
151 + uint8_t ack_buf[48];
152 + uint8_t pkt[80];
153 + nipc_header_t ack_hdr = {
154 + .magic = NIPC_MAGIC_MSG,
155 + .version = NIPC_VERSION,
156 + .header_len = NIPC_HEADER_LEN,
157 + .kind = NIPC_KIND_CONTROL,
158 + .code = NIPC_CODE_HELLO_ACK,
159 + .transport_status = status,
160 + .payload_len = sizeof(ack_buf),
161 + .item_count = 1,
162 + };
163 +
164 + nipc_hello_ack_encode(&ack, ack_buf, sizeof(ack_buf));
165 + nipc_header_encode(&ack_hdr, pkt, sizeof(pkt));
166 + memcpy(pkt + NIPC_HEADER_LEN, ack_buf, sizeof(ack_buf));
167 + raw_send(fd, pkt, NIPC_HEADER_LEN + sizeof(ack_buf));
168 +}
169 +
170 +/* ------------------------------------------------------------------ */
171 +/* Low-level send/recv (one SEQPACKET datagram) */
172 +/* ------------------------------------------------------------------ */
173 +
174 +/* Send exactly len bytes as one SEQPACKET message. */
175 +static nipc_uds_error_t raw_send(int fd, const void *data, size_t len)
176 +{
177 + ssize_t n = send(fd, data, len, MSG_NOSIGNAL);
178 + if (n < 0 || (size_t)n != len)
179 + return NIPC_UDS_ERR_SEND;
180 + return NIPC_UDS_OK;
181 +}
182 +
183 +/* Send header + payload as one SEQPACKET message using sendmsg. */
184 +static nipc_uds_error_t raw_send_iov(int fd, const void *hdr, size_t hdr_len,
185 + const void *payload, size_t payload_len)
186 +{
187 + struct iovec iov[2];
188 + struct msghdr msg;
189 + int iovcnt = 0;
190 +
191 + memset(&msg, 0, sizeof(msg));
192 +
193 + iov[0].iov_base = (void *)hdr;
194 + iov[0].iov_len = hdr_len;
195 + iovcnt = 1;
196 +
197 + if (payload && payload_len > 0) {
198 + iov[1].iov_base = (void *)payload;
199 + iov[1].iov_len = payload_len;
200 + iovcnt = 2;
201 + }
202 +
203 + msg.msg_iov = iov;
204 + msg.msg_iovlen = iovcnt;
205 +
206 + size_t total = hdr_len + payload_len;
207 + ssize_t n = sendmsg(fd, &msg, MSG_NOSIGNAL);
208 + if (n < 0 || (size_t)n != total)
209 + return NIPC_UDS_ERR_SEND;
210 +
211 + return NIPC_UDS_OK;
212 +}
213 +
214 +/* Receive one SEQPACKET message into buf. Returns bytes received, 0 on
215 + * disconnect, -1 on error. */
216 +static ssize_t raw_recv(int fd, void *buf, size_t buf_len)
217 +{
218 + ssize_t n = recv(fd, buf, buf_len, 0);
219 + return n;
220 +}
221 +
222 +/* ------------------------------------------------------------------ */
223 +/* Handshake: client side */
224 +/* ------------------------------------------------------------------ */
225 +
226 +static nipc_uds_error_t client_handshake(int fd,
227 + const nipc_uds_client_config_t *cfg,
228 + nipc_uds_session_t *session)
229 +{
230 + uint8_t buf[128]; /* enough for header(32) + hello(44) = 76, and ack */
231 + nipc_uds_error_t err;
232 +
233 + /* Detect packet size if not specified */
234 + uint32_t pkt_size = cfg->packet_size;
235 + if (pkt_size == 0)
236 + pkt_size = detect_packet_size(fd);
237 +
238 + /* Build HELLO payload */
239 + nipc_hello_t hello = {
240 + .layout_version = 1,
241 + .flags = 0,
242 + .supported_profiles = cfg->supported_profiles ? cfg->supported_profiles : NIPC_PROFILE_BASELINE,
243 + .preferred_profiles = cfg->preferred_profiles,
244 + .max_request_payload_bytes = apply_default(cfg->max_request_payload_bytes, NIPC_MAX_PAYLOAD_DEFAULT),
245 + .max_request_batch_items = apply_default(cfg->max_request_batch_items, UDS_DEFAULT_BATCH_ITEMS),
246 + .max_response_payload_bytes = apply_default(cfg->max_response_payload_bytes, NIPC_MAX_PAYLOAD_DEFAULT),
247 + .max_response_batch_items = apply_default(cfg->max_response_batch_items, UDS_DEFAULT_BATCH_ITEMS),
248 + .auth_token = cfg->auth_token,
249 + .packet_size = pkt_size,
250 + };
251 +
252 + uint8_t hello_buf[44];
253 + nipc_hello_encode(&hello, hello_buf, sizeof(hello_buf));
254 +
255 + /* Build outer CONTROL header */
256 + nipc_header_t hdr = {
257 + .magic = NIPC_MAGIC_MSG,
258 + .version = NIPC_VERSION,
259 + .header_len = NIPC_HEADER_LEN,
260 + .kind = NIPC_KIND_CONTROL,
261 + .flags = 0,
262 + .code = NIPC_CODE_HELLO,
263 + .transport_status = NIPC_STATUS_OK,
264 + .payload_len = sizeof(hello_buf),
265 + .item_count = 1,
266 + .message_id = 0,
267 + };
268 +
269 + nipc_header_encode(&hdr, buf, sizeof(buf));
270 + memcpy(buf + NIPC_HEADER_LEN, hello_buf, sizeof(hello_buf));
271 +
272 + /* Send HELLO */
273 + err = raw_send(fd, buf, NIPC_HEADER_LEN + sizeof(hello_buf));
274 + if (err != NIPC_UDS_OK)
275 + return err;
276 +
277 + /* Receive HELLO_ACK */
278 + ssize_t n = raw_recv(fd, buf, sizeof(buf));
279 + if (n <= 0)
280 + return NIPC_UDS_ERR_RECV;
281 +
282 + /* Decode outer header */
283 + nipc_header_t ack_hdr;
284 + nipc_error_t perr = nipc_header_decode(buf, (size_t)n, &ack_hdr);
285 + if (perr == NIPC_ERR_BAD_VERSION)
286 + return NIPC_UDS_ERR_INCOMPATIBLE;
287 + if (perr != NIPC_OK)
288 + return NIPC_UDS_ERR_PROTOCOL;
289 +
290 + if (ack_hdr.kind != NIPC_KIND_CONTROL || ack_hdr.code != NIPC_CODE_HELLO_ACK)
291 + return NIPC_UDS_ERR_PROTOCOL;
292 +
293 + /* Check transport_status for rejection */
294 + if (ack_hdr.transport_status == NIPC_STATUS_AUTH_FAILED)
295 + return NIPC_UDS_ERR_AUTH_FAILED;
296 + if (ack_hdr.transport_status == NIPC_STATUS_UNSUPPORTED)
297 + return NIPC_UDS_ERR_NO_PROFILE;
298 + if (ack_hdr.transport_status == NIPC_STATUS_INCOMPATIBLE)
299 + return NIPC_UDS_ERR_INCOMPATIBLE;
300 + if (ack_hdr.transport_status == NIPC_STATUS_LIMIT_EXCEEDED)
301 + return NIPC_UDS_ERR_LIMIT_EXCEEDED;
302 + if (ack_hdr.transport_status != NIPC_STATUS_OK)
303 + return NIPC_UDS_ERR_HANDSHAKE;
304 +
305 + /* Decode hello-ack payload */
306 + nipc_hello_ack_t ack;
307 + perr = nipc_hello_ack_decode(buf + NIPC_HEADER_LEN,
308 + (size_t)n - NIPC_HEADER_LEN, &ack);
309 + if (perr == NIPC_ERR_BAD_LAYOUT &&
310 + hello_ack_layout_incompatible(buf + NIPC_HEADER_LEN,
311 + (size_t)n - NIPC_HEADER_LEN))
312 + return NIPC_UDS_ERR_INCOMPATIBLE;
313 + if (perr != NIPC_OK)
314 + return NIPC_UDS_ERR_PROTOCOL;
315 +
316 + /* Fill session */
317 + session->fd = fd;
318 + session->role = NIPC_UDS_ROLE_CLIENT;
319 + session->max_request_payload_bytes = ack.agreed_max_request_payload_bytes;
320 + session->max_request_batch_items = ack.agreed_max_request_batch_items;
321 + session->max_response_payload_bytes = ack.agreed_max_response_payload_bytes;
322 + session->max_response_batch_items = ack.agreed_max_response_batch_items;
323 + session->packet_size = ack.agreed_packet_size;
324 + session->selected_profile = ack.selected_profile;
325 + session->session_id = ack.session_id;
326 + session->recv_buf = NULL;
327 + session->recv_buf_size = 0;
328 +
329 + /* Sanity: reject a packet_size too small for chunking arithmetic */
330 + if (session->packet_size <= NIPC_HEADER_LEN)
331 + return NIPC_UDS_ERR_PROTOCOL;
332 +
333 + return NIPC_UDS_OK;
334 +}
335 +
336 +/* ------------------------------------------------------------------ */
337 +/* Handshake: server side */
338 +/* ------------------------------------------------------------------ */
339 +
340 +static nipc_uds_error_t server_handshake(int fd,
341 + const nipc_uds_server_config_t *cfg,
342 + uint64_t session_id,
343 + nipc_uds_session_t *session)
344 +{
345 + uint8_t buf[128];
346 +
347 + /* Detect server packet size */
348 + uint32_t server_pkt_size = cfg->packet_size;
349 + if (server_pkt_size == 0)
350 + server_pkt_size = detect_packet_size(fd);
351 +
352 + /* Server limits with defaults applied */
353 + uint32_t s_req_pay = apply_default(cfg->max_request_payload_bytes, NIPC_MAX_PAYLOAD_DEFAULT);
354 + uint32_t s_req_bat = apply_default(cfg->max_request_batch_items, UDS_DEFAULT_BATCH_ITEMS);
355 + uint32_t s_resp_pay = apply_default(cfg->max_response_payload_bytes, NIPC_MAX_PAYLOAD_DEFAULT);
356 + uint32_t s_resp_bat = apply_default(cfg->max_response_batch_items, UDS_DEFAULT_BATCH_ITEMS);
357 + uint32_t s_profiles = cfg->supported_profiles ? cfg->supported_profiles : NIPC_PROFILE_BASELINE;
358 + uint32_t s_preferred = cfg->preferred_profiles;
359 +
360 + /* Receive HELLO */
361 + ssize_t n = raw_recv(fd, buf, sizeof(buf));
362 + if (n <= 0)
363 + return NIPC_UDS_ERR_RECV;
364 +
365 + nipc_header_t hdr;
366 + nipc_error_t perr = nipc_header_decode(buf, (size_t)n, &hdr);
367 + if (perr == NIPC_ERR_BAD_VERSION &&
368 + header_version_incompatible(buf, (size_t)n, NIPC_CODE_HELLO)) {
369 + send_rejection_ack(fd, NIPC_STATUS_INCOMPATIBLE);
370 + return NIPC_UDS_ERR_INCOMPATIBLE;
371 + }
372 + if (perr != NIPC_OK)
373 + return NIPC_UDS_ERR_PROTOCOL;
374 +
375 + if (hdr.kind != NIPC_KIND_CONTROL || hdr.code != NIPC_CODE_HELLO)
376 + return NIPC_UDS_ERR_PROTOCOL;
377 +
378 + nipc_hello_t hello;
379 + perr = nipc_hello_decode(buf + NIPC_HEADER_LEN,
380 + (size_t)n - NIPC_HEADER_LEN, &hello);
381 + if (perr == NIPC_ERR_BAD_LAYOUT &&
382 + hello_layout_incompatible(buf + NIPC_HEADER_LEN,
383 + (size_t)n - NIPC_HEADER_LEN)) {
384 + send_rejection_ack(fd, NIPC_STATUS_INCOMPATIBLE);
385 + return NIPC_UDS_ERR_INCOMPATIBLE;
386 + }
387 + if (perr != NIPC_OK)
388 + return NIPC_UDS_ERR_PROTOCOL;
389 +
390 + /* Compute intersection */
391 + uint32_t intersection = hello.supported_profiles & s_profiles;
392 +
393 + /* Check intersection */
394 + if (intersection == 0) {
395 + send_rejection_ack(fd, NIPC_STATUS_UNSUPPORTED);
396 + return NIPC_UDS_ERR_NO_PROFILE;
397 + }
398 +
399 + /* Check auth */
400 + if (hello.auth_token != cfg->auth_token) {
401 + send_rejection_ack(fd, NIPC_STATUS_AUTH_FAILED);
402 + return NIPC_UDS_ERR_AUTH_FAILED;
403 + }
404 +
405 + /* Select profile: prefer preferred_intersection, then intersection */
406 + uint32_t preferred_intersection = intersection &
407 + hello.preferred_profiles & s_preferred;
408 + uint32_t selected;
409 + if (preferred_intersection != 0)
410 + selected = highest_bit(preferred_intersection);
411 + else
412 + selected = highest_bit(intersection);
413 +
414 + if (hello.max_request_payload_bytes > NIPC_MAX_PAYLOAD_CAP) {
415 + send_rejection_ack(fd, NIPC_STATUS_LIMIT_EXCEEDED);
416 + return NIPC_UDS_ERR_LIMIT_EXCEEDED;
417 + }
418 +
419 + /* Negotiate limits:
420 + * - request payload and batch size are client-proposed and echoed
421 + * - response payload is server-authoritative
422 + * - response batch size is symmetric with request batch size */
423 + uint32_t agreed_req_pay = hello.max_request_payload_bytes;
424 + uint32_t agreed_req_bat = hello.max_request_batch_items;
425 + uint32_t agreed_resp_pay = s_resp_pay;
426 + uint32_t agreed_resp_bat = agreed_req_bat;
427 + uint32_t agreed_pkt = min_u32(hello.packet_size, server_pkt_size);
428 +
429 + /* packet_size must be large enough for a usable message packet */
430 + if (agreed_pkt <= NIPC_HEADER_LEN) {
431 + send_rejection_ack(fd, NIPC_STATUS_INCOMPATIBLE);
432 + return NIPC_UDS_ERR_INCOMPATIBLE;
433 + }
434 +
435 + /* Send HELLO_ACK (success) */
436 + nipc_hello_ack_t ack = {
437 + .layout_version = 1,
438 + .flags = 0,
439 + .server_supported_profiles = s_profiles,
440 + .intersection_profiles = intersection,
441 + .selected_profile = selected,
442 + .agreed_max_request_payload_bytes = agreed_req_pay,
443 + .agreed_max_request_batch_items = agreed_req_bat,
444 + .agreed_max_response_payload_bytes = agreed_resp_pay,
445 + .agreed_max_response_batch_items = agreed_resp_bat,
446 + .agreed_packet_size = agreed_pkt,
447 + .session_id = session_id,
448 + };
449 +
450 + uint8_t ack_buf[48];
451 + nipc_hello_ack_encode(&ack, ack_buf, sizeof(ack_buf));
452 +
453 + nipc_header_t ack_hdr = {
454 + .magic = NIPC_MAGIC_MSG,
455 + .version = NIPC_VERSION,
456 + .header_len = NIPC_HEADER_LEN,
457 + .kind = NIPC_KIND_CONTROL,
458 + .flags = 0,
459 + .code = NIPC_CODE_HELLO_ACK,
460 + .transport_status = NIPC_STATUS_OK,
461 + .payload_len = sizeof(ack_buf),
462 + .item_count = 1,
463 + .message_id = 0,
464 + };
465 +
466 + uint8_t pkt[80];
467 + nipc_header_encode(&ack_hdr, pkt, sizeof(pkt));
468 + memcpy(pkt + NIPC_HEADER_LEN, ack_buf, sizeof(ack_buf));
469 +
470 + nipc_uds_error_t send_ack_err = raw_send(fd, pkt, NIPC_HEADER_LEN + sizeof(ack_buf));
471 + if (send_ack_err != NIPC_UDS_OK)
472 + return send_ack_err;
473 +
474 + /* Fill session */
475 + session->fd = fd;
476 + session->role = NIPC_UDS_ROLE_SERVER;
477 + session->max_request_payload_bytes = agreed_req_pay;
478 + session->max_request_batch_items = agreed_req_bat;
479 + session->max_response_payload_bytes = agreed_resp_pay;
480 + session->max_response_batch_items = agreed_resp_bat;
481 + session->packet_size = agreed_pkt;
482 + session->selected_profile = selected;
483 + session->session_id = session_id;
484 + session->recv_buf = NULL;
485 + session->recv_buf_size = 0;
486 +
487 + return NIPC_UDS_OK;
488 +}
489 +
490 +/* ------------------------------------------------------------------ */
491 +/* Stale endpoint recovery */
492 +/* ------------------------------------------------------------------ */
493 +
494 +/* Returns: 0 = stale (unlinked), 1 = live server, -1 = doesn't exist */
495 +static int check_and_recover_stale(const char *path)
496 +{
497 + struct stat st;
498 + if (stat(path, &st) != 0)
499 + return -1; /* doesn't exist */
500 +
501 + /* Try connecting to check if a live server is there */
502 + int probe = socket(AF_UNIX, SOCK_SEQPACKET, 0);
503 + if (probe < 0)
504 + return -1;
505 +
506 + struct sockaddr_un addr;
507 + memset(&addr, 0, sizeof(addr));
508 + addr.sun_family = AF_UNIX;
509 + strncpy(addr.sun_path, path, sizeof(addr.sun_path) - 1);
510 +
511 + int ret;
512 + if (connect(probe, (struct sockaddr *)&addr, sizeof(addr)) == 0) {
513 + /* Connected -> live server */
514 + close(probe);
515 + ret = 1;
516 + } else {
517 + int saved_errno = errno;
518 + close(probe);
519 + /* Only unlink on ECONNREFUSED/ENOENT (stale socket).
520 + * Other errors (EACCES, etc.) should not remove the file. */
521 + if (saved_errno == ECONNREFUSED || saved_errno == ENOENT) {
522 + unlink(path);
523 + ret = 0;
524 + } else {
525 + /* Can't determine ownership — treat as live to prevent overwriting */
526 + ret = 1;
527 + }
528 + }
529 + return ret;
530 +}
531 +
532 +/* ------------------------------------------------------------------ */
533 +/* Public API: listen */
534 +/* ------------------------------------------------------------------ */
535 +
536 +nipc_uds_error_t nipc_uds_listen(const char *run_dir,
537 + const char *service_name,
538 + const nipc_uds_server_config_t *config,
539 + nipc_uds_listener_t *out)
540 +{
541 + memset(out, 0, sizeof(*out));
542 + out->fd = -1;
543 +
544 + /* Build path */
545 + char path[sizeof(((struct sockaddr_un *)0)->sun_path)];
546 + int path_rc = build_socket_path(path, sizeof(path), run_dir, service_name);
547 + if (path_rc == -2)
548 + return NIPC_UDS_ERR_BAD_PARAM;
549 + if (path_rc < 0)
550 + return NIPC_UDS_ERR_PATH_TOO_LONG;
551 +
552 + /* Stale recovery */
553 + int stale = check_and_recover_stale(path);
554 + if (stale == 1)
555 + return NIPC_UDS_ERR_ADDR_IN_USE;
556 +
557 + /* Create socket */
558 + int fd = socket(AF_UNIX, SOCK_SEQPACKET, 0);
559 + if (fd < 0)
560 + return NIPC_UDS_ERR_SOCKET;
561 +
562 + struct sockaddr_un addr;
563 + memset(&addr, 0, sizeof(addr));
564 + addr.sun_family = AF_UNIX;
565 + strncpy(addr.sun_path, path, sizeof(addr.sun_path) - 1);
566 +
567 + if (bind(fd, (struct sockaddr *)&addr, sizeof(addr)) < 0) {
568 + close(fd);
569 + return NIPC_UDS_ERR_SOCKET;
570 + }
571 +
572 + int backlog = config->backlog > 0 ? config->backlog : UDS_DEFAULT_BACKLOG;
573 + if (listen(fd, backlog) < 0) {
574 + close(fd);
575 + unlink(path);
576 + return NIPC_UDS_ERR_SOCKET;
577 + }
578 +
579 + out->fd = fd;
580 + out->config = *config;
581 + strncpy(out->path, path, sizeof(out->path) - 1);
582 + out->path[sizeof(out->path) - 1] = '\0';
583 +
584 + return NIPC_UDS_OK;
585 +}
586 +
587 +/* ------------------------------------------------------------------ */
588 +/* Public API: accept */
589 +/* ------------------------------------------------------------------ */
590 +
591 +nipc_uds_error_t nipc_uds_accept(nipc_uds_listener_t *listener,
592 + uint64_t session_id,
593 + nipc_uds_session_t *out)
594 +{
595 + memset(out, 0, sizeof(*out));
596 + out->fd = -1;
597 +
598 + int client_fd = accept(listener->fd, NULL, NULL);
599 + if (client_fd < 0)
600 + return NIPC_UDS_ERR_ACCEPT;
601 +
602 + nipc_uds_error_t err = server_handshake(client_fd, &listener->config,
603 + session_id, out);
604 + if (err != NIPC_UDS_OK) {
605 + close(client_fd);
606 + out->fd = -1;
607 + return err;
608 + }
609 +
610 + return NIPC_UDS_OK;
611 +}
612 +
613 +/* ------------------------------------------------------------------ */
614 +/* Public API: connect */
615 +/* ------------------------------------------------------------------ */
616 +
617 +nipc_uds_error_t nipc_uds_connect(const char *run_dir,
618 + const char *service_name,
619 + const nipc_uds_client_config_t *config,
620 + nipc_uds_session_t *out)
621 +{
622 + memset(out, 0, sizeof(*out));
623 + out->fd = -1;
624 +
625 + /* Build path */
626 + char path[sizeof(((struct sockaddr_un *)0)->sun_path)];
627 + int path_rc2 = build_socket_path(path, sizeof(path), run_dir, service_name);
628 + if (path_rc2 == -2)
629 + return NIPC_UDS_ERR_BAD_PARAM;
630 + if (path_rc2 < 0)
631 + return NIPC_UDS_ERR_PATH_TOO_LONG;
632 +
633 + /* Create socket */
634 + int fd = socket(AF_UNIX, SOCK_SEQPACKET, 0);
635 + if (fd < 0)
636 + return NIPC_UDS_ERR_SOCKET;
637 +
638 + struct sockaddr_un addr;
639 + memset(&addr, 0, sizeof(addr));
640 + addr.sun_family = AF_UNIX;
641 + strncpy(addr.sun_path, path, sizeof(addr.sun_path) - 1);
642 +
643 + if (connect(fd, (struct sockaddr *)&addr, sizeof(addr)) < 0) {
644 + close(fd);
645 + return NIPC_UDS_ERR_CONNECT;
646 + }
647 +
648 + nipc_uds_error_t err = client_handshake(fd, config, out);
649 + if (err != NIPC_UDS_OK) {
650 + close(fd);
651 + out->fd = -1;
652 + return err;
653 + }
654 +
655 + return NIPC_UDS_OK;
656 +}
657 +
658 +/* ------------------------------------------------------------------ */
659 +/* Public API: close */
660 +/* ------------------------------------------------------------------ */
661 +
662 +void nipc_uds_close_session(nipc_uds_session_t *session)
663 +{
664 + if (!session)
665 + return;
666 +
667 + if (session->fd >= 0) {
668 + close(session->fd);
669 + session->fd = -1;
670 + }
671 +
672 + free(session->recv_buf);
673 + session->recv_buf = NULL;
674 + session->recv_buf_size = 0;
675 +
676 + free(session->inflight_ids);
677 + session->inflight_ids = NULL;
678 + session->inflight_count = 0;
679 + session->inflight_capacity = 0;
680 +}
681 +
682 +void nipc_uds_close_listener(nipc_uds_listener_t *listener)
683 +{
684 + if (!listener)
685 + return;
686 +
687 + if (listener->fd >= 0) {
688 + close(listener->fd);
689 + listener->fd = -1;
690 + }
691 +
692 + if (listener->path[0]) {
693 + unlink(listener->path);
694 + listener->path[0] = '\0';
695 + }
696 +}
697 +
698 +/* ------------------------------------------------------------------ */
699 +/* In-flight message_id tracking (client side) */
700 +/* ------------------------------------------------------------------ */
701 +
702 +/* Add message_id to in-flight set. Returns 0 on success, -1 if duplicate,
703 + * -2 on allocation failure. */
704 +static int inflight_add(nipc_uds_session_t *s, uint64_t id)
705 +{
706 + for (uint32_t i = 0; i < s->inflight_count; i++) {
707 + if (s->inflight_ids[i] == id)
708 + return -1; /* duplicate */
709 + }
710 + if (s->inflight_count >= s->inflight_capacity) {
711 + uint32_t new_cap = s->inflight_capacity ? s->inflight_capacity * 2 : 16;
712 + uint64_t *new_ids = realloc(s->inflight_ids, (size_t)new_cap * sizeof(uint64_t));
713 + if (!new_ids)
714 + return -2; /* allocation failure */
715 + s->inflight_ids = new_ids;
716 + s->inflight_capacity = new_cap;
717 + }
718 + s->inflight_ids[s->inflight_count++] = id;
719 + return 0;
720 +}
721 +
722 +/* Remove message_id from in-flight set. Returns 0 on success, -1 if
723 + * not found. */
724 +static int inflight_remove(nipc_uds_session_t *s, uint64_t id)
725 +{
726 + for (uint32_t i = 0; i < s->inflight_count; i++) {
727 + if (s->inflight_ids[i] == id) {
728 + /* Swap with last */
729 + s->inflight_ids[i] = s->inflight_ids[s->inflight_count - 1];
730 + s->inflight_count--;
731 + return 0;
732 + }
733 + }
734 + return -1; /* not found */
735 +}
736 +
737 +static void inflight_fail_all(nipc_uds_session_t *s)
738 +{
739 + if (!s || s->role != NIPC_UDS_ROLE_CLIENT)
740 + return;
741 +
742 + /* A broken session invalidates every in-flight request on it. */
743 + s->inflight_count = 0;
744 +}
745 +
746 +/* ------------------------------------------------------------------ */
747 +/* Public API: send */
748 +/* ------------------------------------------------------------------ */
749 +
750 +nipc_uds_error_t nipc_uds_send(nipc_uds_session_t *session,
751 + nipc_header_t *hdr,
752 + const void *payload,
753 + size_t payload_len)
754 +{
755 + if (!session || session->fd < 0)
756 + return NIPC_UDS_ERR_BAD_PARAM;
757 +
758 + int tracked = (session->role == NIPC_UDS_ROLE_CLIENT &&
759 + hdr->kind == NIPC_KIND_REQUEST);
760 +
761 + /* Client-side: track in-flight message_ids for requests */
762 + if (tracked) {
763 + int rc = inflight_add(session, hdr->message_id);
764 + if (rc == -1)
765 + return NIPC_UDS_ERR_DUPLICATE_MSG_ID;
766 + if (rc == -2)
767 + return NIPC_UDS_ERR_LIMIT_EXCEEDED;
768 + }
769 +
770 + uint32_t max_payload = 0;
771 + uint32_t max_batch = 0;
772 + if (session->role == NIPC_UDS_ROLE_CLIENT &&
773 + hdr->kind == NIPC_KIND_REQUEST) {
774 + max_payload = session->max_request_payload_bytes;
775 + max_batch = session->max_request_batch_items;
776 + } else if (session->role == NIPC_UDS_ROLE_SERVER &&
777 + hdr->kind == NIPC_KIND_RESPONSE) {
778 + max_payload = session->max_response_payload_bytes;
779 + max_batch = session->max_response_batch_items;
780 + }
781 + if (payload_len > UINT32_MAX ||
782 + (max_payload > 0 && payload_len > max_payload) ||
783 + (max_batch > 0 && hdr->item_count > max_batch)) {
784 + if (tracked)
785 + inflight_remove(session, hdr->message_id);
786 + return NIPC_UDS_ERR_LIMIT_EXCEEDED;
787 + }
788 +
789 + /* Fill envelope fields the caller shouldn't set */
790 + hdr->magic = NIPC_MAGIC_MSG;
791 + hdr->version = NIPC_VERSION;
792 + hdr->header_len = NIPC_HEADER_LEN;
793 + hdr->payload_len = (uint32_t)payload_len;
794 +
795 + size_t total_msg = NIPC_HEADER_LEN + payload_len;
796 +
797 + /* Does it fit in one packet? */
798 + if (total_msg <= session->packet_size) {
799 + /* Single packet: encode header then send header+payload together */
800 + uint8_t hdr_buf[NIPC_HEADER_LEN];
801 + nipc_header_encode(hdr, hdr_buf, sizeof(hdr_buf));
802 + nipc_uds_error_t send_err = raw_send_iov(session->fd, hdr_buf, NIPC_HEADER_LEN,
803 + payload, payload_len);
804 + if (send_err != NIPC_UDS_OK) {
805 + if (session->role == NIPC_UDS_ROLE_CLIENT &&
806 + hdr->kind == NIPC_KIND_REQUEST) {
807 + if (send_err == NIPC_UDS_ERR_SEND)
808 + inflight_fail_all(session);
809 + else
810 + inflight_remove(session, hdr->message_id);
811 + }
812 + }
813 + return send_err;
814 + }
815 +
816 + /* Chunked send */
817 + size_t chunk_payload_budget = session->packet_size - NIPC_HEADER_LEN;
818 + if (chunk_payload_budget == 0)
819 + return NIPC_UDS_ERR_BAD_PARAM;
820 +
821 + /* Calculate chunk count:
822 + * First chunk: header(32) + up to chunk_payload_budget payload bytes.
823 + * But the first chunk carries the outer header, so payload in first
824 + * chunk = chunk_payload_budget. */
825 + size_t remaining = payload_len;
826 + size_t first_chunk_payload = remaining < chunk_payload_budget
827 + ? remaining : chunk_payload_budget;
828 +
829 + remaining -= first_chunk_payload;
830 + uint32_t continuation_chunks = 0;
831 + if (remaining > 0) {
832 + continuation_chunks = (uint32_t)((remaining + chunk_payload_budget - 1)
833 + / chunk_payload_budget);
834 + }
835 + uint32_t chunk_count = 1 + continuation_chunks;
836 +
837 + /* Send first chunk: outer header + first part of payload */
838 + uint8_t hdr_buf[NIPC_HEADER_LEN];
839 + nipc_header_encode(hdr, hdr_buf, sizeof(hdr_buf));
840 +
841 + nipc_uds_error_t err = raw_send_iov(session->fd, hdr_buf, NIPC_HEADER_LEN,
842 + payload, first_chunk_payload);
843 + if (err != NIPC_UDS_OK) {
844 + if (session->role == NIPC_UDS_ROLE_CLIENT &&
845 + hdr->kind == NIPC_KIND_REQUEST) {
846 + if (err == NIPC_UDS_ERR_SEND)
847 + inflight_fail_all(session);
848 + else
849 + inflight_remove(session, hdr->message_id);
850 + }
851 + return err;
852 + }
853 +
854 + /* Send continuation chunks */
855 + const uint8_t *src = (const uint8_t *)payload + first_chunk_payload;
856 + remaining = payload_len - first_chunk_payload;
857 +
858 + for (uint32_t ci = 1; ci < chunk_count; ci++) {
859 + size_t this_chunk = remaining < chunk_payload_budget
860 + ? remaining : chunk_payload_budget;
861 +
862 + nipc_chunk_header_t chk = {
863 + .magic = NIPC_MAGIC_CHUNK,
864 + .version = NIPC_VERSION,
865 + .flags = 0,
866 + .message_id = hdr->message_id,
867 + .total_message_len = (uint32_t)total_msg,
868 + .chunk_index = ci,
869 + .chunk_count = chunk_count,
870 + .chunk_payload_len = (uint32_t)this_chunk,
871 + };
872 +
873 + uint8_t chk_buf[NIPC_HEADER_LEN];
874 + nipc_chunk_header_encode(&chk, chk_buf, sizeof(chk_buf));
875 +
876 + err = raw_send_iov(session->fd, chk_buf, NIPC_HEADER_LEN,
877 + src, this_chunk);
878 + if (err != NIPC_UDS_OK) {
879 + if (session->role == NIPC_UDS_ROLE_CLIENT &&
880 + hdr->kind == NIPC_KIND_REQUEST) {
881 + if (err == NIPC_UDS_ERR_SEND)
882 + inflight_fail_all(session);
883 + else
884 + inflight_remove(session, hdr->message_id);
885 + }
886 + return err;
887 + }
888 +
889 + src += this_chunk;
890 + remaining -= this_chunk;
891 + }
892 +
893 + return NIPC_UDS_OK;
894 +}
895 +
896 +/* ------------------------------------------------------------------ */
897 +/* Public API: receive */
898 +/* ------------------------------------------------------------------ */
899 +
900 +/* Ensure session recv_buf can hold at least `needed` bytes. */
901 +static nipc_uds_error_t ensure_recv_buf(nipc_uds_session_t *session,
902 + size_t needed)
903 +{
904 + if (session->recv_buf_size >= needed)
905 + return NIPC_UDS_OK;
906 +
907 + uint8_t *p = realloc(session->recv_buf, needed);
908 + if (!p)
909 + return NIPC_UDS_ERR_ALLOC;
910 +
911 + session->recv_buf = p;
912 + session->recv_buf_size = needed;
913 + return NIPC_UDS_OK;
914 +}
915 +
916 +/* Validate batch directory if the message has BATCH flag and item_count > 1.
917 + * Called after the full payload is assembled, before returning to caller. */
918 +static nipc_uds_error_t validate_batch(const nipc_header_t *hdr,
919 + const void *payload, size_t payload_len)
920 +{
921 + if (!(hdr->flags & NIPC_FLAG_BATCH) || hdr->item_count <= 1)
922 + return NIPC_UDS_OK;
923 +
924 + uint32_t dir_bytes = hdr->item_count * 8;
925 + uint32_t dir_aligned = (uint32_t)nipc_align8(dir_bytes);
926 + if (payload_len < dir_aligned)
927 + return NIPC_UDS_ERR_PROTOCOL;
928 +
929 + uint32_t packed_area_len = (uint32_t)(payload_len - dir_aligned);
930 + nipc_error_t perr = nipc_batch_dir_validate(payload, dir_bytes,
931 + hdr->item_count,
932 + packed_area_len);
933 + return (perr == NIPC_OK) ? NIPC_UDS_OK : NIPC_UDS_ERR_PROTOCOL;
934 +}
935 +
936 +nipc_uds_error_t nipc_uds_receive(nipc_uds_session_t *session,
937 + void *buf, size_t buf_size,
938 + nipc_header_t *hdr_out,
939 + const void **payload_out,
940 + size_t *payload_len_out)
941 +{
942 + if (!session || session->fd < 0)
943 + return NIPC_UDS_ERR_BAD_PARAM;
944 +
945 + /* Read first packet into the caller's buffer */
946 + ssize_t n = raw_recv(session->fd, buf, buf_size);
947 + if (n <= 0) {
948 + inflight_fail_all(session);
949 + return NIPC_UDS_ERR_RECV;
950 + }
951 +
952 + if ((size_t)n < NIPC_HEADER_LEN)
953 + return NIPC_UDS_ERR_PROTOCOL;
954 +
955 + /* Decode outer header */
956 + nipc_error_t perr = nipc_header_decode(buf, (size_t)n, hdr_out);
957 + if (perr != NIPC_OK)
958 + return NIPC_UDS_ERR_PROTOCOL;
959 +
960 + /* Validate payload_len against negotiated directional limit.
961 + * Server receives requests; client receives responses. */
962 + uint32_t max_payload = (session->role == NIPC_UDS_ROLE_SERVER)
963 + ? session->max_request_payload_bytes
964 + : session->max_response_payload_bytes;
965 + if (hdr_out->payload_len > max_payload)
966 + return NIPC_UDS_ERR_LIMIT_EXCEEDED;
967 +
968 + /* Validate item_count against negotiated directional batch limit. */
969 + uint32_t max_batch = (session->role == NIPC_UDS_ROLE_SERVER)
970 + ? session->max_request_batch_items
971 + : session->max_response_batch_items;
972 + if (hdr_out->item_count > max_batch)
973 + return NIPC_UDS_ERR_LIMIT_EXCEEDED;
974 +
975 + /* Client-side: validate response message_id is in-flight */
976 + if (session->role == NIPC_UDS_ROLE_CLIENT &&
977 + hdr_out->kind == NIPC_KIND_RESPONSE) {
978 + if (inflight_remove(session, hdr_out->message_id) < 0)
979 + return NIPC_UDS_ERR_UNKNOWN_MSG_ID;
980 + }
981 +
982 + size_t total_msg = NIPC_HEADER_LEN + hdr_out->payload_len;
983 +
984 + /* Non-chunked: entire message arrived in one packet */
985 + if ((size_t)n >= total_msg) {
986 + *payload_out = (const uint8_t *)buf + NIPC_HEADER_LEN;
987 + *payload_len_out = hdr_out->payload_len;
988 +
989 + nipc_uds_error_t berr = validate_batch(hdr_out, *payload_out, *payload_len_out);
990 + if (berr != NIPC_UDS_OK)
991 + return berr;
992 +
993 + return NIPC_UDS_OK;
994 + }
995 +
996 + /* Chunked: first packet has partial payload. The total message
997 + * size is NIPC_HEADER_LEN + payload_len from the header. */
998 + size_t first_payload_bytes = (size_t)n - NIPC_HEADER_LEN;
999 +
1000 + /* We need a buffer for the full payload. Use the session recv_buf. */
1001 + nipc_uds_error_t err = ensure_recv_buf(session, hdr_out->payload_len);
1002 + if (err != NIPC_UDS_OK)
1003 + return err;
1004 +
1005 + /* Copy first chunk's payload into recv_buf */
1006 + memcpy(session->recv_buf, (uint8_t *)buf + NIPC_HEADER_LEN,
1007 + first_payload_bytes);
1008 +
1009 + size_t assembled = first_payload_bytes;
1010 + size_t chunk_payload_budget = session->packet_size - NIPC_HEADER_LEN;
1011 +
1012 + /* Calculate expected chunk count */
1013 + size_t remaining_after_first = hdr_out->payload_len - first_payload_bytes;
1014 + uint32_t expected_continuations = 0;
1015 + if (remaining_after_first > 0 && chunk_payload_budget > 0) {
1016 + expected_continuations = (uint32_t)((remaining_after_first +
1017 + chunk_payload_budget - 1)
1018 + / chunk_payload_budget);
1019 + }
1020 + uint32_t expected_chunk_count = 1 + expected_continuations;
1021 +
1022 + /* A temporary buffer for reading continuation packets */
1023 + size_t pkt_buf_size = session->packet_size;
1024 + uint8_t *pkt_buf = malloc(pkt_buf_size);
1025 + if (!pkt_buf)
1026 + return NIPC_UDS_ERR_ALLOC;
1027 +
1028 + for (uint32_t ci = 1; assembled < hdr_out->payload_len; ci++) {
1029 + ssize_t cn = raw_recv(session->fd, pkt_buf, pkt_buf_size);
1030 + if (cn <= 0) {
1031 + free(pkt_buf);
1032 + inflight_fail_all(session);
1033 + return NIPC_UDS_ERR_RECV;
1034 + }
1035 +
1036 + if ((size_t)cn < NIPC_HEADER_LEN) {
1037 + free(pkt_buf);
1038 + return NIPC_UDS_ERR_CHUNK;
1039 + }
1040 +
1041 + nipc_chunk_header_t chk;
1042 + perr = nipc_chunk_header_decode(pkt_buf, (size_t)cn, &chk);
1043 + if (perr != NIPC_OK) {
1044 + free(pkt_buf);
1045 + return NIPC_UDS_ERR_CHUNK;
1046 + }
1047 +
1048 + /* Validate chunk header */
1049 + if (chk.message_id != hdr_out->message_id ||
1050 + chk.chunk_index != ci ||
1051 + chk.chunk_count != expected_chunk_count ||
1052 + chk.total_message_len != (uint32_t)total_msg) {
1053 + free(pkt_buf);
1054 + return NIPC_UDS_ERR_CHUNK;
1055 + }
1056 +
1057 + size_t chunk_data = (size_t)cn - NIPC_HEADER_LEN;
1058 + if (chunk_data != chk.chunk_payload_len) {
1059 + free(pkt_buf);
1060 + return NIPC_UDS_ERR_CHUNK;
1061 + }
1062 +
1063 + if (assembled + chunk_data > hdr_out->payload_len) {
1064 + free(pkt_buf);
1065 + return NIPC_UDS_ERR_CHUNK;
1066 + }
1067 +
1068 + memcpy(session->recv_buf + assembled,
1069 + pkt_buf + NIPC_HEADER_LEN, chunk_data);
1070 + assembled += chunk_data;
1071 + }
1072 +
1073 + free(pkt_buf);
1074 +
1075 + *payload_out = session->recv_buf;
1076 + *payload_len_out = hdr_out->payload_len;
1077 +
1078 + nipc_uds_error_t berr = validate_batch(hdr_out, *payload_out, *payload_len_out);
1079 + if (berr != NIPC_UDS_OK)
1080 + return berr;
1081 +
1082 + return NIPC_UDS_OK;
1083 +}
src/libnetdata/netipc/src/transport/windows/netipc_named_pipe.c new
+1222
@@ -0,0 +1,1222 @@
1 +/*
2 + * netipc_named_pipe.c - L1 Windows Named Pipe transport.
3 + *
4 + * Implements connection lifecycle, handshake with profile/limit negotiation,
5 + * and send/receive with transparent chunking over Win32 Named Pipes in
6 + * message mode. Wire-compatible with all language implementations.
7 + */
8 +
9 +#if defined(_WIN32) || defined(__MSYS__)
10 +
11 +#include "netipc/netipc_named_pipe.h"
12 +#include "netipc/netipc_protocol.h"
13 +
14 +#include <stdio.h>
15 +#include <stdlib.h>
16 +#include <string.h>
17 +
18 +/* ------------------------------------------------------------------ */
19 +/* Win32 error constants for disconnect detection */
20 +/* ------------------------------------------------------------------ */
21 +
22 +#ifndef ERROR_BROKEN_PIPE
23 +#define ERROR_BROKEN_PIPE 109
24 +#endif
25 +#ifndef ERROR_NO_DATA
26 +#define ERROR_NO_DATA 232
27 +#endif
28 +#ifndef ERROR_PIPE_NOT_CONNECTED
29 +#define ERROR_PIPE_NOT_CONNECTED 233
30 +#endif
31 +
32 +/* ------------------------------------------------------------------ */
33 +/* Internal helpers */
34 +/* ------------------------------------------------------------------ */
35 +
36 +static inline uint32_t min_u32(uint32_t a, uint32_t b)
37 +{
38 + return a < b ? a : b;
39 +}
40 +
41 +static inline uint32_t max_u32(uint32_t a, uint32_t b)
42 +{
43 + return a > b ? a : b;
44 +}
45 +
46 +static inline uint32_t apply_default(uint32_t val, uint32_t def)
47 +{
48 + return val == 0 ? def : val;
49 +}
50 +
51 +static inline uint32_t pipe_buffer_size(uint32_t packet_size)
52 +{
53 + /* The protocol packet size controls logical framing and chunk size. The
54 + * underlying pipe quota must stay large enough for full-duplex pipelining
55 + * even when tests force a tiny protocol packet size. */
56 + return max_u32(apply_default(packet_size, NIPC_NP_DEFAULT_PIPE_BUF_SIZE),
57 + NIPC_NP_DEFAULT_PIPE_BUF_SIZE);
58 +}
59 +
60 +static bool header_version_incompatible(const void *buf, size_t buf_len,
61 + uint16_t expected_code)
62 +{
63 + if (buf_len < NIPC_HEADER_LEN)
64 + return false;
65 +
66 + nipc_header_t hdr;
67 + memcpy(&hdr, buf, sizeof(hdr));
68 + return hdr.magic == NIPC_MAGIC_MSG &&
69 + hdr.version != NIPC_VERSION &&
70 + hdr.header_len == NIPC_HEADER_LEN &&
71 + hdr.kind == NIPC_KIND_CONTROL &&
72 + hdr.code == expected_code;
73 +}
74 +
75 +static bool hello_layout_incompatible(const void *buf, size_t buf_len)
76 +{
77 + if (buf_len < sizeof(uint16_t))
78 + return false;
79 +
80 + nipc_hello_t hello;
81 + memset(&hello, 0, sizeof(hello));
82 + memcpy(&hello, buf, sizeof(uint16_t));
83 + return hello.layout_version != 1;
84 +}
85 +
86 +static bool hello_ack_layout_incompatible(const void *buf, size_t buf_len)
87 +{
88 + if (buf_len < sizeof(uint16_t))
89 + return false;
90 +
91 + nipc_hello_ack_t ack;
92 + memset(&ack, 0, sizeof(ack));
93 + memcpy(&ack, buf, sizeof(uint16_t));
94 + return ack.layout_version != 1;
95 +}
96 +
97 +static nipc_np_error_t raw_send(HANDLE pipe, const void *data, size_t len);
98 +
99 +static void send_rejection_ack(HANDLE pipe, uint16_t status)
100 +{
101 + nipc_hello_ack_t ack = { .layout_version = 1 };
102 + uint8_t ack_buf[48];
103 + uint8_t pkt[80];
104 + nipc_header_t ack_hdr = {
105 + .magic = NIPC_MAGIC_MSG,
106 + .version = NIPC_VERSION,
107 + .header_len = NIPC_HEADER_LEN,
108 + .kind = NIPC_KIND_CONTROL,
109 + .code = NIPC_CODE_HELLO_ACK,
110 + .transport_status = status,
111 + .payload_len = sizeof(ack_buf),
112 + .item_count = 1,
113 + };
114 +
115 + nipc_hello_ack_encode(&ack, ack_buf, sizeof(ack_buf));
116 + nipc_header_encode(&ack_hdr, pkt, sizeof(pkt));
117 + memcpy(pkt + NIPC_HEADER_LEN, ack_buf, sizeof(ack_buf));
118 + raw_send(pipe, pkt, NIPC_HEADER_LEN + sizeof(ack_buf));
119 +}
120 +
121 +/* Highest set bit in a bitmask (0 if empty). */
122 +static uint32_t highest_bit(uint32_t mask)
123 +{
124 + if (mask == 0)
125 + return 0;
126 +
127 + uint32_t bit = 1u << 31;
128 + while (!(mask & bit))
129 + bit >>= 1;
130 + return bit;
131 +}
132 +
133 +/* ------------------------------------------------------------------ */
134 +/* FNV-1a 64-bit hash */
135 +/* ------------------------------------------------------------------ */
136 +
137 +uint64_t nipc_fnv1a_64(const void *data, size_t len)
138 +{
139 + uint64_t hash = NIPC_FNV1A_OFFSET_BASIS;
140 + const uint8_t *p = (const uint8_t *)data;
141 +
142 + for (size_t i = 0; i < len; i++) {
143 + hash ^= (uint64_t)p[i];
144 + hash *= NIPC_FNV1A_PRIME;
145 + }
146 +
147 + return hash;
148 +}
149 +
150 +/* ------------------------------------------------------------------ */
151 +/* Service name validation */
152 +/* ------------------------------------------------------------------ */
153 +
154 +/* Validate service_name: only [a-zA-Z0-9._-], non-empty, not "." or "..". */
155 +static int validate_service_name(const char *name)
156 +{
157 + if (!name || name[0] == '\0')
158 + return -1;
159 +
160 + /* Reject "." and ".." */
161 + if (name[0] == '.' && (name[1] == '\0' || (name[1] == '.' && name[2] == '\0')))
162 + return -1;
163 +
164 + for (const char *p = name; *p; p++) {
165 + char c = *p;
166 + if ((c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') ||
167 + (c >= '0' && c <= '9') || c == '.' || c == '_' || c == '-')
168 + continue;
169 + return -1;
170 + }
171 + return 0;
172 +}
173 +
174 +/* ------------------------------------------------------------------ */
175 +/* Pipe name derivation */
176 +/* ------------------------------------------------------------------ */
177 +
178 +int nipc_np_build_pipe_name(wchar_t *dst, size_t dst_chars,
179 + const char *run_dir,
180 + const char *service_name)
181 +{
182 + if (!dst || dst_chars == 0 || !run_dir || !service_name)
183 + return -1;
184 +
185 + if (validate_service_name(service_name) < 0)
186 + return -1;
187 +
188 + /* Hash the run_dir */
189 + uint64_t hash = nipc_fnv1a_64(run_dir, strlen(run_dir));
190 +
191 + /* Build the pipe name as narrow first, then widen */
192 + char narrow[NIPC_NP_MAX_PIPE_NAME];
193 + int n = snprintf(narrow, sizeof(narrow),
194 + "\\\\.\\pipe\\netipc-%016llx-%s",
195 + (unsigned long long)hash, service_name);
196 + if (n < 0 || (size_t)n >= sizeof(narrow))
197 + return -1;
198 +
199 + /* Convert to wide string */
200 + if ((size_t)(n + 1) > dst_chars)
201 + return -1;
202 +
203 + for (int i = 0; i <= n; i++)
204 + dst[i] = (wchar_t)(unsigned char)narrow[i];
205 +
206 + return 0;
207 +}
208 +
209 +/* ------------------------------------------------------------------ */
210 +/* Create a new pipe instance */
211 +/* ------------------------------------------------------------------ */
212 +
213 +static HANDLE create_pipe_instance(const wchar_t *pipe_name,
214 + uint32_t buf_size,
215 + BOOL first_instance)
216 +{
217 + DWORD open_mode = PIPE_ACCESS_DUPLEX;
218 + if (first_instance)
219 + open_mode |= FILE_FLAG_FIRST_PIPE_INSTANCE;
220 +
221 + HANDLE h = CreateNamedPipeW(
222 + pipe_name,
223 + open_mode,
224 + PIPE_TYPE_MESSAGE | PIPE_READMODE_MESSAGE | PIPE_WAIT,
225 + NIPC_NP_MAX_INSTANCES,
226 + buf_size,
227 + buf_size,
228 + 0, /* default timeout */
229 + NULL /* default security */
230 + );
231 +
232 + return h;
233 +}
234 +
235 +/* ------------------------------------------------------------------ */
236 +/* Low-level send / recv */
237 +/* ------------------------------------------------------------------ */
238 +
239 +/* Check if a Win32 error indicates peer disconnect (graceful). */
240 +static int is_disconnect_error(DWORD err)
241 +{
242 + return err == ERROR_BROKEN_PIPE ||
243 + err == ERROR_NO_DATA ||
244 + err == ERROR_PIPE_NOT_CONNECTED;
245 +}
246 +
247 +/* Write exactly len bytes as one pipe message. */
248 +static nipc_np_error_t raw_send(HANDLE pipe, const void *data, size_t len)
249 +{
250 + DWORD written = 0;
251 + BOOL ok = WriteFile(pipe, data, (DWORD)len, &written, NULL);
252 + if (!ok) {
253 + if (is_disconnect_error(GetLastError()))
254 + return NIPC_NP_ERR_DISCONNECTED;
255 + return NIPC_NP_ERR_SEND;
256 + }
257 + if (written != (DWORD)len)
258 + return NIPC_NP_ERR_SEND;
259 + return NIPC_NP_OK;
260 +}
261 +
262 +/* Send header + payload as one pipe message (concatenated). */
263 +static nipc_np_error_t raw_send_msg(HANDLE pipe,
264 + const void *hdr, size_t hdr_len,
265 + const void *payload, size_t payload_len)
266 +{
267 + /* Named Pipes don't support scatter-gather like sendmsg.
268 + * Concatenate into one buffer for a single WriteFile call. */
269 + size_t total = hdr_len + payload_len;
270 + uint8_t stack_buf[256];
271 + uint8_t *buf;
272 + int allocated = 0;
273 +
274 + if (total <= sizeof(stack_buf)) {
275 + buf = stack_buf;
276 + } else {
277 + buf = (uint8_t *)malloc(total);
278 + if (!buf)
279 + return NIPC_NP_ERR_ALLOC;
280 + allocated = 1;
281 + }
282 +
283 + memcpy(buf, hdr, hdr_len);
284 + if (payload && payload_len > 0)
285 + memcpy(buf + hdr_len, payload, payload_len);
286 +
287 + nipc_np_error_t err = raw_send(pipe, buf, total);
288 +
289 + if (allocated)
290 + free(buf);
291 +
292 + return err;
293 +}
294 +
295 +/* Read one pipe message. Returns bytes read, 0 on disconnect, -1 on error. */
296 +static int raw_recv(HANDLE pipe, void *buf, size_t buf_len, DWORD *bytes_read)
297 +{
298 + DWORD read = 0;
299 + BOOL ok = ReadFile(pipe, buf, (DWORD)buf_len, &read, NULL);
300 + if (!ok) {
301 + DWORD err = GetLastError();
302 + if (is_disconnect_error(err)) {
303 + *bytes_read = 0;
304 + return 0; /* disconnect */
305 + }
306 + return -1; /* error */
307 + }
308 + if (read == 0) {
309 + *bytes_read = 0;
310 + return 0; /* disconnect */
311 + }
312 + *bytes_read = read;
313 + return 1; /* success */
314 +}
315 +
316 +/* Poll a synchronous pipe handle until bytes are available. Go/Rust already
317 + * use PeekNamedPipe here; waiting on the pipe HANDLE itself is not reliable
318 + * for request readability. */
319 +static nipc_np_error_t wait_readable(HANDLE pipe,
320 + uint32_t timeout_ms,
321 + bool *readable_out)
322 +{
323 + ULONGLONG deadline = GetTickCount64() + timeout_ms;
324 + int yielded = 0;
325 +
326 + if (readable_out)
327 + *readable_out = false;
328 +
329 + for (;;) {
330 + DWORD available = 0;
331 + BOOL ok = PeekNamedPipe(pipe, NULL, 0, NULL, &available, NULL);
332 + if (!ok) {
333 + DWORD err = GetLastError();
334 + if (is_disconnect_error(err))
335 + return NIPC_NP_ERR_DISCONNECTED;
336 + return NIPC_NP_ERR_RECV;
337 + }
338 +
339 + if (available > 0) {
340 + if (readable_out)
341 + *readable_out = true;
342 + return NIPC_NP_OK;
343 + }
344 +
345 + if (GetTickCount64() >= deadline)
346 + return NIPC_NP_OK;
347 +
348 + if (!yielded) {
349 + yielded = 1;
350 + for (int i = 0; i < 256; i++) {
351 + SwitchToThread();
352 +
353 + available = 0;
354 + ok = PeekNamedPipe(pipe, NULL, 0, NULL, &available, NULL);
355 + if (!ok) {
356 + DWORD err = GetLastError();
357 + if (is_disconnect_error(err))
358 + return NIPC_NP_ERR_DISCONNECTED;
359 + return NIPC_NP_ERR_RECV;
360 + }
361 +
362 + if (available > 0) {
363 + if (readable_out)
364 + *readable_out = true;
365 + return NIPC_NP_OK;
366 + }
367 +
368 + if (GetTickCount64() >= deadline)
369 + return NIPC_NP_OK;
370 + }
371 + continue;
372 + }
373 +
374 + Sleep(1);
375 + }
376 +}
377 +
378 +/* ------------------------------------------------------------------ */
379 +/* In-flight message_id tracking (client side) */
380 +/* ------------------------------------------------------------------ */
381 +
382 +static int inflight_add(nipc_np_session_t *s, uint64_t id)
383 +{
384 + for (uint32_t i = 0; i < s->inflight_count; i++) {
385 + if (s->inflight_ids[i] == id)
386 + return -1; /* duplicate */
387 + }
388 + if (s->inflight_count >= s->inflight_capacity) {
389 + uint32_t new_cap = s->inflight_capacity ? s->inflight_capacity * 2 : 16;
390 + uint64_t *new_ids = realloc(s->inflight_ids, (size_t)new_cap * sizeof(uint64_t));
391 + if (!new_ids)
392 + return -2; /* allocation failure */
393 + s->inflight_ids = new_ids;
394 + s->inflight_capacity = new_cap;
395 + }
396 + s->inflight_ids[s->inflight_count++] = id;
397 + return 0;
398 +}
399 +
400 +static int inflight_remove(nipc_np_session_t *s, uint64_t id)
401 +{
402 + for (uint32_t i = 0; i < s->inflight_count; i++) {
403 + if (s->inflight_ids[i] == id) {
404 + s->inflight_ids[i] = s->inflight_ids[s->inflight_count - 1];
405 + s->inflight_count--;
406 + return 0;
407 + }
408 + }
409 + return -1; /* not found */
410 +}
411 +
412 +static void inflight_fail_all(nipc_np_session_t *s)
413 +{
414 + if (!s || s->role != NIPC_NP_ROLE_CLIENT)
415 + return;
416 +
417 + /* A broken session invalidates every in-flight request on it. */
418 + s->inflight_count = 0;
419 +}
420 +
421 +/* ------------------------------------------------------------------ */
422 +/* Handshake: client side */
423 +/* ------------------------------------------------------------------ */
424 +
425 +static nipc_np_error_t client_handshake(HANDLE pipe,
426 + const nipc_np_client_config_t *cfg,
427 + nipc_np_session_t *session)
428 +{
429 + uint8_t buf[128];
430 +
431 + uint32_t pkt_size = apply_default(cfg->packet_size, NIPC_NP_DEFAULT_PACKET_SIZE);
432 +
433 + /* Build HELLO payload */
434 + nipc_hello_t hello = {
435 + .layout_version = 1,
436 + .flags = 0,
437 + .supported_profiles = cfg->supported_profiles ? cfg->supported_profiles : NIPC_PROFILE_BASELINE,
438 + .preferred_profiles = cfg->preferred_profiles,
439 + .max_request_payload_bytes = apply_default(cfg->max_request_payload_bytes, NIPC_MAX_PAYLOAD_DEFAULT),
440 + .max_request_batch_items = apply_default(cfg->max_request_batch_items, NIPC_NP_DEFAULT_BATCH_ITEMS),
441 + .max_response_payload_bytes = apply_default(cfg->max_response_payload_bytes, NIPC_MAX_PAYLOAD_DEFAULT),
442 + .max_response_batch_items = apply_default(cfg->max_response_batch_items, NIPC_NP_DEFAULT_BATCH_ITEMS),
443 + .auth_token = cfg->auth_token,
444 + .packet_size = pkt_size,
445 + };
446 +
447 + uint8_t hello_buf[44];
448 + nipc_hello_encode(&hello, hello_buf, sizeof(hello_buf));
449 +
450 + /* Build outer CONTROL header */
451 + nipc_header_t hdr = {
452 + .magic = NIPC_MAGIC_MSG,
453 + .version = NIPC_VERSION,
454 + .header_len = NIPC_HEADER_LEN,
455 + .kind = NIPC_KIND_CONTROL,
456 + .flags = 0,
457 + .code = NIPC_CODE_HELLO,
458 + .transport_status = NIPC_STATUS_OK,
459 + .payload_len = sizeof(hello_buf),
460 + .item_count = 1,
461 + .message_id = 0,
462 + };
463 +
464 + nipc_header_encode(&hdr, buf, sizeof(buf));
465 + memcpy(buf + NIPC_HEADER_LEN, hello_buf, sizeof(hello_buf));
466 +
467 + /* Send HELLO */
468 + nipc_np_error_t err = raw_send(pipe, buf, NIPC_HEADER_LEN + sizeof(hello_buf));
469 + if (err != NIPC_NP_OK)
470 + return err;
471 +
472 + /* Receive HELLO_ACK */
473 + DWORD bytes_read = 0;
474 + int rc = raw_recv(pipe, buf, sizeof(buf), &bytes_read);
475 + if (rc <= 0)
476 + return NIPC_NP_ERR_RECV;
477 +
478 + /* Decode outer header */
479 + nipc_header_t ack_hdr;
480 + nipc_error_t perr = nipc_header_decode(buf, (size_t)bytes_read, &ack_hdr);
481 + if (perr == NIPC_ERR_BAD_VERSION)
482 + return NIPC_NP_ERR_INCOMPATIBLE;
483 + if (perr != NIPC_OK)
484 + return NIPC_NP_ERR_PROTOCOL;
485 +
486 + if (ack_hdr.kind != NIPC_KIND_CONTROL || ack_hdr.code != NIPC_CODE_HELLO_ACK)
487 + return NIPC_NP_ERR_PROTOCOL;
488 +
489 + /* Check transport_status for rejection */
490 + if (ack_hdr.transport_status == NIPC_STATUS_AUTH_FAILED)
491 + return NIPC_NP_ERR_AUTH_FAILED;
492 + if (ack_hdr.transport_status == NIPC_STATUS_UNSUPPORTED)
493 + return NIPC_NP_ERR_NO_PROFILE;
494 + if (ack_hdr.transport_status == NIPC_STATUS_INCOMPATIBLE)
495 + return NIPC_NP_ERR_INCOMPATIBLE;
496 + if (ack_hdr.transport_status == NIPC_STATUS_LIMIT_EXCEEDED)
497 + return NIPC_NP_ERR_LIMIT_EXCEEDED;
498 + if (ack_hdr.transport_status != NIPC_STATUS_OK)
499 + return NIPC_NP_ERR_HANDSHAKE;
500 +
501 + /* Decode hello-ack payload */
502 + nipc_hello_ack_t ack;
503 + perr = nipc_hello_ack_decode(buf + NIPC_HEADER_LEN,
504 + (size_t)bytes_read - NIPC_HEADER_LEN, &ack);
505 + if (perr == NIPC_ERR_BAD_LAYOUT &&
506 + hello_ack_layout_incompatible(buf + NIPC_HEADER_LEN,
507 + (size_t)bytes_read - NIPC_HEADER_LEN))
508 + return NIPC_NP_ERR_INCOMPATIBLE;
509 + if (perr != NIPC_OK)
510 + return NIPC_NP_ERR_PROTOCOL;
511 +
512 + /* Fill session */
513 + session->pipe = pipe;
514 + session->role = NIPC_NP_ROLE_CLIENT;
515 + session->max_request_payload_bytes = ack.agreed_max_request_payload_bytes;
516 + session->max_request_batch_items = ack.agreed_max_request_batch_items;
517 + session->max_response_payload_bytes = ack.agreed_max_response_payload_bytes;
518 + session->max_response_batch_items = ack.agreed_max_response_batch_items;
519 + session->packet_size = ack.agreed_packet_size;
520 + session->selected_profile = ack.selected_profile;
521 + session->session_id = ack.session_id;
522 + session->recv_buf = NULL;
523 + session->recv_buf_size = 0;
524 + session->inflight_count = 0;
525 +
526 + return NIPC_NP_OK;
527 +}
528 +
529 +/* ------------------------------------------------------------------ */
530 +/* Handshake: server side */
531 +/* ------------------------------------------------------------------ */
532 +
533 +static nipc_np_error_t server_handshake(HANDLE pipe,
534 + const nipc_np_server_config_t *cfg,
535 + uint64_t session_id,
536 + nipc_np_session_t *session)
537 +{
538 + uint8_t buf[128];
539 +
540 + uint32_t server_pkt_size = apply_default(cfg->packet_size, NIPC_NP_DEFAULT_PACKET_SIZE);
541 + uint32_t s_req_pay = apply_default(cfg->max_request_payload_bytes, NIPC_MAX_PAYLOAD_DEFAULT);
542 + uint32_t s_req_bat = apply_default(cfg->max_request_batch_items, NIPC_NP_DEFAULT_BATCH_ITEMS);
543 + uint32_t s_resp_pay = apply_default(cfg->max_response_payload_bytes, NIPC_MAX_PAYLOAD_DEFAULT);
544 + uint32_t s_resp_bat = apply_default(cfg->max_response_batch_items, NIPC_NP_DEFAULT_BATCH_ITEMS);
545 + uint32_t s_profiles = cfg->supported_profiles ? cfg->supported_profiles : NIPC_PROFILE_BASELINE;
546 + uint32_t s_preferred = cfg->preferred_profiles;
547 +
548 + /* Receive HELLO */
549 + DWORD bytes_read = 0;
550 + int rc = raw_recv(pipe, buf, sizeof(buf), &bytes_read);
551 + if (rc <= 0)
552 + return NIPC_NP_ERR_RECV;
553 +
554 + nipc_header_t hdr;
555 + nipc_error_t perr = nipc_header_decode(buf, (size_t)bytes_read, &hdr);
556 + if (perr == NIPC_ERR_BAD_VERSION &&
557 + header_version_incompatible(buf, (size_t)bytes_read, NIPC_CODE_HELLO)) {
558 + send_rejection_ack(pipe, NIPC_STATUS_INCOMPATIBLE);
559 + return NIPC_NP_ERR_INCOMPATIBLE;
560 + }
561 + if (perr != NIPC_OK)
562 + return NIPC_NP_ERR_PROTOCOL;
563 +
564 + if (hdr.kind != NIPC_KIND_CONTROL || hdr.code != NIPC_CODE_HELLO)
565 + return NIPC_NP_ERR_PROTOCOL;
566 +
567 + nipc_hello_t hello;
568 + perr = nipc_hello_decode(buf + NIPC_HEADER_LEN,
569 + (size_t)bytes_read - NIPC_HEADER_LEN, &hello);
570 + if (perr == NIPC_ERR_BAD_LAYOUT &&
571 + hello_layout_incompatible(buf + NIPC_HEADER_LEN,
572 + (size_t)bytes_read - NIPC_HEADER_LEN)) {
573 + send_rejection_ack(pipe, NIPC_STATUS_INCOMPATIBLE);
574 + return NIPC_NP_ERR_INCOMPATIBLE;
575 + }
576 + if (perr != NIPC_OK)
577 + return NIPC_NP_ERR_PROTOCOL;
578 +
579 + /* Compute intersection */
580 + uint32_t intersection = hello.supported_profiles & s_profiles;
581 +
582 + /* Check intersection */
583 + if (intersection == 0) {
584 + send_rejection_ack(pipe, NIPC_STATUS_UNSUPPORTED);
585 + return NIPC_NP_ERR_NO_PROFILE;
586 + }
587 +
588 + /* Check auth */
589 + if (hello.auth_token != cfg->auth_token) {
590 + send_rejection_ack(pipe, NIPC_STATUS_AUTH_FAILED);
591 + return NIPC_NP_ERR_AUTH_FAILED;
592 + }
593 +
594 + /* Select profile */
595 + uint32_t preferred_intersection = intersection &
596 + hello.preferred_profiles & s_preferred;
597 + uint32_t selected;
598 + if (preferred_intersection != 0)
599 + selected = highest_bit(preferred_intersection);
600 + else
601 + selected = highest_bit(intersection);
602 +
603 + if (hello.max_request_payload_bytes > NIPC_MAX_PAYLOAD_CAP) {
604 + send_rejection_ack(pipe, NIPC_STATUS_LIMIT_EXCEEDED);
605 + return NIPC_NP_ERR_LIMIT_EXCEEDED;
606 + }
607 +
608 + /* Negotiate limits:
609 + * - request payload and batch size are client-proposed and echoed
610 + * - response payload is server-authoritative
611 + * - response batch size is symmetric with request batch size */
612 + uint32_t agreed_req_pay = hello.max_request_payload_bytes;
613 + uint32_t agreed_req_bat = hello.max_request_batch_items;
614 + uint32_t agreed_resp_pay = s_resp_pay;
615 + uint32_t agreed_resp_bat = agreed_req_bat;
616 + uint32_t agreed_pkt = min_u32(hello.packet_size, server_pkt_size);
617 +
618 + /* Reject packet sizes too small for a usable message packet */
619 + if (agreed_pkt <= NIPC_HEADER_LEN) {
620 + send_rejection_ack(pipe, NIPC_STATUS_INCOMPATIBLE);
621 + return NIPC_NP_ERR_INCOMPATIBLE;
622 + }
623 +
624 + /* Send HELLO_ACK (success) */
625 + nipc_hello_ack_t ack = {
626 + .layout_version = 1,
627 + .flags = 0,
628 + .server_supported_profiles = s_profiles,
629 + .intersection_profiles = intersection,
630 + .selected_profile = selected,
631 + .agreed_max_request_payload_bytes = agreed_req_pay,
632 + .agreed_max_request_batch_items = agreed_req_bat,
633 + .agreed_max_response_payload_bytes = agreed_resp_pay,
634 + .agreed_max_response_batch_items = agreed_resp_bat,
635 + .agreed_packet_size = agreed_pkt,
636 + .session_id = session_id,
637 + };
638 +
639 + uint8_t ack_buf[48];
640 + nipc_hello_ack_encode(&ack, ack_buf, sizeof(ack_buf));
641 +
642 + nipc_header_t ack_hdr = {
643 + .magic = NIPC_MAGIC_MSG,
644 + .version = NIPC_VERSION,
645 + .header_len = NIPC_HEADER_LEN,
646 + .kind = NIPC_KIND_CONTROL,
647 + .flags = 0,
648 + .code = NIPC_CODE_HELLO_ACK,
649 + .transport_status = NIPC_STATUS_OK,
650 + .payload_len = sizeof(ack_buf),
651 + .item_count = 1,
652 + .message_id = 0,
653 + };
654 +
655 + uint8_t pkt[80];
656 + nipc_header_encode(&ack_hdr, pkt, sizeof(pkt));
657 + memcpy(pkt + NIPC_HEADER_LEN, ack_buf, sizeof(ack_buf));
658 +
659 + nipc_np_error_t send_err = raw_send(pipe, pkt, NIPC_HEADER_LEN + sizeof(ack_buf));
660 + if (send_err != NIPC_NP_OK)
661 + return send_err;
662 +
663 + /* Fill session */
664 + session->pipe = pipe;
665 + session->role = NIPC_NP_ROLE_SERVER;
666 + session->max_request_payload_bytes = agreed_req_pay;
667 + session->max_request_batch_items = agreed_req_bat;
668 + session->max_response_payload_bytes = agreed_resp_pay;
669 + session->max_response_batch_items = agreed_resp_bat;
670 + session->packet_size = agreed_pkt;
671 + session->selected_profile = selected;
672 + session->session_id = session_id;
673 + session->recv_buf = NULL;
674 + session->recv_buf_size = 0;
675 + session->inflight_count = 0;
676 +
677 + return NIPC_NP_OK;
678 +}
679 +
680 +/* ------------------------------------------------------------------ */
681 +/* Public API: listen */
682 +/* ------------------------------------------------------------------ */
683 +
684 +nipc_np_error_t nipc_np_listen(const char *run_dir,
685 + const char *service_name,
686 + const nipc_np_server_config_t *config,
687 + nipc_np_listener_t *out)
688 +{
689 + if (!config || !out)
690 + return NIPC_NP_ERR_BAD_PARAM;
691 +
692 + memset(out, 0, sizeof(*out));
693 + out->pipe = INVALID_HANDLE_VALUE;
694 +
695 + /* Build pipe name */
696 + if (nipc_np_build_pipe_name(out->pipe_name, NIPC_NP_MAX_PIPE_NAME,
697 + run_dir, service_name) < 0)
698 + return NIPC_NP_ERR_BAD_PARAM;
699 +
700 + /* Determine buffer size */
701 + uint32_t buf_size = pipe_buffer_size(config->packet_size);
702 +
703 + /* Create first pipe instance with FILE_FLAG_FIRST_PIPE_INSTANCE
704 + * to detect if another server already owns this pipe name. */
705 + HANDLE h = create_pipe_instance(out->pipe_name, buf_size, TRUE);
706 + if (h == INVALID_HANDLE_VALUE) {
707 + DWORD err = GetLastError();
708 + if (err == ERROR_ACCESS_DENIED || err == ERROR_PIPE_BUSY)
709 + return NIPC_NP_ERR_ADDR_IN_USE;
710 + return NIPC_NP_ERR_CREATE_PIPE;
711 + }
712 +
713 + out->pipe = h;
714 + out->config = *config;
715 +
716 + return NIPC_NP_OK;
717 +}
718 +
719 +/* ------------------------------------------------------------------ */
720 +/* Public API: accept */
721 +/* ------------------------------------------------------------------ */
722 +
723 +nipc_np_error_t nipc_np_accept(nipc_np_listener_t *listener,
724 + uint64_t session_id,
725 + nipc_np_session_t *out)
726 +{
727 + if (!listener || !out)
728 + return NIPC_NP_ERR_BAD_PARAM;
729 + if (listener->pipe == INVALID_HANDLE_VALUE || listener->pipe == NULL)
730 + return NIPC_NP_ERR_ACCEPT;
731 +
732 + memset(out, 0, sizeof(*out));
733 + out->pipe = INVALID_HANDLE_VALUE;
734 +
735 + /* Wait for client to connect to the current pipe instance */
736 + BOOL connected = ConnectNamedPipe(listener->pipe, NULL);
737 + if (!connected) {
738 + DWORD err = GetLastError();
739 + /* ERROR_PIPE_CONNECTED means a client connected between
740 + * CreateNamedPipe and ConnectNamedPipe — that's fine. */
741 + if (err != ERROR_PIPE_CONNECTED) {
742 + return NIPC_NP_ERR_ACCEPT;
743 + }
744 + }
745 +
746 + /* The listener's current pipe instance is now the session pipe. */
747 + HANDLE session_pipe = listener->pipe;
748 +
749 + /* Create a new pipe instance for the next client. */
750 + uint32_t buf_size = pipe_buffer_size(listener->config.packet_size);
751 + HANDLE next = create_pipe_instance(listener->pipe_name, buf_size, FALSE);
752 + if (next == INVALID_HANDLE_VALUE) {
753 + /* Cannot create next instance — close the session pipe too. */
754 + DisconnectNamedPipe(session_pipe);
755 + CloseHandle(session_pipe);
756 + return NIPC_NP_ERR_CREATE_PIPE;
757 + }
758 + listener->pipe = next;
759 +
760 + /* Perform handshake */
761 + nipc_np_error_t herr = server_handshake(session_pipe, &listener->config,
762 + session_id, out);
763 + if (herr != NIPC_NP_OK) {
764 + DisconnectNamedPipe(session_pipe);
765 + CloseHandle(session_pipe);
766 + out->pipe = INVALID_HANDLE_VALUE;
767 + return herr;
768 + }
769 +
770 + return NIPC_NP_OK;
771 +}
772 +
773 +/* ------------------------------------------------------------------ */
774 +/* Public API: connect */
775 +/* ------------------------------------------------------------------ */
776 +
777 +nipc_np_error_t nipc_np_connect(const char *run_dir,
778 + const char *service_name,
779 + const nipc_np_client_config_t *config,
780 + nipc_np_session_t *out)
781 +{
782 + if (!config || !out)
783 + return NIPC_NP_ERR_BAD_PARAM;
784 +
785 + memset(out, 0, sizeof(*out));
786 + out->pipe = INVALID_HANDLE_VALUE;
787 +
788 + /* Build pipe name */
789 + wchar_t pipe_name[NIPC_NP_MAX_PIPE_NAME];
790 + if (nipc_np_build_pipe_name(pipe_name, NIPC_NP_MAX_PIPE_NAME,
791 + run_dir, service_name) < 0)
792 + return NIPC_NP_ERR_BAD_PARAM;
793 +
794 + /* Connect to the server pipe */
795 + HANDLE h = CreateFileW(
796 + pipe_name,
797 + GENERIC_READ | GENERIC_WRITE,
798 + 0, /* no sharing */
799 + NULL, /* default security */
800 + OPEN_EXISTING,
801 + 0, /* default attributes */
802 + NULL /* no template */
803 + );
804 +
805 + if (h == INVALID_HANDLE_VALUE)
806 + return NIPC_NP_ERR_CONNECT;
807 +
808 + /* Set read mode to message mode */
809 + DWORD mode = PIPE_READMODE_MESSAGE;
810 + if (!SetNamedPipeHandleState(h, &mode, NULL, NULL)) {
811 + CloseHandle(h);
812 + return NIPC_NP_ERR_CONNECT;
813 + }
814 +
815 + /* Perform handshake */
816 + nipc_np_error_t err = client_handshake(h, config, out);
817 + if (err != NIPC_NP_OK) {
818 + CloseHandle(h);
819 + out->pipe = INVALID_HANDLE_VALUE;
820 + return err;
821 + }
822 +
823 + return NIPC_NP_OK;
824 +}
825 +
826 +/* ------------------------------------------------------------------ */
827 +/* Public API: close */
828 +/* ------------------------------------------------------------------ */
829 +
830 +void nipc_np_close_session(nipc_np_session_t *session)
831 +{
832 + if (!session)
833 + return;
834 +
835 + if (session->pipe != INVALID_HANDLE_VALUE && session->pipe != NULL) {
836 + /* Only the server side should flush before disconnecting.
837 + * Client-side flush on a broken session can block shutdown. */
838 + if (session->role == NIPC_NP_ROLE_SERVER) {
839 + FlushFileBuffers(session->pipe);
840 + DisconnectNamedPipe(session->pipe);
841 + }
842 + CloseHandle(session->pipe);
843 + }
844 + session->pipe = INVALID_HANDLE_VALUE;
845 +
846 + free(session->recv_buf);
847 + session->recv_buf = NULL;
848 + session->recv_buf_size = 0;
849 +
850 + free(session->inflight_ids);
851 + session->inflight_ids = NULL;
852 + session->inflight_count = 0;
853 + session->inflight_capacity = 0;
854 +}
855 +
856 +void nipc_np_close_listener(nipc_np_listener_t *listener)
857 +{
858 + if (!listener)
859 + return;
860 +
861 + if (listener->pipe != INVALID_HANDLE_VALUE && listener->pipe != NULL) {
862 + /* A loopback connect reliably wakes a blocking ConnectNamedPipe()
863 + * so the accept thread can observe shutdown and exit. */
864 + if (listener->pipe_name[0] != L'\0') {
865 + HANDLE wake = CreateFileW(
866 + listener->pipe_name,
867 + GENERIC_READ | GENERIC_WRITE,
868 + 0,
869 + NULL,
870 + OPEN_EXISTING,
871 + 0,
872 + NULL);
873 + if (wake != INVALID_HANDLE_VALUE)
874 + CloseHandle(wake);
875 + }
876 + CloseHandle(listener->pipe);
877 + }
878 + listener->pipe = INVALID_HANDLE_VALUE;
879 +}
880 +
881 +/* ------------------------------------------------------------------ */
882 +/* Public API: send */
883 +/* ------------------------------------------------------------------ */
884 +
885 +nipc_np_error_t nipc_np_send(nipc_np_session_t *session,
886 + nipc_header_t *hdr,
887 + const void *payload,
888 + size_t payload_len)
889 +{
890 + if (!session || session->pipe == INVALID_HANDLE_VALUE)
891 + return NIPC_NP_ERR_BAD_PARAM;
892 +
893 + int tracked = (session->role == NIPC_NP_ROLE_CLIENT &&
894 + hdr->kind == NIPC_KIND_REQUEST);
895 +
896 + /* Client-side: track in-flight message_ids for requests */
897 + if (tracked) {
898 + int rc = inflight_add(session, hdr->message_id);
899 + if (rc == -1)
900 + return NIPC_NP_ERR_DUPLICATE_MSG_ID;
901 + if (rc == -2)
902 + return NIPC_NP_ERR_LIMIT_EXCEEDED;
903 + }
904 +
905 + uint32_t max_payload = 0;
906 + uint32_t max_batch = 0;
907 + if (session->role == NIPC_NP_ROLE_CLIENT &&
908 + hdr->kind == NIPC_KIND_REQUEST) {
909 + max_payload = session->max_request_payload_bytes;
910 + max_batch = session->max_request_batch_items;
911 + } else if (session->role == NIPC_NP_ROLE_SERVER &&
912 + hdr->kind == NIPC_KIND_RESPONSE) {
913 + max_payload = session->max_response_payload_bytes;
914 + max_batch = session->max_response_batch_items;
915 + }
916 + if (payload_len > UINT32_MAX ||
917 + (max_payload > 0 && payload_len > max_payload) ||
918 + (max_batch > 0 && hdr->item_count > max_batch)) {
919 + if (tracked)
920 + inflight_remove(session, hdr->message_id);
921 + return NIPC_NP_ERR_LIMIT_EXCEEDED;
922 + }
923 +
924 + /* Fill envelope fields */
925 + hdr->magic = NIPC_MAGIC_MSG;
926 + hdr->version = NIPC_VERSION;
927 + hdr->header_len = NIPC_HEADER_LEN;
928 + hdr->payload_len = (uint32_t)payload_len;
929 +
930 + size_t total_msg = NIPC_HEADER_LEN + payload_len;
931 +
932 + /* Single packet? */
933 + if (total_msg <= session->packet_size) {
934 + uint8_t hdr_buf[NIPC_HEADER_LEN];
935 + nipc_header_encode(hdr, hdr_buf, sizeof(hdr_buf));
936 + nipc_np_error_t err = raw_send_msg(session->pipe, hdr_buf,
937 + NIPC_HEADER_LEN, payload, payload_len);
938 + if (err != NIPC_NP_OK && tracked) {
939 + if (err == NIPC_NP_ERR_SEND || err == NIPC_NP_ERR_DISCONNECTED)
940 + inflight_fail_all(session);
941 + else
942 + inflight_remove(session, hdr->message_id);
943 + }
944 + return err;
945 + }
946 +
947 + /* Chunked send */
948 + size_t chunk_payload_budget = session->packet_size - NIPC_HEADER_LEN;
949 + if (chunk_payload_budget == 0) {
950 + if (tracked) inflight_remove(session, hdr->message_id);
951 + return NIPC_NP_ERR_BAD_PARAM;
952 + }
953 +
954 + size_t remaining = payload_len;
955 + size_t first_chunk_payload = remaining < chunk_payload_budget
956 + ? remaining : chunk_payload_budget;
957 +
958 + remaining -= first_chunk_payload;
959 + uint32_t continuation_chunks = 0;
960 + if (remaining > 0) {
961 + continuation_chunks = (uint32_t)((remaining + chunk_payload_budget - 1)
962 + / chunk_payload_budget);
963 + }
964 + uint32_t chunk_count = 1 + continuation_chunks;
965 +
966 + /* Send first chunk */
967 + uint8_t hdr_buf[NIPC_HEADER_LEN];
968 + nipc_header_encode(hdr, hdr_buf, sizeof(hdr_buf));
969 +
970 + nipc_np_error_t err = raw_send_msg(session->pipe, hdr_buf, NIPC_HEADER_LEN,
971 + payload, first_chunk_payload);
972 + if (err != NIPC_NP_OK) {
973 + if (tracked) {
974 + if (err == NIPC_NP_ERR_SEND || err == NIPC_NP_ERR_DISCONNECTED)
975 + inflight_fail_all(session);
976 + else
977 + inflight_remove(session, hdr->message_id);
978 + }
979 + return err;
980 + }
981 +
982 + /* Send continuation chunks */
983 + const uint8_t *src = (const uint8_t *)payload + first_chunk_payload;
984 + remaining = payload_len - first_chunk_payload;
985 +
986 + for (uint32_t ci = 1; ci < chunk_count; ci++) {
987 + size_t this_chunk = remaining < chunk_payload_budget
988 + ? remaining : chunk_payload_budget;
989 +
990 + nipc_chunk_header_t chk = {
991 + .magic = NIPC_MAGIC_CHUNK,
992 + .version = NIPC_VERSION,
993 + .flags = 0,
994 + .message_id = hdr->message_id,
995 + .total_message_len = (uint32_t)total_msg,
996 + .chunk_index = ci,
997 + .chunk_count = chunk_count,
998 + .chunk_payload_len = (uint32_t)this_chunk,
999 + };
1000 +
1001 + uint8_t chk_buf[NIPC_HEADER_LEN];
1002 + nipc_chunk_header_encode(&chk, chk_buf, sizeof(chk_buf));
1003 +
1004 + err = raw_send_msg(session->pipe, chk_buf, NIPC_HEADER_LEN,
1005 + src, this_chunk);
1006 + if (err != NIPC_NP_OK) {
1007 + if (tracked) {
1008 + if (err == NIPC_NP_ERR_SEND || err == NIPC_NP_ERR_DISCONNECTED)
1009 + inflight_fail_all(session);
1010 + else
1011 + inflight_remove(session, hdr->message_id);
1012 + }
1013 + return err;
1014 + }
1015 +
1016 + src += this_chunk;
1017 + remaining -= this_chunk;
1018 + }
1019 +
1020 + return NIPC_NP_OK;
1021 +}
1022 +
1023 +/* ------------------------------------------------------------------ */
1024 +/* Public API: receive */
1025 +/* ------------------------------------------------------------------ */
1026 +
1027 +static nipc_np_error_t ensure_recv_buf(nipc_np_session_t *session,
1028 + size_t needed)
1029 +{
1030 + if (session->recv_buf_size >= needed)
1031 + return NIPC_NP_OK;
1032 +
1033 + uint8_t *p = (uint8_t *)realloc(session->recv_buf, needed);
1034 + if (!p)
1035 + return NIPC_NP_ERR_ALLOC;
1036 +
1037 + session->recv_buf = p;
1038 + session->recv_buf_size = needed;
1039 + return NIPC_NP_OK;
1040 +}
1041 +
1042 +/* Validate batch directory if the message has BATCH flag and item_count > 1.
1043 + * Called after the full payload is assembled, before returning to caller. */
1044 +static nipc_np_error_t validate_batch(const nipc_header_t *hdr,
1045 + const void *payload, size_t payload_len)
1046 +{
1047 + if (!(hdr->flags & NIPC_FLAG_BATCH) || hdr->item_count <= 1)
1048 + return NIPC_NP_OK;
1049 +
1050 + uint32_t dir_bytes = hdr->item_count * 8;
1051 + uint32_t dir_aligned = (uint32_t)nipc_align8(dir_bytes);
1052 + if (payload_len < dir_aligned)
1053 + return NIPC_NP_ERR_PROTOCOL;
1054 +
1055 + uint32_t packed_area_len = (uint32_t)(payload_len - dir_aligned);
1056 + nipc_error_t perr = nipc_batch_dir_validate(payload, dir_bytes,
1057 + hdr->item_count,
1058 + packed_area_len);
1059 + return (perr == NIPC_OK) ? NIPC_NP_OK : NIPC_NP_ERR_PROTOCOL;
1060 +}
1061 +
1062 +nipc_np_error_t nipc_np_receive(nipc_np_session_t *session,
1063 + void *buf, size_t buf_size,
1064 + nipc_header_t *hdr_out,
1065 + const void **payload_out,
1066 + size_t *payload_len_out)
1067 +{
1068 + if (!session || session->pipe == INVALID_HANDLE_VALUE)
1069 + return NIPC_NP_ERR_BAD_PARAM;
1070 +
1071 + /* Read first message */
1072 + DWORD bytes_read = 0;
1073 + int rc = raw_recv(session->pipe, buf, buf_size, &bytes_read);
1074 + if (rc <= 0) {
1075 + inflight_fail_all(session);
1076 + return NIPC_NP_ERR_RECV;
1077 + }
1078 +
1079 + size_t n = (size_t)bytes_read;
1080 + if (n < NIPC_HEADER_LEN)
1081 + return NIPC_NP_ERR_PROTOCOL;
1082 +
1083 + /* Decode outer header */
1084 + nipc_error_t perr = nipc_header_decode(buf, n, hdr_out);
1085 + if (perr != NIPC_OK)
1086 + return NIPC_NP_ERR_PROTOCOL;
1087 +
1088 + /* Validate payload_len against negotiated directional limit */
1089 + uint32_t max_payload = (session->role == NIPC_NP_ROLE_SERVER)
1090 + ? session->max_request_payload_bytes
1091 + : session->max_response_payload_bytes;
1092 + if (hdr_out->payload_len > max_payload)
1093 + return NIPC_NP_ERR_LIMIT_EXCEEDED;
1094 +
1095 + /* Validate item_count */
1096 + uint32_t max_batch = (session->role == NIPC_NP_ROLE_SERVER)
1097 + ? session->max_request_batch_items
1098 + : session->max_response_batch_items;
1099 + if (hdr_out->item_count > max_batch)
1100 + return NIPC_NP_ERR_LIMIT_EXCEEDED;
1101 +
1102 + /* Client-side: validate response message_id is in-flight */
1103 + if (session->role == NIPC_NP_ROLE_CLIENT &&
1104 + hdr_out->kind == NIPC_KIND_RESPONSE) {
1105 + if (inflight_remove(session, hdr_out->message_id) < 0)
1106 + return NIPC_NP_ERR_UNKNOWN_MSG_ID;
1107 + }
1108 +
1109 + size_t total_msg = NIPC_HEADER_LEN + hdr_out->payload_len;
1110 +
1111 + /* Non-chunked: entire message in one read */
1112 + if (n >= total_msg) {
1113 + *payload_out = (const uint8_t *)buf + NIPC_HEADER_LEN;
1114 + *payload_len_out = hdr_out->payload_len;
1115 +
1116 + nipc_np_error_t berr = validate_batch(hdr_out, *payload_out, *payload_len_out);
1117 + if (berr != NIPC_NP_OK)
1118 + return berr;
1119 +
1120 + return NIPC_NP_OK;
1121 + }
1122 +
1123 + /* Chunked: first message has partial payload */
1124 + size_t first_payload_bytes = n - NIPC_HEADER_LEN;
1125 +
1126 + nipc_np_error_t err = ensure_recv_buf(session, hdr_out->payload_len);
1127 + if (err != NIPC_NP_OK)
1128 + return err;
1129 +
1130 + memcpy(session->recv_buf, (uint8_t *)buf + NIPC_HEADER_LEN,
1131 + first_payload_bytes);
1132 +
1133 + size_t assembled = first_payload_bytes;
1134 + size_t chunk_payload_budget = session->packet_size - NIPC_HEADER_LEN;
1135 +
1136 + /* Calculate expected chunk count */
1137 + size_t remaining_after_first = hdr_out->payload_len - first_payload_bytes;
1138 + uint32_t expected_continuations = 0;
1139 + if (remaining_after_first > 0 && chunk_payload_budget > 0) {
1140 + expected_continuations = (uint32_t)((remaining_after_first +
1141 + chunk_payload_budget - 1)
1142 + / chunk_payload_budget);
1143 + }
1144 + uint32_t expected_chunk_count = 1 + expected_continuations;
1145 +
1146 + /* Temporary buffer for continuation messages */
1147 + size_t pkt_buf_size = session->packet_size;
1148 + uint8_t *pkt_buf = (uint8_t *)malloc(pkt_buf_size);
1149 + if (!pkt_buf)
1150 + return NIPC_NP_ERR_ALLOC;
1151 +
1152 + for (uint32_t ci = 1; assembled < hdr_out->payload_len; ci++) {
1153 + DWORD cn = 0;
1154 + int crc = raw_recv(session->pipe, pkt_buf, pkt_buf_size, &cn);
1155 + if (crc <= 0) {
1156 + free(pkt_buf);
1157 + inflight_fail_all(session);
1158 + return NIPC_NP_ERR_RECV;
1159 + }
1160 +
1161 + if ((size_t)cn < NIPC_HEADER_LEN) {
1162 + free(pkt_buf);
1163 + return NIPC_NP_ERR_CHUNK;
1164 + }
1165 +
1166 + nipc_chunk_header_t chk;
1167 + perr = nipc_chunk_header_decode(pkt_buf, (size_t)cn, &chk);
1168 + if (perr != NIPC_OK) {
1169 + free(pkt_buf);
1170 + return NIPC_NP_ERR_CHUNK;
1171 + }
1172 +
1173 + /* Validate chunk header */
1174 + if (chk.message_id != hdr_out->message_id ||
1175 + chk.chunk_index != ci ||
1176 + chk.chunk_count != expected_chunk_count ||
1177 + chk.total_message_len != (uint32_t)total_msg) {
1178 + free(pkt_buf);
1179 + return NIPC_NP_ERR_CHUNK;
1180 + }
1181 +
1182 + size_t chunk_data = (size_t)cn - NIPC_HEADER_LEN;
1183 + if (chunk_data != chk.chunk_payload_len) {
1184 + free(pkt_buf);
1185 + return NIPC_NP_ERR_CHUNK;
1186 + }
1187 +
1188 + if (assembled + chunk_data > hdr_out->payload_len) {
1189 + free(pkt_buf);
1190 + return NIPC_NP_ERR_CHUNK;
1191 + }
1192 +
1193 + memcpy(session->recv_buf + assembled,
1194 + pkt_buf + NIPC_HEADER_LEN, chunk_data);
1195 + assembled += chunk_data;
1196 + }
1197 +
1198 + free(pkt_buf);
1199 +
1200 + *payload_out = session->recv_buf;
1201 + *payload_len_out = hdr_out->payload_len;
1202 +
1203 + nipc_np_error_t berr = validate_batch(hdr_out, *payload_out, *payload_len_out);
1204 + if (berr != NIPC_NP_OK)
1205 + return berr;
1206 +
1207 + return NIPC_NP_OK;
1208 +}
1209 +
1210 +nipc_np_error_t nipc_np_wait_readable(nipc_np_session_t *session,
1211 + uint32_t timeout_ms,
1212 + bool *readable_out)
1213 +{
1214 + if (!session || !readable_out || session->pipe == INVALID_HANDLE_VALUE)
1215 + return NIPC_NP_ERR_BAD_PARAM;
1216 + nipc_np_error_t err = wait_readable(session->pipe, timeout_ms, readable_out);
1217 + if (err == NIPC_NP_ERR_DISCONNECTED || err == NIPC_NP_ERR_RECV)
1218 + inflight_fail_all(session);
1219 + return err;
1220 +}
1221 +
1222 +#endif /* _WIN32 || __MSYS__ */
src/libnetdata/netipc/src/transport/windows/netipc_win_shm.c new
+948
@@ -0,0 +1,948 @@
1 +/*
2 + * netipc_win_shm.c - L1 Windows SHM transport.
3 + *
4 + * Shared memory data plane with spin + kernel event synchronization.
5 + * Uses CreateFileMappingW / MapViewOfFile for the shared region and
6 + * auto-reset kernel events for signaling (SHM_HYBRID profile).
7 + *
8 + * Wire-compatible with all language implementations.
9 + */
10 +
11 +#if defined(_WIN32) || defined(__MSYS__)
12 +
13 +#include "netipc/netipc_win_shm.h"
14 +#include "netipc/netipc_named_pipe.h" /* nipc_fnv1a_64 */
15 +
16 +#include <stdio.h>
17 +#include <stdlib.h>
18 +#include <string.h>
19 +
20 +/* ------------------------------------------------------------------ */
21 +/* Internal helpers */
22 +/* ------------------------------------------------------------------ */
23 +
24 +typedef enum {
25 + NIPC_WIN_SHM_TEST_FAULT_NONE = 0,
26 + NIPC_WIN_SHM_TEST_FAULT_CREATE_MAPPING,
27 + NIPC_WIN_SHM_TEST_FAULT_OPEN_MAPPING,
28 + NIPC_WIN_SHM_TEST_FAULT_MAP_VIEW,
29 + NIPC_WIN_SHM_TEST_FAULT_CREATE_EVENT,
30 + NIPC_WIN_SHM_TEST_FAULT_OPEN_EVENT,
31 +} nipc_win_shm_test_fault_site_t_internal;
32 +
33 +typedef struct {
34 + nipc_win_shm_test_fault_site_t_internal site;
35 + DWORD error_code;
36 + uint32_t skip_matches;
37 +} nipc_win_shm_test_fault_t;
38 +
39 +static nipc_win_shm_test_fault_t g_win_shm_test_fault = {
40 + .site = NIPC_WIN_SHM_TEST_FAULT_NONE,
41 + .error_code = ERROR_SUCCESS,
42 + .skip_matches = 0,
43 +};
44 +
45 +void nipc_win_shm_test_fault_set(int site,
46 + DWORD error_code,
47 + uint32_t skip_matches)
48 +{
49 + g_win_shm_test_fault.site = (nipc_win_shm_test_fault_site_t_internal)site;
50 + g_win_shm_test_fault.error_code = error_code;
51 + g_win_shm_test_fault.skip_matches = skip_matches;
52 +}
53 +
54 +void nipc_win_shm_test_fault_clear(void)
55 +{
56 + g_win_shm_test_fault.site = NIPC_WIN_SHM_TEST_FAULT_NONE;
57 + g_win_shm_test_fault.error_code = ERROR_SUCCESS;
58 + g_win_shm_test_fault.skip_matches = 0;
59 +}
60 +
61 +static int should_inject_fault(nipc_win_shm_test_fault_site_t_internal site)
62 +{
63 + if (g_win_shm_test_fault.site != site)
64 + return 0;
65 +
66 + if (g_win_shm_test_fault.skip_matches > 0) {
67 + g_win_shm_test_fault.skip_matches--;
68 + return 0;
69 + }
70 +
71 + SetLastError(g_win_shm_test_fault.error_code);
72 + g_win_shm_test_fault.site = NIPC_WIN_SHM_TEST_FAULT_NONE;
73 + return 1;
74 +}
75 +
76 +static HANDLE test_create_file_mapping(HANDLE file,
77 + LPSECURITY_ATTRIBUTES attrs,
78 + DWORD protect,
79 + DWORD size_high,
80 + DWORD size_low,
81 + LPCWSTR name)
82 +{
83 + if (should_inject_fault(NIPC_WIN_SHM_TEST_FAULT_CREATE_MAPPING))
84 + return NULL;
85 +
86 + return CreateFileMappingW(file, attrs, protect, size_high, size_low, name);
87 +}
88 +
89 +static HANDLE test_open_file_mapping(DWORD access,
90 + BOOL inherit,
91 + LPCWSTR name)
92 +{
93 + if (should_inject_fault(NIPC_WIN_SHM_TEST_FAULT_OPEN_MAPPING))
94 + return NULL;
95 +
96 + return OpenFileMappingW(access, inherit, name);
97 +}
98 +
99 +static void *test_map_view_of_file(HANDLE mapping,
100 + DWORD access,
101 + DWORD offset_high,
102 + DWORD offset_low,
103 + SIZE_T bytes)
104 +{
105 + if (should_inject_fault(NIPC_WIN_SHM_TEST_FAULT_MAP_VIEW))
106 + return NULL;
107 +
108 + return MapViewOfFile(mapping, access, offset_high, offset_low, bytes);
109 +}
110 +
111 +static HANDLE test_create_event(LPSECURITY_ATTRIBUTES attrs,
112 + BOOL manual_reset,
113 + BOOL initial_state,
114 + LPCWSTR name)
115 +{
116 + if (should_inject_fault(NIPC_WIN_SHM_TEST_FAULT_CREATE_EVENT))
117 + return NULL;
118 +
119 + return CreateEventW(attrs, manual_reset, initial_state, name);
120 +}
121 +
122 +static HANDLE test_open_event(DWORD access,
123 + BOOL inherit,
124 + LPCWSTR name)
125 +{
126 + if (should_inject_fault(NIPC_WIN_SHM_TEST_FAULT_OPEN_EVENT))
127 + return NULL;
128 +
129 + return OpenEventW(access, inherit, name);
130 +}
131 +
132 +/* Round up to 64-byte cache-line alignment. */
133 +static inline uint32_t align_cacheline(uint32_t v)
134 +{
135 + return (v + (NIPC_WIN_SHM_CACHELINE - 1)) & ~(uint32_t)(NIPC_WIN_SHM_CACHELINE - 1);
136 +}
137 +
138 +/* Atomic reads without write contention.
139 + * InterlockedCompareExchange64(ptr, 0, 0) emits LOCK CMPXCHG8B which writes
140 + * the cache line on every call — catastrophic in spin loops. MemoryBarrier()
141 + * emits MFENCE which flushes the store buffer — also expensive.
142 + *
143 + * On x86-64, aligned 64-bit reads are naturally atomic and have acquire
144 + * semantics (loads are never reordered with loads). A volatile read with
145 + * a compiler barrier (no hardware barrier) is sufficient and matches what
146 + * Go's atomic.LoadInt64 compiles to (plain MOV). */
147 +static inline LONG64 atomic_load_64(volatile LONG64 *ptr)
148 +{
149 + LONG64 val = *ptr;
150 + __asm__ volatile("" ::: "memory"); /* compiler barrier only */
151 + return val;
152 +}
153 +
154 +static inline LONG atomic_load_32(volatile LONG *ptr)
155 +{
156 + LONG val = *ptr;
157 + __asm__ volatile("" ::: "memory"); /* compiler barrier only */
158 + return val;
159 +}
160 +
161 +/* Validate service_name: only [a-zA-Z0-9._-], non-empty, not "." or "..". */
162 +static int validate_service_name(const char *name)
163 +{
164 + if (!name || name[0] == '\0')
165 + return -1;
166 +
167 + if (name[0] == '.' && (name[1] == '\0' || (name[1] == '.' && name[2] == '\0')))
168 + return -1;
169 +
170 + for (const char *p = name; *p; p++) {
171 + char c = *p;
172 + if ((c >= 'a' && c <= 'z') || (c >= 'A' && c <= 'Z') ||
173 + (c >= '0' && c <= '9') || c == '.' || c == '_' || c == '-')
174 + continue;
175 + return -1;
176 + }
177 + return 0;
178 +}
179 +
180 +/* CPU pause hint for spin loops. */
181 +static inline void cpu_relax(void)
182 +{
183 +#if defined(__x86_64__) || defined(__i386__) || defined(_M_X64) || defined(_M_IX86)
184 + YieldProcessor();
185 +#else
186 + /* generic fallback */
187 + MemoryBarrier();
188 +#endif
189 +}
190 +
191 +/* Pointer into the mapped region at a byte offset. */
192 +static inline void *region_ptr(const nipc_win_shm_ctx_t *ctx, uint32_t offset)
193 +{
194 + return (uint8_t *)ctx->base + offset;
195 +}
196 +
197 +/* Pointer to the header. */
198 +static inline nipc_win_shm_region_header_t *region_hdr(const nipc_win_shm_ctx_t *ctx)
199 +{
200 + return (nipc_win_shm_region_header_t *)ctx->base;
201 +}
202 +
203 +/* ------------------------------------------------------------------ */
204 +/* Kernel object name derivation */
205 +/* ------------------------------------------------------------------ */
206 +
207 +/*
208 + * Compute FNV-1a hash for SHM object names.
209 + * Input: run_dir + "\n" + service_name + "\n" + auth_token_decimal
210 + */
211 +static uint64_t compute_shm_hash(const char *run_dir,
212 + const char *service_name,
213 + uint64_t auth_token)
214 +{
215 + char buf[512];
216 + int n = snprintf(buf, sizeof(buf), "%s\n%s\n%llu",
217 + run_dir, service_name,
218 + (unsigned long long)auth_token);
219 + if (n < 0 || (size_t)n >= sizeof(buf))
220 + return 0;
221 +
222 + return nipc_fnv1a_64(buf, (size_t)n);
223 +}
224 +
225 +/*
226 + * Build a kernel object name.
227 + * Format: Local\netipc-{hash:016llx}-{service}-p{profile}-s{session_id:016llx}-{suffix}
228 + */
229 +static int build_object_name(wchar_t *dst, size_t dst_chars,
230 + uint64_t hash,
231 + const char *service_name,
232 + uint32_t profile,
233 + uint64_t session_id,
234 + const char *suffix)
235 +{
236 + char narrow[NIPC_WIN_SHM_MAX_NAME];
237 + int n = snprintf(narrow, sizeof(narrow),
238 + "Local\\netipc-%016llx-%s-p%u-s%016llx-%s",
239 + (unsigned long long)hash, service_name,
240 + (unsigned)profile,
241 + (unsigned long long)session_id, suffix);
242 + if (n < 0 || (size_t)n >= sizeof(narrow))
243 + return -1;
244 +
245 + if ((size_t)(n + 1) > dst_chars)
246 + return -1;
247 +
248 + for (int i = 0; i <= n; i++)
249 + dst[i] = (wchar_t)(unsigned char)narrow[i];
250 +
251 + return 0;
252 +}
253 +
254 +/* ------------------------------------------------------------------ */
255 +/* Server: create */
256 +/* ------------------------------------------------------------------ */
257 +
258 +nipc_win_shm_error_t nipc_win_shm_server_create(
259 + const char *run_dir,
260 + const char *service_name,
261 + uint64_t auth_token,
262 + uint64_t session_id,
263 + uint32_t profile,
264 + uint32_t req_capacity,
265 + uint32_t resp_capacity,
266 + nipc_win_shm_ctx_t *ctx)
267 +{
268 + if (!ctx)
269 + return NIPC_WIN_SHM_ERR_BAD_PARAM;
270 +
271 + memset(ctx, 0, sizeof(*ctx));
272 + ctx->mapping = NULL;
273 + ctx->req_event = INVALID_HANDLE_VALUE;
274 + ctx->resp_event = INVALID_HANDLE_VALUE;
275 +
276 + if (!run_dir || !service_name)
277 + return NIPC_WIN_SHM_ERR_BAD_PARAM;
278 +
279 + if (validate_service_name(service_name) < 0)
280 + return NIPC_WIN_SHM_ERR_BAD_PARAM;
281 +
282 + if (profile != NIPC_WIN_SHM_PROFILE_HYBRID &&
283 + profile != NIPC_WIN_SHM_PROFILE_BUSYWAIT)
284 + return NIPC_WIN_SHM_ERR_BAD_PARAM;
285 +
286 + /* Compute hash and build object names */
287 + uint64_t hash = compute_shm_hash(run_dir, service_name, auth_token);
288 + if (hash == 0)
289 + return NIPC_WIN_SHM_ERR_BAD_PARAM;
290 +
291 + wchar_t mapping_name[NIPC_WIN_SHM_MAX_NAME];
292 + if (build_object_name(mapping_name, NIPC_WIN_SHM_MAX_NAME,
293 + hash, service_name, profile, session_id,
294 + "mapping") < 0)
295 + return NIPC_WIN_SHM_ERR_BAD_PARAM;
296 +
297 + /* Compute aligned offsets */
298 + req_capacity = align_cacheline(req_capacity);
299 + resp_capacity = align_cacheline(resp_capacity);
300 +
301 + uint32_t req_off = align_cacheline(NIPC_WIN_SHM_HEADER_LEN);
302 + /* Guard against uint32 overflow before computing resp_off */
303 + if (req_capacity > UINT32_MAX - req_off)
304 + return NIPC_WIN_SHM_ERR_BAD_PARAM;
305 + uint32_t resp_off = align_cacheline(req_off + req_capacity);
306 + size_t region_size = (size_t)resp_off + resp_capacity;
307 +
308 + /* Create file mapping backed by page file */
309 + SetLastError(ERROR_SUCCESS);
310 + HANDLE mapping = test_create_file_mapping(
311 + INVALID_HANDLE_VALUE, /* page file backed */
312 + NULL, /* default security */
313 + PAGE_READWRITE,
314 + (DWORD)(region_size >> 32),
315 + (DWORD)(region_size & 0xFFFFFFFF),
316 + mapping_name);
317 + if (!mapping)
318 + return NIPC_WIN_SHM_ERR_CREATE_MAPPING;
319 + if (GetLastError() == ERROR_ALREADY_EXISTS) {
320 + CloseHandle(mapping);
321 + return NIPC_WIN_SHM_ERR_ADDR_IN_USE;
322 + }
323 +
324 + /* Map the view */
325 + void *base = test_map_view_of_file(mapping, FILE_MAP_ALL_ACCESS, 0, 0, region_size);
326 + if (!base) {
327 + CloseHandle(mapping);
328 + return NIPC_WIN_SHM_ERR_MAP_VIEW;
329 + }
330 +
331 + /* Zero the region */
332 + memset(base, 0, region_size);
333 +
334 + /* Write header */
335 + nipc_win_shm_region_header_t *hdr = (nipc_win_shm_region_header_t *)base;
336 + hdr->magic = NIPC_WIN_SHM_MAGIC;
337 + hdr->version = NIPC_WIN_SHM_VERSION;
338 + hdr->header_len = NIPC_WIN_SHM_HEADER_LEN;
339 + hdr->profile = profile;
340 + hdr->request_offset = req_off;
341 + hdr->request_capacity = req_capacity;
342 + hdr->response_offset = resp_off;
343 + hdr->response_capacity = resp_capacity;
344 + hdr->spin_tries = NIPC_WIN_SHM_DEFAULT_SPIN;
345 +
346 + /* Memory barrier to ensure header is visible */
347 + MemoryBarrier();
348 +
349 + /* Create kernel events for HYBRID profile */
350 + HANDLE req_event = INVALID_HANDLE_VALUE;
351 + HANDLE resp_event = INVALID_HANDLE_VALUE;
352 +
353 + if (profile == NIPC_WIN_SHM_PROFILE_HYBRID) {
354 + wchar_t req_event_name[NIPC_WIN_SHM_MAX_NAME];
355 + wchar_t resp_event_name[NIPC_WIN_SHM_MAX_NAME];
356 +
357 + if (build_object_name(req_event_name, NIPC_WIN_SHM_MAX_NAME,
358 + hash, service_name, profile, session_id,
359 + "req_event") < 0 ||
360 + build_object_name(resp_event_name, NIPC_WIN_SHM_MAX_NAME,
361 + hash, service_name, profile, session_id,
362 + "resp_event") < 0) {
363 + UnmapViewOfFile(base);
364 + CloseHandle(mapping);
365 + return NIPC_WIN_SHM_ERR_BAD_PARAM;
366 + }
367 +
368 + /* Auto-reset events (bManualReset = FALSE) */
369 + SetLastError(ERROR_SUCCESS);
370 + req_event = test_create_event(NULL, FALSE, FALSE, req_event_name);
371 + if (!req_event) {
372 + UnmapViewOfFile(base);
373 + CloseHandle(mapping);
374 + return NIPC_WIN_SHM_ERR_CREATE_EVENT;
375 + }
376 + if (GetLastError() == ERROR_ALREADY_EXISTS) {
377 + CloseHandle(req_event);
378 + UnmapViewOfFile(base);
379 + CloseHandle(mapping);
380 + return NIPC_WIN_SHM_ERR_ADDR_IN_USE;
381 + }
382 +
383 + SetLastError(ERROR_SUCCESS);
384 + resp_event = test_create_event(NULL, FALSE, FALSE, resp_event_name);
385 + if (!resp_event) {
386 + CloseHandle(req_event);
387 + UnmapViewOfFile(base);
388 + CloseHandle(mapping);
389 + return NIPC_WIN_SHM_ERR_CREATE_EVENT;
390 + }
391 + if (GetLastError() == ERROR_ALREADY_EXISTS) {
392 + CloseHandle(resp_event);
393 + CloseHandle(req_event);
394 + UnmapViewOfFile(base);
395 + CloseHandle(mapping);
396 + return NIPC_WIN_SHM_ERR_ADDR_IN_USE;
397 + }
398 + }
399 +
400 + /* Fill context */
401 + ctx->role = NIPC_WIN_SHM_ROLE_SERVER;
402 + ctx->mapping = mapping;
403 + ctx->base = base;
404 + ctx->region_size = region_size;
405 + ctx->req_event = req_event;
406 + ctx->resp_event = resp_event;
407 + ctx->profile = profile;
408 + ctx->request_offset = req_off;
409 + ctx->request_capacity = req_capacity;
410 + ctx->response_offset = resp_off;
411 + ctx->response_capacity = resp_capacity;
412 + ctx->spin_tries = NIPC_WIN_SHM_DEFAULT_SPIN;
413 + ctx->local_req_seq = 0;
414 + ctx->local_resp_seq = 0;
415 +
416 + return NIPC_WIN_SHM_OK;
417 +}
418 +
419 +/* ------------------------------------------------------------------ */
420 +/* Server: destroy */
421 +/* ------------------------------------------------------------------ */
422 +
423 +void nipc_win_shm_destroy(nipc_win_shm_ctx_t *ctx)
424 +{
425 + if (!ctx)
426 + return;
427 +
428 + /* Set close flag and signal waiting client */
429 + if (ctx->base) {
430 + nipc_win_shm_region_header_t *hdr = region_hdr(ctx);
431 + InterlockedExchange(&hdr->resp_server_closed, 1);
432 + MemoryBarrier();
433 + }
434 +
435 + if (ctx->profile == NIPC_WIN_SHM_PROFILE_HYBRID &&
436 + ctx->resp_event != INVALID_HANDLE_VALUE)
437 + SetEvent(ctx->resp_event);
438 +
439 + if (ctx->base) {
440 + UnmapViewOfFile(ctx->base);
441 + ctx->base = NULL;
442 + }
443 +
444 + if (ctx->mapping) {
445 + CloseHandle(ctx->mapping);
446 + ctx->mapping = NULL;
447 + }
448 +
449 + if (ctx->req_event != INVALID_HANDLE_VALUE) {
450 + CloseHandle(ctx->req_event);
451 + ctx->req_event = INVALID_HANDLE_VALUE;
452 + }
453 +
454 + if (ctx->resp_event != INVALID_HANDLE_VALUE) {
455 + CloseHandle(ctx->resp_event);
456 + ctx->resp_event = INVALID_HANDLE_VALUE;
457 + }
458 +
459 + ctx->region_size = 0;
460 +}
461 +
462 +/* ------------------------------------------------------------------ */
463 +/* Client: attach */
464 +/* ------------------------------------------------------------------ */
465 +
466 +nipc_win_shm_error_t nipc_win_shm_client_attach(
467 + const char *run_dir,
468 + const char *service_name,
469 + uint64_t auth_token,
470 + uint64_t session_id,
471 + uint32_t profile,
472 + nipc_win_shm_ctx_t *ctx)
473 +{
474 + if (!ctx)
475 + return NIPC_WIN_SHM_ERR_BAD_PARAM;
476 +
477 + memset(ctx, 0, sizeof(*ctx));
478 + ctx->mapping = NULL;
479 + ctx->req_event = INVALID_HANDLE_VALUE;
480 + ctx->resp_event = INVALID_HANDLE_VALUE;
481 +
482 + if (!run_dir || !service_name)
483 + return NIPC_WIN_SHM_ERR_BAD_PARAM;
484 +
485 + if (validate_service_name(service_name) < 0)
486 + return NIPC_WIN_SHM_ERR_BAD_PARAM;
487 +
488 + if (profile != NIPC_WIN_SHM_PROFILE_HYBRID &&
489 + profile != NIPC_WIN_SHM_PROFILE_BUSYWAIT)
490 + return NIPC_WIN_SHM_ERR_BAD_PARAM;
491 +
492 + uint64_t hash = compute_shm_hash(run_dir, service_name, auth_token);
493 + if (hash == 0)
494 + return NIPC_WIN_SHM_ERR_BAD_PARAM;
495 +
496 + wchar_t mapping_name[NIPC_WIN_SHM_MAX_NAME];
497 + if (build_object_name(mapping_name, NIPC_WIN_SHM_MAX_NAME,
498 + hash, service_name, profile, session_id,
499 + "mapping") < 0)
500 + return NIPC_WIN_SHM_ERR_BAD_PARAM;
501 +
502 + /* Open existing file mapping */
503 + HANDLE mapping = test_open_file_mapping(FILE_MAP_ALL_ACCESS, FALSE, mapping_name);
504 + if (!mapping)
505 + return NIPC_WIN_SHM_ERR_OPEN_MAPPING;
506 +
507 + /* Map the view -- map header first to read region size */
508 + void *base = test_map_view_of_file(mapping, FILE_MAP_ALL_ACCESS, 0, 0, 0);
509 + if (!base) {
510 + CloseHandle(mapping);
511 + return NIPC_WIN_SHM_ERR_MAP_VIEW;
512 + }
513 +
514 + /* Acquire barrier to see server's header writes */
515 + MemoryBarrier();
516 +
517 + const nipc_win_shm_region_header_t *hdr =
518 + (const nipc_win_shm_region_header_t *)base;
519 +
520 + /* Validate header */
521 + if (hdr->magic != NIPC_WIN_SHM_MAGIC) {
522 + UnmapViewOfFile(base);
523 + CloseHandle(mapping);
524 + return NIPC_WIN_SHM_ERR_BAD_MAGIC;
525 + }
526 +
527 + if (hdr->version != NIPC_WIN_SHM_VERSION) {
528 + UnmapViewOfFile(base);
529 + CloseHandle(mapping);
530 + return NIPC_WIN_SHM_ERR_BAD_VERSION;
531 + }
532 +
533 + if (hdr->header_len != NIPC_WIN_SHM_HEADER_LEN) {
534 + UnmapViewOfFile(base);
535 + CloseHandle(mapping);
536 + return NIPC_WIN_SHM_ERR_BAD_HEADER;
537 + }
538 +
539 + if (hdr->profile != profile) {
540 + UnmapViewOfFile(base);
541 + CloseHandle(mapping);
542 + return NIPC_WIN_SHM_ERR_BAD_PROFILE;
543 + }
544 +
545 + /* Cache and validate header values */
546 + uint32_t req_off = hdr->request_offset;
547 + uint32_t req_cap = hdr->request_capacity;
548 + uint32_t resp_off = hdr->response_offset;
549 + uint32_t resp_cap = hdr->response_capacity;
550 + uint32_t spin = hdr->spin_tries;
551 +
552 + /* Validate offsets and capacities from shared memory */
553 + if (req_off == 0 || req_cap == 0 || resp_off == 0 || resp_cap == 0 ||
554 + req_off % 64 != 0 || req_cap % 64 != 0 ||
555 + resp_off % 64 != 0 || resp_cap % 64 != 0 ||
556 + req_off > UINT32_MAX - req_cap ||
557 + resp_off < req_off + req_cap) {
558 + UnmapViewOfFile(base);
559 + CloseHandle(mapping);
560 + return NIPC_WIN_SHM_ERR_BAD_HEADER;
561 + }
562 +
563 + /* Read current sequence numbers */
564 + LONG64 cur_req_seq = atomic_load_64(
565 + (volatile LONG64 *)&((nipc_win_shm_region_header_t *)base)->req_seq);
566 + LONG64 cur_resp_seq = atomic_load_64(
567 + (volatile LONG64 *)&((nipc_win_shm_region_header_t *)base)->resp_seq);
568 +
569 + size_t region_size = (size_t)resp_off + resp_cap;
570 +
571 + /* Open kernel events for HYBRID profile */
572 + HANDLE req_event = INVALID_HANDLE_VALUE;
573 + HANDLE resp_event = INVALID_HANDLE_VALUE;
574 +
575 + if (profile == NIPC_WIN_SHM_PROFILE_HYBRID) {
576 + wchar_t req_event_name[NIPC_WIN_SHM_MAX_NAME];
577 + wchar_t resp_event_name[NIPC_WIN_SHM_MAX_NAME];
578 +
579 + if (build_object_name(req_event_name, NIPC_WIN_SHM_MAX_NAME,
580 + hash, service_name, profile, session_id,
581 + "req_event") < 0 ||
582 + build_object_name(resp_event_name, NIPC_WIN_SHM_MAX_NAME,
583 + hash, service_name, profile, session_id,
584 + "resp_event") < 0) {
585 + UnmapViewOfFile(base);
586 + CloseHandle(mapping);
587 + return NIPC_WIN_SHM_ERR_BAD_PARAM;
588 + }
589 +
590 + req_event = test_open_event(EVENT_MODIFY_STATE | SYNCHRONIZE, FALSE, req_event_name);
591 + if (!req_event) {
592 + UnmapViewOfFile(base);
593 + CloseHandle(mapping);
594 + return NIPC_WIN_SHM_ERR_OPEN_EVENT;
595 + }
596 +
597 + resp_event = test_open_event(EVENT_MODIFY_STATE | SYNCHRONIZE, FALSE, resp_event_name);
598 + if (!resp_event) {
599 + CloseHandle(req_event);
600 + UnmapViewOfFile(base);
601 + CloseHandle(mapping);
602 + return NIPC_WIN_SHM_ERR_OPEN_EVENT;
603 + }
604 + }
605 +
606 + /* Fill context */
607 + ctx->role = NIPC_WIN_SHM_ROLE_CLIENT;
608 + ctx->mapping = mapping;
609 + ctx->base = base;
610 + ctx->region_size = region_size;
611 + ctx->req_event = req_event;
612 + ctx->resp_event = resp_event;
613 + ctx->profile = profile;
614 + ctx->request_offset = req_off;
615 + ctx->request_capacity = req_cap;
616 + ctx->response_offset = resp_off;
617 + ctx->response_capacity = resp_cap;
618 + ctx->spin_tries = spin;
619 + ctx->local_req_seq = cur_req_seq;
620 + ctx->local_resp_seq = cur_resp_seq;
621 +
622 + return NIPC_WIN_SHM_OK;
623 +}
624 +
625 +/* ------------------------------------------------------------------ */
626 +/* Client: close */
627 +/* ------------------------------------------------------------------ */
628 +
629 +void nipc_win_shm_close(nipc_win_shm_ctx_t *ctx)
630 +{
631 + if (!ctx)
632 + return;
633 +
634 + /* Set close flag and signal waiting server */
635 + if (ctx->base) {
636 + nipc_win_shm_region_header_t *hdr = region_hdr(ctx);
637 + InterlockedExchange(&hdr->req_client_closed, 1);
638 + MemoryBarrier();
639 + }
640 +
641 + if (ctx->profile == NIPC_WIN_SHM_PROFILE_HYBRID &&
642 + ctx->req_event != INVALID_HANDLE_VALUE)
643 + SetEvent(ctx->req_event);
644 +
645 + if (ctx->base) {
646 + UnmapViewOfFile(ctx->base);
647 + ctx->base = NULL;
648 + }
649 +
650 + if (ctx->mapping) {
651 + CloseHandle(ctx->mapping);
652 + ctx->mapping = NULL;
653 + }
654 +
655 + if (ctx->req_event != INVALID_HANDLE_VALUE) {
656 + CloseHandle(ctx->req_event);
657 + ctx->req_event = INVALID_HANDLE_VALUE;
658 + }
659 +
660 + if (ctx->resp_event != INVALID_HANDLE_VALUE) {
661 + CloseHandle(ctx->resp_event);
662 + ctx->resp_event = INVALID_HANDLE_VALUE;
663 + }
664 +
665 + ctx->region_size = 0;
666 +}
667 +
668 +/* ------------------------------------------------------------------ */
669 +/* Data plane: send */
670 +/* ------------------------------------------------------------------ */
671 +
672 +nipc_win_shm_error_t nipc_win_shm_send(
673 + nipc_win_shm_ctx_t *ctx,
674 + const void *msg,
675 + size_t msg_len)
676 +{
677 + if (!ctx || !ctx->base || !msg || msg_len == 0)
678 + return NIPC_WIN_SHM_ERR_BAD_PARAM;
679 +
680 + uint32_t area_offset;
681 + uint32_t area_capacity;
682 + volatile LONG *len_ptr;
683 + volatile LONG64 *seq_ptr;
684 + volatile LONG *peer_waiting_ptr;
685 + HANDLE peer_event;
686 +
687 + nipc_win_shm_region_header_t *hdr = region_hdr(ctx);
688 +
689 + if (ctx->role == NIPC_WIN_SHM_ROLE_CLIENT) {
690 + area_offset = ctx->request_offset;
691 + area_capacity = ctx->request_capacity;
692 + len_ptr = &hdr->req_len;
693 + seq_ptr = &hdr->req_seq;
694 + peer_waiting_ptr = &hdr->req_server_waiting;
695 + peer_event = ctx->req_event;
696 + } else {
697 + area_offset = ctx->response_offset;
698 + area_capacity = ctx->response_capacity;
699 + len_ptr = &hdr->resp_len;
700 + seq_ptr = &hdr->resp_seq;
701 + peer_waiting_ptr = &hdr->resp_client_waiting;
702 + peer_event = ctx->resp_event;
703 + }
704 +
705 + if (msg_len > area_capacity || msg_len > 0x7FFFFFFFu)
706 + return NIPC_WIN_SHM_ERR_MSG_TOO_LARGE;
707 +
708 + /* 1. Write message data into the area */
709 + memcpy(region_ptr(ctx, area_offset), msg, msg_len);
710 +
711 + /* 2. Store message length (interlocked exchange) */
712 + InterlockedExchange(len_ptr, (LONG)msg_len);
713 +
714 + /* 3. Increment sequence number (interlocked increment) */
715 + InterlockedIncrement64(seq_ptr);
716 +
717 + /* 4. If HYBRID and peer is waiting, signal the event */
718 + if (ctx->profile == NIPC_WIN_SHM_PROFILE_HYBRID) {
719 + if (atomic_load_32(peer_waiting_ptr) != 0)
720 + SetEvent(peer_event);
721 + }
722 +
723 + /* Track locally */
724 + if (ctx->role == NIPC_WIN_SHM_ROLE_CLIENT)
725 + ctx->local_req_seq++;
726 + else
727 + ctx->local_resp_seq++;
728 +
729 + return NIPC_WIN_SHM_OK;
730 +}
731 +
732 +/* ------------------------------------------------------------------ */
733 +/* Data plane: receive */
734 +/* ------------------------------------------------------------------ */
735 +
736 +nipc_win_shm_error_t nipc_win_shm_receive(
737 + nipc_win_shm_ctx_t *ctx,
738 + void *buf,
739 + size_t buf_size,
740 + size_t *msg_len_out,
741 + uint32_t timeout_ms)
742 +{
743 + if (!ctx || !ctx->base || !buf || !msg_len_out || buf_size == 0)
744 + return NIPC_WIN_SHM_ERR_BAD_PARAM;
745 +
746 + uint32_t area_offset;
747 + uint32_t area_capacity;
748 + volatile LONG *len_ptr;
749 + volatile LONG64 *seq_ptr;
750 + volatile LONG *self_waiting_ptr;
751 + volatile LONG *peer_closed_ptr;
752 + HANDLE wait_event;
753 + LONG64 expected_seq;
754 +
755 + nipc_win_shm_region_header_t *hdr = region_hdr(ctx);
756 +
757 + if (ctx->role == NIPC_WIN_SHM_ROLE_SERVER) {
758 + area_offset = ctx->request_offset;
759 + area_capacity = ctx->request_capacity;
760 + len_ptr = &hdr->req_len;
761 + seq_ptr = &hdr->req_seq;
762 + self_waiting_ptr = &hdr->req_server_waiting;
763 + peer_closed_ptr = &hdr->req_client_closed;
764 + wait_event = ctx->req_event;
765 + expected_seq = ctx->local_req_seq + 1;
766 + } else {
767 + area_offset = ctx->response_offset;
768 + area_capacity = ctx->response_capacity;
769 + len_ptr = &hdr->resp_len;
770 + seq_ptr = &hdr->resp_seq;
771 + self_waiting_ptr = &hdr->resp_client_waiting;
772 + peer_closed_ptr = &hdr->resp_server_closed;
773 + wait_event = ctx->resp_event;
774 + expected_seq = ctx->local_resp_seq + 1;
775 + }
776 +
777 + /* The copy ceiling is the smaller of the caller buffer and the
778 + * SHM area capacity. This prevents out-of-bounds reads even if
779 + * the peer writes a forged length value. */
780 + uint32_t max_copy = (buf_size < area_capacity) ? (uint32_t)buf_size : area_capacity;
781 +
782 + void *data_ptr = region_ptr(ctx, area_offset);
783 +
784 + /*
785 + * Phase 1: spin. Copy immediately upon observing the sequence advance.
786 + */
787 + bool observed = false;
788 + LONG mlen = 0;
789 + for (uint32_t i = 0; i < ctx->spin_tries; i++) {
790 + LONG64 cur = atomic_load_64(seq_ptr);
791 + if (cur >= expected_seq) {
792 + mlen = atomic_load_32(len_ptr);
793 + if (mlen > 0 && (size_t)mlen <= max_copy)
794 + memcpy(buf, data_ptr, (size_t)mlen);
795 + observed = true;
796 + break;
797 + }
798 + cpu_relax();
799 + }
800 +
801 + /* Phase 2: kernel wait or busy-wait (deadline-based retry for
802 + * spurious wakes — same pattern as POSIX SHM Phase H6 fix). */
803 + if (!observed) {
804 + if (ctx->profile == NIPC_WIN_SHM_PROFILE_HYBRID) {
805 + DWORD deadline_ms = (timeout_ms == 0) ? INFINITE : timeout_ms;
806 + DWORD start_tick = GetTickCount64() & 0xFFFFFFFF;
807 +
808 + for (;;) {
809 + InterlockedExchange(self_waiting_ptr, 1);
810 + MemoryBarrier();
811 +
812 + /* Recheck after setting flag to avoid race */
813 + LONG64 cur = atomic_load_64(seq_ptr);
814 + if (cur >= expected_seq) {
815 + InterlockedExchange(self_waiting_ptr, 0);
816 + break; /* data available */
817 + }
818 +
819 + /* Compute remaining wait time */
820 + DWORD wait_ms;
821 + if (deadline_ms == INFINITE) {
822 + wait_ms = INFINITE;
823 + } else {
824 + DWORD elapsed = (GetTickCount64() & 0xFFFFFFFF) - start_tick;
825 + if (elapsed >= deadline_ms) {
826 + InterlockedExchange(self_waiting_ptr, 0);
827 + return NIPC_WIN_SHM_ERR_TIMEOUT;
828 + }
829 + wait_ms = deadline_ms - elapsed;
830 + }
831 +
832 + DWORD ret = WaitForSingleObject(wait_event, wait_ms);
833 + InterlockedExchange(self_waiting_ptr, 0);
834 +
835 + /* Check sequence — data may have arrived */
836 + cur = atomic_load_64(seq_ptr);
837 + if (cur >= expected_seq)
838 + break; /* data available */
839 +
840 + /* No data — check peer close */
841 + if (atomic_load_32(peer_closed_ptr) != 0) {
842 + cur = atomic_load_64(seq_ptr);
843 + if (cur >= expected_seq)
844 + break;
845 + if (ctx->role == NIPC_WIN_SHM_ROLE_SERVER)
846 + ctx->local_req_seq = expected_seq;
847 + else
848 + ctx->local_resp_seq = expected_seq;
849 + return NIPC_WIN_SHM_ERR_DISCONNECTED;
850 + }
851 +
852 + /* Actual timeout (not spurious) */
853 + if (ret == WAIT_TIMEOUT)
854 + return NIPC_WIN_SHM_ERR_TIMEOUT;
855 +
856 + /* Spurious wake — retry with remaining deadline */
857 + }
858 +
859 + /* Copy immediately after waking */
860 + mlen = atomic_load_32(len_ptr);
861 + if (mlen > 0 && (size_t)mlen <= max_copy)
862 + memcpy(buf, data_ptr, (size_t)mlen);
863 +
864 + } else {
865 + /* SHM_BUSYWAIT: spin indefinitely with periodic deadline checks */
866 + ULONGLONG start = GetTickCount64();
867 + for (;;) {
868 + LONG64 cur = atomic_load_64(seq_ptr);
869 + if (cur >= expected_seq) {
870 + mlen = atomic_load_32(len_ptr);
871 + if (mlen > 0 && (size_t)mlen <= max_copy)
872 + memcpy(buf, data_ptr, (size_t)mlen);
873 + break;
874 + }
875 +
876 + /* Periodic timeout check */
877 + if (timeout_ms > 0) {
878 + ULONGLONG elapsed = GetTickCount64() - start;
879 + if (elapsed >= timeout_ms)
880 + return NIPC_WIN_SHM_ERR_TIMEOUT;
881 + }
882 +
883 + /* Check peer close */
884 + if (atomic_load_32(peer_closed_ptr) != 0) {
885 + cur = atomic_load_64(seq_ptr);
886 + if (cur >= expected_seq) {
887 + mlen = atomic_load_32(len_ptr);
888 + if (mlen > 0 && (size_t)mlen <= max_copy)
889 + memcpy(buf, data_ptr, (size_t)mlen);
890 + break;
891 + }
892 + if (ctx->role == NIPC_WIN_SHM_ROLE_SERVER)
893 + ctx->local_req_seq = expected_seq;
894 + else
895 + ctx->local_resp_seq = expected_seq;
896 + return NIPC_WIN_SHM_ERR_DISCONNECTED;
897 + }
898 +
899 + cpu_relax();
900 + }
901 + }
902 + }
903 +
904 + /* Message larger than caller buffer or area capacity */
905 + if ((size_t)mlen > max_copy) {
906 + *msg_len_out = (size_t)mlen;
907 + /* Still advance tracking -- message is consumed from SHM perspective */
908 + if (ctx->role == NIPC_WIN_SHM_ROLE_SERVER)
909 + ctx->local_req_seq = expected_seq;
910 + else
911 + ctx->local_resp_seq = expected_seq;
912 + return NIPC_WIN_SHM_ERR_MSG_TOO_LARGE;
913 + }
914 +
915 + /* mlen==0 after sequence advance indicates SHM corruption (send rejects 0-length) */
916 + if (mlen == 0) {
917 + if (ctx->role == NIPC_WIN_SHM_ROLE_SERVER)
918 + ctx->local_req_seq = expected_seq;
919 + else
920 + ctx->local_resp_seq = expected_seq;
921 + *msg_len_out = 0;
922 + return NIPC_WIN_SHM_ERR_BAD_HEADER;
923 + }
924 +
925 + *msg_len_out = (size_t)mlen;
926 +
927 + /* Advance local tracking */
928 + if (ctx->role == NIPC_WIN_SHM_ROLE_SERVER)
929 + ctx->local_req_seq = expected_seq;
930 + else
931 + ctx->local_resp_seq = expected_seq;
932 +
933 + return NIPC_WIN_SHM_OK;
934 +}
935 +
936 +/* ------------------------------------------------------------------ */
937 +/* Stale cleanup (no-op on Windows) */
938 +/* ------------------------------------------------------------------ */
939 +
940 +void nipc_win_shm_cleanup_stale(const char *run_dir, const char *service_name)
941 +{
942 + /* Windows kernel objects are reference-counted and auto-cleaned
943 + * when all handles close. No filesystem artifacts to scan. */
944 + (void)run_dir;
945 + (void)service_name;
946 +}
947 +
948 +#endif /* _WIN32 || __MSYS__ */