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
Original file line number Diff line number Diff line change
Expand Up @@ -486,8 +486,14 @@ protected virtual string GetDefaultClientName() =>
public virtual bool GetDefaultSsl(EndPointCollection endPoints) => false;

/// <summary>
/// Gets the SSL Host to check for when connecting to endpoints (customizable in case of internal certificate shenanigans.
/// Gets the SSL host to infer when <see cref="ConfigurationOptions.SslHost"/> is not explicitly set, for
/// endpoints that don't already carry their own host name (customizable in case of internal certificate shenanigans).
/// </summary>
/// <remarks>
/// Only consulted for non-<see cref="DnsEndPoint"/> connections - a <see cref="DnsEndPoint"/> always uses its
/// own <see cref="DnsEndPoint.Host"/> instead, so in practice this applies to IP endpoints, such as ones
/// discovered via cluster topology.
/// </remarks>
/// <param name="endPoints">The configured endpoints to determine SSL host from (e.g. from the port).</param>
/// <returns>The common host, if any, detected from the endpoint collection.</returns>
public virtual string? GetSslHostFromEndpoints(EndPointCollection endPoints)
Expand Down
6 changes: 1 addition & 5 deletions src/StackExchange.Redis/Configuration/LoggingTunnel.cs
Original file line number Diff line number Diff line change
Expand Up @@ -407,11 +407,7 @@ private async Task<Stream> TlsHandshakeAsync(Stream stream, EndPoint endpoint)
#pragma warning restore CS1998 // Async method lacks 'await' operators and will run synchronously
{
// mirrors TLS handshake from PhysicalConnection, but wouldn't help to share code here
var host = _options.SslHost;
if (host.IsNullOrWhiteSpace())
{
host = Format.ToStringHostOnly(endpoint);
}
var host = _options.ResolveTlsHostName(endpoint);

var ssl = new SslStream(
innerStream: stream,
Expand Down
10 changes: 4 additions & 6 deletions src/StackExchange.Redis/Configuration/TlsOptions.cs
Original file line number Diff line number Diff line change
Expand Up @@ -96,12 +96,10 @@ public LocalCertificateSelectionCallback? CertificateSelectionCallback
#endif

/// <summary>
/// The TLS host name to use for the given endpoint: the configured <see cref="SslHost"/> if there is
/// one, otherwise the host portion of the endpoint - which is what the library's own TLS path does.
/// The TLS host name to use for the given endpoint. An explicitly configured <see cref="SslHost"/> always
/// wins. Otherwise, a <see cref="DnsEndPoint"/> uses its own host, and an address endpoint uses the inferred
/// configuration host when available before falling back to its address.
/// </summary>
public string ResolveHost(EndPoint endpoint)
{
var host = SslHost;
return host.IsNullOrWhiteSpace() ? Format.ToStringHostOnly(endpoint) : host!;
}
=> _options?.ResolveTlsHostName(endpoint) ?? Format.ToStringHostOnly(endpoint);
}
19 changes: 17 additions & 2 deletions src/StackExchange.Redis/ConfigurationOptions.cs
Original file line number Diff line number Diff line change
Expand Up @@ -957,14 +957,29 @@ public bool Ssl
}

/// <summary>
/// The target-host to use when validating SSL certificate; setting a value here enables SSL mode.
/// The target host to use for SNI and certificate validation; setting a value here enables SSL mode.
/// </summary>
/// <remarks>
/// When explicitly configured, this overrides the host for every endpoint. When unset, connections to a
/// <see cref="DnsEndPoint"/> use that endpoint's host, allowing hostname-routed clusters to use a distinct
/// SNI name for each node.
/// </remarks>
public string? SslHost
{
get => sslHost ?? Defaults.GetSslHostFromEndpoints(EndPoints);
get => sslHost;
set => sslHost = value;
}

internal string ResolveTlsHostName(EndPoint endpoint)
{
var host = sslHost;
if (host.IsNullOrWhiteSpace() && endpoint is not DnsEndPoint)
{
host = Defaults.GetSslHostFromEndpoints(EndPoints);
}
return Format.GetTlsHostName(endpoint, host);
}

/// <summary>
/// Configures which SSL/TLS protocols should be allowed. If not set, defaults are chosen by the .NET framework.
/// </summary>
Expand Down
22 changes: 22 additions & 0 deletions src/StackExchange.Redis/Format.cs
Original file line number Diff line number Diff line change
Expand Up @@ -131,6 +131,28 @@ internal static string ToStringHostOnly(EndPoint endpoint) =>
_ => "",
};

/// <summary>
/// Gets the TLS SNI and certificate-validation host for a connection to <paramref name="endpoint"/>.
/// </summary>
/// <remarks>
/// An explicitly configured <paramref name="sslHost"/> wins. Otherwise, a <see cref="DnsEndPoint"/>
/// uses its own host and any other endpoint uses its address.
/// </remarks>
internal static string GetTlsHostName(EndPoint endpoint, string? sslHost)
{
if (!sslHost.IsNullOrWhiteSpace())
{
return sslHost;
}

if (endpoint is DnsEndPoint dns && !dns.Host.IsNullOrWhiteSpace())
{
return dns.Host;
}

return ToStringHostOnly(endpoint);
}

