Skip to content
Closed
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
94 changes: 60 additions & 34 deletions Dapper.Contrib/SqlMapperExtensions.cs
Original file line number Diff line number Diff line change
Expand Up @@ -6,8 +6,8 @@
using System.Text;
using System.Collections.Concurrent;
using System.Reflection.Emit;

using Dapper;
using static Dapper.Contrib.Extensions.SqlMapperExtensions;

#if COREFX
using DataException = System.InvalidOperationException;
Expand All @@ -32,6 +32,7 @@ public interface ITableNameMapper

public delegate string GetDatabaseTypeDelegate(IDbConnection connection);
public delegate string TableNameMapperDelegate(Type type);
public delegate string ColumNameMapperDelegate(PropertyInfo propertyInfo);

private static readonly ConcurrentDictionary<RuntimeTypeHandle, IEnumerable<PropertyInfo>> KeyProperties = new ConcurrentDictionary<RuntimeTypeHandle, IEnumerable<PropertyInfo>>();
private static readonly ConcurrentDictionary<RuntimeTypeHandle, IEnumerable<PropertyInfo>> ExplicitKeyProperties = new ConcurrentDictionary<RuntimeTypeHandle, IEnumerable<PropertyInfo>>();
Expand All @@ -40,6 +41,9 @@ public interface ITableNameMapper
private static readonly ConcurrentDictionary<RuntimeTypeHandle, string> GetQueries = new ConcurrentDictionary<RuntimeTypeHandle, string>();
private static readonly ConcurrentDictionary<RuntimeTypeHandle, string> TypeTableName = new ConcurrentDictionary<RuntimeTypeHandle, string>();

private static readonly ConcurrentDictionary<PropertyInfo, string> PropertyInfoToColumnName = new ConcurrentDictionary<PropertyInfo, string>();


private static readonly ISqlAdapter DefaultAdapter = new SqlServerAdapter();
private static readonly Dictionary<string, ISqlAdapter> AdapterDictionary
= new Dictionary<string, ISqlAdapter>
Expand All @@ -51,6 +55,28 @@ private static readonly Dictionary<string, ISqlAdapter> AdapterDictionary
{"mysqlconnection", new MySqlAdapter()},
};


public static ColumNameMapperDelegate ColumnNameMapper;
private static string PropertyInfoToColumnNameCache(PropertyInfo propertyInfo)
{
string name = null;
if (PropertyInfoToColumnName.TryGetValue(propertyInfo, out name))
{
return name;
}
if (ColumnNameMapper == null)
{
name = propertyInfo.Name;
}
else
{
name = ColumnNameMapper(propertyInfo);
}
PropertyInfoToColumnName[propertyInfo] = name;
return name;
}


