diff --git a/packages/http-client-csharp/generator/Microsoft.TypeSpec.Generator.ClientModel/src/Providers/MrwSerializationTypeDefinition.cs b/packages/http-client-csharp/generator/Microsoft.TypeSpec.Generator.ClientModel/src/Providers/MrwSerializationTypeDefinition.cs index 590d2d4b387..c747e93db7b 100644 --- a/packages/http-client-csharp/generator/Microsoft.TypeSpec.Generator.ClientModel/src/Providers/MrwSerializationTypeDefinition.cs +++ b/packages/http-client-csharp/generator/Microsoft.TypeSpec.Generator.ClientModel/src/Providers/MrwSerializationTypeDefinition.cs @@ -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 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]; } /// diff --git a/packages/http-client-csharp/generator/Microsoft.TypeSpec.Generator.ClientModel/test/Providers/MrwSerializationTypeDefinitions/DiscriminatorTests.cs b/packages/http-client-csharp/generator/Microsoft.TypeSpec.Generator.ClientModel/test/Providers/MrwSerializationTypeDefinitions/DiscriminatorTests.cs index f8ccf8586b8..d2e99a056a0 100644 --- a/packages/http-client-csharp/generator/Microsoft.TypeSpec.Generator.ClientModel/test/Providers/MrwSerializationTypeDefinitions/DiscriminatorTests.cs +++ b/packages/http-client-csharp/generator/Microsoft.TypeSpec.Generator.ClientModel/test/Providers/MrwSerializationTypeDefinitions/DiscriminatorTests.cs @@ -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() { diff --git a/packages/http-client-csharp/generator/Microsoft.TypeSpec.Generator/src/Providers/ModelProvider.cs b/packages/http-client-csharp/generator/Microsoft.TypeSpec.Generator/src/Providers/ModelProvider.cs index 22d30b35bcb..9e726fb171f 100644 --- a/packages/http-client-csharp/generator/Microsoft.TypeSpec.Generator/src/Providers/ModelProvider.cs +++ b/packages/http-client-csharp/generator/Microsoft.TypeSpec.Generator/src/Providers/ModelProvider.cs @@ -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( @@ -1310,7 +1314,7 @@ private IEnumerable 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; @@ -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; diff --git a/packages/http-client-csharp/generator/Microsoft.TypeSpec.Generator/test/Providers/ModelProviders/DiscriminatorTests.cs b/packages/http-client-csharp/generator/Microsoft.TypeSpec.Generator/test/Providers/ModelProviders/DiscriminatorTests.cs index 5e5e6cb31ec..7e42ef5030c 100644 --- a/packages/http-client-csharp/generator/Microsoft.TypeSpec.Generator/test/Providers/ModelProviders/DiscriminatorTests.cs +++ b/packages/http-client-csharp/generator/Microsoft.TypeSpec.Generator/test/Providers/ModelProviders/DiscriminatorTests.cs @@ -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() {