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
292 changes: 281 additions & 11 deletions Fluid.SourceGenerator/MemberAccessorGenerator.cs

Large diffs are not rendered by default.

4 changes: 3 additions & 1 deletion Fluid.Tests/Fluid.Tests.csproj
Original file line number Diff line number Diff line change
Expand Up @@ -24,7 +24,9 @@
<PackageReference Include="Microsoft.Bcl.AsyncInterfaces" />
<ProjectReference Include="..\Fluid.MvcViewEngine\Fluid.MvcViewEngine.csproj" />
<ProjectReference Include="..\MinimalApis.LiquidViews\MinimalApis.LiquidViews.csproj" />
<ProjectReference Include="..\Fluid.SourceGenerator\Fluid.SourceGenerator.csproj" />
<ProjectReference Include="..\Fluid.SourceGenerator\Fluid.SourceGenerator.csproj"
OutputItemType="Analyzer"
ReferenceOutputAssembly="true" />
<ProjectReference Include="..\Fluid.ViewEngine\Fluid.ViewEngine.csproj" />
<ProjectReference Include="..\Fluid\Fluid.csproj" />
</ItemGroup>
Expand Down
91 changes: 91 additions & 0 deletions Fluid.Tests/MemberAccessStrategyTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -392,6 +392,79 @@ public void ShouldAllowRuntimeRegistrationsToOverrideGeneratedMemberAccessors()

Assert.Equal("runtime", template.Render(new TemplateContext(model, options)));
}

[Fact]
public void ShouldUseGeneratedMemberAccessorInferredFromTemplateContextConstructor()
{
var options = new TemplateOptions();
var model = new InferredModel();
var context = new TemplateContext(model, options);

var accessor = options.MemberAccessStrategy.GetAccessor(
typeof(InferredModel),
nameof(InferredModel.Inferred),
options.ModelNamesComparer);

Assert.Equal("Fluid.SourceGenerated", accessor.GetType().Namespace);
Assert.Equal("inferred", _parser.Parse("{{ Inferred }}").Render(context));
Assert.Equal("", _parser.Parse("{{ NotAProperty }}").Render(context));
Assert.Null(options.MemberAccessStrategy.GetAccessor(
typeof(InferredModel),
nameof(InferredModel.NotAProperty),
options.ModelNamesComparer));
}

[Fact]
public void InferredAccessorShouldPreserveBaseTypeRegistrationFallback()
{
var options = new TemplateOptions();
options.MemberAccessStrategy.Register<InferredModelBase, object>(
static (_, name) => name == "Custom" ? "custom" : null);
var context = new TemplateContext(new InferredModel(), options);

Assert.Equal("custom", _parser.Parse("{{ Custom }}").Render(context));
}

[Fact]
public void ShouldNotActivateInferredAccessorOnTemplateOptionsDefault()
{
var options = TemplateOptions.Default;
var context = new TemplateContext(new DefaultInferredModel(), options);
var accessor = options.MemberAccessStrategy.GetAccessor(
typeof(DefaultInferredModel),
nameof(DefaultInferredModel.Value),
options.ModelNamesComparer);

Assert.Equal("Fluid.Accessors", accessor.GetType().Namespace);
Assert.Equal("default", _parser.Parse("{{ Value }}").Render(context));
}

[Fact]
public void InferredAccessorShouldUseConfiguredValueConverters()
{
var options = new TemplateOptions();
options.ValueConverters.Add(static value => value is int ? "converted" : null);
var context = new TemplateContext(new InferredModel(), options);

Assert.Equal("converted", _parser.Parse("{{ Count }}").Render(context));
}

[Fact]
public void InferredAccessorShouldPreserveNonGenericAsyncMemberBehavior()
{
var options = new TemplateOptions();
_ = new TemplateContext(new InferredModel(), options);

Assert.Equal("Fluid.Accessors", options.MemberAccessStrategy.GetAccessor(
typeof(InferredModel),
nameof(InferredModel.PlainTask),
options.ModelNamesComparer).GetType().Namespace);
Assert.Equal("Fluid.Accessors", options.MemberAccessStrategy.GetAccessor(
typeof(InferredModel),
nameof(InferredModel.ValueTask),
options.ModelNamesComparer).GetType().Namespace);
}

}

public class ModelWithStaticNull
Expand Down Expand Up @@ -424,6 +497,24 @@ public sealed class GeneratedModel
public string Generated { get; set; } = "model";
}

public class InferredModelBase
{
}

public sealed class InferredModel : InferredModelBase
{
public string Inferred { get; set; } = "inferred";
public int Count { get; set; } = 42;
public Task PlainTask { get; set; } = Task.CompletedTask;
public ValueTask<int> ValueTask { get; set; }
public string NotAProperty() => "method";
}

public sealed class DefaultInferredModel
{
public string Value => "default";
}

