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:
coslog 2026-07-01 14:46:48 +08:00 committed by GitHub
parent 3390be3978
commit 005bc54c38
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
2 changed files with 101 additions and 32 deletions

View File

@ -1,10 +1,73 @@
from collections.abc import Mapping 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.app.app_config.entities import SensitiveWordAvoidanceEntity
from core.moderation.factory import ModerationFactory 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: class SensitiveWordAvoidanceConfigManager:
@classmethod @classmethod
def convert(cls, config: Mapping[str, Any]) -> SensitiveWordAvoidanceEntity | None: def convert(cls, config: Mapping[str, Any]) -> SensitiveWordAvoidanceEntity | None:
@ -24,30 +87,24 @@ class SensitiveWordAvoidanceConfigManager:
def validate_and_set_defaults( def validate_and_set_defaults(
cls, tenant_id: str, config: dict[str, Any], only_structure_validate: bool = False cls, tenant_id: str, config: dict[str, Any], only_structure_validate: bool = False
) -> tuple[dict[str, Any], list[str]]: ) -> tuple[dict[str, Any], list[str]]:
if not config.get("sensitive_word_avoidance"): raw = config.get("sensitive_word_avoidance") or {"enabled": False}
config["sensitive_word_avoidance"] = {"enabled": False} if not isinstance(raw, dict):
if not isinstance(config["sensitive_word_avoidance"], dict):
raise ValueError("sensitive_word_avoidance must be of dict type") 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"]: try:
config["sensitive_word_avoidance"]["enabled"] = False validated = _sensitive_word_avoidance_adapter.validate_python(_normalize_raw(raw))
except ValidationError:
raise
if config["sensitive_word_avoidance"]["enabled"]: if not only_structure_validate and isinstance(
if not config["sensitive_word_avoidance"].get("type"): validated,
raise ValueError("sensitive_word_avoidance.type is required") (
SensitiveWordAvoidanceKeywordsConfig,
if not only_structure_validate: SensitiveWordAvoidanceOpenAIConfig,
typ = config["sensitive_word_avoidance"]["type"] SensitiveWordAvoidanceAPIConfig,
if not isinstance(typ, str): ),
raise ValueError("sensitive_word_avoidance.type must be a string") ):
validated.run_provider_validation(tenant_id)
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)
config["sensitive_word_avoidance"] = validated.model_dump()
return config, ["sensitive_word_avoidance"] return config, ["sensitive_word_avoidance"]

View File

@ -1,6 +1,7 @@
from unittest.mock import MagicMock from unittest.mock import MagicMock
import pytest import pytest
from pydantic import ValidationError
from pytest_mock import MockerFixture from pytest_mock import MockerFixture
from core.app.app_config.common.sensitive_word_avoidance.manager import ( 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): def test_validate_raises_when_enabled_true_without_type(self):
config = {"sensitive_word_avoidance": {"enabled": True}} 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) SensitiveWordAvoidanceConfigManager.validate_and_set_defaults(tenant_id="tenant1", config=config)
def test_validate_raises_when_type_not_string(self): 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) SensitiveWordAvoidanceConfigManager.validate_and_set_defaults(tenant_id="tenant1", config=config)
def test_validate_raises_when_config_not_dict(self): def test_validate_raises_when_config_not_dict(self):
config = { config = {
"sensitive_word_avoidance": { "sensitive_word_avoidance": {
"enabled": True, "enabled": True,
"type": "mock_type", "type": "keywords",
"config": "invalid", "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) SensitiveWordAvoidanceConfigManager.validate_and_set_defaults(tenant_id="tenant1", config=config)
def test_validate_calls_moderation_factory(self, mocker: MockerFixture): def test_validate_calls_moderation_factory(self, mocker: MockerFixture):
@ -145,7 +157,7 @@ class TestSensitiveWordAvoidanceConfigManagerValidateAndSetDefaults:
config = { config = {
"sensitive_word_avoidance": { "sensitive_word_avoidance": {
"enabled": True, "enabled": True,
"type": "mock_type", "type": "keywords",
"config": {"k": "v"}, "config": {"k": "v"},
} }
} }
@ -156,7 +168,7 @@ class TestSensitiveWordAvoidanceConfigManagerValidateAndSetDefaults:
) )
# Assert # 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 result_config["sensitive_word_avoidance"]["enabled"] is True
assert fields == ["sensitive_word_avoidance"] assert fields == ["sensitive_word_avoidance"]
@ -169,7 +181,7 @@ class TestSensitiveWordAvoidanceConfigManagerValidateAndSetDefaults:
config = { config = {
"sensitive_word_avoidance": { "sensitive_word_avoidance": {
"enabled": True, "enabled": True,
"type": "mock_type", "type": "keywords",
"config": None, "config": None,
} }
} }
@ -178,7 +190,7 @@ class TestSensitiveWordAvoidanceConfigManagerValidateAndSetDefaults:
SensitiveWordAvoidanceConfigManager.validate_and_set_defaults(tenant_id="tenant1", config=config) SensitiveWordAvoidanceConfigManager.validate_and_set_defaults(tenant_id="tenant1", config=config)
# Assert # 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): def test_validate_only_structure_validate_skips_factory(self, mocker: MockerFixture):
# Arrange # Arrange
@ -189,7 +201,7 @@ class TestSensitiveWordAvoidanceConfigManagerValidateAndSetDefaults:
config = { config = {
"sensitive_word_avoidance": { "sensitive_word_avoidance": {
"enabled": True, "enabled": True,
"type": "mock_type", "type": "keywords",
"config": {"k": "v"}, "config": {"k": "v"},
} }
} }