diff --git a/adapters/simulation/grid.py b/adapters/simulation/grid.py index 407ee66..e8d9938 100644 --- a/adapters/simulation/grid.py +++ b/adapters/simulation/grid.py @@ -11,10 +11,11 @@ import random from dataclasses import dataclass +from typing import ClassVar -from core.models.state import WorldState, AgentState, Position -from core.models.actions import Action, ActionType, ActionResult -from core.events.bus import EventBus, Event, EventType +from core.events.bus import Event, EventBus, EventType +from core.models.actions import Action, ActionResult, ActionType +from core.models.state import AgentState, Position, WorldState @dataclass @@ -175,7 +176,7 @@ def _handle_retreat(self, action: Action) -> ActionResult: energy_cost=self._config.move_cost, tick=self._world.tick, ) - _action_handlers: dict[ActionType, object] = { + _action_handlers: ClassVar[dict[ActionType, object]] = { ActionType.MOVE: _handle_move, ActionType.WAIT: _handle_wait, ActionType.ATTACK: _handle_attack, diff --git a/core/events/bus.py b/core/events/bus.py index 1e5e5d3..5feb35b 100644 --- a/core/events/bus.py +++ b/core/events/bus.py @@ -10,14 +10,12 @@ from __future__ import annotations from collections import defaultdict +from collections.abc import Callable from dataclasses import dataclass, field -from enum import Enum -from typing import Callable +from enum import StrEnum -from core.models.state import Position - -class EventType(str, Enum): +class EventType(StrEnum): """All valid event types in the runtime.""" # Simulation events diff --git a/core/models/actions.py b/core/models/actions.py index aaf01d9..d067c56 100644 --- a/core/models/actions.py +++ b/core/models/actions.py @@ -2,14 +2,16 @@ from __future__ import annotations -from enum import Enum +from enum import StrEnum +from typing import TYPE_CHECKING from pydantic import BaseModel, Field -from core.models.state import Position +if TYPE_CHECKING: + from core.models.state import Position -class ActionType(str, Enum): +class ActionType(StrEnum): """Available agent actions.""" MOVE = "move" diff --git a/core/models/goals.py b/core/models/goals.py index 4d54e6b..651761b 100644 --- a/core/models/goals.py +++ b/core/models/goals.py @@ -2,12 +2,12 @@ from __future__ import annotations -from enum import Enum +from enum import StrEnum from pydantic import BaseModel -class GoalStatus(str, Enum): +class GoalStatus(StrEnum): """Goal lifecycle status.""" ACTIVE = "active" diff --git a/core/models/traces.py b/core/models/traces.py index b04e4ef..8b18aa4 100644 --- a/core/models/traces.py +++ b/core/models/traces.py @@ -6,9 +6,12 @@ from __future__ import annotations +from typing import TYPE_CHECKING + from pydantic import BaseModel, Field -from core.models.actions import Action +if TYPE_CHECKING: + from core.models.actions import Action class DecisionReason(BaseModel): diff --git a/core/runtime/loop.py b/core/runtime/loop.py index 2be3365..d579e8d 100644 --- a/core/runtime/loop.py +++ b/core/runtime/loop.py @@ -14,13 +14,15 @@ from __future__ import annotations -from core.models.state import AgentState, WorldState -from core.models.actions import Action, ActionType, ActionResult -from core.models.traces import DecisionTrace, DecisionReason -from core.events.bus import EventBus, Event, EventType -from core.constraints.engine import ConstraintEngine -from adapters.simulation.grid import GridSimulation -from policies.fsm.agent_fsm import AgentFSM +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from adapters.simulation.grid import GridSimulation + from core.constraints.engine import ConstraintEngine + from policies.fsm.agent_fsm import AgentFSM +from core.events.bus import Event, EventBus, EventType +from core.models.actions import Action, ActionType +from core.models.traces import DecisionReason, DecisionTrace class RuntimeLoop: diff --git a/policies/fsm/agent_fsm.py b/policies/fsm/agent_fsm.py index 83da361..287310b 100644 --- a/policies/fsm/agent_fsm.py +++ b/policies/fsm/agent_fsm.py @@ -10,14 +10,13 @@ from __future__ import annotations -from enum import Enum +from enum import StrEnum -from core.models.state import AgentState, WorldState from core.models.actions import Action, ActionType -from core.models.state import Position +from core.models.state import AgentState, Position, WorldState -class AgentMode(str, Enum): +class AgentMode(StrEnum): """Top-level agent behavior modes.""" EXPLORE = "explore" diff --git a/replay/trace_log.py b/replay/trace_log.py index cda0412..3456ec3 100644 --- a/replay/trace_log.py +++ b/replay/trace_log.py @@ -7,8 +7,10 @@ from __future__ import annotations import json -from pathlib import Path -from dataclasses import asdict +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from pathlib import Path from core.models.traces import DecisionTrace diff --git a/tests/integration/test_scenario.py b/tests/integration/test_scenario.py index 94bc261..a515ff9 100644 --- a/tests/integration/test_scenario.py +++ b/tests/integration/test_scenario.py @@ -8,10 +8,10 @@ - Same seed produces same outcome (replayability) """ -from core.events.bus import EventBus, EventType +from adapters.simulation.grid import GridSimulation, SimulationConfig from core.constraints.engine import create_default_engine +from core.events.bus import EventBus, EventType from core.runtime.loop import RuntimeLoop -from adapters.simulation.grid import GridSimulation, SimulationConfig from policies.fsm.agent_fsm import AgentFSM @@ -43,7 +43,7 @@ def test_deterministic_replay(self) -> None: run1 = self._run_scenario(seed=42) run2 = self._run_scenario(seed=42) assert len(run1.traces) == len(run2.traces) - for t1, t2 in zip(run1.traces, run2.traces): + for t1, t2 in zip(run1.traces, run2.traces, strict=True): assert t1.tick == t2.tick assert t1.selected_action.action_type == t2.selected_action.action_type @@ -52,9 +52,9 @@ def test_different_seed_different_result(self) -> None: run2 = self._run_scenario(seed=99) # Different seeds should produce different worlds # (traces may differ in length or actions) - traces_match = all( + all( t1.selected_action.action_type == t2.selected_action.action_type - for t1, t2 in zip(run1.traces, run2.traces) + for t1, t2 in zip(run1.traces, run2.traces, strict=False) ) # Not a guarantee they differ, but very likely with different seeds assert len(run1.traces) > 0 @@ -85,7 +85,7 @@ def test_agent_cannot_move_out_of_bounds(self) -> None: runtime = RuntimeLoop(sim, constraints, fsm, bus, max_ticks=30) runtime.run() # No constraint violation should crash the system - violations = bus.get_history(event_type=EventType.CONSTRAINT_VIOLATED) + bus.get_history(event_type=EventType.CONSTRAINT_VIOLATED) # Any violations were handled gracefully (fallback to wait) for trace in runtime.traces: if len(trace.constraints_violated) > 0: