microsoft/healthcare-shared-components

Public

mirrored from https://github.com/microsoft/healthcare-shared-componentsAvailable

CodeCommitsIssuesPull requestsActionsInsightsSecurity
1.0.0-1.0.79

Branches

Tags

  • No tags available.
0Branches0Tags
Go to file
Add file
Code

Clone

HTTPS

Download ZIP

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
6using System;
7using System.Collections.Generic;
8using System.Data;
9using System.Linq;
10using Microsoft.CodeAnalysis.CSharp;
11using Microsoft.CodeAnalysis.CSharp.Syntax;
12using Microsoft.SqlServer.TransactSql.ScriptDom;
13using static Microsoft.CodeAnalysis.CSharp.SyntaxFactory;
14
15namespace 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