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 @@ -29,7 +29,7 @@ internal static partial class OpenSsl
// Special value of 0 means unlimited, -1 means the implementation (OpenSSL) default, which is currently 20 * 1024.
private const string TlsCacheSizeCtxName = "System.Net.Security.TlsCacheSize";
private const string TlsCacheSizeEnvironmentVariable = "DOTNET_SYSTEM_NET_SECURITY_TLSCACHESIZE";
private const int DefaultTlsCacheSizeClient = 500; // since we keep only one TLS Session per hostname, 500 should be enough to cover most scenarios
private const int DefaultTlsCacheSizeClient = 500; // bounds the total number of pooled client sessions across all hostnames
private const int DefaultTlsCacheSizeServer = -1; // use implementation default
private const SslProtocols FakeAlpnSslProtocol = (SslProtocols)1; // used to distinguish server sessions with ALPN
private static readonly Lazy<string[]> s_defaultSigAlgs = new(GetDefaultSignatureAlgorithms);
Expand Down Expand Up @@ -1299,7 +1299,11 @@ private static unsafe int NewSessionCallback(IntPtr ssl, IntPtr session)

if (ctxHandle != null)
{
if (ctxHandle.TryAddSession(name, session))
// TLS 1.3 tickets are single-use, TLS 1.2 sessions are not, and the two
// need opposite caching policies.
ReadOnlySpan<byte> version = MemoryMarshal.CreateReadOnlySpanFromNullTerminated(Ssl.SslGetVersion(ssl));

if (ctxHandle.TryAddSession(name, session, version.SequenceEqual("TLSv1.3"u8)))
{
// offered session was stored in our cache.
return 1;
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -50,6 +50,9 @@ internal static partial class Ssl
[LibraryImport(Libraries.CryptoNative, EntryPoint = "CryptoNative_SslGetVersion")]
internal static partial IntPtr SslGetVersion(SafeSslHandle ssl);

[LibraryImport(Libraries.CryptoNative, EntryPoint = "CryptoNative_SslGetVersion")]
internal static unsafe partial byte* SslGetVersion(IntPtr ssl);

[LibraryImport(Libraries.CryptoNative, EntryPoint = "CryptoNative_SslSetTlsExtHostName", StringMarshalling = StringMarshalling.Utf8)]
[return: MarshalAs(UnmanagedType.Bool)]
internal static partial bool SslSetTlsExtHostName(SafeSslHandle ssl, string host);
Expand Down Expand Up @@ -304,6 +307,9 @@ internal static SafeSharedX509StackHandle SslGetPeerCertChain(SafeSslHandle ssl)
[LibraryImport(Libraries.CryptoNative, EntryPoint = "CryptoNative_SslSessionFree")]
internal static partial void SessionFree(IntPtr session);

[LibraryImport(Libraries.CryptoNative, EntryPoint = "CryptoNative_SslSessionUpRef")]
internal static partial int SessionUpRef(IntPtr session);

[LibraryImport(Libraries.CryptoNative, EntryPoint = "CryptoNative_SslSessionSetHostname")]
internal static unsafe partial int SessionSetHostname(IntPtr session, byte* name);

Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -79,8 +79,21 @@ namespace Microsoft.Win32.SafeHandles
{
internal sealed class SafeSslContextHandle : SafeHandle, ISafeHandleCachable
{
// OpenSSL retires a TLS 1.3 session when the handshake using it finishes
// (tls_finish_handshake calls SSL_CTX_remove_session, which sets not_resumable on
// the shared object), so offering one session to several concurrent handshakes
// silently downgrades all but the first to a full handshake. Pooling several
// tickets per host lets concurrent connections each take a distinct one.
private const int TlsResumePoolSize = 8;

private readonly struct CachedSession(IntPtr session, bool isTls13)
{
public IntPtr Session { get; } = session;
public bool IsTls13 { get; } = isTls13;
}

// This is session cache keyed by SNI e.g. TargetHost
private Dictionary<string, IntPtr>? _sslSessions;
private Dictionary<string, List<CachedSession>>? _sslSessions;
private GCHandle _gch;

// SSL_CTX handles are cached, so we need to keep track of the
Expand Down Expand Up @@ -146,9 +159,12 @@ protected override bool ReleaseHandle()

lock (_sslSessions)
{
foreach (IntPtr session in _sslSessions.Values)
foreach (List<CachedSession> sessions in _sslSessions.Values)
{
Interop.Ssl.SessionFree(session);
foreach (CachedSession cached in sessions)
{
Interop.Ssl.SessionFree(cached.Session);
}
}

_sslSessions.Clear();
Expand All @@ -168,14 +184,14 @@ internal void EnableSessionCache()
{
Debug.Assert(_sslSessions == null);

_sslSessions = new Dictionary<string, IntPtr>();
_sslSessions = new Dictionary<string, List<CachedSession>>();
_gch = GCHandle.Alloc(this);
Debug.Assert(_gch.IsAllocated);
// This is needed so we can find the handle from session in SessionRemove callback.
Interop.Ssl.SslCtxSetData(this, (IntPtr)_gch);
}

internal unsafe bool TryAddSession(byte* namePtr, IntPtr session)
internal unsafe bool TryAddSession(byte* namePtr, IntPtr session, bool isTls13)
Comment thread
wfurt marked this conversation as resolved.
{
Debug.Assert(_sslSessions != null && session != IntPtr.Zero);

Expand All @@ -187,70 +203,109 @@ internal unsafe bool TryAddSession(byte* namePtr, IntPtr session)
string? targetName = Utf8StringMarshaller.ConvertToManaged(namePtr);
Debug.Assert(targetName != null);

if (!string.IsNullOrEmpty(targetName))
if (string.IsNullOrEmpty(targetName))
{
// We do this only for lookup in RemoveSession.
// Since this is part of cache manipulation and no function impact it is done here.
// This will use strdup() so it is safe to pass in raw pointer.
Interop.Ssl.SessionSetHostname(session, namePtr);
return false;
}

IntPtr oldSession = IntPtr.Zero;
// We do this only for lookup in RemoveSession.
// Since this is part of cache manipulation and no function impact it is done here.
// This will use strdup() so it is safe to pass in raw pointer.
Interop.Ssl.SessionSetHostname(session, namePtr);

lock (_sslSessions)
// A TLS 1.2 session stays usable after a resumption and is never replaced by a
// new one (OpenSSL skips new_session_cb on resumed TLS 1.2 handshakes), so a
// single entry is both sufficient and all we will ever be given.
int limit = isTls13 ? TlsResumePoolSize : 1;

IntPtr[]? evicted = null;
int evictedCount = 0;

lock (_sslSessions)
{
ref List<CachedSession>? sessions = ref CollectionsMarshal.GetValueRefOrAddDefault(_sslSessions, targetName, out _);
sessions ??= new List<CachedSession>();

// Pooled tickets are only usable by the protocol version that produced them,
// so a change of negotiated version drops the pool rather than leaving a
// stale entry at the head masking everything behind it.
int toEvict = sessions.Count > 0 && sessions[0].IsTls13 != isTls13
? sessions.Count
: Math.Max(0, sessions.Count - limit + 1);

if (toEvict > 0)
{
if (!_sslSessions.TryAdd(targetName, session))
evicted = new IntPtr[toEvict];
for (; evictedCount < toEvict; evictedCount++)
{
// session to this target host exists, replace it
_sslSessions.Remove(targetName, out oldSession);
bool added = _sslSessions.TryAdd(targetName, session);
Debug.Assert(added);
evicted[evictedCount] = sessions[evictedCount].Session;
}
}

if (oldSession != IntPtr.Zero)
{
// remove old session also from the internal OpenSSL cache
// and drop reference count. Since SSL_CTX_remove_session
// will call session_remove_cb, we need to do this outside
// of _sslSessions lock to avoid deadlock with another thread
// which could be holding SSL_CTX lock and trying to acquire
// _sslSessions lock.
Interop.Ssl.SslCtxRemoveSession(this, oldSession);
Interop.Ssl.SessionFree(oldSession);
sessions.RemoveRange(0, toEvict);
}

return true;
sessions.Add(new CachedSession(session, isTls13));
}

for (int i = 0; i < evictedCount; i++)
{
// Remove the evicted session also from the internal OpenSSL cache and drop
// the reference count. Since SSL_CTX_remove_session will call
// session_remove_cb, we need to do this outside of the _sslSessions lock to
// avoid deadlock with another thread which could be holding the SSL_CTX lock
// and trying to acquire _sslSessions.
Interop.Ssl.SslCtxRemoveSession(this, evicted![i]);
Interop.Ssl.SessionFree(evicted[i]);
}

return false;
return true;
}

internal unsafe void RemoveSession(byte* namePtr, IntPtr session)
{
Debug.Assert(_sslSessions != null);

if (_sslSessions == null || namePtr == null)
{
return;
}

string? targetName = Utf8StringMarshaller.ConvertToManaged(namePtr);
Debug.Assert(targetName != null);

if (_sslSessions != null && targetName != null)
if (targetName == null)
{
IntPtr oldSession = IntPtr.Zero;
bool removed = false;
lock (_sslSessions)
return;
}

bool removed = false;

lock (_sslSessions)
{
if (_sslSessions.TryGetValue(targetName, out List<CachedSession>? sessions))
{
if (_sslSessions.TryGetValue(targetName, out IntPtr existingSession) && existingSession == session)
for (int i = 0; i < sessions.Count; i++)
{
removed = _sslSessions.Remove(targetName, out oldSession);
if (sessions[i].Session == session)
{
sessions.RemoveAt(i);
removed = true;
break;
}
}
}

if (removed)
{
// It seems like we may be called more than once. Since we grabbed only one refference
// when added to Dictionary, we will also drop exactly one when removed.
Interop.Ssl.SessionFree(oldSession);
if (sessions.Count == 0)
{
_sslSessions.Remove(targetName);
}
}
}

if (removed)
{
// It seems like we may be called more than once. Since we grabbed only one
// reference when added to the cache, we will also drop exactly one when removed.
Interop.Ssl.SessionFree(session);
}
}

Expand All @@ -263,18 +318,44 @@ internal bool TrySetSession(SafeSslHandle sslHandle, string name)
return false;
}

IntPtr owned;

lock (_sslSessions)
{
if (_sslSessions.TryGetValue(name, out IntPtr session))
if (!_sslSessions.TryGetValue(name, out List<CachedSession>? sessions) || sessions.Count == 0)
{
// This will increase reference count on the session as needed.
// We need to hold lock here to prevent session being deleted before the call is done.
Interop.Ssl.SslSetSession(sslHandle, session);
return true;
return false;
}

CachedSession cached = sessions[0];

// While the pool holds more than one ticket each concurrent handshake can
// take its own. The last one is still shared rather than withheld, since a
// shared ticket only costs a fallback to a full handshake, while withholding
// it guarantees one.
bool singleUse = cached.IsTls13 && sessions.Count > 1;

if (singleUse)
{
// Taking the entry out of the cache transfers the cache's reference to us.
// The pool holds more than one entry here, so it cannot become empty.
sessions.RemoveAt(0);
}
else if (Interop.Ssl.SessionUpRef(cached.Session) != 1)
{
return false;
}

owned = cached.Session;
}

return false;
// Held outside the lock: RemoveSession frees on a callback OpenSSL raises while
// holding the SSL_CTX lock, so the reference taken above, not the lock, is what
// keeps the session alive across this call.
bool set = Interop.Ssl.SslSetSession(sslHandle, owned) == 1;
Interop.Ssl.SessionFree(owned);

return set;
}
}
}
Original file line number Diff line number Diff line change
Expand Up @@ -408,6 +408,7 @@ static const Entry s_cryptoNative[] =
DllImportEntry(CryptoNative_SslSessionFree)
DllImportEntry(CryptoNative_SslSessionGetHostname)
DllImportEntry(CryptoNative_SslSessionSetHostname)
DllImportEntry(CryptoNative_SslSessionUpRef)
DllImportEntry(CryptoNative_SslSessionReused)
DllImportEntry(CryptoNative_SslSessionGetData)
DllImportEntry(CryptoNative_SslSessionSetData)
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -771,6 +771,7 @@ extern bool g_libSslUses32BitTime;
REQUIRED_FUNCTION(SSL_SESSION_set_ex_data) \
REQUIRED_FUNCTION(SSL_SESSION_get0_hostname) \
REQUIRED_FUNCTION(SSL_SESSION_set1_hostname) \
REQUIRED_FUNCTION(SSL_SESSION_up_ref) \
REQUIRED_FUNCTION(SSL_session_reused) \
REQUIRED_FUNCTION(SSL_set_accept_state) \
REQUIRED_FUNCTION(SSL_set_bio) \
Expand Down Expand Up @@ -1379,6 +1380,7 @@ extern TYPEOF(OPENSSL_gmtime)* OPENSSL_gmtime_ptr;
#define SSL_SESSION_free SSL_SESSION_free_ptr
#define SSL_SESSION_get0_hostname SSL_SESSION_get0_hostname_ptr
#define SSL_SESSION_set1_hostname SSL_SESSION_set1_hostname_ptr
#define SSL_SESSION_up_ref SSL_SESSION_up_ref_ptr
#define SSL_session_reused SSL_session_reused_ptr
#define SSL_SESSION_get_ex_data SSL_SESSION_get_ex_data_ptr
#define SSL_SESSION_set_ex_data SSL_SESSION_set_ex_data_ptr
Expand Down
5 changes: 5 additions & 0 deletions src/native/libs/System.Security.Cryptography.Native/pal_ssl.c
Original file line number Diff line number Diff line change
Expand Up @@ -874,6 +874,11 @@ void CryptoNative_SslSessionFree(SSL_SESSION* session)
SSL_SESSION_free(session);
}

int32_t CryptoNative_SslSessionUpRef(SSL_SESSION* session)
{
return SSL_SESSION_up_ref(session);
}

const char* CryptoNative_SslSessionGetHostname(SSL_SESSION* session)
{
return SSL_SESSION_get0_hostname(session);
Expand Down
5 changes: 5 additions & 0 deletions src/native/libs/System.Security.Cryptography.Native/pal_ssl.h
Original file line number Diff line number Diff line change
Expand Up @@ -207,6 +207,11 @@ PALEXPORT int32_t CryptoNative_SslSetSession(SSL* ssl, SSL_SESSION* session);
*/
PALEXPORT void CryptoNative_SslSessionFree(SSL_SESSION* session);

/*
* Takes an additional reference on an SSL session. Returns 1 on success, 0 on failure.
*/
PALEXPORT int32_t CryptoNative_SslSessionUpRef(SSL_SESSION* session);

/*
* Get name associated with given SSL_SESSION.
*/
Expand Down