Danilo Fragoso 9 mesiacov pred
rodič
commit
8c35017f18
4 zmenil súbory, kde vykonal 363 pridanie a 20 odobranie
  1. 16 5
      index.zig
  2. 81 4
      main.zig
  3. 35 11
      persistence.zig
  4. 231 0
      redis.zig

+ 16 - 5
index.zig

@@ -181,12 +181,13 @@ fn countKeys(node: *RadixNode) usize {
 }
 
 const MAX_KEYS_RETURN = 10000;
+const MAX_KEY_LENGTH = 1024;
 
-fn collectKeys(node: *RadixNode, prefix: []const u8, keys: *std.ArrayListUnmanaged([]const u8), max_keys: usize) void {
+fn collectKeysWithBuffer(node: *RadixNode, prefix_buffer: []u8, prefix_len: usize, keys: *std.ArrayListUnmanaged([]const u8), max_keys: usize) void {
     if (keys.items.len >= max_keys) return;
 
     if (node.is_terminal) {
-        const key = temp_allocator.dupe(u8, prefix) catch return;
+        const key = temp_allocator.dupe(u8, prefix_buffer[0..prefix_len]) catch return;
         keys.append(temp_allocator, key) catch return;
     }
 
@@ -194,12 +195,22 @@ fn collectKeys(node: *RadixNode, prefix: []const u8, keys: *std.ArrayListUnmanag
     while (it.next()) |entry| {
         if (keys.items.len >= max_keys) break;
         const child = entry.value_ptr.*;
-        const new_prefix = std.mem.concat(temp_allocator, u8, &[_][]const u8{ prefix, child.edge }) catch return;
-        collectKeys(child, new_prefix, keys, max_keys);
-        temp_allocator.free(new_prefix);
+        const edge_len = child.edge.len;
+
+        if (prefix_len + edge_len > MAX_KEY_LENGTH) continue;
+
+        @memcpy(prefix_buffer[prefix_len .. prefix_len + edge_len], child.edge);
+        collectKeysWithBuffer(child, prefix_buffer, prefix_len + edge_len, keys, max_keys);
     }
 }
 
+fn collectKeys(node: *RadixNode, prefix: []const u8, keys: *std.ArrayListUnmanaged([]const u8), max_keys: usize) 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);
+}
+
 pub fn getKeysFromNode(node: *RadixNode, prefix: []const u8) [][]const u8 {
     var keys_list = std.ArrayListUnmanaged([]const u8){};
     collectKeys(node, prefix, &keys_list, MAX_KEYS_RETURN);

+ 81 - 4
main.zig

@@ -8,15 +8,28 @@ const socket = @import("socket.zig");
 const command = @import("command.zig");
 const storage = @import("storage.zig");
 const persistence = @import("persistence.zig");
+const redis = @import("redis.zig");
 
 const PORT = 8085;
 var should_exit = std.atomic.Value(bool).init(false);
+var redis_mode = false;
+
 fn handleSignal(sig: c_int) callconv(.c) void {
     _ = sig;
     should_exit.store(true, .seq_cst);
 }
 
 pub fn main() !void {
+    var args = try std.process.argsWithAllocator(std.heap.page_allocator);
+    defer args.deinit();
+
+    _ = args.skip();
+    while (args.next()) |arg| {
+        if (std.mem.eql(u8, arg, "-redis")) {
+            redis_mode = true;
+        }
+    }
+
     const empty_mask = std.mem.zeroes(posix.sigset_t);
     const act = posix.Sigaction{
         .handler = .{ .handler = handleSignal },
@@ -31,7 +44,11 @@ pub fn main() !void {
     defer posix.close(listener);
 
     std.debug.print("2025 pizzakv! TCP Listening on port {any}\n<danilo@fragoso.dev>\n---------\n", .{PORT});
-    std.debug.print("Commands:\n\nread key\nwrite key|value\ndelete key\nkeys\nreads prefix\nstatus\n", .{});
+    if (redis_mode) {
+        std.debug.print("Mode: Redis Protocol (RESP)\nCommands: SET, GET, DEL\n", .{});
+    } else {
+        std.debug.print("Commands:\n\nread key\nwrite key|value\ndelete key\nkeys\nreads prefix\nstatus\n", .{});
+    }
     std.debug.print("---------\n", .{});
 
     storage.init();
@@ -46,7 +63,7 @@ pub fn main() !void {
             },
         };
 
-        const ready = posix.poll(&poll_fds, 1000) catch |err| {
+        const ready = posix.poll(&poll_fds, 100) catch |err| {
             if (should_exit.load(.seq_cst)) break;
             std.debug.print("poll error: {any}\n", .{err});
             continue;
@@ -74,8 +91,13 @@ pub fn main() !void {
 
         posix.setsockopt(conn, posix.IPPROTO.TCP, posix.TCP.NODELAY, &std.mem.toBytes(@as(c_int, 1))) catch {};
 
-        const thread = try std.Thread.spawn(.{}, handleConnection, .{conn});
-        thread.detach();
+        if (redis_mode) {
+            const thread = try std.Thread.spawn(.{}, handleRedisConnection, .{conn});
+            thread.detach();
+        } else {
+            const thread = try std.Thread.spawn(.{}, handleConnection, .{conn});
+            thread.detach();
+        }
     }
 
     std.debug.print("\nShutdown signal received...\n", .{});
@@ -115,3 +137,58 @@ pub fn handleConnection(conn: posix.socket_t) !void {
         };
     }
 }
+
+pub fn handleRedisConnection(conn: posix.socket_t) !void {
+    defer posix.close(conn);
+
+    var requestBuffer: [2 * 1024 * 1024]u8 = undefined;
+    var responseBuffer: [2 * 1024 * 1024]u8 = undefined;
+    var buffered_len: usize = 0;
+
+    const is_darwin = @import("builtin").target.os.tag == .macos;
+    const cork_option = if (is_darwin) posix.TCP.NOPUSH else posix.TCP.CORK;
+
+    while (true) {
+        const n = posix.read(conn, requestBuffer[buffered_len..]) catch |err| {
+            if (err == error.ConnectionResetByPeer) break;
+            return err;
+        };
+        if (n == 0) break;
+
+        const total_len = buffered_len + n;
+        var offset: usize = 0;
+        var response_offset: usize = 0;
+
+        posix.setsockopt(conn, posix.IPPROTO.TCP, cork_option, &std.mem.toBytes(@as(c_int, 1))) catch {};
+
+        while (offset < total_len) {
+            const result = redis.parseCommand(requestBuffer[offset..total_len]) orelse {
+                break;
+            };
+
+            const response = redis.executeCommand(result.cmd, responseBuffer[response_offset..]);
+            response_offset += 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});
+            };
+        }
+
+        if (offset < total_len) {
+            const remaining = total_len - offset;
+            if (remaining > 0 and remaining < requestBuffer.len / 2) {
+                @memcpy(requestBuffer[0..remaining], requestBuffer[offset..total_len]);
+                buffered_len = remaining;
+            } else {
+                buffered_len = 0;
+            }
+        } else {
+            buffered_len = 0;
+        }
+    }
+}

