authorgravatar for andrew@ziglang.orgAndrew Kelley <andrew@ziglang.org> 2025-08-07 19:54:25-07:00
committergravatar for noreply@github.comGitHub <noreply@github.com> 2025-08-07 19:54:25-07:00
log5998a8cebe3973d70c258b2a1440c5c3252d3539
treea889e12970e72420c07c9e988536997479d16a7c
parent2cf15bee0325321e9da496580b55310d4ba1053f
parent46b34949c3b40d286233c98c97bd2e0c221c1518
signaturebadge-check Signed by PGP key B5690EEEBB952194

Merge pull request #24698 from ziglang/http

std: rework HTTP and TLS for new I/O API

31 files changed, 3743 insertions(+), 5295 deletions(-)

lib/compiler/resinator/cli.zig+7-5
...@@ -1141,6 +1141,8 @@ pub fn parse(allocator: Allocator, args: []const []const u8, diagnostics: *Diagn...@@ -1141,6 +1141,8 @@ pub fn parse(allocator: Allocator, args: []const []const u8, diagnostics: *Diagn
1141 }1141 }
1142 output_format = .res;1142 output_format = .res;
1143 }1143 }
1144 } else {
1145 output_format_source = .output_format_arg;
1144 }1146 }
1145 options.output_source = .{ .filename = try filepathWithExtension(allocator, options.input_source.filename, output_format.?.extension()) };1147 options.output_source = .{ .filename = try filepathWithExtension(allocator, options.input_source.filename, output_format.?.extension()) };
1146 } else {1148 } else {
...@@ -1529,21 +1531,21 @@ fn testParseOutput(args: []const []const u8, expected_output: []const u8) !?Opti...@@ -1529,21 +1531,21 @@ fn testParseOutput(args: []const []const u8, expected_output: []const u8) !?Opti
1529 var diagnostics = Diagnostics.init(std.testing.allocator);1531 var diagnostics = Diagnostics.init(std.testing.allocator);
1530 defer diagnostics.deinit();1532 defer diagnostics.deinit();
15311533
1532 var output = std.ArrayList(u8).init(std.testing.allocator);1534 var output: std.io.Writer.Allocating = .init(std.testing.allocator);
1533 defer output.deinit();1535 defer output.deinit();
15341536
1535 var options = parse(std.testing.allocator, args, &diagnostics) catch |err| switch (err) {1537 var options = parse(std.testing.allocator, args, &diagnostics) catch |err| switch (err) {
1536 error.ParseError => {1538 error.ParseError => {
1537 try diagnostics.renderToWriter(args, output.writer(), .no_color);1539 try diagnostics.renderToWriter(args, &output.writer, .no_color);
1538 try std.testing.expectEqualStrings(expected_output, output.items);1540 try std.testing.expectEqualStrings(expected_output, output.getWritten());
1539 return null;1541 return null;
1540 },1542 },
1541 else => |e| return e,1543 else => |e| return e,
1542 };1544 };
1543 errdefer options.deinit();1545 errdefer options.deinit();
15441546
1545 try diagnostics.renderToWriter(args, output.writer(), .no_color);1547 try diagnostics.renderToWriter(args, &output.writer, .no_color);
1546 try std.testing.expectEqualStrings(expected_output, output.items);1548 try std.testing.expectEqualStrings(expected_output, output.getWritten());
1547 return options;1549 return options;
1548}1550}
15491551
lib/compiler/resinator/compile.zig+50-48
...@@ -550,7 +550,7 @@ pub const Compiler = struct {...@@ -550,7 +550,7 @@ pub const Compiler = struct {
550 // so get it here to simplify future usage.550 // so get it here to simplify future usage.
551 const filename_token = node.filename.getFirstToken();551 const filename_token = node.filename.getFirstToken();
552552
553 const file = self.searchForFile(filename_utf8) catch |err| switch (err) {553 const file_handle = self.searchForFile(filename_utf8) catch |err| switch (err) {
554 error.OutOfMemory => |e| return e,554 error.OutOfMemory => |e| return e,
555 else => |e| {555 else => |e| {
556 const filename_string_index = try self.diagnostics.putString(filename_utf8);556 const filename_string_index = try self.diagnostics.putString(filename_utf8);
...@@ -564,13 +564,15 @@ pub const Compiler = struct {...@@ -564,13 +564,15 @@ pub const Compiler = struct {
564 });564 });
565 },565 },
566 };566 };
567 defer file.close();567 defer file_handle.close();
568 var file_buffer: [2048]u8 = undefined;
569 var file_reader = file_handle.reader(&file_buffer);
568570
569 if (maybe_predefined_type) |predefined_type| {571 if (maybe_predefined_type) |predefined_type| {
570 switch (predefined_type) {572 switch (predefined_type) {
571 .GROUP_ICON, .GROUP_CURSOR => {573 .GROUP_ICON, .GROUP_CURSOR => {
572 // Check for animated icon first574 // Check for animated icon first
573 if (ani.isAnimatedIcon(file.deprecatedReader())) {575 if (ani.isAnimatedIcon(file_reader.interface.adaptToOldInterface())) {
574 // Animated icons are just put into the resource unmodified,576 // Animated icons are just put into the resource unmodified,
575 // and the resource type changes to ANIICON/ANICURSOR577 // and the resource type changes to ANIICON/ANICURSOR
576578
...@@ -582,18 +584,18 @@ pub const Compiler = struct {...@@ -582,18 +584,18 @@ pub const Compiler = struct {
582 header.type_value.ordinal = @intFromEnum(new_predefined_type);584 header.type_value.ordinal = @intFromEnum(new_predefined_type);
583 header.memory_flags = MemoryFlags.defaults(new_predefined_type);585 header.memory_flags = MemoryFlags.defaults(new_predefined_type);
584 header.applyMemoryFlags(node.common_resource_attributes, self.source);586 header.applyMemoryFlags(node.common_resource_attributes, self.source);
585 header.data_size = @intCast(try file.getEndPos());587 header.data_size = @intCast(try file_reader.getSize());
586588
587 try header.write(writer, self.errContext(node.id));589 try header.write(writer, self.errContext(node.id));
588 try file.seekTo(0);590 try file_reader.seekTo(0);
589 try writeResourceData(writer, file.deprecatedReader(), header.data_size);591 try writeResourceData(writer, &file_reader.interface, header.data_size);
590 return;592 return;
591 }593 }
592594
593 // isAnimatedIcon moved the file cursor so reset to the start595 // isAnimatedIcon moved the file cursor so reset to the start
594 try file.seekTo(0);596 try file_reader.seekTo(0);
595597
596 const icon_dir = ico.read(self.allocator, file.deprecatedReader(), try file.getEndPos()) catch |err| switch (err) {598 const icon_dir = ico.read(self.allocator, file_reader.interface.adaptToOldInterface(), try file_reader.getSize()) catch |err| switch (err) {
597 error.OutOfMemory => |e| return e,599 error.OutOfMemory => |e| return e,
598 else => |e| {600 else => |e| {
599 return self.iconReadError(601 return self.iconReadError(
...@@ -671,15 +673,15 @@ pub const Compiler = struct {...@@ -671,15 +673,15 @@ pub const Compiler = struct {
671 try writer.writeInt(u16, entry.type_specific_data.cursor.hotspot_y, .little);673 try writer.writeInt(u16, entry.type_specific_data.cursor.hotspot_y, .little);
672 }674 }
673675
674 try file.seekTo(entry.data_offset_from_start_of_file);676 try file_reader.seekTo(entry.data_offset_from_start_of_file);
675 var header_bytes = file.deprecatedReader().readBytesNoEof(16) catch {677 var header_bytes = (file_reader.interface.takeArray(16) catch {
676 return self.iconReadError(678 return self.iconReadError(
677 error.UnexpectedEOF,679 error.UnexpectedEOF,
678 filename_utf8,680 filename_utf8,
679 filename_token,681 filename_token,
680 predefined_type,682 predefined_type,
681 );683 );
682 };684 }).*;
683685
684 const image_format = ico.ImageFormat.detect(&header_bytes);686 const image_format = ico.ImageFormat.detect(&header_bytes);
685 if (!image_format.validate(&header_bytes)) {687 if (!image_format.validate(&header_bytes)) {
...@@ -802,8 +804,8 @@ pub const Compiler = struct {...@@ -802,8 +804,8 @@ pub const Compiler = struct {
802 },804 },
803 }805 }
804806
805 try file.seekTo(entry.data_offset_from_start_of_file);807 try file_reader.seekTo(entry.data_offset_from_start_of_file);
806 try writeResourceDataNoPadding(writer, file.deprecatedReader(), entry.data_size_in_bytes);808 try writeResourceDataNoPadding(writer, &file_reader.interface, entry.data_size_in_bytes);
807 try writeDataPadding(writer, full_data_size);809 try writeDataPadding(writer, full_data_size);
808810
809 if (self.state.icon_id == std.math.maxInt(u16)) {811 if (self.state.icon_id == std.math.maxInt(u16)) {
...@@ -857,9 +859,9 @@ pub const Compiler = struct {...@@ -857,9 +859,9 @@ pub const Compiler = struct {
857 },859 },
858 .BITMAP => {860 .BITMAP => {
859 header.applyMemoryFlags(node.common_resource_attributes, self.source);861 header.applyMemoryFlags(node.common_resource_attributes, self.source);
860 const file_size = try file.getEndPos();862 const file_size = try file_reader.getSize();
861863
862 const bitmap_info = bmp.read(file.deprecatedReader(), file_size) catch |err| {864 const bitmap_info = bmp.read(file_reader.interface.adaptToOldInterface(), file_size) catch |err| {
863 const filename_string_index = try self.diagnostics.putString(filename_utf8);865 const filename_string_index = try self.diagnostics.putString(filename_utf8);
864 return self.addErrorDetailsAndFail(.{866 return self.addErrorDetailsAndFail(.{
865 .err = .bmp_read_error,867 .err = .bmp_read_error,
...@@ -921,18 +923,17 @@ pub const Compiler = struct {...@@ -921,18 +923,17 @@ pub const Compiler = struct {
921923
922 header.data_size = bmp_bytes_to_write;924 header.data_size = bmp_bytes_to_write;
923 try header.write(writer, self.errContext(node.id));925 try header.write(writer, self.errContext(node.id));
924 try file.seekTo(bmp.file_header_len);926 try file_reader.seekTo(bmp.file_header_len);
925 const file_reader = file.deprecatedReader();927 try writeResourceDataNoPadding(writer, &file_reader.interface, bitmap_info.dib_header_size);
926 try writeResourceDataNoPadding(writer, file_reader, bitmap_info.dib_header_size);
927 if (bitmap_info.getBitmasksByteLen() > 0) {928 if (bitmap_info.getBitmasksByteLen() > 0) {
928 try writeResourceDataNoPadding(writer, file_reader, bitmap_info.getBitmasksByteLen());929 try writeResourceDataNoPadding(writer, &file_reader.interface, bitmap_info.getBitmasksByteLen());
929 }930 }
930 if (bitmap_info.getExpectedPaletteByteLen() > 0) {931 if (bitmap_info.getExpectedPaletteByteLen() > 0) {
931 try writeResourceDataNoPadding(writer, file_reader, @intCast(bitmap_info.getActualPaletteByteLen()));932 try writeResourceDataNoPadding(writer, &file_reader.interface, @intCast(bitmap_info.getActualPaletteByteLen()));
932 }933 }
933 try file.seekTo(bitmap_info.pixel_data_offset);934 try file_reader.seekTo(bitmap_info.pixel_data_offset);
934 const pixel_bytes: u32 = @intCast(file_size - bitmap_info.pixel_data_offset);935 const pixel_bytes: u32 = @intCast(file_size - bitmap_info.pixel_data_offset);
935 try writeResourceDataNoPadding(writer, file_reader, pixel_bytes);936 try writeResourceDataNoPadding(writer, &file_reader.interface, pixel_bytes);
936 try writeDataPadding(writer, bmp_bytes_to_write);937 try writeDataPadding(writer, bmp_bytes_to_write);
937 return;938 return;
938 },939 },
...@@ -956,7 +957,7 @@ pub const Compiler = struct {...@@ -956,7 +957,7 @@ pub const Compiler = struct {
956 return;957 return;
957 }958 }
958 header.applyMemoryFlags(node.common_resource_attributes, self.source);959 header.applyMemoryFlags(node.common_resource_attributes, self.source);
959 const file_size = try file.getEndPos();960 const file_size = try file_reader.getSize();
960 if (file_size > std.math.maxInt(u32)) {961 if (file_size > std.math.maxInt(u32)) {
961 return self.addErrorDetailsAndFail(.{962 return self.addErrorDetailsAndFail(.{
962 .err = .resource_data_size_exceeds_max,963 .err = .resource_data_size_exceeds_max,
...@@ -968,8 +969,9 @@ pub const Compiler = struct {...@@ -968,8 +969,9 @@ pub const Compiler = struct {
968 header.data_size = @intCast(file_size);969 header.data_size = @intCast(file_size);
969 try header.write(writer, self.errContext(node.id));970 try header.write(writer, self.errContext(node.id));
970971
971 var header_slurping_reader = headerSlurpingReader(148, file.deprecatedReader());972 var header_slurping_reader = headerSlurpingReader(148, file_reader.interface.adaptToOldInterface());
972 try writeResourceData(writer, header_slurping_reader.reader(), header.data_size);973 var adapter = header_slurping_reader.reader().adaptToNewApi(&.{});
974 try writeResourceData(writer, &adapter.new_interface, header.data_size);
973975
974 try self.state.font_dir.add(self.arena, FontDir.Font{976 try self.state.font_dir.add(self.arena, FontDir.Font{
975 .id = header.name_value.ordinal,977 .id = header.name_value.ordinal,
...@@ -992,7 +994,7 @@ pub const Compiler = struct {...@@ -992,7 +994,7 @@ pub const Compiler = struct {
992 }994 }
993995
994 // Fallback to just writing out the entire contents of the file996 // Fallback to just writing out the entire contents of the file
995 const data_size = try file.getEndPos();997 const data_size = try file_reader.getSize();
996 if (data_size > std.math.maxInt(u32)) {998 if (data_size > std.math.maxInt(u32)) {
997 return self.addErrorDetailsAndFail(.{999 return self.addErrorDetailsAndFail(.{
998 .err = .resource_data_size_exceeds_max,1000 .err = .resource_data_size_exceeds_max,
...@@ -1002,7 +1004,7 @@ pub const Compiler = struct {...@@ -1002,7 +1004,7 @@ pub const Compiler = struct {
1002 // We now know that the data size will fit in a u321004 // We now know that the data size will fit in a u32
1003 header.data_size = @intCast(data_size);1005 header.data_size = @intCast(data_size);
1004 try header.write(writer, self.errContext(node.id));1006 try header.write(writer, self.errContext(node.id));
1005 try writeResourceData(writer, file.deprecatedReader(), header.data_size);1007 try writeResourceData(writer, &file_reader.interface, header.data_size);
1006 }1008 }
10071009
1008 fn iconReadError(1010 fn iconReadError(
...@@ -1250,8 +1252,8 @@ pub const Compiler = struct {...@@ -1250,8 +1252,8 @@ pub const Compiler = struct {
1250 const data_len: u32 = @intCast(data_buffer.items.len);1252 const data_len: u32 = @intCast(data_buffer.items.len);
1251 try self.writeResourceHeader(writer, node.id, node.type, data_len, node.common_resource_attributes, self.state.language);1253 try self.writeResourceHeader(writer, node.id, node.type, data_len, node.common_resource_attributes, self.state.language);
12521254
1253 var data_fbs = std.io.fixedBufferStream(data_buffer.items);1255 var data_fbs: std.Io.Reader = .fixed(data_buffer.items);
1254 try writeResourceData(writer, data_fbs.reader(), data_len);1256 try writeResourceData(writer, &data_fbs, data_len);
1255 }1257 }
12561258
1257 pub fn writeResourceHeader(self: *Compiler, writer: anytype, id_token: Token, type_token: Token, data_size: u32, common_resource_attributes: []Token, language: res.Language) !void {1259 pub fn writeResourceHeader(self: *Compiler, writer: anytype, id_token: Token, type_token: Token, data_size: u32, common_resource_attributes: []Token, language: res.Language) !void {
...@@ -1266,15 +1268,15 @@ pub const Compiler = struct {...@@ -1266,15 +1268,15 @@ pub const Compiler = struct {
1266 try header.write(writer, self.errContext(id_token));1268 try header.write(writer, self.errContext(id_token));
1267 }1269 }
12681270
1269 pub fn writeResourceDataNoPadding(writer: anytype, data_reader: anytype, data_size: u32) !void {1271 pub fn writeResourceDataNoPadding(writer: anytype, data_reader: *std.Io.Reader, data_size: u32) !void {
1270 var limited_reader = std.io.limitedReader(data_reader, data_size);1272 var adapted = writer.adaptToNewApi();
12711273 var buffer: [128]u8 = undefined;
1272 const FifoBuffer = std.fifo.LinearFifo(u8, .{ .Static = 4096 });1274 adapted.new_interface.buffer = &buffer;
1273 var fifo = FifoBuffer.init();1275 try data_reader.streamExact(&adapted.new_interface, data_size);
1274 try fifo.pump(limited_reader.reader(), writer);1276 try adapted.new_interface.flush();
1275 }1277 }
12761278
1277 pub fn writeResourceData(writer: anytype, data_reader: anytype, data_size: u32) !void {1279 pub fn writeResourceData(writer: anytype, data_reader: *std.Io.Reader, data_size: u32) !void {
1278 try writeResourceDataNoPadding(writer, data_reader, data_size);1280 try writeResourceDataNoPadding(writer, data_reader, data_size);
1279 try writeDataPadding(writer, data_size);1281 try writeDataPadding(writer, data_size);
1280 }1282 }
...@@ -1339,8 +1341,8 @@ pub const Compiler = struct {...@@ -1339,8 +1341,8 @@ pub const Compiler = struct {
13391341
1340 try header.write(writer, self.errContext(node.id));1342 try header.write(writer, self.errContext(node.id));
13411343
1342 var data_fbs = std.io.fixedBufferStream(data_buffer.items);1344 var data_fbs: std.Io.Reader = .fixed(data_buffer.items);
1343 try writeResourceData(writer, data_fbs.reader(), data_size);1345 try writeResourceData(writer, &data_fbs, data_size);
1344 }1346 }
13451347
1346 /// Expects `data_writer` to be a LimitedWriter limited to u32, meaning all writes to1348 /// Expects `data_writer` to be a LimitedWriter limited to u32, meaning all writes to
...@@ -1732,8 +1734,8 @@ pub const Compiler = struct {...@@ -1732,8 +1734,8 @@ pub const Compiler = struct {
17321734
1733 try header.write(writer, self.errContext(node.id));1735 try header.write(writer, self.errContext(node.id));
17341736
1735 var data_fbs = std.io.fixedBufferStream(data_buffer.items);1737 var data_fbs: std.Io.Reader = .fixed(data_buffer.items);
1736 try writeResourceData(writer, data_fbs.reader(), data_size);1738 try writeResourceData(writer, &data_fbs, data_size);
1737 }1739 }
17381740
1739 fn writeDialogHeaderAndStrings(1741 fn writeDialogHeaderAndStrings(
...@@ -2046,8 +2048,8 @@ pub const Compiler = struct {...@@ -2046,8 +2048,8 @@ pub const Compiler = struct {
20462048
2047 try header.write(writer, self.errContext(node.id));2049 try header.write(writer, self.errContext(node.id));
20482050
2049 var data_fbs = std.io.fixedBufferStream(data_buffer.items);2051 var data_fbs: std.Io.Reader = .fixed(data_buffer.items);
2050 try writeResourceData(writer, data_fbs.reader(), data_size);2052 try writeResourceData(writer, &data_fbs, data_size);
2051 }2053 }
20522054
2053 /// Weight and italic carry over from previous FONT statements within a single resource,2055 /// Weight and italic carry over from previous FONT statements within a single resource,
...@@ -2121,8 +2123,8 @@ pub const Compiler = struct {...@@ -2121,8 +2123,8 @@ pub const Compiler = struct {
21212123
2122 try header.write(writer, self.errContext(node.id));2124 try header.write(writer, self.errContext(node.id));
21232125
2124 var data_fbs = std.io.fixedBufferStream(data_buffer.items);2126 var data_fbs: std.Io.Reader = .fixed(data_buffer.items);
2125 try writeResourceData(writer, data_fbs.reader(), data_size);2127 try writeResourceData(writer, &data_fbs, data_size);
2126 }2128 }
21272129
2128 /// Expects `data_writer` to be a LimitedWriter limited to u32, meaning all writes to2130 /// Expects `data_writer` to be a LimitedWriter limited to u32, meaning all writes to
...@@ -2386,8 +2388,8 @@ pub const Compiler = struct {...@@ -2386,8 +2388,8 @@ pub const Compiler = struct {
23862388
2387 try header.write(writer, self.errContext(node.id));2389 try header.write(writer, self.errContext(node.id));
23882390
2389 var data_fbs = std.io.fixedBufferStream(data_buffer.items);2391 var data_fbs: std.Io.Reader = .fixed(data_buffer.items);
2390 try writeResourceData(writer, data_fbs.reader(), data_size);2392 try writeResourceData(writer, &data_fbs, data_size);
2391 }2393 }
23922394
2393 /// Expects writer to be a LimitedWriter limited to u16, meaning all writes to2395 /// Expects writer to be a LimitedWriter limited to u16, meaning all writes to
...@@ -3321,8 +3323,8 @@ pub const StringTable = struct {...@@ -3321,8 +3323,8 @@ pub const StringTable = struct {
3321 // we fully control and know are numbers, so they have a fixed size.3323 // we fully control and know are numbers, so they have a fixed size.
3322 try header.writeAssertNoOverflow(writer);3324 try header.writeAssertNoOverflow(writer);
33233325
3324 var data_fbs = std.io.fixedBufferStream(data_buffer.items);3326 var data_fbs: std.Io.Reader = .fixed(data_buffer.items);
3325 try Compiler.writeResourceData(writer, data_fbs.reader(), data_size);3327 try Compiler.writeResourceData(writer, &data_fbs, data_size);
3326 }3328 }
3327 };3329 };
33283330
lib/compiler/resinator/cvtres.zig+24-29
...@@ -65,7 +65,7 @@ pub const ParseResOptions = struct {...@@ -65,7 +65,7 @@ pub const ParseResOptions = struct {
65};65};
6666
67/// The returned ParsedResources should be freed by calling its `deinit` function.67/// The returned ParsedResources should be freed by calling its `deinit` function.
68pub fn parseRes(allocator: Allocator, reader: anytype, options: ParseResOptions) !ParsedResources {68pub fn parseRes(allocator: Allocator, reader: *std.Io.Reader, options: ParseResOptions) !ParsedResources {
69 var resources = ParsedResources.init(allocator);69 var resources = ParsedResources.init(allocator);
70 errdefer resources.deinit();70 errdefer resources.deinit();
7171
...@@ -74,7 +74,7 @@ pub fn parseRes(allocator: Allocator, reader: anytype, options: ParseResOptions)...@@ -74,7 +74,7 @@ pub fn parseRes(allocator: Allocator, reader: anytype, options: ParseResOptions)
74 return resources;74 return resources;
75}75}
7676
77pub fn parseResInto(resources: *ParsedResources, reader: anytype, options: ParseResOptions) !void {77pub fn parseResInto(resources: *ParsedResources, reader: *std.Io.Reader, options: ParseResOptions) !void {
78 const allocator = resources.allocator;78 const allocator = resources.allocator;
79 var bytes_remaining: u64 = options.max_size;79 var bytes_remaining: u64 = options.max_size;
80 {80 {
...@@ -103,43 +103,38 @@ pub const ResourceAndSize = struct {...@@ -103,43 +103,38 @@ pub const ResourceAndSize = struct {
103 total_size: u64,103 total_size: u64,
104};104};
105105
106pub fn parseResource(allocator: Allocator, reader: anytype, max_size: u64) !ResourceAndSize {106pub fn parseResource(allocator: Allocator, reader: *std.Io.Reader, max_size: u64) !ResourceAndSize {
107 var header_counting_reader = std.io.countingReader(reader);107 const data_size = try reader.takeInt(u32, .little);
108 const header_reader = header_counting_reader.reader();108 const header_size = try reader.takeInt(u32, .little);
109 const data_size = try header_reader.readInt(u32, .little);
110 const header_size = try header_reader.readInt(u32, .little);
111 const total_size: u64 = @as(u64, header_size) + data_size;109 const total_size: u64 = @as(u64, header_size) + data_size;
112 if (total_size > max_size) return error.ImpossibleSize;110 if (total_size > max_size) return error.ImpossibleSize;
113111
114 var header_bytes_available = header_size -| 8;112 const remaining_header_bytes = try reader.take(header_size -| 8);
115 var type_reader = std.io.limitedReader(header_reader, header_bytes_available);113 var remaining_header_reader: std.Io.Reader = .fixed(remaining_header_bytes);
116 const type_value = try parseNameOrOrdinal(allocator, type_reader.reader());114 const type_value = try parseNameOrOrdinal(allocator, &remaining_header_reader);
117 errdefer type_value.deinit(allocator);115 errdefer type_value.deinit(allocator);
118116
119 header_bytes_available -|= @intCast(type_value.byteLen());117 const name_value = try parseNameOrOrdinal(allocator, &remaining_header_reader);
120 var name_reader = std.io.limitedReader(header_reader, header_bytes_available);
121 const name_value = try parseNameOrOrdinal(allocator, name_reader.reader());
122 errdefer name_value.deinit(allocator);118 errdefer name_value.deinit(allocator);
123119
124 const padding_after_name = numPaddingBytesNeeded(@intCast(header_counting_reader.bytes_read));120 const padding_after_name = numPaddingBytesNeeded(@intCast(remaining_header_reader.seek));
125 try header_reader.skipBytes(padding_after_name, .{ .buf_size = 3 });121 try remaining_header_reader.discardAll(padding_after_name);
126122
127 std.debug.assert(header_counting_reader.bytes_read % 4 == 0);123 std.debug.assert(remaining_header_reader.seek % 4 == 0);
128 const data_version = try header_reader.readInt(u32, .little);124 const data_version = try remaining_header_reader.takeInt(u32, .little);
129 const memory_flags: MemoryFlags = @bitCast(try header_reader.readInt(u16, .little));125 const memory_flags: MemoryFlags = @bitCast(try remaining_header_reader.takeInt(u16, .little));
130 const language: Language = @bitCast(try header_reader.readInt(u16, .little));126 const language: Language = @bitCast(try remaining_header_reader.takeInt(u16, .little));
131 const version = try header_reader.readInt(u32, .little);127 const version = try remaining_header_reader.takeInt(u32, .little);
132 const characteristics = try header_reader.readInt(u32, .little);128 const characteristics = try remaining_header_reader.takeInt(u32, .little);
133129
134 const header_bytes_read = header_counting_reader.bytes_read;130 if (remaining_header_reader.seek != remaining_header_reader.end) return error.HeaderSizeMismatch;
135 if (header_size != header_bytes_read) return error.HeaderSizeMismatch;
136131
137 const data = try allocator.alloc(u8, data_size);132 const data = try allocator.alloc(u8, data_size);
138 errdefer allocator.free(data);133 errdefer allocator.free(data);
139 try reader.readNoEof(data);134 try reader.readSliceAll(data);
140135
141 const padding_after_data = numPaddingBytesNeeded(@intCast(data_size));136 const padding_after_data = numPaddingBytesNeeded(@intCast(data_size));
142 try reader.skipBytes(padding_after_data, .{ .buf_size = 3 });137 try reader.discardAll(padding_after_data);
143138
144 return .{139 return .{
145 .resource = .{140 .resource = .{
...@@ -156,10 +151,10 @@ pub fn parseResource(allocator: Allocator, reader: anytype, max_size: u64) !Reso...@@ -156,10 +151,10 @@ pub fn parseResource(allocator: Allocator, reader: anytype, max_size: u64) !Reso
156 };151 };
157}152}
158153
159pub fn parseNameOrOrdinal(allocator: Allocator, reader: anytype) !NameOrOrdinal {154pub fn parseNameOrOrdinal(allocator: Allocator, reader: *std.Io.Reader) !NameOrOrdinal {
160 const first_code_unit = try reader.readInt(u16, .little);155 const first_code_unit = try reader.takeInt(u16, .little);
161 if (first_code_unit == 0xFFFF) {156 if (first_code_unit == 0xFFFF) {
162 const ordinal_value = try reader.readInt(u16, .little);157 const ordinal_value = try reader.takeInt(u16, .little);
163 return .{ .ordinal = ordinal_value };158 return .{ .ordinal = ordinal_value };
164 }159 }
165 var name_buf = try std.ArrayListUnmanaged(u16).initCapacity(allocator, 16);160 var name_buf = try std.ArrayListUnmanaged(u16).initCapacity(allocator, 16);
...@@ -167,7 +162,7 @@ pub fn parseNameOrOrdinal(allocator: Allocator, reader: anytype) !NameOrOrdinal...@@ -167,7 +162,7 @@ pub fn parseNameOrOrdinal(allocator: Allocator, reader: anytype) !NameOrOrdinal
167 var code_unit = first_code_unit;162 var code_unit = first_code_unit;
168 while (code_unit != 0) {163 while (code_unit != 0) {
169 try name_buf.append(allocator, std.mem.nativeToLittle(u16, code_unit));164 try name_buf.append(allocator, std.mem.nativeToLittle(u16, code_unit));
170 code_unit = try reader.readInt(u16, .little);165 code_unit = try reader.takeInt(u16, .little);
171 }166 }
172 return .{ .name = try name_buf.toOwnedSliceSentinel(allocator, 0) };167 return .{ .name = try name_buf.toOwnedSliceSentinel(allocator, 0) };
173}168}
lib/compiler/resinator/errors.zig+4-8
...@@ -1078,11 +1078,9 @@ const CorrespondingLines = struct {...@@ -1078,11 +1078,9 @@ const CorrespondingLines = struct {
1078 at_eof: bool = false,1078 at_eof: bool = false,
1079 span: SourceMappings.CorrespondingSpan,1079 span: SourceMappings.CorrespondingSpan,
1080 file: std.fs.File,1080 file: std.fs.File,
1081 buffered_reader: BufferedReaderType,1081 buffered_reader: std.fs.File.Reader,
1082 code_page: SupportedCodePage,1082 code_page: SupportedCodePage,
10831083
1084 const BufferedReaderType = std.io.BufferedReader(512, std.fs.File.DeprecatedReader);
1085
1086 pub fn init(cwd: std.fs.Dir, err_details: ErrorDetails, line_for_comparison: []const u8, corresponding_span: SourceMappings.CorrespondingSpan, corresponding_file: []const u8) !CorrespondingLines {1084 pub fn init(cwd: std.fs.Dir, err_details: ErrorDetails, line_for_comparison: []const u8, corresponding_span: SourceMappings.CorrespondingSpan, corresponding_file: []const u8) !CorrespondingLines {
1087 // We don't do line comparison for this error, so don't print the note if the line1085 // We don't do line comparison for this error, so don't print the note if the line
1088 // number is different1086 // number is different
...@@ -1101,9 +1099,7 @@ const CorrespondingLines = struct {...@@ -1101,9 +1099,7 @@ const CorrespondingLines = struct {
1101 .buffered_reader = undefined,1099 .buffered_reader = undefined,
1102 .code_page = err_details.code_page,1100 .code_page = err_details.code_page,
1103 };1101 };
1104 corresponding_lines.buffered_reader = BufferedReaderType{1102 corresponding_lines.buffered_reader = corresponding_lines.file.reader(&.{});
1105 .unbuffered_reader = corresponding_lines.file.deprecatedReader(),
1106 };
1107 errdefer corresponding_lines.deinit();1103 errdefer corresponding_lines.deinit();
11081104
1109 var fbs = std.io.fixedBufferStream(&corresponding_lines.line_buf);1105 var fbs = std.io.fixedBufferStream(&corresponding_lines.line_buf);
...@@ -1111,7 +1107,7 @@ const CorrespondingLines = struct {...@@ -1111,7 +1107,7 @@ const CorrespondingLines = struct {
11111107
1112 try corresponding_lines.writeLineFromStreamVerbatim(1108 try corresponding_lines.writeLineFromStreamVerbatim(
1113 writer,1109 writer,
1114 corresponding_lines.buffered_reader.reader(),1110 corresponding_lines.buffered_reader.interface.adaptToOldInterface(),
1115 corresponding_span.start_line,1111 corresponding_span.start_line,
1116 );1112 );
11171113
...@@ -1154,7 +1150,7 @@ const CorrespondingLines = struct {...@@ -1154,7 +1150,7 @@ const CorrespondingLines = struct {
11541150
1155 try self.writeLineFromStreamVerbatim(1151 try self.writeLineFromStreamVerbatim(
1156 writer,1152 writer,
1157 self.buffered_reader.reader(),1153 self.buffered_reader.interface.adaptToOldInterface(),
1158 self.line_num,1154 self.line_num,
1159 );1155 );
11601156
lib/compiler/resinator/ico.zig+2-1
...@@ -14,8 +14,9 @@ pub fn read(allocator: std.mem.Allocator, reader: anytype, max_size: u64) ReadEr...@@ -14,8 +14,9 @@ pub fn read(allocator: std.mem.Allocator, reader: anytype, max_size: u64) ReadEr
14 // Some Reader implementations have an empty ReadError error set which would14 // Some Reader implementations have an empty ReadError error set which would
15 // cause 'unreachable else' if we tried to use an else in the switch, so we15 // cause 'unreachable else' if we tried to use an else in the switch, so we
16 // need to detect this case and not try to translate to ReadError16 // need to detect this case and not try to translate to ReadError
17 const anyerror_reader_errorset = @TypeOf(reader).Error == anyerror;
17 const empty_reader_errorset = @typeInfo(@TypeOf(reader).Error).error_set == null or @typeInfo(@TypeOf(reader).Error).error_set.?.len == 0;18 const empty_reader_errorset = @typeInfo(@TypeOf(reader).Error).error_set == null or @typeInfo(@TypeOf(reader).Error).error_set.?.len == 0;
18 if (empty_reader_errorset) {19 if (empty_reader_errorset and !anyerror_reader_errorset) {
19 return readAnyError(allocator, reader, max_size) catch |err| switch (err) {20 return readAnyError(allocator, reader, max_size) catch |err| switch (err) {
20 error.EndOfStream => error.UnexpectedEOF,21 error.EndOfStream => error.UnexpectedEOF,
21 else => |e| return e,22 else => |e| return e,
lib/compiler/resinator/main.zig+2-2
...@@ -325,8 +325,8 @@ pub fn main() !void {...@@ -325,8 +325,8 @@ pub fn main() !void {
325 std.debug.assert(options.output_format == .coff);325 std.debug.assert(options.output_format == .coff);
326326
327 // TODO: Maybe use a buffered file reader instead of reading file into memory -> fbs327 // TODO: Maybe use a buffered file reader instead of reading file into memory -> fbs
328 var fbs = std.io.fixedBufferStream(res_data.bytes);328 var res_reader: std.Io.Reader = .fixed(res_data.bytes);
329 break :resources cvtres.parseRes(allocator, fbs.reader(), .{ .max_size = res_data.bytes.len }) catch |err| {329 break :resources cvtres.parseRes(allocator, &res_reader, .{ .max_size = res_data.bytes.len }) catch |err| {
330 // TODO: Better errors330 // TODO: Better errors
331 try error_handler.emitMessage(allocator, .err, "unable to parse res from '{s}': {s}", .{ res_stream.name, @errorName(err) });331 try error_handler.emitMessage(allocator, .err, "unable to parse res from '{s}': {s}", .{ res_stream.name, @errorName(err) });
332 std.process.exit(1);332 std.process.exit(1);
lib/docs/wasm/markdown.zig+6-7
...@@ -145,13 +145,12 @@ fn mainImpl() !void {...@@ -145,13 +145,12 @@ fn mainImpl() !void {
145 var parser = try Parser.init(gpa);145 var parser = try Parser.init(gpa);
146 defer parser.deinit();146 defer parser.deinit();
147147
148 var stdin_buf = std.io.bufferedReader(std.fs.File.stdin().deprecatedReader());148 var stdin_buffer: [1024]u8 = undefined;
149 var line_buf = std.ArrayList(u8).init(gpa);149 var stdin_reader = std.fs.File.stdin().reader(&stdin_buffer);
150 defer line_buf.deinit();150
151 while (stdin_buf.reader().streamUntilDelimiter(line_buf.writer(), '\n', null)) {151 while (stdin_reader.takeDelimiterExclusive('\n')) |line| {
152 if (line_buf.getLastOrNull() == '\r') _ = line_buf.pop();152 const trimmed = std.mem.trimRight(u8, line, '\r');
153 try parser.feedLine(line_buf.items);153 try parser.feedLine(trimmed);
154 line_buf.clearRetainingCapacity();
155 } else |err| switch (err) {154 } else |err| switch (err) {
156 error.EndOfStream => {},155 error.EndOfStream => {},
157 else => |e| return e,156 else => |e| return e,
lib/std/Build/Fuzz.zig+16-23
...@@ -234,7 +234,7 @@ pub const Previous = struct {...@@ -234,7 +234,7 @@ pub const Previous = struct {
234};234};
235pub fn sendUpdate(235pub fn sendUpdate(
236 fuzz: *Fuzz,236 fuzz: *Fuzz,
237 socket: *std.http.WebSocket,237 socket: *std.http.Server.WebSocket,
238 prev: *Previous,238 prev: *Previous,
239) !void {239) !void {
240 fuzz.coverage_mutex.lock();240 fuzz.coverage_mutex.lock();
...@@ -263,36 +263,36 @@ pub fn sendUpdate(...@@ -263,36 +263,36 @@ pub fn sendUpdate(
263 .string_bytes_len = @intCast(coverage_map.coverage.string_bytes.items.len),263 .string_bytes_len = @intCast(coverage_map.coverage.string_bytes.items.len),
264 .start_timestamp = coverage_map.start_timestamp,264 .start_timestamp = coverage_map.start_timestamp,
265 };265 };
266 const iovecs: [5]std.posix.iovec_const = .{266 var iovecs: [5][]const u8 = .{
267 makeIov(@ptrCast(&header)),267 @ptrCast(&header),
268 makeIov(@ptrCast(coverage_map.coverage.directories.keys())),268 @ptrCast(coverage_map.coverage.directories.keys()),
269 makeIov(@ptrCast(coverage_map.coverage.files.keys())),269 @ptrCast(coverage_map.coverage.files.keys()),
270 makeIov(@ptrCast(coverage_map.source_locations)),270 @ptrCast(coverage_map.source_locations),
271 makeIov(coverage_map.coverage.string_bytes.items),271 coverage_map.coverage.string_bytes.items,
272 };272 };
273 try socket.writeMessagev(&iovecs, .binary);273 try socket.writeMessageVec(&iovecs, .binary);
274 }274 }
275275
276 const header: abi.CoverageUpdateHeader = .{276 const header: abi.CoverageUpdateHeader = .{
277 .n_runs = n_runs,277 .n_runs = n_runs,
278 .unique_runs = unique_runs,278 .unique_runs = unique_runs,
279 };279 };
280 const iovecs: [2]std.posix.iovec_const = .{280 var iovecs: [2][]const u8 = .{
281 makeIov(@ptrCast(&header)),281 @ptrCast(&header),
282 makeIov(@ptrCast(seen_pcs)),282 @ptrCast(seen_pcs),
283 };283 };
284 try socket.writeMessagev(&iovecs, .binary);284 try socket.writeMessageVec(&iovecs, .binary);
285285
286 prev.unique_runs = unique_runs;286 prev.unique_runs = unique_runs;
287 }287 }
288288
289 if (prev.entry_points != coverage_map.entry_points.items.len) {289 if (prev.entry_points != coverage_map.entry_points.items.len) {
290 const header: abi.EntryPointHeader = .init(@intCast(coverage_map.entry_points.items.len));290 const header: abi.EntryPointHeader = .init(@intCast(coverage_map.entry_points.items.len));
291 const iovecs: [2]std.posix.iovec_const = .{291 var iovecs: [2][]const u8 = .{
292 makeIov(@ptrCast(&header)),292 @ptrCast(&header),
293 makeIov(@ptrCast(coverage_map.entry_points.items)),293 @ptrCast(coverage_map.entry_points.items),
294 };294 };
295 try socket.writeMessagev(&iovecs, .binary);295 try socket.writeMessageVec(&iovecs, .binary);
296296
297 prev.entry_points = coverage_map.entry_points.items.len;297 prev.entry_points = coverage_map.entry_points.items.len;
298 }298 }
...@@ -448,10 +448,3 @@ fn addEntryPoint(fuzz: *Fuzz, coverage_id: u64, addr: u64) error{ AlreadyReporte...@@ -448,10 +448,3 @@ fn addEntryPoint(fuzz: *Fuzz, coverage_id: u64, addr: u64) error{ AlreadyReporte
448 }448 }
449 try coverage_map.entry_points.append(fuzz.ws.gpa, @intCast(index));449 try coverage_map.entry_points.append(fuzz.ws.gpa, @intCast(index));
450}450}
451
452fn makeIov(s: []const u8) std.posix.iovec_const {
453 return .{
454 .base = s.ptr,
455 .len = s.len,
456 };
457}
lib/std/Build/WebServer.zig+42-53
...@@ -251,48 +251,44 @@ pub fn now(s: *const WebServer) i64 {...@@ -251,48 +251,44 @@ pub fn now(s: *const WebServer) i64 {
251fn accept(ws: *WebServer, connection: std.net.Server.Connection) void {251fn accept(ws: *WebServer, connection: std.net.Server.Connection) void {
252 defer connection.stream.close();252 defer connection.stream.close();
253253
254 var read_buf: [0x4000]u8 = undefined;254 var send_buffer: [4096]u8 = undefined;
255 var server: std.http.Server = .init(connection, &read_buf);255 var recv_buffer: [4096]u8 = undefined;
256 var connection_reader = connection.stream.reader(&recv_buffer);
257 var connection_writer = connection.stream.writer(&send_buffer);
258 var server: http.Server = .init(connection_reader.interface(), &connection_writer.interface);
256259
257 while (true) {260 while (true) {
258 var request = server.receiveHead() catch |err| switch (err) {261 var request = server.receiveHead() catch |err| switch (err) {
259 error.HttpConnectionClosing => return,262 error.HttpConnectionClosing => return,
260 else => {263 else => return log.err("failed to receive http request: {t}", .{err}),
261 log.err("failed to receive http request: {s}", .{@errorName(err)});
262 return;
263 },
264 };264 };
265 var ws_send_buf: [0x4000]u8 = undefined;265 switch (request.upgradeRequested()) {
266 var ws_recv_buf: [0x4000]u8 align(4) = undefined;266 .websocket => |opt_key| {
267 if (std.http.WebSocket.init(&request, &ws_send_buf, &ws_recv_buf) catch |err| {267 const key = opt_key orelse return log.err("missing websocket key", .{});
268 log.err("failed to initialize websocket connection: {s}", .{@errorName(err)});268 var web_socket = request.respondWebSocket(.{ .key = key }) catch {
269 return;269 return log.err("failed to respond web socket: {t}", .{connection_writer.err.?});
270 }) |ws_init| {270 };
271 var web_socket = ws_init;271 ws.serveWebSocket(&web_socket) catch |err| {
272 ws.serveWebSocket(&web_socket) catch |err| {272 log.err("failed to serve websocket: {t}", .{err});
273 log.err("failed to serve websocket: {s}", .{@errorName(err)});
274 return;
275 };
276 comptime unreachable;
277 } else {
278 ws.serveRequest(&request) catch |err| switch (err) {
279 error.AlreadyReported => return,
280 else => {
281 log.err("failed to serve '{s}': {s}", .{ request.head.target, @errorName(err) });
282 return;273 return;
283 },274 };
284 };275 comptime unreachable;
276 },
277 .other => |name| return log.err("unknown upgrade request: {s}", .{name}),
278 .none => {
279 ws.serveRequest(&request) catch |err| switch (err) {
280 error.AlreadyReported => return,
281 else => {
282 log.err("failed to serve '{s}': {t}", .{ request.head.target, err });
283 return;
284 },
285 };
286 },
285 }287 }
286 }288 }
287}289}
288290
289fn makeIov(s: []const u8) std.posix.iovec_const {291fn serveWebSocket(ws: *WebServer, sock: *http.Server.WebSocket) !noreturn {
290 return .{
291 .base = s.ptr,
292 .len = s.len,
293 };
294}
295fn serveWebSocket(ws: *WebServer, sock: *std.http.WebSocket) !noreturn {
296 var prev_build_status = ws.build_status.load(.monotonic);292 var prev_build_status = ws.build_status.load(.monotonic);
297293
298 const prev_step_status_bits = try ws.gpa.alloc(u8, ws.step_status_bits.len);294 const prev_step_status_bits = try ws.gpa.alloc(u8, ws.step_status_bits.len);
...@@ -312,11 +308,8 @@ fn serveWebSocket(ws: *WebServer, sock: *std.http.WebSocket) !noreturn {...@@ -312,11 +308,8 @@ fn serveWebSocket(ws: *WebServer, sock: *std.http.WebSocket) !noreturn {
312 .timestamp = ws.now(),308 .timestamp = ws.now(),
313 .steps_len = @intCast(ws.all_steps.len),309 .steps_len = @intCast(ws.all_steps.len),
314 };310 };
315 try sock.writeMessagev(&.{311 var bufs: [3][]const u8 = .{ @ptrCast(&hello_header), ws.step_names_trailing, prev_step_status_bits };
316 makeIov(@ptrCast(&hello_header)),312 try sock.writeMessageVec(&bufs, .binary);
317 makeIov(ws.step_names_trailing),
318 makeIov(prev_step_status_bits),
319 }, .binary);
320 }313 }
321314
322 var prev_fuzz: Fuzz.Previous = .init;315 var prev_fuzz: Fuzz.Previous = .init;
...@@ -380,7 +373,7 @@ fn serveWebSocket(ws: *WebServer, sock: *std.http.WebSocket) !noreturn {...@@ -380,7 +373,7 @@ fn serveWebSocket(ws: *WebServer, sock: *std.http.WebSocket) !noreturn {
380 std.Thread.Futex.timedWait(&ws.update_id, start_update_id, std.time.ns_per_ms * default_update_interval_ms) catch {};373 std.Thread.Futex.timedWait(&ws.update_id, start_update_id, std.time.ns_per_ms * default_update_interval_ms) catch {};
381 }374 }
382}375}
383fn recvWebSocketMessages(ws: *WebServer, sock: *std.http.WebSocket) void {376fn recvWebSocketMessages(ws: *WebServer, sock: *http.Server.WebSocket) void {
384 while (true) {377 while (true) {
385 const msg = sock.readSmallMessage() catch return;378 const msg = sock.readSmallMessage() catch return;
386 if (msg.opcode != .binary) continue;379 if (msg.opcode != .binary) continue;
...@@ -402,7 +395,7 @@ fn recvWebSocketMessages(ws: *WebServer, sock: *std.http.WebSocket) void {...@@ -402,7 +395,7 @@ fn recvWebSocketMessages(ws: *WebServer, sock: *std.http.WebSocket) void {
402 }395 }
403}396}
404397
405fn serveRequest(ws: *WebServer, req: *std.http.Server.Request) !void {398fn serveRequest(ws: *WebServer, req: *http.Server.Request) !void {
406 // Strip an optional leading '/debug' component from the request.399 // Strip an optional leading '/debug' component from the request.
407 const target: []const u8, const debug: bool = target: {400 const target: []const u8, const debug: bool = target: {
408 if (mem.eql(u8, req.head.target, "/debug")) break :target .{ "/", true };401 if (mem.eql(u8, req.head.target, "/debug")) break :target .{ "/", true };
...@@ -431,7 +424,7 @@ fn serveRequest(ws: *WebServer, req: *std.http.Server.Request) !void {...@@ -431,7 +424,7 @@ fn serveRequest(ws: *WebServer, req: *std.http.Server.Request) !void {
431424
432fn serveLibFile(425fn serveLibFile(
433 ws: *WebServer,426 ws: *WebServer,
434 request: *std.http.Server.Request,427 request: *http.Server.Request,
435 sub_path: []const u8,428 sub_path: []const u8,
436 content_type: []const u8,429 content_type: []const u8,
437) !void {430) !void {
...@@ -442,7 +435,7 @@ fn serveLibFile(...@@ -442,7 +435,7 @@ fn serveLibFile(
442}435}
443fn serveClientWasm(436fn serveClientWasm(
444 ws: *WebServer,437 ws: *WebServer,
445 req: *std.http.Server.Request,438 req: *http.Server.Request,
446 optimize_mode: std.builtin.OptimizeMode,439 optimize_mode: std.builtin.OptimizeMode,
447) !void {440) !void {
448 var arena_state: std.heap.ArenaAllocator = .init(ws.gpa);441 var arena_state: std.heap.ArenaAllocator = .init(ws.gpa);
...@@ -456,12 +449,12 @@ fn serveClientWasm(...@@ -456,12 +449,12 @@ fn serveClientWasm(
456449
457pub fn serveFile(450pub fn serveFile(
458 ws: *WebServer,451 ws: *WebServer,
459 request: *std.http.Server.Request,452 request: *http.Server.Request,
460 path: Cache.Path,453 path: Cache.Path,
461 content_type: []const u8,454 content_type: []const u8,
462) !void {455) !void {
463 const gpa = ws.gpa;456 const gpa = ws.gpa;
464 // The desired API is actually sendfile, which will require enhancing std.http.Server.457 // The desired API is actually sendfile, which will require enhancing http.Server.
465 // We load the file with every request so that the user can make changes to the file458 // We load the file with every request so that the user can make changes to the file
466 // and refresh the HTML page without restarting this server.459 // and refresh the HTML page without restarting this server.
467 const file_contents = path.root_dir.handle.readFileAlloc(gpa, path.sub_path, 10 * 1024 * 1024) catch |err| {460 const file_contents = path.root_dir.handle.readFileAlloc(gpa, path.sub_path, 10 * 1024 * 1024) catch |err| {
...@@ -478,14 +471,13 @@ pub fn serveFile(...@@ -478,14 +471,13 @@ pub fn serveFile(
478}471}
479pub fn serveTarFile(472pub fn serveTarFile(
480 ws: *WebServer,473 ws: *WebServer,
481 request: *std.http.Server.Request,474 request: *http.Server.Request,
482 paths: []const Cache.Path,475 paths: []const Cache.Path,
483) !void {476) !void {
484 const gpa = ws.gpa;477 const gpa = ws.gpa;
485478
486 var send_buf: [0x4000]u8 = undefined;479 var send_buffer: [0x4000]u8 = undefined;
487 var response = request.respondStreaming(.{480 var response = try request.respondStreaming(&send_buffer, .{
488 .send_buffer = &send_buf,
489 .respond_options = .{481 .respond_options = .{
490 .extra_headers = &.{482 .extra_headers = &.{
491 .{ .name = "Content-Type", .value = "application/x-tar" },483 .{ .name = "Content-Type", .value = "application/x-tar" },
...@@ -497,10 +489,7 @@ pub fn serveTarFile(...@@ -497,10 +489,7 @@ pub fn serveTarFile(
497 var cached_cwd_path: ?[]const u8 = null;489 var cached_cwd_path: ?[]const u8 = null;
498 defer if (cached_cwd_path) |p| gpa.free(p);490 defer if (cached_cwd_path) |p| gpa.free(p);
499491
500 var response_buf: [1024]u8 = undefined;492 var archiver: std.tar.Writer = .{ .underlying_writer = &response.writer };
501 var adapter = response.writer().adaptToNewApi();
502 adapter.new_interface.buffer = &response_buf;
503 var archiver: std.tar.Writer = .{ .underlying_writer = &adapter.new_interface };
504493
505 for (paths) |path| {494 for (paths) |path| {
506 var file = path.root_dir.handle.openFile(path.sub_path, .{}) catch |err| {495 var file = path.root_dir.handle.openFile(path.sub_path, .{}) catch |err| {
...@@ -526,7 +515,6 @@ pub fn serveTarFile(...@@ -526,7 +515,6 @@ pub fn serveTarFile(
526 }515 }
527516
528 // intentionally not calling `archiver.finishPedantically`517 // intentionally not calling `archiver.finishPedantically`
529 try adapter.new_interface.flush();
530 try response.end();518 try response.end();
531}519}
532520
...@@ -804,7 +792,7 @@ pub fn wait(ws: *WebServer) RunnerRequest {...@@ -804,7 +792,7 @@ pub fn wait(ws: *WebServer) RunnerRequest {
804 }792 }
805}793}
806794
807const cache_control_header: std.http.Header = .{795const cache_control_header: http.Header = .{
808 .name = "Cache-Control",796 .name = "Cache-Control",
809 .value = "max-age=0, must-revalidate",797 .value = "max-age=0, must-revalidate",
810};798};
...@@ -819,5 +807,6 @@ const Build = std.Build;...@@ -819,5 +807,6 @@ const Build = std.Build;
819const Cache = Build.Cache;807const Cache = Build.Cache;
820const Fuzz = Build.Fuzz;808const Fuzz = Build.Fuzz;
821const abi = Build.abi;809const abi = Build.abi;
810const http = std.http;
822811
823const WebServer = @This();812const WebServer = @This();
lib/std/Io.zig-11
...@@ -428,19 +428,9 @@ pub const BufferedWriter = @import("Io/buffered_writer.zig").BufferedWriter;...@@ -428,19 +428,9 @@ pub const BufferedWriter = @import("Io/buffered_writer.zig").BufferedWriter;
428/// Deprecated in favor of `Writer`.428/// Deprecated in favor of `Writer`.
429pub const bufferedWriter = @import("Io/buffered_writer.zig").bufferedWriter;429pub const bufferedWriter = @import("Io/buffered_writer.zig").bufferedWriter;
430/// Deprecated in favor of `Reader`.430/// Deprecated in favor of `Reader`.
431pub const BufferedReader = @import("Io/buffered_reader.zig").BufferedReader;
432/// Deprecated in favor of `Reader`.
433pub const bufferedReader = @import("Io/buffered_reader.zig").bufferedReader;
434/// Deprecated in favor of `Reader`.
435pub const bufferedReaderSize = @import("Io/buffered_reader.zig").bufferedReaderSize;
436/// Deprecated in favor of `Reader`.
437pub const FixedBufferStream = @import("Io/fixed_buffer_stream.zig").FixedBufferStream;431pub const FixedBufferStream = @import("Io/fixed_buffer_stream.zig").FixedBufferStream;
438/// Deprecated in favor of `Reader`.432/// Deprecated in favor of `Reader`.
439pub const fixedBufferStream = @import("Io/fixed_buffer_stream.zig").fixedBufferStream;433pub const fixedBufferStream = @import("Io/fixed_buffer_stream.zig").fixedBufferStream;
440/// Deprecated in favor of `Reader.Limited`.
441pub const LimitedReader = @import("Io/limited_reader.zig").LimitedReader;
442/// Deprecated in favor of `Reader.Limited`.
443pub const limitedReader = @import("Io/limited_reader.zig").limitedReader;
444/// Deprecated with no replacement; inefficient pattern434/// Deprecated with no replacement; inefficient pattern
445pub const CountingWriter = @import("Io/counting_writer.zig").CountingWriter;435pub const CountingWriter = @import("Io/counting_writer.zig").CountingWriter;
446/// Deprecated with no replacement; inefficient pattern436/// Deprecated with no replacement; inefficient pattern
...@@ -926,7 +916,6 @@ pub fn PollFiles(comptime StreamEnum: type) type {...@@ -926,7 +916,6 @@ pub fn PollFiles(comptime StreamEnum: type) type {
926test {916test {
927 _ = Reader;917 _ = Reader;
928 _ = Writer;918 _ = Writer;
929 _ = BufferedReader;
930 _ = BufferedWriter;919 _ = BufferedWriter;
931 _ = CountingWriter;920 _ = CountingWriter;
932 _ = CountingReader;921 _ = CountingReader;
lib/std/Io/Reader.zig+5-27
...@@ -367,8 +367,11 @@ pub fn appendRemainingUnlimited(...@@ -367,8 +367,11 @@ pub fn appendRemainingUnlimited(
367 const buffer_contents = r.buffer[r.seek..r.end];367 const buffer_contents = r.buffer[r.seek..r.end];
368 try list.ensureUnusedCapacity(gpa, buffer_contents.len + bump);368 try list.ensureUnusedCapacity(gpa, buffer_contents.len + bump);
369 list.appendSliceAssumeCapacity(buffer_contents);369 list.appendSliceAssumeCapacity(buffer_contents);
370 r.seek = 0;370 // If statement protects `ending`.
371 r.end = 0;371 if (r.end != 0) {
372 r.seek = 0;
373 r.end = 0;
374 }
372 // From here, we leave `buffer` empty, appending directly to `list`.375 // From here, we leave `buffer` empty, appending directly to `list`.
373 var writer: Writer = .{376 var writer: Writer = .{
374 .buffer = undefined,377 .buffer = undefined,
...@@ -1306,31 +1309,6 @@ pub fn defaultRebase(r: *Reader, capacity: usize) RebaseError!void {...@@ -1306,31 +1309,6 @@ pub fn defaultRebase(r: *Reader, capacity: usize) RebaseError!void {
1306 r.end = data.len;1309 r.end = data.len;
1307}1310}
13081311
1309/// Advances the stream and decreases the size of the storage buffer by `n`,
1310/// returning the range of bytes no longer accessible by `r`.
1311///
1312/// This action can be undone by `restitute`.
1313///
1314/// Asserts there are at least `n` buffered bytes already.
1315///
1316/// Asserts that `r.seek` is zero, i.e. the buffer is in a rebased state.
1317pub fn steal(r: *Reader, n: usize) []u8 {
1318 assert(r.seek == 0);
1319 assert(n <= r.end);
1320 const stolen = r.buffer[0..n];
1321 r.buffer = r.buffer[n..];
1322 r.end -= n;
1323 return stolen;
1324}
1325
1326/// Expands the storage buffer, undoing the effects of `steal`
1327/// Assumes that `n` does not exceed the total number of stolen bytes.
1328pub fn restitute(r: *Reader, n: usize) void {
1329 r.buffer = (r.buffer.ptr - n)[0 .. r.buffer.len + n];
1330 r.end += n;
1331 r.seek += n;
1332}
1333
1334test fixed {1312test fixed {
1335 var r: Reader = .fixed("a\x02");1313 var r: Reader = .fixed("a\x02");
1336 try testing.expect((try r.takeByte()) == 'a');1314 try testing.expect((try r.takeByte()) == 'a');
lib/std/Io/Writer.zig+77-19
...@@ -191,29 +191,87 @@ pub fn writeSplatHeader(...@@ -191,29 +191,87 @@ pub fn writeSplatHeader(
191 data: []const []const u8,191 data: []const []const u8,
192 splat: usize,192 splat: usize,
193) Error!usize {193) Error!usize {
194 const new_end = w.end + header.len;194 return writeSplatHeaderLimit(w, header, data, splat, .unlimited);
195 if (new_end <= w.buffer.len) {195}
196 @memcpy(w.buffer[w.end..][0..header.len], header);196
197 w.end = new_end;197/// Equivalent to `writeSplatHeader` but writes at most `limit` bytes.
198 return header.len + try writeSplat(w, data, splat);198pub fn writeSplatHeaderLimit(
199 }199 w: *Writer,
200 var vecs: [8][]const u8 = undefined; // Arbitrarily chosen size.200 header: []const u8,
201 var i: usize = 1;201 data: []const []const u8,
202 vecs[0] = header;202 splat: usize,
203 for (data[0 .. data.len - 1]) |buf| {203 limit: Limit,
204 if (buf.len == 0) continue;204) Error!usize {
205 vecs[i] = buf;205 var remaining = @intFromEnum(limit);
206 i += 1;206 {
207 if (vecs.len - i == 0) break;207 const copy_len = @min(header.len, w.buffer.len - w.end, remaining);
208 if (header.len - copy_len != 0) return writeSplatHeaderLimitFinish(w, header, data, splat, remaining);
209 @memcpy(w.buffer[w.end..][0..copy_len], header[0..copy_len]);
210 w.end += copy_len;
211 remaining -= copy_len;
212 }
213 for (data[0 .. data.len - 1], 0..) |buf, i| {
214 const copy_len = @min(buf.len, w.buffer.len - w.end, remaining);
215 if (buf.len - copy_len != 0) return @intFromEnum(limit) - remaining +
216 try writeSplatHeaderLimitFinish(w, &.{}, data[i..], splat, remaining);
217 @memcpy(w.buffer[w.end..][0..copy_len], buf[0..copy_len]);
218 w.end += copy_len;
219 remaining -= copy_len;
208 }220 }
209 const pattern = data[data.len - 1];221 const pattern = data[data.len - 1];
210 const new_splat = s: {222 const splat_n = pattern.len * splat;
211 if (pattern.len == 0 or vecs.len - i == 0) break :s 1;223 if (splat_n > @min(w.buffer.len - w.end, remaining)) {
224 const buffered_n = @intFromEnum(limit) - remaining;
225 const written = try writeSplatHeaderLimitFinish(w, &.{}, data[data.len - 1 ..][0..1], splat, remaining);
226 return buffered_n + written;
227 }
228
229 for (0..splat) |_| {
230 @memcpy(w.buffer[w.end..][0..pattern.len], pattern);
231 w.end += pattern.len;
232 }
233
234 remaining -= splat_n;
235 return @intFromEnum(limit) - remaining;
236}
237
238fn writeSplatHeaderLimitFinish(
239 w: *Writer,
240 header: []const u8,
241 data: []const []const u8,
242 splat: usize,
243 limit: usize,
244) Error!usize {
245 var remaining = limit;
246 var vecs: [8][]const u8 = undefined;
247 var i: usize = 0;
248 v: {
249 if (header.len != 0) {
250 const copy_len = @min(header.len, remaining);
251 vecs[i] = header[0..copy_len];
252 i += 1;
253 remaining -= copy_len;
254 if (remaining == 0) break :v;
255 }
256 for (data[0 .. data.len - 1]) |buf| if (buf.len != 0) {
257 const copy_len = @min(header.len, remaining);
258 vecs[i] = buf;
259 i += 1;
260 remaining -= copy_len;
261 if (remaining == 0) break :v;
262 if (vecs.len - i == 0) break :v;
263 };
264 const pattern = data[data.len - 1];
265 if (splat == 1) {
266 vecs[i] = pattern[0..@min(remaining, pattern.len)];
267 i += 1;
268 break :v;
269 }
212 vecs[i] = pattern;270 vecs[i] = pattern;
213 i += 1;271 i += 1;
214 break :s splat;272 return w.vtable.drain(w, (&vecs)[0..i], @min(remaining / pattern.len, splat));
215 };273 }
216 return w.vtable.drain(w, vecs[0..i], new_splat);274 return w.vtable.drain(w, (&vecs)[0..i], 1);
217}275}
218276
219test "writeSplatHeader splatting avoids buffer aliasing temptation" {277test "writeSplatHeader splatting avoids buffer aliasing temptation" {
lib/std/Io/buffered_reader.zig deleted-201
...@@ -1,201 +0,0 @@
1const std = @import("../std.zig");
2const io = std.io;
3const mem = std.mem;
4const assert = std.debug.assert;
5const testing = std.testing;
6
7pub fn BufferedReader(comptime buffer_size: usize, comptime ReaderType: type) type {
8 return struct {
9 unbuffered_reader: ReaderType,
10 buf: [buffer_size]u8 = undefined,
11 start: usize = 0,
12 end: usize = 0,
13
14 pub const Error = ReaderType.Error;
15 pub const Reader = io.GenericReader(*Self, Error, read);
16
17 const Self = @This();
18
19 pub fn read(self: *Self, dest: []u8) Error!usize {
20 // First try reading from the already buffered data onto the destination.
21 const current = self.buf[self.start..self.end];
22 if (current.len != 0) {
23 const to_transfer = @min(current.len, dest.len);
24 @memcpy(dest[0..to_transfer], current[0..to_transfer]);
25 self.start += to_transfer;
26 return to_transfer;
27 }
28
29 // If dest is large, read from the unbuffered reader directly into the destination.
30 if (dest.len >= buffer_size) {
31 return self.unbuffered_reader.read(dest);
32 }
33
34 // If dest is small, read from the unbuffered reader into our own internal buffer,
35 // and then transfer to destination.
36 self.end = try self.unbuffered_reader.read(&self.buf);
37 const to_transfer = @min(self.end, dest.len);
38 @memcpy(dest[0..to_transfer], self.buf[0..to_transfer]);
39 self.start = to_transfer;
40 return to_transfer;
41 }
42
43 pub fn reader(self: *Self) Reader {
44 return .{ .context = self };
45 }
46 };
47}
48
49pub fn bufferedReader(reader: anytype) BufferedReader(4096, @TypeOf(reader)) {
50 return .{ .unbuffered_reader = reader };
51}
52
53pub fn bufferedReaderSize(comptime size: usize, reader: anytype) BufferedReader(size, @TypeOf(reader)) {
54 return .{ .unbuffered_reader = reader };
55}
56
57test "OneByte" {
58 const OneByteReadReader = struct {
59 str: []const u8,
60 curr: usize,
61
62 const Error = error{NoError};
63 const Self = @This();
64 const Reader = io.GenericReader(*Self, Error, read);
65
66 fn init(str: []const u8) Self {
67 return Self{
68 .str = str,
69 .curr = 0,
70 };
71 }
72
73 fn read(self: *Self, dest: []u8) Error!usize {
74 if (self.str.len <= self.curr or dest.len == 0)
75 return 0;
76
77 dest[0] = self.str[self.curr];
78 self.curr += 1;
79 return 1;
80 }
81
82 fn reader(self: *Self) Reader {
83 return .{ .context = self };
84 }
85 };
86
87 const str = "This is a test";
88 var one_byte_stream = OneByteReadReader.init(str);
89 var buf_reader = bufferedReader(one_byte_stream.reader());
90 const stream = buf_reader.reader();
91
92 const res = try stream.readAllAlloc(testing.allocator, str.len + 1);
93 defer testing.allocator.free(res);
94 try testing.expectEqualSlices(u8, str, res);
95}
96
97fn smallBufferedReader(underlying_stream: anytype) BufferedReader(8, @TypeOf(underlying_stream)) {
98 return .{ .unbuffered_reader = underlying_stream };
99}
100test "Block" {
101 const BlockReader = struct {
102 block: []const u8,
103 reads_allowed: usize,
104 curr_read: usize,
105
106 const Error = error{NoError};
107 const Self = @This();
108 const Reader = io.GenericReader(*Self, Error, read);
109
110 fn init(block: []const u8, reads_allowed: usize) Self {
111 return Self{
112 .block = block,
113 .reads_allowed = reads_allowed,
114 .curr_read = 0,
115 };
116 }
117
118 fn read(self: *Self, dest: []u8) Error!usize {
119 if (self.curr_read >= self.reads_allowed) return 0;
120 @memcpy(dest[0..self.block.len], self.block);
121
122 self.curr_read += 1;
123 return self.block.len;
124 }
125
126 fn reader(self: *Self) Reader {
127 return .{ .context = self };
128 }
129 };
130
131 const block = "0123";
132
133 // len out == block
134 {
135 var test_buf_reader: BufferedReader(4, BlockReader) = .{
136 .unbuffered_reader = BlockReader.init(block, 2),
137 };
138 const reader = test_buf_reader.reader();
139 var out_buf: [4]u8 = undefined;
140 _ = try reader.readAll(&out_buf);
141 try testing.expectEqualSlices(u8, &out_buf, block);
142 _ = try reader.readAll(&out_buf);
143 try testing.expectEqualSlices(u8, &out_buf, block);
144 try testing.expectEqual(try reader.readAll(&out_buf), 0);
145 }
146
147 // len out < block
148 {
149 var test_buf_reader: BufferedReader(4, BlockReader) = .{
150 .unbuffered_reader = BlockReader.init(block, 2),
151 };
152 const reader = test_buf_reader.reader();
153 var out_buf: [3]u8 = undefined;
154 _ = try reader.readAll(&out_buf);
155 try testing.expectEqualSlices(u8, &out_buf, "012");
156 _ = try reader.readAll(&out_buf);
157 try testing.expectEqualSlices(u8, &out_buf, "301");
158 const n = try reader.readAll(&out_buf);
159 try testing.expectEqualSlices(u8, out_buf[0..n], "23");
160 try testing.expectEqual(try reader.readAll(&out_buf), 0);
161 }
162
163 // len out > block
164 {
165 var test_buf_reader: BufferedReader(4, BlockReader) = .{
166 .unbuffered_reader = BlockReader.init(block, 2),
167 };
168 const reader = test_buf_reader.reader();
169 var out_buf: [5]u8 = undefined;
170 _ = try reader.readAll(&out_buf);
171 try testing.expectEqualSlices(u8, &out_buf, "01230");
172 const n = try reader.readAll(&out_buf);
173 try testing.expectEqualSlices(u8, out_buf[0..n], "123");
174 try testing.expectEqual(try reader.readAll(&out_buf), 0);
175 }
176
177 // len out == 0
178 {
179 var test_buf_reader: BufferedReader(4, BlockReader) = .{
180 .unbuffered_reader = BlockReader.init(block, 2),
181 };
182 const reader = test_buf_reader.reader();
183 var out_buf: [0]u8 = undefined;
184 _ = try reader.readAll(&out_buf);
185 try testing.expectEqualSlices(u8, &out_buf, "");
186 }
187
188 // len bufreader buf > block
189 {
190 var test_buf_reader: BufferedReader(5, BlockReader) = .{
191 .unbuffered_reader = BlockReader.init(block, 2),
192 };
193 const reader = test_buf_reader.reader();
194 var out_buf: [4]u8 = undefined;
195 _ = try reader.readAll(&out_buf);
196 try testing.expectEqualSlices(u8, &out_buf, block);
197 _ = try reader.readAll(&out_buf);
198 try testing.expectEqualSlices(u8, &out_buf, block);
199 try testing.expectEqual(try reader.readAll(&out_buf), 0);
200 }
201}
lib/std/Io/limited_reader.zig deleted-45
...@@ -1,45 +0,0 @@
1const std = @import("../std.zig");
2const io = std.io;
3const assert = std.debug.assert;
4const testing = std.testing;
5
6pub fn LimitedReader(comptime ReaderType: type) type {
7 return struct {
8 inner_reader: ReaderType,
9 bytes_left: u64,
10
11 pub const Error = ReaderType.Error;
12 pub const Reader = io.GenericReader(*Self, Error, read);
13
14 const Self = @This();
15
16 pub fn read(self: *Self, dest: []u8) Error!usize {
17 const max_read = @min(self.bytes_left, dest.len);
18 const n = try self.inner_reader.read(dest[0..max_read]);
19 self.bytes_left -= n;
20 return n;
21 }
22
23 pub fn reader(self: *Self) Reader {
24 return .{ .context = self };
25 }
26 };
27}
28
29/// Returns an initialised `LimitedReader`.
30/// `bytes_left` is a `u64` to be able to take 64 bit file offsets
31pub fn limitedReader(inner_reader: anytype, bytes_left: u64) LimitedReader(@TypeOf(inner_reader)) {
32 return .{ .inner_reader = inner_reader, .bytes_left = bytes_left };
33}
34
35test "basic usage" {
36 const data = "hello world";
37 var fbs = std.io.fixedBufferStream(data);
38 var early_stream = limitedReader(fbs.reader(), 3);
39
40 var buf: [5]u8 = undefined;
41 try testing.expectEqual(@as(usize, 3), try early_stream.reader().read(&buf));
42 try testing.expectEqualSlices(u8, data[0..3], buf[0..3]);
43 try testing.expectEqual(@as(usize, 0), try early_stream.reader().read(&buf));
44 try testing.expectError(error.EndOfStream, early_stream.reader().skipBytes(10, .{}));
45}
lib/std/Io/test.zig+3-3
...@@ -45,9 +45,9 @@ test "write a file, read it, then delete it" {...@@ -45,9 +45,9 @@ test "write a file, read it, then delete it" {
45 const expected_file_size: u64 = "begin".len + data.len + "end".len;45 const expected_file_size: u64 = "begin".len + data.len + "end".len;
46 try expectEqual(expected_file_size, file_size);46 try expectEqual(expected_file_size, file_size);
4747
48 var buf_stream = io.bufferedReader(file.deprecatedReader());48 var file_buffer: [1024]u8 = undefined;
49 const st = buf_stream.reader();49 var file_reader = file.reader(&file_buffer);
50 const contents = try st.readAllAlloc(std.testing.allocator, 2 * 1024);50 const contents = try file_reader.interface.allocRemaining(std.testing.allocator, .limited(2 * 1024));
51 defer std.testing.allocator.free(contents);51 defer std.testing.allocator.free(contents);
5252
53 try expect(mem.eql(u8, contents[0.."begin".len], "begin"));53 try expect(mem.eql(u8, contents[0.."begin".len], "begin"));
lib/std/Uri.zig+92-147
...@@ -4,6 +4,8 @@...@@ -4,6 +4,8 @@
4const std = @import("std.zig");4const std = @import("std.zig");
5const testing = std.testing;5const testing = std.testing;
6const Uri = @This();6const Uri = @This();
7const Allocator = std.mem.Allocator;
8const Writer = std.Io.Writer;
79
8scheme: []const u8,10scheme: []const u8,
9user: ?Component = null,11user: ?Component = null,
...@@ -14,6 +16,32 @@ path: Component = Component.empty,...@@ -14,6 +16,32 @@ path: Component = Component.empty,
14query: ?Component = null,16query: ?Component = null,
15fragment: ?Component = null,17fragment: ?Component = null,
1618
19pub const host_name_max = 255;
20
21/// Returned value may point into `buffer` or be the original string.
22///
23/// Suggested buffer length: `host_name_max`.
24///
25/// See also:
26/// * `getHostAlloc`
27pub fn getHost(uri: Uri, buffer: []u8) error{ UriMissingHost, UriHostTooLong }![]const u8 {
28 const component = uri.host orelse return error.UriMissingHost;
29 return component.toRaw(buffer) catch |err| switch (err) {
30 error.NoSpaceLeft => return error.UriHostTooLong,
31 };
32}
33
34/// Returned value may point into `buffer` or be the original string.
35///
36/// See also:
37/// * `getHost`
38pub fn getHostAlloc(uri: Uri, arena: Allocator) error{ UriMissingHost, UriHostTooLong, OutOfMemory }![]const u8 {
39 const component = uri.host orelse return error.UriMissingHost;
40 const result = try component.toRawMaybeAlloc(arena);
41 if (result.len > host_name_max) return error.UriHostTooLong;
42 return result;
43}
44
17pub const Component = union(enum) {45pub const Component = union(enum) {
18 /// Invalid characters in this component must be percent encoded46 /// Invalid characters in this component must be percent encoded
19 /// before being printed as part of a URI.47 /// before being printed as part of a URI.
...@@ -30,11 +58,19 @@ pub const Component = union(enum) {...@@ -30,11 +58,19 @@ pub const Component = union(enum) {
30 };58 };
31 }59 }
3260
61 /// Returned value may point into `buffer` or be the original string.
62 pub fn toRaw(component: Component, buffer: []u8) error{NoSpaceLeft}![]const u8 {
63 return switch (component) {
64 .raw => |raw| raw,
65 .percent_encoded => |percent_encoded| if (std.mem.indexOfScalar(u8, percent_encoded, '%')) |_|
66 try std.fmt.bufPrint(buffer, "{f}", .{std.fmt.alt(component, .formatRaw)})
67 else
68 percent_encoded,
69 };
70 }
71
33 /// Allocates the result with `arena` only if needed, so the result should not be freed.72 /// Allocates the result with `arena` only if needed, so the result should not be freed.
34 pub fn toRawMaybeAlloc(73 pub fn toRawMaybeAlloc(component: Component, arena: Allocator) Allocator.Error![]const u8 {
35 component: Component,
36 arena: std.mem.Allocator,
37 ) std.mem.Allocator.Error![]const u8 {
38 return switch (component) {74 return switch (component) {
39 .raw => |raw| raw,75 .raw => |raw| raw,
40 .percent_encoded => |percent_encoded| if (std.mem.indexOfScalar(u8, percent_encoded, '%')) |_|76 .percent_encoded => |percent_encoded| if (std.mem.indexOfScalar(u8, percent_encoded, '%')) |_|
...@@ -44,7 +80,7 @@ pub const Component = union(enum) {...@@ -44,7 +80,7 @@ pub const Component = union(enum) {
44 };80 };
45 }81 }
4682
47 pub fn formatRaw(component: Component, w: *std.io.Writer) std.io.Writer.Error!void {83 pub fn formatRaw(component: Component, w: *Writer) Writer.Error!void {
48 switch (component) {84 switch (component) {
49 .raw => |raw| try w.writeAll(raw),85 .raw => |raw| try w.writeAll(raw),
50 .percent_encoded => |percent_encoded| {86 .percent_encoded => |percent_encoded| {
...@@ -67,56 +103,56 @@ pub const Component = union(enum) {...@@ -67,56 +103,56 @@ pub const Component = union(enum) {
67 }103 }
68 }104 }
69105
70 pub fn formatEscaped(component: Component, w: *std.io.Writer) std.io.Writer.Error!void {106 pub fn formatEscaped(component: Component, w: *Writer) Writer.Error!void {
71 switch (component) {107 switch (component) {
72 .raw => |raw| try percentEncode(w, raw, isUnreserved),108 .raw => |raw| try percentEncode(w, raw, isUnreserved),
73 .percent_encoded => |percent_encoded| try w.writeAll(percent_encoded),109 .percent_encoded => |percent_encoded| try w.writeAll(percent_encoded),
74 }110 }
75 }111 }
76112
77 pub fn formatUser(component: Component, w: *std.io.Writer) std.io.Writer.Error!void {113 pub fn formatUser(component: Component, w: *Writer) Writer.Error!void {
78 switch (component) {114 switch (component) {
79 .raw => |raw| try percentEncode(w, raw, isUserChar),115 .raw => |raw| try percentEncode(w, raw, isUserChar),
80 .percent_encoded => |percent_encoded| try w.writeAll(percent_encoded),116 .percent_encoded => |percent_encoded| try w.writeAll(percent_encoded),
81 }117 }
82 }118 }
83119
84 pub fn formatPassword(component: Component, w: *std.io.Writer) std.io.Writer.Error!void {120 pub fn formatPassword(component: Component, w: *Writer) Writer.Error!void {
85 switch (component) {121 switch (component) {
86 .raw => |raw| try percentEncode(w, raw, isPasswordChar),122 .raw => |raw| try percentEncode(w, raw, isPasswordChar),
87 .percent_encoded => |percent_encoded| try w.writeAll(percent_encoded),123 .percent_encoded => |percent_encoded| try w.writeAll(percent_encoded),
88 }124 }
89 }125 }
90126
91 pub fn formatHost(component: Component, w: *std.io.Writer) std.io.Writer.Error!void {127 pub fn formatHost(component: Component, w: *Writer) Writer.Error!void {
92 switch (component) {128 switch (component) {
93 .raw => |raw| try percentEncode(w, raw, isHostChar),129 .raw => |raw| try percentEncode(w, raw, isHostChar),
94 .percent_encoded => |percent_encoded| try w.writeAll(percent_encoded),130 .percent_encoded => |percent_encoded| try w.writeAll(percent_encoded),
95 }131 }
96 }132 }
97133
98 pub fn formatPath(component: Component, w: *std.io.Writer) std.io.Writer.Error!void {134 pub fn formatPath(component: Component, w: *Writer) Writer.Error!void {
99 switch (component) {135 switch (component) {
100 .raw => |raw| try percentEncode(w, raw, isPathChar),136 .raw => |raw| try percentEncode(w, raw, isPathChar),
101 .percent_encoded => |percent_encoded| try w.writeAll(percent_encoded),137 .percent_encoded => |percent_encoded| try w.writeAll(percent_encoded),
102 }138 }
103 }139 }
104140
105 pub fn formatQuery(component: Component, w: *std.io.Writer) std.io.Writer.Error!void {141 pub fn formatQuery(component: Component, w: *Writer) Writer.Error!void {
106 switch (component) {142 switch (component) {
107 .raw => |raw| try percentEncode(w, raw, isQueryChar),143 .raw => |raw| try percentEncode(w, raw, isQueryChar),
108 .percent_encoded => |percent_encoded| try w.writeAll(percent_encoded),144 .percent_encoded => |percent_encoded| try w.writeAll(percent_encoded),
109 }145 }
110 }146 }
111147
112 pub fn formatFragment(component: Component, w: *std.io.Writer) std.io.Writer.Error!void {148 pub fn formatFragment(component: Component, w: *Writer) Writer.Error!void {
113 switch (component) {149 switch (component) {
114 .raw => |raw| try percentEncode(w, raw, isFragmentChar),150 .raw => |raw| try percentEncode(w, raw, isFragmentChar),
115 .percent_encoded => |percent_encoded| try w.writeAll(percent_encoded),151 .percent_encoded => |percent_encoded| try w.writeAll(percent_encoded),
116 }152 }
117 }153 }
118154
119 pub fn percentEncode(w: *std.io.Writer, raw: []const u8, comptime isValidChar: fn (u8) bool) std.io.Writer.Error!void {155 pub fn percentEncode(w: *Writer, raw: []const u8, comptime isValidChar: fn (u8) bool) Writer.Error!void {
120 var start: usize = 0;156 var start: usize = 0;
121 for (raw, 0..) |char, index| {157 for (raw, 0..) |char, index| {
122 if (isValidChar(char)) continue;158 if (isValidChar(char)) continue;
...@@ -165,17 +201,15 @@ pub const ParseError = error{ UnexpectedCharacter, InvalidFormat, InvalidPort };...@@ -165,17 +201,15 @@ pub const ParseError = error{ UnexpectedCharacter, InvalidFormat, InvalidPort };
165/// The return value will contain strings pointing into the original `text`.201/// The return value will contain strings pointing into the original `text`.
166/// Each component that is provided, will be non-`null`.202/// Each component that is provided, will be non-`null`.
167pub fn parseAfterScheme(scheme: []const u8, text: []const u8) ParseError!Uri {203pub fn parseAfterScheme(scheme: []const u8, text: []const u8) ParseError!Uri {
168 var reader = SliceReader{ .slice = text };
169
170 var uri: Uri = .{ .scheme = scheme, .path = undefined };204 var uri: Uri = .{ .scheme = scheme, .path = undefined };
205 var i: usize = 0;
171206
172 if (reader.peekPrefix("//")) a: { // authority part207 if (std.mem.startsWith(u8, text, "//")) a: {
173 std.debug.assert(reader.get().? == '/');208 i = std.mem.indexOfAnyPos(u8, text, 2, &authority_sep) orelse text.len;
174 std.debug.assert(reader.get().? == '/');209 const authority = text[2..i];
175
176 const authority = reader.readUntil(isAuthoritySeparator);
177 if (authority.len == 0) {210 if (authority.len == 0) {
178 if (reader.peekPrefix("/")) break :a else return error.InvalidFormat;211 if (!std.mem.startsWith(u8, text[2..], "/")) return error.InvalidFormat;
212 break :a;
179 }213 }
180214
181 var start_of_host: usize = 0;215 var start_of_host: usize = 0;
...@@ -225,26 +259,28 @@ pub fn parseAfterScheme(scheme: []const u8, text: []const u8) ParseError!Uri {...@@ -225,26 +259,28 @@ pub fn parseAfterScheme(scheme: []const u8, text: []const u8) ParseError!Uri {
225 uri.host = .{ .percent_encoded = authority[start_of_host..end_of_host] };259 uri.host = .{ .percent_encoded = authority[start_of_host..end_of_host] };
226 }260 }
227261
228 uri.path = .{ .percent_encoded = reader.readUntil(isPathSeparator) };262 const path_start = i;
263 i = std.mem.indexOfAnyPos(u8, text, path_start, &path_sep) orelse text.len;
264 uri.path = .{ .percent_encoded = text[path_start..i] };
229265
230 if ((reader.peek() orelse 0) == '?') { // query part266 if (std.mem.startsWith(u8, text[i..], "?")) {
231 std.debug.assert(reader.get().? == '?');267 const query_start = i + 1;
232 uri.query = .{ .percent_encoded = reader.readUntil(isQuerySeparator) };268 i = std.mem.indexOfScalarPos(u8, text, query_start, '#') orelse text.len;
269 uri.query = .{ .percent_encoded = text[query_start..i] };
233 }270 }
234271
235 if ((reader.peek() orelse 0) == '#') { // fragment part272 if (std.mem.startsWith(u8, text[i..], "#")) {
236 std.debug.assert(reader.get().? == '#');273 uri.fragment = .{ .percent_encoded = text[i + 1 ..] };
237 uri.fragment = .{ .percent_encoded = reader.readUntilEof() };
238 }274 }
239275
240 return uri;276 return uri;
241}277}
242278
243pub fn format(uri: *const Uri, writer: *std.io.Writer) std.io.Writer.Error!void {279pub fn format(uri: *const Uri, writer: *Writer) Writer.Error!void {
244 return writeToStream(uri, writer, .all);280 return writeToStream(uri, writer, .all);
245}281}
246282
247pub fn writeToStream(uri: *const Uri, writer: *std.io.Writer, flags: Format.Flags) std.io.Writer.Error!void {283pub fn writeToStream(uri: *const Uri, writer: *Writer, flags: Format.Flags) Writer.Error!void {
248 if (flags.scheme) {284 if (flags.scheme) {
249 try writer.print("{s}:", .{uri.scheme});285 try writer.print("{s}:", .{uri.scheme});
250 if (flags.authority and uri.host != null) {286 if (flags.authority and uri.host != null) {
...@@ -318,7 +354,7 @@ pub const Format = struct {...@@ -318,7 +354,7 @@ pub const Format = struct {
318 };354 };
319 };355 };
320356
321 pub fn default(f: Format, writer: *std.io.Writer) std.io.Writer.Error!void {357 pub fn default(f: Format, writer: *Writer) Writer.Error!void {
322 return writeToStream(f.uri, writer, f.flags);358 return writeToStream(f.uri, writer, f.flags);
323 }359 }
324};360};
...@@ -327,41 +363,34 @@ pub fn fmt(uri: *const Uri, flags: Format.Flags) std.fmt.Formatter(Format, Forma...@@ -327,41 +363,34 @@ pub fn fmt(uri: *const Uri, flags: Format.Flags) std.fmt.Formatter(Format, Forma
327 return .{ .data = .{ .uri = uri, .flags = flags } };363 return .{ .data = .{ .uri = uri, .flags = flags } };
328}364}
329365
330/// Parses the URI or returns an error.366/// The return value will contain strings pointing into the original `text`.
331/// The return value will contain strings pointing into the367/// Each component that is provided will be non-`null`.
332/// original `text`. Each component that is provided, will be non-`null`.
333pub fn parse(text: []const u8) ParseError!Uri {368pub fn parse(text: []const u8) ParseError!Uri {
334 var reader: SliceReader = .{ .slice = text };369 const end = for (text, 0..) |byte, i| {
335 const scheme = reader.readWhile(isSchemeChar);370 if (!isSchemeChar(byte)) break i;
336371 } else text.len;
337 // after the scheme, a ':' must appear372 // After the scheme, a ':' must appear.
338 if (reader.get()) |c| {373 if (end >= text.len) return error.InvalidFormat;
339 if (c != ':')374 if (text[end] != ':') return error.UnexpectedCharacter;
340 return error.UnexpectedCharacter;375 return parseAfterScheme(text[0..end], text[end + 1 ..]);
341 } else {
342 return error.InvalidFormat;
343 }
344
345 return parseAfterScheme(scheme, reader.readUntilEof());
346}376}
347377
348pub const ResolveInPlaceError = ParseError || error{NoSpaceLeft};378pub const ResolveInPlaceError = ParseError || error{NoSpaceLeft};
349379
350/// Resolves a URI against a base URI, conforming to RFC 3986, Section 5.380/// Resolves a URI against a base URI, conforming to
351/// Copies `new` to the beginning of `aux_buf.*`, allowing the slices to overlap,381/// [RFC 3986, Section 5](https://www.rfc-editor.org/rfc/rfc3986#section-5)
352/// then parses `new` as a URI, and then resolves the path in place.382///
383/// Assumes new location is already copied to the beginning of `aux_buf.*`.
384/// Parses that new location as a URI, and then resolves the path in place.
385///
353/// If a merge needs to take place, the newly constructed path will be stored386/// If a merge needs to take place, the newly constructed path will be stored
354/// in `aux_buf.*` just after the copied `new`, and `aux_buf.*` will be modified387/// in `aux_buf.*` just after the copied location, and `aux_buf.*` will be
355/// to only contain the remaining unused space.388/// modified to only contain the remaining unused space.
356pub fn resolve_inplace(base: Uri, new: []const u8, aux_buf: *[]u8) ResolveInPlaceError!Uri {389pub fn resolveInPlace(base: Uri, new_len: usize, aux_buf: *[]u8) ResolveInPlaceError!Uri {
357 std.mem.copyForwards(u8, aux_buf.*, new);390 const new = aux_buf.*[0..new_len];
358 // At this point, new is an invalid pointer.391 const new_parsed = parse(new) catch |err| (parseAfterScheme("", new) catch return err);
359 const new_mut = aux_buf.*[0..new.len];392 aux_buf.* = aux_buf.*[new_len..];
360 aux_buf.* = aux_buf.*[new.len..];393 // As you can see above, `new` is not a const pointer.
361
362 const new_parsed = parse(new_mut) catch |err|
363 (parseAfterScheme("", new_mut) catch return err);
364 // As you can see above, `new_mut` is not a const pointer.
365 const new_path: []u8 = @constCast(new_parsed.path.percent_encoded);394 const new_path: []u8 = @constCast(new_parsed.path.percent_encoded);
366395
367 if (new_parsed.scheme.len > 0) return .{396 if (new_parsed.scheme.len > 0) return .{
...@@ -461,7 +490,7 @@ test remove_dot_segments {...@@ -461,7 +490,7 @@ test remove_dot_segments {
461490
462/// 5.2.3. Merge Paths491/// 5.2.3. Merge Paths
463fn merge_paths(base: Component, new: []u8, aux_buf: *[]u8) error{NoSpaceLeft}!Component {492fn merge_paths(base: Component, new: []u8, aux_buf: *[]u8) error{NoSpaceLeft}!Component {
464 var aux: std.io.Writer = .fixed(aux_buf.*);493 var aux: Writer = .fixed(aux_buf.*);
465 if (!base.isEmpty()) {494 if (!base.isEmpty()) {
466 base.formatPath(&aux) catch return error.NoSpaceLeft;495 base.formatPath(&aux) catch return error.NoSpaceLeft;
467 aux.end = std.mem.lastIndexOfScalar(u8, aux.buffered(), '/') orelse return remove_dot_segments(new);496 aux.end = std.mem.lastIndexOfScalar(u8, aux.buffered(), '/') orelse return remove_dot_segments(new);
...@@ -472,59 +501,6 @@ fn merge_paths(base: Component, new: []u8, aux_buf: *[]u8) error{NoSpaceLeft}!Co...@@ -472,59 +501,6 @@ fn merge_paths(base: Component, new: []u8, aux_buf: *[]u8) error{NoSpaceLeft}!Co
472 return merged_path;501 return merged_path;
473}502}
474503
475const SliceReader = struct {
476 const Self = @This();
477
478 slice: []const u8,
479 offset: usize = 0,
480
481 fn get(self: *Self) ?u8 {
482 if (self.offset >= self.slice.len)
483 return null;
484 const c = self.slice[self.offset];
485 self.offset += 1;
486 return c;
487 }
488
489 fn peek(self: Self) ?u8 {
490 if (self.offset >= self.slice.len)
491 return null;
492 return self.slice[self.offset];
493 }
494
495 fn readWhile(self: *Self, comptime predicate: fn (u8) bool) []const u8 {
496 const start = self.offset;
497 var end = start;
498 while (end < self.slice.len and predicate(self.slice[end])) {
499 end += 1;
500 }
501 self.offset = end;
502 return self.slice[start..end];
503 }
504
505 fn readUntil(self: *Self, comptime predicate: fn (u8) bool) []const u8 {
506 const start = self.offset;
507 var end = start;
508 while (end < self.slice.len and !predicate(self.slice[end])) {
509 end += 1;
510 }
511 self.offset = end;
512 return self.slice[start..end];
513 }
514
515 fn readUntilEof(self: *Self) []const u8 {
516 const start = self.offset;
517 self.offset = self.slice.len;
518 return self.slice[start..];
519 }
520
521 fn peekPrefix(self: Self, prefix: []const u8) bool {
522 if (self.offset + prefix.len > self.slice.len)
523 return false;
524 return std.mem.eql(u8, self.slice[self.offset..][0..prefix.len], prefix);
525 }
526};
527
528/// scheme = ALPHA *( ALPHA / DIGIT / "+" / "-" / "." )504/// scheme = ALPHA *( ALPHA / DIGIT / "+" / "-" / "." )
529fn isSchemeChar(c: u8) bool {505fn isSchemeChar(c: u8) bool {
530 return switch (c) {506 return switch (c) {
...@@ -533,19 +509,6 @@ fn isSchemeChar(c: u8) bool {...@@ -533,19 +509,6 @@ fn isSchemeChar(c: u8) bool {
533 };509 };
534}510}
535511
536/// reserved = gen-delims / sub-delims
537fn isReserved(c: u8) bool {
538 return isGenLimit(c) or isSubLimit(c);
539}
540
541/// gen-delims = ":" / "/" / "?" / "#" / "[" / "]" / "@"
542fn isGenLimit(c: u8) bool {
543 return switch (c) {
544 ':', ',', '?', '#', '[', ']', '@' => true,
545 else => false,
546 };
547}
548
549/// sub-delims = "!" / "$" / "&" / "'" / "(" / ")"512/// sub-delims = "!" / "$" / "&" / "'" / "(" / ")"
550/// / "*" / "+" / "," / ";" / "="513/// / "*" / "+" / "," / ";" / "="
551fn isSubLimit(c: u8) bool {514fn isSubLimit(c: u8) bool {
...@@ -585,26 +548,8 @@ fn isQueryChar(c: u8) bool {...@@ -585,26 +548,8 @@ fn isQueryChar(c: u8) bool {
585548
586const isFragmentChar = isQueryChar;549const isFragmentChar = isQueryChar;
587550
588fn isAuthoritySeparator(c: u8) bool {551const authority_sep: [3]u8 = .{ '/', '?', '#' };
589 return switch (c) {552const path_sep: [2]u8 = .{ '?', '#' };
590 '/', '?', '#' => true,
591 else => false,
592 };
593}
594
595fn isPathSeparator(c: u8) bool {
596 return switch (c) {
597 '?', '#' => true,
598 else => false,
599 };
600}
601
602fn isQuerySeparator(c: u8) bool {
603 return switch (c) {
604 '#' => true,
605 else => false,
606 };
607}
608553
609test "basic" {554test "basic" {
610 const parsed = try parse("https://ziglang.org/download");555 const parsed = try parse("https://ziglang.org/download");
lib/std/crypto/tls.zig+106-99
...@@ -49,8 +49,8 @@ pub const hello_retry_request_sequence = [32]u8{...@@ -49,8 +49,8 @@ pub const hello_retry_request_sequence = [32]u8{
49};49};
5050
51pub const close_notify_alert = [_]u8{51pub const close_notify_alert = [_]u8{
52 @intFromEnum(AlertLevel.warning),52 @intFromEnum(Alert.Level.warning),
53 @intFromEnum(AlertDescription.close_notify),53 @intFromEnum(Alert.Description.close_notify),
54};54};
5555
56pub const ProtocolVersion = enum(u16) {56pub const ProtocolVersion = enum(u16) {
...@@ -138,103 +138,108 @@ pub const ExtensionType = enum(u16) {...@@ -138,103 +138,108 @@ pub const ExtensionType = enum(u16) {
138 _,138 _,
139};139};
140140
141pub const AlertLevel = enum(u8) {141pub const Alert = struct {
142 warning = 1,142 level: Level,
143 fatal = 2,143 description: Description,
144 _,
145};
146144
147pub const AlertDescription = enum(u8) {145 pub const Level = enum(u8) {
148 pub const Error = error{146 warning = 1,
149 TlsAlertUnexpectedMessage,147 fatal = 2,
150 TlsAlertBadRecordMac,148 _,
151 TlsAlertRecordOverflow,
152 TlsAlertHandshakeFailure,
153 TlsAlertBadCertificate,
154 TlsAlertUnsupportedCertificate,
155 TlsAlertCertificateRevoked,
156 TlsAlertCertificateExpired,
157 TlsAlertCertificateUnknown,
158 TlsAlertIllegalParameter,
159 TlsAlertUnknownCa,
160 TlsAlertAccessDenied,
161 TlsAlertDecodeError,
162 TlsAlertDecryptError,
163 TlsAlertProtocolVersion,
164 TlsAlertInsufficientSecurity,
165 TlsAlertInternalError,
166 TlsAlertInappropriateFallback,
167 TlsAlertMissingExtension,
168 TlsAlertUnsupportedExtension,
169 TlsAlertUnrecognizedName,
170 TlsAlertBadCertificateStatusResponse,
171 TlsAlertUnknownPskIdentity,
172 TlsAlertCertificateRequired,
173 TlsAlertNoApplicationProtocol,
174 TlsAlertUnknown,
175 };149 };
176150
177 close_notify = 0,151 pub const Description = enum(u8) {
178 unexpected_message = 10,152 pub const Error = error{
179 bad_record_mac = 20,153 TlsAlertUnexpectedMessage,
180 record_overflow = 22,154 TlsAlertBadRecordMac,
181 handshake_failure = 40,155 TlsAlertRecordOverflow,
182 bad_certificate = 42,156 TlsAlertHandshakeFailure,
183 unsupported_certificate = 43,157 TlsAlertBadCertificate,
184 certificate_revoked = 44,158 TlsAlertUnsupportedCertificate,
185 certificate_expired = 45,159 TlsAlertCertificateRevoked,
186 certificate_unknown = 46,160 TlsAlertCertificateExpired,
187 illegal_parameter = 47,161 TlsAlertCertificateUnknown,
188 unknown_ca = 48,162 TlsAlertIllegalParameter,
189 access_denied = 49,163 TlsAlertUnknownCa,
190 decode_error = 50,164 TlsAlertAccessDenied,
191 decrypt_error = 51,165 TlsAlertDecodeError,
192 protocol_version = 70,166 TlsAlertDecryptError,
193 insufficient_security = 71,167 TlsAlertProtocolVersion,
194 internal_error = 80,168 TlsAlertInsufficientSecurity,
195 inappropriate_fallback = 86,169 TlsAlertInternalError,
196 user_canceled = 90,170 TlsAlertInappropriateFallback,
197 missing_extension = 109,171 TlsAlertMissingExtension,
198 unsupported_extension = 110,172 TlsAlertUnsupportedExtension,
199 unrecognized_name = 112,173 TlsAlertUnrecognizedName,
200 bad_certificate_status_response = 113,174 TlsAlertBadCertificateStatusResponse,
201 unknown_psk_identity = 115,175 TlsAlertUnknownPskIdentity,
202 certificate_required = 116,176 TlsAlertCertificateRequired,
203 no_application_protocol = 120,177 TlsAlertNoApplicationProtocol,
204 _,178 TlsAlertUnknown,
179 };
205180
206 pub fn toError(alert: AlertDescription) Error!void {181 close_notify = 0,
207 switch (alert) {182 unexpected_message = 10,
208 .close_notify => {}, // not an error183 bad_record_mac = 20,
209 .unexpected_message => return error.TlsAlertUnexpectedMessage,184 record_overflow = 22,
210 .bad_record_mac => return error.TlsAlertBadRecordMac,185 handshake_failure = 40,
211 .record_overflow => return error.TlsAlertRecordOverflow,186 bad_certificate = 42,
212 .handshake_failure => return error.TlsAlertHandshakeFailure,187 unsupported_certificate = 43,
213 .bad_certificate => return error.TlsAlertBadCertificate,188 certificate_revoked = 44,
214 .unsupported_certificate => return error.TlsAlertUnsupportedCertificate,189 certificate_expired = 45,
215 .certificate_revoked => return error.TlsAlertCertificateRevoked,190 certificate_unknown = 46,
216 .certificate_expired => return error.TlsAlertCertificateExpired,191 illegal_parameter = 47,
217 .certificate_unknown => return error.TlsAlertCertificateUnknown,192 unknown_ca = 48,
218 .illegal_parameter => return error.TlsAlertIllegalParameter,193 access_denied = 49,
219 .unknown_ca => return error.TlsAlertUnknownCa,194 decode_error = 50,
220 .access_denied => return error.TlsAlertAccessDenied,195 decrypt_error = 51,
221 .decode_error => return error.TlsAlertDecodeError,196 protocol_version = 70,
222 .decrypt_error => return error.TlsAlertDecryptError,197 insufficient_security = 71,
223 .protocol_version => return error.TlsAlertProtocolVersion,198 internal_error = 80,
224 .insufficient_security => return error.TlsAlertInsufficientSecurity,199 inappropriate_fallback = 86,
225 .internal_error => return error.TlsAlertInternalError,200 user_canceled = 90,
226 .inappropriate_fallback => return error.TlsAlertInappropriateFallback,201 missing_extension = 109,
227 .user_canceled => {}, // not an error202 unsupported_extension = 110,
228 .missing_extension => return error.TlsAlertMissingExtension,203 unrecognized_name = 112,
229 .unsupported_extension => return error.TlsAlertUnsupportedExtension,204 bad_certificate_status_response = 113,
230 .unrecognized_name => return error.TlsAlertUnrecognizedName,205 unknown_psk_identity = 115,
231 .bad_certificate_status_response => return error.TlsAlertBadCertificateStatusResponse,206 certificate_required = 116,
232 .unknown_psk_identity => return error.TlsAlertUnknownPskIdentity,207 no_application_protocol = 120,
233 .certificate_required => return error.TlsAlertCertificateRequired,208 _,
234 .no_application_protocol => return error.TlsAlertNoApplicationProtocol,209
235 _ => return error.TlsAlertUnknown,210 pub fn toError(description: Description) Error!void {
211 switch (description) {
212 .close_notify => {}, // not an error
213 .unexpected_message => return error.TlsAlertUnexpectedMessage,
214 .bad_record_mac => return error.TlsAlertBadRecordMac,
215 .record_overflow => return error.TlsAlertRecordOverflow,
216 .handshake_failure => return error.TlsAlertHandshakeFailure,
217 .bad_certificate => return error.TlsAlertBadCertificate,
218 .unsupported_certificate => return error.TlsAlertUnsupportedCertificate,
219 .certificate_revoked => return error.TlsAlertCertificateRevoked,
220 .certificate_expired => return error.TlsAlertCertificateExpired,
221 .certificate_unknown => return error.TlsAlertCertificateUnknown,
222 .illegal_parameter => return error.TlsAlertIllegalParameter,
223 .unknown_ca => return error.TlsAlertUnknownCa,
224 .access_denied => return error.TlsAlertAccessDenied,
225 .decode_error => return error.TlsAlertDecodeError,
226 .decrypt_error => return error.TlsAlertDecryptError,
227 .protocol_version => return error.TlsAlertProtocolVersion,
228 .insufficient_security => return error.TlsAlertInsufficientSecurity,
229 .internal_error => return error.TlsAlertInternalError,
230 .inappropriate_fallback => return error.TlsAlertInappropriateFallback,
231 .user_canceled => {}, // not an error
232 .missing_extension => return error.TlsAlertMissingExtension,
233 .unsupported_extension => return error.TlsAlertUnsupportedExtension,
234 .unrecognized_name => return error.TlsAlertUnrecognizedName,
235 .bad_certificate_status_response => return error.TlsAlertBadCertificateStatusResponse,
236 .unknown_psk_identity => return error.TlsAlertUnknownPskIdentity,
237 .certificate_required => return error.TlsAlertCertificateRequired,
238 .no_application_protocol => return error.TlsAlertNoApplicationProtocol,
239 _ => return error.TlsAlertUnknown,
240 }
236 }241 }
237 }242 };
238};243};
239244
240pub const SignatureScheme = enum(u16) {245pub const SignatureScheme = enum(u16) {
...@@ -650,7 +655,7 @@ pub const Decoder = struct {...@@ -650,7 +655,7 @@ pub const Decoder = struct {
650 }655 }
651656
652 /// Use this function to increase `their_end`.657 /// Use this function to increase `their_end`.
653 pub fn readAtLeast(d: *Decoder, stream: anytype, their_amt: usize) !void {658 pub fn readAtLeast(d: *Decoder, stream: *std.io.Reader, their_amt: usize) !void {
654 assert(!d.disable_reads);659 assert(!d.disable_reads);
655 const existing_amt = d.cap - d.idx;660 const existing_amt = d.cap - d.idx;
656 d.their_end = d.idx + their_amt;661 d.their_end = d.idx + their_amt;
...@@ -658,14 +663,16 @@ pub const Decoder = struct {...@@ -658,14 +663,16 @@ pub const Decoder = struct {
658 const request_amt = their_amt - existing_amt;663 const request_amt = their_amt - existing_amt;
659 const dest = d.buf[d.cap..];664 const dest = d.buf[d.cap..];
660 if (request_amt > dest.len) return error.TlsRecordOverflow;665 if (request_amt > dest.len) return error.TlsRecordOverflow;
661 const actual_amt = try stream.readAtLeast(dest, request_amt);666 stream.readSlice(dest[0..request_amt]) catch |err| switch (err) {
662 if (actual_amt < request_amt) return error.TlsConnectionTruncated;667 error.EndOfStream => return error.TlsConnectionTruncated,
663 d.cap += actual_amt;668 error.ReadFailed => return error.ReadFailed,
669 };
670 d.cap += request_amt;
664 }671 }
665672
666 /// Same as `readAtLeast` but also increases `our_end` by exactly `our_amt`.673 /// Same as `readAtLeast` but also increases `our_end` by exactly `our_amt`.
667 /// Use when `our_amt` is calculated by us, not by them.674 /// Use when `our_amt` is calculated by us, not by them.
668 pub fn readAtLeastOurAmt(d: *Decoder, stream: anytype, our_amt: usize) !void {675 pub fn readAtLeastOurAmt(d: *Decoder, stream: *std.io.Reader, our_amt: usize) !void {
669 assert(!d.disable_reads);676 assert(!d.disable_reads);
670 try readAtLeast(d, stream, our_amt);677 try readAtLeast(d, stream, our_amt);
671 d.our_end = d.idx + our_amt;678 d.our_end = d.idx + our_amt;
lib/std/crypto/tls/Client.zig+417-805
...@@ -1,11 +1,15 @@...@@ -1,11 +1,15 @@
1const builtin = @import("builtin");
2const native_endian = builtin.cpu.arch.endian();
3
1const std = @import("../../std.zig");4const std = @import("../../std.zig");
2const tls = std.crypto.tls;5const tls = std.crypto.tls;
3const Client = @This();6const Client = @This();
4const net = std.net;
5const mem = std.mem;7const mem = std.mem;
6const crypto = std.crypto;8const crypto = std.crypto;
7const assert = std.debug.assert;9const assert = std.debug.assert;
8const Certificate = std.crypto.Certificate;10const Certificate = std.crypto.Certificate;
11const Reader = std.Io.Reader;
12const Writer = std.Io.Writer;
913
10const max_ciphertext_len = tls.max_ciphertext_len;14const max_ciphertext_len = tls.max_ciphertext_len;
11const hmacExpandLabel = tls.hmacExpandLabel;15const hmacExpandLabel = tls.hmacExpandLabel;
...@@ -13,44 +17,60 @@ const hkdfExpandLabel = tls.hkdfExpandLabel;...@@ -13,44 +17,60 @@ const hkdfExpandLabel = tls.hkdfExpandLabel;
13const int = tls.int;17const int = tls.int;
14const array = tls.array;18const array = tls.array;
1519
20/// The encrypted stream from the server to the client. Bytes are pulled from
21/// here via `reader`.
22///
23/// The buffer is asserted to have capacity at least `min_buffer_len`.
24input: *Reader,
25/// Decrypted stream from the server to the client.
26reader: Reader,
27
28/// The encrypted stream from the client to the server. Bytes are pushed here
29/// via `writer`.
30///
31/// The buffer is asserted to have capacity at least `min_buffer_len`.
32output: *Writer,
33/// The plaintext stream from the client to the server.
34writer: Writer,
35
36/// Populated when `error.TlsAlert` is returned.
37alert: ?tls.Alert = null,
38read_err: ?ReadError = null,
16tls_version: tls.ProtocolVersion,39tls_version: tls.ProtocolVersion,
17read_seq: u64,40read_seq: u64,
18write_seq: u64,41write_seq: u64,
19/// The starting index of cleartext bytes inside `partially_read_buffer`.
20partial_cleartext_idx: u15,
21/// The ending index of cleartext bytes inside `partially_read_buffer` as well
22/// as the starting index of ciphertext bytes.
23partial_ciphertext_idx: u15,
24/// The ending index of ciphertext bytes inside `partially_read_buffer`.
25partial_ciphertext_end: u15,
26/// When this is true, the stream may still not be at the end because there42/// When this is true, the stream may still not be at the end because there
27/// may be data in `partially_read_buffer`.43/// may be data in the input buffer.
28received_close_notify: bool,44received_close_notify: bool,
29/// By default, reaching the end-of-stream when reading from the server will
30/// cause `error.TlsConnectionTruncated` to be returned, unless a close_notify
31/// message has been received. By setting this flag to `true`, instead, the
32/// end-of-stream will be forwarded to the application layer above TLS.
33/// This makes the application vulnerable to truncation attacks unless the
34/// application layer itself verifies that the amount of data received equals
35/// the amount of data expected, such as HTTP with the Content-Length header.
36allow_truncation_attacks: bool,45allow_truncation_attacks: bool,
37application_cipher: tls.ApplicationCipher,46application_cipher: tls.ApplicationCipher,
38/// The size is enough to contain exactly one TLSCiphertext record.47
39/// This buffer is segmented into four parts:48/// If non-null, ssl secrets are logged to a stream. Creating such a log file
40/// 0. unused49/// allows other programs with access to that file to decrypt all traffic over
41/// 1. cleartext50/// this connection.
42/// 2. ciphertext51ssl_key_log: ?*SslKeyLog,
43/// 3. unused52
44/// The fields `partial_cleartext_idx`, `partial_ciphertext_idx`, and53pub const ReadError = error{
45/// `partial_ciphertext_end` describe the span of the segments.54 /// The alert description will be stored in `alert`.
46partially_read_buffer: [tls.max_ciphertext_record_len]u8,55 TlsAlert,
47/// If non-null, ssl secrets are logged to a file. Creating such a log file allows other56 TlsBadLength,
48/// programs with access to that file to decrypt all traffic over this connection.57 TlsBadRecordMac,
49ssl_key_log: ?struct {58 TlsConnectionTruncated,
59 TlsDecodeError,
60 TlsRecordOverflow,
61 TlsUnexpectedMessage,
62 TlsIllegalParameter,
63 TlsSequenceOverflow,
64 /// The buffer provided to the read function was not at least
65 /// `min_buffer_len`.
66 OutputBufferUndersize,
67};
68
69pub const SslKeyLog = struct {
50 client_key_seq: u64,70 client_key_seq: u64,
51 server_key_seq: u64,71 server_key_seq: u64,
52 client_random: [32]u8,72 client_random: [32]u8,
53 file: std.fs.File,73 writer: *Writer,
5474
55 fn clientCounter(key_log: *@This()) u64 {75 fn clientCounter(key_log: *@This()) u64 {
56 defer key_log.client_key_seq += 1;76 defer key_log.client_key_seq += 1;
...@@ -61,51 +81,12 @@ ssl_key_log: ?struct {...@@ -61,51 +81,12 @@ ssl_key_log: ?struct {
61 defer key_log.server_key_seq += 1;81 defer key_log.server_key_seq += 1;
62 return key_log.server_key_seq;82 return key_log.server_key_seq;
63 }83 }
64},
65
66/// This is an example of the type that is needed by the read and write
67/// functions. It can have any fields but it must at least have these
68/// functions.
69///
70/// Note that `std.net.Stream` conforms to this interface.
71///
72/// This declaration serves as documentation only.
73pub const StreamInterface = struct {
74 /// Can be any error set.
75 pub const ReadError = error{};
76
77 /// Returns the number of bytes read. The number read may be less than the
78 /// buffer space provided. End-of-stream is indicated by a return value of 0.
79 ///
80 /// The `iovecs` parameter is mutable because so that function may to
81 /// mutate the fields in order to handle partial reads from the underlying
82 /// stream layer.
83 pub fn readv(this: @This(), iovecs: []std.posix.iovec) ReadError!usize {
84 _ = .{ this, iovecs };
85 @panic("unimplemented");
86 }
87
88 /// Can be any error set.
89 pub const WriteError = error{};
90
91 /// Returns the number of bytes read, which may be less than the buffer
92 /// space provided. A short read does not indicate end-of-stream.
93 pub fn writev(this: @This(), iovecs: []const std.posix.iovec_const) WriteError!usize {
94 _ = .{ this, iovecs };
95 @panic("unimplemented");
96 }
97
98 /// Returns the number of bytes read, which may be less than the buffer
99 /// space provided, indicating end-of-stream.
100 /// The `iovecs` parameter is mutable in case this function needs to mutate
101 /// the fields in order to handle partial writes from the underlying layer.
102 pub fn writevAll(this: @This(), iovecs: []std.posix.iovec_const) WriteError!usize {
103 // This can be implemented in terms of writev, or specialized if desired.
104 _ = .{ this, iovecs };
105 @panic("unimplemented");
106 }
107};84};
10885
86/// The `Reader` supplied to `init` requires a buffer capacity
87/// at least this amount.
88pub const min_buffer_len = tls.max_ciphertext_record_len;
89
109pub const Options = struct {90pub const Options = struct {
110 /// How to perform host verification of server certificates.91 /// How to perform host verification of server certificates.
111 host: union(enum) {92 host: union(enum) {
...@@ -127,64 +108,85 @@ pub const Options = struct {...@@ -127,64 +108,85 @@ pub const Options = struct {
127 /// Verify that the server certificate is authorized by a given ca bundle.108 /// Verify that the server certificate is authorized by a given ca bundle.
128 bundle: Certificate.Bundle,109 bundle: Certificate.Bundle,
129 },110 },
130 /// If non-null, ssl secrets are logged to this file. Creating such a log file allows111 /// If non-null, ssl secrets are logged to this stream. Creating such a log file allows
131 /// other programs with access to that file to decrypt all traffic over this connection.112 /// other programs with access to that file to decrypt all traffic over this connection.
132 ssl_key_log_file: ?std.fs.File = null,113 ///
114 /// Only the `writer` field is observed during the handshake (`init`).
115 /// After that, the other fields are populated.
116 ssl_key_log: ?*SslKeyLog = null,
117 /// By default, reaching the end-of-stream when reading from the server will
118 /// cause `error.TlsConnectionTruncated` to be returned, unless a close_notify
119 /// message has been received. By setting this flag to `true`, instead, the
120 /// end-of-stream will be forwarded to the application layer above TLS.
121 ///
122 /// This makes the application vulnerable to truncation attacks unless the
123 /// application layer itself verifies that the amount of data received equals
124 /// the amount of data expected, such as HTTP with the Content-Length header.
125 allow_truncation_attacks: bool = false,
126 write_buffer: []u8,
127 read_buffer: []u8,
128 /// Populated when `error.TlsAlert` is returned from `init`.
129 alert: ?*tls.Alert = null,
133};130};
134131
135pub fn InitError(comptime Stream: type) type {132const InitError = error{
136 return std.mem.Allocator.Error || Stream.WriteError || Stream.ReadError || tls.AlertDescription.Error || error{133 WriteFailed,
137 InsufficientEntropy,134 ReadFailed,
138 DiskQuota,135 InsufficientEntropy,
139 LockViolation,136 DiskQuota,
140 NotOpenForWriting,137 LockViolation,
141 TlsUnexpectedMessage,138 NotOpenForWriting,
142 TlsIllegalParameter,139 /// The alert description will be stored in `alert`.
143 TlsDecryptFailure,140 TlsAlert,
144 TlsRecordOverflow,141 TlsUnexpectedMessage,
145 TlsBadRecordMac,142 TlsIllegalParameter,
146 CertificateFieldHasInvalidLength,143 TlsDecryptFailure,
147 CertificateHostMismatch,144 TlsRecordOverflow,
148 CertificatePublicKeyInvalid,145 TlsBadRecordMac,
149 CertificateExpired,146 CertificateFieldHasInvalidLength,
150 CertificateFieldHasWrongDataType,147 CertificateHostMismatch,
151 CertificateIssuerMismatch,148 CertificatePublicKeyInvalid,
152 CertificateNotYetValid,149 CertificateExpired,
153 CertificateSignatureAlgorithmMismatch,150 CertificateFieldHasWrongDataType,
154 CertificateSignatureAlgorithmUnsupported,151 CertificateIssuerMismatch,
155 CertificateSignatureInvalid,152 CertificateNotYetValid,
156 CertificateSignatureInvalidLength,153 CertificateSignatureAlgorithmMismatch,
157 CertificateSignatureNamedCurveUnsupported,154 CertificateSignatureAlgorithmUnsupported,
158 CertificateSignatureUnsupportedBitCount,155 CertificateSignatureInvalid,
159 TlsCertificateNotVerified,156 CertificateSignatureInvalidLength,
160 TlsBadSignatureScheme,157 CertificateSignatureNamedCurveUnsupported,
161 TlsBadRsaSignatureBitCount,158 CertificateSignatureUnsupportedBitCount,
162 InvalidEncoding,159 TlsCertificateNotVerified,
163 IdentityElement,160 TlsBadSignatureScheme,
164 SignatureVerificationFailed,161 TlsBadRsaSignatureBitCount,
165 TlsDecryptError,162 InvalidEncoding,
166 TlsConnectionTruncated,163 IdentityElement,
167 TlsDecodeError,164 SignatureVerificationFailed,
168 UnsupportedCertificateVersion,165 TlsDecryptError,
169 CertificateTimeInvalid,166 TlsConnectionTruncated,
170 CertificateHasUnrecognizedObjectId,167 TlsDecodeError,
171 CertificateHasInvalidBitString,168 UnsupportedCertificateVersion,
172 MessageTooLong,169 CertificateTimeInvalid,
173 NegativeIntoUnsigned,170 CertificateHasUnrecognizedObjectId,
174 TargetTooSmall,171 CertificateHasInvalidBitString,
175 BufferTooSmall,172 MessageTooLong,
176 InvalidSignature,173 NegativeIntoUnsigned,
177 NotSquare,174 TargetTooSmall,
178 NonCanonical,175 BufferTooSmall,
179 WeakPublicKey,176 InvalidSignature,
180 };177 NotSquare,
181}178 NonCanonical,
179 WeakPublicKey,
180};
182181
183/// Initiates a TLS handshake and establishes a TLSv1.2 or TLSv1.3 session with `stream`, which182/// Initiates a TLS handshake and establishes a TLSv1.2 or TLSv1.3 session.
184/// must conform to `StreamInterface`.
185///183///
186/// `host` is only borrowed during this function call.184/// `host` is only borrowed during this function call.
187pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client {185///
186/// `input` is asserted to have buffer capacity at least `min_buffer_len`.
187pub fn init(input: *Reader, output: *Writer, options: Options) InitError!Client {
188 assert(input.buffer.len >= min_buffer_len);
189 assert(output.buffer.len >= min_buffer_len);
188 const host = switch (options.host) {190 const host = switch (options.host) {
189 .no_verification => "",191 .no_verification => "",
190 .explicit => |host| host,192 .explicit => |host| host,
...@@ -276,11 +278,9 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client...@@ -276,11 +278,9 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client
276 };278 };
277279
278 {280 {
279 var iovecs = [_]std.posix.iovec_const{281 var iovecs: [2][]const u8 = .{ cleartext_header, host };
280 .{ .base = cleartext_header.ptr, .len = cleartext_header.len },282 try output.writeVecAll(iovecs[0..if (host.len == 0) 1 else 2]);
281 .{ .base = host.ptr, .len = host.len },283 try output.flush();
282 };
283 try stream.writevAll(iovecs[0..if (host.len == 0) 1 else 2]);
284 }284 }
285285
286 var tls_version: tls.ProtocolVersion = undefined;286 var tls_version: tls.ProtocolVersion = undefined;
...@@ -329,20 +329,28 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client...@@ -329,20 +329,28 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client
329 var cleartext_fragment_start: usize = 0;329 var cleartext_fragment_start: usize = 0;
330 var cleartext_fragment_end: usize = 0;330 var cleartext_fragment_end: usize = 0;
331 var cleartext_bufs: [2][tls.max_ciphertext_inner_record_len]u8 = undefined;331 var cleartext_bufs: [2][tls.max_ciphertext_inner_record_len]u8 = undefined;
332 var handshake_buffer: [tls.max_ciphertext_record_len]u8 = undefined;
333 var d: tls.Decoder = .{ .buf = &handshake_buffer };
334 fragment: while (true) {332 fragment: while (true) {
335 try d.readAtLeastOurAmt(stream, tls.record_header_len);333 // Ensure the input buffer pointer is stable in this scope.
336 const record_header = d.buf[d.idx..][0..tls.record_header_len];334 input.rebase(tls.max_ciphertext_record_len) catch |err| switch (err) {
337 const record_ct = d.decode(tls.ContentType);335 error.EndOfStream => {}, // We have assurance the remainder of stream can be buffered.
338 d.skip(2); // legacy_version336 };
339 const record_len = d.decode(u16);337 const record_header = input.peek(tls.record_header_len) catch |err| switch (err) {
340 try d.readAtLeast(stream, record_len);338 error.EndOfStream => return error.TlsConnectionTruncated,
341 var record_decoder = try d.sub(record_len);339 error.ReadFailed => return error.ReadFailed,
340 };
341 const record_ct = input.takeEnumNonexhaustive(tls.ContentType, .big) catch unreachable; // already peeked
342 input.toss(2); // legacy_version
343 const record_len = input.takeInt(u16, .big) catch unreachable; // already peeked
344 if (record_len > tls.max_ciphertext_len) return error.TlsRecordOverflow;
345 const record_buffer = input.take(record_len) catch |err| switch (err) {
346 error.EndOfStream => return error.TlsConnectionTruncated,
347 error.ReadFailed => return error.ReadFailed,
348 };
349 var record_decoder: tls.Decoder = .fromTheirSlice(record_buffer);
342 var ctd, const ct = content: switch (cipher_state) {350 var ctd, const ct = content: switch (cipher_state) {
343 .cleartext => .{ record_decoder, record_ct },351 .cleartext => .{ record_decoder, record_ct },
344 .handshake => {352 .handshake => {
345 std.debug.assert(tls_version == .tls_1_3);353 assert(tls_version == .tls_1_3);
346 if (record_ct != .application_data) return error.TlsUnexpectedMessage;354 if (record_ct != .application_data) return error.TlsUnexpectedMessage;
347 try record_decoder.ensure(record_len);355 try record_decoder.ensure(record_len);
348 const cleartext_buf = &cleartext_bufs[cert_buf_index % 2];356 const cleartext_buf = &cleartext_bufs[cert_buf_index % 2];
...@@ -374,7 +382,7 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client...@@ -374,7 +382,7 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client
374 break :content .{ tls.Decoder.fromTheirSlice(@constCast(cleartext_buf[cleartext_fragment_start..cleartext_fragment_end])), ct };382 break :content .{ tls.Decoder.fromTheirSlice(@constCast(cleartext_buf[cleartext_fragment_start..cleartext_fragment_end])), ct };
375 },383 },
376 .application => {384 .application => {
377 std.debug.assert(tls_version == .tls_1_2);385 assert(tls_version == .tls_1_2);
378 if (record_ct != .handshake) return error.TlsUnexpectedMessage;386 if (record_ct != .handshake) return error.TlsUnexpectedMessage;
379 try record_decoder.ensure(record_len);387 try record_decoder.ensure(record_len);
380 const cleartext_buf = &cleartext_bufs[cert_buf_index % 2];388 const cleartext_buf = &cleartext_bufs[cert_buf_index % 2];
...@@ -412,14 +420,11 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client...@@ -412,14 +420,11 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client
412 switch (ct) {420 switch (ct) {
413 .alert => {421 .alert => {
414 ctd.ensure(2) catch continue :fragment;422 ctd.ensure(2) catch continue :fragment;
415 const level = ctd.decode(tls.AlertLevel);423 if (options.alert) |a| a.* = .{
416 const desc = ctd.decode(tls.AlertDescription);424 .level = ctd.decode(tls.Alert.Level),
417 _ = level;425 .description = ctd.decode(tls.Alert.Description),
418426 };
419 // if this isn't a error alert, then it's a closure alert, which makes no sense in a handshake427 return error.TlsAlert;
420 try desc.toError();
421 // TODO: handle server-side closures
422 return error.TlsUnexpectedMessage;
423 },428 },
424 .change_cipher_spec => {429 .change_cipher_spec => {
425 ctd.ensure(1) catch continue :fragment;430 ctd.ensure(1) catch continue :fragment;
...@@ -533,7 +538,7 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client...@@ -533,7 +538,7 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client
533 pv.master_secret = P.Hkdf.extract(&ap_derived_secret, &zeroes);538 pv.master_secret = P.Hkdf.extract(&ap_derived_secret, &zeroes);
534 const client_secret = hkdfExpandLabel(P.Hkdf, pv.handshake_secret, "c hs traffic", &hello_hash, P.Hash.digest_length);539 const client_secret = hkdfExpandLabel(P.Hkdf, pv.handshake_secret, "c hs traffic", &hello_hash, P.Hash.digest_length);
535 const server_secret = hkdfExpandLabel(P.Hkdf, pv.handshake_secret, "s hs traffic", &hello_hash, P.Hash.digest_length);540 const server_secret = hkdfExpandLabel(P.Hkdf, pv.handshake_secret, "s hs traffic", &hello_hash, P.Hash.digest_length);
536 if (options.ssl_key_log_file) |key_log_file| logSecrets(key_log_file, .{541 if (options.ssl_key_log) |key_log| logSecrets(key_log.writer, .{
537 .client_random = &client_hello_rand,542 .client_random = &client_hello_rand,
538 }, .{543 }, .{
539 .SERVER_HANDSHAKE_TRAFFIC_SECRET = &server_secret,544 .SERVER_HANDSHAKE_TRAFFIC_SECRET = &server_secret,
...@@ -707,7 +712,7 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client...@@ -707,7 +712,7 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client
707 &client_hello_rand,712 &client_hello_rand,
708 &server_hello_rand,713 &server_hello_rand,
709 }, 48);714 }, 48);
710 if (options.ssl_key_log_file) |key_log_file| logSecrets(key_log_file, .{715 if (options.ssl_key_log) |key_log| logSecrets(key_log.writer, .{
711 .client_random = &client_hello_rand,716 .client_random = &client_hello_rand,
712 }, .{717 }, .{
713 .CLIENT_RANDOM = &master_secret,718 .CLIENT_RANDOM = &master_secret,
...@@ -755,11 +760,13 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client...@@ -755,11 +760,13 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client
755 nonce,760 nonce,
756 pv.app_cipher.client_write_key,761 pv.app_cipher.client_write_key,
757 );762 );
758 const all_msgs = client_key_exchange_msg ++ client_change_cipher_spec_msg ++ client_verify_msg;763 var all_msgs_vec: [3][]const u8 = .{
759 var all_msgs_vec = [_]std.posix.iovec_const{764 &client_key_exchange_msg,
760 .{ .base = &all_msgs, .len = all_msgs.len },765 &client_change_cipher_spec_msg,
766 &client_verify_msg,
761 };767 };
762 try stream.writevAll(&all_msgs_vec);768 try output.writeVecAll(&all_msgs_vec);
769 try output.flush();
763 },770 },
764 }771 }
765 write_seq += 1;772 write_seq += 1;
...@@ -820,15 +827,16 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client...@@ -820,15 +827,16 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client
820 const nonce = pv.client_handshake_iv;827 const nonce = pv.client_handshake_iv;
821 P.AEAD.encrypt(ciphertext, auth_tag, &out_cleartext, ad, nonce, pv.client_handshake_key);828 P.AEAD.encrypt(ciphertext, auth_tag, &out_cleartext, ad, nonce, pv.client_handshake_key);
822829
823 const all_msgs = client_change_cipher_spec_msg ++ finished_msg;830 var all_msgs_vec: [2][]const u8 = .{
824 var all_msgs_vec = [_]std.posix.iovec_const{831 &client_change_cipher_spec_msg,
825 .{ .base = &all_msgs, .len = all_msgs.len },832 &finished_msg,
826 };833 };
827 try stream.writevAll(&all_msgs_vec);834 try output.writeVecAll(&all_msgs_vec);
835 try output.flush();
828836
829 const client_secret = hkdfExpandLabel(P.Hkdf, pv.master_secret, "c ap traffic", &handshake_hash, P.Hash.digest_length);837 const client_secret = hkdfExpandLabel(P.Hkdf, pv.master_secret, "c ap traffic", &handshake_hash, P.Hash.digest_length);
830 const server_secret = hkdfExpandLabel(P.Hkdf, pv.master_secret, "s ap traffic", &handshake_hash, P.Hash.digest_length);838 const server_secret = hkdfExpandLabel(P.Hkdf, pv.master_secret, "s ap traffic", &handshake_hash, P.Hash.digest_length);
831 if (options.ssl_key_log_file) |key_log_file| logSecrets(key_log_file, .{839 if (options.ssl_key_log) |key_log| logSecrets(key_log.writer, .{
832 .counter = key_seq,840 .counter = key_seq,
833 .client_random = &client_hello_rand,841 .client_random = &client_hello_rand,
834 }, .{842 }, .{
...@@ -855,8 +863,28 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client...@@ -855,8 +863,28 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client
855 else => unreachable,863 else => unreachable,
856 },864 },
857 };865 };
858 const leftover = d.rest();866 if (options.ssl_key_log) |ssl_key_log| ssl_key_log.* = .{
859 var client: Client = .{867 .client_key_seq = key_seq,
868 .server_key_seq = key_seq,
869 .client_random = client_hello_rand,
870 .writer = ssl_key_log.writer,
871 };
872 return .{
873 .input = input,
874 .reader = .{
875 .buffer = options.read_buffer,
876 .vtable = &.{ .stream = stream },
877 .seek = 0,
878 .end = 0,
879 },
880 .output = output,
881 .writer = .{
882 .buffer = options.write_buffer,
883 .vtable = &.{
884 .drain = drain,
885 .flush = flush,
886 },
887 },
860 .tls_version = tls_version,888 .tls_version = tls_version,
861 .read_seq = switch (tls_version) {889 .read_seq = switch (tls_version) {
862 .tls_1_3 => 0,890 .tls_1_3 => 0,
...@@ -868,22 +896,11 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client...@@ -868,22 +896,11 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client
868 .tls_1_2 => write_seq,896 .tls_1_2 => write_seq,
869 else => unreachable,897 else => unreachable,
870 },898 },
871 .partial_cleartext_idx = 0,
872 .partial_ciphertext_idx = 0,
873 .partial_ciphertext_end = @intCast(leftover.len),
874 .received_close_notify = false,899 .received_close_notify = false,
875 .allow_truncation_attacks = false,900 .allow_truncation_attacks = options.allow_truncation_attacks,
876 .application_cipher = app_cipher,901 .application_cipher = app_cipher,
877 .partially_read_buffer = undefined,902 .ssl_key_log = options.ssl_key_log,
878 .ssl_key_log = if (options.ssl_key_log_file) |key_log_file| .{
879 .client_key_seq = key_seq,
880 .server_key_seq = key_seq,
881 .client_random = client_hello_rand,
882 .file = key_log_file,
883 } else null,
884 };903 };
885 @memcpy(client.partially_read_buffer[0..leftover.len], leftover);
886 return client;
887 },904 },
888 else => return error.TlsUnexpectedMessage,905 else => return error.TlsUnexpectedMessage,
889 }906 }
...@@ -897,94 +914,73 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client...@@ -897,94 +914,73 @@ pub fn init(stream: anytype, options: Options) InitError(@TypeOf(stream))!Client
897 }914 }
898}915}
899916
900/// Sends TLS-encrypted data to `stream`, which must conform to `StreamInterface`.917fn drain(w: *Writer, data: []const []const u8, splat: usize) Writer.Error!usize {
901/// Returns the number of cleartext bytes sent, which may be fewer than `bytes.len`.918 const c: *Client = @alignCast(@fieldParentPtr("writer", w));
902pub fn write(c: *Client, stream: anytype, bytes: []const u8) !usize {919 const output = c.output;
903 return writeEnd(c, stream, bytes, false);920 const ciphertext_buf = try output.writableSliceGreedy(min_buffer_len);
904}921 var ciphertext_end: usize = 0;
905922 var total_clear: usize = 0;
906/// Sends TLS-encrypted data to `stream`, which must conform to `StreamInterface`.923 done: {
907pub fn writeAll(c: *Client, stream: anytype, bytes: []const u8) !void {924 {
908 var index: usize = 0;925 const buf = w.buffered();
909 while (index < bytes.len) {926 const prepared = prepareCiphertextRecord(c, ciphertext_buf[ciphertext_end..], buf, .application_data);
910 index += try c.write(stream, bytes[index..]);927 total_clear += prepared.cleartext_len;
928 ciphertext_end += prepared.ciphertext_end;
929 if (prepared.cleartext_len < buf.len) break :done;
930 }
931 for (data[0 .. data.len - 1]) |buf| {
932 if (buf.len < min_buffer_len) break :done;
933 const prepared = prepareCiphertextRecord(c, ciphertext_buf[ciphertext_end..], buf, .application_data);
934 total_clear += prepared.cleartext_len;
935 ciphertext_end += prepared.ciphertext_end;
936 if (prepared.cleartext_len < buf.len) break :done;
937 }
938 const buf = data[data.len - 1];
939 for (0..splat) |_| {
940 if (buf.len < min_buffer_len) break :done;
941 const prepared = prepareCiphertextRecord(c, ciphertext_buf[ciphertext_end..], buf, .application_data);
942 total_clear += prepared.cleartext_len;
943 ciphertext_end += prepared.ciphertext_end;
944 if (prepared.cleartext_len < buf.len) break :done;
945 }
911 }946 }
947 output.advance(ciphertext_end);
948 return w.consume(total_clear);
912}949}
913950
914/// Sends TLS-encrypted data to `stream`, which must conform to `StreamInterface`.951fn flush(w: *Writer) Writer.Error!void {
915/// If `end` is true, then this function additionally sends a `close_notify` alert,952 const c: *Client = @alignCast(@fieldParentPtr("writer", w));
916/// which is necessary for the server to distinguish between a properly finished953 const output = c.output;
917/// TLS session, or a truncation attack.954 const ciphertext_buf = try output.writableSliceGreedy(min_buffer_len);
918pub fn writeAllEnd(c: *Client, stream: anytype, bytes: []const u8, end: bool) !void {955 const prepared = prepareCiphertextRecord(c, ciphertext_buf, w.buffered(), .application_data);
919 var index: usize = 0;956 output.advance(prepared.ciphertext_end);
920 while (index < bytes.len) {957 w.end = 0;
921 index += try c.writeEnd(stream, bytes[index..], end);
922 }
923}958}
924959
925/// Sends TLS-encrypted data to `stream`, which must conform to `StreamInterface`.960/// Sends a `close_notify` alert, which is necessary for the server to
926/// Returns the number of cleartext bytes sent, which may be fewer than `bytes.len`.961/// distinguish between a properly finished TLS session, or a truncation
927/// If `end` is true, then this function additionally sends a `close_notify` alert,962/// attack.
928/// which is necessary for the server to distinguish between a properly finished963pub fn end(c: *Client) Writer.Error!void {
929/// TLS session, or a truncation attack.964 try flush(&c.writer);
930pub fn writeEnd(c: *Client, stream: anytype, bytes: []const u8, end: bool) !usize {965 const output = c.output;
931 var ciphertext_buf: [tls.max_ciphertext_record_len * 4]u8 = undefined;966 const ciphertext_buf = try output.writableSliceGreedy(min_buffer_len);
932 var iovecs_buf: [6]std.posix.iovec_const = undefined;967 const prepared = prepareCiphertextRecord(c, ciphertext_buf, &tls.close_notify_alert, .alert);
933 var prepared = prepareCiphertextRecord(c, &iovecs_buf, &ciphertext_buf, bytes, .application_data);968 output.advance(prepared.ciphertext_end);
934 if (end) {
935 prepared.iovec_end += prepareCiphertextRecord(
936 c,
937 iovecs_buf[prepared.iovec_end..],
938 ciphertext_buf[prepared.ciphertext_end..],
939 &tls.close_notify_alert,
940 .alert,
941 ).iovec_end;
942 }
943
944 const iovec_end = prepared.iovec_end;
945 const overhead_len = prepared.overhead_len;
946
947 // Ideally we would call writev exactly once here, however, we must ensure
948 // that we don't return with a record partially written.
949 var i: usize = 0;
950 var total_amt: usize = 0;
951 while (true) {
952 var amt = try stream.writev(iovecs_buf[i..iovec_end]);
953 while (amt >= iovecs_buf[i].len) {
954 const encrypted_amt = iovecs_buf[i].len;
955 total_amt += encrypted_amt - overhead_len;
956 amt -= encrypted_amt;
957 i += 1;
958 // Rely on the property that iovecs delineate records, meaning that
959 // if amt equals zero here, we have fortunately found ourselves
960 // with a short read that aligns at the record boundary.
961 if (i >= iovec_end) return total_amt;
962 // We also cannot return on a vector boundary if the final close_notify is
963 // not sent; otherwise the caller would not know to retry the call.
964 if (amt == 0 and (!end or i < iovec_end - 1)) return total_amt;
965 }
966 iovecs_buf[i].base += amt;
967 iovecs_buf[i].len -= amt;
968 }
969}969}
970970
971fn prepareCiphertextRecord(971fn prepareCiphertextRecord(
972 c: *Client,972 c: *Client,
973 iovecs: []std.posix.iovec_const,
974 ciphertext_buf: []u8,973 ciphertext_buf: []u8,
975 bytes: []const u8,974 bytes: []const u8,
976 inner_content_type: tls.ContentType,975 inner_content_type: tls.ContentType,
977) struct {976) struct {
978 iovec_end: usize,
979 ciphertext_end: usize,977 ciphertext_end: usize,
980 /// How many bytes are taken up by overhead per record.978 cleartext_len: usize,
981 overhead_len: usize,
982} {979} {
983 // Due to the trailing inner content type byte in the ciphertext, we need980 // Due to the trailing inner content type byte in the ciphertext, we need
984 // an additional buffer for storing the cleartext into before encrypting.981 // an additional buffer for storing the cleartext into before encrypting.
985 var cleartext_buf: [max_ciphertext_len]u8 = undefined;982 var cleartext_buf: [max_ciphertext_len]u8 = undefined;
986 var ciphertext_end: usize = 0;983 var ciphertext_end: usize = 0;
987 var iovec_end: usize = 0;
988 var bytes_i: usize = 0;984 var bytes_i: usize = 0;
989 switch (c.application_cipher) {985 switch (c.application_cipher) {
990 inline else => |*p| switch (c.tls_version) {986 inline else => |*p| switch (c.tls_version) {
...@@ -992,18 +988,15 @@ fn prepareCiphertextRecord(...@@ -992,18 +988,15 @@ fn prepareCiphertextRecord(
992 const pv = &p.tls_1_3;988 const pv = &p.tls_1_3;
993 const P = @TypeOf(p.*);989 const P = @TypeOf(p.*);
994 const overhead_len = tls.record_header_len + P.AEAD.tag_length + 1;990 const overhead_len = tls.record_header_len + P.AEAD.tag_length + 1;
995 const close_notify_alert_reserved = tls.close_notify_alert.len + overhead_len;
996 while (true) {991 while (true) {
997 const encrypted_content_len: u16 = @min(992 const encrypted_content_len: u16 = @min(
998 bytes.len - bytes_i,993 bytes.len - bytes_i,
999 tls.max_ciphertext_inner_record_len,994 tls.max_ciphertext_inner_record_len,
1000 ciphertext_buf.len -|995 ciphertext_buf.len -| (overhead_len + ciphertext_end),
1001 (close_notify_alert_reserved + overhead_len + ciphertext_end),
1002 );996 );
1003 if (encrypted_content_len == 0) return .{997 if (encrypted_content_len == 0) return .{
1004 .iovec_end = iovec_end,
1005 .ciphertext_end = ciphertext_end,998 .ciphertext_end = ciphertext_end,
1006 .overhead_len = overhead_len,999 .cleartext_len = bytes_i,
1007 };1000 };
10081001
1009 @memcpy(cleartext_buf[0..encrypted_content_len], bytes[bytes_i..][0..encrypted_content_len]);1002 @memcpy(cleartext_buf[0..encrypted_content_len], bytes[bytes_i..][0..encrypted_content_len]);
...@@ -1012,7 +1005,6 @@ fn prepareCiphertextRecord(...@@ -1012,7 +1005,6 @@ fn prepareCiphertextRecord(
1012 const ciphertext_len = encrypted_content_len + 1;1005 const ciphertext_len = encrypted_content_len + 1;
1013 const cleartext = cleartext_buf[0..ciphertext_len];1006 const cleartext = cleartext_buf[0..ciphertext_len];
10141007
1015 const record_start = ciphertext_end;
1016 const ad = ciphertext_buf[ciphertext_end..][0..tls.record_header_len];1008 const ad = ciphertext_buf[ciphertext_end..][0..tls.record_header_len];
1017 ad.* = .{@intFromEnum(tls.ContentType.application_data)} ++1009 ad.* = .{@intFromEnum(tls.ContentType.application_data)} ++
1018 int(u16, @intFromEnum(tls.ProtocolVersion.tls_1_2)) ++1010 int(u16, @intFromEnum(tls.ProtocolVersion.tls_1_2)) ++
...@@ -1030,38 +1022,27 @@ fn prepareCiphertextRecord(...@@ -1030,38 +1022,27 @@ fn prepareCiphertextRecord(
1030 };1022 };
1031 P.AEAD.encrypt(ciphertext, auth_tag, cleartext, ad, nonce, pv.client_key);1023 P.AEAD.encrypt(ciphertext, auth_tag, cleartext, ad, nonce, pv.client_key);
1032 c.write_seq += 1; // TODO send key_update on overflow1024 c.write_seq += 1; // TODO send key_update on overflow
1033
1034 const record = ciphertext_buf[record_start..ciphertext_end];
1035 iovecs[iovec_end] = .{
1036 .base = record.ptr,
1037 .len = record.len,
1038 };
1039 iovec_end += 1;
1040 }1025 }
1041 },1026 },
1042 .tls_1_2 => {1027 .tls_1_2 => {
1043 const pv = &p.tls_1_2;1028 const pv = &p.tls_1_2;
1044 const P = @TypeOf(p.*);1029 const P = @TypeOf(p.*);
1045 const overhead_len = tls.record_header_len + P.record_iv_length + P.mac_length;1030 const overhead_len = tls.record_header_len + P.record_iv_length + P.mac_length;
1046 const close_notify_alert_reserved = tls.close_notify_alert.len + overhead_len;
1047 while (true) {1031 while (true) {
1048 const message_len: u16 = @min(1032 const message_len: u16 = @min(
1049 bytes.len - bytes_i,1033 bytes.len - bytes_i,
1050 tls.max_ciphertext_inner_record_len,1034 tls.max_ciphertext_inner_record_len,
1051 ciphertext_buf.len -|1035 ciphertext_buf.len -| (overhead_len + ciphertext_end),
1052 (close_notify_alert_reserved + overhead_len + ciphertext_end),
1053 );1036 );
1054 if (message_len == 0) return .{1037 if (message_len == 0) return .{
1055 .iovec_end = iovec_end,
1056 .ciphertext_end = ciphertext_end,1038 .ciphertext_end = ciphertext_end,
1057 .overhead_len = overhead_len,1039 .cleartext_len = bytes_i,
1058 };1040 };
10591041
1060 @memcpy(cleartext_buf[0..message_len], bytes[bytes_i..][0..message_len]);1042 @memcpy(cleartext_buf[0..message_len], bytes[bytes_i..][0..message_len]);
1061 bytes_i += message_len;1043 bytes_i += message_len;
1062 const cleartext = cleartext_buf[0..message_len];1044 const cleartext = cleartext_buf[0..message_len];
10631045
1064 const record_start = ciphertext_end;
1065 const record_header = ciphertext_buf[ciphertext_end..][0..tls.record_header_len];1046 const record_header = ciphertext_buf[ciphertext_end..][0..tls.record_header_len];
1066 ciphertext_end += tls.record_header_len;1047 ciphertext_end += tls.record_header_len;
1067 record_header.* = .{@intFromEnum(inner_content_type)} ++1048 record_header.* = .{@intFromEnum(inner_content_type)} ++
...@@ -1083,13 +1064,6 @@ fn prepareCiphertextRecord(...@@ -1083,13 +1064,6 @@ fn prepareCiphertextRecord(
1083 ciphertext_end += P.mac_length;1064 ciphertext_end += P.mac_length;
1084 P.AEAD.encrypt(ciphertext, auth_tag, cleartext, ad, nonce, pv.client_write_key);1065 P.AEAD.encrypt(ciphertext, auth_tag, cleartext, ad, nonce, pv.client_write_key);
1085 c.write_seq += 1; // TODO send key_update on overflow1066 c.write_seq += 1; // TODO send key_update on overflow
1086
1087 const record = ciphertext_buf[record_start..ciphertext_end];
1088 iovecs[iovec_end] = .{
1089 .base = record.ptr,
1090 .len = record.len,
1091 };
1092 iovec_end += 1;
1093 }1067 }
1094 },1068 },
1095 else => unreachable,1069 else => unreachable,
...@@ -1098,421 +1072,194 @@ fn prepareCiphertextRecord(...@@ -1098,421 +1072,194 @@ fn prepareCiphertextRecord(
1098}1072}
10991073
1100pub fn eof(c: Client) bool {1074pub fn eof(c: Client) bool {
1101 return c.received_close_notify and1075 return c.received_close_notify;
1102 c.partial_cleartext_idx >= c.partial_ciphertext_idx and
1103 c.partial_ciphertext_idx >= c.partial_ciphertext_end;
1104}
1105
1106/// Receives TLS-encrypted data from `stream`, which must conform to `StreamInterface`.
1107/// Returns the number of bytes read, calling the underlying read function the
1108/// minimal number of times until the buffer has at least `len` bytes filled.
1109/// If the number read is less than `len` it means the stream reached the end.
1110/// Reaching the end of the stream is not an error condition.
1111pub fn readAtLeast(c: *Client, stream: anytype, buffer: []u8, len: usize) !usize {
1112 var iovecs = [1]std.posix.iovec{.{ .base = buffer.ptr, .len = buffer.len }};
1113 return readvAtLeast(c, stream, &iovecs, len);
1114}
1115
1116/// Receives TLS-encrypted data from `stream`, which must conform to `StreamInterface`.
1117pub fn read(c: *Client, stream: anytype, buffer: []u8) !usize {
1118 return readAtLeast(c, stream, buffer, 1);
1119}1076}
11201077
1121/// Receives TLS-encrypted data from `stream`, which must conform to `StreamInterface`.1078fn stream(r: *Reader, w: *Writer, limit: std.Io.Limit) Reader.StreamError!usize {
1122/// Returns the number of bytes read. If the number read is smaller than1079 const c: *Client = @alignCast(@fieldParentPtr("reader", r));
1123/// `buffer.len`, it means the stream reached the end. Reaching the end of the1080 if (c.eof()) return error.EndOfStream;
1124/// stream is not an error condition.1081 const input = c.input;
1125pub fn readAll(c: *Client, stream: anytype, buffer: []u8) !usize {1082 // If at least one full encrypted record is not buffered, read once.
1126 return readAtLeast(c, stream, buffer, buffer.len);1083 const record_header = input.peek(tls.record_header_len) catch |err| switch (err) {
1127}1084 error.EndOfStream => {
11281085 // This is either a truncation attack, a bug in the server, or an
1129/// Receives TLS-encrypted data from `stream`, which must conform to `StreamInterface`.1086 // intentional omission of the close_notify message due to truncation
1130/// Returns the number of bytes read. If the number read is less than the space1087 // detection handled above the TLS layer.
1131/// provided it means the stream reached the end. Reaching the end of the1088 if (c.allow_truncation_attacks) {
1132/// stream is not an error condition.1089 c.received_close_notify = true;
1133/// The `iovecs` parameter is mutable because this function needs to mutate the fields in1090 return error.EndOfStream;
1134/// order to handle partial reads from the underlying stream layer.1091 } else {
1135pub fn readv(c: *Client, stream: anytype, iovecs: []std.posix.iovec) !usize {1092 return failRead(c, error.TlsConnectionTruncated);
1136 return readvAtLeast(c, stream, iovecs, 1);1093 }
1137}1094 },
11381095 error.ReadFailed => return error.ReadFailed,
1139/// Receives TLS-encrypted data from `stream`, which must conform to `StreamInterface`.1096 };
1140/// Returns the number of bytes read, calling the underlying read function the1097 const ct: tls.ContentType = @enumFromInt(record_header[0]);
1141/// minimal number of times until the iovecs have at least `len` bytes filled.1098 const legacy_version = mem.readInt(u16, record_header[1..][0..2], .big);
1142/// If the number read is less than `len` it means the stream reached the end.1099 _ = legacy_version;
1143/// Reaching the end of the stream is not an error condition.1100 const record_len = mem.readInt(u16, record_header[3..][0..2], .big);
1144/// The `iovecs` parameter is mutable because this function needs to mutate the fields in1101 if (record_len > max_ciphertext_len) return failRead(c, error.TlsRecordOverflow);
1145/// order to handle partial reads from the underlying stream layer.1102 const record_end = 5 + record_len;
1146pub fn readvAtLeast(c: *Client, stream: anytype, iovecs: []std.posix.iovec, len: usize) !usize {1103 if (record_end > input.buffered().len) {
1147 if (c.eof()) return 0;1104 input.fillMore() catch |err| switch (err) {
11481105 error.EndOfStream => return failRead(c, error.TlsConnectionTruncated),
1149 var off_i: usize = 0;1106 error.ReadFailed => return error.ReadFailed,
1150 var vec_i: usize = 0;1107 };
1151 while (true) {1108 if (record_end > input.buffered().len) return 0;
1152 var amt = try c.readvAdvanced(stream, iovecs[vec_i..]);
1153 off_i += amt;
1154 if (c.eof() or off_i >= len) return off_i;
1155 while (amt >= iovecs[vec_i].len) {
1156 amt -= iovecs[vec_i].len;
1157 vec_i += 1;
1158 }
1159 iovecs[vec_i].base += amt;
1160 iovecs[vec_i].len -= amt;
1161 }
1162}
1163
1164/// Receives TLS-encrypted data from `stream`, which must conform to `StreamInterface`.
1165/// Returns number of bytes that have been read, populated inside `iovecs`. A
1166/// return value of zero bytes does not mean end of stream. Instead, check the `eof()`
1167/// for the end of stream. The `eof()` may be true after any call to
1168/// `read`, including when greater than zero bytes are returned, and this
1169/// function asserts that `eof()` is `false`.
1170/// See `readv` for a higher level function that has the same, familiar API as
1171/// other read functions, such as `std.fs.File.read`.
1172pub fn readvAdvanced(c: *Client, stream: anytype, iovecs: []const std.posix.iovec) !usize {
1173 var vp: VecPut = .{ .iovecs = iovecs };
1174
1175 // Give away the buffered cleartext we have, if any.
1176 const partial_cleartext = c.partially_read_buffer[c.partial_cleartext_idx..c.partial_ciphertext_idx];
1177 if (partial_cleartext.len > 0) {
1178 const amt: u15 = @intCast(vp.put(partial_cleartext));
1179 c.partial_cleartext_idx += amt;
1180
1181 if (c.partial_cleartext_idx == c.partial_ciphertext_idx and
1182 c.partial_ciphertext_end == c.partial_ciphertext_idx)
1183 {
1184 // The buffer is now empty.
1185 c.partial_cleartext_idx = 0;
1186 c.partial_ciphertext_idx = 0;
1187 c.partial_ciphertext_end = 0;
1188 }
1189
1190 if (c.received_close_notify) {
1191 c.partial_ciphertext_end = 0;
1192 assert(vp.total == amt);
1193 return amt;
1194 } else if (amt > 0) {
1195 // We don't need more data, so don't call read.
1196 assert(vp.total == amt);
1197 return amt;
1198 }
1199 }1109 }
12001110
1201 assert(!c.received_close_notify);
1202
1203 // Ideally, this buffer would never be used. It is needed when `iovecs` are
1204 // too small to fit the cleartext, which may be as large as `max_ciphertext_len`.
1205 var cleartext_stack_buffer: [max_ciphertext_len]u8 = undefined;1111 var cleartext_stack_buffer: [max_ciphertext_len]u8 = undefined;
1206 // Temporarily stores ciphertext before decrypting it and giving it to `iovecs`.1112 const cleartext, const inner_ct: tls.ContentType = cleartext: switch (c.application_cipher) {
1207 var in_stack_buffer: [max_ciphertext_len * 4]u8 = undefined;1113 inline else => |*p| switch (c.tls_version) {
1208 // How many bytes left in the user's buffer.1114 .tls_1_3 => {
1209 const free_size = vp.freeSize();1115 const pv = &p.tls_1_3;
1210 // The amount of the user's buffer that we need to repurpose for storing1116 const P = @TypeOf(p.*);
1211 // ciphertext. The end of the buffer will be used for such purposes.1117 const ad = input.take(tls.record_header_len) catch unreachable; // already peeked
1212 const ciphertext_buf_len = (free_size / 2) -| in_stack_buffer.len;1118 const ciphertext_len = record_len - P.AEAD.tag_length;
1213 // The amount of the user's buffer that will be used to give cleartext. The1119 const ciphertext = input.take(ciphertext_len) catch unreachable; // already peeked
1214 // beginning of the buffer will be used for such purposes.1120 const auth_tag = (input.takeArray(P.AEAD.tag_length) catch unreachable).*; // already peeked
1215 const cleartext_buf_len = free_size - ciphertext_buf_len;1121 const nonce = nonce: {
12161122 const V = @Vector(P.AEAD.nonce_length, u8);
1217 // Recoup `partially_read_buffer` space. This is necessary because it is assumed1123 const pad = [1]u8{0} ** (P.AEAD.nonce_length - 8);
1218 // below that `frag0` is big enough to hold at least one record.1124 const operand: V = pad ++ std.mem.toBytes(big(c.read_seq));
1219 limitedOverlapCopy(c.partially_read_buffer[0..c.partial_ciphertext_end], c.partial_ciphertext_idx);1125 break :nonce @as(V, pv.server_iv) ^ operand;
1220 c.partial_ciphertext_end -= c.partial_ciphertext_idx;1126 };
1221 c.partial_ciphertext_idx = 0;1127 const cleartext = cleartext_stack_buffer[0..ciphertext.len];
1222 c.partial_cleartext_idx = 0;1128 P.AEAD.decrypt(cleartext, ciphertext, auth_tag, ad, nonce, pv.server_key) catch
1223 const first_iov = c.partially_read_buffer[c.partial_ciphertext_end..];1129 return failRead(c, error.TlsBadRecordMac);
12241130 const msg = mem.trimRight(u8, cleartext, "\x00");
1225 var ask_iovecs_buf: [2]std.posix.iovec = .{1131 break :cleartext .{ msg[0 .. msg.len - 1], @enumFromInt(msg[msg.len - 1]) };
1226 .{1132 },
1227 .base = first_iov.ptr,1133 .tls_1_2 => {
1228 .len = first_iov.len,1134 const pv = &p.tls_1_2;
1229 },1135 const P = @TypeOf(p.*);
1230 .{1136 const message_len: u16 = record_len - P.record_iv_length - P.mac_length;
1231 .base = &in_stack_buffer,1137 const ad_header = input.take(tls.record_header_len) catch unreachable; // already peeked
1232 .len = in_stack_buffer.len,1138 const ad = std.mem.toBytes(big(c.read_seq)) ++
1139 ad_header[0 .. 1 + 2] ++
1140 std.mem.toBytes(big(message_len));
1141 const record_iv = (input.takeArray(P.record_iv_length) catch unreachable).*; // already peeked
1142 const masked_read_seq = c.read_seq &
1143 comptime std.math.shl(u64, std.math.maxInt(u64), 8 * P.record_iv_length);
1144 const nonce: [P.AEAD.nonce_length]u8 = nonce: {
1145 const V = @Vector(P.AEAD.nonce_length, u8);
1146 const pad = [1]u8{0} ** (P.AEAD.nonce_length - 8);
1147 const operand: V = pad ++ @as([8]u8, @bitCast(big(masked_read_seq)));
1148 break :nonce @as(V, pv.server_write_IV ++ record_iv) ^ operand;
1149 };
1150 const ciphertext = input.take(message_len) catch unreachable; // already peeked
1151 const auth_tag = (input.takeArray(P.mac_length) catch unreachable).*; // already peeked
1152 const cleartext = cleartext_stack_buffer[0..ciphertext.len];
1153 P.AEAD.decrypt(cleartext, ciphertext, auth_tag, ad, nonce, pv.server_write_key) catch
1154 return failRead(c, error.TlsBadRecordMac);
1155 break :cleartext .{ cleartext, ct };
1156 },
1157 else => unreachable,
1233 },1158 },
1234 };1159 };
12351160 c.read_seq = std.math.add(u64, c.read_seq, 1) catch return failRead(c, error.TlsSequenceOverflow);
1236 // Cleartext capacity of output buffer, in records. Minimum one full record.1161 switch (inner_ct) {
1237 const buf_cap = @max(cleartext_buf_len / max_ciphertext_len, 1);1162 .alert => {
1238 const wanted_read_len = buf_cap * (max_ciphertext_len + tls.record_header_len);1163 if (cleartext.len != 2) return failRead(c, error.TlsDecodeError);
1239 const ask_len = @max(wanted_read_len, cleartext_stack_buffer.len) - c.partial_ciphertext_end;1164 const alert: tls.Alert = .{
1240 const ask_iovecs = limitVecs(&ask_iovecs_buf, ask_len);1165 .level = @enumFromInt(cleartext[0]),
1241 const actual_read_len = try stream.readv(ask_iovecs);1166 .description = @enumFromInt(cleartext[1]),
1242 if (actual_read_len == 0) {1167 };
1243 // This is either a truncation attack, a bug in the server, or an1168 switch (alert.description) {
1244 // intentional omission of the close_notify message due to truncation1169 .close_notify => {
1245 // detection handled above the TLS layer.1170 c.received_close_notify = true;
1246 if (c.allow_truncation_attacks) {1171 return 0;
1247 c.received_close_notify = true;
1248 } else {
1249 return error.TlsConnectionTruncated;
1250 }
1251 }
1252
1253 // There might be more bytes inside `in_stack_buffer` that need to be processed,
1254 // but at least frag0 will have one complete ciphertext record.
1255 const frag0_end = @min(c.partially_read_buffer.len, c.partial_ciphertext_end + actual_read_len);
1256 const frag0 = c.partially_read_buffer[c.partial_ciphertext_idx..frag0_end];
1257 var frag1 = in_stack_buffer[0..actual_read_len -| first_iov.len];
1258 // We need to decipher frag0 and frag1 but there may be a ciphertext record
1259 // straddling the boundary. We can handle this with two memcpy() calls to
1260 // assemble the straddling record in between handling the two sides.
1261 var frag = frag0;
1262 var in: usize = 0;
1263 while (true) {
1264 if (in == frag.len) {
1265 // Perfect split.
1266 if (frag.ptr == frag1.ptr) {
1267 c.partial_ciphertext_end = c.partial_ciphertext_idx;
1268 return vp.total;
1269 }
1270 frag = frag1;
1271 in = 0;
1272 continue;
1273 }
1274
1275 if (in + tls.record_header_len > frag.len) {
1276 if (frag.ptr == frag1.ptr)
1277 return finishRead(c, frag, in, vp.total);
1278
1279 const first = frag[in..];
1280
1281 if (frag1.len < tls.record_header_len)
1282 return finishRead2(c, first, frag1, vp.total);
1283
1284 // A record straddles the two fragments. Copy into the now-empty first fragment.
1285 const record_len_byte_0: u16 = straddleByte(frag, frag1, in + 3);
1286 const record_len_byte_1: u16 = straddleByte(frag, frag1, in + 4);
1287 const record_len = (record_len_byte_0 << 8) | record_len_byte_1;
1288 if (record_len > max_ciphertext_len) return error.TlsRecordOverflow;
1289
1290 const full_record_len = record_len + tls.record_header_len;
1291 const second_len = full_record_len - first.len;
1292 if (frag1.len < second_len)
1293 return finishRead2(c, first, frag1, vp.total);
1294
1295 limitedOverlapCopy(frag, in);
1296 @memcpy(frag[first.len..][0..second_len], frag1[0..second_len]);
1297 frag = frag[0..full_record_len];
1298 frag1 = frag1[second_len..];
1299 in = 0;
1300 continue;
1301 }
1302 const ct: tls.ContentType = @enumFromInt(frag[in]);
1303 in += 1;
1304 const legacy_version = mem.readInt(u16, frag[in..][0..2], .big);
1305 in += 2;
1306 _ = legacy_version;
1307 const record_len = mem.readInt(u16, frag[in..][0..2], .big);
1308 if (record_len > max_ciphertext_len) return error.TlsRecordOverflow;
1309 in += 2;
1310 const end = in + record_len;
1311 if (end > frag.len) {
1312 // We need the record header on the next iteration of the loop.
1313 in -= tls.record_header_len;
1314
1315 if (frag.ptr == frag1.ptr)
1316 return finishRead(c, frag, in, vp.total);
1317
1318 // A record straddles the two fragments. Copy into the now-empty first fragment.
1319 const first = frag[in..];
1320 const full_record_len = record_len + tls.record_header_len;
1321 const second_len = full_record_len - first.len;
1322 if (frag1.len < second_len)
1323 return finishRead2(c, first, frag1, vp.total);
1324
1325 limitedOverlapCopy(frag, in);
1326 @memcpy(frag[first.len..][0..second_len], frag1[0..second_len]);
1327 frag = frag[0..full_record_len];
1328 frag1 = frag1[second_len..];
1329 in = 0;
1330 continue;
1331 }
1332 const cleartext, const inner_ct: tls.ContentType = cleartext: switch (c.application_cipher) {
1333 inline else => |*p| switch (c.tls_version) {
1334 .tls_1_3 => {
1335 const pv = &p.tls_1_3;
1336 const P = @TypeOf(p.*);
1337 const ad = frag[in - tls.record_header_len ..][0..tls.record_header_len];
1338 const ciphertext_len = record_len - P.AEAD.tag_length;
1339 const ciphertext = frag[in..][0..ciphertext_len];
1340 in += ciphertext_len;
1341 const auth_tag = frag[in..][0..P.AEAD.tag_length].*;
1342 const nonce = nonce: {
1343 const V = @Vector(P.AEAD.nonce_length, u8);
1344 const pad = [1]u8{0} ** (P.AEAD.nonce_length - 8);
1345 const operand: V = pad ++ std.mem.toBytes(big(c.read_seq));
1346 break :nonce @as(V, pv.server_iv) ^ operand;
1347 };
1348 const out_buf = vp.peek();
1349 const cleartext_buf = if (ciphertext.len <= out_buf.len)
1350 out_buf
1351 else
1352 &cleartext_stack_buffer;
1353 const cleartext = cleartext_buf[0..ciphertext.len];
1354 P.AEAD.decrypt(cleartext, ciphertext, auth_tag, ad, nonce, pv.server_key) catch
1355 return error.TlsBadRecordMac;
1356 const msg = mem.trimEnd(u8, cleartext, "\x00");
1357 break :cleartext .{ msg[0 .. msg.len - 1], @enumFromInt(msg[msg.len - 1]) };
1358 },1172 },
1359 .tls_1_2 => {1173 .user_canceled => {
1360 const pv = &p.tls_1_2;1174 // TODO: handle server-side closures
1361 const P = @TypeOf(p.*);1175 return failRead(c, error.TlsUnexpectedMessage);
1362 const message_len: u16 = record_len - P.record_iv_length - P.mac_length;
1363 const ad = std.mem.toBytes(big(c.read_seq)) ++
1364 frag[in - tls.record_header_len ..][0 .. 1 + 2] ++
1365 std.mem.toBytes(big(message_len));
1366 const record_iv = frag[in..][0..P.record_iv_length].*;
1367 in += P.record_iv_length;
1368 const masked_read_seq = c.read_seq &
1369 comptime std.math.shl(u64, std.math.maxInt(u64), 8 * P.record_iv_length);
1370 const nonce: [P.AEAD.nonce_length]u8 = nonce: {
1371 const V = @Vector(P.AEAD.nonce_length, u8);
1372 const pad = [1]u8{0} ** (P.AEAD.nonce_length - 8);
1373 const operand: V = pad ++ @as([8]u8, @bitCast(big(masked_read_seq)));
1374 break :nonce @as(V, pv.server_write_IV ++ record_iv) ^ operand;
1375 };
1376 const ciphertext = frag[in..][0..message_len];
1377 in += message_len;
1378 const auth_tag = frag[in..][0..P.mac_length].*;
1379 in += P.mac_length;
1380 const out_buf = vp.peek();
1381 const cleartext_buf = if (message_len <= out_buf.len)
1382 out_buf
1383 else
1384 &cleartext_stack_buffer;
1385 const cleartext = cleartext_buf[0..ciphertext.len];
1386 P.AEAD.decrypt(cleartext, ciphertext, auth_tag, ad, nonce, pv.server_write_key) catch
1387 return error.TlsBadRecordMac;
1388 break :cleartext .{ cleartext, ct };
1389 },1176 },
1390 else => unreachable,1177 else => {
1391 },1178 c.alert = alert;
1392 };1179 return failRead(c, error.TlsAlert);
1393 c.read_seq = try std.math.add(u64, c.read_seq, 1);1180 },
1394 switch (inner_ct) {1181 }
1395 .alert => {1182 },
1396 if (cleartext.len != 2) return error.TlsDecodeError;1183 .handshake => {
1397 const level: tls.AlertLevel = @enumFromInt(cleartext[0]);1184 var ct_i: usize = 0;
1398 const desc: tls.AlertDescription = @enumFromInt(cleartext[1]);1185 while (true) {
1399 if (desc == .close_notify) {1186 const handshake_type: tls.HandshakeType = @enumFromInt(cleartext[ct_i]);
1400 c.received_close_notify = true;1187 ct_i += 1;
1401 c.partial_ciphertext_end = c.partial_ciphertext_idx;1188 const handshake_len = mem.readInt(u24, cleartext[ct_i..][0..3], .big);
1402 return vp.total;1189 ct_i += 3;
1403 }1190 const next_handshake_i = ct_i + handshake_len;
1404 _ = level;1191 if (next_handshake_i > cleartext.len) return failRead(c, error.TlsBadLength);
14051192 const handshake = cleartext[ct_i..next_handshake_i];
1406 try desc.toError();1193 switch (handshake_type) {
1407 // TODO: handle server-side closures1194 .new_session_ticket => {
1408 return error.TlsUnexpectedMessage;1195 // This client implementation ignores new session tickets.
1409 },1196 },
1410 .handshake => {1197 .key_update => {
1411 var ct_i: usize = 0;1198 switch (c.application_cipher) {
1412 while (true) {1199 inline else => |*p| {
1413 const handshake_type: tls.HandshakeType = @enumFromInt(cleartext[ct_i]);1200 const pv = &p.tls_1_3;
1414 ct_i += 1;1201 const P = @TypeOf(p.*);
1415 const handshake_len = mem.readInt(u24, cleartext[ct_i..][0..3], .big);1202 const server_secret = hkdfExpandLabel(P.Hkdf, pv.server_secret, "traffic upd", "", P.Hash.digest_length);
1416 ct_i += 3;1203 if (c.ssl_key_log) |key_log| logSecrets(key_log.writer, .{
1417 const next_handshake_i = ct_i + handshake_len;1204 .counter = key_log.serverCounter(),
1418 if (next_handshake_i > cleartext.len)1205 .client_random = &key_log.client_random,
1419 return error.TlsBadLength;1206 }, .{
1420 const handshake = cleartext[ct_i..next_handshake_i];1207 .SERVER_TRAFFIC_SECRET = &server_secret,
1421 switch (handshake_type) {1208 });
1422 .new_session_ticket => {1209 pv.server_secret = server_secret;
1423 // This client implementation ignores new session tickets.1210 pv.server_key = hkdfExpandLabel(P.Hkdf, server_secret, "key", "", P.AEAD.key_length);
1424 },1211 pv.server_iv = hkdfExpandLabel(P.Hkdf, server_secret, "iv", "", P.AEAD.nonce_length);
1425 .key_update => {1212 },
1426 switch (c.application_cipher) {1213 }
1427 inline else => |*p| {1214 c.read_seq = 0;
1428 const pv = &p.tls_1_3;1215
1429 const P = @TypeOf(p.*);1216 switch (@as(tls.KeyUpdateRequest, @enumFromInt(handshake[0]))) {
1430 const server_secret = hkdfExpandLabel(P.Hkdf, pv.server_secret, "traffic upd", "", P.Hash.digest_length);1217 .update_requested => {
1431 if (c.ssl_key_log) |*key_log| logSecrets(key_log.file, .{1218 switch (c.application_cipher) {
1432 .counter = key_log.serverCounter(),1219 inline else => |*p| {
1433 .client_random = &key_log.client_random,1220 const pv = &p.tls_1_3;
1434 }, .{1221 const P = @TypeOf(p.*);
1435 .SERVER_TRAFFIC_SECRET = &server_secret,1222 const client_secret = hkdfExpandLabel(P.Hkdf, pv.client_secret, "traffic upd", "", P.Hash.digest_length);
1436 });1223 if (c.ssl_key_log) |key_log| logSecrets(key_log.writer, .{
1437 pv.server_secret = server_secret;1224 .counter = key_log.clientCounter(),
1438 pv.server_key = hkdfExpandLabel(P.Hkdf, server_secret, "key", "", P.AEAD.key_length);1225 .client_random = &key_log.client_random,
1439 pv.server_iv = hkdfExpandLabel(P.Hkdf, server_secret, "iv", "", P.AEAD.nonce_length);1226 }, .{
1440 },1227 .CLIENT_TRAFFIC_SECRET = &client_secret,
1441 }1228 });
1442 c.read_seq = 0;1229 pv.client_secret = client_secret;
14431230 pv.client_key = hkdfExpandLabel(P.Hkdf, client_secret, "key", "", P.AEAD.key_length);
1444 switch (@as(tls.KeyUpdateRequest, @enumFromInt(handshake[0]))) {1231 pv.client_iv = hkdfExpandLabel(P.Hkdf, client_secret, "iv", "", P.AEAD.nonce_length);
1445 .update_requested => {1232 },
1446 switch (c.application_cipher) {1233 }
1447 inline else => |*p| {1234 c.write_seq = 0;
1448 const pv = &p.tls_1_3;1235 },
1449 const P = @TypeOf(p.*);1236 .update_not_requested => {},
1450 const client_secret = hkdfExpandLabel(P.Hkdf, pv.client_secret, "traffic upd", "", P.Hash.digest_length);1237 _ => return failRead(c, error.TlsIllegalParameter),
1451 if (c.ssl_key_log) |*key_log| logSecrets(key_log.file, .{
1452 .counter = key_log.clientCounter(),
1453 .client_random = &key_log.client_random,
1454 }, .{
1455 .CLIENT_TRAFFIC_SECRET = &client_secret,
1456 });
1457 pv.client_secret = client_secret;
1458 pv.client_key = hkdfExpandLabel(P.Hkdf, client_secret, "key", "", P.AEAD.key_length);
1459 pv.client_iv = hkdfExpandLabel(P.Hkdf, client_secret, "iv", "", P.AEAD.nonce_length);
1460 },
1461 }
1462 c.write_seq = 0;
1463 },
1464 .update_not_requested => {},
1465 _ => return error.TlsIllegalParameter,
1466 }
1467 },
1468 else => {
1469 return error.TlsUnexpectedMessage;
1470 },
1471 }
1472 ct_i = next_handshake_i;
1473 if (ct_i >= cleartext.len) break;
1474 }
1475 },
1476 .application_data => {
1477 // Determine whether the output buffer or a stack
1478 // buffer was used for storing the cleartext.
1479 if (cleartext.ptr == &cleartext_stack_buffer) {
1480 // Stack buffer was used, so we must copy to the output buffer.
1481 if (c.partial_ciphertext_idx > c.partial_cleartext_idx) {
1482 // We have already run out of room in iovecs. Continue
1483 // appending to `partially_read_buffer`.
1484 @memcpy(
1485 c.partially_read_buffer[c.partial_ciphertext_idx..][0..cleartext.len],
1486 cleartext,
1487 );
1488 c.partial_ciphertext_idx = @intCast(c.partial_ciphertext_idx + cleartext.len);
1489 } else {
1490 const amt = vp.put(cleartext);
1491 if (amt < cleartext.len) {
1492 const rest = cleartext[amt..];
1493 c.partial_cleartext_idx = 0;
1494 c.partial_ciphertext_idx = @intCast(rest.len);
1495 @memcpy(c.partially_read_buffer[0..rest.len], rest);
1496 }1238 }
1497 }1239 },
1498 } else {1240 else => return failRead(c, error.TlsUnexpectedMessage),
1499 // Output buffer was used directly which means no
1500 // memory copying needs to occur, and we can move
1501 // on to the next ciphertext record.
1502 vp.next(cleartext.len);
1503 }1241 }
1504 },1242 ct_i = next_handshake_i;
1505 else => return error.TlsUnexpectedMessage,1243 if (ct_i >= cleartext.len) break;
1506 }1244 }
1507 in = end;1245 return 0;
1246 },
1247 .application_data => {
1248 if (@intFromEnum(limit) < cleartext.len) return failRead(c, error.OutputBufferUndersize);
1249 try w.writeAll(cleartext);
1250 return cleartext.len;
1251 },
1252 else => return failRead(c, error.TlsUnexpectedMessage),
1508 }1253 }
1509}1254}
15101255
1511fn logSecrets(key_log_file: std.fs.File, context: anytype, secrets: anytype) void {1256fn failRead(c: *Client, err: ReadError) error{ReadFailed} {
1512 const locked = if (key_log_file.lock(.exclusive)) |_| true else |_| false;1257 c.read_err = err;
1513 defer if (locked) key_log_file.unlock();1258 return error.ReadFailed;
1514 key_log_file.seekFromEnd(0) catch {};1259}
1515 inline for (@typeInfo(@TypeOf(secrets)).@"struct".fields) |field| key_log_file.deprecatedWriter().print("{s}" ++1260
1261fn logSecrets(w: *Writer, context: anytype, secrets: anytype) void {
1262 inline for (@typeInfo(@TypeOf(secrets)).@"struct".fields) |field| w.print("{s}" ++
1516 (if (@hasField(@TypeOf(context), "counter")) "_{d}" else "") ++ " {x} {x}\n", .{field.name} ++1263 (if (@hasField(@TypeOf(context), "counter")) "_{d}" else "") ++ " {x} {x}\n", .{field.name} ++
1517 (if (@hasField(@TypeOf(context), "counter")) .{context.counter} else .{}) ++ .{1264 (if (@hasField(@TypeOf(context), "counter")) .{context.counter} else .{}) ++ .{
1518 context.client_random,1265 context.client_random,
...@@ -1520,62 +1267,6 @@ fn logSecrets(key_log_file: std.fs.File, context: anytype, secrets: anytype) voi...@@ -1520,62 +1267,6 @@ fn logSecrets(key_log_file: std.fs.File, context: anytype, secrets: anytype) voi
1520 }) catch {};1267 }) catch {};
1521}1268}
15221269
1523fn finishRead(c: *Client, frag: []const u8, in: usize, out: usize) usize {
1524 const saved_buf = frag[in..];
1525 if (c.partial_ciphertext_idx > c.partial_cleartext_idx) {
1526 // There is cleartext at the beginning already which we need to preserve.
1527 c.partial_ciphertext_end = @intCast(c.partial_ciphertext_idx + saved_buf.len);
1528 @memcpy(c.partially_read_buffer[c.partial_ciphertext_idx..][0..saved_buf.len], saved_buf);
1529 } else {
1530 c.partial_cleartext_idx = 0;
1531 c.partial_ciphertext_idx = 0;
1532 c.partial_ciphertext_end = @intCast(saved_buf.len);
1533 @memcpy(c.partially_read_buffer[0..saved_buf.len], saved_buf);
1534 }
1535 return out;
1536}
1537
1538/// Note that `first` usually overlaps with `c.partially_read_buffer`.
1539fn finishRead2(c: *Client, first: []const u8, frag1: []const u8, out: usize) usize {
1540 if (c.partial_ciphertext_idx > c.partial_cleartext_idx) {
1541 // There is cleartext at the beginning already which we need to preserve.
1542 c.partial_ciphertext_end = @intCast(c.partial_ciphertext_idx + first.len + frag1.len);
1543 // TODO: eliminate this call to copyForwards
1544 std.mem.copyForwards(u8, c.partially_read_buffer[c.partial_ciphertext_idx..][0..first.len], first);
1545 @memcpy(c.partially_read_buffer[c.partial_ciphertext_idx + first.len ..][0..frag1.len], frag1);
1546 } else {
1547 c.partial_cleartext_idx = 0;
1548 c.partial_ciphertext_idx = 0;
1549 c.partial_ciphertext_end = @intCast(first.len + frag1.len);
1550 // TODO: eliminate this call to copyForwards
1551 std.mem.copyForwards(u8, c.partially_read_buffer[0..first.len], first);
1552 @memcpy(c.partially_read_buffer[first.len..][0..frag1.len], frag1);
1553 }
1554 return out;
1555}
1556
1557fn limitedOverlapCopy(frag: []u8, in: usize) void {
1558 const first = frag[in..];
1559 if (first.len <= in) {
1560 // A single, non-overlapping memcpy suffices.
1561 @memcpy(frag[0..first.len], first);
1562 } else {
1563 // One memcpy call would overlap, so just do this instead.
1564 std.mem.copyForwards(u8, frag, first);
1565 }
1566}
1567
1568fn straddleByte(s1: []const u8, s2: []const u8, index: usize) u8 {
1569 if (index < s1.len) {
1570 return s1[index];
1571 } else {
1572 return s2[index - s1.len];
1573 }
1574}
1575
1576const builtin = @import("builtin");
1577const native_endian = builtin.cpu.arch.endian();
1578
1579fn big(x: anytype) @TypeOf(x) {1270fn big(x: anytype) @TypeOf(x) {
1580 return switch (native_endian) {1271 return switch (native_endian) {
1581 .big => x,1272 .big => x,
...@@ -1836,81 +1527,6 @@ const CertificatePublicKey = struct {...@@ -1836,81 +1527,6 @@ const CertificatePublicKey = struct {
1836 }1527 }
1837};1528};
18381529
1839/// Abstraction for sending multiple byte buffers to a slice of iovecs.
1840const VecPut = struct {
1841 iovecs: []const std.posix.iovec,
1842 idx: usize = 0,
1843 off: usize = 0,
1844 total: usize = 0,
1845
1846 /// Returns the amount actually put which is always equal to bytes.len
1847 /// unless the vectors ran out of space.
1848 fn put(vp: *VecPut, bytes: []const u8) usize {
1849 if (vp.idx >= vp.iovecs.len) return 0;
1850 var bytes_i: usize = 0;
1851 while (true) {
1852 const v = vp.iovecs[vp.idx];
1853 const dest = v.base[vp.off..v.len];
1854 const src = bytes[bytes_i..][0..@min(dest.len, bytes.len - bytes_i)];
1855 @memcpy(dest[0..src.len], src);
1856 bytes_i += src.len;
1857 vp.off += src.len;
1858 if (vp.off >= v.len) {
1859 vp.off = 0;
1860 vp.idx += 1;
1861 if (vp.idx >= vp.iovecs.len) {
1862 vp.total += bytes_i;
1863 return bytes_i;
1864 }
1865 }
1866 if (bytes_i >= bytes.len) {
1867 vp.total += bytes_i;
1868 return bytes_i;
1869 }
1870 }
1871 }
1872
1873 /// Returns the next buffer that consecutive bytes can go into.
1874 fn peek(vp: VecPut) []u8 {
1875 if (vp.idx >= vp.iovecs.len) return &.{};
1876 const v = vp.iovecs[vp.idx];
1877 return v.base[vp.off..v.len];
1878 }
1879
1880 // After writing to the result of peek(), one can call next() to
1881 // advance the cursor.
1882 fn next(vp: *VecPut, len: usize) void {
1883 vp.total += len;
1884 vp.off += len;
1885 if (vp.off >= vp.iovecs[vp.idx].len) {
1886 vp.off = 0;
1887 vp.idx += 1;
1888 }
1889 }
1890
1891 fn freeSize(vp: VecPut) usize {
1892 if (vp.idx >= vp.iovecs.len) return 0;
1893 var total: usize = 0;
1894 total += vp.iovecs[vp.idx].len - vp.off;
1895 if (vp.idx + 1 >= vp.iovecs.len) return total;
1896 for (vp.iovecs[vp.idx + 1 ..]) |v| total += v.len;
1897 return total;
1898 }
1899};
1900
1901/// Limit iovecs to a specific byte size.
1902fn limitVecs(iovecs: []std.posix.iovec, len: usize) []std.posix.iovec {
1903 var bytes_left: usize = len;
1904 for (iovecs, 0..) |*iovec, vec_i| {
1905 if (bytes_left <= iovec.len) {
1906 iovec.len = bytes_left;
1907 return iovecs[0 .. vec_i + 1];
1908 }
1909 bytes_left -= iovec.len;
1910 }
1911 return iovecs;
1912}
1913
1914/// The priority order here is chosen based on what crypto algorithms Zig has1530/// The priority order here is chosen based on what crypto algorithms Zig has
1915/// available in the standard library as well as what is faster. Following are1531/// available in the standard library as well as what is faster. Following are
1916/// a few data points on the relative performance of these algorithms.1532/// a few data points on the relative performance of these algorithms.
...@@ -1954,7 +1570,3 @@ else...@@ -1954,7 +1570,3 @@ else
1954 .AES_256_GCM_SHA384,1570 .AES_256_GCM_SHA384,
1955 .ECDHE_RSA_WITH_AES_256_GCM_SHA384,1571 .ECDHE_RSA_WITH_AES_256_GCM_SHA384,
1956 });1572 });
1957
1958test {
1959 _ = StreamInterface;
1960}
lib/std/fifo.zig deleted-548
...@@ -1,548 +0,0 @@
1// FIFO of fixed size items
2// Usually used for e.g. byte buffers
3
4const std = @import("std");
5const math = std.math;
6const mem = std.mem;
7const Allocator = mem.Allocator;
8const assert = std.debug.assert;
9const testing = std.testing;
10
11pub const LinearFifoBufferType = union(enum) {
12 /// The buffer is internal to the fifo; it is of the specified size.
13 Static: usize,
14
15 /// The buffer is passed as a slice to the initialiser.
16 Slice,
17
18 /// The buffer is managed dynamically using a `mem.Allocator`.
19 Dynamic,
20};
21
22pub fn LinearFifo(
23 comptime T: type,
24 comptime buffer_type: LinearFifoBufferType,
25) type {
26 const autoalign = false;
27
28 const powers_of_two = switch (buffer_type) {
29 .Static => std.math.isPowerOfTwo(buffer_type.Static),
30 .Slice => false, // Any size slice could be passed in
31 .Dynamic => true, // This could be configurable in future
32 };
33
34 return struct {
35 allocator: if (buffer_type == .Dynamic) Allocator else void,
36 buf: if (buffer_type == .Static) [buffer_type.Static]T else []T,
37 head: usize,
38 count: usize,
39
40 const Self = @This();
41 pub const Reader = std.io.GenericReader(*Self, error{}, readFn);
42 pub const Writer = std.io.GenericWriter(*Self, error{OutOfMemory}, appendWrite);
43
44 // Type of Self argument for slice operations.
45 // If buffer is inline (Static) then we need to ensure we haven't
46 // returned a slice into a copy on the stack
47 const SliceSelfArg = if (buffer_type == .Static) *Self else Self;
48
49 pub const init = switch (buffer_type) {
50 .Static => initStatic,
51 .Slice => initSlice,
52 .Dynamic => initDynamic,
53 };
54
55 fn initStatic() Self {
56 comptime assert(buffer_type == .Static);
57 return .{
58 .allocator = {},
59 .buf = undefined,
60 .head = 0,
61 .count = 0,
62 };
63 }
64
65 fn initSlice(buf: []T) Self {
66 comptime assert(buffer_type == .Slice);
67 return .{
68 .allocator = {},
69 .buf = buf,
70 .head = 0,
71 .count = 0,
72 };
73 }
74
75 fn initDynamic(allocator: Allocator) Self {
76 comptime assert(buffer_type == .Dynamic);
77 return .{
78 .allocator = allocator,
79 .buf = &.{},
80 .head = 0,
81 .count = 0,
82 };
83 }
84
85 pub fn deinit(self: Self) void {
86 if (buffer_type == .Dynamic) self.allocator.free(self.buf);
87 }
88
89 pub fn realign(self: *Self) void {
90 if (self.buf.len - self.head >= self.count) {
91 mem.copyForwards(T, self.buf[0..self.count], self.buf[self.head..][0..self.count]);
92 self.head = 0;
93 } else {
94 var tmp: [4096 / 2 / @sizeOf(T)]T = undefined;
95
96 while (self.head != 0) {
97 const n = @min(self.head, tmp.len);
98 const m = self.buf.len - n;
99 @memcpy(tmp[0..n], self.buf[0..n]);
100 mem.copyForwards(T, self.buf[0..m], self.buf[n..][0..m]);
101 @memcpy(self.buf[m..][0..n], tmp[0..n]);
102 self.head -= n;
103 }
104 }
105 { // set unused area to undefined
106 const unused = mem.sliceAsBytes(self.buf[self.count..]);
107 @memset(unused, undefined);
108 }
109 }
110
111 /// Reduce allocated capacity to `size`.
112 pub fn shrink(self: *Self, size: usize) void {
113 assert(size >= self.count);
114 if (buffer_type == .Dynamic) {
115 self.realign();
116 self.buf = self.allocator.realloc(self.buf, size) catch |e| switch (e) {
117 error.OutOfMemory => return, // no problem, capacity is still correct then.
118 };
119 }
120 }
121
122 /// Ensure that the buffer can fit at least `size` items
123 pub fn ensureTotalCapacity(self: *Self, size: usize) !void {
124 if (self.buf.len >= size) return;
125 if (buffer_type == .Dynamic) {
126 self.realign();
127 const new_size = if (powers_of_two) math.ceilPowerOfTwo(usize, size) catch return error.OutOfMemory else size;
128 self.buf = try self.allocator.realloc(self.buf, new_size);
129 } else {
130 return error.OutOfMemory;
131 }
132 }
133
134 /// Makes sure at least `size` items are unused
135 pub fn ensureUnusedCapacity(self: *Self, size: usize) error{OutOfMemory}!void {
136 if (self.writableLength() >= size) return;
137
138 return try self.ensureTotalCapacity(math.add(usize, self.count, size) catch return error.OutOfMemory);
139 }
140
141 /// Returns number of items currently in fifo
142 pub fn readableLength(self: Self) usize {
143 return self.count;
144 }
145
146 /// Returns a writable slice from the 'read' end of the fifo
147 fn readableSliceMut(self: SliceSelfArg, offset: usize) []T {
148 if (offset > self.count) return &[_]T{};
149
150 var start = self.head + offset;
151 if (start >= self.buf.len) {
152 start -= self.buf.len;
153 return self.buf[start .. start + (self.count - offset)];
154 } else {
155 const end = @min(self.head + self.count, self.buf.len);
156 return self.buf[start..end];
157 }
158 }
159
160 /// Returns a readable slice from `offset`
161 pub fn readableSlice(self: SliceSelfArg, offset: usize) []const T {
162 return self.readableSliceMut(offset);
163 }
164
165 pub fn readableSliceOfLen(self: *Self, len: usize) []const T {
166 assert(len <= self.count);
167 const buf = self.readableSlice(0);
168 if (buf.len >= len) {
169 return buf[0..len];
170 } else {
171 self.realign();
172 return self.readableSlice(0)[0..len];
173 }
174 }
175
176 /// Discard first `count` items in the fifo
177 pub fn discard(self: *Self, count: usize) void {
178 assert(count <= self.count);
179 { // set old range to undefined. Note: may be wrapped around
180 const slice = self.readableSliceMut(0);
181 if (slice.len >= count) {
182 const unused = mem.sliceAsBytes(slice[0..count]);
183 @memset(unused, undefined);
184 } else {
185 const unused = mem.sliceAsBytes(slice[0..]);
186 @memset(unused, undefined);
187 const unused2 = mem.sliceAsBytes(self.readableSliceMut(slice.len)[0 .. count - slice.len]);
188 @memset(unused2, undefined);
189 }
190 }
191 if (autoalign and self.count == count) {
192 self.head = 0;
193 self.count = 0;
194 } else {
195 var head = self.head + count;
196 if (powers_of_two) {
197 // Note it is safe to do a wrapping subtract as
198 // bitwise & with all 1s is a noop
199 head &= self.buf.len -% 1;
200 } else {
201 head %= self.buf.len;
202 }
203 self.head = head;
204 self.count -= count;
205 }
206 }
207
208 /// Read the next item from the fifo
209 pub fn readItem(self: *Self) ?T {
210 if (self.count == 0) return null;
211
212 const c = self.buf[self.head];
213 self.discard(1);
214 return c;
215 }
216
217 /// Read data from the fifo into `dst`, returns number of items copied.
218 pub fn read(self: *Self, dst: []T) usize {
219 var dst_left = dst;
220
221 while (dst_left.len > 0) {
222 const slice = self.readableSlice(0);
223 if (slice.len == 0) break;
224 const n = @min(slice.len, dst_left.len);
225 @memcpy(dst_left[0..n], slice[0..n]);
226 self.discard(n);
227 dst_left = dst_left[n..];
228 }
229
230 return dst.len - dst_left.len;
231 }
232
233 /// Same as `read` except it returns an error union
234 /// The purpose of this function existing is to match `std.io.GenericReader` API.
235 fn readFn(self: *Self, dest: []u8) error{}!usize {
236 return self.read(dest);
237 }
238
239 pub fn reader(self: *Self) Reader {
240 return .{ .context = self };
241 }
242
243 /// Returns number of items available in fifo
244 pub fn writableLength(self: Self) usize {
245 return self.buf.len - self.count;
246 }
247
248 /// Returns the first section of writable buffer.
249 /// Note that this may be of length 0
250 pub fn writableSlice(self: SliceSelfArg, offset: usize) []T {
251 if (offset > self.buf.len) return &[_]T{};
252
253 const tail = self.head + offset + self.count;
254 if (tail < self.buf.len) {
255 return self.buf[tail..];
256 } else {
257 return self.buf[tail - self.buf.len ..][0 .. self.writableLength() - offset];
258 }
259 }
260
261 /// Returns a writable buffer of at least `size` items, allocating memory as needed.
262 /// Use `fifo.update` once you've written data to it.
263 pub fn writableWithSize(self: *Self, size: usize) ![]T {
264 try self.ensureUnusedCapacity(size);
265
266 // try to avoid realigning buffer
267 var slice = self.writableSlice(0);
268 if (slice.len < size) {
269 self.realign();
270 slice = self.writableSlice(0);
271 }
272 return slice;
273 }
274
275 /// Update the tail location of the buffer (usually follows use of writable/writableWithSize)
276 pub fn update(self: *Self, count: usize) void {
277 assert(self.count + count <= self.buf.len);
278 self.count += count;
279 }
280
281 /// Appends the data in `src` to the fifo.
282 /// You must have ensured there is enough space.
283 pub fn writeAssumeCapacity(self: *Self, src: []const T) void {
284 assert(self.writableLength() >= src.len);
285
286 var src_left = src;
287 while (src_left.len > 0) {
288 const writable_slice = self.writableSlice(0);
289 assert(writable_slice.len != 0);
290 const n = @min(writable_slice.len, src_left.len);
291 @memcpy(writable_slice[0..n], src_left[0..n]);
292 self.update(n);
293 src_left = src_left[n..];
294 }
295 }
296
297 /// Write a single item to the fifo
298 pub fn writeItem(self: *Self, item: T) !void {
299 try self.ensureUnusedCapacity(1);
300 return self.writeItemAssumeCapacity(item);
301 }
302
303 pub fn writeItemAssumeCapacity(self: *Self, item: T) void {
304 var tail = self.head + self.count;
305 if (powers_of_two) {
306 tail &= self.buf.len - 1;
307 } else {
308 tail %= self.buf.len;
309 }
310 self.buf[tail] = item;
311 self.update(1);
312 }
313
314 /// Appends the data in `src` to the fifo.
315 /// Allocates more memory as necessary
316 pub fn write(self: *Self, src: []const T) !void {
317 try self.ensureUnusedCapacity(src.len);
318
319 return self.writeAssumeCapacity(src);
320 }
321
322 /// Same as `write` except it returns the number of bytes written, which is always the same
323 /// as `bytes.len`. The purpose of this function existing is to match `std.io.GenericWriter` API.
324 fn appendWrite(self: *Self, bytes: []const u8) error{OutOfMemory}!usize {
325 try self.write(bytes);
326 return bytes.len;
327 }
328
329 pub fn writer(self: *Self) Writer {
330 return .{ .context = self };
331 }
332
333 /// Make `count` items available before the current read location
334 fn rewind(self: *Self, count: usize) void {
335 assert(self.writableLength() >= count);
336
337 var head = self.head + (self.buf.len - count);
338 if (powers_of_two) {
339 head &= self.buf.len - 1;
340 } else {
341 head %= self.buf.len;
342 }
343 self.head = head;
344 self.count += count;
345 }
346
347 /// Place data back into the read stream
348 pub fn unget(self: *Self, src: []const T) !void {
349 try self.ensureUnusedCapacity(src.len);
350
351 self.rewind(src.len);
352
353 const slice = self.readableSliceMut(0);
354 if (src.len < slice.len) {
355 @memcpy(slice[0..src.len], src);
356 } else {
357 @memcpy(slice, src[0..slice.len]);
358 const slice2 = self.readableSliceMut(slice.len);
359 @memcpy(slice2[0 .. src.len - slice.len], src[slice.len..]);
360 }
361 }
362
363 /// Returns the item at `offset`.
364 /// Asserts offset is within bounds.
365 pub fn peekItem(self: Self, offset: usize) T {
366 assert(offset < self.count);
367
368 var index = self.head + offset;
369 if (powers_of_two) {
370 index &= self.buf.len - 1;
371 } else {
372 index %= self.buf.len;
373 }
374 return self.buf[index];
375 }
376
377 /// Pump data from a reader into a writer.
378 /// Stops when reader returns 0 bytes (EOF).
379 /// Buffer size must be set before calling; a buffer length of 0 is invalid.
380 pub fn pump(self: *Self, src_reader: anytype, dest_writer: anytype) !void {
381 assert(self.buf.len > 0);
382 while (true) {
383 if (self.writableLength() > 0) {
384 const n = try src_reader.read(self.writableSlice(0));
385 if (n == 0) break; // EOF
386 self.update(n);
387 }
388 self.discard(try dest_writer.write(self.readableSlice(0)));
389 }
390 // flush remaining data
391 while (self.readableLength() > 0) {
392 self.discard(try dest_writer.write(self.readableSlice(0)));
393 }
394 }
395
396 pub fn toOwnedSlice(self: *Self) Allocator.Error![]T {
397 if (self.head != 0) self.realign();
398 assert(self.head == 0);
399 assert(self.count <= self.buf.len);
400 const allocator = self.allocator;
401 if (allocator.resize(self.buf, self.count)) {
402 const result = self.buf[0..self.count];
403 self.* = Self.init(allocator);
404 return result;
405 }
406 const new_memory = try allocator.dupe(T, self.buf[0..self.count]);
407 allocator.free(self.buf);
408 self.* = Self.init(allocator);
409 return new_memory;
410 }
411 };
412}
413
414test "LinearFifo(u8, .Dynamic) discard(0) from empty buffer should not error on overflow" {
415 var fifo = LinearFifo(u8, .Dynamic).init(testing.allocator);
416 defer fifo.deinit();
417
418 // If overflow is not explicitly allowed this will crash in debug / safe mode
419 fifo.discard(0);
420}
421
422test "LinearFifo(u8, .Dynamic)" {
423 var fifo = LinearFifo(u8, .Dynamic).init(testing.allocator);
424 defer fifo.deinit();
425
426 try fifo.write("HELLO");
427 try testing.expectEqual(@as(usize, 5), fifo.readableLength());
428 try testing.expectEqualSlices(u8, "HELLO", fifo.readableSlice(0));
429
430 {
431 var i: usize = 0;
432 while (i < 5) : (i += 1) {
433 try fifo.write(&[_]u8{fifo.peekItem(i)});
434 }
435 try testing.expectEqual(@as(usize, 10), fifo.readableLength());
436 try testing.expectEqualSlices(u8, "HELLOHELLO", fifo.readableSlice(0));
437 }
438
439 {
440 try testing.expectEqual(@as(u8, 'H'), fifo.readItem().?);
441 try testing.expectEqual(@as(u8, 'E'), fifo.readItem().?);
442 try testing.expectEqual(@as(u8, 'L'), fifo.readItem().?);
443 try testing.expectEqual(@as(u8, 'L'), fifo.readItem().?);
444 try testing.expectEqual(@as(u8, 'O'), fifo.readItem().?);
445 }
446 try testing.expectEqual(@as(usize, 5), fifo.readableLength());
447
448 { // Writes that wrap around
449 try testing.expectEqual(@as(usize, 11), fifo.writableLength());
450 try testing.expectEqual(@as(usize, 6), fifo.writableSlice(0).len);
451 fifo.writeAssumeCapacity("6<chars<11");
452 try testing.expectEqualSlices(u8, "HELLO6<char", fifo.readableSlice(0));
453 try testing.expectEqualSlices(u8, "s<11", fifo.readableSlice(11));
454 try testing.expectEqualSlices(u8, "11", fifo.readableSlice(13));
455 try testing.expectEqualSlices(u8, "", fifo.readableSlice(15));
456 fifo.discard(11);
457 try testing.expectEqualSlices(u8, "s<11", fifo.readableSlice(0));
458 fifo.discard(4);
459 try testing.expectEqual(@as(usize, 0), fifo.readableLength());
460 }
461
462 {
463 const buf = try fifo.writableWithSize(12);
464 try testing.expectEqual(@as(usize, 12), buf.len);
465 var i: u8 = 0;
466 while (i < 10) : (i += 1) {
467 buf[i] = i + 'a';
468 }
469 fifo.update(10);
470 try testing.expectEqualSlices(u8, "abcdefghij", fifo.readableSlice(0));
471 }
472
473 {
474 try fifo.unget("prependedstring");
475 var result: [30]u8 = undefined;
476 try testing.expectEqualSlices(u8, "prependedstringabcdefghij", result[0..fifo.read(&result)]);
477 try fifo.unget("b");
478 try fifo.unget("a");
479 try testing.expectEqualSlices(u8, "ab", result[0..fifo.read(&result)]);
480 }
481
482 fifo.shrink(0);
483
484 {
485 try fifo.writer().print("{s}, {s}!", .{ "Hello", "World" });
486 var result: [30]u8 = undefined;
487 try testing.expectEqualSlices(u8, "Hello, World!", result[0..fifo.read(&result)]);
488 try testing.expectEqual(@as(usize, 0), fifo.readableLength());
489 }
490
491 {
492 try fifo.writer().writeAll("This is a test");
493 var result: [30]u8 = undefined;
494 try testing.expectEqualSlices(u8, "This", (try fifo.reader().readUntilDelimiterOrEof(&result, ' ')).?);
495 try testing.expectEqualSlices(u8, "is", (try fifo.reader().readUntilDelimiterOrEof(&result, ' ')).?);
496 try testing.expectEqualSlices(u8, "a", (try fifo.reader().readUntilDelimiterOrEof(&result, ' ')).?);
497 try testing.expectEqualSlices(u8, "test", (try fifo.reader().readUntilDelimiterOrEof(&result, ' ')).?);
498 }
499
500 {
501 try fifo.ensureTotalCapacity(1);
502 var in_fbs = std.io.fixedBufferStream("pump test");
503 var out_buf: [50]u8 = undefined;
504 var out_fbs = std.io.fixedBufferStream(&out_buf);
505 try fifo.pump(in_fbs.reader(), out_fbs.writer());
506 try testing.expectEqualSlices(u8, in_fbs.buffer, out_fbs.getWritten());
507 }
508}
509
510test LinearFifo {
511 inline for ([_]type{ u1, u8, u16, u64 }) |T| {
512 inline for ([_]LinearFifoBufferType{ LinearFifoBufferType{ .Static = 32 }, .Slice, .Dynamic }) |bt| {
513 const FifoType = LinearFifo(T, bt);
514 var buf: if (bt == .Slice) [32]T else void = undefined;
515 var fifo = switch (bt) {
516 .Static => FifoType.init(),
517 .Slice => FifoType.init(buf[0..]),
518 .Dynamic => FifoType.init(testing.allocator),
519 };
520 defer fifo.deinit();
521
522 try fifo.write(&[_]T{ 0, 1, 1, 0, 1 });
523 try testing.expectEqual(@as(usize, 5), fifo.readableLength());
524
525 {
526 try testing.expectEqual(@as(T, 0), fifo.readItem().?);
527 try testing.expectEqual(@as(T, 1), fifo.readItem().?);
528 try testing.expectEqual(@as(T, 1), fifo.readItem().?);
529 try testing.expectEqual(@as(T, 0), fifo.readItem().?);
530 try testing.expectEqual(@as(T, 1), fifo.readItem().?);
531 try testing.expectEqual(@as(usize, 0), fifo.readableLength());
532 }
533
534 {
535 try fifo.writeItem(1);
536 try fifo.writeItem(1);
537 try fifo.writeItem(1);
538 try testing.expectEqual(@as(usize, 3), fifo.readableLength());
539 }
540
541 {
542 var readBuf: [3]T = undefined;
543 const n = fifo.read(&readBuf);
544 try testing.expectEqual(@as(usize, 3), n); // NOTE: It should be the number of items.
545 }
546 }
547 }
548}
lib/std/fs/File.zig+2-4
...@@ -1351,8 +1351,7 @@ pub const Reader = struct {...@@ -1351,8 +1351,7 @@ pub const Reader = struct {
1351 }1351 }
1352 r.pos += n;1352 r.pos += n;
1353 if (n > data_size) {1353 if (n > data_size) {
1354 io_reader.seek = 0;1354 io_reader.end += n - data_size;
1355 io_reader.end = n - data_size;
1356 return data_size;1355 return data_size;
1357 }1356 }
1358 return n;1357 return n;
...@@ -1386,8 +1385,7 @@ pub const Reader = struct {...@@ -1386,8 +1385,7 @@ pub const Reader = struct {
1386 }1385 }
1387 r.pos += n;1386 r.pos += n;
1388 if (n > data_size) {1387 if (n > data_size) {
1389 io_reader.seek = 0;1388 io_reader.end += n - data_size;
1390 io_reader.end = n - data_size;
1391 return data_size;1389 return data_size;
1392 }1390 }
1393 return n;1391 return n;
lib/std/http.zig+820-56
...@@ -1,14 +1,14 @@...@@ -1,14 +1,14 @@
1const builtin = @import("builtin");1const builtin = @import("builtin");
2const std = @import("std.zig");2const std = @import("std.zig");
3const assert = std.debug.assert;3const assert = std.debug.assert;
4const Writer = std.Io.Writer;
5const File = std.fs.File;
46
5pub const Client = @import("http/Client.zig");7pub const Client = @import("http/Client.zig");
6pub const Server = @import("http/Server.zig");8pub const Server = @import("http/Server.zig");
7pub const protocol = @import("http/protocol.zig");
8pub const HeadParser = @import("http/HeadParser.zig");9pub const HeadParser = @import("http/HeadParser.zig");
9pub const ChunkParser = @import("http/ChunkParser.zig");10pub const ChunkParser = @import("http/ChunkParser.zig");
10pub const HeaderIterator = @import("http/HeaderIterator.zig");11pub const HeaderIterator = @import("http/HeaderIterator.zig");
11pub const WebSocket = @import("http/WebSocket.zig");
1212
13pub const Version = enum {13pub const Version = enum {
14 @"HTTP/1.0",14 @"HTTP/1.0",
...@@ -20,51 +20,32 @@ pub const Version = enum {...@@ -20,51 +20,32 @@ pub const Version = enum {
20/// https://datatracker.ietf.org/doc/html/rfc7231#section-4 Initial definition20/// https://datatracker.ietf.org/doc/html/rfc7231#section-4 Initial definition
21///21///
22/// https://datatracker.ietf.org/doc/html/rfc5789#section-2 PATCH22/// https://datatracker.ietf.org/doc/html/rfc5789#section-2 PATCH
23pub const Method = enum(u64) {23pub const Method = enum {
24 GET = parse("GET"),24 GET,
25 HEAD = parse("HEAD"),25 HEAD,
26 POST = parse("POST"),26 POST,
27 PUT = parse("PUT"),27 PUT,
28 DELETE = parse("DELETE"),28 DELETE,
29 CONNECT = parse("CONNECT"),29 CONNECT,
30 OPTIONS = parse("OPTIONS"),30 OPTIONS,
31 TRACE = parse("TRACE"),31 TRACE,
32 PATCH = parse("PATCH"),32 PATCH,
33
34 _,
35
36 /// Converts `s` into a type that may be used as a `Method` field.
37 /// Asserts that `s` is 24 or fewer bytes.
38 pub fn parse(s: []const u8) u64 {
39 var x: u64 = 0;
40 const len = @min(s.len, @sizeOf(@TypeOf(x)));
41 @memcpy(std.mem.asBytes(&x)[0..len], s[0..len]);
42 return x;
43 }
44
45 pub fn format(self: Method, w: *std.io.Writer) std.io.Writer.Error!void {
46 const bytes: []const u8 = @ptrCast(&@intFromEnum(self));
47 const str = std.mem.sliceTo(bytes, 0);
48 try w.writeAll(str);
49 }
5033
51 /// Returns true if a request of this method is allowed to have a body34 /// Returns true if a request of this method is allowed to have a body
52 /// Actual behavior from servers may vary and should still be checked35 /// Actual behavior from servers may vary and should still be checked
53 pub fn requestHasBody(self: Method) bool {36 pub fn requestHasBody(m: Method) bool {
54 return switch (self) {37 return switch (m) {
55 .POST, .PUT, .PATCH => true,38 .POST, .PUT, .PATCH => true,
56 .GET, .HEAD, .DELETE, .CONNECT, .OPTIONS, .TRACE => false,39 .GET, .HEAD, .DELETE, .CONNECT, .OPTIONS, .TRACE => false,
57 else => true,
58 };40 };
59 }41 }
6042
61 /// Returns true if a response to this method is allowed to have a body43 /// Returns true if a response to this method is allowed to have a body
62 /// Actual behavior from clients may vary and should still be checked44 /// Actual behavior from clients may vary and should still be checked
63 pub fn responseHasBody(self: Method) bool {45 pub fn responseHasBody(m: Method) bool {
64 return switch (self) {46 return switch (m) {
65 .GET, .POST, .DELETE, .CONNECT, .OPTIONS, .PATCH => true,47 .GET, .POST, .DELETE, .CONNECT, .OPTIONS, .PATCH => true,
66 .HEAD, .PUT, .TRACE => false,48 .HEAD, .PUT, .TRACE => false,
67 else => true,
68 };49 };
69 }50 }
7051
...@@ -73,11 +54,10 @@ pub const Method = enum(u64) {...@@ -73,11 +54,10 @@ pub const Method = enum(u64) {
73 /// https://developer.mozilla.org/en-US/docs/Glossary/Safe/HTTP54 /// https://developer.mozilla.org/en-US/docs/Glossary/Safe/HTTP
74 ///55 ///
75 /// https://datatracker.ietf.org/doc/html/rfc7231#section-4.2.156 /// https://datatracker.ietf.org/doc/html/rfc7231#section-4.2.1
76 pub fn safe(self: Method) bool {57 pub fn safe(m: Method) bool {
77 return switch (self) {58 return switch (m) {
78 .GET, .HEAD, .OPTIONS, .TRACE => true,59 .GET, .HEAD, .OPTIONS, .TRACE => true,
79 .POST, .PUT, .DELETE, .CONNECT, .PATCH => false,60 .POST, .PUT, .DELETE, .CONNECT, .PATCH => false,
80 else => false,
81 };61 };
82 }62 }
8363
...@@ -88,11 +68,10 @@ pub const Method = enum(u64) {...@@ -88,11 +68,10 @@ pub const Method = enum(u64) {
88 /// https://developer.mozilla.org/en-US/docs/Glossary/Idempotent68 /// https://developer.mozilla.org/en-US/docs/Glossary/Idempotent
89 ///69 ///
90 /// https://datatracker.ietf.org/doc/html/rfc7231#section-4.2.270 /// https://datatracker.ietf.org/doc/html/rfc7231#section-4.2.2
91 pub fn idempotent(self: Method) bool {71 pub fn idempotent(m: Method) bool {
92 return switch (self) {72 return switch (m) {
93 .GET, .HEAD, .PUT, .DELETE, .OPTIONS, .TRACE => true,73 .GET, .HEAD, .PUT, .DELETE, .OPTIONS, .TRACE => true,
94 .CONNECT, .POST, .PATCH => false,74 .CONNECT, .POST, .PATCH => false,
95 else => false,
96 };75 };
97 }76 }
9877
...@@ -102,11 +81,10 @@ pub const Method = enum(u64) {...@@ -102,11 +81,10 @@ pub const Method = enum(u64) {
102 /// https://developer.mozilla.org/en-US/docs/Glossary/cacheable81 /// https://developer.mozilla.org/en-US/docs/Glossary/cacheable
103 ///82 ///
104 /// https://datatracker.ietf.org/doc/html/rfc7231#section-4.2.383 /// https://datatracker.ietf.org/doc/html/rfc7231#section-4.2.3
105 pub fn cacheable(self: Method) bool {84 pub fn cacheable(m: Method) bool {
106 return switch (self) {85 return switch (m) {
107 .GET, .HEAD => true,86 .GET, .HEAD => true,
108 .POST, .PUT, .DELETE, .CONNECT, .OPTIONS, .TRACE, .PATCH => false,87 .POST, .PUT, .DELETE, .CONNECT, .OPTIONS, .TRACE, .PATCH => false,
109 else => false,
110 };88 };
111 }89 }
112};90};
...@@ -296,13 +274,24 @@ pub const TransferEncoding = enum {...@@ -296,13 +274,24 @@ pub const TransferEncoding = enum {
296};274};
297275
298pub const ContentEncoding = enum {276pub const ContentEncoding = enum {
299 identity,
300 compress,
301 @"x-compress",
302 deflate,
303 gzip,
304 @"x-gzip",
305 zstd,277 zstd,
278 gzip,
279 deflate,
280 compress,
281 identity,
282
283 pub fn fromString(s: []const u8) ?ContentEncoding {
284 const map = std.StaticStringMap(ContentEncoding).initComptime(.{
285 .{ "zstd", .zstd },
286 .{ "gzip", .gzip },
287 .{ "x-gzip", .gzip },
288 .{ "deflate", .deflate },
289 .{ "compress", .compress },
290 .{ "x-compress", .compress },
291 .{ "identity", .identity },
292 });
293 return map.get(s);
294 }
306};295};
307296
308pub const Connection = enum {297pub const Connection = enum {
...@@ -315,15 +304,790 @@ pub const Header = struct {...@@ -315,15 +304,790 @@ pub const Header = struct {
315 value: []const u8,304 value: []const u8,
316};305};
317306
307pub const Reader = struct {
308 in: *std.Io.Reader,
309 /// This is preallocated memory that might be used by `bodyReader`. That
310 /// function might return a pointer to this field, or a different
311 /// `*std.Io.Reader`. Advisable to not access this field directly.
312 interface: std.Io.Reader,
313 /// Keeps track of whether the stream is ready to accept a new request,
314 /// making invalid API usage cause assertion failures rather than HTTP
315 /// protocol violations.
316 state: State,
317 /// HTTP trailer bytes. These are at the end of a transfer-encoding:
318 /// chunked message. This data is available only after calling one of the
319 /// "end" functions and points to data inside the buffer of `in`, and is
320 /// therefore invalidated on the next call to `receiveHead`, or any other
321 /// read from `in`.
322 trailers: []const u8 = &.{},
323 body_err: ?BodyError = null,
324
325 pub const RemainingChunkLen = enum(u64) {
326 head = 0,
327 n = 1,
328 rn = 2,
329 _,
330
331 pub fn init(integer: u64) RemainingChunkLen {
332 return @enumFromInt(integer);
333 }
334
335 pub fn int(rcl: RemainingChunkLen) u64 {
336 return @intFromEnum(rcl);
337 }
338 };
339
340 pub const State = union(enum) {
341 /// The stream is available to be used for the first time, or reused.
342 ready,
343 received_head,
344 /// The stream goes until the connection is closed.
345 body_none,
346 body_remaining_content_length: u64,
347 body_remaining_chunk_len: RemainingChunkLen,
348 /// The stream would be eligible for another HTTP request, however the
349 /// client and server did not negotiate a persistent connection.
350 closing,
351 };
352
353 pub const BodyError = error{
354 HttpChunkInvalid,
355 HttpChunkTruncated,
356 HttpHeadersOversize,
357 };
358
359 pub const HeadError = error{
360 /// Too many bytes of HTTP headers.
361 ///
362 /// The HTTP specification suggests to respond with a 431 status code
363 /// before closing the connection.
364 HttpHeadersOversize,
365 /// Partial HTTP request was received but the connection was closed
366 /// before fully receiving the headers.
367 HttpRequestTruncated,
368 /// The client sent 0 bytes of headers before closing the stream. This
369 /// happens when a keep-alive connection is finally closed.
370 HttpConnectionClosing,
371 /// Transitive error occurred reading from `in`.
372 ReadFailed,
373 };
374
375 /// Buffers the entire head inside `in`.
376 ///
377 /// The resulting memory is invalidated by any subsequent consumption of
378 /// the input stream.
379 pub fn receiveHead(reader: *Reader) HeadError![]const u8 {
380 reader.trailers = &.{};
381 const in = reader.in;
382 var hp: HeadParser = .{};
383 var head_len: usize = 0;
384 while (true) {
385 if (in.buffer.len - head_len == 0) return error.HttpHeadersOversize;
386 const remaining = in.buffered()[head_len..];
387 if (remaining.len == 0) {
388 in.fillMore() catch |err| switch (err) {
389 error.EndOfStream => switch (head_len) {
390 0 => return error.HttpConnectionClosing,
391 else => return error.HttpRequestTruncated,
392 },
393 error.ReadFailed => return error.ReadFailed,
394 };
395 continue;
396 }
397 head_len += hp.feed(remaining);
398 if (hp.state == .finished) {
399 reader.state = .received_head;
400 const head_buffer = in.buffered()[0..head_len];
401 in.toss(head_len);
402 return head_buffer;
403 }
404 }
405 }
406
407 /// If compressed body has been negotiated this will return compressed bytes.
408 ///
409 /// Asserts only called once and after `receiveHead`.
410 ///
411 /// See also:
412 /// * `interfaceDecompressing`
413 pub fn bodyReader(
414 reader: *Reader,
415 buffer: []u8,
416 transfer_encoding: TransferEncoding,
417 content_length: ?u64,
418 ) *std.Io.Reader {
419 assert(reader.state == .received_head);
420 switch (transfer_encoding) {
421 .chunked => {
422 reader.state = .{ .body_remaining_chunk_len = .head };
423 reader.interface = .{
424 .buffer = buffer,
425 .seek = 0,
426 .end = 0,
427 .vtable = &.{
428 .stream = chunkedStream,
429 .discard = chunkedDiscard,
430 },
431 };
432 return &reader.interface;
433 },
434 .none => {
435 if (content_length) |len| {
436 reader.state = .{ .body_remaining_content_length = len };
437 reader.interface = .{
438 .buffer = buffer,
439 .seek = 0,
440 .end = 0,
441 .vtable = &.{
442 .stream = contentLengthStream,
443 .discard = contentLengthDiscard,
444 },
445 };
446 return &reader.interface;
447 } else {
448 reader.state = .body_none;
449 return reader.in;
450 }
451 },
452 }
453 }
454
455 /// If compressed body has been negotiated this will return decompressed bytes.
456 ///
457 /// Asserts only called once and after `receiveHead`.
458 ///
459 /// See also:
460 /// * `interface`
461 pub fn bodyReaderDecompressing(
462 reader: *Reader,
463 transfer_encoding: TransferEncoding,
464 content_length: ?u64,
465 content_encoding: ContentEncoding,
466 decompressor: *Decompressor,
467 decompression_buffer: []u8,
468 ) *std.Io.Reader {
469 if (transfer_encoding == .none and content_length == null) {
470 assert(reader.state == .received_head);
471 reader.state = .body_none;
472 switch (content_encoding) {
473 .identity => {
474 return reader.in;
475 },
476 .deflate => {
477 decompressor.* = .{ .flate = .init(reader.in, .zlib, decompression_buffer) };
478 return &decompressor.flate.reader;
479 },
480 .gzip => {
481 decompressor.* = .{ .flate = .init(reader.in, .gzip, decompression_buffer) };
482 return &decompressor.flate.reader;
483 },
484 .zstd => {
485 decompressor.* = .{ .zstd = .init(reader.in, decompression_buffer, .{ .verify_checksum = false }) };
486 return &decompressor.zstd.reader;
487 },
488 .compress => unreachable,
489 }
490 }
491 const transfer_reader = bodyReader(reader, &.{}, transfer_encoding, content_length);
492 return decompressor.init(transfer_reader, decompression_buffer, content_encoding);
493 }
494
495 fn contentLengthStream(
496 io_r: *std.Io.Reader,
497 w: *Writer,
498 limit: std.Io.Limit,
499 ) std.Io.Reader.StreamError!usize {
500 const reader: *Reader = @alignCast(@fieldParentPtr("interface", io_r));
501 const remaining_content_length = &reader.state.body_remaining_content_length;
502 const remaining = remaining_content_length.*;
503 if (remaining == 0) {
504 reader.state = .ready;
505 return error.EndOfStream;
506 }
507 const n = try reader.in.stream(w, limit.min(.limited64(remaining)));
508 remaining_content_length.* = remaining - n;
509 return n;
510 }
511
512 fn contentLengthDiscard(io_r: *std.Io.Reader, limit: std.Io.Limit) std.Io.Reader.Error!usize {
513 const reader: *Reader = @alignCast(@fieldParentPtr("interface", io_r));
514 const remaining_content_length = &reader.state.body_remaining_content_length;
515 const remaining = remaining_content_length.*;
516 if (remaining == 0) {
517 reader.state = .ready;
518 return error.EndOfStream;
519 }
520 const n = try reader.in.discard(limit.min(.limited64(remaining)));
521 remaining_content_length.* = remaining - n;
522 return n;
523 }
524
525 fn chunkedStream(io_r: *std.Io.Reader, w: *Writer, limit: std.Io.Limit) std.Io.Reader.StreamError!usize {
526 const reader: *Reader = @alignCast(@fieldParentPtr("interface", io_r));
527 const chunk_len_ptr = switch (reader.state) {
528 .ready => return error.EndOfStream,
529 .body_remaining_chunk_len => |*x| x,
530 else => unreachable,
531 };
532 return chunkedReadEndless(reader, w, limit, chunk_len_ptr) catch |err| switch (err) {
533 error.ReadFailed => return error.ReadFailed,
534 error.WriteFailed => return error.WriteFailed,
535 error.EndOfStream => {
536 reader.body_err = error.HttpChunkTruncated;
537 return error.ReadFailed;
538 },
539 else => |e| {
540 reader.body_err = e;
541 return error.ReadFailed;
542 },
543 };
544 }
545
546 fn chunkedReadEndless(
547 reader: *Reader,
548 w: *Writer,
549 limit: std.Io.Limit,
550 chunk_len_ptr: *RemainingChunkLen,
551 ) (BodyError || std.Io.Reader.StreamError)!usize {
552 const in = reader.in;
553 len: switch (chunk_len_ptr.*) {
554 .head => {
555 var cp: ChunkParser = .init;
556 while (true) {
557 const i = cp.feed(in.buffered());
558 switch (cp.state) {
559 .invalid => return error.HttpChunkInvalid,
560 .data => {
561 in.toss(i);
562 break;
563 },
564 else => {
565 in.toss(i);
566 try in.fillMore();
567 continue;
568 },
569 }
570 }
571 if (cp.chunk_len == 0) return parseTrailers(reader, 0);
572 const n = try in.stream(w, limit.min(.limited64(cp.chunk_len)));
573 chunk_len_ptr.* = .init(cp.chunk_len + 2 - n);
574 return n;
575 },
576 .n => {
577 if ((try in.peekByte()) != '\n') return error.HttpChunkInvalid;
578 in.toss(1);
579 continue :len .head;
580 },
581 .rn => {
582 const rn = try in.peekArray(2);
583 if (rn[0] != '\r' or rn[1] != '\n') return error.HttpChunkInvalid;
584 in.toss(2);
585 continue :len .head;
586 },
587 else => |remaining_chunk_len| {
588 const n = try in.stream(w, limit.min(.limited64(@intFromEnum(remaining_chunk_len) - 2)));
589 chunk_len_ptr.* = .init(@intFromEnum(remaining_chunk_len) - n);
590 return n;
591 },
592 }
593 }
594
595 fn chunkedDiscard(io_r: *std.Io.Reader, limit: std.Io.Limit) std.Io.Reader.Error!usize {
596 const reader: *Reader = @alignCast(@fieldParentPtr("interface", io_r));
597 const chunk_len_ptr = switch (reader.state) {
598 .ready => return error.EndOfStream,
599 .body_remaining_chunk_len => |*x| x,
600 else => unreachable,
601 };
602 return chunkedDiscardEndless(reader, limit, chunk_len_ptr) catch |err| switch (err) {
603 error.ReadFailed => return error.ReadFailed,
604 error.EndOfStream => {
605 reader.body_err = error.HttpChunkTruncated;
606 return error.ReadFailed;
607 },
608 else => |e| {
609 reader.body_err = e;
610 return error.ReadFailed;
611 },
612 };
613 }
614
615 fn chunkedDiscardEndless(
616 reader: *Reader,
617 limit: std.Io.Limit,
618 chunk_len_ptr: *RemainingChunkLen,
619 ) (BodyError || std.Io.Reader.Error)!usize {
620 const in = reader.in;
621 len: switch (chunk_len_ptr.*) {
622 .head => {
623 var cp: ChunkParser = .init;
624 while (true) {
625 const i = cp.feed(in.buffered());
626 switch (cp.state) {
627 .invalid => return error.HttpChunkInvalid,
628 .data => {
629 in.toss(i);
630 break;
631 },
632 else => {
633 in.toss(i);
634 try in.fillMore();
635 continue;
636 },
637 }
638 }
639 if (cp.chunk_len == 0) return parseTrailers(reader, 0);
640 const n = try in.discard(limit.min(.limited64(cp.chunk_len)));
641 chunk_len_ptr.* = .init(cp.chunk_len + 2 - n);
642 return n;
643 },
644 .n => {
645 if ((try in.peekByte()) != '\n') return error.HttpChunkInvalid;
646 in.toss(1);
647 continue :len .head;
648 },
649 .rn => {
650 const rn = try in.peekArray(2);
651 if (rn[0] != '\r' or rn[1] != '\n') return error.HttpChunkInvalid;
652 in.toss(2);
653 continue :len .head;
654 },
655 else => |remaining_chunk_len| {
656 const n = try in.discard(limit.min(.limited64(remaining_chunk_len.int() - 2)));
657 chunk_len_ptr.* = .init(remaining_chunk_len.int() - n);
658 return n;
659 },
660 }
661 }
662
663 /// Called when next bytes in the stream are trailers, or "\r\n" to indicate
664 /// end of chunked body.
665 fn parseTrailers(reader: *Reader, amt_read: usize) (BodyError || std.Io.Reader.Error)!usize {
666 const in = reader.in;
667 const rn = try in.peekArray(2);
668 if (rn[0] == '\r' and rn[1] == '\n') {
669 in.toss(2);
670 reader.state = .ready;
671 assert(reader.trailers.len == 0);
672 return amt_read;
673 }
674 var hp: HeadParser = .{ .state = .seen_rn };
675 var trailers_len: usize = 2;
676 while (true) {
677 if (in.buffer.len - trailers_len == 0) return error.HttpHeadersOversize;
678 const remaining = in.buffered()[trailers_len..];
679 if (remaining.len == 0) {
680 try in.fillMore();
681 continue;
682 }
683 trailers_len += hp.feed(remaining);
684 if (hp.state == .finished) {
685 reader.state = .ready;
686 reader.trailers = in.buffered()[0..trailers_len];
687 in.toss(trailers_len);
688 return amt_read;
689 }
690 }
691 }
692};
693
694pub const Decompressor = union(enum) {
695 flate: std.compress.flate.Decompress,
696 zstd: std.compress.zstd.Decompress,
697 none: *std.Io.Reader,
698
699 pub fn init(
700 decompressor: *Decompressor,
701 transfer_reader: *std.Io.Reader,
702 buffer: []u8,
703 content_encoding: ContentEncoding,
704 ) *std.Io.Reader {
705 switch (content_encoding) {
706 .identity => {
707 decompressor.* = .{ .none = transfer_reader };
708 return transfer_reader;
709 },
710 .deflate => {
711 decompressor.* = .{ .flate = .init(transfer_reader, .zlib, buffer) };
712 return &decompressor.flate.reader;
713 },
714 .gzip => {
715 decompressor.* = .{ .flate = .init(transfer_reader, .gzip, buffer) };
716 return &decompressor.flate.reader;
717 },
718 .zstd => {
719 decompressor.* = .{ .zstd = .init(transfer_reader, buffer, .{ .verify_checksum = false }) };
720 return &decompressor.zstd.reader;
721 },
722 .compress => unreachable,
723 }
724 }
725};
726
727/// Request or response body.
728pub const BodyWriter = struct {
729 /// Until the lifetime of `BodyWriter` ends, it is illegal to modify the
730 /// state of this other than via methods of `BodyWriter`.
731 http_protocol_output: *Writer,
732 state: State,
733 writer: Writer,
734
735 pub const Error = Writer.Error;
736
737 /// How many zeroes to reserve for hex-encoded chunk length.
738 const chunk_len_digits = 8;
739 const max_chunk_len: usize = std.math.pow(u64, 16, chunk_len_digits) - 1;
740 const chunk_header_template = ("0" ** chunk_len_digits) ++ "\r\n";
741
742 comptime {
743 assert(max_chunk_len == std.math.maxInt(u32));
744 }
745
746 pub const State = union(enum) {
747 /// End of connection signals the end of the stream.
748 none,
749 /// As a debugging utility, counts down to zero as bytes are written.
750 content_length: u64,
751 /// Each chunk is wrapped in a header and trailer.
752 chunked: Chunked,
753 /// Cleanly finished stream; connection can be reused.
754 end,
755
756 pub const Chunked = union(enum) {
757 /// Index to the start of the hex-encoded chunk length in the chunk
758 /// header within the buffer of `BodyWriter.http_protocol_output`.
759 /// Buffered chunk data starts here plus length of `chunk_header_template`.
760 offset: usize,
761 /// We are in the middle of a chunk and this is how many bytes are
762 /// left until the next header. This includes +2 for "\r"\n", and
763 /// is zero for the beginning of the stream.
764 chunk_len: usize,
765
766 pub const init: Chunked = .{ .chunk_len = 0 };
767 };
768 };
769
770 pub fn isEliding(w: *const BodyWriter) bool {
771 return w.writer.vtable.drain == elidingDrain;
772 }
773
774 /// Sends all buffered data across `BodyWriter.http_protocol_output`.
775 pub fn flush(w: *BodyWriter) Error!void {
776 const out = w.http_protocol_output;
777 switch (w.state) {
778 .end, .none, .content_length => return out.flush(),
779 .chunked => |*chunked| switch (chunked.*) {
780 .offset => |offset| {
781 const chunk_len = out.end - offset - chunk_header_template.len;
782 if (chunk_len > 0) {
783 writeHex(out.buffer[offset..][0..chunk_len_digits], chunk_len);
784 chunked.* = .{ .chunk_len = 2 };
785 } else {
786 out.end = offset;
787 chunked.* = .{ .chunk_len = 0 };
788 }
789 try out.flush();
790 },
791 .chunk_len => return out.flush(),
792 },
793 }
794 }
795
796 /// When using content-length, asserts that the amount of data sent matches
797 /// the value sent in the header, then flushes.
798 ///
799 /// When using transfer-encoding: chunked, writes the end-of-stream message
800 /// with empty trailers, then flushes the stream to the system. Asserts any
801 /// started chunk has been completely finished.
802 ///
803 /// Respects the value of `isEliding` to omit all data after the headers.
804 ///
805 /// See also:
806 /// * `endUnflushed`
807 /// * `endChunked`
808 pub fn end(w: *BodyWriter) Error!void {
809 try endUnflushed(w);
810 try w.http_protocol_output.flush();
811 }
812
813 /// When using content-length, asserts that the amount of data sent matches
814 /// the value sent in the header.
815 ///
816 /// Otherwise, transfer-encoding: chunked is being used, and it writes the
817 /// end-of-stream message with empty trailers.
818 ///
819 /// Respects the value of `isEliding` to omit all data after the headers.
820 ///
821 /// See also:
822 /// * `end`
823 /// * `endChunked`
824 pub fn endUnflushed(w: *BodyWriter) Error!void {
825 switch (w.state) {
826 .end => unreachable,
827 .content_length => |len| {
828 assert(len == 0); // Trips when end() called before all bytes written.
829 w.state = .end;
830 },
831 .none => {},
832 .chunked => return endChunkedUnflushed(w, .{}),
833 }
834 }
835
836 pub const EndChunkedOptions = struct {
837 trailers: []const Header = &.{},
838 };
839
840 /// Writes the end-of-stream message and any optional trailers, flushing
841 /// the underlying stream.
842 ///
843 /// Asserts that the BodyWriter is using transfer-encoding: chunked.
844 ///
845 /// Respects the value of `isEliding` to omit all data after the headers.
846 ///
847 /// See also:
848 /// * `endChunkedUnflushed`
849 /// * `end`
850 pub fn endChunked(w: *BodyWriter, options: EndChunkedOptions) Error!void {
851 try endChunkedUnflushed(w, options);
852 try w.http_protocol_output.flush();
853 }
854
855 /// Writes the end-of-stream message and any optional trailers.
856 ///
857 /// Does not flush.
858 ///
859 /// Asserts that the BodyWriter is using transfer-encoding: chunked.
860 ///
861 /// Respects the value of `isEliding` to omit all data after the headers.
862 ///
863 /// See also:
864 /// * `endChunked`
865 /// * `endUnflushed`
866 /// * `end`
867 pub fn endChunkedUnflushed(w: *BodyWriter, options: EndChunkedOptions) Error!void {
868 const chunked = &w.state.chunked;
869 if (w.isEliding()) {
870 w.state = .end;
871 return;
872 }
873 const bw = w.http_protocol_output;
874 switch (chunked.*) {
875 .offset => |offset| {
876 const chunk_len = bw.end - offset - chunk_header_template.len;
877 writeHex(bw.buffer[offset..][0..chunk_len_digits], chunk_len);
878 try bw.writeAll("\r\n");
879 },
880 .chunk_len => |chunk_len| switch (chunk_len) {
881 0 => {},
882 1 => try bw.writeByte('\n'),
883 2 => try bw.writeAll("\r\n"),
884 else => unreachable, // An earlier write call indicated more data would follow.
885 },
886 }
887 try bw.writeAll("0\r\n");
888 for (options.trailers) |trailer| {
889 try bw.writeAll(trailer.name);
890 try bw.writeAll(": ");
891 try bw.writeAll(trailer.value);
892 try bw.writeAll("\r\n");
893 }
894 try bw.writeAll("\r\n");
895 w.state = .end;
896 }
897
898 pub fn contentLengthDrain(w: *Writer, data: []const []const u8, splat: usize) Error!usize {
899 const bw: *BodyWriter = @alignCast(@fieldParentPtr("writer", w));
900 assert(!bw.isEliding());
901 const out = bw.http_protocol_output;
902 const n = try out.writeSplatHeader(w.buffered(), data, splat);
903 bw.state.content_length -= n;
904 return w.consume(n);
905 }
906
907 pub fn noneDrain(w: *Writer, data: []const []const u8, splat: usize) Error!usize {
908 const bw: *BodyWriter = @alignCast(@fieldParentPtr("writer", w));
909 assert(!bw.isEliding());
910 const out = bw.http_protocol_output;
911 const n = try out.writeSplatHeader(w.buffered(), data, splat);
912 return w.consume(n);
913 }
914
915 pub fn elidingDrain(w: *Writer, data: []const []const u8, splat: usize) Error!usize {
916 const bw: *BodyWriter = @alignCast(@fieldParentPtr("writer", w));
917 const slice = data[0 .. data.len - 1];
918 const pattern = data[slice.len];
919 var written: usize = pattern.len * splat;
920 for (slice) |bytes| written += bytes.len;
921 switch (bw.state) {
922 .content_length => |*len| len.* -= written + w.end,
923 else => {},
924 }
925 w.end = 0;
926 return written;
927 }
928
929 pub fn elidingSendFile(w: *Writer, file_reader: *File.Reader, limit: std.Io.Limit) Writer.FileError!usize {
930 const bw: *BodyWriter = @alignCast(@fieldParentPtr("writer", w));
931 if (File.Handle == void) return error.Unimplemented;
932 if (builtin.zig_backend == .stage2_aarch64) return error.Unimplemented;
933 switch (bw.state) {
934 .content_length => |*len| len.* -= w.end,
935 else => {},
936 }
937 w.end = 0;
938 if (limit == .nothing) return 0;
939 if (file_reader.getSize()) |size| {
940 const n = limit.minInt64(size - file_reader.pos);
941 if (n == 0) return error.EndOfStream;
942 file_reader.seekBy(@intCast(n)) catch return error.Unimplemented;
943 switch (bw.state) {
944 .content_length => |*len| len.* -= n,
945 else => {},
946 }
947 return n;
948 } else |_| {
949 // Error is observable on `file_reader` instance, and it is better to
950 // treat the file as a pipe.
951 return error.Unimplemented;
952 }
953 }
954
955 /// Returns `null` if size cannot be computed without making any syscalls.
956 pub fn noneSendFile(w: *Writer, file_reader: *File.Reader, limit: std.Io.Limit) Writer.FileError!usize {
957 const bw: *BodyWriter = @alignCast(@fieldParentPtr("writer", w));
958 assert(!bw.isEliding());
959 const out = bw.http_protocol_output;
960 const n = try out.sendFileHeader(w.buffered(), file_reader, limit);
961 return w.consume(n);
962 }
963
964 pub fn contentLengthSendFile(w: *Writer, file_reader: *File.Reader, limit: std.Io.Limit) Writer.FileError!usize {
965 const bw: *BodyWriter = @alignCast(@fieldParentPtr("writer", w));
966 assert(!bw.isEliding());
967 const out = bw.http_protocol_output;
968 const n = try out.sendFileHeader(w.buffered(), file_reader, limit);
969 bw.state.content_length -= n;
970 return w.consume(n);
971 }
972
973 pub fn chunkedSendFile(w: *Writer, file_reader: *File.Reader, limit: std.Io.Limit) Writer.FileError!usize {
974 const bw: *BodyWriter = @alignCast(@fieldParentPtr("writer", w));
975 assert(!bw.isEliding());
976 const data_len = Writer.countSendFileLowerBound(w.end, file_reader, limit) orelse {
977 // If the file size is unknown, we cannot lower to a `sendFile` since we would
978 // have to flush the chunk header before knowing the chunk length.
979 return error.Unimplemented;
980 };
981 const out = bw.http_protocol_output;
982 const chunked = &bw.state.chunked;
983 state: switch (chunked.*) {
984 .offset => |off| {
985 // TODO: is it better perf to read small files into the buffer?
986 const buffered_len = out.end - off - chunk_header_template.len;
987 const chunk_len = data_len + buffered_len;
988 writeHex(out.buffer[off..][0..chunk_len_digits], chunk_len);
989 const n = try out.sendFileHeader(w.buffered(), file_reader, limit);
990 chunked.* = .{ .chunk_len = data_len + 2 - n };
991 return w.consume(n);
992 },
993 .chunk_len => |chunk_len| l: switch (chunk_len) {
994 0 => {
995 const off = out.end;
996 const header_buf = try out.writableArray(chunk_header_template.len);
997 @memcpy(header_buf, chunk_header_template);
998 chunked.* = .{ .offset = off };
999 continue :state .{ .offset = off };
1000 },
1001 1 => {
1002 try out.writeByte('\n');
1003 chunked.chunk_len = 0;
1004 continue :l 0;
1005 },
1006 2 => {
1007 try out.writeByte('\r');
1008 chunked.chunk_len = 1;
1009 continue :l 1;
1010 },
1011 else => {
1012 const new_limit = limit.min(.limited(chunk_len - 2));
1013 const n = try out.sendFileHeader(w.buffered(), file_reader, new_limit);
1014 chunked.chunk_len = chunk_len - n;
1015 return w.consume(n);
1016 },
1017 },
1018 }
1019 }
1020
1021 pub fn chunkedDrain(w: *Writer, data: []const []const u8, splat: usize) Error!usize {
1022 const bw: *BodyWriter = @alignCast(@fieldParentPtr("writer", w));
1023 assert(!bw.isEliding());
1024 const out = bw.http_protocol_output;
1025 const data_len = w.end + Writer.countSplat(data, splat);
1026 const chunked = &bw.state.chunked;
1027 state: switch (chunked.*) {
1028 .offset => |offset| {
1029 if (out.unusedCapacityLen() >= data_len) {
1030 return w.consume(out.writeSplatHeader(w.buffered(), data, splat) catch unreachable);
1031 }
1032 const buffered_len = out.end - offset - chunk_header_template.len;
1033 const chunk_len = data_len + buffered_len;
1034 writeHex(out.buffer[offset..][0..chunk_len_digits], chunk_len);
1035 const n = try out.writeSplatHeader(w.buffered(), data, splat);
1036 chunked.* = .{ .chunk_len = data_len + 2 - n };
1037 return w.consume(n);
1038 },
1039 .chunk_len => |chunk_len| l: switch (chunk_len) {
1040 0 => {
1041 const offset = out.end;
1042 const header_buf = try out.writableArray(chunk_header_template.len);
1043 @memcpy(header_buf, chunk_header_template);
1044 chunked.* = .{ .offset = offset };
1045 continue :state .{ .offset = offset };
1046 },
1047 1 => {
1048 try out.writeByte('\n');
1049 chunked.chunk_len = 0;
1050 continue :l 0;
1051 },
1052 2 => {
1053 try out.writeByte('\r');
1054 chunked.chunk_len = 1;
1055 continue :l 1;
1056 },
1057 else => {
1058 const n = try out.writeSplatHeaderLimit(w.buffered(), data, splat, .limited(chunk_len - 2));
1059 chunked.chunk_len = chunk_len - n;
1060 return w.consume(n);
1061 },
1062 },
1063 }
1064 }
1065
1066 /// Writes an integer as base 16 to `buf`, right-aligned, assuming the
1067 /// buffer has already been filled with zeroes.
1068 fn writeHex(buf: []u8, x: usize) void {
1069 assert(std.mem.allEqual(u8, buf, '0'));
1070 const base = 16;
1071 var index: usize = buf.len;
1072 var a = x;
1073 while (a > 0) {
1074 const digit = a % base;
1075 index -= 1;
1076 buf[index] = std.fmt.digitToChar(@intCast(digit), .lower);
1077 a /= base;
1078 }
1079 }
1080};
1081
318test {1082test {
1083 _ = Server;
1084 _ = Status;
1085 _ = Method;
1086 _ = ChunkParser;
1087 _ = HeadParser;
1088
319 if (builtin.os.tag != .wasi) {1089 if (builtin.os.tag != .wasi) {
320 _ = Client;1090 _ = Client;
321 _ = Method;
322 _ = Server;
323 _ = Status;
324 _ = HeadParser;
325 _ = ChunkParser;
326 _ = WebSocket;
327 _ = @import("http/test.zig");1091 _ = @import("http/test.zig");
328 }1092 }
329}1093}
lib/std/http/ChunkParser.zig+3-3
...@@ -1,5 +1,8 @@...@@ -1,5 +1,8 @@
1//! Parser for transfer-encoding: chunked.1//! Parser for transfer-encoding: chunked.
22
3const ChunkParser = @This();
4const std = @import("std");
5
3state: State,6state: State,
4chunk_len: u64,7chunk_len: u64,
58
...@@ -97,9 +100,6 @@ pub fn feed(p: *ChunkParser, bytes: []const u8) usize {...@@ -97,9 +100,6 @@ pub fn feed(p: *ChunkParser, bytes: []const u8) usize {
97 return bytes.len;100 return bytes.len;
98}101}
99102
100const ChunkParser = @This();
101const std = @import("std");
102
103test feed {103test feed {
104 const testing = std.testing;104 const testing = std.testing;
105105
lib/std/http/Client.zig+1039-1011
...@@ -13,9 +13,10 @@ const net = std.net;...@@ -13,9 +13,10 @@ const net = std.net;
13const Uri = std.Uri;13const Uri = std.Uri;
14const Allocator = mem.Allocator;14const Allocator = mem.Allocator;
15const assert = std.debug.assert;15const assert = std.debug.assert;
16const Writer = std.io.Writer;
17const Reader = std.io.Reader;
1618
17const Client = @This();19const Client = @This();
18const proto = @import("protocol.zig");
1920
20pub const disable_tls = std.options.http_disable_tls;21pub const disable_tls = std.options.http_disable_tls;
2122
...@@ -24,6 +25,12 @@ allocator: Allocator,...@@ -24,6 +25,12 @@ allocator: Allocator,
2425
25ca_bundle: if (disable_tls) void else std.crypto.Certificate.Bundle = if (disable_tls) {} else .{},26ca_bundle: if (disable_tls) void else std.crypto.Certificate.Bundle = if (disable_tls) {} else .{},
26ca_bundle_mutex: std.Thread.Mutex = .{},27ca_bundle_mutex: std.Thread.Mutex = .{},
28/// Used both for the reader and writer buffers.
29tls_buffer_size: if (disable_tls) u0 else usize = if (disable_tls) 0 else std.crypto.tls.Client.min_buffer_len,
30/// If non-null, ssl secrets are logged to a stream. Creating such a stream
31/// allows other processes with access to that stream to decrypt all
32/// traffic over connections created with this `Client`.
33ssl_key_log: ?*std.crypto.tls.Client.SslKeyLog = null,
2734
28/// When this is `true`, the next time this client performs an HTTPS request,35/// When this is `true`, the next time this client performs an HTTPS request,
29/// it will first rescan the system for root certificates.36/// it will first rescan the system for root certificates.
...@@ -31,6 +38,13 @@ next_https_rescan_certs: bool = true,...@@ -31,6 +38,13 @@ next_https_rescan_certs: bool = true,
3138
32/// The pool of connections that can be reused (and currently in use).39/// The pool of connections that can be reused (and currently in use).
33connection_pool: ConnectionPool = .{},40connection_pool: ConnectionPool = .{},
41/// Each `Connection` allocates this amount for the reader buffer.
42///
43/// If the entire HTTP header cannot fit in this amount of bytes,
44/// `error.HttpHeadersOversize` will be returned from `Request.wait`.
45read_buffer_size: usize = 4096 + if (disable_tls) 0 else std.crypto.tls.Client.min_buffer_len,
46/// Each `Connection` allocates this amount for the writer buffer.
47write_buffer_size: usize = 1024,
3448
35/// If populated, all http traffic travels through this third party.49/// If populated, all http traffic travels through this third party.
36/// This field cannot be modified while the client has active connections.50/// This field cannot be modified while the client has active connections.
...@@ -41,7 +55,7 @@ http_proxy: ?*Proxy = null,...@@ -41,7 +55,7 @@ http_proxy: ?*Proxy = null,
41/// Pointer to externally-owned memory.55/// Pointer to externally-owned memory.
42https_proxy: ?*Proxy = null,56https_proxy: ?*Proxy = null,
4357
44/// A set of linked lists of connections that can be reused.58/// A Least-Recently-Used cache of open connections to be reused.
45pub const ConnectionPool = struct {59pub const ConnectionPool = struct {
46 mutex: std.Thread.Mutex = .{},60 mutex: std.Thread.Mutex = .{},
47 /// Open connections that are currently in use.61 /// Open connections that are currently in use.
...@@ -55,23 +69,25 @@ pub const ConnectionPool = struct {...@@ -55,23 +69,25 @@ pub const ConnectionPool = struct {
55 pub const Criteria = struct {69 pub const Criteria = struct {
56 host: []const u8,70 host: []const u8,
57 port: u16,71 port: u16,
58 protocol: Connection.Protocol,72 protocol: Protocol,
59 };73 };
6074
61 /// Finds and acquires a connection from the connection pool matching the criteria. This function is threadsafe.75 /// Finds and acquires a connection from the connection pool matching the criteria.
62 /// If no connection is found, null is returned.76 /// If no connection is found, null is returned.
77 ///
78 /// Threadsafe.
63 pub fn findConnection(pool: *ConnectionPool, criteria: Criteria) ?*Connection {79 pub fn findConnection(pool: *ConnectionPool, criteria: Criteria) ?*Connection {
64 pool.mutex.lock();80 pool.mutex.lock();
65 defer pool.mutex.unlock();81 defer pool.mutex.unlock();
6682
67 var next = pool.free.last;83 var next = pool.free.last;
68 while (next) |node| : (next = node.prev) {84 while (next) |node| : (next = node.prev) {
69 const connection: *Connection = @fieldParentPtr("pool_node", node);85 const connection: *Connection = @alignCast(@fieldParentPtr("pool_node", node));
70 if (connection.protocol != criteria.protocol) continue;86 if (connection.protocol != criteria.protocol) continue;
71 if (connection.port != criteria.port) continue;87 if (connection.port != criteria.port) continue;
7288
73 // Domain names are case-insensitive (RFC 5890, Section 2.3.2.4)89 // Domain names are case-insensitive (RFC 5890, Section 2.3.2.4)
74 if (!std.ascii.eqlIgnoreCase(connection.host, criteria.host)) continue;90 if (!std.ascii.eqlIgnoreCase(connection.host(), criteria.host)) continue;
7591
76 pool.acquireUnsafe(connection);92 pool.acquireUnsafe(connection);
77 return connection;93 return connection;
...@@ -96,28 +112,23 @@ pub const ConnectionPool = struct {...@@ -96,28 +112,23 @@ pub const ConnectionPool = struct {
96 return pool.acquireUnsafe(connection);112 return pool.acquireUnsafe(connection);
97 }113 }
98114
99 /// Tries to release a connection back to the connection pool. This function is threadsafe.115 /// Tries to release a connection back to the connection pool.
100 /// If the connection is marked as closing, it will be closed instead.116 /// If the connection is marked as closing, it will be closed instead.
101 ///117 ///
102 /// The allocator must be the owner of all nodes in this pool.118 /// Threadsafe.
103 /// The allocator must be the owner of all resources associated with the connection.119 pub fn release(pool: *ConnectionPool, connection: *Connection) void {
104 pub fn release(pool: *ConnectionPool, allocator: Allocator, connection: *Connection) void {
105 pool.mutex.lock();120 pool.mutex.lock();
106 defer pool.mutex.unlock();121 defer pool.mutex.unlock();
107122
108 pool.used.remove(&connection.pool_node);123 pool.used.remove(&connection.pool_node);
109124
110 if (connection.closing or pool.free_size == 0) {125 if (connection.closing or pool.free_size == 0) return connection.destroy();
111 connection.close(allocator);
112 return allocator.destroy(connection);
113 }
114126
115 if (pool.free_len >= pool.free_size) {127 if (pool.free_len >= pool.free_size) {
116 const popped: *Connection = @fieldParentPtr("pool_node", pool.free.popFirst().?);128 const popped: *Connection = @alignCast(@fieldParentPtr("pool_node", pool.free.popFirst().?));
117 pool.free_len -= 1;129 pool.free_len -= 1;
118130
119 popped.close(allocator);131 popped.destroy();
120 allocator.destroy(popped);
121 }132 }
122133
123 if (connection.proxied) {134 if (connection.proxied) {
...@@ -138,9 +149,11 @@ pub const ConnectionPool = struct {...@@ -138,9 +149,11 @@ pub const ConnectionPool = struct {
138 pool.used.append(&connection.pool_node);149 pool.used.append(&connection.pool_node);
139 }150 }
140151
141 /// Resizes the connection pool. This function is threadsafe.152 /// Resizes the connection pool.
142 ///153 ///
143 /// If the new size is smaller than the current size, then idle connections will be closed until the pool is the new size.154 /// If the new size is smaller than the current size, then idle connections will be closed until the pool is the new size.
155 ///
156 /// Threadsafe.
144 pub fn resize(pool: *ConnectionPool, allocator: Allocator, new_size: usize) void {157 pub fn resize(pool: *ConnectionPool, allocator: Allocator, new_size: usize) void {
145 pool.mutex.lock();158 pool.mutex.lock();
146 defer pool.mutex.unlock();159 defer pool.mutex.unlock();
...@@ -158,538 +171,612 @@ pub const ConnectionPool = struct {...@@ -158,538 +171,612 @@ pub const ConnectionPool = struct {
158 pool.free_size = new_size;171 pool.free_size = new_size;
159 }172 }
160173
161 /// Frees the connection pool and closes all connections within. This function is threadsafe.174 /// Frees the connection pool and closes all connections within.
162 ///175 ///
163 /// All future operations on the connection pool will deadlock.176 /// All future operations on the connection pool will deadlock.
164 pub fn deinit(pool: *ConnectionPool, allocator: Allocator) void {177 ///
178 /// Threadsafe.
179 pub fn deinit(pool: *ConnectionPool) void {
165 pool.mutex.lock();180 pool.mutex.lock();
166181
167 var next = pool.free.first;182 var next = pool.free.first;
168 while (next) |node| {183 while (next) |node| {
169 const connection: *Connection = @fieldParentPtr("pool_node", node);184 const connection: *Connection = @alignCast(@fieldParentPtr("pool_node", node));
170 next = node.next;185 next = node.next;
171 connection.close(allocator);186 connection.destroy();
172 allocator.destroy(connection);
173 }187 }
174188
175 next = pool.used.first;189 next = pool.used.first;
176 while (next) |node| {190 while (next) |node| {
177 const connection: *Connection = @fieldParentPtr("pool_node", node);191 const connection: *Connection = @alignCast(@fieldParentPtr("pool_node", node));
178 next = node.next;192 next = node.next;
179 connection.close(allocator);193 connection.destroy();
180 allocator.destroy(node);
181 }194 }
182195
183 pool.* = undefined;196 pool.* = undefined;
184 }197 }
185};198};
186199
187/// An interface to either a plain or TLS connection.200pub const Protocol = enum {
188pub const Connection = struct {201 plain,
189 stream: net.Stream,202 tls,
190 /// undefined unless protocol is tls.
191 tls_client: if (!disable_tls) *std.crypto.tls.Client else void,
192
193 /// Entry in `ConnectionPool.used` or `ConnectionPool.free`.
194 pool_node: std.DoublyLinkedList.Node,
195
196 /// The protocol that this connection is using.
197 protocol: Protocol,
198
199 /// The host that this connection is connected to.
200 host: []u8,
201203
202 /// The port that this connection is connected to.204 fn port(protocol: Protocol) u16 {
203 port: u16,205 return switch (protocol) {
204206 .plain => 80,
205 /// Whether this connection is proxied and is not directly connected.207 .tls => 443,
206 proxied: bool = false,
207
208 /// Whether this connection is closing when we're done with it.
209 closing: bool = false,
210
211 read_start: BufferSize = 0,
212 read_end: BufferSize = 0,
213 write_end: BufferSize = 0,
214 read_buf: [buffer_size]u8 = undefined,
215 write_buf: [buffer_size]u8 = undefined,
216
217 pub const buffer_size = std.crypto.tls.max_ciphertext_record_len;
218 const BufferSize = std.math.IntFittingRange(0, buffer_size);
219
220 pub const Protocol = enum { plain, tls };
221
222 pub fn readvDirectTls(conn: *Connection, buffers: []std.posix.iovec) ReadError!usize {
223 return conn.tls_client.readv(conn.stream, buffers) catch |err| {
224 // https://github.com/ziglang/zig/issues/2473
225 if (mem.startsWith(u8, @errorName(err), "TlsAlert")) return error.TlsAlert;
226
227 switch (err) {
228 error.TlsConnectionTruncated, error.TlsRecordOverflow, error.TlsDecodeError, error.TlsBadRecordMac, error.TlsBadLength, error.TlsIllegalParameter, error.TlsUnexpectedMessage => return error.TlsFailure,
229 error.ConnectionTimedOut => return error.ConnectionTimedOut,
230 error.ConnectionResetByPeer, error.BrokenPipe => return error.ConnectionResetByPeer,
231 else => return error.UnexpectedReadFailure,
232 }
233 };208 };
234 }209 }
235210
236 pub fn readvDirect(conn: *Connection, buffers: []std.posix.iovec) ReadError!usize {211 pub fn fromScheme(scheme: []const u8) ?Protocol {
237 if (conn.protocol == .tls) {212 const protocol_map = std.StaticStringMap(Protocol).initComptime(.{
238 if (disable_tls) unreachable;213 .{ "http", .plain },
239214 .{ "ws", .plain },
240 return conn.readvDirectTls(buffers);215 .{ "https", .tls },
241 }216 .{ "wss", .tls },
242217 });
243 return conn.stream.readv(buffers) catch |err| switch (err) {218 return protocol_map.get(scheme);
244 error.ConnectionTimedOut => return error.ConnectionTimedOut,
245 error.ConnectionResetByPeer, error.BrokenPipe => return error.ConnectionResetByPeer,
246 else => return error.UnexpectedReadFailure,
247 };
248 }
249
250 /// Refills the read buffer with data from the connection.
251 pub fn fill(conn: *Connection) ReadError!void {
252 if (conn.read_end != conn.read_start) return;
253
254 var iovecs = [1]std.posix.iovec{
255 .{ .base = &conn.read_buf, .len = conn.read_buf.len },
256 };
257 const nread = try conn.readvDirect(&iovecs);
258 if (nread == 0) return error.EndOfStream;
259 conn.read_start = 0;
260 conn.read_end = @intCast(nread);
261 }219 }
262220
263 /// Returns the current slice of buffered data.221 pub fn fromUri(uri: Uri) ?Protocol {
264 pub fn peek(conn: *Connection) []const u8 {222 return fromScheme(uri.scheme);
265 return conn.read_buf[conn.read_start..conn.read_end];
266 }223 }
224};
267225
268 /// Discards the given number of bytes from the read buffer.226pub const Connection = struct {
269 pub fn drop(conn: *Connection, num: BufferSize) void {227 client: *Client,
270 conn.read_start += num;228 stream_writer: net.Stream.Writer,
271 }229 stream_reader: net.Stream.Reader,
230 /// Entry in `ConnectionPool.used` or `ConnectionPool.free`.
231 pool_node: std.DoublyLinkedList.Node,
232 port: u16,
233 host_len: u8,
234 proxied: bool,
235 closing: bool,
236 protocol: Protocol,
272237
273 /// Reads data from the connection into the given buffer.238 const Plain = struct {
274 pub fn read(conn: *Connection, buffer: []u8) ReadError!usize {239 connection: Connection,
275 const available_read = conn.read_end - conn.read_start;240
276 const available_buffer = buffer.len;241 fn create(
242 client: *Client,
243 remote_host: []const u8,
244 port: u16,
245 stream: net.Stream,
246 ) error{OutOfMemory}!*Plain {
247 const gpa = client.allocator;
248 const alloc_len = allocLen(client, remote_host.len);
249 const base = try gpa.alignedAlloc(u8, .of(Plain), alloc_len);
250 errdefer gpa.free(base);
251 const host_buffer = base[@sizeOf(Plain)..][0..remote_host.len];
252 const socket_read_buffer = host_buffer.ptr[host_buffer.len..][0..client.read_buffer_size];
253 const socket_write_buffer = socket_read_buffer.ptr[socket_read_buffer.len..][0..client.write_buffer_size];
254 assert(base.ptr + alloc_len == socket_write_buffer.ptr + socket_write_buffer.len);
255 @memcpy(host_buffer, remote_host);
256 const plain: *Plain = @ptrCast(base);
257 plain.* = .{
258 .connection = .{
259 .client = client,
260 .stream_writer = stream.writer(socket_write_buffer),
261 .stream_reader = stream.reader(socket_read_buffer),
262 .pool_node = .{},
263 .port = port,
264 .host_len = @intCast(remote_host.len),
265 .proxied = false,
266 .closing = false,
267 .protocol = .plain,
268 },
269 };
270 return plain;
271 }
277272
278 if (available_read > available_buffer) { // partially read buffered data273 fn destroy(plain: *Plain) void {
279 @memcpy(buffer[0..available_buffer], conn.read_buf[conn.read_start..conn.read_end][0..available_buffer]);274 const c = &plain.connection;
280 conn.read_start += @intCast(available_buffer);275 const gpa = c.client.allocator;
276 const base: [*]align(@alignOf(Plain)) u8 = @ptrCast(plain);
277 gpa.free(base[0..allocLen(c.client, c.host_len)]);
278 }
281279
282 return available_buffer;280 fn allocLen(client: *Client, host_len: usize) usize {
283 } else if (available_read > 0) { // fully read buffered data281 return @sizeOf(Plain) + host_len + client.read_buffer_size + client.write_buffer_size;
284 @memcpy(buffer[0..available_read], conn.read_buf[conn.read_start..conn.read_end]);282 }
285 conn.read_start += available_read;
286283
287 return available_read;284 fn host(plain: *Plain) []u8 {
285 const base: [*]u8 = @ptrCast(plain);
286 return base[@sizeOf(Plain)..][0..plain.connection.host_len];
288 }287 }
288 };
289289
290 var iovecs = [2]std.posix.iovec{290 const Tls = struct {
291 .{ .base = buffer.ptr, .len = buffer.len },291 client: std.crypto.tls.Client,
292 .{ .base = &conn.read_buf, .len = conn.read_buf.len },292 connection: Connection,
293 };293
294 const nread = try conn.readvDirect(&iovecs);294 fn create(
295 client: *Client,
296 remote_host: []const u8,
297 port: u16,
298 stream: net.Stream,
299 ) error{ OutOfMemory, TlsInitializationFailed }!*Tls {
300 const gpa = client.allocator;
301 const alloc_len = allocLen(client, remote_host.len);
302 const base = try gpa.alignedAlloc(u8, .of(Tls), alloc_len);
303 errdefer gpa.free(base);
304 const host_buffer = base[@sizeOf(Tls)..][0..remote_host.len];
305 const tls_read_buffer = host_buffer.ptr[host_buffer.len..][0..client.tls_buffer_size];
306 const tls_write_buffer = tls_read_buffer.ptr[tls_read_buffer.len..][0..client.tls_buffer_size];
307 const write_buffer = tls_write_buffer.ptr[tls_write_buffer.len..][0..client.write_buffer_size];
308 const read_buffer = write_buffer.ptr[write_buffer.len..][0..client.read_buffer_size];
309 assert(base.ptr + alloc_len == read_buffer.ptr + read_buffer.len);
310 @memcpy(host_buffer, remote_host);
311 const tls: *Tls = @ptrCast(base);
312 tls.* = .{
313 .connection = .{
314 .client = client,
315 .stream_writer = stream.writer(tls_write_buffer),
316 .stream_reader = stream.reader(tls_read_buffer),
317 .pool_node = .{},
318 .port = port,
319 .host_len = @intCast(remote_host.len),
320 .proxied = false,
321 .closing = false,
322 .protocol = .tls,
323 },
324 // TODO data race here on ca_bundle if the user sets next_https_rescan_certs to true
325 .client = std.crypto.tls.Client.init(
326 tls.connection.stream_reader.interface(),
327 &tls.connection.stream_writer.interface,
328 .{
329 .host = .{ .explicit = remote_host },
330 .ca = .{ .bundle = client.ca_bundle },
331 .ssl_key_log = client.ssl_key_log,
332 .read_buffer = read_buffer,
333 .write_buffer = write_buffer,
334 // This is appropriate for HTTPS because the HTTP headers contain
335 // the content length which is used to detect truncation attacks.
336 .allow_truncation_attacks = true,
337 },
338 ) catch return error.TlsInitializationFailed,
339 };
340 return tls;
341 }
295342
296 if (nread > buffer.len) {343 fn destroy(tls: *Tls) void {
297 conn.read_start = 0;344 const c = &tls.connection;
298 conn.read_end = @intCast(nread - buffer.len);345 const gpa = c.client.allocator;
299 return buffer.len;346 const base: [*]align(@alignOf(Tls)) u8 = @ptrCast(tls);
347 gpa.free(base[0..allocLen(c.client, c.host_len)]);
300 }348 }
301349
302 return nread;350 fn allocLen(client: *Client, host_len: usize) usize {
303 }351 return @sizeOf(Tls) + host_len + client.tls_buffer_size + client.tls_buffer_size +
352 client.write_buffer_size + client.read_buffer_size;
353 }
304354
305 pub const ReadError = error{355 fn host(tls: *Tls) []u8 {
306 TlsFailure,356 const base: [*]u8 = @ptrCast(tls);
307 TlsAlert,357 return base[@sizeOf(Tls)..][0..tls.connection.host_len];
308 ConnectionTimedOut,358 }
309 ConnectionResetByPeer,
310 UnexpectedReadFailure,
311 EndOfStream,
312 };359 };
313360
314 pub const Reader = std.io.GenericReader(*Connection, ReadError, read);361 pub const ReadError = std.crypto.tls.Client.ReadError || std.net.Stream.ReadError;
315
316 pub fn reader(conn: *Connection) Reader {
317 return Reader{ .context = conn };
318 }
319362
320 pub fn writeAllDirectTls(conn: *Connection, buffer: []const u8) WriteError!void {363 pub fn getReadError(c: *const Connection) ?ReadError {
321 return conn.tls_client.writeAll(conn.stream, buffer) catch |err| switch (err) {364 return switch (c.protocol) {
322 error.BrokenPipe, error.ConnectionResetByPeer => return error.ConnectionResetByPeer,365 .tls => {
323 else => return error.UnexpectedWriteFailure,366 if (disable_tls) unreachable;
367 const tls: *const Tls = @alignCast(@fieldParentPtr("connection", c));
368 return tls.client.read_err orelse c.stream_reader.getError();
369 },
370 .plain => {
371 return c.stream_reader.getError();
372 },
324 };373 };
325 }374 }
326375
327 pub fn writeAllDirect(conn: *Connection, buffer: []const u8) WriteError!void {376 fn getStream(c: *Connection) net.Stream {
328 if (conn.protocol == .tls) {377 return c.stream_reader.getStream();
329 if (disable_tls) unreachable;378 }
330
331 return conn.writeAllDirectTls(buffer);
332 }
333379
334 return conn.stream.writeAll(buffer) catch |err| switch (err) {380 fn host(c: *Connection) []u8 {
335 error.BrokenPipe, error.ConnectionResetByPeer => return error.ConnectionResetByPeer,381 return switch (c.protocol) {
336 else => return error.UnexpectedWriteFailure,382 .tls => {
383 if (disable_tls) unreachable;
384 const tls: *Tls = @alignCast(@fieldParentPtr("connection", c));
385 return tls.host();
386 },
387 .plain => {
388 const plain: *Plain = @alignCast(@fieldParentPtr("connection", c));
389 return plain.host();
390 },
337 };391 };
338 }392 }
339393
340 /// Writes the given buffer to the connection.394 /// If this is called without calling `flush` or `end`, data will be
341 pub fn write(conn: *Connection, buffer: []const u8) WriteError!usize {395 /// dropped unsent.
342 if (conn.write_buf.len - conn.write_end < buffer.len) {396 pub fn destroy(c: *Connection) void {
343 try conn.flush();397 c.getStream().close();
344398 switch (c.protocol) {
345 if (buffer.len > conn.write_buf.len) {399 .tls => {
346 try conn.writeAllDirect(buffer);400 if (disable_tls) unreachable;
347 return buffer.len;401 const tls: *Tls = @alignCast(@fieldParentPtr("connection", c));
348 }402 tls.destroy();
403 },
404 .plain => {
405 const plain: *Plain = @alignCast(@fieldParentPtr("connection", c));
406 plain.destroy();
407 },
349 }408 }
350
351 @memcpy(conn.write_buf[conn.write_end..][0..buffer.len], buffer);
352 conn.write_end += @intCast(buffer.len);
353
354 return buffer.len;
355 }409 }
356410
357 /// Returns a buffer to be filled with exactly len bytes to write to the connection.411 /// HTTP protocol from client to server.
358 pub fn allocWriteBuffer(conn: *Connection, len: BufferSize) WriteError![]u8 {412 /// This either goes directly to `stream_writer`, or to a TLS client.
359 if (conn.write_buf.len - conn.write_end < len) try conn.flush();413 pub fn writer(c: *Connection) *Writer {
360 defer conn.write_end += len;414 return switch (c.protocol) {
361 return conn.write_buf[conn.write_end..][0..len];415 .tls => {
416 if (disable_tls) unreachable;
417 const tls: *Tls = @alignCast(@fieldParentPtr("connection", c));
418 return &tls.client.writer;
419 },
420 .plain => &c.stream_writer.interface,
421 };
362 }422 }
363423
364 /// Flushes the write buffer to the connection.424 /// HTTP protocol from server to client.
365 pub fn flush(conn: *Connection) WriteError!void {425 /// This either comes directly from `stream_reader`, or from a TLS client.
366 if (conn.write_end == 0) return;426 pub fn reader(c: *Connection) *Reader {
367427 return switch (c.protocol) {
368 try conn.writeAllDirect(conn.write_buf[0..conn.write_end]);428 .tls => {
369 conn.write_end = 0;429 if (disable_tls) unreachable;
430 const tls: *Tls = @alignCast(@fieldParentPtr("connection", c));
431 return &tls.client.reader;
432 },
433 .plain => c.stream_reader.interface(),
434 };
370 }435 }
371436
372 pub const WriteError = error{437 pub fn flush(c: *Connection) Writer.Error!void {
373 ConnectionResetByPeer,438 if (c.protocol == .tls) {
374 UnexpectedWriteFailure,439 if (disable_tls) unreachable;
375 };440 const tls: *Tls = @alignCast(@fieldParentPtr("connection", c));
376441 try tls.client.writer.flush();
377 pub const Writer = std.io.GenericWriter(*Connection, WriteError, write);442 }
378443 try c.stream_writer.interface.flush();
379 pub fn writer(conn: *Connection) Writer {
380 return Writer{ .context = conn };
381 }444 }
382445
383 /// Closes the connection.446 /// If the connection is a TLS connection, sends the close_notify alert.
384 pub fn close(conn: *Connection, allocator: Allocator) void {447 ///
385 if (conn.protocol == .tls) {448 /// Flushes all buffers.
449 pub fn end(c: *Connection) Writer.Error!void {
450 if (c.protocol == .tls) {
386 if (disable_tls) unreachable;451 if (disable_tls) unreachable;
387452 const tls: *Tls = @alignCast(@fieldParentPtr("connection", c));
388 // try to cleanly close the TLS connection, for any server that cares.453 try tls.client.end();
389 _ = conn.tls_client.writeEnd(conn.stream, "", true) catch {};
390 if (conn.tls_client.ssl_key_log) |key_log| key_log.file.close();
391 allocator.destroy(conn.tls_client);
392 }454 }
393455 try c.stream_writer.interface.flush();
394 conn.stream.close();
395 allocator.free(conn.host);
396 }456 }
397};457};
398458
399/// The mode of transport for requests.
400pub const RequestTransfer = union(enum) {
401 content_length: u64,
402 chunked: void,
403 none: void,
404};
405
406/// The decompressor for response messages.
407pub const Compression = union(enum) {
408 //deflate: std.compress.flate.Decompress,
409 //gzip: std.compress.flate.Decompress,
410 // https://github.com/ziglang/zig/issues/18937
411 //zstd: ZstdDecompressor,
412 none: void,
413};
414
415/// A HTTP response originating from a server.
416pub const Response = struct {459pub const Response = struct {
417 version: http.Version,460 request: *Request,
418 status: http.Status,461 /// Pointers in this struct are invalidated when the response body stream
419 reason: []const u8,462 /// is initialized.
463 head: Head,
464
465 pub const Head = struct {
466 bytes: []const u8,
467 version: http.Version,
468 status: http.Status,
469 reason: []const u8,
470 location: ?[]const u8 = null,
471 content_type: ?[]const u8 = null,
472 content_disposition: ?[]const u8 = null,
473
474 keep_alive: bool,
475
476 /// If present, the number of bytes in the response body.
477 content_length: ?u64 = null,
478
479 transfer_encoding: http.TransferEncoding = .none,
480 content_encoding: http.ContentEncoding = .identity,
481
482 pub const ParseError = error{
483 HttpConnectionHeaderUnsupported,
484 HttpContentEncodingUnsupported,
485 HttpHeaderContinuationsUnsupported,
486 HttpHeadersInvalid,
487 HttpTransferEncodingUnsupported,
488 InvalidContentLength,
489 };
420490
421 /// Points into the user-provided `server_header_buffer`.491 pub fn parse(bytes: []const u8) ParseError!Head {
422 location: ?[]const u8 = null,492 var res: Head = .{
423 /// Points into the user-provided `server_header_buffer`.493 .bytes = bytes,
424 content_type: ?[]const u8 = null,494 .status = undefined,
425 /// Points into the user-provided `server_header_buffer`.495 .reason = undefined,
426 content_disposition: ?[]const u8 = null,496 .version = undefined,
497 .keep_alive = false,
498 };
499 var it = mem.splitSequence(u8, bytes, "\r\n");
427500
428 keep_alive: bool,501 const first_line = it.first();
502 if (first_line.len < 12) return error.HttpHeadersInvalid;
429503
430 /// If present, the number of bytes in the response body.504 const version: http.Version = switch (int64(first_line[0..8])) {
431 content_length: ?u64 = null,505 int64("HTTP/1.0") => .@"HTTP/1.0",
506 int64("HTTP/1.1") => .@"HTTP/1.1",
507 else => return error.HttpHeadersInvalid,
508 };
509 if (first_line[8] != ' ') return error.HttpHeadersInvalid;
510 const status: http.Status = @enumFromInt(parseInt3(first_line[9..12]));
511 const reason = mem.trimLeft(u8, first_line[12..], " ");
512
513 res.version = version;
514 res.status = status;
515 res.reason = reason;
516 res.keep_alive = switch (version) {
517 .@"HTTP/1.0" => false,
518 .@"HTTP/1.1" => true,
519 };
432520
433 /// If present, the transfer encoding of the response body, otherwise none.521 while (it.next()) |line| {
434 transfer_encoding: http.TransferEncoding = .none,522 if (line.len == 0) return res;
523 switch (line[0]) {
524 ' ', '\t' => return error.HttpHeaderContinuationsUnsupported,
525 else => {},
526 }
435527
436 /// If present, the compression of the response body, otherwise identity (no compression).528 var line_it = mem.splitScalar(u8, line, ':');
437 transfer_compression: http.ContentEncoding = .identity,529 const header_name = line_it.next().?;
530 const header_value = mem.trim(u8, line_it.rest(), " \t");
531 if (header_name.len == 0) return error.HttpHeadersInvalid;
532
533 if (std.ascii.eqlIgnoreCase(header_name, "connection")) {
534 res.keep_alive = !std.ascii.eqlIgnoreCase(header_value, "close");
535 } else if (std.ascii.eqlIgnoreCase(header_name, "content-type")) {
536 res.content_type = header_value;
537 } else if (std.ascii.eqlIgnoreCase(header_name, "location")) {
538 res.location = header_value;
539 } else if (std.ascii.eqlIgnoreCase(header_name, "content-disposition")) {
540 res.content_disposition = header_value;
541 } else if (std.ascii.eqlIgnoreCase(header_name, "transfer-encoding")) {
542 // Transfer-Encoding: second, first
543 // Transfer-Encoding: deflate, chunked
544 var iter = mem.splitBackwardsScalar(u8, header_value, ',');
545
546 const first = iter.first();
547 const trimmed_first = mem.trim(u8, first, " ");
548
549 var next: ?[]const u8 = first;
550 if (std.meta.stringToEnum(http.TransferEncoding, trimmed_first)) |transfer| {
551 if (res.transfer_encoding != .none) return error.HttpHeadersInvalid; // we already have a transfer encoding
552 res.transfer_encoding = transfer;
553
554 next = iter.next();
555 }
438556
439 parser: proto.HeadersParser,557 if (next) |second| {
440 compression: Compression = .none,558 const trimmed_second = mem.trim(u8, second, " ");
441559
442 /// Whether the response body should be skipped. Any data read from the560 if (http.ContentEncoding.fromString(trimmed_second)) |transfer| {
443 /// response body will be discarded.561 if (res.content_encoding != .identity) return error.HttpHeadersInvalid; // double compression is not supported
444 skip: bool = false,562 res.content_encoding = transfer;
563 } else {
564 return error.HttpTransferEncodingUnsupported;
565 }
566 }
445567
446 pub const ParseError = error{568 if (iter.next()) |_| return error.HttpTransferEncodingUnsupported;
447 HttpHeadersInvalid,569 } else if (std.ascii.eqlIgnoreCase(header_name, "content-length")) {
448 HttpHeaderContinuationsUnsupported,570 const content_length = std.fmt.parseInt(u64, header_value, 10) catch return error.InvalidContentLength;
449 HttpTransferEncodingUnsupported,
450 HttpConnectionHeaderUnsupported,
451 InvalidContentLength,
452 CompressionUnsupported,
453 };
454571
455 pub fn parse(res: *Response, bytes: []const u8) ParseError!void {572 if (res.content_length != null and res.content_length != content_length) return error.HttpHeadersInvalid;
456 var it = mem.splitSequence(u8, bytes, "\r\n");
457573
458 const first_line = it.next().?;574 res.content_length = content_length;
459 if (first_line.len < 12) {575 } else if (std.ascii.eqlIgnoreCase(header_name, "content-encoding")) {
460 return error.HttpHeadersInvalid;576 if (res.content_encoding != .identity) return error.HttpHeadersInvalid;
461 }
462577
463 const version: http.Version = switch (int64(first_line[0..8])) {578 const trimmed = mem.trim(u8, header_value, " ");
464 int64("HTTP/1.0") => .@"HTTP/1.0",
465 int64("HTTP/1.1") => .@"HTTP/1.1",
466 else => return error.HttpHeadersInvalid,
467 };
468 if (first_line[8] != ' ') return error.HttpHeadersInvalid;
469 const status: http.Status = @enumFromInt(parseInt3(first_line[9..12]));
470 const reason = mem.trimStart(u8, first_line[12..], " ");
471
472 res.version = version;
473 res.status = status;
474 res.reason = reason;
475 res.keep_alive = switch (version) {
476 .@"HTTP/1.0" => false,
477 .@"HTTP/1.1" => true,
478 };
479579
480 while (it.next()) |line| {580 if (http.ContentEncoding.fromString(trimmed)) |ce| {
481 if (line.len == 0) return;581 res.content_encoding = ce;
482 switch (line[0]) {
483 ' ', '\t' => return error.HttpHeaderContinuationsUnsupported,
484 else => {},
485 }
486
487 var line_it = mem.splitScalar(u8, line, ':');
488 const header_name = line_it.next().?;
489 const header_value = mem.trim(u8, line_it.rest(), " \t");
490 if (header_name.len == 0) return error.HttpHeadersInvalid;
491
492 if (std.ascii.eqlIgnoreCase(header_name, "connection")) {
493 res.keep_alive = !std.ascii.eqlIgnoreCase(header_value, "close");
494 } else if (std.ascii.eqlIgnoreCase(header_name, "content-type")) {
495 res.content_type = header_value;
496 } else if (std.ascii.eqlIgnoreCase(header_name, "location")) {
497 res.location = header_value;
498 } else if (std.ascii.eqlIgnoreCase(header_name, "content-disposition")) {
499 res.content_disposition = header_value;
500 } else if (std.ascii.eqlIgnoreCase(header_name, "transfer-encoding")) {
501 // Transfer-Encoding: second, first
502 // Transfer-Encoding: deflate, chunked
503 var iter = mem.splitBackwardsScalar(u8, header_value, ',');
504
505 const first = iter.first();
506 const trimmed_first = mem.trim(u8, first, " ");
507
508 var next: ?[]const u8 = first;
509 if (std.meta.stringToEnum(http.TransferEncoding, trimmed_first)) |transfer| {
510 if (res.transfer_encoding != .none) return error.HttpHeadersInvalid; // we already have a transfer encoding
511 res.transfer_encoding = transfer;
512
513 next = iter.next();
514 }
515
516 if (next) |second| {
517 const trimmed_second = mem.trim(u8, second, " ");
518
519 if (std.meta.stringToEnum(http.ContentEncoding, trimmed_second)) |transfer| {
520 if (res.transfer_compression != .identity) return error.HttpHeadersInvalid; // double compression is not supported
521 res.transfer_compression = transfer;
522 } else {582 } else {
523 return error.HttpTransferEncodingUnsupported;583 return error.HttpContentEncodingUnsupported;
524 }584 }
525 }585 }
526
527 if (iter.next()) |_| return error.HttpTransferEncodingUnsupported;
528 } else if (std.ascii.eqlIgnoreCase(header_name, "content-length")) {
529 const content_length = std.fmt.parseInt(u64, header_value, 10) catch return error.InvalidContentLength;
530
531 if (res.content_length != null and res.content_length != content_length) return error.HttpHeadersInvalid;
532
533 res.content_length = content_length;
534 } else if (std.ascii.eqlIgnoreCase(header_name, "content-encoding")) {
535 if (res.transfer_compression != .identity) return error.HttpHeadersInvalid;
536
537 const trimmed = mem.trim(u8, header_value, " ");
538
539 if (std.meta.stringToEnum(http.ContentEncoding, trimmed)) |ce| {
540 res.transfer_compression = ce;
541 } else {
542 return error.HttpTransferEncodingUnsupported;
543 }
544 }586 }
587 return error.HttpHeadersInvalid; // missing empty line
545 }588 }
546 return error.HttpHeadersInvalid; // missing empty line
547 }
548589
549 test parse {590 test parse {
550 const response_bytes = "HTTP/1.1 200 OK\r\n" ++591 const response_bytes = "HTTP/1.1 200 OK\r\n" ++
551 "LOcation:url\r\n" ++592 "LOcation:url\r\n" ++
552 "content-tYpe: text/plain\r\n" ++593 "content-tYpe: text/plain\r\n" ++
553 "content-disposition:attachment; filename=example.txt \r\n" ++594 "content-disposition:attachment; filename=example.txt \r\n" ++
554 "content-Length:10\r\n" ++595 "content-Length:10\r\n" ++
555 "TRansfer-encoding:\tdeflate, chunked \r\n" ++596 "TRansfer-encoding:\tdeflate, chunked \r\n" ++
556 "connectioN:\t keep-alive \r\n\r\n";597 "connectioN:\t keep-alive \r\n\r\n";
557598
558 var header_buffer: [1024]u8 = undefined;599 const head = try Head.parse(response_bytes);
559 var res = Response{600
560 .status = undefined,601 try testing.expectEqual(.@"HTTP/1.1", head.version);
561 .reason = undefined,602 try testing.expectEqualStrings("OK", head.reason);
562 .version = undefined,603 try testing.expectEqual(.ok, head.status);
563 .keep_alive = false,604
564 .parser = .init(&header_buffer),605 try testing.expectEqualStrings("url", head.location.?);
565 };606 try testing.expectEqualStrings("text/plain", head.content_type.?);
607 try testing.expectEqualStrings("attachment; filename=example.txt", head.content_disposition.?);
608
609 try testing.expectEqual(true, head.keep_alive);
610 try testing.expectEqual(10, head.content_length.?);
611 try testing.expectEqual(.chunked, head.transfer_encoding);
612 try testing.expectEqual(.deflate, head.content_encoding);
613 }
566614
567 @memcpy(header_buffer[0..response_bytes.len], response_bytes);615 pub fn iterateHeaders(h: Head) http.HeaderIterator {
568 res.parser.header_bytes_len = response_bytes.len;616 return .init(h.bytes);
617 }
569618
570 try res.parse(response_bytes);619 test iterateHeaders {
620 const response_bytes = "HTTP/1.1 200 OK\r\n" ++
621 "LOcation:url\r\n" ++
622 "content-tYpe: text/plain\r\n" ++
623 "content-disposition:attachment; filename=example.txt \r\n" ++
624 "content-Length:10\r\n" ++
625 "TRansfer-encoding:\tdeflate, chunked \r\n" ++
626 "connectioN:\t keep-alive \r\n\r\n";
627
628 const head = try Head.parse(response_bytes);
629 var it = head.iterateHeaders();
630 {
631 const header = it.next().?;
632 try testing.expectEqualStrings("LOcation", header.name);
633 try testing.expectEqualStrings("url", header.value);
634 try testing.expect(!it.is_trailer);
635 }
636 {
637 const header = it.next().?;
638 try testing.expectEqualStrings("content-tYpe", header.name);
639 try testing.expectEqualStrings("text/plain", header.value);
640 try testing.expect(!it.is_trailer);
641 }
642 {
643 const header = it.next().?;
644 try testing.expectEqualStrings("content-disposition", header.name);
645 try testing.expectEqualStrings("attachment; filename=example.txt", header.value);
646 try testing.expect(!it.is_trailer);
647 }
648 {
649 const header = it.next().?;
650 try testing.expectEqualStrings("content-Length", header.name);
651 try testing.expectEqualStrings("10", header.value);
652 try testing.expect(!it.is_trailer);
653 }
654 {
655 const header = it.next().?;
656 try testing.expectEqualStrings("TRansfer-encoding", header.name);
657 try testing.expectEqualStrings("deflate, chunked", header.value);
658 try testing.expect(!it.is_trailer);
659 }
660 {
661 const header = it.next().?;
662 try testing.expectEqualStrings("connectioN", header.name);
663 try testing.expectEqualStrings("keep-alive", header.value);
664 try testing.expect(!it.is_trailer);
665 }
666 try testing.expectEqual(null, it.next());
667 }
571668
572 try testing.expectEqual(.@"HTTP/1.1", res.version);669 inline fn int64(array: *const [8]u8) u64 {
573 try testing.expectEqualStrings("OK", res.reason);670 return @bitCast(array.*);
574 try testing.expectEqual(.ok, res.status);671 }
575672
576 try testing.expectEqualStrings("url", res.location.?);673 fn parseInt3(text: *const [3]u8) u10 {
577 try testing.expectEqualStrings("text/plain", res.content_type.?);674 const nnn: @Vector(3, u8) = text.*;
578 try testing.expectEqualStrings("attachment; filename=example.txt", res.content_disposition.?);675 const zero: @Vector(3, u8) = .{ '0', '0', '0' };
676 const mmm: @Vector(3, u10) = .{ 100, 10, 1 };
677 return @reduce(.Add, (nnn -% zero) *% mmm);
678 }
579679
580 try testing.expectEqual(true, res.keep_alive);680 test parseInt3 {
581 try testing.expectEqual(10, res.content_length.?);681 const expectEqual = testing.expectEqual;
582 try testing.expectEqual(.chunked, res.transfer_encoding);682 try expectEqual(@as(u10, 0), parseInt3("000"));
583 try testing.expectEqual(.deflate, res.transfer_compression);683 try expectEqual(@as(u10, 418), parseInt3("418"));
584 }684 try expectEqual(@as(u10, 999), parseInt3("999"));
685 }
585686
586 inline fn int64(array: *const [8]u8) u64 {687 /// Help the programmer avoid bugs by calling this when the string
587 return @bitCast(array.*);688 /// memory of `Head` becomes invalidated.
588 }689 fn invalidateStrings(h: *Head) void {
690 h.bytes = undefined;
691 h.reason = undefined;
692 if (h.location) |*s| s.* = undefined;
693 if (h.content_type) |*s| s.* = undefined;
694 if (h.content_disposition) |*s| s.* = undefined;
695 }
696 };
589697
590 fn parseInt3(text: *const [3]u8) u10 {698 /// If compressed body has been negotiated this will return compressed bytes.
591 const nnn: @Vector(3, u8) = text.*;699 ///
592 const zero: @Vector(3, u8) = .{ '0', '0', '0' };700 /// If the returned `Reader` returns `error.ReadFailed` the error is
593 const mmm: @Vector(3, u10) = .{ 100, 10, 1 };701 /// available via `bodyErr`.
594 return @reduce(.Add, (nnn -% zero) *% mmm);702 ///
703 /// Asserts that this function is only called once.
704 ///
705 /// See also:
706 /// * `readerDecompressing`
707 pub fn reader(response: *Response, buffer: []u8) *Reader {
708 response.head.invalidateStrings();
709 const req = response.request;
710 if (!req.method.responseHasBody()) return .ending;
711 const head = &response.head;
712 return req.reader.bodyReader(buffer, head.transfer_encoding, head.content_length);
595 }713 }
596714
597 test parseInt3 {715 /// If compressed body has been negotiated this will return decompressed bytes.
598 const expectEqual = testing.expectEqual;716 ///
599 try expectEqual(@as(u10, 0), parseInt3("000"));717 /// If the returned `Reader` returns `error.ReadFailed` the error is
600 try expectEqual(@as(u10, 418), parseInt3("418"));718 /// available via `bodyErr`.
601 try expectEqual(@as(u10, 999), parseInt3("999"));719 ///
720 /// Asserts that this function is only called once.
721 ///
722 /// See also:
723 /// * `reader`
724 pub fn readerDecompressing(
725 response: *Response,
726 decompressor: *http.Decompressor,
727 decompression_buffer: []u8,
728 ) *Reader {
729 response.head.invalidateStrings();
730 const head = &response.head;
731 return response.request.reader.bodyReaderDecompressing(
732 head.transfer_encoding,
733 head.content_length,
734 head.content_encoding,
735 decompressor,
736 decompression_buffer,
737 );
602 }738 }
603739
604 pub fn iterateHeaders(r: Response) http.HeaderIterator {740 /// After receiving `error.ReadFailed` from the `Reader` returned by
605 return .init(r.parser.get());741 /// `reader` or `readerDecompressing`, this function accesses the
742 /// more specific error code.
743 pub fn bodyErr(response: *const Response) ?http.Reader.BodyError {
744 return response.request.reader.body_err;
606 }745 }
607746
608 test iterateHeaders {747 pub fn iterateTrailers(response: *const Response) http.HeaderIterator {
609 const response_bytes = "HTTP/1.1 200 OK\r\n" ++748 const r = &response.request.reader;
610 "LOcation:url\r\n" ++749 assert(r.state == .ready);
611 "content-tYpe: text/plain\r\n" ++750 return .{
612 "content-disposition:attachment; filename=example.txt \r\n" ++751 .bytes = r.trailers,
613 "content-Length:10\r\n" ++752 .index = 0,
614 "TRansfer-encoding:\tdeflate, chunked \r\n" ++753 .is_trailer = true,
615 "connectioN:\t keep-alive \r\n\r\n";
616
617 var header_buffer: [1024]u8 = undefined;
618 var res = Response{
619 .status = undefined,
620 .reason = undefined,
621 .version = undefined,
622 .keep_alive = false,
623 .parser = .init(&header_buffer),
624 };754 };
625
626 @memcpy(header_buffer[0..response_bytes.len], response_bytes);
627 res.parser.header_bytes_len = response_bytes.len;
628
629 var it = res.iterateHeaders();
630 {
631 const header = it.next().?;
632 try testing.expectEqualStrings("LOcation", header.name);
633 try testing.expectEqualStrings("url", header.value);
634 try testing.expect(!it.is_trailer);
635 }
636 {
637 const header = it.next().?;
638 try testing.expectEqualStrings("content-tYpe", header.name);
639 try testing.expectEqualStrings("text/plain", header.value);
640 try testing.expect(!it.is_trailer);
641 }
642 {
643 const header = it.next().?;
644 try testing.expectEqualStrings("content-disposition", header.name);
645 try testing.expectEqualStrings("attachment; filename=example.txt", header.value);
646 try testing.expect(!it.is_trailer);
647 }
648 {
649 const header = it.next().?;
650 try testing.expectEqualStrings("content-Length", header.name);
651 try testing.expectEqualStrings("10", header.value);
652 try testing.expect(!it.is_trailer);
653 }
654 {
655 const header = it.next().?;
656 try testing.expectEqualStrings("TRansfer-encoding", header.name);
657 try testing.expectEqualStrings("deflate, chunked", header.value);
658 try testing.expect(!it.is_trailer);
659 }
660 {
661 const header = it.next().?;
662 try testing.expectEqualStrings("connectioN", header.name);
663 try testing.expectEqualStrings("keep-alive", header.value);
664 try testing.expect(!it.is_trailer);
665 }
666 try testing.expectEqual(null, it.next());
667 }755 }
668};756};
669757
670/// A HTTP request that has been sent.
671///
672/// Order of operations: open -> send[ -> write -> finish] -> wait -> read
673pub const Request = struct {758pub const Request = struct {
759 /// This field is provided so that clients can observe redirected URIs.
760 ///
761 /// Its backing memory is externally provided by API users when creating a
762 /// request, and then again provided externally via `redirect_buffer` to
763 /// `receiveHead`.
674 uri: Uri,764 uri: Uri,
675 client: *Client,765 client: *Client,
676 /// This is null when the connection is released.766 /// This is null when the connection is released.
677 connection: ?*Connection,767 connection: ?*Connection,
768 reader: http.Reader,
678 keep_alive: bool,769 keep_alive: bool,
679770
680 method: http.Method,771 method: http.Method,
681 version: http.Version = .@"HTTP/1.1",772 version: http.Version = .@"HTTP/1.1",
682 transfer_encoding: RequestTransfer,773 transfer_encoding: TransferEncoding,
683 redirect_behavior: RedirectBehavior,774 redirect_behavior: RedirectBehavior,
775 accept_encoding: @TypeOf(default_accept_encoding) = default_accept_encoding,
684776
685 /// Whether the request should handle a 100-continue response before sending the request body.777 /// Whether the request should handle a 100-continue response before sending the request body.
686 handle_continue: bool,778 handle_continue: bool,
687779
688 /// The response associated with this request.
689 ///
690 /// This field is undefined until `wait` is called.
691 response: Response,
692
693 /// Standard headers that have default, but overridable, behavior.780 /// Standard headers that have default, but overridable, behavior.
694 headers: Headers,781 headers: Headers,
695782
...@@ -703,6 +790,20 @@ pub const Request = struct {...@@ -703,6 +790,20 @@ pub const Request = struct {
703 /// Externally-owned; must outlive the Request.790 /// Externally-owned; must outlive the Request.
704 privileged_headers: []const http.Header,791 privileged_headers: []const http.Header,
705792
793 pub const default_accept_encoding: [@typeInfo(http.ContentEncoding).@"enum".fields.len]bool = b: {
794 var result: [@typeInfo(http.ContentEncoding).@"enum".fields.len]bool = @splat(false);
795 result[@intFromEnum(http.ContentEncoding.gzip)] = true;
796 result[@intFromEnum(http.ContentEncoding.deflate)] = true;
797 result[@intFromEnum(http.ContentEncoding.identity)] = true;
798 break :b result;
799 };
800
801 pub const TransferEncoding = union(enum) {
802 content_length: u64,
803 chunked: void,
804 none: void,
805 };
806
706 pub const Headers = struct {807 pub const Headers = struct {
707 host: Value = .default,808 host: Value = .default,
708 authorization: Value = .default,809 authorization: Value = .default,
...@@ -728,6 +829,11 @@ pub const Request = struct {...@@ -728,6 +829,11 @@ pub const Request = struct {
728 unhandled = std.math.maxInt(u16),829 unhandled = std.math.maxInt(u16),
729 _,830 _,
730831
832 pub fn init(n: u16) RedirectBehavior {
833 assert(n != std.math.maxInt(u16));
834 return @enumFromInt(n);
835 }
836
731 pub fn subtractOne(rb: *RedirectBehavior) void {837 pub fn subtractOne(rb: *RedirectBehavior) void {
732 switch (rb.*) {838 switch (rb.*) {
733 .not_allowed => unreachable,839 .not_allowed => unreachable,
...@@ -742,98 +848,110 @@ pub const Request = struct {...@@ -742,98 +848,110 @@ pub const Request = struct {
742 }848 }
743 };849 };
744850
745 /// Frees all resources associated with the request.851 /// Returns the request's `Connection` back to the pool of the `Client`.
746 pub fn deinit(req: *Request) void {852 pub fn deinit(r: *Request) void {
747 if (req.connection) |connection| {853 if (r.connection) |connection| {
748 if (!req.response.parser.done) {854 connection.closing = connection.closing or switch (r.reader.state) {
749 // If the response wasn't fully read, then we need to close the connection.855 .ready => false,
750 connection.closing = true;856 .received_head => r.method.requestHasBody(),
751 }857 else => true,
752 req.client.connection_pool.release(req.client.allocator, connection);858 };
859 r.client.connection_pool.release(connection);
753 }860 }
754 req.* = undefined;861 r.* = undefined;
755 }862 }
756863
757 // This function must deallocate all resources associated with the request,864 /// Sends and flushes a complete request as only HTTP head, no body.
758 // or keep those which will be used.865 pub fn sendBodiless(r: *Request) Writer.Error!void {
759 // This needs to be kept in sync with deinit and request.866 try sendBodilessUnflushed(r);
760 fn redirect(req: *Request, uri: Uri) !void {867 try r.connection.?.flush();
761 assert(req.response.parser.done);868 }
762
763 req.client.connection_pool.release(req.client.allocator, req.connection.?);
764 req.connection = null;
765
766 var server_header: std.heap.FixedBufferAllocator = .init(req.response.parser.header_bytes_buffer);
767 defer req.response.parser.header_bytes_buffer = server_header.buffer[server_header.end_index..];
768 const protocol, const valid_uri = try validateUri(uri, server_header.allocator());
769
770 const new_host = valid_uri.host.?.raw;
771 const prev_host = req.uri.host.?.raw;
772 const keep_privileged_headers =
773 std.ascii.eqlIgnoreCase(valid_uri.scheme, req.uri.scheme) and
774 std.ascii.endsWithIgnoreCase(new_host, prev_host) and
775 (new_host.len == prev_host.len or new_host[new_host.len - prev_host.len - 1] == '.');
776 if (!keep_privileged_headers) {
777 // When redirecting to a different domain, strip privileged headers.
778 req.privileged_headers = &.{};
779 }
780
781 if (switch (req.response.status) {
782 .see_other => true,
783 .moved_permanently, .found => req.method == .POST,
784 else => false,
785 }) {
786 // A redirect to a GET must change the method and remove the body.
787 req.method = .GET;
788 req.transfer_encoding = .none;
789 req.headers.content_type = .omit;
790 }
791
792 if (req.transfer_encoding != .none) {
793 // The request body has already been sent. The request is
794 // still in a valid state, but the redirect must be handled
795 // manually.
796 return error.RedirectRequiresResend;
797 }
798869
799 req.uri = valid_uri;870 /// Sends but does not flush a complete request as only HTTP head, no body.
800 req.connection = try req.client.connect(new_host, uriPort(valid_uri, protocol), protocol);871 pub fn sendBodilessUnflushed(r: *Request) Writer.Error!void {
801 req.redirect_behavior.subtractOne();872 assert(r.transfer_encoding == .none);
802 req.response.parser.reset();873 assert(!r.method.requestHasBody());
803874 try sendHead(r);
804 req.response = .{
805 .version = undefined,
806 .status = undefined,
807 .reason = undefined,
808 .keep_alive = undefined,
809 .parser = req.response.parser,
810 };
811 }875 }
812876
813 pub const SendError = Connection.WriteError || error{ InvalidContentLength, UnsupportedTransferEncoding };877 /// Transfers the HTTP head over the connection and flushes.
878 ///
879 /// See also:
880 /// * `sendBodyUnflushed`
881 pub fn sendBody(r: *Request, buffer: []u8) Writer.Error!http.BodyWriter {
882 const result = try sendBodyUnflushed(r, buffer);
883 try r.connection.?.flush();
884 return result;
885 }
814886
815 /// Send the HTTP request headers to the server.887 /// Transfers the HTTP head and body over the connection and flushes.
816 pub fn send(req: *Request) SendError!void {888 pub fn sendBodyComplete(r: *Request, body: []u8) Writer.Error!void {
817 if (!req.method.requestHasBody() and req.transfer_encoding != .none)889 r.transfer_encoding = .{ .content_length = body.len };
818 return error.UnsupportedTransferEncoding;890 var bw = try sendBodyUnflushed(r, body);
891 bw.writer.end = body.len;
892 try bw.end();
893 try r.connection.?.flush();
894 }
819895
820 const connection = req.connection.?;896 /// Transfers the HTTP head over the connection, which is not flushed until
821 var connection_writer_adapter = connection.writer().adaptToNewApi();897 /// `BodyWriter.flush` or `BodyWriter.end` is called.
822 const w = &connection_writer_adapter.new_interface;898 ///
823 sendAdapted(req, connection, w) catch |err| switch (err) {899 /// See also:
824 error.WriteFailed => return connection_writer_adapter.err.?,900 /// * `sendBody`
825 else => |e| return e,901 pub fn sendBodyUnflushed(r: *Request, buffer: []u8) Writer.Error!http.BodyWriter {
902 assert(r.method.requestHasBody());
903 try sendHead(r);
904 const http_protocol_output = r.connection.?.writer();
905 return switch (r.transfer_encoding) {
906 .chunked => .{
907 .http_protocol_output = http_protocol_output,
908 .state = .{ .chunked = .init },
909 .writer = .{
910 .buffer = buffer,
911 .vtable = &.{
912 .drain = http.BodyWriter.chunkedDrain,
913 .sendFile = http.BodyWriter.chunkedSendFile,
914 },
915 },
916 },
917 .content_length => |len| .{
918 .http_protocol_output = http_protocol_output,
919 .state = .{ .content_length = len },
920 .writer = .{
921 .buffer = buffer,
922 .vtable = &.{
923 .drain = http.BodyWriter.contentLengthDrain,
924 .sendFile = http.BodyWriter.contentLengthSendFile,
925 },
926 },
927 },
928 .none => .{
929 .http_protocol_output = http_protocol_output,
930 .state = .none,
931 .writer = .{
932 .buffer = buffer,
933 .vtable = &.{
934 .drain = http.BodyWriter.noneDrain,
935 .sendFile = http.BodyWriter.noneSendFile,
936 },
937 },
938 },
826 };939 };
827 }940 }
828941
829 fn sendAdapted(req: *Request, connection: *Connection, w: *std.io.Writer) !void {942 /// Sends HTTP headers without flushing.
830 try req.method.format(w);943 fn sendHead(r: *Request) Writer.Error!void {
944 const uri = r.uri;
945 const connection = r.connection.?;
946 const w = connection.writer();
947
948 try w.writeAll(@tagName(r.method));
831 try w.writeByte(' ');949 try w.writeByte(' ');
832950
833 if (req.method == .CONNECT) {951 if (r.method == .CONNECT) {
834 try req.uri.writeToStream(w, .{ .authority = true });952 try uri.writeToStream(w, .{ .authority = true });
835 } else {953 } else {
836 try req.uri.writeToStream(w, .{954 try uri.writeToStream(w, .{
837 .scheme = connection.proxied,955 .scheme = connection.proxied,
838 .authentication = connection.proxied,956 .authentication = connection.proxied,
839 .authority = connection.proxied,957 .authority = connection.proxied,
...@@ -842,58 +960,64 @@ pub const Request = struct {...@@ -842,58 +960,64 @@ pub const Request = struct {
842 });960 });
843 }961 }
844 try w.writeByte(' ');962 try w.writeByte(' ');
845 try w.writeAll(@tagName(req.version));963 try w.writeAll(@tagName(r.version));
846 try w.writeAll("\r\n");964 try w.writeAll("\r\n");
847965
848 if (try emitOverridableHeader("host: ", req.headers.host, w)) {966 if (try emitOverridableHeader("host: ", r.headers.host, w)) {
849 try w.writeAll("host: ");967 try w.writeAll("host: ");
850 try req.uri.writeToStream(w, .{ .authority = true });968 try uri.writeToStream(w, .{ .authority = true });
851 try w.writeAll("\r\n");969 try w.writeAll("\r\n");
852 }970 }
853971
854 if (try emitOverridableHeader("authorization: ", req.headers.authorization, w)) {972 if (try emitOverridableHeader("authorization: ", r.headers.authorization, w)) {
855 if (req.uri.user != null or req.uri.password != null) {973 if (uri.user != null or uri.password != null) {
856 try w.writeAll("authorization: ");974 try w.writeAll("authorization: ");
857 const authorization = try connection.allocWriteBuffer(975 try basic_authorization.write(uri, w);
858 @intCast(basic_authorization.valueLengthFromUri(req.uri)),
859 );
860 assert(basic_authorization.value(req.uri, authorization).len == authorization.len);
861 try w.writeAll("\r\n");976 try w.writeAll("\r\n");
862 }977 }
863 }978 }
864979
865 if (try emitOverridableHeader("user-agent: ", req.headers.user_agent, w)) {980 if (try emitOverridableHeader("user-agent: ", r.headers.user_agent, w)) {
866 try w.writeAll("user-agent: zig/");981 try w.writeAll("user-agent: zig/");
867 try w.writeAll(builtin.zig_version_string);982 try w.writeAll(builtin.zig_version_string);
868 try w.writeAll(" (std.http)\r\n");983 try w.writeAll(" (std.http)\r\n");
869 }984 }
870985
871 if (try emitOverridableHeader("connection: ", req.headers.connection, w)) {986 if (try emitOverridableHeader("connection: ", r.headers.connection, w)) {
872 if (req.keep_alive) {987 if (r.keep_alive) {
873 try w.writeAll("connection: keep-alive\r\n");988 try w.writeAll("connection: keep-alive\r\n");
874 } else {989 } else {
875 try w.writeAll("connection: close\r\n");990 try w.writeAll("connection: close\r\n");
876 }991 }
877 }992 }
878993
879 if (try emitOverridableHeader("accept-encoding: ", req.headers.accept_encoding, w)) {994 if (try emitOverridableHeader("accept-encoding: ", r.headers.accept_encoding, w)) {
880 // https://github.com/ziglang/zig/issues/18937995 try w.writeAll("accept-encoding: ");
881 //try w.writeAll("accept-encoding: gzip, deflate, zstd\r\n");996 for (r.accept_encoding, 0..) |enabled, i| {
882 try w.writeAll("accept-encoding: gzip, deflate\r\n");997 if (!enabled) continue;
998 const tag: http.ContentEncoding = @enumFromInt(i);
999 if (tag == .identity) continue;
1000 const tag_name = @tagName(tag);
1001 try w.ensureUnusedCapacity(tag_name.len + 2);
1002 try w.writeAll(tag_name);
1003 try w.writeAll(", ");
1004 }
1005 w.undo(2);
1006 try w.writeAll("\r\n");
883 }1007 }
8841008
885 switch (req.transfer_encoding) {1009 switch (r.transfer_encoding) {
886 .chunked => try w.writeAll("transfer-encoding: chunked\r\n"),1010 .chunked => try w.writeAll("transfer-encoding: chunked\r\n"),
887 .content_length => |len| try w.print("content-length: {d}\r\n", .{len}),1011 .content_length => |len| try w.print("content-length: {d}\r\n", .{len}),
888 .none => {},1012 .none => {},
889 }1013 }
8901014
891 if (try emitOverridableHeader("content-type: ", req.headers.content_type, w)) {1015 if (try emitOverridableHeader("content-type: ", r.headers.content_type, w)) {
892 // The default is to omit content-type if not provided because1016 // The default is to omit content-type if not provided because
893 // "application/octet-stream" is redundant.1017 // "application/octet-stream" is redundant.
894 }1018 }
8951019
896 for (req.extra_headers) |header| {1020 for (r.extra_headers) |header| {
897 assert(header.name.len != 0);1021 assert(header.name.len != 0);
8981022
899 try w.writeAll(header.name);1023 try w.writeAll(header.name);
...@@ -904,8 +1028,8 @@ pub const Request = struct {...@@ -904,8 +1028,8 @@ pub const Request = struct {
9041028
905 if (connection.proxied) proxy: {1029 if (connection.proxied) proxy: {
906 const proxy = switch (connection.protocol) {1030 const proxy = switch (connection.protocol) {
907 .plain => req.client.http_proxy,1031 .plain => r.client.http_proxy,
908 .tls => req.client.https_proxy,1032 .tls => r.client.https_proxy,
909 } orelse break :proxy;1033 } orelse break :proxy;
9101034
911 const authorization = proxy.authorization orelse break :proxy;1035 const authorization = proxy.authorization orelse break :proxy;
...@@ -915,282 +1039,200 @@ pub const Request = struct {...@@ -915,282 +1039,200 @@ pub const Request = struct {
915 }1039 }
9161040
917 try w.writeAll("\r\n");1041 try w.writeAll("\r\n");
918
919 try connection.flush();
920 }
921
922 /// Returns true if the default behavior is required, otherwise handles
923 /// writing (or not writing) the header.
924 fn emitOverridableHeader(prefix: []const u8, v: Headers.Value, w: anytype) !bool {
925 switch (v) {
926 .default => return true,
927 .omit => return false,
928 .override => |x| {
929 try w.writeAll(prefix);
930 try w.writeAll(x);
931 try w.writeAll("\r\n");
932 return false;
933 },
934 }
935 }1042 }
9361043
937 const TransferReadError = Connection.ReadError || proto.HeadersParser.ReadError;1044 pub const ReceiveHeadError = http.Reader.HeadError || ConnectError || error{
9381045 /// Server sent headers that did not conform to the HTTP protocol.
939 const TransferReader = std.io.GenericReader(*Request, TransferReadError, transferRead);1046 ///
9401047 /// To find out more detailed diagnostics, `http.Reader.head_buffer` can be
941 fn transferReader(req: *Request) TransferReader {1048 /// passed directly to `Request.Head.parse`.
942 return .{ .context = req };1049 HttpHeadersInvalid,
943 }1050 TooManyHttpRedirects,
9441051 /// This can be avoided by calling `receiveHead` before sending the
945 fn transferRead(req: *Request, buf: []u8) TransferReadError!usize {1052 /// request body.
946 if (req.response.parser.done) return 0;1053 RedirectRequiresResend,
9471054 HttpRedirectLocationMissing,
948 var index: usize = 0;1055 HttpRedirectLocationOversize,
949 while (index == 0) {1056 HttpRedirectLocationInvalid,
950 const amt = try req.response.parser.read(req.connection.?, buf[index..], req.response.skip);1057 HttpContentEncodingUnsupported,
951 if (amt == 0 and req.response.parser.done) break;1058 HttpChunkInvalid,
952 index += amt;1059 HttpChunkTruncated,
953 }1060 HttpHeadersOversize,
9541061 UnsupportedUriScheme,
955 return index;
956 }
9571062
958 pub const WaitError = RequestError || SendError || TransferReadError ||1063 /// Sending the request failed. Error code can be found on the
959 proto.HeadersParser.CheckCompleteHeadError || Response.ParseError ||1064 /// `Connection` object.
960 error{1065 WriteFailed,
961 TooManyHttpRedirects,1066 };
962 RedirectRequiresResend,
963 HttpRedirectLocationMissing,
964 HttpRedirectLocationInvalid,
965 CompressionInitializationFailed,
966 CompressionUnsupported,
967 };
9681067
969 /// Waits for a response from the server and parses any headers that are sent.
970 /// This function will block until the final response is received.
971 ///
972 /// If handling redirects and the request has no payload, then this1068 /// If handling redirects and the request has no payload, then this
973 /// function will automatically follow redirects. If a request payload is1069 /// function will automatically follow redirects.
974 /// present, then this function will error with1070 ///
975 /// error.RedirectRequiresResend.1071 /// If a request payload is present, then this function will error with
1072 /// `error.RedirectRequiresResend`.
1073 ///
1074 /// This function takes an auxiliary buffer to store the arbitrarily large
1075 /// URI which may need to be merged with the previous URI, and that data
1076 /// needs to survive across different connections, which is where the input
1077 /// buffer lives.
976 ///1078 ///
977 /// Must be called after `send` and, if any data was written to the request1079 /// `redirect_buffer` must outlive accesses to `Request.uri`. If this
978 /// body, then also after `finish`.1080 /// buffer capacity would be exceeded, `error.HttpRedirectLocationOversize`
979 pub fn wait(req: *Request) WaitError!void {1081 /// is returned instead. This buffer may be empty if no redirects are to be
1082 /// handled.
1083 ///
1084 /// If this fails with `error.ReadFailed` then the `Connection.getReadError`
1085 /// method of `r.connection` can be used to get more detailed information.
1086 pub fn receiveHead(r: *Request, redirect_buffer: []u8) ReceiveHeadError!Response {
1087 var aux_buf = redirect_buffer;
980 while (true) {1088 while (true) {
981 // This while loop is for handling redirects, which means the request's1089 const head_buffer = try r.reader.receiveHead();
982 // connection may be different than the previous iteration. However, it1090 const response: Response = .{
983 // is still guaranteed to be non-null with each iteration of this loop.1091 .request = r,
984 const connection = req.connection.?;1092 .head = Response.Head.parse(head_buffer) catch return error.HttpHeadersInvalid,
9851093 };
986 while (true) { // read headers1094 const head = &response.head;
987 try connection.fill();
988
989 const nchecked = try req.response.parser.checkCompleteHead(connection.peek());
990 connection.drop(@intCast(nchecked));
9911095
992 if (req.response.parser.state.isContent()) break;1096 if (head.status == .@"continue") {
1097 if (r.handle_continue) continue;
1098 return response; // we're not handling the 100-continue
993 }1099 }
9941100
995 try req.response.parse(req.response.parser.get());1101 // This while loop is for handling redirects, which means the request's
9961102 // connection may be different than the previous iteration. However, it
997 if (req.response.status == .@"continue") {1103 // is still guaranteed to be non-null with each iteration of this loop.
998 // We're done parsing the continue response; reset to prepare1104 const connection = r.connection.?;
999 // for the real response.
1000 req.response.parser.done = true;
1001 req.response.parser.reset();
1002
1003 if (req.handle_continue)
1004 continue;
1005
1006 return; // we're not handling the 100-continue
1007 }
10081105
1009 // we're switching protocols, so this connection is no longer doing http1106 if (r.method == .CONNECT and head.status.class() == .success) {
1010 if (req.method == .CONNECT and req.response.status.class() == .success) {1107 // This connection is no longer doing HTTP.
1011 connection.closing = false;1108 connection.closing = false;
1012 req.response.parser.done = true;1109 return response;
1013 return; // the connection is not HTTP past this point
1014 }1110 }
10151111
1016 connection.closing = !req.response.keep_alive or !req.keep_alive;1112 connection.closing = !head.keep_alive or !r.keep_alive;
10171113
1018 // Any response to a HEAD request and any response with a 1xx1114 // Any response to a HEAD request and any response with a 1xx
1019 // (Informational), 204 (No Content), or 304 (Not Modified) status1115 // (Informational), 204 (No Content), or 304 (Not Modified) status
1020 // code is always terminated by the first empty line after the1116 // code is always terminated by the first empty line after the
1021 // header fields, regardless of the header fields present in the1117 // header fields, regardless of the header fields present in the
1022 // message.1118 // message.
1023 if (req.method == .HEAD or req.response.status.class() == .informational or1119 if (r.method == .HEAD or head.status.class() == .informational or
1024 req.response.status == .no_content or req.response.status == .not_modified)1120 head.status == .no_content or head.status == .not_modified)
1025 {1121 {
1026 req.response.parser.done = true;1122 return response;
1027 return; // The response is empty; no further setup or redirection is necessary.
1028 }
1029
1030 switch (req.response.transfer_encoding) {
1031 .none => {
1032 if (req.response.content_length) |cl| {
1033 req.response.parser.next_chunk_length = cl;
1034
1035 if (cl == 0) req.response.parser.done = true;
1036 } else {
1037 // read until the connection is closed
1038 req.response.parser.next_chunk_length = std.math.maxInt(u64);
1039 }
1040 },
1041 .chunked => {
1042 req.response.parser.next_chunk_length = 0;
1043 req.response.parser.state = .chunk_head_size;
1044 },
1045 }1123 }
10461124
1047 if (req.response.status.class() == .redirect and req.redirect_behavior != .unhandled) {1125 if (head.status.class() == .redirect and r.redirect_behavior != .unhandled) {
1048 // skip the body of the redirect response, this will at least1126 if (r.redirect_behavior == .not_allowed) {
1049 // leave the connection in a known good state.1127 // Connection can still be reused by skipping the body.
1050 req.response.skip = true;1128 const reader = r.reader.bodyReader(&.{}, head.transfer_encoding, head.content_length);
1051 assert(try req.transferRead(&.{}) == 0); // we're skipping, no buffer is necessary1129 _ = reader.discardRemaining() catch |err| switch (err) {
10521130 error.ReadFailed => connection.closing = true,
1053 if (req.redirect_behavior == .not_allowed) return error.TooManyHttpRedirects;1131 };
10541132 return error.TooManyHttpRedirects;
1055 const location = req.response.location orelse
1056 return error.HttpRedirectLocationMissing;
1057
1058 // This mutates the beginning of header_bytes_buffer and uses that
1059 // for the backing memory of the returned Uri.
1060 try req.redirect(req.uri.resolve_inplace(
1061 location,
1062 &req.response.parser.header_bytes_buffer,
1063 ) catch |err| switch (err) {
1064 error.UnexpectedCharacter,
1065 error.InvalidFormat,
1066 error.InvalidPort,
1067 => return error.HttpRedirectLocationInvalid,
1068 error.NoSpaceLeft => return error.HttpHeadersOversize,
1069 });
1070 try req.send();
1071 } else {
1072 req.response.skip = false;
1073 if (!req.response.parser.done) {
1074 switch (req.response.transfer_compression) {
1075 .identity => req.response.compression = .none,
1076 .compress, .@"x-compress" => return error.CompressionUnsupported,
1077 // I'm about to upstream my http.Client rewrite
1078 .deflate => return error.CompressionUnsupported,
1079 // I'm about to upstream my http.Client rewrite
1080 .gzip, .@"x-gzip" => return error.CompressionUnsupported,
1081 // https://github.com/ziglang/zig/issues/18937
1082 //.zstd => req.response.compression = .{
1083 // .zstd = std.compress.zstd.decompressStream(req.client.allocator, req.transferReader()),
1084 //},
1085 .zstd => return error.CompressionUnsupported,
1086 }
1087 }1133 }
10881134 try r.redirect(head, &aux_buf);
1089 break;1135 try r.sendBodiless();
1136 continue;
1090 }1137 }
1091 }
1092 }
1093
1094 pub const ReadError = TransferReadError || proto.HeadersParser.CheckCompleteHeadError ||
1095 error{ DecompressionFailure, InvalidTrailers };
10961138
1097 pub const Reader = std.io.GenericReader(*Request, ReadError, read);1139 if (!r.accept_encoding[@intFromEnum(head.content_encoding)])
1140 return error.HttpContentEncodingUnsupported;
10981141
1099 pub fn reader(req: *Request) Reader {1142 return response;
1100 return .{ .context = req };
1101 }
1102
1103 /// Reads data from the response body. Must be called after `wait`.
1104 pub fn read(req: *Request, buffer: []u8) ReadError!usize {
1105 const out_index = switch (req.response.compression) {
1106 // I'm about to upstream my http client rewrite
1107 //.deflate => |*deflate| deflate.readSlice(buffer) catch return error.DecompressionFailure,
1108 //.gzip => |*gzip| gzip.read(buffer) catch return error.DecompressionFailure,
1109 // https://github.com/ziglang/zig/issues/18937
1110 //.zstd => |*zstd| zstd.read(buffer) catch return error.DecompressionFailure,
1111 else => try req.transferRead(buffer),
1112 };
1113 if (out_index > 0) return out_index;
1114
1115 while (!req.response.parser.state.isContent()) { // read trailing headers
1116 try req.connection.?.fill();
1117
1118 const nchecked = try req.response.parser.checkCompleteHead(req.connection.?.peek());
1119 req.connection.?.drop(@intCast(nchecked));
1120 }1143 }
1121
1122 return 0;
1123 }1144 }
11241145
1125 /// Reads data from the response body. Must be called after `wait`.1146 /// This function takes an auxiliary buffer to store the arbitrarily large
1126 pub fn readAll(req: *Request, buffer: []u8) !usize {1147 /// URI which may need to be merged with the previous URI, and that data
1127 var index: usize = 0;1148 /// needs to survive across different connections, which is where the input
1128 while (index < buffer.len) {1149 /// buffer lives.
1129 const amt = try read(req, buffer[index..]);1150 ///
1130 if (amt == 0) break;1151 /// `aux_buf` must outlive accesses to `Request.uri`.
1131 index += amt;1152 fn redirect(r: *Request, head: *const Response.Head, aux_buf: *[]u8) !void {
1153 const new_location = head.location orelse return error.HttpRedirectLocationMissing;
1154 if (new_location.len > aux_buf.*.len) return error.HttpRedirectLocationOversize;
1155 const location = aux_buf.*[0..new_location.len];
1156 @memcpy(location, new_location);
1157 {
1158 // Skip the body of the redirect response to leave the connection in
1159 // the correct state. This causes `new_location` to be invalidated.
1160 const reader = r.reader.bodyReader(&.{}, head.transfer_encoding, head.content_length);
1161 _ = reader.discardRemaining() catch |err| switch (err) {
1162 error.ReadFailed => return r.reader.body_err.?,
1163 };
1132 }1164 }
1133 return index;1165 const new_uri = r.uri.resolveInPlace(location.len, aux_buf) catch |err| switch (err) {
1134 }1166 error.UnexpectedCharacter => return error.HttpRedirectLocationInvalid,
11351167 error.InvalidFormat => return error.HttpRedirectLocationInvalid,
1136 pub const WriteError = Connection.WriteError || error{ NotWriteable, MessageTooLong };1168 error.InvalidPort => return error.HttpRedirectLocationInvalid,
11371169 error.NoSpaceLeft => return error.HttpRedirectLocationOversize,
1138 pub const Writer = std.io.GenericWriter(*Request, WriteError, write);1170 };
11391171
1140 pub fn writer(req: *Request) Writer {1172 const protocol = Protocol.fromUri(new_uri) orelse return error.UnsupportedUriScheme;
1141 return .{ .context = req };1173 const old_connection = r.connection.?;
1142 }1174 const old_host = old_connection.host();
1175 var new_host_name_buffer: [Uri.host_name_max]u8 = undefined;
1176 const new_host = try new_uri.getHost(&new_host_name_buffer);
1177 const keep_privileged_headers =
1178 std.ascii.eqlIgnoreCase(r.uri.scheme, new_uri.scheme) and
1179 sameParentDomain(old_host, new_host);
11431180
1144 /// Write `bytes` to the server. The `transfer_encoding` field determines how data will be sent.1181 r.client.connection_pool.release(old_connection);
1145 /// Must be called after `send` and before `finish`.1182 r.connection = null;
1146 pub fn write(req: *Request, bytes: []const u8) WriteError!usize {
1147 switch (req.transfer_encoding) {
1148 .chunked => {
1149 if (bytes.len > 0) {
1150 try req.connection.?.writer().print("{x}\r\n", .{bytes.len});
1151 try req.connection.?.writer().writeAll(bytes);
1152 try req.connection.?.writer().writeAll("\r\n");
1153 }
11541183
1155 return bytes.len;1184 if (!keep_privileged_headers) {
1156 },1185 // When redirecting to a different domain, strip privileged headers.
1157 .content_length => |*len| {1186 r.privileged_headers = &.{};
1158 if (len.* < bytes.len) return error.MessageTooLong;1187 }
11591188
1160 const amt = try req.connection.?.write(bytes);1189 if (switch (head.status) {
1161 len.* -= amt;1190 .see_other => true,
1162 return amt;1191 .moved_permanently, .found => r.method == .POST,
1163 },1192 else => false,
1164 .none => return error.NotWriteable,1193 }) {
1194 // A redirect to a GET must change the method and remove the body.
1195 r.method = .GET;
1196 r.transfer_encoding = .none;
1197 r.headers.content_type = .omit;
1165 }1198 }
1166 }
11671199
1168 /// Write `bytes` to the server. The `transfer_encoding` field determines how data will be sent.1200 if (r.transfer_encoding != .none) {
1169 /// Must be called after `send` and before `finish`.1201 // The request body has already been sent. The request is
1170 pub fn writeAll(req: *Request, bytes: []const u8) WriteError!void {1202 // still in a valid state, but the redirect must be handled
1171 var index: usize = 0;1203 // manually.
1172 while (index < bytes.len) {1204 return error.RedirectRequiresResend;
1173 index += try write(req, bytes[index..]);
1174 }1205 }
1175 }
11761206
1177 pub const FinishError = WriteError || error{MessageNotCompleted};1207 const new_connection = try r.client.connect(new_host, uriPort(new_uri, protocol), protocol);
1208 r.uri = new_uri;
1209 r.connection = new_connection;
1210 r.reader = .{
1211 .in = new_connection.reader(),
1212 .state = .ready,
1213 // Populated when `http.Reader.bodyReader` is called.
1214 .interface = undefined,
1215 };
1216 r.redirect_behavior.subtractOne();
1217 }
11781218
1179 /// Finish the body of a request. This notifies the server that you have no more data to send.1219 /// Returns true if the default behavior is required, otherwise handles
1180 /// Must be called after `send`.1220 /// writing (or not writing) the header.
1181 pub fn finish(req: *Request) FinishError!void {1221 fn emitOverridableHeader(prefix: []const u8, v: Headers.Value, bw: *Writer) Writer.Error!bool {
1182 switch (req.transfer_encoding) {1222 switch (v) {
1183 .chunked => try req.connection.?.writer().writeAll("0\r\n\r\n"),1223 .default => return true,
1184 .content_length => |len| if (len != 0) return error.MessageNotCompleted,1224 .omit => return false,
1185 .none => {},1225 .override => |x| {
1226 var vecs: [3][]const u8 = .{ prefix, x, "\r\n" };
1227 try bw.writeVecAll(&vecs);
1228 return false;
1229 },
1186 }1230 }
1187
1188 try req.connection.?.flush();
1189 }1231 }
1190};1232};
11911233
1192pub const Proxy = struct {1234pub const Proxy = struct {
1193 protocol: Connection.Protocol,1235 protocol: Protocol,
1194 host: []const u8,1236 host: []const u8,
1195 authorization: ?[]const u8,1237 authorization: ?[]const u8,
1196 port: u16,1238 port: u16,
...@@ -1204,10 +1246,8 @@ pub const Proxy = struct {...@@ -1204,10 +1246,8 @@ pub const Proxy = struct {
1204pub fn deinit(client: *Client) void {1246pub fn deinit(client: *Client) void {
1205 assert(client.connection_pool.used.first == null); // There are still active requests.1247 assert(client.connection_pool.used.first == null); // There are still active requests.
12061248
1207 client.connection_pool.deinit(client.allocator);1249 client.connection_pool.deinit();
12081250 if (!disable_tls) client.ca_bundle.deinit(client.allocator);
1209 if (!disable_tls)
1210 client.ca_bundle.deinit(client.allocator);
12111251
1212 client.* = undefined;1252 client.* = undefined;
1213}1253}
...@@ -1249,24 +1289,21 @@ fn createProxyFromEnvVar(arena: Allocator, env_var_names: []const []const u8) !?...@@ -1249,24 +1289,21 @@ fn createProxyFromEnvVar(arena: Allocator, env_var_names: []const []const u8) !?
1249 } else return null;1289 } else return null;
12501290
1251 const uri = Uri.parse(content) catch try Uri.parseAfterScheme("http", content);1291 const uri = Uri.parse(content) catch try Uri.parseAfterScheme("http", content);
1252 const protocol, const valid_uri = validateUri(uri, arena) catch |err| switch (err) {1292 const protocol = Protocol.fromUri(uri) orelse return null;
1253 error.UnsupportedUriScheme => return null,1293 const raw_host = try uri.getHostAlloc(arena);
1254 error.UriMissingHost => return error.HttpProxyMissingHost,
1255 error.OutOfMemory => |e| return e,
1256 };
12571294
1258 const authorization: ?[]const u8 = if (valid_uri.user != null or valid_uri.password != null) a: {1295 const authorization: ?[]const u8 = if (uri.user != null or uri.password != null) a: {
1259 const authorization = try arena.alloc(u8, basic_authorization.valueLengthFromUri(valid_uri));1296 const authorization = try arena.alloc(u8, basic_authorization.valueLengthFromUri(uri));
1260 assert(basic_authorization.value(valid_uri, authorization).len == authorization.len);1297 assert(basic_authorization.value(uri, authorization).len == authorization.len);
1261 break :a authorization;1298 break :a authorization;
1262 } else null;1299 } else null;
12631300
1264 const proxy = try arena.create(Proxy);1301 const proxy = try arena.create(Proxy);
1265 proxy.* = .{1302 proxy.* = .{
1266 .protocol = protocol,1303 .protocol = protocol,
1267 .host = valid_uri.host.?.raw,1304 .host = raw_host,
1268 .authorization = authorization,1305 .authorization = authorization,
1269 .port = uriPort(valid_uri, protocol),1306 .port = uriPort(uri, protocol),
1270 .supports_connect = true,1307 .supports_connect = true,
1271 };1308 };
1272 return proxy;1309 return proxy;
...@@ -1277,10 +1314,8 @@ pub const basic_authorization = struct {...@@ -1277,10 +1314,8 @@ pub const basic_authorization = struct {
1277 pub const max_password_len = 255;1314 pub const max_password_len = 255;
1278 pub const max_value_len = valueLength(max_user_len, max_password_len);1315 pub const max_value_len = valueLength(max_user_len, max_password_len);
12791316
1280 const prefix = "Basic ";
1281
1282 pub fn valueLength(user_len: usize, password_len: usize) usize {1317 pub fn valueLength(user_len: usize, password_len: usize) usize {
1283 return prefix.len + std.base64.standard.Encoder.calcSize(user_len + 1 + password_len);1318 return "Basic ".len + std.base64.standard.Encoder.calcSize(user_len + 1 + password_len);
1284 }1319 }
12851320
1286 pub fn valueLengthFromUri(uri: Uri) usize {1321 pub fn valueLengthFromUri(uri: Uri) usize {
...@@ -1300,37 +1335,70 @@ pub const basic_authorization = struct {...@@ -1300,37 +1335,70 @@ pub const basic_authorization = struct {
1300 }1335 }
13011336
1302 pub fn value(uri: Uri, out: []u8) []u8 {1337 pub fn value(uri: Uri, out: []u8) []u8 {
1303 const user: Uri.Component = uri.user orelse .empty;1338 var bw: Writer = .fixed(out);
1304 const password: Uri.Component = uri.password orelse .empty;1339 write(uri, &bw) catch unreachable;
13051340 return bw.buffered();
1306 var buf: [max_user_len + ":".len + max_password_len]u8 = undefined;1341 }
1307 var w: std.io.Writer = .fixed(&buf);
1308 user.formatUser(&w) catch unreachable; // fixed
1309 password.formatPassword(&w) catch unreachable; // fixed
13101342
1311 @memcpy(out[0..prefix.len], prefix);1343 pub fn write(uri: Uri, out: *Writer) Writer.Error!void {
1312 const base64 = std.base64.standard.Encoder.encode(out[prefix.len..], w.buffered());1344 var buf: [max_user_len + 1 + max_password_len]u8 = undefined;
1313 return out[0 .. prefix.len + base64.len];1345 var w: Writer = .fixed(&buf);
1346 const user: Uri.Component = uri.user orelse .empty;
1347 const password: Uri.Component = uri.user orelse .empty;
1348 user.formatUser(&w) catch unreachable;
1349 w.writeByte(':') catch unreachable;
1350 password.formatPassword(&w) catch unreachable;
1351 try out.print("Basic {b64}", .{w.buffered()});
1314 }1352 }
1315};1353};
13161354
1317pub const ConnectTcpError = Allocator.Error || error{ ConnectionRefused, NetworkUnreachable, ConnectionTimedOut, ConnectionResetByPeer, TemporaryNameServerFailure, NameServerFailure, UnknownHostName, HostLacksNetworkAddresses, UnexpectedConnectFailure, TlsInitializationFailed };1355pub const ConnectTcpError = Allocator.Error || error{
1356 ConnectionRefused,
1357 NetworkUnreachable,
1358 ConnectionTimedOut,
1359 ConnectionResetByPeer,
1360 TemporaryNameServerFailure,
1361 NameServerFailure,
1362 UnknownHostName,
1363 HostLacksNetworkAddresses,
1364 UnexpectedConnectFailure,
1365 TlsInitializationFailed,
1366};
13181367
1319/// Connect to `host:port` using the specified protocol. This will reuse a connection if one is already open.1368/// Reuses a `Connection` if one matching `host` and `port` is already open.
1320///1369///
1321/// This function is threadsafe.1370/// Threadsafe.
1322pub fn connectTcp(client: *Client, host: []const u8, port: u16, protocol: Connection.Protocol) ConnectTcpError!*Connection {1371pub fn connectTcp(
1323 if (client.connection_pool.findConnection(.{1372 client: *Client,
1324 .host = host,1373 host: []const u8,
1325 .port = port,1374 port: u16,
1326 .protocol = protocol,1375 protocol: Protocol,
1327 })) |node| return node;1376) ConnectTcpError!*Connection {
1377 return connectTcpOptions(client, .{ .host = host, .port = port, .protocol = protocol });
1378}
1379
1380pub const ConnectTcpOptions = struct {
1381 host: []const u8,
1382 port: u16,
1383 protocol: Protocol,
13281384
1329 if (disable_tls and protocol == .tls)1385 proxied_host: ?[]const u8 = null,
1330 return error.TlsInitializationFailed;1386 proxied_port: ?u16 = null,
1387};
13311388
1332 const conn = try client.allocator.create(Connection);1389pub fn connectTcpOptions(client: *Client, options: ConnectTcpOptions) ConnectTcpError!*Connection {
1333 errdefer client.allocator.destroy(conn);1390 const host = options.host;
1391 const port = options.port;
1392 const protocol = options.protocol;
1393
1394 const proxied_host = options.proxied_host orelse host;
1395 const proxied_port = options.proxied_port orelse port;
1396
1397 if (client.connection_pool.findConnection(.{
1398 .host = proxied_host,
1399 .port = proxied_port,
1400 .protocol = protocol,
1401 })) |conn| return conn;
13341402
1335 const stream = net.tcpConnectToHost(client.allocator, host, port) catch |err| switch (err) {1403 const stream = net.tcpConnectToHost(client.allocator, host, port) catch |err| switch (err) {
1336 error.ConnectionRefused => return error.ConnectionRefused,1404 error.ConnectionRefused => return error.ConnectionRefused,
...@@ -1345,53 +1413,19 @@ pub fn connectTcp(client: *Client, host: []const u8, port: u16, protocol: Connec...@@ -1345,53 +1413,19 @@ pub fn connectTcp(client: *Client, host: []const u8, port: u16, protocol: Connec
1345 };1413 };
1346 errdefer stream.close();1414 errdefer stream.close();
13471415
1348 conn.* = .{1416 switch (protocol) {
1349 .stream = stream,1417 .tls => {
1350 .tls_client = undefined,1418 if (disable_tls) return error.TlsInitializationFailed;
13511419 const tc = try Connection.Tls.create(client, proxied_host, proxied_port, stream);
1352 .protocol = protocol,1420 client.connection_pool.addUsed(&tc.connection);
1353 .host = try client.allocator.dupe(u8, host),1421 return &tc.connection;
1354 .port = port,1422 },
13551423 .plain => {
1356 .pool_node = .{},1424 const pc = try Connection.Plain.create(client, proxied_host, proxied_port, stream);
1357 };1425 client.connection_pool.addUsed(&pc.connection);
1358 errdefer client.allocator.free(conn.host);1426 return &pc.connection;
13591427 },
1360 if (protocol == .tls) {
1361 if (disable_tls) unreachable;
1362
1363 conn.tls_client = try client.allocator.create(std.crypto.tls.Client);
1364 errdefer client.allocator.destroy(conn.tls_client);
1365
1366 const ssl_key_log_file: ?std.fs.File = if (std.options.http_enable_ssl_key_log_file) ssl_key_log_file: {
1367 const ssl_key_log_path = std.process.getEnvVarOwned(client.allocator, "SSLKEYLOGFILE") catch |err| switch (err) {
1368 error.EnvironmentVariableNotFound, error.InvalidWtf8 => break :ssl_key_log_file null,
1369 error.OutOfMemory => return error.OutOfMemory,
1370 };
1371 defer client.allocator.free(ssl_key_log_path);
1372 break :ssl_key_log_file std.fs.cwd().createFile(ssl_key_log_path, .{
1373 .truncate = false,
1374 .mode = switch (builtin.os.tag) {
1375 .windows, .wasi => 0,
1376 else => 0o600,
1377 },
1378 }) catch null;
1379 } else null;
1380 errdefer if (ssl_key_log_file) |key_log_file| key_log_file.close();
1381
1382 conn.tls_client.* = std.crypto.tls.Client.init(stream, .{
1383 .host = .{ .explicit = host },
1384 .ca = .{ .bundle = client.ca_bundle },
1385 .ssl_key_log_file = ssl_key_log_file,
1386 }) catch return error.TlsInitializationFailed;
1387 // This is appropriate for HTTPS because the HTTP headers contain
1388 // the content length which is used to detect truncation attacks.
1389 conn.tls_client.allow_truncation_attacks = true;
1390 }1428 }
1391
1392 client.connection_pool.addUsed(conn);
1393
1394 return conn;
1395}1429}
13961430
1397pub const ConnectUnixError = Allocator.Error || std.posix.SocketError || error{NameTooLong} || std.posix.ConnectError;1431pub const ConnectUnixError = Allocator.Error || std.posix.SocketError || error{NameTooLong} || std.posix.ConnectError;
...@@ -1429,69 +1463,67 @@ pub fn connectUnix(client: *Client, path: []const u8) ConnectUnixError!*Connecti...@@ -1429,69 +1463,67 @@ pub fn connectUnix(client: *Client, path: []const u8) ConnectUnixError!*Connecti
1429 return &conn.data;1463 return &conn.data;
1430}1464}
14311465
1432/// Connect to `tunnel_host:tunnel_port` using the specified proxy with HTTP1466/// Connect to `proxied_host:proxied_port` using the specified proxy with HTTP
1433/// CONNECT. This will reuse a connection if one is already open.1467/// CONNECT. This will reuse a connection if one is already open.
1434///1468///
1435/// This function is threadsafe.1469/// This function is threadsafe.
1436pub fn connectTunnel(1470pub fn connectProxied(
1437 client: *Client,1471 client: *Client,
1438 proxy: *Proxy,1472 proxy: *Proxy,
1439 tunnel_host: []const u8,1473 proxied_host: []const u8,
1440 tunnel_port: u16,1474 proxied_port: u16,
1441) !*Connection {1475) !*Connection {
1442 if (!proxy.supports_connect) return error.TunnelNotSupported;1476 if (!proxy.supports_connect) return error.TunnelNotSupported;
14431477
1444 if (client.connection_pool.findConnection(.{1478 if (client.connection_pool.findConnection(.{
1445 .host = tunnel_host,1479 .host = proxied_host,
1446 .port = tunnel_port,1480 .port = proxied_port,
1447 .protocol = proxy.protocol,1481 .protocol = proxy.protocol,
1448 })) |node|1482 })) |node| return node;
1449 return node;
14501483
1451 var maybe_valid = false;1484 var maybe_valid = false;
1452 (tunnel: {1485 (tunnel: {
1453 const conn = try client.connectTcp(proxy.host, proxy.port, proxy.protocol);1486 const connection = try client.connectTcpOptions(.{
1487 .host = proxy.host,
1488 .port = proxy.port,
1489 .protocol = proxy.protocol,
1490 .proxied_host = proxied_host,
1491 .proxied_port = proxied_port,
1492 });
1454 errdefer {1493 errdefer {
1455 conn.closing = true;1494 connection.closing = true;
1456 client.connection_pool.release(client.allocator, conn);1495 client.connection_pool.release(connection);
1457 }1496 }
14581497
1459 var buffer: [8096]u8 = undefined;1498 var req = client.request(.CONNECT, .{
1460 var req = client.open(.CONNECT, .{
1461 .scheme = "http",1499 .scheme = "http",
1462 .host = .{ .raw = tunnel_host },1500 .host = .{ .raw = proxied_host },
1463 .port = tunnel_port,1501 .port = proxied_port,
1464 }, .{1502 }, .{
1465 .redirect_behavior = .unhandled,1503 .redirect_behavior = .unhandled,
1466 .connection = conn,1504 .connection = connection,
1467 .server_header_buffer = &buffer,
1468 }) catch |err| {1505 }) catch |err| {
1469 std.log.debug("err {}", .{err});
1470 break :tunnel err;1506 break :tunnel err;
1471 };1507 };
1472 defer req.deinit();1508 defer req.deinit();
14731509
1474 req.send() catch |err| break :tunnel err;1510 req.sendBodiless() catch |err| break :tunnel err;
1475 req.wait() catch |err| break :tunnel err;1511 const response = req.receiveHead(&.{}) catch |err| break :tunnel err;
14761512
1477 if (req.response.status.class() == .server_error) {1513 if (response.head.status.class() == .server_error) {
1478 maybe_valid = true;1514 maybe_valid = true;
1479 break :tunnel error.ServerError;1515 break :tunnel error.ServerError;
1480 }1516 }
14811517
1482 if (req.response.status != .ok) break :tunnel error.ConnectionRefused;1518 if (response.head.status != .ok) break :tunnel error.ConnectionRefused;
14831519
1484 // this connection is now a tunnel, so we can't use it for anything else, it will only be released when the client is de-initialized.1520 // this connection is now a tunnel, so we can't use it for anything
1521 // else, it will only be released when the client is de-initialized.
1485 req.connection = null;1522 req.connection = null;
14861523
1487 client.allocator.free(conn.host);1524 connection.closing = false;
1488 conn.host = try client.allocator.dupe(u8, tunnel_host);
1489 errdefer client.allocator.free(conn.host);
14901525
1491 conn.port = tunnel_port;1526 return connection;
1492 conn.closing = false;
1493
1494 return conn;
1495 }) catch {1527 }) catch {
1496 // something went wrong with the tunnel1528 // something went wrong with the tunnel
1497 proxy.supports_connect = maybe_valid;1529 proxy.supports_connect = maybe_valid;
...@@ -1499,12 +1531,11 @@ pub fn connectTunnel(...@@ -1499,12 +1531,11 @@ pub fn connectTunnel(
1499 };1531 };
1500}1532}
15011533
1502// Prevents a dependency loop in open()1534pub const ConnectError = ConnectTcpError || RequestError;
1503const ConnectErrorPartial = ConnectTcpError || error{ UnsupportedUriScheme, ConnectionRefused };
1504pub const ConnectError = ConnectErrorPartial || RequestError;
15051535
1506/// Connect to `host:port` using the specified protocol. This will reuse a1536/// Connect to `host:port` using the specified protocol. This will reuse a
1507/// connection if one is already open.1537/// connection if one is already open.
1538///
1508/// If a proxy is configured for the client, then the proxy will be used to1539/// If a proxy is configured for the client, then the proxy will be used to
1509/// connect to the host.1540/// connect to the host.
1510///1541///
...@@ -1513,7 +1544,7 @@ pub fn connect(...@@ -1513,7 +1544,7 @@ pub fn connect(
1513 client: *Client,1544 client: *Client,
1514 host: []const u8,1545 host: []const u8,
1515 port: u16,1546 port: u16,
1516 protocol: Connection.Protocol,1547 protocol: Protocol,
1517) ConnectError!*Connection {1548) ConnectError!*Connection {
1518 const proxy = switch (protocol) {1549 const proxy = switch (protocol) {
1519 .plain => client.http_proxy,1550 .plain => client.http_proxy,
...@@ -1528,32 +1559,24 @@ pub fn connect(...@@ -1528,32 +1559,24 @@ pub fn connect(
1528 }1559 }
15291560
1530 if (proxy.supports_connect) tunnel: {1561 if (proxy.supports_connect) tunnel: {
1531 return connectTunnel(client, proxy, host, port) catch |err| switch (err) {1562 return connectProxied(client, proxy, host, port) catch |err| switch (err) {
1532 error.TunnelNotSupported => break :tunnel,1563 error.TunnelNotSupported => break :tunnel,
1533 else => |e| return e,1564 else => |e| return e,
1534 };1565 };
1535 }1566 }
15361567
1537 // fall back to using the proxy as a normal http proxy1568 // fall back to using the proxy as a normal http proxy
1538 const conn = try client.connectTcp(proxy.host, proxy.port, proxy.protocol);1569 const connection = try client.connectTcp(proxy.host, proxy.port, proxy.protocol);
1539 errdefer {1570 connection.proxied = true;
1540 conn.closing = true;1571 return connection;
1541 client.connection_pool.release(conn);
1542 }
1543
1544 conn.proxied = true;
1545 return conn;
1546}1572}
15471573
1548pub const RequestError = ConnectTcpError || ConnectErrorPartial || Request.SendError ||1574pub const RequestError = ConnectTcpError || error{
1549 std.fmt.ParseIntError || Connection.WriteError ||1575 UnsupportedUriScheme,
1550 error{1576 UriMissingHost,
1551 UnsupportedUriScheme,1577 UriHostTooLong,
1552 UriMissingHost,1578 CertificateBundleLoadFailure,
15531579};
1554 CertificateBundleLoadFailure,
1555 UnsupportedTransferEncoding,
1556 };
15571580
1558pub const RequestOptions = struct {1581pub const RequestOptions = struct {
1559 version: http.Version = .@"HTTP/1.1",1582 version: http.Version = .@"HTTP/1.1",
...@@ -1578,11 +1601,6 @@ pub const RequestOptions = struct {...@@ -1578,11 +1601,6 @@ pub const RequestOptions = struct {
1578 /// payload or the server has acknowledged the payload).1601 /// payload or the server has acknowledged the payload).
1579 redirect_behavior: Request.RedirectBehavior = @enumFromInt(3),1602 redirect_behavior: Request.RedirectBehavior = @enumFromInt(3),
15801603
1581 /// Externally-owned memory used to store the server's entire HTTP header.
1582 /// `error.HttpHeadersOversize` is returned from read() when a
1583 /// client sends too many bytes of HTTP headers.
1584 server_header_buffer: []u8,
1585
1586 /// Must be an already acquired connection.1604 /// Must be an already acquired connection.
1587 connection: ?*Connection = null,1605 connection: ?*Connection = null,
15881606
...@@ -1598,38 +1616,17 @@ pub const RequestOptions = struct {...@@ -1598,38 +1616,17 @@ pub const RequestOptions = struct {
1598 privileged_headers: []const http.Header = &.{},1616 privileged_headers: []const http.Header = &.{},
1599};1617};
16001618
1601fn validateUri(uri: Uri, arena: Allocator) !struct { Connection.Protocol, Uri } {1619fn uriPort(uri: Uri, protocol: Protocol) u16 {
1602 const protocol_map = std.StaticStringMap(Connection.Protocol).initComptime(.{1620 return uri.port orelse protocol.port();
1603 .{ "http", .plain },
1604 .{ "ws", .plain },
1605 .{ "https", .tls },
1606 .{ "wss", .tls },
1607 });
1608 const protocol = protocol_map.get(uri.scheme) orelse return error.UnsupportedUriScheme;
1609 var valid_uri = uri;
1610 // The host is always going to be needed as a raw string for hostname resolution anyway.
1611 valid_uri.host = .{
1612 .raw = try (uri.host orelse return error.UriMissingHost).toRawMaybeAlloc(arena),
1613 };
1614 return .{ protocol, valid_uri };
1615}
1616
1617fn uriPort(uri: Uri, protocol: Connection.Protocol) u16 {
1618 return uri.port orelse switch (protocol) {
1619 .plain => 80,
1620 .tls => 443,
1621 };
1622}1621}
16231622
1624/// Open a connection to the host specified by `uri` and prepare to send a HTTP request.1623/// Open a connection to the host specified by `uri` and prepare to send a HTTP request.
1625///1624///
1626/// `uri` must remain alive during the entire request.
1627///
1628/// The caller is responsible for calling `deinit()` on the `Request`.1625/// The caller is responsible for calling `deinit()` on the `Request`.
1629/// This function is threadsafe.1626/// This function is threadsafe.
1630///1627///
1631/// Asserts that "\r\n" does not occur in any header name or value.1628/// Asserts that "\r\n" does not occur in any header name or value.
1632pub fn open(1629pub fn request(
1633 client: *Client,1630 client: *Client,
1634 method: http.Method,1631 method: http.Method,
1635 uri: Uri,1632 uri: Uri,
...@@ -1649,59 +1646,58 @@ pub fn open(...@@ -1649,59 +1646,58 @@ pub fn open(
1649 }1646 }
1650 }1647 }
16511648
1652 var server_header: std.heap.FixedBufferAllocator = .init(options.server_header_buffer);1649 const protocol = Protocol.fromUri(uri) orelse return error.UnsupportedUriScheme;
1653 const protocol, const valid_uri = try validateUri(uri, server_header.allocator());
16541650
1655 if (protocol == .tls and @atomicLoad(bool, &client.next_https_rescan_certs, .acquire)) {1651 if (protocol == .tls) {
1656 if (disable_tls) unreachable;1652 if (disable_tls) unreachable;
16571653 if (@atomicLoad(bool, &client.next_https_rescan_certs, .acquire)) {
1658 client.ca_bundle_mutex.lock();1654 client.ca_bundle_mutex.lock();
1659 defer client.ca_bundle_mutex.unlock();1655 defer client.ca_bundle_mutex.unlock();
16601656
1661 if (client.next_https_rescan_certs) {1657 if (client.next_https_rescan_certs) {
1662 client.ca_bundle.rescan(client.allocator) catch1658 client.ca_bundle.rescan(client.allocator) catch
1663 return error.CertificateBundleLoadFailure;1659 return error.CertificateBundleLoadFailure;
1664 @atomicStore(bool, &client.next_https_rescan_certs, false, .release);1660 @atomicStore(bool, &client.next_https_rescan_certs, false, .release);
1661 }
1665 }1662 }
1666 }1663 }
16671664
1668 const conn = options.connection orelse1665 const connection = options.connection orelse c: {
1669 try client.connect(valid_uri.host.?.raw, uriPort(valid_uri, protocol), protocol);1666 var host_name_buffer: [Uri.host_name_max]u8 = undefined;
1667 const host_name = try uri.getHost(&host_name_buffer);
1668 break :c try client.connect(host_name, uriPort(uri, protocol), protocol);
1669 };
16701670
1671 var req: Request = .{1671 return .{
1672 .uri = valid_uri,1672 .uri = uri,
1673 .client = client,1673 .client = client,
1674 .connection = conn,1674 .connection = connection,
1675 .reader = .{
1676 .in = connection.reader(),
1677 .state = .ready,
1678 // Populated when `http.Reader.bodyReader` is called.
1679 .interface = undefined,
1680 },
1675 .keep_alive = options.keep_alive,1681 .keep_alive = options.keep_alive,
1676 .method = method,1682 .method = method,
1677 .version = options.version,1683 .version = options.version,
1678 .transfer_encoding = .none,1684 .transfer_encoding = .none,
1679 .redirect_behavior = options.redirect_behavior,1685 .redirect_behavior = options.redirect_behavior,
1680 .handle_continue = options.handle_continue,1686 .handle_continue = options.handle_continue,
1681 .response = .{
1682 .version = undefined,
1683 .status = undefined,
1684 .reason = undefined,
1685 .keep_alive = undefined,
1686 .parser = .init(server_header.buffer[server_header.end_index..]),
1687 },
1688 .headers = options.headers,1687 .headers = options.headers,
1689 .extra_headers = options.extra_headers,1688 .extra_headers = options.extra_headers,
1690 .privileged_headers = options.privileged_headers,1689 .privileged_headers = options.privileged_headers,
1691 };1690 };
1692 errdefer req.deinit();
1693
1694 return req;
1695}1691}
16961692
1697pub const FetchOptions = struct {1693pub const FetchOptions = struct {
1698 server_header_buffer: ?[]u8 = null,1694 /// `null` means it will be heap-allocated.
1695 redirect_buffer: ?[]u8 = null,
1696 /// `null` means it will be heap-allocated.
1697 decompress_buffer: ?[]u8 = null,
1699 redirect_behavior: ?Request.RedirectBehavior = null,1698 redirect_behavior: ?Request.RedirectBehavior = null,
17001699 /// If the server sends a body, it will be stored here.
1701 /// If the server sends a body, it will be appended to this ArrayList.1700 response_storage: ?ResponseStorage = null,
1702 /// `max_append_size` provides an upper limit for how much they can grow.
1703 response_storage: ResponseStorage = .ignore,
1704 max_append_size: ?usize = null,
17051701
1706 location: Location,1702 location: Location,
1707 method: ?http.Method = null,1703 method: ?http.Method = null,
...@@ -1725,11 +1721,11 @@ pub const FetchOptions = struct {...@@ -1725,11 +1721,11 @@ pub const FetchOptions = struct {
1725 uri: Uri,1721 uri: Uri,
1726 };1722 };
17271723
1728 pub const ResponseStorage = union(enum) {1724 pub const ResponseStorage = struct {
1729 ignore,1725 list: *std.ArrayListUnmanaged(u8),
1730 /// Only the existing capacity will be used.1726 /// If null then only the existing capacity will be used.
1731 static: *std.ArrayListUnmanaged(u8),1727 allocator: ?Allocator = null,
1732 dynamic: *std.ArrayList(u8),1728 append_limit: std.io.Limit = .unlimited,
1733 };1729 };
1734};1730};
17351731
...@@ -1737,23 +1733,29 @@ pub const FetchResult = struct {...@@ -1737,23 +1733,29 @@ pub const FetchResult = struct {
1737 status: http.Status,1733 status: http.Status,
1738};1734};
17391735
1736pub const FetchError = Uri.ParseError || RequestError || Request.ReceiveHeadError || error{
1737 StreamTooLong,
1738 /// TODO provide optional diagnostics when this occurs or break into more error codes
1739 WriteFailed,
1740 UnsupportedCompressionMethod,
1741};
1742
1740/// Perform a one-shot HTTP request with the provided options.1743/// Perform a one-shot HTTP request with the provided options.
1741///1744///
1742/// This function is threadsafe.1745/// This function is threadsafe.
1743pub fn fetch(client: *Client, options: FetchOptions) !FetchResult {1746pub fn fetch(client: *Client, options: FetchOptions) FetchError!FetchResult {
1744 const uri = switch (options.location) {1747 const uri = switch (options.location) {
1745 .url => |u| try Uri.parse(u),1748 .url => |u| try Uri.parse(u),
1746 .uri => |u| u,1749 .uri => |u| u,
1747 };1750 };
1748 var server_header_buffer: [16 * 1024]u8 = undefined;
1749
1750 const method: http.Method = options.method orelse1751 const method: http.Method = options.method orelse
1751 if (options.payload != null) .POST else .GET;1752 if (options.payload != null) .POST else .GET;
17521753
1753 var req = try open(client, method, uri, .{1754 const redirect_behavior: Request.RedirectBehavior = options.redirect_behavior orelse
1754 .server_header_buffer = options.server_header_buffer orelse &server_header_buffer,1755 if (options.payload == null) @enumFromInt(3) else .unhandled;
1755 .redirect_behavior = options.redirect_behavior orelse1756
1756 if (options.payload == null) @enumFromInt(3) else .unhandled,1757 var req = try request(client, method, uri, .{
1758 .redirect_behavior = redirect_behavior,
1757 .headers = options.headers,1759 .headers = options.headers,
1758 .extra_headers = options.extra_headers,1760 .extra_headers = options.extra_headers,
1759 .privileged_headers = options.privileged_headers,1761 .privileged_headers = options.privileged_headers,
...@@ -1761,44 +1763,70 @@ pub fn fetch(client: *Client, options: FetchOptions) !FetchResult {...@@ -1761,44 +1763,70 @@ pub fn fetch(client: *Client, options: FetchOptions) !FetchResult {
1761 });1763 });
1762 defer req.deinit();1764 defer req.deinit();
17631765
1764 if (options.payload) |payload| req.transfer_encoding = .{ .content_length = payload.len };1766 if (options.payload) |payload| {
1767 req.transfer_encoding = .{ .content_length = payload.len };
1768 var body = try req.sendBody(&.{});
1769 try body.writer.writeAll(payload);
1770 try body.end();
1771 } else {
1772 try req.sendBodiless();
1773 }
17651774
1766 try req.send();1775 const redirect_buffer: []u8 = if (redirect_behavior == .unhandled) &.{} else options.redirect_buffer orelse
1776 try client.allocator.alloc(u8, 8 * 1024);
1777 defer if (options.redirect_buffer == null) client.allocator.free(redirect_buffer);
17671778
1768 if (options.payload) |payload| try req.writeAll(payload);1779 var response = try req.receiveHead(redirect_buffer);
17691780
1770 try req.finish();1781 const storage = options.response_storage orelse {
1771 try req.wait();1782 const reader = response.reader(&.{});
1783 _ = reader.discardRemaining() catch |err| switch (err) {
1784 error.ReadFailed => return response.bodyErr().?,
1785 };
1786 return .{ .status = response.head.status };
1787 };
17721788
1773 switch (options.response_storage) {1789 const decompress_buffer: []u8 = switch (response.head.content_encoding) {
1774 .ignore => {1790 .identity => &.{},
1775 // Take advantage of request internals to discard the response body1791 .zstd => options.decompress_buffer orelse try client.allocator.alloc(u8, std.compress.zstd.default_window_len),
1776 // and make the connection available for another request.1792 .deflate, .gzip => options.decompress_buffer orelse try client.allocator.alloc(u8, std.compress.flate.max_window_len),
1777 req.response.skip = true;1793 .compress => return error.UnsupportedCompressionMethod,
1778 assert(try req.transferRead(&.{}) == 0); // No buffer is necessary when skipping.1794 };
1779 },1795 defer if (options.decompress_buffer == null) client.allocator.free(decompress_buffer);
1780 .dynamic => |list| {1796
1781 const max_append_size = options.max_append_size orelse 2 * 1024 * 1024;1797 var decompressor: http.Decompressor = undefined;
1782 try req.reader().readAllArrayList(list, max_append_size);1798 const reader = response.readerDecompressing(&decompressor, decompress_buffer);
1783 },1799 const list = storage.list;
1784 .static => |list| {1800
1785 const buf = b: {1801 if (storage.allocator) |allocator| {
1786 const buf = list.unusedCapacitySlice();1802 reader.appendRemaining(allocator, null, list, storage.append_limit) catch |err| switch (err) {
1787 if (options.max_append_size) |len| {1803 error.ReadFailed => return response.bodyErr().?,
1788 if (len < buf.len) break :b buf[0..len];1804 else => |e| return e,
1789 }1805 };
1790 break :b buf;1806 } else {
1791 };1807 const buf = storage.append_limit.slice(list.unusedCapacitySlice());
1792 list.items.len += try req.reader().readAll(buf);1808 list.items.len += reader.readSliceShort(buf) catch |err| switch (err) {
1793 },1809 error.ReadFailed => return response.bodyErr().?,
1810 };
1794 }1811 }
17951812
1796 return .{1813 return .{ .status = response.head.status };
1797 .status = req.response.status,1814}
1798 };1815
1816pub fn sameParentDomain(parent_host: []const u8, child_host: []const u8) bool {
1817 if (!std.ascii.endsWithIgnoreCase(child_host, parent_host)) return false;
1818 if (child_host.len == parent_host.len) return true;
1819 if (parent_host.len > child_host.len) return false;
1820 return child_host[child_host.len - parent_host.len - 1] == '.';
1821}
1822
1823test sameParentDomain {
1824 try testing.expect(!sameParentDomain("foo.com", "bar.com"));
1825 try testing.expect(sameParentDomain("foo.com", "foo.com"));
1826 try testing.expect(sameParentDomain("foo.com", "bar.foo.com"));
1827 try testing.expect(!sameParentDomain("bar.foo.com", "foo.com"));
1799}1828}
18001829
1801test {1830test {
1802 _ = Response;1831 _ = Response;
1803 _ = &initDefaultProxies;
1804}1832}
lib/std/http/Server.zig+406-751
...@@ -1,139 +1,69 @@...@@ -1,139 +1,69 @@
1//! Blocking HTTP server implementation.1//! Handles a single connection lifecycle.
2//! Handles a single connection's lifecycle.2
33const std = @import("../std.zig");
4connection: net.Server.Connection,4const http = std.http;
5/// Keeps track of whether the Server is ready to accept a new request on the5const mem = std.mem;
6/// same connection, and makes invalid API usage cause assertion failures6const Uri = std.Uri;
7/// rather than HTTP protocol violations.7const assert = std.debug.assert;
8state: State,8const testing = std.testing;
9/// User-provided buffer that must outlive this Server.9const Writer = std.Io.Writer;
10/// Used to store the client's entire HTTP header.10const Reader = std.Io.Reader;
11read_buffer: []u8,11
12/// Amount of available data inside read_buffer.12const Server = @This();
13read_buffer_len: usize,13
14/// Index into `read_buffer` of the first byte of the next HTTP request.14/// Data from the HTTP server to the HTTP client.
15next_request_start: usize,15out: *Writer,
1616reader: http.Reader,
17pub const State = enum {
18 /// The connection is available to be used for the first time, or reused.
19 ready,
20 /// An error occurred in `receiveHead`.
21 receiving_head,
22 /// A Request object has been obtained and from there a Response can be
23 /// opened.
24 received_head,
25 /// The client is uploading something to this Server.
26 receiving_body,
27 /// The connection is eligible for another HTTP request, however the client
28 /// and server did not negotiate a persistent connection.
29 closing,
30};
3117
32/// Initialize an HTTP server that can respond to multiple requests on the same18/// Initialize an HTTP server that can respond to multiple requests on the same
33/// connection.19/// connection.
20///
21/// The buffer of `in` must be large enough to store the client's entire HTTP
22/// header, otherwise `receiveHead` returns `error.HttpHeadersOversize`.
23///
34/// The returned `Server` is ready for `receiveHead` to be called.24/// The returned `Server` is ready for `receiveHead` to be called.
35pub fn init(connection: net.Server.Connection, read_buffer: []u8) Server {25pub fn init(in: *Reader, out: *Writer) Server {
36 return .{26 return .{
37 .connection = connection,27 .reader = .{
38 .state = .ready,28 .in = in,
39 .read_buffer = read_buffer,29 .state = .ready,
40 .read_buffer_len = 0,30 // Populated when `http.Reader.bodyReader` is called.
41 .next_request_start = 0,31 .interface = undefined,
32 },
33 .out = out,
42 };34 };
43}35}
4436
45pub const ReceiveHeadError = error{37pub const ReceiveHeadError = http.Reader.HeadError || error{
46 /// Client sent too many bytes of HTTP headers.
47 /// The HTTP specification suggests to respond with a 431 status code
48 /// before closing the connection.
49 HttpHeadersOversize,
50 /// Client sent headers that did not conform to the HTTP protocol.38 /// Client sent headers that did not conform to the HTTP protocol.
39 ///
40 /// To find out more detailed diagnostics, `Request.head_buffer` can be
41 /// passed directly to `Request.Head.parse`.
51 HttpHeadersInvalid,42 HttpHeadersInvalid,
52 /// A low level I/O error occurred trying to read the headers.
53 HttpHeadersUnreadable,
54 /// Partial HTTP request was received but the connection was closed before
55 /// fully receiving the headers.
56 HttpRequestTruncated,
57 /// The client sent 0 bytes of headers before closing the stream.
58 /// In other words, a keep-alive connection was finally closed.
59 HttpConnectionClosing,
60};43};
6144
62/// The header bytes reference the read buffer that Server was initialized with
63/// and remain alive until the next call to receiveHead.
64pub fn receiveHead(s: *Server) ReceiveHeadError!Request {45pub fn receiveHead(s: *Server) ReceiveHeadError!Request {
65 assert(s.state == .ready);46 const head_buffer = try s.reader.receiveHead();
66 s.state = .received_head;
67 errdefer s.state = .receiving_head;
68
69 // In case of a reused connection, move the next request's bytes to the
70 // beginning of the buffer.
71 if (s.next_request_start > 0) {
72 if (s.read_buffer_len > s.next_request_start) {
73 rebase(s, 0);
74 } else {
75 s.read_buffer_len = 0;
76 }
77 }
78
79 var hp: http.HeadParser = .{};
80
81 if (s.read_buffer_len > 0) {
82 const bytes = s.read_buffer[0..s.read_buffer_len];
83 const end = hp.feed(bytes);
84 if (hp.state == .finished)
85 return finishReceivingHead(s, end);
86 }
87
88 while (true) {
89 const buf = s.read_buffer[s.read_buffer_len..];
90 if (buf.len == 0)
91 return error.HttpHeadersOversize;
92 const read_n = s.connection.stream.read(buf) catch
93 return error.HttpHeadersUnreadable;
94 if (read_n == 0) {
95 if (s.read_buffer_len > 0) {
96 return error.HttpRequestTruncated;
97 } else {
98 return error.HttpConnectionClosing;
99 }
100 }
101 s.read_buffer_len += read_n;
102 const bytes = buf[0..read_n];
103 const end = hp.feed(bytes);
104 if (hp.state == .finished)
105 return finishReceivingHead(s, s.read_buffer_len - bytes.len + end);
106 }
107}
108
109fn finishReceivingHead(s: *Server, head_end: usize) ReceiveHeadError!Request {
110 return .{47 return .{
111 .server = s,48 .server = s,
112 .head_end = head_end,49 .head_buffer = head_buffer,
113 .head = Request.Head.parse(s.read_buffer[0..head_end]) catch50 // No need to track the returned error here since users can repeat the
114 return error.HttpHeadersInvalid,51 // parse with the header buffer to get detailed diagnostics.
115 .reader_state = undefined,52 .head = Request.Head.parse(head_buffer) catch return error.HttpHeadersInvalid,
116 };53 };
117}54}
11855
119pub const Request = struct {56pub const Request = struct {
120 server: *Server,57 server: *Server,
121 /// Index into Server's read_buffer.58 /// Pointers in this struct are invalidated when the request body stream is
122 head_end: usize,59 /// initialized.
123 head: Head,60 head: Head,
124 reader_state: union {61 head_buffer: []const u8,
125 remaining_content_length: u64,62 respond_err: ?RespondError = null,
126 chunk_parser: http.ChunkParser,63
127 },64 pub const RespondError = error{
12865 /// The request contained an `expect` header with an unrecognized value.
129 pub const Compression = union(enum) {66 HttpExpectationFailed,
130 pub const DeflateDecompressor = std.compress.zlib.Decompressor(std.io.AnyReader);
131 pub const GzipDecompressor = std.compress.gzip.Decompressor(std.io.AnyReader);
132
133 deflate: std.compress.flate.Decompress,
134 gzip: std.compress.flate.Decompress,
135 zstd: std.compress.zstd.Decompress,
136 none: void,
137 };67 };
13868
139 pub const Head = struct {69 pub const Head = struct {
...@@ -146,7 +76,6 @@ pub const Request = struct {...@@ -146,7 +76,6 @@ pub const Request = struct {
146 transfer_encoding: http.TransferEncoding,76 transfer_encoding: http.TransferEncoding,
147 transfer_compression: http.ContentEncoding,77 transfer_compression: http.ContentEncoding,
148 keep_alive: bool,78 keep_alive: bool,
149 compression: Compression,
15079
151 pub const ParseError = error{80 pub const ParseError = error{
152 UnknownHttpMethod,81 UnknownHttpMethod,
...@@ -168,10 +97,9 @@ pub const Request = struct {...@@ -168,10 +97,9 @@ pub const Request = struct {
16897
169 const method_end = mem.indexOfScalar(u8, first_line, ' ') orelse98 const method_end = mem.indexOfScalar(u8, first_line, ' ') orelse
170 return error.HttpHeadersInvalid;99 return error.HttpHeadersInvalid;
171 if (method_end > 24) return error.HttpHeadersInvalid;
172100
173 const method_str = first_line[0..method_end];101 const method = std.meta.stringToEnum(http.Method, first_line[0..method_end]) orelse
174 const method: http.Method = @enumFromInt(http.Method.parse(method_str));102 return error.UnknownHttpMethod;
175103
176 const version_start = mem.lastIndexOfScalar(u8, first_line, ' ') orelse104 const version_start = mem.lastIndexOfScalar(u8, first_line, ' ') orelse
177 return error.HttpHeadersInvalid;105 return error.HttpHeadersInvalid;
...@@ -200,7 +128,6 @@ pub const Request = struct {...@@ -200,7 +128,6 @@ pub const Request = struct {
200 .@"HTTP/1.0" => false,128 .@"HTTP/1.0" => false,
201 .@"HTTP/1.1" => true,129 .@"HTTP/1.1" => true,
202 },130 },
203 .compression = .none,
204 };131 };
205132
206 while (it.next()) |line| {133 while (it.next()) |line| {
...@@ -230,7 +157,7 @@ pub const Request = struct {...@@ -230,7 +157,7 @@ pub const Request = struct {
230157
231 const trimmed = mem.trim(u8, header_value, " ");158 const trimmed = mem.trim(u8, header_value, " ");
232159
233 if (std.meta.stringToEnum(http.ContentEncoding, trimmed)) |ce| {160 if (http.ContentEncoding.fromString(trimmed)) |ce| {
234 head.transfer_compression = ce;161 head.transfer_compression = ce;
235 } else {162 } else {
236 return error.HttpTransferEncodingUnsupported;163 return error.HttpTransferEncodingUnsupported;
...@@ -255,7 +182,7 @@ pub const Request = struct {...@@ -255,7 +182,7 @@ pub const Request = struct {
255 if (next) |second| {182 if (next) |second| {
256 const trimmed_second = mem.trim(u8, second, " ");183 const trimmed_second = mem.trim(u8, second, " ");
257184
258 if (std.meta.stringToEnum(http.ContentEncoding, trimmed_second)) |transfer| {185 if (http.ContentEncoding.fromString(trimmed_second)) |transfer| {
259 if (head.transfer_compression != .identity)186 if (head.transfer_compression != .identity)
260 return error.HttpHeadersInvalid; // double compression is not supported187 return error.HttpHeadersInvalid; // double compression is not supported
261 head.transfer_compression = transfer;188 head.transfer_compression = transfer;
...@@ -296,10 +223,19 @@ pub const Request = struct {...@@ -296,10 +223,19 @@ pub const Request = struct {
296 inline fn int64(array: *const [8]u8) u64 {223 inline fn int64(array: *const [8]u8) u64 {
297 return @bitCast(array.*);224 return @bitCast(array.*);
298 }225 }
226
227 /// Help the programmer avoid bugs by calling this when the string
228 /// memory of `Head` becomes invalidated.
229 fn invalidateStrings(h: *Head) void {
230 h.target = undefined;
231 if (h.expect) |*s| s.* = undefined;
232 if (h.content_type) |*s| s.* = undefined;
233 }
299 };234 };
300235
301 pub fn iterateHeaders(r: *Request) http.HeaderIterator {236 pub fn iterateHeaders(r: *const Request) http.HeaderIterator {
302 return http.HeaderIterator.init(r.server.read_buffer[0..r.head_end]);237 assert(r.server.reader.state == .received_head);
238 return http.HeaderIterator.init(r.head_buffer);
303 }239 }
304240
305 test iterateHeaders {241 test iterateHeaders {
...@@ -310,22 +246,19 @@ pub const Request = struct {...@@ -310,22 +246,19 @@ pub const Request = struct {
310 "TRansfer-encoding:\tdeflate, chunked \r\n" ++246 "TRansfer-encoding:\tdeflate, chunked \r\n" ++
311 "connectioN:\t keep-alive \r\n\r\n";247 "connectioN:\t keep-alive \r\n\r\n";
312248
313 var read_buffer: [500]u8 = undefined;
314 @memcpy(read_buffer[0..request_bytes.len], request_bytes);
315
316 var server: Server = .{249 var server: Server = .{
317 .connection = undefined,250 .reader = .{
318 .state = .ready,251 .in = undefined,
319 .read_buffer = &read_buffer,252 .state = .received_head,
320 .read_buffer_len = request_bytes.len,253 .interface = undefined,
321 .next_request_start = 0,254 },
255 .out = undefined,
322 };256 };
323257
324 var request: Request = .{258 var request: Request = .{
325 .server = &server,259 .server = &server,
326 .head_end = request_bytes.len,
327 .head = undefined,260 .head = undefined,
328 .reader_state = undefined,261 .head_buffer = @constCast(request_bytes),
329 };262 };
330263
331 var it = request.iterateHeaders();264 var it = request.iterateHeaders();
...@@ -384,16 +317,22 @@ pub const Request = struct {...@@ -384,16 +317,22 @@ pub const Request = struct {
384 /// no error is surfaced.317 /// no error is surfaced.
385 ///318 ///
386 /// Asserts status is not `continue`.319 /// Asserts status is not `continue`.
387 /// Asserts there are at most 25 extra_headers.
388 /// Asserts that "\r\n" does not occur in any header name or value.320 /// Asserts that "\r\n" does not occur in any header name or value.
389 pub fn respond(321 pub fn respond(
390 request: *Request,322 request: *Request,
391 content: []const u8,323 content: []const u8,
392 options: RespondOptions,324 options: RespondOptions,
393 ) Response.WriteError!void {325 ) ExpectContinueError!void {
394 const max_extra_headers = 25;326 try respondUnflushed(request, content, options);
327 try request.server.out.flush();
328 }
329
330 pub fn respondUnflushed(
331 request: *Request,
332 content: []const u8,
333 options: RespondOptions,
334 ) ExpectContinueError!void {
395 assert(options.status != .@"continue");335 assert(options.status != .@"continue");
396 assert(options.extra_headers.len <= max_extra_headers);
397 if (std.debug.runtime_safety) {336 if (std.debug.runtime_safety) {
398 for (options.extra_headers) |header| {337 for (options.extra_headers) |header| {
399 assert(header.name.len != 0);338 assert(header.name.len != 0);
...@@ -402,6 +341,7 @@ pub const Request = struct {...@@ -402,6 +341,7 @@ pub const Request = struct {
402 assert(std.mem.indexOfPosLinear(u8, header.value, 0, "\r\n") == null);341 assert(std.mem.indexOfPosLinear(u8, header.value, 0, "\r\n") == null);
403 }342 }
404 }343 }
344 try writeExpectContinue(request);
405345
406 const transfer_encoding_none = (options.transfer_encoding orelse .chunked) == .none;346 const transfer_encoding_none = (options.transfer_encoding orelse .chunked) == .none;
407 const server_keep_alive = !transfer_encoding_none and options.keep_alive;347 const server_keep_alive = !transfer_encoding_none and options.keep_alive;
...@@ -409,130 +349,42 @@ pub const Request = struct {...@@ -409,130 +349,42 @@ pub const Request = struct {
409349
410 const phrase = options.reason orelse options.status.phrase() orelse "";350 const phrase = options.reason orelse options.status.phrase() orelse "";
411351
412 var first_buffer: [500]u8 = undefined;352 const out = request.server.out;
413 var h = std.ArrayListUnmanaged(u8).initBuffer(&first_buffer);353 try out.print("{s} {d} {s}\r\n", .{
414 if (request.head.expect != null) {
415 // reader() and hence discardBody() above sets expect to null if it
416 // is handled. So the fact that it is not null here means unhandled.
417 h.appendSliceAssumeCapacity("HTTP/1.1 417 Expectation Failed\r\n");
418 if (!keep_alive) h.appendSliceAssumeCapacity("connection: close\r\n");
419 h.appendSliceAssumeCapacity("content-length: 0\r\n\r\n");
420 try request.server.connection.stream.writeAll(h.items);
421 return;
422 }
423 h.fixedWriter().print("{s} {d} {s}\r\n", .{
424 @tagName(options.version), @intFromEnum(options.status), phrase,354 @tagName(options.version), @intFromEnum(options.status), phrase,
425 }) catch unreachable;355 });
426356
427 switch (options.version) {357 switch (options.version) {
428 .@"HTTP/1.0" => if (keep_alive) h.appendSliceAssumeCapacity("connection: keep-alive\r\n"),358 .@"HTTP/1.0" => if (keep_alive) try out.writeAll("connection: keep-alive\r\n"),
429 .@"HTTP/1.1" => if (!keep_alive) h.appendSliceAssumeCapacity("connection: close\r\n"),359 .@"HTTP/1.1" => if (!keep_alive) try out.writeAll("connection: close\r\n"),
430 }360 }
431361
432 if (options.transfer_encoding) |transfer_encoding| switch (transfer_encoding) {362 if (options.transfer_encoding) |transfer_encoding| switch (transfer_encoding) {
433 .none => {},363 .none => {},
434 .chunked => h.appendSliceAssumeCapacity("transfer-encoding: chunked\r\n"),364 .chunked => try out.writeAll("transfer-encoding: chunked\r\n"),
435 } else {365 } else {
436 h.fixedWriter().print("content-length: {d}\r\n", .{content.len}) catch unreachable;366 try out.print("content-length: {d}\r\n", .{content.len});
437 }367 }
438368
439 var chunk_header_buffer: [18]u8 = undefined;
440 var iovecs: [max_extra_headers * 4 + 3]std.posix.iovec_const = undefined;
441 var iovecs_len: usize = 0;
442
443 iovecs[iovecs_len] = .{
444 .base = h.items.ptr,
445 .len = h.items.len,
446 };
447 iovecs_len += 1;
448
449 for (options.extra_headers) |header| {369 for (options.extra_headers) |header| {
450 iovecs[iovecs_len] = .{370 var vecs: [4][]const u8 = .{ header.name, ": ", header.value, "\r\n" };
451 .base = header.name.ptr,371 try out.writeVecAll(&vecs);
452 .len = header.name.len,
453 };
454 iovecs_len += 1;
455
456 iovecs[iovecs_len] = .{
457 .base = ": ",
458 .len = 2,
459 };
460 iovecs_len += 1;
461
462 if (header.value.len != 0) {
463 iovecs[iovecs_len] = .{
464 .base = header.value.ptr,
465 .len = header.value.len,
466 };
467 iovecs_len += 1;
468 }
469
470 iovecs[iovecs_len] = .{
471 .base = "\r\n",
472 .len = 2,
473 };
474 iovecs_len += 1;
475 }372 }
476373
477 iovecs[iovecs_len] = .{374 try out.writeAll("\r\n");
478 .base = "\r\n",
479 .len = 2,
480 };
481 iovecs_len += 1;
482375
483 if (request.head.method != .HEAD) {376 if (request.head.method != .HEAD) {
484 const is_chunked = (options.transfer_encoding orelse .none) == .chunked;377 const is_chunked = (options.transfer_encoding orelse .none) == .chunked;
485 if (is_chunked) {378 if (is_chunked) {
486 if (content.len > 0) {379 if (content.len > 0) try out.print("{x}\r\n{s}\r\n", .{ content.len, content });
487 const chunk_header = std.fmt.bufPrint(380 try out.writeAll("0\r\n\r\n");
488 &chunk_header_buffer,
489 "{x}\r\n",
490 .{content.len},
491 ) catch unreachable;
492
493 iovecs[iovecs_len] = .{
494 .base = chunk_header.ptr,
495 .len = chunk_header.len,
496 };
497 iovecs_len += 1;
498
499 iovecs[iovecs_len] = .{
500 .base = content.ptr,
501 .len = content.len,
502 };
503 iovecs_len += 1;
504
505 iovecs[iovecs_len] = .{
506 .base = "\r\n",
507 .len = 2,
508 };
509 iovecs_len += 1;
510 }
511
512 iovecs[iovecs_len] = .{
513 .base = "0\r\n\r\n",
514 .len = 5,
515 };
516 iovecs_len += 1;
517 } else if (content.len > 0) {381 } else if (content.len > 0) {
518 iovecs[iovecs_len] = .{382 try out.writeAll(content);
519 .base = content.ptr,
520 .len = content.len,
521 };
522 iovecs_len += 1;
523 }383 }
524 }384 }
525
526 try request.server.connection.stream.writevAll(iovecs[0..iovecs_len]);
527 }385 }
528386
529 pub const RespondStreamingOptions = struct {387 pub const RespondStreamingOptions = struct {
530 /// An externally managed slice of memory used to batch bytes before
531 /// sending. `respondStreaming` asserts this is large enough to store
532 /// the full HTTP response head.
533 ///
534 /// Must outlive the returned Response.
535 send_buffer: []u8,
536 /// If provided, the response will use the content-length header;388 /// If provided, the response will use the content-length header;
537 /// otherwise it will use transfer-encoding: chunked.389 /// otherwise it will use transfer-encoding: chunked.
538 content_length: ?u64 = null,390 content_length: ?u64 = null,
...@@ -540,254 +392,227 @@ pub const Request = struct {...@@ -540,254 +392,227 @@ pub const Request = struct {
540 respond_options: RespondOptions = .{},392 respond_options: RespondOptions = .{},
541 };393 };
542394
543 /// The header is buffered but not sent until Response.flush is called.395 /// The header is not guaranteed to be sent until `BodyWriter.flush` or
396 /// `BodyWriter.end` is called.
544 ///397 ///
545 /// If the request contains a body and the connection is to be reused,398 /// If the request contains a body and the connection is to be reused,
546 /// discards the request body, leaving the Server in the `ready` state. If399 /// discards the request body, leaving the Server in the `ready` state. If
547 /// this discarding fails, the connection is marked as not to be reused and400 /// this discarding fails, the connection is marked as not to be reused and
548 /// no error is surfaced.401 /// no error is surfaced.
549 ///402 ///
550 /// HEAD requests are handled transparently by setting a flag on the403 /// HEAD requests are handled transparently by setting the
551 /// returned Response to omit the body. However it may be worth noticing404 /// `BodyWriter.elide` flag on the returned `BodyWriter`, causing
405 /// the response stream to omit the body. However, it may be worth noticing
552 /// that flag and skipping any expensive work that would otherwise need to406 /// that flag and skipping any expensive work that would otherwise need to
553 /// be done to satisfy the request.407 /// be done to satisfy the request.
554 ///408 ///
555 /// Asserts `send_buffer` is large enough to store the entire response header.
556 /// Asserts status is not `continue`.409 /// Asserts status is not `continue`.
557 pub fn respondStreaming(request: *Request, options: RespondStreamingOptions) Response {410 pub fn respondStreaming(
411 request: *Request,
412 buffer: []u8,
413 options: RespondStreamingOptions,
414 ) ExpectContinueError!http.BodyWriter {
415 try writeExpectContinue(request);
558 const o = options.respond_options;416 const o = options.respond_options;
559 assert(o.status != .@"continue");417 assert(o.status != .@"continue");
560 const transfer_encoding_none = (o.transfer_encoding orelse .chunked) == .none;418 const transfer_encoding_none = (o.transfer_encoding orelse .chunked) == .none;
561 const server_keep_alive = !transfer_encoding_none and o.keep_alive;419 const server_keep_alive = !transfer_encoding_none and o.keep_alive;
562 const keep_alive = request.discardBody(server_keep_alive);420 const keep_alive = request.discardBody(server_keep_alive);
563 const phrase = o.reason orelse o.status.phrase() orelse "";421 const phrase = o.reason orelse o.status.phrase() orelse "";
422 const out = request.server.out;
564423
565 var h = std.ArrayListUnmanaged(u8).initBuffer(options.send_buffer);424 try out.print("{s} {d} {s}\r\n", .{
566425 @tagName(o.version), @intFromEnum(o.status), phrase,
567 const elide_body = if (request.head.expect != null) eb: {426 });
568 // reader() and hence discardBody() above sets expect to null if it
569 // is handled. So the fact that it is not null here means unhandled.
570 h.appendSliceAssumeCapacity("HTTP/1.1 417 Expectation Failed\r\n");
571 if (!keep_alive) h.appendSliceAssumeCapacity("connection: close\r\n");
572 h.appendSliceAssumeCapacity("content-length: 0\r\n\r\n");
573 break :eb true;
574 } else eb: {
575 h.fixedWriter().print("{s} {d} {s}\r\n", .{
576 @tagName(o.version), @intFromEnum(o.status), phrase,
577 }) catch unreachable;
578
579 switch (o.version) {
580 .@"HTTP/1.0" => if (keep_alive) h.appendSliceAssumeCapacity("connection: keep-alive\r\n"),
581 .@"HTTP/1.1" => if (!keep_alive) h.appendSliceAssumeCapacity("connection: close\r\n"),
582 }
583427
584 if (o.transfer_encoding) |transfer_encoding| switch (transfer_encoding) {428 switch (o.version) {
585 .chunked => h.appendSliceAssumeCapacity("transfer-encoding: chunked\r\n"),429 .@"HTTP/1.0" => if (keep_alive) try out.writeAll("connection: keep-alive\r\n"),
586 .none => {},430 .@"HTTP/1.1" => if (!keep_alive) try out.writeAll("connection: close\r\n"),
587 } else if (options.content_length) |len| {431 }
588 h.fixedWriter().print("content-length: {d}\r\n", .{len}) catch unreachable;
589 } else {
590 h.appendSliceAssumeCapacity("transfer-encoding: chunked\r\n");
591 }
592432
593 for (o.extra_headers) |header| {433 if (o.transfer_encoding) |transfer_encoding| switch (transfer_encoding) {
594 assert(header.name.len != 0);434 .chunked => try out.writeAll("transfer-encoding: chunked\r\n"),
595 h.appendSliceAssumeCapacity(header.name);435 .none => {},
596 h.appendSliceAssumeCapacity(": ");436 } else if (options.content_length) |len| {
597 h.appendSliceAssumeCapacity(header.value);437 try out.print("content-length: {d}\r\n", .{len});
598 h.appendSliceAssumeCapacity("\r\n");438 } else {
599 }439 try out.writeAll("transfer-encoding: chunked\r\n");
440 }
600441
601 h.appendSliceAssumeCapacity("\r\n");442 for (o.extra_headers) |header| {
602 break :eb request.head.method == .HEAD;443 assert(header.name.len != 0);
603 };444 var bufs: [4][]const u8 = .{ header.name, ": ", header.value, "\r\n" };
445 try out.writeVecAll(&bufs);
446 }
604447
605 return .{448 try out.writeAll("\r\n");
606 .stream = request.server.connection.stream,449 const elide_body = request.head.method == .HEAD;
607 .send_buffer = options.send_buffer,450 const state: http.BodyWriter.State = if (o.transfer_encoding) |te| switch (te) {
608 .send_buffer_start = 0,451 .chunked => .{ .chunked = .init },
609 .send_buffer_end = h.items.len,452 .none => .none,
610 .transfer_encoding = if (o.transfer_encoding) |te| switch (te) {453 } else if (options.content_length) |len| .{
611 .chunked => .chunked,454 .content_length = len,
612 .none => .none,455 } else .{ .chunked = .init };
613 } else if (options.content_length) |len| .{456
614 .content_length = len,457 return if (elide_body) .{
615 } else .chunked,458 .http_protocol_output = request.server.out,
616 .elide_body = elide_body,459 .state = state,
617 .chunk_len = 0,460 .writer = .{
461 .buffer = buffer,
462 .vtable = &.{
463 .drain = http.BodyWriter.elidingDrain,
464 .sendFile = http.BodyWriter.elidingSendFile,
465 },
466 },
467 } else .{
468 .http_protocol_output = request.server.out,
469 .state = state,
470 .writer = .{
471 .buffer = buffer,
472 .vtable = switch (state) {
473 .none => &.{
474 .drain = http.BodyWriter.noneDrain,
475 .sendFile = http.BodyWriter.noneSendFile,
476 },
477 .content_length => &.{
478 .drain = http.BodyWriter.contentLengthDrain,
479 .sendFile = http.BodyWriter.contentLengthSendFile,
480 },
481 .chunked => &.{
482 .drain = http.BodyWriter.chunkedDrain,
483 .sendFile = http.BodyWriter.chunkedSendFile,
484 },
485 .end => unreachable,
486 },
487 },
618 };488 };
619 }489 }
620490
621 pub const ReadError = net.Stream.ReadError || error{491 pub const UpgradeRequest = union(enum) {
622 HttpChunkInvalid,492 websocket: ?[]const u8,
623 HttpHeadersOversize,493 other: []const u8,
494 none,
624 };495 };
625496
626 fn read_cl(context: *const anyopaque, buffer: []u8) ReadError!usize {497 /// Does not invalidate `request.head`.
627 const request: *Request = @ptrCast(@alignCast(@constCast(context)));498 pub fn upgradeRequested(request: *const Request) UpgradeRequest {
628 const s = request.server;499 switch (request.head.version) {
629500 .@"HTTP/1.0" => return .none,
630 const remaining_content_length = &request.reader_state.remaining_content_length;501 .@"HTTP/1.1" => if (request.head.method != .GET) return .none,
631 if (remaining_content_length.* == 0) {
632 s.state = .ready;
633 return 0;
634 }502 }
635 assert(s.state == .receiving_body);
636 const available = try fill(s, request.head_end);
637 const len = @min(remaining_content_length.*, available.len, buffer.len);
638 @memcpy(buffer[0..len], available[0..len]);
639 remaining_content_length.* -= len;
640 s.next_request_start += len;
641 if (remaining_content_length.* == 0)
642 s.state = .ready;
643 return len;
644 }
645
646 fn fill(s: *Server, head_end: usize) ReadError![]u8 {
647 const available = s.read_buffer[s.next_request_start..s.read_buffer_len];
648 if (available.len > 0) return available;
649 s.next_request_start = head_end;
650 s.read_buffer_len = head_end + try s.connection.stream.read(s.read_buffer[head_end..]);
651 return s.read_buffer[head_end..s.read_buffer_len];
652 }
653503
654 fn read_chunked(context: *const anyopaque, buffer: []u8) ReadError!usize {504 var sec_websocket_key: ?[]const u8 = null;
655 const request: *Request = @ptrCast(@alignCast(@constCast(context)));505 var upgrade_name: ?[]const u8 = null;
656 const s = request.server;506 var it = request.iterateHeaders();
657507 while (it.next()) |header| {
658 const cp = &request.reader_state.chunk_parser;508 if (std.ascii.eqlIgnoreCase(header.name, "sec-websocket-key")) {
659 const head_end = request.head_end;509 sec_websocket_key = header.value;
660510 } else if (std.ascii.eqlIgnoreCase(header.name, "upgrade")) {
661 // Protect against returning 0 before the end of stream.511 upgrade_name = header.value;
662 var out_end: usize = 0;
663 while (out_end == 0) {
664 switch (cp.state) {
665 .invalid => return 0,
666 .data => {
667 assert(s.state == .receiving_body);
668 const available = try fill(s, head_end);
669 const len = @min(cp.chunk_len, available.len, buffer.len);
670 @memcpy(buffer[0..len], available[0..len]);
671 cp.chunk_len -= len;
672 if (cp.chunk_len == 0)
673 cp.state = .data_suffix;
674 out_end += len;
675 s.next_request_start += len;
676 continue;
677 },
678 else => {
679 assert(s.state == .receiving_body);
680 const available = try fill(s, head_end);
681 const n = cp.feed(available);
682 switch (cp.state) {
683 .invalid => return error.HttpChunkInvalid,
684 .data => {
685 if (cp.chunk_len == 0) {
686 // The next bytes in the stream are trailers,
687 // or \r\n to indicate end of chunked body.
688 //
689 // This function must append the trailers at
690 // head_end so that headers and trailers are
691 // together.
692 //
693 // Since returning 0 would indicate end of
694 // stream, this function must read all the
695 // trailers before returning.
696 if (s.next_request_start > head_end) rebase(s, head_end);
697 var hp: http.HeadParser = .{};
698 {
699 const bytes = s.read_buffer[head_end..s.read_buffer_len];
700 const end = hp.feed(bytes);
701 if (hp.state == .finished) {
702 cp.state = .invalid;
703 s.state = .ready;
704 s.next_request_start = s.read_buffer_len - bytes.len + end;
705 return out_end;
706 }
707 }
708 while (true) {
709 const buf = s.read_buffer[s.read_buffer_len..];
710 if (buf.len == 0)
711 return error.HttpHeadersOversize;
712 const read_n = try s.connection.stream.read(buf);
713 s.read_buffer_len += read_n;
714 const bytes = buf[0..read_n];
715 const end = hp.feed(bytes);
716 if (hp.state == .finished) {
717 cp.state = .invalid;
718 s.state = .ready;
719 s.next_request_start = s.read_buffer_len - bytes.len + end;
720 return out_end;
721 }
722 }
723 }
724 const data = available[n..];
725 const len = @min(cp.chunk_len, data.len, buffer.len);
726 @memcpy(buffer[0..len], data[0..len]);
727 cp.chunk_len -= len;
728 if (cp.chunk_len == 0)
729 cp.state = .data_suffix;
730 out_end += len;
731 s.next_request_start += n + len;
732 continue;
733 },
734 else => continue,
735 }
736 },
737 }512 }
738 }513 }
739 return out_end;514
515 const name = upgrade_name orelse return .none;
516 if (std.ascii.eqlIgnoreCase(name, "websocket")) return .{ .websocket = sec_websocket_key };
517 return .{ .other = name };
740 }518 }
741519
742 pub const ReaderError = Response.WriteError || error{520 pub const WebSocketOptions = struct {
743 /// The client sent an expect HTTP header value other than521 /// The value from `UpgradeRequest.websocket` (sec-websocket-key header value).
744 /// "100-continue".522 key: []const u8,
745 HttpExpectationFailed,523 reason: ?[]const u8 = null,
524 extra_headers: []const http.Header = &.{},
746 };525 };
747526
527 /// The header is not guaranteed to be sent until `WebSocket.flush` is
528 /// called on the returned struct.
529 pub fn respondWebSocket(request: *Request, options: WebSocketOptions) ExpectContinueError!WebSocket {
530 if (request.head.expect != null) return error.HttpExpectationFailed;
531
532 const out = request.server.out;
533 const version: http.Version = .@"HTTP/1.1";
534 const status: http.Status = .switching_protocols;
535 const phrase = options.reason orelse status.phrase() orelse "";
536
537 assert(request.head.version == version);
538 assert(request.head.method == .GET);
539
540 var sha1 = std.crypto.hash.Sha1.init(.{});
541 sha1.update(options.key);
542 sha1.update("258EAFA5-E914-47DA-95CA-C5AB0DC85B11");
543 var digest: [std.crypto.hash.Sha1.digest_length]u8 = undefined;
544 sha1.final(&digest);
545 try out.print("{s} {d} {s}\r\n", .{ @tagName(version), @intFromEnum(status), phrase });
546 try out.writeAll("connection: upgrade\r\nupgrade: websocket\r\nsec-websocket-accept: ");
547 const base64_digest = try out.writableArray(28);
548 assert(std.base64.standard.Encoder.encode(base64_digest, &digest).len == base64_digest.len);
549 out.advance(base64_digest.len);
550 try out.writeAll("\r\n");
551
552 for (options.extra_headers) |header| {
553 assert(header.name.len != 0);
554 var bufs: [4][]const u8 = .{ header.name, ": ", header.value, "\r\n" };
555 try out.writeVecAll(&bufs);
556 }
557
558 try out.writeAll("\r\n");
559
560 return .{
561 .input = request.server.reader.in,
562 .output = request.server.out,
563 .key = options.key,
564 };
565 }
566
748 /// In the case that the request contains "expect: 100-continue", this567 /// In the case that the request contains "expect: 100-continue", this
749 /// function writes the continuation header, which means it can fail with a568 /// function writes the continuation header, which means it can fail with a
750 /// write error. After sending the continuation header, it sets the569 /// write error. After sending the continuation header, it sets the
751 /// request's expect field to `null`.570 /// request's expect field to `null`.
752 ///571 ///
753 /// Asserts that this function is only called once.572 /// Asserts that this function is only called once.
754 pub fn reader(request: *Request) ReaderError!std.io.AnyReader {573 ///
755 const s = request.server;574 /// See `readerExpectNone` for an infallible alternative that cannot write
756 assert(s.state == .received_head);575 /// to the server output stream.
757 s.state = .receiving_body;576 pub fn readerExpectContinue(request: *Request, buffer: []u8) ExpectContinueError!*Reader {
758 s.next_request_start = request.head_end;577 const flush = request.head.expect != null;
759578 try writeExpectContinue(request);
760 if (request.head.expect) |expect| {579 if (flush) try request.server.out.flush();
761 if (mem.eql(u8, expect, "100-continue")) {580 return readerExpectNone(request, buffer);
762 try request.server.connection.stream.writeAll("HTTP/1.1 100 Continue\r\n\r\n");581 }
763 request.head.expect = null;
764 } else {
765 return error.HttpExpectationFailed;
766 }
767 }
768582
769 switch (request.head.transfer_encoding) {583 /// Asserts the expect header is `null`. The caller must handle the
770 .chunked => {584 /// expectation manually and then set the value to `null` prior to calling
771 request.reader_state = .{ .chunk_parser = http.ChunkParser.init };585 /// this function.
772 return .{586 ///
773 .readFn = read_chunked,587 /// Asserts that this function is only called once.
774 .context = request,588 ///
775 };589 /// Invalidates the string memory inside `Head`.
776 },590 pub fn readerExpectNone(request: *Request, buffer: []u8) *Reader {
777 .none => {591 assert(request.server.reader.state == .received_head);
778 request.reader_state = .{592 assert(request.head.expect == null);
779 .remaining_content_length = request.head.content_length orelse 0,593 request.head.invalidateStrings();
780 };594 if (!request.head.method.requestHasBody()) return .ending;
781 return .{595 return request.server.reader.bodyReader(buffer, request.head.transfer_encoding, request.head.content_length);
782 .readFn = read_cl,596 }
783 .context = request,597
784 };598 pub const ExpectContinueError = error{
785 },599 /// Failed to write "HTTP/1.1 100 Continue\r\n\r\n" to the stream.
786 }600 WriteFailed,
601 /// The client sent an expect HTTP header value other than
602 /// "100-continue".
603 HttpExpectationFailed,
604 };
605
606 pub fn writeExpectContinue(request: *Request) ExpectContinueError!void {
607 const expect = request.head.expect orelse return;
608 if (!mem.eql(u8, expect, "100-continue")) return error.HttpExpectationFailed;
609 try request.server.out.writeAll("HTTP/1.1 100 Continue\r\n\r\n");
610 request.head.expect = null;
787 }611 }
788612
789 /// Returns whether the connection should remain persistent.613 /// Returns whether the connection should remain persistent.
790 /// If it would fail, it instead sets the Server state to `receiving_body`614 ///
615 /// If it would fail, it instead sets the Server state to receiving body
791 /// and returns false.616 /// and returns false.
792 fn discardBody(request: *Request, keep_alive: bool) bool {617 fn discardBody(request: *Request, keep_alive: bool) bool {
793 // Prepare to receive another request on the same connection.618 // Prepare to receive another request on the same connection.
...@@ -798,350 +623,180 @@ pub const Request = struct {...@@ -798,350 +623,180 @@ pub const Request = struct {
798 // or the request body.623 // or the request body.
799 // If the connection won't be kept alive, then none of this matters624 // If the connection won't be kept alive, then none of this matters
800 // because the connection will be severed after the response is sent.625 // because the connection will be severed after the response is sent.
801 const s = request.server;626 const r = &request.server.reader;
802 if (keep_alive and request.head.keep_alive) switch (s.state) {627 if (keep_alive and request.head.keep_alive) switch (r.state) {
803 .received_head => {628 .received_head => {
804 const r = request.reader() catch return false;629 if (request.head.method.requestHasBody()) {
805 _ = r.discard() catch return false;630 assert(request.head.transfer_encoding != .none or request.head.content_length != null);
806 assert(s.state == .ready);631 const reader_interface = request.readerExpectContinue(&.{}) catch return false;
632 _ = reader_interface.discardRemaining() catch return false;
633 assert(r.state == .ready);
634 } else {
635 r.state = .ready;
636 }
807 return true;637 return true;
808 },638 },
809 .receiving_body, .ready => return true,639 .body_remaining_content_length, .body_remaining_chunk_len, .body_none, .ready => return true,
810 else => unreachable,640 else => unreachable,
811 };641 };
812642
813 // Avoid clobbering the state in case a reading stream already exists.643 // Avoid clobbering the state in case a reading stream already exists.
814 switch (s.state) {644 switch (r.state) {
815 .received_head => s.state = .closing,645 .received_head => r.state = .closing,
816 else => {},646 else => {},
817 }647 }
818 return false;648 return false;
819 }649 }
820};650};
821651
822pub const Response = struct {652/// See https://tools.ietf.org/html/rfc6455
823 stream: net.Stream,653pub const WebSocket = struct {
824 send_buffer: []u8,654 key: []const u8,
825 /// Index of the first byte in `send_buffer`.655 input: *Reader,
826 /// This is 0 unless a short write happens in `write`.656 output: *Writer,
827 send_buffer_start: usize,657
828 /// Index of the last byte + 1 in `send_buffer`.658 pub const Header0 = packed struct(u8) {
829 send_buffer_end: usize,659 opcode: Opcode,
830 /// `null` means transfer-encoding: chunked.660 rsv3: u1 = 0,
831 /// As a debugging utility, counts down to zero as bytes are written.661 rsv2: u1 = 0,
832 transfer_encoding: TransferEncoding,662 rsv1: u1 = 0,
833 elide_body: bool,663 fin: bool,
834 /// Indicates how much of the end of the `send_buffer` corresponds to a
835 /// chunk. This amount of data will be wrapped by an HTTP chunk header.
836 chunk_len: usize,
837
838 pub const TransferEncoding = union(enum) {
839 /// End of connection signals the end of the stream.
840 none,
841 /// As a debugging utility, counts down to zero as bytes are written.
842 content_length: u64,
843 /// Each chunk is wrapped in a header and trailer.
844 chunked,
845 };664 };
846665
847 pub const WriteError = net.Stream.WriteError;666 pub const Header1 = packed struct(u8) {
848667 payload_len: enum(u7) {
849 /// When using content-length, asserts that the amount of data sent matches668 len16 = 126,
850 /// the value sent in the header, then calls `flush`.669 len64 = 127,
851 /// Otherwise, transfer-encoding: chunked is being used, and it writes the670 _,
852 /// end-of-stream message, then flushes the stream to the system.671 },
853 /// Respects the value of `elide_body` to omit all data after the headers.672 mask: bool,
854 pub fn end(r: *Response) WriteError!void {
855 switch (r.transfer_encoding) {
856 .content_length => |len| {
857 assert(len == 0); // Trips when end() called before all bytes written.
858 try flush_cl(r);
859 },
860 .none => {
861 try flush_cl(r);
862 },
863 .chunked => {
864 try flush_chunked(r, &.{});
865 },
866 }
867 r.* = undefined;
868 }
869
870 pub const EndChunkedOptions = struct {
871 trailers: []const http.Header = &.{},
872 };673 };
873674
874 /// Asserts that the Response is using transfer-encoding: chunked.675 pub const Opcode = enum(u4) {
875 /// Writes the end-of-stream message and any optional trailers, then676 continuation = 0,
876 /// flushes the stream to the system.677 text = 1,
877 /// Respects the value of `elide_body` to omit all data after the headers.678 binary = 2,
878 /// Asserts there are at most 25 trailers.679 connection_close = 8,
879 pub fn endChunked(r: *Response, options: EndChunkedOptions) WriteError!void {680 ping = 9,
880 assert(r.transfer_encoding == .chunked);681 /// "A Pong frame MAY be sent unsolicited. This serves as a unidirectional
881 try flush_chunked(r, options.trailers);682 /// heartbeat. A response to an unsolicited Pong frame is not expected."
882 r.* = undefined;683 pong = 10,
883 }684 _,
884685 };
885 /// If using content-length, asserts that writing these bytes to the client
886 /// would not exceed the content-length value sent in the HTTP header.
887 /// May return 0, which does not indicate end of stream. The caller decides
888 /// when the end of stream occurs by calling `end`.
889 pub fn write(r: *Response, bytes: []const u8) WriteError!usize {
890 switch (r.transfer_encoding) {
891 .content_length, .none => return write_cl(r, bytes),
892 .chunked => return write_chunked(r, bytes),
893 }
894 }
895
896 fn write_cl(context: *const anyopaque, bytes: []const u8) WriteError!usize {
897 const r: *Response = @ptrCast(@alignCast(@constCast(context)));
898686
899 var trash: u64 = std.math.maxInt(u64);687 pub const ReadSmallTextMessageError = error{
900 const len = switch (r.transfer_encoding) {688 ConnectionClose,
901 .content_length => |*len| len,689 UnexpectedOpCode,
902 else => &trash,690 MessageTooBig,
903 };691 MissingMaskBit,
692 ReadFailed,
693 EndOfStream,
694 };
904695
905 if (r.elide_body) {696 pub const SmallMessage = struct {
906 len.* -= bytes.len;697 /// Can be text, binary, or ping.
907 return bytes.len;698 opcode: Opcode,
908 }699 data: []u8,
700 };
909701
910 if (bytes.len + r.send_buffer_end > r.send_buffer.len) {702 /// Reads the next message from the WebSocket stream, failing if the
911 const send_buffer_len = r.send_buffer_end - r.send_buffer_start;703 /// message does not fit into the input buffer. The returned memory points
912 var iovecs: [2]std.posix.iovec_const = .{704 /// into the input buffer and is invalidated on the next read.
913 .{705 pub fn readSmallMessage(ws: *WebSocket) ReadSmallTextMessageError!SmallMessage {
914 .base = r.send_buffer.ptr + r.send_buffer_start,706 const in = ws.input;
915 .len = send_buffer_len,707 while (true) {
916 },708 const header = try in.takeArray(2);
917 .{709 const h0: Header0 = @bitCast(header[0]);
918 .base = bytes.ptr,710 const h1: Header1 = @bitCast(header[1]);
919 .len = bytes.len,711
920 },712 switch (h0.opcode) {
921 };713 .text, .binary, .pong, .ping => {},
922 const n = try r.stream.writev(&iovecs);714 .connection_close => return error.ConnectionClose,
923715 .continuation => return error.UnexpectedOpCode,
924 if (n >= send_buffer_len) {716 _ => return error.UnexpectedOpCode,
925 // It was enough to reset the buffer.
926 r.send_buffer_start = 0;
927 r.send_buffer_end = 0;
928 const bytes_n = n - send_buffer_len;
929 len.* -= bytes_n;
930 return bytes_n;
931 }717 }
932718
933 // It didn't even make it through the existing buffer, let719 if (!h0.fin) return error.MessageTooBig;
934 // alone the new bytes provided.720 if (!h1.mask) return error.MissingMaskBit;
935 r.send_buffer_start += n;
936 return 0;
937 }
938
939 // All bytes can be stored in the remaining space of the buffer.
940 @memcpy(r.send_buffer[r.send_buffer_end..][0..bytes.len], bytes);
941 r.send_buffer_end += bytes.len;
942 len.* -= bytes.len;
943 return bytes.len;
944 }
945721
946 fn write_chunked(context: *const anyopaque, bytes: []const u8) WriteError!usize {722 const len: usize = switch (h1.payload_len) {
947 const r: *Response = @ptrCast(@alignCast(@constCast(context)));723 .len16 => try in.takeInt(u16, .big),
948 assert(r.transfer_encoding == .chunked);724 .len64 => std.math.cast(usize, try in.takeInt(u64, .big)) orelse return error.MessageTooBig,
949725 else => @intFromEnum(h1.payload_len),
950 if (r.elide_body)726 };
951 return bytes.len;727 if (len > in.buffer.len) return error.MessageTooBig;
952728 const mask: u32 = @bitCast((try in.takeArray(4)).*);
953 if (bytes.len + r.send_buffer_end > r.send_buffer.len) {729 const payload = try in.take(len);
954 const send_buffer_len = r.send_buffer_end - r.send_buffer_start;730
955 const chunk_len = r.chunk_len + bytes.len;731 // Skip pongs.
956 var header_buf: [18]u8 = undefined;732 if (h0.opcode == .pong) continue;
957 const chunk_header = std.fmt.bufPrint(&header_buf, "{x}\r\n", .{chunk_len}) catch unreachable;733
958734 // The last item may contain a partial word of unused data.
959 var iovecs: [5]std.posix.iovec_const = .{735 const floored_len = (payload.len / 4) * 4;
960 .{736 const u32_payload: []align(1) u32 = @ptrCast(payload[0..floored_len]);
961 .base = r.send_buffer.ptr + r.send_buffer_start,737 for (u32_payload) |*elem| elem.* ^= mask;
962 .len = send_buffer_len - r.chunk_len,738 const mask_bytes: []const u8 = @ptrCast(&mask);
963 },739 for (payload[floored_len..], mask_bytes[0 .. payload.len - floored_len]) |*leftover, m|
964 .{740 leftover.* ^= m;
965 .base = chunk_header.ptr,741
966 .len = chunk_header.len,742 return .{
967 },743 .opcode = h0.opcode,
968 .{744 .data = payload,
969 .base = r.send_buffer.ptr + r.send_buffer_end - r.chunk_len,
970 .len = r.chunk_len,
971 },
972 .{
973 .base = bytes.ptr,
974 .len = bytes.len,
975 },
976 .{
977 .base = "\r\n",
978 .len = 2,
979 },
980 };745 };
981 // TODO make this writev instead of writevAll, which involves
982 // complicating the logic of this function.
983 try r.stream.writevAll(&iovecs);
984 r.send_buffer_start = 0;
985 r.send_buffer_end = 0;
986 r.chunk_len = 0;
987 return bytes.len;
988 }746 }
989
990 // All bytes can be stored in the remaining space of the buffer.
991 @memcpy(r.send_buffer[r.send_buffer_end..][0..bytes.len], bytes);
992 r.send_buffer_end += bytes.len;
993 r.chunk_len += bytes.len;
994 return bytes.len;
995 }747 }
996748
997 /// If using content-length, asserts that writing these bytes to the client749 pub fn writeMessage(ws: *WebSocket, data: []const u8, op: Opcode) Writer.Error!void {
998 /// would not exceed the content-length value sent in the HTTP header.750 var bufs: [1][]const u8 = .{data};
999 pub fn writeAll(r: *Response, bytes: []const u8) WriteError!void {751 try writeMessageVecUnflushed(ws, &bufs, op);
1000 var index: usize = 0;752 try ws.output.flush();
1001 while (index < bytes.len) {
1002 index += try write(r, bytes[index..]);
1003 }
1004 }753 }
1005754
1006 /// Sends all buffered data to the client.755 pub fn writeMessageUnflushed(ws: *WebSocket, data: []const u8, op: Opcode) Writer.Error!void {
1007 /// This is redundant after calling `end`.756 var bufs: [1][]const u8 = .{data};
1008 /// Respects the value of `elide_body` to omit all data after the headers.757 try writeMessageVecUnflushed(ws, &bufs, op);
1009 pub fn flush(r: *Response) WriteError!void {
1010 switch (r.transfer_encoding) {
1011 .none, .content_length => return flush_cl(r),
1012 .chunked => return flush_chunked(r, null),
1013 }
1014 }758 }
1015759
1016 fn flush_cl(r: *Response) WriteError!void {760 pub fn writeMessageVec(ws: *WebSocket, data: [][]const u8, op: Opcode) Writer.Error!void {
1017 try r.stream.writeAll(r.send_buffer[r.send_buffer_start..r.send_buffer_end]);761 try writeMessageVecUnflushed(ws, data, op);
1018 r.send_buffer_start = 0;762 try ws.output.flush();
1019 r.send_buffer_end = 0;
1020 }763 }
1021764
1022 fn flush_chunked(r: *Response, end_trailers: ?[]const http.Header) WriteError!void {765 pub fn writeMessageVecUnflushed(ws: *WebSocket, data: [][]const u8, op: Opcode) Writer.Error!void {
1023 const max_trailers = 25;766 const total_len = l: {
1024 if (end_trailers) |trailers| assert(trailers.len <= max_trailers);767 var total_len: u64 = 0;
1025 assert(r.transfer_encoding == .chunked);768 for (data) |iovec| total_len += iovec.len;
1026769 break :l total_len;
1027 const http_headers = r.send_buffer[r.send_buffer_start .. r.send_buffer_end - r.chunk_len];
1028
1029 if (r.elide_body) {
1030 try r.stream.writeAll(http_headers);
1031 r.send_buffer_start = 0;
1032 r.send_buffer_end = 0;
1033 r.chunk_len = 0;
1034 return;
1035 }
1036
1037 var header_buf: [18]u8 = undefined;
1038 const chunk_header = std.fmt.bufPrint(&header_buf, "{x}\r\n", .{r.chunk_len}) catch unreachable;
1039
1040 var iovecs: [max_trailers * 4 + 5]std.posix.iovec_const = undefined;
1041 var iovecs_len: usize = 0;
1042
1043 iovecs[iovecs_len] = .{
1044 .base = http_headers.ptr,
1045 .len = http_headers.len,
1046 };770 };
1047 iovecs_len += 1;771 const out = ws.output;
1048772 try out.writeByte(@bitCast(@as(Header0, .{
1049 if (r.chunk_len > 0) {773 .opcode = op,
1050 iovecs[iovecs_len] = .{774 .fin = true,
1051 .base = chunk_header.ptr,775 })));
1052 .len = chunk_header.len,776 switch (total_len) {
1053 };777 0...125 => try out.writeByte(@bitCast(@as(Header1, .{
1054 iovecs_len += 1;778 .payload_len = @enumFromInt(total_len),
1055779 .mask = false,
1056 iovecs[iovecs_len] = .{780 }))),
1057 .base = r.send_buffer.ptr + r.send_buffer_end - r.chunk_len,781 126...0xffff => {
1058 .len = r.chunk_len,782 try out.writeByte(@bitCast(@as(Header1, .{
1059 };783 .payload_len = .len16,
1060 iovecs_len += 1;784 .mask = false,
1061785 })));
1062 iovecs[iovecs_len] = .{786 try out.writeInt(u16, @intCast(total_len), .big);
1063 .base = "\r\n",787 },
1064 .len = 2,788 else => {
1065 };789 try out.writeByte(@bitCast(@as(Header1, .{
1066 iovecs_len += 1;790 .payload_len = .len64,
1067 }791 .mask = false,
1068792 })));
1069 if (end_trailers) |trailers| {793 try out.writeInt(u64, total_len, .big);
1070 iovecs[iovecs_len] = .{794 },
1071 .base = "0\r\n",
1072 .len = 3,
1073 };
1074 iovecs_len += 1;
1075
1076 for (trailers) |trailer| {
1077 iovecs[iovecs_len] = .{
1078 .base = trailer.name.ptr,
1079 .len = trailer.name.len,
1080 };
1081 iovecs_len += 1;
1082
1083 iovecs[iovecs_len] = .{
1084 .base = ": ",
1085 .len = 2,
1086 };
1087 iovecs_len += 1;
1088
1089 if (trailer.value.len != 0) {
1090 iovecs[iovecs_len] = .{
1091 .base = trailer.value.ptr,
1092 .len = trailer.value.len,
1093 };
1094 iovecs_len += 1;
1095 }
1096
1097 iovecs[iovecs_len] = .{
1098 .base = "\r\n",
1099 .len = 2,
1100 };
1101 iovecs_len += 1;
1102 }
1103
1104 iovecs[iovecs_len] = .{
1105 .base = "\r\n",
1106 .len = 2,
1107 };
1108 iovecs_len += 1;
1109 }795 }
1110796 try out.writeVecAll(data);
1111 try r.stream.writevAll(iovecs[0..iovecs_len]);
1112 r.send_buffer_start = 0;
1113 r.send_buffer_end = 0;
1114 r.chunk_len = 0;
1115 }797 }
1116798
1117 pub fn writer(r: *Response) std.io.AnyWriter {799 pub fn flush(ws: *WebSocket) Writer.Error!void {
1118 return .{800 try ws.output.flush();
1119 .writeFn = switch (r.transfer_encoding) {
1120 .none, .content_length => write_cl,
1121 .chunked => write_chunked,
1122 },
1123 .context = r,
1124 };
1125 }801 }
1126};802};
1127
1128fn rebase(s: *Server, index: usize) void {
1129 const leftover = s.read_buffer[s.next_request_start..s.read_buffer_len];
1130 const dest = s.read_buffer[index..][0..leftover.len];
1131 if (leftover.len <= s.next_request_start - index) {
1132 @memcpy(dest, leftover);
1133 } else {
1134 mem.copyBackwards(u8, dest, leftover);
1135 }
1136 s.read_buffer_len = index + leftover.len;
1137}
1138
1139const std = @import("../std.zig");
1140const http = std.http;
1141const mem = std.mem;
1142const net = std.net;
1143const Uri = std.Uri;
1144const assert = std.debug.assert;
1145const testing = std.testing;
1146
1147const Server = @This();
lib/std/http/WebSocket.zig deleted-246
...@@ -1,246 +0,0 @@
1//! See https://tools.ietf.org/html/rfc6455
2
3const builtin = @import("builtin");
4const std = @import("std");
5const WebSocket = @This();
6const assert = std.debug.assert;
7const native_endian = builtin.cpu.arch.endian();
8
9key: []const u8,
10request: *std.http.Server.Request,
11recv_fifo: std.fifo.LinearFifo(u8, .Slice),
12reader: std.io.AnyReader,
13response: std.http.Server.Response,
14/// Number of bytes that have been peeked but not discarded yet.
15outstanding_len: usize,
16
17pub const InitError = error{WebSocketUpgradeMissingKey} ||
18 std.http.Server.Request.ReaderError;
19
20pub fn init(
21 request: *std.http.Server.Request,
22 send_buffer: []u8,
23 recv_buffer: []align(4) u8,
24) InitError!?WebSocket {
25 switch (request.head.version) {
26 .@"HTTP/1.0" => return null,
27 .@"HTTP/1.1" => if (request.head.method != .GET) return null,
28 }
29
30 var sec_websocket_key: ?[]const u8 = null;
31 var upgrade_websocket: bool = false;
32 var it = request.iterateHeaders();
33 while (it.next()) |header| {
34 if (std.ascii.eqlIgnoreCase(header.name, "sec-websocket-key")) {
35 sec_websocket_key = header.value;
36 } else if (std.ascii.eqlIgnoreCase(header.name, "upgrade")) {
37 if (!std.ascii.eqlIgnoreCase(header.value, "websocket"))
38 return null;
39 upgrade_websocket = true;
40 }
41 }
42 if (!upgrade_websocket)
43 return null;
44
45 const key = sec_websocket_key orelse return error.WebSocketUpgradeMissingKey;
46
47 var sha1 = std.crypto.hash.Sha1.init(.{});
48 sha1.update(key);
49 sha1.update("258EAFA5-E914-47DA-95CA-C5AB0DC85B11");
50 var digest: [std.crypto.hash.Sha1.digest_length]u8 = undefined;
51 sha1.final(&digest);
52 var base64_digest: [28]u8 = undefined;
53 assert(std.base64.standard.Encoder.encode(&base64_digest, &digest).len == base64_digest.len);
54
55 request.head.content_length = std.math.maxInt(u64);
56
57 return .{
58 .key = key,
59 .recv_fifo = std.fifo.LinearFifo(u8, .Slice).init(recv_buffer),
60 .reader = try request.reader(),
61 .response = request.respondStreaming(.{
62 .send_buffer = send_buffer,
63 .respond_options = .{
64 .status = .switching_protocols,
65 .extra_headers = &.{
66 .{ .name = "upgrade", .value = "websocket" },
67 .{ .name = "connection", .value = "upgrade" },
68 .{ .name = "sec-websocket-accept", .value = &base64_digest },
69 },
70 .transfer_encoding = .none,
71 },
72 }),
73 .request = request,
74 .outstanding_len = 0,
75 };
76}
77
78pub const Header0 = packed struct(u8) {
79 opcode: Opcode,
80 rsv3: u1 = 0,
81 rsv2: u1 = 0,
82 rsv1: u1 = 0,
83 fin: bool,
84};
85
86pub const Header1 = packed struct(u8) {
87 payload_len: enum(u7) {
88 len16 = 126,
89 len64 = 127,
90 _,
91 },
92 mask: bool,
93};
94
95pub const Opcode = enum(u4) {
96 continuation = 0,
97 text = 1,
98 binary = 2,
99 connection_close = 8,
100 ping = 9,
101 /// "A Pong frame MAY be sent unsolicited. This serves as a unidirectional
102 /// heartbeat. A response to an unsolicited Pong frame is not expected."
103 pong = 10,
104 _,
105};
106
107pub const ReadSmallTextMessageError = error{
108 ConnectionClose,
109 UnexpectedOpCode,
110 MessageTooBig,
111 MissingMaskBit,
112} || RecvError;
113
114pub const SmallMessage = struct {
115 /// Can be text, binary, or ping.
116 opcode: Opcode,
117 data: []u8,
118};
119
120/// Reads the next message from the WebSocket stream, failing if the message does not fit
121/// into `recv_buffer`.
122pub fn readSmallMessage(ws: *WebSocket) ReadSmallTextMessageError!SmallMessage {
123 while (true) {
124 const header_bytes = (try recv(ws, 2))[0..2];
125 const h0: Header0 = @bitCast(header_bytes[0]);
126 const h1: Header1 = @bitCast(header_bytes[1]);
127
128 switch (h0.opcode) {
129 .text, .binary, .pong, .ping => {},
130 .connection_close => return error.ConnectionClose,
131 .continuation => return error.UnexpectedOpCode,
132 _ => return error.UnexpectedOpCode,
133 }
134
135 if (!h0.fin) return error.MessageTooBig;
136 if (!h1.mask) return error.MissingMaskBit;
137
138 const len: usize = switch (h1.payload_len) {
139 .len16 => try recvReadInt(ws, u16),
140 .len64 => std.math.cast(usize, try recvReadInt(ws, u64)) orelse return error.MessageTooBig,
141 else => @intFromEnum(h1.payload_len),
142 };
143 if (len > ws.recv_fifo.buf.len) return error.MessageTooBig;
144
145 const mask: u32 = @bitCast((try recv(ws, 4))[0..4].*);
146 const payload = try recv(ws, len);
147
148 // Skip pongs.
149 if (h0.opcode == .pong) continue;
150
151 // The last item may contain a partial word of unused data.
152 const floored_len = (payload.len / 4) * 4;
153 const u32_payload: []align(1) u32 = @alignCast(std.mem.bytesAsSlice(u32, payload[0..floored_len]));
154 for (u32_payload) |*elem| elem.* ^= mask;
155 const mask_bytes = std.mem.asBytes(&mask)[0 .. payload.len - floored_len];
156 for (payload[floored_len..], mask_bytes) |*leftover, m| leftover.* ^= m;
157
158 return .{
159 .opcode = h0.opcode,
160 .data = payload,
161 };
162 }
163}
164
165const RecvError = std.http.Server.Request.ReadError || error{EndOfStream};
166
167fn recv(ws: *WebSocket, len: usize) RecvError![]u8 {
168 ws.recv_fifo.discard(ws.outstanding_len);
169 assert(len <= ws.recv_fifo.buf.len);
170 if (len > ws.recv_fifo.count) {
171 const small_buf = ws.recv_fifo.writableSlice(0);
172 const needed = len - ws.recv_fifo.count;
173 const buf = if (small_buf.len >= needed) small_buf else b: {
174 ws.recv_fifo.realign();
175 break :b ws.recv_fifo.writableSlice(0);
176 };
177 const n = try @as(RecvError!usize, @errorCast(ws.reader.readAtLeast(buf, needed)));
178 if (n < needed) return error.EndOfStream;
179 ws.recv_fifo.update(n);
180 }
181 ws.outstanding_len = len;
182 // TODO: improve the std lib API so this cast isn't necessary.
183 return @constCast(ws.recv_fifo.readableSliceOfLen(len));
184}
185
186fn recvReadInt(ws: *WebSocket, comptime I: type) !I {
187 const unswapped: I = @bitCast((try recv(ws, @sizeOf(I)))[0..@sizeOf(I)].*);
188 return switch (native_endian) {
189 .little => @byteSwap(unswapped),
190 .big => unswapped,
191 };
192}
193
194pub const WriteError = std.http.Server.Response.WriteError;
195
196pub fn writeMessage(ws: *WebSocket, message: []const u8, opcode: Opcode) WriteError!void {
197 const iovecs: [1]std.posix.iovec_const = .{
198 .{ .base = message.ptr, .len = message.len },
199 };
200 return writeMessagev(ws, &iovecs, opcode);
201}
202
203pub fn writeMessagev(ws: *WebSocket, message: []const std.posix.iovec_const, opcode: Opcode) WriteError!void {
204 const total_len = l: {
205 var total_len: u64 = 0;
206 for (message) |iovec| total_len += iovec.len;
207 break :l total_len;
208 };
209
210 var header_buf: [2 + 8]u8 = undefined;
211 header_buf[0] = @bitCast(@as(Header0, .{
212 .opcode = opcode,
213 .fin = true,
214 }));
215 const header = switch (total_len) {
216 0...125 => blk: {
217 header_buf[1] = @bitCast(@as(Header1, .{
218 .payload_len = @enumFromInt(total_len),
219 .mask = false,
220 }));
221 break :blk header_buf[0..2];
222 },
223 126...0xffff => blk: {
224 header_buf[1] = @bitCast(@as(Header1, .{
225 .payload_len = .len16,
226 .mask = false,
227 }));
228 std.mem.writeInt(u16, header_buf[2..4], @intCast(total_len), .big);
229 break :blk header_buf[0..4];
230 },
231 else => blk: {
232 header_buf[1] = @bitCast(@as(Header1, .{
233 .payload_len = .len64,
234 .mask = false,
235 }));
236 std.mem.writeInt(u64, header_buf[2..10], total_len, .big);
237 break :blk header_buf[0..10];
238 },
239 };
240
241 const response = &ws.response;
242 try response.writeAll(header);
243 for (message) |iovec|
244 try response.writeAll(iovec.base[0..iovec.len]);
245 try response.flush();
246}
lib/std/http/protocol.zig deleted-464
...@@ -1,464 +0,0 @@
1const std = @import("../std.zig");
2const builtin = @import("builtin");
3const testing = std.testing;
4const mem = std.mem;
5
6const assert = std.debug.assert;
7
8pub const State = enum {
9 invalid,
10
11 // Begin header and trailer parsing states.
12
13 start,
14 seen_n,
15 seen_r,
16 seen_rn,
17 seen_rnr,
18 finished,
19
20 // Begin transfer-encoding: chunked parsing states.
21
22 chunk_head_size,
23 chunk_head_ext,
24 chunk_head_r,
25 chunk_data,
26 chunk_data_suffix,
27 chunk_data_suffix_r,
28
29 /// Returns true if the parser is in a content state (ie. not waiting for more headers).
30 pub fn isContent(self: State) bool {
31 return switch (self) {
32 .invalid, .start, .seen_n, .seen_r, .seen_rn, .seen_rnr => false,
33 .finished, .chunk_head_size, .chunk_head_ext, .chunk_head_r, .chunk_data, .chunk_data_suffix, .chunk_data_suffix_r => true,
34 };
35 }
36};
37
38pub const HeadersParser = struct {
39 state: State = .start,
40 /// A fixed buffer of len `max_header_bytes`.
41 /// Pointers into this buffer are not stable until after a message is complete.
42 header_bytes_buffer: []u8,
43 header_bytes_len: u32,
44 next_chunk_length: u64,
45 /// `false`: headers. `true`: trailers.
46 done: bool,
47
48 /// Initializes the parser with a provided buffer `buf`.
49 pub fn init(buf: []u8) HeadersParser {
50 return .{
51 .header_bytes_buffer = buf,
52 .header_bytes_len = 0,
53 .done = false,
54 .next_chunk_length = 0,
55 };
56 }
57
58 /// Reinitialize the parser.
59 /// Asserts the parser is in the "done" state.
60 pub fn reset(hp: *HeadersParser) void {
61 assert(hp.done);
62 hp.* = .{
63 .state = .start,
64 .header_bytes_buffer = hp.header_bytes_buffer,
65 .header_bytes_len = 0,
66 .done = false,
67 .next_chunk_length = 0,
68 };
69 }
70
71 pub fn get(hp: HeadersParser) []u8 {
72 return hp.header_bytes_buffer[0..hp.header_bytes_len];
73 }
74
75 pub fn findHeadersEnd(r: *HeadersParser, bytes: []const u8) u32 {
76 var hp: std.http.HeadParser = .{
77 .state = switch (r.state) {
78 .start => .start,
79 .seen_n => .seen_n,
80 .seen_r => .seen_r,
81 .seen_rn => .seen_rn,
82 .seen_rnr => .seen_rnr,
83 .finished => .finished,
84 else => unreachable,
85 },
86 };
87 const result = hp.feed(bytes);
88 r.state = switch (hp.state) {
89 .start => .start,
90 .seen_n => .seen_n,
91 .seen_r => .seen_r,
92 .seen_rn => .seen_rn,
93 .seen_rnr => .seen_rnr,
94 .finished => .finished,
95 };
96 return @intCast(result);
97 }
98
99 pub fn findChunkedLen(r: *HeadersParser, bytes: []const u8) u32 {
100 var cp: std.http.ChunkParser = .{
101 .state = switch (r.state) {
102 .chunk_head_size => .head_size,
103 .chunk_head_ext => .head_ext,
104 .chunk_head_r => .head_r,
105 .chunk_data => .data,
106 .chunk_data_suffix => .data_suffix,
107 .chunk_data_suffix_r => .data_suffix_r,
108 .invalid => .invalid,
109 else => unreachable,
110 },
111 .chunk_len = r.next_chunk_length,
112 };
113 const result = cp.feed(bytes);
114 r.state = switch (cp.state) {
115 .head_size => .chunk_head_size,
116 .head_ext => .chunk_head_ext,
117 .head_r => .chunk_head_r,
118 .data => .chunk_data,
119 .data_suffix => .chunk_data_suffix,
120 .data_suffix_r => .chunk_data_suffix_r,
121 .invalid => .invalid,
122 };
123 r.next_chunk_length = cp.chunk_len;
124 return @intCast(result);
125 }
126
127 /// Returns whether or not the parser has finished parsing a complete
128 /// message. A message is only complete after the entire body has been read
129 /// and any trailing headers have been parsed.
130 pub fn isComplete(r: *HeadersParser) bool {
131 return r.done and r.state == .finished;
132 }
133
134 pub const CheckCompleteHeadError = error{HttpHeadersOversize};
135
136 /// Pushes `in` into the parser. Returns the number of bytes consumed by
137 /// the header. Any header bytes are appended to `header_bytes_buffer`.
138 pub fn checkCompleteHead(hp: *HeadersParser, in: []const u8) CheckCompleteHeadError!u32 {
139 if (hp.state.isContent()) return 0;
140
141 const i = hp.findHeadersEnd(in);
142 const data = in[0..i];
143 if (hp.header_bytes_len + data.len > hp.header_bytes_buffer.len)
144 return error.HttpHeadersOversize;
145
146 @memcpy(hp.header_bytes_buffer[hp.header_bytes_len..][0..data.len], data);
147 hp.header_bytes_len += @intCast(data.len);
148
149 return i;
150 }
151
152 pub const ReadError = error{
153 HttpChunkInvalid,
154 };
155
156 /// Reads the body of the message into `buffer`. Returns the number of
157 /// bytes placed in the buffer.
158 ///
159 /// If `skip` is true, the buffer will be unused and the body will be skipped.
160 ///
161 /// See `std.http.Client.Connection for an example of `conn`.
162 pub fn read(r: *HeadersParser, conn: anytype, buffer: []u8, skip: bool) !usize {
163 assert(r.state.isContent());
164 if (r.done) return 0;
165
166 var out_index: usize = 0;
167 while (true) {
168 switch (r.state) {
169 .invalid, .start, .seen_n, .seen_r, .seen_rn, .seen_rnr => unreachable,
170 .finished => {
171 const data_avail = r.next_chunk_length;
172
173 if (skip) {
174 conn.fill() catch |err| switch (err) {
175 error.EndOfStream => {
176 r.done = true;
177 return 0;
178 },
179 else => |e| return e,
180 };
181
182 const nread = @min(conn.peek().len, data_avail);
183 conn.drop(@intCast(nread));
184 r.next_chunk_length -= nread;
185
186 if (r.next_chunk_length == 0 or nread == 0) r.done = true;
187
188 return out_index;
189 } else if (out_index < buffer.len) {
190 const out_avail = buffer.len - out_index;
191
192 const can_read = @as(usize, @intCast(@min(data_avail, out_avail)));
193 const nread = try conn.read(buffer[0..can_read]);
194 r.next_chunk_length -= nread;
195
196 if (r.next_chunk_length == 0 or nread == 0) r.done = true;
197
198 return nread;
199 } else {
200 return out_index;
201 }
202 },
203 .chunk_data_suffix, .chunk_data_suffix_r, .chunk_head_size, .chunk_head_ext, .chunk_head_r => {
204 conn.fill() catch |err| switch (err) {
205 error.EndOfStream => {
206 r.done = true;
207 return 0;
208 },
209 else => |e| return e,
210 };
211
212 const i = r.findChunkedLen(conn.peek());
213 conn.drop(@intCast(i));
214
215 switch (r.state) {
216 .invalid => return error.HttpChunkInvalid,
217 .chunk_data => if (r.next_chunk_length == 0) {
218 if (std.mem.eql(u8, conn.peek(), "\r\n")) {
219 r.state = .finished;
220 conn.drop(2);
221 } else {
222 // The trailer section is formatted identically
223 // to the header section.
224 r.state = .seen_rn;
225 }
226 r.done = true;
227
228 return out_index;
229 },
230 else => return out_index,
231 }
232
233 continue;
234 },
235 .chunk_data => {
236 const data_avail = r.next_chunk_length;
237 const out_avail = buffer.len - out_index;
238
239 if (skip) {
240 conn.fill() catch |err| switch (err) {
241 error.EndOfStream => {
242 r.done = true;
243 return 0;
244 },
245 else => |e| return e,
246 };
247
248 const nread = @min(conn.peek().len, data_avail);
249 conn.drop(@intCast(nread));
250 r.next_chunk_length -= nread;
251 } else if (out_avail > 0) {
252 const can_read: usize = @intCast(@min(data_avail, out_avail));
253 const nread = try conn.read(buffer[out_index..][0..can_read]);
254 r.next_chunk_length -= nread;
255 out_index += nread;
256 }
257
258 if (r.next_chunk_length == 0) {
259 r.state = .chunk_data_suffix;
260 continue;
261 }
262
263 return out_index;
264 },
265 }
266 }
267 }
268};
269
270inline fn int16(array: *const [2]u8) u16 {
271 return @as(u16, @bitCast(array.*));
272}
273
274inline fn int24(array: *const [3]u8) u24 {
275 return @as(u24, @bitCast(array.*));
276}
277
278inline fn int32(array: *const [4]u8) u32 {
279 return @as(u32, @bitCast(array.*));
280}
281
282inline fn intShift(comptime T: type, x: anytype) T {
283 switch (@import("builtin").cpu.arch.endian()) {
284 .little => return @as(T, @truncate(x >> (@bitSizeOf(@TypeOf(x)) - @bitSizeOf(T)))),
285 .big => return @as(T, @truncate(x)),
286 }
287}
288
289/// A buffered (and peekable) Connection.
290const MockBufferedConnection = struct {
291 pub const buffer_size = 0x2000;
292
293 conn: std.io.FixedBufferStream([]const u8),
294 buf: [buffer_size]u8 = undefined,
295 start: u16 = 0,
296 end: u16 = 0,
297
298 pub fn fill(conn: *MockBufferedConnection) ReadError!void {
299 if (conn.end != conn.start) return;
300
301 const nread = try conn.conn.read(conn.buf[0..]);
302 if (nread == 0) return error.EndOfStream;
303 conn.start = 0;
304 conn.end = @as(u16, @truncate(nread));
305 }
306
307 pub fn peek(conn: *MockBufferedConnection) []const u8 {
308 return conn.buf[conn.start..conn.end];
309 }
310
311 pub fn drop(conn: *MockBufferedConnection, num: u16) void {
312 conn.start += num;
313 }
314
315 pub fn readAtLeast(conn: *MockBufferedConnection, buffer: []u8, len: usize) ReadError!usize {
316 var out_index: u16 = 0;
317 while (out_index < len) {
318 const available = conn.end - conn.start;
319 const left = buffer.len - out_index;
320
321 if (available > 0) {
322 const can_read = @as(u16, @truncate(@min(available, left)));
323
324 @memcpy(buffer[out_index..][0..can_read], conn.buf[conn.start..][0..can_read]);
325 out_index += can_read;
326 conn.start += can_read;
327
328 continue;
329 }
330
331 if (left > conn.buf.len) {
332 // skip the buffer if the output is large enough
333 return conn.conn.read(buffer[out_index..]);
334 }
335
336 try conn.fill();
337 }
338
339 return out_index;
340 }
341
342 pub fn read(conn: *MockBufferedConnection, buffer: []u8) ReadError!usize {
343 return conn.readAtLeast(buffer, 1);
344 }
345
346 pub const ReadError = std.io.FixedBufferStream([]const u8).ReadError || error{EndOfStream};
347 pub const Reader = std.io.GenericReader(*MockBufferedConnection, ReadError, read);
348
349 pub fn reader(conn: *MockBufferedConnection) Reader {
350 return Reader{ .context = conn };
351 }
352
353 pub fn writeAll(conn: *MockBufferedConnection, buffer: []const u8) WriteError!void {
354 return conn.conn.writeAll(buffer);
355 }
356
357 pub fn write(conn: *MockBufferedConnection, buffer: []const u8) WriteError!usize {
358 return conn.conn.write(buffer);
359 }
360
361 pub const WriteError = std.io.FixedBufferStream([]const u8).WriteError;
362 pub const Writer = std.io.GenericWriter(*MockBufferedConnection, WriteError, write);
363
364 pub fn writer(conn: *MockBufferedConnection) Writer {
365 return Writer{ .context = conn };
366 }
367};
368
369test "HeadersParser.read length" {
370 // mock BufferedConnection for read
371 var headers_buf: [256]u8 = undefined;
372
373 var r = HeadersParser.init(&headers_buf);
374 const data = "GET / HTTP/1.1\r\nHost: localhost\r\nContent-Length: 5\r\n\r\nHello";
375
376 var conn: MockBufferedConnection = .{
377 .conn = std.io.fixedBufferStream(data),
378 };
379
380 while (true) { // read headers
381 try conn.fill();
382
383 const nchecked = try r.checkCompleteHead(conn.peek());
384 conn.drop(@intCast(nchecked));
385
386 if (r.state.isContent()) break;
387 }
388
389 var buf: [8]u8 = undefined;
390
391 r.next_chunk_length = 5;
392 const len = try r.read(&conn, &buf, false);
393 try std.testing.expectEqual(@as(usize, 5), len);
394 try std.testing.expectEqualStrings("Hello", buf[0..len]);
395
396 try std.testing.expectEqualStrings("GET / HTTP/1.1\r\nHost: localhost\r\nContent-Length: 5\r\n\r\n", r.get());
397}
398
399test "HeadersParser.read chunked" {
400 // mock BufferedConnection for read
401
402 var headers_buf: [256]u8 = undefined;
403 var r = HeadersParser.init(&headers_buf);
404 const data = "GET / HTTP/1.1\r\nHost: localhost\r\n\r\n2\r\nHe\r\n2\r\nll\r\n1\r\no\r\n0\r\n\r\n";
405
406 var conn: MockBufferedConnection = .{
407 .conn = std.io.fixedBufferStream(data),
408 };
409
410 while (true) { // read headers
411 try conn.fill();
412
413 const nchecked = try r.checkCompleteHead(conn.peek());
414 conn.drop(@intCast(nchecked));
415
416 if (r.state.isContent()) break;
417 }
418 var buf: [8]u8 = undefined;
419
420 r.state = .chunk_head_size;
421 const len = try r.read(&conn, &buf, false);
422 try std.testing.expectEqual(@as(usize, 5), len);
423 try std.testing.expectEqualStrings("Hello", buf[0..len]);
424
425 try std.testing.expectEqualStrings("GET / HTTP/1.1\r\nHost: localhost\r\n\r\n", r.get());
426}
427
428test "HeadersParser.read chunked trailer" {
429 // mock BufferedConnection for read
430
431 var headers_buf: [256]u8 = undefined;
432 var r = HeadersParser.init(&headers_buf);
433 const data = "GET / HTTP/1.1\r\nHost: localhost\r\n\r\n2\r\nHe\r\n2\r\nll\r\n1\r\no\r\n0\r\nContent-Type: text/plain\r\n\r\n";
434
435 var conn: MockBufferedConnection = .{
436 .conn = std.io.fixedBufferStream(data),
437 };
438
439 while (true) { // read headers
440 try conn.fill();
441
442 const nchecked = try r.checkCompleteHead(conn.peek());
443 conn.drop(@intCast(nchecked));
444
445 if (r.state.isContent()) break;
446 }
447 var buf: [8]u8 = undefined;
448
449 r.state = .chunk_head_size;
450 const len = try r.read(&conn, &buf, false);
451 try std.testing.expectEqual(@as(usize, 5), len);
452 try std.testing.expectEqualStrings("Hello", buf[0..len]);
453
454 while (true) { // read headers
455 try conn.fill();
456
457 const nchecked = try r.checkCompleteHead(conn.peek());
458 conn.drop(@intCast(nchecked));
459
460 if (r.state.isContent()) break;
461 }
462
463 try std.testing.expectEqualStrings("GET / HTTP/1.1\r\nHost: localhost\r\n\r\nContent-Type: text/plain\r\n\r\n", r.get());
464}
lib/std/http/test.zig+315-353
...@@ -10,32 +10,33 @@ const expectError = std.testing.expectError;...@@ -10,32 +10,33 @@ const expectError = std.testing.expectError;
1010
11test "trailers" {11test "trailers" {
12 const test_server = try createTestServer(struct {12 const test_server = try createTestServer(struct {
13 fn run(net_server: *std.net.Server) anyerror!void {13 fn run(test_server: *TestServer) anyerror!void {
14 var header_buffer: [1024]u8 = undefined;14 const net_server = &test_server.net_server;
15 var recv_buffer: [1024]u8 = undefined;
16 var send_buffer: [1024]u8 = undefined;
15 var remaining: usize = 1;17 var remaining: usize = 1;
16 while (remaining != 0) : (remaining -= 1) {18 while (remaining != 0) : (remaining -= 1) {
17 const conn = try net_server.accept();19 const connection = try net_server.accept();
18 defer conn.stream.close();20 defer connection.stream.close();
1921
20 var server = http.Server.init(conn, &header_buffer);22 var connection_br = connection.stream.reader(&recv_buffer);
23 var connection_bw = connection.stream.writer(&send_buffer);
24 var server = http.Server.init(connection_br.interface(), &connection_bw.interface);
2125
22 try expectEqual(.ready, server.state);26 try expectEqual(.ready, server.reader.state);
23 var request = try server.receiveHead();27 var request = try server.receiveHead();
24 try serve(&request);28 try serve(&request);
25 try expectEqual(.ready, server.state);29 try expectEqual(.ready, server.reader.state);
26 }30 }
27 }31 }
2832
29 fn serve(request: *http.Server.Request) !void {33 fn serve(request: *http.Server.Request) !void {
30 try expectEqualStrings(request.head.target, "/trailer");34 try expectEqualStrings(request.head.target, "/trailer");
3135
32 var send_buffer: [1024]u8 = undefined;36 var response = try request.respondStreaming(&.{}, .{});
33 var response = request.respondStreaming(.{37 try response.writer.writeAll("Hello, ");
34 .send_buffer = &send_buffer,
35 });
36 try response.writeAll("Hello, ");
37 try response.flush();38 try response.flush();
38 try response.writeAll("World!\n");39 try response.writer.writeAll("World!\n");
39 try response.flush();40 try response.flush();
40 try response.endChunked(.{41 try response.endChunked(.{
41 .trailers = &.{42 .trailers = &.{
...@@ -58,34 +59,32 @@ test "trailers" {...@@ -58,34 +59,32 @@ test "trailers" {
58 const uri = try std.Uri.parse(location);59 const uri = try std.Uri.parse(location);
5960
60 {61 {
61 var server_header_buffer: [1024]u8 = undefined;62 var req = try client.request(.GET, uri, .{});
62 var req = try client.open(.GET, uri, .{
63 .server_header_buffer = &server_header_buffer,
64 });
65 defer req.deinit();63 defer req.deinit();
6664
67 try req.send();65 try req.sendBodiless();
68 try req.wait();66 var response = try req.receiveHead(&.{});
69
70 const body = try req.reader().readAllAlloc(gpa, 8192);
71 defer gpa.free(body);
7267
73 try expectEqualStrings("Hello, World!\n", body);
74
75 var it = req.response.iterateHeaders();
76 {68 {
69 var it = response.head.iterateHeaders();
77 const header = it.next().?;70 const header = it.next().?;
78 try expect(!it.is_trailer);
79 try expectEqualStrings("transfer-encoding", header.name);71 try expectEqualStrings("transfer-encoding", header.name);
80 try expectEqualStrings("chunked", header.value);72 try expectEqualStrings("chunked", header.value);
73 try expectEqual(null, it.next());
81 }74 }
75
76 const body = try response.reader(&.{}).allocRemaining(gpa, .unlimited);
77 defer gpa.free(body);
78
79 try expectEqualStrings("Hello, World!\n", body);
80
82 {81 {
82 var it = response.iterateTrailers();
83 const header = it.next().?;83 const header = it.next().?;
84 try expect(it.is_trailer);
85 try expectEqualStrings("X-Checksum", header.name);84 try expectEqualStrings("X-Checksum", header.name);
86 try expectEqualStrings("aaaa", header.value);85 try expectEqualStrings("aaaa", header.value);
86 try expectEqual(null, it.next());
87 }87 }
88 try expectEqual(null, it.next());
89 }88 }
9089
91 // connection has been kept alive90 // connection has been kept alive
...@@ -94,19 +93,24 @@ test "trailers" {...@@ -94,19 +93,24 @@ test "trailers" {
9493
95test "HTTP server handles a chunked transfer coding request" {94test "HTTP server handles a chunked transfer coding request" {
96 const test_server = try createTestServer(struct {95 const test_server = try createTestServer(struct {
97 fn run(net_server: *std.net.Server) !void {96 fn run(test_server: *TestServer) anyerror!void {
98 var header_buffer: [8192]u8 = undefined;97 const net_server = &test_server.net_server;
99 const conn = try net_server.accept();98 var recv_buffer: [8192]u8 = undefined;
100 defer conn.stream.close();99 var send_buffer: [500]u8 = undefined;
101100 const connection = try net_server.accept();
102 var server = http.Server.init(conn, &header_buffer);101 defer connection.stream.close();
102
103 var connection_br = connection.stream.reader(&recv_buffer);
104 var connection_bw = connection.stream.writer(&send_buffer);
105 var server = http.Server.init(connection_br.interface(), &connection_bw.interface);
103 var request = try server.receiveHead();106 var request = try server.receiveHead();
104107
105 try expect(request.head.transfer_encoding == .chunked);108 try expect(request.head.transfer_encoding == .chunked);
106109
107 var buf: [128]u8 = undefined;110 var buf: [128]u8 = undefined;
108 const n = try (try request.reader()).readAll(&buf);111 var br = try request.readerExpectContinue(&.{});
109 try expect(mem.eql(u8, buf[0..n], "ABCD"));112 const n = try br.readSliceShort(&buf);
113 try expectEqualStrings("ABCD", buf[0..n]);
110114
111 try request.respond("message from server!\n", .{115 try request.respond("message from server!\n", .{
112 .extra_headers = &.{116 .extra_headers = &.{
...@@ -154,16 +158,20 @@ test "HTTP server handles a chunked transfer coding request" {...@@ -154,16 +158,20 @@ test "HTTP server handles a chunked transfer coding request" {
154158
155test "echo content server" {159test "echo content server" {
156 const test_server = try createTestServer(struct {160 const test_server = try createTestServer(struct {
157 fn run(net_server: *std.net.Server) anyerror!void {161 fn run(test_server: *TestServer) anyerror!void {
158 var read_buffer: [1024]u8 = undefined;162 const net_server = &test_server.net_server;
163 var recv_buffer: [1024]u8 = undefined;
164 var send_buffer: [100]u8 = undefined;
159165
160 accept: while (true) {166 accept: while (!test_server.shutting_down) {
161 const conn = try net_server.accept();167 const connection = try net_server.accept();
162 defer conn.stream.close();168 defer connection.stream.close();
163169
164 var http_server = http.Server.init(conn, &read_buffer);170 var connection_br = connection.stream.reader(&recv_buffer);
171 var connection_bw = connection.stream.writer(&send_buffer);
172 var http_server = http.Server.init(connection_br.interface(), &connection_bw.interface);
165173
166 while (http_server.state == .ready) {174 while (http_server.reader.state == .ready) {
167 var request = http_server.receiveHead() catch |err| switch (err) {175 var request = http_server.receiveHead() catch |err| switch (err) {
168 error.HttpConnectionClosing => continue :accept,176 error.HttpConnectionClosing => continue :accept,
169 else => |e| return e,177 else => |e| return e,
...@@ -173,8 +181,12 @@ test "echo content server" {...@@ -173,8 +181,12 @@ test "echo content server" {
173 }181 }
174 if (request.head.expect) |expect_header_value| {182 if (request.head.expect) |expect_header_value| {
175 if (mem.eql(u8, expect_header_value, "garbage")) {183 if (mem.eql(u8, expect_header_value, "garbage")) {
176 try expectError(error.HttpExpectationFailed, request.reader());184 try expectError(error.HttpExpectationFailed, request.readerExpectContinue(&.{}));
177 try request.respond("", .{ .keep_alive = false });185 request.head.expect = null;
186 try request.respond("", .{
187 .keep_alive = false,
188 .status = .expectation_failed,
189 });
178 continue;190 continue;
179 }191 }
180 }192 }
...@@ -195,16 +207,16 @@ test "echo content server" {...@@ -195,16 +207,16 @@ test "echo content server" {
195 // request.head.target,207 // request.head.target,
196 //});208 //});
197209
198 const body = try (try request.reader()).readAllAlloc(std.testing.allocator, 8192);210 try expect(mem.startsWith(u8, request.head.target, "/echo-content"));
211 try expectEqualStrings("text/plain", request.head.content_type.?);
212
213 // head strings expire here
214 const body = try (try request.readerExpectContinue(&.{})).allocRemaining(std.testing.allocator, .unlimited);
199 defer std.testing.allocator.free(body);215 defer std.testing.allocator.free(body);
200216
201 try expect(mem.startsWith(u8, request.head.target, "/echo-content"));
202 try expectEqualStrings("Hello, World!\n", body);217 try expectEqualStrings("Hello, World!\n", body);
203 try expectEqualStrings("text/plain", request.head.content_type.?);
204218
205 var send_buffer: [100]u8 = undefined;219 var response = try request.respondStreaming(&.{}, .{
206 var response = request.respondStreaming(.{
207 .send_buffer = &send_buffer,
208 .content_length = switch (request.head.transfer_encoding) {220 .content_length = switch (request.head.transfer_encoding) {
209 .chunked => null,221 .chunked => null,
210 .none => len: {222 .none => len: {
...@@ -213,9 +225,8 @@ test "echo content server" {...@@ -213,9 +225,8 @@ test "echo content server" {
213 },225 },
214 },226 },
215 });227 });
216
217 try response.flush(); // Test an early flush to send the HTTP headers before the body.228 try response.flush(); // Test an early flush to send the HTTP headers before the body.
218 const w = response.writer();229 const w = &response.writer;
219 try w.writeAll("Hello, ");230 try w.writeAll("Hello, ");
220 try w.writeAll("World!\n");231 try w.writeAll("World!\n");
221 try response.end();232 try response.end();
...@@ -241,35 +252,35 @@ test "Server.Request.respondStreaming non-chunked, unknown content-length" {...@@ -241,35 +252,35 @@ test "Server.Request.respondStreaming non-chunked, unknown content-length" {
241 // In this case, the response is expected to stream until the connection is252 // In this case, the response is expected to stream until the connection is
242 // closed, indicating the end of the body.253 // closed, indicating the end of the body.
243 const test_server = try createTestServer(struct {254 const test_server = try createTestServer(struct {
244 fn run(net_server: *std.net.Server) anyerror!void {255 fn run(test_server: *TestServer) anyerror!void {
245 var header_buffer: [1000]u8 = undefined;256 const net_server = &test_server.net_server;
257 var recv_buffer: [1000]u8 = undefined;
258 var send_buffer: [500]u8 = undefined;
246 var remaining: usize = 1;259 var remaining: usize = 1;
247 while (remaining != 0) : (remaining -= 1) {260 while (remaining != 0) : (remaining -= 1) {
248 const conn = try net_server.accept();261 const connection = try net_server.accept();
249 defer conn.stream.close();262 defer connection.stream.close();
250263
251 var server = http.Server.init(conn, &header_buffer);264 var connection_br = connection.stream.reader(&recv_buffer);
265 var connection_bw = connection.stream.writer(&send_buffer);
266 var server = http.Server.init(connection_br.interface(), &connection_bw.interface);
252267
253 try expectEqual(.ready, server.state);268 try expectEqual(.ready, server.reader.state);
254 var request = try server.receiveHead();269 var request = try server.receiveHead();
255 try expectEqualStrings(request.head.target, "/foo");270 try expectEqualStrings(request.head.target, "/foo");
256 var send_buffer: [500]u8 = undefined;271 var buf: [30]u8 = undefined;
257 var response = request.respondStreaming(.{272 var response = try request.respondStreaming(&buf, .{
258 .send_buffer = &send_buffer,
259 .respond_options = .{273 .respond_options = .{
260 .transfer_encoding = .none,274 .transfer_encoding = .none,
261 },275 },
262 });276 });
263 var total: usize = 0;277 const w = &response.writer;
264 for (0..500) |i| {278 for (0..500) |i| {
265 var buf: [30]u8 = undefined;279 try w.print("{d}, ah ha ha!\n", .{i});
266 const line = try std.fmt.bufPrint(&buf, "{d}, ah ha ha!\n", .{i});
267 try response.writeAll(line);
268 total += line.len;
269 }280 }
270 try expectEqual(7390, total);281 try w.flush();
271 try response.end();282 try response.end();
272 try expectEqual(.closing, server.state);283 try expectEqual(.closing, server.reader.state);
273 }284 }
274 }285 }
275 });286 });
...@@ -284,7 +295,7 @@ test "Server.Request.respondStreaming non-chunked, unknown content-length" {...@@ -284,7 +295,7 @@ test "Server.Request.respondStreaming non-chunked, unknown content-length" {
284295
285 var tiny_buffer: [1]u8 = undefined; // allows allocRemaining to detect limit exceeded296 var tiny_buffer: [1]u8 = undefined; // allows allocRemaining to detect limit exceeded
286 var stream_reader = stream.reader(&tiny_buffer);297 var stream_reader = stream.reader(&tiny_buffer);
287 const response = try stream_reader.interface().allocRemaining(gpa, .limited(8192));298 const response = try stream_reader.interface().allocRemaining(gpa, .unlimited);
288 defer gpa.free(response);299 defer gpa.free(response);
289300
290 var expected_response = std.ArrayList(u8).init(gpa);301 var expected_response = std.ArrayList(u8).init(gpa);
...@@ -308,15 +319,20 @@ test "Server.Request.respondStreaming non-chunked, unknown content-length" {...@@ -308,15 +319,20 @@ test "Server.Request.respondStreaming non-chunked, unknown content-length" {
308319
309test "receiving arbitrary http headers from the client" {320test "receiving arbitrary http headers from the client" {
310 const test_server = try createTestServer(struct {321 const test_server = try createTestServer(struct {
311 fn run(net_server: *std.net.Server) anyerror!void {322 fn run(test_server: *TestServer) anyerror!void {
312 var read_buffer: [666]u8 = undefined;323 const net_server = &test_server.net_server;
324 var recv_buffer: [666]u8 = undefined;
325 var send_buffer: [777]u8 = undefined;
313 var remaining: usize = 1;326 var remaining: usize = 1;
314 while (remaining != 0) : (remaining -= 1) {327 while (remaining != 0) : (remaining -= 1) {
315 const conn = try net_server.accept();328 const connection = try net_server.accept();
316 defer conn.stream.close();329 defer connection.stream.close();
317330
318 var server = http.Server.init(conn, &read_buffer);331 var connection_br = connection.stream.reader(&recv_buffer);
319 try expectEqual(.ready, server.state);332 var connection_bw = connection.stream.writer(&send_buffer);
333 var server = http.Server.init(connection_br.interface(), &connection_bw.interface);
334
335 try expectEqual(.ready, server.reader.state);
320 var request = try server.receiveHead();336 var request = try server.receiveHead();
321 try expectEqualStrings("/bar", request.head.target);337 try expectEqualStrings("/bar", request.head.target);
322 var it = request.iterateHeaders();338 var it = request.iterateHeaders();
...@@ -350,7 +366,7 @@ test "receiving arbitrary http headers from the client" {...@@ -350,7 +366,7 @@ test "receiving arbitrary http headers from the client" {
350366
351 var tiny_buffer: [1]u8 = undefined; // allows allocRemaining to detect limit exceeded367 var tiny_buffer: [1]u8 = undefined; // allows allocRemaining to detect limit exceeded
352 var stream_reader = stream.reader(&tiny_buffer);368 var stream_reader = stream.reader(&tiny_buffer);
353 const response = try stream_reader.interface().allocRemaining(gpa, .limited(8192));369 const response = try stream_reader.interface().allocRemaining(gpa, .unlimited);
354 defer gpa.free(response);370 defer gpa.free(response);
355371
356 var expected_response = std.ArrayList(u8).init(gpa);372 var expected_response = std.ArrayList(u8).init(gpa);
...@@ -368,19 +384,21 @@ test "general client/server API coverage" {...@@ -368,19 +384,21 @@ test "general client/server API coverage" {
368 return error.SkipZigTest;384 return error.SkipZigTest;
369 }385 }
370386
371 const global = struct {
372 var handle_new_requests = true;
373 };
374 const test_server = try createTestServer(struct {387 const test_server = try createTestServer(struct {
375 fn run(net_server: *std.net.Server) anyerror!void {388 fn run(test_server: *TestServer) anyerror!void {
376 var client_header_buffer: [1024]u8 = undefined;389 const net_server = &test_server.net_server;
377 outer: while (global.handle_new_requests) {390 var recv_buffer: [1024]u8 = undefined;
391 var send_buffer: [100]u8 = undefined;
392
393 outer: while (!test_server.shutting_down) {
378 var connection = try net_server.accept();394 var connection = try net_server.accept();
379 defer connection.stream.close();395 defer connection.stream.close();
380396
381 var http_server = http.Server.init(connection, &client_header_buffer);397 var connection_br = connection.stream.reader(&recv_buffer);
398 var connection_bw = connection.stream.writer(&send_buffer);
399 var http_server = http.Server.init(connection_br.interface(), &connection_bw.interface);
382400
383 while (http_server.state == .ready) {401 while (http_server.reader.state == .ready) {
384 var request = http_server.receiveHead() catch |err| switch (err) {402 var request = http_server.receiveHead() catch |err| switch (err) {
385 error.HttpConnectionClosing => continue :outer,403 error.HttpConnectionClosing => continue :outer,
386 else => |e| return e,404 else => |e| return e,
...@@ -393,21 +411,19 @@ test "general client/server API coverage" {...@@ -393,21 +411,19 @@ test "general client/server API coverage" {
393411
394 fn handleRequest(request: *http.Server.Request, listen_port: u16) !void {412 fn handleRequest(request: *http.Server.Request, listen_port: u16) !void {
395 const log = std.log.scoped(.server);413 const log = std.log.scoped(.server);
414 const gpa = std.testing.allocator;
396415
397 log.info("{f} {s} {s}", .{416 log.info("{t} {t} {s}", .{ request.head.method, request.head.version, request.head.target });
398 request.head.method, @tagName(request.head.version), request.head.target,417 const target = try gpa.dupe(u8, request.head.target);
399 });418 defer gpa.free(target);
400419
401 const gpa = std.testing.allocator;420 const reader = (try request.readerExpectContinue(&.{}));
402 const body = try (try request.reader()).readAllAlloc(gpa, 8192);421 const body = try reader.allocRemaining(gpa, .unlimited);
403 defer gpa.free(body);422 defer gpa.free(body);
404423
405 var send_buffer: [100]u8 = undefined;424 if (mem.startsWith(u8, target, "/get")) {
406425 var response = try request.respondStreaming(&.{}, .{
407 if (mem.startsWith(u8, request.head.target, "/get")) {426 .content_length = if (mem.indexOf(u8, target, "?chunked") == null)
408 var response = request.respondStreaming(.{
409 .send_buffer = &send_buffer,
410 .content_length = if (mem.indexOf(u8, request.head.target, "?chunked") == null)
411 14427 14
412 else428 else
413 null,429 null,
...@@ -417,27 +433,27 @@ test "general client/server API coverage" {...@@ -417,27 +433,27 @@ test "general client/server API coverage" {
417 },433 },
418 },434 },
419 });435 });
420 const w = response.writer();436 const w = &response.writer;
421 try w.writeAll("Hello, ");437 try w.writeAll("Hello, ");
422 try w.writeAll("World!\n");438 try w.writeAll("World!\n");
423 try response.end();439 try response.end();
424 // Writing again would cause an assertion failure.440 // Writing again would cause an assertion failure.
425 } else if (mem.startsWith(u8, request.head.target, "/large")) {441 } else if (mem.startsWith(u8, target, "/large")) {
426 var response = request.respondStreaming(.{442 var response = try request.respondStreaming(&.{}, .{
427 .send_buffer = &send_buffer,
428 .content_length = 14 * 1024 + 14 * 10,443 .content_length = 14 * 1024 + 14 * 10,
429 });444 });
430445
431 try response.flush(); // Test an early flush to send the HTTP headers before the body.446 try response.flush(); // Test an early flush to send the HTTP headers before the body.
432447
433 const w = response.writer();448 const w = &response.writer;
434449
435 var i: u32 = 0;450 var i: u32 = 0;
436 while (i < 5) : (i += 1) {451 while (i < 5) : (i += 1) {
437 try w.writeAll("Hello, World!\n");452 try w.writeAll("Hello, World!\n");
438 }453 }
439454
440 try w.writeAll("Hello, World!\n" ** 1024);455 var vec: [1][]const u8 = .{"Hello, World!\n"};
456 try w.writeSplatAll(&vec, 1024);
441457
442 i = 0;458 i = 0;
443 while (i < 5) : (i += 1) {459 while (i < 5) : (i += 1) {
...@@ -445,9 +461,8 @@ test "general client/server API coverage" {...@@ -445,9 +461,8 @@ test "general client/server API coverage" {
445 }461 }
446462
447 try response.end();463 try response.end();
448 } else if (mem.eql(u8, request.head.target, "/redirect/1")) {464 } else if (mem.eql(u8, target, "/redirect/1")) {
449 var response = request.respondStreaming(.{465 var response = try request.respondStreaming(&.{}, .{
450 .send_buffer = &send_buffer,
451 .respond_options = .{466 .respond_options = .{
452 .status = .found,467 .status = .found,
453 .extra_headers = &.{468 .extra_headers = &.{
...@@ -456,18 +471,18 @@ test "general client/server API coverage" {...@@ -456,18 +471,18 @@ test "general client/server API coverage" {
456 },471 },
457 });472 });
458473
459 const w = response.writer();474 const w = &response.writer;
460 try w.writeAll("Hello, ");475 try w.writeAll("Hello, ");
461 try w.writeAll("Redirected!\n");476 try w.writeAll("Redirected!\n");
462 try response.end();477 try response.end();
463 } else if (mem.eql(u8, request.head.target, "/redirect/2")) {478 } else if (mem.eql(u8, target, "/redirect/2")) {
464 try request.respond("Hello, Redirected!\n", .{479 try request.respond("Hello, Redirected!\n", .{
465 .status = .found,480 .status = .found,
466 .extra_headers = &.{481 .extra_headers = &.{
467 .{ .name = "location", .value = "/redirect/1" },482 .{ .name = "location", .value = "/redirect/1" },
468 },483 },
469 });484 });
470 } else if (mem.eql(u8, request.head.target, "/redirect/3")) {485 } else if (mem.eql(u8, target, "/redirect/3")) {
471 const location = try std.fmt.allocPrint(gpa, "http://127.0.0.1:{d}/redirect/2", .{486 const location = try std.fmt.allocPrint(gpa, "http://127.0.0.1:{d}/redirect/2", .{
472 listen_port,487 listen_port,
473 });488 });
...@@ -479,23 +494,23 @@ test "general client/server API coverage" {...@@ -479,23 +494,23 @@ test "general client/server API coverage" {
479 .{ .name = "location", .value = location },494 .{ .name = "location", .value = location },
480 },495 },
481 });496 });
482 } else if (mem.eql(u8, request.head.target, "/redirect/4")) {497 } else if (mem.eql(u8, target, "/redirect/4")) {
483 try request.respond("Hello, Redirected!\n", .{498 try request.respond("Hello, Redirected!\n", .{
484 .status = .found,499 .status = .found,
485 .extra_headers = &.{500 .extra_headers = &.{
486 .{ .name = "location", .value = "/redirect/3" },501 .{ .name = "location", .value = "/redirect/3" },
487 },502 },
488 });503 });
489 } else if (mem.eql(u8, request.head.target, "/redirect/5")) {504 } else if (mem.eql(u8, target, "/redirect/5")) {
490 try request.respond("Hello, Redirected!\n", .{505 try request.respond("Hello, Redirected!\n", .{
491 .status = .found,506 .status = .found,
492 .extra_headers = &.{507 .extra_headers = &.{
493 .{ .name = "location", .value = "/%2525" },508 .{ .name = "location", .value = "/%2525" },
494 },509 },
495 });510 });
496 } else if (mem.eql(u8, request.head.target, "/%2525")) {511 } else if (mem.eql(u8, target, "/%2525")) {
497 try request.respond("Encoded redirect successful!\n", .{});512 try request.respond("Encoded redirect successful!\n", .{});
498 } else if (mem.eql(u8, request.head.target, "/redirect/invalid")) {513 } else if (mem.eql(u8, target, "/redirect/invalid")) {
499 const invalid_port = try getUnusedTcpPort();514 const invalid_port = try getUnusedTcpPort();
500 const location = try std.fmt.allocPrint(gpa, "http://127.0.0.1:{d}", .{invalid_port});515 const location = try std.fmt.allocPrint(gpa, "http://127.0.0.1:{d}", .{invalid_port});
501 defer gpa.free(location);516 defer gpa.free(location);
...@@ -506,7 +521,7 @@ test "general client/server API coverage" {...@@ -506,7 +521,7 @@ test "general client/server API coverage" {
506 .{ .name = "location", .value = location },521 .{ .name = "location", .value = location },
507 },522 },
508 });523 });
509 } else if (mem.eql(u8, request.head.target, "/empty")) {524 } else if (mem.eql(u8, target, "/empty")) {
510 try request.respond("", .{525 try request.respond("", .{
511 .extra_headers = &.{526 .extra_headers = &.{
512 .{ .name = "empty", .value = "" },527 .{ .name = "empty", .value = "" },
...@@ -524,17 +539,13 @@ test "general client/server API coverage" {...@@ -524,17 +539,13 @@ test "general client/server API coverage" {
524 return s.listen_address.in.getPort();539 return s.listen_address.in.getPort();
525 }540 }
526 });541 });
527 defer {542 defer test_server.destroy();
528 global.handle_new_requests = false;
529 test_server.destroy();
530 }
531543
532 const log = std.log.scoped(.client);544 const log = std.log.scoped(.client);
533545
534 const gpa = std.testing.allocator;546 const gpa = std.testing.allocator;
535 var client: http.Client = .{ .allocator = gpa };547 var client: http.Client = .{ .allocator = gpa };
536 errdefer client.deinit();548 defer client.deinit();
537 // defer client.deinit(); handled below
538549
539 const port = test_server.port();550 const port = test_server.port();
540551
...@@ -544,20 +555,19 @@ test "general client/server API coverage" {...@@ -544,20 +555,19 @@ test "general client/server API coverage" {
544 const uri = try std.Uri.parse(location);555 const uri = try std.Uri.parse(location);
545556
546 log.info("{s}", .{location});557 log.info("{s}", .{location});
547 var server_header_buffer: [1024]u8 = undefined;558 var redirect_buffer: [1024]u8 = undefined;
548 var req = try client.open(.GET, uri, .{559 var req = try client.request(.GET, uri, .{});
549 .server_header_buffer = &server_header_buffer,
550 });
551 defer req.deinit();560 defer req.deinit();
552561
553 try req.send();562 try req.sendBodiless();
554 try req.wait();563 var response = try req.receiveHead(&redirect_buffer);
564
565 try expectEqualStrings("text/plain", response.head.content_type.?);
555566
556 const body = try req.reader().readAllAlloc(gpa, 8192);567 const body = try response.reader(&.{}).allocRemaining(gpa, .unlimited);
557 defer gpa.free(body);568 defer gpa.free(body);
558569
559 try expectEqualStrings("Hello, World!\n", body);570 try expectEqualStrings("Hello, World!\n", body);
560 try expectEqualStrings("text/plain", req.response.content_type.?);
561 }571 }
562572
563 // connection has been kept alive573 // connection has been kept alive
...@@ -569,16 +579,14 @@ test "general client/server API coverage" {...@@ -569,16 +579,14 @@ test "general client/server API coverage" {
569 const uri = try std.Uri.parse(location);579 const uri = try std.Uri.parse(location);
570580
571 log.info("{s}", .{location});581 log.info("{s}", .{location});
572 var server_header_buffer: [1024]u8 = undefined;582 var redirect_buffer: [1024]u8 = undefined;
573 var req = try client.open(.GET, uri, .{583 var req = try client.request(.GET, uri, .{});
574 .server_header_buffer = &server_header_buffer,
575 });
576 defer req.deinit();584 defer req.deinit();
577585
578 try req.send();586 try req.sendBodiless();
579 try req.wait();587 var response = try req.receiveHead(&redirect_buffer);
580588
581 const body = try req.reader().readAllAlloc(gpa, 8192 * 1024);589 const body = try response.reader(&.{}).allocRemaining(gpa, .unlimited);
582 defer gpa.free(body);590 defer gpa.free(body);
583591
584 try expectEqual(@as(usize, 14 * 1024 + 14 * 10), body.len);592 try expectEqual(@as(usize, 14 * 1024 + 14 * 10), body.len);
...@@ -593,21 +601,20 @@ test "general client/server API coverage" {...@@ -593,21 +601,20 @@ test "general client/server API coverage" {
593 const uri = try std.Uri.parse(location);601 const uri = try std.Uri.parse(location);
594602
595 log.info("{s}", .{location});603 log.info("{s}", .{location});
596 var server_header_buffer: [1024]u8 = undefined;604 var redirect_buffer: [1024]u8 = undefined;
597 var req = try client.open(.HEAD, uri, .{605 var req = try client.request(.HEAD, uri, .{});
598 .server_header_buffer = &server_header_buffer,
599 });
600 defer req.deinit();606 defer req.deinit();
601607
602 try req.send();608 try req.sendBodiless();
603 try req.wait();609 var response = try req.receiveHead(&redirect_buffer);
610
611 try expectEqualStrings("text/plain", response.head.content_type.?);
612 try expectEqual(14, response.head.content_length.?);
604613
605 const body = try req.reader().readAllAlloc(gpa, 8192);614 const body = try response.reader(&.{}).allocRemaining(gpa, .unlimited);
606 defer gpa.free(body);615 defer gpa.free(body);
607616
608 try expectEqualStrings("", body);617 try expectEqualStrings("", body);
609 try expectEqualStrings("text/plain", req.response.content_type.?);
610 try expectEqual(14, req.response.content_length.?);
611 }618 }
612619
613 // connection has been kept alive620 // connection has been kept alive
...@@ -619,20 +626,19 @@ test "general client/server API coverage" {...@@ -619,20 +626,19 @@ test "general client/server API coverage" {
619 const uri = try std.Uri.parse(location);626 const uri = try std.Uri.parse(location);
620627
621 log.info("{s}", .{location});628 log.info("{s}", .{location});
622 var server_header_buffer: [1024]u8 = undefined;629 var redirect_buffer: [1024]u8 = undefined;
623 var req = try client.open(.GET, uri, .{630 var req = try client.request(.GET, uri, .{});
624 .server_header_buffer = &server_header_buffer,
625 });
626 defer req.deinit();631 defer req.deinit();
627632
628 try req.send();633 try req.sendBodiless();
629 try req.wait();634 var response = try req.receiveHead(&redirect_buffer);
635
636 try expectEqualStrings("text/plain", response.head.content_type.?);
630637
631 const body = try req.reader().readAllAlloc(gpa, 8192);638 const body = try response.reader(&.{}).allocRemaining(gpa, .unlimited);
632 defer gpa.free(body);639 defer gpa.free(body);
633640
634 try expectEqualStrings("Hello, World!\n", body);641 try expectEqualStrings("Hello, World!\n", body);
635 try expectEqualStrings("text/plain", req.response.content_type.?);
636 }642 }
637643
638 // connection has been kept alive644 // connection has been kept alive
...@@ -644,21 +650,20 @@ test "general client/server API coverage" {...@@ -644,21 +650,20 @@ test "general client/server API coverage" {
644 const uri = try std.Uri.parse(location);650 const uri = try std.Uri.parse(location);
645651
646 log.info("{s}", .{location});652 log.info("{s}", .{location});
647 var server_header_buffer: [1024]u8 = undefined;653 var redirect_buffer: [1024]u8 = undefined;
648 var req = try client.open(.HEAD, uri, .{654 var req = try client.request(.HEAD, uri, .{});
649 .server_header_buffer = &server_header_buffer,
650 });
651 defer req.deinit();655 defer req.deinit();
652656
653 try req.send();657 try req.sendBodiless();
654 try req.wait();658 var response = try req.receiveHead(&redirect_buffer);
655659
656 const body = try req.reader().readAllAlloc(gpa, 8192);660 try expectEqualStrings("text/plain", response.head.content_type.?);
661 try expect(response.head.transfer_encoding == .chunked);
662
663 const body = try response.reader(&.{}).allocRemaining(gpa, .unlimited);
657 defer gpa.free(body);664 defer gpa.free(body);
658665
659 try expectEqualStrings("", body);666 try expectEqualStrings("", body);
660 try expectEqualStrings("text/plain", req.response.content_type.?);
661 try expect(req.response.transfer_encoding == .chunked);
662 }667 }
663668
664 // connection has been kept alive669 // connection has been kept alive
...@@ -670,21 +675,21 @@ test "general client/server API coverage" {...@@ -670,21 +675,21 @@ test "general client/server API coverage" {
670 const uri = try std.Uri.parse(location);675 const uri = try std.Uri.parse(location);
671676
672 log.info("{s}", .{location});677 log.info("{s}", .{location});
673 var server_header_buffer: [1024]u8 = undefined;678 var redirect_buffer: [1024]u8 = undefined;
674 var req = try client.open(.GET, uri, .{679 var req = try client.request(.GET, uri, .{
675 .server_header_buffer = &server_header_buffer,
676 .keep_alive = false,680 .keep_alive = false,
677 });681 });
678 defer req.deinit();682 defer req.deinit();
679683
680 try req.send();684 try req.sendBodiless();
681 try req.wait();685 var response = try req.receiveHead(&redirect_buffer);
686
687 try expectEqualStrings("text/plain", response.head.content_type.?);
682688
683 const body = try req.reader().readAllAlloc(gpa, 8192);689 const body = try response.reader(&.{}).allocRemaining(gpa, .unlimited);
684 defer gpa.free(body);690 defer gpa.free(body);
685691
686 try expectEqualStrings("Hello, World!\n", body);692 try expectEqualStrings("Hello, World!\n", body);
687 try expectEqualStrings("text/plain", req.response.content_type.?);
688 }693 }
689694
690 // connection has been closed695 // connection has been closed
...@@ -696,32 +701,32 @@ test "general client/server API coverage" {...@@ -696,32 +701,32 @@ test "general client/server API coverage" {
696 const uri = try std.Uri.parse(location);701 const uri = try std.Uri.parse(location);
697702
698 log.info("{s}", .{location});703 log.info("{s}", .{location});
699 var server_header_buffer: [1024]u8 = undefined;704 var redirect_buffer: [1024]u8 = undefined;
700 var req = try client.open(.GET, uri, .{705 var req = try client.request(.GET, uri, .{
701 .server_header_buffer = &server_header_buffer,
702 .extra_headers = &.{706 .extra_headers = &.{
703 .{ .name = "empty", .value = "" },707 .{ .name = "empty", .value = "" },
704 },708 },
705 });709 });
706 defer req.deinit();710 defer req.deinit();
707711
708 try req.send();712 try req.sendBodiless();
709 try req.wait();713 var response = try req.receiveHead(&redirect_buffer);
710714
711 try std.testing.expectEqual(.ok, req.response.status);715 try std.testing.expectEqual(.ok, response.head.status);
712
713 const body = try req.reader().readAllAlloc(gpa, 8192);
714 defer gpa.free(body);
715716
716 try expectEqualStrings("", body);717 var it = response.head.iterateHeaders();
717
718 var it = req.response.iterateHeaders();
719 {718 {
720 const header = it.next().?;719 const header = it.next().?;
721 try expect(!it.is_trailer);720 try expect(!it.is_trailer);
722 try expectEqualStrings("content-length", header.name);721 try expectEqualStrings("content-length", header.name);
723 try expectEqualStrings("0", header.value);722 try expectEqualStrings("0", header.value);
724 }723 }
724
725 const body = try response.reader(&.{}).allocRemaining(gpa, .unlimited);
726 defer gpa.free(body);
727
728 try expectEqualStrings("", body);
729
725 {730 {
726 const header = it.next().?;731 const header = it.next().?;
727 try expect(!it.is_trailer);732 try expect(!it.is_trailer);
...@@ -740,16 +745,14 @@ test "general client/server API coverage" {...@@ -740,16 +745,14 @@ test "general client/server API coverage" {
740 const uri = try std.Uri.parse(location);745 const uri = try std.Uri.parse(location);
741746
742 log.info("{s}", .{location});747 log.info("{s}", .{location});
743 var server_header_buffer: [1024]u8 = undefined;748 var redirect_buffer: [1024]u8 = undefined;
744 var req = try client.open(.GET, uri, .{749 var req = try client.request(.GET, uri, .{});
745 .server_header_buffer = &server_header_buffer,
746 });
747 defer req.deinit();750 defer req.deinit();
748751
749 try req.send();752 try req.sendBodiless();
750 try req.wait();753 var response = try req.receiveHead(&redirect_buffer);
751754
752 const body = try req.reader().readAllAlloc(gpa, 8192);755 const body = try response.reader(&.{}).allocRemaining(gpa, .unlimited);
753 defer gpa.free(body);756 defer gpa.free(body);
754757
755 try expectEqualStrings("Hello, World!\n", body);758 try expectEqualStrings("Hello, World!\n", body);
...@@ -764,16 +767,14 @@ test "general client/server API coverage" {...@@ -764,16 +767,14 @@ test "general client/server API coverage" {
764 const uri = try std.Uri.parse(location);767 const uri = try std.Uri.parse(location);
765768
766 log.info("{s}", .{location});769 log.info("{s}", .{location});
767 var server_header_buffer: [1024]u8 = undefined;770 var redirect_buffer: [1024]u8 = undefined;
768 var req = try client.open(.GET, uri, .{771 var req = try client.request(.GET, uri, .{});
769 .server_header_buffer = &server_header_buffer,
770 });
771 defer req.deinit();772 defer req.deinit();
772773
773 try req.send();774 try req.sendBodiless();
774 try req.wait();775 var response = try req.receiveHead(&redirect_buffer);
775776
776 const body = try req.reader().readAllAlloc(gpa, 8192);777 const body = try response.reader(&.{}).allocRemaining(gpa, .unlimited);
777 defer gpa.free(body);778 defer gpa.free(body);
778779
779 try expectEqualStrings("Hello, World!\n", body);780 try expectEqualStrings("Hello, World!\n", body);
...@@ -788,16 +789,14 @@ test "general client/server API coverage" {...@@ -788,16 +789,14 @@ test "general client/server API coverage" {
788 const uri = try std.Uri.parse(location);789 const uri = try std.Uri.parse(location);
789790
790 log.info("{s}", .{location});791 log.info("{s}", .{location});
791 var server_header_buffer: [1024]u8 = undefined;792 var redirect_buffer: [1024]u8 = undefined;
792 var req = try client.open(.GET, uri, .{793 var req = try client.request(.GET, uri, .{});
793 .server_header_buffer = &server_header_buffer,
794 });
795 defer req.deinit();794 defer req.deinit();
796795
797 try req.send();796 try req.sendBodiless();
798 try req.wait();797 var response = try req.receiveHead(&redirect_buffer);
799798
800 const body = try req.reader().readAllAlloc(gpa, 8192);799 const body = try response.reader(&.{}).allocRemaining(gpa, .unlimited);
801 defer gpa.free(body);800 defer gpa.free(body);
802801
803 try expectEqualStrings("Hello, World!\n", body);802 try expectEqualStrings("Hello, World!\n", body);
...@@ -812,17 +811,17 @@ test "general client/server API coverage" {...@@ -812,17 +811,17 @@ test "general client/server API coverage" {
812 const uri = try std.Uri.parse(location);811 const uri = try std.Uri.parse(location);
813812
814 log.info("{s}", .{location});813 log.info("{s}", .{location});
815 var server_header_buffer: [1024]u8 = undefined;814 var redirect_buffer: [1024]u8 = undefined;
816 var req = try client.open(.GET, uri, .{815 var req = try client.request(.GET, uri, .{});
817 .server_header_buffer = &server_header_buffer,
818 });
819 defer req.deinit();816 defer req.deinit();
820817
821 try req.send();818 try req.sendBodiless();
822 req.wait() catch |err| switch (err) {819 if (req.receiveHead(&redirect_buffer)) |_| {
820 return error.TestFailed;
821 } else |err| switch (err) {
823 error.TooManyHttpRedirects => {},822 error.TooManyHttpRedirects => {},
824 else => return err,823 else => return err,
825 };824 }
826 }825 }
827826
828 { // redirect to encoded url827 { // redirect to encoded url
...@@ -831,16 +830,14 @@ test "general client/server API coverage" {...@@ -831,16 +830,14 @@ test "general client/server API coverage" {
831 const uri = try std.Uri.parse(location);830 const uri = try std.Uri.parse(location);
832831
833 log.info("{s}", .{location});832 log.info("{s}", .{location});
834 var server_header_buffer: [1024]u8 = undefined;833 var redirect_buffer: [1024]u8 = undefined;
835 var req = try client.open(.GET, uri, .{834 var req = try client.request(.GET, uri, .{});
836 .server_header_buffer = &server_header_buffer,
837 });
838 defer req.deinit();835 defer req.deinit();
839836
840 try req.send();837 try req.sendBodiless();
841 try req.wait();838 var response = try req.receiveHead(&redirect_buffer);
842839
843 const body = try req.reader().readAllAlloc(gpa, 8192);840 const body = try response.reader(&.{}).allocRemaining(gpa, .unlimited);
844 defer gpa.free(body);841 defer gpa.free(body);
845842
846 try expectEqualStrings("Encoded redirect successful!\n", body);843 try expectEqualStrings("Encoded redirect successful!\n", body);
...@@ -855,14 +852,12 @@ test "general client/server API coverage" {...@@ -855,14 +852,12 @@ test "general client/server API coverage" {
855 const uri = try std.Uri.parse(location);852 const uri = try std.Uri.parse(location);
856853
857 log.info("{s}", .{location});854 log.info("{s}", .{location});
858 var server_header_buffer: [1024]u8 = undefined;855 var redirect_buffer: [1024]u8 = undefined;
859 var req = try client.open(.GET, uri, .{856 var req = try client.request(.GET, uri, .{});
860 .server_header_buffer = &server_header_buffer,
861 });
862 defer req.deinit();857 defer req.deinit();
863858
864 try req.send();859 try req.sendBodiless();
865 const result = req.wait();860 const result = req.receiveHead(&redirect_buffer);
866861
867 // a proxy without an upstream is likely to return a 5xx status.862 // a proxy without an upstream is likely to return a 5xx status.
868 if (client.http_proxy == null) {863 if (client.http_proxy == null) {
...@@ -872,77 +867,40 @@ test "general client/server API coverage" {...@@ -872,77 +867,40 @@ test "general client/server API coverage" {
872867
873 // connection has been kept alive868 // connection has been kept alive
874 try expect(client.http_proxy != null or client.connection_pool.free_len == 1);869 try expect(client.http_proxy != null or client.connection_pool.free_len == 1);
875
876 { // issue 16282 *** This test leaves the client in an invalid state, it must be last ***
877 const location = try std.fmt.allocPrint(gpa, "http://127.0.0.1:{d}/get", .{port});
878 defer gpa.free(location);
879 const uri = try std.Uri.parse(location);
880
881 const total_connections = client.connection_pool.free_size + 64;
882 var requests = try gpa.alloc(http.Client.Request, total_connections);
883 defer gpa.free(requests);
884
885 var header_bufs = std.ArrayList([]u8).init(gpa);
886 defer header_bufs.deinit();
887 defer for (header_bufs.items) |item| gpa.free(item);
888
889 for (0..total_connections) |i| {
890 const headers_buf = try gpa.alloc(u8, 1024);
891 try header_bufs.append(headers_buf);
892 var req = try client.open(.GET, uri, .{
893 .server_header_buffer = headers_buf,
894 });
895 req.response.parser.done = true;
896 req.connection.?.closing = false;
897 requests[i] = req;
898 }
899
900 for (0..total_connections) |i| {
901 requests[i].deinit();
902 }
903
904 // free connections should be full now
905 try expect(client.connection_pool.free_len == client.connection_pool.free_size);
906 }
907
908 client.deinit();
909
910 {
911 global.handle_new_requests = false;
912
913 const conn = try std.net.tcpConnectToAddress(test_server.net_server.listen_address);
914 conn.close();
915 }
916}870}
917871
918test "Server streams both reading and writing" {872test "Server streams both reading and writing" {
919 const test_server = try createTestServer(struct {873 const test_server = try createTestServer(struct {
920 fn run(net_server: *std.net.Server) anyerror!void {874 fn run(test_server: *TestServer) anyerror!void {
921 var header_buffer: [1024]u8 = undefined;875 const net_server = &test_server.net_server;
922 const conn = try net_server.accept();876 var recv_buffer: [1024]u8 = undefined;
923 defer conn.stream.close();877 var send_buffer: [777]u8 = undefined;
924878
925 var server = http.Server.init(conn, &header_buffer);879 const connection = try net_server.accept();
926 var request = try server.receiveHead();880 defer connection.stream.close();
927 const reader = try request.reader();
928881
929 var send_buffer: [777]u8 = undefined;882 var connection_br = connection.stream.reader(&recv_buffer);
930 var response = request.respondStreaming(.{883 var connection_bw = connection.stream.writer(&send_buffer);
931 .send_buffer = &send_buffer,884 var server = http.Server.init(connection_br.interface(), &connection_bw.interface);
885 var request = try server.receiveHead();
886 var read_buffer: [100]u8 = undefined;
887 var br = try request.readerExpectContinue(&read_buffer);
888 var response = try request.respondStreaming(&.{}, .{
932 .respond_options = .{889 .respond_options = .{
933 .transfer_encoding = .none, // Causes keep_alive=false890 .transfer_encoding = .none, // Causes keep_alive=false
934 },891 },
935 });892 });
936 const writer = response.writer();893 const w = &response.writer;
937894
938 while (true) {895 while (true) {
939 try response.flush();896 try response.flush();
940 var buf: [100]u8 = undefined;897 const buf = br.peekGreedy(1) catch |err| switch (err) {
941 const n = try reader.read(&buf);898 error.EndOfStream => break,
942 if (n == 0) break;899 error.ReadFailed => return error.ReadFailed,
943 const sub_buf = buf[0..n];900 };
944 for (sub_buf) |*b| b.* = std.ascii.toUpper(b.*);901 br.toss(buf.len);
945 try writer.writeAll(sub_buf);902 for (buf) |*b| b.* = std.ascii.toUpper(b.*);
903 try w.writeAll(buf);
946 }904 }
947 try response.end();905 try response.end();
948 }906 }
...@@ -952,27 +910,24 @@ test "Server streams both reading and writing" {...@@ -952,27 +910,24 @@ test "Server streams both reading and writing" {
952 var client: http.Client = .{ .allocator = std.testing.allocator };910 var client: http.Client = .{ .allocator = std.testing.allocator };
953 defer client.deinit();911 defer client.deinit();
954912
955 var server_header_buffer: [555]u8 = undefined;913 var redirect_buffer: [555]u8 = undefined;
956 var req = try client.open(.POST, .{914 var req = try client.request(.POST, .{
957 .scheme = "http",915 .scheme = "http",
958 .host = .{ .raw = "127.0.0.1" },916 .host = .{ .raw = "127.0.0.1" },
959 .port = test_server.port(),917 .port = test_server.port(),
960 .path = .{ .percent_encoded = "/" },918 .path = .{ .percent_encoded = "/" },
961 }, .{919 }, .{});
962 .server_header_buffer = &server_header_buffer,
963 });
964 defer req.deinit();920 defer req.deinit();
965921
966 req.transfer_encoding = .chunked;922 req.transfer_encoding = .chunked;
967 try req.send();923 var body_writer = try req.sendBody(&.{});
968 try req.wait();924 var response = try req.receiveHead(&redirect_buffer);
969
970 try req.writeAll("one ");
971 try req.writeAll("fish");
972925
973 try req.finish();926 try body_writer.writer.writeAll("one ");
927 try body_writer.writer.writeAll("fish");
928 try body_writer.end();
974929
975 const body = try req.reader().readAllAlloc(std.testing.allocator, 8192);930 const body = try response.reader(&.{}).allocRemaining(std.testing.allocator, .unlimited);
976 defer std.testing.allocator.free(body);931 defer std.testing.allocator.free(body);
977932
978 try expectEqualStrings("ONE FISH", body);933 try expectEqualStrings("ONE FISH", body);
...@@ -987,9 +942,8 @@ fn echoTests(client: *http.Client, port: u16) !void {...@@ -987,9 +942,8 @@ fn echoTests(client: *http.Client, port: u16) !void {
987 defer gpa.free(location);942 defer gpa.free(location);
988 const uri = try std.Uri.parse(location);943 const uri = try std.Uri.parse(location);
989944
990 var server_header_buffer: [1024]u8 = undefined;945 var redirect_buffer: [1024]u8 = undefined;
991 var req = try client.open(.POST, uri, .{946 var req = try client.request(.POST, uri, .{
992 .server_header_buffer = &server_header_buffer,
993 .extra_headers = &.{947 .extra_headers = &.{
994 .{ .name = "content-type", .value = "text/plain" },948 .{ .name = "content-type", .value = "text/plain" },
995 },949 },
...@@ -998,14 +952,14 @@ fn echoTests(client: *http.Client, port: u16) !void {...@@ -998,14 +952,14 @@ fn echoTests(client: *http.Client, port: u16) !void {
998952
999 req.transfer_encoding = .{ .content_length = 14 };953 req.transfer_encoding = .{ .content_length = 14 };
1000954
1001 try req.send();955 var body_writer = try req.sendBody(&.{});
1002 try req.writeAll("Hello, ");956 try body_writer.writer.writeAll("Hello, ");
1003 try req.writeAll("World!\n");957 try body_writer.writer.writeAll("World!\n");
1004 try req.finish();958 try body_writer.end();
1005959
1006 try req.wait();960 var response = try req.receiveHead(&redirect_buffer);
1007961
1008 const body = try req.reader().readAllAlloc(gpa, 8192);962 const body = try response.reader(&.{}).allocRemaining(gpa, .unlimited);
1009 defer gpa.free(body);963 defer gpa.free(body);
1010964
1011 try expectEqualStrings("Hello, World!\n", body);965 try expectEqualStrings("Hello, World!\n", body);
...@@ -1021,9 +975,8 @@ fn echoTests(client: *http.Client, port: u16) !void {...@@ -1021,9 +975,8 @@ fn echoTests(client: *http.Client, port: u16) !void {
1021 .{port},975 .{port},
1022 ));976 ));
1023977
1024 var server_header_buffer: [1024]u8 = undefined;978 var redirect_buffer: [1024]u8 = undefined;
1025 var req = try client.open(.POST, uri, .{979 var req = try client.request(.POST, uri, .{
1026 .server_header_buffer = &server_header_buffer,
1027 .extra_headers = &.{980 .extra_headers = &.{
1028 .{ .name = "content-type", .value = "text/plain" },981 .{ .name = "content-type", .value = "text/plain" },
1029 },982 },
...@@ -1032,14 +985,14 @@ fn echoTests(client: *http.Client, port: u16) !void {...@@ -1032,14 +985,14 @@ fn echoTests(client: *http.Client, port: u16) !void {
1032985
1033 req.transfer_encoding = .chunked;986 req.transfer_encoding = .chunked;
1034987
1035 try req.send();988 var body_writer = try req.sendBody(&.{});
1036 try req.writeAll("Hello, ");989 try body_writer.writer.writeAll("Hello, ");
1037 try req.writeAll("World!\n");990 try body_writer.writer.writeAll("World!\n");
1038 try req.finish();991 try body_writer.end();
1039992
1040 try req.wait();993 var response = try req.receiveHead(&redirect_buffer);
1041994
1042 const body = try req.reader().readAllAlloc(gpa, 8192);995 const body = try response.reader(&.{}).allocRemaining(gpa, .unlimited);
1043 defer gpa.free(body);996 defer gpa.free(body);
1044997
1045 try expectEqualStrings("Hello, World!\n", body);998 try expectEqualStrings("Hello, World!\n", body);
...@@ -1053,8 +1006,8 @@ fn echoTests(client: *http.Client, port: u16) !void {...@@ -1053,8 +1006,8 @@ fn echoTests(client: *http.Client, port: u16) !void {
1053 const location = try std.fmt.allocPrint(gpa, "http://127.0.0.1:{d}/echo-content#fetch", .{port});1006 const location = try std.fmt.allocPrint(gpa, "http://127.0.0.1:{d}/echo-content#fetch", .{port});
1054 defer gpa.free(location);1007 defer gpa.free(location);
10551008
1056 var body = std.ArrayList(u8).init(gpa);1009 var body: std.ArrayListUnmanaged(u8) = .empty;
1057 defer body.deinit();1010 defer body.deinit(gpa);
10581011
1059 const res = try client.fetch(.{1012 const res = try client.fetch(.{
1060 .location = .{ .url = location },1013 .location = .{ .url = location },
...@@ -1063,7 +1016,7 @@ fn echoTests(client: *http.Client, port: u16) !void {...@@ -1063,7 +1016,7 @@ fn echoTests(client: *http.Client, port: u16) !void {
1063 .extra_headers = &.{1016 .extra_headers = &.{
1064 .{ .name = "content-type", .value = "text/plain" },1017 .{ .name = "content-type", .value = "text/plain" },
1065 },1018 },
1066 .response_storage = .{ .dynamic = &body },1019 .response_storage = .{ .allocator = gpa, .list = &body },
1067 });1020 });
1068 try expectEqual(.ok, res.status);1021 try expectEqual(.ok, res.status);
1069 try expectEqualStrings("Hello, World!\n", body.items);1022 try expectEqualStrings("Hello, World!\n", body.items);
...@@ -1074,9 +1027,8 @@ fn echoTests(client: *http.Client, port: u16) !void {...@@ -1074,9 +1027,8 @@ fn echoTests(client: *http.Client, port: u16) !void {
1074 defer gpa.free(location);1027 defer gpa.free(location);
1075 const uri = try std.Uri.parse(location);1028 const uri = try std.Uri.parse(location);
10761029
1077 var server_header_buffer: [1024]u8 = undefined;1030 var redirect_buffer: [1024]u8 = undefined;
1078 var req = try client.open(.POST, uri, .{1031 var req = try client.request(.POST, uri, .{
1079 .server_header_buffer = &server_header_buffer,
1080 .extra_headers = &.{1032 .extra_headers = &.{
1081 .{ .name = "expect", .value = "100-continue" },1033 .{ .name = "expect", .value = "100-continue" },
1082 .{ .name = "content-type", .value = "text/plain" },1034 .{ .name = "content-type", .value = "text/plain" },
...@@ -1086,15 +1038,15 @@ fn echoTests(client: *http.Client, port: u16) !void {...@@ -1086,15 +1038,15 @@ fn echoTests(client: *http.Client, port: u16) !void {
10861038
1087 req.transfer_encoding = .chunked;1039 req.transfer_encoding = .chunked;
10881040
1089 try req.send();1041 var body_writer = try req.sendBody(&.{});
1090 try req.writeAll("Hello, ");1042 try body_writer.writer.writeAll("Hello, ");
1091 try req.writeAll("World!\n");1043 try body_writer.writer.writeAll("World!\n");
1092 try req.finish();1044 try body_writer.end();
10931045
1094 try req.wait();1046 var response = try req.receiveHead(&redirect_buffer);
1095 try expectEqual(.ok, req.response.status);1047 try expectEqual(.ok, response.head.status);
10961048
1097 const body = try req.reader().readAllAlloc(gpa, 8192);1049 const body = try response.reader(&.{}).allocRemaining(gpa, .unlimited);
1098 defer gpa.free(body);1050 defer gpa.free(body);
10991051
1100 try expectEqualStrings("Hello, World!\n", body);1052 try expectEqualStrings("Hello, World!\n", body);
...@@ -1105,9 +1057,8 @@ fn echoTests(client: *http.Client, port: u16) !void {...@@ -1105,9 +1057,8 @@ fn echoTests(client: *http.Client, port: u16) !void {
1105 defer gpa.free(location);1057 defer gpa.free(location);
1106 const uri = try std.Uri.parse(location);1058 const uri = try std.Uri.parse(location);
11071059
1108 var server_header_buffer: [1024]u8 = undefined;1060 var redirect_buffer: [1024]u8 = undefined;
1109 var req = try client.open(.POST, uri, .{1061 var req = try client.request(.POST, uri, .{
1110 .server_header_buffer = &server_header_buffer,
1111 .extra_headers = &.{1062 .extra_headers = &.{
1112 .{ .name = "content-type", .value = "text/plain" },1063 .{ .name = "content-type", .value = "text/plain" },
1113 .{ .name = "expect", .value = "garbage" },1064 .{ .name = "expect", .value = "garbage" },
...@@ -1117,23 +1068,24 @@ fn echoTests(client: *http.Client, port: u16) !void {...@@ -1117,23 +1068,24 @@ fn echoTests(client: *http.Client, port: u16) !void {
11171068
1118 req.transfer_encoding = .chunked;1069 req.transfer_encoding = .chunked;
11191070
1120 try req.send();1071 var body_writer = try req.sendBody(&.{});
1121 try req.wait();1072 try body_writer.flush();
1122 try expectEqual(.expectation_failed, req.response.status);1073 var response = try req.receiveHead(&redirect_buffer);
1074 try expectEqual(.expectation_failed, response.head.status);
1075 _ = try response.reader(&.{}).discardRemaining();
1123 }1076 }
1124
1125 _ = try client.fetch(.{
1126 .location = .{
1127 .url = try std.fmt.bufPrint(&location_buffer, "http://127.0.0.1:{d}/end", .{port}),
1128 },
1129 });
1130}1077}
11311078
1132const TestServer = struct {1079const TestServer = struct {
1080 shutting_down: bool,
1133 server_thread: std.Thread,1081 server_thread: std.Thread,
1134 net_server: std.net.Server,1082 net_server: std.net.Server,
11351083
1136 fn destroy(self: *@This()) void {1084 fn destroy(self: *@This()) void {
1085 self.shutting_down = true;
1086 const conn = std.net.tcpConnectToAddress(self.net_server.listen_address) catch @panic("shutdown failure");
1087 conn.close();
1088
1137 self.server_thread.join();1089 self.server_thread.join();
1138 self.net_server.deinit();1090 self.net_server.deinit();
1139 std.testing.allocator.destroy(self);1091 std.testing.allocator.destroy(self);
...@@ -1153,20 +1105,27 @@ fn createTestServer(S: type) !*TestServer {...@@ -1153,20 +1105,27 @@ fn createTestServer(S: type) !*TestServer {
11531105
1154 const address = try std.net.Address.parseIp("127.0.0.1", 0);1106 const address = try std.net.Address.parseIp("127.0.0.1", 0);
1155 const test_server = try std.testing.allocator.create(TestServer);1107 const test_server = try std.testing.allocator.create(TestServer);
1156 test_server.net_server = try address.listen(.{ .reuse_address = true });1108 test_server.* = .{
1157 test_server.server_thread = try std.Thread.spawn(.{}, S.run, .{&test_server.net_server});1109 .net_server = try address.listen(.{ .reuse_address = true }),
1110 .server_thread = try std.Thread.spawn(.{}, S.run, .{test_server}),
1111 .shutting_down = false,
1112 };
1158 return test_server;1113 return test_server;
1159}1114}
11601115
1161test "redirect to different connection" {1116test "redirect to different connection" {
1162 const test_server_new = try createTestServer(struct {1117 const test_server_new = try createTestServer(struct {
1163 fn run(net_server: *std.net.Server) anyerror!void {1118 fn run(test_server: *TestServer) anyerror!void {
1164 var header_buffer: [888]u8 = undefined;1119 const net_server = &test_server.net_server;
1120 var recv_buffer: [888]u8 = undefined;
1121 var send_buffer: [777]u8 = undefined;
11651122
1166 const conn = try net_server.accept();1123 const connection = try net_server.accept();
1167 defer conn.stream.close();1124 defer connection.stream.close();
11681125
1169 var server = http.Server.init(conn, &header_buffer);1126 var connection_br = connection.stream.reader(&recv_buffer);
1127 var connection_bw = connection.stream.writer(&send_buffer);
1128 var server = http.Server.init(connection_br.interface(), &connection_bw.interface);
1170 var request = try server.receiveHead();1129 var request = try server.receiveHead();
1171 try expectEqualStrings(request.head.target, "/ok");1130 try expectEqualStrings(request.head.target, "/ok");
1172 try request.respond("good job, you pass", .{});1131 try request.respond("good job, you pass", .{});
...@@ -1180,18 +1139,22 @@ test "redirect to different connection" {...@@ -1180,18 +1139,22 @@ test "redirect to different connection" {
1180 global.other_port = test_server_new.port();1139 global.other_port = test_server_new.port();
11811140
1182 const test_server_orig = try createTestServer(struct {1141 const test_server_orig = try createTestServer(struct {
1183 fn run(net_server: *std.net.Server) anyerror!void {1142 fn run(test_server: *TestServer) anyerror!void {
1184 var header_buffer: [999]u8 = undefined;1143 const net_server = &test_server.net_server;
1144 var recv_buffer: [999]u8 = undefined;
1185 var send_buffer: [100]u8 = undefined;1145 var send_buffer: [100]u8 = undefined;
11861146
1187 const conn = try net_server.accept();1147 const connection = try net_server.accept();
1188 defer conn.stream.close();1148 defer connection.stream.close();
11891149
1190 const new_loc = try std.fmt.bufPrint(&send_buffer, "http://127.0.0.1:{d}/ok", .{1150 var loc_buf: [50]u8 = undefined;
1151 const new_loc = try std.fmt.bufPrint(&loc_buf, "http://127.0.0.1:{d}/ok", .{
1191 global.other_port.?,1152 global.other_port.?,
1192 });1153 });
11931154
1194 var server = http.Server.init(conn, &header_buffer);1155 var connection_br = connection.stream.reader(&recv_buffer);
1156 var connection_bw = connection.stream.writer(&send_buffer);
1157 var server = http.Server.init(connection_br.interface(), &connection_bw.interface);
1195 var request = try server.receiveHead();1158 var request = try server.receiveHead();
1196 try expectEqualStrings(request.head.target, "/help");1159 try expectEqualStrings(request.head.target, "/help");
1197 try request.respond("", .{1160 try request.respond("", .{
...@@ -1216,16 +1179,15 @@ test "redirect to different connection" {...@@ -1216,16 +1179,15 @@ test "redirect to different connection" {
1216 const uri = try std.Uri.parse(location);1179 const uri = try std.Uri.parse(location);
12171180
1218 {1181 {
1219 var server_header_buffer: [666]u8 = undefined;1182 var redirect_buffer: [666]u8 = undefined;
1220 var req = try client.open(.GET, uri, .{1183 var req = try client.request(.GET, uri, .{});
1221 .server_header_buffer = &server_header_buffer,
1222 });
1223 defer req.deinit();1184 defer req.deinit();
12241185
1225 try req.send();1186 try req.sendBodiless();
1226 try req.wait();1187 var response = try req.receiveHead(&redirect_buffer);
1188 var reader = response.reader(&.{});
12271189
1228 const body = try req.reader().readAllAlloc(gpa, 8192);1190 const body = try reader.allocRemaining(gpa, .unlimited);
1229 defer gpa.free(body);1191 defer gpa.free(body);
12301192
1231 try expectEqualStrings("good job, you pass", body);1193 try expectEqualStrings("good job, you pass", body);
lib/std/net.zig+1-1
...@@ -1944,7 +1944,7 @@ pub const Stream = struct {...@@ -1944,7 +1944,7 @@ pub const Stream = struct {
1944 pub const Error = ReadError;1944 pub const Error = ReadError;
19451945
1946 pub fn getStream(r: *const Reader) Stream {1946 pub fn getStream(r: *const Reader) Stream {
1947 return r.stream;1947 return r.net_stream;
1948 }1948 }
19491949
1950 pub fn getError(r: *const Reader) ?Error {1950 pub fn getError(r: *const Reader) ?Error {
lib/std/std.zig-1
...@@ -57,7 +57,6 @@ pub const debug = @import("debug.zig");...@@ -57,7 +57,6 @@ pub const debug = @import("debug.zig");
57pub const dwarf = @import("dwarf.zig");57pub const dwarf = @import("dwarf.zig");
58pub const elf = @import("elf.zig");58pub const elf = @import("elf.zig");
59pub const enums = @import("enums.zig");59pub const enums = @import("enums.zig");
60pub const fifo = @import("fifo.zig");
61pub const fmt = @import("fmt.zig");60pub const fmt = @import("fmt.zig");
62pub const fs = @import("fs.zig");61pub const fs = @import("fs.zig");
63pub const gpu = @import("gpu.zig");62pub const gpu = @import("gpu.zig");
src/Package/Fetch.zig+148-193
...@@ -385,21 +385,23 @@ pub fn run(f: *Fetch) RunError!void {...@@ -385,21 +385,23 @@ pub fn run(f: *Fetch) RunError!void {
385 var resource: Resource = .{ .dir = dir };385 var resource: Resource = .{ .dir = dir };
386 return f.runResource(path_or_url, &resource, null);386 return f.runResource(path_or_url, &resource, null);
387 } else |dir_err| {387 } else |dir_err| {
388 var server_header_buffer: [init_resource_buffer_size]u8 = undefined;
389
388 const file_err = if (dir_err == error.NotDir) e: {390 const file_err = if (dir_err == error.NotDir) e: {
389 if (fs.cwd().openFile(path_or_url, .{})) |file| {391 if (fs.cwd().openFile(path_or_url, .{})) |file| {
390 var resource: Resource = .{ .file = file };392 var resource: Resource = .{ .file = file.reader(&server_header_buffer) };
391 return f.runResource(path_or_url, &resource, null);393 return f.runResource(path_or_url, &resource, null);
392 } else |err| break :e err;394 } else |err| break :e err;
393 } else dir_err;395 } else dir_err;
394396
395 const uri = std.Uri.parse(path_or_url) catch |uri_err| {397 const uri = std.Uri.parse(path_or_url) catch |uri_err| {
396 return f.fail(0, try eb.printString(398 return f.fail(0, try eb.printString(
397 "'{s}' could not be recognized as a file path ({s}) or an URL ({s})",399 "'{s}' could not be recognized as a file path ({t}) or an URL ({t})",
398 .{ path_or_url, @errorName(file_err), @errorName(uri_err) },400 .{ path_or_url, file_err, uri_err },
399 ));401 ));
400 };402 };
401 var server_header_buffer: [header_buffer_size]u8 = undefined;403 var resource: Resource = undefined;
402 var resource = try f.initResource(uri, &server_header_buffer);404 try f.initResource(uri, &resource, &server_header_buffer);
403 return f.runResource(try uri.path.toRawMaybeAlloc(arena), &resource, null);405 return f.runResource(try uri.path.toRawMaybeAlloc(arena), &resource, null);
404 }406 }
405 },407 },
...@@ -464,8 +466,9 @@ pub fn run(f: *Fetch) RunError!void {...@@ -464,8 +466,9 @@ pub fn run(f: *Fetch) RunError!void {
464 f.location_tok,466 f.location_tok,
465 try eb.printString("invalid URI: {s}", .{@errorName(err)}),467 try eb.printString("invalid URI: {s}", .{@errorName(err)}),
466 );468 );
467 var server_header_buffer: [header_buffer_size]u8 = undefined;469 var buffer: [init_resource_buffer_size]u8 = undefined;
468 var resource = try f.initResource(uri, &server_header_buffer);470 var resource: Resource = undefined;
471 try f.initResource(uri, &resource, &buffer);
469 return f.runResource(try uri.path.toRawMaybeAlloc(arena), &resource, remote.hash);472 return f.runResource(try uri.path.toRawMaybeAlloc(arena), &resource, remote.hash);
470}473}
471474
...@@ -866,8 +869,8 @@ fn fail(f: *Fetch, msg_tok: std.zig.Ast.TokenIndex, msg_str: u32) RunError {...@@ -866,8 +869,8 @@ fn fail(f: *Fetch, msg_tok: std.zig.Ast.TokenIndex, msg_str: u32) RunError {
866}869}
867870
868const Resource = union(enum) {871const Resource = union(enum) {
869 file: fs.File,872 file: fs.File.Reader,
870 http_request: std.http.Client.Request,873 http_request: HttpRequest,
871 git: Git,874 git: Git,
872 dir: fs.Dir,875 dir: fs.Dir,
873876
...@@ -877,10 +880,16 @@ const Resource = union(enum) {...@@ -877,10 +880,16 @@ const Resource = union(enum) {
877 want_oid: git.Oid,880 want_oid: git.Oid,
878 };881 };
879882
883 const HttpRequest = struct {
884 request: std.http.Client.Request,
885 response: std.http.Client.Response,
886 buffer: []u8,
887 };
888
880 fn deinit(resource: *Resource) void {889 fn deinit(resource: *Resource) void {
881 switch (resource.*) {890 switch (resource.*) {
882 .file => |*file| file.close(),891 .file => |*file_reader| file_reader.file.close(),
883 .http_request => |*req| req.deinit(),892 .http_request => |*http_request| http_request.request.deinit(),
884 .git => |*git_resource| {893 .git => |*git_resource| {
885 git_resource.fetch_stream.deinit();894 git_resource.fetch_stream.deinit();
886 git_resource.session.deinit();895 git_resource.session.deinit();
...@@ -890,21 +899,13 @@ const Resource = union(enum) {...@@ -890,21 +899,13 @@ const Resource = union(enum) {
890 resource.* = undefined;899 resource.* = undefined;
891 }900 }
892901
893 fn reader(resource: *Resource) std.io.AnyReader {902 fn reader(resource: *Resource) *std.Io.Reader {
894 return .{903 return switch (resource.*) {
895 .context = resource,904 .file => |*file_reader| return &file_reader.interface,
896 .readFn = read,905 .http_request => |*http_request| return http_request.response.reader(http_request.buffer),
897 };906 .git => |*g| return &g.fetch_stream.reader,
898 }
899
900 fn read(context: *const anyopaque, buffer: []u8) anyerror!usize {
901 const resource: *Resource = @ptrCast(@alignCast(@constCast(context)));
902 switch (resource.*) {
903 .file => |*f| return f.read(buffer),
904 .http_request => |*r| return r.read(buffer),
905 .git => |*g| return g.fetch_stream.read(buffer),
906 .dir => unreachable,907 .dir => unreachable,
907 }908 };
908 }909 }
909};910};
910911
...@@ -967,20 +968,22 @@ const FileType = enum {...@@ -967,20 +968,22 @@ const FileType = enum {
967 }968 }
968};969};
969970
970const header_buffer_size = 16 * 1024;971const init_resource_buffer_size = git.Packet.max_data_length;
971972
972fn initResource(f: *Fetch, uri: std.Uri, server_header_buffer: []u8) RunError!Resource {973fn initResource(f: *Fetch, uri: std.Uri, resource: *Resource, reader_buffer: []u8) RunError!void {
973 const gpa = f.arena.child_allocator;974 const gpa = f.arena.child_allocator;
974 const arena = f.arena.allocator();975 const arena = f.arena.allocator();
975 const eb = &f.error_bundle;976 const eb = &f.error_bundle;
976977
977 if (ascii.eqlIgnoreCase(uri.scheme, "file")) {978 if (ascii.eqlIgnoreCase(uri.scheme, "file")) {
978 const path = try uri.path.toRawMaybeAlloc(arena);979 const path = try uri.path.toRawMaybeAlloc(arena);
979 return .{ .file = f.parent_package_root.openFile(path, .{}) catch |err| {980 const file = f.parent_package_root.openFile(path, .{}) catch |err| {
980 return f.fail(f.location_tok, try eb.printString("unable to open '{f}{s}': {s}", .{981 return f.fail(f.location_tok, try eb.printString("unable to open '{f}{s}': {t}", .{
981 f.parent_package_root, path, @errorName(err),982 f.parent_package_root, path, err,
982 }));983 }));
983 } };984 };
985 resource.* = .{ .file = file.reader(reader_buffer) };
986 return;
984 }987 }
985988
986 const http_client = f.job_queue.http_client;989 const http_client = f.job_queue.http_client;
...@@ -988,37 +991,35 @@ fn initResource(f: *Fetch, uri: std.Uri, server_header_buffer: []u8) RunError!Re...@@ -988,37 +991,35 @@ fn initResource(f: *Fetch, uri: std.Uri, server_header_buffer: []u8) RunError!Re
988 if (ascii.eqlIgnoreCase(uri.scheme, "http") or991 if (ascii.eqlIgnoreCase(uri.scheme, "http") or
989 ascii.eqlIgnoreCase(uri.scheme, "https"))992 ascii.eqlIgnoreCase(uri.scheme, "https"))
990 {993 {
991 var req = http_client.open(.GET, uri, .{994 resource.* = .{ .http_request = .{
992 .server_header_buffer = server_header_buffer,995 .request = http_client.request(.GET, uri, .{}) catch |err|
993 }) catch |err| {996 return f.fail(f.location_tok, try eb.printString("unable to connect to server: {t}", .{err})),
994 return f.fail(f.location_tok, try eb.printString(997 .response = undefined,
995 "unable to connect to server: {s}",998 .buffer = reader_buffer,
996 .{@errorName(err)},999 } };
997 ));1000 const request = &resource.http_request.request;
998 };1001 errdefer request.deinit();
999 errdefer req.deinit(); // releases more than memory1002
10001003 request.sendBodiless() catch |err|
1001 req.send() catch |err| {1004 return f.fail(f.location_tok, try eb.printString("HTTP request failed: {t}", .{err}));
1002 return f.fail(f.location_tok, try eb.printString(1005
1003 "HTTP request failed: {s}",1006 var redirect_buffer: [1024]u8 = undefined;
1004 .{@errorName(err)},1007 const response = &resource.http_request.response;
1005 ));1008 response.* = request.receiveHead(&redirect_buffer) catch |err| switch (err) {
1006 };1009 error.ReadFailed => {
1007 req.wait() catch |err| {1010 return f.fail(f.location_tok, try eb.printString("HTTP response read failure: {t}", .{
1008 return f.fail(f.location_tok, try eb.printString(1011 request.connection.?.getReadError().?,
1009 "invalid HTTP response: {s}",1012 }));
1010 .{@errorName(err)},1013 },
1011 ));1014 else => |e| return f.fail(f.location_tok, try eb.printString("invalid HTTP response: {t}", .{e})),
1012 };1015 };
10131016
1014 if (req.response.status != .ok) {1017 if (response.head.status != .ok) return f.fail(f.location_tok, try eb.printString(
1015 return f.fail(f.location_tok, try eb.printString(1018 "bad HTTP response code: '{d} {s}'",
1016 "bad HTTP response code: '{d} {s}'",1019 .{ response.head.status, response.head.status.phrase() orelse "" },
1017 .{ @intFromEnum(req.response.status), req.response.status.phrase() orelse "" },1020 ));
1018 ));
1019 }
10201021
1021 return .{ .http_request = req };1022 return;
1022 }1023 }
10231024
1024 if (ascii.eqlIgnoreCase(uri.scheme, "git+http") or1025 if (ascii.eqlIgnoreCase(uri.scheme, "git+http") or
...@@ -1026,7 +1027,7 @@ fn initResource(f: *Fetch, uri: std.Uri, server_header_buffer: []u8) RunError!Re...@@ -1026,7 +1027,7 @@ fn initResource(f: *Fetch, uri: std.Uri, server_header_buffer: []u8) RunError!Re
1026 {1027 {
1027 var transport_uri = uri;1028 var transport_uri = uri;
1028 transport_uri.scheme = uri.scheme["git+".len..];1029 transport_uri.scheme = uri.scheme["git+".len..];
1029 var session = git.Session.init(gpa, http_client, transport_uri, server_header_buffer) catch |err| {1030 var session = git.Session.init(gpa, http_client, transport_uri, reader_buffer) catch |err| {
1030 return f.fail(f.location_tok, try eb.printString(1031 return f.fail(f.location_tok, try eb.printString(
1031 "unable to discover remote git server capabilities: {s}",1032 "unable to discover remote git server capabilities: {s}",
1032 .{@errorName(err)},1033 .{@errorName(err)},
...@@ -1042,16 +1043,12 @@ fn initResource(f: *Fetch, uri: std.Uri, server_header_buffer: []u8) RunError!Re...@@ -1042,16 +1043,12 @@ fn initResource(f: *Fetch, uri: std.Uri, server_header_buffer: []u8) RunError!Re
1042 const want_ref_head = try std.fmt.allocPrint(arena, "refs/heads/{s}", .{want_ref});1043 const want_ref_head = try std.fmt.allocPrint(arena, "refs/heads/{s}", .{want_ref});
1043 const want_ref_tag = try std.fmt.allocPrint(arena, "refs/tags/{s}", .{want_ref});1044 const want_ref_tag = try std.fmt.allocPrint(arena, "refs/tags/{s}", .{want_ref});
10441045
1045 var ref_iterator = session.listRefs(.{1046 var ref_iterator: git.Session.RefIterator = undefined;
1047 session.listRefs(&ref_iterator, .{
1046 .ref_prefixes = &.{ want_ref, want_ref_head, want_ref_tag },1048 .ref_prefixes = &.{ want_ref, want_ref_head, want_ref_tag },
1047 .include_peeled = true,1049 .include_peeled = true,
1048 .server_header_buffer = server_header_buffer,1050 .buffer = reader_buffer,
1049 }) catch |err| {1051 }) catch |err| return f.fail(f.location_tok, try eb.printString("unable to list refs: {t}", .{err}));
1050 return f.fail(f.location_tok, try eb.printString(
1051 "unable to list refs: {s}",
1052 .{@errorName(err)},
1053 ));
1054 };
1055 defer ref_iterator.deinit();1052 defer ref_iterator.deinit();
1056 while (ref_iterator.next() catch |err| {1053 while (ref_iterator.next() catch |err| {
1057 return f.fail(f.location_tok, try eb.printString(1054 return f.fail(f.location_tok, try eb.printString(
...@@ -1089,25 +1086,21 @@ fn initResource(f: *Fetch, uri: std.Uri, server_header_buffer: []u8) RunError!Re...@@ -1089,25 +1086,21 @@ fn initResource(f: *Fetch, uri: std.Uri, server_header_buffer: []u8) RunError!Re
10891086
1090 var want_oid_buf: [git.Oid.max_formatted_length]u8 = undefined;1087 var want_oid_buf: [git.Oid.max_formatted_length]u8 = undefined;
1091 _ = std.fmt.bufPrint(&want_oid_buf, "{f}", .{want_oid}) catch unreachable;1088 _ = std.fmt.bufPrint(&want_oid_buf, "{f}", .{want_oid}) catch unreachable;
1092 var fetch_stream = session.fetch(&.{&want_oid_buf}, server_header_buffer) catch |err| {1089 var fetch_stream: git.Session.FetchStream = undefined;
1093 return f.fail(f.location_tok, try eb.printString(1090 session.fetch(&fetch_stream, &.{&want_oid_buf}, reader_buffer) catch |err| {
1094 "unable to create fetch stream: {s}",1091 return f.fail(f.location_tok, try eb.printString("unable to create fetch stream: {t}", .{err}));
1095 .{@errorName(err)},
1096 ));
1097 };1092 };
1098 errdefer fetch_stream.deinit();1093 errdefer fetch_stream.deinit();
10991094
1100 return .{ .git = .{1095 resource.* = .{ .git = .{
1101 .session = session,1096 .session = session,
1102 .fetch_stream = fetch_stream,1097 .fetch_stream = fetch_stream,
1103 .want_oid = want_oid,1098 .want_oid = want_oid,
1104 } };1099 } };
1100 return;
1105 }1101 }
11061102
1107 return f.fail(f.location_tok, try eb.printString(1103 return f.fail(f.location_tok, try eb.printString("unsupported URL scheme: {s}", .{uri.scheme}));
1108 "unsupported URL scheme: {s}",
1109 .{uri.scheme},
1110 ));
1111}1104}
11121105
1113fn unpackResource(1106fn unpackResource(
...@@ -1121,9 +1114,11 @@ fn unpackResource(...@@ -1121,9 +1114,11 @@ fn unpackResource(
1121 .file => FileType.fromPath(uri_path) orelse1114 .file => FileType.fromPath(uri_path) orelse
1122 return f.fail(f.location_tok, try eb.printString("unknown file type: '{s}'", .{uri_path})),1115 return f.fail(f.location_tok, try eb.printString("unknown file type: '{s}'", .{uri_path})),
11231116
1124 .http_request => |req| ft: {1117 .http_request => |*http_request| ft: {
1118 const head = &http_request.response.head;
1119
1125 // Content-Type takes first precedence.1120 // Content-Type takes first precedence.
1126 const content_type = req.response.content_type orelse1121 const content_type = head.content_type orelse
1127 return f.fail(f.location_tok, try eb.addString("missing 'Content-Type' header"));1122 return f.fail(f.location_tok, try eb.addString("missing 'Content-Type' header"));
11281123
1129 // Extract the MIME type, ignoring charset and boundary directives1124 // Extract the MIME type, ignoring charset and boundary directives
...@@ -1165,7 +1160,7 @@ fn unpackResource(...@@ -1165,7 +1160,7 @@ fn unpackResource(
1165 }1160 }
11661161
1167 // Next, the filename from 'content-disposition: attachment' takes precedence.1162 // Next, the filename from 'content-disposition: attachment' takes precedence.
1168 if (req.response.content_disposition) |cd_header| {1163 if (head.content_disposition) |cd_header| {
1169 break :ft FileType.fromContentDisposition(cd_header) orelse {1164 break :ft FileType.fromContentDisposition(cd_header) orelse {
1170 return f.fail(f.location_tok, try eb.printString(1165 return f.fail(f.location_tok, try eb.printString(
1171 "unsupported Content-Disposition header value: '{s}' for Content-Type=application/octet-stream",1166 "unsupported Content-Disposition header value: '{s}' for Content-Type=application/octet-stream",
...@@ -1176,10 +1171,7 @@ fn unpackResource(...@@ -1176,10 +1171,7 @@ fn unpackResource(
11761171
1177 // Finally, the path from the URI is used.1172 // Finally, the path from the URI is used.
1178 break :ft FileType.fromPath(uri_path) orelse {1173 break :ft FileType.fromPath(uri_path) orelse {
1179 return f.fail(f.location_tok, try eb.printString(1174 return f.fail(f.location_tok, try eb.printString("unknown file type: '{s}'", .{uri_path}));
1180 "unknown file type: '{s}'",
1181 .{uri_path},
1182 ));
1183 };1175 };
1184 },1176 },
11851177
...@@ -1187,10 +1179,9 @@ fn unpackResource(...@@ -1187,10 +1179,9 @@ fn unpackResource(
11871179
1188 .dir => |dir| {1180 .dir => |dir| {
1189 f.recursiveDirectoryCopy(dir, tmp_directory.handle) catch |err| {1181 f.recursiveDirectoryCopy(dir, tmp_directory.handle) catch |err| {
1190 return f.fail(f.location_tok, try eb.printString(1182 return f.fail(f.location_tok, try eb.printString("unable to copy directory '{s}': {t}", .{
1191 "unable to copy directory '{s}': {s}",1183 uri_path, err,
1192 .{ uri_path, @errorName(err) },1184 }));
1193 ));
1194 };1185 };
1195 return .{};1186 return .{};
1196 },1187 },
...@@ -1198,27 +1189,17 @@ fn unpackResource(...@@ -1198,27 +1189,17 @@ fn unpackResource(
11981189
1199 switch (file_type) {1190 switch (file_type) {
1200 .tar => {1191 .tar => {
1201 var adapter_buffer: [1024]u8 = undefined;1192 return unpackTarball(f, tmp_directory.handle, resource.reader());
1202 var adapter = resource.reader().adaptToNewApi(&adapter_buffer);
1203 return unpackTarball(f, tmp_directory.handle, &adapter.new_interface);
1204 },1193 },
1205 .@"tar.gz" => {1194 .@"tar.gz" => {
1206 var adapter_buffer: [std.crypto.tls.max_ciphertext_record_len]u8 = undefined;
1207 var adapter = resource.reader().adaptToNewApi(&adapter_buffer);
1208 var flate_buffer: [std.compress.flate.max_window_len]u8 = undefined;1195 var flate_buffer: [std.compress.flate.max_window_len]u8 = undefined;
1209 var decompress: std.compress.flate.Decompress = .init(&adapter.new_interface, .gzip, &flate_buffer);1196 var decompress: std.compress.flate.Decompress = .init(resource.reader(), .gzip, &flate_buffer);
1210 return try unpackTarball(f, tmp_directory.handle, &decompress.reader);1197 return try unpackTarball(f, tmp_directory.handle, &decompress.reader);
1211 },1198 },
1212 .@"tar.xz" => {1199 .@"tar.xz" => {
1213 const gpa = f.arena.child_allocator;1200 const gpa = f.arena.child_allocator;
1214 const reader = resource.reader();1201 var dcp = std.compress.xz.decompress(gpa, resource.reader().adaptToOldInterface()) catch |err|
1215 var br = std.io.bufferedReaderSize(std.crypto.tls.max_ciphertext_record_len, reader);1202 return f.fail(f.location_tok, try eb.printString("unable to decompress tarball: {t}", .{err}));
1216 var dcp = std.compress.xz.decompress(gpa, br.reader()) catch |err| {
1217 return f.fail(f.location_tok, try eb.printString(
1218 "unable to decompress tarball: {s}",
1219 .{@errorName(err)},
1220 ));
1221 };
1222 defer dcp.deinit();1203 defer dcp.deinit();
1223 var adapter_buffer: [1024]u8 = undefined;1204 var adapter_buffer: [1024]u8 = undefined;
1224 var adapter = dcp.reader().adaptToNewApi(&adapter_buffer);1205 var adapter = dcp.reader().adaptToNewApi(&adapter_buffer);
...@@ -1227,9 +1208,7 @@ fn unpackResource(...@@ -1227,9 +1208,7 @@ fn unpackResource(
1227 .@"tar.zst" => {1208 .@"tar.zst" => {
1228 const window_size = std.compress.zstd.default_window_len;1209 const window_size = std.compress.zstd.default_window_len;
1229 const window_buffer = try f.arena.allocator().create([window_size]u8);1210 const window_buffer = try f.arena.allocator().create([window_size]u8);
1230 var adapter_buffer: [std.crypto.tls.max_ciphertext_record_len]u8 = undefined;1211 var decompress: std.compress.zstd.Decompress = .init(resource.reader(), window_buffer, .{
1231 var adapter = resource.reader().adaptToNewApi(&adapter_buffer);
1232 var decompress: std.compress.zstd.Decompress = .init(&adapter.new_interface, window_buffer, .{
1233 .verify_checksum = false,1212 .verify_checksum = false,
1234 });1213 });
1235 return try unpackTarball(f, tmp_directory.handle, &decompress.reader);1214 return try unpackTarball(f, tmp_directory.handle, &decompress.reader);
...@@ -1237,12 +1216,15 @@ fn unpackResource(...@@ -1237,12 +1216,15 @@ fn unpackResource(
1237 .git_pack => return unpackGitPack(f, tmp_directory.handle, &resource.git) catch |err| switch (err) {1216 .git_pack => return unpackGitPack(f, tmp_directory.handle, &resource.git) catch |err| switch (err) {
1238 error.FetchFailed => return error.FetchFailed,1217 error.FetchFailed => return error.FetchFailed,
1239 error.OutOfMemory => return error.OutOfMemory,1218 error.OutOfMemory => return error.OutOfMemory,
1240 else => |e| return f.fail(f.location_tok, try eb.printString(1219 else => |e| return f.fail(f.location_tok, try eb.printString("unable to unpack git files: {t}", .{e})),
1241 "unable to unpack git files: {s}",1220 },
1242 .{@errorName(e)},1221 .zip => return unzip(f, tmp_directory.handle, resource.reader()) catch |err| switch (err) {
1222 error.ReadFailed => return f.fail(f.location_tok, try eb.printString(
1223 "failed reading resource: {t}",
1224 .{err},
1243 )),1225 )),
1226 else => |e| return e,
1244 },1227 },
1245 .zip => return try unzip(f, tmp_directory.handle, resource.reader()),
1246 }1228 }
1247}1229}
12481230
...@@ -1277,99 +1259,69 @@ fn unpackTarball(f: *Fetch, out_dir: fs.Dir, reader: *std.Io.Reader) RunError!Un...@@ -1277,99 +1259,69 @@ fn unpackTarball(f: *Fetch, out_dir: fs.Dir, reader: *std.Io.Reader) RunError!Un
1277 return res;1259 return res;
1278}1260}
12791261
1280fn unzip(f: *Fetch, out_dir: fs.Dir, reader: anytype) RunError!UnpackResult {1262fn unzip(f: *Fetch, out_dir: fs.Dir, reader: *std.Io.Reader) error{ ReadFailed, OutOfMemory, FetchFailed }!UnpackResult {
1281 // We write the entire contents to a file first because zip files1263 // We write the entire contents to a file first because zip files
1282 // must be processed back to front and they could be too large to1264 // must be processed back to front and they could be too large to
1283 // load into memory.1265 // load into memory.
12841266
1285 const cache_root = f.job_queue.global_cache;1267 const cache_root = f.job_queue.global_cache;
1286
1287 // TODO: the downside of this solution is if we get a failure/crash/oom/power out
1288 // during this process, we leave behind a zip file that would be
1289 // difficult to know if/when it can be cleaned up.
1290 // Might be worth it to use a mechanism that enables other processes
1291 // to see if the owning process of a file is still alive (on linux this
1292 // can be done with file locks).
1293 // Coupled with this mechansism, we could also use slots (i.e. zig-cache/tmp/0,
1294 // zig-cache/tmp/1, etc) which would mean that subsequent runs would
1295 // automatically clean up old dead files.
1296 // This could all be done with a simple TmpFile abstraction.
1297 const prefix = "tmp/";1268 const prefix = "tmp/";
1298 const suffix = ".zip";1269 const suffix = ".zip";
1299
1300 const random_bytes_count = 20;
1301 const random_path_len = comptime std.fs.base64_encoder.calcSize(random_bytes_count);
1302 var zip_path: [prefix.len + random_path_len + suffix.len]u8 = undefined;
1303 @memcpy(zip_path[0..prefix.len], prefix);
1304 @memcpy(zip_path[prefix.len + random_path_len ..], suffix);
1305 {
1306 var random_bytes: [random_bytes_count]u8 = undefined;
1307 std.crypto.random.bytes(&random_bytes);
1308 _ = std.fs.base64_encoder.encode(
1309 zip_path[prefix.len..][0..random_path_len],
1310 &random_bytes,
1311 );
1312 }
1313
1314 defer cache_root.handle.deleteFile(&zip_path) catch {};
1315
1316 const eb = &f.error_bundle;1270 const eb = &f.error_bundle;
13171271 const random_len = @sizeOf(u64) * 2;
1318 {1272
1319 var zip_file = cache_root.handle.createFile(1273 var zip_path: [prefix.len + random_len + suffix.len]u8 = undefined;
1320 &zip_path,1274 zip_path[0..prefix.len].* = prefix.*;
1321 .{},1275 zip_path[prefix.len + random_len ..].* = suffix.*;
1322 ) catch |err| return f.fail(f.location_tok, try eb.printString(1276
1323 "failed to create tmp zip file: {s}",1277 var zip_file = while (true) {
1324 .{@errorName(err)},1278 const random_integer = std.crypto.random.int(u64);
1325 ));1279 zip_path[prefix.len..][0..random_len].* = std.fmt.hex(random_integer);
1326 defer zip_file.close();1280
1327 var buf: [4096]u8 = undefined;1281 break cache_root.handle.createFile(&zip_path, .{
1328 while (true) {1282 .exclusive = true,
1329 const len = reader.readAll(&buf) catch |err| return f.fail(f.location_tok, try eb.printString(1283 .read = true,
1330 "read zip stream failed: {s}",1284 }) catch |err| switch (err) {
1331 .{@errorName(err)},1285 error.PathAlreadyExists => continue,
1332 ));1286 else => |e| return f.fail(
1333 if (len == 0) break;1287 f.location_tok,
1334 zip_file.deprecatedWriter().writeAll(buf[0..len]) catch |err| return f.fail(f.location_tok, try eb.printString(1288 try eb.printString("failed to create temporary zip file: {t}", .{e}),
1335 "write temporary zip file failed: {s}",1289 ),
1336 .{@errorName(err)},1290 };
1337 ));1291 };
1338 }1292 defer zip_file.close();
1339 }1293 var zip_file_buffer: [4096]u8 = undefined;
1294 var zip_file_reader = b: {
1295 var zip_file_writer = zip_file.writer(&zip_file_buffer);
1296
1297 _ = reader.streamRemaining(&zip_file_writer.interface) catch |err| switch (err) {
1298 error.ReadFailed => return error.ReadFailed,
1299 error.WriteFailed => return f.fail(
1300 f.location_tok,
1301 try eb.printString("failed writing temporary zip file: {t}", .{err}),
1302 ),
1303 };
1304 zip_file_writer.interface.flush() catch |err| return f.fail(
1305 f.location_tok,
1306 try eb.printString("failed writing temporary zip file: {t}", .{err}),
1307 );
1308 break :b zip_file_writer.moveToReader();
1309 };
13401310
1341 var diagnostics: std.zip.Diagnostics = .{ .allocator = f.arena.allocator() };1311 var diagnostics: std.zip.Diagnostics = .{ .allocator = f.arena.allocator() };
1342 // no need to deinit since we are using an arena allocator1312 // no need to deinit since we are using an arena allocator
13431313
1344 {1314 zip_file_reader.seekTo(0) catch |err|
1345 var zip_file = cache_root.handle.openFile(1315 return f.fail(f.location_tok, try eb.printString("failed to seek temporary zip file: {t}", .{err}));
1346 &zip_path,1316 std.zip.extract(out_dir, &zip_file_reader, .{
1347 .{},1317 .allow_backslashes = true,
1348 ) catch |err| return f.fail(f.location_tok, try eb.printString(1318 .diagnostics = &diagnostics,
1349 "failed to open temporary zip file: {s}",1319 }) catch |err| return f.fail(f.location_tok, try eb.printString("zip extract failed: {t}", .{err}));
1350 .{@errorName(err)},
1351 ));
1352 defer zip_file.close();
1353
1354 var zip_file_buffer: [1024]u8 = undefined;
1355 var zip_file_reader = zip_file.reader(&zip_file_buffer);
1356
1357 std.zip.extract(out_dir, &zip_file_reader, .{
1358 .allow_backslashes = true,
1359 .diagnostics = &diagnostics,
1360 }) catch |err| return f.fail(f.location_tok, try eb.printString(
1361 "zip extract failed: {s}",
1362 .{@errorName(err)},
1363 ));
1364 }
13651320
1366 cache_root.handle.deleteFile(&zip_path) catch |err| return f.fail(f.location_tok, try eb.printString(1321 cache_root.handle.deleteFile(&zip_path) catch |err|
1367 "delete temporary zip failed: {s}",1322 return f.fail(f.location_tok, try eb.printString("delete temporary zip failed: {t}", .{err}));
1368 .{@errorName(err)},
1369 ));
13701323
1371 const res: UnpackResult = .{ .root_dir = diagnostics.root_dir };1324 return .{ .root_dir = diagnostics.root_dir };
1372 return res;
1373}1325}
13741326
1375fn unpackGitPack(f: *Fetch, out_dir: fs.Dir, resource: *Resource.Git) anyerror!UnpackResult {1327fn unpackGitPack(f: *Fetch, out_dir: fs.Dir, resource: *Resource.Git) anyerror!UnpackResult {
...@@ -1387,10 +1339,13 @@ fn unpackGitPack(f: *Fetch, out_dir: fs.Dir, resource: *Resource.Git) anyerror!U...@@ -1387,10 +1339,13 @@ fn unpackGitPack(f: *Fetch, out_dir: fs.Dir, resource: *Resource.Git) anyerror!U
1387 var pack_file = try pack_dir.createFile("pkg.pack", .{ .read = true });1339 var pack_file = try pack_dir.createFile("pkg.pack", .{ .read = true });
1388 defer pack_file.close();1340 defer pack_file.close();
1389 var pack_file_buffer: [4096]u8 = undefined;1341 var pack_file_buffer: [4096]u8 = undefined;
1390 var fifo = std.fifo.LinearFifo(u8, .{ .Slice = {} }).init(&pack_file_buffer);1342 var pack_file_reader = b: {
1391 try fifo.pump(resource.fetch_stream.reader(), pack_file.deprecatedWriter());1343 var pack_file_writer = pack_file.writer(&pack_file_buffer);
13921344 const fetch_reader = &resource.fetch_stream.reader;
1393 var pack_file_reader = pack_file.reader(&pack_file_buffer);1345 _ = try fetch_reader.streamRemaining(&pack_file_writer.interface);
1346 try pack_file_writer.interface.flush();
1347 break :b pack_file_writer.moveToReader();
1348 };
13941349
1395 var index_file = try pack_dir.createFile("pkg.idx", .{ .read = true });1350 var index_file = try pack_dir.createFile("pkg.idx", .{ .read = true });
1396 defer index_file.close();1351 defer index_file.close();
src/Package/Fetch/git.zig+156-131
...@@ -585,17 +585,17 @@ const ObjectCache = struct {...@@ -585,17 +585,17 @@ const ObjectCache = struct {
585/// [protocol-common](https://git-scm.com/docs/protocol-common). The special585/// [protocol-common](https://git-scm.com/docs/protocol-common). The special
586/// meanings of the delimiter and response-end packets are documented in586/// meanings of the delimiter and response-end packets are documented in
587/// [protocol-v2](https://git-scm.com/docs/protocol-v2).587/// [protocol-v2](https://git-scm.com/docs/protocol-v2).
588const Packet = union(enum) {588pub const Packet = union(enum) {
589 flush,589 flush,
590 delimiter,590 delimiter,
591 response_end,591 response_end,
592 data: []const u8,592 data: []const u8,
593593
594 const max_data_length = 65516;594 pub const max_data_length = 65516;
595595
596 /// Reads a packet in pkt-line format.596 /// Reads a packet in pkt-line format.
597 fn read(reader: anytype, buf: *[max_data_length]u8) !Packet {597 fn read(reader: *std.Io.Reader) !Packet {
598 const length = std.fmt.parseUnsigned(u16, &try reader.readBytesNoEof(4), 16) catch return error.InvalidPacket;598 const length = std.fmt.parseUnsigned(u16, try reader.take(4), 16) catch return error.InvalidPacket;
599 switch (length) {599 switch (length) {
600 0 => return .flush,600 0 => return .flush,
601 1 => return .delimiter,601 1 => return .delimiter,
...@@ -603,13 +603,11 @@ const Packet = union(enum) {...@@ -603,13 +603,11 @@ const Packet = union(enum) {
603 3 => return error.InvalidPacket,603 3 => return error.InvalidPacket,
604 else => if (length - 4 > max_data_length) return error.InvalidPacket,604 else => if (length - 4 > max_data_length) return error.InvalidPacket,
605 }605 }
606 const data = buf[0 .. length - 4];606 return .{ .data = try reader.take(length - 4) };
607 try reader.readNoEof(data);
608 return .{ .data = data };
609 }607 }
610608
611 /// Writes a packet in pkt-line format.609 /// Writes a packet in pkt-line format.
612 fn write(packet: Packet, writer: anytype) !void {610 fn write(packet: Packet, writer: *std.Io.Writer) !void {
613 switch (packet) {611 switch (packet) {
614 .flush => try writer.writeAll("0000"),612 .flush => try writer.writeAll("0000"),
615 .delimiter => try writer.writeAll("0001"),613 .delimiter => try writer.writeAll("0001"),
...@@ -657,8 +655,10 @@ pub const Session = struct {...@@ -657,8 +655,10 @@ pub const Session = struct {
657 allocator: Allocator,655 allocator: Allocator,
658 transport: *std.http.Client,656 transport: *std.http.Client,
659 uri: std.Uri,657 uri: std.Uri,
660 http_headers_buffer: []u8,658 /// Asserted to be at least `Packet.max_data_length`
659 response_buffer: []u8,
661 ) !Session {660 ) !Session {
661 assert(response_buffer.len >= Packet.max_data_length);
662 var session: Session = .{662 var session: Session = .{
663 .transport = transport,663 .transport = transport,
664 .location = try .init(allocator, uri),664 .location = try .init(allocator, uri),
...@@ -668,7 +668,8 @@ pub const Session = struct {...@@ -668,7 +668,8 @@ pub const Session = struct {
668 .allocator = allocator,668 .allocator = allocator,
669 };669 };
670 errdefer session.deinit();670 errdefer session.deinit();
671 var capability_iterator = try session.getCapabilities(http_headers_buffer);671 var capability_iterator: CapabilityIterator = undefined;
672 try session.getCapabilities(&capability_iterator, response_buffer);
672 defer capability_iterator.deinit();673 defer capability_iterator.deinit();
673 while (try capability_iterator.next()) |capability| {674 while (try capability_iterator.next()) |capability| {
674 if (mem.eql(u8, capability.key, "agent")) {675 if (mem.eql(u8, capability.key, "agent")) {
...@@ -743,7 +744,8 @@ pub const Session = struct {...@@ -743,7 +744,8 @@ pub const Session = struct {
743 ///744 ///
744 /// The `session.location` is updated if the server returns a redirect, so745 /// The `session.location` is updated if the server returns a redirect, so
745 /// that subsequent session functions do not need to handle redirects.746 /// that subsequent session functions do not need to handle redirects.
746 fn getCapabilities(session: *Session, http_headers_buffer: []u8) !CapabilityIterator {747 fn getCapabilities(session: *Session, it: *CapabilityIterator, response_buffer: []u8) !void {
748 assert(response_buffer.len >= Packet.max_data_length);
747 var info_refs_uri = session.location.uri;749 var info_refs_uri = session.location.uri;
748 {750 {
749 const session_uri_path = try std.fmt.allocPrint(session.allocator, "{f}", .{751 const session_uri_path = try std.fmt.allocPrint(session.allocator, "{f}", .{
...@@ -757,19 +759,22 @@ pub const Session = struct {...@@ -757,19 +759,22 @@ pub const Session = struct {
757 info_refs_uri.fragment = null;759 info_refs_uri.fragment = null;
758760
759 const max_redirects = 3;761 const max_redirects = 3;
760 var request = try session.transport.open(.GET, info_refs_uri, .{762 it.* = .{
761 .redirect_behavior = @enumFromInt(max_redirects),763 .request = try session.transport.request(.GET, info_refs_uri, .{
762 .server_header_buffer = http_headers_buffer,764 .redirect_behavior = .init(max_redirects),
763 .extra_headers = &.{765 .extra_headers = &.{
764 .{ .name = "Git-Protocol", .value = "version=2" },766 .{ .name = "Git-Protocol", .value = "version=2" },
765 },767 },
766 });768 }),
767 errdefer request.deinit();769 .reader = undefined,
768 try request.send();770 };
769 try request.finish();771 errdefer it.deinit();
772 const request = &it.request;
773 try request.sendBodiless();
770774
771 try request.wait();775 var redirect_buffer: [1024]u8 = undefined;
772 if (request.response.status != .ok) return error.ProtocolError;776 var response = try request.receiveHead(&redirect_buffer);
777 if (response.head.status != .ok) return error.ProtocolError;
773 const any_redirects_occurred = request.redirect_behavior.remaining() < max_redirects;778 const any_redirects_occurred = request.redirect_behavior.remaining() < max_redirects;
774 if (any_redirects_occurred) {779 if (any_redirects_occurred) {
775 const request_uri_path = try std.fmt.allocPrint(session.allocator, "{f}", .{780 const request_uri_path = try std.fmt.allocPrint(session.allocator, "{f}", .{
...@@ -784,8 +789,7 @@ pub const Session = struct {...@@ -784,8 +789,7 @@ pub const Session = struct {
784 session.location = new_location;789 session.location = new_location;
785 }790 }
786791
787 const reader = request.reader();792 it.reader = response.reader(response_buffer);
788 var buf: [Packet.max_data_length]u8 = undefined;
789 var state: enum { response_start, response_content } = .response_start;793 var state: enum { response_start, response_content } = .response_start;
790 while (true) {794 while (true) {
791 // Some Git servers (at least GitHub) include an additional795 // Some Git servers (at least GitHub) include an additional
...@@ -795,15 +799,15 @@ pub const Session = struct {...@@ -795,15 +799,15 @@ pub const Session = struct {
795 // Thus, we need to skip any such useless additional responses799 // Thus, we need to skip any such useless additional responses
796 // before we get the one we're actually looking for. The responses800 // before we get the one we're actually looking for. The responses
797 // will be delimited by flush packets.801 // will be delimited by flush packets.
798 const packet = Packet.read(reader, &buf) catch |e| switch (e) {802 const packet = Packet.read(it.reader) catch |err| switch (err) {
799 error.EndOfStream => return error.UnsupportedProtocol, // 'version 2' packet not found803 error.EndOfStream => return error.UnsupportedProtocol, // 'version 2' packet not found
800 else => |other| return other,804 else => |e| return e,
801 };805 };
802 switch (packet) {806 switch (packet) {
803 .flush => state = .response_start,807 .flush => state = .response_start,
804 .data => |data| switch (state) {808 .data => |data| switch (state) {
805 .response_start => if (mem.eql(u8, Packet.normalizeText(data), "version 2")) {809 .response_start => if (mem.eql(u8, Packet.normalizeText(data), "version 2")) {
806 return .{ .request = request };810 return;
807 } else {811 } else {
808 state = .response_content;812 state = .response_content;
809 },813 },
...@@ -816,7 +820,7 @@ pub const Session = struct {...@@ -816,7 +820,7 @@ pub const Session = struct {
816820
817 const CapabilityIterator = struct {821 const CapabilityIterator = struct {
818 request: std.http.Client.Request,822 request: std.http.Client.Request,
819 buf: [Packet.max_data_length]u8 = undefined,823 reader: *std.Io.Reader,
820824
821 const Capability = struct {825 const Capability = struct {
822 key: []const u8,826 key: []const u8,
...@@ -830,13 +834,13 @@ pub const Session = struct {...@@ -830,13 +834,13 @@ pub const Session = struct {
830 }834 }
831 };835 };
832836
833 fn deinit(iterator: *CapabilityIterator) void {837 fn deinit(it: *CapabilityIterator) void {
834 iterator.request.deinit();838 it.request.deinit();
835 iterator.* = undefined;839 it.* = undefined;
836 }840 }
837841
838 fn next(iterator: *CapabilityIterator) !?Capability {842 fn next(it: *CapabilityIterator) !?Capability {
839 switch (try Packet.read(iterator.request.reader(), &iterator.buf)) {843 switch (try Packet.read(it.reader)) {
840 .flush => return null,844 .flush => return null,
841 .data => |data| return Capability.parse(Packet.normalizeText(data)),845 .data => |data| return Capability.parse(Packet.normalizeText(data)),
842 else => return error.UnexpectedPacket,846 else => return error.UnexpectedPacket,
...@@ -854,11 +858,13 @@ pub const Session = struct {...@@ -854,11 +858,13 @@ pub const Session = struct {
854 include_symrefs: bool = false,858 include_symrefs: bool = false,
855 /// Whether to include the peeled object ID for returned tag refs.859 /// Whether to include the peeled object ID for returned tag refs.
856 include_peeled: bool = false,860 include_peeled: bool = false,
857 server_header_buffer: []u8,861 /// Asserted to be at least `Packet.max_data_length`.
862 buffer: []u8,
858 };863 };
859864
860 /// Returns an iterator over refs known to the server.865 /// Returns an iterator over refs known to the server.
861 pub fn listRefs(session: Session, options: ListRefsOptions) !RefIterator {866 pub fn listRefs(session: Session, it: *RefIterator, options: ListRefsOptions) !void {
867 assert(options.buffer.len >= Packet.max_data_length);
862 var upload_pack_uri = session.location.uri;868 var upload_pack_uri = session.location.uri;
863 {869 {
864 const session_uri_path = try std.fmt.allocPrint(session.allocator, "{f}", .{870 const session_uri_path = try std.fmt.allocPrint(session.allocator, "{f}", .{
...@@ -871,59 +877,56 @@ pub const Session = struct {...@@ -871,59 +877,56 @@ pub const Session = struct {
871 upload_pack_uri.query = null;877 upload_pack_uri.query = null;
872 upload_pack_uri.fragment = null;878 upload_pack_uri.fragment = null;
873879
874 var body: std.ArrayListUnmanaged(u8) = .empty;880 var body: std.Io.Writer = .fixed(options.buffer);
875 defer body.deinit(session.allocator);881 try Packet.write(.{ .data = "command=ls-refs\n" }, &body);
876 const body_writer = body.writer(session.allocator);
877 try Packet.write(.{ .data = "command=ls-refs\n" }, body_writer);
878 if (session.supports_agent) {882 if (session.supports_agent) {
879 try Packet.write(.{ .data = agent_capability }, body_writer);883 try Packet.write(.{ .data = agent_capability }, &body);
880 }884 }
881 {885 {
882 const object_format_packet = try std.fmt.allocPrint(session.allocator, "object-format={s}\n", .{@tagName(session.object_format)});886 const object_format_packet = try std.fmt.allocPrint(session.allocator, "object-format={t}\n", .{
887 session.object_format,
888 });
883 defer session.allocator.free(object_format_packet);889 defer session.allocator.free(object_format_packet);
884 try Packet.write(.{ .data = object_format_packet }, body_writer);890 try Packet.write(.{ .data = object_format_packet }, &body);
885 }891 }
886 try Packet.write(.delimiter, body_writer);892 try Packet.write(.delimiter, &body);
887 for (options.ref_prefixes) |ref_prefix| {893 for (options.ref_prefixes) |ref_prefix| {
888 const ref_prefix_packet = try std.fmt.allocPrint(session.allocator, "ref-prefix {s}\n", .{ref_prefix});894 const ref_prefix_packet = try std.fmt.allocPrint(session.allocator, "ref-prefix {s}\n", .{ref_prefix});
889 defer session.allocator.free(ref_prefix_packet);895 defer session.allocator.free(ref_prefix_packet);
890 try Packet.write(.{ .data = ref_prefix_packet }, body_writer);896 try Packet.write(.{ .data = ref_prefix_packet }, &body);
891 }897 }
892 if (options.include_symrefs) {898 if (options.include_symrefs) {
893 try Packet.write(.{ .data = "symrefs\n" }, body_writer);899 try Packet.write(.{ .data = "symrefs\n" }, &body);
894 }900 }
895 if (options.include_peeled) {901 if (options.include_peeled) {
896 try Packet.write(.{ .data = "peel\n" }, body_writer);902 try Packet.write(.{ .data = "peel\n" }, &body);
897 }903 }
898 try Packet.write(.flush, body_writer);904 try Packet.write(.flush, &body);
899905
900 var request = try session.transport.open(.POST, upload_pack_uri, .{906 it.* = .{
901 .redirect_behavior = .unhandled,907 .request = try session.transport.request(.POST, upload_pack_uri, .{
902 .server_header_buffer = options.server_header_buffer,908 .redirect_behavior = .unhandled,
903 .extra_headers = &.{909 .extra_headers = &.{
904 .{ .name = "Content-Type", .value = "application/x-git-upload-pack-request" },910 .{ .name = "Content-Type", .value = "application/x-git-upload-pack-request" },
905 .{ .name = "Git-Protocol", .value = "version=2" },911 .{ .name = "Git-Protocol", .value = "version=2" },
906 },912 },
907 });913 }),
908 errdefer request.deinit();914 .reader = undefined,
909 request.transfer_encoding = .{ .content_length = body.items.len };
910 try request.send();
911 try request.writeAll(body.items);
912 try request.finish();
913
914 try request.wait();
915 if (request.response.status != .ok) return error.ProtocolError;
916
917 return .{
918 .format = session.object_format,915 .format = session.object_format,
919 .request = request,
920 };916 };
917 const request = &it.request;
918 errdefer request.deinit();
919 try request.sendBodyComplete(body.buffered());
920
921 var response = try request.receiveHead(options.buffer);
922 if (response.head.status != .ok) return error.ProtocolError;
923 it.reader = response.reader(options.buffer);
921 }924 }
922925
923 pub const RefIterator = struct {926 pub const RefIterator = struct {
924 format: Oid.Format,927 format: Oid.Format,
925 request: std.http.Client.Request,928 request: std.http.Client.Request,
926 buf: [Packet.max_data_length]u8 = undefined,929 reader: *std.Io.Reader,
927930
928 pub const Ref = struct {931 pub const Ref = struct {
929 oid: Oid,932 oid: Oid,
...@@ -937,13 +940,13 @@ pub const Session = struct {...@@ -937,13 +940,13 @@ pub const Session = struct {
937 iterator.* = undefined;940 iterator.* = undefined;
938 }941 }
939942
940 pub fn next(iterator: *RefIterator) !?Ref {943 pub fn next(it: *RefIterator) !?Ref {
941 switch (try Packet.read(iterator.request.reader(), &iterator.buf)) {944 switch (try Packet.read(it.reader)) {
942 .flush => return null,945 .flush => return null,
943 .data => |data| {946 .data => |data| {
944 const ref_data = Packet.normalizeText(data);947 const ref_data = Packet.normalizeText(data);
945 const oid_sep_pos = mem.indexOfScalar(u8, ref_data, ' ') orelse return error.InvalidRefPacket;948 const oid_sep_pos = mem.indexOfScalar(u8, ref_data, ' ') orelse return error.InvalidRefPacket;
946 const oid = Oid.parse(iterator.format, data[0..oid_sep_pos]) catch return error.InvalidRefPacket;949 const oid = Oid.parse(it.format, data[0..oid_sep_pos]) catch return error.InvalidRefPacket;
947950
948 const name_sep_pos = mem.indexOfScalarPos(u8, ref_data, oid_sep_pos + 1, ' ') orelse ref_data.len;951 const name_sep_pos = mem.indexOfScalarPos(u8, ref_data, oid_sep_pos + 1, ' ') orelse ref_data.len;
949 const name = ref_data[oid_sep_pos + 1 .. name_sep_pos];952 const name = ref_data[oid_sep_pos + 1 .. name_sep_pos];
...@@ -957,7 +960,7 @@ pub const Session = struct {...@@ -957,7 +960,7 @@ pub const Session = struct {
957 if (mem.startsWith(u8, attribute, "symref-target:")) {960 if (mem.startsWith(u8, attribute, "symref-target:")) {
958 symref_target = attribute["symref-target:".len..];961 symref_target = attribute["symref-target:".len..];
959 } else if (mem.startsWith(u8, attribute, "peeled:")) {962 } else if (mem.startsWith(u8, attribute, "peeled:")) {
960 peeled = Oid.parse(iterator.format, attribute["peeled:".len..]) catch return error.InvalidRefPacket;963 peeled = Oid.parse(it.format, attribute["peeled:".len..]) catch return error.InvalidRefPacket;
961 }964 }
962 last_sep_pos = next_sep_pos;965 last_sep_pos = next_sep_pos;
963 }966 }
...@@ -973,9 +976,12 @@ pub const Session = struct {...@@ -973,9 +976,12 @@ pub const Session = struct {
973 /// performed if the server supports it.976 /// performed if the server supports it.
974 pub fn fetch(977 pub fn fetch(
975 session: Session,978 session: Session,
979 fs: *FetchStream,
976 wants: []const []const u8,980 wants: []const []const u8,
977 http_headers_buffer: []u8,981 /// Asserted to be at least `Packet.max_data_length`.
978 ) !FetchStream {982 response_buffer: []u8,
983 ) !void {
984 assert(response_buffer.len >= Packet.max_data_length);
979 var upload_pack_uri = session.location.uri;985 var upload_pack_uri = session.location.uri;
980 {986 {
981 const session_uri_path = try std.fmt.allocPrint(session.allocator, "{f}", .{987 const session_uri_path = try std.fmt.allocPrint(session.allocator, "{f}", .{
...@@ -988,63 +994,71 @@ pub const Session = struct {...@@ -988,63 +994,71 @@ pub const Session = struct {
988 upload_pack_uri.query = null;994 upload_pack_uri.query = null;
989 upload_pack_uri.fragment = null;995 upload_pack_uri.fragment = null;
990996
991 var body: std.ArrayListUnmanaged(u8) = .empty;997 var body: std.Io.Writer = .fixed(response_buffer);
992 defer body.deinit(session.allocator);998 try Packet.write(.{ .data = "command=fetch\n" }, &body);
993 const body_writer = body.writer(session.allocator);
994 try Packet.write(.{ .data = "command=fetch\n" }, body_writer);
995 if (session.supports_agent) {999 if (session.supports_agent) {
996 try Packet.write(.{ .data = agent_capability }, body_writer);1000 try Packet.write(.{ .data = agent_capability }, &body);
997 }1001 }
998 {1002 {
999 const object_format_packet = try std.fmt.allocPrint(session.allocator, "object-format={s}\n", .{@tagName(session.object_format)});1003 const object_format_packet = try std.fmt.allocPrint(session.allocator, "object-format={s}\n", .{@tagName(session.object_format)});
1000 defer session.allocator.free(object_format_packet);1004 defer session.allocator.free(object_format_packet);
1001 try Packet.write(.{ .data = object_format_packet }, body_writer);1005 try Packet.write(.{ .data = object_format_packet }, &body);
1002 }1006 }
1003 try Packet.write(.delimiter, body_writer);1007 try Packet.write(.delimiter, &body);
1004 // Our packfile parser supports the OFS_DELTA object type1008 // Our packfile parser supports the OFS_DELTA object type
1005 try Packet.write(.{ .data = "ofs-delta\n" }, body_writer);1009 try Packet.write(.{ .data = "ofs-delta\n" }, &body);
1006 // We do not currently convey server progress information to the user1010 // We do not currently convey server progress information to the user
1007 try Packet.write(.{ .data = "no-progress\n" }, body_writer);1011 try Packet.write(.{ .data = "no-progress\n" }, &body);
1008 if (session.supports_shallow) {1012 if (session.supports_shallow) {
1009 try Packet.write(.{ .data = "deepen 1\n" }, body_writer);1013 try Packet.write(.{ .data = "deepen 1\n" }, &body);
1010 }1014 }
1011 for (wants) |want| {1015 for (wants) |want| {
1012 var buf: [Packet.max_data_length]u8 = undefined;1016 var buf: [Packet.max_data_length]u8 = undefined;
1013 const arg = std.fmt.bufPrint(&buf, "want {s}\n", .{want}) catch unreachable;1017 const arg = std.fmt.bufPrint(&buf, "want {s}\n", .{want}) catch unreachable;
1014 try Packet.write(.{ .data = arg }, body_writer);1018 try Packet.write(.{ .data = arg }, &body);
1015 }1019 }
1016 try Packet.write(.{ .data = "done\n" }, body_writer);1020 try Packet.write(.{ .data = "done\n" }, &body);
1017 try Packet.write(.flush, body_writer);1021 try Packet.write(.flush, &body);
10181022
1019 var request = try session.transport.open(.POST, upload_pack_uri, .{1023 fs.* = .{
1020 .redirect_behavior = .not_allowed,1024 .request = try session.transport.request(.POST, upload_pack_uri, .{
1021 .server_header_buffer = http_headers_buffer,1025 .redirect_behavior = .not_allowed,
1022 .extra_headers = &.{1026 .extra_headers = &.{
1023 .{ .name = "Content-Type", .value = "application/x-git-upload-pack-request" },1027 .{ .name = "Content-Type", .value = "application/x-git-upload-pack-request" },
1024 .{ .name = "Git-Protocol", .value = "version=2" },1028 .{ .name = "Git-Protocol", .value = "version=2" },
1025 },1029 },
1026 });1030 }),
1031 .input = undefined,
1032 .reader = undefined,
1033 .remaining_len = undefined,
1034 };
1035 const request = &fs.request;
1027 errdefer request.deinit();1036 errdefer request.deinit();
1028 request.transfer_encoding = .{ .content_length = body.items.len };
1029 try request.send();
1030 try request.writeAll(body.items);
1031 try request.finish();
10321037
1033 try request.wait();1038 try request.sendBodyComplete(body.buffered());
1034 if (request.response.status != .ok) return error.ProtocolError;1039
1040 var response = try request.receiveHead(&.{});
1041 if (response.head.status != .ok) return error.ProtocolError;
10351042
1036 const reader = request.reader();1043 const reader = response.reader(response_buffer);
1037 // We are not interested in any of the sections of the returned fetch1044 // We are not interested in any of the sections of the returned fetch
1038 // data other than the packfile section, since we aren't doing anything1045 // data other than the packfile section, since we aren't doing anything
1039 // complex like ref negotiation (this is a fresh clone).1046 // complex like ref negotiation (this is a fresh clone).
1040 var state: enum { section_start, section_content } = .section_start;1047 var state: enum { section_start, section_content } = .section_start;
1041 while (true) {1048 while (true) {
1042 var buf: [Packet.max_data_length]u8 = undefined;1049 const packet = try Packet.read(reader);
1043 const packet = try Packet.read(reader, &buf);
1044 switch (state) {1050 switch (state) {
1045 .section_start => switch (packet) {1051 .section_start => switch (packet) {
1046 .data => |data| if (mem.eql(u8, Packet.normalizeText(data), "packfile")) {1052 .data => |data| if (mem.eql(u8, Packet.normalizeText(data), "packfile")) {
1047 return .{ .request = request };1053 fs.input = reader;
1054 fs.reader = .{
1055 .buffer = &.{},
1056 .vtable = &.{ .stream = FetchStream.stream },
1057 .seek = 0,
1058 .end = 0,
1059 };
1060 fs.remaining_len = 0;
1061 return;
1048 } else {1062 } else {
1049 state = .section_content;1063 state = .section_content;
1050 },1064 },
...@@ -1061,20 +1075,23 @@ pub const Session = struct {...@@ -1061,20 +1075,23 @@ pub const Session = struct {
10611075
1062 pub const FetchStream = struct {1076 pub const FetchStream = struct {
1063 request: std.http.Client.Request,1077 request: std.http.Client.Request,
1064 buf: [Packet.max_data_length]u8 = undefined,1078 input: *std.Io.Reader,
1065 pos: usize = 0,1079 reader: std.Io.Reader,
1066 len: usize = 0,1080 err: ?Error = null,
1081 remaining_len: usize,
10671082
1068 pub fn deinit(stream: *FetchStream) void {1083 pub fn deinit(fs: *FetchStream) void {
1069 stream.request.deinit();1084 fs.request.deinit();
1070 }1085 }
10711086
1072 pub const ReadError = std.http.Client.Request.ReadError || error{1087 pub const Error = error{
1073 InvalidPacket,1088 InvalidPacket,
1074 ProtocolError,1089 ProtocolError,
1075 UnexpectedPacket,1090 UnexpectedPacket,
1091 WriteFailed,
1092 ReadFailed,
1093 EndOfStream,
1076 };1094 };
1077 pub const Reader = std.io.GenericReader(*FetchStream, ReadError, read);
10781095
1079 const StreamCode = enum(u8) {1096 const StreamCode = enum(u8) {
1080 pack_data = 1,1097 pack_data = 1,
...@@ -1083,33 +1100,41 @@ pub const Session = struct {...@@ -1083,33 +1100,41 @@ pub const Session = struct {
1083 _,1100 _,
1084 };1101 };
10851102
1086 pub fn reader(stream: *FetchStream) Reader {1103 pub fn stream(r: *std.Io.Reader, w: *std.Io.Writer, limit: std.Io.Limit) std.Io.Reader.StreamError!usize {
1087 return .{ .context = stream };1104 const fs: *FetchStream = @alignCast(@fieldParentPtr("reader", r));
1088 }1105 const input = fs.input;
10891106 if (fs.remaining_len == 0) {
1090 pub fn read(stream: *FetchStream, buf: []u8) !usize {
1091 if (stream.pos == stream.len) {
1092 while (true) {1107 while (true) {
1093 switch (try Packet.read(stream.request.reader(), &stream.buf)) {1108 switch (Packet.read(input) catch |err| {
1094 .flush => return 0,1109 fs.err = err;
1110 return error.ReadFailed;
1111 }) {
1112 .flush => return error.EndOfStream,
1095 .data => |data| if (data.len > 1) switch (@as(StreamCode, @enumFromInt(data[0]))) {1113 .data => |data| if (data.len > 1) switch (@as(StreamCode, @enumFromInt(data[0]))) {
1096 .pack_data => {1114 .pack_data => {
1097 stream.pos = 1;1115 input.toss(1);
1098 stream.len = data.len;1116 fs.remaining_len = data.len;
1099 break;1117 break;
1100 },1118 },
1101 .fatal_error => return error.ProtocolError,1119 .fatal_error => {
1120 fs.err = error.ProtocolError;
1121 return error.ReadFailed;
1122 },
1102 else => {},1123 else => {},
1103 },1124 },
1104 else => return error.UnexpectedPacket,1125 else => {
1126 fs.err = error.UnexpectedPacket;
1127 return error.ReadFailed;
1128 },
1105 }1129 }
1106 }1130 }
1107 }1131 }
11081132 const buf = limit.slice(try w.writableSliceGreedy(1));
1109 const size = @min(buf.len, stream.len - stream.pos);1133 const n = @min(buf.len, fs.remaining_len);
1110 @memcpy(buf[0..size], stream.buf[stream.pos .. stream.pos + size]);1134 @memcpy(buf[0..n], input.buffered()[0..n]);
1111 stream.pos += size;1135 input.toss(n);
1112 return size;1136 fs.remaining_len -= n;
1137 return n;
1113 }1138 }
1114 };1139 };
1115};1140};