Skip to content

Commit 9dd51a8

Browse files
authored
Merge pull request #20 from AVSystem/fix/discriminated-union-subtypes
fix: generate discriminated union subtypes for oneOf schemas
2 parents 4e1d413 + 17ea730 commit 9dd51a8

4 files changed

Lines changed: 295 additions & 7 deletions

File tree

core/src/main/kotlin/com/avsystem/justworks/core/parser/SpecParser.kt

Lines changed: 28 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
package com.avsystem.justworks.core.parser
22

3+
import arrow.core.compareTo
34
import arrow.core.fold
45
import arrow.core.merge
56
import arrow.core.raise.context.Raise
@@ -117,11 +118,29 @@ object SpecParser {
117118
}
118119
}
119120

121+
// Pick up synthetic schemas added by detectAndUnwrapOneOfWrappers.
122+
// Iterate until stable, since processing a synthetic schema could register more.
123+
tailrec fun collectModels(processed: Set<String>, acc: List<SchemaModel>): List<SchemaModel> {
124+
val currentKeys = componentSchemas.keys - allSchemas.keys - processed
125+
return if (currentKeys.isEmpty()) {
126+
acc
127+
} else {
128+
val newModels = currentKeys
129+
.asSequence()
130+
.mapNotNull { name -> componentSchemas[name]?.let { name to it } }
131+
.filterNot { (_, schema) -> schema.isEnumSchema }
132+
.map { (name, schema) -> extractSchemaModel(name, schema) }
133+
134+
collectModels(processed + currentKeys, acc + newModels)
135+
}
136+
}
137+
138+
val syntheticModels = collectModels(emptySet(), emptyList())
120139
return ApiSpec(
121140
title = info?.title ?: "Untitled",
122141
version = info?.version ?: "0.0.0",
123142
endpoints = endpoints,
124-
schemas = schemaModels,
143+
schemas = schemaModels + syntheticModels,
125144
enums = enumModels,
126145
)
127146
}
@@ -301,12 +320,14 @@ object SpecParser {
301320
)
302321

303322
val schemaName = ensureNotNull(
304-
propertySchema.resolveName() ?: propertyName
305-
.takeIf { propertySchema.isInlineObject }
306-
?.also { name ->
307-
componentSchemas[name] = propertySchema
308-
componentSchemaIdentity[propertySchema] = name
309-
},
323+
propertySchema.resolveName()
324+
?: propertyName
325+
.takeIf { propertySchema.isInlineObject }
326+
?.let { rawName ->
327+
componentSchemas[rawName] = propertySchema
328+
componentSchemaIdentity[propertySchema] = rawName
329+
rawName
330+
},
310331
)
311332

312333
propertyName to schemaName

core/src/test/kotlin/com/avsystem/justworks/core/gen/ModelGeneratorPolymorphicTest.kt

Lines changed: 164 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -612,6 +612,170 @@ class ModelGeneratorPolymorphicTest {
612612
)
613613
}
614614

