forked from BasedHardware/omi
-
Notifications
You must be signed in to change notification settings - Fork 0
Expand file tree
/
Copy pathmemory_non_active_routes.py
More file actions
159 lines (124 loc) · 5.14 KB
/
Copy pathmemory_non_active_routes.py
File metadata and controls
159 lines (124 loc) · 5.14 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
"""Canonical non-active memory route persistence (WS-G7)."""
from __future__ import annotations
import hashlib
import json
from datetime import datetime, timezone
from enum import Enum
from typing import Any, Dict, List, Optional
try:
from google.cloud.firestore_v1 import transactional
except ImportError: # pragma: no cover - local unit tests mock Firestore.
def transactional(func):
def wrapper(transaction, *args, **kwargs):
if hasattr(transaction, "_begin"):
transaction._begin()
try:
result = func(transaction, *args, **kwargs)
if hasattr(transaction, "_commit"):
transaction._commit()
return result
except Exception:
if hasattr(transaction, "_rollback"):
transaction._rollback()
raise
finally:
if hasattr(transaction, "_clean_up"):
transaction._clean_up()
return wrapper
from pydantic import BaseModel, ConfigDict, Field, field_validator
from database._client import db
from database.memory_collections import MemoryCollections
from database.read_boundary import parse_snapshot_strict
class NonActiveRoute(str, Enum):
review = "review"
archive = "archive"
context_only = "context_only"
reject = "reject"
hidden = "hidden"
skip = "skip"
class NonActiveRouteStoreConflict(Exception):
pass
class _NonActiveRouteOutcomeBase(BaseModel):
model_config = ConfigDict(validate_assignment=True)
uid: str
route: NonActiveRoute
idempotency_key: str
source_ids: List[str]
reason: str
run_id: str
patch_id: Optional[str] = None
audit_metadata: Dict[str, Any] = Field(default_factory=dict)
created_at: datetime = Field(default_factory=lambda: datetime.now(timezone.utc))
default_long_term_visible: bool = False
@field_validator("uid", "idempotency_key", "reason", "run_id")
@classmethod
def validate_nonblank(cls, value: str) -> str:
if not value or not value.strip():
raise ValueError("required fields must not be blank")
return value
@field_validator("source_ids")
@classmethod
def validate_source_ids(cls, value: List[str]) -> List[str]:
normalized = sorted({source_id.strip() for source_id in value if source_id and source_id.strip()})
if not normalized:
raise ValueError("source_ids must not be empty")
return normalized
@field_validator("created_at")
@classmethod
def validate_timezone(cls, value: datetime) -> datetime:
if value.tzinfo is None or value.utcoffset() is None:
raise ValueError("timestamps must be timezone-aware")
return value
class NonActiveRouteOutcome(_NonActiveRouteOutcomeBase):
outcome_id: Optional[str] = None
payload_fingerprint: Optional[str] = None
class PersistedNonActiveRouteOutcome(_NonActiveRouteOutcomeBase):
outcome_id: str
payload_fingerprint: str
def persist_non_active_route_outcome(
outcome: NonActiveRouteOutcome,
*,
db_client=db,
) -> PersistedNonActiveRouteOutcome:
transaction = db_client.transaction()
return _persist_non_active_route_outcome_transaction(transaction, db_client, outcome)
@transactional
def _persist_non_active_route_outcome_transaction(
transaction,
db_client,
outcome: NonActiveRouteOutcome,
) -> PersistedNonActiveRouteOutcome:
persisted = _with_persistence_fields(outcome)
collections = MemoryCollections(uid=persisted.uid)
outcome_ref = db_client.document(f"{collections.non_active_memory_routes}/{persisted.outcome_id}")
snapshot = outcome_ref.get(transaction=transaction)
if snapshot.exists:
existing = parse_snapshot_strict(PersistedNonActiveRouteOutcome, snapshot)
if existing.payload_fingerprint != persisted.payload_fingerprint:
raise NonActiveRouteStoreConflict("idempotency key payload mismatch")
return existing
transaction.set(outcome_ref, persisted.model_dump(mode="json"))
return persisted
def _with_persistence_fields(outcome: NonActiveRouteOutcome) -> PersistedNonActiveRouteOutcome:
data = outcome.model_dump(mode="python")
outcome_id = outcome.outcome_id or _stable_outcome_id(outcome.uid, outcome.idempotency_key)
data["outcome_id"] = outcome_id
data["default_long_term_visible"] = False
data["payload_fingerprint"] = _payload_fingerprint(data)
return PersistedNonActiveRouteOutcome(**data)
def _stable_outcome_id(uid: str, idempotency_key: str) -> str:
digest = hashlib.sha256(f"{uid}:{idempotency_key}".encode("utf-8")).hexdigest()
return f"nar_{digest[:32]}"
def _payload_fingerprint(data: Dict[str, Any]) -> str:
comparable = dict(data)
comparable.pop("payload_fingerprint", None)
comparable.pop("created_at", None)
payload = json.dumps(comparable, sort_keys=True, default=str, separators=(",", ":"))
return hashlib.sha256(payload.encode("utf-8")).hexdigest()
__all__ = [
"NonActiveRoute",
"NonActiveRouteOutcome",
"NonActiveRouteStoreConflict",
"PersistedNonActiveRouteOutcome",
"persist_non_active_route_outcome",
]