diff --git a/lib/src/main/java/graphql/nadel/definition/coordinates/NadelSchemaMemberCoordinatesFactory.kt b/lib/src/main/java/graphql/nadel/definition/coordinates/NadelSchemaMemberCoordinatesFactory.kt index afcb1b292..fe34aac48 100644 --- a/lib/src/main/java/graphql/nadel/definition/coordinates/NadelSchemaMemberCoordinatesFactory.kt +++ b/lib/src/main/java/graphql/nadel/definition/coordinates/NadelSchemaMemberCoordinatesFactory.kt @@ -1,9 +1,10 @@ package graphql.nadel.definition.coordinates import graphql.Directives +import graphql.language.Definition import graphql.language.Document -import graphql.language.NamedNode import graphql.nadel.engine.util.AnySDLDefinition +import graphql.nadel.engine.util.AnySDLNamedDefinition import graphql.nadel.engine.util.unwrapAll import graphql.nadel.schema.NadelSchemaDefinitionTraverser import graphql.nadel.schema.NadelSchemaDefinitionTraverserElement @@ -62,22 +63,33 @@ class NadelSchemaMemberCoordinatesFactory { fun create( schema: Document, + resolveTypeReferences: Boolean, ): Set { - val definitions = schema.definitions + return create( + schema = schema.definitions, + resolveTypeReferences = resolveTypeReferences, + ) + } + + fun create( + schema: Iterable>, + resolveTypeReferences: Boolean, + ): Set { + val definitions = schema .asSequence() - .filterIsInstance() + .filterIsInstance() // There can be multiple definitions per name, but in this scenario we don't care val definitionByName = definitions .associateBy { - (it as NamedNode<*>).name + it.name } val roots = definitions .mapNotNull(NadelSchemaDefinitionTraverserElement::from) .toList() - return createImpl(roots, definitionByName) + return createImpl(roots, definitionByName, resolveTypeReferences) } private fun createImpl( @@ -97,13 +109,18 @@ class NadelSchemaMemberCoordinatesFactory { private fun createImpl( roots: List, definitionByName: Map, + resolveTypeReferences: Boolean, ): Set { val coordinates = mutableSetOf() NadelSchemaDefinitionTraverser() .traverse( roots, - NadelSchemaDefinitionCoordinateCollectorTraverserVisitor(coordinates, definitionByName), + NadelSchemaDefinitionCoordinateCollectorTraverserVisitor( + coordinates, + definitionByName, + resolveTypeReferences, + ), ) return coordinates @@ -225,6 +242,7 @@ internal class NadelSchemaCoordinateCollectorTraverserVisitor( internal class NadelSchemaDefinitionCoordinateCollectorTraverserVisitor( private val coordinates: MutableCollection, private val definitionByName: Map, + private val resolveTypeReferences: Boolean, ) : NadelSchemaDefinitionTraverserVisitor { override fun visitGraphQLAppliedDirective(element: NadelSchemaDefinitionTraverserElement.AppliedDirective): Boolean { coordinates.add(element.coordinates()) @@ -302,6 +320,14 @@ internal class NadelSchemaDefinitionCoordinateCollectorTraverserVisitor( } override fun visitTypeReference(element: NadelSchemaDefinitionTraverserElement.TypeReference): Boolean { + return if (resolveTypeReferences) { + resolveTypeReference(element) + } else { + false + } + } + + private fun resolveTypeReference(element: NadelSchemaDefinitionTraverserElement.TypeReference): Boolean { // Resolve definition then traverse val typeName = element.node.unwrapAll().name diff --git a/lib/src/test/kotlin/graphql/nadel/definition/coordinates/NadelSchemaMemberCoordinatesFactoryTest.kt b/lib/src/test/kotlin/graphql/nadel/definition/coordinates/NadelSchemaMemberCoordinatesFactoryTest.kt index afea16a97..18d11e4af 100644 --- a/lib/src/test/kotlin/graphql/nadel/definition/coordinates/NadelSchemaMemberCoordinatesFactoryTest.kt +++ b/lib/src/test/kotlin/graphql/nadel/definition/coordinates/NadelSchemaMemberCoordinatesFactoryTest.kt @@ -30,7 +30,10 @@ abstract class NadelSchemaMemberCoordinatesFactoryTest { class DocumentDefinitionExtractorTest : NadelSchemaMemberCoordinatesFactoryTest() { override fun extractCoordinates(schema: String): Set { - return NadelSchemaMemberCoordinatesFactory().create(Parser().parseDocument(schema)) + return NadelSchemaMemberCoordinatesFactory().create( + Parser().parseDocument(schema), + resolveTypeReferences = true, + ) } @Test