ba-agent-rl-env / server /ba_agent_environment.py
naman-cen's picture
Upload folder using huggingface_hub
c23cc1b verified
Raw
History Blame Contribute Delete
8.86 kB
"""
BA Agent RL Environment — full BA Toolkit pipeline (5 stages) over one feature.
Pipeline order:
EXTRACT -> INTERVIEW -> GRAPH -> STORY_GEN -> PRD -> FINISH
Reward shape (dense):
every step -> shaped reward in [-0.10, +0.17]
on FINISH -> composite terminal reward in [-0.20, +1.05]
"""
from __future__ import annotations
import json
import random
from typing import Any, Dict, List, Optional
from uuid import uuid4
from openenv.core.env_server.interfaces import Environment
from ba_agent_env.models import (
ActionType,
BAAgentAction,
BAAgentObservation,
BAAgentState,
GeneratedStory,
PipelineStage,
TaskSample,
TaskStatus,
)
from .dataset import get_task_by_id, sample_task
from .reward import step_reward, terminal_reward
_STAGE_FOR_ACTION = {
ActionType.EXTRACT: PipelineStage.AWAITING_EXTRACT,
ActionType.INTERVIEW: PipelineStage.AWAITING_INTERVIEW,
ActionType.GRAPH: PipelineStage.AWAITING_GRAPH,
ActionType.STORY_GEN: PipelineStage.AWAITING_STORY_GEN,
}
_NEXT_STAGE = {
PipelineStage.AWAITING_EXTRACT: PipelineStage.AWAITING_INTERVIEW,
PipelineStage.AWAITING_INTERVIEW: PipelineStage.AWAITING_GRAPH,
PipelineStage.AWAITING_GRAPH: PipelineStage.AWAITING_STORY_GEN,
PipelineStage.AWAITING_STORY_GEN: PipelineStage.READY_TO_FINISH,
}
_STAGE_TO_KEY = {
ActionType.EXTRACT: "extract",
ActionType.INTERVIEW: "interview",
ActionType.GRAPH: "graph",
ActionType.STORY_GEN: "story_gen",
}
def _parse_stories(payload: str) -> List[GeneratedStory]:
text = (payload or "").strip()
if not text:
return []
try:
data = json.loads(text)
except Exception:
return []
if isinstance(data, dict):
data = data.get("stories") or data.get("generated_stories") or data.get("user_stories") or []
if not isinstance(data, list):
return []
out: List[GeneratedStory] = []
for i, s in enumerate(data):
if not isinstance(s, dict):
continue
out.append(
GeneratedStory(
story_id=str(s.get("story_id") or s.get("id") or f"S-{i+1:03d}"),
title=str(s.get("title") or s.get("name") or ""),
description=str(s.get("description") or s.get("narrative") or ""),
acceptance_criteria=s.get("acceptance_criteria") or s.get("ac") or "",
story_points=s.get("story_points"),
state=s.get("state"),
)
)
return out
def _allowed_actions(stage: PipelineStage) -> List[ActionType]:
if stage == PipelineStage.AWAITING_EXTRACT:
return [ActionType.EXTRACT]
if stage == PipelineStage.AWAITING_INTERVIEW:
return [ActionType.INTERVIEW]
if stage == PipelineStage.AWAITING_GRAPH:
return [ActionType.GRAPH]
if stage == PipelineStage.AWAITING_STORY_GEN:
return [ActionType.STORY_GEN]
if stage == PipelineStage.READY_TO_FINISH:
return [ActionType.FINISH]
return []
class BAAgentEnvironment(Environment[BAAgentAction, BAAgentObservation, BAAgentState]):
"""OpenEnv environment wrapping the full BA Toolkit pipeline over one feature."""
def __init__(
self,
task_id: Optional[str] = None,
max_steps: int = 12,
seed: Optional[int] = None,
):
self._task_id_override = task_id
self._max_steps = max_steps
self._rng = random.Random(seed)
self._state: BAAgentState = BAAgentState()
self.reset()
def reset(self, seed: Optional[int] = None, episode_id: Optional[str] = None, **kwargs: Any) -> BAAgentObservation: # type: ignore[override]
if seed is not None:
self._rng = random.Random(seed)
if self._task_id_override:
task = get_task_by_id(self._task_id_override) or sample_task(self._rng)
else:
task = sample_task(self._rng)
self._state = BAAgentState(
task=task,
current_stage=PipelineStage.AWAITING_EXTRACT,
pipeline_outputs={},
generated_stories=[],
step_count=0,
max_steps=self._max_steps,
status=TaskStatus.RUNNING,
last_reward=0.0,
episode_id=episode_id or str(uuid4()),
)
return self._build_obs(feedback="Episode reset. Run EXTRACT to begin.")
def step(self, action: BAAgentAction, timeout_s: Optional[float] = None, **kwargs: Any) -> BAAgentObservation: # type: ignore[override]
state = self._state
state.step_count += 1
task = state.task
assert task is not None
if state.step_count > state.max_steps:
state.status = TaskStatus.TASK_LIMIT_REACHED
return self._build_obs(feedback="Step limit exceeded.", reward=-0.10, done=True, metadata={"status": state.status.value, "step_breakdown": {"step_limit": -0.10}})
if action.action_type == ActionType.FINISH:
if state.current_stage != PipelineStage.READY_TO_FINISH:
return self._build_obs(
feedback=f"FINISH called too early. Current stage: {state.current_stage.value}.",
reward=-0.10, done=False,
metadata={"status": "early_finish", "step_breakdown": {"early_finish": -0.10}},
)
story_gen_present = bool(state.generated_stories)
reward, details = terminal_reward(state.generated_stories, task, story_gen_present=story_gen_present)
state.last_reward = reward
state.status = TaskStatus.COMPLETED
state.current_stage = PipelineStage.DONE
return self._build_obs(
feedback=f"Episode complete. Composite {details['composite_0_to_100']:.1f}/100. Terminal {details['terminal_total']:.3f}.",
reward=reward, done=True, metadata=details,
)
out_of_order = False
expected_stage = _STAGE_FOR_ACTION.get(action.action_type)
if expected_stage is None or state.current_stage != expected_stage:
out_of_order = True
prior_outputs = dict(state.pipeline_outputs)
parsed = _parse_stories(action.payload) if action.action_type == ActionType.STORY_GEN else None
r, bd = step_reward(
action_type=action.action_type.value,
payload=action.payload or "",
task=task,
prior_outputs=prior_outputs,
out_of_order=out_of_order,
parsed_stories=parsed,
)
if out_of_order:
state.status = TaskStatus.INVALID_ACTION_ORDER
return self._build_obs(
feedback=(
f"Invalid action {action.action_type.value} for stage {state.current_stage.value}. "
f"Allowed: {[a.value for a in _allowed_actions(state.current_stage)]}"
),
reward=r, done=False, metadata={"status": state.status.value, "step_breakdown": bd},
)
if action.action_type == ActionType.STORY_GEN and not parsed:
state.status = TaskStatus.INVALID_PAYLOAD
return self._build_obs(
feedback="STORY_GEN payload must be a JSON list of stories.",
reward=r, done=False, metadata={"status": state.status.value, "step_breakdown": bd},
)
key = _STAGE_TO_KEY[action.action_type]
state.pipeline_outputs[key] = action.payload or ""
if action.action_type == ActionType.STORY_GEN:
state.generated_stories = parsed or []
state.status = TaskStatus.RUNNING
state.current_stage = _NEXT_STAGE[state.current_stage]
return self._build_obs(
feedback=f"Stage {action.action_type.value} accepted. Next: {state.current_stage.value}.",
reward=r, done=False, metadata={"status": "ok", "step_breakdown": bd},
)
@property
def state(self) -> BAAgentState: # type: ignore[override]
return self._state
def _build_obs(
self, feedback: str = "", reward: float = 0.0, done: bool = False,
metadata: Optional[Dict[str, Any]] = None,
) -> BAAgentObservation:
s = self._state
task = s.task
assert task is not None
return BAAgentObservation(
task_id=task.task_id,
title=task.title,
description=task.description,
input_documents=task.input_documents,
current_stage=s.current_stage,
available_actions=_allowed_actions(s.current_stage),
pipeline_outputs=dict(s.pipeline_outputs),
step_count=s.step_count,
max_steps=s.max_steps,
feedback=feedback,
status=s.status,
reward=reward,
done=done,
metadata=metadata or {},
)