Skip to content
Merged
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 @@ -24,7 +24,7 @@ class ClassDefinitionGenerator {
): Option[GeneratedClassDefinitions] = {
val allSchemas: Map[String, OpenapiSchemaType] = doc.components.toSeq.flatMap(_.schemas).toMap
val allOneOfSchemas = allSchemas.collect { case (name, oneOf: OpenapiSchemaOneOf) => name -> oneOf }.toSeq
val adtInheritanceMap: Map[String, Seq[String]] = mkMapParentsByChild(allOneOfSchemas)
val adtInheritanceMap: Map[String, Seq[(String, OpenapiSchemaOneOf)]] = mkMapParentsByChild(allOneOfSchemas)
val generatesQueryOrPathParamEnums = enumsDefinedOnEndpointParams ||
allSchemas
.collect { case (name, _: OpenapiSchemaEnum) => name }
Expand All @@ -40,7 +40,7 @@ class ClassDefinitionGenerator {
jsonParamRefs.toSeq.flatMap(ref => allSchemas.get(ref.stripPrefix("#/components/schemas/")))
)

val adtTypes = adtInheritanceMap.flatMap(_._2).toSeq.distinct.map(name => s"sealed trait $name").mkString("", "\n", "\n")
val adtTypes = adtInheritanceMap.flatMap(_._2).toSeq.map(_._1).distinct.map(name => s"sealed trait $name").mkString("", "\n", "\n")
val enumSerdeHelper = if (!generatesQueryOrPathParamEnums) "" else enumSerdeHelperDefn(targetScala3)
val schemas = SchemaGenerator.generateSchemas(doc, allSchemas, fullModelPath, jsonSerdeLib, maxSchemasPerFile)
val jsonSerdes = JsonSerdeGenerator.serdeDefs(
Expand All @@ -50,7 +50,7 @@ class ClassDefinitionGenerator {
allTransitiveJsonParamRefs,
fullModelPath,
validateNonDiscriminatedOneOfs,
adtInheritanceMap,
adtInheritanceMap.mapValues(_.map(_._1)),
targetScala3
)
val defns = doc.components
Expand All @@ -71,7 +71,7 @@ class ClassDefinitionGenerator {
defns.map(helpers + "\n" + _).map(defStr => GeneratedClassDefinitions(defStr, jsonSerdes, schemas))
}

private def mkMapParentsByChild(allOneOfSchemas: Seq[(String, OpenapiSchemaOneOf)]): Map[String, Seq[String]] =
private def mkMapParentsByChild(allOneOfSchemas: Seq[(String, OpenapiSchemaOneOf)]): Map[String, Seq[(String, OpenapiSchemaOneOf)]] =
allOneOfSchemas
.flatMap { case (name, schema) =>
val validatedChildren = schema.types.map {
Expand All @@ -92,7 +92,7 @@ class ClassDefinitionGenerator {
s"Discriminator values $targetClassNames did not match schema variants $validatedChildren for oneOf defn $name"
)
}
validatedChildren.map(_ -> name)
validatedChildren.map(_ -> ((name, schema)))
}
.groupBy(_._1)
.mapValues(_.map(_._2))
Expand Down Expand Up @@ -203,7 +203,7 @@ class ClassDefinitionGenerator {
name: String,
obj: OpenapiSchemaObject,
jsonParamRefs: Set[String],
adtInheritanceMap: Map[String, Seq[String]],
adtInheritanceMap: Map[String, Seq[(String, OpenapiSchemaOneOf)]],
jsonSerdeLib: JsonSerdeLib.JsonSerdeLib,
targetScala3: Boolean
): Seq[String] = {
Expand All @@ -226,25 +226,44 @@ class ClassDefinitionGenerator {
.flatten
.toList

val (properties, maybeEnums) = obj.properties.map { case (key, OpenapiSchemaField(schemaType, maybeDefault)) =>
val (tpe, maybeEnum) = mapSchemaTypeToType(name, key, obj.required.contains(key), schemaType, isJson, jsonSerdeLib, targetScala3)
val fixedKey = fixKey(key)
val optional = schemaType.nullable || !obj.required.contains(key)
val maybeExplicitDefault =
maybeDefault.map(" = " + DefaultValueRenderer.render(allModels = allSchemas, thisType = schemaType, optional)(_))
val default = maybeExplicitDefault getOrElse (if (optional) " = None" else "")
s"$fixedKey: $tpe$default" -> maybeEnum
}.unzip

val parents = adtInheritanceMap.getOrElse(name, Nil) match {
case Nil => ""
case ps => ps.mkString(" extends ", " with ", "")
case ps => ps.map(_._1).mkString(" extends ", " with ", "")
}
val discriminatorDefFields = adtInheritanceMap
.getOrElse(name, Nil)
.flatMap { case (_, parent) =>
parent.discriminator.map { d =>
d.propertyName -> d.mapping.flatMap(_.find(_._2.stripPrefix("#/components/schemas/") == name).map(_._1)).getOrElse(name)
}
}
.distinct
val discriminatorDefBody = discriminatorDefFields.filter { case (n, _) => obj.properties.map(_._1).toSet.contains(n) } match {
case Nil => ""
case fields =>
val fs = fields.map { case (k, v) => s"""def `$k`: String = "$v"""" }.mkString("\n")
s""" {
|${indent(2)(fs)}
|}""".stripMargin
}

val (properties, maybeEnums) = obj.properties
.filterNot(discriminatorDefFields.map(_._1) contains _._1)
.map { case (key, OpenapiSchemaField(schemaType, maybeDefault)) =>
val (tpe, maybeEnum) = mapSchemaTypeToType(name, key, obj.required.contains(key), schemaType, isJson, jsonSerdeLib, targetScala3)
val fixedKey = fixKey(key)
val optional = schemaType.nullable || !obj.required.contains(key)
val maybeExplicitDefault =
maybeDefault.map(" = " + DefaultValueRenderer.render(allModels = allSchemas, thisType = schemaType, optional)(_))
val default = maybeExplicitDefault getOrElse (if (optional) " = None" else "")
s"$fixedKey: $tpe$default" -> maybeEnum
}
.unzip

val enumDefn = maybeEnums.flatten.toList
s"""|case class $name (
|${indent(2)(properties.mkString(",\n"))}
|)$parents""".stripMargin :: innerClasses ::: enumDefn ::: acc
|)$parents$discriminatorDefBody""".stripMargin :: innerClasses ::: enumDefn ::: acc
}

rec(addName("", name), obj, Nil)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -76,7 +76,9 @@ object TapirGeneratedEndpoints {
s: String,
i: Option[Int] = None,
d: Option[Double] = None
) extends ADTWithDiscriminator with ADTWithDiscriminatorNoMapping
) extends ADTWithDiscriminator with ADTWithDiscriminatorNoMapping {
def `type`: String = "SubA"
}
case class SubtypeWithoutD3 (
s: String,
i: Option[Int] = None,
Expand All @@ -103,7 +105,9 @@ object TapirGeneratedEndpoints {
case class SubtypeWithD2 (
s: String,
a: Option[Seq[String]] = None
) extends ADTWithDiscriminator with ADTWithDiscriminatorNoMapping
) extends ADTWithDiscriminator with ADTWithDiscriminatorNoMapping {
def `type`: String = "SubB"
}

sealed trait AnEnum extends enumeratum.EnumEntry
object AnEnum extends enumeratum.Enum[AnEnum] with enumeratum.CirceEnum[AnEnum] {
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -21,12 +21,12 @@ object TapirGeneratedEndpointsJsonSerdes {
implicit lazy val subtypeWithD1JsonDecoder: io.circe.Decoder[SubtypeWithD1] = io.circe.generic.semiauto.deriveDecoder[SubtypeWithD1]
implicit lazy val subtypeWithD1JsonEncoder: io.circe.Encoder[SubtypeWithD1] = io.circe.generic.semiauto.deriveEncoder[SubtypeWithD1]
implicit lazy val aDTWithDiscriminatorNoMappingJsonEncoder: io.circe.Encoder[ADTWithDiscriminatorNoMapping] = io.circe.Encoder.instance {
case x: SubtypeWithD1 => io.circe.Encoder[SubtypeWithD1].apply(x).mapObject(_.add("type", io.circe.Json.fromString("SubtypeWithD1")))
case x: SubtypeWithD2 => io.circe.Encoder[SubtypeWithD2].apply(x).mapObject(_.add("type", io.circe.Json.fromString("SubtypeWithD2")))
case x: SubtypeWithD1 => io.circe.Encoder[SubtypeWithD1].apply(x).mapObject(_.add("noMapType", io.circe.Json.fromString("SubtypeWithD1")))
case x: SubtypeWithD2 => io.circe.Encoder[SubtypeWithD2].apply(x).mapObject(_.add("noMapType", io.circe.Json.fromString("SubtypeWithD2")))
}
implicit lazy val aDTWithDiscriminatorNoMappingJsonDecoder: io.circe.Decoder[ADTWithDiscriminatorNoMapping] = io.circe.Decoder { (c: io.circe.HCursor) =>
for {
discriminator <- c.downField("type").as[String]
discriminator <- c.downField("noMapType").as[String]
res <- discriminator match {
case "SubtypeWithD1" => c.as[SubtypeWithD1]
case "SubtypeWithD2" => c.as[SubtypeWithD2]
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -28,7 +28,7 @@ object TapirGeneratedEndpointsSchemas {
val derived = implicitly[sttp.tapir.generic.Derived[sttp.tapir.Schema[ADTWithDiscriminatorNoMapping]]].value
derived.schemaType match {
case s: sttp.tapir.SchemaType.SCoproduct[_] => derived.copy(schemaType = s.addDiscriminatorField(
sttp.tapir.FieldName("type"),
sttp.tapir.FieldName("noMapType"),
sttp.tapir.Schema.string,
Map(
"SubtypeWithD1" -> sttp.tapir.SchemaType.SRef(sttp.tapir.Schema.SName("sttp.tapir.generated.TapirGeneratedEndpoints.SubtypeWithD1")),
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -106,7 +106,7 @@ class JsonRoundtrip extends AnyFreeSpec with Matchers {
val reqJsonBody = TapirGeneratedEndpointsJsonSerdes.aDTWithDiscriminatorNoMappingJsonEncoder(reqBody).noSpacesSortKeys
val respBody = SubtypeWithD1("a string+SubtypeWithD1", Some(123), Some(23.4))
val respJsonBody = TapirGeneratedEndpointsJsonSerdes.aDTWithDiscriminatorJsonEncoder(respBody).noSpacesSortKeys
reqJsonBody shouldEqual """{"d":23.4,"i":123,"s":"a string","type":"SubtypeWithD1"}"""
reqJsonBody shouldEqual """{"d":23.4,"i":123,"noMapType":"SubtypeWithD1","s":"a string"}"""
respJsonBody shouldEqual """{"d":23.4,"i":123,"s":"a string+SubtypeWithD1","type":"SubA"}"""
Await.result(
sttp.client3.basicRequest
Expand All @@ -126,7 +126,7 @@ class JsonRoundtrip extends AnyFreeSpec with Matchers {
val reqJsonBody = TapirGeneratedEndpointsJsonSerdes.aDTWithDiscriminatorNoMappingJsonEncoder(reqBody).noSpacesSortKeys
val respBody = SubtypeWithD2("a string+SubtypeWithD2", Some(Seq("string 1", "string 2")))
val respJsonBody = TapirGeneratedEndpointsJsonSerdes.aDTWithDiscriminatorJsonEncoder(respBody).noSpacesSortKeys
reqJsonBody shouldEqual """{"a":["string 1","string 2"],"s":"a string","type":"SubtypeWithD2"}"""
reqJsonBody shouldEqual """{"a":["string 1","string 2"],"noMapType":"SubtypeWithD2","s":"a string"}"""
respJsonBody shouldEqual """{"a":["string 1","string 2"],"s":"a string+SubtypeWithD2","type":"SubB"}"""
Await.result(
sttp.client3.basicRequest
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -119,12 +119,15 @@ components:
- $ref: '#/components/schemas/SubtypeWithD1'
- $ref: '#/components/schemas/SubtypeWithD2'
discriminator:
propertyName: type
propertyName: noMapType
SubtypeWithD1:
type: object
required:
- type
- s
properties:
type:
type: string
s:
type: string
i:
Expand All @@ -135,8 +138,11 @@ components:
SubtypeWithD2:
type: object
required:
- type
- s
properties:
type:
type: string
s:
type: string
a:
Expand Down