authorgravatar for andrew@ziglang.orgAndrew Kelley <andrew@ziglang.org> 2025-06-26 16:58:58-07:00
committergravatar for andrew@ziglang.orgAndrew Kelley <andrew@ziglang.org> 2025-07-01 16:35:30-07:00
log5b5243b5b7f5d01267afa1241c5847b3507694c1
treee124010f48228e8c2b7cb5c9146d2b50e92479ee
parent9f8486170c5a5dda56381098cf10ecf8c6364d9b

std: fix drain bugs in Writer.Allocating and net.Stream


2 files changed, 47 insertions(+), 35 deletions(-)

lib/std/io/Writer.zig+7-4
...@@ -2172,6 +2172,7 @@ pub const Allocating = struct {...@@ -2172,6 +2172,7 @@ pub const Allocating = struct {
2172 }2172 }
21732173
2174 fn drain(w: *Writer, data: []const []const u8, splat: usize) Error!usize {2174 fn drain(w: *Writer, data: []const []const u8, splat: usize) Error!usize {
2175 if (data.len == 0) return 0; // flush
2175 const a: *Allocating = @fieldParentPtr("interface", w);2176 const a: *Allocating = @fieldParentPtr("interface", w);
2176 const gpa = a.allocator;2177 const gpa = a.allocator;
2177 const pattern = data[data.len - 1];2178 const pattern = data[data.len - 1];
...@@ -2179,14 +2180,16 @@ pub const Allocating = struct {...@@ -2179,14 +2180,16 @@ pub const Allocating = struct {
2179 var list = a.toArrayList();2180 var list = a.toArrayList();
2180 defer setArrayList(a, list);2181 defer setArrayList(a, list);
2181 const start_len = list.items.len;2182 const start_len = list.items.len;
2182 for (data[0 .. data.len - 1]) |bytes| {2183 for (data) |bytes| {
2183 list.ensureUnusedCapacity(gpa, bytes.len + splat_len) catch return error.WriteFailed;2184 list.ensureUnusedCapacity(gpa, bytes.len + splat_len) catch return error.WriteFailed;
2184 list.appendSliceAssumeCapacity(bytes);2185 list.appendSliceAssumeCapacity(bytes);
2185 }2186 }
2186 switch (pattern.len) {2187 if (splat == 0) {
2188 list.items.len -= pattern.len;
2189 } else switch (pattern.len) {
2187 0 => {},2190 0 => {},
2188 1 => list.appendNTimesAssumeCapacity(pattern[0], splat),2191 1 => list.appendNTimesAssumeCapacity(pattern[0], splat - 1),
2189 else => for (0..splat) |_| list.appendSliceAssumeCapacity(pattern),2192 else => for (0..splat - 1) |_| list.appendSliceAssumeCapacity(pattern),
2190 }2193 }
2191 return list.items.len - start_len;2194 return list.items.len - start_len;
2192 }2195 }
lib/std/net.zig+40-31
...@@ -2104,7 +2104,6 @@ pub const Stream = struct {...@@ -2104,7 +2104,6 @@ pub const Stream = struct {
2104 fn drain(io_w: *io.Writer, data: []const []const u8, splat: usize) io.Writer.Error!usize {2104 fn drain(io_w: *io.Writer, data: []const []const u8, splat: usize) io.Writer.Error!usize {
2105 const w: *Writer = @fieldParentPtr("interface", io_w);2105 const w: *Writer = @fieldParentPtr("interface", io_w);
2106 const buffered = io_w.buffered();2106 const buffered = io_w.buffered();
2107 var splat_buffer: [splat_buffer_len]u8 = undefined;
2108 var iovecs: [max_buffers_len]std.posix.iovec_const = undefined;2107 var iovecs: [max_buffers_len]std.posix.iovec_const = undefined;
2109 var msg: posix.msghdr_const = msg: {2108 var msg: posix.msghdr_const = msg: {
2110 var i: usize = 0;2109 var i: usize = 0;
...@@ -2115,7 +2114,7 @@ pub const Stream = struct {...@@ -2115,7 +2114,7 @@ pub const Stream = struct {
2115 };2114 };
2116 i += 1;2115 i += 1;
2117 }2116 }
2118 for (data[0..data.len]) |bytes| {2117 for (data) |bytes| {
2119 // OS checks ptr addr before length so zero length vectors must be omitted.2118 // OS checks ptr addr before length so zero length vectors must be omitted.
2120 if (bytes.len == 0) continue;2119 if (bytes.len == 0) continue;
2121 iovecs[i] = .{2120 iovecs[i] = .{
...@@ -2135,38 +2134,48 @@ pub const Stream = struct {...@@ -2135,38 +2134,48 @@ pub const Stream = struct {
2135 .flags = 0,2134 .flags = 0,
2136 };2135 };
2137 };2136 };
2138 const pattern = data[data.len - 1];2137 if (data.len != 0) {
2139 switch (splat) {2138 const pattern = data[data.len - 1];
2140 0 => msg.iovlen -= 1,2139 switch (splat) {
2141 1 => {},2140 0 => if (msg.iovlen != 0 and iovecs[msg.iovlen - 1].base == data[data.len - 1].ptr) {
2142 else => switch (pattern.len) {2141 msg.iovlen -= 1;
2143 0 => {},2142 },
2144 1 => {2143 1 => {},
2145 // Replace the 1-byte buffer with a bigger one.2144 else => switch (pattern.len) {
2146 const memset_len = @min(splat_buffer.len, splat);2145 0 => {},
2147 const buf = splat_buffer[0..memset_len];2146 1 => memset: {
2148 @memset(buf, pattern[0]);2147 // Replace the 1-byte buffer with a bigger one.
2149 iovecs[msg.iovlen - 1] = .{ .base = buf.ptr, .len = buf.len };2148 if (msg.iovlen != 0 and iovecs[msg.iovlen - 1].base == data[data.len - 1].ptr)
2150 var remaining_splat = splat - buf.len;2149 msg.iovlen -= 1;
2151 while (remaining_splat > splat_buffer.len and msg.iovlen < iovecs.len) {2150 if (iovecs.len - msg.iovlen == 0) break :memset;
2152 iovecs[msg.iovlen] = .{ .base = &splat_buffer, .len = splat_buffer.len };2151 const splat_buffer = io_w.buffer[io_w.end..];
2153 remaining_splat -= splat_buffer.len;2152 const memset_len = @min(splat_buffer.len, splat);
2153 const buf = splat_buffer[0..memset_len];
2154 @memset(buf, pattern[0]);
2155 iovecs[msg.iovlen] = .{ .base = buf.ptr, .len = buf.len };
2154 msg.iovlen += 1;2156 msg.iovlen += 1;
2155 }2157 var remaining_splat = splat - buf.len;
2156 if (remaining_splat > 0 and msg.iovlen < iovecs.len) {2158 while (remaining_splat > splat_buffer.len and iovecs.len - msg.iovlen != 0) {
2157 iovecs[msg.iovlen] = .{ .base = &splat_buffer, .len = remaining_splat };2159 assert(buf.len == splat_buffer.len);
2160 iovecs[msg.iovlen] = .{ .base = splat_buffer.ptr, .len = splat_buffer.len };
2161 msg.iovlen += 1;
2162 remaining_splat -= splat_buffer.len;
2163 }
2164 if (remaining_splat > 0 and iovecs.len - msg.iovlen != 0) {
2165 iovecs[msg.iovlen] = .{ .base = splat_buffer.ptr, .len = remaining_splat };
2166 msg.iovlen += 1;
2167 }
2168 },
2169 else => for (0..splat - 1) |_| {
2170 if (iovecs.len - msg.iovlen == 0) break;
2171 iovecs[msg.iovlen] = .{
2172 .base = pattern.ptr,
2173 .len = pattern.len,
2174 };
2158 msg.iovlen += 1;2175 msg.iovlen += 1;
2159 }2176 },
2160 },2177 },
2161 else => for (0..splat - 1) |_| {2178 }
2162 if (iovecs.len - msg.iovlen == 0) break;
2163 iovecs[msg.iovlen] = .{
2164 .base = pattern.ptr,
2165 .len = pattern.len,
2166 };
2167 msg.iovlen += 1;
2168 },
2169 },
2170 }2179 }
2171 const flags = posix.MSG.NOSIGNAL;2180 const flags = posix.MSG.NOSIGNAL;
2172 return io_w.consume(std.posix.sendmsg(w.file_writer.file.handle, &msg, flags) catch |err| {2181 return io_w.consume(std.posix.sendmsg(w.file_writer.file.handle, &msg, flags) catch |err| {