48 lines
2.0 KiB
Python
48 lines
2.0 KiB
Python
from hashlib import sha256
|
|
from math import sqrt
|
|
from typing import Any
|
|
|
|
|
|
def embed_text(text: str, size: int = 32) -> list[float]:
|
|
"""Deterministic local embedding used for offline demo and tests."""
|
|
|
|
buckets = [0.0] * size
|
|
for word in text.lower().split():
|
|
digest = sha256(word.encode()).digest()
|
|
buckets[digest[0] % size] += 1.0
|
|
norm = sqrt(sum(v * v for v in buckets)) or 1.0
|
|
return [v / norm for v in buckets]
|
|
|
|
|
|
class MemoryManager:
|
|
"""Stores semantic memories behind a Qdrant-like interface with local fallback."""
|
|
|
|
def __init__(self) -> None:
|
|
self._items: list[dict[str, Any]] = []
|
|
|
|
def upsert(self, memory_type: str, text: str, metadata: dict[str, Any] | None = None) -> dict[str, Any]:
|
|
item = {
|
|
"id": sha256(f"{memory_type}:{text}".encode()).hexdigest()[:16],
|
|
"type": memory_type,
|
|
"text": text,
|
|
"metadata": metadata or {},
|
|
"embedding": embed_text(text),
|
|
}
|
|
self._items = [existing for existing in self._items if existing["id"] != item["id"]]
|
|
self._items.append(item)
|
|
return item
|
|
|
|
def retrieve(self, query: str, limit: int = 5) -> list[dict[str, Any]]:
|
|
query_vec = embed_text(query)
|
|
|
|
def score(item: dict[str, Any]) -> float:
|
|
return sum(a * b for a, b in zip(query_vec, item["embedding"], strict=True))
|
|
|
|
ranked = sorted(self._items, key=score, reverse=True)
|
|
return [{k: v for k, v in item.items() if k != "embedding"} | {"score": score(item)} for item in ranked[:limit]]
|
|
|
|
def seed_demo(self) -> None:
|
|
self.upsert("customer_preference", "Customer prefers white flowers and a classic minimal arrangement.", {"customer": "Demo"})
|
|
self.upsert("vendor", "Aegean Blooms is the preferred florist for white flower arrangements.", {"vendor": "Aegean Blooms"})
|
|
self.upsert("meeting", "Previous meeting agreed to compare venues before reserving July dates.", {"meeting_id": "prior-1"})
|