2
0

redis.zig 15 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381382383384385386387388389390391392393394395396397398399400401402403404405406407408409410411412413414415416417418419420421422423424425426427428429430431432433434435436437438439440441442443444445446447448449450451452453454455456457458459460461462463464465466467468
  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. const digit: usize = c - '0';
  24. result = std.math.mul(usize, result, 10) catch return null;
  25. result = std.math.add(usize, result, digit) catch return null;
  26. }
  27. return result;
  28. }
  29. fn parseBulkString(buf: []const u8, pos: *usize) ?[]const u8 {
  30. if (pos.* >= buf.len or buf[pos.*] != '$') return null;
  31. pos.* += 1;
  32. const len_end = std.mem.indexOfScalarPos(u8, buf, pos.*, '\r') orelse return null;
  33. const len = parseInteger(buf, pos.*, len_end) orelse return null;
  34. pos.* = len_end + 2;
  35. const str_start = pos.*;
  36. const str_end = std.math.add(usize, str_start, len) catch return null;
  37. if (str_end > buf.len) return null;
  38. const result = buf[str_start..str_end];
  39. pos.* = str_end + 2;
  40. return result;
  41. }
  42. pub fn parseCommand(buf: []const u8) ?ParseResult {
  43. if (buf.len == 0) return null;
  44. var pos: usize = 0;
  45. if (buf[pos] != '*') return null;
  46. pos += 1;
  47. const array_len_end = std.mem.indexOfScalarPos(u8, buf, pos, '\r') orelse return null;
  48. const array_len = parseInteger(buf, pos, array_len_end) orelse return null;
  49. pos = array_len_end + 2;
  50. if (array_len < 1 or array_len > 16) return null;
  51. var elements: [16][]const u8 = undefined;
  52. for (0..array_len) |i| {
  53. elements[i] = parseBulkString(buf, &pos) orelse return null;
  54. }
  55. const cmd_str = elements[0];
  56. var cmd: RedisCommand = undefined;
  57. if (cmd_str.len == 3) {
  58. const upper: u32 = (@as(u32, cmd_str[0]) & 0xDF) << 16 | (@as(u32, cmd_str[1]) & 0xDF) << 8 | (@as(u32, cmd_str[2]) & 0xDF);
  59. if (upper == (@as(u32, 'S') << 16 | @as(u32, 'E') << 8 | @as(u32, 'T'))) {
  60. if (array_len < 3) return null;
  61. cmd = RedisCommand{
  62. .cmd_type = .SET,
  63. .key = elements[1],
  64. .value = elements[2],
  65. };
  66. } else if (upper == (@as(u32, 'G') << 16 | @as(u32, 'E') << 8 | @as(u32, 'T'))) {
  67. if (array_len < 2) return null;
  68. cmd = RedisCommand{
  69. .cmd_type = .GET,
  70. .key = elements[1],
  71. .value = "",
  72. };
  73. } else if (upper == (@as(u32, 'D') << 16 | @as(u32, 'E') << 8 | @as(u32, 'L'))) {
  74. if (array_len < 2) return null;
  75. cmd = RedisCommand{
  76. .cmd_type = .DEL,
  77. .key = elements[1],
  78. .value = "",
  79. };
  80. } else {
  81. cmd = RedisCommand{
  82. .cmd_type = .UNKNOWN,
  83. .key = "",
  84. .value = "",
  85. };
  86. }
  87. } else {
  88. cmd = RedisCommand{
  89. .cmd_type = .UNKNOWN,
  90. .key = "",
  91. .value = "",
  92. };
  93. }
  94. return ParseResult{
  95. .cmd = cmd,
  96. .bytes_consumed = pos,
  97. };
  98. }
  99. fn intDigits(value: usize) usize {
  100. if (value == 0) return 1;
  101. var v = value;
  102. var d: usize = 0;
  103. while (v > 0) : (v /= 10) d += 1;
  104. return d;
  105. }
  106. fn formatInt(buf: []u8, value: usize) ?usize {
  107. if (value == 0) {
  108. if (buf.len < 1) return null;
  109. buf[0] = '0';
  110. return 1;
  111. }
  112. const len = intDigits(value);
  113. if (buf.len < len) return null;
  114. var v = value;
  115. var i: usize = len;
  116. while (i > 0) {
  117. i -= 1;
  118. buf[i] = @intCast('0' + (v % 10));
  119. v /= 10;
  120. }
  121. return len;
  122. }
  123. fn formatSimpleString(buf: []u8, str: []const u8) ?[]const u8 {
  124. const needed = 1 + str.len + 2;
  125. if (buf.len < needed) return null;
  126. buf[0] = '+';
  127. @memcpy(buf[1 .. 1 + str.len], str);
  128. buf[1 + str.len] = '\r';
  129. buf[2 + str.len] = '\n';
  130. return buf[0..needed];
  131. }
  132. fn formatBulkString(buf: []u8, str: []const u8) ?[]const u8 {
  133. const needed = 1 + intDigits(str.len) + 2 + str.len + 2;
  134. if (buf.len < needed) return null;
  135. buf[0] = '$';
  136. var pos: usize = 1;
  137. pos += formatInt(buf[pos..], str.len) orelse return null;
  138. buf[pos] = '\r';
  139. buf[pos + 1] = '\n';
  140. pos += 2;
  141. @memcpy(buf[pos .. pos + str.len], str);
  142. pos += str.len;
  143. buf[pos] = '\r';
  144. buf[pos + 1] = '\n';
  145. return buf[0 .. pos + 2];
  146. }
  147. fn formatNullBulkString(buf: []u8) ?[]const u8 {
  148. if (buf.len < 5) return null;
  149. buf[0] = '$';
  150. buf[1] = '-';
  151. buf[2] = '1';
  152. buf[3] = '\r';
  153. buf[4] = '\n';
  154. return buf[0..5];
  155. }
  156. fn formatInteger(buf: []u8, value: i64) ?[]const u8 {
  157. var pos: usize = 1;
  158. if (value < 0) {
  159. const needed = 1 + 1 + intDigits(@intCast(-value)) + 2;
  160. if (buf.len < needed) return null;
  161. buf[0] = ':';
  162. buf[1] = '-';
  163. pos = 2;
  164. pos += formatInt(buf[pos..], @intCast(-value)) orelse return null;
  165. } else {
  166. const needed = 1 + intDigits(@intCast(value)) + 2;
  167. if (buf.len < needed) return null;
  168. buf[0] = ':';
  169. pos = 1;
  170. pos += formatInt(buf[pos..], @intCast(value)) orelse return null;
  171. }
  172. buf[pos] = '\r';
  173. buf[pos + 1] = '\n';
  174. return buf[0 .. pos + 2];
  175. }
  176. pub fn formatError(buf: []u8, msg: []const u8) ?[]const u8 {
  177. const needed = 1 + msg.len + 2;
  178. if (buf.len < needed) return null;
  179. buf[0] = '-';
  180. @memcpy(buf[1 .. 1 + msg.len], msg);
  181. buf[1 + msg.len] = '\r';
  182. buf[2 + msg.len] = '\n';
  183. return buf[0..needed];
  184. }
  185. pub fn executeCommand(cmd: RedisCommand, response_buf: []u8) ?[]const u8 {
  186. switch (cmd.cmd_type) {
  187. .SET => {
  188. if (storage.write(cmd.key, cmd.value)) {
  189. return formatSimpleString(response_buf, "OK");
  190. } else {
  191. return formatError(response_buf, "ERR write failed");
  192. }
  193. },
  194. .GET => {
  195. if (storage.read(cmd.key)) |value| {
  196. return formatBulkString(response_buf, value);
  197. } else {
  198. return formatNullBulkString(response_buf);
  199. }
  200. },
  201. .DEL => {
  202. const deleted = storage.delete(cmd.key);
  203. return formatInteger(response_buf, if (deleted) 1 else 0);
  204. },
  205. .UNKNOWN => {
  206. return formatError(response_buf, "ERR unknown command");
  207. },
  208. }
  209. }
  210. // -- Tests --
  211. fn buildRedisArray(parts: []const []const u8) []u8 {
  212. var buf: [4096]u8 = undefined;
  213. var pos: usize = 0;
  214. buf[pos] = '*';
  215. pos += 1;
  216. pos += formatInt(buf[pos..], parts.len).?;
  217. buf[pos] = '\r';
  218. buf[pos + 1] = '\n';
  219. pos += 2;
  220. for (parts) |part| {
  221. buf[pos] = '$';
  222. pos += 1;
  223. pos += formatInt(buf[pos..], part.len).?;
  224. buf[pos] = '\r';
  225. buf[pos + 1] = '\n';
  226. pos += 2;
  227. @memcpy(buf[pos .. pos + part.len], part);
  228. pos += part.len;
  229. buf[pos] = '\r';
  230. buf[pos + 1] = '\n';
  231. pos += 2;
  232. }
  233. return buf[0..pos];
  234. }
  235. test "parseCommand SET" {
  236. const input = buildRedisArray(&.{ "SET", "mykey", "myvalue" });
  237. const result = parseCommand(input) orelse return error.TestUnexpectedResult;
  238. try std.testing.expectEqual(CommandType.SET, result.cmd.cmd_type);
  239. try std.testing.expectEqualStrings("mykey", result.cmd.key);
  240. try std.testing.expectEqualStrings("myvalue", result.cmd.value);
  241. }
  242. test "parseCommand GET" {
  243. const input = buildRedisArray(&.{ "GET", "mykey" });
  244. const result = parseCommand(input) orelse return error.TestUnexpectedResult;
  245. try std.testing.expectEqual(CommandType.GET, result.cmd.cmd_type);
  246. try std.testing.expectEqualStrings("mykey", result.cmd.key);
  247. }
  248. test "parseCommand DEL" {
  249. const input = buildRedisArray(&.{ "DEL", "mykey" });
  250. const result = parseCommand(input) orelse return error.TestUnexpectedResult;
  251. try std.testing.expectEqual(CommandType.DEL, result.cmd.cmd_type);
  252. try std.testing.expectEqualStrings("mykey", result.cmd.key);
  253. }
  254. test "parseCommand case insensitive" {
  255. const input = buildRedisArray(&.{ "set", "k", "v" });
  256. const result = parseCommand(input) orelse return error.TestUnexpectedResult;
  257. try std.testing.expectEqual(CommandType.SET, result.cmd.cmd_type);
  258. }
  259. test "parseCommand unknown command" {
  260. const input = buildRedisArray(&.{ "FOO", "bar" });
  261. const result = parseCommand(input) orelse return error.TestUnexpectedResult;
  262. try std.testing.expectEqual(CommandType.UNKNOWN, result.cmd.cmd_type);
  263. }
  264. test "parseCommand empty input" {
  265. try std.testing.expectEqual(@as(?ParseResult, null), parseCommand(""));
  266. }
  267. test "parseCommand malformed input" {
  268. try std.testing.expectEqual(@as(?ParseResult, null), parseCommand("garbage"));
  269. try std.testing.expectEqual(@as(?ParseResult, null), parseCommand("*"));
  270. try std.testing.expectEqual(@as(?ParseResult, null), parseCommand("*1\r\n"));
  271. }
  272. test "parseCommand bytes_consumed" {
  273. const input = buildRedisArray(&.{ "GET", "key1" });
  274. const result = parseCommand(input) orelse return error.TestUnexpectedResult;
  275. try std.testing.expectEqual(input.len, result.bytes_consumed);
  276. }
  277. test "parseInteger max usize fits" {
  278. const s = "18446744073709551615";
  279. try std.testing.expectEqual(@as(?usize, std.math.maxInt(usize)), parseInteger(s, 0, s.len));
  280. }
  281. test "parseInteger overflow returns null" {
  282. try std.testing.expectEqual(@as(?usize, null), parseInteger("18446744073709551616", 0, 20));
  283. try std.testing.expectEqual(@as(?usize, null), parseInteger("999999999999999999999999999999", 0, 30));
  284. }
  285. test "parseCommand huge bulk length returns null" {
  286. const input = "*2\r\n$3\r\nGET\r\n$18446744073709551615\r\n";
  287. try std.testing.expectEqual(@as(?ParseResult, null), parseCommand(input));
  288. }
  289. test "formatInt zero" {
  290. var buf: [20]u8 = undefined;
  291. const len = formatInt(&buf, 0) orelse return error.TestUnexpectedResult;
  292. try std.testing.expectEqualStrings("0", buf[0..len]);
  293. }
  294. test "formatInt positive" {
  295. var buf: [20]u8 = undefined;
  296. const len = formatInt(&buf, 12345) orelse return error.TestUnexpectedResult;
  297. try std.testing.expectEqualStrings("12345", buf[0..len]);
  298. }
  299. test "formatInt respects buffer bounds" {
  300. var buf: [4]u8 = undefined;
  301. try std.testing.expectEqual(@as(?usize, null), formatInt(buf[0..3], 1234));
  302. const len = formatInt(buf[0..4], 1234) orelse return error.TestUnexpectedResult;
  303. try std.testing.expectEqualStrings("1234", buf[0..len]);
  304. }
  305. test "formatSimpleString" {
  306. var buf: [64]u8 = undefined;
  307. const result = formatSimpleString(&buf, "OK") orelse return error.TestUnexpectedResult;
  308. try std.testing.expectEqualStrings("+OK\r\n", result);
  309. }
  310. test "formatSimpleString respects buffer bounds" {
  311. var exact: [5]u8 = undefined;
  312. const ok = formatSimpleString(exact[0..5], "OK") orelse return error.TestUnexpectedResult;
  313. try std.testing.expectEqualStrings("+OK\r\n", ok);
  314. var short: [4]u8 = undefined;
  315. try std.testing.expectEqual(@as(?[]const u8, null), formatSimpleString(short[0..4], "OK"));
  316. }
  317. test "formatBulkString" {
  318. var buf: [64]u8 = undefined;
  319. const result = formatBulkString(&buf, "hello") orelse return error.TestUnexpectedResult;
  320. try std.testing.expectEqualStrings("$5\r\nhello\r\n", result);
  321. }
  322. test "formatBulkString respects buffer bounds" {
  323. const value = "hello world";
  324. const needed = 1 + intDigits(value.len) + 2 + value.len + 2;
  325. var exact: [64]u8 = undefined;
  326. const ok = formatBulkString(exact[0..needed], value) orelse return error.TestUnexpectedResult;
  327. try std.testing.expectEqualStrings("$11\r\nhello world\r\n", ok);
  328. var short: [64]u8 = undefined;
  329. try std.testing.expectEqual(@as(?[]const u8, null), formatBulkString(short[0 .. needed - 1], value));
  330. }
  331. test "formatNullBulkString" {
  332. var buf: [64]u8 = undefined;
  333. const result = formatNullBulkString(&buf) orelse return error.TestUnexpectedResult;
  334. try std.testing.expectEqualStrings("$-1\r\n", result);
  335. }
  336. test "formatNullBulkString respects buffer bounds" {
  337. var exact: [5]u8 = undefined;
  338. const ok = formatNullBulkString(exact[0..5]) orelse return error.TestUnexpectedResult;
  339. try std.testing.expectEqualStrings("$-1\r\n", ok);
  340. var short: [4]u8 = undefined;
  341. try std.testing.expectEqual(@as(?[]const u8, null), formatNullBulkString(short[0..4]));
  342. }
  343. test "formatError" {
  344. var buf: [64]u8 = undefined;
  345. const result = formatError(&buf, "ERR bad") orelse return error.TestUnexpectedResult;
  346. try std.testing.expectEqualStrings("-ERR bad\r\n", result);
  347. }
  348. test "formatError respects buffer bounds" {
  349. const msg = "ERR bad";
  350. const needed = 1 + msg.len + 2;
  351. var exact: [16]u8 = undefined;
  352. const ok = formatError(exact[0..needed], msg) orelse return error.TestUnexpectedResult;
  353. try std.testing.expectEqualStrings("-ERR bad\r\n", ok);
  354. var short: [16]u8 = undefined;
  355. try std.testing.expectEqual(@as(?[]const u8, null), formatError(short[0 .. needed - 1], msg));
  356. }
  357. test "formatInteger positive" {
  358. var buf: [64]u8 = undefined;
  359. const result = formatInteger(&buf, 42) orelse return error.TestUnexpectedResult;
  360. try std.testing.expectEqualStrings(":42\r\n", result);
  361. }
  362. test "formatInteger zero" {
  363. var buf: [64]u8 = undefined;
  364. const result = formatInteger(&buf, 0) orelse return error.TestUnexpectedResult;
  365. try std.testing.expectEqualStrings(":0\r\n", result);
  366. }
  367. test "formatInteger negative" {
  368. var buf: [64]u8 = undefined;
  369. const result = formatInteger(&buf, -7) orelse return error.TestUnexpectedResult;
  370. try std.testing.expectEqualStrings(":-7\r\n", result);
  371. }
  372. test "formatInteger respects buffer bounds" {
  373. var exact: [16]u8 = undefined;
  374. const ok = formatInteger(exact[0..5], 42) orelse return error.TestUnexpectedResult;
  375. try std.testing.expectEqualStrings(":42\r\n", ok);
  376. var short: [16]u8 = undefined;
  377. try std.testing.expectEqual(@as(?[]const u8, null), formatInteger(short[0..4], 42));
  378. }
  379. test "executeCommand UNKNOWN" {
  380. var buf: [256]u8 = undefined;
  381. const cmd = RedisCommand{ .cmd_type = .UNKNOWN, .key = "", .value = "" };
  382. const result = executeCommand(cmd, &buf) orelse return error.TestUnexpectedResult;
  383. try std.testing.expectEqualStrings("-ERR unknown command\r\n", result);
  384. }
  385. test "executeCommand GET with insufficient buffer returns null" {
  386. storage.init();
  387. _ = storage.restore("overflow_key", "this value is far too long to fit in a tiny buffer");
  388. const cmd = RedisCommand{ .cmd_type = .GET, .key = "overflow_key", .value = "" };
  389. var tiny: [16]u8 = undefined;
  390. try std.testing.expectEqual(@as(?[]const u8, null), executeCommand(cmd, tiny[0..]));
  391. var enough: [512]u8 = undefined;
  392. const resp = executeCommand(cmd, enough[0..]) orelse return error.TestUnexpectedResult;
  393. try std.testing.expect(std.mem.startsWith(u8, resp, "$"));
  394. }