public sealed class GeneratedTemplateOptions : TemplateOptions, ITemplateOptionsMemberAccessorRegistrar
{
void ITemplateOptionsMemberAccessorRegistrar.RegisterMemberAccessors(TemplateOptions options)
Expand Down
83 changes: 82 additions & 1 deletion Fluid.Tests/MemberAccessorGeneratorTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -109,6 +109,87 @@ public partial class PublicTemplateOptions : TemplateOptions
Assert.Contains("strategy.Register(typeof(global::Address), \"*\", new global::Fluid.SourceGenerated.Address_GeneratedMemberAccessor());", generated);
}

[Fact]
public void ShouldInferModelTypeFromTemplateContextWithCustomOptions()
{
var source = """
using Fluid;

public class Person
{
public string FirstName { get; set; } = "";
public string Hidden { private get; set; } = "";
public string Method() => "";
public System.Threading.Tasks.Task PlainTask { get; set; } = System.Threading.Tasks.Task.CompletedTask;
public System.Threading.Tasks.ValueTask<int> ValueTask { get; set; }
}

public static class ContextFactory
{
public static TemplateContext Create(Person person, TemplateOptions options)
=> new TemplateContext(person, options);
}
""";

var generated = RunGenerator(source);

Assert.Contains("internal sealed class Person_GeneratedMemberAccessor", generated);
Assert.Contains("[global::System.Runtime.CompilerServices.ModuleInitializer]", generated);
Assert.Contains(
"DefaultMemberAccessStrategy.RegisterSourceGeneratedAccessor(typeof(global::Person), new Person_GeneratedMemberAccessor_Inferred0(), new string[] { \"FirstName\" });",
generated);
Assert.DoesNotContain("typed.Hidden", generated);
Assert.Contains("typed.Method()", generated);
Assert.DoesNotContain("new string[] { \"PlainTask\" }", generated);
Assert.DoesNotContain("new string[] { \"ValueTask\" }", generated);
}

[Fact]
public void ShouldNotInferModelTypeFromTemplateContextUsingDefaultOptions()
{
var source = """
using Fluid;

public class Person
{
public string FirstName { get; set; } = "";
}

public static class ContextFactory
{
public static TemplateContext Create(Person person)
=> new TemplateContext(person);
}
""";

var generated = RunGenerator(source);

Assert.DoesNotContain("Person_GeneratedMemberAccessor", generated);
}

[Fact]
public void ShouldNotInferModelTypeFromExplicitTemplateOptionsDefault()
{
var source = """
using Fluid;

public class Person
{
public string FirstName { get; set; } = "";
}

public static class ContextFactory
{
public static TemplateContext Create(Person person)
=> new TemplateContext(person, TemplateOptions.Default);
}
""";

var generated = RunGenerator(source);

Assert.DoesNotContain("Person_GeneratedMemberAccessor", generated);
}

private static string RunGenerator(string source)
{
var syntaxTree = CSharpSyntaxTree.ParseText(source);
Expand All @@ -122,7 +203,7 @@ private static string RunGenerator(string source)
var driver = CSharpGeneratorDriver.Create(generator).RunGenerators(compilation);
var runResult = driver.GetRunResult();

Assert.Equal(2, runResult.GeneratedTrees.Length);
Assert.InRange(runResult.GeneratedTrees.Length, 1, 2);
Assert.Empty(runResult.Diagnostics.Where(static x => x.Severity == DiagnosticSeverity.Error));

var outputCompilation = compilation.AddSyntaxTrees(runResult.GeneratedTrees);
Expand Down
121 changes: 117 additions & 4 deletions Fluid/DefaultMemberAccessStrategy.cs
Original file line number Diff line number Diff line change
@@ -1,3 +1,4 @@
using System.ComponentModel;
using Fluid.Accessors;
using System.Reflection;
using System.Runtime.CompilerServices;
Expand All @@ -20,10 +21,26 @@ public ReflectionCache(object registrationToken)
public volatile Dictionary<ReflectedAccessorKey, MemberAccessor> Accessors = [];
}

private sealed class AccessorCacheState
{
public AccessorCacheState(object registrations, object generatedRegistrations)
{
Registrations = registrations;
GeneratedRegistrations = generatedRegistrations;
}

public object Registrations { get; }
public object GeneratedRegistrations { get; }
}

private static readonly bool _dynamicCodeSupported = IsDynamicCodeSupported();

private volatile Dictionary<AccessorKey, MemberAccessor> _registrations = [];
private volatile Dictionary<Type, GeneratedMemberAccessorRegistration[]> _generatedRegistrations = [];
private volatile object _generatedRegistryToken = GeneratedMemberAccessorRegistry.CacheToken;
private volatile Type _lastGeneratedType;
private volatile ReflectionCache _reflectionCache;
private volatile AccessorCacheState _accessorCacheState;

