Skip to content

Commit 7f18869

Browse files
authored
Allow sending queries from INpgsqlDatabaseInfoFactory.Load (#6634)
Fixes #6633
1 parent cd209e5 commit 7f18869

7 files changed

Lines changed: 83 additions & 20 deletions

File tree

src/Npgsql/Internal/INpgsqlDatabaseInfoFactory.cs

Lines changed: 15 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,6 @@
1+
using System;
12
using System.Diagnostics.CodeAnalysis;
3+
using System.Threading;
24
using System.Threading.Tasks;
35
using Npgsql.Util;
46

@@ -20,5 +22,17 @@ public interface INpgsqlDatabaseInfoFactory
2022
/// An object describing the database to which <paramref name="conn"/> is connected, or null if the
2123
/// database isn't of the correct type and isn't handled by this factory.
2224
/// </returns>
23-
Task<NpgsqlDatabaseInfo?> Load(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async);
25+
Task<NpgsqlDatabaseInfo?> Load(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async)
26+
=> throw new NotImplementedException();
27+
28+
/// <summary>
29+
/// Given a connection, loads all necessary information about the connected database, e.g. its types.
30+
/// A factory should only handle the exact database type it was meant for, and return null otherwise.
31+
/// </summary>
32+
/// <returns>
33+
/// An object describing the database to which <paramref name="conn"/> is connected, or null if the
34+
/// database isn't of the correct type and isn't handled by this factory.
35+
/// </returns>
36+
Task<NpgsqlDatabaseInfo?> Load(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async, CancellationToken cancellationToken)
37+
=> Load(conn, timeout, async);
2438
}

src/Npgsql/Internal/NpgsqlDatabaseInfo.cs

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
using System;
22
using System.Collections.Generic;
33
using System.Diagnostics.CodeAnalysis;
4+
using System.Threading;
45
using System.Threading.Tasks;
56
using Npgsql.Internal.Postgres;
67
using Npgsql.PostgresTypes;
@@ -321,11 +322,11 @@ public static void RegisterFactory(INpgsqlDatabaseInfoFactory factory)
321322
Factories = factories;
322323
}
323324

