ashishMenon05 commited on
Commit
399ee59
·
1 Parent(s): be84c71

fix(backend): resolve AttributeError in episode reset payload

Browse files
Files changed (1) hide show
  1. 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
- "agent_a_model": settings.AGENT_A_MODEL,
38
- "agent_b_model": settings.AGENT_B_MODEL
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()