// Only the exact type opts in. A derived strategy may override GetAccessor to resolve from its
// own source, which these maps -- and therefore the token -- would not reflect; it would then serve
Expand All @@ -33,10 +50,19 @@ public ReflectionCache(object registrationToken)
public DefaultMemberAccessStrategy()
{
_accessorCachingSupported = GetType() == typeof(DefaultMemberAccessStrategy);
_reflectionCache = new ReflectionCache(_registrations);
_accessorCacheState = new AccessorCacheState(_registrations, _generatedRegistrations);
_reflectionCache = new ReflectionCache(_accessorCacheState);
}

protected internal override object AccessorCacheToken => _accessorCachingSupported ? _registrations : null;
protected internal override object AccessorCacheToken
=> _accessorCachingSupported ? GetAccessorCacheState(_registrations, _generatedRegistrations) : null;

/// <summary>
/// Registers an accessor emitted by the Fluid source generator.
/// </summary>
[EditorBrowsable(EditorBrowsableState.Never)]
public static void RegisterSourceGeneratedAccessor(Type type, MemberAccessor accessor, params string[] memberNames)
=> GeneratedMemberAccessorRegistry.Register(type, accessor, memberNames);

public override MemberAccessor GetAccessor(Type type, string name, StringComparer stringComparer)
{
Expand All @@ -45,13 +71,25 @@ public override MemberAccessor GetAccessor(Type type, string name, StringCompare
ArgumentNullException.ThrowIfNull(stringComparer);

var registrations = _registrations;
var generatedRegistrations = _generatedRegistrations;

if (TryGetRegisteredAccessor(registrations, type, name, out var accessor))
{
return accessor;
}

var reflectionCache = GetReflectionCache(registrations);
if (generatedRegistrations.TryGetValue(type, out var generatedAccessors))
{
foreach (var generatedAccessor in generatedAccessors)
{
if (generatedAccessor.CanAccess(name, stringComparer))
{
return generatedAccessor.Accessor;
}
}
}

var reflectionCache = GetReflectionCache(GetAccessorCacheState(registrations, generatedRegistrations));
var key = new ReflectedAccessorKey(type, name, stringComparer);
var reflectedAccessors = reflectionCache.Accessors;

Expand Down Expand Up @@ -231,7 +269,60 @@ [new AccessorKey(type, name)] = accessor
Interlocked.CompareExchange(ref _registrations, updated, registrations),
registrations))
{
_reflectionCache = new ReflectionCache(updated);
_reflectionCache = new ReflectionCache(GetAccessorCacheState(updated, _generatedRegistrations));
return;
}
}
}

internal override void RegisterGeneratedAccessor(Type type)
{
while (true)
{
var registryToken = GeneratedMemberAccessorRegistry.CacheToken;
var generatedRegistrations = _generatedRegistrations;

if (ReferenceEquals(_generatedRegistryToken, registryToken))
{
if (ReferenceEquals(_lastGeneratedType, type))
{
return;
}

if (generatedRegistrations.ContainsKey(type))
{
_lastGeneratedType = type;
return;
}
}

var accessors = GeneratedMemberAccessorRegistry.GetAccessors(type, out registryToken);
if (accessors is null)
{
_generatedRegistryToken = registryToken;
return;
}

if (generatedRegistrations.TryGetValue(type, out var registeredAccessors) &&
ReferenceEquals(registeredAccessors, accessors))
{
_generatedRegistryToken = registryToken;
_lastGeneratedType = type;
return;
}

var updated = new Dictionary<Type, GeneratedMemberAccessorRegistration[]>(generatedRegistrations)
{
[type] = accessors
};

if (ReferenceEquals(
Interlocked.CompareExchange(ref _generatedRegistrations, updated, generatedRegistrations),
generatedRegistrations))
{
_generatedRegistryToken = registryToken;
_lastGeneratedType = type;
_reflectionCache = new ReflectionCache(GetAccessorCacheState(_registrations, updated));
return;
}
}
Expand All @@ -257,6 +348,28 @@ private ReflectionCache GetReflectionCache(object registrationToken)
}
}

private AccessorCacheState GetAccessorCacheState(object registrations, object generatedRegistrations)
{
while (true)
{
var state = _accessorCacheState;

if (ReferenceEquals(state.Registrations, registrations) &&
ReferenceEquals(state.GeneratedRegistrations, generatedRegistrations))
{
return state;
}

var updated = new AccessorCacheState(registrations, generatedRegistrations);
if (ReferenceEquals(
Interlocked.CompareExchange(ref _accessorCacheState, updated, state),
state))
{
return updated;
}
}
}

private static void AddReflectedAccessor(
ReflectionCache reflectionCache,
ReflectedAccessorKey key,
Expand Down
Loading