Merge pull request #81 from fabioscarsi/feat/wire-skill-ranker-cache-to-search

fix(search): wire search_skills to SkillRanker embedding cache
This commit is contained in:
Xu Lingrui 2026-05-13 03:08:11 +08:00 committed by GitHub
commit e865030ea8
No known key found for this signature in database
GPG key ID: B5690EEEBB952194
2 changed files with 196 additions and 50 deletions

View file

@ -21,6 +21,21 @@ logger = logging.getLogger("openspace.cloud")
CLOUD_EMBEDDING_SEARCH_MAX_LIMIT = 300
# Shared SkillRanker singleton. Its pickle cache survives process restarts;
# the singleton itself is per-process and avoids reloading the pickle on every
# search_skills invocation.
_shared_ranker = None
def _get_shared_ranker():
"""Lazy-init shared ``SkillRanker`` (with persistent embedding cache)."""
global _shared_ranker
if _shared_ranker is None:
from openspace.skill_engine.skill_ranker import SkillRanker
_shared_ranker = SkillRanker(enable_cache=True)
return _shared_ranker
def _check_safety(text: str) -> list[str]:
"""Lazy wrapper — avoids importing skill_engine at module load time."""
from openspace.skill_engine.skill_utils import check_skill_safety
@ -129,9 +144,9 @@ class SkillSearchEngine:
limit: int,
) -> List[Dict[str, Any]]:
"""BM25 rough-rank to keep top candidates for embedding stage."""
from openspace.skill_engine.skill_ranker import SkillRanker, SkillCandidate
from openspace.skill_engine.skill_ranker import SkillCandidate
ranker = SkillRanker(enable_cache=True)
ranker = _get_shared_ranker()
bm25_candidates = [
SkillCandidate(
skill_id=c.get("skill_id", ""),
@ -424,15 +439,33 @@ async def hybrid_search_skills(
try:
query_embedding = await asyncio.to_thread(generate_embedding, normalized_query)
if query_embedding:
# Route candidate embedding generation through SkillRanker's persistent
# cache (pickle on disk) instead of re-computing on every query.
# Cloud candidates that already carry ``_embedding`` (from server) are
# left untouched.
from openspace.skill_engine.skill_ranker import SkillCandidate
ranker = _get_shared_ranker()
for candidate in candidates:
if not candidate.get("_embedding") and candidate.get("_embedding_text"):
candidate_embedding = await asyncio.to_thread(
generate_embedding, candidate["_embedding_text"],
)
if candidate_embedding:
candidate["_embedding"] = candidate_embedding
except Exception:
pass
if candidate.get("_embedding") or not candidate.get("_embedding_text"):
continue
sid = candidate.get("skill_id") or ""
if not sid:
# Without a stable skill_id the cache would collide; skip.
continue
cand = SkillCandidate(
skill_id=sid,
name=candidate.get("name", ""),
description=candidate.get("description", ""),
body="",
embedding_text=candidate["_embedding_text"],
)
candidate_embedding = await asyncio.to_thread(
ranker.get_or_compute_embedding, cand,
)
if candidate_embedding:
candidate["_embedding"] = candidate_embedding
except Exception as e:
logger.warning(f"hybrid_search_skills: embedding unavailable: {e}")
engine = SkillSearchEngine()
return engine.search(normalized_query, candidates, query_embedding=query_embedding, limit=limit)

View file

@ -18,11 +18,15 @@ Reused by:
from __future__ import annotations
import base64
import hashlib
import json
import math
import os
import pickle
import re
import tempfile
import threading
from dataclasses import dataclass, field
from datetime import datetime
from pathlib import Path
@ -43,8 +47,11 @@ PREFILTER_THRESHOLD = 10
# How many candidates to keep after BM25 rough-rank (before embedding re-rank)
BM25_CANDIDATES_MULTIPLIER = 3 # top_k * 3
# Cache version — increment when format changes
_CACHE_VERSION = 1
# Cache version — increment when format changes.
#
# v2 intentionally does not migrate v1 skill_id-only entries: old entries did
# not store the source text hash, so they cannot be proven fresh.
_CACHE_VERSION = 2
@dataclass
@ -81,9 +88,10 @@ class SkillRanker:
cache_dir: Optional[Path] = None,
enable_cache: bool = True,
) -> None:
# Embedding cache: skill_id → List[float]
# Embedding cache: encoded skill_id + content hash → List[float]
self._embedding_cache: Dict[str, List[float]] = {}
self._enable_cache = enable_cache
self._cache_lock = threading.RLock()
if cache_dir is None:
try:
@ -146,6 +154,36 @@ class SkillRanker:
"""Embedding-only ranking."""
return self._embedding_rank(query, candidates, top_k)
@staticmethod
def _skill_key_prefix(skill_id: str) -> str:
"""Return a collision-safe prefix for all cache entries of a skill."""
encoded = base64.urlsafe_b64encode(skill_id.encode("utf-8")).decode("ascii")
return f"{encoded}:"
@classmethod
def _content_key(cls, skill_id: str, embedding_text: str) -> str:
"""Compose a content-addressed cache key.
The cache stores embeddings under ``"{encoded_skill_id}:{sha256(text)[:16]}"``
so that any change to the text used for embedding (name, description,
or body) automatically invalidates the cached entry without relying
on an external trigger.
"""
digest = hashlib.sha256(embedding_text.encode("utf-8")).hexdigest()[:16]
return f"{cls._skill_key_prefix(skill_id)}{digest}"
def _drop_stale_entries_locked(self, skill_id: str, keep_key: str) -> None:
"""Drop older content-addressed entries for a skill.
Caller must hold ``self._cache_lock``.
"""
prefix = self._skill_key_prefix(skill_id)
for stale in [
k for k in self._embedding_cache
if k != keep_key and k.startswith(prefix)
]:
self._embedding_cache.pop(stale, None)
def get_or_compute_embedding(
self, candidate: SkillCandidate,
) -> Optional[List[float]]:
@ -157,30 +195,65 @@ class SkillRanker:
if candidate.embedding:
return candidate.embedding
# Check cache
cached = self._embedding_cache.get(candidate.skill_id)
if cached:
candidate.embedding = cached
return cached
text = self._build_embedding_text(candidate)
if not candidate.skill_id:
emb = self._generate_embedding(text)
if emb:
candidate.embedding = emb
return emb
content_key = self._content_key(candidate.skill_id, text)
dropped_legacy = False
with self._cache_lock:
# Check content-addressed cache. If a legacy in-memory entry exists,
# discard it instead of migrating unverifiable stale data.
cached = self._embedding_cache.get(content_key)
if cached:
candidate.embedding = cached
return cached
dropped_legacy = self._embedding_cache.pop(candidate.skill_id, None) is not None
# Compute
text = self._build_embedding_text(candidate)
emb = self._generate_embedding(text)
if emb:
candidate.embedding = emb
self._embedding_cache[candidate.skill_id] = emb
self._save_cache()
with self._cache_lock:
cached = self._embedding_cache.get(content_key)
if cached:
candidate.embedding = cached
return cached
if emb:
candidate.embedding = emb
self._embedding_cache[content_key] = emb
# Bound cache growth: previous versions of this skill are now
# obsolete, drop them in the same write.
self._drop_stale_entries_locked(candidate.skill_id, content_key)
self._save_cache()
elif dropped_legacy:
self._save_cache()
return emb
def invalidate_cache(self, skill_id: str) -> None:
"""Remove a skill's cached embedding (e.g. after evolution)."""
self._embedding_cache.pop(skill_id, None)
self._save_cache()
"""Remove all cached embeddings for a skill (e.g. after evolution).
Removes every cache entry whose key matches either the exact
``skill_id`` (legacy format) or the ``"{skill_id}:*"`` content-addressed
prefix, covering any historical content version that might linger.
"""
prefix = self._skill_key_prefix(skill_id)
with self._cache_lock:
keys_to_drop = [
k for k in self._embedding_cache
if k == skill_id or k.startswith(prefix)
]
for k in keys_to_drop:
self._embedding_cache.pop(k, None)
if keys_to_drop:
self._save_cache()
def clear_cache(self) -> None:
"""Clear all cached embeddings."""
self._embedding_cache.clear()
self._save_cache()
with self._cache_lock:
self._embedding_cache.clear()
self._save_cache()
@staticmethod
def _tokenize(text: str) -> List[str]:
@ -273,21 +346,41 @@ class SkillRanker:
if not query_emb:
return []
# Ensure all candidates have embeddings
# Ensure all candidates have embeddings (content-addressed cache
# with backward-compat fallback for legacy skill_id-only keys).
cache_dirty = False
for c in candidates:
if not c.embedding:
cached = self._embedding_cache.get(c.skill_id)
if c.embedding:
continue
text = self._build_embedding_text(c)
if not c.skill_id:
emb = self._generate_embedding(text, api_key=api_key)
if emb:
c.embedding = emb
continue
content_key = self._content_key(c.skill_id, text)
with self._cache_lock:
cached = self._embedding_cache.get(content_key)
if cached:
c.embedding = cached
else:
text = self._build_embedding_text(c)
emb = self._generate_embedding(text, api_key=api_key)
if emb:
c.embedding = emb
self._embedding_cache[c.skill_id] = emb
continue
# Do not migrate legacy skill_id-only entries; they may be
# stale because the old format did not store a text hash.
if self._embedding_cache.pop(c.skill_id, None) is not None:
cache_dirty = True
emb = self._generate_embedding(text, api_key=api_key)
if emb:
c.embedding = emb
with self._cache_lock:
self._embedding_cache[content_key] = emb
# Bound cache growth: previous versions of this skill are
# obsolete, drop them.
self._drop_stale_entries_locked(c.skill_id, content_key)
cache_dirty = True
# Save newly computed embeddings
self._save_cache()
# Save newly computed / migrated embeddings
if cache_dirty:
self._save_cache()
# Score
for c in candidates:
@ -361,8 +454,14 @@ class SkillRanker:
try:
with open(path, "rb") as f:
data = pickle.load(f)
if isinstance(data, dict) and data.get("version") == _CACHE_VERSION:
self._embedding_cache = data.get("embeddings", {})
if (
isinstance(data, dict)
and data.get("version") == _CACHE_VERSION
and data.get("model") == SKILL_EMBEDDING_MODEL
and isinstance(data.get("embeddings"), dict)
):
with self._cache_lock:
self._embedding_cache = data.get("embeddings", {})
logger.debug(f"Loaded {len(self._embedding_cache)} skill embeddings from cache")
except Exception as e:
logger.warning(f"Failed to load skill embedding cache: {e}")
@ -370,18 +469,32 @@ class SkillRanker:
def _save_cache(self) -> None:
"""Persist embedding cache to disk."""
if not self._enable_cache or not self._embedding_cache:
if not self._enable_cache:
return
try:
self._cache_dir.mkdir(parents=True, exist_ok=True)
data = {
"version": _CACHE_VERSION,
"model": SKILL_EMBEDDING_MODEL,
"last_updated": datetime.now().isoformat(),
"embeddings": self._embedding_cache,
}
with open(self._cache_file(), "wb") as f:
pickle.dump(data, f, protocol=pickle.HIGHEST_PROTOCOL)
with self._cache_lock:
data = {
"version": _CACHE_VERSION,
"model": SKILL_EMBEDDING_MODEL,
"last_updated": datetime.now().isoformat(),
"embeddings": dict(self._embedding_cache),
}
tmp_name = ""
try:
with tempfile.NamedTemporaryFile(
"wb",
dir=self._cache_dir,
prefix=".skill_embeddings_",
suffix=".tmp",
delete=False,
) as f:
tmp_name = f.name
pickle.dump(data, f, protocol=pickle.HIGHEST_PROTOCOL)
os.replace(tmp_name, self._cache_file())
finally:
if tmp_name and os.path.exists(tmp_name):
os.unlink(tmp_name)
except Exception as e:
logger.warning(f"Failed to save skill embedding cache: {e}")