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
10 changes: 9 additions & 1 deletion binding.go
Original file line number Diff line number Diff line change
Expand Up @@ -696,7 +696,15 @@ func tryBindNestedField(
if rawMap, ok := entry.value.(map[string]any); ok {
nestedData := make(map[string]mergedEntry)
for k, v := range rawMap {
nestedData[k] = mergedEntry{value: v, sourceName: entry.sourceName}
normalizedKey := strings.ToLower(k)
if _, exists := nestedData[normalizedKey]; exists {
return true, []FieldError{{
FieldPath: fieldPath + "." + normalizedKey,
Code: ErrCodeInvalidType,
Message: "duplicate configuration keys after normalization",
}}
}
nestedData[normalizedKey] = mergedEntry{value: v, sourceName: entry.sourceName}
}
return true, bindNested(fieldValue, nestedData, "", fieldPath)
}
Expand Down
235 changes: 231 additions & 4 deletions loader.go
Original file line number Diff line number Diff line change
Expand Up @@ -143,18 +143,23 @@ func (l *Loader[T]) loadInternal(ctx context.Context, store bool) (*T, *Provenan

// Step 3: In strict mode, detect unknown keys
if l.strict {
validKeys := collectValidKeys(rootType, "")
dynamicMapKeyPatterns := collectDynamicMapKeyPatterns(rootType, "")
strictMetadataCache := make(map[strictValidationMetadataKey]strictValidationMetadata)
strictMetadata := getStrictValidationMetadata(rootType, "", strictMetadataCache)

// Check for unknown keys
var unknownKeyErrors []FieldError
for key := range mergedData {
if !validKeys[key] && !matchesDynamicMapKeyPattern(key, dynamicMapKeyPatterns) {
for key, entry := range mergedData {
if !strictMetadata.validKeys[key] && !matchesDynamicMapKeyPattern(key, strictMetadata.dynamicMapKeyPatterns) {
unknownKeyErrors = append(unknownKeyErrors, FieldError{
FieldPath: key,
Code: ErrCodeUnknownKey,
Message: "unknown configuration key (strict mode)",
})
continue
}

if fieldType, ok := strictMetadata.validKeyTypes[key]; ok {
unknownKeyErrors = append(unknownKeyErrors, collectStructuredUnknownKeyErrors(entry.value, fieldType, key, strictMetadataCache)...)
}
}

Expand Down Expand Up @@ -305,6 +310,87 @@ func collectValidKeys(t reflect.Type, prefix string) map[string]bool {
return validKeys
}

func collectValidKeyTypes(t reflect.Type, prefix string) map[string]reflect.Type {
validKeyTypes := make(map[string]reflect.Type)

if t.Kind() == reflect.Ptr {
t = t.Elem()
}
if t.Kind() != reflect.Struct {
return validKeyTypes
}

for _, meta := range getStructFieldMeta(t) {
field := meta.field
tagCfg := meta.tagCfg
keyPath := determineKeyPath(field.Name, tagCfg, prefix)
fieldType := field.Type

validKeyTypes[keyPath] = fieldType

if isOptionalType(fieldType) {
innerType := fieldType.Field(0).Type
if innerType.Kind() == reflect.Struct {
for nestedKey, nestedType := range collectValidKeyTypes(innerType, keyPath) {
validKeyTypes[nestedKey] = nestedType
}
}
continue
}

if fieldType.Kind() != reflect.Struct || fieldType.PkgPath() == "time" {
continue
}

nestedPrefix := keyPath
if tagCfg.prefix != "" {
nestedPrefix = tagCfg.prefix
}
for nestedKey, nestedType := range collectValidKeyTypes(fieldType, nestedPrefix) {
validKeyTypes[nestedKey] = nestedType
}
}

return validKeyTypes
}

type strictValidationMetadata struct {
validKeys map[string]bool
validKeyTypes map[string]reflect.Type
dynamicMapKeyPatterns []string
}

type strictValidationMetadataKey struct {
targetType reflect.Type
prefix string
}

func getStrictValidationMetadata(
targetType reflect.Type,
prefix string,
cache map[strictValidationMetadataKey]strictValidationMetadata,
) strictValidationMetadata {
for targetType.Kind() == reflect.Ptr {
targetType = targetType.Elem()
}

cacheKey := strictValidationMetadataKey{
targetType: targetType,
prefix: prefix,
}
if metadata, ok := cache[cacheKey]; ok {
return metadata
}

metadata := strictValidationMetadata{
validKeys: collectValidKeys(targetType, prefix),
validKeyTypes: collectValidKeyTypes(targetType, prefix),
dynamicMapKeyPatterns: collectDynamicMapKeyPatterns(targetType, prefix),
}
cache[cacheKey] = metadata
return metadata
}

func collectDynamicMapKeyPatterns(t reflect.Type, prefix string) []string {
patternSet := make(map[string]struct{})

Expand Down Expand Up @@ -364,6 +450,147 @@ func collectDynamicMapKeyPatterns(t reflect.Type, prefix string) []string {
return patterns
}

func collectStructuredUnknownKeyErrors(
rawValue any,
targetType reflect.Type,
keyPath string,
metadataCache map[strictValidationMetadataKey]strictValidationMetadata,
) []FieldError {
if rawValue == nil {
return nil
}

targetType = unwrapStrictType(targetType)

switch targetType.Kind() {
case reflect.Struct:
if targetType == timeType || targetType == durationType {
return nil
}

rawMap, ok := rawStringMap(rawValue)
if !ok {
return nil
}

return collectUnknownStructMapKeys(rawMap, targetType, keyPath, metadataCache)

case reflect.Slice, reflect.Array:
elemType := unwrapStrictType(targetType.Elem())
if elemType.Kind() != reflect.Struct || elemType == timeType || elemType == durationType {
return nil
}

rawRef := reflect.ValueOf(rawValue)
if rawRef.Kind() != reflect.Slice && rawRef.Kind() != reflect.Array {
return nil
}

var fieldErrors []FieldError
for i := 0; i < rawRef.Len(); i++ {
elemMap, ok := rawStringMap(rawRef.Index(i).Interface())
if !ok {
continue
}

fieldErrors = append(fieldErrors, collectUnknownStructMapKeys(elemMap, elemType, fmt.Sprintf("%s.%d", keyPath, i), metadataCache)...)
}
return fieldErrors

case reflect.Map:
if targetType.Key().Kind() != reflect.String {
return nil
}

elemType := unwrapStrictType(targetType.Elem())
if elemType.Kind() != reflect.Struct || elemType == timeType || elemType == durationType {
return nil
}

rawMapRef := reflect.ValueOf(rawValue)
if rawMapRef.Kind() != reflect.Map {
return nil
}

var fieldErrors []FieldError
for _, rawKey := range rawMapRef.MapKeys() {
if rawKey.Kind() != reflect.String {
continue
}

elemMap, ok := rawStringMap(rawMapRef.MapIndex(rawKey).Interface())
if !ok {
continue
}

fieldErrors = append(fieldErrors, collectUnknownStructMapKeys(elemMap, elemType, keyPath+"."+strings.ToLower(rawKey.String()), metadataCache)...)
}
return fieldErrors
}

return nil
}

func collectUnknownStructMapKeys(
rawMap map[string]any,
targetType reflect.Type,
keyPath string,
metadataCache map[strictValidationMetadataKey]strictValidationMetadata,
) []FieldError {
strictMetadata := getStrictValidationMetadata(targetType, "", metadataCache)

var fieldErrors []FieldError
for rawKey, rawValue := range rawMap {
normalizedKey := strings.ToLower(rawKey)
if !strictMetadata.validKeys[normalizedKey] && !matchesDynamicMapKeyPattern(normalizedKey, strictMetadata.dynamicMapKeyPatterns) {
fieldErrors = append(fieldErrors, FieldError{
FieldPath: keyPath + "." + normalizedKey,
Code: ErrCodeUnknownKey,
Message: "unknown configuration key (strict mode)",
})
Comment thread
Azhovan marked this conversation as resolved.
continue
}

if nestedType, ok := strictMetadata.validKeyTypes[normalizedKey]; ok {
fieldErrors = append(fieldErrors, collectStructuredUnknownKeyErrors(rawValue, nestedType, keyPath+"."+normalizedKey, metadataCache)...)
}
}

return fieldErrors
}

func unwrapStrictType(t reflect.Type) reflect.Type {
for t.Kind() == reflect.Ptr {
t = t.Elem()
}
if isOptionalType(t) {
return unwrapStrictType(t.Field(0).Type)
}
return t
}

func rawStringMap(rawValue any) (map[string]any, bool) {
rawMap, ok := rawValue.(map[string]any)
if ok {
return rawMap, true
}

rawRef := reflect.ValueOf(rawValue)
if !rawRef.IsValid() || rawRef.Kind() != reflect.Map {
return nil, false
}

result := make(map[string]any, rawRef.Len())
for _, rawKey := range rawRef.MapKeys() {
if rawKey.Kind() != reflect.String {
return nil, false
}
result[rawKey.String()] = rawRef.MapIndex(rawKey).Interface()
}

return result, true
}

func collectMapValueKeyPatterns(baseKey string, valueType reflect.Type) []string {
if valueType.Kind() == reflect.Ptr {
valueType = valueType.Elem()
Expand Down
Loading
Loading