dhruv-punia-bits/memory-compaction-openenv
0
1from __future__ import annotations2 3from enum import Enum4from typing import Any5 6from pydantic import BaseModel, Field, field_validator7 8 9class Difficulty(str, Enum):10 EASY = "easy"11 MEDIUM = "medium"12 HARD = "hard"13 14 15class MemoryType(str, Enum):16 PREFERENCE = "preference"17 FACT = "fact"18 TASK = "task"19 CONSTRAINT = "constraint"20 PLAN = "plan"21 CORRECTION = "correction"22 23 24class MemoryStatus(str, Enum):25 ACTIVE = "active"26 SUPERSEDED = "superseded"27 UNCERTAIN = "uncertain"28 29 30class Operation(str, Enum):31 APPEND_MEMORY = "append_memory"32 UPDATE_MEMORY = "update_memory"33 DELETE_MEMORY = "delete_memory"34 REPLACE_SUMMARY = "replace_summary"35 NOOP = "noop"36 37 38class Turn(BaseModel):39 turn_id: int40 speaker: str41 text: str42 43 44class MemoryItem(BaseModel):45 memory_id: str = Field(min_length=1)46 type: MemoryType47 subject: str = Field(min_length=1)48 predicate: str = Field(min_length=1)49 object: str = Field(min_length=1)50 confidence: float = Field(default=1.0, ge=0.0, le=1.0)51 source_turn_ids: list[int] = Field(default_factory=list)52 source_text: str = ""53 status: MemoryStatus = MemoryStatus.ACTIVE54 updated_from_memory_id: str | None = None55 expires_at: str | None = None56 importance: float = Field(default=0.5, ge=0.0, le=1.0)57 task_relevance: float = Field(default=0.5, ge=0.0, le=1.0)58 requires_confirmation: bool = False59 60 def key(self) -> tuple[str, str, str]:61 return (self.subject.lower(), self.predicate.lower(), self.object.lower())62 63 def relation_key(self) -> tuple[str, str]:64 return (self.subject.lower(), self.predicate.lower())65 66 67class FutureQuery(BaseModel):68 query_id: str69 prompt: str70 expected_memory_keys: list[str]71 forbidden_memory_keys: list[str] = Field(default_factory=list)72 73 74class Observation(BaseModel):75 episode_id: str76 difficulty: Difficulty77 phase: str78 current_turn: Turn | None = None79 recent_turns: list[Turn] = Field(default_factory=list)80 working_summary: str = ""81 durable_memory: list[MemoryItem] = Field(default_factory=list)82 token_budget_remaining: int83 step_count: int84 future_query_queue_size: int85 done: bool = False86 info: dict[str, Any] = Field(default_factory=dict)87 88 89class ResetRequest(BaseModel):90 difficulty: Difficulty = Difficulty.EASY91 seed: int | None = None92 93 94class StepAction(BaseModel):95 operation: Operation = Operation.NOOP96 memory_items: list[MemoryItem] = Field(default_factory=list)97 summary_text: str | None = None98 rationale: str | None = None99 100 @field_validator("summary_text")101 @classmethod102 def normalize_summary(cls, value: str | None) -> str | None:103 if value is None:104 return value105 return " ".join(value.split())106 107 108class RewardBreakdown(BaseModel):109 schema_score: float = 0.0110 coverage_score: float = 0.0111 precision_score: float = 0.0112 consistency_score: float = 0.0113 trust_score: float = 0.0114 budget_score: float = 0.0115 efficiency_score: float = 0.0116 terminal_score: float = 0.0117 total_score: float = 0.0118 119 120class StepResponse(BaseModel):121 observation: Observation122 reward: float = Field(ge=0.0, le=1.0)123 done: bool124 reward_breakdown: RewardBreakdown125 metadata: dict[str, Any] = Field(default_factory=dict)126 127 128class ResetResponse(BaseModel):129 observation: Observation130 131 132class StateResponse(BaseModel):133 state: dict[str, Any]134 135 136class MetadataResponse(BaseModel):137 name: str138 description: str139 version: str140 benchmark_task_count: int141 142 143class SchemaResponse(BaseModel):144 action: dict[str, Any]145 observation: dict[str, Any]146 state: dict[str, Any]147 148 149class MCPResponse(BaseModel):150 jsonrpc: str = "2.0"151 id: str | None = None152 result: dict[str, Any] = Field(default_factory=dict)153 154 155class TaskDefinition(BaseModel):156 difficulty: Difficulty157 title: str158 description: str159 seed: int160 token_budget: int161 turn_count: int162 evaluation_query_count: int163 164 165class EpisodeState(BaseModel):166 episode_id: str167 difficulty: Difficulty168 seed: int169 phase: str170 current_index: int171 token_budget: int172 step_count: int173 recent_turn_window: int174 turns: list[Turn]175 future_queries: list[FutureQuery]176 working_summary: str = ""177 durable_memory: list[MemoryItem] = Field(default_factory=list)178 latest_reward_breakdown: RewardBreakdown = Field(default_factory=RewardBreakdown)179 done: bool = False180 hidden_gold_memories: list[MemoryItem] = Field(default_factory=list)181 hidden_new_memory_ids: list[str] = Field(default_factory=list)182 hidden_ambiguous_turn_ids: list[int] = Field(default_factory=list)183 hidden_low_value_turn_ids: list[int] = Field(default_factory=list)184 185 def token_usage(self) -> int:186 summary_tokens = len(self.working_summary.split())187 memory_tokens = sum(188 len(item.subject.split()) + len(item.predicate.split()) + len(item.object.split())189 for item in self.durable_memory190 )191 return summary_tokens + memory_tokens192 193 def token_budget_remaining(self) -> int:194 return max(0, self.token_budget - self.token_usage())195 