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
53 changes: 53 additions & 0 deletions src/Marten.Testing/Linq/dictionary_is_translated.cs
Original file line number Diff line number Diff line change
@@ -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<NulloIdentityMap>
{
public dictionary_is_translated()
{
theStore.BulkInsert(Target.GenerateRandomData(100).ToArray());
}

[Fact]
public void dictionary_containskey_is_translated_to_json_map()
{
var query = theSession.Query<Target>().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<Target>().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<string, string>("key0", "value0");
var results = theSession.Query<Target>().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<Target>().Where(t => t.StringDict.Contains(new KeyValuePair<string, string>("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();
}
}
}
2 changes: 1 addition & 1 deletion src/Marten.Testing/Linq/query_for_json_format.cs
Original file line number Diff line number Diff line change
Expand Up @@ -36,7 +36,7 @@ public string ToJson()
}},
""UserName"": ""{UserName}"",
""Birthdate"": ""{Birthdate.ToString("s")}""
}}".Replace("\r\n", "");
}}".Replace("\r\n", "").Replace("\n", "");

Copy link
Copy Markdown
Contributor Author

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

needed this to make the tests green up on my machine

}
}

Expand Down
6 changes: 6 additions & 0 deletions src/Marten.Testing/Target.cs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand All @@ -91,6 +93,7 @@ public static Target Random(bool deep = false)
public Target()
{
Id = Guid.NewGuid();
StringDict = new Dictionary<string, string>();
}

public Guid Id { get; set; }
Expand Down Expand Up @@ -125,6 +128,9 @@ public Target()
public int? NullableNumber { get; set; }
public DateTime? NullableDateTime { get; set; }
public bool? NullableBoolean { get; set; }

public IDictionary<string,string> StringDict { get; set; }

}

public class Address
Expand Down
5 changes: 4 additions & 1 deletion src/Marten/Linq/MartenExpressionParser.cs
Original file line number Diff line number Diff line change
Expand Up @@ -69,7 +69,10 @@ public IWhereFragment ParseWhereFragment(IQueryableDocument mapping, Expression

// soft deletes
new MaybeDeletedParser(),
new IsDeletedParser()
new IsDeletedParser(),

// dictionaries
new DictionaryExpressions()
};


Expand Down
70 changes: 70 additions & 0 deletions src/Marten/Linq/Parsing/DictionaryExpressions.cs
Original file line number Diff line number Diff line change
@@ -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<object, NpgsqlTypes.NpgsqlDbType?>(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<string, unknown>
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<object, NpgsqlTypes.NpgsqlDbType?>(json, NpgsqlTypes.NpgsqlDbType.Jsonb));
}
}
}
21 changes: 15 additions & 6 deletions src/Marten/Linq/WhereFragment.cs
Original file line number Diff line number Diff line change
@@ -1,28 +1,32 @@
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<object, NpgsqlTypes.NpgsqlDbType?>[] _parameters;
private readonly string _token;

public WhereFragment(string sql, params object[] parameters)

public CustomizableWhereFragment(string sql, string paramReplacementToken, params Tuple<object, NpgsqlTypes.NpgsqlDbType?>[] parameters)
{
_sql = sql;
_parameters = parameters;
_token = paramReplacementToken;
}

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;
Expand All @@ -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<object, NpgsqlTypes.NpgsqlDbType?>(x, null)).ToArray()) { }
}
}
5 changes: 4 additions & 1 deletion src/Marten/Util/TypeMappings.cs
Original file line number Diff line number Diff line change
Expand Up @@ -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;
Expand Down Expand Up @@ -105,6 +106,8 @@ public static string GetPgType(Type memberType)
return GetPgType(memberType.GetInnerTypeFromNullable());
}

if (memberType.IsConstructedGenericType) return PgTypes[memberType.GetGenericTypeDefinition()];

return PgTypes[memberType];
}

Expand Down