diff --git a/README.md b/README.md index 47420e8..80beca5 100644 --- a/README.md +++ b/README.md @@ -98,6 +98,7 @@ The following validations will be implemented: - ltefield (less than or equal field): field must be less than or equal to another field - ltfield (less than field): field must be less than another field - dive: tags after `dive` apply to each slice element, array element, or map value +- keys, endkeys: after `dive` on a map, tags between `keys` and `endkeys` apply to each key and tags after `endkeys` apply to each value ### dive @@ -110,7 +111,20 @@ type User struct { } ``` -`required` before `dive` checks that `Addresses` is not empty. After `dive`, each pointer must be non-nil and `Address` field tags run. Each `Labels` value must be non-empty. Map keys are not validated. `keys` and `endkeys` are not implemented. +`required` before `dive` checks that `Addresses` is not empty. After `dive`, each pointer must be non-nil and `Address` field tags run. Each `Labels` value must be non-empty. Plain `dive` validates map values. + +### keys and endkeys + +`keys` follows `dive` immediately and applies only to maps. Tags between `keys` and `endkeys` validate each map key. Tags after `endkeys` validate each map value. + +```go +type User struct { + Labels map[string]string `valid:"dive,keys,min=2,endkeys,required"` + Scores map[uint8]string `valid:"dive,keys,gte=1,endkeys,required"` +} +``` + +`min=2` checks each `Labels` key. `required` checks each `Labels` value. `gte=1` checks each `Scores` key. A missing `endkeys`, an extra `endkeys`, `keys` on a non-map, `dive` inside `keys`, another `keys` block, and a non-scalar map key are rejected. The following table shows the validations and possible types, where: diff --git a/docs/internals.md b/docs/internals.md index 3435549..68da345 100644 --- a/docs/internals.md +++ b/docs/internals.md @@ -48,7 +48,7 @@ An `*ast.StructType` node fills that struct and appends it. The field loop keeps A pointer appends `*` and then walks the inner expression. A slice walks the element and then appends `[]`. An array walks the element, then sets `Size` from the length literal and appends `[N]`. A map appends `map` and then walks `ast.MapType.Key`. -The map value is stored on `FieldType.MapValue`. Parser tests record `map[string]uint8` with `BaseType` `string` and `MapValue` `uint8`, and `map[uint8]string` with `BaseType` `uint8` and `MapValue` `string`. Collection checks still classify the map from the key type. `FieldType.ToType` still prints a map as `map[BaseType]BaseType`. `dive` reads `MapValue`. +The map value is stored on `FieldType.MapValue`. Parser tests record `map[string]uint8` with `BaseType` `string` and `MapValue` `uint8`, and `map[uint8]string` with `BaseType` `uint8` and `MapValue` `string`. Collection checks still classify the map from the key type. `FieldType.ToType` still prints a map as `map[BaseType]BaseType`. `dive` reads `MapValue`. `keys` reads the key through `FieldType.MapKey`, which is `BaseType` for a Go scalar key. The same walk records a pointer on an element separately from a pointer to the container. `[]*int64` is `BaseType` `int64`, `ComposedType` `*[]`, and `ElemPointer` true. `*[]int64` is the same `ComposedType` with `ElemPointer` false. `*map[string]bool` is `BaseType` `string` and `ComposedType` `*map`. A star that remains glued to another container after one `dive`, such as the element of `[][]*string`, is rejected. Nested arrays are rejected because only the outer length is kept. @@ -148,7 +148,7 @@ Those slice and map copies call `types.SliceOnlyContains`, `types.SliceNotContai When the field is a struct or a pointer to a struct, each validation on that field appends a nested call instead of a condition-table test. The call is `TypeValidate(&obj.Field)`, where `Type` is `BaseType`. If `BaseType` starts with the struct's own package name and a dot, that prefix is removed. A same-package field whose `BaseType` is `main.InnerStructType` calls `InnerStructTypeValidate`. A field whose `BaseType` is `mypkg.InnerStructType` calls `mypkg.InnerStructTypeValidate`. The call is emitted when `BaseType` is in the parsed-struct index. A missing type returns `no validator found for struct type`. -`dive` splits the tag list. Tags before `dive` use the field type. Tags after `dive` use `FieldType.DiveInto`, which is one slice element, one array element, or the map value. The generator writes a `for` loop. A struct element calls the nested validator once. A pointer element also checks `required` as non-nil and skips a nil pointer. `required`, `min`, `max`, and `len` on a slice or map whose element is not a Go type use the length check for `[]string` or `map[string]string`. The same length tags apply to `[]*T`. `in` and `nin` do not, because those helpers expect `[]T`. Field comparisons after `dive`, `keys`, `endkeys`, nested arrays, and a pointer glued to another container are rejected. +`dive` splits the tag list. Tags before `dive` use the field type. Tags after `dive` use `FieldType.DiveInto`, which is one slice element, one array element, or the map value. When the next tag is `keys`, `SplitKeysBlock` keeps that same list: tags between `keys` and `endkeys` use `FieldType.MapKey`, and tags after `endkeys` use the map value. The generator writes one `for` loop over the key and the value. A plain `dive` still validates map values only. A struct element calls the nested validator once. A pointer element also checks `required` as non-nil and skips a nil pointer. `required`, `min`, `max`, and `len` on a slice or map whose element is not a Go type use the length check for `[]string` or `map[string]string`. The same length tags apply to `[]*T`. `in` and `nin` do not, because those helpers expect `[]T`. Field comparisons after `dive`, a missing or extra `endkeys`, `keys` on a non-map, `dive` inside `keys`, another `keys` block, a non-scalar map key, nested arrays, and a pointer glued to another container are rejected. Imports kept on the generated package are the struct file's imports whose local name is a package name parsed in this run. `buildImportPath` writes each of those paths as a quoted import and always adds `github.com/opencodeco/validgen/types`. When any struct in the package has `UnmarshalJSON` source, it also adds `encoding/json` and `errors`. When generated code calls `strings.EqualFold`, it also adds `strings`. diff --git a/internal/analyzer/analyzer.go b/internal/analyzer/analyzer.go index a3ddd0e..acee3a2 100644 --- a/internal/analyzer/analyzer.go +++ b/internal/analyzer/analyzer.go @@ -89,34 +89,113 @@ func checkForInvalidOperations(structs []*Struct) error { for _, st := range structs { for i, fd := range st.Fields { - current := fd.Type - dived := false - for _, val := range st.FieldsValidations[i].Validations { - op := val.Operation - if !ops.IsValid(op) { - return types.NewValidationError("unsupported operation %s", op) - } + err := checkValidations(ops, fd.FieldName, fd.Type, st.FieldsValidations[i].Validations, structsWithValidation, false, 0) + if err != nil { + return err + } + } + } - if op == "dive" { - next, err := current.DiveInto() - if err != nil { - return types.NewValidationError("operation dive: field %s: %s", fd.FieldName, err.Error()) - } - current = next - dived = true - continue - } + return nil +} - if err := validateOperation(ops, op, current, structsWithValidation, dived); err != nil { - return err +func checkValidations(ops *operations.Operations, fieldName string, current common.FieldType, validations []*Validation, structs map[string]bool, dived bool, keysDepth int) error { + for i, val := range validations { + op := val.Operation + if !ops.IsValid(op) { + return types.NewValidationError("unsupported operation %s", op) + } + + switch op { + case "keys": + return types.NewValidationError("operation keys: field %s: keys must immediately follow dive", fieldName) + case "endkeys": + return types.NewValidationError("operation endkeys: field %s: endkeys without keys", fieldName) + case "dive": + if i+1 < len(validations) && validations[i+1].Operation == "keys" { + if keysDepth > 0 { + return types.NewValidationError("operation keys: field %s: nested keys are not supported", fieldName) } + return checkMapKeys(ops, fieldName, current, validations[i+1:], structs) } + + next, err := current.DiveInto() + if err != nil { + return types.NewValidationError("operation dive: field %s: %s", fieldName, err.Error()) + } + current = next + dived = true + continue + } + + if err := validateOperation(ops, op, current, structs, dived); err != nil { + return err } } return nil } +func checkMapKeys(ops *operations.Operations, fieldName string, current common.FieldType, validations []*Validation, structs map[string]bool) error { + keyType, err := current.MapKey() + if err != nil { + return types.NewValidationError("operation keys: field %s: %s", fieldName, err.Error()) + } + + valueType, err := current.DiveInto() + if err != nil { + return types.NewValidationError("operation dive: field %s: %s", fieldName, err.Error()) + } + + keyVals, valueVals, err := SplitKeysBlock(fieldName, validations) + if err != nil { + return err + } + + if err := checkValidations(ops, fieldName, keyType, keyVals, structs, true, 1); err != nil { + return err + } + + return checkValidations(ops, fieldName, valueType, valueVals, structs, true, 1) +} + +// SplitKeysBlock splits the validations that follow dive, starting at keys. +// Tags between keys and endkeys apply to the map key. Tags after endkeys apply to the map value. +func SplitKeysBlock(fieldName string, rest []*Validation) ([]*Validation, []*Validation, error) { + if len(rest) == 0 || rest[0].Operation != "keys" { + return nil, nil, types.NewValidationError("operation keys: field %s: keys must immediately follow dive", fieldName) + } + + end := -1 + for i := 1; i < len(rest); i++ { + switch rest[i].Operation { + case "endkeys": + end = i + case "dive": + return nil, nil, types.NewValidationError("operation dive: field %s: nested dive inside keys is not supported", fieldName) + case "keys": + return nil, nil, types.NewValidationError("operation keys: field %s: nested keys are not supported", fieldName) + } + if end != -1 { + break + } + } + if end == -1 { + return nil, nil, types.NewValidationError("operation keys: field %s: missing endkeys", fieldName) + } + + for _, val := range rest[end+1:] { + switch val.Operation { + case "keys": + return nil, nil, types.NewValidationError("operation keys: field %s: nested keys are not supported", fieldName) + case "endkeys": + return nil, nil, types.NewValidationError("operation endkeys: field %s: extra endkeys", fieldName) + } + } + + return rest[1:end], rest[end+1:], nil +} + func validateOperation(ops *operations.Operations, op string, ft common.FieldType, structs map[string]bool, dived bool) error { if dived && ops.IsFieldOperation(op) { return types.NewValidationError("operation %s: field comparisons are not supported after dive", op) diff --git a/internal/analyzer/dive_test.go b/internal/analyzer/dive_test.go index 54fbd7d..4dbec1a 100644 --- a/internal/analyzer/dive_test.go +++ b/internal/analyzer/dive_test.go @@ -135,7 +135,7 @@ func TestAnalyzeDiveRejected(t *testing.T) { wantErr: types.NewValidationError("operation dive: field Matrix: unsupported pointer composition *[]string"), }, { - name: "keys stays unsupported", + name: "missing endkeys", field: parser.Field{ FieldName: "Labels", Type: common.FieldType{ @@ -145,7 +145,7 @@ func TestAnalyzeDiveRejected(t *testing.T) { }, Tag: `valid:"dive,keys,required"`, }, - wantErr: types.NewValidationError("parser validation keys: unsupported validation keys"), + wantErr: types.NewValidationError("operation keys: field Labels: missing endkeys"), }, { name: "in on a slice of pointers", @@ -181,3 +181,165 @@ func TestAnalyzeDiveRejected(t *testing.T) { }) } } + +func TestAnalyzeMapKeys(t *testing.T) { + stringMap := common.FieldType{ + BaseType: "string", + ComposedType: "map", + MapValue: &common.FieldType{BaseType: "string"}, + } + uintMap := common.FieldType{ + BaseType: "uint8", + ComposedType: "map", + MapValue: &common.FieldType{BaseType: "string"}, + } + tests := []struct { + name string + field parser.Field + wantErr error + }{ + { + name: "string keys and values", + field: parser.Field{ + FieldName: "Labels", + Type: stringMap, + Tag: `valid:"dive,keys,min=2,endkeys,required"`, + }, + }, + { + name: "typed keys and values", + field: parser.Field{ + FieldName: "Scores", + Type: uintMap, + Tag: `valid:"dive,keys,gte=1,endkeys,required"`, + }, + }, + { + name: "value validation after endkeys", + field: parser.Field{ + FieldName: "Labels", + Type: stringMap, + Tag: `valid:"required,dive,keys,min=2,endkeys,email"`, + }, + }, + { + name: "dive into values after endkeys", + field: parser.Field{ + FieldName: "Labels", + Type: common.FieldType{ + BaseType: "string", + ComposedType: "map", + MapValue: &common.FieldType{BaseType: "string", ComposedType: "[]"}, + }, + Tag: `valid:"dive,keys,min=2,endkeys,dive,required"`, + }, + }, + { + name: "plain dive still validates map values", + field: parser.Field{ + FieldName: "Labels", + Type: stringMap, + Tag: `valid:"dive,required"`, + }, + }, + { + name: "extra endkeys", + field: parser.Field{ + FieldName: "Labels", + Type: stringMap, + Tag: `valid:"dive,keys,min=2,endkeys,required,endkeys"`, + }, + wantErr: types.NewValidationError("operation endkeys: field Labels: extra endkeys"), + }, + { + name: "endkeys without keys", + field: parser.Field{ + FieldName: "Labels", + Type: stringMap, + Tag: `valid:"dive,endkeys,required"`, + }, + wantErr: types.NewValidationError("operation endkeys: field Labels: endkeys without keys"), + }, + { + name: "keys does not follow dive", + field: parser.Field{ + FieldName: "Labels", + Type: stringMap, + Tag: `valid:"dive,required,keys,min=2,endkeys"`, + }, + wantErr: types.NewValidationError("operation keys: field Labels: keys must immediately follow dive"), + }, + { + name: "keys on a slice", + field: parser.Field{ + FieldName: "Names", + Type: common.FieldType{BaseType: "string", ComposedType: "[]"}, + Tag: `valid:"dive,keys,min=2,endkeys,required"`, + }, + wantErr: types.NewValidationError("operation keys: field Names: []string is not a map"), + }, + { + name: "keys on a string", + field: parser.Field{ + FieldName: "Name", + Type: common.FieldType{BaseType: "string"}, + Tag: `valid:"dive,keys,min=2,endkeys"`, + }, + wantErr: types.NewValidationError("operation keys: field Name: string is not a map"), + }, + { + name: "dive inside keys", + field: parser.Field{ + FieldName: "Labels", + Type: stringMap, + Tag: `valid:"dive,keys,dive,min=2,endkeys,required"`, + }, + wantErr: types.NewValidationError("operation dive: field Labels: nested dive inside keys is not supported"), + }, + { + name: "nested keys", + field: parser.Field{ + FieldName: "Labels", + Type: common.FieldType{ + BaseType: "string", + ComposedType: "map", + MapValue: &common.FieldType{ + BaseType: "string", + ComposedType: "map", + MapValue: &common.FieldType{BaseType: "string"}, + }, + }, + Tag: `valid:"dive,keys,min=2,endkeys,dive,keys,min=2,endkeys,required"`, + }, + wantErr: types.NewValidationError("operation keys: field Labels: nested keys are not supported"), + }, + { + name: "array key", + field: parser.Field{ + FieldName: "Labels", + Type: common.FieldType{ + BaseType: "string", + ComposedType: "map[N]", + Size: "2", + MapValue: &common.FieldType{BaseType: "string"}, + }, + Tag: `valid:"dive,keys,min=2,endkeys,required"`, + }, + wantErr: types.NewValidationError("operation keys: field Labels: nested map keys are not supported"), + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + user := &parser.Struct{ + PackageName: "main", + StructName: "User", + Fields: []parser.Field{tt.field}, + } + _, err := AnalyzeStructs([]*parser.Struct{addressStruct(), user}) + if err != tt.wantErr { + t.Fatalf("AnalyzeStructs() error = %v, want %v", err, tt.wantErr) + } + }) + } +} diff --git a/internal/analyzer/operations/operations_list.go b/internal/analyzer/operations/operations_list.go index c565813..9d0bdcb 100644 --- a/internal/analyzer/operations/operations_list.go +++ b/internal/analyzer/operations/operations_list.go @@ -130,10 +130,21 @@ var operationsList = map[string]Operation{ IsFieldOperation: true, ValidTypes: []string{"", ""}, }, - // dive is a level separator. The analyzer checks the container, then the element. + // dive, keys, and endkeys are level separators. + // The analyzer checks the container, then the map key, then the element. "dive": { CountValues: common.ZeroValue, IsFieldOperation: false, ValidTypes: nil, }, + "keys": { + CountValues: common.ZeroValue, + IsFieldOperation: false, + ValidTypes: nil, + }, + "endkeys": { + CountValues: common.ZeroValue, + IsFieldOperation: false, + ValidTypes: nil, + }, } diff --git a/internal/analyzer/operations/operations_test.go b/internal/analyzer/operations/operations_test.go index c9b8606..8c6c70d 100644 --- a/internal/analyzer/operations/operations_test.go +++ b/internal/analyzer/operations/operations_test.go @@ -35,6 +35,8 @@ func TestOperationsIsValid(t *testing.T) { {op: "ltefield", want: true}, {op: "ltfield", want: true}, {op: "dive", want: true}, + {op: "keys", want: true}, + {op: "endkeys", want: true}, {op: "invalid_op", want: false}, } @@ -500,6 +502,8 @@ func TestOperationsIsFieldOperation(t *testing.T) { {op: "ltefield", want: true}, {op: "ltfield", want: true}, {op: "dive", want: false}, + {op: "keys", want: false}, + {op: "endkeys", want: false}, {op: "invalid_op", want: false}, } @@ -541,6 +545,8 @@ func TestOperationsArgsCount(t *testing.T) { {op: "ltefield", want: common.OneValue}, {op: "ltfield", want: common.OneValue}, {op: "dive", want: common.ZeroValue}, + {op: "keys", want: common.ZeroValue}, + {op: "endkeys", want: common.ZeroValue}, {op: "invalid_op", want: common.UndefinedValue}, } diff --git a/internal/analyzer/parser_validation_test.go b/internal/analyzer/parser_validation_test.go index 59cdaaf..8772239 100644 --- a/internal/analyzer/parser_validation_test.go +++ b/internal/analyzer/parser_validation_test.go @@ -141,6 +141,24 @@ func TestValidParserValidation(t *testing.T) { Values: []string{}, }, }, + { + name: "keys tag", + validation: "keys", + want: &Validation{ + Operation: "keys", + ExpectedValues: common.ZeroValue, + Values: []string{}, + }, + }, + { + name: "endkeys tag", + validation: "endkeys", + want: &Validation{ + Operation: "endkeys", + ExpectedValues: common.ZeroValue, + Values: []string{}, + }, + }, } for _, tt := range tests { @@ -185,14 +203,14 @@ func TestParserInvalidValidation(t *testing.T) { expectedErr: types.NewValidationError("unsupported validation xpto"), }, { - name: "keys is not implemented", - validation: "keys", - expectedErr: types.NewValidationError("unsupported validation keys"), + name: "keys with a target", + validation: "keys=a", + expectedErr: types.NewValidationError("expected zero target, but has a"), }, { - name: "endkeys is not implemented", - validation: "endkeys", - expectedErr: types.NewValidationError("unsupported validation endkeys"), + name: "endkeys with a target", + validation: "endkeys=a", + expectedErr: types.NewValidationError("expected zero target, but has a"), }, { name: "malformed value", diff --git a/internal/codegenerator/build_validator.go b/internal/codegenerator/build_validator.go index 8a2d492..c555d24 100644 --- a/internal/codegenerator/build_validator.go +++ b/internal/codegenerator/build_validator.go @@ -8,6 +8,7 @@ import ( "github.com/opencodeco/validgen/internal/analyzer" "github.com/opencodeco/validgen/internal/common" + "github.com/opencodeco/validgen/types" ) var funcValidatorTpl = `func {{.StructName}}Validate(obj *{{.StructName}}) []error { @@ -66,18 +67,34 @@ func (gv *GenValidations) BuildUnmarshalJSONCode() string { } func (gv *GenValidations) BuildValidationCode(fieldName string, fieldType common.FieldType, fieldValidations []*analyzer.Validation) (string, error) { - return gv.emitValidations("obj."+fieldName, fieldName, fieldType, fieldValidations, false, 0) + return gv.emitValidations("obj."+fieldName, fieldName, fieldType, fieldValidations, false, 0, 0) } -func (gv *GenValidations) emitValidations(expr, fieldName string, fieldType common.FieldType, fieldValidations []*analyzer.Validation, dived bool, depth int) (string, error) { +func (gv *GenValidations) emitValidations(expr, fieldName string, fieldType common.FieldType, fieldValidations []*analyzer.Validation, dived bool, depth, keysDepth int) (string, error) { tests := "" for i, fieldValidation := range fieldValidations { - if fieldValidation.Operation == "dive" { - loop, err := gv.emitDive(expr, fieldName, fieldType, fieldValidations[i+1:], depth) + switch fieldValidation.Operation { + case "dive": + if i+1 < len(fieldValidations) && fieldValidations[i+1].Operation == "keys" { + if keysDepth > 0 { + return "", types.NewValidationError("operation keys: field %s: nested keys are not supported", fieldName) + } + loop, err := gv.emitMapKeys(expr, fieldName, fieldType, fieldValidations[i+1:], depth) + if err != nil { + return "", err + } + return tests + loop, nil + } + + loop, err := gv.emitDive(expr, fieldName, fieldType, fieldValidations[i+1:], depth, keysDepth) if err != nil { return "", err } return tests + loop, nil + case "keys": + return "", types.NewValidationError("operation keys: field %s: keys must immediately follow dive", fieldName) + case "endkeys": + return "", types.NewValidationError("operation endkeys: field %s: endkeys without keys", fieldName) } testCode, err := gv.emitOne(expr, fieldName, fieldType, fieldValidation, dived) @@ -98,7 +115,7 @@ func (gv *GenValidations) emitValidations(expr, fieldName string, fieldType comm return tests, nil } -func (gv *GenValidations) emitDive(expr, fieldName string, fieldType common.FieldType, rest []*analyzer.Validation, depth int) (string, error) { +func (gv *GenValidations) emitDive(expr, fieldName string, fieldType common.FieldType, rest []*analyzer.Validation, depth, keysDepth int) (string, error) { elemType, err := fieldType.DiveInto() if err != nil { return "", fmt.Errorf("field %s: %w", fieldName, err) @@ -106,7 +123,7 @@ func (gv *GenValidations) emitDive(expr, fieldName string, fieldType common.Fiel depth++ elemExpr := fmt.Sprintf("elem%d", depth) - body, err := gv.emitValidations(elemExpr, fieldName, elemType, rest, true, depth) + body, err := gv.emitValidations(elemExpr, fieldName, elemType, rest, true, depth, keysDepth) if err != nil { return "", err } @@ -126,6 +143,63 @@ func (gv *GenValidations) emitDive(expr, fieldName string, fieldType common.Fiel return prefix + fmt.Sprintf("for _, %s := range %s {\n%s}\n", elemExpr, rangeExpr, body) + suffix, nil } +func (gv *GenValidations) emitMapKeys(expr, fieldName string, fieldType common.FieldType, rest []*analyzer.Validation, depth int) (string, error) { + keyType, err := fieldType.MapKey() + if err != nil { + return "", types.NewValidationError("operation keys: field %s: %s", fieldName, err.Error()) + } + valueType, err := fieldType.DiveInto() + if err != nil { + return "", fmt.Errorf("field %s: %w", fieldName, err) + } + + keyVals, valueVals, err := analyzer.SplitKeysBlock(fieldName, rest) + if err != nil { + return "", err + } + + depth++ + keyExpr := fmt.Sprintf("key%d", depth) + elemExpr := fmt.Sprintf("elem%d", depth) + keyBody, err := gv.emitValidations(keyExpr, fieldName, keyType, keyVals, true, depth, 1) + if err != nil { + return "", err + } + valueBody, err := gv.emitValidations(elemExpr, fieldName, valueType, valueVals, true, depth, 1) + if err != nil { + return "", err + } + body := keyBody + valueBody + if strings.TrimSpace(body) == "" { + return "", nil + } + + keyName := "_" + if strings.TrimSpace(keyBody) != "" { + keyName = keyExpr + } + elemName := "_" + if strings.TrimSpace(valueBody) != "" { + elemName = elemExpr + } + + rangeExpr := expr + prefix := "" + suffix := "" + if fieldType.PointerToContainer() { + rangeExpr = "*" + expr + prefix = fmt.Sprintf("if %s != nil {\n", expr) + suffix = "}\n" + } + + loop := fmt.Sprintf("for %s, %s := range %s {\n%s}\n", keyName, elemName, rangeExpr, body) + if elemName == "_" { + loop = fmt.Sprintf("for %s := range %s {\n%s}\n", keyName, rangeExpr, body) + } + + return prefix + loop + suffix, nil +} + func (gv *GenValidations) emitOne(expr, fieldName string, fieldType common.FieldType, fieldValidation *analyzer.Validation, dived bool) (string, error) { if fieldType.IsNestedStruct() { if !dived { diff --git a/internal/codegenerator/dive_test.go b/internal/codegenerator/dive_test.go index 609c7d8..107438b 100644 --- a/internal/codegenerator/dive_test.go +++ b/internal/codegenerator/dive_test.go @@ -124,6 +124,196 @@ errs = append(errs, types.NewValidationError("Matrix is required")) } } +func TestBuildMapKeyValidationCode(t *testing.T) { + stringMap := common.FieldType{ + BaseType: "string", + ComposedType: "map", + MapValue: &common.FieldType{BaseType: "string"}, + } + uintMap := common.FieldType{ + BaseType: "uint8", + ComposedType: "map", + MapValue: &common.FieldType{BaseType: "string"}, + } + tests := []struct { + name string + fieldName string + fieldType common.FieldType + fieldValidation string + want string + wantErr string + }{ + { + name: "string keys and values", + fieldName: "Labels", + fieldType: stringMap, + fieldValidation: "dive,keys,min=2,endkeys,required", + want: `for key1, elem1 := range obj.Labels { +if !(len(key1) >= 2) { +errs = append(errs, types.NewValidationError("Labels length must be >= 2")) +} +if !(elem1 != "") { +errs = append(errs, types.NewValidationError("Labels is required")) +} +} +`, + }, + { + name: "typed keys and values", + fieldName: "Scores", + fieldType: uintMap, + fieldValidation: "dive,keys,gte=1,endkeys,required", + want: `for key1, elem1 := range obj.Scores { +if !(key1 >= 1) { +errs = append(errs, types.NewValidationError("Scores must be >= 1")) +} +if !(elem1 != "") { +errs = append(errs, types.NewValidationError("Scores is required")) +} +} +`, + }, + { + name: "value validation after endkeys", + fieldName: "Labels", + fieldType: stringMap, + fieldValidation: "required,dive,keys,min=2,endkeys,email", + want: `if !(len(obj.Labels) != 0) { +errs = append(errs, types.NewValidationError("Labels must not be empty")) +} +for key1, elem1 := range obj.Labels { +if !(len(key1) >= 2) { +errs = append(errs, types.NewValidationError("Labels length must be >= 2")) +} +if !(types.IsValidEmail(elem1)) { +errs = append(errs, types.NewValidationError("Labels must be a valid email")) +} +} +`, + }, + { + name: "dive into values after endkeys", + fieldName: "Labels", + fieldType: common.FieldType{ + BaseType: "string", + ComposedType: "map", + MapValue: &common.FieldType{BaseType: "string", ComposedType: "[]"}, + }, + fieldValidation: "dive,keys,min=2,endkeys,dive,required", + want: `for key1, elem1 := range obj.Labels { +if !(len(key1) >= 2) { +errs = append(errs, types.NewValidationError("Labels length must be >= 2")) +} +for _, elem2 := range elem1 { +if !(elem2 != "") { +errs = append(errs, types.NewValidationError("Labels is required")) +} +} +} +`, + }, + { + name: "keys only", + fieldName: "Labels", + fieldType: stringMap, + fieldValidation: "dive,keys,min=2,endkeys", + want: `for key1 := range obj.Labels { +if !(len(key1) >= 2) { +errs = append(errs, types.NewValidationError("Labels length must be >= 2")) +} +} +`, + }, + { + name: "pointer to a map", + fieldName: "Labels", + fieldType: common.FieldType{BaseType: "string", ComposedType: "*map", MapValue: &common.FieldType{BaseType: "string"}}, + fieldValidation: "dive,keys,min=2,endkeys,required", + want: `if obj.Labels != nil { +for key1, elem1 := range *obj.Labels { +if !(len(key1) >= 2) { +errs = append(errs, types.NewValidationError("Labels length must be >= 2")) +} +if !(elem1 != "") { +errs = append(errs, types.NewValidationError("Labels is required")) +} +} +} +`, + }, + { + name: "missing endkeys", + fieldName: "Labels", + fieldType: stringMap, + fieldValidation: "dive,keys,required", + wantErr: "operation keys: field Labels: missing endkeys", + }, + { + name: "extra endkeys", + fieldName: "Labels", + fieldType: stringMap, + fieldValidation: "dive,keys,min=2,endkeys,endkeys", + wantErr: "operation endkeys: field Labels: extra endkeys", + }, + { + name: "keys on a slice", + fieldName: "Names", + fieldType: common.FieldType{BaseType: "string", ComposedType: "[]"}, + fieldValidation: "dive,keys,min=2,endkeys,required", + wantErr: "operation keys: field Names: []string is not a map", + }, + { + name: "nested keys", + fieldName: "Labels", + fieldType: common.FieldType{ + BaseType: "string", + ComposedType: "map", + MapValue: &common.FieldType{ + BaseType: "string", + ComposedType: "map", + MapValue: &common.FieldType{BaseType: "string"}, + }, + }, + fieldValidation: "dive,keys,min=2,endkeys,dive,keys,min=2,endkeys,required", + wantErr: "operation keys: field Labels: nested keys are not supported", + }, + { + name: "dive inside keys", + fieldName: "Labels", + fieldType: stringMap, + fieldValidation: "dive,keys,dive,min=2,endkeys,required", + wantErr: "operation dive: field Labels: nested dive inside keys is not supported", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + gv := GenValidations{ + Struct: &analyzer.Struct{ + Struct: parser.Struct{PackageName: "main"}, + }, + } + var validations []*analyzer.Validation + for _, part := range splitValidations(tt.fieldValidation) { + validations = append(validations, AssertParserValidation(t, part)) + } + got, err := gv.BuildValidationCode(tt.fieldName, tt.fieldType, validations) + if tt.wantErr != "" { + if err == nil || err.Error() != tt.wantErr { + t.Fatalf("BuildValidationCode() error = %v, want %s", err, tt.wantErr) + } + return + } + if err != nil { + t.Fatalf("BuildValidationCode() error = %v", err) + } + if got != tt.want { + t.Fatalf("BuildValidationCode() = %q, want %q", got, tt.want) + } + }) + } +} + func splitValidations(tag string) []string { if tag == "" { return nil diff --git a/internal/common/dive.go b/internal/common/dive.go index e2f91f4..15722f9 100644 --- a/internal/common/dive.go +++ b/internal/common/dive.go @@ -93,6 +93,24 @@ func (ft FieldType) OperationType(op string, accept func(FieldType) bool) (Field return FieldType{}, false } +// MapKey returns the key of a map. +// The key must be a Go scalar. Array, slice, map, and struct keys are rejected. +func (ft FieldType) MapKey() (FieldType, error) { + ct := ft.ComposedType + if strings.HasPrefix(ct, "*") && !ft.ElemPointer { + ct = strings.TrimPrefix(ct, "*") + } + if !strings.HasPrefix(ct, "map") { + return FieldType{}, fmt.Errorf("%s is not a map", ft.ToType()) + } + key := FieldType{BaseType: ft.BaseType} + if ct != "map" || !key.IsGoType() { + return FieldType{}, fmt.Errorf("nested map keys are not supported") + } + + return key, nil +} + // DiveInto returns the type of one slice element, array element, or map value. func (ft FieldType) DiveInto() (FieldType, error) { ct := ft.ComposedType diff --git a/internal/common/dive_test.go b/internal/common/dive_test.go index 5b9f994..042910c 100644 --- a/internal/common/dive_test.go +++ b/internal/common/dive_test.go @@ -2,6 +2,80 @@ package common import "testing" +func TestMapKey(t *testing.T) { + tests := []struct { + name string + in FieldType + want FieldType + wantErr string + }{ + { + name: "string key", + in: FieldType{ + BaseType: "string", + ComposedType: "map", + MapValue: &FieldType{BaseType: "string"}, + }, + want: FieldType{BaseType: "string"}, + }, + { + name: "uint8 key", + in: FieldType{ + BaseType: "uint8", + ComposedType: "map", + MapValue: &FieldType{BaseType: "string"}, + }, + want: FieldType{BaseType: "uint8"}, + }, + { + name: "pointer to a map", + in: FieldType{ + BaseType: "string", + ComposedType: "*map", + MapValue: &FieldType{BaseType: "int"}, + }, + want: FieldType{BaseType: "string"}, + }, + { + name: "slice", + in: FieldType{BaseType: "string", ComposedType: "[]"}, + wantErr: "[]string is not a map", + }, + { + name: "array key", + in: FieldType{BaseType: "string", ComposedType: "map[N]", Size: "2"}, + wantErr: "nested map keys are not supported", + }, + { + name: "struct key", + in: FieldType{ + BaseType: "main.Address", + ComposedType: "map", + MapValue: &FieldType{BaseType: "string"}, + }, + wantErr: "nested map keys are not supported", + }, + } + + for _, tt := range tests { + t.Run(tt.name, func(t *testing.T) { + got, err := tt.in.MapKey() + if tt.wantErr != "" { + if err == nil || err.Error() != tt.wantErr { + t.Fatalf("MapKey() error = %v, want %s", err, tt.wantErr) + } + return + } + if err != nil { + t.Fatalf("MapKey() error = %v", err) + } + if got != tt.want { + t.Fatalf("MapKey() = %#v, want %#v", got, tt.want) + } + }) + } +} + func TestDiveInto(t *testing.T) { tests := []struct { name string diff --git a/internal/parser/dive_test.go b/internal/parser/dive_test.go index 4841ecf..4367fad 100644 --- a/internal/parser/dive_test.go +++ b/internal/parser/dive_test.go @@ -55,6 +55,42 @@ func TestParseDiveFieldTypes(t *testing.T) { } } +func TestParseMapKeyTags(t *testing.T) { + src := "package main\n" + + "type User struct {\n" + + " Labels map[string]string `valid:\"dive,keys,min=2,endkeys,required\"`\n" + + " Scores map[uint8]string `valid:\"dive,keys,gte=1,endkeys,required\"`\n" + + " Broken map[string]string `valid:\"dive,keys,required\"`\n" + + "}\n" + + got, err := parseStructs("example/main.go", src) + if err != nil { + t.Fatal(err) + } + + stringValue := &common.FieldType{BaseType: "string"} + want := []Field{ + { + FieldName: "Labels", + Type: common.FieldType{BaseType: "string", ComposedType: "map", MapValue: stringValue}, + Tag: `valid:"dive,keys,min=2,endkeys,required"`, + }, + { + FieldName: "Scores", + Type: common.FieldType{BaseType: "uint8", ComposedType: "map", MapValue: stringValue}, + Tag: `valid:"dive,keys,gte=1,endkeys,required"`, + }, + { + FieldName: "Broken", + Type: common.FieldType{BaseType: "string", ComposedType: "map", MapValue: stringValue}, + Tag: `valid:"dive,keys,required"`, + }, + } + if !reflect.DeepEqual(got[0].Fields, want) { + t.Fatalf("fields mismatch\ngot:\n%swant:\n%s", formatStructs(got[:1]), formatFields(want)) + } +} + func formatFields(fields []Field) string { s := &Struct{Fields: fields} return formatStructs([]*Struct{s}) diff --git a/tests/endtoend/dive.go b/tests/endtoend/dive.go index 7dca52c..8ea1e44 100644 --- a/tests/endtoend/dive.go +++ b/tests/endtoend/dive.go @@ -15,6 +15,11 @@ type PointerDiveUser struct { Addresses []*Address `valid:"required,dive,required"` } +type KeyedDiveUser struct { + Labels map[string]string `valid:"dive,keys,min=2,endkeys,required"` + Scores map[uint8]string `valid:"dive,keys,gte=1,endkeys,required"` +} + func diveTests() { log.Println("starting dive tests") @@ -61,5 +66,55 @@ func diveTests() { pointerOK := &PointerDiveUser{Addresses: []*Address{{Street: "av 123", City: "city 123"}}} assertExpectedErrorMsgs("pointer dive valid", PointerDiveUserValidate(pointerOK), nil) + keyTests() + log.Println("dive tests ok") } + +func keyTests() { + log.Println("starting keys tests") + + shortKey := &KeyedDiveUser{ + Labels: map[string]string{"a": "earth"}, + Scores: map[uint8]string{1: "ok"}, + } + assertExpectedErrorMsgs("string key failure", KeyedDiveUserValidate(shortKey), []string{ + "Labels length must be >= 2", + }) + + lowScore := &KeyedDiveUser{ + Labels: map[string]string{"home": "earth"}, + Scores: map[uint8]string{0: "ok"}, + } + assertExpectedErrorMsgs("typed key failure", KeyedDiveUserValidate(lowScore), []string{ + "Scores must be >= 1", + }) + + emptyValue := &KeyedDiveUser{ + Labels: map[string]string{"home": ""}, + Scores: map[uint8]string{1: ""}, + } + assertExpectedErrorMsgs("value validation after endkeys", KeyedDiveUserValidate(emptyValue), []string{ + "Labels is required", + "Scores is required", + }) + + both := &KeyedDiveUser{ + Labels: map[string]string{"a": ""}, + Scores: map[uint8]string{0: ""}, + } + assertExpectedErrorMsgs("key and value failure", KeyedDiveUserValidate(both), []string{ + "Labels length must be >= 2", + "Labels is required", + "Scores must be >= 1", + "Scores is required", + }) + + ok := &KeyedDiveUser{ + Labels: map[string]string{"home": "earth"}, + Scores: map[uint8]string{1: "ok"}, + } + assertExpectedErrorMsgs("keys valid", KeyedDiveUserValidate(ok), nil) + + log.Println("keys tests ok") +} diff --git a/tests/endtoend/validator__.go b/tests/endtoend/validator__.go index e854a3d..776ac39 100755 --- a/tests/endtoend/validator__.go +++ b/tests/endtoend/validator__.go @@ -326,6 +326,26 @@ func DiveUserValidate(obj *DiveUser) []error { } return errs } +func KeyedDiveUserValidate(obj *KeyedDiveUser) []error { + var errs []error + for key1, elem1 := range obj.Labels { + if !(len(key1) >= 2) { + errs = append(errs, types.NewValidationError("Labels length must be >= 2")) + } + if !(elem1 != "") { + errs = append(errs, types.NewValidationError("Labels is required")) + } + } + for key1, elem1 := range obj.Scores { + if !(key1 >= 1) { + errs = append(errs, types.NewValidationError("Scores must be >= 1")) + } + if !(elem1 != "") { + errs = append(errs, types.NewValidationError("Scores is required")) + } + } + return errs +} func PointerDiveUserValidate(obj *PointerDiveUser) []error { var errs []error if !(len(obj.Addresses) != 0) {