Files

252 lines
7.7 KiB
Python

# -*- coding: utf-8 -*-
"""Vector store — Qdrant remoto (DS920) con fallback SQLite locale su Pi ARM."""
import json
import logging
import math
import os
import sqlite3
import uuid
from typing import Optional
LOGGER = logging.getLogger("loogle_mcp.vector_store")
VECTOR_SIZE = 768
_local = sqlite3.connect(":memory:", check_same_thread=False) # placeholder
_qdrant_client = None
_qdrant_checked = False
_use_fallback = False
def _fallback_path() -> str:
return os.environ.get("MCP_VECTOR_FALLBACK", "/data/vector_fallback.db")
def _fallback_conn() -> sqlite3.Connection:
path = _fallback_path()
os.makedirs(os.path.dirname(path), exist_ok=True)
conn = sqlite3.connect(path, timeout=30)
conn.execute(
"""
CREATE TABLE IF NOT EXISTS vectors (
id TEXT PRIMARY KEY,
collection TEXT NOT NULL,
vector TEXT NOT NULL,
payload TEXT NOT NULL
)
"""
)
conn.execute("CREATE INDEX IF NOT EXISTS idx_vectors_collection ON vectors(collection)")
conn.commit()
return conn
def _get_qdrant():
global _qdrant_client, _qdrant_checked, _use_fallback
if _qdrant_checked:
return None if _use_fallback else _qdrant_client
_qdrant_checked = True
url = os.environ.get("QDRANT_URL", "").strip()
if not url:
_use_fallback = True
LOGGER.warning("QDRANT_URL non impostato — fallback SQLite")
return None
try:
from qdrant_client import QdrantClient
from qdrant_client.http import models as qm
client = QdrantClient(url=url, timeout=60)
client.get_collections()
_qdrant_client = client
globals()["qm"] = qm
LOGGER.info("Qdrant connesso: %s", url)
return client
except Exception as exc:
_use_fallback = True
LOGGER.warning("Qdrant non disponibile (%s) — fallback SQLite", exc)
return None
def ensure_collection(name: str, vector_size: int = VECTOR_SIZE) -> None:
client = _get_qdrant()
if client is None:
return
from qdrant_client.http import models as qm
names = {c.name for c in client.get_collections().collections}
if name in names:
return
client.create_collection(
collection_name=name,
vectors_config=qm.VectorParams(size=vector_size, distance=qm.Distance.COSINE),
)
def kb_collection(username: str) -> str:
return f"kb_personal_{username}"
def ctx_collection(username: str) -> str:
return f"ctx_{username}"
SHARED_COLLECTION = "kb_shared_family"
GITEA_SHARED_COLLECTION = "gitea_shared_family"
APPS_SHARED_COLLECTION = "apps_shared_family"
def gitea_collection(username: str) -> str:
return f"gitea_personal_{username}"
def _point_id(name: str) -> str:
return str(uuid.uuid5(uuid.NAMESPACE_URL, name))
def _cosine(a: list[float], b: list[float]) -> float:
dot = sum(x * y for x, y in zip(a, b))
na = math.sqrt(sum(x * x for x in a)) or 1.0
nb = math.sqrt(sum(x * x for x in b)) or 1.0
return dot / (na * nb)
def upsert_chunks(
collection: str,
ids: list[str],
vectors: list[list[float]],
payloads: list[dict],
) -> None:
if not vectors:
return
ensure_collection(collection, len(vectors[0]))
client = _get_qdrant()
if client is not None:
from qdrant_client.http import models as qm
points = [
qm.PointStruct(id=_point_id(pid), vector=vec, payload=payload)
for pid, vec, payload in zip(ids, vectors, payloads)
]
client.upsert(collection_name=collection, points=points)
return
conn = _fallback_conn()
for pid, vec, payload in zip(ids, vectors, payloads):
conn.execute(
"INSERT OR REPLACE INTO vectors(id,collection,vector,payload) VALUES (?,?,?,?)",
(_point_id(pid), collection, json.dumps(vec), json.dumps(payload, ensure_ascii=False)),
)
conn.commit()
def delete_by_doc(collection: str, doc_id: int) -> None:
client = _get_qdrant()
if client is not None:
from qdrant_client.http import models as qm
ensure_collection(collection)
client.delete(
collection_name=collection,
points_selector=qm.FilterSelector(
filter=qm.Filter(
must=[qm.FieldCondition(key="doc_id", match=qm.MatchValue(value=doc_id))]
)
),
)
return
conn = _fallback_conn()
rows = conn.execute("SELECT id,payload FROM vectors WHERE collection=?", (collection,)).fetchall()
for row_id, payload_raw in rows:
payload = json.loads(payload_raw)
if payload.get("doc_id") == doc_id:
conn.execute("DELETE FROM vectors WHERE id=?", (row_id,))
conn.commit()
def count_by_doc(collection: str, doc_id: int) -> int:
client = _get_qdrant()
if client is not None:
from qdrant_client.http import models as qm
ensure_collection(collection)
result = client.count(
collection_name=collection,
count_filter=qm.Filter(
must=[qm.FieldCondition(key="doc_id", match=qm.MatchValue(value=doc_id))]
),
exact=True,
)
return int(result.count)
conn = _fallback_conn()
rows = conn.execute("SELECT payload FROM vectors WHERE collection=?", (collection,)).fetchall()
count = 0
for (payload_raw,) in rows:
payload = json.loads(payload_raw)
if payload.get("doc_id") == doc_id:
count += 1
return count
def collection_point_count(collection: str) -> int:
client = _get_qdrant()
if client is not None:
try:
info = client.get_collection(collection)
return int(info.points_count or 0)
except Exception:
return 0
conn = _fallback_conn()
row = conn.execute(
"SELECT COUNT(*) FROM vectors WHERE collection=?", (collection,)
).fetchone()
return int(row[0] if row else 0)
def search(
collections: list[str],
vector: list[float],
limit: int = 8,
visibility_filter: Optional[dict] = None,
) -> list[dict]:
results: list[dict] = []
client = _get_qdrant()
if client is not None:
from qdrant_client.http import models as qm
for collection in collections:
ensure_collection(collection, len(vector))
flt = None
if visibility_filter:
must = [
qm.FieldCondition(key=k, match=qm.MatchValue(value=v))
for k, v in visibility_filter.items()
]
if must:
flt = qm.Filter(must=must)
hits = client.search(
collection_name=collection,
query_vector=vector,
limit=limit,
query_filter=flt,
)
for hit in hits:
payload = dict(hit.payload or {})
payload["score"] = hit.score
payload["collection"] = collection
results.append(payload)
else:
conn = _fallback_conn()
for collection in collections:
rows = conn.execute(
"SELECT vector,payload FROM vectors WHERE collection=?", (collection,)
).fetchall()
for vec_raw, payload_raw in rows:
payload = dict(json.loads(payload_raw))
if visibility_filter:
if any(payload.get(k) != v for k, v in visibility_filter.items()):
continue
score = _cosine(vector, json.loads(vec_raw))
payload["score"] = score
payload["collection"] = collection
results.append(payload)
results.sort(key=lambda x: x.get("score", 0), reverse=True)
return results[:limit]