aether-taskflow / env /aether_env.py
Nithin1026's picture
Initial submit
a74cbe6 verified
Raw
History Blame
25.7 kB
from __future__ import annotations
import random
import re
import sys
import uuid
from pathlib import Path
from typing import Any, Dict, List, Optional
# Allow running from repo root or server/
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
from openenv.core import Environment
from models import (
AetherTaskFlowAction,
AetherTaskFlowObservation,
AetherTaskFlowState,
ActionType,
)
from env.tasks import generate_tasks, get_profile, apply_dynamic_updates
from env.grader import grade
class AetherTaskFlowEnvironment(Environment):
"""
AETHER-TaskFlow: Adaptive Workflow Management RL Environment.
The agent manages a dynamic task queue under resource constraints
and system uncertainty. Three task scenarios of increasing difficulty
test baseline reasoning, adaptation, and robustness.
"""
SUPPORTS_CONCURRENT_SESSIONS = True
DEFAULT_SEED = 42
def __init__(self, difficulty: str = "easy", default_seed: int = DEFAULT_SEED) -> None:
super().__init__()
if difficulty not in ("easy", "medium", "hard"):
raise ValueError(f"difficulty must be easy/medium/hard, got '{difficulty}'")
self._difficulty = difficulty
self._profile = get_profile(difficulty)
self._default_seed = int(default_seed)
self._state: AetherTaskFlowState = AetherTaskFlowState()
self._rng = random.Random()
self._deferred_tasks: List[Dict[str, Any]] = []
def reset(
self,
seed: Optional[int] = None,
episode_id: Optional[str] = None,
**kwargs: Any,
) -> AetherTaskFlowObservation:
self._reset_rubric()
seed = self._default_seed if seed is None else int(seed)
# === CRITICAL: Ensure determinism ===
random.seed(seed)
self._rng = random.Random(seed)
ep_id = episode_id or str(uuid.uuid4())
profile = self._profile
tasks = generate_tasks(self._difficulty, seed=seed)
task_dicts = [t.to_dict() for t in tasks]
resources = {
"time": float(profile["initial_time"]),
"energy": float(profile["initial_energy"]),
"budget": float(profile["initial_budget"]),
}
self._state = AetherTaskFlowState(
episode_id=ep_id,
step_count=0,
difficulty=self._difficulty,
tasks=task_dicts,
completed_tasks=[],
failed_tasks=[],
deferred_tasks=[],
resources=dict(resources),
initial_resources=dict(resources),
system_health=1.0,
cumulative_value=0.0,
cumulative_reward=0.0,
tasks_completed=0,
tasks_failed=0,
episode_done=False,
seed=seed,
)
self._deferred_tasks = []
self._sync_state_queues()
return self._build_obs(
last_action_type=None,
last_action_task_id=None,
last_action_outcome="Episode started. Select a task to act on.",
reward=0.0,
done=False,
)
def step(
self,
action: AetherTaskFlowAction | Dict[str, Any] | str,
timeout_s: Optional[float] = None,
**kwargs: Any,
) -> AetherTaskFlowObservation:
if self._state.episode_id is None:
self.reset(seed=self._default_seed)
try:
parsed_action = self._coerce_action(action)
return self._step_impl(parsed_action)
except Exception as exc:
return self._safe_step_failure(action, exc)
def _step_impl(self, action: AetherTaskFlowAction) -> AetherTaskFlowObservation:
s = self._state
if s.episode_done:
return self._build_obs(
last_action_type=None,
last_action_task_id=None,
last_action_outcome="Episode already finished.",
reward=0.0,
done=True,
)
s.step_count += 1
profile = self._profile
max_steps: int = profile["max_steps"]
# ---- Apply dynamic task updates (medium/hard) ----
if self._difficulty in ("medium", "hard"):
from env.tasks import apply_dynamic_updates as _upd
task_objs_updated = _upd(
[self._make_task_info(t) for t in s.tasks],
s.step_count,
self._difficulty,
self._rng,
)
s.tasks = [t.to_dict() for t in task_objs_updated]
# ---- Deadline expiry check (before acting) ----
still_alive, newly_failed = [], []
for t in s.tasks:
if t.get("deadline", 1) <= 0 and t["status"] == "pending":
t["status"] = "failed"
newly_failed.append(t)
s.system_health = max(0.0, s.system_health - 0.05)
else:
still_alive.append(t)
s.tasks = still_alive
s.failed_tasks.extend(newly_failed)
s.tasks_failed += len(newly_failed)
# ---- Find the target task ----
task = self._find_task(action.task_id, s.tasks)
if task is None:
# Try deferred list
task = self._find_task(action.task_id, self._deferred_tasks)
raw_reward = 0.0
outcome = ""
if task is None:
raw_reward = -0.5
outcome = (
f"Task {action.task_id} not found in active queue. "
"Choose a valid task_id from the observation."
)
s.system_health = max(0.0, s.system_health - 0.02)
else:
raw_reward, outcome = self._execute_action(action.action_type, task, s, max_steps)
s.cumulative_reward += raw_reward
# ---- Decrement deadlines each step ----
for t in s.tasks:
if t["status"] == "pending":
t["deadline"] = max(0, t["deadline"] - 1)
# ---- Recycle deferred tasks if resources improve ----
from env.algorithms import AWFROX
recycler = AWFROX()
resources_dict = {
"energy": s.resources["energy"],
"budget": s.resources["budget"],
}
active_updated, still_deferred = recycler.recycle_deferred(
s.tasks, self._deferred_tasks, resources_dict, s.step_count
)
s.tasks = active_updated
self._deferred_tasks = still_deferred
self._sync_state_queues()
# ---- Done condition ----
no_more_tasks = len(s.tasks) == 0 and len(self._deferred_tasks) == 0
out_of_time = s.step_count >= max_steps
out_of_resources = (
s.resources["energy"] <= 0 or s.resources["time"] <= 0
)
system_collapse = s.system_health <= 0.0
done = no_more_tasks or out_of_time or out_of_resources or system_collapse
s.episode_done = done
self._sync_state_queues()
return self._build_obs(
last_action_type=action.action_type.value,
last_action_task_id=action.task_id,
last_action_outcome=outcome,
reward=self._normalize_step_reward(raw_reward),
done=done,
)
def message_to_action(self, message: str) -> AetherTaskFlowAction:
"""Convert free-form UI text into a valid environment action."""
return self._coerce_action(message)
def _coerce_action(
self,
action: AetherTaskFlowAction | Dict[str, Any] | str | None,
) -> AetherTaskFlowAction:
if isinstance(action, AetherTaskFlowAction):
return action
if action is None:
return self._recommended_action("No action provided; selected a safe default.")
if isinstance(action, str):
return self._parse_action_message(action)
if isinstance(action, dict):
if "message" in action and isinstance(action["message"], str):
return self._parse_action_message(action["message"])
if "input" in action and isinstance(action["input"], str):
return self._parse_action_message(action["input"])
if "action" in action:
nested_action = action["action"]
if isinstance(nested_action, (dict, str)) or nested_action is None:
return self._coerce_action(nested_action)
recommended = self._recommended_action("Filled missing action fields from the current state.")
normalized_payload = {
"action_type": action.get("action_type", recommended.action_type.value),
"task_id": action.get("task_id", recommended.task_id),
"reasoning": action.get("reasoning", recommended.reasoning),
}
return AetherTaskFlowAction.model_validate(normalized_payload)
raise TypeError(f"Unsupported action input: {type(action)!r}")
def _parse_action_message(self, message: str) -> AetherTaskFlowAction:
normalized = (message or "").strip().lower()
recommended = self._recommended_action(
"Selected the top-ranked task from the current observation."
)
if not normalized:
return recommended
keyword_map = (
(ActionType.OPTIMIZE, ("optimize", "optimise", "tune", "analyze", "analyse")),
(ActionType.DELEGATE, ("delegate", "assign", "handoff", "hand off", "offload")),
(ActionType.DEFER, ("defer", "later", "wait", "skip", "postpone")),
(ActionType.EXECUTE, ("execute", "run", "do", "complete", "process", "start")),
)
chosen_action = recommended.action_type
for action_type, keywords in keyword_map:
if any(keyword in normalized for keyword in keywords):
chosen_action = action_type
break
requested_task_id = self._extract_task_id(normalized)
if requested_task_id is not None and self._task_exists(requested_task_id):
task_id = requested_task_id
else:
task_id = recommended.task_id
return AetherTaskFlowAction(
action_type=chosen_action,
task_id=task_id,
reasoning=f"parsed from '{message.strip()[:80]}'",
)
def _recommended_action(self, reasoning: str) -> AetherTaskFlowAction:
candidates = self._iter_candidate_tasks()
if not candidates:
return AetherTaskFlowAction(
action_type=ActionType.DEFER,
task_id=0,
reasoning=reasoning,
)
from env.algorithms import AETHER, RAPTOR
resources = {
"energy": self._state.resources.get("energy", 0.0),
"budget": self._state.resources.get("budget", 0.0),
"time": self._state.resources.get("time", 0.0),
}
max_steps = self._profile["max_steps"]
ranked = AETHER().rank_tasks(candidates, resources, self._state.step_count, max_steps)
best_task_id, _ = ranked[0]
best_task = next(task for task in candidates if task["task_id"] == best_task_id)
action_type = ActionType(
RAPTOR().decide(best_task, resources, self._state.step_count, max_steps)
)
return AetherTaskFlowAction(
action_type=action_type,
task_id=best_task_id,
reasoning=reasoning,
)
def _iter_candidate_tasks(self) -> List[Dict[str, Any]]:
active_tasks = [task for task in self._state.tasks if task.get("status") == "pending"]
if active_tasks:
return active_tasks
deferred_tasks = [
task for task in self._deferred_tasks if task.get("status") in ("pending", "deferred")
]
return deferred_tasks
def _task_exists(self, task_id: int) -> bool:
return self._find_task(task_id, self._state.tasks) is not None or self._find_task(
task_id, self._deferred_tasks
) is not None
def _extract_task_id(self, text: str) -> Optional[int]:
explicit_match = re.search(r"(?:task|id|#)\s*(\d+)", text)
if explicit_match:
return int(explicit_match.group(1))
loose_match = re.search(r"\b(\d+)\b", text)
if loose_match:
return int(loose_match.group(1))
return None
def _safe_step_failure(
self,
action: AetherTaskFlowAction | Dict[str, Any] | str,
exc: Exception,
) -> AetherTaskFlowObservation:
self._state.episode_done = True
self._state.system_health = max(0.0, self._state.system_health - 0.1)
self._sync_state_queues()
last_action_type = None
last_action_task_id = None
if isinstance(action, AetherTaskFlowAction):
last_action_type = action.action_type.value
last_action_task_id = action.task_id
elif isinstance(action, dict):
raw_action_type = action.get("action_type")
if isinstance(raw_action_type, str):
last_action_type = raw_action_type
raw_task_id = action.get("task_id")
if isinstance(raw_task_id, int):
last_action_task_id = raw_task_id
return self._build_obs(
last_action_type=last_action_type,
last_action_task_id=last_action_task_id,
last_action_outcome=(
f"Step failed safely: {type(exc).__name__}: {str(exc)[:160]}"
),
reward=self._normalize_step_reward(-1.0),
done=True,
)
def _reset_rubric(self) -> None:
"""Called at the start of every reset() — OpenEnv lifecycle hook."""
# No persistent rubric state in this env; this hook satisfies the
# openenv.core.Environment base-class interface.
pass
def get_metadata(self) -> dict:
"""Return environment metadata (used by WebInterfaceManager on startup)."""
return {
"name": "aether_taskflow",
"description": (
"AETHER-TaskFlow: Adaptive Workflow Management RL Environment. "
"Agent manages a dynamic task queue under resource constraints, "
"uncertainty, and time pressure. Real-world enterprise tasks."
),
"difficulty": self._difficulty,
"max_steps": self._profile["max_steps"],
"action_types": ["execute", "defer", "delegate", "optimize"],
"version": "1.0.0",
}
def close(self) -> None:
"""Clean up environment resources (no-op for this in-memory env)."""
pass
# ------------------------------------------------------------------
# state property (OpenEnv required)
# ------------------------------------------------------------------
@property
def state(self) -> AetherTaskFlowState:
self._sync_state_queues()
return self._state
def _execute_action(
self,
action_type: ActionType,
task: Dict[str, Any],
s: AetherTaskFlowState,
max_steps: int,
) -> tuple[float, str]:
"""Execute the chosen action on a task. Returns (reward, outcome_str)."""
energy_cost = task["required_energy"]
budget_cost = task["required_budget"]
uncertainty = task["uncertainty"]
value = task["value"]
priority = task["priority"]
deadline = task["deadline"]
time_left = max_steps - s.step_count
if action_type == ActionType.EXECUTE:
# Check resource sufficiency
if s.resources["energy"] < energy_cost or s.resources["budget"] < budget_cost:
s.system_health = max(0.0, s.system_health - 0.08)
return -1.0, (
f"Cannot execute '{task['name']}': insufficient resources "
f"(need E={energy_cost:.1f}/B={budget_cost:.1f}, "
f"have E={s.resources['energy']:.1f}/B={s.resources['budget']:.1f})."
)
# Uncertainty-based failure chance
success_prob = 1.0 - uncertainty * 0.4
if self._rng.random() > success_prob:
# Partial failure — lose resources, get partial reward
s.resources["energy"] = max(0.0, s.resources["energy"] - energy_cost * 0.5)
s.resources["budget"] = max(0.0, s.resources["budget"] - budget_cost * 0.5)
s.system_health = max(0.0, s.system_health - 0.06)
self._remove_task(task["task_id"], s)
partial_reward = value * priority * 0.25
s.cumulative_value += partial_reward
return partial_reward, (
f"Partial failure on '{task['name']}' (uncertainty={uncertainty:.2f}). "
f"Partial reward: {partial_reward:.2f}."
)
# Success
s.resources["energy"] = max(0.0, s.resources["energy"] - energy_cost)
s.resources["budget"] = max(0.0, s.resources["budget"] - budget_cost)
s.resources["time"] = max(0.0, s.resources["time"] - 1.0)
task["status"] = "completed"
self._remove_task(task["task_id"], s)
s.completed_tasks.append(task)
s.tasks_completed += 1
# Reward: base value × priority, bonus for early completion
deadline_bonus = max(0.0, deadline / max(time_left, 1)) * 0.5
reward = value * priority + deadline_bonus
s.cumulative_value += reward
return reward, (
f"Successfully executed '{task['name']}'. "
f"Reward: {reward:.2f} (value={value:.1f}, priority={priority:.2f})."
)
elif action_type == ActionType.DEFER:
# Low penalty; task goes to deferred queue
task["status"] = "deferred"
self._remove_task(task["task_id"], s)
self._deferred_tasks.append(task)
self._sync_state_queues()
s.resources["time"] = max(0.0, s.resources["time"] - 0.5)
defer_penalty = -0.2 * priority # higher priority = bigger penalty for deferring
return defer_penalty, (
f"Deferred '{task['name']}'. "
f"Penalty: {defer_penalty:.2f}. Will retry when resources recover."
)
elif action_type == ActionType.DELEGATE:
# Offload — no resource cost, reduced reward
task["status"] = "completed"
self._remove_task(task["task_id"], s)
s.completed_tasks.append(task)
s.tasks_completed += 1
delegate_reward = value * priority * 0.35
s.cumulative_value += delegate_reward
return delegate_reward, (
f"Delegated '{task['name']}'. "
f"Reward: {delegate_reward:.2f} (35% of full value)."
)
elif action_type == ActionType.OPTIMIZE:
# Spend a small energy/budget to reduce uncertainty
opt_energy = max(0.3, energy_cost * 0.2)
opt_budget = max(1.0, budget_cost * 0.15)
if s.resources["energy"] < opt_energy:
return -0.1, f"Cannot optimize '{task['name']}': not enough energy."
s.resources["energy"] = max(0.0, s.resources["energy"] - opt_energy)
s.resources["budget"] = max(0.0, s.resources["budget"] - opt_budget)
s.resources["time"] = max(0.0, s.resources["time"] - 0.5)
# Reduce uncertainty significantly
reduction = self._rng.uniform(0.2, 0.45)
old_unc = task["uncertainty"]
task["uncertainty"] = max(0.02, task["uncertainty"] - reduction)
return 0.1, (
f"Optimized '{task['name']}': uncertainty {old_unc:.2f}{task['uncertainty']:.2f}. "
f"Small positive reward for risk reduction."
)
return 0.0, "Unknown action type."
def _find_task(self, task_id: int, task_list: List[Dict[str, Any]]) -> Optional[Dict[str, Any]]:
for t in task_list:
if t["task_id"] == task_id and t["status"] in ("pending", "deferred"):
return t
return None
def _remove_task(self, task_id: int, s: AetherTaskFlowState) -> None:
s.tasks = [t for t in s.tasks if t["task_id"] != task_id]
self._deferred_tasks = [t for t in self._deferred_tasks if t["task_id"] != task_id]
self._sync_state_queues()
def _sync_state_queues(self) -> None:
self._state.deferred_tasks = [dict(task) for task in self._deferred_tasks]
def _max_positive_step_reward(self) -> float:
"""Upper-bound the raw positive reward for the current difficulty profile."""
priority_max = float(self._profile["priority_range"][1])
value_max = float(self._profile["value_range"][1])
deadline_max = float(self._profile["deadline_range"][1])
return max(1.0, (value_max * priority_max) + (deadline_max * 0.5))
def _normalize_step_reward(self, raw_reward: float) -> float:
"""
Map raw action rewards into [0, 1] for OpenEnv-facing observations.
Negative rewards occupy [0.0, 0.5), zero remains a neutral midpoint in
the action-reward scale, and positive rewards occupy (0.5, 1.0].
Non-action lifecycle observations (reset/already-done) still emit their
explicit reward values directly via _build_obs.
"""
min_reward = -1.0
if raw_reward <= 0.0:
normalized = ((raw_reward - min_reward) / (0.0 - min_reward)) * 0.5
else:
max_reward = self._max_positive_step_reward()
normalized = 0.5 + 0.5 * min(raw_reward / max_reward, 1.0)
return round(max(0.0, min(1.0, normalized)), 4)
def _get_obs(self) -> Dict[str, Any]:
"""Return a compact state snapshot for manual debugging and simple UIs."""
s = self._state
self._sync_state_queues()
return {
"episode_id": s.episode_id,
"difficulty": s.difficulty,
"step_count": s.step_count,
"num_tasks": len(s.tasks),
"num_deferred_tasks": len(self._deferred_tasks),
"tasks_completed": s.tasks_completed,
"tasks_failed": s.tasks_failed,
"resources": {
"time": round(s.resources.get("time", 0.0), 2),
"energy": round(s.resources.get("energy", 0.0), 2),
"budget": round(s.resources.get("budget", 0.0), 2),
},
"system_health": round(s.system_health, 2),
"done": s.episode_done,
}
def _make_task_info(self, t: Dict[str, Any]):
from env.tasks import TaskInfo as _TI
return _TI(
task_id=t["task_id"],
name=t["name"],
priority=t["priority"],
deadline=t["deadline"],
uncertainty=t["uncertainty"],
value=t["value"],
required_energy=t["required_energy"],
required_budget=t["required_budget"],
category=t["category"],
status=t["status"],
)
def _build_obs(
self,
last_action_type: Optional[str],
last_action_task_id: Optional[int],
last_action_outcome: Optional[str],
reward: float,
done: bool,
) -> AetherTaskFlowObservation:
s = self._state
return AetherTaskFlowObservation(
done=done,
reward=reward,
metadata={
"difficulty": s.difficulty,
"episode_id": s.episode_id,
"step_count": s.step_count,
"summary": self._get_obs(),
},
tasks=list(s.tasks),
time_remaining=int(s.resources.get("time", 0)),
energy_remaining=round(s.resources.get("energy", 0.0), 2),
budget_remaining=round(s.resources.get("budget", 0.0), 2),
system_health=round(s.system_health, 4),
step_number=s.step_count,
tasks_completed=s.tasks_completed,
tasks_failed=s.tasks_failed,
cumulative_value=round(s.cumulative_value, 4),
last_action_type=last_action_type,
last_action_task_id=last_action_task_id,
last_action_outcome=last_action_outcome,
difficulty=s.difficulty,
episode_id=s.episode_id,
)
def _build_grade_result(self) -> Dict[str, Any]:
s = self._state
profile = self._profile
return {
"difficulty": s.difficulty,
"tasks_completed": s.tasks_completed,
"tasks_failed": s.tasks_failed,
"total_tasks": s.tasks_completed + s.tasks_failed + len(s.tasks) + len(self._deferred_tasks),
"remaining_time": s.resources.get("time", 0),
"remaining_energy": s.resources.get("energy", 0),
"remaining_budget": s.resources.get("budget", 0),
"initial_time": s.initial_resources.get("time", profile["initial_time"]),
"initial_energy": s.initial_resources.get("energy", profile["initial_energy"]),
"initial_budget": s.initial_resources.get("budget", profile["initial_budget"]),
"system_health": s.system_health,
"steps_used": s.step_count,
"max_steps": profile["max_steps"],
"cumulative_value": s.cumulative_value,
}
def compute_final_score(self) -> float:
"""Compute the final grade [0, 1] for the completed episode."""
result = self._build_grade_result()
score = grade(self._difficulty, result)
return max(0.0, min(1.0, score))