refactor(tools): reuse skill_usage._atomic_write for sync state; AST-neutral bracket hugging
This commit is contained in:
+16
-37
@@ -19,7 +19,6 @@ from __future__ import annotations
|
||||
import json
|
||||
import logging
|
||||
import os
|
||||
import tempfile
|
||||
from contextlib import suppress
|
||||
from pathlib import Path, PurePosixPath
|
||||
from typing import Any, Callable, Dict, List, Optional, Tuple
|
||||
@@ -31,8 +30,7 @@ from tools.skills_sync_client_wire import ( # noqa: F401 (re-exports)
|
||||
assemble_root_from_skill_trees, build_commit, build_root_tree, build_sync_manifest_bytes,
|
||||
build_tree, canonical_json_bytes, materialize_tree, merge_skill, nest_skill_tree,
|
||||
parse_sync_manifest, read_manifest_of_root, read_ref_hash, root_tree_of_commit,
|
||||
skill_trees_of_root, wire_address,
|
||||
)
|
||||
skill_trees_of_root, wire_address)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -73,8 +71,7 @@ def resolve_identity() -> Dict[str, Any]:
|
||||
owner = claims.get("sub") or claims.get("privy_did") or claims.get("tid") or "unknown"
|
||||
return {
|
||||
"api_key": api_key, "base_url": creds.get("base_url"), "owner": str(owner),
|
||||
"nous_admin": claims.get(NOUS_ADMIN_CLAIM) is True, "claims": claims,
|
||||
}
|
||||
"nous_admin": claims.get(NOUS_ADMIN_CLAIM) is True, "claims": claims}
|
||||
|
||||
|
||||
# Configuration -- env-first so a Hermes Cloud instance can enable sync purely through the
|
||||
@@ -356,20 +353,12 @@ def read_sync_state() -> Dict[str, Any]:
|
||||
|
||||
def write_sync_state(data: Dict[str, Any]) -> None:
|
||||
"""Write the local sync state atomically. Best-effort."""
|
||||
path = _sync_state_path()
|
||||
try:
|
||||
path.parent.mkdir(parents=True, exist_ok=True)
|
||||
fd, tmp = tempfile.mkstemp(dir=str(path.parent), prefix=".sync_state_", suffix=".tmp")
|
||||
try:
|
||||
with os.fdopen(fd, "w", encoding="utf-8") as f:
|
||||
json.dump(data, f, indent=2, sort_keys=True, ensure_ascii=False)
|
||||
f.flush()
|
||||
os.fsync(f.fileno())
|
||||
os.replace(tmp, path)
|
||||
except BaseException:
|
||||
with suppress(OSError):
|
||||
os.unlink(tmp)
|
||||
raise
|
||||
from tools.skill_usage import _atomic_write
|
||||
|
||||
_atomic_write(
|
||||
_sync_state_path(), ".sync_state_",
|
||||
lambda f: json.dump(data, f, indent=2, sort_keys=True, ensure_ascii=False))
|
||||
except Exception as e:
|
||||
logger.debug("skills_sync_client: sync state write failed: %s", e)
|
||||
|
||||
@@ -435,8 +424,7 @@ def push_skills(
|
||||
*,
|
||||
skill_names: Optional[List[str]] = None,
|
||||
identity: Optional[Dict[str, Any]] = None,
|
||||
message: str = "hermes skill sync",
|
||||
) -> Dict[str, Any]:
|
||||
message: str = "hermes skill sync") -> Dict[str, Any]:
|
||||
"""Push opted-in skills to ``refs/user/<owner>/HEAD``: upload new objects, CAS HEAD. A 409 with
|
||||
an actual head -> three-way merge + one retry; a 409 against a NON-EXISTENT ref (stale local
|
||||
head, e.g. from another plane) -> redo the CAS as a create. Never raises for inert/no-op cases."""
|
||||
@@ -461,8 +449,7 @@ def push_skills(
|
||||
|
||||
commit_hash = build_commit(
|
||||
root_hash, [base_head] if base_head else [], owner=owner, device=stable_device_id(),
|
||||
message=message, objects=objects,
|
||||
)
|
||||
message=message, objects=objects)
|
||||
client.put_objects(objects.objects)
|
||||
ref = user_head_ref(owner)
|
||||
result = {"ok": True, "head": commit_hash, "pushed_objects": len(objects)}
|
||||
@@ -487,8 +474,7 @@ def _resolve_push_conflict(
|
||||
our_commit: str,
|
||||
objects: ObjectSet,
|
||||
message: str,
|
||||
base_head: Optional[str],
|
||||
) -> Dict[str, Any]:
|
||||
base_head: Optional[str]) -> Dict[str, Any]:
|
||||
"""Three-way merge per skill against the base we forked from (see merge_skill): each side
|
||||
changing a DIFFERENT skill -> merge commit (2 parents) + CAS retry; both changing the SAME
|
||||
skill -> TRUE OVERLAP -> written to refs/user/<owner>/conflict/<n> for out-of-band resolution."""
|
||||
@@ -519,9 +505,7 @@ def _resolve_push_conflict(
|
||||
"overlapping_skills": sorted(overlaps), "actual_head": actual_head,
|
||||
"message": (
|
||||
f"{len(overlaps)} skill(s) changed on both sides; wrote "
|
||||
f"{conflict_ref}. Resolve out-of-band (hermes sync / NAS UI)."
|
||||
),
|
||||
}
|
||||
f"{conflict_ref}. Resolve out-of-band (hermes sync / NAS UI).")}
|
||||
|
||||
# Merge commit (parents: actual, ours); re-add our objects so the merge push is self-contained.
|
||||
merge_objects = ObjectSet()
|
||||
@@ -529,16 +513,14 @@ def _resolve_push_conflict(
|
||||
merged_root = assemble_root_from_skill_trees(merged, merge_objects)
|
||||
merge_commit = build_commit(
|
||||
merged_root, [actual_head, our_commit], owner=owner, device=stable_device_id(),
|
||||
message=f"merge: {message}", objects=merge_objects,
|
||||
)
|
||||
message=f"merge: {message}", objects=merge_objects)
|
||||
client.put_objects(merge_objects.objects)
|
||||
try:
|
||||
client.cas_ref(user_head_ref(owner), actual_head, merge_commit)
|
||||
except SyncConflict as c2:
|
||||
return {
|
||||
"ok": False, "conflict": True, "actual_head": c2.actual,
|
||||
"message": f"merge CAS lost again (head now {c2.actual}); retry sync.",
|
||||
}
|
||||
"message": f"merge CAS lost again (head now {c2.actual}); retry sync."}
|
||||
_record_head(read_sync_state(), merge_commit, merged_root)
|
||||
return {"ok": True, "head": merge_commit, "merged": True}
|
||||
|
||||
@@ -619,8 +601,7 @@ def sync_status() -> Dict[str, Any]:
|
||||
"nous_admin": False, "logged_in": False, "feature_enabled": sync_feature_enabled(),
|
||||
"default_opt_in": sync_default_opt_in(), "base_url": resolve_sync_base_url(),
|
||||
"opted_in_skills": [], "local_head": None, "owner": None, "org_available": False,
|
||||
"org_id": None, "org_role": None, "org_skills": [], "org_skills_modified": [],
|
||||
}
|
||||
"org_id": None, "org_role": None, "org_skills": [], "org_skills_modified": []}
|
||||
try:
|
||||
identity = resolve_identity()
|
||||
status.update(logged_in=True, owner=identity.get("owner"), nous_admin=bool(identity.get("nous_admin")))
|
||||
@@ -638,8 +619,7 @@ def sync_status() -> Dict[str, Any]:
|
||||
status.update(
|
||||
org_available=True, org_id=org_identity.get("org_id"), org_role=org_identity.get("org_role"),
|
||||
org_skills=list_org_skill_names(),
|
||||
org_skills_modified=list_locally_modified_org_skills(org_identity.get("org_id")),
|
||||
)
|
||||
org_skills_modified=list_locally_modified_org_skills(org_identity.get("org_id")))
|
||||
except SyncInertError:
|
||||
pass
|
||||
except Exception as e:
|
||||
@@ -654,5 +634,4 @@ from tools.skills_sync_client_org import ( # noqa: E402,F401 (re-exports)
|
||||
_read_org_baseline, _read_org_head, _skill_dir_fingerprint, _write_active_org_marker,
|
||||
_write_org_baseline, _write_org_provenance, list_locally_modified_org_skills,
|
||||
list_org_skill_names, maybe_pull_org_skills, org_head_ref, org_skill_is_locally_modified,
|
||||
propose_skill, pull_org_skills, resolve_org_identity,
|
||||
)
|
||||
propose_skill, pull_org_skills, resolve_org_identity)
|
||||
|
||||
@@ -21,8 +21,7 @@ from typing import Any, Callable, Dict, List, Optional
|
||||
from tools.skills_sync_client_wire import (
|
||||
DEFAULT_MAX_OBJECT_BYTES, ObjectSet, SyncClient, SyncConflict, SyncError, _check_version,
|
||||
assemble_root_from_skill_trees, build_commit, build_tree, materialize_tree, read_ref_hash,
|
||||
root_tree_of_commit, skill_trees_of_root,
|
||||
)
|
||||
root_tree_of_commit, skill_trees_of_root)
|
||||
|
||||
logger = logging.getLogger("tools.skills_sync_client")
|
||||
|
||||
@@ -155,8 +154,7 @@ def _clear_active_org_marker() -> None:
|
||||
marker.unlink()
|
||||
logger.info(
|
||||
"skills_sync_client: cleared active-org marker "
|
||||
"(token has no org workflow); org skills no longer resolve"
|
||||
)
|
||||
"(token has no org workflow); org skills no longer resolve")
|
||||
except Exception as e:
|
||||
logger.debug("skills_sync_client: marker clear failed: %s", e)
|
||||
|
||||
@@ -258,8 +256,7 @@ def pull_org_skills(
|
||||
if conflicted:
|
||||
logger.warning(
|
||||
"skills_sync_client: %d org skill(s) have local edits AND upstream "
|
||||
"changes; left untouched: %s", len(conflicted), ", ".join(conflicted),
|
||||
)
|
||||
"changes; left untouched: %s", len(conflicted), ", ".join(conflicted))
|
||||
return {"ok": True, "org_id": org_id, "head": head, "updated": updated, "conflicted": conflicted}
|
||||
|
||||
|
||||
@@ -268,8 +265,7 @@ def propose_skill(
|
||||
client: Optional[SyncClient] = None,
|
||||
*,
|
||||
identity: Optional[Dict[str, Any]] = None,
|
||||
message: Optional[str] = None,
|
||||
) -> Dict[str, Any]:
|
||||
message: Optional[str] = None) -> Dict[str, Any]:
|
||||
"""Propose a local (personal) skill's content to the org canonical set.
|
||||
|
||||
Snapshots the skill dir as an org-scoped commit splicing that ONE skill subtree into the
|
||||
@@ -298,8 +294,7 @@ def propose_skill(
|
||||
base_head = _read_org_head(client, org_id)
|
||||
skill_map = (
|
||||
skill_trees_of_root(client, root_tree_of_commit(client, base_head, org_scope=True), org_scope=True)
|
||||
if base_head else {}
|
||||
)
|
||||
if base_head else {})
|
||||
skill_map[str(rel)] = skill_tree
|
||||
root_hash = assemble_root_from_skill_trees(skill_map, objects)
|
||||
commit_hash = build_commit(
|
||||
@@ -323,8 +318,7 @@ def propose_skill(
|
||||
if result.get("proposal_pending"):
|
||||
return {
|
||||
"ok": True, "proposal_pending": True, "proposal_id": result.get("proposal_id"),
|
||||
"ref": result.get("ref"), "commit": commit_hash, "org_id": org_id,
|
||||
}
|
||||
"ref": result.get("ref"), "commit": commit_hash, "org_id": org_id}
|
||||
return {"ok": True, "merged": True, "head": result.get("hash", commit_hash), "commit": commit_hash, "org_id": org_id}
|
||||
|
||||
|
||||
|
||||
@@ -145,16 +145,14 @@ def build_tree(dir_path: Path, objects: ObjectSet, *, max_object_bytes: int) ->
|
||||
|
||||
def build_commit(
|
||||
tree_hash: str, parents: List[str], *, owner: str, device: str, message: str,
|
||||
objects: ObjectSet, ts: Optional[str] = None,
|
||||
) -> str:
|
||||
objects: ObjectSet, ts: Optional[str] = None) -> str:
|
||||
"""Build a commit object and return its address. ``parents``: 0 for the first commit, 1 for
|
||||
an edit, 2 for a merge (parents[0] = base fast-forwarded from, parents[1] = other head)."""
|
||||
return objects.add(KIND_COMMIT, canonical_json_bytes({
|
||||
"type": KIND_COMMIT, "tree": tree_hash, "parents": list(parents),
|
||||
"author": {"owner": owner, "device": device},
|
||||
"ts": ts or datetime.now(timezone.utc).strftime("%Y-%m-%dT%H:%M:%SZ"),
|
||||
"message": message, "artifact_type": ARTIFACT_TYPE_SKILL,
|
||||
}))
|
||||
"message": message, "artifact_type": ARTIFACT_TYPE_SKILL}))
|
||||
|
||||
|
||||
def build_root_tree(node: Dict[str, Any], objects: ObjectSet, *, manifest_hash: Optional[str] = None) -> str:
|
||||
@@ -217,8 +215,7 @@ def _check_version(caps: Dict[str, Any]) -> None:
|
||||
if ver.split(".", 1)[0] != WIRE_VERSION:
|
||||
raise SyncError(
|
||||
f"this server speaks sync version {ver!r}, but this Hermes speaks "
|
||||
f"{WIRE_VERSION} — update Hermes to sync with it"
|
||||
)
|
||||
f"{WIRE_VERSION} — update Hermes to sync with it")
|
||||
|
||||
|
||||
def _body(r) -> Dict[str, Any]:
|
||||
@@ -272,8 +269,7 @@ class SyncClient:
|
||||
"""GET objects/:hash -> ``(kind, bytes)``; kind from the object-type header, blob default."""
|
||||
r = self._request(
|
||||
"GET", f"org/objects/{obj_hash}" if org_scope else f"objects/{obj_hash}", "get_object",
|
||||
errors={404: f"object {obj_hash} not found", 403: f"object {obj_hash} not readable"},
|
||||
)
|
||||
errors={404: f"object {obj_hash} not found", 403: f"object {obj_hash} not readable"})
|
||||
return r.headers.get("X-HSP-Object-Type") or KIND_BLOB, r.content
|
||||
|
||||
def _get_json_of_kind(self, obj_hash: str, expected: str, org_scope: bool) -> Dict[str, Any]:
|
||||
@@ -298,8 +294,7 @@ class SyncClient:
|
||||
r = self._request(
|
||||
"POST", "objects", "put_objects", ok=(200, 201), files=files,
|
||||
params={"scope": "org"} if org_scope else None,
|
||||
errors={413: "object too large (413)", 422: lambda r: f"hash_mismatch (422): {r.text}"},
|
||||
)
|
||||
errors={413: "object too large (413)", 422: lambda r: f"hash_mismatch (422): {r.text}"})
|
||||
return _body(r)
|
||||
|
||||
def cas_ref(self, name: str, from_hash: Optional[str], to_hash: str) -> Dict[str, Any]:
|
||||
@@ -308,8 +303,7 @@ class SyncClient:
|
||||
``{"proposal_pending": True, ...}``: SUCCESS-shaped but never to be presented as live/merged."""
|
||||
r = self._request(
|
||||
"POST", f"refs/{name}", "cas_ref", ok=(200, 202, 409), json={"from": from_hash, "to": to_hash},
|
||||
errors={403: "forbidden (403) -- owner/permission"},
|
||||
)
|
||||
errors={403: "forbidden (403) -- owner/permission"})
|
||||
if r.status_code == 202:
|
||||
return {"proposal_pending": True, **_body(r)}
|
||||
if r.status_code == 409: # "" actual = the ref does not exist server-side (-> None)
|
||||
|
||||
@@ -95,8 +95,7 @@ _PATTERNS: List[Tuple[str, str, str]] = [
|
||||
# plus, BOM, LTR/RTL embedding + pop + overrides, LTR/RTL/first-strong isolates + pop.
|
||||
INVISIBLE_CHARS = frozenset(
|
||||
"\u200b\u200c\u200d\u2060\u2062\u2063\u2064\ufeff"
|
||||
"\u202a\u202b\u202c\u202d\u202e\u2066\u2067\u2068\u2069"
|
||||
)
|
||||
"\u202a\u202b\u202c\u202d\u202e\u2066\u2067\u2068\u2069")
|
||||
|
||||
# Compiled per scope at import; inclusion is cumulative (all ⊂ context ⊂ strict).
|
||||
_SCOPE_SETS = {"all": ("all", "context", "strict"), "context": ("context", "strict"), "strict": ("strict",)}
|
||||
@@ -146,8 +145,7 @@ def first_threat_message(content: str, scope: str = "strict") -> Optional[str]:
|
||||
return (
|
||||
f"Blocked: content matches threat pattern '{pid}'. "
|
||||
f"Content is injected into the system prompt and must not contain "
|
||||
f"injection or exfiltration payloads."
|
||||
)
|
||||
f"injection or exfiltration payloads.")
|
||||
|
||||
|
||||
__all__ = ["INVISIBLE_CHARS", "MAX_SCAN_CHARS", "scan_for_threats", "first_threat_message"]
|
||||
|
||||
@@ -58,8 +58,7 @@ def _load_security_config() -> dict:
|
||||
"tirith_enabled": _env_bool("TIRITH_ENABLED", cfg.get("tirith_enabled", True)),
|
||||
"tirith_path": os.getenv("TIRITH_BIN", cfg.get("tirith_path", "tirith")),
|
||||
"tirith_timeout": _env_int("TIRITH_TIMEOUT", cfg.get("tirith_timeout", 5)),
|
||||
"tirith_fail_open": _env_bool("TIRITH_FAIL_OPEN", cfg.get("tirith_fail_open", True)),
|
||||
}
|
||||
"tirith_fail_open": _env_bool("TIRITH_FAIL_OPEN", cfg.get("tirith_fail_open", True))}
|
||||
|
||||
|
||||
# --- Module state ---
|
||||
@@ -97,8 +96,7 @@ def _record_tirith_crash() -> None:
|
||||
logger.warning(
|
||||
"tirith circuit breaker opened after %d consecutive failures; "
|
||||
"disabling for the rest of the process",
|
||||
_crash_count,
|
||||
)
|
||||
_crash_count)
|
||||
|
||||
|
||||
def _warn_once(key: str, message: str, *args) -> None:
|
||||
@@ -243,8 +241,7 @@ def _verify_cosign(checksums_path: str, sig_path: str, cert_path: str) -> bool |
|
||||
"--certificate-identity-regexp", _COSIGN_IDENTITY_REGEXP,
|
||||
"--certificate-oidc-issuer", _COSIGN_ISSUER, checksums_path],
|
||||
capture_output=True, text=True, encoding='utf-8', errors='replace',
|
||||
timeout=15, stdin=subprocess.DEVNULL,
|
||||
)
|
||||
timeout=15, stdin=subprocess.DEVNULL)
|
||||
except (OSError, subprocess.TimeoutExpired) as exc:
|
||||
logger.warning("cosign execution failed: %s", exc)
|
||||
return None
|
||||
@@ -491,8 +488,7 @@ def ensure_installed(*, log_failures: bool = True):
|
||||
return found
|
||||
if _install_thread is None or not _install_thread.is_alive():
|
||||
_install_thread = threading.Thread(
|
||||
target=_background_install, kwargs={"log_failures": log_failures}, daemon=True
|
||||
)
|
||||
target=_background_install, kwargs={"log_failures": log_failures}, daemon=True)
|
||||
_install_thread.start()
|
||||
return None # not available yet; commands fail-open until ready
|
||||
|
||||
@@ -505,8 +501,7 @@ _EXIT_ACTIONS = {0: "allow", 1: "block", 2: "warn"}
|
||||
# Summary when tirith's JSON is unparseable and only the exit code is known.
|
||||
_NO_DETAILS_SUMMARY = {
|
||||
"block": "security issue detected (details unavailable)",
|
||||
"warn": "security warning detected (details unavailable)",
|
||||
}
|
||||
"warn": "security warning detected (details unavailable)"}
|
||||
|
||||
|
||||
def _verdict(action: str, summary: str = "", findings: list | None = None) -> dict:
|
||||
@@ -539,8 +534,7 @@ def check_command_security(command: str) -> dict:
|
||||
result = subprocess.run(
|
||||
[tirith_path, "check", "--json", "--non-interactive", "--shell", "posix", "--", command],
|
||||
capture_output=True, text=True, encoding='utf-8', errors='replace',
|
||||
timeout=timeout, stdin=subprocess.DEVNULL,
|
||||
)
|
||||
timeout=timeout, stdin=subprocess.DEVNULL)
|
||||
except OSError as exc:
|
||||
# FileNotFoundError / PermissionError / exec format error: dedupe by (class, errno)
|
||||
# so each failure mode surfaces once, not per command.
|
||||
@@ -584,5 +578,4 @@ def _is_app_tld_finding(finding: dict) -> bool:
|
||||
return False
|
||||
return any(
|
||||
val is not None and ".app" in str(val).lower()
|
||||
for val in (finding.get(k) for k in ("value", "tld", "detail", "description", "message"))
|
||||
)
|
||||
for val in (finding.get(k) for k in ("value", "tld", "detail", "description", "message")))
|
||||
|
||||
Reference in New Issue
Block a user