Skip to content

Commit 4763a95

Browse files
feat: add support for CancellationToken in generated methods
- implement CancellationGen to handle optional CancellationToken threading - update DbDriver to expose Cancellation property - modify CommonGen to support cancellation arguments in method parameters and reader calls - update multiple generator types (One, Many, Exec, ExecRows, ExecLastId, CopyFrom) to include CancellationToken in signatures and DB calls - update MySqlConnector, Npgsql, and Sqlite drivers to pass CancellationToken to connection and command execution - add withCancellationToken option to PluginOptions and RawOptions - update README.md with documentation and usage examples for cancellation - add unit tests for CancellationGen and integration tests for generated code across all engines
1 parent 5f6252a commit 4763a95

18 files changed

Lines changed: 437 additions & 42 deletions

Drivers/DbDriver.cs

Lines changed: 4 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
using Microsoft.CodeAnalysis.CSharp.Syntax;
22
using Plugin;
3+
using SqlcGenCsharp.Drivers.Generators;
34
using System;
45
using System.Collections.Generic;
56
using System.Linq;
@@ -42,6 +43,8 @@ public abstract class DbDriver
4243

4344
public Options Options { get; }
4445

46+
public CancellationGen Cancellation => new(Options.WithCancellationToken);
47+
4548
public string DefaultSchema { get; }
4649

4750
public abstract string TransactionClassName { get; }
@@ -296,7 +299,7 @@ public virtual string[] GetLastIdStatement(Query query)
296299
var convertFuncCall = convertFunc(Variable.Result.AsVarName());
297300
return
298301
[
299-
$"var {Variable.Result.AsVarName()} = await {Variable.Command.AsVarName()}.ExecuteScalarAsync();",
302+
$"var {Variable.Result.AsVarName()} = await {Variable.Command.AsVarName()}.ExecuteScalarAsync({Cancellation.Argument()});",
300303
$"return {convertFuncCall};"
301304
];
302305
}
Lines changed: 38 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,38 @@
1+
namespace SqlcGenCsharp.Drivers.Generators;
2+
3+
// Owns every way an optional CancellationToken shows up in generated code.
4+
// When disabled, every fragment is empty / a pass-through, so output is unchanged.
5+
public class CancellationGen(bool enabled)
6+
{
7+
private static string TokenName => Variable.CancellationToken.AsVarName();
8+
9+
// Trailing optional parameter for a generated method signature.
10+
public string MethodParameter()
11+
{
12+
return enabled ? $"CancellationToken {TokenName} = default" : string.Empty;
13+
}
14+
15+
// Trailing optional parameter appended after an existing parameter list, e.g. (List<T> args{TrailingMethodParameter()}).
16+
public string TrailingMethodParameter()
17+
{
18+
return enabled ? $", CancellationToken {TokenName} = default" : string.Empty;
19+
}
20+
21+
// Sole argument for an otherwise parameterless async call, e.g. ReadAsync(Argument()).
22+
public string Argument()
23+
{
24+
return enabled ? TokenName : string.Empty;
25+
}
26+
27+
// Trailing argument appended after existing call arguments, e.g. WriteAsync(value{TrailingArgument()}).
28+
public string TrailingArgument()
29+
{
30+
return enabled ? $", {TokenName}" : string.Empty;
31+
}
32+
33+
// Wraps flat Dapper call arguments in a CommandDefinition carrying the token; pass-through when disabled.
34+
public string WrapDapperArgs(string flatArgs)
35+
{
36+
return enabled ? $"new CommandDefinition({flatArgs}, cancellationToken: {TokenName})" : flatArgs;
37+
}
38+
}

Drivers/Generators/CommonGen.cs

Lines changed: 9 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -6,11 +6,13 @@ namespace SqlcGenCsharp.Drivers.Generators;
66

77
public class CommonGen(DbDriver dbDriver)
88
{
9-
public static string GetMethodParameterList(string argInterface, IEnumerable<Parameter> parameters)
9+
public static string GetMethodParameterList(string argInterface, IEnumerable<Parameter> parameters,
10+
string cancellationParam = "")
1011
{
11-
return $"{(string.IsNullOrEmpty(argInterface) || !parameters.Any()
12+
var argsParam = string.IsNullOrEmpty(argInterface) || !parameters.Any()
1213
? string.Empty
13-
: $"{argInterface} {Variable.Args.AsVarName()}")}";
14+
: $"{argInterface} {Variable.Args.AsVarName()}";
15+
return string.Join(", ", new[] { argsParam, cancellationParam }.Where(s => !string.IsNullOrEmpty(s)));
1416
}
1517

