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
Original file line number Diff line number Diff line change
Expand Up @@ -797,20 +797,22 @@ private static MethodBodyStatement BuildDiscriminatedModelsCondition(

private SwitchCaseStatement[] GetDiscriminatorSwitchCases(ModelProvider unknownVariant)
{
SwitchCaseStatement[] cases = new SwitchCaseStatement[_model.DerivedModels.Count - 1];
int index = 0;
for (int i = 0; i < cases.Length; i++)
{
var model = _model.DerivedModels[i];
if (ReferenceEquals(model, unknownVariant))
// Enumerate every derived model rather than the first (Count - 1) entries. The unknown
// variant is not guaranteed to be last - rebasing a hierarchy (for example via
// hierarchyBuilding) can append derived models after it - which would otherwise leave
// unassigned entries in the array.
List<SwitchCaseStatement> cases = new(_model.DerivedModels.Count);
foreach (var model in _model.DerivedModels)
{
if (ReferenceEquals(model, unknownVariant) || model.DiscriminatorValue is null)
{
continue;
}
cases[index++] = new SwitchCaseStatement(
Literal(model.DiscriminatorValue!),
Return(GetDeserializationMethodInvocationForType(model, _jsonElementParameterSnippet, _dataParameter, _serializationOptionsParameter)));
cases.Add(new SwitchCaseStatement(
Literal(model.DiscriminatorValue),
Return(GetDeserializationMethodInvocationForType(model, _jsonElementParameterSnippet, _dataParameter, _serializationOptionsParameter))));
}
return cases;
return [.. cases];
}

/// <summary>
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -73,6 +73,49 @@ public void BaseSerializationContainsSwitchStatement()
Assert.AreEqual(2, switchStatement!.Cases.Count);
}

[Test]
public void RebasedDerivedModelDoesNotProduceNullSwitchCase()
{
// A model rebased onto this hierarchy (for example via hierarchyBuilding) is appended to
// DerivedModels after the generated unknown variant, so the unknown variant is no longer
// last. The discriminator switch cases must still be fully populated.
var rebasedModel = InputFactory.Model(
"voiceItem",
discriminatedKind: "voice",
baseModel: _baseModel,
properties:
[
InputFactory.Property("kind", InputPrimitiveType.String, isRequired: true, isDiscriminator: true)
]);

MockHelpers.LoadMockGenerator(inputModels: () => [_baseModel, _catModel, _dogModel, rebasedModel]);
var baseModel = ScmCodeModelGenerator.Instance.TypeFactory.CreateModel(_baseModel);
Assert.IsNotNull(baseModel);

var unknownVariantIndex = baseModel!.DerivedModels
.Select((m, i) => (m, i))
.First(t => t.m.IsUnknownDiscriminatorModel).i;
Assert.AreNotEqual(
baseModel.DerivedModels.Count - 1,
unknownVariantIndex,
"The unknown variant should not be last, otherwise this scenario is not covered.");

var serialization = baseModel.SerializationProviders.First();
var deserializeMethod = serialization.Methods.First(m => m.Signature.Name == "DeserializePet");
var statements = (MethodBodyStatements)deserializeMethod.BodyStatements!;
var ifStatement = (IfStatement)statements.Statements[1];
var switchStatement = (SwitchStatement)((MethodBodyStatements)ifStatement.Body).Statements[0];

Assert.AreEqual(3, switchStatement.Cases.Count);
for (int i = 0; i < switchStatement.Cases.Count; i++)
{
Assert.IsNotNull(switchStatement.Cases[i], $"Switch case at index {i} should not be null.");
}

// Writing the type must not throw - a null case previously caused a NullReferenceException.
Assert.DoesNotThrow(() => new TypeProviderWriter(serialization).Write());
}

[Test]
public void UnknownVariantJsonCreateCoreShouldReturnDeserializeBase()
{
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -760,7 +760,11 @@ protected internal override ConstructorProvider[] BuildConstructors()
: _inputModel.Usage.HasFlag(InputModelTypeUsage.Input)
? MethodSignatureModifiers.Public
: MethodSignatureModifiers.Internal;
var (constructorParameters, constructorInitializer) = BuildConstructorParameters(true);
var includeDiscriminatorParameter = _isDiscriminatedBaseType
&& BaseModelProvider?._inputModel.DiscriminatorProperty is not null;
var (constructorParameters, constructorInitializer) = BuildConstructorParameters(
true,
includeDiscriminatorParameter);

var constructor = new ConstructorProvider(
signature: new ConstructorSignature(
Expand Down Expand Up @@ -1310,7 +1314,7 @@ private IEnumerable<FieldProvider> GetAllBaseFieldsForConstructorInitialization(
? baseParameters
: baseParameters.Where(p =>
p.Property is null
|| (!overriddenProperties.Contains(p.Property!) && (!p.Property.IsDiscriminator || !isInitializationConstructor || (includeDiscriminatorParameter && IsMultiLevelDiscriminator)))));
|| (!overriddenProperties.Contains(p.Property!) && (!p.Property.IsDiscriminator || !isInitializationConstructor || includeDiscriminatorParameter))));

// construct the initializer using the parameters from base signature
ConstructorInitializer? constructorInitializer = null;
Expand All @@ -1322,9 +1326,8 @@ p.Property is null
if (isInitializationConstructor && (IsMultiLevelDiscriminator || BaseModelProvider.IsMultiLevelDiscriminator))
{
var baseDiscriminatorParam = baseParameters.FirstOrDefault(p => p.Property?.IsDiscriminator == true);
var hasDiscriminatorProperty = BaseModelProvider.CanonicalView.Properties.Any(p => p.IsDiscriminator);

ValueExpression discriminatorExpression = (hasDiscriminatorProperty && baseDiscriminatorParam is not null && includeDiscriminatorParameter)
ValueExpression discriminatorExpression = (baseDiscriminatorParam is not null && includeDiscriminatorParameter)
? constructorParameters.FirstOrDefault(p => p.Property?.IsDiscriminator == true) ?? baseDiscriminatorParam
: DiscriminatorLiteral;

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -177,6 +177,101 @@ public void BaseConstructorShouldBePrivateProtected()
Assert.IsTrue(serializationCtor!.Signature.Modifiers.HasFlag(MethodSignatureModifiers.Internal));
}

[TestCase(false)]
[TestCase(true)]
public void RebasedDiscriminatedBaseConstructorForwardsDeclaredParameter(bool useMultiLevelBase)
{
var rootModel = InputFactory.Model(
"conversationItem",
properties:
[
InputFactory.Property(
"type",
InputPrimitiveType.String,
isRequired: true,
isDiscriminator: true)
]);
var rebasedModel = rootModel;

if (useMultiLevelBase)
{
rebasedModel = InputFactory.Model(
"realtimeConversationItem",
properties:
[
InputFactory.Property(
"type",
InputPrimitiveType.String,
isRequired: true,
isDiscriminator: true)
],
baseModel: rootModel,
discriminatedKind: "realtime");
}

var voiceModel = InputFactory.Model(
"voiceConversationItem",
properties:
[
InputFactory.Property(
"type",
InputPrimitiveType.String,
isRequired: true,
isDiscriminator: true)
],
baseModel: rebasedModel);

MockHelpers.LoadMockGenerator(inputModelTypes: [rootModel, rebasedModel, voiceModel]);
var model = CodeModelGenerator.Instance.TypeFactory.CreateModel(voiceModel);

Assert.IsNotNull(model);
var constructor = model!.Constructors.Single(c =>
c.Signature.Modifiers.HasFlag(MethodSignatureModifiers.Private)
&& c.Signature.Modifiers.HasFlag(MethodSignatureModifiers.Protected));
Assert.AreEqual(1, constructor.Signature.Parameters.Count);
Assert.AreEqual("type", constructor.Signature.Parameters[0].Name);
Assert.IsNotNull(constructor.Signature.Initializer);
Assert.AreEqual(1, constructor.Signature.Initializer!.Arguments.Count);
Assert.AreEqual("@type", constructor.Signature.Initializer.Arguments[0].ToDisplayString());
}

[Test]
public void RebasedDiscriminatedBaseWritesDeclaredConstructorParameter()
{
var rootModel = InputFactory.Model(
"conversationItem",
properties:
[
InputFactory.Property(
"type",
InputPrimitiveType.String,
isRequired: true,
isDiscriminator: true)
]);
var voiceModel = InputFactory.Model(
"voiceConversationItem",
properties:
[
InputFactory.Property(
"type",
InputPrimitiveType.String,
isRequired: true,
isDiscriminator: true)
],
baseModel: rootModel);

MockHelpers.LoadMockGenerator(inputModelTypes: [rootModel, voiceModel]);
var model = CodeModelGenerator.Instance.TypeFactory.CreateModel(voiceModel);

Assert.IsNotNull(model);
var content = new TypeProviderWriter(model!).Write().Content;

// The base call must only forward parameters that the constructor actually declares.
StringAssert.Contains(
"private protected VoiceConversationItem(string @type) : base(@type)",
content);
}

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