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
65 changes: 21 additions & 44 deletions src/NATS.Extensions.Microsoft.DependencyInjection/NatsBuilder.cs
Original file line number Diff line number Diff line change
Expand Up @@ -13,7 +13,7 @@ public class NatsBuilder
{
private readonly IServiceCollection _services;

private int _poolSize = 1;
private Func<IServiceProvider, int> _poolSizeConfigurer = _ => 1;
private Func<IServiceProvider, NatsOpts, NatsOpts>? _configureOpts;
private Action<IServiceProvider, NatsConnection>? _configureConnection;
private object? _diKey = null;
Expand All @@ -25,7 +25,14 @@ public NatsBuilder(IServiceCollection services)

public NatsBuilder WithPoolSize(int size)
{
_poolSize = Math.Max(size, 1);
_poolSizeConfigurer = _ => Math.Max(size, 1);

return this;
}

public NatsBuilder WithPoolSize(Func<IServiceProvider, int> sizeConfigurer)
{
_poolSizeConfigurer = sp => Math.Max(sizeConfigurer(sp), 1);

return this;
}
Expand Down Expand Up @@ -120,43 +127,23 @@ public NatsBuilder WithSerializerRegistry(INatsSerializerRegistry registry)

internal IServiceCollection Build()
{
if (_poolSize != 1)
if (_diKey == null)
{
if (_diKey == null)
{
_services.TryAddSingleton<NatsConnectionPool>(provider => PoolFactory(provider));
_services.TryAddSingleton<INatsConnectionPool>(static provider => provider.GetRequiredService<NatsConnectionPool>());
_services.TryAddTransient<NatsConnection>(static provider => PooledConnectionFactory(provider, null));
_services.TryAddTransient<INatsConnection>(static provider => provider.GetRequiredService<NatsConnection>());
_services.TryAddTransient<INatsClient>(static provider => provider.GetRequiredService<NatsConnection>());
}
else
{
#if NET8_0_OR_GREATER
_services.TryAddKeyedSingleton<NatsConnectionPool>(_diKey, PoolFactory);
_services.TryAddKeyedSingleton<INatsConnectionPool>(_diKey, static (provider, key) => provider.GetRequiredKeyedService<NatsConnectionPool>(key));
_services.TryAddKeyedTransient<NatsConnection>(_diKey, PooledConnectionFactory);
_services.TryAddKeyedTransient<INatsConnection>(_diKey, static (provider, key) => provider.GetRequiredKeyedService<NatsConnection>(key));
_services.TryAddKeyedTransient<INatsClient>(_diKey, static (provider, key) => provider.GetRequiredKeyedService<NatsConnection>(key));
#endif
}
_services.TryAddSingleton<NatsConnectionPool>(provider => PoolFactory(provider));
_services.TryAddSingleton<INatsConnectionPool>(static provider => provider.GetRequiredService<NatsConnectionPool>());
_services.TryAddTransient<NatsConnection>(static provider => PooledConnectionFactory(provider, null));
_services.TryAddTransient<INatsConnection>(static provider => provider.GetRequiredService<NatsConnection>());
_services.TryAddTransient<INatsClient>(static provider => provider.GetRequiredService<NatsConnection>());
}
else
{
if (_diKey == null)
{
_services.TryAddSingleton<NatsConnection>(provider => SingleConnectionFactory(provider));
_services.TryAddSingleton<INatsConnection>(static provider => provider.GetRequiredService<NatsConnection>());
_services.TryAddSingleton<INatsClient>(static provider => provider.GetRequiredService<NatsConnection>());
}
else
{
#if NET8_0_OR_GREATER
_services.TryAddKeyedSingleton(_diKey, SingleConnectionFactory);
_services.TryAddKeyedSingleton<INatsConnection>(_diKey, static (provider, key) => provider.GetRequiredKeyedService<NatsConnection>(key));
_services.TryAddKeyedSingleton<INatsClient>(_diKey, static (provider, key) => provider.GetRequiredKeyedService<NatsConnection>(key));
_services.TryAddKeyedSingleton<NatsConnectionPool>(_diKey, PoolFactory);
_services.TryAddKeyedSingleton<INatsConnectionPool>(_diKey, static (provider, key) => provider.GetRequiredKeyedService<NatsConnectionPool>(key));
_services.TryAddKeyedTransient<NatsConnection>(_diKey, PooledConnectionFactory);
_services.TryAddKeyedTransient<INatsConnection>(_diKey, static (provider, key) => provider.GetRequiredKeyedService<NatsConnection>(key));
_services.TryAddKeyedTransient<INatsClient>(_diKey, static (provider, key) => provider.GetRequiredKeyedService<NatsConnection>(key));
#endif
}
}

return _services;
Expand All @@ -179,17 +166,7 @@ private NatsConnectionPool PoolFactory(IServiceProvider provider, object? diKey
{
var options = GetNatsOpts(provider);

return new NatsConnectionPool(_poolSize, options, con => _configureConnection?.Invoke(provider, con));
}

private NatsConnection SingleConnectionFactory(IServiceProvider provider, object? diKey = null)
{
var options = GetNatsOpts(provider);

var conn = new NatsConnection(options);
_configureConnection?.Invoke(provider, conn);

return conn;
return new NatsConnectionPool(_poolSizeConfigurer(provider), options, con => _configureConnection?.Invoke(provider, con));
}

private NatsOpts GetNatsOpts(IServiceProvider provider)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -177,6 +177,21 @@ public async Task AddNatsClient_WithDefaultSerializer()
}
}

[Fact]
public void AddNatsClient_RegistersNatsConnectionAsTransient_WhenPoolSizeFuncIsGreaterThanOne()
{
var services = new ServiceCollection();
services.AddSingleton<ILoggerFactory, NullLoggerFactory>();
services.AddNatsClient(nats => nats.WithPoolSize(_ => 2));

var provider = services.BuildServiceProvider();
var natsConnection1 = provider.GetRequiredService<INatsConnection>();
var natsConnection2 = provider.GetRequiredService<INatsConnection>();

Assert.NotNull(natsConnection1);
Assert.NotSame(natsConnection1, natsConnection2); // Transient should return different instances
}

[Fact]
public async Task AddNatsClient_WithJsonSerializer()
{
Expand Down Expand Up @@ -347,6 +362,56 @@ public void AddNats_RegistersKeyedNatsConnection_WhenKeyIsProvided_pooled()
Assert.NotSame(obj1, obj2);
}
}

