openenv-claw1 / app /routes.py
vishal harkal
Fix suite success/reward metrics and baseline payload
4138efa
Raw History Blame Contribute Delete
11.8 kB
from __future__ import annotations
from threading import RLock
from typing import Any, Dict, List, Optional
from fastapi import APIRouter, Body, HTTPException
from pydantic import BaseModel, Field, ValidationError
from baseline.baseline_agent import BaselineGreedyAgent
from env.environment import LastMileDeliveryEnvironment
from env.models import Action, EnvironmentConfig, EnvironmentState, Observation, ScenarioDefinition, StepResult
from grader.grader import DeliveryEpisodeGrader, GradeReport
from tasks.registry import get_all_task_metadata, get_task_config, get_task_definition
router = APIRouter()
class ResetRequest(BaseModel):
task: str = "easy"
seed: int | None = None
class ScenarioResetRequest(BaseModel):
seed: int | None = None
success_condition: Dict[str, Any] = Field(default_factory=dict)
scenario: ScenarioDefinition
class GraderResponse(BaseModel):
task: str | None
steps_recorded: int
report: GradeReport
class BaselineResponse(BaseModel):
task: str
seed: int | None
agent: str
steps_executed: int = Field(..., ge=0)
total_reward: float
done: bool
success: bool
score: float = Field(..., ge=0.0, le=1.0)
report: GradeReport
class TaskDescriptor(BaseModel):
name: str
description: str
difficulty: int
parameters: Dict[str, Any]
success_condition: Dict[str, Any]
class TasksResponse(BaseModel):
tasks: List[TaskDescriptor]
action_schema: Dict[str, Any]
TASKS_ACTION_SCHEMA: Dict[str, Any] = {
"type": "object",
"required": ["move", "accept_order", "deliver_order", "wait"],
"additionalProperties": False,
"properties": {
"move": {
"type": ["string", "null"],
"enum": ["up", "down", "left", "right", "stay", None],
},
"accept_order": {
"type": ["integer", "null"],
},
"deliver_order": {
"type": "boolean",
},
"wait": {
"type": "boolean",
},
},
}
class EnvironmentService:
BASELINE_MAX_STEPS = 256
def __init__(self) -> None:
self._lock = RLock()
self._env: Optional[LastMileDeliveryEnvironment] = None
self._task: Optional[str] = None
self._success_condition: Dict[str, Any] = {}
self._trajectory: list[StepResult] = []
self._grader = DeliveryEpisodeGrader()
self._baseline_agent = BaselineGreedyAgent()
def reset(self, task: str, seed: int | None) -> Observation:
with self._lock:
config = get_task_config(task_name=task, seed=seed)
task_definition = get_task_definition(task)
self._env = LastMileDeliveryEnvironment(config=config)
self._task = task
self._success_condition = dict(task_definition.get("success_condition", {}))
self._trajectory = []
return self._env.reset(seed=seed)
def reset_from_scenario(
self,
scenario: ScenarioDefinition,
seed: int | None,
success_condition: Dict[str, Any] | None,
) -> Observation:
with self._lock:
config = EnvironmentConfig(
width=scenario.width,
height=scenario.height,
max_steps=scenario.max_steps,
max_orders=max(1, len(scenario.orders)),
obstacle_density=0.0,
traffic_density=0.0,
dynamic_obstacles_enabled=bool(scenario.dynamic_obstacles),
dynamic_obstacle_ratio=1.0 if scenario.dynamic_obstacles else 0.0,
battery_enabled=scenario.battery_profile.enabled,
battery_capacity=scenario.battery_profile.capacity,
battery_recharge_rate=scenario.battery_profile.recharge_rate,
charging_stations=[item.model_copy(deep=True) for item in scenario.charging_stations],
seed=seed,
)
self._env = LastMileDeliveryEnvironment(config=config)
self._task = "scenario"
self._success_condition = dict(success_condition or {})
self._trajectory = []
return self._env.reset_with_scenario(scenario=scenario, seed=seed)
def step(self, action: Action) -> StepResult:
with self._lock:
if self._env is None:
raise RuntimeError("Environment not initialized. Call /reset first.")
result = self._env.step_result(action)
self._trajectory.append(result)
return result
def state(self) -> EnvironmentState:
with self._lock:
if self._env is None:
return EnvironmentState(
initialized=False,
done=False,
total_reward=0.0,
step_count=0,
delivered_orders=0,
current_order_id=None,
observation=None,
)
return self._env.state()
def current_task(self) -> Optional[str]:
with self._lock:
return self._task
def grade(self) -> GradeReport:
with self._lock:
if self._env is None:
raise RuntimeError("Environment not initialized. Call /reset first.")
final_observation = self._env.current_observation()
success_condition = dict(self._success_condition)
if not success_condition:
resolved_task = self._task or "easy"
try:
task_definition = get_task_definition(resolved_task)
success_condition = dict(task_definition.get("success_condition", {}))
except ValueError:
success_condition = {}
return self._grader.grade_episode(
self._trajectory,
final_observation,
success_condition=success_condition,
)
def baseline_rollout(
self,
task: str | None = None,
seed: int | None = None,
max_steps: int | None = None,
) -> BaselineResponse:
with self._lock:
resolved_task = task or self._task or "easy"
config = get_task_config(task_name=resolved_task, seed=seed)
rollout_env = LastMileDeliveryEnvironment(config=config)
observation = rollout_env.reset(seed=seed)
step_cap = config.max_steps
if max_steps is not None:
step_cap = min(step_cap, max_steps)
step_cap = max(1, min(step_cap, self.BASELINE_MAX_STEPS))
trajectory: list[StepResult] = []
for _ in range(step_cap):
action = self._baseline_agent.act(observation)
result = rollout_env.step_result(action)
trajectory.append(result)
observation = result.observation
if result.done:
break
task_definition = get_task_definition(resolved_task)
success_condition = task_definition.get("success_condition", {})
report = self._grader.grade_episode(
trajectory,
rollout_env.current_observation(),
success_condition=success_condition,
)
final_state = rollout_env.state()
return BaselineResponse(
task=resolved_task,
seed=seed,
agent="BaselineGreedyAgent",
steps_executed=len(trajectory),
total_reward=final_state.total_reward,
done=final_state.done,
success=report.success,
score=report.score,
report=report,
)
def steps_recorded(self) -> int:
with self._lock:
return len(self._trajectory)
def _http_error(status_code: int, message: str) -> HTTPException:
return HTTPException(
status_code=status_code,
detail={
"error": "bad_request" if status_code == 400 else "internal_server_error",
"message": message,
},
)
def _parse_step_action(payload: Any) -> Action:
if payload is None:
raise ValueError("Missing step payload")
if not isinstance(payload, dict):
raise ValueError("Step payload must be a JSON object")
allowed_keys = {"move", "accept_order", "deliver_order", "wait"}
unexpected_keys = sorted(set(payload.keys()) - allowed_keys)
if unexpected_keys:
raise ValueError(f"Unexpected action fields: {', '.join(unexpected_keys)}")
missing_keys = sorted(allowed_keys - set(payload.keys()))
if missing_keys:
raise ValueError(f"Missing action fields: {', '.join(missing_keys)}")
try:
return Action.model_validate(payload)
except ValidationError as exc:
raise ValueError("Invalid action payload") from exc
def _parse_scenario_reset_request(payload: Any) -> ScenarioResetRequest:
if payload is None:
raise ValueError("Missing scenario reset payload")
if not isinstance(payload, dict):
raise ValueError("Scenario reset payload must be a JSON object")
try:
return ScenarioResetRequest.model_validate(payload)
except ValidationError as exc:
raise ValueError("Invalid scenario payload") from exc
def _normalize_task_metadata(raw_tasks: List[Dict[str, Any]]) -> List[TaskDescriptor]:
normalized: List[TaskDescriptor] = []
for item in raw_tasks:
missing = [
field
for field in ("name", "description", "difficulty", "parameters", "success_condition")
if field not in item
]
if missing:
raise ValueError(f"Task metadata missing required fields: {', '.join(missing)}")
normalized.append(
TaskDescriptor(
name=str(item["name"]),
description=str(item["description"]),
difficulty=int(item["difficulty"]),
parameters=dict(item["parameters"]),
success_condition=dict(item["success_condition"]),
)
)
return normalized
env_service = EnvironmentService()
@router.get("/")
def root() -> dict:
return {
"name": "last-mile-delivery-optimization",
"status": "ok",
"health": "/health",
"docs": "/docs",
}
@router.get("/health")
def health() -> dict:
return {"status": "ok"}
@router.get("/tasks", response_model=TasksResponse)
def list_tasks() -> TasksResponse:
try:
raw_tasks = get_all_task_metadata()
tasks = _normalize_task_metadata(raw_tasks)
return TasksResponse(tasks=tasks, action_schema=TASKS_ACTION_SCHEMA)
except ValueError as exc:
raise _http_error(500, str(exc)) from exc
except Exception as exc:
raise _http_error(500, "Failed to retrieve task metadata") from exc
@router.post("/reset", response_model=Observation)
def reset_environment(request: ResetRequest | None = Body(default=None)) -> Observation:
try:
payload = request or ResetRequest()
return env_service.reset(task=payload.task, seed=payload.seed)
except ValueError as exc:
raise _http_error(400, str(exc)) from exc
except Exception as exc:
raise _http_error(500, "Failed to reset environment") from exc
@router.post("/reset_from_scenario", response_model=Observation)
def reset_environment_from_scenario(payload: Any = Body(default=None)) -> Observation:
try:
request = _parse_scenario_reset_request(payload)
return env_service.reset_from_scenario(
scenario=request.scenario,
seed=request.seed,
success_condition=request.success_condition,
)
except ValueError as exc:
raise _http_error(400, str(exc)) from exc
except Exception as exc:
raise _http_error(500, "Failed to reset environment from scenario") from exc
@router.post("/step", response_model=StepResult)
def environment_step(payload: Any = Body(default=None)) -> StepResult:
try:
action = _parse_step_action(payload)
return env_service.step(action=action)
except ValueError as exc:
raise _http_error(400, str(exc)) from exc
except RuntimeError as exc:
raise _http_error(400, str(exc)) from exc
except Exception as exc:
raise _http_error(500, "Failed to execute step") from exc
@router.get("/state", response_model=EnvironmentState)
def environment_state() -> EnvironmentState:
try:
return env_service.state()
except RuntimeError as exc:
raise _http_error(400, str(exc)) from exc
except Exception as exc:
raise _http_error(500, "Failed to retrieve state") from exc
@router.get("/grader", response_model=GraderResponse)
def environment_grader() -> GraderResponse:
try:
report = env_service.grade()
return GraderResponse(
task=env_service.current_task(),
steps_recorded=env_service.steps_recorded(),
report=report,
)
except RuntimeError as exc:
raise _http_error(400, str(exc)) from exc
except Exception as exc:
raise _http_error(500, "Failed to compute grade") from exc
@router.get("/baseline", response_model=BaselineResponse)
def environment_baseline(task: str | None = None, seed: int | None = None, max_steps: int | None = None) -> BaselineResponse:
try:
return env_service.baseline_rollout(task=task, seed=seed, max_steps=max_steps)
except ValueError as exc:
raise _http_error(400, str(exc)) from exc
except RuntimeError as exc:
raise _http_error(400, str(exc)) from exc
except Exception as exc:
raise _http_error(500, "Failed to compute baseline rollout") from exc