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

1"""Global configuration models for storage, compute, and error handling. 

2 

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""" 

9 

10from typing import Any, Literal 

11 

12from pydantic import BaseModel, ConfigDict, Field, model_validator 

13 

14 

15class RetryConfig(BaseModel): 

16 """Retry policy for storage operations.""" 

17 

18 model_config = ConfigDict(extra="forbid") 

19 

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.") 

25 

26 

27class StorageConfig(BaseModel): 

28 """Remote or local storage configuration.""" 

29 

30 model_config = ConfigDict(extra="forbid") 

31 

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').") 

34 

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.") 

43 

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.") 

54 

55 retry: RetryConfig | None = None 

56 

57 checksum_validation: Literal["when_required", "always", "never"] = Field( 

58 default="when_required", 

59 description="Controls S3 checksum validation behaviour.", 

60 ) 

61 

62 @model_validator(mode="after") 

63 def _require_bucket_for_s3(self) -> "StorageConfig": 

64 """Validate that S3 storage configuration includes a bucket name. 

65 

66 Returns: 

67 The validated StorageConfig instance. 

68 

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 

75 

76 

77class ComputeConfig(BaseModel): 

78 """Global compute / runtime settings.""" 

79 

80 model_config = ConfigDict(extra="forbid") 

81 

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.") 

94 

95 

96class ImageErrorsConfig(BaseModel): 

97 """Error-handling policy for image-processing failures.""" 

98 

99 model_config = ConfigDict(extra="forbid") 

100 

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 ) 

113 

114 

115class TabularErrorsConfig(BaseModel): 

116 """Error-handling policy for tabular-data failures.""" 

117 

118 model_config = ConfigDict(extra="forbid") 

119 

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 ) 

128 

129 

130class ErrorsConfig(BaseModel): 

131 """Aggregate error-handling configuration.""" 

132 

133 model_config = ConfigDict(extra="forbid") 

134 

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 )