Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
35 changes: 31 additions & 4 deletions assert/assertions.go
Original file line number Diff line number Diff line change
Expand Up @@ -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
}
Expand All @@ -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()
Expand All @@ -102,15 +126,16 @@ 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))
}
}
return result.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()

Expand All @@ -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()
Expand Down
189 changes: 189 additions & 0 deletions assert/assertions_test.go
Original file line number Diff line number Diff line change
Expand Up @@ -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()

Expand Down
Loading