diff --git a/components/clp-py-utils/clp_py_utils/clp_config.py b/components/clp-py-utils/clp_py_utils/clp_config.py index 4d408dcc75..87e0edb82d 100644 --- a/components/clp-py-utils/clp_py_utils/clp_config.py +++ b/components/clp-py-utils/clp_py_utils/clp_config.py @@ -10,6 +10,7 @@ Field, field_validator, model_validator, + PlainSerializer, PrivateAttr, ) from strenum import KebabCaseStrEnum, LowercaseStrEnum @@ -21,6 +22,7 @@ read_yaml_config_file, validate_path_could_be_dir, ) +from .serialization_utils import serialize_path, serialize_str_enum # Constants # Component names @@ -98,6 +100,8 @@ CLP_QUEUE_PASS_ENV_VAR_NAME = "CLP_QUEUE_PASS" CLP_REDIS_PASS_ENV_VAR_NAME = "CLP_REDIS_PASS" +# Serializer +StrEnumSerializer = PlainSerializer(serialize_str_enum) # Generic types NonEmptyStr = Annotated[str, Field(min_length=1)] PositiveFloat = Annotated[float, Field(gt=0)] @@ -106,6 +110,7 @@ # TODO: Replace this with pydantic_extra_types.domain.DomainStr. DomainStr = NonEmptyStr Port = Annotated[int, Field(gt=0, lt=2**16)] +SerializablePath = Annotated[pathlib.Path, PlainSerializer(serialize_path)] ZstdCompressionLevel = Annotated[int, Field(ge=1, le=19)] @@ -114,17 +119,26 @@ class StorageEngine(KebabCaseStrEnum): CLP_S = auto() +StorageEngineStr = Annotated[StorageEngine, StrEnumSerializer] + + class DatabaseEngine(KebabCaseStrEnum): MARIADB = auto() MYSQL = auto() +DatabaseEngineStr = Annotated[DatabaseEngine, StrEnumSerializer] + + class QueryEngine(KebabCaseStrEnum): CLP = auto() CLP_S = auto() PRESTO = auto() +QueryEngineStr = Annotated[QueryEngine, StrEnumSerializer] + + class StorageType(LowercaseStrEnum): FS = auto() S3 = auto() @@ -137,9 +151,12 @@ class AwsAuthType(LowercaseStrEnum): ec2 = auto() +AwsAuthTypeStr = Annotated[AwsAuthType, StrEnumSerializer] + + class Package(BaseModel): - storage_engine: StorageEngine = StorageEngine.CLP - query_engine: QueryEngine = QueryEngine.CLP + storage_engine: StorageEngineStr = StorageEngine.CLP + query_engine: QueryEngineStr = QueryEngine.CLP @model_validator(mode="after") def validate_query_engine_package_compatibility(self): @@ -163,15 +180,9 @@ def validate_query_engine_package_compatibility(self): return self - def dump_to_primitive_dict(self): - d = self.model_dump() - d["storage_engine"] = d["storage_engine"].value - d["query_engine"] = d["query_engine"].value - return d - class Database(BaseModel): - type: DatabaseEngine = DatabaseEngine.MARIADB + type: DatabaseEngineStr = DatabaseEngine.MARIADB host: DomainStr = "localhost" port: Port = 3306 name: NonEmptyStr = "clp-db" @@ -232,7 +243,6 @@ def get_clp_connection_params_and_type(self, disable_localhost_socket_connection def dump_to_primitive_dict(self): d = self.model_dump(exclude={"username", "password"}) - d["type"] = d["type"].value return d def load_credentials_from_file(self, credentials_file_path: pathlib.Path): @@ -360,12 +370,7 @@ class S3Credentials(BaseModel): class AwsAuthentication(BaseModel): - type: Literal[ - AwsAuthType.credentials.value, - AwsAuthType.profile.value, - AwsAuthType.env_vars.value, - AwsAuthType.ec2.value, - ] + type: AwsAuthTypeStr profile: Optional[NonEmptyStr] = None credentials: Optional[S3Credentials] = None @@ -408,13 +413,10 @@ class S3IngestionConfig(BaseModel): type: Literal[StorageType.S3.value] = StorageType.S3.value aws_authentication: AwsAuthentication - def dump_to_primitive_dict(self): - return self.model_dump() - class FsStorage(BaseModel): type: Literal[StorageType.FS.value] = StorageType.FS.value - directory: pathlib.Path + directory: SerializablePath @field_validator("directory", mode="before") @classmethod @@ -425,16 +427,11 @@ def validate_directory(cls, value): def make_config_paths_absolute(self, clp_home: pathlib.Path): self.directory = make_config_path_absolute(clp_home, self.directory) - def dump_to_primitive_dict(self): - d = self.model_dump() - d["directory"] = str(d["directory"]) - return d - class S3Storage(BaseModel): type: Literal[StorageType.S3.value] = StorageType.S3.value s3_config: S3Config - staging_directory: pathlib.Path + staging_directory: SerializablePath @field_validator("staging_directory", mode="before") @classmethod @@ -455,30 +452,25 @@ def validate_key_prefix(cls, value): def make_config_paths_absolute(self, clp_home: pathlib.Path): self.staging_directory = make_config_path_absolute(clp_home, self.staging_directory) - def dump_to_primitive_dict(self): - d = self.model_dump() - d["staging_directory"] = str(d["staging_directory"]) - return d - class FsIngestionConfig(FsStorage): - directory: pathlib.Path = pathlib.Path("/") + directory: SerializablePath = pathlib.Path("/") class ArchiveFsStorage(FsStorage): - directory: pathlib.Path = CLP_DEFAULT_DATA_DIRECTORY_PATH / "archives" + directory: SerializablePath = CLP_DEFAULT_DATA_DIRECTORY_PATH / "archives" class StreamFsStorage(FsStorage): - directory: pathlib.Path = CLP_DEFAULT_DATA_DIRECTORY_PATH / "streams" + directory: SerializablePath = CLP_DEFAULT_DATA_DIRECTORY_PATH / "streams" class ArchiveS3Storage(S3Storage): - staging_directory: pathlib.Path = CLP_DEFAULT_DATA_DIRECTORY_PATH / "staged-archives" + staging_directory: SerializablePath = CLP_DEFAULT_DATA_DIRECTORY_PATH / "staged-archives" class StreamS3Storage(S3Storage): - staging_directory: pathlib.Path = CLP_DEFAULT_DATA_DIRECTORY_PATH / "staged-streams" + staging_directory: SerializablePath = CLP_DEFAULT_DATA_DIRECTORY_PATH / "staged-streams" def _get_directory_from_storage_config( @@ -520,11 +512,6 @@ def set_directory(self, directory: pathlib.Path): def get_directory(self) -> pathlib.Path: return _get_directory_from_storage_config(self.storage) - def dump_to_primitive_dict(self): - d = self.model_dump() - d["storage"] = self.storage.dump_to_primitive_dict() - return d - class StreamOutput(BaseModel): storage: Union[StreamFsStorage, StreamS3Storage] = StreamFsStorage() @@ -536,11 +523,6 @@ def set_directory(self, directory: pathlib.Path): def get_directory(self) -> pathlib.Path: return _get_directory_from_storage_config(self.storage) - def dump_to_primitive_dict(self): - d = self.model_dump() - d["storage"] = self.storage.dump_to_primitive_dict() - return d - class WebUi(BaseModel): host: DomainStr = "localhost" @@ -590,24 +572,26 @@ class CLPConfig(BaseModel): query_worker: QueryWorker = QueryWorker() webui: WebUi = WebUi() garbage_collector: GarbageCollector = GarbageCollector() - credentials_file_path: pathlib.Path = CLP_DEFAULT_CREDENTIALS_FILE_PATH + credentials_file_path: SerializablePath = CLP_DEFAULT_CREDENTIALS_FILE_PATH presto: Optional[Presto] = None archive_output: ArchiveOutput = ArchiveOutput() stream_output: StreamOutput = StreamOutput() - data_directory: pathlib.Path = pathlib.Path("var") / "data" - logs_directory: pathlib.Path = pathlib.Path("var") / "log" - aws_config_directory: Optional[pathlib.Path] = None + data_directory: SerializablePath = pathlib.Path("var") / "data" + logs_directory: SerializablePath = pathlib.Path("var") / "log" + aws_config_directory: Optional[SerializablePath] = None - _container_image_id_path: pathlib.Path = PrivateAttr( + _container_image_id_path: SerializablePath = PrivateAttr( default=CLP_PACKAGE_CONTAINER_IMAGE_ID_PATH ) - _version_file_path: pathlib.Path = PrivateAttr(default=CLP_VERSION_FILE_PATH) + _version_file_path: SerializablePath = PrivateAttr(default=CLP_VERSION_FILE_PATH) @field_validator("aws_config_directory") @classmethod - def expand_profile_user_home(cls, value: Optional[pathlib.Path]): + def expand_profile_user_home( + cls, value: Optional[SerializablePath] + ) -> Optional[SerializablePath]: if value is not None: value = value.expanduser() return value @@ -693,7 +677,7 @@ def validate_aws_config_dir(self): auth_configs.append(self.stream_output.storage.s3_config.aws_authentication) for auth in auth_configs: - if AwsAuthType.profile.value == auth.type: + if AwsAuthType.profile == auth.type: profile_auth_used = True break @@ -735,27 +719,14 @@ def get_runnable_components(self) -> Set[str]: def dump_to_primitive_dict(self): custom_serialized_fields = { - "package", "database", "queue", "redis", - "logs_input", - "archive_output", - "stream_output", } d = self.model_dump(exclude=custom_serialized_fields) for key in custom_serialized_fields: d[key] = getattr(self, key).dump_to_primitive_dict() - # Turn paths into primitive strings - d["credentials_file_path"] = str(self.credentials_file_path) - d["data_directory"] = str(self.data_directory) - d["logs_directory"] = str(self.logs_directory) - if self.aws_config_directory is not None: - d["aws_config_directory"] = str(self.aws_config_directory) - else: - d["aws_config_directory"] = None - return d @model_validator(mode="after") @@ -772,22 +743,12 @@ def validate_presto_config(self): class WorkerConfig(BaseModel): package: Package = Package() archive_output: ArchiveOutput = ArchiveOutput() - data_directory: pathlib.Path = CLPConfig().data_directory + data_directory: SerializablePath = CLPConfig().data_directory # Only needed by query workers. stream_output: StreamOutput = StreamOutput() stream_collection_name: str = ResultsCache().stream_collection_name - def dump_to_primitive_dict(self): - d = self.model_dump() - d["archive_output"] = self.archive_output.dump_to_primitive_dict() - - # Turn paths into primitive strings - d["data_directory"] = str(self.data_directory) - d["stream_output"] = self.stream_output.dump_to_primitive_dict() - - return d - def get_components_for_target(target: str) -> Set[str]: if target in TARGET_TO_COMPONENTS: diff --git a/components/clp-py-utils/clp_py_utils/serialization_utils.py b/components/clp-py-utils/clp_py_utils/serialization_utils.py new file mode 100644 index 0000000000..2ebd0e0e78 --- /dev/null +++ b/components/clp-py-utils/clp_py_utils/serialization_utils.py @@ -0,0 +1,23 @@ +import pathlib + +from strenum import StrEnum + + +def serialize_str_enum(member: StrEnum) -> str: + """ + Serializes a `strenum.StrEnum` member to its underlying value. + + :param member: + :return: The underlying string value of the enum member. + """ + return member.value + + +def serialize_path(path: pathlib.Path) -> str: + """ + Serializes a `pathlib.Path` to its string representation. + + :param path: + :return: The string representation of the path. + """ + return str(path)