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
32 changes: 7 additions & 25 deletions packages/django-cf/django_cf/db/base_engine.py
Original file line number Diff line number Diff line change
Expand Up @@ -242,34 +242,16 @@ def from_object(
except ImportError:
jsnull = None

result = []

for row in data:
row_items = ()
if isinstance(row, list):
for v in row:
if v is jsnull:
row_items += (None,)
else:
row_items += (v,)
else:
for v in row.values():
if v is jsnull:
row_items += (None,)
else:
row_items += (v,)

result.append(row_items)
def to_row(row):
values = row if isinstance(row, list) else row.values()
return tuple(None if v is jsnull else v for v in values)

instance = CFResult(result)
instance = CFResult([to_row(row) for row in data])

if rows_read or rows_written:
if "INSERT" in query.upper():
instance.set_rowcount(rows_written or 0)
elif "UPDATE" in query.upper() or "DELETE" in query.upper():
instance.set_rowcount(rows_written or 0)
else:
instance.set_rowcount(rows_read or 0)
upper_query = query.upper()
is_write = any(kw in upper_query for kw in ("INSERT", "UPDATE", "DELETE"))
instance.set_rowcount((rows_written if is_write else rows_read) or 0)

if last_row_id is not None:
instance.set_lastrowid(last_row_id)
Expand Down
175 changes: 91 additions & 84 deletions packages/django-cf/django_cf/middleware/CloudflareAccessMiddleware.py
Original file line number Diff line number Diff line change
Expand Up @@ -137,71 +137,74 @@ def _authenticate_cloudflare_access(self, request):
logger.error("Failed to retrieve Cloudflare public keys")
return None

email = None
name = None

# Validate and decode JWT
try:
# Try each public key until one works
decoded_token = None
for key_data in public_keys:
try:
decoded_token = self._decode_and_verify_jwt(jwt_token, key_data)
if decoded_token:
break
except Exception as e:
logger.debug(f"Key {key_data.get('kid')} failed: {str(e)}")
continue

decoded_token = self._decode_with_any_key(jwt_token, public_keys)
if not decoded_token:
logger.warning("JWT token validation failed with all available keys")
return None

# Validate AUD if configured
if self.aud:
token_aud = decoded_token.get("aud")
if isinstance(token_aud, list):
if self.aud not in token_aud:
logger.warning(
f"JWT audience mismatch. Expected: {self.aud}, Got: {token_aud}"
)
return None
elif token_aud != self.aud:
logger.warning(
f"JWT audience mismatch. Expected: {self.aud}, Got: {token_aud}"
)
return None

# Validate AUD if not configured but we have a team name
if not self.aud and self.team_name:
token_aud = decoded_token.get("aud")
if not token_aud:
logger.warning("No audience found in JWT token")
return None
if not self._validate_audience(decoded_token):
return None

# Extract user information from JWT claims
email = decoded_token.get("email")
name = decoded_token.get("name", "")

# Try to get name from custom claims if not in standard claims
if not name:
custom_claims = decoded_token.get("custom", {})
first_name = custom_claims.get("firstName", "")
last_name = custom_claims.get("lastName", "")
if first_name or last_name:
name = f"{first_name} {last_name}".strip()

if not email:
logger.warning("No email found in JWT token")
return None
name = self._extract_name(decoded_token)

except Exception as e:
logger.warning(f"JWT token validation error: {repr(e)}")
return None

# Get or create user
user = self._get_or_create_user(email, name)
return user
return self._get_or_create_user(email, name)

def _decode_with_any_key(self, jwt_token, public_keys):
"""Try each public key in turn; return the decoded payload or None."""
for key_data in public_keys:
try:
decoded_token = self._decode_and_verify_jwt(jwt_token, key_data)
except Exception as e:
logger.debug(f"Key {key_data.get('kid')} failed: {str(e)}")
continue
if decoded_token:
return decoded_token
return None

def _validate_audience(self, decoded_token):
"""Check the token's audience against the configured AUD, if any."""
token_aud = decoded_token.get("aud")

if self.aud:
matches = (
self.aud in token_aud
if isinstance(token_aud, list)
else token_aud == self.aud
)
if not matches:
logger.warning(
f"JWT audience mismatch. Expected: {self.aud}, Got: {token_aud}"
)
return matches

# AUD not configured but we have a team name: require some audience
if self.team_name and not token_aud:
logger.warning("No audience found in JWT token")
return False

return True

@staticmethod
def _extract_name(decoded_token):
"""Extract the display name from standard or custom JWT claims."""
name = decoded_token.get("name", "")
if name:
return name

custom_claims = decoded_token.get("custom", {})
first_name = custom_claims.get("firstName", "")
last_name = custom_claims.get("lastName", "")
return f"{first_name} {last_name}".strip()

def _extract_jwt_token(self, request):
"""Extract JWT token from CF-Access-Jwt-Assertion header or cf_authorization cookie."""
Expand Down Expand Up @@ -267,49 +270,53 @@ def _get_cloudflare_public_keys(self):
if cached_keys:
return cached_keys

data = self._fetch_certs()
if data is None:
return None

processed_keys = self._process_jwks(data.get("keys", []))

# Cache the keys
cache.set(cache_key, processed_keys, self.cache_timeout)
return processed_keys

def _fetch_certs(self):
"""Fetch the JWKS document from the certs URL. Returns dict or None."""
if IS_WORKER:
response = run_sync(fetch(self.certs_url))
if response.status == 200:
data = run_sync(response.json()).to_py()
else:
if response.status != 200:
logger.error(f"Failed to fetch Cloudflare keys: HTTP {response.status}")
return None
else:
try:
with urllib.request.urlopen(self.certs_url) as response:
if response.status == 200:
data = json.loads(response.read().decode("utf-8"))
else:
logger.error(
f"Failed to fetch Cloudflare keys: HTTP {response.status}"
)
return None
except urllib.error.URLError as e:
logger.error(f"Network error fetching Cloudflare keys: {str(e)}")
return None
except Exception as e:
logger.error(f"Unexpected error fetching Cloudflare keys: {str(e)}")
return None
return run_sync(response.json()).to_py()

keys = data.get("keys", [])
try:
with urllib.request.urlopen(self.certs_url) as response:
if response.status != 200:
logger.error(
f"Failed to fetch Cloudflare keys: HTTP {response.status}"
)
return None
return json.loads(response.read().decode("utf-8"))
except urllib.error.URLError as e:
logger.error(f"Network error fetching Cloudflare keys: {str(e)}")
return None
except Exception as e:
logger.error(f"Unexpected error fetching Cloudflare keys: {str(e)}")
return None

# Process keys for JWT validation
def _process_jwks(self, keys):
"""Convert RSA JWKs into the internal key format used for verification."""
processed_keys = []
for key_info in keys:
if key_info.get("kty") == "RSA":
# Extract RSA components
try:
processed_key = self._process_rsa_key(key_info)
if processed_key:
processed_keys.append(processed_key)
except Exception as e:
logger.warning(
f"Failed to process key {key_info.get('kid')}: {str(e)}"
)
continue

# Cache the keys
cache.set(cache_key, processed_keys, self.cache_timeout)
if key_info.get("kty") != "RSA":
continue
try:
processed_key = self._process_rsa_key(key_info)
except Exception as e:
logger.warning(f"Failed to process key {key_info.get('kid')}: {str(e)}")
continue
if processed_key:
processed_keys.append(processed_key)
return processed_keys

def _process_rsa_key(self, key_info):
Expand Down
1 change: 0 additions & 1 deletion packages/django-cf/pyproject.toml
Original file line number Diff line number Diff line change
Expand Up @@ -58,7 +58,6 @@ lint.extend-ignore = [
# Deferred: satisfying these requires behavioural or structural changes to
# code imported from https://github.com/G4brym/django-cf, so they are turned
# off to keep the lint adoption commit mechanical. Re-enable one rule per PR.
"PLR0912", # too many branches; needs decomposition
"PLR0913", # too many arguments; signature change
"PLR2004", # magic value comparison; needs named constants
"PLW0603", # global statement; architectural, see db/backends/do/storage.py
Expand Down
Loading