refactor(openviking): inline one-site helpers, fold turn-upload fallbacks into _TurnUpload

This commit is contained in:
Teknium
2026-09-02 23:47:17 -07:00
parent b2d09087c8
commit e41d4ded1f

View File

@@ -259,14 +259,6 @@ class _VikingClient:
h["Authorization"] = "Bearer " + self._api_key
return h
def _url(self, path: str) -> str:
return f"{self._endpoint}{path}"
def _multipart_headers(self, *, include_tenant: bool | None = None) -> dict:
headers = self._headers(include_tenant=include_tenant)
headers.pop("Content-Type", None)
return headers
@staticmethod
def _needs_trusted_identity_retry(exc: Exception) -> bool:
"""Trusted mode asks for X-OpenViking-Account/User with wording that varies across
@@ -278,6 +270,11 @@ class _VikingClient:
return False
return getattr(exc, "status_code", None) in (None, 400)
def _multipart_headers(self, *, include_tenant: bool | None = None) -> dict:
headers = self._headers(include_tenant=include_tenant)
headers.pop("Content-Type", None)
return headers
def _send_with_trusted_identity_retry(self, send, *, multipart: bool = False) -> dict:
build = self._multipart_headers if multipart else self._headers
try:
@@ -309,7 +306,7 @@ class _VikingClient:
def _request(self, method: str, path: str, kwargs: dict) -> dict:
timeout = kwargs.pop("timeout", _TIMEOUT)
fn = getattr(self._httpx, method)
return self._send_with_trusted_identity_retry(lambda headers: fn(self._url(path), headers=headers, timeout=timeout, **kwargs))
return self._send_with_trusted_identity_retry(lambda headers: fn(f"{self._endpoint}{path}", headers=headers, timeout=timeout, **kwargs))
def get(self, path: str, **kwargs) -> dict:
return self._request("get", path, kwargs)
@@ -325,7 +322,7 @@ class _VikingClient:
def _send(headers):
with file_path.open("rb") as f:
return self._httpx.post(self._url("/api/v1/resources/temp_upload"),
return self._httpx.post(f"{self._endpoint}/api/v1/resources/temp_upload",
files={"file": (file_path.name, f, mime_type)}, headers=headers, timeout=_TIMEOUT)
temp_file_id = self._send_with_trusted_identity_retry(_send, multipart=True).get("result", {}).get("temp_file_id", "")
@@ -341,7 +338,7 @@ class _VikingClient:
def _anonymous_json(self, path: str) -> dict:
"""Probe server identity without disclosing credentials or tenant IDs."""
return self._parse_response(self._httpx.get(self._url(path), headers={"Accept": "application/json"}, timeout=3.0))
return self._parse_response(self._httpx.get(f"{self._endpoint}{path}", headers={"Accept": "application/json"}, timeout=3.0))
def health_payload(self) -> dict:
"""``GET /health``, anonymous first so credentials never reach an unknown host.
@@ -352,7 +349,7 @@ class _VikingClient:
except _OpenVikingHTTPError as exc:
if not self._api_key or _status_code_from_error(exc) not in {401, 403}:
raise
return self._parse_response(self._httpx.get(self._url("/health"), headers=self._headers(include_tenant=False), timeout=3.0))
return self._parse_response(self._httpx.get(f"{self._endpoint}/health", headers=self._headers(include_tenant=False), timeout=3.0))
def openapi_payload(self) -> dict:
return self._anonymous_json("/openapi.json")
@@ -472,10 +469,6 @@ def _resolve_user_space(client, *, timeout: Optional[float] = None) -> Optional[
return None
def _user_scoped_uri(user_space: str, suffix: str) -> str:
return f"viking://user/{user_space}/{suffix}"
def _zip_directory(dir_path: Path) -> Path:
"""Zip a directory tree into a temp file, skipping symlinks, escapes, and read-blocked files."""
from agent.file_safety import raise_if_read_blocked
@@ -500,10 +493,6 @@ def _is_windows_absolute_path(value: str) -> bool:
return len(value) >= 3 and value[0].isalpha() and value[1] == ":" and value[2] in {"/", "\\"}
def _is_remote_resource_source(value: str) -> bool:
return value.startswith(_REMOTE_RESOURCE_PREFIXES)
def _memory_segment_index(parts: List[str]) -> Optional[int]:
"""Index of the ``memories`` segment for the user / user-uid / peer / uid-peer layouts."""
if not parts or parts[0] != "user":
@@ -535,7 +524,7 @@ def _validate_forget_memory_uri(raw_uri: Any) -> tuple[Optional[str], Optional[s
def _is_local_path_reference(value: str) -> bool:
if not value or "\n" in value or "\r" in value or _is_remote_resource_source(value):
if not value or "\n" in value or "\r" in value or value.startswith(_REMOTE_RESOURCE_PREFIXES):
return False
if _is_windows_absolute_path(value):
return True
@@ -783,10 +772,6 @@ def _load_hermes_openviking_config() -> dict:
return {}
def _env_value(name: str) -> Optional[str]:
return os.environ[name].strip() if name in os.environ else None
def _ovcli_values_for(provider_config: dict) -> dict:
"""Connection values from the linked ovcli profile, or {} when none is linked."""
if not provider_config.get("use_ovcli_config"):
@@ -803,15 +788,17 @@ def _resolve_connection_settings(provider_config: Optional[dict] = None) -> dict
ovcli_values = _ovcli_values_for(provider_config)
def layered(key: str, default: str = "", *, env_authoritative: bool = False) -> str:
env = _env_value(f"OPENVIKING_{key.upper()}")
if env is not None and env_authoritative:
return env
env = os.environ.get(f"OPENVIKING_{key.upper()}")
if env is not None:
env = env.strip()
if env_authoritative:
return env
return env or ovcli_values.get(key) or _clean_config_value(provider_config.get(key)) or default
api_key_env = _env_value("OPENVIKING_API_KEY")
api_key_env = os.environ.get("OPENVIKING_API_KEY")
return {
"endpoint": _normalize_openviking_url(layered("endpoint", _DEFAULT_ENDPOINT)),
"api_key": api_key_env if api_key_env is not None else ovcli_values.get("api_key", ""),
"api_key": api_key_env.strip() if api_key_env is not None else ovcli_values.get("api_key", ""),
"account": layered("account", env_authoritative=True),
"user": layered("user", env_authoritative=True),
"agent": layered("agent", _DEFAULT_AGENT),
@@ -1109,10 +1096,6 @@ from ._setup import ( # noqa: E402,F401 re-exported: tests and callers patch t
_message_text = flatten_message_text # OpenAI-style string/list content -> text
def _text_part(content: str) -> Dict[str, str]:
return {"type": "text", "text": content}
def _tool_part(tool_id: str, tool_name: str, tool_input: Dict[str, Any], tool_status: str, **extra) -> Dict[str, Any]:
return {"type": "tool", "tool_id": tool_id, "tool_name": tool_name, "tool_input": tool_input, **extra, "tool_status": tool_status}
@@ -1222,14 +1205,6 @@ class _TurnUpload:
if env_var_enabled(_SYNC_TRACE_ENV):
logger.info("OpenViking sync_turn trace: " + fmt, *args)
def post_remaining_individually(self, client: _VikingClient) -> None:
path = f"/api/v1/sessions/{self.sid}/messages"
while self.next_index < len(self.batch_messages):
payload = self.batch_messages[self.next_index]
self._trace("POST %s message_index=%d payload=%s", path, self.next_index, json.dumps(payload, ensure_ascii=False))
client.post(path, payload)
self.next_index += 1
def post(self, client: _VikingClient) -> None:
if self.batch_messages:
while self.next_index < len(self.batch_messages):
@@ -1247,7 +1222,12 @@ class _TurnUpload:
self.next_index = batch_end
if self.next_index == len(self.batch_messages):
return
self.provider._post_session_turn(client, self.sid, self.user_content[:4000], _message_text(self.assistant_content)[:4000])
# Plain-text fallback: one user + one assistant message.
assistant_message: Dict[str, Any] = {"role": "assistant", "parts": [{"type": "text", "text": _message_text(self.assistant_content)[:4000]}]}
if self.provider._agent:
assistant_message["peer_id"] = self.provider._agent
client.post(f"/api/v1/sessions/{self.sid}/messages/batch",
{"messages": [{"role": "user", "parts": [{"type": "text", "text": self.user_content[:4000]}]}, assistant_message]})
def run(self) -> None:
try:
@@ -1266,7 +1246,12 @@ class _TurnUpload:
logger.warning("OpenViking structured sync retry failed; writing %d remaining messages individually: %s",
len(self.batch_messages) - self.next_index, retry_error)
try:
self.post_remaining_individually(retry_client)
path = f"/api/v1/sessions/{self.sid}/messages"
while self.next_index < len(self.batch_messages):
payload = self.batch_messages[self.next_index]
self._trace("POST %s message_index=%d payload=%s", path, self.next_index, json.dumps(payload, ensure_ascii=False))
retry_client.post(path, payload)
self.next_index += 1
except Exception as fallback_error:
logger.warning("OpenViking sync_turn failed during individual-message fallback: %s", fallback_error)
@@ -1372,7 +1357,7 @@ class OpenVikingMemoryProvider(MemoryProvider):
return display
display["endpoint"] = settings.get("endpoint") or _DEFAULT_ENDPOINT
display.update({key: settings[key] for key in ("agent", "account", "user") if settings.get(key)})
env_overrides = [key for key in _OPENVIKING_ENV_KEYS if _env_value(key) is not None]
env_overrides = [key for key in _OPENVIKING_ENV_KEYS if key in os.environ]
if env_overrides:
display["env_overrides"] = ", ".join(env_overrides)
return display
@@ -1383,10 +1368,6 @@ class OpenVikingMemoryProvider(MemoryProvider):
# -- connection lifecycle ------------------------------------------------
def _runtime_start_active(self) -> bool:
# Caller holds _runtime_start_lock.
return self._runtime_start_pending or bool(self._runtime_start_thread and self._runtime_start_thread.is_alive())
def _start_runtime_openviking_waiter(self, *, endpoint: str, status_callback=None, warning_callback=None) -> None:
# Caller holds _runtime_start_lock and reserved ownership via _runtime_start_pending.
if self._runtime_start_thread and self._runtime_start_thread.is_alive():
@@ -1459,7 +1440,7 @@ class OpenVikingMemoryProvider(MemoryProvider):
return
with self._runtime_start_lock:
if self._shutting_down or self._runtime_start_active():
if self._shutting_down or self._runtime_start_pending or (self._runtime_start_thread and self._runtime_start_thread.is_alive()):
return
self._runtime_start_pending = True
start_state, start_message = _start_local_openviking_server(endpoint)
@@ -1485,7 +1466,7 @@ class OpenVikingMemoryProvider(MemoryProvider):
except _OpenVikingEndpointError as exc:
connection_error = str(exc)
settings = dict.fromkeys(_CONNECTION_KEYS, "")
self._apply_settings(settings)
self._endpoint, self._api_key, self._account, self._user, self._agent = (settings[k] for k in _CONNECTION_KEYS)
# Baseline established — set here, not at the end, so an exception in the
# connection attempt (swallowed by MemoryManager) can't leave the provider
# stuck in never-refresh mode.
@@ -1520,9 +1501,6 @@ class OpenVikingMemoryProvider(MemoryProvider):
global _last_active_provider # atexit safety net
_last_active_provider = self
def _apply_settings(self, settings: dict) -> None:
self._endpoint, self._api_key, self._account, self._user, self._agent = (settings[k] for k in _CONNECTION_KEYS)
def _ensure_client(self) -> Optional["_VikingClient"]:
"""Active client, rebuilt if the resolved config changed.
@@ -1560,14 +1538,14 @@ class OpenVikingMemoryProvider(MemoryProvider):
if self._client is not None:
return self._client
with self._runtime_start_lock:
if self._runtime_start_active():
if self._runtime_start_pending or (self._runtime_start_thread and self._runtime_start_thread.is_alive()):
return self._client
# Last attempt at this exact config failed: skip the 3s probe until the
# cooldown elapses or the resolved config changes.
if self._in_cooldown(settings_key):
return None
self._apply_settings(settings)
self._endpoint, self._api_key, self._account, self._user, self._agent = settings_key
try:
client = self._build_client()
except ImportError:
@@ -1594,11 +1572,8 @@ class OpenVikingMemoryProvider(MemoryProvider):
"""Client from the published snapshot (one tuple load: background writers run
without _client_refresh_lock and must not see torn fields); falls back to the
raw fields for legacy/hand-wired paths with no snapshot."""
snapshot = self._conn_snapshot
if snapshot is not None:
endpoint, api_key, account, user, agent = snapshot
return _VikingClient(endpoint, api_key, account=account, user=user, agent=agent)
return self._build_client()
endpoint, api_key, account, user, agent = self._conn_snapshot or self._settings_tuple()
return _VikingClient(endpoint, api_key, account=account, user=user, agent=agent)
# -- prompt / prefetch ---------------------------------------------------
@@ -1772,22 +1747,15 @@ class OpenVikingMemoryProvider(MemoryProvider):
return resp.get("result") if isinstance(resp, dict) and "result" in resp else resp
@classmethod
def _extract_text_content(cls, resp: Any) -> str:
"""Text body from a content endpoint (plain string or {content|text} object)."""
result = cls._unwrap_result(resp)
if isinstance(result, str):
return result.strip()
if isinstance(result, dict):
return str(result.get("content") or result.get("text") or "").strip()
return ""
@classmethod
def _extract_read_content(cls, resp: Any) -> str:
"""Like _extract_text_content but only accepts string content/text fields."""
def _extract_text_content(cls, resp: Any, *, strict: bool = False) -> str:
"""Text body from a content endpoint (plain string or {content|text} object);
``strict`` accepts only non-blank string fields."""
result = cls._unwrap_result(resp)
if isinstance(result, str):
return result.strip()
if isinstance(result, dict):
if not strict:
return str(result.get("content") or result.get("text") or "").strip()
for key in ("content", "text"):
value = result.get(key)
if isinstance(value, str) and value.strip():
@@ -1893,7 +1861,7 @@ class OpenVikingMemoryProvider(MemoryProvider):
user = self._user_space(active_client, timeout=self._remaining_recall_timeout(deadline, request_timeout))
except Exception:
return empty
uris = tuple(_user_scoped_uri(user, suffix) for suffix in _SESSION_START_SUFFIXES)
uris = tuple(f"viking://user/{user}/{suffix}" for suffix in _SESSION_START_SUFFIXES)
try:
profile = self._extract_text_content(budgeted_get("/api/v1/content/read", {"uri": uris[0]}))
except Exception as e:
@@ -1951,7 +1919,7 @@ class OpenVikingMemoryProvider(MemoryProvider):
def _build_session_start_memory_block(cls, *, profile: str, preferences: List[Dict[str, str]],
entities: List[Dict[str, str]], token_budget: int, uris: Optional[tuple] = None) -> str:
"""Profile (<= half the budget) then preferences/entities listings sharing the rest."""
profile_uri, preferences_uri, entities_uri = uris or tuple(_user_scoped_uri("default", suffix) for suffix in _SESSION_START_SUFFIXES)
profile_uri, preferences_uri, entities_uri = uris or tuple(f"viking://user/default/{suffix}" for suffix in _SESSION_START_SUFFIXES)
profile = profile.strip()
if not profile and not preferences and not entities:
return ""
@@ -2057,7 +2025,7 @@ class OpenVikingMemoryProvider(MemoryProvider):
try:
timeout = self._remaining_recall_timeout(deadline, request_timeout)
read_state["full_reads"] += 1
content = self._extract_read_content(client.get("/api/v1/content/read", params={"uri": uri}, timeout=timeout))
content = self._extract_text_content(client.get("/api/v1/content/read", params={"uri": uri}, timeout=timeout), strict=True)
if content:
return content
except Exception as e:
@@ -2145,7 +2113,7 @@ class OpenVikingMemoryProvider(MemoryProvider):
continue
flush_tool_parts()
text = _message_text(message.get("content"))
parts: List[Dict[str, Any]] = [_text_part(text)] if text else []
parts: List[Dict[str, Any]] = [{"type": "text", "text": text}] if text else []
if role == "assistant":
for tool_call in message.get("tool_calls") or []:
if not isinstance(tool_call, dict):
@@ -2164,13 +2132,6 @@ class OpenVikingMemoryProvider(MemoryProvider):
flush_tool_parts()
return payload_messages
def _post_session_turn(self, client: _VikingClient, sid: str, user_content: str, assistant_content: str) -> None:
assistant_message: Dict[str, Any] = {"role": "assistant", "parts": [_text_part(assistant_content)]}
if self._agent:
assistant_message["peer_id"] = self._agent
client.post(f"/api/v1/sessions/{sid}/messages/batch",
{"messages": [{"role": "user", "parts": [_text_part(user_content)]}, assistant_message]})
def sync_turn(self, user_content: str, assistant_content: str, *, session_id: str = "",
messages: Optional[List[Dict[str, Any]]] = None) -> None:
"""Record the conversation turn in OpenViking's session (non-blocking)."""
@@ -2245,10 +2206,6 @@ class OpenVikingMemoryProvider(MemoryProvider):
self._spawn_tracked(name, target, self._inflight_lock, lambda: self._inflight_writers.setdefault(sid, set()), after_discard=drop_empty)
def _spawn_deferred_commit(self, name: str, body: Callable[[], None]) -> None:
"""Run ``body`` on a tracked daemon thread (joined by shutdown / _drain_finalizers)."""
self._spawn_tracked(name, body, self._deferred_commit_lock, lambda: self._deferred_commit_threads)
@staticmethod
def _join_all(alive: Callable[[], List[threading.Thread]], timeout: float, *, slice_cap: Optional[float] = None) -> bool:
"""Join threads from ``alive()`` until none remain or the shared budget runs out."""
@@ -2285,13 +2242,6 @@ class OpenVikingMemoryProvider(MemoryProvider):
# -- session commit / pending-session recovery --------------------------
def _session_has_pending_tokens(self, sid: str) -> bool:
try:
session = self._unwrap_result(self._client.get(f"/api/v1/sessions/{sid}"))
return isinstance(session, dict) and int(session.get("pending_tokens") or 0) > 0
except Exception:
return False
def _has_committed_session(self, sid: str) -> bool:
with self._committed_session_lock:
return sid in self._committed_session_ids
@@ -2473,14 +2423,20 @@ class OpenVikingMemoryProvider(MemoryProvider):
finally:
self._flock_close(lock_file, None if owner == self._run_id else self._state_path("lock", owner), "owner run lock")
self._spawn_deferred_commit(f"openviking-recover-owner-{owner_run_id or 'legacy'}", _recover_owner)
self._spawn_tracked(f"openviking-recover-owner-{owner_run_id or 'legacy'}", _recover_owner, self._deferred_commit_lock, lambda: self._deferred_commit_threads)
def _session_needs_commit(self, sid: str, turn_count: int) -> bool:
# The committed-guard wins over turn_count: a racing sync_turn can re-increment
# _turn_count after a commit+reset.
if self._has_committed_session(sid):
return False
return turn_count > 0 or self._session_has_pending_tokens(sid)
if turn_count > 0:
return True
try:
session = self._unwrap_result(self._client.get(f"/api/v1/sessions/{sid}"))
return isinstance(session, dict) and int(session.get("pending_tokens") or 0) > 0
except Exception:
return False
def _commit_session(self, sid: str, turn_count: int, *, context: str, clear_missing: bool = False) -> bool:
try:
@@ -2516,7 +2472,7 @@ class OpenVikingMemoryProvider(MemoryProvider):
finally:
self._claim_deferred_sid(sid, release=True)
self._spawn_deferred_commit(f"openviking-finalize-{sid}", _finalize)
self._spawn_tracked(f"openviking-finalize-{sid}", _finalize, self._deferred_commit_lock, lambda: self._deferred_commit_threads)
def on_session_end(self, messages: List[Dict[str, Any]]) -> None:
"""Commit the session (synchronously — it must land before process exit) to
@@ -2593,7 +2549,7 @@ class OpenVikingMemoryProvider(MemoryProvider):
active_client = client if client is not None else getattr(self, "_client", None)
agent = str(getattr(active_client, "_agent", getattr(self, "_agent", "")) or "").strip()
peer_prefix = f"peers/{agent}/" if agent else ""
return _user_scoped_uri(self._user_space(active_client, timeout=timeout), f"{peer_prefix}memories/{subdir}/mem_{uuid.uuid4().hex[:12]}.md")
return f"viking://user/{self._user_space(active_client, timeout=timeout)}/{peer_prefix}memories/{subdir}/mem_{uuid.uuid4().hex[:12]}.md"
def on_memory_write(self, action: str, target: str, content: str, metadata: Optional[Dict[str, Any]] = None) -> None:
"""Mirror successful built-in memory additions to OpenViking."""
@@ -2792,7 +2748,7 @@ class OpenVikingMemoryProvider(MemoryProvider):
return tool_error("OpenViking server not connected")
session_id = f"hermes-remember-{uuid.uuid4().hex[:12]}"
session_uri = _user_scoped_uri(self._user_space(client), f"sessions/{session_id}")
session_uri = f"viking://user/{self._user_space(client)}/sessions/{session_id}"
def failure(message: str, *, stage: str, message_status: str) -> str:
return tool_error(
@@ -2805,7 +2761,7 @@ class OpenVikingMemoryProvider(MemoryProvider):
),
)
try:
client.post(f"/api/v1/sessions/{session_id}/messages", {"role": "user", "parts": [_text_part(content)]})
client.post(f"/api/v1/sessions/{session_id}/messages", {"role": "user", "parts": [{"type": "text", "text": content}]})
except Exception as e:
logger.error("OpenViking remember message failed for %s: %s", session_id, e)
return failure(f"Memory message submission failed for session {session_id}: {e}", stage="message", message_status="unknown")
@@ -2848,11 +2804,13 @@ class OpenVikingMemoryProvider(MemoryProvider):
parsed_url = urlparse(url)
source_path = None
if parsed_url.scheme == "file" and not _is_remote_resource_source(url):
if url.startswith(_REMOTE_RESOURCE_PREFIXES):
pass
elif parsed_url.scheme == "file":
source_path = _path_from_file_uri(url)
if isinstance(source_path, str):
return tool_error(source_path)
elif not _is_remote_resource_source(url) and (not parsed_url.scheme or _is_windows_absolute_path(url)):
elif not parsed_url.scheme or _is_windows_absolute_path(url):
source_path = Path(url).expanduser()
cleanup_path: Optional[Path] = None