Skip to content

Commit 4f353c8

Browse files
feat: support data type override in plugin options
1 parent 02ba667 commit 4f353c8

30 files changed

Lines changed: 681 additions & 50 deletions

File tree

CodeGenerator/Generators/DataClassesGen.cs

Lines changed: 7 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -10,24 +10,24 @@ namespace SqlcGenCsharp.Generators;
1010

1111
internal class DataClassesGen(DbDriver dbDriver)
1212
{
13-
public MemberDeclarationSyntax Generate(string name, ClassMember? classMember, IList<Column> columns, Options options)
13+
public MemberDeclarationSyntax Generate(string name, ClassMember? classMember, IList<Column> columns, Options options, Query? query = null)
1414
{
1515
var className = classMember is null ? name : classMember.Value.Name(name);
1616
if (options.DotnetFramework.IsDotnetCore() && !options.UseDapper)
17-
return GenerateAsRecord(className, columns);
18-
return GenerateAsCLass(className, columns);
17+
return GenerateAsRecord(className, columns, query);
18+
return GenerateAsCLass(className, columns, query);
1919
}
2020

21-
private MemberDeclarationSyntax GenerateAsRecord(string className, IList<Column> columns)
21+
private MemberDeclarationSyntax GenerateAsRecord(string className, IList<Column> columns, Query? query)
2222
{
2323
var seenEmbed = new Dictionary<string, int>();
2424
var recordParameters = columns
25-
.Select(column => $"{dbDriver.GetCsharpType(column)} {GetFieldName(column, seenEmbed)}")
25+
.Select(column => $"{dbDriver.GetCsharpType(column, query)} {GetFieldName(column, seenEmbed)}")
2626
.JoinByComma();
2727
return ParseMemberDeclaration($"public readonly record struct {className} ({recordParameters});")!;
2828
}
2929

30-
private ClassDeclarationSyntax GenerateAsCLass(string className, IList<Column> columns)
30+
private ClassDeclarationSyntax GenerateAsCLass(string className, IList<Column> columns, Query? query)
3131
{
3232
var modernDotnetSupported = dbDriver.Options.DotnetFramework.IsDotnetCore();
3333
return ClassDeclaration(className)
@@ -40,7 +40,7 @@ MemberDeclarationSyntax[] ColumnsToProperties()
4040
var seenEmbed = new Dictionary<string, int>();
4141
return columns.Select(column =>
4242
{
43-
var csharpType = dbDriver.GetCsharpType(column);
43+
var csharpType = dbDriver.GetCsharpType(column, query);
4444
var optionalRequiredModifier = RequiredModifierNeeded(column) ? "required" : string.Empty;
4545
var setterMethod = modernDotnetSupported ? "init" : "set";
4646
return ParseMemberDeclaration(

CodeGenerator/Generators/QueriesGen.cs

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -95,7 +95,7 @@ private IEnumerable<MemberDeclarationSyntax> GetMembersForSingleQuery(Query quer
9595
private MemberDeclarationSyntax? GetQueryParamsDataclass(Query query)
9696
{
9797
if (query.Params.Count <= 0) return null;
98-
var columns = query.Params.Select(dbDriver.GetColumnFromParam).ToList();
98+
var columns = query.Params.Select(p => dbDriver.GetColumnFromParam(p, query)).ToList();
9999
return DataClassesGen.Generate(query.Name, ClassMember.Args, columns, dbDriver.Options);
100100
}
101101

Lines changed: 44 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,44 @@
1+
using Google.Protobuf;
2+
using Plugin;
3+
using SqlcGenCsharp;
4+
using System.Text;
5+
using System.Xml;
6+
7+
namespace CodegenTests;
8+
9+
public class CodegenTypeOverrideTests
10+
{
11+
private readonly Settings _postgresSettings = new()
12+
{
13+
Engine = "postgresql",
14+
Codegen = new Codegen { Out = "DummyProject" }
15+
};
16+
17+
private readonly Catalog _emptyCatalog = new()
18+
{
19+
Schemas =
20+
{
21+
new Schema
22+
{
23+
Name = string.Empty,
24+
Tables = { Capacity = 0 },
25+
Enums = { Capacity = 0 },
26+
}
27+
}
28+
};
29+
30+
private CodeGenerator CodeGenerator { get; } = new();
31+
32+
[Test]
33+
public void TestOverrideQueryColumnDataType()
34+
{
35+
var request = new GenerateRequest
36+
{
37+
Settings = _postgresSettings,
38+
Catalog = _emptyCatalog,
39+
PluginOptions = ByteString.CopyFrom("{\"overrides\":[{\"column\":\"GetPostgresFunctions:max_integer\",\"csharp_type\":{\"type\":\"int\"}},{\"column\":\"GetPostgresFunctions:max_varchar\",\"csharp_type\":{\"type\":\"string\"}},{\"column\":\"GetPostgresFunctions:max_timestamp\",\"csharp_type\":{\"type\":\"DateTime\"}}]}", Encoding.UTF8)
40+
};
41+
42+
var response = CodeGenerator.Generate(request);
43+
}
44+
}

Drivers/DbDriver.cs

Lines changed: 35 additions & 8 deletions
Original file line numberDiff line numberDiff line change
@@ -130,13 +130,13 @@ public string AddNullableSuffixIfNeeded(string csharpType, bool notNull)
130130
return IsTypeNullable(csharpType) ? $"{csharpType}?" : csharpType;
131131
}
132132

133-
public string GetCsharpType(Column column)
133+
public string GetCsharpType(Column column, Query? query)
134134
{
135-
var csharpType = GetCsharpTypeWithoutNullableSuffix(column);
135+
var csharpType = GetCsharpTypeWithoutNullableSuffix(column, query);
136136
return AddNullableSuffixIfNeeded(csharpType, column.NotNull);
137137
}
138138

139-
private string GetCsharpTypeWithoutNullableSuffix(Column column)
139+
private string GetCsharpTypeWithoutNullableSuffix(Column column, Query? query)
140140
{
141141
if (column.EmbedTable != null)
142142
return column.EmbedTable.Name.ToModelName(column.EmbedTable.Schema, DefaultSchema);
@@ -147,6 +147,13 @@ private string GetCsharpTypeWithoutNullableSuffix(Column column)
147147
if (IsEnumType(column))
148148
return column.Type.Name.ToModelName(column.Table.Schema, DefaultSchema);
149149

150+
if (query is not null)
151+
{
152+
var foundOverride = FindOverrideForQueryColumn(query, column);
153+
if (foundOverride is not null)
154+
return foundOverride.CsharpType.Type;
155+
}
156+
150157
foreach (var columnMapping in ColumnMappings
151158
.Where(columnMapping => DoesColumnMappingApply(columnMapping, column)))
152159
{
@@ -177,14 +184,29 @@ private static bool DoesColumnMappingApply(ColumnMapping columnMapping, Column c
177184
return typeInfo.Length.Value == column.Length;
178185
}
179186

180-
public string GetColumnReader(Column column, int ordinal)
187+
private string GetColumnReader(OverrideOption overrideOption, int ordinal)
188+
{
189+
var columnMapping = ColumnMappings.Find(c => c.CsharpType == overrideOption.CsharpType.Type);
190+
if (columnMapping is not null)
191+
return columnMapping.ReaderFn(ordinal);
192+
throw new NotSupportedException($"Column {overrideOption.Column} has unsupported column type: {overrideOption.CsharpType.Type}");
193+
}
194+
195+
public string GetColumnReader(Column column, int ordinal, Query? query)
181196
{
182197
if (IsEnumType(column))
183198
{
184199
var enumName = column.Type.Name.ToModelName(column.Table.Schema, DefaultSchema);
185200
return $"{Variable.Reader.AsVarName()}.GetString({ordinal}).To{enumName}()";
186201
}
187202

203+
if (query is not null)
204+
{
205+
var foundOverride = FindOverrideForQueryColumn(query, column);
206+
if (foundOverride is not null)
207+
return GetColumnReader(foundOverride, ordinal);
208+
}
209+
188210
foreach (var columnMapping in ColumnMappings
189211
.Where(columnMapping => DoesColumnMappingApply(columnMapping, column)))
190212
{
@@ -227,10 +249,10 @@ public string GetIdColumnType(Query query)
227249
var tableColumns = Tables[query.InsertIntoTable.Schema][query.InsertIntoTable.Name].Columns;
228250
var idColumn = tableColumns.First(c => c.Name.Equals("id", StringComparison.OrdinalIgnoreCase));
229251
if (idColumn is not null)
230-
return GetCsharpType(idColumn);
252+
return GetCsharpType(idColumn, query);
231253

232254
idColumn = tableColumns.First(c => c.Name.Contains("id", StringComparison.CurrentCultureIgnoreCase));
233-
return GetCsharpType(idColumn ?? tableColumns[0]);
255+
return GetCsharpType(idColumn ?? tableColumns[0], query);
234256
}
235257

236258
public virtual string[] GetLastIdStatement(Query query)
@@ -243,10 +265,10 @@ public virtual string[] GetLastIdStatement(Query query)
243265
];
244266
}
245267

246-
public Column GetColumnFromParam(Parameter queryParam)
268+
public Column GetColumnFromParam(Parameter queryParam, Query query)
247269
{
248270
if (string.IsNullOrEmpty(queryParam.Column.Name))
249-
queryParam.Column.Name = $"{GetCsharpType(queryParam.Column).Replace("[]", "Arr")}_{queryParam.Number}";
271+
queryParam.Column.Name = $"{GetCsharpType(queryParam.Column, query).Replace("[]", "Arr")}_{queryParam.Number}";
250272
return queryParam.Column;
251273
}
252274

@@ -262,4 +284,9 @@ protected bool BatchQueryExists()
262284
{
263285
return Queries.Any(q => q.Cmd is ":copyfrom");
264286
}
287+
288+
private OverrideOption? FindOverrideForQueryColumn(Query query, Column column)
289+
{
290+
return Options.Overrides.FirstOrDefault(o => o.Column.Equals($"{query.Name}:{column.Name}"));
291+
}
265292
}

Drivers/Generators/CommonGen.cs

Lines changed: 12 additions & 12 deletions
Original file line numberDiff line numberDiff line change
@@ -92,7 +92,7 @@ public static string GetSqlTransformations(Query query, string queryTextConstant
9292
""";
9393
}
9494

95-
public string InstantiateDataclass(Column[] columns, string returnInterface)
95+
public string InstantiateDataclass(Column[] columns, string returnInterface, Query? query)
9696
{
9797
var columnsInit = new List<string>();
9898
var actualOrdinal = 0;
@@ -102,7 +102,7 @@ public string InstantiateDataclass(Column[] columns, string returnInterface)
102102
{
103103
if (column.EmbedTable is null)
104104
{
105-
columnsInit.Add(GetAsSimpleAssignment(column, actualOrdinal));
105+
columnsInit.Add(GetAsSimpleAssignment(column, actualOrdinal, query));
106106
actualOrdinal++;
107107
continue;
108108
}
@@ -113,38 +113,38 @@ public string InstantiateDataclass(Column[] columns, string returnInterface)
113113
seenEmbed.TryAdd(tableFieldType, 1);
114114
seenEmbed[tableFieldType]++;
115115

116-
var tableColumnsInit = GetAsEmbeddedTableColumnAssignment(column, actualOrdinal);
116+
var tableColumnsInit = GetAsEmbeddedTableColumnAssignment(column, actualOrdinal, query);
117117
columnsInit.Add($"{tableFieldName} = {InstantiateDataclassInternal(tableFieldType, tableColumnsInit)}");
118118
actualOrdinal += tableColumnsInit.Length;
119119
}
120120

121121
return InstantiateDataclassInternal(returnInterface, columnsInit);
122122

123-
string[] GetAsEmbeddedTableColumnAssignment(Column tableColumn, int ordinal)
123+
string[] GetAsEmbeddedTableColumnAssignment(Column tableColumn, int ordinal, Query? query)
124124
{
125125
var schemaName = tableColumn.EmbedTable.Schema == dbDriver.DefaultSchema ? string.Empty : tableColumn.EmbedTable.Schema;
126126
var tableColumns = dbDriver.Tables[schemaName][tableColumn.EmbedTable.Name].Columns;
127127
return tableColumns
128-
.Select((c, o) => GetAsSimpleAssignment(c, o + ordinal))
128+
.Select((c, o) => GetAsSimpleAssignment(c, o + ordinal, query))
129129
.ToArray();
130130
}
131131

132-
string GetAsSimpleAssignment(Column column, int ordinal)
132+
string GetAsSimpleAssignment(Column column, int ordinal, Query query)
133133
{
134-
var readExpression = GetReadExpression(column, ordinal);
134+
var readExpression = GetReadExpression(column, ordinal, query);
135135
return $"{column.Name.ToPascalCase()} = {readExpression}";
136136
}
137137

138-
string GetReadExpression(Column column, int ordinal)
138+
string GetReadExpression(Column column, int ordinal, Query query)
139139
{
140140
return column.NotNull
141-
? dbDriver.GetColumnReader(column, ordinal)
142-
: $"{CheckNullExpression(ordinal)} ? {GetNullExpression(column)} : {dbDriver.GetColumnReader(column, ordinal)}";
141+
? dbDriver.GetColumnReader(column, ordinal, query)
142+
: $"{CheckNullExpression(ordinal)} ? {GetNullExpression(column, query)} : {dbDriver.GetColumnReader(column, ordinal, query)}";
143143
}
144144

145-
string GetNullExpression(Column column)
145+
string GetNullExpression(Column column, Query? query)
146146
{
147-
var csharpType = dbDriver.GetCsharpType(column);
147+
var csharpType = dbDriver.GetCsharpType(column, query);
148148
if (dbDriver.Options.DotnetFramework.IsDotnetCore()) return "null";
149149
return dbDriver.IsTypeNullable(csharpType) ? $"({csharpType}) null" : "null";
150150
}

Drivers/Generators/ManyDeclareGen.cs

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -54,7 +54,7 @@ string GetAsDriver()
5454
var commandParameters = CommonGen.AddParametersToCommand(query.Params);
5555
var initDataReader = CommonGen.InitDataReader();
5656
var awaitReaderRow = CommonGen.AwaitReaderRow();
57-
var dataclassInit = CommonGen.InstantiateDataclass(query.Columns.ToArray(), returnInterface);
57+
var dataclassInit = CommonGen.InstantiateDataclass(query.Columns.ToArray(), returnInterface, query);
5858
var readWhileExists = $$"""
5959
while ({{awaitReaderRow}})
6060
{

Drivers/Generators/OneDeclareGen.cs

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -54,7 +54,7 @@ string GetAsDriver()
5454
var commandParameters = CommonGen.AddParametersToCommand(query.Params);
5555
var initDataReader = CommonGen.InitDataReader();
5656
var awaitReaderRow = CommonGen.AwaitReaderRow();
57-
var returnDataclass = CommonGen.InstantiateDataclass(query.Columns.ToArray(), returnInterface);
57+
var returnDataclass = CommonGen.InstantiateDataclass(query.Columns.ToArray(), returnInterface, query);
5858
return $$"""
5959
using ({{establishConnection}})
6060
{

Drivers/MySqlConnectorDriver.cs

Lines changed: 7 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -92,7 +92,13 @@ public partial class MySqlConnectorDriver(
9292
new Dictionary<string, DbTypeInfo>
9393
{
9494
{ "decimal", new DbTypeInfo() }
95-
}, ordinal => $"reader.GetDecimal({ordinal})")
95+
}, ordinal => $"reader.GetDecimal({ordinal})"),
96+
// last item in the dictionary - enforce TODO
97+
new("object",
98+
new Dictionary<string, DbTypeInfo>
99+
{
100+
{ "any", new DbTypeInfo() }
101+
}, ordinal => $"reader.GetValue({ordinal})")
96102
];
97103

98104
public override UsingDirectiveSyntax[] GetUsingDirectivesForQueries()

Drivers/NpgsqlDriver.cs

Lines changed: 8 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -150,7 +150,13 @@ public NpgsqlDriver(
150150
new Dictionary<string, DbTypeInfo>
151151
{
152152
{ "circle", new DbTypeInfo(NpgsqlTypeOverride: "NpgsqlDbType.Circle") }
153-
}, ordinal => $"reader.GetFieldValue<NpgsqlCircle>({ordinal})")
153+
}, ordinal => $"reader.GetFieldValue<NpgsqlCircle>({ordinal})"),
154+
// last item in the dictionary - enforce TODO
155+
new("object",
156+
new Dictionary<string, DbTypeInfo>
157+
{
158+
{ "anyarray", new DbTypeInfo() }
159+
}, ordinal => $"reader.GetValue({ordinal})")
154160
];
155161

156162
public override UsingDirectiveSyntax[] GetUsingDirectivesForQueries()
@@ -286,7 +292,7 @@ public override string TransformQueryText(Query query)
286292
for (var i = 0; i < query.Params.Count; i++)
287293
{
288294
var currentParameter = query.Params[i];
289-
var column = GetColumnFromParam(currentParameter);
295+
var column = GetColumnFromParam(currentParameter, query);
290296
queryText = Regex.Replace(queryText, $@"\$\s*{i + 1}\b", $"@{column.Name}");
291297
}
292298

Drivers/SqliteDriver.cs

Lines changed: 6 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -37,6 +37,12 @@ public partial class SqliteDriver(
3737
{
3838
{"real", new DbTypeInfo()}
3939
}, ordinal => $"reader.GetDecimal({ordinal})"),
40+
// last item in the dictionary - enforce TODO
41+
new("object",
42+
new Dictionary<string, DbTypeInfo>
43+
{
44+
{ "any", new DbTypeInfo() }
45+
}, ordinal => $"reader.GetValue({ordinal})")
4046
];
4147

4248
public override UsingDirectiveSyntax[] GetUsingDirectivesForQueries()

0 commit comments

Comments
 (0)