Selaa lähdekoodia

fix memory leak

Danilo Fragoso 7 kuukautta sitten
vanhempi
sitoutus
3920048461
6 muutettua tiedostoa jossa 103 lisäystä ja 65 poistoa
  1. BIN
      bin/pizzakv_amd64_static
  2. 5 5
      command.zig
  3. 43 39
      index.zig
  4. 9 1
      main.zig
  5. 16 0
      socket.zig
  6. 30 20
      storage.zig

BIN
bin/pizzakv_amd64_static


+ 5 - 5
command.zig

@@ -20,17 +20,17 @@ fn parseKeyValue(buf: []const u8) ?[2][]const u8 {
     return [2][]const u8{ key, kvIterator.rest() };
 }
 
-pub fn parse(msg: []const u8) ?[]const u8 {
+pub fn parse(msg: []const u8, allocator: std.mem.Allocator) ?[]const u8 {
     const trimSet = [_]u8{ '\n', ' ', '\r' };
     const cleanMsg = std.mem.trim(u8, msg, &trimSet);
     var messageIterator = std.mem.splitAny(u8, cleanMsg, " ");
 
     const cmdString = messageIterator.first();
-    const command = std.meta.stringToEnum(Command, cmdString) orelse {
+    const cmd = std.meta.stringToEnum(Command, cmdString) orelse {
         return null;
     };
 
-    switch (command) {
+    switch (cmd) {
         .read => {
             const key = messageIterator.rest();
 
@@ -62,11 +62,11 @@ pub fn parse(msg: []const u8) ?[]const u8 {
             return SUCCESS_RESPONSE;
         },
         .keys => {
-            return index.getAllKeys();
+            return index.getAllKeys(allocator);
         },
         .reads => {
             const prefix = messageIterator.rest();
-            return index.getValuesByPrefix(prefix);
+            return index.getValuesByPrefix(prefix, allocator);
         },
         .status => {
             return "well going our operation";

+ 43 - 39
index.zig

@@ -3,7 +3,6 @@ const storage = @import("storage.zig");
 
 var tree_arena = std.heap.ArenaAllocator.init(std.heap.page_allocator);
 const tree_allocator = tree_arena.allocator();
-const temp_allocator = std.heap.c_allocator;
 
 var tree_mutex: std.Thread.Mutex = .{};
 
@@ -133,16 +132,15 @@ pub fn delete(key: []const u8) void {
     }
 }
 
-fn findNodeForPrefix(node: *RadixNode, prefix: []const u8, actual_path: *[]u8) ?*RadixNode {
+fn findNodeForPrefix(node: *RadixNode, prefix: []const u8, path_buf: *[MAX_KEY_LENGTH]u8, path_len: *usize) ?*RadixNode {
     if (prefix.len == 0) {
-        actual_path.* = &[_]u8{};
+        path_len.* = 0;
         return node;
     }
 
     var current = node;
     var remaining = prefix;
-    var path_buffer: [MAX_KEY_LENGTH]u8 = undefined;
-    var path_len: usize = 0;
+    path_len.* = 0;
 
     while (remaining.len > 0) {
         var found = false;
@@ -154,13 +152,13 @@ fn findNodeForPrefix(node: *RadixNode, prefix: []const u8, actual_path: *[]u8) ?
 
             if (prefix_match_len > 0) {
                 if (prefix_match_len == remaining.len) {
-                    actual_path.* = temp_allocator.dupe(u8, path_buffer[0..path_len]) catch &[_]u8{};
                     return child;
                 }
 
                 if (prefix_match_len == child.edge.len) {
-                    @memcpy(path_buffer[path_len .. path_len + prefix_match_len], child.edge[0..prefix_match_len]);
-                    path_len += prefix_match_len;
+                    if (path_len.* + prefix_match_len > MAX_KEY_LENGTH) return null;
+                    @memcpy(path_buf[path_len.* .. path_len.* + prefix_match_len], child.edge[0..prefix_match_len]);
+                    path_len.* += prefix_match_len;
                     remaining = remaining[prefix_match_len..];
                     current = child;
                     found = true;
@@ -176,13 +174,13 @@ fn findNodeForPrefix(node: *RadixNode, prefix: []const u8, actual_path: *[]u8) ?
         }
     }
 
-    actual_path.* = temp_allocator.dupe(u8, path_buffer[0..path_len]) catch &[_]u8{};
     return current;
 }
 
 fn findNode(node: *RadixNode, key: []const u8) ?*RadixNode {
-    var dummy_path: []u8 = &[_]u8{};
-    return findNodeForPrefix(node, key, &dummy_path);
+    var dummy_buf: [MAX_KEY_LENGTH]u8 = undefined;
+    var dummy_len: usize = 0;
+    return findNodeForPrefix(node, key, &dummy_buf, &dummy_len);
 }
 
 pub fn searchByPrefix(prefix: []const u8) ?*RadixNode {
@@ -207,7 +205,7 @@ fn countKeys(node: *RadixNode) usize {
 const MAX_KEYS_RETURN = 100_000_000;
 const MAX_KEY_LENGTH = 1024;
 
-fn collectKeysWithBuffer(node: *RadixNode, prefix_buffer: []u8, prefix_len: usize, keys: *std.ArrayListUnmanaged([]const u8), max_keys: usize, search_prefix: []const u8, include_node_edge: bool) void {
+fn collectKeysWithBuffer(node: *RadixNode, prefix_buffer: []u8, prefix_len: usize, keys: *std.ArrayListUnmanaged([]const u8), max_keys: usize, search_prefix: []const u8, include_node_edge: bool, allocator: std.mem.Allocator) void {
     if (keys.items.len >= max_keys) return;
 
     var current_len = prefix_len;
@@ -221,8 +219,8 @@ fn collectKeysWithBuffer(node: *RadixNode, prefix_buffer: []u8, prefix_len: usiz
     if (node.is_terminal) {
         const key = prefix_buffer[0..current_len];
         if (key.len >= search_prefix.len and std.mem.eql(u8, key[0..search_prefix.len], search_prefix)) {
-            const key_copy = temp_allocator.dupe(u8, key) catch return;
-            keys.append(temp_allocator, key_copy) catch return;
+            const key_copy = allocator.dupe(u8, key) catch return;
+            keys.append(allocator, key_copy) catch return;
         }
     }
 
@@ -231,66 +229,72 @@ fn collectKeysWithBuffer(node: *RadixNode, prefix_buffer: []u8, prefix_len: usiz
         if (keys.items.len >= max_keys) break;
         const child = entry.value_ptr.*;
 
-        collectKeysWithBuffer(child, prefix_buffer, current_len, keys, max_keys, search_prefix, true);
+        collectKeysWithBuffer(child, prefix_buffer, current_len, keys, max_keys, search_prefix, true, allocator);
     }
 }
 
-fn collectKeys(node: *RadixNode, prefix: []const u8, keys: *std.ArrayListUnmanaged([]const u8), max_keys: usize, search_prefix: []const u8) void {
+fn collectKeys(node: *RadixNode, prefix: []const u8, keys: *std.ArrayListUnmanaged([]const u8), max_keys: usize, search_prefix: []const u8, allocator: std.mem.Allocator) void {
     var prefix_buffer: [MAX_KEY_LENGTH]u8 = undefined;
     if (prefix.len > MAX_KEY_LENGTH) return;
     @memcpy(prefix_buffer[0..prefix.len], prefix);
-    collectKeysWithBuffer(node, &prefix_buffer, prefix.len, keys, max_keys, search_prefix, true);
+    collectKeysWithBuffer(node, &prefix_buffer, prefix.len, keys, max_keys, search_prefix, true, allocator);
 }
 
-pub fn getKeysFromNode(node: *RadixNode, prefix: []const u8) [][]const u8 {
+pub fn getKeysFromNode(node: *RadixNode, prefix: []const u8, allocator: std.mem.Allocator) [][]const u8 {
     var keys_list = std.ArrayListUnmanaged([]const u8){};
-    collectKeys(node, prefix, &keys_list, MAX_KEYS_RETURN, prefix);
-    return keys_list.toOwnedSlice(temp_allocator) catch &[_][]const u8{};
+    collectKeys(node, prefix, &keys_list, MAX_KEYS_RETURN, prefix, allocator);
+    return keys_list.toOwnedSlice(allocator) catch &[_][]const u8{};
 }
 
-pub fn getKeysByPrefix(prefix: []const u8) []const u8 {
+pub fn getKeysByPrefix(prefix: []const u8, allocator: std.mem.Allocator) []const u8 {
     tree_mutex.lock();
     defer tree_mutex.unlock();
 
     ensureRoot();
     const node = searchByPrefix(prefix) orelse return "";
-    const keys = getKeysFromNode(node, prefix);
+    const keys = getKeysFromNode(node, prefix, allocator);
     if (keys.len == 0) return "";
-    return std.mem.join(temp_allocator, "\n", keys) catch "";
+    return std.mem.join(allocator, "\n", keys) catch "";
 }
 
-pub fn getValuesByPrefix(prefix: []const u8) []const u8 {
-    tree_mutex.lock();
-    defer tree_mutex.unlock();
+pub fn getValuesByPrefix(prefix: []const u8, allocator: std.mem.Allocator) []const u8 {
+    // Phase 1: Collect matching keys under tree_mutex
+    var keys: [][]const u8 = &[_][]const u8{};
+    {
+        tree_mutex.lock();
+        defer tree_mutex.unlock();
 
-    ensureRoot();
+        ensureRoot();
 
-    var actual_path: []u8 = &[_]u8{};
-    const node = findNodeForPrefix(root, prefix, &actual_path) orelse {
-        return "";
-    };
+        var path_buf: [MAX_KEY_LENGTH]u8 = undefined;
+        var path_len: usize = 0;
+        const node = findNodeForPrefix(root, prefix, &path_buf, &path_len) orelse {
+            return "";
+        };
 
-    var keys_list = std.ArrayListUnmanaged([]const u8){};
-    collectKeys(node, actual_path, &keys_list, MAX_KEYS_RETURN, prefix);
-    const keys = keys_list.toOwnedSlice(temp_allocator) catch &[_][]const u8{};
+        var keys_list = std.ArrayListUnmanaged([]const u8){};
+        collectKeys(node, path_buf[0..path_len], &keys_list, MAX_KEYS_RETURN, prefix, allocator);
+        keys = keys_list.toOwnedSlice(allocator) catch &[_][]const u8{};
+    }
 
     if (keys.len == 0) return "";
 
-    const values = temp_allocator.alloc([]const u8, keys.len) catch return "";
+    // Phase 2: Read values without tree_mutex to avoid deadlock with write/delete
+    const values = allocator.alloc([]const u8, keys.len) catch return "";
     for (keys, 0..) |key, i| {
         const value = storage.read(key) orelse "";
         values[i] = value;
     }
 
-    return std.mem.join(temp_allocator, "\n", values) catch "";
+    return std.mem.join(allocator, "\n", values) catch "";
 }
 
-pub fn getAllKeys() []const u8 {
+pub fn getAllKeys(allocator: std.mem.Allocator) []const u8 {
     tree_mutex.lock();
     defer tree_mutex.unlock();
 
     ensureRoot();
-    const keys = getKeysFromNode(root, &[_]u8{});
+    const keys = getKeysFromNode(root, &[_]u8{}, allocator);
     if (keys.len == 0) return "";
-    return std.mem.join(temp_allocator, "\n", keys) catch "";
+    return std.mem.join(allocator, "\n", keys) catch "";
 }

+ 9 - 1
main.zig

@@ -109,6 +109,8 @@ pub fn main() !void {
         }
 
         posix.setsockopt(conn, posix.IPPROTO.TCP, TCP.NODELAY, &std.mem.toBytes(@as(c_int, 1))) catch {};
+        socket.setReadTimeout(conn, 300) catch {};  // 5 minutes
+        socket.setWriteTimeout(conn, 300) catch {};  // 5 minutes
 
         if (redis_mode) {
             const thread = try std.Thread.spawn(.{}, handleRedisConnection, .{conn});
@@ -146,17 +148,20 @@ pub fn handleConnection(conn: posix.socket_t) !void {
     defer posix.close(conn);
 
     var requestBuffer: [1024 * 1024]u8 = undefined;
+    var response_arena = std.heap.ArenaAllocator.init(std.heap.c_allocator);
+    defer response_arena.deinit();
 
     while (true) {
         const n = socket.readUntilCR(conn, &requestBuffer) catch |err| {
             if (err == error.ConnectionClosed) break;
+            if (err == error.WouldBlock) continue;
             return err;
         };
         if (n == 0) {
             break;
         }
 
-        const cmdResponse = command.parse(requestBuffer[0..n]) orelse {
+        const cmdResponse = command.parse(requestBuffer[0..n], response_arena.allocator()) orelse {
             socket.write(conn, "error\r") catch |err| {
                 std.debug.print("error writing: {any}", .{err});
             };
@@ -171,6 +176,9 @@ pub fn handleConnection(conn: posix.socket_t) !void {
         socket.writev(conn, &iovecs) catch |err| {
             std.debug.print("error writing: {any}", .{err});
         };
+
+        // Free temporary allocations from this request
+        _ = response_arena.reset(.retain_capacity);
     }
 }
 

+ 16 - 0
socket.zig

@@ -2,6 +2,22 @@ const std = @import("std");
 const net = std.net;
 const posix = std.posix;
 
+pub fn setReadTimeout(conn: posix.socket_t, seconds: u32) !void {
+    const timeout = posix.timeval{
+        .sec = @intCast(seconds),
+        .usec = 0,
+    };
+    try posix.setsockopt(conn, posix.SOL.SOCKET, posix.SO.RCVTIMEO, &std.mem.toBytes(timeout));
+}
+
+pub fn setWriteTimeout(conn: posix.socket_t, seconds: u32) !void {
+    const timeout = posix.timeval{
+        .sec = @intCast(seconds),
+        .usec = 0,
+    };
+    try posix.setsockopt(conn, posix.SOL.SOCKET, posix.SO.SNDTIMEO, &std.mem.toBytes(timeout));
+}
+
 pub fn init(port: u16) !posix.socket_t {
     const address = try std.net.Address.parseIp("0.0.0.0", port);
 

+ 30 - 20
storage.zig

@@ -53,25 +53,29 @@ pub fn restore(key: []const u8, value: []const u8) bool {
     const hash = hashing.hashKey(key);
     const shard_idx = getShardIndex(hash);
 
-    shards[shard_idx].rwlock.lock();
-    defer shards[shard_idx].rwlock.unlock();
+    var entry_key: []const u8 = undefined;
+    {
+        shards[shard_idx].rwlock.lock();
+        defer shards[shard_idx].rwlock.unlock();
 
-    const entry = writeVolatile(hash, key, value);
-    if (entry != null) {
-        index.insert(entry.?.key);
-        return true;
+        const entry = writeVolatile(hash, key, value) orelse return false;
+        entry_key = entry.key;
     }
-    return false;
+
+    index.insert(entry_key);
+    return true;
 }
 
 pub fn restoreDelete(key: []const u8) bool {
     const hash = hashing.hashKey(key);
     const shard_idx = getShardIndex(hash);
 
-    shards[shard_idx].rwlock.lock();
-    defer shards[shard_idx].rwlock.unlock();
+    const deleted = blk: {
+        shards[shard_idx].rwlock.lock();
+        defer shards[shard_idx].rwlock.unlock();
+        break :blk deleteVolatile(hash, key);
+    };
 
-    const deleted = deleteVolatile(hash, key);
     if (deleted) {
         index.delete(key);
         return true;
@@ -112,16 +116,19 @@ pub fn write(key: []const u8, value: []const u8) bool {
     const hash = hashing.hashKey(key);
     const shard_idx = getShardIndex(hash);
 
-    shards[shard_idx].rwlock.lock();
-    defer shards[shard_idx].rwlock.unlock();
+    var entry_key: []const u8 = undefined;
+    var entry_value: []const u8 = undefined;
+    {
+        shards[shard_idx].rwlock.lock();
+        defer shards[shard_idx].rwlock.unlock();
 
-    const entry = writeVolatile(hash, key, value);
-    if (entry == null) {
-        return false;
+        const entry = writeVolatile(hash, key, value) orelse return false;
+        entry_key = entry.key;
+        entry_value = entry.value;
     }
 
-    index.insert(entry.?.key);
-    persistence.persist('W', entry.?.key, entry.?.value);
+    index.insert(entry_key);
+    persistence.persist('W', entry_key, entry_value);
     return true;
 }
 
@@ -178,10 +185,13 @@ pub fn deleteVolatile(hash: u32, key: []const u8) bool {
 pub fn delete(key: []const u8) bool {
     const hash = hashing.hashKey(key);
     const shard_idx = getShardIndex(hash);
-    shards[shard_idx].rwlock.lock();
-    defer shards[shard_idx].rwlock.unlock();
 
-    const deleted = deleteVolatile(hash, key);
+    const deleted = blk: {
+        shards[shard_idx].rwlock.lock();
+        defer shards[shard_idx].rwlock.unlock();
+        break :blk deleteVolatile(hash, key);
+    };
+
     if (deleted) {
         index.delete(key);
         persistence.persist('D', key, "");