Skip to content

Commit cf1b822

Browse files
feat: support set data type in mysql batch insert
1 parent 78e8c01 commit cf1b822

20 files changed

Lines changed: 242 additions & 104 deletions

File tree

CodegenTests/CodegenUtilsTests.cs

Lines changed: 0 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -145,9 +145,6 @@ public void TestMysqlCopyFromGenerateUtilsMembers()
145145
var expected = new HashSet<string>
146146
{
147147
MySqlConnectorDriver.NullToStringCsvConverter,
148-
MySqlConnectorDriver.BoolToBitCsvConverter,
149-
MySqlConnectorDriver.ByteCsvConverter,
150-
MySqlConnectorDriver.ByteArrayCsvConverter
151148
};
152149
var actual = members
153150
.FindAll(m => m is ClassDeclarationSyntax)

Drivers/DbDriver.cs

Lines changed: 12 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -59,15 +59,13 @@ public abstract class DbDriver
5959
public abstract Dictionary<string, ColumnMapping> ColumnMappings { get; }
6060

6161
protected const string JsonElementTypeHandler =
62-
"""
63-
public class JsonElementTypeHandler : SqlMapper.TypeHandler<JsonElement>
62+
"""
63+
private class JsonElementTypeHandler : SqlMapper.TypeHandler<JsonElement>
6464
{
6565
public override JsonElement Parse(object value)
6666
{
6767
if (value is string s)
6868
return JsonDocument.Parse(s).RootElement;
69-
if (value is null)
70-
return default;
7169
throw new DataException($"Cannot convert {value?.GetType()} to JsonElement");
7270
}
7371
@@ -76,7 +74,7 @@ public override void SetValue(IDbDataParameter parameter, JsonElement value)
7674
parameter.Value = value.GetRawText();
7775
}
7876
}
79-
""";
77+
""";
8078

