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
Original file line number Diff line number Diff line change
Expand Up @@ -523,11 +523,10 @@ class ClientApiGenerator(
}
}

val concreteTypesResult = createConcreteTypes(type, javaType.build(), javaType, mutableSetOf(), 0)
val unionTypesResult = createUnionTypes(type, javaType, javaType.build(), mutableSetOf(), 0)
val fragmentTypesResult = addFragmentProjectionMethods(type, javaType.build(), javaType, mutableSetOf(), 0)

val javaFile = JavaFile.builder(getPackageName(), javaType.build()).build()
return CodeGenResult(clientProjections = listOf(javaFile)).merge(codeGenResult).merge(concreteTypesResult).merge(unionTypesResult)
return CodeGenResult(clientProjections = listOf(javaFile)).merge(codeGenResult).merge(fragmentTypesResult)
}

private fun addFieldSelectionMethodWithArguments(
Expand Down Expand Up @@ -674,44 +673,17 @@ class ClientApiGenerator(
return CodeGenResult(clientProjections = listOf(javaFile)).merge(codeGenResult)
}

private fun createConcreteTypes(
private fun addFragmentProjectionMethods(
type: TypeDefinition<*>,
root: TypeSpec,
javaType: TypeSpec.Builder,
processedEdges: Set<Pair<String, String>>,
queryDepth: Int,
): CodeGenResult =
if (type is InterfaceTypeDefinition) {
val concreteTypes =
document
.getDefinitionsOfType(ObjectTypeDefinition::class.java)
.filter {
it.implements.filterIsInstance<NamedNode<*>>().any { iface -> iface.name == type.name }
}.distinctBy { it.name }
concreteTypes
.map {
addFragmentProjectionMethod(javaType, root, it, processedEdges, queryDepth)
}.fold(CodeGenResult.EMPTY) { total, current -> total.merge(current) }
} else {
CodeGenResult.EMPTY
}

private fun createUnionTypes(
type: TypeDefinition<*>,
javaType: TypeSpec.Builder,
rootType: TypeSpec,
processedEdges: Set<Pair<String, String>>,
queryDepth: Int,
): CodeGenResult =
if (type is UnionTypeDefinition) {
val memberTypes = type.memberTypes.mapNotNull { it.findTypeDefinition(document, true) }.toList()
memberTypes
.map {
addFragmentProjectionMethod(javaType, rootType, it, processedEdges, queryDepth)
}.fold(CodeGenResult.EMPTY) { total, current -> total.merge(current) }
} else {
CodeGenResult.EMPTY
}
getFragmentTypes(type)
.map {
addFragmentProjectionMethod(javaType, root, it, processedEdges, queryDepth)
}.fold(CodeGenResult.EMPTY) { total, current -> total.merge(current) }

private fun addFragmentProjectionMethod(
javaType: TypeSpec.Builder,
Expand Down Expand Up @@ -740,18 +712,18 @@ class ClientApiGenerator(
).build(),
)

return createFragment(it as ObjectTypeDefinition, rootType, projectionName, processedEdges, queryDepth)
return createFragment(it, rootType, projectionName, processedEdges, queryDepth)
}

private fun createFragment(
type: ObjectTypeDefinition,
type: TypeDefinition<*>,
root: TypeSpec,
prefix: String,
processedEdges: Set<Pair<String, String>>,
queryDepth: Int,
): CodeGenResult {
val subProjection =
createSubProjectionType(type, root, prefix, processedEdges, queryDepth)
createSubProjectionType(type, root, prefix, processedEdges, queryDepth, fragment = true)
?: return CodeGenResult.EMPTY
val javaType = subProjection.first
val codeGenResult = subProjection.second
Expand Down Expand Up @@ -794,6 +766,27 @@ class ClientApiGenerator(
return CodeGenResult(clientProjections = listOf(javaFile)).merge(codeGenResult)
}

private fun getFragmentTypes(type: TypeDefinition<*>): List<TypeDefinition<*>> =
when (type) {
is InterfaceTypeDefinition -> {
document
.getDefinitionsOfType(InterfaceTypeDefinition::class.java)

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

This is the heart of the change -- we now include other interface types in addition to concrete object types.

.plus(document.getDefinitionsOfType(ObjectTypeDefinition::class.java))
.filter {
it.implements.filterIsInstance<NamedNode<*>>().any { iface -> iface.name == type.name }
}.distinctBy { it.name }
}

is UnionTypeDefinition -> {
type.memberTypes
.mapNotNull { it.findTypeDefinition(document, true) }
}

else -> {
emptyList()
}
}

private fun createSubProjection(
type: TypeDefinition<*>,
root: TypeSpec,
Expand All @@ -817,6 +810,7 @@ class ClientApiGenerator(
prefix: String,
processedEdges: Set<Pair<String, String>>,
queryDepth: Int,
fragment: Boolean = false,

Copy link
Copy Markdown
Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

We need to know if the projection is for a fragment so that we don't generate fragments within fragments.

): Pair<TypeSpec.Builder, CodeGenResult>? {
val className = ClassName.get(BaseSubProjectionNode::class.java)
val clazzName = "${prefix}Projection"
Expand Down Expand Up @@ -959,10 +953,14 @@ class ClientApiGenerator(
}
}

val concreteTypesResult = createConcreteTypes(type, root, javaType, processedEdges, queryDepth)
val unionTypesResult = createUnionTypes(type, javaType, root, processedEdges, queryDepth)
val fragmentTypesResult =
if (!fragment) {
addFragmentProjectionMethods(type, root, javaType, processedEdges, queryDepth)
} else {
CodeGenResult.EMPTY
}

return javaType to codeGenResult.merge(concreteTypesResult).merge(unionTypesResult)
return javaType to codeGenResult.merge(fragmentTypesResult)
}

private fun getDeprecateDirective(node: DirectivesContainer<*>): Directive? {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -145,6 +145,66 @@ class ClientApiGenFragmentTest {
)
}

@Test
fun interfaceFragmentWithInterfaceHierarchy() {
val schema =
"""
type Query {
searchForSeries(title: String): [Series]
}

interface Show {
title: String
relatedShows: [Show]
}

interface Series implements Show {
title: String
relatedShows: [Show]
}

interface MiniSeries implements Series & Show {
title: String
relatedShows: [Show]
episodes: Int
}
""".trimIndent()

val codeGenResult =
CodeGen(
CodeGenConfig(
schemas = setOf(schema),
packageName = BASE_PACKAGE_NAME,
generateClientApi = true,
),
).generate()

assertThat(codeGenResult.clientProjections.size).isEqualTo(4)
assertThat(codeGenResult.clientProjections[0].typeSpec().name()).isEqualTo("SearchForSeriesProjectionRoot")
assertThat(codeGenResult.clientProjections[0].typeSpec().methodSpecs()).extracting("name").contains("title")
assertThat(codeGenResult.clientProjections[0].typeSpec().methodSpecs()).extracting("name").contains("__typename")
assertThat(codeGenResult.clientProjections[0].typeSpec().methodSpecs()).extracting("name").contains("onMiniSeries")
assertThat(codeGenResult.clientProjections[0].typeSpec().methodSpecs()).extracting("name").doesNotContain("onSeries")
assertThat(codeGenResult.clientProjections[1].typeSpec().name()).isEqualTo("ShowProjection")
assertThat(codeGenResult.clientProjections[1].typeSpec().methodSpecs()).extracting("name").contains("title")
assertThat(codeGenResult.clientProjections[1].typeSpec().methodSpecs()).extracting("name").contains("relatedShows")
assertThat(codeGenResult.clientProjections[1].typeSpec().methodSpecs()).extracting("name").contains("onSeries")
assertThat(codeGenResult.clientProjections[1].typeSpec().methodSpecs()).extracting("name").contains("onMiniSeries")
assertThat(codeGenResult.clientProjections[1].typeSpec().methodSpecs()).extracting("name").doesNotContain("onShow")
assertThat(codeGenResult.clientProjections[2].typeSpec().name()).isEqualTo("SeriesFragmentProjection")
assertThat(codeGenResult.clientProjections[2].typeSpec().methodSpecs()).extracting("name").contains("title")
assertThat(codeGenResult.clientProjections[2].typeSpec().methodSpecs()).extracting("name").contains("relatedShows")
assertThat(codeGenResult.clientProjections[2].typeSpec().methodSpecs()).extracting("name").doesNotContain("onMiniSeries")
assertThat(codeGenResult.clientProjections[3].typeSpec().name()).isEqualTo("MiniSeriesFragmentProjection")
assertThat(codeGenResult.clientProjections[3].typeSpec().methodSpecs()).extracting("name").contains("title")
assertThat(codeGenResult.clientProjections[3].typeSpec().methodSpecs()).extracting("name").contains("relatedShows")

assertCompilesJava(
codeGenResult.clientProjections + codeGenResult.javaQueryTypes + codeGenResult.javaEnumTypes + codeGenResult.javaDataTypes +
codeGenResult.javaInterfaces,
)
}

@Test
fun unionFragment() {
val schema =
Expand Down
Loading