324-
internal static async Task<NpgsqlDatabaseInfo> Load(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async)
325+
internal static async Task<NpgsqlDatabaseInfo> Load(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async, CancellationToken cancellationToken)
325326
{
326327
foreach (var factory in Factories)
327328
{
328-
var dbInfo = await factory.Load(conn, timeout, async).ConfigureAwait(false);
329+
var dbInfo = await factory.Load(conn, timeout, async, cancellationToken).ConfigureAwait(false);
329330
if (dbInfo != null)
330331
{
331332
dbInfo.ProcessTypes();

src/Npgsql/NpgsqlDataSource.cs

Lines changed: 1 addition & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -311,10 +311,7 @@ internal async Task Bootstrap(
311311
},
312312
dbTypeResolver: null);
313313

314-
NpgsqlDatabaseInfo databaseInfo;
315-
316-
using (connector.StartUserAction(ConnectorState.Executing, cancellationToken))
317-
databaseInfo = await NpgsqlDatabaseInfo.Load(connector, timeout, async).ConfigureAwait(false);
314+
var databaseInfo = await NpgsqlDatabaseInfo.Load(connector, timeout, async, cancellationToken).ConfigureAwait(false);
318315

319316
var serializerOptions = new PgSerializerOptions(databaseInfo, _resolverChain, CreateTimeZoneProvider(connector.Timezone), textEncoding: connector.TextEncoding)
320317
{

src/Npgsql/PostgresDatabaseInfo.cs

Lines changed: 14 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -3,6 +3,7 @@
33
using System.Diagnostics;
44
using System.Globalization;
55
using System.Text;
6+
using System.Threading;
67
using System.Threading.Tasks;
78
using Microsoft.Extensions.Logging;
89
using Microsoft.Extensions.Logging.Abstractions;
@@ -24,10 +25,10 @@ namespace Npgsql;
2425
sealed class PostgresDatabaseInfoFactory : INpgsqlDatabaseInfoFactory
2526
{
2627
/// <inheritdoc />
27-
public async Task<NpgsqlDatabaseInfo?> Load(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async)
28+
public async Task<NpgsqlDatabaseInfo?> Load(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async, CancellationToken cancellationToken)
2829
{
2930
var db = new PostgresDatabaseInfo(conn);
30-
await db.LoadPostgresInfo(conn, timeout, async).ConfigureAwait(false);
31+
await db.LoadPostgresInfo(conn, timeout, async, cancellationToken).ConfigureAwait(false);
3132
Debug.Assert(db.LongVersion != null);
3233
return db;
3334
}
@@ -80,7 +81,7 @@ class PostgresDatabaseInfo : NpgsqlDatabaseInfo
8081
public virtual bool HasTypeCategory => Version.IsGreaterOrEqual(8, 4);
8182

8283
internal PostgresDatabaseInfo(NpgsqlConnector conn)
83-
: base(conn.Host!, conn.Port, conn.Database!, conn.PostgresParameters["server_version"])
84+
: base(conn.Host, conn.Port, conn.Database!, conn.PostgresParameters["server_version"])
8485
=> _connectionLogger = conn.LoggingConfiguration.ConnectionLogger;
8586

8687
private protected PostgresDatabaseInfo(string host, int port, string databaseName, string serverVersion)
@@ -93,16 +94,20 @@ private protected PostgresDatabaseInfo(string host, int port, string databaseNam
9394
/// <param name="conn">The database connection.</param>
9495
/// <param name="timeout">The timeout while loading types from the backend.</param>
9596
/// <param name="async">True to load types asynchronously.</param>
97+
/// <param name="cancellationToken">An optional token to cancel the asynchronous operation. The default value is <see cref="CancellationToken.None"/>.</param>
9698
/// <returns>
9799
/// A task representing the asynchronous operation.
98100
/// </returns>
99-
internal async Task LoadPostgresInfo(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async)
101+
internal async Task LoadPostgresInfo(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async, CancellationToken cancellationToken)
100102
{
101-
HasIntegerDateTimes =
102-
conn.PostgresParameters.TryGetValue("integer_datetimes", out var intDateTimes) &&
103-
intDateTimes == "on";
103+
using (conn.StartUserAction(ConnectorState.Executing, cancellationToken))
104+
{
105+
HasIntegerDateTimes =
106+
conn.PostgresParameters.TryGetValue("integer_datetimes", out var intDateTimes) &&
107+
intDateTimes == "on";
104108

105-
_types = await LoadBackendTypes(conn, timeout, async).ConfigureAwait(false);
109+
_types = await LoadBackendTypes(conn, timeout, async).ConfigureAwait(false);
110+
}
106111
}
107112

108113
const string BuiltinSchemaListSqlFragment = "'pg_catalog', 'information_schema', 'pg_toast'";
@@ -205,7 +210,7 @@ FROM pg_enum
205210
/// </returns>
206211
/// <exception cref="TimeoutException" />
207212
/// <exception cref="ArgumentOutOfRangeException">Unknown typtype for type '{internalName}' in pg_type: {typeChar}.</exception>
208-
internal async Task<List<PostgresType>> LoadBackendTypes(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async)
213+
async Task<List<PostgresType>> LoadBackendTypes(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async)
209214
{
210215
var versionQuery = "SELECT version();";
211216
var typeLoading = conn.DataSource.Configuration.TypeLoading;

src/Npgsql/PostgresMinimalDatabaseInfo.cs

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,4 +1,5 @@
11
using System.Collections.Generic;
2+
using System.Threading;
23
using System.Threading.Tasks;
34
using Npgsql.Internal;
45
using Npgsql.Internal.Postgres;
@@ -9,7 +10,7 @@ namespace Npgsql;
910

1011
sealed class PostgresMinimalDatabaseInfoFactory : INpgsqlDatabaseInfoFactory
1112
{
12-
public Task<NpgsqlDatabaseInfo?> Load(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async)
13+
public Task<NpgsqlDatabaseInfo?> Load(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async, CancellationToken cancellationToken)
1314
=> Task.FromResult(
1415
!conn.DataSource.Configuration.TypeLoading.LoadTypes
1516
? (NpgsqlDatabaseInfo)new PostgresMinimalDatabaseInfo(conn)

test/Npgsql.Tests/ConnectionTests.cs

Lines changed: 45 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1902,10 +1902,54 @@ public async Task Breaking_connection_while_loading_database_info()
19021902

19031903
class BreakingDatabaseInfoFactory : INpgsqlDatabaseInfoFactory
19041904
{
1905-
public Task<NpgsqlDatabaseInfo?> Load(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async)
1905+
public Task<NpgsqlDatabaseInfo?> Load(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async, CancellationToken cancellationToken)
19061906
=> throw conn.Break(new IOException());
19071907
}
19081908

1909+
[Test]
1910+
[NonParallelizable] // Modifies global database info factories
1911+
[IssueLink("https://github.com/npgsql/npgsql/issues/6633")]
1912+
public async Task Allow_to_query_while_loading_database_info([Values] bool async)
1913+
{
1914+
await using var dataSource = CreateDataSource();
1915+
1916+
var factory = new QueryingDatabaseInfoFactory();
1917+
NpgsqlDatabaseInfo.RegisterFactory(factory);
1918+
try
1919+
{
1920+
await using var _ = async
1921+
? await dataSource.OpenConnectionAsync()
1922+
: dataSource.OpenConnection();
1923+
Assert.That(factory.QueryExecuted, Is.True);
1924+
}
1925+
finally
1926+
{
1927+
NpgsqlDatabaseInfo.ResetFactories();
1928+
}
1929+
}
1930+
1931+
class QueryingDatabaseInfoFactory : INpgsqlDatabaseInfoFactory
1932+
{
1933+
public async Task<NpgsqlDatabaseInfo?> Load(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async,
1934+
CancellationToken cancellationToken)
1935+
{
1936+
using var cmd = conn.CreateCommand();
1937+
cmd.CommandText = "SELECT 1";
1938+
1939+
var result = async
1940+
? await cmd.ExecuteScalarAsync(cancellationToken)
1941+
: cmd.ExecuteScalar();
1942+
1943+
Assert.That(result, Is.EqualTo(1));
1944+
1945+
QueryExecuted = true;
1946+
1947+
return null;
1948+
}
1949+
1950+
public bool QueryExecuted { get; private set; }
1951+
}
1952+
19091953
#region Logging tests
19101954

19111955
[Test]

test/Npgsql.Tests/TransactionTests.cs

Lines changed: 3 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,6 @@
11
using System;
22
using System.Data;
3+
using System.Threading;
34
using System.Threading.Tasks;
45
using Npgsql.Internal;
56
using Npgsql.Tests.Support;
@@ -682,10 +683,10 @@ public async Task Bug3686()
682683

683684
class NoTransactionDatabaseInfoFactory : INpgsqlDatabaseInfoFactory
684685
{
685-
public async Task<NpgsqlDatabaseInfo?> Load(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async)
686+
public async Task<NpgsqlDatabaseInfo?> Load(NpgsqlConnector conn, NpgsqlTimeout timeout, bool async, CancellationToken cancellationToken)
686687
{
687688
var db = new NoTransactionDatabaseInfo(conn);
688-
await db.LoadPostgresInfo(conn, timeout, async);
689+
await db.LoadPostgresInfo(conn, timeout, async, cancellationToken);
689690
return db;
690691
}
691692
}

0 commit comments

Comments
 (0)