615+
// -- CEM-01: boolean discriminator names (KotlinPoet handles escaping) --
616+
617+
@Test
618+
fun `boolean discriminator names produce valid data classes`() {
619+
val deviceStatusSchema = schema(
620+
name = "DeviceStatus",
621+
oneOf = listOf(
622+
TypeRef.Reference("true"),
623+
TypeRef.Reference("false"),
624+
),
625+
discriminator = Discriminator(
626+
propertyName = "online",
627+
mapping = mapOf(
628+
"true" to "#/components/schemas/true",
629+
"false" to "#/components/schemas/false",
630+
),
631+
),
632+
)
633+
val trueSchema = schema(
634+
name = "true",
635+
properties = listOf(
636+
PropertyModel("connectedSince", TypeRef.Primitive(PrimitiveType.STRING), null, false),
637+
),
638+
requiredProperties = setOf("connectedSince"),
639+
)
640+
val falseSchema = schema(
641+
name = "false",
642+
properties = listOf(
643+
PropertyModel("lastSeen", TypeRef.Primitive(PrimitiveType.STRING), null, false),
644+
),
645+
requiredProperties = setOf("lastSeen"),
646+
)
647+
648+
val files = generator.generate(
649+
spec(schemas = listOf(deviceStatusSchema, trueSchema, falseSchema)),
650+
)
651+
652+
val trueType = findType(files, "true")
653+
assertTrue(KModifier.DATA in trueType.modifiers, "'true' should be data class")
654+
655+
val falseType = findType(files, "false")
656+
assertTrue(KModifier.DATA in falseType.modifiers, "'false' should be data class")
657+
658+
// Both implement DeviceStatus sealed interface
659+
val trueSuperinterfaces = trueType.superinterfaces.keys.map { it.toString() }
660+
assertTrue(
661+
"$modelPackage.DeviceStatus" in trueSuperinterfaces,
662+
"'true' should implement DeviceStatus. Superinterfaces: $trueSuperinterfaces",
663+
)
664+
val falseSuperinterfaces = falseType.superinterfaces.keys.map { it.toString() }
665+
assertTrue(
666+
"$modelPackage.DeviceStatus" in falseSuperinterfaces,
667+
"'false' should implement DeviceStatus. Superinterfaces: $falseSuperinterfaces",
668+
)
669+
}
670+
671+
@Test
672+
fun `all oneOf variant schemas generate data classes even with many subtypes`() {
673+
val variantNames = listOf(
674+
"ExtenderDevice",
675+
"EthernetDevice",
676+
"WanDevice",
677+
"USBDevice",
678+
"WiFiDevice",
679+
"OtherDevice",
680+
)
681+
682+
val networkMeshSchema = schema(
683+
name = "NetworkMeshDevice",
684+
oneOf = variantNames.map { TypeRef.Reference(it) },
685+
discriminator = Discriminator(
686+
propertyName = "deviceType",
687+
mapping = variantNames.associateWith { "#/components/schemas/$it" },
688+
),
689+
)
690+
691+
val variantSchemas = variantNames.map { name ->
692+
schema(
693+
name = name,
694+
properties = listOf(
695+
PropertyModel("deviceId", TypeRef.Primitive(PrimitiveType.STRING), null, false),
696+
),
697+
requiredProperties = setOf("deviceId"),
698+
)
699+
}
700+
701+
val files = generator.generate(
702+
spec(schemas = listOf(networkMeshSchema) + variantSchemas),
703+
)
704+
705+
// All 6 variants generated
706+
for (name in variantNames) {
707+
val variantType = findType(files, name)
708+
assertTrue(
709+
KModifier.DATA in variantType.modifiers,
710+
"$name should be a data class",
711+
)
712+
val superinterfaces = variantType.superinterfaces.keys.map { it.toString() }
713+
assertTrue(
714+
"$modelPackage.NetworkMeshDevice" in superinterfaces,
715+
"$name should implement NetworkMeshDevice. Superinterfaces: $superinterfaces",
716+
)
717+
}
718+
719+
// SerializersModule contains all variants
720+
val serializersModuleFile = files.find { it.name == "SerializersModule" }
721+
assertNotNull(serializersModuleFile, "SerializersModule file should be generated")
722+
val moduleCode = serializersModuleFile.toString()
723+
for (name in variantNames) {
724+
assertTrue(
725+
name in moduleCode,
726+
"SerializersModule should reference $name. Code: $moduleCode",
727+
)
728+
}
729+
}
730+
731+
@Test
732+
fun `SerializersModule includes boolean variant names`() {
733+
val deviceStatusSchema = schema(
734+
name = "DeviceStatus",
735+
oneOf = listOf(
736+
TypeRef.Reference("true"),
737+
TypeRef.Reference("false"),
738+
),
739+
discriminator = Discriminator(
740+
propertyName = "online",
741+
mapping = mapOf(
742+
"true" to "#/components/schemas/true",
743+
"false" to "#/components/schemas/false",
744+
),
745+
),
746+
)
747+
val trueSchema = schema(
748+
name = "true",
749+
properties = listOf(
750+
PropertyModel("connectedSince", TypeRef.Primitive(PrimitiveType.STRING), null, false),
751+
),
752+
requiredProperties = setOf("connectedSince"),
753+
)
754+
val falseSchema = schema(
755+
name = "false",
756+
properties = listOf(
757+
PropertyModel("lastSeen", TypeRef.Primitive(PrimitiveType.STRING), null, false),
758+
),
759+
requiredProperties = setOf("lastSeen"),
760+
)
761+
762+
val files = generator.generate(
763+
spec(schemas = listOf(deviceStatusSchema, trueSchema, falseSchema)),
764+
)
765+
766+
val serializersModuleFile = files.find { it.name == "SerializersModule" }
767+
assertNotNull(serializersModuleFile, "SerializersModule file should be generated")
768+
val moduleCode = serializersModuleFile.toString()
769+
assertTrue(
770+
"`true`" in moduleCode,
771+
"SerializersModule should reference `true`. Code: $moduleCode",
772+
)
773+
assertTrue(
774+
"`false`" in moduleCode,
775+
"SerializersModule should reference `false`. Code: $moduleCode",
776+
)
777+
}
778+
615779
// -- POLY-06: allOf with sealed parent --
616780

