91 lines
3.8 KiB
Python
91 lines
3.8 KiB
Python
from __future__ import annotations
|
|
|
|
import os
|
|
from pathlib import Path
|
|
|
|
from pydantic import Field, SecretStr, model_validator
|
|
from pydantic_settings import BaseSettings, SettingsConfigDict
|
|
|
|
|
|
def _secret(name: str, *, required: bool = True) -> str | None:
|
|
"""Read a secret only from NAME_FILE; values never enter repr/log output."""
|
|
file_name = os.getenv(f"{name}_FILE")
|
|
if not file_name:
|
|
if required:
|
|
raise ValueError(f"{name}_FILE is required")
|
|
return None
|
|
path = Path(file_name)
|
|
value = path.read_text(encoding="utf-8").rstrip("\r\n")
|
|
if not value or value.startswith("<"):
|
|
raise ValueError(f"{name}_FILE contains an invalid value")
|
|
return value
|
|
|
|
|
|
class BootstrapSettings(BaseSettings):
|
|
model_config = SettingsConfigDict(extra="ignore", populate_by_name=True)
|
|
|
|
app_env: str = Field(default="development", alias="APP_ENV")
|
|
process_role: str = Field(default="api", alias="MESSAGE_SAFETY_PROCESS_ROLE")
|
|
host: str = Field(default="0.0.0.0", alias="MESSAGE_SAFETY_HOST") # noqa: S104
|
|
port: int = Field(default=8080, alias="MESSAGE_SAFETY_PORT")
|
|
worker_concurrency: int = Field(
|
|
default=5, ge=1, le=32, alias="MESSAGE_SAFETY_WORKER_CONCURRENCY"
|
|
)
|
|
dns_resolvers: str = Field(default="", alias="MESSAGE_SAFETY_DNS_RESOLVERS")
|
|
antivirus_socket: Path = Field(
|
|
default=Path("/run/han-kesl/scan.sock"),
|
|
alias="MESSAGE_SAFETY_ANTIVIRUS_SOCKET",
|
|
)
|
|
s3_endpoint_url: str = Field(alias="SELECTEL_S3_ENDPOINT_URL")
|
|
s3_bucket: str = Field(alias="SELECTEL_S3_BUCKET_QUARANTINE")
|
|
artifacts_dir: Path = Field(
|
|
default=Path("/app/app/artifacts"), alias="MESSAGE_SAFETY_ARTIFACTS_DIR"
|
|
)
|
|
mode_file: Path = Field(
|
|
default=Path("/etc/han-chat/message-safety-mode.env"),
|
|
alias="MESSAGE_SAFETY_MODE_FILE",
|
|
)
|
|
database_url: SecretStr | None = None
|
|
redis_url: SecretStr | None = None
|
|
service_token: SecretStr | None = None
|
|
s3_access_key: SecretStr | None = None
|
|
s3_secret_key: SecretStr | None = None
|
|
|
|
@model_validator(mode="after")
|
|
def load_secret_files(self) -> BootstrapSettings:
|
|
self.database_url = SecretStr(_secret("MESSAGE_SAFETY_DATABASE_URL"))
|
|
self.redis_url = SecretStr(_secret("MESSAGE_SAFETY_REDIS_URL", required=False) or "")
|
|
if self.process_role == "api":
|
|
self.service_token = SecretStr(_secret("MESSAGE_SAFETY_SERVICE_TOKEN"))
|
|
elif self.process_role == "worker":
|
|
self.s3_access_key = SecretStr(_secret("SELECTEL_S3_QUARANTINE_READ_ACCESS_KEY"))
|
|
self.s3_secret_key = SecretStr(_secret("SELECTEL_S3_QUARANTINE_READ_SECRET_KEY"))
|
|
else:
|
|
raise ValueError("MESSAGE_SAFETY_PROCESS_ROLE must be api or worker")
|
|
return self
|
|
|
|
|
|
class EmergencyMode(BaseSettings):
|
|
model_config = SettingsConfigDict(extra="forbid", populate_by_name=True)
|
|
mock: bool = Field(default=False, alias="MESSAGE_SAFETY_MOCK_ENABLED")
|
|
text_free: bool = Field(default=False, alias="MESSAGE_SAFETY_MOCK_TEXT_FREE")
|
|
file_free: bool = Field(default=False, alias="MESSAGE_SAFETY_MOCK_FILE_FREE")
|
|
|
|
@model_validator(mode="after")
|
|
def valid_flags(self) -> EmergencyMode:
|
|
if not self.mock and (self.text_free or self.file_free):
|
|
raise ValueError("free flags require MOCK=true")
|
|
return self
|
|
|
|
@classmethod
|
|
def from_file(cls, path: Path) -> EmergencyMode:
|
|
values: dict[str, str] = {}
|
|
if path.exists():
|
|
for line in path.read_text(encoding="utf-8").splitlines():
|
|
if line and not line.startswith("#"):
|
|
key, sep, value = line.partition("=")
|
|
if not sep or key in values:
|
|
raise ValueError("invalid emergency mode file")
|
|
values[key] = value
|
|
return cls.model_validate(values)
|