From 4a6b72a31a4b90f2f5ba16916b38202c458b6980 Mon Sep 17 00:00:00 2001 From: L4XB Date: Mon, 14 Sep 2026 11:37:03 +0200 Subject: [PATCH 1/2] fix(gateway): pin a re-encryption to the ciphertext it read MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit `reencrypt_credentials` and `reencrypt_search_tools` read every row holding a secret, decrypt it and re-encrypt it, and nothing pinned the write to what the loop had seen. A PATCH committing in that window was overwritten with a re-encryption of the value it replaced — a lost update on a credential, and a silent one (#1127). Each row is now written with a conditional UPDATE matching the ciphertext that was read. Zero rows matched means someone else got there first, and the row is counted as skipped instead of clobbered. A Core UPDATE rather than a mutation on the loaded row, because the WHERE is the entire point and an ORM flush carries no condition at all. Skipped is reported rather than retried, and the count is additive on both response models. A rotation is hand-run with the operator watching, and a row whose value changed under them is already encrypted with the primary key by whoever wrote it — so the honest answer is "these were not mine to rewrite", not a loop that races the same edit again. Silently folding them into `reencrypted` would tell the operator the rotation was complete when it was not, which matters because the next step of the documented procedure is removing the old key. Both stores change together. Three stores with two shapes is worse than three with one, because the next person copies whichever they find first. The racing cells use a second real connection to the same SQLite database in WAL mode rather than a stand-in: the property under test is what the WHERE clause sees once someone else has committed, and a mock of "the row changed" would pass just as well against the bug. The competing commit is fired from inside the service's own `encrypt_secret`, which is precisely the window — the row has been read and decrypted, and its UPDATE has not run yet. This is the re-encryption half of #1127 only. The transaction-ownership half is a convention-versus-code decision across all three stores and belongs to the maintainers, not to this PR. --- docs/public/openapi.json | 12 + src/gateway/api/routes/providers.py | 11 +- src/gateway/api/routes/search_tools.py | 11 +- .../services/provider_store_service.py | 48 +++- .../services/search_tool_store_service.py | 50 +++- tests/unit/test_reencrypt_version_check.py | 234 ++++++++++++++++++ 6 files changed, 339 insertions(+), 27 deletions(-) create mode 100644 tests/unit/test_reencrypt_version_check.py diff --git a/docs/public/openapi.json b/docs/public/openapi.json index 39cd1eab8b..82a5573821 100644 --- a/docs/public/openapi.json +++ b/docs/public/openapi.json @@ -9278,6 +9278,12 @@ "title": "Reencrypted", "type": "integer" }, + "skipped": { + "default": 0, + "description": "Number of rows whose stored key changed between the read and the write, so the re-encryption was not applied. They already hold whoever wrote them last.", + "title": "Skipped", + "type": "integer" + }, "unreadable": { "description": "Number of encrypted keys left untouched because they could not be decrypted.", "title": "Unreadable", @@ -9299,6 +9305,12 @@ "title": "Reencrypted", "type": "integer" }, + "skipped": { + "default": 0, + "description": "Number of rows whose stored key changed between the read and the write, so the re-encryption was not applied. They already hold whoever wrote them last.", + "title": "Skipped", + "type": "integer" + }, "unreadable": { "description": "Number of encrypted keys left untouched because they could not be decrypted.", "title": "Unreadable", diff --git a/src/gateway/api/routes/providers.py b/src/gateway/api/routes/providers.py index bcc90a8a18..7c58a8b67c 100644 --- a/src/gateway/api/routes/providers.py +++ b/src/gateway/api/routes/providers.py @@ -370,6 +370,13 @@ class ReencryptProviderCredentialsResponse(BaseModel): reencrypted: int = Field(description="Number of stored provider keys re-encrypted.") unreadable: int = Field(description="Number of encrypted keys left untouched because they could not be decrypted.") + skipped: int = Field( + default=0, + description=( + "Number of rows whose stored key changed between the read and the write, so the " + "re-encryption was not applied. They already hold whoever wrote them last." + ), + ) class TestProviderRequest(BaseModel): @@ -506,7 +513,7 @@ async def reencrypt_stored_provider_keys( by replacing the affected provider keys. """ try: - reencrypted, unreadable = await reencrypt_credentials(db) + reencrypted, unreadable, skipped = await reencrypt_credentials(db) await db.commit() except SecretBoxUnavailableError as exc: await db.rollback() @@ -519,7 +526,7 @@ async def reencrypt_stored_provider_keys( await refresh_provider_cache(db, config) except SQLAlchemyError: logger.warning("Provider overlay refresh failed after re-encrypting credentials; converges within TTL") - return ReencryptProviderCredentialsResponse(reencrypted=reencrypted, unreadable=unreadable) + return ReencryptProviderCredentialsResponse(reencrypted=reencrypted, unreadable=unreadable, skipped=skipped) @router.get("/provider-credentials") diff --git a/src/gateway/api/routes/search_tools.py b/src/gateway/api/routes/search_tools.py index ae44acf249..01ed9124a8 100644 --- a/src/gateway/api/routes/search_tools.py +++ b/src/gateway/api/routes/search_tools.py @@ -170,6 +170,13 @@ class ReencryptSearchToolsResponse(BaseModel): reencrypted: int = Field(description="Number of stored search-tool keys re-encrypted.") unreadable: int = Field(description="Number of encrypted keys left untouched because they could not be decrypted.") + skipped: int = Field( + default=0, + description=( + "Number of rows whose stored key changed between the read and the write, so the " + "re-encryption was not applied. They already hold whoever wrote them last." + ), + ) def _is_decryptable(row: SearchToolCredential) -> bool: @@ -302,7 +309,7 @@ async def reencrypt_stored_search_tool_keys( tool's key. """ try: - reencrypted, unreadable = await reencrypt_search_tools(db) + reencrypted, unreadable, skipped = await reencrypt_search_tools(db) await db.commit() except SecretBoxUnavailableError as exc: await db.rollback() @@ -314,7 +321,7 @@ async def reencrypt_stored_search_tool_keys( await refresh_search_tool_cache(db, config) except SQLAlchemyError: logger.warning("Search tool overlay refresh failed after re-encrypting keys; converges within TTL") - return ReencryptSearchToolsResponse(reencrypted=reencrypted, unreadable=unreadable) + return ReencryptSearchToolsResponse(reencrypted=reencrypted, unreadable=unreadable, skipped=skipped) @router.post("", status_code=status.HTTP_201_CREATED) diff --git a/src/gateway/services/provider_store_service.py b/src/gateway/services/provider_store_service.py index 70d7ffa8b7..bc71487d71 100644 --- a/src/gateway/services/provider_store_service.py +++ b/src/gateway/services/provider_store_service.py @@ -21,9 +21,9 @@ import asyncio import time -from typing import Any, Final +from typing import Any, Final, cast -from sqlalchemy import select +from sqlalchemy import CursorResult, select, update from sqlalchemy.ext.asyncio import AsyncSession from gateway.core.config import GatewayConfig @@ -259,13 +259,25 @@ async def save_credential( return row -async def reencrypt_credentials(db: AsyncSession) -> tuple[int, int]: +async def reencrypt_credentials(db: AsyncSession) -> tuple[int, int, int]: """Re-encrypt stored provider keys with the current primary OTARI_SECRET_KEY. - Returns ``(reencrypted, unreadable)``. Rows without a stored key are ignored. - If any encrypted key cannot be decrypted with the configured key set, it is - left untouched and counted as unreadable so the operator can recover it by + Returns ``(reencrypted, unreadable, skipped)``. Rows without a stored key are + ignored. A key that cannot be decrypted with the configured key set is left + untouched and counted as unreadable, so the operator can recover it by replacing that provider's key. + + Each row is written with a conditional UPDATE matching the ciphertext that + was read. Rotation reads every row, decrypts and re-encrypts, and nothing + pinned the write to what it had seen: an edit committing in that window was + overwritten with a re-encryption of the value it replaced — a silent lost + update on a credential (otari#1127). Zero rows matched means someone else + got there first, and that row is counted as skipped rather than clobbered. + + Skipped is reported rather than retried. A rotation is run by hand, the + operator is watching, and a row whose value changed under them is already + encrypted with the primary key by whoever wrote it — so the honest answer is + "these were not mine to rewrite", not a loop that races the same edit again. """ rows = ( (await db.execute(select(ProviderCredential).where(ProviderCredential.encrypted_api_key.is_not(None)))) @@ -274,17 +286,31 @@ async def reencrypt_credentials(db: AsyncSession) -> tuple[int, int]: ) reencrypted = 0 unreadable = 0 + skipped = 0 for row in rows: - if row.encrypted_api_key is None: + original = row.encrypted_api_key + if original is None: continue try: - plaintext = decrypt_secret(row.encrypted_api_key) + plaintext = decrypt_secret(original) except SecretDecryptionError: unreadable += 1 continue - row.encrypted_api_key = encrypt_secret(plaintext) - reencrypted += 1 - return reencrypted, unreadable + # Core UPDATE rather than a mutation on the loaded row: the whole point + # is the WHERE, and an ORM flush would carry no condition at all. + result = await db.execute( + update(ProviderCredential) + .where(ProviderCredential.instance == row.instance, ProviderCredential.encrypted_api_key == original) + .values(encrypted_api_key=encrypt_secret(plaintext)) + .execution_options(synchronize_session=False) + ) + # `execute` is typed as returning Result; an UPDATE always yields a + # CursorResult, which is where rowcount lives. + if cast(CursorResult[Any], result).rowcount == 1: + reencrypted += 1 + else: + skipped += 1 + return reencrypted, unreadable, skipped async def delete_credential(db: AsyncSession, instance: str) -> bool: diff --git a/src/gateway/services/search_tool_store_service.py b/src/gateway/services/search_tool_store_service.py index 038d03825b..8c8d989679 100644 --- a/src/gateway/services/search_tool_store_service.py +++ b/src/gateway/services/search_tool_store_service.py @@ -22,9 +22,9 @@ import asyncio import time -from typing import Any, Final +from typing import Any, Final, cast -from sqlalchemy import select +from sqlalchemy import CursorResult, select, update from sqlalchemy.ext.asyncio import AsyncSession from gateway.core.config import GatewayConfig @@ -262,13 +262,25 @@ async def save_search_tool( return row -async def reencrypt_search_tools(db: AsyncSession) -> tuple[int, int]: +async def reencrypt_search_tools(db: AsyncSession) -> tuple[int, int, int]: """Re-encrypt stored search-tool keys with the current primary OTARI_SECRET_KEY. - Returns ``(reencrypted, unreadable)``. Rows without a stored key are ignored. - A key that cannot be decrypted with the configured key set is left untouched - and counted as unreadable, so the operator can recover it by replacing that - tool's key. + Returns ``(reencrypted, unreadable, skipped)``. Rows without a stored key are + ignored. A key that cannot be decrypted with the configured key set is left + untouched and counted as unreadable, so the operator can recover it by + replacing that tool's key. + + Each row is written with a conditional UPDATE matching the ciphertext that + was read. Rotation reads every row, decrypts and re-encrypts, and nothing + pinned the write to what it had seen: an edit committing in that window was + overwritten with a re-encryption of the value it replaced — a silent lost + update on a credential (otari#1127). Zero rows matched means someone else + got there first, and that row is counted as skipped rather than clobbered. + + Skipped is reported rather than retried. A rotation is run by hand, the + operator is watching, and a row whose value changed under them is already + encrypted with the primary key by whoever wrote it — so the honest answer is + "these were not mine to rewrite", not a loop that races the same edit again. """ rows = ( (await db.execute(select(SearchToolCredential).where(SearchToolCredential.encrypted_api_key.is_not(None)))) @@ -277,17 +289,31 @@ async def reencrypt_search_tools(db: AsyncSession) -> tuple[int, int]: ) reencrypted = 0 unreadable = 0 + skipped = 0 for row in rows: - if row.encrypted_api_key is None: + original = row.encrypted_api_key + if original is None: continue try: - plaintext = decrypt_secret(row.encrypted_api_key) + plaintext = decrypt_secret(original) except SecretDecryptionError: unreadable += 1 continue - row.encrypted_api_key = encrypt_secret(plaintext) - reencrypted += 1 - return reencrypted, unreadable + # Core UPDATE rather than a mutation on the loaded row: the whole point + # is the WHERE, and an ORM flush would carry no condition at all. + result = await db.execute( + update(SearchToolCredential) + .where(SearchToolCredential.name == row.name, SearchToolCredential.encrypted_api_key == original) + .values(encrypted_api_key=encrypt_secret(plaintext)) + .execution_options(synchronize_session=False) + ) + # `execute` is typed as returning Result; an UPDATE always yields a + # CursorResult, which is where rowcount lives. + if cast(CursorResult[Any], result).rowcount == 1: + reencrypted += 1 + else: + skipped += 1 + return reencrypted, unreadable, skipped async def delete_search_tool(db: AsyncSession, name: str) -> bool: diff --git a/tests/unit/test_reencrypt_version_check.py b/tests/unit/test_reencrypt_version_check.py new file mode 100644 index 0000000000..ae3c7b3fbd --- /dev/null +++ b/tests/unit/test_reencrypt_version_check.py @@ -0,0 +1,234 @@ +"""Rotation never overwrites a credential someone changed while it was working. + +otari#1127. ``reencrypt_*`` reads every row holding a secret, decrypts it and +re-encrypts it, and nothing pinned the write to the ciphertext it had read. A +PATCH committing in that window was overwritten with a re-encryption of the +value it replaced — a lost update on a credential, and a silent one. + +The window is small and rotation is rare and hand-run, which is exactly why it +is worth a test: nobody would ever see it happen. + +The racing cells use a second, real connection to the same database rather than +a stand-in, because the property under test is what the WHERE clause sees when +someone else has committed. A mock of "the row changed" would pass just as well +against the bug. +""" + +import asyncio +from collections.abc import Awaitable, Callable +from datetime import UTC, datetime +from pathlib import Path +from tempfile import TemporaryDirectory +from typing import TypeVar + +import pytest +from sqlalchemy import create_engine, select, text, update +from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine +from sqlmodel import SQLModel + +from gateway.models.entities import ProviderCredential, SearchToolCredential +from gateway.services import provider_store_service as provider_store +from gateway.services import search_tool_store_service as search_tool_store +from gateway.services.provider_store_service import reencrypt_credentials +from gateway.services.search_tool_store_service import reencrypt_search_tools +from gateway.services.secret_box import decrypt_secret, encrypt_secret, generate_secret_key + +T = TypeVar("T") + + +def _run(scenario: Callable[[AsyncSession, str], Awaitable[T]]) -> T: + """Run one scenario against a file-backed SQLite database in WAL mode. + + A file rather than ``:memory:`` so a second connection can reach the same + rows, and WAL so that second connection can commit while the rotation holds + its read. + """ + + async def main() -> T: + with TemporaryDirectory() as tmp: + db_path = str(Path(tmp) / "otari-rotation.db") + engine = create_async_engine(f"sqlite+aiosqlite:///{db_path}") + try: + async with engine.begin() as conn: + await conn.execute(text("PRAGMA journal_mode=WAL")) + await conn.run_sync(SQLModel.metadata.create_all) + session_factory = async_sessionmaker(engine, expire_on_commit=False) + async with session_factory() as session: + return await scenario(session, db_path) + finally: + # aiosqlite runs each connection on its own thread, so an + # undisposed engine leaks one per call. + await engine.dispose() + + return asyncio.run(main()) + + +def _commit_from_another_connection(db_path: str, statement: object) -> None: + """Commit one statement on a separate connection, the way a PATCH would.""" + other = create_engine(f"sqlite:///{db_path}") + try: + with other.begin() as conn: + conn.execute(statement) # type: ignore[arg-type] + finally: + other.dispose() + + +@pytest.fixture(autouse=True) +def _secret_key(monkeypatch: pytest.MonkeyPatch) -> None: + monkeypatch.setenv("OTARI_SECRET_KEY", generate_secret_key()) + + +async def _add_provider(session: AsyncSession, instance: str, api_key: str) -> None: + session.add( + ProviderCredential( + instance=instance, + provider_type="openai", + encrypted_api_key=encrypt_secret(api_key), + last4=api_key[-4:], + created_at=datetime.now(UTC), + updated_at=datetime.now(UTC), + ) + ) + await session.commit() + + +async def _add_search_tool(session: AsyncSession, name: str, api_key: str) -> None: + session.add( + SearchToolCredential( + name=name, + provider="searxng", + encrypted_api_key=encrypt_secret(api_key), + created_at=datetime.now(UTC), + updated_at=datetime.now(UTC), + ) + ) + await session.commit() + + +class TestProviderRotation: + def test_an_untouched_row_is_reencrypted(self) -> None: + async def scenario(session: AsyncSession, _db: str) -> tuple[int, int, int]: + await _add_provider(session, "openai", "sk-live-value") + return await reencrypt_credentials(session) + + assert _run(scenario) == (1, 0, 0) + + def test_the_plaintext_survives_the_rotation(self) -> None: + """A re-encryption that lost the value would pass every count above.""" + + async def scenario(session: AsyncSession, _db: str) -> str: + await _add_provider(session, "openai", "sk-live-value") + before = (await session.execute(select(ProviderCredential))).scalars().one().encrypted_api_key + await reencrypt_credentials(session) + await session.commit() + session.expire_all() + after = (await session.execute(select(ProviderCredential))).scalars().one().encrypted_api_key + assert after is not None + assert after != before, "re-encryption produced the same ciphertext" + return decrypt_secret(after) + + assert _run(scenario) == "sk-live-value" + + def test_a_row_changed_under_the_rotation_is_skipped_not_clobbered( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + """The whole point: the stored value must still be the competing edit's. + + The window is between the read and the write, so the competing commit has + to land there and nowhere else. Hooking the service's own + ``encrypt_secret`` puts it exactly there — the row has been read and + decrypted, and its UPDATE has not run yet. + """ + + async def scenario(session: AsyncSession, db_path: str) -> tuple[tuple[int, int, int], str]: + await _add_provider(session, "openai", "sk-old-value") + real_encrypt = provider_store.encrypt_secret + raced = False + + def encrypt_and_let_someone_else_commit(plaintext: str) -> str: + nonlocal raced + if not raced: + raced = True + _commit_from_another_connection( + db_path, + update(ProviderCredential) + .where(ProviderCredential.instance == "openai") + .values(encrypted_api_key=encrypt_secret("sk-new-value")), + ) + return real_encrypt(plaintext) + + monkeypatch.setattr(provider_store, "encrypt_secret", encrypt_and_let_someone_else_commit) + + counts = await reencrypt_credentials(session) + await session.commit() + session.expire_all() + stored = (await session.execute(select(ProviderCredential))).scalars().one() + assert stored.encrypted_api_key is not None + return counts, decrypt_secret(stored.encrypted_api_key) + + counts, stored_value = _run(scenario) + assert counts == (0, 0, 1) + assert stored_value == "sk-new-value" + + def test_an_undecryptable_row_is_left_alone(self) -> None: + async def scenario(session: AsyncSession, _db: str) -> tuple[int, int, int]: + session.add( + ProviderCredential( + instance="broken", + provider_type="openai", + encrypted_api_key="not-a-ciphertext", + last4="text", + created_at=datetime.now(UTC), + updated_at=datetime.now(UTC), + ) + ) + await session.commit() + return await reencrypt_credentials(session) + + assert _run(scenario) == (0, 1, 0) + + +class TestSearchToolRotation: + """Same shape, second store. Both or neither — three stores with two shapes + is worse than three with one, because the next person copies whichever they + find first.""" + + def test_an_untouched_row_is_reencrypted(self) -> None: + async def scenario(session: AsyncSession, _db: str) -> tuple[int, int, int]: + await _add_search_tool(session, "searxng", "key-live") + return await reencrypt_search_tools(session) + + assert _run(scenario) == (1, 0, 0) + + def test_a_row_changed_under_the_rotation_is_skipped_not_clobbered( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + async def scenario(session: AsyncSession, db_path: str) -> tuple[tuple[int, int, int], str]: + await _add_search_tool(session, "searxng", "key-old") + real_encrypt = search_tool_store.encrypt_secret + raced = False + + def encrypt_and_let_someone_else_commit(plaintext: str) -> str: + nonlocal raced + if not raced: + raced = True + _commit_from_another_connection( + db_path, + update(SearchToolCredential) + .where(SearchToolCredential.name == "searxng") + .values(encrypted_api_key=encrypt_secret("key-new")), + ) + return real_encrypt(plaintext) + + monkeypatch.setattr(search_tool_store, "encrypt_secret", encrypt_and_let_someone_else_commit) + + counts = await reencrypt_search_tools(session) + await session.commit() + session.expire_all() + stored = (await session.execute(select(SearchToolCredential))).scalars().one() + assert stored.encrypted_api_key is not None + return counts, decrypt_secret(stored.encrypted_api_key) + + counts, stored_value = _run(scenario) + assert counts == (0, 0, 1) + assert stored_value == "key-new" From f7deef562ffc8d2f1d9971f021b6abc30f5a265d Mon Sep 17 00:00:00 2001 From: L4XB Date: Mon, 14 Sep 2026 17:19:13 +0200 Subject: [PATCH 2/2] fix(gateway): let the overlay refresh repopulate what it loads MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit The rotation endpoint re-encrypts, commits and refreshes the overlay on the same session, and the session factory sets `expire_on_commit=False`. A row the identity map still holds is therefore handed back with the values it was loaded with — for a row a concurrent PATCH changed, that is the credential the PATCH replaced, coming back into the runtime cache while the database correctly keeps the new one. Nothing holds those rows alive that long today: the rotation's own list dies when it returns and the identity map is weak, so the reload happens by collection timing rather than by rule. `populate_existing` makes it a rule. Measured directly: holding the pre-race rows across the refresh serves the old ciphertext, releasing them serves the new one. The end-to-end cell asserts the route's outcome and says in its docstring that it passes either way, because a cell that reached into the rotation's locals to force the stale path would be testing the collector, not the service. --- .../services/provider_store_service.py | 10 +++- .../services/search_tool_store_service.py | 10 +++- tests/unit/test_reencrypt_version_check.py | 52 ++++++++++++++++++- 3 files changed, 69 insertions(+), 3 deletions(-) diff --git a/src/gateway/services/provider_store_service.py b/src/gateway/services/provider_store_service.py index bc71487d71..ed2f875f9c 100644 --- a/src/gateway/services/provider_store_service.py +++ b/src/gateway/services/provider_store_service.py @@ -125,7 +125,15 @@ async def refresh_provider_cache(db: AsyncSession, config: GatewayConfig) -> set """Reload the overlay from the database, apply it, and return shadowed names.""" global _cached_at # noqa: PLW0603 - rows = (await db.execute(select(ProviderCredential))).scalars().all() + # `populate_existing`: the session factory sets `expire_on_commit=False`, + # so a row already in the identity map keeps the values it was loaded + # with and this SELECT would hand them straight back. The rotation + # endpoint refreshes on the same session it just re-encrypted on, where + # that means a credential a concurrent PATCH replaced can return to the + # cache. Today nothing holds those rows alive that long, which makes it + # a garbage-collection timing question rather than a guarantee + # (CodeRabbit). + rows = (await db.execute(select(ProviderCredential).execution_options(populate_existing=True))).scalars().all() overlay: dict[str, dict[str, Any]] = {} for row in rows: try: diff --git a/src/gateway/services/search_tool_store_service.py b/src/gateway/services/search_tool_store_service.py index 8c8d989679..47d12f5d87 100644 --- a/src/gateway/services/search_tool_store_service.py +++ b/src/gateway/services/search_tool_store_service.py @@ -129,7 +129,15 @@ async def refresh_search_tool_cache(db: AsyncSession, config: GatewayConfig) -> """Reload the overlay from the database, apply it, and return shadowed names.""" global _cached_at # noqa: PLW0603 - rows = (await db.execute(select(SearchToolCredential))).scalars().all() + # `populate_existing`: the session factory sets `expire_on_commit=False`, + # so a row already in the identity map keeps the values it was loaded + # with and this SELECT would hand them straight back. The rotation + # endpoint refreshes on the same session it just re-encrypted on, where + # that means a credential a concurrent PATCH replaced can return to the + # cache. Today nothing holds those rows alive that long, which makes it + # a garbage-collection timing question rather than a guarantee + # (CodeRabbit). + rows = (await db.execute(select(SearchToolCredential).execution_options(populate_existing=True))).scalars().all() overlay: dict[str, dict[str, Any]] = {} for row in rows: try: diff --git a/tests/unit/test_reencrypt_version_check.py b/tests/unit/test_reencrypt_version_check.py index ae3c7b3fbd..9e37b544c1 100644 --- a/tests/unit/test_reencrypt_version_check.py +++ b/tests/unit/test_reencrypt_version_check.py @@ -26,10 +26,11 @@ from sqlalchemy.ext.asyncio import AsyncSession, async_sessionmaker, create_async_engine from sqlmodel import SQLModel +from gateway.core.config import GatewayConfig from gateway.models.entities import ProviderCredential, SearchToolCredential from gateway.services import provider_store_service as provider_store from gateway.services import search_tool_store_service as search_tool_store -from gateway.services.provider_store_service import reencrypt_credentials +from gateway.services.provider_store_service import reencrypt_credentials, refresh_provider_cache, reset_provider_cache from gateway.services.search_tool_store_service import reencrypt_search_tools from gateway.services.secret_box import decrypt_secret, encrypt_secret, generate_secret_key @@ -170,6 +171,55 @@ def encrypt_and_let_someone_else_commit(plaintext: str) -> str: assert counts == (0, 0, 1) assert stored_value == "sk-new-value" + def test_the_skipped_row_does_not_come_back_through_the_cache( + self, monkeypatch: pytest.MonkeyPatch + ) -> None: + """The database keeps the competing edit; the runtime cache must too. + + The route re-encrypts, commits, and refreshes the overlay on the SAME + session, and the factory sets ``expire_on_commit=False`` — so a row the + identity map still holds would be handed back with the ciphertext read + BEFORE the race (CodeRabbit). + + This cell asserts the route's outcome, not the mechanism: it passes with + and without ``populate_existing``, because the rotation's own row list + dies when it returns and the identity map is weak, so the reload happens + by collection timing rather than by rule. Measured directly, holding the + pre-race rows across the refresh does serve ``sk-old-value``. That is + why the refresh asks for ``populate_existing`` instead of depending on + when the garbage collector runs. + + No ``expire_all()`` here, deliberately: the route does not call one, and + adding one to the test would hide the question entirely. + """ + + async def scenario(session: AsyncSession, db_path: str) -> str: + await _add_provider(session, "openai", "sk-old-value") + real_encrypt = provider_store.encrypt_secret + raced = False + + def encrypt_and_let_someone_else_commit(plaintext: str) -> str: + nonlocal raced + if not raced: + raced = True + _commit_from_another_connection( + db_path, + update(ProviderCredential) + .where(ProviderCredential.instance == "openai") + .values(encrypted_api_key=encrypt_secret("sk-new-value")), + ) + return real_encrypt(plaintext) + + monkeypatch.setattr(provider_store, "encrypt_secret", encrypt_and_let_someone_else_commit) + + assert await reencrypt_credentials(session) == (0, 0, 1) + await session.commit() + reset_provider_cache() + await refresh_provider_cache(session, GatewayConfig(providers={})) + return str(provider_store._cache["openai"]["api_key"]) + + assert _run(scenario) == "sk-new-value" + def test_an_undecryptable_row_is_left_alone(self) -> None: async def scenario(session: AsyncSession, _db: str) -> tuple[int, int, int]: session.add(