diff --git a/src/MiniValidation/MiniValidation.csproj b/src/MiniValidation/MiniValidation.csproj index b48eb90..ed40f1a 100644 --- a/src/MiniValidation/MiniValidation.csproj +++ b/src/MiniValidation/MiniValidation.csproj @@ -5,7 +5,7 @@ netstandard2.0;net8.0 ComponentModel DataAnnotations validation README.md - 10.0 + 11.0 diff --git a/src/MiniValidation/MiniValidator.cs b/src/MiniValidation/MiniValidator.cs index 99c8f36..51e6db2 100644 --- a/src/MiniValidation/MiniValidator.cs +++ b/src/MiniValidation/MiniValidator.cs @@ -1,4 +1,4 @@ -using System; +using System; using System.Collections; using System.Collections.Generic; using System.Collections.ObjectModel; @@ -54,6 +54,12 @@ public static bool RequiresValidation(Type targetType, bool recurse = true) /// A dictionary that contains details of each failed validation. /// true if is valid; otherwise false. /// Thrown when is null. + /// + /// + /// var widget = new Widget { Name = "" }; + /// var isValid = MiniValidator.TryValidate(widget, out var errors); + /// + /// public static bool TryValidate(TTarget target, out IDictionary errors) { return TryValidateImpl(target, null, recurse: true, allowAsync: false, out errors); diff --git a/src/MiniValidation/TypeDetailsCache.cs b/src/MiniValidation/TypeDetailsCache.cs index 2086de3..2e7e67e 100644 --- a/src/MiniValidation/TypeDetailsCache.cs +++ b/src/MiniValidation/TypeDetailsCache.cs @@ -1,4 +1,4 @@ -using System; +using System; using System.Collections.Concurrent; using System.Collections.Generic; using System.ComponentModel; @@ -6,6 +6,9 @@ using System.Diagnostics.CodeAnalysis; using System.Linq; using System.Reflection; +using System.Runtime.CompilerServices; + +[assembly: InternalsVisibleTo("MiniValidation.UnitTests")] namespace MiniValidation; @@ -236,6 +239,8 @@ private static (ValidationAttribute[]?, DisplayAttribute?, SkipRecursionAttribut .Where(attr => !IsDuplicateTypeDescriptorAttribute(attr, propertyAttributes))); } + var hasRequiredMemberAttribute = false; + foreach (var attr in customAttributes) { if (attr is ValidationAttribute validationAttr) @@ -251,11 +256,113 @@ private static (ValidationAttribute[]?, DisplayAttribute?, SkipRecursionAttribut { skipRecursionAttribute = skipRecursionAttr; } + else if (string.Equals(attr.GetType().FullName, "System.Runtime.CompilerServices.RequiredMemberAttribute", StringComparison.Ordinal)) + { + hasRequiredMemberAttribute = true; + } + } + + if (hasRequiredMemberAttribute && !property.PropertyType.IsValueType && !IsReferenceTypeNullable(property)) + { + validationAttributes ??= new(); + if (!validationAttributes.OfType().Any()) + { + validationAttributes.Add(new RequiredAttribute()); + } } return new(validationAttributes?.ToArray(), displayAttribute, skipRecursionAttribute); } + internal static bool IsReferenceTypeNullable(PropertyInfo property) + { +#if NET6_0_OR_GREATER + // Create context per lookup for thread safety during concurrent cache initialization + var nullabilityContext = new NullabilityInfoContext(); + var nullabilityInfo = nullabilityContext.Create(property); + return nullabilityInfo.WriteState == NullabilityState.Nullable || nullabilityInfo.ReadState == NullabilityState.Nullable; +#else + return IsReferenceTypeNullableFallback(property); +#endif + } + + internal static bool IsReferenceTypeNullableFallback(PropertyInfo property) + { + if (HasNullableFlowAttribute(property)) + { + return true; + } + + var nullableAttr = property.GetCustomAttributes(false) + .FirstOrDefault(attr => string.Equals(attr.GetType().FullName, "System.Runtime.CompilerServices.NullableAttribute", StringComparison.Ordinal)); + + if (nullableAttr != null) + { + var flagsField = nullableAttr.GetType().GetField("NullableFlags"); + if (flagsField?.GetValue(nullableAttr) is byte[] flags && flags.Length > 0) + { + return flags[0] == 2; + } + } + + var declaringType = property.DeclaringType; + while (declaringType != null) + { + var nullableContextAttr = declaringType.GetCustomAttributes(false) + .FirstOrDefault(attr => string.Equals(attr.GetType().FullName, "System.Runtime.CompilerServices.NullableContextAttribute", StringComparison.Ordinal)); + + if (nullableContextAttr != null) + { + var flagField = nullableContextAttr.GetType().GetField("Flag"); + if (flagField?.GetValue(nullableContextAttr) is byte flag) + { + return flag == 2; + } + } + declaringType = declaringType.DeclaringType; + } + + return false; + } + + private static bool HasNullableFlowAttribute(PropertyInfo property) + { + if (HasAllowOrMaybeNullAttribute(property.GetCustomAttributes(false))) + { + return true; + } + + if (property.GetMethod is { } getMethod && HasAllowOrMaybeNullAttribute(getMethod.ReturnParameter.GetCustomAttributes(false))) + { + return true; + } + + if (property.SetMethod is { } setMethod) + { + var setParams = setMethod.GetParameters(); + if (setParams.Length > 0 && HasAllowOrMaybeNullAttribute(setParams[setParams.Length - 1].GetCustomAttributes(false))) + { + return true; + } + } + + return false; + } + + private static bool HasAllowOrMaybeNullAttribute(object[] attributes) + { + foreach (var attr in attributes) + { + var fullName = attr.GetType().FullName; + if (string.Equals(fullName, "System.Diagnostics.CodeAnalysis.AllowNullAttribute", StringComparison.Ordinal) + || string.Equals(fullName, "System.Diagnostics.CodeAnalysis.MaybeNullAttribute", StringComparison.Ordinal)) + { + return true; + } + } + return false; + } + private static bool IsDuplicateTypeDescriptorAttribute(Attribute typeDescriptorAttribute, Attribute[] propertyAttributes) { foreach (var propertyAttribute in propertyAttributes) diff --git a/tests/MiniValidation.UnitTests/MiniValidation.UnitTests.csproj b/tests/MiniValidation.UnitTests/MiniValidation.UnitTests.csproj index a640ad8..e07edf8 100644 --- a/tests/MiniValidation.UnitTests/MiniValidation.UnitTests.csproj +++ b/tests/MiniValidation.UnitTests/MiniValidation.UnitTests.csproj @@ -1,8 +1,8 @@ - + net8.0;net9.0;net10.0 - 10.0 + 11.0 enable enable diff --git a/tests/MiniValidation.UnitTests/TryValidate.cs b/tests/MiniValidation.UnitTests/TryValidate.cs index 70a1759..95e72d6 100644 --- a/tests/MiniValidation.UnitTests/TryValidate.cs +++ b/tests/MiniValidation.UnitTests/TryValidate.cs @@ -558,4 +558,108 @@ public AlwaysInvalidAttribute(string id) public override bool IsValid(object? value) => false; } + + [Fact] + public void RequiredMemberAttribute_On_NonNullable_Member_Treated_As_Required() + { + var thingToValidate = new TestTypeWithNonNullableRequiredMember { Name = null! }; + + var result = MiniValidator.TryValidate(thingToValidate, out var errors); + + Assert.False(result); + var entry = Assert.Single(errors); + Assert.Equal(nameof(TestTypeWithNonNullableRequiredMember.Name), entry.Key); + } + + [Fact] + public void RequiredMemberAttribute_On_Nullable_Members_Ignored() + { + var thingToValidate = new TestTypeWithNullableRequiredMembers { Name = null, Count = null }; + + var result = MiniValidator.TryValidate(thingToValidate, out var errors); + + Assert.True(result); + Assert.Empty(errors); + } + + [Fact] + public void Unrelated_RequiredMemberAttribute_Does_Not_Add_Required_Validation() + { + var thingToValidate = new TestTypeWithCustomRequiredMemberAttr { Name = null }; + + var result = MiniValidator.TryValidate(thingToValidate, out var errors); + + Assert.True(result); + Assert.Empty(errors); + } + + [Fact] + public void Required_Value_Types_Do_Not_Trigger_RequiresValidation() + { + Assert.False(MiniValidator.RequiresValidation(typeof(TestTypeWithRequiredValueType))); + Assert.False(MiniValidator.RequiresValidation(typeof(TestTypeWithNullableRequiredMembers))); + } + + [Fact] + public void RequiredMemberAttribute_With_AllowNull_Or_MaybeNull_Ignored() + { + var thingToValidate = new TestTypeWithAllowNullRequiredMember { Value = null! }; + var result = MiniValidator.TryValidate(thingToValidate, out var errors); + Assert.True(result); + Assert.Empty(errors); + + var thingToValidateMaybeNull = new TestTypeWithMaybeNullRequiredMember { Value = null! }; + var resultMaybeNull = MiniValidator.TryValidate(thingToValidateMaybeNull, out errors); + Assert.True(resultMaybeNull); + Assert.Empty(errors); + } + + [Fact] + public void IsReferenceTypeNullableFallback_Matches_Modern_Behavior() + { + var propAllowNull = typeof(TestTypeWithAllowNullRequiredMember).GetProperty(nameof(TestTypeWithAllowNullRequiredMember.Value))!; + var propMaybeNull = typeof(TestTypeWithMaybeNullRequiredMember).GetProperty(nameof(TestTypeWithMaybeNullRequiredMember.Value))!; + var propNonNullable = typeof(TestTypeWithNonNullableRequiredMember).GetProperty(nameof(TestTypeWithNonNullableRequiredMember.Name))!; + + Assert.True(TypeDetailsCache.IsReferenceTypeNullableFallback(propAllowNull)); + Assert.True(TypeDetailsCache.IsReferenceTypeNullableFallback(propMaybeNull)); + Assert.False(TypeDetailsCache.IsReferenceTypeNullableFallback(propNonNullable)); + } + + class TestTypeWithNonNullableRequiredMember + { + public required string Name { get; set; } + } + + class TestTypeWithAllowNullRequiredMember + { + [System.Diagnostics.CodeAnalysis.AllowNull] + public required string Value { get; set; } + } + + class TestTypeWithMaybeNullRequiredMember + { + [System.Diagnostics.CodeAnalysis.MaybeNull] + public required string Value { get; set; } + } + + class TestTypeWithNullableRequiredMembers + { + public required string? Name { get; set; } + public required int? Count { get; set; } + } + + class TestTypeWithRequiredValueType + { + public required int Value { get; set; } + } + + class TestTypeWithCustomRequiredMemberAttr + { + [CustomRequiredMember] + public string? Name { get; set; } + } + + [AttributeUsage(AttributeTargets.Property)] + class CustomRequiredMemberAttribute : Attribute { } }