Skip to content

Commit e6f79b4

Browse files
committed
Add faster schema definition transformer
1 parent 052c804 commit e6f79b4

5 files changed

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

0 commit comments

Comments
 (0)