add MessageEmbedding table (#144)

* fix (sync): Add sync script between public and private remotes

* add embedding column to messages

* add semantic search and tests

* undo db.py change

* use embedding client

* types

* rm .github/workflows/sync-public-changes.yml

* CodeRabbit comments

* compute token count with pydantic

* semantic default None + fix tests

* types and fix make token_count private

* add MessageEmbedding table

* CR and type

* undo change to schema

* fix session / peer where

* add tests to validate embedding creation + search

* CR comment, add chunking todo

* fix get_or_create_collection with peer/target in agent.chat

* move embedding client and implement chunking

* rm comments

* fix bug in migration

* add script to generate message embeddings

* default all workspaces

* CR comments

---------

Co-authored-by: Vineeth Voruganti <13438633+VVoruganti@users.noreply.github.com>
This commit is contained in:
Rajat Ahuja 2025-06-26 14:52:18 -04:00 committed by GitHub
parent 41bf5adc92
commit c36d2ac449
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194
17 changed files with 1460 additions and 117 deletions

View File

@ -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
# =============================================================================

View File

@ -1,4 +1,3 @@
import sqlalchemy as sa
from alembic import op

View File

@ -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)

View File

@ -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())

View File

@ -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(

View File

@ -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

View File

@ -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

View File

@ -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,

228
src/embeddings.py Normal file
View File

@ -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
]

View File

@ -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

View File

@ -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)

View File

@ -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)

View File

@ -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"

View File

@ -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(

View File

@ -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]

View File

@ -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", ""
)

View File

@ -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", ""
)