fix(api): stop locking the API key row on every authenticated request (#12882)

This commit is contained in:
César Arroba
2026-09-28 11:16:13 +02:00
committed by GitHub
parent c2b8092461
commit c114aa304b
4 changed files with 215 additions and 71 deletions
@@ -0,0 +1 @@
API key authentication no longer locks the key row on every request and now throttles `last_used_at` updates to once per 60 seconds, preventing a hot key from serializing all its requests onto a single locked row
+45 -33
View File
@@ -1,4 +1,5 @@
import logging
from datetime import timedelta
from math import isfinite
from uuid import UUID
@@ -6,7 +7,7 @@ from api.db_router import MainRouter
from api.models import TenantAPIKey, TenantAPIKeyManager
from cryptography.fernet import InvalidToken
from django.core.exceptions import ObjectDoesNotExist
from django.db import transaction
from django.db.models import Q
from django.utils import timezone
from drf_simple_apikey.backends import APIKeyAuthentication as BaseAPIKeyAuth
from drf_simple_apikey.crypto import get_crypto
@@ -18,12 +19,15 @@ from rest_framework_simplejwt.authentication import JWTAuthentication
logger = logging.getLogger(__name__)
# Writing on every request makes all requests of a busy key contend on one row
API_KEY_LAST_USED_AT_THROTTLE_SECONDS = 60
class OrphanedAPIKeyError(Exception):
"""Raised when an API key outlived the user that owns it.
Handled by `authenticate`, which commits the revocation written while detecting it
and then rejects the request with `AuthenticationFailed`.
The revocation is written by a plain `update()` before this is raised, so it is
already persisted by the time `authenticate` catches it and rejects the request.
"""
@@ -37,8 +41,9 @@ class TenantAPIKeyAuthentication(BaseAPIKeyAuth):
"""
Override to use admin connection, bypassing RLS during authentication.
Returns the validated API key row, locked with `select_for_update`, so callers
must run inside `transaction.atomic(using=MainRouter.admin_db)`.
Returns the validated API key row from a single read. `authenticate` builds
the auth claims from that same row instead of looking it up again, so a key
revoked or orphaned right after validation can't still authenticate.
"""
try:
payload = self.key_crypto.decrypt(key)
@@ -67,9 +72,11 @@ class TenantAPIKeyAuthentication(BaseAPIKeyAuth):
raise AuthenticationFailed("API Key has already expired.")
try:
# Loading `entity` in the same query keeps a user deleted after this read
# from turning the later `api_key.entity` access into a 500
api_key = (
self.model.objects.using(MainRouter.admin_db)
.select_for_update()
.select_related("entity")
.get(id=api_key_pk)
)
except ObjectDoesNotExist:
@@ -85,8 +92,9 @@ class TenantAPIKeyAuthentication(BaseAPIKeyAuth):
# Revoke it as well, so it stops showing up as active and later attempts fail
# the `revoked` check above like any other revoked key.
if api_key.entity_id is None:
api_key.revoked = True
api_key.save(update_fields=["revoked"], using=MainRouter.admin_db)
self.model.objects.using(MainRouter.admin_db).filter(
id=api_key.id, revoked=False
).update(revoked=True)
logger.warning(
"Revoked orphaned API key: prefix=%s tenant=%s",
api_key.prefix,
@@ -112,34 +120,38 @@ class TenantAPIKeyAuthentication(BaseAPIKeyAuth):
except ValueError:
raise AuthenticationFailed("Invalid API Key.")
# Validation, the `last_used_at` update and the auth claims all read the same
# row, locked until the transaction ends. Looking the key up a second time to
# build the claims used to leave a window where a key revoked or orphaned right
# after passing validation still authenticated.
with transaction.atomic(using=MainRouter.admin_db):
try:
api_key = self._authenticate_credentials(request, key)
except OrphanedAPIKeyError:
# Rejected below instead of here: leaving the block normally commits
# the revocation `_authenticate_credentials` wrote, while raising from
# inside would roll it back.
pass
else:
# The prefix used to be checked by the second lookup
if api_key.prefix != prefix:
raise AuthenticationFailed("Invalid API Key.")
try:
api_key = self._authenticate_credentials(request, key)
except OrphanedAPIKeyError:
raise AuthenticationFailed("No entity matching this api key.")
api_key.last_used_at = timezone.now()
api_key.save(update_fields=["last_used_at"], using=MainRouter.admin_db)
# The prefix used to be checked by the second lookup
if api_key.prefix != prefix:
raise AuthenticationFailed("Invalid API Key.")
entity = api_key.entity
return entity, {
"tenant_id": str(api_key.tenant_id),
"sub": str(entity.id),
"api_key_prefix": api_key.prefix,
}
self._throttled_touch_last_used_at(api_key)
raise AuthenticationFailed("No entity matching this api key.")
entity = api_key.entity
return entity, {
"tenant_id": str(api_key.tenant_id),
"sub": str(entity.id),
"api_key_prefix": api_key.prefix,
}
@staticmethod
def _throttled_touch_last_used_at(api_key: TenantAPIKey) -> None:
"""Write `last_used_at` at most once per throttle interval, without locking the row."""
now = timezone.now()
stale_before = now - timedelta(seconds=API_KEY_LAST_USED_AT_THROTTLE_SECONDS)
if api_key.last_used_at is not None and api_key.last_used_at >= stale_before:
return
TenantAPIKey.objects.using(MainRouter.admin_db).filter(
id=api_key.id, revoked=False
).filter(
Q(last_used_at__isnull=True) | Q(last_used_at__lt=stale_before)
).update(last_used_at=now)
class CombinedJWTOrAPIKeyAuthentication(BaseAuthentication):
@@ -1,9 +1,9 @@
import json
import time
from datetime import UTC, datetime, timedelta
from uuid import uuid4
import pytest
from api.authentication import API_KEY_LAST_USED_AT_THROTTLE_SECONDS
from api.db_router import MainRouter
from api.models import Membership, Role, TenantAPIKey, User, UserRoleRelationship
from api.signals import revoke_membership_api_keys, revoke_user_api_keys
@@ -11,6 +11,7 @@ from conftest import TEST_PASSWORD, get_api_tokens, get_authorization_header
from django.db.utils import ConnectionDoesNotExist
from django.urls import reverse
from drf_simple_apikey.crypto import get_crypto
from freezegun import freeze_time
from rest_framework.test import APIClient
from rest_framework_simplejwt.token_blacklist.models import (
BlacklistedToken,
@@ -527,7 +528,7 @@ class TestAPIKeyAuthentication:
def test_last_used_at_tracking(
self, create_test_user, tenants_fixture, api_keys_fixture
):
"""Verify last_used_at timestamp updates on each authentication."""
"""Verify last_used_at timestamp is set on first use and throttled after that."""
client = APIClient()
api_key = api_keys_fixture[0]
@@ -536,7 +537,11 @@ class TestAPIKeyAuthentication:
# Use API key to authenticate
api_key_headers = get_api_key_header(api_key._raw_key)
first_response = client.get(reverse("provider-list"), headers=api_key_headers)
start = datetime.now(UTC)
with freeze_time(start):
first_response = client.get(
reverse("provider-list"), headers=api_key_headers
)
assert first_response.status_code == 200
# Reload from database and check last_used_at is set
@@ -544,17 +549,23 @@ class TestAPIKeyAuthentication:
first_used_at = api_key.last_used_at
assert first_used_at is not None
# Use the same key again after a small delay
time.sleep(0.1)
# Using the same key again within the throttle interval does not rewrite it
second_response = client.get(reverse("provider-list"), headers=api_key_headers)
assert second_response.status_code == 200
# Reload and verify last_used_at was updated
api_key.refresh_from_db()
second_used_at = api_key.last_used_at
assert second_used_at is not None
assert second_used_at > first_used_at
assert api_key.last_used_at == first_used_at
# Past the throttle interval, the next use refreshes it
later = start + timedelta(seconds=API_KEY_LAST_USED_AT_THROTTLE_SECONDS + 1)
with freeze_time(later):
third_response = client.get(
reverse("provider-list"), headers=api_key_headers
)
assert third_response.status_code == 200
api_key.refresh_from_db()
assert api_key.last_used_at > first_used_at
@pytest.mark.django_db
@@ -1441,6 +1452,7 @@ class TestAPIKeyRLSBypass:
The update to last_used_at during authentication must also use the
admin database since it occurs before RLS context is established.
Past the throttle interval, using the key again refreshes the timestamp.
"""
client = APIClient()
api_key = api_keys_fixture[0]
@@ -1448,7 +1460,11 @@ class TestAPIKeyRLSBypass:
assert api_key.last_used_at is None
api_key_headers = get_api_key_header(api_key._raw_key)
first_response = client.get(reverse("provider-list"), headers=api_key_headers)
start = datetime.now(UTC)
with freeze_time(start):
first_response = client.get(
reverse("provider-list"), headers=api_key_headers
)
assert first_response.status_code == 200
@@ -1456,9 +1472,11 @@ class TestAPIKeyRLSBypass:
first_timestamp = api_key.last_used_at
assert first_timestamp is not None
time.sleep(0.1)
second_response = client.get(reverse("provider-list"), headers=api_key_headers)
later = start + timedelta(seconds=API_KEY_LAST_USED_AT_THROTTLE_SECONDS + 1)
with freeze_time(later):
second_response = client.get(
reverse("provider-list"), headers=api_key_headers
)
assert second_response.status_code == 200
api_key.refresh_from_db()
+137 -24
View File
@@ -5,16 +5,18 @@ from uuid import uuid4
import pytest
from api.authentication import (
API_KEY_LAST_USED_AT_THROTTLE_SECONDS,
OrphanedAPIKeyError,
SSEAuthentication,
TenantAPIKeyAuthentication,
)
from api.db_router import MainRouter
from api.models import TenantAPIKey
from api.models import TenantAPIKey, User
from django.db import connections
from django.db.models.query import QuerySet
from django.test import RequestFactory
from django.test.utils import CaptureQueriesContext
from freezegun import freeze_time
from rest_framework.exceptions import AuthenticationFailed
@@ -286,14 +288,15 @@ class TestTenantAPIKeyAuthentication:
assert str(exc_info.value.detail) == "This API Key has been revoked."
def test_authenticate_reads_the_api_key_once_under_a_row_lock(
def test_authenticate_reads_the_api_key_once_without_a_row_lock(
self, auth_backend, api_keys_fixture, request_factory
):
"""Test the API key is read a single time and the row is locked.
"""Test the API key is read a single time and no row is locked.
Validation, the `last_used_at` update and the claims must all come from the
same authoritative row: a second, unlocked lookup would reopen the window
where a key revoked in between still authenticates.
same authoritative row: a second lookup would reopen the window where a key
revoked in between still authenticates. `SELECT ... FOR UPDATE` serialized
every request for a hot key onto one locked row and is not used any more.
"""
api_key = api_keys_fixture[0]
@@ -310,33 +313,40 @@ class TestTenantAPIKeyAuthentication:
]
assert len(api_key_selects) == 1
assert "FOR UPDATE" in api_key_selects[0]
assert "FOR UPDATE" not in api_key_selects[0]
def test_authenticate_ignores_revocation_after_the_locked_read(
def test_authenticate_ignores_revocation_after_the_single_read(
self, auth_backend, api_keys_fixture, request_factory
):
"""Test the claims describe the row that was validated, not a later state.
Regression test: the key used to be looked up again to build the auth dict,
without rechecking `revoked` or `entity`. A key revoked or orphaned between
both reads still authenticated, and the claims came from that stale row. With
a single locked read the write below cannot land mid-authentication, and the
revocation only takes effect on the next request.
both reads still authenticated, and the claims came from that stale row.
There is now only a single read, so this race is closed by construction and
the revocation only takes effect on the next request.
"""
api_key = api_keys_fixture[0]
entity_at_validation = api_key.entity
original_save = TenantAPIKey.save
original_authenticate_credentials = (
TenantAPIKeyAuthentication._authenticate_credentials
)
def revoke_and_orphan_before_saving(instance, *args, **kwargs):
# Runs after validation, right before the claims are built: the exact
# window a concurrent revocation or user deletion used to slip into
def revoke_and_orphan_after_reading(self, request, key):
# Runs right after the single read `authenticate` will use to build the
# claims: the exact window a concurrent revocation used to slip into
result = original_authenticate_credentials(self, request, key)
TenantAPIKey.objects.filter(id=api_key.id).update(revoked=True, entity=None)
return original_save(instance, *args, **kwargs)
return result
request = request_factory.get("/")
request.META["HTTP_AUTHORIZATION"] = f"Api-Key {api_key._raw_key}"
with patch.object(TenantAPIKey, "save", revoke_and_orphan_before_saving):
with patch.object(
TenantAPIKeyAuthentication,
"_authenticate_credentials",
revoke_and_orphan_after_reading,
):
entity, auth_dict = auth_backend.authenticate(request)
assert entity == entity_at_validation
@@ -350,6 +360,43 @@ class TestTenantAPIKeyAuthentication:
assert str(exc_info.value.detail) == "This API Key has been revoked."
def test_authenticate_survives_owner_deleted_after_the_single_read(
self, auth_backend, api_keys_fixture, request_factory
):
"""Test a user deleted right after the read does not turn into a 500.
Without the row lock a concurrent user deletion can land between the read
and building the claims. `entity` is loaded by the same query, so no later
lookup can raise `DoesNotExist`.
"""
api_key = api_keys_fixture[0]
owner_id = api_key.entity_id
original_authenticate_credentials = (
TenantAPIKeyAuthentication._authenticate_credentials
)
def delete_owner_after_reading(self, request, key):
result = original_authenticate_credentials(self, request, key)
User.objects.using(MainRouter.admin_db).filter(id=owner_id).delete()
return result
request = request_factory.get("/")
request.META["HTTP_AUTHORIZATION"] = f"Api-Key {api_key._raw_key}"
with patch.object(
TenantAPIKeyAuthentication,
"_authenticate_credentials",
delete_owner_after_reading,
):
entity, auth_dict = auth_backend.authenticate(request)
assert auth_dict["sub"] == str(owner_id)
assert entity.id == owner_id
# From the next request on, the orphaned key is rejected with a 401
with pytest.raises(AuthenticationFailed):
auth_backend.authenticate(request)
def test_authenticate_expired_api_key(
self, auth_backend, create_test_user, tenants_fixture, request_factory
):
@@ -421,24 +468,90 @@ class TestTenantAPIKeyAuthentication:
if original_last_used:
assert api_key.last_used_at > original_last_used
def test_authenticate_saves_to_admin_database(
def test_authenticate_updates_last_used_at_on_admin_database(
self, auth_backend, api_keys_fixture, request_factory
):
"""Test that the API key save operation uses admin database."""
"""Test that the `last_used_at` update runs against the admin database."""
api_key = api_keys_fixture[0]
raw_key = api_key._raw_key
request = request_factory.get("/")
request.META["HTTP_AUTHORIZATION"] = f"Api-Key {raw_key}"
# Mock the save method to verify it's called with using='admin'
with patch.object(TenantAPIKey, "save") as mock_save:
with CaptureQueriesContext(connections[MainRouter.admin_db]) as captured:
auth_backend.authenticate(request)
# Verify save was called with using=admin_db
mock_save.assert_called_once_with(
update_fields=["last_used_at"], using=MainRouter.admin_db
)
api_key_updates = [
query["sql"]
for query in captured.captured_queries
if query["sql"].startswith("UPDATE") and '"api_keys"' in query["sql"]
]
assert len(api_key_updates) == 1
assert "last_used_at" in api_key_updates[0]
def test_authenticate_does_not_rewrite_last_used_at_within_throttle_interval(
self, auth_backend, api_keys_fixture, request_factory
):
"""Test that a second authentication within the throttle interval is a no-op write."""
api_key = api_keys_fixture[0]
raw_key = api_key._raw_key
request = request_factory.get("/")
request.META["HTTP_AUTHORIZATION"] = f"Api-Key {raw_key}"
# First call sets last_used_at
auth_backend.authenticate(request)
api_key.refresh_from_db()
first_used_at = api_key.last_used_at
assert first_used_at is not None
# Second call, still within the throttle interval, must issue no UPDATE
with CaptureQueriesContext(connections[MainRouter.admin_db]) as captured:
auth_backend.authenticate(request)
api_key_updates = [
query["sql"]
for query in captured.captured_queries
if query["sql"].startswith("UPDATE") and '"api_keys"' in query["sql"]
]
assert api_key_updates == []
api_key.refresh_from_db()
assert api_key.last_used_at == first_used_at
def test_authenticate_rewrites_last_used_at_after_throttle_interval(
self, auth_backend, api_keys_fixture, request_factory
):
"""Test that `last_used_at` is refreshed once it is older than the throttle interval."""
api_key = api_keys_fixture[0]
raw_key = api_key._raw_key
request = request_factory.get("/")
request.META["HTTP_AUTHORIZATION"] = f"Api-Key {raw_key}"
start = datetime.now(UTC)
with freeze_time(start):
auth_backend.authenticate(request)
api_key.refresh_from_db()
first_used_at = api_key.last_used_at
assert first_used_at is not None
later = start + timedelta(seconds=API_KEY_LAST_USED_AT_THROTTLE_SECONDS + 1)
with freeze_time(later):
with CaptureQueriesContext(connections[MainRouter.admin_db]) as captured:
auth_backend.authenticate(request)
api_key_updates = [
query["sql"]
for query in captured.captured_queries
if query["sql"].startswith("UPDATE") and '"api_keys"' in query["sql"]
]
assert len(api_key_updates) == 1
api_key.refresh_from_db()
assert api_key.last_used_at > first_used_at
def test_authenticate_returns_correct_auth_dict(
self, auth_backend, api_keys_fixture, request_factory