Skip to content
Open
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
3 changes: 3 additions & 0 deletions .changes/unreleased/Fixed-20260920-205039.yaml
Original file line number Diff line number Diff line change
@@ -0,0 +1,3 @@
kind: Fixed
body: Reuse TLS contexts for Core and third-party provider requests to avoid repeatedly loading trusted certificates.
time: 2026-09-20T20:50:39.0338+01:00
201 changes: 201 additions & 0 deletions scripts/benchmark_tls.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,201 @@
# Copyright (c) 2021, VRAI Labs and/or its affiliates. All rights reserved.
#
# This software is licensed under the Apache License, Version 2.0 (the
# "License") as published by the Apache Software Foundation.
#
# You may not use this file except in compliance with the License. You may
# obtain a copy of the License at http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS, WITHOUT
# WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the
# License for the specific language governing permissions and limitations
# under the License.
"""Compare TLS/client reuse: python scripts/benchmark_tls.py after make dev-install.

Uses local HTTP/HTTPS servers and temporary certificates. Each variant performs
one cold request followed by 60 warm requests, with real HTTPX transports.
"""

import asyncio
import json
import os
import platform
import ssl
import statistics
import tempfile
import time
from datetime import datetime, timedelta, timezone
from ipaddress import ip_address
from pathlib import Path
from typing import List, Optional, Union
from unittest.mock import patch

import certifi
import httpx
from cryptography import x509
from cryptography.hazmat.primitives import hashes, serialization
from cryptography.hazmat.primitives.asymmetric import rsa
from cryptography.x509.oid import NameOID
from supertokens_python.ssl_utils import get_ssl_context, reset_ssl_context


async def main() -> None:
print(
json.dumps(
{
"python": platform.python_version(),
"httpx": httpx.__version__,
"openssl": ssl.OPENSSL_VERSION,
"os": platform.platform(),
"requests_per_variant": 61,
"concurrency": 1,
"trust": "certifi bundle plus ephemeral local CA",
}
)
)
with tempfile.TemporaryDirectory() as directory:
key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
name = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, "localhost")])
now = datetime.now(timezone.utc)
cert = (
x509.CertificateBuilder()
.subject_name(name)
.issuer_name(name)
.public_key(key.public_key())
.serial_number(x509.random_serial_number())
.not_valid_before(now - timedelta(days=1))
.not_valid_after(now + timedelta(days=1))
.add_extension(
x509.SubjectAlternativeName(
[x509.DNSName("localhost"), x509.IPAddress(ip_address("127.0.0.1"))]
),
False,
)
.add_extension(x509.BasicConstraints(ca=True, path_length=None), True)
.sign(key, hashes.SHA256())
)
cert_path = Path(directory) / "server.pem"
cert_path.write_bytes(cert.public_bytes(serialization.Encoding.PEM))
key_path = Path(directory) / "server.key"
key_path.write_bytes(
key.private_bytes(
serialization.Encoding.PEM,
serialization.PrivateFormat.PKCS8,
serialization.NoEncryption(),
)
)
bundle = Path(directory) / "ca.pem"
bundle.write_bytes(Path(certifi.where()).read_bytes() + cert_path.read_bytes())
server_context = ssl.SSLContext(ssl.PROTOCOL_TLS_SERVER)
server_context.load_cert_chain(cert_path, key_path)
connections = 0

async def serve(
reader: asyncio.StreamReader, writer: asyncio.StreamWriter
) -> None:
nonlocal connections
connections += 1
try:
while True:
await reader.readuntil(b"\r\n\r\n")
writer.write(
b"HTTP/1.1 200 OK\r\nContent-Length: 2\r\nContent-Type: application/json\r\n\r\n{}"
)
await writer.drain()
except (asyncio.IncompleteReadError, ConnectionError):
pass
finally:
writer.close()
await writer.wait_closed()

with patch.dict(
os.environ,
{"SSL_CERT_FILE": str(bundle), "NO_PROXY": "localhost,127.0.0.1"},
):
for scheme in ("http", "https"):
server = await asyncio.start_server(
serve,
"127.0.0.1",
0,
ssl=server_context if scheme == "https" else None,
)
url = f"{scheme}://127.0.0.1:{server.sockets[0].getsockname()[1]}/"
async with server:
for variant in ("fresh", "cached-context", "shared-client"):
reset_ssl_context()
connections = 0
contexts = 0
context_ms = 0.0
original = ssl.create_default_context

def counted(
purpose: ssl.Purpose = ssl.Purpose.SERVER_AUTH,
*,
cafile: Optional[str] = None,
capath: Optional[str] = None,
cadata: Optional[Union[str, bytes]] = None,
) -> ssl.SSLContext:
nonlocal contexts, context_ms
started = time.perf_counter()
result = original(
purpose, cafile=cafile, capath=capath, cadata=cadata
)
contexts += 1
context_ms += (time.perf_counter() - started) * 1000
return result

