Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
2 changes: 1 addition & 1 deletion .github/workflows/build.yml
Original file line number Diff line number Diff line change
Expand Up @@ -126,4 +126,4 @@ jobs:
- name: Verify pushed codegen requests are synced
run: |
dotnet publish LocalRunner -c release --output dist/
sqlc -f sqlc.local.yaml diff
sqlc -f sqlc.requests.yaml diff
6 changes: 6 additions & 0 deletions .pre-commit-config.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,6 @@
repos:
- repo: https://github.com/sqlfluff/sqlfluff
rev: 3.4.1
hooks:
- id: sqlfluff-fix
args: [--FIX-EVEN-UNPARSABLE]
17 changes: 17 additions & 0 deletions .sqlfluff
Original file line number Diff line number Diff line change
@@ -0,0 +1,17 @@
[sqlfluff]
exclude_rules = AM04,AL03,RF02,RF04,AM05,AL01,ST06
dialect = ansi

[sqlfluff:rules]

[sqlfluff:rules:LT02]
capitalisation_policy = upper

[sqlfluff:paths:examples/config/postgresql/]
dialect = postgres

[sqlfluff:paths:examples/config/mysql/]
dialect = mysql

[sqlfluff:paths:examples/config/sqlite/]
dialect = sqlite
16 changes: 8 additions & 8 deletions CodeGenerator/Generators/DataClassesGen.cs
Original file line number Diff line number Diff line change
Expand Up @@ -10,24 +10,24 @@ namespace SqlcGenCsharp.Generators;

