diff --git a/assert/assertions.go b/assert/assertions.go index 166f63726..0e28d8b03 100644 --- a/assert/assertions.go +++ b/assert/assertions.go @@ -83,6 +83,16 @@ func ObjectsAreEqual(expected, actual interface{}) bool { // copyExportedFields iterates downward through nested data structures and creates a copy // that only contains the exported struct fields. func copyExportedFields(expected interface{}) interface{} { + return copyExportedFieldsWithVisited(expected, make(map[copyExportedFieldsVisit]reflect.Value)) +} + +type copyExportedFieldsVisit struct { + typ reflect.Type + ptr uintptr + length int +} + +func copyExportedFieldsWithVisited(expected interface{}, visited map[copyExportedFieldsVisit]reflect.Value) interface{} { if isNil(expected) { return expected } @@ -91,6 +101,20 @@ func copyExportedFields(expected interface{}) interface{} { expectedKind := expectedType.Kind() expectedValue := reflect.ValueOf(expected) + // Pointers, slices and maps can form cycles. Reuse the copy already made + // for a value seen earlier so that recursive structures terminate. + var visit copyExportedFieldsVisit + switch expectedKind { + case reflect.Ptr, reflect.Slice, reflect.Map: + visit = copyExportedFieldsVisit{typ: expectedType, ptr: expectedValue.Pointer()} + if expectedKind == reflect.Slice { + visit.length = expectedValue.Len() + } + if result, ok := visited[visit]; ok { + return result.Interface() + } + } + switch expectedKind { case reflect.Struct: result := reflect.New(expectedType).Elem() @@ -102,7 +126,7 @@ func copyExportedFields(expected interface{}) interface{} { if isNil(fieldValue) || isNil(fieldValue.Interface()) { continue } - newValue := copyExportedFields(fieldValue.Interface()) + newValue := copyExportedFieldsWithVisited(fieldValue.Interface(), visited) result.Field(i).Set(reflect.ValueOf(newValue)) } } @@ -110,7 +134,8 @@ func copyExportedFields(expected interface{}) interface{} { case reflect.Ptr: result := reflect.New(expectedType.Elem()) - unexportedRemoved := copyExportedFields(expectedValue.Elem().Interface()) + visited[visit] = result + unexportedRemoved := copyExportedFieldsWithVisited(expectedValue.Elem().Interface(), visited) result.Elem().Set(reflect.ValueOf(unexportedRemoved)) return result.Interface() @@ -120,22 +145,24 @@ func copyExportedFields(expected interface{}) interface{} { result = reflect.New(reflect.ArrayOf(expectedValue.Len(), expectedType.Elem())).Elem() } else { result = reflect.MakeSlice(expectedType, expectedValue.Len(), expectedValue.Len()) + visited[visit] = result } for i := 0; i < expectedValue.Len(); i++ { index := expectedValue.Index(i) if isNil(index) { continue } - unexportedRemoved := copyExportedFields(index.Interface()) + unexportedRemoved := copyExportedFieldsWithVisited(index.Interface(), visited) result.Index(i).Set(reflect.ValueOf(unexportedRemoved)) } return result.Interface() case reflect.Map: result := reflect.MakeMap(expectedType) + visited[visit] = result for _, k := range expectedValue.MapKeys() { index := expectedValue.MapIndex(k) - unexportedRemoved := copyExportedFields(index.Interface()) + unexportedRemoved := copyExportedFieldsWithVisited(index.Interface(), visited) result.SetMapIndex(k, reflect.ValueOf(unexportedRemoved)) } return result.Interface() diff --git a/assert/assertions_test.go b/assert/assertions_test.go index 11642e096..db03499b7 100644 --- a/assert/assertions_test.go +++ b/assert/assertions_test.go @@ -372,6 +372,195 @@ func TestCopyExportedFields(t *testing.T) { } } +func TestEqualExportedValuesRecursiveStruct(t *testing.T) { + t.Parallel() + + type Node struct { + Value int + Self *Node + Left *Node + Right *Node + unexported string + } + + selfCycle := func(value int, unexported string) *Node { + n := &Node{Value: value, unexported: unexported} + n.Self = n + return n + } + // ring builds n1 -> n2 -> ... -> nN -> n1 through the Self field. + ring := func(values ...int) *Node { + nodes := make([]*Node, len(values)) + for i, v := range values { + nodes[i] = &Node{Value: v} + } + for i := range nodes { + nodes[i].Self = nodes[(i+1)%len(nodes)] + } + return nodes[0] + } + chain := func(values ...int) *Node { + var head *Node + for i := len(values) - 1; i >= 0; i-- { + head = &Node{Value: values[i], Self: head} + } + return head + } + shared := &Node{Value: 3} + + type SliceCycle struct { + Items []interface{} + unexported string + } + sliceCycle := func(unexported string) SliceCycle { + items := make([]interface{}, 1) + items[0] = items + return SliceCycle{Items: items, unexported: unexported} + } + + type MapCycle struct { + Children map[string]*MapCycle + unexported string + } + mapCycle := func(unexported string) *MapCycle { + m := &MapCycle{Children: map[string]*MapCycle{}, unexported: unexported} + m.Children["self"] = m + return m + } + + cases := []struct { + name string + value1 interface{} + value2 interface{} + expectedEqual bool + }{ + { + name: "direct self-cycle", + value1: selfCycle(1, "a"), + value2: selfCycle(1, "b"), + expectedEqual: true, + }, + { + name: "two-node cycle", + value1: ring(1, 2), + value2: ring(1, 2), + expectedEqual: true, + }, + { + name: "three-node cycle", + value1: ring(1, 2, 3), + value2: ring(1, 2, 3), + expectedEqual: true, + }, + { + name: "acyclic chain", + value1: chain(1, 2, 3), + value2: chain(1, 2, 3), + expectedEqual: true, + }, + { + name: "shared reference without a cycle", + value1: &Node{Value: 1, Left: shared, Right: shared}, + value2: &Node{Value: 1, Left: &Node{Value: 3}, Right: &Node{Value: 3}}, + expectedEqual: true, + }, + { + name: "struct containing recursive pointers", + value1: S3{Exported1: &Nested{Exported: selfCycle(1, "a")}}, + value2: S3{Exported1: &Nested{Exported: selfCycle(1, "b")}}, + expectedEqual: true, + }, + { + name: "array of cyclic pointers", + value1: [2]*Node{selfCycle(1, "a"), ring(1, 2)}, + value2: [2]*Node{selfCycle(1, "b"), ring(1, 2)}, + expectedEqual: true, + }, + { + name: "slice cycle", + value1: sliceCycle("a"), + value2: sliceCycle("b"), + expectedEqual: true, + }, + { + name: "map cycle", + value1: mapCycle("a"), + value2: mapCycle("b"), + expectedEqual: true, + }, + { + name: "self-cycles with different exported fields", + value1: selfCycle(1, "a"), + value2: selfCycle(2, "a"), + expectedEqual: false, + }, + { + name: "three-node cycles with different exported fields", + value1: ring(1, 2, 3), + value2: ring(1, 2, 4), + expectedEqual: false, + }, + { + name: "acyclic chains with different exported fields", + value1: chain(1, 2, 3), + value2: chain(1, 2, 4), + expectedEqual: false, + }, + { + name: "self-cycle and nil pointer", + value1: selfCycle(1, "a"), + value2: &Node{Value: 1}, + expectedEqual: false, + }, + { + name: "cycle and acyclic chain", + value1: ring(1, 2), + value2: chain(1, 2), + expectedEqual: false, + }, + } + + for _, c := range cases { + t.Run(c.name, func(t *testing.T) { + mockT := new(mockTestingT) + + actual := EqualExportedValues(mockT, c.value1, c.value2) + if actual != c.expectedEqual { + t.Errorf("Expected EqualExportedValues to be %t, but was %t", c.expectedEqual, actual) + } + if !c.expectedEqual && !strings.Contains(mockT.errorString(), "Not equal (comparing only exported fields)") { + t.Errorf("Expected a failure message, got %q", mockT.errorString()) + } + }) + } +} + +func TestCopyExportedFieldsRecursiveStruct(t *testing.T) { + t.Parallel() + + type Node struct { + Self *Node + unexported string + } + + input := &Node{unexported: "a"} + input.Self = input + + output, ok := copyExportedFields(input).(*Node) + if !ok { + t.Fatalf("Expected a *Node, got %T", copyExportedFields(input)) + } + if output == input { + t.Error("Expected a copy, got the input pointer") + } + if output.Self != output { + t.Error("Expected the copy to point back to itself") + } + if output.unexported != "" { + t.Errorf("Expected unexported field to be cleared, got %q", output.unexported) + } +} + func TestEqualExportedValues(t *testing.T) { t.Parallel()