diff --git a/src/Marten.Testing/Linq/dictionary_is_translated.cs b/src/Marten.Testing/Linq/dictionary_is_translated.cs new file mode 100644 index 0000000000..11f03ac609 --- /dev/null +++ b/src/Marten.Testing/Linq/dictionary_is_translated.cs @@ -0,0 +1,53 @@ +using Marten.Services; +using System; +using System.Collections.Generic; +using System.Linq; +using Xunit; + +namespace Marten.Testing.Linq +{ + public class dictionary_is_translated : DocumentSessionFixture + { + public dictionary_is_translated() + { + theStore.BulkInsert(Target.GenerateRandomData(100).ToArray()); + } + + [Fact] + public void dictionary_containskey_is_translated_to_json_map() + { + var query = theSession.Query().Where(t => t.StringDict.ContainsKey("foo")); + var command = query.ToCommand(Marten.Linq.FetchType.FetchMany); + var dictParam = command.Parameters[0]; + (dictParam.DbType == System.Data.DbType.String).ShouldBeTrue(); + (dictParam.Value.ToString() == "foo").ShouldBeTrue(); + } + + // using key0 and value0 for these because the last node, which is deep, should have at least a single dict node + + [Fact] + public void dict_can_query_using_containskey() + { + var results = theSession.Query().Where(x => x.StringDict.ContainsKey("key0")).ToList(); + results.All(r => r.StringDict.ContainsKey("key0")).ShouldBeTrue(); + } + + [Fact] + public void dict_can_query_using_containsKVP() + { + var kvp = new KeyValuePair("key0", "value0"); + var results = theSession.Query().Where(x => x.StringDict.Contains(kvp)).ToList(); + results.All(r => r.StringDict.Contains(kvp)).ShouldBeTrue(); + } + + [Fact] + public void icollection_keyvaluepair_contains_is_translated_to_json_map() + { + var query = theSession.Query().Where(t => t.StringDict.Contains(new KeyValuePair("foo", "bar"))); + var command = query.ToCommand(Marten.Linq.FetchType.FetchMany); + var dictParam = command.Parameters[0]; + (dictParam.NpgsqlDbType == NpgsqlTypes.NpgsqlDbType.Jsonb).ShouldBeTrue(); + (dictParam.Value.ToString() == "{\"foo\":\"bar\"}").ShouldBeTrue(); + } + } +} diff --git a/src/Marten.Testing/Linq/query_for_json_format.cs b/src/Marten.Testing/Linq/query_for_json_format.cs index d4ab80ea8b..b4e45ca1ce 100644 --- a/src/Marten.Testing/Linq/query_for_json_format.cs +++ b/src/Marten.Testing/Linq/query_for_json_format.cs @@ -36,7 +36,7 @@ public string ToJson() }}, ""UserName"": ""{UserName}"", ""Birthdate"": ""{Birthdate.ToString("s")}"" -}}".Replace("\r\n", ""); +}}".Replace("\r\n", "").Replace("\n", ""); } } diff --git a/src/Marten.Testing/Target.cs b/src/Marten.Testing/Target.cs index 078e48ca81..0cc935fbce 100644 --- a/src/Marten.Testing/Target.cs +++ b/src/Marten.Testing/Target.cs @@ -83,6 +83,8 @@ public static Target Random(bool deep = false) { target.Children[i] = Random(); } + + target.StringDict = Enumerable.Range(0, _random.Next(1, 10)).ToDictionary(i => $"key{i}", i => $"value{i}"); } return target; @@ -91,6 +93,7 @@ public static Target Random(bool deep = false) public Target() { Id = Guid.NewGuid(); + StringDict = new Dictionary(); } public Guid Id { get; set; } @@ -125,6 +128,9 @@ public Target() public int? NullableNumber { get; set; } public DateTime? NullableDateTime { get; set; } public bool? NullableBoolean { get; set; } + + public IDictionary StringDict { get; set; } + } public class Address diff --git a/src/Marten/Linq/MartenExpressionParser.cs b/src/Marten/Linq/MartenExpressionParser.cs index 9c9cf80606..1621834bc5 100644 --- a/src/Marten/Linq/MartenExpressionParser.cs +++ b/src/Marten/Linq/MartenExpressionParser.cs @@ -69,7 +69,10 @@ public IWhereFragment ParseWhereFragment(IQueryableDocument mapping, Expression // soft deletes new MaybeDeletedParser(), - new IsDeletedParser() + new IsDeletedParser(), + + // dictionaries + new DictionaryExpressions() }; diff --git a/src/Marten/Linq/Parsing/DictionaryExpressions.cs b/src/Marten/Linq/Parsing/DictionaryExpressions.cs new file mode 100644 index 0000000000..047875579e --- /dev/null +++ b/src/Marten/Linq/Parsing/DictionaryExpressions.cs @@ -0,0 +1,70 @@ +using System; +using System.Linq.Expressions; +using Marten.Schema; +using System.Reflection; +using System.Collections.Generic; + +namespace Marten.Linq.Parsing +{ + public class DictionaryExpressions : IMethodCallParser + { + static bool IsCollectionContainsWithStringKey(MethodInfo m) => + m.Name == "Contains" + && m.DeclaringType.IsConstructedGenericType + && m.DeclaringType.GetGenericTypeDefinition() == typeof(ICollection<>) + && m.DeclaringType.GenericTypeArguments[0].IsConstructedGenericType + && m.DeclaringType.GenericTypeArguments[0].GetGenericTypeDefinition() == typeof(KeyValuePair<,>) + && m.DeclaringType.GenericTypeArguments[0].GenericTypeArguments[0] == typeof(string); + + static bool IsDictionaryContainsKey(MethodInfo m) => + m.Name == "ContainsKey" + && m.DeclaringType.IsConstructedGenericType + && m.DeclaringType.GetGenericTypeDefinition() == typeof(IDictionary<,>) + && m.DeclaringType.GenericTypeArguments[0] == typeof(string); + + public bool Matches(MethodCallExpression expression) + { + return IsCollectionContainsWithStringKey(expression.Method) + || IsDictionaryContainsKey(expression.Method); + } + + public IWhereFragment Parse(IQueryableDocument mapping, ISerializer serializer, MethodCallExpression expression) + { + var finder = new FindMembers(); + finder.Visit(expression); + var members = finder.Members; + var fieldlocator = mapping.FieldFor(members).SqlLocator; + + if (IsCollectionContainsWithStringKey(expression.Method)) + { + return QueryFromICollectionContains(expression, fieldlocator, serializer); + } + else if (IsDictionaryContainsKey(expression.Method)) + { + return QueryFromDictionaryContainsKey(expression, fieldlocator); + } + else throw new NotImplementedException("Could not understand the format of the dictionary access"); + } + + static IWhereFragment QueryFromDictionaryContainsKey(MethodCallExpression expression, string fieldLocator) + { + var key = (string)expression.Arguments[0].Value(); + // have to use different token here because we actually want the `?` character as the operator! + return new CustomizableWhereFragment($"{fieldLocator} ? @1", "@1", Tuple.Create(key, NpgsqlTypes.NpgsqlDbType.Text)); + } + + static IWhereFragment QueryFromICollectionContains(MethodCallExpression expression, string fieldPath, ISerializer serializer) + { + var constant = expression.Arguments[0] as ConstantExpression; + var kvp = constant.Value; // is kvp + var kvpType = kvp.GetType(); + var key = kvpType.GetProperty("Key").GetValue(kvp); + var value = kvpType.GetProperty("Value").GetValue(kvp); + var dictType = typeof(Dictionary<,>).MakeGenericType(kvpType.GenericTypeArguments[0], kvpType.GenericTypeArguments[1]); + var dict = dictType.GetConstructors()[0].Invoke(null); + dictType.GetMethod("Add").Invoke(dict, new[] { key, value }); + var json = serializer.ToJson(dict); + return new CustomizableWhereFragment($"{fieldPath} @> ?", "?", Tuple.Create(json, NpgsqlTypes.NpgsqlDbType.Jsonb)); + } + } +} diff --git a/src/Marten/Linq/WhereFragment.cs b/src/Marten/Linq/WhereFragment.cs index 6a811a2fe0..12df71ba08 100644 --- a/src/Marten/Linq/WhereFragment.cs +++ b/src/Marten/Linq/WhereFragment.cs @@ -1,19 +1,23 @@ -using System.Collections.Generic; using Baseline; using Marten.Util; using Npgsql; +using System; +using System.Linq; namespace Marten.Linq { - public class WhereFragment : IWhereFragment + public class CustomizableWhereFragment : IWhereFragment { private readonly string _sql; - private readonly object[] _parameters; + private readonly Tuple[] _parameters; + private readonly string _token; - public WhereFragment(string sql, params object[] parameters) + + public CustomizableWhereFragment(string sql, string paramReplacementToken, params Tuple[] parameters) { _sql = sql; _parameters = parameters; + _token = paramReplacementToken; } public string ToSql(NpgsqlCommand command) @@ -21,8 +25,8 @@ public string ToSql(NpgsqlCommand command) var sql = _sql; _parameters.Each(x => { - var param = command.AddParameter(x); - sql = sql.ReplaceFirst("?", ":" + param.ParameterName); + var param = command.AddParameter(x.Item1, x.Item2); + sql = sql.ReplaceFirst(_token, ":" + param.ParameterName); }); return sql; @@ -33,4 +37,9 @@ public bool Contains(string sqlText) return _sql.Contains(sqlText); } } + + public class WhereFragment : CustomizableWhereFragment + { + public WhereFragment(string sql, params object[] parameters) : base(sql, "?", parameters.Select(x => Tuple.Create(x, null)).ToArray()) { } + } } \ No newline at end of file diff --git a/src/Marten/Util/TypeMappings.cs b/src/Marten/Util/TypeMappings.cs index b184cad4e1..ff584ec735 100644 --- a/src/Marten/Util/TypeMappings.cs +++ b/src/Marten/Util/TypeMappings.cs @@ -20,7 +20,8 @@ public static class TypeMappings {typeof (decimal), "decimal"}, {typeof(float), "decimal" }, {typeof(DateTime), "timestamp without time zone" }, - {typeof (DateTimeOffset), "timestamp with time zone"} + {typeof (DateTimeOffset), "timestamp with time zone"}, + {typeof (IDictionary<,>), "jsonb" }, }; private static readonly MethodInfo _getNgpsqlDbTypeMethod; @@ -105,6 +106,8 @@ public static string GetPgType(Type memberType) return GetPgType(memberType.GetInnerTypeFromNullable()); } + if (memberType.IsConstructedGenericType) return PgTypes[memberType.GetGenericTypeDefinition()]; + return PgTypes[memberType]; }