Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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
26 changes: 26 additions & 0 deletions TUnit.Mocks.SourceGenerator.Tests/MockGeneratorTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -242,6 +242,32 @@ void M()
return VerifyGeneratorOutput(source);
}

[Test]
public Task Interface_With_RefStruct_Parameters()
{
var source = """
using System;
using TUnit.Mocks;

public interface IBufferProcessor
{
void Process(ReadOnlySpan<byte> data);
int Parse(ReadOnlySpan<char> text);
string GetName();
}

public class TestUsage
{
void M()
{
var mock = Mock.Of<IBufferProcessor>();
}
}
""";

return VerifyGeneratorOutput(source);
}

[Test]
public Task Interface_With_Mixed_Members()
{
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,93 @@
// <auto-generated/>
#nullable enable

namespace TUnit.Mocks.Generated
{
internal static class IBufferProcessor_MockFactory
{
[global::System.Runtime.CompilerServices.ModuleInitializer]
internal static void Register()
{
global::TUnit.Mocks.Mock.RegisterFactory<global::IBufferProcessor>(Create);
}

private static global::TUnit.Mocks.Mock<global::IBufferProcessor> Create(global::TUnit.Mocks.MockBehavior behavior)
{
var engine = new global::TUnit.Mocks.MockEngine<global::IBufferProcessor>(behavior);
var impl = new IBufferProcessor_MockImpl(engine);
engine.Raisable = impl;
var mock = new global::TUnit.Mocks.Mock<global::IBufferProcessor>(impl, engine);
return mock;
}
}
}


// ===== FILE SEPARATOR =====

// <auto-generated/>
#nullable enable

namespace TUnit.Mocks.Generated
{
internal sealed class IBufferProcessor_MockImpl : global::IBufferProcessor, global::TUnit.Mocks.IRaisable
{
private readonly global::TUnit.Mocks.MockEngine<global::IBufferProcessor> _engine;

internal IBufferProcessor_MockImpl(global::TUnit.Mocks.MockEngine<global::IBufferProcessor> engine)
{
_engine = engine;
}

public void Process(global::System.ReadOnlySpan<byte> data)
{
_engine.HandleCall(0, "Process", global::System.Array.Empty<object?>());
}

public int Parse(global::System.ReadOnlySpan<char> text)
{
return _engine.HandleCallWithReturn<int>(1, "Parse", global::System.Array.Empty<object?>(), default);
}

public string GetName()
{
return _engine.HandleCallWithReturn<string>(2, "GetName", global::System.Array.Empty<object?>(), "");
}

[global::System.ComponentModel.EditorBrowsable(global::System.ComponentModel.EditorBrowsableState.Never)]
public void RaiseEvent(string eventName, object? args)
{
throw new global::System.InvalidOperationException($"No event named '{eventName}' exists on this mock.");
}
}
}


// ===== FILE SEPARATOR =====

// <auto-generated/>
#nullable enable

namespace TUnit.Mocks.Generated
{
public static class IBufferProcessor_MockMemberExtensions
{
public static global::TUnit.Mocks.VoidMockMethodCall Process(this global::TUnit.Mocks.Mock<global::IBufferProcessor> mock)
{
var matchers = global::System.Array.Empty<global::TUnit.Mocks.Arguments.IArgumentMatcher>();
return new global::TUnit.Mocks.VoidMockMethodCall(mock.Engine, 0, "Process", matchers);
}

public static global::TUnit.Mocks.MockMethodCall<int> Parse(this global::TUnit.Mocks.Mock<global::IBufferProcessor> mock)
{
var matchers = global::System.Array.Empty<global::TUnit.Mocks.Arguments.IArgumentMatcher>();
return new global::TUnit.Mocks.MockMethodCall<int>(mock.Engine, 1, "Parse", matchers);
}

public static global::TUnit.Mocks.MockMethodCall<string> GetName(this global::TUnit.Mocks.Mock<global::IBufferProcessor> mock)
{
var matchers = global::System.Array.Empty<global::TUnit.Mocks.Arguments.IArgumentMatcher>();
return new global::TUnit.Mocks.MockMethodCall<string>(mock.Engine, 2, "GetName", matchers);
}
}
}
62 changes: 58 additions & 4 deletions TUnit.Mocks.SourceGenerator/Builders/MockImplBuilder.cs
Original file line number Diff line number Diff line change
Expand Up @@ -238,6 +238,17 @@ private static void GenerateWrapMethodBody(CodeWriter writer, MockMemberModel me
writer.AppendLine("}");
writer.AppendLine($"return _wrappedInstance.{method.Name}({argPassList});");
}
else if (method.IsRefStructReturn)
{
writer.AppendLine($"if (_engine.TryHandleCall({method.MemberId}, \"{method.Name}\", {argsArray}))");
writer.AppendLine("{");
writer.IncreaseIndent();
EmitOutRefReadback(writer, method);
writer.AppendLine("return default;");
writer.DecreaseIndent();
writer.AppendLine("}");
writer.AppendLine($"return _wrappedInstance.{method.Name}({argPassList});");
}
else
{
writer.AppendLine($"if (_engine.TryHandleCallWithReturn<{method.ReturnType}>({method.MemberId}, \"{method.Name}\", {argsArray}, {method.SmartDefault}, out var __result))");
Expand Down Expand Up @@ -475,6 +486,18 @@ private static void GeneratePartialMethodBody(CodeWriter writer, MockMemberModel
writer.AppendLine("}");
writer.AppendLine($"return base.{method.Name}({argPassList});");
}
else if (method.IsRefStructReturn)
{
// synchronous method returning ref struct — use void dispatch, fall back to base
writer.AppendLine($"if (_engine.TryHandleCall({method.MemberId}, \"{method.Name}\", {argsArray}))");
writer.AppendLine("{");
writer.IncreaseIndent();
EmitOutRefReadback(writer, method);
writer.AppendLine("return default;");
writer.DecreaseIndent();
writer.AppendLine("}");
writer.AppendLine($"return base.{method.Name}({argPassList});");
}
else
{
// synchronous method with return value
Expand Down Expand Up @@ -566,6 +589,15 @@ private static void GenerateEngineDispatchBody(CodeWriter writer, MockMemberMode
}
}
}
else if (method.IsRefStructReturn)
{
// Synchronous method returning a ref struct — can't use HandleCallWithReturn<T> because
// ref structs can't be generic type arguments. Use void dispatch for call tracking,
// callbacks, and throws. Return default (e.g. ReadOnlySpan<byte>.Empty).
writer.AppendLine($"_engine.HandleCall({method.MemberId}, \"{method.Name}\", {argsArray});");
EmitOutRefReadback(writer, method);
writer.AppendLine("return default;");
}
else
{
// Synchronous method with return value — need to read back out/ref before returning
Expand All @@ -589,12 +621,32 @@ private static void GenerateInterfaceProperty(CodeWriter writer, MockMemberModel

if (prop.HasGetter)
{
writer.AppendLine($"get => _engine.HandleCallWithReturn<{prop.ReturnType}>({prop.MemberId}, \"get_{prop.Name}\", global::System.Array.Empty<object?>(), {prop.SmartDefault});");
if (prop.IsRefStructReturn)
{
// ref struct property — can't use HandleCallWithReturn<T>, use void dispatch + return default
writer.AppendLine("get");
writer.OpenBrace();
writer.AppendLine($"_engine.HandleCall({prop.MemberId}, \"get_{prop.Name}\", global::System.Array.Empty<object?>());");
writer.AppendLine("return default;");
writer.CloseBrace();
}
else
{
writer.AppendLine($"get => _engine.HandleCallWithReturn<{prop.ReturnType}>({prop.MemberId}, \"get_{prop.Name}\", global::System.Array.Empty<object?>(), {prop.SmartDefault});");
}
}

if (prop.HasSetter)
{
writer.AppendLine($"set => _engine.HandleCall({prop.SetterMemberId}, \"set_{prop.Name}\", new object?[] {{ value }});");
if (prop.IsRefStructReturn)
{
// ref struct property — can't box value, use empty args
writer.AppendLine($"set => _engine.HandleCall({prop.SetterMemberId}, \"set_{prop.Name}\", global::System.Array.Empty<object?>());");
}
else
{
writer.AppendLine($"set => _engine.HandleCall({prop.SetterMemberId}, \"set_{prop.Name}\", new object?[] {{ value }});");
}
}

writer.CloseBrace();
Expand Down Expand Up @@ -838,6 +890,7 @@ private static void EmitOutRefReadback(CodeWriter writer, MockMemberModel method
for (int i = 0; i < method.Parameters.Length; i++)
{
var p = method.Parameters[i];
if (p.IsRefStruct) continue; // ref structs can't be cast from object
if (p.Direction == ParameterDirection.Out || p.Direction == ParameterDirection.Ref)
{
writer.AppendLine($"if (__outRef.TryGetValue({i}, out var __v{i})) {p.Name} = ({p.FullyQualifiedType})__v{i}!;");
Expand All @@ -848,8 +901,9 @@ private static void EmitOutRefReadback(CodeWriter writer, MockMemberModel method

private static string GetArgsArrayExpression(MockMemberModel method)
{
// Only include non-out parameters in args array
var matchableParams = method.Parameters.Where(p => p.Direction != ParameterDirection.Out).ToList();
// Only include non-out, non-ref-struct parameters in args array
// (ref structs cannot be boxed into object?[])
var matchableParams = method.Parameters.Where(p => p.Direction != ParameterDirection.Out && !p.IsRefStruct).ToList();
if (matchableParams.Count == 0) return "global::System.Array.Empty<object?>()";
var args = string.Join(", ", matchableParams.Select(p => p.Name));
return $"new object?[] {{ {args} }}";
Expand Down
41 changes: 24 additions & 17 deletions TUnit.Mocks.SourceGenerator/Builders/MockMembersBuilder.cs
Original file line number Diff line number Diff line change
Expand Up @@ -48,8 +48,9 @@ public static string Build(MockTypeModel model)
}

// Properties -- extension properties via C# 14 extension blocks
// (skip ref struct properties — can't use PropertyMockCall<RefStruct>)
var memberProps = model.Properties
.Where(p => !p.IsIndexer && (p.HasGetter || p.HasSetter))
.Where(p => !p.IsIndexer && !p.IsRefStructReturn && (p.HasGetter || p.HasSetter))
.ToList();
if (memberProps.Count > 0)
{
Expand Down Expand Up @@ -83,13 +84,15 @@ private static bool ShouldGenerateTypedWrapper(MockMemberModel method, bool hasE
{
if (method.IsGenericMethod) return false;

var nonOutParams = method.Parameters.Where(p => p.Direction != ParameterDirection.Out).ToList();
if (nonOutParams.Count == 0)
// Exclude out params and ref struct params (can't be boxed or used as type args)
var matchableParams = method.Parameters.Where(p => p.Direction != ParameterDirection.Out && !p.IsRefStruct).ToList();
if (matchableParams.Count == 0)
{
var hasOutRefParams = method.Parameters.Any(p => p.Direction == ParameterDirection.Out || p.Direction == ParameterDirection.Ref);
var hasOutRefParams = method.Parameters.Any(p =>
!p.IsRefStruct && (p.Direction == ParameterDirection.Out || p.Direction == ParameterDirection.Ref));
return hasEvents || hasOutRefParams;
}
return nonOutParams.Count <= MaxTypedParams;
return matchableParams.Count <= MaxTypedParams;
}

private static string GetWrapperName(string safeName, MockMemberModel method)
Expand All @@ -106,15 +109,16 @@ private static void GenerateUnifiedSealedClass(CodeWriter writer, MockMemberMode
: method.ReturnType;

var wrapperName = GetWrapperName(safeName, method);
var nonOutParams = method.Parameters.Where(p => p.Direction != ParameterDirection.Out).ToList();
var matchableParams = method.Parameters.Where(p => p.Direction != ParameterDirection.Out && !p.IsRefStruct).ToList();

if (method.IsVoid)
// Ref struct returns use the void wrapper (can't use generic type args with ref structs)
if (method.IsVoid || method.IsRefStructReturn)
{
GenerateVoidUnifiedClass(writer, wrapperName, nonOutParams, events, method.Parameters);
GenerateVoidUnifiedClass(writer, wrapperName, matchableParams, events, method.Parameters);
}
else
{
GenerateReturnUnifiedClass(writer, wrapperName, nonOutParams, setupReturnType, events, method.Parameters);
GenerateReturnUnifiedClass(writer, wrapperName, matchableParams, setupReturnType, events, method.Parameters);
}
}

Expand Down Expand Up @@ -460,8 +464,9 @@ private static void GenerateMemberMethod(CodeWriter writer, MockMemberModel meth
{
returnType = GetWrapperName(safeName, method);
}
else if (method.IsVoid)
else if (method.IsVoid || method.IsRefStructReturn)
{
// Ref struct returns use VoidMockMethodCall (can't use ref struct as generic type arg)
returnType = "global::TUnit.Mocks.VoidMockMethodCall";
}
else
Expand All @@ -479,16 +484,17 @@ private static void GenerateMemberMethod(CodeWriter writer, MockMemberModel meth

using (writer.Block($"public static {returnType} {safeMemberName}{typeParams}({fullParamList}){constraints}"))
{
// Build matchers array
var nonOutParams = method.Parameters.Where(p => p.Direction != ParameterDirection.Out).ToList();
// Build matchers array (exclude out and ref struct params)
var matchableParams = method.Parameters
.Where(p => p.Direction != ParameterDirection.Out && !p.IsRefStruct).ToList();

if (nonOutParams.Count == 0)
if (matchableParams.Count == 0)
{
writer.AppendLine("var matchers = global::System.Array.Empty<global::TUnit.Mocks.Arguments.IArgumentMatcher>();");
}
else
{
var matcherArgs = string.Join(", ", nonOutParams.Select(p => $"{p.Name}.Matcher"));
var matcherArgs = string.Join(", ", matchableParams.Select(p => $"{p.Name}.Matcher"));
writer.AppendLine($"var matchers = new global::TUnit.Mocks.Arguments.IArgumentMatcher[] {{ {matcherArgs} }};");
}

Expand All @@ -497,7 +503,7 @@ private static void GenerateMemberMethod(CodeWriter writer, MockMemberModel meth
var wrapperName = GetWrapperName(safeName, method);
writer.AppendLine($"return new {wrapperName}(mock.Engine, {method.MemberId}, \"{method.Name}\", matchers);");
}
else if (method.IsVoid)
else if (method.IsVoid || method.IsRefStructReturn)
{
writer.AppendLine($"return new global::TUnit.Mocks.VoidMockMethodCall(mock.Engine, {method.MemberId}, \"{method.Name}\", matchers);");
}
Expand Down Expand Up @@ -572,9 +578,10 @@ private static void GenerateRaiseExtensionMethods(CodeWriter writer, MockTypeMod

private static string GetArgParameterList(MockMemberModel method)
{
// Only include non-out parameters as Arg<T> in setup
// Only include non-out, non-ref-struct parameters as Arg<T> in setup
// (ref structs cannot be used as generic type arguments)
return string.Join(", ", method.Parameters
.Where(p => p.Direction != ParameterDirection.Out)
.Where(p => p.Direction != ParameterDirection.Out && !p.IsRefStruct)
.Select(p => $"global::TUnit.Mocks.Arguments.Arg<{p.FullyQualifiedType}> {p.Name}"));
}

Expand Down
9 changes: 6 additions & 3 deletions TUnit.Mocks.SourceGenerator/Discovery/MemberDiscovery.cs
Original file line number Diff line number Diff line change
Expand Up @@ -299,7 +299,8 @@ private static MockMemberModel CreateMethodModel(IMethodSymbol method, ref int m
FullyQualifiedType = p.Type.GetFullyQualifiedName(),
Direction = p.GetParameterDirection(),
HasDefaultValue = p.HasExplicitDefaultValue,
DefaultValueExpression = p.HasExplicitDefaultValue ? FormatDefaultValue(p) : null
DefaultValueExpression = p.HasExplicitDefaultValue ? FormatDefaultValue(p) : null,
IsRefStruct = p.Type.IsRefLikeType
}).ToImmutableArray()
),
TypeParameters = new EquatableArray<MockTypeParameterModel>(
Expand All @@ -316,7 +317,8 @@ private static MockMemberModel CreateMethodModel(IMethodSymbol method, ref int m
IsAbstractMember = method.IsAbstract,
IsVirtualMember = method.IsVirtual || method.IsOverride,
IsProtected = method.DeclaredAccessibility == Accessibility.Protected
|| method.DeclaredAccessibility == Accessibility.ProtectedOrInternal
|| method.DeclaredAccessibility == Accessibility.ProtectedOrInternal,
IsRefStructReturn = returnType.IsRefLikeType
};
}

Expand Down Expand Up @@ -365,7 +367,8 @@ private static MockMemberModel CreatePropertyModel(IPropertySymbol property, ref
IsAbstractMember = property.IsAbstract,
IsVirtualMember = property.IsVirtual || property.IsOverride,
IsProtected = property.DeclaredAccessibility == Accessibility.Protected
|| property.DeclaredAccessibility == Accessibility.ProtectedOrInternal
|| property.DeclaredAccessibility == Accessibility.ProtectedOrInternal,
IsRefStructReturn = property.Type.IsRefLikeType
};
}

Expand Down
4 changes: 3 additions & 1 deletion TUnit.Mocks.SourceGenerator/Models/MockMemberModel.cs
Original file line number Diff line number Diff line change
Expand Up @@ -29,6 +29,7 @@ internal sealed record MockMemberModel : IEquatable<MockMemberModel>
public bool IsAbstractMember { get; init; }
public bool IsVirtualMember { get; init; }
public bool IsProtected { get; init; }
public bool IsRefStructReturn { get; init; }

public bool Equals(MockMemberModel? other)
{
Expand All @@ -54,7 +55,8 @@ public bool Equals(MockMemberModel? other)
&& UnwrappedSmartDefault == other.UnwrappedSmartDefault
&& IsAbstractMember == other.IsAbstractMember
&& IsVirtualMember == other.IsVirtualMember
&& IsProtected == other.IsProtected;
&& IsProtected == other.IsProtected
&& IsRefStructReturn == other.IsRefStructReturn;
}

public override int GetHashCode()
Expand Down
Loading
Loading