mirror of
https://github.com/langgenius/dify.git
synced 2026-07-21 18:58:35 +08:00
refactor: use Pydantic for sensitive word avoidance config (Fixes #37… (#37660)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com> Co-authored-by: Asuka Minato <i@asukaminato.eu.org>
This commit is contained in:
parent
3390be3978
commit
005bc54c38
@ -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"]
|
||||
|
||||
@ -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"},
|
||||
}
|
||||
}
|
||||
|
||||
Loading…
Reference in New Issue
Block a user