Compare commits

...
2 Commits
Author SHA1 Message Date
Nicolò Boschi 5c7010db98 fix doc build 2025-12-16 13:43:22 +01:00
Nicolò Boschi 7f453d0761 fix doc build 2025-12-16 13:42:42 +01:00
124 changed files with 5584 additions and 5317 deletions
+10 -9
View File
@@ -3,23 +3,24 @@ Memory System for AI Agents.
Temporal + Semantic Memory Architecture using PostgreSQL with pgvector.
"""
from .config import HindsightConfig, get_config
from .engine.cross_encoder import CrossEncoderModel, LocalSTCrossEncoder, RemoteTEICrossEncoder
from .engine.embeddings import Embeddings, LocalSTEmbeddings, RemoteTEIEmbeddings
from .engine.llm_wrapper import LLMConfig
from .engine.memory_engine import MemoryEngine
from .engine.search.trace import (
SearchTrace,
QueryInfo,
EntryPoint,
NodeVisit,
WeightComponents,
LinkInfo,
NodeVisit,
PruningDecision,
SearchSummary,
QueryInfo,
SearchPhaseMetrics,
SearchSummary,
SearchTrace,
WeightComponents,
)
from .engine.search.tracer import SearchTracer
from .engine.embeddings import Embeddings, LocalSTEmbeddings, RemoteTEIEmbeddings
from .engine.cross_encoder import CrossEncoderModel, LocalSTCrossEncoder, RemoteTEICrossEncoder
from .engine.llm_wrapper import LLMConfig
from .config import HindsightConfig, get_config
__all__ = [
"MemoryEngine",
+5 -8
View File
@@ -2,20 +2,19 @@
Alembic environment configuration for SQLAlchemy with pgvector.
Uses synchronous psycopg2 driver for migrations to avoid pgbouncer issues.
"""
import logging
import os
import sys
from pathlib import Path
from sqlalchemy import pool, engine_from_config
from sqlalchemy.engine import Connection
from alembic import context
from dotenv import load_dotenv
from sqlalchemy import engine_from_config, pool
# Import your models here
from hindsight_api.models import Base
# Load environment variables based on HINDSIGHT_API_DATABASE_URL env var or default to local
def load_env():
"""Load environment variables from .env"""
@@ -30,6 +29,7 @@ def load_env():
if env_file.exists():
load_dotenv(env_file)
load_env()
# this is the Alembic Config object, which provides
@@ -128,10 +128,7 @@ def run_migrations_online() -> None:
connection.execute(text("SET SESSION CHARACTERISTICS AS TRANSACTION READ WRITE"))
connection.commit() # Commit the SET command
context.configure(
connection=connection,
target_metadata=target_metadata
)
context.configure(connection=connection, target_metadata=target_metadata)
with context.begin_transaction():
context.run_migrations()
@@ -5,120 +5,150 @@ Revises:
Create Date: 2025-11-27 11:54:19.228030
"""
from typing import Sequence, Union
from alembic import op
from collections.abc import Sequence
import sqlalchemy as sa
from sqlalchemy.dialects import postgresql
from alembic import op
from pgvector.sqlalchemy import Vector
from sqlalchemy.dialects import postgresql
# revision identifiers, used by Alembic.
revision: str = '5a366d414dce'
down_revision: Union[str, Sequence[str], None] = None
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
revision: str = "5a366d414dce"
down_revision: str | Sequence[str] | None = None
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
"""Upgrade schema - create all tables from scratch."""
# Enable required extensions
op.execute('CREATE EXTENSION IF NOT EXISTS vector')
op.execute("CREATE EXTENSION IF NOT EXISTS vector")
# Create banks table
op.create_table(
'banks',
sa.Column('bank_id', sa.Text(), nullable=False),
sa.Column('name', sa.Text(), nullable=True),
sa.Column('personality', postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
sa.Column('background', sa.Text(), nullable=True),
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.Column('updated_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.PrimaryKeyConstraint('bank_id', name=op.f('pk_banks'))
"banks",
sa.Column("bank_id", sa.Text(), nullable=False),
sa.Column("name", sa.Text(), nullable=True),
sa.Column(
"personality",
postgresql.JSONB(astext_type=sa.Text()),
server_default=sa.text("'{}'::jsonb"),
nullable=False,
),
sa.Column("background", sa.Text(), nullable=True),
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.Column("updated_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.PrimaryKeyConstraint("bank_id", name=op.f("pk_banks")),
)
# Create documents table
op.create_table(
'documents',
sa.Column('id', sa.Text(), nullable=False),
sa.Column('bank_id', sa.Text(), nullable=False),
sa.Column('original_text', sa.Text(), nullable=True),
sa.Column('content_hash', sa.Text(), nullable=True),
sa.Column('metadata', postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.Column('updated_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.PrimaryKeyConstraint('id', 'bank_id', name=op.f('pk_documents'))
"documents",
sa.Column("id", sa.Text(), nullable=False),
sa.Column("bank_id", sa.Text(), nullable=False),
sa.Column("original_text", sa.Text(), nullable=True),
sa.Column("content_hash", sa.Text(), nullable=True),
sa.Column(
"metadata", postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False
),
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.Column("updated_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.PrimaryKeyConstraint("id", "bank_id", name=op.f("pk_documents")),
)
op.create_index('idx_documents_bank_id', 'documents', ['bank_id'])
op.create_index('idx_documents_content_hash', 'documents', ['content_hash'])
op.create_index("idx_documents_bank_id", "documents", ["bank_id"])
op.create_index("idx_documents_content_hash", "documents", ["content_hash"])
# Create async_operations table
op.create_table(
'async_operations',
sa.Column('operation_id', postgresql.UUID(as_uuid=True), server_default=sa.text('gen_random_uuid()'), nullable=False),
sa.Column('bank_id', sa.Text(), nullable=False),
sa.Column('operation_type', sa.Text(), nullable=False),
sa.Column('status', sa.Text(), server_default='pending', nullable=False),
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.Column('updated_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.Column('completed_at', postgresql.TIMESTAMP(timezone=True), nullable=True),
sa.Column('error_message', sa.Text(), nullable=True),
sa.Column('result_metadata', postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
sa.PrimaryKeyConstraint('operation_id', name=op.f('pk_async_operations')),
sa.CheckConstraint("status IN ('pending', 'processing', 'completed', 'failed')", name='async_operations_status_check')
"async_operations",
sa.Column(
"operation_id", postgresql.UUID(as_uuid=True), server_default=sa.text("gen_random_uuid()"), nullable=False
),
sa.Column("bank_id", sa.Text(), nullable=False),
sa.Column("operation_type", sa.Text(), nullable=False),
sa.Column("status", sa.Text(), server_default="pending", nullable=False),
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.Column("updated_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.Column("completed_at", postgresql.TIMESTAMP(timezone=True), nullable=True),
sa.Column("error_message", sa.Text(), nullable=True),
sa.Column(
"result_metadata",
postgresql.JSONB(astext_type=sa.Text()),
server_default=sa.text("'{}'::jsonb"),
nullable=False,
),
sa.PrimaryKeyConstraint("operation_id", name=op.f("pk_async_operations")),
sa.CheckConstraint(
"status IN ('pending', 'processing', 'completed', 'failed')", name="async_operations_status_check"
),
)
op.create_index('idx_async_operations_bank_id', 'async_operations', ['bank_id'])
op.create_index('idx_async_operations_status', 'async_operations', ['status'])
op.create_index('idx_async_operations_bank_status', 'async_operations', ['bank_id', 'status'])
op.create_index("idx_async_operations_bank_id", "async_operations", ["bank_id"])
op.create_index("idx_async_operations_status", "async_operations", ["status"])
op.create_index("idx_async_operations_bank_status", "async_operations", ["bank_id", "status"])
# Create entities table
op.create_table(
'entities',
sa.Column('id', postgresql.UUID(as_uuid=True), server_default=sa.text('gen_random_uuid()'), nullable=False),
sa.Column('canonical_name', sa.Text(), nullable=False),
sa.Column('bank_id', sa.Text(), nullable=False),
sa.Column('metadata', postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
sa.Column('first_seen', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.Column('last_seen', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.Column('mention_count', sa.Integer(), server_default='1', nullable=False),
sa.PrimaryKeyConstraint('id', name=op.f('pk_entities'))
"entities",
sa.Column("id", postgresql.UUID(as_uuid=True), server_default=sa.text("gen_random_uuid()"), nullable=False),
sa.Column("canonical_name", sa.Text(), nullable=False),
sa.Column("bank_id", sa.Text(), nullable=False),
sa.Column(
"metadata", postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False
),
sa.Column("first_seen", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.Column("last_seen", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.Column("mention_count", sa.Integer(), server_default="1", nullable=False),
sa.PrimaryKeyConstraint("id", name=op.f("pk_entities")),
)
op.create_index('idx_entities_bank_id', 'entities', ['bank_id'])
op.create_index('idx_entities_canonical_name', 'entities', ['canonical_name'])
op.create_index('idx_entities_bank_name', 'entities', ['bank_id', 'canonical_name'])
op.create_index("idx_entities_bank_id", "entities", ["bank_id"])
op.create_index("idx_entities_canonical_name", "entities", ["canonical_name"])
op.create_index("idx_entities_bank_name", "entities", ["bank_id", "canonical_name"])
# Create unique index on (bank_id, LOWER(canonical_name)) for entity resolution
op.execute('CREATE UNIQUE INDEX idx_entities_bank_lower_name ON entities (bank_id, LOWER(canonical_name))')
op.execute("CREATE UNIQUE INDEX idx_entities_bank_lower_name ON entities (bank_id, LOWER(canonical_name))")
# Create memory_units table
op.create_table(
'memory_units',
sa.Column('id', postgresql.UUID(as_uuid=True), server_default=sa.text('gen_random_uuid()'), nullable=False),
sa.Column('bank_id', sa.Text(), nullable=False),
sa.Column('document_id', sa.Text(), nullable=True),
sa.Column('text', sa.Text(), nullable=False),
sa.Column('embedding', Vector(384), nullable=True),
sa.Column('context', sa.Text(), nullable=True),
sa.Column('event_date', postgresql.TIMESTAMP(timezone=True), nullable=False),
sa.Column('occurred_start', postgresql.TIMESTAMP(timezone=True), nullable=True),
sa.Column('occurred_end', postgresql.TIMESTAMP(timezone=True), nullable=True),
sa.Column('mentioned_at', postgresql.TIMESTAMP(timezone=True), nullable=True),
sa.Column('fact_type', sa.Text(), server_default='world', nullable=False),
sa.Column('confidence_score', sa.Float(), nullable=True),
sa.Column('access_count', sa.Integer(), server_default='0', nullable=False),
sa.Column('metadata', postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False),
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.Column('updated_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.ForeignKeyConstraint(['document_id', 'bank_id'], ['documents.id', 'documents.bank_id'], name='memory_units_document_fkey', ondelete='CASCADE'),
sa.PrimaryKeyConstraint('id', name=op.f('pk_memory_units')),
sa.CheckConstraint("fact_type IN ('world', 'bank', 'opinion', 'observation')", name='memory_units_fact_type_check'),
sa.CheckConstraint("confidence_score IS NULL OR (confidence_score >= 0.0 AND confidence_score <= 1.0)", name='memory_units_confidence_range_check'),
"memory_units",
sa.Column("id", postgresql.UUID(as_uuid=True), server_default=sa.text("gen_random_uuid()"), nullable=False),
sa.Column("bank_id", sa.Text(), nullable=False),
sa.Column("document_id", sa.Text(), nullable=True),
sa.Column("text", sa.Text(), nullable=False),
sa.Column("embedding", Vector(384), nullable=True),
sa.Column("context", sa.Text(), nullable=True),
sa.Column("event_date", postgresql.TIMESTAMP(timezone=True), nullable=False),
sa.Column("occurred_start", postgresql.TIMESTAMP(timezone=True), nullable=True),
sa.Column("occurred_end", postgresql.TIMESTAMP(timezone=True), nullable=True),
sa.Column("mentioned_at", postgresql.TIMESTAMP(timezone=True), nullable=True),
sa.Column("fact_type", sa.Text(), server_default="world", nullable=False),
sa.Column("confidence_score", sa.Float(), nullable=True),
sa.Column("access_count", sa.Integer(), server_default="0", nullable=False),
sa.Column(
"metadata", postgresql.JSONB(astext_type=sa.Text()), server_default=sa.text("'{}'::jsonb"), nullable=False
),
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.Column("updated_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.ForeignKeyConstraint(
["document_id", "bank_id"],
["documents.id", "documents.bank_id"],
name="memory_units_document_fkey",
ondelete="CASCADE",
),
sa.PrimaryKeyConstraint("id", name=op.f("pk_memory_units")),
sa.CheckConstraint(
"fact_type IN ('world', 'bank', 'opinion', 'observation')", name="memory_units_fact_type_check"
),
sa.CheckConstraint(
"confidence_score IS NULL OR (confidence_score >= 0.0 AND confidence_score <= 1.0)",
name="memory_units_confidence_range_check",
),
sa.CheckConstraint(
"(fact_type = 'opinion' AND confidence_score IS NOT NULL) OR "
"(fact_type = 'observation') OR "
"(fact_type NOT IN ('opinion', 'observation') AND confidence_score IS NULL)",
name='confidence_score_fact_type_check'
)
name="confidence_score_fact_type_check",
),
)
# Add search_vector column for full-text search
@@ -128,18 +158,41 @@ def upgrade() -> None:
GENERATED ALWAYS AS (to_tsvector('english', COALESCE(text, '') || ' ' || COALESCE(context, ''))) STORED
""")
op.create_index('idx_memory_units_bank_id', 'memory_units', ['bank_id'])
op.create_index('idx_memory_units_document_id', 'memory_units', ['document_id'])
op.create_index('idx_memory_units_event_date', 'memory_units', [sa.text('event_date DESC')])
op.create_index('idx_memory_units_bank_date', 'memory_units', ['bank_id', sa.text('event_date DESC')])
op.create_index('idx_memory_units_access_count', 'memory_units', [sa.text('access_count DESC')])
op.create_index('idx_memory_units_fact_type', 'memory_units', ['fact_type'])
op.create_index('idx_memory_units_bank_fact_type', 'memory_units', ['bank_id', 'fact_type'])
op.create_index('idx_memory_units_bank_type_date', 'memory_units', ['bank_id', 'fact_type', sa.text('event_date DESC')])
op.create_index('idx_memory_units_opinion_confidence', 'memory_units', ['bank_id', sa.text('confidence_score DESC')], postgresql_where=sa.text("fact_type = 'opinion'"))
op.create_index('idx_memory_units_opinion_date', 'memory_units', ['bank_id', sa.text('event_date DESC')], postgresql_where=sa.text("fact_type = 'opinion'"))
op.create_index('idx_memory_units_observation_date', 'memory_units', ['bank_id', sa.text('event_date DESC')], postgresql_where=sa.text("fact_type = 'observation'"))
op.create_index('idx_memory_units_embedding', 'memory_units', ['embedding'], postgresql_using='hnsw', postgresql_ops={'embedding': 'vector_cosine_ops'})
op.create_index("idx_memory_units_bank_id", "memory_units", ["bank_id"])
op.create_index("idx_memory_units_document_id", "memory_units", ["document_id"])
op.create_index("idx_memory_units_event_date", "memory_units", [sa.text("event_date DESC")])
op.create_index("idx_memory_units_bank_date", "memory_units", ["bank_id", sa.text("event_date DESC")])
op.create_index("idx_memory_units_access_count", "memory_units", [sa.text("access_count DESC")])
op.create_index("idx_memory_units_fact_type", "memory_units", ["fact_type"])
op.create_index("idx_memory_units_bank_fact_type", "memory_units", ["bank_id", "fact_type"])
op.create_index(
"idx_memory_units_bank_type_date", "memory_units", ["bank_id", "fact_type", sa.text("event_date DESC")]
)
op.create_index(
"idx_memory_units_opinion_confidence",
"memory_units",
["bank_id", sa.text("confidence_score DESC")],
postgresql_where=sa.text("fact_type = 'opinion'"),
)
op.create_index(
"idx_memory_units_opinion_date",
"memory_units",
["bank_id", sa.text("event_date DESC")],
postgresql_where=sa.text("fact_type = 'opinion'"),
)
op.create_index(
"idx_memory_units_observation_date",
"memory_units",
["bank_id", sa.text("event_date DESC")],
postgresql_where=sa.text("fact_type = 'observation'"),
)
op.create_index(
"idx_memory_units_embedding",
"memory_units",
["embedding"],
postgresql_using="hnsw",
postgresql_ops={"embedding": "vector_cosine_ops"},
)
# Create BM25 full-text search index on search_vector
op.execute("""
@@ -158,116 +211,149 @@ def upgrade() -> None:
FROM memory_units
""")
op.create_index('idx_memory_units_bm25_bank', 'memory_units_bm25', ['bank_id'])
op.create_index('idx_memory_units_bm25_text_vector', 'memory_units_bm25', ['text_vector'], postgresql_using='gin')
op.create_index("idx_memory_units_bm25_bank", "memory_units_bm25", ["bank_id"])
op.create_index("idx_memory_units_bm25_text_vector", "memory_units_bm25", ["text_vector"], postgresql_using="gin")
# Create entity_cooccurrences table
op.create_table(
'entity_cooccurrences',
sa.Column('entity_id_1', postgresql.UUID(as_uuid=True), nullable=False),
sa.Column('entity_id_2', postgresql.UUID(as_uuid=True), nullable=False),
sa.Column('cooccurrence_count', sa.Integer(), server_default='1', nullable=False),
sa.Column('last_cooccurred', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.ForeignKeyConstraint(['entity_id_1'], ['entities.id'], name=op.f('fk_entity_cooccurrences_entity_id_1_entities'), ondelete='CASCADE'),
sa.ForeignKeyConstraint(['entity_id_2'], ['entities.id'], name=op.f('fk_entity_cooccurrences_entity_id_2_entities'), ondelete='CASCADE'),
sa.PrimaryKeyConstraint('entity_id_1', 'entity_id_2', name=op.f('pk_entity_cooccurrences')),
sa.CheckConstraint('entity_id_1 < entity_id_2', name='entity_cooccurrence_order_check')
"entity_cooccurrences",
sa.Column("entity_id_1", postgresql.UUID(as_uuid=True), nullable=False),
sa.Column("entity_id_2", postgresql.UUID(as_uuid=True), nullable=False),
sa.Column("cooccurrence_count", sa.Integer(), server_default="1", nullable=False),
sa.Column(
"last_cooccurred", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False
),
sa.ForeignKeyConstraint(
["entity_id_1"],
["entities.id"],
name=op.f("fk_entity_cooccurrences_entity_id_1_entities"),
ondelete="CASCADE",
),
sa.ForeignKeyConstraint(
["entity_id_2"],
["entities.id"],
name=op.f("fk_entity_cooccurrences_entity_id_2_entities"),
ondelete="CASCADE",
),
sa.PrimaryKeyConstraint("entity_id_1", "entity_id_2", name=op.f("pk_entity_cooccurrences")),
sa.CheckConstraint("entity_id_1 < entity_id_2", name="entity_cooccurrence_order_check"),
)
op.create_index('idx_entity_cooccurrences_entity1', 'entity_cooccurrences', ['entity_id_1'])
op.create_index('idx_entity_cooccurrences_entity2', 'entity_cooccurrences', ['entity_id_2'])
op.create_index('idx_entity_cooccurrences_count', 'entity_cooccurrences', [sa.text('cooccurrence_count DESC')])
op.create_index("idx_entity_cooccurrences_entity1", "entity_cooccurrences", ["entity_id_1"])
op.create_index("idx_entity_cooccurrences_entity2", "entity_cooccurrences", ["entity_id_2"])
op.create_index("idx_entity_cooccurrences_count", "entity_cooccurrences", [sa.text("cooccurrence_count DESC")])
# Create memory_links table
op.create_table(
'memory_links',
sa.Column('from_unit_id', postgresql.UUID(as_uuid=True), nullable=False),
sa.Column('to_unit_id', postgresql.UUID(as_uuid=True), nullable=False),
sa.Column('link_type', sa.Text(), nullable=False),
sa.Column('entity_id', postgresql.UUID(as_uuid=True), nullable=True),
sa.Column('weight', sa.Float(), server_default='1.0', nullable=False),
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.ForeignKeyConstraint(['entity_id'], ['entities.id'], name=op.f('fk_memory_links_entity_id_entities'), ondelete='CASCADE'),
sa.ForeignKeyConstraint(['from_unit_id'], ['memory_units.id'], name=op.f('fk_memory_links_from_unit_id_memory_units'), ondelete='CASCADE'),
sa.ForeignKeyConstraint(['to_unit_id'], ['memory_units.id'], name=op.f('fk_memory_links_to_unit_id_memory_units'), ondelete='CASCADE'),
sa.CheckConstraint("link_type IN ('temporal', 'semantic', 'entity', 'causes', 'caused_by', 'enables', 'prevents')", name='memory_links_link_type_check'),
sa.CheckConstraint('weight >= 0.0 AND weight <= 1.0', name='memory_links_weight_check')
"memory_links",
sa.Column("from_unit_id", postgresql.UUID(as_uuid=True), nullable=False),
sa.Column("to_unit_id", postgresql.UUID(as_uuid=True), nullable=False),
sa.Column("link_type", sa.Text(), nullable=False),
sa.Column("entity_id", postgresql.UUID(as_uuid=True), nullable=True),
sa.Column("weight", sa.Float(), server_default="1.0", nullable=False),
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.ForeignKeyConstraint(
["entity_id"], ["entities.id"], name=op.f("fk_memory_links_entity_id_entities"), ondelete="CASCADE"
),
sa.ForeignKeyConstraint(
["from_unit_id"],
["memory_units.id"],
name=op.f("fk_memory_links_from_unit_id_memory_units"),
ondelete="CASCADE",
),
sa.ForeignKeyConstraint(
["to_unit_id"],
["memory_units.id"],
name=op.f("fk_memory_links_to_unit_id_memory_units"),
ondelete="CASCADE",
),
sa.CheckConstraint(
"link_type IN ('temporal', 'semantic', 'entity', 'causes', 'caused_by', 'enables', 'prevents')",
name="memory_links_link_type_check",
),
sa.CheckConstraint("weight >= 0.0 AND weight <= 1.0", name="memory_links_weight_check"),
)
# Create unique constraint using COALESCE for nullable entity_id
op.execute("CREATE UNIQUE INDEX idx_memory_links_unique ON memory_links (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid))")
op.create_index('idx_memory_links_from_unit', 'memory_links', ['from_unit_id'])
op.create_index('idx_memory_links_to_unit', 'memory_links', ['to_unit_id'])
op.create_index('idx_memory_links_entity', 'memory_links', ['entity_id'])
op.create_index('idx_memory_links_link_type', 'memory_links', ['link_type'])
op.execute(
"CREATE UNIQUE INDEX idx_memory_links_unique ON memory_links (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid))"
)
op.create_index("idx_memory_links_from_unit", "memory_links", ["from_unit_id"])
op.create_index("idx_memory_links_to_unit", "memory_links", ["to_unit_id"])
op.create_index("idx_memory_links_entity", "memory_links", ["entity_id"])
op.create_index("idx_memory_links_link_type", "memory_links", ["link_type"])
# Create unit_entities table
op.create_table(
'unit_entities',
sa.Column('unit_id', postgresql.UUID(as_uuid=True), nullable=False),
sa.Column('entity_id', postgresql.UUID(as_uuid=True), nullable=False),
sa.ForeignKeyConstraint(['entity_id'], ['entities.id'], name=op.f('fk_unit_entities_entity_id_entities'), ondelete='CASCADE'),
sa.ForeignKeyConstraint(['unit_id'], ['memory_units.id'], name=op.f('fk_unit_entities_unit_id_memory_units'), ondelete='CASCADE'),
sa.PrimaryKeyConstraint('unit_id', 'entity_id', name=op.f('pk_unit_entities'))
"unit_entities",
sa.Column("unit_id", postgresql.UUID(as_uuid=True), nullable=False),
sa.Column("entity_id", postgresql.UUID(as_uuid=True), nullable=False),
sa.ForeignKeyConstraint(
["entity_id"], ["entities.id"], name=op.f("fk_unit_entities_entity_id_entities"), ondelete="CASCADE"
),
sa.ForeignKeyConstraint(
["unit_id"], ["memory_units.id"], name=op.f("fk_unit_entities_unit_id_memory_units"), ondelete="CASCADE"
),
sa.PrimaryKeyConstraint("unit_id", "entity_id", name=op.f("pk_unit_entities")),
)
op.create_index('idx_unit_entities_unit', 'unit_entities', ['unit_id'])
op.create_index('idx_unit_entities_entity', 'unit_entities', ['entity_id'])
op.create_index("idx_unit_entities_unit", "unit_entities", ["unit_id"])
op.create_index("idx_unit_entities_entity", "unit_entities", ["entity_id"])
def downgrade() -> None:
"""Downgrade schema - drop all tables."""
# Drop tables in reverse dependency order
op.drop_index('idx_unit_entities_entity', table_name='unit_entities')
op.drop_index('idx_unit_entities_unit', table_name='unit_entities')
op.drop_table('unit_entities')
op.drop_index("idx_unit_entities_entity", table_name="unit_entities")
op.drop_index("idx_unit_entities_unit", table_name="unit_entities")
op.drop_table("unit_entities")
op.drop_index('idx_memory_links_link_type', table_name='memory_links')
op.drop_index('idx_memory_links_entity', table_name='memory_links')
op.drop_index('idx_memory_links_to_unit', table_name='memory_links')
op.drop_index('idx_memory_links_from_unit', table_name='memory_links')
op.execute('DROP INDEX IF EXISTS idx_memory_links_unique')
op.drop_table('memory_links')
op.drop_index("idx_memory_links_link_type", table_name="memory_links")
op.drop_index("idx_memory_links_entity", table_name="memory_links")
op.drop_index("idx_memory_links_to_unit", table_name="memory_links")
op.drop_index("idx_memory_links_from_unit", table_name="memory_links")
op.execute("DROP INDEX IF EXISTS idx_memory_links_unique")
op.drop_table("memory_links")
op.drop_index('idx_entity_cooccurrences_count', table_name='entity_cooccurrences')
op.drop_index('idx_entity_cooccurrences_entity2', table_name='entity_cooccurrences')
op.drop_index('idx_entity_cooccurrences_entity1', table_name='entity_cooccurrences')
op.drop_table('entity_cooccurrences')
op.drop_index("idx_entity_cooccurrences_count", table_name="entity_cooccurrences")
op.drop_index("idx_entity_cooccurrences_entity2", table_name="entity_cooccurrences")
op.drop_index("idx_entity_cooccurrences_entity1", table_name="entity_cooccurrences")
op.drop_table("entity_cooccurrences")
# Drop BM25 materialized view and index
op.drop_index('idx_memory_units_bm25_text_vector', table_name='memory_units_bm25')
op.drop_index('idx_memory_units_bm25_bank', table_name='memory_units_bm25')
op.execute('DROP MATERIALIZED VIEW IF EXISTS memory_units_bm25')
op.drop_index("idx_memory_units_bm25_text_vector", table_name="memory_units_bm25")
op.drop_index("idx_memory_units_bm25_bank", table_name="memory_units_bm25")
op.execute("DROP MATERIALIZED VIEW IF EXISTS memory_units_bm25")
op.drop_index('idx_memory_units_embedding', table_name='memory_units')
op.drop_index('idx_memory_units_observation_date', table_name='memory_units')
op.drop_index('idx_memory_units_opinion_date', table_name='memory_units')
op.drop_index('idx_memory_units_opinion_confidence', table_name='memory_units')
op.drop_index('idx_memory_units_bank_type_date', table_name='memory_units')
op.drop_index('idx_memory_units_bank_fact_type', table_name='memory_units')
op.drop_index('idx_memory_units_fact_type', table_name='memory_units')
op.drop_index('idx_memory_units_access_count', table_name='memory_units')
op.drop_index('idx_memory_units_bank_date', table_name='memory_units')
op.drop_index('idx_memory_units_event_date', table_name='memory_units')
op.drop_index('idx_memory_units_document_id', table_name='memory_units')
op.drop_index('idx_memory_units_bank_id', table_name='memory_units')
op.execute('DROP INDEX IF EXISTS idx_memory_units_text_search')
op.drop_table('memory_units')
op.drop_index("idx_memory_units_embedding", table_name="memory_units")
op.drop_index("idx_memory_units_observation_date", table_name="memory_units")
op.drop_index("idx_memory_units_opinion_date", table_name="memory_units")
op.drop_index("idx_memory_units_opinion_confidence", table_name="memory_units")
op.drop_index("idx_memory_units_bank_type_date", table_name="memory_units")
op.drop_index("idx_memory_units_bank_fact_type", table_name="memory_units")
op.drop_index("idx_memory_units_fact_type", table_name="memory_units")
op.drop_index("idx_memory_units_access_count", table_name="memory_units")
op.drop_index("idx_memory_units_bank_date", table_name="memory_units")
op.drop_index("idx_memory_units_event_date", table_name="memory_units")
op.drop_index("idx_memory_units_document_id", table_name="memory_units")
op.drop_index("idx_memory_units_bank_id", table_name="memory_units")
op.execute("DROP INDEX IF EXISTS idx_memory_units_text_search")
op.drop_table("memory_units")
op.execute('DROP INDEX IF EXISTS idx_entities_bank_lower_name')
op.drop_index('idx_entities_bank_name', table_name='entities')
op.drop_index('idx_entities_canonical_name', table_name='entities')
op.drop_index('idx_entities_bank_id', table_name='entities')
op.drop_table('entities')
op.execute("DROP INDEX IF EXISTS idx_entities_bank_lower_name")
op.drop_index("idx_entities_bank_name", table_name="entities")
op.drop_index("idx_entities_canonical_name", table_name="entities")
op.drop_index("idx_entities_bank_id", table_name="entities")
op.drop_table("entities")
op.drop_index('idx_async_operations_bank_status', table_name='async_operations')
op.drop_index('idx_async_operations_status', table_name='async_operations')
op.drop_index('idx_async_operations_bank_id', table_name='async_operations')
op.drop_table('async_operations')
op.drop_index("idx_async_operations_bank_status", table_name="async_operations")
op.drop_index("idx_async_operations_status", table_name="async_operations")
op.drop_index("idx_async_operations_bank_id", table_name="async_operations")
op.drop_table("async_operations")
op.drop_index('idx_documents_content_hash', table_name='documents')
op.drop_index('idx_documents_bank_id', table_name='documents')
op.drop_table('documents')
op.drop_index("idx_documents_content_hash", table_name="documents")
op.drop_index("idx_documents_bank_id", table_name="documents")
op.drop_table("documents")
op.drop_table('banks')
op.drop_table("banks")
# Drop extensions (optional - comment out if you want to keep them)
# op.execute('DROP EXTENSION IF EXISTS vector')
@@ -5,18 +5,18 @@ Revises: 5a366d414dce
Create Date: 2025-11-28 00:00:00.000000
"""
from typing import Sequence, Union
from alembic import op
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
from sqlalchemy.dialects import postgresql
# revision identifiers, used by Alembic.
revision: str = 'b7c4d8e9f1a2'
down_revision: Union[str, Sequence[str], None] = '5a366d414dce'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
revision: str = "b7c4d8e9f1a2"
down_revision: str | Sequence[str] | None = "5a366d414dce"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
@@ -24,47 +24,47 @@ def upgrade() -> None:
# Create chunks table with single text PK (bank_id_document_id_chunk_index)
op.create_table(
'chunks',
sa.Column('chunk_id', sa.Text(), nullable=False),
sa.Column('document_id', sa.Text(), nullable=False),
sa.Column('bank_id', sa.Text(), nullable=False),
sa.Column('chunk_index', sa.Integer(), nullable=False),
sa.Column('chunk_text', sa.Text(), nullable=False),
sa.Column('created_at', postgresql.TIMESTAMP(timezone=True), server_default=sa.text('now()'), nullable=False),
sa.ForeignKeyConstraint(['document_id', 'bank_id'], ['documents.id', 'documents.bank_id'], name='chunks_document_fkey', ondelete='CASCADE'),
sa.PrimaryKeyConstraint('chunk_id', name=op.f('pk_chunks'))
"chunks",
sa.Column("chunk_id", sa.Text(), nullable=False),
sa.Column("document_id", sa.Text(), nullable=False),
sa.Column("bank_id", sa.Text(), nullable=False),
sa.Column("chunk_index", sa.Integer(), nullable=False),
sa.Column("chunk_text", sa.Text(), nullable=False),
sa.Column("created_at", postgresql.TIMESTAMP(timezone=True), server_default=sa.text("now()"), nullable=False),
sa.ForeignKeyConstraint(
["document_id", "bank_id"],
["documents.id", "documents.bank_id"],
name="chunks_document_fkey",
ondelete="CASCADE",
),
sa.PrimaryKeyConstraint("chunk_id", name=op.f("pk_chunks")),
)
# Add indexes for efficient queries
op.create_index('idx_chunks_document_id', 'chunks', ['document_id'])
op.create_index('idx_chunks_bank_id', 'chunks', ['bank_id'])
op.create_index("idx_chunks_document_id", "chunks", ["document_id"])
op.create_index("idx_chunks_bank_id", "chunks", ["bank_id"])
# Add chunk_id column to memory_units (nullable, as existing records won't have chunks)
op.add_column('memory_units', sa.Column('chunk_id', sa.Text(), nullable=True))
op.add_column("memory_units", sa.Column("chunk_id", sa.Text(), nullable=True))
# Add foreign key constraint to chunks table
op.create_foreign_key(
'memory_units_chunk_fkey',
'memory_units',
'chunks',
['chunk_id'],
['chunk_id'],
ondelete='SET NULL'
"memory_units_chunk_fkey", "memory_units", "chunks", ["chunk_id"], ["chunk_id"], ondelete="SET NULL"
)
# Add index on chunk_id for efficient lookups
op.create_index('idx_memory_units_chunk_id', 'memory_units', ['chunk_id'])
op.create_index("idx_memory_units_chunk_id", "memory_units", ["chunk_id"])
def downgrade() -> None:
"""Remove chunks table and chunk_id from memory_units."""
# Drop index and foreign key from memory_units
op.drop_index('idx_memory_units_chunk_id', table_name='memory_units')
op.drop_constraint('memory_units_chunk_fkey', 'memory_units', type_='foreignkey')
op.drop_column('memory_units', 'chunk_id')
op.drop_index("idx_memory_units_chunk_id", table_name="memory_units")
op.drop_constraint("memory_units_chunk_fkey", "memory_units", type_="foreignkey")
op.drop_column("memory_units", "chunk_id")
# Drop chunks table indexes and table
op.drop_index('idx_chunks_bank_id', table_name='chunks')
op.drop_index('idx_chunks_document_id', table_name='chunks')
op.drop_table('chunks')
op.drop_index("idx_chunks_bank_id", table_name="chunks")
op.drop_index("idx_chunks_document_id", table_name="chunks")
op.drop_table("chunks")
@@ -5,35 +5,35 @@ Revises: b7c4d8e9f1a2
Create Date: 2025-12-02 00:00:00.000000
"""
from typing import Sequence, Union
from alembic import op
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
from sqlalchemy.dialects import postgresql
# revision identifiers, used by Alembic.
revision: str = 'c8e5f2a3b4d1'
down_revision: Union[str, Sequence[str], None] = 'b7c4d8e9f1a2'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
revision: str = "c8e5f2a3b4d1"
down_revision: str | Sequence[str] | None = "b7c4d8e9f1a2"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
"""Add retain_params JSONB column to documents table."""
# Add retain_params column to store parameters passed during retain
op.add_column('documents', sa.Column('retain_params', postgresql.JSONB(), nullable=True))
op.add_column("documents", sa.Column("retain_params", postgresql.JSONB(), nullable=True))
# Add index for efficient queries on retain_params
op.create_index('idx_documents_retain_params', 'documents', ['retain_params'], postgresql_using='gin')
op.create_index("idx_documents_retain_params", "documents", ["retain_params"], postgresql_using="gin")
def downgrade() -> None:
"""Remove retain_params column from documents table."""
# Drop index
op.drop_index('idx_documents_retain_params', table_name='documents')
op.drop_index("idx_documents_retain_params", table_name="documents")
# Drop column
op.drop_column('documents', 'retain_params')
op.drop_column("documents", "retain_params")
@@ -5,20 +5,19 @@ Revises: c8e5f2a3b4d1
Create Date: 2024-12-04 15:00:00.000000
"""
from alembic import op
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision = 'd9f6a3b4c5e2'
down_revision = 'c8e5f2a3b4d1'
revision = "d9f6a3b4c5e2"
down_revision = "c8e5f2a3b4d1"
branch_labels = None
depends_on = None
def upgrade():
# Drop old check constraint FIRST (before updating data)
op.drop_constraint('memory_units_fact_type_check', 'memory_units', type_='check')
op.drop_constraint("memory_units_fact_type_check", "memory_units", type_="check")
# Update existing 'bank' values to 'experience'
op.execute("UPDATE memory_units SET fact_type = 'experience' WHERE fact_type = 'bank'")
@@ -27,22 +26,18 @@ def upgrade():
# Create new check constraint with 'experience' instead of 'bank'
op.create_check_constraint(
'memory_units_fact_type_check',
'memory_units',
"fact_type IN ('world', 'experience', 'opinion', 'observation')"
"memory_units_fact_type_check", "memory_units", "fact_type IN ('world', 'experience', 'opinion', 'observation')"
)
def downgrade():
# Drop new check constraint FIRST
op.drop_constraint('memory_units_fact_type_check', 'memory_units', type_='check')
op.drop_constraint("memory_units_fact_type_check", "memory_units", type_="check")
# Update 'experience' back to 'bank'
op.execute("UPDATE memory_units SET fact_type = 'bank' WHERE fact_type = 'experience'")
# Recreate old check constraint
op.create_check_constraint(
'memory_units_fact_type_check',
'memory_units',
"fact_type IN ('world', 'bank', 'opinion', 'observation')"
"memory_units_fact_type_check", "memory_units", "fact_type IN ('world', 'bank', 'opinion', 'observation')"
)
@@ -8,17 +8,17 @@ Migrate disposition traits from Big Five (openness, conscientiousness, extravers
agreeableness, neuroticism, bias_strength with 0-1 float values) to the new 3-trait
system (skepticism, literalism, empathy with 1-5 integer values).
"""
from typing import Sequence, Union
from alembic import op
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
# revision identifiers, used by Alembic.
revision: str = 'e0a1b2c3d4e5'
down_revision: Union[str, Sequence[str], None] = 'rename_personality'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
revision: str = "e0a1b2c3d4e5"
down_revision: str | Sequence[str] | None = "rename_personality"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
@@ -31,17 +31,21 @@ def upgrade() -> None:
# - literalism: derived from conscientiousness (detail-oriented people are more literal)
# - empathy: derived from agreeableness + inverse of neuroticism
# Default all to 3 (neutral) for simplicity
conn.execute(sa.text("""
conn.execute(
sa.text("""
UPDATE banks
SET disposition = '{"skepticism": 3, "literalism": 3, "empathy": 3}'::jsonb
WHERE disposition IS NOT NULL
"""))
""")
)
# Update the default for new banks
conn.execute(sa.text("""
conn.execute(
sa.text("""
ALTER TABLE banks
ALTER COLUMN disposition SET DEFAULT '{"skepticism": 3, "literalism": 3, "empathy": 3}'::jsonb
"""))
""")
)
def downgrade() -> None:
@@ -49,14 +53,18 @@ def downgrade() -> None:
conn = op.get_bind()
# Revert to Big Five format with default values
conn.execute(sa.text("""
conn.execute(
sa.text("""
UPDATE banks
SET disposition = '{"openness": 0.5, "conscientiousness": 0.5, "extraversion": 0.5, "agreeableness": 0.5, "neuroticism": 0.5, "bias_strength": 0.5}'::jsonb
WHERE disposition IS NOT NULL
"""))
""")
)
# Update the default for new banks
conn.execute(sa.text("""
conn.execute(
sa.text("""
ALTER TABLE banks
ALTER COLUMN disposition SET DEFAULT '{"openness": 0.5, "conscientiousness": 0.5, "extraversion": 0.5, "agreeableness": 0.5, "neuroticism": 0.5, "bias_strength": 0.5}'::jsonb
"""))
""")
)
@@ -5,18 +5,18 @@ Revises: d9f6a3b4c5e2
Create Date: 2024-12-04
"""
from typing import Sequence, Union
from alembic import op
from collections.abc import Sequence
import sqlalchemy as sa
from alembic import op
from sqlalchemy.dialects import postgresql
# revision identifiers, used by Alembic.
revision: str = 'rename_personality'
down_revision: Union[str, Sequence[str], None] = 'd9f6a3b4c5e2'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
revision: str = "rename_personality"
down_revision: str | Sequence[str] | None = "d9f6a3b4c5e2"
branch_labels: str | Sequence[str] | None = None
depends_on: str | Sequence[str] | None = None
def upgrade() -> None:
@@ -24,42 +24,51 @@ def upgrade() -> None:
conn = op.get_bind()
# Check if 'personality' column exists (old database)
result = conn.execute(sa.text("""
result = conn.execute(
sa.text("""
SELECT column_name
FROM information_schema.columns
WHERE table_name = 'banks' AND column_name = 'personality'
"""))
""")
)
has_personality = result.fetchone() is not None
# Check if 'disposition' column exists (new database)
result = conn.execute(sa.text("""
result = conn.execute(
sa.text("""
SELECT column_name
FROM information_schema.columns
WHERE table_name = 'banks' AND column_name = 'disposition'
"""))
""")
)
has_disposition = result.fetchone() is not None
if has_personality and not has_disposition:
# Old database: rename personality -> disposition
op.alter_column('banks', 'personality', new_column_name='disposition')
op.alter_column("banks", "personality", new_column_name="disposition")
elif not has_personality and not has_disposition:
# Neither exists (shouldn't happen, but be safe): add disposition column
op.add_column('banks', sa.Column(
'disposition',
postgresql.JSONB(astext_type=sa.Text()),
server_default=sa.text("'{}'::jsonb"),
nullable=False
))
op.add_column(
"banks",
sa.Column(
"disposition",
postgresql.JSONB(astext_type=sa.Text()),
server_default=sa.text("'{}'::jsonb"),
nullable=False,
),
)
# else: disposition already exists, nothing to do
def downgrade() -> None:
"""Revert disposition column back to personality."""
conn = op.get_bind()
result = conn.execute(sa.text("""
result = conn.execute(
sa.text("""
SELECT column_name
FROM information_schema.columns
WHERE table_name = 'banks' AND column_name = 'disposition'
"""))
""")
)
if result.fetchone():
op.alter_column('banks', 'disposition', new_column_name='personality')
op.alter_column("banks", "disposition", new_column_name="personality")
+12 -12
View File
@@ -3,8 +3,10 @@ Unified API module for Hindsight.
Provides both HTTP REST API and MCP (Model Context Protocol) server.
"""
import logging
from typing import Optional
from fastapi import FastAPI
from hindsight_api import MemoryEngine
@@ -17,7 +19,7 @@ def create_app(
http_api_enabled: bool = True,
mcp_api_enabled: bool = False,
mcp_mount_path: str = "/mcp",
initialize_memory: bool = True
initialize_memory: bool = True,
) -> FastAPI:
"""
Create and configure the unified Hindsight API application.
@@ -47,10 +49,8 @@ def create_app(
# Import and create HTTP API if enabled
if http_api_enabled:
from .http import create_app as create_http_app
app = create_http_app(
memory=memory,
initialize_memory=initialize_memory
)
app = create_http_app(memory=memory, initialize_memory=initialize_memory)
logger.info("HTTP REST API enabled")
else:
# Create minimal FastAPI app
@@ -77,15 +77,15 @@ def create_app(
# Re-export commonly used items for backwards compatibility
from .http import (
RecallRequest,
RecallResult,
RecallResponse,
MemoryItem,
RetainRequest,
ReflectRequest,
ReflectResponse,
CreateBankRequest,
DispositionTraits,
MemoryItem,
RecallRequest,
RecallResponse,
RecallResult,
ReflectRequest,
ReflectResponse,
RetainRequest,
)
__all__ = [
File diff suppressed because it is too large Load Diff
+30 -28
View File
@@ -4,27 +4,33 @@ import json
import logging
import os
from contextvars import ContextVar
from typing import Optional
from fastmcp import FastMCP
from hindsight_api import MemoryEngine
from hindsight_api.engine.response_models import VALID_RECALL_FACT_TYPES
# Configure logging from HINDSIGHT_API_LOG_LEVEL environment variable
_log_level_str = os.environ.get("HINDSIGHT_API_LOG_LEVEL", "info").lower()
_log_level_map = {"critical": logging.CRITICAL, "error": logging.ERROR, "warning": logging.WARNING,
"info": logging.INFO, "debug": logging.DEBUG, "trace": logging.DEBUG}
_log_level_map = {
"critical": logging.CRITICAL,
"error": logging.ERROR,
"warning": logging.WARNING,
"info": logging.INFO,
"debug": logging.DEBUG,
"trace": logging.DEBUG,
}
logging.basicConfig(
level=_log_level_map.get(_log_level_str, logging.INFO),
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s"
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s",
)
logger = logging.getLogger(__name__)
# Context variable to hold the current bank_id from the URL path
_current_bank_id: ContextVar[Optional[str]] = ContextVar("current_bank_id", default=None)
_current_bank_id: ContextVar[str | None] = ContextVar("current_bank_id", default=None)
def get_current_bank_id() -> Optional[str]:
def get_current_bank_id() -> str | None:
"""Get the current bank_id from context (set from URL path)."""
return _current_bank_id.get()
@@ -61,10 +67,7 @@ def create_mcp_server(memory: MemoryEngine) -> FastMCP:
"""
try:
bank_id = get_current_bank_id()
await memory.retain_batch_async(
bank_id=bank_id,
contents=[{"content": content, "context": context}]
)
await memory.retain_batch_async(bank_id=bank_id, contents=[{"content": content, "context": context}])
return "Memory stored successfully"
except Exception as e:
logger.error(f"Error storing memory: {e}", exc_info=True)
@@ -88,11 +91,9 @@ def create_mcp_server(memory: MemoryEngine) -> FastMCP:
try:
bank_id = get_current_bank_id()
from hindsight_api.engine.memory_engine import Budget
search_result = await memory.recall_async(
bank_id=bank_id,
query=query,
fact_type=list(VALID_RECALL_FACT_TYPES),
budget=Budget.LOW
bank_id=bank_id, query=query, fact_type=list(VALID_RECALL_FACT_TYPES), budget=Budget.LOW
)
results = [
@@ -133,7 +134,7 @@ class MCPMiddleware:
# Strip any mount prefix (e.g., /mcp) that FastAPI might not have stripped
root_path = scope.get("root_path", "")
if root_path and path.startswith(root_path):
path = path[len(root_path):] or "/"
path = path[len(root_path) :] or "/"
# Also handle case where mount path wasn't stripped (e.g., /mcp/...)
if path.startswith("/mcp/"):
@@ -169,10 +170,7 @@ class MCPMiddleware:
body = message.get("body", b"")
if body and b"/messages" in body:
# Rewrite /messages to /{bank_id}/messages in SSE endpoint event
body = body.replace(
b"data: /messages",
f"data: /{bank_id}/messages".encode()
)
body = body.replace(b"data: /messages", f"data: /{bank_id}/messages".encode())
message = {**message, "body": body}
await send(message)
@@ -183,15 +181,19 @@ class MCPMiddleware:
async def _send_error(self, send, status: int, message: str):
"""Send an error response."""
body = json.dumps({"error": message}).encode()
await send({
"type": "http.response.start",
"status": status,
"headers": [(b"content-type", b"application/json")],
})
await send({
"type": "http.response.body",
"body": body,
})
await send(
{
"type": "http.response.start",
"status": status,
"headers": [(b"content-type", b"application/json")],
}
)
await send(
{
"type": "http.response.body",
"body": body,
}
)
def create_mcp_app(memory: MemoryEngine):
-1
View File
@@ -28,7 +28,6 @@ def _interpolate_color(start: tuple, end: tuple, t: float) -> tuple:
def gradient_text(text: str, start: tuple = GRADIENT_START, end: tuple = GRADIENT_END) -> str:
"""Render text with a gradient color effect."""
result = []
length = len(text)
for i, char in enumerate(text):
+7 -13
View File
@@ -3,10 +3,10 @@ Centralized configuration for Hindsight API.
All environment variables and their defaults are defined here.
"""
import logging
import os
from dataclasses import dataclass
from typing import Optional
import logging
logger = logging.getLogger(__name__)
@@ -63,19 +63,19 @@ class HindsightConfig:
# LLM
llm_provider: str
llm_api_key: Optional[str]
llm_api_key: str | None
llm_model: str
llm_base_url: Optional[str]
llm_base_url: str | None
# Embeddings
embeddings_provider: str
embeddings_local_model: str
embeddings_tei_url: Optional[str]
embeddings_tei_url: str | None
# Reranker
reranker_provider: str
reranker_local_model: str
reranker_tei_url: Optional[str]
reranker_tei_url: str | None
# Server
host: str
@@ -92,29 +92,24 @@ class HindsightConfig:
return cls(
# Database
database_url=os.getenv(ENV_DATABASE_URL, DEFAULT_DATABASE_URL),
# LLM
llm_provider=os.getenv(ENV_LLM_PROVIDER, DEFAULT_LLM_PROVIDER),
llm_api_key=os.getenv(ENV_LLM_API_KEY),
llm_model=os.getenv(ENV_LLM_MODEL, DEFAULT_LLM_MODEL),
llm_base_url=os.getenv(ENV_LLM_BASE_URL) or None,
# Embeddings
embeddings_provider=os.getenv(ENV_EMBEDDINGS_PROVIDER, DEFAULT_EMBEDDINGS_PROVIDER),
embeddings_local_model=os.getenv(ENV_EMBEDDINGS_LOCAL_MODEL, DEFAULT_EMBEDDINGS_LOCAL_MODEL),
embeddings_tei_url=os.getenv(ENV_EMBEDDINGS_TEI_URL),
# Reranker
reranker_provider=os.getenv(ENV_RERANKER_PROVIDER, DEFAULT_RERANKER_PROVIDER),
reranker_local_model=os.getenv(ENV_RERANKER_LOCAL_MODEL, DEFAULT_RERANKER_LOCAL_MODEL),
reranker_tei_url=os.getenv(ENV_RERANKER_TEI_URL),
# Server
host=os.getenv(ENV_HOST, DEFAULT_HOST),
port=int(os.getenv(ENV_PORT, DEFAULT_PORT)),
log_level=os.getenv(ENV_LOG_LEVEL, DEFAULT_LOG_LEVEL),
mcp_enabled=os.getenv(ENV_MCP_ENABLED, str(DEFAULT_MCP_ENABLED)).lower() == "true",
# Recall
graph_retriever=os.getenv(ENV_GRAPH_RETRIEVER, DEFAULT_GRAPH_RETRIEVER),
)
@@ -147,8 +142,7 @@ class HindsightConfig:
def configure_logging(self) -> None:
"""Configure Python logging based on the log level."""
logging.basicConfig(
level=self.get_python_log_level(),
format="%(asctime)s - %(levelname)s - %(name)s - %(message)s"
level=self.get_python_log_level(), format="%(asctime)s - %(levelname)s - %(name)s - %(message)s"
)
def log_config(self) -> None:
@@ -7,24 +7,24 @@ This package contains all the implementation details of the memory engine:
- Supporting modules: embeddings, cross_encoder, entity_resolver, etc.
"""
from .memory_engine import MemoryEngine
from .cross_encoder import CrossEncoderModel, LocalSTCrossEncoder, RemoteTEICrossEncoder
from .db_utils import acquire_with_retry
from .embeddings import Embeddings, LocalSTEmbeddings, RemoteTEIEmbeddings
from .cross_encoder import CrossEncoderModel, LocalSTCrossEncoder, RemoteTEICrossEncoder
from .llm_wrapper import LLMConfig
from .memory_engine import MemoryEngine
from .response_models import MemoryFact, RecallResult, ReflectResult
from .search.trace import (
SearchTrace,
QueryInfo,
EntryPoint,
NodeVisit,
WeightComponents,
LinkInfo,
NodeVisit,
PruningDecision,
SearchSummary,
QueryInfo,
SearchPhaseMetrics,
SearchSummary,
SearchTrace,
WeightComponents,
)
from .search.tracer import SearchTracer
from .llm_wrapper import LLMConfig
from .response_models import RecallResult, ReflectResult, MemoryFact
__all__ = [
"MemoryEngine",
@@ -5,19 +5,19 @@ Provides an interface for reranking with different backends.
Configuration via environment variables - see hindsight_api.config for all env var names.
"""
from abc import ABC, abstractmethod
from typing import List, Tuple, Optional
import logging
import os
from abc import ABC, abstractmethod
import httpx
from ..config import (
ENV_RERANKER_PROVIDER,
ENV_RERANKER_LOCAL_MODEL,
ENV_RERANKER_TEI_URL,
DEFAULT_RERANKER_PROVIDER,
DEFAULT_RERANKER_LOCAL_MODEL,
DEFAULT_RERANKER_PROVIDER,
ENV_RERANKER_LOCAL_MODEL,
ENV_RERANKER_PROVIDER,
ENV_RERANKER_TEI_URL,
)
logger = logging.getLogger(__name__)
@@ -47,7 +47,7 @@ class CrossEncoderModel(ABC):
pass
@abstractmethod
def predict(self, pairs: List[Tuple[str, str]]) -> List[float]:
def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Score query-document pairs for relevance.
@@ -72,7 +72,7 @@ class LocalSTCrossEncoder(CrossEncoderModel):
- Trained for passage re-ranking
"""
def __init__(self, model_name: Optional[str] = None):
def __init__(self, model_name: str | None = None):
"""
Initialize local SentenceTransformers cross-encoder.
@@ -104,7 +104,7 @@ class LocalSTCrossEncoder(CrossEncoderModel):
self._model = CrossEncoder(self.model_name)
logger.info("Reranker: local provider initialized")
def predict(self, pairs: List[Tuple[str, str]]) -> List[float]:
def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Score query-document pairs for relevance.
@@ -117,7 +117,7 @@ class LocalSTCrossEncoder(CrossEncoderModel):
if self._model is None:
raise RuntimeError("Reranker not initialized. Call initialize() first.")
scores = 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)
class RemoteTEICrossEncoder(CrossEncoderModel):
@@ -153,8 +153,8 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
self.batch_size = batch_size
self.max_retries = max_retries
self.retry_delay = retry_delay
self._client: Optional[httpx.Client] = None
self._model_id: Optional[str] = None
self._client: httpx.Client | None = None
self._model_id: str | None = None
@property
def provider_name(self) -> str:
@@ -163,6 +163,7 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
def _request_with_retry(self, method: str, url: str, **kwargs) -> httpx.Response:
"""Make an HTTP request with automatic retries on transient errors."""
import time
last_error = None
delay = self.retry_delay
@@ -177,14 +178,18 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
except (httpx.ConnectError, httpx.ReadTimeout, httpx.WriteTimeout) as e:
last_error = e
if attempt < self.max_retries:
logger.warning(f"TEI request failed (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s...")
logger.warning(
f"TEI request failed (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s..."
)
time.sleep(delay)
delay *= 2 # Exponential backoff
except httpx.HTTPStatusError as e:
# Retry on 5xx server errors
if e.response.status_code >= 500 and attempt < self.max_retries:
last_error = e
logger.warning(f"TEI server error (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s...")
logger.warning(
f"TEI server error (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s..."
)
time.sleep(delay)
delay *= 2
else:
@@ -209,7 +214,7 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
except httpx.HTTPError as e:
raise RuntimeError(f"Failed to connect to TEI server at {self.base_url}: {e}")
def predict(self, pairs: List[Tuple[str, str]]) -> List[float]:
def predict(self, pairs: list[tuple[str, str]]) -> list[float]:
"""
Score query-document pairs using the remote TEI reranker.
@@ -229,7 +234,7 @@ class RemoteTEICrossEncoder(CrossEncoderModel):
# Process in batches
for i in range(0, len(pairs), self.batch_size):
batch = pairs[i:i + self.batch_size]
batch = pairs[i : i + self.batch_size]
# TEI rerank endpoint expects query and texts separately
# All pairs in a batch should have the same query for optimal performance
@@ -287,15 +292,11 @@ def create_cross_encoder_from_env() -> CrossEncoderModel:
if provider == "tei":
url = os.environ.get(ENV_RERANKER_TEI_URL)
if not url:
raise ValueError(
f"{ENV_RERANKER_TEI_URL} is required when {ENV_RERANKER_PROVIDER} is 'tei'"
)
raise ValueError(f"{ENV_RERANKER_TEI_URL} is required when {ENV_RERANKER_PROVIDER} is 'tei'")
return RemoteTEICrossEncoder(base_url=url)
elif provider == "local":
model = os.environ.get(ENV_RERANKER_LOCAL_MODEL)
model_name = model or DEFAULT_RERANKER_LOCAL_MODEL
return LocalSTCrossEncoder(model_name=model_name)
else:
raise ValueError(
f"Unknown reranker provider: {provider}. Supported: 'local', 'tei'"
)
raise ValueError(f"Unknown reranker provider: {provider}. Supported: 'local', 'tei'")
@@ -1,9 +1,11 @@
"""
Database utility functions for connection management with retry logic.
"""
import asyncio
import logging
from contextlib import asynccontextmanager
import asyncpg
logger = logging.getLogger(__name__)
@@ -54,16 +56,14 @@ async def retry_with_backoff(
except retryable_exceptions as e:
last_exception = e
if attempt < max_retries:
delay = min(base_delay * (2 ** attempt), max_delay)
delay = min(base_delay * (2**attempt), max_delay)
logger.warning(
f"Database operation failed (attempt {attempt + 1}/{max_retries + 1}): {e}. "
f"Retrying in {delay:.1f}s..."
)
await asyncio.sleep(delay)
else:
logger.error(
f"Database operation failed after {max_retries + 1} attempts: {e}"
)
logger.error(f"Database operation failed after {max_retries + 1} attempts: {e}")
raise last_exception
@@ -83,6 +83,7 @@ async def acquire_with_retry(pool: asyncpg.Pool, max_retries: int = DEFAULT_MAX_
Yields:
An asyncpg connection
"""
async def acquire():
return await pool.acquire()
@@ -8,20 +8,20 @@ the database schema (pgvector column defined as vector(384)).
Configuration via environment variables - see hindsight_api.config for all env var names.
"""
from abc import ABC, abstractmethod
from typing import List, Optional
import logging
import os
from abc import ABC, abstractmethod
import httpx
from ..config import (
ENV_EMBEDDINGS_PROVIDER,
ENV_EMBEDDINGS_LOCAL_MODEL,
ENV_EMBEDDINGS_TEI_URL,
DEFAULT_EMBEDDINGS_PROVIDER,
DEFAULT_EMBEDDINGS_LOCAL_MODEL,
DEFAULT_EMBEDDINGS_PROVIDER,
EMBEDDING_DIMENSION,
ENV_EMBEDDINGS_LOCAL_MODEL,
ENV_EMBEDDINGS_PROVIDER,
ENV_EMBEDDINGS_TEI_URL,
)
logger = logging.getLogger(__name__)
@@ -52,7 +52,7 @@ class Embeddings(ABC):
pass
@abstractmethod
def encode(self, texts: List[str]) -> List[List[float]]:
def encode(self, texts: list[str]) -> list[list[float]]:
"""
Generate 384-dimensional embeddings for a list of texts.
@@ -75,7 +75,7 @@ class LocalSTEmbeddings(Embeddings):
embeddings matching the database schema.
"""
def __init__(self, model_name: Optional[str] = None):
def __init__(self, model_name: str | None = None):
"""
Initialize local SentenceTransformers embeddings.
@@ -123,7 +123,7 @@ class LocalSTEmbeddings(Embeddings):
logger.info(f"Embeddings: local provider initialized (dim: {model_dim})")
def encode(self, texts: List[str]) -> List[List[float]]:
def encode(self, texts: list[str]) -> list[list[float]]:
"""
Generate 384-dimensional embeddings for a list of texts.
@@ -172,8 +172,8 @@ class RemoteTEIEmbeddings(Embeddings):
self.batch_size = batch_size
self.max_retries = max_retries
self.retry_delay = retry_delay
self._client: Optional[httpx.Client] = None
self._model_id: Optional[str] = None
self._client: httpx.Client | None = None
self._model_id: str | None = None
@property
def provider_name(self) -> str:
@@ -182,6 +182,7 @@ class RemoteTEIEmbeddings(Embeddings):
def _request_with_retry(self, method: str, url: str, **kwargs) -> httpx.Response:
"""Make an HTTP request with automatic retries on transient errors."""
import time
last_error = None
delay = self.retry_delay
@@ -196,14 +197,18 @@ class RemoteTEIEmbeddings(Embeddings):
except (httpx.ConnectError, httpx.ReadTimeout, httpx.WriteTimeout) as e:
last_error = e
if attempt < self.max_retries:
logger.warning(f"TEI request failed (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s...")
logger.warning(
f"TEI request failed (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s..."
)
time.sleep(delay)
delay *= 2 # Exponential backoff
except httpx.HTTPStatusError as e:
# Retry on 5xx server errors
if e.response.status_code >= 500 and attempt < self.max_retries:
last_error = e
logger.warning(f"TEI server error (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s...")
logger.warning(
f"TEI server error (attempt {attempt + 1}/{self.max_retries + 1}): {e}. Retrying in {delay}s..."
)
time.sleep(delay)
delay *= 2
else:
@@ -228,7 +233,7 @@ class RemoteTEIEmbeddings(Embeddings):
except httpx.HTTPError as e:
raise RuntimeError(f"Failed to connect to TEI server at {self.base_url}: {e}")
def encode(self, texts: List[str]) -> List[List[float]]:
def encode(self, texts: list[str]) -> list[list[float]]:
"""
Generate embeddings using the remote TEI server.
@@ -248,7 +253,7 @@ class RemoteTEIEmbeddings(Embeddings):
# Process in batches
for i in range(0, len(texts), self.batch_size):
batch = texts[i:i + self.batch_size]
batch = texts[i : i + self.batch_size]
try:
response = self._request_with_retry(
@@ -278,15 +283,11 @@ def create_embeddings_from_env() -> Embeddings:
if provider == "tei":
url = os.environ.get(ENV_EMBEDDINGS_TEI_URL)
if not url:
raise ValueError(
f"{ENV_EMBEDDINGS_TEI_URL} is required when {ENV_EMBEDDINGS_PROVIDER} is 'tei'"
)
raise ValueError(f"{ENV_EMBEDDINGS_TEI_URL} is required when {ENV_EMBEDDINGS_PROVIDER} is 'tei'")
return RemoteTEIEmbeddings(base_url=url)
elif provider == "local":
model = os.environ.get(ENV_EMBEDDINGS_LOCAL_MODEL)
model_name = model or DEFAULT_EMBEDDINGS_LOCAL_MODEL
return LocalSTEmbeddings(model_name=model_name)
else:
raise ValueError(
f"Unknown embeddings provider: {provider}. Supported: 'local', 'tei'"
)
raise ValueError(f"Unknown embeddings provider: {provider}. Supported: 'local', 'tei'")
@@ -4,12 +4,13 @@ Entity extraction and resolution for memory system.
Uses spaCy for entity extraction and implements resolution logic
to disambiguate entities across memory units.
"""
import asyncpg
from typing import List, Dict, Optional, Set, Any
from difflib import SequenceMatcher
from datetime import datetime, timezone
from .db_utils import acquire_with_retry
from datetime import UTC, datetime
from difflib import SequenceMatcher
import asyncpg
from .db_utils import acquire_with_retry
# Load spaCy model (singleton)
_nlp = None
@@ -32,11 +33,11 @@ class EntityResolver:
async def resolve_entities_batch(
self,
bank_id: str,
entities_data: List[Dict],
entities_data: list[dict],
context: str,
unit_event_date,
conn=None,
) -> List[str]:
) -> list[str]:
"""
Resolve multiple entities in batch (MUCH faster than sequential).
@@ -62,7 +63,9 @@ class EntityResolver:
else:
return await self._resolve_entities_batch_impl(conn, bank_id, entities_data, context, unit_event_date)
async def _resolve_entities_batch_impl(self, conn, bank_id: str, entities_data: List[Dict], context: str, unit_event_date) -> List[str]:
async def _resolve_entities_batch_impl(
self, conn, bank_id: str, entities_data: list[dict], context: str, unit_event_date
) -> list[str]:
# Query ALL candidates for this bank
all_entities = await conn.fetch(
"""
@@ -70,11 +73,11 @@ class EntityResolver:
FROM entities
WHERE bank_id = $1
""",
bank_id
bank_id,
)
# Build entity ID to name mapping for co-occurrence lookups
entity_id_to_name = {row['id']: row['canonical_name'].lower() for row in all_entities}
entity_id_to_name = {row["id"]: row["canonical_name"].lower() for row in all_entities}
# Query ALL co-occurrences for this bank's entities in one query
# This builds a map of entity_id -> set of co-occurring entity names
@@ -85,13 +88,13 @@ class EntityResolver:
WHERE ec.entity_id_1 IN (SELECT id FROM entities WHERE bank_id = $1)
OR ec.entity_id_2 IN (SELECT id FROM entities WHERE bank_id = $1)
""",
bank_id
bank_id,
)
# Build co-occurrence map: entity_id -> set of co-occurring entity names (lowercase)
cooccurrence_map: Dict[str, Set[str]] = {}
cooccurrence_map: dict[str, set[str]] = {}
for row in all_cooccurrences:
eid1, eid2 = row['entity_id_1'], row['entity_id_2']
eid1, eid2 = row["entity_id_1"], row["entity_id_2"]
# Add both directions
if eid1 not in cooccurrence_map:
cooccurrence_map[eid1] = set()
@@ -105,22 +108,24 @@ class EntityResolver:
# Build candidate map for each entity text
all_candidates = {} # Maps entity_text -> list of candidates
entity_texts = list(set(e['text'] for e in entities_data))
entity_texts = list(set(e["text"] for e in entities_data))
for entity_text in entity_texts:
matching = []
entity_text_lower = entity_text.lower()
for row in all_entities:
canonical_name = row['canonical_name']
ent_id = row['id']
metadata = row['metadata']
last_seen = row['last_seen']
mention_count = row['mention_count']
canonical_name = row["canonical_name"]
ent_id = row["id"]
metadata = row["metadata"]
last_seen = row["last_seen"]
mention_count = row["mention_count"]
canonical_lower = canonical_name.lower()
# Match if exact or substring match
if (entity_text_lower == canonical_lower or
entity_text_lower in canonical_lower or
canonical_lower in entity_text_lower):
if (
entity_text_lower == canonical_lower
or entity_text_lower in canonical_lower
or canonical_lower in entity_text_lower
):
matching.append((ent_id, canonical_name, metadata, last_seen, mention_count))
all_candidates[entity_text] = matching
@@ -130,10 +135,10 @@ class EntityResolver:
entities_to_create = [] # (idx, entity_data, event_date)
for idx, entity_data in enumerate(entities_data):
entity_text = entity_data['text']
nearby_entities = entity_data.get('nearby_entities', [])
entity_text = entity_data["text"]
nearby_entities = entity_data.get("nearby_entities", [])
# Use per-entity date if available, otherwise fall back to batch-level date
entity_event_date = entity_data.get('event_date', unit_event_date)
entity_event_date = entity_data.get("event_date", unit_event_date)
candidates = all_candidates.get(entity_text, [])
@@ -146,17 +151,13 @@ class EntityResolver:
best_candidate = None
best_score = 0.0
nearby_entity_set = {e['text'].lower() for e in nearby_entities if e['text'] != entity_text}
nearby_entity_set = {e["text"].lower() for e in nearby_entities if e["text"] != entity_text}
for candidate_id, canonical_name, metadata, last_seen, mention_count in candidates:
score = 0.0
# 1. Name similarity (0-0.5)
name_similarity = SequenceMatcher(
None,
entity_text.lower(),
canonical_name.lower()
).ratio()
name_similarity = SequenceMatcher(None, entity_text.lower(), canonical_name.lower()).ratio()
score += name_similarity * 0.5
# 2. Co-occurring entities (0-0.3)
@@ -169,8 +170,10 @@ class EntityResolver:
# 3. Temporal proximity (0-0.2)
if last_seen and entity_event_date:
# Normalize timezone awareness for comparison
event_date_utc = entity_event_date if entity_event_date.tzinfo else entity_event_date.replace(tzinfo=timezone.utc)
last_seen_utc = last_seen if last_seen.tzinfo else last_seen.replace(tzinfo=timezone.utc)
event_date_utc = (
entity_event_date if entity_event_date.tzinfo else entity_event_date.replace(tzinfo=UTC)
)
last_seen_utc = last_seen if last_seen.tzinfo else last_seen.replace(tzinfo=UTC)
days_diff = abs((event_date_utc - last_seen_utc).total_seconds() / 86400)
if days_diff < 7:
temporal_score = max(0, 1.0 - (days_diff / 7))
@@ -198,7 +201,7 @@ class EntityResolver:
last_seen = $2
WHERE id = $1::uuid
""",
entities_to_update
entities_to_update,
)
# Batch create new entities using COPY + INSERT for maximum speed
@@ -208,7 +211,7 @@ class EntityResolver:
# For duplicates, we only insert once and reuse the ID
unique_entities = {} # lowercase_name -> (entity_data, event_date, [indices])
for idx, entity_data, event_date in entities_to_create:
name_lower = entity_data['text'].lower()
name_lower = entity_data["text"].lower()
if name_lower not in unique_entities:
unique_entities[name_lower] = (entity_data, event_date, [idx])
else:
@@ -222,7 +225,7 @@ class EntityResolver:
indices_map = [] # Maps result index -> list of original indices
for name_lower, (entity_data, event_date, indices) in unique_entities.items():
entity_names.append(entity_data['text'])
entity_names.append(entity_data["text"])
entity_dates.append(event_date)
indices_map.append(indices)
@@ -241,12 +244,12 @@ class EntityResolver:
""",
bank_id,
entity_names,
entity_dates
entity_dates,
)
# Map returned IDs back to original indices
for result_idx, row in enumerate(rows):
entity_id = row['id']
entity_id = row["id"]
for original_idx in indices_map[result_idx]:
entity_ids[original_idx] = entity_id
@@ -257,7 +260,7 @@ class EntityResolver:
bank_id: str,
entity_text: str,
context: str,
nearby_entities: List[Dict],
nearby_entities: list[dict],
unit_event_date,
) -> str:
"""
@@ -287,14 +290,14 @@ class EntityResolver:
)
ORDER BY mention_count DESC
""",
bank_id, entity_text, f"%{entity_text}%"
bank_id,
entity_text,
f"%{entity_text}%",
)
if not candidates:
# New entity - create it
return await self._create_entity(
conn, bank_id, entity_text, unit_event_date
)
return await self._create_entity(conn, bank_id, entity_text, unit_event_date)
# Score candidates based on:
# 1. Name similarity
@@ -306,21 +309,17 @@ class EntityResolver:
best_score = 0.0
best_name_similarity = 0.0
nearby_entity_set = {e['text'].lower() for e in nearby_entities if e['text'] != entity_text}
nearby_entity_set = {e["text"].lower() for e in nearby_entities if e["text"] != entity_text}
for row in candidates:
candidate_id = row['id']
canonical_name = row['canonical_name']
metadata = row['metadata']
last_seen = row['last_seen']
candidate_id = row["id"]
canonical_name = row["canonical_name"]
metadata = row["metadata"]
last_seen = row["last_seen"]
score = 0.0
# 1. Name similarity (0-1)
name_similarity = SequenceMatcher(
None,
entity_text.lower(),
canonical_name.lower()
).ratio()
name_similarity = SequenceMatcher(None, entity_text.lower(), canonical_name.lower()).ratio()
score += name_similarity * 0.5
# 2. Co-occurring entities (0-0.5)
@@ -338,9 +337,9 @@ class EntityResolver:
)
WHERE ec.entity_id_1 = $1 OR ec.entity_id_2 = $1
""",
candidate_id
candidate_id,
)
co_entities = {r['canonical_name'].lower() for r in co_entity_rows}
co_entities = {r["canonical_name"].lower() for r in co_entity_rows}
# Check overlap with nearby entities
overlap = len(nearby_entity_set & co_entities)
@@ -372,14 +371,13 @@ class EntityResolver:
last_seen = $1
WHERE id = $2
""",
unit_event_date, best_candidate
unit_event_date,
best_candidate,
)
return best_candidate
else:
# Not confident - create new entity
return await self._create_entity(
conn, bank_id, entity_text, unit_event_date
)
return await self._create_entity(conn, bank_id, entity_text, unit_event_date)
async def _create_entity(
self,
@@ -413,7 +411,10 @@ class EntityResolver:
last_seen = EXCLUDED.last_seen
RETURNING id
""",
bank_id, entity_text, event_date, event_date
bank_id,
entity_text,
event_date,
event_date,
)
return entity_id
@@ -434,7 +435,8 @@ class EntityResolver:
VALUES ($1, $2)
ON CONFLICT DO NOTHING
""",
unit_id, entity_id
unit_id,
entity_id,
)
# Update co-occurrence cache: find other entities in this unit
@@ -444,10 +446,11 @@ class EntityResolver:
FROM unit_entities
WHERE unit_id = $1 AND entity_id != $2
""",
unit_id, entity_id
unit_id,
entity_id,
)
other_entities = [row['entity_id'] for row in rows]
other_entities = [row["entity_id"] for row in rows]
# Update co-occurrences for each pair
for other_entity_id in other_entities:
@@ -477,10 +480,11 @@ class EntityResolver:
cooccurrence_count = entity_cooccurrences.cooccurrence_count + 1,
last_cooccurred = NOW()
""",
entity_id_1, entity_id_2
entity_id_1,
entity_id_2,
)
async def link_units_to_entities_batch(self, unit_entity_pairs: List[tuple[str, str]], conn=None):
async def link_units_to_entities_batch(self, unit_entity_pairs: list[tuple[str, str]], conn=None):
"""
Link multiple memory units to entities in batch (MUCH faster than sequential).
@@ -499,7 +503,7 @@ class EntityResolver:
else:
return await self._link_units_to_entities_batch_impl(conn, unit_entity_pairs)
async def _link_units_to_entities_batch_impl(self, conn, unit_entity_pairs: List[tuple[str, str]]):
async def _link_units_to_entities_batch_impl(self, conn, unit_entity_pairs: list[tuple[str, str]]):
# Batch insert all unit-entity links
await conn.executemany(
"""
@@ -507,7 +511,7 @@ class EntityResolver:
VALUES ($1, $2)
ON CONFLICT DO NOTHING
""",
unit_entity_pairs
unit_entity_pairs,
)
# Build map of unit -> entities for co-occurrence calculation
@@ -524,7 +528,7 @@ class EntityResolver:
entity_list = list(entity_ids) # Convert set to list for iteration
# For each pair of entities in this unit, create co-occurrence
for i, entity_id_1 in enumerate(entity_list):
for entity_id_2 in entity_list[i+1:]:
for entity_id_2 in entity_list[i + 1 :]:
# Skip if same entity (shouldn't happen with set, but be safe)
if entity_id_1 == entity_id_2:
continue
@@ -535,7 +539,7 @@ class EntityResolver:
# Batch update co-occurrences
if cooccurrence_pairs:
now = datetime.now(timezone.utc)
now = datetime.now(UTC)
await conn.executemany(
"""
INSERT INTO entity_cooccurrences (entity_id_1, entity_id_2, cooccurrence_count, last_cooccurred)
@@ -545,10 +549,10 @@ class EntityResolver:
cooccurrence_count = entity_cooccurrences.cooccurrence_count + 1,
last_cooccurred = EXCLUDED.last_cooccurred
""",
[(e1, e2, 1, now) for e1, e2 in cooccurrence_pairs]
[(e1, e2, 1, now) for e1, e2 in cooccurrence_pairs],
)
async def get_units_by_entity(self, entity_id: str, limit: int = 100) -> List[str]:
async def get_units_by_entity(self, entity_id: str, limit: int = 100) -> list[str]:
"""
Get all units that mention an entity.
@@ -568,15 +572,16 @@ class EntityResolver:
ORDER BY unit_id
LIMIT $2
""",
entity_id, limit
entity_id,
limit,
)
return [row['unit_id'] for row in rows]
return [row["unit_id"] for row in rows]
async def get_entity_by_text(
self,
bank_id: str,
entity_text: str,
) -> Optional[str]:
) -> str | None:
"""
Find an entity by text (for query resolution).
@@ -596,7 +601,8 @@ class EntityResolver:
ORDER BY mention_count DESC
LIMIT 1
""",
bank_id, entity_text
bank_id,
entity_text,
)
return row['id'] if row else None
return row["id"] if row else None
@@ -1,15 +1,17 @@
"""
LLM wrapper for unified configuration across providers.
"""
import asyncio
import logging
import os
import time
import asyncio
from typing import Optional, Any, Dict, List
from openai import AsyncOpenAI, RateLimitError, APIError, APIStatusError, APIConnectionError, LengthFinishReasonError
from typing import Any
from google import genai
from google.genai import types as genai_types
from google.genai import errors as genai_errors
import logging
from google.genai import types as genai_types
from openai import APIConnectionError, APIStatusError, AsyncOpenAI, LengthFinishReasonError
# Seed applied to every Groq request for deterministic behavior.
DEFAULT_LLM_SEED = 4242
@@ -31,6 +33,7 @@ class OutputTooLongError(Exception):
to allow callers to handle output length issues without depending on
provider-specific implementations.
"""
pass
@@ -68,9 +71,7 @@ class LLMProvider:
# Validate provider
valid_providers = ["openai", "groq", "ollama", "gemini"]
if self.provider not in valid_providers:
raise ValueError(
f"Invalid LLM provider: {self.provider}. Must be one of: {', '.join(valid_providers)}"
)
raise ValueError(f"Invalid LLM provider: {self.provider}. Must be one of: {', '.join(valid_providers)}")
# Set default base URLs
if not self.base_url:
@@ -106,7 +107,9 @@ class LLMProvider:
RuntimeError: If the connection test fails.
"""
try:
logger.info(f"Verifying LLM: provider={self.provider}, model={self.model}, base_url={self.base_url or 'default'}...")
logger.info(
f"Verifying LLM: provider={self.provider}, model={self.model}, base_url={self.base_url or 'default'}..."
)
await self.call(
messages=[{"role": "user", "content": "Say 'ok'"}],
max_completion_tokens=10,
@@ -117,16 +120,14 @@ class LLMProvider:
# If we get here without exception, the connection is working
logger.info(f"LLM verified: {self.provider}/{self.model}")
except Exception as e:
raise RuntimeError(
f"LLM connection verification failed for {self.provider}/{self.model}: {e}"
) from e
raise RuntimeError(f"LLM connection verification failed for {self.provider}/{self.model}: {e}") from e
async def call(
self,
messages: List[Dict[str, str]],
response_format: Optional[Any] = None,
max_completion_tokens: Optional[int] = None,
temperature: Optional[float] = None,
messages: list[dict[str, str]],
response_format: Any | None = None,
max_completion_tokens: int | None = None,
temperature: float | None = None,
scope: str = "memory",
max_retries: int = 10,
initial_backoff: float = 1.0,
@@ -161,8 +162,7 @@ class LLMProvider:
# Handle Gemini provider separately
if self.provider == "gemini":
return await self._call_gemini(
messages, response_format, max_retries, initial_backoff,
max_backoff, skip_validation, start_time
messages, response_format, max_retries, initial_backoff, max_backoff, skip_validation, start_time
)
call_params = {
@@ -213,16 +213,18 @@ class LLMProvider:
try:
if response_format is not None:
# Add schema to system message for JSON mode
if hasattr(response_format, 'model_json_schema'):
if hasattr(response_format, "model_json_schema"):
schema = response_format.model_json_schema()
schema_msg = f"\n\nYou must respond with valid JSON matching this schema:\n{json.dumps(schema, indent=2)}"
if call_params['messages'] and call_params['messages'][0].get('role') == 'system':
call_params['messages'][0]['content'] += schema_msg
elif call_params['messages']:
call_params['messages'][0]['content'] = schema_msg + "\n\n" + call_params['messages'][0]['content']
if call_params["messages"] and call_params["messages"][0].get("role") == "system":
call_params["messages"][0]["content"] += schema_msg
elif call_params["messages"]:
call_params["messages"][0]["content"] = (
schema_msg + "\n\n" + call_params["messages"][0]["content"]
)
call_params['response_format'] = {"type": "json_object"}
call_params["response_format"] = {"type": "json_object"}
response = await self._client.chat.completions.create(**call_params)
content = response.choices[0].message.content
@@ -242,8 +244,8 @@ class LLMProvider:
if duration > 10.0:
ratio = max(1, usage.completion_tokens) / usage.prompt_tokens
cached_tokens = 0
if hasattr(usage, 'prompt_tokens_details') and usage.prompt_tokens_details:
cached_tokens = getattr(usage.prompt_tokens_details, 'cached_tokens', 0) or 0
if hasattr(usage, "prompt_tokens_details") and usage.prompt_tokens_details:
cached_tokens = getattr(usage.prompt_tokens_details, "cached_tokens", 0) or 0
cache_info = f", cached_tokens={cached_tokens}" if cached_tokens > 0 else ""
logger.info(
f"slow llm call: model={self.provider}/{self.model}, "
@@ -256,15 +258,19 @@ class LLMProvider:
except LengthFinishReasonError as e:
logger.warning(f"LLM output exceeded token limits: {str(e)}")
raise OutputTooLongError(
f"LLM output exceeded token limits. Input may need to be split into smaller chunks."
"LLM output exceeded token limits. Input may need to be split into smaller chunks."
) from e
except APIConnectionError as e:
last_exception = e
if attempt < max_retries:
status_code = getattr(e, 'status_code', None) or getattr(getattr(e, 'response', None), 'status_code', None)
logger.warning(f"Connection error, retrying... (attempt {attempt + 1}/{max_retries + 1}) - status_code={status_code}, message={e}")
backoff = min(initial_backoff * (2 ** attempt), max_backoff)
status_code = getattr(e, "status_code", None) or getattr(
getattr(e, "response", None), "status_code", None
)
logger.warning(
f"Connection error, retrying... (attempt {attempt + 1}/{max_retries + 1}) - status_code={status_code}, message={e}"
)
backoff = min(initial_backoff * (2**attempt), max_backoff)
await asyncio.sleep(backoff)
continue
else:
@@ -279,7 +285,7 @@ class LLMProvider:
last_exception = e
if attempt < max_retries:
backoff = min(initial_backoff * (2 ** attempt), max_backoff)
backoff = min(initial_backoff * (2**attempt), max_backoff)
jitter = backoff * 0.2 * (2 * (time.time() % 1) - 1)
sleep_time = backoff + jitter
await asyncio.sleep(sleep_time)
@@ -293,12 +299,12 @@ class LLMProvider:
if last_exception:
raise last_exception
raise RuntimeError(f"LLM call failed after all retries with no exception captured")
raise RuntimeError("LLM call failed after all retries with no exception captured")
async def _call_gemini(
self,
messages: List[Dict[str, str]],
response_format: Optional[Any],
messages: list[dict[str, str]],
response_format: Any | None,
max_retries: int,
initial_backoff: float,
max_backoff: float,
@@ -313,27 +319,21 @@ class LLMProvider:
gemini_contents = []
for msg in messages:
role = msg.get('role', 'user')
content = msg.get('content', '')
role = msg.get("role", "user")
content = msg.get("content", "")
if role == 'system':
if role == "system":
if system_instruction:
system_instruction += "\n\n" + content
else:
system_instruction = content
elif role == 'assistant':
gemini_contents.append(genai_types.Content(
role="model",
parts=[genai_types.Part(text=content)]
))
elif role == "assistant":
gemini_contents.append(genai_types.Content(role="model", parts=[genai_types.Part(text=content)]))
else:
gemini_contents.append(genai_types.Content(
role="user",
parts=[genai_types.Part(text=content)]
))
gemini_contents.append(genai_types.Content(role="user", parts=[genai_types.Part(text=content)]))
# Add JSON schema instruction if response_format is provided
if response_format is not None and hasattr(response_format, 'model_json_schema'):
if response_format is not None and hasattr(response_format, "model_json_schema"):
schema = response_format.model_json_schema()
schema_msg = f"\n\nYou must respond with valid JSON matching this schema:\n{json.dumps(schema, indent=2)}"
if system_instruction:
@@ -344,10 +344,10 @@ class LLMProvider:
# Build generation config
config_kwargs = {}
if system_instruction:
config_kwargs['system_instruction'] = system_instruction
config_kwargs["system_instruction"] = system_instruction
if response_format is not None:
config_kwargs['response_mime_type'] = 'application/json'
config_kwargs['response_schema'] = response_format
config_kwargs["response_mime_type"] = "application/json"
config_kwargs["response_schema"] = response_format
generation_config = genai_types.GenerateContentConfig(**config_kwargs) if config_kwargs else None
@@ -366,14 +366,14 @@ class LLMProvider:
# Handle empty response
if content is None:
block_reason = None
if hasattr(response, 'candidates') and response.candidates:
if hasattr(response, "candidates") and response.candidates:
candidate = response.candidates[0]
if hasattr(candidate, 'finish_reason'):
if hasattr(candidate, "finish_reason"):
block_reason = candidate.finish_reason
if attempt < max_retries:
logger.warning(f"Gemini returned empty response (reason: {block_reason}), retrying...")
backoff = min(initial_backoff * (2 ** attempt), max_backoff)
backoff = min(initial_backoff * (2**attempt), max_backoff)
await asyncio.sleep(backoff)
continue
else:
@@ -390,7 +390,7 @@ class LLMProvider:
# Log slow calls
duration = time.time() - start_time
if duration > 10.0 and hasattr(response, 'usage_metadata') and response.usage_metadata:
if duration > 10.0 and hasattr(response, "usage_metadata") and response.usage_metadata:
usage = response.usage_metadata
logger.info(
f"slow llm call: model={self.provider}/{self.model}, "
@@ -403,8 +403,8 @@ class LLMProvider:
except json.JSONDecodeError as e:
last_exception = e
if attempt < max_retries:
logger.warning(f"Gemini returned invalid JSON, retrying...")
backoff = min(initial_backoff * (2 ** attempt), max_backoff)
logger.warning("Gemini returned invalid JSON, retrying...")
backoff = min(initial_backoff * (2**attempt), max_backoff)
await asyncio.sleep(backoff)
continue
else:
@@ -421,7 +421,7 @@ class LLMProvider:
if e.code in (400, 429, 500, 502, 503, 504) or (e.code and e.code >= 500):
last_exception = e
if attempt < max_retries:
backoff = min(initial_backoff * (2 ** attempt), max_backoff)
backoff = min(initial_backoff * (2**attempt), max_backoff)
jitter = backoff * 0.2 * (2 * (time.time() % 1) - 1)
await asyncio.sleep(backoff + jitter)
else:
@@ -437,7 +437,7 @@ class LLMProvider:
if last_exception:
raise last_exception
raise RuntimeError(f"Gemini call failed after all retries")
raise RuntimeError("Gemini call failed after all retries")
@classmethod
def for_memory(cls) -> "LLMProvider":
@@ -447,13 +447,7 @@ class LLMProvider:
base_url = os.getenv("HINDSIGHT_API_LLM_BASE_URL", "")
model = os.getenv("HINDSIGHT_API_LLM_MODEL", "openai/gpt-oss-120b")
return cls(
provider=provider,
api_key=api_key,
base_url=base_url,
model=model,
reasoning_effort="low"
)
return cls(provider=provider, api_key=api_key, base_url=base_url, model=model, reasoning_effort="low")
@classmethod
def for_answer_generation(cls) -> "LLMProvider":
@@ -463,13 +457,7 @@ class LLMProvider:
base_url = os.getenv("HINDSIGHT_API_ANSWER_LLM_BASE_URL", os.getenv("HINDSIGHT_API_LLM_BASE_URL", ""))
model = os.getenv("HINDSIGHT_API_ANSWER_LLM_MODEL", os.getenv("HINDSIGHT_API_LLM_MODEL", "openai/gpt-oss-120b"))
return cls(
provider=provider,
api_key=api_key,
base_url=base_url,
model=model,
reasoning_effort="high"
)
return cls(provider=provider, api_key=api_key, base_url=base_url, model=model, reasoning_effort="high")
@classmethod
def for_judge(cls) -> "LLMProvider":
@@ -479,13 +467,7 @@ class LLMProvider:
base_url = os.getenv("HINDSIGHT_API_JUDGE_LLM_BASE_URL", os.getenv("HINDSIGHT_API_LLM_BASE_URL", ""))
model = os.getenv("HINDSIGHT_API_JUDGE_LLM_MODEL", os.getenv("HINDSIGHT_API_LLM_MODEL", "openai/gpt-oss-120b"))
return cls(
provider=provider,
api_key=api_key,
base_url=base_url,
model=model,
reasoning_effort="high"
)
return cls(provider=provider, api_key=api_key, base_url=base_url, model=model, reasoning_effort="high")
# Backwards compatibility alias
File diff suppressed because it is too large Load Diff
@@ -4,11 +4,12 @@ Query analysis abstraction for the memory system.
Provides an interface for analyzing natural language queries to extract
structured information like temporal constraints.
"""
from abc import ABC, abstractmethod
from typing import Optional
from datetime import datetime, timedelta
import logging
import re
from abc import ABC, abstractmethod
from datetime import datetime, timedelta
from pydantic import BaseModel, Field
logger = logging.getLogger(__name__)
@@ -20,6 +21,7 @@ class TemporalConstraint(BaseModel):
Represents a time range with start and end dates.
"""
start_date: datetime = Field(description="Start of the time range (inclusive)")
end_date: datetime = Field(description="End of the time range (inclusive)")
@@ -33,9 +35,9 @@ class QueryAnalysis(BaseModel):
Contains extracted structured information like temporal constraints.
"""
temporal_constraint: Optional[TemporalConstraint] = Field(
default=None,
description="Extracted temporal constraint, if any"
temporal_constraint: TemporalConstraint | None = Field(
default=None, description="Extracted temporal constraint, if any"
)
@@ -58,9 +60,7 @@ class QueryAnalyzer(ABC):
pass
@abstractmethod
def analyze(
self, query: str, reference_date: Optional[datetime] = None
) -> QueryAnalysis:
def analyze(self, query: str, reference_date: datetime | None = None) -> QueryAnalysis:
"""
Analyze a natural language query.
@@ -95,11 +95,10 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
"""Load dateparser (lazy import)."""
if self._search_dates is None:
from dateparser.search import search_dates
self._search_dates = search_dates
def analyze(
self, query: str, reference_date: Optional[datetime] = None
) -> QueryAnalysis:
def analyze(self, query: str, reference_date: datetime | None = None) -> QueryAnalysis:
"""
Analyze query using dateparser.
@@ -126,9 +125,9 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
# Use dateparser's search_dates to find temporal expressions
settings = {
'RELATIVE_BASE': reference_date,
'PREFER_DATES_FROM': 'past',
'RETURN_AS_TIMEZONE_AWARE': False,
"RELATIVE_BASE": reference_date,
"PREFER_DATES_FROM": "past",
"RETURN_AS_TIMEZONE_AWARE": False,
}
results = self._search_dates(query, settings=settings)
@@ -137,11 +136,8 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
return QueryAnalysis(temporal_constraint=None)
# Filter out false positives (common words parsed as dates)
false_positives = {'do', 'may', 'march', 'will', 'can', 'sat', 'sun', 'mon', 'tue', 'wed', 'thu', 'fri'}
valid_results = [
(text, date) for text, date in results
if text.lower() not in false_positives or len(text) > 3
]
false_positives = {"do", "may", "march", "will", "can", "sat", "sun", "mon", "tue", "wed", "thu", "fri"}
valid_results = [(text, date) for text, date in results if text.lower() not in false_positives or len(text) > 3]
if not valid_results:
return QueryAnalysis(temporal_constraint=None)
@@ -153,84 +149,94 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
start_date = parsed_date.replace(hour=0, minute=0, second=0, microsecond=0)
end_date = parsed_date.replace(hour=23, minute=59, second=59, microsecond=999999)
return QueryAnalysis(
temporal_constraint=TemporalConstraint(
start_date=start_date,
end_date=end_date
)
)
return QueryAnalysis(temporal_constraint=TemporalConstraint(start_date=start_date, end_date=end_date))
def _extract_period(
self, query: str, reference_date: datetime
) -> Optional[TemporalConstraint]:
def _extract_period(self, query: str, reference_date: datetime) -> TemporalConstraint | None:
"""
Extract period-based temporal expressions (week, month, year, weekend).
These need special handling as they represent date ranges, not single dates.
Supports multiple languages.
"""
def constraint(start: datetime, end: datetime) -> TemporalConstraint:
return TemporalConstraint(
start_date=start.replace(hour=0, minute=0, second=0, microsecond=0),
end_date=end.replace(hour=23, minute=59, second=59, microsecond=999999)
end_date=end.replace(hour=23, minute=59, second=59, microsecond=999999),
)
# Yesterday patterns (English, Spanish, Italian, French, German)
if re.search(r'\b(yesterday|ayer|ieri|hier|gestern)\b', query, re.IGNORECASE):
if re.search(r"\b(yesterday|ayer|ieri|hier|gestern)\b", query, re.IGNORECASE):
d = reference_date - timedelta(days=1)
return constraint(d, d)
# Today patterns
if re.search(r'\b(today|hoy|oggi|aujourd\'?hui|heute)\b', query, re.IGNORECASE):
if re.search(r"\b(today|hoy|oggi|aujourd\'?hui|heute)\b", query, re.IGNORECASE):
return constraint(reference_date, reference_date)
# "a couple of days ago" / "a few days ago" patterns
# These are imprecise so we create a range
if re.search(r'\b(a\s+)?couple\s+(of\s+)?days?\s+ago\b', query, re.IGNORECASE):
if re.search(r"\b(a\s+)?couple\s+(of\s+)?days?\s+ago\b", query, re.IGNORECASE):
# "a couple of days" = approximately 2 days, give range of 1-3 days
return constraint(reference_date - timedelta(days=3), reference_date - timedelta(days=1))
if re.search(r'\b(a\s+)?few\s+days?\s+ago\b', query, re.IGNORECASE):
if re.search(r"\b(a\s+)?few\s+days?\s+ago\b", query, re.IGNORECASE):
# "a few days" = approximately 3-4 days, give range of 2-5 days
return constraint(reference_date - timedelta(days=5), reference_date - timedelta(days=2))
# "a couple of weeks ago" / "a few weeks ago" patterns
if re.search(r'\b(a\s+)?couple\s+(of\s+)?weeks?\s+ago\b', query, re.IGNORECASE):
if re.search(r"\b(a\s+)?couple\s+(of\s+)?weeks?\s+ago\b", query, re.IGNORECASE):
# "a couple of weeks" = approximately 2 weeks, give range of 1-3 weeks
return constraint(reference_date - timedelta(weeks=3), reference_date - timedelta(weeks=1))
if re.search(r'\b(a\s+)?few\s+weeks?\s+ago\b', query, re.IGNORECASE):
if re.search(r"\b(a\s+)?few\s+weeks?\s+ago\b", query, re.IGNORECASE):
# "a few weeks" = approximately 3-4 weeks, give range of 2-5 weeks
return constraint(reference_date - timedelta(weeks=5), reference_date - timedelta(weeks=2))
# "a couple of months ago" / "a few months ago" patterns
if re.search(r'\b(a\s+)?couple\s+(of\s+)?months?\s+ago\b', query, re.IGNORECASE):
if re.search(r"\b(a\s+)?couple\s+(of\s+)?months?\s+ago\b", query, re.IGNORECASE):
# "a couple of months" = approximately 2 months, give range of 1-3 months
return constraint(reference_date - timedelta(days=90), reference_date - timedelta(days=30))
if re.search(r'\b(a\s+)?few\s+months?\s+ago\b', query, re.IGNORECASE):
if re.search(r"\b(a\s+)?few\s+months?\s+ago\b", query, re.IGNORECASE):
# "a few months" = approximately 3-4 months, give range of 2-5 months
return constraint(reference_date - timedelta(days=150), reference_date - timedelta(days=60))
# Last week patterns (English, Spanish, Italian, French, German)
if re.search(r'\b(last\s+week|la\s+semana\s+pasada|la\s+settimana\s+scorsa|la\s+semaine\s+derni[eè]re|letzte\s+woche)\b', query, re.IGNORECASE):
if re.search(
r"\b(last\s+week|la\s+semana\s+pasada|la\s+settimana\s+scorsa|la\s+semaine\s+derni[eè]re|letzte\s+woche)\b",
query,
re.IGNORECASE,
):
start = reference_date - timedelta(days=reference_date.weekday() + 7)
return constraint(start, start + timedelta(days=6))
# Last month patterns
if re.search(r'\b(last\s+month|el\s+mes\s+pasado|il\s+mese\s+scorso|le\s+mois\s+dernier|letzten?\s+monat)\b', query, re.IGNORECASE):
if re.search(
r"\b(last\s+month|el\s+mes\s+pasado|il\s+mese\s+scorso|le\s+mois\s+dernier|letzten?\s+monat)\b",
query,
re.IGNORECASE,
):
first = reference_date.replace(day=1)
end = first - timedelta(days=1)
start = end.replace(day=1)
return constraint(start, end)
# Last year patterns
if re.search(r'\b(last\s+year|el\s+a[ñn]o\s+pasado|l\'anno\s+scorso|l\'ann[ée]e\s+derni[eè]re|letztes?\s+jahr)\b', query, re.IGNORECASE):
if re.search(
r"\b(last\s+year|el\s+a[ñn]o\s+pasado|l\'anno\s+scorso|l\'ann[ée]e\s+derni[eè]re|letztes?\s+jahr)\b",
query,
re.IGNORECASE,
):
year = reference_date.year - 1
return constraint(datetime(year, 1, 1), datetime(year, 12, 31))
# Last weekend patterns
if re.search(r'\b(last\s+weekend|el\s+fin\s+de\s+semana\s+pasado|lo\s+scorso\s+fine\s+settimana|le\s+week-?end\s+dernier|letztes?\s+wochenende)\b', query, re.IGNORECASE):
if re.search(
r"\b(last\s+weekend|el\s+fin\s+de\s+semana\s+pasado|lo\s+scorso\s+fine\s+settimana|le\s+week-?end\s+dernier|letztes?\s+wochenende)\b",
query,
re.IGNORECASE,
):
days_since_sat = (reference_date.weekday() + 2) % 7
if days_since_sat == 0:
days_since_sat = 7
@@ -239,22 +245,22 @@ class DateparserQueryAnalyzer(QueryAnalyzer):
# Month + Year patterns (e.g., "June 2024", "junio 2024", "giugno 2024")
month_patterns = {
'january|enero|gennaio|janvier|januar': 1,
'february|febrero|febbraio|f[ée]vrier|februar': 2,
'march|marzo|mars|m[äa]rz': 3,
'april|abril|aprile|avril': 4,
'may|mayo|maggio|mai': 5,
'june|junio|giugno|juin|juni': 6,
'july|julio|luglio|juillet|juli': 7,
'august|agosto|ao[uû]t': 8,
'september|septiembre|settembre|septembre': 9,
'october|octubre|ottobre|octobre|oktober': 10,
'november|noviembre|novembre': 11,
'december|diciembre|dicembre|d[ée]cembre|dezember': 12,
"january|enero|gennaio|janvier|januar": 1,
"february|febrero|febbraio|f[ée]vrier|februar": 2,
"march|marzo|mars|m[äa]rz": 3,
"april|abril|aprile|avril": 4,
"may|mayo|maggio|mai": 5,
"june|junio|giugno|juin|juni": 6,
"july|julio|luglio|juillet|juli": 7,
"august|agosto|ao[uû]t": 8,
"september|septiembre|settembre|septembre": 9,
"october|octubre|ottobre|octobre|oktober": 10,
"november|noviembre|novembre": 11,
"december|diciembre|dicembre|d[ée]cembre|dezember": 12,
}
for pattern, month_num in month_patterns.items():
match = re.search(rf'\b({pattern})\s+(\d{{4}})\b', query, re.IGNORECASE)
match = re.search(rf"\b({pattern})\s+(\d{{4}})\b", query, re.IGNORECASE)
if match:
year = int(match.group(2))
start = datetime(year, month_num, 1)
@@ -279,11 +285,7 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
- Model size: ~80M params (~300MB download)
"""
def __init__(
self,
model_name: str = "google/flan-t5-small",
device: str = "cpu"
):
def __init__(self, model_name: str = "google/flan-t5-small", device: str = "cpu"):
"""
Initialize T5 query analyzer.
@@ -304,11 +306,10 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
return
try:
from transformers import AutoTokenizer, AutoModelForSeq2SeqLM
from transformers import AutoModelForSeq2SeqLM, AutoTokenizer
except ImportError:
raise ImportError(
"transformers is required for TransformerQueryAnalyzer. "
"Install it with: pip install transformers"
"transformers is required for TransformerQueryAnalyzer. Install it with: pip install transformers"
)
logger.info(f"Loading query analyzer model: {self.model_name}...")
@@ -322,9 +323,7 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
"""Lazy load the T5 model for temporal extraction (calls load())."""
self.load()
def _extract_with_rules(
self, query: str, reference_date: datetime
) -> Optional[TemporalConstraint]:
def _extract_with_rules(self, query: str, reference_date: datetime) -> TemporalConstraint | None:
"""
Extract temporal expressions using rule-based patterns.
@@ -332,6 +331,7 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
patterns that need model-based extraction.
"""
import re
query_lower = query.lower()
def get_last_weekday(weekday: int) -> datetime:
@@ -343,50 +343,60 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
def constraint(start: datetime, end: datetime) -> TemporalConstraint:
return TemporalConstraint(
start_date=start.replace(hour=0, minute=0, second=0, microsecond=0),
end_date=end.replace(hour=23, minute=59, second=59, microsecond=999999)
end_date=end.replace(hour=23, minute=59, second=59, microsecond=999999),
)
# Yesterday
if re.search(r'\byesterday\b', query_lower):
if re.search(r"\byesterday\b", query_lower):
d = reference_date - timedelta(days=1)
return constraint(d, d)
# Last week
if re.search(r'\blast\s+week\b', query_lower):
if re.search(r"\blast\s+week\b", query_lower):
start = reference_date - timedelta(days=reference_date.weekday() + 7)
return constraint(start, start + timedelta(days=6))
# Last month
if re.search(r'\blast\s+month\b', query_lower):
if re.search(r"\blast\s+month\b", query_lower):
first = reference_date.replace(day=1)
end = first - timedelta(days=1)
start = end.replace(day=1)
return constraint(start, end)
# Last year
if re.search(r'\blast\s+year\b', query_lower):
if re.search(r"\blast\s+year\b", query_lower):
y = reference_date.year - 1
return constraint(datetime(y, 1, 1), datetime(y, 12, 31))
# Last weekend
if re.search(r'\blast\s+weekend\b', query_lower):
if re.search(r"\blast\s+weekend\b", query_lower):
sat = get_last_weekday(5)
return constraint(sat, sat + timedelta(days=1))
# Last <weekday>
weekdays = {'monday': 0, 'tuesday': 1, 'wednesday': 2, 'thursday': 3,
'friday': 4, 'saturday': 5, 'sunday': 6}
weekdays = {"monday": 0, "tuesday": 1, "wednesday": 2, "thursday": 3, "friday": 4, "saturday": 5, "sunday": 6}
for name, num in weekdays.items():
if re.search(rf'\blast\s+{name}\b', query_lower):
if re.search(rf"\blast\s+{name}\b", query_lower):
d = get_last_weekday(num)
return constraint(d, d)
# Month + Year: "June 2024", "in March 2023"
months = {'january': 1, 'february': 2, 'march': 3, 'april': 4, 'may': 5,
'june': 6, 'july': 7, 'august': 8, 'september': 9, 'october': 10,
'november': 11, 'december': 12}
months = {
"january": 1,
"february": 2,
"march": 3,
"april": 4,
"may": 5,
"june": 6,
"july": 7,
"august": 8,
"september": 9,
"october": 10,
"november": 11,
"december": 12,
}
for name, num in months.items():
match = re.search(rf'\b{name}\s+(\d{{4}})\b', query_lower)
match = re.search(rf"\b{name}\s+(\d{{4}})\b", query_lower)
if match:
year = int(match.group(1))
if num == 12:
@@ -397,9 +407,7 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
return None
def analyze(
self, query: str, reference_date: Optional[datetime] = None
) -> QueryAnalysis:
def analyze(self, query: str, reference_date: datetime | None = None) -> QueryAnalysis:
"""
Analyze query for temporal expressions.
@@ -435,11 +443,11 @@ class TransformerQueryAnalyzer(QueryAnalyzer):
last_saturday = get_last_weekday(5)
# Build prompt for T5
prompt = f"""Today is {reference_date.strftime('%Y-%m-%d')}. Extract date range or "none".
prompt = f"""Today is {reference_date.strftime("%Y-%m-%d")}. Extract date range or "none".
June 2024 = 2024-06-01 to 2024-06-30
yesterday = {yesterday.strftime('%Y-%m-%d')} to {yesterday.strftime('%Y-%m-%d')}
last Saturday = {last_saturday.strftime('%Y-%m-%d')} to {last_saturday.strftime('%Y-%m-%d')}
yesterday = {yesterday.strftime("%Y-%m-%d")} to {yesterday.strftime("%Y-%m-%d")}
last Saturday = {last_saturday.strftime("%Y-%m-%d")} to {last_saturday.strftime("%Y-%m-%d")}
what is the weather = none
{query} ="""
@@ -448,13 +456,7 @@ what is the weather = none
inputs = {k: v.to(self.device) for k, v in inputs.items()}
with self._no_grad():
outputs = self._model.generate(
**inputs,
max_new_tokens=30,
num_beams=3,
do_sample=False,
temperature=1.0
)
outputs = self._model.generate(**inputs, max_new_tokens=30, num_beams=3, do_sample=False, temperature=1.0)
result = self._tokenizer.decode(outputs[0], skip_special_tokens=True).strip()
@@ -466,14 +468,14 @@ what is the weather = none
"""Get torch.no_grad context manager."""
try:
import torch
return torch.no_grad()
except ImportError:
from contextlib import nullcontext
return nullcontext()
def _parse_generated_output(
self, result: str, reference_date: datetime
) -> Optional[TemporalConstraint]:
def _parse_generated_output(self, result: str, reference_date: datetime) -> TemporalConstraint | None:
"""
Parse T5 generated output into TemporalConstraint.
@@ -492,7 +494,8 @@ what is the weather = none
try:
# Parse "YYYY-MM-DD to YYYY-MM-DD"
import re
pattern = r'(\d{4}-\d{2}-\d{2})\s+to\s+(\d{4}-\d{2}-\d{2})'
pattern = r"(\d{4}-\d{2}-\d{2})\s+to\s+(\d{4}-\d{2}-\d{2})"
match = re.search(pattern, result, re.IGNORECASE)
if match:
@@ -513,7 +516,7 @@ what is the weather = none
return TemporalConstraint(start_date=start_date, end_date=end_date)
except (ValueError, AttributeError) as e:
except (ValueError, AttributeError):
return None
return None
@@ -6,9 +6,9 @@ API response models should be kept separate and convert from these core models t
API stability even if internal models change.
"""
from typing import Optional, List, Dict, Any
from pydantic import BaseModel, Field, ConfigDict
from typing import Any
from pydantic import BaseModel, ConfigDict, Field
# Valid fact types for recall operations (excludes 'observation' which is internal)
VALID_RECALL_FACT_TYPES = frozenset(["world", "experience", "opinion"])
@@ -23,17 +23,12 @@ class DispositionTraits(BaseModel):
- literalism: 1=flexible interpretation, 5=literal interpretation (how strictly to interpret information)
- empathy: 1=detached, 5=empathetic (how much to consider emotional context)
"""
skepticism: int = Field(ge=1, le=5, description="How skeptical vs trusting (1=trusting, 5=skeptical)")
literalism: int = Field(ge=1, le=5, description="How literally to interpret information (1=flexible, 5=literal)")
empathy: int = Field(ge=1, le=5, description="How much to consider emotional context (1=detached, 5=empathetic)")
model_config = ConfigDict(json_schema_extra={
"example": {
"skepticism": 3,
"literalism": 3,
"empathy": 3
}
})
model_config = ConfigDict(json_schema_extra={"example": {"skepticism": 3, "literalism": 3, "empathy": 3}})
class MemoryFact(BaseModel):
@@ -43,38 +38,44 @@ class MemoryFact(BaseModel):
This represents a unit of information stored in the memory system,
including both the content and metadata.
"""
model_config = ConfigDict(json_schema_extra={
"example": {
"id": "123e4567-e89b-12d3-a456-426614174000",
"text": "Alice works at Google on the AI team",
"fact_type": "world",
"entities": ["Alice", "Google"],
"context": "work info",
"occurred_start": "2024-01-15T10:30:00Z",
"occurred_end": "2024-01-15T10:30:00Z",
"mentioned_at": "2024-01-15T10:30:00Z",
"document_id": "session_abc123",
"metadata": {"source": "slack"},
"chunk_id": "bank123_session_abc123_0",
"activation": 0.95
model_config = ConfigDict(
json_schema_extra={
"example": {
"id": "123e4567-e89b-12d3-a456-426614174000",
"text": "Alice works at Google on the AI team",
"fact_type": "world",
"entities": ["Alice", "Google"],
"context": "work info",
"occurred_start": "2024-01-15T10:30:00Z",
"occurred_end": "2024-01-15T10:30:00Z",
"mentioned_at": "2024-01-15T10:30:00Z",
"document_id": "session_abc123",
"metadata": {"source": "slack"},
"chunk_id": "bank123_session_abc123_0",
"activation": 0.95,
}
}
})
)
id: str = Field(description="Unique identifier for the memory fact")
text: str = Field(description="The actual text content of the memory")
fact_type: str = Field(description="Type of fact: 'world', 'experience', 'opinion', or 'observation'")
entities: Optional[List[str]] = Field(None, description="Entity names mentioned in this fact")
context: Optional[str] = Field(None, description="Additional context for the memory")
occurred_start: Optional[str] = Field(None, description="ISO format date when the event started occurring")
occurred_end: Optional[str] = Field(None, description="ISO format date when the event ended occurring")
mentioned_at: Optional[str] = Field(None, description="ISO format date when the fact was mentioned/learned")
document_id: Optional[str] = Field(None, description="ID of the document this memory belongs to")
metadata: Optional[Dict[str, str]] = Field(None, description="User-defined metadata")
chunk_id: Optional[str] = Field(None, description="ID of the chunk this fact was extracted from (format: bank_id_document_id_chunk_index)")
entities: list[str] | None = Field(None, description="Entity names mentioned in this fact")
context: str | None = Field(None, description="Additional context for the memory")
occurred_start: str | None = Field(None, description="ISO format date when the event started occurring")
occurred_end: str | None = Field(None, description="ISO format date when the event ended occurring")
mentioned_at: str | None = Field(None, description="ISO format date when the fact was mentioned/learned")
document_id: str | None = Field(None, description="ID of the document this memory belongs to")
metadata: dict[str, str] | None = Field(None, description="User-defined metadata")
chunk_id: str | None = Field(
None, description="ID of the chunk this fact was extracted from (format: bank_id_document_id_chunk_index)"
)
class ChunkInfo(BaseModel):
"""Information about a chunk."""
chunk_text: str = Field(description="The raw chunk text")
chunk_index: int = Field(description="Index of the chunk within the document")
truncated: bool = Field(default=False, description="Whether the chunk was truncated due to token limits")
@@ -87,35 +88,33 @@ class RecallResult(BaseModel):
Contains a list of matching memory facts and optional trace information
for debugging and transparency.
"""
model_config = ConfigDict(json_schema_extra={
"example": {
"results": [
{
"id": "123e4567-e89b-12d3-a456-426614174000",
"text": "Alice works at Google on the AI team",
"fact_type": "world",
"context": "work info",
"occurred_start": "2024-01-15T10:30:00Z",
"occurred_end": "2024-01-15T10:30:00Z",
"activation": 0.95
}
],
"trace": {
"query": "What did Alice say about machine learning?",
"num_results": 1
model_config = ConfigDict(
json_schema_extra={
"example": {
"results": [
{
"id": "123e4567-e89b-12d3-a456-426614174000",
"text": "Alice works at Google on the AI team",
"fact_type": "world",
"context": "work info",
"occurred_start": "2024-01-15T10:30:00Z",
"occurred_end": "2024-01-15T10:30:00Z",
"activation": 0.95,
}
],
"trace": {"query": "What did Alice say about machine learning?", "num_results": 1},
}
}
})
results: List[MemoryFact] = Field(description="List of memory facts matching the query")
trace: Optional[Dict[str, Any]] = Field(None, description="Trace information for debugging")
entities: Optional[Dict[str, "EntityState"]] = Field(
None,
description="Entity states for entities mentioned in results (keyed by canonical name)"
)
chunks: Optional[Dict[str, ChunkInfo]] = Field(
None,
description="Chunks for facts, keyed by '{document_id}_{chunk_index}'"
results: list[MemoryFact] = Field(description="List of memory facts matching the query")
trace: dict[str, Any] | None = Field(None, description="Trace information for debugging")
entities: dict[str, "EntityState"] | None = Field(
None, description="Entity states for entities mentioned in results (keyed by canonical name)"
)
chunks: dict[str, ChunkInfo] | None = Field(
None, description="Chunks for facts, keyed by '{document_id}_{chunk_index}'"
)
@@ -126,37 +125,35 @@ class ReflectResult(BaseModel):
Contains the formulated answer, the facts it was based on (organized by type),
and any new opinions that were formed during the reflection process.
"""
model_config = ConfigDict(json_schema_extra={
"example": {
"text": "Based on my knowledge, machine learning is being actively used in healthcare...",
"based_on": {
"world": [
{
"id": "123e4567-e89b-12d3-a456-426614174000",
"text": "Machine learning is used in medical diagnosis",
"fact_type": "world",
"context": "healthcare",
"occurred_start": "2024-01-15T10:30:00Z",
"occurred_end": "2024-01-15T10:30:00Z"
}
],
"experience": [],
"opinion": []
},
"new_opinions": [
"Machine learning has great potential in healthcare"
]
model_config = ConfigDict(
json_schema_extra={
"example": {
"text": "Based on my knowledge, machine learning is being actively used in healthcare...",
"based_on": {
"world": [
{
"id": "123e4567-e89b-12d3-a456-426614174000",
"text": "Machine learning is used in medical diagnosis",
"fact_type": "world",
"context": "healthcare",
"occurred_start": "2024-01-15T10:30:00Z",
"occurred_end": "2024-01-15T10:30:00Z",
}
],
"experience": [],
"opinion": [],
},
"new_opinions": ["Machine learning has great potential in healthcare"],
}
}
})
)
text: str = Field(description="The formulated answer text")
based_on: Dict[str, List[MemoryFact]] = Field(
based_on: dict[str, list[MemoryFact]] = Field(
description="Facts used to formulate the answer, organized by type (world, experience, opinion)"
)
new_opinions: List[str] = Field(
default_factory=list,
description="List of newly formed opinions during reflection"
)
new_opinions: list[str] = Field(default_factory=list, description="List of newly formed opinions during reflection")
class Opinion(BaseModel):
@@ -166,12 +163,12 @@ class Opinion(BaseModel):
Opinions represent the bank's formed perspectives on topics,
with a confidence level indicating strength of belief.
"""
model_config = ConfigDict(json_schema_extra={
"example": {
"text": "Machine learning has great potential in healthcare",
"confidence": 0.85
model_config = ConfigDict(
json_schema_extra={
"example": {"text": "Machine learning has great potential in healthcare", "confidence": 0.85}
}
})
)
text: str = Field(description="The opinion text")
confidence: float = Field(description="Confidence score between 0.0 and 1.0")
@@ -184,15 +181,15 @@ class EntityObservation(BaseModel):
Observations are objective facts synthesized from multiple memory facts
about an entity, without personality influence.
"""
model_config = ConfigDict(json_schema_extra={
"example": {
"text": "John is detail-oriented and works at Google",
"mentioned_at": "2024-01-15T10:30:00Z"
model_config = ConfigDict(
json_schema_extra={
"example": {"text": "John is detail-oriented and works at Google", "mentioned_at": "2024-01-15T10:30:00Z"}
}
})
)
text: str = Field(description="The observation text")
mentioned_at: Optional[str] = Field(None, description="ISO format date when this observation was created")
mentioned_at: str | None = Field(None, description="ISO format date when this observation was created")
class EntityState(BaseModel):
@@ -201,20 +198,22 @@ class EntityState(BaseModel):
Contains observations synthesized from facts about the entity.
"""
model_config = ConfigDict(json_schema_extra={
"example": {
"entity_id": "123e4567-e89b-12d3-a456-426614174000",
"canonical_name": "John",
"observations": [
{"text": "John is detail-oriented", "mentioned_at": "2024-01-15T10:30:00Z"},
{"text": "John works at Google on the AI team", "mentioned_at": "2024-01-14T09:00:00Z"}
]
model_config = ConfigDict(
json_schema_extra={
"example": {
"entity_id": "123e4567-e89b-12d3-a456-426614174000",
"canonical_name": "John",
"observations": [
{"text": "John is detail-oriented", "mentioned_at": "2024-01-15T10:30:00Z"},
{"text": "John works at Google on the AI team", "mentioned_at": "2024-01-14T09:00:00Z"},
],
}
}
})
)
entity_id: str = Field(description="Unique identifier for the entity")
canonical_name: str = Field(description="Canonical name of the entity")
observations: List[EntityObservation] = Field(
default_factory=list,
description="List of observations about this entity"
observations: list[EntityObservation] = Field(
default_factory=list, description="List of observations about this entity"
)
@@ -12,23 +12,16 @@ This package contains modular components for the retain operation:
- fact_storage: Handle fact insertion into database
"""
from .types import (
RetainContent,
ExtractedFact,
ProcessedFact,
ChunkMetadata,
EntityRef,
CausalRelation,
RetainBatch
from . import (
chunk_storage,
deduplication,
embedding_processing,
entity_processing,
fact_extraction,
fact_storage,
link_creation,
)
from . import fact_extraction
from . import embedding_processing
from . import deduplication
from . import entity_processing
from . import link_creation
from . import chunk_storage
from . import fact_storage
from .types import CausalRelation, ChunkMetadata, EntityRef, ExtractedFact, ProcessedFact, RetainBatch, RetainContent
__all__ = [
# Types
@@ -5,8 +5,10 @@ bank profile utilities for disposition and background management.
import json
import logging
import re
from typing import Dict, Optional, TypedDict
from typing import TypedDict
from pydantic import BaseModel, Field
from ..db_utils import acquire_with_retry
from ..response_models import DispositionTraits
@@ -21,6 +23,7 @@ DEFAULT_DISPOSITION = {
class BankProfile(TypedDict):
"""Type for bank profile data."""
name: str
disposition: DispositionTraits
background: str
@@ -28,6 +31,7 @@ class BankProfile(TypedDict):
class BackgroundMergeResponse(BaseModel):
"""LLM response for background merge with disposition inference."""
background: str = Field(description="Merged background in first person perspective")
disposition: DispositionTraits = Field(description="Inferred disposition traits (skepticism, literalism, empathy)")
@@ -51,7 +55,7 @@ async def get_bank_profile(pool, bank_id: str) -> BankProfile:
SELECT name, disposition, background
FROM banks WHERE bank_id = $1
""",
bank_id
bank_id,
)
if row:
@@ -61,9 +65,7 @@ async def get_bank_profile(pool, bank_id: str) -> BankProfile:
disposition_data = json.loads(disposition_data)
return BankProfile(
name=row["name"],
disposition=DispositionTraits(**disposition_data),
background=row["background"]
name=row["name"], disposition=DispositionTraits(**disposition_data), background=row["background"]
)
# Bank doesn't exist, create with defaults
@@ -76,21 +78,13 @@ async def get_bank_profile(pool, bank_id: str) -> BankProfile:
bank_id,
bank_id, # Default name is the bank_id
json.dumps(DEFAULT_DISPOSITION),
""
"",
)
return BankProfile(
name=bank_id,
disposition=DispositionTraits(**DEFAULT_DISPOSITION),
background=""
)
return BankProfile(name=bank_id, disposition=DispositionTraits(**DEFAULT_DISPOSITION), background="")
async def update_bank_disposition(
pool,
bank_id: str,
disposition: Dict[str, int]
) -> None:
async def update_bank_disposition(pool, bank_id: str, disposition: dict[str, int]) -> None:
"""
Update bank disposition traits.
@@ -111,17 +105,11 @@ async def update_bank_disposition(
WHERE bank_id = $1
""",
bank_id,
json.dumps(disposition)
json.dumps(disposition),
)
async def merge_bank_background(
pool,
llm_config,
bank_id: str,
new_info: str,
update_disposition: bool = True
) -> dict:
async def merge_bank_background(pool, llm_config, bank_id: str, new_info: str, update_disposition: bool = True) -> dict:
"""
Merge new background information with existing background using LLM.
Normalizes to first person ("I") and resolves conflicts.
@@ -142,12 +130,7 @@ async def merge_bank_background(
current_background = profile["background"]
# Use LLM to merge backgrounds and optionally infer disposition
result = await _llm_merge_background(
llm_config,
current_background,
new_info,
infer_disposition=update_disposition
)
result = await _llm_merge_background(llm_config, current_background, new_info, infer_disposition=update_disposition)
merged_background = result["background"]
inferred_disposition = result.get("disposition")
@@ -166,7 +149,7 @@ async def merge_bank_background(
""",
bank_id,
merged_background,
json.dumps(inferred_disposition)
json.dumps(inferred_disposition),
)
else:
# Update only background
@@ -178,7 +161,7 @@ async def merge_bank_background(
WHERE bank_id = $1
""",
bank_id,
merged_background
merged_background,
)
response = {"background": merged_background}
@@ -188,12 +171,7 @@ async def merge_bank_background(
return response
async def _llm_merge_background(
llm_config,
current: str,
new_info: str,
infer_disposition: bool = False
) -> dict:
async def _llm_merge_background(llm_config, current: str, new_info: str, infer_disposition: bool = False) -> dict:
"""
Use LLM to intelligently merge background information.
Optionally infer Big Five disposition traits from the merged background.
@@ -273,25 +251,19 @@ Merged background:"""
response_format=BackgroundMergeResponse,
scope="bank_background",
temperature=0.3,
max_completion_tokens=8192
max_completion_tokens=8192,
)
logger.info(f"Successfully got structured response: background={parsed.background[:100]}")
# Convert Pydantic model to dict format
return {
"background": parsed.background,
"disposition": parsed.disposition.model_dump()
}
return {"background": parsed.background, "disposition": parsed.disposition.model_dump()}
except Exception as e:
logger.warning(f"Structured output failed, falling back to manual parsing: {e}")
# Fall through to manual parsing below
# Manual parsing fallback or non-disposition merge
content = await llm_config.call(
messages=messages,
scope="bank_background",
temperature=0.3,
max_completion_tokens=8192
messages=messages, scope="bank_background", temperature=0.3, max_completion_tokens=8192
)
logger.info(f"LLM response for background merge (first 500 chars): {content[:500]}")
@@ -310,7 +282,7 @@ Merged background:"""
# Method 2: Extract from markdown code blocks
if result is None:
# Remove markdown code blocks
code_block_match = re.search(r'```(?:json)?\s*(\{.*?\})\s*```', content, re.DOTALL)
code_block_match = re.search(r"```(?:json)?\s*(\{.*?\})\s*```", content, re.DOTALL)
if code_block_match:
try:
result = json.loads(code_block_match.group(1))
@@ -321,7 +293,9 @@ Merged background:"""
# Method 3: Find nested JSON structure
if result is None:
# Look for JSON object with nested structure
json_match = re.search(r'\{[^{}]*"background"[^{}]*"disposition"[^{}]*\{[^{}]*\}[^{}]*\}', content, re.DOTALL)
json_match = re.search(
r'\{[^{}]*"background"[^{}]*"disposition"[^{}]*\{[^{}]*\}[^{}]*\}', content, re.DOTALL
)
if json_match:
try:
result = json.loads(json_match.group())
@@ -335,7 +309,7 @@ Merged background:"""
# Fallback: use new_info as background with default disposition
return {
"background": new_info if new_info else current if current else "",
"disposition": DEFAULT_DISPOSITION.copy()
"disposition": DEFAULT_DISPOSITION.copy(),
}
# Validate disposition values
@@ -401,13 +375,15 @@ async def list_banks(pool) -> list:
if isinstance(disposition_data, str):
disposition_data = json.loads(disposition_data)
result.append({
"bank_id": row["bank_id"],
"name": row["name"],
"disposition": disposition_data,
"background": row["background"],
"created_at": row["created_at"].isoformat() if row["created_at"] else None,
"updated_at": row["updated_at"].isoformat() if row["updated_at"] else None,
})
result.append(
{
"bank_id": row["bank_id"],
"name": row["name"],
"disposition": disposition_data,
"background": row["background"],
"created_at": row["created_at"].isoformat() if row["created_at"] else None,
"updated_at": row["updated_at"].isoformat() if row["updated_at"] else None,
}
)
return result
@@ -3,20 +3,15 @@ Chunk storage for retain pipeline.
Handles storage of document chunks in the database.
"""
import logging
from typing import List, Dict, Optional
from .types import ChunkMetadata
logger = logging.getLogger(__name__)
async def store_chunks_batch(
conn,
bank_id: str,
document_id: str,
chunks: List[ChunkMetadata]
) -> Dict[int, str]:
async def store_chunks_batch(conn, bank_id: str, document_id: str, chunks: list[ChunkMetadata]) -> dict[int, str]:
"""
Store document chunks in the database.
@@ -55,16 +50,13 @@ async def store_chunks_batch(
[document_id] * len(chunk_texts),
[bank_id] * len(chunk_texts),
chunk_texts,
chunk_indices
chunk_indices,
)
return chunk_id_map
def map_facts_to_chunks(
facts_chunk_indices: List[int],
chunk_id_map: Dict[int, str]
) -> List[Optional[str]]:
def map_facts_to_chunks(facts_chunk_indices: list[int], chunk_id_map: dict[int, str]) -> list[str | None]:
"""
Map fact chunk indices to chunk IDs.
@@ -3,22 +3,17 @@ Deduplication logic for retain pipeline.
Checks for duplicate facts using semantic similarity and temporal proximity.
"""
import logging
from datetime import datetime
from typing import List
from collections import defaultdict
from datetime import UTC
from .types import ProcessedFact
logger = logging.getLogger(__name__)
async def check_duplicates_batch(
conn,
bank_id: str,
facts: List[ProcessedFact],
duplicate_checker_fn
) -> List[bool]:
async def check_duplicates_batch(conn, bank_id: str, facts: list[ProcessedFact], duplicate_checker_fn) -> list[bool]:
"""
Check which facts are duplicates using batched time-window queries.
@@ -47,16 +42,12 @@ async def check_duplicates_batch(
# Defensive: if both are None (shouldn't happen), use now()
if fact_date is None:
from datetime import datetime, timezone
fact_date = datetime.now(timezone.utc)
from datetime import datetime
fact_date = datetime.now(UTC)
# Round to 12-hour bucket to group similar times
bucket_key = fact_date.replace(
hour=(fact_date.hour // 12) * 12,
minute=0,
second=0,
microsecond=0
)
bucket_key = fact_date.replace(hour=(fact_date.hour // 12) * 12, minute=0, second=0, microsecond=0)
time_buckets[bucket_key].append((idx, fact))
# Process each bucket in batch
@@ -68,14 +59,7 @@ async def check_duplicates_batch(
embeddings = [item[1].embedding for item in bucket_items]
# Check duplicates for this time bucket
dup_flags = await duplicate_checker_fn(
conn,
bank_id,
texts,
embeddings,
bucket_date,
time_window_hours=24
)
dup_flags = await duplicate_checker_fn(conn, bank_id, texts, embeddings, bucket_date, time_window_hours=24)
# Map results back to original indices
for idx, is_dup in zip(indices, dup_flags):
@@ -84,10 +68,7 @@ async def check_duplicates_batch(
return all_is_duplicate
def filter_duplicates(
facts: List[ProcessedFact],
is_duplicate_flags: List[bool]
) -> List[ProcessedFact]:
def filter_duplicates(facts: list[ProcessedFact], is_duplicate_flags: list[bool]) -> list[ProcessedFact]:
"""
Filter out duplicate facts based on duplicate flags.
@@ -3,9 +3,8 @@ Embedding processing for retain pipeline.
Handles augmenting fact texts with temporal information and generating embeddings.
"""
import logging
from typing import List
from datetime import datetime
from . import embedding_utils
from .types import ExtractedFact
@@ -13,7 +12,7 @@ from .types import ExtractedFact
logger = logging.getLogger(__name__)
def augment_texts_with_dates(facts: List[ExtractedFact], format_date_fn) -> List[str]:
def augment_texts_with_dates(facts: list[ExtractedFact], format_date_fn) -> list[str]:
"""
Augment fact texts with readable dates for better temporal matching.
@@ -37,10 +36,7 @@ def augment_texts_with_dates(facts: List[ExtractedFact], format_date_fn) -> List
return augmented_texts
async def generate_embeddings_batch(
embeddings_model,
texts: List[str]
) -> List[List[float]]:
async def generate_embeddings_batch(embeddings_model, texts: list[str]) -> list[list[float]]:
"""
Generate embeddings for a batch of texts.
@@ -54,9 +50,6 @@ async def generate_embeddings_batch(
if not texts:
return []
embeddings = await embedding_utils.generate_embeddings_batch(
embeddings_model,
texts
)
embeddings = await embedding_utils.generate_embeddings_batch(embeddings_model, texts)
return embeddings
@@ -4,12 +4,11 @@ Embedding generation utilities for memory units.
import asyncio
import logging
from typing import List
logger = logging.getLogger(__name__)
def generate_embedding(embeddings_backend, text: str) -> List[float]:
def generate_embedding(embeddings_backend, text: str) -> list[float]:
"""
Generate embedding for text using the provided embeddings backend.
@@ -27,7 +26,7 @@ def generate_embedding(embeddings_backend, text: str) -> List[float]:
raise Exception(f"Failed to generate embedding: {str(e)}")
async def generate_embeddings_batch(embeddings_backend, texts: List[str]) -> List[List[float]]:
async def generate_embeddings_batch(embeddings_backend, texts: list[str]) -> list[list[float]]:
"""
Generate embeddings for multiple texts using the provided embeddings backend.
@@ -47,7 +46,7 @@ async def generate_embeddings_batch(embeddings_backend, texts: List[str]) -> Lis
embeddings = await loop.run_in_executor(
None, # Use default thread pool
embeddings_backend.encode,
texts
texts,
)
return embeddings
except Exception as e:
@@ -3,24 +3,18 @@ Entity processing for retain pipeline.
Handles entity extraction, resolution, and link creation for stored facts.
"""
import logging
from typing import List, Tuple, Dict, Any
from uuid import UUID
from .types import ProcessedFact, EntityRef, EntityLink
import logging
from . import link_utils
from .types import EntityLink, ProcessedFact
logger = logging.getLogger(__name__)
async def process_entities_batch(
entity_resolver,
conn,
bank_id: str,
unit_ids: List[str],
facts: List[ProcessedFact],
log_buffer: List[str] = None
) -> List[EntityLink]:
entity_resolver, conn, bank_id: str, unit_ids: list[str], facts: list[ProcessedFact], log_buffer: list[str] = None
) -> list[EntityLink]:
"""
Process entities for all facts and create entity links.
@@ -53,8 +47,7 @@ async def process_entities_batch(
fact_dates = [fact.occurred_start if fact.occurred_start is not None else fact.mentioned_at for fact in facts]
# Convert EntityRef objects to dict format expected by link_utils
entities_per_fact = [
[{'text': entity.name, 'type': 'CONCEPT'} for entity in (fact.entities or [])]
for fact in facts
[{"text": entity.name, "type": "CONCEPT"} for entity in (fact.entities or [])] for fact in facts
]
# Use existing link_utils function for entity processing
@@ -67,16 +60,13 @@ async def process_entities_batch(
"", # context (not used in current implementation)
fact_dates,
entities_per_fact,
log_buffer # Pass log_buffer for detailed logging
log_buffer, # Pass log_buffer for detailed logging
)
return entity_links
async def insert_entity_links_batch(
conn,
entity_links: List[EntityLink]
) -> None:
async def insert_entity_links_batch(conn, entity_links: list[EntityLink]) -> None:
"""
Insert entity links in batch.
@@ -4,16 +4,17 @@ Fact extraction from text using LLM.
Extracts semantic facts, entities, and temporal information from text.
Uses the LLMConfig wrapper for all LLM calls.
"""
import logging
import os
import json
import re
import asyncio
import json
import logging
import re
from datetime import datetime, timedelta
from typing import List, Dict, Optional, Literal
from openai import AsyncOpenAI
from pydantic import BaseModel, Field, field_validator, ConfigDict
from ..llm_wrapper import OutputTooLongError, LLMConfig
from typing import Literal
from pydantic import BaseModel, ConfigDict, Field, field_validator
from ..llm_wrapper import LLMConfig, OutputTooLongError
def _sanitize_text(text: str) -> str:
@@ -31,11 +32,12 @@ def _sanitize_text(text: str) -> str:
return text
# Remove surrogate characters (U+D800 to U+DFFF) using regex
# These are invalid in UTF-8 and cause encoding errors
return re.sub(r'[\ud800-\udfff]', '', text)
return re.sub(r"[\ud800-\udfff]", "", text)
class Entity(BaseModel):
"""An entity extracted from text."""
text: str = Field(
description="The specific, named entity as it appears in the fact. Must be a proper noun or specific identifier."
)
@@ -48,42 +50,46 @@ class Fact(BaseModel):
This is what fact_extraction returns and what the rest of the pipeline expects.
Combined fact text format: "what | when | where | who | why"
"""
# Required fields
fact: str = Field(description="Combined fact text: what | when | where | who | why")
fact_type: Literal["world", "experience", "opinion"] = Field(description="Perspective: world/experience/opinion")
# Optional temporal fields
occurred_start: Optional[str] = None
occurred_end: Optional[str] = None
mentioned_at: Optional[str] = None
occurred_start: str | None = None
occurred_end: str | None = None
mentioned_at: str | None = None
# Optional location field
where: Optional[str] = Field(None, description="WHERE the fact occurred or is about (specific location, place, or area)")
where: str | None = Field(
None, description="WHERE the fact occurred or is about (specific location, place, or area)"
)
# Optional structured data
entities: Optional[List[Entity]] = None
causal_relations: Optional[List['CausalRelation']] = None
entities: list[Entity] | None = None
causal_relations: list["CausalRelation"] | None = None
class CausalRelation(BaseModel):
"""Causal relationship between facts."""
target_fact_index: int = Field(
description="Index of the related fact in the facts array (0-based). "
"This creates a directed causal link to another fact in the extraction."
"This creates a directed causal link to another fact in the extraction."
)
relation_type: Literal["causes", "caused_by", "enables", "prevents"] = Field(
description="Type of causal relationship: "
"'causes' = this fact directly causes the target fact, "
"'caused_by' = this fact was caused by the target fact, "
"'enables' = this fact enables/allows the target fact, "
"'prevents' = this fact prevents/blocks the target fact"
"'causes' = this fact directly causes the target fact, "
"'caused_by' = this fact was caused by the target fact, "
"'enables' = this fact enables/allows the target fact, "
"'prevents' = this fact prevents/blocks the target fact"
)
strength: float = Field(
description="Strength of causal relationship (0.0 to 1.0). "
"1.0 = direct/strong causation, 0.5 = moderate, 0.3 = weak/indirect",
"1.0 = direct/strong causation, 0.5 = moderate, 0.3 = weak/indirect",
ge=0.0,
le=1.0,
default=1.0
default=1.0,
)
@@ -92,9 +98,7 @@ class ExtractedFact(BaseModel):
model_config = ConfigDict(
json_schema_mode="validation",
json_schema_extra={
"required": ["what", "when", "where", "who", "why", "fact_type"]
}
json_schema_extra={"required": ["what", "when", "where", "who", "why", "fact_type"]},
)
# ==========================================================================
@@ -103,43 +107,43 @@ class ExtractedFact(BaseModel):
what: str = Field(
description="WHAT happened - COMPLETE, DETAILED description with ALL specifics. "
"NEVER summarize or omit details. Include: exact actions, objects, quantities, specifics. "
"BE VERBOSE - capture every detail that was mentioned. "
"Example: 'Emily got married to Sarah at a rooftop garden ceremony with 50 guests attending and a live jazz band playing' "
"NOT: 'A wedding happened' or 'Emily got married'"
"NEVER summarize or omit details. Include: exact actions, objects, quantities, specifics. "
"BE VERBOSE - capture every detail that was mentioned. "
"Example: 'Emily got married to Sarah at a rooftop garden ceremony with 50 guests attending and a live jazz band playing' "
"NOT: 'A wedding happened' or 'Emily got married'"
)
when: str = Field(
description="WHEN it happened - ALWAYS include temporal information if mentioned. "
"Include: specific dates, times, durations, relative time references. "
"Examples: 'on June 15th, 2024 at 3pm', 'last weekend', 'for the past 3 years', 'every morning at 6am'. "
"Write 'N/A' ONLY if absolutely no temporal context exists. Prefer converting to absolute dates when possible."
"Include: specific dates, times, durations, relative time references. "
"Examples: 'on June 15th, 2024 at 3pm', 'last weekend', 'for the past 3 years', 'every morning at 6am'. "
"Write 'N/A' ONLY if absolutely no temporal context exists. Prefer converting to absolute dates when possible."
)
where: str = Field(
description="WHERE it happened or is about - SPECIFIC locations, places, areas, regions if applicable. "
"Include: cities, neighborhoods, venues, buildings, countries, specific addresses when mentioned. "
"Examples: 'downtown San Francisco at a rooftop garden venue', 'at the user's home in Brooklyn', 'online via Zoom', 'Paris, France'. "
"Write 'N/A' ONLY if absolutely no location context exists or if the fact is completely location-agnostic."
"Include: cities, neighborhoods, venues, buildings, countries, specific addresses when mentioned. "
"Examples: 'downtown San Francisco at a rooftop garden venue', 'at the user's home in Brooklyn', 'online via Zoom', 'Paris, France'. "
"Write 'N/A' ONLY if absolutely no location context exists or if the fact is completely location-agnostic."
)
who: str = Field(
description="WHO is involved - ALL people/entities with FULL context and relationships. "
"Include: names, roles, relationships to user, background details. "
"Resolve coreferences (if 'my roommate' is later named 'Emily', write 'Emily, the user's college roommate'). "
"BE DETAILED about relationships and roles. "
"Example: 'Emily (user's college roommate from Stanford, now works at Google), Sarah (Emily's partner of 5 years, software engineer)' "
"NOT: 'my friend' or 'Emily and Sarah'"
"Include: names, roles, relationships to user, background details. "
"Resolve coreferences (if 'my roommate' is later named 'Emily', write 'Emily, the user's college roommate'). "
"BE DETAILED about relationships and roles. "
"Example: 'Emily (user's college roommate from Stanford, now works at Google), Sarah (Emily's partner of 5 years, software engineer)' "
"NOT: 'my friend' or 'Emily and Sarah'"
)
why: str = Field(
description="WHY it matters - ALL emotional, contextual, and motivational details. "
"Include EVERYTHING: feelings, preferences, motivations, observations, context, background, significance. "
"BE VERBOSE - capture all the nuance and meaning. "
"FOR ASSISTANT FACTS: MUST include what the user asked/requested that led to this interaction! "
"Example (world): 'The user felt thrilled and inspired, has always dreamed of an outdoor ceremony, mentioned wanting a similar garden venue, was particularly moved by the intimate atmosphere and personal vows' "
"Example (assistant): 'User asked how to fix slow API performance with 1000+ concurrent users, expected 70-80% reduction in database load' "
"NOT: 'User liked it' or 'To help user'"
"Include EVERYTHING: feelings, preferences, motivations, observations, context, background, significance. "
"BE VERBOSE - capture all the nuance and meaning. "
"FOR ASSISTANT FACTS: MUST include what the user asked/requested that led to this interaction! "
"Example (world): 'The user felt thrilled and inspired, has always dreamed of an outdoor ceremony, mentioned wanting a similar garden venue, was particularly moved by the intimate atmosphere and personal vows' "
"Example (assistant): 'User asked how to fix slow API performance with 1000+ concurrent users, expected 70-80% reduction in database load' "
"NOT: 'User liked it' or 'To help user'"
)
# ==========================================================================
@@ -148,17 +152,17 @@ class ExtractedFact(BaseModel):
fact_kind: str = Field(
default="conversation",
description="'event' = specific datable occurrence (set occurred dates), 'conversation' = general info (no occurred dates)"
description="'event' = specific datable occurrence (set occurred dates), 'conversation' = general info (no occurred dates)",
)
# Temporal fields - optional
occurred_start: Optional[str] = Field(
occurred_start: str | None = Field(
default=None,
description="WHEN the event happened (ISO timestamp). Only for fact_kind='event'. Leave null for conversations."
description="WHEN the event happened (ISO timestamp). Only for fact_kind='event'. Leave null for conversations.",
)
occurred_end: Optional[str] = Field(
occurred_end: str | None = Field(
default=None,
description="WHEN the event ended (ISO timestamp). Only for events with duration. Leave null for conversations."
description="WHEN the event ended (ISO timestamp). Only for events with duration. Leave null for conversations.",
)
# Classification (CRITICAL - required)
@@ -168,16 +172,15 @@ class ExtractedFact(BaseModel):
)
# Entities - extracted from fact content
entities: Optional[List[Entity]] = Field(
entities: list[Entity] | None = Field(
default=None,
description="Named entities, objects, AND abstract concepts from the fact. Include: people names, organizations, places, significant objects (e.g., 'coffee maker', 'car'), AND abstract concepts/themes (e.g., 'friendship', 'career growth', 'loss', 'celebration'). Extract anything that could help link related facts together."
description="Named entities, objects, AND abstract concepts from the fact. Include: people names, organizations, places, significant objects (e.g., 'coffee maker', 'car'), AND abstract concepts/themes (e.g., 'friendship', 'career growth', 'loss', 'celebration'). Extract anything that could help link related facts together.",
)
causal_relations: Optional[List[CausalRelation]] = Field(
default=None,
description="Causal links to other facts. Can be null."
causal_relations: list[CausalRelation] | None = Field(
default=None, description="Causal links to other facts. Can be null."
)
@field_validator('entities', mode='before')
@field_validator("entities", mode="before")
@classmethod
def ensure_entities_list(cls, v):
"""Ensure entities is always a list (convert None to empty list)."""
@@ -185,7 +188,7 @@ class ExtractedFact(BaseModel):
return []
return v
@field_validator('causal_relations', mode='before')
@field_validator("causal_relations", mode="before")
@classmethod
def ensure_causal_relations_list(cls, v):
"""Ensure causal_relations is always a list (convert None to empty list)."""
@@ -198,11 +201,11 @@ class ExtractedFact(BaseModel):
parts = [self.what]
# Add 'who' if not N/A
if self.who and self.who.upper() != 'N/A':
if self.who and self.who.upper() != "N/A":
parts.append(f"Involving: {self.who}")
# Add 'why' if not N/A
if self.why and self.why.upper() != 'N/A':
if self.why and self.why.upper() != "N/A":
parts.append(self.why)
if len(parts) == 1:
@@ -213,12 +216,11 @@ class ExtractedFact(BaseModel):
class FactExtractionResponse(BaseModel):
"""Response containing all extracted facts."""
facts: List[ExtractedFact] = Field(
description="List of extracted factual statements"
)
facts: list[ExtractedFact] = Field(description="List of extracted factual statements")
def chunk_text(text: str, max_chars: int) -> List[str]:
def chunk_text(text: str, max_chars: int) -> list[str]:
"""
Split text into chunks, preserving conversation structure when possible.
@@ -232,7 +234,6 @@ def chunk_text(text: str, max_chars: int) -> List[str]:
Returns:
List of text chunks, roughly under max_chars
"""
import json
from langchain_text_splitters import RecursiveCharacterTextSplitter
# If text is small enough, return as-is
@@ -256,21 +257,21 @@ def chunk_text(text: str, max_chars: int) -> List[str]:
is_separator_regex=False,
separators=[
"\n\n", # Paragraph breaks
"\n", # Line breaks
". ", # Sentence endings
"! ", # Exclamations
"? ", # Questions
"; ", # Semicolons
", ", # Commas
" ", # Words
"", # Characters (last resort)
"\n", # Line breaks
". ", # Sentence endings
"! ", # Exclamations
"? ", # Questions
"; ", # Semicolons
", ", # Commas
" ", # Words
"", # Characters (last resort)
],
)
return splitter.split_text(text)
def _chunk_conversation(turns: List[dict], max_chars: int) -> List[str]:
def _chunk_conversation(turns: list[dict], max_chars: int) -> list[str]:
"""
Chunk a conversation array at turn boundaries, preserving complete turns.
@@ -281,7 +282,6 @@ def _chunk_conversation(turns: List[dict], max_chars: int) -> List[str]:
Returns:
List of JSON-serialized chunks, each containing complete turns
"""
import json
chunks = []
current_chunk = []
@@ -315,10 +315,10 @@ async def _extract_facts_from_chunk(
total_chunks: int,
event_date: datetime,
context: str,
llm_config: 'LLMConfig',
llm_config: "LLMConfig",
agent_name: str = None,
extract_opinions: bool = False
) -> List[Dict[str, str]]:
extract_opinions: bool = False,
) -> list[dict[str, str]]:
"""
Extract facts from a single chunk (internal helper for parallel processing).
@@ -333,7 +333,9 @@ async def _extract_facts_from_chunk(
# Opinion extraction uses a separate prompt (not this one)
fact_types_instruction = "Extract ONLY 'opinion' type facts (formed opinions, beliefs, and perspectives). DO NOT extract 'world' or 'assistant' facts."
else:
fact_types_instruction = "Extract ONLY 'world' and 'assistant' type facts. DO NOT extract opinions - those are extracted separately."
fact_types_instruction = (
"Extract ONLY 'world' and 'assistant' type facts. DO NOT extract opinions - those are extracted separately."
)
prompt = f"""Extract facts from text into structured format with FOUR required dimensions - BE EXTREMELY DETAILED.
@@ -534,10 +536,8 @@ WHAT TO EXTRACT vs SKIP
✅ EXTRACT: User preferences (ALWAYS as separate facts!), feelings, plans, events, relationships, achievements
❌ SKIP: Greetings, filler ("thanks", "cool"), purely structural statements"""
import logging
from openai import BadRequestError
logger = logging.getLogger(__name__)
@@ -548,11 +548,11 @@ WHAT TO EXTRACT vs SKIP
# Sanitize input text to prevent Unicode encoding errors (e.g., unpaired surrogates)
sanitized_chunk = _sanitize_text(chunk)
sanitized_context = _sanitize_text(context) if context else 'none'
sanitized_context = _sanitize_text(context) if context else "none"
# Build user message with metadata and chunk content in a clear format
# Format event_date with day of week for better temporal reasoning
event_date_formatted = event_date.strftime('%A, %B %d, %Y') # e.g., "Monday, June 10, 2024"
event_date_formatted = event_date.strftime("%A, %B %d, %Y") # e.g., "Monday, June 10, 2024"
user_message = f"""Extract facts from the following text chunk.
{memory_bank_context}
@@ -566,16 +566,7 @@ Text:
for attempt in range(max_retries):
try:
extraction_response_json = await llm_config.call(
messages=[
{
"role": "system",
"content": prompt
},
{
"role": "user",
"content": user_message
}
],
messages=[{"role": "system", "content": prompt}, {"role": "user", "content": user_message}],
response_format=FactExtractionResponse,
scope="memory_extract_facts",
temperature=0.1,
@@ -601,7 +592,7 @@ Text:
)
return []
raw_facts = extraction_response_json.get('facts', [])
raw_facts = extraction_response_json.get("facts", [])
if not raw_facts:
logger.debug(
f"LLM response missing 'facts' field or returned empty list. "
@@ -622,48 +613,48 @@ Text:
# Helper to get non-empty value
def get_value(field_name):
value = llm_fact.get(field_name)
if value and value != '' and value != [] and value != {} and str(value).upper() != 'N/A':
if value and value != "" and value != [] and value != {} and str(value).upper() != "N/A":
return value
return None
# NEW FORMAT: what, when, who, why (all required)
what = get_value('what')
when = get_value('when')
who = get_value('who')
why = get_value('why')
what = get_value("what")
when = get_value("when")
who = get_value("who")
why = get_value("why")
# Fallback to old format if new fields not present
if not what:
what = get_value('factual_core')
what = get_value("factual_core")
if not what:
logger.warning(f"Skipping fact {i}: missing 'what' field")
continue
# Critical field: fact_type
# LLM uses "assistant" but we convert to "experience" for storage
fact_type = llm_fact.get('fact_type')
fact_type = llm_fact.get("fact_type")
# Convert "assistant" → "experience" for storage
if fact_type == 'assistant':
fact_type = 'experience'
if fact_type == "assistant":
fact_type = "experience"
# Validate fact_type (after conversion)
if fact_type not in ['world', 'experience', 'opinion']:
if fact_type not in ["world", "experience", "opinion"]:
# Try to fix common mistakes - check if they swapped fact_type and fact_kind
fact_kind = llm_fact.get('fact_kind')
if fact_kind == 'assistant':
fact_type = 'experience'
elif fact_kind in ['world', 'experience', 'opinion']:
fact_kind = llm_fact.get("fact_kind")
if fact_kind == "assistant":
fact_type = "experience"
elif fact_kind in ["world", "experience", "opinion"]:
fact_type = fact_kind
else:
# Default to 'world' if we can't determine
fact_type = 'world'
fact_type = "world"
logger.warning(f"Fact {i}: defaulting to fact_type='world'")
# Get fact_kind for temporal handling (but don't store it)
fact_kind = llm_fact.get('fact_kind', 'conversation')
if fact_kind not in ['conversation', 'event', 'other']:
fact_kind = 'conversation'
fact_kind = llm_fact.get("fact_kind", "conversation")
if fact_kind not in ["conversation", "event", "other"]:
fact_kind = "conversation"
# Build combined fact text from the 4 dimensions: what | when | who | why
fact_data = {}
@@ -682,20 +673,20 @@ Text:
# Add temporal fields
# For events: occurred_start/occurred_end (when the event happened)
if fact_kind == 'event':
occurred_start = get_value('occurred_start')
occurred_end = get_value('occurred_end')
if fact_kind == "event":
occurred_start = get_value("occurred_start")
occurred_end = get_value("occurred_end")
if occurred_start:
fact_data['occurred_start'] = occurred_start
fact_data["occurred_start"] = occurred_start
# For point events: if occurred_end not set, default to occurred_start
if occurred_end:
fact_data['occurred_end'] = occurred_end
fact_data["occurred_end"] = occurred_end
else:
fact_data['occurred_end'] = occurred_start
fact_data["occurred_end"] = occurred_start
# Add entities if present (validate as Entity objects)
# LLM sometimes returns strings instead of {"text": "..."} format
entities = get_value('entities')
entities = get_value("entities")
if entities:
# Validate and normalize each entity
validated_entities = []
@@ -703,38 +694,34 @@ Text:
if isinstance(ent, str):
# Normalize string to Entity object
validated_entities.append(Entity(text=ent))
elif isinstance(ent, dict) and 'text' in ent:
elif isinstance(ent, dict) and "text" in ent:
try:
validated_entities.append(Entity.model_validate(ent))
except Exception as e:
logger.warning(f"Invalid entity {ent}: {e}")
if validated_entities:
fact_data['entities'] = validated_entities
fact_data["entities"] = validated_entities
# Add causal relations if present (validate as CausalRelation objects)
# Filter out invalid relations (missing required fields)
causal_relations = get_value('causal_relations')
causal_relations = get_value("causal_relations")
if causal_relations:
validated_relations = []
for rel in causal_relations:
if isinstance(rel, dict) and 'target_fact_index' in rel and 'relation_type' in rel:
if isinstance(rel, dict) and "target_fact_index" in rel and "relation_type" in rel:
try:
validated_relations.append(CausalRelation.model_validate(rel))
except Exception as e:
logger.warning(f"Invalid causal relation {rel}: {e}")
if validated_relations:
fact_data['causal_relations'] = validated_relations
fact_data["causal_relations"] = validated_relations
# Always set mentioned_at to the event_date (when the conversation/document occurred)
fact_data['mentioned_at'] = event_date.isoformat()
fact_data["mentioned_at"] = event_date.isoformat()
# Build Fact model instance
try:
fact = Fact(
fact=combined_text,
fact_type=fact_type,
**fact_data
)
fact = Fact(fact=combined_text, fact_type=fact_type, **fact_data)
chunk_facts.append(fact)
except Exception as e:
logger.error(f"Failed to create Fact model for fact {i}: {e}")
@@ -753,7 +740,9 @@ Text:
except BadRequestError as e:
last_error = e
if "json_validate_failed" in str(e):
logger.warning(f" [1.3.{chunk_index + 1}] Attempt {attempt + 1}/{max_retries} failed with JSON validation error: {e}")
logger.warning(
f" [1.3.{chunk_index + 1}] Attempt {attempt + 1}/{max_retries} failed with JSON validation error: {e}"
)
if attempt < max_retries - 1:
logger.info(f" [1.3.{chunk_index + 1}] Retrying...")
continue
@@ -772,8 +761,8 @@ async def _extract_facts_with_auto_split(
context: str,
llm_config: LLMConfig,
agent_name: str = None,
extract_opinions: bool = False
) -> List[Dict[str, str]]:
extract_opinions: bool = False,
) -> list[dict[str, str]]:
"""
Extract facts from a chunk with automatic splitting if output exceeds token limits.
@@ -794,6 +783,7 @@ async def _extract_facts_with_auto_split(
List of fact dictionaries extracted from the chunk (possibly from sub-chunks)
"""
import logging
logger = logging.getLogger(__name__)
try:
@@ -806,9 +796,9 @@ async def _extract_facts_with_auto_split(
context=context,
llm_config=llm_config,
agent_name=agent_name,
extract_opinions=extract_opinions
extract_opinions=extract_opinions,
)
except OutputTooLongError as e:
except OutputTooLongError:
# Output exceeded token limits - split the chunk in half and retry
logger.warning(
f"Output too long for chunk {chunk_index + 1}/{total_chunks} "
@@ -824,7 +814,7 @@ async def _extract_facts_with_auto_split(
search_start = max(0, mid_point - search_range)
search_end = min(len(chunk), mid_point + search_range)
sentence_endings = ['. ', '! ', '? ', '\n\n']
sentence_endings = [". ", "! ", "? ", "\n\n"]
best_split = mid_point
for ending in sentence_endings:
@@ -838,8 +828,7 @@ async def _extract_facts_with_auto_split(
second_half = chunk[best_split:].strip()
logger.info(
f"Split chunk {chunk_index + 1} into two sub-chunks: "
f"{len(first_half)} chars and {len(second_half)} chars"
f"Split chunk {chunk_index + 1} into two sub-chunks: {len(first_half)} chars and {len(second_half)} chars"
)
# Process both halves recursively (in parallel)
@@ -852,7 +841,7 @@ async def _extract_facts_with_auto_split(
context=context,
llm_config=llm_config,
agent_name=agent_name,
extract_opinions=extract_opinions
extract_opinions=extract_opinions,
),
_extract_facts_with_auto_split(
chunk=second_half,
@@ -862,8 +851,8 @@ async def _extract_facts_with_auto_split(
context=context,
llm_config=llm_config,
agent_name=agent_name,
extract_opinions=extract_opinions
)
extract_opinions=extract_opinions,
),
]
sub_results = await asyncio.gather(*sub_tasks)
@@ -873,9 +862,7 @@ async def _extract_facts_with_auto_split(
for sub_result in sub_results:
all_facts.extend(sub_result)
logger.info(
f"Successfully extracted {len(all_facts)} facts from split chunk {chunk_index + 1}"
)
logger.info(f"Successfully extracted {len(all_facts)} facts from split chunk {chunk_index + 1}")
return all_facts
@@ -887,7 +874,7 @@ async def extract_facts_from_text(
agent_name: str,
context: str = "",
extract_opinions: bool = False,
) -> tuple[List[Fact], List[tuple[str, int]]]:
) -> tuple[list[Fact], list[tuple[str, int]]]:
"""
Extract semantic facts from conversational or narrative text using LLM.
@@ -920,7 +907,7 @@ async def extract_facts_from_text(
context=context,
llm_config=llm_config,
agent_name=agent_name,
extract_opinions=extract_opinions
extract_opinions=extract_opinions,
)
for i, chunk in enumerate(chunks)
]
@@ -938,8 +925,10 @@ async def extract_facts_from_text(
# ============================================================================
# Import types for the orchestration layer (note: ExtractedFact here is different from the Pydantic model above)
from .types import RetainContent, ExtractedFact as ExtractedFactType, ChunkMetadata, CausalRelation as CausalRelationType
from typing import Tuple
from .types import CausalRelation as CausalRelationType
from .types import ChunkMetadata, RetainContent
from .types import ExtractedFact as ExtractedFactType
logger = logging.getLogger(__name__)
@@ -948,11 +937,8 @@ SECONDS_PER_FACT = 10
async def extract_facts_from_contents(
contents: List[RetainContent],
llm_config,
agent_name: str,
extract_opinions: bool = False
) -> Tuple[List[ExtractedFactType], List[ChunkMetadata]]:
contents: list[RetainContent], llm_config, agent_name: str, extract_opinions: bool = False
) -> tuple[list[ExtractedFactType], list[ChunkMetadata]]:
"""
Extract facts from multiple content items in parallel.
@@ -985,7 +971,7 @@ async def extract_facts_from_contents(
context=item.context,
llm_config=llm_config,
agent_name=agent_name,
extract_opinions=extract_opinions
extract_opinions=extract_opinions,
)
fact_extraction_tasks.append(task)
@@ -993,8 +979,8 @@ async def extract_facts_from_contents(
all_fact_results = await asyncio.gather(*fact_extraction_tasks)
# Step 3: Flatten and convert to typed objects
extracted_facts: List[ExtractedFactType] = []
chunks_metadata: List[ChunkMetadata] = []
extracted_facts: list[ExtractedFactType] = []
chunks_metadata: list[ChunkMetadata] = []
global_chunk_idx = 0
global_fact_idx = 0
@@ -1008,7 +994,7 @@ async def extract_facts_from_contents(
chunk_text=chunk_text,
fact_count=chunk_fact_count,
content_index=content_index,
chunk_index=global_chunk_idx
chunk_index=global_chunk_idx,
)
chunks_metadata.append(chunk_metadata)
global_chunk_idx += 1
@@ -1029,18 +1015,21 @@ async def extract_facts_from_contents(
fact_type=fact_from_llm.fact_type,
entities=[e.text for e in (fact_from_llm.entities or [])],
# occurred_start/end: from LLM only, leave None if not provided
occurred_start=_parse_datetime(fact_from_llm.occurred_start) if fact_from_llm.occurred_start else None,
occurred_end=_parse_datetime(fact_from_llm.occurred_end) if fact_from_llm.occurred_end else None,
occurred_start=_parse_datetime(fact_from_llm.occurred_start)
if fact_from_llm.occurred_start
else None,
occurred_end=_parse_datetime(fact_from_llm.occurred_end)
if fact_from_llm.occurred_end
else None,
causal_relations=_convert_causal_relations(
fact_from_llm.causal_relations or [],
global_fact_idx
fact_from_llm.causal_relations or [], global_fact_idx
),
content_index=content_index,
chunk_index=chunk_global_idx,
context=content.context,
# mentioned_at: always the event_date (when the conversation/document occurred)
mentioned_at=content.event_date,
metadata=content.metadata
metadata=content.metadata,
)
extracted_facts.append(extracted_fact)
@@ -1056,13 +1045,14 @@ async def extract_facts_from_contents(
def _parse_datetime(date_str: str):
"""Parse ISO datetime string."""
from dateutil import parser as date_parser
try:
return date_parser.isoparse(date_str)
except Exception:
return None
def _convert_causal_relations(relations_from_llm, fact_start_idx: int) -> List[CausalRelationType]:
def _convert_causal_relations(relations_from_llm, fact_start_idx: int) -> list[CausalRelationType]:
"""
Convert causal relations from LLM format to ExtractedFact format.
@@ -1073,13 +1063,13 @@ def _convert_causal_relations(relations_from_llm, fact_start_idx: int) -> List[C
causal_relation = CausalRelationType(
relation_type=rel.relation_type,
target_fact_index=fact_start_idx + rel.target_fact_index,
strength=rel.strength
strength=rel.strength,
)
causal_relations.append(causal_relation)
return causal_relations
def _add_temporal_offsets(facts: List[ExtractedFactType], contents: List[RetainContent]) -> None:
def _add_temporal_offsets(facts: list[ExtractedFactType], contents: list[RetainContent]) -> None:
"""
Add time offsets to preserve fact ordering within each content.
@@ -3,10 +3,9 @@ Fact storage for retain pipeline.
Handles insertion of facts into the database.
"""
import logging
import json
from typing import List, Optional
from uuid import UUID
import logging
from .types import ProcessedFact
@@ -14,11 +13,8 @@ logger = logging.getLogger(__name__)
async def insert_facts_batch(
conn,
bank_id: str,
facts: List[ProcessedFact],
document_id: Optional[str] = None
) -> List[str]:
conn, bank_id: str, facts: list[ProcessedFact], document_id: str | None = None
) -> list[str]:
"""
Insert facts into the database in batch.
@@ -62,7 +58,7 @@ async def insert_facts_batch(
contexts.append(fact.context)
fact_types.append(fact.fact_type)
# confidence_score is only for opinion facts
confidence_scores.append(1.0 if fact.fact_type == 'opinion' else None)
confidence_scores.append(1.0 if fact.fact_type == "opinion" else None)
access_counts.append(0) # Initial access count
metadata_jsons.append(json.dumps(fact.metadata))
chunk_ids.append(fact.chunk_id)
@@ -93,10 +89,10 @@ async def insert_facts_batch(
access_counts,
metadata_jsons,
chunk_ids,
document_ids
document_ids,
)
unit_ids = [str(row['id']) for row in results]
unit_ids = [str(row["id"]) for row in results]
return unit_ids
@@ -119,17 +115,12 @@ async def ensure_bank_exists(conn, bank_id: str) -> None:
""",
bank_id,
'{"skepticism": 3, "literalism": 3, "empathy": 3}',
""
"",
)
async def handle_document_tracking(
conn,
bank_id: str,
document_id: str,
combined_content: str,
is_first_batch: bool,
retain_params: Optional[dict] = None
conn, bank_id: str, document_id: str, combined_content: str, is_first_batch: bool, retain_params: dict | None = None
) -> None:
"""
Handle document tracking in the database.
@@ -150,10 +141,7 @@ async def handle_document_tracking(
# Always delete old document first if it exists (cascades to units and links)
# Only delete on the first batch to avoid deleting data we just inserted
if is_first_batch:
await conn.fetchval(
"DELETE FROM documents WHERE id = $1 AND bank_id = $2 RETURNING id",
document_id, bank_id
)
await conn.fetchval("DELETE FROM documents WHERE id = $1 AND bank_id = $2 RETURNING id", document_id, bank_id)
# Insert document (or update if exists from concurrent operations)
await conn.execute(
@@ -172,5 +160,5 @@ async def handle_document_tracking(
combined_content,
content_hash,
json.dumps({}), # Empty metadata dict
json.dumps(retain_params) if retain_params else None
json.dumps(retain_params) if retain_params else None,
)
@@ -3,20 +3,16 @@ Link creation for retain pipeline.
Handles creation of temporal, semantic, and causal links between facts.
"""
import logging
from typing import List
from .types import ProcessedFact, CausalRelation
import logging
from . import link_utils
from .types import ProcessedFact
logger = logging.getLogger(__name__)
async def create_temporal_links_batch(
conn,
bank_id: str,
unit_ids: List[str]
) -> int:
async def create_temporal_links_batch(conn, bank_id: str, unit_ids: list[str]) -> int:
"""
Create temporal links between facts.
@@ -33,20 +29,10 @@ async def create_temporal_links_batch(
if not unit_ids:
return 0
return await link_utils.create_temporal_links_batch_per_fact(
conn,
bank_id,
unit_ids,
log_buffer=[]
)
return await link_utils.create_temporal_links_batch_per_fact(conn, bank_id, unit_ids, log_buffer=[])
async def create_semantic_links_batch(
conn,
bank_id: str,
unit_ids: List[str],
embeddings: List[List[float]]
) -> int:
async def create_semantic_links_batch(conn, bank_id: str, unit_ids: list[str], embeddings: list[list[float]]) -> int:
"""
Create semantic links between facts.
@@ -67,20 +53,10 @@ async def create_semantic_links_batch(
if len(unit_ids) != len(embeddings):
raise ValueError(f"Mismatch between unit_ids ({len(unit_ids)}) and embeddings ({len(embeddings)})")
return await link_utils.create_semantic_links_batch(
conn,
bank_id,
unit_ids,
embeddings,
log_buffer=[]
)
return await link_utils.create_semantic_links_batch(conn, bank_id, unit_ids, embeddings, log_buffer=[])
async def create_causal_links_batch(
conn,
unit_ids: List[str],
facts: List[ProcessedFact]
) -> int:
async def create_causal_links_batch(conn, unit_ids: list[str], facts: list[ProcessedFact]) -> int:
"""
Create causal links between facts.
@@ -108,9 +84,9 @@ async def create_causal_links_batch(
# Convert CausalRelation objects to dicts
relations_dicts = [
{
'relation_type': rel.relation_type,
'target_fact_index': rel.target_fact_index,
'strength': rel.strength
"relation_type": rel.relation_type,
"target_fact_index": rel.target_fact_index,
"strength": rel.strength,
}
for rel in fact.causal_relations
]
@@ -118,10 +94,6 @@ async def create_causal_links_batch(
else:
causal_relations_per_fact.append([])
link_count = await link_utils.create_causal_links_batch(
conn,
unit_ids,
causal_relations_per_fact
)
link_count = await link_utils.create_causal_links_batch(conn, unit_ids, causal_relations_per_fact)
return link_count
@@ -2,10 +2,9 @@
Link creation utilities for temporal, semantic, and entity links.
"""
import time
import logging
from typing import List
from datetime import timedelta, datetime, timezone
import time
from datetime import UTC, datetime, timedelta
from uuid import UUID
from .types import EntityLink
@@ -19,7 +18,7 @@ def _normalize_datetime(dt):
return None
if dt.tzinfo is None:
# Naive datetime - assume UTC
return dt.replace(tzinfo=timezone.utc)
return dt.replace(tzinfo=UTC)
return dt
@@ -54,24 +53,26 @@ def compute_temporal_links(
try:
time_lower = unit_event_date_norm - timedelta(hours=time_window_hours)
except OverflowError:
time_lower = datetime.min.replace(tzinfo=timezone.utc)
time_lower = datetime.min.replace(tzinfo=UTC)
try:
time_upper = unit_event_date_norm + timedelta(hours=time_window_hours)
except OverflowError:
time_upper = datetime.max.replace(tzinfo=timezone.utc)
time_upper = datetime.max.replace(tzinfo=UTC)
# Filter candidates within this unit's time window
matching_neighbors = [
(row['id'], row['event_date'])
(row["id"], row["event_date"])
for row in candidates
if time_lower <= _normalize_datetime(row['event_date']) <= time_upper
if time_lower <= _normalize_datetime(row["event_date"]) <= time_upper
][:10] # Limit to top 10
for recent_id, recent_event_date in matching_neighbors:
# Calculate temporal proximity weight
time_diff_hours = abs((unit_event_date_norm - _normalize_datetime(recent_event_date)).total_seconds() / 3600)
time_diff_hours = abs(
(unit_event_date_norm - _normalize_datetime(recent_event_date)).total_seconds() / 3600
)
weight = max(0.3, 1.0 - (time_diff_hours / time_window_hours))
links.append((unit_id, str(recent_id), 'temporal', weight, None))
links.append((unit_id, str(recent_id), "temporal", weight, None))
return links
@@ -99,17 +100,17 @@ def compute_temporal_query_bounds(
try:
min_date = min(all_dates) - timedelta(hours=time_window_hours)
except OverflowError:
min_date = datetime.min.replace(tzinfo=timezone.utc)
min_date = datetime.min.replace(tzinfo=UTC)
try:
max_date = max(all_dates) + timedelta(hours=time_window_hours)
except OverflowError:
max_date = datetime.max.replace(tzinfo=timezone.utc)
max_date = datetime.max.replace(tzinfo=UTC)
return min_date, max_date
def _log(log_buffer, message, level='info'):
def _log(log_buffer, message, level="info"):
"""Helper to log to buffer if available, otherwise use logger.
Args:
@@ -117,7 +118,7 @@ def _log(log_buffer, message, level='info'):
message: The log message
level: 'info', 'debug', 'warning', or 'error'. Debug messages are not added to buffer.
"""
if level == 'debug':
if level == "debug":
# Debug messages only go to logger, not to buffer
logger.debug(message)
return
@@ -125,23 +126,23 @@ def _log(log_buffer, message, level='info'):
if log_buffer is not None:
log_buffer.append(message)
else:
if level == 'info':
if level == "info":
logger.info(message)
else:
logger.log(logging.WARNING if level == 'warning' else logging.ERROR, message)
logger.log(logging.WARNING if level == "warning" else logging.ERROR, message)
async def extract_entities_batch_optimized(
entity_resolver,
conn,
bank_id: str,
unit_ids: List[str],
sentences: List[str],
unit_ids: list[str],
sentences: list[str],
context: str,
fact_dates: List,
llm_entities: List[List[dict]],
log_buffer: List[str] = None,
) -> List[tuple]:
fact_dates: list,
llm_entities: list[list[dict]],
log_buffer: list[str] = None,
) -> list[tuple]:
"""
Process LLM-extracted entities for ALL facts in batch.
@@ -171,15 +172,19 @@ async def extract_entities_batch_optimized(
formatted_entities = []
for ent in entity_list:
# Handle both Entity objects and dicts
if hasattr(ent, 'text'):
if hasattr(ent, "text"):
# Entity objects only have 'text', default type to 'CONCEPT'
formatted_entities.append({'text': ent.text, 'type': 'CONCEPT'})
formatted_entities.append({"text": ent.text, "type": "CONCEPT"})
elif isinstance(ent, dict):
formatted_entities.append({'text': ent.get('text', ''), 'type': ent.get('type', 'CONCEPT')})
formatted_entities.append({"text": ent.get("text", ""), "type": ent.get("type", "CONCEPT")})
all_entities.append(formatted_entities)
total_entities = sum(len(ents) for ents in all_entities)
_log(log_buffer, f" [6.1] Process LLM entities: {total_entities} entities from {len(sentences)} facts in {time.time() - substep_start:.3f}s", level='debug')
_log(
log_buffer,
f" [6.1] Process LLM entities: {total_entities} entities from {len(sentences)} facts in {time.time() - substep_start:.3f}s",
level="debug",
)
# Step 2: Resolve entities in BATCH (much faster!)
substep_start = time.time()
@@ -195,13 +200,19 @@ async def extract_entities_batch_optimized(
continue
for local_idx, entity in enumerate(entities):
all_entities_flat.append({
'text': entity['text'],
'type': entity['type'],
'nearby_entities': entities,
})
all_entities_flat.append(
{
"text": entity["text"],
"type": entity["type"],
"nearby_entities": entities,
}
)
entity_to_unit.append((unit_id, local_idx, fact_date))
_log(log_buffer, f" [6.2.1] Prepare entities: {len(all_entities_flat)} entities in {time.time() - substep_6_2_1_start:.3f}s", level='debug')
_log(
log_buffer,
f" [6.2.1] Prepare entities: {len(all_entities_flat)} entities in {time.time() - substep_6_2_1_start:.3f}s",
level="debug",
)
# Resolve ALL entities in one batch call
if all_entities_flat:
@@ -210,7 +221,7 @@ async def extract_entities_batch_optimized(
# Add per-entity dates to entity data for batch resolution
for idx, (unit_id, local_idx, fact_date) in enumerate(entity_to_unit):
all_entities_flat[idx]['event_date'] = fact_date
all_entities_flat[idx]["event_date"] = fact_date
# Resolve ALL entities in ONE batch call (much faster than sequential buckets)
# INSERT ... ON CONFLICT handles any race conditions at the DB level
@@ -219,10 +230,14 @@ async def extract_entities_batch_optimized(
entities_data=all_entities_flat,
context=context,
unit_event_date=None, # Not used when per-entity dates provided
conn=conn # Use main transaction connection
conn=conn, # Use main transaction connection
)
_log(log_buffer, f" [6.2.2] Resolve entities: {len(all_entities_flat)} entities in single batch in {time.time() - substep_6_2_2_start:.3f}s", level='debug')
_log(
log_buffer,
f" [6.2.2] Resolve entities: {len(all_entities_flat)} entities in single batch in {time.time() - substep_6_2_2_start:.3f}s",
level="debug",
)
# [6.2.3] Create unit-entity links in BATCH
substep_6_2_3_start = time.time()
@@ -239,12 +254,24 @@ async def extract_entities_batch_optimized(
# Batch insert all unit-entity links (MUCH faster!)
await entity_resolver.link_units_to_entities_batch(unit_entity_pairs, conn=conn)
_log(log_buffer, f" [6.2.3] Create unit-entity links (batched): {len(unit_entity_pairs)} links in {time.time() - substep_6_2_3_start:.3f}s", level='debug')
_log(
log_buffer,
f" [6.2.3] Create unit-entity links (batched): {len(unit_entity_pairs)} links in {time.time() - substep_6_2_3_start:.3f}s",
level="debug",
)
_log(log_buffer, f" [6.2] Entity resolution (batched): {len(all_entities_flat)} entities resolved in {time.time() - step_6_2_start:.3f}s", level='debug')
_log(
log_buffer,
f" [6.2] Entity resolution (batched): {len(all_entities_flat)} entities resolved in {time.time() - step_6_2_start:.3f}s",
level="debug",
)
else:
unit_to_entity_ids = {}
_log(log_buffer, f" [6.2] Entity resolution (batched): 0 entities in {time.time() - step_6_2_start:.3f}s", level='debug')
_log(
log_buffer,
f" [6.2] Entity resolution (batched): 0 entities in {time.time() - step_6_2_start:.3f}s",
level="debug",
)
# Step 3: Create entity links between units that share entities
substep_start = time.time()
@@ -253,13 +280,14 @@ async def extract_entities_batch_optimized(
for entity_ids in unit_to_entity_ids.values():
all_entity_ids.update(entity_ids)
_log(log_buffer, f" [6.3] Creating entity links for {len(all_entity_ids)} unique entities...", level='debug')
_log(log_buffer, f" [6.3] Creating entity links for {len(all_entity_ids)} unique entities...", level="debug")
# Find all units that reference these entities (ONE batched query)
entity_to_units = {}
if all_entity_ids:
query_start = time.time()
import uuid
entity_id_list = [uuid.UUID(eid) if isinstance(eid, str) else eid for eid in all_entity_ids]
rows = await conn.fetch(
"""
@@ -267,25 +295,29 @@ async def extract_entities_batch_optimized(
FROM unit_entities
WHERE entity_id = ANY($1::uuid[])
""",
entity_id_list
entity_id_list,
)
_log(
log_buffer,
f" [6.3.1] Query unit_entities: {len(rows)} rows in {time.time() - query_start:.3f}s",
level="debug",
)
_log(log_buffer, f" [6.3.1] Query unit_entities: {len(rows)} rows in {time.time() - query_start:.3f}s", level='debug')
# Group by entity_id
group_start = time.time()
for row in rows:
entity_id = row['entity_id']
entity_id = row["entity_id"]
if entity_id not in entity_to_units:
entity_to_units[entity_id] = []
entity_to_units[entity_id].append(row['unit_id'])
_log(log_buffer, f" [6.3.2] Group by entity_id: {time.time() - group_start:.3f}s", level='debug')
entity_to_units[entity_id].append(row["unit_id"])
_log(log_buffer, f" [6.3.2] Group by entity_id: {time.time() - group_start:.3f}s", level="debug")
# Create bidirectional links between units that share entities
# OPTIMIZATION: Limit links per entity to avoid N² explosion
# Only link each new unit to the most recent MAX_LINKS_PER_ENTITY units
MAX_LINKS_PER_ENTITY = 50 # Limit to prevent explosion when entity appears in many facts
link_gen_start = time.time()
links: List[EntityLink] = []
links: list[EntityLink] = []
new_unit_set = set(unit_ids) # Units from this batch
def to_uuid(val) -> UUID:
@@ -299,27 +331,52 @@ async def extract_entities_batch_optimized(
# Link new units to each other (within batch) - also limited
# For very common entities, limit within-batch links too
new_units_to_link = new_units[-MAX_LINKS_PER_ENTITY:] if len(new_units) > MAX_LINKS_PER_ENTITY else new_units
new_units_to_link = (
new_units[-MAX_LINKS_PER_ENTITY:] if len(new_units) > MAX_LINKS_PER_ENTITY else new_units
)
for i, unit_id_1 in enumerate(new_units_to_link):
for unit_id_2 in new_units_to_link[i+1:]:
links.append(EntityLink(from_unit_id=to_uuid(unit_id_1), to_unit_id=to_uuid(unit_id_2), entity_id=entity_uuid))
links.append(EntityLink(from_unit_id=to_uuid(unit_id_2), to_unit_id=to_uuid(unit_id_1), entity_id=entity_uuid))
for unit_id_2 in new_units_to_link[i + 1 :]:
links.append(
EntityLink(
from_unit_id=to_uuid(unit_id_1), to_unit_id=to_uuid(unit_id_2), entity_id=entity_uuid
)
)
links.append(
EntityLink(
from_unit_id=to_uuid(unit_id_2), to_unit_id=to_uuid(unit_id_1), entity_id=entity_uuid
)
)
# Link new units to LIMITED existing units (most recent)
existing_to_link = existing_units[-MAX_LINKS_PER_ENTITY:] # Take most recent
for new_unit in new_units:
for existing_unit in existing_to_link:
links.append(EntityLink(from_unit_id=to_uuid(new_unit), to_unit_id=to_uuid(existing_unit), entity_id=entity_uuid))
links.append(EntityLink(from_unit_id=to_uuid(existing_unit), to_unit_id=to_uuid(new_unit), entity_id=entity_uuid))
links.append(
EntityLink(
from_unit_id=to_uuid(new_unit), to_unit_id=to_uuid(existing_unit), entity_id=entity_uuid
)
)
links.append(
EntityLink(
from_unit_id=to_uuid(existing_unit), to_unit_id=to_uuid(new_unit), entity_id=entity_uuid
)
)
_log(log_buffer, f" [6.3.3] Generate {len(links)} links: {time.time() - link_gen_start:.3f}s", level='debug')
_log(log_buffer, f" [6.3] Entity link creation: {len(links)} links for {len(all_entity_ids)} unique entities in {time.time() - substep_start:.3f}s", level='debug')
_log(
log_buffer, f" [6.3.3] Generate {len(links)} links: {time.time() - link_gen_start:.3f}s", level="debug"
)
_log(
log_buffer,
f" [6.3] Entity link creation: {len(links)} links for {len(all_entity_ids)} unique entities in {time.time() - substep_start:.3f}s",
level="debug",
)
return links
except Exception as e:
logger.error(f"Failed to extract entities in batch: {str(e)}")
import traceback
traceback.print_exc()
raise
@@ -327,9 +384,9 @@ async def extract_entities_batch_optimized(
async def create_temporal_links_batch_per_fact(
conn,
bank_id: str,
unit_ids: List[str],
unit_ids: list[str],
time_window_hours: int = 24,
log_buffer: List[str] = None,
log_buffer: list[str] = None,
) -> int:
"""
Create temporal links for multiple units, each with their own event_date.
@@ -361,10 +418,13 @@ async def create_temporal_links_batch_per_fact(
FROM memory_units
WHERE id::text = ANY($1)
""",
unit_ids
unit_ids,
)
new_units = {str(row["id"]): row["event_date"] for row in rows}
_log(
log_buffer,
f" [7.1] Fetch event_dates for {len(unit_ids)} units: {time_mod.time() - fetch_dates_start:.3f}s",
)
new_units = {str(row['id']): row['event_date'] for row in rows}
_log(log_buffer, f" [7.1] Fetch event_dates for {len(unit_ids)} units: {time_mod.time() - fetch_dates_start:.3f}s")
# Fetch ALL potential temporal neighbors in ONE query (much faster!)
# Get time range across all units with overflow protection
@@ -383,9 +443,12 @@ async def create_temporal_links_batch_per_fact(
bank_id,
min_date,
max_date,
unit_ids
unit_ids,
)
_log(
log_buffer,
f" [7.2] Fetch {len(all_candidates)} candidate neighbors (1 query): {time_mod.time() - fetch_neighbors_start:.3f}s",
)
_log(log_buffer, f" [7.2] Fetch {len(all_candidates)} candidate neighbors (1 query): {time_mod.time() - fetch_neighbors_start:.3f}s")
# Filter and create links in memory (much faster than N queries)
link_gen_start = time_mod.time()
@@ -408,8 +471,8 @@ async def create_temporal_links_batch_per_fact(
if time_diff_hours <= time_window_hours:
weight = max(0.3, 1.0 - (time_diff_hours / time_window_hours))
# Create bidirectional links
links.append((unit_id, other_id, 'temporal', weight, None))
links.append((other_id, unit_id, 'temporal', weight, None))
links.append((unit_id, other_id, "temporal", weight, None))
links.append((other_id, unit_id, "temporal", weight, None))
_log(log_buffer, f" [7.3] Generate {len(links)} temporal links: {time_mod.time() - link_gen_start:.3f}s")
@@ -421,7 +484,7 @@ async def create_temporal_links_batch_per_fact(
VALUES ($1, $2, $3, $4, $5)
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
""",
links
links,
)
_log(log_buffer, f" [7.4] Insert {len(links)} temporal links: {time_mod.time() - insert_start:.3f}s")
@@ -430,6 +493,7 @@ async def create_temporal_links_batch_per_fact(
except Exception as e:
logger.error(f"Failed to create temporal links: {str(e)}")
import traceback
traceback.print_exc()
raise
@@ -437,11 +501,11 @@ async def create_temporal_links_batch_per_fact(
async def create_semantic_links_batch(
conn,
bank_id: str,
unit_ids: List[str],
embeddings: List[List[float]],
unit_ids: list[str],
embeddings: list[list[float]],
top_k: int = 5,
threshold: float = 0.7,
log_buffer: List[str] = None,
log_buffer: list[str] = None,
) -> int:
"""
Create semantic links for multiple units efficiently.
@@ -465,6 +529,7 @@ async def create_semantic_links_batch(
try:
import time as time_mod
import numpy as np
# Fetch ALL existing units with embeddings in ONE query
@@ -478,9 +543,12 @@ async def create_semantic_links_batch(
AND id::text != ALL($2)
""",
bank_id,
unit_ids
unit_ids,
)
_log(
log_buffer,
f" [8.1] Fetch {len(all_existing)} existing embeddings (1 query): {time_mod.time() - fetch_start:.3f}s",
)
_log(log_buffer, f" [8.1] Fetch {len(all_existing)} existing embeddings (1 query): {time_mod.time() - fetch_start:.3f}s")
# Convert to numpy for vectorized similarity computation
compute_start = time_mod.time()
@@ -488,15 +556,16 @@ async def create_semantic_links_batch(
if all_existing:
# Convert existing embeddings to numpy array
existing_ids = [str(row['id']) for row in all_existing]
existing_ids = [str(row["id"]) for row in all_existing]
# Stack embeddings as 2D array: (num_embeddings, embedding_dim)
embedding_arrays = []
for row in all_existing:
raw_emb = row['embedding']
raw_emb = row["embedding"]
# Handle different pgvector formats
if isinstance(raw_emb, str):
# Parse string format: "[1.0, 2.0, ...]"
import json
emb = np.array(json.loads(raw_emb), dtype=np.float32)
elif isinstance(raw_emb, (list, tuple)):
emb = np.array(raw_emb, dtype=np.float32)
@@ -537,7 +606,7 @@ async def create_semantic_links_batch(
similar_id = existing_ids[idx]
# Clamp to [0, 1] to handle floating point precision issues
similarity = float(min(1.0, max(0.0, similarities[idx])))
all_links.append((unit_id, similar_id, 'semantic', similarity, None))
all_links.append((unit_id, similar_id, "semantic", similarity, None))
# Also compute similarities WITHIN the new batch (new units to each other)
# Apply the same top_k limit per unit as we do for existing units
@@ -565,9 +634,12 @@ async def create_semantic_links_batch(
other_id = unit_ids[other_idx]
# Clamp to [0, 1] to handle floating point precision issues
similarity = float(min(1.0, max(0.0, similarities[local_idx])))
all_links.append((unit_id, other_id, 'semantic', similarity, None))
all_links.append((unit_id, other_id, "semantic", similarity, None))
_log(log_buffer, f" [8.2] Compute similarities & generate {len(all_links)} semantic links: {time_mod.time() - compute_start:.3f}s")
_log(
log_buffer,
f" [8.2] Compute similarities & generate {len(all_links)} semantic links: {time_mod.time() - compute_start:.3f}s",
)
if all_links:
insert_start = time_mod.time()
@@ -577,20 +649,23 @@ async def create_semantic_links_batch(
VALUES ($1, $2, $3, $4, $5)
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
""",
all_links
all_links,
)
_log(
log_buffer, f" [8.3] Insert {len(all_links)} semantic links: {time_mod.time() - insert_start:.3f}s"
)
_log(log_buffer, f" [8.3] Insert {len(all_links)} semantic links: {time_mod.time() - insert_start:.3f}s")
return len(all_links)
except Exception as e:
logger.error(f"Failed to create semantic links: {str(e)}")
import traceback
traceback.print_exc()
raise
async def insert_entity_links_batch(conn, links: List[EntityLink], chunk_size: int = 50000):
async def insert_entity_links_batch(conn, links: list[EntityLink], chunk_size: int = 50000):
"""
Insert all entity links using COPY to temp table + INSERT for maximum speed.
@@ -606,7 +681,6 @@ async def insert_entity_links_batch(conn, links: List[EntityLink], chunk_size: i
if not links:
return
import uuid as uuid_mod
import time as time_mod
total_start = time_mod.time()
@@ -633,21 +707,15 @@ async def insert_entity_links_batch(conn, links: List[EntityLink], chunk_size: i
convert_start = time_mod.time()
records = []
for link in links:
records.append((
link.from_unit_id,
link.to_unit_id,
link.link_type,
link.weight,
link.entity_id
))
records.append((link.from_unit_id, link.to_unit_id, link.link_type, link.weight, link.entity_id))
logger.debug(f" [9.3] Convert {len(records)} records: {time_mod.time() - convert_start:.3f}s")
# Bulk load using COPY (fastest method)
copy_start = time_mod.time()
await conn.copy_records_to_table(
'_temp_entity_links',
"_temp_entity_links",
records=records,
columns=['from_unit_id', 'to_unit_id', 'link_type', 'weight', 'entity_id']
columns=["from_unit_id", "to_unit_id", "link_type", "weight", "entity_id"],
)
logger.debug(f" [9.4] COPY {len(records)} records to temp table: {time_mod.time() - copy_start:.3f}s")
@@ -665,8 +733,8 @@ async def insert_entity_links_batch(conn, links: List[EntityLink], chunk_size: i
async def create_causal_links_batch(
conn,
unit_ids: List[str],
causal_relations_per_fact: List[List[dict]],
unit_ids: list[str],
causal_relations_per_fact: list[list[dict]],
) -> int:
"""
Create causal links between facts based on LLM-extracted causal relationships.
@@ -694,6 +762,7 @@ async def create_causal_links_batch(
try:
import time as time_mod
create_start = time_mod.time()
# Build links list
@@ -705,12 +774,12 @@ async def create_causal_links_batch(
from_unit_id = unit_ids[fact_idx]
for relation in causal_relations:
target_idx = relation['target_fact_index']
relation_type = relation['relation_type']
strength = relation.get('strength', 1.0)
target_idx = relation["target_fact_index"]
relation_type = relation["relation_type"]
strength = relation.get("strength", 1.0)
# Validate relation_type - must match database constraint
valid_types = {'causes', 'caused_by', 'enables', 'prevents'}
valid_types = {"causes", "caused_by", "enables", "prevents"}
if relation_type not in valid_types:
logger.error(
f"Invalid relation_type '{relation_type}' (type: {type(relation_type).__name__}) "
@@ -735,7 +804,6 @@ async def create_causal_links_batch(
# weight is the strength of the relationship
links.append((from_unit_id, to_unit_id, relation_type, strength, None))
if links:
insert_start = time_mod.time()
try:
@@ -745,14 +813,16 @@ async def create_causal_links_batch(
VALUES ($1, $2, $3, $4, $5)
ON CONFLICT (from_unit_id, to_unit_id, link_type, COALESCE(entity_id, '00000000-0000-0000-0000-000000000000'::uuid)) DO NOTHING
""",
links
links,
)
except Exception as db_error:
# Log the actual data being inserted for debugging
logger.error(f"Database insert failed for causal links. Error: {db_error}")
logger.error(f"Attempted to insert {len(links)} links. First few:")
for i, link in enumerate(links[:3]):
logger.error(f" Link {i}: from={link[0]}, to={link[1]}, type='{link[2]}' (repr={repr(link[2])}), weight={link[3]}, entity={link[4]}")
logger.error(
f" Link {i}: from={link[0]}, to={link[1]}, type='{link[2]}' (repr={repr(link[2])}), weight={link[3]}, entity={link[4]}"
)
raise
return len(links)
@@ -760,5 +830,6 @@ async def create_causal_links_batch(
except Exception as e:
logger.error(f"Failed to create causal links: {str(e)}")
import traceback
traceback.print_exc()
raise
@@ -3,15 +3,14 @@ Observation regeneration for retain pipeline.
Regenerates entity observations as part of the retain transaction.
"""
import logging
import time
import uuid
from datetime import datetime, timezone
from typing import List, Dict, Optional
from datetime import UTC, datetime
from ..search import observation_utils
from . import embedding_utils
from ..db_utils import acquire_with_retry
from .types import EntityLink
logger = logging.getLogger(__name__)
@@ -19,12 +18,12 @@ logger = logging.getLogger(__name__)
def utcnow():
"""Get current UTC time."""
return datetime.now(timezone.utc)
return datetime.now(UTC)
# Simple dataclass-like container for facts (avoid importing from memory_engine)
class MemoryFactForObservation:
def __init__(self, id: str, text: str, fact_type: str, context: str, occurred_start: Optional[str]):
def __init__(self, id: str, text: str, fact_type: str, context: str, occurred_start: str | None):
self.id = id
self.text = text
self.fact_type = fact_type
@@ -33,12 +32,7 @@ class MemoryFactForObservation:
async def regenerate_observations_batch(
conn,
embeddings_model,
llm_config,
bank_id: str,
entity_links: List[EntityLink],
log_buffer: List[str] = None
conn, embeddings_model, llm_config, bank_id: str, entity_links: list[EntityLink], log_buffer: list[str] = None
) -> None:
"""
Regenerate observations for top entities in this batch.
@@ -61,7 +55,7 @@ async def regenerate_observations_batch(
return
# Count mentions per entity in this batch
entity_mention_counts: Dict[str, int] = {}
entity_mention_counts: dict[str, int] = {}
for link in entity_links:
if link.entity_id:
entity_id = str(link.entity_id)
@@ -71,11 +65,7 @@ async def regenerate_observations_batch(
return
# Sort by mention count descending and take top N
sorted_entities = sorted(
entity_mention_counts.items(),
key=lambda x: x[1],
reverse=True
)
sorted_entities = sorted(entity_mention_counts.items(), key=lambda x: x[1], reverse=True)
entities_to_process = [e[0] for e in sorted_entities[:TOP_N_ENTITIES]]
obs_start = time.time()
@@ -89,9 +79,10 @@ async def regenerate_observations_batch(
SELECT id, canonical_name FROM entities
WHERE id = ANY($1) AND bank_id = $2
""",
entity_uuids, bank_id
entity_uuids,
bank_id,
)
entity_names = {row['id']: row['canonical_name'] for row in entity_rows}
entity_names = {row["id"]: row["canonical_name"] for row in entity_rows}
# Batch query for fact counts
fact_counts = await conn.fetch(
@@ -102,9 +93,10 @@ async def regenerate_observations_batch(
WHERE ue.entity_id = ANY($1) AND mu.bank_id = $2
GROUP BY ue.entity_id
""",
entity_uuids, bank_id
entity_uuids,
bank_id,
)
entity_fact_counts = {row['entity_id']: row['cnt'] for row in fact_counts}
entity_fact_counts = {row["entity_id"]: row["cnt"] for row in fact_counts}
# Filter entities that meet the threshold
entities_with_names = []
@@ -126,8 +118,7 @@ async def regenerate_observations_batch(
for entity_id, entity_name in entities_with_names:
try:
obs_ids = await _regenerate_entity_observations(
conn, embeddings_model, llm_config,
bank_id, entity_id, entity_name
conn, embeddings_model, llm_config, bank_id, entity_id, entity_name
)
total_observations += len(obs_ids)
except Exception as e:
@@ -135,17 +126,14 @@ async def regenerate_observations_batch(
obs_time = time.time() - obs_start
if log_buffer is not None:
log_buffer.append(f"[11] Observations: {total_observations} observations for {len(entities_with_names)} entities in {obs_time:.3f}s")
log_buffer.append(
f"[11] Observations: {total_observations} observations for {len(entities_with_names)} entities in {obs_time:.3f}s"
)
async def _regenerate_entity_observations(
conn,
embeddings_model,
llm_config,
bank_id: str,
entity_id: str,
entity_name: str
) -> List[str]:
conn, embeddings_model, llm_config, bank_id: str, entity_id: str, entity_name: str
) -> list[str]:
"""
Regenerate observations for a single entity.
@@ -176,7 +164,8 @@ async def _regenerate_entity_observations(
ORDER BY mu.occurred_start DESC
LIMIT 50
""",
bank_id, entity_uuid
bank_id,
entity_uuid,
)
if not rows:
@@ -185,21 +174,19 @@ async def _regenerate_entity_observations(
# Convert to fact objects for observation extraction
facts = []
for row in rows:
occurred_start = row['occurred_start'].isoformat() if row['occurred_start'] else None
facts.append(MemoryFactForObservation(
id=str(row['id']),
text=row['text'],
fact_type=row['fact_type'],
context=row['context'],
occurred_start=occurred_start
))
occurred_start = row["occurred_start"].isoformat() if row["occurred_start"] else None
facts.append(
MemoryFactForObservation(
id=str(row["id"]),
text=row["text"],
fact_type=row["fact_type"],
context=row["context"],
occurred_start=occurred_start,
)
)
# Extract observations using LLM
observations = await observation_utils.extract_observations_from_facts(
llm_config,
entity_name,
facts
)
observations = await observation_utils.extract_observations_from_facts(llm_config, entity_name, facts)
if not observations:
return []
@@ -217,13 +204,12 @@ async def _regenerate_entity_observations(
AND ue.entity_id = $2
)
""",
bank_id, entity_uuid
bank_id,
entity_uuid,
)
# Generate embeddings for new observations
embeddings = await embedding_utils.generate_embeddings_batch(
embeddings_model, observations
)
embeddings = await embedding_utils.generate_embeddings_batch(embeddings_model, observations)
# Insert new observations
current_time = utcnow()
@@ -247,9 +233,9 @@ async def _regenerate_entity_observations(
current_time,
current_time,
current_time,
current_time
current_time,
)
obs_id = str(result['id'])
obs_id = str(result["id"])
created_ids.append(obs_id)
# Link observation to entity
@@ -258,7 +244,8 @@ async def _regenerate_entity_observations(
INSERT INTO unit_entities (unit_id, entity_id)
VALUES ($1, $2)
""",
uuid.UUID(obs_id), entity_uuid
uuid.UUID(obs_id),
entity_uuid,
)
return created_ids
@@ -3,31 +3,33 @@ Main orchestrator for the retain pipeline.
Coordinates all retain pipeline modules to store memories efficiently.
"""
import logging
import time
import uuid
from datetime import datetime, timezone
from typing import List, Dict, Any, Optional
from datetime import UTC, datetime
from typing import Any
from . import bank_utils
from ..db_utils import acquire_with_retry
from . import bank_utils
def utcnow():
"""Get current UTC time."""
return datetime.now(timezone.utc)
return datetime.now(UTC)
from .types import RetainContent, ExtractedFact, ProcessedFact, EntityLink
from . import (
fact_extraction,
embedding_processing,
deduplication,
chunk_storage,
fact_storage,
deduplication,
embedding_processing,
entity_processing,
fact_extraction,
fact_storage,
link_creation,
observation_regeneration
observation_regeneration,
)
from .types import ExtractedFact, ProcessedFact, RetainContent
logger = logging.getLogger(__name__)
@@ -41,12 +43,12 @@ async def retain_batch(
format_date_fn,
duplicate_checker_fn,
bank_id: str,
contents_dicts: List[Dict[str, Any]],
document_id: Optional[str] = None,
contents_dicts: list[dict[str, Any]],
document_id: str | None = None,
is_first_batch: bool = True,
fact_type_override: Optional[str] = None,
confidence_score: Optional[float] = None,
) -> List[List[str]]:
fact_type_override: str | None = None,
confidence_score: float | None = None,
) -> list[list[str]]:
"""
Process a batch of content through the retain pipeline.
@@ -73,10 +75,10 @@ async def retain_batch(
# Buffer all logs
log_buffer = []
log_buffer.append(f"{'='*60}")
log_buffer.append(f"{'=' * 60}")
log_buffer.append(f"RETAIN_BATCH START: {bank_id}")
log_buffer.append(f"Batch size: {len(contents_dicts)} content items, {total_chars:,} chars")
log_buffer.append(f"{'='*60}")
log_buffer.append(f"{'=' * 60}")
# Get bank profile
profile = await bank_utils.get_bank_profile(pool, bank_id)
@@ -89,21 +91,20 @@ async def retain_batch(
content=item["content"],
context=item.get("context", ""),
event_date=item.get("event_date") or utcnow(),
metadata=item.get("metadata", {})
metadata=item.get("metadata", {}),
)
contents.append(content)
# Step 1: Extract facts from all contents
step_start = time.time()
extract_opinions = (fact_type_override == 'opinion')
extract_opinions = fact_type_override == "opinion"
extracted_facts, chunks = await fact_extraction.extract_facts_from_contents(
contents,
llm_config,
agent_name,
extract_opinions
contents, llm_config, agent_name, extract_opinions
)
log_buffer.append(
f"[1] Extract facts: {len(extracted_facts)} facts, {len(chunks)} chunks from {len(contents)} contents in {time.time() - step_start:.3f}s"
)
log_buffer.append(f"[1] Extract facts: {len(extracted_facts)} facts, {len(chunks)} chunks from {len(contents)} contents in {time.time() - step_start:.3f}s")
if not extracted_facts:
return [[] for _ in contents]
@@ -130,6 +131,7 @@ async def retain_batch(
# Group contents by document_id for document tracking and chunk storage
from collections import defaultdict
contents_by_doc = defaultdict(list)
for idx, content_dict in enumerate(contents_dicts):
doc_id = content_dict.get("document_id")
@@ -155,7 +157,11 @@ async def retain_batch(
if first_item.get("context"):
retain_params["context"] = first_item["context"]
if first_item.get("event_date"):
retain_params["event_date"] = first_item["event_date"].isoformat() if hasattr(first_item["event_date"], "isoformat") else str(first_item["event_date"])
retain_params["event_date"] = (
first_item["event_date"].isoformat()
if hasattr(first_item["event_date"], "isoformat")
else str(first_item["event_date"])
)
if first_item.get("metadata"):
retain_params["metadata"] = first_item["metadata"]
@@ -195,7 +201,11 @@ async def retain_batch(
if first_item.get("context"):
retain_params["context"] = first_item["context"]
if first_item.get("event_date"):
retain_params["event_date"] = first_item["event_date"].isoformat() if hasattr(first_item["event_date"], "isoformat") else str(first_item["event_date"])
retain_params["event_date"] = (
first_item["event_date"].isoformat()
if hasattr(first_item["event_date"], "isoformat")
else str(first_item["event_date"])
)
if first_item.get("metadata"):
retain_params["metadata"] = first_item["metadata"]
@@ -205,7 +215,9 @@ async def retain_batch(
document_ids_added.append(actual_doc_id)
if document_ids_added:
log_buffer.append(f"[2.5] Document tracking: {len(document_ids_added)} documents in {time.time() - step_start:.3f}s")
log_buffer.append(
f"[2.5] Document tracking: {len(document_ids_added)} documents in {time.time() - step_start:.3f}s"
)
# Store chunks and map to facts for all documents
step_start = time.time()
@@ -230,7 +242,9 @@ async def retain_batch(
for chunk_idx, chunk_id in chunk_id_map.items():
chunk_id_map_by_doc[(doc_id, chunk_idx)] = chunk_id
log_buffer.append(f"[3] Store chunks: {len(chunks)} chunks for {len(chunks_by_doc)} documents in {time.time() - step_start:.3f}s")
log_buffer.append(
f"[3] Store chunks: {len(chunks)} chunks for {len(chunks_by_doc)} documents in {time.time() - step_start:.3f}s"
)
# Map chunk_ids and document_ids to facts
for fact, processed_fact in zip(extracted_facts, processed_facts):
@@ -265,7 +279,9 @@ async def retain_batch(
is_duplicate_flags = await deduplication.check_duplicates_batch(
conn, bank_id, processed_facts, duplicate_checker_fn
)
log_buffer.append(f"[4] Deduplication: {sum(is_duplicate_flags)} duplicates in {time.time() - step_start:.3f}s")
log_buffer.append(
f"[4] Deduplication: {sum(is_duplicate_flags)} duplicates in {time.time() - step_start:.3f}s"
)
# Filter out duplicates
non_duplicate_facts = deduplication.filter_duplicates(processed_facts, is_duplicate_flags)
@@ -293,14 +309,18 @@ async def retain_batch(
# Create semantic links
step_start = time.time()
embeddings_for_links = [fact.embedding for fact in non_duplicate_facts]
semantic_link_count = await link_creation.create_semantic_links_batch(conn, bank_id, unit_ids, embeddings_for_links)
semantic_link_count = await link_creation.create_semantic_links_batch(
conn, bank_id, unit_ids, embeddings_for_links
)
log_buffer.append(f"[8] Semantic links: {semantic_link_count} links in {time.time() - step_start:.3f}s")
# Insert entity links
step_start = time.time()
if entity_links:
await entity_processing.insert_entity_links_batch(conn, entity_links)
log_buffer.append(f"[9] Entity links: {len(entity_links) if entity_links else 0} links in {time.time() - step_start:.3f}s")
log_buffer.append(
f"[9] Entity links: {len(entity_links) if entity_links else 0} links in {time.time() - step_start:.3f}s"
)
# Create causal links
step_start = time.time()
@@ -309,34 +329,22 @@ async def retain_batch(
# Regenerate observations INSIDE transaction for atomicity
await observation_regeneration.regenerate_observations_batch(
conn,
embeddings_model,
llm_config,
bank_id,
entity_links,
log_buffer
conn, embeddings_model, llm_config, bank_id, entity_links, log_buffer
)
# 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)
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)
# Log final summary
total_time = time.time() - start_time
log_buffer.append(f"{'='*60}")
log_buffer.append(f"{'=' * 60}")
log_buffer.append(f"RETAIN_BATCH COMPLETE: {len(unit_ids)} units in {total_time:.3f}s")
if document_ids_added:
log_buffer.append(f"Documents: {', '.join(document_ids_added)}")
log_buffer.append(f"{'='*60}")
log_buffer.append(f"{'=' * 60}")
logger.info("\n" + "\n".join(log_buffer) + "\n")
@@ -344,11 +352,11 @@ async def retain_batch(
def _map_results_to_contents(
contents: List[RetainContent],
extracted_facts: List[ExtractedFact],
is_duplicate_flags: List[bool],
unit_ids: List[str]
) -> List[List[str]]:
contents: list[RetainContent],
extracted_facts: list[ExtractedFact],
is_duplicate_flags: list[bool],
unit_ids: list[str],
) -> list[list[str]]:
"""
Map created unit IDs back to original content items.
@@ -376,17 +384,19 @@ def _map_results_to_contents(
async def _trigger_background_tasks(
task_backend,
bank_id: str,
unit_ids: List[str],
facts: List[ProcessedFact],
unit_ids: list[str],
facts: list[ProcessedFact],
) -> None:
"""Trigger opinion reinforcement as background task (after transaction commits)."""
# Trigger opinion reinforcement if there are entities
fact_entities = [[e.name for e in fact.entities] for fact in facts]
if any(fact_entities):
await task_backend.submit_task({
'type': 'reinforce_opinion',
'bank_id': bank_id,
'created_unit_ids': unit_ids,
'unit_texts': [fact.fact_text for fact in facts],
'unit_entities': fact_entities
})
await task_backend.submit_task(
{
"type": "reinforce_opinion",
"bank_id": bank_id,
"created_unit_ids": unit_ids,
"unit_texts": [fact.fact_text for fact in facts],
"unit_entities": fact_entities,
}
)
@@ -6,8 +6,7 @@ from content input to fact storage.
"""
from dataclasses import dataclass, field
from typing import List, Optional, Dict, Any
from datetime import datetime
from datetime import UTC, datetime
from uuid import UUID
@@ -18,16 +17,18 @@ class RetainContent:
Represents a single piece of content to extract facts from.
"""
content: str
context: str = ""
event_date: Optional[datetime] = None
metadata: Dict[str, str] = field(default_factory=dict)
event_date: datetime | None = None
metadata: dict[str, str] = field(default_factory=dict)
def __post_init__(self):
"""Ensure event_date is set."""
if self.event_date is None:
from datetime import datetime, timezone
self.event_date = datetime.now(timezone.utc)
from datetime import datetime
self.event_date = datetime.now(UTC)
@dataclass
@@ -37,6 +38,7 @@ class ChunkMetadata:
Used to track which facts were extracted from which chunks.
"""
chunk_text: str
fact_count: int
content_index: int # Index of the source content
@@ -50,9 +52,10 @@ class EntityRef:
Entities are extracted by the LLM during fact extraction.
"""
name: str
canonical_name: Optional[str] = None # Resolved canonical name
entity_id: Optional[UUID] = None # Resolved entity ID
canonical_name: str | None = None # Resolved canonical name
entity_id: UUID | None = None # Resolved entity ID
@dataclass
@@ -62,6 +65,7 @@ class CausalRelation:
Represents how one fact causes, enables, or prevents another.
"""
relation_type: str # "causes", "enables", "prevents", "caused_by"
target_fact_index: int # Index of the target fact in the batch
strength: float = 1.0 # Strength of the causal relationship
@@ -74,20 +78,21 @@ class ExtractedFact:
This is the raw output from fact extraction before processing.
"""
fact_text: str
fact_type: str # "world", "experience", "opinion", "observation"
entities: List[str] = field(default_factory=list)
occurred_start: Optional[datetime] = None
occurred_end: Optional[datetime] = None
where: Optional[str] = None # WHERE the fact occurred or is about
causal_relations: List[CausalRelation] = field(default_factory=list)
entities: list[str] = field(default_factory=list)
occurred_start: datetime | None = None
occurred_end: datetime | None = None
where: str | None = None # WHERE the fact occurred or is about
causal_relations: list[CausalRelation] = field(default_factory=list)
# Context from the content item
content_index: int = 0 # Which content this fact came from
chunk_index: int = 0 # Which chunk this fact came from
context: str = ""
mentioned_at: Optional[datetime] = None
metadata: Dict[str, str] = field(default_factory=dict)
mentioned_at: datetime | None = None
metadata: dict[str, str] = field(default_factory=dict)
@dataclass
@@ -97,37 +102,38 @@ class ProcessedFact:
Includes resolved entities, embeddings, and all necessary fields.
"""
# Core fact data
fact_text: str
fact_type: str
embedding: List[float]
embedding: list[float]
# Temporal data
occurred_start: Optional[datetime]
occurred_end: Optional[datetime]
occurred_start: datetime | None
occurred_end: datetime | None
mentioned_at: datetime
# Context and metadata
context: str
metadata: Dict[str, str]
metadata: dict[str, str]
# Location data
where: Optional[str] = None
where: str | None = None
# Entities
entities: List[EntityRef] = field(default_factory=list)
entities: list[EntityRef] = field(default_factory=list)
# Causal relations
causal_relations: List[CausalRelation] = field(default_factory=list)
causal_relations: list[CausalRelation] = field(default_factory=list)
# Chunk reference
chunk_id: Optional[str] = None
chunk_id: str | None = None
# Document reference (denormalized for query performance)
document_id: Optional[str] = None
document_id: str | None = None
# DB fields (set after insertion)
unit_id: Optional[UUID] = None
unit_id: UUID | None = None
@property
def is_duplicate(self) -> bool:
@@ -136,10 +142,8 @@ class ProcessedFact:
@staticmethod
def from_extracted_fact(
extracted_fact: 'ExtractedFact',
embedding: List[float],
chunk_id: Optional[str] = None
) -> 'ProcessedFact':
extracted_fact: "ExtractedFact", embedding: list[float], chunk_id: str | None = None
) -> "ProcessedFact":
"""
Create ProcessedFact from ExtractedFact.
@@ -151,12 +155,12 @@ class ProcessedFact:
Returns:
ProcessedFact ready for storage
"""
from datetime import datetime, timezone
from datetime import datetime
# Use occurred dates only if explicitly provided by LLM
occurred_start = extracted_fact.occurred_start
occurred_end = extracted_fact.occurred_end
mentioned_at = extracted_fact.mentioned_at or datetime.now(timezone.utc)
mentioned_at = extracted_fact.mentioned_at or datetime.now(UTC)
# Convert entity strings to EntityRef objects
entities = [EntityRef(name=name) for name in extracted_fact.entities]
@@ -172,7 +176,7 @@ class ProcessedFact:
metadata=extracted_fact.metadata,
entities=entities,
causal_relations=extracted_fact.causal_relations,
chunk_id=chunk_id
chunk_id=chunk_id,
)
@@ -183,10 +187,11 @@ class EntityLink:
Used for entity-based graph connections in the memory graph.
"""
from_unit_id: UUID
to_unit_id: UUID
entity_id: UUID
link_type: str = 'entity'
link_type: str = "entity"
weight: float = 1.0
@@ -197,24 +202,25 @@ class RetainBatch:
Tracks all facts, chunks, and metadata for a batch operation.
"""
bank_id: str
contents: List[RetainContent]
document_id: Optional[str] = None
fact_type_override: Optional[str] = None
confidence_score: Optional[float] = None
contents: list[RetainContent]
document_id: str | None = None
fact_type_override: str | None = None
confidence_score: float | None = None
# Extracted data (populated during processing)
extracted_facts: List[ExtractedFact] = field(default_factory=list)
processed_facts: List[ProcessedFact] = field(default_factory=list)
chunks: List[ChunkMetadata] = field(default_factory=list)
extracted_facts: list[ExtractedFact] = field(default_factory=list)
processed_facts: list[ProcessedFact] = field(default_factory=list)
chunks: list[ChunkMetadata] = field(default_factory=list)
# Results (populated after storage)
unit_ids_by_content: List[List[str]] = field(default_factory=list)
unit_ids_by_content: list[list[str]] = field(default_factory=list)
def get_facts_for_content(self, content_index: int) -> List[ExtractedFact]:
def get_facts_for_content(self, content_index: int) -> list[ExtractedFact]:
"""Get all extracted facts for a specific content item."""
return [f for f in self.extracted_facts if f.content_index == content_index]
def get_chunks_for_content(self, content_index: int) -> List[ChunkMetadata]:
def get_chunks_for_content(self, content_index: int) -> list[ChunkMetadata]:
"""Get all chunks for a specific content item."""
return [c for c in self.chunks if c.content_index == content_index]
@@ -7,15 +7,15 @@ Provides modular search architecture:
- Reranking: Pluggable strategies (heuristic, cross-encoder)
"""
from .retrieval import (
retrieve_parallel,
get_default_graph_retriever,
set_default_graph_retriever,
ParallelRetrievalResult,
)
from .graph_retrieval import GraphRetriever, BFSGraphRetriever
from .graph_retrieval import BFSGraphRetriever, GraphRetriever
from .mpfp_retrieval import MPFPGraphRetriever
from .reranking import CrossEncoderReranker
from .retrieval import (
ParallelRetrievalResult,
get_default_graph_retriever,
retrieve_parallel,
set_default_graph_retriever,
)
__all__ = [
"retrieve_parallel",
@@ -2,15 +2,12 @@
Helper functions for hybrid search (semantic + BM25 + graph).
"""
from typing import List, Dict, Any, Tuple
import asyncio
from .types import RetrievalResult, MergedCandidate
from typing import Any
from .types import MergedCandidate, RetrievalResult
def reciprocal_rank_fusion(
result_lists: List[List[RetrievalResult]],
k: int = 60
) -> List[MergedCandidate]:
def reciprocal_rank_fusion(result_lists: list[list[RetrievalResult]], k: int = 60) -> list[MergedCandidate]:
"""
Merge multiple ranked result lists using Reciprocal Rank Fusion.
@@ -73,20 +70,14 @@ def reciprocal_rank_fusion(
sorted(rrf_scores.items(), key=lambda x: x[1], reverse=True), start=1
):
merged_candidate = MergedCandidate(
retrieval=all_retrievals[doc_id],
rrf_score=rrf_score,
rrf_rank=rrf_rank,
source_ranks=source_ranks[doc_id]
retrieval=all_retrievals[doc_id], rrf_score=rrf_score, rrf_rank=rrf_rank, source_ranks=source_ranks[doc_id]
)
merged_results.append(merged_candidate)
return merged_results
def normalize_scores_on_deltas(
results: List[Dict[str, Any]],
score_keys: List[str]
) -> List[Dict[str, Any]]:
def normalize_scores_on_deltas(results: list[dict[str, Any]], score_keys: list[str]) -> list[dict[str, Any]]:
"""
Normalize scores based on deltas (min-max normalization within result set).
@@ -6,13 +6,11 @@ allowing different algorithms (BFS spreading activation, PPR, etc.) to be
swapped without changing the rest of the recall pipeline.
"""
from abc import ABC, abstractmethod
from typing import List, Optional
from datetime import datetime
import logging
from abc import ABC, abstractmethod
from .types import RetrievalResult
from ..db_utils import acquire_with_retry
from .types import RetrievalResult
logger = logging.getLogger(__name__)
@@ -40,10 +38,10 @@ class GraphRetriever(ABC):
bank_id: str,
fact_type: str,
budget: int,
query_text: Optional[str] = None,
semantic_seeds: Optional[List[RetrievalResult]] = None,
temporal_seeds: Optional[List[RetrievalResult]] = None,
) -> List[RetrievalResult]:
query_text: str | None = None,
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
) -> list[RetrievalResult]:
"""
Retrieve relevant facts via graph traversal.
@@ -109,10 +107,10 @@ class BFSGraphRetriever(GraphRetriever):
bank_id: str,
fact_type: str,
budget: int,
query_text: Optional[str] = None,
semantic_seeds: Optional[List[RetrievalResult]] = None,
temporal_seeds: Optional[List[RetrievalResult]] = None,
) -> List[RetrievalResult]:
query_text: str | None = None,
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
) -> list[RetrievalResult]:
"""
Retrieve facts using BFS spreading activation.
@@ -127,9 +125,7 @@ class BFSGraphRetriever(GraphRetriever):
for interface compatibility but not used.
"""
async with acquire_with_retry(pool) as conn:
return await self._retrieve_with_conn(
conn, query_embedding_str, bank_id, fact_type, budget
)
return await self._retrieve_with_conn(conn, query_embedding_str, bank_id, fact_type, budget)
async def _retrieve_with_conn(
self,
@@ -138,7 +134,7 @@ class BFSGraphRetriever(GraphRetriever):
bank_id: str,
fact_type: str,
budget: int,
) -> List[RetrievalResult]:
) -> list[RetrievalResult]:
"""Internal implementation with connection."""
# Step 1: Find entry points
@@ -155,8 +151,11 @@ class BFSGraphRetriever(GraphRetriever):
ORDER BY embedding <=> $1::vector
LIMIT $5
""",
query_embedding_str, bank_id, fact_type,
self.entry_point_threshold, self.entry_point_limit
query_embedding_str,
bank_id,
fact_type,
self.entry_point_threshold,
self.entry_point_limit,
)
if not entry_points:
@@ -165,10 +164,7 @@ class BFSGraphRetriever(GraphRetriever):
# Step 2: BFS spreading activation
visited = set()
results = []
queue = [
(RetrievalResult.from_db_row(dict(r)), r["similarity"])
for r in entry_points
]
queue = [(RetrievalResult.from_db_row(dict(r)), r["similarity"]) for r in entry_points]
budget_remaining = budget
while queue and budget_remaining > 0:
@@ -205,7 +201,10 @@ class BFSGraphRetriever(GraphRetriever):
ORDER BY ml.weight DESC
LIMIT $4
""",
batch_nodes, self.min_activation, fact_type, max_neighbors
batch_nodes,
self.min_activation,
fact_type,
max_neighbors,
)
for n in neighbors:
@@ -16,13 +16,12 @@ Key properties:
import asyncio
import logging
from dataclasses import dataclass, field
from typing import List, Dict, Optional, Tuple
from collections import defaultdict
from dataclasses import dataclass, field
from .types import RetrievalResult
from .graph_retrieval import GraphRetriever
from ..db_utils import acquire_with_retry
from .graph_retrieval import GraphRetriever
from .types import RetrievalResult
logger = logging.getLogger(__name__)
@@ -31,9 +30,11 @@ logger = logging.getLogger(__name__)
# Data Classes
# -----------------------------------------------------------------------------
@dataclass
class EdgeTarget:
"""A neighbor node with its edge weight."""
node_id: str
weight: float
@@ -41,19 +42,15 @@ class EdgeTarget:
@dataclass
class TypedAdjacency:
"""Adjacency lists split by edge type."""
# edge_type -> from_node_id -> list of (to_node_id, weight)
graphs: Dict[str, Dict[str, List[EdgeTarget]]] = field(default_factory=dict)
def get_neighbors(self, edge_type: str, node_id: str) -> List[EdgeTarget]:
# edge_type -> from_node_id -> list of (to_node_id, weight)
graphs: dict[str, dict[str, list[EdgeTarget]]] = field(default_factory=dict)
def get_neighbors(self, edge_type: str, node_id: str) -> list[EdgeTarget]:
"""Get neighbors for a node via a specific edge type."""
return self.graphs.get(edge_type, {}).get(node_id, [])
def get_normalized_neighbors(
self,
edge_type: str,
node_id: str,
top_k: int
) -> List[EdgeTarget]:
def get_normalized_neighbors(self, edge_type: str, node_id: str, top_k: int) -> list[EdgeTarget]:
"""Get top-k neighbors with weights normalized to sum to 1."""
neighbors = self.get_neighbors(edge_type, node_id)[:top_k]
if not neighbors:
@@ -63,45 +60,49 @@ class TypedAdjacency:
if total == 0:
return []
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]
@dataclass
class PatternResult:
"""Result from a single pattern traversal."""
pattern: List[str]
scores: Dict[str, float] # node_id -> accumulated mass
pattern: list[str]
scores: dict[str, float] # node_id -> accumulated mass
@dataclass
class MPFPConfig:
"""Configuration for MPFP algorithm."""
alpha: float = 0.15 # teleport/keep probability
threshold: float = 1e-6 # mass pruning threshold (lower = explore more)
top_k_neighbors: int = 20 # fan-out limit per node
alpha: float = 0.15 # teleport/keep probability
threshold: float = 1e-6 # mass pruning threshold (lower = explore more)
top_k_neighbors: int = 20 # fan-out limit per node
# Patterns from semantic seeds
patterns_semantic: List[List[str]] = field(default_factory=lambda: [
['semantic', 'semantic'], # topic expansion
['entity', 'temporal'], # entity timeline
['semantic', 'causes'], # reasoning chains (forward)
['semantic', 'caused_by'], # reasoning chains (backward)
['entity', 'semantic'], # entity context
])
patterns_semantic: list[list[str]] = field(
default_factory=lambda: [
["semantic", "semantic"], # topic expansion
["entity", "temporal"], # entity timeline
["semantic", "causes"], # reasoning chains (forward)
["semantic", "caused_by"], # reasoning chains (backward)
["entity", "semantic"], # entity context
]
)
# Patterns from temporal seeds
patterns_temporal: List[List[str]] = field(default_factory=lambda: [
['temporal', 'semantic'], # what was happening then
['temporal', 'entity'], # who was involved then
])
patterns_temporal: list[list[str]] = field(
default_factory=lambda: [
["temporal", "semantic"], # what was happening then
["temporal", "entity"], # who was involved then
]
)
@dataclass
class SeedNode:
"""An entry point node with its initial score."""
node_id: str
score: float # initial mass (e.g., similarity score)
@@ -110,9 +111,10 @@ class SeedNode:
# Core Algorithm
# -----------------------------------------------------------------------------
def mpfp_traverse(
seeds: List[SeedNode],
pattern: List[str],
seeds: list[SeedNode],
pattern: list[str],
adjacency: TypedAdjacency,
config: MPFPConfig,
) -> PatternResult:
@@ -131,20 +133,18 @@ def mpfp_traverse(
if not seeds:
return PatternResult(pattern=pattern, scores={})
scores: Dict[str, float] = {}
scores: dict[str, float] = {}
# Initialize frontier with seed masses (normalized)
total_seed_score = sum(s.score for s in seeds)
if total_seed_score == 0:
total_seed_score = len(seeds) # fallback to uniform
frontier: Dict[str, float] = {
s.node_id: s.score / total_seed_score for s in seeds
}
frontier: dict[str, float] = {s.node_id: s.score / total_seed_score for s in seeds}
# Follow pattern hop by hop
for edge_type in pattern:
next_frontier: Dict[str, float] = {}
next_frontier: dict[str, float] = {}
for node_id, mass in frontier.items():
if mass < config.threshold:
@@ -155,15 +155,10 @@ def mpfp_traverse(
# Push (1-α) to neighbors
push_mass = (1 - config.alpha) * mass
neighbors = adjacency.get_normalized_neighbors(
edge_type, node_id, config.top_k_neighbors
)
neighbors = adjacency.get_normalized_neighbors(edge_type, node_id, config.top_k_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
frontier = next_frontier
@@ -176,10 +171,10 @@ def mpfp_traverse(
def rrf_fusion(
results: List[PatternResult],
results: list[PatternResult],
k: int = 60,
top_k: int = 50,
) -> List[Tuple[str, float]]:
) -> list[tuple[str, float]]:
"""
Reciprocal Rank Fusion to combine pattern results.
@@ -191,28 +186,20 @@ def rrf_fusion(
Returns:
List of (node_id, fused_score) tuples, sorted by score descending
"""
fused: Dict[str, float] = {}
fused: dict[str, float] = {}
for result in results:
if not result.scores:
continue
# Rank nodes by their score in this pattern
ranked = sorted(
result.scores.keys(),
key=lambda n: result.scores[n],
reverse=True
)
ranked = sorted(result.scores.keys(), key=lambda n: result.scores[n], reverse=True)
for rank, node_id in enumerate(ranked):
fused[node_id] = fused.get(node_id, 0) + 1.0 / (k + rank + 1)
# Sort by fused score and return top-k
sorted_results = sorted(
fused.items(),
key=lambda x: x[1],
reverse=True
)
sorted_results = sorted(fused.items(), key=lambda x: x[1], reverse=True)
return sorted_results[:top_k]
@@ -221,6 +208,7 @@ def rrf_fusion(
# Database Loading
# -----------------------------------------------------------------------------
async def load_typed_adjacency(pool, bank_id: str) -> TypedAdjacency:
"""
Load all edges for a bank, split by edge type.
@@ -237,31 +225,27 @@ async def load_typed_adjacency(pool, bank_id: str) -> TypedAdjacency:
AND ml.weight >= 0.1
ORDER BY ml.from_unit_id, ml.weight DESC
""",
bank_id
bank_id,
)
graphs: Dict[str, Dict[str, List[EdgeTarget]]] = defaultdict(
lambda: defaultdict(list)
)
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']
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)
)
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(
pool,
node_ids: List[str],
node_ids: list[str],
fact_type: str,
) -> List[RetrievalResult]:
) -> list[RetrievalResult]:
"""Fetch full memory unit details for a list of node IDs."""
if not node_ids:
return []
@@ -276,7 +260,7 @@ async def fetch_memory_units_by_ids(
AND fact_type = $2
""",
node_ids,
fact_type
fact_type,
)
return [RetrievalResult.from_db_row(dict(r)) for r in rows]
@@ -286,6 +270,7 @@ async def fetch_memory_units_by_ids(
# Graph Retriever Implementation
# -----------------------------------------------------------------------------
class MPFPGraphRetriever(GraphRetriever):
"""
Graph retrieval using Meta-Path Forward Push.
@@ -294,7 +279,7 @@ class MPFPGraphRetriever(GraphRetriever):
then fuses results via RRF.
"""
def __init__(self, config: Optional[MPFPConfig] = None):
def __init__(self, config: MPFPConfig | None = None):
"""
Initialize MPFP retriever.
@@ -302,7 +287,7 @@ class MPFPGraphRetriever(GraphRetriever):
config: Algorithm configuration (uses defaults if None)
"""
self.config = config or MPFPConfig()
self._adjacency_cache: Dict[str, TypedAdjacency] = {}
self._adjacency_cache: dict[str, TypedAdjacency] = {}
@property
def name(self) -> str:
@@ -315,10 +300,10 @@ class MPFPGraphRetriever(GraphRetriever):
bank_id: str,
fact_type: str,
budget: int,
query_text: Optional[str] = None,
semantic_seeds: Optional[List[RetrievalResult]] = None,
temporal_seeds: Optional[List[RetrievalResult]] = None,
) -> List[RetrievalResult]:
query_text: str | None = None,
semantic_seeds: list[RetrievalResult] | None = None,
temporal_seeds: list[RetrievalResult] | None = None,
) -> list[RetrievalResult]:
"""
Retrieve facts using MPFP algorithm.
@@ -339,14 +324,12 @@ class MPFPGraphRetriever(GraphRetriever):
adjacency = await load_typed_adjacency(pool, bank_id)
# Convert seeds to SeedNode format
semantic_seed_nodes = self._convert_seeds(semantic_seeds, 'similarity')
temporal_seed_nodes = self._convert_seeds(temporal_seeds, 'temporal_score')
semantic_seed_nodes = self._convert_seeds(semantic_seeds, "similarity")
temporal_seed_nodes = self._convert_seeds(temporal_seeds, "temporal_score")
# If no semantic seeds provided, fall back to finding our own
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
tasks = []
@@ -407,9 +390,9 @@ class MPFPGraphRetriever(GraphRetriever):
def _convert_seeds(
self,
seeds: Optional[List[RetrievalResult]],
seeds: list[RetrievalResult] | None,
score_attr: str,
) -> List[SeedNode]:
) -> list[SeedNode]:
"""Convert RetrievalResult seeds to SeedNode format."""
if not seeds:
return []
@@ -431,7 +414,7 @@ class MPFPGraphRetriever(GraphRetriever):
fact_type: str,
limit: int = 20,
threshold: float = 0.3,
) -> List[SeedNode]:
) -> list[SeedNode]:
"""Fallback: find semantic seeds via embedding search."""
async with acquire_with_retry(pool) as conn:
rows = await conn.fetch(
@@ -445,10 +428,11 @@ class MPFPGraphRetriever(GraphRetriever):
ORDER BY embedding <=> $1::vector
LIMIT $5
""",
query_embedding_str, bank_id, fact_type, threshold, limit
query_embedding_str,
bank_id,
fact_type,
threshold,
limit,
)
return [
SeedNode(node_id=str(r['id']), score=r['similarity'])
for r in rows
]
return [SeedNode(node_id=str(r["id"]), score=r["similarity"]) for r in rows]
@@ -6,7 +6,7 @@ about an entity, without personality influence.
"""
import logging
from typing import List, Dict, Any
from pydantic import BaseModel, Field
from ..response_models import MemoryFact
@@ -16,18 +16,17 @@ logger = logging.getLogger(__name__)
class Observation(BaseModel):
"""An observation about an entity."""
observation: str = Field(description="The observation text - a factual statement about the entity")
class ObservationExtractionResponse(BaseModel):
"""Response containing extracted observations."""
observations: List[Observation] = Field(
default_factory=list,
description="List of observations about the entity"
)
observations: list[Observation] = Field(default_factory=list, description="List of observations about the entity")
def format_facts_for_observation_prompt(facts: List[MemoryFact]) -> str:
def format_facts_for_observation_prompt(facts: list[MemoryFact]) -> str:
"""Format facts as text for observation extraction prompt."""
import json
@@ -35,9 +34,7 @@ def format_facts_for_observation_prompt(facts: List[MemoryFact]) -> str:
return "[]"
formatted = []
for fact in facts:
fact_obj = {
"text": fact.text
}
fact_obj = {"text": fact.text}
# Add context if available
if fact.context:
@@ -92,11 +89,7 @@ def get_observation_system_message() -> str:
return "You are an objective observer synthesizing facts about an entity. Generate clear, factual observations without opinions or personality influence. Be concise and accurate."
async def extract_observations_from_facts(
llm_config,
entity_name: str,
facts: List[MemoryFact]
) -> List[str]:
async def extract_observations_from_facts(llm_config, entity_name: str, facts: list[MemoryFact]) -> list[str]:
"""
Extract observations from facts about an entity using LLM.
@@ -118,10 +111,10 @@ async def extract_observations_from_facts(
result = await llm_config.call(
messages=[
{"role": "system", "content": get_observation_system_message()},
{"role": "user", "content": prompt}
{"role": "user", "content": prompt},
],
response_format=ObservationExtractionResponse,
scope="memory_extract_observation"
scope="memory_extract_observation",
)
observations = [op.observation for op in result.observations]
@@ -2,7 +2,6 @@
Cross-encoder neural reranking for search results.
"""
from typing import List
from .types import MergedCandidate, ScoredResult
@@ -24,14 +23,11 @@ class CrossEncoderReranker:
"""
if cross_encoder is None:
from hindsight_api.engine.cross_encoder import create_cross_encoder_from_env
cross_encoder = create_cross_encoder_from_env()
self.cross_encoder = cross_encoder
def rerank(
self,
query: str,
candidates: List[MergedCandidate]
) -> List[ScoredResult]:
def rerank(self, query: str, candidates: list[MergedCandidate]) -> list[ScoredResult]:
"""
Rerank candidates using cross-encoder scores.
@@ -77,6 +73,7 @@ class CrossEncoderReranker:
# Normalize scores using sigmoid to [0, 1] range
# Cross-encoder returns logits which can be negative
import numpy as np
def sigmoid(x):
return 1 / (1 + np.exp(-x))
@@ -89,7 +86,7 @@ class CrossEncoderReranker:
candidate=candidate,
cross_encoder_score=float(raw_score),
cross_encoder_score_normalized=float(norm_score),
weight=float(norm_score) # Initial weight is just cross-encoder score
weight=float(norm_score), # Initial weight is just cross-encoder score
)
scored_results.append(scored_result)
@@ -8,16 +8,17 @@ Implements:
4. Temporal retrieval (time-aware search with spreading)
"""
from typing import List, Dict, Optional
from dataclasses import dataclass, field
from datetime import datetime
import asyncio
import logging
from ..db_utils import acquire_with_retry
from .types import RetrievalResult
from .graph_retrieval import GraphRetriever, BFSGraphRetriever
from .mpfp_retrieval import MPFPGraphRetriever
from dataclasses import dataclass, field
from datetime import UTC, datetime
from typing import Optional
from ...config import get_config
from ..db_utils import acquire_with_retry
from .graph_retrieval import BFSGraphRetriever, GraphRetriever
from .mpfp_retrieval import MPFPGraphRetriever
from .types import RetrievalResult
logger = logging.getLogger(__name__)
@@ -25,16 +26,17 @@ logger = logging.getLogger(__name__)
@dataclass
class ParallelRetrievalResult:
"""Result from parallel retrieval across all methods."""
semantic: List[RetrievalResult]
bm25: List[RetrievalResult]
graph: List[RetrievalResult]
temporal: Optional[List[RetrievalResult]]
timings: Dict[str, float] = field(default_factory=dict)
temporal_constraint: Optional[tuple] = None # (start_date, end_date)
semantic: list[RetrievalResult]
bm25: list[RetrievalResult]
graph: list[RetrievalResult]
temporal: list[RetrievalResult] | None
timings: dict[str, float] = field(default_factory=dict)
temporal_constraint: tuple | None = None # (start_date, end_date)
# Default graph retriever instance (can be overridden)
_default_graph_retriever: Optional[GraphRetriever] = None
_default_graph_retriever: GraphRetriever | None = None
def get_default_graph_retriever() -> GraphRetriever:
@@ -62,12 +64,8 @@ def set_default_graph_retriever(retriever: GraphRetriever) -> None:
async def retrieve_semantic(
conn,
query_emb_str: str,
bank_id: str,
fact_type: str,
limit: int
) -> List[RetrievalResult]:
conn, query_emb_str: str, bank_id: str, fact_type: str, limit: int
) -> list[RetrievalResult]:
"""
Semantic retrieval via vector similarity.
@@ -93,18 +91,15 @@ async def retrieve_semantic(
ORDER BY embedding <=> $1::vector
LIMIT $4
""",
query_emb_str, bank_id, fact_type, limit
query_emb_str,
bank_id,
fact_type,
limit,
)
return [RetrievalResult.from_db_row(dict(r)) for r in results]
async def retrieve_bm25(
conn,
query_text: str,
bank_id: str,
fact_type: str,
limit: int
) -> List[RetrievalResult]:
async def retrieve_bm25(conn, query_text: str, bank_id: str, fact_type: str, limit: int) -> list[RetrievalResult]:
"""
BM25 keyword retrieval via full-text search.
@@ -122,7 +117,7 @@ async def retrieve_bm25(
# Sanitize query text: remove special characters that have meaning in tsquery
# Keep only alphanumeric characters and spaces
sanitized_text = re.sub(r'[^\w\s]', ' ', query_text.lower())
sanitized_text = re.sub(r"[^\w\s]", " ", query_text.lower())
# Split and filter empty strings
tokens = [token for token in sanitized_text.split() if token]
@@ -146,7 +141,10 @@ async def retrieve_bm25(
ORDER BY bm25_score DESC
LIMIT $4
""",
query_tsquery, bank_id, fact_type, limit
query_tsquery,
bank_id,
fact_type,
limit,
)
return [RetrievalResult.from_db_row(dict(r)) for r in results]
@@ -159,8 +157,8 @@ async def retrieve_temporal(
start_date: datetime,
end_date: datetime,
budget: int,
semantic_threshold: float = 0.1
) -> List[RetrievalResult]:
semantic_threshold: float = 0.1,
) -> list[RetrievalResult]:
"""
Temporal retrieval with spreading activation.
@@ -182,13 +180,12 @@ async def retrieve_temporal(
Returns:
List of RetrievalResult objects with temporal scores
"""
from datetime import timezone
# Ensure start_date and end_date are timezone-aware (UTC) to match database datetimes
if start_date.tzinfo is None:
start_date = start_date.replace(tzinfo=timezone.utc)
start_date = start_date.replace(tzinfo=UTC)
if end_date.tzinfo is None:
end_date = end_date.replace(tzinfo=timezone.utc)
end_date = end_date.replace(tzinfo=UTC)
entry_points = await conn.fetch(
"""
@@ -215,7 +212,12 @@ async def retrieve_temporal(
ORDER BY COALESCE(occurred_start, mentioned_at, occurred_end) DESC, (embedding <=> $1::vector) ASC
LIMIT 10
""",
query_emb_str, bank_id, fact_type, start_date, end_date, semantic_threshold
query_emb_str,
bank_id,
fact_type,
start_date,
end_date,
semantic_threshold,
)
if not entry_points:
@@ -258,7 +260,9 @@ async def retrieve_temporal(
results.append(ep_result)
# Spread through temporal links
queue = [(RetrievalResult.from_db_row(dict(ep)), ep["similarity"], 1.0) for ep in entry_points] # (unit, semantic_sim, temporal_score)
queue = [
(RetrievalResult.from_db_row(dict(ep)), ep["similarity"], 1.0) for ep in entry_points
] # (unit, semantic_sim, temporal_score)
budget_remaining = budget - len(entry_points)
while queue and budget_remaining > 0:
@@ -283,7 +287,10 @@ async def retrieve_temporal(
ORDER BY ml.weight DESC
LIMIT 10
""",
query_emb_str, current.id, fact_type, semantic_threshold
query_emb_str,
current.id,
fact_type,
semantic_threshold,
)
for n in neighbors:
@@ -307,7 +314,9 @@ async def retrieve_temporal(
if neighbor_best_date:
days_from_mid = abs((neighbor_best_date - mid_date).total_seconds() / 86400)
neighbor_temporal_proximity = 1.0 - min(days_from_mid / (total_days / 2), 1.0) if total_days > 0 else 1.0
neighbor_temporal_proximity = (
1.0 - min(days_from_mid / (total_days / 2), 1.0) if total_days > 0 else 1.0
)
else:
neighbor_temporal_proximity = 0.3 # Lower score if no temporal data
@@ -349,9 +358,9 @@ async def retrieve_parallel(
bank_id: str,
fact_type: str,
thinking_budget: int,
question_date: Optional[datetime] = None,
question_date: datetime | None = None,
query_analyzer: Optional["QueryAnalyzer"] = None,
graph_retriever: Optional[GraphRetriever] = None,
graph_retriever: GraphRetriever | None = None,
) -> ParallelRetrievalResult:
"""
Run 3-way or 4-way parallel retrieval (adds temporal if detected).
@@ -372,29 +381,26 @@ async def retrieve_parallel(
"""
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()
if retriever.name == "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:
return await _retrieve_parallel_bfs(
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
)
@dataclass
class _SemanticGraphResult:
"""Internal result from semantic→graph chain."""
semantic: List[RetrievalResult]
graph: List[RetrievalResult]
semantic: list[RetrievalResult]
graph: list[RetrievalResult]
semantic_time: float
graph_time: float
@@ -402,7 +408,8 @@ class _SemanticGraphResult:
@dataclass
class _TimedResult:
"""Internal result with timing."""
results: List[RetrievalResult]
results: list[RetrievalResult]
time: float
@@ -413,7 +420,7 @@ async def _retrieve_parallel_mpfp(
bank_id: str,
fact_type: str,
thinking_budget: int,
temporal_constraint: Optional[tuple],
temporal_constraint: tuple | None,
retriever: GraphRetriever,
) -> ParallelRetrievalResult:
"""
@@ -430,9 +437,7 @@ async def _retrieve_parallel_mpfp(
"""Chain: semantic retrieval → graph retrieval (using semantic as seeds)."""
start = time.time()
async with acquire_with_retry(pool) as conn:
semantic = await retrieve_semantic(
conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget
)
semantic = await retrieve_semantic(conn, query_embedding_str, bank_id, fact_type, limit=thinking_budget)
semantic_time = time.time() - start
# Get temporal seeds if needed (quick query, part of this chain)
@@ -441,8 +446,7 @@ async def _retrieve_parallel_mpfp(
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
conn, query_embedding_str, bank_id, fact_type, tc_start, tc_end, limit=20
)
# Run graph with seeds
@@ -473,8 +477,14 @@ async def _retrieve_parallel_mpfp(
start = time.time()
async with acquire_with_retry(pool) as conn:
results = await retrieve_temporal(
conn, query_embedding_str, bank_id, fact_type,
tc_start, tc_end, budget=thinking_budget, semantic_threshold=0.1
conn,
query_embedding_str,
bank_id,
fact_type,
tc_start,
tc_end,
budget=thinking_budget,
semantic_threshold=0.1,
)
return _TimedResult(results, time.time() - start)
@@ -527,14 +537,13 @@ async def _get_temporal_entry_points(
end_date: datetime,
limit: int = 20,
semantic_threshold: float = 0.1,
) -> List[RetrievalResult]:
) -> list[RetrievalResult]:
"""Get temporal entry points (facts in date range with semantic relevance)."""
from datetime import timezone
if start_date.tzinfo is None:
start_date = start_date.replace(tzinfo=timezone.utc)
start_date = start_date.replace(tzinfo=UTC)
if end_date.tzinfo is None:
end_date = end_date.replace(tzinfo=timezone.utc)
end_date = end_date.replace(tzinfo=UTC)
rows = await conn.fetch(
"""
@@ -557,7 +566,13 @@ async def _get_temporal_entry_points(
(embedding <=> $1::vector) ASC
LIMIT $7
""",
query_embedding_str, bank_id, fact_type, start_date, end_date, semantic_threshold, limit
query_embedding_str,
bank_id,
fact_type,
start_date,
end_date,
semantic_threshold,
limit,
)
results = []
@@ -597,7 +612,7 @@ async def _retrieve_parallel_bfs(
bank_id: str,
fact_type: str,
thinking_budget: int,
temporal_constraint: Optional[tuple],
temporal_constraint: tuple | None,
retriever: GraphRetriever,
) -> ParallelRetrievalResult:
"""BFS retrieval: all methods run in parallel (original behavior)."""
@@ -631,8 +646,14 @@ async def _retrieve_parallel_bfs(
start = time.time()
async with acquire_with_retry(pool) as conn:
results = await retrieve_temporal(
conn, query_embedding_str, bank_id, fact_type,
tc_start, tc_end, budget=thinking_budget, semantic_threshold=0.1
conn,
query_embedding_str,
bank_id,
fact_type,
tc_start,
tc_end,
budget=thinking_budget,
semantic_threshold=0.1,
)
return _TimedResult(results, time.time() - start)
@@ -4,11 +4,11 @@ Scoring functions for memory search and retrieval.
Includes recency weighting, frequency weighting, temporal proximity,
and similarity calculations used in memory activation and ranking.
"""
from datetime import datetime
from typing import List
def cosine_similarity(vec1: List[float], vec2: List[float]) -> float:
def cosine_similarity(vec1: list[float], vec2: list[float]) -> float:
"""
Calculate cosine similarity between two vectors.
@@ -58,6 +58,7 @@ def calculate_recency_weight(days_since: float, half_life_days: float = 365.0) -
Weight between 0 and 1
"""
import math
# Logarithmic decay: 1 / (1 + log(1 + days_since/half_life))
# This decays much slower than exponential, giving better long-term differentiation
normalized_age = days_since / half_life_days
@@ -79,6 +80,7 @@ def calculate_frequency_weight(access_count: int, max_boost: float = 2.0) -> flo
Weight between 1.0 and max_boost
"""
import math
if access_count <= 0:
return 1.0
@@ -116,11 +118,7 @@ def calculate_temporal_anchor(occurred_start: datetime, occurred_end: datetime)
return midpoint
def calculate_temporal_proximity(
anchor_a: datetime,
anchor_b: datetime,
half_life_days: float = 30.0
) -> float:
def calculate_temporal_proximity(anchor_a: datetime, anchor_b: datetime, half_life_days: float = 30.0) -> float:
"""
Calculate temporal proximity between two temporal anchors.
@@ -4,16 +4,16 @@ Temporal extraction for time-aware search queries.
Handles natural language temporal expressions using transformer-based query analysis.
"""
from typing import Optional, Tuple
from datetime import datetime
import logging
from hindsight_api.engine.query_analyzer import QueryAnalyzer, DateparserQueryAnalyzer
from datetime import datetime
from hindsight_api.engine.query_analyzer import DateparserQueryAnalyzer, QueryAnalyzer
logger = logging.getLogger(__name__)
# Global default analyzer instance
# Can be overridden by passing a custom analyzer to extract_temporal_constraint
_default_analyzer: Optional[QueryAnalyzer] = None
_default_analyzer: QueryAnalyzer | None = None
def get_default_analyzer() -> QueryAnalyzer:
@@ -33,9 +33,9 @@ def get_default_analyzer() -> QueryAnalyzer:
def extract_temporal_constraint(
query: str,
reference_date: Optional[datetime] = None,
analyzer: Optional[QueryAnalyzer] = None,
) -> Optional[Tuple[datetime, datetime]]:
reference_date: datetime | None = None,
analyzer: QueryAnalyzer | None = None,
) -> tuple[datetime, datetime] | None:
"""
Extract temporal constraint from query.
@@ -55,10 +55,7 @@ def extract_temporal_constraint(
analysis = analyzer.analyze(query, reference_date)
if analysis.temporal_constraint:
result = (
analysis.temporal_constraint.start_date,
analysis.temporal_constraint.end_date
)
result = (analysis.temporal_constraint.start_date, analysis.temporal_constraint.end_date)
return result
return None
@@ -2,41 +2,35 @@
Think operation utilities for formulating answers based on agent and world facts.
"""
import asyncio
import logging
import re
from datetime import datetime, timezone
from typing import Dict, List, Any
from datetime import datetime
from pydantic import BaseModel, Field
from ..response_models import ReflectResult, MemoryFact, DispositionTraits
from ..response_models import DispositionTraits, MemoryFact
logger = logging.getLogger(__name__)
class Opinion(BaseModel):
"""An opinion formed by the bank."""
opinion: str = Field(description="The opinion or perspective with reasoning included")
confidence: float = Field(description="Confidence score for this opinion (0.0 to 1.0, where 1.0 is very confident)")
class OpinionExtractionResponse(BaseModel):
"""Response containing extracted opinions."""
opinions: List[Opinion] = Field(
default_factory=list,
description="List of opinions formed with their supporting reasons and confidence scores"
opinions: list[Opinion] = Field(
default_factory=list, description="List of opinions formed with their supporting reasons and confidence scores"
)
def describe_trait_level(value: int) -> str:
"""Convert trait value (1-5) to descriptive text."""
levels = {
1: "very low",
2: "low",
3: "moderate",
4: "high",
5: "very high"
}
levels = {1: "very low", 2: "low", 3: "moderate", 4: "high", 5: "very high"}
return levels.get(value, "moderate")
@@ -47,7 +41,7 @@ def build_disposition_description(disposition: DispositionTraits) -> str:
2: "You tend to trust information but may question obvious inconsistencies.",
3: "You have a balanced approach to information, neither too trusting nor too skeptical.",
4: "You are somewhat skeptical and often question the reliability of information.",
5: "You are highly skeptical and critically examine all information for accuracy and hidden motives."
5: "You are highly skeptical and critically examine all information for accuracy and hidden motives.",
}
literalism_desc = {
@@ -55,7 +49,7 @@ def build_disposition_description(disposition: DispositionTraits) -> str:
2: "You tend to consider context and implied meaning alongside literal statements.",
3: "You balance literal interpretation with contextual understanding.",
4: "You prefer to interpret information more literally and precisely.",
5: "You interpret information very literally and focus on exact wording and commitments."
5: "You interpret information very literally and focus on exact wording and commitments.",
}
empathy_desc = {
@@ -63,7 +57,7 @@ def build_disposition_description(disposition: DispositionTraits) -> str:
2: "You consider facts first but acknowledge emotional factors exist.",
3: "You balance factual analysis with emotional understanding.",
4: "You give significant weight to emotional context and human factors.",
5: "You strongly consider the emotional state and circumstances of others when forming memories."
5: "You strongly consider the emotional state and circumstances of others when forming memories.",
}
return f"""Your disposition traits:
@@ -72,7 +66,7 @@ def build_disposition_description(disposition: DispositionTraits) -> str:
- Empathy ({describe_trait_level(disposition.empathy)}): {empathy_desc.get(disposition.empathy, empathy_desc[3])}"""
def format_facts_for_prompt(facts: List[MemoryFact]) -> str:
def format_facts_for_prompt(facts: list[MemoryFact]) -> str:
"""Format facts as JSON for LLM prompt."""
import json
@@ -80,9 +74,7 @@ def format_facts_for_prompt(facts: List[MemoryFact]) -> str:
return "[]"
formatted = []
for fact in facts:
fact_obj = {
"text": fact.text
}
fact_obj = {"text": fact.text}
# Add context if available
if fact.context:
@@ -94,7 +86,7 @@ def format_facts_for_prompt(facts: List[MemoryFact]) -> str:
if isinstance(occurred_start, str):
fact_obj["occurred_start"] = occurred_start
elif isinstance(occurred_start, datetime):
fact_obj["occurred_start"] = occurred_start.strftime('%Y-%m-%d %H:%M:%S')
fact_obj["occurred_start"] = occurred_start.strftime("%Y-%m-%d %H:%M:%S")
formatted.append(fact_obj)
@@ -176,16 +168,14 @@ def get_system_message(disposition: DispositionTraits) -> str:
elif disposition.empathy <= 2:
instructions.append("Focus on facts and outcomes rather than emotional context.")
disposition_instruction = " ".join(instructions) if instructions else "Balance your disposition traits when interpreting information."
disposition_instruction = (
" ".join(instructions) if instructions else "Balance your disposition traits when interpreting information."
)
return f"You are a person with your own thoughts, experiences, opinions, and disposition. When asked a question, respond naturally like a human would - share what you genuinely think based on what you know and have experienced. {disposition_instruction} Be direct, express your views confidently, and use 'I think', 'I believe', 'in my view', etc. Respond in plain text without markdown formatting."
async def extract_opinions_from_text(
llm_config,
text: str,
query: str
) -> List[Opinion]:
async def extract_opinions_from_text(llm_config, text: str, query: str) -> list[Opinion]:
"""
Extract opinions with reasons and confidence from text using LLM.
@@ -238,11 +228,14 @@ If no genuine opinions are expressed (e.g., the response just says "I don't know
try:
result = await llm_config.call(
messages=[
{"role": "system", "content": "You are converting opinions from text into first-person statements. Always use 'I think', 'I believe', 'I feel', etc. NEVER use third-person like 'The speaker' or 'They'."},
{"role": "user", "content": extraction_prompt}
{
"role": "system",
"content": "You are converting opinions from text into first-person statements. Always use 'I think', 'I believe', 'I feel', etc. NEVER use third-person like 'The speaker' or 'They'.",
},
{"role": "user", "content": extraction_prompt},
],
response_format=OpinionExtractionResponse,
scope="memory_extract_opinion"
scope="memory_extract_opinion",
)
# Format opinions with confidence score and convert to first-person
@@ -253,14 +246,18 @@ If no genuine opinions are expressed (e.g., the response just says "I don't know
# Replace common third-person patterns with first-person
def singularize_verb(verb):
if verb.endswith('es'):
if verb.endswith("es"):
return verb[:-1] # believes -> believe
elif verb.endswith('s'):
elif verb.endswith("s"):
return verb[:-1] # thinks -> think
return verb
# Pattern: "The speaker/user [verb]..." -> "I [verb]..."
match = re.match(r'^(The speaker|The user|They|It is believed) (believes?|thinks?|feels?|says|asserts?|considers?)(\s+that)?(.*)$', opinion_text, re.IGNORECASE)
match = re.match(
r"^(The speaker|The user|They|It is believed) (believes?|thinks?|feels?|says|asserts?|considers?)(\s+that)?(.*)$",
opinion_text,
re.IGNORECASE,
)
if match:
verb = singularize_verb(match.group(2))
that_part = match.group(3) or "" # Keep " that" if present
@@ -268,14 +265,18 @@ If no genuine opinions are expressed (e.g., the response just says "I don't know
opinion_text = f"I {verb}{that_part}{rest}"
# If still doesn't start with first-person, prepend "I believe that "
first_person_starters = ["I think", "I believe", "I feel", "In my view", "I've come to believe", "Previously I"]
first_person_starters = [
"I think",
"I believe",
"I feel",
"In my view",
"I've come to believe",
"Previously I",
]
if not any(opinion_text.startswith(starter) for starter in first_person_starters):
opinion_text = "I believe that " + opinion_text[0].lower() + opinion_text[1:]
formatted_opinions.append(Opinion(
opinion=opinion_text,
confidence=op.confidence
))
formatted_opinions.append(Opinion(opinion=opinion_text, confidence=op.confidence))
return formatted_opinions
@@ -287,9 +288,9 @@ If no genuine opinions are expressed (e.g., the response just says "I don't know
async def reflect(
llm_config,
query: str,
experience_facts: List[str] = None,
world_facts: List[str] = None,
opinion_facts: List[str] = None,
experience_facts: list[str] = None,
world_facts: list[str] = None,
opinion_facts: list[str] = None,
name: str = "Assistant",
disposition: DispositionTraits = None,
background: str = "",
@@ -320,7 +321,7 @@ async def reflect(
disposition = DispositionTraits(skepticism=3, literalism=3, empathy=3)
# Convert string lists to MemoryFact format for formatting
def to_memory_facts(facts: List[str], fact_type: str) -> List[MemoryFact]:
def to_memory_facts(facts: list[str], fact_type: str) -> list[MemoryFact]:
if not facts:
return []
return [MemoryFact(id=f"test-{i}", text=f, fact_type=fact_type) for i, f in enumerate(facts)]
@@ -350,13 +351,10 @@ async def reflect(
# Call LLM
answer_text = await llm_config.call(
messages=[
{"role": "system", "content": system_message},
{"role": "user", "content": prompt}
],
messages=[{"role": "system", "content": system_message}, {"role": "user", "content": prompt}],
scope="memory_think",
temperature=0.9,
max_completion_tokens=1000
max_completion_tokens=1000,
)
return answer_text.strip()
@@ -4,15 +4,18 @@ Search trace models for debugging and visualization.
These Pydantic models define the structure of search traces, capturing
every step of the spreading activation search process for analysis.
"""
from datetime import datetime
from typing import List, Optional, Dict, Any, Literal
from typing import Any, Literal
from pydantic import BaseModel, Field
class QueryInfo(BaseModel):
"""Information about the search query."""
query_text: str = Field(description="Original query text")
query_embedding: List[float] = Field(description="Generated query embedding vector")
query_embedding: list[float] = Field(description="Generated query embedding vector")
timestamp: datetime = Field(description="When the query was executed")
budget: int = Field(description="Maximum nodes to explore")
max_tokens: int = Field(description="Maximum tokens to return in results")
@@ -20,6 +23,7 @@ class QueryInfo(BaseModel):
class EntryPoint(BaseModel):
"""An entry point node selected for search."""
node_id: str = Field(description="Memory unit ID")
text: str = Field(description="Memory unit text content")
similarity_score: float = Field(description="Cosine similarity to query", ge=0.0, le=1.0)
@@ -28,6 +32,7 @@ class EntryPoint(BaseModel):
class WeightComponents(BaseModel):
"""Breakdown of weight calculation components."""
activation: float = Field(description="Activation from spreading (can exceed 1.0 through accumulation)", ge=0.0)
semantic_similarity: float = Field(description="Semantic similarity to query", ge=0.0, le=1.0)
recency: float = Field(description="Recency weight", ge=0.0, le=1.0)
@@ -43,99 +48,120 @@ class WeightComponents(BaseModel):
class LinkInfo(BaseModel):
"""Information about a link to a neighbor."""
to_node_id: str = Field(description="Target node ID")
link_type: Literal["temporal", "semantic", "entity"] = Field(description="Type of link")
link_weight: float = Field(description="Weight of the link (can exceed 1.0 when aggregating multiple connections)", ge=0.0)
entity_id: Optional[str] = Field(default=None, description="Entity ID if link_type is 'entity'")
new_activation: Optional[float] = Field(default=None, description="Activation that would be passed to neighbor (None for supplementary links)")
link_weight: float = Field(
description="Weight of the link (can exceed 1.0 when aggregating multiple connections)", ge=0.0
)
entity_id: str | None = Field(default=None, description="Entity ID if link_type is 'entity'")
new_activation: float | None = Field(
default=None, description="Activation that would be passed to neighbor (None for supplementary links)"
)
followed: bool = Field(description="Whether this link was followed (or pruned)")
prune_reason: Optional[str] = Field(default=None, description="Why link was not followed (if not followed)")
is_supplementary: bool = Field(default=False, description="Whether this is a supplementary link (multiple connections to same node)")
prune_reason: str | None = Field(default=None, description="Why link was not followed (if not followed)")
is_supplementary: bool = Field(
default=False, description="Whether this is a supplementary link (multiple connections to same node)"
)
class NodeVisit(BaseModel):
"""Information about visiting a node during search."""
step: int = Field(description="Step number in search (1-based)")
node_id: str = Field(description="Memory unit ID")
text: str = Field(description="Memory unit text content")
context: str = Field(description="Memory unit context")
event_date: Optional[datetime] = Field(default=None, description="When the memory occurred")
event_date: datetime | None = Field(default=None, description="When the memory occurred")
access_count: int = Field(description="Number of times accessed before this search")
# How this node was reached
is_entry_point: bool = Field(description="Whether this is an entry point")
parent_node_id: Optional[str] = Field(default=None, description="Node that led to this one")
link_type: Optional[Literal["temporal", "semantic", "entity"]] = Field(default=None, description="Type of link from parent")
link_weight: Optional[float] = Field(default=None, description="Weight of link from parent")
parent_node_id: str | None = Field(default=None, description="Node that led to this one")
link_type: Literal["temporal", "semantic", "entity"] | None = Field(
default=None, description="Type of link from parent"
)
link_weight: float | None = Field(default=None, description="Weight of link from parent")
# Weights
weights: WeightComponents = Field(description="Weight calculation breakdown")
# Neighbors discovered from this node
neighbors_explored: List[LinkInfo] = Field(default_factory=list, description="Links explored from this node")
neighbors_explored: list[LinkInfo] = Field(default_factory=list, description="Links explored from this node")
# Ranking
final_rank: Optional[int] = Field(default=None, description="Final rank in results (1-based, None if not in top-k)")
final_rank: int | None = Field(default=None, description="Final rank in results (1-based, None if not in top-k)")
class PruningDecision(BaseModel):
"""Records when a node was considered but not visited."""
node_id: str = Field(description="Node that was pruned")
reason: Literal["already_visited", "activation_too_low", "budget_exhausted"] = Field(description="Why it was pruned")
reason: Literal["already_visited", "activation_too_low", "budget_exhausted"] = Field(
description="Why it was pruned"
)
activation: float = Field(description="Activation value when pruned")
would_have_been_step: int = Field(description="What step it would have been if visited")
class SearchPhaseMetrics(BaseModel):
"""Performance metrics for a search phase."""
phase_name: str = Field(description="Name of the phase")
duration_seconds: float = Field(description="Time taken in seconds")
details: Dict[str, Any] = Field(default_factory=dict, description="Additional phase-specific metrics")
details: dict[str, Any] = Field(default_factory=dict, description="Additional phase-specific metrics")
class RetrievalResult(BaseModel):
"""A single result from a retrieval method."""
rank: int = Field(description="Rank in this retrieval method (1-based)")
node_id: str = Field(description="Memory unit ID")
text: str = Field(description="Memory unit text content")
context: str = Field(default="", description="Memory unit context")
event_date: Optional[datetime] = Field(default=None, description="When the memory occurred")
fact_type: Optional[str] = Field(default=None, description="Fact type (world, experience, opinion)")
event_date: datetime | None = Field(default=None, description="When the memory occurred")
fact_type: str | None = Field(default=None, description="Fact type (world, experience, opinion)")
score: float = Field(description="Score from this retrieval method")
score_name: str = Field(description="Name of the score (e.g., 'similarity', 'bm25_score', 'activation')")
class RetrievalMethodResults(BaseModel):
"""Results from a single retrieval method."""
method_name: Literal["semantic", "bm25", "graph", "temporal"] = Field(description="Name of retrieval method")
fact_type: Optional[str] = Field(default=None, description="Fact type this retrieval was for (world, experience, opinion)")
results: List[RetrievalResult] = Field(description="Retrieved results with ranks")
fact_type: str | None = Field(
default=None, description="Fact type this retrieval was for (world, experience, opinion)"
)
results: list[RetrievalResult] = Field(description="Retrieved results with ranks")
duration_seconds: float = Field(description="Time taken for this retrieval")
metadata: Dict[str, Any] = Field(default_factory=dict, description="Method-specific metadata")
metadata: dict[str, Any] = Field(default_factory=dict, description="Method-specific metadata")
class RRFMergeResult(BaseModel):
"""A result after RRF merging."""
node_id: str = Field(description="Memory unit ID")
text: str = Field(description="Memory unit text content")
rrf_score: float = Field(description="Reciprocal Rank Fusion score")
source_ranks: Dict[str, int] = Field(description="Rank in each source that contributed (method_name -> rank)")
source_ranks: dict[str, int] = Field(description="Rank in each source that contributed (method_name -> rank)")
final_rrf_rank: int = Field(description="Rank after RRF merge (1-based)")
class RerankedResult(BaseModel):
"""A result after reranking."""
node_id: str = Field(description="Memory unit ID")
text: str = Field(description="Memory unit text content")
rerank_score: float = Field(description="Final reranking score")
rerank_rank: int = Field(description="Rank after reranking (1-based)")
rrf_rank: int = Field(description="Original RRF rank before reranking")
rank_change: int = Field(description="Change in rank (positive = moved up)")
score_components: Dict[str, float] = Field(default_factory=dict, description="Score breakdown")
score_components: dict[str, float] = Field(default_factory=dict, description="Score breakdown")
class SearchSummary(BaseModel):
"""Summary statistics about the search."""
total_nodes_visited: int = Field(description="Total nodes visited")
total_nodes_pruned: int = Field(description="Total nodes pruned")
entry_points_found: int = Field(description="Number of entry points")
@@ -150,33 +176,36 @@ class SearchSummary(BaseModel):
entity_links_followed: int = Field(default=0, description="Entity links followed")
# Phase timings
phase_metrics: List[SearchPhaseMetrics] = Field(default_factory=list, description="Metrics for each phase")
phase_metrics: list[SearchPhaseMetrics] = Field(default_factory=list, description="Metrics for each phase")
class SearchTrace(BaseModel):
"""Complete trace of a search operation."""
query: QueryInfo = Field(description="Query information")
# New 4-way retrieval architecture
retrieval_results: List[RetrievalMethodResults] = Field(default_factory=list, description="Results from each retrieval method")
rrf_merged: List[RRFMergeResult] = Field(default_factory=list, description="Results after RRF merging")
reranked: List[RerankedResult] = Field(default_factory=list, description="Results after reranking")
retrieval_results: list[RetrievalMethodResults] = Field(
default_factory=list, description="Results from each retrieval method"
)
rrf_merged: list[RRFMergeResult] = Field(default_factory=list, description="Results after RRF merging")
reranked: list[RerankedResult] = Field(default_factory=list, description="Results after reranking")
# Legacy fields (kept for backward compatibility with graph/temporal visualizations)
entry_points: List[EntryPoint] = Field(default_factory=list, description="Entry points selected for search (legacy)")
visits: List[NodeVisit] = Field(default_factory=list, description="All nodes visited during search (legacy, for graph viz)")
pruned: List[PruningDecision] = Field(default_factory=list, description="Nodes that were pruned (legacy)")
entry_points: list[EntryPoint] = Field(
default_factory=list, description="Entry points selected for search (legacy)"
)
visits: list[NodeVisit] = Field(
default_factory=list, description="All nodes visited during search (legacy, for graph viz)"
)
pruned: list[PruningDecision] = Field(default_factory=list, description="Nodes that were pruned (legacy)")
summary: SearchSummary = Field(description="Summary statistics")
# Final results (for comparison with visits)
final_results: List[Dict[str, Any]] = Field(description="Final ranked results returned to user")
final_results: list[dict[str, Any]] = Field(description="Final ranked results returned to user")
model_config = {
"json_encoders": {
datetime: lambda v: v.isoformat()
}
}
model_config = {"json_encoders": {datetime: lambda v: v.isoformat()}}
def to_json(self, **kwargs) -> str:
"""Export trace as JSON string."""
@@ -186,14 +215,14 @@ class SearchTrace(BaseModel):
"""Export trace as dictionary."""
return self.model_dump()
def get_visit_by_node_id(self, node_id: str) -> Optional[NodeVisit]:
def get_visit_by_node_id(self, node_id: str) -> NodeVisit | None:
"""Find a visit by node ID."""
for visit in self.visits:
if visit.node_id == node_id:
return visit
return None
def get_search_path_to_node(self, node_id: str) -> List[NodeVisit]:
def get_search_path_to_node(self, node_id: str) -> list[NodeVisit]:
"""Get the path from entry point to a specific node."""
path = []
current_visit = self.get_visit_by_node_id(node_id)
@@ -207,10 +236,10 @@ class SearchTrace(BaseModel):
return path
def get_nodes_by_link_type(self, link_type: Literal["temporal", "semantic", "entity"]) -> List[NodeVisit]:
def get_nodes_by_link_type(self, link_type: Literal["temporal", "semantic", "entity"]) -> list[NodeVisit]:
"""Get all nodes reached via a specific link type."""
return [v for v in self.visits if v.link_type == link_type]
def get_entry_point_nodes(self) -> List[NodeVisit]:
def get_entry_point_nodes(self) -> list[NodeVisit]:
"""Get all entry point visits."""
return [v for v in self.visits if v.is_entry_point]
@@ -4,24 +4,25 @@ Search tracer for collecting detailed search execution traces.
The SearchTracer collects comprehensive information about each step
of the spreading activation search process for debugging and visualization.
"""
import time
from datetime import datetime, timezone
from typing import List, Optional, Dict, Any, Literal
from datetime import UTC, datetime
from typing import Any, Literal
from .trace import (
SearchTrace,
QueryInfo,
EntryPoint,
NodeVisit,
WeightComponents,
LinkInfo,
NodeVisit,
PruningDecision,
SearchSummary,
SearchPhaseMetrics,
RetrievalResult,
RetrievalMethodResults,
RRFMergeResult,
QueryInfo,
RerankedResult,
RetrievalMethodResults,
RetrievalResult,
RRFMergeResult,
SearchPhaseMetrics,
SearchSummary,
SearchTrace,
WeightComponents,
)
@@ -58,17 +59,17 @@ class SearchTracer:
self.max_tokens = max_tokens
# Trace data
self.query_embedding: Optional[List[float]] = None
self.start_time: Optional[float] = None
self.entry_points: List[EntryPoint] = []
self.visits: List[NodeVisit] = []
self.pruned: List[PruningDecision] = []
self.phase_metrics: List[SearchPhaseMetrics] = []
self.query_embedding: list[float] | None = None
self.start_time: float | None = None
self.entry_points: list[EntryPoint] = []
self.visits: list[NodeVisit] = []
self.pruned: list[PruningDecision] = []
self.phase_metrics: list[SearchPhaseMetrics] = []
# New 4-way retrieval tracking
self.retrieval_results: List[RetrievalMethodResults] = []
self.rrf_merged: List[RRFMergeResult] = []
self.reranked: List[RerankedResult] = []
self.retrieval_results: list[RetrievalMethodResults] = []
self.rrf_merged: list[RRFMergeResult] = []
self.reranked: list[RerankedResult] = []
# Tracking state
self.current_step = 0
@@ -83,7 +84,7 @@ class SearchTracer:
"""Start timing the search."""
self.start_time = time.time()
def record_query_embedding(self, embedding: List[float]):
def record_query_embedding(self, embedding: list[float]):
"""Record the query embedding."""
self.query_embedding = embedding
@@ -117,9 +118,9 @@ class SearchTracer:
event_date: datetime,
access_count: int,
is_entry_point: bool,
parent_node_id: Optional[str],
link_type: Optional[Literal["temporal", "semantic", "entity"]],
link_weight: Optional[float],
parent_node_id: str | None,
link_type: Literal["temporal", "semantic", "entity"] | None,
link_weight: float | None,
activation: float,
semantic_similarity: float,
recency: float,
@@ -199,10 +200,10 @@ class SearchTracer:
to_node_id: str,
link_type: Literal["temporal", "semantic", "entity"],
link_weight: float,
entity_id: Optional[str],
new_activation: Optional[float],
entity_id: str | None,
new_activation: float | None,
followed: bool,
prune_reason: Optional[str] = None,
prune_reason: str | None = None,
is_supplementary: bool = False,
):
"""
@@ -266,7 +267,7 @@ class SearchTracer:
)
)
def add_phase_metric(self, phase_name: str, duration_seconds: float, details: Optional[Dict[str, Any]] = None):
def add_phase_metric(self, phase_name: str, duration_seconds: float, details: dict[str, Any] | None = None):
"""
Record metrics for a search phase.
@@ -286,11 +287,11 @@ class SearchTracer:
def add_retrieval_results(
self,
method_name: Literal["semantic", "bm25", "graph", "temporal"],
results: List[tuple], # List of (doc_id, data) tuples
results: list[tuple], # List of (doc_id, data) tuples
duration_seconds: float,
score_field: str, # e.g., "similarity", "bm25_score"
metadata: Optional[Dict[str, Any]] = None,
fact_type: Optional[str] = None
metadata: dict[str, Any] | None = None,
fact_type: str | None = None,
):
"""
Record results from a single retrieval method.
@@ -331,7 +332,7 @@ class SearchTracer:
)
)
def add_rrf_merged(self, merged_results: List[tuple]):
def add_rrf_merged(self, merged_results: list[tuple]):
"""
Record RRF merged results.
@@ -350,7 +351,7 @@ class SearchTracer:
)
)
def add_reranked(self, reranked_results: List[Dict[str, Any]], rrf_merged: List):
def add_reranked(self, reranked_results: list[dict[str, Any]], rrf_merged: list):
"""
Record reranked results.
@@ -373,7 +374,15 @@ class SearchTracer:
# Keys from ScoredResult.to_dict(): cross_encoder_score, cross_encoder_score_normalized,
# rrf_normalized, temporal, recency, combined_score, weight
score_components = {}
for key in ["cross_encoder_score", "cross_encoder_score_normalized", "rrf_score", "rrf_normalized", "temporal", "recency", "combined_score"]:
for key in [
"cross_encoder_score",
"cross_encoder_score_normalized",
"rrf_score",
"rrf_normalized",
"temporal",
"recency",
"combined_score",
]:
if key in result and result[key] is not None:
score_components[key] = result[key]
@@ -389,7 +398,7 @@ class SearchTracer:
)
)
def finalize(self, final_results: List[Dict[str, Any]]) -> SearchTrace:
def finalize(self, final_results: list[dict[str, Any]]) -> SearchTrace:
"""
Finalize the trace and return the complete SearchTrace object.
@@ -416,7 +425,7 @@ class SearchTracer:
query_info = QueryInfo(
query_text=self.query_text,
query_embedding=self.query_embedding or [],
timestamp=datetime.now(timezone.utc),
timestamp=datetime.now(UTC),
budget=self.budget,
max_tokens=self.max_tokens,
)
@@ -6,8 +6,8 @@ providing type safety and making data flow explicit.
"""
from dataclasses import dataclass, field
from typing import Optional, List, Dict, Any
from datetime import datetime
from typing import Any
@dataclass
@@ -17,28 +17,29 @@ class RetrievalResult:
This represents a raw result from the database query, before merging or reranking.
"""
id: str
text: str
fact_type: str
context: Optional[str] = None
event_date: Optional[datetime] = None
occurred_start: Optional[datetime] = None
occurred_end: Optional[datetime] = None
mentioned_at: Optional[datetime] = None
document_id: Optional[str] = None
chunk_id: Optional[str] = None
context: str | None = None
event_date: datetime | None = None
occurred_start: datetime | None = None
occurred_end: datetime | None = None
mentioned_at: datetime | None = None
document_id: str | None = None
chunk_id: str | None = None
access_count: int = 0
embedding: Optional[List[float]] = None
embedding: list[float] | None = None
# Retrieval-specific scores (only one will be set depending on retrieval method)
similarity: Optional[float] = None # Semantic retrieval
bm25_score: Optional[float] = None # BM25 retrieval
activation: Optional[float] = None # Graph retrieval (spreading activation)
temporal_score: Optional[float] = None # Temporal retrieval
temporal_proximity: Optional[float] = None # Temporal retrieval
similarity: float | None = None # Semantic retrieval
bm25_score: float | None = None # BM25 retrieval
activation: float | None = None # Graph retrieval (spreading activation)
temporal_score: float | None = None # Temporal retrieval
temporal_proximity: float | None = None # Temporal retrieval
@classmethod
def from_db_row(cls, row: Dict[str, Any]) -> "RetrievalResult":
def from_db_row(cls, row: dict[str, Any]) -> "RetrievalResult":
"""Create from a database row (asyncpg Record converted to dict)."""
return cls(
id=str(row["id"]),
@@ -68,13 +69,14 @@ class MergedCandidate:
Contains the original retrieval data plus RRF metadata.
"""
# Original retrieval data
retrieval: RetrievalResult
# RRF metadata
rrf_score: float
rrf_rank: int = 0
source_ranks: Dict[str, int] = field(default_factory=dict) # method_name -> rank
source_ranks: dict[str, int] = field(default_factory=dict) # method_name -> rank
@property
def id(self) -> str:
@@ -89,6 +91,7 @@ class ScoredResult:
Contains all retrieval/merge data plus reranking scores and combined score.
"""
# Original merged candidate
candidate: MergedCandidate
@@ -115,7 +118,7 @@ class ScoredResult:
"""Convenience property to access retrieval data."""
return self.candidate.retrieval
def to_dict(self) -> Dict[str, Any]:
def to_dict(self) -> dict[str, Any]:
"""
Convert to dict for backwards compatibility.
@@ -6,10 +6,12 @@ This provides an abstraction that can be adapted to different execution models:
- Pub/Sub architectures (future)
- Message brokers (future)
"""
from abc import ABC, abstractmethod
from typing import Any, Dict, Optional, Callable, Awaitable
import asyncio
import logging
from abc import ABC, abstractmethod
from collections.abc import Awaitable, Callable
from typing import Any
logger = logging.getLogger(__name__)
@@ -29,10 +31,10 @@ class TaskBackend(ABC):
def __init__(self):
"""Initialize the task backend."""
self._executor: Optional[Callable[[Dict[str, Any]], Awaitable[None]]] = None
self._executor: Callable[[dict[str, Any]], Awaitable[None]] | None = None
self._initialized = False
def set_executor(self, executor: Callable[[Dict[str, Any]], Awaitable[None]]):
def set_executor(self, executor: Callable[[dict[str, Any]], Awaitable[None]]):
"""
Set the executor callback for processing tasks.
@@ -49,7 +51,7 @@ class TaskBackend(ABC):
pass
@abstractmethod
async def submit_task(self, task_dict: Dict[str, Any]):
async def submit_task(self, task_dict: dict[str, Any]):
"""
Submit a task for execution.
@@ -65,7 +67,7 @@ class TaskBackend(ABC):
"""
pass
async def _execute_task(self, task_dict: Dict[str, Any]):
async def _execute_task(self, task_dict: dict[str, Any]):
"""
Execute a task through the registered executor.
@@ -73,16 +75,17 @@ class TaskBackend(ABC):
task_dict: Task dictionary to execute
"""
if self._executor is None:
task_type = task_dict.get('type', 'unknown')
task_type = task_dict.get("type", "unknown")
logger.warning(f"No executor registered, skipping task {task_type}")
return
try:
await self._executor(task_dict)
except Exception as e:
task_type = task_dict.get('type', 'unknown')
task_type = task_dict.get("type", "unknown")
logger.error(f"Error executing task {task_type}: {e}")
import traceback
traceback.print_exc()
@@ -94,11 +97,7 @@ class AsyncIOQueueBackend(TaskBackend):
and a periodic consumer worker.
"""
def __init__(
self,
batch_size: int = 100,
batch_interval: float = 1.0
):
def __init__(self, batch_size: int = 100, batch_interval: float = 1.0):
"""
Initialize AsyncIO queue backend.
@@ -107,9 +106,9 @@ class AsyncIOQueueBackend(TaskBackend):
batch_interval: Maximum time (seconds) to wait before processing batch
"""
super().__init__()
self._queue: Optional[asyncio.Queue] = None
self._worker_task: Optional[asyncio.Task] = None
self._shutdown_event: Optional[asyncio.Event] = None
self._queue: asyncio.Queue | None = None
self._worker_task: asyncio.Task | None = None
self._shutdown_event: asyncio.Event | None = None
self._batch_size = batch_size
self._batch_interval = batch_interval
@@ -124,7 +123,7 @@ class AsyncIOQueueBackend(TaskBackend):
self._initialized = True
logger.info("AsyncIOQueueBackend initialized")
async def submit_task(self, task_dict: Dict[str, Any]):
async def submit_task(self, task_dict: dict[str, Any]):
"""
Submit a task by putting it in the queue.
@@ -135,8 +134,8 @@ class AsyncIOQueueBackend(TaskBackend):
await self.initialize()
await self._queue.put(task_dict)
task_type = task_dict.get('type', 'unknown')
task_id = task_dict.get('id')
task_type = task_dict.get("type", "unknown")
task_id = task_dict.get("id")
async def wait_for_pending_tasks(self, timeout: float = 5.0):
"""
@@ -200,20 +199,16 @@ class AsyncIOQueueBackend(TaskBackend):
while len(tasks) < self._batch_size and asyncio.get_event_loop().time() < deadline:
try:
remaining_time = max(0.1, deadline - asyncio.get_event_loop().time())
task_dict = await asyncio.wait_for(
self._queue.get(),
timeout=remaining_time
)
task_dict = await asyncio.wait_for(self._queue.get(), timeout=remaining_time)
tasks.append(task_dict)
except asyncio.TimeoutError:
except TimeoutError:
break
# Process batch
if tasks:
# Execute tasks concurrently
await asyncio.gather(
*[self._execute_task(task_dict) for task_dict in tasks],
return_exceptions=True
*[self._execute_task(task_dict) for task_dict in tasks], return_exceptions=True
)
except asyncio.CancelledError:
+25 -10
View File
@@ -1,9 +1,10 @@
"""
Utility functions for memory system.
"""
import logging
from datetime import datetime
from typing import List, Dict, TYPE_CHECKING
from typing import TYPE_CHECKING
if TYPE_CHECKING:
from .llm_wrapper import LLMConfig
@@ -12,7 +13,14 @@ if TYPE_CHECKING:
from .retain.fact_extraction import extract_facts_from_text
async def extract_facts(text: str, event_date: datetime, context: str = "", llm_config: 'LLMConfig' = None, agent_name: str = None, extract_opinions: bool = False) -> tuple[List['Fact'], List[tuple[str, int]]]:
async def extract_facts(
text: str,
event_date: datetime,
context: str = "",
llm_config: "LLMConfig" = None,
agent_name: str = None,
extract_opinions: bool = False,
) -> tuple[list["Fact"], list[tuple[str, int]]]:
"""
Extract semantic facts from text using LLM.
@@ -41,16 +49,25 @@ async def extract_facts(text: str, event_date: datetime, context: str = "", llm_
if not text or not text.strip():
return [], []
facts, chunks = await extract_facts_from_text(text, event_date, context=context, llm_config=llm_config, agent_name=agent_name, extract_opinions=extract_opinions)
facts, chunks = await extract_facts_from_text(
text,
event_date,
context=context,
llm_config=llm_config,
agent_name=agent_name,
extract_opinions=extract_opinions,
)
if not facts:
logging.warning(f"LLM extracted 0 facts from text of length {len(text)}. This may indicate the text contains no meaningful information, or the LLM failed to extract facts. Full text: {text}")
logging.warning(
f"LLM extracted 0 facts from text of length {len(text)}. This may indicate the text contains no meaningful information, or the LLM failed to extract facts. Full text: {text}"
)
return [], chunks
return facts, chunks
def cosine_similarity(vec1: List[float], vec2: List[float]) -> float:
def cosine_similarity(vec1: list[float], vec2: list[float]) -> float:
"""
Calculate cosine similarity between two vectors.
@@ -100,6 +117,7 @@ def calculate_recency_weight(days_since: float, half_life_days: float = 365.0) -
Weight between 0 and 1
"""
import math
# Logarithmic decay: 1 / (1 + log(1 + days_since/half_life))
# This decays much slower than exponential, giving better long-term differentiation
normalized_age = days_since / half_life_days
@@ -121,6 +139,7 @@ def calculate_frequency_weight(access_count: int, max_boost: float = 2.0) -> flo
Weight between 1.0 and max_boost
"""
import math
if access_count <= 0:
return 1.0
@@ -158,11 +177,7 @@ def calculate_temporal_anchor(occurred_start: datetime, occurred_end: datetime)
return midpoint
def calculate_temporal_proximity(
anchor_a: datetime,
anchor_b: datetime,
half_life_days: float = 30.0
) -> float:
def calculate_temporal_proximity(anchor_a: datetime, anchor_b: datetime, half_life_days: float = 30.0) -> float:
"""
Calculate temporal proximity between two temporal anchors.
+21 -40
View File
@@ -6,6 +6,7 @@ Run the server with:
Stop with Ctrl+C.
"""
import argparse
import asyncio
import atexit
@@ -13,15 +14,14 @@ import os
import signal
import sys
import warnings
from typing import Optional
import uvicorn
from . import MemoryEngine
from .api import create_app
from .config import get_config, HindsightConfig
from .banner import print_banner
from .config import HindsightConfig, get_config
print()
print_banner()
@@ -33,7 +33,7 @@ warnings.filterwarnings("ignore", message="websockets.server.WebSocketServerProt
os.environ["TOKENIZERS_PARALLELISM"] = "false"
# Global reference for cleanup
_memory: Optional[MemoryEngine] = None
_memory: MemoryEngine | None = None
def _cleanup():
@@ -70,59 +70,41 @@ def main():
# Server options
parser.add_argument(
"--host", default=config.host,
help=f"Host to bind to (default: {config.host}, env: HINDSIGHT_API_HOST)"
"--host", default=config.host, help=f"Host to bind to (default: {config.host}, env: HINDSIGHT_API_HOST)"
)
parser.add_argument(
"--port", type=int, default=config.port,
help=f"Port to bind to (default: {config.port}, env: HINDSIGHT_API_PORT)"
"--port",
type=int,
default=config.port,
help=f"Port to bind to (default: {config.port}, env: HINDSIGHT_API_PORT)",
)
parser.add_argument(
"--log-level", default=config.log_level,
"--log-level",
default=config.log_level,
choices=["critical", "error", "warning", "info", "debug", "trace"],
help=f"Log level (default: {config.log_level}, env: HINDSIGHT_API_LOG_LEVEL)"
help=f"Log level (default: {config.log_level}, env: HINDSIGHT_API_LOG_LEVEL)",
)
# Development options
parser.add_argument(
"--reload", action="store_true",
help="Enable auto-reload on code changes (development only)"
)
parser.add_argument(
"--workers", type=int, default=1,
help="Number of worker processes (default: 1)"
)
parser.add_argument("--reload", action="store_true", help="Enable auto-reload on code changes (development only)")
parser.add_argument("--workers", type=int, default=1, help="Number of worker processes (default: 1)")
# Access log options
parser.add_argument(
"--access-log", action="store_true",
help="Enable access log"
)
parser.add_argument(
"--no-access-log", dest="access_log", action="store_false",
help="Disable access log (default)"
)
parser.add_argument("--access-log", action="store_true", help="Enable access log")
parser.add_argument("--no-access-log", dest="access_log", action="store_false", help="Disable access log (default)")
parser.set_defaults(access_log=False)
# Proxy options
parser.add_argument(
"--proxy-headers", action="store_true",
help="Enable X-Forwarded-Proto, X-Forwarded-For headers"
"--proxy-headers", action="store_true", help="Enable X-Forwarded-Proto, X-Forwarded-For headers"
)
parser.add_argument(
"--forwarded-allow-ips", default=None,
help="Comma separated list of IPs to trust with proxy headers"
"--forwarded-allow-ips", default=None, help="Comma separated list of IPs to trust with proxy headers"
)
# SSL options
parser.add_argument(
"--ssl-keyfile", default=None,
help="SSL key file"
)
parser.add_argument(
"--ssl-certfile", default=None,
help="SSL certificate file"
)
parser.add_argument("--ssl-keyfile", default=None, help="SSL key file")
parser.add_argument("--ssl-certfile", default=None, help="SSL certificate file")
args = parser.parse_args()
@@ -188,9 +170,8 @@ def main():
if args.ssl_certfile:
uvicorn_config["ssl_certfile"] = args.ssl_certfile
from .banner import print_startup_info
print_startup_info(
host=args.host,
port=args.port,
+5 -7
View File
@@ -38,8 +38,8 @@ import sys
from mcp.server.fastmcp import FastMCP
from hindsight_api.config import (
ENV_MCP_LOCAL_BANK_ID,
DEFAULT_MCP_LOCAL_BANK_ID,
ENV_MCP_LOCAL_BANK_ID,
)
# Configure logging - default to info
@@ -103,10 +103,7 @@ def create_local_mcp_server(bank_id: str, memory=None) -> FastMCP:
async def _retain():
try:
await memory.retain_batch_async(
bank_id=bank_id,
contents=[{"content": content, "context": context}]
)
await memory.retain_batch_async(bank_id=bank_id, contents=[{"content": content, "context": context}])
except Exception as e:
logger.error(f"Error storing memory: {e}", exc_info=True)
@@ -140,7 +137,7 @@ def create_local_mcp_server(bank_id: str, memory=None) -> FastMCP:
query=query,
fact_type=list(VALID_RECALL_FACT_TYPES),
budget=budget_enum,
max_tokens=max_tokens
max_tokens=max_tokens,
)
return search_result.model_dump()
@@ -169,7 +166,8 @@ async def _initialize_and_run(bank_id: str):
def main():
"""Main entry point for the stdio MCP server."""
import asyncio
from hindsight_api.config import get_config, ENV_LLM_API_KEY
from hindsight_api.config import ENV_LLM_API_KEY, get_config
# Check for required environment variables
config = get_config()
+44 -30
View File
@@ -6,16 +6,15 @@ This module provides metrics for:
- Token usage (input/output) per operation
- Per-bank granularity via labels
"""
import logging
from typing import Dict, Any, Optional
from contextlib import contextmanager
import time
from contextlib import contextmanager
from opentelemetry import metrics
from opentelemetry.exporter.prometheus import PrometheusMetricReader
from opentelemetry.sdk.metrics import MeterProvider
from opentelemetry.sdk.resources import Resource
from opentelemetry.exporter.prometheus import PrometheusMetricReader
from prometheus_client import REGISTRY
logger = logging.getLogger(__name__)
@@ -39,19 +38,18 @@ def initialize_metrics(service_name: str = "hindsight-api", service_version: str
global _meter
# Create resource with service information
resource = Resource.create({
"service.name": service_name,
"service.version": service_version,
})
resource = Resource.create(
{
"service.name": service_name,
"service.version": service_version,
}
)
# Create Prometheus metric reader
prometheus_reader = PrometheusMetricReader()
# Create meter provider with Prometheus exporter
provider = MeterProvider(
resource=resource,
metric_readers=[prometheus_reader]
)
provider = MeterProvider(resource=resource, metric_readers=[prometheus_reader])
# Set the global meter provider
metrics.set_meter_provider(provider)
@@ -73,11 +71,19 @@ class MetricsCollectorBase:
"""Base class for metrics collectors."""
@contextmanager
def record_operation(self, operation: str, bank_id: str, budget: Optional[str] = None, max_tokens: Optional[int] = None):
def record_operation(self, operation: str, bank_id: str, budget: str | None = None, max_tokens: int | None = None):
"""Context manager to record operation duration and status."""
raise NotImplementedError
def record_tokens(self, operation: str, bank_id: str, input_tokens: int = 0, output_tokens: int = 0, budget: Optional[str] = None, max_tokens: Optional[int] = None):
def record_tokens(
self,
operation: str,
bank_id: str,
input_tokens: int = 0,
output_tokens: int = 0,
budget: str | None = None,
max_tokens: int | None = None,
):
"""Record token usage for an operation."""
raise NotImplementedError
@@ -86,11 +92,19 @@ class NoOpMetricsCollector(MetricsCollectorBase):
"""No-op metrics collector that does nothing. Used when metrics are disabled."""
@contextmanager
def record_operation(self, operation: str, bank_id: str, budget: Optional[str] = None, max_tokens: Optional[int] = None):
def record_operation(self, operation: str, bank_id: str, budget: str | None = None, max_tokens: int | None = None):
"""No-op context manager."""
yield
def record_tokens(self, operation: str, bank_id: str, input_tokens: int = 0, output_tokens: int = 0, budget: Optional[str] = None, max_tokens: Optional[int] = None):
def record_tokens(
self,
operation: str,
bank_id: str,
input_tokens: int = 0,
output_tokens: int = 0,
budget: str | None = None,
max_tokens: int | None = None,
):
"""No-op token recording."""
pass
@@ -108,33 +122,25 @@ class MetricsCollector(MetricsCollectorBase):
# Operation latency histogram (in seconds)
# Records duration of retain, recall, reflect operations
self.operation_duration = self.meter.create_histogram(
name="hindsight.operation.duration",
description="Duration of Hindsight operations in seconds",
unit="s"
name="hindsight.operation.duration", description="Duration of Hindsight operations in seconds", unit="s"
)
# Token usage counters
self.tokens_input = self.meter.create_counter(
name="hindsight.tokens.input",
description="Number of input tokens consumed",
unit="tokens"
name="hindsight.tokens.input", description="Number of input tokens consumed", unit="tokens"
)
self.tokens_output = self.meter.create_counter(
name="hindsight.tokens.output",
description="Number of output tokens generated",
unit="tokens"
name="hindsight.tokens.output", description="Number of output tokens generated", unit="tokens"
)
# Operation counter (success/failure)
self.operation_total = self.meter.create_counter(
name="hindsight.operation.total",
description="Total number of operations executed",
unit="operations"
name="hindsight.operation.total", description="Total number of operations executed", unit="operations"
)
@contextmanager
def record_operation(self, operation: str, bank_id: str, budget: Optional[str] = None, max_tokens: Optional[int] = None):
def record_operation(self, operation: str, bank_id: str, budget: str | None = None, max_tokens: int | None = None):
"""
Context manager to record operation duration and status.
@@ -175,7 +181,15 @@ class MetricsCollector(MetricsCollectorBase):
# Record operation count
self.operation_total.add(1, attributes)
def record_tokens(self, operation: str, bank_id: str, input_tokens: int = 0, output_tokens: int = 0, budget: Optional[str] = None, max_tokens: Optional[int] = None):
def record_tokens(
self,
operation: str,
bank_id: str,
input_tokens: int = 0,
output_tokens: int = 0,
budget: str | None = None,
max_tokens: int | None = None,
):
"""
Record token usage for an operation.
+10 -8
View File
@@ -11,11 +11,10 @@ safe rolling deployments.
No alembic.ini required - all configuration is done programmatically.
"""
import logging
import os
import shutil
from pathlib import Path
from typing import Optional
from alembic import command
from alembic.config import Config
@@ -31,7 +30,7 @@ def _run_migrations_internal(database_url: str, script_location: str) -> None:
"""
Internal function to run migrations without locking.
"""
logger.info(f"Running database migrations to head...")
logger.info("Running database migrations to head...")
logger.info(f"Database URL: {database_url}")
logger.info(f"Script location: {script_location}")
@@ -57,7 +56,7 @@ def _run_migrations_internal(database_url: str, script_location: str) -> None:
logger.info("Database migrations completed successfully")
def run_migrations(database_url: str, script_location: Optional[str] = None) -> None:
def run_migrations(database_url: str, script_location: str | None = None) -> None:
"""
Run database migrations to the latest version using programmatic Alembic configuration.
@@ -97,8 +96,7 @@ def run_migrations(database_url: str, script_location: Optional[str] = None) ->
script_path = Path(script_location)
if not script_path.exists():
raise FileNotFoundError(
f"Alembic script location not found at {script_location}. "
"Database migrations cannot be run."
f"Alembic script location not found at {script_location}. Database migrations cannot be run."
)
# Use PostgreSQL advisory lock to coordinate between distributed workers
@@ -130,7 +128,9 @@ def run_migrations(database_url: str, script_location: Optional[str] = None) ->
raise RuntimeError("Database migration failed") from e
def check_migration_status(database_url: Optional[str] = None, script_location: Optional[str] = None) -> tuple[str | None, str | None]:
def check_migration_status(
database_url: str | None = None, script_location: str | None = None
) -> tuple[str | None, str | None]:
"""
Check current database schema version and latest available version.
@@ -151,7 +151,9 @@ def check_migration_status(database_url: Optional[str] = None, script_location:
if database_url is None:
database_url = os.getenv("HINDSIGHT_API_DATABASE_URL")
if not database_url:
logger.warning("Database URL not provided and HINDSIGHT_API_DATABASE_URL not set, cannot check migration status")
logger.warning(
"Database URL not provided and HINDSIGHT_API_DATABASE_URL not set, cannot check migration status"
)
return None, None
# Get current revision from database
+61 -73
View File
@@ -1,49 +1,47 @@
"""
SQLAlchemy models for the memory system.
"""
from datetime import datetime
from typing import Optional
from uuid import UUID as PyUUID, uuid4
from datetime import datetime
from uuid import UUID as PyUUID
from pgvector.sqlalchemy import Vector
from sqlalchemy import (
CheckConstraint,
Column,
Float,
ForeignKey,
ForeignKeyConstraint,
Index,
Integer,
PrimaryKeyConstraint,
Text,
func,
)
from sqlalchemy import (
text as sql_text,
)
from sqlalchemy.dialects.postgresql import JSONB, TIMESTAMP, UUID
from sqlalchemy.ext.asyncio import AsyncAttrs
from sqlalchemy.orm import DeclarativeBase, Mapped, mapped_column, relationship
from pgvector.sqlalchemy import Vector
class Base(AsyncAttrs, DeclarativeBase):
"""Base class for all models."""
pass
class Document(Base):
"""Source documents for memory units."""
__tablename__ = "documents"
id: Mapped[str] = mapped_column(Text, primary_key=True)
bank_id: Mapped[str] = mapped_column(Text, primary_key=True)
original_text: Mapped[Optional[str]] = mapped_column(Text)
content_hash: Mapped[Optional[str]] = mapped_column(Text)
original_text: Mapped[str | None] = mapped_column(Text)
content_hash: Mapped[str | None] = mapped_column(Text)
doc_metadata: Mapped[dict] = mapped_column("metadata", JSONB, server_default=sql_text("'{}'::jsonb"))
created_at: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), server_default=func.now()
)
updated_at: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), server_default=func.now()
)
created_at: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
updated_at: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
# Relationships
memory_units = relationship("MemoryUnit", back_populates="document", cascade="all, delete-orphan")
@@ -56,45 +54,42 @@ class Document(Base):
class MemoryUnit(Base):
"""Individual sentence-level memories."""
__tablename__ = "memory_units"
id: Mapped[PyUUID] = mapped_column(
UUID(as_uuid=True), primary_key=True, server_default=sql_text("gen_random_uuid()")
)
bank_id: Mapped[str] = mapped_column(Text, nullable=False)
document_id: Mapped[Optional[str]] = mapped_column(Text)
document_id: Mapped[str | None] = mapped_column(Text)
text: Mapped[str] = mapped_column(Text, nullable=False)
embedding = mapped_column(Vector(384)) # pgvector type
context: Mapped[Optional[str]] = mapped_column(Text)
event_date: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), nullable=False) # Kept for backward compatibility
occurred_start: Mapped[Optional[datetime]] = mapped_column(TIMESTAMP(timezone=True)) # When fact occurred (range start)
occurred_end: Mapped[Optional[datetime]] = mapped_column(TIMESTAMP(timezone=True)) # When fact occurred (range end)
mentioned_at: Mapped[Optional[datetime]] = mapped_column(TIMESTAMP(timezone=True)) # When fact was mentioned
context: Mapped[str | None] = mapped_column(Text)
event_date: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), nullable=False
) # Kept for backward compatibility
occurred_start: Mapped[datetime | None] = mapped_column(
TIMESTAMP(timezone=True)
) # When fact occurred (range start)
occurred_end: Mapped[datetime | None] = mapped_column(TIMESTAMP(timezone=True)) # When fact occurred (range end)
mentioned_at: Mapped[datetime | None] = mapped_column(TIMESTAMP(timezone=True)) # When fact was mentioned
fact_type: Mapped[str] = mapped_column(Text, nullable=False, server_default="world")
confidence_score: Mapped[Optional[float]] = mapped_column(Float)
confidence_score: Mapped[float | None] = mapped_column(Float)
access_count: Mapped[int] = mapped_column(Integer, server_default="0")
unit_metadata: Mapped[dict] = mapped_column("metadata", JSONB, server_default=sql_text("'{}'::jsonb")) # User-defined metadata (str->str)
created_at: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), server_default=func.now()
)
updated_at: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), server_default=func.now()
)
unit_metadata: Mapped[dict] = mapped_column(
"metadata", JSONB, server_default=sql_text("'{}'::jsonb")
) # User-defined metadata (str->str)
created_at: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
updated_at: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
# Relationships
document = relationship("Document", back_populates="memory_units")
unit_entities = relationship("UnitEntity", back_populates="memory_unit", cascade="all, delete-orphan")
outgoing_links = relationship(
"MemoryLink",
foreign_keys="MemoryLink.from_unit_id",
back_populates="from_unit",
cascade="all, delete-orphan"
"MemoryLink", foreign_keys="MemoryLink.from_unit_id", back_populates="from_unit", cascade="all, delete-orphan"
)
incoming_links = relationship(
"MemoryLink",
foreign_keys="MemoryLink.to_unit_id",
back_populates="to_unit",
cascade="all, delete-orphan"
"MemoryLink", foreign_keys="MemoryLink.to_unit_id", back_populates="to_unit", cascade="all, delete-orphan"
)
__table_args__ = (
@@ -110,7 +105,7 @@ class MemoryUnit(Base):
"(fact_type = 'opinion' AND confidence_score IS NOT NULL) OR "
"(fact_type = 'observation') OR "
"(fact_type NOT IN ('opinion', 'observation') AND confidence_score IS NULL)",
name="confidence_score_fact_type_check"
name="confidence_score_fact_type_check",
),
Index("idx_memory_units_bank_id", "bank_id"),
Index("idx_memory_units_document_id", "document_id"),
@@ -119,39 +114,46 @@ class MemoryUnit(Base):
Index("idx_memory_units_access_count", "access_count", postgresql_ops={"access_count": "DESC"}),
Index("idx_memory_units_fact_type", "fact_type"),
Index("idx_memory_units_bank_fact_type", "bank_id", "fact_type"),
Index("idx_memory_units_bank_type_date", "bank_id", "fact_type", "event_date", postgresql_ops={"event_date": "DESC"}),
Index(
"idx_memory_units_bank_type_date",
"bank_id",
"fact_type",
"event_date",
postgresql_ops={"event_date": "DESC"},
),
Index(
"idx_memory_units_opinion_confidence",
"bank_id",
"confidence_score",
postgresql_where=sql_text("fact_type = 'opinion'"),
postgresql_ops={"confidence_score": "DESC"}
postgresql_ops={"confidence_score": "DESC"},
),
Index(
"idx_memory_units_opinion_date",
"bank_id",
"event_date",
postgresql_where=sql_text("fact_type = 'opinion'"),
postgresql_ops={"event_date": "DESC"}
postgresql_ops={"event_date": "DESC"},
),
Index(
"idx_memory_units_observation_date",
"bank_id",
"event_date",
postgresql_where=sql_text("fact_type = 'observation'"),
postgresql_ops={"event_date": "DESC"}
postgresql_ops={"event_date": "DESC"},
),
Index(
"idx_memory_units_embedding",
"embedding",
postgresql_using="hnsw",
postgresql_ops={"embedding": "vector_cosine_ops"}
postgresql_ops={"embedding": "vector_cosine_ops"},
),
)
class Entity(Base):
"""Resolved entities (people, organizations, locations, etc.)."""
__tablename__ = "entities"
id: Mapped[PyUUID] = mapped_column(
@@ -160,12 +162,8 @@ class Entity(Base):
canonical_name: Mapped[str] = mapped_column(Text, nullable=False)
bank_id: Mapped[str] = mapped_column(Text, nullable=False)
entity_metadata: Mapped[dict] = mapped_column("metadata", JSONB, server_default=sql_text("'{}'::jsonb"))
first_seen: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), server_default=func.now()
)
last_seen: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), server_default=func.now()
)
first_seen: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
last_seen: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
mention_count: Mapped[int] = mapped_column(Integer, server_default="1")
# Relationships
@@ -175,13 +173,13 @@ class Entity(Base):
"EntityCooccurrence",
foreign_keys="EntityCooccurrence.entity_id_1",
back_populates="entity_1",
cascade="all, delete-orphan"
cascade="all, delete-orphan",
)
cooccurrences_2 = relationship(
"EntityCooccurrence",
foreign_keys="EntityCooccurrence.entity_id_2",
back_populates="entity_2",
cascade="all, delete-orphan"
cascade="all, delete-orphan",
)
__table_args__ = (
@@ -193,6 +191,7 @@ class Entity(Base):
class UnitEntity(Base):
"""Association between memory units and entities."""
__tablename__ = "unit_entities"
unit_id: Mapped[PyUUID] = mapped_column(
@@ -214,6 +213,7 @@ class UnitEntity(Base):
class EntityCooccurrence(Base):
"""Materialized cache of entity co-occurrences."""
__tablename__ = "entity_cooccurrences"
entity_id_1: Mapped[PyUUID] = mapped_column(
@@ -223,9 +223,7 @@ class EntityCooccurrence(Base):
UUID(as_uuid=True), ForeignKey("entities.id", ondelete="CASCADE"), primary_key=True
)
cooccurrence_count: Mapped[int] = mapped_column(Integer, server_default="1")
last_cooccurred: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), server_default=func.now()
)
last_cooccurred: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
# Relationships
entity_1 = relationship("Entity", foreign_keys=[entity_id_1], back_populates="cooccurrences_1")
@@ -241,6 +239,7 @@ class EntityCooccurrence(Base):
class MemoryLink(Base):
"""Links between memory units (temporal, semantic, entity)."""
__tablename__ = "memory_links"
from_unit_id: Mapped[PyUUID] = mapped_column(
@@ -250,13 +249,11 @@ class MemoryLink(Base):
UUID(as_uuid=True), ForeignKey("memory_units.id", ondelete="CASCADE"), primary_key=True
)
link_type: Mapped[str] = mapped_column(Text, primary_key=True)
entity_id: Mapped[Optional[PyUUID]] = mapped_column(
entity_id: Mapped[PyUUID | None] = mapped_column(
UUID(as_uuid=True), ForeignKey("entities.id", ondelete="CASCADE"), primary_key=True
)
weight: Mapped[float] = mapped_column(Float, nullable=False, server_default="1.0")
created_at: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), server_default=func.now()
)
created_at: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
# Relationships
from_unit = relationship("MemoryUnit", foreign_keys=[from_unit_id], back_populates="outgoing_links")
@@ -266,7 +263,7 @@ class MemoryLink(Base):
__table_args__ = (
CheckConstraint(
"link_type IN ('temporal', 'semantic', 'entity', 'causes', 'caused_by', 'enables', 'prevents')",
name="memory_links_link_type_check"
name="memory_links_link_type_check",
),
CheckConstraint("weight >= 0.0 AND weight <= 1.0", name="memory_links_weight_check"),
Index("idx_memory_links_from", "from_unit_id"),
@@ -278,31 +275,22 @@ class MemoryLink(Base):
"from_unit_id",
"weight",
postgresql_where=sql_text("weight >= 0.1"),
postgresql_ops={"weight": "DESC"}
postgresql_ops={"weight": "DESC"},
),
)
class Bank(Base):
"""Memory bank profiles with disposition traits and background."""
__tablename__ = "banks"
bank_id: Mapped[str] = mapped_column(Text, primary_key=True)
disposition: Mapped[dict] = mapped_column(
JSONB,
nullable=False,
server_default=sql_text(
'\'{"skepticism": 3, "literalism": 3, "empathy": 3}\'::jsonb'
)
JSONB, nullable=False, server_default=sql_text('\'{"skepticism": 3, "literalism": 3, "empathy": 3}\'::jsonb')
)
background: Mapped[str] = mapped_column(Text, nullable=False, server_default="")
created_at: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), server_default=func.now()
)
updated_at: Mapped[datetime] = mapped_column(
TIMESTAMP(timezone=True), server_default=func.now()
)
created_at: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
updated_at: Mapped[datetime] = mapped_column(TIMESTAMP(timezone=True), server_default=func.now())
__table_args__ = (
Index("idx_banks_bank_id", "bank_id"),
)
__table_args__ = (Index("idx_banks_bank_id", "bank_id"),)
+4 -6
View File
@@ -1,6 +1,5 @@
import asyncio
import logging
from typing import Optional
from pg0 import Pg0
@@ -16,7 +15,7 @@ class EmbeddedPostgres:
def __init__(
self,
port: Optional[int] = None,
port: int | None = None,
username: str = DEFAULT_USERNAME,
password: str = DEFAULT_PASSWORD,
database: str = DEFAULT_DATABASE,
@@ -28,7 +27,7 @@ class EmbeddedPostgres:
self.password = password
self.database = database
self.name = name
self._pg0: Optional[Pg0] = None
self._pg0: Pg0 | None = None
def _get_pg0(self) -> Pg0:
if self._pg0 is None:
@@ -71,8 +70,7 @@ class EmbeddedPostgres:
logger.debug(f"pg0 start attempt {attempt}/{max_retries} failed: {last_error}")
raise RuntimeError(
f"Failed to start embedded PostgreSQL after {max_retries} attempts. "
f"Last error: {last_error}"
f"Failed to start embedded PostgreSQL after {max_retries} attempts. Last error: {last_error}"
)
async def stop(self) -> None:
@@ -113,7 +111,7 @@ class EmbeddedPostgres:
return await self.start()
_default_instance: Optional[EmbeddedPostgres] = None
_default_instance: EmbeddedPostgres | None = None
def get_embedded_postgres() -> EmbeddedPostgres:
+3 -6
View File
@@ -6,6 +6,7 @@ This module provides the ASGI app for uvicorn import string usage:
For CLI usage, use the hindsight-api command instead.
"""
import os
import warnings
@@ -29,15 +30,11 @@ config.configure_logging()
_memory = MemoryEngine()
# Create unified app with both HTTP and optionally MCP
app = create_app(
memory=_memory,
http_api_enabled=True,
mcp_api_enabled=config.mcp_enabled,
mcp_mount_path="/mcp"
)
app = create_app(memory=_memory, http_api_enabled=True, mcp_api_enabled=config.mcp_enabled, mcp_mount_path="/mcp")
if __name__ == "__main__":
# When run directly, delegate to the CLI
from hindsight_api.main import main
main()
+9 -3
View File
@@ -97,6 +97,10 @@ dev = [
[tool.ruff]
line-length = 120
target-version = "py311"
exclude = [
"tests/",
"**/tests/",
]
[tool.ruff.lint]
select = [
@@ -104,12 +108,14 @@ select = [
"W", # pycodestyle warnings
"F", # Pyflakes
"I", # isort
"B", # flake8-bugbear
"UP", # pyupgrade
]
ignore = [
"E501", # line too long (handled by formatter)
"B008", # do not perform function calls in argument defaults
"E402", # module import not at top of file
"F401", # unused import (too noisy during development)
"F841", # unused variable (too noisy during development)
"F811", # redefined while unused
"F821", # undefined name (forward references in type hints)
]
[tool.ruff.format]
@@ -12,12 +12,12 @@ This comprehensive test suite validates that the fact extraction system:
These are quality/accuracy tests that verify the LLM-based extraction
produces semantically correct and complete facts.
"""
import pytest
import re
from datetime import datetime, timezone
from hindsight_api.engine.retain.fact_extraction import extract_facts_from_text
from hindsight_api import LLMConfig
from datetime import UTC, datetime
import pytest
from hindsight_api import LLMConfig
from hindsight_api.engine.retain.fact_extraction import extract_facts_from_text
# =============================================================================
# DIMENSION PRESERVATION TESTS
@@ -432,6 +432,7 @@ with a concert surrounded by music, joy and the warm summer breeze.
assert birthday_fact is not None, "Should extract fact about birthday celebration"
fact_date_str = birthday_fact.occurred_start
assert fact_date_str is not None, "occurred_start should not be None for temporal events"
if 'T' in fact_date_str:
fact_date = datetime.fromisoformat(fact_date_str.replace('Z', '+00:00'))
@@ -497,7 +498,7 @@ Yesterday I went for a morning jog for the first time in a nearby park.
async def test_extract_facts_with_relative_dates(self):
"""Test that relative dates are converted to absolute dates."""
reference_date = datetime(2024, 3, 20, 14, 0, 0, tzinfo=timezone.utc)
reference_date = datetime(2024, 3, 20, 14, 0, 0, tzinfo=UTC)
llm_config = LLMConfig.for_memory()
text = """
@@ -531,7 +532,7 @@ Yesterday I went for a morning jog for the first time in a nearby park.
async def test_extract_facts_with_no_temporal_info(self):
"""Test that facts without temporal info are still extracted."""
reference_date = datetime(2024, 3, 20, 14, 0, 0, tzinfo=timezone.utc)
reference_date = datetime(2024, 3, 20, 14, 0, 0, tzinfo=UTC)
llm_config = LLMConfig.for_memory()
text = "Alice works at Google. She loves Python programming."
@@ -555,7 +556,7 @@ Yesterday I went for a morning jog for the first time in a nearby park.
async def test_extract_facts_with_absolute_dates(self):
"""Test that absolute dates in text are preserved."""
reference_date = datetime(2024, 3, 20, 14, 0, 0, tzinfo=timezone.utc)
reference_date = datetime(2024, 3, 20, 14, 0, 0, tzinfo=UTC)
llm_config = LLMConfig.for_memory()
text = """
@@ -1047,4 +1048,4 @@ class TestDispositionInference:
assert "texas" in background.lower()
# Higher skepticism expected from "very skeptical of people"
assert disposition["skepticism"] >= 3
assert disposition["skepticism"] >= 3
+37
View File
@@ -0,0 +1,37 @@
import js from "@eslint/js";
import tseslint from "typescript-eslint";
import reactPlugin from "eslint-plugin-react";
import reactHooksPlugin from "eslint-plugin-react-hooks";
export default [
js.configs.recommended,
...tseslint.configs.recommended,
{
files: ["**/*.{ts,tsx}"],
plugins: {
react: reactPlugin,
"react-hooks": reactHooksPlugin,
},
languageOptions: {
parserOptions: {
ecmaFeatures: {
jsx: true,
},
},
},
rules: {
"@typescript-eslint/no-unused-vars": "warn",
"@typescript-eslint/no-explicit-any": "warn",
"react/react-in-jsx-scope": "off",
"no-case-declarations": "off",
},
settings: {
react: {
version: "detect",
},
},
},
{
ignores: [".next/", "node_modules/"],
},
];
+6 -1
View File
@@ -50,6 +50,11 @@
"typescript": "^5.9.3"
},
"devDependencies": {
"prettier": "^3.7.4"
"@eslint/eslintrc": "^3.3.3",
"@eslint/js": "^9.39.2",
"eslint-plugin-react": "^7.37.5",
"eslint-plugin-react-hooks": "^7.0.1",
"prettier": "^3.7.4",
"typescript-eslint": "^8.50.0"
}
}
@@ -1,16 +1,13 @@
import { NextResponse } from 'next/server';
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
import { NextResponse } from "next/server";
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
export async function GET() {
try {
const response = await sdk.listBanks({ client: lowLevelClient });
return NextResponse.json(response.data, { status: 200 });
} catch (error) {
console.error('Error fetching banks:', error);
return NextResponse.json(
{ error: 'Failed to fetch banks' },
{ status: 500 }
);
console.error("Error fetching banks:", error);
return NextResponse.json({ error: "Failed to fetch banks" }, { status: 500 });
}
}
@@ -20,10 +17,7 @@ export async function POST(request: Request) {
const { bank_id } = body;
if (!bank_id) {
return NextResponse.json(
{ error: 'bank_id is required' },
{ status: 400 }
);
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
}
const response = await sdk.createOrUpdateBank({
@@ -34,10 +28,7 @@ export async function POST(request: Request) {
return NextResponse.json(response.data, { status: 201 });
} catch (error) {
console.error('Error creating bank:', error);
return NextResponse.json(
{ error: 'Failed to create bank' },
{ status: 500 }
);
console.error("Error creating bank:", error);
return NextResponse.json({ error: "Failed to create bank" }, { status: 500 });
}
}
@@ -1,5 +1,5 @@
import { NextRequest, NextResponse } from 'next/server';
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
import { NextRequest, NextResponse } from "next/server";
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
export async function GET(
request: NextRequest,
@@ -10,15 +10,12 @@ export async function GET(
const response = await sdk.getChunk({
client: lowLevelClient,
path: { chunk_id: chunkId }
path: { chunk_id: chunkId },
});
return NextResponse.json(response.data, { status: 200 });
} catch (error) {
console.error('Error fetching chunk:', error);
return NextResponse.json(
{ error: 'Failed to fetch chunk' },
{ status: 500 }
);
console.error("Error fetching chunk:", error);
return NextResponse.json({ error: "Failed to fetch chunk" }, { status: 500 });
}
}
@@ -1,5 +1,5 @@
import { NextRequest, NextResponse } from 'next/server';
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
import { NextRequest, NextResponse } from "next/server";
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
export async function GET(
request: NextRequest,
@@ -8,26 +8,20 @@ export async function GET(
try {
const { documentId } = await params;
const searchParams = request.nextUrl.searchParams;
const bankId = searchParams.get('bank_id');
const bankId = searchParams.get("bank_id");
if (!bankId) {
return NextResponse.json(
{ error: 'bank_id is required' },
{ status: 400 }
);
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
}
const response = await sdk.getDocument({
client: lowLevelClient,
path: { bank_id: bankId, document_id: documentId }
path: { bank_id: bankId, document_id: documentId },
});
return NextResponse.json(response.data, { status: 200 });
} catch (error) {
console.error('Error fetching document:', error);
return NextResponse.json(
{ error: 'Failed to fetch document' },
{ status: 500 }
);
console.error("Error fetching document:", error);
return NextResponse.json({ error: "Failed to fetch document" }, { status: 500 });
}
}
@@ -1,33 +1,27 @@
import { NextRequest, NextResponse } from 'next/server';
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
import { NextRequest, NextResponse } from "next/server";
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
export async function GET(request: NextRequest) {
try {
const searchParams = request.nextUrl.searchParams;
const bankId = searchParams.get('bank_id');
const bankId = searchParams.get("bank_id");
if (!bankId) {
return NextResponse.json(
{ error: 'bank_id is required' },
{ status: 400 }
);
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
}
const limit = searchParams.get('limit') ? Number(searchParams.get('limit')) : undefined;
const offset = searchParams.get('offset') ? Number(searchParams.get('offset')) : undefined;
const limit = searchParams.get("limit") ? Number(searchParams.get("limit")) : undefined;
const offset = searchParams.get("offset") ? Number(searchParams.get("offset")) : undefined;
const response = await sdk.listDocuments({
client: lowLevelClient,
path: { bank_id: bankId },
query: { limit, offset }
query: { limit, offset },
});
return NextResponse.json(response.data, { status: 200 });
} catch (error) {
console.error('Error fetching documents:', error);
return NextResponse.json(
{ error: 'Failed to fetch documents' },
{ status: 500 }
);
console.error("Error fetching documents:", error);
return NextResponse.json({ error: "Failed to fetch documents" }, { status: 500 });
}
}
@@ -1,5 +1,5 @@
import { NextRequest, NextResponse } from 'next/server';
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
import { NextRequest, NextResponse } from "next/server";
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
export async function POST(
request: NextRequest,
@@ -8,13 +8,10 @@ export async function POST(
try {
const { entityId } = await params;
const searchParams = request.nextUrl.searchParams;
const bankId = searchParams.get('bank_id');
const bankId = searchParams.get("bank_id");
if (!bankId) {
return NextResponse.json(
{ error: 'bank_id is required' },
{ status: 400 }
);
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
}
const decodedEntityId = decodeURIComponent(entityId);
@@ -23,15 +20,15 @@ export async function POST(
client: lowLevelClient,
path: {
bank_id: bankId,
entity_id: decodedEntityId
}
entity_id: decodedEntityId,
},
});
return NextResponse.json(response.data, { status: 200 });
} catch (error) {
console.error('Error regenerating entity observations:', error);
console.error("Error regenerating entity observations:", error);
return NextResponse.json(
{ error: 'Failed to regenerate entity observations' },
{ error: "Failed to regenerate entity observations" },
{ status: 500 }
);
}
@@ -1,5 +1,5 @@
import { NextRequest, NextResponse } from 'next/server';
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
import { NextRequest, NextResponse } from "next/server";
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
export async function GET(
request: NextRequest,
@@ -8,13 +8,10 @@ export async function GET(
try {
const { entityId } = await params;
const searchParams = request.nextUrl.searchParams;
const bankId = searchParams.get('bank_id');
const bankId = searchParams.get("bank_id");
if (!bankId) {
return NextResponse.json(
{ error: 'bank_id is required' },
{ status: 400 }
);
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
}
// Decode URL-encoded entityId in case it contains special chars
@@ -24,23 +21,17 @@ export async function GET(
client: lowLevelClient,
path: {
bank_id: bankId,
entity_id: decodedEntityId
}
entity_id: decodedEntityId,
},
});
if (response.error) {
return NextResponse.json(
{ error: response.error },
{ status: 500 }
);
return NextResponse.json({ error: response.error }, { status: 500 });
}
return NextResponse.json(response.data, { status: 200 });
} catch (error) {
console.error('Error getting entity:', error);
return NextResponse.json(
{ error: 'Failed to get entity' },
{ status: 500 }
);
console.error("Error getting entity:", error);
return NextResponse.json({ error: "Failed to get entity" }, { status: 500 });
}
}
@@ -1,39 +1,30 @@
import { NextRequest, NextResponse } from 'next/server';
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
import { NextRequest, NextResponse } from "next/server";
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
export async function GET(request: NextRequest) {
try {
const searchParams = request.nextUrl.searchParams;
const bankId = searchParams.get('bank_id');
const bankId = searchParams.get("bank_id");
if (!bankId) {
return NextResponse.json(
{ error: 'bank_id is required' },
{ status: 400 }
);
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
}
const limit = searchParams.get('limit') ? Number(searchParams.get('limit')) : undefined;
const limit = searchParams.get("limit") ? Number(searchParams.get("limit")) : undefined;
const response = await sdk.listEntities({
client: lowLevelClient,
path: { bank_id: bankId },
query: { limit }
query: { limit },
});
if (response.error) {
return NextResponse.json(
{ error: response.error },
{ status: 500 }
);
return NextResponse.json({ error: response.error }, { status: 500 });
}
return NextResponse.json(response.data, { status: 200 });
} catch (error) {
console.error('Error listing entities:', error);
return NextResponse.json(
{ error: 'Failed to list entities' },
{ status: 500 }
);
console.error("Error listing entities:", error);
return NextResponse.json({ error: "Failed to list entities" }, { status: 500 });
}
}
@@ -1,35 +1,29 @@
import { NextRequest, NextResponse } from 'next/server';
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
import { NextRequest, NextResponse } from "next/server";
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
export async function GET(request: NextRequest) {
try {
const searchParams = request.nextUrl.searchParams;
const bankId = searchParams.get('bank_id') || searchParams.get('agent_id');
const bankId = searchParams.get("bank_id") || searchParams.get("agent_id");
if (!bankId) {
return NextResponse.json(
{ error: 'bank_id is required' },
{ status: 400 }
);
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
}
// Get optional query parameters
const type = searchParams.get('type') || searchParams.get('fact_type') || undefined;
const type = searchParams.get("type") || searchParams.get("fact_type") || undefined;
const response = await sdk.getGraph({
client: lowLevelClient,
path: { bank_id: bankId },
query: {
type: type
}
type: type,
},
});
return NextResponse.json(response.data, { status: 200 });
} catch (error) {
console.error('Error fetching graph data:', error);
return NextResponse.json(
{ error: 'Failed to fetch graph data' },
{ status: 500 }
);
console.error("Error fetching graph data:", error);
return NextResponse.json({ error: "Failed to fetch graph data" }, { status: 500 });
}
}
@@ -1,37 +1,31 @@
import { NextRequest, NextResponse } from 'next/server';
import { hindsightClient, sdk, lowLevelClient } from '@/lib/hindsight-client';
import { NextRequest, NextResponse } from "next/server";
import { hindsightClient, sdk, lowLevelClient } from "@/lib/hindsight-client";
export async function GET(request: NextRequest) {
try {
const searchParams = request.nextUrl.searchParams;
const bankId = searchParams.get('bank_id') || searchParams.get('agent_id');
const bankId = searchParams.get("bank_id") || searchParams.get("agent_id");
if (!bankId) {
return NextResponse.json(
{ error: 'bank_id is required' },
{ status: 400 }
);
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
}
const limit = searchParams.get('limit') ? Number(searchParams.get('limit')) : undefined;
const offset = searchParams.get('offset') ? Number(searchParams.get('offset')) : undefined;
const type = searchParams.get('type') || searchParams.get('fact_type') || undefined;
const q = searchParams.get('q') || undefined;
const limit = searchParams.get("limit") ? Number(searchParams.get("limit")) : undefined;
const offset = searchParams.get("offset") ? Number(searchParams.get("offset")) : undefined;
const type = searchParams.get("type") || searchParams.get("fact_type") || undefined;
const q = searchParams.get("q") || undefined;
const response = await hindsightClient.listMemories(bankId, {
limit,
offset,
type,
q
q,
});
return NextResponse.json(response, { status: 200 });
} catch (error) {
console.error('Error listing memory units:', error);
return NextResponse.json(
{ error: 'Failed to list memory units' },
{ status: 500 }
);
console.error("Error listing memory units:", error);
return NextResponse.json({ error: "Failed to list memory units" }, { status: 500 });
}
}
@@ -39,7 +33,10 @@ export async function GET(request: NextRequest) {
// Use clearBankMemories to delete all memories for a bank instead
export async function DELETE(request: NextRequest) {
return NextResponse.json(
{ error: 'Individual memory unit deletion is not yet supported. Use clear all memories instead.' },
{
error:
"Individual memory unit deletion is not yet supported. Use clear all memories instead.",
},
{ status: 501 } // Not Implemented
);
}
@@ -1,5 +1,5 @@
import { NextRequest, NextResponse } from 'next/server';
import { hindsightClient } from '@/lib/hindsight-client';
import { NextRequest, NextResponse } from "next/server";
import { hindsightClient } from "@/lib/hindsight-client";
export async function POST(request: NextRequest) {
try {
@@ -7,10 +7,7 @@ export async function POST(request: NextRequest) {
const bankId = body.bank_id || body.agent_id;
if (!bankId) {
return NextResponse.json(
{ error: 'bank_id is required' },
{ status: 400 }
);
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
}
const { items, document_id } = body;
@@ -19,10 +16,7 @@ export async function POST(request: NextRequest) {
return NextResponse.json(response, { status: 200 });
} catch (error) {
console.error('Error batch retain:', error);
return NextResponse.json(
{ error: 'Failed to batch retain' },
{ status: 500 }
);
console.error("Error batch retain:", error);
return NextResponse.json({ error: "Failed to batch retain" }, { status: 500 });
}
}
@@ -1,5 +1,5 @@
import { NextRequest, NextResponse } from 'next/server';
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
import { NextRequest, NextResponse } from "next/server";
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
export async function POST(request: NextRequest) {
try {
@@ -7,10 +7,7 @@ export async function POST(request: NextRequest) {
const bankId = body.bank_id || body.agent_id;
if (!bankId) {
return NextResponse.json(
{ error: 'bank_id is required' },
{ status: 400 }
);
return NextResponse.json({ error: "bank_id is required" }, { status: 400 });
}
const { items } = body;
@@ -18,15 +15,12 @@ export async function POST(request: NextRequest) {
const response = await sdk.retainMemories({
client: lowLevelClient,
path: { bank_id: bankId },
body: { items, async: true }
body: { items, async: true },
});
return NextResponse.json(response.data, { status: 200 });
} catch (error) {
console.error('Error batch retain async:', error);
return NextResponse.json(
{ error: 'Failed to batch retain async' },
{ status: 500 }
);
console.error("Error batch retain async:", error);
return NextResponse.json({ error: "Failed to batch retain async" }, { status: 500 });
}
}
@@ -1,5 +1,5 @@
import { NextRequest, NextResponse } from 'next/server';
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
import { NextRequest, NextResponse } from "next/server";
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
export async function GET(
request: NextRequest,
@@ -9,15 +9,12 @@ export async function GET(
const { agentId } = await params;
const response = await sdk.listOperations({
client: lowLevelClient,
path: { bank_id: agentId }
path: { bank_id: agentId },
});
return NextResponse.json(response.data || {}, { status: 200 });
} catch (error) {
console.error('Error fetching operations:', error);
return NextResponse.json(
{ error: 'Failed to fetch operations' },
{ status: 500 }
);
console.error("Error fetching operations:", error);
return NextResponse.json({ error: "Failed to fetch operations" }, { status: 500 });
}
}
@@ -28,26 +25,20 @@ export async function DELETE(
try {
const { agentId } = await params;
const searchParams = request.nextUrl.searchParams;
const operationId = searchParams.get('operation_id');
const operationId = searchParams.get("operation_id");
if (!operationId) {
return NextResponse.json(
{ error: 'operation_id is required' },
{ status: 400 }
);
return NextResponse.json({ error: "operation_id is required" }, { status: 400 });
}
const response = await sdk.cancelOperation({
client: lowLevelClient,
path: { bank_id: agentId, operation_id: operationId }
path: { bank_id: agentId, operation_id: operationId },
});
return NextResponse.json(response.data || {}, { status: 200 });
} catch (error) {
console.error('Error canceling operation:', error);
return NextResponse.json(
{ error: 'Failed to cancel operation' },
{ status: 500 }
);
console.error("Error canceling operation:", error);
return NextResponse.json({ error: "Failed to cancel operation" }, { status: 500 });
}
}
@@ -1,5 +1,5 @@
import { NextRequest, NextResponse } from 'next/server';
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
import { NextRequest, NextResponse } from "next/server";
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
export async function GET(
request: NextRequest,
@@ -9,15 +9,12 @@ export async function GET(
const { bankId } = await params;
const response = await sdk.getBankProfile({
client: lowLevelClient,
path: { bank_id: bankId }
path: { bank_id: bankId },
});
return NextResponse.json(response.data, { status: 200 });
} catch (error) {
console.error('Error fetching bank profile:', error);
return NextResponse.json(
{ error: 'Failed to fetch bank profile' },
{ status: 500 }
);
console.error("Error fetching bank profile:", error);
return NextResponse.json({ error: "Failed to fetch bank profile" }, { status: 500 });
}
}
@@ -32,14 +29,11 @@ export async function PUT(
const response = await sdk.createOrUpdateBank({
client: lowLevelClient,
path: { bank_id: bankId },
body: body
body: body,
});
return NextResponse.json(response.data, { status: 200 });
} catch (error) {
console.error('Error updating bank profile:', error);
return NextResponse.json(
{ error: 'Failed to update bank profile' },
{ status: 500 }
);
console.error("Error updating bank profile:", error);
return NextResponse.json({ error: "Failed to update bank profile" }, { status: 500 });
}
}
@@ -1,14 +1,22 @@
import { NextRequest, NextResponse } from 'next/server';
import { lowLevelClient, sdk } from '@/lib/hindsight-client';
import { NextRequest, NextResponse } from "next/server";
import { lowLevelClient, sdk } from "@/lib/hindsight-client";
export async function POST(request: NextRequest) {
try {
const body = await request.json();
const bankId = body.bank_id || body.agent_id || 'default';
const bankId = body.bank_id || body.agent_id || "default";
const { query, types, fact_type, max_tokens, trace, budget, include, query_timestamp } = body;
console.log('[Recall API] Request:', { bankId, query, types: types || fact_type, max_tokens, trace, budget, query_timestamp });
console.log('[Recall API] Include options:', JSON.stringify(include, null, 2));
console.log("[Recall API] Request:", {
bankId,
query,
types: types || fact_type,
max_tokens,
trace,
budget,
query_timestamp,
});
console.log("[Recall API] Include options:", JSON.stringify(include, null, 2));
const response = await sdk.recallMemories({
client: lowLevelClient,
@@ -18,18 +26,18 @@ export async function POST(request: NextRequest) {
types: types || fact_type,
max_tokens,
trace,
budget: budget || 'mid',
budget: budget || "mid",
include,
query_timestamp,
},
});
if (!response.data) {
console.error('[Recall API] No data in response', { response, error: response.error });
throw new Error(`API returned no data: ${JSON.stringify(response.error || 'Unknown error')}`);
console.error("[Recall API] No data in response", { response, error: response.error });
throw new Error(`API returned no data: ${JSON.stringify(response.error || "Unknown error")}`);
}
console.log('[Recall API] Response structure:', {
console.log("[Recall API] Response structure:", {
hasResults: !!response.data?.results,
resultsCount: response.data?.results?.length,
hasTrace: !!response.data?.trace,
@@ -52,10 +60,7 @@ export async function POST(request: NextRequest) {
return NextResponse.json(jsonResponse, { status: 200 });
} catch (error) {
console.error('Error recalling:', error);
return NextResponse.json(
{ error: 'Failed to recall' },
{ status: 500 }
);
console.error("Error recalling:", error);
return NextResponse.json({ error: "Failed to recall" }, { status: 500 });
}
}
@@ -1,37 +1,34 @@
import { NextRequest, NextResponse } from 'next/server';
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
import { NextRequest, NextResponse } from "next/server";
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
export async function POST(request: NextRequest) {
try {
const body = await request.json();
const bankId = body.bank_id || body.agent_id || 'default';
const bankId = body.bank_id || body.agent_id || "default";
const { query, context, budget, thinking_budget, include_facts } = body;
const requestBody: any = {
query,
budget: budget || (thinking_budget ? 'mid' : 'low'),
context: context || undefined
budget: budget || (thinking_budget ? "mid" : "low"),
context: context || undefined,
};
// Add include options if specified
if (include_facts) {
requestBody.include = {
facts: {}
facts: {},
};
}
const response = await sdk.reflect({
client: lowLevelClient,
path: { bank_id: bankId },
body: requestBody
body: requestBody,
});
return NextResponse.json(response.data, { status: 200 });
} catch (error) {
console.error('Error reflecting:', error);
return NextResponse.json(
{ error: 'Failed to reflect' },
{ status: 500 }
);
console.error("Error reflecting:", error);
return NextResponse.json({ error: "Failed to reflect" }, { status: 500 });
}
}
@@ -1,5 +1,5 @@
import { NextRequest, NextResponse } from 'next/server';
import { sdk, lowLevelClient } from '@/lib/hindsight-client';
import { NextRequest, NextResponse } from "next/server";
import { sdk, lowLevelClient } from "@/lib/hindsight-client";
export async function GET(
request: NextRequest,
@@ -9,14 +9,11 @@ export async function GET(
const { agentId } = await params;
const response = await sdk.getAgentStats({
client: lowLevelClient,
path: { bank_id: agentId }
path: { bank_id: agentId },
});
return NextResponse.json(response.data, { status: 200 });
} catch (error) {
console.error('Error fetching stats:', error);
return NextResponse.json(
{ error: 'Failed to fetch stats' },
{ status: 500 }
);
console.error("Error fetching stats:", error);
return NextResponse.json({ error: "Failed to fetch stats" }, { status: 500 });
}
}
@@ -1,19 +1,19 @@
'use client';
"use client";
import { useParams, useRouter, useSearchParams } from 'next/navigation';
import { useEffect } from 'react';
import { BankSelector } from '@/components/bank-selector';
import { Sidebar } from '@/components/sidebar';
import { DataView } from '@/components/data-view';
import { DocumentsView } from '@/components/documents-view';
import { EntitiesView } from '@/components/entities-view';
import { ThinkView } from '@/components/think-view';
import { SearchDebugView } from '@/components/search-debug-view';
import { BankProfileView } from '@/components/bank-profile-view';
import { useBank } from '@/lib/bank-context';
import { useParams, useRouter, useSearchParams } from "next/navigation";
import { useEffect } from "react";
import { BankSelector } from "@/components/bank-selector";
import { Sidebar } from "@/components/sidebar";
import { DataView } from "@/components/data-view";
import { DocumentsView } from "@/components/documents-view";
import { EntitiesView } from "@/components/entities-view";
import { ThinkView } from "@/components/think-view";
import { SearchDebugView } from "@/components/search-debug-view";
import { BankProfileView } from "@/components/bank-profile-view";
import { useBank } from "@/lib/bank-context";
type NavItem = 'recall' | 'reflect' | 'data' | 'documents' | 'entities' | 'profile';
type DataSubTab = 'world' | 'experience' | 'opinion';
type NavItem = "recall" | "reflect" | "data" | "documents" | "entities" | "profile";
type DataSubTab = "world" | "experience" | "opinion";
export default function BankPage() {
const params = useParams();
@@ -22,8 +22,8 @@ export default function BankPage() {
const { currentBank, setCurrentBank } = useBank();
const bankId = params.bankId as string;
const view = (searchParams.get('view') || 'profile') as NavItem;
const subTab = (searchParams.get('subTab') || 'world') as DataSubTab;
const view = (searchParams.get("view") || "profile") as NavItem;
const subTab = (searchParams.get("subTab") || "world") as DataSubTab;
// Sync URL bank with context
useEffect(() => {
@@ -50,18 +50,19 @@ export default function BankPage() {
<main className="flex-1 overflow-y-auto">
<div className="p-6">
{/* Profile Tab */}
{view === 'profile' && (
{view === "profile" && (
<div>
<h1 className="text-3xl font-bold mb-2 text-foreground">Bank Profile</h1>
<p className="text-muted-foreground mb-6">
View and edit the memory bank profile, disposition traits, and background information.
View and edit the memory bank profile, disposition traits, and background
information.
</p>
<BankProfileView />
</div>
)}
{/* Recall Tab */}
{view === 'recall' && (
{view === "recall" && (
<div>
<h1 className="text-3xl font-bold mb-2 text-foreground">Recall Analyzer</h1>
<p className="text-muted-foreground mb-6">
@@ -72,7 +73,7 @@ export default function BankPage() {
)}
{/* Reflect Tab */}
{view === 'reflect' && (
{view === "reflect" && (
<div>
<h1 className="text-3xl font-bold mb-2 text-foreground">Reflect</h1>
<p className="text-muted-foreground mb-6">
@@ -83,7 +84,7 @@ export default function BankPage() {
)}
{/* Data/Memories Tab */}
{view === 'data' && (
{view === "data" && (
<div>
<h1 className="text-3xl font-bold mb-2 text-foreground">Memories</h1>
<p className="text-muted-foreground mb-6">
@@ -93,41 +94,41 @@ export default function BankPage() {
<div className="mb-6 border-b border-border">
<div className="flex gap-1">
<button
onClick={() => handleDataSubTabChange('world')}
onClick={() => handleDataSubTabChange("world")}
className={`px-6 py-3 font-semibold text-sm transition-all relative ${
subTab === 'world'
? 'text-primary'
: 'text-muted-foreground hover:text-foreground'
subTab === "world"
? "text-primary"
: "text-muted-foreground hover:text-foreground"
}`}
>
World Facts
{subTab === 'world' && (
{subTab === "world" && (
<div className="absolute bottom-0 left-0 right-0 h-0.5 bg-primary" />
)}
</button>
<button
onClick={() => handleDataSubTabChange('experience')}
onClick={() => handleDataSubTabChange("experience")}
className={`px-6 py-3 font-semibold text-sm transition-all relative ${
subTab === 'experience'
? 'text-primary'
: 'text-muted-foreground hover:text-foreground'
subTab === "experience"
? "text-primary"
: "text-muted-foreground hover:text-foreground"
}`}
>
Experience
{subTab === 'experience' && (
{subTab === "experience" && (
<div className="absolute bottom-0 left-0 right-0 h-0.5 bg-primary" />
)}
</button>
<button
onClick={() => handleDataSubTabChange('opinion')}
onClick={() => handleDataSubTabChange("opinion")}
className={`px-6 py-3 font-semibold text-sm transition-all relative ${
subTab === 'opinion'
? 'text-primary'
: 'text-muted-foreground hover:text-foreground'
subTab === "opinion"
? "text-primary"
: "text-muted-foreground hover:text-foreground"
}`}
>
Opinions
{subTab === 'opinion' && (
{subTab === "opinion" && (
<div className="absolute bottom-0 left-0 right-0 h-0.5 bg-primary" />
)}
</button>
@@ -135,15 +136,15 @@ export default function BankPage() {
</div>
<div>
{subTab === 'world' && <DataView key="world" factType="world" />}
{subTab === 'experience' && <DataView key="experience" factType="experience" />}
{subTab === 'opinion' && <DataView key="opinion" factType="opinion" />}
{subTab === "world" && <DataView key="world" factType="world" />}
{subTab === "experience" && <DataView key="experience" factType="experience" />}
{subTab === "opinion" && <DataView key="opinion" factType="opinion" />}
</div>
</div>
)}
{/* Documents Tab */}
{view === 'documents' && (
{view === "documents" && (
<div>
<h1 className="text-3xl font-bold mb-2 text-foreground">Documents</h1>
<p className="text-muted-foreground mb-6">
@@ -154,7 +155,7 @@ export default function BankPage() {
)}
{/* Entities Tab */}
{view === 'entities' && (
{view === "entities" && (
<div>
<h1 className="text-3xl font-bold mb-2 text-foreground">Entities</h1>
<p className="text-muted-foreground mb-6">
@@ -1,9 +1,9 @@
'use client';
"use client";
import { useEffect } from 'react';
import { useRouter } from 'next/navigation';
import { BankSelector } from '@/components/bank-selector';
import { useBank } from '@/lib/bank-context';
import { useEffect } from "react";
import { useRouter } from "next/navigation";
import { BankSelector } from "@/components/bank-selector";
import { useBank } from "@/lib/bank-context";
export default function DashboardPage() {
const router = useRouter();
+1 -3
View File
@@ -20,9 +20,7 @@ export default function RootLayout({
<html lang="en" suppressHydrationWarning>
<body className="bg-background text-foreground">
<ThemeProvider>
<BankProvider>
{children}
</BankProvider>
<BankProvider>{children}</BankProvider>
</ThemeProvider>
</body>
</html>
+2 -2
View File
@@ -1,5 +1,5 @@
import { redirect } from 'next/navigation';
import { redirect } from "next/navigation";
export default function Home() {
redirect('/dashboard');
redirect("/dashboard");
}
@@ -1,35 +1,35 @@
'use client';
"use client";
import { useState } from 'react';
import { client } from '@/lib/api';
import { useBank } from '@/lib/bank-context';
import { Button } from '@/components/ui/button';
import { Input } from '@/components/ui/input';
import { Textarea } from '@/components/ui/textarea';
import { Checkbox } from '@/components/ui/checkbox';
import { useState } from "react";
import { client } from "@/lib/api";
import { useBank } from "@/lib/bank-context";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import { Textarea } from "@/components/ui/textarea";
import { Checkbox } from "@/components/ui/checkbox";
export function AddMemoryView() {
const { currentBank } = useBank();
const [content, setContent] = useState('');
const [context, setContext] = useState('');
const [eventDate, setEventDate] = useState('');
const [documentId, setDocumentId] = useState('');
const [content, setContent] = useState("");
const [context, setContext] = useState("");
const [eventDate, setEventDate] = useState("");
const [documentId, setDocumentId] = useState("");
const [async, setAsync] = useState(false);
const [loading, setLoading] = useState(false);
const [result, setResult] = useState<string | null>(null);
const clearForm = () => {
setContent('');
setContext('');
setEventDate('');
setDocumentId('');
setContent("");
setContext("");
setEventDate("");
setDocumentId("");
setAsync(false);
setResult(null);
};
const submitMemory = async () => {
if (!currentBank || !content) {
alert('Please enter content');
alert("Please enter content");
return;
}
@@ -49,10 +49,10 @@ export function AddMemoryView() {
});
setResult(data.message as string);
setContent('');
setContent("");
} catch (error) {
console.error('Error submitting memory:', error);
setResult('Error: ' + (error as Error).message);
console.error("Error submitting memory:", error);
setResult("Error: " + (error as Error).message);
} finally {
setLoading(false);
}
@@ -106,7 +106,8 @@ export function AddMemoryView() {
placeholder="Optional document identifier (automatically upserts if document exists)..."
/>
<small className="text-muted-foreground text-xs mt-1 block">
Note: If a document with this ID already exists, it will be automatically replaced with the new content.
Note: If a document with this ID already exists, it will be automatically replaced
with the new content.
</small>
</div>
@@ -124,23 +125,19 @@ export function AddMemoryView() {
</div>
<div className="flex gap-2.5">
<Button
onClick={submitMemory}
disabled={loading}
>
{loading ? 'Retaining...' : 'Retain Memory'}
<Button onClick={submitMemory} disabled={loading}>
{loading ? "Retaining..." : "Retain Memory"}
</Button>
<Button
onClick={clearForm}
variant="secondary"
>
<Button onClick={clearForm} variant="secondary">
Clear Form
</Button>
</div>
</div>
{result && (
<div className={`mt-5 p-5 rounded-lg border-2 ${result.startsWith('Error') ? 'bg-destructive/10 border-destructive text-destructive' : 'bg-primary/10 border-primary text-primary'}`}>
<div
className={`mt-5 p-5 rounded-lg border-2 ${result.startsWith("Error") ? "bg-destructive/10 border-destructive text-destructive" : "bg-primary/10 border-primary text-primary"}`}
>
<div className="font-semibold">{result}</div>
</div>
)}
@@ -1,13 +1,32 @@
'use client';
"use client";
import { useState, useEffect } from 'react';
import { client } from '@/lib/api';
import { useBank } from '@/lib/bank-context';
import { Button } from '@/components/ui/button';
import { Textarea } from '@/components/ui/textarea';
import { Card, CardContent, CardDescription, CardHeader, CardTitle } from '@/components/ui/card';
import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from '@/components/ui/table';
import { RefreshCw, Save, Brain, FileText, Clock, AlertCircle, CheckCircle, Database, Link2, FolderOpen, Activity } from 'lucide-react';
import { useState, useEffect } from "react";
import { client } from "@/lib/api";
import { useBank } from "@/lib/bank-context";
import { Button } from "@/components/ui/button";
import { Textarea } from "@/components/ui/textarea";
import { Card, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card";
import {
Table,
TableBody,
TableCell,
TableHead,
TableHeader,
TableRow,
} from "@/components/ui/table";
import {
RefreshCw,
Save,
Brain,
FileText,
Clock,
AlertCircle,
CheckCircle,
Database,
Link2,
FolderOpen,
Activity,
} from "lucide-react";
interface DispositionTraits {
skepticism: number;
@@ -51,31 +70,39 @@ interface Operation {
error_message?: string;
}
const TRAIT_LABELS: Record<keyof DispositionTraits, { label: string; shortLabel: string; description: string; lowLabel: string; highLabel: string }> = {
const TRAIT_LABELS: Record<
keyof DispositionTraits,
{ label: string; shortLabel: string; description: string; lowLabel: string; highLabel: string }
> = {
skepticism: {
label: 'Skepticism',
shortLabel: 'S',
description: 'How skeptical vs trusting when forming opinions',
lowLabel: 'Trusting',
highLabel: 'Skeptical'
label: "Skepticism",
shortLabel: "S",
description: "How skeptical vs trusting when forming opinions",
lowLabel: "Trusting",
highLabel: "Skeptical",
},
literalism: {
label: 'Literalism',
shortLabel: 'L',
description: 'How literally to interpret information when forming opinions',
lowLabel: 'Flexible',
highLabel: 'Literal'
label: "Literalism",
shortLabel: "L",
description: "How literally to interpret information when forming opinions",
lowLabel: "Flexible",
highLabel: "Literal",
},
empathy: {
label: 'Empathy',
shortLabel: 'E',
description: 'How much to consider emotional context when forming opinions',
lowLabel: 'Detached',
highLabel: 'Empathetic'
}
label: "Empathy",
shortLabel: "E",
description: "How much to consider emotional context when forming opinions",
lowLabel: "Detached",
highLabel: "Empathetic",
},
};
function DispositionEditor({ disposition, editMode, editDisposition, onEditChange }: {
function DispositionEditor({
disposition,
editMode,
editDisposition,
onEditChange,
}: {
disposition: DispositionTraits;
editMode: boolean;
editDisposition: DispositionTraits;
@@ -89,7 +116,9 @@ function DispositionEditor({ disposition, editMode, editDisposition, onEditChang
<div key={trait} className="space-y-2">
<div className="flex justify-between items-center">
<div>
<label className="text-sm font-medium text-foreground">{TRAIT_LABELS[trait].label}</label>
<label className="text-sm font-medium text-foreground">
{TRAIT_LABELS[trait].label}
</label>
<p className="text-xs text-muted-foreground">{TRAIT_LABELS[trait].description}</p>
</div>
<span className="text-sm font-bold text-primary">{data[trait]}/5</span>
@@ -138,11 +167,11 @@ export function BankProfileView() {
const [editMode, setEditMode] = useState(false);
// Edit state
const [editBackground, setEditBackground] = useState('');
const [editBackground, setEditBackground] = useState("");
const [editDisposition, setEditDisposition] = useState<DispositionTraits>({
skepticism: 3,
literalism: 3,
empathy: 3
empathy: 3,
});
const loadData = async () => {
@@ -153,7 +182,7 @@ export function BankProfileView() {
const [profileData, statsData, opsData] = await Promise.all([
client.getBankProfile(currentBank),
client.getBankStats(currentBank),
client.listOperations(currentBank)
client.listOperations(currentBank),
]);
setProfile(profileData);
setStats(statsData as BankStats);
@@ -163,8 +192,8 @@ export function BankProfileView() {
setEditBackground(profileData.background);
setEditDisposition(profileData.disposition);
} catch (error) {
console.error('Error loading bank profile:', error);
alert('Error loading bank profile: ' + (error as Error).message);
console.error("Error loading bank profile:", error);
alert("Error loading bank profile: " + (error as Error).message);
} finally {
setLoading(false);
}
@@ -177,13 +206,13 @@ export function BankProfileView() {
try {
await client.updateBankProfile(currentBank, {
background: editBackground,
disposition: editDisposition
disposition: editDisposition,
});
await loadData();
setEditMode(false);
} catch (error) {
console.error('Error saving bank profile:', error);
alert('Error saving bank profile: ' + (error as Error).message);
console.error("Error saving bank profile:", error);
alert("Error saving bank profile: " + (error as Error).message);
} finally {
setSaving(false);
}
@@ -211,7 +240,9 @@ export function BankProfileView() {
<Card>
<CardContent className="p-10 text-center">
<h3 className="text-xl font-semibold mb-2 text-card-foreground">No Bank Selected</h3>
<p className="text-muted-foreground">Please select a memory bank from the dropdown above to view its profile.</p>
<p className="text-muted-foreground">
Please select a memory bank from the dropdown above to view its profile.
</p>
</CardContent>
</Card>
);
@@ -315,11 +346,17 @@ export function BankProfileView() {
</CardContent>
</Card>
<Card className={`bg-gradient-to-br ${stats.pending_operations > 0 ? 'from-amber-500/10 to-amber-600/5 border-amber-500/20' : 'from-slate-500/10 to-slate-600/5 border-slate-500/20'}`}>
<Card
className={`bg-gradient-to-br ${stats.pending_operations > 0 ? "from-amber-500/10 to-amber-600/5 border-amber-500/20" : "from-slate-500/10 to-slate-600/5 border-slate-500/20"}`}
>
<CardContent className="p-4">
<div className="flex items-center gap-3">
<div className={`p-2 rounded-lg ${stats.pending_operations > 0 ? 'bg-amber-500/20' : 'bg-slate-500/20'}`}>
<Activity className={`w-5 h-5 ${stats.pending_operations > 0 ? 'text-amber-500 animate-pulse' : 'text-slate-500'}`} />
<div
className={`p-2 rounded-lg ${stats.pending_operations > 0 ? "bg-amber-500/20" : "bg-slate-500/20"}`}
>
<Activity
className={`w-5 h-5 ${stats.pending_operations > 0 ? "text-amber-500 animate-pulse" : "text-slate-500"}`}
/>
</div>
<div>
<p className="text-xs text-muted-foreground font-medium">Pending</p>
@@ -335,16 +372,28 @@ export function BankProfileView() {
{stats && (
<div className="grid grid-cols-3 gap-3">
<div className="bg-blue-500/10 border border-blue-500/20 rounded-xl p-4 text-center">
<p className="text-xs text-blue-600 dark:text-blue-400 font-semibold uppercase tracking-wide">World Facts</p>
<p className="text-2xl font-bold text-blue-600 dark:text-blue-400 mt-1">{stats.nodes_by_fact_type?.world || 0}</p>
<p className="text-xs text-blue-600 dark:text-blue-400 font-semibold uppercase tracking-wide">
World Facts
</p>
<p className="text-2xl font-bold text-blue-600 dark:text-blue-400 mt-1">
{stats.nodes_by_fact_type?.world || 0}
</p>
</div>
<div className="bg-purple-500/10 border border-purple-500/20 rounded-xl p-4 text-center">
<p className="text-xs text-purple-600 dark:text-purple-400 font-semibold uppercase tracking-wide">Experience</p>
<p className="text-2xl font-bold text-purple-600 dark:text-purple-400 mt-1">{stats.nodes_by_fact_type?.experience || 0}</p>
<p className="text-xs text-purple-600 dark:text-purple-400 font-semibold uppercase tracking-wide">
Experience
</p>
<p className="text-2xl font-bold text-purple-600 dark:text-purple-400 mt-1">
{stats.nodes_by_fact_type?.experience || 0}
</p>
</div>
<div className="bg-amber-500/10 border border-amber-500/20 rounded-xl p-4 text-center">
<p className="text-xs text-amber-600 dark:text-amber-400 font-semibold uppercase tracking-wide">Opinions</p>
<p className="text-2xl font-bold text-amber-600 dark:text-amber-400 mt-1">{stats.nodes_by_fact_type?.opinion || 0}</p>
<p className="text-xs text-amber-600 dark:text-amber-400 font-semibold uppercase tracking-wide">
Opinions
</p>
<p className="text-2xl font-bold text-amber-600 dark:text-amber-400 mt-1">
{stats.nodes_by_fact_type?.opinion || 0}
</p>
</div>
</div>
)}
@@ -365,7 +414,9 @@ export function BankProfileView() {
disposition={profile.disposition}
editMode={editMode}
editDisposition={editDisposition}
onEditChange={(trait, value) => setEditDisposition(prev => ({ ...prev, [trait]: value }))}
onEditChange={(trait, value) =>
setEditDisposition((prev) => ({ ...prev, [trait]: value }))
}
/>
)}
</CardContent>
@@ -392,7 +443,7 @@ export function BankProfileView() {
/>
) : (
<p className="text-sm text-foreground whitespace-pre-wrap leading-relaxed">
{profile?.background || 'No background information provided.'}
{profile?.background || "No background information provided."}
</p>
)}
</CardContent>
@@ -416,13 +467,17 @@ export function BankProfileView() {
{stats.pending_operations > 0 && (
<div className="flex items-center gap-1.5 px-2.5 py-1 rounded-full bg-amber-500/10 border border-amber-500/20">
<Clock className="w-3.5 h-3.5 text-amber-500" />
<span className="text-xs font-semibold text-amber-600 dark:text-amber-400">{stats.pending_operations} pending</span>
<span className="text-xs font-semibold text-amber-600 dark:text-amber-400">
{stats.pending_operations} pending
</span>
</div>
)}
{stats.failed_operations > 0 && (
<div className="flex items-center gap-1.5 px-2.5 py-1 rounded-full bg-red-500/10 border border-red-500/20">
<AlertCircle className="w-3.5 h-3.5 text-red-500" />
<span className="text-xs font-semibold text-red-600 dark:text-red-400">{stats.failed_operations} failed</span>
<span className="text-xs font-semibold text-red-600 dark:text-red-400">
{stats.failed_operations} failed
</span>
</div>
)}
</div>
@@ -445,32 +500,35 @@ export function BankProfileView() {
</TableHeader>
<TableBody>
{operations.slice(0, 10).map((op) => (
<TableRow key={op.id} className={op.status === 'failed' ? 'bg-red-500/5' : ''}>
<TableRow key={op.id} className={op.status === "failed" ? "bg-red-500/5" : ""}>
<TableCell className="font-mono text-xs text-muted-foreground">
{op.id.substring(0, 8)}
</TableCell>
<TableCell className="font-medium">{op.task_type}</TableCell>
<TableCell className="text-center">{op.items_count}</TableCell>
<TableCell className="font-mono text-xs text-muted-foreground">
{op.document_id ? op.document_id.substring(0, 12) + '...' : '—'}
{op.document_id ? op.document_id.substring(0, 12) + "..." : "—"}
</TableCell>
<TableCell className="text-sm text-muted-foreground">
{new Date(op.created_at).toLocaleString()}
</TableCell>
<TableCell>
{op.status === 'pending' && (
{op.status === "pending" && (
<span className="inline-flex items-center gap-1.5 px-2 py-0.5 rounded-full text-xs font-medium bg-amber-500/10 text-amber-600 dark:text-amber-400 border border-amber-500/20">
<Clock className="w-3 h-3" />
pending
</span>
)}
{op.status === 'failed' && (
<span className="inline-flex items-center gap-1.5 px-2 py-0.5 rounded-full text-xs font-medium bg-red-500/10 text-red-600 dark:text-red-400 border border-red-500/20" title={op.error_message}>
{op.status === "failed" && (
<span
className="inline-flex items-center gap-1.5 px-2 py-0.5 rounded-full text-xs font-medium bg-red-500/10 text-red-600 dark:text-red-400 border border-red-500/20"
title={op.error_message}
>
<AlertCircle className="w-3 h-3" />
failed
</span>
)}
{op.status === 'completed' && (
{op.status === "completed" && (
<span className="inline-flex items-center gap-1.5 px-2 py-0.5 rounded-full text-xs font-medium bg-emerald-500/10 text-emerald-600 dark:text-emerald-400 border border-emerald-500/20">
<CheckCircle className="w-3 h-3" />
done
@@ -483,7 +541,9 @@ export function BankProfileView() {
</Table>
</div>
) : (
<p className="text-muted-foreground text-center py-8 text-sm">No background operations</p>
<p className="text-muted-foreground text-center py-8 text-sm">
No background operations
</p>
)}
</CardContent>
</Card>
@@ -1,11 +1,11 @@
'use client';
"use client";
import * as React from 'react';
import { Suspense } from 'react';
import { useRouter, useSearchParams } from 'next/navigation';
import { useBank } from '@/lib/bank-context';
import { client } from '@/lib/api';
import { Button } from '@/components/ui/button';
import * as React from "react";
import { Suspense } from "react";
import { useRouter, useSearchParams } from "next/navigation";
import { useBank } from "@/lib/bank-context";
import { client } from "@/lib/api";
import { Button } from "@/components/ui/button";
import {
Command,
CommandEmpty,
@@ -13,26 +13,22 @@ import {
CommandInput,
CommandItem,
CommandList,
} from '@/components/ui/command';
import {
Popover,
PopoverContent,
PopoverTrigger,
} from '@/components/ui/popover';
} from "@/components/ui/command";
import { Popover, PopoverContent, PopoverTrigger } from "@/components/ui/popover";
import {
Dialog,
DialogContent,
DialogHeader,
DialogTitle,
DialogFooter,
} from '@/components/ui/dialog';
import { Input } from '@/components/ui/input';
import { Check, ChevronsUpDown, Plus, FileText, Moon, Sun, Github } from 'lucide-react';
import { useTheme } from '@/lib/theme-context';
import Image from 'next/image';
import { Textarea } from '@/components/ui/textarea';
import { Checkbox } from '@/components/ui/checkbox';
import { cn } from '@/lib/utils';
} from "@/components/ui/dialog";
import { Input } from "@/components/ui/input";
import { Check, ChevronsUpDown, Plus, FileText, Moon, Sun, Github } from "lucide-react";
import { useTheme } from "@/lib/theme-context";
import Image from "next/image";
import { Textarea } from "@/components/ui/textarea";
import { Checkbox } from "@/components/ui/checkbox";
import { cn } from "@/lib/utils";
function BankSelectorInner() {
const router = useRouter();
@@ -41,16 +37,16 @@ function BankSelectorInner() {
const { theme, toggleTheme } = useTheme();
const [open, setOpen] = React.useState(false);
const [createDialogOpen, setCreateDialogOpen] = React.useState(false);
const [newBankId, setNewBankId] = React.useState('');
const [newBankId, setNewBankId] = React.useState("");
const [isCreating, setIsCreating] = React.useState(false);
const [createError, setCreateError] = React.useState<string | null>(null);
// Document creation state
const [docDialogOpen, setDocDialogOpen] = React.useState(false);
const [docContent, setDocContent] = React.useState('');
const [docContext, setDocContext] = React.useState('');
const [docEventDate, setDocEventDate] = React.useState('');
const [docDocumentId, setDocDocumentId] = React.useState('');
const [docContent, setDocContent] = React.useState("");
const [docContext, setDocContext] = React.useState("");
const [docEventDate, setDocEventDate] = React.useState("");
const [docDocumentId, setDocDocumentId] = React.useState("");
const [docAsync, setDocAsync] = React.useState(false);
const [isCreatingDoc, setIsCreatingDoc] = React.useState(false);
const [docError, setDocError] = React.useState<string | null>(null);
@@ -69,12 +65,12 @@ function BankSelectorInner() {
await client.createBank(newBankId.trim());
await loadBanks();
setCreateDialogOpen(false);
setNewBankId('');
setNewBankId("");
// Navigate to the new bank
setCurrentBank(newBankId.trim());
router.push(`/banks/${newBankId.trim()}?view=data`);
} catch (error) {
setCreateError(error instanceof Error ? error.message : 'Failed to create bank');
setCreateError(error instanceof Error ? error.message : "Failed to create bank");
} finally {
setIsCreating(false);
}
@@ -106,16 +102,16 @@ function BankSelectorInner() {
// Reset form and close dialog
setDocDialogOpen(false);
setDocContent('');
setDocContext('');
setDocEventDate('');
setDocDocumentId('');
setDocContent("");
setDocContext("");
setDocEventDate("");
setDocDocumentId("");
setDocAsync(false);
// Navigate to documents view to see the new document
router.push(`/banks/${currentBank}?view=documents`);
} catch (error) {
setDocError(error instanceof Error ? error.message : 'Failed to create document');
setDocError(error instanceof Error ? error.message : "Failed to create document");
} finally {
setIsCreatingDoc(false);
}
@@ -125,7 +121,14 @@ function BankSelectorInner() {
<div className="bg-card text-card-foreground px-5 py-3 border-b-4 border-primary-gradient">
<div className="flex items-center gap-4 text-sm">
{/* Logo */}
<Image src="/logo.png" alt="Hindsight" width={40} height={40} className="h-10 w-auto" unoptimized />
<Image
src="/logo.png"
alt="Hindsight"
width={40}
height={40}
className="h-10 w-auto"
unoptimized
/>
{/* Separator */}
<div className="h-8 w-px bg-border" />
@@ -145,9 +148,7 @@ function BankSelectorInner() {
</PopoverTrigger>
<PopoverContent className="w-[250px] p-0">
<Command>
{sortedBanks.length > 0 && (
<CommandInput placeholder="Search memory banks..." />
)}
{sortedBanks.length > 0 && <CommandInput placeholder="Search memory banks..." />}
<CommandList>
<CommandEmpty>No memory banks yet.</CommandEmpty>
<CommandGroup>
@@ -159,9 +160,11 @@ function BankSelectorInner() {
setCurrentBank(value);
setOpen(false);
// Preserve current view and subTab when switching banks
const view = searchParams.get('view') || 'data';
const subTab = searchParams.get('subTab');
const queryString = subTab ? `?view=${view}&subTab=${subTab}` : `?view=${view}`;
const view = searchParams.get("view") || "data";
const subTab = searchParams.get("subTab");
const queryString = subTab
? `?view=${view}&subTab=${subTab}`
: `?view=${view}`;
router.push(`/banks/${value}${queryString}`);
}}
>
@@ -234,13 +237,9 @@ function BankSelectorInner() {
size="icon"
onClick={toggleTheme}
className="h-9 w-9"
title={theme === 'light' ? 'Switch to dark mode' : 'Switch to light mode'}
title={theme === "light" ? "Switch to dark mode" : "Switch to light mode"}
>
{theme === 'light' ? (
<Moon className="h-5 w-5" />
) : (
<Sun className="h-5 w-5" />
)}
{theme === "light" ? <Moon className="h-5 w-5" /> : <Sun className="h-5 w-5" />}
</Button>
<Dialog open={createDialogOpen} onOpenChange={setCreateDialogOpen}>
@@ -254,32 +253,27 @@ function BankSelectorInner() {
value={newBankId}
onChange={(e) => setNewBankId(e.target.value)}
onKeyDown={(e) => {
if (e.key === 'Enter' && !isCreating) {
if (e.key === "Enter" && !isCreating) {
handleCreateBank();
}
}}
autoFocus
/>
{createError && (
<p className="text-sm text-destructive mt-2">{createError}</p>
)}
{createError && <p className="text-sm text-destructive mt-2">{createError}</p>}
</div>
<DialogFooter>
<Button
variant="secondary"
onClick={() => {
setCreateDialogOpen(false);
setNewBankId('');
setNewBankId("");
setCreateError(null);
}}
>
Cancel
</Button>
<Button
onClick={handleCreateBank}
disabled={isCreating || !newBankId.trim()}
>
{isCreating ? 'Creating...' : 'Create'}
<Button onClick={handleCreateBank} disabled={isCreating || !newBankId.trim()}>
{isCreating ? "Creating..." : "Create"}
</Button>
</DialogFooter>
</DialogContent>
@@ -290,7 +284,8 @@ function BankSelectorInner() {
<DialogHeader>
<DialogTitle>Add New Document</DialogTitle>
<p className="text-sm text-muted-foreground">
Add a new document to memory bank: <span className="font-semibold">{currentBank}</span>
Add a new document to memory bank:{" "}
<span className="font-semibold">{currentBank}</span>
</p>
</DialogHeader>
<div className="py-4 space-y-4">
@@ -327,7 +322,9 @@ function BankSelectorInner() {
</div>
<div>
<label className="font-bold block mb-1 text-sm text-foreground">Document ID</label>
<label className="font-bold block mb-1 text-sm text-foreground">
Document ID
</label>
<Input
type="text"
value={docDocumentId}
@@ -348,30 +345,25 @@ function BankSelectorInner() {
</label>
</div>
{docError && (
<p className="text-sm text-destructive">{docError}</p>
)}
{docError && <p className="text-sm text-destructive">{docError}</p>}
</div>
<DialogFooter>
<Button
variant="secondary"
onClick={() => {
setDocDialogOpen(false);
setDocContent('');
setDocContext('');
setDocEventDate('');
setDocDocumentId('');
setDocContent("");
setDocContext("");
setDocEventDate("");
setDocDocumentId("");
setDocAsync(false);
setDocError(null);
}}
>
Cancel
</Button>
<Button
onClick={handleCreateDocument}
disabled={isCreatingDoc || !docContent.trim()}
>
{isCreatingDoc ? 'Adding...' : 'Add Document'}
<Button onClick={handleCreateDocument} disabled={isCreatingDoc || !docContent.trim()}>
{isCreatingDoc ? "Adding..." : "Add Document"}
</Button>
</DialogFooter>
</DialogContent>
@@ -383,36 +375,45 @@ function BankSelectorInner() {
export function BankSelector() {
return (
<Suspense fallback={
<div className="bg-card text-card-foreground px-5 py-3 border-b-4 border-primary-gradient">
<div className="flex items-center gap-4 text-sm">
<Image src="/logo.png" alt="Hindsight" width={40} height={40} className="h-10 w-auto" unoptimized />
<div className="h-8 w-px bg-border" />
<Button
variant="outline"
className="w-[250px] justify-between font-bold border-2 border-primary"
disabled
>
Loading...
<ChevronsUpDown className="ml-2 h-4 w-4 shrink-0 opacity-50" />
</Button>
<div className="flex-1" />
<a
href="https://github.com/vectorize-io/hindsight"
target="_blank"
rel="noopener noreferrer"
className="flex items-center gap-2 px-3 py-2 rounded-lg hover:bg-accent transition-colors text-muted-foreground"
>
<Github className="h-5 w-5" />
<span className="text-sm font-medium">GitHub</span>
</a>
<div className="h-8 w-px bg-border" />
<Button variant="ghost" size="icon" className="h-9 w-9" disabled>
<Moon className="h-5 w-5" />
</Button>
<Suspense
fallback={
<div className="bg-card text-card-foreground px-5 py-3 border-b-4 border-primary-gradient">
<div className="flex items-center gap-4 text-sm">
<Image
src="/logo.png"
alt="Hindsight"
width={40}
height={40}
className="h-10 w-auto"
unoptimized
/>
<div className="h-8 w-px bg-border" />
<Button
variant="outline"
className="w-[250px] justify-between font-bold border-2 border-primary"
disabled
>
Loading...
<ChevronsUpDown className="ml-2 h-4 w-4 shrink-0 opacity-50" />
</Button>
<div className="flex-1" />
<a
href="https://github.com/vectorize-io/hindsight"
target="_blank"
rel="noopener noreferrer"
className="flex items-center gap-2 px-3 py-2 rounded-lg hover:bg-accent transition-colors text-muted-foreground"
>
<Github className="h-5 w-5" />
<span className="text-sm font-medium">GitHub</span>
</a>
<div className="h-8 w-px bg-border" />
<Button variant="ghost" size="icon" className="h-9 w-9" disabled>
<Moon className="h-5 w-5" />
</Button>
</div>
</div>
</div>
}>
}
>
<BankSelectorInner />
</Suspense>
);
@@ -1,20 +1,40 @@
'use client';
"use client";
import { useState, useEffect, useRef, useMemo, useCallback } from 'react';
import { client } from '@/lib/api';
import { useBank } from '@/lib/bank-context';
import { Button } from '@/components/ui/button';
import { Input } from '@/components/ui/input';
import { Copy, Check, Calendar, ZoomIn, ZoomOut, ChevronLeft, ChevronRight, ChevronsLeft, ChevronsRight, Settings2, Eye, EyeOff } from 'lucide-react';
import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from '@/components/ui/table';
import { Label } from '@/components/ui/label';
import { Slider } from '@/components/ui/slider';
import { Switch } from '@/components/ui/switch';
import { MemoryDetailPanel } from './memory-detail-panel';
import { Graph2D, convertHindsightGraphData, GraphNode } from './graph-2d';
import { useState, useEffect, useRef, useMemo, useCallback } from "react";
import { client } from "@/lib/api";
import { useBank } from "@/lib/bank-context";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import {
Copy,
Check,
Calendar,
ZoomIn,
ZoomOut,
ChevronLeft,
ChevronRight,
ChevronsLeft,
ChevronsRight,
Settings2,
Eye,
EyeOff,
} from "lucide-react";
import {
Table,
TableBody,
TableCell,
TableHead,
TableHeader,
TableRow,
} from "@/components/ui/table";
import { Label } from "@/components/ui/label";
import { Slider } from "@/components/ui/slider";
import { Switch } from "@/components/ui/switch";
import { MemoryDetailPanel } from "./memory-detail-panel";
import { Graph2D, convertHindsightGraphData, GraphNode } from "./graph-2d";
type FactType = 'world' | 'experience' | 'opinion';
type ViewMode = 'graph' | 'table' | 'timeline';
type FactType = "world" | "experience" | "opinion";
type ViewMode = "graph" | "table" | "timeline";
interface DataViewProps {
factType: FactType;
@@ -22,10 +42,10 @@ interface DataViewProps {
export function DataView({ factType }: DataViewProps) {
const { currentBank } = useBank();
const [viewMode, setViewMode] = useState<ViewMode>('graph');
const [viewMode, setViewMode] = useState<ViewMode>("graph");
const [data, setData] = useState<any>(null);
const [loading, setLoading] = useState(false);
const [searchQuery, setSearchQuery] = useState('');
const [searchQuery, setSearchQuery] = useState("");
const [copiedId, setCopiedId] = useState<string | null>(null);
const [currentPage, setCurrentPage] = useState(1);
const [selectedGraphNode, setSelectedGraphNode] = useState<any>(null);
@@ -36,10 +56,12 @@ export function DataView({ factType }: DataViewProps) {
const [showLabels, setShowLabels] = useState(true);
const [maxNodes, setMaxNodes] = useState<number | undefined>(50);
const [showControlPanel, setShowControlPanel] = useState(true);
const [visibleLinkTypes, setVisibleLinkTypes] = useState<Set<string>>(new Set(['semantic', 'temporal', 'entity', 'causal']));
const [visibleLinkTypes, setVisibleLinkTypes] = useState<Set<string>>(
new Set(["semantic", "temporal", "entity", "causal"])
);
const toggleLinkType = (type: string) => {
setVisibleLinkTypes(prev => {
setVisibleLinkTypes((prev) => {
const next = new Set(prev);
if (next.has(type)) {
next.delete(type);
@@ -53,12 +75,12 @@ export function DataView({ factType }: DataViewProps) {
// Esc key handler to deselect graph node
useEffect(() => {
const handleKeyDown = (e: KeyboardEvent) => {
if (e.key === 'Escape' && selectedGraphNode) {
if (e.key === "Escape" && selectedGraphNode) {
setSelectedGraphNode(null);
}
};
window.addEventListener('keydown', handleKeyDown);
return () => window.removeEventListener('keydown', handleKeyDown);
window.addEventListener("keydown", handleKeyDown);
return () => window.removeEventListener("keydown", handleKeyDown);
}, [selectedGraphNode]);
const copyToClipboard = async (text: string) => {
@@ -67,7 +89,7 @@ export function DataView({ factType }: DataViewProps) {
setCopiedId(text);
setTimeout(() => setCopiedId(null), 2000);
} catch (err) {
console.error('Failed to copy:', err);
console.error("Failed to copy:", err);
}
};
@@ -80,7 +102,7 @@ export function DataView({ factType }: DataViewProps) {
bank_id: currentBank,
type: factType,
});
console.log('Loaded graph data:', {
console.log("Loaded graph data:", {
total_units: graphData.total_units,
nodes: graphData.nodes?.length,
edges: graphData.edges?.length,
@@ -88,7 +110,7 @@ export function DataView({ factType }: DataViewProps) {
});
setData(graphData);
} catch (error) {
console.error('Error loading data:', error);
console.error("Error loading data:", error);
alert(`Error loading ${factType} data: ` + (error as Error).message);
} finally {
setLoading(false);
@@ -101,9 +123,7 @@ export function DataView({ factType }: DataViewProps) {
if (!searchQuery) return data.table_rows;
const query = searchQuery.toLowerCase();
return data.table_rows.filter((row: any) =>
row.text?.toLowerCase().includes(query)
);
return data.table_rows.filter((row: any) => row.text?.toLowerCase().includes(query));
}, [data, searchQuery]);
// Get filtered node IDs for graph filtering
@@ -113,10 +133,10 @@ export function DataView({ factType }: DataViewProps) {
// Helper to get normalized link type
const getLinkTypeCategory = (type: string | undefined): string => {
if (!type) return 'semantic';
if (type === 'semantic' || type === 'temporal' || type === 'entity') return type;
if (['causes', 'caused_by', 'enables', 'prevents'].includes(type)) return 'causal';
return 'semantic';
if (!type) return "semantic";
if (type === "semantic" || type === "temporal" || type === "entity") return type;
if (["causes", "caused_by", "enables", "prevents"].includes(type)) return "causal";
return "semantic";
};
// Convert data for Graph2D with filtering
@@ -129,16 +149,16 @@ export function DataView({ factType }: DataViewProps) {
// Filter nodes based on search query
if (searchQuery) {
const filteredNodes = fullData.nodes.filter(node => filteredNodeIds.has(node.id));
const filteredNodeIdSet = new Set(filteredNodes.map(n => n.id));
const filteredNodes = fullData.nodes.filter((node) => filteredNodeIds.has(node.id));
const filteredNodeIdSet = new Set(filteredNodes.map((n) => n.id));
nodes = filteredNodes;
links = fullData.links.filter(link =>
filteredNodeIdSet.has(link.source) && filteredNodeIdSet.has(link.target)
links = fullData.links.filter(
(link) => filteredNodeIdSet.has(link.source) && filteredNodeIdSet.has(link.target)
);
}
// Filter links based on visible link types
links = links.filter(link => {
links = links.filter((link) => {
const category = getLinkTypeCategory(link.type);
return visibleLinkTypes.has(category);
});
@@ -148,44 +168,62 @@ export function DataView({ factType }: DataViewProps) {
// Calculate link stats for display
const linkStats = useMemo(() => {
let semantic = 0, temporal = 0, entity = 0, causal = 0, total = 0;
let semantic = 0,
temporal = 0,
entity = 0,
causal = 0,
total = 0;
const otherTypes: Record<string, number> = {};
graph2DData.links.forEach(l => {
graph2DData.links.forEach((l) => {
total++;
const type = l.type || 'unknown';
if (type === 'semantic') semantic++;
else if (type === 'temporal') temporal++;
else if (type === 'entity') entity++;
else if (type === 'causes' || type === 'caused_by' || type === 'enables' || type === 'prevents') causal++;
const type = l.type || "unknown";
if (type === "semantic") semantic++;
else if (type === "temporal") temporal++;
else if (type === "entity") entity++;
else if (
type === "causes" ||
type === "caused_by" ||
type === "enables" ||
type === "prevents"
)
causal++;
else {
otherTypes[type] = (otherTypes[type] || 0) + 1;
}
});
console.log('Graph link stats:', { semantic, temporal, entity, causal, total });
console.log("Graph link stats:", { semantic, temporal, entity, causal, total });
if (Object.keys(otherTypes).length > 0) {
console.log('Other link types:', otherTypes);
console.log("Other link types:", otherTypes);
}
return { semantic, temporal, entity, causal, total, otherTypes };
}, [graph2DData]);
// Handle node click in graph - show in panel
const handleGraphNodeClick = useCallback((node: GraphNode) => {
const nodeData = data?.table_rows?.find((row: any) => row.id === node.id);
if (nodeData) {
setSelectedGraphNode(nodeData);
}
}, [data]);
const handleGraphNodeClick = useCallback(
(node: GraphNode) => {
const nodeData = data?.table_rows?.find((row: any) => row.id === node.id);
if (nodeData) {
setSelectedGraphNode(nodeData);
}
},
[data]
);
// Memoized color functions to prevent graph re-initialization
// Uses brand colors: primary blue (#0074d9), teal (#009296), amber for entity, purple for causal
const nodeColorFn = useCallback((node: GraphNode) => node.color || '#0074d9', []);
const nodeColorFn = useCallback((node: GraphNode) => node.color || "#0074d9", []);
const linkColorFn = useCallback((link: any) => {
if (link.type === 'temporal') return '#009296'; // Brand teal
if (link.type === 'entity') return '#f59e0b'; // Amber
if (link.type === 'causes' || link.type === 'caused_by' || link.type === 'enables' || link.type === 'prevents') {
return '#8b5cf6'; // Purple for causal
if (link.type === "temporal") return "#009296"; // Brand teal
if (link.type === "entity") return "#f59e0b"; // Amber
if (
link.type === "causes" ||
link.type === "caused_by" ||
link.type === "enables" ||
link.type === "prevents"
) {
return "#8b5cf6"; // Purple for causal
}
return '#0074d9'; // Brand primary blue for semantic
return "#0074d9"; // Brand primary blue for semantic
}, []);
// Reset to first page when search query changes
@@ -224,35 +262,37 @@ export function DataView({ factType }: DataViewProps) {
<div className="flex items-center justify-between mb-6">
<div className="text-sm text-muted-foreground">
{searchQuery ? `${filteredTableRows.length} of ${data.total_units} memories` : `${data.total_units} total memories`}
{searchQuery
? `${filteredTableRows.length} of ${data.total_units} memories`
: `${data.total_units} total memories`}
</div>
<div className="flex items-center gap-2 bg-muted rounded-lg p-1">
<button
onClick={() => setViewMode('graph')}
onClick={() => setViewMode("graph")}
className={`px-4 py-2 rounded-md text-sm font-medium transition-all ${
viewMode === 'graph'
? 'bg-background text-foreground shadow-sm'
: 'text-muted-foreground hover:text-foreground'
viewMode === "graph"
? "bg-background text-foreground shadow-sm"
: "text-muted-foreground hover:text-foreground"
}`}
>
Graph View
</button>
<button
onClick={() => setViewMode('table')}
onClick={() => setViewMode("table")}
className={`px-4 py-2 rounded-md text-sm font-medium transition-all ${
viewMode === 'table'
? 'bg-background text-foreground shadow-sm'
: 'text-muted-foreground hover:text-foreground'
viewMode === "table"
? "bg-background text-foreground shadow-sm"
: "text-muted-foreground hover:text-foreground"
}`}
>
Table View
</button>
<button
onClick={() => setViewMode('timeline')}
onClick={() => setViewMode("timeline")}
className={`px-4 py-2 rounded-md text-sm font-medium transition-all ${
viewMode === 'timeline'
? 'bg-background text-foreground shadow-sm'
: 'text-muted-foreground hover:text-foreground'
viewMode === "timeline"
? "bg-background text-foreground shadow-sm"
: "text-muted-foreground hover:text-foreground"
}`}
>
Timeline View
@@ -260,7 +300,7 @@ export function DataView({ factType }: DataViewProps) {
</div>
</div>
{viewMode === 'graph' && (
{viewMode === "graph" && (
<div className="flex gap-0">
{/* Graph */}
<div className="flex-1 min-w-0">
@@ -279,7 +319,7 @@ export function DataView({ factType }: DataViewProps) {
<button
onClick={() => setShowControlPanel(!showControlPanel)}
className="flex-shrink-0 w-5 h-[700px] bg-transparent hover:bg-muted/50 flex items-center justify-center transition-colors"
title={showControlPanel ? 'Hide panel' : 'Show panel'}
title={showControlPanel ? "Hide panel" : "Show panel"}
>
{showControlPanel ? (
<ChevronRight className="w-3 h-3 text-muted-foreground/60" />
@@ -289,7 +329,9 @@ export function DataView({ factType }: DataViewProps) {
</button>
{/* Right Panel - Legend/Controls OR Memory Details */}
<div className={`${showControlPanel ? 'w-80' : 'w-0'} transition-all duration-300 overflow-hidden flex-shrink-0`}>
<div
className={`${showControlPanel ? "w-80" : "w-0"} transition-all duration-300 overflow-hidden flex-shrink-0`}
>
<div className="w-80 h-[700px] bg-card border-l border-border overflow-y-auto">
{selectedGraphNode ? (
/* Memory Detail View */
@@ -308,47 +350,67 @@ export function DataView({ factType }: DataViewProps) {
{/* Nodes */}
<div className="flex items-center justify-between text-sm">
<div className="flex items-center gap-2">
<div className="w-3 h-3 rounded-full" style={{ backgroundColor: '#0074d9' }} />
<div
className="w-3 h-3 rounded-full"
style={{ backgroundColor: "#0074d9" }}
/>
<span className="text-foreground">Nodes</span>
</div>
<span className="font-mono text-foreground">
{Math.min(maxNodes ?? graph2DData.nodes.length, graph2DData.nodes.length)}/{graph2DData.nodes.length}
{Math.min(
maxNodes ?? graph2DData.nodes.length,
graph2DData.nodes.length
)}
/{graph2DData.nodes.length}
</span>
</div>
<div className="text-xs font-medium text-muted-foreground mt-2 mb-1">Links ({linkStats.total}) <span className="text-muted-foreground/60">· click to filter</span></div>
<div className="text-xs font-medium text-muted-foreground mt-2 mb-1">
Links ({linkStats.total}){" "}
<span className="text-muted-foreground/60">· click to filter</span>
</div>
<button
onClick={() => toggleLinkType('semantic')}
onClick={() => toggleLinkType("semantic")}
className={`w-full flex items-center justify-between text-sm px-2 py-1 rounded transition-all ${
visibleLinkTypes.has('semantic') ? 'hover:bg-muted' : 'opacity-40 hover:opacity-60'
visibleLinkTypes.has("semantic")
? "hover:bg-muted"
: "opacity-40 hover:opacity-60"
}`}
>
<div className="flex items-center gap-2">
<div className="w-4 h-0.5 bg-[#0074d9]" />
<span className="text-foreground">Semantic</span>
</div>
<span className={`font-mono ${linkStats.semantic === 0 ? 'text-destructive' : 'text-foreground'}`}>
<span
className={`font-mono ${linkStats.semantic === 0 ? "text-destructive" : "text-foreground"}`}
>
{linkStats.semantic}
</span>
</button>
<button
onClick={() => toggleLinkType('temporal')}
onClick={() => toggleLinkType("temporal")}
className={`w-full flex items-center justify-between text-sm px-2 py-1 rounded transition-all ${
visibleLinkTypes.has('temporal') ? 'hover:bg-muted' : 'opacity-40 hover:opacity-60'
visibleLinkTypes.has("temporal")
? "hover:bg-muted"
: "opacity-40 hover:opacity-60"
}`}
>
<div className="flex items-center gap-2">
<div className="w-4 h-0.5 bg-[#009296]" />
<span className="text-foreground">Temporal</span>
</div>
<span className={`font-mono ${linkStats.temporal === 0 ? 'text-destructive' : 'text-foreground'}`}>
<span
className={`font-mono ${linkStats.temporal === 0 ? "text-destructive" : "text-foreground"}`}
>
{linkStats.temporal}
</span>
</button>
<button
onClick={() => toggleLinkType('entity')}
onClick={() => toggleLinkType("entity")}
className={`w-full flex items-center justify-between text-sm px-2 py-1 rounded transition-all ${
visibleLinkTypes.has('entity') ? 'hover:bg-muted' : 'opacity-40 hover:opacity-60'
visibleLinkTypes.has("entity")
? "hover:bg-muted"
: "opacity-40 hover:opacity-60"
}`}
>
<div className="flex items-center gap-2">
@@ -358,23 +420,29 @@ export function DataView({ factType }: DataViewProps) {
<span className="font-mono text-foreground">{linkStats.entity}</span>
</button>
<button
onClick={() => toggleLinkType('causal')}
onClick={() => toggleLinkType("causal")}
className={`w-full flex items-center justify-between text-sm px-2 py-1 rounded transition-all ${
visibleLinkTypes.has('causal') ? 'hover:bg-muted' : 'opacity-40 hover:opacity-60'
visibleLinkTypes.has("causal")
? "hover:bg-muted"
: "opacity-40 hover:opacity-60"
}`}
>
<div className="flex items-center gap-2">
<div className="w-4 h-0.5 bg-[#8b5cf6]" />
<span className="text-foreground">Causal</span>
</div>
<span className={`font-mono ${linkStats.causal === 0 ? 'text-muted-foreground' : 'text-foreground'}`}>
<span
className={`font-mono ${linkStats.causal === 0 ? "text-muted-foreground" : "text-foreground"}`}
>
{linkStats.causal}
</span>
</button>
{Object.entries(linkStats.otherTypes || {}).map(([type, count]) => (
<div key={type} className="flex items-center justify-between text-sm">
<span className="text-muted-foreground capitalize ml-6">{type}</span>
<span className="font-mono text-muted-foreground">{count as number}</span>
<span className="font-mono text-muted-foreground">
{count as number}
</span>
</div>
))}
</div>
@@ -387,7 +455,9 @@ export function DataView({ factType }: DataViewProps) {
<h3 className="text-sm font-semibold mb-3 text-foreground">Display</h3>
<div className="space-y-4">
<div className="flex items-center justify-between">
<Label htmlFor="show-labels" className="text-sm text-foreground">Show labels</Label>
<Label htmlFor="show-labels" className="text-sm text-foreground">
Show labels
</Label>
<Switch
id="show-labels"
checked={showLabels}
@@ -407,7 +477,7 @@ export function DataView({ factType }: DataViewProps) {
<div className="flex items-center justify-between mb-2">
<Label className="text-sm text-foreground">Max nodes</Label>
<span className="text-xs text-muted-foreground">
{maxNodes ?? 'All'} / {graph2DData.nodes.length}
{maxNodes ?? "All"} / {graph2DData.nodes.length}
</span>
</div>
<Slider
@@ -415,7 +485,9 @@ export function DataView({ factType }: DataViewProps) {
min={10}
max={Math.max(graph2DData.nodes.length, 10)}
step={10}
onValueChange={([v]) => setMaxNodes(v >= graph2DData.nodes.length ? undefined : v)}
onValueChange={([v]) =>
setMaxNodes(v >= graph2DData.nodes.length ? undefined : v)
}
className="w-full"
/>
</div>
@@ -438,7 +510,7 @@ export function DataView({ factType }: DataViewProps) {
</div>
)}
{viewMode === 'table' && (
{viewMode === "table" && (
<div>
<div className="w-full">
<div className="pb-4">
@@ -465,10 +537,16 @@ export function DataView({ factType }: DataViewProps) {
<TableBody>
{paginatedRows.map((row: any, idx: number) => {
const occurredDisplay = row.occurred_start
? new Date(row.occurred_start).toLocaleDateString('en-US', { month: 'short', day: 'numeric' })
? new Date(row.occurred_start).toLocaleDateString("en-US", {
month: "short",
day: "numeric",
})
: null;
const mentionedDisplay = row.mentioned_at
? new Date(row.mentioned_at).toLocaleDateString('en-US', { month: 'short', day: 'numeric' })
? new Date(row.mentioned_at).toLocaleDateString("en-US", {
month: "short",
day: "numeric",
})
: null;
return (
@@ -476,29 +554,36 @@ export function DataView({ factType }: DataViewProps) {
key={row.id || idx}
onClick={() => setSelectedTableMemory(row)}
className={`cursor-pointer hover:bg-muted/50 ${
selectedTableMemory?.id === row.id ? 'bg-primary/10' : ''
selectedTableMemory?.id === row.id ? "bg-primary/10" : ""
}`}
>
<TableCell className="py-2">
<div className="line-clamp-2 text-sm leading-snug text-foreground">{row.text}</div>
<div className="line-clamp-2 text-sm leading-snug text-foreground">
{row.text}
</div>
{row.context && (
<div className="text-xs text-muted-foreground mt-0.5 truncate">{row.context}</div>
<div className="text-xs text-muted-foreground mt-0.5 truncate">
{row.context}
</div>
)}
</TableCell>
<TableCell className="py-2">
{row.entities ? (
<div className="flex gap-1 flex-wrap">
{row.entities.split(', ').slice(0, 2).map((entity: string, i: number) => (
<span
key={i}
className="text-[10px] px-1.5 py-0.5 rounded-full bg-primary/10 text-primary font-medium"
>
{entity}
</span>
))}
{row.entities.split(', ').length > 2 && (
{row.entities
.split(", ")
.slice(0, 2)
.map((entity: string, i: number) => (
<span
key={i}
className="text-[10px] px-1.5 py-0.5 rounded-full bg-primary/10 text-primary font-medium"
>
{entity}
</span>
))}
{row.entities.split(", ").length > 2 && (
<span className="text-[10px] text-muted-foreground">
+{row.entities.split(', ').length - 2}
+{row.entities.split(", ").length - 2}
</span>
)}
</div>
@@ -507,10 +592,14 @@ export function DataView({ factType }: DataViewProps) {
)}
</TableCell>
<TableCell className="text-xs py-2 text-foreground">
{occurredDisplay || <span className="text-muted-foreground">-</span>}
{occurredDisplay || (
<span className="text-muted-foreground">-</span>
)}
</TableCell>
<TableCell className="text-xs py-2 text-foreground">
{mentionedDisplay || <span className="text-muted-foreground">-</span>}
{mentionedDisplay || (
<span className="text-muted-foreground">-</span>
)}
</TableCell>
<TableCell className="py-2">
<Button
@@ -541,7 +630,8 @@ export function DataView({ factType }: DataViewProps) {
{totalPages > 1 && (
<div className="flex items-center justify-between mt-3 pt-3 border-t">
<div className="text-xs text-muted-foreground">
{startIndex + 1}-{Math.min(endIndex, filteredTableRows.length)} of {filteredTableRows.length}
{startIndex + 1}-{Math.min(endIndex, filteredTableRows.length)} of{" "}
{filteredTableRows.length}
</div>
<div className="flex items-center gap-1">
<Button
@@ -556,7 +646,7 @@ export function DataView({ factType }: DataViewProps) {
<Button
variant="outline"
size="sm"
onClick={() => setCurrentPage(p => Math.max(1, p - 1))}
onClick={() => setCurrentPage((p) => Math.max(1, p - 1))}
disabled={currentPage === 1}
className="h-7 w-7 p-0"
>
@@ -568,7 +658,7 @@ export function DataView({ factType }: DataViewProps) {
<Button
variant="outline"
size="sm"
onClick={() => setCurrentPage(p => Math.min(totalPages, p + 1))}
onClick={() => setCurrentPage((p) => Math.min(totalPages, p + 1))}
disabled={currentPage === totalPages}
className="h-7 w-7 p-0"
>
@@ -591,7 +681,9 @@ export function DataView({ factType }: DataViewProps) {
})()
) : (
<div className="text-center py-12 text-muted-foreground">
{data.table_rows?.length > 0 ? 'No memories match your filter' : 'No memories found'}
{data.table_rows?.length > 0
? "No memories match your filter"
: "No memories found"}
</div>
)}
</div>
@@ -610,9 +702,7 @@ export function DataView({ factType }: DataViewProps) {
</div>
)}
{viewMode === 'timeline' && (
<TimelineView data={data} filteredRows={filteredTableRows} />
)}
{viewMode === "timeline" && <TimelineView data={data} filteredRows={filteredTableRows} />}
</>
) : (
<div className="flex items-center justify-center py-20">
@@ -627,17 +717,18 @@ export function DataView({ factType }: DataViewProps) {
}
// Timeline View Component - Custom compact timeline with zoom and navigation
type Granularity = 'year' | 'month' | 'week' | 'day';
type Granularity = "year" | "month" | "week" | "day";
function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }) {
const [selectedItem, setSelectedItem] = useState<any>(null);
const [granularity, setGranularity] = useState<Granularity>('month');
const [granularity, setGranularity] = useState<Granularity>("month");
const [currentIndex, setCurrentIndex] = useState(0);
const timelineRef = useRef<HTMLDivElement>(null);
// Filter and sort items that have occurred_start dates (using filtered data)
const { sortedItems, itemsWithoutDates } = useMemo(() => {
if (!filteredRows || filteredRows.length === 0) return { sortedItems: [], itemsWithoutDates: [] };
if (!filteredRows || filteredRows.length === 0)
return { sortedItems: [], itemsWithoutDates: [] };
const withDates = filteredRows
.filter((row: any) => row.occurred_start)
@@ -662,31 +753,36 @@ function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }
const day = date.getDate();
switch (granularity) {
case 'year':
case "year":
return `${year}`;
case 'month':
return `${year}-${String(month + 1).padStart(2, '0')}`;
case 'week':
case "month":
return `${year}-${String(month + 1).padStart(2, "0")}`;
case "week":
const startOfWeek = new Date(date);
startOfWeek.setDate(day - date.getDay());
return `${startOfWeek.getFullYear()}-W${String(Math.ceil((startOfWeek.getDate()) / 7)).padStart(2, '0')}-${String(startOfWeek.getMonth() + 1).padStart(2, '0')}-${String(startOfWeek.getDate()).padStart(2, '0')}`;
case 'day':
return `${year}-${String(month + 1).padStart(2, '0')}-${String(day).padStart(2, '0')}`;
return `${startOfWeek.getFullYear()}-W${String(Math.ceil(startOfWeek.getDate() / 7)).padStart(2, "0")}-${String(startOfWeek.getMonth() + 1).padStart(2, "0")}-${String(startOfWeek.getDate()).padStart(2, "0")}`;
case "day":
return `${year}-${String(month + 1).padStart(2, "0")}-${String(day).padStart(2, "0")}`;
}
};
const getGroupLabel = (key: string, date: Date): string => {
switch (granularity) {
case 'year':
case "year":
return key;
case 'month':
return date.toLocaleDateString('en-US', { year: 'numeric', month: 'short' });
case 'week':
case "month":
return date.toLocaleDateString("en-US", { year: "numeric", month: "short" });
case "week":
const endOfWeek = new Date(date);
endOfWeek.setDate(date.getDate() + 6);
return `${date.toLocaleDateString('en-US', { month: 'short', day: 'numeric' })} - ${endOfWeek.toLocaleDateString('en-US', { month: 'short', day: 'numeric', year: 'numeric' })}`;
case 'day':
return date.toLocaleDateString('en-US', { weekday: 'short', month: 'short', day: 'numeric', year: 'numeric' });
return `${date.toLocaleDateString("en-US", { month: "short", day: "numeric" })} - ${endOfWeek.toLocaleDateString("en-US", { month: "short", day: "numeric", year: "numeric" })}`;
case "day":
return date.toLocaleDateString("en-US", {
weekday: "short",
month: "short",
day: "numeric",
year: "numeric",
});
}
};
@@ -697,8 +793,8 @@ function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }
if (!groups[key]) {
// For week, parse the start date from key
let groupDate = date;
if (granularity === 'week') {
const parts = key.split('-');
if (granularity === "week") {
const parts = key.split("-");
groupDate = new Date(parseInt(parts[0]), parseInt(parts[2]) - 1, parseInt(parts[3]));
}
groups[key] = { items: [], date: groupDate };
@@ -729,11 +825,11 @@ function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }
const clampedIndex = Math.max(0, Math.min(index, timelineGroups.length - 1));
setCurrentIndex(clampedIndex);
const element = document.getElementById(`timeline-group-${clampedIndex}`);
element?.scrollIntoView({ behavior: 'smooth', block: 'start' });
element?.scrollIntoView({ behavior: "smooth", block: "start" });
};
const zoomIn = () => {
const levels: Granularity[] = ['year', 'month', 'week', 'day'];
const levels: Granularity[] = ["year", "month", "week", "day"];
const currentIdx = levels.indexOf(granularity);
if (currentIdx < levels.length - 1) {
setGranularity(levels[currentIdx + 1]);
@@ -741,7 +837,7 @@ function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }
};
const zoomOut = () => {
const levels: Granularity[] = ['year', 'month', 'week', 'day'];
const levels: Granularity[] = ["year", "month", "week", "day"];
const currentIdx = levels.indexOf(granularity);
if (currentIdx > 0) {
setGranularity(levels[currentIdx - 1]);
@@ -767,16 +863,20 @@ function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }
const formatDateTime = (dateStr: string) => {
const date = new Date(dateStr);
const dateFormatted = date.toLocaleDateString('en-US', { month: 'short', day: 'numeric' });
const timeFormatted = date.toLocaleTimeString('en-US', { hour: '2-digit', minute: '2-digit', hour12: false });
const dateFormatted = date.toLocaleDateString("en-US", { month: "short", day: "numeric" });
const timeFormatted = date.toLocaleTimeString("en-US", {
hour: "2-digit",
minute: "2-digit",
hour12: false,
});
return { date: dateFormatted, time: timeFormatted };
};
const granularityLabels: Record<Granularity, string> = {
year: 'Year',
month: 'Month',
week: 'Week',
day: 'Day',
year: "Year",
month: "Month",
week: "Week",
day: "Day",
};
return (
@@ -790,7 +890,8 @@ function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }
{itemsWithoutDates.length > 0 && ` · ${itemsWithoutDates.length} without dates`}
{dateRange && (
<span className="ml-2 text-foreground">
({dateRange.first.toLocaleDateString('en-US', { month: 'short', year: 'numeric' })} {dateRange.last.toLocaleDateString('en-US', { month: 'short', year: 'numeric' })})
({dateRange.first.toLocaleDateString("en-US", { month: "short", year: "numeric" })}{" "}
{dateRange.last.toLocaleDateString("en-US", { month: "short", year: "numeric" })})
</span>
)}
</div>
@@ -802,7 +903,7 @@ function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }
variant="secondary"
size="sm"
onClick={zoomOut}
disabled={granularity === 'year'}
disabled={granularity === "year"}
className="h-7 w-7 p-0"
title="Zoom out"
>
@@ -815,7 +916,7 @@ function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }
variant="secondary"
size="sm"
onClick={zoomIn}
disabled={granularity === 'day'}
disabled={granularity === "day"}
className="h-7 w-7 p-0"
title="Zoom in"
>
@@ -888,7 +989,7 @@ function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }
</div>
<div className="w-2 h-2 rounded-full bg-primary z-10" />
<span className="ml-2 text-[10px] text-muted-foreground">
{group.items.length} {group.items.length === 1 ? 'item' : 'items'}
{group.items.length} {group.items.length === 1 ? "item" : "items"}
</span>
</div>
@@ -899,7 +1000,7 @@ function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }
key={item.id || idx}
onClick={() => setSelectedItem(item)}
className={`flex items-start cursor-pointer group ${
selectedItem?.id === item.id ? 'opacity-100' : 'hover:opacity-80'
selectedItem?.id === item.id ? "opacity-100" : "hover:opacity-80"
}`}
>
{/* Date & Time */}
@@ -914,17 +1015,23 @@ function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }
{/* Connector dot */}
<div className="flex-shrink-0 pt-2">
<div className={`w-1.5 h-1.5 rounded-full z-10 ${
selectedItem?.id === item.id ? 'bg-primary' : 'bg-muted-foreground/50 group-hover:bg-primary'
}`} />
<div
className={`w-1.5 h-1.5 rounded-full z-10 ${
selectedItem?.id === item.id
? "bg-primary"
: "bg-muted-foreground/50 group-hover:bg-primary"
}`}
/>
</div>
{/* Card */}
<div className={`ml-3 flex-1 p-2 rounded border transition-colors ${
selectedItem?.id === item.id
? 'bg-primary/10 border-primary'
: 'bg-card border-border hover:border-primary/50'
}`}>
<div
className={`ml-3 flex-1 p-2 rounded border transition-colors ${
selectedItem?.id === item.id
? "bg-primary/10 border-primary"
: "bg-card border-border hover:border-primary/50"
}`}
>
<p className="text-xs text-foreground line-clamp-2 leading-relaxed">
{item.text}
</p>
@@ -935,14 +1042,20 @@ function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }
)}
{item.entities && (
<div className="flex gap-1 mt-1 flex-wrap">
{item.entities.split(', ').slice(0, 3).map((entity: string, i: number) => (
<span key={i} className="text-[9px] px-1.5 py-0.5 rounded-full bg-primary/10 text-primary font-medium">
{entity}
</span>
))}
{item.entities.split(', ').length > 3 && (
{item.entities
.split(", ")
.slice(0, 3)
.map((entity: string, i: number) => (
<span
key={i}
className="text-[9px] px-1.5 py-0.5 rounded-full bg-primary/10 text-primary font-medium"
>
{entity}
</span>
))}
{item.entities.split(", ").length > 3 && (
<span className="text-[9px] text-muted-foreground">
+{item.entities.split(', ').length - 3}
+{item.entities.split(", ").length - 3}
</span>
)}
</div>
@@ -959,11 +1072,7 @@ function TimelineView({ data, filteredRows }: { data: any; filteredRows: any[] }
{/* Detail Panel - Fixed on Right */}
{selectedItem && (
<div className="fixed right-0 top-0 h-screen w-[420px] bg-card border-l-2 border-primary shadow-2xl z-50 overflow-y-auto animate-in slide-in-from-right duration-300 ease-out">
<MemoryDetailPanel
memory={selectedItem}
onClose={() => setSelectedItem(null)}
inPanel
/>
<MemoryDetailPanel memory={selectedItem} onClose={() => setSelectedItem(null)} inPanel />
</div>
)}
</div>
@@ -1,18 +1,18 @@
'use client';
"use client";
import { useState, useEffect } from 'react';
import { client } from '@/lib/api';
import { useBank } from '@/lib/bank-context';
import { useState, useEffect } from "react";
import { client } from "@/lib/api";
import { useBank } from "@/lib/bank-context";
import {
Dialog,
DialogContent,
DialogHeader,
DialogTitle,
DialogDescription,
} from '@/components/ui/dialog';
} from "@/components/ui/dialog";
interface DocumentChunkModalProps {
type: 'document' | 'chunk';
type: "document" | "chunk";
id: string | null;
onClose: () => void;
}
@@ -30,9 +30,9 @@ export function DocumentChunkModal({ type, id, onClose }: DocumentChunkModalProp
setLoading(true);
setError(null);
try {
if (type === 'document') {
if (type === "document") {
if (!currentBank) {
setError('No bank selected');
setError("No bank selected");
return;
}
const doc = await client.getDocument(id, currentBank);
@@ -58,13 +58,11 @@ export function DocumentChunkModal({ type, id, onClose }: DocumentChunkModalProp
<Dialog open={isOpen} onOpenChange={(open) => !open && onClose()}>
<DialogContent className="max-w-3xl max-h-[80vh] overflow-hidden flex flex-col">
<DialogHeader>
<DialogTitle>
{type === 'document' ? 'Document Details' : 'Chunk Details'}
</DialogTitle>
<DialogTitle>{type === "document" ? "Document Details" : "Chunk Details"}</DialogTitle>
<DialogDescription>
{type === 'document'
? 'View the original document text and metadata'
: 'View the chunk text and metadata'}
{type === "document"
? "View the original document text and metadata"
: "View the chunk text and metadata"}
</DialogDescription>
</DialogHeader>
@@ -73,9 +71,7 @@ export function DocumentChunkModal({ type, id, onClose }: DocumentChunkModalProp
<div className="flex items-center justify-center py-20">
<div className="text-center">
<div className="text-4xl mb-2"></div>
<div className="text-sm text-muted-foreground">
Loading {type}...
</div>
<div className="text-sm text-muted-foreground">Loading {type}...</div>
</div>
</div>
) : error ? (
@@ -87,7 +83,7 @@ export function DocumentChunkModal({ type, id, onClose }: DocumentChunkModalProp
</div>
) : data ? (
<div className="space-y-4">
{type === 'document' ? (
{type === "document" ? (
<>
<div className="space-y-3">
<div className="p-3 bg-muted rounded-lg">
@@ -128,9 +124,7 @@ export function DocumentChunkModal({ type, id, onClose }: DocumentChunkModalProp
{data.original_text && (
<div>
<div className="text-sm font-bold text-foreground mb-2">
Original Text
</div>
<div className="text-sm font-bold text-foreground mb-2">Original Text</div>
<div className="p-4 bg-muted rounded-lg border border-border max-h-[300px] overflow-y-auto">
<pre className="text-sm whitespace-pre-wrap font-mono text-foreground">
{data.original_text}
@@ -190,9 +184,7 @@ export function DocumentChunkModal({ type, id, onClose }: DocumentChunkModalProp
{data.chunk_text && (
<div>
<div className="text-sm font-bold text-foreground mb-2">
Chunk Text
</div>
<div className="text-sm font-bold text-foreground mb-2">Chunk Text</div>
<div className="p-4 bg-muted rounded-lg border border-border max-h-[300px] overflow-y-auto">
<pre className="text-sm whitespace-pre-wrap font-mono text-foreground">
{data.chunk_text}
@@ -1,18 +1,25 @@
'use client';
"use client";
import { useState, useEffect } from 'react';
import { client } from '@/lib/api';
import { useBank } from '@/lib/bank-context';
import { Button } from '@/components/ui/button';
import { Input } from '@/components/ui/input';
import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from '@/components/ui/table';
import { X } from 'lucide-react';
import { useState, useEffect } from "react";
import { client } from "@/lib/api";
import { useBank } from "@/lib/bank-context";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import {
Table,
TableBody,
TableCell,
TableHead,
TableHeader,
TableRow,
} from "@/components/ui/table";
import { X } from "lucide-react";
export function DocumentsView() {
const { currentBank } = useBank();
const [documents, setDocuments] = useState<any[]>([]);
const [loading, setLoading] = useState(false);
const [searchQuery, setSearchQuery] = useState('');
const [searchQuery, setSearchQuery] = useState("");
const [total, setTotal] = useState(0);
// Document view panel state
@@ -32,8 +39,8 @@ export function DocumentsView() {
setDocuments(data.items || []);
setTotal(data.total || 0);
} catch (error) {
console.error('Error loading documents:', error);
alert('Error loading documents: ' + (error as Error).message);
console.error("Error loading documents:", error);
alert("Error loading documents: " + (error as Error).message);
} finally {
setLoading(false);
}
@@ -49,8 +56,8 @@ export function DocumentsView() {
const doc: any = await client.getDocument(documentId, currentBank);
setSelectedDocument(doc);
} catch (error) {
console.error('Error loading document:', error);
alert('Error loading document: ' + (error as Error).message);
console.error("Error loading document:", error);
alert("Error loading document: " + (error as Error).message);
setSelectedDocument(null);
} finally {
setLoadingDocument(false);
@@ -75,9 +82,7 @@ export function DocumentsView() {
</div>
</div>
) : documents.length > 0 ? (
<div className="mb-4 text-sm text-muted-foreground">
{total} total documents
</div>
<div className="mb-4 text-sm text-muted-foreground">{total} total documents</div>
) : (
<div className="flex items-center justify-center py-20">
<div className="text-center">
@@ -119,20 +124,24 @@ export function DocumentsView() {
documents.map((doc) => (
<TableRow
key={doc.id}
className={`cursor-pointer hover:bg-muted/50 ${selectedDocument?.id === doc.id ? 'bg-primary/10' : ''}`}
className={`cursor-pointer hover:bg-muted/50 ${selectedDocument?.id === doc.id ? "bg-primary/10" : ""}`}
onClick={() => viewDocumentText(doc.id)}
>
<TableCell title={doc.id} className="text-card-foreground">
{doc.id.length > 30 ? doc.id.substring(0, 30) + '...' : doc.id}
{doc.id.length > 30 ? doc.id.substring(0, 30) + "..." : doc.id}
</TableCell>
<TableCell className="text-card-foreground">
{doc.created_at ? new Date(doc.created_at).toLocaleString() : 'N/A'}
{doc.created_at ? new Date(doc.created_at).toLocaleString() : "N/A"}
</TableCell>
<TableCell className="text-card-foreground">
{doc.retain_params?.context || '-'}
{doc.retain_params?.context || "-"}
</TableCell>
<TableCell className="text-card-foreground">
{doc.text_length?.toLocaleString()} chars
</TableCell>
<TableCell className="text-card-foreground">
{doc.memory_unit_count}
</TableCell>
<TableCell className="text-card-foreground">{doc.text_length?.toLocaleString()} chars</TableCell>
<TableCell className="text-card-foreground">{doc.memory_unit_count}</TableCell>
<TableCell>
<Button
onClick={(e) => {
@@ -140,7 +149,7 @@ export function DocumentsView() {
viewDocumentText(doc.id);
}}
size="sm"
variant={selectedDocument?.id === doc.id ? 'default' : 'secondary'}
variant={selectedDocument?.id === doc.id ? "default" : "secondary"}
title="View original text"
>
View Text
@@ -168,7 +177,9 @@ export function DocumentsView() {
<div className="flex justify-between items-center mb-6 pb-4 border-b border-border">
<div>
<h3 className="text-xl font-bold text-foreground">Document Details</h3>
<p className="text-sm text-muted-foreground mt-1">Original document text and metadata</p>
<p className="text-sm text-muted-foreground mt-1">
Original document text and metadata
</p>
</div>
<Button
variant="secondary"
@@ -192,20 +203,32 @@ export function DocumentsView() {
<div className="space-y-5">
{/* Document ID */}
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">Document ID</div>
<div className="text-sm font-mono break-all text-card-foreground">{selectedDocument.id}</div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Document ID
</div>
<div className="text-sm font-mono break-all text-card-foreground">
{selectedDocument.id}
</div>
</div>
{/* Created & Memory Units */}
{selectedDocument.created_at && (
<div className="grid grid-cols-2 gap-4">
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">Created</div>
<div className="text-sm font-medium text-card-foreground">{new Date(selectedDocument.created_at).toLocaleString()}</div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Created
</div>
<div className="text-sm font-medium text-card-foreground">
{new Date(selectedDocument.created_at).toLocaleString()}
</div>
</div>
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">Memory Units</div>
<div className="text-sm font-medium text-card-foreground">{selectedDocument.memory_unit_count}</div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Memory Units
</div>
<div className="text-sm font-medium text-card-foreground">
{selectedDocument.memory_unit_count}
</div>
</div>
</div>
)}
@@ -213,26 +236,40 @@ export function DocumentsView() {
{/* Text Length */}
{selectedDocument.original_text && (
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">Text Length</div>
<div className="text-sm font-medium text-card-foreground">{selectedDocument.original_text.length.toLocaleString()} characters</div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Text Length
</div>
<div className="text-sm font-medium text-card-foreground">
{selectedDocument.original_text.length.toLocaleString()} characters
</div>
</div>
)}
{/* Retain Parameters */}
{selectedDocument.retain_params && (
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">Retain Parameters</div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Retain Parameters
</div>
<div className="text-sm space-y-2 text-card-foreground">
{selectedDocument.retain_params.context && (
<div><span className="font-semibold">Context:</span> {selectedDocument.retain_params.context}</div>
<div>
<span className="font-semibold">Context:</span>{" "}
{selectedDocument.retain_params.context}
</div>
)}
{selectedDocument.retain_params.event_date && (
<div><span className="font-semibold">Event Date:</span> {new Date(selectedDocument.retain_params.event_date).toLocaleString()}</div>
<div>
<span className="font-semibold">Event Date:</span>{" "}
{new Date(selectedDocument.retain_params.event_date).toLocaleString()}
</div>
)}
{selectedDocument.retain_params.metadata && (
<div className="mt-2">
<span className="font-semibold">Metadata:</span>
<pre className="mt-1 text-xs bg-background p-2 rounded text-card-foreground">{JSON.stringify(selectedDocument.retain_params.metadata, null, 2)}</pre>
<pre className="mt-1 text-xs bg-background p-2 rounded text-card-foreground">
{JSON.stringify(selectedDocument.retain_params.metadata, null, 2)}
</pre>
</div>
)}
</div>
@@ -242,9 +279,13 @@ export function DocumentsView() {
{/* Original Text */}
{selectedDocument.original_text && (
<div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">Original Text</div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Original Text
</div>
<div className="p-4 bg-muted/50 rounded-lg border border-border max-h-[400px] overflow-y-auto">
<pre className="text-sm whitespace-pre-wrap font-mono leading-relaxed text-card-foreground">{selectedDocument.original_text}</pre>
<pre className="text-sm whitespace-pre-wrap font-mono leading-relaxed text-card-foreground">
{selectedDocument.original_text}
</pre>
</div>
</div>
)}
@@ -1,10 +1,17 @@
'use client';
"use client";
import { useState, useEffect } from 'react';
import { client } from '@/lib/api';
import { useBank } from '@/lib/bank-context';
import { Button } from '@/components/ui/button';
import { Table, TableBody, TableCell, TableHead, TableHeader, TableRow } from '@/components/ui/table';
import { useState, useEffect } from "react";
import { client } from "@/lib/api";
import { useBank } from "@/lib/bank-context";
import { Button } from "@/components/ui/button";
import {
Table,
TableBody,
TableCell,
TableHead,
TableHeader,
TableRow,
} from "@/components/ui/table";
interface Entity {
id: string;
@@ -41,8 +48,8 @@ export function EntitiesView() {
});
setEntities(result.items || []);
} catch (error) {
console.error('Error loading entities:', error);
alert('Error loading entities: ' + (error as Error).message);
console.error("Error loading entities:", error);
alert("Error loading entities: " + (error as Error).message);
} finally {
setLoading(false);
}
@@ -56,8 +63,8 @@ export function EntitiesView() {
const result: any = await client.getEntity(entityId, currentBank);
setSelectedEntity(result);
} catch (error) {
console.error('Error loading entity detail:', error);
alert('Error loading entity detail: ' + (error as Error).message);
console.error("Error loading entity detail:", error);
alert("Error loading entity detail: " + (error as Error).message);
} finally {
setLoadingDetail(false);
}
@@ -72,8 +79,8 @@ export function EntitiesView() {
// Reload entity detail to show new observations
await loadEntityDetail(selectedEntity.id);
} catch (error) {
console.error('Error regenerating observations:', error);
alert('Error regenerating observations: ' + (error as Error).message);
console.error("Error regenerating observations:", error);
alert("Error regenerating observations: " + (error as Error).message);
} finally {
setRegenerating(false);
}
@@ -87,7 +94,7 @@ export function EntitiesView() {
}, [currentBank]);
const formatDate = (dateStr?: string) => {
if (!dateStr) return 'N/A';
if (!dateStr) return "N/A";
return new Date(dateStr).toLocaleDateString();
};
@@ -104,44 +111,50 @@ export function EntitiesView() {
</div>
) : entities.length > 0 ? (
<>
<div className="mb-4 text-sm text-muted-foreground">
{entities.length} entities
</div>
<div className="overflow-x-auto">
<Table>
<TableHeader>
<TableRow>
<TableHead>Name</TableHead>
<TableHead>Mentions</TableHead>
<TableHead>First Seen</TableHead>
<TableHead>Last Seen</TableHead>
</TableRow>
</TableHeader>
<TableBody>
{entities.map((entity) => (
<TableRow
key={entity.id}
onClick={() => loadEntityDetail(entity.id)}
className={`cursor-pointer hover:bg-muted/50 ${
selectedEntity?.id === entity.id ? 'bg-primary/10' : ''
}`}
>
<TableCell className="font-medium text-card-foreground">{entity.canonical_name}</TableCell>
<TableCell className="text-card-foreground">{entity.mention_count}</TableCell>
<TableCell className="text-card-foreground">{formatDate(entity.first_seen)}</TableCell>
<TableCell className="text-card-foreground">{formatDate(entity.last_seen)}</TableCell>
<div className="mb-4 text-sm text-muted-foreground">{entities.length} entities</div>
<div className="overflow-x-auto">
<Table>
<TableHeader>
<TableRow>
<TableHead>Name</TableHead>
<TableHead>Mentions</TableHead>
<TableHead>First Seen</TableHead>
<TableHead>Last Seen</TableHead>
</TableRow>
))}
</TableBody>
</Table>
</div>
</TableHeader>
<TableBody>
{entities.map((entity) => (
<TableRow
key={entity.id}
onClick={() => loadEntityDetail(entity.id)}
className={`cursor-pointer hover:bg-muted/50 ${
selectedEntity?.id === entity.id ? "bg-primary/10" : ""
}`}
>
<TableCell className="font-medium text-card-foreground">
{entity.canonical_name}
</TableCell>
<TableCell className="text-card-foreground">{entity.mention_count}</TableCell>
<TableCell className="text-card-foreground">
{formatDate(entity.first_seen)}
</TableCell>
<TableCell className="text-card-foreground">
{formatDate(entity.last_seen)}
</TableCell>
</TableRow>
))}
</TableBody>
</Table>
</div>
</>
) : (
<div className="flex items-center justify-center py-20">
<div className="text-center">
<div className="text-4xl mb-2">👥</div>
<div className="text-sm text-muted-foreground">No entities found</div>
<div className="text-xs text-muted-foreground mt-1">Entities are extracted from facts when memories are added.</div>
<div className="text-xs text-muted-foreground mt-1">
Entities are extracted from facts when memories are added.
</div>
</div>
</div>
)}
@@ -154,7 +167,9 @@ export function EntitiesView() {
{/* Header */}
<div className="flex justify-between items-center mb-6 pb-4 border-b border-border">
<div>
<h3 className="text-xl font-bold text-card-foreground">{selectedEntity.canonical_name}</h3>
<h3 className="text-xl font-bold text-card-foreground">
{selectedEntity.canonical_name}
</h3>
<p className="text-sm text-muted-foreground mt-1">Entity details</p>
</div>
<Button
@@ -171,32 +186,46 @@ export function EntitiesView() {
{/* Entity Info */}
<div className="grid grid-cols-2 gap-4">
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">Mentions</div>
<div className="text-lg font-semibold text-card-foreground">{selectedEntity.mention_count}</div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Mentions
</div>
<div className="text-lg font-semibold text-card-foreground">
{selectedEntity.mention_count}
</div>
</div>
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">First Seen</div>
<div className="text-sm font-medium text-card-foreground">{formatDate(selectedEntity.first_seen)}</div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
First Seen
</div>
<div className="text-sm font-medium text-card-foreground">
{formatDate(selectedEntity.first_seen)}
</div>
</div>
</div>
{/* ID */}
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">Entity ID</div>
<code className="text-xs font-mono break-all text-muted-foreground">{selectedEntity.id}</code>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Entity ID
</div>
<code className="text-xs font-mono break-all text-muted-foreground">
{selectedEntity.id}
</code>
</div>
{/* Observations */}
<div>
<div className="flex justify-between items-center mb-3">
<div className="text-xs font-bold text-muted-foreground uppercase">Observations</div>
<div className="text-xs font-bold text-muted-foreground uppercase">
Observations
</div>
<Button
onClick={regenerateObservations}
disabled={regenerating}
variant="outline"
size="sm"
>
{regenerating ? 'Regenerating...' : 'Regenerate'}
{regenerating ? "Regenerating..." : "Regenerate"}
</Button>
</div>
@@ -217,7 +246,8 @@ export function EntitiesView() {
</ul>
) : (
<div className="text-muted-foreground text-sm p-4 bg-muted/50 rounded-lg">
No observations yet. Click &quot;Regenerate&quot; to generate observations from facts.
No observations yet. Click &quot;Regenerate&quot; to generate observations from
facts.
</div>
)}
</div>
@@ -1,7 +1,7 @@
'use client';
"use client";
import { useRef, useEffect, useState, useMemo } from 'react';
import cytoscape, { Core, NodeSingular } from 'cytoscape';
import { useRef, useEffect, useState, useMemo } from "react";
import cytoscape, { Core, NodeSingular } from "cytoscape";
// Hook to detect dark mode
function useIsDarkMode() {
@@ -9,14 +9,14 @@ function useIsDarkMode() {
useEffect(() => {
const checkDark = () => {
setIsDark(document.documentElement.classList.contains('dark'));
setIsDark(document.documentElement.classList.contains("dark"));
};
checkDark();
// Watch for theme changes
const observer = new MutationObserver(checkDark);
observer.observe(document.documentElement, { attributes: true, attributeFilter: ['class'] });
observer.observe(document.documentElement, { attributes: true, attributeFilter: ["class"] });
return () => observer.disconnect();
}, []);
@@ -71,11 +71,11 @@ export interface Graph2DProps {
// ============================================================================
// Brand colors
const BRAND_PRIMARY = '#0074d9';
const BRAND_TEAL = '#009296';
const LINK_SEMANTIC = '#0074d9'; // Primary blue for semantic
const LINK_TEMPORAL = '#009296'; // Teal for temporal
const LINK_ENTITY = '#f59e0b'; // Amber for entity
const BRAND_PRIMARY = "#0074d9";
const BRAND_TEAL = "#009296";
const LINK_SEMANTIC = "#0074d9"; // Primary blue for semantic
const LINK_TEMPORAL = "#009296"; // Teal for temporal
const LINK_ENTITY = "#f59e0b"; // Amber for entity
const DEFAULT_NODE_COLOR = BRAND_PRIMARY;
const DEFAULT_LINK_COLOR = LINK_SEMANTIC;
@@ -128,20 +128,20 @@ export function Graph2D({
}
// Show ALL links between visible nodes (no random link limiting)
const nodeIds = new Set(nodes.map(n => n.id));
const links = data.links.filter(l => nodeIds.has(l.source) && nodeIds.has(l.target));
const nodeIds = new Set(nodes.map((n) => n.id));
const links = data.links.filter((l) => nodeIds.has(l.source) && nodeIds.has(l.target));
return { nodes, links };
}, [data, maxNodes]);
// Convert to Cytoscape format
const cyElements = useMemo(() => {
const nodes = graphData.nodes.map(node => ({
const nodes = graphData.nodes.map((node) => ({
data: {
id: node.id,
label: node.label || node.id.substring(0, 8),
color: nodeColorFn ? nodeColorFn(node) : (node.color || DEFAULT_NODE_COLOR),
size: nodeSizeFn ? nodeSizeFn(node) : (node.size || DEFAULT_NODE_SIZE),
color: nodeColorFn ? nodeColorFn(node) : node.color || DEFAULT_NODE_COLOR,
size: nodeSizeFn ? nodeSizeFn(node) : node.size || DEFAULT_NODE_SIZE,
originalNode: node,
},
}));
@@ -151,8 +151,8 @@ export function Graph2D({
id: `edge-${idx}`,
source: link.source,
target: link.target,
color: linkColorFn ? linkColorFn(link) : (link.color || DEFAULT_LINK_COLOR),
width: linkWidthFn ? linkWidthFn(link) : (link.width || DEFAULT_LINK_WIDTH),
color: linkColorFn ? linkColorFn(link) : link.color || DEFAULT_LINK_COLOR,
width: linkWidthFn ? linkWidthFn(link) : link.width || DEFAULT_LINK_WIDTH,
type: link.type,
entity: link.entity,
weight: link.weight,
@@ -176,98 +176,98 @@ export function Graph2D({
setIsLoading(true);
// Theme-aware colors
const textColor = isDarkMode ? '#ffffff' : '#1f2937';
const textBgColor = isDarkMode ? 'rgba(0,0,0,0.8)' : 'rgba(255,255,255,0.9)';
const borderColor = isDarkMode ? '#ffffff' : '#374151';
const textColor = isDarkMode ? "#ffffff" : "#1f2937";
const textBgColor = isDarkMode ? "rgba(0,0,0,0.8)" : "rgba(255,255,255,0.9)";
const borderColor = isDarkMode ? "#ffffff" : "#374151";
const cy = cytoscape({
container: containerRef.current,
elements: cyElements,
style: [
{
selector: 'node',
selector: "node",
style: {
'background-fill': 'radial-gradient',
'background-gradient-stop-colors': ['#0074d9', '#005bb5'],
'background-gradient-stop-positions': ['0%', '100%'],
'width': 'data(size)',
'height': 'data(size)',
'label': showLabels ? 'data(label)' : '',
'color': textColor,
'text-valign': 'bottom',
'text-halign': 'center',
'font-size': '8px',
'font-weight': 500,
'text-margin-y': 3,
'text-wrap': 'wrap',
'text-max-width': '80px',
'text-background-color': textBgColor,
'text-background-opacity': 0.9,
'text-background-padding': '2px',
'text-background-shape': 'roundrectangle',
'border-width': 0,
'z-index': 0,
"background-fill": "radial-gradient",
"background-gradient-stop-colors": ["#0074d9", "#005bb5"],
"background-gradient-stop-positions": ["0%", "100%"],
width: "data(size)",
height: "data(size)",
label: showLabels ? "data(label)" : "",
color: textColor,
"text-valign": "bottom",
"text-halign": "center",
"font-size": "8px",
"font-weight": 500,
"text-margin-y": 3,
"text-wrap": "wrap",
"text-max-width": "80px",
"text-background-color": textBgColor,
"text-background-opacity": 0.9,
"text-background-padding": "2px",
"text-background-shape": "roundrectangle",
"border-width": 0,
"z-index": 0,
},
},
{
selector: 'node:selected',
selector: "node:selected",
style: {
'border-width': 3,
'border-color': '#0074d9',
'border-opacity': 1,
"border-width": 3,
"border-color": "#0074d9",
"border-opacity": 1,
},
},
{
selector: 'node:active',
selector: "node:active",
style: {
'overlay-opacity': 0,
"overlay-opacity": 0,
},
},
{
selector: 'edge',
selector: "edge",
style: {
'width': 'data(width)',
'line-color': 'data(color)',
'target-arrow-color': 'data(color)',
'curve-style': 'bezier',
'opacity': isDarkMode ? 0.5 : 0.6,
'z-index': 1,
width: "data(width)",
"line-color": "data(color)",
"target-arrow-color": "data(color)",
"curve-style": "bezier",
opacity: isDarkMode ? 0.5 : 0.6,
"z-index": 1,
},
},
{
selector: 'edge:selected',
selector: "edge:selected",
style: {
'opacity': 1,
'width': 3,
opacity: 1,
width: 3,
},
},
// Dimmed state for non-selected elements
{
selector: '.dimmed',
selector: ".dimmed",
style: {
'opacity': 0.15,
opacity: 0.15,
},
},
// Highlighted state for selected node and neighbors
{
selector: 'node.highlighted',
selector: "node.highlighted",
style: {
'opacity': 1,
'border-width': 3,
'border-color': '#0074d9',
'border-opacity': 1,
opacity: 1,
"border-width": 3,
"border-color": "#0074d9",
"border-opacity": 1,
},
},
{
selector: 'edge.highlighted',
selector: "edge.highlighted",
style: {
'opacity': 0.9,
'width': 2,
opacity: 0.9,
width: 2,
},
},
],
layout: {
name: 'cose',
name: "cose",
animate: false,
randomize: true,
nodeRepulsion: () => 100000,
@@ -290,9 +290,9 @@ export function Graph2D({
cyRef.current = cy;
// Event handlers
cy.on('tap', 'node', (evt) => {
cy.on("tap", "node", (evt) => {
const node = evt.target as NodeSingular;
const originalNode = node.data('originalNode') as GraphNode;
const originalNode = node.data("originalNode") as GraphNode;
if (onNodeClickRef.current && originalNode) {
onNodeClickRef.current(originalNode);
}
@@ -303,31 +303,33 @@ export function Graph2D({
// Find all links connected to this node from full data
const connectedLinks = fullData.links.filter(
l => l.source === clickedNodeId || l.target === clickedNodeId
(l) => l.source === clickedNodeId || l.target === clickedNodeId
);
// Find all connected node IDs
const connectedNodeIds = new Set<string>();
connectedLinks.forEach(l => {
connectedLinks.forEach((l) => {
connectedNodeIds.add(l.source);
connectedNodeIds.add(l.target);
});
// Add any missing nodes to the graph
const existingNodeIds = new Set(cy.nodes().map(n => n.id()));
const existingNodeIds = new Set(cy.nodes().map((n) => n.id()));
const nodesToAdd: any[] = [];
const edgesToAdd: any[] = [];
connectedNodeIds.forEach(nodeId => {
connectedNodeIds.forEach((nodeId) => {
if (!existingNodeIds.has(nodeId)) {
const nodeData = fullData.nodes.find(n => n.id === nodeId);
const nodeData = fullData.nodes.find((n) => n.id === nodeId);
if (nodeData) {
nodesToAdd.push({
group: 'nodes',
group: "nodes",
data: {
id: nodeData.id,
label: nodeData.label || nodeData.id.substring(0, 8),
color: nodeColorFnRef.current ? nodeColorFnRef.current(nodeData) : (nodeData.color || DEFAULT_NODE_COLOR),
color: nodeColorFnRef.current
? nodeColorFnRef.current(nodeData)
: nodeData.color || DEFAULT_NODE_COLOR,
size: nodeData.size || DEFAULT_NODE_SIZE,
originalNode: nodeData,
isTemporary: true, // Mark as temporarily added
@@ -338,18 +340,22 @@ export function Graph2D({
});
// Add missing edges
const existingEdgeIds = new Set(cy.edges().map(e => `${e.data('source')}-${e.data('target')}`));
const existingEdgeIds = new Set(
cy.edges().map((e) => `${e.data("source")}-${e.data("target")}`)
);
connectedLinks.forEach((link, idx) => {
const edgeKey = `${link.source}-${link.target}`;
const reverseKey = `${link.target}-${link.source}`;
if (!existingEdgeIds.has(edgeKey) && !existingEdgeIds.has(reverseKey)) {
edgesToAdd.push({
group: 'edges',
group: "edges",
data: {
id: `temp-edge-${idx}-${Date.now()}`,
source: link.source,
target: link.target,
color: linkColorFnRef.current ? linkColorFnRef.current(link) : (link.color || DEFAULT_LINK_COLOR),
color: linkColorFnRef.current
? linkColorFnRef.current(link)
: link.color || DEFAULT_LINK_COLOR,
width: link.width || DEFAULT_LINK_WIDTH,
type: link.type,
isTemporary: true,
@@ -364,7 +370,7 @@ export function Graph2D({
// Position new nodes near the clicked node
const clickedPos = node.position();
cy.nodes('[?isTemporary]').forEach((n, i) => {
cy.nodes("[?isTemporary]").forEach((n, i) => {
const angle = (2 * Math.PI * i) / nodesToAdd.length;
const radius = 150;
n.position({
@@ -378,71 +384,77 @@ export function Graph2D({
const neighborhood = node.neighborhood().add(node);
// Dim all elements first
cy.elements().addClass('dimmed');
cy.elements().addClass("dimmed");
// Highlight the neighborhood
neighborhood.removeClass('dimmed');
neighborhood.addClass('highlighted');
neighborhood.removeClass("dimmed");
neighborhood.addClass("highlighted");
// Center on the neighborhood without changing positions
cy.animate({
fit: { eles: neighborhood, padding: 50 },
}, { duration: 400 });
cy.animate(
{
fit: { eles: neighborhood, padding: 50 },
},
{ duration: 400 }
);
});
// Click on background to reset
cy.on('tap', (evt) => {
cy.on("tap", (evt) => {
if (evt.target === cy) {
// Remove temporary nodes and edges
cy.elements('[?isTemporary]').remove();
cy.elements("[?isTemporary]").remove();
cy.elements().removeClass('dimmed highlighted');
cy.animate({
fit: { eles: cy.elements(), padding: 50 },
}, { duration: 400 });
cy.elements().removeClass("dimmed highlighted");
cy.animate(
{
fit: { eles: cy.elements(), padding: 50 },
},
{ duration: 400 }
);
}
});
cy.on('mouseover', 'node', (evt) => {
cy.on("mouseover", "node", (evt) => {
const node = evt.target as NodeSingular;
const originalNode = node.data('originalNode') as GraphNode;
const originalNode = node.data("originalNode") as GraphNode;
setHoveredNode(originalNode);
if (onNodeHoverRef.current && originalNode) {
onNodeHoverRef.current(originalNode);
}
containerRef.current!.style.cursor = 'pointer';
containerRef.current!.style.cursor = "pointer";
});
cy.on('mouseout', 'node', () => {
cy.on("mouseout", "node", () => {
setHoveredNode(null);
if (onNodeHoverRef.current) {
onNodeHoverRef.current(null);
}
containerRef.current!.style.cursor = 'default';
containerRef.current!.style.cursor = "default";
});
// Edge hover handlers
cy.on('mouseover', 'edge', (evt) => {
cy.on("mouseover", "edge", (evt) => {
const edge = evt.target;
const originalLink = edge.data('originalLink') as GraphLink;
const originalLink = edge.data("originalLink") as GraphLink;
if (originalLink) {
setHoveredLink(originalLink);
// Get position for tooltip
const renderedPos = edge.renderedMidpoint();
setLinkTooltipPos({ x: renderedPos.x, y: renderedPos.y });
}
containerRef.current!.style.cursor = 'pointer';
containerRef.current!.style.cursor = "pointer";
});
cy.on('mouseout', 'edge', () => {
cy.on("mouseout", "edge", () => {
setHoveredLink(null);
setLinkTooltipPos(null);
containerRef.current!.style.cursor = 'default';
containerRef.current!.style.cursor = "default";
});
// Run layout
cy.layout({
name: 'cose',
name: "cose",
animate: false,
randomize: true,
nodeRepulsion: () => 100000,
@@ -476,12 +488,15 @@ export function Graph2D({
}
};
window.addEventListener('resize', handleResize);
return () => window.removeEventListener('resize', handleResize);
window.addEventListener("resize", handleResize);
return () => window.removeEventListener("resize", handleResize);
}, []);
return (
<div className="relative w-full rounded-lg overflow-hidden border border-border" style={{ height }}>
<div
className="relative w-full rounded-lg overflow-hidden border border-border"
style={{ height }}
>
{/* Loading state */}
{isLoading && (
<div className="absolute inset-0 flex items-center justify-center bg-background z-10">
@@ -498,10 +513,10 @@ export function Graph2D({
className="w-full h-full"
style={{
background: isDarkMode
? 'radial-gradient(circle at 1px 1px, rgba(255,255,255,0.08) 1px, transparent 0)'
: 'radial-gradient(circle at 1px 1px, rgba(0,0,0,0.06) 1px, transparent 0)',
backgroundSize: '20px 20px',
backgroundColor: isDarkMode ? '#0f1419' : '#f8fafc',
? "radial-gradient(circle at 1px 1px, rgba(255,255,255,0.08) 1px, transparent 0)"
: "radial-gradient(circle at 1px 1px, rgba(0,0,0,0.06) 1px, transparent 0)",
backgroundSize: "20px 20px",
backgroundColor: isDarkMode ? "#0f1419" : "#f8fafc",
}}
/>
@@ -521,17 +536,21 @@ export function Graph2D({
style={{
left: linkTooltipPos.x,
top: linkTooltipPos.y,
transform: 'translate(-50%, -100%) translateY(-8px)',
transform: "translate(-50%, -100%) translateY(-8px)",
}}
>
<div className={`px-3 py-2 rounded-lg shadow-lg text-sm ${
isDarkMode ? 'bg-gray-800 text-white' : 'bg-white text-gray-900 border border-gray-200'
}`}>
<div
className={`px-3 py-2 rounded-lg shadow-lg text-sm ${
isDarkMode
? "bg-gray-800 text-white"
: "bg-white text-gray-900 border border-gray-200"
}`}
>
<div className="font-medium capitalize mb-1">
{(() => {
const type = hoveredLink.type || 'semantic';
if (['causes', 'caused_by', 'enables', 'prevents'].includes(type)) {
return `Causal (${type.replace('_', ' ')})`;
const type = hoveredLink.type || "semantic";
if (["causes", "caused_by", "enables", "prevents"].includes(type)) {
return `Causal (${type.replace("_", " ")})`;
}
return `${type} link`;
})()}
@@ -564,15 +583,26 @@ export function Graph2D({
export function convertHindsightGraphData(hindsightData: {
nodes?: Array<{ data: { id: string; label?: string; color?: string } }>;
edges?: Array<{ data: { source: string; target: string; color?: string; lineStyle?: string; linkType?: string; entityName?: string; weight?: number; similarity?: number } }>;
edges?: Array<{
data: {
source: string;
target: string;
color?: string;
lineStyle?: string;
linkType?: string;
entityName?: string;
weight?: number;
similarity?: number;
};
}>;
table_rows?: Array<{ id: string; text: string; entities?: string; context?: string }>;
}): GraphData {
const nodes: GraphNode[] = (hindsightData.nodes || []).map(n => {
const tableRow = hindsightData.table_rows?.find(r => r.id === n.data.id);
const nodes: GraphNode[] = (hindsightData.nodes || []).map((n) => {
const tableRow = hindsightData.table_rows?.find((r) => r.id === n.data.id);
// Use memory text as label, truncated to ~40 chars
let label = n.data.label;
if (!label && tableRow?.text) {
label = tableRow.text.length > 40 ? tableRow.text.substring(0, 40) + '...' : tableRow.text;
label = tableRow.text.length > 40 ? tableRow.text.substring(0, 40) + "..." : tableRow.text;
}
if (!label) {
label = n.data.id.substring(0, 8);
@@ -585,13 +615,13 @@ export function convertHindsightGraphData(hindsightData: {
};
});
const links: GraphLink[] = (hindsightData.edges || []).map(e => ({
const links: GraphLink[] = (hindsightData.edges || []).map((e) => ({
source: e.data.source,
target: e.data.target,
color: e.data.color,
// Use linkType directly from API, fallback to lineStyle check, default to semantic
type: e.data.linkType || (e.data.lineStyle === 'dashed' ? 'temporal' : 'semantic'),
entity: e.data.entityName, // API returns entityName
type: e.data.linkType || (e.data.lineStyle === "dashed" ? "temporal" : "semantic"),
entity: e.data.entityName, // API returns entityName
weight: e.data.weight ?? e.data.similarity,
}));
@@ -1,9 +1,9 @@
'use client';
"use client";
import { useState } from 'react';
import { Button } from '@/components/ui/button';
import { Copy, Check, X } from 'lucide-react';
import { DocumentChunkModal } from './document-chunk-modal';
import { useState } from "react";
import { Button } from "@/components/ui/button";
import { Copy, Check, X } from "lucide-react";
import { DocumentChunkModal } from "./document-chunk-modal";
interface MemoryDetailPanelProps {
memory: any;
@@ -19,7 +19,7 @@ export function MemoryDetailPanel({
inPanel = false,
}: MemoryDetailPanelProps) {
const [copiedId, setCopiedId] = useState<string | null>(null);
const [modalType, setModalType] = useState<'document' | 'chunk' | null>(null);
const [modalType, setModalType] = useState<"document" | "chunk" | null>(null);
const [modalId, setModalId] = useState<string | null>(null);
const copyToClipboard = async (text: string) => {
@@ -28,17 +28,17 @@ export function MemoryDetailPanel({
setCopiedId(text);
setTimeout(() => setCopiedId(null), 2000);
} catch (err) {
console.error('Failed to copy:', err);
console.error("Failed to copy:", err);
}
};
const openDocumentModal = (docId: string) => {
setModalType('document');
setModalType("document");
setModalId(docId);
};
const openChunkModal = (chunkId: string) => {
setModalType('chunk');
setModalType("chunk");
setModalId(chunkId);
};
@@ -52,8 +52,8 @@ export function MemoryDetailPanel({
// Handle both 'id' and 'node_id' (trace results use node_id)
const memoryId = memory.id || memory.node_id;
const labelSize = compact ? 'text-[10px]' : 'text-xs';
const textSize = compact ? 'text-xs' : 'text-sm';
const labelSize = compact ? "text-[10px]" : "text-xs";
const textSize = compact ? "text-xs" : "text-sm";
// Panel mode: no outer border/bg, larger padding, prominent close button
if (inPanel) {
@@ -66,12 +66,7 @@ export function MemoryDetailPanel({
<h3 className="text-xl font-bold text-foreground">Memory Details</h3>
<p className="text-sm text-muted-foreground mt-1">Full memory content and metadata</p>
</div>
<Button
variant="secondary"
size="sm"
onClick={onClose}
className="h-8 w-8 p-0"
>
<Button variant="secondary" size="sm" onClick={onClose} className="h-8 w-8 p-0">
<X className="h-5 w-5" />
</Button>
</div>
@@ -79,14 +74,20 @@ export function MemoryDetailPanel({
<div className="space-y-5">
{/* Full Text */}
<div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">Full Text</div>
<div className="text-sm whitespace-pre-wrap leading-relaxed text-foreground">{memory.text}</div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Full Text
</div>
<div className="text-sm whitespace-pre-wrap leading-relaxed text-foreground">
{memory.text}
</div>
</div>
{/* Context */}
{memory.context && (
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">Context</div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Context
</div>
<div className="text-sm text-foreground">{memory.context}</div>
</div>
)}
@@ -94,19 +95,19 @@ export function MemoryDetailPanel({
{/* Dates */}
<div className="grid grid-cols-2 gap-4">
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">Occurred</div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Occurred
</div>
<div className="text-sm font-medium text-foreground">
{memory.occurred_start
? new Date(memory.occurred_start).toLocaleString()
: 'N/A'}
{memory.occurred_start ? new Date(memory.occurred_start).toLocaleString() : "N/A"}
</div>
</div>
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">Mentioned</div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Mentioned
</div>
<div className="text-sm font-medium text-foreground">
{memory.mentioned_at
? new Date(memory.mentioned_at).toLocaleString()
: 'N/A'}
{memory.mentioned_at ? new Date(memory.mentioned_at).toLocaleString() : "N/A"}
</div>
</div>
</div>
@@ -114,10 +115,16 @@ export function MemoryDetailPanel({
{/* Entities */}
{memory.entities && (
<div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-3">Entities</div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-3">
Entities
</div>
<div className="flex flex-wrap gap-2">
{(Array.isArray(memory.entities) ? memory.entities : String(memory.entities).split(', ')).map((entity: any, i: number) => {
const entityText = typeof entity === 'string' ? entity : (entity?.name || JSON.stringify(entity));
{(Array.isArray(memory.entities)
? memory.entities
: String(memory.entities).split(", ")
).map((entity: any, i: number) => {
const entityText =
typeof entity === "string" ? entity : entity?.name || JSON.stringify(entity);
return (
<span
key={i}
@@ -134,9 +141,13 @@ export function MemoryDetailPanel({
{/* ID */}
{memoryId && (
<div className="p-4 bg-muted/50 rounded-lg">
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">Memory ID</div>
<div className="text-xs font-bold text-muted-foreground uppercase mb-2">
Memory ID
</div>
<div className="flex items-center gap-2">
<code className="text-xs font-mono break-all flex-1 text-muted-foreground">{memoryId}</code>
<code className="text-xs font-mono break-all flex-1 text-muted-foreground">
{memoryId}
</code>
<Button
variant="ghost"
size="sm"
@@ -181,24 +192,22 @@ export function MemoryDetailPanel({
{/* Document/Chunk Modal */}
{modalType && modalId && (
<DocumentChunkModal
type={modalType}
id={modalId}
onClose={closeModal}
/>
<DocumentChunkModal type={modalType} id={modalId} onClose={closeModal} />
)}
</>
);
}
// Original compact/default mode
const padding = compact ? 'p-3' : 'p-4';
const titleSize = compact ? 'text-sm' : 'text-lg';
const gap = compact ? 'space-y-2' : 'space-y-4';
const padding = compact ? "p-3" : "p-4";
const titleSize = compact ? "text-sm" : "text-lg";
const gap = compact ? "space-y-2" : "space-y-4";
return (
<>
<div className={`bg-card border-2 border-primary rounded-lg ${padding} sticky top-4 max-h-[calc(100vh-120px)] overflow-y-auto`}>
<div
className={`bg-card border-2 border-primary rounded-lg ${padding} sticky top-4 max-h-[calc(100vh-120px)] overflow-y-auto`}
>
<div className="flex justify-between items-start mb-4">
<div>
<h3 className={`${titleSize} font-bold text-card-foreground`}>Memory Details</h3>
@@ -210,58 +219,68 @@ export function MemoryDetailPanel({
variant="ghost"
size="sm"
onClick={onClose}
className={compact ? 'h-6 w-6 p-0' : 'h-8 w-8 p-0'}
className={compact ? "h-6 w-6 p-0" : "h-8 w-8 p-0"}
>
<X className={compact ? 'h-3 w-3' : 'h-4 w-4'} />
<X className={compact ? "h-3 w-3" : "h-4 w-4"} />
</Button>
</div>
<div className={gap}>
{/* Full Text */}
<div className={`${compact ? 'p-2' : 'p-3'} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>Full Text</div>
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Full Text
</div>
<div className={`${textSize} whitespace-pre-wrap`}>{memory.text}</div>
</div>
{/* Context */}
{memory.context && (
<div className={`${compact ? 'p-2' : 'p-3'} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>Context</div>
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Context
</div>
<div className={textSize}>{memory.context}</div>
</div>
)}
{/* Dates */}
<div className="grid grid-cols-2 gap-2">
<div className={`${compact ? 'p-2' : 'p-3'} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>Occurred</div>
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Occurred
</div>
<div className={textSize}>
{memory.occurred_start
? new Date(memory.occurred_start).toLocaleString()
: 'N/A'}
{memory.occurred_start ? new Date(memory.occurred_start).toLocaleString() : "N/A"}
</div>
</div>
<div className={`${compact ? 'p-2' : 'p-3'} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>Mentioned</div>
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Mentioned
</div>
<div className={textSize}>
{memory.mentioned_at
? new Date(memory.mentioned_at).toLocaleString()
: 'N/A'}
{memory.mentioned_at ? new Date(memory.mentioned_at).toLocaleString() : "N/A"}
</div>
</div>
</div>
{/* Entities */}
{memory.entities && (
<div className={`${compact ? 'p-2' : 'p-3'} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-2`}>Entities</div>
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-2`}>
Entities
</div>
<div className="flex flex-wrap gap-1">
{(Array.isArray(memory.entities) ? memory.entities : String(memory.entities).split(', ')).map((entity: any, i: number) => {
const entityText = typeof entity === 'string' ? entity : (entity?.name || JSON.stringify(entity));
{(Array.isArray(memory.entities)
? memory.entities
: String(memory.entities).split(", ")
).map((entity: any, i: number) => {
const entityText =
typeof entity === "string" ? entity : entity?.name || JSON.stringify(entity);
return (
<span
key={i}
className={`${compact ? 'text-[10px] px-1.5 py-0.5' : 'text-xs px-2 py-1'} rounded bg-secondary text-secondary-foreground`}
className={`${compact ? "text-[10px] px-1.5 py-0.5" : "text-xs px-2 py-1"} rounded bg-secondary text-secondary-foreground`}
>
{entityText}
</span>
@@ -273,10 +292,14 @@ export function MemoryDetailPanel({
{/* ID */}
{memoryId && (
<div className={`${compact ? 'p-2' : 'p-3'} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>Memory ID</div>
<div className={`${compact ? "p-2" : "p-3"} bg-muted rounded-lg`}>
<div className={`${labelSize} font-bold text-muted-foreground uppercase mb-1`}>
Memory ID
</div>
<div className="flex items-center gap-2">
<span className={`${compact ? 'text-[10px]' : 'text-sm'} font-mono break-all`}>{memoryId}</span>
<span className={`${compact ? "text-[10px]" : "text-sm"} font-mono break-all`}>
{memoryId}
</span>
<Button
variant="ghost"
size="sm"
@@ -295,13 +318,13 @@ export function MemoryDetailPanel({
{/* Document/Chunk buttons */}
{(memory.document_id || memory.chunk_id) && (
<div className={`flex gap-2 ${compact ? 'pt-1' : ''}`}>
<div className={`flex gap-2 ${compact ? "pt-1" : ""}`}>
{memory.document_id && (
<Button
onClick={() => openDocumentModal(memory.document_id)}
size="sm"
variant="secondary"
className={`flex-1 ${compact ? 'h-7 text-xs' : ''}`}
className={`flex-1 ${compact ? "h-7 text-xs" : ""}`}
>
View Document
</Button>
@@ -311,7 +334,7 @@ export function MemoryDetailPanel({
onClick={() => openChunkModal(memory.chunk_id)}
size="sm"
variant="secondary"
className={`flex-1 ${compact ? 'h-7 text-xs' : ''}`}
className={`flex-1 ${compact ? "h-7 text-xs" : ""}`}
>
View Chunk
</Button>
@@ -323,11 +346,7 @@ export function MemoryDetailPanel({
{/* Document/Chunk Modal */}
{modalType && modalId && (
<DocumentChunkModal
type={modalType}
id={modalId}
onClose={closeModal}
/>
<DocumentChunkModal type={modalType} id={modalId} onClose={closeModal} />
)}
</>
);
File diff suppressed because it is too large Load Diff
@@ -1,12 +1,21 @@
'use client';
"use client";
import { useState } from 'react';
import { useBank } from '@/lib/bank-context';
import { Search, Sparkles, Database, FileText, Users, ChevronLeft, ChevronRight, Box } from 'lucide-react';
import { cn } from '@/lib/utils';
import Link from 'next/link';
import { useState } from "react";
import { useBank } from "@/lib/bank-context";
import {
Search,
Sparkles,
Database,
FileText,
Users,
ChevronLeft,
ChevronRight,
Box,
} from "lucide-react";
import { cn } from "@/lib/utils";
import Link from "next/link";
type NavItem = 'recall' | 'reflect' | 'data' | 'documents' | 'entities' | 'profile';
type NavItem = "recall" | "reflect" | "data" | "documents" | "entities" | "profile";
interface SidebarProps {
currentTab: NavItem;
@@ -22,19 +31,21 @@ export function Sidebar({ currentTab, onTabChange }: SidebarProps) {
}
const navItems = [
{ id: 'data' as NavItem, label: 'Memories', icon: Database },
{ id: 'recall' as NavItem, label: 'Recall', icon: Search },
{ id: 'reflect' as NavItem, label: 'Reflect', icon: Sparkles },
{ id: 'documents' as NavItem, label: 'Documents', icon: FileText },
{ id: 'entities' as NavItem, label: 'Entities', icon: Users },
{ id: 'profile' as NavItem, label: 'Memory Bank', icon: Box },
{ id: "data" as NavItem, label: "Memories", icon: Database },
{ id: "recall" as NavItem, label: "Recall", icon: Search },
{ id: "reflect" as NavItem, label: "Reflect", icon: Sparkles },
{ id: "documents" as NavItem, label: "Documents", icon: FileText },
{ id: "entities" as NavItem, label: "Entities", icon: Users },
{ id: "profile" as NavItem, label: "Memory Bank", icon: Box },
];
return (
<aside className={cn(
'bg-card border-r border-border flex flex-col transition-all duration-300',
isCollapsed ? 'w-16' : 'w-64'
)}>
<aside
className={cn(
"bg-card border-r border-border flex flex-col transition-all duration-300",
isCollapsed ? "w-16" : "w-64"
)}
>
<nav className="flex-1 p-3 pt-4">
<ul className="space-y-1">
{navItems.map((item) => {
@@ -56,11 +67,11 @@ export function Sidebar({ currentTab, onTabChange }: SidebarProps) {
// Middle-click or Ctrl/Cmd+click will naturally open in new tab
}}
className={cn(
'w-full flex items-center gap-3 px-4 py-3 rounded-lg text-sm font-medium transition-all',
"w-full flex items-center gap-3 px-4 py-3 rounded-lg text-sm font-medium transition-all",
isActive
? 'bg-primary-gradient text-white shadow-sm'
: 'text-muted-foreground hover:bg-accent hover:text-accent-foreground',
isCollapsed && 'justify-center px-0'
? "bg-primary-gradient text-white shadow-sm"
: "text-muted-foreground hover:bg-accent hover:text-accent-foreground",
isCollapsed && "justify-center px-0"
)}
title={isCollapsed ? item.label : undefined}
>
@@ -78,10 +89,10 @@ export function Sidebar({ currentTab, onTabChange }: SidebarProps) {
<button
onClick={() => setIsCollapsed(!isCollapsed)}
className={cn(
'w-full flex items-center gap-3 px-4 py-2 rounded-lg text-sm text-muted-foreground hover:bg-accent hover:text-accent-foreground transition-colors',
isCollapsed && 'justify-center px-0'
"w-full flex items-center gap-3 px-4 py-2 rounded-lg text-sm text-muted-foreground hover:bg-accent hover:text-accent-foreground transition-colors",
isCollapsed && "justify-center px-0"
)}
title={isCollapsed ? 'Expand sidebar' : 'Collapse sidebar'}
title={isCollapsed ? "Expand sidebar" : "Collapse sidebar"}
>
{isCollapsed ? (
<ChevronRight className="w-5 h-5" />
@@ -1,23 +1,29 @@
'use client';
"use client";
import { useState } from 'react';
import { client } from '@/lib/api';
import { useBank } from '@/lib/bank-context';
import { Button } from '@/components/ui/button';
import { Input } from '@/components/ui/input';
import { Textarea } from '@/components/ui/textarea';
import { Select, SelectContent, SelectItem, SelectTrigger, SelectValue } from '@/components/ui/select';
import { Checkbox } from '@/components/ui/checkbox';
import { Card, CardContent, CardDescription, CardHeader, CardTitle } from '@/components/ui/card';
import { Sparkles, Info } from 'lucide-react';
import JsonView from 'react18-json-view';
import 'react18-json-view/src/style.css';
import { useState } from "react";
import { client } from "@/lib/api";
import { useBank } from "@/lib/bank-context";
import { Button } from "@/components/ui/button";
import { Input } from "@/components/ui/input";
import { Textarea } from "@/components/ui/textarea";
import {
Select,
SelectContent,
SelectItem,
SelectTrigger,
SelectValue,
} from "@/components/ui/select";
import { Checkbox } from "@/components/ui/checkbox";
import { Card, CardContent, CardDescription, CardHeader, CardTitle } from "@/components/ui/card";
import { Sparkles, Info } from "lucide-react";
import JsonView from "react18-json-view";
import "react18-json-view/src/style.css";
export function ThinkView() {
const { currentBank } = useBank();
const [query, setQuery] = useState('');
const [context, setContext] = useState('');
const [budget, setBudget] = useState<'low' | 'mid' | 'high'>('mid');
const [query, setQuery] = useState("");
const [context, setContext] = useState("");
const [budget, setBudget] = useState<"low" | "mid" | "high">("mid");
const [includeFacts, setIncludeFacts] = useState(true);
const [showRawJson, setShowRawJson] = useState(false);
const [result, setResult] = useState<any>(null);
@@ -38,8 +44,8 @@ export function ThinkView() {
});
setResult(data);
} catch (error) {
console.error('Error running reflect:', error);
alert('Error running reflect: ' + (error as Error).message);
console.error("Error running reflect:", error);
alert("Error running reflect: " + (error as Error).message);
} finally {
setLoading(false);
}
@@ -57,7 +63,7 @@ export function ThinkView() {
value={query}
onChange={(e) => setQuery(e.target.value)}
placeholder="Enter your question..."
onKeyDown={(e) => e.key === 'Enter' && runReflect()}
onKeyDown={(e) => e.key === "Enter" && runReflect()}
/>
</div>
<div>
@@ -83,10 +89,7 @@ export function ThinkView() {
Include Facts
</label>
</div>
<Button
onClick={runReflect}
disabled={loading || !query}
>
<Button onClick={runReflect} disabled={loading || !query}>
<Sparkles className="w-4 h-4 mr-2" />
Reflect
</Button>
@@ -134,14 +137,14 @@ export function ThinkView() {
</div>
<div className="flex gap-2">
<Button
variant={!showRawJson ? 'default' : 'outline'}
variant={!showRawJson ? "default" : "outline"}
size="sm"
onClick={() => setShowRawJson(false)}
>
Based On
</Button>
<Button
variant={showRawJson ? 'default' : 'outline'}
variant={showRawJson ? "default" : "outline"}
size="sm"
onClick={() => setShowRawJson(true)}
>
@@ -155,9 +158,11 @@ export function ThinkView() {
includeFacts && result.based_on && result.based_on.length > 0 ? (
(() => {
// Group facts by type
const worldFacts = result.based_on.filter((f: any) => f.type === 'world');
const experienceFacts = result.based_on.filter((f: any) => f.type === 'experience');
const opinionFacts = result.based_on.filter((f: any) => f.type === 'opinion');
const worldFacts = result.based_on.filter((f: any) => f.type === "world");
const experienceFacts = result.based_on.filter(
(f: any) => f.type === "experience"
);
const opinionFacts = result.based_on.filter((f: any) => f.type === "opinion");
return (
<div className="grid grid-cols-1 md:grid-cols-3 gap-4">
@@ -172,7 +177,11 @@ export function ThinkView() {
{worldFacts.map((fact: any, i: number) => (
<li key={i} className="p-2 bg-muted rounded">
{fact.text}
{fact.context && <div className="text-xs text-muted-foreground mt-1">{fact.context}</div>}
{fact.context && (
<div className="text-xs text-muted-foreground mt-1">
{fact.context}
</div>
)}
</li>
))}
</ul>
@@ -185,7 +194,9 @@ export function ThinkView() {
<Card>
<CardHeader className="pb-3">
<CardTitle className="text-base">Experience</CardTitle>
<CardDescription className="text-xs">Conversations & Events</CardDescription>
<CardDescription className="text-xs">
Conversations & Events
</CardDescription>
</CardHeader>
<CardContent>
{experienceFacts.length > 0 ? (
@@ -193,7 +204,11 @@ export function ThinkView() {
{experienceFacts.map((fact: any, i: number) => (
<li key={i} className="p-2 bg-muted rounded">
{fact.text}
{fact.context && <div className="text-xs text-muted-foreground mt-1">{fact.context}</div>}
{fact.context && (
<div className="text-xs text-muted-foreground mt-1">
{fact.context}
</div>
)}
</li>
))}
</ul>
@@ -206,7 +221,9 @@ export function ThinkView() {
<Card>
<CardHeader className="pb-3">
<CardTitle className="text-base">Opinions</CardTitle>
<CardDescription className="text-xs">Beliefs & Preferences</CardDescription>
<CardDescription className="text-xs">
Beliefs & Preferences
</CardDescription>
</CardHeader>
<CardContent>
{opinionFacts.length > 0 ? (
@@ -214,7 +231,11 @@ export function ThinkView() {
{opinionFacts.map((fact: any, i: number) => (
<li key={i} className="p-2 bg-muted rounded">
{fact.text}
{fact.context && <div className="text-xs text-muted-foreground mt-1">{fact.context}</div>}
{fact.context && (
<div className="text-xs text-muted-foreground mt-1">
{fact.context}
</div>
)}
</li>
))}
</ul>
@@ -230,7 +251,9 @@ export function ThinkView() {
<div className="flex items-start gap-3 p-4 bg-amber-50 dark:bg-amber-950 border border-amber-200 dark:border-amber-800 rounded-lg">
<Info className="w-5 h-5 text-amber-600 dark:text-amber-400 mt-0.5 flex-shrink-0" />
<div>
<p className="font-semibold text-amber-900 dark:text-amber-100">No facts found</p>
<p className="font-semibold text-amber-900 dark:text-amber-100">
No facts found
</p>
<p className="text-sm text-amber-700 dark:text-amber-300 mt-1">
No memories were found or used to generate this answer.
</p>
@@ -240,20 +263,19 @@ export function ThinkView() {
<div className="flex items-start gap-3 p-4 bg-amber-50 dark:bg-amber-950 border border-amber-200 dark:border-amber-800 rounded-lg">
<Info className="w-5 h-5 text-amber-600 dark:text-amber-400 mt-0.5 flex-shrink-0" />
<div>
<p className="font-semibold text-amber-900 dark:text-amber-100">Facts not included</p>
<p className="font-semibold text-amber-900 dark:text-amber-100">
Facts not included
</p>
<p className="text-sm text-amber-700 dark:text-amber-300 mt-1">
Enable "Include Facts" above to see which memories were used to generate this answer.
Enable "Include Facts" above to see which memories were used to generate
this answer.
</p>
</div>
</div>
)
) : (
<div className="bg-muted p-4 rounded border border-border overflow-auto max-h-[600px]">
<JsonView
src={result}
collapsed={1}
theme="default"
/>
<JsonView src={result} collapsed={1} theme="default" />
</div>
)}
</CardContent>
@@ -273,7 +295,9 @@ export function ThinkView() {
{result.new_opinions.map((opinion: any, i: number) => (
<div key={i} className="p-3 bg-muted rounded-lg border border-border">
<div className="font-semibold text-foreground">{opinion.text}</div>
<div className="text-sm text-muted-foreground mt-1">Confidence: {opinion.confidence?.toFixed(2)}</div>
<div className="text-sm text-muted-foreground mt-1">
Confidence: {opinion.confidence?.toFixed(2)}
</div>
</div>
))}
</div>
@@ -1,8 +1,8 @@
import * as React from "react"
import { Slot } from "@radix-ui/react-slot"
import { cva, type VariantProps } from "class-variance-authority"
import * as React from "react";
import { Slot } from "@radix-ui/react-slot";
import { cva, type VariantProps } from "class-variance-authority";
import { cn } from "@/lib/utils"
import { cn } from "@/lib/utils";
const buttonVariants = cva(
"inline-flex items-center justify-center gap-2 whitespace-nowrap rounded-md text-sm font-medium ring-offset-background transition-colors focus-visible:outline-none focus-visible:ring-2 focus-visible:ring-ring focus-visible:ring-offset-2 disabled:pointer-events-none disabled:opacity-50 [&_svg]:pointer-events-none [&_svg]:size-4 [&_svg]:shrink-0",
@@ -10,12 +10,9 @@ const buttonVariants = cva(
variants: {
variant: {
default: "bg-primary-gradient text-white hover:opacity-90",
destructive:
"bg-destructive text-destructive-foreground hover:bg-destructive/90",
outline:
"border border-input bg-background hover:bg-accent hover:text-accent-foreground",
secondary:
"bg-secondary text-secondary-foreground hover:bg-secondary/80",
destructive: "bg-destructive text-destructive-foreground hover:bg-destructive/90",
outline: "border border-input bg-background hover:bg-accent hover:text-accent-foreground",
secondary: "bg-secondary text-secondary-foreground hover:bg-secondary/80",
ghost: "hover:bg-accent hover:text-accent-foreground",
link: "text-primary underline-offset-4 hover:underline",
},
@@ -31,26 +28,21 @@ const buttonVariants = cva(
size: "default",
},
}
)
);
export interface ButtonProps
extends React.ButtonHTMLAttributes<HTMLButtonElement>,
VariantProps<typeof buttonVariants> {
asChild?: boolean
extends React.ButtonHTMLAttributes<HTMLButtonElement>, VariantProps<typeof buttonVariants> {
asChild?: boolean;
}
const Button = React.forwardRef<HTMLButtonElement, ButtonProps>(
({ className, variant, size, asChild = false, ...props }, ref) => {
const Comp = asChild ? Slot : "button"
const Comp = asChild ? Slot : "button";
return (
<Comp
className={cn(buttonVariants({ variant, size, className }))}
ref={ref}
{...props}
/>
)
<Comp className={cn(buttonVariants({ variant, size, className }))} ref={ref} {...props} />
);
}
)
Button.displayName = "Button"
);
Button.displayName = "Button";
export { Button, buttonVariants }
export { Button, buttonVariants };
@@ -1,79 +1,55 @@
import * as React from "react"
import * as React from "react";
import { cn } from "@/lib/utils"
import { cn } from "@/lib/utils";
const Card = React.forwardRef<
HTMLDivElement,
React.HTMLAttributes<HTMLDivElement>
>(({ className, ...props }, ref) => (
<div
ref={ref}
className={cn(
"rounded-lg border bg-card text-card-foreground shadow-sm",
className
)}
{...props}
/>
))
Card.displayName = "Card"
const Card = React.forwardRef<HTMLDivElement, React.HTMLAttributes<HTMLDivElement>>(
({ className, ...props }, ref) => (
<div
ref={ref}
className={cn("rounded-lg border bg-card text-card-foreground shadow-sm", className)}
{...props}
/>
)
);
Card.displayName = "Card";
const CardHeader = React.forwardRef<
HTMLDivElement,
React.HTMLAttributes<HTMLDivElement>
>(({ className, ...props }, ref) => (
<div
ref={ref}
className={cn("flex flex-col space-y-1.5 p-6", className)}
{...props}
/>
))
CardHeader.displayName = "CardHeader"
const CardHeader = React.forwardRef<HTMLDivElement, React.HTMLAttributes<HTMLDivElement>>(
({ className, ...props }, ref) => (
<div ref={ref} className={cn("flex flex-col space-y-1.5 p-6", className)} {...props} />
)
);
CardHeader.displayName = "CardHeader";
const CardTitle = React.forwardRef<
HTMLDivElement,
React.HTMLAttributes<HTMLDivElement>
>(({ className, ...props }, ref) => (
<div
ref={ref}
className={cn(
"text-2xl font-semibold leading-none tracking-tight",
className
)}
{...props}
/>
))
CardTitle.displayName = "CardTitle"
const CardTitle = React.forwardRef<HTMLDivElement, React.HTMLAttributes<HTMLDivElement>>(
({ className, ...props }, ref) => (
<div
ref={ref}
className={cn("text-2xl font-semibold leading-none tracking-tight", className)}
{...props}
/>
)
);
CardTitle.displayName = "CardTitle";
const CardDescription = React.forwardRef<
HTMLDivElement,
React.HTMLAttributes<HTMLDivElement>
>(({ className, ...props }, ref) => (
<div
ref={ref}
className={cn("text-sm text-muted-foreground", className)}
{...props}
/>
))
CardDescription.displayName = "CardDescription"
const CardDescription = React.forwardRef<HTMLDivElement, React.HTMLAttributes<HTMLDivElement>>(
({ className, ...props }, ref) => (
<div ref={ref} className={cn("text-sm text-muted-foreground", className)} {...props} />
)
);
CardDescription.displayName = "CardDescription";
const CardContent = React.forwardRef<
HTMLDivElement,
React.HTMLAttributes<HTMLDivElement>
>(({ className, ...props }, ref) => (
<div ref={ref} className={cn("p-6 pt-0", className)} {...props} />
))
CardContent.displayName = "CardContent"
const CardContent = React.forwardRef<HTMLDivElement, React.HTMLAttributes<HTMLDivElement>>(
({ className, ...props }, ref) => (
<div ref={ref} className={cn("p-6 pt-0", className)} {...props} />
)
);
CardContent.displayName = "CardContent";
const CardFooter = React.forwardRef<
HTMLDivElement,
React.HTMLAttributes<HTMLDivElement>
>(({ className, ...props }, ref) => (
<div
ref={ref}
className={cn("flex items-center p-6 pt-0", className)}
{...props}
/>
))
CardFooter.displayName = "CardFooter"
const CardFooter = React.forwardRef<HTMLDivElement, React.HTMLAttributes<HTMLDivElement>>(
({ className, ...props }, ref) => (
<div ref={ref} className={cn("flex items-center p-6 pt-0", className)} {...props} />
)
);
CardFooter.displayName = "CardFooter";
export { Card, CardHeader, CardFooter, CardTitle, CardDescription, CardContent }
export { Card, CardHeader, CardFooter, CardTitle, CardDescription, CardContent };
@@ -1,10 +1,10 @@
"use client"
"use client";
import * as React from "react"
import * as CheckboxPrimitive from "@radix-ui/react-checkbox"
import { Check } from "lucide-react"
import * as React from "react";
import * as CheckboxPrimitive from "@radix-ui/react-checkbox";
import { Check } from "lucide-react";
import { cn } from "@/lib/utils"
import { cn } from "@/lib/utils";
const Checkbox = React.forwardRef<
React.ElementRef<typeof CheckboxPrimitive.Root>,
@@ -18,13 +18,11 @@ const Checkbox = React.forwardRef<
)}
{...props}
>
<CheckboxPrimitive.Indicator
className={cn("grid place-content-center text-current")}
>
<CheckboxPrimitive.Indicator className={cn("grid place-content-center text-current")}>
<Check className="h-4 w-4" />
</CheckboxPrimitive.Indicator>
</CheckboxPrimitive.Root>
))
Checkbox.displayName = CheckboxPrimitive.Root.displayName
));
Checkbox.displayName = CheckboxPrimitive.Root.displayName;
export { Checkbox }
export { Checkbox };
@@ -1,12 +1,12 @@
"use client"
"use client";
import * as React from "react"
import { type DialogProps } from "@radix-ui/react-dialog"
import { Command as CommandPrimitive } from "cmdk"
import { Search } from "lucide-react"
import * as React from "react";
import { type DialogProps } from "@radix-ui/react-dialog";
import { Command as CommandPrimitive } from "cmdk";
import { Search } from "lucide-react";
import { cn } from "@/lib/utils"
import { Dialog, DialogContent } from "@/components/ui/dialog"
import { cn } from "@/lib/utils";
import { Dialog, DialogContent } from "@/components/ui/dialog";
const Command = React.forwardRef<
React.ElementRef<typeof CommandPrimitive>,
@@ -20,8 +20,8 @@ const Command = React.forwardRef<
)}
{...props}
/>
))
Command.displayName = CommandPrimitive.displayName
));
Command.displayName = CommandPrimitive.displayName;
const CommandDialog = ({ children, ...props }: DialogProps) => {
return (
@@ -32,8 +32,8 @@ const CommandDialog = ({ children, ...props }: DialogProps) => {
</Command>
</DialogContent>
</Dialog>
)
}
);
};
const CommandInput = React.forwardRef<
React.ElementRef<typeof CommandPrimitive.Input>,
@@ -50,9 +50,9 @@ const CommandInput = React.forwardRef<
{...props}
/>
</div>
))
));
CommandInput.displayName = CommandPrimitive.Input.displayName
CommandInput.displayName = CommandPrimitive.Input.displayName;
const CommandList = React.forwardRef<
React.ElementRef<typeof CommandPrimitive.List>,
@@ -63,22 +63,18 @@ const CommandList = React.forwardRef<
className={cn("max-h-[300px] overflow-y-auto overflow-x-hidden", className)}
{...props}
/>
))
));
CommandList.displayName = CommandPrimitive.List.displayName
CommandList.displayName = CommandPrimitive.List.displayName;
const CommandEmpty = React.forwardRef<
React.ElementRef<typeof CommandPrimitive.Empty>,
React.ComponentPropsWithoutRef<typeof CommandPrimitive.Empty>
>((props, ref) => (
<CommandPrimitive.Empty
ref={ref}
className="py-6 text-center text-sm"
{...props}
/>
))
<CommandPrimitive.Empty ref={ref} className="py-6 text-center text-sm" {...props} />
));
CommandEmpty.displayName = CommandPrimitive.Empty.displayName
CommandEmpty.displayName = CommandPrimitive.Empty.displayName;
const CommandGroup = React.forwardRef<
React.ElementRef<typeof CommandPrimitive.Group>,
@@ -92,9 +88,9 @@ const CommandGroup = React.forwardRef<
)}
{...props}
/>
))
));
CommandGroup.displayName = CommandPrimitive.Group.displayName
CommandGroup.displayName = CommandPrimitive.Group.displayName;
const CommandSeparator = React.forwardRef<
React.ElementRef<typeof CommandPrimitive.Separator>,
@@ -105,8 +101,8 @@ const CommandSeparator = React.forwardRef<
className={cn("-mx-1 h-px bg-border", className)}
{...props}
/>
))
CommandSeparator.displayName = CommandPrimitive.Separator.displayName
));
CommandSeparator.displayName = CommandPrimitive.Separator.displayName;
const CommandItem = React.forwardRef<
React.ElementRef<typeof CommandPrimitive.Item>,
@@ -120,25 +116,19 @@ const CommandItem = React.forwardRef<
)}
{...props}
/>
))
));
CommandItem.displayName = CommandPrimitive.Item.displayName
CommandItem.displayName = CommandPrimitive.Item.displayName;
const CommandShortcut = ({
className,
...props
}: React.HTMLAttributes<HTMLSpanElement>) => {
const CommandShortcut = ({ className, ...props }: React.HTMLAttributes<HTMLSpanElement>) => {
return (
<span
className={cn(
"ml-auto text-xs tracking-widest text-muted-foreground",
className
)}
className={cn("ml-auto text-xs tracking-widest text-muted-foreground", className)}
{...props}
/>
)
}
CommandShortcut.displayName = "CommandShortcut"
);
};
CommandShortcut.displayName = "CommandShortcut";
export {
Command,
@@ -150,4 +140,4 @@ export {
CommandItem,
CommandShortcut,
CommandSeparator,
}
};
@@ -1,18 +1,18 @@
"use client"
"use client";
import * as React from "react"
import * as DialogPrimitive from "@radix-ui/react-dialog"
import { X } from "lucide-react"
import * as React from "react";
import * as DialogPrimitive from "@radix-ui/react-dialog";
import { X } from "lucide-react";
import { cn } from "@/lib/utils"
import { cn } from "@/lib/utils";
const Dialog = DialogPrimitive.Root
const Dialog = DialogPrimitive.Root;
const DialogTrigger = DialogPrimitive.Trigger
const DialogTrigger = DialogPrimitive.Trigger;
const DialogPortal = DialogPrimitive.Portal
const DialogPortal = DialogPrimitive.Portal;
const DialogClose = DialogPrimitive.Close
const DialogClose = DialogPrimitive.Close;
const DialogOverlay = React.forwardRef<
React.ElementRef<typeof DialogPrimitive.Overlay>,
@@ -26,8 +26,8 @@ const DialogOverlay = React.forwardRef<
)}
{...props}
/>
))
DialogOverlay.displayName = DialogPrimitive.Overlay.displayName
));
DialogOverlay.displayName = DialogPrimitive.Overlay.displayName;
const DialogContent = React.forwardRef<
React.ElementRef<typeof DialogPrimitive.Content>,
@@ -50,36 +50,21 @@ const DialogContent = React.forwardRef<
</DialogPrimitive.Close>
</DialogPrimitive.Content>
</DialogPortal>
))
DialogContent.displayName = DialogPrimitive.Content.displayName
));
DialogContent.displayName = DialogPrimitive.Content.displayName;
const DialogHeader = ({
className,
...props
}: React.HTMLAttributes<HTMLDivElement>) => (
const DialogHeader = ({ className, ...props }: React.HTMLAttributes<HTMLDivElement>) => (
<div className={cn("flex flex-col space-y-1.5 text-center sm:text-left", className)} {...props} />
);
DialogHeader.displayName = "DialogHeader";
const DialogFooter = ({ className, ...props }: React.HTMLAttributes<HTMLDivElement>) => (
<div
className={cn(
"flex flex-col space-y-1.5 text-center sm:text-left",
className
)}
className={cn("flex flex-col-reverse sm:flex-row sm:justify-end sm:space-x-2", className)}
{...props}
/>
)
DialogHeader.displayName = "DialogHeader"
const DialogFooter = ({
className,
...props
}: React.HTMLAttributes<HTMLDivElement>) => (
<div
className={cn(
"flex flex-col-reverse sm:flex-row sm:justify-end sm:space-x-2",
className
)}
{...props}
/>
)
DialogFooter.displayName = "DialogFooter"
);
DialogFooter.displayName = "DialogFooter";
const DialogTitle = React.forwardRef<
React.ElementRef<typeof DialogPrimitive.Title>,
@@ -87,14 +72,11 @@ const DialogTitle = React.forwardRef<
>(({ className, ...props }, ref) => (
<DialogPrimitive.Title
ref={ref}
className={cn(
"text-lg font-semibold leading-none tracking-tight text-foreground",
className
)}
className={cn("text-lg font-semibold leading-none tracking-tight text-foreground", className)}
{...props}
/>
))
DialogTitle.displayName = DialogPrimitive.Title.displayName
));
DialogTitle.displayName = DialogPrimitive.Title.displayName;
const DialogDescription = React.forwardRef<
React.ElementRef<typeof DialogPrimitive.Description>,
@@ -105,8 +87,8 @@ const DialogDescription = React.forwardRef<
className={cn("text-sm text-muted-foreground", className)}
{...props}
/>
))
DialogDescription.displayName = DialogPrimitive.Description.displayName
));
DialogDescription.displayName = DialogPrimitive.Description.displayName;
export {
Dialog,
@@ -119,4 +101,4 @@ export {
DialogFooter,
DialogTitle,
DialogDescription,
}
};
@@ -1,6 +1,6 @@
import * as React from "react"
import * as React from "react";
import { cn } from "@/lib/utils"
import { cn } from "@/lib/utils";
const Input = React.forwardRef<HTMLInputElement, React.ComponentProps<"input">>(
({ className, type, ...props }, ref) => {
@@ -14,9 +14,9 @@ const Input = React.forwardRef<HTMLInputElement, React.ComponentProps<"input">>(
ref={ref}
{...props}
/>
)
);
}
)
Input.displayName = "Input"
);
Input.displayName = "Input";
export { Input }
export { Input };

Some files were not shown because too many files have changed in this diff Show More