from __future__ import annotations from dataclasses import dataclass from typing import Any, Dict, List, Sequence from pydantic import BaseModel, Field from env.models import Observation, Order, OrderPriority, StepResult from grader.metrics import ( clamp, round_score, ) DEFAULT_SUCCESS_CONDITION: Dict[str, Any] = { "completion_rate_min": 1.0, "max_steps": 200, "invalid_action_rate_max": 0.10, } DEFAULT_HIGH_PRIORITY_DEADLINE = 12 class GradeReport(BaseModel): score: float = Field(..., ge=0.0, le=1.0) success: bool delivered_orders: int = Field(..., ge=0) total_orders: int = Field(..., ge=0) high_priority_total_orders: int = Field(..., ge=0) high_priority_delivered_orders: int = Field(..., ge=0) high_priority_on_time_deliveries: int = Field(..., ge=0) steps_taken: int = Field(..., ge=0) max_steps_target: int = Field(..., ge=0) optimal_steps: int = Field(..., ge=0) completion_rate: float = Field(..., ge=0.0, le=1.0) high_priority_on_time_rate: float = Field(..., ge=0.0, le=1.0) efficiency_ratio: float = Field(..., ge=0.0, le=1.0) invalid_action_rate: float = Field(..., ge=0.0, le=1.0) completion_component: float = Field(..., ge=0.0, le=1.0) priority_component: float = Field(..., ge=0.0, le=1.0) efficiency_component: float = Field(..., ge=0.0, le=1.0) penalty_component: float = Field(..., ge=0.0, le=1.0) invalid_actions: int = Field(..., ge=0) delay_events: int = Field(..., ge=0) battery_depletion_events: int = Field(..., ge=0) no_progress_events: int = Field(..., ge=0) battery_remaining: int | None = None success_condition_used: Dict[str, Any] = Field(default_factory=dict) scoring_logic: str safeguards_applied: List[str] = Field(default_factory=list) edge_case_handling: List[str] = Field(default_factory=list) @dataclass(frozen=True) class EpisodeStats: steps_taken: int invalid_actions: int delay_events: int battery_depletion_events: int no_progress_events: int @dataclass(frozen=True) class SuccessTargets: completion_target: float priority_target: float | None max_steps_target: int invalid_action_rate_max: float require_no_battery_depletion: bool high_priority_deadline_steps: int class DeliveryEpisodeGrader: """Deterministic task-aware grader aligned with task success_condition metrics.""" def __init__(self) -> None: self.base_weights = { "completion": 0.55, "priority": 0.25, "efficiency": 0.20, } def grade_episode( self, trajectory: Sequence[StepResult], final_observation: Observation, success_condition: Dict[str, Any] | None = None, ) -> GradeReport: order_catalog = self._collect_order_catalog(trajectory, final_observation) delivered_steps = self._collect_delivered_steps(trajectory) delivered_ids = set(delivered_steps.keys()) delivered_orders = len(delivered_ids) remaining_ids = {order.order_id for order in final_observation.pending_orders} if final_observation.current_order is not None: remaining_ids.add(final_observation.current_order.order_id) known_ids = set(order_catalog.keys()) all_ids = known_ids.union(delivered_ids).union(remaining_ids) total_orders = len(all_ids) episode_stats = self._collect_episode_stats(trajectory) edge_case_notes: List[str] = [] safeguards: List[str] = [] targets = self._resolve_success_targets( success_condition=success_condition, fallback_max_steps=max(1, final_observation.step_count), edge_case_notes=edge_case_notes, ) if episode_stats.steps_taken == 0: edge_case_notes.append("No steps were recorded for this episode.") if total_orders == 0: completion_rate_value = 1.0 edge_case_notes.append("No orders were present; completion_rate set to 1.0 by convention.") else: completion_rate_value = delivered_orders / total_orders ( high_priority_total, high_priority_delivered, high_priority_on_time, high_priority_on_time_rate, ) = self._high_priority_metrics( order_catalog=order_catalog, delivered_steps=delivered_steps, deadline_steps=targets.high_priority_deadline_steps, edge_case_notes=edge_case_notes, ) efficiency_ratio = self._efficiency_ratio( steps_taken=episode_stats.steps_taken, max_steps_target=targets.max_steps_target, ) invalid_action_rate = self._invalid_action_rate( invalid_actions=episode_stats.invalid_actions, steps_taken=episode_stats.steps_taken, ) completion_component = self._target_progress( value=completion_rate_value, target=targets.completion_target, edge_case_notes=edge_case_notes, metric_name="completion_rate", ) if targets.priority_target is None: priority_component = 1.0 edge_case_notes.append("high_priority_on_time_rate_min not set; priority component treated as neutral.") else: priority_component = self._target_progress( value=high_priority_on_time_rate, target=targets.priority_target, edge_case_notes=edge_case_notes, metric_name="high_priority_on_time_rate", ) efficiency_component = efficiency_ratio penalty_component = self._invalid_action_penalty_component( invalid_action_rate=invalid_action_rate, invalid_action_rate_max=targets.invalid_action_rate_max, edge_case_notes=edge_case_notes, ) if targets.require_no_battery_depletion and episode_stats.battery_depletion_events > 0: penalty_component *= 0.5 safeguards.append("battery_depletion_penalty_applied") completion_weight = self.base_weights["completion"] priority_weight = self.base_weights["priority"] if targets.priority_target is not None else 0.0 efficiency_weight = self.base_weights["efficiency"] if targets.priority_target is None: completion_weight += 0.15 efficiency_weight += 0.10 raw_base = ( completion_weight * completion_component + priority_weight * priority_component + efficiency_weight * efficiency_component ) raw_score = raw_base * penalty_component score = round_score(raw_score, decimals=4) meets_completion = completion_rate_value >= targets.completion_target meets_priority = True if targets.priority_target is not None: meets_priority = high_priority_on_time_rate >= targets.priority_target meets_steps = episode_stats.steps_taken <= targets.max_steps_target meets_invalid_actions = invalid_action_rate <= targets.invalid_action_rate_max meets_battery_rule = (not targets.require_no_battery_depletion) or (episode_stats.battery_depletion_events == 0) success = all( [ meets_completion, meets_priority, meets_steps, meets_invalid_actions, meets_battery_rule, ] ) logic = ( "metrics follow success_condition keys: completion_rate(_min), " "high_priority_on_time_rate_min, max_steps, invalid_action_rate_max. " "components: completion (highest), priority SLA (medium when configured), " "efficiency from steps/max_steps (lower), and penalties reduce via invalid_action_rate." ) if not edge_case_notes: edge_case_notes.append("No edge-case adjustments were needed.") success_condition_used = { "completion_rate_target": targets.completion_target, "high_priority_on_time_rate_target": targets.priority_target, "max_steps": targets.max_steps_target, "invalid_action_rate_max": targets.invalid_action_rate_max, "battery_depletion": not targets.require_no_battery_depletion, "high_priority_deadline_steps": targets.high_priority_deadline_steps, } return GradeReport( score=score, success=success, delivered_orders=delivered_orders, total_orders=total_orders, high_priority_total_orders=high_priority_total, high_priority_delivered_orders=high_priority_delivered, high_priority_on_time_deliveries=high_priority_on_time, steps_taken=episode_stats.steps_taken, max_steps_target=targets.max_steps_target, optimal_steps=targets.max_steps_target, completion_rate=completion_rate_value, high_priority_on_time_rate=high_priority_on_time_rate, efficiency_ratio=efficiency_ratio, invalid_action_rate=invalid_action_rate, completion_component=completion_component, priority_component=priority_component, efficiency_component=efficiency_component, penalty_component=penalty_component, invalid_actions=episode_stats.invalid_actions, delay_events=episode_stats.delay_events, battery_depletion_events=episode_stats.battery_depletion_events, no_progress_events=episode_stats.no_progress_events, battery_remaining=final_observation.battery_level, success_condition_used=success_condition_used, scoring_logic=logic, safeguards_applied=safeguards, edge_case_handling=edge_case_notes, ) def _resolve_success_targets( self, success_condition: Dict[str, Any] | None, fallback_max_steps: int, edge_case_notes: List[str], ) -> SuccessTargets: condition = dict(DEFAULT_SUCCESS_CONDITION) if success_condition is not None: condition.update(success_condition) completion_raw = condition.get("completion_rate") if completion_raw is None: completion_raw = condition.get("completion_rate_min", 1.0) completion_target = self._safe_ratio_target( raw_value=completion_raw, default=1.0, metric_name="completion_rate_target", edge_case_notes=edge_case_notes, ) priority_raw = condition.get("high_priority_on_time_rate_min") priority_target: float | None = None if priority_raw is not None: priority_target = self._safe_ratio_target( raw_value=priority_raw, default=1.0, metric_name="high_priority_on_time_rate_target", edge_case_notes=edge_case_notes, ) max_steps_raw = condition.get("max_steps", fallback_max_steps) try: max_steps_target = int(max_steps_raw) except Exception: max_steps_target = fallback_max_steps edge_case_notes.append("Invalid max_steps in success_condition; fallback max_steps used.") if max_steps_target <= 0: max_steps_target = max(1, fallback_max_steps) edge_case_notes.append("Non-positive max_steps in success_condition; fallback max_steps used.") invalid_rate_raw = condition.get("invalid_action_rate_max", 0.10) invalid_action_rate_max = self._safe_ratio_target( raw_value=invalid_rate_raw, default=0.10, metric_name="invalid_action_rate_max", edge_case_notes=edge_case_notes, ) battery_depletion_rule = condition.get("battery_depletion") require_no_battery_depletion = battery_depletion_rule is False deadline_raw = condition.get("high_priority_deadline_steps", DEFAULT_HIGH_PRIORITY_DEADLINE) try: deadline_steps = int(deadline_raw) except Exception: deadline_steps = DEFAULT_HIGH_PRIORITY_DEADLINE edge_case_notes.append("Invalid high_priority_deadline_steps; default deadline used.") if deadline_steps <= 0: deadline_steps = DEFAULT_HIGH_PRIORITY_DEADLINE edge_case_notes.append("Non-positive high_priority_deadline_steps; default deadline used.") return SuccessTargets( completion_target=completion_target, priority_target=priority_target, max_steps_target=max_steps_target, invalid_action_rate_max=invalid_action_rate_max, require_no_battery_depletion=require_no_battery_depletion, high_priority_deadline_steps=deadline_steps, ) def _safe_ratio_target( self, raw_value: Any, default: float, metric_name: str, edge_case_notes: List[str], ) -> float: try: parsed = float(raw_value) except Exception: edge_case_notes.append(f"Invalid {metric_name}; default value used.") return default return clamp(parsed, 0.0, 1.0) def _collect_episode_stats(self, trajectory: Sequence[StepResult]) -> EpisodeStats: return EpisodeStats( steps_taken=len(trajectory), invalid_actions=sum(1 for item in trajectory if item.info.invalid_action), delay_events=sum(1 for item in trajectory if item.info.delay_penalty_applied), battery_depletion_events=sum(1 for item in trajectory if item.info.battery_depleted), no_progress_events=sum(1 for item in trajectory if item.info.made_progress is False), ) def _collect_delivered_steps(self, trajectory: Sequence[StepResult]) -> Dict[str, int]: delivered_steps: Dict[str, int] = {} for index, item in enumerate(trajectory, start=1): delivered_id = item.info.delivered_order_id if delivered_id is None: continue delivered_step = item.observation.step_count if item.observation.step_count > 0 else index delivered_steps[delivered_id] = delivered_step return delivered_steps def _high_priority_metrics( self, order_catalog: Dict[str, Order], delivered_steps: Dict[str, int], deadline_steps: int, edge_case_notes: List[str], ) -> tuple[int, int, int, float]: high_priority_orders = [ order for order in order_catalog.values() if order.priority == OrderPriority.HIGH ] high_priority_total = len(high_priority_orders) if high_priority_total == 0: edge_case_notes.append("No high-priority orders were present; high-priority SLA treated as 1.0.") return 0, 0, 0, 1.0 high_priority_delivered = 0 high_priority_on_time = 0 for order in high_priority_orders: delivered_step = delivered_steps.get(order.order_id) if delivered_step is None: continue high_priority_delivered += 1 if order.accepted_step is None: continue if delivered_step - order.accepted_step <= deadline_steps: high_priority_on_time += 1 on_time_rate = high_priority_on_time / high_priority_total return high_priority_total, high_priority_delivered, high_priority_on_time, on_time_rate def _efficiency_ratio(self, steps_taken: int, max_steps_target: int) -> float: if max_steps_target <= 0: return 1.0 if steps_taken <= 0: return 1.0 return clamp((max_steps_target - steps_taken) / max_steps_target, 0.0, 1.0) def _invalid_action_rate(self, invalid_actions: int, steps_taken: int) -> float: if steps_taken <= 0: return 0.0 return clamp(invalid_actions / steps_taken, 0.0, 1.0) def _target_progress( self, value: float, target: float, edge_case_notes: List[str], metric_name: str, ) -> float: if target <= 0.0: edge_case_notes.append(f"{metric_name} target was 0.0; component treated as 1.0.") return 1.0 return clamp(value / target, 0.0, 1.0) def _invalid_action_penalty_component( self, invalid_action_rate: float, invalid_action_rate_max: float, edge_case_notes: List[str], ) -> float: if invalid_action_rate_max <= 0.0: if invalid_action_rate > 0.0: edge_case_notes.append( "invalid_action_rate_max is 0.0 and invalid actions occurred; penalty set to 0.0." ) return 0.0 return 1.0 normalized = invalid_action_rate / invalid_action_rate_max if invalid_action_rate <= invalid_action_rate_max: return clamp(1.0 - 0.25 * normalized, 0.75, 1.0) overflow = (invalid_action_rate - invalid_action_rate_max) / max(1.0 - invalid_action_rate_max, 1e-9) return clamp(0.75 - 0.75 * overflow, 0.0, 0.75) if total_orders == 0: edge_case_notes.append("No orders were present; completion is evaluated as neutral.") if penalty_stats.steps_taken == 0 and total_orders > 0: edge_case_notes.append("No steps were recorded for a non-empty episode.") unknown_orders = max(0, total_orders - len(order_catalog)) if unknown_orders > 0: edge_case_notes.append( "Some order geometries were missing from observations; overhead-only fallback used." ) ordered_orders = [order_catalog[key] for key in sorted(order_catalog.keys())] optimal_steps_known = optimal_steps_single_agent( orders=ordered_orders, start_location=self.start_location, ) optimal_steps = optimal_steps_known + (unknown_orders * 3) if total_orders > 0 and optimal_steps == 0: optimal_steps = max(1, penalty_stats.steps_taken) safeguards.append("optimal_steps_zero_guard") delivered_weight, total_weight = self._weighted_completion_masses( delivered_ids=delivered_ids, all_ids=all_ids, order_catalog=order_catalog, ) completion, efficiency_component, penalty_component = self._compute_components( delivered_weight=delivered_weight, total_weight=total_weight, penalty_stats=penalty_stats, optimal_steps=optimal_steps, ) raw_score = self._combine_components( completion=completion, efficiency_component=efficiency_component, penalty_component=penalty_component, ) raw_score = self._apply_safeguards( raw_score=raw_score, safeguards=safeguards, total_orders=total_orders, delivered_orders=delivered_orders, steps_taken=penalty_stats.steps_taken, invalid_actions=penalty_stats.invalid_actions, delay_events=penalty_stats.delay_events, battery_depletion_events=penalty_stats.battery_depletion_events, no_progress_events=penalty_stats.no_progress_events, completion=completion, efficiency_component=efficiency_component, penalty_component=penalty_component, ) score = round_score(raw_score, decimals=4) logic = ( "score = 0.50*completion + 0.30*efficiency + 0.20*penalty_quality; " "completion is weighted by order priority (high>low), " "efficiency uses optimal_steps/steps_taken and is gated by completion, " "penalty_quality decreases with invalid, delay, battery-depletion, and sustained no-progress events." ) if not edge_case_notes: edge_case_notes.append("No edge-case adjustments were needed.") return GradeReport( score=score, delivered_orders=delivered_orders, total_orders=total_orders, steps_taken=penalty_stats.steps_taken, optimal_steps=optimal_steps, completion_component=completion, efficiency_component=efficiency_component, penalty_component=penalty_component, invalid_actions=penalty_stats.invalid_actions, delay_events=penalty_stats.delay_events, battery_depletion_events=penalty_stats.battery_depletion_events, no_progress_events=penalty_stats.no_progress_events, battery_remaining=final_observation.battery_level, scoring_logic=logic, safeguards_applied=safeguards, edge_case_handling=edge_case_notes, ) def _collect_order_catalog( self, trajectory: Sequence[StepResult], final_observation: Observation, ) -> Dict[str, Order]: orders: Dict[str, Order] = {} for item in trajectory: for order in item.observation.pending_orders: orders[order.order_id] = order if item.observation.current_order is not None: orders[item.observation.current_order.order_id] = item.observation.current_order for order in final_observation.pending_orders: orders[order.order_id] = order if final_observation.current_order is not None: orders[final_observation.current_order.order_id] = final_observation.current_order return orders