microsoft/healthcare-shared-components
Publicmirrored from https://github.com/microsoft/healthcare-shared-componentsAvailable
tools/Microsoft.Health.Extensions.BuildTimeCodeGenerator/Sql/CreateProcedureVisitor.cs
401lines · modecode
| 1 | // ------------------------------------------------------------------------------------------------- |
| 2 | // Copyright (c) Microsoft Corporation. All rights reserved. |
| 3 | // Licensed under the MIT License (MIT). See LICENSE in the repo root for license information. |
| 4 | // ------------------------------------------------------------------------------------------------- |
| 5 | |
| 6 | using System; |
| 7 | using System.Collections.Generic; |
| 8 | using System.Data; |
| 9 | using System.Linq; |
| 10 | using Microsoft.CodeAnalysis.CSharp; |
| 11 | using Microsoft.CodeAnalysis.CSharp.Syntax; |
| 12 | using Microsoft.SqlServer.TransactSql.ScriptDom; |
| 13 | using static Microsoft.CodeAnalysis.CSharp.SyntaxFactory; |
| 14 | |
| 15 | namespace Microsoft.Health.Extensions.BuildTimeCodeGenerator.Sql |
| 16 | { |
| 17 | /// <summary> |
| 18 | /// Visits a SQL AST, creating a class for each CREATE PROCEDURE statement. Classes have a PopulateCommand method, with a signature derived from the procedure's signature. |
| 19 | /// </summary> |
| 20 | internal class CreateProcedureVisitor : SqlVisitor |
| 21 | { |
| 22 | private const string TvpGeneratorGenericTypeName = "TInput"; |
| 23 | private const string PopulateCommandMethodName = "PopulateCommand"; |
| 24 | private const string CommandParameterName = "command"; |
| 25 | |
| 26 | public override int ArtifactSortOder => 1; |
| 27 | |
| 28 | public override void Visit(CreateProcedureStatement node) |
| 29 | { |
| 30 | var rawProcedureName = node.ProcedureReference.Name.BaseIdentifier.Value; |
| 31 | string procedureName = GetMemberNameWithoutVersionSuffix(node.ProcedureReference.Name); |
| 32 | string schemaQualifiedProcedureName = $"{node.ProcedureReference.Name.SchemaIdentifier.Value}.{rawProcedureName}"; |
| 33 | string className = $"{procedureName}Procedure"; |
| 34 | |
| 35 | ClassDeclarationSyntax classDeclarationSyntax = |
| 36 | ClassDeclaration(className) |
| 37 | .WithModifiers(TokenList(Token(SyntaxKind.InternalKeyword))) |
| 38 | |
| 39 | // derive from StoredProcedure |
| 40 | .WithBaseList( |
| 41 | BaseList( |
| 42 | SingletonSeparatedList<BaseTypeSyntax>( |
| 43 | SimpleBaseType( |
| 44 | IdentifierName("StoredProcedure"))))) |
| 45 | |
| 46 | // call base("dbo.StoredProcedure") |
| 47 | .AddMembers( |
| 48 | ConstructorDeclaration( |
| 49 | Identifier(className)) |
| 50 | .WithModifiers( |
| 51 | TokenList( |
| 52 | Token(SyntaxKind.InternalKeyword))) |
| 53 | .WithInitializer( |
| 54 | ConstructorInitializer( |
| 55 | SyntaxKind.BaseConstructorInitializer, |
| 56 | ArgumentList( |
| 57 | SingletonSeparatedList( |
| 58 | Argument( |
| 59 | LiteralExpression( |
| 60 | SyntaxKind.StringLiteralExpression, |
| 61 | Literal(schemaQualifiedProcedureName))))))) |
| 62 | .WithBody(Block())) |
| 63 | |
| 64 | // add fields for each parameter |
| 65 | .AddMembers(node.Parameters.Select(CreateFieldForParameter).ToArray()) |
| 66 | |
| 67 | // add the PopulateCommand method |
| 68 | .AddMembers(AddPopulateCommandMethod(node, schemaQualifiedProcedureName), AddPopulateCommandMethodForTableValuedParameters(node, procedureName)); |
| 69 | |
| 70 | FieldDeclarationSyntax fieldDeclarationSyntax = CreateStaticFieldForClass(className, procedureName); |
| 71 | |
| 72 | var (tvpGeneratorClass, tvpHolderStruct) = CreateTvpGeneratorTypes(node, procedureName); |
| 73 | |
| 74 | MembersToAdd.Add(classDeclarationSyntax.AddSortingKey(this, procedureName)); |
| 75 | MembersToAdd.Add(fieldDeclarationSyntax.AddSortingKey(this, procedureName)); |
| 76 | MembersToAdd.Add(tvpGeneratorClass.AddSortingKey(this, procedureName)); |
| 77 | MembersToAdd.Add(tvpHolderStruct.AddSortingKey(this, procedureName)); |
| 78 | |
| 79 | base.Visit(node); |
| 80 | } |
| 81 | |
| 82 | /// <summary> |
| 83 | /// Creates a Column-derived field for a stored procedure parameter. |
| 84 | /// </summary> |
| 85 | /// <param name="parameter">The stored procedure parameter</param> |
| 86 | /// <returns>The field declaration</returns> |
| 87 | private MemberDeclarationSyntax CreateFieldForParameter(ProcedureParameter parameter) |
| 88 | { |
| 89 | TypeSyntax typeName; |
| 90 | List<ArgumentSyntax> arguments = new List<ArgumentSyntax> |
| 91 | { |
| 92 | Argument(LiteralExpression(SyntaxKind.StringLiteralExpression, Literal(parameter.VariableName.Value))), |
| 93 | }; |
| 94 | if (TryGetSqlDbTypeForParameter(parameter, out SqlDbType sqlDbType)) |
| 95 | { |
| 96 | // new ParameterDefinition<int>("@paramName", SqlDbType.Int, nullable, maxlength,...) |
| 97 | typeName = GenericName("ParameterDefinition") |
| 98 | .AddTypeArgumentListArguments(SqlDbTypeToClrType(sqlDbType, nullable: parameter.Value != null).ToTypeSyntax(true)); |
| 99 | |
| 100 | arguments.Add( |
| 101 | Argument( |
| 102 | MemberAccessExpression( |
| 103 | SyntaxKind.SimpleMemberAccessExpression, |
| 104 | typeof(SqlDbType).ToTypeSyntax(true), |
| 105 | IdentifierName(sqlDbType.ToString())))); |
| 106 | arguments.Add( |
| 107 | Argument( |
| 108 | LiteralExpression(parameter.Value == null ? SyntaxKind.FalseLiteralExpression : SyntaxKind.TrueLiteralExpression))); |
| 109 | |
| 110 | arguments.AddRange(GetDataTypeSpecificConstructorArguments(parameter.DataType, null)); |
| 111 | } |
| 112 | else |
| 113 | { |
| 114 | // new MyTableValuedParameterDefinition("@paramName"); |
| 115 | typeName = IdentifierName(GetClassNameForTableValuedParameterDefinition(parameter.DataType.Name)); |
| 116 | } |
| 117 | |
| 118 | return FieldDeclaration( |
| 119 | VariableDeclaration(typeName) |
| 120 | |
| 121 | // call it "_myParam" when the parameter is named "@myParam" |
| 122 | .AddVariables(VariableDeclarator(FieldNameForParameter(parameter)) |
| 123 | .WithInitializer( |
| 124 | EqualsValueClause( |
| 125 | ObjectCreationExpression(typeName) |
| 126 | .AddArgumentListArguments(arguments.ToArray()))))) |
| 127 | .AddModifiers(Token(SyntaxKind.PrivateKeyword), Token(SyntaxKind.ReadOnlyKeyword)); |
| 128 | } |
| 129 | |
| 130 | /// <summary> |
| 131 | /// Creates a PopulateCommand method taking a SqlCommand and parameters for each sproc parameter. |
| 132 | /// </summary> |
| 133 | /// <param name="node">The CREATE STORED PROCEDURE statement</param> |
| 134 | /// <param name="schemaQualifiedProcedureName">The full name of the stored procedure</param> |
| 135 | /// <returns>The method declaration</returns> |
| 136 | private MethodDeclarationSyntax AddPopulateCommandMethod(CreateProcedureStatement node, string schemaQualifiedProcedureName) |
| 137 | { |
| 138 | return MethodDeclaration( |
| 139 | typeof(void).ToTypeSyntax(), |
| 140 | Identifier(PopulateCommandMethodName)) |
| 141 | .AddModifiers(Token(SyntaxKind.PublicKeyword)) |
| 142 | |
| 143 | // first parameter is the SqlCommand |
| 144 | .AddParameterListParameters(Parameter(Identifier(CommandParameterName)).WithType(ParseTypeName("SqlCommandWrapper"))) |
| 145 | |
| 146 | // Add a parameter for each stored procedure parameter |
| 147 | .AddParameterListParameters(node.Parameters.Select(selector: p => |
| 148 | Parameter(Identifier(ParameterNameForParameter(p))) |
| 149 | .WithType(DataTypeReferenceToClrType(p.DataType, p.Value != null))).ToArray()) |
| 150 | |
| 151 | // start the body with: |
| 152 | // command.CommandType = CommandType.StoredProcedure |
| 153 | // command.CommandText = "dbo.MySproc" |
| 154 | .AddBodyStatements( |
| 155 | ExpressionStatement( |
| 156 | AssignmentExpression( |
| 157 | SyntaxKind.SimpleAssignmentExpression, |
| 158 | MemberAccessExpression( |
| 159 | SyntaxKind.SimpleMemberAccessExpression, |
| 160 | IdentifierName(CommandParameterName), |
| 161 | IdentifierName("CommandType")), |
| 162 | MemberAccessExpression( |
| 163 | SyntaxKind.SimpleMemberAccessExpression, |
| 164 | typeof(CommandType).ToTypeSyntax(useGlobalAlias: true), |
| 165 | IdentifierName("StoredProcedure")))), |
| 166 | ExpressionStatement( |
| 167 | AssignmentExpression( |
| 168 | SyntaxKind.SimpleAssignmentExpression, |
| 169 | MemberAccessExpression( |
| 170 | SyntaxKind.SimpleMemberAccessExpression, |
| 171 | IdentifierName(CommandParameterName), |
| 172 | IdentifierName("CommandText")), |
| 173 | LiteralExpression(SyntaxKind.StringLiteralExpression, Literal(schemaQualifiedProcedureName))))) |
| 174 | |
| 175 | // now for each parameter generate: |
| 176 | // _fieldForParameter.AddParameter(command, parameterValue) |
| 177 | .AddBodyStatements(node.Parameters.Select(p => (StatementSyntax)ExpressionStatement( |
| 178 | InvocationExpression( |
| 179 | MemberAccessExpression( |
| 180 | SyntaxKind.SimpleMemberAccessExpression, |
| 181 | IdentifierName(FieldNameForParameter(p)), |
| 182 | IdentifierName("AddParameter"))) |
| 183 | .AddArgumentListArguments( |
| 184 | Argument(MemberAccessExpression( |
| 185 | SyntaxKind.SimpleMemberAccessExpression, |
| 186 | IdentifierName(CommandParameterName), |
| 187 | IdentifierName("Parameters"))), |
| 188 | Argument(IdentifierName(ParameterNameForParameter(p)))))).ToArray()); |
| 189 | } |
| 190 | |
| 191 | private MemberDeclarationSyntax AddPopulateCommandMethodForTableValuedParameters(CreateProcedureStatement node, string procedureName) |
| 192 | { |
| 193 | var nonTableParameters = new List<ProcedureParameter>(); |
| 194 | var tableParameters = new List<ProcedureParameter>(); |
| 195 | |
| 196 | foreach (var procedureParameter in node.Parameters) |
| 197 | { |
| 198 | if (TryGetSqlDbTypeForParameter(procedureParameter, out _)) |
| 199 | { |
| 200 | nonTableParameters.Add(procedureParameter); |
| 201 | } |
| 202 | else |
| 203 | { |
| 204 | tableParameters.Add(procedureParameter); |
| 205 | } |
| 206 | } |
| 207 | |
| 208 | if (tableParameters.Count == 0) |
| 209 | { |
| 210 | return IncompleteMember(); |
| 211 | } |
| 212 | |
| 213 | string tableValuedParametersParameterName = "tableValuedParameters"; |
| 214 | |
| 215 | return MethodDeclaration( |
| 216 | typeof(void).ToTypeSyntax(), |
| 217 | Identifier(PopulateCommandMethodName)) |
| 218 | .AddModifiers(Token(SyntaxKind.PublicKeyword)) |
| 219 | |
| 220 | // first parameter is the SqlCommand |
| 221 | .AddParameterListParameters(Parameter(Identifier(CommandParameterName)).WithType(ParseTypeName("SqlCommandWrapper"))) |
| 222 | |
| 223 | // Add a parameter for each non-TVP |
| 224 | .AddParameterListParameters(nonTableParameters.Select(selector: p => |
| 225 | Parameter(Identifier(ParameterNameForParameter(p))) |
| 226 | .WithType(DataTypeReferenceToClrType(p.DataType, p.Value != null))).ToArray()) |
| 227 | |
| 228 | // Add a parameter for the TVP set |
| 229 | .AddParameterListParameters( |
| 230 | Parameter(Identifier(tableValuedParametersParameterName)).WithType(IdentifierName(TableValuedParametersStructName(procedureName)))) |
| 231 | |
| 232 | // Call the overload |
| 233 | .AddBodyStatements( |
| 234 | ExpressionStatement( |
| 235 | InvocationExpression( |
| 236 | IdentifierName(PopulateCommandMethodName)) |
| 237 | .AddArgumentListArguments(Argument(IdentifierName(CommandParameterName))) |
| 238 | .AddArgumentListArguments( |
| 239 | nonTableParameters.Select(p => |
| 240 | Argument(IdentifierName(ParameterNameForParameter(p))) |
| 241 | .WithNameColon(NameColon(ParameterNameForParameter(p)))).ToArray()) |
| 242 | .AddArgumentListArguments( |
| 243 | tableParameters.Select(p => |
| 244 | Argument(MemberAccessExpression( |
| 245 | SyntaxKind.SimpleMemberAccessExpression, |
| 246 | IdentifierName(tableValuedParametersParameterName), |
| 247 | IdentifierName(PropertyNameForParameter(p)))) |
| 248 | .WithNameColon(NameColon(ParameterNameForParameter(p)))).ToArray()))); |
| 249 | } |
| 250 | |
| 251 | private (MemberDeclarationSyntax tvpGeneratorClass, MemberDeclarationSyntax tvpHolderStruct) CreateTvpGeneratorTypes(CreateProcedureStatement node, string procedureName) |
| 252 | { |
| 253 | List<(string parameterName, string rowStructName)> rowTypes = node.Parameters |
| 254 | .Where(p => !TryGetSqlDbTypeForParameter(p, out _)) |
| 255 | .Select(p => (parameterName: PropertyNameForParameter(p), rowStructName: GetRowStructNameForTableType(p.DataType.Name))) |
| 256 | .ToList(); |
| 257 | |
| 258 | if (rowTypes.Count == 0) |
| 259 | { |
| 260 | // no table-valued parameters on this procedure |
| 261 | return (IncompleteMember(), IncompleteMember()); |
| 262 | } |
| 263 | |
| 264 | var holderStructName = TableValuedParametersStructName(procedureName); |
| 265 | |
| 266 | // create a struct with properties for each table-valued parameter |
| 267 | |
| 268 | var structDeclaration = StructDeclaration(holderStructName) |
| 269 | .AddModifiers(Token(SyntaxKind.InternalKeyword)) |
| 270 | |
| 271 | // Add a constructor with parameters for each column, setting the associated property for each column. |
| 272 | .AddMembers( |
| 273 | ConstructorDeclaration( |
| 274 | Identifier(holderStructName)) |
| 275 | .WithModifiers( |
| 276 | TokenList( |
| 277 | Token(SyntaxKind.InternalKeyword))) |
| 278 | .AddParameterListParameters( |
| 279 | rowTypes.Select(p => |
| 280 | Parameter(Identifier(p.parameterName)) |
| 281 | .WithType(TypeExtensions.CreateGenericTypeFromGenericTypeDefinition( |
| 282 | typeof(IEnumerable<>).ToTypeSyntax(true), |
| 283 | IdentifierName(p.rowStructName)))).ToArray()) |
| 284 | .WithBody( |
| 285 | Block(rowTypes.Select(p => |
| 286 | ExpressionStatement( |
| 287 | AssignmentExpression( |
| 288 | SyntaxKind.SimpleAssignmentExpression, |
| 289 | left: MemberAccessExpression( |
| 290 | SyntaxKind.SimpleMemberAccessExpression, |
| 291 | ThisExpression(), |
| 292 | IdentifierName(p.parameterName)), |
| 293 | right: IdentifierName(p.parameterName))))))) |
| 294 | |
| 295 | // Add a property for each column |
| 296 | .AddMembers(rowTypes.Select(p => |
| 297 | (MemberDeclarationSyntax)PropertyDeclaration( |
| 298 | TypeExtensions.CreateGenericTypeFromGenericTypeDefinition( |
| 299 | typeof(IEnumerable<>).ToTypeSyntax(true), |
| 300 | IdentifierName(p.rowStructName)), |
| 301 | Identifier(p.parameterName)) |
| 302 | .AddModifiers(Token(SyntaxKind.InternalKeyword)) |
| 303 | .AddAccessorListAccessors(AccessorDeclaration(SyntaxKind.GetAccessorDeclaration) |
| 304 | .WithSemicolonToken(Token(SyntaxKind.SemicolonToken)))).ToArray()); |
| 305 | |
| 306 | string className = $"{procedureName}TvpGenerator"; |
| 307 | |
| 308 | List<string> distinctTvpTypeNames = rowTypes.Select(r => r.rowStructName).Distinct().ToList(); |
| 309 | |
| 310 | var classDeclaration = ClassDeclaration(className) |
| 311 | .AddTypeParameterListParameters(TypeParameter(TvpGeneratorGenericTypeName)) |
| 312 | .AddBaseListTypes( |
| 313 | SimpleBaseType( |
| 314 | GenericName("IStoredProcedureTableValuedParametersGenerator") |
| 315 | .AddTypeArgumentListArguments( |
| 316 | IdentifierName(TvpGeneratorGenericTypeName), |
| 317 | IdentifierName(holderStructName)))) |
| 318 | .AddModifiers(Token(SyntaxKind.InternalKeyword)) |
| 319 | .AddMembers( |
| 320 | ConstructorDeclaration(Identifier(className)) |
| 321 | .WithModifiers(TokenList(Token(SyntaxKind.PublicKeyword))) |
| 322 | .AddParameterListParameters(distinctTvpTypeNames |
| 323 | .Select(t => |
| 324 | Parameter(Identifier(GeneratorFieldName(t))) |
| 325 | .WithType(GeneratorType(t))).ToArray()) |
| 326 | .WithBody( |
| 327 | Block(distinctTvpTypeNames |
| 328 | .Select(t => |
| 329 | ExpressionStatement( |
| 330 | AssignmentExpression( |
| 331 | SyntaxKind.SimpleAssignmentExpression, |
| 332 | left: MemberAccessExpression( |
| 333 | SyntaxKind.SimpleMemberAccessExpression, |
| 334 | ThisExpression(), |
| 335 | IdentifierName(GeneratorFieldName(t))), |
| 336 | right: IdentifierName(GeneratorFieldName(t)))))))) |
| 337 | .AddMembers( |
| 338 | distinctTvpTypeNames |
| 339 | .Select(t => (MemberDeclarationSyntax)FieldDeclaration( |
| 340 | VariableDeclaration(GeneratorType(t)) |
| 341 | .AddVariables(VariableDeclarator(GeneratorFieldName(t)))) |
| 342 | .AddModifiers(Token(SyntaxKind.PrivateKeyword), Token(SyntaxKind.ReadOnlyKeyword))).ToArray()) |
| 343 | .AddMembers( |
| 344 | MethodDeclaration(IdentifierName(holderStructName), Identifier("Generate")) |
| 345 | .AddModifiers(Token(SyntaxKind.PublicKeyword)) |
| 346 | .AddParameterListParameters(Parameter(Identifier("input")).WithType(IdentifierName(TvpGeneratorGenericTypeName))) |
| 347 | .AddBodyStatements( |
| 348 | ReturnStatement( |
| 349 | ObjectCreationExpression(IdentifierName(holderStructName)) |
| 350 | .AddArgumentListArguments( |
| 351 | rowTypes.Select(p => Argument( |
| 352 | InvocationExpression( |
| 353 | MemberAccessExpression( |
| 354 | SyntaxKind.SimpleMemberAccessExpression, |
| 355 | IdentifierName(GeneratorFieldName(p.rowStructName)), |
| 356 | IdentifierName("GenerateRows"))) |
| 357 | .AddArgumentListArguments(Argument(IdentifierName("input"))))).ToArray())))); |
| 358 | |
| 359 | return (classDeclaration, structDeclaration); |
| 360 | } |
| 361 | |
| 362 | private static string TableValuedParametersStructName(string procedureName) |
| 363 | { |
| 364 | return $"{procedureName}TableValuedParameters"; |
| 365 | } |
| 366 | |
| 367 | private static string GeneratorFieldName(string tableTypeName) |
| 368 | { |
| 369 | return $"{tableTypeName}Generator"; |
| 370 | } |
| 371 | |
| 372 | private static TypeSyntax GeneratorType(string rowStructName) |
| 373 | { |
| 374 | return GenericName("ITableValuedParameterRowGenerator") |
| 375 | .AddTypeArgumentListArguments( |
| 376 | IdentifierName(TvpGeneratorGenericTypeName), |
| 377 | IdentifierName(rowStructName)); |
| 378 | } |
| 379 | |
| 380 | private static bool TryGetSqlDbTypeForParameter(ProcedureParameter parameter, out SqlDbType sqlDbType) |
| 381 | { |
| 382 | return Enum.TryParse<SqlDbType>(parameter.DataType.Name.BaseIdentifier.Value, ignoreCase: true, out sqlDbType); |
| 383 | } |
| 384 | |
| 385 | private static string ParameterNameForParameter(ProcedureParameter parameter) |
| 386 | { |
| 387 | return parameter.VariableName.Value.Substring(1); |
| 388 | } |
| 389 | |
| 390 | private static string FieldNameForParameter(ProcedureParameter parameter) |
| 391 | { |
| 392 | return $"_{parameter.VariableName.Value.Substring(1)}"; |
| 393 | } |
| 394 | |
| 395 | private static string PropertyNameForParameter(ProcedureParameter parameter) |
| 396 | { |
| 397 | string parameterName = parameter.VariableName.Value; |
| 398 | return $"{char.ToUpperInvariant(parameterName[1])}{(parameterName.Length > 2 ? parameterName.Substring(2) : string.Empty)}"; |
| 399 | } |
| 400 | } |
| 401 | } |
| 402 | |