diff --git a/src/main.zig b/src/main.zig index faa5228..4b1d9f5 100644 --- a/src/main.zig +++ b/src/main.zig @@ -213,15 +213,15 @@ fn installSignalHandlers() void { fn handleHttpConn(stream: std.net.Stream, stats: *broadcaster.Stats, persist: *event_log_mod.DiskPersist, slurper: *slurper_mod.Slurper, ci: *collection_index_mod.CollectionIndex) void { defer stream.close(); - // read request (headers + body for small POSTs) + // read request — may need multiple reads if proxy splits headers/body var buf: [8192]u8 = undefined; - const n = stream.read(&buf) catch return; - if (n == 0) return; - const request = buf[0..n]; + var total: usize = 0; + total = stream.read(&buf) catch return; + if (total == 0) return; // parse first line: "METHOD /path HTTP/1.1" - const line_end = std.mem.indexOfScalar(u8, request, '\n') orelse return; - const first_line = request[0..line_end]; + const line_end = std.mem.indexOfScalar(u8, buf[0..total], '\n') orelse return; + const first_line = buf[0..line_end]; const method_end = std.mem.indexOfScalar(u8, first_line, ' ') orelse return; const method = first_line[0..method_end]; @@ -231,19 +231,54 @@ fn handleHttpConn(stream: std.net.Stream, stats: *broadcaster.Stats, persist: *e const path_end = std.mem.indexOfScalar(u8, rest, ' ') orelse rest.len; const path = rest[0..path_end]; - // find body (after \r\n\r\n) - const header_end = std.mem.indexOf(u8, request, "\r\n\r\n"); - const body: []const u8 = if (header_end) |he| request[he + 4 ..] else ""; + // find end of headers + const header_end = std.mem.indexOf(u8, buf[0..total], "\r\n\r\n"); + + // for POSTs: if we have headers, parse Content-Length and read remaining body + if (std.mem.eql(u8, method, "POST")) { + if (header_end) |he| { + const headers = buf[0..he]; + const content_length = parseContentLength(headers) orelse 0; + const body_start = he + 4; + const body_needed = body_start + content_length; + + // keep reading until we have the full body (or buffer is full) + while (total < body_needed and total < buf.len) { + const m = stream.read(buf[total..]) catch break; + if (m == 0) break; + total += m; + } + } + } + + const request = buf[0..total]; + // re-find header_end in case more data shifted things (it won't, but be safe) + const he = std.mem.indexOf(u8, request, "\r\n\r\n"); + const body: []const u8 = if (he) |h| request[h + 4 ..] else ""; if (std.mem.eql(u8, method, "GET")) { handleGet(stream, path, stats, persist, slurper, ci); } else if (std.mem.eql(u8, method, "POST")) { - handlePost(stream, path, request[0 .. header_end orelse n], body, persist, slurper); + handlePost(stream, path, request[0 .. he orelse total], body, persist, slurper); } else { httpRespond(stream, "405 Method Not Allowed", "text/plain", "method not allowed"); } } +fn parseContentLength(headers: []const u8) ?usize { + var iter = std.mem.splitScalar(u8, headers, '\n'); + while (iter.next()) |line| { + const trimmed = std.mem.trimRight(u8, line, "\r"); + const colon = std.mem.indexOfScalar(u8, trimmed, ':') orelse continue; + const key = std.mem.trim(u8, trimmed[0..colon], " "); + if (std.ascii.eqlIgnoreCase(key, "content-length")) { + const val = std.mem.trim(u8, trimmed[colon + 1 ..], " "); + return std.fmt.parseInt(usize, val, 10) catch null; + } + } + return null; +} + fn handleGet(stream: std.net.Stream, full_path: []const u8, stats: *broadcaster.Stats, persist: *event_log_mod.DiskPersist, slurper: *slurper_mod.Slurper, ci: *collection_index_mod.CollectionIndex) void { // split path from query string const qmark = std.mem.indexOfScalar(u8, full_path, '?');