Coverage for packages/dqm-ml-core/src/dqm_ml_core/models/global_.py: 91%
55 statements
« prev ^ index » next coverage.py v7.14.1, created at 2026-07-21 08:27 +0000
« prev ^ index » next coverage.py v7.14.1, created at 2026-07-21 08:27 +0000
1"""Global configuration models for storage, compute, and error handling.
3Defines shared configuration classes used across pipeline interfaces:
4- RetryConfig: Retry policies for storage operations
5- StorageConfig: S3/local storage with credential management
6- ComputeConfig: Runtime settings (device, memory, threads, logging)
7- ImageErrorsConfig/TabularErrorsConfig/ErrorsConfig: Error handling policies
8"""
10from typing import Any, Literal
12from pydantic import BaseModel, ConfigDict, Field, model_validator
15class RetryConfig(BaseModel):
16 """Retry policy for storage operations."""
18 model_config = ConfigDict(extra="forbid")
20 mode: Literal["default", "standard"] = Field(
21 default="standard",
22 description="Retry mode: 'default' uses exponential backoff, 'standard' uses fixed intervals.",
23 )
24 max_attempts: int = Field(default=3, gt=0, description="Maximum number of retry attempts.")
27class StorageConfig(BaseModel):
28 """Remote or local storage configuration."""
30 model_config = ConfigDict(extra="forbid")
32 type: Literal["s3", "local"] = Field(description="Storage backend type.")
33 bucket: str | None = Field(default=None, description="S3 bucket name (required when type='s3').")
35 access_key: str | None = Field(default=None, description="AWS access key ID.")
36 secret_key: str | None = Field(default=None, description="AWS secret access key.")
37 session_token: str | None = Field(default=None, description="AWS session token.")
38 anonymous: bool = Field(default=False, description="Use anonymous (unsigned) requests.")
39 role_arn: str | None = Field(default=None, description="ARN of IAM role to assume.")
40 session_name: str | None = Field(default=None, description="Name for the assumed role session.")
41 external_id: str | None = Field(default=None, description="External ID for role assumption.")
42 load_frequency: int = Field(default=900, gt=0, description="Frequency (seconds) to refresh credentials / role.")
44 region: str | None = Field(default=None, description="AWS region (e.g. 'us-east-1').")
45 endpoint: str | None = Field(default=None, description="Custom S3 endpoint URL.")
46 request_timeout: float | None = Field(default=None, description="Request timeout in seconds.")
47 connect_timeout: float | None = Field(default=None, description="Connection timeout in seconds.")
48 scheme: str | None = Field(default=None, description="URI scheme (e.g. 'https').")
49 proxy_options: dict[str, Any] | str | None = Field(
50 default=None,
51 description="Proxy configuration as a dict or URL string.",
52 )
53 tls_ca_file_path: str | None = Field(default=None, description="Path to a custom TLS CA bundle.")
55 retry: RetryConfig | None = None
57 checksum_validation: Literal["when_required", "always", "never"] = Field(
58 default="when_required",
59 description="Controls S3 checksum validation behaviour.",
60 )
62 @model_validator(mode="after")
63 def _require_bucket_for_s3(self) -> "StorageConfig":
64 """Validate that S3 storage configuration includes a bucket name.
66 Returns:
67 The validated StorageConfig instance.
69 Raises:
70 ValueError: If storage type is 's3' but no bucket is provided.
71 """
72 if self.type == "s3" and self.bucket is None:
73 raise ValueError("StorageConfig with type='s3' requires a 'bucket'")
74 return self
77class ComputeConfig(BaseModel):
78 """Global compute / runtime settings."""
80 model_config = ConfigDict(extra="forbid")
82 seed: int = Field(default=42, description="Random seed for reproducibility.")
83 log_level: Literal["debug", "info", "warning", "error"] = Field(
84 default="warning",
85 description="Logging verbosity.",
86 )
87 max_memory: str | None = Field(default=None, description="Maximum memory per worker (e.g. '4Gi').")
88 device: Literal["auto", "cpu", "cuda"] = Field(
89 default="auto",
90 description="Compute device: 'auto' picks cuda if available.",
91 )
92 progress_bar: bool = Field(default=True, description="Show tqdm progress bars.")
93 threads: int = Field(default=4, gt=0, description="Number of worker threads.")
96class ImageErrorsConfig(BaseModel):
97 """Error-handling policy for image-processing failures."""
99 model_config = ConfigDict(extra="forbid")
101 on_decode_failure: Literal["silent_fail", "fail_fast"] = Field(
102 default="silent_fail",
103 description="Action when an image cannot be decoded.",
104 )
105 on_transform_error: Literal["silent_fail", "fail_fast"] = Field(
106 default="silent_fail",
107 description="Action when an image transform fails.",
108 )
109 on_unsupported_format: Literal["silent_fail", "fail_fast"] = Field(
110 default="fail_fast",
111 description="Action on unsupported image format.",
112 )
115class TabularErrorsConfig(BaseModel):
116 """Error-handling policy for tabular-data failures."""
118 model_config = ConfigDict(extra="forbid")
120 on_missing_column: Literal["silent_fail", "fail_fast"] = Field(
121 default="fail_fast",
122 description="Action when a required column is missing.",
123 )
124 on_file_not_found: Literal["silent_fail", "fail_fast"] = Field(
125 default="fail_fast",
126 description="Action when a data file is not found.",
127 )
130class ErrorsConfig(BaseModel):
131 """Aggregate error-handling configuration."""
133 model_config = ConfigDict(extra="forbid")
135 default: Literal["silent_fail", "fail_fast"] = Field(
136 default="silent_fail",
137 description="Default error action when no specific policy is set.",
138 )
139 images: ImageErrorsConfig | None = None
140 tabular: TabularErrorsConfig | None = None
141 max_failure_rate: float = Field(
142 default=0.05,
143 ge=0,
144 le=1,
145 description="Maximum tolerated failure rate before the job aborts (from 0.0 to 1.0).",
146 )