authorgravatar for mitchell.hashimoto@gmail.comMitchell Hashimoto <mitchell.hashimoto@gmail.com> 2022-03-08 20:44:58-08:00
committergravatar for andrew@ziglang.orgAndrew Kelley <andrew@ziglang.org> 2022-03-10 14:20:16-07:00
log569870ca41e73c64d8dc9f1eccfef3529caf2266
tree8a4ec47628cacc1723efbff01ba36ebc147a86ad
parent0b82c02945c69e2e0465b5a4d9de471ea3c76d50

stage2: error_set_merged type equality

This implements type equality for error sets. This is done through element-wise error set comparison. Inferred error sets are always distinct types and other error sets are always sorted. See #11022.

8 files changed, 185 insertions(+), 58 deletions(-)

src/Module.zig+13-1
......@@ -824,7 +824,7 @@ pub const ErrorSet = struct {
824824 /// Offset from Decl node index, points to the error set AST node.
825825 node_offset: i32,
826826 /// The string bytes are stored in the owner Decl arena.
827 /// They are in the same order they appear in the AST.
827 /// These must be in sorted order. See sortNames.
828828 names: NameMap,
829829
830830 pub const NameMap = std.StringArrayHashMapUnmanaged(void);
......@@ -836,6 +836,18 @@ pub const ErrorSet = struct {
836836 .lazy = .{ .node_offset = self.node_offset },
837837 };
838838 }
839
840 /// sort the NameMap. This should be called whenever the map is modified.
841 /// alloc should be the allocator used for the NameMap data.
842 pub fn sortNames(names: *NameMap) void {
843 const Context = struct {
844 keys: [][]const u8,
845 pub fn lessThan(ctx: @This(), a_index: usize, b_index: usize) bool {
846 return std.mem.lessThan(u8, ctx.keys[a_index], ctx.keys[b_index]);
847 }
848 };
849 names.sort(Context{ .keys = names.keys() });
850 }
839851};
840852
841853pub const RequiresComptime = enum { no, yes, unknown, wip };
src/Sema.zig+4
......@@ -2212,6 +2212,10 @@ fn zirErrorSetDecl(
22122212 return sema.fail(block, src, "duplicate error set field {s}", .{name});
22132213 }
22142214 }
2215
2216 // names must be sorted.
2217 Module.ErrorSet.sortNames(&names);
2218
22152219 error_set.* = .{
22162220 .owner_decl = new_decl,
22172221 .node_offset = inst_data.src_node,
src/type.zig+45-21
......@@ -564,27 +564,30 @@ pub const Type = extern union {
564564 => {
565565 if (b.zigTypeTag() != .ErrorSet) return false;
566566
567 // TODO: revisit the language specification for how to evaluate equality
568 // for error set types.
569
570 if (a.tag() == .anyerror and b.tag() == .anyerror) {
571 return true;
567 // inferred error sets are only equal if both are inferred
568 // and they originate from the exact same function.
569 if (a.castTag(.error_set_inferred)) |a_pl| {
570 if (b.castTag(.error_set_inferred)) |b_pl| {
571 return a_pl.data.func == b_pl.data.func;
572 }
573 return false;
572574 }
573
574 if (a.tag() == .error_set and b.tag() == .error_set) {
575 return a.castTag(.error_set).?.data.owner_decl == b.castTag(.error_set).?.data.owner_decl;
575 if (b.tag() == .error_set_inferred) return false;
576
577 // anyerror matches exactly.
578 const a_is_any = a.isAnyError();
579 const b_is_any = b.isAnyError();
580 if (a_is_any or b_is_any) return a_is_any and b_is_any;
581
582 // two resolved sets match if their error set names match.
583 const a_set = a.errorSetNames();
584 const b_set = b.errorSetNames();
585 if (a_set.len != b_set.len) return false;
586 for (b_set) |b_val| {
587 if (!a.errorSetHasField(b_val)) return false;
576588 }
577589
578 if (a.tag() == .error_set_inferred and b.tag() == .error_set_inferred) {
579 return a.castTag(.error_set_inferred).?.data == b.castTag(.error_set_inferred).?.data;
580 }
581
582 if (a.tag() == .error_set_single and b.tag() == .error_set_single) {
583 const a_data = a.castTag(.error_set_single).?.data;
584 const b_data = b.castTag(.error_set_single).?.data;
585 return std.mem.eql(u8, a_data, b_data);
586 }
587 return false;
590 return true;
588591 },
589592
590593 .@"opaque" => {
......@@ -961,12 +964,30 @@ pub const Type = extern union {
961964
962965 .error_set,
963966 .error_set_single,
964 .anyerror,
965 .error_set_inferred,
966967 .error_set_merged,
967968 => {
969 // all are treated like an "error set" for hashing
970 std.hash.autoHash(hasher, std.builtin.TypeId.ErrorSet);
971 std.hash.autoHash(hasher, Tag.error_set);
972
973 const names = ty.errorSetNames();
974 std.hash.autoHash(hasher, names.len);
975 assert(std.sort.isSorted([]const u8, names, u8, std.mem.lessThan));
976 for (names) |name| hasher.update(name);
977 },
978
979 .anyerror => {
980 // anyerror is distinct from other error sets
968981 std.hash.autoHash(hasher, std.builtin.TypeId.ErrorSet);
969 // TODO implement this after revisiting Type.Eql for error sets
982 std.hash.autoHash(hasher, Tag.anyerror);
983 },
984
985 .error_set_inferred => {
986 // inferred error sets are compared using their data pointer
987 const data = ty.castTag(.error_set_inferred).?.data.func;
988 std.hash.autoHash(hasher, std.builtin.TypeId.ErrorSet);
989 std.hash.autoHash(hasher, Tag.error_set_inferred);
990 std.hash.autoHash(hasher, data);
970991 },
971992
972993 .@"opaque" => {
......@@ -4365,6 +4386,9 @@ pub const Type = extern union {
43654386 try names.put(arena, name, {});
43664387 }
43674388
4389 // names must be sorted
4390 Module.ErrorSet.sortNames(&names);
4391
43684392 return try Tag.error_set_merged.create(arena, names);
43694393 }
43704394
src/value.zig+10
......@@ -1870,6 +1870,16 @@ pub const Value = extern union {
18701870
18711871 return eql(a_payload.container_ptr, b_payload.container_ptr, ty);
18721872 },
1873 .@"error" => {
1874 const a_name = a.castTag(.@"error").?.data.name;
1875 const b_name = b.castTag(.@"error").?.data.name;
1876 return std.mem.eql(u8, a_name, b_name);
1877 },
1878 .eu_payload => {
1879 const a_payload = a.castTag(.eu_payload).?.data;
1880 const b_payload = b.castTag(.eu_payload).?.data;
1881 return eql(a_payload, b_payload, ty.errorUnionPayload());
1882 },
18731883 .eu_payload_ptr => @panic("TODO: Implement more pointer eql cases"),
18741884 .opt_payload_ptr => @panic("TODO: Implement more pointer eql cases"),
18751885 .array => {
test/behavior/cast.zig+8-8
......@@ -669,8 +669,8 @@ test "peer type resolution: disjoint error sets" {
669669 try expect(error_set_info == .ErrorSet);
670670 try expect(error_set_info.ErrorSet.?.len == 3);
671671 try expect(mem.eql(u8, error_set_info.ErrorSet.?[0].name, "One"));
672 try expect(mem.eql(u8, error_set_info.ErrorSet.?[1].name, "Two"));
673 try expect(mem.eql(u8, error_set_info.ErrorSet.?[2].name, "Three"));
672 try expect(mem.eql(u8, error_set_info.ErrorSet.?[1].name, "Three"));
673 try expect(mem.eql(u8, error_set_info.ErrorSet.?[2].name, "Two"));
674674 }
675675
676676 {
......@@ -678,8 +678,8 @@ test "peer type resolution: disjoint error sets" {
678678 const error_set_info = @typeInfo(ty);
679679 try expect(error_set_info == .ErrorSet);
680680 try expect(error_set_info.ErrorSet.?.len == 3);
681 try expect(mem.eql(u8, error_set_info.ErrorSet.?[0].name, "Three"));
682 try expect(mem.eql(u8, error_set_info.ErrorSet.?[1].name, "One"));
681 try expect(mem.eql(u8, error_set_info.ErrorSet.?[0].name, "One"));
682 try expect(mem.eql(u8, error_set_info.ErrorSet.?[1].name, "Three"));
683683 try expect(mem.eql(u8, error_set_info.ErrorSet.?[2].name, "Two"));
684684 }
685685}
......@@ -704,8 +704,8 @@ test "peer type resolution: error union and error set" {
704704
705705 const error_set_info = @typeInfo(info.ErrorUnion.error_set);
706706 try expect(error_set_info.ErrorSet.?.len == 3);
707 try expect(mem.eql(u8, error_set_info.ErrorSet.?[0].name, "Three"));
708 try expect(mem.eql(u8, error_set_info.ErrorSet.?[1].name, "One"));
707 try expect(mem.eql(u8, error_set_info.ErrorSet.?[0].name, "One"));
708 try expect(mem.eql(u8, error_set_info.ErrorSet.?[1].name, "Three"));
709709 try expect(mem.eql(u8, error_set_info.ErrorSet.?[2].name, "Two"));
710710 }
711711
......@@ -717,8 +717,8 @@ test "peer type resolution: error union and error set" {
717717 const error_set_info = @typeInfo(info.ErrorUnion.error_set);
718718 try expect(error_set_info.ErrorSet.?.len == 3);
719719 try expect(mem.eql(u8, error_set_info.ErrorSet.?[0].name, "One"));
720 try expect(mem.eql(u8, error_set_info.ErrorSet.?[1].name, "Two"));
721 try expect(mem.eql(u8, error_set_info.ErrorSet.?[2].name, "Three"));
720 try expect(mem.eql(u8, error_set_info.ErrorSet.?[1].name, "Three"));
721 try expect(mem.eql(u8, error_set_info.ErrorSet.?[2].name, "Two"));
722722 }
723723}
724724
test/behavior/error.zig+100-7
......@@ -330,7 +330,11 @@ fn intLiteral(str: []const u8) !?i64 {
330330}
331331
332332test "nested error union function call in optional unwrap" {
333 if (builtin.zig_backend != .stage1) return error.SkipZigTest; // TODO
333 if (builtin.zig_backend == .stage2_c) return error.SkipZigTest; // TODO
334 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
335 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest;
336 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest;
337 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest;
334338
335339 const S = struct {
336340 const Foo = struct {
......@@ -375,7 +379,11 @@ test "nested error union function call in optional unwrap" {
375379}
376380
377381test "return function call to error set from error union function" {
378 if (builtin.zig_backend != .stage1) return error.SkipZigTest; // TODO
382 if (builtin.zig_backend == .stage2_c) return error.SkipZigTest; // TODO
383 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
384 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest;
385 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest;
386 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest;
379387
380388 const S = struct {
381389 fn errorable() anyerror!i32 {
......@@ -404,7 +412,11 @@ test "optional error set is the same size as error set" {
404412}
405413
406414test "nested catch" {
407 if (builtin.zig_backend != .stage1) return error.SkipZigTest; // TODO
415 if (builtin.zig_backend == .stage2_c) return error.SkipZigTest; // TODO
416 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
417 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest;
418 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest;
419 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest;
408420
409421 const S = struct {
410422 fn entry() !void {
......@@ -428,11 +440,18 @@ test "nested catch" {
428440}
429441
430442test "function pointer with return type that is error union with payload which is pointer of parent struct" {
431 if (builtin.zig_backend != .stage1) return error.SkipZigTest; // TODO
443 // This test uses the stage2 const fn pointer
444 if (builtin.zig_backend == .stage1) return error.SkipZigTest;
445
446 if (builtin.zig_backend == .stage2_c) return error.SkipZigTest; // TODO
447 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
448 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest;
449 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest;
450 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest;
432451
433452 const S = struct {
434453 const Foo = struct {
435 fun: fn (a: i32) (anyerror!*Foo),
454 fun: *const fn (a: i32) (anyerror!*Foo),
436455 };
437456
438457 const Err = error{UnspecifiedErr};
......@@ -480,7 +499,11 @@ test "return result loc as peer result loc in inferred error set function" {
480499}
481500
482501test "error payload type is correctly resolved" {
483 if (builtin.zig_backend != .stage1) return error.SkipZigTest; // TODO
502 if (builtin.zig_backend == .stage2_c) return error.SkipZigTest; // TODO
503 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
504 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest;
505 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest;
506 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest;
484507
485508 const MyIntWrapper = struct {
486509 const Self = @This();
......@@ -496,7 +519,11 @@ test "error payload type is correctly resolved" {
496519}
497520
498521test "error union comptime caching" {
499 if (builtin.zig_backend != .stage1) return error.SkipZigTest; // TODO
522 if (builtin.zig_backend == .stage2_c) return error.SkipZigTest; // TODO
523 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
524 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest;
525 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest;
526 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest;
500527
501528 const S = struct {
502529 fn quux(comptime arg: anytype) void {
......@@ -539,3 +566,69 @@ test "@errorName sentinel length matches slice length" {
539566pub fn testBuiltinErrorName(err: anyerror) [:0]const u8 {
540567 return @errorName(err);
541568}
569
570test "error set equality" {
571 // This tests using stage2 logic (#11022)
572 if (builtin.zig_backend == .stage1) return error.SkipZigTest;
573
574 if (builtin.zig_backend == .stage2_c) return error.SkipZigTest; // TODO
575 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
576 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest;
577 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest;
578 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest;
579
580 const a = error{One};
581 const b = error{One};
582
583 try expect(a == a);
584 try expect(a == b);
585 try expect(a == error{One});
586
587 // should treat as a set
588 const c = error{ One, Two };
589 const d = error{ Two, One };
590
591 try expect(c == d);
592}
593
594test "inferred error set equality" {
595 if (builtin.zig_backend == .stage2_c) return error.SkipZigTest; // TODO
596 if (builtin.zig_backend == .stage2_wasm) return error.SkipZigTest; // TODO
597 if (builtin.zig_backend == .stage2_x86_64) return error.SkipZigTest;
598 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest;
599 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest;
600
601 const S = struct {
602 fn foo() !void {
603 return @This().bar();
604 }
605
606 fn bar() !void {
607 return error.Bad;
608 }
609
610 fn baz() !void {
611 return quux();
612 }
613
614 fn quux() anyerror!void {}
615 };
616
617 const FooError = @typeInfo(@typeInfo(@TypeOf(S.foo)).Fn.return_type.?).ErrorUnion.error_set;
618 const BarError = @typeInfo(@typeInfo(@TypeOf(S.bar)).Fn.return_type.?).ErrorUnion.error_set;
619 const BazError = @typeInfo(@typeInfo(@TypeOf(S.baz)).Fn.return_type.?).ErrorUnion.error_set;
620
621 try expect(BarError != error{Bad});
622
623 try expect(FooError != anyerror);
624 try expect(BarError != anyerror);
625 try expect(BazError != anyerror);
626
627 try expect(FooError != BarError);
628 try expect(FooError != BazError);
629 try expect(BarError != BazError);
630
631 try expect(FooError == FooError);
632 try expect(BarError == BarError);
633 try expect(BazError == BazError);
634}
test/behavior/type_info.zig+5-2
......@@ -205,6 +205,9 @@ test "type info: error set single value" {
205205}
206206
207207test "type info: error set merged" {
208 // #11022 forces ordering of error sets in stage2
209 if (builtin.zig_backend == .stage1) return error.SkipZigTest;
210
208211 if (builtin.zig_backend == .stage2_aarch64) return error.SkipZigTest;
209212 if (builtin.zig_backend == .stage2_c) return error.SkipZigTest; // TODO
210213 if (builtin.zig_backend == .stage2_arm) return error.SkipZigTest; // TODO
......@@ -217,8 +220,8 @@ test "type info: error set merged" {
217220 try expect(error_set_info == .ErrorSet);
218221 try expect(error_set_info.ErrorSet.?.len == 3);
219222 try expect(mem.eql(u8, error_set_info.ErrorSet.?[0].name, "One"));
220 try expect(mem.eql(u8, error_set_info.ErrorSet.?[1].name, "Two"));
221 try expect(mem.eql(u8, error_set_info.ErrorSet.?[2].name, "Three"));
223 try expect(mem.eql(u8, error_set_info.ErrorSet.?[1].name, "Three"));
224 try expect(mem.eql(u8, error_set_info.ErrorSet.?[2].name, "Two"));
222225}
223226
224227test "type info: enum info" {
test/stage2/x86_64.zig-19
......@@ -1412,25 +1412,6 @@ pub fn addCases(ctx: *TestContext) !void {
14121412 });
14131413 }
14141414
1415 {
1416 var case = ctx.exe("error set equality", target);
1417
1418 case.addCompareOutput(
1419 \\pub fn main() void {
1420 \\ assert(@TypeOf(error.Foo) == @TypeOf(error.Foo));
1421 \\ assert(@TypeOf(error.Bar) != @TypeOf(error.Foo));
1422 \\ assert(anyerror == anyerror);
1423 \\ assert(error{Foo} != error{Foo});
1424 \\ // TODO put inferred error sets here when @typeInfo works
1425 \\}
1426 \\fn assert(b: bool) void {
1427 \\ if (!b) unreachable;
1428 \\}
1429 ,
1430 "",
1431 );
1432 }
1433
14341415 {
14351416 var case = ctx.exe("comptime var", target);
14361417