CoolFace
Apppublic

dhruv-punia-bits/memory-compaction-openenv

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
models.py195 linesDownload Raw Back to server
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