//! MongoDB wire protocol — OP_MSG framing, request parsing, reply building. //! All integers are little-endian. A message is: //! int32 messageLength | int32 requestID | int32 responseTo | int32 opCode //! [OP_MSG only] uint32 flagBits | sectionKind(0x00 body | 0x01 sequence) ... const std = @import("std"); const bson = @import("bson.zig"); pub const op_code_msg: i32 = 2013; pub const op_code_query: i32 = 2004; pub const op_code_reply: i32 = 2001; pub const max_message_size: usize = 48 * 1024 * 1024; pub const max_bson_object_size: i32 = 16 * 1024 * 1024; pub const flag_checksum_present: u32 = 1 << 1; pub const flag_more_to_come: u32 = 1 << 0; pub const Seq = struct { name: []const u8, docs: []bson.Document, }; /// A parsed request message. Owns the body document and all sequence /// documents via `arena` plus each document's own arena. pub const Message = struct { arena: std.heap.ArenaAllocator, request_id: u32, flags: u32, op_code: i32, body: bson.Document, seqs: []const Seq, pub fn parse(gpa: std.mem.Allocator, bytes: []const u8) !Message { if (bytes.len < 16) return error.InvalidMessage; const total: u32 = std.mem.readInt(u32, bytes[0..4], .little); if (total != bytes.len or total < 16) return error.InvalidMessage; const op_code: i32 = std.mem.readInt(i32, bytes[12..16], .little); if (op_code == op_code_msg) { return parse_msg(gpa, bytes); } if (op_code == op_code_query) { return parse_query(gpa, bytes); } return error.UnsupportedOpCode; } fn parse_msg(gpa: std.mem.Allocator, bytes: []const u8) !Message { const flags: u32 = std.mem.readInt(u32, bytes[16..20], .little); var arena = std.heap.ArenaAllocator.init(gpa); errdefer arena.deinit(); var idx: usize = 20; var body: ?bson.Document = null; var seqs: std.ArrayListUnmanaged(Seq) = .empty; errdefer seqs.deinit(arena.allocator()); // Every successfully parsed document owns its own arena allocated // from gpa. If parsing fails partway, they must all be freed — // the message arena cleanup alone leaks them. var parsed_docs: std.ArrayListUnmanaged(bson.Document) = .empty; defer parsed_docs.deinit(gpa); errdefer { for (parsed_docs.items) |*d| d.deinit(); } while (idx < bytes.len) { const kind = bytes[idx]; idx += 1; switch (kind) { 0x00 => { if (body != null) return error.InvalidMessage; // only one body if (idx + 4 > bytes.len) return error.InvalidMessage; const doc_len: u32 = std.mem.readInt(u32, bytes[idx..][0..4], .little); if (doc_len < 5 or doc_len > bytes.len - idx) return error.InvalidMessage; const parsed = try bson.Document.parse(gpa, bytes[idx..]); try parsed_docs.append(gpa, parsed); if (parsed.pairs.len == 0) return error.InvalidMessage; body = parsed; idx += doc_len; }, 0x01 => { if (idx + 4 > bytes.len) return error.InvalidMessage; const size: u32 = std.mem.readInt(u32, bytes[idx..][0..4], .little); if (size < 5 or size > bytes.len - idx) return error.InvalidMessage; const seq_end = idx + size; var name_end = idx + 4; while (name_end < seq_end and bytes[name_end] != 0) name_end += 1; if (name_end >= seq_end) return error.InvalidMessage; const name = bytes[idx + 4 .. name_end]; var doc_idx = name_end + 1; var docs: std.ArrayListUnmanaged(bson.Document) = .empty; errdefer docs.deinit(arena.allocator()); while (doc_idx < seq_end) { if (doc_idx + 4 > seq_end) return error.InvalidMessage; const doc_len: u32 = std.mem.readInt(u32, bytes[doc_idx..][0..4], .little); if (doc_len < 5 or doc_len > seq_end - doc_idx) return error.InvalidMessage; const parsed = try bson.Document.parse(gpa, bytes[doc_idx..]); try parsed_docs.append(gpa, parsed); try docs.append(arena.allocator(), parsed); doc_idx += doc_len; } try seqs.append(arena.allocator(), .{ .name = try arena.allocator().dupe(u8, name), .docs = try docs.toOwnedSlice(arena.allocator()), }); idx = seq_end; }, else => return error.InvalidMessage, } } // Trailing CRC32C checksum, if present, is skipped by not consuming // it — we treat it as part of the message length and ignore it. return .{ .arena = arena, .request_id = std.mem.readInt(u32, bytes[4..8], .little), .flags = flags, .op_code = op_code_msg, .body = body orelse return error.InvalidMessage, .seqs = seqs.items, }; } /// OP_QUERY — legacy framing, still used by drivers for the initial /// handshake commands. Layout: flags | fullCollectionName(cstring) | /// numberToSkip | numberToReturn | query BSON doc. fn parse_query(gpa: std.mem.Allocator, bytes: []const u8) !Message { var arena = std.heap.ArenaAllocator.init(gpa); errdefer arena.deinit(); var idx: usize = 20; if (idx + 1 > bytes.len) return error.InvalidMessage; while (idx < bytes.len and bytes[idx] != 0) idx += 1; // skip flags, ns is a cstring after flags if (idx >= bytes.len) return error.InvalidMessage; idx += 1; // numberToSkip + numberToReturn if (idx + 8 > bytes.len) return error.InvalidMessage; idx += 8; const doc_len: u32 = std.mem.readInt(u32, bytes[idx..][0..4], .little); if (doc_len < 5 or doc_len > bytes.len - idx) return error.InvalidMessage; const body = try bson.Document.parse(gpa, bytes[idx..]); if (body.pairs.len == 0) return error.InvalidMessage; return .{ .arena = arena, .request_id = std.mem.readInt(u32, bytes[4..8], .little), .flags = std.mem.readInt(u32, bytes[16..20], .little), .op_code = op_code_query, .body = body, .seqs = &.{}, }; } pub fn deinit(self: *Message) void { for (self.seqs) |seq| { for (seq.docs) |*doc| doc.deinit(); } self.body.deinit(); self.arena.deinit(); } /// The command name is the first field of the command document. pub fn command_name(self: *const Message) []const u8 { return self.body.pairs[0].key; } /// Database from the `$db` field (present in OP_MSG commands). pub fn db_name(self: *const Message) ?[]const u8 { const v = bson.get_pair(self.body.pairs, "$db") orelse return null; return switch (v) { .string => |s| s, else => null, }; } }; /// Builder for a command reply document. Strings for keys and values must /// live long enough — allocate them from `arena()`. pub const Reply = struct { arena: std.heap.ArenaAllocator, pairs: std.ArrayListUnmanaged(bson.Pair), pub fn init(gpa: std.mem.Allocator) Reply { return .{ .arena = std.heap.ArenaAllocator.init(gpa), .pairs = .empty, }; } pub fn deinit(self: *Reply) void { self.pairs.deinit(self.arena.allocator()); self.arena.deinit(); } pub fn arena_alloc(self: *Reply) std.mem.Allocator { return self.arena.allocator(); } pub fn put(self: *Reply, key: []const u8, value: bson.Value) !void { const key_owned = try self.arena.allocator().dupe(u8, key); try self.pairs.append(self.arena.allocator(), .{ .key = key_owned, .value = value }); } pub fn put_ok(self: *Reply) !void { try self.put("ok", .{ .double = 1.0 }); } pub fn put_error(self: *Reply, code: i32, code_name: []const u8, message: []const u8) !void { try self.put("ok", .{ .double = 0.0 }); try self.put("errmsg", .{ .string = try self.arena.allocator().dupe(u8, message) }); try self.put("code", .{ .int32 = code }); try self.put("codeName", .{ .string = try self.arena.allocator().dupe(u8, code_name) }); } /// Serialize this reply as a full OP_MSG message. pub fn build(self: *Reply, gpa: std.mem.Allocator, request_id: u32, response_to: u32, out: *std.ArrayListUnmanaged(u8)) !void { try write_message(gpa, request_id, response_to, 0, self.pairs.items, out); } }; /// Serialize an OP_MSG reply: header + flags + single body section. pub fn write_message( gpa: std.mem.Allocator, request_id: u32, response_to: u32, flags: u32, body: []const bson.Pair, out: *std.ArrayListUnmanaged(u8), ) !void { const len_pos = out.items.len; var header: [20]u8 = undefined; std.mem.writeInt(u32, header[4..8], request_id, .little); std.mem.writeInt(u32, header[8..12], response_to, .little); std.mem.writeInt(i32, header[12..16], op_code_msg, .little); std.mem.writeInt(u32, header[16..20], flags, .little); try out.appendSlice(gpa, &header); try out.append(gpa, 0x00); // single body section try bson.write_doc(body, gpa, out); const total = out.items.len - len_pos; if (total > std.math.maxInt(u32)) return error.MessageTooLarge; std.mem.writeInt(u32, out.items[len_pos..][0..4], @intCast(total), .little); } /// Serialize an OP_REPLY (legacy): header + responseFlags | cursorID | /// startingFrom | numberReturned + one document. Used only to answer /// OP_QUERY handshake requests. pub fn write_reply_query( gpa: std.mem.Allocator, request_id: u32, response_to: u32, body: []const bson.Pair, out: *std.ArrayListUnmanaged(u8), ) !void { const len_pos = out.items.len; var header: [16]u8 = undefined; std.mem.writeInt(u32, header[4..8], request_id, .little); std.mem.writeInt(u32, header[8..12], response_to, .little); std.mem.writeInt(i32, header[12..16], op_code_reply, .little); try out.appendSlice(gpa, &header); var reply_fields: [20]u8 = [_]u8{0} ** 20; // responseFlags | cursorID | startingFrom | numberReturned std.mem.writeInt(i32, reply_fields[16..20], 1, .little); // numberReturned = 1 try out.appendSlice(gpa, &reply_fields); try bson.write_doc(body, gpa, out); const total = out.items.len - len_pos; if (total > std.math.maxInt(u32)) return error.MessageTooLarge; std.mem.writeInt(u32, out.items[len_pos..][0..4], @intCast(total), .little); } // --------------------------------------------------------------------------- // Tests // --------------------------------------------------------------------------- const testing = std.testing; test "reply serializes to a parseable message" { var reply = Reply.init(testing.allocator); defer reply.deinit(); try reply.put_ok(); try reply.put("version", .{ .string = "4.4.0" }); var out: std.ArrayListUnmanaged(u8) = .empty; defer out.deinit(testing.allocator); try reply.build(testing.allocator, 1, 42, &out); // Parse the reply as if it were a request — body section only. var msg = try Message.parse(testing.allocator, out.items); defer msg.deinit(); try testing.expectEqual(@as(u32, 1), msg.request_id); try testing.expectEqualStrings("ok", msg.command_name()); try testing.expectEqual(@as(f64, 1.0), msg.body.get("ok").?.double); try testing.expectEqualStrings("4.4.0", msg.body.get("version").?.string); try testing.expectEqual(@as(usize, 0), msg.seqs.len); } test "reject non-OP_MSG non-OP_QUERY opcodes" { var buf: [20]u8 = undefined; std.mem.writeInt(u32, buf[0..4], 20, .little); std.mem.writeInt(i32, buf[12..16], 2005, .little); // OP_GET_MORE try testing.expectError(error.UnsupportedOpCode, Message.parse(testing.allocator, &buf)); } test "broken message with a parsed body section does not leak" { // Build a valid body section followed by a garbage section; the body // document must be freed on the error path. var doc_buf: std.ArrayListUnmanaged(u8) = .empty; defer doc_buf.deinit(testing.allocator); try bson.write_doc(&.{.{ .key = "ping", .value = .{ .int32 = 1 } }}, testing.allocator, &doc_buf); var msg: std.ArrayListUnmanaged(u8) = .empty; defer msg.deinit(testing.allocator); try msg.appendSlice(testing.allocator, &[_]u8{ 0, 0, 0, 0, 7, 0, 0, 0, 0, 0, 0, 0, 0xDD, 0x07, 0, 0, 0, 0, 0, 0 }); try msg.append(testing.allocator, 0x00); try msg.appendSlice(testing.allocator, doc_buf.items); try msg.append(testing.allocator, 0x7F); // invalid section kind std.mem.writeInt(u32, msg.items[0..4], @intCast(msg.items.len), .little); try testing.expectError(error.InvalidMessage, Message.parse(testing.allocator, msg.items)); // testing.allocator flags any leaked body document arena here. } test "parse OP_QUERY handshake" { var out: std.ArrayListUnmanaged(u8) = .empty; defer out.deinit(testing.allocator); // header (opcode 2004) + flags + ns cstring + skip/return + query doc var msg_buf: std.ArrayListUnmanaged(u8) = .empty; defer msg_buf.deinit(testing.allocator); try msg_buf.appendSlice(testing.allocator, &[_]u8{ 0, 0, 0, 0, 7, 0, 0, 0, 0, 0, 0, 0, 0xD4, 0x07, 0, 0, 0, 0, 0, 0 }); try msg_buf.appendSlice(testing.allocator, "admin.$cmd"); try msg_buf.append(testing.allocator, 0); try msg_buf.appendSlice(testing.allocator, &[_]u8{ 0, 0, 0, 0, 1, 0, 0, 0 }); // skip=0, return=1 try bson.write_doc(&.{.{ .key = "isMaster", .value = .{ .int32 = 1 } }}, testing.allocator, &msg_buf); std.mem.writeInt(u32, msg_buf.items[0..4], @intCast(msg_buf.items.len), .little); var msg = try Message.parse(testing.allocator, msg_buf.items); defer msg.deinit(); try testing.expectEqual(@as(i32, op_code_query), msg.op_code); try testing.expectEqualStrings("isMaster", msg.command_name()); }