Skip to content
Draft
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
38 changes: 37 additions & 1 deletion src/Npgsql/NpgsqlConnectionStringBuilder.cs
Original file line number Diff line number Diff line change
Expand Up @@ -212,6 +212,32 @@ void SetValue(string propertyName, object? value)

#region Properties - Connection

/// <summary>
/// DNS SRV cluster domain for service discovery. When set, the driver looks up
/// <c>_postgresql._tcp.&lt;SrvHost&gt;</c> SRV records at <see cref="NpgsqlDataSource"/> build
/// time and uses the returned host/port pairs — sorted by priority then weight per RFC 2782 —
/// as the list of hosts to connect to.
/// <para>
/// Can be set via the <c>SrvHost=cluster.example.com</c> connection property.
/// Mutually exclusive with an explicit <see cref="Host"/> value.
/// </para>
/// </summary>
[Category("Connection")]
[Description("DNS SRV domain for PostgreSQL service discovery (_postgresql._tcp.<SrvHost>).")]
[DisplayName("SrvHost")]
[NpgsqlConnectionStringProperty]
public string? SrvHost
{
get => _srvHost;
set
{
_srvHost = value;
SetValue(nameof(SrvHost), value);
_dataSourceCached = null;
}
}
string? _srvHost;

/// <summary>
/// The hostname or IP address of the PostgreSQL server to connect to.
/// </summary>
Expand Down Expand Up @@ -1455,10 +1481,20 @@ public int InternalCommandTimeout

