redis.zig 10 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372
  1. const std = @import("std");
  2. const storage = @import("storage.zig");
  3. const CommandType = enum {
  4. SET,
  5. GET,
  6. DEL,
  7. UNKNOWN,
  8. };
  9. pub const RedisCommand = struct {
  10. cmd_type: CommandType,
  11. key: []const u8,
  12. value: []const u8,
  13. };
  14. pub const ParseResult = struct {
  15. cmd: RedisCommand,
  16. bytes_consumed: usize,
  17. };
  18. fn parseInteger(buf: []const u8, start: usize, end: usize) ?usize {
  19. if (start >= end) return null;
  20. var result: usize = 0;
  21. for (buf[start..end]) |c| {
  22. if (c < '0' or c > '9') return null;
  23. result = result * 10 + (c - '0');
  24. }
  25. return result;
  26. }
  27. fn parseBulkString(buf: []const u8, pos: *usize) ?[]const u8 {
  28. if (pos.* >= buf.len or buf[pos.*] != '$') return null;
  29. pos.* += 1;
  30. const len_end = std.mem.indexOfScalarPos(u8, buf, pos.*, '\r') orelse return null;
  31. const len = parseInteger(buf, pos.*, len_end) orelse return null;
  32. pos.* = len_end + 2;
  33. const str_start = pos.*;
  34. const str_end = str_start + len;
  35. if (str_end > buf.len) return null;
  36. const result = buf[str_start..str_end];
  37. pos.* = str_end + 2;
  38. return result;
  39. }
  40. pub fn parseCommand(buf: []const u8) ?ParseResult {
  41. if (buf.len == 0) return null;
  42. var pos: usize = 0;
  43. if (buf[pos] != '*') return null;
  44. pos += 1;
  45. const array_len_end = std.mem.indexOfScalarPos(u8, buf, pos, '\r') orelse return null;
  46. const array_len = parseInteger(buf, pos, array_len_end) orelse return null;
  47. pos = array_len_end + 2;
  48. if (array_len < 1 or array_len > 16) return null;
  49. var elements: [16][]const u8 = undefined;
  50. for (0..array_len) |i| {
  51. elements[i] = parseBulkString(buf, &pos) orelse return null;
  52. }
  53. const cmd_str = elements[0];
  54. var cmd: RedisCommand = undefined;
  55. if (cmd_str.len == 3) {
  56. const upper: u32 = (@as(u32, cmd_str[0]) & 0xDF) << 16 | (@as(u32, cmd_str[1]) & 0xDF) << 8 | (@as(u32, cmd_str[2]) & 0xDF);
  57. if (upper == (@as(u32, 'S') << 16 | @as(u32, 'E') << 8 | @as(u32, 'T'))) {
  58. if (array_len < 3) return null;
  59. cmd = RedisCommand{
  60. .cmd_type = .SET,
  61. .key = elements[1],
  62. .value = elements[2],
  63. };
  64. } else if (upper == (@as(u32, 'G') << 16 | @as(u32, 'E') << 8 | @as(u32, 'T'))) {
  65. if (array_len < 2) return null;
  66. cmd = RedisCommand{
  67. .cmd_type = .GET,
  68. .key = elements[1],
  69. .value = "",
  70. };
  71. } else if (upper == (@as(u32, 'D') << 16 | @as(u32, 'E') << 8 | @as(u32, 'L'))) {
  72. if (array_len < 2) return null;
  73. cmd = RedisCommand{
  74. .cmd_type = .DEL,
  75. .key = elements[1],
  76. .value = "",
  77. };
  78. } else {
  79. cmd = RedisCommand{
  80. .cmd_type = .UNKNOWN,
  81. .key = "",
  82. .value = "",
  83. };
  84. }
  85. } else {
  86. cmd = RedisCommand{
  87. .cmd_type = .UNKNOWN,
  88. .key = "",
  89. .value = "",
  90. };
  91. }
  92. return ParseResult{
  93. .cmd = cmd,
  94. .bytes_consumed = pos,
  95. };
  96. }
  97. fn formatSimpleString(buf: []u8, str: []const u8) []const u8 {
  98. var pos: usize = 0;
  99. buf[pos] = '+';
  100. pos += 1;
  101. @memcpy(buf[pos .. pos + str.len], str);
  102. pos += str.len;
  103. buf[pos] = '\r';
  104. buf[pos + 1] = '\n';
  105. return buf[0 .. pos + 2];
  106. }
  107. fn formatBulkString(buf: []u8, str: []const u8) []const u8 {
  108. var pos: usize = 0;
  109. buf[pos] = '$';
  110. pos += 1;
  111. pos += formatInt(buf[pos..], str.len);
  112. buf[pos] = '\r';
  113. buf[pos + 1] = '\n';
  114. pos += 2;
  115. @memcpy(buf[pos .. pos + str.len], str);
  116. pos += str.len;
  117. buf[pos] = '\r';
  118. buf[pos + 1] = '\n';
  119. return buf[0 .. pos + 2];
  120. }
  121. fn formatInt(buf: []u8, value: usize) usize {
  122. if (value == 0) {
  123. buf[0] = '0';
  124. return 1;
  125. }
  126. var v = value;
  127. var len: usize = 0;
  128. var temp: [20]u8 = undefined;
  129. while (v > 0) {
  130. temp[len] = @intCast('0' + (v % 10));
  131. v /= 10;
  132. len += 1;
  133. }
  134. var i: usize = 0;
  135. while (i < len) : (i += 1) {
  136. buf[i] = temp[len - 1 - i];
  137. }
  138. return len;
  139. }
  140. fn formatNullBulkString(buf: []u8) []const u8 {
  141. buf[0] = '$';
  142. buf[1] = '-';
  143. buf[2] = '1';
  144. buf[3] = '\r';
  145. buf[4] = '\n';
  146. return buf[0..5];
  147. }
  148. fn formatInteger(buf: []u8, value: i64) []const u8 {
  149. var pos: usize = 0;
  150. buf[pos] = ':';
  151. pos += 1;
  152. if (value < 0) {
  153. buf[pos] = '-';
  154. pos += 1;
  155. pos += formatInt(buf[pos..], @intCast(-value));
  156. } else {
  157. pos += formatInt(buf[pos..], @intCast(value));
  158. }
  159. buf[pos] = '\r';
  160. buf[pos + 1] = '\n';
  161. return buf[0 .. pos + 2];
  162. }
  163. fn formatError(buf: []u8, msg: []const u8) []const u8 {
  164. var pos: usize = 0;
  165. buf[pos] = '-';
  166. pos += 1;
  167. @memcpy(buf[pos .. pos + msg.len], msg);
  168. pos += msg.len;
  169. buf[pos] = '\r';
  170. buf[pos + 1] = '\n';
  171. return buf[0 .. pos + 2];
  172. }
  173. pub fn executeCommand(cmd: RedisCommand, response_buf: []u8) []const u8 {
  174. switch (cmd.cmd_type) {
  175. .SET => {
  176. if (storage.write(cmd.key, cmd.value)) {
  177. return formatSimpleString(response_buf, "OK");
  178. } else {
  179. return formatError(response_buf, "ERR write failed");
  180. }
  181. },
  182. .GET => {
  183. if (storage.read(cmd.key)) |value| {
  184. return formatBulkString(response_buf, value);
  185. } else {
  186. return formatNullBulkString(response_buf);
  187. }
  188. },
  189. .DEL => {
  190. const deleted = storage.delete(cmd.key);
  191. return formatInteger(response_buf, if (deleted) 1 else 0);
  192. },
  193. .UNKNOWN => {
  194. return formatError(response_buf, "ERR unknown command");
  195. },
  196. }
  197. }
  198. // -- Tests --
  199. fn buildRedisArray(parts: []const []const u8) []u8 {
  200. var buf: [4096]u8 = undefined;
  201. var pos: usize = 0;
  202. buf[pos] = '*';
  203. pos += 1;
  204. pos += formatInt(buf[pos..], parts.len);
  205. buf[pos] = '\r';
  206. buf[pos + 1] = '\n';
  207. pos += 2;
  208. for (parts) |part| {
  209. buf[pos] = '$';
  210. pos += 1;
  211. pos += formatInt(buf[pos..], part.len);
  212. buf[pos] = '\r';
  213. buf[pos + 1] = '\n';
  214. pos += 2;
  215. @memcpy(buf[pos .. pos + part.len], part);
  216. pos += part.len;
  217. buf[pos] = '\r';
  218. buf[pos + 1] = '\n';
  219. pos += 2;
  220. }
  221. return buf[0..pos];
  222. }
  223. test "parseCommand SET" {
  224. const input = buildRedisArray(&.{ "SET", "mykey", "myvalue" });
  225. const result = parseCommand(input) orelse return error.TestUnexpectedResult;
  226. try std.testing.expectEqual(CommandType.SET, result.cmd.cmd_type);
  227. try std.testing.expectEqualStrings("mykey", result.cmd.key);
  228. try std.testing.expectEqualStrings("myvalue", result.cmd.value);
  229. }
  230. test "parseCommand GET" {
  231. const input = buildRedisArray(&.{ "GET", "mykey" });
  232. const result = parseCommand(input) orelse return error.TestUnexpectedResult;
  233. try std.testing.expectEqual(CommandType.GET, result.cmd.cmd_type);
  234. try std.testing.expectEqualStrings("mykey", result.cmd.key);
  235. }
  236. test "parseCommand DEL" {
  237. const input = buildRedisArray(&.{ "DEL", "mykey" });
  238. const result = parseCommand(input) orelse return error.TestUnexpectedResult;
  239. try std.testing.expectEqual(CommandType.DEL, result.cmd.cmd_type);
  240. try std.testing.expectEqualStrings("mykey", result.cmd.key);
  241. }
  242. test "parseCommand case insensitive" {
  243. const input = buildRedisArray(&.{ "set", "k", "v" });
  244. const result = parseCommand(input) orelse return error.TestUnexpectedResult;
  245. try std.testing.expectEqual(CommandType.SET, result.cmd.cmd_type);
  246. }
  247. test "parseCommand unknown command" {
  248. const input = buildRedisArray(&.{ "FOO", "bar" });
  249. const result = parseCommand(input) orelse return error.TestUnexpectedResult;
  250. try std.testing.expectEqual(CommandType.UNKNOWN, result.cmd.cmd_type);
  251. }
  252. test "parseCommand empty input" {
  253. try std.testing.expectEqual(@as(?ParseResult, null), parseCommand(""));
  254. }
  255. test "parseCommand malformed input" {
  256. try std.testing.expectEqual(@as(?ParseResult, null), parseCommand("garbage"));
  257. try std.testing.expectEqual(@as(?ParseResult, null), parseCommand("*"));
  258. try std.testing.expectEqual(@as(?ParseResult, null), parseCommand("*1\r\n"));
  259. }
  260. test "parseCommand bytes_consumed" {
  261. const input = buildRedisArray(&.{ "GET", "key1" });
  262. const result = parseCommand(input) orelse return error.TestUnexpectedResult;
  263. try std.testing.expectEqual(input.len, result.bytes_consumed);
  264. }
  265. test "formatInt zero" {
  266. var buf: [20]u8 = undefined;
  267. const len = formatInt(&buf, 0);
  268. try std.testing.expectEqualStrings("0", buf[0..len]);
  269. }
  270. test "formatInt positive" {
  271. var buf: [20]u8 = undefined;
  272. const len = formatInt(&buf, 12345);
  273. try std.testing.expectEqualStrings("12345", buf[0..len]);
  274. }
  275. test "formatSimpleString" {
  276. var buf: [64]u8 = undefined;
  277. const result = formatSimpleString(&buf, "OK");
  278. try std.testing.expectEqualStrings("+OK\r\n", result);
  279. }
  280. test "formatBulkString" {
  281. var buf: [64]u8 = undefined;
  282. const result = formatBulkString(&buf, "hello");
  283. try std.testing.expectEqualStrings("$5\r\nhello\r\n", result);
  284. }
  285. test "formatNullBulkString" {
  286. var buf: [64]u8 = undefined;
  287. const result = formatNullBulkString(&buf);
  288. try std.testing.expectEqualStrings("$-1\r\n", result);
  289. }
  290. test "formatError" {
  291. var buf: [64]u8 = undefined;
  292. const result = formatError(&buf, "ERR bad");
  293. try std.testing.expectEqualStrings("-ERR bad\r\n", result);
  294. }
  295. test "formatInteger positive" {
  296. var buf: [64]u8 = undefined;
  297. const result = formatInteger(&buf, 42);
  298. try std.testing.expectEqualStrings(":42\r\n", result);
  299. }
  300. test "formatInteger zero" {
  301. var buf: [64]u8 = undefined;
  302. const result = formatInteger(&buf, 0);
  303. try std.testing.expectEqualStrings(":0\r\n", result);
  304. }
  305. test "formatInteger negative" {
  306. var buf: [64]u8 = undefined;
  307. const result = formatInteger(&buf, -7);
  308. try std.testing.expectEqualStrings(":-7\r\n", result);
  309. }
  310. test "executeCommand UNKNOWN" {
  311. var buf: [256]u8 = undefined;
  312. const cmd = RedisCommand{ .cmd_type = .UNKNOWN, .key = "", .value = "" };
  313. const result = executeCommand(cmd, &buf);
  314. try std.testing.expectEqualStrings("-ERR unknown command\r\n", result);
  315. }