1618
public static string GetDapperArgs(Query query)
@@ -52,14 +54,14 @@ public string ConstructDapperParamsDict(Query query)
5254
""";
5355
}
5456

55-
public static string AwaitReaderRow()
57+
public static string AwaitReaderRow(string cancellationArg)
5658
{
57-
return $"await {Variable.Reader.AsVarName()}.ReadAsync()";
59+
return $"await {Variable.Reader.AsVarName()}.ReadAsync({cancellationArg})";
5860
}
5961

60-
public static string InitDataReader()
62+
public static string InitDataReader(string cancellationArg)
6163
{
62-
return $"var {Variable.Reader.AsVarName()} = await {Variable.Command.AsVarName()}.ExecuteReaderAsync()";
64+
return $"var {Variable.Reader.AsVarName()} = await {Variable.Command.AsVarName()}.ExecuteReaderAsync({cancellationArg})";
6365
}
6466

6567
public static string GetSqlTransformations(Query query, string queryTextConstant)

Drivers/Generators/CopyFromDeclareGen.cs

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -9,7 +9,7 @@ public class CopyFromDeclareGen(DbDriver dbDriver)
99
public MemberDeclarationSyntax Generate(string queryTextConstant, string argInterface, Query query)
1010
{
1111
return ParseMemberDeclaration($$"""
12-
public async Task {{query.Name.ToMethodName(dbDriver.Options.WithAsyncSuffix)}}(List<{{argInterface}}> args)
12+
public async Task {{query.Name.ToMethodName(dbDriver.Options.WithAsyncSuffix)}}(List<{{argInterface}}> args{{dbDriver.Cancellation.TrailingMethodParameter()}})
1313
{
1414
{{((ICopyFrom)dbDriver).GetCopyFromImpl(query, queryTextConstant)}}
1515
}

Drivers/Generators/ExecDeclareGen.cs

Lines changed: 14 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@ public class ExecDeclareGen(DbDriver dbDriver)
1010

