Skip to content

Commit 58f5750

Browse files
authored
Schema definition transformer (#718)
* Add faster schema definition transformer * Refactor slightly * Add comment
1 parent d883a01 commit 58f5750

5 files changed

Lines changed: 326 additions & 55 deletions

File tree

lib/src/main/java/graphql/nadel/NadelSchemas.kt

Lines changed: 21 additions & 11 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
package graphql.nadel
22

3+
import graphql.nadel.schema.NadelSchemaDefinitionTransformationHook
34
import graphql.nadel.schema.NeverWiringFactory
45
import graphql.nadel.schema.OverallSchemaGenerator
56
import graphql.nadel.schema.SchemaTransformationHook
@@ -23,11 +24,12 @@ data class NadelSchemas(
2324

2425
class Builder {
2526
internal var schemaTransformationHook: SchemaTransformationHook = SchemaTransformationHook.Identity
27+
internal var schemaDefinitionTransformationHook: NadelSchemaDefinitionTransformationHook =
28+
NadelSchemaDefinitionTransformationHook.Identity
2629

2730
internal var overallWiringFactory: WiringFactory = NeverWiringFactory()
2831
internal var underlyingWiringFactory: WiringFactory = NeverWiringFactory()
2932

30-
3133
internal var serviceExecutionFactory: ServiceExecutionFactory? = null
3234

3335
// .nadel files
@@ -47,6 +49,10 @@ data class NadelSchemas(
4749
schemaTransformationHook = value
4850
}
4951

52+
fun schemaDefinitionTransformationHook(value: NadelSchemaDefinitionTransformationHook): Builder = also {
53+
schemaDefinitionTransformationHook = value
54+
}
55+
5056
fun overallWiringFactory(value: WiringFactory): Builder = also {
5157
overallWiringFactory = value
5258
}
@@ -162,10 +168,12 @@ data class NadelSchemas(
162168
// Combine readers & type defs
163169
val readersToTypeDefs = underlyingSchemaReaders
164170
.mapValues { (_, reader) ->
165-
SchemaUtil.parseTypeDefinitionRegistry(
166-
reader,
167-
captureSourceLocation = captureSourceLocation,
168-
)
171+
reader.use {
172+
SchemaUtil.parseTypeDefinitionRegistry(
173+
reader,
174+
captureSourceLocation = captureSourceLocation,
175+
)
176+
}
169177
}
170178
val resolvedUnderlyingTypeDefs = readersToTypeDefs + underlyingTypeDefs
171179

@@ -203,10 +211,12 @@ data class NadelSchemas(
203211
val underlyingSchemaGenerator = UnderlyingSchemaGenerator()
204212

205213
return builder.overallSchemaReaders.map { (serviceName, reader) ->
206-
val schemaDefinitions = SchemaUtil.parseSchemaDefinitions(
207-
reader,
208-
captureSourceLocation = captureSourceLocation,
209-
)
214+
val schemaDefinitions = reader.use {
215+
SchemaUtil.parseSchemaDefinitions(
216+
reader,
217+
captureSourceLocation = captureSourceLocation,
218+
)
219+
}
210220
val typeDefinitionRegistry = NadelTypeDefinitionRegistry.from(schemaDefinitions)
211221

212222
// Builder should enforce non-null entry
@@ -227,7 +237,8 @@ data class NadelSchemas(
227237
val serviceRegistries = services.map(Service::definitionRegistry)
228238
val schema = overallSchemaGenerator.buildOverallSchema(
229239
serviceRegistries,
230-
builder.overallWiringFactory
240+
builder.overallWiringFactory,
241+
builder.schemaDefinitionTransformationHook,
231242
)
232243
val newSchema = builder.schemaTransformationHook.apply(schema, services)
233244

@@ -240,4 +251,3 @@ data class NadelSchemas(
240251
}
241252
}
242253
}
243-
Lines changed: 240 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,240 @@
1+
package graphql.nadel.schema
2+
3+
import graphql.language.FieldDefinition
4+
import graphql.language.ImplementingTypeDefinition
5+
import graphql.language.InterfaceTypeDefinition
6+
import graphql.language.InterfaceTypeExtensionDefinition
7+
import graphql.language.ObjectTypeDefinition
8+
import graphql.language.ObjectTypeExtensionDefinition
9+
import graphql.nadel.NadelOperationKind
10+
import graphql.nadel.engine.util.unwrapAll
11+
import graphql.nadel.util.AnyImplementingTypeDefinition
12+
import graphql.nadel.util.AnyNamedNode
13+
import graphql.nadel.util.AnySDLDefinition
14+
import graphql.schema.idl.DirectiveInfo
15+
import graphql.schema.idl.ScalarInfo
16+
17+
fun interface NadelFieldDefinitionVisibilityTransformationPredicate {
18+
/**
19+
* @return `true` to keep field and `false` to delete
20+
*/
21+
operator fun invoke(
22+
parent: ImplementingTypeDefinition<*>,
23+
field: FieldDefinition,
24+
): Boolean
25+
}
26+
27+
class NadelFieldDefinitionVisibilityTransformation(
28+
val fieldPredicate: NadelFieldDefinitionVisibilityTransformationPredicate,
29+
) : NadelSchemaDefinitionTransformationHook {
30+
override fun invoke(definitions: List<AnySDLDefinition>): List<AnySDLDefinition> {
31+
return deleteFields(definitions)
32+
}
33+
34+
private fun deleteFields(definitions: List<AnySDLDefinition>): List<AnySDLDefinition> {
35+
val newDefinitions = definitions
36+
.map { definition ->
37+
when (definition) {
38+
is InterfaceTypeExtensionDefinition -> transformExtensionType(definition)
39+
is InterfaceTypeDefinition -> transformType(definition)
40+
is ObjectTypeExtensionDefinition -> transformExtensionType(definition)
41+
is ObjectTypeDefinition -> transformType(definition)
42+
else -> definition
43+
}
44+
}
45+
46+
val observedBeforeTransform = getStronglyReferencedTypes(definitions)
47+
val observedAfterTransform = getStronglyReferencedTypes(newDefinitions)
48+
49+
return newDefinitions
50+
.filter { definition ->
51+
if (definition is AnyNamedNode) {
52+
if (definition.name in observedBeforeTransform) {
53+
definition.name in observedAfterTransform
54+
} else {
55+
true
56+
}
57+
} else {
58+
true
59+
}
60+
}
61+
}
62+
63+
private fun transformExtensionType(definition: InterfaceTypeExtensionDefinition): InterfaceTypeExtensionDefinition {
64+
val newFields = filterFields(definition)
65+
if (newFields.size == definition.fieldDefinitions.size) {
66+
return definition
67+
}
68+
69+
return definition.transformExtension { builder ->
70+
builder.definitions(newFields)
71+
}
72+
}
73+
74+
private fun transformType(definition: InterfaceTypeDefinition): InterfaceTypeDefinition {
75+
val newFields = filterFields(definition)
76+
if (newFields.size == definition.fieldDefinitions.size) {
77+
return definition
78+
}
79+
80+
return definition.transform { builder ->
81+
builder.definitions(newFields)
82+
}
83+
}
84+
85+
private fun transformExtensionType(definition: ObjectTypeExtensionDefinition): ObjectTypeExtensionDefinition {
86+
val newFields = filterFields(definition)
87+
if (newFields.size == definition.fieldDefinitions.size) {
88+
return definition
89+
}
90+
91+
return definition.transformExtension { builder ->
92+
builder.fieldDefinitions(newFields)
93+
}
94+
}
95+
96+
private fun transformType(definition: ObjectTypeDefinition): ObjectTypeDefinition {
97+
val newFields = filterFields(definition)
98+
if (newFields.size == definition.fieldDefinitions.size) {
99+
return definition
100+
}
101+
102+
return definition.transform { builder ->
103+
builder.fieldDefinitions(newFields)
104+
}
105+
}
106+
107+
private fun filterFields(
108+
parent: ImplementingTypeDefinition<*>,
109+
): List<FieldDefinition> {
110+
return parent.fieldDefinitions
111+
.filter { field ->
112+
fieldPredicate(parent, field)
113+
}
114+
}
115+
116+
private fun getStronglyReferencedTypes(types: List<AnySDLDefinition>): Set<String> {
117+
val typesByName = types
118+
.groupBy {
119+
(it as AnyNamedNode).name
120+
}
121+
122+
val typeQueue = NadelOperationKind.entries
123+
.asSequence()
124+
.filter { it.name in typesByName }
125+
.mapTo(mutableListOf()) { it.name }
126+
127+
val typeReferences = mutableSetOf<String>()
128+
129+
while (typeQueue.isNotEmpty()) {
130+
// Exhaust queue
131+
while (typeQueue.isNotEmpty()) {
132+
val typeName = typeQueue.removeLast()
133+
if (ScalarInfo.isGraphqlSpecifiedScalar(typeName) || DirectiveInfo.isGraphqlSpecifiedDirective(typeName)) {
134+
continue
135+
}
136+
137+
val types = typesByName[typeName] ?: throw NullPointerException(typeName)
138+
collectTypeReferences(types) { typeReference ->
139+
if (typeReference !in typeReferences) {
140+
typeReferences.add(typeReference)
141+
typeQueue.add(typeReference)
142+
}
143+
}
144+
}
145+
146+
// Populate queue up with interface implementations
147+
types.forEach { type ->
148+
if (type is AnyImplementingTypeDefinition) {
149+
if (type.name !in typeReferences) {
150+
if (type.implements.any { it.unwrapAll().name in typeReferences }) {
151+
typeReferences.add(type.name)
152+
typeQueue.add(type.name)
153+
}
154+
}
155+
}
156+
}
157+
}
158+
159+
return typeReferences
160+
}
161+
162+
private fun collectTypeReferences(
163+
type: List<AnySDLDefinition>,
164+
onTypeReferenced: (String) -> Unit,
165+
) {
166+
NadelSchemaDefinitionTraverser()
167+
.traverse(
168+
roots = type
169+
.asSequence()
170+
.mapNotNull(NadelSchemaDefinitionTraverserElement::from)
171+
.asIterable(),
172+
object : NadelSchemaDefinitionTraverserVisitor {
173+
override fun visitGraphQLArgument(element: NadelSchemaDefinitionTraverserElement.Argument): Boolean {
174+
return true
175+
}
176+
177+
override fun visitGraphQLUnionType(element: NadelSchemaDefinitionTraverserElement.UnionType): Boolean {
178+
onTypeReferenced(element.node.name)
179+
return true
180+
}
181+
182+
override fun visitGraphQLInterfaceType(element: NadelSchemaDefinitionTraverserElement.InterfaceType): Boolean {
183+
onTypeReferenced(element.node.name)
184+
return true
185+
}
186+
187+
override fun visitGraphQLEnumType(element: NadelSchemaDefinitionTraverserElement.EnumType): Boolean {
188+
onTypeReferenced(element.node.name)
189+
return true
190+
}
191+
192+
override fun visitGraphQLEnumValueDefinition(element: NadelSchemaDefinitionTraverserElement.EnumValueDefinition): Boolean {
193+
return true
194+
}
195+
196+
override fun visitGraphQLFieldDefinition(element: NadelSchemaDefinitionTraverserElement.FieldDefinition): Boolean {
197+
return true
198+
}
199+
200+
override fun visitGraphQLInputObjectField(element: NadelSchemaDefinitionTraverserElement.InputObjectField): Boolean {
201+
return true
202+
}
203+
204+
override fun visitGraphQLInputObjectType(element: NadelSchemaDefinitionTraverserElement.InputObjectType): Boolean {
205+
onTypeReferenced(element.node.name)
206+
return true
207+
}
208+
209+
override fun visitGraphQLObjectType(element: NadelSchemaDefinitionTraverserElement.ObjectType): Boolean {
210+
onTypeReferenced(element.node.name)
211+
return true
212+
}
213+
214+
override fun visitGraphQLScalarType(element: NadelSchemaDefinitionTraverserElement.ScalarType): Boolean {
215+
onTypeReferenced(element.node.name)
216+
return true
217+
}
218+
219+
override fun visitGraphQLDirective(element: NadelSchemaDefinitionTraverserElement.Directive): Boolean {
220+
onTypeReferenced(element.node.name)
221+
return true
222+
}
223+
224+
override fun visitGraphQLAppliedDirective(element: NadelSchemaDefinitionTraverserElement.AppliedDirective): Boolean {
225+
onTypeReferenced(element.node.name)
226+
return true
227+
}
228+
229+
override fun visitGraphQLAppliedDirectiveArgument(element: NadelSchemaDefinitionTraverserElement.AppliedDirectiveArgument): Boolean {
230+
return true
231+
}
232+
233+
override fun visitTypeReference(element: NadelSchemaDefinitionTraverserElement.TypeReference): Boolean {
234+
onTypeReferenced(element.node.unwrapAll().name)
235+
return true
236+
}
237+
}
238+
)
239+
}
240+
}
Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,11 @@
1+
package graphql.nadel.schema
2+
3+
import graphql.nadel.util.AnySDLDefinition
4+
5+
fun interface NadelSchemaDefinitionTransformationHook {
6+
operator fun invoke(definitions: List<AnySDLDefinition>): List<AnySDLDefinition>
7+
8+
companion object {
9+
val Identity = NadelSchemaDefinitionTransformationHook { originalSchema -> originalSchema }
10+
}
11+
}

0 commit comments

Comments
 (0)