diff --git a/api/services/configuration/masking.py b/api/services/configuration/masking.py index a7e1af6a..bbcd983a 100644 --- a/api/services/configuration/masking.py +++ b/api/services/configuration/masking.py @@ -28,7 +28,19 @@ def contains_masked_key(value: str | list[str] | None) -> bool: if value is None: return False keys = value if isinstance(value, list) else [value] - return any(MASK_MARKER in k for k in keys) + return any(MASK_MARKER in k or _matches_mask_shape(k) for k in keys) + + +def _matches_mask_shape(key: str) -> bool: + """Return True when *key* has the exact shape produced by mask_key().""" + if not key: + return False + + if len(key) <= VISIBLE_CHARS: + return key == MASK_CHAR * len(key) + + masked_prefix_length = len(key) - VISIBLE_CHARS + return key[:masked_prefix_length] == MASK_CHAR * masked_prefix_length def check_for_masked_keys(config: "EffectiveAIModelConfiguration") -> None: diff --git a/api/tests/test_masked_key_rejection.py b/api/tests/test_masked_key_rejection.py index 437d1cc9..b1c8fe5a 100644 --- a/api/tests/test_masked_key_rejection.py +++ b/api/tests/test_masked_key_rejection.py @@ -7,7 +7,7 @@ from fastapi.testclient import TestClient from api.routes.user import router from api.schemas.ai_model_configuration import EffectiveAIModelConfiguration from api.services.auth.depends import get_user -from api.services.configuration.masking import mask_key +from api.services.configuration.masking import contains_masked_key, mask_key from api.services.configuration.registry import ( GoogleLLMService, GoogleVertexLLMConfiguration, @@ -43,6 +43,26 @@ def _existing_openai_config(): class TestMaskedKeyRejection: + def test_detects_short_masked_api_keys(self): + """Short masked API keys should be rejected even without the legacy marker.""" + assert contains_masked_key(mask_key("EMPTY")) + assert contains_masked_key(mask_key("X")) + assert contains_masked_key(mask_key("mykey")) + + def test_detects_masked_api_key_with_mask_char_in_visible_suffix(self): + """The visible suffix can contain MASK_CHAR and still be masked.""" + masked = mask_key("vsk*v") + + assert masked == "*sk*v" + assert contains_masked_key(masked) + + def test_allows_unmasked_short_api_keys(self): + """Unmasked short keys should not be rejected as masked placeholders.""" + assert not contains_masked_key("EMPTY") + assert not contains_masked_key("X") + assert not contains_masked_key("mykey") + assert not contains_masked_key("vsk*v") + def test_rejects_masked_api_key_on_provider_change(self): """Changing provider with a masked API key should return 400.""" app = _make_test_app()