@@ -30,8 +30,12 @@ const ServerMapping = struct {
3030 addr : net.Address ,
3131};
3232
33- /// Mapping: client address → backend address (for return traffic)
33+ /// Mapping: connection ID → client address (for return traffic)
34+ /// We store the client's SCID (which becomes the server's DCID in responses)
3435const ClientMapping = struct {
36+ /// Client's SCID (the DCID the server will use when responding)
37+ cid : [20 ]u8 = .{0 } ** 20 ,
38+ cid_len : u8 = 0 ,
3539 client_addr : posix.sockaddr.storage = std .mem .zeroes (posix .sockaddr .storage ),
3640 backend_addr : posix.sockaddr.storage = std .mem .zeroes (posix .sockaddr .storage ),
3741 occupied : bool = false ,
@@ -195,6 +199,7 @@ pub fn main() !void {
195199 // Client mappings for return traffic
196200 var client_mappings : [MAX_CLIENTS ]ClientMapping = .{ClientMapping {}} ** MAX_CLIENTS ;
197201
202+ var next_rr : usize = 0 ; // round-robin index for Initial packets
198203 var recv_buf : [MAX_PACKET_SIZE ]u8 = undefined ;
199204 var src_addr : posix.sockaddr.storage = std .mem .zeroes (posix .sockaddr .storage );
200205 var addr_len : posix.socklen_t = @sizeOf (posix .sockaddr .storage );
@@ -218,17 +223,20 @@ pub fn main() !void {
218223 const now : i64 = @intCast (std .time .nanoTimestamp ());
219224
220225 // Check if this is from a known backend (return traffic)
221- if (findClientForBackend (& client_mappings , & src_addr )) | client_entry | {
222- // Forward back to client
223- const client_sa : * const posix.sockaddr = @ptrCast (& client_entry .client_addr );
224- const client_addr_len : posix.socklen_t = if (client_entry .client_addr .family == posix .AF .INET6 )
225- @sizeOf (posix .sockaddr .in6 )
226- else
227- @sizeOf (posix .sockaddr .in );
228- _ = posix .sendto (sock , recv_buf [0.. n ], 0 , client_sa , client_addr_len ) catch | err | {
229- std .log .err ("sendto client error: {any}" , .{err });
230- continue ;
231- };
226+ if (isFromBackend (& config , & src_addr )) {
227+ // Extract DCID from the response packet to find the client
228+ const resp_dcid = extractDcid (recv_buf [0.. n ]) orelse continue ;
229+ if (findClientByCid (& client_mappings , resp_dcid )) | client_entry | {
230+ const client_sa : * const posix.sockaddr = @ptrCast (& client_entry .client_addr );
231+ const client_addr_len : posix.socklen_t = if (client_entry .client_addr .family == posix .AF .INET6 )
232+ @sizeOf (posix .sockaddr .in6 )
233+ else
234+ @sizeOf (posix .sockaddr .in );
235+ _ = posix .sendto (sock , recv_buf [0.. n ], 0 , client_sa , client_addr_len ) catch | err | {
236+ std .log .err ("sendto client error: {any}" , .{err });
237+ continue ;
238+ };
239+ }
232240 continue ;
233241 }
234242
@@ -255,62 +263,132 @@ pub fn main() !void {
255263
256264 const dcid = recv_buf [dcid_offset .. dcid_offset + dcid_len ];
257265
258- // Check config_id matches
259- if (dcid .len < 1 ) continue ;
260- const pkt_config_id = quic_lb .extractConfigId (dcid [0 ]);
261- if (pkt_config_id != config .lb_config .config_id ) {
262- std .log .warn ("unknown config_id {d} in packet, dropping" , .{pkt_config_id });
263- continue ;
266+ // Try to extract server_id from CID for routing
267+ var backend_addr : ? net.Address = null ;
268+
269+ if (dcid .len >= quic_lb .cidLength (& config .lb_config )) {
270+ var server_id : [15 ]u8 = undefined ;
271+ if (quic_lb .extractServerId (& config .lb_config , dcid , & server_id )) {
272+ backend_addr = findBackend (& config , & server_id );
273+ }
264274 }
265275
266- // Extract server_id
267- var server_id : [15 ]u8 = undefined ;
268- if (! quic_lb .extractServerId (& config .lb_config , dcid , & server_id )) {
269- std .log .warn ("failed to extract server_id from CID, dropping" , .{});
270- continue ;
276+ // Fallback: Initial packets have random client-generated DCIDs.
277+ // Route by round-robin, or reuse existing client mapping.
278+ if (backend_addr == null ) {
279+ // Check if we already have a mapping for this client
280+ if (findMappingForClient (& client_mappings , & src_addr )) | existing | {
281+ // Find the server entry matching this backend address
282+ for (config .servers [0.. config .server_count ]) | mapping | {
283+ if (sockaddrEql (& existing .backend_addr , & addrToStorage (& mapping .addr .any ))) {
284+ backend_addr = mapping .addr ;
285+ break ;
286+ }
287+ }
288+ }
289+ if (backend_addr == null and config .server_count > 0 ) {
290+ // Round-robin for new clients
291+ backend_addr = config .servers [next_rr ].addr ;
292+ next_rr = (next_rr + 1 ) % config .server_count ;
293+ }
271294 }
272295
273- // Look up backend
274- const backend_addr = findBackend (& config , & server_id ) orelse {
275- std .log .warn ("no backend for server_id (first 2 bytes: {x:0>2}{x:0>2})" , .{ server_id [0 ], server_id [1 ] });
276- continue ;
277- };
296+ const dest = backend_addr orelse continue ;
278297
279298 // Store client → backend mapping for return traffic
280- storeClientMapping (& client_mappings , & src_addr , & backend_addr .any , now );
299+ // Extract client's SCID from long header (Initial/Handshake)
300+ if (is_long_header and n > dcid_offset + dcid_len + 1 ) {
301+ const scid_len_pos = dcid_offset + dcid_len ;
302+ const scid_len_val = recv_buf [scid_len_pos ];
303+ if (n > scid_len_pos + 1 + scid_len_val ) {
304+ const scid = recv_buf [scid_len_pos + 1 .. scid_len_pos + 1 + scid_len_val ];
305+ storeClientMapping (& client_mappings , scid , & src_addr , & dest .any , now );
306+ }
307+ }
281308
282309 // Forward to backend
283- _ = posix .sendto (sock , recv_buf [0.. n ], 0 , & backend_addr .any , backend_addr .getOsSockLen ()) catch | err | {
310+ _ = posix .sendto (sock , recv_buf [0.. n ], 0 , & dest .any , dest .getOsSockLen ()) catch | err | {
284311 std .log .err ("sendto backend error: {any}" , .{err });
285312 continue ;
286313 };
287314 }
288315}
289316
290- fn findClientForBackend (mappings : []ClientMapping , backend_addr : * const posix.sockaddr.storage ) ? * ClientMapping {
317+ fn isFromBackend (config : * const LbConfig , addr : * const posix.sockaddr.storage ) bool {
318+ const src_port = std .mem .readInt (u16 , std .mem .asBytes (addr )[2.. 4], .big );
319+ for (config .servers [0.. config .server_count ]) | mapping | {
320+ if (mapping .addr .getPort () == src_port ) {
321+ std .log .info ("LB: return traffic from backend port {d}" , .{src_port });
322+ return true ;
323+ }
324+ }
325+ std .log .debug ("LB: packet from port {d} (not a backend)" , .{src_port });
326+ return false ;
327+ }
328+
329+ fn extractDcid (pkt : []const u8 ) ? []const u8 {
330+ if (pkt .len < 2 ) return null ;
331+ if (pkt [0 ] & 0x80 != 0 ) {
332+ // Long header: DCID len at byte 5, DCID at byte 6
333+ if (pkt .len < 6 ) return null ;
334+ const dcid_len = pkt [5 ];
335+ if (pkt .len < 6 + dcid_len ) return null ;
336+ return pkt [6 .. 6 + dcid_len ];
337+ } else {
338+ // Short header: DCID at byte 1, but we don't know the length.
339+ // Use the client's known CID length — typically 8 bytes for our implementation.
340+ const dcid_len : usize = @min (8 , pkt .len - 1 );
341+ return pkt [1 .. 1 + dcid_len ];
342+ }
343+ }
344+
345+ fn findClientByCid (mappings : []ClientMapping , dcid : []const u8 ) ? * ClientMapping {
291346 for (mappings ) | * m | {
292- if (m .occupied and sockaddrEql ( & m . backend_addr , backend_addr ) ) {
293- return m ;
347+ if (m .occupied and m . cid_len > 0 and m . cid_len <= dcid . len ) {
348+ if ( std . mem . eql ( u8 , m . cid [0 .. m . cid_len ], dcid [0 .. m . cid_len ])) return m ;
294349 }
295350 }
351+ std .log .debug ("LB: findClientByCid failed for dcid len={d}, stored mappings: {d} occupied" , .{
352+ dcid .len ,
353+ countOccupied (mappings ),
354+ });
296355 return null ;
297356}
298357
299- fn storeClientMapping (mappings : []ClientMapping , client_addr : * const posix.sockaddr.storage , backend_addr : * const posix.sockaddr , now : i64 ) void {
300- // Try to find existing or empty slot
358+ fn countOccupied (mappings : []const ClientMapping ) usize {
359+ var n : usize = 0 ;
360+ for (mappings ) | m | {
361+ if (m .occupied ) n += 1 ;
362+ }
363+ return n ;
364+ }
365+
366+ fn findMappingForClient (mappings : []ClientMapping , client_addr : * const posix.sockaddr.storage ) ? * ClientMapping {
367+ for (mappings ) | * m | {
368+ if (m .occupied and sockaddrEql (& m .client_addr , client_addr )) return m ;
369+ }
370+ return null ;
371+ }
372+
373+ fn storeClientMapping (mappings : []ClientMapping , client_scid : []const u8 , client_addr : * const posix.sockaddr.storage , backend_addr : * const posix.sockaddr , now : i64 ) void {
301374 var oldest_idx : usize = 0 ;
302375 var oldest_time : i64 = std .math .maxInt (i64 );
303376
304377 for (mappings , 0.. ) | * m , i | {
305- if (! m .occupied ) {
306- m .occupied = true ;
378+ // Update existing by CID match
379+ if (m .occupied and m .cid_len > 0 and m .cid_len == client_scid .len and
380+ std .mem .eql (u8 , m .cid [0.. m .cid_len ], client_scid ))
381+ {
307382 m .client_addr = client_addr .* ;
308383 @memcpy (std .mem .asBytes (& m .backend_addr )[0.. @sizeOf (posix .sockaddr )], std .mem .asBytes (backend_addr ));
309384 m .last_seen = now ;
310385 return ;
311386 }
312- if (sockaddrEql (& m .client_addr , client_addr )) {
313- // Update existing mapping
387+ if (! m .occupied ) {
388+ m .occupied = true ;
389+ m .cid_len = @intCast (@min (client_scid .len , 20 ));
390+ @memcpy (m .cid [0.. m .cid_len ], client_scid [0.. m .cid_len ]);
391+ m .client_addr = client_addr .* ;
314392 @memcpy (std .mem .asBytes (& m .backend_addr )[0.. @sizeOf (posix .sockaddr )], std .mem .asBytes (backend_addr ));
315393 m .last_seen = now ;
316394 return ;
@@ -322,10 +400,22 @@ fn storeClientMapping(mappings: []ClientMapping, client_addr: *const posix.socka
322400 }
323401
324402 // Evict oldest
325- mappings [oldest_idx ].occupied = true ;
326- mappings [oldest_idx ].client_addr = client_addr .* ;
327- @memcpy (std .mem .asBytes (& mappings [oldest_idx ].backend_addr )[0.. @sizeOf (posix .sockaddr )], std .mem .asBytes (backend_addr ));
328- mappings [oldest_idx ].last_seen = now ;
403+ const m = & mappings [oldest_idx ];
404+ m .occupied = true ;
405+ m .cid_len = @intCast (@min (client_scid .len , 20 ));
406+ @memcpy (m .cid [0.. m .cid_len ], client_scid [0.. m .cid_len ]);
407+ m .client_addr = client_addr .* ;
408+ @memcpy (std .mem .asBytes (& m .backend_addr )[0.. @sizeOf (posix .sockaddr )], std .mem .asBytes (backend_addr ));
409+ m .last_seen = now ;
410+ }
411+
412+ fn addrToStorage (sa : * const posix.sockaddr ) posix.sockaddr.storage {
413+ var storage : posix.sockaddr.storage = std .mem .zeroes (posix .sockaddr .storage );
414+ const src = std .mem .asBytes (sa );
415+ const dst = std .mem .asBytes (& storage );
416+ const len : usize = if (sa .family == posix .AF .INET6 ) @sizeOf (posix .sockaddr .in6 ) else @sizeOf (posix .sockaddr .in );
417+ @memcpy (dst [0.. len ], src [0.. len ]);
418+ return storage ;
329419}
330420
331421fn sockaddrEql (a : * const posix.sockaddr.storage , b : * const posix.sockaddr.storage ) bool {
0 commit comments