diff --git a/api/core/app/app_config/common/sensitive_word_avoidance/manager.py b/api/core/app/app_config/common/sensitive_word_avoidance/manager.py index c8ec7cb44dc..b02b4077338 100644 --- a/api/core/app/app_config/common/sensitive_word_avoidance/manager.py +++ b/api/core/app/app_config/common/sensitive_word_avoidance/manager.py @@ -1,10 +1,73 @@ from collections.abc import Mapping -from typing import Any +from typing import Annotated, Any, Literal + +from pydantic import BaseModel, ConfigDict, Field, TypeAdapter, ValidationError from core.app.app_config.entities import SensitiveWordAvoidanceEntity from core.moderation.factory import ModerationFactory +class SensitiveWordAvoidanceDisabledConfig(BaseModel): + model_config = ConfigDict(extra="forbid") + + enabled: Literal[False] = False + + +class SensitiveWordAvoidanceKeywordsConfig(BaseModel): + model_config = ConfigDict(extra="forbid") + + enabled: Literal[True] = True + type: Literal["keywords"] + config: dict[str, Any] = Field(default_factory=dict) + + def run_provider_validation(self, tenant_id: str) -> None: + ModerationFactory.validate_config(name=self.type, tenant_id=tenant_id, config=self.config) + + +class SensitiveWordAvoidanceOpenAIConfig(BaseModel): + model_config = ConfigDict(extra="forbid") + + enabled: Literal[True] = True + type: Literal["openai_moderation"] + config: dict[str, Any] = Field(default_factory=dict) + + def run_provider_validation(self, tenant_id: str) -> None: + ModerationFactory.validate_config(name=self.type, tenant_id=tenant_id, config=self.config) + + +class SensitiveWordAvoidanceAPIConfig(BaseModel): + model_config = ConfigDict(extra="forbid") + + enabled: Literal[True] = True + type: Literal["api"] + config: dict[str, Any] = Field(default_factory=dict) + + def run_provider_validation(self, tenant_id: str) -> None: + ModerationFactory.validate_config(name=self.type, tenant_id=tenant_id, config=self.config) + + +EnabledSensitiveWordAvoidanceConfig = Annotated[ + SensitiveWordAvoidanceKeywordsConfig | SensitiveWordAvoidanceOpenAIConfig | SensitiveWordAvoidanceAPIConfig, + Field(discriminator="type"), +] + +SensitiveWordAvoidanceConfig = Annotated[ + SensitiveWordAvoidanceDisabledConfig | EnabledSensitiveWordAvoidanceConfig, + Field(discriminator="enabled"), +] + +_sensitive_word_avoidance_adapter: TypeAdapter[SensitiveWordAvoidanceConfig] = TypeAdapter(SensitiveWordAvoidanceConfig) + + +def _normalize_raw(raw: Any) -> Any: + if isinstance(raw, dict): + if raw.get("enabled") is None: + raw = {**raw, "enabled": False} + elif raw.get("enabled") is True and raw.get("config") is None: + raw = {**raw, "config": {}} + return raw + + class SensitiveWordAvoidanceConfigManager: @classmethod def convert(cls, config: Mapping[str, Any]) -> SensitiveWordAvoidanceEntity | None: @@ -24,30 +87,24 @@ class SensitiveWordAvoidanceConfigManager: def validate_and_set_defaults( cls, tenant_id: str, config: dict[str, Any], only_structure_validate: bool = False ) -> tuple[dict[str, Any], list[str]]: - if not config.get("sensitive_word_avoidance"): - config["sensitive_word_avoidance"] = {"enabled": False} - - if not isinstance(config["sensitive_word_avoidance"], dict): + raw = config.get("sensitive_word_avoidance") or {"enabled": False} + if not isinstance(raw, dict): raise ValueError("sensitive_word_avoidance must be of dict type") - if "enabled" not in config["sensitive_word_avoidance"] or not config["sensitive_word_avoidance"]["enabled"]: - config["sensitive_word_avoidance"]["enabled"] = False + try: + validated = _sensitive_word_avoidance_adapter.validate_python(_normalize_raw(raw)) + except ValidationError: + raise - if config["sensitive_word_avoidance"]["enabled"]: - if not config["sensitive_word_avoidance"].get("type"): - raise ValueError("sensitive_word_avoidance.type is required") - - if not only_structure_validate: - typ = config["sensitive_word_avoidance"]["type"] - if not isinstance(typ, str): - raise ValueError("sensitive_word_avoidance.type must be a string") - - sensitive_word_avoidance_config = config["sensitive_word_avoidance"].get("config") - if sensitive_word_avoidance_config is None: - sensitive_word_avoidance_config = {} - if not isinstance(sensitive_word_avoidance_config, dict): - raise ValueError("sensitive_word_avoidance.config must be a dict") - - ModerationFactory.validate_config(name=typ, tenant_id=tenant_id, config=sensitive_word_avoidance_config) + if not only_structure_validate and isinstance( + validated, + ( + SensitiveWordAvoidanceKeywordsConfig, + SensitiveWordAvoidanceOpenAIConfig, + SensitiveWordAvoidanceAPIConfig, + ), + ): + validated.run_provider_validation(tenant_id) + config["sensitive_word_avoidance"] = validated.model_dump() return config, ["sensitive_word_avoidance"] diff --git a/api/tests/unit_tests/core/app/app_config/common/test_sensitive_word_avoidance_manager.py b/api/tests/unit_tests/core/app/app_config/common/test_sensitive_word_avoidance_manager.py index bd4ca5ff85a..0c8073793b2 100644 --- a/api/tests/unit_tests/core/app/app_config/common/test_sensitive_word_avoidance_manager.py +++ b/api/tests/unit_tests/core/app/app_config/common/test_sensitive_word_avoidance_manager.py @@ -1,6 +1,7 @@ from unittest.mock import MagicMock import pytest +from pydantic import ValidationError from pytest_mock import MockerFixture from core.app.app_config.common.sensitive_word_avoidance.manager import ( @@ -110,7 +111,7 @@ class TestSensitiveWordAvoidanceConfigManagerValidateAndSetDefaults: def test_validate_raises_when_enabled_true_without_type(self): config = {"sensitive_word_avoidance": {"enabled": True}} - with pytest.raises(ValueError, match="type is required"): + with pytest.raises(ValidationError, match="discriminator 'type'"): SensitiveWordAvoidanceConfigManager.validate_and_set_defaults(tenant_id="tenant1", config=config) def test_validate_raises_when_type_not_string(self): @@ -121,19 +122,30 @@ class TestSensitiveWordAvoidanceConfigManagerValidateAndSetDefaults: } } - with pytest.raises(ValueError, match="must be a string"): + with pytest.raises(ValidationError, match="does not match any of the expected tags"): + SensitiveWordAvoidanceConfigManager.validate_and_set_defaults(tenant_id="tenant1", config=config) + + def test_validate_raises_when_type_invalid(self): + config = { + "sensitive_word_avoidance": { + "enabled": True, + "type": "mock_type", + } + } + + with pytest.raises(ValidationError, match="does not match any of the expected tags"): SensitiveWordAvoidanceConfigManager.validate_and_set_defaults(tenant_id="tenant1", config=config) def test_validate_raises_when_config_not_dict(self): config = { "sensitive_word_avoidance": { "enabled": True, - "type": "mock_type", + "type": "keywords", "config": "invalid", } } - with pytest.raises(ValueError, match="must be a dict"): + with pytest.raises(ValidationError, match="valid dictionary"): SensitiveWordAvoidanceConfigManager.validate_and_set_defaults(tenant_id="tenant1", config=config) def test_validate_calls_moderation_factory(self, mocker: MockerFixture): @@ -145,7 +157,7 @@ class TestSensitiveWordAvoidanceConfigManagerValidateAndSetDefaults: config = { "sensitive_word_avoidance": { "enabled": True, - "type": "mock_type", + "type": "keywords", "config": {"k": "v"}, } } @@ -156,7 +168,7 @@ class TestSensitiveWordAvoidanceConfigManagerValidateAndSetDefaults: ) # Assert - mock_validate.assert_called_once_with(name="mock_type", tenant_id="tenant1", config={"k": "v"}) + mock_validate.assert_called_once_with(name="keywords", tenant_id="tenant1", config={"k": "v"}) assert result_config["sensitive_word_avoidance"]["enabled"] is True assert fields == ["sensitive_word_avoidance"] @@ -169,7 +181,7 @@ class TestSensitiveWordAvoidanceConfigManagerValidateAndSetDefaults: config = { "sensitive_word_avoidance": { "enabled": True, - "type": "mock_type", + "type": "keywords", "config": None, } } @@ -178,7 +190,7 @@ class TestSensitiveWordAvoidanceConfigManagerValidateAndSetDefaults: SensitiveWordAvoidanceConfigManager.validate_and_set_defaults(tenant_id="tenant1", config=config) # Assert - mock_validate.assert_called_once_with(name="mock_type", tenant_id="tenant1", config={}) + mock_validate.assert_called_once_with(name="keywords", tenant_id="tenant1", config={}) def test_validate_only_structure_validate_skips_factory(self, mocker: MockerFixture): # Arrange @@ -189,7 +201,7 @@ class TestSensitiveWordAvoidanceConfigManagerValidateAndSetDefaults: config = { "sensitive_word_avoidance": { "enabled": True, - "type": "mock_type", + "type": "keywords", "config": {"k": "v"}, } }