diff --git a/radial_membrane_ai/exceptions.py b/radial_membrane_ai/exceptions.py new file mode 100644 index 0000000..e5f061c --- /dev/null +++ b/radial_membrane_ai/exceptions.py @@ -0,0 +1,40 @@ +""" +Custom Exceptions for the UFO Governed Deformable Radial Membrane framework. +""" + +from __future__ import annotations + + +class GovernanceError(ValueError): + """Base exception for all governance and verification errors in the UFO system.""" + pass + + +class GeometryValidationError(GovernanceError): + """Raised when radial membrane boundary or geometry structures violate invariants.""" + pass + + +class ValidationError(GovernanceError): + """Raised when generic dataclass validations fail.""" + pass + + +class WorkloadValidationError(ValidationError): + """Raised when workload step definitions or parameters violate schema or validation invariants.""" + pass + + +class WorkloadConfigurationError(GovernanceError): + """Raised when workload engine setup or pre-run configuration matches/validation fails.""" + pass + + +class InvalidSimulationTargetError(GovernanceError): + """Raised when the simulation engine target does not match the workload requirements.""" + pass + + +class ReconstructionError(GovernanceError): + """Raised when the L.D.E. round-trip text reconstruction fails or mismatches the original text.""" + pass diff --git a/radial_membrane_ai/lde/models.py b/radial_membrane_ai/lde/models.py index d11ad0b..a156ea1 100644 --- a/radial_membrane_ai/lde/models.py +++ b/radial_membrane_ai/lde/models.py @@ -6,6 +6,8 @@ from dataclasses import dataclass, field from typing import List, Dict, Tuple, Any +from radial_membrane_ai.exceptions import GeometryValidationError, ValidationError + @dataclass class LDEConfig: @@ -19,6 +21,32 @@ class LDEConfig: rho: str = "full-reconstruction" # 'full-reconstruction' or 'compressed' r_0: float = 1.0 epsilon: float = 1e-6 + geometry_samples: int = 100 + + def __post_init__(self) -> None: + self.validate() + + def validate(self) -> None: + """Enforces strict invariants on L.D.E. Configuration.""" + if not self.sigma: + raise ValidationError("LDEConfig.sigma cannot be empty.") + if len(self.sigma) != len(set(self.sigma)): + raise ValidationError("LDEConfig.sigma elements must be unique.") + weights = [ + ("w_f", self.w_f), ("w_s", self.w_s), ("w_b", self.w_b), + ("w_r", self.w_r), ("w_c", self.w_c) + ] + for w_name, w_val in weights: + if w_val < 0.0: + raise ValidationError(f"Weight {w_name} must be non-negative, got {w_val}.") + if self.r_0 <= 0.0: + raise ValidationError(f"r_0 must be strictly positive, got {self.r_0}.") + if self.epsilon <= 0.0: + raise ValidationError(f"epsilon must be strictly positive, got {self.epsilon}.") + if self.rho not in {"full-reconstruction", "compressed"}: + raise ValidationError(f"rho must be 'full-reconstruction' or 'compressed', got {self.rho}.") + if self.geometry_samples <= 0: + raise ValidationError(f"geometry_samples must be strictly positive, got {self.geometry_samples}.") @dataclass @@ -53,6 +81,32 @@ class LDEBoundaryGeometry: tangent_map: List[float] asymmetry: float + def __post_init__(self) -> None: + self.validate() + + def validate(self) -> None: + """Enforces strict invariants on L.D.E. Boundary Geometry.""" + if not self.radius_map: + raise GeometryValidationError("radius_map cannot be empty.") + if not self.curvature_map: + raise GeometryValidationError("curvature_map cannot be empty.") + if not self.tangent_map: + raise GeometryValidationError("tangent_map cannot be empty.") + + n_rad = len(self.radius_map) + if len(self.curvature_map) != n_rad or len(self.tangent_map) != n_rad: + raise GeometryValidationError( + f"Geometry lists must be of equal length. " + f"Got radius={n_rad}, curvature={len(self.curvature_map)}, tangent={len(self.tangent_map)}." + ) + + for idx, r in enumerate(self.radius_map): + if r < 0.0: + raise GeometryValidationError(f"Radius map capacity cannot be negative. Got {r} at index {idx}.") + + if self.asymmetry < 0.0: + raise GeometryValidationError(f"Asymmetry must be non-negative, got {self.asymmetry}.") + @dataclass class LDEState: diff --git a/radial_membrane_ai/lde/pipeline.py b/radial_membrane_ai/lde/pipeline.py index f12b317..d4e6217 100644 --- a/radial_membrane_ai/lde/pipeline.py +++ b/radial_membrane_ai/lde/pipeline.py @@ -14,6 +14,7 @@ LDEBoundaryGeometry, LDEState ) +from radial_membrane_ai.exceptions import ReconstructionError def lde_encode(text: str, config: Optional[LDEConfig] = None) -> LDEState: @@ -203,12 +204,13 @@ def lde_encode(text: str, config: Optional[LDEConfig] = None) -> LDEState: ) links[i][j] = P_i_to_j - # 8. Compute boundary geometry (100 sample angles) + # 8. Compute boundary geometry (configurable sample angles) # Basis function: phi_l(theta) = max(0, cos(angular_distance(theta, theta_l))) radius_map: List[float] = [] - dtheta = (2.0 * math.pi) / 100.0 + samples = config.geometry_samples + dtheta = (2.0 * math.pi) / samples - for k in range(100): + for k in range(samples): theta = k * dtheta r_val = config.r_0 for l in sigma: @@ -228,10 +230,10 @@ def lde_encode(text: str, config: Optional[LDEConfig] = None) -> LDEState: tangent_map: List[float] = [] curvature_map: List[float] = [] - for k in range(100): + for k in range(samples): r_k = radius_map[k] - r_prev = radius_map[(k - 1) % 100] - r_next = radius_map[(k + 1) % 100] + r_prev = radius_map[(k - 1) % samples] + r_next = radius_map[(k + 1) % samples] # tangent = (r_{k+1} - r_{k-1}) / (2 * dtheta) tangent_val = (r_next - r_prev) / (2.0 * dtheta) @@ -290,7 +292,7 @@ def lde_encode(text: str, config: Optional[LDEConfig] = None) -> LDEState: reconstructed.append(char) reconstructed_text = "".join(reconstructed) if reconstructed_text != text: - raise ValueError("Reconstruction verification failed!") + raise ReconstructionError("Reconstruction verification failed!") return LDEState( strings=strings, diff --git a/radial_membrane_ai/lde/visualizer.py b/radial_membrane_ai/lde/visualizer.py index b337ce6..b428535 100644 --- a/radial_membrane_ai/lde/visualizer.py +++ b/radial_membrane_ai/lde/visualizer.py @@ -29,7 +29,7 @@ def render_strings(self, state: LDEState) -> plt.Figure: """ Renders activation, spread, tension, and stiffness of each letter string as a bar plot. """ - fig = plt.figure(figsize=(10, 5), dpi=100) + fig = plt.figure(figsize=(10, 5), dpi=120) ax = fig.add_subplot(1, 1, 1) ax.set_title("L.D.E. String Local Metrics", fontsize=12, fontweight="bold") @@ -58,7 +58,7 @@ def render_channels(self, state: LDEState) -> plt.Figure: """ Renders letter-pair coherence corridors (V-Channels) in a polar routing layout. """ - fig = plt.figure(figsize=(8, 8), dpi=100) + fig = plt.figure(figsize=(8, 8), dpi=120) ax = fig.add_subplot(1, 1, 1, projection="polar") ax.set_title("L.D.E. V-Channel Routing Corridors", fontsize=12, fontweight="bold", pad=20) @@ -92,12 +92,13 @@ def render_boundary(self, state: LDEState) -> plt.Figure: """ Renders polar boundary radius map r(theta). """ - fig = plt.figure(figsize=(8, 8), dpi=100) + fig = plt.figure(figsize=(8, 8), dpi=120) ax = fig.add_subplot(1, 1, 1, projection="polar") ax.set_title("L.D.E. Polar Boundary Geometry", fontsize=12, fontweight="bold", pad=20) - dtheta = (2.0 * math.pi) / 100.0 - angles = [k * dtheta for k in range(100)] + samples = len(state.boundary.radius_map) + dtheta = (2.0 * math.pi) / samples + angles = [k * dtheta for k in range(samples)] radii = state.boundary.radius_map # Close the loop @@ -115,7 +116,7 @@ def render_boundary(self, state: LDEState) -> plt.Figure: if string_obj.depth > 0.1: theta_l = string_obj.phase # Find closest index - best_k = min(range(100), key=lambda k: abs(k * dtheta - theta_l)) + best_k = min(range(samples), key=lambda k: abs(k * dtheta - theta_l)) r_l = radii[best_k] ax.scatter(theta_l, r_l, color="#F44336", s=50, edgecolors="black", zorder=4) ax.text(theta_l, r_l + 0.1, l, fontsize=9, fontweight="bold") @@ -127,7 +128,7 @@ def render_depth_distribution(self, state: LDEState) -> plt.Figure: """ Renders the sorted depth values of all letter strings. """ - fig = plt.figure(figsize=(10, 5), dpi=100) + fig = plt.figure(figsize=(10, 5), dpi=120) ax = fig.add_subplot(1, 1, 1) ax.set_title("L.D.E. Letter Depth Distribution", fontsize=12, fontweight="bold") @@ -166,8 +167,9 @@ def render_lde_dashboard(self, state: LDEState) -> plt.Figure: # Panel A: Boundary Geometry # ------------------------------------------------------------- ax_a.set_title("A. Boundary Geometry", fontsize=10, fontweight="bold", pad=10) - dtheta = (2.0 * math.pi) / 100.0 - angles = [k * dtheta for k in range(100)] + samples = len(state.boundary.radius_map) + dtheta = (2.0 * math.pi) / samples + angles = [k * dtheta for k in range(samples)] angles_closed = angles + [angles[0]] radii_closed = state.boundary.radius_map + [state.boundary.radius_map[0]] ax_a.plot(angles_closed, radii_closed, color="#3F51B5", linewidth=1.5) @@ -180,7 +182,7 @@ def render_lde_dashboard(self, state: LDEState) -> plt.Figure: for l in deepest_letters: s_obj = state.strings[l] if s_obj.depth > 0: - best_k = min(range(100), key=lambda k: abs(k * dtheta - s_obj.phase)) + best_k = min(range(samples), key=lambda k: abs(k * dtheta - s_obj.phase)) ax_a.scatter(s_obj.phase, state.boundary.radius_map[best_k], color="#F44336", s=30, edgecolors="black") ax_a.text(s_obj.phase, state.boundary.radius_map[best_k] + 0.1, l, fontsize=8, fontweight="bold") @@ -219,7 +221,8 @@ def render_lde_dashboard(self, state: LDEState) -> plt.Figure: # Panel D: Boundary Deformation # ------------------------------------------------------------- ax_d.set_title("D. Boundary Deformation", fontsize=10, fontweight="bold") - sample_indices = np.arange(100) + samples = len(state.boundary.radius_map) + sample_indices = np.arange(samples) radius_deviation = [r - 1.0 for r in state.boundary.radius_map] ax_d.plot(sample_indices, radius_deviation, color="#9C27B0", label="Deviation (r - r_0)", linewidth=1.2) ax_d.plot(sample_indices, state.boundary.tangent_map, color="#00BCD4", label="Tangent", linewidth=1.0) diff --git a/radial_membrane_ai/tests/test_workloads.py b/radial_membrane_ai/tests/test_workloads.py index 273ac6a..4120296 100644 --- a/radial_membrane_ai/tests/test_workloads.py +++ b/radial_membrane_ai/tests/test_workloads.py @@ -279,7 +279,8 @@ def test_unsupported_simulation_target() -> None: stability_expectation=StabilityBand.GREEN, coherence_expectation=1.0, envelope_expectation=PolicyEnvelope(), - steps=[] + steps=[], + _bypass_validation=True ) with pytest.raises(ValueError, match="Unknown target:"): engine.run(bad_workload) @@ -422,7 +423,8 @@ def custom_get_reg(agent_id: str) -> Any: stability_expectation=StabilityBand.GREEN, coherence_expectation=1.1, # mismatch to cover avg_coherence < expected (line 707) envelope_expectation=PolicyEnvelope(), - steps=sa_steps + steps=sa_steps, + _bypass_validation=True ) # Monkeypatch semantic store to simulate a deletion right before step 2 execution @@ -510,7 +512,8 @@ def custom_tick_ma(task_value: float, excitation: Any, *args: Any, **kwargs: Any stability_expectation=StabilityBand.GREEN, coherence_expectation=0.0, envelope_expectation=PolicyEnvelope(), - steps=ma_steps + steps=ma_steps, + _bypass_validation=True ) _ = engine_ma.run(ma_workload) @@ -580,7 +583,8 @@ def custom_tick_ma(task_value: float, excitation: Any, *args: Any, **kwargs: Any stability_expectation=StabilityBand.GREEN, coherence_expectation=0.0, envelope_expectation=PolicyEnvelope(), - steps=mc_steps + steps=mc_steps, + _bypass_validation=True ) # Monkeypatch to release cluster and agent quarantine right before step 1. diff --git a/radial_membrane_ai/ufo_engine/multi_agent.py b/radial_membrane_ai/ufo_engine/multi_agent.py index 5a7c42c..9c8088a 100644 --- a/radial_membrane_ai/ufo_engine/multi_agent.py +++ b/radial_membrane_ai/ufo_engine/multi_agent.py @@ -12,6 +12,7 @@ from radial_membrane_ai.multi_agent.agent import UFOAgent from radial_membrane_ai.multi_agent.coupling import InterAgentVChannel, GlobalHolisticGovernor +from radial_membrane_ai.utils import set_deterministic_env from radial_membrane_ai.multi_agent.governance import MultiAgentMeshGovernance from radial_membrane_ai.shard import ShardState from radial_membrane_ai.ufo_engine.config import CostWeights, StabilityBandConfig @@ -103,6 +104,9 @@ def __init__( self.sao_events: List[Dict[str, Any]] = [] self.residual_history: List[float] = [] + # Call deterministic environment seeding + set_deterministic_env() + # Kernel Regime Expansion Layer components self.regime_manager = RegimeManager() self.tick_count: int = 0 @@ -706,6 +710,22 @@ def tick(self, task_value: float, excitation: np.ndarray) -> str: f"Regime Stability Violation: Complete Rollback triggered globally (band was '{band}')." ) + # Post-tick invariant assertions and near-violation warning logs + hard_tension_limit = 10.0 + latest_tension = self.temporal_state.accumulated_tension if self.temporal_state is not None else 0.0 + + if latest_tension > hard_tension_limit: + from radial_membrane_ai.exceptions import GovernanceError + raise GovernanceError( + f"Governed bound violated: Global accumulated tension ({latest_tension:.4f}) " + f"exceeded hard limit ({hard_tension_limit})." + ) + elif latest_tension >= 0.95 * hard_tension_limit: + self.interventions.append( + f"Near-violation warning: Global accumulated tension ({latest_tension:.4f}) " + f"is within 5% of hard limit ({hard_tension_limit})." + ) + finally: # Restore agent boundaries get_radius scale for agent in self.agents: diff --git a/radial_membrane_ai/ufo_engine/multi_cluster.py b/radial_membrane_ai/ufo_engine/multi_cluster.py index 4000e9e..c3fd029 100644 --- a/radial_membrane_ai/ufo_engine/multi_cluster.py +++ b/radial_membrane_ai/ufo_engine/multi_cluster.py @@ -29,6 +29,7 @@ # Kernel Regime imports from radial_membrane_ai.kernel_regimes.manager import RegimeManager from radial_membrane_ai.kernel_regimes.regime import KernelRegimeType +from radial_membrane_ai.utils import set_deterministic_env class MultiClusterEngine: @@ -58,6 +59,9 @@ def __init__(self) -> None: self.global_band_history: List[str] = [] self.interventions: List[str] = [] + # Call deterministic environment seeding + set_deterministic_env() + # Kernel Regime Expansion Layer components self.regime_manager = RegimeManager() self.tick_count: int = 0 @@ -765,6 +769,22 @@ def tick(self, task_value: float, default_excitation: np.ndarray) -> str: f"Regime Stability Violation: Global Multi-Cluster Rollback triggered (band was '{band}')." ) + # Post-tick invariant assertions and near-violation warning logs + hard_tension_limit = 10.0 + latest_tension = self.temporal_state.accumulated_tension if self.temporal_state is not None else 0.0 + + if latest_tension > hard_tension_limit: + from radial_membrane_ai.exceptions import GovernanceError + raise GovernanceError( + f"Governed bound violated: Global mesh tension ({latest_tension:.4f}) " + f"exceeded hard limit ({hard_tension_limit})." + ) + elif latest_tension >= 0.95 * hard_tension_limit: + self.interventions.append( + f"Near-violation warning: Global mesh tension ({latest_tension:.4f}) " + f"is within 5% of hard limit ({hard_tension_limit})." + ) + finally: for cluster in self.clusters.values(): for agent in cluster.agents: diff --git a/radial_membrane_ai/ufo_engine/single_agent.py b/radial_membrane_ai/ufo_engine/single_agent.py index 0c4e878..e9dd3ae 100644 --- a/radial_membrane_ai/ufo_engine/single_agent.py +++ b/radial_membrane_ai/ufo_engine/single_agent.py @@ -25,6 +25,7 @@ from radial_membrane_ai.facet import FacetVector, TensionAutomaton, TensionState from radial_membrane_ai.coherence import closure_coherence from radial_membrane_ai.ufo_engine.config import CostWeights, StabilityBandConfig +from radial_membrane_ai.utils import set_deterministic_env # Kernel Regime imports from radial_membrane_ai.kernel_regimes.manager import RegimeManager @@ -95,6 +96,9 @@ def __init__( self.radius_history: List[np.ndarray] = [] self.residual_history: List[float] = [] + # Call deterministic environment seeding + set_deterministic_env() + # Kernel Regime Expansion Layer components self.regime_manager = RegimeManager() self.quarantine_timer: int = 0 @@ -588,6 +592,38 @@ def tick( obs_cost = reduced_cost.weighted_cost(self.cost_weights.to_dict(), quality_signal=avg_q_coh) self.observable_cost_history.append(obs_cost) + # Post-tick invariant assertions and near-violation warning logs + hard_energy_limit = 15.0 + hard_tension_limit = 10.0 + + # Check Lyapunov energy + latest_energy = self.v_history[-1] if self.v_history else 0.0 + if latest_energy > hard_energy_limit: + from radial_membrane_ai.exceptions import GovernanceError + raise GovernanceError( + f"Governed bound violated: Lyapunov energy ({latest_energy:.4f}) " + f"exceeded hard limit ({hard_energy_limit})." + ) + elif latest_energy >= 0.95 * hard_energy_limit: + self.interventions.append( + f"Near-violation warning: Lyapunov energy ({latest_energy:.4f}) " + f"is within 5% of hard limit ({hard_energy_limit})." + ) + + # Check Tension + latest_tension = t_state.accumulated_tension if t_state is not None else 0.0 + if latest_tension > hard_tension_limit: + from radial_membrane_ai.exceptions import GovernanceError + raise GovernanceError( + f"Governed bound violated: Temporal tension ({latest_tension:.4f}) " + f"exceeded hard limit ({hard_tension_limit})." + ) + elif latest_tension >= 0.95 * hard_tension_limit: + self.interventions.append( + f"Near-violation warning: Temporal tension ({latest_tension:.4f}) " + f"is within 5% of hard limit ({hard_tension_limit})." + ) + self._record_state() finally: diff --git a/radial_membrane_ai/utils.py b/radial_membrane_ai/utils.py new file mode 100644 index 0000000..f9d632e --- /dev/null +++ b/radial_membrane_ai/utils.py @@ -0,0 +1,16 @@ +""" +Centralized Utilities for the UFO Governed Deformable Radial Membrane framework. +""" + +from __future__ import annotations +import random +import numpy as np + + +def set_deterministic_env(seed: int = 0) -> None: + """ + Sets environment seeds and properties across Python's `random` module + and NumPy's generator system to ensure absolute reproducibility. + """ + random.seed(seed) + np.random.seed(seed) diff --git a/radial_membrane_ai/workloads/engine.py b/radial_membrane_ai/workloads/engine.py index 21211f2..d38e766 100644 --- a/radial_membrane_ai/workloads/engine.py +++ b/radial_membrane_ai/workloads/engine.py @@ -26,6 +26,8 @@ WorkloadStep, Workload ) +from radial_membrane_ai.exceptions import InvalidSimulationTargetError, WorkloadConfigurationError +from radial_membrane_ai.utils import set_deterministic_env @dataclass @@ -128,10 +130,98 @@ def __init__( self.multi_agent_engine = multi_agent_engine or MultiAgentEngine() self.multi_cluster_engine = multi_cluster_engine or MultiClusterEngine() + def validate_workload(self, workload: Workload) -> None: + """ + Validates the workload up front before execution, raising structured errors. + """ + if getattr(workload, "_bypass_validation", False): + return + + # 1. Target Engine matching check + if workload.target == SimulationTarget.SINGLE_AGENT: + if not self.single_agent_engine: + raise InvalidSimulationTargetError("SingleAgentEngine not initialized.") + elif workload.target == SimulationTarget.MULTI_AGENT: + if not self.multi_agent_engine: + raise InvalidSimulationTargetError("MultiAgentEngine not initialized.") + elif workload.target == SimulationTarget.MULTI_CLUSTER: + if not self.multi_cluster_engine: + raise InvalidSimulationTargetError("MultiClusterEngine not initialized.") + elif workload.target is None: + raise InvalidSimulationTargetError("Unknown target: None") + else: + raise InvalidSimulationTargetError(f"Unknown target: {workload.target}") + + # 2. Get active engine to check agent presence + engine = self._select_engine(workload.target) + + # Gather expected active entity IDs based on target + valid_entity_ids = {"global", "single_agent"} + if workload.target == SimulationTarget.MULTI_AGENT: + for a in getattr(engine, "agents", []): + valid_entity_ids.add(a.agent_id) + elif workload.target == SimulationTarget.MULTI_CLUSTER: + for c_id, c in getattr(engine, "clusters", {}).items(): + valid_entity_ids.add(c_id) + for a in getattr(c, "agents", []): + valid_entity_ids.add(a.agent_id) + + # 3. Check actions payload structures and valid entity IDs references + for step_idx, step in enumerate(workload.steps): + # Verify agent actions target valid IDs + for action_key, act in step.agent_actions.items(): + target_id = ( + act.payload.get("entity_id") + or act.payload.get("scope") + or act.payload.get("agent_id") + or act.payload.get("cluster_id") + ) + if (target_id is not None + and target_id not in valid_entity_ids + and not target_id.startswith("agent_") + and not target_id.startswith("cluster_") + and not target_id.startswith("single_")): + raise WorkloadConfigurationError( + f"Step {step_idx}: Action '{action_key}' references invalid entity '{target_id}'." + ) + # Verify cluster actions target valid IDs + for action_key, act in step.cluster_actions.items(): + target_id = ( + act.payload.get("entity_id") + or act.payload.get("scope") + or act.payload.get("agent_id") + or act.payload.get("cluster_id") + ) + if (target_id is not None + and target_id not in valid_entity_ids + and not target_id.startswith("cluster_") + and not target_id.startswith("agent_")): + raise WorkloadConfigurationError( + f"Step {step_idx}: Action '{action_key}' references invalid entity '{target_id}'." + ) + # Verify global actions payload structure + for act in step.global_actions: + target_id = ( + act.payload.get("entity_id") + or act.payload.get("scope") + or act.payload.get("agent_id") + or act.payload.get("cluster_id") + ) + if (target_id is not None + and target_id != "global" + and target_id not in valid_entity_ids + and not target_id.startswith("agent_") + and not target_id.startswith("cluster_")): + raise WorkloadConfigurationError( + f"Step {step_idx}: Global action references invalid entity '{target_id}'." + ) + def run(self, workload: Workload) -> WorkloadResult: """ Executes a complete workload and returns a full diagnostic WorkloadResult. """ + set_deterministic_env() + self.validate_workload(workload) engine = self._select_engine(workload.target) result = WorkloadResult() @@ -340,6 +430,8 @@ def run_trace(self, workload: Workload) -> WorkloadTrace: """ Executes a complete workload and returns a full diagnostic WorkloadTrace. """ + set_deterministic_env() + self.validate_workload(workload) engine = self._select_engine(workload.target) trace_frames: List[WorkloadFrame] = [] diff --git a/radial_membrane_ai/workloads/workload.py b/radial_membrane_ai/workloads/workload.py index 56be741..4979fc0 100644 --- a/radial_membrane_ai/workloads/workload.py +++ b/radial_membrane_ai/workloads/workload.py @@ -9,6 +9,7 @@ from radial_membrane_ai.kernel_regimes.regime import KernelRegimeType from radial_membrane_ai.collective_reasoning.policy_envelope import PolicyEnvelope +from radial_membrane_ai.exceptions import WorkloadValidationError # Type aliases as specified AgentID = str @@ -115,6 +116,54 @@ class WorkloadStep: expected_stability_band: StabilityBand = StabilityBand.GREEN expected_coherence_range: Tuple[float, float] = (0.0, 1.0) expected_envelope_state: EnvelopeState = EnvelopeState.ADMIT + _bypass_validation: bool = False + + def __post_init__(self) -> None: + if not self._bypass_validation: + self.validate() + + def validate(self) -> None: + if not isinstance(self.agent_actions, dict): + raise WorkloadValidationError("agent_actions must be a dictionary.") + if not isinstance(self.cluster_actions, dict): + raise WorkloadValidationError("cluster_actions must be a dictionary.") + if not isinstance(self.global_actions, list): + raise WorkloadValidationError("global_actions must be a list.") + + # Check expected_coherence_range + if not isinstance(self.expected_coherence_range, tuple) or len(self.expected_coherence_range) != 2: + raise WorkloadValidationError("expected_coherence_range must be a tuple of length 2.") + min_coh, max_coh = self.expected_coherence_range + if not (0.0 <= min_coh <= max_coh <= 1.0): + raise WorkloadValidationError( + f"expected_coherence_range must be between 0.0 and 1.0 with min <= max. " + f"Got {self.expected_coherence_range}." + ) + + # Validate Actions + for key, act in self.agent_actions.items(): + if not isinstance(act, Action): + raise WorkloadValidationError(f"Value for agent action '{key}' must be an Action instance.") + if not isinstance(act.type, ActionType): + raise WorkloadValidationError(f"Action '{key}' has an invalid type: {act.type}.") + if not isinstance(act.payload, dict): + raise WorkloadValidationError(f"Action '{key}' payload must be a dictionary.") + + for key, act in self.cluster_actions.items(): + if not isinstance(act, Action): + raise WorkloadValidationError(f"Value for cluster action '{key}' must be an Action instance.") + if not isinstance(act.type, ActionType): + raise WorkloadValidationError(f"Action '{key}' has an invalid type: {act.type}.") + if not isinstance(act.payload, dict): + raise WorkloadValidationError(f"Action '{key}' payload must be a dictionary.") + + for idx, act in enumerate(self.global_actions): + if not isinstance(act, Action): + raise WorkloadValidationError(f"Global action at index {idx} must be an Action instance.") + if not isinstance(act.type, ActionType): + raise WorkloadValidationError(f"Global action at index {idx} has an invalid type: {act.type}.") + if not isinstance(act.payload, dict): + raise WorkloadValidationError(f"Global action at index {idx} payload must be a dictionary.") @dataclass @@ -127,6 +176,36 @@ class Workload: coherence_expectation: float envelope_expectation: PolicyEnvelope steps: List[WorkloadStep] + _bypass_validation: bool = False + + def __post_init__(self) -> None: + if not self._bypass_validation: + self.validate() + + def validate(self) -> None: + if not self.name or not isinstance(self.name, str): + raise WorkloadValidationError("Workload name must be a non-empty string.") + if not isinstance(self.target, SimulationTarget): + raise WorkloadValidationError(f"Workload target must be a valid SimulationTarget. Got {self.target}.") + if not isinstance(self.regime_expectation, KernelRegimeType): + raise WorkloadValidationError( + f"Workload regime_expectation must be a valid KernelRegimeType enum. Got {self.regime_expectation}." + ) + if not isinstance(self.stability_expectation, StabilityBand): + raise WorkloadValidationError( + f"Workload stability_expectation must be a valid StabilityBand enum. Got {self.stability_expectation}." + ) + if not (0.0 <= self.coherence_expectation <= 1.0): + raise WorkloadValidationError( + f"Workload coherence_expectation must be between 0.0 and 1.0. Got {self.coherence_expectation}." + ) + if not isinstance(self.envelope_expectation, PolicyEnvelope): + raise WorkloadValidationError("Workload envelope_expectation must be a valid PolicyEnvelope instance.") + if not isinstance(self.steps, list): + raise WorkloadValidationError("Workload steps must be a list of WorkloadStep instances.") + for idx, step in enumerate(self.steps): + if not isinstance(step, WorkloadStep): + raise WorkloadValidationError(f"Workload step at index {idx} must be a WorkloadStep instance.") # =============================================================================