dhruv-punia-bits/memory-compaction-openenv
0
1from __future__ import annotations2 3import uuid4from typing import Any5 6from server.grader import score_step, score_terminal7from server.models import EpisodeState, MemoryItem, MemoryStatus, Observation, Operation, ResetRequest, StepAction, StepResponse, TaskDefinition8from server.tasks import build_episode_payload, build_task_definitions9 10 11class MemoryCompactionEnvironment:12 def __init__(self) -> None:13 self._current_state: EpisodeState | None = None14 15 def list_tasks(self) -> list[TaskDefinition]:16 return build_task_definitions()17 18 def reset(self, request: ResetRequest) -> Observation:19 payload = build_episode_payload(request.difficulty, request.seed)20 self._current_state = EpisodeState(21 episode_id=f"{request.difficulty.value}-{payload['seed']}-{uuid.uuid4().hex[:8]}",22 difficulty=request.difficulty,23 seed=int(payload["seed"]),24 phase="ingest",25 current_index=0,26 token_budget=int(payload["token_budget"]),27 step_count=0,28 recent_turn_window=int(payload["recent_turn_window"]),29 turns=list(payload["turns"]),30 future_queries=list(payload["future_queries"]),31 hidden_gold_memories=list(payload["gold_memories"]),32 hidden_new_memory_ids=[],33 hidden_ambiguous_turn_ids=list(payload.get("ambiguous_turn_ids", [])),34 hidden_low_value_turn_ids=list(payload.get("low_value_turn_ids", [])),35 )36 self._populate_hidden_new_memory_ids()37 return self._build_observation()38 39 def state(self) -> EpisodeState:40 if self._current_state is None:41 raise ValueError("No active episode. Call reset first.")42 return self._current_state43 44 def step(self, action: StepAction) -> StepResponse:45 state = self.state()46 if state.done:47 terminal = score_terminal(state)48 state.latest_reward_breakdown = terminal49 return StepResponse(50 observation=self._build_observation(),51 reward=terminal.total_score,52 done=True,53 reward_breakdown=terminal,54 metadata={"phase": state.phase},55 )56 57 self._apply_action(action)58 state.step_count += 159 step_breakdown = score_step(state, action)60 reward = step_breakdown.total_score61 62 if state.phase == "ingest":63 state.current_index += 164 if state.current_index >= len(state.turns):65 state.phase = "evaluation"66 self._populate_hidden_new_memory_ids()67 68 if state.phase == "evaluation":69 terminal = score_terminal(state)70 reward = round((reward * 0.35) + (terminal.total_score * 0.65), 6)71 step_breakdown.terminal_score = terminal.total_score72 step_breakdown.total_score = reward73 state.done = True74 75 state.latest_reward_breakdown = step_breakdown76 final_reward = round(_strict_reward(reward), 6)77 step_breakdown.total_score = final_reward78 return StepResponse(79 observation=self._build_observation(),80 reward=final_reward,81 done=state.done,82 reward_breakdown=step_breakdown,83 metadata={84 "phase": state.phase,85 "token_usage": state.token_usage(),86 "token_budget": state.token_budget,87 },88 )89 90 def _apply_action(self, action: StepAction) -> None:91 state = self.state()92 if action.summary_text is not None:93 state.working_summary = self._truncate_summary(action.summary_text)94 95 if action.operation == Operation.APPEND_MEMORY:96 for item in action.memory_items:97 self._upsert_memory(item, preserve_history=True)98 elif action.operation == Operation.UPDATE_MEMORY:99 for incoming in action.memory_items:100 self._upsert_memory(incoming, preserve_history=False)101 elif action.operation == Operation.DELETE_MEMORY:102 ids_to_delete = {item.memory_id for item in action.memory_items}103 state.durable_memory = [item for item in state.durable_memory if item.memory_id not in ids_to_delete]104 elif action.operation == Operation.REPLACE_SUMMARY:105 state.working_summary = self._truncate_summary(action.summary_text or "")106 107 def _build_observation(self) -> Observation:108 state = self.state()109 current_turn = state.turns[state.current_index] if state.phase == "ingest" and state.current_index < len(state.turns) else None110 start = max(0, state.current_index - state.recent_turn_window + 1)111 recent_turns = state.turns[start : state.current_index + 1]112 return Observation(113 episode_id=state.episode_id,114 difficulty=state.difficulty,115 phase=state.phase,116 current_turn=current_turn,117 recent_turns=recent_turns,118 working_summary=state.working_summary,119 durable_memory=state.durable_memory,120 token_budget_remaining=state.token_budget_remaining(),121 step_count=state.step_count,122 future_query_queue_size=len(state.future_queries),123 done=state.done,124 info={125 "future_queries": [query.prompt for query in state.future_queries] if state.phase == "evaluation" else [],126 "token_budget": state.token_budget,127 "memory_write_policy": [128 "Store only durable information likely to matter later.",129 "Prefer uncertain or summary-only handling for ambiguous statements.",130 "When a value changes, keep the latest value active and mark the old one superseded.",131 ],132 "ambiguous_turn_ids": state.hidden_ambiguous_turn_ids,133 "low_value_turn_ids": state.hidden_low_value_turn_ids,134 },135 )136 137 def _populate_hidden_new_memory_ids(self) -> None:138 state = self.state()139 payload = build_episode_payload(state.difficulty, state.seed)140 mapping: dict[int, list[str]] = payload["new_memory_ids_by_turn"] # type: ignore[assignment]141 if state.phase != "ingest" or state.current_index >= len(state.turns):142 state.hidden_new_memory_ids = []143 return144 turn_id = state.turns[state.current_index].turn_id145 state.hidden_new_memory_ids = list(mapping.get(turn_id, []))146 147 def _upsert_memory(self, item: MemoryItem, preserve_history: bool) -> None:148 state = self.state()149 incoming = item.model_copy(deep=True)150 incoming.source_turn_ids = sorted(set(incoming.source_turn_ids))151 152 conflicting_indices: list[int] = []153 for index, existing in enumerate(state.durable_memory):154 if existing.memory_id == incoming.memory_id:155 state.durable_memory[index] = self._merge_memory(existing, incoming)156 self._supersede_relation_conflicts(state.durable_memory[index])157 return158 if (159 existing.status == MemoryStatus.ACTIVE160 and existing.relation_key() == incoming.relation_key()161 and existing.object.lower() != incoming.object.lower()162 ):163 conflicting_indices.append(index)164 165 if conflicting_indices and not incoming.requires_confirmation:166 for index in conflicting_indices:167 conflicting = state.durable_memory[index]168 conflicting.status = MemoryStatus.SUPERSEDED169 conflicting.updated_from_memory_id = conflicting.updated_from_memory_id or incoming.memory_id170 171 if preserve_history or not any(memory.memory_id == incoming.memory_id for memory in state.durable_memory):172 state.durable_memory.append(incoming)173 else:174 for index, existing in enumerate(state.durable_memory):175 if existing.memory_id == incoming.memory_id:176 state.durable_memory[index] = self._merge_memory(existing, incoming)177 break178 179 self._supersede_relation_conflicts(incoming)180 181 def _merge_memory(self, existing: MemoryItem, incoming: MemoryItem) -> MemoryItem:182 merged = existing.model_copy(deep=True)183 merged.type = incoming.type184 merged.subject = incoming.subject185 merged.predicate = incoming.predicate186 merged.object = incoming.object187 merged.confidence = incoming.confidence188 merged.source_turn_ids = sorted(set(existing.source_turn_ids + incoming.source_turn_ids))189 merged.source_text = incoming.source_text or existing.source_text190 merged.status = incoming.status191 merged.updated_from_memory_id = incoming.updated_from_memory_id or existing.updated_from_memory_id192 merged.expires_at = incoming.expires_at or existing.expires_at193 merged.importance = incoming.importance194 merged.task_relevance = incoming.task_relevance195 merged.requires_confirmation = incoming.requires_confirmation196 return merged197 198 def _supersede_relation_conflicts(self, incoming: MemoryItem) -> None:199 state = self.state()200 if incoming.status != MemoryStatus.ACTIVE or incoming.requires_confirmation:201 return202 for existing in state.durable_memory:203 if existing.memory_id == incoming.memory_id:204 continue205 if (206 existing.status == MemoryStatus.ACTIVE207 and existing.relation_key() == incoming.relation_key()208 and existing.object.lower() != incoming.object.lower()209 ):210 existing.status = MemoryStatus.SUPERSEDED211 existing.updated_from_memory_id = existing.updated_from_memory_id or incoming.memory_id212 213 def _truncate_summary(self, text: str, max_words: int = 40) -> str:214 return " ".join(text.split()[:max_words])215 216 217def serialise_state(state: EpisodeState) -> dict[str, Any]:218 return state.model_dump()219 220 221def _strict_reward(value: float) -> float:222 epsilon = 0.001223 return max(epsilon, min(1.0 - epsilon, value))224 