internal class DataClassesGen(DbDriver dbDriver)
{
public MemberDeclarationSyntax Generate(string name, ClassMember? classMember, IList<Column> columns, Options options)
public MemberDeclarationSyntax Generate(string name, ClassMember? classMember, IList<Column> columns, Options options, Query? query)
{
var className = classMember is null ? name : classMember.Value.Name(name);
if (options.DotnetFramework.IsDotnetCore() && !options.UseDapper)
return GenerateAsRecord(className, columns);
return GenerateAsCLass(className, columns);
return GenerateAsRecord(className, columns, query);
return GenerateAsCLass(className, columns, query);
}

private MemberDeclarationSyntax GenerateAsRecord(string className, IList<Column> columns)
private MemberDeclarationSyntax GenerateAsRecord(string className, IList<Column> columns, Query? query)
{
var seenEmbed = new Dictionary<string, int>();
var recordParameters = columns
.Select(column => $"{dbDriver.GetCsharpType(column)} {GetFieldName(column, seenEmbed)}")
.Select(column => $"{dbDriver.GetCsharpType(column, query)} {GetFieldName(column, seenEmbed)}")
.JoinByComma();
return ParseMemberDeclaration($"public readonly record struct {className} ({recordParameters});")!;
}

private ClassDeclarationSyntax GenerateAsCLass(string className, IList<Column> columns)
private ClassDeclarationSyntax GenerateAsCLass(string className, IList<Column> columns, Query? query)
{
var modernDotnetSupported = dbDriver.Options.DotnetFramework.IsDotnetCore();
return ClassDeclaration(className)
Expand All @@ -40,7 +40,7 @@ MemberDeclarationSyntax[] ColumnsToProperties()
var seenEmbed = new Dictionary<string, int>();
return columns.Select(column =>
{
var csharpType = dbDriver.GetCsharpType(column);
var csharpType = dbDriver.GetCsharpType(column, query);
var optionalRequiredModifier = RequiredModifierNeeded(column) ? "required" : string.Empty;
var setterMethod = modernDotnetSupported ? "init" : "set";
return ParseMemberDeclaration(
Expand All @@ -58,7 +58,7 @@ bool RequiredModifierNeeded(Column column)
return false;
if (column.EmbedTable != null)
return true;
return column.NotNull;
return dbDriver.IsColumnNotNull(column, query);
}
}

Expand Down
2 changes: 1 addition & 1 deletion CodeGenerator/Generators/ModelsGen.cs
Original file line number Diff line number Diff line change
Expand Up @@ -46,7 +46,7 @@ private MemberDeclarationSyntax[] GenerateDataClasses(Dictionary<string, Diction
from schemaTables in tables
from table in schemaTables.Value
let className = table.Value.Rel.Name.ToModelName(table.Value.Rel.Schema, dbDriver.DefaultSchema)
select DataClassesGen.Generate(className, null, table.Value.Columns, dbDriver.Options)
select DataClassesGen.Generate(className, null, table.Value.Columns, dbDriver.Options, null)
).ToArray();
}

Expand Down
6 changes: 3 additions & 3 deletions CodeGenerator/Generators/QueriesGen.cs
Original file line number Diff line number Diff line change
Expand Up @@ -89,14 +89,14 @@ private IEnumerable<MemberDeclarationSyntax> GetMembersForSingleQuery(Query quer
private MemberDeclarationSyntax? GetQueryColumnsDataclass(Query query)
{
if (query.Columns.Count <= 0) return null;
return DataClassesGen.Generate(query.Name, ClassMember.Row, query.Columns, dbDriver.Options);
return DataClassesGen.Generate(query.Name, ClassMember.Row, query.Columns, dbDriver.Options, query);
}

private MemberDeclarationSyntax? GetQueryParamsDataclass(Query query)
{
if (query.Params.Count <= 0) return null;
var columns = query.Params.Select(dbDriver.GetColumnFromParam).ToList();
return DataClassesGen.Generate(query.Name, ClassMember.Args, columns, dbDriver.Options);
var columns = query.Params.Select(p => dbDriver.GetColumnFromParam(p, query)).ToList();
return DataClassesGen.Generate(query.Name, ClassMember.Args, columns, dbDriver.Options, query);
}

private MemberDeclarationSyntax? GetQueryTextConstant(Query query)
Expand Down
44 changes: 44 additions & 0 deletions CodegenTests/CodegenTypeOverrideTests.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,44 @@
using Google.Protobuf;
using Plugin;
using SqlcGenCsharp;
using System.Text;
using System.Xml;

namespace CodegenTests;

public class CodegenTypeOverrideTests
{
private readonly Settings _postgresSettings = new()
{
Engine = "postgresql",
Codegen = new Codegen { Out = "DummyProject" }
};

private readonly Catalog _emptyCatalog = new()
{
Schemas =
{
new Schema
{
Name = string.Empty,
Tables = { Capacity = 0 },
Enums = { Capacity = 0 },
}
}
};

private CodeGenerator CodeGenerator { get; } = new();

[Test]
public void TestOverrideQueryColumnDataType()
{
var request = new GenerateRequest
{
Settings = _postgresSettings,
Catalog = _emptyCatalog,
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)
};

var response = CodeGenerator.Generate(request);
}
}
2 changes: 1 addition & 1 deletion CodegenTests/test-requests/DefaultSchemaEnum/query.sql
Original file line number Diff line number Diff line change
Expand Up @@ -2,4 +2,4 @@
SELECT * FROM dummy_table LIMIT 1;

-- name: TestInsert :exec
INSERT INTO dummy_table (dummy_column) VALUES (?);
INSERT INTO dummy_table (dummy_column) VALUES (?);
4 changes: 2 additions & 2 deletions CodegenTests/test-requests/DefaultSchemaEnum/schema.sql
Original file line number Diff line number Diff line change
@@ -1,4 +1,4 @@
CREATE TABLE dummy_table
(
dummy_column ENUM ('x', 'y')
);
dummy_column ENUM('x', 'y')
);
4 changes: 2 additions & 2 deletions CodegenTests/test-requests/SchemaScopedEnum/schema.sql
Original file line number Diff line number Diff line change
Expand Up @@ -2,5 +2,5 @@ CREATE SCHEMA dummy_schema;

CREATE TABLE dummy_schema.dummy_table
(
dummy_column ENUM ('x', 'y')
);
dummy_column ENUM('x', 'y')
);
61 changes: 52 additions & 9 deletions Drivers/DbDriver.cs
Original file line number Diff line number Diff line change
Expand Up @@ -130,13 +130,13 @@ public string AddNullableSuffixIfNeeded(string csharpType, bool notNull)
return IsTypeNullable(csharpType) ? $"{csharpType}?" : csharpType;
}

public string GetCsharpType(Column column)
public string GetCsharpType(Column column, Query? query)
{
var csharpType = GetCsharpTypeWithoutNullableSuffix(column);
return AddNullableSuffixIfNeeded(csharpType, column.NotNull);
var csharpType = GetCsharpTypeWithoutNullableSuffix(column, query);
return AddNullableSuffixIfNeeded(csharpType, IsColumnNotNull(column, query));
}

