diff --git a/.env.template b/.env.template index 74e6d787..2540f84b 100644 --- a/.env.template +++ b/.env.template @@ -80,6 +80,10 @@ LLM_ANTHROPIC_API_KEY=your-anthropic-api-key-here # LLM_SUMMARY_MAX_TOKENS_SHORT=1000 # LLM_SUMMARY_MAX_TOKENS_LONG=2000 +# Embedding settings +# LLM_MAX_EMBEDDING_TOKENS=8192 +# LLM_MAX_EMBEDDING_TOKENS_PER_REQUEST=300000 + # ============================================================================= # Agent Settings # ============================================================================= diff --git a/migrations/utils.py b/migrations/utils.py index db98970d..4333ae46 100644 --- a/migrations/utils.py +++ b/migrations/utils.py @@ -1,4 +1,3 @@ - import sqlalchemy as sa from alembic import op diff --git a/migrations/versions/917195d9b5e9_add_messageembedding_table.py b/migrations/versions/917195d9b5e9_add_messageembedding_table.py new file mode 100644 index 00000000..8d8f78f4 --- /dev/null +++ b/migrations/versions/917195d9b5e9_add_messageembedding_table.py @@ -0,0 +1,116 @@ +"""add messageembedding table + +Revision ID: 917195d9b5e9 +Revises: d429de0e5338 +Create Date: 2024-01-01 12:00:00.000000 + +""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op +from pgvector.sqlalchemy import Vector + +from migrations.utils import index_exists, table_exists +from src.config import settings + +# revision identifiers, used by Alembic. +revision: str = "917195d9b5e9" +down_revision: str | None = "d429de0e5338" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None +schema = settings.DB.SCHEMA + + +def upgrade() -> None: + op.create_table( + "message_embeddings", + sa.Column("id", sa.BigInteger(), sa.Identity(), nullable=False), + sa.Column("content", sa.Text(), nullable=False), + sa.Column("embedding", Vector(1536), nullable=False), + sa.Column("message_id", sa.Text(), nullable=False), + sa.Column("workspace_name", sa.Text(), nullable=False), + sa.Column("session_name", sa.Text(), nullable=True), + sa.Column("peer_name", sa.Text(), nullable=False), + sa.Column( + "created_at", + sa.DateTime(timezone=True), + nullable=False, + server_default=sa.func.now(), + ), + # Foreign key constraints + sa.ForeignKeyConstraint(["message_id"], ["messages.public_id"]), + sa.ForeignKeyConstraint(["workspace_name"], ["workspaces.name"]), + sa.ForeignKeyConstraint( + ["session_name", "workspace_name"], + ["sessions.name", "sessions.workspace_name"], + ), + sa.ForeignKeyConstraint( + ["peer_name", "workspace_name"], ["peers.name", "peers.workspace_name"] + ), + schema=schema, + ) + + # Create indexes + op.create_index( + "idx_message_embeddings_message_id", + "message_embeddings", + ["message_id"], + schema=schema, + ) + op.create_index( + "idx_message_embeddings_workspace_name", + "message_embeddings", + ["workspace_name"], + schema=schema, + ) + op.create_index( + "idx_message_embeddings_session_name", + "message_embeddings", + ["session_name"], + schema=schema, + ) + op.create_index( + "idx_message_embeddings_peer_name", + "message_embeddings", + ["peer_name"], + schema=schema, + ) + op.create_index( + "idx_message_embeddings_created_at", + "message_embeddings", + ["created_at"], + schema=schema, + ) + + # Create HNSW index for vector similarity search + op.execute(f""" + CREATE INDEX idx_message_embeddings_embedding_hnsw + ON {schema}.message_embeddings + USING hnsw (embedding vector_cosine_ops) + WITH (m = 16, ef_construction = 64) + """) + + +def downgrade() -> None: + inspector = sa.inspect(op.get_bind()) + if not table_exists("message_embeddings", inspector): + return + + # Drop indexes defensively + indexes_to_drop = [ + "idx_message_embeddings_embedding_hnsw", + "idx_message_embeddings_message_id", + "idx_message_embeddings_workspace_name", + "idx_message_embeddings_session_name", + "idx_message_embeddings_peer_name", + "idx_message_embeddings_created_at", + ] + + for index_name in indexes_to_drop: + if index_exists("message_embeddings", index_name, inspector): + op.drop_index(index_name, table_name="message_embeddings", schema=schema) + + # Drop table (this will also drop foreign keys and check constraints) + op.drop_table("message_embeddings", schema=schema) diff --git a/scripts/generate_message_embeddings.py b/scripts/generate_message_embeddings.py new file mode 100644 index 00000000..f8afa299 --- /dev/null +++ b/scripts/generate_message_embeddings.py @@ -0,0 +1,222 @@ +""" +Script to generate embeddings for existing messages that don't already have embeddings. + +# Note: When generating embeddings for messages, we need to consider two limits defined in the settings: +# 1. MAX_EMBEDDING_TOKENS: This is the maximum number of tokens that can be included in a single message for which an embedding is generated. +# If a message exceeds this limit, it will be chunked into multiple embeddings. +# 2. MAX_EMBEDDING_TOKENS_PER_REQUEST: This is the maximum total number of tokens that can be included in a single request to the embedding provider. +# If the total number of tokens across all messages in a batch exceeds this limit, the batch will need to be split into multiple batches. + +Usage: + python scripts/generate_message_embeddings.py [--workspace-name WORKSPACE] [--session-name SESSION] [--peer-name PEER] +""" + +import argparse +import asyncio +import os +import sys + +# Add the project root to the path +project_root = os.path.abspath(os.path.join(os.path.dirname(__file__), "..")) +sys.path.insert(0, project_root) + +import tiktoken # noqa: E402 +from sqlalchemy import select # noqa: E402 +from sqlalchemy.ext.asyncio import AsyncSession # noqa: E402 + +from src import models # noqa: E402 +from src.config import settings # noqa: E402 +from src.dependencies import tracked_db # noqa: E402 +from src.embeddings import EmbeddingClient # noqa: E402 + + +async def get_messages_without_embeddings( + db: AsyncSession, + workspace_name: str | None = None, + session_name: str | None = None, + peer_name: str | None = None, +) -> list[models.Message]: + """ + Get all messages that don't have embeddings yet. + + Args: + db: Database session + workspace_name: Optional workspace name filter + session_name: Optional session name filter + peer_name: Optional peer name filter + + Returns: + List of messages without embeddings + """ + # Query messages that don't have embeddings + stmt = ( + select(models.Message) + .outerjoin( + models.MessageEmbedding, + models.Message.public_id == models.MessageEmbedding.message_id, + ) + .where(models.MessageEmbedding.message_id.is_(None)) # No embedding exists + .order_by(models.Message.id) + ) + + # Apply filters if provided + if workspace_name: + stmt = stmt.where(models.Message.workspace_name == workspace_name) + + if session_name: + stmt = stmt.where(models.Message.session_name == session_name) + + if peer_name: + stmt = stmt.where(models.Message.peer_name == peer_name) + + result = await db.execute(stmt) + return list(result.scalars().all()) + + +async def create_embeddings_for_messages( + db: AsyncSession, + messages: list[models.Message], + embedding_client: EmbeddingClient, +) -> int: + """ + Create embeddings for a batch of messages. + + Args: + db: Database session + messages: List of messages to create embeddings for + embedding_client: Embedding client instance + + Returns: + Number of embeddings created + """ + if not messages: + return 0 + + # Initialize tiktoken encoding (same as used in MessageCreate schema) + encoding = tiktoken.get_encoding("cl100k_base") + + # Prepare data for batch embedding with proper token encoding + id_resource_dict = { + message.public_id: ( + message.content, + encoding.encode(message.content), # Properly encode the content + ) + for message in messages + } + + # Generate embeddings + embedding_dict = await embedding_client.batch_embed(id_resource_dict) + + # Create MessageEmbedding objects + embedding_objects: list[models.MessageEmbedding] = [] + embeddings_created = 0 + + for message in messages: + embeddings = embedding_dict.get(message.public_id, []) + for embedding in embeddings: + embedding_obj = models.MessageEmbedding( + content=message.content, + embedding=embedding, + message_id=message.public_id, + workspace_name=message.workspace_name, + session_name=message.session_name, + peer_name=message.peer_name, + ) + embedding_objects.append(embedding_obj) + embeddings_created += 1 + + # Add to database + if embedding_objects: + db.add_all(embedding_objects) + await db.commit() + + return embeddings_created + + +async def main() -> None: + parser = argparse.ArgumentParser( + description="Generate embeddings for messages that don't already have them", + ) + + parser.add_argument( + "--batch-size", + type=int, + default=50, + help="Number of messages to process in each batch (default: 50)", + ) + parser.add_argument( + "--workspace-name", + help="Only process messages from this workspace", + ) + parser.add_argument( + "--session-name", + help="Only process messages from this session", + ) + parser.add_argument( + "--peer-name", + help="Only process messages from this peer", + ) + + args = parser.parse_args() + + # Initialize embedding client + embedding_client = EmbeddingClient(settings.LLM.OPENAI_API_KEY) + + print("Generating embeddings for messages...") + if args.workspace_name: + print(f" Filtering by workspace: {args.workspace_name}") + else: + print(" Processing all workspaces") + if args.session_name: + print(f" Filtering by session: {args.session_name}") + if args.peer_name: + print(f" Filtering by peer: {args.peer_name}") + + # Use tracked_db context manager for proper database session handling + async with tracked_db("generate_embeddings") as db: + try: + # Get messages without embeddings + print("Finding messages without embeddings...") + messages = await get_messages_without_embeddings( + db, args.workspace_name, args.session_name, args.peer_name + ) + + if not messages: + print("No messages found that need embeddings.") + return + + print(f"Found {len(messages)} messages without embeddings.") + + # Process in batches + batch_size = args.batch_size + total_embeddings = 0 + + for i in range(0, len(messages), batch_size): + batch = messages[i : i + batch_size] + batch_num = (i // batch_size) + 1 + total_batches = (len(messages) + batch_size - 1) // batch_size + + print( + f"Processing batch {batch_num}/{total_batches} ({len(batch)} messages)..." + ) + + embeddings_created = await create_embeddings_for_messages( + db, batch, embedding_client + ) + total_embeddings += embeddings_created + + print( + f" Created {embeddings_created} embeddings for batch {batch_num}" + ) + + print( + f"\nCompleted! Created {total_embeddings} embeddings for {len(messages)} messages." + ) + + except Exception as e: + print(f"Error: {e}") + sys.exit(1) + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/src/agent.py b/src/agent.py index 894813c4..e6aea57f 100644 --- a/src/agent.py +++ b/src/agent.py @@ -106,6 +106,7 @@ async def chat( session_name: str | None, queries: str | list[str], stream: bool = False, + target: str | None = None, ) -> llm.Stream | llm.CallResponse: """ Chat with the Dialectic API using on-demand user representation generation. @@ -175,9 +176,13 @@ async def chat( # Run short-term inference and long-term facts in parallel async def fetch_long_term(): async with tracked_db("chat.get_collection") as db_embed: - name = "global_representation" if session_name is None else "" + name = ( + "global_representation" + if target is None + else crud.construct_collection_name(peer_name, target) + ) collection = await crud.get_or_create_collection( - db_embed, workspace_name, peer_name, collection_name=name + db_embed, workspace_name, collection_name=name, peer_name=peer_name ) collection_name = collection.name # Extract the ID while session is active facts = await get_long_term_facts( diff --git a/src/config.py b/src/config.py index a7a7f454..ae36be99 100644 --- a/src/config.py +++ b/src/config.py @@ -203,6 +203,13 @@ class LLMSettings(HonchoSettings): # SUMMARY_SYSTEM_PROMPT_SHORT_FILE: Optional[str] = "prompts/summary_short_system.txt" # SUMMARY_SYSTEM_PROMPT_LONG_FILE: Optional[str] = "prompts/summary_long_system.txt" + # Embed all messages that are sent by peers + EMBED_MESSAGES: bool = False + MAX_EMBEDDING_TOKENS: Annotated[int, Field(default=8192, gt=0)] = 8192 + MAX_EMBEDDING_TOKENS_PER_REQUEST: Annotated[int, Field(default=300000, gt=0)] = ( + 300000 + ) + class AgentSettings(HonchoSettings): model_config = SettingsConfigDict(env_prefix="AGENT_") # pyright: ignore diff --git a/src/crud.py b/src/crud.py index b0860211..bef92afb 100644 --- a/src/crud.py +++ b/src/crud.py @@ -1,10 +1,9 @@ from collections.abc import Sequence from logging import getLogger -from typing import Any, final +from typing import Any from dotenv import load_dotenv from nanoid import generate as generate_nanoid -from openai import AsyncOpenAI from sqlalchemy import Select, cast, func, insert, select, update from sqlalchemy.dialects.postgresql import insert as pg_insert from sqlalchemy.engine import Row @@ -12,30 +11,19 @@ from sqlalchemy.ext.asyncio import AsyncSession from sqlalchemy.types import BigInteger from src.config import settings +from src.embeddings import EmbeddingClient from . import models, schemas from .exceptions import ( + DisabledException, ResourceNotFoundException, + ValidationException, ) from .utils.filter import apply_filter load_dotenv(override=True) -@final -class EmbeddingClient: - def __init__(self, api_key: str | None): - if api_key is None: - raise ValueError("API key is required") - self.client = AsyncOpenAI(api_key=api_key) - - async def embed(self, query: str) -> list[float]: - response = await self.client.embeddings.create( - input=query, model="text-embedding-3-small" - ) - return response.data[0].embedding - - embedding_client = EmbeddingClient(settings.LLM.OPENAI_API_KEY) logger = getLogger(__name__) @@ -989,11 +977,14 @@ async def search( workspace_name: str, session_name: str | None = None, peer_name: str | None = None, + semantic: bool | None = None, ) -> Select[tuple[models.Message]]: """ Search across message content using a hybrid approach: + - Uses semantic search if embed_messages is set, else fall back to full text - Uses PostgreSQL full text search for natural language queries - Falls back to exact string matching for queries with special characters + - Optionally uses semantic search with embeddings If a session or peer is provided, the search will be scoped to that session or peer. Otherwise, it will search across all messages in the workspace. @@ -1003,6 +994,10 @@ async def search( workspace_name: Name of the workspace session_name: Optional name of the session peer_name: Optional name of the peer + semantic: Optional boolean to configure semantic search: + - None: try semantic search if embed_messages is set, else fall back to full text + - True: try semantic search if embed_messages is set, else throw error + - False: use full text search Returns: List of messages that match the search query, ordered by relevance @@ -1011,58 +1006,105 @@ async def search( from sqlalchemy import func, or_ - # Check if query contains special characters that FTS might not handle well - has_special_chars = bool( - re.search(r'[~`!@#$%^&*()_+=\[\]{};\':"\\|,.<>/?-]', query) - ) - # Base query conditions base_conditions = [models.Message.workspace_name == workspace_name] - if has_special_chars: - # For queries with special characters, use exact string matching (ILIKE) - # This ensures we can find exact matches like "~special-uuid~" - search_condition = models.Message.content.ilike(f"%{query}%") + should_use_semantic_search = False # Default to full text search + if semantic is None: + # Try semantic search if embed_messages is set, else fall back to full text + should_use_semantic_search = settings.LLM.EMBED_MESSAGES + elif semantic is True: + # Try semantic search if embed_messages is set, else throw error + if settings.LLM.EMBED_MESSAGES: + should_use_semantic_search = True + else: + raise DisabledException( + "Semantic search requires EMBED_MESSAGES flag to be enabled" + ) + + if should_use_semantic_search: + # Generate embedding for the search query + try: + embedding_query = await embedding_client.embed(query) + except ValueError as e: + raise ValidationException( + f"Query exceeds maximum token limit of {settings.LLM.MAX_EMBEDDING_TOKENS}." + ) from e + + # Use cosine distance for semantic search on MessageEmbedding table + # Join with Message table to get the actual message data base_query = ( select(models.Message) - .where(*base_conditions, search_condition) - .order_by(models.Message.created_at.desc()) - ) - else: - # For natural language queries, use full text search with ranking - fts_condition = func.to_tsvector("english", models.Message.content).op("@@")( - func.plainto_tsquery("english", query) - ) - - # Combine FTS with ILIKE as fallback for better coverage - combined_condition = or_( - fts_condition, models.Message.content.ilike(f"%{query}%") - ) - - base_query = ( - select(models.Message) - .where(*base_conditions, combined_condition) + .join( + models.MessageEmbedding, + models.Message.public_id == models.MessageEmbedding.message_id, + ) + .where(models.MessageEmbedding.workspace_name == workspace_name) .order_by( - # Order by FTS relevance first, then by creation time - func.coalesce( - func.ts_rank( - func.to_tsvector("english", models.Message.content), - func.plainto_tsquery("english", query), - ), - 0, - ).desc(), - models.Message.created_at.desc(), + models.MessageEmbedding.embedding.cosine_distance(embedding_query) ) ) - # Add additional filters based on parameters - if session_name is not None: - stmt = base_query.where(models.Message.session_name == session_name) - elif peer_name is not None: - stmt = base_query.where(models.Message.peer_name == peer_name) + if session_name is not None: + stmt = base_query.where( + models.MessageEmbedding.session_name == session_name + ) + elif peer_name is not None: + stmt = base_query.where(models.MessageEmbedding.peer_name == peer_name) + else: + stmt = base_query + else: - stmt = base_query + # Check if query contains special characters that FTS might not handle well + has_special_chars = bool( + re.search(r'[~`!@#$%^&*()_+=\[\]{};\':"\\|,.<>/?-]', query) + ) + + if has_special_chars: + # For queries with special characters, use exact string matching (ILIKE) + # This ensures we can find exact matches like "~special-uuid~" + search_condition = models.Message.content.ilike(f"%{query}%") + + base_query = ( + select(models.Message) + .where(*base_conditions, search_condition) + .order_by(models.Message.created_at.desc()) + ) + else: + # For natural language queries, use full text search with ranking + fts_condition = func.to_tsvector("english", models.Message.content).op( + "@@" + )(func.plainto_tsquery("english", query)) + + # Combine FTS with ILIKE as fallback for better coverage + combined_condition = or_( + fts_condition, models.Message.content.ilike(f"%{query}%") + ) + + base_query = ( + select(models.Message) + .where(*base_conditions, combined_condition) + .order_by( + # Order by FTS relevance first, then by creation time + func.coalesce( + func.ts_rank( + func.to_tsvector("english", models.Message.content), + func.plainto_tsquery("english", query), + ), + 0, + ).desc(), + models.Message.created_at.desc(), + ) + ) + + # Add additional filters based on parameters + if session_name is not None: + stmt = base_query.where(models.Message.session_name == session_name) + elif peer_name is not None: + stmt = base_query.where(models.Message.peer_name == peer_name) + else: + stmt = base_query return stmt @@ -1184,19 +1226,55 @@ async def create_messages( ) # Create list of message objects (this will trigger the before_insert event) - message_objects = [ - models.Message( + message_objects: list[models.Message] = [] + for message in messages: + message_obj = models.Message( session_name=session_name, peer_name=message.peer_name, content=message.content, h_metadata=message.metadata or {}, workspace_name=workspace_name, + public_id=generate_nanoid(), + token_count=len(message.encoded_message), ) - for message in messages - ] + message_objects.append(message_obj) - # Add all messages and commit db.add_all(message_objects) + await db.flush() + + if settings.LLM.EMBED_MESSAGES: + encoded_message_lookup = { + msg.public_id: orig_msg.encoded_message + for msg, orig_msg in zip(message_objects, messages, strict=True) + } + id_resource_dict = { + message.public_id: ( + message.content, + encoded_message_lookup[message.public_id], + ) + for message in message_objects + } + embedding_dict = await embedding_client.batch_embed(id_resource_dict) + + # Create MessageEmbedding entries for each embedded message + embedding_objects: list[models.MessageEmbedding] = [] + for message_obj in message_objects: + embeddings = embedding_dict.get(message_obj.public_id, []) + for embedding in embeddings: + embedding_obj = models.MessageEmbedding( + content=message_obj.content, + embedding=embedding, + message_id=message_obj.public_id, + workspace_name=workspace_name, + session_name=session_name, + peer_name=message_obj.peer_name, + ) + embedding_objects.append(embedding_obj) + + # Add all embedding objects to the session + if embedding_objects: + db.add_all(embedding_objects) + await db.commit() return message_objects @@ -1227,19 +1305,55 @@ async def create_messages_for_peer( db, workspace_name=workspace_name, peers=[schemas.PeerCreate(name=peer_name)] ) # Create list of message objects (this will trigger the before_insert event) - message_objects = [ - models.Message( + message_objects: list[models.Message] = [] + + for message in messages: + message_obj = models.Message( session_name=None, peer_name=peer_name, content=message.content, h_metadata=message.metadata or {}, workspace_name=workspace_name, + public_id=generate_nanoid(), + token_count=len(message.encoded_message), ) - for message in messages - ] + message_objects.append(message_obj) - # Add all messages and commit db.add_all(message_objects) + await db.flush() + + if settings.LLM.EMBED_MESSAGES: + encoded_message_lookup = { + msg.public_id: orig_msg.encoded_message + for msg, orig_msg in zip(message_objects, messages, strict=True) + } + id_resource_dict = { + message.public_id: ( + message.content, + encoded_message_lookup[message.public_id], + ) + for message in message_objects + } + embedding_dict = await embedding_client.batch_embed(id_resource_dict) + + # Create MessageEmbedding entries for each embedded message + embedding_objects: list[models.MessageEmbedding] = [] + for message_obj in message_objects: + embeddings = embedding_dict.get(message_obj.public_id, []) + for embedding in embeddings: + embedding_obj = models.MessageEmbedding( + content=message_obj.content, + embedding=embedding, + message_id=message_obj.public_id, + workspace_name=workspace_name, + peer_name=peer_name, + ) + embedding_objects.append(embedding_obj) + + # Add all embedding objects to the session + if embedding_objects: + db.add_all(embedding_objects) + await db.commit() return message_objects @@ -1450,7 +1564,10 @@ async def update_message( async def get_collection( - db: AsyncSession, workspace_name: str, peer_name: str, collection_name: str + db: AsyncSession, + workspace_name: str, + collection_name: str, + peer_name: str | None = None, ) -> models.Collection: """ Get a collection by name for a specific peer and workspace. @@ -1470,9 +1587,10 @@ async def get_collection( stmt = ( select(models.Collection) .where(models.Collection.workspace_name == workspace_name) - .where(models.Collection.peer_name == peer_name) .where(models.Collection.name == collection_name) ) + if peer_name: + stmt = stmt.where(models.Collection.peer_name == peer_name) result = await db.execute(stmt) collection = result.scalar_one_or_none() if collection is None: @@ -1485,12 +1603,12 @@ async def get_collection( async def get_or_create_collection( db: AsyncSession, workspace_name: str, - peer_name: str, collection_name: str, + peer_name: str | None = None, ) -> models.Collection: try: honcho_collection = await get_collection( - db, workspace_name, peer_name, collection_name + db, workspace_name, collection_name, peer_name ) return honcho_collection except ResourceNotFoundException: @@ -1520,7 +1638,13 @@ async def query_documents( top_k: int = 5, ) -> Sequence[models.Document]: # Using ModelClient for embeddings - embedding_query = await embedding_client.embed(query) + try: + embedding_query = await embedding_client.embed(query) + except ValueError as e: + raise ValidationException( + f"Query exceeds maximum token limit of {settings.LLM.MAX_EMBEDDING_TOKENS}." + ) from e + stmt = ( select(models.Document) .where(models.Document.workspace_name == workspace_name) @@ -1570,8 +1694,8 @@ async def create_document( await get_collection( db, workspace_name=workspace_name, - peer_name=peer_name, collection_name=collection_name, + peer_name=peer_name, ) # Using ModelClient for embeddings diff --git a/src/deriver/consumer.py b/src/deriver/consumer.py index c249a984..40e6b8f7 100644 --- a/src/deriver/consumer.py +++ b/src/deriver/consumer.py @@ -155,7 +155,7 @@ async def process_message( else "global_representation" ) collection = await crud.get_or_create_collection( - db, workspace_name, peer_name, collection_name + db, workspace_name, collection_name, peer_name ) embedding_store = CollectionEmbeddingStore( workspace_name=workspace_name, diff --git a/src/embeddings.py b/src/embeddings.py new file mode 100644 index 00000000..8b084d10 --- /dev/null +++ b/src/embeddings.py @@ -0,0 +1,228 @@ +import asyncio +import logging +from collections import defaultdict +from typing import NamedTuple + +import tiktoken +from openai import AsyncOpenAI + +from .config import settings + +logger = logging.getLogger(__name__) + + +class BatchItem(NamedTuple): + """A single item in a batch with its metadata.""" + + text: str + text_id: str + chunk_index: int + + +class EmbeddingClient: + """ + Embedding client for OpenAI with chunking and batching support. + """ + + def __init__(self, api_key: str | None = None): + if api_key is None: + api_key = settings.LLM.OPENAI_API_KEY + if not api_key: + raise ValueError("API key is required") + self.client: AsyncOpenAI = AsyncOpenAI(api_key=api_key) + self.encoding: tiktoken.Encoding = tiktoken.get_encoding("cl100k_base") + self.max_embedding_tokens: int = settings.LLM.MAX_EMBEDDING_TOKENS + self.max_embedding_tokens_per_request: int = ( + settings.LLM.MAX_EMBEDDING_TOKENS_PER_REQUEST + ) + + async def embed(self, query: str) -> list[float]: + token_count = len(self.encoding.encode(query)) + + if token_count > self.max_embedding_tokens: + raise ValueError( + f"Query exceeds maximum token limit of {self.max_embedding_tokens} tokens (got {token_count} tokens)" + ) + + response = await self.client.embeddings.create( + model="text-embedding-3-small", input=query + ) + return response.data[0].embedding + + async def batch_embed( + self, id_resource_dict: dict[str, tuple[str, list[int]]] + ) -> dict[str, list[list[float]]]: + """ + Embed multiple texts, chunking long ones and batching API calls. + + Args: + id_resource_dict: Maps text IDs to (text, encoded_tokens) tuples + + Returns: + Maps text IDs to lists of embedding vectors (one per chunk) + """ + if not id_resource_dict: + return {} + + # 1. Prepare chunks for all texts if needed + text_chunks = self._prepare_chunks(id_resource_dict) + + # 2. Create batches that fit API limits (max 2048 embeddings per request, max 300,000 tokens per request) + batches = self._create_batches(text_chunks) + + # 3. Process all batches concurrently + batch_results = await asyncio.gather( + *[self._process_batch(batch) for batch in batches], + ) + + # 4. Accumulate results preserving chunk order + return self._accumulate_embeddings(batch_results) + + def _prepare_chunks( + self, id_resource_dict: dict[str, tuple[str, list[int]]] + ) -> dict[str, list[tuple[str, int]]]: + """ + Chunk texts that exceed token limits. + + Args: + id_resource_dict: Maps text IDs to (text, encoded_tokens) tuples + + Returns: + Maps text IDs to lists of (chunk_text, token_count) tuples + """ + return { + text_id: ( + _chunk_text_with_tokens( + text, encoded_tokens, self.max_embedding_tokens, self.encoding + ) + if len(encoded_tokens) > self.max_embedding_tokens + else [(text, len(encoded_tokens))] + ) + for text_id, (text, encoded_tokens) in id_resource_dict.items() + } + + def _create_batches( + self, text_chunks: dict[str, list[tuple[str, int]]] + ) -> list[list[BatchItem]]: + """ + Group chunks into batches that fit API limits. + + Args: + text_chunks: Maps text IDs to lists of (chunk_text, token_count) tuples + + Returns: + List of batches, each containing BatchItem objects + """ + batches: list[list[BatchItem]] = [] + current_batch: list[BatchItem] = [] + current_tokens = 0 + + for text_id, chunks in text_chunks.items(): + for chunk_idx, (chunk_text, chunk_tokens) in enumerate(chunks): + # Check if adding this chunk would exceed limits + would_exceed_tokens = ( + current_tokens + chunk_tokens + > self.max_embedding_tokens_per_request + ) + would_exceed_count = len(current_batch) >= 2048 # OpenAI's input limit + + if current_batch and (would_exceed_tokens or would_exceed_count): + batches.append(current_batch) + current_batch = [] + current_tokens = 0 + + current_batch.append(BatchItem(chunk_text, text_id, chunk_idx)) + current_tokens += chunk_tokens + + if current_batch: + batches.append(current_batch) + + return batches + + async def _process_batch( + self, batch: list[BatchItem] + ) -> dict[str, dict[int, list[float]]]: + """ + Process a single batch through the embeddings API. + + Args: + batch: List of BatchItem objects to embed + + Returns: + Maps text IDs to {chunk_index: embedding_vector} dictionaries + """ + try: + response = await self.client.embeddings.create( + model="text-embedding-3-small", input=[item.text for item in batch] + ) + + # Organize embeddings by text_id and chunk_index + result: dict[str, dict[int, list[float]]] = defaultdict(dict) + for item, embedding_data in zip(batch, response.data, strict=True): + result[item.text_id][item.chunk_index] = embedding_data.embedding + + return dict(result) + + except Exception: + logger.exception("Error processing batch") + raise + + def _accumulate_embeddings( + self, batch_results: list[dict[str, dict[int, list[float]]]] + ) -> dict[str, list[list[float]]]: + """ + Combine batch results into final output, preserving chunk order. + + Args: + batch_results: List of batch results from _process_batch + + Returns: + Maps text IDs to ordered lists of embedding vectors + """ + all_embeddings: dict[str, dict[int, list[float]]] = defaultdict(dict) + + # Collect all embeddings by text_id and chunk_index + for batch_result in batch_results: + for text_id, chunk_dict in batch_result.items(): + all_embeddings[text_id].update(chunk_dict) + + # Convert to ordered lists + return { + text_id: [chunk_dict[i] for i in sorted(chunk_dict.keys())] + for text_id, chunk_dict in all_embeddings.items() + } + + +def _chunk_text_with_tokens( + text: str, + encoded_tokens: list[int], + max_tokens: int, + encoding: tiktoken.Encoding, +) -> list[tuple[str, int]]: + """ + Split text into chunks that fit within token limits, with 20% overlap. + + Args: + text: Original text to chunk + encoded_tokens: Pre-encoded tokens for the text + max_tokens: Maximum tokens per chunk + encoding: Tiktoken encoding model + + Returns: + List of (chunk_text, token_count) tuples + """ + if len(encoded_tokens) <= max_tokens: + return [(text, len(encoded_tokens))] + + # Use 20% overlap for better semantic continuity + overlap_tokens = int(max_tokens * 0.2) + step_size = max_tokens - overlap_tokens + + return [ + ( + encoding.decode(encoded_tokens[i : i + max_tokens]), + min(max_tokens, len(encoded_tokens) - i), + ) + for i in range(0, len(encoded_tokens), step_size) + if i < len(encoded_tokens) # Ensure we don't create empty chunks + ] diff --git a/src/models.py b/src/models.py index 9aee9cf0..07de89a6 100644 --- a/src/models.py +++ b/src/models.py @@ -2,7 +2,6 @@ import datetime from logging import getLogger from typing import Any, final -import tiktoken from dotenv import load_dotenv from nanoid import generate as generate_nanoid from pgvector.sqlalchemy import Vector @@ -19,7 +18,6 @@ from sqlalchemy import ( Integer, Table, UniqueConstraint, - event, text, ) from sqlalchemy.dialects.postgresql import JSONB, TEXT @@ -34,23 +32,6 @@ load_dotenv(override=True) logger = getLogger(__name__) -# Initialize tiktoken encoder for token counting for message content -tokenizer = tiktoken.get_encoding("cl100k_base") - - -def count_tokens(text: str) -> int: - """Count tokens in a text string using tiktoken.""" - if not text: - return 0 - try: - return len(tokenizer.encode(text)) - except Exception as e: - # Fallback: rough estimation (4 chars per token) - logger.warning( - f"Error counting tokens for text: {text[:50]}{'...' if len(text) > 50 else ''}, using fallback (4 chars per token). Error: {str(e)}" - ) - return len(text) // 4 - # Association table for many-to-many relationship between sessions and peers session_peers_table = Table( @@ -246,10 +227,49 @@ class Message(Base): return f"Message(id={self.id}, session_name={self.session_name}, peer_name={self.peer_name}, content={self.content})" -@event.listens_for(Message, "before_insert") -def calculate_token_count_on_insert(_mapper: Any, _connection: Any, target: Message): - """Calculate token count before inserting a new message.""" - target.token_count = count_tokens(target.content) +@final +class MessageEmbedding(Base): + __tablename__: str = "message_embeddings" + + id: Mapped[int] = mapped_column( + BigInteger, Identity(), primary_key=True, autoincrement=True + ) + content: Mapped[str] = mapped_column(TEXT) + embedding: MappedColumn[Any] = mapped_column(Vector(1536)) + message_id: Mapped[str] = mapped_column( + ForeignKey("messages.public_id"), index=True + ) + workspace_name: Mapped[str] = mapped_column( + ForeignKey("workspaces.name"), index=True + ) + session_name: Mapped[str | None] = mapped_column(TEXT, index=True, nullable=True) + peer_name: Mapped[str | None] = mapped_column(TEXT, index=True) + created_at: Mapped[datetime.datetime] = mapped_column( + DateTime(timezone=True), index=True, default=func.now() + ) + + # Relationship to Message + message = relationship("Message", backref="embeddings") + + __table_args__ = ( + # Compound foreign key constraints + ForeignKeyConstraint( + ["session_name", "workspace_name"], + ["sessions.name", "sessions.workspace_name"], + ), + ForeignKeyConstraint( + ["peer_name", "workspace_name"], + ["peers.name", "peers.workspace_name"], + ), + # HNSW index on embedding column for efficient similarity search + Index( + "idx_message_embeddings_embedding_hnsw", + "embedding", + postgresql_using="hnsw", + postgresql_with={"m": 16, "ef_construction": 64}, + postgresql_ops={"embedding": "vector_cosine_ops"}, + ), + ) @final diff --git a/src/routers/peers.py b/src/routers/peers.py index 7592ae0d..dcf51f13 100644 --- a/src/routers/peers.py +++ b/src/routers/peers.py @@ -172,7 +172,12 @@ async def chat( if not options.stream: return await agent.chat( - workspace_id, peer_id, options.session_id, options.queries, options.stream + workspace_id, + peer_id, + options.session_id, + options.queries, + options.stream, + options.target, ) async def parse_stream(): @@ -183,6 +188,7 @@ async def chat( options.session_id, options.queries, stream=True, + target=options.target, ) if isinstance(stream, Stream): async for chunk, _ in stream: @@ -326,10 +332,17 @@ async def get_working_representation( async def search_peer( workspace_id: str = Path(..., description="ID of the workspace"), peer_id: str = Path(..., description="ID of the peer"), - query: str = Body(..., description="Search query"), + search: schemas.MessageSearchOptions = Body( + ..., description="Message search parameters " + ), db: AsyncSession = db, ): """Search a Peer""" - stmt = await crud.search(query, workspace_name=workspace_id, peer_name=peer_id) + stmt = await crud.search( + search.query, + workspace_name=workspace_id, + peer_name=peer_id, + semantic=search.semantic, + ) return await apaginate(db, stmt) diff --git a/src/routers/sessions.py b/src/routers/sessions.py index 720514ee..ff103445 100644 --- a/src/routers/sessions.py +++ b/src/routers/sessions.py @@ -451,12 +451,19 @@ async def get_session_context( async def search_session( workspace_id: str = Path(..., description="ID of the workspace"), session_id: str = Path(..., description="ID of the session"), - query: str = Body(..., description="Search query"), + search: schemas.MessageSearchOptions = Body( + ..., description="Message search parameters " + ), db: AsyncSession = db, ): """Search a Session""" + query, semantic = search.query, search.semantic + stmt = await crud.search( - query, workspace_name=workspace_id, session_name=session_id + query, + workspace_name=workspace_id, + session_name=session_id, + semantic=semantic, ) return await apaginate(db, stmt) diff --git a/src/schemas.py b/src/schemas.py index 7cf189d5..3e3837ad 100644 --- a/src/schemas.py +++ b/src/schemas.py @@ -1,8 +1,16 @@ # pyright: reportUnannotatedClassAttribute=false # pyright: ignore import datetime -from typing import Annotated, Any +from typing import Annotated, Any, Self -from pydantic import BaseModel, ConfigDict, Field, field_validator +import tiktoken +from pydantic import ( + BaseModel, + ConfigDict, + Field, + PrivateAttr, + field_validator, + model_validator, +) RESOURCE_NAME_PATTERN = r"^[a-zA-Z0-9_-]+$" @@ -108,6 +116,20 @@ class MessageCreate(MessageBase): peer_name: str = Field(alias="peer_id") metadata: dict[str, Any] | None = None + _encoded_message: list[int] = PrivateAttr(default=[]) + + @property + def encoded_message(self) -> list[int]: + return self._encoded_message + + @model_validator(mode="after") + def validate_and_set_token_count(self) -> Self: + encoding = tiktoken.get_encoding("cl100k_base") + encoded_message = encoding.encode(self.content) + + self._encoded_message = encoded_message + return self + class MessageGet(MessageBase): filter: dict[str, Any] | None = None @@ -215,6 +237,14 @@ class DocumentUpdate(DocumentBase): metadata: dict[str, Any] | None = None +class MessageSearchOptions(BaseModel): + query: str = Field(..., description="Search query") + semantic: bool | None = Field( + default=None, + description="Whether to explicitly use semantic search to filter the results", + ) + + class DialecticOptions(BaseModel): session_id: str | None = Field( None, description="ID of the session to scope the representation to" diff --git a/tests/conftest.py b/tests/conftest.py index 5fef5087..839427f2 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -290,9 +290,25 @@ def mock_langfuse(): @pytest.fixture(autouse=True) def mock_openai_embeddings(): """Mock OpenAI embeddings API calls for testing""" - with patch("src.crud.embedding_client.embed") as mock_create: - mock_create.return_value = [0.1] * 1536 - yield mock_create + with ( + patch("src.crud.embedding_client.embed") as mock_embed, + patch("src.crud.embedding_client.batch_embed") as mock_batch_embed, + ): + # Mock the embed method to return a fake embedding vector + mock_embed.return_value = [0.1] * 1536 + + # Mock the batch_embed method to return a dict of fake embedding vectors + # Updated to support chunking - each text_id maps to a list of embedding vectors + async def mock_batch_embed_func( + id_resource_dict: dict[str, tuple[str, list[int]]], + ) -> dict[str, list[list[float]]]: + return { + text_id: [[0.1] * 1536] for text_id in id_resource_dict + } # Single chunk per text + + mock_batch_embed.side_effect = mock_batch_embed_func + + yield {"embed": mock_embed, "batch_embed": mock_batch_embed} @pytest.fixture(autouse=True) @@ -396,7 +412,10 @@ def mock_crud_collection_operations(): from src import models async def mock_get_or_create_collection( - _: AsyncSession, workspace_name: str, peer_name: str, collection_name: str + _: AsyncSession, + workspace_name: str, + collection_name: str, + peer_name: str | None = None, ): # Create a mock collection object that doesn't require database commit mock_collection = models.Collection( diff --git a/tests/integration/test_message_embeddings.py b/tests/integration/test_message_embeddings.py new file mode 100644 index 00000000..4d3baa25 --- /dev/null +++ b/tests/integration/test_message_embeddings.py @@ -0,0 +1,394 @@ +""" +Tests for message embedding functionality. + +These tests verify that message embeddings are created, stored, and can be searched. +""" + +from typing import Any + +import pytest +from nanoid import generate as generate_nanoid +from sqlalchemy import select +from sqlalchemy.ext.asyncio import AsyncSession + +from src import models +from src.crud import create_messages, create_messages_for_peer, search +from src.models import Peer, Workspace +from src.schemas import MessageCreate + + +@pytest.mark.asyncio +async def test_message_embedding_created_when_setting_enabled( + db_session: AsyncSession, + sample_data: tuple[Workspace, Peer], + monkeypatch: pytest.MonkeyPatch, +): + """Test that MessageEmbedding is created when EMBED_MESSAGES setting is True""" + # Monkeypatch the setting to enable message embeddings + monkeypatch.setattr("src.config.settings.LLM.EMBED_MESSAGES", True) + + test_workspace, test_peer = sample_data + + # Create a test session + test_session = models.Session( + workspace_name=test_workspace.name, name=str(generate_nanoid()) + ) + db_session.add(test_session) + await db_session.commit() + + # Create a message using the CRUD function directly + test_message_content = "This is a test message for embedding" + messages = [ + MessageCreate( + content=test_message_content, + peer_id=test_peer.name, + metadata={"test": "embedding_enabled"}, + ) + ] + + created_messages = await create_messages( + db=db_session, + messages=messages, + workspace_name=test_workspace.name, + session_name=test_session.name, + ) + + assert len(created_messages) == 1 + created_message = created_messages[0] + + # Query the MessageEmbedding table to verify an embedding was created + stmt = select(models.MessageEmbedding).where( + models.MessageEmbedding.message_id == created_message.public_id + ) + result = await db_session.execute(stmt) + embedding_record = result.scalar_one_or_none() + + # Verify the embedding was created + assert embedding_record is not None + assert embedding_record.message_id == created_message.public_id + assert embedding_record.content == test_message_content + assert embedding_record.workspace_name == test_workspace.name + assert embedding_record.session_name == test_session.name + assert embedding_record.peer_name == test_peer.name + # Verify embedding vector exists and is not empty + assert embedding_record.embedding is not None + assert len(embedding_record.embedding) > 0 + + +@pytest.mark.asyncio +async def test_message_embedding_not_created_when_setting_disabled( + db_session: AsyncSession, + sample_data: tuple[Workspace, Peer], + monkeypatch: pytest.MonkeyPatch, +): + """Test that MessageEmbedding is NOT created when EMBED_MESSAGES setting is False""" + # Monkeypatch the setting to disable message embeddings + monkeypatch.setattr("src.config.settings.LLM.EMBED_MESSAGES", False) + + test_workspace, test_peer = sample_data + + # Create a test session + test_session = models.Session( + workspace_name=test_workspace.name, name=str(generate_nanoid()) + ) + db_session.add(test_session) + await db_session.commit() + + # Create a message using the CRUD function directly + test_message_content = "This is a test message without embedding" + messages = [ + MessageCreate( + content=test_message_content, + peer_id=test_peer.name, + metadata={"test": "embedding_disabled"}, + ) + ] + + created_messages = await create_messages( + db=db_session, + messages=messages, + workspace_name=test_workspace.name, + session_name=test_session.name, + ) + + assert len(created_messages) == 1 + created_message = created_messages[0] + + # Query the MessageEmbedding table to verify NO embedding was created + stmt = select(models.MessageEmbedding).where( + models.MessageEmbedding.message_id == created_message.public_id + ) + result = await db_session.execute(stmt) + embedding_record = result.scalar_one_or_none() + + # Verify no embedding was created + assert embedding_record is None + + +@pytest.mark.asyncio +async def test_multiple_message_embeddings_created_when_setting_enabled( + db_session: AsyncSession, + sample_data: tuple[Workspace, Peer], + monkeypatch: pytest.MonkeyPatch, +): + """Test that multiple MessageEmbeddings are created for batch message creation""" + # Monkeypatch the setting to enable message embeddings + monkeypatch.setattr("src.config.settings.LLM.EMBED_MESSAGES", True) + + test_workspace, test_peer = sample_data + + # Create a test session + test_session = models.Session( + workspace_name=test_workspace.name, name=str(generate_nanoid()) + ) + db_session.add(test_session) + await db_session.commit() + + # Create multiple messages + messages = [ + MessageCreate( + content="First test message", + peer_id=test_peer.name, + metadata={"order": 1}, + ), + MessageCreate( + content="Second test message", + peer_id=test_peer.name, + metadata={"order": 2}, + ), + ] + + created_messages = await create_messages( + db=db_session, + messages=messages, + workspace_name=test_workspace.name, + session_name=test_session.name, + ) + + assert len(created_messages) == 2 + + # Query the MessageEmbedding table to verify embeddings were created for both messages + for i, created_message in enumerate(created_messages): + stmt = select(models.MessageEmbedding).where( + models.MessageEmbedding.message_id == created_message.public_id + ) + result = await db_session.execute(stmt) + embedding_record = result.scalar_one_or_none() + + # Verify the embedding was created + assert embedding_record is not None + assert embedding_record.message_id == created_message.public_id + assert embedding_record.content == messages[i].content + assert embedding_record.workspace_name == test_workspace.name + assert embedding_record.session_name == test_session.name + assert embedding_record.peer_name == test_peer.name + assert embedding_record.embedding is not None + assert len(embedding_record.embedding) > 0 + + +@pytest.mark.asyncio +async def test_message_embedding_with_peer_only_messages( + db_session: AsyncSession, + sample_data: tuple[Workspace, Peer], + monkeypatch: pytest.MonkeyPatch, +): + """Test that MessageEmbedding is created for peer-only messages (no session)""" + # Monkeypatch the setting to enable message embeddings + monkeypatch.setattr("src.config.settings.LLM.EMBED_MESSAGES", True) + + test_workspace, test_peer = sample_data + + # Create a message for peer only (no session) + test_message_content = "This is a peer-only message with embedding" + messages = [ + MessageCreate( + content=test_message_content, + peer_id=test_peer.name, # This will be overridden by the function + metadata={"test": "peer_only"}, + ) + ] + + created_messages = await create_messages_for_peer( + db=db_session, + messages=messages, + workspace_name=test_workspace.name, + peer_name=test_peer.name, + ) + + assert len(created_messages) == 1 + created_message = created_messages[0] + + # Verify the message was created with peer but no session + assert created_message.peer_name == test_peer.name + assert created_message.session_name is None + + # Query the MessageEmbedding table to verify an embedding was created + stmt = select(models.MessageEmbedding).where( + models.MessageEmbedding.message_id == created_message.public_id + ) + result = await db_session.execute(stmt) + embedding_record = result.scalar_one_or_none() + + # Verify the embedding was created + assert embedding_record is not None + assert embedding_record.message_id == created_message.public_id + assert embedding_record.content == test_message_content + assert embedding_record.workspace_name == test_workspace.name + assert ( + embedding_record.session_name is None + ) # Should be None for peer-only messages + assert embedding_record.peer_name == test_peer.name + assert embedding_record.embedding is not None + assert len(embedding_record.embedding) > 0 + + +@pytest.mark.asyncio +async def test_semantic_search_when_embeddings_enabled( + db_session: AsyncSession, + sample_data: tuple[Workspace, Peer], + monkeypatch: pytest.MonkeyPatch, + mock_openai_embeddings: dict[str, Any], +): + """Test that search uses semantic search by default when EMBED_MESSAGES is True""" + # Monkeypatch the setting to enable message embeddings + monkeypatch.setattr("src.config.settings.LLM.EMBED_MESSAGES", True) + + test_workspace, test_peer = sample_data + + test_session = models.Session( + workspace_name=test_workspace.name, name=str(generate_nanoid()) + ) + db_session.add(test_session) + await db_session.commit() + + test_message_content = ( + "I love programming with Python and building web applications" + ) + messages = [ + MessageCreate( + content=test_message_content, + peer_id=test_peer.name, + metadata={"test": "semantic_search"}, + ) + ] + + created_messages = await create_messages( + db=db_session, + messages=messages, + workspace_name=test_workspace.name, + session_name=test_session.name, + ) + + assert len(created_messages) == 1 + created_message = created_messages[0] + + # Verify the embedding was created + stmt = select(models.MessageEmbedding).where( + models.MessageEmbedding.message_id == created_message.public_id + ) + result = await db_session.execute(stmt) + embedding_record = result.scalar_one_or_none() + assert embedding_record is not None + + # Now test semantic search without explicitly setting semantic=True + # This should use semantic search because EMBED_MESSAGES is True + search_query = ( + "Python development and web apps" # Similar meaning to the message content + ) + + # Check the call count before search + initial_call_count: int = mock_openai_embeddings["embed"].call_count + + search_stmt = await search( + query=search_query, + workspace_name=test_workspace.name, + session_name=test_session.name, + ) + + # Verify that the embed method was called during search - e.g. we used semantic search + assert mock_openai_embeddings["embed"].call_count == initial_call_count + 1 + + search_result = await db_session.execute(search_stmt) + found_messages = list(search_result.scalars().all()) + + # Verify that our message was found via semantic search + assert len(found_messages) > 0 + found_message_ids = [msg.public_id for msg in found_messages] + assert created_message.public_id in found_message_ids + + +@pytest.mark.asyncio +async def test_message_chunking_creates_multiple_embeddings( + db_session: AsyncSession, + sample_data: tuple[Workspace, Peer], + monkeypatch: pytest.MonkeyPatch, + mock_openai_embeddings: dict[str, Any], +): + """Test that messages exceeding token limits are chunked and create multiple embeddings""" + # Monkeypatch the setting to enable message embeddings + monkeypatch.setattr("src.config.settings.LLM.EMBED_MESSAGES", True) + + # Mock a low token limit to force chunking + monkeypatch.setattr("src.config.settings.LLM.MAX_EMBEDDING_TOKENS", 10) + + test_workspace, test_peer = sample_data + + # Create a test session + test_session = models.Session( + workspace_name=test_workspace.name, name=str(generate_nanoid()) + ) + db_session.add(test_session) + await db_session.commit() + + test_message_content = "This is a very long message that should be chunked into multiple pieces because it exceeds the token limit that we set for testing purposes. This message contains many words and should definitely be split into multiple chunks." + + def mock_batch_embed_chunked( + id_resource_dict: dict[str, tuple[str, list[int]]], + ) -> dict[str, list[list[float]]]: + return { + text_id: [[0.1] * 1536, [0.2] * 1536, [0.3] * 1536] # 3 chunks per message + for text_id in id_resource_dict + } + + mock_openai_embeddings["batch_embed"].side_effect = mock_batch_embed_chunked + + messages = [ + MessageCreate( + content=test_message_content, + peer_id=test_peer.name, + metadata={"test": "chunking"}, + ) + ] + + created_messages = await create_messages( + db=db_session, + messages=messages, + workspace_name=test_workspace.name, + session_name=test_session.name, + ) + + assert len(created_messages) == 1 + created_message = created_messages[0] + + # Query the MessageEmbedding table to verify multiple embeddings were created + stmt = select(models.MessageEmbedding).where( + models.MessageEmbedding.message_id == created_message.public_id + ) + result = await db_session.execute(stmt) + embedding_records = list(result.scalars().all()) + + # Verify multiple embeddings were created (one per chunk) + assert len(embedding_records) == 3 # Should have 3 embeddings for 3 chunks + + for _, embedding_record in enumerate(embedding_records): + assert embedding_record.message_id == created_message.public_id + assert ( + embedding_record.content == test_message_content + ) # Full content stored in each + assert embedding_record.workspace_name == test_workspace.name + assert embedding_record.session_name == test_session.name + assert embedding_record.peer_name == test_peer.name + assert embedding_record.embedding is not None + assert len(embedding_record.embedding) == 1536 + # Each chunk should have a different embedding vector (0.1, 0.2, 0.3) + assert embedding_record.embedding[0] in [0.1, 0.2, 0.3] diff --git a/tests/routes/test_peers.py b/tests/routes/test_peers.py index 1c51f6d1..66e039c0 100644 --- a/tests/routes/test_peers.py +++ b/tests/routes/test_peers.py @@ -1,3 +1,4 @@ +import pytest from fastapi.testclient import TestClient from nanoid import generate as generate_nanoid @@ -506,7 +507,7 @@ def test_search_peer(client: TestClient, sample_data: tuple[Workspace, Peer]): # Search with a query response = client.post( f"/v2/workspaces/{test_workspace.name}/peers/{test_peer.name}/search", - json="search query", + json={"query": "search query"}, ) assert response.status_code == 200 data = response.json() @@ -527,7 +528,8 @@ def test_search_peer_empty_query( # Search with empty query response = client.post( - f"/v2/workspaces/{test_workspace.name}/peers/{test_peer.name}/search", json="" + f"/v2/workspaces/{test_workspace.name}/peers/{test_peer.name}/search", + json={"query": ""}, ) assert response.status_code == 200 data = response.json() @@ -546,8 +548,77 @@ def test_search_peer_nonexistent( response = client.post( f"/v2/workspaces/{test_workspace.name}/peers/{nonexistent_peer_id}/search", - json="test query", + json={"query": "test query"}, ) # This should probably return 404 or handle gracefully # The exact behavior depends on the crud.search implementation assert response.status_code in [200, 404, 422] + + +def test_search_peer_with_semantic_search_false( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test peer search with semantic=false""" + test_workspace, test_peer = sample_data + + # Add some messages to search through + client.post( + f"/v2/workspaces/{test_workspace.name}/peers/{test_peer.name}/messages", + json={ + "messages": [ + {"content": "Search this content", "peer_id": test_peer.name}, + {"content": "Another searchable message", "peer_id": test_peer.name}, + ] + }, + ) + + # Search with semantic=false + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/{test_peer.name}/search", + json={"query": "search", "semantic": False}, + ) + assert response.status_code == 200 + data = response.json() + + # Response should have pagination structure + assert "items" in data + assert "total" in data + assert "page" in data + assert "size" in data + assert isinstance(data["items"], list) + + +def test_search_peer_with_semantic_search_true_disabled( + client: TestClient, + sample_data: tuple[Workspace, Peer], + monkeypatch: pytest.MonkeyPatch, +): + """Test peer search with semantic=true when EMBED_MESSAGES is disabled""" + # Override the EMBED_MESSAGES setting to False for this test + monkeypatch.setattr("src.config.settings.LLM.EMBED_MESSAGES", False) + + test_workspace, test_peer = sample_data + + # Add some messages to search through + client.post( + f"/v2/workspaces/{test_workspace.name}/peers/{test_peer.name}/messages", + json={ + "messages": [ + {"content": "Search this content", "peer_id": test_peer.name}, + {"content": "Another searchable message", "peer_id": test_peer.name}, + ] + }, + ) + + # Search with semantic=true (should fail if EMBED_MESSAGES is disabled) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/peers/{test_peer.name}/search", + json={"query": "search", "semantic": True}, + ) + + assert response.status_code == 405 + + data = response.json() + assert "Semantic search requires EMBED_MESSAGES flag to be enabled" in data.get( + "detail", "" + ) diff --git a/tests/routes/test_sessions.py b/tests/routes/test_sessions.py index 0e3e3349..55d78811 100644 --- a/tests/routes/test_sessions.py +++ b/tests/routes/test_sessions.py @@ -773,7 +773,7 @@ def test_search_session(client: TestClient, sample_data: tuple[Workspace, Peer]) # Search with a query response = client.post( f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/search", - json="search query", + json={"query": "search query"}, ) assert response.status_code == 200 data = response.json() @@ -801,7 +801,8 @@ def test_search_session_empty_query( # Search with empty query response = client.post( - f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/search", json="" + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/search", + json={"query": ""}, ) assert response.status_code == 200 data = response.json() @@ -820,8 +821,91 @@ def test_search_session_nonexistent( response = client.post( f"/v2/workspaces/{test_workspace.name}/sessions/{nonexistent_session_id}/search", - json="test query", + json={"query": "test query"}, ) # This should probably return 404 or handle gracefully # The exact behavior depends on the crud.search implementation assert response.status_code in [200, 404, 422] + + +def test_search_session_with_semantic_search_false( + client: TestClient, sample_data: tuple[Workspace, Peer] +): + """Test session search with semantic=false""" + test_workspace, test_peer = sample_data + session_id = str(generate_nanoid()) + + # Create session + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={"id": session_id, "peer_names": {test_peer.name: {}}}, + ) + + # Add messages to search through + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages", + json={ + "messages": [ + {"content": "Search this content", "peer_id": test_peer.name}, + {"content": "Another message to find", "peer_id": test_peer.name}, + ] + }, + ) + + # Search with semantic=false + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/search", + json={"query": "search", "semantic": False}, + ) + assert response.status_code == 200 + data = response.json() + + # Response should have pagination structure + assert "items" in data + assert "total" in data + assert "page" in data + assert "size" in data + assert isinstance(data["items"], list) + + +def test_search_session_with_semantic_search_true_disabled( + client: TestClient, + sample_data: tuple[Workspace, Peer], + monkeypatch: pytest.MonkeyPatch, +): + """Test session search with semantic=true when EMBED_MESSAGES is disabled""" + # Override the EMBED_MESSAGES setting to False for this test + monkeypatch.setattr("src.config.settings.LLM.EMBED_MESSAGES", False) + + test_workspace, test_peer = sample_data + session_id = str(generate_nanoid()) + + # Create session + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions", + json={"id": session_id, "peer_names": {test_peer.name: {}}}, + ) + + # Add messages to search through + client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/messages", + json={ + "messages": [ + {"content": "Search this content", "peer_id": test_peer.name}, + {"content": "Another message to find", "peer_id": test_peer.name}, + ] + }, + ) + + # Search with semantic=true (should fail if EMBED_MESSAGES is disabled) + response = client.post( + f"/v2/workspaces/{test_workspace.name}/sessions/{session_id}/search", + json={"query": "search", "semantic": True}, + ) + + assert response.status_code == 405 + + data = response.json() + assert "Semantic search requires EMBED_MESSAGES flag to be enabled" in data.get( + "detail", "" + )