diff --git a/src/main/kotlin/graphql/kickstart/tools/ResolverInfo.kt b/src/main/kotlin/graphql/kickstart/tools/ResolverInfo.kt index 4d734af7..00042740 100644 --- a/src/main/kotlin/graphql/kickstart/tools/ResolverInfo.kt +++ b/src/main/kotlin/graphql/kickstart/tools/ResolverInfo.kt @@ -3,7 +3,10 @@ package graphql.kickstart.tools import graphql.kickstart.tools.resolver.FieldResolverScanner import graphql.kickstart.tools.util.GraphQLRootResolver import graphql.kickstart.tools.util.JavaType +import graphql.kickstart.tools.util.eraseUnboundedWildcards +import graphql.kickstart.tools.util.unwrap import org.apache.commons.lang3.reflect.TypeUtils +import java.lang.reflect.ParameterizedType internal abstract class ResolverInfo { abstract fun getFieldSearches(): List @@ -28,6 +31,11 @@ internal class NormalResolverInfo( private fun findDataClass(): Class { val type = TypeUtils.getTypeArguments(resolverType, GraphQLResolver::class.java)[GraphQLResolver::class.java.typeParameters[0]] + ?.eraseUnboundedWildcards() + + if (type is ParameterizedType) { + throw ResolverError("Resolver '${resolverType.name}' may not have a parameterized type (${type.typeName}) as its type, use the raw type or unbounded wildcards ( in Java, <*> in Kotlin) instead.") + } if (type == null || type !is Class<*>) { throw ResolverError("Unable to determine data class for resolver '${resolverType.name}' from generic interface! This is most likely a bug with graphql-java-tools.") @@ -54,14 +62,16 @@ internal class NormalResolverInfo( */ internal class MultiResolverInfo( val resolverInfoList: List, - override val dataClassType: Class + private val dataClass: JavaType ) : DataClassTypeResolverInfo, ResolverInfo() { + override val dataClassType = dataClass.unwrap() + override fun getFieldSearches(): List { return resolverInfoList .asSequence() .map { FieldResolverScanner.Search(it.resolverType, this, it.resolver, it.dataClassType) } - .plus(FieldResolverScanner.Search(dataClassType, this, null)) + .plus(FieldResolverScanner.Search(dataClass, this, null)) .toList() } } diff --git a/src/main/kotlin/graphql/kickstart/tools/SchemaClassScanner.kt b/src/main/kotlin/graphql/kickstart/tools/SchemaClassScanner.kt index 6c3fbce1..7a6398da 100644 --- a/src/main/kotlin/graphql/kickstart/tools/SchemaClassScanner.kt +++ b/src/main/kotlin/graphql/kickstart/tools/SchemaClassScanner.kt @@ -308,19 +308,20 @@ internal class SchemaClassScanner( * Find all resolvers for the data class or any of its supertypes, most specific first. */ private fun getResolverInfoFromDataClass(dataClass: JavaType): ResolverInfo { + val rawDataClass = dataClass.unwrap() val resolverInfoList = resolverInfos - .filter { it.dataClassType == dataClass || isResolverForSupertype(it, dataClass) } + .filter { it.dataClassType == rawDataClass || isResolverForSupertype(it, rawDataClass) } .sortedByDescending { ClassUtils.getAllSuperclasses(it.dataClassType).size + ClassUtils.getAllInterfaces(it.dataClassType).size } return when { resolverInfoList.isEmpty() -> DataClassResolverInfo(dataClass) resolverInfoList.size == 1 && resolverInfoList.single().dataClassType == dataClass -> resolverInfoList.single() - else -> MultiResolverInfo(resolverInfoList, dataClass.unwrap()) + else -> MultiResolverInfo(resolverInfoList, dataClass) } } - private fun isResolverForSupertype(resolverInfo: NormalResolverInfo, dataClass: JavaType) = - dataClass is Class<*> && resolverInfo.dataClassType != Object::class.java && resolverInfo.dataClassType.isAssignableFrom(dataClass) + private fun isResolverForSupertype(resolverInfo: NormalResolverInfo, dataClass: Class<*>) = + resolverInfo.dataClassType != Object::class.java && resolverInfo.dataClassType.isAssignableFrom(dataClass) private fun scanResolverInfoForPotentialMatches(type: ObjectTypeDefinition, resolverInfo: ResolverInfo) { type.getExtendedFieldDefinitions(extensionDefinitions).forEach { field -> diff --git a/src/main/kotlin/graphql/kickstart/tools/resolver/FieldResolverScanner.kt b/src/main/kotlin/graphql/kickstart/tools/resolver/FieldResolverScanner.kt index 173cc4a3..73a3c91b 100644 --- a/src/main/kotlin/graphql/kickstart/tools/resolver/FieldResolverScanner.kt +++ b/src/main/kotlin/graphql/kickstart/tools/resolver/FieldResolverScanner.kt @@ -163,7 +163,7 @@ internal class FieldResolverScanner(val options: SchemaParserOptions) { private fun verifyMethodArguments(method: Method, requiredCount: Int, search: Search): Boolean { val appropriateFirstParameter = if (search.requiredFirstParameterType != null) { method.genericParameterTypes.firstOrNull()?.let { - it == search.requiredFirstParameterType || method.declaringClass.typeParameters.contains(it) + it.eraseUnboundedWildcards() == search.requiredFirstParameterType || method.declaringClass.typeParameters.contains(it) } ?: false } else { // an extension receiver can only take the source object diff --git a/src/main/kotlin/graphql/kickstart/tools/util/Utils.kt b/src/main/kotlin/graphql/kickstart/tools/util/Utils.kt index ae8dd5c4..24a1e8f6 100644 --- a/src/main/kotlin/graphql/kickstart/tools/util/Utils.kt +++ b/src/main/kotlin/graphql/kickstart/tools/util/Utils.kt @@ -41,6 +41,16 @@ internal fun JavaType.unwrap(): Class = this as Class<*> } +/** + * Replaces a parameterized type whose type arguments are all unbounded wildcards, e.g. Kotlin's Page<*>, by its raw type. + */ +internal fun JavaType.eraseUnboundedWildcards(): JavaType = + if (this is ParameterizedType && this.actualTypeArguments.all { TypeUtils.equals(it, TypeUtils.WILDCARD_ALL) }) { + this.rawType + } else { + this + } + internal fun JavaType.typeArgument(type: Class<*>): JavaType? = TypeUtils.getTypeArguments(this, type)?.get(type.typeParameters.first()) diff --git a/src/test/kotlin/graphql/kickstart/tools/GenericResolverTest.kt b/src/test/kotlin/graphql/kickstart/tools/GenericResolverTest.kt index cc8444df..0fc34a45 100644 --- a/src/test/kotlin/graphql/kickstart/tools/GenericResolverTest.kt +++ b/src/test/kotlin/graphql/kickstart/tools/GenericResolverTest.kt @@ -1,5 +1,8 @@ package graphql.kickstart.tools +import graphql.GraphQL +import graphql.kickstart.tools.resolver.FieldResolverError +import org.junit.Assert.assertThrows import org.junit.Test class GenericResolverTest { @@ -64,4 +67,99 @@ class GenericResolverTest { class Car class CarResolver : FooGraphQLResolver() + + @Test + fun `star projected resolvers are applied to parameterized data classes`() { + val gql = GraphQL.newGraphQL(pageSchema(PageResolver())).build() + + val data = assertNoGraphQlErrors(gql) { + """ + query { + page { + content { name } + size + } + } + """ + } + + assertEquals(data["page"], mapOf("content" to listOf(mapOf("name" to "item")), "size" to 1)) + } + + @Test + fun `supertype resolvers are applied to parameterized data classes`() { + val gql = GraphQL.newGraphQL(pageSchema(CountableResolver())).build() + + val data = assertNoGraphQlErrors(gql) { + """ + query { + page { + content { name } + size + } + } + """ + } + + assertEquals(data["page"], mapOf("content" to listOf(mapOf("name" to "item")), "size" to 1)) + } + + @Test + fun `resolvers for a specific parameterization of a data class are rejected`() { + assertThrows(FieldResolverError::class.java) { pageSchema(ItemPageSourceResolver()) } + + val error = assertThrows(ResolverError::class.java) { pageSchema(ItemPageResolver()) } + assertEquals(error.message, "Resolver '${ItemPageResolver::class.java.name}' may not have a parameterized type " + + "(${Page::class.java.name}<${Item::class.java.name}>) as its type, use the raw type or unbounded wildcards ( in Java, <*> in Kotlin) instead.") + } + + private fun pageSchema(resolver: GraphQLResolver<*>) = SchemaParser.newParser() + .schemaString( + """ + type Query { + page: ItemPage! + } + + type ItemPage { + content: [Item!]! + size: Int! + } + + type Item { + name: String! + } + """) + .resolvers(QueryResolver3(), resolver) + .build() + .makeExecutableSchema() + + class QueryResolver3 : GraphQLQueryResolver { + fun getPage(): Page = Page(listOf(Item("item"))) + } + + interface Countable { + fun count(): Int + } + + class Page(val content: List) : Countable { + override fun count(): Int = content.size + } + + class Item(val name: String) + + class PageResolver : GraphQLResolver> { + fun getSize(page: Page<*>): Int = page.content.size + } + + class CountableResolver : GraphQLResolver { + fun getSize(countable: Countable): Int = countable.count() + } + + class ItemPageSourceResolver : GraphQLResolver> { + fun getSize(page: Page): Int = page.content.size + } + + class ItemPageResolver : GraphQLResolver> { + fun getSize(page: Page): Int = page.content.size + } } diff --git a/src/test/kotlin/graphql/kickstart/tools/ResolverMethodsTest.java b/src/test/kotlin/graphql/kickstart/tools/ResolverMethodsTest.java index 28a3c293..864e8c69 100644 --- a/src/test/kotlin/graphql/kickstart/tools/ResolverMethodsTest.java +++ b/src/test/kotlin/graphql/kickstart/tools/ResolverMethodsTest.java @@ -6,6 +6,7 @@ import graphql.schema.GraphQLSchema; import org.junit.Test; +import java.util.List; import java.util.Map; import static org.junit.Assert.assertEquals; @@ -76,9 +77,65 @@ public String name(Product product) { assertEquals(Map.of("name", "product"), data.get("product")); } + // Raw types can't be expressed in Kotlin, so this resolver must stay in Java. + @Test + public void testRawResolverForParameterizedDataClass() { + GraphQLSchema schema = SchemaParser.newParser() + .schemaString("type Query { page: ItemPage! } type ItemPage { content: [Item!]! size: Int! } type Item { name: String! }") + .resolvers(new PageQueryResolver(), new RawPageResolver()) + .build() + .makeExecutableSchema(); + + GraphQL gql = GraphQL.newGraphQL(schema).build(); + + ExecutionResult result = gql + .execute(ExecutionInput.newExecutionInput() + .query("query { page { content { name } size } }") + .root(new Object())); + + assertTrue(result.getErrors().isEmpty()); + Map data = result.getData(); + assertEquals(Map.of("content", List.of(Map.of("name", "item")), "size", 1), data.get("page")); + } + static class Product { } + static class Page { + private final List content; + + Page(List content) { + this.content = content; + } + + public List getContent() { + return content; + } + } + + static class Item { + public String getName() { + return "item"; + } + } + + static class PageQueryResolver implements GraphQLQueryResolver { + + @SuppressWarnings("unused") + public Page page() { + return new Page<>(List.of(new Item())); + } + } + + @SuppressWarnings("rawtypes") + static class RawPageResolver implements GraphQLResolver { + + @SuppressWarnings("unused") + public int size(Page page) { + return page.getContent().size(); + } + } + static class Resolver implements GraphQLQueryResolver { @SuppressWarnings("unused")