617781
@Test

core/src/test/kotlin/com/avsystem/justworks/core/parser/SpecParserPolymorphicTest.kt

Lines changed: 78 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -56,4 +56,82 @@ class SpecParserPolymorphicTest : SpecParserTestBase() {
5656
assertEquals("#/components/schemas/Circle", discriminator.mapping["circle"])
5757
assertEquals("#/components/schemas/Square", discriminator.mapping["square"])
5858
}
59+
60+
// -- Synthetic schemas from wrapper unwrapping --
61+
62+
@Test
63+
fun `boolean discriminator spec preserves original schema names`() {
64+
val spec = parseSpec(loadResource("boolean-discriminator-spec.yaml"))
65+
val schemaNames = spec.schemas.map { it.name }.toSet()
66+
67+
assertTrue(
68+
"true" in schemaNames,
69+
"Expected schema 'true' in output — KotlinPoet handles escaping. Schemas: $schemaNames",
70+
)
71+
assertTrue(
72+
"false" in schemaNames,
73+
"Expected schema 'false' in output — KotlinPoet handles escaping. Schemas: $schemaNames",
74+
)
75+
}
76+
77+
@Test
78+
fun `boolean discriminator mapping preserves original values as keys`() {
79+
val spec = parseSpec(loadResource("boolean-discriminator-spec.yaml"))
80+
81+
val deviceStatus =
82+
spec.schemas.find { it.name == "DeviceStatus" }
83+
?: fail("DeviceStatus schema not found. Schemas: ${spec.schemas.map { it.name }}")
84+
85+
val discriminator = assertNotNull(deviceStatus.discriminator, "DeviceStatus should have discriminator")
86+
val mappingKeys = discriminator.mapping.keys
87+
88+
assertTrue(
89+
"true" in mappingKeys,
90+
"Discriminator mapping should have 'true' as key. Keys: $mappingKeys",
91+
)
92+
assertTrue(
93+
"false" in mappingKeys,
94+
"Discriminator mapping should have 'false' as key. Keys: $mappingKeys",
95+
)
96+
97+
// Values reference original schema names
98+
assertTrue(
99+
discriminator.mapping["true"]!!.endsWith("true"),
100+
"Mapping for 'true' should reference 'true'. Value: ${discriminator.mapping["true"]}",
101+
)
102+
assertTrue(
103+
discriminator.mapping["false"]!!.endsWith("false"),
104+
"Mapping for 'false' should reference 'false'. Value: ${discriminator.mapping["false"]}",
105+
)
106+
}
107+
108+
@Test
109+
fun `wrapper-unwrapped synthetic schemas appear in parsed output`() {
110+
val spec = parseSpec(loadResource("boolean-discriminator-spec.yaml"))
111+
val schemaNames = spec.schemas.map { it.name }.toSet()
112+
113+
// DeviceStatus parent + 2 synthetic variants
114+
assertTrue(
115+
"DeviceStatus" in schemaNames,
116+
"Parent schema 'DeviceStatus' should be in output. Schemas: $schemaNames",
117+
)
118+
assertTrue(
119+
spec.schemas.size >= 3,
120+
"Should have at least 3 schemas (parent + 2 variants). Got: ${spec.schemas.size}. Schemas: $schemaNames",
121+
)
122+
}
123+
124+
@Test
125+
fun `polymorphic-spec regression test - all schemas present`() {
126+
val spec = parseSpec(loadResource("polymorphic-spec.yaml"))
127+
val schemaNames = spec.schemas.map { it.name }.toSet()
128+
129+
val expectedSchemas = setOf("Shape", "Circle", "Square", "Pet", "Cat", "Dog", "ExtendedDog")
130+
for (expected in expectedSchemas) {
131+
assertTrue(
132+
expected in schemaNames,
133+
"Expected schema '$expected' in output. Schemas: $schemaNames",
134+
)
135+
}
136+
}
59137
}
Lines changed: 25 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,25 @@
1+
openapi: '3.0.0'
2+
info:
3+
title: Boolean Discriminator Test
4+
version: '1.0'
5+
paths: {}
6+
components:
7+
schemas:
8+
DeviceStatus:
9+
oneOf:
10+
- type: object
11+
properties:
12+
"true":
13+
type: object
14+
properties:
15+
connectedSince:
16+
type: string
17+
format: date-time
18+
- type: object
19+
properties:
20+
"false":
21+
type: object
22+
properties:
23+
lastSeen:
24+
type: string
25+
format: date-time

0 commit comments

Comments
 (0)