forked from BasedHardware/omi
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmemory_operations.py
More file actions
347 lines (306 loc) · 13 KB
/
Copy pathmemory_operations.py
File metadata and controls
347 lines (306 loc) · 13 KB
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
164
165
166
167
168
169
170
171
172
173
174
175
176
177
178
179
180
181
182
183
184
185
186
187
188
189
190
191
192
193
194
195
196
197
198
199
200
201
202
203
204
205
206
207
208
209
210
211
212
213
214
215
216
217
218
219
220
221
222
223
224
225
226
227
228
229
230
231
232
233
234
235
236
237
238
239
240
241
242
243
244
245
246
247
248
249
250
251
252
253
254
255
256
257
258
259
260
261
262
263
264
265
266
267
268
269
270
271
272
273
274
275
276
277
278
279
280
281
282
283
284
285
286
287
288
289
290
291
292
293
294
295
296
297
298
299
300
301
302
303
304
305
306
307
308
309
310
311
312
313
314
315
316
317
318
319
320
321
322
323
324
325
326
327
328
329
330
331
332
333
334
335
336
337
338
339
340
341
342
343
344
345
346
347
from __future__ import annotations
from datetime import datetime, timezone
from enum import Enum
from typing import Any, Dict, List, Optional
from pydantic import BaseModel, ConfigDict, Field, field_validator, model_validator
from models.memory_contracts import deterministic_contract_id
class MemoryOperationType(str, Enum):
source_candidate = "source_candidate"
source_replacement = "source_replacement"
synthesis = "synthesis"
long_term_apply = "long_term_apply"
user_mutation = "user_mutation"
archive_transition = "archive_transition"
projection_sync = "projection_sync"
vector_sync = "vector_sync"
graph_enrichment = "graph_enrichment"
deletion = "deletion"
ledger_mutation = "ledger_mutation"
class MemoryOperationStatus(str, Enum):
pending = "pending"
committed = "committed"
skipped_idempotent = "skipped_idempotent"
retryable_failure = "retryable_failure"
permanent_failure = "permanent_failure"
stale_generation = "stale_generation"
_TERMINAL_STATUSES = {
MemoryOperationStatus.committed,
MemoryOperationStatus.skipped_idempotent,
MemoryOperationStatus.permanent_failure,
MemoryOperationStatus.stale_generation,
}
class MemoryLedgerReopenReceipt(BaseModel):
"""Atomic source-to-tail receipt for standalone ledger reopening.
This is journal metadata, not a second memory authority. One receipt is
keyed by the closed source memory id so concurrent requests with different
client operation UUIDs cannot create multiple current tails.
"""
schema_version: str = "memory_ledger_reopen_receipt.v1"
uid: str
source_memory_id: str
replacement_memory_id: str
operation_id: str
account_generation: int
source_generation: int
source_item_revision: int
source_content_hash: str
committed_at: datetime = Field(default_factory=lambda: datetime.now(timezone.utc))
@field_validator(
"uid",
"source_memory_id",
"replacement_memory_id",
"operation_id",
"source_content_hash",
)
@classmethod
def validate_required_nonblank(cls, value: str) -> str:
if not value or not value.strip():
raise ValueError("reopen receipt identifiers must not be blank")
return value
@field_validator("account_generation", "source_generation", "source_item_revision")
@classmethod
def validate_nonnegative(cls, value: int) -> int:
if value < 0:
raise ValueError("reopen receipt generations and revisions must be nonnegative")
return value
class OperationLogicalPayload(BaseModel):
model_config = ConfigDict(extra="forbid")
decision: str
memory_text: Optional[str] = None
target_memory_id: Optional[str] = None
result_status: Optional[str] = None
supersedes: List[str] = Field(default_factory=list)
subject_entity_id: Optional[str] = None
predicate: Optional[str] = None
arguments: Dict[str, Any] = Field(default_factory=dict)
target_tier: Optional[str] = None
target_visibility: Optional[str] = None
target_user_asserted: Optional[bool] = None
clear_graph_assertion: Optional[bool] = None
mutation_metadata: Optional[Dict[str, Any]] = None
metadata: Dict[str, Any] = Field(default_factory=dict)
def canonical(self) -> Dict[str, Any]:
return self.model_dump(exclude_none=True)
def _coerce_logical_payload(value: OperationLogicalPayload | Dict[str, Any]) -> OperationLogicalPayload:
if isinstance(value, OperationLogicalPayload):
return value
known = {
key: value[key]
for key in [
"decision",
"memory_text",
"target_memory_id",
"result_status",
"supersedes",
"subject_entity_id",
"predicate",
"arguments",
"target_tier",
"target_visibility",
"target_user_asserted",
"clear_graph_assertion",
"mutation_metadata",
]
if key in value
}
metadata = {key: val for key, val in value.items() if key not in known}
return OperationLogicalPayload(**known, metadata=metadata)
def build_operation_id(
*,
uid: str,
operation_type: MemoryOperationType | str,
source_packet_id: Optional[str],
target_memory_id: Optional[str],
evidence_ids: List[str],
logical_payload: OperationLogicalPayload | Dict[str, Any],
account_generation: int,
source_generation: int,
observed_head_commit_id: Optional[str] = None,
output_index: Optional[int] = None,
) -> str:
"""Build a server-owned logical idempotency ID.
`observed_head_commit_id` and model output order/index are intentionally ignored.
Account/source generations are included so deletes/purges reset identity space.
"""
resolved_type = operation_type.value if isinstance(operation_type, MemoryOperationType) else operation_type
payload_model = _coerce_logical_payload(logical_payload)
payload = {
"uid": uid,
"operation_type": resolved_type,
"source_packet_id": source_packet_id,
"target_memory_id": target_memory_id,
"evidence_ids": sorted(evidence_ids or []),
"logical_payload": payload_model.canonical(),
"account_generation": account_generation,
"source_generation": source_generation,
}
return "op_" + deterministic_contract_id("memory-operation", payload)[:32]
def logical_payload_digest(value: OperationLogicalPayload | Dict[str, Any]) -> str:
return deterministic_contract_id("memory-operation-logical-payload", _coerce_logical_payload(value).canonical())
class MemoryOperation(BaseModel):
operation_id: str
uid: str
operation_type: MemoryOperationType
status: MemoryOperationStatus
source_packet_id: Optional[str] = None
target_memory_id: Optional[str] = None
evidence_ids: List[str] = Field(default_factory=list)
logical_payload: OperationLogicalPayload
logical_payload_digest: str
account_generation: int
source_generation: int
observed_head_commit_id: Optional[str] = None
committed_head_commit_id: Optional[str] = None
committed_sequence: Optional[int] = None
committed_memory_item_ids: List[str] = Field(default_factory=list)
committed_outbox_event_ids: List[str] = Field(default_factory=list)
attempt_count: int = 0
error_code: Optional[str] = None
untrusted_proposed_operation_id: Optional[str] = None
created_at: datetime = Field(default_factory=lambda: datetime.now(timezone.utc))
updated_at: datetime = Field(default_factory=lambda: datetime.now(timezone.utc))
@classmethod
def new(
cls,
*,
uid: str,
operation_type: MemoryOperationType,
source_packet_id: Optional[str],
target_memory_id: Optional[str],
evidence_ids: List[str],
logical_payload: OperationLogicalPayload | Dict[str, Any],
account_generation: int,
source_generation: int,
observed_head_commit_id: Optional[str] = None,
proposed_operation_id: Optional[str] = None,
) -> "MemoryOperation":
payload_model = _coerce_logical_payload(logical_payload)
operation_id = build_operation_id(
uid=uid,
operation_type=operation_type,
source_packet_id=source_packet_id,
target_memory_id=target_memory_id,
evidence_ids=evidence_ids,
logical_payload=payload_model,
account_generation=account_generation,
source_generation=source_generation,
observed_head_commit_id=observed_head_commit_id,
)
now = datetime.now(timezone.utc)
return cls(
operation_id=operation_id,
uid=uid,
operation_type=operation_type,
status=MemoryOperationStatus.pending,
source_packet_id=source_packet_id,
target_memory_id=target_memory_id,
evidence_ids=evidence_ids,
logical_payload=payload_model,
logical_payload_digest=logical_payload_digest(payload_model),
account_generation=account_generation,
source_generation=source_generation,
observed_head_commit_id=observed_head_commit_id,
untrusted_proposed_operation_id=proposed_operation_id if proposed_operation_id else None,
created_at=now,
updated_at=now,
)
@field_validator("operation_id", "uid")
@classmethod
def validate_required_nonblank(cls, value: str) -> str:
if not value or not value.strip():
raise ValueError("required operation fields must not be blank")
return value
@field_validator("committed_head_commit_id", "error_code")
@classmethod
def validate_optional_nonblank(cls, value: Optional[str]) -> Optional[str]:
if value is not None and not value.strip():
raise ValueError("optional operation fields must not be blank")
return value
@field_validator("account_generation", "source_generation", "attempt_count")
@classmethod
def validate_nonnegative(cls, value: int) -> int:
if value < 0:
raise ValueError("generation and attempt counts must be nonnegative")
return value
@field_validator("committed_sequence")
@classmethod
def validate_optional_nonnegative(cls, value: Optional[int]) -> Optional[int]:
if value is not None and value < 0:
raise ValueError("committed_sequence must be nonnegative")
return value
@field_validator("created_at", "updated_at")
@classmethod
def validate_timezone(cls, value: datetime) -> datetime:
if value.tzinfo is None or value.utcoffset() is None:
raise ValueError("operation timestamps must be timezone-aware")
return value
@model_validator(mode="after")
def validate_integrity(self):
if self.updated_at < self.created_at:
raise ValueError("updated_at must be >= created_at")
expected = build_operation_id(
uid=self.uid,
operation_type=self.operation_type,
source_packet_id=self.source_packet_id,
target_memory_id=self.target_memory_id,
evidence_ids=self.evidence_ids,
logical_payload=self.logical_payload,
account_generation=self.account_generation,
source_generation=self.source_generation,
observed_head_commit_id=self.observed_head_commit_id,
)
if self.operation_id != expected:
raise ValueError("operation_id does not match server-computed logical identity")
if self.logical_payload_digest != logical_payload_digest(self.logical_payload):
raise ValueError("logical_payload_digest does not match canonical logical payload")
if self.status == MemoryOperationStatus.committed and not self.committed_head_commit_id:
raise ValueError("committed operations require committed_head_commit_id")
if self.status == MemoryOperationStatus.committed and self.committed_sequence is None:
raise ValueError("committed operations require committed_sequence")
if (
self.status in {MemoryOperationStatus.retryable_failure, MemoryOperationStatus.permanent_failure}
and not self.error_code
):
raise ValueError("failure operations require error_code")
return self
def _transition(self, *, status: MemoryOperationStatus, **updates: Any) -> "MemoryOperation":
if self.status in _TERMINAL_STATUSES:
raise ValueError(f"cannot transition terminal operation from {self.status.value}")
data = self.model_dump(mode="python")
data.update(updates)
data["status"] = status
data["updated_at"] = datetime.now(timezone.utc)
return MemoryOperation(**data)
def mark_retryable(self, error_code: str) -> "MemoryOperation":
return self._transition(
status=MemoryOperationStatus.retryable_failure,
attempt_count=self.attempt_count + 1,
error_code=error_code,
)
def mark_committed(
self,
committed_head_commit_id: str,
*,
committed_sequence: int,
committed_memory_item_ids: Optional[List[str]] = None,
committed_outbox_event_ids: Optional[List[str]] = None,
) -> "MemoryOperation":
return self._transition(
status=MemoryOperationStatus.committed,
committed_head_commit_id=committed_head_commit_id,
committed_sequence=committed_sequence,
committed_memory_item_ids=committed_memory_item_ids or [],
committed_outbox_event_ids=committed_outbox_event_ids or [],
error_code=None,
)
def is_stale(self, *, account_generation: int, source_generation: int) -> bool:
return account_generation != self.account_generation or source_generation != self.source_generation
__all__ = [
"MemoryLedgerReopenReceipt",
"MemoryOperation",
"MemoryOperationStatus",
"MemoryOperationType",
"OperationLogicalPayload",
"build_operation_id",
"logical_payload_digest",
]