8179
protected const string TransformQueryForSliceArgsImpl = """
8280
public static string TransformQueryForSliceArgs(string originalSql, int sliceSize, string paramName)
@@ -204,9 +202,15 @@ public static void ConfigureSqlMapper()
204202

205203
protected bool TypeExistsInQueries(string csharpType)
206204
{
207-
return Queries
208-
.SelectMany(query => query.Columns)
209-
.Any(column => csharpType == GetCsharpTypeWithoutNullableSuffix(column, null));
205+
return Queries.Any(q => TypeExistsInQuery(csharpType, q));
206+
}
207+
208+
protected bool TypeExistsInQuery(string csharpType, Query query)
209+
{
210+
return query.Columns
211+
.Any(column => csharpType == GetCsharpTypeWithoutNullableSuffix(column, query)) ||
212+
query.Params
213+
.Any(p => csharpType == GetCsharpTypeWithoutNullableSuffix(p.Column, query));
210214
}
211215

212216
public string AddNullableSuffixIfNeeded(string csharpType, bool notNull)

Drivers/MySqlConnectorDriver.cs

Lines changed: 93 additions & 39 deletions
Original file line numberDiff line numberDiff line change
@@ -146,22 +146,22 @@ public partial class MySqlConnectorDriver(
146146
public override string TransactionClassName => "MySqlTransaction";
147147

148148
private readonly Func<string, string> _setTypeHandlerFunc = x =>
149-
$$"""
150-
public class {{x}}TypeHandler : SqlMapper.TypeHandler<{{x}}[]>
151-
{
152-
public override {{x}}[] Parse(object value)
153-
{
154-
if (value is string s)
155-
return s.To{{x}}Arr();
156-
throw new DataException($"Cannot convert {value?.GetType()} to {{x}}[]");
157-
}
158-
159-
public override void SetValue(IDbDataParameter parameter, {{x}}[] value)
160-
{
161-
parameter.Value = string.Join(",", value);
162-
}
163-
}
164-
""";
149+
$$"""
150+
private class {{x}}TypeHandler : SqlMapper.TypeHandler<{{x}}[]>
151+
{
152+
public override {{x}}[] Parse(object value)
153+
{
154+
if (value is string s)
155+
return s.To{{x}}Arr();
156+
throw new DataException($"Cannot convert {value?.GetType()} to {{x}}[]");
157+
}
158+
159+
public override void SetValue(IDbDataParameter parameter, {{x}}[] value)
160+
{
161+
parameter.Value = string.Join(",", value);
162+
}
163+
}
164+
""";
165165

166166
public override ISet<string> GetUsingDirectivesForQueries()
167167
{
@@ -211,15 +211,17 @@ public override ISet<string> GetUsingDirectivesForUtils()
211211
);
212212
}
213213

214+
private bool IsSetType(Column column)
215+
{
216+
var enumType = GetEnumType(column);
217+
return enumType is not null && IsEnumOfTypeSet(column, enumType);
218+
}
219+
214220
protected override ISet<string> GetConfigureSqlMappings()
215221
{
216222
var setSqlMappings = Queries
217223
.SelectMany(q => q.Columns)
218-
.Where(c =>
219-
{
220-
var enumType = GetEnumType(c);
221-
return enumType is not null && IsEnumOfTypeSet(c, enumType);
222-
})
224+
.Where(IsSetType)
223225
.Select(c =>
224226
{
225227
var enumName = c.Type.Name.ToModelName(GetColumnSchema(c), DefaultSchema);
@@ -255,11 +257,24 @@ public override MemberDeclarationSyntax[] GetMemberDeclarationsForUtils()
255257
.AddRangeIf(GetSetTypeHandlers(), Options.UseDapper);
256258

257259
if (!CopyFromQueryExists())
258-
return [.. memberDeclarations];
260+
return memberDeclarations.ToArray();
259261

260-
var csvConverters = new List<MemberDeclarationSyntax>
262+
foreach (var query in Queries)
261263
{
262-
ParseMemberDeclaration($$"""
264+
if (query.Cmd != ":copyfrom")
265+
continue;
266+
foreach (var p in query.Params)
267+
{
268+
if (!IsSetType(p.Column))
269+
continue;
270+
var enumName = p.Column.Type.Name.ToModelName(GetColumnSchema(p.Column), DefaultSchema);
271+
memberDeclarations = memberDeclarations.AddRangeExcludeNulls([ParseMemberDeclaration(SetCsvConverterFunc(enumName))!]);
272+
}
273+
}
274+
275+
return memberDeclarations
276+
.AddRangeIf([
277+
ParseMemberDeclaration($$"""
263278
public class {{NullToStringCsvConverter}} : DefaultTypeConverter
264279
{
265280
public override {{AddNullableSuffixIfNeeded("string", false)}} ConvertToString(
@@ -269,6 +284,8 @@ public class {{NullToStringCsvConverter}} : DefaultTypeConverter
269284
}
270285
}
271286
""")!,
287+
], CopyFromQueryExists())
288+
.AddRangeIf([
272289
ParseMemberDeclaration($$"""
273290
public class BoolToBitCsvConverter : DefaultTypeConverter
274291
{
@@ -287,7 +304,9 @@ public class BoolToBitCsvConverter : DefaultTypeConverter
287304
}
288305
}
289306
""")!,
290-
ParseMemberDeclaration($$"""
307+
], CopyFromQueryExists() && TypeExistsInQueries("bool"))
308+
.AddRangeIf([
309+
ParseMemberDeclaration($$"""
291310
public class ByteCsvConverter : DefaultTypeConverter
292311
{
293312
public override {{AddNullableSuffixIfNeeded("string", false)}} ConvertToString(
@@ -301,7 +320,9 @@ public class ByteCsvConverter : DefaultTypeConverter
301320
}
302321
}
303322
""")!,
304-
ParseMemberDeclaration($$"""
323+
], CopyFromQueryExists() && TypeExistsInQueries("byte"))
324+
.AddRangeIf([
325+
ParseMemberDeclaration($$"""
305326
public class ByteArrayCsvConverter : DefaultTypeConverter
306327
{
307328
public override {{AddNullableSuffixIfNeeded("string", false)}} ConvertToString(
@@ -314,10 +335,25 @@ public class ByteArrayCsvConverter : DefaultTypeConverter
314335
return base.ConvertToString(value, row, memberMapData);
315336
}
316337
}
317-
""")!
318-
};
338+
""")!,
339+
], CopyFromQueryExists() && TypeExistsInQueries("byte[]"))
340+
.ToArray();
319341

320-
return [.. memberDeclarations, .. csvConverters];
342+
string SetCsvConverterFunc(string x) =>
343+
$$"""
344+
public class {{x}}CsvConverter : DefaultTypeConverter
345+
{
346+
public override {{AddNullableSuffixIfNeeded("string", false)}} ConvertToString(
347+
{{AddNullableSuffixIfNeeded("object", false)}} value, IWriterRow row, MemberMapData memberMapData)
348+
{
349+
if (value == null)
350+
return @"\N";
351+
if (value is {{x}}[] arrVal)
352+
return string.Join(",", arrVal);
353+
return base.ConvertToString(value, row, memberMapData);
354+
}
355+
}
356+
""";
321357
}
322358

323359
public override ConnectionGenCommands EstablishConnection(Query query)
@@ -419,7 +455,7 @@ public string GetCopyFromImpl(Query query, string queryTextConstant)
419455
var {{optionsVar}} = new TypeConverterOptions { Formats = new[] { supportedDateTimeFormat } };
420456
{{csvWriterVar}}.Context.TypeConverterOptionsCache.AddOptions<DateTime>({{optionsVar}});
421457
{{csvWriterVar}}.Context.TypeConverterOptionsCache.AddOptions<DateTime?>({{optionsVar}});
422-
{{GetBoolAndByteConverters().JoinByNewLine()}}
458+
{{GetBoolAndByteConverters(query).JoinByNewLine()}}
423459
{{GetCsvNullConverters(query).JoinByNewLine()}}
424460
await {{csvWriterVar}}.WriteRecordsAsync({{Variable.Args.AsVarName()}});
425461
}
@@ -459,7 +495,10 @@ private ISet<string> GetCsvNullConverters(Query query)
459495
foreach (var p in query.Params)
460496
{
461497
var csharpType = GetCsharpTypeWithoutNullableSuffix(p.Column, query);
462-
if (!BoolAndByteTypes.Contains(csharpType) && TypeExistsInQueries(csharpType))
498+
if (
499+
!BoolAndByteTypes.Contains(csharpType) &&
500+
!IsSetType(p.Column) &&
501+
TypeExistsInQuery(csharpType, query))
463502
{
464503
var nullableCsharpType = AddNullableSuffixIfNeeded(csharpType, false);
465504
converters.Add($"{Variable.CsvWriter.AsVarName()}.Context.TypeConverterCache.AddConverter<{nullableCsharpType}>({nullConverterFn});");
@@ -468,7 +507,7 @@ private ISet<string> GetCsvNullConverters(Query query)
468507
return converters;
469508
}
470509

471-
private ISet<string> GetBoolAndByteConverters()
510+
private ISet<string> GetBoolAndByteConverters(Query query)
472511
{
473512
var csvWriterVar = Variable.CsvWriter.AsVarName();
474513
return new HashSet<string>()
@@ -477,34 +516,49 @@ private ISet<string> GetBoolAndByteConverters()
477516
$"{csvWriterVar}.Context.TypeConverterCache.AddConverter<{AddNullableSuffixIfNeeded("bool", true)}>(new Utils.{BoolToBitCsvConverter}());",
478517
$"{csvWriterVar}.Context.TypeConverterCache.AddConverter<{AddNullableSuffixIfNeeded("bool", false)}>(new Utils.{BoolToBitCsvConverter}());"
479518
],
480-
TypeExistsInQueries("bool")
519+
TypeExistsInQuery("bool", query)
481520
)
482521
.AddRangeIf(
483522
[
484523
$"{csvWriterVar}.Context.TypeConverterCache.AddConverter<{AddNullableSuffixIfNeeded("byte", true)}>(new Utils.{ByteCsvConverter}());",
485524
$"{csvWriterVar}.Context.TypeConverterCache.AddConverter<{AddNullableSuffixIfNeeded("byte", false)}>(new Utils.{ByteCsvConverter}());",
486525
],
487-
TypeExistsInQueries("byte")
526+
TypeExistsInQuery("byte", query)
488527
)
489528
.AddRangeIf(
490529
[
491530
$"{csvWriterVar}.Context.TypeConverterCache.AddConverter<{AddNullableSuffixIfNeeded("byte[]", true)}>(new Utils.{ByteArrayCsvConverter}());",
492531
$"{csvWriterVar}.Context.TypeConverterCache.AddConverter<{AddNullableSuffixIfNeeded("byte[]", false)}>(new Utils.{ByteArrayCsvConverter}());",
493532
],
494-
TypeExistsInQueries("byte[]")
495-
);
533+
TypeExistsInQuery("byte[]", query)
534+
)
535+
.AddRangeExcludeNulls(GetSetConverters(query));
536+
}
537+
538+
private ISet<string> GetSetConverters(Query query)
539+
{
540+
var converters = new HashSet<string>();
541+
foreach (var p in query.Params)
542+
{
543+
if (!IsSetType(p.Column))
544+
continue;
545+
546+
var enumName = p.Column.Type.Name.ToModelName(GetColumnSchema(p.Column), DefaultSchema);
547+
var csvWriterVar = Variable.CsvWriter.AsVarName();
548+
converters.Add($"{csvWriterVar}.Context.TypeConverterCache.AddConverter<{AddNullableSuffixIfNeeded($"{enumName}[]", true)}>(new Utils.{enumName}CsvConverter());");
549+
converters.Add($"{csvWriterVar}.Context.TypeConverterCache.AddConverter<{AddNullableSuffixIfNeeded($"{enumName}[]", false)}>(new Utils.{enumName}CsvConverter());");
550+
}
551+
return converters;
496552
}
497553

498-
private bool IsEnumOfTypeSet(Column column, Plugin.Enum enumType)
554+
private static bool IsEnumOfTypeSet(Column column, Plugin.Enum enumType)
499555
{
500556
return column.Length > enumType.Vals.Select(v => v.Length).Sum();
501557
}
502558

503559
public override string GetEnumTypeAsCsharpType(Column column, Plugin.Enum enumType)
504560
{
505561
var enumName = column.Type.Name.ToModelName(GetColumnSchema(column), DefaultSchema);
506-
if (this is not MySqlConnectorDriver)
507-
return enumName;
508562
return IsEnumOfTypeSet(column, enumType) ? $"{enumName}[]" : enumName;
509563
}
510564
}

Drivers/Variable.cs

Lines changed: 1 addition & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -14,6 +14,7 @@ public enum Variable
1414
Loader,
1515
CsvWriter,
1616
NullConverterFn,
17+
SetConverterFn,
1718

1819
Args,
1920
QueryParams,

docs/05_MySql.md

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -70,7 +70,7 @@ we consider support for the different data types separately for batch inserts an
7070
| mediumblob |||
7171
| longblob |||
7272
| enum |||
73-
| set || |
73+
| set || |
7474
| json |||
7575
| geometry |||
7676
| point |||

end2end/EndToEndScaffold/Templates/MySqlTests.cs

Lines changed: 19 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -602,25 +602,37 @@ public async Task TestMySqlTransactionRollback()
602602
{
603603
Impl = $$"""
604604
[Test]
605-
[TestCase(100, MysqlTypesCEnum.Big)]
606-
[TestCase(500, MysqlTypesCEnum.Small)]
607-
[TestCase(10, null)]
608-
public async Task TestCopyFrom(int batchSize, MysqlTypesCEnum? cEnum)
605+
[TestCase(100, MysqlTypesCEnum.Big, new[] { MysqlTypesCSet.Tea, MysqlTypesCSet.Coffee })]
606+
[TestCase(500, MysqlTypesCEnum.Small, new[] { MysqlTypesCSet.Milk })]
607+
[TestCase(10, null, null)]
608+
public async Task TestCopyFrom(
609+
int batchSize,
610+
MysqlTypesCEnum? cEnum,
611+
MysqlTypesCSet[] cSet)
609612
{
610613
var batchArgs = Enumerable.Range(0, batchSize)
611614
.Select(_ => new QuerySql.InsertMysqlTypesBatchArgs
612615
{
613-
CEnum = cEnum
616+
CEnum = cEnum,
617+
CSet = cSet
614618
})
615619
.ToList();
616620
await QuerySql.InsertMysqlTypesBatch(batchArgs);
617621
var expected = new QuerySql.GetMysqlTypesCntRow
618622
{
619623
Cnt = batchSize,
620-
CEnum = cEnum
624+
CEnum = cEnum,
625+
CSet = cSet
621626
};
622627
var actual = await QuerySql.GetMysqlTypesCnt();
623-
Assert.That(actual{{Consts.UnknownRecordValuePlaceholder}}.CEnum, Is.EqualTo(expected.CEnum));
628+
AssertSingularEquals(expected, actual{{Consts.UnknownRecordValuePlaceholder}});
629+
630+
void AssertSingularEquals(QuerySql.GetMysqlTypesCntRow x, QuerySql.GetMysqlTypesCntRow y)
631+
{
632+
Assert.That(x.Cnt, Is.EqualTo(y.Cnt));
633+
Assert.That(x.CEnum, Is.EqualTo(y.CEnum));
634+
Assert.That(x.CSet, Is.EqualTo(y.CSet));
635+
}
624636
}
625637
"""
626638
},

end2end/EndToEndTests/MySqlConnectorDapperTester.generated.cs

Lines changed: 14 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -744,20 +744,27 @@ public async Task TestBinaryCopyFrom(int batchSize, byte? cBit, byte[] cBinary,
744744
}
745745

746746
[Test]
747-
[TestCase(100, MysqlTypesCEnum.Big)]
748-
[TestCase(500, MysqlTypesCEnum.Small)]
749-
[TestCase(10, null)]
750-
public async Task TestCopyFrom(int batchSize, MysqlTypesCEnum? cEnum)
747+
[TestCase(100, MysqlTypesCEnum.Big, new[] { MysqlTypesCSet.Tea, MysqlTypesCSet.Coffee })]
748+
[TestCase(500, MysqlTypesCEnum.Small, new[] { MysqlTypesCSet.Milk })]
749+
[TestCase(10, null, null)]
750+
public async Task TestCopyFrom(int batchSize, MysqlTypesCEnum? cEnum, MysqlTypesCSet[] cSet)
751751
{
752-
var batchArgs = Enumerable.Range(0, batchSize).Select(_ => new QuerySql.InsertMysqlTypesBatchArgs { CEnum = cEnum }).ToList();
752+
var batchArgs = Enumerable.Range(0, batchSize).Select(_ => new QuerySql.InsertMysqlTypesBatchArgs { CEnum = cEnum, CSet = cSet }).ToList();
753753
await QuerySql.InsertMysqlTypesBatch(batchArgs);
754754
var expected = new QuerySql.GetMysqlTypesCntRow
755755
{
756756
Cnt = batchSize,
757-
CEnum = cEnum
757+
CEnum = cEnum,
758+
CSet = cSet
758759
};
759760
var actual = await QuerySql.GetMysqlTypesCnt();
760-
Assert.That(actual.CEnum, Is.EqualTo(expected.CEnum));
761+
AssertSingularEquals(expected, actual);
762+
void AssertSingularEquals(QuerySql.GetMysqlTypesCntRow x, QuerySql.GetMysqlTypesCntRow y)
763+
{
764+
Assert.That(x.Cnt, Is.EqualTo(y.Cnt));
765+
Assert.That(x.CEnum, Is.EqualTo(y.CEnum));
766+
Assert.That(x.CSet, Is.EqualTo(y.CSet));
767+
}
761768
}
762769
}
763770
}

0 commit comments

Comments
 (0)