Переглянути джерело

harden protocol and WAL writes

Danilo Fragoso 2 тижнів тому
батько
коміт
c726cf8efe
6 змінених файлів з 294 додано та 161 видалено
  1. 38 7
      main.zig
  2. 1 0
      makefile
  3. 65 71
      persistence.zig
  4. 168 72
      redis.zig
  5. 2 2
      socket.zig
  6. 20 9
      storage.zig

+ 38 - 7
main.zig

@@ -22,6 +22,7 @@ const TCP = switch (builtin.target.os.tag) {
 //main.zig:23:12: error: variable of type 'comptime_int' must be const or comptime
 // var PORT = 8085;
 var PORT: u16 = 8085;
+var HOST: []const u8 = "127.0.0.1";
 var should_exit = std.atomic.Value(bool).init(false);
 var active_connections = std.atomic.Value(u32).init(0);
 var redis_mode = false;
@@ -54,6 +55,8 @@ pub fn main() !void {
                 std.debug.print("Invalid port number: {any}\n", .{port_str});
                 return;
             }
+        } else if (arg.len > 6 and std.mem.eql(u8, arg[0..6], "-host=")) {
+            HOST = arg[6..];
         } else {
             std.debug.print("Unknown argument: {any}\n", .{arg});
             return;
@@ -79,14 +82,14 @@ pub fn main() !void {
     }
 
     const unix_path = ".pizzakv.sock";
-    const listener = if (unix_mode) try socket.initUnix(unix_path) else try socket.init(PORT);
+    const listener = if (unix_mode) try socket.initUnix(unix_path) else try socket.init(HOST, PORT);
     defer posix.close(listener);
     defer if (unix_mode) posix.unlink(unix_path) catch {};
 
     if (unix_mode) {
         std.debug.print("\n2025 pizzakv! Unix socket at {s}\n<danilo@fragoso.dev>\n---------\n", .{unix_path});
     } else {
-        std.debug.print("\n2025 pizzakv! TCP Listening on port {any}\n<danilo@fragoso.dev>\n---------\n", .{PORT});
+        std.debug.print("\n2025 pizzakv! TCP Listening on {s}:{any}\n<danilo@fragoso.dev>\n---------\n", .{ HOST, PORT });
     }
     if (redis_mode) {
         std.debug.print("Mode: Redis Protocol (RESP)\nCommands: SET, GET, DEL\n", .{});
@@ -204,6 +207,15 @@ pub fn handleConnection(conn: posix.socket_t) !void {
     }
 }
 
+fn sendAll(conn: posix.socket_t, buf: []const u8) !void {
+    var offset: usize = 0;
+    while (offset < buf.len) {
+        const sent = try posix.send(conn, buf[offset..], posix.MSG.NOSIGNAL);
+        if (sent == 0) return error.ConnectionClosed;
+        offset += sent;
+    }
+}
+
 pub fn handleRedisConnection(conn: posix.socket_t) !void {
     _ = active_connections.fetchAdd(1, .seq_cst);
     defer _ = active_connections.fetchSub(1, .seq_cst);
@@ -234,17 +246,36 @@ pub fn handleRedisConnection(conn: posix.socket_t) !void {
                 break;
             };
 
-            const response = redis.executeCommand(result.cmd, responseBuffer[response_offset..]);
-            response_offset += response.len;
+            // Mutating commands have small responses. Flush first when the
+            // remaining output slice is small so SET/DEL are never retried
+            // after their side effect has already happened.
+            if (response_offset > 0 and responseBuffer.len - response_offset < 64 and result.cmd.cmd_type != .GET) {
+                try sendAll(conn, responseBuffer[0..response_offset]);
+                response_offset = 0;
+            }
+
+            var response = redis.executeCommand(result.cmd, responseBuffer[response_offset..]);
+            if (response == null) {
+                if (response_offset > 0) {
+                    try sendAll(conn, responseBuffer[0..response_offset]);
+                    response_offset = 0;
+                }
+                response = redis.executeCommand(result.cmd, responseBuffer[0..]) orelse
+                    redis.formatError(responseBuffer[0..], "ERR response too large");
+            }
+
+            const final_response = response orelse {
+                offset += result.bytes_consumed;
+                continue;
+            };
+            response_offset += final_response.len;
             offset += result.bytes_consumed;
         }
 
         posix.setsockopt(conn, posix.IPPROTO.TCP, cork_option, &std.mem.toBytes(@as(c_int, 0))) catch {};
 
         if (response_offset > 0) {
-            _ = posix.send(conn, responseBuffer[0..response_offset], posix.MSG.NOSIGNAL) catch |err| {
-                std.debug.print("error writing: {any}", .{err});
-            };
+            try sendAll(conn, responseBuffer[0..response_offset]);
         }
 
         if (offset < total_len) {

+ 1 - 0
makefile

@@ -31,6 +31,7 @@ test:
 	zig test storage.zig
 	zig test index.zig
 	zig test command.zig
+	zig test persistence.zig
 
 bench:
 	node tools/test_nov.js

+ 65 - 71
persistence.zig

@@ -20,35 +20,24 @@ const OPCode = enum {
 
 pub fn init() !void {
     const cwd = std.fs.cwd();
-    storage_file = cwd.openFile(".db", .{ .mode = .read_write }) catch |err| {
+    storage_file = cwd.openFile(".db", .{ .mode = .read_write }) catch |err| blk: {
         if (err == std.fs.File.OpenError.FileNotFound) {
             std.debug.print("No persisted data found, starting fresh...\n", .{});
-
-            storage_file = cwd.createFile(".db", .{ .read = true }) catch |ierr| {
-                std.debug.print("Failed to create storage file: {any}\n", .{ierr});
-                return;
-            };
-
+            const file = try cwd.createFile(".db", .{ .read = true });
             std.debug.print("Created new storage file .db\n", .{});
+            break :blk file;
+        } else {
+            return err;
         }
-        return;
     };
 
     var record_count: usize = 0;
-    restoreFromFile(storage_file.?, &record_count) catch |err| {
-        std.debug.print("Failed to read storage file: {any}\n", .{err});
-        return;
-    };
+    try restoreFromFile(storage_file.?, &record_count);
     std.debug.print("Restored {d} records from persistence", .{record_count});
 
     storage_file.?.close();
-    storage_file = cwd.openFile(".db", .{ .mode = .write_only }) catch |err| {
-        std.debug.print("Failed to reopen storage file in append mode: {any}\n", .{err});
-        return;
-    };
+    storage_file = try cwd.openFile(".db", .{ .mode = .write_only });
     try storage_file.?.seekFromEnd(0);
-
-    return;
 }
 
 fn restoreFromFile(file: std.fs.File, record_count: *usize) !void {
@@ -106,69 +95,52 @@ pub fn setInstantWal(enabled: bool) void {
     instant_wal = enabled;
 }
 
-pub fn persist(opcode: u8, key: []const u8, value: []const u8) void {
-    const record_len = 1 + 1 + key.len + 1 + value.len + 1;
+fn recordLen(key: []const u8, value: []const u8) usize {
+    return 1 + 1 + key.len + 1 + value.len + 1;
+}
 
-    mutex.lock();
-    defer mutex.unlock();
+fn encodeRecord(buf: []u8, opcode: u8, key: []const u8, value: []const u8) usize {
+    var pos: usize = 0;
+    buf[pos] = opcode;
+    pos += 1;
+    buf[pos] = '|';
+    pos += 1;
+    @memcpy(buf[pos .. pos + key.len], key);
+    pos += key.len;
+    buf[pos] = '|';
+    pos += 1;
+    @memcpy(buf[pos .. pos + value.len], value);
+    pos += value.len;
+    buf[pos] = '\r';
+    pos += 1;
+    return pos;
+}
 
-    if (buffer_position + record_len > FLUSH_THRESHOLD) {
-        flushBuffer() catch |err| {
-            std.debug.print("Failed to flush buffer: {any}\n", .{err});
-            return;
-        };
+fn syncFile() !void {
+    if (storage_file) |f| {
+        try f.sync();
     }
+}
+
+pub fn persist(opcode: u8, key: []const u8, value: []const u8) !void {
+    const record_len = recordLen(key, value);
+
+    mutex.lock();
+    defer mutex.unlock();
 
     if (record_len > BUFFER_SIZE) {
-        var temp_buffer: [BUFFER_SIZE]u8 = undefined;
-        var pos: usize = 0;
-        temp_buffer[pos] = opcode;
-        pos += 1;
-        temp_buffer[pos] = '|';
-        pos += 1;
-        @memcpy(temp_buffer[pos .. pos + key.len], key);
-        pos += key.len;
-        temp_buffer[pos] = '|';
-        pos += 1;
-        @memcpy(temp_buffer[pos .. pos + value.len], value);
-        pos += value.len;
-        temp_buffer[pos] = '\r';
-        pos += 1;
-
-        _ = storage_file.?.write(temp_buffer[0..pos]) catch |err| {
-            std.debug.print("Failed to write large record to storage file: {any}\n", .{err});
-            return;
-        };
-        return;
+        return error.RecordTooLarge;
     }
 
-    if (buffer_position + record_len > BUFFER_SIZE) {
-        flushBuffer() catch |err| {
-            std.debug.print("Failed to flush buffer: {any}\n", .{err});
-            return;
-        };
+    if (buffer_position + record_len > FLUSH_THRESHOLD) {
+        try flushBuffer();
     }
 
-    var pos = buffer_position;
-    write_buffer[pos] = opcode;
-    pos += 1;
-    write_buffer[pos] = '|';
-    pos += 1;
-    @memcpy(write_buffer[pos .. pos + key.len], key);
-    pos += key.len;
-    write_buffer[pos] = '|';
-    pos += 1;
-    @memcpy(write_buffer[pos .. pos + value.len], value);
-    pos += value.len;
-    write_buffer[pos] = '\r';
-    pos += 1;
-
-    buffer_position = pos;
+    buffer_position += encodeRecord(write_buffer[buffer_position..], opcode, key, value);
 
     if (instant_wal) {
-        flushBuffer() catch |err| {
-            std.debug.print("Failed to flush buffer in instant WAL mode: {any}\n", .{err});
-        };
+        try flushBuffer();
+        try syncFile();
     }
 }
 
@@ -176,6 +148,7 @@ pub fn flush() !void {
     mutex.lock();
     defer mutex.unlock();
     try flushBuffer();
+    try syncFile();
 }
 
 fn flushBuffer() !void {
@@ -183,6 +156,27 @@ fn flushBuffer() !void {
         return;
     }
 
-    _ = try storage_file.?.write(write_buffer[0..buffer_position]);
+    const f = storage_file orelse return error.StorageFileNotOpen;
+    try f.writeAll(write_buffer[0..buffer_position]);
     buffer_position = 0;
 }
+
+// -- Tests --
+
+test "recordLen matches encoded record length" {
+    try std.testing.expectEqual(@as(usize, 4 + 3 + 5), recordLen("key", "value"));
+    try std.testing.expectEqual(@as(usize, 4), recordLen("", ""));
+}
+
+test "encodeRecord produces WAL framing" {
+    var buf: [64]u8 = undefined;
+    const written = encodeRecord(&buf, 'W', "key1", "value1");
+    try std.testing.expectEqualStrings("W|key1|value1\r", buf[0..written]);
+    try std.testing.expectEqual(@as(usize, recordLen("key1", "value1")), written);
+}
+
+test "encodeRecord delete framing" {
+    var buf: [64]u8 = undefined;
+    const written = encodeRecord(&buf, 'D', "key2", "");
+    try std.testing.expectEqualStrings("D|key2|\r", buf[0..written]);
+}

+ 168 - 72
redis.zig

@@ -24,7 +24,9 @@ fn parseInteger(buf: []const u8, start: usize, end: usize) ?usize {
     var result: usize = 0;
     for (buf[start..end]) |c| {
         if (c < '0' or c > '9') return null;
-        result = result * 10 + (c - '0');
+        const digit: usize = c - '0';
+        result = std.math.mul(usize, result, 10) catch return null;
+        result = std.math.add(usize, result, digit) catch return null;
     }
     return result;
 }
@@ -38,7 +40,7 @@ fn parseBulkString(buf: []const u8, pos: *usize) ?[]const u8 {
     pos.* = len_end + 2;
 
     const str_start = pos.*;
-    const str_end = str_start + len;
+    const str_end = std.math.add(usize, str_start, len) catch return null;
     if (str_end > buf.len) return null;
 
     const result = buf[str_start..str_end];
@@ -113,23 +115,51 @@ pub fn parseCommand(buf: []const u8) ?ParseResult {
     };
 }
 
-fn formatSimpleString(buf: []u8, str: []const u8) []const u8 {
-    var pos: usize = 0;
-    buf[pos] = '+';
-    pos += 1;
-    @memcpy(buf[pos .. pos + str.len], str);
-    pos += str.len;
-    buf[pos] = '\r';
-    buf[pos + 1] = '\n';
-    return buf[0 .. pos + 2];
+fn intDigits(value: usize) usize {
+    if (value == 0) return 1;
+    var v = value;
+    var d: usize = 0;
+    while (v > 0) : (v /= 10) d += 1;
+    return d;
 }
 
-fn formatBulkString(buf: []u8, str: []const u8) []const u8 {
-    var pos: usize = 0;
-    buf[pos] = '$';
-    pos += 1;
+fn formatInt(buf: []u8, value: usize) ?usize {
+    if (value == 0) {
+        if (buf.len < 1) return null;
+        buf[0] = '0';
+        return 1;
+    }
+
+    const len = intDigits(value);
+    if (buf.len < len) return null;
 
-    pos += formatInt(buf[pos..], str.len);
+    var v = value;
+    var i: usize = len;
+    while (i > 0) {
+        i -= 1;
+        buf[i] = @intCast('0' + (v % 10));
+        v /= 10;
+    }
+    return len;
+}
+
+fn formatSimpleString(buf: []u8, str: []const u8) ?[]const u8 {
+    const needed = 1 + str.len + 2;
+    if (buf.len < needed) return null;
+    buf[0] = '+';
+    @memcpy(buf[1 .. 1 + str.len], str);
+    buf[1 + str.len] = '\r';
+    buf[2 + str.len] = '\n';
+    return buf[0..needed];
+}
+
+fn formatBulkString(buf: []u8, str: []const u8) ?[]const u8 {
+    const needed = 1 + intDigits(str.len) + 2 + str.len + 2;
+    if (buf.len < needed) return null;
+
+    buf[0] = '$';
+    var pos: usize = 1;
+    pos += formatInt(buf[pos..], str.len) orelse return null;
     buf[pos] = '\r';
     buf[pos + 1] = '\n';
     pos += 2;
@@ -142,31 +172,8 @@ fn formatBulkString(buf: []u8, str: []const u8) []const u8 {
     return buf[0 .. pos + 2];
 }
 
-fn formatInt(buf: []u8, value: usize) usize {
-    if (value == 0) {
-        buf[0] = '0';
-        return 1;
-    }
-
-    var v = value;
-    var len: usize = 0;
-    var temp: [20]u8 = undefined;
-
-    while (v > 0) {
-        temp[len] = @intCast('0' + (v % 10));
-        v /= 10;
-        len += 1;
-    }
-
-    var i: usize = 0;
-    while (i < len) : (i += 1) {
-        buf[i] = temp[len - 1 - i];
-    }
-
-    return len;
-}
-
-fn formatNullBulkString(buf: []u8) []const u8 {
+fn formatNullBulkString(buf: []u8) ?[]const u8 {
+    if (buf.len < 5) return null;
     buf[0] = '$';
     buf[1] = '-';
     buf[2] = '1';
@@ -175,17 +182,21 @@ fn formatNullBulkString(buf: []u8) []const u8 {
     return buf[0..5];
 }
 
-fn formatInteger(buf: []u8, value: i64) []const u8 {
-    var pos: usize = 0;
-    buf[pos] = ':';
-    pos += 1;
-
+fn formatInteger(buf: []u8, value: i64) ?[]const u8 {
+    var pos: usize = 1;
     if (value < 0) {
-        buf[pos] = '-';
-        pos += 1;
-        pos += formatInt(buf[pos..], @intCast(-value));
+        const needed = 1 + 1 + intDigits(@intCast(-value)) + 2;
+        if (buf.len < needed) return null;
+        buf[0] = ':';
+        buf[1] = '-';
+        pos = 2;
+        pos += formatInt(buf[pos..], @intCast(-value)) orelse return null;
     } else {
-        pos += formatInt(buf[pos..], @intCast(value));
+        const needed = 1 + intDigits(@intCast(value)) + 2;
+        if (buf.len < needed) return null;
+        buf[0] = ':';
+        pos = 1;
+        pos += formatInt(buf[pos..], @intCast(value)) orelse return null;
     }
 
     buf[pos] = '\r';
@@ -193,18 +204,17 @@ fn formatInteger(buf: []u8, value: i64) []const u8 {
     return buf[0 .. pos + 2];
 }
 
-fn formatError(buf: []u8, msg: []const u8) []const u8 {
-    var pos: usize = 0;
-    buf[pos] = '-';
-    pos += 1;
-    @memcpy(buf[pos .. pos + msg.len], msg);
-    pos += msg.len;
-    buf[pos] = '\r';
-    buf[pos + 1] = '\n';
-    return buf[0 .. pos + 2];
+pub fn formatError(buf: []u8, msg: []const u8) ?[]const u8 {
+    const needed = 1 + msg.len + 2;
+    if (buf.len < needed) return null;
+    buf[0] = '-';
+    @memcpy(buf[1 .. 1 + msg.len], msg);
+    buf[1 + msg.len] = '\r';
+    buf[2 + msg.len] = '\n';
+    return buf[0..needed];
 }
 
-pub fn executeCommand(cmd: RedisCommand, response_buf: []u8) []const u8 {
+pub fn executeCommand(cmd: RedisCommand, response_buf: []u8) ?[]const u8 {
     switch (cmd.cmd_type) {
         .SET => {
             if (storage.write(cmd.key, cmd.value)) {
@@ -238,7 +248,7 @@ fn buildRedisArray(parts: []const []const u8) []u8 {
 
     buf[pos] = '*';
     pos += 1;
-    pos += formatInt(buf[pos..], parts.len);
+    pos += formatInt(buf[pos..], parts.len).?;
     buf[pos] = '\r';
     buf[pos + 1] = '\n';
     pos += 2;
@@ -246,7 +256,7 @@ fn buildRedisArray(parts: []const []const u8) []u8 {
     for (parts) |part| {
         buf[pos] = '$';
         pos += 1;
-        pos += formatInt(buf[pos..], part.len);
+        pos += formatInt(buf[pos..], part.len).?;
         buf[pos] = '\r';
         buf[pos + 1] = '\n';
         pos += 2;
@@ -310,63 +320,149 @@ test "parseCommand bytes_consumed" {
     try std.testing.expectEqual(input.len, result.bytes_consumed);
 }
 
+test "parseInteger max usize fits" {
+    const s = "18446744073709551615";
+    try std.testing.expectEqual(@as(?usize, std.math.maxInt(usize)), parseInteger(s, 0, s.len));
+}
+
+test "parseInteger overflow returns null" {
+    try std.testing.expectEqual(@as(?usize, null), parseInteger("18446744073709551616", 0, 20));
+    try std.testing.expectEqual(@as(?usize, null), parseInteger("999999999999999999999999999999", 0, 30));
+}
+
+test "parseCommand huge bulk length returns null" {
+    const input = "*2\r\n$3\r\nGET\r\n$18446744073709551615\r\n";
+    try std.testing.expectEqual(@as(?ParseResult, null), parseCommand(input));
+}
+
 test "formatInt zero" {
     var buf: [20]u8 = undefined;
-    const len = formatInt(&buf, 0);
+    const len = formatInt(&buf, 0) orelse return error.TestUnexpectedResult;
     try std.testing.expectEqualStrings("0", buf[0..len]);
 }
 
 test "formatInt positive" {
     var buf: [20]u8 = undefined;
-    const len = formatInt(&buf, 12345);
+    const len = formatInt(&buf, 12345) orelse return error.TestUnexpectedResult;
     try std.testing.expectEqualStrings("12345", buf[0..len]);
 }
 
+test "formatInt respects buffer bounds" {
+    var buf: [4]u8 = undefined;
+    try std.testing.expectEqual(@as(?usize, null), formatInt(buf[0..3], 1234));
+    const len = formatInt(buf[0..4], 1234) orelse return error.TestUnexpectedResult;
+    try std.testing.expectEqualStrings("1234", buf[0..len]);
+}
+
 test "formatSimpleString" {
     var buf: [64]u8 = undefined;
-    const result = formatSimpleString(&buf, "OK");
+    const result = formatSimpleString(&buf, "OK") orelse return error.TestUnexpectedResult;
     try std.testing.expectEqualStrings("+OK\r\n", result);
 }
 
+test "formatSimpleString respects buffer bounds" {
+    var exact: [5]u8 = undefined;
+    const ok = formatSimpleString(exact[0..5], "OK") orelse return error.TestUnexpectedResult;
+    try std.testing.expectEqualStrings("+OK\r\n", ok);
+
+    var short: [4]u8 = undefined;
+    try std.testing.expectEqual(@as(?[]const u8, null), formatSimpleString(short[0..4], "OK"));
+}
+
 test "formatBulkString" {
     var buf: [64]u8 = undefined;
-    const result = formatBulkString(&buf, "hello");
+    const result = formatBulkString(&buf, "hello") orelse return error.TestUnexpectedResult;
     try std.testing.expectEqualStrings("$5\r\nhello\r\n", result);
 }
 
+test "formatBulkString respects buffer bounds" {
+    const value = "hello world";
+    const needed = 1 + intDigits(value.len) + 2 + value.len + 2;
+
+    var exact: [64]u8 = undefined;
+    const ok = formatBulkString(exact[0..needed], value) orelse return error.TestUnexpectedResult;
+    try std.testing.expectEqualStrings("$11\r\nhello world\r\n", ok);
+
+    var short: [64]u8 = undefined;
+    try std.testing.expectEqual(@as(?[]const u8, null), formatBulkString(short[0 .. needed - 1], value));
+}
+
 test "formatNullBulkString" {
     var buf: [64]u8 = undefined;
-    const result = formatNullBulkString(&buf);
+    const result = formatNullBulkString(&buf) orelse return error.TestUnexpectedResult;
     try std.testing.expectEqualStrings("$-1\r\n", result);
 }
 
+test "formatNullBulkString respects buffer bounds" {
+    var exact: [5]u8 = undefined;
+    const ok = formatNullBulkString(exact[0..5]) orelse return error.TestUnexpectedResult;
+    try std.testing.expectEqualStrings("$-1\r\n", ok);
+
+    var short: [4]u8 = undefined;
+    try std.testing.expectEqual(@as(?[]const u8, null), formatNullBulkString(short[0..4]));
+}
+
 test "formatError" {
     var buf: [64]u8 = undefined;
-    const result = formatError(&buf, "ERR bad");
+    const result = formatError(&buf, "ERR bad") orelse return error.TestUnexpectedResult;
     try std.testing.expectEqualStrings("-ERR bad\r\n", result);
 }
 
+test "formatError respects buffer bounds" {
+    const msg = "ERR bad";
+    const needed = 1 + msg.len + 2;
+
+    var exact: [16]u8 = undefined;
+    const ok = formatError(exact[0..needed], msg) orelse return error.TestUnexpectedResult;
+    try std.testing.expectEqualStrings("-ERR bad\r\n", ok);
+
+    var short: [16]u8 = undefined;
+    try std.testing.expectEqual(@as(?[]const u8, null), formatError(short[0 .. needed - 1], msg));
+}
+
 test "formatInteger positive" {
     var buf: [64]u8 = undefined;
-    const result = formatInteger(&buf, 42);
+    const result = formatInteger(&buf, 42) orelse return error.TestUnexpectedResult;
     try std.testing.expectEqualStrings(":42\r\n", result);
 }
 
 test "formatInteger zero" {
     var buf: [64]u8 = undefined;
-    const result = formatInteger(&buf, 0);
+    const result = formatInteger(&buf, 0) orelse return error.TestUnexpectedResult;
     try std.testing.expectEqualStrings(":0\r\n", result);
 }
 
 test "formatInteger negative" {
     var buf: [64]u8 = undefined;
-    const result = formatInteger(&buf, -7);
+    const result = formatInteger(&buf, -7) orelse return error.TestUnexpectedResult;
     try std.testing.expectEqualStrings(":-7\r\n", result);
 }
 
+test "formatInteger respects buffer bounds" {
+    var exact: [16]u8 = undefined;
+    const ok = formatInteger(exact[0..5], 42) orelse return error.TestUnexpectedResult;
+    try std.testing.expectEqualStrings(":42\r\n", ok);
+
+    var short: [16]u8 = undefined;
+    try std.testing.expectEqual(@as(?[]const u8, null), formatInteger(short[0..4], 42));
+}
+
 test "executeCommand UNKNOWN" {
     var buf: [256]u8 = undefined;
     const cmd = RedisCommand{ .cmd_type = .UNKNOWN, .key = "", .value = "" };
-    const result = executeCommand(cmd, &buf);
+    const result = executeCommand(cmd, &buf) orelse return error.TestUnexpectedResult;
     try std.testing.expectEqualStrings("-ERR unknown command\r\n", result);
 }
+
+test "executeCommand GET with insufficient buffer returns null" {
+    storage.init();
+    _ = storage.restore("overflow_key", "this value is far too long to fit in a tiny buffer");
+    const cmd = RedisCommand{ .cmd_type = .GET, .key = "overflow_key", .value = "" };
+
+    var tiny: [16]u8 = undefined;
+    try std.testing.expectEqual(@as(?[]const u8, null), executeCommand(cmd, tiny[0..]));
+
+    var enough: [512]u8 = undefined;
+    const resp = executeCommand(cmd, enough[0..]) orelse return error.TestUnexpectedResult;
+    try std.testing.expect(std.mem.startsWith(u8, resp, "$"));
+}

+ 2 - 2
socket.zig

@@ -18,8 +18,8 @@ pub fn setWriteTimeout(conn: posix.socket_t, seconds: u32) !void {
     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);
+pub fn init(host: []const u8, port: u16) !posix.socket_t {
+    const address = try std.net.Address.parseIp(host, port);
 
     const tpe: u32 = posix.SOCK.STREAM;
     const protocol = posix.IPPROTO.TCP;

+ 20 - 9
storage.zig

@@ -120,11 +120,11 @@ pub fn write(key: []const u8, value: []const u8) bool {
         shards[shard_idx].rwlock.lock();
         defer shards[shard_idx].rwlock.unlock();
 
+        // Persist before publishing the mutation so a durability failure is
+        // never returned after the new value has become visible in memory.
+        persistence.persist('W', key, value) catch return false;
         const entry = writeVolatile(hash, key, value) orelse return false;
-        // Prefix scans must see the index entry before this write becomes
-        // visible to another connection.
         index.insert(entry.key);
-        persistence.persist('W', entry.key, entry.value);
     }
     return true;
 }
@@ -201,15 +201,26 @@ pub fn delete(key: []const u8) bool {
     const hash = hashing.hashKey(key);
     const shard_idx = getShardIndex(hash);
 
-    const deleted = blk: {
-        shards[shard_idx].rwlock.lock();
-        defer shards[shard_idx].rwlock.unlock();
-        break :blk deleteVolatile(hash, key);
-    };
+    shards[shard_idx].rwlock.lock();
+    defer shards[shard_idx].rwlock.unlock();
+
+    const bucket_idx = hash % shards[shard_idx].buckets.len;
+    var current = shards[shard_idx].buckets[bucket_idx];
+    var exists = false;
+    while (current) |entry| {
+        if (entry.hash == hash and std.mem.eql(u8, entry.key, key)) {
+            exists = true;
+            break;
+        }
+        current = entry.next;
+    }
+    if (!exists) return false;
 
+    // Keep the hash table and radix index unchanged if persistence fails.
+    persistence.persist('D', key, "") catch return false;
+    const deleted = deleteVolatile(hash, key);
     if (deleted) {
         index.delete(key);
-        persistence.persist('D', key, "");
     }
 
     return deleted;