Spaces:
Sleeping
Sleeping
ashishMenon05 commited on
Commit ·
399ee59
1
Parent(s): be84c71
fix(backend): resolve AttributeError in episode reset payload
Browse files- backend/core/episode_manager.py +93 -94
backend/core/episode_manager.py
CHANGED
|
@@ -1,95 +1,94 @@
|
|
| 1 |
-
import asyncio
|
| 2 |
-
from core.environment import NexusEnvironment
|
| 3 |
-
from api.routes.websocket import broadcast
|
| 4 |
-
|
| 5 |
-
class EpisodeManager:
|
| 6 |
-
"""Manages active episodes and coordinates the WebSocket emissions."""
|
| 7 |
-
def __init__(self):
|
| 8 |
-
self.env = NexusEnvironment()
|
| 9 |
-
self.is_paused = False
|
| 10 |
-
self.simulation_task = None
|
| 11 |
-
|
| 12 |
-
async def reset(self, task: str, custom_scenario: dict = None, seed: int = None, max_steps: int = None, broadcast_episode: bool = True):
|
| 13 |
-
# Cancel any active simulation loop
|
| 14 |
-
if hasattr(self, 'simulation_task') and self.simulation_task and not self.simulation_task.done():
|
| 15 |
-
self.simulation_task.cancel()
|
| 16 |
-
try:
|
| 17 |
-
await self.simulation_task
|
| 18 |
-
except asyncio.CancelledError:
|
| 19 |
-
pass
|
| 20 |
-
self.simulation_task = None
|
| 21 |
-
|
| 22 |
-
obs = await self.env.reset(task=task, custom_scenario=custom_scenario, seed=seed, max_steps=max_steps)
|
| 23 |
-
|
| 24 |
-
if broadcast_episode:
|
| 25 |
-
# Broadcast episode_start
|
| 26 |
-
sc_safe = self.env.active_scenario.copy()
|
| 27 |
-
if "root_cause" in sc_safe: del sc_safe["root_cause"]
|
| 28 |
-
if "correct_fix" in sc_safe: del sc_safe["correct_fix"]
|
| 29 |
-
if "clue_map" in sc_safe: del sc_safe["clue_map"]
|
| 30 |
-
|
| 31 |
-
from config import settings
|
| 32 |
-
await broadcast("episode_start", {
|
| 33 |
-
"episode_id": self.env.active_episode.episode_id,
|
| 34 |
-
"scenario": sc_safe,
|
| 35 |
-
"task": task,
|
| 36 |
-
"difficulty": self.env.active_episode.difficulty,
|
| 37 |
-
"
|
| 38 |
-
|
| 39 |
-
})
|
| 40 |
return obs
|
| 41 |
-
|
| 42 |
-
async def step(self, action):
|
| 43 |
-
obs, reward, done, info = await self.env.step(action)
|
| 44 |
-
|
| 45 |
-
# Broadcast agent message
|
| 46 |
-
await broadcast("agent_message", {
|
| 47 |
-
"agent_id": action.agent_id,
|
| 48 |
-
"message": action.message,
|
| 49 |
-
"step": self.env.active_episode.steps_taken
|
| 50 |
-
})
|
| 51 |
-
|
| 52 |
-
# Broadcast tool calls
|
| 53 |
-
for tc in action.tool_calls:
|
| 54 |
-
await broadcast("tool_call", {
|
| 55 |
-
"agent_id": action.agent_id,
|
| 56 |
-
"tool_name": tc.tool_name,
|
| 57 |
-
"params": tc.params,
|
| 58 |
-
"step": self.env.active_episode.steps_taken
|
| 59 |
-
})
|
| 60 |
-
|
| 61 |
-
# Broadcast tool results
|
| 62 |
-
for tr in obs.tool_results:
|
| 63 |
-
await broadcast("tool_result", {
|
| 64 |
-
"tool_name": tr.tool_name,
|
| 65 |
-
"result": tr.result,
|
| 66 |
-
"success": tr.success,
|
| 67 |
-
"step": self.env.active_episode.steps_taken
|
| 68 |
-
})
|
| 69 |
-
|
| 70 |
-
# Broadcast reward
|
| 71 |
-
await broadcast("reward_update", {
|
| 72 |
-
"agent_id": action.agent_id,
|
| 73 |
-
"reward": reward,
|
| 74 |
-
"breakdown": info.get("breakdown", {}),
|
| 75 |
-
"cumulative": self.env.active_episode.cumulative_reward,
|
| 76 |
-
"step": self.env.active_episode.steps_taken
|
| 77 |
-
})
|
| 78 |
-
|
| 79 |
-
if done:
|
| 80 |
-
await broadcast("episode_end", {
|
| 81 |
-
"episode_id": self.env.active_episode.episode_id,
|
| 82 |
-
"success": info.get("success", False),
|
| 83 |
-
"steps_taken": self.env.active_episode.steps_taken,
|
| 84 |
-
"final_score": info.get("final_score", getattr(self.env.active_episode, "cumulative_reward", 0)),
|
| 85 |
-
"final_breakdown": info.get("breakdown", {}),
|
| 86 |
-
"clues_found": self.env.active_episode.clues_found,
|
| 87 |
-
"root_cause_found": self.env.active_episode.fix_correct,
|
| 88 |
-
"fix_verified": self.env.active_episode.fix_verified,
|
| 89 |
-
"time_taken_seconds": 0,
|
| 90 |
-
"reward_history": self.env.active_episode.reward_history
|
| 91 |
-
})
|
| 92 |
-
|
| 93 |
-
return obs, reward, done, info
|
| 94 |
-
|
| 95 |
-
episode_manager = EpisodeManager()
|
|
|
|
| 1 |
+
import asyncio
|
| 2 |
+
from core.environment import NexusEnvironment
|
| 3 |
+
from api.routes.websocket import broadcast
|
| 4 |
+
|
| 5 |
+
class EpisodeManager:
|
| 6 |
+
"""Manages active episodes and coordinates the WebSocket emissions."""
|
| 7 |
+
def __init__(self):
|
| 8 |
+
self.env = NexusEnvironment()
|
| 9 |
+
self.is_paused = False
|
| 10 |
+
self.simulation_task = None
|
| 11 |
+
|
| 12 |
+
async def reset(self, task: str, custom_scenario: dict = None, seed: int = None, max_steps: int = None, broadcast_episode: bool = True):
|
| 13 |
+
# Cancel any active simulation loop
|
| 14 |
+
if hasattr(self, 'simulation_task') and self.simulation_task and not self.simulation_task.done():
|
| 15 |
+
self.simulation_task.cancel()
|
| 16 |
+
try:
|
| 17 |
+
await self.simulation_task
|
| 18 |
+
except asyncio.CancelledError:
|
| 19 |
+
pass
|
| 20 |
+
self.simulation_task = None
|
| 21 |
+
|
| 22 |
+
obs = await self.env.reset(task=task, custom_scenario=custom_scenario, seed=seed, max_steps=max_steps)
|
| 23 |
+
|
| 24 |
+
if broadcast_episode:
|
| 25 |
+
# Broadcast episode_start
|
| 26 |
+
sc_safe = self.env.active_scenario.copy()
|
| 27 |
+
if "root_cause" in sc_safe: del sc_safe["root_cause"]
|
| 28 |
+
if "correct_fix" in sc_safe: del sc_safe["correct_fix"]
|
| 29 |
+
if "clue_map" in sc_safe: del sc_safe["clue_map"]
|
| 30 |
+
|
| 31 |
+
from config import settings
|
| 32 |
+
await broadcast("episode_start", {
|
| 33 |
+
"episode_id": self.env.active_episode.episode_id,
|
| 34 |
+
"scenario": sc_safe,
|
| 35 |
+
"task": task,
|
| 36 |
+
"difficulty": self.env.active_episode.difficulty,
|
| 37 |
+
"agents": getattr(settings, "AGENTS", [])
|
| 38 |
+
})
|
|
|
|
| 39 |
return obs
|
| 40 |
+
|
| 41 |
+
async def step(self, action):
|
| 42 |
+
obs, reward, done, info = await self.env.step(action)
|
| 43 |
+
|
| 44 |
+
# Broadcast agent message
|
| 45 |
+
await broadcast("agent_message", {
|
| 46 |
+
"agent_id": action.agent_id,
|
| 47 |
+
"message": action.message,
|
| 48 |
+
"step": self.env.active_episode.steps_taken
|
| 49 |
+
})
|
| 50 |
+
|
| 51 |
+
# Broadcast tool calls
|
| 52 |
+
for tc in action.tool_calls:
|
| 53 |
+
await broadcast("tool_call", {
|
| 54 |
+
"agent_id": action.agent_id,
|
| 55 |
+
"tool_name": tc.tool_name,
|
| 56 |
+
"params": tc.params,
|
| 57 |
+
"step": self.env.active_episode.steps_taken
|
| 58 |
+
})
|
| 59 |
+
|
| 60 |
+
# Broadcast tool results
|
| 61 |
+
for tr in obs.tool_results:
|
| 62 |
+
await broadcast("tool_result", {
|
| 63 |
+
"tool_name": tr.tool_name,
|
| 64 |
+
"result": tr.result,
|
| 65 |
+
"success": tr.success,
|
| 66 |
+
"step": self.env.active_episode.steps_taken
|
| 67 |
+
})
|
| 68 |
+
|
| 69 |
+
# Broadcast reward
|
| 70 |
+
await broadcast("reward_update", {
|
| 71 |
+
"agent_id": action.agent_id,
|
| 72 |
+
"reward": reward,
|
| 73 |
+
"breakdown": info.get("breakdown", {}),
|
| 74 |
+
"cumulative": self.env.active_episode.cumulative_reward,
|
| 75 |
+
"step": self.env.active_episode.steps_taken
|
| 76 |
+
})
|
| 77 |
+
|
| 78 |
+
if done:
|
| 79 |
+
await broadcast("episode_end", {
|
| 80 |
+
"episode_id": self.env.active_episode.episode_id,
|
| 81 |
+
"success": info.get("success", False),
|
| 82 |
+
"steps_taken": self.env.active_episode.steps_taken,
|
| 83 |
+
"final_score": info.get("final_score", getattr(self.env.active_episode, "cumulative_reward", 0)),
|
| 84 |
+
"final_breakdown": info.get("breakdown", {}),
|
| 85 |
+
"clues_found": self.env.active_episode.clues_found,
|
| 86 |
+
"root_cause_found": self.env.active_episode.fix_correct,
|
| 87 |
+
"fix_verified": self.env.active_episode.fix_verified,
|
| 88 |
+
"time_taken_seconds": 0,
|
| 89 |
+
"reward_history": self.env.active_episode.reward_history
|
| 90 |
+
})
|
| 91 |
+
|
| 92 |
+
return obs, reward, done, info
|
| 93 |
+
|
| 94 |
+
episode_manager = EpisodeManager()
|