+ 35 - 11
persistence.zig

@@ -75,39 +75,63 @@ pub fn init() !void {
 }
 
 pub fn persist(opcode: u8, key: []const u8, value: []const u8) void {
-    const record = std.fmt.allocPrint(c_allocator, "{c}|{s}|{s}\r", .{ opcode, key, value }) catch {
-        std.debug.print("Failed to format record for persistence\n", .{});
-        return;
-    };
-    defer c_allocator.free(record);
+    const record_len = 1 + 1 + key.len + 1 + value.len + 1;
 
     mutex.lock();
     defer mutex.unlock();
 
-    if (buffer_position + record.len > FLUSH_THRESHOLD) {
+    if (buffer_position + record_len > FLUSH_THRESHOLD) {
         flushBuffer() catch |err| {
             std.debug.print("Failed to flush buffer: {any}\n", .{err});
             return;
         };
     }
 
-    if (record.len > BUFFER_SIZE) {
-        _ = storage_file.?.write(record) catch |err| {
+    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;
     }
 
-    if (buffer_position + record.len > BUFFER_SIZE) {
+    if (buffer_position + record_len > BUFFER_SIZE) {
         flushBuffer() catch |err| {
             std.debug.print("Failed to flush buffer: {any}\n", .{err});
             return;
         };
     }
 
-    @memcpy(write_buffer[buffer_position .. buffer_position + record.len], record);
-    buffer_position += record.len;
+    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;
 }
 
 pub fn flush() !void {

+ 231 - 0
redis.zig

@@ -0,0 +1,231 @@
+const std = @import("std");
+const storage = @import("storage.zig");
+
+const CommandType = enum {
+    SET,
+    GET,
+    DEL,
+    UNKNOWN,
+};
+
+pub const RedisCommand = struct {
+    cmd_type: CommandType,
+    key: []const u8,
+    value: []const u8,
+};
+
+pub const ParseResult = struct {
+    cmd: RedisCommand,
+    bytes_consumed: usize,
+};
+
+fn parseInteger(buf: []const u8, start: usize, end: usize) ?usize {
+    if (start >= end) return null;
+    var result: usize = 0;
+    for (buf[start..end]) |c| {
+        if (c < '0' or c > '9') return null;
+        result = result * 10 + (c - '0');
+    }
+    return result;
+}
+
+fn parseBulkString(buf: []const u8, pos: *usize) ?[]const u8 {
+    if (pos.* >= buf.len or buf[pos.*] != '$') return null;
+    pos.* += 1;
+
+    const len_end = std.mem.indexOfScalarPos(u8, buf, pos.*, '\r') orelse return null;
+    const len = parseInteger(buf, pos.*, len_end) orelse return null;
+    pos.* = len_end + 2;
+
+    const str_start = pos.*;
+    const str_end = str_start + len;
+    if (str_end > buf.len) return null;
+
+    const result = buf[str_start..str_end];
+    pos.* = str_end + 2;
+
+    return result;
+}
+
+pub fn parseCommand(buf: []const u8) ?ParseResult {
+    if (buf.len == 0) return null;
+
+    var pos: usize = 0;
+
+    if (buf[pos] != '*') return null;
+    pos += 1;
+
+    const array_len_end = std.mem.indexOfScalarPos(u8, buf, pos, '\r') orelse return null;
+    const array_len = parseInteger(buf, pos, array_len_end) orelse return null;
+    pos = array_len_end + 2;
+
+    if (array_len < 1 or array_len > 16) return null;
+
+    var elements: [16][]const u8 = undefined;
+    for (0..array_len) |i| {
+        elements[i] = parseBulkString(buf, &pos) orelse return null;
+    }
+
+    const cmd_str = elements[0];
+    var cmd: RedisCommand = undefined;
+
+    if (cmd_str.len == 3) {
+        const upper: u32 = (@as(u32, cmd_str[0]) & 0xDF) << 16 | (@as(u32, cmd_str[1]) & 0xDF) << 8 | (@as(u32, cmd_str[2]) & 0xDF);
+        if (upper == (@as(u32, 'S') << 16 | @as(u32, 'E') << 8 | @as(u32, 'T'))) {
+            if (array_len < 3) return null;
+            cmd = RedisCommand{
+                .cmd_type = .SET,
+                .key = elements[1],
+                .value = elements[2],
+            };
+        } else if (upper == (@as(u32, 'G') << 16 | @as(u32, 'E') << 8 | @as(u32, 'T'))) {
+            if (array_len < 2) return null;
+            cmd = RedisCommand{
+                .cmd_type = .GET,
+                .key = elements[1],
+                .value = "",
+            };
+        } else if (upper == (@as(u32, 'D') << 16 | @as(u32, 'E') << 8 | @as(u32, 'L'))) {
+            if (array_len < 2) return null;
+            cmd = RedisCommand{
+                .cmd_type = .DEL,
+                .key = elements[1],
+                .value = "",
+            };
+        } else {
+            cmd = RedisCommand{
+                .cmd_type = .UNKNOWN,
+                .key = "",
+                .value = "",
+            };
+        }
+    } else {
+        cmd = RedisCommand{
+            .cmd_type = .UNKNOWN,
+            .key = "",
+            .value = "",
+        };
+    }
+
+    return ParseResult{
+        .cmd = cmd,
+        .bytes_consumed = pos,
+    };
+}
+
+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 formatBulkString(buf: []u8, str: []const u8) []const u8 {
+    var pos: usize = 0;
+    buf[pos] = '$';
+    pos += 1;
+
+    pos += formatInt(buf[pos..], str.len);
+    buf[pos] = '\r';
+    buf[pos + 1] = '\n';
+    pos += 2;
+
+    @memcpy(buf[pos .. pos + str.len], str);
+    pos += str.len;
+    buf[pos] = '\r';
+    buf[pos + 1] = '\n';
+
+    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 {
+    buf[0] = '$';
+    buf[1] = '-';
+    buf[2] = '1';
+    buf[3] = '\r';
+    buf[4] = '\n';
+    return buf[0..5];
+}
+
+fn formatInteger(buf: []u8, value: i64) []const u8 {
+    var pos: usize = 0;
+    buf[pos] = ':';
+    pos += 1;
+
+    if (value < 0) {
+        buf[pos] = '-';
+        pos += 1;
+        pos += formatInt(buf[pos..], @intCast(-value));
+    } else {
+        pos += formatInt(buf[pos..], @intCast(value));
+    }
+
+    buf[pos] = '\r';
+    buf[pos + 1] = '\n';
+    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 executeCommand(cmd: RedisCommand, response_buf: []u8) []const u8 {
+    switch (cmd.cmd_type) {
+        .SET => {
+            if (storage.write(cmd.key, cmd.value)) {
+                return formatSimpleString(response_buf, "OK");
+            } else {
+                return formatError(response_buf, "ERR write failed");
+            }
+        },
+        .GET => {
+            if (storage.read(cmd.key)) |value| {
+                return formatBulkString(response_buf, value);
+            } else {
+                return formatNullBulkString(response_buf);
+            }
+        },
+        .DEL => {
+            const deleted = storage.delete(cmd.key);
+            return formatInteger(response_buf, if (deleted) 1 else 0);
+        },
+        .UNKNOWN => {
+            return formatError(response_buf, "ERR unknown command");
+        },
+    }
+}