misc: performance improvements (#140)
* misc: performance improvements * misc: performance improvements * misc: performance improvements
This commit is contained in:
parent
0d0abaaa9f
commit
eb2702bcba
25 changed files with 2981 additions and 384 deletions
|
|
@ -957,6 +957,12 @@ def create_app(
|
||||||
await memory.initialize()
|
await memory.initialize()
|
||||||
logging.info("Memory system initialized")
|
logging.info("Memory system initialized")
|
||||||
|
|
||||||
|
# Set up DB pool metrics after memory initialization
|
||||||
|
metrics_collector = get_metrics_collector()
|
||||||
|
if memory._pool is not None and hasattr(metrics_collector, "set_db_pool"):
|
||||||
|
metrics_collector.set_db_pool(memory._pool)
|
||||||
|
logging.info("DB pool metrics configured")
|
||||||
|
|
||||||
# Call HTTP extension startup hook
|
# Call HTTP extension startup hook
|
||||||
if http_extension:
|
if http_extension:
|
||||||
await http_extension.on_startup()
|
await http_extension.on_startup()
|
||||||
|
|
@ -993,6 +999,30 @@ def create_app(
|
||||||
# This is required for mounted sub-applications where lifespan may not fire
|
# This is required for mounted sub-applications where lifespan may not fire
|
||||||
app.state.memory = memory
|
app.state.memory = memory
|
||||||
|
|
||||||
|
# Add HTTP metrics middleware
|
||||||
|
@app.middleware("http")
|
||||||
|
async def http_metrics_middleware(request, call_next):
|
||||||
|
"""Record HTTP request metrics."""
|
||||||
|
# Normalize endpoint path to reduce cardinality
|
||||||
|
# Replace UUIDs and numeric IDs with placeholders
|
||||||
|
import re
|
||||||
|
|
||||||
|
from starlette.requests import Request
|
||||||
|
|
||||||
|
path = request.url.path
|
||||||
|
# Replace UUIDs
|
||||||
|
path = re.sub(r"/[0-9a-f]{8}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{4}-[0-9a-f]{12}", "/{id}", path)
|
||||||
|
# Replace numeric IDs
|
||||||
|
path = re.sub(r"/\d+(?=/|$)", "/{id}", path)
|
||||||
|
|
||||||
|
status_code = [500] # Default to 500, will be updated
|
||||||
|
metrics_collector = get_metrics_collector()
|
||||||
|
|
||||||
|
with metrics_collector.record_http_request(request.method, path, lambda: status_code[0]):
|
||||||
|
response = await call_next(request)
|
||||||
|
status_code[0] = response.status_code
|
||||||
|
return response
|
||||||
|
|
||||||
# Register all routes
|
# Register all routes
|
||||||
_register_routes(app)
|
_register_routes(app)
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -48,6 +48,7 @@ ENV_RERANKER_COHERE_MODEL = "HINDSIGHT_API_RERANKER_COHERE_MODEL"
|
||||||
|
|
||||||
ENV_RERANKER_PROVIDER = "HINDSIGHT_API_RERANKER_PROVIDER"
|
ENV_RERANKER_PROVIDER = "HINDSIGHT_API_RERANKER_PROVIDER"
|
||||||
ENV_RERANKER_LOCAL_MODEL = "HINDSIGHT_API_RERANKER_LOCAL_MODEL"
|
ENV_RERANKER_LOCAL_MODEL = "HINDSIGHT_API_RERANKER_LOCAL_MODEL"
|
||||||
|
ENV_RERANKER_LOCAL_MAX_CONCURRENT = "HINDSIGHT_API_RERANKER_LOCAL_MAX_CONCURRENT"
|
||||||
ENV_RERANKER_TEI_URL = "HINDSIGHT_API_RERANKER_TEI_URL"
|
ENV_RERANKER_TEI_URL = "HINDSIGHT_API_RERANKER_TEI_URL"
|
||||||
ENV_RERANKER_TEI_BATCH_SIZE = "HINDSIGHT_API_RERANKER_TEI_BATCH_SIZE"
|
ENV_RERANKER_TEI_BATCH_SIZE = "HINDSIGHT_API_RERANKER_TEI_BATCH_SIZE"
|
||||||
ENV_RERANKER_TEI_MAX_CONCURRENT = "HINDSIGHT_API_RERANKER_TEI_MAX_CONCURRENT"
|
ENV_RERANKER_TEI_MAX_CONCURRENT = "HINDSIGHT_API_RERANKER_TEI_MAX_CONCURRENT"
|
||||||
|
|
@ -69,6 +70,7 @@ ENV_RETAIN_MAX_COMPLETION_TOKENS = "HINDSIGHT_API_RETAIN_MAX_COMPLETION_TOKENS"
|
||||||
ENV_RETAIN_CHUNK_SIZE = "HINDSIGHT_API_RETAIN_CHUNK_SIZE"
|
ENV_RETAIN_CHUNK_SIZE = "HINDSIGHT_API_RETAIN_CHUNK_SIZE"
|
||||||
ENV_RETAIN_EXTRACT_CAUSAL_LINKS = "HINDSIGHT_API_RETAIN_EXTRACT_CAUSAL_LINKS"
|
ENV_RETAIN_EXTRACT_CAUSAL_LINKS = "HINDSIGHT_API_RETAIN_EXTRACT_CAUSAL_LINKS"
|
||||||
ENV_RETAIN_EXTRACTION_MODE = "HINDSIGHT_API_RETAIN_EXTRACTION_MODE"
|
ENV_RETAIN_EXTRACTION_MODE = "HINDSIGHT_API_RETAIN_EXTRACTION_MODE"
|
||||||
|
ENV_RETAIN_OBSERVATIONS_ASYNC = "HINDSIGHT_API_RETAIN_OBSERVATIONS_ASYNC"
|
||||||
|
|
||||||
# Optimization flags
|
# Optimization flags
|
||||||
ENV_SKIP_LLM_VERIFICATION = "HINDSIGHT_API_SKIP_LLM_VERIFICATION"
|
ENV_SKIP_LLM_VERIFICATION = "HINDSIGHT_API_SKIP_LLM_VERIFICATION"
|
||||||
|
|
@ -84,8 +86,9 @@ ENV_DB_COMMAND_TIMEOUT = "HINDSIGHT_API_DB_COMMAND_TIMEOUT"
|
||||||
ENV_DB_ACQUIRE_TIMEOUT = "HINDSIGHT_API_DB_ACQUIRE_TIMEOUT"
|
ENV_DB_ACQUIRE_TIMEOUT = "HINDSIGHT_API_DB_ACQUIRE_TIMEOUT"
|
||||||
|
|
||||||
# Background task processing
|
# Background task processing
|
||||||
ENV_TASK_BATCH_SIZE = "HINDSIGHT_API_TASK_BATCH_SIZE"
|
ENV_TASK_BACKEND = "HINDSIGHT_API_TASK_BACKEND"
|
||||||
ENV_TASK_BATCH_INTERVAL = "HINDSIGHT_API_TASK_BATCH_INTERVAL"
|
ENV_TASK_BACKEND_MEMORY_BATCH_SIZE = "HINDSIGHT_API_TASK_BACKEND_MEMORY_BATCH_SIZE"
|
||||||
|
ENV_TASK_BACKEND_MEMORY_BATCH_INTERVAL = "HINDSIGHT_API_TASK_BACKEND_MEMORY_BATCH_INTERVAL"
|
||||||
|
|
||||||
# Default values
|
# Default values
|
||||||
DEFAULT_DATABASE_URL = "pg0"
|
DEFAULT_DATABASE_URL = "pg0"
|
||||||
|
|
@ -101,6 +104,7 @@ DEFAULT_EMBEDDING_DIMENSION = 384
|
||||||
|
|
||||||
DEFAULT_RERANKER_PROVIDER = "local"
|
DEFAULT_RERANKER_PROVIDER = "local"
|
||||||
DEFAULT_RERANKER_LOCAL_MODEL = "cross-encoder/ms-marco-MiniLM-L-6-v2"
|
DEFAULT_RERANKER_LOCAL_MODEL = "cross-encoder/ms-marco-MiniLM-L-6-v2"
|
||||||
|
DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT = 4 # Limit concurrent CPU-bound reranking to prevent thrashing
|
||||||
DEFAULT_RERANKER_TEI_BATCH_SIZE = 128
|
DEFAULT_RERANKER_TEI_BATCH_SIZE = 128
|
||||||
DEFAULT_RERANKER_TEI_MAX_CONCURRENT = 8
|
DEFAULT_RERANKER_TEI_MAX_CONCURRENT = 8
|
||||||
|
|
||||||
|
|
@ -111,7 +115,7 @@ DEFAULT_HOST = "0.0.0.0"
|
||||||
DEFAULT_PORT = 8888
|
DEFAULT_PORT = 8888
|
||||||
DEFAULT_LOG_LEVEL = "info"
|
DEFAULT_LOG_LEVEL = "info"
|
||||||
DEFAULT_MCP_ENABLED = True
|
DEFAULT_MCP_ENABLED = True
|
||||||
DEFAULT_GRAPH_RETRIEVER = "bfs" # Options: "bfs", "mpfp"
|
DEFAULT_GRAPH_RETRIEVER = "mpfp" # Options: "mpfp", "bfs"
|
||||||
DEFAULT_MCP_LOCAL_BANK_ID = "mcp"
|
DEFAULT_MCP_LOCAL_BANK_ID = "mcp"
|
||||||
|
|
||||||
# Observation thresholds
|
# Observation thresholds
|
||||||
|
|
@ -124,6 +128,7 @@ DEFAULT_RETAIN_CHUNK_SIZE = 3000 # Max chars per chunk for fact extraction
|
||||||
DEFAULT_RETAIN_EXTRACT_CAUSAL_LINKS = True # Extract causal links between facts
|
DEFAULT_RETAIN_EXTRACT_CAUSAL_LINKS = True # Extract causal links between facts
|
||||||
DEFAULT_RETAIN_EXTRACTION_MODE = "concise" # Extraction mode: "concise" or "verbose"
|
DEFAULT_RETAIN_EXTRACTION_MODE = "concise" # Extraction mode: "concise" or "verbose"
|
||||||
RETAIN_EXTRACTION_MODES = ("concise", "verbose") # Allowed extraction modes
|
RETAIN_EXTRACTION_MODES = ("concise", "verbose") # Allowed extraction modes
|
||||||
|
DEFAULT_RETAIN_OBSERVATIONS_ASYNC = False # Run observation generation async (after retain completes)
|
||||||
|
|
||||||
# Database migrations
|
# Database migrations
|
||||||
DEFAULT_RUN_MIGRATIONS_ON_STARTUP = True
|
DEFAULT_RUN_MIGRATIONS_ON_STARTUP = True
|
||||||
|
|
@ -135,8 +140,9 @@ DEFAULT_DB_COMMAND_TIMEOUT = 60 # seconds
|
||||||
DEFAULT_DB_ACQUIRE_TIMEOUT = 30 # seconds
|
DEFAULT_DB_ACQUIRE_TIMEOUT = 30 # seconds
|
||||||
|
|
||||||
# Background task processing
|
# Background task processing
|
||||||
DEFAULT_TASK_BATCH_SIZE = 10
|
DEFAULT_TASK_BACKEND = "memory" # Options: "memory", "noop"
|
||||||
DEFAULT_TASK_BATCH_INTERVAL = 1.0 # seconds
|
DEFAULT_TASK_BACKEND_MEMORY_BATCH_SIZE = 10
|
||||||
|
DEFAULT_TASK_BACKEND_MEMORY_BATCH_INTERVAL = 1.0 # seconds
|
||||||
|
|
||||||
# Default MCP tool descriptions (can be customized via env vars)
|
# Default MCP tool descriptions (can be customized via env vars)
|
||||||
DEFAULT_MCP_RETAIN_DESCRIPTION = """Store important information to long-term memory.
|
DEFAULT_MCP_RETAIN_DESCRIPTION = """Store important information to long-term memory.
|
||||||
|
|
@ -230,6 +236,7 @@ class HindsightConfig:
|
||||||
retain_chunk_size: int
|
retain_chunk_size: int
|
||||||
retain_extract_causal_links: bool
|
retain_extract_causal_links: bool
|
||||||
retain_extraction_mode: str
|
retain_extraction_mode: str
|
||||||
|
retain_observations_async: bool
|
||||||
|
|
||||||
# Optimization flags
|
# Optimization flags
|
||||||
skip_llm_verification: bool
|
skip_llm_verification: bool
|
||||||
|
|
@ -245,8 +252,9 @@ class HindsightConfig:
|
||||||
db_acquire_timeout: int
|
db_acquire_timeout: int
|
||||||
|
|
||||||
# Background task processing
|
# Background task processing
|
||||||
task_batch_size: int
|
task_backend: str
|
||||||
task_batch_interval: float
|
task_backend_memory_batch_size: int
|
||||||
|
task_backend_memory_batch_interval: float
|
||||||
|
|
||||||
@classmethod
|
@classmethod
|
||||||
def from_env(cls) -> "HindsightConfig":
|
def from_env(cls) -> "HindsightConfig":
|
||||||
|
|
@ -309,6 +317,10 @@ class HindsightConfig:
|
||||||
retain_extraction_mode=_validate_extraction_mode(
|
retain_extraction_mode=_validate_extraction_mode(
|
||||||
os.getenv(ENV_RETAIN_EXTRACTION_MODE, DEFAULT_RETAIN_EXTRACTION_MODE)
|
os.getenv(ENV_RETAIN_EXTRACTION_MODE, DEFAULT_RETAIN_EXTRACTION_MODE)
|
||||||
),
|
),
|
||||||
|
retain_observations_async=os.getenv(
|
||||||
|
ENV_RETAIN_OBSERVATIONS_ASYNC, str(DEFAULT_RETAIN_OBSERVATIONS_ASYNC)
|
||||||
|
).lower()
|
||||||
|
== "true",
|
||||||
# Database migrations
|
# Database migrations
|
||||||
run_migrations_on_startup=os.getenv(ENV_RUN_MIGRATIONS_ON_STARTUP, "true").lower() == "true",
|
run_migrations_on_startup=os.getenv(ENV_RUN_MIGRATIONS_ON_STARTUP, "true").lower() == "true",
|
||||||
# Database connection pool
|
# Database connection pool
|
||||||
|
|
@ -317,8 +329,13 @@ class HindsightConfig:
|
||||||
db_command_timeout=int(os.getenv(ENV_DB_COMMAND_TIMEOUT, str(DEFAULT_DB_COMMAND_TIMEOUT))),
|
db_command_timeout=int(os.getenv(ENV_DB_COMMAND_TIMEOUT, str(DEFAULT_DB_COMMAND_TIMEOUT))),
|
||||||
db_acquire_timeout=int(os.getenv(ENV_DB_ACQUIRE_TIMEOUT, str(DEFAULT_DB_ACQUIRE_TIMEOUT))),
|
db_acquire_timeout=int(os.getenv(ENV_DB_ACQUIRE_TIMEOUT, str(DEFAULT_DB_ACQUIRE_TIMEOUT))),
|
||||||
# Background task processing
|
# Background task processing
|
||||||
task_batch_size=int(os.getenv(ENV_TASK_BATCH_SIZE, str(DEFAULT_TASK_BATCH_SIZE))),
|
task_backend=os.getenv(ENV_TASK_BACKEND, DEFAULT_TASK_BACKEND),
|
||||||
task_batch_interval=float(os.getenv(ENV_TASK_BATCH_INTERVAL, str(DEFAULT_TASK_BATCH_INTERVAL))),
|
task_backend_memory_batch_size=int(
|
||||||
|
os.getenv(ENV_TASK_BACKEND_MEMORY_BATCH_SIZE, str(DEFAULT_TASK_BACKEND_MEMORY_BATCH_SIZE))
|
||||||
|
),
|
||||||
|
task_backend_memory_batch_interval=float(
|
||||||
|
os.getenv(ENV_TASK_BACKEND_MEMORY_BATCH_INTERVAL, str(DEFAULT_TASK_BACKEND_MEMORY_BATCH_INTERVAL))
|
||||||
|
),
|
||||||
)
|
)
|
||||||
|
|
||||||
def get_llm_base_url(self) -> str:
|
def get_llm_base_url(self) -> str:
|
||||||
|
|
|
||||||
|
|
@ -10,17 +10,20 @@ import asyncio
|
||||||
import logging
|
import logging
|
||||||
import os
|
import os
|
||||||
from abc import ABC, abstractmethod
|
from abc import ABC, abstractmethod
|
||||||
|
from concurrent.futures import ThreadPoolExecutor
|
||||||
|
|
||||||
import httpx
|
import httpx
|
||||||
|
|
||||||
from ..config import (
|
from ..config import (
|
||||||
DEFAULT_RERANKER_COHERE_MODEL,
|
DEFAULT_RERANKER_COHERE_MODEL,
|
||||||
|
DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT,
|
||||||
DEFAULT_RERANKER_LOCAL_MODEL,
|
DEFAULT_RERANKER_LOCAL_MODEL,
|
||||||
DEFAULT_RERANKER_PROVIDER,
|
DEFAULT_RERANKER_PROVIDER,
|
||||||
DEFAULT_RERANKER_TEI_BATCH_SIZE,
|
DEFAULT_RERANKER_TEI_BATCH_SIZE,
|
||||||
DEFAULT_RERANKER_TEI_MAX_CONCURRENT,
|
DEFAULT_RERANKER_TEI_MAX_CONCURRENT,
|
||||||
ENV_COHERE_API_KEY,
|
ENV_COHERE_API_KEY,
|
||||||
ENV_RERANKER_COHERE_MODEL,
|
ENV_RERANKER_COHERE_MODEL,
|
||||||
|
ENV_RERANKER_LOCAL_MAX_CONCURRENT,
|
||||||
ENV_RERANKER_LOCAL_MODEL,
|
ENV_RERANKER_LOCAL_MODEL,
|
||||||
ENV_RERANKER_PROVIDER,
|
ENV_RERANKER_PROVIDER,
|
||||||
ENV_RERANKER_TEI_BATCH_SIZE,
|
ENV_RERANKER_TEI_BATCH_SIZE,
|
||||||
|
|
@ -78,25 +81,34 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||||
- Fast inference (~80ms for 100 pairs on CPU)
|
- Fast inference (~80ms for 100 pairs on CPU)
|
||||||
- Small model (80MB)
|
- Small model (80MB)
|
||||||
- Trained for passage re-ranking
|
- Trained for passage re-ranking
|
||||||
|
|
||||||
|
Uses a dedicated thread pool to limit concurrent CPU-bound work.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, model_name: str | None = None):
|
# Shared executor across all instances (one model loaded anyway)
|
||||||
|
_executor: ThreadPoolExecutor | None = None
|
||||||
|
_max_concurrent: int = 4 # Limit concurrent CPU-bound reranking calls
|
||||||
|
|
||||||
|
def __init__(self, model_name: str | None = None, max_concurrent: int = 4):
|
||||||
"""
|
"""
|
||||||
Initialize local SentenceTransformers cross-encoder.
|
Initialize local SentenceTransformers cross-encoder.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
model_name: Name of the CrossEncoder model to use.
|
model_name: Name of the CrossEncoder model to use.
|
||||||
Default: cross-encoder/ms-marco-MiniLM-L-6-v2
|
Default: cross-encoder/ms-marco-MiniLM-L-6-v2
|
||||||
|
max_concurrent: Maximum concurrent reranking calls (default: 2).
|
||||||
|
Higher values may cause CPU thrashing under load.
|
||||||
"""
|
"""
|
||||||
self.model_name = model_name or DEFAULT_RERANKER_LOCAL_MODEL
|
self.model_name = model_name or DEFAULT_RERANKER_LOCAL_MODEL
|
||||||
self._model = None
|
self._model = None
|
||||||
|
LocalSTCrossEncoder._max_concurrent = max_concurrent
|
||||||
|
|
||||||
@property
|
@property
|
||||||
def provider_name(self) -> str:
|
def provider_name(self) -> str:
|
||||||
return "local"
|
return "local"
|
||||||
|
|
||||||
async def initialize(self) -> None:
|
async def initialize(self) -> None:
|
||||||
"""Load the cross-encoder model."""
|
"""Load the cross-encoder model and initialize the executor."""
|
||||||
if self._model is not None:
|
if self._model is not None:
|
||||||
return
|
return
|
||||||
|
|
||||||
|
|
@ -108,14 +120,30 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||||
"Install it with: pip install sentence-transformers"
|
"Install it with: pip install sentence-transformers"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Note: We use CPU even when GPU/MPS is available because:
|
||||||
|
# 1. The reranker model (MiniLM) is tiny (~22M params)
|
||||||
|
# 2. Batch sizes are small (~100-200 pairs)
|
||||||
|
# 3. Data transfer overhead to GPU outweighs compute benefit
|
||||||
|
# 4. CPU inference is actually faster for this workload
|
||||||
logger.info(f"Reranker: initializing local provider with model {self.model_name}")
|
logger.info(f"Reranker: initializing local provider with model {self.model_name}")
|
||||||
self._model = CrossEncoder(self.model_name)
|
self._model = CrossEncoder(self.model_name)
|
||||||
logger.info("Reranker: local provider initialized")
|
|
||||||
|
# Initialize shared executor (limited workers naturally limits concurrency)
|
||||||
|
if LocalSTCrossEncoder._executor is None:
|
||||||
|
LocalSTCrossEncoder._executor = ThreadPoolExecutor(
|
||||||
|
max_workers=LocalSTCrossEncoder._max_concurrent,
|
||||||
|
thread_name_prefix="reranker",
|
||||||
|
)
|
||||||
|
logger.info(f"Reranker: local provider initialized (max_concurrent={LocalSTCrossEncoder._max_concurrent})")
|
||||||
|
else:
|
||||||
|
logger.info("Reranker: local provider initialized (using existing executor)")
|
||||||
|
|
||||||
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
async def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
|
||||||
"""
|
"""
|
||||||
Score query-document pairs for relevance.
|
Score query-document pairs for relevance.
|
||||||
|
|
||||||
|
Uses a dedicated thread pool with limited workers to prevent CPU thrashing.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
pairs: List of (query, document) tuples to score
|
pairs: List of (query, document) tuples to score
|
||||||
|
|
||||||
|
|
@ -125,9 +153,12 @@ class LocalSTCrossEncoder(CrossEncoderModel):
|
||||||
if self._model is None:
|
if self._model is None:
|
||||||
raise RuntimeError("Reranker not initialized. Call initialize() first.")
|
raise RuntimeError("Reranker not initialized. Call initialize() first.")
|
||||||
|
|
||||||
# Run CPU-bound inference in thread pool to avoid blocking event loop
|
# Use dedicated executor - limited workers naturally limits concurrency
|
||||||
loop = asyncio.get_event_loop()
|
loop = asyncio.get_event_loop()
|
||||||
scores = await loop.run_in_executor(None, lambda: self._model.predict(pairs, show_progress_bar=False))
|
scores = await loop.run_in_executor(
|
||||||
|
LocalSTCrossEncoder._executor,
|
||||||
|
lambda: self._model.predict(pairs, show_progress_bar=False),
|
||||||
|
)
|
||||||
return scores.tolist() if hasattr(scores, "tolist") else list(scores)
|
return scores.tolist() if hasattr(scores, "tolist") else list(scores)
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -301,8 +332,7 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
|
||||||
semaphore = asyncio.Semaphore(self.max_concurrent)
|
semaphore = asyncio.Semaphore(self.max_concurrent)
|
||||||
|
|
||||||
tasks = [
|
tasks = [
|
||||||
self._rerank_query_group(self._async_client, semaphore, query, texts)
|
self._rerank_query_group(self._async_client, semaphore, query, texts) for query, _, texts in tasks_info
|
||||||
for query, _, texts in tasks_info
|
|
||||||
]
|
]
|
||||||
results = await asyncio.gather(*tasks)
|
results = await asyncio.gather(*tasks)
|
||||||
|
|
||||||
|
|
@ -449,7 +479,10 @@ def create_cross_encoder_from_env() -> CrossEncoderModel:
|
||||||
elif provider == "local":
|
elif provider == "local":
|
||||||
model = os.environ.get(ENV_RERANKER_LOCAL_MODEL)
|
model = os.environ.get(ENV_RERANKER_LOCAL_MODEL)
|
||||||
model_name = model or DEFAULT_RERANKER_LOCAL_MODEL
|
model_name = model or DEFAULT_RERANKER_LOCAL_MODEL
|
||||||
return LocalSTCrossEncoder(model_name=model_name)
|
max_concurrent = int(
|
||||||
|
os.environ.get(ENV_RERANKER_LOCAL_MAX_CONCURRENT, str(DEFAULT_RERANKER_LOCAL_MAX_CONCURRENT))
|
||||||
|
)
|
||||||
|
return LocalSTCrossEncoder(model_name=model_name, max_concurrent=max_concurrent)
|
||||||
elif provider == "cohere":
|
elif provider == "cohere":
|
||||||
api_key = os.environ.get(ENV_COHERE_API_KEY)
|
api_key = os.environ.get(ENV_COHERE_API_KEY)
|
||||||
if not api_key:
|
if not api_key:
|
||||||
|
|
|
||||||
|
|
@ -83,11 +83,22 @@ async def acquire_with_retry(pool: asyncpg.Pool, max_retries: int = DEFAULT_MAX_
|
||||||
Yields:
|
Yields:
|
||||||
An asyncpg connection
|
An asyncpg connection
|
||||||
"""
|
"""
|
||||||
|
import time
|
||||||
|
|
||||||
|
start = time.time()
|
||||||
|
|
||||||
async def acquire():
|
async def acquire():
|
||||||
return await pool.acquire()
|
return await pool.acquire()
|
||||||
|
|
||||||
conn = await retry_with_backoff(acquire, max_retries=max_retries)
|
conn = await retry_with_backoff(acquire, max_retries=max_retries)
|
||||||
|
acquire_time = time.time() - start
|
||||||
|
|
||||||
|
# Log slow connection acquisitions (indicates pool contention)
|
||||||
|
if acquire_time > 0.05: # 50ms threshold
|
||||||
|
pool_size = pool.get_size()
|
||||||
|
pool_free = pool.get_idle_size()
|
||||||
|
logger.warning(f"[DB POOL] Slow acquire: {acquire_time:.3f}s | size={pool_size}, idle={pool_free}")
|
||||||
|
|
||||||
try:
|
try:
|
||||||
yield conn
|
yield conn
|
||||||
finally:
|
finally:
|
||||||
|
|
|
||||||
|
|
@ -150,7 +150,7 @@ from .retain import bank_utils, embedding_utils
|
||||||
from .retain.types import RetainContentDict
|
from .retain.types import RetainContentDict
|
||||||
from .search import observation_utils, think_utils
|
from .search import observation_utils, think_utils
|
||||||
from .search.reranking import CrossEncoderReranker
|
from .search.reranking import CrossEncoderReranker
|
||||||
from .task_backend import AsyncIOQueueBackend, TaskBackend
|
from .task_backend import AsyncIOQueueBackend, NoopTaskBackend, TaskBackend
|
||||||
|
|
||||||
|
|
||||||
class Budget(str, Enum):
|
class Budget(str, Enum):
|
||||||
|
|
@ -257,8 +257,8 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
db_command_timeout: PostgreSQL command timeout in seconds. Defaults to HINDSIGHT_API_DB_COMMAND_TIMEOUT.
|
db_command_timeout: PostgreSQL command timeout in seconds. Defaults to HINDSIGHT_API_DB_COMMAND_TIMEOUT.
|
||||||
db_acquire_timeout: Connection acquisition timeout in seconds. Defaults to HINDSIGHT_API_DB_ACQUIRE_TIMEOUT.
|
db_acquire_timeout: Connection acquisition timeout in seconds. Defaults to HINDSIGHT_API_DB_ACQUIRE_TIMEOUT.
|
||||||
task_backend: Custom task backend. If not provided, uses AsyncIOQueueBackend.
|
task_backend: Custom task backend. If not provided, uses AsyncIOQueueBackend.
|
||||||
task_batch_size: Background task batch size. Defaults to HINDSIGHT_API_TASK_BATCH_SIZE.
|
task_batch_size: Background task batch size. Defaults to HINDSIGHT_API_TASK_BACKEND_MEMORY_BATCH_SIZE.
|
||||||
task_batch_interval: Background task batch interval in seconds. Defaults to HINDSIGHT_API_TASK_BATCH_INTERVAL.
|
task_batch_interval: Background task batch interval in seconds. Defaults to HINDSIGHT_API_TASK_BACKEND_MEMORY_BATCH_INTERVAL.
|
||||||
run_migrations: Whether to run database migrations during initialize(). Default: True
|
run_migrations: Whether to run database migrations during initialize(). Default: True
|
||||||
operation_validator: Optional extension to validate operations before execution.
|
operation_validator: Optional extension to validate operations before execution.
|
||||||
If provided, retain/recall/reflect operations will be validated.
|
If provided, retain/recall/reflect operations will be validated.
|
||||||
|
|
@ -396,11 +396,17 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
self._cross_encoder_reranker = CrossEncoderReranker(cross_encoder=cross_encoder)
|
self._cross_encoder_reranker = CrossEncoderReranker(cross_encoder=cross_encoder)
|
||||||
|
|
||||||
# Initialize task backend
|
# Initialize task backend
|
||||||
_task_batch_size = task_batch_size if task_batch_size is not None else config.task_batch_size
|
if task_backend:
|
||||||
_task_batch_interval = task_batch_interval if task_batch_interval is not None else config.task_batch_interval
|
self._task_backend = task_backend
|
||||||
self._task_backend = task_backend or AsyncIOQueueBackend(
|
elif config.task_backend == "noop":
|
||||||
batch_size=_task_batch_size, batch_interval=_task_batch_interval
|
self._task_backend = NoopTaskBackend()
|
||||||
)
|
else:
|
||||||
|
# Default to memory (AsyncIOQueueBackend)
|
||||||
|
_task_batch_size = task_batch_size if task_batch_size is not None else config.task_backend_memory_batch_size
|
||||||
|
_task_batch_interval = (
|
||||||
|
task_batch_interval if task_batch_interval is not None else config.task_backend_memory_batch_interval
|
||||||
|
)
|
||||||
|
self._task_backend = AsyncIOQueueBackend(batch_size=_task_batch_size, batch_interval=_task_batch_interval)
|
||||||
|
|
||||||
# Backpressure mechanism: limit concurrent searches to prevent overwhelming the database
|
# Backpressure mechanism: limit concurrent searches to prevent overwhelming the database
|
||||||
# Limit concurrent searches to prevent connection pool exhaustion
|
# Limit concurrent searches to prevent connection pool exhaustion
|
||||||
|
|
@ -1605,19 +1611,38 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
step_start = time.time()
|
step_start = time.time()
|
||||||
query_embedding_str = str(query_embedding)
|
query_embedding_str = str(query_embedding)
|
||||||
|
|
||||||
from .search.retrieval import retrieve_parallel
|
from .search.retrieval import get_default_graph_retriever, retrieve_parallel
|
||||||
|
from .search.temporal_extraction import extract_temporal_constraint
|
||||||
|
|
||||||
# Track each retrieval start time
|
# Track each retrieval start time
|
||||||
retrieval_start = time.time()
|
retrieval_start = time.time()
|
||||||
|
|
||||||
|
# Pre-extract temporal constraint once (shared across all fact types)
|
||||||
|
tc_start = time.time()
|
||||||
|
temporal_constraint = extract_temporal_constraint(
|
||||||
|
query, reference_date=question_date, analyzer=self.query_analyzer
|
||||||
|
)
|
||||||
|
tc_duration = time.time() - tc_start
|
||||||
|
|
||||||
# Run retrieval for each fact type in parallel
|
# Run retrieval for each fact type in parallel
|
||||||
|
# MPFP does lazy edge loading internally, no need to pre-load adjacency
|
||||||
retrieval_tasks = [
|
retrieval_tasks = [
|
||||||
retrieve_parallel(
|
retrieve_parallel(
|
||||||
pool, query, query_embedding_str, bank_id, ft, thinking_budget, question_date, self.query_analyzer
|
pool,
|
||||||
|
query,
|
||||||
|
query_embedding_str,
|
||||||
|
bank_id,
|
||||||
|
ft,
|
||||||
|
thinking_budget,
|
||||||
|
question_date,
|
||||||
|
self.query_analyzer,
|
||||||
|
temporal_constraint=temporal_constraint,
|
||||||
)
|
)
|
||||||
for ft in fact_type
|
for ft in fact_type
|
||||||
]
|
]
|
||||||
|
parallel_start = time.time()
|
||||||
all_retrievals = await asyncio.gather(*retrieval_tasks)
|
all_retrievals = await asyncio.gather(*retrieval_tasks)
|
||||||
|
parallel_duration = time.time() - parallel_start
|
||||||
|
|
||||||
# Combine all results from all fact types and aggregate timings
|
# Combine all results from all fact types and aggregate timings
|
||||||
semantic_results = []
|
semantic_results = []
|
||||||
|
|
@ -1625,6 +1650,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
graph_results = []
|
graph_results = []
|
||||||
temporal_results = []
|
temporal_results = []
|
||||||
aggregated_timings = {"semantic": 0.0, "bm25": 0.0, "graph": 0.0, "temporal": 0.0}
|
aggregated_timings = {"semantic": 0.0, "bm25": 0.0, "graph": 0.0, "temporal": 0.0}
|
||||||
|
all_mpfp_timings = []
|
||||||
|
|
||||||
detected_temporal_constraint = None
|
detected_temporal_constraint = None
|
||||||
for idx, retrieval_result in enumerate(all_retrievals):
|
for idx, retrieval_result in enumerate(all_retrievals):
|
||||||
|
|
@ -1645,6 +1671,8 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
# Capture temporal constraint (same across all fact types)
|
# Capture temporal constraint (same across all fact types)
|
||||||
if retrieval_result.temporal_constraint:
|
if retrieval_result.temporal_constraint:
|
||||||
detected_temporal_constraint = retrieval_result.temporal_constraint
|
detected_temporal_constraint = retrieval_result.temporal_constraint
|
||||||
|
# Collect MPFP timings
|
||||||
|
all_mpfp_timings.extend(retrieval_result.mpfp_timings)
|
||||||
|
|
||||||
# If no temporal results from any fact type, set to None
|
# If no temporal results from any fact type, set to None
|
||||||
if not temporal_results:
|
if not temporal_results:
|
||||||
|
|
@ -1663,8 +1691,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
retrieval_duration = time.time() - retrieval_start
|
retrieval_duration = time.time() - retrieval_start
|
||||||
|
|
||||||
step_duration = time.time() - step_start
|
step_duration = time.time() - step_start
|
||||||
total_retrievals = len(fact_type) * (4 if temporal_results else 3)
|
# Format per-method timings (these are the actual parallel retrieval times)
|
||||||
# Format per-method timings
|
|
||||||
timing_parts = [
|
timing_parts = [
|
||||||
f"semantic={len(semantic_results)}({aggregated_timings['semantic']:.3f}s)",
|
f"semantic={len(semantic_results)}({aggregated_timings['semantic']:.3f}s)",
|
||||||
f"bm25={len(bm25_results)}({aggregated_timings['bm25']:.3f}s)",
|
f"bm25={len(bm25_results)}({aggregated_timings['bm25']:.3f}s)",
|
||||||
|
|
@ -1676,8 +1703,10 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
temporal_count = len(temporal_results) if temporal_results else 0
|
temporal_count = len(temporal_results) if temporal_results else 0
|
||||||
timing_parts.append(f"temporal={temporal_count}({aggregated_timings['temporal']:.3f}s)")
|
timing_parts.append(f"temporal={temporal_count}({aggregated_timings['temporal']:.3f}s)")
|
||||||
temporal_info = f" | temporal_range={start_dt.strftime('%Y-%m-%d')} to {end_dt.strftime('%Y-%m-%d')}"
|
temporal_info = f" | temporal_range={start_dt.strftime('%Y-%m-%d')} to {end_dt.strftime('%Y-%m-%d')}"
|
||||||
|
# Only tc is sequential setup now (adjacency loads in parallel with retrieval)
|
||||||
|
setup_info = f", tc={tc_duration:.3f}s" if tc_duration > 0.01 else ""
|
||||||
log_buffer.append(
|
log_buffer.append(
|
||||||
f" [2] {total_retrievals}-way retrieval ({len(fact_type)} fact_types): {', '.join(timing_parts)} in {step_duration:.3f}s{temporal_info}"
|
f" [2] Parallel retrieval ({len(fact_type)} fact_types): {', '.join(timing_parts)} in {parallel_duration:.3f}s{setup_info}{temporal_info}"
|
||||||
)
|
)
|
||||||
|
|
||||||
# Record retrieval results for tracer - per fact type
|
# Record retrieval results for tracer - per fact type
|
||||||
|
|
@ -1831,9 +1860,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
|
|
||||||
# Re-sort by combined score
|
# Re-sort by combined score
|
||||||
scored_results.sort(key=lambda x: x.weight, reverse=True)
|
scored_results.sort(key=lambda x: x.weight, reverse=True)
|
||||||
log_buffer.append(
|
|
||||||
" [4.6] Combined scoring: cross_encoder(0.6) + rrf(0.2) + temporal(0.1) + recency(0.1)"
|
|
||||||
)
|
|
||||||
|
|
||||||
# Add reranked results to tracer AFTER combined scoring (so normalized values are included)
|
# Add reranked results to tracer AFTER combined scoring (so normalized values are included)
|
||||||
if tracer:
|
if tracer:
|
||||||
|
|
@ -1852,7 +1878,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
# Step 5: Truncate to thinking_budget * 2 for token filtering
|
# Step 5: Truncate to thinking_budget * 2 for token filtering
|
||||||
rerank_limit = thinking_budget * 2
|
rerank_limit = thinking_budget * 2
|
||||||
top_scored = scored_results[:rerank_limit]
|
top_scored = scored_results[:rerank_limit]
|
||||||
log_buffer.append(f" [5] Truncated to top {len(top_scored)} results")
|
|
||||||
|
|
||||||
# Step 6: Token budget filtering
|
# Step 6: Token budget filtering
|
||||||
step_start = time.time()
|
step_start = time.time()
|
||||||
|
|
@ -1867,7 +1892,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
|
|
||||||
step_duration = time.time() - step_start
|
step_duration = time.time() - step_start
|
||||||
log_buffer.append(
|
log_buffer.append(
|
||||||
f" [6] Token filtering: {len(top_scored)} results, {total_tokens}/{max_tokens} tokens in {step_duration:.3f}s"
|
f" [5] Token filtering: {len(top_scored)} results, {total_tokens}/{max_tokens} tokens in {step_duration:.3f}s"
|
||||||
)
|
)
|
||||||
|
|
||||||
if tracer:
|
if tracer:
|
||||||
|
|
@ -1901,7 +1926,6 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
visited_ids = list(set([sr.id for sr in scored_results[:50]])) # Top 50
|
visited_ids = list(set([sr.id for sr in scored_results[:50]])) # Top 50
|
||||||
if visited_ids:
|
if visited_ids:
|
||||||
await self._task_backend.submit_task({"type": "access_count_update", "node_ids": visited_ids})
|
await self._task_backend.submit_task({"type": "access_count_update", "node_ids": visited_ids})
|
||||||
log_buffer.append(f" [7] Queued access count updates for {len(visited_ids)} nodes")
|
|
||||||
|
|
||||||
# Log fact_type distribution in results
|
# Log fact_type distribution in results
|
||||||
fact_type_counts = {}
|
fact_type_counts = {}
|
||||||
|
|
@ -1934,6 +1958,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
top_results_dicts.append(result_dict)
|
top_results_dicts.append(result_dict)
|
||||||
|
|
||||||
# Get entities for each fact if include_entities is requested
|
# Get entities for each fact if include_entities is requested
|
||||||
|
step_start = time.time()
|
||||||
fact_entity_map = {} # unit_id -> list of (entity_id, entity_name)
|
fact_entity_map = {} # unit_id -> list of (entity_id, entity_name)
|
||||||
if include_entities and top_scored:
|
if include_entities and top_scored:
|
||||||
unit_ids = [uuid.UUID(sr.id) for sr in top_scored]
|
unit_ids = [uuid.UUID(sr.id) for sr in top_scored]
|
||||||
|
|
@ -1955,6 +1980,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
fact_entity_map[unit_id].append(
|
fact_entity_map[unit_id].append(
|
||||||
{"entity_id": str(row["entity_id"]), "canonical_name": row["canonical_name"]}
|
{"entity_id": str(row["entity_id"]), "canonical_name": row["canonical_name"]}
|
||||||
)
|
)
|
||||||
|
entity_map_duration = time.time() - step_start
|
||||||
|
|
||||||
# Convert results to MemoryFact objects
|
# Convert results to MemoryFact objects
|
||||||
memory_facts = []
|
memory_facts = []
|
||||||
|
|
@ -1981,6 +2007,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
)
|
)
|
||||||
|
|
||||||
# Fetch entity observations if requested
|
# Fetch entity observations if requested
|
||||||
|
step_start = time.time()
|
||||||
entities_dict = None
|
entities_dict = None
|
||||||
total_entity_tokens = 0
|
total_entity_tokens = 0
|
||||||
total_chunk_tokens = 0
|
total_chunk_tokens = 0
|
||||||
|
|
@ -2001,7 +2028,13 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
entities_ordered.append((entity_id, entity_name))
|
entities_ordered.append((entity_id, entity_name))
|
||||||
seen_entity_ids.add(entity_id)
|
seen_entity_ids.add(entity_id)
|
||||||
|
|
||||||
# Fetch observations for each entity (respect token budget, in order)
|
# Fetch all observations in a single batched query
|
||||||
|
entity_ids = [eid for eid, _ in entities_ordered]
|
||||||
|
all_observations = await self.get_entity_observations_batch(
|
||||||
|
bank_id, entity_ids, limit_per_entity=5, request_context=request_context
|
||||||
|
)
|
||||||
|
|
||||||
|
# Build entities_dict respecting token budget, in relevance order
|
||||||
entities_dict = {}
|
entities_dict = {}
|
||||||
encoding = _get_tiktoken_encoding()
|
encoding = _get_tiktoken_encoding()
|
||||||
|
|
||||||
|
|
@ -2009,9 +2042,7 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
if total_entity_tokens >= max_entity_tokens:
|
if total_entity_tokens >= max_entity_tokens:
|
||||||
break
|
break
|
||||||
|
|
||||||
observations = await self.get_entity_observations(
|
observations = all_observations.get(entity_id, [])
|
||||||
bank_id, entity_id, limit=5, request_context=request_context
|
|
||||||
)
|
|
||||||
|
|
||||||
# Calculate tokens for this entity's observations
|
# Calculate tokens for this entity's observations
|
||||||
entity_tokens = 0
|
entity_tokens = 0
|
||||||
|
|
@ -2029,8 +2060,10 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
entity_id=entity_id, canonical_name=entity_name, observations=included_observations
|
entity_id=entity_id, canonical_name=entity_name, observations=included_observations
|
||||||
)
|
)
|
||||||
total_entity_tokens += entity_tokens
|
total_entity_tokens += entity_tokens
|
||||||
|
entity_obs_duration = time.time() - step_start
|
||||||
|
|
||||||
# Fetch chunks if requested
|
# Fetch chunks if requested
|
||||||
|
step_start = time.time()
|
||||||
chunks_dict = None
|
chunks_dict = None
|
||||||
if include_chunks and top_scored:
|
if include_chunks and top_scored:
|
||||||
from .response_models import ChunkInfo
|
from .response_models import ChunkInfo
|
||||||
|
|
@ -2090,6 +2123,12 @@ class MemoryEngine(MemoryEngineInterface):
|
||||||
chunk_text=chunk_text, chunk_index=row["chunk_index"], truncated=False
|
chunk_text=chunk_text, chunk_index=row["chunk_index"], truncated=False
|
||||||
)
|
)
|
||||||
total_chunk_tokens += chunk_tokens
|
total_chunk_tokens += chunk_tokens
|
||||||
|
chunks_duration = time.time() - step_start
|
||||||
|
|
||||||
|
# Log entity/chunk fetch timing (only if any enrichment was requested)
|
||||||
|
log_buffer.append(
|
||||||
|
f" [6] Response enrichment: entity_map={entity_map_duration:.3f}s, entity_obs={entity_obs_duration:.3f}s, chunks={chunks_duration:.3f}s"
|
||||||
|
)
|
||||||
|
|
||||||
# Finalize trace if enabled
|
# Finalize trace if enabled
|
||||||
trace_dict = None
|
trace_dict = None
|
||||||
|
|
@ -3485,6 +3524,64 @@ Guidelines:
|
||||||
observations.append(EntityObservation(text=row["text"], mentioned_at=mentioned_at))
|
observations.append(EntityObservation(text=row["text"], mentioned_at=mentioned_at))
|
||||||
return observations
|
return observations
|
||||||
|
|
||||||
|
async def get_entity_observations_batch(
|
||||||
|
self,
|
||||||
|
bank_id: str,
|
||||||
|
entity_ids: list[str],
|
||||||
|
*,
|
||||||
|
limit_per_entity: int = 5,
|
||||||
|
request_context: "RequestContext",
|
||||||
|
) -> dict[str, list[Any]]:
|
||||||
|
"""
|
||||||
|
Get observations for multiple entities in a single query.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
bank_id: bank IDentifier
|
||||||
|
entity_ids: List of entity UUIDs to get observations for
|
||||||
|
limit_per_entity: Maximum observations per entity
|
||||||
|
request_context: Request context for authentication.
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict mapping entity_id -> list of EntityObservation objects
|
||||||
|
"""
|
||||||
|
if not entity_ids:
|
||||||
|
return {}
|
||||||
|
|
||||||
|
await self._authenticate_tenant(request_context)
|
||||||
|
pool = await self._get_pool()
|
||||||
|
async with acquire_with_retry(pool) as conn:
|
||||||
|
# Use window function to limit observations per entity
|
||||||
|
rows = await conn.fetch(
|
||||||
|
f"""
|
||||||
|
WITH ranked AS (
|
||||||
|
SELECT
|
||||||
|
ue.entity_id,
|
||||||
|
mu.text,
|
||||||
|
mu.mentioned_at,
|
||||||
|
ROW_NUMBER() OVER (PARTITION BY ue.entity_id ORDER BY mu.mentioned_at DESC) as rn
|
||||||
|
FROM {fq_table("memory_units")} mu
|
||||||
|
JOIN {fq_table("unit_entities")} ue ON mu.id = ue.unit_id
|
||||||
|
WHERE mu.bank_id = $1
|
||||||
|
AND mu.fact_type = 'observation'
|
||||||
|
AND ue.entity_id = ANY($2::uuid[])
|
||||||
|
)
|
||||||
|
SELECT entity_id, text, mentioned_at
|
||||||
|
FROM ranked
|
||||||
|
WHERE rn <= $3
|
||||||
|
ORDER BY entity_id, rn
|
||||||
|
""",
|
||||||
|
bank_id,
|
||||||
|
[uuid.UUID(eid) for eid in entity_ids],
|
||||||
|
limit_per_entity,
|
||||||
|
)
|
||||||
|
|
||||||
|
result: dict[str, list[Any]] = {eid: [] for eid in entity_ids}
|
||||||
|
for row in rows:
|
||||||
|
entity_id = str(row["entity_id"])
|
||||||
|
mentioned_at = row["mentioned_at"].isoformat() if row["mentioned_at"] else None
|
||||||
|
result[entity_id].append(EntityObservation(text=row["text"], mentioned_at=mentioned_at))
|
||||||
|
return result
|
||||||
|
|
||||||
async def list_entities(
|
async def list_entities(
|
||||||
self,
|
self,
|
||||||
bank_id: str,
|
bank_id: str,
|
||||||
|
|
|
||||||
|
|
@ -9,6 +9,7 @@ import time
|
||||||
import uuid
|
import uuid
|
||||||
from datetime import UTC, datetime
|
from datetime import UTC, datetime
|
||||||
|
|
||||||
|
from ...config import get_config
|
||||||
from ..db_utils import acquire_with_retry
|
from ..db_utils import acquire_with_retry
|
||||||
from . import bank_utils
|
from . import bank_utils
|
||||||
|
|
||||||
|
|
@ -395,16 +396,26 @@ async def retain_batch(
|
||||||
causal_link_count = await link_creation.create_causal_links_batch(conn, unit_ids, non_duplicate_facts)
|
causal_link_count = await link_creation.create_causal_links_batch(conn, unit_ids, non_duplicate_facts)
|
||||||
log_buffer.append(f"[10] Causal links: {causal_link_count} links in {time.time() - step_start:.3f}s")
|
log_buffer.append(f"[10] Causal links: {causal_link_count} links in {time.time() - step_start:.3f}s")
|
||||||
|
|
||||||
# Regenerate observations INSIDE transaction for atomicity
|
# Regenerate observations - sync (in transaction) or async (background task)
|
||||||
await observation_regeneration.regenerate_observations_batch(
|
config = get_config()
|
||||||
conn, embeddings_model, llm_config, bank_id, entity_links, log_buffer
|
if config.retain_observations_async:
|
||||||
)
|
# Queue for async processing after transaction commits
|
||||||
|
entity_ids_for_async = list(set(link.entity_id for link in entity_links)) if entity_links else []
|
||||||
|
log_buffer.append(
|
||||||
|
f"[11] Observations: queued {len(entity_ids_for_async)} entities for async processing"
|
||||||
|
)
|
||||||
|
else:
|
||||||
|
# Run synchronously inside transaction for atomicity
|
||||||
|
await observation_regeneration.regenerate_observations_batch(
|
||||||
|
conn, embeddings_model, llm_config, bank_id, entity_links, log_buffer
|
||||||
|
)
|
||||||
|
entity_ids_for_async = []
|
||||||
|
|
||||||
# Map results back to original content items
|
# Map results back to original content items
|
||||||
result_unit_ids = _map_results_to_contents(contents, extracted_facts, is_duplicate_flags, unit_ids)
|
result_unit_ids = _map_results_to_contents(contents, extracted_facts, is_duplicate_flags, unit_ids)
|
||||||
|
|
||||||
# Trigger background tasks AFTER transaction commits (opinion reinforcement only)
|
# Trigger background tasks AFTER transaction commits
|
||||||
await _trigger_background_tasks(task_backend, bank_id, unit_ids, non_duplicate_facts)
|
await _trigger_background_tasks(task_backend, bank_id, unit_ids, non_duplicate_facts, entity_ids_for_async)
|
||||||
|
|
||||||
# Log final summary
|
# Log final summary
|
||||||
total_time = time.time() - start_time
|
total_time = time.time() - start_time
|
||||||
|
|
@ -454,8 +465,9 @@ async def _trigger_background_tasks(
|
||||||
bank_id: str,
|
bank_id: str,
|
||||||
unit_ids: list[str],
|
unit_ids: list[str],
|
||||||
facts: list[ProcessedFact],
|
facts: list[ProcessedFact],
|
||||||
|
entity_ids_for_observations: list[str] | None = None,
|
||||||
) -> None:
|
) -> None:
|
||||||
"""Trigger opinion reinforcement as background task (after transaction commits)."""
|
"""Trigger background tasks after transaction commits."""
|
||||||
# Trigger opinion reinforcement if there are entities
|
# Trigger opinion reinforcement if there are entities
|
||||||
fact_entities = [[e.name for e in fact.entities] for fact in facts]
|
fact_entities = [[e.name for e in fact.entities] for fact in facts]
|
||||||
if any(fact_entities):
|
if any(fact_entities):
|
||||||
|
|
@ -468,3 +480,13 @@ async def _trigger_background_tasks(
|
||||||
"unit_entities": fact_entities,
|
"unit_entities": fact_entities,
|
||||||
}
|
}
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Trigger observation regeneration if async mode is enabled
|
||||||
|
if entity_ids_for_observations:
|
||||||
|
await task_backend.submit_task(
|
||||||
|
{
|
||||||
|
"type": "regenerate_observations",
|
||||||
|
"bank_id": bank_id,
|
||||||
|
"entity_ids": entity_ids_for_observations,
|
||||||
|
}
|
||||||
|
)
|
||||||
|
|
|
||||||
|
|
@ -11,7 +11,7 @@ from abc import ABC, abstractmethod
|
||||||
|
|
||||||
from ..db_utils import acquire_with_retry
|
from ..db_utils import acquire_with_retry
|
||||||
from ..memory_engine import fq_table
|
from ..memory_engine import fq_table
|
||||||
from .types import RetrievalResult
|
from .types import MPFPTimings, RetrievalResult
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
@ -42,7 +42,8 @@ class GraphRetriever(ABC):
|
||||||
query_text: str | None = None,
|
query_text: str | None = None,
|
||||||
semantic_seeds: list[RetrievalResult] | None = None,
|
semantic_seeds: list[RetrievalResult] | None = None,
|
||||||
temporal_seeds: list[RetrievalResult] | None = None,
|
temporal_seeds: list[RetrievalResult] | None = None,
|
||||||
) -> list[RetrievalResult]:
|
adjacency=None, # TypedAdjacency, optional pre-loaded graph
|
||||||
|
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
|
||||||
"""
|
"""
|
||||||
Retrieve relevant facts via graph traversal.
|
Retrieve relevant facts via graph traversal.
|
||||||
|
|
||||||
|
|
@ -55,9 +56,10 @@ class GraphRetriever(ABC):
|
||||||
query_text: Original query text (optional, for some strategies)
|
query_text: Original query text (optional, for some strategies)
|
||||||
semantic_seeds: Pre-computed semantic entry points (from semantic retrieval)
|
semantic_seeds: Pre-computed semantic entry points (from semantic retrieval)
|
||||||
temporal_seeds: Pre-computed temporal entry points (from temporal retrieval)
|
temporal_seeds: Pre-computed temporal entry points (from temporal retrieval)
|
||||||
|
adjacency: Pre-loaded typed adjacency graph (optional, for MPFP)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
List of RetrievalResult objects with activation scores set
|
Tuple of (List of RetrievalResult with activation scores, optional timing info)
|
||||||
"""
|
"""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
|
@ -111,7 +113,8 @@ class BFSGraphRetriever(GraphRetriever):
|
||||||
query_text: str | None = None,
|
query_text: str | None = None,
|
||||||
semantic_seeds: list[RetrievalResult] | None = None,
|
semantic_seeds: list[RetrievalResult] | None = None,
|
||||||
temporal_seeds: list[RetrievalResult] | None = None,
|
temporal_seeds: list[RetrievalResult] | None = None,
|
||||||
) -> list[RetrievalResult]:
|
adjacency=None, # Not used by BFS
|
||||||
|
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
|
||||||
"""
|
"""
|
||||||
Retrieve facts using BFS spreading activation.
|
Retrieve facts using BFS spreading activation.
|
||||||
|
|
||||||
|
|
@ -122,11 +125,12 @@ class BFSGraphRetriever(GraphRetriever):
|
||||||
4. Return visited nodes up to budget
|
4. Return visited nodes up to budget
|
||||||
|
|
||||||
Note: BFS finds its own entry points via embedding search.
|
Note: BFS finds its own entry points via embedding search.
|
||||||
The semantic_seeds and temporal_seeds parameters are accepted
|
The semantic_seeds, temporal_seeds, and adjacency parameters are accepted
|
||||||
for interface compatibility but not used.
|
for interface compatibility but not used.
|
||||||
"""
|
"""
|
||||||
async with acquire_with_retry(pool) as conn:
|
async with acquire_with_retry(pool) as conn:
|
||||||
return await self._retrieve_with_conn(conn, query_embedding_str, bank_id, fact_type, budget)
|
results = await self._retrieve_with_conn(conn, query_embedding_str, bank_id, fact_type, budget)
|
||||||
|
return results, None
|
||||||
|
|
||||||
async def _retrieve_with_conn(
|
async def _retrieve_with_conn(
|
||||||
self,
|
self,
|
||||||
|
|
|
||||||
|
|
@ -9,6 +9,7 @@ propagation from Approximate PPR.
|
||||||
|
|
||||||
Key properties:
|
Key properties:
|
||||||
- Sublinear in graph size (threshold pruning bounds active nodes)
|
- Sublinear in graph size (threshold pruning bounds active nodes)
|
||||||
|
- Lazy edge loading: only loads edges for frontier nodes, not entire graph
|
||||||
- Predefined patterns capture different retrieval intents
|
- Predefined patterns capture different retrieval intents
|
||||||
- All patterns run in parallel, results fused via RRF
|
- All patterns run in parallel, results fused via RRF
|
||||||
- No LLM in the loop during traversal
|
- No LLM in the loop during traversal
|
||||||
|
|
@ -22,7 +23,7 @@ from dataclasses import dataclass, field
|
||||||
from ..db_utils import acquire_with_retry
|
from ..db_utils import acquire_with_retry
|
||||||
from ..memory_engine import fq_table
|
from ..memory_engine import fq_table
|
||||||
from .graph_retrieval import GraphRetriever
|
from .graph_retrieval import GraphRetriever
|
||||||
from .types import RetrievalResult
|
from .types import MPFPTimings, RetrievalResult
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
@ -41,11 +42,18 @@ class EdgeTarget:
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class TypedAdjacency:
|
class EdgeCache:
|
||||||
"""Adjacency lists split by edge type."""
|
"""
|
||||||
|
Cache for lazily-loaded edges.
|
||||||
|
|
||||||
# edge_type -> from_node_id -> list of (to_node_id, weight)
|
Grows per-hop as edges are loaded for frontier nodes.
|
||||||
|
Shared across patterns to avoid redundant loads.
|
||||||
|
"""
|
||||||
|
|
||||||
|
# edge_type -> from_node_id -> list of EdgeTarget
|
||||||
graphs: dict[str, dict[str, list[EdgeTarget]]] = field(default_factory=dict)
|
graphs: dict[str, dict[str, list[EdgeTarget]]] = field(default_factory=dict)
|
||||||
|
# Track which (edge_type, node_id) have been loaded
|
||||||
|
_loaded: set[tuple[str, str]] = field(default_factory=set)
|
||||||
|
|
||||||
def get_neighbors(self, edge_type: str, node_id: str) -> list[EdgeTarget]:
|
def get_neighbors(self, edge_type: str, node_id: str) -> list[EdgeTarget]:
|
||||||
"""Get neighbors for a node via a specific edge type."""
|
"""Get neighbors for a node via a specific edge type."""
|
||||||
|
|
@ -63,6 +71,33 @@ class TypedAdjacency:
|
||||||
|
|
||||||
return [EdgeTarget(node_id=n.node_id, weight=n.weight / total) for n in neighbors]
|
return [EdgeTarget(node_id=n.node_id, weight=n.weight / total) for n in neighbors]
|
||||||
|
|
||||||
|
def is_loaded(self, edge_type: str, node_id: str) -> bool:
|
||||||
|
"""Check if edges for this node+type have been loaded."""
|
||||||
|
return (edge_type, node_id) in self._loaded
|
||||||
|
|
||||||
|
def get_uncached(self, edge_type: str, node_ids: list[str]) -> list[str]:
|
||||||
|
"""Get node IDs that haven't been loaded yet for this edge type."""
|
||||||
|
return [n for n in node_ids if not self.is_loaded(edge_type, n)]
|
||||||
|
|
||||||
|
def add_edges(self, edge_type: str, edges: dict[str, list[EdgeTarget]], all_queried: list[str]):
|
||||||
|
"""
|
||||||
|
Add loaded edges to the cache.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
edge_type: Type of edges
|
||||||
|
edges: Dict mapping from_node_id -> list of EdgeTarget
|
||||||
|
all_queried: All node IDs that were queried (marks them as loaded even if no edges)
|
||||||
|
"""
|
||||||
|
if edge_type not in self.graphs:
|
||||||
|
self.graphs[edge_type] = {}
|
||||||
|
|
||||||
|
for node_id, neighbors in edges.items():
|
||||||
|
self.graphs[edge_type][node_id] = neighbors
|
||||||
|
|
||||||
|
# Mark all queried nodes as loaded (even if they have no edges)
|
||||||
|
for node_id in all_queried:
|
||||||
|
self._loaded.add((edge_type, node_id))
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class PatternResult:
|
class PatternResult:
|
||||||
|
|
@ -109,24 +144,80 @@ class SeedNode:
|
||||||
|
|
||||||
|
|
||||||
# -----------------------------------------------------------------------------
|
# -----------------------------------------------------------------------------
|
||||||
# Core Algorithm
|
# Lazy Edge Loading
|
||||||
# -----------------------------------------------------------------------------
|
# -----------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
def mpfp_traverse(
|
async def load_edges_for_frontier(
|
||||||
seeds: list[SeedNode],
|
pool,
|
||||||
pattern: list[str],
|
bank_id: str,
|
||||||
adjacency: TypedAdjacency,
|
edge_type: str,
|
||||||
config: MPFPConfig,
|
node_ids: list[str],
|
||||||
) -> PatternResult:
|
) -> dict[str, list[EdgeTarget]]:
|
||||||
"""
|
"""
|
||||||
Forward Push traversal following a meta-path pattern.
|
Load edges for specific frontier nodes only.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
|
pool: Database connection pool
|
||||||
|
bank_id: Memory bank ID
|
||||||
|
edge_type: Type of edges to load
|
||||||
|
node_ids: Frontier node IDs to load edges for
|
||||||
|
|
||||||
|
Returns:
|
||||||
|
Dict mapping from_node_id -> list of EdgeTarget
|
||||||
|
"""
|
||||||
|
if not node_ids:
|
||||||
|
return {}
|
||||||
|
|
||||||
|
async with acquire_with_retry(pool) as conn:
|
||||||
|
rows = await conn.fetch(
|
||||||
|
f"""
|
||||||
|
SELECT ml.from_unit_id, ml.to_unit_id, ml.weight
|
||||||
|
FROM {fq_table("memory_links")} ml
|
||||||
|
WHERE ml.from_unit_id = ANY($1::uuid[])
|
||||||
|
AND ml.link_type = $2
|
||||||
|
AND ml.weight >= 0.1
|
||||||
|
ORDER BY ml.from_unit_id, ml.weight DESC
|
||||||
|
""",
|
||||||
|
node_ids,
|
||||||
|
edge_type,
|
||||||
|
)
|
||||||
|
|
||||||
|
result: dict[str, list[EdgeTarget]] = defaultdict(list)
|
||||||
|
for row in rows:
|
||||||
|
from_id = str(row["from_unit_id"])
|
||||||
|
to_id = str(row["to_unit_id"])
|
||||||
|
weight = row["weight"]
|
||||||
|
result[from_id].append(EdgeTarget(node_id=to_id, weight=weight))
|
||||||
|
|
||||||
|
return dict(result)
|
||||||
|
|
||||||
|
|
||||||
|
# -----------------------------------------------------------------------------
|
||||||
|
# Core Algorithm (Async with Lazy Loading)
|
||||||
|
# -----------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
|
async def mpfp_traverse_async(
|
||||||
|
pool,
|
||||||
|
bank_id: str,
|
||||||
|
seeds: list[SeedNode],
|
||||||
|
pattern: list[str],
|
||||||
|
config: MPFPConfig,
|
||||||
|
cache: EdgeCache,
|
||||||
|
) -> PatternResult:
|
||||||
|
"""
|
||||||
|
Async Forward Push traversal with lazy edge loading.
|
||||||
|
|
||||||
|
Loads edges on-demand per hop, only for frontier nodes.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
pool: Database connection pool
|
||||||
|
bank_id: Memory bank ID
|
||||||
seeds: Entry point nodes with initial scores
|
seeds: Entry point nodes with initial scores
|
||||||
pattern: Sequence of edge types to follow
|
pattern: Sequence of edge types to follow
|
||||||
adjacency: Typed adjacency structure
|
|
||||||
config: Algorithm parameters
|
config: Algorithm parameters
|
||||||
|
cache: Shared edge cache (grows as edges are loaded)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
PatternResult with accumulated scores per node
|
PatternResult with accumulated scores per node
|
||||||
|
|
@ -145,6 +236,21 @@ def mpfp_traverse(
|
||||||
|
|
||||||
# Follow pattern hop by hop
|
# Follow pattern hop by hop
|
||||||
for edge_type in pattern:
|
for edge_type in pattern:
|
||||||
|
# Collect frontier nodes above threshold
|
||||||
|
active_nodes = [node_id for node_id, mass in frontier.items() if mass >= config.threshold]
|
||||||
|
|
||||||
|
if not active_nodes:
|
||||||
|
break
|
||||||
|
|
||||||
|
# Find nodes that need edge loading
|
||||||
|
uncached = cache.get_uncached(edge_type, active_nodes)
|
||||||
|
|
||||||
|
# Batch load edges for uncached nodes
|
||||||
|
if uncached:
|
||||||
|
edges = await load_edges_for_frontier(pool, bank_id, edge_type, uncached)
|
||||||
|
cache.add_edges(edge_type, edges, uncached)
|
||||||
|
|
||||||
|
# Propagate mass
|
||||||
next_frontier: dict[str, float] = {}
|
next_frontier: dict[str, float] = {}
|
||||||
|
|
||||||
for node_id, mass in frontier.items():
|
for node_id, mass in frontier.items():
|
||||||
|
|
@ -156,7 +262,7 @@ def mpfp_traverse(
|
||||||
|
|
||||||
# Push (1-α) to neighbors
|
# Push (1-α) to neighbors
|
||||||
push_mass = (1 - config.alpha) * mass
|
push_mass = (1 - config.alpha) * mass
|
||||||
neighbors = adjacency.get_normalized_neighbors(edge_type, node_id, config.top_k_neighbors)
|
neighbors = cache.get_normalized_neighbors(edge_type, node_id, config.top_k_neighbors)
|
||||||
|
|
||||||
for neighbor in neighbors:
|
for neighbor in neighbors:
|
||||||
next_frontier[neighbor.node_id] = next_frontier.get(neighbor.node_id, 0) + push_mass * neighbor.weight
|
next_frontier[neighbor.node_id] = next_frontier.get(neighbor.node_id, 0) + push_mass * neighbor.weight
|
||||||
|
|
@ -210,38 +316,6 @@ def rrf_fusion(
|
||||||
# -----------------------------------------------------------------------------
|
# -----------------------------------------------------------------------------
|
||||||
|
|
||||||
|
|
||||||
async def load_typed_adjacency(pool, bank_id: str) -> TypedAdjacency:
|
|
||||||
"""
|
|
||||||
Load all edges for a bank, split by edge type.
|
|
||||||
|
|
||||||
Single query, then organize in-memory for fast traversal.
|
|
||||||
"""
|
|
||||||
async with acquire_with_retry(pool) as conn:
|
|
||||||
rows = await conn.fetch(
|
|
||||||
f"""
|
|
||||||
SELECT ml.from_unit_id, ml.to_unit_id, ml.link_type, ml.weight
|
|
||||||
FROM {fq_table("memory_links")} ml
|
|
||||||
JOIN {fq_table("memory_units")} mu ON ml.from_unit_id = mu.id
|
|
||||||
WHERE mu.bank_id = $1
|
|
||||||
AND ml.weight >= 0.1
|
|
||||||
ORDER BY ml.from_unit_id, ml.weight DESC
|
|
||||||
""",
|
|
||||||
bank_id,
|
|
||||||
)
|
|
||||||
|
|
||||||
graphs: dict[str, dict[str, list[EdgeTarget]]] = defaultdict(lambda: defaultdict(list))
|
|
||||||
|
|
||||||
for row in rows:
|
|
||||||
from_id = str(row["from_unit_id"])
|
|
||||||
to_id = str(row["to_unit_id"])
|
|
||||||
link_type = row["link_type"]
|
|
||||||
weight = row["weight"]
|
|
||||||
|
|
||||||
graphs[link_type][from_id].append(EdgeTarget(node_id=to_id, weight=weight))
|
|
||||||
|
|
||||||
return TypedAdjacency(graphs=dict(graphs))
|
|
||||||
|
|
||||||
|
|
||||||
async def fetch_memory_units_by_ids(
|
async def fetch_memory_units_by_ids(
|
||||||
pool,
|
pool,
|
||||||
node_ids: list[str],
|
node_ids: list[str],
|
||||||
|
|
@ -274,10 +348,10 @@ async def fetch_memory_units_by_ids(
|
||||||
|
|
||||||
class MPFPGraphRetriever(GraphRetriever):
|
class MPFPGraphRetriever(GraphRetriever):
|
||||||
"""
|
"""
|
||||||
Graph retrieval using Meta-Path Forward Push.
|
Graph retrieval using Meta-Path Forward Push with lazy edge loading.
|
||||||
|
|
||||||
Runs predefined patterns in parallel from semantic and temporal seeds,
|
Runs predefined patterns in parallel from semantic and temporal seeds,
|
||||||
then fuses results via RRF.
|
loading edges on-demand per hop instead of loading entire graph upfront.
|
||||||
"""
|
"""
|
||||||
|
|
||||||
def __init__(self, config: MPFPConfig | None = None):
|
def __init__(self, config: MPFPConfig | None = None):
|
||||||
|
|
@ -303,9 +377,10 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||||
query_text: str | None = None,
|
query_text: str | None = None,
|
||||||
semantic_seeds: list[RetrievalResult] | None = None,
|
semantic_seeds: list[RetrievalResult] | None = None,
|
||||||
temporal_seeds: list[RetrievalResult] | None = None,
|
temporal_seeds: list[RetrievalResult] | None = None,
|
||||||
) -> list[RetrievalResult]:
|
adjacency=None, # Ignored - kept for interface compatibility
|
||||||
|
) -> tuple[list[RetrievalResult], MPFPTimings | None]:
|
||||||
"""
|
"""
|
||||||
Retrieve facts using MPFP algorithm.
|
Retrieve facts using MPFP algorithm with lazy edge loading.
|
||||||
|
|
||||||
Args:
|
Args:
|
||||||
pool: Database connection pool
|
pool: Database connection pool
|
||||||
|
|
@ -316,12 +391,14 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||||
query_text: Original query text (optional)
|
query_text: Original query text (optional)
|
||||||
semantic_seeds: Pre-computed semantic entry points
|
semantic_seeds: Pre-computed semantic entry points
|
||||||
temporal_seeds: Pre-computed temporal entry points
|
temporal_seeds: Pre-computed temporal entry points
|
||||||
|
adjacency: Ignored (kept for interface compatibility)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
List of RetrievalResult with activation scores
|
Tuple of (List of RetrievalResult with activation scores, MPFPTimings)
|
||||||
"""
|
"""
|
||||||
# Load typed adjacency (could cache per bank_id with TTL)
|
import time
|
||||||
adjacency = await load_typed_adjacency(pool, bank_id)
|
|
||||||
|
timings = MPFPTimings(fact_type=fact_type)
|
||||||
|
|
||||||
# Convert seeds to SeedNode format
|
# Convert seeds to SeedNode format
|
||||||
semantic_seed_nodes = self._convert_seeds(semantic_seeds, "similarity")
|
semantic_seed_nodes = self._convert_seeds(semantic_seeds, "similarity")
|
||||||
|
|
@ -331,52 +408,54 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||||
if not semantic_seed_nodes:
|
if not semantic_seed_nodes:
|
||||||
semantic_seed_nodes = await self._find_semantic_seeds(pool, query_embedding_str, bank_id, fact_type)
|
semantic_seed_nodes = await self._find_semantic_seeds(pool, query_embedding_str, bank_id, fact_type)
|
||||||
|
|
||||||
# Run all patterns in parallel
|
# Collect all pattern jobs
|
||||||
tasks = []
|
pattern_jobs = []
|
||||||
|
|
||||||
# Patterns from semantic seeds
|
# Patterns from semantic seeds
|
||||||
for pattern in self.config.patterns_semantic:
|
for pattern in self.config.patterns_semantic:
|
||||||
if semantic_seed_nodes:
|
if semantic_seed_nodes:
|
||||||
tasks.append(
|
pattern_jobs.append((semantic_seed_nodes, pattern))
|
||||||
asyncio.to_thread(
|
|
||||||
mpfp_traverse,
|
|
||||||
semantic_seed_nodes,
|
|
||||||
pattern,
|
|
||||||
adjacency,
|
|
||||||
self.config,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
# Patterns from temporal seeds
|
# Patterns from temporal seeds
|
||||||
for pattern in self.config.patterns_temporal:
|
for pattern in self.config.patterns_temporal:
|
||||||
if temporal_seed_nodes:
|
if temporal_seed_nodes:
|
||||||
tasks.append(
|
pattern_jobs.append((temporal_seed_nodes, pattern))
|
||||||
asyncio.to_thread(
|
|
||||||
mpfp_traverse,
|
|
||||||
temporal_seed_nodes,
|
|
||||||
pattern,
|
|
||||||
adjacency,
|
|
||||||
self.config,
|
|
||||||
)
|
|
||||||
)
|
|
||||||
|
|
||||||
if not tasks:
|
if not pattern_jobs:
|
||||||
return []
|
return [], timings
|
||||||
|
|
||||||
# Gather pattern results
|
timings.pattern_count = len(pattern_jobs)
|
||||||
pattern_results = await asyncio.gather(*tasks)
|
|
||||||
|
# Shared edge cache across all patterns
|
||||||
|
cache = EdgeCache()
|
||||||
|
|
||||||
|
# Run all patterns in parallel (each does lazy edge loading)
|
||||||
|
step_start = time.time()
|
||||||
|
pattern_tasks = [
|
||||||
|
mpfp_traverse_async(pool, bank_id, seeds, pattern, self.config, cache) for seeds, pattern in pattern_jobs
|
||||||
|
]
|
||||||
|
pattern_results = await asyncio.gather(*pattern_tasks)
|
||||||
|
timings.traverse = time.time() - step_start
|
||||||
|
|
||||||
|
# Count edges loaded
|
||||||
|
timings.edge_count = sum(len(neighbors) for g in cache.graphs.values() for neighbors in g.values())
|
||||||
|
|
||||||
# Fuse results
|
# Fuse results
|
||||||
|
step_start = time.time()
|
||||||
fused = rrf_fusion(pattern_results, top_k=budget)
|
fused = rrf_fusion(pattern_results, top_k=budget)
|
||||||
|
timings.fusion = time.time() - step_start
|
||||||
|
|
||||||
if not fused:
|
if not fused:
|
||||||
return []
|
return [], timings
|
||||||
|
|
||||||
# Get top result IDs (don't exclude seeds - they may be highly relevant)
|
# Get top result IDs
|
||||||
result_ids = [node_id for node_id, score in fused][:budget]
|
result_ids = [node_id for node_id, score in fused][:budget]
|
||||||
|
|
||||||
# Fetch full details
|
# Fetch full details
|
||||||
|
step_start = time.time()
|
||||||
results = await fetch_memory_units_by_ids(pool, result_ids, fact_type)
|
results = await fetch_memory_units_by_ids(pool, result_ids, fact_type)
|
||||||
|
timings.fetch = time.time() - step_start
|
||||||
|
timings.result_count = len(results)
|
||||||
|
|
||||||
# Add activation scores from fusion
|
# Add activation scores from fusion
|
||||||
score_map = {node_id: score for node_id, score in fused}
|
score_map = {node_id: score for node_id, score in fused}
|
||||||
|
|
@ -386,7 +465,7 @@ class MPFPGraphRetriever(GraphRetriever):
|
||||||
# Sort by activation
|
# Sort by activation
|
||||||
results.sort(key=lambda r: r.activation or 0, reverse=True)
|
results.sort(key=lambda r: r.activation or 0, reverse=True)
|
||||||
|
|
||||||
return results
|
return results, timings
|
||||||
|
|
||||||
def _convert_seeds(
|
def _convert_seeds(
|
||||||
self,
|
self,
|
||||||
|
|
|
||||||
|
|
@ -19,7 +19,7 @@ from ..db_utils import acquire_with_retry
|
||||||
from ..memory_engine import fq_table
|
from ..memory_engine import fq_table
|
||||||
from .graph_retrieval import BFSGraphRetriever, GraphRetriever
|
from .graph_retrieval import BFSGraphRetriever, GraphRetriever
|
||||||
from .mpfp_retrieval import MPFPGraphRetriever
|
from .mpfp_retrieval import MPFPGraphRetriever
|
||||||
from .types import RetrievalResult
|
from .types import MPFPTimings, RetrievalResult
|
||||||
|
|
||||||
logger = logging.getLogger(__name__)
|
logger = logging.getLogger(__name__)
|
||||||
|
|
||||||
|
|
@ -34,6 +34,7 @@ class ParallelRetrievalResult:
|
||||||
temporal: list[RetrievalResult] | None
|
temporal: list[RetrievalResult] | None
|
||||||
timings: dict[str, float] = field(default_factory=dict)
|
timings: dict[str, float] = field(default_factory=dict)
|
||||||
temporal_constraint: tuple | None = None # (start_date, end_date)
|
temporal_constraint: tuple | None = None # (start_date, end_date)
|
||||||
|
mpfp_timings: list[MPFPTimings] = field(default_factory=list) # MPFP sub-step timings per fact type
|
||||||
|
|
||||||
|
|
||||||
# Default graph retriever instance (can be overridden)
|
# Default graph retriever instance (can be overridden)
|
||||||
|
|
@ -260,94 +261,101 @@ async def retrieve_temporal(
|
||||||
ep_result.temporal_proximity = temporal_proximity
|
ep_result.temporal_proximity = temporal_proximity
|
||||||
results.append(ep_result)
|
results.append(ep_result)
|
||||||
|
|
||||||
# Spread through temporal links
|
# Spread through temporal links using BATCHED neighbor fetching
|
||||||
queue = [
|
# Map node_id -> (semantic_sim, temporal_score) for propagation
|
||||||
(RetrievalResult.from_db_row(dict(ep)), ep["similarity"], 1.0) for ep in entry_points
|
node_scores = {str(ep["id"]): (ep["similarity"], 1.0) for ep in entry_points}
|
||||||
] # (unit, semantic_sim, temporal_score)
|
frontier = list(node_scores.keys()) # Current batch of nodes to expand
|
||||||
budget_remaining = budget - len(entry_points)
|
budget_remaining = budget - len(entry_points)
|
||||||
|
batch_size = 20 # Process this many nodes per DB query
|
||||||
|
|
||||||
while queue and budget_remaining > 0:
|
while frontier and budget_remaining > 0:
|
||||||
current, semantic_sim, temporal_score = queue.pop(0)
|
# Take a batch from frontier
|
||||||
current_id = current.id
|
batch_ids = frontier[:batch_size]
|
||||||
|
frontier = frontier[batch_size:]
|
||||||
|
|
||||||
# Get neighbors via temporal and causal links
|
# Batch fetch all neighbors for this batch of nodes
|
||||||
if budget_remaining > 0:
|
neighbors = await conn.fetch(
|
||||||
neighbors = await conn.fetch(
|
f"""
|
||||||
f"""
|
SELECT mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start, mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding, mu.fact_type, mu.document_id, mu.chunk_id,
|
||||||
SELECT mu.id, mu.text, mu.context, mu.event_date, mu.occurred_start, mu.occurred_end, mu.mentioned_at, mu.access_count, mu.embedding, mu.fact_type, mu.document_id, mu.chunk_id,
|
ml.weight, ml.link_type, ml.from_unit_id,
|
||||||
ml.weight, ml.link_type,
|
1 - (mu.embedding <=> $1::vector) AS similarity
|
||||||
1 - (mu.embedding <=> $1::vector) AS similarity
|
FROM {fq_table("memory_links")} ml
|
||||||
FROM {fq_table("memory_links")} ml
|
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
|
||||||
JOIN {fq_table("memory_units")} mu ON ml.to_unit_id = mu.id
|
WHERE ml.from_unit_id = ANY($2::uuid[])
|
||||||
WHERE ml.from_unit_id = $2
|
AND ml.link_type IN ('temporal', 'causes', 'caused_by', 'enables', 'prevents')
|
||||||
AND ml.link_type IN ('temporal', 'causes', 'caused_by', 'enables', 'prevents')
|
AND ml.weight >= 0.1
|
||||||
AND ml.weight >= 0.1
|
AND mu.fact_type = $3
|
||||||
AND mu.fact_type = $3
|
AND mu.embedding IS NOT NULL
|
||||||
AND mu.embedding IS NOT NULL
|
AND (1 - (mu.embedding <=> $1::vector)) >= $4
|
||||||
AND (1 - (mu.embedding <=> $1::vector)) >= $4
|
ORDER BY ml.weight DESC
|
||||||
ORDER BY ml.weight DESC
|
LIMIT $5
|
||||||
LIMIT 10
|
""",
|
||||||
""",
|
query_emb_str,
|
||||||
query_emb_str,
|
batch_ids,
|
||||||
current.id,
|
fact_type,
|
||||||
fact_type,
|
semantic_threshold,
|
||||||
semantic_threshold,
|
batch_size * 10, # Allow up to 10 neighbors per node in batch
|
||||||
)
|
)
|
||||||
|
|
||||||
for n in neighbors:
|
for n in neighbors:
|
||||||
neighbor_id = str(n["id"])
|
neighbor_id = str(n["id"])
|
||||||
if neighbor_id in visited:
|
if neighbor_id in visited:
|
||||||
continue
|
continue
|
||||||
|
|
||||||
visited.add(neighbor_id)
|
visited.add(neighbor_id)
|
||||||
budget_remaining -= 1
|
budget_remaining -= 1
|
||||||
|
|
||||||
# Calculate temporal score for neighbor using best available date
|
# Get parent's scores for propagation
|
||||||
neighbor_best_date = None
|
parent_id = str(n["from_unit_id"])
|
||||||
if n["occurred_start"] is not None and n["occurred_end"] is not None:
|
_, parent_temporal_score = node_scores.get(parent_id, (0.5, 0.5))
|
||||||
neighbor_best_date = n["occurred_start"] + (n["occurred_end"] - n["occurred_start"]) / 2
|
|
||||||
elif n["occurred_start"] is not None:
|
|
||||||
neighbor_best_date = n["occurred_start"]
|
|
||||||
elif n["occurred_end"] is not None:
|
|
||||||
neighbor_best_date = n["occurred_end"]
|
|
||||||
elif n["mentioned_at"] is not None:
|
|
||||||
neighbor_best_date = n["mentioned_at"]
|
|
||||||
|
|
||||||
if neighbor_best_date:
|
# Calculate temporal score for neighbor using best available date
|
||||||
days_from_mid = abs((neighbor_best_date - mid_date).total_seconds() / 86400)
|
neighbor_best_date = None
|
||||||
neighbor_temporal_proximity = (
|
if n["occurred_start"] is not None and n["occurred_end"] is not None:
|
||||||
1.0 - min(days_from_mid / (total_days / 2), 1.0) if total_days > 0 else 1.0
|
neighbor_best_date = n["occurred_start"] + (n["occurred_end"] - n["occurred_start"]) / 2
|
||||||
)
|
elif n["occurred_start"] is not None:
|
||||||
else:
|
neighbor_best_date = n["occurred_start"]
|
||||||
neighbor_temporal_proximity = 0.3 # Lower score if no temporal data
|
elif n["occurred_end"] is not None:
|
||||||
|
neighbor_best_date = n["occurred_end"]
|
||||||
|
elif n["mentioned_at"] is not None:
|
||||||
|
neighbor_best_date = n["mentioned_at"]
|
||||||
|
|
||||||
# Boost causal links (same as graph retrieval)
|
if neighbor_best_date:
|
||||||
link_type = n["link_type"]
|
days_from_mid = abs((neighbor_best_date - mid_date).total_seconds() / 86400)
|
||||||
if link_type in ("causes", "caused_by"):
|
neighbor_temporal_proximity = (
|
||||||
causal_boost = 2.0
|
1.0 - min(days_from_mid / (total_days / 2), 1.0) if total_days > 0 else 1.0
|
||||||
elif link_type in ("enables", "prevents"):
|
)
|
||||||
causal_boost = 1.5
|
else:
|
||||||
else:
|
neighbor_temporal_proximity = 0.3 # Lower score if no temporal data
|
||||||
causal_boost = 1.0
|
|
||||||
|
|
||||||
# Propagate temporal score through links (decay, with causal boost)
|
# Boost causal links (same as graph retrieval)
|
||||||
propagated_temporal = temporal_score * n["weight"] * causal_boost * 0.7
|
link_type = n["link_type"]
|
||||||
|
if link_type in ("causes", "caused_by"):
|
||||||
|
causal_boost = 2.0
|
||||||
|
elif link_type in ("enables", "prevents"):
|
||||||
|
causal_boost = 1.5
|
||||||
|
else:
|
||||||
|
causal_boost = 1.0
|
||||||
|
|
||||||
# Combined temporal score
|
# Propagate temporal score through links (decay, with causal boost)
|
||||||
combined_temporal = max(neighbor_temporal_proximity, propagated_temporal)
|
propagated_temporal = parent_temporal_score * n["weight"] * causal_boost * 0.7
|
||||||
|
|
||||||
# Create RetrievalResult with temporal scores
|
# Combined temporal score
|
||||||
neighbor_result = RetrievalResult.from_db_row(dict(n))
|
combined_temporal = max(neighbor_temporal_proximity, propagated_temporal)
|
||||||
neighbor_result.temporal_score = combined_temporal
|
|
||||||
neighbor_result.temporal_proximity = neighbor_temporal_proximity
|
|
||||||
results.append(neighbor_result)
|
|
||||||
|
|
||||||
# Add to queue for further spreading
|
# Create RetrievalResult with temporal scores
|
||||||
if budget_remaining > 0 and combined_temporal > 0.2:
|
neighbor_result = RetrievalResult.from_db_row(dict(n))
|
||||||
queue.append((neighbor_result, n["similarity"], combined_temporal))
|
neighbor_result.temporal_score = combined_temporal
|
||||||
|
neighbor_result.temporal_proximity = neighbor_temporal_proximity
|
||||||
|
results.append(neighbor_result)
|
||||||
|
|
||||||
if budget_remaining <= 0:
|
# Track scores for propagation and add to frontier
|
||||||
break
|
if budget_remaining > 0 and combined_temporal > 0.2:
|
||||||
|
node_scores[neighbor_id] = (n["similarity"], combined_temporal)
|
||||||
|
frontier.append(neighbor_id)
|
||||||
|
|
||||||
|
if budget_remaining <= 0:
|
||||||
|
break
|
||||||
|
|
||||||
return results
|
return results
|
||||||
|
|
||||||
|
|
@ -362,6 +370,7 @@ async def retrieve_parallel(
|
||||||
question_date: datetime | None = None,
|
question_date: datetime | None = None,
|
||||||
query_analyzer: Optional["QueryAnalyzer"] = None,
|
query_analyzer: Optional["QueryAnalyzer"] = None,
|
||||||
graph_retriever: GraphRetriever | None = None,
|
graph_retriever: GraphRetriever | None = None,
|
||||||
|
temporal_constraint: tuple | None = None, # Pre-extracted temporal constraint
|
||||||
) -> ParallelRetrievalResult:
|
) -> ParallelRetrievalResult:
|
||||||
"""
|
"""
|
||||||
Run 3-way or 4-way parallel retrieval (adds temporal if detected).
|
Run 3-way or 4-way parallel retrieval (adds temporal if detected).
|
||||||
|
|
@ -376,19 +385,31 @@ async def retrieve_parallel(
|
||||||
question_date: Optional date when question was asked (for temporal filtering)
|
question_date: Optional date when question was asked (for temporal filtering)
|
||||||
query_analyzer: Query analyzer to use (defaults to TransformerQueryAnalyzer)
|
query_analyzer: Query analyzer to use (defaults to TransformerQueryAnalyzer)
|
||||||
graph_retriever: Graph retrieval strategy (defaults to configured retriever)
|
graph_retriever: Graph retrieval strategy (defaults to configured retriever)
|
||||||
|
temporal_constraint: Pre-extracted temporal constraint (optional)
|
||||||
|
|
||||||
Returns:
|
Returns:
|
||||||
ParallelRetrievalResult with semantic, bm25, graph, temporal results and timings
|
ParallelRetrievalResult with semantic, bm25, graph, temporal results and timings
|
||||||
"""
|
"""
|
||||||
from .temporal_extraction import extract_temporal_constraint
|
# Extract temporal constraint if not pre-provided
|
||||||
|
if temporal_constraint is None:
|
||||||
|
from .temporal_extraction import extract_temporal_constraint
|
||||||
|
|
||||||
temporal_constraint = extract_temporal_constraint(query_text, reference_date=question_date, analyzer=query_analyzer)
|
temporal_constraint = extract_temporal_constraint(
|
||||||
|
query_text, reference_date=question_date, analyzer=query_analyzer
|
||||||
|
)
|
||||||
|
|
||||||
retriever = graph_retriever or get_default_graph_retriever()
|
retriever = graph_retriever or get_default_graph_retriever()
|
||||||
|
|
||||||
if retriever.name == "mpfp":
|
if retriever.name == "mpfp":
|
||||||
return await _retrieve_parallel_mpfp(
|
return await _retrieve_parallel_mpfp(
|
||||||
pool, query_text, query_embedding_str, bank_id, fact_type, thinking_budget, temporal_constraint, retriever
|
pool,
|
||||||
|
query_text,
|
||||||
|
query_embedding_str,
|
||||||
|
bank_id,
|
||||||
|
fact_type,
|
||||||
|
thinking_budget,
|
||||||
|
temporal_constraint,
|
||||||
|
retriever,
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
return await _retrieve_parallel_bfs(
|
return await _retrieve_parallel_bfs(
|
||||||
|
|
@ -396,16 +417,6 @@ async def retrieve_parallel(
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
|
||||||
class _SemanticGraphResult:
|
|
||||||
"""Internal result from semantic→graph chain."""
|
|
||||||
|
|
||||||
semantic: list[RetrievalResult]
|
|
||||||
graph: list[RetrievalResult]
|
|
||||||
semantic_time: float
|
|
||||||
graph_time: float
|
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class _TimedResult:
|
class _TimedResult:
|
||||||
"""Internal result with timing."""
|
"""Internal result with timing."""
|
||||||
|
|
@ -425,46 +436,24 @@ async def _retrieve_parallel_mpfp(
|
||||||
retriever: GraphRetriever,
|
retriever: GraphRetriever,
|
||||||
) -> ParallelRetrievalResult:
|
) -> ParallelRetrievalResult:
|
||||||
"""
|
"""
|
||||||
MPFP retrieval with optimized parallelization.
|
MPFP retrieval with true parallelization.
|
||||||
|
|
||||||
Runs 2-3 parallel task chains:
|
All methods run independently in parallel:
|
||||||
- Task 1: Semantic → Graph (chained, graph uses semantic seeds)
|
- Semantic: vector similarity search
|
||||||
- Task 2: BM25 (independent)
|
- BM25: keyword search
|
||||||
- Task 3: Temporal (if constraint detected)
|
- Graph: MPFP traversal (does its own semantic seeds internally)
|
||||||
|
- Temporal: date-range search (if constraint detected)
|
||||||
|
|
||||||
|
Graph does its own semantic query for seeds, avoiding chain dependency.
|
||||||
"""
|
"""
|
||||||
import time
|
import time
|
||||||
|
|
||||||
async def run_semantic_then_graph() -> _SemanticGraphResult:
|
async def run_semantic() -> _TimedResult:
|
||||||
"""Chain: semantic retrieval → graph retrieval (using semantic as seeds)."""
|
"""Independent semantic retrieval."""
|
||||||
start = time.time()
|
start = time.time()
|
||||||
async with acquire_with_retry(pool) as conn:
|
async with acquire_with_retry(pool) as conn:
|
||||||
semantic = await retrieve_semantic(conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget)
|
results = await retrieve_semantic(conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget)
|
||||||
semantic_time = time.time() - start
|
return _TimedResult(results, time.time() - start)
|
||||||
|
|
||||||
# Get temporal seeds if needed (quick query, part of this chain)
|
|
||||||
temporal_seeds = None
|
|
||||||
if temporal_constraint:
|
|
||||||
tc_start, tc_end = temporal_constraint
|
|
||||||
async with acquire_with_retry(pool) as conn:
|
|
||||||
temporal_seeds = await _get_temporal_entry_points(
|
|
||||||
conn, query_embedding_str, bank_id, fact_type, tc_start, tc_end, limit=20
|
|
||||||
)
|
|
||||||
|
|
||||||
# Run graph with seeds
|
|
||||||
start = time.time()
|
|
||||||
graph = await retriever.retrieve(
|
|
||||||
pool=pool,
|
|
||||||
query_embedding_str=query_embedding_str,
|
|
||||||
bank_id=bank_id,
|
|
||||||
fact_type=fact_type,
|
|
||||||
budget=thinking_budget,
|
|
||||||
query_text=query_text,
|
|
||||||
semantic_seeds=semantic,
|
|
||||||
temporal_seeds=temporal_seeds,
|
|
||||||
)
|
|
||||||
graph_time = time.time() - start
|
|
||||||
|
|
||||||
return _SemanticGraphResult(semantic, graph, semantic_time, graph_time)
|
|
||||||
|
|
||||||
async def run_bm25() -> _TimedResult:
|
async def run_bm25() -> _TimedResult:
|
||||||
"""Independent BM25 retrieval."""
|
"""Independent BM25 retrieval."""
|
||||||
|
|
@ -473,8 +462,34 @@ async def _retrieve_parallel_mpfp(
|
||||||
results = await retrieve_bm25(conn, query_text, bank_id, fact_type, limit=thinking_budget)
|
results = await retrieve_bm25(conn, query_text, bank_id, fact_type, limit=thinking_budget)
|
||||||
return _TimedResult(results, time.time() - start)
|
return _TimedResult(results, time.time() - start)
|
||||||
|
|
||||||
|
async def run_graph() -> tuple[list[RetrievalResult], float, MPFPTimings | None]:
|
||||||
|
"""Independent graph retrieval - does its own semantic seeds."""
|
||||||
|
start = time.time()
|
||||||
|
|
||||||
|
# Get temporal seeds if needed (graph uses them for temporal patterns)
|
||||||
|
temporal_seeds = None
|
||||||
|
if temporal_constraint:
|
||||||
|
tc_start, tc_end = temporal_constraint
|
||||||
|
async with acquire_with_retry(pool) as conn:
|
||||||
|
temporal_seeds = await _get_temporal_entry_points(
|
||||||
|
conn, query_embedding_str, bank_id, fact_type, tc_start, tc_end, limit=20
|
||||||
|
)
|
||||||
|
|
||||||
|
# MPFP does its own semantic seeds via _find_semantic_seeds
|
||||||
|
results, mpfp_timing = await retriever.retrieve(
|
||||||
|
pool=pool,
|
||||||
|
query_embedding_str=query_embedding_str,
|
||||||
|
bank_id=bank_id,
|
||||||
|
fact_type=fact_type,
|
||||||
|
budget=thinking_budget,
|
||||||
|
query_text=query_text,
|
||||||
|
semantic_seeds=None, # Let MPFP find its own seeds
|
||||||
|
temporal_seeds=temporal_seeds,
|
||||||
|
)
|
||||||
|
return results, time.time() - start, mpfp_timing
|
||||||
|
|
||||||
async def run_temporal(tc_start, tc_end) -> _TimedResult:
|
async def run_temporal(tc_start, tc_end) -> _TimedResult:
|
||||||
"""Temporal retrieval (uses its own entry point finding)."""
|
"""Independent temporal retrieval."""
|
||||||
start = time.time()
|
start = time.time()
|
||||||
async with acquire_with_retry(pool) as conn:
|
async with acquire_with_retry(pool) as conn:
|
||||||
results = await retrieve_temporal(
|
results = await retrieve_temporal(
|
||||||
|
|
@ -489,43 +504,49 @@ async def _retrieve_parallel_mpfp(
|
||||||
)
|
)
|
||||||
return _TimedResult(results, time.time() - start)
|
return _TimedResult(results, time.time() - start)
|
||||||
|
|
||||||
# Run parallel task chains
|
# Run all methods in parallel (no chain dependencies)
|
||||||
if temporal_constraint:
|
if temporal_constraint:
|
||||||
tc_start, tc_end = temporal_constraint
|
tc_start, tc_end = temporal_constraint
|
||||||
sg_result, bm25_result, temporal_result = await asyncio.gather(
|
semantic_result, bm25_result, graph_result, temporal_result = await asyncio.gather(
|
||||||
run_semantic_then_graph(),
|
run_semantic(),
|
||||||
run_bm25(),
|
run_bm25(),
|
||||||
|
run_graph(),
|
||||||
run_temporal(tc_start, tc_end),
|
run_temporal(tc_start, tc_end),
|
||||||
)
|
)
|
||||||
|
graph_results, graph_time, mpfp_timing = graph_result
|
||||||
return ParallelRetrievalResult(
|
return ParallelRetrievalResult(
|
||||||
semantic=sg_result.semantic,
|
semantic=semantic_result.results,
|
||||||
bm25=bm25_result.results,
|
bm25=bm25_result.results,
|
||||||
graph=sg_result.graph,
|
graph=graph_results,
|
||||||
temporal=temporal_result.results,
|
temporal=temporal_result.results,
|
||||||
timings={
|
timings={
|
||||||
"semantic": sg_result.semantic_time,
|
"semantic": semantic_result.time,
|
||||||
"graph": sg_result.graph_time,
|
|
||||||
"bm25": bm25_result.time,
|
"bm25": bm25_result.time,
|
||||||
|
"graph": graph_time,
|
||||||
"temporal": temporal_result.time,
|
"temporal": temporal_result.time,
|
||||||
},
|
},
|
||||||
temporal_constraint=temporal_constraint,
|
temporal_constraint=temporal_constraint,
|
||||||
|
mpfp_timings=[mpfp_timing] if mpfp_timing else [],
|
||||||
)
|
)
|
||||||
else:
|
else:
|
||||||
sg_result, bm25_result = await asyncio.gather(
|
semantic_result, bm25_result, graph_result = await asyncio.gather(
|
||||||
run_semantic_then_graph(),
|
run_semantic(),
|
||||||
run_bm25(),
|
run_bm25(),
|
||||||
|
run_graph(),
|
||||||
)
|
)
|
||||||
|
graph_results, graph_time, mpfp_timing = graph_result
|
||||||
return ParallelRetrievalResult(
|
return ParallelRetrievalResult(
|
||||||
semantic=sg_result.semantic,
|
semantic=semantic_result.results,
|
||||||
bm25=bm25_result.results,
|
bm25=bm25_result.results,
|
||||||
graph=sg_result.graph,
|
graph=graph_results,
|
||||||
temporal=None,
|
temporal=None,
|
||||||
timings={
|
timings={
|
||||||
"semantic": sg_result.semantic_time,
|
"semantic": semantic_result.time,
|
||||||
"graph": sg_result.graph_time,
|
|
||||||
"bm25": bm25_result.time,
|
"bm25": bm25_result.time,
|
||||||
|
"graph": graph_time,
|
||||||
},
|
},
|
||||||
temporal_constraint=None,
|
temporal_constraint=None,
|
||||||
|
mpfp_timings=[mpfp_timing] if mpfp_timing else [],
|
||||||
)
|
)
|
||||||
|
|
||||||
|
|
||||||
|
|
@ -633,7 +654,7 @@ async def _retrieve_parallel_bfs(
|
||||||
|
|
||||||
async def run_graph() -> _TimedResult:
|
async def run_graph() -> _TimedResult:
|
||||||
start = time.time()
|
start = time.time()
|
||||||
results = await retriever.retrieve(
|
results, _ = await retriever.retrieve(
|
||||||
pool=pool,
|
pool=pool,
|
||||||
query_embedding_str=query_embedding_str,
|
query_embedding_str=query_embedding_str,
|
||||||
bank_id=bank_id,
|
bank_id=bank_id,
|
||||||
|
|
|
||||||
|
|
@ -10,6 +10,21 @@ from datetime import datetime
|
||||||
from typing import Any
|
from typing import Any
|
||||||
|
|
||||||
|
|
||||||
|
@dataclass
|
||||||
|
class MPFPTimings:
|
||||||
|
"""Timing breakdown for a single MPFP retrieval call."""
|
||||||
|
|
||||||
|
fact_type: str
|
||||||
|
adjacency_query: float = 0.0
|
||||||
|
adjacency_process: float = 0.0
|
||||||
|
edge_count: int = 0
|
||||||
|
traverse: float = 0.0
|
||||||
|
pattern_count: int = 0
|
||||||
|
fusion: float = 0.0
|
||||||
|
fetch: float = 0.0
|
||||||
|
result_count: int = 0
|
||||||
|
|
||||||
|
|
||||||
@dataclass
|
@dataclass
|
||||||
class RetrievalResult:
|
class RetrievalResult:
|
||||||
"""
|
"""
|
||||||
|
|
|
||||||
|
|
@ -121,6 +121,29 @@ class SyncTaskBackend(TaskBackend):
|
||||||
logger.debug("SyncTaskBackend shutdown")
|
logger.debug("SyncTaskBackend shutdown")
|
||||||
|
|
||||||
|
|
||||||
|
class NoopTaskBackend(TaskBackend):
|
||||||
|
"""
|
||||||
|
No-op task backend that discards all tasks.
|
||||||
|
|
||||||
|
This is useful for tests where background task execution is not needed
|
||||||
|
and would only slow down the test suite.
|
||||||
|
"""
|
||||||
|
|
||||||
|
async def initialize(self):
|
||||||
|
"""No-op."""
|
||||||
|
self._initialized = True
|
||||||
|
logger.debug("NoopTaskBackend initialized")
|
||||||
|
|
||||||
|
async def submit_task(self, task_dict: dict[str, Any]):
|
||||||
|
"""Discard the task (do nothing)."""
|
||||||
|
pass
|
||||||
|
|
||||||
|
async def shutdown(self):
|
||||||
|
"""No-op."""
|
||||||
|
self._initialized = False
|
||||||
|
logger.debug("NoopTaskBackend shutdown")
|
||||||
|
|
||||||
|
|
||||||
class AsyncIOQueueBackend(TaskBackend):
|
class AsyncIOQueueBackend(TaskBackend):
|
||||||
"""
|
"""
|
||||||
Task backend implementation using asyncio queues.
|
Task backend implementation using asyncio queues.
|
||||||
|
|
|
||||||
|
|
@ -198,6 +198,7 @@ def main():
|
||||||
retain_chunk_size=config.retain_chunk_size,
|
retain_chunk_size=config.retain_chunk_size,
|
||||||
retain_extract_causal_links=config.retain_extract_causal_links,
|
retain_extract_causal_links=config.retain_extract_causal_links,
|
||||||
retain_extraction_mode=config.retain_extraction_mode,
|
retain_extraction_mode=config.retain_extraction_mode,
|
||||||
|
retain_observations_async=config.retain_observations_async,
|
||||||
skip_llm_verification=config.skip_llm_verification,
|
skip_llm_verification=config.skip_llm_verification,
|
||||||
lazy_reranker=config.lazy_reranker,
|
lazy_reranker=config.lazy_reranker,
|
||||||
run_migrations_on_startup=config.run_migrations_on_startup,
|
run_migrations_on_startup=config.run_migrations_on_startup,
|
||||||
|
|
@ -205,8 +206,9 @@ def main():
|
||||||
db_pool_max_size=config.db_pool_max_size,
|
db_pool_max_size=config.db_pool_max_size,
|
||||||
db_command_timeout=config.db_command_timeout,
|
db_command_timeout=config.db_command_timeout,
|
||||||
db_acquire_timeout=config.db_acquire_timeout,
|
db_acquire_timeout=config.db_acquire_timeout,
|
||||||
task_batch_size=config.task_batch_size,
|
task_backend=config.task_backend,
|
||||||
task_batch_interval=config.task_batch_interval,
|
task_backend_memory_batch_size=config.task_backend_memory_batch_size,
|
||||||
|
task_backend_memory_batch_interval=config.task_backend_memory_batch_interval,
|
||||||
)
|
)
|
||||||
config.configure_logging()
|
config.configure_logging()
|
||||||
if not args.daemon:
|
if not args.daemon:
|
||||||
|
|
|
||||||
|
|
@ -6,11 +6,18 @@ This module provides metrics for:
|
||||||
- Token usage (input/output) per operation
|
- Token usage (input/output) per operation
|
||||||
- Per-bank granularity via labels
|
- Per-bank granularity via labels
|
||||||
- LLM call latency and token usage with scope dimension
|
- LLM call latency and token usage with scope dimension
|
||||||
|
- HTTP request metrics (latency, count by endpoint/method/status)
|
||||||
|
- Process metrics (CPU, memory, file descriptors, threads)
|
||||||
|
- Database connection pool metrics
|
||||||
"""
|
"""
|
||||||
|
|
||||||
import logging
|
import logging
|
||||||
|
import os
|
||||||
|
import resource
|
||||||
|
import threading
|
||||||
import time
|
import time
|
||||||
from contextlib import contextmanager
|
from contextlib import contextmanager
|
||||||
|
from typing import TYPE_CHECKING, Callable
|
||||||
|
|
||||||
from opentelemetry import metrics
|
from opentelemetry import metrics
|
||||||
from opentelemetry.exporter.prometheus import PrometheusMetricReader
|
from opentelemetry.exporter.prometheus import PrometheusMetricReader
|
||||||
|
|
@ -18,6 +25,9 @@ from opentelemetry.sdk.metrics import MeterProvider
|
||||||
from opentelemetry.sdk.metrics.view import ExplicitBucketHistogramAggregation, View
|
from opentelemetry.sdk.metrics.view import ExplicitBucketHistogramAggregation, View
|
||||||
from opentelemetry.sdk.resources import Resource
|
from opentelemetry.sdk.resources import Resource
|
||||||
|
|
||||||
|
if TYPE_CHECKING:
|
||||||
|
import asyncpg
|
||||||
|
|
||||||
# Custom bucket boundaries for operation duration (in seconds)
|
# Custom bucket boundaries for operation duration (in seconds)
|
||||||
# Fine granularity in 0-30s range where most operations complete
|
# Fine granularity in 0-30s range where most operations complete
|
||||||
DURATION_BUCKETS = (0.1, 0.25, 0.5, 0.75, 1.0, 2.0, 3.0, 5.0, 7.5, 10.0, 15.0, 20.0, 30.0, 60.0, 120.0)
|
DURATION_BUCKETS = (0.1, 0.25, 0.5, 0.75, 1.0, 2.0, 3.0, 5.0, 7.5, 10.0, 15.0, 20.0, 30.0, 60.0, 120.0)
|
||||||
|
|
@ -25,6 +35,9 @@ DURATION_BUCKETS = (0.1, 0.25, 0.5, 0.75, 1.0, 2.0, 3.0, 5.0, 7.5, 10.0, 15.0, 2
|
||||||
# LLM duration buckets (finer granularity for faster LLM calls)
|
# LLM duration buckets (finer granularity for faster LLM calls)
|
||||||
LLM_DURATION_BUCKETS = (0.1, 0.25, 0.5, 1.0, 2.0, 3.0, 5.0, 10.0, 15.0, 30.0, 60.0, 120.0)
|
LLM_DURATION_BUCKETS = (0.1, 0.25, 0.5, 1.0, 2.0, 3.0, 5.0, 10.0, 15.0, 30.0, 60.0, 120.0)
|
||||||
|
|
||||||
|
# HTTP request duration buckets (millisecond-level for fast endpoints)
|
||||||
|
HTTP_DURATION_BUCKETS = (0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0, 30.0)
|
||||||
|
|
||||||
|
|
||||||
def get_token_bucket(token_count: int) -> str:
|
def get_token_bucket(token_count: int) -> str:
|
||||||
"""
|
"""
|
||||||
|
|
@ -107,9 +120,17 @@ def initialize_metrics(service_name: str = "hindsight-api", service_version: str
|
||||||
aggregation=ExplicitBucketHistogramAggregation(boundaries=LLM_DURATION_BUCKETS),
|
aggregation=ExplicitBucketHistogramAggregation(boundaries=LLM_DURATION_BUCKETS),
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# Create view with custom bucket boundaries for HTTP request duration histogram
|
||||||
|
http_duration_view = View(
|
||||||
|
instrument_name="hindsight.http.duration",
|
||||||
|
aggregation=ExplicitBucketHistogramAggregation(boundaries=HTTP_DURATION_BUCKETS),
|
||||||
|
)
|
||||||
|
|
||||||
# Create meter provider with Prometheus exporter and custom views
|
# Create meter provider with Prometheus exporter and custom views
|
||||||
provider = MeterProvider(
|
provider = MeterProvider(
|
||||||
resource=resource, metric_readers=[prometheus_reader], views=[duration_view, llm_duration_view]
|
resource=resource,
|
||||||
|
metric_readers=[prometheus_reader],
|
||||||
|
views=[duration_view, llm_duration_view, http_duration_view],
|
||||||
)
|
)
|
||||||
|
|
||||||
# Set the global meter provider
|
# Set the global meter provider
|
||||||
|
|
@ -167,6 +188,15 @@ class MetricsCollectorBase:
|
||||||
"""
|
"""
|
||||||
raise NotImplementedError
|
raise NotImplementedError
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def record_http_request(self, method: str, endpoint: str, status_code_getter: Callable[[], int]):
|
||||||
|
"""Context manager to record HTTP request metrics."""
|
||||||
|
raise NotImplementedError
|
||||||
|
|
||||||
|
def set_db_pool(self, pool: "asyncpg.Pool"):
|
||||||
|
"""Set the database pool for metrics collection."""
|
||||||
|
pass
|
||||||
|
|
||||||
|
|
||||||
class NoOpMetricsCollector(MetricsCollectorBase):
|
class NoOpMetricsCollector(MetricsCollectorBase):
|
||||||
"""No-op metrics collector that does nothing. Used when metrics are disabled."""
|
"""No-op metrics collector that does nothing. Used when metrics are disabled."""
|
||||||
|
|
@ -196,6 +226,11 @@ class NoOpMetricsCollector(MetricsCollectorBase):
|
||||||
"""No-op LLM call recording."""
|
"""No-op LLM call recording."""
|
||||||
pass
|
pass
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def record_http_request(self, method: str, endpoint: str, status_code_getter: Callable[[], int]):
|
||||||
|
"""No-op HTTP request recording."""
|
||||||
|
yield
|
||||||
|
|
||||||
|
|
||||||
class MetricsCollector(MetricsCollectorBase):
|
class MetricsCollector(MetricsCollectorBase):
|
||||||
"""
|
"""
|
||||||
|
|
@ -238,6 +273,27 @@ class MetricsCollector(MetricsCollectorBase):
|
||||||
name="hindsight.llm.calls.total", description="Total number of LLM API calls", unit="calls"
|
name="hindsight.llm.calls.total", description="Total number of LLM API calls", unit="calls"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
# HTTP request metrics
|
||||||
|
self.http_request_duration = self.meter.create_histogram(
|
||||||
|
name="hindsight.http.duration", description="Duration of HTTP requests in seconds", unit="s"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.http_requests_total = self.meter.create_counter(
|
||||||
|
name="hindsight.http.requests.total", description="Total number of HTTP requests", unit="requests"
|
||||||
|
)
|
||||||
|
|
||||||
|
self.http_requests_in_progress = self.meter.create_up_down_counter(
|
||||||
|
name="hindsight.http.requests.in_progress",
|
||||||
|
description="Number of HTTP requests in progress",
|
||||||
|
unit="requests",
|
||||||
|
)
|
||||||
|
|
||||||
|
# Process metrics (observable gauges - collected on scrape)
|
||||||
|
self._setup_process_metrics()
|
||||||
|
|
||||||
|
# DB pool metrics holder (set via set_db_pool)
|
||||||
|
self._db_pool: "asyncpg.Pool | None" = None
|
||||||
|
|
||||||
@contextmanager
|
@contextmanager
|
||||||
def record_operation(
|
def record_operation(
|
||||||
self,
|
self,
|
||||||
|
|
@ -340,6 +396,196 @@ class MetricsCollector(MetricsCollectorBase):
|
||||||
}
|
}
|
||||||
self.llm_tokens_output.add(output_tokens, output_attributes)
|
self.llm_tokens_output.add(output_tokens, output_attributes)
|
||||||
|
|
||||||
|
@contextmanager
|
||||||
|
def record_http_request(self, method: str, endpoint: str, status_code_getter: Callable[[], int]):
|
||||||
|
"""
|
||||||
|
Context manager to record HTTP request metrics.
|
||||||
|
|
||||||
|
Usage:
|
||||||
|
status_code = [200] # Use list for mutability
|
||||||
|
with metrics.record_http_request("GET", "/api/banks", lambda: status_code[0]):
|
||||||
|
# ... handle request
|
||||||
|
status_code[0] = response.status_code
|
||||||
|
|
||||||
|
Args:
|
||||||
|
method: HTTP method (GET, POST, etc.)
|
||||||
|
endpoint: Request endpoint path
|
||||||
|
status_code_getter: Callable that returns the status code after request completes
|
||||||
|
"""
|
||||||
|
start_time = time.time()
|
||||||
|
base_attributes = {"method": method, "endpoint": endpoint}
|
||||||
|
|
||||||
|
# Track in-progress
|
||||||
|
self.http_requests_in_progress.add(1, base_attributes)
|
||||||
|
|
||||||
|
try:
|
||||||
|
yield
|
||||||
|
finally:
|
||||||
|
duration = time.time() - start_time
|
||||||
|
status_code = status_code_getter()
|
||||||
|
status_class = f"{status_code // 100}xx"
|
||||||
|
|
||||||
|
attributes = {
|
||||||
|
**base_attributes,
|
||||||
|
"status_code": str(status_code),
|
||||||
|
"status_class": status_class,
|
||||||
|
}
|
||||||
|
|
||||||
|
# Record duration and count
|
||||||
|
self.http_request_duration.record(duration, attributes)
|
||||||
|
self.http_requests_total.add(1, attributes)
|
||||||
|
|
||||||
|
# Decrement in-progress
|
||||||
|
self.http_requests_in_progress.add(-1, base_attributes)
|
||||||
|
|
||||||
|
def _setup_process_metrics(self):
|
||||||
|
"""Set up observable gauges for process metrics."""
|
||||||
|
|
||||||
|
def get_cpu_times(_options):
|
||||||
|
"""Get process CPU times."""
|
||||||
|
try:
|
||||||
|
rusage = resource.getrusage(resource.RUSAGE_SELF)
|
||||||
|
yield metrics.Observation(rusage.ru_utime, {"type": "user"})
|
||||||
|
yield metrics.Observation(rusage.ru_stime, {"type": "system"})
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def get_memory_usage(_options):
|
||||||
|
"""Get process memory usage in bytes."""
|
||||||
|
try:
|
||||||
|
rusage = resource.getrusage(resource.RUSAGE_SELF)
|
||||||
|
# ru_maxrss is in kilobytes on Linux, bytes on macOS
|
||||||
|
max_rss = rusage.ru_maxrss
|
||||||
|
if os.uname().sysname == "Linux":
|
||||||
|
max_rss *= 1024 # Convert KB to bytes
|
||||||
|
yield metrics.Observation(max_rss, {"type": "rss_max"})
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def get_open_file_descriptors(_options):
|
||||||
|
"""Get number of open file descriptors."""
|
||||||
|
try:
|
||||||
|
# Try to count open FDs by checking /proc on Linux
|
||||||
|
if os.path.exists("/proc/self/fd"):
|
||||||
|
count = len(os.listdir("/proc/self/fd"))
|
||||||
|
yield metrics.Observation(count)
|
||||||
|
else:
|
||||||
|
# Fallback: use resource limits
|
||||||
|
soft, hard = resource.getrlimit(resource.RLIMIT_NOFILE)
|
||||||
|
yield metrics.Observation(soft, {"limit": "soft"})
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def get_thread_count(_options):
|
||||||
|
"""Get number of active threads."""
|
||||||
|
try:
|
||||||
|
yield metrics.Observation(threading.active_count())
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Create observable gauges
|
||||||
|
self.meter.create_observable_gauge(
|
||||||
|
name="hindsight.process.cpu.seconds",
|
||||||
|
callbacks=[get_cpu_times],
|
||||||
|
description="Process CPU time in seconds",
|
||||||
|
unit="s",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.meter.create_observable_gauge(
|
||||||
|
name="hindsight.process.memory.bytes",
|
||||||
|
callbacks=[get_memory_usage],
|
||||||
|
description="Process memory usage in bytes",
|
||||||
|
unit="By",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.meter.create_observable_gauge(
|
||||||
|
name="hindsight.process.open_fds",
|
||||||
|
callbacks=[get_open_file_descriptors],
|
||||||
|
description="Number of open file descriptors",
|
||||||
|
unit="{fds}",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.meter.create_observable_gauge(
|
||||||
|
name="hindsight.process.threads",
|
||||||
|
callbacks=[get_thread_count],
|
||||||
|
description="Number of active threads",
|
||||||
|
unit="{threads}",
|
||||||
|
)
|
||||||
|
|
||||||
|
def set_db_pool(self, pool: "asyncpg.Pool"):
|
||||||
|
"""
|
||||||
|
Set the database pool for metrics collection.
|
||||||
|
|
||||||
|
Args:
|
||||||
|
pool: asyncpg connection pool instance
|
||||||
|
"""
|
||||||
|
self._db_pool = pool
|
||||||
|
self._setup_db_pool_metrics()
|
||||||
|
|
||||||
|
def _setup_db_pool_metrics(self):
|
||||||
|
"""Set up observable gauges for database pool metrics."""
|
||||||
|
|
||||||
|
def get_pool_size(_options):
|
||||||
|
"""Get current pool size."""
|
||||||
|
if self._db_pool is not None:
|
||||||
|
try:
|
||||||
|
yield metrics.Observation(self._db_pool.get_size())
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def get_pool_free_size(_options):
|
||||||
|
"""Get number of free connections in pool."""
|
||||||
|
if self._db_pool is not None:
|
||||||
|
try:
|
||||||
|
yield metrics.Observation(self._db_pool.get_idle_size())
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def get_pool_min_size(_options):
|
||||||
|
"""Get pool minimum size."""
|
||||||
|
if self._db_pool is not None:
|
||||||
|
try:
|
||||||
|
yield metrics.Observation(self._db_pool.get_min_size())
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
def get_pool_max_size(_options):
|
||||||
|
"""Get pool maximum size."""
|
||||||
|
if self._db_pool is not None:
|
||||||
|
try:
|
||||||
|
yield metrics.Observation(self._db_pool.get_max_size())
|
||||||
|
except Exception:
|
||||||
|
pass
|
||||||
|
|
||||||
|
# Create observable gauges for pool metrics
|
||||||
|
self.meter.create_observable_gauge(
|
||||||
|
name="hindsight.db.pool.size",
|
||||||
|
callbacks=[get_pool_size],
|
||||||
|
description="Current number of connections in the pool",
|
||||||
|
unit="{connections}",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.meter.create_observable_gauge(
|
||||||
|
name="hindsight.db.pool.idle",
|
||||||
|
callbacks=[get_pool_free_size],
|
||||||
|
description="Number of idle connections in the pool",
|
||||||
|
unit="{connections}",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.meter.create_observable_gauge(
|
||||||
|
name="hindsight.db.pool.min",
|
||||||
|
callbacks=[get_pool_min_size],
|
||||||
|
description="Minimum pool size",
|
||||||
|
unit="{connections}",
|
||||||
|
)
|
||||||
|
|
||||||
|
self.meter.create_observable_gauge(
|
||||||
|
name="hindsight.db.pool.max",
|
||||||
|
callbacks=[get_pool_max_size],
|
||||||
|
description="Maximum pool size",
|
||||||
|
unit="{connections}",
|
||||||
|
)
|
||||||
|
|
||||||
|
|
||||||
# Global metrics collector instance (defaults to no-op)
|
# Global metrics collector instance (defaults to no-op)
|
||||||
_metrics_collector: MetricsCollectorBase = NoOpMetricsCollector()
|
_metrics_collector: MetricsCollectorBase = NoOpMetricsCollector()
|
||||||
|
|
|
||||||
|
|
@ -64,12 +64,12 @@ class TestMetricsCollector:
|
||||||
def mock_meter(self):
|
def mock_meter(self):
|
||||||
"""Create a mock meter for testing."""
|
"""Create a mock meter for testing."""
|
||||||
meter = MagicMock()
|
meter = MagicMock()
|
||||||
# Create separate mocks for each histogram (operation_duration, llm_duration)
|
# Create separate mocks for each histogram (operation_duration, llm_duration, http_request_duration)
|
||||||
histogram_mocks = [MagicMock(), MagicMock()]
|
histogram_mocks = [MagicMock(), MagicMock(), MagicMock()]
|
||||||
meter.create_histogram.side_effect = histogram_mocks
|
meter.create_histogram.side_effect = histogram_mocks
|
||||||
# Create separate mocks for each counter
|
# Create separate mocks for each counter
|
||||||
# (operation_total, llm_tokens_input, llm_tokens_output, llm_calls_total)
|
# (operation_total, llm_tokens_input, llm_tokens_output, llm_calls_total, http_requests_total)
|
||||||
counter_mocks = [MagicMock() for _ in range(4)]
|
counter_mocks = [MagicMock() for _ in range(5)]
|
||||||
meter.create_counter.side_effect = counter_mocks
|
meter.create_counter.side_effect = counter_mocks
|
||||||
return meter
|
return meter
|
||||||
|
|
||||||
|
|
@ -257,12 +257,12 @@ class TestLLMMetrics:
|
||||||
def mock_meter(self):
|
def mock_meter(self):
|
||||||
"""Create a mock meter for testing."""
|
"""Create a mock meter for testing."""
|
||||||
meter = MagicMock()
|
meter = MagicMock()
|
||||||
# Create separate mocks for each histogram (operation_duration, llm_duration)
|
# Create separate mocks for each histogram (operation_duration, llm_duration, http_request_duration)
|
||||||
histogram_mocks = [MagicMock(), MagicMock()]
|
histogram_mocks = [MagicMock(), MagicMock(), MagicMock()]
|
||||||
meter.create_histogram.side_effect = histogram_mocks
|
meter.create_histogram.side_effect = histogram_mocks
|
||||||
# Create separate mocks for each counter
|
# Create separate mocks for each counter
|
||||||
# (operation_total, llm_tokens_input, llm_tokens_output, llm_calls_total)
|
# (operation_total, llm_tokens_input, llm_tokens_output, llm_calls_total, http_requests_total)
|
||||||
counter_mocks = [MagicMock() for _ in range(4)]
|
counter_mocks = [MagicMock() for _ in range(5)]
|
||||||
meter.create_counter.side_effect = counter_mocks
|
meter.create_counter.side_effect = counter_mocks
|
||||||
return meter
|
return meter
|
||||||
|
|
||||||
|
|
|
||||||
553
hindsight-api/tests/test_mpfp_retrieval.py
Normal file
553
hindsight-api/tests/test_mpfp_retrieval.py
Normal file
|
|
@ -0,0 +1,553 @@
|
||||||
|
"""
|
||||||
|
Tests for MPFP (Meta-Path Forward Push) graph retrieval.
|
||||||
|
|
||||||
|
Tests cover:
|
||||||
|
1. EdgeCache - lazy caching behavior
|
||||||
|
2. mpfp_traverse_async - core traversal algorithm
|
||||||
|
3. load_edges_for_frontier - lazy edge loading
|
||||||
|
4. rrf_fusion - result fusion
|
||||||
|
5. MPFPGraphRetriever - full integration
|
||||||
|
"""
|
||||||
|
|
||||||
|
import pytest
|
||||||
|
from unittest.mock import AsyncMock, MagicMock, patch
|
||||||
|
from datetime import datetime, timezone
|
||||||
|
|
||||||
|
from hindsight_api.engine.search.mpfp_retrieval import (
|
||||||
|
EdgeCache,
|
||||||
|
EdgeTarget,
|
||||||
|
MPFPConfig,
|
||||||
|
MPFPGraphRetriever,
|
||||||
|
PatternResult,
|
||||||
|
SeedNode,
|
||||||
|
load_edges_for_frontier,
|
||||||
|
mpfp_traverse_async,
|
||||||
|
rrf_fusion,
|
||||||
|
)
|
||||||
|
from hindsight_api.engine.search.types import RetrievalResult
|
||||||
|
|
||||||
|
|
||||||
|
class TestEdgeCache:
|
||||||
|
"""Tests for the EdgeCache lazy loading cache."""
|
||||||
|
|
||||||
|
def test_empty_cache_returns_empty_neighbors(self):
|
||||||
|
"""Empty cache should return empty list for any node."""
|
||||||
|
cache = EdgeCache()
|
||||||
|
neighbors = cache.get_neighbors("semantic", "node-1")
|
||||||
|
assert neighbors == []
|
||||||
|
|
||||||
|
def test_is_loaded_false_for_uncached(self):
|
||||||
|
"""is_loaded should return False for nodes not yet loaded."""
|
||||||
|
cache = EdgeCache()
|
||||||
|
assert cache.is_loaded("semantic", "node-1") is False
|
||||||
|
|
||||||
|
def test_add_edges_marks_as_loaded(self):
|
||||||
|
"""Adding edges should mark nodes as loaded."""
|
||||||
|
cache = EdgeCache()
|
||||||
|
|
||||||
|
edges = {
|
||||||
|
"node-1": [EdgeTarget("node-2", 0.8), EdgeTarget("node-3", 0.6)],
|
||||||
|
}
|
||||||
|
cache.add_edges("semantic", edges, ["node-1", "node-4"]) # node-4 has no edges
|
||||||
|
|
||||||
|
assert cache.is_loaded("semantic", "node-1") is True
|
||||||
|
assert cache.is_loaded("semantic", "node-4") is True # Marked even with no edges
|
||||||
|
assert cache.is_loaded("semantic", "node-2") is False # Target, not source
|
||||||
|
|
||||||
|
def test_get_neighbors_returns_added_edges(self):
|
||||||
|
"""get_neighbors should return edges after add_edges."""
|
||||||
|
cache = EdgeCache()
|
||||||
|
|
||||||
|
edges = {
|
||||||
|
"node-1": [EdgeTarget("node-2", 0.8), EdgeTarget("node-3", 0.6)],
|
||||||
|
}
|
||||||
|
cache.add_edges("semantic", edges, ["node-1"])
|
||||||
|
|
||||||
|
neighbors = cache.get_neighbors("semantic", "node-1")
|
||||||
|
assert len(neighbors) == 2
|
||||||
|
assert neighbors[0].node_id == "node-2"
|
||||||
|
assert neighbors[0].weight == 0.8
|
||||||
|
|
||||||
|
def test_get_uncached_filters_loaded_nodes(self):
|
||||||
|
"""get_uncached should only return nodes not yet loaded."""
|
||||||
|
cache = EdgeCache()
|
||||||
|
|
||||||
|
# Load some nodes
|
||||||
|
cache.add_edges("semantic", {"node-1": []}, ["node-1", "node-2"])
|
||||||
|
|
||||||
|
# Check uncached
|
||||||
|
uncached = cache.get_uncached("semantic", ["node-1", "node-2", "node-3", "node-4"])
|
||||||
|
assert set(uncached) == {"node-3", "node-4"}
|
||||||
|
|
||||||
|
def test_get_normalized_neighbors_normalizes_weights(self):
|
||||||
|
"""get_normalized_neighbors should normalize weights to sum to 1."""
|
||||||
|
cache = EdgeCache()
|
||||||
|
|
||||||
|
edges = {
|
||||||
|
"node-1": [
|
||||||
|
EdgeTarget("node-2", 0.8),
|
||||||
|
EdgeTarget("node-3", 0.4),
|
||||||
|
EdgeTarget("node-4", 0.2),
|
||||||
|
],
|
||||||
|
}
|
||||||
|
cache.add_edges("semantic", edges, ["node-1"])
|
||||||
|
|
||||||
|
# Get top 2, normalized
|
||||||
|
neighbors = cache.get_normalized_neighbors("semantic", "node-1", top_k=2)
|
||||||
|
assert len(neighbors) == 2
|
||||||
|
|
||||||
|
# Weights should sum to 1
|
||||||
|
total = sum(n.weight for n in neighbors)
|
||||||
|
assert abs(total - 1.0) < 0.001
|
||||||
|
|
||||||
|
# node-2 should have higher normalized weight than node-3
|
||||||
|
assert neighbors[0].node_id == "node-2"
|
||||||
|
assert neighbors[1].node_id == "node-3"
|
||||||
|
# Original: 0.8 and 0.4, so normalized: 0.8/1.2 and 0.4/1.2
|
||||||
|
assert abs(neighbors[0].weight - 0.8 / 1.2) < 0.001
|
||||||
|
assert abs(neighbors[1].weight - 0.4 / 1.2) < 0.001
|
||||||
|
|
||||||
|
def test_different_edge_types_are_separate(self):
|
||||||
|
"""Different edge types should be stored separately."""
|
||||||
|
cache = EdgeCache()
|
||||||
|
|
||||||
|
cache.add_edges("semantic", {"node-1": [EdgeTarget("node-2", 0.8)]}, ["node-1"])
|
||||||
|
cache.add_edges("temporal", {"node-1": [EdgeTarget("node-3", 0.5)]}, ["node-1"])
|
||||||
|
|
||||||
|
semantic_neighbors = cache.get_neighbors("semantic", "node-1")
|
||||||
|
temporal_neighbors = cache.get_neighbors("temporal", "node-1")
|
||||||
|
|
||||||
|
assert len(semantic_neighbors) == 1
|
||||||
|
assert semantic_neighbors[0].node_id == "node-2"
|
||||||
|
|
||||||
|
assert len(temporal_neighbors) == 1
|
||||||
|
assert temporal_neighbors[0].node_id == "node-3"
|
||||||
|
|
||||||
|
|
||||||
|
class TestRRFFusion:
|
||||||
|
"""Tests for RRF (Reciprocal Rank Fusion)."""
|
||||||
|
|
||||||
|
def test_empty_results(self):
|
||||||
|
"""Empty results should return empty fusion."""
|
||||||
|
fused = rrf_fusion([])
|
||||||
|
assert fused == []
|
||||||
|
|
||||||
|
def test_single_pattern_ranking(self):
|
||||||
|
"""Single pattern should preserve ranking order."""
|
||||||
|
result = PatternResult(
|
||||||
|
pattern=["semantic"],
|
||||||
|
scores={"node-1": 0.9, "node-2": 0.7, "node-3": 0.5},
|
||||||
|
)
|
||||||
|
|
||||||
|
fused = rrf_fusion([result], top_k=3)
|
||||||
|
assert len(fused) == 3
|
||||||
|
# node-1 should be first (highest score)
|
||||||
|
assert fused[0][0] == "node-1"
|
||||||
|
assert fused[1][0] == "node-2"
|
||||||
|
assert fused[2][0] == "node-3"
|
||||||
|
|
||||||
|
def test_multiple_patterns_boost_common_nodes(self):
|
||||||
|
"""Nodes appearing in multiple patterns should get boosted."""
|
||||||
|
result1 = PatternResult(
|
||||||
|
pattern=["semantic", "semantic"],
|
||||||
|
scores={"node-1": 0.9, "node-2": 0.7},
|
||||||
|
)
|
||||||
|
result2 = PatternResult(
|
||||||
|
pattern=["entity", "temporal"],
|
||||||
|
scores={"node-1": 0.8, "node-3": 0.6}, # node-1 in both
|
||||||
|
)
|
||||||
|
|
||||||
|
fused = rrf_fusion([result1, result2], top_k=3)
|
||||||
|
|
||||||
|
# node-1 should be first (appears in both patterns)
|
||||||
|
assert fused[0][0] == "node-1"
|
||||||
|
# Its score should be higher than others
|
||||||
|
assert fused[0][1] > fused[1][1]
|
||||||
|
|
||||||
|
def test_top_k_limits_results(self):
|
||||||
|
"""top_k should limit the number of results."""
|
||||||
|
result = PatternResult(
|
||||||
|
pattern=["semantic"],
|
||||||
|
scores={f"node-{i}": 1.0 / (i + 1) for i in range(10)},
|
||||||
|
)
|
||||||
|
|
||||||
|
fused = rrf_fusion([result], top_k=3)
|
||||||
|
assert len(fused) == 3
|
||||||
|
|
||||||
|
def test_empty_pattern_scores_ignored(self):
|
||||||
|
"""Patterns with empty scores should be ignored."""
|
||||||
|
result1 = PatternResult(pattern=["semantic"], scores={})
|
||||||
|
result2 = PatternResult(
|
||||||
|
pattern=["entity"],
|
||||||
|
scores={"node-1": 0.5},
|
||||||
|
)
|
||||||
|
|
||||||
|
fused = rrf_fusion([result1, result2], top_k=3)
|
||||||
|
assert len(fused) == 1
|
||||||
|
assert fused[0][0] == "node-1"
|
||||||
|
|
||||||
|
|
||||||
|
class TestMPFPTraverseAsync:
|
||||||
|
"""Tests for the async MPFP traversal algorithm."""
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_empty_seeds_returns_empty(self):
|
||||||
|
"""Empty seeds should return empty result."""
|
||||||
|
cache = EdgeCache()
|
||||||
|
config = MPFPConfig()
|
||||||
|
|
||||||
|
result = await mpfp_traverse_async(
|
||||||
|
pool=None, # Not used when no seeds
|
||||||
|
bank_id="test",
|
||||||
|
seeds=[],
|
||||||
|
pattern=["semantic"],
|
||||||
|
config=config,
|
||||||
|
cache=cache,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.scores == {}
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_single_hop_no_edges(self):
|
||||||
|
"""Single hop with no edges should deposit mass at seeds."""
|
||||||
|
cache = EdgeCache()
|
||||||
|
config = MPFPConfig(alpha=0.15, threshold=1e-6)
|
||||||
|
|
||||||
|
# Pre-populate cache with empty edges for seed
|
||||||
|
cache.add_edges("semantic", {}, ["seed-1"])
|
||||||
|
|
||||||
|
seeds = [SeedNode("seed-1", 1.0)]
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"hindsight_api.engine.search.mpfp_retrieval.load_edges_for_frontier",
|
||||||
|
new_callable=AsyncMock,
|
||||||
|
return_value={},
|
||||||
|
):
|
||||||
|
result = await mpfp_traverse_async(
|
||||||
|
pool=MagicMock(),
|
||||||
|
bank_id="test",
|
||||||
|
seeds=seeds,
|
||||||
|
pattern=["semantic"],
|
||||||
|
config=config,
|
||||||
|
cache=cache,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Seed should have alpha portion of its mass
|
||||||
|
assert "seed-1" in result.scores
|
||||||
|
assert result.scores["seed-1"] == pytest.approx(config.alpha, rel=0.01)
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_single_hop_with_edges(self):
|
||||||
|
"""Single hop should spread mass to neighbors."""
|
||||||
|
cache = EdgeCache()
|
||||||
|
config = MPFPConfig(alpha=0.15, threshold=1e-6, top_k_neighbors=10)
|
||||||
|
|
||||||
|
seeds = [SeedNode("seed-1", 1.0)]
|
||||||
|
|
||||||
|
# Mock edge loading
|
||||||
|
async def mock_load_edges(pool, bank_id, edge_type, node_ids):
|
||||||
|
if "seed-1" in node_ids:
|
||||||
|
return {
|
||||||
|
"seed-1": [
|
||||||
|
EdgeTarget("neighbor-1", 0.8),
|
||||||
|
EdgeTarget("neighbor-2", 0.4),
|
||||||
|
]
|
||||||
|
}
|
||||||
|
return {}
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"hindsight_api.engine.search.mpfp_retrieval.load_edges_for_frontier",
|
||||||
|
side_effect=mock_load_edges,
|
||||||
|
):
|
||||||
|
result = await mpfp_traverse_async(
|
||||||
|
pool=MagicMock(),
|
||||||
|
bank_id="test",
|
||||||
|
seeds=seeds,
|
||||||
|
pattern=["semantic"],
|
||||||
|
config=config,
|
||||||
|
cache=cache,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Seed keeps alpha portion
|
||||||
|
assert "seed-1" in result.scores
|
||||||
|
assert result.scores["seed-1"] == pytest.approx(config.alpha, rel=0.01)
|
||||||
|
|
||||||
|
# Neighbors get remaining mass (normalized)
|
||||||
|
assert "neighbor-1" in result.scores
|
||||||
|
assert "neighbor-2" in result.scores
|
||||||
|
|
||||||
|
# neighbor-1 should get more (higher weight)
|
||||||
|
assert result.scores["neighbor-1"] > result.scores["neighbor-2"]
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_two_hops(self):
|
||||||
|
"""Two-hop pattern should traverse through neighbors."""
|
||||||
|
cache = EdgeCache()
|
||||||
|
config = MPFPConfig(alpha=0.15, threshold=1e-6, top_k_neighbors=10)
|
||||||
|
|
||||||
|
seeds = [SeedNode("seed-1", 1.0)]
|
||||||
|
|
||||||
|
# Mock edge loading for two hops
|
||||||
|
async def mock_load_edges(pool, bank_id, edge_type, node_ids):
|
||||||
|
edges = {}
|
||||||
|
if "seed-1" in node_ids:
|
||||||
|
edges["seed-1"] = [EdgeTarget("hop1-node", 1.0)]
|
||||||
|
if "hop1-node" in node_ids:
|
||||||
|
edges["hop1-node"] = [EdgeTarget("hop2-node", 1.0)]
|
||||||
|
return edges
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"hindsight_api.engine.search.mpfp_retrieval.load_edges_for_frontier",
|
||||||
|
side_effect=mock_load_edges,
|
||||||
|
):
|
||||||
|
result = await mpfp_traverse_async(
|
||||||
|
pool=MagicMock(),
|
||||||
|
bank_id="test",
|
||||||
|
seeds=seeds,
|
||||||
|
pattern=["semantic", "semantic"], # Two hops
|
||||||
|
config=config,
|
||||||
|
cache=cache,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should have scores for all three nodes
|
||||||
|
assert "seed-1" in result.scores
|
||||||
|
assert "hop1-node" in result.scores
|
||||||
|
assert "hop2-node" in result.scores
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_cache_reuse(self):
|
||||||
|
"""Cache should prevent redundant edge loading."""
|
||||||
|
cache = EdgeCache()
|
||||||
|
config = MPFPConfig(alpha=0.15, threshold=1e-6)
|
||||||
|
|
||||||
|
# Pre-load cache
|
||||||
|
cache.add_edges("semantic", {"seed-1": [EdgeTarget("neighbor-1", 1.0)]}, ["seed-1"])
|
||||||
|
|
||||||
|
seeds = [SeedNode("seed-1", 1.0)]
|
||||||
|
|
||||||
|
load_mock = AsyncMock(return_value={})
|
||||||
|
|
||||||
|
with patch(
|
||||||
|
"hindsight_api.engine.search.mpfp_retrieval.load_edges_for_frontier",
|
||||||
|
load_mock,
|
||||||
|
):
|
||||||
|
await mpfp_traverse_async(
|
||||||
|
pool=MagicMock(),
|
||||||
|
bank_id="test",
|
||||||
|
seeds=seeds,
|
||||||
|
pattern=["semantic"],
|
||||||
|
config=config,
|
||||||
|
cache=cache,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should not call load_edges_for_frontier since seed-1 is already cached
|
||||||
|
load_mock.assert_not_called()
|
||||||
|
|
||||||
|
|
||||||
|
class TestMPFPGraphRetriever:
|
||||||
|
"""Tests for the MPFPGraphRetriever class."""
|
||||||
|
|
||||||
|
def test_name_is_mpfp(self):
|
||||||
|
"""Retriever name should be 'mpfp'."""
|
||||||
|
retriever = MPFPGraphRetriever()
|
||||||
|
assert retriever.name == "mpfp"
|
||||||
|
|
||||||
|
def test_default_config(self):
|
||||||
|
"""Default config should have expected patterns."""
|
||||||
|
retriever = MPFPGraphRetriever()
|
||||||
|
|
||||||
|
assert len(retriever.config.patterns_semantic) > 0
|
||||||
|
assert len(retriever.config.patterns_temporal) > 0
|
||||||
|
assert retriever.config.alpha == 0.15
|
||||||
|
assert retriever.config.top_k_neighbors == 20
|
||||||
|
|
||||||
|
def test_custom_config(self):
|
||||||
|
"""Custom config should be used."""
|
||||||
|
config = MPFPConfig(alpha=0.3, top_k_neighbors=10)
|
||||||
|
retriever = MPFPGraphRetriever(config=config)
|
||||||
|
|
||||||
|
assert retriever.config.alpha == 0.3
|
||||||
|
assert retriever.config.top_k_neighbors == 10
|
||||||
|
|
||||||
|
def test_convert_seeds_from_retrieval_results(self):
|
||||||
|
"""_convert_seeds should extract scores from RetrievalResult."""
|
||||||
|
retriever = MPFPGraphRetriever()
|
||||||
|
|
||||||
|
results = [
|
||||||
|
RetrievalResult(id="id-1", text="text1", fact_type="world", similarity=0.9),
|
||||||
|
RetrievalResult(id="id-2", text="text2", fact_type="world", similarity=0.7),
|
||||||
|
]
|
||||||
|
|
||||||
|
seeds = retriever._convert_seeds(results, "similarity")
|
||||||
|
|
||||||
|
assert len(seeds) == 2
|
||||||
|
assert seeds[0].node_id == "id-1"
|
||||||
|
assert seeds[0].score == 0.9
|
||||||
|
assert seeds[1].node_id == "id-2"
|
||||||
|
assert seeds[1].score == 0.7
|
||||||
|
|
||||||
|
def test_convert_seeds_empty(self):
|
||||||
|
"""_convert_seeds should handle empty/None input."""
|
||||||
|
retriever = MPFPGraphRetriever()
|
||||||
|
|
||||||
|
assert retriever._convert_seeds(None, "similarity") == []
|
||||||
|
assert retriever._convert_seeds([], "similarity") == []
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_retrieve_no_seeds_returns_empty(self):
|
||||||
|
"""Retrieve with no seeds should return empty results."""
|
||||||
|
retriever = MPFPGraphRetriever()
|
||||||
|
|
||||||
|
# Mock _find_semantic_seeds to return empty
|
||||||
|
with patch.object(retriever, "_find_semantic_seeds", new_callable=AsyncMock, return_value=[]):
|
||||||
|
results, timings = await retriever.retrieve(
|
||||||
|
pool=MagicMock(),
|
||||||
|
query_embedding_str="[0.1, 0.2]",
|
||||||
|
bank_id="test",
|
||||||
|
fact_type="world",
|
||||||
|
budget=10,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert results == []
|
||||||
|
assert timings is not None
|
||||||
|
assert timings.pattern_count == 0
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_retrieve_with_semantic_seeds(self):
|
||||||
|
"""Retrieve with semantic seeds should run patterns and return results."""
|
||||||
|
retriever = MPFPGraphRetriever()
|
||||||
|
|
||||||
|
semantic_seeds = [
|
||||||
|
RetrievalResult(id="seed-1", text="seed text", fact_type="world", similarity=0.9),
|
||||||
|
]
|
||||||
|
|
||||||
|
# Mock the internal functions
|
||||||
|
async def mock_traverse(*args, **kwargs):
|
||||||
|
return PatternResult(pattern=["semantic"], scores={"seed-1": 0.5, "result-1": 0.3})
|
||||||
|
|
||||||
|
async def mock_fetch(pool, node_ids, fact_type):
|
||||||
|
return [
|
||||||
|
RetrievalResult(id="seed-1", text="seed text", fact_type="world"),
|
||||||
|
RetrievalResult(id="result-1", text="result text", fact_type="world"),
|
||||||
|
]
|
||||||
|
|
||||||
|
with (
|
||||||
|
patch(
|
||||||
|
"hindsight_api.engine.search.mpfp_retrieval.mpfp_traverse_async",
|
||||||
|
side_effect=mock_traverse,
|
||||||
|
),
|
||||||
|
patch(
|
||||||
|
"hindsight_api.engine.search.mpfp_retrieval.fetch_memory_units_by_ids",
|
||||||
|
side_effect=mock_fetch,
|
||||||
|
),
|
||||||
|
):
|
||||||
|
results, timings = await retriever.retrieve(
|
||||||
|
pool=MagicMock(),
|
||||||
|
query_embedding_str="[0.1, 0.2]",
|
||||||
|
bank_id="test",
|
||||||
|
fact_type="world",
|
||||||
|
budget=10,
|
||||||
|
semantic_seeds=semantic_seeds,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert len(results) == 2
|
||||||
|
assert timings is not None
|
||||||
|
assert timings.pattern_count > 0
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_mpfp_integration(memory, request_context):
|
||||||
|
"""Integration test: MPFP retrieval with real database."""
|
||||||
|
bank_id = f"test_mpfp_{datetime.now(timezone.utc).timestamp()}"
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Store memories with entity relationships
|
||||||
|
await memory.retain_async(
|
||||||
|
bank_id=bank_id,
|
||||||
|
content="Alice works at TechCorp as a software engineer",
|
||||||
|
context="employee info",
|
||||||
|
request_context=request_context,
|
||||||
|
)
|
||||||
|
await memory.retain_async(
|
||||||
|
bank_id=bank_id,
|
||||||
|
content="TechCorp is located in San Francisco",
|
||||||
|
context="company info",
|
||||||
|
request_context=request_context,
|
||||||
|
)
|
||||||
|
await memory.retain_async(
|
||||||
|
bank_id=bank_id,
|
||||||
|
content="Bob is Alice's manager at TechCorp",
|
||||||
|
context="employee info",
|
||||||
|
request_context=request_context,
|
||||||
|
)
|
||||||
|
await memory.retain_async(
|
||||||
|
bank_id=bank_id,
|
||||||
|
content="San Francisco has many tech companies",
|
||||||
|
context="city info",
|
||||||
|
request_context=request_context,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Query should find related facts via graph traversal
|
||||||
|
from hindsight_api.engine.memory_engine import Budget
|
||||||
|
|
||||||
|
result = await memory.recall_async(
|
||||||
|
bank_id=bank_id,
|
||||||
|
query="Tell me about Alice",
|
||||||
|
fact_type=["world"],
|
||||||
|
budget=Budget.MID,
|
||||||
|
max_tokens=2048,
|
||||||
|
request_context=request_context,
|
||||||
|
)
|
||||||
|
|
||||||
|
# Should return results
|
||||||
|
assert result.results is not None
|
||||||
|
assert len(result.results) > 0
|
||||||
|
|
||||||
|
# Should find Alice-related facts
|
||||||
|
fact_texts = [f.text for f in result.results]
|
||||||
|
alice_facts = [t for t in fact_texts if "Alice" in t or "TechCorp" in t]
|
||||||
|
assert len(alice_facts) > 0, f"Should find Alice-related facts, got: {fact_texts}"
|
||||||
|
|
||||||
|
print(f"\n✓ MPFP integration test passed! Found {len(result.results)} facts")
|
||||||
|
|
||||||
|
finally:
|
||||||
|
await memory.delete_bank(bank_id, request_context=request_context)
|
||||||
|
|
||||||
|
|
||||||
|
@pytest.mark.asyncio
|
||||||
|
async def test_mpfp_lazy_loading_efficiency(memory, request_context):
|
||||||
|
"""Test that MPFP loads edges lazily, not upfront."""
|
||||||
|
bank_id = f"test_mpfp_lazy_{datetime.now(timezone.utc).timestamp()}"
|
||||||
|
|
||||||
|
try:
|
||||||
|
# Store many memories to create a larger graph
|
||||||
|
for i in range(20):
|
||||||
|
await memory.retain_async(
|
||||||
|
bank_id=bank_id,
|
||||||
|
content=f"Fact number {i} about topic {i % 5}",
|
||||||
|
context=f"context {i}",
|
||||||
|
request_context=request_context,
|
||||||
|
)
|
||||||
|
|
||||||
|
from hindsight_api.engine.memory_engine import Budget
|
||||||
|
|
||||||
|
# Query - MPFP should only load edges for relevant frontier nodes
|
||||||
|
result = await memory.recall_async(
|
||||||
|
bank_id=bank_id,
|
||||||
|
query="topic 0",
|
||||||
|
fact_type=["world"],
|
||||||
|
budget=Budget.LOW,
|
||||||
|
max_tokens=1024,
|
||||||
|
enable_trace=True,
|
||||||
|
request_context=request_context,
|
||||||
|
)
|
||||||
|
|
||||||
|
assert result.results is not None
|
||||||
|
|
||||||
|
# Check trace for timing info
|
||||||
|
if result.trace:
|
||||||
|
print(f"\n✓ MPFP lazy loading test passed!")
|
||||||
|
print(f" - Facts returned: {len(result.results)}")
|
||||||
|
|
||||||
|
finally:
|
||||||
|
await memory.delete_bank(bank_id, request_context=request_context)
|
||||||
|
|
@ -189,7 +189,10 @@ Supported OpenAI embedding dimensions:
|
||||||
|----------|-------------|---------|
|
|----------|-------------|---------|
|
||||||
| `HINDSIGHT_API_RERANKER_PROVIDER` | Provider: `local`, `tei`, or `cohere` | `local` |
|
| `HINDSIGHT_API_RERANKER_PROVIDER` | Provider: `local`, `tei`, or `cohere` | `local` |
|
||||||
| `HINDSIGHT_API_RERANKER_LOCAL_MODEL` | Model for local provider | `cross-encoder/ms-marco-MiniLM-L-6-v2` |
|
| `HINDSIGHT_API_RERANKER_LOCAL_MODEL` | Model for local provider | `cross-encoder/ms-marco-MiniLM-L-6-v2` |
|
||||||
|
| `HINDSIGHT_API_RERANKER_LOCAL_MAX_CONCURRENT` | Max concurrent local reranking (prevents CPU thrashing under load) | `4` |
|
||||||
| `HINDSIGHT_API_RERANKER_TEI_URL` | TEI server URL | - |
|
| `HINDSIGHT_API_RERANKER_TEI_URL` | TEI server URL | - |
|
||||||
|
| `HINDSIGHT_API_RERANKER_TEI_BATCH_SIZE` | Batch size for TEI reranking | `128` |
|
||||||
|
| `HINDSIGHT_API_RERANKER_TEI_MAX_CONCURRENT` | Max concurrent TEI reranking requests | `8` |
|
||||||
| `HINDSIGHT_API_RERANKER_COHERE_MODEL` | Cohere rerank model | `rerank-english-v3.0` |
|
| `HINDSIGHT_API_RERANKER_COHERE_MODEL` | Cohere rerank model | `rerank-english-v3.0` |
|
||||||
|
|
||||||
```bash
|
```bash
|
||||||
|
|
@ -264,6 +267,7 @@ Controls the retain (memory ingestion) pipeline.
|
||||||
| `HINDSIGHT_API_RETAIN_CHUNK_SIZE` | Max characters per chunk for fact extraction. Larger chunks extract fewer LLM calls but may lose context. | `3000` |
|
| `HINDSIGHT_API_RETAIN_CHUNK_SIZE` | Max characters per chunk for fact extraction. Larger chunks extract fewer LLM calls but may lose context. | `3000` |
|
||||||
| `HINDSIGHT_API_RETAIN_EXTRACTION_MODE` | Fact extraction mode: `concise` (selective, fewer high-quality facts) or `verbose` (detailed, more facts) | `concise` |
|
| `HINDSIGHT_API_RETAIN_EXTRACTION_MODE` | Fact extraction mode: `concise` (selective, fewer high-quality facts) or `verbose` (detailed, more facts) | `concise` |
|
||||||
| `HINDSIGHT_API_RETAIN_EXTRACT_CAUSAL_LINKS` | Extract causal relationships between facts | `true` |
|
| `HINDSIGHT_API_RETAIN_EXTRACT_CAUSAL_LINKS` | Extract causal relationships between facts | `true` |
|
||||||
|
| `HINDSIGHT_API_RETAIN_OBSERVATIONS_ASYNC` | Run entity observation generation asynchronously (after retain completes) | `false` |
|
||||||
|
|
||||||
#### Extraction Modes
|
#### Extraction Modes
|
||||||
|
|
||||||
|
|
@ -293,8 +297,9 @@ Controls background task processing for async operations like opinion formation
|
||||||
|
|
||||||
| Variable | Description | Default |
|
| Variable | Description | Default |
|
||||||
|----------|-------------|---------|
|
|----------|-------------|---------|
|
||||||
| `HINDSIGHT_API_TASK_BATCH_SIZE` | Max tasks to process in one batch | `10` |
|
| `HINDSIGHT_API_TASK_BACKEND` | Task backend implementation: `memory` (in-process queue) or `noop` (discard tasks, useful for tests) | `memory` |
|
||||||
| `HINDSIGHT_API_TASK_BATCH_INTERVAL` | Interval between batch processing in seconds | `1.0` |
|
| `HINDSIGHT_API_TASK_BACKEND_MEMORY_BATCH_SIZE` | Max tasks to process in one batch (memory backend only) | `10` |
|
||||||
|
| `HINDSIGHT_API_TASK_BACKEND_MEMORY_BATCH_INTERVAL` | Interval between batch processing in seconds (memory backend only) | `1.0` |
|
||||||
|
|
||||||
### Performance Optimization
|
### Performance Optimization
|
||||||
|
|
||||||
|
|
|
||||||
|
|
@ -171,4 +171,4 @@ PORT=80 HINDSIGHT_CP_DATAPLANE_API_URL=https://api.hindsight.io npx @vectorize-i
|
||||||
|
|
||||||
- [Configuration](./configuration.md) — Environment variables and settings
|
- [Configuration](./configuration.md) — Environment variables and settings
|
||||||
- [Models](./models.md) — ML models and providers
|
- [Models](./models.md) — ML models and providers
|
||||||
- [Metrics](./metrics.md) — Monitoring and observability
|
- [Monitoring](./monitoring.md) — Metrics and observability
|
||||||
|
|
|
||||||
|
|
@ -1,95 +0,0 @@
|
||||||
# Metrics
|
|
||||||
|
|
||||||
Hindsight exposes Prometheus metrics at `/metrics` for monitoring.
|
|
||||||
|
|
||||||
```bash
|
|
||||||
curl http://localhost:8888/metrics
|
|
||||||
```
|
|
||||||
|
|
||||||
## Available Metrics
|
|
||||||
|
|
||||||
### Operation Metrics
|
|
||||||
|
|
||||||
| Metric | Type | Labels | Description |
|
|
||||||
|--------|------|--------|-------------|
|
|
||||||
| `hindsight.operation.duration` | Histogram | operation, bank_id, source, budget, max_tokens, success | Duration of operations in seconds |
|
|
||||||
| `hindsight.operation.total` | Counter | operation, bank_id, source, budget, max_tokens, success | Total number of operations executed |
|
|
||||||
|
|
||||||
**Labels:**
|
|
||||||
- `operation`: Operation type (`retain`, `recall`, `reflect`)
|
|
||||||
- `bank_id`: Memory bank identifier
|
|
||||||
- `source`: Where the operation was triggered from (`api`, `reflect`, `internal`)
|
|
||||||
- `budget`: Budget level if specified (`low`, `mid`, `high`)
|
|
||||||
- `max_tokens`: Max tokens if specified
|
|
||||||
- `success`: Whether the operation succeeded (`true`, `false`)
|
|
||||||
|
|
||||||
The `source` label allows distinguishing between:
|
|
||||||
- `api`: Direct API calls from clients
|
|
||||||
- `reflect`: Internal recall calls made during reflect operations
|
|
||||||
- `internal`: Other internal operations
|
|
||||||
|
|
||||||
### LLM Metrics
|
|
||||||
|
|
||||||
| Metric | Type | Labels | Description |
|
|
||||||
|--------|------|--------|-------------|
|
|
||||||
| `hindsight.llm.duration` | Histogram | provider, model, scope, success | Duration of LLM API calls in seconds |
|
|
||||||
| `hindsight.llm.calls.total` | Counter | provider, model, scope, success | Total number of LLM API calls |
|
|
||||||
| `hindsight.llm.tokens.input` | Counter | provider, model, scope, success, token_bucket | Input tokens for LLM calls |
|
|
||||||
| `hindsight.llm.tokens.output` | Counter | provider, model, scope, success, token_bucket | Output tokens from LLM calls |
|
|
||||||
|
|
||||||
**Labels:**
|
|
||||||
- `provider`: LLM provider (`openai`, `anthropic`, `gemini`, `groq`, `ollama`, `lmstudio`)
|
|
||||||
- `model`: Model name (e.g., `gpt-4`, `claude-3-sonnet`)
|
|
||||||
- `scope`: What the LLM call is for (`memory`, `reflect`, `entity_observation`, `answer`)
|
|
||||||
- `success`: Whether the call succeeded (`true`, `false`)
|
|
||||||
- `token_bucket`: Token count bucket for cardinality control (`0-100`, `100-500`, `500-1k`, `1k-5k`, `5k-10k`, `10k-50k`, `50k+`)
|
|
||||||
|
|
||||||
### Histogram Buckets
|
|
||||||
|
|
||||||
Custom bucket boundaries are configured for better percentile accuracy:
|
|
||||||
|
|
||||||
**Operation Duration Buckets (seconds):**
|
|
||||||
```
|
|
||||||
0.1, 0.25, 0.5, 0.75, 1.0, 2.0, 3.0, 5.0, 7.5, 10.0, 15.0, 20.0, 30.0, 60.0, 120.0
|
|
||||||
```
|
|
||||||
|
|
||||||
**LLM Duration Buckets (seconds):**
|
|
||||||
```
|
|
||||||
0.1, 0.25, 0.5, 1.0, 2.0, 3.0, 5.0, 10.0, 15.0, 30.0, 60.0, 120.0
|
|
||||||
```
|
|
||||||
|
|
||||||
## Prometheus Configuration
|
|
||||||
|
|
||||||
```yaml
|
|
||||||
scrape_configs:
|
|
||||||
- job_name: 'hindsight'
|
|
||||||
static_configs:
|
|
||||||
- targets: ['localhost:8888']
|
|
||||||
```
|
|
||||||
|
|
||||||
## Example Queries
|
|
||||||
|
|
||||||
### Average operation latency by type
|
|
||||||
```promql
|
|
||||||
rate(hindsight_operation_duration_sum[5m]) / rate(hindsight_operation_duration_count[5m])
|
|
||||||
```
|
|
||||||
|
|
||||||
### LLM calls per minute by provider
|
|
||||||
```promql
|
|
||||||
rate(hindsight_llm_calls_total[1m]) * 60
|
|
||||||
```
|
|
||||||
|
|
||||||
### P95 LLM latency
|
|
||||||
```promql
|
|
||||||
histogram_quantile(0.95, rate(hindsight_llm_duration_bucket[5m]))
|
|
||||||
```
|
|
||||||
|
|
||||||
### Total tokens consumed by model
|
|
||||||
```promql
|
|
||||||
sum by (model) (hindsight_llm_tokens_input_total + hindsight_llm_tokens_output_total)
|
|
||||||
```
|
|
||||||
|
|
||||||
### Internal vs API recall operations
|
|
||||||
```promql
|
|
||||||
sum by (source) (rate(hindsight_operation_total{operation="recall"}[5m]))
|
|
||||||
```
|
|
||||||
199
hindsight-docs/docs/developer/monitoring.md
Normal file
199
hindsight-docs/docs/developer/monitoring.md
Normal file
|
|
@ -0,0 +1,199 @@
|
||||||
|
# Monitoring
|
||||||
|
|
||||||
|
Hindsight provides comprehensive monitoring through Prometheus metrics and pre-built Grafana dashboards.
|
||||||
|
|
||||||
|
## Local Development
|
||||||
|
|
||||||
|
For local metrics visualization, a convenience script downloads and runs Prometheus and Grafana:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
./scripts/dev/start-monitoring.sh
|
||||||
|
```
|
||||||
|
|
||||||
|
This will start:
|
||||||
|
- **Grafana**: http://localhost:8890 (anonymous access enabled)
|
||||||
|
- **Prometheus**: http://localhost:8889
|
||||||
|
- **API Metrics**: http://localhost:8888/metrics
|
||||||
|
|
||||||
|
:::note Production Deployment
|
||||||
|
The local monitoring script is for development only. In production, you need to install and configure Prometheus and Grafana separately, then point Prometheus to scrape your Hindsight API's `/metrics` endpoint.
|
||||||
|
:::
|
||||||
|
|
||||||
|
## Grafana Dashboards
|
||||||
|
|
||||||
|
Pre-built dashboards are available in [`monitoring/grafana/dashboards/`](https://github.com/anthropics/hindsight/tree/main/monitoring/grafana/dashboards). Import these JSON files into your Grafana instance:
|
||||||
|
|
||||||
|
| Dashboard | Description |
|
||||||
|
|-----------|-------------|
|
||||||
|
| **Hindsight Operations** | Operation rates, latency percentiles, per-bank metrics |
|
||||||
|
| **Hindsight LLM Metrics** | LLM calls, token usage, latency by scope/provider |
|
||||||
|
| **Hindsight API Service** | HTTP requests, error rates, DB pool, process metrics |
|
||||||
|
|
||||||
|
The dashboards are automatically provisioned when using the monitoring stack script.
|
||||||
|
|
||||||
|
## Metrics Endpoint
|
||||||
|
|
||||||
|
Hindsight exposes Prometheus metrics at `/metrics`:
|
||||||
|
|
||||||
|
```bash
|
||||||
|
curl http://localhost:8888/metrics
|
||||||
|
```
|
||||||
|
|
||||||
|
## Available Metrics
|
||||||
|
|
||||||
|
### Operation Metrics
|
||||||
|
|
||||||
|
| Metric | Type | Labels | Description |
|
||||||
|
|--------|------|--------|-------------|
|
||||||
|
| `hindsight.operation.duration` | Histogram | operation, bank_id, source, budget, max_tokens, success | Duration of operations in seconds |
|
||||||
|
| `hindsight.operation.total` | Counter | operation, bank_id, source, budget, max_tokens, success | Total number of operations executed |
|
||||||
|
|
||||||
|
**Labels:**
|
||||||
|
- `operation`: Operation type (`retain`, `recall`, `reflect`)
|
||||||
|
- `bank_id`: Memory bank identifier
|
||||||
|
- `source`: Where the operation was triggered from (`api`, `reflect`, `internal`)
|
||||||
|
- `budget`: Budget level if specified (`low`, `mid`, `high`)
|
||||||
|
- `max_tokens`: Max tokens if specified
|
||||||
|
- `success`: Whether the operation succeeded (`true`, `false`)
|
||||||
|
|
||||||
|
The `source` label allows distinguishing between:
|
||||||
|
- `api`: Direct API calls from clients
|
||||||
|
- `reflect`: Internal recall calls made during reflect operations
|
||||||
|
- `internal`: Other internal operations
|
||||||
|
|
||||||
|
### LLM Metrics
|
||||||
|
|
||||||
|
| Metric | Type | Labels | Description |
|
||||||
|
|--------|------|--------|-------------|
|
||||||
|
| `hindsight.llm.duration` | Histogram | provider, model, scope, success | Duration of LLM API calls in seconds |
|
||||||
|
| `hindsight.llm.calls.total` | Counter | provider, model, scope, success | Total number of LLM API calls |
|
||||||
|
| `hindsight.llm.tokens.input` | Counter | provider, model, scope, success, token_bucket | Input tokens for LLM calls |
|
||||||
|
| `hindsight.llm.tokens.output` | Counter | provider, model, scope, success, token_bucket | Output tokens from LLM calls |
|
||||||
|
|
||||||
|
**Labels:**
|
||||||
|
- `provider`: LLM provider (`openai`, `anthropic`, `gemini`, `groq`, `ollama`, `lmstudio`)
|
||||||
|
- `model`: Model name (e.g., `gpt-4`, `claude-3-sonnet`)
|
||||||
|
- `scope`: What the LLM call is for (`memory`, `reflect`, `entity_observation`, `answer`)
|
||||||
|
- `success`: Whether the call succeeded (`true`, `false`)
|
||||||
|
- `token_bucket`: Token count bucket for cardinality control (`0-100`, `100-500`, `500-1k`, `1k-5k`, `5k-10k`, `10k-50k`, `50k+`)
|
||||||
|
|
||||||
|
### HTTP Request Metrics
|
||||||
|
|
||||||
|
| Metric | Type | Labels | Description |
|
||||||
|
|--------|------|--------|-------------|
|
||||||
|
| `hindsight.http.duration` | Histogram | method, endpoint, status_code, status_class | Duration of HTTP requests in seconds |
|
||||||
|
| `hindsight.http.requests.total` | Counter | method, endpoint, status_code, status_class | Total number of HTTP requests |
|
||||||
|
| `hindsight.http.requests.in_progress` | UpDownCounter | method, endpoint | Number of HTTP requests currently being processed |
|
||||||
|
|
||||||
|
**Labels:**
|
||||||
|
- `method`: HTTP method (`GET`, `POST`, `PUT`, `DELETE`)
|
||||||
|
- `endpoint`: Request path (normalized to reduce cardinality - UUIDs replaced with `{id}`)
|
||||||
|
- `status_code`: HTTP status code (`200`, `400`, `500`, etc.)
|
||||||
|
- `status_class`: Status code class (`2xx`, `4xx`, `5xx`)
|
||||||
|
|
||||||
|
### Database Pool Metrics
|
||||||
|
|
||||||
|
| Metric | Type | Labels | Description |
|
||||||
|
|--------|------|--------|-------------|
|
||||||
|
| `hindsight.db.pool.size` | Gauge | - | Current number of connections in the pool |
|
||||||
|
| `hindsight.db.pool.idle` | Gauge | - | Number of idle connections in the pool |
|
||||||
|
| `hindsight.db.pool.min` | Gauge | - | Minimum pool size |
|
||||||
|
| `hindsight.db.pool.max` | Gauge | - | Maximum pool size |
|
||||||
|
|
||||||
|
### Process Metrics
|
||||||
|
|
||||||
|
| Metric | Type | Labels | Description |
|
||||||
|
|--------|------|--------|-------------|
|
||||||
|
| `hindsight.process.cpu.seconds` | Gauge | type | Process CPU time in seconds |
|
||||||
|
| `hindsight.process.memory.bytes` | Gauge | type | Process memory usage in bytes |
|
||||||
|
| `hindsight.process.open_fds` | Gauge | - | Number of open file descriptors |
|
||||||
|
| `hindsight.process.threads` | Gauge | - | Number of active threads |
|
||||||
|
|
||||||
|
**Labels:**
|
||||||
|
- `type` (CPU): `user` or `system`
|
||||||
|
- `type` (Memory): `rss_max` (maximum resident set size)
|
||||||
|
|
||||||
|
### Histogram Buckets
|
||||||
|
|
||||||
|
Custom bucket boundaries are configured for better percentile accuracy:
|
||||||
|
|
||||||
|
**Operation Duration Buckets (seconds):**
|
||||||
|
```
|
||||||
|
0.1, 0.25, 0.5, 0.75, 1.0, 2.0, 3.0, 5.0, 7.5, 10.0, 15.0, 20.0, 30.0, 60.0, 120.0
|
||||||
|
```
|
||||||
|
|
||||||
|
**LLM Duration Buckets (seconds):**
|
||||||
|
```
|
||||||
|
0.1, 0.25, 0.5, 1.0, 2.0, 3.0, 5.0, 10.0, 15.0, 30.0, 60.0, 120.0
|
||||||
|
```
|
||||||
|
|
||||||
|
**HTTP Duration Buckets (seconds):**
|
||||||
|
```
|
||||||
|
0.005, 0.01, 0.025, 0.05, 0.1, 0.25, 0.5, 1.0, 2.5, 5.0, 10.0, 30.0
|
||||||
|
```
|
||||||
|
|
||||||
|
## Prometheus Configuration
|
||||||
|
|
||||||
|
```yaml
|
||||||
|
scrape_configs:
|
||||||
|
- job_name: 'hindsight'
|
||||||
|
static_configs:
|
||||||
|
- targets: ['localhost:8888']
|
||||||
|
```
|
||||||
|
|
||||||
|
## Example Queries
|
||||||
|
|
||||||
|
### Average operation latency by type
|
||||||
|
```promql
|
||||||
|
rate(hindsight_operation_duration_sum[5m]) / rate(hindsight_operation_duration_count[5m])
|
||||||
|
```
|
||||||
|
|
||||||
|
### LLM calls per minute by provider
|
||||||
|
```promql
|
||||||
|
rate(hindsight_llm_calls_total[1m]) * 60
|
||||||
|
```
|
||||||
|
|
||||||
|
### P95 LLM latency
|
||||||
|
```promql
|
||||||
|
histogram_quantile(0.95, rate(hindsight_llm_duration_bucket[5m]))
|
||||||
|
```
|
||||||
|
|
||||||
|
### Total tokens consumed by model
|
||||||
|
```promql
|
||||||
|
sum by (model) (hindsight_llm_tokens_input_total + hindsight_llm_tokens_output_total)
|
||||||
|
```
|
||||||
|
|
||||||
|
### Internal vs API recall operations
|
||||||
|
```promql
|
||||||
|
sum by (source) (rate(hindsight_operation_total{operation="recall"}[5m]))
|
||||||
|
```
|
||||||
|
|
||||||
|
### HTTP requests per second by endpoint
|
||||||
|
```promql
|
||||||
|
sum by (endpoint) (rate(hindsight_http_requests_total[1m]))
|
||||||
|
```
|
||||||
|
|
||||||
|
### HTTP error rate (5xx)
|
||||||
|
```promql
|
||||||
|
sum(rate(hindsight_http_requests_total{status_class="5xx"}[5m])) / sum(rate(hindsight_http_requests_total[5m]))
|
||||||
|
```
|
||||||
|
|
||||||
|
### P95 HTTP latency
|
||||||
|
```promql
|
||||||
|
histogram_quantile(0.95, sum by (le) (rate(hindsight_http_duration_seconds_bucket[5m])))
|
||||||
|
```
|
||||||
|
|
||||||
|
### Database pool utilization
|
||||||
|
```promql
|
||||||
|
hindsight_db_pool_size / hindsight_db_pool_max
|
||||||
|
```
|
||||||
|
|
||||||
|
### Active database connections
|
||||||
|
```promql
|
||||||
|
hindsight_db_pool_size - hindsight_db_pool_idle
|
||||||
|
```
|
||||||
|
|
||||||
|
### CPU usage rate
|
||||||
|
```promql
|
||||||
|
rate(hindsight_process_cpu_seconds{type="user"}[1m])
|
||||||
|
```
|
||||||
|
|
@ -229,6 +229,49 @@ Budget and max_tokens control different aspects of recall:
|
||||||
|
|
||||||
---
|
---
|
||||||
|
|
||||||
|
## Graph Retrieval Algorithms
|
||||||
|
|
||||||
|
Hindsight supports two graph traversal algorithms, each optimized for different scenarios:
|
||||||
|
|
||||||
|
| Algorithm | Default | Best For | Complexity |
|
||||||
|
|-----------|---------|----------|------------|
|
||||||
|
| **MPFP** | ✓ | Large graphs, production | O(P × H × F × K) |
|
||||||
|
| **BFS** | | Small graphs, debugging | O(V + E) |
|
||||||
|
|
||||||
|
### MPFP (Meta-Path Forward Push)
|
||||||
|
|
||||||
|
A sublinear graph traversal algorithm that follows predefined meta-paths (patterns of edge types) using lazy edge loading.
|
||||||
|
|
||||||
|
**How it works:**
|
||||||
|
1. Starts from semantic entry points (top similar facts)
|
||||||
|
2. Follows multiple meta-path patterns in parallel:
|
||||||
|
- `semantic → semantic` (topic expansion)
|
||||||
|
- `entity → temporal` (entity timeline)
|
||||||
|
- `semantic → causes` (causal reasoning)
|
||||||
|
- `entity → semantic` (entity context)
|
||||||
|
3. Loads edges lazily per hop, only for active frontier nodes
|
||||||
|
4. Fuses results from all patterns via Reciprocal Rank Fusion (RRF)
|
||||||
|
|
||||||
|
**Complexity:** O(P × H × F × K) where P = patterns (~7), H = hops (2), F = frontier size (~20-100), K = neighbors per node (20).
|
||||||
|
|
||||||
|
**Use case:** Production workloads with large memory banks (10k+ facts). Only loads the edges it needs, avoiding full graph scans.
|
||||||
|
|
||||||
|
### BFS (Breadth-First Spreading Activation)
|
||||||
|
|
||||||
|
Classic spreading activation that propagates relevance scores through the graph using breadth-first traversal.
|
||||||
|
|
||||||
|
**How it works:**
|
||||||
|
1. Starts from semantic entry points with initial activation scores
|
||||||
|
2. Spreads activation to neighbors with decay (α = 0.8 per hop)
|
||||||
|
3. Boosts causal links (causes, enables, prevents)
|
||||||
|
4. Continues until budget exhausted or activation below threshold
|
||||||
|
|
||||||
|
**Complexity:** O(V + E) where V and E are visited nodes and edges, bounded by budget.
|
||||||
|
|
||||||
|
**Use case:** Small memory banks, debugging, or when you need to understand exactly how results were found.
|
||||||
|
|
||||||
|
---
|
||||||
|
|
||||||
## Next Steps
|
## Next Steps
|
||||||
|
|
||||||
- [**Retain**](./retain) — How memories are stored with rich context
|
- [**Retain**](./retain) — How memories are stored with rich context
|
||||||
|
|
|
||||||
|
|
@ -133,8 +133,8 @@ const sidebars: SidebarsConfig = {
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
type: 'doc',
|
type: 'doc',
|
||||||
id: 'developer/metrics',
|
id: 'developer/monitoring',
|
||||||
label: 'Metrics',
|
label: 'Monitoring',
|
||||||
},
|
},
|
||||||
{
|
{
|
||||||
type: 'doc',
|
type: 'doc',
|
||||||
|
|
|
||||||
1291
monitoring/grafana/dashboards/hindsight-api-service.json
Normal file
1291
monitoring/grafana/dashboards/hindsight-api-service.json
Normal file
File diff suppressed because it is too large
Load diff
|
|
@ -107,8 +107,8 @@ mkdir -p "$GRAFANA_PROV_DIR/dashboards"
|
||||||
mkdir -p "$GRAFANA_DIR/dashboards"
|
mkdir -p "$GRAFANA_DIR/dashboards"
|
||||||
mkdir -p "$GRAFANA_DIR/data"
|
mkdir -p "$GRAFANA_DIR/data"
|
||||||
|
|
||||||
# Copy dashboards
|
# Copy dashboards from project root monitoring directory
|
||||||
cp "$SCRIPT_DIR/grafana/dashboards/"*.json "$GRAFANA_DIR/dashboards/"
|
cp "$PROJECT_ROOT/monitoring/grafana/dashboards/"*.json "$GRAFANA_DIR/dashboards/"
|
||||||
|
|
||||||
# Create Grafana datasource config
|
# Create Grafana datasource config
|
||||||
cat > "$GRAFANA_PROV_DIR/datasources/prometheus.yaml" <<EOF
|
cat > "$GRAFANA_PROV_DIR/datasources/prometheus.yaml" <<EOF
|
||||||
|
|
@ -181,6 +181,7 @@ echo ""
|
||||||
echo " Dashboards:"
|
echo " Dashboards:"
|
||||||
echo " - Hindsight Operations"
|
echo " - Hindsight Operations"
|
||||||
echo " - Hindsight LLM Metrics"
|
echo " - Hindsight LLM Metrics"
|
||||||
|
echo " - Hindsight API Service"
|
||||||
echo ""
|
echo ""
|
||||||
echo "=================================="
|
echo "=================================="
|
||||||
echo ""
|
echo ""
|
||||||
|
|
|
||||||
Loading…
Reference in a new issue