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 4e86d04df0.

* 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 <olayinkaadelakun@Olayinkas-MacBook-Pro.local>
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
Co-authored-by: Olayinka Adelakun <olayinkaadelakun@mac.war.can.ibm.com>
Co-authored-by: Himavarsha <40851462+HimavarshaVS@users.noreply.github.com>
This commit is contained in:
olayinkaadelakun
2026-01-28 12:41:35 -05:00
committed by GitHub
parent 554bfc7651
commit b89bd76e22
6 changed files with 308 additions and 56 deletions

View File

@ -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

View File

@ -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()

View File

@ -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

View File

@ -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"],
)

View File

@ -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"

View File

@ -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