diff --git a/api/changelog.d/api-key-auth-no-row-lock.fixed.md b/api/changelog.d/api-key-auth-no-row-lock.fixed.md new file mode 100644 index 0000000000..3d0a40260c --- /dev/null +++ b/api/changelog.d/api-key-auth-no-row-lock.fixed.md @@ -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 diff --git a/api/src/backend/api/authentication.py b/api/src/backend/api/authentication.py index af1a35c7c0..2c99edfbad 100644 --- a/api/src/backend/api/authentication.py +++ b/api/src/backend/api/authentication.py @@ -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): diff --git a/api/src/backend/api/tests/integration/test_authentication.py b/api/src/backend/api/tests/integration/test_authentication.py index 926ba3a133..8e9040e117 100644 --- a/api/src/backend/api/tests/integration/test_authentication.py +++ b/api/src/backend/api/tests/integration/test_authentication.py @@ -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() diff --git a/api/src/backend/api/tests/test_authentication.py b/api/src/backend/api/tests/test_authentication.py index 782bfe670e..b01b709146 100644 --- a/api/src/backend/api/tests/test_authentication.py +++ b/api/src/backend/api/tests/test_authentication.py @@ -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