From b89bd76e22af1fb7d60c318d01f73eb160c2a8ff Mon Sep 17 00:00:00 2001 From: olayinkaadelakun Date: Wed, 28 Jan 2026 12:41:35 -0500 Subject: [PATCH] fix: Encrypt API KEY (#11335) * fix: encrypt api key * fix: remove unnecessary if condition * fix: remove unnecessary if condition * [autofix.ci] apply automated fixes * fix: added better exception laballing * [autofix.ci] apply automated fixes * [autofix.ci] apply automated fixes * improve functions and clean up * improve functions and clean up * improve functions and clean up * fix: add benchmark test for large api_key set * Revert "fix: add benchmark test for large api_key set" This reverts commit 4e86d04df0c8d9f6379332165f8f22787946d389. * fix: add benchmark test for large api_key set * fix: use real DB for better simulated testing * [autofix.ci] apply automated fixes * fix testcases * fix testcases * fix testcases * fix testcases * fix testcases * fix ruff * improved space complexity by adding dependency injection * this fix helps ensure that testcases don't have to use api_key that start with sk- * fix: update testcase --------- Co-authored-by: Olayinka Adelakun Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> Co-authored-by: Olayinka Adelakun Co-authored-by: Himavarsha <40851462+HimavarshaVS@users.noreply.github.com> --- src/backend/base/langflow/api/v1/api_key.py | 5 +- .../base/langflow/services/auth/utils.py | 8 +- .../services/database/models/api_key/crud.py | 81 ++++++++++--- .../tests/performance/check_key_benchmark.py | 111 ++++++++++++++++++ src/backend/tests/unit/test_api_key_source.py | 81 +++++++------ src/backend/tests/unit/test_get_api_key.py | 78 ++++++++++++ 6 files changed, 308 insertions(+), 56 deletions(-) create mode 100644 src/backend/tests/performance/check_key_benchmark.py create mode 100644 src/backend/tests/unit/test_get_api_key.py diff --git a/src/backend/base/langflow/api/v1/api_key.py b/src/backend/base/langflow/api/v1/api_key.py index 45da582aef..bcd18f0272 100644 --- a/src/backend/base/langflow/api/v1/api_key.py +++ b/src/backend/base/langflow/api/v1/api_key.py @@ -21,9 +21,8 @@ async def get_api_keys_route( ) -> ApiKeysResponse: try: user_id = current_user.id - keys = await get_api_keys(db, user_id) - - return ApiKeysResponse(total_count=len(keys), user_id=user_id, api_keys=keys) + api_keys = await get_api_keys(db, user_id) + return ApiKeysResponse(total_count=len(api_keys), user_id=user_id, api_keys=api_keys) except Exception as exc: raise HTTPException(status_code=400, detail=str(exc)) from exc diff --git a/src/backend/base/langflow/services/auth/utils.py b/src/backend/base/langflow/services/auth/utils.py index 73231a7269..9719855723 100644 --- a/src/backend/base/langflow/services/auth/utils.py +++ b/src/backend/base/langflow/services/auth/utils.py @@ -686,7 +686,7 @@ def encrypt_api_key(api_key: str, settings_service: SettingsService): return encrypted_key.decode() -def decrypt_api_key(encrypted_api_key: str, settings_service: SettingsService): +def decrypt_api_key(encrypted_api_key: str, settings_service: SettingsService, fernet_obj: Fernet | None = None) -> str: """Decrypt the provided encrypted API key using Fernet decryption. This function supports both encrypted and plain text values. It first attempts @@ -699,12 +699,16 @@ def decrypt_api_key(encrypted_api_key: str, settings_service: SettingsService): Args: encrypted_api_key (str): The encrypted API key or plain text value. settings_service (SettingsService): Service providing authentication settings. + fernet_obj (Fernet | None): Optional pre-initialized Fernet object. Returns: str: The decrypted API key, the original value if plain text, or empty string if it's encrypted with a different key. """ - fernet = get_fernet(settings_service) + fernet = fernet_obj + if fernet is None: + fernet = get_fernet(settings_service) + if isinstance(encrypted_api_key, str): try: return fernet.decrypt(encrypted_api_key.encode()).decode() diff --git a/src/backend/base/langflow/services/database/models/api_key/crud.py b/src/backend/base/langflow/services/database/models/api_key/crud.py index eb122da780..3564a6eee5 100644 --- a/src/backend/base/langflow/services/database/models/api_key/crud.py +++ b/src/backend/base/langflow/services/database/models/api_key/crud.py @@ -1,13 +1,15 @@ +import binascii import datetime import os import secrets from typing import TYPE_CHECKING from uuid import UUID -from sqlalchemy.orm import selectinload -from sqlmodel import select +from cryptography.fernet import InvalidToken +from sqlmodel import select, update from sqlmodel.ext.asyncio.session import AsyncSession +from langflow.services.auth import utils as auth_utils from langflow.services.database.models.api_key.model import ApiKey, ApiKeyCreate, ApiKeyRead, UnmaskedApiKeyRead from langflow.services.database.models.user.model import User from langflow.services.deps import get_settings_service @@ -17,17 +19,42 @@ if TYPE_CHECKING: async def get_api_keys(session: AsyncSession, user_id: UUID) -> list[ApiKeyRead]: + """Get all API keys for a user with decrypted values.""" + settings_service = get_settings_service() query: SelectOfScalar = select(ApiKey).where(ApiKey.user_id == user_id) - api_keys = (await session.exec(query)).all() - return [ApiKeyRead.model_validate(api_key) for api_key in api_keys] + api_key_objects = (await session.exec(query)).all() + + fernet = auth_utils.get_fernet(settings_service) + api_keys = [] + for api_key_obj in api_key_objects: + data = api_key_obj.model_dump() + + api_key = data.get("api_key") + if api_key: + try: + actual_key = auth_utils.decrypt_api_key(api_key, settings_service=settings_service, fernet_obj=fernet) + except (ValueError, TypeError, InvalidToken, UnicodeDecodeError, AttributeError, binascii.Error): + # Fallback to stored value for legacy entries + actual_key = api_key + else: + actual_key = api_key + + data["api_key"] = actual_key + api_keys.append(ApiKeyRead.model_validate(data)) + + return api_keys async def create_api_key(session: AsyncSession, api_key_create: ApiKeyCreate, user_id: UUID) -> UnmaskedApiKeyRead: # Generate a random API key with 32 bytes of randomness generated_api_key = f"sk-{secrets.token_urlsafe(32)}" + settings_service = get_settings_service() + + stored_api_key = auth_utils.encrypt_api_key(generated_api_key, settings_service=settings_service) + api_key = ApiKey( - api_key=generated_api_key, + api_key=stored_api_key, name=api_key_create.name, user_id=user_id, created_at=api_key_create.created_at or datetime.datetime.now(datetime.timezone.utc), @@ -73,15 +100,41 @@ async def check_key(session: AsyncSession, api_key: str) -> User | None: async def _check_key_from_db(session: AsyncSession, api_key: str, settings_service) -> User | None: """Validate API key against the database.""" - query: SelectOfScalar = select(ApiKey).options(selectinload(ApiKey.user)).where(ApiKey.api_key == api_key) - api_key_object: ApiKey | None = (await session.exec(query)).first() - if api_key_object is not None: - if settings_service.settings.disable_track_apikey_usage is not True: - api_key_object.total_uses += 1 - api_key_object.last_used_at = datetime.datetime.now(datetime.timezone.utc) - session.add(api_key_object) - await session.flush() - return api_key_object.user + query = select(ApiKey.id, ApiKey.api_key, ApiKey.user_id) + rows = (await session.exec(query)).all() # list of tuples (id, api_key, user_id) + + if not rows: + return None + + fernet = auth_utils.get_fernet(settings_service) + + for api_key_id, stored_value, user_id in rows: + if stored_value is None: + continue + + if stored_value == api_key: + matched = True + else: + try: + candidate = auth_utils.decrypt_api_key( + stored_value, settings_service=settings_service, fernet_obj=fernet + ) + except (ValueError, TypeError, InvalidToken): + candidate = stored_value + matched = candidate == api_key + + if matched: + if settings_service.settings.disable_track_apikey_usage is not True: + await session.exec( + update(ApiKey) + .where(ApiKey.id == api_key_id) + .values( + total_uses=ApiKey.total_uses + 1, + last_used_at=datetime.datetime.now(datetime.timezone.utc), + ) + ) + return await session.get(User, user_id) + return None diff --git a/src/backend/tests/performance/check_key_benchmark.py b/src/backend/tests/performance/check_key_benchmark.py new file mode 100644 index 0000000000..c7da0057b3 --- /dev/null +++ b/src/backend/tests/performance/check_key_benchmark.py @@ -0,0 +1,111 @@ +# src/backend/tests/perf/check_key_benchmark.py +import logging +import statistics +import time +import uuid + +import pytest +from langflow.services.auth import utils as auth_utils +from langflow.services.database.models.api_key import crud as api_key_crud +from langflow.services.database.models.api_key.model import ApiKey +from langflow.services.database.models.user.model import User +from langflow.services.deps import get_settings_service +from sqlmodel.ext.asyncio.session import AsyncSession + +logger = logging.getLogger(__name__) + + +def _get_test_password() -> str: + """Generate a unique test password for benchmark runs.""" + return str(uuid.uuid4()) + + +class DummyResult: + def __init__(self, items): + self._items = items + + def all(self): + return self._items + + +async def benchmark_once( + n_keys: int, + iterations: int = 100, + async_db_session: AsyncSession | None = None, +): + settings_service = None + try: + settings_service = get_settings_service() + except Exception: + settings_service = None + + stored_rows = [] + # generate N keys, keep one matching candidate_key to test hit + candidate_raw = f"sk-test-{uuid.uuid4()}" + + for i in range(n_keys): + raw = f"sk-test-{uuid.uuid4()}" + if i == n_keys - 1: + raw = candidate_raw + try: + stored = auth_utils.encrypt_api_key(raw, settings_service=settings_service) + except Exception: + stored = f"enc-{raw}" + stored_rows.append((str(i), stored, str(uuid.uuid4()))) + + if async_db_session is not None: + # use provided async session fixture to mimic DB + db_session = async_db_session + # create a user + user = User(username=f"u-{uuid.uuid4()}", password=_get_test_password()) + db_session.add(user) + await db_session.flush() + await db_session.refresh(user) + + for i, (_, stored, _uid) in enumerate(stored_rows): + api = ApiKey(api_key=stored, name=f"k-{i}", user_id=user.id) + db_session.add(api) + await db_session.commit() + + timings = [] + for _ in range(iterations): + t0 = time.perf_counter() + await api_key_crud._check_key_from_db(db_session, candidate_raw, settings_service) + t1 = time.perf_counter() + timings.append((t1 - t0) * 1000.0) # ms + + mean = statistics.mean(timings) + p50 = statistics.median(timings) + total_ms = sum(timings) + return { + "n_keys": n_keys, + "iterations": iterations, + "mean_ms": mean, + "p50_ms": p50, + "total_ms": total_ms, + } + + +@pytest.mark.parametrize("n_keys", [1, 10, 50, 100, 1000]) +async def test_benchmark_check_key_from_db_smoke(async_session: AsyncSession, n_keys): + """Run a quick smoke benchmark using simulated stored values (no real crypto). + + This test doesn't assert strict performance thresholds — it ensures the + benchmark runner works under pytest and returns sensible metrics. + """ + # keep iterations small for CI-friendly run time - use async session fixture + r = await benchmark_once(n_keys=n_keys, iterations=5, async_db_session=async_session) + + # basic sanity checks + assert r["n_keys"] == n_keys + assert r["mean_ms"] >= 0.0 + assert r["p50_ms"] >= 0.0 + + # log results so they are captured by pytest's logging capture + logger.info( + "perf n=%s mean=%.2fms p50=%.2fms total=%.2fms", + n_keys, + r["mean_ms"], + r["p50_ms"], + r["total_ms"], + ) diff --git a/src/backend/tests/unit/test_api_key_source.py b/src/backend/tests/unit/test_api_key_source.py index 11700cf2f1..715c91c02f 100644 --- a/src/backend/tests/unit/test_api_key_source.py +++ b/src/backend/tests/unit/test_api_key_source.py @@ -199,25 +199,27 @@ class TestCheckKeyFromDb: @pytest.mark.asyncio async def test_valid_key_returns_user(self, mock_session, mock_user, mock_settings_service_db): """Valid API key should return the associated user.""" - mock_api_key = MagicMock() - mock_api_key.user = mock_user - mock_api_key.total_uses = 0 + api_key_id = uuid4() + user_id = mock_user.id mock_result = MagicMock() - mock_result.first.return_value = mock_api_key - mock_session.exec.return_value = mock_result + mock_result.all.return_value = [(api_key_id, "sk-valid-key", user_id)] + + mock_session.exec = AsyncMock(return_value=mock_result) + + mock_session.get = AsyncMock(return_value=mock_user) result = await _check_key_from_db(mock_session, "sk-valid-key", mock_settings_service_db) assert result == mock_user - assert mock_api_key.total_uses == 1 + mock_session.get.assert_called_once_with(User, user_id) @pytest.mark.asyncio async def test_invalid_key_returns_none(self, mock_session, mock_settings_service_db): """Invalid API key should return None.""" mock_result = MagicMock() - mock_result.first.return_value = None - mock_session.exec.return_value = mock_result + mock_result.all.return_value = [] # No keys in DB + mock_session.exec = AsyncMock(return_value=mock_result) result = await _check_key_from_db(mock_session, "sk-invalid-key", mock_settings_service_db) @@ -226,44 +228,43 @@ class TestCheckKeyFromDb: @pytest.mark.asyncio async def test_usage_tracking_increments(self, mock_session, mock_user, mock_settings_service_db): """API key usage should be tracked when not disabled.""" - mock_api_key = MagicMock() - mock_api_key.user = mock_user - mock_api_key.total_uses = 5 + api_key_id = uuid4() + user_id = mock_user.id mock_result = MagicMock() - mock_result.first.return_value = mock_api_key - mock_session.exec.return_value = mock_result + mock_result.all.return_value = [(api_key_id, "sk-valid-key", user_id)] + mock_session.exec = AsyncMock(return_value=mock_result) + mock_session.get = AsyncMock(return_value=mock_user) await _check_key_from_db(mock_session, "sk-valid-key", mock_settings_service_db) - assert mock_api_key.total_uses == 6 - mock_session.add.assert_called_once_with(mock_api_key) - mock_session.flush.assert_called_once() + # Verify exec was called twice (select + update) + assert mock_session.exec.call_count == 2 @pytest.mark.asyncio async def test_usage_tracking_disabled(self, mock_session, mock_user, mock_settings_service_db): """API key usage should not be tracked when disabled.""" mock_settings_service_db.settings.disable_track_apikey_usage = True - mock_api_key = MagicMock() - mock_api_key.user = mock_user - mock_api_key.total_uses = 5 + api_key_id = uuid4() + user_id = mock_user.id mock_result = MagicMock() - mock_result.first.return_value = mock_api_key - mock_session.exec.return_value = mock_result + mock_result.all.return_value = [(api_key_id, "sk-valid-key", user_id)] + mock_session.exec = AsyncMock(return_value=mock_result) + mock_session.get = AsyncMock(return_value=mock_user) await _check_key_from_db(mock_session, "sk-valid-key", mock_settings_service_db) - assert mock_api_key.total_uses == 5 # Not incremented - mock_session.add.assert_not_called() + # Verify exec was called only once (select, no update) + assert mock_session.exec.call_count == 1 @pytest.mark.asyncio async def test_empty_key_returns_none(self, mock_session, mock_settings_service_db): """Empty API key should return None.""" mock_result = MagicMock() - mock_result.first.return_value = None - mock_session.exec.return_value = mock_result + mock_result.all.return_value = [] # No keys match + mock_session.exec = AsyncMock(return_value=mock_result) result = await _check_key_from_db(mock_session, "", mock_settings_service_db) @@ -479,13 +480,13 @@ class TestCheckKeyIntegration: @pytest.mark.asyncio async def test_full_flow_db_mode_valid_key(self, mock_session, mock_user): """Full flow test: db mode with valid key.""" - mock_api_key = MagicMock() - mock_api_key.user = mock_user - mock_api_key.total_uses = 0 + api_key_id = uuid4() + user_id = mock_user.id mock_result = MagicMock() - mock_result.first.return_value = mock_api_key - mock_session.exec.return_value = mock_result + mock_result.all.return_value = [(api_key_id, "sk-valid-key", user_id)] + mock_session.exec = AsyncMock(return_value=mock_result) + mock_session.get = AsyncMock(return_value=mock_user) mock_settings = MagicMock() mock_settings.auth_settings.API_KEY_SOURCE = "db" @@ -498,6 +499,7 @@ class TestCheckKeyIntegration: result = await check_key(mock_session, "sk-valid-key") assert result == mock_user + mock_session.get.assert_called_once_with(User, user_id) @pytest.mark.asyncio async def test_full_flow_env_mode_valid_key(self, mock_session, mock_superuser, monkeypatch): @@ -530,13 +532,18 @@ class TestCheckKeyIntegration: monkeypatch.setenv("LANGFLOW_API_KEY", "sk-correct-key") # Setup mock for db fallback - mock_api_key = MagicMock() - mock_api_key.user = mock_user - mock_api_key.total_uses = 0 + api_key_id = uuid4() + user_id = mock_user.id + + monkeypatch.setattr( + "langflow.services.database.models.api_key.crud.auth_utils.decrypt_api_key", + lambda v, _settings_service=None: "sk-wrong-key" if v == "sk-wrong-key" else v, + ) mock_result = MagicMock() - mock_result.first.return_value = mock_api_key - mock_session.exec.return_value = mock_result + mock_result.all.return_value = [(api_key_id, "sk-wrong-key", user_id)] + mock_session.exec = AsyncMock(return_value=mock_result) + mock_session.get = AsyncMock(return_value=mock_user) mock_settings = MagicMock() mock_settings.auth_settings.API_KEY_SOURCE = "env" @@ -560,8 +567,8 @@ class TestCheckKeyIntegration: # Setup mock for db - key not found mock_result = MagicMock() - mock_result.first.return_value = None - mock_session.exec.return_value = mock_result + mock_result.all.return_value = [] + mock_session.exec = AsyncMock(return_value=mock_result) mock_settings = MagicMock() mock_settings.auth_settings.API_KEY_SOURCE = "env" diff --git a/src/backend/tests/unit/test_get_api_key.py b/src/backend/tests/unit/test_get_api_key.py new file mode 100644 index 0000000000..f5ceeaa4ac --- /dev/null +++ b/src/backend/tests/unit/test_get_api_key.py @@ -0,0 +1,78 @@ +import asyncio +from uuid import uuid4 + +import langflow.services.database.models.api_key.crud as crud_module +import pytest +from cryptography.fernet import InvalidToken + + +class DummyResult: + def __init__(self, items): + self._items = items + + def all(self): + return self._items + + +class MockSession: + def __init__(self, items): + self._items = items + + async def exec(self, _query=None): + # emulate SQLModel AsyncSession.exec returning a result with .all() + await asyncio.sleep(0) # ensure it's truly async + return DummyResult(self._items) + + +class MockApiKeyObj: + def __init__(self, data: dict): + self._data = data + + def model_dump(self): + return dict(self._data) + + +@pytest.mark.asyncio +async def test_get_api_keys_decrypts_and_falls_back(monkeypatch): + user_id = uuid4() + + items = [ + MockApiKeyObj({"id": "1", "api_key": "enc-1", "name": "k1", "user_id": str(user_id)}), + MockApiKeyObj({"id": "2", "api_key": "bad-enc", "name": "k2", "user_id": str(user_id)}), + MockApiKeyObj({"id": "3", "api_key": None, "name": "k3", "user_id": str(user_id)}), + ] + + session = MockSession(items) + + # Ensure get_settings_service returns a dummy settings (decrypt stub ignores it, but function expects it) + monkeypatch.setattr(crud_module, "get_settings_service", lambda: object()) + + monkeypatch.setattr(crud_module.auth_utils, "get_fernet", lambda _settings_service: None) + + # Patch decrypt_api_key to: + # - return 'sk-decrypted' for 'enc-1' + # - raise InvalidToken for 'bad-enc' to trigger fallback + def fake_decrypt(val, *, settings_service=None, fernet_obj=None): # noqa: ARG001 + if val == "enc-1": + return "sk-decrypted" + if val == "bad-enc": + raise InvalidToken + return val + + monkeypatch.setattr(crud_module.auth_utils, "decrypt_api_key", fake_decrypt) + + # Patch ApiKeyRead.model_validate to just return the provided dict for easy assertions + monkeypatch.setattr(crud_module.ApiKeyRead, "model_validate", staticmethod(lambda data: data)) + + result = await crud_module.get_api_keys(session, user_id) + + # three entries returned + assert isinstance(result, list) + assert len(result) == 3 + + # first decrypted + assert result[0]["api_key"] == "sk-decrypted" + # second fell back to stored value 'bad-enc' + assert result[1]["api_key"] == "bad-enc" + # third remains None + assert result[2]["api_key"] is None