Skip to content
Draft
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
Original file line number Diff line number Diff line change
Expand Up @@ -7,9 +7,12 @@
using System.Linq;
using System.Threading.Tasks;
using Microsoft.CodeAnalysis;
using Microsoft.TypeSpec.Generator.ClientModel.Providers;
using Microsoft.TypeSpec.Generator.Expressions;
using Microsoft.TypeSpec.Generator.Input;
using Microsoft.TypeSpec.Generator.Primitives;
using Microsoft.TypeSpec.Generator.Providers;
using Microsoft.TypeSpec.Generator.Statements;
using Microsoft.TypeSpec.Generator.Tests.Common;
using NUnit.Framework;

Expand Down Expand Up @@ -73,6 +76,56 @@ public async Task OperationResponseBodyModelRemainsPublicAsRootOutputModel()
await GenerateAndAssertPublicModels([responseModel], [client], ["ResponseBody"]);
}

[TestCase("ErrorResponse", true, false)]
[TestCase("ErrorResponse", true, true)]
[TestCase("ServiceErrorResponse", true, false)]
[TestCase("ServiceErrorResponse", true, true)]
[TestCase("ErrorResult", false, false)]
[TestCase("ErrorResult", false, true)]
public async Task ErrorResultHelperDoesNotKeepServiceModel(string modelName, bool isError, bool referenced)
{
var model = InputFactory.Model(modelName, @namespace: "Sample", access: null!,
usage: InputModelTypeUsage.Json | (isError ? InputModelTypeUsage.Error : InputModelTypeUsage.None));
var responseModel = InputFactory.Model("ResponseBody",
properties: referenced ? [InputFactory.Property("Details", model)] : []);
var operation = InputFactory.Operation("Get", responses: [InputFactory.OperationResponse(bodytype: responseModel)]);
var method = InputFactory.BasicServiceMethod("Get", operation,
response: InputFactory.ServiceMethodResponse(responseModel, []));
var client = InputFactory.Client("TestClient", methods: [method]);

await GenerateAndAssertFiles(
enums: [],
models: [responseModel, model],
clients: [client],
customFiles: [],
expectedFiles: [Path.Combine("src", "Generated", "Internal", "ErrorResult.cs")],
publicModelNames: referenced ? ["ResponseBody", modelName] : ["ResponseBody"],
assertProviders: (session, providers) =>
{
var modelProvider = CodeModelGenerator.Instance.TypeFactory.CreateModel(model)!;
Assert.AreEqual(modelName, modelProvider.Name);
Assert.AreEqual(referenced, session.ShouldWriteProvider(modelProvider));
foreach (var serialization in modelProvider.SerializationProviders)
{
Assert.AreEqual(referenced, session.ShouldWriteProvider(serialization));
}

var helper = providers.OfType<ErrorResultDefinition>().Single();
Assert.AreEqual(modelProvider.Type.Namespace, helper.Type.Namespace);
Assert.AreEqual(1, helper.Type.Arguments.Count);
Assert.IsTrue(helper.DeclarationModifiers.HasFlag(TypeSignatureModifiers.Internal));
Assert.IsFalse(helper.DeclarationModifiers.HasFlag(TypeSignatureModifiers.Public));

var factory = providers.OfType<ModelFactoryProvider>().Single();
Assert.AreEqual(referenced, factory.Methods.Any(m => m.Signature.ReturnType?.Equals(modelProvider.Type) == true));

var context = providers.OfType<ModelReaderWriterContextDefinition>().Single();
var buildableTypes = context.GetAttributesForWrite().OfType<AttributeStatement>()
.SelectMany(a => a.Arguments).OfType<TypeOfExpression>().Select(e => e.Type);
Assert.AreEqual(referenced, buildableTypes.Contains(modelProvider.Type));
});
}

