2
0

main.zig 18 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242243244245246247248249250251252253254255256257258259260261262263264265266267268269270271272273274275276277278279280281282283284285286287288289290291292293294295296297298299300301302303304305306307308309310311312313314315316317318319320321322323324325326327328329330331332333334335336337338339340341342343344345346347348349350351352353354355356357358359360361362363364365366367368369370371372373374375376377378379380381
  1. const std = @import("std");
  2. const builtin = @import("builtin");
  3. const posix = std.posix;
  4. const socket = @import("socket.zig");
  5. const Engine = @import("engine.zig").Engine;
  6. const pizzaria = @import("command.zig");
  7. const resp = @import("redis.zig");
  8. const pkbfi = @import("pkbfi.zig");
  9. const migration = @import("migration.zig");
  10. const initial_buffer = 64 * 1024;
  11. const max_connection_bytes = 512 * 1024 * 1024;
  12. const max_connections = 256;
  13. var should_exit = std.atomic.Value(bool).init(false);
  14. var active_connections = std.atomic.Value(u32).init(0);
  15. var memory_mutex: std.Thread.Mutex = .{};
  16. var allocated_connection_bytes: usize = 0;
  17. fn signalHandler(_: c_int) callconv(.c) void {
  18. should_exit.store(true, .release);
  19. }
  20. fn reserveMemory(engine: *Engine, amount: usize) !void {
  21. memory_mutex.lock();
  22. defer memory_mutex.unlock();
  23. if (amount > max_connection_bytes - allocated_connection_bytes) return error.Backpressure;
  24. allocated_connection_bytes += amount;
  25. engine.addConnectionBytes(amount);
  26. }
  27. fn releaseMemory(engine: *Engine, amount: usize) void {
  28. memory_mutex.lock();
  29. std.debug.assert(amount <= allocated_connection_bytes);
  30. allocated_connection_bytes -= amount;
  31. memory_mutex.unlock();
  32. engine.removeConnectionBytes(amount);
  33. }
  34. fn sendAll(connection: posix.socket_t, bytes: []const u8) !void {
  35. var position: usize = 0;
  36. while (position < bytes.len) {
  37. const written = posix.send(connection, bytes[position..], posix.MSG.NOSIGNAL) catch |err| switch (err) {
  38. error.WouldBlock => return error.WriteTimedOut,
  39. else => return err,
  40. };
  41. if (written == 0) return error.ConnectionClosed;
  42. position += written;
  43. }
  44. }
  45. fn handleResp(engine: *Engine, connection: posix.socket_t, command: resp.Command) !void {
  46. if (command.command_type == .get) return handleRespGet(engine, connection, command.key);
  47. var response = try resp.execute(engine, std.heap.smp_allocator, command);
  48. defer response.deinit(std.heap.smp_allocator);
  49. if (response == .bulk) {
  50. try reserveMemory(engine, response.bulk.bytes.len);
  51. defer releaseMemory(engine, response.bulk.bytes.len);
  52. }
  53. var prefix: [64]u8 = undefined;
  54. const encoded = try resp.encodePrefix(response, &prefix);
  55. try sendAll(connection, encoded);
  56. if (response == .bulk) {
  57. try sendAll(connection, response.bulk.bytes);
  58. try sendAll(connection, "\r\n");
  59. }
  60. }
  61. fn handleRespGet(engine: *Engine, connection: posix.socket_t, key: []const u8) !void {
  62. const record = try engine.getRef(key) orelse {
  63. try sendAll(connection, "$-1\r\n");
  64. return;
  65. };
  66. var header: [32]u8 = undefined;
  67. try sendAll(connection, try std.fmt.bufPrint(&header, "${d}\r\n", .{record.value_len}));
  68. const buffer_size = @min(@as(usize, 64 * 1024), record.value_len);
  69. try reserveMemory(engine, buffer_size);
  70. defer releaseMemory(engine, buffer_size);
  71. const buffer = try std.heap.smp_allocator.alloc(u8, buffer_size);
  72. defer std.heap.smp_allocator.free(buffer);
  73. var position: u32 = 0;
  74. while (position < record.value_len) {
  75. const amount = try engine.readValue(record, buffer, position);
  76. try sendAll(connection, buffer[0..amount]);
  77. position += @intCast(amount);
  78. }
  79. try sendAll(connection, "\r\n");
  80. }
  81. fn handleRespGets(engine: *Engine, connection: posix.socket_t, keys: []const []const u8) !void {
  82. const capacity = 1024 * 1024;
  83. try reserveMemory(engine, capacity);
  84. defer releaseMemory(engine, capacity);
  85. const output = try std.heap.smp_allocator.alloc(u8, capacity);
  86. defer std.heap.smp_allocator.free(output);
  87. var position: usize = 0;
  88. for (keys) |key| {
  89. const record = try engine.getRef(key) orelse {
  90. if (position + 5 > output.len) {
  91. try sendAll(connection, output[0..position]);
  92. position = 0;
  93. }
  94. @memcpy(output[position .. position + 5], "$-1\r\n");
  95. position += 5;
  96. continue;
  97. };
  98. var header: [32]u8 = undefined;
  99. const encoded_header = try std.fmt.bufPrint(&header, "${d}\r\n", .{record.value_len});
  100. const needed = encoded_header.len + record.value_len + 2;
  101. if (needed > output.len) {
  102. if (position != 0) {
  103. try sendAll(connection, output[0..position]);
  104. position = 0;
  105. }
  106. try handleRespGet(engine, connection, key);
  107. continue;
  108. }
  109. if (position + needed > output.len) {
  110. try sendAll(connection, output[0..position]);
  111. position = 0;
  112. }
  113. @memcpy(output[position .. position + encoded_header.len], encoded_header);
  114. position += encoded_header.len;
  115. const amount = try engine.readValue(record, output[position .. position + record.value_len], 0);
  116. position += amount;
  117. @memcpy(output[position .. position + 2], "\r\n");
  118. position += 2;
  119. }
  120. if (position != 0) try sendAll(connection, output[0..position]);
  121. }
  122. fn handleConnection(engine: *Engine, connection: posix.socket_t) void {
  123. defer _ = active_connections.fetchSub(1, .monotonic);
  124. defer posix.close(connection);
  125. reserveMemory(engine, initial_buffer) catch return;
  126. var buffer = std.heap.smp_allocator.alloc(u8, initial_buffer) catch {
  127. releaseMemory(engine, initial_buffer);
  128. return;
  129. };
  130. defer {
  131. std.heap.smp_allocator.free(buffer);
  132. releaseMemory(engine, buffer.len);
  133. }
  134. var session = pkbfi.Session.init(std.heap.smp_allocator);
  135. defer session.deinit();
  136. var buffered: usize = 0;
  137. while (!should_exit.load(.acquire)) {
  138. if (buffered == buffer.len) {
  139. const maximum: usize = if (buffered >= 4 and std.mem.eql(u8, buffer[0..4], "PKBF")) pkbfi.max_frame_size + pkbfi.header_size else if (buffered > 0 and buffer[0] == '*') pkbfi.max_frame_size + 1024 else 1024 * 1024;
  140. if (buffer.len >= maximum) return;
  141. const next = @min(maximum, buffer.len * 2);
  142. reserveMemory(engine, next - buffer.len) catch return;
  143. buffer = std.heap.smp_allocator.realloc(buffer, next) catch {
  144. releaseMemory(engine, next - buffer.len);
  145. return;
  146. };
  147. }
  148. const amount = posix.read(connection, buffer[buffered..]) catch |err| switch (err) {
  149. error.WouldBlock => return,
  150. error.ConnectionResetByPeer => return,
  151. else => return,
  152. };
  153. if (amount == 0) return;
  154. buffered += amount;
  155. var consumed: usize = 0;
  156. while (consumed < buffered) {
  157. const input = buffer[consumed..buffered];
  158. if (input.len >= 4 and std.mem.eql(u8, input[0..4], "PKBF")) {
  159. const frame = pkbfi.parse(input) catch |err| switch (err) {
  160. error.Incomplete => break,
  161. else => return,
  162. };
  163. if (frame.opcode == .put) {
  164. var operations: [256]@import("engine.zig").Operation = undefined;
  165. var request_ids: [256]u64 = undefined;
  166. var lsns: [256]u64 = undefined;
  167. var count: usize = 0;
  168. var pipeline_consumed: usize = 0;
  169. while (count < operations.len and pipeline_consumed < input.len) {
  170. const next = pkbfi.parse(input[pipeline_consumed..]) catch |err| switch (err) {
  171. error.Incomplete => break,
  172. else => return,
  173. };
  174. if (next.opcode != .put) break;
  175. operations[count] = pkbfi.putOperation(next) catch return;
  176. request_ids[count] = next.request_id;
  177. count += 1;
  178. pipeline_consumed += next.consumed;
  179. }
  180. engine.beginRequest();
  181. engine.putMany(operations[0..count], lsns[0..count]) catch {
  182. engine.endRequest();
  183. return;
  184. };
  185. engine.endRequest();
  186. var responses = std.ArrayListUnmanaged(u8){};
  187. defer responses.deinit(std.heap.smp_allocator);
  188. for (0..count) |index| {
  189. var body: [10]u8 = [_]u8{0} ** 10;
  190. std.mem.writeInt(u64, body[2..10], lsns[index], .little);
  191. const response = pkbfi.encode(std.heap.smp_allocator, @intFromEnum(pkbfi.Opcode.put) | 0x8000, 1, request_ids[index], &body) catch return;
  192. defer std.heap.smp_allocator.free(response);
  193. responses.appendSlice(std.heap.smp_allocator, response) catch return;
  194. }
  195. reserveMemory(engine, responses.items.len) catch return;
  196. sendAll(connection, responses.items) catch {
  197. releaseMemory(engine, responses.items.len);
  198. return;
  199. };
  200. releaseMemory(engine, responses.items.len);
  201. consumed += pipeline_consumed;
  202. } else {
  203. engine.beginRequest();
  204. const response = session.execute(engine, frame) catch {
  205. engine.endRequest();
  206. return;
  207. };
  208. engine.endRequest();
  209. reserveMemory(engine, response.len) catch {
  210. std.heap.smp_allocator.free(response);
  211. return;
  212. };
  213. sendAll(connection, response) catch {
  214. std.heap.smp_allocator.free(response);
  215. releaseMemory(engine, response.len);
  216. return;
  217. };
  218. std.heap.smp_allocator.free(response);
  219. releaseMemory(engine, response.len);
  220. consumed += frame.consumed;
  221. }
  222. } else if (input[0] == '*') {
  223. const parsed = resp.parse(input) catch |err| switch (err) {
  224. error.Incomplete => break,
  225. else => return,
  226. };
  227. if (parsed.command.command_type == .set) {
  228. var operations: [256]@import("engine.zig").Operation = undefined;
  229. var lsns: [256]u64 = undefined;
  230. var count: usize = 0;
  231. var pipeline_consumed: usize = 0;
  232. while (count < operations.len and pipeline_consumed < input.len) {
  233. const next = resp.parse(input[pipeline_consumed..]) catch |err| switch (err) {
  234. error.Incomplete => break,
  235. else => return,
  236. };
  237. if (next.command.command_type != .set) break;
  238. operations[count] = .{ .opcode = .put, .key = next.command.key, .value = next.command.value };
  239. count += 1;
  240. pipeline_consumed += next.consumed;
  241. }
  242. engine.beginRequest();
  243. engine.putMany(operations[0..count], lsns[0..count]) catch {
  244. engine.endRequest();
  245. return;
  246. };
  247. engine.endRequest();
  248. var responses: [256 * 5]u8 = undefined;
  249. for (0..count) |index| @memcpy(responses[index * 5 ..][0..5], "+OK\r\n");
  250. sendAll(connection, responses[0 .. count * 5]) catch return;
  251. consumed += pipeline_consumed;
  252. } else if (parsed.command.command_type == .get) {
  253. var keys: [256][]const u8 = undefined;
  254. var count: usize = 0;
  255. var pipeline_consumed: usize = 0;
  256. while (count < keys.len and pipeline_consumed < input.len) {
  257. const next = resp.parse(input[pipeline_consumed..]) catch |err| switch (err) {
  258. error.Incomplete => break,
  259. else => return,
  260. };
  261. if (next.command.command_type != .get) break;
  262. keys[count] = next.command.key;
  263. count += 1;
  264. pipeline_consumed += next.consumed;
  265. }
  266. engine.beginRequest();
  267. handleRespGets(engine, connection, keys[0..count]) catch {
  268. engine.endRequest();
  269. return;
  270. };
  271. engine.endRequest();
  272. consumed += pipeline_consumed;
  273. } else {
  274. engine.beginRequest();
  275. handleResp(engine, connection, parsed.command) catch {
  276. engine.endRequest();
  277. return;
  278. };
  279. engine.endRequest();
  280. consumed += parsed.consumed;
  281. }
  282. } else {
  283. const end = std.mem.indexOfScalar(u8, input, '\r') orelse break;
  284. engine.beginRequest();
  285. const response = pizzaria.execute(engine, std.heap.smp_allocator, input[0..end]) catch {
  286. engine.endRequest();
  287. return;
  288. };
  289. engine.endRequest();
  290. reserveMemory(engine, response.len) catch {
  291. std.heap.smp_allocator.free(response);
  292. return;
  293. };
  294. sendAll(connection, response) catch {
  295. std.heap.smp_allocator.free(response);
  296. releaseMemory(engine, response.len);
  297. return;
  298. };
  299. sendAll(connection, "\r") catch {
  300. std.heap.smp_allocator.free(response);
  301. releaseMemory(engine, response.len);
  302. return;
  303. };
  304. std.heap.smp_allocator.free(response);
  305. releaseMemory(engine, response.len);
  306. consumed += end + 1 + @intFromBool(input.len > end + 1 and input[end + 1] == '\n');
  307. }
  308. }
  309. if (consumed != 0) {
  310. const remaining = buffered - consumed;
  311. std.mem.copyForwards(u8, buffer[0..remaining], buffer[consumed..buffered]);
  312. buffered = remaining;
  313. }
  314. if (buffer.len > initial_buffer and buffered <= initial_buffer) {
  315. var replacement = std.heap.smp_allocator.alloc(u8, initial_buffer) catch return;
  316. @memcpy(replacement[0..buffered], buffer[0..buffered]);
  317. const released = buffer.len - initial_buffer;
  318. std.heap.smp_allocator.free(buffer);
  319. buffer = replacement;
  320. releaseMemory(engine, released);
  321. }
  322. }
  323. }
  324. pub fn main() !void {
  325. var host: []const u8 = "127.0.0.1";
  326. var port: u16 = 8085;
  327. var path: []const u8 = ".pkvdb";
  328. var unix_path: ?[]const u8 = null;
  329. var migration_source: ?[]const u8 = null;
  330. var args = try std.process.argsWithAllocator(std.heap.page_allocator);
  331. defer args.deinit();
  332. _ = args.skip();
  333. while (args.next()) |argument| {
  334. if (std.mem.startsWith(u8, argument, "-host=")) host = argument[6..] else if (std.mem.startsWith(u8, argument, "-port=")) port = try std.fmt.parseInt(u16, argument[6..], 10) else if (std.mem.startsWith(u8, argument, "-path=")) path = argument[6..] else if (std.mem.startsWith(u8, argument, "-migrate=")) migration_source = argument[9..] else if (std.mem.eql(u8, argument, "-unix")) unix_path = ".pizzakv.sock" else if (std.mem.startsWith(u8, argument, "-unix=")) unix_path = argument[6..] else if (std.mem.eql(u8, argument, "-redis") or std.mem.eql(u8, argument, "-pkbfi") or std.mem.eql(u8, argument, "-iwal")) {} else return error.InvalidArgument;
  335. }
  336. if (migration_source) |source| {
  337. const result = try migration.migrate(std.heap.smp_allocator, source, path);
  338. std.debug.print("Migrated {d} keys from {d} records checksum={x}\n", .{ result.keys, result.records, result.checksum });
  339. return;
  340. }
  341. const action = posix.Sigaction{ .handler = .{ .handler = signalHandler }, .mask = std.mem.zeroes(posix.sigset_t), .flags = 0 };
  342. _ = posix.sigaction(posix.SIG.TERM, &action, null);
  343. _ = posix.sigaction(posix.SIG.INT, &action, null);
  344. var engine = try Engine.open(std.heap.smp_allocator, path);
  345. defer engine.close();
  346. const listener = if (unix_path) |name| try socket.initUnix(name) else try socket.init(host, port);
  347. defer posix.close(listener);
  348. defer if (unix_path) |name| posix.unlink(name) catch {};
  349. std.debug.print("PizzaKV {s} Pizzaria/RESP/PKBFI\n", .{path});
  350. while (!should_exit.load(.acquire)) {
  351. var descriptors = [_]posix.pollfd{.{ .fd = listener, .events = posix.POLL.IN, .revents = 0 }};
  352. if ((posix.poll(&descriptors, 100) catch continue) == 0) continue;
  353. const connection = posix.accept(listener, null, null, 0) catch continue;
  354. if (active_connections.fetchAdd(1, .monotonic) >= max_connections) {
  355. _ = active_connections.fetchSub(1, .monotonic);
  356. posix.close(connection);
  357. continue;
  358. }
  359. if (unix_path == null and (builtin.target.os.tag == .linux or builtin.target.os.tag == .macos)) posix.setsockopt(connection, posix.IPPROTO.TCP, posix.TCP.NODELAY, &std.mem.toBytes(@as(c_int, 1))) catch {};
  360. socket.setReadTimeout(connection, 30) catch {};
  361. socket.setWriteTimeout(connection, 30) catch {};
  362. const thread = std.Thread.spawn(.{}, handleConnection, .{ &engine, connection }) catch {
  363. _ = active_connections.fetchSub(1, .monotonic);
  364. posix.close(connection);
  365. continue;
  366. };
  367. thread.detach();
  368. }
  369. while (active_connections.load(.monotonic) != 0) std.Thread.sleep(10 * std.time.ns_per_ms);
  370. }