private string GetCsharpTypeWithoutNullableSuffix(Column column)
private string GetCsharpTypeWithoutNullableSuffix(Column column, Query? query)
{
if (column.EmbedTable != null)
return column.EmbedTable.Name.ToModelName(column.EmbedTable.Schema, DefaultSchema);
Expand All @@ -147,6 +147,13 @@ private string GetCsharpTypeWithoutNullableSuffix(Column column)
if (IsEnumType(column))
return column.Type.Name.ToModelName(column.Table.Schema, DefaultSchema);

if (query is not null)
{
var foundOverride = FindOverrideForQueryColumn(query, column);
if (foundOverride is not null)
return foundOverride.CsharpType.Type;
}

foreach (var columnMapping in ColumnMappings
.Where(columnMapping => DoesColumnMappingApply(columnMapping, column)))
{
Expand Down Expand Up @@ -177,14 +184,29 @@ private static bool DoesColumnMappingApply(ColumnMapping columnMapping, Column c
return typeInfo.Length.Value == column.Length;
}

public string GetColumnReader(Column column, int ordinal)
private string GetColumnReader(OverrideOption overrideOption, int ordinal)
{
var columnMapping = ColumnMappings.Find(c => c.CsharpType == overrideOption.CsharpType.Type);
if (columnMapping is not null)
return columnMapping.ReaderFn(ordinal);
throw new NotSupportedException($"Column {overrideOption.Column} has unsupported column type: {overrideOption.CsharpType.Type}");
}

public string GetColumnReader(Column column, int ordinal, Query? query)
{
if (IsEnumType(column))
{
var enumName = column.Type.Name.ToModelName(column.Table.Schema, DefaultSchema);
return $"{Variable.Reader.AsVarName()}.GetString({ordinal}).To{enumName}()";
}

if (query is not null)
{
var foundOverride = FindOverrideForQueryColumn(query, column);
if (foundOverride is not null)
return GetColumnReader(foundOverride, ordinal);
}

foreach (var columnMapping in ColumnMappings
.Where(columnMapping => DoesColumnMappingApply(columnMapping, column)))
{
Expand Down Expand Up @@ -227,10 +249,10 @@ public string GetIdColumnType(Query query)
var tableColumns = Tables[query.InsertIntoTable.Schema][query.InsertIntoTable.Name].Columns;
var idColumn = tableColumns.First(c => c.Name.Equals("id", StringComparison.OrdinalIgnoreCase));
if (idColumn is not null)
return GetCsharpType(idColumn);
return GetCsharpType(idColumn, query);

idColumn = tableColumns.First(c => c.Name.Contains("id", StringComparison.CurrentCultureIgnoreCase));
return GetCsharpType(idColumn ?? tableColumns[0]);
return GetCsharpType(idColumn ?? tableColumns[0], query);
}

public virtual string[] GetLastIdStatement(Query query)
Expand All @@ -243,10 +265,10 @@ public virtual string[] GetLastIdStatement(Query query)
];
}

public Column GetColumnFromParam(Parameter queryParam)
public Column GetColumnFromParam(Parameter queryParam, Query query)
{
if (string.IsNullOrEmpty(queryParam.Column.Name))
queryParam.Column.Name = $"{GetCsharpType(queryParam.Column).Replace("[]", "Arr")}_{queryParam.Number}";
queryParam.Column.Name = $"{GetCsharpType(queryParam.Column, query).Replace("[]", "Arr")}_{queryParam.Number}";
return queryParam.Column;
}

Expand All @@ -262,4 +284,25 @@ protected bool BatchQueryExists()
{
return Queries.Any(q => q.Cmd is ":copyfrom");
}

public OverrideOption? FindOverrideForQueryColumn(Query query, Column column)
{
return Options.Overrides.FirstOrDefault(o => o.Column.Equals($"{query.Name}:{column.Name}"));
}

/// <summary>
/// If the column data type is overridden, we need to check for nulls in generated code
/// </summary>
/// <param name="column"></param>
/// <param name="query"></param>
/// <returns>Adjusted not null value</returns>
public bool IsColumnNotNull(Column column, Query? query)
{
if (query is null)
return column.NotNull;
var overrideColumn = FindOverrideForQueryColumn(query, column);
if (overrideColumn is not null)
return overrideColumn.CsharpType.NotNull;
return column.NotNull;
}
}
52 changes: 26 additions & 26 deletions Drivers/Generators/CommonGen.cs
Original file line number Diff line number Diff line change
Expand Up @@ -92,7 +92,7 @@
""";
}

public string InstantiateDataclass(Column[] columns, string returnInterface)
public string InstantiateDataclass(Column[] columns, string returnInterface, Query? query)
{
var columnsInit = new List<string>();
var actualOrdinal = 0;
Expand All @@ -102,7 +102,7 @@
{
if (column.EmbedTable is null)
{
columnsInit.Add(GetAsSimpleAssignment(column, actualOrdinal));
columnsInit.Add(GetAsSimpleAssignment(column, actualOrdinal, query));

Check warning on line 105 in Drivers/Generators/CommonGen.cs

View workflow job for this annotation

GitHub Actions / Build (WASM)

Possible null reference argument for parameter 'query' in 'string GetAsSimpleAssignment(Column column, int ordinal, Query query)'.

Check warning on line 105 in Drivers/Generators/CommonGen.cs

View workflow job for this annotation

GitHub Actions / Codegen Tests

Possible null reference argument for parameter 'query' in 'string GetAsSimpleAssignment(Column column, int ordinal, Query query)'.
actualOrdinal++;
continue;
}
Expand All @@ -113,47 +113,28 @@
seenEmbed.TryAdd(tableFieldType, 1);
seenEmbed[tableFieldType]++;

var tableColumnsInit = GetAsEmbeddedTableColumnAssignment(column, actualOrdinal);
var tableColumnsInit = GetAsEmbeddedTableColumnAssignment(column, actualOrdinal, query);
columnsInit.Add($"{tableFieldName} = {InstantiateDataclassInternal(tableFieldType, tableColumnsInit)}");
actualOrdinal += tableColumnsInit.Length;
}

return InstantiateDataclassInternal(returnInterface, columnsInit);

string[] GetAsEmbeddedTableColumnAssignment(Column tableColumn, int ordinal)
string[] GetAsEmbeddedTableColumnAssignment(Column tableColumn, int ordinal, Query? query)
{
var schemaName = tableColumn.EmbedTable.Schema == dbDriver.DefaultSchema ? string.Empty : tableColumn.EmbedTable.Schema;
var tableColumns = dbDriver.Tables[schemaName][tableColumn.EmbedTable.Name].Columns;
return tableColumns
.Select((c, o) => GetAsSimpleAssignment(c, o + ordinal))
.Select((c, o) => GetAsSimpleAssignment(c, o + ordinal, query))

Check warning on line 128 in Drivers/Generators/CommonGen.cs

View workflow job for this annotation

GitHub Actions / Build (WASM)

Possible null reference argument for parameter 'query' in 'string GetAsSimpleAssignment(Column column, int ordinal, Query query)'.

Check warning on line 128 in Drivers/Generators/CommonGen.cs

View workflow job for this annotation

GitHub Actions / Codegen Tests

Possible null reference argument for parameter 'query' in 'string GetAsSimpleAssignment(Column column, int ordinal, Query query)'.
.ToArray();
}

string GetAsSimpleAssignment(Column column, int ordinal)
string GetAsSimpleAssignment(Column column, int ordinal, Query query)
{
var readExpression = GetReadExpression(column, ordinal);
var readExpression = GetReadExpression(column, ordinal, query);
return $"{column.Name.ToPascalCase()} = {readExpression}";
}

string GetReadExpression(Column column, int ordinal)
{
return column.NotNull
? dbDriver.GetColumnReader(column, ordinal)
: $"{CheckNullExpression(ordinal)} ? {GetNullExpression(column)} : {dbDriver.GetColumnReader(column, ordinal)}";
}

string GetNullExpression(Column column)
{
var csharpType = dbDriver.GetCsharpType(column);
if (dbDriver.Options.DotnetFramework.IsDotnetCore()) return "null";
return dbDriver.IsTypeNullable(csharpType) ? $"({csharpType}) null" : "null";
}

string CheckNullExpression(int ordinal)
{
return $"{Variable.Reader.AsVarName()}.IsDBNull({ordinal})";
}

string InstantiateDataclassInternal(string name, IEnumerable<string> fieldsInit)
{
return $$"""
Expand All @@ -164,4 +145,23 @@
""";
}
}

private string GetNullExpression(Column column, Query? query)
{
var csharpType = dbDriver.GetCsharpType(column, query);
if (dbDriver.Options.DotnetFramework.IsDotnetCore()) return "null";
return dbDriver.IsTypeNullable(csharpType) ? $"({csharpType}) null" : "null";
}

private static string CheckNullExpression(int ordinal)
{
return $"{Variable.Reader.AsVarName()}.IsDBNull({ordinal})";
}

private string GetReadExpression(Column column, int ordinal, Query query)
{
if (dbDriver.IsColumnNotNull(column, query))
return dbDriver.GetColumnReader(column, ordinal, query);
return $"{CheckNullExpression(ordinal)} ? {GetNullExpression(column, query)} : {dbDriver.GetColumnReader(column, ordinal, query)}";
}
}
Loading