using System.Diagnostics; using System.Text; using Microsoft.CodeAnalysis; using Microsoft.CodeAnalysis.CSharp; using Microsoft.CodeAnalysis.CSharp.Syntax; using Microsoft.CodeAnalysis.Text; using static Microsoft.CodeAnalysis.SymbolDisplayFormat; using static Microsoft.CodeAnalysis.SymbolDisplayMiscellaneousOptions; using Robust.Roslyn.Shared; // Yes dude I know this source generator isn't incremental, I'll fix it eventually. #pragma warning disable RS1035 namespace Robust.Shared.CompNetworkGenerator { [Generator] #pragma warning disable RS1042 public class ComponentNetworkGenerator : ISourceGenerator #pragma warning restore RS1042 { private const string ClassAttributeName = "Robust.Shared.Analyzers.AutoGenerateComponentStateAttribute"; private const string MemberAttributeName = "Robust.Shared.Analyzers.AutoNetworkedFieldAttribute"; private const string GlobalEntityUidName = "global::Robust.Shared.GameObjects.EntityUid"; private const string GlobalNullableEntityUidName = "global::Robust.Shared.GameObjects.EntityUid?"; private const string GlobalNetEntityName = "global::Robust.Shared.GameObjects.NetEntity"; private const string GlobalNetEntityNullableName = "global::Robust.Shared.GameObjects.NetEntity?"; private const string GlobalEntityCoordinatesName = "global::Robust.Shared.Map.EntityCoordinates"; private const string GlobalNullableEntityCoordinatesName = "global::Robust.Shared.Map.EntityCoordinates?"; private const string GlobalEntityUidSetName = "global::System.Collections.Generic.HashSet"; private const string GlobalNetEntityUidSetName = $"global::System.Collections.Generic.HashSet<{GlobalNetEntityName}>"; private const string GlobalEntityUidListName = "global::System.Collections.Generic.List"; private const string GlobalNetEntityUidListName = $"global::System.Collections.Generic.List<{GlobalNetEntityName}>"; private const string GlobalDictionaryName = "global::System.Collections.Generic.Dictionary"; private const string GlobalHashSetName = "global::System.Collections.Generic.HashSet"; private const string GlobalListName = "global::System.Collections.Generic.List"; private const string GlobalIRobustCloneableName = "global::Robust.Shared.Serialization.IRobustCloneable"; private static readonly SymbolDisplayFormat FullNullableFormat = FullyQualifiedFormat.WithMiscellaneousOptions(IncludeNullableReferenceTypeModifier); private static string? GenerateSource( in GeneratorExecutionContext context, INamedTypeSymbol classSymbol, TypeDeclarationSyntax classSyntax, CSharpCompilation comp, bool raiseAfterAutoHandle, bool fieldDeltas, bool excludeReplays) { var partialInfo = PartialTypeInfo.FromSymbol(classSymbol, classSyntax); var componentName = classSymbol.Name; var stateName = $"{componentName}_AutoState"; var members = TypeSymbolHelper.GetAllMembersIncludingInherited(classSymbol); var fields = new List<(ITypeSymbol Type, string FieldName)>(); var fieldAttr = comp.GetTypeByMetadataName(MemberAttributeName); foreach (var mem in members) { var attribute = mem.GetAttributes().FirstOrDefault(a => a.AttributeClass != null && a.AttributeClass.Equals(fieldAttr, SymbolEqualityComparer.Default)); if (attribute == null) { continue; } switch (mem) { case IFieldSymbol field: fields.Add((field.Type, field.Name)); break; case IPropertySymbol prop: { if (prop.SetMethod == null || prop.SetMethod.DeclaredAccessibility != Accessibility.Public) { var msg = "Property is marked with [AutoNetworkedField], but has no accessible setter method."; context.ReportDiagnostic( Diagnostic.Create( new DiagnosticDescriptor( "RXN0008", msg, msg, "Usage", DiagnosticSeverity.Error, true), classSymbol.Locations[0])); continue; } if (prop.GetMethod == null || prop.GetMethod.DeclaredAccessibility != Accessibility.Public) { var msg = "Property is marked with [AutoNetworkedField], but has no accessible getter method."; context.ReportDiagnostic( Diagnostic.Create( new DiagnosticDescriptor( "RXN0008", msg, msg, "Usage", DiagnosticSeverity.Error, true), classSymbol.Locations[0])); continue; } fields.Add((prop.Type, prop.Name)); break; } } } if (fields.Count == 0) { var msg = "Component is marked with [AutoGenerateComponentState], but has no valid members marked with [AutoNetworkedField]."; context.ReportDiagnostic( Diagnostic.Create( new DiagnosticDescriptor( "RXN0007", msg, msg, "Usage", DiagnosticSeverity.Error, true), classSymbol.Locations[0])); return null; } // eg: // public string Name = default!; // public int Count = default!; var stateFields = new StringBuilder(); // eg: // Name = component.Name, // Count = component.Count, var getStateInit = new StringBuilder(); var clientGetStateInit = new StringBuilder(); // eg: // component.Name = state.Name; // component.Count = state.Count; var handleStateSetters = new StringBuilder(); // Builds the string for duplicating a full component state, in preparation for applying a delta state state // without modifying the original. var shallowClone = new StringBuilder(); // Delta field states var deltaGetFields = new StringBuilder(); var clientDeltaGetFields = new StringBuilder(); var deltaHandleFields = new StringBuilder(); // Apply the delta field to the full state. var deltaApply = new List(); var index = -1; var fieldsStr = new StringBuilder(); var fieldStates = new StringBuilder(); var networkedTypes = new List(); var usesClientCollectionCopy = false; void AppendShallowClone(string fieldName) { shallowClone.Append($@" {fieldName} = this.{fieldName},"); } void AppendCollectionClone(string fieldName, bool nullable) { var value = nullable ? $"this.{fieldName} == null ? null! : new(this.{fieldName})" : $"new(this.{fieldName})"; shallowClone.Append($@" {fieldName} = {value},"); } string GetClientCollectionField(string fieldName, bool nullable) { usesClientCollectionCopy = true; return nullable ? $"component.{fieldName} == null ? null! : new(component.{fieldName})" : $"new(component.{fieldName})"; } string GetCollectionRefill(ITypeSymbol type, string target, string source, string indentation) { var named = (INamedTypeSymbol) type; return named.ConstructedFrom.ToDisplayString(FullyQualifiedFormat) switch { GlobalDictionaryName => $@"foreach (var (key, value) in {source}) {indentation} {target}.Add(key, value);", GlobalHashSetName => $"{target}.UnionWith({source});", GlobalListName => $"{target}.AddRange({source});", _ => throw new InvalidOperationException($"Unsupported collection type {type}") }; } foreach (var (type, name) in fields) { index++; if (index == 0) { fieldsStr.Append(@$"""{name}"""); } else { fieldsStr.Append(@$", ""{name}"""); } var typeDisplayStr = type.ToDisplayString(FullNullableFormat); var nullable = type.NullableAnnotation == NullableAnnotation.Annotated; var nullableAnnotation = nullable ? "?" : string.Empty; string deltaStateName = $"{name}_FieldComponentState"; // The type used for networking, e.g. EntityUid -> NetEntity string networkedType; string getField; string? clientGetField = null; string? cast; // TODO: Uhh I just need casts or something. var castString = typeDisplayStr.Substring(8); deltaGetFields.Append(@$" case {Math.Pow(2, index)}: args.State = new {deltaStateName}() {{ "); clientDeltaGetFields.Append(@$" case {Math.Pow(2, index)}: args.State = new {deltaStateName}() {{ "); deltaHandleFields.Append(@$" case {deltaStateName} {deltaStateName}_State: {{"); var fieldHandleValue = $"{deltaStateName}_State.{name}!"; switch (typeDisplayStr) { case GlobalEntityUidName: case GlobalNullableEntityUidName: networkedType = $"NetEntity{nullableAnnotation}"; stateFields.Append($@" public {networkedType} {name} = default!;"); getField = $"GetNetEntity(component.{name})"; cast = $"(NetEntity{nullableAnnotation})"; handleStateSetters.Append($@" component.{name} = EnsureEntity<{componentName}>(state.{name}, uid);"); deltaHandleFields.Append($@" component.{name} = EnsureEntity<{componentName}>({cast} {fieldHandleValue}, uid);"); AppendShallowClone(name); deltaApply.Add($"fullState.{name} = {name};"); break; case GlobalEntityCoordinatesName: case GlobalNullableEntityCoordinatesName: networkedType = $"NetCoordinates{nullableAnnotation}"; stateFields.Append($@" public {networkedType} {name} = default!;"); getField = $"GetNetCoordinates(component.{name})"; cast = $"(NetCoordinates{nullableAnnotation})"; handleStateSetters.Append($@" component.{name} = EnsureCoordinates<{componentName}>(state.{name}, uid);"); deltaHandleFields.Append($@" component.{name} = EnsureCoordinates<{componentName}>({cast} {fieldHandleValue}, uid);"); AppendShallowClone(name); deltaApply.Add($@"fullState.{name} = {name};"); break; case GlobalEntityUidSetName: networkedType = $"{GlobalNetEntityUidSetName}"; stateFields.Append($@" public {networkedType} {name} = default!;"); getField = $"GetNetEntitySet(component.{name})"; cast = $"({GlobalNetEntityUidSetName})"; handleStateSetters.Append($@" EnsureEntitySet<{componentName}>(state.{name}, uid, component.{name});"); deltaHandleFields.Append($@" EnsureEntitySet<{componentName}>({cast} {fieldHandleValue}, uid, component.{name});"); AppendCollectionClone(name, nullable); deltaApply.Add($@"fullState.{name} = {name};"); break; case GlobalEntityUidListName: networkedType = $"{GlobalNetEntityUidListName}"; stateFields.Append($@" public {networkedType} {name} = default!;"); getField = $"GetNetEntityList(component.{name})"; cast = $"({GlobalNetEntityUidListName})"; handleStateSetters.Append($@" EnsureEntityList<{componentName}>(state.{name}, uid, component.{name});"); deltaHandleFields.Append($@" EnsureEntityList<{componentName}>({cast} {fieldHandleValue}, uid, component.{name});"); AppendCollectionClone(name, nullable); deltaApply.Add($@"fullState.{name} = {name};"); break; default: if (type is INamedTypeSymbol { TypeArguments.Length: 2 } named && named.ConstructedFrom.ToDisplayString(FullyQualifiedFormat) == GlobalDictionaryName) { var key = named.TypeArguments[0].ToDisplayString(FullNullableFormat); var keyNullable = key.EndsWith("?"); var value = named.TypeArguments[1].ToDisplayString(FullNullableFormat); var valueNullable = value.EndsWith("?"); if (key is GlobalEntityUidName or GlobalNullableEntityUidName) { key = keyNullable ? GlobalNetEntityNullableName : GlobalNetEntityName; var ensureGeneric = $"{componentName}, {value}"; if (value is GlobalEntityUidName or GlobalNullableEntityUidName) { value = valueNullable ? GlobalNetEntityNullableName : GlobalNetEntityName; ensureGeneric = componentName; } networkedType = $"Dictionary<{key}, {value}>"; stateFields.Append($@" public {networkedType} {name} = default!;"); getField = $"GetNetEntityDictionary(component.{name})"; if (valueNullable && value is not GlobalNetEntityName and not GlobalNetEntityNullableName) { cast = $"(Dictionary<{key}, {value}>)"; handleStateSetters.Append($@" EnsureEntityDictionaryNullableValue<{componentName}, {value}>(state.{name}, uid, component.{name});"); deltaHandleFields.Append($@" EnsureEntityDictionaryNullableValue<{componentName}, {value}>({cast} {fieldHandleValue}, uid, component.{name});"); } else { cast = $"({castString})"; handleStateSetters.Append($@" EnsureEntityDictionary<{ensureGeneric}>(state.{name}, uid, component.{name});"); deltaHandleFields.Append($@" EnsureEntityDictionary<{ensureGeneric}>({cast} {fieldHandleValue}, uid, component.{name});"); } AppendCollectionClone(name, nullable); deltaApply.Add($@"fullState.{name} = {name};"); break; } if (value is GlobalEntityUidName or GlobalNullableEntityUidName) { value = valueNullable ? GlobalNetEntityNullableName : GlobalNetEntityName; networkedType = $"Dictionary<{key}, {value}>"; stateFields.Append($@" public {networkedType} {name} = default!;"); getField = $"GetNetEntityDictionary(component.{name})"; cast = $"(Dictionary<{key}, {value}>)"; handleStateSetters.Append($@" EnsureEntityDictionary<{componentName}, {key}>(state.{name}, uid, component.{name});"); deltaHandleFields.Append($@" EnsureEntityDictionary<{componentName}, {key}>({cast} {fieldHandleValue}, uid, component.{name});"); AppendCollectionClone(name, nullable); deltaApply.Add($@"fullState.{name} = {name};"); break; } } networkedType = $"{typeDisplayStr}"; stateFields.Append($@" public {networkedType} {name} = default!;"); if (ImplementsInterface(type, GlobalIRobustCloneableName)) { getField = $"component.{name}"; cast = $"({castString})"; var nullCast = nullable ? castString.Substring(0, castString.Length - 1) : castString; if (nullable) { handleStateSetters.Append($@" component.{name} = state.{name} == null ? null! : state.{name}.Clone();"); deltaHandleFields.Append($@" var {name}Value = {cast} {fieldHandleValue}; if ({name}Value == null) component.{name} = null!; else component.{name} = ({nullCast})({name}Value.Clone());"); AppendShallowClone(name); deltaApply.Add($"fullState.{name} = {name} == null ? null! : {name}.Clone();"); } else { handleStateSetters.Append($@" component.{name} = state.{name}.Clone();"); deltaHandleFields.Append($@" component.{name} = {cast}({fieldHandleValue}.Clone());"); AppendShallowClone(name); deltaApply.Add($"fullState.{name} = {name}.Clone();"); } } else if (IsCloneType(type)) { getField = $"component.{name}"; clientGetField = GetClientCollectionField(name, nullable); cast = $"({castString})"; var nullCast = nullable ? castString.Substring(0, castString.Length - 1) : castString; var handleRefill = GetCollectionRefill(type, $"component.{name}", $"state.{name}", " "); var deltaRefill = GetCollectionRefill(type, $"component.{name}", $"{name}Value", " "); if (nullable) { handleStateSetters.Append($@" if (state.{name} == null) component.{name} = null!; else if (component.{name} == null) component.{name} = new(state.{name}); else if (!ReferenceEquals(component.{name}, state.{name})) {{ component.{name}.Clear(); {handleRefill} }}"); deltaHandleFields.Append($@" var {name}Value = {cast} {fieldHandleValue}; if ({name}Value == null) component.{name} = null!; else if (component.{name} == null) component.{name} = new {nullCast}({name}Value); else if (!ReferenceEquals(component.{name}, {name}Value)) {{ component.{name}.Clear(); {deltaRefill} }}"); deltaApply.Add($"fullState.{name} = {name} == null ? null! : new({name});"); } else { handleStateSetters.Append($@" if (!ReferenceEquals(component.{name}, state.{name})) {{ component.{name}.Clear(); {handleRefill} }}"); deltaHandleFields.Append($@" var {name}Value = {cast} {fieldHandleValue}; if (!ReferenceEquals(component.{name}, {name}Value)) {{ component.{name}.Clear(); {deltaRefill} }}"); deltaApply.Add($"fullState.{name} = new({name});"); } AppendCollectionClone(name, nullable); } else { getField = $"component.{name}"; cast = $"({castString})"; handleStateSetters.Append($@" component.{name} = state.{name};"); deltaHandleFields.Append($@" component.{name} = {cast} {fieldHandleValue};"); AppendShallowClone(name); deltaApply.Add($"fullState.{name} = {name};"); } break; } /* * End loop stuff */ networkedTypes.Add(networkedType); clientGetField ??= getField; getStateInit.Append($@" {name} = {getField},"); clientGetStateInit.Append($@" {name} = {clientGetField},"); deltaGetFields.Append(@$" {name} = {getField} }}; return;"); clientDeltaGetFields.Append(@$" {name} = {clientGetField} }}; return;"); deltaHandleFields.Append(@" break; } "); } var deltaGetState = ""; var clientDeltaGetState = ""; var deltaInterface = ""; var deltaCompFields = ""; var deltaNetRegister = ""; var cloneMethod = ""; if (fieldDeltas) { cloneMethod = $@" public {stateName} ShallowClone() {{ return new {stateName}() {{{shallowClone} }}; }} "; for (var i = 0; i < fields.Count; i++) { var name = fields[i].FieldName; string deltaStateName = $"{name}_FieldComponentState"; var networkedType = networkedTypes[i]; var apply = deltaApply[i]; // Creates a state per field fieldStates.Append($@" [Serializable, NetSerializable] public sealed class {deltaStateName} : IComponentDeltaState<{stateName}> {{ public {networkedType} {name} = default!; public void ApplyToFullState({stateName} fullState) {{ {apply} }} public {stateName} CreateNewFullState({stateName} fullState) {{ var newState = fullState.ShallowClone(); ApplyToFullState(newState); return newState; }} }} "); } deltaNetRegister = $@"EntityManager.ComponentFactory.RegisterNetworkedFields<{classSymbol}>({fieldsStr});"; deltaGetState = @$"// Delta state if (component is IComponentDelta delta && args.FromTick > component.CreationTick) {{ var aspects = EntityManager.GetModifiedAspects(component, args.FromTick); // Try and get a matching delta state for the relevant dirty fields, otherwise fall back to full state. switch (aspects) {{ case >= DeltaAspect.Unclassified: break;{deltaGetFields} default: break; }} }}"; clientDeltaGetState = @$"// Delta state if (component is IComponentDelta delta && args.FromTick > component.CreationTick) {{ var aspects = EntityManager.GetModifiedAspects(component, args.FromTick); // Try and get a matching delta state for the relevant dirty fields, otherwise fall back to full state. switch (aspects) {{ case >= DeltaAspect.Unclassified: break;{clientDeltaGetFields} default: break; }} }}"; deltaInterface = " : IComponentDelta"; deltaCompFields = @$"/// public GameTick LastUnclassifiedDirty {{ get; set; }} /// public GameTick[] LastModifiedFields {{ get; set; }}"; } string handleState; if (!fieldDeltas) { var eventRaise = ""; var stateSetters = TrimNewLines(handleStateSetters); if (raiseAfterAutoHandle) { eventRaise = @" var ev = new AfterAutoHandleStateEvent(args.Current); EntityManager.EventBus.RaiseComponentEvent(uid, component, ref ev);"; } handleState = $@" if (args.Current is not {stateName} state) return; {stateSetters}{eventRaise}"; } else { // Re-indent handleStateSetters so it aligns with the switch block var stateSetters = TrimNewLines(handleStateSetters); stateSetters = stateSetters.Replace(" ", " "); var eventRaise = ""; if (raiseAfterAutoHandle) { eventRaise = @" if (args.Current is not {} current) return; var ev = new AfterAutoHandleStateEvent(current); EntityManager.EventBus.RaiseComponentEvent(uid, component, ref ev);"; } handleState = $@" switch(args.Current) {{{deltaHandleFields} case {stateName} state: {{{stateSetters} break; }} default: return; }}{eventRaise}"; } var excludeReplaysStr = string.Empty; if (excludeReplays) { excludeReplaysStr = @" if (args.ReplayState) { args.ExcludeReplays = true; return; } "; } var outSb = new StringBuilder(); var stateFieldsText = TrimNewLines(stateFields); var getStateInitText = TrimNewLines(getStateInit); var clientGetStateInitText = TrimNewLines(clientGetStateInit); var cloneMethodText = TrimNewLines(cloneMethod); var excludeReplaysText = TrimNewLines(excludeReplaysStr); var deltaGetStateText = TrimNewLines(deltaGetState); var clientDeltaGetStateText = TrimNewLines(clientDeltaGetState); var deltaCompFieldsText = TrimNewLines(deltaCompFields); var fieldStatesText = TrimNewLines(fieldStates); var netManagerDependency = usesClientCollectionCopy ? "[global::Robust.Shared.IoC.Dependency] private global::Robust.Shared.Network.INetManager _net = default!;" : string.Empty; var getStateSubscription = usesClientCollectionCopy ? $@" if (_net.IsClient) SubscribeLocalEvent<{componentName}, ComponentGetState>(OnGetStateClient); else SubscribeLocalEvent<{componentName}, ComponentGetState>(OnGetState);" : $@" SubscribeLocalEvent<{componentName}, ComponentGetState>(OnGetState);"; outSb.Append(""" // #nullable enable using System; using Robust.Shared.GameStates; using Robust.Shared.GameObjects; using Robust.Shared.Analyzers; using Robust.Shared.Collections; using Robust.Shared.Serialization; using Robust.Shared.Map; using Robust.Shared.Timing; using Robust.Shared.Utility; using System.Collections.Generic; """); partialInfo.WriteHeader(outSb); outSb.AppendLine(deltaInterface); outSb.AppendLine("{"); if (deltaCompFieldsText.Length != 0) { outSb.AppendLine(deltaCompFieldsText); outSb.AppendLine(); } outSb.AppendLine(" [System.Serializable, NetSerializable]"); outSb.AppendLine(" [global::System.ComponentModel.EditorBrowsable(global::System.ComponentModel.EditorBrowsableState.Never)]"); outSb.AppendLine(" [RobustAutoGenerated]"); outSb.AppendLine($" public sealed class {stateName} : IComponentState"); outSb.AppendLine(" {"); outSb.AppendLine(stateFieldsText); if (cloneMethodText.Length != 0) { outSb.AppendLine(); outSb.AppendLine(cloneMethodText); } outSb.AppendLine(" }"); outSb.AppendLine(); outSb.AppendLine(" [RobustAutoGenerated]"); outSb.AppendLine(" [global::System.ComponentModel.EditorBrowsable(global::System.ComponentModel.EditorBrowsableState.Never)]"); outSb.AppendLine($" public sealed class {componentName}_AutoNetworkSystem : EntitySystem"); outSb.AppendLine(" {"); if (netManagerDependency.Length != 0) { outSb.AppendLine($" {netManagerDependency}"); outSb.AppendLine(); } outSb.AppendLine(" public override void Initialize()"); outSb.AppendLine(" {"); if (deltaNetRegister.Length != 0) outSb.AppendLine($" {deltaNetRegister}"); outSb.AppendLine(getStateSubscription); outSb.AppendLine($" SubscribeLocalEvent<{componentName}, ComponentHandleState>(OnHandleState);"); outSb.AppendLine(" }"); outSb.AppendLine(); outSb.AppendLine($" private void OnGetState(EntityUid uid, {componentName} component, ref ComponentGetState args)"); outSb.AppendLine(" {"); if (excludeReplaysStr.Length != 0) { outSb.AppendLine(IndentFirstLine(excludeReplaysText, 12)); outSb.AppendLine(); } if (deltaGetStateText.Length != 0) { outSb.AppendLine(IndentFirstLine(deltaGetStateText, 12)); outSb.AppendLine(); } outSb.AppendLine(" // Get full state"); outSb.AppendLine($" args.State = new {stateName}"); outSb.AppendLine(" {"); outSb.AppendLine(getStateInitText); outSb.AppendLine(" };"); outSb.AppendLine(" }"); if (usesClientCollectionCopy) { outSb.AppendLine(); outSb.AppendLine($" private void OnGetStateClient(EntityUid uid, {componentName} component, ref ComponentGetState args)"); outSb.AppendLine(" {"); if (clientDeltaGetStateText.Length != 0) { outSb.AppendLine(IndentFirstLine(clientDeltaGetStateText, 12)); outSb.AppendLine(); } outSb.AppendLine(" // Get full state"); outSb.AppendLine($" args.State = new {stateName}"); outSb.AppendLine(" {"); outSb.AppendLine(clientGetStateInitText); outSb.AppendLine(" };"); outSb.AppendLine(" }"); } outSb.AppendLine(); outSb.AppendLine($" private void OnHandleState(EntityUid uid, {componentName} component, ref ComponentHandleState args)"); outSb.AppendLine(" {"); outSb.AppendLine(TrimNewLines(handleState)); outSb.AppendLine(" }"); outSb.AppendLine(" }"); if (fieldStatesText.Length != 0) { outSb.AppendLine(); outSb.AppendLine(fieldStatesText); } outSb.AppendLine("}"); partialInfo.WriteFooter(outSb); return outSb.ToString(); } private static string TrimNewLines(StringBuilder source) { return source.ToString().Trim('\r', '\n'); } private static string TrimNewLines(string source) { return source.Trim('\r', '\n'); } private static string IndentFirstLine(string source, int spaces) { if (source.Length == 0) return source; return new string(' ', spaces) + source; } public void Execute(GeneratorExecutionContext context) { var comp = (CSharpCompilation) context.Compilation; if (!(context.SyntaxReceiver is NameReferenceSyntaxReceiver receiver)) { return; } var symbols = GetAnnotatedTypes(context, comp, receiver); // Generate component sources and add foreach (var (classType, classSyntax, attribute) in symbols) { try { var raiseEv = false; var fieldDeltas = false; var excludeReplays = false; if (attribute.ConstructorArguments is [{Value: bool raise}, {Value: bool fields}, {Value: bool exclude}]) { // Get the afterautohandle bool, which is first constructor arg raiseEv = raise; fieldDeltas = fields; excludeReplays = exclude; } var source = GenerateSource(context, classType, classSyntax, comp, raiseEv, fieldDeltas, excludeReplays); // can be null if no members marked with network field, which already has a diagnostic, so // just continue if (source == null) continue; context.AddSource($"{classType.Name}_CompNetwork.g.cs", SourceText.From(source, Encoding.UTF8)); } catch (Exception e) { context.ReportDiagnostic( Diagnostic.Create( new DiagnosticDescriptor( "RXN0003", "Unhandled exception occured while generating automatic component state handling.", $"Unhandled exception occured while generating automatic component state handling: {e}", "Usage", DiagnosticSeverity.Error, true), classType.Locations[0])); } } } private IReadOnlyList<(INamedTypeSymbol Type, TypeDeclarationSyntax Syntax, AttributeData Attribute)> GetAnnotatedTypes( in GeneratorExecutionContext context, CSharpCompilation comp, NameReferenceSyntaxReceiver receiver) { var symbols = new List<(INamedTypeSymbol, TypeDeclarationSyntax, AttributeData)>(); var attributeSymbol = comp.GetTypeByMetadataName(ClassAttributeName); var fieldAttr = comp.GetTypeByMetadataName(MemberAttributeName); foreach (var candidateClass in receiver.CandidateClasses) { var model = comp.GetSemanticModel(candidateClass.SyntaxTree); var typeSymbol = model.GetDeclaredSymbol(candidateClass); var relevantAttribute = typeSymbol?.GetAttributes().FirstOrDefault(attr => attr.AttributeClass != null && attr.AttributeClass.Equals(attributeSymbol, SymbolEqualityComparer.Default)); if (typeSymbol == null) continue; if (relevantAttribute == null) { foreach (var mem in TypeSymbolHelper.GetAllMembersIncludingInherited(typeSymbol)) { var attribute = mem.GetAttributes().FirstOrDefault(a => a.AttributeClass != null && a.AttributeClass.Equals(fieldAttr, SymbolEqualityComparer.Default)); if (attribute == null) continue; var msg = "Field is marked with [AutoNetworkedField], but its class has no [AutoGenerateComponentState] attribute."; context.ReportDiagnostic( Diagnostic.Create( new DiagnosticDescriptor( "RXN0007", msg, msg, "Usage", DiagnosticSeverity.Error, true), candidateClass.Keyword.GetLocation())); } continue; } var isPartial = candidateClass.Modifiers.Any(m => m.IsKind(SyntaxKind.PartialKeyword)); if (isPartial) { symbols.Add((typeSymbol, candidateClass, relevantAttribute)); } else { var missingPartialKeywordMessage = $"The type {typeSymbol.Name} should be declared with the 'partial' keyword " + "as it is annotated with the [AutoGenerateComponentState] attribute."; context.ReportDiagnostic( Diagnostic.Create( new DiagnosticDescriptor( "RXN0006", missingPartialKeywordMessage, missingPartialKeywordMessage, "Usage", DiagnosticSeverity.Error, true), candidateClass.Keyword.GetLocation())); } } return symbols; } public void Initialize(GeneratorInitializationContext context) { if (!Debugger.IsAttached) { //Debugger.Launch(); } context.RegisterForSyntaxNotifications(() => new NameReferenceSyntaxReceiver()); } private static bool IsCloneType(ITypeSymbol type) { if (type is not INamedTypeSymbol named || !named.IsGenericType) { return false; } var constructed = named.ConstructedFrom.ToDisplayString(FullyQualifiedFormat); return constructed switch { GlobalDictionaryName or GlobalHashSetName or GlobalListName => true, _ => false }; } private static bool ImplementsInterface(ITypeSymbol type, string interfaceName) { foreach (var interfaceType in type.AllInterfaces) { if (interfaceType.ToDisplayString(FullyQualifiedFormat).Contains(interfaceName) || interfaceType.ConstructedFrom.ToDisplayString(FullyQualifiedFormat).Contains(interfaceName)) { return true; } } return false; } } }