aether-taskflow / tests /test_aether_taskflow.py
Nithin1026's picture
Initial submit (#1)
9a28110
Raw
History Blame Contribute Delete
22.5 kB
from __future__ import annotations
import sys
from pathlib import Path
sys.path.insert(0, str(Path(__file__).resolve().parents[1]))
import pytest
from fastapi.testclient import TestClient
from models import AetherTaskFlowAction, AetherTaskFlowObservation, AetherTaskFlowState, ActionType
from env.tasks import generate_tasks, get_profile
from env.algorithms import AETHER, RAPTOR, AWFROX
from env.grader import grade, grade_easy, grade_medium, grade_hard
from env.aether_env import AetherTaskFlowEnvironment
from server.app import app
@pytest.fixture
def easy_env():
return AetherTaskFlowEnvironment(difficulty="easy")
@pytest.fixture
def medium_env():
return AetherTaskFlowEnvironment(difficulty="medium")
@pytest.fixture
def hard_env():
return AetherTaskFlowEnvironment(difficulty="hard")
@pytest.fixture
def resources():
return {"energy": 10.0, "budget": 50.0, "time": 10.0}
@pytest.fixture
def api_client():
return TestClient(app)
class TestTaskGeneration:
def test_easy_task_count(self):
tasks = generate_tasks("easy", seed=42)
assert len(tasks) == 5
def test_medium_task_count(self):
tasks = generate_tasks("medium", seed=42)
assert len(tasks) == 8
def test_hard_task_count(self):
tasks = generate_tasks("hard", seed=42)
assert len(tasks) == 12
def test_tasks_have_required_fields(self):
tasks = generate_tasks("easy", seed=1)
for t in tasks:
d = t.to_dict()
assert "task_id" in d
assert "priority" in d
assert "deadline" in d
assert "uncertainty" in d
assert "value" in d
assert "required_energy" in d
assert "required_budget" in d
assert "category" in d
assert "status" in d
def test_priority_range(self):
tasks = generate_tasks("easy", seed=7)
for t in tasks:
assert 0.0 <= t.priority <= 1.0
def test_uncertainty_range(self):
tasks = generate_tasks("hard", seed=7)
for t in tasks:
assert 0.0 <= t.uncertainty <= 1.0
def test_reproducibility(self):
a = generate_tasks("medium", seed=99)
b = generate_tasks("medium", seed=99)
assert [t.task_id for t in a] == [t.task_id for t in b]
assert [round(t.priority, 5) for t in a] == [round(t.priority, 5) for t in b]
def test_different_seeds_differ(self):
a = generate_tasks("easy", seed=1)
b = generate_tasks("easy", seed=2)
# At least one task should differ
priorities_a = [t.priority for t in a]
priorities_b = [t.priority for t in b]
assert priorities_a != priorities_b
class TestAETHER:
def test_score_returns_float(self, resources):
aether = AETHER()
task = {
"task_id": 0, "priority": 0.8, "deadline": 3,
"uncertainty": 0.2, "value": 15.0,
"required_energy": 1.0, "required_budget": 5.0,
}
score = aether.score(task, resources, step=0, max_steps=10)
assert isinstance(score, float)
def test_higher_priority_scores_higher(self, resources):
aether = AETHER()
low = {"task_id": 0, "priority": 0.2, "deadline": 5, "uncertainty": 0.1, "value": 10.0, "required_energy": 1.0, "required_budget": 5.0}
high = {"task_id": 1, "priority": 0.9, "deadline": 5, "uncertainty": 0.1, "value": 10.0, "required_energy": 1.0, "required_budget": 5.0}
assert aether.score(high, resources, 0, 10) > aether.score(low, resources, 0, 10)
def test_high_uncertainty_penalised(self, resources):
aether = AETHER()
base = {"task_id": 0, "priority": 0.7, "deadline": 4, "value": 10.0, "required_energy": 1.0, "required_budget": 5.0}
low_unc = {**base, "uncertainty": 0.1}
high_unc = {**base, "uncertainty": 0.9}
assert aether.score(low_unc, resources, 0, 10) > aether.score(high_unc, resources, 0, 10)
def test_rank_tasks_sorted_descending(self, resources):
aether = AETHER()
tasks = [
{"task_id": i, "priority": 0.3 + i * 0.2, "deadline": 5,
"uncertainty": 0.1, "value": 10.0, "required_energy": 1.0, "required_budget": 5.0}
for i in range(4)
]
ranked = aether.rank_tasks(tasks, resources, 0, 10)
scores = [s for _, s in ranked]
assert scores == sorted(scores, reverse=True)
def test_update_modifies_weights(self):
aether = AETHER()
weights_before = dict(aether.weights)
aether.update(5.0)
aether.update(5.0)
assert aether.weights != weights_before
def test_weights_stay_in_range(self):
aether = AETHER()
for _ in range(50):
aether.update(10.0)
assert aether.weights["priority"] >= 0.5
assert aether.weights["uncertainty_penalty"] <= -0.1
class TestRAPTOR:
def test_defers_on_low_energy(self):
raptor = RAPTOR()
task = {"priority": 0.8, "deadline": 3, "uncertainty": 0.2, "value": 15.0,
"required_energy": 5.0, "required_budget": 5.0}
resources = {"energy": 1.0, "budget": 50.0} # energy too low
action = raptor.decide(task, resources, step=0, max_steps=10)
assert action in ("defer", "delegate")
def test_optimizes_high_uncertainty(self):
raptor = RAPTOR()
task = {"priority": 0.8, "deadline": 4, "uncertainty": 0.9, "value": 15.0,
"required_energy": 1.0, "required_budget": 5.0}
resources = {"energy": 10.0, "budget": 50.0}
action = raptor.decide(task, resources, step=0, max_steps=10)
assert action == "optimize"
def test_executes_with_good_resources(self):
raptor = RAPTOR()
task = {"priority": 0.8, "deadline": 3, "uncertainty": 0.1, "value": 15.0,
"required_energy": 1.0, "required_budget": 5.0}
resources = {"energy": 10.0, "budget": 50.0}
action = raptor.decide(task, resources, step=0, max_steps=10)
assert action == "execute"
def test_all_return_valid_action(self):
raptor = RAPTOR()
valid = {"execute", "defer", "delegate", "optimize"}
for seed in range(20):
import random
rng = random.Random(seed)
task = {"priority": rng.random(), "deadline": rng.randint(0, 8),
"uncertainty": rng.random(), "value": rng.uniform(3, 30),
"required_energy": rng.uniform(0.5, 4), "required_budget": rng.uniform(1, 20)}
resources = {"energy": rng.uniform(0, 12), "budget": rng.uniform(0, 60)}
action = raptor.decide(task, resources, rng.randint(0, 9), 10)
assert action in valid
class TestAWFROX:
def test_removes_expired_tasks(self):
recycler = AWFROX()
tasks = [
{"task_id": 0, "deadline": -1, "status": "pending"},
{"task_id": 1, "deadline": 3, "status": "pending"},
]
viable = recycler.filter_viable(tasks, {}, step=0, max_steps=10)
assert len(viable) == 1
assert viable[0]["task_id"] == 1
def test_recycles_deferred_when_resources_available(self):
recycler = AWFROX()
active = []
deferred = [{"task_id": 5, "status": "deferred", "deadline": 3,
"required_energy": 1.0, "required_budget": 5.0}]
resources = {"energy": 10.0, "budget": 50.0}
new_active, new_deferred = recycler.recycle_deferred(active, deferred, resources, step=2)
assert len(new_active) == 1
assert new_active[0]["status"] == "pending"
assert len(new_deferred) == 0
def test_keeps_deferred_when_resources_insufficient(self):
recycler = AWFROX()
active = []
deferred = [{"task_id": 5, "status": "deferred", "deadline": 3,
"required_energy": 10.0, "required_budget": 50.0}]
resources = {"energy": 0.5, "budget": 1.0} # insufficient
new_active, new_deferred = recycler.recycle_deferred(active, deferred, resources, step=2)
assert len(new_active) == 0
assert len(new_deferred) == 1
class TestGrader:
def _result(self, **kwargs):
base = {
"tasks_completed": 4, "tasks_failed": 1, "total_tasks": 5,
"remaining_time": 3.0, "remaining_energy": 5.0, "remaining_budget": 30.0,
"initial_time": 10.0, "initial_energy": 12.0, "initial_budget": 60.0,
"system_health": 0.9, "steps_used": 7, "max_steps": 10,
}
base.update(kwargs)
return base
def test_score_in_range(self):
for diff in ["easy", "medium", "hard"]:
s = grade(diff, self._result())
assert 0.0 <= s <= 1.0, f"{diff}: {s}"
def test_perfect_score_near_one(self):
perfect = {
"tasks_completed": 10, "tasks_failed": 0, "total_tasks": 10,
"remaining_time": 5.0, "remaining_energy": 8.0, "remaining_budget": 40.0,
"initial_time": 10.0, "initial_energy": 12.0, "initial_budget": 60.0,
"system_health": 1.0, "steps_used": 5, "max_steps": 10,
}
for diff in ["easy", "medium", "hard"]:
s = grade(diff, perfect)
assert s >= 0.6, f"{diff}: {s}"
def test_zero_score_on_all_failed(self):
worst = {
"tasks_completed": 0, "tasks_failed": 10, "total_tasks": 10,
"remaining_time": 0.0, "remaining_energy": 0.0, "remaining_budget": 0.0,
"initial_time": 10.0, "initial_energy": 12.0, "initial_budget": 60.0,
"system_health": 0.0, "steps_used": 10, "max_steps": 10,
}
for diff in ["easy", "medium", "hard"]:
s = grade(diff, worst)
assert s == 0.0, f"{diff}: {s}"
def test_deterministic(self):
result = self._result()
s1 = grade("medium", result)
s2 = grade("medium", result)
assert s1 == s2
def test_hard_collapse_penalty(self):
collapsed = self._result(system_health=0.1)
normal = self._result(system_health=0.8)
assert grade("hard", collapsed) < grade("hard", normal)
class TestEnvironmentReset:
def test_reset_returns_observation(self, easy_env):
obs = easy_env.reset(seed=42)
assert isinstance(obs, AetherTaskFlowObservation)
def test_reset_provides_tasks(self, easy_env):
obs = easy_env.reset(seed=42)
assert len(obs.tasks) > 0
def test_reset_has_full_resources(self, easy_env):
obs = easy_env.reset(seed=42)
assert obs.time_remaining == 10
assert obs.energy_remaining > 0
assert obs.budget_remaining > 0
def test_reset_health_is_one(self, easy_env):
obs = easy_env.reset(seed=42)
assert obs.system_health == 1.0
def test_reset_not_done(self, easy_env):
obs = easy_env.reset(seed=42)
assert obs.done is False
def test_reset_is_reproducible(self, easy_env):
obs1 = easy_env.reset(seed=7)
obs2 = easy_env.reset(seed=7)
assert len(obs1.tasks) == len(obs2.tasks)
assert obs1.tasks[0]["task_id"] == obs2.tasks[0]["task_id"]
def test_reset_names_are_reproducible(self, easy_env):
obs1 = easy_env.reset(seed=42)
obs2 = easy_env.reset(seed=42)
assert [task["name"] for task in obs1.tasks] == [task["name"] for task in obs2.tasks]
def test_reset_episode_id_provided(self, easy_env):
obs = easy_env.reset(seed=1, episode_id="test-ep-001")
assert obs.episode_id == "test-ep-001"
def test_reset_generates_episode_id_if_missing(self, easy_env):
obs = easy_env.reset(seed=1)
assert obs.episode_id is not None
assert len(obs.episode_id) > 0
class TestEnvironmentStep:
def test_execute_reduces_resources(self, easy_env):
obs = easy_env.reset(seed=42)
energy_before = obs.energy_remaining
task = obs.tasks[0]
action = AetherTaskFlowAction(action_type=ActionType.EXECUTE, task_id=task["task_id"])
obs2 = easy_env.step(action)
assert obs2.energy_remaining <= energy_before
def test_execute_valid_task_earns_positive_reward(self, easy_env):
obs = easy_env.reset(seed=42)
task = obs.tasks[0]
action = AetherTaskFlowAction(action_type=ActionType.EXECUTE, task_id=task["task_id"])
obs2 = easy_env.step(action)
# Positive action rewards are normalized into the upper half of [0, 1].
assert obs2.reward is not None
assert 0.5 < obs2.reward <= 1.0
def test_delegate_earns_positive_reward(self, easy_env):
obs = easy_env.reset(seed=42)
task = obs.tasks[0]
action = AetherTaskFlowAction(action_type=ActionType.DELEGATE, task_id=task["task_id"])
obs2 = easy_env.step(action)
assert obs2.reward is not None
assert 0.5 < obs2.reward <= 1.0
def test_defer_earns_negative_reward(self, easy_env):
obs = easy_env.reset(seed=42)
task = obs.tasks[0]
action = AetherTaskFlowAction(action_type=ActionType.DEFER, task_id=task["task_id"])
obs2 = easy_env.step(action)
assert obs2.reward is not None
assert 0.0 <= obs2.reward < 0.5
def test_optimize_returns_small_positive(self, easy_env):
obs = easy_env.reset(seed=42)
task = obs.tasks[0]
action = AetherTaskFlowAction(action_type=ActionType.OPTIMIZE, task_id=task["task_id"])
obs2 = easy_env.step(action)
assert 0.5 < obs2.reward < 0.55
def test_invalid_task_id_penalised(self, easy_env):
obs = easy_env.reset(seed=42)
action = AetherTaskFlowAction(action_type=ActionType.EXECUTE, task_id=9999)
obs2 = easy_env.step(action)
assert obs2.reward is not None
assert 0.0 <= obs2.reward < 0.5
def test_step_reward_is_normalized(self, easy_env):
obs = easy_env.reset(seed=42)
task = obs.tasks[0]
action = AetherTaskFlowAction(action_type=ActionType.EXECUTE, task_id=task["task_id"])
obs2 = easy_env.step(action)
assert 0.0 <= obs2.reward <= 1.0
def test_step_accepts_string_action(self, easy_env):
easy_env.reset(seed=42)
obs = easy_env.step("execute")
assert isinstance(obs, AetherTaskFlowObservation)
assert obs.last_action_type in {"execute", "defer", "delegate", "optimize"}
def test_step_safe_failure_returns_terminal_observation(self, easy_env):
easy_env.reset(seed=42)
obs = easy_env.step({"task_id": "not-an-int"})
assert obs.done is True
assert obs.reward == 0.0
assert "failed safely" in (obs.last_action_outcome or "").lower()
def test_last_action_feedback_populated(self, easy_env):
obs = easy_env.reset(seed=42)
task = obs.tasks[0]
action = AetherTaskFlowAction(action_type=ActionType.EXECUTE, task_id=task["task_id"])
obs2 = easy_env.step(action)
assert obs2.last_action_type == "execute"
assert obs2.last_action_task_id == task["task_id"]
assert obs2.last_action_outcome is not None
def test_step_after_done_returns_done(self, easy_env):
obs = easy_env.reset(seed=42)
# Exhaust all tasks
for _ in range(15):
if obs.done:
break
tasks = obs.tasks
if not tasks:
break
action = AetherTaskFlowAction(
action_type=ActionType.EXECUTE,
task_id=tasks[0]["task_id"]
)
obs = easy_env.step(action)
# Extra step after done should return done
if obs.done:
action = AetherTaskFlowAction(action_type=ActionType.EXECUTE, task_id=0)
obs2 = easy_env.step(action)
assert obs2.done is True
def test_defer_is_visible_in_state(self, easy_env):
obs = easy_env.reset(seed=42)
task = obs.tasks[0]
easy_env.state.resources["energy"] = 0.0
easy_env.state.resources["budget"] = 0.0
easy_env.step(
AetherTaskFlowAction(action_type=ActionType.DEFER, task_id=task["task_id"])
)
assert len(easy_env.state.deferred_tasks) == 1
class TestFullEpisode:
def _run_episode(self, difficulty: str, seed: int = 42) -> dict:
env = AetherTaskFlowEnvironment(difficulty=difficulty)
obs = env.reset(seed=seed)
rewards = []
steps = 0
while not obs.done and steps < 15:
tasks = obs.tasks
if not tasks:
break
task = tasks[0]
action = AetherTaskFlowAction(
action_type=ActionType.EXECUTE,
task_id=task["task_id"],
)
obs = env.step(action)
rewards.append(obs.reward or 0)
steps += 1
score = env.compute_final_score()
return {"score": score, "steps": steps, "rewards": rewards}
def test_easy_episode_completes(self):
result = self._run_episode("easy")
assert result["score"] >= 0.0
assert result["steps"] > 0
def test_medium_episode_completes(self):
result = self._run_episode("medium")
assert result["score"] >= 0.0
def test_hard_episode_completes(self):
result = self._run_episode("hard")
assert result["score"] >= 0.0
def test_score_in_range_all_difficulties(self):
for diff in ["easy", "medium", "hard"]:
result = self._run_episode(diff)
assert 0.0 <= result["score"] <= 1.0, f"{diff}: {result['score']}"
def test_rewards_in_range_all_difficulties(self):
for diff in ["easy", "medium", "hard"]:
result = self._run_episode(diff)
assert all(0.0 <= reward <= 1.0 for reward in result["rewards"]), (
f"{diff}: {result['rewards']}"
)
def test_easy_score_higher_than_hard(self):
easy = self._run_episode("easy")
hard = self._run_episode("hard")
# Easy should generally score higher than hard with naive agent
assert easy["score"] >= hard["score"]
class TestStateProperty:
def test_state_is_aether_state(self, easy_env):
easy_env.reset(seed=42)
state = easy_env.state
assert isinstance(state, AetherTaskFlowState)
def test_state_tracks_steps(self, easy_env):
obs = easy_env.reset(seed=42)
assert easy_env.state.step_count == 0
task = obs.tasks[0]
action = AetherTaskFlowAction(action_type=ActionType.EXECUTE, task_id=task["task_id"])
easy_env.step(action)
assert easy_env.state.step_count == 1
def test_state_tracks_completions(self, easy_env):
obs = easy_env.reset(seed=42)
assert easy_env.state.tasks_completed == 0
task = obs.tasks[0]
action = AetherTaskFlowAction(action_type=ActionType.DELEGATE, task_id=task["task_id"])
easy_env.step(action)
assert easy_env.state.tasks_completed == 1
def test_state_difficulty_matches_env(self):
for diff in ["easy", "medium", "hard"]:
env = AetherTaskFlowEnvironment(difficulty=diff)
env.reset(seed=1)
assert env.state.difficulty == diff
def test_debug_snapshot_is_readable(self, easy_env):
easy_env.reset(seed=42)
snapshot = easy_env._get_obs()
assert snapshot["num_tasks"] > 0
assert "resources" in snapshot
assert "system_health" in snapshot
class TestOpenEnvCompliance:
def test_observation_is_pydantic_model(self, easy_env):
obs = easy_env.reset(seed=1)
assert hasattr(obs, "model_dump")
d = obs.model_dump()
assert isinstance(d, dict)
def test_observation_has_done_field(self, easy_env):
obs = easy_env.reset(seed=1)
assert hasattr(obs, "done")
assert isinstance(obs.done, bool)
def test_observation_has_reward_field(self, easy_env):
obs = easy_env.reset(seed=1)
assert hasattr(obs, "reward")
def test_state_has_episode_id(self, easy_env):
easy_env.reset(seed=1, episode_id="abc-123")
assert easy_env.state.episode_id == "abc-123"
def test_state_has_step_count(self, easy_env):
easy_env.reset(seed=1)
assert hasattr(easy_env.state, "step_count")
def test_invalid_difficulty_raises(self):
with pytest.raises(ValueError):
AetherTaskFlowEnvironment(difficulty="impossible")
def test_action_coerces_freeform_action_type(self):
action = AetherTaskFlowAction(action_type="hi", task_id=0)
assert action.action_type == ActionType.EXECUTE
assert action.task_id == 0
def test_action_extracts_task_id_from_freeform_text(self):
action = AetherTaskFlowAction(action_type="delegate task 3")
assert action.action_type == ActionType.DELEGATE
assert action.task_id == 3
def test_action_accepts_message_payload_shape(self):
action = AetherTaskFlowAction.model_validate({"message": "optimize 2"})
assert action.action_type == ActionType.OPTIMIZE
assert action.task_id == 2
class TestPersistentServerRoutes:
def test_reset_step_state_share_same_session(self, api_client):
reset_response = api_client.post("/reset", json={"seed": 42})
assert reset_response.status_code == 200
reset_payload = reset_response.json()
first_task_id = reset_payload["observation"]["tasks"][0]["task_id"]
episode_id = reset_payload["observation"]["episode_id"]
step_response = api_client.post(
"/step",
json={"action": {"action_type": "execute", "task_id": first_task_id}},
)
assert step_response.status_code == 200
state_response = api_client.get("/state")
assert state_response.status_code == 200
state_payload = state_response.json()
assert state_payload["episode_id"] == episode_id
assert state_payload["step_count"] == 1
def test_step_accepts_message_payload(self, api_client):
api_client.post("/reset", json={"seed": 42})
response = api_client.post("/step", json={"message": "execute"})
assert response.status_code == 200
payload = response.json()
assert payload["observation"]["last_action_type"] in {
"execute",
"defer",
"delegate",
"optimize",
}