[Fact]
public void AddNats_RegistersKeyedNatsConnection_WhenKeyIsProvided_pooledFunc()
{
var key1 = "TestKey1";
var key2 = "TestKey2";

var services = new ServiceCollection();
services.AddSingleton<ILoggerFactory, NullLoggerFactory>();

services.AddNatsClient(builder => builder.WithPoolSize(_ => 2).WithKey(key1));
services.AddNatsClient(builder => builder.WithPoolSize(_ => 2).WithKey(key2));
var provider = services.BuildServiceProvider();

Dictionary<string, List<object>> connections = new();
foreach (var key in new[] { key1, key2 })
{
var nats1 = provider.GetKeyedService<INatsConnection>(key);
Assert.NotNull(nats1);
var nats2 = provider.GetKeyedService<INatsConnection>(key);
Assert.NotNull(nats2);
var nats3 = provider.GetKeyedService<INatsConnection>(key);
Assert.NotNull(nats3);
var nats4 = provider.GetKeyedService<INatsConnection>(key);
Assert.NotNull(nats4);

// relying on the fact that the pool size is 2 and connections are returned in a round-robin fashion
Assert.NotSame(nats1, nats2);
Assert.Same(nats1, nats3);
Assert.NotSame(nats2, nats3);
Assert.Same(nats2, nats4);

if (!connections.TryGetValue(key, out var list))
{
list = new List<object>();
connections.Add(key, list);
}

list.Add(nats1);
list.Add(nats2);
list.Add(nats3);
list.Add(nats4);
}

foreach (var obj1 in connections[key1])
{
foreach (var obj2 in connections[key2])
Assert.NotSame(obj1, obj2);
}
}
#endif
}

Expand Down