Matching a DbSet with a T when mocking with Moq framework

Viewed 144

I am using Entity Framework 6 and Moq framework. Currently I am writing a few Unit Tests where in each test I need to set up a dataset with the appropriate type for each test. One of the unit tests looks like this:

[TestMethod]
public async Task GetAllCaseCategories_WithEmptyDataset_ReturnsEmpty()
{
    var data = new List<DataAccessLayer.Tables.CaseCategory>(){};
    var mockContext = GetMockContextWithCaseCategoryDataSetAsync(data);
    var mockLogger = new Mock<IDatabaseContextLogging>();
    DatabaseTablesAccess databaseAccess = new DatabaseTablesAccess(mockDatabaseContext.Object, mockLogger.Object);
    List<OneCaseCategory> caseCategories = await databaseAccess.GetAllCaseCategories();
    Assert.IsTrue(caseCategories.Count().Equals(0), "caseCategories should not contain any items.");
}

The method that generates the database context mock object is contain in the method called "GetMockContextWithCaseCategoryDataSetAsync". Here it is:

public Mock<Context> GetMockContextWithCaseCategoryDataSetAsync(List<CaseCategory> data)
{
    var mockCaseCategorySet = new Mock<DbSet<CaseCategory>>() { };
    mockCaseCategorySet.As<IQueryable<CaseCategory>>().Setup(x => x.Provider).Returns(new TestDbAsyncQueryProvider<CaseCategory>(data.AsQueryable().Provider));
    mockCaseCategorySet.As<IQueryable<CaseCategory>>().Setup(x => x.Expression).Returns(data.AsQueryable().Expression);
    mockCaseCategorySet.As<IQueryable<CaseCategory>>().Setup(x => x.ElementType).Returns(data.AsQueryable().ElementType);
    mockCaseCategorySet.As<IDbAsyncEnumerable<CaseCategory>>().Setup(x => x.GetAsyncEnumerator()).Returns(new TestDbAsyncEnumerator<CaseCategory>(data.AsQueryable().GetEnumerator()));
    mockCaseCategorySet.As<IQueryable<CaseCategory>>().Setup(x => x.GetEnumerator()).Returns(data.AsQueryable().GetEnumerator());
    Mock<Context> m = new Mock<Context>();
    m.Setup(x => x.CaseCategories).Returns(mockCaseCategorySet.Object);
    return m;
}

I have multiple tables in the data access object and I have written a method for each one that returns the data access object with the appropriate type of data bound to the correct table in the mock object. I would like to generalize this and here is my attempt:

public Mock<Context> GetMockContextWithCaseCategoryDataSetAsync<T>(List<T> data) where T : class
{
    var mockCaseCategorySet = new Mock<DbSet<T>>() { };
    mockCaseCategorySet.As<IQueryable<T>>().Setup(x => x.Provider).Returns(new TestDbAsyncQueryProvider<T>(data.AsQueryable().Provider));
    mockCaseCategorySet.As<IQueryable<T>>().Setup(x => x.Expression).Returns(data.AsQueryable().Expression);
    mockCaseCategorySet.As<IQueryable<T>>().Setup(x => x.ElementType).Returns(data.AsQueryable().ElementType);
    mockCaseCategorySet.As<IDbAsyncEnumerable<T>>().Setup(x => x.GetAsyncEnumerator()).Returns(new TestDbAsyncEnumerator<T>(data.AsQueryable().GetEnumerator()));
    mockCaseCategorySet.As<IQueryable<T>>().Setup(x => x.GetEnumerator()).Returns(data.AsQueryable().GetEnumerator());
    Mock<Context> m = new Mock<Context>();
    m.Setup(x => x.?).Returns(mockCaseCategorySet.Object);
    return m;
}

Everything is good up until that second last line in the generic version of the method where I am setting up which table returns the mocked object. I am lost as to how (if even possible) I control which table is used based on the generic passed in. I would like to use it like this:

var context = GetMockContextWithCaseCategoryDataSetAsync<CaseCategory>(data);
1 Answers

I don't see a way to do what you're trying to do unless you also access the DbSets by the generic Set<TEntity> method (like _dbContext.Set<Salary>().Select(x => x.Id)).

An alternative setup that I have been using and that is sort of generic is to create a "MockedDbContext" class that I reuse in all tests accessing the database. This usually looks something like this;

public class MockedDbContext
{
    public MockedDbContext() => MockContextProperties();

    public Mock<MyDbContext> ContextMock { get; } = new Mock<MyDbContext>(MockBehavior.Loose);

    public HashSet<Employee> Employees { get; } = new HashSet<Employee>();
    public HashSet<Salary> Salaries { get; } = new HashSet<Salary>();
    // etc...

