forked from BasedHardware/omi
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathaccount_generation_source.py
More file actions
126 lines (101 loc) · 5.46 KB
/
Copy pathaccount_generation_source.py
File metadata and controls
126 lines (101 loc) · 5.46 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
"""Canonical module for ``utils.memory.v3.account_generation_source`` (WS-G8b).
This module owns the trusted V3 account-generation read contract.
"""
from __future__ import annotations
from dataclasses import dataclass
from enum import Enum
from typing import Any, cast
from database.memory_collections import MemoryCollections
from models.memory_state_head import MEMORY_STATE_HEAD_SCHEMA_VERSION, MEMORY_STATE_HEAD_SOURCE
V3_TRUSTED_ACCOUNT_GENERATION_SCHEMA_VERSION = MEMORY_STATE_HEAD_SCHEMA_VERSION
V3_TRUSTED_ACCOUNT_GENERATION_SOURCE = MEMORY_STATE_HEAD_SOURCE
class V3AccountGenerationFailureReason(str, Enum):
MISSING_STATE_HEAD = 'missing_state_head'
MALFORMED_STATE_HEAD = 'malformed_state_head'
UNSUPPORTED_SCHEMA = 'unsupported_schema'
UID_MISMATCH = 'uid_mismatch'
SOURCE_MISMATCH = 'source_mismatch'
MALFORMED_ACCOUNT_GENERATION = 'malformed_account_generation'
READ_FAILED = 'read_failed'
class V3TrustedAccountGenerationReadError(RuntimeError):
def __init__(self, reason: V3AccountGenerationFailureReason, message: str | None = None):
super().__init__(message or reason.value)
self.reason = reason
@dataclass(frozen=True)
class V3TrustedAccountGenerationResult:
uid: str
source_path: str
account_generation: int | None = None
head_commit_id: str | None = None
commit_sequence: int | None = None
source: str | None = None
schema_version: int | None = None
read_error_reason: V3AccountGenerationFailureReason | None = None
def require_account_generation(self) -> int:
if self.read_error_reason is not None or self.account_generation is None:
raise V3TrustedAccountGenerationReadError(
self.read_error_reason or V3AccountGenerationFailureReason.MALFORMED_ACCOUNT_GENERATION
)
return self.account_generation
_MALFORMED_SNAPSHOT_DATA = object()
def _snapshot_data(snapshot: Any) -> dict[str, Any] | None | object:
if snapshot is None or getattr(snapshot, 'exists', False) is False:
return None
data = snapshot.to_dict()
return cast(dict[str, Any], data) if isinstance(data, dict) else _MALFORMED_SNAPSHOT_DATA
def _fail(*, uid: str, source_path: str, reason: V3AccountGenerationFailureReason) -> V3TrustedAccountGenerationResult:
return V3TrustedAccountGenerationResult(uid=uid, source_path=source_path, read_error_reason=reason)
def read_memory_v3_trusted_account_generation(
*,
uid: str,
db_client: Any,
transaction: Any | None = None,
) -> V3TrustedAccountGenerationResult:
"""Read and validate the independent account-generation state-head source.
Future `/v3` GET wiring must feed this returned generation into projection
reads and then compare it with control/projection/cursor generations. It must
not derive ``expected_account_generation`` from the projection state being
verified or from the memory control decision document.
"""
source_path = MemoryCollections(uid=uid).memory_state_head
try:
document = db_client.document(source_path)
snapshot = document.get(transaction=transaction) if transaction is not None else document.get()
data = _snapshot_data(snapshot)
except Exception:
return _fail(uid=uid, source_path=source_path, reason=V3AccountGenerationFailureReason.READ_FAILED)
if data is None:
return _fail(uid=uid, source_path=source_path, reason=V3AccountGenerationFailureReason.MISSING_STATE_HEAD)
if data is _MALFORMED_SNAPSHOT_DATA:
return _fail(uid=uid, source_path=source_path, reason=V3AccountGenerationFailureReason.MALFORMED_STATE_HEAD)
if not isinstance(data, dict):
return _fail(uid=uid, source_path=source_path, reason=V3AccountGenerationFailureReason.MALFORMED_STATE_HEAD)
payload = cast(dict[str, Any], data)
if payload.get('schema_version') != V3_TRUSTED_ACCOUNT_GENERATION_SCHEMA_VERSION:
return _fail(uid=uid, source_path=source_path, reason=V3AccountGenerationFailureReason.UNSUPPORTED_SCHEMA)
if payload.get('uid') != uid:
return _fail(uid=uid, source_path=source_path, reason=V3AccountGenerationFailureReason.UID_MISMATCH)
if payload.get('source') != V3_TRUSTED_ACCOUNT_GENERATION_SOURCE:
return _fail(uid=uid, source_path=source_path, reason=V3AccountGenerationFailureReason.SOURCE_MISMATCH)
account_generation = payload.get('account_generation')
if isinstance(account_generation, bool) or not isinstance(account_generation, int) or account_generation < 0:
return _fail(
uid=uid,
source_path=source_path,
reason=V3AccountGenerationFailureReason.MALFORMED_ACCOUNT_GENERATION,
)
head_commit_id = payload.get('head_commit_id')
commit_sequence = payload.get('commit_sequence')
if not isinstance(head_commit_id, str) or not head_commit_id:
return _fail(uid=uid, source_path=source_path, reason=V3AccountGenerationFailureReason.MALFORMED_STATE_HEAD)
if isinstance(commit_sequence, bool) or not isinstance(commit_sequence, int) or commit_sequence < 0:
return _fail(uid=uid, source_path=source_path, reason=V3AccountGenerationFailureReason.MALFORMED_STATE_HEAD)
return V3TrustedAccountGenerationResult(
uid=uid,
source_path=source_path,
account_generation=account_generation,
head_commit_id=head_commit_id,
commit_sequence=commit_sequence,
source=V3_TRUSTED_ACCOUNT_GENERATION_SOURCE,
schema_version=V3_TRUSTED_ACCOUNT_GENERATION_SCHEMA_VERSION,
)