fix(memory): keep OpenViking identity operations consistent

This commit is contained in:
ehz0ah
2026-08-25 11:48:19 +08:00
committed by kshitij
parent 5ff03cb0c4
commit 4387e03960
2 changed files with 134 additions and 15 deletions
+38 -15
View File
@@ -105,7 +105,7 @@ _PREFERENCES_SUFFIX = "memories/preferences"
_ENTITIES_SUFFIX = "memories/entities"
def _resolve_user_space(client) -> Optional[str]:
def _resolve_user_space(client, *, timeout: Optional[float] = None) -> Optional[str]:
"""Server-asserted current user for explicit-uid URIs.
Return ``None`` when the probe fails or reports no user. Callers can use a
@@ -113,7 +113,8 @@ def _resolve_user_space(client) -> Optional[str]:
identity because a later probe can succeed.
"""
try:
status = client.get("/api/v1/system/status")
kwargs = {"timeout": timeout} if timeout is not None else {}
status = client.get("/api/v1/system/status", **kwargs)
result = (status or {}).get("result") or {}
user = str(result.get("user") or "").strip()
if user:
@@ -3910,7 +3911,7 @@ class OpenVikingMemoryProvider(MemoryProvider):
).lstrip()
return f"{head}{marker}{tail}" if tail else _head_only()
def _user_space(self, client=None) -> str:
def _user_space(self, client=None, *, timeout: Optional[float] = None) -> str:
"""Resolve the user space, caching only a confirmed client identity."""
active = client if client is not None else getattr(self, "_client", None)
cached = getattr(self, "_user_space_cache", None)
@@ -3918,7 +3919,7 @@ class OpenVikingMemoryProvider(MemoryProvider):
return cached[1]
if active is not None:
resolved = _resolve_user_space(active)
resolved = _resolve_user_space(active, timeout=timeout)
if resolved:
# The probe can overlap a config reload. Only publish it when
# this is still the provider's active client. Old in-flight
@@ -3938,8 +3939,8 @@ class OpenVikingMemoryProvider(MemoryProvider):
def _user_scoped_uri(self, suffix: str, client=None) -> str:
return _user_scoped_uri(self._user_space(client), suffix)
def _session_start_uris(self) -> tuple:
user = self._user_space()
def _session_start_uris(self, user: Optional[str] = None) -> tuple:
user = user or self._user_space()
return (
_user_scoped_uri(user, _PROFILE_SUFFIX),
_user_scoped_uri(user, _PREFERENCES_SUFFIX),
@@ -3949,6 +3950,7 @@ class OpenVikingMemoryProvider(MemoryProvider):
def _read_session_start_profile(
self,
client: _VikingClient,
uri: str,
*,
deadline: float,
request_timeout: float,
@@ -3957,7 +3959,7 @@ class OpenVikingMemoryProvider(MemoryProvider):
timeout = self._remaining_recall_timeout(deadline, request_timeout)
resp = client.get(
"/api/v1/content/read",
params={"uri": self._user_scoped_uri(_PROFILE_SUFFIX, client)},
params={"uri": uri},
timeout=timeout,
)
except Exception as e:
@@ -3996,8 +3998,16 @@ class OpenVikingMemoryProvider(MemoryProvider):
if not active_client:
return {}
try:
identity_timeout = self._remaining_recall_timeout(deadline, request_timeout)
user = self._user_space(active_client, timeout=identity_timeout)
except Exception:
return {"profile": None, "preferences": [], "entities": []}
uris = self._session_start_uris(user)
profile = self._read_session_start_profile(
active_client,
uris[0],
deadline=deadline,
request_timeout=request_timeout,
)
@@ -4007,16 +4017,17 @@ class OpenVikingMemoryProvider(MemoryProvider):
"profile": profile,
"preferences": self._list_session_start_memories(
active_client,
self._user_scoped_uri(_PREFERENCES_SUFFIX, active_client),
uris[1],
deadline=deadline,
request_timeout=request_timeout,
),
"entities": self._list_session_start_memories(
active_client,
self._user_scoped_uri(_ENTITIES_SUFFIX, active_client),
uris[2],
deadline=deadline,
request_timeout=request_timeout,
),
"uris": uris,
}
@staticmethod
@@ -4161,7 +4172,7 @@ class OpenVikingMemoryProvider(MemoryProvider):
preferences=raw_parts.get("preferences") or [],
entities=raw_parts.get("entities") or [],
token_budget=self._profile_token_budget(),
uris=self._session_start_uris(),
uris=raw_parts["uris"],
)
@staticmethod
@@ -4861,15 +4872,21 @@ class OpenVikingMemoryProvider(MemoryProvider):
old_session_id, new_id, parent_session_id, reset,
)
def _build_memory_uri(self, subdir: str) -> str:
def _build_memory_uri(self, subdir: str, *, client=None) -> str:
"""Build a viking:// memory URI under the configured peer namespace."""
slug = uuid.uuid4().hex[:12]
# Explicit-uid URIs are canonical under every auth mode; the uid-less
# `viking://user/peers/...` shorthand was removed upstream (#4196) and
# `viking://~/...` only expands for USER/ADMIN roles, not dev/ROOT.
active_client = client if client is not None else getattr(self, "_client", None)
agent = str(
getattr(active_client, "_agent", "")
or getattr(self, "_agent", "")
or _DEFAULT_AGENT
).strip()
return _user_scoped_uri(
self._user_space(),
f"peers/{self._agent}/memories/{subdir}/mem_{slug}.md",
self._user_space(active_client),
f"peers/{agent}/memories/{subdir}/mem_{slug}.md",
)
def on_memory_write(
@@ -4884,11 +4901,17 @@ class OpenVikingMemoryProvider(MemoryProvider):
return
subdir = _MEMORY_WRITE_TARGET_SUBDIR_MAP.get(target, _DEFAULT_MEMORY_SUBDIR)
uri = self._build_memory_uri(subdir)
try:
# Keep identity resolution, URI construction, and the write on one
# connection snapshot even if the active profile reloads.
client = self._new_client()
except Exception as e:
logger.debug("OpenViking memory mirror client creation failed: %s", e)
return
def _write():
try:
client = self._new_client()
uri = self._build_memory_uri(subdir, client=client)
client.post("/api/v1/content/write", {
"uri": uri,
"content": content,
@@ -1313,6 +1313,60 @@ def test_shutdown_waits_for_memory_write_worker(monkeypatch):
assert provider._memory_write_threads == set()
def test_memory_write_uses_one_connection_for_identity_uri_and_post(monkeypatch):
import threading
provider = OpenVikingMemoryProvider()
provider._agent = "alice-agent"
provider._ensure_client = lambda: True
identity_started = threading.Event()
release_identity = threading.Event()
write_finished = threading.Event()
writes = []
class StubClient:
def __init__(self, user, agent):
self._user = user
self._agent = agent
def get(self, path, **kwargs):
assert path == "/api/v1/system/status"
identity_started.set()
assert release_identity.wait(timeout=2.0)
return {"status": "ok", "result": {"user": self._user}}
def post(self, path, payload=None, **kwargs):
writes.append((self._user, self._agent, path, payload))
write_finished.set()
return {"status": "ok"}
alice = StubClient("alice", "alice-agent")
bob = StubClient("bob", "bob-agent")
provider._client = alice
monkeypatch.setattr(provider, "_new_client", lambda: alice)
provider.on_memory_write("add", "user", "remember this")
assert identity_started.wait(timeout=2.0), "identity probe did not start"
# Simulate a profile reload while the write worker is resolving identity.
provider._client = bob
provider._agent = "bob-agent"
release_identity.set()
assert write_finished.wait(timeout=2.0), "memory write did not finish"
for worker in list(provider._memory_write_threads):
worker.join(timeout=2.0)
assert len(writes) == 1
user, agent, path, payload = writes[0]
assert (user, agent, path) == ("alice", "alice-agent", "/api/v1/content/write")
assert payload["uri"].startswith(
"viking://user/alice/peers/alice-agent/memories/preferences/mem_"
)
assert provider._memory_write_threads == set()
def _make_prefetch_provider() -> OpenVikingMemoryProvider:
provider = OpenVikingMemoryProvider()
provider._client = MagicMock()
@@ -1432,6 +1486,48 @@ def test_prefetch_prepends_session_start_memory_context_once_per_session():
assert provider._search_prefetch_context.call_count == 2
def test_session_start_reuses_one_fallback_user_after_status_probe_failure():
provider = _make_prefetch_provider()
provider._user = "configured-user"
provider._client._user = "configured-user"
provider._search_prefetch_context = MagicMock(return_value="")
status_calls = 0
status_timeouts = []
read_uris = []
def fake_get(path, params=None, **kwargs):
nonlocal status_calls
if path == "/api/v1/system/status":
status_calls += 1
status_timeouts.append(kwargs.get("timeout"))
if status_calls == 1:
raise RuntimeError("temporary status failure")
return {"status": "ok", "result": {"user": "alice"}}
uri = (params or {}).get("uri", "")
read_uris.append(uri)
if path == "/api/v1/content/read":
return {"result": "Configured-user profile."}
return {"result": []}
provider._client.get.side_effect = fake_get
block = provider.prefetch("What should we recall?", session_id="sid-fallback")
assert status_calls == 1
assert len(status_timeouts) == 1
assert 0 < status_timeouts[0] <= 3.0
assert read_uris == [
"viking://user/configured-user/memories/profile.md",
"viking://user/configured-user/memories/preferences",
"viking://user/configured-user/memories/entities",
]
assert (
'<user-profile uri="viking://user/configured-user/memories/profile.md">'
in block
)
def test_prefetch_reinjects_after_in_place_compression_same_session():
provider = _make_prefetch_provider()
provider._session_id = "sid-123"