diff --git a/lib/src/main/java/graphql/nadel/engine/util/NadelPseudoSealedType.kt b/lib/src/main/java/graphql/nadel/engine/util/NadelPseudoSealedType.kt index f23a3e2cb..4de97bf34 100644 --- a/lib/src/main/java/graphql/nadel/engine/util/NadelPseudoSealedType.kt +++ b/lib/src/main/java/graphql/nadel/engine/util/NadelPseudoSealedType.kt @@ -21,6 +21,7 @@ import graphql.language.ScalarTypeExtensionDefinition import graphql.language.SchemaDefinition import graphql.language.SchemaExtensionDefinition import graphql.language.Type +import graphql.language.TypeDefinition import graphql.language.TypeName import graphql.language.UnionTypeDefinition import graphql.language.UnionTypeExtensionDefinition @@ -48,7 +49,7 @@ inline fun GraphQLFieldsContainer.whenType( return when (this) { is GraphQLInterfaceType -> interfaceType(this) is GraphQLObjectType -> objectType(this) - else -> throw IllegalStateException("Should never happen") + else -> throw IllegalStateException(javaClass.name) } } @@ -61,7 +62,7 @@ inline fun GraphQLNamedInputType.whenType( is GraphQLEnumType -> enumType(this) is GraphQLInputObjectType -> inputObjectType(this) is GraphQLScalarType -> scalarType(this) - else -> throw IllegalStateException("Should never happen") + else -> throw IllegalStateException(javaClass.name) } } @@ -78,7 +79,7 @@ inline fun GraphQLNamedOutputType.whenType( is GraphQLObjectType -> objectType(this) is GraphQLScalarType -> scalarType(this) is GraphQLUnionType -> unionType(this) - else -> throw IllegalStateException("Should never happen") + else -> throw IllegalStateException(javaClass.name) } } @@ -97,7 +98,7 @@ inline fun GraphQLNamedType.whenType( is GraphQLObjectType -> objectType(this) is GraphQLScalarType -> scalarType(this) is GraphQLUnionType -> unionType(this) - else -> throw IllegalStateException("Should never happen") + else -> throw IllegalStateException(javaClass.name) } } @@ -114,7 +115,7 @@ inline fun GraphQLOutputType.whenUnmodifiedType( is GraphQLObjectType -> objectType(unmodifiedType) is GraphQLScalarType -> scalarType(unmodifiedType) is GraphQLUnionType -> unionType(unmodifiedType) - else -> throw IllegalStateException("Should never happen") + else -> throw IllegalStateException(javaClass.name) } } @@ -127,7 +128,7 @@ inline fun GraphQLInputType.whenUnmodifiedType( is GraphQLEnumType -> enumType(unmodifiedType) is GraphQLInputObjectType -> inputObjectType(unmodifiedType) is GraphQLScalarType -> scalarType(unmodifiedType) - else -> throw IllegalStateException("Should never happen") + else -> throw IllegalStateException(javaClass.name) } } @@ -140,7 +141,7 @@ inline fun GraphQLType.whenType( is GraphQLList -> listType(this) is GraphQLNonNull -> nonNull(this) is GraphQLUnmodifiedType -> unmodifiedType(this) - else -> throw IllegalStateException("Should never happen") + else -> throw IllegalStateException(javaClass.name) } } @@ -153,7 +154,7 @@ inline fun Type<*>.whenType( is ListType -> listType(this) is NonNullType -> nonNull(this) is TypeName -> unmodifiedType(this) - else -> throw IllegalStateException("Should never happen") + else -> throw IllegalStateException(javaClass.name) } } @@ -190,6 +191,70 @@ inline fun AnySDLDefinition.whenType( is SchemaDefinition -> schemaDefinition(this) is UnionTypeExtensionDefinition -> unionTypeExtensionDefinition(this) is UnionTypeDefinition -> unionTypeDefinition(this) - else -> throw IllegalStateException("Should never happen") + else -> throw IllegalStateException(javaClass.name) + } +} + +inline fun AnySDLNamedDefinition.whenType( + directiveDefinition: (DirectiveDefinition) -> T, + enumTypeDefinition: (EnumTypeDefinition) -> T, + enumTypeExtensionDefinition: (EnumTypeExtensionDefinition) -> T, + inputObjectTypeDefinition: (InputObjectTypeDefinition) -> T, + inputObjectTypeExtensionDefinition: (InputObjectTypeExtensionDefinition) -> T, + interfaceTypeDefinition: (InterfaceTypeDefinition) -> T, + interfaceTypeExtensionDefinition: (InterfaceTypeExtensionDefinition) -> T, + objectTypeDefinition: (ObjectTypeDefinition) -> T, + objectTypeExtensionDefinition: (ObjectTypeExtensionDefinition) -> T, + scalarTypeDefinition: (ScalarTypeDefinition) -> T, + scalarTypeExtensionDefinition: (ScalarTypeExtensionDefinition) -> T, + unionTypeDefinition: (UnionTypeDefinition) -> T, + unionTypeExtensionDefinition: (UnionTypeExtensionDefinition) -> T, +): T { + return when (this) { + is DirectiveDefinition -> directiveDefinition(this) + is EnumTypeExtensionDefinition -> enumTypeExtensionDefinition(this) + is EnumTypeDefinition -> enumTypeDefinition(this) + is InputObjectTypeExtensionDefinition -> inputObjectTypeExtensionDefinition(this) + is InputObjectTypeDefinition -> inputObjectTypeDefinition(this) + is InterfaceTypeExtensionDefinition -> interfaceTypeExtensionDefinition(this) + is InterfaceTypeDefinition -> interfaceTypeDefinition(this) + is ObjectTypeExtensionDefinition -> objectTypeExtensionDefinition(this) + is ObjectTypeDefinition -> objectTypeDefinition(this) + is ScalarTypeExtensionDefinition -> scalarTypeExtensionDefinition(this) + is ScalarTypeDefinition -> scalarTypeDefinition(this) + is UnionTypeExtensionDefinition -> unionTypeExtensionDefinition(this) + is UnionTypeDefinition -> unionTypeDefinition(this) + else -> throw IllegalStateException(javaClass.name) + } +} + +inline fun TypeDefinition<*>.whenType( + enumTypeDefinition: (EnumTypeDefinition) -> T, + enumTypeExtensionDefinition: (EnumTypeExtensionDefinition) -> T, + inputObjectTypeDefinition: (InputObjectTypeDefinition) -> T, + inputObjectTypeExtensionDefinition: (InputObjectTypeExtensionDefinition) -> T, + interfaceTypeDefinition: (InterfaceTypeDefinition) -> T, + interfaceTypeExtensionDefinition: (InterfaceTypeExtensionDefinition) -> T, + objectTypeDefinition: (ObjectTypeDefinition) -> T, + objectTypeExtensionDefinition: (ObjectTypeExtensionDefinition) -> T, + scalarTypeDefinition: (ScalarTypeDefinition) -> T, + scalarTypeExtensionDefinition: (ScalarTypeExtensionDefinition) -> T, + unionTypeDefinition: (UnionTypeDefinition) -> T, + unionTypeExtensionDefinition: (UnionTypeExtensionDefinition) -> T, +): T { + return when (this) { + is EnumTypeExtensionDefinition -> enumTypeExtensionDefinition(this) + is EnumTypeDefinition -> enumTypeDefinition(this) + is InputObjectTypeExtensionDefinition -> inputObjectTypeExtensionDefinition(this) + is InputObjectTypeDefinition -> inputObjectTypeDefinition(this) + is InterfaceTypeExtensionDefinition -> interfaceTypeExtensionDefinition(this) + is InterfaceTypeDefinition -> interfaceTypeDefinition(this) + is ObjectTypeExtensionDefinition -> objectTypeExtensionDefinition(this) + is ObjectTypeDefinition -> objectTypeDefinition(this) + is ScalarTypeExtensionDefinition -> scalarTypeExtensionDefinition(this) + is ScalarTypeDefinition -> scalarTypeDefinition(this) + is UnionTypeExtensionDefinition -> unionTypeExtensionDefinition(this) + is UnionTypeDefinition -> unionTypeDefinition(this) + else -> throw IllegalStateException(javaClass.name) } } diff --git a/lib/src/test/kotlin/graphql/nadel/archunit/NadelPseudoSealedTypeKtTest.kt b/lib/src/test/kotlin/graphql/nadel/archunit/NadelPseudoSealedTypeKtTest.kt index 6003b7ea9..b04479fdd 100644 --- a/lib/src/test/kotlin/graphql/nadel/archunit/NadelPseudoSealedTypeKtTest.kt +++ b/lib/src/test/kotlin/graphql/nadel/archunit/NadelPseudoSealedTypeKtTest.kt @@ -15,11 +15,13 @@ import graphql.language.NonNullType import graphql.language.ObjectTypeDefinition import graphql.language.ObjectTypeExtensionDefinition import graphql.language.SDLDefinition +import graphql.language.SDLNamedDefinition import graphql.language.ScalarTypeDefinition import graphql.language.ScalarTypeExtensionDefinition import graphql.language.SchemaDefinition import graphql.language.SchemaExtensionDefinition import graphql.language.Type +import graphql.language.TypeDefinition import graphql.language.TypeName import graphql.language.UnionTypeDefinition import graphql.language.UnionTypeExtensionDefinition @@ -194,4 +196,53 @@ class NadelPseudoSealedTypeKtTest { ) .check(schemaClasses) } + + @Test + fun `whenType(AnySDLNamedDefinition)`() { + classes() + .that() + .areAssignableTo(SDLNamedDefinition::class.java) + .and() + .areNotInterfaces() + .equalsExactly( + DirectiveDefinition::class, + EnumTypeDefinition::class, + EnumTypeExtensionDefinition::class, + InputObjectTypeDefinition::class, + InputObjectTypeExtensionDefinition::class, + InterfaceTypeDefinition::class, + InterfaceTypeExtensionDefinition::class, + ObjectTypeDefinition::class, + ObjectTypeExtensionDefinition::class, + ScalarTypeDefinition::class, + ScalarTypeExtensionDefinition::class, + UnionTypeDefinition::class, + UnionTypeExtensionDefinition::class, + ) + .check(schemaClasses) + } + + @Test + fun `whenType(TypeDefinition)`() { + classes() + .that() + .areAssignableTo(TypeDefinition::class.java) + .and() + .areNotInterfaces() + .equalsExactly( + EnumTypeDefinition::class, + EnumTypeExtensionDefinition::class, + InputObjectTypeDefinition::class, + InputObjectTypeExtensionDefinition::class, + InterfaceTypeDefinition::class, + InterfaceTypeExtensionDefinition::class, + ObjectTypeDefinition::class, + ObjectTypeExtensionDefinition::class, + ScalarTypeDefinition::class, + ScalarTypeExtensionDefinition::class, + UnionTypeDefinition::class, + UnionTypeExtensionDefinition::class, + ) + .check(schemaClasses) + } }