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..ed2f875f9c 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 @@ -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: @@ -259,13 +267,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 +294,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..47d12f5d87 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 @@ -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: @@ -262,13 +270,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 +297,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..9e37b544c1 --- /dev/null +++ b/tests/unit/test_reencrypt_version_check.py @@ -0,0 +1,284 @@ +"""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.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, 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 + +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_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( + 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"