feat(system): DPLAN-0170 memory system fix — vectorization, dedup, config rename, rollover hook
Co-Authored-By: @devpulse <devpulse@aipass>
This commit is contained in:
@@ -0,0 +1 @@
|
||||
{"file": "/home/patrick/Projects/AIPass/src/aipass/memory/apps/handlers/search/vector_search.py", "errors": [{"line": 42, "message": "E402: Module level import not at top of file"}]}
|
||||
Executable
+173
@@ -0,0 +1,173 @@
|
||||
#!/usr/bin/env python3
|
||||
"""
|
||||
Pre-Compact Rollover Hook — check branch memory files and run rollover if overdue.
|
||||
|
||||
Runs alongside pre_compact.py on PreCompact events. Scans all branches'
|
||||
.trinity files for over-limit conditions and executes rollover via drone
|
||||
if any are found. Stdout stays clean (pre_compact.py owns stdout for
|
||||
context injection). All logging goes to stderr.
|
||||
|
||||
Version: 1.0.0
|
||||
"""
|
||||
|
||||
import json
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
|
||||
def _find_repo_root():
|
||||
"""Find the AIPass repo root (contains AIPASS_REGISTRY.json)."""
|
||||
current = Path(__file__).resolve().parent
|
||||
for parent in [current] + list(current.parents):
|
||||
if (parent / "AIPASS_REGISTRY.json").exists():
|
||||
return parent
|
||||
cwd = Path.cwd()
|
||||
for parent in [cwd] + list(cwd.parents):
|
||||
if (parent / "AIPASS_REGISTRY.json").exists():
|
||||
return parent
|
||||
return None
|
||||
|
||||
|
||||
def _read_registry(repo_root):
|
||||
"""Read branch list from AIPASS_REGISTRY.json."""
|
||||
registry_path = repo_root / "AIPASS_REGISTRY.json"
|
||||
if not registry_path.exists():
|
||||
return []
|
||||
try:
|
||||
data = json.loads(registry_path.read_text(encoding="utf-8"))
|
||||
branches = data.get("branches", [])
|
||||
for branch in branches:
|
||||
raw_path = branch.get("path", "")
|
||||
resolved = Path(raw_path)
|
||||
if not resolved.is_absolute():
|
||||
resolved = repo_root / raw_path
|
||||
branch["_resolved_path"] = resolved
|
||||
return branches
|
||||
except Exception:
|
||||
return []
|
||||
|
||||
|
||||
def _check_file(file_path):
|
||||
"""Check if a .trinity memory file is overdue for rollover.
|
||||
|
||||
Returns (overdue: bool, description: str) or (False, "") if not overdue.
|
||||
"""
|
||||
if not file_path.is_file():
|
||||
return False, ""
|
||||
|
||||
try:
|
||||
raw = file_path.read_text(encoding="utf-8")
|
||||
data = json.loads(raw)
|
||||
except Exception:
|
||||
return False, ""
|
||||
|
||||
metadata = data.get("document_metadata", {})
|
||||
schema_version = metadata.get("schema_version", "1.0.0")
|
||||
limits = metadata.get("limits", {})
|
||||
|
||||
if schema_version.startswith("2"):
|
||||
reasons = []
|
||||
max_sessions = limits.get("max_sessions")
|
||||
if max_sessions is not None:
|
||||
sessions = data.get("sessions", [])
|
||||
if isinstance(sessions, list) and len(sessions) >= max_sessions:
|
||||
reasons.append(f"{len(sessions)}/{max_sessions} sessions")
|
||||
|
||||
max_key_learnings = limits.get("max_key_learnings")
|
||||
if max_key_learnings is not None:
|
||||
key_learnings = data.get("key_learnings", {})
|
||||
if isinstance(key_learnings, dict) and len(key_learnings) >= max_key_learnings:
|
||||
reasons.append(f"{len(key_learnings)}/{max_key_learnings} learnings")
|
||||
|
||||
max_observations = limits.get("max_observations")
|
||||
if max_observations is not None:
|
||||
observations = data.get("observations", [])
|
||||
if isinstance(observations, list) and len(observations) >= max_observations:
|
||||
reasons.append(f"{len(observations)}/{max_observations} observations")
|
||||
|
||||
if reasons:
|
||||
return True, ", ".join(reasons)
|
||||
return False, ""
|
||||
|
||||
# v1: line-count based
|
||||
max_lines = limits.get("max_lines", 600)
|
||||
current_lines = raw.count("\n") + 1
|
||||
if current_lines >= max_lines:
|
||||
return True, f"{current_lines}/{max_lines} lines"
|
||||
return False, ""
|
||||
|
||||
|
||||
def _find_overdue(repo_root):
|
||||
"""Scan all branches for overdue memory files. Returns list of (branch, type, reason)."""
|
||||
branches = _read_registry(repo_root)
|
||||
overdue = []
|
||||
|
||||
for branch in branches:
|
||||
name = branch.get("name", "unknown")
|
||||
branch_path = branch.get("_resolved_path")
|
||||
if not branch_path or not branch_path.is_dir():
|
||||
continue
|
||||
|
||||
for memory_type in ["local", "observations"]:
|
||||
file_path = branch_path / ".trinity" / f"{memory_type}.json"
|
||||
is_overdue, reason = _check_file(file_path)
|
||||
if is_overdue:
|
||||
overdue.append((name, memory_type, reason))
|
||||
|
||||
return overdue
|
||||
|
||||
|
||||
def _run_rollover(repo_root):
|
||||
"""Execute rollover via drone subprocess. Returns (success, output)."""
|
||||
try:
|
||||
result = subprocess.run(
|
||||
["drone", "@memory", "rollover", "run"],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
timeout=110,
|
||||
cwd=str(repo_root),
|
||||
)
|
||||
return result.returncode == 0, result.stdout + result.stderr
|
||||
except subprocess.TimeoutExpired:
|
||||
return False, "Rollover timed out (110s)"
|
||||
except Exception as e:
|
||||
return False, str(e)
|
||||
|
||||
|
||||
def main():
|
||||
"""Main hook entry point."""
|
||||
try:
|
||||
json.load(sys.stdin)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
try:
|
||||
repo_root = _find_repo_root()
|
||||
if not repo_root:
|
||||
sys.exit(0)
|
||||
|
||||
overdue = _find_overdue(repo_root)
|
||||
if not overdue:
|
||||
sys.exit(0)
|
||||
|
||||
summary = "; ".join(f"{name}.{mtype} ({reason})" for name, mtype, reason in overdue)
|
||||
print(f"Pre-compact rollover: {len(overdue)} overdue — {summary}", file=sys.stderr)
|
||||
|
||||
success, output = _run_rollover(repo_root)
|
||||
if success:
|
||||
print(f"Pre-compact rollover: complete ({len(overdue)} files processed)", file=sys.stderr)
|
||||
else:
|
||||
print(f"Pre-compact rollover: failed — {output[:200]}", file=sys.stderr)
|
||||
|
||||
except Exception as e:
|
||||
print(f"Pre-compact rollover error: {e}", file=sys.stderr)
|
||||
|
||||
sys.exit(0)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
sys.path.insert(0, str(Path(__file__).resolve().parent))
|
||||
from hook_log import run_and_log
|
||||
|
||||
run_and_log("PreCompact", "provider", __file__, main)
|
||||
@@ -17,7 +17,9 @@
|
||||
{"script": "stop_sound.py", "event": "Stop", "source": "repo"},
|
||||
{"script": "notification_sound.py", "event": "Notification", "source": "repo"},
|
||||
{"script": "pre_compact.py", "event": "PreCompact", "matcher": "manual", "source": "repo", "timeout": 60},
|
||||
{"script": "pre_compact.py", "event": "PreCompact", "matcher": "auto", "source": "repo", "timeout": 60}
|
||||
{"script": "pre_compact.py", "event": "PreCompact", "matcher": "auto", "source": "repo", "timeout": 60},
|
||||
{"script": "pre_compact_rollover.py", "event": "PreCompact", "matcher": "manual", "source": "repo", "timeout": 120},
|
||||
{"script": "pre_compact_rollover.py", "event": "PreCompact", "matcher": "auto", "source": "repo", "timeout": 120}
|
||||
],
|
||||
"env": {
|
||||
"AIPASS_HOME": "{{REPO_ROOT}}",
|
||||
|
||||
@@ -22,6 +22,7 @@ STANDARD_FOOTER = """
|
||||
⚠️ TASK CHECKLIST (before marking complete):
|
||||
□ SEEDGO CHECK → drone @seedgo audit @branch (80%+)
|
||||
□ UPDATE MEMORIES → Your .trinity/local.json records this work
|
||||
□ UPDATE STATUS → Your STATUS.local.md reflects current state
|
||||
□ CLOSE FPLAN → drone @flow close <plan_id>
|
||||
□ EMAIL SENDER → drone @ai_mail email @<sender> "Subject" "Summary"
|
||||
|
||||
|
||||
@@ -2,12 +2,12 @@
|
||||
|
||||
## Identity
|
||||
|
||||
Memory is the central archive — vector search, rollover, and memory management for all AIPass branches. ChromaDB + sentence-transformers for semantic search. Rollover archives old `.trinity/` entries when files exceed 600 lines.
|
||||
Memory is the central archive — vector search, rollover, and memory management for all AIPass branches. ChromaDB + fastembed (ONNX) for semantic search. Rollover archives old `.trinity/` entries when files exceed 600 lines.
|
||||
|
||||
## Key Commands
|
||||
|
||||
```
|
||||
drone @memory search "query" # Semantic search (requires torch)
|
||||
drone @memory search "query" # Semantic search (requires fastembed)
|
||||
drone @memory search "q" --branch X # Filter by branch
|
||||
drone @memory rollover # Execute rollover for triggered files
|
||||
drone @memory status # Show rollover stats per branch
|
||||
@@ -26,7 +26,7 @@ Handlers implement domain logic under `apps/handlers/` (archive, json, learnings
|
||||
|
||||
## Known Issues
|
||||
|
||||
- `search` fails without `torch`/`sentence-transformers` installed
|
||||
- `search` fails without `fastembed` installed
|
||||
- 5 commands in `--help` have no backing module: push-templates, diff-templates, template-status, symbolic demo, symbolic fragments
|
||||
- `status` shows 0 branches — may need registry path investigation
|
||||
|
||||
|
||||
@@ -5,7 +5,7 @@
|
||||
**Purpose:** Central memory archive — vector search, rollover, and memory management for all AIPass branches.
|
||||
**Module:** `aipass.memory`
|
||||
**Created:** 2026-03-07
|
||||
**Last Updated:** 2026-05-02
|
||||
**Last Updated:** 2026-05-10
|
||||
**Citizen Class:** builder
|
||||
|
||||
---
|
||||
@@ -15,7 +15,7 @@
|
||||
Memory is the archival backbone of AIPass. Every branch accumulates session history and learnings in `.trinity/` files. When those files reach capacity, Memory archives the oldest entries into ChromaDB vectors — searchable, permanent, never lost.
|
||||
|
||||
What Memory does:
|
||||
- **Rollover** — detects when `.trinity/local.json` or `observations.json` exceed limits, extracts oldest entries, embeds them via sentence-transformers, stores in ChromaDB, trims the source file
|
||||
- **Rollover** — detects when `.trinity/local.json` or `observations.json` exceed limits, extracts oldest entries, embeds them via fastembed, stores in ChromaDB, trims the source file
|
||||
- **Search** — semantic search across all archived branch memories (4+ collections, 2200+ vectors)
|
||||
- **Templates** — distributes `.trinity/` schema updates across all branches (push, diff, status)
|
||||
- **Symbolic** — fragmented memory extraction from conversations (demo, analyze, extract, fragments, bootstrap, hook-test)
|
||||
@@ -99,9 +99,9 @@ memory/
|
||||
│ ├── symbolic/ # 6 handlers — chroma_client, deduplicator, extractor, hook, retriever, storage
|
||||
│ ├── templates/ # pusher.py, differ.py, spawn_pusher.py — template distribution
|
||||
│ ├── tracking/ # line_counter.py — metadata line count tracking
|
||||
│ ├── vector/ # embedder.py, embed_subprocess.py — sentence-transformer embeddings
|
||||
│ ├── vector/ # embedder.py, embed_subprocess.py — fastembed embeddings
|
||||
│ └── central_writer.py # Central memory write operations
|
||||
├── config/ # memory_bank.config.json — per-branch rollover limits
|
||||
├── config/ # memory.config.json — per-branch rollover limits
|
||||
├── templates/ # LOCAL.template.json, OBS.template — schema templates
|
||||
├── tests/ # 450 tests (16/16 module coverage)
|
||||
├── .chroma/ # Global ChromaDB vector store
|
||||
@@ -118,7 +118,7 @@ startup trigger → check_and_rollover()
|
||||
→ orchestrator.execute_rollover()
|
||||
→ create_rollover_backup() # safety copy to branch/.backup/
|
||||
→ extract_items() # v2: max(excess, 1) oldest entries
|
||||
→ embed via subprocess # sentence-transformers in memory .venv
|
||||
→ embed via subprocess # fastembed (ONNX) in memory .venv
|
||||
→ store in ChromaDB # global + local collections
|
||||
→ trim source file # write back with oldest removed
|
||||
```
|
||||
@@ -130,7 +130,7 @@ startup trigger → check_and_rollover()
|
||||
|
||||
### Subprocess Isolation
|
||||
|
||||
All ML operations (torch, sentence-transformers, chromadb) run via subprocess. The main process never imports these heavy libraries. Each embedding call resolves a Python interpreter via `_get_memory_python()` (env var `AIPASS_MEMORY_PYTHON` → `memory/.venv/bin/python` → `sys.executable`) and runs a self-contained script that reads stdin JSON and writes stdout JSON.
|
||||
All ML operations (fastembed, chromadb) run via subprocess. The main process never imports these libraries. Each embedding call resolves a Python interpreter via `_get_memory_python()` (env var `AIPASS_MEMORY_PYTHON` → `memory/.venv/bin/python` → `sys.executable`) and runs a self-contained script that reads stdin JSON and writes stdout JSON.
|
||||
|
||||
---
|
||||
|
||||
@@ -142,7 +142,7 @@ All ML operations (torch, sentence-transformers, chromadb) run via subprocess. T
|
||||
- `prax` (internal) — logging via `get_system_logger()`
|
||||
|
||||
### ML (in memory `.venv/` only)
|
||||
- `torch` + `sentence-transformers` — embedding generation
|
||||
- `fastembed` — embedding generation (ONNX, no torch required)
|
||||
- `chromadb` — vector storage and semantic search
|
||||
- `numpy` — numerical operations
|
||||
|
||||
@@ -165,7 +165,7 @@ All ML operations (torch, sentence-transformers, chromadb) run via subprocess. T
|
||||
|
||||
## Known Issues
|
||||
|
||||
- `search` requires torch/sentence-transformers in memory `.venv/` — fails without them
|
||||
- `search` requires fastembed in memory `.venv/` — fails without it
|
||||
- memory_watcher.py at 704 lines (near 700 threshold, bypassed in seedgo)
|
||||
- symbolic.py at 1604 lines (legacy port, bypassed)
|
||||
- manager.py at 1076 lines (complex learning extraction, bypassed)
|
||||
@@ -183,7 +183,7 @@ All ML operations (torch, sentence-transformers, chromadb) run via subprocess. T
|
||||
|
||||
---
|
||||
|
||||
*Last Updated: 2026-04-22*
|
||||
*Last Updated: 2026-05-10*
|
||||
|
||||
---
|
||||
[← Back to AIPass](../../../README.md)
|
||||
|
||||
@@ -230,7 +230,7 @@ def process_plans() -> Dict[str, Any]:
|
||||
Dict with success, files_processed, total_chunks
|
||||
"""
|
||||
# Load config
|
||||
config_path = _MEMORY_ROOT / "config" / "memory_bank.config.json"
|
||||
config_path = _MEMORY_ROOT / "config" / "memory.config.json"
|
||||
try:
|
||||
config = json.loads(config_path.read_text(encoding="utf-8"))
|
||||
plans_config = config.get("plans", {})
|
||||
@@ -242,7 +242,7 @@ def process_plans() -> Dict[str, Any]:
|
||||
return {"success": True, "skipped": True, "reason": "plans disabled"}
|
||||
|
||||
# Resolve plans directory (relative to repo root)
|
||||
plans_dir = plans_config.get("path", "src/aipass/flow/processed_plans")
|
||||
plans_dir = plans_config.get("path", ".backup/processed_plans")
|
||||
repo_root = _find_repo_root()
|
||||
plans_path = Path(plans_dir) if Path(plans_dir).is_absolute() else repo_root / plans_dir
|
||||
extensions = plans_config.get("supported_extensions", [".md"])
|
||||
|
||||
@@ -15,7 +15,7 @@ Processes markdown/text files from memory_pool directory:
|
||||
3. Vectorizes via ChromaDB
|
||||
4. Archives old files based on retention config
|
||||
|
||||
All settings read from memory_bank.config.json
|
||||
All settings read from memory.config.json
|
||||
"""
|
||||
|
||||
import json
|
||||
@@ -30,7 +30,7 @@ from aipass.memory.apps.handlers.json import json_handler
|
||||
|
||||
# Paths
|
||||
_MEMORY_ROOT = Path(__file__).resolve().parent.parent.parent.parent # handlers/intake/ → handlers/ → apps/ → memory/
|
||||
CONFIG_PATH = _MEMORY_ROOT / "config" / "memory_bank.config.json"
|
||||
CONFIG_PATH = _MEMORY_ROOT / "config" / "memory.config.json"
|
||||
MEMORY_POOL_PATH = _MEMORY_ROOT / "memory_pool"
|
||||
CHROMA_PATH = _MEMORY_ROOT / ".chroma"
|
||||
|
||||
@@ -103,7 +103,7 @@ def find_source_file(filename: str) -> Path | None:
|
||||
|
||||
|
||||
def load_config() -> dict:
|
||||
"""Load memory_pool config from memory_bank.config.json"""
|
||||
"""Load memory_pool config from memory.config.json"""
|
||||
try:
|
||||
with open(CONFIG_PATH) as f:
|
||||
config = json.load(f)
|
||||
|
||||
@@ -171,13 +171,13 @@ def _get_memory_file_path(branch: Dict, memory_type: str) -> Path | None:
|
||||
|
||||
def _load_config() -> Dict[str, Any]:
|
||||
"""
|
||||
Load memory_bank.config.json
|
||||
Load memory.config.json
|
||||
|
||||
Returns:
|
||||
Config dict, or empty dict on error
|
||||
"""
|
||||
# Look for config relative to this handler's location
|
||||
config_path = Path(__file__).resolve().parents[3] / "config" / "memory_bank.config.json"
|
||||
config_path = Path(__file__).resolve().parents[3] / "config" / "memory.config.json"
|
||||
|
||||
if not config_path.exists():
|
||||
return {}
|
||||
|
||||
@@ -103,7 +103,7 @@ def _get_rollover_threshold(branch_name: str, file_path: Path | None = None) ->
|
||||
logger.warning(f"[memory_watcher] Failed to read file-level threshold from {file_path}: {e}")
|
||||
|
||||
# 2. Check per-branch config override
|
||||
config_path = _MEMORY_ROOT / "config" / "memory_bank.config.json"
|
||||
config_path = _MEMORY_ROOT / "config" / "memory.config.json"
|
||||
|
||||
try:
|
||||
with open(config_path) as f:
|
||||
@@ -288,7 +288,7 @@ def _check_memory_pool() -> Dict[str, Any]:
|
||||
"""
|
||||
import json
|
||||
|
||||
config_path = _MEMORY_ROOT / "config" / "memory_bank.config.json"
|
||||
config_path = _MEMORY_ROOT / "config" / "memory.config.json"
|
||||
pool_path = _MEMORY_ROOT / "memory_pool"
|
||||
|
||||
# Load config
|
||||
@@ -348,7 +348,7 @@ def _check_plans() -> Dict[str, Any]:
|
||||
"""
|
||||
import json
|
||||
|
||||
config_path = _MEMORY_ROOT / "config" / "memory_bank.config.json"
|
||||
config_path = _MEMORY_ROOT / "config" / "memory.config.json"
|
||||
|
||||
# Load config
|
||||
try:
|
||||
|
||||
@@ -143,7 +143,7 @@ def store_vectors_subprocess(
|
||||
|
||||
def encode_batch_subprocess(texts: list) -> dict:
|
||||
"""
|
||||
Encode texts via subprocess using memory venv's sentence-transformers.
|
||||
Encode texts via subprocess using memory venv's fastembed.
|
||||
|
||||
Args:
|
||||
texts: List of text strings to encode
|
||||
|
||||
@@ -61,7 +61,7 @@ MIN_SIMILARITY_THRESHOLD = 0.40 # 40% minimum relevance
|
||||
|
||||
def encode_query_subprocess(query: str) -> dict:
|
||||
"""
|
||||
Encode query text via subprocess using memory venv's sentence-transformers.
|
||||
Encode query text via subprocess using memory venv's fastembed.
|
||||
|
||||
Args:
|
||||
query: Search query text
|
||||
|
||||
@@ -42,70 +42,6 @@ _MEMORY_ROOT = Path(__file__).resolve().parents[3]
|
||||
from aipass.memory.apps.handlers.storage.chroma import get_client
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# QUERY ENCODING SERVICE (Singleton)
|
||||
# =============================================================================
|
||||
|
||||
|
||||
class QueryEncoder:
|
||||
"""
|
||||
Query encoding service using same model as embedder.py
|
||||
|
||||
Ensures query embeddings are compatible with stored embeddings.
|
||||
Uses all-MiniLM-L6-v2 model with same settings.
|
||||
"""
|
||||
|
||||
def __init__(self, model_name: str = "sentence-transformers/all-MiniLM-L6-v2"):
|
||||
"""
|
||||
Initialize query encoder
|
||||
|
||||
Args:
|
||||
model_name: HuggingFace model identifier (must match embedder.py)
|
||||
|
||||
Raises:
|
||||
ImportError: If fastembed is not installed
|
||||
"""
|
||||
# Late imports (heavy optional dependencies)
|
||||
try:
|
||||
from fastembed import TextEmbedding
|
||||
except ImportError as e:
|
||||
logger.info(f"[vector_search] Optional ML dependencies not available: {e}")
|
||||
raise ImportError(f"Search requires fastembed. Install with: pip install fastembed. Original error: {e}")
|
||||
|
||||
self.model_name = model_name
|
||||
self.model = TextEmbedding(model_name)
|
||||
self.dimension = 384 # all-MiniLM-L6-v2 output dimension
|
||||
|
||||
def encode(self, query: str) -> List[float]:
|
||||
"""
|
||||
Encode query text to embedding
|
||||
|
||||
Args:
|
||||
query: Query text string
|
||||
|
||||
Returns:
|
||||
384-dimensional embedding as list
|
||||
"""
|
||||
embeddings = list(self.model.embed([query]))
|
||||
return embeddings[0].tolist()
|
||||
|
||||
|
||||
# Global encoder instance (singleton pattern)
|
||||
_query_encoder = None
|
||||
|
||||
|
||||
def _get_encoder() -> QueryEncoder:
|
||||
"""
|
||||
Get or create query encoder singleton
|
||||
|
||||
Lazy initialization - model loaded on first use
|
||||
"""
|
||||
global _query_encoder
|
||||
if _query_encoder is None:
|
||||
_query_encoder = QueryEncoder()
|
||||
return _query_encoder
|
||||
|
||||
|
||||
# =============================================================================
|
||||
# CHROMA SEARCH SERVICE
|
||||
# =============================================================================
|
||||
@@ -133,44 +69,6 @@ class SearchService:
|
||||
self.client = get_client(db_path)
|
||||
self.db_path = db_path
|
||||
|
||||
def query_collection(
|
||||
self,
|
||||
collection_name: str,
|
||||
query_embedding: List[float],
|
||||
n_results: int = 5,
|
||||
where: Dict[str, Any] | None = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Query a specific collection
|
||||
|
||||
Args:
|
||||
collection_name: Name of collection to query
|
||||
query_embedding: Query embedding vector
|
||||
n_results: Number of results to return
|
||||
where: Metadata filter (e.g., {"branch": "SEEDGO"})
|
||||
|
||||
Returns:
|
||||
Dict with query results
|
||||
"""
|
||||
try:
|
||||
collection = self.client.get_collection(collection_name, embedding_function=None)
|
||||
except Exception as e:
|
||||
logger.warning(f"[vector_search] Collection lookup failed for '{collection_name}': {e}")
|
||||
return {"collection": collection_name, "exists": False, "error": f"Collection not found: {e}"}
|
||||
|
||||
# Query collection
|
||||
results = collection.query(query_embeddings=[query_embedding], n_results=n_results, where=where)
|
||||
|
||||
return {
|
||||
"collection": collection_name,
|
||||
"exists": True,
|
||||
"ids": results["ids"][0] if results["ids"] else [],
|
||||
"documents": results["documents"][0] if results["documents"] else [],
|
||||
"metadatas": results["metadatas"][0] if results["metadatas"] else [],
|
||||
"distances": results["distances"][0] if results["distances"] else [],
|
||||
"count": len(results["ids"][0]) if results["ids"] else 0,
|
||||
}
|
||||
|
||||
def list_collections(self) -> List[str]:
|
||||
"""
|
||||
List all collections in database
|
||||
@@ -219,108 +117,6 @@ def _get_service(db_path: Path | None = None) -> SearchService:
|
||||
# =============================================================================
|
||||
|
||||
|
||||
def search_collection(
|
||||
query_embedding: List[float],
|
||||
collection_name: str,
|
||||
n_results: int = 5,
|
||||
where: Dict[str, Any] | None = None,
|
||||
db_path: Path | None = None,
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Query a ChromaDB collection with embedding
|
||||
|
||||
Search a specific collection for semantically similar memories.
|
||||
|
||||
Args:
|
||||
query_embedding: Pre-encoded query embedding (384-dim list)
|
||||
collection_name: Name of collection to search
|
||||
n_results: Number of results to return (default: 5)
|
||||
where: Optional metadata filter (e.g., {"branch": "SEEDGO"})
|
||||
db_path: Path to ChromaDB database (None = global memory/.chroma)
|
||||
|
||||
Returns:
|
||||
Dict with success status and search results
|
||||
|
||||
Example:
|
||||
# Encode query first
|
||||
query_result = encode_query("how does rollover work?")
|
||||
|
||||
# Search collection
|
||||
result = search_collection(
|
||||
query_embedding=query_result['embedding'],
|
||||
collection_name="seed_observations",
|
||||
n_results=10
|
||||
)
|
||||
|
||||
if result['success']:
|
||||
for i, doc in enumerate(result['documents']):
|
||||
print(f"{i+1}. {doc[:100]}...")
|
||||
"""
|
||||
# Convert string db_path to Path
|
||||
if db_path is not None and isinstance(db_path, str):
|
||||
db_path = Path(db_path)
|
||||
|
||||
if not query_embedding:
|
||||
return {"success": False, "error": "No query embedding provided"}
|
||||
|
||||
try:
|
||||
service = _get_service(db_path)
|
||||
result = service.query_collection(
|
||||
collection_name=collection_name, query_embedding=query_embedding, n_results=n_results, where=where
|
||||
)
|
||||
|
||||
# Check if collection exists
|
||||
if not result.get("exists", False):
|
||||
return {
|
||||
"success": False,
|
||||
"error": result.get("error", "Collection not found"),
|
||||
"collection": collection_name,
|
||||
}
|
||||
|
||||
json_handler.log_operation(
|
||||
"vector_search_collection",
|
||||
{"collection": collection_name, "count": result.get("count", 0), "success": True},
|
||||
)
|
||||
return {"success": True, **result}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"[vector_search] Collection search failed for '{collection_name}': {e}")
|
||||
return {"success": False, "error": f"Search failed: {e}"}
|
||||
|
||||
|
||||
def encode_query(query: str) -> Dict[str, Any]:
|
||||
"""
|
||||
Encode query text to embedding using same model as storage
|
||||
|
||||
Uses all-MiniLM-L6-v2 model (same as embedder.py) to ensure
|
||||
query embeddings are compatible with stored embeddings.
|
||||
|
||||
Args:
|
||||
query: Query text string
|
||||
|
||||
Returns:
|
||||
Dict with success status and embedding
|
||||
|
||||
Example:
|
||||
result = encode_query("how does memory compression work?")
|
||||
if result['success']:
|
||||
embedding = result['embedding'] # 384-dim list
|
||||
dimension = result['dimension'] # 384
|
||||
"""
|
||||
if not query or not query.strip():
|
||||
return {"success": False, "error": "Empty query string"}
|
||||
|
||||
try:
|
||||
encoder = _get_encoder()
|
||||
embedding = encoder.encode(query)
|
||||
|
||||
return {"success": True, "embedding": embedding, "dimension": len(embedding), "model": encoder.model_name}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"[vector_search] Query encoding failed: {e}")
|
||||
return {"success": False, "error": f"Encoding failed: {e}"}
|
||||
|
||||
|
||||
def list_collections(db_path: Path | None = None) -> Dict[str, Any]:
|
||||
"""
|
||||
List available collections in database
|
||||
@@ -345,81 +141,12 @@ def list_collections(db_path: Path | None = None) -> Dict[str, Any]:
|
||||
service = _get_service(db_path)
|
||||
collections = service.list_collections()
|
||||
|
||||
json_handler.log_operation(
|
||||
"vector_list_collections",
|
||||
{"count": len(collections), "success": True},
|
||||
)
|
||||
return {"success": True, "collections": collections, "count": len(collections), "db_path": str(service.db_path)}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"[vector_search] Failed to list collections: {e}")
|
||||
return {"success": False, "error": f"Failed to list collections: {e}"}
|
||||
|
||||
|
||||
def search_all_collections(
|
||||
query_embedding: List[float], n_results: int = 5, where: Dict[str, Any] | None = None, db_path: Path | None = None
|
||||
) -> Dict[str, Any]:
|
||||
"""
|
||||
Search across all collections in database
|
||||
|
||||
Queries every collection and aggregates results.
|
||||
Useful for cross-branch semantic search.
|
||||
|
||||
Args:
|
||||
query_embedding: Pre-encoded query embedding
|
||||
n_results: Number of results per collection
|
||||
where: Optional metadata filter
|
||||
db_path: Path to ChromaDB database (None = global default)
|
||||
|
||||
Returns:
|
||||
Dict with success status and aggregated results
|
||||
|
||||
Example:
|
||||
# Find similar memories across all branches
|
||||
query_result = encode_query("deployment process")
|
||||
result = search_all_collections(
|
||||
query_embedding=query_result['embedding'],
|
||||
n_results=3
|
||||
)
|
||||
|
||||
if result['success']:
|
||||
for coll_name, coll_results in result['results'].items():
|
||||
print(f"\n{coll_name}:")
|
||||
for doc in coll_results['documents']:
|
||||
print(f" - {doc[:80]}...")
|
||||
"""
|
||||
# Convert string db_path to Path
|
||||
if db_path is not None and isinstance(db_path, str):
|
||||
db_path = Path(db_path)
|
||||
|
||||
if not query_embedding:
|
||||
return {"success": False, "error": "No query embedding provided"}
|
||||
|
||||
try:
|
||||
service = _get_service(db_path)
|
||||
collections = service.list_collections()
|
||||
|
||||
if not collections:
|
||||
return {"success": True, "results": {}, "message": "No collections found"}
|
||||
|
||||
# Query each collection
|
||||
results = {}
|
||||
for collection_name in collections:
|
||||
coll_result = service.query_collection(
|
||||
collection_name=collection_name, query_embedding=query_embedding, n_results=n_results, where=where
|
||||
)
|
||||
|
||||
if coll_result.get("exists", False):
|
||||
results[collection_name] = {
|
||||
"documents": coll_result["documents"],
|
||||
"metadatas": coll_result["metadatas"],
|
||||
"distances": coll_result["distances"],
|
||||
"ids": coll_result["ids"],
|
||||
"count": coll_result["count"],
|
||||
}
|
||||
|
||||
total = sum(r["count"] for r in results.values())
|
||||
json_handler.log_operation(
|
||||
"vector_search_all", {"collections": len(results), "total_results": total, "success": True}
|
||||
)
|
||||
return {"success": True, "results": results, "collections_searched": len(results), "total_results": total}
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"[vector_search] Multi-collection search failed: {e}")
|
||||
return {"success": False, "error": f"Multi-collection search failed: {e}"}
|
||||
|
||||
@@ -29,7 +29,7 @@ Dependencies (optional):
|
||||
|
||||
from typing import List, Dict, Any
|
||||
from pathlib import Path
|
||||
from datetime import datetime
|
||||
import hashlib
|
||||
|
||||
from aipass.prax.apps.modules.logger import get_system_logger
|
||||
from aipass.memory.apps.handlers.json import json_handler
|
||||
@@ -138,18 +138,14 @@ class ChromaService:
|
||||
embedding_function=None,
|
||||
)
|
||||
|
||||
# Get existing count for ID generation
|
||||
existing_count = collection.count()
|
||||
|
||||
# Generate unique IDs
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
ids = [f"{branch}_{memory_type}_{existing_count + i}_{timestamp}" for i in range(len(embeddings))]
|
||||
# Content-hash IDs prevent duplicates across rollover runs
|
||||
ids = [f"{branch}_{memory_type}_{hashlib.sha256(doc.encode()).hexdigest()[:16]}" for doc in documents]
|
||||
|
||||
# Convert embeddings to list format (Chroma requirement)
|
||||
embeddings_list = [emb.tolist() if hasattr(emb, "tolist") else emb for emb in embeddings]
|
||||
|
||||
# Batch insert (optimal size: 100-150)
|
||||
collection.add(embeddings=embeddings_list, documents=documents, metadatas=metadatas, ids=ids)
|
||||
# Upsert: idempotent — same content gets same ID, no duplicates
|
||||
collection.upsert(embeddings=embeddings_list, documents=documents, metadatas=metadatas, ids=ids)
|
||||
|
||||
new_count = collection.count()
|
||||
|
||||
|
||||
@@ -21,8 +21,8 @@ Output: JSON on stdout with result
|
||||
import sys
|
||||
import json
|
||||
import logging
|
||||
import hashlib
|
||||
from pathlib import Path
|
||||
from datetime import datetime
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -67,14 +67,14 @@ def _store_vectors(branch, memory_type, embeddings, documents, metadatas, db_pat
|
||||
embedding_function=None,
|
||||
)
|
||||
|
||||
existing_count = collection.count()
|
||||
timestamp = datetime.now().strftime("%Y%m%d_%H%M%S")
|
||||
ids = [f"{branch}_{memory_type}_{existing_count + i}_{timestamp}" for i in range(len(embeddings))]
|
||||
# Content-hash IDs prevent duplicates across rollover runs
|
||||
ids = [f"{branch}_{memory_type}_{hashlib.sha256(doc.encode()).hexdigest()[:16]}" for doc in documents]
|
||||
|
||||
# Chroma expects lists, not numpy arrays
|
||||
embeddings_list = [emb.tolist() if hasattr(emb, "tolist") else emb for emb in embeddings]
|
||||
|
||||
collection.add(embeddings=embeddings_list, documents=documents, metadatas=metadatas, ids=ids)
|
||||
# Upsert: idempotent — same content gets same ID, no duplicates
|
||||
collection.upsert(embeddings=embeddings_list, documents=documents, metadatas=metadatas, ids=ids)
|
||||
|
||||
new_count = collection.count()
|
||||
|
||||
@@ -107,7 +107,7 @@ def _check_plan(plan_label, db_path=None):
|
||||
"""
|
||||
client = _get_client(db_path)
|
||||
|
||||
collection_name = "flow_flow_plans"
|
||||
collection_name = "flow_plans"
|
||||
try:
|
||||
collection = client.get_collection(collection_name, embedding_function=None)
|
||||
except Exception as e:
|
||||
|
||||
@@ -167,7 +167,7 @@ def print_help() -> None:
|
||||
console.print("[bold]WORKFLOW:[/bold]")
|
||||
console.print(" 1. Detect files exceeding limits (line count or entry count)")
|
||||
console.print(" 2. Extract oldest entries")
|
||||
console.print(" 3. Generate embeddings via sentence-transformers")
|
||||
console.print(" 3. Generate embeddings via fastembed")
|
||||
console.print(" 4. Store vectors in local + global ChromaDB")
|
||||
console.print()
|
||||
|
||||
|
||||
@@ -1501,13 +1501,8 @@ def bootstrap_from_jsonl(max_sessions: int = 8) -> None:
|
||||
for i, jsonl_path in enumerate(sessions, 1):
|
||||
# Derive branch name from parent directory
|
||||
branch_dir = jsonl_path.parent.name
|
||||
# New layout: -home-patrick-Projects-AIPass-src-aipass-<branch>
|
||||
# Layout: -home-patrick-Projects-AIPass-src-aipass-<branch>
|
||||
branch_name = branch_dir.rsplit("-aipass-", 1)[-1].replace("-", "_").upper()
|
||||
# Legacy layout fallback
|
||||
if branch_name.startswith("AIPASS_CORE_"):
|
||||
branch_name = branch_name.replace("AIPASS_CORE_", "")
|
||||
if branch_name.startswith("AIPASS_OS_"):
|
||||
branch_name = branch_name.replace("AIPASS_OS_", "")
|
||||
|
||||
file_size_kb = jsonl_path.stat().st_size / 1024
|
||||
console.print(
|
||||
|
||||
@@ -13,7 +13,7 @@ Checks whether a plan has been vectorized in ChromaDB.
|
||||
|
||||
Purpose:
|
||||
Thin orchestration layer - calls chroma_subprocess via subprocess
|
||||
to query the flow_flow_plans collection for a given plan label.
|
||||
to query the flow_plans collection for a given plan label.
|
||||
"""
|
||||
|
||||
import subprocess
|
||||
@@ -261,7 +261,7 @@ def print_help() -> None:
|
||||
console.print(" [dim]drone @memory verify HPLAN-0001[/dim]")
|
||||
console.print()
|
||||
console.print("[bold]HOW IT WORKS:[/bold]")
|
||||
console.print(" 1. Query the flow_flow_plans ChromaDB collection")
|
||||
console.print(" 1. Query the flow_plans ChromaDB collection")
|
||||
console.print(" 2. Filter entries by source_file metadata matching the plan label")
|
||||
console.print(" 3. Report vectorization status and chunk count")
|
||||
console.print()
|
||||
|
||||
@@ -0,0 +1,24 @@
|
||||
{
|
||||
"memory_pool": {
|
||||
"enabled": false,
|
||||
"process_on_startup": false,
|
||||
"extensions": [".md", ".txt"]
|
||||
},
|
||||
"rollover": {
|
||||
"defaults": {
|
||||
"max_lines": 500,
|
||||
"archive_oldest": 100
|
||||
},
|
||||
"per_branch": {}
|
||||
},
|
||||
"plans": {
|
||||
"enabled": true,
|
||||
"path": ".backup/processed_plans",
|
||||
"collection_name": "plans",
|
||||
"supported_extensions": [".md"]
|
||||
},
|
||||
"intake": {
|
||||
"enabled": false,
|
||||
"pool_dir": "memory_pool"
|
||||
}
|
||||
}
|
||||
@@ -126,7 +126,7 @@ class TestLoadConfig:
|
||||
def test_returns_config_dict_when_file_exists(self, tmp_path: Path, monkeypatch):
|
||||
config_dir = tmp_path / "config"
|
||||
config_dir.mkdir()
|
||||
config_file = config_dir / "memory_bank.config.json"
|
||||
config_file = config_dir / "memory.config.json"
|
||||
config_data = {
|
||||
"rollover": {
|
||||
"defaults": {"max_lines": 500},
|
||||
@@ -166,7 +166,7 @@ class TestLoadConfig:
|
||||
def test_returns_empty_dict_on_invalid_json(self, tmp_path: Path, monkeypatch):
|
||||
config_dir = tmp_path / "config"
|
||||
config_dir.mkdir()
|
||||
config_file = config_dir / "memory_bank.config.json"
|
||||
config_file = config_dir / "memory.config.json"
|
||||
config_file.write_text("NOT VALID JSON {{", encoding="utf-8")
|
||||
|
||||
from aipass.memory.apps.handlers.monitor import detector
|
||||
|
||||
@@ -113,7 +113,7 @@ class TestLoadConfig:
|
||||
|
||||
def test_loads_valid_config(self, monkeypatch, tmp_path):
|
||||
mod = _import_pool_processor(monkeypatch)
|
||||
config_file = tmp_path / "memory_bank.config.json"
|
||||
config_file = tmp_path / "memory.config.json"
|
||||
config_file.write_text(
|
||||
json.dumps({"memory_pool": {"enabled": True, "keep_recent": 5, "collection_name": "test_pool"}}),
|
||||
encoding="utf-8",
|
||||
@@ -137,7 +137,7 @@ class TestLoadConfig:
|
||||
|
||||
def test_returns_empty_when_no_memory_pool_key(self, monkeypatch, tmp_path):
|
||||
mod = _import_pool_processor(monkeypatch)
|
||||
config_file = tmp_path / "memory_bank.config.json"
|
||||
config_file = tmp_path / "memory.config.json"
|
||||
config_file.write_text(json.dumps({"rollover": {}}), encoding="utf-8")
|
||||
monkeypatch.setattr(mod, "CONFIG_PATH", config_file)
|
||||
|
||||
|
||||
@@ -394,10 +394,10 @@ class TestProcessPlans:
|
||||
"""Test process_plans main entry point."""
|
||||
|
||||
def _setup_config(self, tmp_path, config_data):
|
||||
"""Write a memory_bank.config.json and return its path."""
|
||||
"""Write a memory.config.json and return its path."""
|
||||
config_dir = tmp_path / "config"
|
||||
config_dir.mkdir(parents=True, exist_ok=True)
|
||||
config_path = config_dir / "memory_bank.config.json"
|
||||
config_path = config_dir / "memory.config.json"
|
||||
config_path.write_text(json.dumps(config_data), encoding="utf-8")
|
||||
return config_path
|
||||
|
||||
|
||||
@@ -11,12 +11,8 @@
|
||||
Covers:
|
||||
from aipass.memory.apps.handlers.search.query_executor import encode_query_subprocess
|
||||
from aipass.memory.apps.handlers.search.query_executor import search_vectors_subprocess
|
||||
from aipass.memory.apps.handlers.search.vector_search import search_collection
|
||||
from aipass.memory.apps.handlers.search.vector_search import encode_query
|
||||
from aipass.memory.apps.handlers.search.vector_search import search_all_collections
|
||||
|
||||
Tests subprocess-based encoding/search, ChromaDB collection queries,
|
||||
query encoding via the singleton QueryEncoder, and multi-collection search.
|
||||
Tests subprocess-based encoding/search.
|
||||
All tests use mocks -- no live subprocess, ML model, or ChromaDB access.
|
||||
"""
|
||||
|
||||
@@ -67,70 +63,6 @@ def _import_query_executor(monkeypatch):
|
||||
return query_executor, mocks
|
||||
|
||||
|
||||
def _prepare_vector_search_mocks(monkeypatch):
|
||||
"""Insert mocks for vector_search module-level imports.
|
||||
|
||||
Returns a dict of key mock objects for assertions.
|
||||
"""
|
||||
# Mock chromadb client via the chroma module
|
||||
mock_client = MagicMock()
|
||||
mock_chroma = MagicMock()
|
||||
mock_chroma.get_client = MagicMock(return_value=mock_client)
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"aipass.memory.apps.handlers.storage.chroma",
|
||||
mock_chroma,
|
||||
)
|
||||
monkeypatch.setitem(
|
||||
sys.modules,
|
||||
"aipass.memory.apps.handlers.storage",
|
||||
MagicMock(),
|
||||
)
|
||||
|
||||
# Mock fastembed for QueryEncoder
|
||||
mock_model = MagicMock()
|
||||
mock_model.embed.side_effect = lambda texts: iter(
|
||||
[MagicMock(tolist=MagicMock(return_value=[0.1] * 384)) for _ in texts]
|
||||
)
|
||||
|
||||
mock_te_cls = MagicMock(return_value=mock_model)
|
||||
|
||||
mock_fastembed = MagicMock()
|
||||
mock_fastembed.TextEmbedding = mock_te_cls
|
||||
monkeypatch.setitem(sys.modules, "fastembed", mock_fastembed)
|
||||
|
||||
return {
|
||||
"client": mock_client,
|
||||
"model": mock_model,
|
||||
"st_cls": mock_te_cls,
|
||||
}
|
||||
|
||||
|
||||
def _import_vector_search(monkeypatch):
|
||||
"""Prepare mocks and import (or reimport) vector_search.
|
||||
|
||||
Returns (module, mocks_dict).
|
||||
"""
|
||||
mocks = _prepare_vector_search_mocks(monkeypatch)
|
||||
|
||||
# Remove cached modules so they get re-imported with our mocks
|
||||
sys.modules.pop("aipass.memory.apps.handlers.search.vector_search", None)
|
||||
|
||||
# Clear parent package attribute so Python re-executes module code
|
||||
parent = sys.modules.get("aipass.memory.apps.handlers.search")
|
||||
if parent is not None and hasattr(parent, "vector_search"):
|
||||
delattr(parent, "vector_search")
|
||||
|
||||
from aipass.memory.apps.handlers.search import vector_search # noqa: E402
|
||||
|
||||
# Reset singletons for a clean slate
|
||||
setattr(vector_search, "_query_encoder", None)
|
||||
setattr(vector_search, "_search_service", None)
|
||||
setattr(vector_search, "_local_services", {})
|
||||
|
||||
return vector_search, mocks
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Tests: encode_query_subprocess
|
||||
# ===========================================================================
|
||||
@@ -316,277 +248,3 @@ class TestSearchVectorsSubprocess:
|
||||
|
||||
assert result["success"] is False
|
||||
assert "Cannot execute" in result["error"]
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Tests: search_collection (vector_search)
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
class TestSearchCollection:
|
||||
"""Verify search_collection wraps SearchService.query_collection."""
|
||||
|
||||
def test_search_returns_results_on_success(self, monkeypatch, tmp_path):
|
||||
"""Successful collection query returns success with results."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
mock_collection = MagicMock()
|
||||
mock_collection.query.return_value = {
|
||||
"ids": [["id1", "id2"]],
|
||||
"documents": [["doc1", "doc2"]],
|
||||
"metadatas": [[{"branch": "TEST"}, {"branch": "TEST"}]],
|
||||
"distances": [[0.1, 0.2]],
|
||||
}
|
||||
mocks["client"].get_collection.return_value = mock_collection
|
||||
|
||||
result = mod.search_collection(
|
||||
query_embedding=[0.1] * 384,
|
||||
collection_name="test_collection",
|
||||
n_results=5,
|
||||
db_path=tmp_path / ".chroma",
|
||||
)
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["count"] == 2
|
||||
assert "doc1" in result["documents"]
|
||||
|
||||
def test_search_returns_error_for_missing_collection(self, monkeypatch, tmp_path):
|
||||
"""Missing collection returns success=False with error message."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
mocks["client"].get_collection.side_effect = ValueError("Collection not found")
|
||||
|
||||
result = mod.search_collection(
|
||||
query_embedding=[0.1] * 384,
|
||||
collection_name="nonexistent",
|
||||
db_path=tmp_path / ".chroma",
|
||||
)
|
||||
|
||||
assert result["success"] is False
|
||||
assert "not found" in result["error"].lower()
|
||||
|
||||
def test_search_returns_error_for_empty_embedding(self, monkeypatch):
|
||||
"""Empty query embedding returns error."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
result = mod.search_collection(
|
||||
query_embedding=[],
|
||||
collection_name="test_collection",
|
||||
)
|
||||
|
||||
assert result["success"] is False
|
||||
assert "no query embedding" in result["error"].lower()
|
||||
|
||||
def test_search_accepts_string_db_path(self, monkeypatch, tmp_path):
|
||||
"""String db_path is converted to Path internally."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
mock_collection = MagicMock()
|
||||
mock_collection.query.return_value = {
|
||||
"ids": [["id1"]],
|
||||
"documents": [["doc1"]],
|
||||
"metadatas": [[{}]],
|
||||
"distances": [[0.05]],
|
||||
}
|
||||
mocks["client"].get_collection.return_value = mock_collection
|
||||
|
||||
result = mod.search_collection(
|
||||
query_embedding=[0.1] * 384,
|
||||
collection_name="test_collection",
|
||||
db_path=str(tmp_path / ".chroma"),
|
||||
)
|
||||
|
||||
assert result["success"] is True
|
||||
|
||||
def test_search_passes_where_filter(self, monkeypatch, tmp_path):
|
||||
"""Metadata where filter is forwarded to collection.query."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
mock_collection = MagicMock()
|
||||
mock_collection.query.return_value = {
|
||||
"ids": [[]],
|
||||
"documents": [[]],
|
||||
"metadatas": [[]],
|
||||
"distances": [[]],
|
||||
}
|
||||
mocks["client"].get_collection.return_value = mock_collection
|
||||
|
||||
mod.search_collection(
|
||||
query_embedding=[0.1] * 384,
|
||||
collection_name="test_collection",
|
||||
where={"branch": "SEEDGO"},
|
||||
db_path=tmp_path / ".chroma",
|
||||
)
|
||||
|
||||
call_kwargs = mock_collection.query.call_args.kwargs
|
||||
assert call_kwargs["where"] == {"branch": "SEEDGO"}
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Tests: encode_query (vector_search)
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
class TestEncodeQuery:
|
||||
"""Verify encode_query uses QueryEncoder singleton."""
|
||||
|
||||
def test_encode_returns_embedding(self, monkeypatch):
|
||||
"""Successful encoding returns embedding with dimension."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
result = mod.encode_query("test query")
|
||||
|
||||
assert result["success"] is True
|
||||
assert len(result["embedding"]) == 384
|
||||
assert result["dimension"] == 384
|
||||
assert result["model"] == "sentence-transformers/all-MiniLM-L6-v2"
|
||||
|
||||
def test_encode_rejects_empty_query(self, monkeypatch):
|
||||
"""Empty query string returns error."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
result = mod.encode_query("")
|
||||
|
||||
assert result["success"] is False
|
||||
assert "empty" in result["error"].lower()
|
||||
|
||||
def test_encode_rejects_whitespace_only_query(self, monkeypatch):
|
||||
"""Whitespace-only query string returns error."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
result = mod.encode_query(" ")
|
||||
|
||||
assert result["success"] is False
|
||||
assert "empty" in result["error"].lower()
|
||||
|
||||
def test_encode_handles_model_error(self, monkeypatch):
|
||||
"""If model.encode raises, error is caught and returned."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
mocks["model"].embed.side_effect = RuntimeError("CUDA out of memory")
|
||||
|
||||
result = mod.encode_query("test query")
|
||||
|
||||
assert result["success"] is False
|
||||
assert "CUDA" in result["error"]
|
||||
|
||||
def test_encode_uses_singleton(self, monkeypatch):
|
||||
"""Successive calls reuse the same QueryEncoder instance."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
mod.encode_query("first query")
|
||||
mod.encode_query("second query")
|
||||
|
||||
# TextEmbedding should only be constructed once
|
||||
mocks["st_cls"].assert_called_once()
|
||||
|
||||
|
||||
# ===========================================================================
|
||||
# Tests: search_all_collections (vector_search)
|
||||
# ===========================================================================
|
||||
|
||||
|
||||
class TestSearchAllCollections:
|
||||
"""Verify search_all_collections aggregates results across collections."""
|
||||
|
||||
def test_search_all_returns_aggregated_results(self, monkeypatch, tmp_path):
|
||||
"""Searching all collections returns results from each."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
mock_coll_a = MagicMock()
|
||||
mock_coll_a.name = "coll_a"
|
||||
mock_coll_b = MagicMock()
|
||||
mock_coll_b.name = "coll_b"
|
||||
mocks["client"].list_collections.return_value = [mock_coll_a, mock_coll_b]
|
||||
|
||||
mock_collection = MagicMock()
|
||||
mock_collection.query.return_value = {
|
||||
"ids": [["id1"]],
|
||||
"documents": [["doc1"]],
|
||||
"metadatas": [[{"branch": "TEST"}]],
|
||||
"distances": [[0.1]],
|
||||
}
|
||||
mocks["client"].get_collection.return_value = mock_collection
|
||||
|
||||
result = mod.search_all_collections(
|
||||
query_embedding=[0.1] * 384,
|
||||
n_results=3,
|
||||
db_path=tmp_path / ".chroma",
|
||||
)
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["collections_searched"] == 2
|
||||
assert result["total_results"] == 2
|
||||
assert "coll_a" in result["results"]
|
||||
assert "coll_b" in result["results"]
|
||||
|
||||
def test_search_all_returns_empty_when_no_collections(self, monkeypatch, tmp_path):
|
||||
"""No collections returns success with empty results."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
mocks["client"].list_collections.return_value = []
|
||||
|
||||
result = mod.search_all_collections(
|
||||
query_embedding=[0.1] * 384,
|
||||
db_path=tmp_path / ".chroma",
|
||||
)
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["results"] == {}
|
||||
assert "no collections" in result["message"].lower()
|
||||
|
||||
def test_search_all_returns_error_for_empty_embedding(self, monkeypatch):
|
||||
"""Empty embedding returns error."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
result = mod.search_all_collections(query_embedding=[])
|
||||
|
||||
assert result["success"] is False
|
||||
assert "no query embedding" in result["error"].lower()
|
||||
|
||||
def test_search_all_handles_exception(self, monkeypatch, tmp_path):
|
||||
"""Exception during search returns error dict."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
mocks["client"].list_collections.side_effect = RuntimeError("DB corrupted")
|
||||
|
||||
result = mod.search_all_collections(
|
||||
query_embedding=[0.1] * 384,
|
||||
db_path=tmp_path / ".chroma",
|
||||
)
|
||||
|
||||
assert result["success"] is False
|
||||
assert "DB corrupted" in result["error"]
|
||||
|
||||
def test_search_all_accepts_string_db_path(self, monkeypatch, tmp_path):
|
||||
"""String db_path is accepted and converted."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
mocks["client"].list_collections.return_value = []
|
||||
|
||||
result = mod.search_all_collections(
|
||||
query_embedding=[0.1] * 384,
|
||||
db_path=str(tmp_path / ".chroma"),
|
||||
)
|
||||
|
||||
assert result["success"] is True
|
||||
|
||||
def test_search_all_skips_nonexistent_collections(self, monkeypatch, tmp_path):
|
||||
"""Collections that fail get_collection are excluded from results."""
|
||||
mod, mocks = _import_vector_search(monkeypatch)
|
||||
|
||||
mock_coll_a = MagicMock()
|
||||
mock_coll_a.name = "coll_a"
|
||||
mocks["client"].list_collections.return_value = [mock_coll_a]
|
||||
|
||||
# get_collection raises => query_collection returns exists=False
|
||||
mocks["client"].get_collection.side_effect = ValueError("Collection not found")
|
||||
|
||||
result = mod.search_all_collections(
|
||||
query_embedding=[0.1] * 384,
|
||||
db_path=tmp_path / ".chroma",
|
||||
)
|
||||
|
||||
assert result["success"] is True
|
||||
assert result["collections_searched"] == 0
|
||||
assert result["total_results"] == 0
|
||||
|
||||
@@ -103,7 +103,7 @@ class TestChromaServiceStoreVectors:
|
||||
_reset_globals(chroma)
|
||||
|
||||
mock_collection = MagicMock()
|
||||
mock_collection.count.side_effect = [0, 3] # before and after add
|
||||
mock_collection.count.return_value = 3 # count after upsert
|
||||
|
||||
mock_client = MagicMock()
|
||||
mock_client.get_or_create_collection.return_value = mock_collection
|
||||
@@ -121,7 +121,7 @@ class TestChromaServiceStoreVectors:
|
||||
assert result["count"] == 3
|
||||
assert result["total_vectors"] == 3
|
||||
assert len(result["ids"]) == 3
|
||||
mock_collection.add.assert_called_once()
|
||||
mock_collection.upsert.assert_called_once()
|
||||
|
||||
|
||||
class TestChromaServiceGetCollectionStats:
|
||||
|
||||
Reference in New Issue
Block a user