Files

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)