diff --git a/schema/enum_test.go b/schema/enum_test.go new file mode 100644 index 000000000..8737fa551 --- /dev/null +++ b/schema/enum_test.go @@ -0,0 +1,55 @@ +package schema + +import ( + "encoding/json" + "reflect" + "testing" + + "github.com/stretchr/testify/require" +) + +func TestGenerateTypedEnums(t *testing.T) { + type level int + type options struct { + Count int `json:"count" enum:"1, 2"` + Ratio float64 `json:"ratio" enum:"0.5, 1.5"` + Enabled bool `json:"enabled" enum:"true, false"` + Level *level `json:"level" enum:"1, 2"` + Label string `json:"label" enum:"true, 1"` + } + s := Generate(reflect.TypeFor[options]()) + _, err := ParseAndValidate(`{"count":1,"ratio":0.5,"enabled":false,"level":2,"label":"true"}`, s) + require.NoError(t, err) + + _, err = ParseAndValidate(`{"count":3,"ratio":0.5,"enabled":false,"level":2,"label":"true"}`, s) + require.Error(t, err) + _, err = ParseAndValidate(`{"count":1,"ratio":2.5,"enabled":false,"level":2,"label":"true"}`, s) + require.Error(t, err) + require.Equal(t, []any{"true", "1"}, s.Properties["label"].Enum) +} + +func TestGenerateIntegerEnumsPreservePrecision(t *testing.T) { + type options struct { + Signed int64 `json:"signed" enum:"-9223372036854775808,9223372036854775807"` + Unsigned uint64 `json:"unsigned" enum:"18446744073709551615"` + } + s := Generate(reflect.TypeFor[options]()) + data, err := json.Marshal(ToParameters(s)) + require.NoError(t, err) + require.Contains(t, string(data), `"enum":[-9223372036854775808,9223372036854775807]`) + require.Contains(t, string(data), `"enum":[18446744073709551615]`) +} + +func TestGenerateInvalidNumericEnumsKeepOriginalValues(t *testing.T) { + type options struct { + Count int `json:"count" enum:"invalid"` + Ratio float64 `json:"ratio" enum:"NaN,null"` + Enabled bool `json:"enabled" enum:"invalid"` + } + s := Generate(reflect.TypeFor[options]()) + require.Equal(t, []any{"invalid"}, s.Properties["count"].Enum) + require.Equal(t, []any{"NaN", "null"}, s.Properties["ratio"].Enum) + require.Equal(t, []any{"invalid"}, s.Properties["enabled"].Enum) + _, err := json.Marshal(s) + require.NoError(t, err) +} diff --git a/schema/schema.go b/schema/schema.go index 92ccd7b34..24407c560 100644 --- a/schema/schema.go +++ b/schema/schema.go @@ -8,6 +8,7 @@ import ( "fmt" "reflect" "slices" + "strconv" "strings" "charm.land/fantasy/jsonrepair" @@ -160,7 +161,7 @@ func generateSchemaRecursive(t reflect.Type, visited map[reflect.Type]bool) Sche enumValues := strings.Split(enumTag, ",") fieldSchema.Enum = make([]any, len(enumValues)) for i, v := range enumValues { - fieldSchema.Enum[i] = strings.TrimSpace(v) + fieldSchema.Enum[i] = parseEnumValue(strings.TrimSpace(v), fieldSchema.Type) } } @@ -179,6 +180,28 @@ func generateSchemaRecursive(t reflect.Type, visited map[reflect.Type]bool) Sche } } +func parseEnumValue(value, schemaType string) any { + switch schemaType { + case "integer": + if n, err := strconv.ParseInt(value, 10, 64); err == nil { + return n + } + if n, err := strconv.ParseUint(value, 10, 64); err == nil { + return n + } + case "number": + var n json.Number + if err := json.Unmarshal([]byte(value), &n); err == nil && n != "" { + return n + } + case "boolean": + if value == "true" || value == "false" { + return value == "true" + } + } + return value +} + // ToMap converts a Schema to a map representation suitable for JSON Schema. func ToMap(schema Schema) map[string]any { result := make(map[string]any)