    private void MockContextProperties()
    {
        SimulateGet(x => x.Employees, Employees);
        SimulateGet(x => x.Salaries, Salaries);
        // etc...
    }

    private void SimulateGet<TModel>(
        Expression<Func<MyDbContext, DbSet<TModel>>> dbSetExpression,
        HashSet<TModel> mockedData) where TModel : class
    {
        ContextMock.SetupGet(dbSetExpression).Returns(new InMemoryDbSet<TModel>(mockedData));
    }

    public void VerifyThatSaveWasCalled() => ContextMock.Verify(x => x.SaveChanges());
    public void VerifyThatSaveWasNotCalled() => ContextMock.Verify(x => x.SaveChanges(), Times.Never);
}

I use it like this;

[TestFixture]
public class When_getting_some_data
{
    private readonly MockedDbContext _mockedDbContext = new MockedDbContext();
    private readonly Fixture _fixture = new Fixture();
    private string _organizationNumber = "yyyyyy-xxx";
    private readonly DateTime _created = new DateTime(2022, 2, 2);
    
    private SomeInformation _result;

    [OneTimeSetUp]
    public void Initialize()
    {
        var salary = _fixture.Build<Salary>().With(o => o.Created, _created).Create();
        _mockedDbContext.Salaries.Add(salary);

        _result = new SomeClass(_mockedDbContext.ContextMock.Object).GetSomeInformation(_organizationNumber);
    }

    [Test]
    public void Then_returned_data_is_as_expected()
    {
        _result.Should().NotBeNull();
        _result.Created.Should().Be(_created);
    }
}

And It internally uses an InMemoryDbSet<T> that looks like this;

/// <summary>
/// The in-memory database set, taken from Microsoft's online example (http://msdn.microsoft.com/en-us/ff714955.aspx)
/// and modified to be based on DbSet instead of ObjectSet.
/// </summary>
/// <typeparam name="T">The type of DbSet.</typeparam>
public class InMemoryDbSet<T> : DbSet<T>, IQueryable<T> where T : class
{
    private static readonly HashSet<T> StaticData = new HashSet<T>();
    private readonly HashSet<T> _nonStaticData;

    /// <summary>
    /// Creates an instance of the InMemoryDbSet using the default static backing store.This means
    /// that data persists between test runs, like it would do with a database unless you
    /// cleared it down.
    /// </summary>
    public InMemoryDbSet() : this(true) { }

    /// <summary>
    /// This constructor allows you to pass in your own data store, instead of using
    /// the static backing store.
    /// </summary>
    /// <param name="data">A place to store data.</param>
    public InMemoryDbSet(HashSet<T> data) => _nonStaticData = data;

    /// <summary>
    /// Creates an instance of the InMemoryDbSet using the default static backing store.This means
    /// that data persists between test runs, like it would do with a database unless you
    /// cleared it down.
    /// </summary>
    /// <param name="clearDownExistingData">True to clear existing data</param>
    public InMemoryDbSet(bool clearDownExistingData)
    {
        if (clearDownExistingData)
            Clear();
    }

    public Func<IEnumerable<T>, object[], T> FindFunction { get; set; }
    public Type ElementType => Data.AsQueryable().ElementType;
    public Expression Expression => Data.AsQueryable().Expression;
    public IQueryProvider Provider => Data.AsQueryable().Provider;
    IQueryProvider IQueryable.Provider => Data.AsQueryable().Provider;
    public new ObservableCollection<T> Local => new ObservableCollection<T>(Data);

    /// <summary>
    /// The non static backing store data for the InMemoryDbSet.
    /// </summary>
    private HashSet<T> Data => _nonStaticData ?? StaticData;

    public void Clear() => Data.Clear();

    public override EntityEntry<T> Add(T entity)
    {
        Data.Add(entity);
        return null;
    }

    public new T Attach(T entity)
    {
        Data.Add(entity);
        return entity;
    }

    public TDerivedEntity Create<TDerivedEntity>() where TDerivedEntity : class, T => Activator.CreateInstance<TDerivedEntity>();

    public T Create() => Activator.CreateInstance<T>();

    public new virtual T Find(params object[] keyValues)
    {
        if (FindFunction != null)
        {
            return FindFunction(Data, keyValues);
        }

        throw new NotImplementedException("Derive from InMemoryDbSet and override Find, or provide a FindFunction.");
    }

    public new T Remove(T entity)
    {
        Data.Remove(entity);
        return entity;
    }

    public IEnumerator<T> GetEnumerator() => Data.GetEnumerator();
}

This saves a lot of work when writing unit tests and is kind of a one time setup which is what I think you were looking for.

Related