main.zig 18 KB

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