internal static bool TryGetHostPort(EndPoint? endpoint, [NotNullWhen(true)] out string? host, [NotNullWhen(true)] out int? port)
{
if (endpoint is not null)
Expand Down
6 changes: 1 addition & 5 deletions src/StackExchange.Redis/PhysicalConnection.cs
Original file line number Diff line number Diff line change
Expand Up @@ -1226,11 +1226,7 @@ static Stream DemandSocketStream(Socket? socket)
if (config.Ssl)
{
log?.LogInformationConfiguringTLS();
var host = config.SslHost;
if (host.IsNullOrWhiteSpace())
{
host = Format.ToStringHostOnly(bridge.ServerEndPoint.EndPoint);
}
var host = config.ResolveTlsHostName(bridge.ServerEndPoint.EndPoint);

stream ??= DemandSocketStream(socket);
var ssl = new SslStream(
Expand Down
6 changes: 4 additions & 2 deletions tests/StackExchange.Redis.Tests/SSLTests.cs
Original file line number Diff line number Diff line change
Expand Up @@ -364,7 +364,7 @@ public async Task RedisLabsEnvironmentVariableClientCertificate(bool setEnv)
}

[Fact]
public void SSLHostInferredFromEndpoints()
public void SSLHostIsOnlyExplicitlyConfigured()
{
var options = new ConfigurationOptions
{
Expand All @@ -376,7 +376,9 @@ public void SSLHostInferredFromEndpoints()
},
Ssl = true,
};
Assert.Equal("mycache.rediscache.windows.net", options.SslHost);
Assert.Null(options.SslHost);
options.SslHost = "override.rediscache.windows.net";
Assert.Equal("override.rediscache.windows.net", options.SslHost);
options = new ConfigurationOptions()
{
EndPoints = { { "121.23.23.45", 15000 } },
Expand Down
88 changes: 88 additions & 0 deletions tests/StackExchange.Redis.Tests/TlsHostNameUnitTests.cs
Original file line number Diff line number Diff line change
@@ -0,0 +1,88 @@
using System.Net;
using Xunit;

namespace StackExchange.Redis.Tests;

/// <summary>
/// TLS target-host selection for SNI and certificate validation.
/// </summary>
public class TlsHostNameUnitTests
{
[Fact]
public void ExplicitSslHostOverridesDnsEndPoint()
{
var endpoint = new DnsEndPoint("host-2.redis.example.com", 443);
var options = new ConfigurationOptions { SslHost = "host-1.redis.example.com" };
Assert.Equal("host-1.redis.example.com", options.ResolveTlsHostName(endpoint));
}

[Fact]
public void DnsEndPointUsesItsOwnHostInsteadOfInferredHost()
{
var options = new ConfigurationOptions
{
EndPoints = { "host-1.redis.example.com:443" },
};

Assert.Null(options.SslHost);
Assert.Equal("host-2.redis.example.com", options.ResolveTlsHostName(new DnsEndPoint("host-2.redis.example.com", 443)));
}

[Fact]
public void IpEndPointPrefersExplicitSslHost()
{
var endpoint = new IPEndPoint(IPAddress.Parse("10.0.0.1"), 6379);
var options = new ConfigurationOptions
{
EndPoints = { "host-1.redis.example.com:443" },
SslHost = "mycache.redis.example.com",
};
Assert.Equal("mycache.redis.example.com", options.ResolveTlsHostName(endpoint));
}

[Fact]
public void IpEndPointUsesInferredHost()
{
var options = new ConfigurationOptions
{
EndPoints = { "host-1.redis.example.com:443" },
};
Assert.Equal("host-1.redis.example.com", options.ResolveTlsHostName(new IPEndPoint(IPAddress.Parse("10.0.0.1"), 6379)));
}

[Fact]
public void IpEndPointFallsBackToAddressWhenHostsAreMissing()
{
var endpoint = new IPEndPoint(IPAddress.Parse("10.0.0.1"), 6379);
Assert.Equal("10.0.0.1", new ConfigurationOptions().ResolveTlsHostName(endpoint));
}

[Fact]
public void TlsOptionsResolveHostHonorsExplicitOverride()
{
var options = new ConfigurationOptions
{
EndPoints = { "host-1.redis.example.com:443" },
Ssl = true,
SslHost = "override.redis.example.com",
};
var tls = new Configuration.TlsOptions(options);

Assert.Equal("override.redis.example.com", tls.ResolveHost(new DnsEndPoint("host-2.redis.example.com", 443)));
Assert.Equal("override.redis.example.com", tls.ResolveHost(new IPEndPoint(IPAddress.Parse("10.0.0.1"), 443)));
}

[Fact]
public void TlsOptionsResolveHostUsesEndpointBeforeInferredHost()
{
var options = new ConfigurationOptions
{
EndPoints = { "host-1.redis.example.com:443" },
Ssl = true,
};
var tls = new Configuration.TlsOptions(options);

Assert.Equal("host-2.redis.example.com", tls.ResolveHost(new DnsEndPoint("host-2.redis.example.com", 443)));
Assert.Equal("host-1.redis.example.com", tls.ResolveHost(new IPEndPoint(IPAddress.Parse("10.0.0.1"), 443)));
}
}
Loading