Files
scadaproj/ZB.MOM.WW.LocalDb/tests/ZB.MOM.WW.LocalDb.Tests/GrpcAdapterTests.cs
T

413 lines
17 KiB
C#

using Grpc.Core;
using Grpc.Core.Interceptors;
using Grpc.Net.Client;
using Microsoft.AspNetCore.Builder;
using Microsoft.AspNetCore.Hosting;
using Microsoft.AspNetCore.TestHost;
using Microsoft.Extensions.Configuration;
using Microsoft.Extensions.DependencyInjection;
using Microsoft.Extensions.Hosting;
using ZB.MOM.WW.LocalDb.Contracts;
using ZB.MOM.WW.LocalDb.Replication;
namespace ZB.MOM.WW.LocalDb.Tests;
public sealed class GrpcAdapterTests : IAsyncLifetime
{
private static readonly TimeSpan RunTimeout = TimeSpan.FromSeconds(30);
private readonly List<string> _paths = [];
private readonly List<IHost> _hosts = [];
private readonly List<ServiceProvider> _providers = [];
public Task InitializeAsync() => Task.CompletedTask;
public async Task DisposeAsync()
{
foreach (var provider in _providers)
await provider.DisposeAsync();
foreach (var host in _hosts)
{
try { await host.StopAsync(TimeSpan.FromSeconds(5)); } catch { /* teardown */ }
host.Dispose();
}
Microsoft.Data.Sqlite.SqliteConnection.ClearAllPools();
foreach (var p in _paths)
{
if (File.Exists(p)) File.Delete(p);
if (File.Exists(p + "-wal")) File.Delete(p + "-wal");
if (File.Exists(p + "-shm")) File.Delete(p + "-shm");
}
}
// ---- end-to-end wire tests -------------------------------------------------------------
[Fact]
public async Task EndToEnd_TwoDbs_InsertOnA_AppearsOnB()
{
var serverHost = await BuildServerAsync(NewDbPath(), apiKey: null, sink: null);
var serverDb = serverHost.Services.GetRequiredService<ILocalDb>();
await serverDb.ExecuteAsync("INSERT INTO orders (id, sku, qty) VALUES (2, 'FROM_SERVER', 20)");
var handler = serverHost.GetTestServer().CreateHandler();
var (clientProvider, bg) = BuildClient(NewDbPath(), "http://localhost", apiKey: null, () => ChannelOver(handler));
var clientDb = clientProvider.GetRequiredService<ILocalDb>();
await clientDb.ExecuteAsync("INSERT INTO orders (id, sku, qty) VALUES (1, 'FROM_CLIENT', 10)");
using var cts = new CancellationTokenSource(RunTimeout);
await bg.StartAsync(cts.Token);
await WaitForAsync(
async () => (await ReadOrders(clientDb)).Count == 2 && (await ReadOrders(serverDb)).Count == 2,
RunTimeout);
var clientRows = await ReadOrders(clientDb);
var serverRows = await ReadOrders(serverDb);
Assert.Equal(clientRows, serverRows);
Assert.Equal([(1L, "FROM_CLIENT", (long?)10), (2L, "FROM_SERVER", (long?)20)], clientRows);
await bg.StopAsync(CancellationToken.None);
}
[Fact]
public async Task SecondConcurrentStream_Rejected()
{
var serverHost = await BuildServerAsync(NewDbPath(), apiKey: null, sink: null);
var serverDb = serverHost.Services.GetRequiredService<ILocalDb>();
var handler = serverHost.GetTestServer().CreateHandler();
var (clientProvider, bg) = BuildClient(NewDbPath(), "http://localhost", apiKey: null, () => ChannelOver(handler));
var clientDb = clientProvider.GetRequiredService<ILocalDb>();
await clientDb.ExecuteAsync("INSERT INTO orders (id, sku, qty) VALUES (1, 'C', 10)");
using var cts = new CancellationTokenSource(RunTimeout);
await bg.StartAsync(cts.Token);
// Convergence proves the first stream's server handler is active and holding the single-stream guard.
await WaitForAsync(async () => (await ReadOrders(serverDb)).Count == 1, RunTimeout);
var client = new LocalDbSync.LocalDbSyncClient(ChannelOver(handler));
using var call2 = client.Sync(cancellationToken: cts.Token);
var ex = await Assert.ThrowsAsync<RpcException>(() => call2.ResponseStream.MoveNext(cts.Token));
Assert.Equal(StatusCode.AlreadyExists, ex.StatusCode);
await bg.StopAsync(CancellationToken.None);
}
// Shipped variant: the channel factory throws on its first call, then returns a live channel.
// This deterministically drives the backoff/retry loop (no flaky server restart in TestHost)
// and proves convergence after recovery.
[Fact]
public async Task ClientReconnects_AfterTransientFault()
{
var serverHost = await BuildServerAsync(NewDbPath(), apiKey: null, sink: null);
var serverDb = serverHost.Services.GetRequiredService<ILocalDb>();
await serverDb.ExecuteAsync("INSERT INTO orders (id, sku, qty) VALUES (2, 'S', 20)");
var handler = serverHost.GetTestServer().CreateHandler();
var calls = 0;
GrpcChannel Factory()
{
if (Interlocked.Increment(ref calls) == 1)
throw new InvalidOperationException("transient dial failure");
return ChannelOver(handler);
}
var (clientProvider, bg) = BuildClient(NewDbPath(), "http://localhost", apiKey: null, Factory);
var clientDb = clientProvider.GetRequiredService<ILocalDb>();
await clientDb.ExecuteAsync("INSERT INTO orders (id, sku, qty) VALUES (1, 'C', 10)");
using var cts = new CancellationTokenSource(RunTimeout);
await bg.StartAsync(cts.Token);
await WaitForAsync(
async () => (await ReadOrders(clientDb)).Count == 2 && (await ReadOrders(serverDb)).Count == 2,
RunTimeout);
Assert.True(bg.ConnectionAttempts >= 2, $"expected >= 2 connection attempts, got {bg.ConnectionAttempts}");
await bg.StopAsync(CancellationToken.None);
}
[Fact]
public async Task ApiKeyHeader_SentWhenConfigured()
{
var sink = new AuthCaptureSink();
var serverHost = await BuildServerAsync(NewDbPath(), apiKey: "test-key", sink: sink);
var handler = serverHost.GetTestServer().CreateHandler();
var (clientProvider, bg) = BuildClient(NewDbPath(), "http://localhost", apiKey: "test-key", () => ChannelOver(handler));
var clientDb = clientProvider.GetRequiredService<ILocalDb>();
await clientDb.ExecuteAsync("INSERT INTO orders (id, sku, qty) VALUES (1, 'C', 10)");
using var cts = new CancellationTokenSource(RunTimeout);
await bg.StartAsync(cts.Token);
await WaitForAsync(() => Task.FromResult(sink.Authorization is not null), RunTimeout);
Assert.Equal("Bearer test-key", sink.Authorization);
await bg.StopAsync(CancellationToken.None);
}
// Shipped variant: in-proc against GrpcSyncDuplex with a flow-control-stalled fake writer
// (WriteAsync never completes until its token fires) — pins requirement 2 (Send honors its
// CancellationToken even when transport-blocked) deterministically; a true stalled-gRPC-peer
// rig in TestHost cannot force flow-control blockage reliably.
[Fact]
public async Task Send_StalledOnFlowControl_CancelUnblocksPromptly()
{
using var cts = new CancellationTokenSource();
var (duplex, readerTask) = Replication.Internal.GrpcSyncDuplex.Create(
new HangingReader(), new HangingWriter(), cts.Token);
var send = duplex.Send(new SyncMessage(), cts.Token);
Assert.False(send.IsCompleted);
var sw = System.Diagnostics.Stopwatch.StartNew();
cts.Cancel();
await Assert.ThrowsAnyAsync<OperationCanceledException>(() => send);
// The reader pump must also tear down: its fault travels into the inbox completion.
await readerTask.WaitAsync(TimeSpan.FromSeconds(5));
await Assert.ThrowsAnyAsync<OperationCanceledException>(
async () => await duplex.Inbox.Completion);
sw.Stop();
Assert.True(sw.Elapsed < TimeSpan.FromSeconds(3), $"teardown took {sw.Elapsed}");
}
// ---- validator -------------------------------------------------------------------------
[Fact]
public void Validator_InvalidPeerAddress_Fails()
{
var result = new ReplicationOptionsValidator().Validate(
null, new ReplicationOptions { PeerAddress = "not-a-uri" });
Assert.True(result.Failed);
Assert.Contains(result.Failures, f => f.Contains("PeerAddress", StringComparison.Ordinal));
}
[Fact]
public void Validator_NonPositiveIntervals_Fail()
{
var result = new ReplicationOptionsValidator().Validate(null, new ReplicationOptions
{
FlushInterval = TimeSpan.Zero,
MaxBatchSize = 0,
MaxOplogRows = 0,
MaxOplogAge = TimeSpan.Zero,
TombstoneRetention = TimeSpan.Zero,
ReconnectBackoffMax = TimeSpan.Zero,
MaxHlcDriftAhead = TimeSpan.Zero,
});
Assert.True(result.Failed);
Assert.Contains(result.Failures, f => f.Contains("FlushInterval", StringComparison.Ordinal));
Assert.Contains(result.Failures, f => f.Contains("MaxBatchSize", StringComparison.Ordinal));
Assert.Contains(result.Failures, f => f.Contains("MaxOplogRows", StringComparison.Ordinal));
Assert.Contains(result.Failures, f => f.Contains("MaxOplogAge", StringComparison.Ordinal));
Assert.Contains(result.Failures, f => f.Contains("TombstoneRetention", StringComparison.Ordinal));
Assert.Contains(result.Failures, f => f.Contains("ReconnectBackoffMax", StringComparison.Ordinal));
Assert.Contains(result.Failures, f => f.Contains("MaxHlcDriftAhead", StringComparison.Ordinal));
}
[Fact]
public void Validator_EmptyPeerAddress_Allowed()
{
// A passive node leaves PeerAddress empty; the validator can't know the role, so it only
// checks format when present. All other defaults are valid.
var result = new ReplicationOptionsValidator().Validate(null, new ReplicationOptions { PeerAddress = "" });
Assert.True(result.Succeeded);
}
[Fact]
public void Validator_ReconnectBackoffMax_CappedAtOneDay()
{
var atCeiling = new ReplicationOptionsValidator().Validate(
null, new ReplicationOptions { ReconnectBackoffMax = TimeSpan.FromDays(1) });
Assert.True(atCeiling.Succeeded);
var overCeiling = new ReplicationOptionsValidator().Validate(
null, new ReplicationOptions { ReconnectBackoffMax = TimeSpan.FromDays(1) + TimeSpan.FromTicks(1) });
Assert.True(overCeiling.Failed);
Assert.Contains(overCeiling.Failures, f => f.Contains("ReconnectBackoffMax", StringComparison.Ordinal));
}
// ---- DI --------------------------------------------------------------------------------
[Fact]
public async Task AddZbLocalDbReplication_PassiveConfig_InitiatorIdles()
{
var passive = new ServiceCollection();
passive.AddLogging();
var passiveCfg = ConfigFor(NewDbPath(), peerAddress: null, apiKey: null);
passive.AddZbLocalDb(passiveCfg, OnReady);
passive.AddZbLocalDbReplication(passiveCfg);
var p1 = passive.BuildServiceProvider();
_providers.Add(p1);
Assert.NotNull(p1.GetService<LocalDbSyncService>());
// The hosted service is always registered; with no PeerAddress it starts, idles, and
// completes without dialing anything.
var bg = p1.GetServices<IHostedService>().OfType<SyncBackgroundService>().Single();
await bg.StartAsync(CancellationToken.None);
await bg.ExecuteTask!.WaitAsync(TimeSpan.FromSeconds(5));
Assert.True(bg.ExecuteTask.IsCompletedSuccessfully);
Assert.Equal(0, bg.ConnectionAttempts);
await bg.StopAsync(CancellationToken.None);
var initiator = new ServiceCollection();
initiator.AddLogging();
var initiatorCfg = ConfigFor(NewDbPath(), peerAddress: "http://localhost", apiKey: null);
initiator.AddZbLocalDb(initiatorCfg, OnReady);
initiator.AddZbLocalDbReplication(initiatorCfg);
var p2 = initiator.BuildServiceProvider();
_providers.Add(p2);
Assert.Contains(p2.GetServices<IHostedService>(), s => s is SyncBackgroundService);
}
// ---- harness ---------------------------------------------------------------------------
private string NewDbPath()
{
var path = Path.Combine(Path.GetTempPath(), Guid.NewGuid() + ".db");
_paths.Add(path);
return path;
}
private static void OnReady(ILocalDb db)
{
using var conn = db.CreateConnection();
using var cmd = conn.CreateCommand();
cmd.CommandText = "CREATE TABLE orders (id INTEGER PRIMARY KEY, sku TEXT, qty INTEGER)";
cmd.ExecuteNonQuery();
db.RegisterReplicated("orders");
}
private static IConfiguration ConfigFor(string path, string? peerAddress, string? apiKey)
{
var dict = new Dictionary<string, string?>
{
["LocalDb:Path"] = path,
["LocalDb:Replication:FlushInterval"] = "00:00:00.020",
};
if (peerAddress is not null) dict["LocalDb:Replication:PeerAddress"] = peerAddress;
if (apiKey is not null) dict["LocalDb:Replication:ApiKey"] = apiKey;
return new ConfigurationBuilder().AddInMemoryCollection(dict).Build();
}
private static GrpcChannel ChannelOver(HttpMessageHandler handler) =>
GrpcChannel.ForAddress("http://localhost", new GrpcChannelOptions { HttpHandler = handler });
private async Task<IHost> BuildServerAsync(string path, string? apiKey, AuthCaptureSink? sink)
{
var config = ConfigFor(path, peerAddress: null, apiKey: apiKey);
var host = await new HostBuilder()
.ConfigureWebHost(web =>
{
web.UseTestServer();
web.ConfigureServices(services =>
{
services.AddRouting();
if (sink is not null)
{
services.AddSingleton(sink);
services.AddGrpc(o => o.Interceptors.Add<AuthCaptureInterceptor>());
}
else
{
services.AddGrpc();
}
services.AddZbLocalDb(config, OnReady);
services.AddZbLocalDbReplication(config);
});
web.Configure(app =>
{
app.UseRouting();
app.UseEndpoints(e => e.MapZbLocalDbSync());
});
})
.StartAsync();
_hosts.Add(host);
return host;
}
private (ServiceProvider Provider, SyncBackgroundService Bg) BuildClient(
string path, string peerAddress, string? apiKey, Func<GrpcChannel> channelFactory)
{
var config = ConfigFor(path, peerAddress, apiKey);
var services = new ServiceCollection();
services.AddLogging();
services.AddZbLocalDb(config, OnReady);
services.AddZbLocalDbReplication(config);
var provider = services.BuildServiceProvider();
_providers.Add(provider);
var bg = provider.GetServices<IHostedService>().OfType<SyncBackgroundService>().Single();
bg.ChannelFactory = channelFactory;
return (provider, bg);
}
private static async Task<List<(long Id, string? Sku, long? Qty)>> ReadOrders(ILocalDb db)
{
var rows = await db.QueryAsync(
"SELECT id, sku, qty FROM orders ORDER BY id",
x => (x.GetInt64(0), x.IsDBNull(1) ? null : x.GetString(1), (long?)(x.IsDBNull(2) ? null : x.GetInt64(2))));
return rows.ToList();
}
private static async Task WaitForAsync(Func<Task<bool>> predicate, TimeSpan timeout)
{
var deadline = DateTime.UtcNow + timeout;
while (DateTime.UtcNow < deadline)
{
if (await predicate()) return;
await Task.Delay(20);
}
throw new TimeoutException("Condition not reached within " + timeout);
}
private sealed class HangingReader : IAsyncStreamReader<SyncMessage>
{
public SyncMessage Current => throw new InvalidOperationException("no message");
public async Task<bool> MoveNext(CancellationToken cancellationToken)
{
await Task.Delay(Timeout.Infinite, cancellationToken);
return false;
}
}
private sealed class HangingWriter : IAsyncStreamWriter<SyncMessage>
{
public WriteOptions? WriteOptions { get; set; }
public Task WriteAsync(SyncMessage message) => Task.Delay(Timeout.Infinite);
public Task WriteAsync(SyncMessage message, CancellationToken cancellationToken) =>
Task.Delay(Timeout.Infinite, cancellationToken);
}
private sealed class AuthCaptureSink
{
public volatile string? Authorization;
}
private sealed class AuthCaptureInterceptor(AuthCaptureSink sink) : Interceptor
{
public override Task DuplexStreamingServerHandler<TRequest, TResponse>(
IAsyncStreamReader<TRequest> requestStream,
IServerStreamWriter<TResponse> responseStream,
ServerCallContext context,
DuplexStreamingServerMethod<TRequest, TResponse> continuation)
{
var auth = context.RequestHeaders.GetValue("authorization");
if (auth is not null)
sink.Authorization = auth;
return continuation(requestStream, responseStream, context);
}
}
}