internal void PostProcessAndValidate()
{
ArgumentException.ThrowIfNullOrWhiteSpace(Host);
if (!string.IsNullOrWhiteSpace(SrvHost) && !string.IsNullOrWhiteSpace(Host))
throw new ArgumentException(
"SrvHost and Host are mutually exclusive. Use SrvHost for DNS SRV discovery or Host for direct connections, not both.");

// When SrvHost is set, Host will be populated by SRV resolution in the DataSource builder.
if (string.IsNullOrWhiteSpace(SrvHost))
ArgumentException.ThrowIfNullOrWhiteSpace(Host);

if (SslNegotiation == SslNegotiation.Direct && SslMode is not SslMode.Require and not SslMode.VerifyCA and not SslMode.VerifyFull)
throw new ArgumentException("SSL Mode has to be Require or higher to be used with direct SSL Negotiation");

if (string.IsNullOrWhiteSpace(Host))
return; // SRV resolution not done yet, skip host-based checks

if (!Host.Contains(','))
{
if (TargetSessionAttributesParsed is not null &&
Expand Down
1 change: 1 addition & 0 deletions src/Npgsql/NpgsqlDataSourceBuilder.cs
Original file line number Diff line number Diff line change
Expand Up @@ -589,6 +589,7 @@ public NpgsqlDataSourceBuilder UsePhysicalConnectionInitializer(
return this;
}


/// <summary>
/// Builds and returns an <see cref="NpgsqlDataSource" /> which is ready for use.
/// </summary>
Expand Down
11 changes: 11 additions & 0 deletions src/Npgsql/NpgsqlSlimDataSourceBuilder.cs
Original file line number Diff line number Diff line change
Expand Up @@ -789,6 +789,7 @@ public NpgsqlSlimDataSourceBuilder UsePhysicalConnectionInitializer(
/// </summary>
public NpgsqlDataSource Build()
{
ResolveSrvIfNeeded();
ConnectionStringBuilder.PostProcessAndValidate();
var (connectionStringBuilder, config) = PrepareConfiguration();

Expand All @@ -809,6 +810,7 @@ public NpgsqlDataSource Build()
/// </summary>
public NpgsqlMultiHostDataSource BuildMultiHost()
{
ResolveSrvIfNeeded();
ConnectionStringBuilder.PostProcessAndValidate();
var (connectionStringBuilder, config) = PrepareConfiguration();

Expand All @@ -817,6 +819,15 @@ public NpgsqlMultiHostDataSource BuildMultiHost()
return new(connectionStringBuilder, config);
}

void ResolveSrvIfNeeded()
{
var srvHost = ConnectionStringBuilder.SrvHost;
if (string.IsNullOrWhiteSpace(srvHost))
return;

ConnectionStringBuilder.Host = SrvLookup.ResolveToHostString(srvHost);
}

// Used in testing.
internal (NpgsqlConnectionStringBuilder, NpgsqlDataSourceConfiguration) PrepareConfiguration()
{
Expand Down
7 changes: 7 additions & 0 deletions src/Npgsql/Properties/AssemblyInfo.cs
Original file line number Diff line number Diff line change
Expand Up @@ -24,6 +24,13 @@
"7aa16153bcea2ae9a471145624826f60d7c8e71cd025b554a0177bd935a78096" +
"29f0a7afc778ebb4ad033e1bf512c1a9c6ceea26b077bc46cac93800435e77ee")]

[assembly: InternalsVisibleTo("Npgsql.SrvTests, PublicKey=" +
"0024000004800000940000000602000000240000525341310004000001000100" +
"2b3c590b2a4e3d347e6878dc0ff4d21eb056a50420250c6617044330701d35c9" +
"8078a5df97a62d83c9a2db2d072523a8fc491398254c6b89329b8c1dcef43a1e" +
"7aa16153bcea2ae9a471145624826f60d7c8e71cd025b554a0177bd935a78096" +
"29f0a7afc778ebb4ad033e1bf512c1a9c6ceea26b077bc46cac93800435e77ee")]

[assembly: InternalsVisibleTo("Npgsql.Benchmarks, PublicKey=" +
"0024000004800000940000000602000000240000525341310004000001000100" +
"2b3c590b2a4e3d347e6878dc0ff4d21eb056a50420250c6617044330701d35c9" +
Expand Down
2 changes: 2 additions & 0 deletions src/Npgsql/PublicAPI.Unshipped.txt
Original file line number Diff line number Diff line change
@@ -1,4 +1,6 @@
#nullable enable
Npgsql.NpgsqlConnectionStringBuilder.SrvHost.get -> string?
Npgsql.NpgsqlConnectionStringBuilder.SrvHost.set -> void
*REMOVED*Npgsql.NpgsqlConnectionStringBuilder.Multiplexing.get -> bool
*REMOVED*Npgsql.NpgsqlConnectionStringBuilder.Multiplexing.set -> void
*REMOVED*Npgsql.NpgsqlConnectionStringBuilder.WriteCoalescingBufferThresholdBytes.get -> int
Expand Down
227 changes: 227 additions & 0 deletions src/Npgsql/SrvLookup.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,227 @@
using System;
using System.Collections.Generic;
using System.Net;
using System.Net.NetworkInformation;
using System.Net.Sockets;
using System.Text;
using System.Threading;
using System.Threading.Tasks;

namespace Npgsql;

/// <summary>
/// Resolves DNS SRV records for PostgreSQL service discovery without external dependencies.
/// Looks up <c>_postgresql._tcp.&lt;srvHost&gt;</c> and returns host:port pairs
/// sorted by priority ascending then weight descending per RFC 2782.
/// </summary>
static class SrvLookup
{
const string ServicePrefix = "_postgresql._tcp.";
const int DnsPort = 53;
const int TypeSRV = 33;
const int ClassIN = 1;
const int TimeoutMs = 5000;

/// <summary>Internal SRV record representation used in unit tests.</summary>
internal readonly record struct SrvRecord(ushort Priority, ushort Weight, ushort Port, string Target);

internal static string ResolveToHostString(string srvHost)
{
var qname = ServicePrefix + srvHost;
var records = QuerySrv(qname);
return SortAndBuild(records) ?? throw new NpgsqlException($"No SRV records found for {qname}");
}

internal static async Task<string> ResolveToHostStringAsync(
string srvHost,
CancellationToken cancellationToken = default)
{
var qname = ServicePrefix + srvHost;
var records = await QuerySrvAsync(qname, cancellationToken).ConfigureAwait(false);
return SortAndBuild(records) ?? throw new NpgsqlException($"No SRV records found for {qname}");
}

/// <summary>
/// Sorts <paramref name="records"/> by priority ascending then weight descending (RFC 2782)
/// and returns a comma-separated <c>host:port</c> string, or <see langword="null"/> when empty.
/// </summary>
/// <remarks>Internal so unit tests can exercise sorting without a live DNS server.</remarks>
internal static string? SortAndBuild(IEnumerable<SrvRecord> records)
{
var list = new List<SrvRecord>(records);
if (list.Count == 0)
return null;

list.Sort((a, b) =>
{
var cmp = a.Priority.CompareTo(b.Priority);
return cmp != 0 ? cmp : b.Weight.CompareTo(a.Weight);
});

var sb = new StringBuilder();
for (var i = 0; i < list.Count; i++)
{
if (i > 0) sb.Append(',');
sb.Append(list[i].Target).Append(':').Append(list[i].Port);
}
return sb.ToString();
}

static List<SrvRecord> QuerySrv(string qname)
{
var query = BuildDnsQuery(qname);
Exception? last = null;
foreach (var server in GetSystemDnsServers())
{
try
{
using var udp = new UdpClient(server.AddressFamily);
udp.Client.SendTimeout = TimeoutMs;
udp.Client.ReceiveTimeout = TimeoutMs;
udp.Connect(server, DnsPort);
udp.Send(query, query.Length);
var remote = new IPEndPoint(IPAddress.Any, 0);
return ParseDnsResponse(qname, udp.Receive(ref remote));
}
catch (NpgsqlException) { throw; }
catch (Exception ex) { last = ex; }
}
throw new NpgsqlException($"DNS SRV lookup failed for {qname}", last);
}

static async Task<List<SrvRecord>> QuerySrvAsync(string qname, CancellationToken ct)
{
var query = BuildDnsQuery(qname);
Exception? last = null;
foreach (var server in GetSystemDnsServers())
{
try
{
using var udp = new UdpClient(server.AddressFamily);
udp.Connect(server, DnsPort);
await udp.SendAsync(new ReadOnlyMemory<byte>(query), ct).ConfigureAwait(false);
var result = await udp.ReceiveAsync(ct).ConfigureAwait(false);
return ParseDnsResponse(qname, result.Buffer);
}
catch (NpgsqlException) { throw; }
catch (Exception ex) { last = ex; }
}
throw new NpgsqlException($"DNS SRV lookup failed for {qname}", last);
}

static IPAddress[] GetSystemDnsServers()
{
var servers = new List<IPAddress>();
foreach (var iface in NetworkInterface.GetAllNetworkInterfaces())
{
if (iface.OperationalStatus != OperationalStatus.Up)
continue;
foreach (var addr in iface.GetIPProperties().DnsAddresses)
if (!servers.Contains(addr))
servers.Add(addr);
}
return servers.Count > 0 ? [.. servers] : [IPAddress.Loopback];
}

// Encode a DNS query packet for qname/SRV/IN.
static byte[] BuildDnsQuery(string qname)
{
var buf = new List<byte>(64);
var id = (ushort)Random.Shared.Next(65536);
buf.Add((byte)(id >> 8)); buf.Add((byte)id);
buf.Add(0x01); buf.Add(0x00); // flags: QR=0 RD=1
buf.Add(0x00); buf.Add(0x01); // QDCOUNT=1
buf.Add(0x00); buf.Add(0x00); // ANCOUNT=0
buf.Add(0x00); buf.Add(0x00); // NSCOUNT=0
buf.Add(0x00); buf.Add(0x00); // ARCOUNT=0
foreach (var label in qname.Split('.'))
{
buf.Add((byte)label.Length);
foreach (var c in label) buf.Add((byte)c);
}
buf.Add(0x00); // root label
buf.Add(0x00); buf.Add(TypeSRV); // QTYPE=SRV
buf.Add(0x00); buf.Add(ClassIN); // QCLASS=IN
return [.. buf];
}

// Parse a DNS response and extract SRV records from the answer section.
static List<SrvRecord> ParseDnsResponse(string qname, byte[] buf)
{
if (buf.Length < 12)
throw new NpgsqlException($"DNS response too short for {qname}");

var rcode = buf[3] & 0x0F;
if (rcode != 0)
throw new NpgsqlException($"DNS SRV lookup failed for {qname}: RCODE={rcode}");

var qdcount = (buf[4] << 8) | buf[5];
var ancount = (buf[6] << 8) | buf[7];
var pos = 12;

for (var i = 0; i < qdcount; i++)
{
pos = SkipName(buf, pos);
pos += 4; // QTYPE + QCLASS
}

var records = new List<SrvRecord>(ancount);
for (var i = 0; i < ancount; i++)
{
pos = SkipName(buf, pos);
if (pos + 10 > buf.Length) break;

var rtype = (buf[pos] << 8) | buf[pos + 1];
var rdlen = (buf[pos + 8] << 8) | buf[pos + 9];
pos += 10;
if (pos + rdlen > buf.Length) break;

if (rtype == TypeSRV && rdlen >= 7)
{
var priority = (ushort)((buf[pos] << 8) | buf[pos + 1]);
var weight = (ushort)((buf[pos + 2] << 8) | buf[pos + 3]);
var port = (ushort)((buf[pos + 4] << 8) | buf[pos + 5]);
var target = ReadName(buf, pos + 6).TrimEnd('.');
records.Add(new SrvRecord(priority, weight, port, target));
}
pos += rdlen;
}
return records;
}

// Skip a (possibly compressed) DNS name; returns position after it.
static int SkipName(byte[] buf, int pos)
{
while (pos < buf.Length)
{
var len = buf[pos];
if (len == 0) return pos + 1;
if ((len & 0xC0) == 0xC0) return pos + 2;
pos += len + 1;
}
return pos;
}

// Decode a (possibly compressed) DNS name starting at pos.
static string ReadName(byte[] buf, int pos)
{
var sb = new StringBuilder();
var jumped = false;
while (pos < buf.Length)
{
var len = buf[pos];
if (len == 0) break;
if ((len & 0xC0) == 0xC0)
{
pos = ((len & 0x3F) << 8) | buf[pos + 1];
jumped = true;
continue;
}
if (sb.Length > 0) sb.Append('.');
sb.Append(Encoding.ASCII.GetString(buf, pos + 1, len));
pos += len + 1;
if (jumped) continue;
}
return sb.ToString();
}
}
10 changes: 10 additions & 0 deletions test/Npgsql.SrvTests/Npgsql.SrvTests.csproj
Original file line number Diff line number Diff line change
@@ -0,0 +1,10 @@
<Project Sdk="Microsoft.NET.Sdk">
<ItemGroup>
<PackageReference Include="NUnit" />
<PackageReference Include="Microsoft.NET.Test.Sdk" />
<PackageReference Include="NUnit3TestAdapter" />
</ItemGroup>
<ItemGroup>
<ProjectReference Include="../../src/Npgsql/Npgsql.csproj" />
</ItemGroup>
</Project>
Loading