elapsed: List[float] = []
initialization: List[float] = []
context: Optional[ssl.SSLContext] = None
shared: Optional[httpx.AsyncClient] = None
with patch("ssl.create_default_context", counted):
for _ in range(61):
started = time.perf_counter()
if variant == "cached-context":
context = get_ssl_context()
client = shared or httpx.AsyncClient(
timeout=30.0,
verify=context if context is not None else True,
)
initialization.append(
(time.perf_counter() - started) * 1000
)
if variant == "shared-client":
shared = client
try:
response = await client.get(url)
assert response.json() == {}
finally:
if shared is None:
await client.aclose()
elapsed.append((time.perf_counter() - started) * 1000)
if shared is not None:
await shared.aclose()
warm = sorted(elapsed[1:])
print(
json.dumps(
{
"scheme": scheme,
"variant": variant,
"cold_ms": round(elapsed[0], 3),
"warm_median_ms": round(statistics.median(warm), 3),
"warm_p95_ms": round(
warm[int(len(warm) * 0.95) - 1], 3
),
"warm_init_median_ms": round(
statistics.median(initialization[1:]), 3
),
"contexts": contexts,
"context_total_ms": round(context_ms, 3),
"connections": connections,
"tls_handshakes": connections
if scheme == "https"
else 0,
}
)
)


if __name__ == "__main__":
asyncio.run(main())
4 changes: 3 additions & 1 deletion supertokens_python/querier.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,8 @@

from httpx import AsyncClient, ConnectTimeout, NetworkError, Response

from supertokens_python.ssl_utils import get_ssl_context

from .constants import (
API_KEY_HEADER,
API_VERSION,
Expand Down Expand Up @@ -107,7 +109,7 @@ async def api_request(
raise Exception("Retry request failed")

try:
async with AsyncClient(timeout=30.0) as client:
async with AsyncClient(timeout=30.0, verify=get_ssl_context()) as client:
if method == "GET":
return await client.get(url, *args, **kwargs) # type: ignore
if method == "POST":
Expand Down
3 changes: 2 additions & 1 deletion supertokens_python/recipe/thirdparty/providers/custom.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
get_actual_client_id_from_development_client_id,
is_using_oauth_development_client_id,
)
from supertokens_python.ssl_utils import get_ssl_context

from ..provider import (
AuthorisationRedirect,
Expand Down Expand Up @@ -155,7 +156,7 @@ async def verify_id_token_from_jwks_endpoint_and_get_payload(
id_token: str, jwks_uri: str, audience: str
):
public_keys: List[RSAAlgorithm] = []
async with AsyncClient(timeout=30.0) as client:
async with AsyncClient(timeout=30.0, verify=get_ssl_context()) as client:
response = await client.get(jwks_uri) # type:ignore
key_payload = response.json()
for key in key_payload["keys"]:
Expand Down
5 changes: 3 additions & 2 deletions supertokens_python/recipe/thirdparty/providers/utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,6 +5,7 @@
from supertokens_python.logger import log_debug_message
from supertokens_python.normalised_url_domain import NormalisedURLDomain
from supertokens_python.normalised_url_path import NormalisedURLPath
from supertokens_python.ssl_utils import get_ssl_context

DEV_OAUTH_CLIENT_IDS = [
"1060725074195-kmeum4crr01uirfl2op9kd5acmi9jutn.apps.googleusercontent.com",
Expand Down Expand Up @@ -36,7 +37,7 @@ async def do_get_request(
if headers is None:
headers = {}

async with AsyncClient(timeout=30.0) as client:
async with AsyncClient(timeout=30.0, verify=get_ssl_context()) as client:
res = await client.get(url, params=query_params, headers=headers) # type:ignore

log_debug_message(
Expand All @@ -59,7 +60,7 @@ async def do_post_request(
headers["content-type"] = "application/x-www-form-urlencoded"
headers["accept"] = "application/json"

async with AsyncClient(timeout=30.0) as client:
async with AsyncClient(timeout=30.0, verify=get_ssl_context()) as client:
res = await client.post(url, data=body_params, headers=headers) # type:ignore
log_debug_message(
"Received response with status %s and body %s", res.status_code, res.text
Expand Down
52 changes: 52 additions & 0 deletions supertokens_python/ssl_utils.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,52 @@
# Copyright (c) 2021, VRAI Labs and/or its affiliates. All rights reserved.
#
# This software is licensed under the Apache License, Version 2.0 (the
# "License") as published by the Apache Software Foundation.
#
# You may not use this file except in compliance with the License. You may
# obtain a copy of the License at http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS, WITHOUT
# WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. See the
# License for the specific language governing permissions and limitations
# under the License.
"""Internal TLS configuration shared by short-lived HTTP clients."""

from functools import lru_cache
from os import environ
from ssl import SSLContext
from typing import Optional

import httpx


@lru_cache(maxsize=1)
def _create_ssl_context(
ssl_cert_file: Optional[str],
ssl_cert_dir: Optional[str],
ssl_key_log_file: Optional[str],
) -> SSLContext:
# Arguments form the cache key; HTTPX reads the environment itself so its
# certificate selection and TLS defaults remain authoritative.
return httpx.create_ssl_context(verify=True, trust_env=True)


def get_ssl_context() -> SSLContext:
"""Reuse TLS configuration without sharing clients across event loops.

Environment changes invalidate the single-entry cache. Changes to certificate
files at the same path require a process restart or an explicit cache reset.
Concurrent cold calls may each build a context; subsequent calls reuse the
cached result. Clients must not modify the context's verification settings.
"""
return _create_ssl_context(
environ.get("SSL_CERT_FILE"),
environ.get("SSL_CERT_DIR"),
environ.get("SSLKEYLOGFILE"),
)


def reset_ssl_context() -> None:
"""Clear cached TLS configuration when resetting SDK state in tests."""
_create_ssl_context.cache_clear()
Loading