[Test]
public async Task InternalModelReferencedByPublicModelPropertyIsPublicized()
{
Expand Down Expand Up @@ -834,7 +887,8 @@ private static async Task GenerateAndAssertFiles(
string[] internalModelNames = null!,
string[] internalClientNames = null!,
string packageName = "Sample",
Action? configureGenerator = null)
Action? configureGenerator = null,
Action<ProviderReferenceMapSession, IReadOnlyList<TypeProvider>>? assertProviders = null)
{
publicModelNames ??= [];
internalModelNames ??= [];
Expand Down Expand Up @@ -930,6 +984,7 @@ await MockHelpers.LoadMockGeneratorAsync(
{
AssertProviderWritten(session, allProviders, unexpectedFile, expected: false);
}
assertProviders?.Invoke(session, providers);
}

private static IEnumerable<TypeProvider> EnumerateAllProviders(IEnumerable<TypeProvider> providers)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -229,7 +229,7 @@ private static HashSet<string> MaterializeKeepSet(KeptTypesInfo info)
var result = new HashSet<string>(info.TypeNames);
foreach (var provider in info.TypeProviders)
{
result.Add(provider.Type.FullyQualifiedName);
Comment thread
jorgerangel-msft marked this conversation as resolved.
result.Add(ProviderReferenceMapAnalyzer.GetProviderTypeName(provider.Type));
}
return result;
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -321,7 +321,8 @@ protected override string BuildName()
private protected override string NormalizeTypeName(string name)
{
var normalizedName = base.NormalizeTypeName(name);
if (!normalizedName.EndsWith(ResponseSuffix, StringComparison.Ordinal))
if (_inputModel.Usage.HasFlag(InputModelTypeUsage.Error) ||
!normalizedName.EndsWith(ResponseSuffix, StringComparison.Ordinal))
{
return normalizedName;
}
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -598,7 +598,7 @@ private static string GetSimpleName(string fullyQualifiedName)
return lastDot < 0 ? null : fullyQualifiedName.Substring(0, lastDot);
}

private static string GetProviderTypeName(CSharpType type)
internal static string GetProviderTypeName(CSharpType type)
{
var name = type.Arguments.Count > 0 && !type.Name.Contains('`', StringComparison.Ordinal)
? $"{type.Name}`{type.Arguments.Count}"
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -154,6 +154,22 @@ public void TestBuildName_ResponseSuffix(string inputName, bool isExactName, str
Assert.AreEqual("ServiceResponse", model.Properties[0].Name);
}

[TestCase("ErrorResponse", false, "ErrorResponse")]
[TestCase("ServiceErrorResponse", false, "ServiceErrorResponse")]
[TestCase("IpResponse", false, "IPResponse")]
[TestCase("IpResponse", true, "IpResponse")]
public void TestBuildName_ErrorResponseSuffix(string inputName, bool isExactName, string expectedName)
{
var inputModel = InputFactory.Model(
inputName,
usage: InputModelTypeUsage.Error | InputModelTypeUsage.Output | InputModelTypeUsage.Json,
isExactName: isExactName);
var model = new ModelProvider(inputModel);

Assert.AreEqual(expectedName, model.Name);
Assert.AreEqual($"{expectedName}.cs", Path.GetFileName(model.RelativeFilePath));
}

[TestCase("WidgetResponse", "WidgetResponse", false, false)]
[TestCase("WidgetResponse", "WidgetResponse", true, false)]
[TestCase("WidgetResponse", "WidgetResponse", false, true)]
Expand Down Expand Up @@ -188,12 +204,17 @@ await MockHelpers.LoadMockGeneratorAsync(
Assert.IsNotNull(lastContract ? model.LastContractView : model.CustomCodeView);
}

[TestCase("WidgetResponse", "CustomizedWidget")]
[TestCase("IpResponse", "CustomizedIP")]
[TestCase("GadgetResponse", "CustomizedGadget")]
public async Task TestBuildName_ResponseSuffixPreservesCustomName(string inputName, string expectedName)
[TestCase("WidgetResponse", "CustomizedWidget", false)]
[TestCase("WidgetResponse", "CustomizedWidget", true)]
[TestCase("IpResponse", "CustomizedIP", false)]
[TestCase("IpResponse", "CustomizedIP", true)]
[TestCase("GadgetResponse", "CustomizedGadget", false)]
public async Task TestBuildName_ResponseSuffixPreservesCustomName(
string inputName, string expectedName, bool isError)
{
var inputModel = InputFactory.Model(inputName);
var inputModel = InputFactory.Model(inputName,
usage: InputModelTypeUsage.Output | InputModelTypeUsage.Json |
(isError ? InputModelTypeUsage.Error : InputModelTypeUsage.None));
await MockHelpers.LoadMockGeneratorAsync(
inputModelTypes: [inputModel],
compilation: async () => await Helpers.GetCompilationFromDirectoryAsync());
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -64,6 +64,47 @@ public void NonRootKeptTypesAreWrittenWithoutRootingOtherTypes()
Assert.IsFalse(ProviderReferenceMapAnalyzer.ShouldWriteProvider(unusedModel));
}

[Test]
public void KeptProvidersDistinguishGenericArity(
[Values(false, true)] bool isRoot,
[Values(0, 1, 2)] int keptArity,
[Values("Sample", "Sample.Nested.Models")] string ns)
{
var argument = CreateNamedType("T", string.Empty);
TypeProvider[] providers =
[
new TestTypeProvider("ErrorResult", TypeSignatureModifiers.Public, ns: ns),
new GenericTestTypeProvider("ErrorResult", TypeSignatureModifiers.Internal, ns, argument),
new GenericTestTypeProvider("ErrorResult", TypeSignatureModifiers.Internal, ns, argument, CreateNamedType("U", string.Empty))
];
MockHelpers.LoadMockGenerator(createOutputLibrary: () => new TestOutputLibrary(providers));
CodeModelGenerator.Instance.AddTypeToKeep(providers[keptArity], isRoot);

var expectedName = keptArity == 0 ? $"{ns}.ErrorResult" : $"{ns}.ErrorResult`{keptArity}";
var keepSet = isRoot ? CodeModelGenerator.Instance.AdditionalRootTypes : CodeModelGenerator.Instance.NonRootTypes;
Assert.That(keepSet, Contains.Item(expectedName));
if (keptArity == 0)
{
Assert.AreEqual(providers[keptArity].Type.FullyQualifiedName, expectedName);
}
else
{
Assert.That(keepSet, Does.Not.Contain(providers[keptArity].Type.FullyQualifiedName));
}

using var session = ProviderReferenceMapAnalyzer.PrepareForGeneration(providers);

for (var i = 0; i < providers.Length; i++)
{
Assert.AreEqual(i == keptArity, session.ShouldWriteProvider(providers[i]), $"Arity {i}");
}
if (!isRoot && keptArity > 0)
{
Assert.IsTrue(providers[keptArity].DeclarationModifiers.HasFlag(TypeSignatureModifiers.Internal));
Assert.IsFalse(providers[keptArity].DeclarationModifiers.HasFlag(TypeSignatureModifiers.Public));
}
}

[Test]
public void ProviderNamedClientProviderIsNotTreatedAsClientWithoutCapability()
{
Expand Down
Loading