1111
public MemberDeclarationSyntax Generate(string queryTextConstant, string argInterface, Query query)
1212
{
13-
var parametersStr = CommonGen.GetMethodParameterList(argInterface, query.Params);
13+
var parametersStr = CommonGen.GetMethodParameterList(argInterface, query.Params, dbDriver.Cancellation.MethodParameter());
1414
return ParseMemberDeclaration($$"""
1515
public async Task {{query.Name.ToMethodName(dbDriver.Options.WithAsyncSuffix)}}({{parametersStr}})
1616
{
@@ -46,9 +46,10 @@ private string GetDapperNoTxBody(string sqlVar, Query query)
4646
{
4747
var connectionCommands = dbDriver.EstablishConnection(query);
4848
var dapperArgs = CommonGen.GetDapperArgs(query);
49+
var callArgs = dbDriver.Cancellation.WrapDapperArgs($"{sqlVar}{dapperArgs}");
4950
return connectionCommands.GetConnectionOrDataSource.WrapBlock(
5051
$"""
51-
await {Variable.Connection.AsVarName()}.ExecuteAsync({sqlVar}{dapperArgs});
52+
await {Variable.Connection.AsVarName()}.ExecuteAsync({callArgs});
5253
return;
5354
"""
5455
);
@@ -58,6 +59,15 @@ private string GetDapperWithTxBody(string sqlVar, Query query)
5859
{
5960
var transactionProperty = Variable.Transaction.AsPropertyName();
6061
var dapperArgs = CommonGen.GetDapperArgs(query);
62+
if (dbDriver.Options.WithCancellationToken)
63+
{
64+
var callArgs = dbDriver.Cancellation.WrapDapperArgs($"{sqlVar}{dapperArgs}, transaction: this.{transactionProperty}");
65+
return $$"""
66+
{{dbDriver.TransactionConnectionNullExcetionThrow}}
67+
await this.{{transactionProperty}}.Connection.ExecuteAsync({{callArgs}});
68+
""";
69+
}
70+
6171
return $$"""
6272
{{dbDriver.TransactionConnectionNullExcetionThrow}}
6373
await this.{{transactionProperty}}.Connection.ExecuteAsync(
@@ -75,7 +85,7 @@ private string GetDriverNoTxBody(string sqlVar, Query query)
7585
{sqlCommands.SetCommandText.AppendSemicolonUnlessEmpty()}
7686
{dbDriver.AddParametersToCommand(query)}
7787
{sqlCommands.PrepareCommand.AppendSemicolonUnlessEmpty()}
78-
await {Variable.Command.AsVarName()}.ExecuteNonQueryAsync();
88+
await {Variable.Command.AsVarName()}.ExecuteNonQueryAsync({dbDriver.Cancellation.Argument()});
7989
"""
8090
);
8191
return connectionCommands.GetConnectionOrDataSource.WrapBlock(
@@ -99,7 +109,7 @@ private string GetDriverWithTxBody(string sqlVar, Query query)
99109
{{commandVar}}.CommandText = {{sqlVar}};
100110
{{commandVar}}.Transaction = this.{{transactionProperty}};
101111
{{dbDriver.AddParametersToCommand(query)}}
102-
await {{commandVar}}.ExecuteNonQueryAsync();
112+
await {{commandVar}}.ExecuteNonQueryAsync({{dbDriver.Cancellation.Argument()}});
103113
}
104114
""";
105115
}

Drivers/Generators/ExecLastIdDeclareGen.cs

Lines changed: 5 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -11,7 +11,7 @@ public class ExecLastIdDeclareGen(DbDriver dbDriver)
1111

1212
public MemberDeclarationSyntax Generate(string queryTextConstant, string argInterface, Query query)
1313
{
14-
var parametersStr = CommonGen.GetMethodParameterList(argInterface, query.Params);
14+
var parametersStr = CommonGen.GetMethodParameterList(argInterface, query.Params, dbDriver.Cancellation.MethodParameter());
1515
return ParseMemberDeclaration($$"""
1616
public async Task<{{dbDriver.GetIdColumnType(query)}}> {{query.Name.ToMethodName(dbDriver.Options.WithAsyncSuffix)}}({{parametersStr}})
1717
{
@@ -48,18 +48,20 @@ private string GetDapperNoTxBody(string sqlVar, Query query)
4848
var connectionCommands = dbDriver.EstablishConnection(query);
4949
var dapperArgs = CommonGen.GetDapperArgs(query);
5050
var idColumnType = dbDriver.GetIdColumnType(query);
51+
var callArgs = dbDriver.Cancellation.WrapDapperArgs($"{sqlVar}{dapperArgs}");
5152
return connectionCommands.GetConnectionOrDataSource.WrapBlock($"""
52-
return await {Variable.Connection.AsVarName()}.QuerySingleAsync<{idColumnType}>({sqlVar}{dapperArgs});
53+
return await {Variable.Connection.AsVarName()}.QuerySingleAsync<{idColumnType}>({callArgs});
5354
""");
5455
}
5556

5657
private string GetDapperWithTxBody(string sqlVar, Query query)
5758
{
5859
var transactionProperty = Variable.Transaction.AsPropertyName();
5960
var dapperArgs = query.Params.Any() ? $", {Variable.QueryParams.AsVarName()}" : string.Empty;
61+
var callArgs = dbDriver.Cancellation.WrapDapperArgs($"{sqlVar}{dapperArgs}, transaction: this.{transactionProperty}");
6062
return $$"""
6163
{{dbDriver.TransactionConnectionNullExcetionThrow}}
62-
return await this.{{transactionProperty}}.Connection.QuerySingleAsync<{{dbDriver.GetIdColumnType(query)}}>({{sqlVar}}{{dapperArgs}}, transaction: this.{{transactionProperty}});
64+
return await this.{{transactionProperty}}.Connection.QuerySingleAsync<{{dbDriver.GetIdColumnType(query)}}>({{callArgs}});
6365
""";
6466
}
6567

Drivers/Generators/ExecRowsDeclareGen.cs

Lines changed: 14 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -10,7 +10,7 @@ public class ExecRowsDeclareGen(DbDriver dbDriver)
1010

1111
public MemberDeclarationSyntax Generate(string queryTextConstant, string argInterface, Query query)
1212
{
13-
var parametersStr = CommonGen.GetMethodParameterList(argInterface, query.Params);
13+
var parametersStr = CommonGen.GetMethodParameterList(argInterface, query.Params, dbDriver.Cancellation.MethodParameter());
1414
return ParseMemberDeclaration($$"""
1515
public async Task<long> {{query.Name.ToMethodName(dbDriver.Options.WithAsyncSuffix)}}({{parametersStr}})
1616
{
@@ -46,15 +46,25 @@ private string GetDapperNoTxBody(string sqlVar, Query query)
4646
{
4747
var connectionCommands = dbDriver.EstablishConnection(query);
4848
var dapperArgs = CommonGen.GetDapperArgs(query);
49+
var callArgs = dbDriver.Cancellation.WrapDapperArgs($"{sqlVar}{dapperArgs}");
4950
return connectionCommands.GetConnectionOrDataSource.WrapBlock(
50-
$"return await {Variable.Connection.AsVarName()}.ExecuteAsync({sqlVar}{dapperArgs});"
51+
$"return await {Variable.Connection.AsVarName()}.ExecuteAsync({callArgs});"
5152
);
5253
}
5354

5455
private string GetDapperWithTxBody(string sqlVar, Query query)
5556
{
5657
var transactionProperty = Variable.Transaction.AsPropertyName();
5758
var dapperArgs = CommonGen.GetDapperArgs(query);
59+
if (dbDriver.Options.WithCancellationToken)
60+
{
61+
var callArgs = dbDriver.Cancellation.WrapDapperArgs($"{sqlVar}{dapperArgs}, transaction: this.{transactionProperty}");
62+
return $$"""
63+
{{dbDriver.TransactionConnectionNullExcetionThrow}}
64+
return await this.{{transactionProperty}}.Connection.ExecuteAsync({{callArgs}});
65+
""";
66+
}
67+
5868
return $$"""
5969
{{dbDriver.TransactionConnectionNullExcetionThrow}}
6070
return await this.{{transactionProperty}}.Connection.ExecuteAsync(
@@ -72,7 +82,7 @@ private string GetDriverNoTxBody(string sqlVar, Query query)
7282
{sqlCommands.SetCommandText.AppendSemicolonUnlessEmpty()}
7383
{dbDriver.AddParametersToCommand(query)}
7484
{sqlCommands.PrepareCommand.AppendSemicolonUnlessEmpty()}
75-
return await {Variable.Command.AsVarName()}.ExecuteNonQueryAsync();
85+
return await {Variable.Command.AsVarName()}.ExecuteNonQueryAsync({dbDriver.Cancellation.Argument()});
7686
"""
7787
);
7888
return connectionCommands.GetConnectionOrDataSource.WrapBlock(
@@ -95,7 +105,7 @@ private string GetDriverWithTxBody(string sqlVar, Query query)
95105
{{commandVar}}.CommandText = {{sqlVar}};
96106
{{commandVar}}.Transaction = this.{{transactionProperty}};
97107
{{dbDriver.AddParametersToCommand(query)}}
98-
return await {{commandVar}}.ExecuteNonQueryAsync();
108+
return await {{commandVar}}.ExecuteNonQueryAsync({{dbDriver.Cancellation.Argument()}});
99109
}
100110
""";
101111
}

Drivers/Generators/ManyDeclareGen.cs

Lines changed: 16 additions & 6 deletions
Original file line numberDiff line numberDiff line change
@@ -12,7 +12,7 @@ public class ManyDeclareGen(DbDriver dbDriver)
1212

1313
public MemberDeclarationSyntax Generate(string queryTextConstant, string argInterface, string returnInterface, Query query)
1414
{
15-
var parametersStr = CommonGen.GetMethodParameterList(argInterface, query.Params);
15+
var parametersStr = CommonGen.GetMethodParameterList(argInterface, query.Params, dbDriver.Cancellation.MethodParameter());
1616
var returnType = $"Task<List<{returnInterface}>>";
1717
return ParseMemberDeclaration($$"""
1818
public async {{returnType}} {{query.Name.ToMethodName(dbDriver.Options.WithAsyncSuffix)}}({{parametersStr}})
@@ -52,9 +52,10 @@ private string GetDapperNoTxBody(string sqlVar, string returnInterface, Query qu
5252
var dapperArgs = CommonGen.GetDapperArgs(query);
5353
var returnType = dbDriver.AddNullableSuffixIfNeeded(returnInterface, true);
5454
var resultVar = Variable.Result.AsVarName();
55+
var callArgs = dbDriver.Cancellation.WrapDapperArgs($"{sqlVar}{dapperArgs}");
5556
return connectionCommands.GetConnectionOrDataSource.WrapBlock(
5657
$"""
57-
var {resultVar} = await {Variable.Connection.AsVarName()}.QueryAsync<{returnType}>({sqlVar}{dapperArgs});
58+
var {resultVar} = await {Variable.Connection.AsVarName()}.QueryAsync<{returnType}>({callArgs});
5859
return {resultVar}.AsList();
5960
"""
6061
);
@@ -66,6 +67,15 @@ private string GetDapperWithTxBody(string sqlVar, string returnInterface, Query
6667
var dapperArgs = CommonGen.GetDapperArgs(query);
6768
var returnType = dbDriver.AddNullableSuffixIfNeeded(returnInterface, true);
6869

70+
if (dbDriver.Options.WithCancellationToken)
71+
{
72+
var callArgs = dbDriver.Cancellation.WrapDapperArgs($"{sqlVar}{dapperArgs}, transaction: this.{transactionProperty}");
73+
return $$"""
74+
{{dbDriver.TransactionConnectionNullExcetionThrow}}
75+
return (await this.{{transactionProperty}}.Connection.QueryAsync<{{returnType}}>({{callArgs}})).AsList();
76+
""";
77+
}
78+
6979
return $$"""
7080
{{dbDriver.TransactionConnectionNullExcetionThrow}}
7181
return (await this.{{transactionProperty}}.Connection.QueryAsync<{{returnType}}>(
@@ -80,7 +90,7 @@ private string GetDriverNoTxBody(string sqlVar, string returnInterface, Query qu
8090
var dataclassInit = CommonGen.InstantiateDataclass([.. query.Columns], returnInterface, query);
8191
var resultVar = Variable.Result.AsVarName();
8292
var readWhileExists = $$"""
83-
while ({{CommonGen.AwaitReaderRow()}})
93+
while ({{CommonGen.AwaitReaderRow(dbDriver.Cancellation.Argument())}})
8494
{{resultVar}}.Add({{dataclassInit}});
8595
""";
8696
var sqlCommands = dbDriver.CreateSqlCommand(sqlVar);
@@ -89,7 +99,7 @@ private string GetDriverNoTxBody(string sqlVar, string returnInterface, Query qu
8999
{{sqlCommands.SetCommandText.AppendSemicolonUnlessEmpty()}}
90100
{{dbDriver.AddParametersToCommand(query)}}
91101
{{sqlCommands.PrepareCommand.AppendSemicolonUnlessEmpty()}}
92-
using ({{CommonGen.InitDataReader()}})
102+
using ({{CommonGen.InitDataReader(dbDriver.Cancellation.Argument())}})
93103
{
94104
var {{resultVar}} = new List<{{returnInterface}}>();
95105
{{readWhileExists}}
@@ -118,10 +128,10 @@ private string GetDriverWithTxBody(string sqlVar, string returnInterface, Query
118128
{{commandVar}}.CommandText = {{sqlVar}};
119129
{{commandVar}}.Transaction = this.{{transactionProperty}};
120130
{{dbDriver.AddParametersToCommand(query)}}
121-
using ({{CommonGen.InitDataReader()}})
131+
using ({{CommonGen.InitDataReader(dbDriver.Cancellation.Argument())}})
122132
{
123133
var {{resultVar}} = new List<{{returnInterface}}>();
124-
while ({{CommonGen.AwaitReaderRow()}})
134+
while ({{CommonGen.AwaitReaderRow(dbDriver.Cancellation.Argument())}})
125135
{{resultVar}}.Add({{CommonGen.InstantiateDataclass([.. query.Columns], returnInterface, query)}});
126136
return {{resultVar}};
127137
}

0 commit comments

Comments
 (0)