commit 4de42091a40b46414f8dfbacf741f09b9f8e745a Author: mongo-light Date: Sun Aug 2 10:29:01 2026 +0300 baseline: mongo-light working tree before concurrency refactor diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..ab56755 --- /dev/null +++ b/.gitignore @@ -0,0 +1,3 @@ +.zig-cache/ +zig-out/ +*.log diff --git a/README.md b/README.md new file mode 100644 index 0000000..118954b --- /dev/null +++ b/README.md @@ -0,0 +1,80 @@ +# mongo-light + +A lightweight, embedded MongoDB-compatible document database written in +Zig 0.16. Like SQLite, it stores everything in a single file; unlike SQLite, +it speaks the MongoDB wire protocol, so real clients — `mongosh`, the Node.js +driver, PyMongo — connect over TCP and just work. + +## Quick start + +```sh +zig build # build the server +zig build test # run the unit test suite + +zig-out/bin/mongo-light --port 27017 --db data.log + +# in another terminal: +mongosh --port 27017 +> db.users.insertOne({name: "alice", age: 30}) +> db.users.find({age: {$gt: 25}}).toArray() +> db.users.updateOne({name: "alice"}, {$set: {vip: true}}) +> db.users.deleteOne({name: "bob"}) +``` + +## Features + +- **Wire protocol**: OP_MSG (2013) plus legacy OP_QUERY/OP_REPLY (2004/2001) + for the driver handshake; hello/isMaster with `maxWireVersion: 8`, so + modern drivers (Node, Python, mongosh) connect without workarounds. +- **BSON**: full parse/serialize round-trip for all common types + (including binary, regex, timestamps, ObjectId), canonical MongoDB + comparison order for sorting and range queries. +- **CRUD**: `insert`, `find` (filter, sort, skip/limit, projection), + `update` (multi/upsert), `delete`, `findAndModify`, `count`, + `aggregate` (`$match`, `$sort`, `$skip`, `$limit`, `$project`, `$count`, + `$group` with `$sum`), plus `create`/`drop`/`listCollections`/ + `listDatabases`/`dropDatabase`. +- **Query operators**: `$eq` `$ne` `$gt` `$gte` `$lt` `$lte` `$in` `$nin` + `$exists` `$regex` (hand-rolled engine: anchors, `.`, `* + ?`, character + classes, groups, alternation, `i`/`s` options) `$not` `$and` `$or` `$nor` + `$size` `$all` `$elemMatch`, with dot paths and array multikey semantics. +- **Update operators**: `$set` `$unset` `$inc` `$push` (`$each`) `$pull` + `$rename`, with dot-path creation (including array indices). +- **Storage**: append-only record log (CRC32-checked, `fsync` per write, + torn-tail tolerant) with in-memory indexes rebuilt on open and automatic + compaction (rewrite + atomic rename when the log grows past 16 MB). + Killed mid-write (`kill -9`), the database recovers all committed writes; + the log and compaction both work with relative or absolute `--db` paths. + Records up to the announced 16 MB `maxBsonObjectSize` replay correctly. +- **Concurrency**: one mutex serializes commands end-to-end (command-level + atomicity); the Io worker pool only overlaps connection I/O with command + execution — there is no read parallelism. Fine for light workloads. + +## Layout + +``` +src/ + bson.zig BSON parse/serialize, ObjectId, canonical comparison order + wire.zig OP_MSG/OP_QUERY framing, message + reply builders + commands.zig command dispatch (hello, CRUD, aggregate, admin) + server.zig TCP accept loop, per-connection handlers + db.zig in-memory engine: db → collection → _id → document maps + storage.zig append-only log: records, replay, CRC validation + query.zig filter matcher, regex engine, sort, projection + update.zig update operators with dot-path navigation + main.zig CLI: --port, --bind, --db +``` + +## Not (yet) implemented + +- Authentication (SCRAM) — run without credentials +- Real cursors (all results are returned in one batch, cursor id 0) +- Indexes (O(n) scans) +- Transactions, change streams, replicasets +- Compression (OP_COMPRESSED) + +## Code style + +Zig 0.16 idioms (`std.Io` threaded through everything, unmanaged +containers); user-declared functions use `snake_case` per this repo's house +style. diff --git a/build.zig b/build.zig new file mode 100644 index 0000000..f1f43ed --- /dev/null +++ b/build.zig @@ -0,0 +1,46 @@ +const std = @import("std"); + +pub fn build(b: *std.Build) void { + const target = b.standardTargetOptions(.{}); + const optimize = b.standardOptimizeOption(.{}); + + const lib_mod = b.createModule(.{ + .root_source_file = b.path("src/lib.zig"), + .target = target, + .optimize = optimize, + }); + + const exe_mod = b.createModule(.{ + .root_source_file = b.path("src/main.zig"), + .target = target, + .optimize = optimize, + .imports = &.{ + .{ .name = "mongo", .module = lib_mod }, + }, + }); + + const exe = b.addExecutable(.{ + .name = "mongo-light", + .root_module = exe_mod, + }); + b.installArtifact(exe); + + const run_cmd = b.addRunArtifact(exe); + run_cmd.step.dependOn(b.getInstallStep()); + if (b.args) |args| run_cmd.addArgs(args); + const run_step = b.step("run", "Run mongo-light server"); + run_step.dependOn(&run_cmd.step); + + const test_mod = b.createModule(.{ + .root_source_file = b.path("src/lib.zig"), + .target = target, + .optimize = optimize, + }); + + const test_step = b.addTest(.{ + .root_module = test_mod, + }); + const run_tests = b.addRunArtifact(test_step); + const test_help = b.step("test", "Run unit tests"); + test_help.dependOn(&run_tests.step); +} diff --git a/build.zig.zon b/build.zig.zon new file mode 100644 index 0000000..cf84ce4 --- /dev/null +++ b/build.zig.zon @@ -0,0 +1,7 @@ +.{ + .name = .mongo_light, + .version = "0.0.1", + .minimum_zig_version = "0.16.0", + .paths = .{""}, + .fingerprint = 0xb3baeb7ec5369577, +} diff --git a/src/bson.zig b/src/bson.zig new file mode 100644 index 0000000..35b57ab --- /dev/null +++ b/src/bson.zig @@ -0,0 +1,767 @@ +//! BSON — Binary JSON. The foundation of the MongoDB wire protocol and the +//! storage engine. Documents parse into an arena-backed value tree; the tree +//! serializes back to canonical BSON bytes. + +const std = @import("std"); + +pub const ObjectId = [12]u8; + +pub const Binary = struct { + subtype: u8, + data: []const u8, +}; + +pub const Regex = struct { + pattern: []const u8, + options: []const u8, +}; + +/// Values we round-trip but never interpret: db_pointer (0x0C), +/// code_with_scope (0x0F), undefined (0x06). `data` is the raw payload that +/// follows the element type byte, re-emitted verbatim on serialize. +pub const Opaque = struct { + kind: u8, + data: []const u8, +}; + +pub const Pair = struct { + key: []const u8, + value: Value, +}; + +pub const Value = union(enum) { + double: f64, + string: []const u8, + doc: []const Pair, + array: []const Value, + binary: Binary, + object_id: ObjectId, + bool: bool, + datetime: i64, + null, + regex: Regex, + code: []const u8, + symbol: []const u8, + int32: i32, + timestamp: u64, + int64: i64, + decimal128: [16]u8, + min_key, + max_key, + opaque_val: Opaque, + + pub fn type_tag(self: Value) u8 { + return switch (self) { + .double => 0x01, + .string => 0x02, + .doc => 0x03, + .array => 0x04, + .binary => 0x05, + .object_id => 0x07, + .bool => 0x08, + .datetime => 0x09, + .null => 0x0A, + .regex => 0x0B, + .code => 0x0D, + .symbol => 0x0E, + .int32 => 0x10, + .timestamp => 0x11, + .int64 => 0x12, + .decimal128 => 0x13, + .min_key => 0xFF, // not serializable; rank only + .max_key => 0x7F, // not serializable; rank only + .opaque_val => |o| o.kind, + }; + } + + pub fn is_number(self: Value) bool { + return switch (self) { + .double, .int32, .int64 => true, + else => false, + }; + } + + /// Numeric value widened to f128 — exact for i64 and f64. + pub fn as_f128(self: Value) f128 { + return switch (self) { + .double => |d| @as(f128, @floatCast(d)), + .int32 => |i| @as(f128, @floatFromInt(i)), + .int64 => |i| @as(f128, @floatFromInt(i)), + else => unreachable, + }; + } +}; + +/// A parsed document owns everything it references via its arena. Not +/// copyable — pass by pointer. +pub const Document = struct { + arena: std.heap.ArenaAllocator, + pairs: []const Pair, + + pub fn parse(allocator: std.mem.Allocator, bytes: []const u8) !Document { + var arena = std.heap.ArenaAllocator.init(allocator); + errdefer arena.deinit(); + var idx: usize = 0; + const pairs = try parse_doc_into(&arena, bytes, &idx); + return .{ .arena = arena, .pairs = pairs }; + } + + pub fn deinit(self: *Document) void { + self.arena.deinit(); + } + + pub fn get(self: *const Document, key: []const u8) ?Value { + return get_pair(self.pairs, key); + } + + pub fn alloc(allocator: std.mem.Allocator, pairs: []const Pair) !Document { + var arena = std.heap.ArenaAllocator.init(allocator); + errdefer arena.deinit(); + const copied = try arena.allocator().dupe(Pair, pairs); + return .{ .arena = arena, .pairs = copied }; + } + + /// Serialize the full document (length-prefixed) into `out`. + pub fn to_bytes(self: *const Document, gpa: std.mem.Allocator, out: *std.ArrayListUnmanaged(u8)) !void { + try write_doc(self.pairs, gpa, out); + } +}; + +pub fn get_pair(pairs: []const Pair, key: []const u8) ?Value { + for (pairs) |p| { + if (std.mem.eql(u8, p.key, key)) return p.value; + } + return null; +} + +pub fn parse_doc(allocator: std.mem.Allocator, bytes: []const u8) !Document { + return Document.parse(allocator, bytes); +} + +// --------------------------------------------------------------------------- +// Parsing +// --------------------------------------------------------------------------- + +const Parser = struct { + arena: *std.heap.ArenaAllocator, + bytes: []const u8, + + fn fail() error{InvalidBson} { + return error.InvalidBson; + } +}; + +const ParseError = error{ InvalidBson, OutOfMemory }; + +fn parse_doc_into(arena: *std.heap.ArenaAllocator, bytes: []const u8, idx: *usize) ParseError![]const Pair { + const p = Parser{ .arena = arena, .bytes = bytes }; + return parse_doc_inner(p, idx); +} + +fn ensure_available(bytes: []const u8, idx: usize, n: usize) error{InvalidBson}!void { + if (bytes.len -| idx < n) return error.InvalidBson; +} + +fn parse_doc_inner(p: Parser, idx: *usize) ParseError![]const Pair { + const start = idx.*; + ensure_available(p.bytes, start, 4) catch return Parser.fail(); + const total: u32 = std.mem.readInt(u32, p.bytes[start..][0..4], .little); + if (total < 5) return Parser.fail(); + if (p.bytes.len - start < total) return Parser.fail(); + const end = start + total; + if (p.bytes[end - 1] != 0x00) return Parser.fail(); + + const gpa = p.arena.allocator(); + var pairs: std.ArrayListUnmanaged(Pair) = .empty; + errdefer pairs.deinit(gpa); + + idx.* = start + 4; + while (idx.* < end - 1) { + const value = try parse_element(p, idx); + try pairs.append(gpa, value); + } + if (idx.* != end - 1) return Parser.fail(); + idx.* = end; + return pairs.toOwnedSlice(gpa); +} + +fn parse_element(p: Parser, idx: *usize) ParseError!Pair { + ensure_available(p.bytes, idx.*, 1) catch return Parser.fail(); + const tag = p.bytes[idx.*]; + idx.* += 1; + const key = try parse_cstring(p, idx); + const value = try parse_value(p, tag, idx); + return .{ .key = key, .value = value }; +} + +fn parse_cstring(p: Parser, idx: *usize) ParseError![]const u8 { + const start = idx.*; + while (idx.* < p.bytes.len and p.bytes[idx.*] != 0) idx.* += 1; + if (idx.* >= p.bytes.len) return Parser.fail(); + idx.* += 1; + // Strings are copied into the arena so documents are self-contained and + // outlive the input buffer (wire messages and log records are transient). + return p.arena.allocator().dupe(u8, p.bytes[start .. idx.* - 1]); +} + +fn parse_string(p: Parser, idx: *usize) ParseError![]const u8 { + ensure_available(p.bytes, idx.*, 4) catch return Parser.fail(); + const len: u32 = std.mem.readInt(u32, p.bytes[idx.*..][0..4], .little); + if (len == 0 or len > p.bytes.len - (idx.* + 4)) return Parser.fail(); + const str = p.bytes[idx.* + 4 .. idx.* + 4 + len]; + if (str[len - 1] != 0) return Parser.fail(); + idx.* += 4 + len; + return p.arena.allocator().dupe(u8, str[0 .. len - 1]); +} + +fn parse_value(p: Parser, tag: u8, idx: *usize) ParseError!Value { + return switch (tag) { + 0x01 => blk: { + ensure_available(p.bytes, idx.*, 8) catch return Parser.fail(); + const v: f64 = @bitCast(std.mem.readInt(u64, p.bytes[idx.*..][0..8], .little)); + idx.* += 8; + break :blk .{ .double = v }; + }, + 0x02 => .{ .string = try parse_string(p, idx) }, + 0x03 => .{ .doc = try parse_doc_inner(p, idx) }, + 0x04 => .{ .array = try parse_array(p, idx) }, + 0x05 => blk: { + ensure_available(p.bytes, idx.*, 5) catch return Parser.fail(); + const len: u32 = std.mem.readInt(u32, p.bytes[idx.*..][0..4], .little); + const subtype = p.bytes[idx.* + 4]; + ensure_available(p.bytes, idx.* + 5, len) catch return Parser.fail(); + const data = p.bytes[idx.* + 5 .. idx.* + 5 + len]; + idx.* += 5 + len; + break :blk .{ .binary = .{ .subtype = subtype, .data = try p.arena.allocator().dupe(u8, data) } }; + }, + 0x06 => .{ .opaque_val = .{ .kind = 0x06, .data = &.{} } }, + 0x07 => blk: { + ensure_available(p.bytes, idx.*, 12) catch return Parser.fail(); + const oid: ObjectId = p.bytes[idx.*..][0..12].*; + idx.* += 12; + break :blk .{ .object_id = oid }; + }, + 0x08 => blk: { + ensure_available(p.bytes, idx.*, 1) catch return Parser.fail(); + const v = p.bytes[idx.*]; + if (v > 1) return Parser.fail(); + idx.* += 1; + break :blk .{ .bool = v == 1 }; + }, + 0x09 => blk: { + ensure_available(p.bytes, idx.*, 8) catch return Parser.fail(); + const v: i64 = std.mem.readInt(i64, p.bytes[idx.*..][0..8], .little); + idx.* += 8; + break :blk .{ .datetime = v }; + }, + 0x0A => .null, + 0x0B => .{ .regex = .{ + .pattern = try parse_cstring(p, idx), + .options = try parse_cstring(p, idx), + } }, + 0x0C => blk: { + const start = idx.*; + _ = try parse_string(p, idx); + ensure_available(p.bytes, idx.*, 12) catch return Parser.fail(); + idx.* += 12; + // Like all other types, the payload is copied into the arena so + // documents stay valid after the input buffer is reused. + break :blk .{ .opaque_val = .{ .kind = 0x0C, .data = try p.arena.allocator().dupe(u8, p.bytes[start..idx.*]) } }; + }, + 0x0D => .{ .code = try parse_string(p, idx) }, + 0x0E => .{ .symbol = try parse_string(p, idx) }, + 0x0F => blk: { + ensure_available(p.bytes, idx.*, 4) catch return Parser.fail(); + const total: u32 = std.mem.readInt(u32, p.bytes[idx.*..][0..4], .little); + if (total < 4 or total > p.bytes.len - idx.*) return Parser.fail(); + const data = p.bytes[idx.* .. idx.* + total]; + idx.* += total; + break :blk .{ .opaque_val = .{ .kind = 0x0F, .data = try p.arena.allocator().dupe(u8, data) } }; + }, + 0x10 => blk: { + ensure_available(p.bytes, idx.*, 4) catch return Parser.fail(); + const v: i32 = std.mem.readInt(i32, p.bytes[idx.*..][0..4], .little); + idx.* += 4; + break :blk .{ .int32 = v }; + }, + 0x11 => blk: { + ensure_available(p.bytes, idx.*, 8) catch return Parser.fail(); + const v: u64 = std.mem.readInt(u64, p.bytes[idx.*..][0..8], .little); + idx.* += 8; + break :blk .{ .timestamp = v }; + }, + 0x12 => blk: { + ensure_available(p.bytes, idx.*, 8) catch return Parser.fail(); + const v: i64 = std.mem.readInt(i64, p.bytes[idx.*..][0..8], .little); + idx.* += 8; + break :blk .{ .int64 = v }; + }, + 0x13 => blk: { + ensure_available(p.bytes, idx.*, 16) catch return Parser.fail(); + const v: [16]u8 = p.bytes[idx.*..][0..16].*; + idx.* += 16; + break :blk .{ .decimal128 = v }; + }, + else => return Parser.fail(), + }; +} + +fn parse_array(p: Parser, idx: *usize) ParseError![]const Value { + const start = idx.*; + ensure_available(p.bytes, start, 4) catch return Parser.fail(); + const total: u32 = std.mem.readInt(u32, p.bytes[start..][0..4], .little); + if (total < 5) return Parser.fail(); + if (p.bytes.len - start < total) return Parser.fail(); + const end = start + total; + if (p.bytes[end - 1] != 0x00) return Parser.fail(); + + const gpa = p.arena.allocator(); + var values: std.ArrayListUnmanaged(Value) = .empty; + errdefer values.deinit(gpa); + + idx.* = start + 4; + while (idx.* < end - 1) { + ensure_available(p.bytes, idx.*, 1) catch return Parser.fail(); + const tag = p.bytes[idx.*]; + idx.* += 1; + _ = try parse_cstring(p, idx); + try values.append(gpa, try parse_value(p, tag, idx)); + } + if (idx.* != end - 1) return Parser.fail(); + idx.* = end; + return values.toOwnedSlice(gpa); +} + +// --------------------------------------------------------------------------- +// Serialization +// --------------------------------------------------------------------------- + +pub const SerializeError = error{ + BsonTooLarge, + BsonNulInKey, + BsonNotSerializable, + OutOfMemory, +}; + +pub fn write_value(v: Value, gpa: std.mem.Allocator, out: *std.ArrayListUnmanaged(u8)) SerializeError!void { + switch (v) { + .double => |d| { + var buf: [8]u8 = undefined; + std.mem.writeInt(u64, &buf, @bitCast(d), .little); + try out.appendSlice(gpa, &buf); + }, + .string => |s| try write_string(s, gpa, out), + .doc => |pairs| try write_doc(pairs, gpa, out), + .array => |items| try write_array(items, gpa, out), + .binary => |b| { + if (b.data.len > std.math.maxInt(u32)) return error.BsonTooLarge; + var buf: [5]u8 = undefined; + std.mem.writeInt(u32, buf[0..4], @intCast(b.data.len), .little); + buf[4] = b.subtype; + try out.appendSlice(gpa, &buf); + try out.appendSlice(gpa, b.data); + }, + .object_id => |oid| try out.appendSlice(gpa, &oid), + .bool => |b| try out.append(gpa, @intFromBool(b)), + .datetime => |t| { + var buf: [8]u8 = undefined; + std.mem.writeInt(i64, &buf, t, .little); + try out.appendSlice(gpa, &buf); + }, + .null => {}, + .regex => |r| { + try write_cstring(r.pattern, gpa, out); + try write_cstring(r.options, gpa, out); + }, + .code => |c| try write_string(c, gpa, out), + .symbol => |s| try write_string(s, gpa, out), + .int32 => |i| { + var buf: [4]u8 = undefined; + std.mem.writeInt(i32, &buf, i, .little); + try out.appendSlice(gpa, &buf); + }, + .timestamp => |t| { + var buf: [8]u8 = undefined; + std.mem.writeInt(u64, &buf, t, .little); + try out.appendSlice(gpa, &buf); + }, + .int64 => |i| { + var buf: [8]u8 = undefined; + std.mem.writeInt(i64, &buf, i, .little); + try out.appendSlice(gpa, &buf); + }, + .decimal128 => |d| try out.appendSlice(gpa, &d), + .opaque_val => |o| try out.appendSlice(gpa, o.data), + .min_key, .max_key => return error.BsonNotSerializable, + } +} + +pub fn write_cstring(s: []const u8, gpa: std.mem.Allocator, out: *std.ArrayListUnmanaged(u8)) SerializeError!void { + if (std.mem.indexOfScalar(u8, s, 0) != null) return error.BsonNulInKey; + try out.appendSlice(gpa, s); + try out.append(gpa, 0); +} + +pub fn write_string(s: []const u8, gpa: std.mem.Allocator, out: *std.ArrayListUnmanaged(u8)) SerializeError!void { + if (s.len + 1 > std.math.maxInt(u32)) return error.BsonTooLarge; + var buf: [4]u8 = undefined; + std.mem.writeInt(u32, &buf, @intCast(s.len + 1), .little); + try out.appendSlice(gpa, &buf); + try out.appendSlice(gpa, s); + try out.append(gpa, 0); +} + +pub fn write_element(pair: Pair, gpa: std.mem.Allocator, out: *std.ArrayListUnmanaged(u8)) SerializeError!void { + try out.append(gpa, pair.value.type_tag()); + try write_cstring(pair.key, gpa, out); + try write_value(pair.value, gpa, out); +} + +/// Write a length-prefixed document. Length is patched in after the body. +pub fn write_doc(pairs: []const Pair, gpa: std.mem.Allocator, out: *std.ArrayListUnmanaged(u8)) SerializeError!void { + const len_pos = out.items.len; + var zero: [4]u8 = [4]u8{ 0, 0, 0, 0 }; + try out.appendSlice(gpa, &zero); + for (pairs) |p| try write_element(p, gpa, out); + try out.append(gpa, 0); + const total = out.items.len - len_pos; + if (total > std.math.maxInt(u32)) return error.BsonTooLarge; + std.mem.writeInt(u32, out.items[len_pos..][0..4], @intCast(total), .little); +} + +fn write_array(items: []const Value, gpa: std.mem.Allocator, out: *std.ArrayListUnmanaged(u8)) !void { + const len_pos = out.items.len; + var zero: [4]u8 = [4]u8{ 0, 0, 0, 0 }; + try out.appendSlice(gpa, &zero); + var buf: [16]u8 = undefined; + for (items, 0..) |item, i| { + try out.append(gpa, item.type_tag()); + const key = std.fmt.bufPrint(&buf, "{d}", .{i}) catch unreachable; + try write_cstring(key, gpa, out); + try write_value(item, gpa, out); + } + try out.append(gpa, 0); + const total = out.items.len - len_pos; + if (total > std.math.maxInt(u32)) return error.BsonTooLarge; + std.mem.writeInt(u32, out.items[len_pos..][0..4], @intCast(total), .little); +} + +/// Serialize a single value with its type byte (no key) — used for `_id` +/// map keys and index entries. +pub fn serialize_value(gpa: std.mem.Allocator, v: Value) ![]u8 { + var out: std.ArrayListUnmanaged(u8) = .empty; + errdefer out.deinit(gpa); + try out.append(gpa, v.type_tag()); + try write_value(v, gpa, &out); + return out.toOwnedSlice(gpa); +} + +/// Deep-copy a value into `arena`, so the copy is self-contained. +pub fn copy_value(arena: std.mem.Allocator, v: Value) std.mem.Allocator.Error!Value { + return switch (v) { + .string => |s| .{ .string = try arena.dupe(u8, s) }, + .symbol => |s| .{ .symbol = try arena.dupe(u8, s) }, + .code => |c| .{ .code = try arena.dupe(u8, c) }, + .doc => |pairs| .{ .doc = try copy_pairs(arena, pairs) }, + .array => |items| .{ .array = try copy_values(arena, items) }, + .binary => |b| .{ .binary = .{ .subtype = b.subtype, .data = try arena.dupe(u8, b.data) } }, + .regex => |r| .{ .regex = .{ + .pattern = try arena.dupe(u8, r.pattern), + .options = try arena.dupe(u8, r.options), + } }, + .opaque_val => |o| .{ .opaque_val = .{ .kind = o.kind, .data = try arena.dupe(u8, o.data) } }, + else => v, + }; +} + +pub fn copy_pairs(arena: std.mem.Allocator, pairs: []const Pair) std.mem.Allocator.Error![]const Pair { + const out = try arena.alloc(Pair, pairs.len); + for (pairs, 0..) |p, i| { + out[i] = .{ .key = try arena.dupe(u8, p.key), .value = try copy_value(arena, p.value) }; + } + return out; +} + +fn copy_values(arena: std.mem.Allocator, items: []const Value) std.mem.Allocator.Error![]const Value { + const out = try arena.alloc(Value, items.len); + for (items, 0..) |item, i| out[i] = try copy_value(arena, item); + return out; +} + +// --------------------------------------------------------------------------- +// ObjectId generation +// --------------------------------------------------------------------------- + +pub const ObjectIdGen = struct { + random_prefix: [5]u8, + counter: u32, + + pub fn init(io: std.Io) ObjectIdGen { + var self: ObjectIdGen = undefined; + io.random(&self.random_prefix); + self.counter = 0; + return self; + } + + pub fn new(self: *ObjectIdGen, io: std.Io) ObjectId { + var oid: ObjectId = undefined; + const now = std.Io.Timestamp.now(io, .real); + const secs: u32 = @truncate(@as(u64, @intCast(now.toSeconds()))); + std.mem.writeInt(u32, oid[0..4], secs, .big); + @memcpy(oid[4..9], &self.random_prefix); + self.counter +%= 1; + std.mem.writeInt(u24, oid[9..12], @truncate(self.counter), .big); + return oid; + } +}; + +// --------------------------------------------------------------------------- +// Canonical BSON comparison order +// --------------------------------------------------------------------------- + +fn rank(v: Value) u8 { + return switch (v) { + .min_key => 0, + .null => 1, + .double, .int32, .int64 => 2, + .string, .symbol, .code => 3, + .doc => 4, + .array => 5, + .binary => 6, + .object_id => 7, + .bool => 8, + .datetime => 9, + .timestamp => 10, + .regex => 11, + .opaque_val => 12, + .decimal128 => 13, + .max_key => 14, + }; +} + +pub fn compare(a: Value, b: Value) std.math.Order { + const ra = rank(a); + const rb = rank(b); + if (ra != rb) return std.math.order(ra, rb); + return switch (ra) { + 0, 1 => .eq, + 2 => compare_f128(a.as_f128(), b.as_f128()), + 3 => std.mem.order(u8, as_str(a), as_str(b)), + 4 => compare_docs(a.doc, b.doc), + 5 => compare_arrays(a.array, b.array), + 6 => compare_binary(a.binary, b.binary), + 7 => std.mem.order(u8, &a.object_id, &b.object_id), + 8 => std.math.order(@intFromBool(a.bool), @intFromBool(b.bool)), + 9 => std.math.order(a.datetime, b.datetime), + 10 => std.math.order(a.timestamp, b.timestamp), + 11 => blk: { + const p = std.mem.order(u8, as_regex(a).pattern, as_regex(b).pattern); + break :blk if (p != .eq) p else std.mem.order(u8, as_regex(a).options, as_regex(b).options); + }, + 12 => std.mem.order(u8, a.opaque_val.data, b.opaque_val.data), + 13 => std.mem.order(u8, &a.decimal128, &b.decimal128), + 14 => .eq, + else => unreachable, + }; +} + +fn as_str(v: Value) []const u8 { + return switch (v) { + .string => |s| s, + .symbol => |s| s, + .code => |c| c, + else => unreachable, + }; +} + +fn as_regex(v: Value) Regex { + return switch (v) { + .regex => |r| r, + else => unreachable, + }; +} + +fn compare_f128(a: f128, b: f128) std.math.Order { + if (a < b) return .lt; + if (a > b) return .gt; + // Distinguish -0.0 from +0.0 like MongoDB does (they are equal); NaN + // sorts greater than every number (MongoDB treats NaN as largest). + if (std.math.isNan(a)) { + if (std.math.isNan(b)) return .eq; + return .gt; + } + if (std.math.isNan(b)) return .lt; + return .eq; +} + +fn compare_docs(a: []const Pair, b: []const Pair) std.math.Order { + const n = @min(a.len, b.len); + for (a[0..n], b[0..n]) |pa, pb| { + // Keys tie-break equal values, so {a: 1} and {b: 1} are distinct. + const ko = std.mem.order(u8, pa.key, pb.key); + if (ko != .eq) return ko; + const o = compare(pa.value, pb.value); + if (o != .eq) return o; + } + return std.math.order(a.len, b.len); +} + +fn compare_arrays(a: []const Value, b: []const Value) std.math.Order { + const n = @min(a.len, b.len); + for (a[0..n], b[0..n]) |va, vb| { + const o = compare(va, vb); + if (o != .eq) return o; + } + return std.math.order(a.len, b.len); +} + +fn compare_binary(a: Binary, b: Binary) std.math.Order { + const l = std.math.order(a.data.len, b.data.len); + if (l != .eq) return l; + const d = std.mem.order(u8, a.data, b.data); + if (d != .eq) return d; + return std.math.order(a.subtype, b.subtype); +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +const testing = std.testing; + +test "document round-trip" { + var out: std.ArrayListUnmanaged(u8) = .empty; + defer out.deinit(testing.allocator); + + try write_doc(&.{ + .{ .key = "_id", .value = .{ .int32 = 7 } }, + .{ .key = "name", .value = .{ .string = "héllo" } }, + .{ .key = "pi", .value = .{ .double = 3.25 } }, + .{ .key = "ok", .value = .{ .bool = true } }, + .{ .key = "nul", .value = .null }, + .{ .key = "big", .value = .{ .int64 = 1 << 40 } }, + .{ .key = "when", .value = .{ .datetime = 1_700_000_000_000 } }, + .{ .key = "re", .value = .{ .regex = .{ .pattern = "^a", .options = "i" } } }, + .{ .key = "bin", .value = .{ .binary = .{ .subtype = 0x80, .data = &[_]u8{ 1, 2, 3 } } } }, + .{ .key = "arr", .value = .{ .array = &.{ .{ .int32 = 1 }, .{ .string = "x" } } } }, + .{ .key = "sub", .value = .{ .doc = &.{.{ .key = "a", .value = .{ .int32 = 1 } }} } }, + .{ .key = "ts", .value = .{ .timestamp = 42 } }, + .{ .key = "oid", .value = .{ .object_id = [_]u8{ 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12 } } }, + }, testing.allocator, &out); + + var doc = try Document.parse(testing.allocator, out.items); + defer doc.deinit(); + + try testing.expectEqual(@as(i64, 7), doc.get("_id").?.int32); + try testing.expectEqualStrings("héllo", doc.get("name").?.string); + try testing.expectEqual(@as(f64, 3.25), doc.get("pi").?.double); + try testing.expect(doc.get("ok").?.bool); + try testing.expectEqual(@as(i64, 1 << 40), doc.get("big").?.int64); + try testing.expectEqual(@as(i64, 1_700_000_000_000), doc.get("when").?.datetime); + try testing.expectEqualStrings("^a", doc.get("re").?.regex.pattern); + try testing.expectEqualStrings("i", doc.get("re").?.regex.options); + try testing.expectEqualSlices(u8, &[_]u8{ 1, 2, 3 }, doc.get("bin").?.binary.data); + try testing.expectEqual(@as(usize, 2), doc.get("arr").?.array.len); + try testing.expectEqual(@as(i64, 1), doc.get("sub").?.doc[0].value.int32); + try testing.expectEqual(@as(u64, 42), doc.get("ts").?.timestamp); + try testing.expectEqualSlices(u8, &[_]u8{ 1, 2, 3, 4, 5, 6, 7, 8, 9, 10, 11, 12 }, &doc.get("oid").?.object_id); +} + +test "array and nested doc round-trip" { + var out: std.ArrayListUnmanaged(u8) = .empty; + defer out.deinit(testing.allocator); + try write_doc(&.{ + .{ .key = "a", .value = .{ .array = &.{ + .{ .int32 = 10 }, + .{ .doc = &.{.{ .key = "deep", .value = .{ .string = "v" } }} }, + } } }, + }, testing.allocator, &out); + + var doc = try Document.parse(testing.allocator, out.items); + defer doc.deinit(); + const a = doc.get("a").?.array; + try testing.expectEqual(@as(i64, 10), a[0].int32); + try testing.expectEqualStrings("v", a[1].doc[0].value.string); +} + +test "reject truncated document" { + try testing.expectError(error.InvalidBson, Document.parse(testing.allocator, &[_]u8{ 6, 0, 0, 0 })); + try testing.expectError(error.InvalidBson, Document.parse(testing.allocator, &[_]u8{ 4, 0, 0, 0, 0 })); + try testing.expectError(error.InvalidBson, Document.parse(testing.allocator, &[_]u8{ 9, 0, 0, 0, 0x01, 'a', 0 })); +} + +test "compare: canonical order" { + const min = Value{ .min_key = {} }; + const nul = Value.null; + const i32a = Value{ .int32 = 5 }; + const i64a = Value{ .int64 = 5 }; + const dbl = Value{ .double = 4.9 }; + const str = Value{ .string = "a" }; + const obj = Value{ .doc = &.{} }; + const arr = Value{ .array = &.{} }; + const oid = Value{ .object_id = [_]u8{0} ** 12 }; + const btrue = Value{ .bool = true }; + const bfalse = Value{ .bool = false }; + const max = Value{ .max_key = {} }; + + try testing.expectEqual(std.math.Order.lt, compare(min, nul)); + try testing.expectEqual(std.math.Order.lt, compare(nul, i32a)); + try testing.expectEqual(std.math.Order.eq, compare(i32a, i64a)); // numeric equality across widths + try testing.expectEqual(std.math.Order.gt, compare(i32a, dbl)); // 5 > 4.9 + try testing.expectEqual(std.math.Order.lt, compare(i32a, str)); + try testing.expectEqual(std.math.Order.lt, compare(str, obj)); + try testing.expectEqual(std.math.Order.lt, compare(obj, arr)); + try testing.expectEqual(std.math.Order.lt, compare(arr, oid)); + try testing.expectEqual(std.math.Order.lt, compare(oid, bfalse)); + try testing.expectEqual(std.math.Order.lt, compare(bfalse, btrue)); + try testing.expectEqual(std.math.Order.gt, compare(max, nul)); +} + +test "compare: strings and docs" { + try testing.expectEqual(std.math.Order.lt, compare(.{ .string = "a" }, .{ .string = "b" })); + try testing.expectEqual(std.math.Order.eq, compare(.{ .string = "x" }, .{ .string = "x" })); + try testing.expectEqual(std.math.Order.lt, compare( + .{ .doc = &.{.{ .key = "a", .value = .{ .int32 = 1 } }} }, + .{ .doc = &.{.{ .key = "a", .value = .{ .int32 = 2 } }} }, + )); + try testing.expectEqual(std.math.Order.lt, compare( + .{ .array = &.{.{ .int32 = 1 }} }, + .{ .array = &.{ .{ .int32 = 1 }, .{ .int32 = 1 } } }, + )); +} + +test "compare: NaN is greatest number" { + try testing.expectEqual(std.math.Order.gt, compare(.{ .double = std.math.nan(f64) }, .{ .int64 = 1 << 62 })); +} + +test "object id generation" { + var threaded: std.Io.Threaded = .init_single_threaded; + defer threaded.deinit(); + const io = threaded.io(); + + var gen = ObjectIdGen.init(io); + const a = gen.new(io); + const b = gen.new(io); + try testing.expect(!std.mem.eql(u8, &a, &b)); + // timestamp bytes match wall clock roughly + try testing.expect(a[0] >= 0x66); // 2024+ in big-endian seconds + // The 5-byte random prefix is untouched by the counter (spec layout: + // 4s timestamp | 5B random | 3B counter). + try testing.expectEqualSlices(u8, a[0..9], b[0..9]); + try testing.expect(!std.mem.eql(u8, a[9..12], b[9..12])); +} + +test "serialize_value deterministic for _id keys" { + const gpa = testing.allocator; + const v1 = try serialize_value(gpa, .{ .doc = &.{ + .{ .key = "a", .value = .{ .int32 = 1 } }, + .{ .key = "b", .value = .{ .string = "s" } }, + } }); + defer gpa.free(v1); + const v2 = try serialize_value(gpa, .{ .doc = &.{ + .{ .key = "a", .value = .{ .int32 = 1 } }, + .{ .key = "b", .value = .{ .string = "s" } }, + } }); + defer gpa.free(v2); + try testing.expectEqualSlices(u8, v1, v2); +} diff --git a/src/commands.zig b/src/commands.zig new file mode 100644 index 0000000..e6d9538 --- /dev/null +++ b/src/commands.zig @@ -0,0 +1,1259 @@ +//! MongoDB command dispatch. Each command fills `reply` with its result; +//! unknown commands and failures produce error replies with real codes. + +const std = @import("std"); +const builtin = @import("builtin"); +const bson = @import("bson.zig"); +const wire = @import("wire.zig"); +const db = @import("db.zig"); +const query = @import("query.zig"); +const update = @import("update.zig"); + +pub const Context = struct { + gpa: std.mem.Allocator, + io: std.Io, + oid_gen: *bson.ObjectIdGen, + connection_id: u32, + client_desc: []const u8, + engine: *db.Engine, + server_start: std.Io.Timestamp, +}; + +pub const ErrorCode = enum(i32) { + command_not_found = 59, + bad_value = 2, + invalid_argument = 72, + namespace_not_found = 26, + duplicate_key = 11000, + namespace_exists = 48, + failed_to_parse = 9, + internal_error = 1, + invalid_pipeline_operator = 40324, +}; + +pub fn dispatch(ctx: *Context, msg: *wire.Message, reply: *wire.Reply) !void { + try ctx.engine.lock(); + defer ctx.engine.unlock(); + + const name = msg.command_name(); + if (std.mem.eql(u8, name, "hello")) return cmd_hello(ctx, reply); + if (std.mem.eql(u8, name, "isMaster") or std.mem.eql(u8, name, "ismaster")) return cmd_is_master(ctx, reply); + if (std.mem.eql(u8, name, "ping")) return cmd_ping(reply); + if (std.mem.eql(u8, name, "buildInfo")) return cmd_build_info(reply); + if (std.mem.eql(u8, name, "getParameter")) return cmd_get_parameter(ctx, msg, reply); + if (std.mem.eql(u8, name, "whatsmyuri")) return cmd_whatsmyuri(ctx, reply); + if (std.mem.eql(u8, name, "hostInfo")) return cmd_host_info(ctx, reply); + if (std.mem.eql(u8, name, "getCmdLineOpts")) return cmd_get_cmd_line_opts(reply); + if (std.mem.eql(u8, name, "serverStatus")) return cmd_server_status(ctx, reply); + if (std.mem.eql(u8, name, "endSessions")) return cmd_end_sessions(reply); + if (std.mem.eql(u8, name, "connectionStatus")) return cmd_connection_status(reply); + if (std.mem.eql(u8, name, "listDatabases")) return cmd_list_databases(ctx, reply); + if (std.mem.eql(u8, name, "listCollections")) return cmd_list_collections(ctx, msg, reply); + if (std.mem.eql(u8, name, "create")) return cmd_create(ctx, msg, reply); + if (std.mem.eql(u8, name, "drop")) return cmd_drop(ctx, msg, reply); + if (std.mem.eql(u8, name, "dropDatabase")) return cmd_drop_database(ctx, msg, reply); + if (std.mem.eql(u8, name, "insert")) return cmd_insert(ctx, msg, reply); + if (std.mem.eql(u8, name, "find")) return cmd_find(ctx, msg, reply); + if (std.mem.eql(u8, name, "update")) return cmd_update(ctx, msg, reply); + if (std.mem.eql(u8, name, "delete")) return cmd_delete(ctx, msg, reply); + if (std.mem.eql(u8, name, "findAndModify")) return cmd_find_and_modify(ctx, msg, reply); + if (std.mem.eql(u8, name, "count")) return cmd_count(ctx, msg, reply); + if (std.mem.eql(u8, name, "aggregate")) return cmd_aggregate(ctx, msg, reply); + if (std.mem.eql(u8, name, "getMore")) return cmd_get_more(reply); + if (std.mem.eql(u8, name, "killCursors")) return cmd_kill_cursors(reply); + + var buf: [256]u8 = undefined; + const errmsg = try std.fmt.bufPrint(&buf, "no such command: '{s}'", .{name}); + return reply.put_error(@intFromEnum(ErrorCode.command_not_found), "CommandNotFound", errmsg); +} + +pub fn reply_error(reply: *wire.Reply, code: i32, code_name: []const u8, message: []const u8) !void { + try reply.put_error(code, code_name, message); +} + +// --------------------------------------------------------------------------- +// Handshake / administration +// --------------------------------------------------------------------------- + +fn add_server_info(ctx: *Context, reply: *wire.Reply) !void { + try reply.put("isWritablePrimary", .{ .bool = true }); + try reply.put("maxBsonObjectSize", .{ .int32 = wire.max_bson_object_size }); + try reply.put("maxMessageSizeBytes", .{ .int32 = 48000000 }); + try reply.put("maxWriteBatchSize", .{ .int32 = 100000 }); + const now = std.Io.Timestamp.now(ctx.io, .real); + try reply.put("localTime", .{ .datetime = now.toMilliseconds() }); + try reply.put("logicalSessionTimeoutMinutes", .{ .int32 = 30 }); + try reply.put("connectionId", .{ .int32 = @intCast(ctx.connection_id) }); + try reply.put("minWireVersion", .{ .int32 = 0 }); + try reply.put("maxWireVersion", .{ .int32 = 8 }); + try reply.put("readOnly", .{ .bool = false }); + + const oid = ctx.oid_gen.new(ctx.io); + const tv = try reply.arena_alloc().alloc(bson.Pair, 2); + tv[0] = .{ .key = try reply.arena_alloc().dupe(u8, "processId"), .value = .{ .object_id = oid } }; + tv[1] = .{ .key = try reply.arena_alloc().dupe(u8, "counter"), .value = .{ .int64 = 0 } }; + try reply.put("topologyVersion", .{ .doc = tv }); +} + +fn cmd_hello(ctx: *Context, reply: *wire.Reply) !void { + try add_server_info(ctx, reply); + try reply.put_ok(); +} + +fn cmd_is_master(ctx: *Context, reply: *wire.Reply) !void { + try add_server_info(ctx, reply); + try reply.put("ismaster", .{ .bool = true }); + try reply.put("helloOk", .{ .bool = true }); + try reply.put_ok(); +} + +fn cmd_ping(reply: *wire.Reply) !void { + try reply.put_ok(); +} + +fn cmd_build_info(reply: *wire.Reply) !void { + try reply.put("version", .{ .string = "4.4.0" }); + try reply.put("gitVersion", .{ .string = "mongo-light" }); + try reply.put("versionArray", .{ .array = try int_array(reply, &.{ 4, 4, 0, 0 }) }); + try reply.put("openssl", .{ .doc = &.{} }); + try reply.put("loaderFlags", .{ .string = "" }); + try reply.put("compilerInfo", .{ .string = "zig 0.16.0" }); + try reply.put("allocator", .{ .string = "system" }); + try reply.put("javascriptEngine", .{ .string = "none" }); + try reply.put("bits", .{ .int32 = 64 }); + try reply.put("debug", .{ .bool = false }); + try reply.put("maxBsonObjectSize", .{ .int32 = wire.max_bson_object_size }); + try reply.put("storageEngines", .{ .array = try str_array(reply, &.{"wiredTiger"}) }); + try reply.put_ok(); +} + +fn cmd_get_parameter(ctx: *Context, msg: *wire.Message, reply: *wire.Reply) !void { + _ = ctx; + // mongosh probes featureCompatibilityVersion; respond per requested key. + const params_value = msg.body.get("getParameter") orelse return reply.put_error( + @intFromEnum(ErrorCode.invalid_argument), + "InvalidArgument", + "getParameter requires a document", + ); + const params = switch (params_value) { + .doc => |d| d, + else => return reply.put_error( + @intFromEnum(ErrorCode.invalid_argument), + "InvalidArgument", + "getParameter requires a document", + ), + }; + var found: bool = false; + for (params) |p| { + if (std.mem.eql(u8, p.key, "featureCompatibilityVersion")) { + const fcv = try reply.arena_alloc().alloc(bson.Pair, 1); + fcv[0] = .{ .key = try reply.arena_alloc().dupe(u8, "version"), .value = .{ .string = "4.4" } }; + try reply.put("featureCompatibilityVersion", .{ .doc = fcv }); + found = true; + } + } + if (!found) { + return reply.put_error( + @intFromEnum(ErrorCode.invalid_argument), + "InvalidArgument", + "no option found to get", + ); + } + try reply.put_ok(); +} + +fn cmd_whatsmyuri(ctx: *Context, reply: *wire.Reply) !void { + try reply.put("you", .{ .string = ctx.client_desc }); + try reply.put_ok(); +} + +fn cmd_host_info(ctx: *Context, reply: *wire.Reply) !void { + const now = std.Io.Timestamp.now(ctx.io, .real); + const cpu_count: usize = std.Thread.getCpuCount() catch 1; + + const system = try reply.arena_alloc().alloc(bson.Pair, 6); + system[0] = .{ .key = try reply.arena_alloc().dupe(u8, "currentTime"), .value = .{ .datetime = now.toMilliseconds() } }; + system[1] = .{ .key = try reply.arena_alloc().dupe(u8, "hostname"), .value = .{ .string = "localhost" } }; + system[2] = .{ .key = try reply.arena_alloc().dupe(u8, "cpuAddrSize"), .value = .{ .int32 = 64 } }; + system[3] = .{ .key = try reply.arena_alloc().dupe(u8, "memSizeMB"), .value = .{ .int32 = 0 } }; + system[4] = .{ .key = try reply.arena_alloc().dupe(u8, "numCores"), .value = .{ .int32 = @intCast(cpu_count) } }; + system[5] = .{ .key = try reply.arena_alloc().dupe(u8, "cpuArch"), .value = .{ .string = @tagName(builtin.cpu.arch) } }; + try reply.put("system", .{ .doc = system }); + + const os = try reply.arena_alloc().alloc(bson.Pair, 3); + os[0] = .{ .key = try reply.arena_alloc().dupe(u8, "type"), .value = .{ .string = @tagName(builtin.os.tag) } }; + os[1] = .{ .key = try reply.arena_alloc().dupe(u8, "name"), .value = .{ .string = @tagName(builtin.os.tag) } }; + os[2] = .{ .key = try reply.arena_alloc().dupe(u8, "version"), .value = .{ .string = "unknown" } }; + try reply.put("os", .{ .doc = os }); + + try reply.put("extra", .{ .doc = &.{} }); + try reply.put_ok(); +} + +fn cmd_get_cmd_line_opts(reply: *wire.Reply) !void { + const argv = try reply.arena_alloc().alloc(bson.Pair, 2); + argv[0] = .{ .key = try reply.arena_alloc().dupe(u8, "dbpath"), .value = .{ .string = "mongo-light.log" } }; + argv[1] = .{ .key = try reply.arena_alloc().dupe(u8, "port"), .value = .{ .int32 = 27017 } }; + try reply.put("argv", .{ .array = &.{} }); + try reply.put("parsed", .{ .doc = argv }); + try reply.put_ok(); +} + +fn cmd_server_status(ctx: *Context, reply: *wire.Reply) !void { + const now = std.Io.Timestamp.now(ctx.io, .real); + const uptime: i64 = std.Io.Timestamp.durationTo(ctx.server_start, now).toSeconds(); + + try reply.put("host", .{ .string = "localhost" }); + try reply.put("version", .{ .string = "4.4.0" }); + try reply.put("process", .{ .string = "mongod" }); + try reply.put("uptime", .{ .double = @floatFromInt(uptime) }); + try reply.put("localTime", .{ .datetime = now.toMilliseconds() }); + const connections = try reply.arena_alloc().alloc(bson.Pair, 1); + connections[0] = .{ .key = try reply.arena_alloc().dupe(u8, "current"), .value = .{ .int32 = @intCast(ctx.connection_id) } }; + try reply.put("connections", .{ .doc = connections }); + try reply.put_ok(); +} + +fn cmd_end_sessions(reply: *wire.Reply) !void { + try reply.put_ok(); +} + +fn cmd_connection_status(reply: *wire.Reply) !void { + const auth_info = try reply.arena_alloc().alloc(bson.Pair, 2); + auth_info[0] = .{ .key = try reply.arena_alloc().dupe(u8, "authenticatedUsers"), .value = .{ .array = &.{} } }; + auth_info[1] = .{ .key = try reply.arena_alloc().dupe(u8, "authenticatedUserRoles"), .value = .{ .array = &.{} } }; + try reply.put("authInfo", .{ .doc = auth_info }); + try reply.put_ok(); +} + +fn cmd_list_databases(ctx: *Context, reply: *wire.Reply) !void { + var names: std.ArrayListUnmanaged([]const u8) = .empty; + defer names.deinit(ctx.gpa); + try ctx.engine.database_names(&names); + + const values = try reply.arena_alloc().alloc(bson.Value, names.items.len); + for (names.items, 0..) |n, i| { + const entry = try reply.arena_alloc().alloc(bson.Pair, 3); + entry[0] = .{ .key = try reply.arena_alloc().dupe(u8, "name"), .value = .{ .string = try reply.arena_alloc().dupe(u8, n) } }; + entry[1] = .{ .key = try reply.arena_alloc().dupe(u8, "sizeOnDisk"), .value = .{ .double = 0 } }; + entry[2] = .{ .key = try reply.arena_alloc().dupe(u8, "empty"), .value = .{ .bool = true } }; + values[i] = .{ .doc = entry }; + } + try reply.put("databases", .{ .array = values }); + try reply.put("totalSize", .{ .int32 = 0 }); + try reply.put("totalSizeMb", .{ .int32 = 0 }); + try reply.put_ok(); +} + +fn cmd_list_collections(ctx: *Context, msg: *wire.Message, reply: *wire.Reply) !void { + const db_name = msg.db_name() orelse return reply.put_error( + @intFromEnum(ErrorCode.invalid_argument), + "InvalidArgument", + "listCollections requires $db", + ); + var names: std.ArrayListUnmanaged([]const u8) = .empty; + defer names.deinit(ctx.gpa); + try ctx.engine.collection_names(db_name, &names); + + const values = try reply.arena_alloc().alloc(bson.Value, names.items.len); + for (names.items, 0..) |n, i| { + const entry = try reply.arena_alloc().alloc(bson.Pair, 3); + entry[0] = .{ .key = try reply.arena_alloc().dupe(u8, "name"), .value = .{ .string = try reply.arena_alloc().dupe(u8, n) } }; + entry[1] = .{ .key = try reply.arena_alloc().dupe(u8, "type"), .value = .{ .string = "collection" } }; + entry[2] = .{ .key = try reply.arena_alloc().dupe(u8, "options"), .value = .{ .doc = &.{} } }; + values[i] = .{ .doc = entry }; + } + try reply.put("cursor", .{ .doc = try cursor_doc(reply, 0, try format_namespace(reply, db_name, ""), values) }); + try reply.put_ok(); +} + +fn cmd_create(ctx: *Context, msg: *wire.Message, reply: *wire.Reply) !void { + const db_name = msg.db_name() orelse return reply.put_error( + @intFromEnum(ErrorCode.invalid_argument), + "InvalidArgument", + "create requires $db", + ); + const coll_name = msg.body.get("create") orelse return reply.put_error( + @intFromEnum(ErrorCode.bad_value), + "BadValue", + "create requires a collection name", + ); + switch (coll_name) { + .string => |s| { + _ = try ctx.engine.get_or_create_collection(db_name, s); + try reply.put_ok(); + }, + else => return reply.put_error(@intFromEnum(ErrorCode.bad_value), "BadValue", "collection name must be a string"), + } +} + +fn cmd_drop(ctx: *Context, msg: *wire.Message, reply: *wire.Reply) !void { + const db_name = msg.db_name() orelse return reply.put_error( + @intFromEnum(ErrorCode.invalid_argument), + "InvalidArgument", + "drop requires $db", + ); + const coll_name = msg.body.get("drop") orelse return reply.put_error( + @intFromEnum(ErrorCode.bad_value), + "BadValue", + "drop requires a collection name", + ); + switch (coll_name) { + .string => |s| { + const dropped = try ctx.engine.drop_collection(db_name, s); + if (!dropped) { + return reply.put_error(@intFromEnum(ErrorCode.namespace_not_found), "NamespaceNotFound", "ns not found"); + } + try reply.put_ok(); + }, + else => return reply.put_error(@intFromEnum(ErrorCode.bad_value), "BadValue", "collection name must be a string"), + } +} + +fn cmd_drop_database(ctx: *Context, msg: *wire.Message, reply: *wire.Reply) !void { + const db_name = msg.db_name() orelse return reply.put_error( + @intFromEnum(ErrorCode.invalid_argument), + "InvalidArgument", + "dropDatabase requires $db", + ); + _ = try ctx.engine.drop_database(db_name); + try reply.put("dropped", .{ .string = try reply.arena_alloc().dupe(u8, db_name) }); + try reply.put_ok(); +} + +// --------------------------------------------------------------------------- +// CRUD +// --------------------------------------------------------------------------- + +fn cmd_insert(ctx: *Context, msg: *wire.Message, reply: *wire.Reply) !void { + const db_name = msg.db_name() orelse return invalid_arg(reply, "insert requires $db"); + const coll_name = str_arg(msg.body.get("insert")) orelse return bad_value(reply, "insert requires a collection name"); + + var docs: []const bson.Document = &.{}; + for (msg.seqs) |seq| { + if (std.mem.eql(u8, seq.name, "documents")) docs = seq.docs; + } + if (docs.len == 0) { + const d = msg.body.get("documents") orelse return bad_value(reply, "insert requires documents"); + switch (d) { + .array => |arr| { + const parsed = try reply.arena_alloc().alloc(bson.Document, arr.len); + for (arr, 0..) |item, i| { + parsed[i] = switch (item) { + .doc => |pairs| bson.Document{ .arena = undefined, .pairs = pairs }, + else => return bad_value(reply, "documents must be documents"), + }; + } + docs = parsed; + }, + else => return bad_value(reply, "documents must be an array"), + } + } + + var inserted: i64 = 0; + var write_errors: std.ArrayListUnmanaged(bson.Value) = .empty; + defer write_errors.deinit(reply.arena_alloc()); + + for (docs, 0..) |*doc, i| { + if (ctx.engine.insert(db_name, coll_name, doc, ctx.oid_gen)) |_| { + inserted += 1; + } else |err| { + switch (err) { + error.DuplicateKey => { + const e = try reply.arena_alloc().alloc(bson.Pair, 3); + e[0] = .{ .key = try reply.arena_alloc().dupe(u8, "index"), .value = .{ .int32 = @intCast(i) } }; + e[1] = .{ .key = try reply.arena_alloc().dupe(u8, "code"), .value = .{ .int32 = @intFromEnum(ErrorCode.duplicate_key) } }; + const id = doc.get("_id") orelse bson.Value.null; + const key = try serialize_value_compact(reply, id); + const msg_text = try std.fmt.allocPrint(reply.arena_alloc(), "E11000 duplicate key error collection: {s}.{s} index: _id_ dup key: {s}", .{ db_name, coll_name, key }); + e[2] = .{ .key = try reply.arena_alloc().dupe(u8, "errmsg"), .value = .{ .string = msg_text } }; + try write_errors.append(reply.arena_alloc(), .{ .doc = e }); + }, + else => return err, + } + } + } + + try reply.put("n", .{ .int32 = @intCast(inserted) }); + if (write_errors.items.len > 0) { + // Copy into the reply arena: write_errors is freed when this command + // returns, before the reply is serialized. + const arr = try reply.arena_alloc().alloc(bson.Value, write_errors.items.len); + @memcpy(arr, write_errors.items); + try reply.put("writeErrors", .{ .array = arr }); + } + try reply.put_ok(); +} + +fn cmd_find(ctx: *Context, msg: *wire.Message, reply: *wire.Reply) !void { + const db_name = msg.db_name() orelse return invalid_arg(reply, "find requires $db"); + const coll_name = str_arg(msg.body.get("find")) orelse return bad_value(reply, "find requires a collection name"); + + const filter = doc_arg(msg.body.get("filter")) orelse return bad_value(reply, "filter must be a document"); + const filter_doc = bson.Document{ .arena = undefined, .pairs = filter }; + + const sort_keys = try parse_sort_keys(reply, msg.body.get("sort")); + const proj_pairs = doc_arg(msg.body.get("projection")); + const skip: u64 = int_arg(msg.body.get("skip")) orelse 0; + const limit_raw: i64 = switch (msg.body.get("limit") orelse bson.Value{ .int32 = 0 }) { + .int32 => |i| i, + .int64 => |i| i, + else => 0, + }; + const limit: usize = if (limit_raw < 0) @intCast(-limit_raw) else @intCast(limit_raw); + + var matched: std.ArrayListUnmanaged(*const bson.Document) = .empty; + defer matched.deinit(ctx.gpa); + + if (ctx.engine.get_collection(db_name, coll_name)) |coll| { + var it = coll.docs.iterator(); + while (it.next()) |entry| { + if (try query.matches(ctx.gpa, &filter_doc, entry.value_ptr.*)) { + try matched.append(ctx.gpa, entry.value_ptr.*); + } + } + } + + if (sort_keys.len > 0) { + try query.sort_docs(ctx.gpa, reply.arena_alloc(), matched.items, sort_keys); + } + if (skip < matched.items.len) { + const start = matched.items[skip..]; + const end = if (limit > 0 and limit < start.len) start[0..limit] else start; + try emit_docs(reply, db_name, coll_name, proj_pairs, end); + } else { + try emit_docs(reply, db_name, coll_name, proj_pairs, &.{}); + } + try reply.put_ok(); +} + +fn emit_docs(reply: *wire.Reply, db_name: []const u8, coll_name: []const u8, proj_pairs: ?[]const bson.Pair, docs: []const *const bson.Document) !void { + const values = try reply.arena_alloc().alloc(bson.Value, docs.len); + for (docs, 0..) |d, i| { + values[i] = try project_doc(reply, d, proj_pairs); + } + try reply.put("cursor", .{ .doc = try cursor_doc(reply, 0, try format_namespace(reply, db_name, coll_name), values) }); +} + +/// Project a stored doc (or deep-copy it) into the reply arena. +fn project_doc(reply: *wire.Reply, doc: *const bson.Document, proj_pairs: ?[]const bson.Pair) !bson.Value { + if (proj_pairs) |pp| { + var out: std.ArrayListUnmanaged(bson.Pair) = .empty; + errdefer out.deinit(reply.arena_alloc()); + try query.project(reply.arena_alloc(), doc, &.{ .arena = undefined, .pairs = pp }, &out); + return .{ .doc = out.items }; + } + return .{ .doc = try bson.copy_pairs(reply.arena_alloc(), doc.pairs) }; +} + +fn cmd_update(ctx: *Context, msg: *wire.Message, reply: *wire.Reply) !void { + const db_name = msg.db_name() orelse return invalid_arg(reply, "update requires $db"); + const coll_name = str_arg(msg.body.get("update")) orelse return bad_value(reply, "update requires a collection name"); + + var specs: []const bson.Document = &.{}; + for (msg.seqs) |seq| { + if (std.mem.eql(u8, seq.name, "updates")) specs = seq.docs; + } + if (specs.len == 0) { + const d = msg.body.get("updates") orelse return bad_value(reply, "update requires updates"); + const arr = switch (d) { + .array => |a| a, + else => return bad_value(reply, "updates must be an array"), + }; + const parsed = try reply.arena_alloc().alloc(bson.Document, arr.len); + for (arr, 0..) |item, i| { + parsed[i] = switch (item) { + .doc => |pairs| bson.Document{ .arena = undefined, .pairs = pairs }, + else => return bad_value(reply, "updates must be documents"), + }; + } + specs = parsed; + } + + var n_matched: i64 = 0; + var n_modified: i64 = 0; + var upserted: std.ArrayListUnmanaged(bson.Pair) = .empty; + defer upserted.deinit(reply.arena_alloc()); + + for (specs, 0..) |*spec, si| { + const q = doc_arg(spec.get("q")) orelse return bad_value(reply, "update spec requires q"); + const u_doc = doc_arg(spec.get("u")) orelse return bad_value(reply, "update spec requires u"); + const multi = bool_arg(spec.get("multi")) orelse false; + const upsert = bool_arg(spec.get("upsert")) orelse false; + + var matched: std.ArrayListUnmanaged(*bson.Document) = .empty; + defer matched.deinit(ctx.gpa); + if (ctx.engine.get_collection(db_name, coll_name)) |coll| { + var it = coll.docs.iterator(); + while (it.next()) |entry| { + if (try query.matches(ctx.gpa, &.{ .arena = undefined, .pairs = q }, entry.value_ptr.*)) { + try matched.append(ctx.gpa, entry.value_ptr.*); + if (!multi) break; + } + } + } + + if (matched.items.len == 0) { + if (upsert) { + const new_doc = try build_upsert_doc(ctx, reply, q, u_doc); + ctx.engine.insert(db_name, coll_name, new_doc, ctx.oid_gen) catch |err| switch (err) { + error.DuplicateKey => return duplicate_key_error(reply, db_name, coll_name, new_doc), + else => return err, + }; + const id = new_doc.get("_id") orelse bson.Value.null; + const u = try reply.arena_alloc().alloc(bson.Pair, 2); + u[0] = .{ .key = try reply.arena_alloc().dupe(u8, "index"), .value = .{ .int32 = @intCast(si) } }; + u[1] = .{ .key = try reply.arena_alloc().dupe(u8, "_id"), .value = try bson.copy_value(reply.arena_alloc(), id) }; + try upserted.append(reply.arena_alloc(), .{ .key = "u", .value = .{ .doc = u } }); + n_matched += 1; + } + continue; + } + + n_matched += @intCast(matched.items.len); + for (matched.items) |doc| { + // Work on a copy: the log write must precede any visible change, + // and a rejected update must not corrupt the stored document. + const copy = try clone_doc(ctx, reply, doc); + update.apply(copy, &.{ .arena = undefined, .pairs = u_doc }) catch |err| switch (err) { + error.ImmutableId, error.InvalidUpdate => return bad_value(reply, "bad update"), + else => return err, + }; + try ctx.engine.replace(db_name, coll_name, copy, ctx.oid_gen); + n_modified += 1; + } + } + + try reply.put("n", .{ .int32 = @intCast(n_matched) }); + try reply.put("nModified", .{ .int32 = @intCast(n_modified) }); + if (upserted.items.len > 0) { + const arr = try reply.arena_alloc().alloc(bson.Value, upserted.items.len); + for (upserted.items, 0..) |u, i| arr[i] = u.value; + try reply.put("upserted", .{ .array = arr }); + } + try reply.put_ok(); +} + +fn cmd_delete(ctx: *Context, msg: *wire.Message, reply: *wire.Reply) !void { + const db_name = msg.db_name() orelse return invalid_arg(reply, "delete requires $db"); + const coll_name = str_arg(msg.body.get("delete")) orelse return bad_value(reply, "delete requires a collection name"); + + var specs: []const bson.Document = &.{}; + for (msg.seqs) |seq| { + if (std.mem.eql(u8, seq.name, "deletes")) specs = seq.docs; + } + if (specs.len == 0) { + const d = msg.body.get("deletes") orelse return bad_value(reply, "delete requires deletes"); + const arr = switch (d) { + .array => |a| a, + else => return bad_value(reply, "deletes must be an array"), + }; + const parsed = try reply.arena_alloc().alloc(bson.Document, arr.len); + for (arr, 0..) |item, i| { + parsed[i] = switch (item) { + .doc => |pairs| bson.Document{ .arena = undefined, .pairs = pairs }, + else => return bad_value(reply, "deletes must be documents"), + }; + } + specs = parsed; + } + + var n_deleted: i64 = 0; + for (specs) |*spec| { + const q = doc_arg(spec.get("q")) orelse return bad_value(reply, "delete spec requires q"); + const limit: i64 = switch (spec.get("limit") orelse bson.Value{ .int32 = 1 }) { + .int32 => |i| i, + .int64 => |i| i, + else => 1, + }; + if (ctx.engine.get_collection(db_name, coll_name)) |coll| { + var to_remove: std.ArrayListUnmanaged([]const u8) = .empty; + defer { + for (to_remove.items) |k| ctx.gpa.free(k); + to_remove.deinit(ctx.gpa); + } + var it = coll.docs.iterator(); + while (it.next()) |entry| { + if (try query.matches(ctx.gpa, &.{ .arena = undefined, .pairs = q }, entry.value_ptr.*)) { + try to_remove.append(ctx.gpa, try ctx.gpa.dupe(u8, entry.key_ptr.*)); + if (limit == 1) break; + } + } + for (to_remove.items) |key| { + if (try ctx.engine.remove(db_name, coll_name, key)) n_deleted += 1; + } + } + } + + try reply.put("n", .{ .int32 = @intCast(n_deleted) }); + try reply.put_ok(); +} + +fn cmd_find_and_modify(ctx: *Context, msg: *wire.Message, reply: *wire.Reply) !void { + const db_name = msg.db_name() orelse return invalid_arg(reply, "findAndModify requires $db"); + const coll_name = str_arg(msg.body.get("findAndModify")) orelse return bad_value(reply, "findAndModify requires a collection name"); + + const q = doc_arg(msg.body.get("query")) orelse &.{}; + const sort_keys = try parse_sort_keys(reply, msg.body.get("sort")); + const remove = bool_arg(msg.body.get("remove")) orelse false; + const do_update = msg.body.get("update") != null; + const upsert = bool_arg(msg.body.get("upsert")) orelse false; + const ret_new = bool_arg(msg.body.get("new")) orelse false; + const proj_pairs = doc_arg(msg.body.get("fields")); + + if (remove and do_update) return bad_value(reply, "remove and update are mutually exclusive"); + if (!remove and !do_update) return bad_value(reply, "must specify update or remove"); + + var matched: std.ArrayListUnmanaged(*const bson.Document) = .empty; + defer matched.deinit(ctx.gpa); + if (ctx.engine.get_collection(db_name, coll_name)) |coll| { + var it = coll.docs.iterator(); + while (it.next()) |entry| { + if (try query.matches(ctx.gpa, &.{ .arena = undefined, .pairs = q }, entry.value_ptr.*)) { + try matched.append(ctx.gpa, entry.value_ptr.*); + } + } + } + if (sort_keys.len > 0) { + try query.sort_docs(ctx.gpa, reply.arena_alloc(), matched.items, sort_keys); + } + + var leo: std.ArrayListUnmanaged(bson.Pair) = .empty; + defer leo.deinit(reply.arena_alloc()); + + if (matched.items.len == 0 and do_update and upsert) { + const u_doc = doc_arg(msg.body.get("update")) orelse return bad_value(reply, "update must be a document"); + const new_doc = try build_upsert_doc(ctx, reply, q, u_doc); + ctx.engine.insert(db_name, coll_name, new_doc, ctx.oid_gen) catch |err| switch (err) { + error.DuplicateKey => return duplicate_key_error(reply, db_name, coll_name, new_doc), + else => return err, + }; + try leo.append(reply.arena_alloc(), .{ .key = "n", .value = .{ .int32 = 1 } }); + try leo.append(reply.arena_alloc(), .{ .key = "updatedExisting", .value = .{ .bool = false } }); + try leo.append(reply.arena_alloc(), .{ .key = "upserted", .value = try bson.copy_value(reply.arena_alloc(), new_doc.get("_id") orelse bson.Value.null) }); + if (ret_new) { + try reply.put("value", try project_doc(reply, new_doc, proj_pairs)); + } else { + try reply.put("value", .null); + } + try reply.put("lastErrorObject", .{ .doc = try leo.toOwnedSlice(reply.arena_alloc()) }); + try reply.put_ok(); + return; + } + + const target = if (matched.items.len > 0) matched.items[0] else null; + if (target != null and remove) { + const id_key = try bson.serialize_value(ctx.gpa, target.?.get("_id") orelse unreachable); + defer ctx.gpa.free(id_key); + try leo.append(reply.arena_alloc(), .{ .key = "n", .value = .{ .int32 = 1 } }); + try leo.append(reply.arena_alloc(), .{ .key = "updatedExisting", .value = .{ .bool = false } }); + try reply.put("value", try project_doc(reply, target.?, proj_pairs)); + _ = try ctx.engine.remove(db_name, coll_name, id_key); + try reply.put("lastErrorObject", .{ .doc = try leo.toOwnedSlice(reply.arena_alloc()) }); + try reply.put_ok(); + return; + } + + if (target != null and do_update) { + const u_doc = doc_arg(msg.body.get("update")) orelse return bad_value(reply, "update must be a document"); + const before = try bson.copy_pairs(reply.arena_alloc(), target.?.pairs); + const copy = try clone_doc(ctx, reply, target.?); + update.apply(copy, &.{ .arena = undefined, .pairs = u_doc }) catch |err| switch (err) { + error.ImmutableId, error.InvalidUpdate => return bad_value(reply, "bad update"), + else => return err, + }; + try ctx.engine.replace(db_name, coll_name, copy, ctx.oid_gen); + try leo.append(reply.arena_alloc(), .{ .key = "n", .value = .{ .int32 = 1 } }); + try leo.append(reply.arena_alloc(), .{ .key = "updatedExisting", .value = .{ .bool = true } }); + if (ret_new) { + try reply.put("value", try project_doc(reply, copy, proj_pairs)); + } else { + try reply.put("value", .{ .doc = before }); + } + try reply.put("lastErrorObject", .{ .doc = try leo.toOwnedSlice(reply.arena_alloc()) }); + try reply.put_ok(); + return; + } + + // No match and no upsert + try leo.append(reply.arena_alloc(), .{ .key = "n", .value = .{ .int32 = 0 } }); + try leo.append(reply.arena_alloc(), .{ .key = "updatedExisting", .value = .{ .bool = false } }); + try reply.put("value", .null); + try reply.put("lastErrorObject", .{ .doc = try leo.toOwnedSlice(reply.arena_alloc()) }); + try reply.put_ok(); +} + +fn cmd_count(ctx: *Context, msg: *wire.Message, reply: *wire.Reply) !void { + const db_name = msg.db_name() orelse return invalid_arg(reply, "count requires $db"); + const coll_name = str_arg(msg.body.get("count")) orelse return bad_value(reply, "count requires a collection name"); + const q = doc_arg(msg.body.get("query")) orelse &.{}; + + var n: i64 = 0; + if (ctx.engine.get_collection(db_name, coll_name)) |coll| { + var it = coll.docs.iterator(); + while (it.next()) |entry| { + if (try query.matches(ctx.gpa, &.{ .arena = undefined, .pairs = q }, entry.value_ptr.*)) n += 1; + } + } + try reply.put("n", .{ .int32 = @intCast(n) }); + try reply.put_ok(); +} + +fn cmd_aggregate(ctx: *Context, msg: *wire.Message, reply: *wire.Reply) !void { + const db_name = msg.db_name() orelse return invalid_arg(reply, "aggregate requires $db"); + const coll_name = str_arg(msg.body.get("aggregate")) orelse return bad_value(reply, "aggregate requires a collection name"); + + const pipeline_value = msg.body.get("pipeline") orelse return bad_value(reply, "aggregate requires pipeline"); + const stages = switch (pipeline_value) { + .array => |a| a, + else => return bad_value(reply, "pipeline must be an array"), + }; + + // The pipeline operates on a stream of documents; each stage transforms + // the current window [start, end) of `stream`, and $group replaces the + // stream entirely (so $sort/$limit after it apply to the groups). + var stream: std.ArrayListUnmanaged(*const bson.Document) = .empty; + defer stream.deinit(ctx.gpa); + if (ctx.engine.get_collection(db_name, coll_name)) |coll| { + var it = coll.docs.iterator(); + while (it.next()) |entry| try stream.append(ctx.gpa, entry.value_ptr.*); + } + + var start: usize = 0; + var end: usize = stream.items.len; + var count_stage: ?[]const u8 = null; + var proj_pairs: ?[]const bson.Pair = null; + + for (stages) |stage_v| { + const stage = switch (stage_v) { + .doc => |pairs| pairs, + else => return bad_value(reply, "pipeline stages must be documents"), + }; + if (stage.len == 0) continue; + const stage_name = stage[0].key; + if (std.mem.eql(u8, stage_name, "$match")) { + const filter = doc_arg(stage[0].value) orelse return bad_value(reply, "$match requires a document"); + var kept: std.ArrayListUnmanaged(*const bson.Document) = .empty; + defer kept.deinit(ctx.gpa); + for (stream.items[start..end]) |d| { + if (try query.matches(ctx.gpa, &.{ .arena = undefined, .pairs = filter }, d)) { + try kept.append(ctx.gpa, d); + } + } + stream.deinit(ctx.gpa); + stream = kept; + kept = .empty; + start = 0; + end = stream.items.len; + } else if (std.mem.eql(u8, stage_name, "$sort")) { + const keys = try parse_sort_keys(reply, stage[0].value); + if (keys.len > 0) { + try query.sort_docs(ctx.gpa, reply.arena_alloc(), stream.items[start..end], keys); + } + } else if (std.mem.eql(u8, stage_name, "$skip")) { + const n: usize = switch (stage[0].value) { + .int32 => |i| @intCast(@max(0, i)), + .int64 => |i| @intCast(@max(0, i)), + else => return bad_value(reply, "$skip requires a number"), + }; + start = @min(start + n, end); + } else if (std.mem.eql(u8, stage_name, "$limit")) { + const n: usize = switch (stage[0].value) { + .int32 => |i| @intCast(@max(0, i)), + .int64 => |i| @intCast(@max(0, i)), + else => return bad_value(reply, "$limit requires a number"), + }; + end = @min(end, start + n); + } else if (std.mem.eql(u8, stage_name, "$project")) { + proj_pairs = doc_arg(stage[0].value); + } else if (std.mem.eql(u8, stage_name, "$group")) { + const gp = doc_arg(stage[0].value) orelse return bad_value(reply, "$group requires a document"); + const grouped_opt = try run_group(ctx, reply, gp, stream.items[start..end]); + const grouped = grouped_opt orelse return; + // Group results replace the stream: later stages see groups. + stream.deinit(ctx.gpa); + stream = grouped; + start = 0; + end = stream.items.len; + } else if (std.mem.eql(u8, stage_name, "$count")) { + count_stage = switch (stage[0].value) { + .string => |s| s, + else => return bad_value(reply, "$count requires a string"), + }; + } else { + const msg_text = try std.fmt.allocPrint(reply.arena_alloc(), "Unrecognized pipeline stage name: '{s}'", .{stage_name}); + return reply.put_error(@intFromEnum(ErrorCode.invalid_pipeline_operator), "InvalidPipelineOperator", msg_text); + } + } + + const slice = stream.items[start..end]; + + if (count_stage) |name| { + const c = try reply.arena_alloc().alloc(bson.Pair, 1); + c[0] = .{ .key = try reply.arena_alloc().dupe(u8, name), .value = .{ .int32 = @intCast(slice.len) } }; + const values = try reply.arena_alloc().alloc(bson.Value, 1); + values[0] = .{ .doc = c }; + try reply.put("cursor", .{ .doc = try cursor_doc(reply, 0, try format_namespace(reply, db_name, coll_name), values) }); + } else { + try emit_docs(reply, db_name, coll_name, proj_pairs, slice); + } + try reply.put_ok(); +} + +/// Minimal $group: supports `_id` of null/literal/"$field" and `$sum` +/// accumulators (constant or "$field"). +fn run_group(ctx: *Context, reply: *wire.Reply, group_pairs: []const bson.Pair, docs: []const *const bson.Document) !?std.ArrayListUnmanaged(*const bson.Document) { + const arena = reply.arena_alloc(); + const id_expr = bson.get_pair(group_pairs, "_id") orelse { + try bad_value(reply, "$group requires _id"); + return null; + }; + + var accs: std.ArrayListUnmanaged([]const bson.Pair) = .empty; + defer accs.deinit(arena); + for (group_pairs) |p| { + if (std.mem.eql(u8, p.key, "_id")) continue; + try accs.append(arena, &.{.{ .key = p.key, .value = p.value }}); + } + + const Group = struct { + id_value: bson.Value, + sums: []f64, + }; + var groups: std.StringHashMapUnmanaged(Group) = .empty; + defer groups.deinit(ctx.gpa); + // StringHashMapUnmanaged does not copy keys; keep them alive until done. + var keys_owned: std.ArrayListUnmanaged([]u8) = .empty; + defer { + for (keys_owned.items) |k| ctx.gpa.free(k); + keys_owned.deinit(ctx.gpa); + } + + for (docs) |doc| { + const id_value: bson.Value = switch (id_expr) { + .null => .null, + else => switch (id_expr) { + .string => |s| if (s.len > 0 and s[0] == '$') query_path_value(doc, s[1..]) orelse .null else id_expr, + else => id_expr, + }, + }; + var id_key_buf: std.ArrayListUnmanaged(u8) = .empty; + defer id_key_buf.deinit(ctx.gpa); + try bson.write_value(id_value, ctx.gpa, &id_key_buf); + const key = try ctx.gpa.dupe(u8, id_key_buf.items); + try keys_owned.append(ctx.gpa, key); + + const gop = try groups.getOrPut(ctx.gpa, key); + if (!gop.found_existing) { + const sums = try ctx.gpa.alloc(f64, accs.items.len); + @memset(sums, 0); + gop.value_ptr.* = .{ .id_value = id_value, .sums = sums }; + } + for (accs.items, 0..) |acc, i| { + var expr = acc[0].value; + // Unwrap {$sum: } accumulator documents. + if (expr == .doc) { + if (bson.get_pair(expr.doc, "$sum")) |inner| { + expr = inner; + } else continue; + } + const term: f64 = switch (expr) { + .int32 => |n| @floatFromInt(n), + .int64 => |n| @floatFromInt(n), + .double => |n| n, + .string => |s| if (s.len > 0 and s[0] == '$') + switch (query_path_value(doc, s[1..]) orelse .null) { + .int32 => |n| @floatFromInt(n), + .int64 => |n| @floatFromInt(n), + .double => |n| n, + else => 0, + } + else + 0, + else => 0, + }; + gop.value_ptr.sums[i] += term; + } + } + + var out: std.ArrayListUnmanaged(*const bson.Document) = .empty; + errdefer out.deinit(ctx.gpa); + var it = groups.iterator(); + // free sum arrays + defer { + var git = groups.iterator(); + while (git.next()) |e| ctx.gpa.free(e.value_ptr.sums); + } + while (it.next()) |entry| { + const npairs = 1 + accs.items.len; + const pairs = try arena.alloc(bson.Pair, npairs); + pairs[0] = .{ .key = try arena.dupe(u8, "_id"), .value = try bson.copy_value(arena, entry.value_ptr.id_value) }; + for (accs.items, 0..) |acc, i| { + const sum: f64 = entry.value_ptr.sums[i]; + const sum_value: bson.Value = if (sum == @floor(sum) and sum <= 2_147_483_647 and sum >= -2_147_483_648) + .{ .int32 = @intFromFloat(sum) } + else + .{ .double = sum }; + pairs[1 + i] = .{ .key = try arena.dupe(u8, acc[0].key), .value = sum_value }; + } + const doc = try arena.create(bson.Document); + doc.* = bson.Document{ .arena = undefined, .pairs = pairs }; + try out.append(ctx.gpa, doc); + } + return out; +} + +/// Resolve a simple "$field" path expression inside a document. +fn query_path_value(doc: *const bson.Document, path: []const u8) ?bson.Value { + var cur: bson.Value = undefined; + var it = std.mem.splitScalar(u8, path, '.'); + const first = it.next() orelse return null; + cur = bson.get_pair(doc.pairs, first) orelse return null; + while (it.next()) |seg| { + cur = switch (cur) { + .doc => |pairs| bson.get_pair(pairs, seg) orelse return null, + else => return null, + }; + } + return cur; +} + +fn cmd_get_more(reply: *wire.Reply) !void { + try reply.put("cursor", .{ .doc = &.{ + .{ .key = "id", .value = .{ .int64 = 0 } }, + .{ .key = "ns", .value = .{ .string = "test.$cmd" } }, + .{ .key = "nextBatch", .value = .{ .array = &.{} } }, + } }); + try reply.put_ok(); +} + +fn cmd_kill_cursors(reply: *wire.Reply) !void { + try reply.put("cursorsKilled", .{ .array = &.{} }); + try reply.put("cursorsNotFound", .{ .array = &.{} }); + try reply.put("cursorsAlive", .{ .array = &.{} }); + try reply.put_ok(); +} + +// --------------------------------------------------------------------------- +// Helpers +// --------------------------------------------------------------------------- + +/// Deep-copy a stored document into the reply arena so updates can be +/// applied off the live doc. +fn clone_doc(ctx: *Context, reply: *wire.Reply, doc: *const bson.Document) !*bson.Document { + var buf: std.ArrayListUnmanaged(u8) = .empty; + defer buf.deinit(ctx.gpa); + try doc.to_bytes(ctx.gpa, &buf); + const owned = try reply.arena_alloc().create(bson.Document); + owned.* = try bson.Document.parse(reply.arena_alloc(), buf.items); + return owned; +} + +/// Build the document for an upsert: equality fields from the filter, then +/// the update operators applied. Owned by the reply arena. +fn build_upsert_doc(ctx: *Context, reply: *wire.Reply, q: []const bson.Pair, u_doc: []const bson.Pair) !*bson.Document { + _ = ctx; + const arena = reply.arena_alloc(); + var pairs: std.ArrayListUnmanaged(bson.Pair) = .empty; + defer pairs.deinit(arena); + for (q) |p| { + const is_operator = p.key.len > 0 and p.key[0] == '$'; + const is_embedded_operator = p.value == .doc and doc_all_operators(p.value.doc); + if (!is_operator and !is_embedded_operator) { + try pairs.append(arena, .{ .key = try arena.dupe(u8, p.key), .value = try bson.copy_value(arena, p.value) }); + } + } + const owned = try arena.create(bson.Document); + owned.* = bson.Document{ .arena = std.heap.ArenaAllocator.init(arena), .pairs = try pairs.toOwnedSlice(arena) }; + // Apply update operators to build the final doc; _id handled by insert. + update.apply(owned, &.{ .arena = undefined, .pairs = u_doc }) catch |err| switch (err) { + error.ImmutableId, error.InvalidUpdate => return error.InvalidUpdate, + else => return err, + }; + _ = ctx; + return owned; +} + +fn doc_all_operators(pairs: []const bson.Pair) bool { + if (pairs.len == 0) return false; + for (pairs) |p| { + if (p.key.len == 0 or p.key[0] != '$') return false; + } + return true; +} + +fn parse_sort_keys(reply: *wire.Reply, value: ?bson.Value) ![]const query.SortKey { + const pairs = doc_arg(value) orelse return &.{}; + const out = try reply.arena_alloc().alloc(query.SortKey, pairs.len); + for (pairs, 0..) |p, i| { + const descending = switch (p.value) { + .int32 => |n| n < 0, + .int64 => |n| n < 0, + .double => |n| n < 0, + .string => |s| std.mem.eql(u8, s, "desc"), + else => false, + }; + out[i] = .{ + .path = try reply.arena_alloc().dupe(u8, p.key), + .descending = descending, + }; + } + return out; +} + +fn doc_arg(v: ?bson.Value) ?[]const bson.Pair { + return switch (v orelse return null) { + .doc => |pairs| pairs, + else => null, + }; +} + +fn str_arg(v: ?bson.Value) ?[]const u8 { + return switch (v orelse return null) { + .string => |s| s, + else => null, + }; +} + +fn bool_arg(v: ?bson.Value) ?bool { + return switch (v orelse return null) { + .bool => |b| b, + .int32 => |i| i != 0, + else => null, + }; +} + +fn int_arg(v: ?bson.Value) ?u64 { + return switch (v orelse return null) { + .int32 => |i| if (i < 0) null else @intCast(i), + .int64 => |i| if (i < 0) null else @intCast(i), + .double => |d| if (d < 0) null else @intFromFloat(d), + else => null, + }; +} + +fn invalid_arg(reply: *wire.Reply, msg: []const u8) !void { + return reply.put_error(@intFromEnum(ErrorCode.invalid_argument), "InvalidArgument", msg); +} + +fn bad_value(reply: *wire.Reply, msg: []const u8) !void { + return reply.put_error(@intFromEnum(ErrorCode.bad_value), "BadValue", msg); +} + +fn duplicate_key_error(reply: *wire.Reply, db_name: []const u8, coll_name: []const u8, doc: *const bson.Document) !void { + const id = doc.get("_id") orelse bson.Value.null; + const key = try serialize_value_compact(reply, id); + const msg_text = try std.fmt.allocPrint(reply.arena_alloc(), "E11000 duplicate key error collection: {s}.{s} index: _id_ dup key: {s}", .{ db_name, coll_name, key }); + return reply.put_error(@intFromEnum(ErrorCode.duplicate_key), "DuplicateKey", msg_text); +} + +/// Compact extended-JSON-ish rendering of a value for error messages. +fn serialize_value_compact(reply: *wire.Reply, v: bson.Value) ![]const u8 { + const arena = reply.arena_alloc(); + return switch (v) { + .int32 => |i| try std.fmt.allocPrint(arena, "{d}", .{i}), + .int64 => |i| try std.fmt.allocPrint(arena, "{d}", .{i}), + .double => |d| try std.fmt.allocPrint(arena, "{d}", .{d}), + .string => |s| try std.fmt.allocPrint(arena, "'{s}'", .{s}), + .bool => |b| try std.fmt.allocPrint(arena, "{}", .{b}), + .object_id => |oid| blk: { + const hex = std.fmt.bytesToHex(oid[0..], .lower); + break :blk try std.fmt.allocPrint(arena, "ObjectId('{s}')", .{hex}); + }, + else => "{ ... }", + }; +} + +/// Build a { id: 0, ns: "...", firstBatch: [...] } cursor document. +pub fn cursor_doc(reply: *wire.Reply, cursor_id: i64, ns: []const u8, docs: []const bson.Value) ![]const bson.Pair { + const c = try reply.arena_alloc().alloc(bson.Pair, 3); + c[0] = .{ .key = try reply.arena_alloc().dupe(u8, "id"), .value = .{ .int64 = cursor_id } }; + c[1] = .{ .key = try reply.arena_alloc().dupe(u8, "ns"), .value = .{ .string = try reply.arena_alloc().dupe(u8, ns) } }; + c[2] = .{ .key = try reply.arena_alloc().dupe(u8, "firstBatch"), .value = .{ .array = docs } }; + return c; +} + +fn format_namespace(reply: *wire.Reply, db_name: []const u8, coll_name: []const u8) ![]const u8 { + return std.fmt.allocPrint(reply.arena_alloc(), "{s}.{s}", .{ db_name, coll_name }); +} + +fn int_array(reply: *wire.Reply, values: []const i32) ![]const bson.Value { + const out = try reply.arena_alloc().alloc(bson.Value, values.len); + for (values, 0..) |v, i| out[i] = .{ .int32 = v }; + return out; +} + +fn str_array(reply: *wire.Reply, values: []const []const u8) ![]const bson.Value { + const out = try reply.arena_alloc().alloc(bson.Value, values.len); + for (values, 0..) |v, i| out[i] = .{ .string = try reply.arena_alloc().dupe(u8, v) }; + return out; +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +const testing = std.testing; + +test "ping and hello replies parse" { + var threaded: std.Io.Threaded = .init_single_threaded; + defer threaded.deinit(); + const io = threaded.io(); + var gen = bson.ObjectIdGen.init(io); + var tmp = std.testing.tmpDir(.{}); + defer tmp.cleanup(); + const path = try std.fmt.allocPrint(testing.allocator, ".zig-cache/tmp/{s}/cmd.log", .{tmp.sub_path}); + defer testing.allocator.free(path); + var engine = try db.Engine.open(testing.allocator, io, path); + defer engine.deinit(); + var ctx = Context{ + .gpa = testing.allocator, + .io = io, + .oid_gen = &gen, + .connection_id = 1, + .client_desc = "127.0.0.1:0", + .engine = &engine, + .server_start = std.Io.Timestamp.now(io, .real), + }; + + var reply = wire.Reply.init(testing.allocator); + defer reply.deinit(); + var msg = try parse_fake_msg("ping", .null, &.{}); + defer msg.deinit(); + try dispatch(&ctx, &msg, &reply); + try testing.expectEqual(@as(f64, 1.0), reply.pairs.items[0].value.double); + + var reply2 = wire.Reply.init(testing.allocator); + defer reply2.deinit(); + var msg2 = try parse_fake_msg("hello", .null, &.{}); + defer msg2.deinit(); + try dispatch(&ctx, &msg2, &reply2); + const ok = bson.get_pair(reply2.pairs.items, "ok").?; + try testing.expectEqual(@as(f64, 1.0), ok.double); + const primary = bson.get_pair(reply2.pairs.items, "isWritablePrimary").?; + try testing.expect(primary.bool); + try testing.expectEqual(@as(i32, 8), bson.get_pair(reply2.pairs.items, "maxWireVersion").?.int32); +} + +test "unknown command gives CommandNotFound" { + var threaded: std.Io.Threaded = .init_single_threaded; + defer threaded.deinit(); + const io = threaded.io(); + var gen = bson.ObjectIdGen.init(io); + var tmp = std.testing.tmpDir(.{}); + defer tmp.cleanup(); + const path = try std.fmt.allocPrint(testing.allocator, ".zig-cache/tmp/{s}/cmd.log", .{tmp.sub_path}); + defer testing.allocator.free(path); + var engine = try db.Engine.open(testing.allocator, io, path); + defer engine.deinit(); + var ctx = Context{ + .gpa = testing.allocator, + .io = io, + .oid_gen = &gen, + .connection_id = 1, + .client_desc = "127.0.0.1:0", + .engine = &engine, + .server_start = std.Io.Timestamp.now(io, .real), + }; + var reply = wire.Reply.init(testing.allocator); + defer reply.deinit(); + var msg = try parse_fake_msg("nonsenseCmd", .null, &.{}); + defer msg.deinit(); + try dispatch(&ctx, &msg, &reply); + try testing.expectEqual(@as(f64, 0.0), bson.get_pair(reply.pairs.items, "ok").?.double); + try testing.expectEqual(@as(i32, 59), bson.get_pair(reply.pairs.items, "code").?.int32); + try testing.expectEqualStrings("CommandNotFound", bson.get_pair(reply.pairs.items, "codeName").?.string); +} + +test "getParameter responds for featureCompatibilityVersion" { + var threaded: std.Io.Threaded = .init_single_threaded; + defer threaded.deinit(); + const io = threaded.io(); + var gen = bson.ObjectIdGen.init(io); + var tmp = std.testing.tmpDir(.{}); + defer tmp.cleanup(); + const path = try std.fmt.allocPrint(testing.allocator, ".zig-cache/tmp/{s}/cmd.log", .{tmp.sub_path}); + defer testing.allocator.free(path); + var engine = try db.Engine.open(testing.allocator, io, path); + defer engine.deinit(); + var ctx = Context{ + .gpa = testing.allocator, + .io = io, + .oid_gen = &gen, + .connection_id = 1, + .client_desc = "127.0.0.1:0", + .engine = &engine, + .server_start = std.Io.Timestamp.now(io, .real), + }; + var reply = wire.Reply.init(testing.allocator); + defer reply.deinit(); + var fcv_doc = try testing.allocator.alloc(bson.Pair, 1); + defer testing.allocator.free(fcv_doc); + fcv_doc[0] = .{ .key = "featureCompatibilityVersion", .value = .{ .int32 = 1 } }; + var msg = try parse_fake_msg("getParameter", .{ .doc = fcv_doc }, &.{}); + defer msg.deinit(); + try dispatch(&ctx, &msg, &reply); + const fcv = bson.get_pair(reply.pairs.items, "featureCompatibilityVersion").?; + try testing.expectEqualStrings("4.4", fcv.doc[0].value.string); +} + +test "insert counts only successful inserts" { + var threaded: std.Io.Threaded = .init_single_threaded; + defer threaded.deinit(); + const io = threaded.io(); + var gen = bson.ObjectIdGen.init(io); + var tmp = std.testing.tmpDir(.{}); + defer tmp.cleanup(); + const path = try std.fmt.allocPrint(testing.allocator, ".zig-cache/tmp/{s}/cmd.log", .{tmp.sub_path}); + defer testing.allocator.free(path); + var engine = try db.Engine.open(testing.allocator, io, path); + defer engine.deinit(); + var ctx = Context{ + .gpa = testing.allocator, + .io = io, + .oid_gen = &gen, + .connection_id = 1, + .client_desc = "127.0.0.1:0", + .engine = &engine, + .server_start = std.Io.Timestamp.now(io, .real), + }; + + // documents: [{_id:1}, {_id:1} (dup), {_id:2}] → n: 2 + one writeError. + const docs = [_]bson.Value{ + .{ .doc = &.{.{ .key = "_id", .value = .{ .int32 = 1 } }} }, + .{ .doc = &.{.{ .key = "_id", .value = .{ .int32 = 1 } }} }, + .{ .doc = &.{.{ .key = "_id", .value = .{ .int32 = 2 } }} }, + }; + var reply = wire.Reply.init(testing.allocator); + defer reply.deinit(); + var msg = try parse_fake_msg("insert", .{ .string = "users" }, &.{ + .{ .key = "documents", .value = .{ .array = &docs } }, + }); + defer msg.deinit(); + try dispatch(&ctx, &msg, &reply); + try testing.expectEqual(@as(i64, 2), bson.get_pair(reply.pairs.items, "n").?.int32); + const errors = bson.get_pair(reply.pairs.items, "writeErrors").?; + try testing.expectEqual(@as(usize, 1), errors.array.len); + const first = errors.array[0].doc[0]; + try testing.expectEqualStrings("index", first.key); + try testing.expectEqual(@as(i64, 1), first.value.int32); +} + +fn parse_fake_msg(name: []const u8, value: bson.Value, extra: []const bson.Pair) !wire.Message { + var out: std.ArrayListUnmanaged(u8) = .empty; + defer out.deinit(testing.allocator); + const pairs = try testing.allocator.alloc(bson.Pair, 2 + extra.len); + defer testing.allocator.free(pairs); + pairs[0] = .{ .key = name, .value = value }; + pairs[1] = .{ .key = "$db", .value = .{ .string = "test" } }; + @memcpy(pairs[2..], extra); + try bson.write_doc(pairs, testing.allocator, &out); + var msg: std.ArrayListUnmanaged(u8) = .empty; + defer msg.deinit(testing.allocator); + try msg.appendSlice(testing.allocator, &[_]u8{ 0, 0, 0, 0, 1, 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, out.items); + std.mem.writeInt(u32, msg.items[0..4], @intCast(msg.items.len), .little); + std.mem.writeInt(i32, msg.items[12..16], wire.op_code_msg, .little); + return wire.Message.parse(testing.allocator, msg.items); +} diff --git a/src/db.zig b/src/db.zig new file mode 100644 index 0000000..7281dc6 --- /dev/null +++ b/src/db.zig @@ -0,0 +1,538 @@ +//! In-memory database engine backed by the append-only log. Maps +//! db -> collection -> _id(serialized) -> owned Document. All mutations are +//! logged and synced before they become visible in memory, so a crash never +//! loses a committed write. Callers must hold `mutex` around a command. + +const std = @import("std"); +const bson = @import("bson.zig"); +const storage = @import("storage.zig"); + +pub const Collection = struct { + docs: std.StringHashMapUnmanaged(*bson.Document), + + fn init() Collection { + return .{ .docs = .empty }; + } +}; + +pub const Db = struct { + collections: std.StringHashMapUnmanaged(Collection), +}; + +pub const Engine = struct { + gpa: std.mem.Allocator, + io: std.Io, + mutex: std.Io.Mutex, + log: storage.Log, + dbs: std.StringHashMapUnmanaged(Db), + seq: u64, + compact_threshold: u64, + + pub fn open(gpa: std.mem.Allocator, io: std.Io, path: []const u8) !Engine { + var engine = Engine{ + .gpa = gpa, + .io = io, + .mutex = std.Io.Mutex.init, + .log = try storage.Log.open(gpa, io, path), + .dbs = .empty, + .seq = 0, + .compact_threshold = 16 * 1024 * 1024, + }; + errdefer { + engine.log.close(); + engine.dbs.deinit(gpa); + } + + try engine.log.replay(&engine, apply_record); + return engine; + } + + pub fn deinit(self: *Engine) void { + var db_it = self.dbs.iterator(); + while (db_it.next()) |db_entry| { + var coll_it = db_entry.value_ptr.collections.iterator(); + while (coll_it.next()) |coll_entry| { + var doc_it = coll_entry.value_ptr.docs.iterator(); + while (doc_it.next()) |doc_entry| { + doc_entry.value_ptr.*.deinit(); + self.gpa.destroy(doc_entry.value_ptr.*); + self.gpa.free(doc_entry.key_ptr.*); + } + coll_entry.value_ptr.docs.deinit(self.gpa); + self.gpa.free(coll_entry.key_ptr.*); + } + db_entry.value_ptr.collections.deinit(self.gpa); + self.gpa.free(db_entry.key_ptr.*); + } + self.dbs.deinit(self.gpa); + self.log.close(); + } + + // -- commands (callers must hold the mutex) ------------------------------ + + pub fn lock(self: *Engine) !void { + try self.mutex.lock(self.io); + } + + pub fn unlock(self: *Engine) void { + self.mutex.unlock(self.io); + } + + /// Insert a document. Fails with error.DuplicateKey if the _id exists. + /// Generates an ObjectId _id when absent. + pub fn insert(self: *Engine, db_name: []const u8, coll_name: []const u8, doc: *const bson.Document, oid_gen: *bson.ObjectIdGen) !void { + const coll = try self.get_or_create_collection(db_name, coll_name); + const owned = try self.own_with_id(doc, oid_gen); + errdefer { + owned.deinit(); + self.gpa.destroy(owned); + } + + const id_value = owned.get("_id") orelse unreachable; + const id_key = try bson.serialize_value(self.gpa, id_value); + defer self.gpa.free(id_key); + + if (coll.docs.contains(id_key)) return error.DuplicateKey; + + const doc_bytes = try serialize_doc(self.gpa, owned); + defer self.gpa.free(doc_bytes); + self.seq += 1; + try self.log.append_upsert(db_name, coll_name, doc_bytes, self.seq); + + const key_owned = try self.gpa.dupe(u8, id_key); + try coll.docs.put(self.gpa, key_owned, owned); + try self.maybe_compact(); + } + + /// Insert or replace a document by _id (upsert without existence check). + pub fn replace(self: *Engine, db_name: []const u8, coll_name: []const u8, doc: *const bson.Document, oid_gen: *bson.ObjectIdGen) !void { + const coll = try self.get_or_create_collection(db_name, coll_name); + const owned = try self.own_with_id(doc, oid_gen); + errdefer { + owned.deinit(); + self.gpa.destroy(owned); + } + + const id_value = owned.get("_id") orelse unreachable; + const id_key = try bson.serialize_value(self.gpa, id_value); + defer self.gpa.free(id_key); + + const doc_bytes = try serialize_doc(self.gpa, owned); + defer self.gpa.free(doc_bytes); + self.seq += 1; + try self.log.append_upsert(db_name, coll_name, doc_bytes, self.seq); + + if (coll.docs.fetchRemove(id_key)) |old| { + old.value.*.deinit(); + self.gpa.destroy(old.value); + self.gpa.free(old.key); + } + const key_owned = try self.gpa.dupe(u8, id_key); + try coll.docs.put(self.gpa, key_owned, owned); + try self.maybe_compact(); + } + + /// Remove a document by _id. Returns true if it existed. + pub fn remove(self: *Engine, db_name: []const u8, coll_name: []const u8, id_key: []const u8) !bool { + const db = self.dbs.get(db_name) orelse return false; + const coll = db.collections.getPtr(coll_name) orelse return false; + const doc = coll.docs.get(id_key) orelse return false; + + // Log (and sync) the delete before removing it from memory, so the + // log always describes at least as much as the in-memory state. + const doc_bytes = try serialize_doc(self.gpa, doc); + defer self.gpa.free(doc_bytes); + self.seq += 1; + try self.log.append_delete(db_name, coll_name, doc_bytes, self.seq); + + const removed = coll.docs.fetchRemove(id_key) orelse unreachable; + removed.value.*.deinit(); + self.gpa.destroy(removed.value); + self.gpa.free(removed.key); + return true; + } + + pub fn get_collection(self: *Engine, db_name: []const u8, coll_name: []const u8) ?*Collection { + const db = self.dbs.get(db_name) orelse return null; + return db.collections.getPtr(coll_name); + } + + pub fn get_doc(self: *Engine, db_name: []const u8, coll_name: []const u8, id_key: []const u8) ?*const bson.Document { + const coll = self.get_collection(db_name, coll_name) orelse return null; + return coll.docs.get(id_key); + } + + pub fn drop_collection(self: *Engine, db_name: []const u8, coll_name: []const u8) !bool { + const db = self.dbs.getPtr(db_name) orelse return false; + var removed = db.collections.fetchRemove(coll_name) orelse return false; + var doc_it = removed.value.docs.iterator(); + while (doc_it.next()) |doc_entry| { + doc_entry.value_ptr.*.deinit(); + self.gpa.destroy(doc_entry.value_ptr.*); + self.gpa.free(doc_entry.key_ptr.*); + } + removed.value.docs.deinit(self.gpa); + self.gpa.free(removed.key); + return true; + } + + pub fn drop_database(self: *Engine, db_name: []const u8) !bool { + var removed = self.dbs.fetchRemove(db_name) orelse return false; + var coll_it = removed.value.collections.iterator(); + while (coll_it.next()) |coll_entry| { + var docs_it = coll_entry.value_ptr.docs.iterator(); + while (docs_it.next()) |doc_entry| { + doc_entry.value_ptr.*.deinit(); + self.gpa.destroy(doc_entry.value_ptr.*); + self.gpa.free(doc_entry.key_ptr.*); + } + coll_entry.value_ptr.docs.deinit(self.gpa); + self.gpa.free(coll_entry.key_ptr.*); + } + removed.value.collections.deinit(self.gpa); + self.gpa.free(removed.key); + return true; + } + + pub fn database_names(self: *Engine, out: *std.ArrayListUnmanaged([]const u8)) !void { + var it = self.dbs.iterator(); + while (it.next()) |entry| try out.append(self.gpa, entry.key_ptr.*); + } + + pub fn collection_names(self: *Engine, db_name: []const u8, out: *std.ArrayListUnmanaged([]const u8)) !void { + const db = self.dbs.get(db_name) orelse return; + var it = db.collections.iterator(); + while (it.next()) |entry| try out.append(self.gpa, entry.key_ptr.*); + } + + // -- internals ----------------------------------------------------------- + + pub fn get_or_create_collection(self: *Engine, db_name: []const u8, coll_name: []const u8) !*Collection { + const db = self.dbs.getPtr(db_name) orelse { + const db_key = try self.gpa.dupe(u8, db_name); + errdefer self.gpa.free(db_key); + try self.dbs.put(self.gpa, db_key, .{ .collections = .empty }); + return self.get_or_create_collection(db_name, coll_name); + }; + if (db.collections.getPtr(coll_name)) |coll| return coll; + const coll_key = try self.gpa.dupe(u8, coll_name); + errdefer self.gpa.free(coll_key); + try db.collections.put(self.gpa, coll_key, Collection.init()); + return db.collections.getPtr(coll_name) orelse unreachable; + } + + /// Deep-copy a document into engine-owned storage, prepending a + /// generated ObjectId `_id` when absent. + fn own_with_id(self: *Engine, doc: *const bson.Document, oid_gen: *bson.ObjectIdGen) !*bson.Document { + var pairs: std.ArrayListUnmanaged(bson.Pair) = .empty; + defer pairs.deinit(self.gpa); + if (doc.get("_id") == null) { + const oid = oid_gen.new(self.io); + try pairs.append(self.gpa, .{ .key = "_id", .value = .{ .object_id = oid } }); + } + try pairs.appendSlice(self.gpa, doc.pairs); + + var out: std.ArrayListUnmanaged(u8) = .empty; + defer out.deinit(self.gpa); + try bson.write_doc(pairs.items, self.gpa, &out); + const owned = try self.gpa.create(bson.Document); + errdefer self.gpa.destroy(owned); + owned.* = try bson.Document.parse(self.gpa, out.items); + return owned; + } + + fn maybe_compact(self: *Engine) !void { + if (self.log.log_bytes < self.compact_threshold) return; + try self.compact(); + } + + /// Rewrite the log with only live documents, atomically swapping the file. + /// Callers must hold the mutex. + pub fn compact(self: *Engine) !void { + const tmp_path = try std.fmt.allocPrint(self.gpa, "{s}.tmp", .{self.log.path}); + defer self.gpa.free(tmp_path); + std.Io.Dir.cwd().deleteFile(self.io, tmp_path) catch {}; + var new_log = try storage.Log.open(self.gpa, self.io, tmp_path); + defer new_log.close(); + + var db_it = self.dbs.iterator(); + while (db_it.next()) |db_entry| { + var coll_it = db_entry.value_ptr.collections.iterator(); + while (coll_it.next()) |coll_entry| { + var doc_it = coll_entry.value_ptr.docs.iterator(); + while (doc_it.next()) |doc_entry| { + const doc_bytes = try serialize_doc(self.gpa, doc_entry.value_ptr.*); + defer self.gpa.free(doc_bytes); + try new_log.append_upsert(db_entry.key_ptr.*, coll_entry.key_ptr.*, doc_bytes, self.seq); + } + } + } + const new_end_pos = new_log.end_pos; + + try std.Io.Dir.renameAbsolute(tmp_path, self.log.path, self.io); + // Persist the rename: fsync the parent directory so the new + // directory entry survives a power loss right after compaction. + const parent = parent_dir(self.log.path); + var dir_file = try std.Io.Dir.cwd().openFile(self.io, parent, .{ .mode = .read_only, .allow_directory = true }); + defer dir_file.close(self.io); + try dir_file.sync(self.io); + + const old_path = try self.gpa.dupe(u8, self.log.path); + self.log.close(); + self.log = try storage.Log.open(self.gpa, self.io, old_path); + // Log.open starts at end_pos 0 and does not replay; continue appending + // where the compacted file actually ends. + self.log.end_pos = new_end_pos; + self.gpa.free(old_path); + } +}; + +fn parent_dir(path: []const u8) []const u8 { + const last = std.mem.lastIndexOfScalar(u8, path, '/') orelse return "."; + if (last == 0) return "/"; + return path[0..last]; +} + +fn serialize_doc(gpa: std.mem.Allocator, doc: *const bson.Document) ![]u8 { + var out: std.ArrayListUnmanaged(u8) = .empty; + errdefer out.deinit(gpa); + try doc.to_bytes(gpa, &out); + return out.toOwnedSlice(gpa); +} + +fn apply_record(ctx: *anyopaque, record: storage.Record, doc: *bson.Document) anyerror!void { + const self: *Engine = @ptrCast(@alignCast(ctx)); + var stored = false; + defer if (!stored) { + doc.deinit(); + self.gpa.destroy(doc); + }; + const id_value = doc.get("_id") orelse { + std.debug.print("mongo-light: log record without _id, skipping\n", .{}); + return; + }; + const id_key = try bson.serialize_value(self.gpa, id_value); + defer self.gpa.free(id_key); + + const coll = self.get_or_create_collection(record.db, record.coll) catch return; + + switch (record.type) { + storage.record_type_upsert => { + if (coll.docs.fetchRemove(id_key)) |old| { + old.value.*.deinit(); + self.gpa.destroy(old.value); + self.gpa.free(old.key); + } + const key_owned = try self.gpa.dupe(u8, id_key); + try coll.docs.put(self.gpa, key_owned, doc); + stored = true; + }, + storage.record_type_delete => { + if (coll.docs.fetchRemove(id_key)) |old| { + old.value.*.deinit(); + self.gpa.destroy(old.value); + self.gpa.free(old.key); + } + }, + else => {}, + } +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +const testing = std.testing; + +const TmpLog = struct { + tmp: std.testing.TmpDir, + path: []u8, + + fn init(gpa: std.mem.Allocator) !TmpLog { + const tmp = std.testing.tmpDir(.{}); + const path = try std.fmt.allocPrint(gpa, ".zig-cache/tmp/{s}/test.log", .{tmp.sub_path}); + return .{ .tmp = tmp, .path = path }; + } + + fn deinit(self: *TmpLog, gpa: std.mem.Allocator) void { + self.tmp.cleanup(); + gpa.free(self.path); + } +}; + +fn test_env(threaded: *std.Io.Threaded) struct { io: std.Io, gen: bson.ObjectIdGen } { + const io = threaded.io(); + const gen = bson.ObjectIdGen.init(io); + return .{ .io = io, .gen = gen }; +} + +fn make_doc(gpa: std.mem.Allocator, id: i32, name: []const u8) !bson.Document { + var arena = std.heap.ArenaAllocator.init(gpa); + errdefer arena.deinit(); + const pairs = try arena.allocator().alloc(bson.Pair, 2); + pairs[0] = .{ .key = try arena.allocator().dupe(u8, "_id"), .value = .{ .int32 = id } }; + pairs[1] = .{ .key = try arena.allocator().dupe(u8, "name"), .value = .{ .string = try arena.allocator().dupe(u8, name) } }; + return .{ .arena = arena, .pairs = pairs }; +} + +test "insert, query, remove" { + var threaded: std.Io.Threaded = .init_single_threaded; + defer threaded.deinit(); + var env = test_env(&threaded); + const io = env.io; + const gpa = testing.allocator; + + var tmp = try TmpLog.init(gpa); + defer tmp.deinit(gpa); + var engine = try Engine.open(gpa, io, tmp.path); + defer engine.deinit(); + + var d1 = try make_doc(gpa, 1, "alice"); + defer d1.deinit(); + var d2 = try make_doc(gpa, 2, "bob"); + defer d2.deinit(); + + try engine.lock(); + try engine.insert("app", "users", &d1, &env.gen); + try engine.insert("app", "users", &d2, &env.gen); + engine.unlock(); + + // duplicate key + var d3 = try make_doc(gpa, 1, "alice2"); + defer d3.deinit(); + try engine.lock(); + try testing.expectError(error.DuplicateKey, engine.insert("app", "users", &d3, &env.gen)); + engine.unlock(); + + // find by id + const id_key = try bson.serialize_value(gpa, bson.Value{ .int32 = 2 }); + defer gpa.free(id_key); + try engine.lock(); + const found = engine.get_doc("app", "users", id_key).?; + try testing.expectEqualStrings("bob", found.get("name").?.string); + const removed = try engine.remove("app", "users", id_key); + try testing.expect(removed); + engine.unlock(); +} + +test "reopen replays log" { + var threaded: std.Io.Threaded = .init_single_threaded; + defer threaded.deinit(); + var env = test_env(&threaded); + const io = env.io; + const gpa = testing.allocator; + + var tmp = try TmpLog.init(gpa); + defer tmp.deinit(gpa); + { + var engine = try Engine.open(gpa, io, tmp.path); + defer engine.deinit(); + var d1 = try make_doc(gpa, 1, "alice"); + defer d1.deinit(); + var d2 = try make_doc(gpa, 2, "bob"); + defer d2.deinit(); + try engine.lock(); + try engine.insert("app", "users", &d1, &env.gen); + try engine.insert("app", "users", &d2, &env.gen); + const id_key2 = try bson.serialize_value(gpa, bson.Value{ .int32 = 2 }); + defer gpa.free(id_key2); + _ = try engine.remove("app", "users", id_key2); + engine.unlock(); + } + + var engine2 = try Engine.open(gpa, io, tmp.path); + defer engine2.deinit(); + try engine2.lock(); + const id_key = try bson.serialize_value(gpa, bson.Value{ .int32 = 2 }); + defer gpa.free(id_key); + try testing.expect(engine2.get_doc("app", "users", id_key) == null); + const id_key1 = try bson.serialize_value(gpa, bson.Value{ .int32 = 1 }); + defer gpa.free(id_key1); + try testing.expectEqualStrings("alice", engine2.get_doc("app", "users", id_key1).?.get("name").?.string); + engine2.unlock(); +} + +test "auto _id generation survives reopen" { + var threaded: std.Io.Threaded = .init_single_threaded; + defer threaded.deinit(); + var env = test_env(&threaded); + const io = env.io; + const gpa = testing.allocator; + + var tmp = try TmpLog.init(gpa); + defer tmp.deinit(gpa); + { + var engine = try Engine.open(gpa, io, tmp.path); + defer engine.deinit(); + + var doc = try make_doc(gpa, 0, "no-id-here"); + defer doc.deinit(); + // strip _id + const stripped = doc.pairs[1..]; + + var arena = std.heap.ArenaAllocator.init(gpa); + defer arena.deinit(); + var d2 = try bson.Document.alloc(gpa, try arena.allocator().dupe(bson.Pair, stripped)); + defer d2.deinit(); + + try engine.lock(); + try engine.insert("app", "no_ids", &d2, &env.gen); + engine.unlock(); + } + + var engine2 = try Engine.open(gpa, io, tmp.path); + defer engine2.deinit(); + const coll = engine2.get_collection("app", "no_ids").?; + var it = coll.docs.iterator(); + var count: usize = 0; + while (it.next()) |entry| { + count += 1; + try testing.expect(entry.value_ptr.*.get("_id").?.object_id.len == 12); + } + try testing.expectEqual(@as(usize, 1), count); +} + +test "compaction rewrites log and keeps data" { + var threaded: std.Io.Threaded = .init_single_threaded; + defer threaded.deinit(); + var env = test_env(&threaded); + const io = env.io; + const gpa = testing.allocator; + + var tmp = try TmpLog.init(gpa); + defer tmp.deinit(gpa); + { + var engine = try Engine.open(gpa, io, tmp.path); + engine.compact_threshold = 1; // always compact + defer engine.deinit(); + + var docs: [4]bson.Document = undefined; + defer for (&docs) |*d| d.deinit(); + try engine.lock(); + for (0..4) |i| { + docs[i] = try make_doc(gpa, @intCast(i + 1), "user-{d}"); + try engine.insert("app", "users", &docs[i], &env.gen); + } + engine.unlock(); + } + + // Reopen after compaction and keep writing: with the log reopened at + // end_pos 0, appends would clobber the compacted records. + var engine2 = try Engine.open(gpa, io, tmp.path); + defer engine2.deinit(); + try engine2.lock(); + var extra = try make_doc(gpa, 5, "eve"); + defer extra.deinit(); + try engine2.insert("app", "users", &extra, &env.gen); + engine2.unlock(); + + var engine3 = try Engine.open(gpa, io, tmp.path); + defer engine3.deinit(); + try engine3.lock(); + for (1..6) |i| { + const id_key = try bson.serialize_value(gpa, bson.Value{ .int32 = @intCast(i) }); + defer gpa.free(id_key); + try testing.expect(engine3.get_doc("app", "users", id_key) != null); + } + engine3.unlock(); +} diff --git a/src/lib.zig b/src/lib.zig new file mode 100644 index 0000000..b8ee928 --- /dev/null +++ b/src/lib.zig @@ -0,0 +1,21 @@ +// mongo-light core library. Public entry point for tests and the server. + +pub const bson = @import("bson.zig"); +pub const wire = @import("wire.zig"); +pub const commands = @import("commands.zig"); +pub const server = @import("server.zig"); +pub const storage = @import("storage.zig"); +pub const db = @import("db.zig"); +pub const query = @import("query.zig"); +pub const update = @import("update.zig"); + +test { + _ = @import("bson.zig"); + _ = @import("wire.zig"); + _ = @import("commands.zig"); + _ = @import("server.zig"); + _ = @import("storage.zig"); + _ = @import("db.zig"); + _ = @import("query.zig"); + _ = @import("update.zig"); +} diff --git a/src/main.zig b/src/main.zig new file mode 100644 index 0000000..af6f7f5 --- /dev/null +++ b/src/main.zig @@ -0,0 +1,58 @@ +const std = @import("std"); +const mongo = @import("mongo"); + +const usage = + \\mongo-light — lightweight MongoDB-compatible document database + \\ + \\usage: mongo-light [options] + \\ --port listen port (default 27017) + \\ --bind bind address (default 127.0.0.1) + \\ --db database file (default mongo-light.log) + \\ --help show this help + \\ +; + +pub fn main(init: std.process.Init) !void { + var port: u16 = 27017; + var bind_ip: []const u8 = "127.0.0.1"; + var db_path: []const u8 = "mongo-light.log"; + + var it = std.process.Args.Iterator.init(init.minimal.args); + defer it.deinit(); + _ = it.next(); // program name + while (it.next()) |arg| { + if (std.mem.eql(u8, arg, "--port")) { + const v = it.next() orelse return error.MissingValue; + port = std.fmt.parseInt(u16, v, 10) catch { + std.debug.print("mongo-light: invalid port '{s}'\n", .{v}); + return error.InvalidPort; + }; + } else if (std.mem.eql(u8, arg, "--bind")) { + bind_ip = it.next() orelse return error.MissingValue; + } else if (std.mem.eql(u8, arg, "--db")) { + db_path = it.next() orelse return error.MissingValue; + } else if (std.mem.eql(u8, arg, "--help") or std.mem.eql(u8, arg, "-h")) { + try std.Io.File.writeStreamingAll(.stdout(), init.io, usage); + return; + } else { + std.debug.print("mongo-light: unknown option '{s}'\n{s}", .{ arg, usage }); + return error.UnknownOption; + } + } + + const oid_gen = mongo.bson.ObjectIdGen.init(init.io); + var engine = try mongo.db.Engine.open(init.gpa, init.io, db_path); + defer engine.deinit(); + std.debug.print("mongo-light: opened database '{s}'\n", .{db_path}); + + var server = mongo.server.Server{ + .gpa = init.gpa, + .port = port, + .bind_ip = bind_ip, + .oid_gen = oid_gen, + .connection_counter = .init(1), + .engine = &engine, + .start_time = std.Io.Timestamp.now(init.io, .real), + }; + try server.run(); +} diff --git a/src/query.zig b/src/query.zig new file mode 100644 index 0000000..4479098 --- /dev/null +++ b/src/query.zig @@ -0,0 +1,929 @@ +//! Query engine: filter matching (MongoDB query operators), sorting by +//! canonical BSON order, and projections. Includes a small backtracking +//! regex engine for $regex. + +const std = @import("std"); +const bson = @import("bson.zig"); + +// --------------------------------------------------------------------------- +// Filter matching +// --------------------------------------------------------------------------- + +pub const QueryError = error{OutOfMemory}; + +pub fn matches(gpa: std.mem.Allocator, filter: *const bson.Document, doc: *const bson.Document) QueryError!bool { + for (filter.pairs) |p| { + if (p.key.len > 0 and p.key[0] == '$') { + if (!try match_top_level(gpa, p.key, p.value, doc)) return false; + } else { + if (!try field_matches(gpa, p.key, p.value, doc)) return false; + } + } + return true; +} + +fn match_top_level(gpa: std.mem.Allocator, op: []const u8, value: bson.Value, doc: *const bson.Document) QueryError!bool { + if (std.mem.eql(u8, op, "$and") or std.mem.eql(u8, op, "$or")) { + const want_and = std.mem.eql(u8, op, "$and"); + const filters = switch (value) { + .array => |arr| arr, + else => return false, + }; + for (filters) |item| { + const f = switch (item) { + .doc => |pairs| pairs, + else => return false, + }; + const matched = try matches(gpa, &.{ .arena = undefined, .pairs = f }, doc); + if (want_and and !matched) return false; + if (!want_and and matched) return true; + } + return want_and; + } + if (std.mem.eql(u8, op, "$nor")) { + const filters = switch (value) { + .array => |arr| arr, + else => return false, + }; + for (filters) |item| { + const f = switch (item) { + .doc => |pairs| pairs, + else => return false, + }; + if (try matches(gpa, &.{ .arena = undefined, .pairs = f }, doc)) return false; + } + return true; + } + return false; +} + +fn is_operator_doc(value: bson.Value) ?[]const bson.Pair { + return switch (value) { + .doc => |pairs| blk: { + for (pairs) |p| { + if (p.key.len == 0 or p.key[0] != '$') break :blk null; + } + break :blk pairs; + }, + else => null, + }; +} + +fn field_matches(gpa: std.mem.Allocator, path: []const u8, expected: bson.Value, doc: *const bson.Document) QueryError!bool { + var candidates: std.ArrayListUnmanaged(bson.Value) = .empty; + defer candidates.deinit(gpa); + try collect_values(gpa, doc.pairs, path, &candidates, 0); + // MongoDB applies queries to array elements as well as the array itself. + // Index the snapshot length, re-reading items each iteration: appending + // may reallocate the buffer, which would invalidate a captured slice. + const direct_count = candidates.items.len; + var i: usize = 0; + while (i < direct_count) : (i += 1) { + const a = candidates.items[i]; + if (a == .array) { + for (a.array) |elem| try candidates.append(gpa, elem); + } + } + + if (is_operator_doc(expected)) |pairs| { + var options: []const u8 = ""; + for (pairs) |p| { + if (std.mem.eql(u8, p.key, "$options")) { + if (p.value == .string) options = p.value.string; + } + } + for (pairs) |p| { + if (std.mem.eql(u8, p.key, "$options")) continue; + if (!try match_operator(gpa, p.key, p.value, candidates.items, options)) return false; + } + return true; + } + // Bare equality — matches if any candidate equals the expected value. + for (candidates.items) |actual| { + if (bson.compare(actual, expected) == .eq) return true; + } + return false; +} + +fn match_operator(gpa: std.mem.Allocator, op: []const u8, value: bson.Value, actuals: []const bson.Value, regex_options: []const u8) QueryError!bool { + if (std.mem.eql(u8, op, "$eq")) { + for (actuals) |a| if (bson.compare(a, value) == .eq) return true; + return false; + } + if (std.mem.eql(u8, op, "$ne")) { + for (actuals) |a| if (bson.compare(a, value) == .eq) return false; + return true; + } + if (std.mem.eql(u8, op, "$gt") or std.mem.eql(u8, op, "$gte") or + std.mem.eql(u8, op, "$lt") or std.mem.eql(u8, op, "$lte")) + { + for (actuals) |a| { + const o = bson.compare(a, value); + if (std.mem.eql(u8, op, "$gt") and o == .gt) return true; + if (std.mem.eql(u8, op, "$gte") and o != .lt) return true; + if (std.mem.eql(u8, op, "$lt") and o == .lt) return true; + if (std.mem.eql(u8, op, "$lte") and o != .gt) return true; + } + return false; + } + if (std.mem.eql(u8, op, "$in") or std.mem.eql(u8, op, "$nin")) { + const members = switch (value) { + .array => |arr| arr, + else => return false, + }; + const want_in = std.mem.eql(u8, op, "$in"); + for (actuals) |a| { + for (members) |m| { + if (bson.compare(a, m) == .eq) return want_in; + } + } + return !want_in; + } + if (std.mem.eql(u8, op, "$exists")) { + const want = switch (value) { + .bool => |b| b, + else => return false, + }; + return (actuals.len > 0) == want; + } + if (std.mem.eql(u8, op, "$regex")) { + const pattern = switch (value) { + .string => |s| s, + .doc => |pairs| blk: { + const pat = bson.get_pair(pairs, "$regex") orelse return false; + break :blk switch (pat) { + .string => |s| s, + else => return false, + }; + }, + else => return false, + }; + for (actuals) |a| { + if (a == .string and regex_match(pattern, regex_options, a.string)) return true; + } + return false; + } + if (std.mem.eql(u8, op, "$not")) { + const pairs = is_operator_doc(value) orelse { + // $not with a bare value means $ne-ish semantics; treat as + // "not equal to this regex or value". + if (value == .string) { + for (actuals) |a| { + if (a == .string and regex_match(value.string, "", a.string)) return false; + } + return true; + } + return false; + }; + for (pairs) |p| { + if (try match_operator(gpa, p.key, p.value, actuals, regex_options)) return false; + } + return true; + } + if (std.mem.eql(u8, op, "$size")) { + const want = switch (value) { + .int32 => |i| i, + .int64 => |i| @as(i32, @intCast(i)), + else => return false, + }; + for (actuals) |a| { + if (a == .array and a.array.len == @as(usize, @intCast(want))) return true; + } + return false; + } + if (std.mem.eql(u8, op, "$all")) { + const members = switch (value) { + .array => |arr| arr, + else => return false, + }; + outer: for (members) |m| { + for (actuals) |a| { + if (a == .array) { + for (a.array) |elem| { + if (bson.compare(elem, m) == .eq) continue :outer; + } + } + } + return false; + } + return true; + } + if (std.mem.eql(u8, op, "$elemMatch")) { + const operand = switch (value) { + .doc => |pairs| pairs, + else => return false, + }; + var all_operators = operand.len > 0; + for (operand) |p| { + if (p.key.len == 0 or p.key[0] != '$') { + all_operators = false; + break; + } + } + for (actuals) |a| { + if (a != .array) continue; + for (a.array) |elem| { + if (all_operators) { + var single: [1]bson.Value = .{elem}; + var ok = true; + for (operand) |p| { + if (!try match_operator(gpa, p.key, p.value, single[0..], "")) { + ok = false; + break; + } + } + if (ok) return true; + } else { + switch (elem) { + .doc => |pairs| { + if (try matches(gpa, &.{ .arena = undefined, .pairs = operand }, &.{ .arena = undefined, .pairs = pairs })) return true; + }, + else => {}, + } + } + } + } + return false; + } + return false; +} + +/// Collect values reachable at `path` (dot-separated), descending into +/// documents and, per MongoDB multikey semantics, into arrays of documents. +/// Appends into `out`; on OOM, collection stops early (the engine is +/// already failing at that point). +fn collect_values(gpa: std.mem.Allocator, pairs: []const bson.Pair, path: []const u8, out: *std.ArrayListUnmanaged(bson.Value), depth: usize) QueryError!void { + var it = std.mem.splitScalar(u8, path, '.'); + const first = it.next() orelse return; + + for (pairs) |p| { + if (!std.mem.eql(u8, p.key, first)) continue; + const rest = it.rest(); + if (rest.len == 0) { + if (depth < 8) { + try out.append(gpa, p.value); + } + } else { + try collect_from_value(gpa, p.value, rest, out, depth + 1); + } + } +} + +fn collect_from_value(gpa: std.mem.Allocator, v: bson.Value, path: []const u8, out: *std.ArrayListUnmanaged(bson.Value), depth: usize) QueryError!void { + if (depth > 8) return; + switch (v) { + .doc => |pairs| try collect_values(gpa, pairs, path, out, depth), + .array => |items| { + for (items) |item| { + switch (item) { + .doc => try collect_values(gpa, item.doc, path, out, depth), + else => {}, + } + } + }, + else => return, + } +} + +// --------------------------------------------------------------------------- +// Regex (subset): ^ $ . * + ? [...] ( ) | and escaped literals +// --------------------------------------------------------------------------- + +pub fn regex_match(pattern: []const u8, options: []const u8, text: []const u8) bool { + // Bound recursion depth: deeply nested groups would overflow the stack. + var depth: usize = 0; + var max_depth: usize = 0; + var ri: usize = 0; + while (ri < pattern.len) : (ri += 1) { + if (pattern[ri] == '\\') { + ri += 1; + continue; + } + if (pattern[ri] == '(') { + depth += 1; + max_depth = @max(max_depth, depth); + } else if (pattern[ri] == ')') { + depth -|= 1; + } + } + if (max_depth > 256) return false; + + const case_insensitive = std.mem.indexOfScalar(u8, options, 'i') != null; + const dot_all = std.mem.indexOfScalar(u8, options, 's') != null; + + const anchored = pattern.len > 0 and pattern[0] == '^'; + const start_pattern = if (anchored) pattern[1..] else pattern; + + var p: usize = 0; + if (anchored) { + if (match_here(start_pattern, &p, text, 0, case_insensitive, dot_all) == null) return false; + return p == start_pattern.len; + } + var t: usize = 0; + while (t <= text.len) : (t += 1) { + p = 0; + if (match_here(start_pattern, &p, text, t, case_insensitive, dot_all) != null and p == start_pattern.len) return true; + } + return false; +} + +/// Match `pattern[p..]` against `text[t..]`, returning the new text +/// position on success (null on failure). Backtracks via recursion. +fn match_here(pattern: []const u8, p: *usize, text: []const u8, t: usize, ci: bool, dot_all: bool) ?usize { + var pos = t; + while (p.* < pattern.len) { + const c = pattern[p.*]; + switch (c) { + '$' => { + if (p.* + 1 == pattern.len) { + p.* += 1; + return if (pos == text.len) pos else null; + } + if (pos >= text.len) return null; + if (!chars_eq(c, text[pos], ci)) return null; + p.* += 1; + pos += 1; + }, + '^' => { + if (pos != 0) return null; + p.* += 1; + }, + '(' => { + const end = find_group_end(pattern, p.*) orelse return null; + const inner = pattern[p.* + 1 .. end - 1]; + var parts: [8][]const u8 = undefined; + var nparts: usize = 0; + var seg_start: usize = 0; + var depth: usize = 0; + var i: usize = 0; + while (i < inner.len) : (i += 1) { + const ic = inner[i]; + if (ic == '\\') { + i += 1; + continue; + } + if (ic == '(') depth +|= 1; + if (ic == ')') depth -|= 1; + if (ic == '|' and depth == 0) { + if (nparts < parts.len) parts[nparts] = inner[seg_start..i]; + nparts += 1; + seg_start = i + 1; + } + } + if (nparts < parts.len) parts[nparts] = inner[seg_start..]; + nparts += 1; + if (nparts == 1) { + var gp: usize = 0; + const after = match_here(parts[0], &gp, text, pos, ci, dot_all) orelse return null; + if (gp != parts[0].len) return null; + p.* = end; + pos = after; + } else { + var matched = false; + for (parts[0..nparts]) |part| { + var gp: usize = 0; + const after = match_here(part, &gp, text, pos, ci, dot_all) orelse continue; + if (gp == part.len) { + p.* = end; + pos = after; + matched = true; + break; + } + } + if (!matched) return null; + } + }, + ')' => return null, // unbalanced + '\\' => { + if (p.* + 1 >= pattern.len) return null; + const lit = pattern[p.* + 1]; + if (pos >= text.len or !chars_eq(lit, text[pos], ci)) return null; + p.* += 2; + pos += 1; + }, + else => { + var element_end: usize = undefined; + if (c == '[') { + var close = p.* + 1; + if (close < pattern.len and pattern[close] == '^') close += 1; + while (close < pattern.len and pattern[close] != ']') close += 1; + if (close >= pattern.len) return null; + element_end = close + 1; + } else { + element_end = p.* + 1; + } + const element = pattern[p.*..element_end]; + + var q_end = element_end; + var min: usize = 1; + var max: usize = 1; + if (element_end < pattern.len and (pattern[element_end] == '*' or pattern[element_end] == '+' or pattern[element_end] == '?')) { + switch (pattern[element_end]) { + '*' => { + min = 0; + max = std.math.maxInt(usize); + }, + '+' => { + min = 1; + max = std.math.maxInt(usize); + }, + '?' => { + min = 0; + }, + else => {}, + } + q_end = element_end + 1; + } + + // Greedy: consume as many as possible, then backtrack. + var consumed: usize = 0; + var pos_cur = pos; + while (max == std.math.maxInt(usize) or consumed < max) { + if (element_matches(element, text, pos_cur, ci, dot_all)) { + pos_cur += 1; + consumed += 1; + } else break; + } + var attempt = consumed; + while (attempt >= min) : (attempt -= 1) { + p.* = q_end; + if (match_here(pattern, p, text, pos_cur - (consumed - attempt), ci, dot_all)) |after| { + return after; + } + if (attempt == 0) break; + } + return null; + }, + } + } + return pos; +} + +fn chars_eq(a: u8, b: u8, ci: bool) bool { + if (ci) return std.ascii.toLower(a) == std.ascii.toLower(b); + return a == b; +} + +fn in_range(lo: u8, hi: u8, c: u8, ci: bool) bool { + if (ci) { + const l = std.ascii.toLower(c); + return l >= std.ascii.toLower(lo) and l <= std.ascii.toLower(hi); + } + return c >= lo and c <= hi; +} + +fn find_group_end(pattern: []const u8, open: usize) ?usize { + var depth: usize = 1; + var i = open + 1; + while (i < pattern.len) : (i += 1) { + if (pattern[i] == '\\') { + i += 1; + continue; + } + if (pattern[i] == '(') depth += 1; + if (pattern[i] == ')') { + depth -= 1; + if (depth == 0) return i + 1; + } + } + return null; +} + +fn element_matches(element: []const u8, text: []const u8, t: usize, ci: bool, dot_all: bool) bool { + if (t >= text.len) return false; + if (element[0] == '.') { + return dot_all or text[t] != '\n'; + } + if (element[0] == '[') { + var negated = false; + var i: usize = 1; + if (i < element.len and element[i] == '^') { + negated = true; + i += 1; + } + var matched = false; + while (i < element.len and element[i] != ']') { + if (i + 2 < element.len and element[i + 1] == '-') { + if (in_range(element[i], element[i + 2], text[t], ci)) matched = true; + i += 3; + } else { + if (chars_eq(element[i], text[t], ci)) matched = true; + i += 1; + } + } + return matched != negated; + } + return chars_eq(element[0], text[t], ci); +} + +// --------------------------------------------------------------------------- +// Sort +// --------------------------------------------------------------------------- + +pub const SortKey = struct { + path: []const u8, + descending: bool, +}; + +/// Sort `docs` in place by `keys`. Candidate values are collected up front +/// (allocations happen before the sort), so the comparator itself is pure +/// and cannot fail — OOM during collection propagates as QueryError. +pub fn sort_docs(gpa: std.mem.Allocator, arena: std.mem.Allocator, docs: []*const bson.Document, keys: []const SortKey) QueryError!void { + if (keys.len == 0 or docs.len < 2) return; + + const SortedDoc = struct { + doc: *const bson.Document, + values: [][]const bson.Value, + }; + const entries = try arena.alloc(SortedDoc, docs.len); + for (docs, 0..) |d, i| { + const values = try arena.alloc([]const bson.Value, keys.len); + for (keys, 0..) |k, ki| { + var list: std.ArrayListUnmanaged(bson.Value) = .empty; + try collect_values(arena, d.pairs, k.path, &list, 0); + values[ki] = list.items; + } + entries[i] = .{ .doc = d, .values = values }; + } + + const Ctx = struct { + keys: []const SortKey, + fn lessThan(ctx: @This(), a: SortedDoc, b: SortedDoc) bool { + for (ctx.keys, 0..) |k, ki| { + const aval: bson.Value = if (a.values[ki].len > 0) a.values[ki][0] else .null; + const bval: bson.Value = if (b.values[ki].len > 0) b.values[ki][0] else .null; + const o = bson.compare(aval, bval); + if (o != .eq) return if (k.descending) o == .gt else o == .lt; + } + return false; + } + }; + std.mem.sort(SortedDoc, entries, Ctx{ .keys = keys }, Ctx.lessThan); + + for (entries, 0..) |e, i| docs[i] = e.doc; + _ = gpa; +} + +// --------------------------------------------------------------------------- +// Projection +// --------------------------------------------------------------------------- + +pub const ProjectionError = error{ OutOfMemory, InvalidProjection }; + +/// Apply a projection document, writing resulting pairs into `out` (which +/// should use the caller's arena so strings are owned). +pub fn project(arena: std.mem.Allocator, doc: *const bson.Document, proj: *const bson.Document, out: *std.ArrayListUnmanaged(bson.Pair)) ProjectionError!void { + var inclusion: ?bool = null; + var non_id_count: usize = 0; + for (proj.pairs) |p| { + if (std.mem.eql(u8, p.key, "_id")) continue; + non_id_count += 1; + const flag = projection_flag(p.value); + inclusion = if (inclusion == null) flag else inclusion; + } + + // {_id: 0} alone means "drop _id, keep everything else". + const include = if (non_id_count > 0) (inclusion orelse true) else false; + if (include) { + // Inclusion list: _id unless excluded, plus listed paths. + var include_id = true; + if (bson.get_pair(proj.pairs, "_id")) |idv| { + include_id = projection_flag(idv); + } + if (include_id) { + if (bson.get_pair(doc.pairs, "_id")) |idv| { + try out.append(arena, .{ .key = try arena.dupe(u8, "_id"), .value = try bson.copy_value(arena, idv) }); + } + } + for (proj.pairs) |p| { + if (std.mem.eql(u8, p.key, "_id")) continue; + if (!projection_flag(p.value)) continue; + try project_path(arena, doc.pairs, p.key, out); + } + } else { + // Exclusion: copy everything except excluded paths (and _id if set). + for (doc.pairs) |p| { + if (std.mem.eql(u8, p.key, "_id")) { + var excluded = false; + if (bson.get_pair(proj.pairs, "_id")) |idv| excluded = !projection_flag(idv); + if (excluded) continue; + } + if (is_excluded(proj, p.key)) continue; + if (p.value == .doc and has_deeper_exclusion(proj, p.key)) { + const filtered = try exclude_doc(arena, p.value.doc, proj, p.key); + try out.append(arena, .{ .key = try arena.dupe(u8, p.key), .value = .{ .doc = filtered } }); + } else { + try out.append(arena, .{ .key = try arena.dupe(u8, p.key), .value = try bson.copy_value(arena, p.value) }); + } + } + } +} + +/// Recursively apply exclusions to a nested document given the parent path. +fn exclude_doc(arena: std.mem.Allocator, pairs: []const bson.Pair, proj: *const bson.Document, parent: []const u8) ProjectionError![]const bson.Pair { + var out: std.ArrayListUnmanaged(bson.Pair) = .empty; + errdefer out.deinit(arena); + for (pairs) |p| { + const full = if (parent.len > 0) try std.fmt.allocPrint(arena, "{s}.{s}", .{ parent, p.key }) else p.key; + if (is_excluded(proj, full)) continue; + if (p.value == .doc and has_deeper_exclusion(proj, full)) { + const filtered = try exclude_doc(arena, p.value.doc, proj, full); + try out.append(arena, .{ .key = try arena.dupe(u8, p.key), .value = .{ .doc = filtered } }); + } else { + try out.append(arena, .{ .key = try arena.dupe(u8, p.key), .value = try bson.copy_value(arena, p.value) }); + } + } + return out.toOwnedSlice(arena); +} + +fn is_excluded(proj: *const bson.Document, key: []const u8) bool { + for (proj.pairs) |pp| { + if (std.mem.eql(u8, pp.key, key)) return true; + } + return false; +} + +fn has_deeper_exclusion(proj: *const bson.Document, key: []const u8) bool { + for (proj.pairs) |pp| { + if (is_prefix_or_equal(key, pp.key)) return true; + } + return false; +} + +fn projection_flag(v: bson.Value) bool { + return switch (v) { + .bool => |b| b, + .int32 => |i| i != 0, + .int64 => |i| i != 0, + .double => |d| d != 0, + else => false, + }; +} + +/// Include a dotted path (e.g. "a.b.c"), creating nested documents as needed. +fn project_path(arena: std.mem.Allocator, pairs: []const bson.Pair, path: []const u8, out: *std.ArrayListUnmanaged(bson.Pair)) ProjectionError!void { + var it = std.mem.splitScalar(u8, path, '.'); + const first = it.next() orelse return; + const rest = it.rest(); + + // Does the doc have this top-level field? + const field = bson.get_pair(pairs, first); + if (rest.len == 0) { + if (field) |f| { + try out.append(arena, .{ .key = try arena.dupe(u8, first), .value = try bson.copy_value(arena, f) }); + } + return; + } + if (field) |f| { + switch (f) { + .doc => |sub| { + var nested: std.ArrayListUnmanaged(bson.Pair) = .empty; + errdefer nested.deinit(arena); + try project_path(arena, sub, rest, &nested); + if (nested.items.len > 0) { + try out.append(arena, .{ .key = try arena.dupe(u8, first), .value = .{ .doc = try nested.toOwnedSlice(arena) } }); + } + }, + else => {}, + } + } +} + +fn is_prefix_or_equal(prefix: []const u8, key: []const u8) bool { + if (std.mem.eql(u8, prefix, key)) return true; + if (prefix.len < key.len and std.mem.eql(u8, prefix, key[0..prefix.len])) { + return key[prefix.len] == '.'; + } + return false; +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +const testing = std.testing; + +fn doc_of(pairs: []const bson.Pair) bson.Document { + return .{ .arena = undefined, .pairs = pairs }; +} + +test "basic filters" { + const d = doc_of(&.{ + .{ .key = "_id", .value = .{ .int32 = 1 } }, + .{ .key = "age", .value = .{ .int32 = 30 } }, + .{ .key = "name", .value = .{ .string = "alice" } }, + .{ .key = "tags", .value = .{ .array = &.{ .{ .string = "a" }, .{ .string = "b" } } } }, + }); + try testing.expect(try matches(testing.allocator, &doc_of(&.{}), &d)); + try testing.expect(try matches(testing.allocator, &doc_of(&.{.{ .key = "age", .value = .{ .int32 = 30 } }}), &d)); + try testing.expect(!try matches(testing.allocator, &doc_of(&.{.{ .key = "age", .value = .{ .int32 = 31 } }}), &d)); + try testing.expect(try matches(testing.allocator, &doc_of(&.{.{ .key = "age", .value = .{ .doc = &.{.{ .key = "$gt", .value = .{ .int32 = 20 } }} } }}), &d)); + try testing.expect(!try matches(testing.allocator, &doc_of(&.{.{ .key = "age", .value = .{ .doc = &.{.{ .key = "$gt", .value = .{ .int32 = 30 } }} } }}), &d)); + try testing.expect(try matches(testing.allocator, &doc_of(&.{.{ .key = "age", .value = .{ .doc = &.{.{ .key = "$in", .value = .{ .array = &.{ .{ .int32 = 29 }, .{ .int32 = 30 } } } }} } }}), &d)); + try testing.expect(try matches(testing.allocator, &doc_of(&.{.{ .key = "missing", .value = .{ .doc = &.{.{ .key = "$exists", .value = .{ .bool = false } }} } }}), &d)); + try testing.expect(try matches(testing.allocator, &doc_of(&.{.{ .key = "tags", .value = .{ .string = "b" } }}), &d)); + try testing.expect(try matches(testing.allocator, &doc_of(&.{.{ .key = "tags", .value = .{ .doc = &.{.{ .key = "$size", .value = .{ .int32 = 2 } }} } }}), &d)); +} + +test "regex filter" { + const d = doc_of(&.{.{ .key = "name", .value = .{ .string = "Alice Smith" } }}); + try testing.expect(try matches(testing.allocator, &doc_of(&.{.{ .key = "name", .value = .{ .doc = &.{ + .{ .key = "$regex", .value = .{ .string = "^al" } }, + .{ .key = "$options", .value = .{ .string = "i" } }, + } } }}), &d)); + try testing.expect(!try matches(testing.allocator, &doc_of(&.{.{ .key = "name", .value = .{ .doc = &.{.{ .key = "$regex", .value = .{ .string = "^X" } }} } }}), &d)); + try testing.expect(!try matches(testing.allocator, &doc_of(&.{.{ .key = "name", .value = .{ .doc = &.{.{ .key = "$regex", .value = .{ .string = "sm[i]th$" } }} } }}), &d)); + try testing.expect(try matches(testing.allocator, &doc_of(&.{.{ .key = "name", .value = .{ .doc = &.{.{ .key = "$regex", .value = .{ .string = "th$" } }} } }}), &d)); +} + +test "regex engine basics" { + try testing.expect(regex_match("abc", "", "xabcx")); + try testing.expect(regex_match("^abc", "", "abc")); + try testing.expect(!regex_match("^abc", "", "xabc")); + try testing.expect(regex_match("a.c", "", "abc")); + try testing.expect(regex_match("a*b", "", "aaab")); + try testing.expect(regex_match("a+b", "", "aab")); + try testing.expect(regex_match("colou?r", "", "color")); + try testing.expect(regex_match("colou?r", "", "colour")); + try testing.expect(regex_match("[0-9]+", "", "abc123def")); + try testing.expect(!regex_match("^[0-9]+$", "", "abc123")); + try testing.expect(regex_match("(ab|cd)e", "", "cde")); + try testing.expect(regex_match("a\\.b", "", "a.b")); + try testing.expect(regex_match("^foo$", "", "foo")); + try testing.expect(!regex_match("^foo$", "", "foobar")); + try testing.expect(regex_match("hello", "i", "HELLO")); + try testing.expect(regex_match("^a+$", "", "aaaa")); + try testing.expect(!regex_match("^a+$", "", "aaab")); +} + +test "regex hostile input does not crash" { + // Escaped paren inside a group used to underflow the alternation scan. + // Pattern is: "(" ++ "\" ++ "))" — a group containing an escaped ')'. + const hostile = "(" ++ "\\" ++ "))"; + try testing.expect(regex_match(hostile, "", ")")); + try testing.expect(!regex_match(hostile, "", "x")); + try testing.expect(regex_match("a" ++ "\\" ++ "|b", "", "a|b")); + // Deeply nested groups must be rejected, not blow the stack. + const deep = "(" ** 300 ++ "x" ++ ")" ** 300; + try testing.expect(!regex_match(deep, "", "x")); +} + +test "long array values are not truncated" { + var items: [30]bson.Value = undefined; + for (0..30) |i| items[i] = .{ .int32 = @intCast(i) }; + const d = doc_of(&.{.{ .key = "tags", .value = .{ .array = &items } }}); + try testing.expect(try matches(testing.allocator, &doc_of(&.{.{ .key = "tags", .value = .{ .int32 = 29 } }}), &d)); + try testing.expect(!try matches(testing.allocator, &doc_of(&.{.{ .key = "tags", .value = .{ .int32 = 30 } }}), &d)); +} + +test "multikey path with several array candidates does not use-after-free" { + // Two array candidates at the path; flattening the first appends past + // the initial capacity, which used to realloc the buffer while a + // captured slice was still being iterated. + var big: [20]bson.Value = undefined; + for (0..20) |j| big[j] = .{ .int32 = @intCast(j) }; + const d = doc_of(&.{.{ .key = "items", .value = .{ .array = &.{ + .{ .doc = &.{.{ .key = "tags", .value = .{ .array = &big } }} }, + .{ .doc = &.{.{ .key = "tags", .value = .{ .array = &.{.{ .string = "needle" }} } }} }, + } } }}); + try testing.expect(try matches(testing.allocator, &doc_of(&.{.{ .key = "items.tags", .value = .{ .string = "needle" } }}), &d)); + // And the flattened long array is searched, not just the first element. + try testing.expect(try matches(testing.allocator, &doc_of(&.{.{ .key = "items.tags", .value = .{ .int32 = 19 } }}), &d)); +} + +test "OOM during value collection propagates, not a false match" { + // A tiny FixedBufferAllocator makes the candidate collection fail; the + // error must surface instead of leaving an empty candidate list, which + // would make negating operators like $ne report a match. + const d = doc_of(&.{.{ .key = "x", .value = .{ .int32 = 5 } }}); + const ne = doc_of(&.{.{ .key = "x", .value = .{ .doc = &.{.{ .key = "$ne", .value = .{ .int32 = 5 } }} } }}); + + var buf: [16]u8 = undefined; + var fba = std.heap.FixedBufferAllocator.init(&buf); + try testing.expectError(error.OutOfMemory, matches(fba.allocator(), &ne, &d)); + + const ex = doc_of(&.{.{ .key = "x", .value = .{ .doc = &.{.{ .key = "$exists", .value = .{ .bool = false } }} } }}); + try testing.expectError(error.OutOfMemory, matches(fba.allocator(), &ex, &d)); +} + +test "documents compare by field name too" { + try testing.expectEqual(std.math.Order.lt, bson.compare( + .{ .doc = &.{.{ .key = "a", .value = .{ .int32 = 1 } }} }, + .{ .doc = &.{.{ .key = "b", .value = .{ .int32 = 1 } }} }, + )); +} + +test "dot path filters" { + const d = doc_of(&.{ + .{ .key = "user", .value = .{ .doc = &.{ + .{ .key = "profile", .value = .{ .doc = &.{.{ .key = "age", .value = .{ .int32 = 25 } }} } }, + } } }, + .{ .key = "items", .value = .{ .array = &.{ + .{ .doc = &.{.{ .key = "sku", .value = .{ .string = "x" } }} }, + .{ .doc = &.{.{ .key = "sku", .value = .{ .string = "y" } }} }, + } } }, + }); + try testing.expect(try matches(testing.allocator, &doc_of(&.{.{ .key = "user.profile.age", .value = .{ .int32 = 25 } }}), &d)); + try testing.expect(!try matches(testing.allocator, &doc_of(&.{.{ .key = "user.profile.age", .value = .{ .int32 = 26 } }}), &d)); + try testing.expect(try matches(testing.allocator, &doc_of(&.{.{ .key = "items.sku", .value = .{ .string = "y" } }}), &d)); +} + +test "sort compares by BSON order" { + const a = doc_of(&.{ .{ .key = "n", .value = .{ .int32 = 2 } }, .{ .key = "x", .value = .{ .string = "a" } } }); + const b = doc_of(&.{ .{ .key = "n", .value = .{ .int32 = 10 } }, .{ .key = "x", .value = .{ .string = "b" } } }); + var arena = std.heap.ArenaAllocator.init(testing.allocator); + defer arena.deinit(); + + var docs = [_]*const bson.Document{ &b, &a }; + const asc = [_]SortKey{.{ .path = "n", .descending = false }}; + try sort_docs(testing.allocator, arena.allocator(), &docs, &asc); + try testing.expect(docs[0] == &a); + try testing.expect(docs[1] == &b); + + var docs2 = [_]*const bson.Document{ &a, &b }; + const desc = [_]SortKey{.{ .path = "n", .descending = true }}; + try sort_docs(testing.allocator, arena.allocator(), &docs2, &desc); + try testing.expect(docs2[0] == &b); + + const missing = [_]SortKey{.{ .path = "zz", .descending = false }}; + try sort_docs(testing.allocator, arena.allocator(), &docs2, &missing); + try testing.expect(docs2[0] == &b); // stable-ish: order untouched by missing key +} + +test "projection inclusion and exclusion" { + const d = doc_of(&.{ + .{ .key = "_id", .value = .{ .int32 = 1 } }, + .{ .key = "a", .value = .{ .int32 = 1 } }, + .{ .key = "b", .value = .{ .int32 = 2 } }, + .{ .key = "nested", .value = .{ .doc = &.{ + .{ .key = "x", .value = .{ .int32 = 3 } }, + .{ .key = "y", .value = .{ .int32 = 4 } }, + } } }, + }); + + var arena = std.heap.ArenaAllocator.init(testing.allocator); + defer arena.deinit(); + + var out: std.ArrayListUnmanaged(bson.Pair) = .empty; + defer out.deinit(arena.allocator()); + try project(arena.allocator(), &d, &doc_of(&.{ + .{ .key = "a", .value = .{ .int32 = 1 } }, + .{ .key = "nested.x", .value = .{ .int32 = 1 } }, + }), &out); + try testing.expectEqual(@as(usize, 3), out.items.len); + try testing.expect(bson.get_pair(out.items, "a") != null); + try testing.expect(bson.get_pair(out.items, "b") == null); + const nx = bson.get_pair(out.items, "nested").?; + try testing.expectEqual(@as(usize, 1), nx.doc.len); + try testing.expectEqualStrings("x", nx.doc[0].key); + + // {_id: 0} alone: everything except _id. + out.clearRetainingCapacity(); + try project(arena.allocator(), &d, &doc_of(&.{.{ .key = "_id", .value = .{ .int32 = 0 } }}), &out); + try testing.expect(bson.get_pair(out.items, "_id") == null); + try testing.expect(bson.get_pair(out.items, "a") != null); + try testing.expect(bson.get_pair(out.items, "b") != null); + + // Mixed inclusion with _id: 0: only listed fields, no _id. + out.clearRetainingCapacity(); + try project(arena.allocator(), &d, &doc_of(&.{ + .{ .key = "_id", .value = .{ .int32 = 0 } }, + .{ .key = "a", .value = .{ .int32 = 1 } }, + }), &out); + try testing.expect(bson.get_pair(out.items, "_id") == null); + try testing.expect(bson.get_pair(out.items, "a") != null); + try testing.expect(bson.get_pair(out.items, "b") == null); + + out.clearRetainingCapacity(); + try project(arena.allocator(), &d, &doc_of(&.{ + .{ .key = "b", .value = .{ .int32 = 0 } }, + .{ .key = "nested.x", .value = .{ .int32 = 0 } }, + }), &out); + try testing.expect(bson.get_pair(out.items, "a") != null); + try testing.expect(bson.get_pair(out.items, "b") == null); + const n = bson.get_pair(out.items, "nested").?; + try testing.expectEqual(@as(usize, 1), n.doc.len); + try testing.expectEqualStrings("y", n.doc[0].key); +} + +test "and/or filters" { + const d = doc_of(&.{ .{ .key = "a", .value = .{ .int32 = 1 } }, .{ .key = "b", .value = .{ .int32 = 2 } } }); + try testing.expect(try matches(testing.allocator, &doc_of(&.{.{ .key = "$and", .value = .{ .array = &.{ + .{ .doc = &.{.{ .key = "a", .value = .{ .int32 = 1 } }} }, + .{ .doc = &.{.{ .key = "b", .value = .{ .int32 = 2 } }} }, + } } }}), &d)); + try testing.expect(try matches(testing.allocator, &doc_of(&.{.{ .key = "$or", .value = .{ .array = &.{ + .{ .doc = &.{.{ .key = "a", .value = .{ .int32 = 99 } }} }, + .{ .doc = &.{.{ .key = "b", .value = .{ .int32 = 2 } }} }, + } } }}), &d)); + try testing.expect(!try matches(testing.allocator, &doc_of(&.{.{ .key = "$or", .value = .{ .array = &.{ + .{ .doc = &.{.{ .key = "a", .value = .{ .int32 = 99 } }} }, + .{ .doc = &.{.{ .key = "b", .value = .{ .int32 = 99 } }} }, + } } }}), &d)); +} + +/// Public single-value operator matcher, used by $pull and $elemMatch. +pub fn value_matches_operator(gpa: std.mem.Allocator, op: []const u8, value: bson.Value, actual: bson.Value) QueryError!bool { + var single: [1]bson.Value = .{actual}; + return match_operator(gpa, op, value, single[0..], ""); +} diff --git a/src/server.zig b/src/server.zig new file mode 100644 index 0000000..547e2e8 --- /dev/null +++ b/src/server.zig @@ -0,0 +1,145 @@ +//! TCP server speaking the MongoDB wire protocol. Accept loop dispatches each +//! connection onto the Io worker pool; a connection is handled until the peer +//! closes it or a protocol error occurs. + +const std = @import("std"); +const bson = @import("bson.zig"); +const wire = @import("wire.zig"); +const commands = @import("commands.zig"); +const db = @import("db.zig"); + +pub const Server = struct { + gpa: std.mem.Allocator, + port: u16, + bind_ip: []const u8, + oid_gen: bson.ObjectIdGen, + connection_counter: std.atomic.Value(u32), + engine: *db.Engine, + start_time: std.Io.Timestamp, + + pub fn run(self: *Server) !void { + var threaded: std.Io.Threaded = std.Io.Threaded.init(self.gpa, .{}); + defer threaded.deinit(); + const io = threaded.io(); + + var addr = try std.Io.net.IpAddress.parse(self.bind_ip, self.port); + var listener = try addr.listen(io, .{ .reuse_address = true }); + defer listener.deinit(io); + + std.debug.print("mongo-light: listening on {s}:{d}\n", .{ self.bind_ip, self.port }); + + var group: std.Io.Group = .init; + defer group.cancel(io); + + while (true) { + const stream = listener.accept(io) catch |err| switch (err) { + error.Canceled => return, + else => { + std.debug.print("mongo-light: accept error: {s}\n", .{@errorName(err)}); + continue; + }, + }; + group.async(io, handle_connection, .{ io, stream, self }); + } + } +}; + +/// Entry point required by `Group.async`: must return only `error.Canceled`. +fn handle_connection(io: std.Io, stream: std.Io.net.Stream, server: *Server) error{Canceled}!void { + handle_connection_inner(io, stream, server) catch {}; +} + +fn handle_connection_inner(io: std.Io, stream: std.Io.net.Stream, server: *Server) !void { + defer stream.close(io); + + const connection_id = server.connection_counter.fetchAdd(1, .monotonic); + var read_scratch: [16 * 1024]u8 = undefined; + var write_scratch: [16 * 1024]u8 = undefined; + var reader = stream.reader(io, &read_scratch); + var writer = stream.writer(io, &write_scratch); + + var msg_buf: std.ArrayListUnmanaged(u8) = .empty; + defer msg_buf.deinit(server.gpa); + var out_buf: std.ArrayListUnmanaged(u8) = .empty; + defer out_buf.deinit(server.gpa); + + var reply_request_id: u32 = 1; + var ctx = commands.Context{ + .gpa = server.gpa, + .io = io, + .oid_gen = &server.oid_gen, + .connection_id = connection_id, + .client_desc = "127.0.0.1:0", + .engine = server.engine, + .server_start = server.start_time, + }; + + while (true) { + var len_bytes: [4]u8 = undefined; + reader.interface.readSliceAll(&len_bytes) catch |err| switch (err) { + error.EndOfStream => return, // clean client disconnect + else => { + std.debug.print("mongo-light: read error on conn {d}: {s}\n", .{ connection_id, @errorName(err) }); + return; + }, + }; + const total: u32 = std.mem.readInt(u32, &len_bytes, .little); + if (total < 16 or total > wire.max_message_size) { + std.debug.print("mongo-light: bad message length {d} on conn {d}\n", .{ total, connection_id }); + return; + } + + msg_buf.clearRetainingCapacity(); + try msg_buf.ensureTotalCapacity(server.gpa, total); + msg_buf.items.len = total; + std.mem.writeInt(u32, msg_buf.items[0..4], total, .little); + reader.interface.readSliceAll(msg_buf.items[4..]) catch |err| { + std.debug.print("mongo-light: read error on conn {d}: {s}\n", .{ connection_id, @errorName(err) }); + return; + }; + + var msg = wire.Message.parse(server.gpa, msg_buf.items) catch |err| { + // Unparseable request: close the connection. + const op: i32 = if (msg_buf.items.len >= 16) std.mem.readInt(i32, msg_buf.items[12..16], .little) else 0; + std.debug.print("mongo-light: bad message on conn {d}: {s} (opCode {d})\n", .{ connection_id, @errorName(err), op }); + return; + }; + defer msg.deinit(); + + var reply = wire.Reply.init(server.gpa); + defer reply.deinit(); + + commands.dispatch(&ctx, &msg, &reply) catch { + // Discard any partial reply (the client would read the first + // ok field, which may already say 1) and send a clean error. + reply.pairs.clearRetainingCapacity(); + reply.put_error( + @intFromEnum(commands.ErrorCode.internal_error), + "InternalError", + "internal error", + ) catch return; + }; + + out_buf.clearRetainingCapacity(); + if (msg.op_code == wire.op_code_query) { + wire.write_reply_query(server.gpa, reply_request_id, msg.request_id, reply.pairs.items, &out_buf) catch |err| { + std.debug.print("mongo-light: reply build error on conn {d}: {s}\n", .{ connection_id, @errorName(err) }); + return; + }; + } else { + reply.build(server.gpa, reply_request_id, msg.request_id, &out_buf) catch |err| { + std.debug.print("mongo-light: reply build error on conn {d}: {s}\n", .{ connection_id, @errorName(err) }); + return; + }; + } + reply_request_id +%= 1; + writer.interface.writeAll(out_buf.items) catch |err| { + std.debug.print("mongo-light: write error on conn {d}: {s}\n", .{ connection_id, @errorName(err) }); + return; + }; + writer.interface.flush() catch |err| { + std.debug.print("mongo-light: flush error on conn {d}: {s}\n", .{ connection_id, @errorName(err) }); + return; + }; + } +} diff --git a/src/storage.zig b/src/storage.zig new file mode 100644 index 0000000..f086ad3 --- /dev/null +++ b/src/storage.zig @@ -0,0 +1,329 @@ +//! Append-only record log. Each record is: +//! [0..4) u32 len — total record bytes +//! [4..8) u32 crc32 over bytes [8..len) +//! [8..16) u64 seq +//! [16] u8 type +//! [17..20) reserved +//! [20..) db\0 coll\0 bson doc +//! All integers little-endian. Reads and writes are positional, so the fd +//! offset never matters. A torn tail record (crash mid-append) is detected +//! during replay and skipped. + +const std = @import("std"); +const bson = @import("bson.zig"); + +pub const record_type_upsert: u8 = 1; +pub const record_type_delete: u8 = 2; + +pub const header_len: usize = 20; // len + crc + seq + type + reserved + +/// Largest record payload we will accept during replay. Matches the +/// announced maxBsonObjectSize with room for names and header. +pub const max_record_payload: usize = 16 * 1024 * 1024 + 64 * 1024; + +pub const Record = struct { + seq: u64, + type: u8, + db: []const u8, // transient: valid only during replay callback + coll: []const u8, + doc: []const u8, // raw bson bytes +}; + +pub const Error = error{ + InvalidLog, // corrupt interior record (bad CRC or impossible length) +}; + +/// Callback receives transient slices and a heap-allocated, freshly parsed +/// document. The callback takes ownership of the document (must deinit). +pub const ReplayFn = *const fn (ctx: *anyopaque, record: Record, doc: *bson.Document) anyerror!void; + +pub const Log = struct { + gpa: std.mem.Allocator, + io: std.Io, + file: std.Io.File, + path: []const u8, + end_pos: u64, + log_bytes: u64, // bytes written since the log was last rewritten + + pub fn open(gpa: std.mem.Allocator, io: std.Io, path: []const u8) !Log { + // Resolve to an absolute path so compaction can rename the file + // without depending on the caller's working directory. + const abs_path = blk: { + if (path.len > 0 and path[0] == '/') break :blk try gpa.dupe(u8, path); + const cwd = try std.process.currentPathAlloc(io, gpa); + defer gpa.free(cwd); + break :blk try std.fmt.allocPrint(gpa, "{s}/{s}", .{ cwd, path }); + }; + errdefer gpa.free(abs_path); + + const dir = std.Io.Dir.cwd(); + const file: std.Io.File = dir.openFile(io, abs_path, .{ .mode = .read_write }) catch |err| switch (err) { + error.FileNotFound => try dir.createFile(io, abs_path, .{ .read = true }), + else => return err, + }; + return .{ + .gpa = gpa, + .io = io, + .file = file, + .path = abs_path, + .end_pos = 0, + .log_bytes = 0, + }; + } + + pub fn close(self: *Log) void { + self.file.close(self.io); + self.gpa.free(self.path); + } + + /// Replay all valid records from the beginning of the file. + pub fn replay(self: *Log, ctx: *anyopaque, callback: ReplayFn) !void { + var chunk: [64 * 1024]u8 = undefined; + var heap_buf: []u8 = &.{}; + defer if (heap_buf.len > 0) self.gpa.free(heap_buf); + var pos: u64 = 0; + + while (true) { + const len_read = self.file.readPositionalAll(self.io, chunk[0..4], pos) catch |err| { + std.debug.print("mongo-light: log read error at {d}: {s}\n", .{ pos, @errorName(err) }); + return error.InvalidLog; + }; + if (len_read == 0) return; // clean end + if (len_read < 4) return; // torn tail + const total: u32 = std.mem.readInt(u32, chunk[0..4], .little); + if (total < header_len) { + std.debug.print("mongo-light: corrupt record length {d} at {d}\n", .{ total, pos }); + return error.InvalidLog; + } + const payload_len: usize = total - 4; + // Documents up to maxBsonObjectSize are legal; anything larger is + // corruption. Covers a hostile length prefix from a truncated file. + if (payload_len > max_record_payload) { + std.debug.print("mongo-light: record too large at {d}\n", .{pos}); + return error.InvalidLog; + } + const payload = if (payload_len <= chunk.len) chunk[0..payload_len] else blk: { + if (heap_buf.len < payload_len) { + if (heap_buf.len > 0) self.gpa.free(heap_buf); + heap_buf = try self.gpa.alloc(u8, payload_len); + } + break :blk heap_buf[0..payload_len]; + }; + const payload_read = self.file.readPositionalAll(self.io, payload, pos + 4) catch return error.InvalidLog; + if (payload_read < payload_len) return; // torn tail — crash during append + const crc_stored: u32 = std.mem.readInt(u32, payload[0..4], .little); + const crc_actual = std.hash.Crc32.hash(payload[4..payload_len]); + if (crc_stored != crc_actual) return error.InvalidLog; + + var idx: usize = 16; // after crc + seq + type + reserved + const seq: u64 = std.mem.readInt(u64, payload[4..12], .little); + const rtype = payload[12]; + const db = read_cstring(payload, &idx) orelse return error.InvalidLog; + const coll = read_cstring(payload, &idx) orelse return error.InvalidLog; + if (idx > payload_len) return error.InvalidLog; + const doc_bytes = payload[idx..payload_len]; + + const doc = try self.gpa.create(bson.Document); + doc.* = bson.Document.parse(self.gpa, doc_bytes) catch { + self.gpa.destroy(doc); + std.debug.print("mongo-light: unparseable doc in log at {d}\n", .{pos}); + return error.InvalidLog; + }; + try callback(ctx, .{ + .seq = seq, + .type = rtype, + .db = db, + .coll = coll, + .doc = doc_bytes, + }, doc); + + pos += total; + self.end_pos = pos; + } + } + + pub fn append_upsert(self: *Log, db: []const u8, coll: []const u8, doc: []const u8, seq: u64) !void { + try self.append(record_type_upsert, db, coll, doc, seq); + } + + pub fn append_delete(self: *Log, db: []const u8, coll: []const u8, doc: []const u8, seq: u64) !void { + try self.append(record_type_delete, db, coll, doc, seq); + } + + fn append(self: *Log, rtype: u8, db: []const u8, coll: []const u8, doc: []const u8, seq: u64) !void { + if (std.mem.indexOfScalar(u8, db, 0) != null or std.mem.indexOfScalar(u8, coll, 0) != null) { + return error.NulInName; + } + var buf: std.ArrayListUnmanaged(u8) = .empty; + defer buf.deinit(self.gpa); + try buf.appendNTimes(self.gpa, 0, header_len); + std.mem.writeInt(u64, buf.items[8..16], seq, .little); + buf.items[16] = rtype; + try buf.appendSlice(self.gpa, db); + try buf.append(self.gpa, 0); + try buf.appendSlice(self.gpa, coll); + try buf.append(self.gpa, 0); + try buf.appendSlice(self.gpa, doc); + if (buf.items.len > std.math.maxInt(u32)) return error.LogTooLarge; + const total: u32 = @intCast(buf.items.len); + std.mem.writeInt(u32, buf.items[0..4], total, .little); + std.mem.writeInt(u32, buf.items[4..8], std.hash.Crc32.hash(buf.items[8..]), .little); + try self.file.writePositionalAll(self.io, buf.items, self.end_pos); + self.end_pos += buf.items.len; + self.log_bytes += buf.items.len; + try self.file.sync(self.io); + } + + fn read_cstring(bytes: []const u8, idx: *usize) ?[]const u8 { + const start = idx.*; + while (idx.* < bytes.len and bytes[idx.*] != 0) idx.* += 1; + if (idx.* >= bytes.len) return null; + idx.* += 1; + return bytes[start .. idx.* - 1]; + } +}; + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +const testing = std.testing; + +const TmpLog = struct { + tmp: std.testing.TmpDir, + path: []u8, + + fn init(gpa: std.mem.Allocator) !TmpLog { + const tmp = std.testing.tmpDir(.{}); + const path = try std.fmt.allocPrint(gpa, ".zig-cache/tmp/{s}/test.log", .{tmp.sub_path}); + return .{ .tmp = tmp, .path = path }; + } + + fn deinit(self: *TmpLog, gpa: std.mem.Allocator) void { + self.tmp.cleanup(); + gpa.free(self.path); + } +}; + +test "append, replay, torn tail" { + var threaded: std.Io.Threaded = .init_single_threaded; + defer threaded.deinit(); + const io = threaded.io(); + const gpa = testing.allocator; + + var tmp = try TmpLog.init(gpa); + defer tmp.deinit(gpa); + var log = try Log.open(gpa, io, tmp.path); + defer log.close(); + + const doc_bytes = [_]u8{ + 0x0E, 0x00, 0x00, 0x00, // len 14 + 0x10, '_', 'i', 'd', 0, 0x2A, 0x00, 0x00, 0x00, // _id: 42 + 0x00, + }; + try log.append_upsert("db1", "coll1", &doc_bytes, 1); + try log.append_delete("db1", "coll1", &doc_bytes, 2); + + var seen: std.ArrayListUnmanaged(u8) = .empty; + defer seen.deinit(gpa); + const Ctx = struct { + seen: *std.ArrayListUnmanaged(u8), + gpa: std.mem.Allocator, + fn apply(ctx: *anyopaque, record: Record, doc: *bson.Document) anyerror!void { + const self: *@This() = @ptrCast(@alignCast(ctx)); + try self.seen.append(self.gpa, record.type); + try self.seen.append(self.gpa, @intCast(record.seq)); + doc.deinit(); + self.gpa.destroy(doc); + } + }; + var ctx = Ctx{ .seen = &seen, .gpa = gpa }; + try log.replay(@ptrCast(&ctx), Ctx.apply); + + try testing.expectEqualSlices(u8, &[_]u8{ record_type_upsert, 1, record_type_delete, 2 }, seen.items); + try testing.expectEqual(@as(u64, log.end_pos), log.end_pos); +} + +test "record larger than the read chunk replays" { + var threaded: std.Io.Threaded = .init_single_threaded; + defer threaded.deinit(); + const io = threaded.io(); + const gpa = testing.allocator; + + var tmp = try TmpLog.init(gpa); + defer tmp.deinit(gpa); + var log = try Log.open(gpa, io, tmp.path); + defer log.close(); + + // Build a doc whose payload pushes the record past the 64 KiB stack chunk. + const big = try gpa.alloc(u8, 80 * 1024); + defer gpa.free(big); + @memset(big, 'x'); + const pairs = [_]bson.Pair{ + .{ .key = "_id", .value = .{ .int32 = 1 } }, + .{ .key = "blob", .value = .{ .binary = .{ .subtype = 0, .data = big } } }, + }; + var out: std.ArrayListUnmanaged(u8) = .empty; + defer out.deinit(gpa); + try bson.write_doc(&pairs, gpa, &out); + try log.append_upsert("db", "big", out.items, 1); + + var count: usize = 0; + const Ctx = struct { + count: *usize, + gpa: std.mem.Allocator, + fn apply(ctx: *anyopaque, _: Record, doc: *bson.Document) anyerror!void { + const self: *@This() = @ptrCast(@alignCast(ctx)); + try testing.expectEqual(@as(usize, 80 * 1024), doc.get("blob").?.binary.data.len); + self.count.* += 1; + doc.deinit(); + self.gpa.destroy(doc); + } + }; + var ctx = Ctx{ .count = &count, .gpa = gpa }; + try log.replay(@ptrCast(&ctx), Ctx.apply); + try testing.expectEqual(@as(usize, 1), count); +} + +test "reject corrupt interior record" { + var threaded: std.Io.Threaded = .init_single_threaded; + defer threaded.deinit(); + const io = threaded.io(); + const gpa = testing.allocator; + + var tmp = try TmpLog.init(gpa); + defer tmp.deinit(gpa); + const path = tmp.path; + + var log = try Log.open(gpa, io, path); + const doc_bytes = [_]u8{ 0x0E, 0, 0, 0, 0x10, '_', 'i', 'd', 0, 42, 0, 0, 0, 0 }; + try log.append_upsert("db", "c", &doc_bytes, 1); + log.close(); + + // Corrupt the file: flip a byte in the middle of the record. + const dir = std.Io.Dir.cwd(); + var f = try dir.openFile(io, path, .{ .mode = .read_write }); + var buf: [64]u8 = undefined; + const n = try f.readPositionalAll(io, &buf, 0); + _ = n; + buf[25] ^= 0xFF; + try f.writePositionalAll(io, buf[0..64], 0); + f.close(io); + + var log2 = try Log.open(gpa, io, path); + defer log2.close(); + defer dir.deleteFile(io, path) catch {}; + var count: usize = 0; + const Ctx = struct { + count: *usize, + gpa: std.mem.Allocator, + fn apply(ctx: *anyopaque, _: Record, doc: *bson.Document) anyerror!void { + const self: *@This() = @ptrCast(@alignCast(ctx)); + self.count.* += 1; + doc.deinit(); + self.gpa.destroy(doc); + } + }; + var ctx = Ctx{ .count = &count, .gpa = gpa }; + try testing.expectError(error.InvalidLog, log2.replay(@ptrCast(&ctx), Ctx.apply)); +} diff --git a/src/update.zig b/src/update.zig new file mode 100644 index 0000000..fc4f4d0 --- /dev/null +++ b/src/update.zig @@ -0,0 +1,427 @@ +//! Update operators: $set, $unset, $inc, $push, $pull, $rename with +//! dot-path navigation. Mutates the document's pairs in place, allocating +//! from the document's own arena. + +const std = @import("std"); +const bson = @import("bson.zig"); +const query = @import("query.zig"); + +pub const UpdateError = error{ ImmutableId, InvalidUpdate, OutOfMemory }; + +const max_path_segments = 16; + +/// Apply an update document (whose fields are operator documents) to `doc`. +pub fn apply(doc: *bson.Document, update: *const bson.Document) UpdateError!void { + const arena = doc.arena.allocator(); + var pairs = try copy_pairs_to_list(arena, doc.pairs); + for (update.pairs) |op| { + if (op.key.len == 0 or op.key[0] != '$') return error.InvalidUpdate; + try apply_operator(arena, &pairs, op.key, op.value); + } + doc.pairs = try pairs.toOwnedSlice(arena); +} + +fn apply_operator(arena: std.mem.Allocator, pairs: *std.ArrayListUnmanaged(bson.Pair), op: []const u8, value: bson.Value) UpdateError!void { + if (std.mem.eql(u8, op, "$set")) { + const ops = doc_pairs(value) orelse return error.InvalidUpdate; + for (ops) |p| { + if (std.mem.eql(u8, p.key, "_id")) return error.ImmutableId; + var segs: [max_path_segments][]const u8 = undefined; + const n = split_path(p.key, &segs) orelse return error.InvalidUpdate; + try set_path(arena, pairs, segs[0..n], try bson.copy_value(arena, p.value)); + } + return; + } + if (std.mem.eql(u8, op, "$unset")) { + const ops = doc_pairs(value) orelse return error.InvalidUpdate; + for (ops) |p| { + var segs: [max_path_segments][]const u8 = undefined; + const n = split_path(p.key, &segs) orelse continue; + unset_path(arena, pairs, segs[0..n]); + } + return; + } + if (std.mem.eql(u8, op, "$inc")) { + const ops = doc_pairs(value) orelse return error.InvalidUpdate; + for (ops) |p| { + var segs: [max_path_segments][]const u8 = undefined; + const n = split_path(p.key, &segs) orelse return error.InvalidUpdate; + const current = get_value(pairs.items, segs[0..n]) orelse bson.Value{ .int32 = 0 }; + if (!current.is_number() or !p.value.is_number()) return error.InvalidUpdate; + const sum = try numeric_add(arena, current, p.value); + try set_path(arena, pairs, segs[0..n], sum); + } + return; + } + if (std.mem.eql(u8, op, "$push")) { + const ops = doc_pairs(value) orelse return error.InvalidUpdate; + for (ops) |p| { + var segs: [max_path_segments][]const u8 = undefined; + const n = split_path(p.key, &segs) orelse return error.InvalidUpdate; + const current_opt = get_value(pairs.items, segs[0..n]); + var items: std.ArrayListUnmanaged(bson.Value) = .empty; + defer items.deinit(arena); + if (current_opt) |current| { + switch (current) { + .array => |arr| try items.appendSlice(arena, arr), + .null => {}, + else => return error.InvalidUpdate, // non-array field + } + } + if (p.value == .doc) { + if (bson.get_pair(p.value.doc, "$each")) |each| { + const arr = switch (each) { + .array => |a| a, + else => return error.InvalidUpdate, + }; + for (arr) |item| try items.append(arena, try bson.copy_value(arena, item)); + try set_path(arena, pairs, segs[0..n], .{ .array = try items.toOwnedSlice(arena) }); + continue; + } + } + try items.append(arena, try bson.copy_value(arena, p.value)); + try set_path(arena, pairs, segs[0..n], .{ .array = try items.toOwnedSlice(arena) }); + } + return; + } + if (std.mem.eql(u8, op, "$pull")) { + const ops = doc_pairs(value) orelse return error.InvalidUpdate; + for (ops) |p| { + var segs: [max_path_segments][]const u8 = undefined; + const n = split_path(p.key, &segs) orelse return error.InvalidUpdate; + const current = get_value(pairs.items, segs[0..n]) orelse continue; + const arr = switch (current) { + .array => |a| a, + else => return error.InvalidUpdate, + }; + var items: std.ArrayListUnmanaged(bson.Value) = .empty; + defer items.deinit(arena); + for (arr) |elem| { + if (!pull_matches(arena, p.value, elem)) { + try items.append(arena, elem); + } + } + try set_path(arena, pairs, segs[0..n], .{ .array = try items.toOwnedSlice(arena) }); + } + return; + } + if (std.mem.eql(u8, op, "$rename")) { + const ops = doc_pairs(value) orelse return error.InvalidUpdate; + for (ops) |p| { + if (p.value != .string) return error.InvalidUpdate; + if (std.mem.eql(u8, p.key, "_id") or std.mem.eql(u8, p.value.string, "_id")) return error.ImmutableId; + var old_segs: [max_path_segments][]const u8 = undefined; + const old_n = split_path(p.key, &old_segs) orelse return error.InvalidUpdate; + const v = get_value(pairs.items, old_segs[0..old_n]) orelse continue; // no-op when absent + unset_path(arena, pairs, old_segs[0..old_n]); + var new_segs: [max_path_segments][]const u8 = undefined; + const new_n = split_path(p.value.string, &new_segs) orelse return error.InvalidUpdate; + try set_path(arena, pairs, new_segs[0..new_n], v); + } + return; + } + return error.InvalidUpdate; +} + +fn doc_pairs(v: bson.Value) ?[]const bson.Pair { + return switch (v) { + .doc => |pairs| pairs, + else => null, + }; +} + +fn split_path(path: []const u8, out: *[max_path_segments][]const u8) ?usize { + var n: usize = 0; + var it = std.mem.splitScalar(u8, path, '.'); + while (it.next()) |seg| { + if (n >= max_path_segments) return null; + out[n] = seg; + n += 1; + } + return n; +} + +fn parse_index(seg: []const u8) ?usize { + return std.fmt.parseInt(usize, seg, 10) catch null; +} + +fn copy_pairs_to_list(arena: std.mem.Allocator, pairs: []const bson.Pair) UpdateError!std.ArrayListUnmanaged(bson.Pair) { + var out: std.ArrayListUnmanaged(bson.Pair) = .empty; + errdefer out.deinit(arena); + try out.appendSlice(arena, pairs); + return out; +} + +fn find_pair(pairs: []const bson.Pair, key: []const u8) ?usize { + for (pairs, 0..) |p, i| { + if (std.mem.eql(u8, p.key, key)) return i; + } + return null; +} + +fn get_value(pairs: []const bson.Pair, segs: []const []const u8) ?bson.Value { + const idx = find_pair(pairs, segs[0]) orelse return null; + if (segs.len == 1) return pairs[idx].value; + return switch (pairs[idx].value) { + .doc => |sub| get_value(sub, segs[1..]), + .array => |arr| blk: { + const index = parse_index(segs[1]) orelse break :blk null; + if (index >= arr.len) break :blk null; + if (segs.len == 2) break :blk arr[index]; + break :blk switch (arr[index]) { + .doc => |sub| get_value(sub, segs[2..]), + else => null, + }; + }, + else => null, + }; +} + +fn set_path(arena: std.mem.Allocator, pairs: *std.ArrayListUnmanaged(bson.Pair), segs: []const []const u8, value: bson.Value) UpdateError!void { + if (segs.len == 1) { + if (find_pair(pairs.items, segs[0])) |idx| { + pairs.items[idx].value = value; + } else { + try pairs.append(arena, .{ .key = try arena.dupe(u8, segs[0]), .value = value }); + } + return; + } + const idx = find_pair(pairs.items, segs[0]) orelse { + const is_array = parse_index(segs[1]) != null; + try pairs.append(arena, .{ .key = try arena.dupe(u8, segs[0]), .value = if (is_array) .{ .array = &.{} } else .{ .doc = &.{} } }); + return set_path(arena, pairs, segs, value); + }; + switch (pairs.items[idx].value) { + .doc => |sub| { + var sub_pairs = try copy_pairs_to_list(arena, sub); + defer sub_pairs.deinit(arena); + try set_path(arena, &sub_pairs, segs[1..], value); + pairs.items[idx].value = .{ .doc = try sub_pairs.toOwnedSlice(arena) }; + }, + .array => |arr| { + const index = parse_index(segs[1]) orelse { + // treat as non-array: replace with a doc + var sub_pairs: std.ArrayListUnmanaged(bson.Pair) = .empty; + defer sub_pairs.deinit(arena); + try set_path(arena, &sub_pairs, segs[1..], value); + pairs.items[idx].value = .{ .doc = try sub_pairs.toOwnedSlice(arena) }; + return; + }; + var items = try copy_array_to_list(arena, arr); + defer items.deinit(arena); + if (index >= items.items.len) { + try items.appendNTimes(arena, .null, index + 1 - items.items.len); + } + if (segs.len == 2) { + items.items[index] = value; + } else { + switch (items.items[index]) { + .doc => |sub| { + var sub_pairs = try copy_pairs_to_list(arena, sub); + defer sub_pairs.deinit(arena); + try set_path(arena, &sub_pairs, segs[2..], value); + items.items[index] = .{ .doc = try sub_pairs.toOwnedSlice(arena) }; + }, + else => { + var sub_pairs: std.ArrayListUnmanaged(bson.Pair) = .empty; + defer sub_pairs.deinit(arena); + try set_path(arena, &sub_pairs, segs[2..], value); + items.items[index] = .{ .doc = try sub_pairs.toOwnedSlice(arena) }; + }, + } + } + pairs.items[idx].value = .{ .array = try items.toOwnedSlice(arena) }; + }, + else => { + var sub_pairs: std.ArrayListUnmanaged(bson.Pair) = .empty; + defer sub_pairs.deinit(arena); + try set_path(arena, &sub_pairs, segs[1..], value); + pairs.items[idx].value = .{ .doc = try sub_pairs.toOwnedSlice(arena) }; + }, + } +} + +fn copy_array_to_list(arena: std.mem.Allocator, arr: []const bson.Value) UpdateError!std.ArrayListUnmanaged(bson.Value) { + var out: std.ArrayListUnmanaged(bson.Value) = .empty; + errdefer out.deinit(arena); + try out.appendSlice(arena, arr); + return out; +} + +fn unset_path(arena: std.mem.Allocator, pairs: *std.ArrayListUnmanaged(bson.Pair), segs: []const []const u8) void { + if (segs.len == 1) { + if (find_pair(pairs.items, segs[0])) |idx| { + _ = pairs.orderedRemove(idx); + } + return; + } + const idx = find_pair(pairs.items, segs[0]) orelse return; + switch (pairs.items[idx].value) { + .doc => |sub| { + var sub_pairs = copy_pairs_to_list(arena, sub) catch return; + unset_path(arena, &sub_pairs, segs[1..]); + pairs.items[idx].value = .{ .doc = sub_pairs.items }; + }, + else => {}, + } +} + +fn pull_matches(arena: std.mem.Allocator, condition: bson.Value, elem: bson.Value) bool { + switch (condition) { + .doc => |cond_pairs| { + const elem_doc = switch (elem) { + .doc => |pairs| pairs, + else => return false, + }; + var all_operators = cond_pairs.len > 0; + for (cond_pairs) |p| { + if (p.key.len == 0 or p.key[0] != '$') { + all_operators = false; + break; + } + } + if (all_operators) { + // Operator condition against the element's value at each + // operator's field — treat element doc as the doc. + var ok = true; + for (cond_pairs) |p| { + const actuals = bson.get_pair(elem_doc, p.key[1..]); + const a: bson.Value = actuals orelse .null; + if (!(query.value_matches_operator(arena, p.key, p.value, a) catch false)) ok = false; + } + return ok; + } + return (query.matches(arena, &.{ .arena = undefined, .pairs = cond_pairs }, &.{ .arena = undefined, .pairs = elem_doc }) catch false); + }, + else => return bson.compare(condition, elem) == .eq, + } +} + +fn numeric_add(arena: std.mem.Allocator, a: bson.Value, b: bson.Value) UpdateError!bson.Value { + _ = arena; + if (a == .double or b == .double) { + const sum: f64 = @floatCast(a.as_f128() + b.as_f128()); + return .{ .double = sum }; + } + if (a == .int64 or b == .int64) { + const av: i64 = switch (a) { + .int32 => |i| i, + .int64 => |i| i, + else => unreachable, + }; + const bv: i64 = switch (b) { + .int32 => |i| i, + .int64 => |i| i, + else => unreachable, + }; + const sum = std.math.add(i64, av, bv) catch return error.InvalidUpdate; + return .{ .int64 = sum }; + } + const sum: i64 = @as(i64, a.int32) + b.int32; + if (sum >= std.math.minInt(i32) and sum <= std.math.maxInt(i32)) { + return .{ .int32 = @intCast(sum) }; + } + return .{ .int64 = sum }; +} + +// --------------------------------------------------------------------------- +// Tests +// --------------------------------------------------------------------------- + +const testing = std.testing; + +fn doc_of(pairs: []const bson.Pair) bson.Document { + return .{ .arena = undefined, .pairs = pairs }; +} + +test "$set, $inc, $unset, $rename" { + var doc = bson.Document{ .arena = std.heap.ArenaAllocator.init(testing.allocator), .pairs = &.{} }; + defer doc.arena.deinit(); + doc.pairs = try doc.arena.allocator().dupe(bson.Pair, &.{ + .{ .key = "a", .value = .{ .int32 = 1 } }, + .{ .key = "user", .value = .{ .doc = &.{ + .{ .key = "name", .value = .{ .string = "bob" } }, + .{ .key = "age", .value = .{ .int32 = 30 } }, + } } }, + .{ .key = "gone", .value = .{ .int32 = 9 } }, + }); + + try apply(&doc, &doc_of(&.{ + .{ .key = "$set", .value = .{ .doc = &.{ + .{ .key = "user.name", .value = .{ .string = "alice" } }, + .{ .key = "user.city", .value = .{ .string = "NYC" } }, + .{ .key = "new", .value = .{ .int32 = 5 } }, + } } }, + .{ .key = "$inc", .value = .{ .doc = &.{.{ .key = "user.age", .value = .{ .int32 = 2 } }} } }, + .{ .key = "$unset", .value = .{ .doc = &.{.{ .key = "gone", .value = .{ .string = "" } }} } }, + .{ .key = "$rename", .value = .{ .doc = &.{.{ .key = "new", .value = .{ .string = "renamed" } }} } }, + })); + + const user = bson.get_pair(doc.pairs, "user").?; + try testing.expectEqualStrings("alice", user.doc[0].value.string); + try testing.expectEqual(@as(i64, 32), bson.get_pair(user.doc, "age").?.int32); + try testing.expectEqualStrings("NYC", bson.get_pair(user.doc, "city").?.string); + try testing.expect(bson.get_pair(doc.pairs, "gone") == null); + try testing.expect(bson.get_pair(doc.pairs, "new") == null); + try testing.expectEqual(@as(i64, 5), bson.get_pair(doc.pairs, "renamed").?.int32); +} + +test "$push and $pull" { + var doc = bson.Document{ .arena = std.heap.ArenaAllocator.init(testing.allocator), .pairs = &.{} }; + defer doc.arena.deinit(); + doc.pairs = try doc.arena.allocator().dupe(bson.Pair, &.{ + .{ .key = "tags", .value = .{ .array = &.{ .{ .string = "a" }, .{ .string = "b" } } } }, + }); + + try apply(&doc, &doc_of(&.{ + .{ .key = "$push", .value = .{ .doc = &.{.{ .key = "tags", .value = .{ .string = "c" } }} } }, + })); + try testing.expectEqual(@as(usize, 3), bson.get_pair(doc.pairs, "tags").?.array.len); + try testing.expectEqualStrings("c", bson.get_pair(doc.pairs, "tags").?.array[2].string); + + try apply(&doc, &doc_of(&.{ + .{ .key = "$pull", .value = .{ .doc = &.{.{ .key = "tags", .value = .{ .string = "b" } }} } }, + })); + const tags = bson.get_pair(doc.pairs, "tags").?.array; + try testing.expectEqual(@as(usize, 2), tags.len); + try testing.expectEqualStrings("a", tags[0].string); + try testing.expectEqualStrings("c", tags[1].string); + + try apply(&doc, &doc_of(&.{ + .{ .key = "$push", .value = .{ .doc = &.{.{ .key = "tags", .value = .{ .doc = &.{.{ .key = "$each", .value = .{ .array = &.{ .{ .string = "x" }, .{ .string = "y" } } } }} } }} } }, + })); + try testing.expectEqual(@as(usize, 4), bson.get_pair(doc.pairs, "tags").?.array.len); +} + +test "$set nested creation and _id protection" { + var doc = bson.Document{ .arena = std.heap.ArenaAllocator.init(testing.allocator), .pairs = &.{} }; + defer doc.arena.deinit(); + doc.pairs = try doc.arena.allocator().dupe(bson.Pair, &.{ + .{ .key = "_id", .value = .{ .int32 = 1 } }, + }); + + try apply(&doc, &doc_of(&.{ + .{ .key = "$set", .value = .{ .doc = &.{ + .{ .key = "a.b.c", .value = .{ .int32 = 42 } }, + .{ .key = "arr.1", .value = .{ .string = "x" } }, + } } }, + })); + const a = bson.get_pair(doc.pairs, "a").?; + const b = bson.get_pair(a.doc, "b").?; + try testing.expectEqual(@as(i64, 42), bson.get_pair(b.doc, "c").?.int32); + const arr = bson.get_pair(doc.pairs, "arr").?.array; + try testing.expectEqual(@as(usize, 2), arr.len); + try testing.expectEqualStrings("x", arr[1].string); + + try testing.expectError(error.ImmutableId, apply(&doc, &doc_of(&.{ + .{ .key = "$set", .value = .{ .doc = &.{.{ .key = "_id", .value = .{ .int32 = 2 } }} } }, + }))); +} + +test "non-operator update rejected" { + var doc = bson.Document{ .arena = std.heap.ArenaAllocator.init(testing.allocator), .pairs = &.{} }; + defer doc.arena.deinit(); + try testing.expectError(error.InvalidUpdate, apply(&doc, &doc_of(&.{ + .{ .key = "plain", .value = .{ .int32 = 1 } }, + }))); +} diff --git a/src/wire.zig b/src/wire.zig new file mode 100644 index 0000000..63bbec6 --- /dev/null +++ b/src/wire.zig @@ -0,0 +1,341 @@ +//! 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()); +}