private static List<PropertyInfo> ComputedPropertiesCache(Type type)
{
IEnumerable<PropertyInfo> pi;
Expand Down Expand Up @@ -131,7 +157,7 @@ private static bool IsWriteable(PropertyInfo pi)

private static PropertyInfo GetSingleKey<T>(string method)
{
var type = typeof (T);
var type = typeof(T);
var keys = KeyPropertiesCache(type);
var explicitKeys = ExplicitKeyPropertiesCache(type);
var keyCount = keys.Count + explicitKeys.Count;
Expand Down Expand Up @@ -165,7 +191,7 @@ public static T Get<T>(this IDbConnection connection, dynamic id, IDbTransaction
var key = GetSingleKey<T>(nameof(Get));
var name = GetTableName(type);

sql = $"select * from {name} where {key.Name} = @id";
sql = $"SELECT * FROM {name} WHERE {PropertyInfoToColumnNameCache(key)} = @id";
GetQueries[type.TypeHandle] = sql;
}

Expand All @@ -185,7 +211,7 @@ public static T Get<T>(this IDbConnection connection, dynamic id, IDbTransaction

foreach (var property in TypePropertiesCache(type))
{
var val = res[property.Name];
var val = res[PropertyInfoToColumnNameCache(property)];
property.SetValue(obj, Convert.ChangeType(val, property.PropertyType), null);
}

Expand Down Expand Up @@ -220,7 +246,7 @@ public static IEnumerable<T> GetAll<T>(this IDbConnection connection, IDbTransac
GetSingleKey<T>(nameof(GetAll));
var name = GetTableName(type);

sql = "select * from " + name;
sql = "SELECT * FROM " + name;
GetQueries[cacheType.TypeHandle] = sql;
}

Expand All @@ -233,7 +259,7 @@ public static IEnumerable<T> GetAll<T>(this IDbConnection connection, IDbTransac
var obj = ProxyGenerator.GetInterfaceProxy<T>();
foreach (var property in TypePropertiesCache(type))
{
var val = res[property.Name];
var val = res[PropertyInfoToColumnNameCache(property)];
property.SetValue(obj, Convert.ChangeType(val, property.PropertyType), null);
}
((IProxy)obj).IsDirty = false; //reset change tracking and return
Expand Down Expand Up @@ -316,7 +342,7 @@ public static long Insert<T>(this IDbConnection connection, T entityToInsert, ID
for (var i = 0; i < allPropertiesExceptKeyAndComputed.Count; i++)
{
var property = allPropertiesExceptKeyAndComputed.ElementAt(i);
adapter.AppendColumnName(sbColumnList, property.Name); //fix for issue #336
adapter.AppendColumnName(sbColumnList, PropertyInfoToColumnNameCache(property)); //fix for issue #336
if (i < allPropertiesExceptKeyAndComputed.Count - 1)
sbColumnList.Append(", ");
}
Expand All @@ -337,12 +363,12 @@ public static long Insert<T>(this IDbConnection connection, T entityToInsert, ID
if (!isList) //single entity
{
returnVal = adapter.Insert(connection, transaction, commandTimeout, name, sbColumnList.ToString(),
sbParameterList.ToString(), keyProperties, entityToInsert);
sbParameterList.ToString(), keyProperties, ColumnNameMapper, entityToInsert);
}
else
{
//insert list of entities
var cmd = $"insert into {name} ({sbColumnList}) values ({sbParameterList})";
var cmd = $"INSERT INTO {name} ({sbColumnList}) VALUES ({sbParameterList})";
returnVal = connection.Execute(cmd, entityToInsert, transaction, commandTimeout);
}
if (wasClosed) connection.Close();
Expand Down Expand Up @@ -385,29 +411,29 @@ public static bool Update<T>(this IDbConnection connection, T entityToUpdate, ID
var name = GetTableName(type);

var sb = new StringBuilder();
sb.AppendFormat("update {0} set ", name);
sb.AppendFormat("UPDATE {0} SET ", name);

var allProperties = TypePropertiesCache(type);
keyProperties.AddRange(explicitKeyProperties);
var computedProperties = ComputedPropertiesCache(type);
var nonIdProps = allProperties.Except(keyProperties.Union(computedProperties)).ToList();

var adapter = GetFormatter(connection);
var adapter = GetFormatter(connection);

for (var i = 0; i < nonIdProps.Count; i++)
{
var property = nonIdProps.ElementAt(i);
adapter.AppendColumnNameEqualsValue(sb, property.Name); //fix for issue #336
adapter.AppendColumnNameEqualsValue(sb, PropertyInfoToColumnNameCache(property)); //fix for issue #336
if (i < nonIdProps.Count - 1)
sb.AppendFormat(", ");
}
sb.Append(" where ");
sb.Append(" WHERE ");
for (var i = 0; i < keyProperties.Count; i++)
{
var property = keyProperties.ElementAt(i);
adapter.AppendColumnNameEqualsValue(sb, property.Name); //fix for issue #336
adapter.AppendColumnNameEqualsValue(sb, PropertyInfoToColumnNameCache(property)); //fix for issue #336
if (i < keyProperties.Count - 1)
sb.AppendFormat(" and ");
sb.AppendFormat(" AND ");
}
var updated = connection.Execute(sb.ToString(), entityToUpdate, commandTimeout: commandTimeout, transaction: transaction);
return updated > 0;
Expand Down Expand Up @@ -447,16 +473,16 @@ public static bool Delete<T>(this IDbConnection connection, T entityToDelete, ID
keyProperties.AddRange(explicitKeyProperties);

var sb = new StringBuilder();
sb.AppendFormat("delete from {0} where ", name);
sb.AppendFormat("DELETE FROM {0} WHERE ", name);

var adapter = GetFormatter(connection);

for (var i = 0; i < keyProperties.Count; i++)
{
var property = keyProperties.ElementAt(i);
adapter.AppendColumnNameEqualsValue(sb, property.Name); //fix for issue #336
adapter.AppendColumnNameEqualsValue(sb, PropertyInfoToColumnNameCache(property)); //fix for issue #336
if (i < keyProperties.Count - 1)
sb.AppendFormat(" and ");
sb.AppendFormat(" AND ");
}
var deleted = connection.Execute(sb.ToString(), entityToDelete, transaction, commandTimeout);
return deleted > 0;
Expand All @@ -474,7 +500,7 @@ public static bool DeleteAll<T>(this IDbConnection connection, IDbTransaction tr
{
var type = typeof(T);
var name = GetTableName(type);
var statement = $"delete from {name}";
var statement = $"DELETE FROM {name}";
var deleted = connection.Execute(statement, null, transaction, commandTimeout);
return deleted > 0;
}
Expand All @@ -484,7 +510,7 @@ public static bool DeleteAll<T>(this IDbConnection connection, IDbTransaction tr
/// Please note that this callback is global and will be used by all the calls that require a database specific adapter.
/// </summary>
public static GetDatabaseTypeDelegate GetDatabaseType;

private static ISqlAdapter GetFormatter(IDbConnection connection)
{
var name = GetDatabaseType?.Invoke(connection).ToLower()
Expand Down Expand Up @@ -688,18 +714,18 @@ public class ComputedAttribute : Attribute

public partial interface ISqlAdapter
{
int Insert(IDbConnection connection, IDbTransaction transaction, int? commandTimeout, string tableName, string columnList, string parameterList, IEnumerable<PropertyInfo> keyProperties, object entityToInsert);
int Insert(IDbConnection connection, IDbTransaction transaction, int? commandTimeout, string tableName, string columnList, string parameterList, IEnumerable<PropertyInfo> keyProperties, ColumNameMapperDelegate columnNameMapper, object entityToInsert);

//new methods for issue #336
void AppendColumnName(StringBuilder sb, string columnName);
void AppendColumnNameEqualsValue(StringBuilder sb, string columnName);
}

public partial class SqlServerAdapter : ISqlAdapter
{
public int Insert(IDbConnection connection, IDbTransaction transaction, int? commandTimeout, string tableName, string columnList, string parameterList, IEnumerable<PropertyInfo> keyProperties, object entityToInsert)
public int Insert(IDbConnection connection, IDbTransaction transaction, int? commandTimeout, string tableName, string columnList, string parameterList, IEnumerable<PropertyInfo> keyProperties, ColumNameMapperDelegate columnNameMapper, object entityToInsert)
{
var cmd = $"insert into {tableName} ({columnList}) values ({parameterList});select SCOPE_IDENTITY() id";
var cmd = $"INSERT INTO {tableName} ({columnList}) VALUES ({parameterList});SELECT SCOPE_IDENTITY() [id]";
var multi = connection.QueryMultiple(cmd, entityToInsert, transaction, commandTimeout);

var first = multi.Read().FirstOrDefault();
Expand Down Expand Up @@ -728,14 +754,14 @@ public void AppendColumnNameEqualsValue(StringBuilder sb, string columnName)

public partial class SqlCeServerAdapter : ISqlAdapter
{
public int Insert(IDbConnection connection, IDbTransaction transaction, int? commandTimeout, string tableName, string columnList, string parameterList, IEnumerable<PropertyInfo> keyProperties, object entityToInsert)
public int Insert(IDbConnection connection, IDbTransaction transaction, int? commandTimeout, string tableName, string columnList, string parameterList, IEnumerable<PropertyInfo> keyProperties, ColumNameMapperDelegate columnNameMapper, object entityToInsert)
{
var cmd = $"insert into {tableName} ({columnList}) values ({parameterList})";
var cmd = $"INSERT INTO {tableName} ({columnList}) VALUES ({parameterList})";
connection.Execute(cmd, entityToInsert, transaction, commandTimeout);
var r = connection.Query("select @@IDENTITY id", transaction: transaction, commandTimeout: commandTimeout).ToList();
var r = connection.Query("SELECT @@IDENTITY [id]", transaction: transaction, commandTimeout: commandTimeout).ToList();

if (r.First().id == null) return 0;
var id = (int) r.First().id;
var id = (int)r.First().id;

var propertyInfos = keyProperties as PropertyInfo[] ?? keyProperties.ToArray();
if (!propertyInfos.Any()) return id;
Expand All @@ -759,7 +785,7 @@ public void AppendColumnNameEqualsValue(StringBuilder sb, string columnName)

public partial class MySqlAdapter : ISqlAdapter
{
public int Insert(IDbConnection connection, IDbTransaction transaction, int? commandTimeout, string tableName, string columnList, string parameterList, IEnumerable<PropertyInfo> keyProperties, object entityToInsert)
public int Insert(IDbConnection connection, IDbTransaction transaction, int? commandTimeout, string tableName, string columnList, string parameterList, IEnumerable<PropertyInfo> keyProperties, ColumNameMapperDelegate columnNameMapper, object entityToInsert)
{
var cmd = $"insert into {tableName} ({columnList}) values ({parameterList})";
connection.Execute(cmd, entityToInsert, transaction, commandTimeout);
Expand Down Expand Up @@ -790,10 +816,10 @@ public void AppendColumnNameEqualsValue(StringBuilder sb, string columnName)

public partial class PostgresAdapter : ISqlAdapter
{
public int Insert(IDbConnection connection, IDbTransaction transaction, int? commandTimeout, string tableName, string columnList, string parameterList, IEnumerable<PropertyInfo> keyProperties, object entityToInsert)
public int Insert(IDbConnection connection, IDbTransaction transaction, int? commandTimeout, string tableName, string columnList, string parameterList, IEnumerable<PropertyInfo> keyProperties, ColumNameMapperDelegate columnNameMapper, object entityToInsert)
{
var sb = new StringBuilder();
sb.AppendFormat("insert into {0} ({1}) values ({2})", tableName, columnList, parameterList);
sb.AppendFormat("INSERT INTO {0} ({1}) VALUES ({2})", tableName, columnList, parameterList);

// If no primary key then safe to assume a join table with not too much data to return
var propertyInfos = keyProperties as PropertyInfo[] ?? keyProperties.ToArray();
Expand All @@ -808,7 +834,7 @@ public int Insert(IDbConnection connection, IDbTransaction transaction, int? com
if (!first)
sb.Append(", ");
first = false;
sb.Append(property.Name);
sb.AppendFormat("\"{0}\"", columnNameMapper(property));
}
}

Expand All @@ -818,7 +844,7 @@ public int Insert(IDbConnection connection, IDbTransaction transaction, int? com
var id = 0;
foreach (var p in propertyInfos)
{
var value = ((IDictionary<string, object>)results.First())[p.Name.ToLower()];
var value = ((IDictionary<string, object>)results.First())[columnNameMapper(p)];
p.SetValue(entityToInsert, value, null);
if (id == 0)
id = Convert.ToInt32(value);
Expand All @@ -839,7 +865,7 @@ public void AppendColumnNameEqualsValue(StringBuilder sb, string columnName)

public partial class SQLiteAdapter : ISqlAdapter
{
public int Insert(IDbConnection connection, IDbTransaction transaction, int? commandTimeout, string tableName, string columnList, string parameterList, IEnumerable<PropertyInfo> keyProperties, object entityToInsert)
public int Insert(IDbConnection connection, IDbTransaction transaction, int? commandTimeout, string tableName, string columnList, string parameterList, IEnumerable<PropertyInfo> keyProperties, ColumNameMapperDelegate columnNameMapper, object entityToInsert)
{
var cmd = $"INSERT INTO {tableName} ({columnList}) VALUES ({parameterList}); SELECT last_insert_rowid() id";
var multi = connection.QueryMultiple(cmd, entityToInsert, transaction, commandTimeout);
Expand Down