From 16c8d0163a39ee4225c8779a5d38830abccc401d Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Wed, 15 Oct 2025 11:17:09 -0400 Subject: [PATCH 01/17] feat: filter validation error from sentry. centralize config --- src/deriver/queue_manager.py | 11 ++----- src/main.py | 34 ++++----------------- src/sentry.py | 58 ++++++++++++++++++++++++++++++++++++ 3 files changed, 65 insertions(+), 38 deletions(-) create mode 100644 src/sentry.py diff --git a/src/deriver/queue_manager.py b/src/deriver/queue_manager.py index fbe463a0..356ae3ba 100644 --- a/src/deriver/queue_manager.py +++ b/src/deriver/queue_manager.py @@ -28,6 +28,7 @@ from src.dreamer.dream_scheduler import ( set_dream_scheduler, ) from src.models import QueueItem +from src.sentry import initialize_sentry from src.utils.work_unit import parse_work_unit_key from src.webhooks.events import ( QueueEmptyEvent, @@ -68,15 +69,7 @@ class QueueManager: # Initialize Sentry if enabled, using settings if settings.SENTRY.ENABLED: - sentry_sdk.init( - dsn=settings.SENTRY.DSN, - enable_tracing=True, - release=settings.SENTRY.RELEASE, - environment=settings.SENTRY.ENVIRONMENT, - traces_sample_rate=settings.SENTRY.TRACES_SAMPLE_RATE, - profiles_sample_rate=settings.SENTRY.PROFILES_SAMPLE_RATE, - integrations=[AsyncioIntegration()], - ) + initialize_sentry(integrations=[AsyncioIntegration()]) def add_task(self, task: asyncio.Task[None]) -> None: """Track a new task""" diff --git a/src/main.py b/src/main.py index c1cb1976..4d317b38 100644 --- a/src/main.py +++ b/src/main.py @@ -3,22 +3,16 @@ import re import uuid from collections.abc import Awaitable, Callable from contextlib import asynccontextmanager -from typing import TYPE_CHECKING import sentry_sdk from fastapi import FastAPI, Request, Response from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import JSONResponse from fastapi_pagination import add_pagination - -from src import prometheus -from src.utils.logging import get_route_template - -if TYPE_CHECKING: - from sentry_sdk._types import Event, Hint from sentry_sdk.integrations.fastapi import FastApiIntegration from sentry_sdk.integrations.starlette import StarletteIntegration +from src import prometheus from src.config import settings from src.db import engine, request_context from src.exceptions import HonchoException @@ -31,6 +25,8 @@ from src.routers import ( workspaces, ) from src.security import create_admin_jwt +from src.sentry import initialize_sentry +from src.utils.logging import get_route_template def get_log_level() -> int: @@ -71,27 +67,7 @@ async def setup_admin_jwt(): # Sentry Setup SENTRY_ENABLED = settings.SENTRY.ENABLED if SENTRY_ENABLED: - - def before_send(event: "Event", hint: "Hint") -> "Event | None": - if "exc_info" in hint: - _, exc_value, _ = hint["exc_info"] - # Filter out HonchoExceptions from being sent to Sentry - if isinstance(exc_value, HonchoException): - return None - - return event - - # Sentry SDK's default behavior: - # - Captures INFO+ level logs as breadcrumbs - # - Captures ERROR+ level logs as Sentry events - # - # For custom log levels, use the LoggingIntegration class: - # sentry_sdk.init(..., integrations=[LoggingIntegration(level=logging.INFO, event_level=logging.ERROR)]) - sentry_sdk.init( - dsn=settings.SENTRY.DSN, - traces_sample_rate=settings.SENTRY.TRACES_SAMPLE_RATE, - profiles_sample_rate=settings.SENTRY.PROFILES_SAMPLE_RATE, - before_send=before_send, + initialize_sentry( integrations=[ StarletteIntegration( transaction_style="endpoint", @@ -99,7 +75,7 @@ if SENTRY_ENABLED: FastApiIntegration( transaction_style="endpoint", ), - ], + ] ) diff --git a/src/sentry.py b/src/sentry.py new file mode 100644 index 00000000..1980f0e8 --- /dev/null +++ b/src/sentry.py @@ -0,0 +1,58 @@ +"""Sentry initialization and configuration.""" + +from __future__ import annotations + +from collections.abc import Sequence +from typing import TYPE_CHECKING, Any + +import sentry_sdk +from pydantic import ValidationError + +from src.config import settings +from src.exceptions import HonchoException + +if TYPE_CHECKING: + from sentry_sdk._types import Event, Hint + + +_UNSET = object() + + +def _filter_sentry_event(event: Event, hint: Hint | None) -> Event | None: + """Filter out events raised from known non-actionable exceptions before Sentry sees them.""" + if not hint: + return event + + exc_info = hint.get("exc_info") + if not exc_info: + return event + + _, exc_value, _ = exc_info + if isinstance( + exc_value, HonchoException | ValidationError + ): # Filters out HonchoExceptions and ValidationErrors (typically coming from Pydantic) + return None + + return event + + +# Sentry SDK's default behavior: +# - Captures INFO+ level logs as breadcrumbs +# - Captures ERROR+ level logs as Sentry events +# +# For custom log levels, use the LoggingIntegration class: +# sentry_sdk.init(..., integrations=[LoggingIntegration(level=logging.INFO, event_level=logging.ERROR)]) +def initialize_sentry( + *, + integrations: Sequence[Any], +) -> None: + sentry_sdk.init( + dsn=settings.SENTRY.DSN, + enable_tracing=True, + release=settings.SENTRY.RELEASE, + environment=settings.SENTRY.ENVIRONMENT, + traces_sample_rate=settings.SENTRY.TRACES_SAMPLE_RATE, + profiles_sample_rate=settings.SENTRY.PROFILES_SAMPLE_RATE, + before_send=_filter_sentry_event, + integrations=integrations, + ) From 49b22c6f14fb00da3faf756d25ab7762461ceb1c Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Wed, 15 Oct 2025 11:31:12 -0400 Subject: [PATCH 02/17] fix: add integration type --- src/sentry.py | 8 +++----- 1 file changed, 3 insertions(+), 5 deletions(-) diff --git a/src/sentry.py b/src/sentry.py index 1980f0e8..7ad467f1 100644 --- a/src/sentry.py +++ b/src/sentry.py @@ -3,7 +3,7 @@ from __future__ import annotations from collections.abc import Sequence -from typing import TYPE_CHECKING, Any +from typing import TYPE_CHECKING import sentry_sdk from pydantic import ValidationError @@ -13,9 +13,7 @@ from src.exceptions import HonchoException if TYPE_CHECKING: from sentry_sdk._types import Event, Hint - - -_UNSET = object() + from sentry_sdk.integrations import Integration def _filter_sentry_event(event: Event, hint: Hint | None) -> Event | None: @@ -44,7 +42,7 @@ def _filter_sentry_event(event: Event, hint: Hint | None) -> Event | None: # sentry_sdk.init(..., integrations=[LoggingIntegration(level=logging.INFO, event_level=logging.ERROR)]) def initialize_sentry( *, - integrations: Sequence[Any], + integrations: Sequence[Integration], ) -> None: sentry_sdk.init( dsn=settings.SENTRY.DSN, From ec144d17eb0bcebe14eb11b4d9d58181c389a5d2 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Wed, 15 Oct 2025 11:44:37 -0400 Subject: [PATCH 03/17] fix: filter out fastapi validation error from sentry --- src/sentry.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/src/sentry.py b/src/sentry.py index 7ad467f1..e1fae034 100644 --- a/src/sentry.py +++ b/src/sentry.py @@ -6,6 +6,7 @@ from collections.abc import Sequence from typing import TYPE_CHECKING import sentry_sdk +from fastapi.exceptions import RequestValidationError from pydantic import ValidationError from src.config import settings @@ -27,7 +28,7 @@ def _filter_sentry_event(event: Event, hint: Hint | None) -> Event | None: _, exc_value, _ = exc_info if isinstance( - exc_value, HonchoException | ValidationError + exc_value, HonchoException | ValidationError | RequestValidationError ): # Filters out HonchoExceptions and ValidationErrors (typically coming from Pydantic) return None @@ -44,6 +45,11 @@ def initialize_sentry( *, integrations: Sequence[Integration], ) -> None: + """Initialize Sentry SDK with project settings. + + Args: + integrations: Sentry SDK integrations to enable (e.g., Starlette, FastAPI). + """ sentry_sdk.init( dsn=settings.SENTRY.DSN, enable_tracing=True, From 4dbf667debda6aa4076d4bcf2fa65cbd20368b8b Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Wed, 15 Oct 2025 11:48:17 -0400 Subject: [PATCH 04/17] fix: add log for sentry filters for debugging --- src/sentry.py | 12 +++++++++--- 1 file changed, 9 insertions(+), 3 deletions(-) diff --git a/src/sentry.py b/src/sentry.py index e1fae034..498392f3 100644 --- a/src/sentry.py +++ b/src/sentry.py @@ -2,6 +2,7 @@ from __future__ import annotations +import logging from collections.abc import Sequence from typing import TYPE_CHECKING @@ -16,6 +17,8 @@ if TYPE_CHECKING: from sentry_sdk._types import Event, Hint from sentry_sdk.integrations import Integration +logger = logging.getLogger(__name__) + def _filter_sentry_event(event: Event, hint: Hint | None) -> Event | None: """Filter out events raised from known non-actionable exceptions before Sentry sees them.""" @@ -27,9 +30,12 @@ def _filter_sentry_event(event: Event, hint: Hint | None) -> Event | None: return event _, exc_value, _ = exc_info - if isinstance( - exc_value, HonchoException | ValidationError | RequestValidationError - ): # Filters out HonchoExceptions and ValidationErrors (typically coming from Pydantic) + if isinstance(exc_value, HonchoException): + return None + + # Filters out ValidationErrors and RequestValidationErrors (typically coming from Pydantic) + if isinstance(exc_value, ValidationError | RequestValidationError): + logger.info(f"Filtering out validation error from Sentry: {exc_value}") return None return event From f0b246197e657f31ffde8c6f2a8832bfed707cac Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Wed, 15 Oct 2025 11:57:56 -0400 Subject: [PATCH 05/17] fix: keep internal validation errors --- src/main.py | 30 +++++++++++++++++++++++++++++- src/sentry.py | 29 +++-------------------------- 2 files changed, 32 insertions(+), 27 deletions(-) diff --git a/src/main.py b/src/main.py index 4d317b38..6ae59192 100644 --- a/src/main.py +++ b/src/main.py @@ -3,12 +3,15 @@ import re import uuid from collections.abc import Awaitable, Callable from contextlib import asynccontextmanager +from typing import TYPE_CHECKING import sentry_sdk from fastapi import FastAPI, Request, Response +from fastapi.exceptions import RequestValidationError from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import JSONResponse from fastapi_pagination import add_pagination +from pydantic import ValidationError from sentry_sdk.integrations.fastapi import FastApiIntegration from sentry_sdk.integrations.starlette import StarletteIntegration @@ -28,6 +31,9 @@ from src.security import create_admin_jwt from src.sentry import initialize_sentry from src.utils.logging import get_route_template +if TYPE_CHECKING: + from sentry_sdk._types import Event, Hint + def get_log_level() -> int: """ @@ -64,6 +70,27 @@ async def setup_admin_jwt(): print(f"\n ADMIN JWT: {token}\n") +def before_send(event: "Event", hint: "Hint | None") -> "Event | None": + """Filter out events raised from known non-actionable exceptions before Sentry sees them.""" + if not hint: + return event + + exc_info = hint.get("exc_info") + if not exc_info: + return event + + _, exc_value, _ = exc_info + if isinstance(exc_value, HonchoException): + return None + + # Filters out ValidationErrors and RequestValidationErrors (typically coming from Pydantic) + if isinstance(exc_value, ValidationError | RequestValidationError): + logger.info(f"Filtering out validation error from Sentry: {exc_value}") + return None + + return event + + # Sentry Setup SENTRY_ENABLED = settings.SENTRY.ENABLED if SENTRY_ENABLED: @@ -75,7 +102,8 @@ if SENTRY_ENABLED: FastApiIntegration( transaction_style="endpoint", ), - ] + ], + before_send=before_send, ) diff --git a/src/sentry.py b/src/sentry.py index 498392f3..f213844c 100644 --- a/src/sentry.py +++ b/src/sentry.py @@ -7,40 +7,16 @@ from collections.abc import Sequence from typing import TYPE_CHECKING import sentry_sdk -from fastapi.exceptions import RequestValidationError -from pydantic import ValidationError from src.config import settings -from src.exceptions import HonchoException if TYPE_CHECKING: - from sentry_sdk._types import Event, Hint + from sentry_sdk._types import EventProcessor from sentry_sdk.integrations import Integration logger = logging.getLogger(__name__) -def _filter_sentry_event(event: Event, hint: Hint | None) -> Event | None: - """Filter out events raised from known non-actionable exceptions before Sentry sees them.""" - if not hint: - return event - - exc_info = hint.get("exc_info") - if not exc_info: - return event - - _, exc_value, _ = exc_info - if isinstance(exc_value, HonchoException): - return None - - # Filters out ValidationErrors and RequestValidationErrors (typically coming from Pydantic) - if isinstance(exc_value, ValidationError | RequestValidationError): - logger.info(f"Filtering out validation error from Sentry: {exc_value}") - return None - - return event - - # Sentry SDK's default behavior: # - Captures INFO+ level logs as breadcrumbs # - Captures ERROR+ level logs as Sentry events @@ -50,6 +26,7 @@ def _filter_sentry_event(event: Event, hint: Hint | None) -> Event | None: def initialize_sentry( *, integrations: Sequence[Integration], + before_send: EventProcessor | None = None, ) -> None: """Initialize Sentry SDK with project settings. @@ -63,6 +40,6 @@ def initialize_sentry( environment=settings.SENTRY.ENVIRONMENT, traces_sample_rate=settings.SENTRY.TRACES_SAMPLE_RATE, profiles_sample_rate=settings.SENTRY.PROFILES_SAMPLE_RATE, - before_send=_filter_sentry_event, + before_send=before_send, integrations=integrations, ) From cffee2bf5003a03b0c3642545b9e21c2b66f2834 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Wed, 15 Oct 2025 12:07:16 -0400 Subject: [PATCH 06/17] Update src/sentry.py Co-authored-by: coderabbitai[bot] <136622811+coderabbitai[bot]@users.noreply.github.com> --- src/sentry.py | 1 + 1 file changed, 1 insertion(+) diff --git a/src/sentry.py b/src/sentry.py index f213844c..b11fe57c 100644 --- a/src/sentry.py +++ b/src/sentry.py @@ -32,6 +32,7 @@ def initialize_sentry( Args: integrations: Sentry SDK integrations to enable (e.g., Starlette, FastAPI). + before_send: Optional event filter callback to suppress specific exceptions. """ sentry_sdk.init( dsn=settings.SENTRY.DSN, From 77a965e97fa60119a7e2e2dec12f1ab601bc7092 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Thu, 16 Oct 2025 11:46:55 -0400 Subject: [PATCH 07/17] feat: fix race condition in message sequence batching (#235) * feat: fix race condition in message sequence batching * fix: CodeRabbit comments; commit early to release the advisory lock before generating embeddings * fix: use index + rm unused method * fix: PR comments * fix: bug in lock timeout * fix: patch tracked_db for peers route within conftest.py --- ...7a643_add_message_seq_in_session_column.py | 135 +++++++++++++++ src/crud/__init__.py | 2 - src/crud/message.py | 160 ++++++++---------- src/crud/session.py | 1 + src/deriver/enqueue.py | 15 +- src/models.py | 7 + src/routers/messages.py | 2 + tests/conftest.py | 1 + tests/crud/test_workspace.py | 4 + tests/deriver/conftest.py | 3 + tests/deriver/test_deriver_processing.py | 1 + tests/deriver/test_queue_processing.py | 14 +- tests/integration/test_enqueue.py | 22 ++- tests/routes/test_messages.py | 17 ++ tests/routes/test_peers.py | 10 -- 15 files changed, 278 insertions(+), 116 deletions(-) create mode 100644 migrations/versions/bb6fb3a7a643_add_message_seq_in_session_column.py diff --git a/migrations/versions/bb6fb3a7a643_add_message_seq_in_session_column.py b/migrations/versions/bb6fb3a7a643_add_message_seq_in_session_column.py new file mode 100644 index 00000000..36627479 --- /dev/null +++ b/migrations/versions/bb6fb3a7a643_add_message_seq_in_session_column.py @@ -0,0 +1,135 @@ +"""add seq_in_session column to messages table + +Revision ID: bb6fb3a7a643 +Revises: 76ffba56fe8c +Create Date: 2025-10-20 12:00:00.000000 + +""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +from migrations.utils import column_exists, constraint_exists, get_schema + +# revision identifiers, used by Alembic. +revision: str = "bb6fb3a7a643" +down_revision: str | None = "76ffba56fe8c" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + +schema = get_schema() + +BATCH_SIZE = 10_000 + + +def upgrade() -> None: + if not column_exists("messages", "seq_in_session"): + op.add_column( + "messages", + sa.Column("seq_in_session", sa.BigInteger(), nullable=True), + schema=schema, + ) + + conn = op.get_bind() + preparer = conn.dialect.identifier_preparer + messages_table = sa.Table("messages", sa.MetaData(), schema=schema) + qualified_messages = preparer.format_table(messages_table) + id_col = preparer.quote("id") + workspace_col = preparer.quote("workspace_name") + session_col = preparer.quote("session_name") + seq_col = preparer.quote("seq_in_session") + + distinct_sessions = conn.execute( + sa.text( + f""" + SELECT DISTINCT + {workspace_col} AS workspace_name, + {session_col} AS session_name + FROM {qualified_messages} + """ + ) + ) + + update_stmt = sa.text( + f""" + WITH params AS ( + SELECT + :workspace_name AS workspace_name, + :session_name AS session_name, + COALESCE( + ( + SELECT MAX({seq_col}) + FROM {qualified_messages} + WHERE {workspace_col} = :workspace_name + AND {session_col} = :session_name + ), + 0 + ) AS offset + ), + batch AS ( + SELECT + m.{id_col} AS id, + ROW_NUMBER() OVER (ORDER BY m.{id_col}) AS rn, + params.offset + FROM {qualified_messages} AS m + JOIN params ON TRUE + WHERE m.{workspace_col} = params.workspace_name + AND m.{session_col} = params.session_name + AND m.{seq_col} IS NULL + ORDER BY m.{id_col} + LIMIT :batch_size + ) + UPDATE {qualified_messages} AS m + SET {seq_col} = batch.rn + batch.offset + FROM batch + WHERE m.{id_col} = batch.id + """ + ) + + for workspace_name, session_name in distinct_sessions: + while True: + result = conn.execute( + update_stmt, + { + "workspace_name": workspace_name, + "session_name": session_name, + "batch_size": BATCH_SIZE, + }, + ) + updated_rows = result.rowcount or 0 + result.close() + if updated_rows == 0: + break + distinct_sessions.close() + + op.alter_column( + "messages", + "seq_in_session", + nullable=False, + schema=schema, + ) + + if not constraint_exists("messages", "uq_messages_session_seq", "unique"): + op.create_unique_constraint( + "uq_messages_session_seq", + "messages", + ["workspace_name", "session_name", "seq_in_session"], + schema=schema, + ) + + +def downgrade() -> None: + schema = get_schema() + + if constraint_exists("messages", "uq_messages_session_seq", "unique"): + op.drop_constraint( + "uq_messages_session_seq", + "messages", + type_="unique", + schema=schema, + ) + + if column_exists("messages", "seq_in_session"): + op.drop_column("messages", "seq_in_session", schema=schema) diff --git a/src/crud/__init__.py b/src/crud/__init__.py index 69a4799e..7e0db845 100644 --- a/src/crud/__init__.py +++ b/src/crud/__init__.py @@ -9,7 +9,6 @@ from .message import ( create_messages, get_message, get_message_seq_in_session, - get_message_seqs_in_session_batch, get_messages, get_messages_id_range, update_message, @@ -67,7 +66,6 @@ __all__ = [ "get_messages_id_range", "get_message", "get_message_seq_in_session", - "get_message_seqs_in_session_batch", "update_message", # Peer "get_or_create_peers", diff --git a/src/crud/message.py b/src/crud/message.py index 143f1aa9..c6d3476d 100644 --- a/src/crud/message.py +++ b/src/crud/message.py @@ -2,7 +2,7 @@ from logging import getLogger from typing import Any from nanoid import generate as generate_nanoid -from sqlalchemy import ColumnElement, Select, and_, func, select +from sqlalchemy import ColumnElement, Select, and_, func, select, text from sqlalchemy.ext.asyncio import AsyncSession from src import models, schemas @@ -78,9 +78,32 @@ async def create_messages( workspace_name=workspace_name, ) + await db.execute(text("SET LOCAL lock_timeout = '5s'")) + await db.execute( + text( + "SELECT pg_advisory_xact_lock(hashtext(:workspace_name), hashtext(:session_name))" + ), + {"workspace_name": workspace_name, "session_name": session_name}, + ) + + # Get the last sequence number on a session - uses (workspace_name, session_name, seq_in_session) index + last_seq = ( + await db.scalar( + select(models.Message.seq_in_session) + .where( + models.Message.workspace_name == workspace_name, + models.Message.session_name == session_name, + ) + .order_by(models.Message.seq_in_session.desc()) + .limit(1) + ) + or 0 + ) + # Create list of message objects (this will trigger the before_insert event) message_objects: list[models.Message] = [] - for message in messages: + for offset, message in enumerate(messages, start=1): + message_seq_in_session = last_seq + offset message_obj = models.Message( session_name=session_name, peer_name=message.peer_name, @@ -90,46 +113,55 @@ async def create_messages( public_id=generate_nanoid(), token_count=len(message.encoded_message), created_at=message.created_at, # Use provided created_at if available + seq_in_session=message_seq_in_session, ) message_objects.append(message_obj) db.add_all(message_objects) - await db.flush() - - if settings.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) + # Commit here to release the advisory lock before generating embeddings await db.commit() + try: + if settings.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() + except Exception: + logger.exception( + "Failed to generate message embeddings for %s messages in workspace %s and session %s.", + len(message_objects), + workspace_name, + session_name, + ) return message_objects @@ -268,61 +300,13 @@ async def get_message_seq_in_session( The sequence number of the message (1-indexed) """ stmt = ( - select(func.count(models.Message.id)) + select(models.Message.seq_in_session) .where(models.Message.workspace_name == workspace_name) .where(models.Message.session_name == session_name) - .where(models.Message.id < message_id) + .where(models.Message.id == message_id) ) - result = await db.execute(stmt) - count = result.scalar() or 0 - return count + 1 - - -async def get_message_seqs_in_session_batch( - db: AsyncSession, - workspace_name: str, - session_name: str, - message_ids: list[int], -) -> dict[int, int]: - """ - Get the sequence numbers for multiple messages within a session in a single query. - - Args: - db: Database session - workspace_name: Name of the workspace - session_name: Name of the session - message_ids: List of message primary key IDs - - Returns: - Dictionary mapping message_id to sequence number (1-indexed). - If a given ID does not exist in the specified session, its value will be 0. - Note: duplicate IDs in the input are de-duplicated in the query and the - resulting mapping will contain a single entry per unique message_id. - """ - if not message_ids: - return {} - - unique_ids = list(set(message_ids)) - - # Rank all messages in the session, then select only the ones we care about - ranked = ( - select( - models.Message.id.label("id"), - func.row_number().over(order_by=models.Message.id).label("seq"), - ) - .where( - models.Message.workspace_name == workspace_name, - models.Message.session_name == session_name, - ) - .subquery() - ) - stmt = select(ranked.c.id, ranked.c.seq).where(ranked.c.id.in_(unique_ids)) - result = await db.execute(stmt) - rows = result.all() - id_to_position = {row[0]: int(row[1]) for row in rows} - - # Return positions for requested IDs; 0 for any missing/out-of-session IDs - return {msg_id: id_to_position.get(msg_id, 0) for msg_id in message_ids} + seq: int | None = await db.scalar(stmt) + return int(seq) if seq is not None else 0 async def get_message( diff --git a/src/crud/session.py b/src/crud/session.py index b3adf591..cfb77660 100644 --- a/src/crud/session.py +++ b/src/crud/session.py @@ -336,6 +336,7 @@ async def clone_session( "h_metadata": message.h_metadata, "workspace_name": workspace_name, "peer_name": message.peer_name, + "seq_in_session": message.seq_in_session, } for message in messages_to_clone ] diff --git a/src/deriver/enqueue.py b/src/deriver/enqueue.py index a9eb736e..9405f6fa 100644 --- a/src/deriver/enqueue.py +++ b/src/deriver/enqueue.py @@ -98,12 +98,6 @@ async def handle_session( db_session, workspace_name, session_name ) - # Get all message IDs to fetch sequences in batch - message_ids = [msg["message_id"] for msg in payload] - message_seq_map = await crud.get_message_seqs_in_session_batch( - db_session, workspace_name, session_name, message_ids - ) - queue_records: list[dict[str, Any]] = [] for message in payload: @@ -114,7 +108,6 @@ async def handle_session( peers_with_configuration, session.id, deriver_disabled=deriver_disabled, - message_seq_map=message_seq_map, ) ) return queue_records @@ -253,7 +246,6 @@ async def generate_queue_records( session_id: str, *, deriver_disabled: bool, - message_seq_map: dict[int, int] | None = None, ) -> list[dict[str, Any]]: """ Process a single message and generate queue records based on configurations. @@ -272,10 +264,9 @@ async def generate_queue_records( observed = message["peer_name"] message_id: int = message["message_id"] - # Use pre-fetched sequence if available, otherwise fall back to individual query - if message_seq_map and message_id in message_seq_map: - message_seq_in_session = message_seq_map[message_id] - else: + # Prefer the sequence captured during message creation; fallback only if missing + message_seq_in_session = int(message.get("seq_in_session") or 0) + if message_seq_in_session <= 0: message_seq_in_session = await crud.get_message_seq_in_session( db_session, workspace_name=message["workspace_name"], diff --git a/src/models.py b/src/models.py index 9b4e6f1d..1cd3b00c 100644 --- a/src/models.py +++ b/src/models.py @@ -186,6 +186,7 @@ class Message(Base): "internal_metadata", JSONB, default=dict ) token_count: Mapped[int] = mapped_column(Integer, default=0, nullable=False) + seq_in_session: Mapped[int] = mapped_column(BigInteger, nullable=False, index=True) created_at: Mapped[datetime.datetime] = mapped_column( DateTime(timezone=True), index=True, default=func.now() @@ -216,6 +217,12 @@ class Message(Base): "id", postgresql_include=["id", "created_at"], ), + UniqueConstraint( + "workspace_name", + "session_name", + "seq_in_session", + name="uq_messages_session_seq", + ), # Full text search index on content column Index( "idx_messages_content_gin", diff --git a/src/routers/messages.py b/src/routers/messages.py index 9c858c8e..259fc9d7 100644 --- a/src/routers/messages.py +++ b/src/routers/messages.py @@ -71,6 +71,7 @@ async def create_messages_for_session( "peer_name": message.peer_name, "created_at": message.created_at, "message_public_id": message.public_id, + "message_seq_in_session": message.seq_in_session, } for message in created_messages ] @@ -134,6 +135,7 @@ async def create_messages_with_file( "peer_name": message.peer_name, "created_at": message.created_at, "message_public_id": message.public_id, + "message_seq_in_session": message.seq_in_session, } for message in created_messages ] diff --git a/tests/conftest.py b/tests/conftest.py index 844a2769..9ba3036e 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -492,6 +492,7 @@ def mock_tracked_db(db_session: AsyncSession): patch("src.deriver.queue_manager.tracked_db", mock_tracked_db_context), patch("src.routers.sessions.tracked_db", mock_tracked_db_context), patch("src.crud.representation.tracked_db", mock_tracked_db_context), + patch("src.routers.peers.tracked_db", mock_tracked_db_context), ): yield diff --git a/tests/crud/test_workspace.py b/tests/crud/test_workspace.py index a9d32fb6..18601128 100644 --- a/tests/crud/test_workspace.py +++ b/tests/crud/test_workspace.py @@ -105,12 +105,14 @@ class TestWorkspaceCRUD: workspace_name=test_workspace.name, session_name=session.name, peer_name=test_peer.name, + seq_in_session=1, ) message2 = models.Message( content="Test message 2", workspace_name=test_workspace.name, session_name=session.name, peer_name=test_peer.name, + seq_in_session=2, ) db_session.add_all([message1, message2]) await db_session.flush() @@ -422,12 +424,14 @@ class TestWorkspaceCRUD: workspace_name=test_workspace.name, session_name=session1.name, peer_name=test_peer.name, + seq_in_session=1, ) message2 = models.Message( content="Test message 2", workspace_name=test_workspace.name, session_name=session2.name, peer_name=peer2.name, + seq_in_session=1, ) db_session.add_all([message1, message2]) diff --git a/tests/deriver/conftest.py b/tests/deriver/conftest.py index 9ade5ef7..89410656 100644 --- a/tests/deriver/conftest.py +++ b/tests/deriver/conftest.py @@ -83,18 +83,21 @@ async def sample_messages( "content": "Hello, this is the first message from peer1", "peer_name": peer1.name, "workspace_name": session.workspace_name, + "seq_in_session": 1, }, { "session_name": session.name, "content": "Hi there! This is a response from peer2", "peer_name": peer2.name, "workspace_name": session.workspace_name, + "seq_in_session": 2, }, { "session_name": session.name, "content": "I'm just observing this conversation as peer3", "peer_name": peer3.name, "workspace_name": session.workspace_name, + "seq_in_session": 3, }, ] diff --git a/tests/deriver/test_deriver_processing.py b/tests/deriver/test_deriver_processing.py index dc986ae2..002377df 100644 --- a/tests/deriver/test_deriver_processing.py +++ b/tests/deriver/test_deriver_processing.py @@ -169,6 +169,7 @@ class TestDeriverProcessing: session_name="test_session", peer_name="alice", content=f"message {message_id}", + seq_in_session=i + 1, token_count=0, created_at=now - timedelta(minutes=7 - i), ) diff --git a/tests/deriver/test_queue_processing.py b/tests/deriver/test_queue_processing.py index 6f1527ad..d620392a 100644 --- a/tests/deriver/test_queue_processing.py +++ b/tests/deriver/test_queue_processing.py @@ -121,13 +121,14 @@ class TestQueueProcessing: # Create and save messages to the database first messages: list[models.Message] = [] - for _ in range(3): + for i in range(3): message = models.Message( session_name=session.name, workspace_name=session.workspace_name, peer_name=peer.name, content="hello", token_count=10, + seq_in_session=i + 1, ) db_session.add(message) messages.append(message) @@ -291,6 +292,7 @@ class TestQueueProcessing: peer_name=peer.name, content=f"Test message {i}", token_count=token_count, + seq_in_session=i + 1, ) db_session.add(message) messages.append(message) @@ -396,13 +398,14 @@ class TestQueueProcessing: ] messages: list[models.Message] = [] - for peer, token_count in messages_data: + for i, (peer, token_count) in enumerate(messages_data): message = models.Message( session_name=session.name, workspace_name=session.workspace_name, peer_name=peer.name, content=f"Message from {peer.name}", token_count=token_count, + seq_in_session=i + 1, ) db_session.add(message) messages.append(message) @@ -561,13 +564,14 @@ class TestQueueProcessing: ] messages: list[models.Message] = [] - for peer, token_count in messages_data: + for i, (peer, token_count) in enumerate(messages_data): message = models.Message( session_name=session.name, workspace_name=session.workspace_name, peer_name=peer.name, content=f"Message from {peer.name}", token_count=token_count, + seq_in_session=i + 1, ) db_session.add(message) messages.append(message) @@ -704,6 +708,7 @@ class TestQueueProcessing: peer_name=peer.name, content="First summary message", public_id=generate_nanoid(), + seq_in_session=1, ), models.Message( id=1000, @@ -712,6 +717,7 @@ class TestQueueProcessing: peer_name=peer.name, content="Second summary message", public_id=generate_nanoid(), + seq_in_session=2, ), ] @@ -820,6 +826,7 @@ class TestQueueProcessing: peer_name=peer.name, content=f"Test message {i}", token_count=token_count, + seq_in_session=i + 1, ) db_session.add(message) messages.append(message) @@ -930,6 +937,7 @@ class TestQueueProcessing: peer_name=peer.name, content=f"Test message {i}", token_count=token_count, + seq_in_session=i + 1, ) db_session.add(message) messages.append(message) diff --git a/tests/integration/test_enqueue.py b/tests/integration/test_enqueue.py index 0f541704..5b1870e8 100644 --- a/tests/integration/test_enqueue.py +++ b/tests/integration/test_enqueue.py @@ -4,7 +4,7 @@ from unittest.mock import AsyncMock, patch import pytest from nanoid import generate as generate_nanoid -from sqlalchemy import select +from sqlalchemy import func, select from sqlalchemy.ext.asyncio import AsyncSession from src import crud, models, schemas @@ -26,6 +26,15 @@ class TestEnqueueFunction: count: int = 1, ) -> list[dict[str, Any]]: """Create real messages in database and return payload with actual IDs""" + # Get the current max sequence number for this session + result = await db_session.execute( + select(func.max(models.Message.seq_in_session)).where( + models.Message.workspace_name == workspace_name, + models.Message.session_name == session_name, + ) + ) + current_max_seq = result.scalar() or 0 + messages: list[models.Message] = [] for i in range(count): message = models.Message( @@ -34,6 +43,7 @@ class TestEnqueueFunction: peer_name=peer_name, content=f"Test message {i}", public_id=generate_nanoid(), + seq_in_session=current_max_seq + i + 1, token_count=10, h_metadata={"test": f"value_{i}"}, ) @@ -1085,6 +1095,15 @@ class TestAdvancedEnqueueEdgeCases: count: int = 1, ) -> list[dict[str, Any]]: """Create real messages in database and return payload with actual IDs""" + # Get the current max sequence number for this session + result = await db_session.execute( + select(func.max(models.Message.seq_in_session)).where( + models.Message.workspace_name == workspace_name, + models.Message.session_name == session_name, + ) + ) + current_max_seq = result.scalar() or 0 + messages: list[models.Message] = [] for i in range(count): message = models.Message( @@ -1093,6 +1112,7 @@ class TestAdvancedEnqueueEdgeCases: peer_name=peer_name, content=f"Test message {i}", public_id=generate_nanoid(), + seq_in_session=current_max_seq + i + 1, token_count=10, h_metadata={"test": f"value_{i}"}, ) diff --git a/tests/routes/test_messages.py b/tests/routes/test_messages.py index 2f222339..5bc3f1d1 100644 --- a/tests/routes/test_messages.py +++ b/tests/routes/test_messages.py @@ -173,6 +173,7 @@ async def test_get_messages( content="Test message", workspace_name=test_workspace.name, peer_name=test_peer.name, + seq_in_session=1, ) db_session.add(test_message) await db_session.commit() @@ -211,12 +212,14 @@ async def test_get_messages_with_reverse( content="First message", workspace_name=test_workspace.name, peer_name=test_peer.name, + seq_in_session=1, ) test_message2 = models.Message( session_name=test_session.name, content="Second message", workspace_name=test_workspace.name, peer_name=test_peer.name, + seq_in_session=2, ) db_session.add(test_message1) db_session.add(test_message2) @@ -266,6 +269,7 @@ async def test_get_messages_with_empty_filter( content="Test message", workspace_name=test_workspace.name, peer_name=test_peer.name, + seq_in_session=1, ) db_session.add(test_message) await db_session.commit() @@ -299,6 +303,7 @@ async def test_get_messages_with_null_filter( content="Test message", workspace_name=test_workspace.name, peer_name=test_peer.name, + seq_in_session=1, ) db_session.add(test_message) await db_session.commit() @@ -332,6 +337,7 @@ async def test_get_messages_no_body( content="Test message", workspace_name=test_workspace.name, peer_name=test_peer.name, + seq_in_session=1, ) db_session.add(test_message) await db_session.commit() @@ -364,6 +370,7 @@ async def test_get_filtered_messages( workspace_name=test_workspace.name, peer_name=test_peer.name, h_metadata={"key": "value"}, + seq_in_session=1, ) test_message2 = models.Message( session_name=test_session.name, @@ -371,6 +378,7 @@ async def test_get_filtered_messages( workspace_name=test_workspace.name, peer_name=test_peer.name, h_metadata={"key": "value2"}, + seq_in_session=2, ) db_session.add(test_message) db_session.add(test_message2) @@ -410,6 +418,7 @@ async def test_get_filtered_messages_with_complex_filter( workspace_name=test_workspace.name, peer_name=test_peer.name, h_metadata={"type": "question", "priority": "high", "category": "technical"}, + seq_in_session=1, ) test_message2 = models.Message( session_name=test_session.name, @@ -417,6 +426,7 @@ async def test_get_filtered_messages_with_complex_filter( workspace_name=test_workspace.name, peer_name=test_peer.name, h_metadata={"type": "answer", "priority": "high", "category": "technical"}, + seq_in_session=2, ) test_message3 = models.Message( session_name=test_session.name, @@ -424,6 +434,7 @@ async def test_get_filtered_messages_with_complex_filter( workspace_name=test_workspace.name, peer_name=test_peer.name, h_metadata={"type": "question", "priority": "low", "category": "general"}, + seq_in_session=3, ) db_session.add(test_message1) db_session.add(test_message2) @@ -494,6 +505,7 @@ async def test_update_message( content="Test message", workspace_name=test_workspace.name, peer_name=test_peer.name, + seq_in_session=1, ) db_session.add(test_message) await db_session.commit() @@ -526,6 +538,7 @@ async def test_update_message_with_complex_metadata( content="Test message", workspace_name=test_workspace.name, peer_name=test_peer.name, + seq_in_session=1, ) db_session.add(test_message) await db_session.commit() @@ -568,6 +581,7 @@ async def test_update_message_empty_metadata( workspace_name=test_workspace.name, peer_name=test_peer.name, h_metadata={"test_key": "test_value"}, + seq_in_session=1, ) db_session.add(test_message) await db_session.commit() @@ -607,6 +621,7 @@ async def test_update_message_with_empty_dict_metadata( workspace_name=test_workspace.name, peer_name=test_peer.name, h_metadata={"old_key": "old_value"}, + seq_in_session=1, ) db_session.add(test_message) await db_session.commit() @@ -638,6 +653,7 @@ async def test_get_single_message( content="Test message", workspace_name=test_workspace.name, peer_name=test_peer.name, + seq_in_session=1, ) db_session.add(test_message) await db_session.commit() @@ -870,6 +886,7 @@ async def test_update_message_handles_crud_value_error( content="Test message", workspace_name=test_workspace.name, peer_name=test_peer.name, + seq_in_session=1, ) db_session.add(test_message) await db_session.commit() diff --git a/tests/routes/test_peers.py b/tests/routes/test_peers.py index 2341e15e..f106abd9 100644 --- a/tests/routes/test_peers.py +++ b/tests/routes/test_peers.py @@ -1,5 +1,4 @@ from typing import Any -from unittest.mock import AsyncMock, patch import pytest from fastapi.testclient import TestClient @@ -309,15 +308,10 @@ def test_get_sessions_for_peer_with_empty_filter( assert isinstance(data["items"], list) -@patch("src.routers.peers.tracked_db") def test_chat( - mock_tracked_db: AsyncMock, client: TestClient, sample_data: tuple[Workspace, Peer], - db_session: AsyncSession, ): - mock_tracked_db.return_value.__aenter__.return_value = db_session - test_workspace, test_peer = sample_data target_peer = str(generate_nanoid()) @@ -335,15 +329,11 @@ def test_chat( assert "content" in data -@patch("src.routers.peers.tracked_db") def test_chat_with_optional_params( - mock_tracked_db: AsyncMock, client: TestClient, sample_data: tuple[Workspace, Peer], - db_session: AsyncSession, ): """Test chat endpoint with optional parameters""" - mock_tracked_db.return_value.__aenter__.return_value = db_session test_workspace, test_peer = sample_data session_id = str(generate_nanoid()) From a7d01d17dfd9a5f39674bd7302a79e4452badf7e Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Tue, 21 Oct 2025 11:30:56 -0400 Subject: [PATCH 08/17] feat: batch large operations in recent migrations (#242) * feat: batch large operations in recent migrations * fix: batch downgrade * fix: CodeRabbit comments * fix: delete user-created collections and documents * fix: guard against infinite loop for internal_metadata.session_name is null * fix: add idempotency guards --- ..._make_session_name_required_on_messages.py | 76 +++++-- ..._replace_collection_name_with_observer_.py | 208 ++++++++++++++---- ...c5_add_session_name_column_to_documents.py | 40 +++- 3 files changed, 247 insertions(+), 77 deletions(-) diff --git a/migrations/versions/05486ce795d5_make_session_name_required_on_messages.py b/migrations/versions/05486ce795d5_make_session_name_required_on_messages.py index 6f239da9..f8bc4486 100644 --- a/migrations/versions/05486ce795d5_make_session_name_required_on_messages.py +++ b/migrations/versions/05486ce795d5_make_session_name_required_on_messages.py @@ -106,27 +106,63 @@ def upgrade() -> None: f"Created session peer association for peer '{peer_name}' in default session '{default_session_name}'" ) - # Step 4: Assign orphaned messages for this peer to the default session - op.execute( - sa.text(f""" - UPDATE {schema}.messages - SET session_name = '{default_session_name}' - WHERE workspace_name = '{workspace_name}' - AND peer_name = '{peer_name}' - AND session_name IS NULL - """) - ) + # Step 4: Assign orphaned messages for this peer to the default session in batches + batch_size = 5000 + while True: + result = conn.execute( + sa.text(f""" + WITH batch AS ( + SELECT id + FROM {schema}.messages + WHERE workspace_name = :workspace_name + AND peer_name = :peer_name + AND session_name IS NULL + ORDER BY id + LIMIT :batch_size + ) + UPDATE {schema}.messages m + SET session_name = :default_session_name + FROM batch + WHERE m.id = batch.id + """), + { + "workspace_name": workspace_name, + "peer_name": peer_name, + "default_session_name": default_session_name, + "batch_size": batch_size, + }, + ) + if result.rowcount == 0: + break - # Step 4.5: Handle orphaned message embeddings for this peer - op.execute( - sa.text(f""" - UPDATE {schema}.message_embeddings - SET session_name = '{default_session_name}' - WHERE workspace_name = '{workspace_name}' - AND peer_name = '{peer_name}' - AND session_name IS NULL - """) - ) + # Step 4.5: Handle orphaned message embeddings for this peer in batches + batch_size = 5000 + while True: + result = conn.execute( + sa.text(f""" + WITH batch AS ( + SELECT id + FROM {schema}.message_embeddings + WHERE workspace_name = :workspace_name + AND peer_name = :peer_name + AND session_name IS NULL + ORDER BY id + LIMIT :batch_size + ) + UPDATE {schema}.message_embeddings me + SET session_name = :default_session_name + FROM batch + WHERE me.id = batch.id + """), + { + "workspace_name": workspace_name, + "peer_name": peer_name, + "default_session_name": default_session_name, + "batch_size": batch_size, + }, + ) + if result.rowcount == 0: + break # Step 5: Sanity check that no orphaned messages remain remaining_orphaned = conn.execute( diff --git a/migrations/versions/08894082221a_replace_collection_name_with_observer_.py b/migrations/versions/08894082221a_replace_collection_name_with_observer_.py index 2e27a5a8..cadad00f 100644 --- a/migrations/versions/08894082221a_replace_collection_name_with_observer_.py +++ b/migrations/versions/08894082221a_replace_collection_name_with_observer_.py @@ -55,16 +55,30 @@ def upgrade() -> None: ), {"session_id": session_id, "workspace_name": workspace_name}, ) - # Update all documents with NULL session_name - connection.execute( - text( - f""" - UPDATE {schema}.documents - SET session_name = '__global_observations__' - WHERE session_name IS NULL - """ - ), - ) + # Update all documents with NULL session_name in batches + batch_size = 5000 + while True: + result = connection.execute( + text( + f""" + WITH batch AS ( + SELECT id + FROM {schema}.documents + WHERE session_name IS NULL + ORDER BY id + LIMIT :batch_size + ) + UPDATE {schema}.documents d + SET session_name = '__global_observations__' + FROM batch + WHERE d.id = batch.id + AND d.session_name IS NULL + """ + ), + {"batch_size": batch_size}, + ) + if result.rowcount == 0: + break op.alter_column("documents", "session_name", nullable=False, schema=schema) @@ -84,29 +98,98 @@ def upgrade() -> None: schema=schema, ) - # Step 2: Populate collections observer and observed from existing name field + # Step 2a: Identify collections that should be deleted (in memory) + # These are user-created collections that are not used in the system: + # - name = 'global_representation' + # - name starts with peer_name + "_" (pattern: observer_observed) + # - name ends with "_" + peer_name (pattern: observed_observer) + collections_to_delete = connection.execute( + text( + f""" + SELECT id, name, peer_name, workspace_name + FROM {schema}.collections + WHERE name != 'global_representation' + AND name NOT LIKE peer_name || '_%' + AND name NOT LIKE '%_' || peer_name + """ + ) + ).fetchall() + + # Step 2b: Delete documents that reference collections marked for deletion (in batches) + if collections_to_delete: + # Delete documents in batches + batch_size = 5000 + for i in range(0, len(collections_to_delete), batch_size): + batch = collections_to_delete[i : i + batch_size] + collection_ids = [row.id for row in batch] + + connection.execute( + text( + f""" + DELETE FROM {schema}.documents d + USING {schema}.collections c + WHERE d.collection_name = c.name + AND d.peer_name = c.peer_name + AND d.workspace_name = c.workspace_name + AND c.id = ANY(:collection_ids) + """ + ), + {"collection_ids": collection_ids}, + ) + + # Step 2c: Delete the collections identified in step 2a (in batches) + if collections_to_delete: + batch_size = 5000 + for i in range(0, len(collections_to_delete), batch_size): + batch = collections_to_delete[i : i + batch_size] + collection_ids = [row.id for row in batch] + + connection.execute( + text( + f""" + DELETE FROM {schema}.collections + WHERE id = ANY(:collection_ids) + """ + ), + {"collection_ids": collection_ids}, + ) + + # Step 2d: Populate collections observer and observed from existing name field in batches # The logic is: # - observer = peer_name (the exact peer ID) # - If name is "global_representation", observed = peer_name (self-observation) # - If name starts with peer_name + "_", extract the observed part (pattern: observer_observed) # - If name ends with "_" + peer_name, extract the first part (pattern: observed_observer) - # - Otherwise (legacy edge cases), observed = name itself - connection.execute( - text( - f""" - UPDATE {schema}.collections - SET - observer = peer_name, - observed = CASE - WHEN name = 'global_representation' THEN peer_name - WHEN name LIKE peer_name || '_%' THEN substring(name from length(peer_name) + 2) - WHEN name LIKE '%_' || peer_name THEN substring(name from 1 for length(name) - length(peer_name) - 1) - ELSE name - END - WHERE observer IS NULL OR observed IS NULL - """ + # - Any legacy edge cases will have been deleted in step 2a. + batch_size = 5000 + while True: + result = connection.execute( + text( + f""" + WITH batch AS ( + SELECT id + FROM {schema}.collections + WHERE observer IS NULL OR observed IS NULL + ORDER BY id + LIMIT :batch_size + ) + UPDATE {schema}.collections c + SET + observer = c.peer_name, + observed = CASE + WHEN c.name = 'global_representation' THEN c.peer_name + WHEN c.name LIKE c.peer_name || '_%' THEN substring(c.name from length(c.peer_name) + 2) + WHEN c.name LIKE '%_' || c.peer_name THEN substring(c.name from 1 for length(c.name) - length(c.peer_name) - 1) + ELSE c.peer_name + END + FROM batch + WHERE c.id = batch.id + """ + ), + {"batch_size": batch_size}, ) - ) + if result.rowcount == 0: + break # Step 3: Make collections observer and observed NOT NULL op.alter_column("collections", "observer", nullable=False, schema=schema) @@ -151,6 +234,7 @@ def upgrade() -> None: AND d.collection_name = c.name AND d.peer_name = c.peer_name AND d.workspace_name = c.workspace_name + AND (d.observer IS NULL OR d.observed IS NULL) """ ), {"batch_size": batch_size}, @@ -415,19 +499,33 @@ def downgrade() -> None: schema=schema, ) - # Step 5: Populate documents collection_name from observer and observed - connection.execute( - text( - f""" - UPDATE {schema}.documents - SET collection_name = CASE - WHEN observer = observed THEN 'global_representation' - ELSE observer || '_' || observed - END - WHERE collection_name IS NULL - """ + # Step 5: Populate documents collection_name from observer and observed in batches + batch_size = 5000 + while True: + result = connection.execute( + text( + f""" + WITH batch AS ( + SELECT id + FROM {schema}.documents + WHERE collection_name IS NULL + ORDER BY id + LIMIT :batch_size + ) + UPDATE {schema}.documents d + SET collection_name = CASE + WHEN d.observer = d.observed THEN 'global_representation' + ELSE d.observer || '_' || d.observed + END + FROM batch + WHERE d.id = batch.id + AND d.collection_name IS NULL + """ + ), + {"batch_size": batch_size}, ) - ) + if result.rowcount == 0: + break # Step 6: Make documents collection_name NOT NULL op.alter_column("documents", "collection_name", nullable=False, schema=schema) @@ -440,16 +538,30 @@ def downgrade() -> None: schema=schema, ) - # Populate peer_name with observed value - connection.execute( - text( - f""" - UPDATE {schema}.documents - SET peer_name = observed - WHERE peer_name IS NULL - """ + # Populate peer_name with observed value in batches + batch_size = 5000 + while True: + result = connection.execute( + text( + f""" + WITH batch AS ( + SELECT id + FROM {schema}.documents + WHERE peer_name IS NULL + ORDER BY id + LIMIT :batch_size + ) + UPDATE {schema}.documents d + SET peer_name = d.observed + FROM batch + WHERE d.id = batch.id + AND d.peer_name IS NULL + """ + ), + {"batch_size": batch_size}, ) - ) + if result.rowcount == 0: + break # Make peer_name NOT NULL op.alter_column("documents", "peer_name", nullable=False, schema=schema) diff --git a/migrations/versions/564ba40505c5_add_session_name_column_to_documents.py b/migrations/versions/564ba40505c5_add_session_name_column_to_documents.py index d6a5fde4..d7238385 100644 --- a/migrations/versions/564ba40505c5_add_session_name_column_to_documents.py +++ b/migrations/versions/564ba40505c5_add_session_name_column_to_documents.py @@ -34,16 +34,38 @@ def upgrade() -> None: schema=schema, ) - # Step 2: Migrate data from internal_metadata to session_name column - op.execute( - sa.text( - f""" - UPDATE {schema}.documents - SET session_name = internal_metadata->>'session_name' - WHERE internal_metadata ? 'session_name' - """ + # Step 2: Migrate data from internal_metadata to session_name column in batches + # Process in batches to avoid timeout with large datasets + # Only migrate documents that have 'session_name' key with a non-null, non-empty value + bind = op.get_bind() + batch_size = 5000 + + while True: + result = bind.execute( + sa.text( + f""" + WITH batch AS ( + SELECT id + FROM {schema}.documents + WHERE session_name IS NULL + AND internal_metadata ? 'session_name' + AND internal_metadata->>'session_name' IS NOT NULL + AND internal_metadata->>'session_name' != '' + ORDER BY id + LIMIT :batch_size + ) + UPDATE {schema}.documents d + SET session_name = d.internal_metadata->>'session_name' + FROM batch b + WHERE d.id = b.id + AND d.session_name IS NULL + """ + ), + {"batch_size": batch_size}, ) - ) + + if result.rowcount == 0: + break # Step 3: Create index on session_name for efficient querying if not index_exists("documents", "idx_documents_session_name", inspector): From 5db7b4948cc1591e401c364027a244d55345afde Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Thu, 23 Oct 2025 16:24:35 -0400 Subject: [PATCH 09/17] feat: introduce alembic migration verification (#238) * feat: introducer migration verification checks * fix: move verification to tests/alembic * feat: add verification steps for all alembic migrations * fix: isolate test runs and implement all migration tests * test: parametrize * fix: CR comments * fix: Add README * feat: add precommit hook for validating alembic * fix: rm pytest-alembic package * test: create bulk resources to test migration batching * fix: add latest test * fix: tests to handle non-standard schema * chore: Code Rabbit nits * fix: CR comments 1 --------- Co-authored-by: Vineeth Voruganti <13438633+VVoruganti@users.noreply.github.com> --- .pre-commit-config.yaml | 21 +- ..._replace_collection_name_with_observer_.py | 16 +- ...564f50_add_user_id_and_app_id_to_tables.py | 70 ++-- .../88b0fb10906f_add_webhooks_table.py | 43 +++ ...917195d9b5e9_add_messageembedding_table.py | 9 +- .../versions/a1b2c3d4e5f6_initial_schema.py | 22 +- .../d429de0e5338_adopt_peer_paradigm.py | 327 +++++++++++----- pyproject.toml | 2 +- scripts/ensure_alembic_tests.py | 68 ++++ scripts/generate_jwt_secret.py | 2 +- scripts/generate_message_embeddings.py | 1 + src/models.py | 4 +- tests/alembic/README.md | 29 ++ tests/alembic/__init__.py | 5 + tests/alembic/conftest.py | 107 ++++++ tests/alembic/registry.py | 73 ++++ tests/alembic/revisions/__init__.py | 35 ++ ..._make_session_name_required_on_messages.py | 104 +++++ ..._replace_collection_name_with_observer_.py | 361 +++++++++++++++++ ...421aff_rename_metamessage_type_to_label.py | 123 ++++++ ...564f50_add_user_id_and_app_id_to_tables.py | 185 +++++++++ ...c5_add_session_name_column_to_documents.py | 211 ++++++++++ ...3cf2cf77_add_indexes_to_documents_table.py | 16 + ...ffba56fe8c_add_error_field_to_queueitem.py | 44 +++ .../test_88b0fb10906f_add_webhooks_table.py | 115 ++++++ ...917195d9b5e9_add_messageembedding_table.py | 26 ++ .../test_a1b2c3d4e5f6_initial_schema.py | 39 ++ ...change_metamessages_to_user_level_with_.py | 114 ++++++ ...7a643_add_message_seq_in_session_column.py | 16 + ...28084f472_add_indexes_for_messages_and_.py | 23 ++ .../test_d429de0e5338_adopt_peer_paradigm.py | 286 ++++++++++++++ tests/alembic/scaffold.py | 148 +++++++ tests/alembic/test_pipeline.py | 85 ++++ tests/alembic/verifier.py | 144 +++++++ tests/conftest.py | 4 + tests/sdk/test_pagination.py | 362 ++++++++---------- 36 files changed, 2876 insertions(+), 364 deletions(-) create mode 100644 scripts/ensure_alembic_tests.py create mode 100644 tests/alembic/README.md create mode 100644 tests/alembic/__init__.py create mode 100644 tests/alembic/conftest.py create mode 100644 tests/alembic/registry.py create mode 100644 tests/alembic/revisions/__init__.py create mode 100644 tests/alembic/revisions/test_05486ce795d5_make_session_name_required_on_messages.py create mode 100644 tests/alembic/revisions/test_08894082221a_replace_collection_name_with_observer_.py create mode 100644 tests/alembic/revisions/test_20f89a421aff_rename_metamessage_type_to_label.py create mode 100644 tests/alembic/revisions/test_556a16564f50_add_user_id_and_app_id_to_tables.py create mode 100644 tests/alembic/revisions/test_564ba40505c5_add_session_name_column_to_documents.py create mode 100644 tests/alembic/revisions/test_66e63cf2cf77_add_indexes_to_documents_table.py create mode 100644 tests/alembic/revisions/test_76ffba56fe8c_add_error_field_to_queueitem.py create mode 100644 tests/alembic/revisions/test_88b0fb10906f_add_webhooks_table.py create mode 100644 tests/alembic/revisions/test_917195d9b5e9_add_messageembedding_table.py create mode 100644 tests/alembic/revisions/test_a1b2c3d4e5f6_initial_schema.py create mode 100644 tests/alembic/revisions/test_b765d82110bd_change_metamessages_to_user_level_with_.py create mode 100644 tests/alembic/revisions/test_bb6fb3a7a643_add_message_seq_in_session_column.py create mode 100644 tests/alembic/revisions/test_c3828084f472_add_indexes_for_messages_and_.py create mode 100644 tests/alembic/revisions/test_d429de0e5338_adopt_peer_paradigm.py create mode 100644 tests/alembic/scaffold.py create mode 100644 tests/alembic/test_pipeline.py create mode 100644 tests/alembic/verifier.py diff --git a/.pre-commit-config.yaml b/.pre-commit-config.yaml index 5ba50470..7cc0de63 100644 --- a/.pre-commit-config.yaml +++ b/.pre-commit-config.yaml @@ -65,12 +65,31 @@ repos: # Run main application tests - id: pytest-main name: pytest (main app) - entry: uv run pytest tests/ + entry: uv run pytest tests/ --ignore=tests/alembic/ language: system files: ^(src/|tests/).*\.py$ stages: [pre-push] pass_filenames: false + # Run Alembic tests only when migrations change + - id: pytest-alembic + name: pytest (alembic migrations) + entry: uv run pytest tests/alembic/ + language: system + files: ^(migrations/.*\.py|tests/alembic/.*\.py)$ + stages: [pre-push] + pass_filenames: false + require_serial: true + + # Ensure each alembic migration revision has a corresponding test file + - id: ensure-alembic-coverage + name: ensure alembic migration test coverage + entry: uv run python scripts/ensure_alembic_tests.py + language: system + files: ^(migrations/versions/.*\.py|tests/alembic/revisions/.*\.py)$ + stages: [pre-push] + pass_filenames: false + # Run Python SDK tests (if they exist) - id: pytest-python-sdk name: pytest (Python SDK) diff --git a/migrations/versions/08894082221a_replace_collection_name_with_observer_.py b/migrations/versions/08894082221a_replace_collection_name_with_observer_.py index cadad00f..87067d67 100644 --- a/migrations/versions/08894082221a_replace_collection_name_with_observer_.py +++ b/migrations/versions/08894082221a_replace_collection_name_with_observer_.py @@ -476,12 +476,10 @@ def downgrade() -> None: # Make peer_name NOT NULL op.alter_column("collections", "peer_name", nullable=False, schema=schema) - # Recreate the foreign key constraint for peer_name - if not fk_exists( - "collections", "collections_peer_name_workspace_name_fkey", inspector - ): + # Recreate the legacy foreign key constraint for peer_name + if not fk_exists("collections", "fk_collections_peer_name_peers", inspector): op.create_foreign_key( - "collections_peer_name_workspace_name_fkey", + "fk_collections_peer_name_peers", "collections", "peers", ["peer_name", "workspace_name"], @@ -530,7 +528,7 @@ def downgrade() -> None: # Step 6: Make documents collection_name NOT NULL op.alter_column("documents", "collection_name", nullable=False, schema=schema) - # Step 6a: Restore peer_name column to documents (set to observed value) + # Step 6a: Restore peer_name column to documents (set to observer value) if not column_exists("documents", "peer_name", inspector): op.add_column( "documents", @@ -637,14 +635,14 @@ def downgrade() -> None: schema=schema, ) - # Step 11: Recreate the old foreign key constraint from documents to collections + # Step 11: Recreate the legacy foreign key constraint from documents to collections if not fk_exists( "documents", - "documents_collection_name_peer_name_workspace_name_fkey", + "fk_documents_collection_name_collections", inspector, ): op.create_foreign_key( - "documents_collection_name_peer_name_workspace_name_fkey", + "fk_documents_collection_name_collections", "documents", "collections", ["collection_name", "peer_name", "workspace_name"], diff --git a/migrations/versions/556a16564f50_add_user_id_and_app_id_to_tables.py b/migrations/versions/556a16564f50_add_user_id_and_app_id_to_tables.py index d766f124..05338d0e 100644 --- a/migrations/versions/556a16564f50_add_user_id_and_app_id_to_tables.py +++ b/migrations/versions/556a16564f50_add_user_id_and_app_id_to_tables.py @@ -435,18 +435,18 @@ def downgrade(): op.alter_column("documents", "user_id", nullable=True, schema=schema) print("Made app_id and user_id nullable again for documents table") - # Drop foreign key constraints - try: + # Drop foreign key constraints (guarded) + if fk_exists("documents", "documents_user_id_fkey", inspector): op.drop_constraint("documents_user_id_fkey", "documents", schema=schema) print("Dropped user_id foreign key for documents table") - except Exception as e: - print(f"Error dropping user_id foreign key for documents table: {e}") + else: + print("user_id foreign key for documents table does not exist; skipping drop") - try: + if fk_exists("documents", "documents_app_id_fkey", inspector): op.drop_constraint("documents_app_id_fkey", "documents", schema=schema) print("Dropped app_id foreign key for documents table") - except Exception as e: - print(f"Error dropping app_id foreign key for documents table: {e}") + else: + print("app_id foreign key for documents table does not exist; skipping drop") # Drop the columns op.drop_column("documents", "user_id", schema=schema) @@ -464,12 +464,12 @@ def downgrade(): op.alter_column("collections", "app_id", nullable=True, schema=schema) print("Made app_id nullable again for collections table") - # Drop foreign key constraint - try: + # Drop foreign key constraint (guarded) + if fk_exists("collections", "collections_app_id_fkey", inspector): op.drop_constraint("collections_app_id_fkey", "collections", schema=schema) print("Dropped app_id foreign key for collections table") - except Exception as e: - print(f"Error dropping app_id foreign key for collections table: {e}") + else: + print("app_id foreign key for collections table does not exist; skipping drop") # Drop the column op.drop_column("collections", "app_id", schema=schema) @@ -484,20 +484,26 @@ def downgrade(): ) print("Dropped index ix_metamessages_app_id") - # Make app_id nullable again - op.alter_column("metamessages", "app_id", nullable=True, schema=schema) - print("Made app_id nullable again for metamessages table") + # Make app_id nullable again (only if column exists) + if column_exists("metamessages", "app_id", inspector): + op.alter_column("metamessages", "app_id", nullable=True, schema=schema) + print("Made app_id nullable again for metamessages table") + else: + print("metamessages.app_id column does not exist; skipping alter nullable") - # Drop foreign key constraint - try: + # Drop foreign key constraint (guarded) + if fk_exists("metamessages", "metamessages_app_id_fkey", inspector): op.drop_constraint("metamessages_app_id_fkey", "metamessages", schema=schema) print("Dropped app_id foreign key for metamessages table") - except Exception as e: - print(f"Error dropping app_id foreign key for metamessages table: {e}") + else: + print("app_id foreign key for metamessages table does not exist; skipping drop") - # Drop the column - op.drop_column("metamessages", "app_id", schema=schema) - print("Dropped app_id column from metamessages table") + # Drop the column (only if present) + if column_exists("metamessages", "app_id", inspector): + op.drop_column("metamessages", "app_id", schema=schema) + print("Dropped app_id column from metamessages table") + else: + print("metamessages.app_id column does not exist; skipping drop column") # 2. Messages table @@ -514,18 +520,18 @@ def downgrade(): op.alter_column("messages", "user_id", nullable=True, schema=schema) print("Made app_id and user_id nullable again for messages table") - # Drop foreign key constraints - try: + # Drop foreign key constraints (guarded) + if fk_exists("messages", "messages_user_id_fkey", inspector): op.drop_constraint("messages_user_id_fkey", "messages", schema=schema) print("Dropped user_id foreign key for messages table") - except Exception as e: - print(f"Error dropping user_id foreign key for messages table: {e}") + else: + print("user_id foreign key for messages table does not exist; skipping drop") - try: + if fk_exists("messages", "messages_app_id_fkey", inspector): op.drop_constraint("messages_app_id_fkey", "messages", schema=schema) print("Dropped app_id foreign key for messages table") - except Exception as e: - print(f"Error dropping app_id foreign key for messages table: {e}") + else: + print("app_id foreign key for messages table does not exist; skipping drop") # Drop the columns op.drop_column("messages", "user_id", schema=schema) @@ -543,12 +549,12 @@ def downgrade(): op.alter_column("sessions", "app_id", nullable=True, schema=schema) print("Made app_id nullable again for sessions table") - # Drop foreign key constraint - try: + # Drop foreign key constraint (guarded) + if fk_exists("sessions", "sessions_app_id_fkey", inspector): op.drop_constraint("sessions_app_id_fkey", "sessions", schema=schema) print("Dropped app_id foreign key for sessions table") - except Exception as e: - print(f"Error dropping app_id foreign key for sessions table: {e}") + else: + print("app_id foreign key for sessions table does not exist; skipping drop") # Drop the column op.drop_column("sessions", "app_id", schema=schema) diff --git a/migrations/versions/88b0fb10906f_add_webhooks_table.py b/migrations/versions/88b0fb10906f_add_webhooks_table.py index da607938..39b5bf51 100644 --- a/migrations/versions/88b0fb10906f_add_webhooks_table.py +++ b/migrations/versions/88b0fb10906f_add_webhooks_table.py @@ -183,3 +183,46 @@ def downgrade() -> None: if column_exists("active_queue_sessions", "work_unit_data", inspector): op.drop_column("active_queue_sessions", "work_unit_data", schema=schema) + + # Re-add columns to active_queue_sessions + if not column_exists("active_queue_sessions", "session_id", inspector): + op.add_column( + "active_queue_sessions", + sa.Column("session_id", sa.TEXT(), nullable=True), + schema=schema, + ) + + if not column_exists("active_queue_sessions", "sender_name", inspector): + op.add_column( + "active_queue_sessions", + sa.Column("sender_name", sa.TEXT(), nullable=True), + schema=schema, + ) + + if not column_exists("active_queue_sessions", "target_name", inspector): + op.add_column( + "active_queue_sessions", + sa.Column("target_name", sa.TEXT(), nullable=True), + schema=schema, + ) + + if not column_exists("active_queue_sessions", "task_type", inspector): + op.add_column( + "active_queue_sessions", + sa.Column("task_type", sa.TEXT(), nullable=False), + schema=schema, + ) + + # Re-add unique constraint + if not constraint_exists( + "active_queue_sessions", + "unique_active_queue_session", + "unique", + inspector, + ): + op.create_unique_constraint( + "unique_active_queue_session", + "active_queue_sessions", + ["session_id", "sender_name", "target_name", "task_type"], + schema=schema, + ) diff --git a/migrations/versions/917195d9b5e9_add_messageembedding_table.py b/migrations/versions/917195d9b5e9_add_messageembedding_table.py index 8d8f78f4..4eb73739 100644 --- a/migrations/versions/917195d9b5e9_add_messageembedding_table.py +++ b/migrations/versions/917195d9b5e9_add_messageembedding_table.py @@ -40,14 +40,15 @@ def upgrade() -> None: server_default=sa.func.now(), ), # Foreign key constraints - sa.ForeignKeyConstraint(["message_id"], ["messages.public_id"]), - sa.ForeignKeyConstraint(["workspace_name"], ["workspaces.name"]), + sa.ForeignKeyConstraint(["message_id"], [f"{schema}.messages.public_id"]), + sa.ForeignKeyConstraint(["workspace_name"], [f"{schema}.workspaces.name"]), sa.ForeignKeyConstraint( ["session_name", "workspace_name"], - ["sessions.name", "sessions.workspace_name"], + [f"{schema}.sessions.name", f"{schema}.sessions.workspace_name"], ), sa.ForeignKeyConstraint( - ["peer_name", "workspace_name"], ["peers.name", "peers.workspace_name"] + ["peer_name", "workspace_name"], + [f"{schema}.peers.name", f"{schema}.peers.workspace_name"], ), schema=schema, ) diff --git a/migrations/versions/a1b2c3d4e5f6_initial_schema.py b/migrations/versions/a1b2c3d4e5f6_initial_schema.py index dfb47843..cafa43ca 100644 --- a/migrations/versions/a1b2c3d4e5f6_initial_schema.py +++ b/migrations/versions/a1b2c3d4e5f6_initial_schema.py @@ -83,7 +83,7 @@ def upgrade() -> None: sa.CheckConstraint("length(name) <= 512", name="name_length"), sa.CheckConstraint("public_id ~ '^[A-Za-z0-9_-]+$'", name="public_id_format"), sa.ForeignKeyConstraint( - ["app_id"], ["apps.public_id"], name=op.f("fk_users_app_id_apps") + ["app_id"], [f"{schema}.apps.public_id"], name=op.f("fk_users_app_id_apps") ), sa.PrimaryKeyConstraint("id", name=op.f("pk_users")), sa.UniqueConstraint("name", "app_id", name="unique_name_app_user"), @@ -130,7 +130,9 @@ def upgrade() -> None: sa.CheckConstraint("length(public_id) = 21", name="public_id_length"), sa.CheckConstraint("public_id ~ '^[A-Za-z0-9_-]+$'", name="public_id_format"), sa.ForeignKeyConstraint( - ["user_id"], ["users.public_id"], name=op.f("fk_sessions_user_id_users") + ["user_id"], + [f"{schema}.users.public_id"], + name=op.f("fk_sessions_user_id_users"), ), sa.PrimaryKeyConstraint("id", name=op.f("pk_sessions")), sa.UniqueConstraint("public_id", name=op.f("uq_sessions_public_id")), @@ -186,7 +188,7 @@ def upgrade() -> None: sa.CheckConstraint("public_id ~ '^[A-Za-z0-9_-]+$'", name="public_id_format"), sa.ForeignKeyConstraint( ["session_id"], - ["sessions.public_id"], + [f"{schema}.sessions.public_id"], name=op.f("fk_messages_session_id_sessions"), ), sa.PrimaryKeyConstraint("id", name=op.f("pk_messages")), @@ -246,7 +248,7 @@ def upgrade() -> None: sa.CheckConstraint("public_id ~ '^[A-Za-z0-9_-]+$'", name="public_id_format"), sa.ForeignKeyConstraint( ["message_id"], - ["messages.public_id"], + [f"{schema}.messages.public_id"], name=op.f("fk_metamessages_message_id_messages"), ), sa.PrimaryKeyConstraint("id", name=op.f("pk_metamessages")), @@ -308,7 +310,9 @@ def upgrade() -> None: sa.CheckConstraint("length(name) <= 512", name="name_length"), sa.CheckConstraint("public_id ~ '^[A-Za-z0-9_-]+$'", name="public_id_format"), sa.ForeignKeyConstraint( - ["user_id"], ["users.public_id"], name=op.f("fk_collections_user_id_users") + ["user_id"], + [f"{schema}.users.public_id"], + name=op.f("fk_collections_user_id_users"), ), sa.PrimaryKeyConstraint("id", name=op.f("pk_collections")), sa.UniqueConstraint("name", "user_id", name="unique_name_collection_user"), @@ -372,7 +376,7 @@ def upgrade() -> None: sa.CheckConstraint("public_id ~ '^[A-Za-z0-9_-]+$'", name="public_id_format"), sa.ForeignKeyConstraint( ["collection_id"], - ["collections.public_id"], + [f"{schema}.collections.public_id"], name=op.f("fk_documents_collection_id_collections"), ), sa.PrimaryKeyConstraint("id", name=op.f("pk_documents")), @@ -412,7 +416,9 @@ def upgrade() -> None: sa.Column("payload", postgresql.JSONB(astext_type=sa.Text()), nullable=False), sa.Column("processed", sa.Boolean(), nullable=False, server_default="false"), sa.ForeignKeyConstraint( - ["session_id"], ["sessions.id"], name=op.f("fk_queue_session_id_sessions") + ["session_id"], + [f"{schema}.sessions.id"], + name=op.f("fk_queue_session_id_sessions"), ), sa.PrimaryKeyConstraint("id", name=op.f("pk_queue")), schema=schema, @@ -437,7 +443,7 @@ def upgrade() -> None: ), sa.ForeignKeyConstraint( ["session_id"], - ["sessions.id"], + [f"{schema}.sessions.id"], name=op.f("fk_active_queue_sessions_session_id_sessions"), ), sa.PrimaryKeyConstraint("session_id", name=op.f("pk_active_queue_sessions")), diff --git a/migrations/versions/d429de0e5338_adopt_peer_paradigm.py b/migrations/versions/d429de0e5338_adopt_peer_paradigm.py index 9a8d9844..6cb42b74 100644 --- a/migrations/versions/d429de0e5338_adopt_peer_paradigm.py +++ b/migrations/versions/d429de0e5338_adopt_peer_paradigm.py @@ -13,7 +13,7 @@ import sqlalchemy as sa import tiktoken from alembic import op from nanoid import generate as generate_nanoid -from sqlalchemy import text +from sqlalchemy import Inspector, text from sqlalchemy.dialects import postgresql from migrations.utils import ( @@ -114,8 +114,117 @@ def downgrade() -> None: # Step 11: Restore foreign keys restore_foreign_keys(schema) + # Step 12: Readd metamessages table + if not table_exists("metamessages", inspector): + op.create_table( + "metamessages", + sa.Column( + "id", + sa.BigInteger(), + sa.Identity(always=False), + primary_key=True, + autoincrement=True, + index=True, + nullable=False, + ), + sa.Column("public_id", sa.TEXT(), nullable=False, index=True, unique=True), + sa.Column("content", sa.TEXT(), nullable=True), + sa.Column( + "message_id", + sa.TEXT(), + sa.ForeignKey("messages.public_id"), + nullable=True, + index=True, + ), + sa.Column( + "created_at", + sa.DateTime(timezone=True), + nullable=False, + index=True, + server_default=sa.text("now()"), + ), + sa.Column("label", sa.TEXT(), nullable=False, index=True), + sa.Column( + "session_id", + sa.TEXT(), + sa.ForeignKey("sessions.public_id"), + nullable=True, + index=True, + ), + sa.Column( + "user_id", + sa.TEXT(), + sa.ForeignKey("users.public_id"), + nullable=True, + index=True, + ), + sa.Column( + "app_id", + sa.TEXT(), + sa.ForeignKey("apps.public_id"), + nullable=True, + index=True, + ), + sa.Column( + "metadata", + postgresql.JSONB(astext_type=sa.Text()), + nullable=False, + server_default=sa.text("'{}'::jsonb"), + ), + sa.Index( + "idx_metamessages_lookup", + "label", + sa.text("id DESC"), + postgresql_include=["public_id", "message_id", "created_at"], + ), + sa.Index( + "idx_metamessages_user_lookup", + "user_id", + "label", + sa.text("id DESC"), + ), + sa.Index( + "idx_metamessages_session_lookup", + "session_id", + "label", + sa.text("id DESC"), + ), + sa.Index( + "idx_metamessages_message_lookup", + "message_id", + "label", + sa.text("id DESC"), + ), + sa.CheckConstraint("length(public_id) = 21", name="public_id_length"), + sa.CheckConstraint( + "public_id ~ '^[A-Za-z0-9_-]+$'", name="public_id_format" + ), + sa.CheckConstraint("length(content) <= 65535", name="content_length"), + sa.CheckConstraint("length(label) <= 512", name="label_length"), + # Added constraints to ensure consistency + sa.CheckConstraint( + "(message_id IS NULL) OR (session_id IS NOT NULL)", + name="message_requires_session", + ), + sa.ForeignKeyConstraint( + ["session_id"], + [f"{schema}.sessions.public_id"], + "fk_metamessages_session_id_sessions", + ), + sa.ForeignKeyConstraint( + ["user_id"], + [f"{schema}.users.public_id"], + "fk_metamessages_user_id_users", + ), + sa.ForeignKeyConstraint( + ["app_id"], + [f"{schema}.apps.public_id"], + "fk_metamessages_app_id_apps", + ), + ) -def rename_tables(schema: str, inspector) -> None: + +def rename_tables(schema: str, inspector: Inspector) -> None: """Rename apps->workspaces and users->peers tables.""" if inspector.has_table("apps", schema=schema): op.rename_table("apps", "workspaces", schema=schema) @@ -123,7 +232,7 @@ def rename_tables(schema: str, inspector) -> None: op.rename_table("users", "peers", schema=schema) -def update_workspaces_table(schema: str, inspector) -> None: +def update_workspaces_table(schema: str, inspector: Inspector) -> None: """Update workspaces table (formerly apps).""" # Add configuration column @@ -176,7 +285,7 @@ def update_workspaces_table(schema: str, inspector) -> None: ) -def update_peers_table(schema: str, inspector) -> None: +def update_peers_table(schema: str, inspector: Inspector) -> None: """Update peers table (formerly users).""" # Add configuration column @@ -266,9 +375,10 @@ def update_peers_table(schema: str, inspector) -> None: ) op.drop_column("peers", "app_id", schema=schema) + op.alter_column("peers", "workspace_name", nullable=False, schema=schema) -def update_sessions_table(schema: str, inspector) -> None: +def update_sessions_table(schema: str, inspector: Inspector) -> None: """Update sessions table.""" # Add configuration column @@ -304,6 +414,7 @@ def update_sessions_table(schema: str, inspector) -> None: f"UPDATE {schema}.sessions SET workspace_name = workspaces.name FROM {schema}.workspaces WHERE sessions.app_id = workspaces.id" ) ) + op.alter_column("sessions", "workspace_name", nullable=False, schema=schema) op.add_column( "sessions", @@ -365,7 +476,7 @@ def update_sessions_table(schema: str, inspector) -> None: ) -def create_and_populate_session_peers_table(schema: str, inspector) -> None: +def create_and_populate_session_peers_table(schema: str, inspector: Inspector) -> None: """Create and populate session_peers table.""" # Create session_peers table @@ -400,11 +511,11 @@ def create_and_populate_session_peers_table(schema: str, inspector) -> None: ), sa.ForeignKeyConstraint( ["peer_name", "workspace_name"], - ["peers.name", "peers.workspace_name"], + [f"{schema}.peers.name", f"{schema}.peers.workspace_name"], ), sa.ForeignKeyConstraint( ["session_name", "workspace_name"], - ["sessions.name", "sessions.workspace_name"], + [f"{schema}.sessions.name", f"{schema}.sessions.workspace_name"], ), sa.PrimaryKeyConstraint("workspace_name", "session_name", "peer_name"), ) @@ -456,7 +567,7 @@ def create_and_populate_session_peers_table(schema: str, inspector) -> None: ) -def update_messages_table(schema: str, inspector) -> None: +def update_messages_table(schema: str, inspector: Inspector) -> None: """Update messages table.""" # Add new columns @@ -618,7 +729,7 @@ def update_messages_table(schema: str, inspector) -> None: backfill_token_counts(schema) -def update_collections_table(schema: str, inspector) -> None: +def update_collections_table(schema: str, inspector: Inspector) -> None: """Update collections table.""" # Add new columns @@ -741,7 +852,7 @@ def update_collections_table(schema: str, inspector) -> None: ) -def update_documents_table(schema: str, inspector) -> None: +def update_documents_table(schema: str, inspector: Inspector) -> None: """Update documents table.""" # Add new columns @@ -873,7 +984,9 @@ def update_documents_table(schema: str, inspector) -> None: ) -def update_queue_and_active_queue_sessions_tables(schema: str, inspector) -> None: +def update_queue_and_active_queue_sessions_tables( + schema: str, inspector: Inspector +) -> None: """Update queue and active_queue_sessions tables.""" # Drop foreign key constraints first, before changing column types @@ -903,7 +1016,7 @@ def update_queue_and_active_queue_sessions_tables(schema: str, inspector) -> Non # Get the mapping of old session.id (integer) to new session.id (text, which is public_id) # At this point, sessions table still has both id (integer) and public_id (text) columns - session_id_mapping = {} + session_id_mapping: dict[int, str] = {} if table_exists("sessions", inspector): sessions_mapping = connection.execute( sa.text(f"SELECT id, public_id FROM {schema}.sessions") @@ -1081,7 +1194,7 @@ def backfill_token_counts(schema: str) -> None: break # Calculate token counts and update messages in batches - batch_updates = [] + batch_updates: list[tuple[int, int]] = [] for message_id, content in messages: token_count = _count_tokens(content) batch_updates.append((message_id, token_count)) @@ -1110,7 +1223,7 @@ def backfill_token_counts(schema: str) -> None: offset += batch_size -def restore_app_user_columns(schema: str, inspector) -> None: +def restore_app_user_columns(schema: str, inspector: Inspector) -> None: """Restore app_id and user_id columns to peers and sessions.""" # Add app_id back to peers if not column_exists("peers", "app_id", inspector): @@ -1154,8 +1267,7 @@ def restore_app_user_columns(schema: str, inspector) -> None: schema=schema, ) # Populate user_id from session_peers - # The user peer is the one that existed before agent peers were created - # Agent peers have names that match their IDs (both are nanoids), user peers have different names and IDs + # Prefer user peers (where id != name), but fall back to any peer if needed op.execute( sa.text(f""" UPDATE {schema}.sessions SET user_id = ( @@ -1172,7 +1284,7 @@ def restore_app_user_columns(schema: str, inspector) -> None: op.alter_column("sessions", "user_id", nullable=False, schema=schema) -def restore_documents_table(schema: str, inspector) -> None: +def restore_documents_table(schema: str, inspector: Inspector) -> None: """Restore documents table to pre-peer paradigm state.""" # Add back id column as primary key if not column_exists("documents", "temp_id", inspector): @@ -1306,7 +1418,7 @@ def restore_documents_table(schema: str, inspector) -> None: ) -def restore_collections_table(schema: str, inspector) -> None: +def restore_collections_table(schema: str, inspector: Inspector) -> None: """Restore collections table to pre-peer paradigm state.""" # Add back id column as primary key if not column_exists("collections", "temp_id", inspector): @@ -1433,7 +1545,7 @@ def restore_collections_table(schema: str, inspector) -> None: ) -def restore_messages_table(schema: str, inspector) -> None: +def restore_messages_table(schema: str, inspector: Inspector) -> None: """Restore messages table to pre-peer paradigm state.""" # Add back old columns if not column_exists("messages", "session_id", inspector): @@ -1553,8 +1665,20 @@ def restore_messages_table(schema: str, inspector) -> None: referent_schema=schema, ) + op.create_index( + "idx_messages_session_lookup", + "messages", + ["session_id", "id"], + postgresql_include=[ + "public_id", + "is_user", + "created_at", + ], + schema=schema, + ) -def restore_sessions_table(schema: str, inspector) -> None: + +def restore_sessions_table(schema: str, inspector: Inspector) -> None: """Restore sessions table to pre-peer paradigm state.""" # Add back id column as BigInteger primary key if not column_exists("sessions", "temp_id", inspector): @@ -1624,8 +1748,12 @@ def restore_sessions_table(schema: str, inspector) -> None: "public_id_format", "sessions", "public_id ~ '^[A-Za-z0-9_-]+$'", schema=schema ) + op.create_index( + "idx_sessions_user_lookup", "sessions", ["user_id", "public_id"], schema=schema + ) -def restore_peers_table(schema: str, inspector) -> None: + +def restore_peers_table(schema: str, inspector: Inspector) -> None: """Restore peers table to pre-user paradigm state.""" # Add back id column as BigInteger primary key if not column_exists("peers", "temp_id", inspector): @@ -1703,7 +1831,7 @@ def restore_peers_table(schema: str, inspector) -> None: ) -def restore_workspaces_table(schema: str, inspector) -> None: +def restore_workspaces_table(schema: str, inspector: Inspector) -> None: """Restore workspaces table to pre-peer paradigm state (apps).""" # Add back id column as BigInteger primary key if not column_exists("workspaces", "temp_id", inspector): @@ -1764,14 +1892,16 @@ def restore_workspaces_table(schema: str, inspector) -> None: ) -def restore_queue_and_active_queue_sessions_tables(schema: str, inspector) -> None: +def restore_queue_and_active_queue_sessions_tables( + schema: str, inspector: Inspector +) -> None: """Restore queue and active_queue_sessions tables to pre-peer paradigm state.""" connection = op.get_bind() # Create reverse mapping from session.public_id (text) back to session.id (BigInteger) # At this point in downgrade, sessions table still has both id (BigInteger) and public_id (text) - session_id_reverse_mapping = {} + session_id_reverse_mapping: dict[str, int] = {} if table_exists("sessions", inspector): sessions_mapping = connection.execute( sa.text(f"SELECT id, public_id FROM {schema}.sessions") @@ -1781,26 +1911,27 @@ def restore_queue_and_active_queue_sessions_tables(schema: str, inspector) -> No session_id_reverse_mapping[text_id] = big_int_id # Update queue table - if table_exists("queue", inspector) and session_id_reverse_mapping: + if table_exists("queue", inspector): # Get current session_id values in queue table (they are text now) - queue_session_ids = connection.execute( - sa.text( - f"SELECT DISTINCT session_id FROM {schema}.queue WHERE session_id IS NOT NULL" - ) - ).fetchall() - - # Convert session_id values back to BigInteger - for (session_id,) in queue_session_ids: - if session_id in session_id_reverse_mapping: - old_id = session_id_reverse_mapping[session_id] - connection.execute( - sa.text( - f"UPDATE {schema}.queue SET session_id = :old_id WHERE session_id = :new_id" - ), - {"old_id": str(old_id), "new_id": str(session_id)}, + if session_id_reverse_mapping: + queue_session_ids = connection.execute( + sa.text( + f"SELECT DISTINCT session_id FROM {schema}.queue WHERE session_id IS NOT NULL" ) + ).fetchall() - # Change column type back to BigInteger + # Convert session_id values back to BigInteger + for (session_id,) in queue_session_ids: + if session_id in session_id_reverse_mapping: + old_id = session_id_reverse_mapping[session_id] + connection.execute( + sa.text( + f"UPDATE {schema}.queue SET session_id = :old_id WHERE session_id = :new_id" + ), + {"old_id": str(old_id), "new_id": str(session_id)}, + ) + + # Change column type back to BigInteger (always) op.alter_column( "queue", "session_id", @@ -1810,36 +1941,40 @@ def restore_queue_and_active_queue_sessions_tables(schema: str, inspector) -> No ) # Update active_queue_sessions table - if table_exists("active_queue_sessions", inspector) and session_id_reverse_mapping: - # Get current session_id values in active_queue_sessions table (they are text now) - active_queue_session_ids = connection.execute( - sa.text( - f"SELECT DISTINCT session_id FROM {schema}.active_queue_sessions WHERE session_id IS NOT NULL" - ) - ).fetchall() + if table_exists("active_queue_sessions", inspector): + if session_id_reverse_mapping: + # Get current session_id values in active_queue_sessions table (they are text now) + active_queue_session_ids = connection.execute( + sa.text( + f"SELECT DISTINCT session_id FROM {schema}.active_queue_sessions WHERE session_id IS NOT NULL" + ) + ).fetchall() - if constraint_exists( - "active_queue_sessions", "pk_active_queue_sessions", "primary", inspector - ): - op.drop_constraint( - "pk_active_queue_sessions", + if constraint_exists( "active_queue_sessions", - type_="primary", - schema=schema, - ) - - # Convert session_id values back to BigInteger - for (session_id,) in active_queue_session_ids: - if session_id in session_id_reverse_mapping: - old_id = session_id_reverse_mapping[session_id] - connection.execute( - sa.text( - f"UPDATE {schema}.active_queue_sessions SET session_id = :old_id WHERE session_id = :new_id" - ), - {"old_id": str(old_id), "new_id": str(session_id)}, + "pk_active_queue_sessions", + "primary", + inspector, + ): + op.drop_constraint( + "pk_active_queue_sessions", + "active_queue_sessions", + type_="primary", + schema=schema, ) - # Change column type back to BigInteger + # Convert session_id values back to BigInteger + for (session_id,) in active_queue_session_ids: + if session_id in session_id_reverse_mapping: + old_id = session_id_reverse_mapping[session_id] + connection.execute( + sa.text( + f"UPDATE {schema}.active_queue_sessions SET session_id = :old_id WHERE session_id = :new_id" + ), + {"old_id": str(old_id), "new_id": str(session_id)}, + ) + + # Change column type back to BigInteger (always) op.alter_column( "active_queue_sessions", "session_id", @@ -1848,12 +1983,16 @@ def restore_queue_and_active_queue_sessions_tables(schema: str, inspector) -> No postgresql_using="session_id::bigint", ) - op.create_primary_key( - "pk_active_queue_sessions", - "active_queue_sessions", - ["session_id"], - schema=schema, - ) + # Ensure primary key exists on session_id for pre-peer shape + if not constraint_exists( + "active_queue_sessions", "pk_active_queue_sessions", "primary", inspector + ): + op.create_primary_key( + "pk_active_queue_sessions", + "active_queue_sessions", + ["session_id"], + schema=schema, + ) if constraint_exists( "active_queue_sessions", @@ -1877,29 +2016,10 @@ def restore_queue_and_active_queue_sessions_tables(schema: str, inspector) -> No if column_exists("active_queue_sessions", "task_type", inspector): op.drop_column("active_queue_sessions", "task_type", schema=schema) - # Restore foreign key constraints - if table_exists("queue", inspector): - op.create_foreign_key( - "fk_queue_session_id_sessions", - "queue", - "sessions", - ["session_id"], - ["id"], - referent_schema=schema, - ) - - if table_exists("active_queue_sessions", inspector): - op.create_foreign_key( - "fk_active_queue_sessions_session_id_sessions", - "active_queue_sessions", - "sessions", - ["session_id"], - ["id"], - referent_schema=schema, - ) + # Defer restoring foreign key constraints to restore_foreign_keys() -def restore_table_names(schema: str, inspector) -> None: +def restore_table_names(schema: str, inspector: Inspector) -> None: """Restore table names: workspaces->apps and peers->users.""" if inspector.has_table("workspaces", schema=schema): op.rename_table("workspaces", "apps", schema=schema) @@ -1957,6 +2077,23 @@ def restore_foreign_keys(schema: str) -> None: ["public_id"], referent_schema=schema, ) + # Restore queue/session FKs last, after all type changes are complete + op.create_foreign_key( + "fk_queue_session_id_sessions", + "queue", + "sessions", + ["session_id"], + ["id"], + referent_schema=schema, + ) + op.create_foreign_key( + "fk_active_queue_sessions_session_id_sessions", + "active_queue_sessions", + "sessions", + ["session_id"], + ["id"], + referent_schema=schema, + ) op.create_foreign_key( "fk_sessions_app_id_apps", "sessions", diff --git a/pyproject.toml b/pyproject.toml index 9b0de576..0022cd1e 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -81,7 +81,7 @@ extend-immutable-calls = ["fastapi.Depends"] [tool.pytest.ini_options] asyncio_mode = "auto" asyncio_default_fixture_loop_scope = "session" -addopts = "--strict-markers --cov=src/ --cov=sdks/python/src/honcho --cov-report=term-missing" +addopts = "--strict-markers --cov=src/ --cov=sdks/python/src/honcho --cov-report=term-missing --ignore=tests/alembic" testpaths = ["tests"] pythonpath = ["src"] diff --git a/scripts/ensure_alembic_tests.py b/scripts/ensure_alembic_tests.py new file mode 100644 index 00000000..fe3cad8b --- /dev/null +++ b/scripts/ensure_alembic_tests.py @@ -0,0 +1,68 @@ +#!/usr/bin/env uv run python +""" +Script that validates that all alembic migration revisions have a corresponding test file. +Note that this script is actively used within our precommit hooks and should not be removed. +If this script is moved, the corresponding precommit hook will need to be updated. +""" + +import sys +from pathlib import Path + +PROJECT_ROOT = Path(__file__).resolve().parents[1] +MIGRATIONS_DIR = PROJECT_ROOT / "migrations" / "versions" +TESTS_DIR = PROJECT_ROOT / "tests" / "alembic" / "revisions" + + +def main() -> int: + missing: list[str] = [] + + # Gather migration base names (without extension) + migration_files = [ + p for p in MIGRATIONS_DIR.glob("*.py") if p.name != "__init__.py" + ] + migration_basenames = {p.stem for p in migration_files} + print(f"Migration basenames: {migration_basenames}") + + # Gather test file base names mapped back to migration names by stripping leading 'test_' + test_files = [p for p in TESTS_DIR.glob("test_*.py") if p.name != "__init__.py"] + test_targets = {p.stem.removeprefix("test_") for p in test_files} + + for migration_basename in sorted(migration_basenames): + if migration_basename not in test_targets: + missing.append(migration_basename) + if missing: + print( + "Missing Alembic tests for the following migration revisions:", + file=sys.stderr, + ) + for name in missing: + print(f" - {name}", file=sys.stderr) + print( + "\nExpected test files under tests/alembic/revisions named as:", + file=sys.stderr, + ) + for name in missing: + print(f" - tests/alembic/revisions/test_{name}.py", file=sys.stderr) + print("\nScaffold helper commands:", file=sys.stderr) + for name in missing: + print( + f" - python -m tests.alembic.scaffold {name.split('_', 1)[0]}", + file=sys.stderr, + ) + return 1 + + # Optional: warn if there are tests without corresponding migrations (stale tests) + stale_tests = sorted(test_targets - migration_basenames) + if stale_tests: + print( + "Warning: Found tests without corresponding migration files (stale?):", + file=sys.stderr, + ) + for name in stale_tests: + print(f" - tests/alembic/revisions/test_{name}.py", file=sys.stderr) + + return 0 + + +if __name__ == "__main__": + raise SystemExit(main()) diff --git a/scripts/generate_jwt_secret.py b/scripts/generate_jwt_secret.py index 449d4cc7..6a7dc5a4 100755 --- a/scripts/generate_jwt_secret.py +++ b/scripts/generate_jwt_secret.py @@ -1,4 +1,4 @@ -#!/usr/bin/env python +#!/usr/bin/env uv run python """ Utility script to generate a JWT secret for use in the .env file. This uses the same logic as the automatically generated version in security.py. diff --git a/scripts/generate_message_embeddings.py b/scripts/generate_message_embeddings.py index fc19c3fe..75c3f22f 100644 --- a/scripts/generate_message_embeddings.py +++ b/scripts/generate_message_embeddings.py @@ -1,3 +1,4 @@ +#!/usr/bin/env uv run python """ Script to generate embeddings for existing messages that don't already have embeddings. diff --git a/src/models.py b/src/models.py index 1cd3b00c..b66e31cf 100644 --- a/src/models.py +++ b/src/models.py @@ -113,7 +113,7 @@ class Peer(Base): DateTime(timezone=True), index=True, default=func.now() ) workspace_name: Mapped[str] = mapped_column( - ForeignKey("workspaces.name"), index=True + ForeignKey("workspaces.name"), index=True, nullable=False ) configuration: Mapped[dict[str, Any]] = mapped_column(JSONB, default=dict) @@ -149,7 +149,7 @@ class Session(Base): ) messages = relationship("Message", back_populates="session") workspace_name: Mapped[str] = mapped_column( - ForeignKey("workspaces.name"), index=True + ForeignKey("workspaces.name"), index=True, nullable=False ) configuration: Mapped[dict[str, Any]] = mapped_column(JSONB, default=dict) diff --git a/tests/alembic/README.md b/tests/alembic/README.md new file mode 100644 index 00000000..2dec3d36 --- /dev/null +++ b/tests/alembic/README.md @@ -0,0 +1,29 @@ +These tests validate Alembic migrations end-to-end for structure, order, and correctness. +They ensure reversibility, expected schema and data, and integration with the registry and pipeline. +The key components are the verifier, the test pipeline, the registry, and the revisions under test. + +### Verifier + +- The verifier runs checks when specific revisions are applied and reverted. +- It validates the schema after upgrade, verifies data migrations such as defaults, backfills, and transforms, and confirms reversibility after downgrade. +- Assertions are grouped per-revision or feature, and helpers use the SQLAlchemy inspector to introspect the database. + +### Test Pipeline + +- The pipeline orchestrates the database lifecycle by creating a database, applying upgrades and downgrades, running verifications, and tearing down resources. +- It typically starts from base, upgrades to the revision immediately before a target revision, seeds the DB, runs the target migration, and then validates the schema + data +- It relies on shared fixtures such as `engine`, `connection`, and `alembic_config`, and it ensures isolation per test + +### Registry + +- The registry declares revisions and test metadata used to drive scenarios. +- It defines ordering and selection, attaches verifier callbacks to specific revisions or ranges, and centralizes per revision expectations. + +### Revisions + +- Revisions are the migration scripts under `alembic/revisions`. +- Each revision should provide functions decorated with `register_before_upgrade()` and `register_after_upgrade()`. These are used to validate schemas and data before and after a migration is run + +### Running the Tests + +- Tests can be run all together `pytest tests/alembic` or individually `pytest tests/alembic -k "revision_number"`. For example, to run the test against a1b2c3d4e5f6_initial_schema.py, you would run the command `pytest tests/alembic -k "a1b2c3d4e5f6"` diff --git a/tests/alembic/__init__.py b/tests/alembic/__init__.py new file mode 100644 index 00000000..571f7c81 --- /dev/null +++ b/tests/alembic/__init__.py @@ -0,0 +1,5 @@ +"""Test package for Alembic migration scenarios.""" + +from . import revisions + +__all__ = ["revisions"] diff --git a/tests/alembic/conftest.py b/tests/alembic/conftest.py new file mode 100644 index 00000000..022ae5fc --- /dev/null +++ b/tests/alembic/conftest.py @@ -0,0 +1,107 @@ +"""Pytest configuration for Alembic-focused tests.""" + +from __future__ import annotations + +import os +from collections.abc import Generator +from pathlib import Path + +import pytest +from alembic.config import Config +from sqlalchemy import create_engine, text +from sqlalchemy.engine import Engine +from sqlalchemy.engine.url import URL +from sqlalchemy_utils import ( + create_database, # pyright: ignore + database_exists, # pyright: ignore + drop_database, # pyright: ignore +) + +from src.config import settings +from tests.conftest import CONNECTION_URI + +ALEMBIC_CONFIG_PATH = Path(__file__).resolve().parents[2] / "alembic.ini" +ALEMBIC_TEST_DB_URL: URL = CONNECTION_URI.set(database="alembic_migration_tests") + + +@pytest.fixture(scope="session", autouse=True) +def configure_alembic_settings(alembic_database: str) -> Generator[None, None, None]: + """Point application settings at the Alembic test database.""" + + previous_uri = settings.DB.CONNECTION_URI + os.environ["DB_CONNECTION_URI"] = alembic_database + settings.DB.CONNECTION_URI = alembic_database + + try: + yield + finally: + settings.DB.CONNECTION_URI = previous_uri + if previous_uri: + os.environ["DB_CONNECTION_URI"] = previous_uri + else: + os.environ.pop("DB_CONNECTION_URI", None) + + +@pytest.fixture +def alembic_cfg(alembic_database: str) -> Config: + """Provide an Alembic Config bound to the alembic test database.""" + cfg = Config(str(ALEMBIC_CONFIG_PATH)) + cfg.set_main_option( + "script_location", str(ALEMBIC_CONFIG_PATH.parent / "migrations") + ) + cfg.set_main_option("sqlalchemy.url", alembic_database) + return cfg + + +@pytest.fixture(scope="session") +def alembic_database() -> Generator[str, None, None]: + """Provision a dedicated DB for Alembic verification tests.""" + + assert ALEMBIC_TEST_DB_URL.database == "alembic_migration_tests", ( + "Can't set up Alembic test database fixture. " + + "ALEMBIC_TEST_DB_URL.database is {ALEMBIC_TEST_DB_URL.database}, " + + "expected 'alembic_migration_tests'. " + ) + + if database_exists(ALEMBIC_TEST_DB_URL): + drop_database(ALEMBIC_TEST_DB_URL) # start fresh + create_database(ALEMBIC_TEST_DB_URL) + + engine = create_engine(str(ALEMBIC_TEST_DB_URL)) + try: + with engine.begin() as conn: + conn.execute(text("CREATE EXTENSION IF NOT EXISTS vector")) + yield str(ALEMBIC_TEST_DB_URL) + finally: + engine.dispose() + if database_exists(ALEMBIC_TEST_DB_URL): + drop_database(ALEMBIC_TEST_DB_URL) + + +@pytest.fixture +def alembic_engine(alembic_database: str) -> Generator[Engine, None, None]: + """Yield an engine bound to the Alembic test database.""" + + engine = create_engine(alembic_database, pool_pre_ping=True) + try: + yield engine + finally: + engine.dispose() + + +@pytest.fixture(autouse=True) +def reset_schema_between_tests( + alembic_engine: Engine, +) -> Generator[None, None, None]: + """Drop and recreate the schema for a clean slate each test (fast reset).""" + + schema = settings.DB.SCHEMA + + def _reset_schema() -> None: + with alembic_engine.begin() as conn: + conn.execute(text(f'DROP SCHEMA IF EXISTS "{schema}" CASCADE')) + conn.execute(text(f'CREATE SCHEMA "{schema}"')) + + _reset_schema() + yield + _reset_schema() diff --git a/tests/alembic/registry.py b/tests/alembic/registry.py new file mode 100644 index 00000000..d2b013bc --- /dev/null +++ b/tests/alembic/registry.py @@ -0,0 +1,73 @@ +"""Registry for migration-specific prepare and verification hooks.""" + +from __future__ import annotations + +from collections.abc import Callable, Mapping +from dataclasses import dataclass +from typing import TYPE_CHECKING + +if TYPE_CHECKING: + from tests.alembic.verifier import MigrationVerifier + + +BeforeUpgradeHook = Callable[["MigrationVerifier"], None] +AfterUpgradeHook = Callable[["MigrationVerifier"], None] + + +@dataclass(slots=True) +class RevisionHooks: + """Container for lifecycle hooks tied to a revision.""" + + before_upgrade: BeforeUpgradeHook | None = None + after_upgrade: AfterUpgradeHook | None = None + + +_REGISTRY: dict[str, RevisionHooks] = {} + + +def register_before_upgrade( + revision: str, +) -> Callable[[BeforeUpgradeHook], BeforeUpgradeHook]: + """Register a callable to execute before upgrading to ``revision``.""" + + def decorator(func: BeforeUpgradeHook) -> BeforeUpgradeHook: + hooks = _REGISTRY.setdefault(revision, RevisionHooks()) + if hooks.before_upgrade is not None: + raise ValueError( + f"before_upgrade hook already registered for revision {revision}" + ) + hooks.before_upgrade = func + return func + + return decorator + + +def register_after_upgrade( + revision: str, +) -> Callable[[AfterUpgradeHook], AfterUpgradeHook]: + """Register a callable to execute after upgrading to ``revision``.""" + + def decorator(func: AfterUpgradeHook) -> AfterUpgradeHook: + hooks = _REGISTRY.setdefault(revision, RevisionHooks()) + if hooks.after_upgrade is not None: + raise ValueError( + f"after_upgrade hook already registered for revision {revision}" + ) + hooks.after_upgrade = func + return func + + return decorator + + +def get_registered_hooks() -> Mapping[str, RevisionHooks]: + """Return a snapshot of the registered revision hooks.""" + + return dict(_REGISTRY) + + +__all__ = [ + "RevisionHooks", + "register_before_upgrade", + "register_after_upgrade", + "get_registered_hooks", +] diff --git a/tests/alembic/revisions/__init__.py b/tests/alembic/revisions/__init__.py new file mode 100644 index 00000000..359e7b1f --- /dev/null +++ b/tests/alembic/revisions/__init__.py @@ -0,0 +1,35 @@ +"""Register revision-specific hooks for migration verification.""" + +from . import ( + test_05486ce795d5_make_session_name_required_on_messages, + test_08894082221a_replace_collection_name_with_observer_, + test_20f89a421aff_rename_metamessage_type_to_label, + test_66e63cf2cf77_add_indexes_to_documents_table, + test_76ffba56fe8c_add_error_field_to_queueitem, + test_88b0fb10906f_add_webhooks_table, + test_556a16564f50_add_user_id_and_app_id_to_tables, + test_564ba40505c5_add_session_name_column_to_documents, + test_917195d9b5e9_add_messageembedding_table, + test_a1b2c3d4e5f6_initial_schema, + test_b765d82110bd_change_metamessages_to_user_level_with_, + test_bb6fb3a7a643_add_message_seq_in_session_column, + test_c3828084f472_add_indexes_for_messages_and_, + test_d429de0e5338_adopt_peer_paradigm, +) + +__all__ = [ + "test_05486ce795d5_make_session_name_required_on_messages", + "test_08894082221a_replace_collection_name_with_observer_", + "test_20f89a421aff_rename_metamessage_type_to_label", + "test_556a16564f50_add_user_id_and_app_id_to_tables", + "test_564ba40505c5_add_session_name_column_to_documents", + "test_66e63cf2cf77_add_indexes_to_documents_table", + "test_76ffba56fe8c_add_error_field_to_queueitem", + "test_88b0fb10906f_add_webhooks_table", + "test_917195d9b5e9_add_messageembedding_table", + "test_a1b2c3d4e5f6_initial_schema", + "test_b765d82110bd_change_metamessages_to_user_level_with_", + "test_bb6fb3a7a643_add_message_seq_in_session_column", + "test_c3828084f472_add_indexes_for_messages_and_", + "test_d429de0e5338_adopt_peer_paradigm", +] diff --git a/tests/alembic/revisions/test_05486ce795d5_make_session_name_required_on_messages.py b/tests/alembic/revisions/test_05486ce795d5_make_session_name_required_on_messages.py new file mode 100644 index 00000000..093ed17b --- /dev/null +++ b/tests/alembic/revisions/test_05486ce795d5_make_session_name_required_on_messages.py @@ -0,0 +1,104 @@ +"""Hooks for revision 05486ce795d5 (session_name required).""" + +from __future__ import annotations + +from nanoid import generate as generate_nanoid +from sqlalchemy import text + +from tests.alembic.registry import register_after_upgrade, register_before_upgrade +from tests.alembic.verifier import MigrationVerifier + +WORKSPACE_ID = generate_nanoid() +WORKSPACE_NAME = "workspace-name" +PEER_ID = generate_nanoid() +PEER_NAME = "peer-name" +SESSION_ID = generate_nanoid() +SESSION_NAME = "session-name" +MESSAGE_ID = generate_nanoid() +MESSAGE_CONTENT = "message-content" + + +@register_before_upgrade("05486ce795d5") +def seed_orphaned_message(verifier: MigrationVerifier) -> None: + """Insert a message without a session_name to exercise the migration.""" + + schema = verifier.schema + connection = verifier.conn + + connection.execute( + text( + f""" + INSERT INTO "{schema}"."workspaces" ("id", "name") + VALUES (:workspace_id, :workspace_name) + """ + ), + {"workspace_id": WORKSPACE_ID, "workspace_name": WORKSPACE_NAME}, + ) + + connection.execute( + text( + f""" + INSERT INTO "{schema}"."peers" ("id", "name", "workspace_name") + VALUES (:peer_id, :peer_name, :workspace_name) + """ + ), + { + "peer_id": PEER_ID, + "peer_name": PEER_NAME, + "workspace_name": WORKSPACE_NAME, + }, + ) + + connection.execute( + text( + f""" + INSERT INTO "{schema}"."messages" + ("public_id", "workspace_name", "peer_name", "session_name", "content") + VALUES (:public_id, :workspace_name, :peer_name, NULL, :content) + """ + ), + { + "public_id": MESSAGE_ID, + "workspace_name": WORKSPACE_NAME, + "peer_name": PEER_NAME, + "content": MESSAGE_CONTENT, + }, + ) + + +@register_after_upgrade("05486ce795d5") +def verify_session_name_enforced(verifier: MigrationVerifier) -> None: + """Ensure session_name is populated and constrained after the migration.""" + + verifier.assert_column_exists("messages", "session_name", nullable=False) + + schema = verifier.schema + conn = verifier.conn + + session = conn.execute( + text( + f'SELECT "name", "workspace_name" FROM "{schema}"."sessions" ' + + 'WHERE "name" = :session_name AND "workspace_name" = :workspace_name' + ), + { + "session_name": PEER_NAME, + "workspace_name": WORKSPACE_NAME, + }, + ).one_or_none() + assert ( + session is not None + ), "Session with expected name and workspace does not exist" + + message = conn.execute( + text( + f'SELECT "public_id", "content" FROM "{schema}"."messages" ' + + 'WHERE "session_name" = :session_name AND "workspace_name" = :workspace_name AND "peer_name" = :peer_name' + ), + { + "session_name": PEER_NAME, + "workspace_name": WORKSPACE_NAME, + "peer_name": PEER_NAME, + }, + ).one_or_none() + assert message is not None, "Message with expected session_name does not exist" + assert message.content == MESSAGE_CONTENT diff --git a/tests/alembic/revisions/test_08894082221a_replace_collection_name_with_observer_.py b/tests/alembic/revisions/test_08894082221a_replace_collection_name_with_observer_.py new file mode 100644 index 00000000..30f9aeb6 --- /dev/null +++ b/tests/alembic/revisions/test_08894082221a_replace_collection_name_with_observer_.py @@ -0,0 +1,361 @@ +"""Hooks for revision 08894082221a (observer/observed refactor).""" + +from __future__ import annotations + +from nanoid import generate as generate_nanoid +from sqlalchemy import inspect, text + +from tests.alembic.registry import register_after_upgrade, register_before_upgrade +from tests.alembic.verifier import MigrationVerifier + +COLLECTION_ID: str = generate_nanoid() +GLOBAL_REP_COLLECTION_NAME = "global_representation" +DOCUMENT_ID = generate_nanoid() +DOCUMENT_NAME = "document-name" +LEGACY_PEER_ID = generate_nanoid() +LEGACY_PEER_NAME = "legacy-peer" +WORKSPACE_ID = generate_nanoid() +WORKSPACE_NAME = "workspace-name" + + +# Based on migration logic: when name='global_representation', observer=peer_name and observed=peer_name +EXPECTED_OBSERVED = LEGACY_PEER_NAME + + +@register_before_upgrade("08894082221a") +def prepare_observer_observed(verifier: MigrationVerifier) -> None: + schema = verifier.schema + conn = verifier.conn + + # Create workspace + conn.execute( + text( + f'INSERT INTO "{schema}"."workspaces" ("id", "name") ' + + "VALUES (:workspace_id, :workspace_name)" + ), + {"workspace_id": WORKSPACE_ID, "workspace_name": WORKSPACE_NAME}, + ) + + # Create peer + conn.execute( + text( + f'INSERT INTO "{schema}"."peers" ("id", "name", "workspace_name") ' + + "VALUES (:peer_id, :peer_name, :workspace_name)" + ), + { + "peer_id": LEGACY_PEER_ID, + "peer_name": LEGACY_PEER_NAME, + "workspace_name": WORKSPACE_NAME, + }, + ) + + # Create collection with the old schema (has 'name' field) + conn.execute( + text( + f'INSERT INTO "{schema}"."collections" ' + + '("id", "name", "peer_name", "workspace_name") ' + + "VALUES (:collection_id, :collection_name, :peer_name, :workspace_name)" + ), + { + "collection_id": COLLECTION_ID, + "collection_name": GLOBAL_REP_COLLECTION_NAME, + "peer_name": LEGACY_PEER_NAME, + "workspace_name": WORKSPACE_NAME, + }, + ) + + # Create document with the old schema (has 'collection_name' field) without session_name + conn.execute( + text( + f'INSERT INTO "{schema}"."documents" ' + + '("id", "collection_name", "peer_name", "workspace_name", "content") ' + + "VALUES (:document_id, :collection_name, :peer_name, :workspace_name, :content)" + ), + { + "document_id": DOCUMENT_ID, + "collection_name": GLOBAL_REP_COLLECTION_NAME, + "peer_name": LEGACY_PEER_NAME, + "workspace_name": WORKSPACE_NAME, + "content": "test content", + }, + ) + + # Bulk create many documents with the old schema (no session_name) + bulk_stmt = text( + f'INSERT INTO "{schema}"."documents" ' + + '("id", "collection_name", "peer_name", "workspace_name", "content") ' + + "VALUES (:document_id, :collection_name, :peer_name, :workspace_name, :content)" + ) + + params_template = { + "collection_name": GLOBAL_REP_COLLECTION_NAME, + "peer_name": LEGACY_PEER_NAME, + "workspace_name": WORKSPACE_NAME, + "content": "test content", + } + + total_docs = 120_000 + batch_size = 10_000 + for start in range(0, total_docs, batch_size): + end = min(start + batch_size, total_docs) + batch_params = [ + {"document_id": generate_nanoid(), **params_template} + for _ in range(start, end) + ] + + conn.execute(bulk_stmt, batch_params) + + # Speed up bulk seeding + conn.execute(text("SET LOCAL synchronous_commit = OFF")) + + # ----------------------------- + # Group 1: self-observation collections (global_representation) + # observer = peer_name, observed = peer_name + # ----------------------------- + peer_insert_stmt = text( + f'INSERT INTO "{schema}"."peers" ("id", "name", "workspace_name") ' + + "VALUES (:peer_id, :peer_name, :workspace_name)" + ) + collection_insert_stmt = text( + f'INSERT INTO "{schema}"."collections" ("id", "name", "peer_name", "workspace_name") ' + + "VALUES (:collection_id, :name, :peer_name, :workspace_name)" + ) + + total = 100_000 + batch_size = 10_000 + + for start in range(0, total, batch_size): + end = min(start + batch_size, total) + # Insert peers for group 1 + peers_params = [ + { + "peer_id": generate_nanoid(), + "peer_name": f"selfpeer_{i}", + "workspace_name": WORKSPACE_NAME, + } + for i in range(start, end) + ] + conn.execute(peer_insert_stmt, peers_params) + + # Insert collections for group 1 + collections_params = [ + { + "collection_id": generate_nanoid(), + "name": GLOBAL_REP_COLLECTION_NAME, + "peer_name": f"selfpeer_{i}", + "workspace_name": WORKSPACE_NAME, + } + for i in range(start, end) + ] + conn.execute(collection_insert_stmt, collections_params) + + # ----------------------------- + # Group 2: prefix pattern (observer_observed) + # observer = 'prefix_observer', observed extracted from name suffix + # ----------------------------- + # Create the fixed observer peer + conn.execute( + peer_insert_stmt, + { + "peer_id": generate_nanoid(), + "peer_name": "prefix_observer", + "workspace_name": WORKSPACE_NAME, + }, + ) + + for start in range(0, total, batch_size): + end = min(start + batch_size, total) + # Observed peers for group 2 + peers_params = [ + { + "peer_id": generate_nanoid(), + "peer_name": f"obs_prefix_{i}", + "workspace_name": WORKSPACE_NAME, + } + for i in range(start, end) + ] + conn.execute(peer_insert_stmt, peers_params) + + # Collections for group 2 (name = 'prefix_observer_obs_prefix_') + collections_params = [ + { + "collection_id": generate_nanoid(), + "name": f"prefix_observer_obs_prefix_{i}", + "peer_name": "prefix_observer", + "workspace_name": WORKSPACE_NAME, + } + for i in range(start, end) + ] + conn.execute(collection_insert_stmt, collections_params) + + # ----------------------------- + # Group 3: suffix pattern (observed_observer) + # observer = 'suffix_observer', observed extracted from name prefix + # ----------------------------- + # Create the fixed observer peer + conn.execute( + peer_insert_stmt, + { + "peer_id": generate_nanoid(), + "peer_name": "suffix_observer", + "workspace_name": WORKSPACE_NAME, + }, + ) + + for start in range(0, total, batch_size): + end = min(start + batch_size, total) + # Observed peers for group 3 + peers_params = [ + { + "peer_id": generate_nanoid(), + "peer_name": f"obs_suffix_{i}", + "workspace_name": WORKSPACE_NAME, + } + for i in range(start, end) + ] + conn.execute(peer_insert_stmt, peers_params) + + # Collections for group 3 (name = 'obs_suffix__suffix_observer') + collections_params = [ + { + "collection_id": generate_nanoid(), + "name": f"obs_suffix_{i}_suffix_observer", + "peer_name": "suffix_observer", + "workspace_name": WORKSPACE_NAME, + } + for i in range(start, end) + ] + conn.execute(collection_insert_stmt, collections_params) + + +@register_after_upgrade("08894082221a") +def verify_observer_observed_migration(verifier: MigrationVerifier) -> None: + """Assert that collections/documents now rely on observer/observed fields.""" + + inspector = inspect(verifier.conn) + + collection_columns = { + col["name"] + for col in inspector.get_columns("collections", schema=verifier.schema) + } + assert "observer" in collection_columns + assert "observed" in collection_columns + assert "name" not in collection_columns + + document_columns = { + col["name"] + for col in inspector.get_columns("documents", schema=verifier.schema) + } + assert "observer" in document_columns + assert "observed" in document_columns + assert "collection_name" not in document_columns + + verifier.assert_indexes_exist( + [ + ("collections", "idx_collections_observer"), + ("collections", "idx_collections_observed"), + ("documents", "idx_documents_observer"), + ("documents", "idx_documents_observed"), + ] + ) + verifier.assert_constraint_exists( + "collections", "unique_observer_observed_collection", "unique" + ) + verifier.assert_constraint_exists( + "documents", "documents_observer_observed_workspace_name_fkey", "foreign_key" + ) + + collection = verifier.conn.execute( + text( + 'SELECT "observer", "observed", "workspace_name" ' + + f'FROM "{verifier.schema}"."collections" ' + + 'WHERE "id" = :collection_id' + ), + {"collection_id": COLLECTION_ID}, + ).one() + assert collection.observer == LEGACY_PEER_NAME + assert collection.observed == EXPECTED_OBSERVED + assert collection.workspace_name == WORKSPACE_NAME + + document = verifier.conn.execute( + text( + 'SELECT "observer", "observed", "workspace_name", "session_name" ' + + f'FROM "{verifier.schema}"."documents" ' + + 'WHERE "id" = :document_id' + ), + {"document_id": DOCUMENT_ID}, + ).one() + assert document.observer == LEGACY_PEER_NAME + assert document.observed == EXPECTED_OBSERVED + assert document.workspace_name == WORKSPACE_NAME + assert document.session_name == "__global_observations__" + + # Verify no NULLs remain in documents.session_name + verifier.assert_no_nulls("documents", "session_name") + + # Verify collections mapping for Group 1 (self-observation) + count_self = verifier.conn.execute( + text( + "SELECT COUNT(1) FROM " + + f'"{verifier.schema}"."collections" ' + + 'WHERE "workspace_name" = :ws ' + + 'AND "observer" LIKE :prefix ' + + 'AND "observed" = "observer"' + ), + {"ws": WORKSPACE_NAME, "prefix": "selfpeer_%"}, + ).scalar() + assert count_self == 100_000 + + # Verify collections mapping for Group 2 (prefix pattern observer_observed) + count_prefix = verifier.conn.execute( + text( + "SELECT COUNT(1) FROM " + + f'"{verifier.schema}"."collections" ' + + 'WHERE "workspace_name" = :ws ' + + 'AND "observer" = :observer ' + + 'AND "observed" LIKE :obs_prefix' + ), + { + "ws": WORKSPACE_NAME, + "observer": "prefix_observer", + "obs_prefix": "obs_prefix_%", + }, + ).scalar() + assert count_prefix == 100_000 + + distinct_prefix_observed = verifier.conn.execute( + text( + 'SELECT COUNT(DISTINCT "observed") FROM ' + + f'"{verifier.schema}"."collections" ' + + 'WHERE "workspace_name" = :ws AND "observer" = :observer' + ), + {"ws": WORKSPACE_NAME, "observer": "prefix_observer"}, + ).scalar() + assert distinct_prefix_observed == 100_000 + + # Verify collections mapping for Group 3 (suffix pattern observed_observer) + count_suffix = verifier.conn.execute( + text( + "SELECT COUNT(1) FROM " + + f'"{verifier.schema}"."collections" ' + + 'WHERE "workspace_name" = :ws ' + + 'AND "observer" = :observer ' + + 'AND "observed" LIKE :obs_prefix' + ), + { + "ws": WORKSPACE_NAME, + "observer": "suffix_observer", + "obs_prefix": "obs_suffix_%", + }, + ).scalar() + assert count_suffix == 100_000 + + distinct_suffix_observed = verifier.conn.execute( + text( + 'SELECT COUNT(DISTINCT "observed") FROM ' + + f'"{verifier.schema}"."collections" ' + + 'WHERE "workspace_name" = :ws AND "observer" = :observer' + ), + {"ws": WORKSPACE_NAME, "observer": "suffix_observer"}, + ).scalar() + assert distinct_suffix_observed == 100_000 diff --git a/tests/alembic/revisions/test_20f89a421aff_rename_metamessage_type_to_label.py b/tests/alembic/revisions/test_20f89a421aff_rename_metamessage_type_to_label.py new file mode 100644 index 00000000..399e70f5 --- /dev/null +++ b/tests/alembic/revisions/test_20f89a421aff_rename_metamessage_type_to_label.py @@ -0,0 +1,123 @@ +"""Hooks for revision 20f89a421aff (metamessage label rename).""" + +from __future__ import annotations + +from nanoid import generate as generate_nanoid +from sqlalchemy import text + +from tests.alembic.registry import register_after_upgrade, register_before_upgrade +from tests.alembic.verifier import MigrationVerifier + +APP_ID = generate_nanoid() +USER_ID = generate_nanoid() +SESSION_ID = generate_nanoid() +MESSAGE_ID = generate_nanoid() +METAMESSAGE_ID = generate_nanoid() + + +@register_before_upgrade("20f89a421aff") +def prepare_metamessage_label(verifier: MigrationVerifier) -> None: + OLD_INDEXES = ( + ("metamessages", "idx_metamessages_lookup"), + ("metamessages", "idx_metamessages_user_lookup"), + ("metamessages", "idx_metamessages_session_lookup"), + ("metamessages", "idx_metamessages_message_lookup"), + ) + verifier.assert_indexes_exist(OLD_INDEXES) + verifier.assert_constraint_exists( + "metamessages", "metamessage_type_length", "check" + ) + + schema = verifier.schema + connection = verifier.conn + inspector = verifier.get_inspector() + + connection.execute( + text( + f'INSERT INTO "{schema}"."apps" ("public_id", "name") ' + + "VALUES (:id, :name)" + ), + {"id": APP_ID, "name": "verify-app"}, + ) + + connection.execute( + text( + f'INSERT INTO "{schema}"."users" ("public_id", "name", "app_id") ' + + "VALUES (:id, :name, :app_id)" + ), + {"id": USER_ID, "name": "verify-user", "app_id": APP_ID}, + ) + + connection.execute( + text( + f'INSERT INTO "{schema}"."sessions" ' + + '("public_id", "user_id", "app_id", "is_active") ' + + "VALUES (:id, :user_id, :app_id, true)" + ), + {"id": SESSION_ID, "user_id": USER_ID, "app_id": APP_ID}, + ) + + connection.execute( + text( + f'INSERT INTO "{schema}"."messages" ' + + '("public_id", "session_id", "app_id", "user_id", "is_user", "content") ' + + "VALUES (:id, :session_id, :app_id, :user_id, true, :content)" + ), + { + "id": MESSAGE_ID, + "session_id": SESSION_ID, + "app_id": APP_ID, + "user_id": USER_ID, + "content": "legacy metamessage seed", + }, + ) + + connection.execute( + text( + f'INSERT INTO "{schema}"."metamessages" ' + + '("public_id", "metamessage_type", "content", "message_id", "user_id", "session_id", "app_id") ' + + "VALUES (:id, :type, :content, :message_id, :user_id, :session_id, :app_id)" + ), + { + "id": METAMESSAGE_ID, + "type": "seed", + "content": "legacy metamessage", + "message_id": MESSAGE_ID, + "user_id": USER_ID, + "session_id": SESSION_ID, + "app_id": APP_ID, + }, + ) + + columns = { + col["name"] for col in inspector.get_columns("metamessages", schema=schema) + } + assert "metamessage_type" in columns + assert "label" not in columns + + +@register_after_upgrade("20f89a421aff") +def verify_metamessage_label(verifier: MigrationVerifier) -> None: + NEW_INDEXES = ( + ("metamessages", "idx_metamessages_lookup"), + ("metamessages", "idx_metamessages_user_lookup"), + ("metamessages", "idx_metamessages_session_lookup"), + ("metamessages", "idx_metamessages_message_lookup"), + ) + verifier.assert_indexes_exist(NEW_INDEXES) + verifier.assert_column_exists("metamessages", "label", nullable=False) + verifier.assert_constraint_exists("metamessages", "label_length", "check") + + verifier.assert_constraint_exists( + "metamessages", "metamessage_type_length", "check", exists=False + ) + verifier.assert_column_exists("metamessages", "metamessage_type", exists=False) + + row = verifier.conn.execute( + text( + f'SELECT "label" FROM "{verifier.schema}"."metamessages" ' + + 'WHERE "public_id" = :public_id' + ), + {"public_id": METAMESSAGE_ID}, + ).one() + assert row.label == "seed" diff --git a/tests/alembic/revisions/test_556a16564f50_add_user_id_and_app_id_to_tables.py b/tests/alembic/revisions/test_556a16564f50_add_user_id_and_app_id_to_tables.py new file mode 100644 index 00000000..c0e56afa --- /dev/null +++ b/tests/alembic/revisions/test_556a16564f50_add_user_id_and_app_id_to_tables.py @@ -0,0 +1,185 @@ +"""Hooks for revision 556a16564f50 (propagate app/user identifiers).""" + +from __future__ import annotations + +from nanoid import generate as generate_nanoid +from sqlalchemy import text + +from tests.alembic.registry import register_after_upgrade, register_before_upgrade +from tests.alembic.verifier import MigrationVerifier + +APP_ID = generate_nanoid() +USER_ID = generate_nanoid() +SESSION_ID = generate_nanoid() +MESSAGE_ID = generate_nanoid() +METAMESSAGE_ID = generate_nanoid() +COLLECTION_ID = generate_nanoid() +DOCUMENT_ID = generate_nanoid() + + +@register_before_upgrade("556a16564f50") +def prepare_user_app_ids(verifier: MigrationVerifier) -> None: + schema = verifier.schema + connection = verifier.conn + + # Create app + connection.execute( + text( + f'INSERT INTO "{schema}"."apps" ("public_id", "name") ' + + "VALUES (:app_id, :name)" + ), + {"app_id": APP_ID, "name": "test-app"}, + ) + + # Create user + connection.execute( + text( + f'INSERT INTO "{schema}"."users" ("public_id", "name", "app_id") ' + + "VALUES (:user_id, :name, :app_id)" + ), + {"user_id": USER_ID, "name": "test-user", "app_id": APP_ID}, + ) + + # Create session + connection.execute( + text( + f'INSERT INTO "{schema}"."sessions" ("public_id", "user_id", "is_active") ' + + "VALUES (:session_id, :user_id, true)" + ), + {"session_id": SESSION_ID, "user_id": USER_ID}, + ) + + # Create message + connection.execute( + text( + f'INSERT INTO "{schema}"."messages" ("public_id", "session_id", "is_user", "content") ' + + "VALUES (:message_id, :session_id, true, :content)" + ), + { + "message_id": MESSAGE_ID, + "session_id": SESSION_ID, + "content": "test-content", + }, + ) + + # Create metamessage + connection.execute( + text( + f'INSERT INTO "{schema}"."metamessages" ("public_id", "user_id", "content", "metamessage_type") ' + + "VALUES (:metamessage_id, :user_id, :content, :metamessage_type)" + ), + { + "metamessage_id": METAMESSAGE_ID, + "user_id": USER_ID, + "content": "test-content", + "metamessage_type": "test-type", + }, + ) + + # Create collection + connection.execute( + text( + f'INSERT INTO "{schema}"."collections" ("public_id", "name", "user_id") ' + + "VALUES (:collection_id, :name, :user_id)" + ), + { + "collection_id": COLLECTION_ID, + "name": "test-collection", + "user_id": USER_ID, + }, + ) + + # Create document + connection.execute( + text( + f'INSERT INTO "{schema}"."documents" ("public_id", "content", "collection_id") ' + + "VALUES (:document_id, :content, :collection_id)" + ), + { + "document_id": DOCUMENT_ID, + "content": "test-document-content", + "collection_id": COLLECTION_ID, + }, + ) + + +@register_after_upgrade("556a16564f50") +def verify_app_and_user_ids(verifier: MigrationVerifier) -> None: + """Ensure app/user identifiers are present after the upgrade.""" + + verifier.assert_column_exists("sessions", "app_id", nullable=False) + verifier.assert_indexes_exist([("sessions", "ix_sessions_app_id")]) + + verifier.assert_column_exists("messages", "app_id", nullable=False) + verifier.assert_column_exists("messages", "user_id", nullable=False) + verifier.assert_indexes_exist( + [ + ("messages", "ix_messages_app_id"), + ("messages", "ix_messages_user_id"), + ] + ) + + verifier.assert_column_exists("metamessages", "app_id", nullable=False) + verifier.assert_indexes_exist([("metamessages", "ix_metamessages_app_id")]) + + verifier.assert_column_exists("collections", "app_id", nullable=False) + verifier.assert_indexes_exist([("collections", "ix_collections_app_id")]) + + verifier.assert_column_exists("documents", "app_id", nullable=False) + verifier.assert_column_exists("documents", "user_id", nullable=False) + verifier.assert_indexes_exist( + [ + ("documents", "ix_documents_app_id"), + ("documents", "ix_documents_user_id"), + ] + ) + + conn = verifier.conn + schema = verifier.schema + + session_row = conn.execute( + text( + f'SELECT "app_id" FROM "{schema}"."sessions" ' + + 'WHERE "public_id" = :session_public_id' + ), + {"session_public_id": SESSION_ID}, + ).one() + assert session_row.app_id == APP_ID + + message_row = conn.execute( + text( + f'SELECT "app_id", "user_id" FROM "{schema}"."messages" ' + + 'WHERE "public_id" = :message_public_id' + ), + {"message_public_id": MESSAGE_ID}, + ).one() + assert message_row.app_id == APP_ID + assert message_row.user_id == USER_ID + + metamessage_row = conn.execute( + text( + f'SELECT "app_id" FROM "{schema}"."metamessages" ' + + 'WHERE "public_id" = :public_id' + ), + {"public_id": METAMESSAGE_ID}, + ).one() + assert metamessage_row.app_id == APP_ID + + collection_row = conn.execute( + text( + f'SELECT "app_id" FROM "{schema}"."collections" ' + + 'WHERE "public_id" = :public_id' + ), + {"public_id": COLLECTION_ID}, + ).one() + assert collection_row.app_id == APP_ID + + document_row = conn.execute( + text( + f'SELECT "app_id", "user_id" FROM "{schema}"."documents" ' + + 'WHERE "public_id" = :public_id' + ), + {"public_id": DOCUMENT_ID}, + ).one() + assert document_row.app_id == APP_ID + assert document_row.user_id == USER_ID diff --git a/tests/alembic/revisions/test_564ba40505c5_add_session_name_column_to_documents.py b/tests/alembic/revisions/test_564ba40505c5_add_session_name_column_to_documents.py new file mode 100644 index 00000000..3bcc16d9 --- /dev/null +++ b/tests/alembic/revisions/test_564ba40505c5_add_session_name_column_to_documents.py @@ -0,0 +1,211 @@ +"""Hooks for revision 564ba40505c5 (documents session_name column).""" + +from __future__ import annotations + +import json +import time + +from nanoid import generate as generate_nanoid +from sqlalchemy import text + +from tests.alembic.registry import register_after_upgrade, register_before_upgrade +from tests.alembic.verifier import MigrationVerifier + +WORKSPACE_ID = generate_nanoid() +WORKSPACE_NAME = generate_nanoid() +PEER_ID = generate_nanoid() +PEER_NAME = generate_nanoid() +SESSION_ID = generate_nanoid() +SESSION_NAME = generate_nanoid() +COLLECTION_ID = generate_nanoid() +COLLECTION_NAME = generate_nanoid() +DOCUMENT_ID = generate_nanoid() + + +@register_before_upgrade("564ba40505c5") +def seed_document_session_metadata(verifier: MigrationVerifier) -> None: + schema = verifier.schema + connection = verifier.conn + verifier.assert_column_exists("documents", "session_name", exists=False) + verifier.assert_indexes_not_exist([("documents", "idx_documents_session_name")]) + verifier.assert_constraint_exists( + "documents", "fk_documents_session_workspace", "foreign_key", exists=False + ) + + # Seed workspaces, peers, and collections + connection.execute( + text( + f""" + INSERT INTO "{schema}"."workspaces" ("id", "name") + VALUES (:workspace_id, :workspace_name) + """ + ), + {"workspace_id": WORKSPACE_ID, "workspace_name": WORKSPACE_NAME}, + ) + + connection.execute( + text( + f""" + INSERT INTO "{schema}"."peers" ("id", "name", "workspace_name") + VALUES (:peer_id, :peer_name, :workspace_name) + """ + ), + { + "peer_id": PEER_ID, + "peer_name": PEER_NAME, + "workspace_name": WORKSPACE_NAME, + }, + ) + + connection.execute( + text( + f""" + INSERT INTO "{schema}"."sessions" + ("id", "name", "workspace_name", "is_active") + VALUES (:session_id, :session_name, :workspace_name, true) + """ + ), + { + "session_id": SESSION_ID, + "session_name": SESSION_NAME, + "workspace_name": WORKSPACE_NAME, + }, + ) + + # Create session_peers entry so downgrade can find user_id + connection.execute( + text( + f""" + INSERT INTO "{schema}"."session_peers" + ("workspace_name", "session_name", "peer_name") + VALUES (:workspace_name, :session_name, :peer_name) + """ + ), + { + "workspace_name": WORKSPACE_NAME, + "session_name": SESSION_NAME, + "peer_name": PEER_NAME, + }, + ) + + connection.execute( + text( + f""" + INSERT INTO "{schema}"."collections" + ("id", "name", "peer_name", "workspace_name") + VALUES (:collection_id, :collection_name, :peer_name, :workspace_name) + """ + ), + { + "collection_id": COLLECTION_ID, + "collection_name": COLLECTION_NAME, + "peer_name": PEER_NAME, + "workspace_name": WORKSPACE_NAME, + }, + ) + + # Bulk insert 500k documents with session_name in internal_metadata to test batched backfill + print("[564ba40505c5] Starting bulk insert of 500k documents...", flush=True) + t0 = time.perf_counter() + connection.execute(text("SET LOCAL synchronous_commit = OFF")) + connection.execute( + text( + f'INSERT INTO "{schema}"."documents" ' + + '("id", "collection_name", "peer_name", "workspace_name", "content", "internal_metadata") ' + + "SELECT " + + " 'bulk-' || substr(md5(random()::text), 1, 16), " + + f" '{COLLECTION_NAME}', " + + f" '{PEER_NAME}', " + + f" '{WORKSPACE_NAME}', " + + " 'seed content', " + + f" jsonb_build_object('session_name', '{SESSION_NAME}') " + + "FROM generate_series(1, :n)" + ), + {"n": 500_000}, + ) + t1 = time.perf_counter() + print(f"[564ba40505c5] Bulk insert completed in {t1 - t0:.2f}s", flush=True) + + # Add document with session_name in internal_metadata + internal_metadata = json.dumps( + { + "session_name": SESSION_NAME, + } + ) + + connection.execute( + text( + f""" + INSERT INTO "{schema}"."documents" + ( + "id", + "collection_name", + "peer_name", + "workspace_name", + "content", + "internal_metadata" + ) + VALUES ( + :document_id, + :collection_name, + :peer_name, + :workspace_name, + :content, + :internal_metadata + ) + """ + ), + { + "document_id": DOCUMENT_ID, + "collection_name": COLLECTION_NAME, + "peer_name": PEER_NAME, + "workspace_name": WORKSPACE_NAME, + "content": "seed content", + "internal_metadata": internal_metadata, + }, + ) + + +@register_after_upgrade("564ba40505c5") +def verify_document_session_column(verifier: MigrationVerifier) -> None: + """Ensure the session_name column reflects migrated metadata.""" + + verifier.assert_column_exists("documents", "session_name", nullable=True) + verifier.assert_indexes_exist([("documents", "idx_documents_session_name")]) + verifier.assert_constraint_exists( + "documents", "fk_documents_session_workspace", "foreign_key" + ) + + # Quick diagnostics: counts before assertions + total_docs = ( + verifier.conn.execute( + text(f'SELECT COUNT(*) FROM "{verifier.schema}"."documents"') + ).scalar() + or 0 + ) + null_sessions = ( + verifier.conn.execute( + text( + f'SELECT COUNT(*) FROM "{verifier.schema}"."documents" ' + + 'WHERE "session_name" IS NULL' + ) + ).scalar() + or 0 + ) + print( + f"[564ba40505c5] Documents total={total_docs}, session_name NULLs={null_sessions}", + flush=True, + ) + + # All documents that contained session_name in internal_metadata should now have session_name populated + verifier.assert_no_nulls("documents", "session_name") + + # Sanity-check seeded document preserved expected session_name + row = verifier.conn.execute( + text( + f'SELECT "session_name" FROM "{verifier.schema}"."documents" ' + + 'WHERE "id" = :document_id' + ), + {"document_id": DOCUMENT_ID}, + ).one() + assert row.session_name == SESSION_NAME diff --git a/tests/alembic/revisions/test_66e63cf2cf77_add_indexes_to_documents_table.py b/tests/alembic/revisions/test_66e63cf2cf77_add_indexes_to_documents_table.py new file mode 100644 index 00000000..524c5368 --- /dev/null +++ b/tests/alembic/revisions/test_66e63cf2cf77_add_indexes_to_documents_table.py @@ -0,0 +1,16 @@ +"""Hooks for revision 66e63cf2cf77 (documents HNSW index).""" + +from __future__ import annotations + +from tests.alembic.registry import register_after_upgrade, register_before_upgrade +from tests.alembic.verifier import MigrationVerifier + + +@register_before_upgrade("66e63cf2cf77") +def prepare_documents_hnsw(verifier: MigrationVerifier) -> None: + verifier.assert_indexes_not_exist([("documents", "idx_documents_embedding_hnsw")]) + + +@register_after_upgrade("66e63cf2cf77") +def verify_documents_hnsw(verifier: MigrationVerifier) -> None: + verifier.assert_indexes_exist([("documents", "idx_documents_embedding_hnsw")]) diff --git a/tests/alembic/revisions/test_76ffba56fe8c_add_error_field_to_queueitem.py b/tests/alembic/revisions/test_76ffba56fe8c_add_error_field_to_queueitem.py new file mode 100644 index 00000000..a88fc2b6 --- /dev/null +++ b/tests/alembic/revisions/test_76ffba56fe8c_add_error_field_to_queueitem.py @@ -0,0 +1,44 @@ +"""Hooks for revision 76ffba56fe8c (queue created_at/error columns).""" + +from __future__ import annotations + +from sqlalchemy import text + +from tests.alembic.registry import register_after_upgrade, register_before_upgrade +from tests.alembic.verifier import MigrationVerifier + + +@register_before_upgrade("76ffba56fe8c") +def prepare_queue_created_at(verifier: MigrationVerifier) -> None: + verifier.assert_column_exists("queue", "created_at", nullable=False, exists=False) + verifier.assert_column_exists("queue", "error", exists=False) + verifier.assert_indexes_not_exist([("queue", "ix_queue_created_at")]) + + conn = verifier.conn + schema = verifier.schema + + # Bulk insert 500k additional queue items efficiently to exercise batched backfill + conn.execute(text("SET LOCAL synchronous_commit = OFF")) + conn.execute( + text( + f'INSERT INTO "{schema}"."queue" ' + + '("session_id", "work_unit_key", "task_type", "payload", "processed") ' + + "SELECT NULL, " + + " 'seed-batch-queue-item-' || gs::text, " + + " 'representation', " + + " jsonb_build_object('marker','bulk'), " + + " false " + + "FROM generate_series(1, :n) AS gs" + ), + {"n": 500_000}, + ) + + +@register_after_upgrade("76ffba56fe8c") +def verify_queue_created_at(verifier: MigrationVerifier) -> None: + verifier.assert_column_exists("queue", "created_at", nullable=False) + verifier.assert_column_exists("queue", "error") + verifier.assert_indexes_exist([("queue", "ix_queue_created_at")]) + + # Ensure all rows have non-null created_at after backfill + verifier.assert_no_nulls("queue", "created_at") diff --git a/tests/alembic/revisions/test_88b0fb10906f_add_webhooks_table.py b/tests/alembic/revisions/test_88b0fb10906f_add_webhooks_table.py new file mode 100644 index 00000000..1b5a0148 --- /dev/null +++ b/tests/alembic/revisions/test_88b0fb10906f_add_webhooks_table.py @@ -0,0 +1,115 @@ +"""Hooks for revision 88b0fb10906f (webhooks and queue updates).""" + +from __future__ import annotations + +import json + +from nanoid import generate as generate_nanoid +from sqlalchemy import text + +from tests.alembic.registry import register_after_upgrade, register_before_upgrade +from tests.alembic.verifier import MigrationVerifier + +PAYLOAD_MARKER = generate_nanoid() +WORKSPACE_NAME = generate_nanoid() +SESSION_NAME = generate_nanoid() +SENDER_NAME = generate_nanoid() +TARGET_NAME = generate_nanoid() + + +@register_before_upgrade("88b0fb10906f") +def prepare_webhooks(verifier: MigrationVerifier) -> None: + # Validate objects introduced by the migration are absent + verifier.assert_table_exists("webhook_endpoints", exists=False) + verifier.assert_column_exists("queue", "task_type", exists=False) + verifier.assert_column_exists("queue", "work_unit_key", exists=False) + + verifier.assert_column_exists( + "active_queue_sessions", "work_unit_key", exists=False + ) + verifier.assert_constraint_exists( + "active_queue_sessions", "unique_work_unit_key", "unique", exists=False + ) + + # Validate objects removed by the migration still exist + verifier.assert_column_exists("active_queue_sessions", "session_id") + verifier.assert_column_exists("active_queue_sessions", "sender_name") + verifier.assert_column_exists("active_queue_sessions", "target_name") + verifier.assert_column_exists("active_queue_sessions", "task_type") + verifier.assert_constraint_exists( + "active_queue_sessions", "unique_active_queue_session", "unique" + ) + + # Bulk insert 500k queue items to test batched backfill + verifier.conn.execute(text("SET LOCAL synchronous_commit = OFF")) + verifier.conn.execute( + text( + f'INSERT INTO "{verifier.schema}"."queue" ("payload") ' + + "SELECT jsonb_build_object(" + + " 'marker', 'bulk-' || substr(md5(random()::text), 1, 21)," + + " 'task_type', 'representation'," + + f" 'workspace_name', '{WORKSPACE_NAME}'," + + f" 'session_name', '{SESSION_NAME}'," + + f" 'sender_name', '{SENDER_NAME}'," + + f" 'target_name', '{TARGET_NAME}'" + + ") FROM generate_series(1, :n)" + ), + {"n": 500_000}, + ) + + payload = json.dumps( + { + "marker": PAYLOAD_MARKER, + "task_type": "representation", + "workspace_name": WORKSPACE_NAME, + "session_name": SESSION_NAME, + "sender_name": SENDER_NAME, + "target_name": TARGET_NAME, + } + ) + + # Add single queue item to test expected data transformation + verifier.conn.execute( + text(f'INSERT INTO "{verifier.schema}"."queue" ("payload") VALUES (:payload)'), + {"payload": payload}, + ) + + +@register_after_upgrade("88b0fb10906f") +def verify_webhooks_and_queue(verifier: MigrationVerifier) -> None: + """Validate queue enrichment and webhook table creation.""" + + verifier.assert_table_exists("webhook_endpoints") + + verifier.assert_column_exists("queue", "task_type", nullable=False) + verifier.assert_column_exists("queue", "work_unit_key", nullable=False) + verifier.assert_column_exists("active_queue_sessions", "work_unit_key") + verifier.assert_constraint_exists( + "active_queue_sessions", "unique_work_unit_key", "unique" + ) + + verifier.assert_column_exists("active_queue_sessions", "session_id", exists=False) + verifier.assert_column_exists("active_queue_sessions", "sender_name", exists=False) + verifier.assert_column_exists("active_queue_sessions", "target_name", exists=False) + verifier.assert_column_exists("active_queue_sessions", "task_type", exists=False) + verifier.assert_constraint_exists( + "active_queue_sessions", "unique_active_queue_session", "unique", exists=False + ) + + # Ensure no NULLs remain after backfill across all rows + verifier.assert_no_nulls("queue", "task_type") + verifier.assert_no_nulls("queue", "work_unit_key") + + row = verifier.conn.execute( + text( + f'SELECT "task_type", "work_unit_key" FROM "{verifier.schema}"."queue" ' + + "WHERE payload->>'marker' = :marker" + ), + {"marker": PAYLOAD_MARKER}, + ).one() + + assert row.task_type == "representation" + expected_key = ( + f"representation:{WORKSPACE_NAME}:{SESSION_NAME}:{SENDER_NAME}:{TARGET_NAME}" + ) + assert row.work_unit_key == expected_key diff --git a/tests/alembic/revisions/test_917195d9b5e9_add_messageembedding_table.py b/tests/alembic/revisions/test_917195d9b5e9_add_messageembedding_table.py new file mode 100644 index 00000000..f9c82848 --- /dev/null +++ b/tests/alembic/revisions/test_917195d9b5e9_add_messageembedding_table.py @@ -0,0 +1,26 @@ +"""Hooks for revision 917195d9b5e9 (message embeddings table).""" + +from __future__ import annotations + +from tests.alembic.registry import register_after_upgrade, register_before_upgrade +from tests.alembic.verifier import MigrationVerifier + +INDEXES = ( + ("message_embeddings", "idx_message_embeddings_message_id"), + ("message_embeddings", "idx_message_embeddings_workspace_name"), + ("message_embeddings", "idx_message_embeddings_session_name"), + ("message_embeddings", "idx_message_embeddings_peer_name"), + ("message_embeddings", "idx_message_embeddings_created_at"), + ("message_embeddings", "idx_message_embeddings_embedding_hnsw"), +) + + +@register_before_upgrade("917195d9b5e9") +def prepare_message_embeddings(verifier: MigrationVerifier) -> None: + verifier.assert_table_exists("message_embeddings", exists=False) + + +@register_after_upgrade("917195d9b5e9") +def verify_message_embeddings_table(verifier: MigrationVerifier) -> None: + verifier.assert_table_exists("message_embeddings") + verifier.assert_indexes_exist(INDEXES) diff --git a/tests/alembic/revisions/test_a1b2c3d4e5f6_initial_schema.py b/tests/alembic/revisions/test_a1b2c3d4e5f6_initial_schema.py new file mode 100644 index 00000000..2dadd2fd --- /dev/null +++ b/tests/alembic/revisions/test_a1b2c3d4e5f6_initial_schema.py @@ -0,0 +1,39 @@ +"""Hooks for initial schema revision a1b2c3d4e5f6.""" + +from __future__ import annotations + +from nanoid import generate as generate_nanoid + +from tests.alembic.registry import register_after_upgrade, register_before_upgrade +from tests.alembic.verifier import MigrationVerifier + +APP_ID = generate_nanoid() +APP_NAME = "seed-app" +USER_ID = generate_nanoid() +USER_NAME = "seed-user" +SESSION_ID = generate_nanoid() +SESSION_NAME = "seed-session" +MESSAGE_PUBLIC_ID = generate_nanoid() +MESSAGE_CONTENT = "seed message" +METAMESSAGE_PUBLIC_ID = generate_nanoid() +METAMESSAGE_CONTENT = "seed metamessage" +COLLECTION_ID = generate_nanoid() +COLLECTION_NAME = "seed-collection" +DOCUMENT_ID = generate_nanoid() +DOCUMENT_CONTENT = "Seed document" + + +@register_before_upgrade("a1b2c3d4e5f6") +def prepare_initial(_verifier: MigrationVerifier) -> None: + pass + + +@register_after_upgrade("a1b2c3d4e5f6") +def verify_initial_schema(verifier: MigrationVerifier) -> None: + verifier.assert_table_exists("apps", exists=True) + verifier.assert_table_exists("users", exists=True) + verifier.assert_table_exists("sessions", exists=True) + verifier.assert_table_exists("messages", exists=True) + verifier.assert_table_exists("collections", exists=True) + verifier.assert_table_exists("documents", exists=True) + verifier.assert_table_exists("metamessages", exists=True) diff --git a/tests/alembic/revisions/test_b765d82110bd_change_metamessages_to_user_level_with_.py b/tests/alembic/revisions/test_b765d82110bd_change_metamessages_to_user_level_with_.py new file mode 100644 index 00000000..524b68b4 --- /dev/null +++ b/tests/alembic/revisions/test_b765d82110bd_change_metamessages_to_user_level_with_.py @@ -0,0 +1,114 @@ +"""Hooks for revision b765d82110bd (metamessages user-level migration).""" + +from __future__ import annotations + +from nanoid import generate as generate_nanoid +from sqlalchemy import text + +from tests.alembic.registry import register_after_upgrade, register_before_upgrade +from tests.alembic.verifier import MigrationVerifier + +METAMESSAGE_ID = generate_nanoid() +USER_ID = generate_nanoid() +SESSION_ID = generate_nanoid() +MESSAGE_ID = generate_nanoid() +_INDEXES = ( + ("metamessages", "idx_metamessages_user_lookup"), + ("metamessages", "idx_metamessages_session_lookup"), + ("metamessages", "idx_metamessages_message_lookup"), +) + + +@register_before_upgrade("b765d82110bd") +def prepare_metamessages(verifier: MigrationVerifier) -> None: + schema = verifier.schema + connection = verifier.conn + # Create app + APP_ID = generate_nanoid() + connection.execute( + text( + f'INSERT INTO "{schema}"."apps" ("public_id", "name") ' + + "VALUES (:app_id, :name)" + ), + {"app_id": APP_ID, "name": "test-app"}, + ) + + # Create user + connection.execute( + text( + f'INSERT INTO "{schema}"."users" ("public_id", "name", "app_id") ' + + "VALUES (:user_id, :name, :app_id)" + ), + {"user_id": USER_ID, "name": "test-user", "app_id": APP_ID}, + ) + + # Create session + connection.execute( + text( + f'INSERT INTO "{schema}"."sessions" ("public_id", "user_id", "is_active") ' + + "VALUES (:session_id, :user_id, true)" + ), + {"session_id": SESSION_ID, "user_id": USER_ID}, + ) + + # Create message + connection.execute( + text( + f'INSERT INTO "{schema}"."messages" ("public_id", "session_id", "is_user", "content") ' + + "VALUES (:message_id, :session_id, true, :content)" + ), + { + "message_id": MESSAGE_ID, + "session_id": SESSION_ID, + "content": "test-content", + }, + ) + + # Create metamessage + connection.execute( + text( + f'INSERT INTO "{schema}"."metamessages" ("public_id", "message_id", "content", "metamessage_type") ' + + "VALUES (:metamessage_id, :message_id, :content, :metamessage_type)" + ), + { + "metamessage_id": METAMESSAGE_ID, + "message_id": MESSAGE_ID, + "content": "test-content", + "metamessage_type": "test-type", + }, + ) + + +@register_after_upgrade("b765d82110bd") +def verify_metamessages_promoted(verifier: MigrationVerifier) -> None: + """Confirm metamessages rows migrated to user-level shape.""" + + verifier.assert_column_exists("metamessages", "user_id", nullable=False) + verifier.assert_column_exists("metamessages", "session_id") + verifier.assert_column_exists("metamessages", "message_id", nullable=True) + verifier.assert_indexes_exist(_INDEXES) + verifier.assert_constraint_exists( + "metamessages", "fk_metamessages_user_id_users", "foreign_key" + ) + verifier.assert_constraint_exists( + "metamessages", "fk_metamessages_session_id_sessions", "foreign_key" + ) + verifier.assert_constraint_exists( + "metamessages", "message_requires_session", "check" + ) + verifier.assert_no_nulls("metamessages", "user_id") + + row = verifier.conn.execute( + text( + 'SELECT "user_id", "session_id", "message_id", "metamessage_type", "content" ' + + f'FROM "{verifier.schema}"."metamessages" ' + + 'WHERE "public_id" = :public_id' + ), + {"public_id": METAMESSAGE_ID}, + ).one() + + assert row.user_id == USER_ID + assert row.session_id == SESSION_ID + assert row.message_id == MESSAGE_ID + assert row.metamessage_type == "test-type" + assert row.content == "test-content" diff --git a/tests/alembic/revisions/test_bb6fb3a7a643_add_message_seq_in_session_column.py b/tests/alembic/revisions/test_bb6fb3a7a643_add_message_seq_in_session_column.py new file mode 100644 index 00000000..dd83cc03 --- /dev/null +++ b/tests/alembic/revisions/test_bb6fb3a7a643_add_message_seq_in_session_column.py @@ -0,0 +1,16 @@ +"""Hooks for revision bb6fb3a7a643 (add_message_seq_in_session_column).""" + +from __future__ import annotations + +from tests.alembic.registry import register_after_upgrade, register_before_upgrade +from tests.alembic.verifier import MigrationVerifier + + +@register_before_upgrade("bb6fb3a7a643") +def prepare_add_message_seq_in_session_column(_verifier: MigrationVerifier) -> None: + """Seed state and assertions before upgrading to bb6fb3a7a643.""" + + +@register_after_upgrade("bb6fb3a7a643") +def verify_add_message_seq_in_session_column(_verifier: MigrationVerifier) -> None: + """Add assertions validating the effects of bb6fb3a7a643.""" diff --git a/tests/alembic/revisions/test_c3828084f472_add_indexes_for_messages_and_.py b/tests/alembic/revisions/test_c3828084f472_add_indexes_for_messages_and_.py new file mode 100644 index 00000000..92bdb6f5 --- /dev/null +++ b/tests/alembic/revisions/test_c3828084f472_add_indexes_for_messages_and_.py @@ -0,0 +1,23 @@ +"""Hooks for revision c3828084f472 (read indexes).""" + +from __future__ import annotations + +from tests.alembic.registry import register_after_upgrade, register_before_upgrade +from tests.alembic.verifier import MigrationVerifier + +_INDEXES = ( + ("users", "idx_users_app_lookup"), + ("sessions", "idx_sessions_user_lookup"), + ("messages", "idx_messages_session_lookup"), + ("metamessages", "idx_metamessages_lookup"), +) + + +@register_before_upgrade("c3828084f472") +def prepare_read_indexes(verifier: MigrationVerifier) -> None: + verifier.assert_indexes_not_exist(_INDEXES) + + +@register_after_upgrade("c3828084f472") +def verify_read_indexes(verifier: MigrationVerifier) -> None: + verifier.assert_indexes_exist(_INDEXES) diff --git a/tests/alembic/revisions/test_d429de0e5338_adopt_peer_paradigm.py b/tests/alembic/revisions/test_d429de0e5338_adopt_peer_paradigm.py new file mode 100644 index 00000000..97c72903 --- /dev/null +++ b/tests/alembic/revisions/test_d429de0e5338_adopt_peer_paradigm.py @@ -0,0 +1,286 @@ +"""Hooks for revision d429de0e5338 (adopt peer paradigm).""" + +from __future__ import annotations + +import json + +from nanoid import generate as generate_nanoid +from sqlalchemy import text + +from tests.alembic.registry import register_after_upgrade, register_before_upgrade +from tests.alembic.verifier import MigrationVerifier + +APP_ID = generate_nanoid() +APP_NAME = "app-name" +USER_ID = generate_nanoid() +USER_NAME = "user-name" +SESSION_ID = generate_nanoid() +MESSAGE_ID = generate_nanoid() +MESSAGE_CONTENT = "message-content" +COLLECTION_ID = generate_nanoid() +COLLECTION_NAME = "collection-name" +DOCUMENT_ID = generate_nanoid() +QUEUE_MARKER = "seed" + + +@register_before_upgrade("d429de0e5338") +def prepare_peer_paradigm(verifier: MigrationVerifier) -> None: + """Seed legacy tables so the migration can transform real data.""" + + schema = verifier.schema + conn = verifier.conn + + conn.execute( + text( + f'INSERT INTO "{schema}"."apps" ("public_id", "name") ' + + "VALUES (:app_id, :app_name)" + ), + {"app_id": APP_ID, "app_name": APP_NAME}, + ) + + conn.execute( + text( + f'INSERT INTO "{schema}"."users" ("public_id", "name", "app_id") ' + + "VALUES (:user_id, :user_name, :app_id)" + ), + {"user_id": USER_ID, "user_name": USER_NAME, "app_id": APP_ID}, + ) + + session_result = conn.execute( + text( + f'INSERT INTO "{schema}"."sessions" ' + + '("public_id", "user_id", "app_id", "is_active") ' + + "VALUES (:session_id, :user_id, :app_id, true) RETURNING id" + ), + {"session_id": SESSION_ID, "user_id": USER_ID, "app_id": APP_ID}, + ).one() + SESSION_DB_ID = session_result.id + + conn.execute( + text( + f'INSERT INTO "{schema}"."messages" ' + + '("public_id", "session_id", "is_user", "content", "user_id", "app_id") ' + + "VALUES (:message_id, :session_id, true, :content, :user_id, :app_id)" + ), + { + "message_id": MESSAGE_ID, + "session_id": SESSION_ID, + "content": MESSAGE_CONTENT, + "user_id": USER_ID, + "app_id": APP_ID, + }, + ) + + conn.execute( + text( + f'INSERT INTO "{schema}"."collections" ' + + '("public_id", "name", "user_id", "app_id") ' + + "VALUES (:collection_id, :collection_name, :user_id, :app_id)" + ), + { + "collection_id": COLLECTION_ID, + "collection_name": COLLECTION_NAME, + "user_id": USER_ID, + "app_id": APP_ID, + }, + ) + + conn.execute( + text( + f'INSERT INTO "{schema}"."documents" ' + + '("public_id", "collection_id", "user_id", "app_id", "content") ' + + "VALUES (:document_id, :collection_id, :user_id, :app_id, :content)" + ), + { + "document_id": DOCUMENT_ID, + "collection_id": COLLECTION_ID, + "user_id": USER_ID, + "app_id": APP_ID, + "content": "document-content", + }, + ) + + queue_payload = json.dumps({"marker": QUEUE_MARKER}) + conn.execute( + text( + f'INSERT INTO "{schema}"."queue" ' + + '("session_id", "payload", "processed") ' + + "VALUES (:session_db_id, :payload, false)" + ), + { + "session_db_id": SESSION_DB_ID, + "payload": queue_payload, + }, + ) + + +@register_after_upgrade("d429de0e5338") +def verify_peer_paradigm(verifier: MigrationVerifier) -> None: + """Validate the large-scale peer paradigm migration.""" + + inspector = verifier.get_inspector() + + assert inspector.has_table("workspaces", schema=verifier.schema) + assert not inspector.has_table("apps", schema=verifier.schema) + + verifier.assert_column_exists("workspaces", "configuration", nullable=False) + verifier.assert_column_exists("workspaces", "internal_metadata", nullable=False) + verifier.assert_column_exists("peers", "workspace_name", nullable=False) + verifier.assert_column_exists("peers", "configuration", nullable=False) + verifier.assert_column_exists("sessions", "name", nullable=False) + verifier.assert_column_exists("sessions", "workspace_name", nullable=False) + verifier.assert_column_exists( + "messages", "session_name", nullable=True + ) # later changed back to required + verifier.assert_column_exists("messages", "workspace_name", nullable=False) + verifier.assert_column_exists("messages", "peer_name", nullable=False) + verifier.assert_column_exists("collections", "peer_name", nullable=False) + verifier.assert_column_exists("collections", "workspace_name", nullable=False) + verifier.assert_column_exists("collections", "internal_metadata", nullable=False) + verifier.assert_column_exists("documents", "peer_name", nullable=False) + verifier.assert_column_exists("documents", "workspace_name", nullable=False) + verifier.assert_column_exists("documents", "collection_name", nullable=False) + verifier.assert_column_exists("documents", "internal_metadata", nullable=False) + verifier.assert_column_exists("active_queue_sessions", "id", nullable=False) + verifier.assert_column_exists("active_queue_sessions", "session_id", nullable=True) + verifier.assert_column_exists("active_queue_sessions", "sender_name", nullable=True) + verifier.assert_column_exists("active_queue_sessions", "target_name", nullable=True) + verifier.assert_column_exists("active_queue_sessions", "task_type", nullable=False) + verifier.assert_column_exists("messages", "session_id", exists=False) + verifier.assert_column_exists("messages", "user_id", exists=False) + verifier.assert_column_exists("messages", "app_id", exists=False) + verifier.assert_column_exists("messages", "is_user", exists=False) + verifier.assert_indexes_exist( + [ + ("messages", "idx_messages_session_lookup"), + ("messages", "ix_messages_peer_name"), + ("messages", "ix_messages_workspace_name"), + ] + ) + + conn = verifier.conn + schema = verifier.schema + + workspace = conn.execute( + text( + 'SELECT "id", "name", "configuration", "internal_metadata" ' + + f'FROM "{schema}"."workspaces" WHERE "id" = :workspace_id' + ), + {"workspace_id": APP_ID}, + ).one() + assert workspace.name == APP_NAME + assert workspace.configuration == {} + assert workspace.internal_metadata == {} + + peer = conn.execute( + text( + 'SELECT "id", "name", "workspace_name", "configuration", "internal_metadata" ' + + f'FROM "{schema}"."peers" WHERE "id" = :peer_id' + ), + {"peer_id": USER_ID}, + ).one() + assert peer.name == USER_NAME + assert peer.workspace_name == APP_NAME + assert peer.configuration == {} + assert peer.internal_metadata == {} + + session = conn.execute( + text( + 'SELECT "id", "name", "workspace_name", "configuration", "internal_metadata" ' + + f'FROM "{schema}"."sessions" WHERE "id" = :session_id' + ), + {"session_id": SESSION_ID}, + ).one() + assert session.name == SESSION_ID + assert session.workspace_name == APP_NAME + assert session.configuration == {} + assert session.internal_metadata == {} + + session_peer = conn.execute( + text( + f'SELECT "peer_name" FROM "{schema}"."session_peers" ' + + 'WHERE "workspace_name" = :workspace_name ' + + 'AND "session_name" = :session_name AND "peer_name" = :peer_name' + ), + { + "workspace_name": APP_NAME, + "session_name": SESSION_ID, + "peer_name": USER_NAME, + }, + ).one() + assert session_peer.peer_name == USER_NAME + + message = conn.execute( + text( + 'SELECT "session_name", "workspace_name", "peer_name", "token_count", "internal_metadata" ' + + f'FROM "{schema}"."messages" WHERE "public_id" = :message_id' + ), + {"message_id": MESSAGE_ID}, + ).one() + assert message.session_name == SESSION_ID + assert message.workspace_name == APP_NAME + assert message.peer_name == USER_NAME + assert message.token_count >= 0 + assert message.internal_metadata == {} + + collection = conn.execute( + text( + 'SELECT "id", "name", "peer_name", "workspace_name", "internal_metadata" ' + + f'FROM "{schema}"."collections" WHERE "id" = :collection_id' + ), + {"collection_id": COLLECTION_ID}, + ).one() + assert collection.name == COLLECTION_NAME + assert collection.peer_name == USER_NAME + assert collection.workspace_name == APP_NAME + assert collection.internal_metadata == {} + + document = conn.execute( + text( + 'SELECT "peer_name", "workspace_name", "collection_name", "internal_metadata" ' + + f'FROM "{schema}"."documents" WHERE "id" = :document_id' + ), + {"document_id": DOCUMENT_ID}, + ).one() + assert document.peer_name == USER_NAME + assert document.workspace_name == APP_NAME + assert document.collection_name == COLLECTION_NAME + assert document.internal_metadata == {} + + queue_row = conn.execute( + text( + f'SELECT "session_id", "payload" FROM "{schema}"."queue" ' + + "WHERE payload->>'marker' = :marker" + ), + {"marker": QUEUE_MARKER}, + ).one() + assert queue_row.session_id == SESSION_ID + + verifier.assert_constraint_exists( + "messages", "fk_messages_session_name_sessions", "foreign_key" + ) + verifier.assert_constraint_exists( + "messages", "fk_messages_peer_name_peers", "foreign_key" + ) + verifier.assert_constraint_exists( + "messages", "fk_messages_workspace_name_workspaces", "foreign_key" + ) + verifier.assert_constraint_exists( + "collections", "fk_collections_peer_name_peers", "foreign_key" + ) + verifier.assert_constraint_exists( + "collections", "fk_collections_workspace_name_workspaces", "foreign_key" + ) + verifier.assert_constraint_exists( + "peers", "fk_peers_workspace_name_workspaces", "foreign_key" + ) + verifier.assert_constraint_exists("peers", "unique_name_workspace_peer", "unique") + verifier.assert_constraint_exists("sessions", "unique_session_name", "unique") + verifier.assert_constraint_exists( + "collections", "unique_name_collection_peer", "unique" + ) + verifier.assert_constraint_exists( + "active_queue_sessions", "unique_active_queue_session", "unique" + ) + + verifier.assert_table_exists("metamessages", exists=False) diff --git a/tests/alembic/scaffold.py b/tests/alembic/scaffold.py new file mode 100644 index 00000000..80cb2d59 --- /dev/null +++ b/tests/alembic/scaffold.py @@ -0,0 +1,148 @@ +"""Utility to scaffold migration hook test modules.""" + +from __future__ import annotations + +import argparse +import re +from dataclasses import dataclass +from pathlib import Path +from textwrap import dedent + +PROJECT_ROOT = Path(__file__).resolve().parents[2] +MIGRATIONS_DIR = PROJECT_ROOT / "migrations" / "versions" +TESTS_REVISION_DIR = Path(__file__).resolve().parent / "revisions" +REVISION_INIT_PATH = TESTS_REVISION_DIR / "__init__.py" + +TEMPLATE = '''"""Hooks for revision {revision}{slug_note}.""" + +from __future__ import annotations + +from tests.alembic.registry import register_after_upgrade, register_before_upgrade +from tests.alembic.verifier import MigrationVerifier + + +@register_before_upgrade("{revision}") +def prepare_{identifier}(_verifier: MigrationVerifier) -> None: + """Seed state and assertions before upgrading to {revision}.""" + + +@register_after_upgrade("{revision}") +def verify_{identifier}(_verifier: MigrationVerifier) -> None: + """Add assertions validating the effects of {revision}.""" +''' + +HEADER = '"""Register revision-specific hooks for migration verification."""' + + +@dataclass(slots=True) +class MigrationInfo: + revision: str + slug: str + path: Path + + @property + def identifier(self) -> str: + """Return a Python-safe identifier derived from the migration slug.""" + + base = re.sub(r"[^0-9a-zA-Z]+", "_", self.slug) + base = base.strip("_").lower() or "revision" + if base[0].isdigit(): + base = f"revision_{base}" + return base + + @property + def test_filename(self) -> str: + return f"test_{self.path.stem}.py" + + +def parse_args() -> argparse.Namespace: + parser = argparse.ArgumentParser( + description=("Generate a revision test stub with before/after upgrade hooks.") + ) + parser.add_argument( + "revision", + help="Revision id (e.g. a1b2c3d4e5f6) or path to a migration file.", + ) + return parser.parse_args() + + +def resolve_migration(path_or_revision: str) -> MigrationInfo: + candidate = Path(path_or_revision) + if candidate.suffix == ".py" and candidate.exists(): + return MigrationInfo( + revision=candidate.stem.split("_", 1)[0], + slug=candidate.stem.split("_", 1)[1] if "_" in candidate.stem else "", + path=candidate.resolve(), + ) + + matches = sorted(MIGRATIONS_DIR.glob(f"{path_or_revision}_*.py")) + if not matches: + raise FileNotFoundError( + f"Could not find migration matching {path_or_revision!r} in {MIGRATIONS_DIR}." + ) + if len(matches) > 1: + options = ", ".join(match.name for match in matches) + raise ValueError( + f"Revision prefix {path_or_revision!r} matches multiple migrations: {options}." + + " Provide a more specific revision or path." + ) + match = matches[0] + stem = match.stem + if "_" in stem: + revision, slug = stem.split("_", 1) + else: + revision, slug = stem, "" + return MigrationInfo(revision=revision, slug=slug, path=match.resolve()) + + +def build_template(info: MigrationInfo) -> str: + slug_note = f" ({info.slug})" if info.slug else "" + content = TEMPLATE.format( + revision=info.revision, + identifier=info.identifier, + slug_note=slug_note, + ) + return dedent(content).rstrip() + "\n" + + +def write_stub(info: MigrationInfo) -> Path: + TESTS_REVISION_DIR.mkdir(parents=True, exist_ok=True) + target = TESTS_REVISION_DIR / info.test_filename + target.write_text(build_template(info), encoding="utf-8") + return target + + +def refresh_revision_init() -> None: + modules = sorted( + path.stem for path in TESTS_REVISION_DIR.glob("test_*.py") if path.is_file() + ) + import_block = "\n".join(f" {module}," for module in modules) + all_block = "\n".join(f' "{module}",' for module in modules) + new_content = ( + dedent( + f"""{HEADER} + +from . import ( +{import_block} +) + +__all__ = [ +{all_block} +] +""" + ).rstrip() + + "\n" + ) + REVISION_INIT_PATH.write_text(new_content, encoding="utf-8") + + +def main() -> None: + args = parse_args() + info = resolve_migration(args.revision) + target = write_stub(info) + refresh_revision_init() + print(f"Created {target.relative_to(PROJECT_ROOT)}") + + +if __name__ == "__main__": + main() diff --git a/tests/alembic/test_pipeline.py b/tests/alembic/test_pipeline.py new file mode 100644 index 00000000..73fba526 --- /dev/null +++ b/tests/alembic/test_pipeline.py @@ -0,0 +1,85 @@ +"""Test pipeline for running alembic migrations and corresponding test hooks in order.""" + +from __future__ import annotations + +import pytest +from alembic import command +from alembic.config import Config +from alembic.script import ScriptDirectory +from sqlalchemy import Engine + +from tests.alembic.conftest import ALEMBIC_CONFIG_PATH +from tests.alembic.registry import get_registered_hooks +from tests.alembic.verifier import MigrationVerifier + + +def _load_revision_sequence() -> tuple[str, ...]: + """Read the Alembic script directory to produce the linear revision order.""" + + script = ScriptDirectory.from_config(Config(str(ALEMBIC_CONFIG_PATH))) + revisions = list(script.walk_revisions()) # newest -> oldest + revisions.reverse() + return tuple(revision.revision for revision in revisions) + + +REVISION_SEQUENCE: tuple[str, ...] = _load_revision_sequence() +REVISION_PARAMS = [ + pytest.param(revision, id=f"{index:02d}_{revision}") + for index, revision in enumerate(REVISION_SEQUENCE, start=1) +] + + +def _test_single_revision( + revision: str, + alembic_cfg: Config, + alembic_engine: Engine, +) -> None: + """ + Test a single migration revision upgrade. + + For each revision: + - Migrate to the previous revision + - Run the before_upgrade hook to seed and validate the state of the DB before the revision + - Migrate to the current revision + - Run the after_upgrade hook to validate the state of the DB after the revision + """ + hooks_map = get_registered_hooks() + revision_order = list(REVISION_SEQUENCE) + + # Find the previous revision in the chain + revision_index = revision_order.index(revision) + previous_revision = ( + revision_order[revision_index - 1] if revision_index > 0 else "base" + ) + + # Migrate up to the previous revision using a shared connection + with alembic_engine.begin() as conn: + alembic_cfg.attributes["connection"] = conn + command.upgrade(alembic_cfg, previous_revision) + + # Run before_upgrade hook if it exists + hooks = hooks_map.get(revision) + if hooks and hooks.before_upgrade: + with alembic_engine.begin() as conn: + verifier = MigrationVerifier(conn, revision) + hooks.before_upgrade(verifier) + + # Migrate to the current revision using the same pattern + with alembic_engine.begin() as conn: + alembic_cfg.attributes["connection"] = conn + command.upgrade(alembic_cfg, revision) + + # Run after_upgrade hook if it exists + if hooks and hooks.after_upgrade: + with alembic_engine.begin() as conn: + verifier = MigrationVerifier(conn, revision) + hooks.after_upgrade(verifier) + + +@pytest.mark.parametrize("revision", REVISION_PARAMS) +def test_migration_revision( + revision: str, + alembic_cfg: Config, + alembic_engine: Engine, +) -> None: + _test_single_revision(revision, alembic_cfg, alembic_engine) diff --git a/tests/alembic/verifier.py b/tests/alembic/verifier.py new file mode 100644 index 00000000..a859f719 --- /dev/null +++ b/tests/alembic/verifier.py @@ -0,0 +1,144 @@ +"""Utilities for asserting alembic migration behaviour inside tests.""" + +from __future__ import annotations + +from collections.abc import Sequence + +from sqlalchemy import inspect, text +from sqlalchemy.engine import Connection +from sqlalchemy.engine.reflection import Inspector + +from migrations.utils import get_schema + + +class MigrationVerifier: + """Helper to run reusable assertions against the migrated schema.""" + + def __init__(self, connection: Connection, revision: str): + self.conn: Connection = connection + self.revision: str = revision + self.schema: str = get_schema() + self._inspector: Inspector | None = None + + def assert_table_exists(self, table: str, *, exists: bool = True) -> None: + """Assert that a table exists in the schema""" + tables = self.get_inspector().get_table_names(schema=self.schema) + + if exists: + assert table in tables + else: + assert table not in tables + + def assert_column_exists( + self, + table: str, + column: str, + *, + exists: bool = True, + nullable: bool | None = None, + ) -> None: + """Assert that a column exists in the schema""" + columns = self.get_inspector().get_columns(table, schema=self.schema) + col_names = [c["name"] for c in columns] + + if exists: + assert column in col_names + else: + assert column not in col_names + + if nullable is not None: + if column not in col_names: + # Column absence was asserted above; nothing further to verify + return + + column_info = next(col for col in columns if col["name"] == column) + actual_nullable = column_info.get("nullable", True) + assert ( + actual_nullable == nullable + ), f"Column {table}.{column} nullability is {actual_nullable}; expected {nullable}" + + def assert_column_type(self, table: str, column: str, expected_type: type) -> None: + """Assert that a column has the expected type""" + columns = self.get_inspector().get_columns(table, schema=self.schema) + column_info = next((col for col in columns if col["name"] == column), None) + assert ( + column_info is not None + ), f"Column {table}.{column} not found after migration {self.revision}" + actual_type = column_info["type"] + assert isinstance( + actual_type, expected_type + ), f"Column {table}.{column} has type {type(actual_type).__name__}; expected {expected_type.__name__}" + + def assert_no_nulls(self, table: str, column: str) -> None: + """Assert that a column has no null values""" + result = self.conn.execute( + text( + f'SELECT COUNT(*) FROM "{self.schema}"."{table}" ' + + f'WHERE "{column}" IS NULL' + ) + ) + count = result.scalar() or 0 + assert ( + count == 0 + ), f"Found {count} NULL values in {table}.{column} after migration {self.revision}" + + def assert_constraint_exists( + self, + table: str, + constraint_name: str, + constraint_type: str, + *, + exists: bool = True, + ) -> None: + """Assert that a constraint exists in the schema""" + names = self.fetch_constraints(table, constraint_type) + + if exists: + assert constraint_name in names + else: + assert constraint_name not in names + + def assert_indexes_exist(self, checks: Sequence[tuple[str, str]]) -> None: + """Assert that indexes exist in the schema""" + for table_name, index_name in checks: + indexes = self.get_inspector().get_indexes(table_name, schema=self.schema) + names = [idx["name"] for idx in indexes] + assert ( + index_name in names + ), f"Index {index_name} not found on {table_name} after migration {self.revision}" + + def assert_indexes_not_exist(self, checks: Sequence[tuple[str, str]]) -> None: + """Assert that indexes do not exist in the schema""" + for table_name, index_name in checks: + indexes = self.get_inspector().get_indexes(table_name, schema=self.schema) + names = [idx["name"] for idx in indexes] + assert ( + index_name not in names + ), f"Index {index_name} still present on {table_name} after migration {self.revision}" + + def get_inspector(self) -> Inspector: + """Get the inspector for the connection""" + if self._inspector is None: + self._inspector = inspect(self.conn) + return self._inspector + + def fetch_constraints(self, table: str, constraint_type: str) -> list[str | None]: + """Collect the names of constraints in a table""" + inspector = self.get_inspector() + + if constraint_type == "unique": + constraints = inspector.get_unique_constraints(table, schema=self.schema) + elif constraint_type == "foreign_key": + constraints = inspector.get_foreign_keys(table, schema=self.schema) + elif constraint_type == "check": + constraints = inspector.get_check_constraints(table, schema=self.schema) + elif constraint_type == "primary_key": + constraint = inspector.get_pk_constraint(table, schema=self.schema) + constraints = [constraint] if constraint else [] + else: + raise ValueError(f"Unknown constraint type: {constraint_type}") + + return [c.get("name") for c in constraints if c] + + +__all__ = ["MigrationVerifier"] diff --git a/tests/conftest.py b/tests/conftest.py index 9ba3036e..4b7cb9c1 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -493,6 +493,10 @@ def mock_tracked_db(db_session: AsyncSession): patch("src.routers.sessions.tracked_db", mock_tracked_db_context), patch("src.crud.representation.tracked_db", mock_tracked_db_context), patch("src.routers.peers.tracked_db", mock_tracked_db_context), + patch("src.dreamer.dreamer.tracked_db", mock_tracked_db_context), + patch("src.dreamer.dream_scheduler.tracked_db", mock_tracked_db_context), + patch("src.dialectic.chat.tracked_db", mock_tracked_db_context), + patch("src.utils.summarizer.tracked_db", mock_tracked_db_context), ): yield diff --git a/tests/sdk/test_pagination.py b/tests/sdk/test_pagination.py index 2dcaceb9..021582cd 100644 --- a/tests/sdk/test_pagination.py +++ b/tests/sdk/test_pagination.py @@ -9,276 +9,216 @@ from sdks.python.src.honcho.peer import Peer @pytest.mark.asyncio -async def test_sync_page_get_next_page( +async def test_page_get_next_page( client_fixture: tuple[Honcho | AsyncHoncho, str], ): """ - Tests that SyncPage.get_next_page() works correctly for sync clients. + Tests that Page.get_next_page() works correctly for both sync and async clients. """ honcho_client, client_type = client_fixture - if client_type != "sync": - pytest.skip("Test only applicable to sync client") + if client_type == "async": + assert isinstance(honcho_client, AsyncHoncho) - assert isinstance(honcho_client, Honcho) + # Create multiple peers to test pagination + for i in range(15): + peer = await honcho_client.peer(id=f"pagination-test-peer-async-{i}") + await peer.get_metadata() # Create the peer - # Create multiple peers to test pagination - for i in range(15): - peer = honcho_client.peer(id=f"pagination-test-peer-{i}") - peer.get_metadata() # Create the peer + # Get first page + first_page = await honcho_client.get_peers() + assert isinstance(first_page, AsyncPage) + assert len(first_page.items) > 0 - # Get first page - SDK doesn't expose size parameter yet, so we get default pagination - first_page = honcho_client.get_peers() - assert isinstance(first_page, SyncPage) - assert len(first_page.items) > 0 + # If there's a next page, test get_next_page() + if first_page.has_next_page(): + second_page = await first_page.get_next_page() + assert second_page is not None + assert isinstance(second_page, AsyncPage) + assert second_page.page == 2 + else: + assert isinstance(honcho_client, Honcho) - # If there's a next page, test get_next_page() - if first_page.has_next_page(): - second_page = first_page.get_next_page() - assert second_page is not None - assert isinstance(second_page, SyncPage) - assert second_page.page == 2 + # Create multiple peers to test pagination + for i in range(15): + peer = honcho_client.peer(id=f"pagination-test-peer-{i}") + peer.get_metadata() # Create the peer + + # Get first page + first_page = honcho_client.get_peers() + assert isinstance(first_page, SyncPage) + assert len(first_page.items) > 0 + + # If there's a next page, test get_next_page() + if first_page.has_next_page(): + second_page = first_page.get_next_page() + assert second_page is not None + assert isinstance(second_page, SyncPage) + assert second_page.page == 2 @pytest.mark.asyncio -async def test_async_page_get_next_page( +async def test_page_transform_preserved_across_pages( client_fixture: tuple[Honcho | AsyncHoncho, str], ): """ - Tests that AsyncPage.get_next_page() works correctly for async clients. + Tests that transformation function is preserved when getting next page. """ honcho_client, client_type = client_fixture - if client_type != "async": - pytest.skip("Test only applicable to async client") + if client_type == "async": + assert isinstance(honcho_client, AsyncHoncho) - assert isinstance(honcho_client, AsyncHoncho) + # Create multiple peers to ensure pagination + for i in range(25): + peer = await honcho_client.peer(id=f"transform-test-peer-async-{i}") + await peer.get_metadata() - # Create multiple peers to test pagination - for i in range(15): - peer = await honcho_client.peer(id=f"pagination-test-peer-async-{i}") - await peer.get_metadata() # Create the peer + # Get first page + first_page = await honcho_client.get_peers() + assert isinstance(first_page, AsyncPage) - # Get first page - SDK doesn't expose size parameter yet, so we get default pagination - first_page = await honcho_client.get_peers() - assert isinstance(first_page, AsyncPage) - assert len(first_page.items) > 0 - - # If there's a next page, test get_next_page() - if first_page.has_next_page(): - second_page = await first_page.get_next_page() - assert second_page is not None - assert isinstance(second_page, AsyncPage) - assert second_page.page == 2 - - -@pytest.mark.asyncio -async def test_sync_page_transform_preserved_across_pages( - client_fixture: tuple[Honcho | AsyncHoncho, str], -): - """ - Tests that transformation function is preserved when getting next page (sync). - """ - honcho_client, client_type = client_fixture - - if client_type != "sync": - pytest.skip("Test only applicable to sync client") - - assert isinstance(honcho_client, Honcho) - - # Create multiple peers to ensure pagination - for i in range(25): - peer = honcho_client.peer(id=f"transform-test-peer-{i}") - peer.get_metadata() - - # Get first page - first_page = honcho_client.get_peers() - assert isinstance(first_page, SyncPage) - - # Verify items are Peer instances (transformed) - for item in first_page.items: - assert isinstance(item, Peer) - - # Get next page and verify transformation is preserved - if first_page.has_next_page(): - second_page = first_page.get_next_page() - assert second_page is not None - for item in second_page.items: - assert isinstance(item, Peer) - - -@pytest.mark.asyncio -async def test_async_page_transform_preserved_across_pages( - client_fixture: tuple[Honcho | AsyncHoncho, str], -): - """ - Tests that transformation function is preserved when getting next page (async). - """ - honcho_client, client_type = client_fixture - - if client_type != "async": - pytest.skip("Test only applicable to async client") - - assert isinstance(honcho_client, AsyncHoncho) - - # Create multiple peers to ensure pagination - for i in range(25): - peer = await honcho_client.peer(id=f"transform-test-peer-async-{i}") - await peer.get_metadata() - - # Get first page - first_page = await honcho_client.get_peers() - assert isinstance(first_page, AsyncPage) - - # Verify items are AsyncPeer instances (transformed) - for item in first_page.items: - assert isinstance(item, AsyncPeer) - - # Get next page and verify transformation is preserved - if first_page.has_next_page(): - second_page = await first_page.get_next_page() - assert second_page is not None - for item in second_page.items: + # Verify items are AsyncPeer instances (transformed) + for item in first_page.items: assert isinstance(item, AsyncPeer) + # Get next page and verify transformation is preserved + if first_page.has_next_page(): + second_page = await first_page.get_next_page() + assert second_page is not None + for item in second_page.items: + assert isinstance(item, AsyncPeer) + else: + assert isinstance(honcho_client, Honcho) -@pytest.mark.asyncio -async def test_sync_page_get_next_page_throws_exception_on_last_page( - client_fixture: tuple[Honcho | AsyncHoncho, str], -): - """ - Tests that get_next_page() returns None when on the last page (sync). - """ - honcho_client, client_type = client_fixture + # Create multiple peers to ensure pagination + for i in range(25): + peer = honcho_client.peer(id=f"transform-test-peer-{i}") + peer.get_metadata() - if client_type != "sync": - pytest.skip("Test only applicable to sync client") + # Get first page + first_page = honcho_client.get_peers() + assert isinstance(first_page, SyncPage) - assert isinstance(honcho_client, Honcho) + # Verify items are Peer instances (transformed) + for item in first_page.items: + assert isinstance(item, Peer) - # Create just a few peers to ensure we're on the last page - for i in range(3): - peer = honcho_client.peer(id=f"last-page-test-peer-{i}") - peer.get_metadata() - - # Get first page - first_page = honcho_client.get_peers() - assert isinstance(first_page, SyncPage) - - # Should be on last page (or only page) - if not first_page.has_next_page(): - # get_next_page should return None - try: - first_page.get_next_page() - except Exception as e: - assert isinstance(e, RuntimeError) + # Get next page and verify transformation is preserved + if first_page.has_next_page(): + second_page = first_page.get_next_page() + assert second_page is not None + for item in second_page.items: + assert isinstance(item, Peer) @pytest.mark.asyncio -async def test_async_page_get_next_page_throws_exception_on_last_page( +async def test_page_get_next_page_throws_exception_on_last_page( client_fixture: tuple[Honcho | AsyncHoncho, str], ): """ - Tests that get_next_page() returns None when on the last page (async). + Tests that get_next_page() throws RuntimeError when on the last page. """ honcho_client, client_type = client_fixture - if client_type != "async": - pytest.skip("Test only applicable to async client") + if client_type == "async": + assert isinstance(honcho_client, AsyncHoncho) - assert isinstance(honcho_client, AsyncHoncho) + # Create just a few peers to ensure we're on the last page + for i in range(3): + peer = await honcho_client.peer(id=f"last-page-test-peer-async-{i}") + await peer.get_metadata() - # Create just a few peers to ensure we're on the last page - for i in range(3): - peer = await honcho_client.peer(id=f"last-page-test-peer-async-{i}") - await peer.get_metadata() + # Get first page + first_page = await honcho_client.get_peers() + assert isinstance(first_page, AsyncPage) - # Get first page - first_page = await honcho_client.get_peers() - assert isinstance(first_page, AsyncPage) + # Should be on last page (or only page) + if not first_page.has_next_page(): + # get_next_page should throw runtime error + try: + await first_page.get_next_page() + except Exception as e: + assert isinstance(e, RuntimeError) + else: + assert isinstance(honcho_client, Honcho) - # Should be on last page (or only page) - if not first_page.has_next_page(): - # get_next_page should throw runtime error - try: - await first_page.get_next_page() - except Exception as e: - assert isinstance(e, RuntimeError) + # Create just a few peers to ensure we're on the last page + for i in range(3): + peer = honcho_client.peer(id=f"last-page-test-peer-{i}") + peer.get_metadata() + + # Get first page + first_page = honcho_client.get_peers() + assert isinstance(first_page, SyncPage) + + # Should be on last page (or only page) + if not first_page.has_next_page(): + # get_next_page should throw runtime error + try: + first_page.get_next_page() + except Exception as e: + assert isinstance(e, RuntimeError) @pytest.mark.asyncio -async def test_sync_page_manual_pagination( +async def test_page_manual_pagination( client_fixture: tuple[Honcho | AsyncHoncho, str], ): """ - Tests manual pagination with get_next_page (sync). + Tests manual pagination with get_next_page. """ honcho_client, client_type = client_fixture - if client_type != "sync": - pytest.skip("Test only applicable to sync client") + if client_type == "async": + assert isinstance(honcho_client, AsyncHoncho) - assert isinstance(honcho_client, Honcho) + # Create enough peers to ensure multiple pages + for i in range(25): + peer = await honcho_client.peer(id=f"manual-pagination-test-peer-async-{i}") + await peer.get_metadata() - # Create enough peers to ensure multiple pages - for i in range(25): - peer = honcho_client.peer(id=f"manual-pagination-test-peer-{i}") - peer.get_metadata() + # Collect items via manual pagination + manual_items = [] + manual_page = await honcho_client.get_peers() + page_count = 0 - # Collect items via manual pagination - manual_items = [] - manual_page = honcho_client.get_peers() - page_count = 0 + while manual_page is not None: + page_count += 1 + manual_items.extend(manual_page.items) # pyright: ignore - while manual_page is not None: - page_count += 1 - manual_items.extend(manual_page.items) # pyright: ignore + if not manual_page.has_next_page(): + break - if not manual_page.has_next_page(): - break + manual_page = await manual_page.get_next_page() - manual_page = manual_page.get_next_page() + # Should have collected all items + assert len(manual_items) >= 25 # pyright: ignore + # All items should be AsyncPeer instances + assert all(isinstance(item, AsyncPeer) for item in manual_items) # pyright: ignore + else: + assert isinstance(honcho_client, Honcho) - # Should have collected all items - assert len(manual_items) >= 25 # pyright: ignore - # All items should be Peer instances - assert all(isinstance(item, Peer) for item in manual_items) # pyright: ignore + # Create enough peers to ensure multiple pages + for i in range(25): + peer = honcho_client.peer(id=f"manual-pagination-test-peer-{i}") + peer.get_metadata() + # Collect items via manual pagination + manual_items = [] + manual_page = honcho_client.get_peers() + page_count = 0 -@pytest.mark.asyncio -async def test_async_page_manual_pagination( - client_fixture: tuple[Honcho | AsyncHoncho, str], -): - """ - Tests manual pagination with get_next_page (async). - """ - honcho_client, client_type = client_fixture + while manual_page is not None: + page_count += 1 + manual_items.extend(manual_page.items) # pyright: ignore - if client_type != "async": - pytest.skip("Test only applicable to async client") + if not manual_page.has_next_page(): + break - assert isinstance(honcho_client, AsyncHoncho) + manual_page = manual_page.get_next_page() - # Create enough peers to ensure multiple pages - for i in range(25): - peer = await honcho_client.peer(id=f"manual-pagination-test-peer-async-{i}") - await peer.get_metadata() - - # Collect items via manual pagination - manual_items = [] - manual_page = await honcho_client.get_peers() - page_count = 0 - - while manual_page is not None: - page_count += 1 - manual_items.extend(manual_page.items) # pyright: ignore - - if not manual_page.has_next_page(): - break - - manual_page = await manual_page.get_next_page() - - # Should have collected all items - assert len(manual_items) >= 25 # pyright: ignore - # All items should be AsyncPeer instances - assert all(isinstance(item, AsyncPeer) for item in manual_items) # pyright: ignore + # Should have collected all items + assert len(manual_items) >= 25 # pyright: ignore + # All items should be Peer instances + assert all(isinstance(item, Peer) for item in manual_items) # pyright: ignore From 96abc49dbed9514652a877b893d0a0a413372a9b Mon Sep 17 00:00:00 2001 From: Vineeth Voruganti <13438633+VVoruganti@users.noreply.github.com> Date: Fri, 24 Oct 2025 10:35:57 -0400 Subject: [PATCH 10/17] chore: v2.4.1 Release Notes --- CHANGELOG.md | 16 ++++++++ README.md | 2 +- docs/changelog/compatibility-guide.mdx | 5 ++- docs/changelog/introduction.mdx | 22 +++++++++-- docs/docs.json | 54 +++++++++++++++++++------- pyproject.toml | 2 +- src/main.py | 2 +- uv.lock | 2 +- 8 files changed, 83 insertions(+), 22 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 5e194440..1fdd2e90 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,6 +5,22 @@ All notable changes to this project will be documented in this file. The format is based on [Keep a Changelog](http://keepachangelog.com/) and this project adheres to [Semantic Versioning](http://semver.org/). +## [2.4.1] - 2025-10-24 + +### Added + +- Alembic migration validation test suite + +### Fixed + +- Alembic migrations to batch changes +- Batch message creation sequence number + +### Changed + +- Logging infrastructure to remove noisy messages +- Sentry integration is centralized + ## [2.4.0] - 2025-10-09 ### Added diff --git a/README.md b/README.md index 86a8b790..452f487b 100644 --- a/README.md +++ b/README.md @@ -8,7 +8,7 @@ --- -![Static Badge](https://img.shields.io/badge/Version-2.4.0-blue) +![Static Badge](https://img.shields.io/badge/Version-2.4.1-blue) [![PyPI version](https://img.shields.io/pypi/v/honcho-ai.svg)](https://pypi.org/project/honcho-ai/) [![NPM version](https://img.shields.io/npm/v/@honcho-ai/sdk.svg)](https://npmjs.org/package/@honcho-ai/sdk) [![Discord](https://img.shields.io/discord/1016845111637839922?style=flat&logo=discord&logoColor=23ffffff&label=Plastic%20Labs&labelColor=235865F2)](https://discord.gg/plasticlabs) diff --git a/docs/changelog/compatibility-guide.mdx b/docs/changelog/compatibility-guide.mdx index f6b3c46e..edcc8fdd 100644 --- a/docs/changelog/compatibility-guide.mdx +++ b/docs/changelog/compatibility-guide.mdx @@ -8,7 +8,7 @@ This guide helps you understand which versions of Honcho's API are compatible wi ## Version Compatibility -### Honcho API v2.4.0 (Current) +### Honcho API v2.4.1 (Current) @@ -35,7 +35,8 @@ This guide helps you understand which versions of Honcho's API are compatible wi | Honcho API Version | TypeScript SDK | Python SDK | |-------------------|---------------|------------| -| v2.4.0 (Current) | v1.5.0 | v1.5.0 | +| v2.4.1 (Current) | v1.5.0 | v1.5.0 | +| v2.4.0 | v1.5.0 | v1.5.0 | | v2.3.3 | v1.4.1 | v1.4.1 | | v2.3.2 | v1.4.0 | v1.4.0 | | v2.3.1 | v1.4.0 | v1.4.0 | diff --git a/docs/changelog/introduction.mdx b/docs/changelog/introduction.mdx index 062ec647..4edab9dc 100644 --- a/docs/changelog/introduction.mdx +++ b/docs/changelog/introduction.mdx @@ -27,7 +27,23 @@ Welcome to the Honcho changelog! This section documents all notable changes to t ### Honcho API and SDK Changelogs - + + ### Added + + - Alembic migration validation test suite + + ### Fixed + + - Alembic migrations to batch changes + - Batch message creation sequence number + + ### Changed + + - Logging infrastructure to remove noisy messages + - Sentry integration is centralized + + + ### Added - Unified `Representation` class @@ -335,7 +351,7 @@ Welcome to the Honcho changelog! This section documents all notable changes to t [Python SDK](https://pypi.org/project/honcho-ai/) - + ### Added - Delete workspace method @@ -409,7 +425,7 @@ Welcome to the Honcho changelog! This section documents all notable changes to t [TypeScript SDK](https://www.npmjs.com/package/@honcho-ai/sdk) - + ### Added - Delete workspace method diff --git a/docs/docs.json b/docs/docs.json index 01dc6bea..4a9a6412 100644 --- a/docs/docs.json +++ b/docs/docs.json @@ -9,14 +9,21 @@ }, "favicon": "/favicon.svg", "contextual": { - "options": ["copy", "view", "chatgpt", "claude"] + "options": [ + "copy", + "view", + "chatgpt", + "claude" + ] }, "navigation": { "versions": [ { - "version": "v2.4.0", + "version": "v2.4.1", "api": { - "openapi": ["openapi.documented.yml"] + "openapi": [ + "openapi.documented.yml" + ] }, "tabs": [ { @@ -54,11 +61,17 @@ "groups": [ { "group": "Getting Started", - "pages": ["v2/guides/overview", "v2/guides/mcp"] + "pages": [ + "v2/guides/overview", + "v2/guides/mcp" + ] }, { "group": "Application Interfaces", - "pages": ["v2/guides/discord", "v2/guides/telegram"] + "pages": [ + "v2/guides/discord", + "v2/guides/telegram" + ] }, { "group": "Design Patterns", @@ -93,7 +106,9 @@ "groups": [ { "group": "API Documentation", - "pages": ["v2/api-reference/introduction"] + "pages": [ + "v2/api-reference/introduction" + ] }, { "group": "workspaces", @@ -148,7 +163,6 @@ "v2/api-reference/endpoint/messages/create-messages-with-file" ] }, - { "group": "webhooks", "pages": [ @@ -184,7 +198,9 @@ { "version": "v1.1.0", "api": { - "openapi": ["openapi.json"] + "openapi": [ + "openapi.json" + ] }, "tabs": [ { @@ -214,15 +230,23 @@ "groups": [ { "group": "Getting Started", - "pages": ["v1/guides/overview", "v1/guides/streaming-response"] + "pages": [ + "v1/guides/overview", + "v1/guides/streaming-response" + ] }, { "group": "Application Interfaces", - "pages": ["v1/guides/discord", "v1/guides/honcho-mcp"] + "pages": [ + "v1/guides/discord", + "v1/guides/honcho-mcp" + ] }, { "group": "Personal Memory", - "pages": ["v1/guides/dialectic-endpoint"] + "pages": [ + "v1/guides/dialectic-endpoint" + ] } ] }, @@ -231,7 +255,9 @@ "groups": [ { "group": "API Documentation", - "pages": ["v1/api-reference/introduction"] + "pages": [ + "v1/api-reference/introduction" + ] }, { "group": "apps", @@ -279,7 +305,9 @@ }, { "group": "keys", - "pages": ["v1/api-reference/endpoint/keys/create-key"] + "pages": [ + "v1/api-reference/endpoint/keys/create-key" + ] }, { "group": "metamessages", diff --git a/pyproject.toml b/pyproject.toml index 0022cd1e..6c9b037c 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "honcho" -version = "2.4.0" +version = "2.4.1" description = "Honcho Server" authors = [ {name = "Plastic Labs", email = "hello@plasticlabs.ai"}, diff --git a/src/main.py b/src/main.py index 6ae59192..412822d2 100644 --- a/src/main.py +++ b/src/main.py @@ -123,7 +123,7 @@ app = FastAPI( title="Honcho API", summary="The Identity Layer for the Agentic World", description="""Honcho is a platform for giving agents user-centric memory and social cognition""", - version="2.4.0", + version="2.4.1", contact={ "name": "Plastic Labs", "url": "https://honcho.dev", diff --git a/uv.lock b/uv.lock index 9d96e1cf..dfe736ee 100644 --- a/uv.lock +++ b/uv.lock @@ -673,7 +673,7 @@ wheels = [ [[package]] name = "honcho" -version = "2.4.0" +version = "2.4.1" source = { virtual = "." } dependencies = [ { name = "alembic" }, From 8a37a9570ba241427e58e7339d5c0782c769f1ff Mon Sep 17 00:00:00 2001 From: Vineeth Voruganti <13438633+VVoruganti@users.noreply.github.com> Date: Mon, 27 Oct 2025 16:58:10 -0400 Subject: [PATCH 11/17] fix: Improve alembic migration reliability --- migrations/env.py | 5 ++++- 1 file changed, 4 insertions(+), 1 deletion(-) diff --git a/migrations/env.py b/migrations/env.py index 8ee74328..6787437e 100644 --- a/migrations/env.py +++ b/migrations/env.py @@ -94,7 +94,10 @@ def run_migrations_online() -> None: prefix="sqlalchemy.", echo=False, poolclass=pool.NullPool, - connect_args={"prepare_threshold": None}, + connect_args={ + "prepare_threshold": None, + "options": "-c statement_timeout=300000", # 5 minutes in milliseconds + }, ) with connectable.connect() as connection: From 6df41265ed495e3c2d775589f027b50e3c09271f Mon Sep 17 00:00:00 2001 From: Vineeth Voruganti <13438633+VVoruganti@users.noreply.github.com> Date: Tue, 28 Oct 2025 17:41:49 -0400 Subject: [PATCH 12/17] fix: Ensure alembic migrations use a session pooler (#247) * fix: Ensure alembic migrations use a session pooler * fix: bump batch size; fix document delete in 08894082221a (#249) * chore: add alembic logging * fix (alembic): remove order by in batches * fix: fkey -> fk * chore: code rabbit --------- Co-authored-by: Rajat Ahuja --- migrations/env.py | 56 +++++- ..._replace_collection_name_with_observer_.py | 160 +++++++++++------- ...c5_add_session_name_column_to_documents.py | 18 +- ..._replace_collection_name_with_observer_.py | 4 +- 4 files changed, 171 insertions(+), 67 deletions(-) diff --git a/migrations/env.py b/migrations/env.py index 6787437e..f0a8197c 100644 --- a/migrations/env.py +++ b/migrations/env.py @@ -2,9 +2,10 @@ import logging import sys from logging.config import fileConfig from pathlib import Path +from urllib.parse import urlparse, urlunparse from alembic import context -from sqlalchemy import engine_from_config, pool, text +from sqlalchemy import engine_from_config, text from src.config import settings @@ -59,7 +60,7 @@ def run_migrations_offline() -> None: script output. """ - url = get_url() + url = ensure_session_pooler(get_url()) context.configure( url=url, @@ -74,6 +75,51 @@ def run_migrations_offline() -> None: context.run_migrations() +def ensure_session_pooler(connection_uri: str) -> str: + """ + Ensure a PostgreSQL connection URI uses the session pooler port (5432). + + Converts transaction pooler port (6543) to session pooler port (5432). + Leaves other ports unchanged. + + Args: + connection_uri: PostgreSQL connection URI + + Returns: + Connection URI with session pooler port (5432) + + Examples: + >>> ensure_session_pooler("postgresql://user:pass@host:6543/db") + 'postgresql://user:pass@host:5432/db' + + >>> ensure_session_pooler("postgresql://user:pass@host:5432/db") + 'postgresql://user:pass@host:5432/db' + + >>> ensure_session_pooler("postgresql+psycopg://user:pass@host.supabase.co:6543/postgres") + 'postgresql+psycopg://user:pass@host.supabase.co:5432/postgres' + """ + parsed = urlparse(connection_uri) + + # Get current port, default to 5432 if not specified + current_port = parsed.port or 5432 + + # If using transaction pooler port (6543), switch to session pooler (5432) + if current_port == 6543: + # Replace the port in the netloc + if parsed.port: + # If port is explicitly in the URL, replace it + new_netloc = parsed.netloc.replace(f":{current_port}", ":5432") + else: + # If port not in URL but somehow detected, add it + new_netloc = f"{parsed.netloc}:5432" + + # Reconstruct the URL with new port + new_parsed = parsed._replace(netloc=new_netloc) + return urlunparse(new_parsed) + + return connection_uri + + def run_migrations_online() -> None: """Run migrations in 'online' mode. @@ -87,13 +133,13 @@ def run_migrations_online() -> None: configuration = {} url = get_url() - configuration["sqlalchemy.url"] = url + validated_url = ensure_session_pooler(url) + configuration["sqlalchemy.url"] = validated_url connectable = engine_from_config( configuration, prefix="sqlalchemy.", - echo=False, - poolclass=pool.NullPool, + echo=True, connect_args={ "prepare_threshold": None, "options": "-c statement_timeout=300000", # 5 minutes in milliseconds diff --git a/migrations/versions/08894082221a_replace_collection_name_with_observer_.py b/migrations/versions/08894082221a_replace_collection_name_with_observer_.py index 87067d67..481ae33f 100644 --- a/migrations/versions/08894082221a_replace_collection_name_with_observer_.py +++ b/migrations/versions/08894082221a_replace_collection_name_with_observer_.py @@ -22,6 +22,9 @@ down_revision: str | None = "564ba40505c5" branch_labels: str | Sequence[str] | None = None depends_on: str | Sequence[str] | None = None +# Batch size for bulk operations +BATCH_SIZE = 10000 + def upgrade() -> None: """Replace collections.name and documents.collection_name with observer and observed fields.""" @@ -56,7 +59,6 @@ def upgrade() -> None: {"session_id": session_id, "workspace_name": workspace_name}, ) # Update all documents with NULL session_name in batches - batch_size = 5000 while True: result = connection.execute( text( @@ -65,17 +67,15 @@ def upgrade() -> None: SELECT id FROM {schema}.documents WHERE session_name IS NULL - ORDER BY id LIMIT :batch_size ) UPDATE {schema}.documents d SET session_name = '__global_observations__' FROM batch WHERE d.id = batch.id - AND d.session_name IS NULL """ ), - {"batch_size": batch_size}, + {"batch_size": BATCH_SIZE}, ) if result.rowcount == 0: break @@ -117,31 +117,37 @@ def upgrade() -> None: # Step 2b: Delete documents that reference collections marked for deletion (in batches) if collections_to_delete: - # Delete documents in batches - batch_size = 5000 - for i in range(0, len(collections_to_delete), batch_size): - batch = collections_to_delete[i : i + batch_size] - collection_ids = [row.id for row in batch] + collection_ids = [row.id for row in collections_to_delete] - connection.execute( + # Delete in smaller chunks of actual documents to avoid locking issues + while True: + result = connection.execute( text( f""" - DELETE FROM {schema}.documents d - USING {schema}.collections c - WHERE d.collection_name = c.name - AND d.peer_name = c.peer_name - AND d.workspace_name = c.workspace_name - AND c.id = ANY(:collection_ids) - """ + WITH to_delete AS ( + SELECT d.id + FROM {schema}.documents d + JOIN {schema}.collections c + ON d.collection_name = c.name + AND d.peer_name = c.peer_name + AND d.workspace_name = c.workspace_name + WHERE c.id = ANY(:collection_ids) + LIMIT :batch_size + ) + DELETE FROM {schema}.documents + WHERE id IN (SELECT id FROM to_delete) + """ ), - {"collection_ids": collection_ids}, + {"collection_ids": collection_ids, "batch_size": BATCH_SIZE}, ) + if result.rowcount == 0: + break # No more rows to delete + # Step 2c: Delete the collections identified in step 2a (in batches) if collections_to_delete: - batch_size = 5000 - for i in range(0, len(collections_to_delete), batch_size): - batch = collections_to_delete[i : i + batch_size] + for i in range(0, len(collections_to_delete), BATCH_SIZE): + batch = collections_to_delete[i : i + BATCH_SIZE] collection_ids = [row.id for row in batch] connection.execute( @@ -161,7 +167,17 @@ def upgrade() -> None: # - If name starts with peer_name + "_", extract the observed part (pattern: observer_observed) # - If name ends with "_" + peer_name, extract the first part (pattern: observed_observer) # - Any legacy edge cases will have been deleted in step 2a. - batch_size = 5000 + + # Create temporary index to speed up batching + if not index_exists("collections", "idx_temp_collections_null_observer"): + op.create_index( + "idx_temp_collections_null_observer", + "collections", + ["id"], + postgresql_where=text("observer IS NULL OR observed IS NULL"), + schema=schema, + ) + while True: result = connection.execute( text( @@ -170,7 +186,6 @@ def upgrade() -> None: SELECT id FROM {schema}.collections WHERE observer IS NULL OR observed IS NULL - ORDER BY id LIMIT :batch_size ) UPDATE {schema}.collections c @@ -186,11 +201,17 @@ def upgrade() -> None: WHERE c.id = batch.id """ ), - {"batch_size": batch_size}, + {"batch_size": BATCH_SIZE}, ) if result.rowcount == 0: break + # Drop temporary index after batching is complete + if index_exists("collections", "idx_temp_collections_null_observer", inspector): + op.drop_index( + "idx_temp_collections_null_observer", "collections", schema=schema + ) + # Step 3: Make collections observer and observed NOT NULL op.alter_column("collections", "observer", nullable=False, schema=schema) op.alter_column("collections", "observed", nullable=False, schema=schema) @@ -213,14 +234,24 @@ def upgrade() -> None: # Step 5: Populate documents observer and observed from collections table # Join to the already-populated collections table to get authoritative values - # Process in batches of 1000 to reduce query size - batch_size = 1000 + # Process in batches to reduce query size + + # Create temporary index to speed up batching + if not index_exists("documents", "idx_temp_docs_null_observer"): + op.create_index( + "idx_temp_docs_null_observer", + "documents", + ["id"], + postgresql_where=text("observer IS NULL OR observed IS NULL"), + schema=schema, + ) + while True: result = connection.execute( text( f""" WITH batch AS ( - SELECT d.ctid + SELECT d.id FROM {schema}.documents d WHERE d.observer IS NULL OR d.observed IS NULL LIMIT :batch_size @@ -230,18 +261,21 @@ def upgrade() -> None: observer = c.observer, observed = c.observed FROM {schema}.collections c, batch - WHERE d.ctid = batch.ctid + WHERE d.id = batch.id AND d.collection_name = c.name AND d.peer_name = c.peer_name AND d.workspace_name = c.workspace_name - AND (d.observer IS NULL OR d.observed IS NULL) """ ), - {"batch_size": batch_size}, + {"batch_size": BATCH_SIZE}, ) if result.rowcount == 0: break + # Drop temporary index after batching is complete + if index_exists("documents", "idx_temp_docs_null_observer", inspector): + op.drop_index("idx_temp_docs_null_observer", "documents", schema=schema) + # Step 6: Make documents observer and observed NOT NULL op.alter_column("documents", "observer", nullable=False, schema=schema) op.alter_column("documents", "observed", nullable=False, schema=schema) @@ -308,10 +342,10 @@ def upgrade() -> None: # Step 10: Add composite foreign key constraint for observer peer on collections if not fk_exists( - "collections", "collections_observer_workspace_name_fkey", inspector + "collections", "fk_collections_observer_workspace_name_peers", inspector ): op.create_foreign_key( - "collections_observer_workspace_name_fkey", + "fk_collections_observer_workspace_name_peers", "collections", "peers", ["observer", "workspace_name"], @@ -322,10 +356,10 @@ def upgrade() -> None: # Step 11: Add composite foreign key constraint for observed peer on collections if not fk_exists( - "collections", "collections_observed_workspace_name_fkey", inspector + "collections", "fk_collections_observed_workspace_name_peers", inspector ): op.create_foreign_key( - "collections_observed_workspace_name_fkey", + "fk_collections_observed_workspace_name_peers", "collections", "peers", ["observed", "workspace_name"], @@ -336,10 +370,12 @@ def upgrade() -> None: # Step 12: Add composite foreign key constraint from documents to collections using observer/observed if not fk_exists( - "documents", "documents_observer_observed_workspace_name_fkey", inspector + "documents", + "fk_documents_observer_observed_workspace_name_collections", + inspector, ): op.create_foreign_key( - "documents_observer_observed_workspace_name_fkey", + "fk_documents_observer_observed_workspace_name_collections", "documents", "collections", ["observer", "observed", "workspace_name"], @@ -349,9 +385,11 @@ def upgrade() -> None: ) # Step 13: Add composite foreign key constraint for observer peer on documents - if not fk_exists("documents", "documents_observer_workspace_name_fkey", inspector): + if not fk_exists( + "documents", "fk_documents_observer_workspace_name_peers", inspector + ): op.create_foreign_key( - "documents_observer_workspace_name_fkey", + "fk_documents_observer_workspace_name_peers", "documents", "peers", ["observer", "workspace_name"], @@ -361,9 +399,11 @@ def upgrade() -> None: ) # Step 14: Add composite foreign key constraint for observed peer on documents - if not fk_exists("documents", "documents_observed_workspace_name_fkey", inspector): + if not fk_exists( + "documents", "fk_documents_observed_workspace_name_peers", inspector + ): op.create_foreign_key( - "documents_observed_workspace_name_fkey", + "fk_documents_observed_workspace_name_peers", "documents", "peers", ["observed", "workspace_name"], @@ -498,7 +538,6 @@ def downgrade() -> None: ) # Step 5: Populate documents collection_name from observer and observed in batches - batch_size = 5000 while True: result = connection.execute( text( @@ -507,7 +546,6 @@ def downgrade() -> None: SELECT id FROM {schema}.documents WHERE collection_name IS NULL - ORDER BY id LIMIT :batch_size ) UPDATE {schema}.documents d @@ -520,7 +558,7 @@ def downgrade() -> None: AND d.collection_name IS NULL """ ), - {"batch_size": batch_size}, + {"batch_size": BATCH_SIZE}, ) if result.rowcount == 0: break @@ -537,7 +575,6 @@ def downgrade() -> None: ) # Populate peer_name with observed value in batches - batch_size = 5000 while True: result = connection.execute( text( @@ -546,7 +583,6 @@ def downgrade() -> None: SELECT id FROM {schema}.documents WHERE peer_name IS NULL - ORDER BY id LIMIT :batch_size ) UPDATE {schema}.documents d @@ -556,7 +592,7 @@ def downgrade() -> None: AND d.peer_name IS NULL """ ), - {"batch_size": batch_size}, + {"batch_size": BATCH_SIZE}, ) if result.rowcount == 0: break @@ -565,9 +601,11 @@ def downgrade() -> None: op.alter_column("documents", "peer_name", nullable=False, schema=schema) # Recreate the foreign key constraint for peer_name on documents - if not fk_exists("documents", "documents_peer_name_workspace_name_fkey", inspector): + if not fk_exists( + "documents", "fk_documents_peer_name_workspace_name_peers", inspector + ): op.create_foreign_key( - "documents_peer_name_workspace_name_fkey", + "fk_documents_peer_name_workspace_name_peers", "documents", "peers", ["peer_name", "workspace_name"], @@ -588,26 +626,28 @@ def downgrade() -> None: # CONSTRAINTS AND INDEXES # Step 8: Drop new foreign key constraints from documents if fk_exists( - "documents", "documents_observer_observed_workspace_name_fkey", inspector + "documents", + "fk_documents_observer_observed_workspace_name_collections", + inspector, ): op.drop_constraint( - "documents_observer_observed_workspace_name_fkey", + "fk_documents_observer_observed_workspace_name_collections", "documents", type_="foreignkey", schema=schema, ) - if fk_exists("documents", "documents_observer_workspace_name_fkey", inspector): + if fk_exists("documents", "fk_documents_observer_workspace_name_peers", inspector): op.drop_constraint( - "documents_observer_workspace_name_fkey", + "fk_documents_observer_workspace_name_peers", "documents", type_="foreignkey", schema=schema, ) - if fk_exists("documents", "documents_observed_workspace_name_fkey", inspector): + if fk_exists("documents", "fk_documents_observed_workspace_name_peers", inspector): op.drop_constraint( - "documents_observed_workspace_name_fkey", + "fk_documents_observed_workspace_name_peers", "documents", type_="foreignkey", schema=schema, @@ -652,17 +692,21 @@ def downgrade() -> None: ) # Step 12: Drop foreign key constraints from collections - if fk_exists("collections", "collections_observer_workspace_name_fkey", inspector): + if fk_exists( + "collections", "fk_collections_observer_workspace_name_peers", inspector + ): op.drop_constraint( - "collections_observer_workspace_name_fkey", + "fk_collections_observer_workspace_name_peers", "collections", type_="foreignkey", schema=schema, ) - if fk_exists("collections", "collections_observed_workspace_name_fkey", inspector): + if fk_exists( + "collections", "fk_collections_observed_workspace_name_peers", inspector + ): op.drop_constraint( - "collections_observed_workspace_name_fkey", + "fk_collections_observed_workspace_name_peers", "collections", type_="foreignkey", schema=schema, diff --git a/migrations/versions/564ba40505c5_add_session_name_column_to_documents.py b/migrations/versions/564ba40505c5_add_session_name_column_to_documents.py index d7238385..36274e12 100644 --- a/migrations/versions/564ba40505c5_add_session_name_column_to_documents.py +++ b/migrations/versions/564ba40505c5_add_session_name_column_to_documents.py @@ -38,7 +38,17 @@ def upgrade() -> None: # Process in batches to avoid timeout with large datasets # Only migrate documents that have 'session_name' key with a non-null, non-empty value bind = op.get_bind() - batch_size = 5000 + batch_size = 10000 + + # Create temporary index to speed up batching + if not index_exists("documents", "idx_temp_docs_null_session", inspector): + op.create_index( + "idx_temp_docs_null_session", + "documents", + ["id"], + postgresql_where=sa.text("session_name IS NULL"), + schema=schema, + ) while True: result = bind.execute( @@ -51,14 +61,12 @@ def upgrade() -> None: AND internal_metadata ? 'session_name' AND internal_metadata->>'session_name' IS NOT NULL AND internal_metadata->>'session_name' != '' - ORDER BY id LIMIT :batch_size ) UPDATE {schema}.documents d SET session_name = d.internal_metadata->>'session_name' FROM batch b WHERE d.id = b.id - AND d.session_name IS NULL """ ), {"batch_size": batch_size}, @@ -67,6 +75,10 @@ def upgrade() -> None: if result.rowcount == 0: break + # Drop temporary index after batching is complete + if index_exists("documents", "idx_temp_docs_null_session"): + op.drop_index("idx_temp_docs_null_session", "documents", schema=schema) + # Step 3: Create index on session_name for efficient querying if not index_exists("documents", "idx_documents_session_name", inspector): op.create_index( diff --git a/tests/alembic/revisions/test_08894082221a_replace_collection_name_with_observer_.py b/tests/alembic/revisions/test_08894082221a_replace_collection_name_with_observer_.py index 30f9aeb6..2cf7951b 100644 --- a/tests/alembic/revisions/test_08894082221a_replace_collection_name_with_observer_.py +++ b/tests/alembic/revisions/test_08894082221a_replace_collection_name_with_observer_.py @@ -262,7 +262,9 @@ def verify_observer_observed_migration(verifier: MigrationVerifier) -> None: "collections", "unique_observer_observed_collection", "unique" ) verifier.assert_constraint_exists( - "documents", "documents_observer_observed_workspace_name_fkey", "foreign_key" + "documents", + "fk_documents_observer_observed_workspace_name_collections", + "foreign_key", ) collection = verifier.conn.execute( From 342d72184aae72e67ab0f8a0c4b5eac77eaf31e2 Mon Sep 17 00:00:00 2001 From: doria <93405247+dr-frmr@users.noreply.github.com> Date: Mon, 3 Nov 2025 11:25:29 -0500 Subject: [PATCH 13/17] feat: rework langfuse setup to work more cleanly; fix bug in dream scheduling (#253) --- src/crud/representation.py | 3 +- src/deriver/consumer.py | 75 ++++++---------------------------- src/deriver/deriver.py | 16 ++------ src/deriver/enqueue.py | 2 +- src/dialectic/chat.py | 14 +------ src/dreamer/dream_scheduler.py | 13 +++--- src/dreamer/dreamer.py | 2 + src/utils/clients.py | 24 ++--------- src/utils/langfuse_client.py | 40 ------------------ src/utils/logging.py | 50 +++++++++++++++++++---- src/utils/summarizer.py | 4 +- tests/utils/test_clients.py | 19 --------- 12 files changed, 77 insertions(+), 185 deletions(-) delete mode 100644 src/utils/langfuse_client.py diff --git a/src/crud/representation.py b/src/crud/representation.py index 09eaeae7..13aefce7 100644 --- a/src/crud/representation.py +++ b/src/crud/representation.py @@ -14,7 +14,7 @@ from src.dependencies import tracked_db from src.dreamer.dream_scheduler import check_and_schedule_dream from src.embedding_client import embedding_client from src.utils.formatting import format_datetime_utc -from src.utils.logging import accumulate_metric, conditional_observe +from src.utils.logging import accumulate_metric from src.utils.representation import ( DeductiveObservation, ExplicitObservation, @@ -41,7 +41,6 @@ class RepresentationManager: self.observer: str = observer self.observed: str = observed - @conditional_observe async def save_representation( self, representation: Representation, diff --git a/src/deriver/consumer.py b/src/deriver/consumer.py index 20e0b1b0..6cafc0fc 100644 --- a/src/deriver/consumer.py +++ b/src/deriver/consumer.py @@ -7,13 +7,11 @@ from rich.console import Console from sqlalchemy import select from src import models -from src.config import settings from src.dependencies import tracked_db from src.deriver.deriver import process_representation_tasks_batch from src.dreamer.dreamer import process_dream from src.models import Message from src.utils import summarizer -from src.utils.langfuse_client import get_langfuse_client from src.utils.logging import log_performance_metrics from src.utils.queue_payload import ( DreamPayload, @@ -27,8 +25,6 @@ logging.getLogger("sqlalchemy.engine.Engine").disabled = True console = Console(markup=True) -lf = get_langfuse_client() if settings.LANGFUSE_PUBLIC_KEY else None - async def process_item(task_type: str, queue_payload: dict[str, Any]) -> None: """Process a single item from the queue.""" @@ -80,39 +76,16 @@ async def process_item(task_type: str, queue_payload: dict[str, Any]) -> None: message_public_id = message.public_id with sentry_sdk.start_transaction(name="process_summary_task", op="deriver"): - if lf: - with lf.start_as_current_span( - name="summary_processing", - input={ - "workspace_name": validated.workspace_name, - "session_name": validated.session_name, - "message_id": validated.message_id, - }, - metadata={ - "summary_model": settings.SUMMARY.MODEL, - }, - ): - await summarizer.summarize_if_needed( - validated.workspace_name, - validated.session_name, - validated.message_id, - validated.message_seq_in_session, - message_public_id, - ) - log_performance_metrics( - "summary", f"{validated.workspace_name}_{validated.message_id}" - ) - else: - await summarizer.summarize_if_needed( - validated.workspace_name, - validated.session_name, - validated.message_id, - validated.message_seq_in_session, - message_public_id, - ) - log_performance_metrics( - "summary", f"{validated.workspace_name}_{validated.message_id}" - ) + await summarizer.summarize_if_needed( + validated.workspace_name, + validated.session_name, + validated.message_id, + validated.message_seq_in_session, + message_public_id, + ) + log_performance_metrics( + "summary", f"{validated.workspace_name}_{validated.message_id}" + ) elif task_type == "dream": with sentry_sdk.start_transaction(name="process_dream_task", op="deriver"): @@ -162,28 +135,6 @@ async def process_representation_batch( len(messages), ) - if lf: - with lf.start_as_current_span( - name="representation_processing", - input={ - "payloads": [ - { - "message_id": msg.id, - "observer": observer, - "observed": observed, - "session_name": msg.session_name, - } - for msg in messages - ] - }, - metadata={ - "critical_analysis_model": settings.DERIVER.MODEL, - }, - ): - await process_representation_tasks_batch( - messages, observer=observer, observed=observed - ) - else: - await process_representation_tasks_batch( - messages, observer=observer, observed=observed - ) + await process_representation_tasks_batch( + messages, observer=observer, observed=observed + ) diff --git a/src/deriver/deriver.py b/src/deriver/deriver.py index 15900e86..5990a29b 100644 --- a/src/deriver/deriver.py +++ b/src/deriver/deriver.py @@ -12,7 +12,6 @@ from src.models import Message from src.utils import summarizer from src.utils.clients import honcho_llm_call from src.utils.formatting import format_new_turn_with_timestamp -from src.utils.langfuse_client import get_langfuse_client from src.utils.logging import ( accumulate_metric, conditional_observe, @@ -33,9 +32,8 @@ from .prompts import ( logger = logging.getLogger(__name__) logging.getLogger("sqlalchemy.engine.Engine").disabled = True -lf = get_langfuse_client() if settings.LANGFUSE_PUBLIC_KEY else None - +@conditional_observe(name="Critical Analysis Call") async def critical_analysis_call( peer_id: str, peer_card: list[str] | None, @@ -78,6 +76,7 @@ async def critical_analysis_call( return response.content +@conditional_observe(name="Peer Card Call") async def peer_card_call( old_peer_card: list[str] | None, new_observations: Representation, @@ -281,9 +280,6 @@ async def process_representation_tasks_batch( log_performance_metrics("deriver", f"{latest_message.id}_{observer}") - if lf: - lf.update_current_trace(output=final_observations.format_as_markdown()) - class CertaintyReasoner: """Certainty reasoner for analyzing and deriving insights.""" @@ -308,7 +304,7 @@ class CertaintyReasoner: self.observer = observer self.estimated_input_tokens: int = estimated_input_tokens - @conditional_observe + @conditional_observe(name="Deriver") @sentry_sdk.trace async def reason( self, @@ -364,11 +360,6 @@ class CertaintyReasoner: latest_message.created_at, ) - if lf: - lf.update_current_generation( - output=reasoning_response.format_as_markdown(), - ) - analysis_duration_ms = (time.perf_counter() - analysis_start) * 1000 accumulate_metric( f"deriver_{latest_message.id}_{self.observer}", @@ -413,7 +404,6 @@ class CertaintyReasoner: return reasoning_response - @conditional_observe @sentry_sdk.trace async def _update_peer_card( self, diff --git a/src/deriver/enqueue.py b/src/deriver/enqueue.py index 9405f6fa..9731b91c 100644 --- a/src/deriver/enqueue.py +++ b/src/deriver/enqueue.py @@ -32,7 +32,7 @@ async def enqueue(payload: list[dict[str, Any]]) -> None: # Generate work unit keys for dreams that might be affected by this message dream_keys: list[str] = get_affected_dream_keys(message) for dream_key in dream_keys: - if dream_scheduler.cancel_dream(dream_key): + if await dream_scheduler.cancel_dream(dream_key): cancelled_dreams.add(dream_key) if cancelled_dreams: diff --git a/src/dialectic/chat.py b/src/dialectic/chat.py index b134b043..d90977e9 100644 --- a/src/dialectic/chat.py +++ b/src/dialectic/chat.py @@ -18,9 +18,9 @@ from src.config import settings from src.dependencies import tracked_db from src.utils import summarizer from src.utils.clients import HonchoLLMCallStreamChunk, honcho_llm_call -from src.utils.langfuse_client import get_langfuse_client from src.utils.logging import ( accumulate_metric, + conditional_observe, log_performance_metrics, ) from src.utils.representation import Representation @@ -34,9 +34,6 @@ logger = logging.getLogger(__name__) # Load environment variables load_dotenv() -# Create langfuse client -lf = get_langfuse_client() if settings.LANGFUSE_PUBLIC_KEY else None - async def dialectic_call( query: str, @@ -151,6 +148,7 @@ async def dialectic_stream( return response +@conditional_observe(name="Dialectic") async def chat( workspace_name: str, session_name: str | None, @@ -189,14 +187,6 @@ async def chat( context_window_size -= estimate_tokens(query) - if lf: - lf.update_current_trace( - metadata={ - "query_generation_model": settings.DIALECTIC.QUERY_GENERATION_MODEL, - "query_generation_provider": settings.DIALECTIC.QUERY_GENERATION_PROVIDER, - "dialectic_model": settings.DIALECTIC.MODEL, - } - ) accumulate_metric( f"dialectic_chat_{dialectic_chat_uuid}", "query", diff --git a/src/dreamer/dream_scheduler.py b/src/dreamer/dream_scheduler.py index 8c572e1c..a56edd3f 100644 --- a/src/dreamer/dream_scheduler.py +++ b/src/dreamer/dream_scheduler.py @@ -1,4 +1,5 @@ import asyncio +import contextlib from datetime import datetime, timezone from logging import getLogger from typing import Any @@ -80,7 +81,7 @@ class DreamScheduler: cls._instance = None cls._initialized = False - def schedule_dream( + async def schedule_dream( self, work_unit_key: str, workspace_name: str, @@ -95,7 +96,7 @@ class DreamScheduler: return # Cancel any existing dream for this collection - self.cancel_dream(work_unit_key) + await self.cancel_dream(work_unit_key) task = asyncio.create_task( self._delayed_dream( @@ -110,12 +111,14 @@ class DreamScheduler: self.pending_dreams[work_unit_key] = task task.add_done_callback(lambda t: self.pending_dreams.pop(work_unit_key, None)) - def cancel_dream(self, work_unit_key: str) -> bool: + async def cancel_dream(self, work_unit_key: str) -> bool: """Cancel a pending dream. Returns True if a dream was cancelled.""" if work_unit_key in self.pending_dreams: task = self.pending_dreams.pop(work_unit_key) task.cancel() - logger.debug(f"Cancelled pending dream for {work_unit_key}") + # Wait for the task to actually finish (including its done callback) + with contextlib.suppress(asyncio.CancelledError): + await task return True return False @@ -331,7 +334,7 @@ async def check_and_schedule_dream( } ) - dream_scheduler.schedule_dream( + await dream_scheduler.schedule_dream( collection_work_unit_key, collection.workspace_name, current_document_count, diff --git a/src/dreamer/dreamer.py b/src/dreamer/dreamer.py index bac32db9..9ee43bb8 100644 --- a/src/dreamer/dreamer.py +++ b/src/dreamer/dreamer.py @@ -11,6 +11,7 @@ from src.dreamer.prompts import consolidation_prompt from src.embedding_client import embedding_client from src.utils.clients import honcho_llm_call from src.utils.formatting import format_datetime_utc +from src.utils.logging import conditional_observe from src.utils.queue_payload import DreamPayload from src.utils.representation import ( ExplicitObservation, @@ -176,6 +177,7 @@ async def _consolidate_cluster( ) +@conditional_observe(name="[Dream] Consolidate Call") async def consolidate_call( representation: Representation, ) -> Representation: diff --git a/src/utils/clients.py b/src/utils/clients.py index 920af15a..9a0628ef 100644 --- a/src/utils/clients.py +++ b/src/utils/clients.py @@ -1,7 +1,6 @@ import json import logging -from collections.abc import AsyncIterator, Callable -from functools import wraps +from collections.abc import AsyncIterator from typing import Any, Generic, Literal, TypeVar, cast, overload from anthropic import AsyncAnthropic @@ -18,7 +17,7 @@ from tenacity import retry, stop_after_attempt, wait_exponential from src.config import settings from src.utils.json_parser import validate_and_repair_json -from src.utils.langfuse_client import get_langfuse_client +from src.utils.logging import conditional_observe from src.utils.representation import PromptRepresentation from src.utils.types import SupportedProviders @@ -27,8 +26,6 @@ logger = logging.getLogger(__name__) T = TypeVar("T") M = TypeVar("M", bound=BaseModel) -lf = get_langfuse_client() if settings.LANGFUSE_PUBLIC_KEY else None - CLIENTS: dict[ SupportedProviders, AsyncAnthropic | AsyncOpenAI | genai.Client | AsyncGroq, @@ -168,6 +165,7 @@ async def honcho_llm_call( ) -> AsyncIterator[HonchoLLMCallStreamChunk]: ... +@conditional_observe(name="LLM Call") async def honcho_llm_call( provider: SupportedProviders, model: str, @@ -191,10 +189,6 @@ async def honcho_llm_call( decorated = honcho_llm_call_inner - # apply langfuse if enabled - if settings.LANGFUSE_PUBLIC_KEY: - decorated = with_langfuse(decorated) - # apply tracking if track_name: decorated = ai_track(track_name)(decorated) @@ -781,15 +775,3 @@ async def handle_streaming_response( is_done=True, finish_reasons=[chunk.choices[0].finish_reason], ) - - -def with_langfuse(func: Callable[..., Any]) -> Callable[..., Any]: - @wraps(func) - async def wrapper(*args: Any, **kwargs: Any) -> Any: - if lf: - with lf.start_as_current_generation(name="LLM Call"): - return await func(*args, **kwargs) - else: - return await func(*args, **kwargs) - - return wrapper diff --git a/src/utils/langfuse_client.py b/src/utils/langfuse_client.py deleted file mode 100644 index c6e68c0b..00000000 --- a/src/utils/langfuse_client.py +++ /dev/null @@ -1,40 +0,0 @@ -""" -Centralized Langfuse client management. - -This module provides a singleton Langfuse client to avoid multiple initialization -errors when modules import get_client() at the module level. -""" - -from typing import Any - -from langfuse import get_client - -_langfuse_client: Any = None - - -def get_langfuse_client() -> Any: - """ - Get the singleton Langfuse client instance. - - This function ensures that get_client() is only called once, regardless of - how many modules import this function. This prevents multiple authentication - error messages when LANGFUSE_PUBLIC_KEY is not configured. - - Returns: - Any: The singleton Langfuse client instance - """ - global _langfuse_client - - if _langfuse_client is None: - _langfuse_client = get_client() - - return _langfuse_client - - -# For backward compatibility, provide the client as a module-level variable -# but only initialize it when first accessed -def __getattr__(name: str): - """Lazy initialization of module-level 'lf' attribute.""" - if name == "lf": - return get_langfuse_client() - raise AttributeError(f"module '{__name__}' has no attribute '{name}'") diff --git a/src/utils/logging.py b/src/utils/logging.py index ca615899..2002b7bd 100644 --- a/src/utils/logging.py +++ b/src/utils/logging.py @@ -6,9 +6,10 @@ and a conditional observe decorator that only applies when Langfuse is configure import datetime from collections.abc import Callable -from typing import Any +from typing import ParamSpec, TypeVar, overload from fastapi import Request +from langfuse import observe # pyright: ignore from rich import box from rich.console import Console, Group, RenderableType from rich.panel import Panel @@ -27,25 +28,56 @@ console = Console(markup=True) COLLECT_METRICS_LOCAL = settings.COLLECT_METRICS_LOCAL +P = ParamSpec("P") +R = TypeVar("R") -def conditional_observe(func: Callable[..., Any]) -> Callable[..., Any]: + +@overload +def conditional_observe( + func: Callable[P, R], +) -> Callable[P, R]: ... + + +@overload +def conditional_observe( + *, + name: str, +) -> Callable[[Callable[P, R]], Callable[P, R]]: ... + + +def conditional_observe( + func: Callable[P, R] | None = None, + *, + name: str | None = None, +) -> Callable[P, R] | Callable[[Callable[P, R]], Callable[P, R]]: """ Conditionally apply the @observe decorator only when LANGFUSE_PUBLIC_KEY is present. + Can be used in two ways: + 1. As a decorator: @conditional_observe + 2. As a decorator factory: @conditional_observe(name="...") + Args: - func: The function to potentially decorate + func: The function to potentially decorate (when used as @conditional_observe) + name: Optional name for the observation (when used as @conditional_observe(name="...")) Returns: The decorated function if Langfuse is configured, otherwise the original function """ - if settings.LANGFUSE_PUBLIC_KEY: - # Import here to avoid circular imports and only import when needed - from langfuse import observe # pyright: ignore - return observe()(func) + def decorator(f: Callable[P, R]) -> Callable[P, R]: + if settings.LANGFUSE_PUBLIC_KEY: + observe_name = name if name is not None else f.__name__ + return observe(name=observe_name)(f) + else: + return f + + if func is not None: + # Used as @conditional_observe (without parentheses) + return decorator(func) else: - # Return the function unchanged if Langfuse is not configured - return func + # Used as @conditional_observe(name="...") (with parentheses and keyword args) + return decorator # dict[task_name, list[tuple[metric_name, metric_value, metric_unit]]] diff --git a/src/utils/summarizer.py b/src/utils/summarizer.py index be21009d..065e3f99 100644 --- a/src/utils/summarizer.py +++ b/src/utils/summarizer.py @@ -14,7 +14,7 @@ from src.dependencies import tracked_db from src.exceptions import ResourceNotFoundException from src.utils.clients import HonchoLLMCallResponse, honcho_llm_call from src.utils.formatting import utc_now_iso -from src.utils.logging import accumulate_metric +from src.utils.logging import accumulate_metric, conditional_observe from .. import crud, models @@ -79,6 +79,7 @@ class SummaryType(Enum): LONG = "honcho_chat_summary_long" +@conditional_observe(name="Create Short Summary") async def create_short_summary( messages: list[models.Message], input_tokens: int, @@ -129,6 +130,7 @@ Produce as thorough a summary as possible in {output_words} words or less. ) +@conditional_observe(name="Create Long Summary") async def create_long_summary( messages: list[models.Message], previous_summary: str | None = None, diff --git a/tests/utils/test_clients.py b/tests/utils/test_clients.py index c6552926..54909d8d 100644 --- a/tests/utils/test_clients.py +++ b/tests/utils/test_clients.py @@ -32,7 +32,6 @@ from src.utils.clients import ( handle_streaming_response, honcho_llm_call, honcho_llm_call_inner, - with_langfuse, ) @@ -89,24 +88,6 @@ class TestLLMCallResponse: assert chunk.finish_reasons == [] -class TestLangfuseIntegration: - """Tests for Langfuse integration""" - - @pytest.mark.asyncio - async def test_with_langfuse_decorator(self): - """Test Langfuse decorator functionality""" - - @with_langfuse - async def test_func(): - return "decorated" - - # Mock the langfuse client - with patch("src.utils.clients.lf") as mock_lf: - result = await test_func() - assert result == "decorated" - mock_lf.start_as_current_generation.assert_called_once_with(name="LLM Call") - - @pytest.mark.asyncio class TestAnthropicClient: """Tests for Anthropic client functionality""" From 1df47e61c8e57ab4d98469bbe7fd7c96884837bd Mon Sep 17 00:00:00 2001 From: doria <93405247+dr-frmr@users.noreply.github.com> Date: Mon, 3 Nov 2025 11:26:44 -0500 Subject: [PATCH 14/17] fix: remove list webhook db call (not needed, causes race condition) (#257) * fix: remove get_workspace call from list_webhooks (not necessary, causes race condition) * fix: remove old test --- src/crud/webhook.py | 5 +---- src/routers/webhooks.py | 2 +- src/webhooks/webhook_delivery.py | 2 +- tests/routes/test_webhooks.py | 7 ------- 4 files changed, 3 insertions(+), 13 deletions(-) diff --git a/src/crud/webhook.py b/src/crud/webhook.py index 1dcc29a8..65b05569 100644 --- a/src/crud/webhook.py +++ b/src/crud/webhook.py @@ -63,7 +63,7 @@ async def get_or_create_webhook_endpoint( async def list_webhook_endpoints( - db: AsyncSession, workspace_name: str + workspace_name: str, ) -> Select[tuple[models.WebhookEndpoint]]: """ List all webhook endpoints, optionally filtered by workspace. @@ -75,9 +75,6 @@ async def list_webhook_endpoints( Returns: List of webhook endpoints """ - # Verify workspace exists - await get_workspace(db, workspace_name) - return select(models.WebhookEndpoint).where( models.WebhookEndpoint.workspace_name == workspace_name ) diff --git a/src/routers/webhooks.py b/src/routers/webhooks.py index c6ae2017..aa3256c1 100644 --- a/src/routers/webhooks.py +++ b/src/routers/webhooks.py @@ -61,7 +61,7 @@ async def list_webhook_endpoints( if not jwt_params.ad and jwt_params.w is not None and jwt_params.w != workspace_id: raise AuthenticationException("Unauthorized access to resource") - stmt = await crud.list_webhook_endpoints(db, workspace_id) + stmt = await crud.list_webhook_endpoints(workspace_id) return await apaginate(db, stmt) diff --git a/src/webhooks/webhook_delivery.py b/src/webhooks/webhook_delivery.py index d9587009..316fdccc 100644 --- a/src/webhooks/webhook_delivery.py +++ b/src/webhooks/webhook_delivery.py @@ -82,7 +82,7 @@ async def _get_webhook_urls(db: AsyncSession, workspace_name: str) -> list[str]: Get all webhook endpoint URLs for a workspace. """ try: - endpoints = await list_webhook_endpoints(db, workspace_name) + endpoints = await list_webhook_endpoints(workspace_name) result = await db.execute(endpoints) return [endpoint.url for endpoint in result.scalars().all()] except Exception: diff --git a/tests/routes/test_webhooks.py b/tests/routes/test_webhooks.py index ed176473..8272dc9f 100644 --- a/tests/routes/test_webhooks.py +++ b/tests/routes/test_webhooks.py @@ -95,13 +95,6 @@ async def test_list_webhook_endpoints_with_data( assert "http://example2.com/webhook" in endpoint_urls -@pytest.mark.asyncio -async def test_list_webhook_endpoints_missing_workspace(client: TestClient): - response = client.get("/v2/workspaces/nonexistent-workspace/webhooks") - assert response.status_code == 404 - assert response.json() == {"detail": "Workspace nonexistent-workspace not found"} - - @pytest.mark.asyncio async def test_delete_webhook_endpoint( client: TestClient, sample_data: tuple[Workspace, Peer] From 097f3b31a0c21393678c6e2fb73637b60508da53 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Mon, 3 Nov 2025 14:58:12 -0500 Subject: [PATCH 15/17] align DB state to sqlalchemy model definitions (#245) * feat: align DB schema with sqlalchemy model definitions * test: add test for new migration file * fix: add pk for message embeddings table * fix: add naming constraint * fix: align more indexes + unique constraints --> rm redundancy * fix: add missing indexes * Fix Server Default Migrations and add in appropriate Server Defaults (#252) * fix: Improve alembic migration reliability * fix: Ensure alembic migrations use a session pooler (#247) * fix: Ensure alembic migrations use a session pooler * fix: bump batch size; fix document delete in 08894082221a (#249) * chore: add alembic logging * fix (alembic): remove order by in batches * fix: fkey -> fk * chore: code rabbit --------- Co-authored-by: Rajat Ahuja * fix: Add server defaults to appropriate columns * chore: (tests) Add alembic migration test --------- Co-authored-by: Rajat Ahuja * fix: align models with alembic check --------- Co-authored-by: Vineeth Voruganti <13438633+VVoruganti@users.noreply.github.com> --- migrations/env.py | 13 +- migrations/script.py.mako | 10 +- migrations/utils.py | 86 ++++ ...07_align_schema_with_declarative_models.py | 341 +++++++++++++++ ..._replace_collection_name_with_observer_.py | 7 +- ...421aff_rename_metamessage_type_to_label.py | 5 +- ...9adf9_add_server_defaults_to_timestamp_.py | 387 ++++++++++++++++++ src/db.py | 2 + src/models.py | 170 +++++--- tests/alembic/revisions/__init__.py | 2 + ...07_align_schema_with_declarative_models.py | 42 ++ ...9adf9_add_server_defaults_to_timestamp_.py | 284 +++++++++++++ 12 files changed, 1285 insertions(+), 64 deletions(-) create mode 100644 migrations/versions/066e87ca5b07_align_schema_with_declarative_models.py create mode 100644 migrations/versions/e9b705f9adf9_add_server_defaults_to_timestamp_.py create mode 100644 tests/alembic/revisions/test_066e87ca5b07_align_schema_with_declarative_models.py create mode 100644 tests/alembic/revisions/test_e9b705f9adf9_add_server_defaults_to_timestamp_.py diff --git a/migrations/env.py b/migrations/env.py index f0a8197c..4ee654bd 100644 --- a/migrations/env.py +++ b/migrations/env.py @@ -1,4 +1,4 @@ -import logging +import logging # noqa: I001 import sys from logging.config import fileConfig from pathlib import Path @@ -12,6 +12,10 @@ from src.config import settings # Import your models from src.db import Base +# Import all models so they register with Base.metadata +import src.models # noqa: F401 + + # Set up logging more verbosely logging.basicConfig() logging.getLogger("sqlalchemy.engine").setLevel(logging.INFO) @@ -166,6 +170,13 @@ def run_migrations_online() -> None: connection=connection, target_metadata=target_metadata, version_table_schema=target_metadata.schema, + include_schemas=True, + include_object=lambda obj, name, type_, reflected, compare_to: ( + # Only include objects from our target schema + getattr(obj, "schema", None) == target_metadata.schema + if hasattr(obj, "schema") + else True + ), ) with context.begin_transaction(): diff --git a/migrations/script.py.mako b/migrations/script.py.mako index fbc4b07d..bba30146 100644 --- a/migrations/script.py.mako +++ b/migrations/script.py.mako @@ -5,18 +5,20 @@ Revises: ${down_revision | comma,n} Create Date: ${create_date} """ -from typing import Sequence, Union +from collections.abc import Sequence from alembic import op import sqlalchemy as sa ${imports if imports else ""} +from migrations.utils import get_schema # revision identifiers, used by Alembic. revision: str = ${repr(up_revision)} -down_revision: Union[str, None] = ${repr(down_revision)} -branch_labels: Union[str, Sequence[str], None] = ${repr(branch_labels)} -depends_on: Union[str, Sequence[str], None] = ${repr(depends_on)} +down_revision: str | None = ${repr(down_revision)} +branch_labels: str | Sequence[str] | None = ${repr(branch_labels)} +depends_on: str | Sequence[str] | None = ${repr(depends_on)} +schema = get_schema() def upgrade() -> None: ${upgrades if upgrades else "pass"} diff --git a/migrations/utils.py b/migrations/utils.py index 4333ae46..1dc74d86 100644 --- a/migrations/utils.py +++ b/migrations/utils.py @@ -71,3 +71,89 @@ def constraint_exists( else: raise ValueError(f"Invalid constraint type: {type}") return any(constraint["name"] == constraint_name for constraint in constraints) + + +def make_column_non_nullable_safe(table_name: str, column_name: str) -> None: + """ + Make a column non-nullable using a non-blocking approach to minimize lock duration. + + WARNING: Only use this if you can guarantee that: + 1. No NULL values currently exist in the column + 2. The application code is already writing non-NULL values to this column or + 3. The column has never accepted NULLs in practice + + This uses a 4-step process to avoid long exclusive locks: + 1. Add CHECK constraint with NOT VALID (instant, no scan) + 2. Validate the constraint (scans but allows concurrent read/writes to the table) + 3. Set column NOT NULL (fast since we've validated the constraint) + 4. Drop the redundant CHECK constraint + + Args: + table_name: The name of the table + column_name: The name of the column to make non-nullable + """ + schema = get_schema() + conn = op.get_bind() + constraint_name = f"{table_name}_{column_name}_not_null" + + # Step 1: Check if the column is already non-nullable + inspector = sa.inspect(op.get_bind()) + columns = inspector.get_columns(table_name, schema=schema) + column_info = next((col for col in columns if col["name"] == column_name), None) + if column_info is None: + raise ValueError(f"Column {table_name}.{column_name} does not exist") + if not column_info["nullable"]: + print(f"Column {table_name}.{column_name} is already non-nullable, skipping...") + return + + # Step 2: Add CHECK constraint without validation (instant) + # Note: op.create_check_constraint() doesn't support NOT VALID, so use raw SQL + + # Get the identifier preparer for safe quoting + dialect = conn.dialect + preparer = dialect.identifier_preparer + + quoted_schema = preparer.quote(schema) + quoted_table = preparer.quote(table_name) + quoted_constraint = preparer.quote(constraint_name) + quoted_column = preparer.quote(column_name) + + # Step 2: Add CHECK constraint without validation (instant) + # Note: op.create_check_constraint() doesn't support NOT VALID, so use raw SQL + if not constraint_exists(table_name, constraint_name, "check"): + conn.execute( + sa.text( + f""" + ALTER TABLE {quoted_schema}.{quoted_table} + ADD CONSTRAINT {quoted_constraint} + CHECK ({quoted_column} IS NOT NULL) + NOT VALID + """ + ) + ) + + # Step 3: Validate constraint (scans but allows concurrent operations) + conn.execute( + sa.text( + f""" + ALTER TABLE {quoted_schema}.{quoted_table} + VALIDATE CONSTRAINT {quoted_constraint} + """ + ) + ) + + # Step 4: Set NOT NULL (fast with validated constraint) + op.alter_column( + table_name, + column_name, + nullable=False, + schema=schema, + ) + + # Step 5: Drop the redundant CHECK constraint + op.drop_constraint( + constraint_name, + table_name, + type_="check", + schema=schema, + ) diff --git a/migrations/versions/066e87ca5b07_align_schema_with_declarative_models.py b/migrations/versions/066e87ca5b07_align_schema_with_declarative_models.py new file mode 100644 index 00000000..72966bf0 --- /dev/null +++ b/migrations/versions/066e87ca5b07_align_schema_with_declarative_models.py @@ -0,0 +1,341 @@ +"""align_schema_with_declarative_models + +Revision ID: 066e87ca5b07 +Revises: bb6fb3a7a643 +Create Date: 2025-10-27 12:36:51.614959 + +""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +from migrations.utils import ( + column_exists, + constraint_exists, + fk_exists, + get_schema, + index_exists, + make_column_non_nullable_safe, +) + +# revision identifiers, used by Alembic. +revision: str = "066e87ca5b07" +down_revision: str | None = "bb6fb3a7a643" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + +schema = get_schema() + + +def upgrade() -> None: + """ + The application code has been previously updated to ensure none of the following columns have NULL values but the actual DB schema is out of sync with our SQLAlchemy model definitions. + This migration fixes this by making the columns non-nullable using a non-blocking approach to minimize lock duration. + """ + conn = op.get_bind() + + # Make peers.workspace_name non-nullable + if column_exists("peers", "workspace_name"): + make_column_non_nullable_safe("peers", "workspace_name") + + # Make sessions.workspace_name non-nullable + if column_exists("sessions", "workspace_name"): + make_column_non_nullable_safe("sessions", "workspace_name") + + # Make active_queue_sessions.work_unit_key non-nullable + if column_exists("active_queue_sessions", "work_unit_key"): + make_column_non_nullable_safe("active_queue_sessions", "work_unit_key") + + # Make documents.embedding non-nullable + if column_exists("documents", "embedding"): + make_column_non_nullable_safe("documents", "embedding") + + # Add primary key constraint to message_embeddings.id + if column_exists("message_embeddings", "id") and not constraint_exists( + "message_embeddings", "pk_message_embeddings", "primary" + ): + conn.execute( + sa.text( + f""" + ALTER TABLE {schema}.message_embeddings + ADD CONSTRAINT pk_message_embeddings + PRIMARY KEY (id) + """ + ) + ) + + # Rename indexes on peers table + inspector = sa.inspect(conn) + index_renames = [ + ("peers", "ix_users_created_at", "ix_peers_created_at"), + ("peers", "ix_users_name", "ix_peers_name"), + ("workspaces", "ix_apps_created_at", "ix_workspaces_created_at"), + ("workspaces", "ix_apps_name", "ix_workspaces_name"), + ] + for table_name, old_name, new_name in index_renames: + if index_exists(table_name, old_name, inspector): + conn.execute( + sa.text(f"ALTER INDEX {schema}.{old_name} RENAME TO {new_name}") + ) + + # Drop redundant indexes + if index_exists("workspaces", "ix_apps_public_id", inspector): + op.drop_index("ix_apps_public_id", table_name="workspaces", schema=schema) + if index_exists("peers", "ix_users_public_id", inspector): + op.drop_index("ix_users_public_id", table_name="peers", schema=schema) + if index_exists("sessions", "ix_sessions_public_id", inspector): + op.drop_index("ix_sessions_public_id", table_name="sessions", schema=schema) + if index_exists("documents", "ix_documents_public_id", inspector): + op.drop_index("ix_documents_public_id", table_name="documents", schema=schema) + if index_exists("collections", "ix_collections_public_id", inspector): + op.drop_index( + "ix_collections_public_id", table_name="collections", schema=schema + ) + + # Drop redundant unique constraints + if constraint_exists("workspaces", "uq_apps_public_id", "unique", inspector): + op.drop_constraint( + "uq_apps_public_id", "workspaces", type_="unique", schema=schema + ) + + if constraint_exists("peers", "uq_users_public_id", "unique", inspector): + op.drop_constraint("uq_users_public_id", "peers", type_="unique", schema=schema) + + if constraint_exists("sessions", "uq_sessions_public_id", "unique", inspector): + op.drop_constraint( + "uq_sessions_public_id", "sessions", type_="unique", schema=schema + ) + + if constraint_exists( + "collections", "uq_collections_public_id", "unique", inspector + ): + op.drop_constraint( + "uq_collections_public_id", "collections", type_="unique", schema=schema + ) + + if constraint_exists("documents", "uq_documents_public_id", "unique", inspector): + op.drop_constraint( + "uq_documents_public_id", "documents", type_="unique", schema=schema + ) + + # Drop unnecessary index on active queue + if index_exists( + "active_queue_sessions", + f"ix_{schema}_active_queue_sessions_work_unit_key", + inspector, + ): + op.drop_index( + f"ix_{schema}_active_queue_sessions_work_unit_key", + table_name="active_queue_sessions", + schema=schema, + ) + + # Add FK constraint on queue.session_id to sessions.id + if not fk_exists("queue", "fk_queue_session_id"): + # Add constraint without validation (fast, doesn't scan) + conn.execute( + sa.text( + f""" + ALTER TABLE {schema}.queue + ADD CONSTRAINT fk_queue_session_id + FOREIGN KEY (session_id) + REFERENCES {schema}.sessions(id) + NOT VALID + """ + ) + ) + # Validate constraint (scans but allows concurrent reads) + conn.execute( + sa.text( + f"ALTER TABLE {schema}.queue VALIDATE CONSTRAINT fk_queue_session_id" + ) + ) + + # Create missing indexes + if not index_exists("peers", "ix_peers_workspace_name", inspector): + op.create_index( + "ix_peers_workspace_name", "peers", ["workspace_name"], schema=schema + ) + + if not index_exists("collections", "ix_collections_workspace_name", inspector): + op.create_index( + "ix_collections_workspace_name", + "collections", + ["workspace_name"], + schema=schema, + ) + + if not index_exists("documents", "ix_documents_workspace_name", inspector): + op.create_index( + "ix_documents_workspace_name", + "documents", + ["workspace_name"], + schema=schema, + ) + + if not fk_exists("session_peers", "fk_session_peers_workspace_name", inspector): + # Add constraint without validation (fast, doesn't scan) + conn.execute( + sa.text( + f""" + ALTER TABLE {schema}.session_peers + ADD CONSTRAINT fk_session_peers_workspace_name + FOREIGN KEY (workspace_name) + REFERENCES {schema}.workspaces(name) + NOT VALID + """ + ) + ) + # Validate constraint (scans but allows concurrent reads) + conn.execute( + sa.text( + f"ALTER TABLE {schema}.session_peers VALIDATE CONSTRAINT fk_session_peers_workspace_name" + ) + ) + + +def downgrade() -> None: + conn = op.get_bind() + inspector = sa.inspect(conn) + + if fk_exists("session_peers", "fk_session_peers_workspace_name", inspector): + op.drop_constraint( + "fk_session_peers_workspace_name", + table_name="session_peers", + type_="foreignkey", + schema=schema, + ) + + if index_exists("documents", "ix_documents_workspace_name", inspector): + op.drop_index( + "ix_documents_workspace_name", table_name="documents", schema=schema + ) + + if index_exists("peers", "ix_peers_workspace_name", inspector): + op.drop_index("ix_peers_workspace_name", table_name="peers", schema=schema) + if index_exists("collections", "ix_collections_workspace_name", inspector): + op.drop_index( + "ix_collections_workspace_name", table_name="collections", schema=schema + ) + + # First, drop the FK constraint (we'll recreate it later if needed) + if fk_exists("queue", "fk_queue_session_id"): + op.drop_constraint( + "fk_queue_session_id", + "queue", + type_="foreignkey", + schema=schema, + ) + + if not index_exists( + "active_queue_sessions", + f"ix_{schema}_active_queue_sessions_work_unit_key", + inspector, + ): + op.create_index( + f"ix_{schema}_active_queue_sessions_work_unit_key", + table_name="active_queue_sessions", + columns=["work_unit_key"], + schema=schema, + ) + + # Recreate the redundant unique constraints + if not constraint_exists("sessions", "uq_sessions_public_id", "unique", inspector): + op.create_unique_constraint( + "uq_sessions_public_id", "sessions", ["id"], schema=schema + ) + if not constraint_exists("peers", "uq_users_public_id", "unique", inspector): + op.create_unique_constraint( + "uq_users_public_id", "peers", ["id"], schema=schema + ) + if not constraint_exists("workspaces", "uq_apps_public_id", "unique", inspector): + op.create_unique_constraint( + "uq_apps_public_id", "workspaces", ["id"], schema=schema + ) + + if not constraint_exists( + "collections", "uq_collections_public_id", "unique", inspector + ): + op.create_unique_constraint( + "uq_collections_public_id", "collections", ["id"], schema=schema + ) + if not constraint_exists( + "documents", "uq_documents_public_id", "unique", inspector + ): + op.create_unique_constraint( + "uq_documents_public_id", "documents", ["id"], schema=schema + ) + + # Recreate the redundant indexes + if not index_exists("sessions", "ix_sessions_public_id", inspector): + op.create_index("ix_sessions_public_id", "sessions", ["id"], schema=schema) + if not index_exists("peers", "ix_users_public_id", inspector): + op.create_index("ix_users_public_id", "peers", ["id"], schema=schema) + if not index_exists("workspaces", "ix_apps_public_id", inspector): + op.create_index("ix_apps_public_id", "workspaces", ["id"], schema=schema) + if not index_exists("documents", "ix_documents_public_id", inspector): + op.create_index("ix_documents_public_id", "documents", ["id"], schema=schema) + if not index_exists("collections", "ix_collections_public_id", inspector): + op.create_index( + "ix_collections_public_id", + "collections", + ["id"], + schema=schema, + ) + + # Rename indexes on peers table back to original names + index_renames = [ + ("peers", "ix_peers_created_at", "ix_users_created_at"), + ("peers", "ix_peers_name", "ix_users_name"), + ("workspaces", "ix_workspaces_created_at", "ix_apps_created_at"), + ("workspaces", "ix_workspaces_name", "ix_apps_name"), + ] + for table_name, new_name, old_name in index_renames: + if index_exists(table_name, new_name, inspector): + conn.execute( + sa.text(f"ALTER INDEX {schema}.{new_name} RENAME TO {old_name}") + ) + + # Drop primary key constraint from message_embeddings.id + if constraint_exists("message_embeddings", "pk_message_embeddings", "primary"): + op.drop_constraint( + "pk_message_embeddings", "message_embeddings", "primary", schema=schema + ) + + # Make documents.embedding nullable + if column_exists("documents", "embedding"): + op.alter_column( + "documents", + "embedding", + nullable=True, + schema=schema, + ) + + # Make active_queue_sessions.work_unit_key nullable + if column_exists("active_queue_sessions", "work_unit_key"): + op.alter_column( + "active_queue_sessions", + "work_unit_key", + nullable=True, + schema=schema, + ) + + # Make sessions.workspace_name nullable + if column_exists("sessions", "workspace_name"): + op.alter_column( + "sessions", + "workspace_name", + nullable=True, + schema=schema, + ) + + # Make peers.workspace_name nullable + if column_exists("peers", "workspace_name"): + op.alter_column( + "peers", + "workspace_name", + nullable=True, + schema=schema, + ) diff --git a/migrations/versions/08894082221a_replace_collection_name_with_observer_.py b/migrations/versions/08894082221a_replace_collection_name_with_observer_.py index 481ae33f..e38111b0 100644 --- a/migrations/versions/08894082221a_replace_collection_name_with_observer_.py +++ b/migrations/versions/08894082221a_replace_collection_name_with_observer_.py @@ -53,7 +53,7 @@ def upgrade() -> None: connection.execute( text( f""" - INSERT INTO {schema}.sessions (id, name, workspace_name, is_active) VALUES (:session_id, '__global_observations__', :workspace_name, true) ON CONFLICT DO NOTHING + INSERT INTO {schema}.sessions (id, name, workspace_name, is_active, metadata, internal_metadata, configuration, created_at) VALUES (:session_id, '__global_observations__', :workspace_name, true, '{{}}', '{{}}', '{{}}', NOW()) ON CONFLICT DO NOTHING """ ), {"session_id": session_id, "workspace_name": workspace_name}, @@ -446,7 +446,10 @@ def upgrade() -> None: schema=schema, ) - # Step 17: Drop the name column from collections + # Step 17: Drop the name_length check constraint before dropping the name column from collections + if constraint_exists("collections", "name_length", "check", inspector): + op.drop_constraint("name_length", "collections", schema=schema) + if column_exists("collections", "name", inspector): op.drop_column("collections", "name", schema=schema) diff --git a/migrations/versions/20f89a421aff_rename_metamessage_type_to_label.py b/migrations/versions/20f89a421aff_rename_metamessage_type_to_label.py index fb04c186..114eefa2 100644 --- a/migrations/versions/20f89a421aff_rename_metamessage_type_to_label.py +++ b/migrations/versions/20f89a421aff_rename_metamessage_type_to_label.py @@ -11,6 +11,8 @@ from collections.abc import Sequence import sqlalchemy as sa from alembic import op +from migrations.utils import constraint_exists + # revision identifiers, used by Alembic. revision: str = "20f89a421aff" down_revision: str | None = "556a16564f50" @@ -61,7 +63,8 @@ def upgrade() -> None: ) # Rename check constraint - op.execute("ALTER TABLE metamessages DROP CONSTRAINT metamessage_type_length;") + if constraint_exists("metamessages", "metamessage_type_length", "check"): + op.execute("ALTER TABLE metamessages DROP CONSTRAINT metamessage_type_length;") op.create_check_constraint("label_length", "metamessages", "length(label) <= 512") # ### end Alembic commands ### diff --git a/migrations/versions/e9b705f9adf9_add_server_defaults_to_timestamp_.py b/migrations/versions/e9b705f9adf9_add_server_defaults_to_timestamp_.py new file mode 100644 index 00000000..b23ba032 --- /dev/null +++ b/migrations/versions/e9b705f9adf9_add_server_defaults_to_timestamp_.py @@ -0,0 +1,387 @@ +"""add server defaults to timestamp boolean and jsonb columns + +Revision ID: e9b705f9adf9 +Revises: 066e87ca5b07 +Create Date: 2025-10-29 12:08:36.803611 + +""" + +from collections.abc import Sequence + +import sqlalchemy as sa +from alembic import op + +from migrations.utils import get_schema + +# revision identifiers, used by Alembic. +revision: str = "e9b705f9adf9" +down_revision: str | None = "066e87ca5b07" +branch_labels: str | Sequence[str] | None = None +depends_on: str | Sequence[str] | None = None + + +schema = get_schema() + + +def upgrade() -> None: + # Add server defaults for timestamp columns + op.alter_column( + "workspaces", + "created_at", + server_default=sa.func.now(), + schema=schema, + ) + op.alter_column( + "peers", + "created_at", + server_default=sa.func.now(), + schema=schema, + ) + op.alter_column( + "sessions", + "created_at", + server_default=sa.func.now(), + schema=schema, + ) + op.alter_column( + "messages", + "created_at", + server_default=sa.func.now(), + schema=schema, + ) + op.alter_column( + "message_embeddings", + "created_at", + server_default=sa.func.now(), + schema=schema, + ) + op.alter_column( + "collections", + "created_at", + server_default=sa.func.now(), + schema=schema, + ) + op.alter_column( + "documents", + "created_at", + server_default=sa.func.now(), + schema=schema, + ) + op.alter_column( + "queue", + "created_at", + server_default=sa.func.now(), + schema=schema, + ) + op.alter_column( + "webhook_endpoints", + "created_at", + server_default=sa.func.now(), + schema=schema, + ) + op.alter_column( + "session_peers", + "joined_at", + server_default=sa.func.now(), + schema=schema, + ) + op.alter_column( + "active_queue_sessions", + "last_updated", + server_default=sa.func.now(), + schema=schema, + ) + + # Add server defaults for JSONB columns + op.alter_column( + "workspaces", + "metadata", + server_default=sa.text("'{}'::jsonb"), + schema=schema, + ) + op.alter_column( + "workspaces", + "internal_metadata", + server_default=sa.text("'{}'::jsonb"), + schema=schema, + ) + op.alter_column( + "workspaces", + "configuration", + server_default=sa.text("'{}'::jsonb"), + schema=schema, + ) + op.alter_column( + "peers", + "metadata", + server_default=sa.text("'{}'::jsonb"), + schema=schema, + ) + op.alter_column( + "peers", + "internal_metadata", + server_default=sa.text("'{}'::jsonb"), + schema=schema, + ) + op.alter_column( + "peers", + "configuration", + server_default=sa.text("'{}'::jsonb"), + schema=schema, + ) + op.alter_column( + "sessions", + "metadata", + server_default=sa.text("'{}'::jsonb"), + schema=schema, + ) + op.alter_column( + "sessions", + "internal_metadata", + server_default=sa.text("'{}'::jsonb"), + schema=schema, + ) + op.alter_column( + "sessions", + "configuration", + server_default=sa.text("'{}'::jsonb"), + schema=schema, + ) + op.alter_column( + "messages", + "metadata", + server_default=sa.text("'{}'::jsonb"), + schema=schema, + ) + op.alter_column( + "messages", + "internal_metadata", + server_default=sa.text("'{}'::jsonb"), + schema=schema, + ) + op.alter_column( + "collections", + "metadata", + server_default=sa.text("'{}'::jsonb"), + schema=schema, + ) + op.alter_column( + "collections", + "internal_metadata", + server_default=sa.text("'{}'::jsonb"), + schema=schema, + ) + op.alter_column( + "documents", + "internal_metadata", + server_default=sa.text("'{}'::jsonb"), + schema=schema, + ) + op.alter_column( + "session_peers", + "configuration", + server_default=sa.text("'{}'::jsonb"), + schema=schema, + ) + op.alter_column( + "session_peers", + "internal_metadata", + server_default=sa.text("'{}'::jsonb"), + schema=schema, + ) + + # Add server defaults for boolean columns + op.alter_column( + "sessions", + "is_active", + server_default=sa.text("true"), + schema=schema, + ) + op.alter_column( + "queue", + "processed", + server_default=sa.text("false"), + schema=schema, + ) + + +def downgrade() -> None: + # Remove server defaults for timestamp columns + op.alter_column( + "workspaces", + "created_at", + server_default=None, + schema=schema, + ) + op.alter_column( + "peers", + "created_at", + server_default=None, + schema=schema, + ) + op.alter_column( + "sessions", + "created_at", + server_default=None, + schema=schema, + ) + op.alter_column( + "messages", + "created_at", + server_default=None, + schema=schema, + ) + op.alter_column( + "message_embeddings", + "created_at", + server_default=None, + schema=schema, + ) + op.alter_column( + "collections", + "created_at", + server_default=None, + schema=schema, + ) + op.alter_column( + "documents", + "created_at", + server_default=None, + schema=schema, + ) + op.alter_column( + "queue", + "created_at", + server_default=None, + schema=schema, + ) + op.alter_column( + "webhook_endpoints", + "created_at", + server_default=None, + schema=schema, + ) + op.alter_column( + "session_peers", + "joined_at", + server_default=None, + schema=schema, + ) + op.alter_column( + "active_queue_sessions", + "last_updated", + server_default=None, + schema=schema, + ) + + # Remove server defaults for JSONB columns + op.alter_column( + "workspaces", + "metadata", + server_default=None, + schema=schema, + ) + op.alter_column( + "workspaces", + "internal_metadata", + server_default=None, + schema=schema, + ) + op.alter_column( + "workspaces", + "configuration", + server_default=None, + schema=schema, + ) + op.alter_column( + "peers", + "metadata", + server_default=None, + schema=schema, + ) + op.alter_column( + "peers", + "internal_metadata", + server_default=None, + schema=schema, + ) + op.alter_column( + "peers", + "configuration", + server_default=None, + schema=schema, + ) + op.alter_column( + "sessions", + "metadata", + server_default=None, + schema=schema, + ) + op.alter_column( + "sessions", + "internal_metadata", + server_default=None, + schema=schema, + ) + op.alter_column( + "sessions", + "configuration", + server_default=None, + schema=schema, + ) + op.alter_column( + "messages", + "metadata", + server_default=None, + schema=schema, + ) + op.alter_column( + "messages", + "internal_metadata", + server_default=None, + schema=schema, + ) + op.alter_column( + "collections", + "metadata", + server_default=None, + schema=schema, + ) + op.alter_column( + "collections", + "internal_metadata", + server_default=None, + schema=schema, + ) + op.alter_column( + "documents", + "internal_metadata", + server_default=None, + schema=schema, + ) + op.alter_column( + "session_peers", + "configuration", + server_default=None, + schema=schema, + ) + op.alter_column( + "session_peers", + "internal_metadata", + server_default=None, + schema=schema, + ) + + # Remove server defaults for boolean columns + op.alter_column( + "sessions", + "is_active", + server_default=None, + schema=schema, + ) + op.alter_column( + "queue", + "processed", + server_default=None, + schema=schema, + ) diff --git a/src/db.py b/src/db.py index 9d37d321..100663d2 100644 --- a/src/db.py +++ b/src/db.py @@ -46,6 +46,8 @@ SessionLocal = async_sessionmaker( ) table_schema = settings.DB.SCHEMA +# Note: column_0_N_name expands to include all columns in multi-column constraints +# e.g., "workspace_id_tenant_id" for a composite constraint on both columns meta = MetaData() meta.schema = table_schema Base = declarative_base(metadata=meta) diff --git a/src/models.py b/src/models.py index b66e31cf..2efad02f 100644 --- a/src/models.py +++ b/src/models.py @@ -51,13 +51,25 @@ session_peers_table = Table( nullable=False, ), Column("peer_name", TEXT, primary_key=True, nullable=False), - Column("configuration", JSONB, default=dict), - Column("internal_metadata", JSONB, default=dict), + Column( + "configuration", + JSONB, + default=dict, + nullable=False, + server_default=text("'{}'::jsonb"), + ), + Column( + "internal_metadata", + JSONB, + default=dict, + nullable=False, + server_default=text("'{}'::jsonb"), + ), Column( "joined_at", DateTime(timezone=True), nullable=False, - default=func.now(), + server_default=func.now(), ), Column( "left_at", @@ -81,22 +93,28 @@ session_peers_table = Table( class Workspace(Base): __tablename__: str = "workspaces" id: Mapped[str] = mapped_column(TEXT, default=generate_nanoid, primary_key=True) - name: Mapped[str] = mapped_column(TEXT, index=True, unique=True) + name: Mapped[str] = mapped_column(TEXT, unique=True) peers = relationship("Peer", back_populates="workspace") webhook_endpoints = relationship("WebhookEndpoint", back_populates="workspace") created_at: Mapped[datetime.datetime] = mapped_column( - DateTime(timezone=True), index=True, default=func.now() + DateTime(timezone=True), server_default=func.now() + ) + h_metadata: Mapped[dict[str, Any]] = mapped_column( + "metadata", JSONB, default=dict, server_default=text("'{}'::jsonb") ) - h_metadata: Mapped[dict[str, Any]] = mapped_column("metadata", JSONB, default=dict) internal_metadata: Mapped[dict[str, Any]] = mapped_column( "internal_metadata", JSONB, default=dict ) - configuration: Mapped[dict[str, Any]] = mapped_column(JSONB, default=dict) + configuration: Mapped[dict[str, Any]] = mapped_column( + JSONB, default=dict, server_default=text("'{}'::jsonb") + ) __table_args__ = ( CheckConstraint("length(id) = 21", name="id_length"), CheckConstraint("length(name) <= 512", name="name_length"), CheckConstraint("id ~ '^[A-Za-z0-9_-]+$'", name="id_format"), + Index("ix_workspaces_created_at", "created_at"), + Index("ix_workspaces_name", "name"), ) @@ -104,18 +122,22 @@ class Workspace(Base): class Peer(Base): __tablename__: str = "peers" id: Mapped[str] = mapped_column(TEXT, default=generate_nanoid, primary_key=True) - name: Mapped[str] = mapped_column(TEXT, index=True) - h_metadata: Mapped[dict[str, Any]] = mapped_column("metadata", JSONB, default=dict) + name: Mapped[str] = mapped_column(TEXT, nullable=False) + h_metadata: Mapped[dict[str, Any]] = mapped_column( + "metadata", JSONB, default=dict, server_default=text("'{}'::jsonb") + ) internal_metadata: Mapped[dict[str, Any]] = mapped_column( - "internal_metadata", JSONB, default=dict + "internal_metadata", JSONB, default=dict, server_default=text("'{}'::jsonb") ) created_at: Mapped[datetime.datetime] = mapped_column( - DateTime(timezone=True), index=True, default=func.now() + DateTime(timezone=True), server_default=func.now() ) workspace_name: Mapped[str] = mapped_column( - ForeignKey("workspaces.name"), index=True, nullable=False + ForeignKey("workspaces.name"), nullable=False + ) + configuration: Mapped[dict[str, Any]] = mapped_column( + JSONB, default=dict, server_default=text("'{}'::jsonb") ) - configuration: Mapped[dict[str, Any]] = mapped_column(JSONB, default=dict) workspace = relationship("Workspace", back_populates="peers") sessions = relationship( @@ -128,6 +150,9 @@ class Peer(Base): CheckConstraint("length(name) <= 512", name="name_length"), CheckConstraint("id ~ '^[A-Za-z0-9_-]+$'", name="id_format"), Index("idx_peers_workspace_lookup", "workspace_name", "name"), + Index("ix_peers_created_at", "created_at"), + Index("ix_peers_name", "name"), + Index("ix_peers_workspace_name", "workspace_name"), ) def __repr__(self) -> str: @@ -138,20 +163,24 @@ class Peer(Base): class Session(Base): __tablename__: str = "sessions" id: Mapped[str] = mapped_column(TEXT, primary_key=True, default=generate_nanoid) - name: Mapped[str] = mapped_column(TEXT, index=True) - is_active: Mapped[bool] = mapped_column(default=True) - h_metadata: Mapped[dict[str, Any]] = mapped_column("metadata", JSONB, default=dict) + name: Mapped[str] = mapped_column(TEXT) + is_active: Mapped[bool] = mapped_column(default=True, server_default=text("true")) + h_metadata: Mapped[dict[str, Any]] = mapped_column( + "metadata", JSONB, default=dict, server_default=text("'{}'::jsonb") + ) internal_metadata: Mapped[dict[str, Any]] = mapped_column( - "internal_metadata", JSONB, default=dict + "internal_metadata", JSONB, default=dict, server_default=text("'{}'::jsonb") ) created_at: Mapped[datetime.datetime] = mapped_column( - DateTime(timezone=True), index=True, default=func.now() + DateTime(timezone=True), server_default=func.now() ) messages = relationship("Message", back_populates="session") workspace_name: Mapped[str] = mapped_column( - ForeignKey("workspaces.name"), index=True, nullable=False + ForeignKey("workspaces.name"), nullable=False + ) + configuration: Mapped[dict[str, Any]] = mapped_column( + JSONB, default=dict, server_default=text("'{}'::jsonb") ) - configuration: Mapped[dict[str, Any]] = mapped_column(JSONB, default=dict) peers = relationship( "Peer", secondary=session_peers_table, back_populates="sessions" @@ -162,6 +191,7 @@ class Session(Base): CheckConstraint("length(name) <= 512", name="name_length"), CheckConstraint("length(id) = 21", name="id_length"), CheckConstraint("id ~ '^[A-Za-z0-9_-]+$'", name="id_format"), + Index("ix_sessions_created_at", "created_at"), ) def __repr__(self) -> str: @@ -175,26 +205,30 @@ class Message(Base): BigInteger, Identity(), primary_key=True, autoincrement=True ) public_id: Mapped[str] = mapped_column( - TEXT, index=True, unique=True, default=generate_nanoid + TEXT, + unique=True, + default=generate_nanoid, ) # NOTE: Messages in Honcho 2.0 could historically be stored outside of a session. # We have since assigned all of these messages to a default session. - session_name: Mapped[str] = mapped_column(index=True, nullable=False) + session_name: Mapped[str] = mapped_column(TEXT, nullable=False) content: Mapped[str] = mapped_column(TEXT) - h_metadata: Mapped[dict[str, Any]] = mapped_column("metadata", JSONB, default=dict) + h_metadata: Mapped[dict[str, Any]] = mapped_column( + "metadata", JSONB, default=dict, server_default=text("'{}'::jsonb") + ) internal_metadata: Mapped[dict[str, Any]] = mapped_column( - "internal_metadata", JSONB, default=dict + "internal_metadata", JSONB, default=dict, server_default=text("'{}'::jsonb") ) token_count: Mapped[int] = mapped_column(Integer, default=0, nullable=False) - seq_in_session: Mapped[int] = mapped_column(BigInteger, nullable=False, index=True) + seq_in_session: Mapped[int] = mapped_column(BigInteger, nullable=False) created_at: Mapped[datetime.datetime] = mapped_column( - DateTime(timezone=True), index=True, default=func.now() + DateTime(timezone=True), server_default=func.now() ) session = relationship("Session", back_populates="messages") - peer_name: Mapped[str] = mapped_column(index=True) + peer_name: Mapped[str] = mapped_column(TEXT) workspace_name: Mapped[str] = mapped_column( - ForeignKey("workspaces.name"), index=True + ForeignKey("workspaces.name"), ) __table_args__ = ( @@ -229,6 +263,11 @@ class Message(Base): text("to_tsvector('english', content)"), postgresql_using="gin", ), + Index("ix_messages_created_at", "created_at"), + Index("ix_messages_id", "id"), + Index("ix_messages_peer_name", "peer_name"), + Index("ix_messages_public_id", "public_id"), + Index("ix_messages_workspace_name", "workspace_name"), ) @override @@ -246,15 +285,15 @@ class MessageEmbedding(Base): 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 + ForeignKey("messages.public_id"), ) workspace_name: Mapped[str] = mapped_column( - ForeignKey("workspaces.name"), index=True + ForeignKey("workspaces.name"), ) - session_name: Mapped[str] = mapped_column(TEXT, index=True, nullable=False) - peer_name: Mapped[str | None] = mapped_column(TEXT, index=True) + session_name: Mapped[str] = mapped_column(TEXT, nullable=False) + peer_name: Mapped[str] = mapped_column(TEXT) created_at: Mapped[datetime.datetime] = mapped_column( - DateTime(timezone=True), index=True, default=func.now() + DateTime(timezone=True), server_default=func.now() ) # Relationship to Message @@ -278,6 +317,11 @@ class MessageEmbedding(Base): postgresql_with={"m": 16, "ef_construction": 64}, postgresql_ops={"embedding": "vector_cosine_ops"}, ), + Index("idx_message_embeddings_created_at", "created_at"), + Index("idx_message_embeddings_message_id", "message_id"), + Index("idx_message_embeddings_peer_name", "peer_name"), + Index("idx_message_embeddings_session_name", "session_name"), + Index("idx_message_embeddings_workspace_name", "workspace_name"), ) @@ -286,20 +330,22 @@ class Collection(Base): __tablename__: str = "collections" id: Mapped[str] = mapped_column(TEXT, default=generate_nanoid, primary_key=True) - observer: Mapped[str] = mapped_column(TEXT, index=True) - observed: Mapped[str] = mapped_column(TEXT, index=True) + observer: Mapped[str] = mapped_column(TEXT) + observed: Mapped[str] = mapped_column(TEXT) created_at: Mapped[datetime.datetime] = mapped_column( - DateTime(timezone=True), index=True, default=func.now() + DateTime(timezone=True), server_default=func.now() + ) + h_metadata: Mapped[dict[str, Any]] = mapped_column( + "metadata", JSONB, default=dict, server_default=text("'{}'::jsonb") ) - h_metadata: Mapped[dict[str, Any]] = mapped_column("metadata", JSONB, default=dict) internal_metadata: Mapped[dict[str, Any]] = mapped_column( - "internal_metadata", JSONB, default=dict + "internal_metadata", JSONB, default=dict, server_default=text("'{}'::jsonb") ) documents = relationship( "Document", back_populates="collection", cascade="all, delete, delete-orphan" ) workspace_name: Mapped[str] = mapped_column( - ForeignKey("workspaces.name"), index=True + ForeignKey("workspaces.name"), ) __table_args__ = ( @@ -321,6 +367,10 @@ class Collection(Base): ["observed", "workspace_name"], ["peers.name", "peers.workspace_name"], ), + Index("idx_collections_observer", "observer"), + Index("idx_collections_observed", "observed"), + Index("ix_collections_created_at", "created_at"), + Index("ix_collections_workspace_name", "workspace_name"), ) @@ -329,20 +379,18 @@ class Document(Base): __tablename__: str = "documents" id: Mapped[str] = mapped_column(TEXT, default=generate_nanoid, primary_key=True) internal_metadata: Mapped[dict[str, Any]] = mapped_column( - "internal_metadata", JSONB, default=dict + "internal_metadata", JSONB, default=dict, server_default=text("'{}'::jsonb") ) content: Mapped[str] = mapped_column(TEXT) embedding: MappedColumn[Any] = mapped_column(Vector(1536)) created_at: Mapped[datetime.datetime] = mapped_column( - DateTime(timezone=True), index=True, default=func.now() + DateTime(timezone=True), server_default=func.now() ) - observer: Mapped[str] = mapped_column(TEXT, index=True) - observed: Mapped[str] = mapped_column(TEXT, index=True) - workspace_name: Mapped[str] = mapped_column( - ForeignKey("workspaces.name"), index=True - ) - session_name: Mapped[str] = mapped_column(TEXT, index=True) + observer: Mapped[str] = mapped_column(TEXT) + observed: Mapped[str] = mapped_column(TEXT) + workspace_name: Mapped[str] = mapped_column(ForeignKey("workspaces.name")) + session_name: Mapped[str] = mapped_column(TEXT) collection = relationship("Collection", back_populates="documents") __table_args__ = ( @@ -383,6 +431,11 @@ class Document(Base): "embedding": "vector_cosine_ops" }, # Cosine distance operator ), + Index("idx_documents_observer", "observer"), + Index("idx_documents_observed", "observed"), + Index("idx_documents_session_name", "session_name"), + Index("ix_documents_created_at", "created_at"), + Index("ix_documents_workspace_name", "workspace_name"), ) @@ -395,17 +448,22 @@ class QueueItem(Base): id: Mapped[int] = mapped_column( BigInteger, Identity(), primary_key=True, autoincrement=True ) - session_id: Mapped[str] = mapped_column( - ForeignKey("sessions.id"), index=True, nullable=True - ) + session_id: Mapped[str] = mapped_column(ForeignKey("sessions.id"), nullable=True) work_unit_key: Mapped[str] = mapped_column(TEXT, nullable=False) task_type: Mapped[TaskType] = mapped_column(TEXT, nullable=False) payload: Mapped[dict[str, Any]] = mapped_column(JSONB, nullable=False) - processed: Mapped[bool] = mapped_column(Boolean, default=False) + processed: Mapped[bool] = mapped_column( + Boolean, default=False, server_default=text("false") + ) error: Mapped[str | None] = mapped_column(TEXT, nullable=True) created_at: Mapped[datetime.datetime] = mapped_column( - DateTime(timezone=True), index=True, default=func.now() + DateTime(timezone=True), server_default=func.now() + ) + + __table_args__ = ( + Index("ix_queue_created_at", "created_at"), + Index("ix_queue_session_id", "session_id"), ) def __repr__(self) -> str: @@ -418,10 +476,10 @@ class ActiveQueueSession(Base): id: Mapped[str] = mapped_column(TEXT, default=generate_nanoid, primary_key=True) - work_unit_key: Mapped[str] = mapped_column(TEXT, unique=True, index=True) + work_unit_key: Mapped[str] = mapped_column(TEXT, unique=True) last_updated: Mapped[datetime.datetime] = mapped_column( - DateTime(timezone=True), default=func.now(), onupdate=func.now() + DateTime(timezone=True), server_default=func.now(), onupdate=func.now() ) @@ -430,11 +488,11 @@ class WebhookEndpoint(Base): __tablename__: str = "webhook_endpoints" id: Mapped[str] = mapped_column(TEXT, default=generate_nanoid, primary_key=True) workspace_name: Mapped[str] = mapped_column( - ForeignKey("workspaces.name"), index=True, nullable=False + ForeignKey("workspaces.name"), nullable=False ) url: Mapped[str] = mapped_column(TEXT, nullable=False) created_at: Mapped[datetime.datetime] = mapped_column( - DateTime(timezone=True), default=func.now() + DateTime(timezone=True), server_default=func.now() ) workspace = relationship("Workspace", back_populates="webhook_endpoints") diff --git a/tests/alembic/revisions/__init__.py b/tests/alembic/revisions/__init__.py index 359e7b1f..988130db 100644 --- a/tests/alembic/revisions/__init__.py +++ b/tests/alembic/revisions/__init__.py @@ -2,6 +2,7 @@ from . import ( test_05486ce795d5_make_session_name_required_on_messages, + test_066e87ca5b07_align_schema_with_declarative_models, test_08894082221a_replace_collection_name_with_observer_, test_20f89a421aff_rename_metamessage_type_to_label, test_66e63cf2cf77_add_indexes_to_documents_table, @@ -19,6 +20,7 @@ from . import ( __all__ = [ "test_05486ce795d5_make_session_name_required_on_messages", + "test_066e87ca5b07_align_schema_with_declarative_models", "test_08894082221a_replace_collection_name_with_observer_", "test_20f89a421aff_rename_metamessage_type_to_label", "test_556a16564f50_add_user_id_and_app_id_to_tables", diff --git a/tests/alembic/revisions/test_066e87ca5b07_align_schema_with_declarative_models.py b/tests/alembic/revisions/test_066e87ca5b07_align_schema_with_declarative_models.py new file mode 100644 index 00000000..40d396bc --- /dev/null +++ b/tests/alembic/revisions/test_066e87ca5b07_align_schema_with_declarative_models.py @@ -0,0 +1,42 @@ +"""Hooks for revision 066e87ca5b07 (align_schema_with_declarative_models).""" + +from __future__ import annotations + +from tests.alembic.registry import register_after_upgrade, register_before_upgrade +from tests.alembic.verifier import MigrationVerifier + + +@register_before_upgrade("066e87ca5b07") +def prepare_align_schema_with_declarative_models(verifier: MigrationVerifier) -> None: + """Seed state and assertions before upgrading to 066e87ca5b07.""" + # Assert columns exist but are nullable before migration + verifier.assert_column_exists( + "active_queue_sessions", "work_unit_key", exists=True, nullable=True + ) + verifier.assert_column_exists("documents", "embedding", exists=True, nullable=True) + + # Assert FK constraint does not exist yet + verifier.assert_constraint_exists( + "queue", "fk_queue_session_id", "foreign_key", exists=False + ) + + +@register_after_upgrade("066e87ca5b07") +def verify_align_schema_with_declarative_models(verifier: MigrationVerifier) -> None: + """Add assertions validating the effects of 066e87ca5b07.""" + # Assert columns are now non-nullable + verifier.assert_column_exists( + "peers", "workspace_name", exists=True, nullable=False + ) + verifier.assert_column_exists( + "sessions", "workspace_name", exists=True, nullable=False + ) + verifier.assert_column_exists( + "active_queue_sessions", "work_unit_key", exists=True, nullable=False + ) + verifier.assert_column_exists("documents", "embedding", exists=True, nullable=False) + + # Assert FK constraint now exists + verifier.assert_constraint_exists( + "queue", "fk_queue_session_id", "foreign_key", exists=True + ) diff --git a/tests/alembic/revisions/test_e9b705f9adf9_add_server_defaults_to_timestamp_.py b/tests/alembic/revisions/test_e9b705f9adf9_add_server_defaults_to_timestamp_.py new file mode 100644 index 00000000..31d9c6c4 --- /dev/null +++ b/tests/alembic/revisions/test_e9b705f9adf9_add_server_defaults_to_timestamp_.py @@ -0,0 +1,284 @@ +"""Hooks for revision e9b705f9adf9 (add server defaults to timestamp, boolean, and jsonb columns).""" + +from __future__ import annotations + +from nanoid import generate as generate_nanoid +from sqlalchemy import text + +from tests.alembic.registry import register_after_upgrade, register_before_upgrade +from tests.alembic.verifier import MigrationVerifier + +# Test data IDs +WORKSPACE_ID = generate_nanoid() +PEER_ID = generate_nanoid() +SESSION_ID = generate_nanoid() +MESSAGE_ID = generate_nanoid() +COLLECTION_ID = generate_nanoid() +DOCUMENT_ID = generate_nanoid() + + +@register_before_upgrade("e9b705f9adf9") +def prepare_add_server_defaults(verifier: MigrationVerifier) -> None: + """Seed state before upgrading to e9b705f9adf9. + + This migration adds server defaults to timestamp, JSONB, and boolean columns. + We verify that columns exist but don't have server defaults before the migration. + """ + conn = verifier.conn + schema = verifier.schema + inspector = verifier.get_inspector() + + # Sample timestamp columns to check - they should exist but without server defaults + for table, column in [ + ("workspaces", "created_at"), + ("peers", "created_at"), + ("sessions", "created_at"), + ("messages", "created_at"), + ("collections", "created_at"), + ("documents", "created_at"), + ("queue", "created_at"), + ]: + columns = inspector.get_columns(table, schema=schema) + col_info = next((c for c in columns if c["name"] == column), None) + assert ( + col_info is not None + ), f"Column {table}.{column} should exist before migration" + + # Create test data to ensure existing rows work after migration + conn.execute( + text( + f'INSERT INTO "{schema}"."workspaces" ' + + '("id", "name", "created_at", "metadata", "internal_metadata", "configuration") ' + + "VALUES (:id, :name, NOW(), :metadata, :internal_metadata, :configuration)" + ), + { + "id": WORKSPACE_ID, + "name": "test-workspace", + "metadata": "{}", + "internal_metadata": "{}", + "configuration": "{}", + }, + ) + + conn.execute( + text( + f'INSERT INTO "{schema}"."peers" ' + + '("id", "name", "workspace_name", "created_at", "metadata", "internal_metadata", "configuration") ' + + "VALUES (:id, :name, :workspace_name, NOW(), :metadata, :internal_metadata, :configuration)" + ), + { + "id": PEER_ID, + "name": "test-peer", + "workspace_name": "test-workspace", + "metadata": "{}", + "internal_metadata": "{}", + "configuration": "{}", + }, + ) + + conn.execute( + text( + f'INSERT INTO "{schema}"."sessions" ' + + '("id", "name", "workspace_name", "created_at", "is_active", "metadata", "internal_metadata", "configuration") ' + + "VALUES (:id, :name, :workspace_name, NOW(), true, :metadata, :internal_metadata, :configuration)" + ), + { + "id": SESSION_ID, + "name": "test-session", + "workspace_name": "test-workspace", + "metadata": "{}", + "internal_metadata": "{}", + "configuration": "{}", + }, + ) + + +@register_after_upgrade("e9b705f9adf9") +def verify_add_server_defaults(verifier: MigrationVerifier) -> None: + """Validate server defaults were added correctly to all columns.""" + conn = verifier.conn + schema = verifier.schema + inspector = verifier.get_inspector() + + # Verify timestamp columns have server defaults (now() function) + timestamp_columns = [ + ("workspaces", "created_at"), + ("peers", "created_at"), + ("sessions", "created_at"), + ("messages", "created_at"), + ("message_embeddings", "created_at"), + ("collections", "created_at"), + ("documents", "created_at"), + ("queue", "created_at"), + ("webhook_endpoints", "created_at"), + ("session_peers", "joined_at"), + ("active_queue_sessions", "last_updated"), + ] + + for table, column in timestamp_columns: + columns = inspector.get_columns(table, schema=schema) + col_info = next((c for c in columns if c["name"] == column), None) + assert ( + col_info is not None + ), f"Column {table}.{column} not found after migration" + + # Check that a server default exists + default = col_info.get("default") + assert default is not None, ( + f"Column {table}.{column} should have a server default after migration, " + f"but default is None" + ) + + # Verify JSONB columns have server defaults (empty object '{}') + jsonb_columns = [ + ("workspaces", "metadata"), + ("workspaces", "internal_metadata"), + ("workspaces", "configuration"), + ("peers", "metadata"), + ("peers", "internal_metadata"), + ("peers", "configuration"), + ("sessions", "metadata"), + ("sessions", "internal_metadata"), + ("sessions", "configuration"), + ("messages", "metadata"), + ("messages", "internal_metadata"), + ("collections", "metadata"), + ("collections", "internal_metadata"), + ("documents", "internal_metadata"), + ("session_peers", "configuration"), + ("session_peers", "internal_metadata"), + ] + + for table, column in jsonb_columns: + columns = inspector.get_columns(table, schema=schema) + col_info = next((c for c in columns if c["name"] == column), None) + assert ( + col_info is not None + ), f"Column {table}.{column} not found after migration" + + # Check that a server default exists + default = col_info.get("default") + assert default is not None, ( + f"Column {table}.{column} should have a server default after migration, " + f"but default is None" + ) + + # Verify boolean columns have server defaults + boolean_columns = [ + ("sessions", "is_active", "true"), + ("queue", "processed", "false"), + ] + + for table, column, expected_default in boolean_columns: + columns = inspector.get_columns(table, schema=schema) + col_info = next((c for c in columns if c["name"] == column), None) + assert ( + col_info is not None + ), f"Column {table}.{column} not found after migration" + + # Check that a server default exists + default = col_info.get("default") + assert default is not None, ( + f"Column {table}.{column} should have a server default after migration, " + f"but default is None" + ) + + assert ( + default == expected_default + ), f"Column {table}.{column} should have a server default of {expected_default} after migration, but default is {default}" + + # Test that defaults actually work by inserting rows without explicit values + test_workspace_id = generate_nanoid() + conn.execute( + text( + f'INSERT INTO "{schema}"."workspaces" ("id", "name") ' + + "VALUES (:id, :name)" + ), + {"id": test_workspace_id, "name": "test-defaults-workspace"}, + ) + + # Verify the inserted workspace has default values + workspace = conn.execute( + text( + 'SELECT "created_at", "metadata", "internal_metadata", "configuration" ' + + f'FROM "{schema}"."workspaces" WHERE "id" = :id' + ), + {"id": test_workspace_id}, + ).one() + + assert workspace.created_at is not None, "created_at should be auto-populated" + assert workspace.metadata == {}, "metadata should default to empty object" + assert ( + workspace.internal_metadata == {} + ), "internal_metadata should default to empty object" + assert workspace.configuration == {}, "configuration should default to empty object" + + # Test peer defaults + test_peer_id = generate_nanoid() + conn.execute( + text( + f'INSERT INTO "{schema}"."peers" ("id", "name", "workspace_name") ' + + "VALUES (:id, :name, :workspace_name)" + ), + { + "id": test_peer_id, + "name": "test-defaults-peer", + "workspace_name": "test-defaults-workspace", + }, + ) + + peer = conn.execute( + text( + 'SELECT "created_at", "metadata", "internal_metadata", "configuration" ' + + f'FROM "{schema}"."peers" WHERE "id" = :id' + ), + {"id": test_peer_id}, + ).one() + + assert peer.created_at is not None, "peer created_at should be auto-populated" + assert peer.metadata == {}, "peer metadata should default to empty object" + assert ( + peer.internal_metadata == {} + ), "peer internal_metadata should default to empty object" + assert peer.configuration == {}, "peer configuration should default to empty object" + + # Test session defaults (including boolean is_active) + test_session_id = generate_nanoid() + conn.execute( + text( + f'INSERT INTO "{schema}"."sessions" ("id", "name", "workspace_name") ' + + "VALUES (:id, :name, :workspace_name)" + ), + { + "id": test_session_id, + "name": "test-defaults-session", + "workspace_name": "test-defaults-workspace", + }, + ) + + session = conn.execute( + text( + 'SELECT "created_at", "is_active", "metadata", "internal_metadata", "configuration" ' + + f'FROM "{schema}"."sessions" WHERE "id" = :id' + ), + {"id": test_session_id}, + ).one() + + assert session.created_at is not None, "session created_at should be auto-populated" + assert session.is_active is True, "session is_active should default to true" + assert session.metadata == {}, "session metadata should default to empty object" + assert ( + session.internal_metadata == {} + ), "session internal_metadata should default to empty object" + assert ( + session.configuration == {} + ), "session configuration should default to empty object" + + # Verify pre-existing data still exists + existing_workspace = conn.execute( + text(f'SELECT "id" FROM "{schema}"."workspaces" WHERE "id" = :id'), + {"id": WORKSPACE_ID}, + ).one_or_none() + assert ( + existing_workspace is not None + ), "Pre-existing workspace should still exist after migration" From 1d0934a5685cb73111e910e58f07e25bb7bd88d2 Mon Sep 17 00:00:00 2001 From: Rajat Ahuja Date: Mon, 3 Nov 2025 15:33:48 -0500 Subject: [PATCH 16/17] fix: message seq in session N+1 (#261) * fix: message seq in session N+1 * test: behavior of enqueue * fix: test * chore: Code Rabbit Comments --------- Co-authored-by: Vineeth Voruganti <13438633+VVoruganti@users.noreply.github.com> --- src/routers/messages.py | 4 +- tests/integration/test_enqueue.py | 144 ++++++++++++++++++++++++++++++ 2 files changed, 146 insertions(+), 2 deletions(-) diff --git a/src/routers/messages.py b/src/routers/messages.py index 259fc9d7..5cbc9efe 100644 --- a/src/routers/messages.py +++ b/src/routers/messages.py @@ -71,7 +71,7 @@ async def create_messages_for_session( "peer_name": message.peer_name, "created_at": message.created_at, "message_public_id": message.public_id, - "message_seq_in_session": message.seq_in_session, + "seq_in_session": message.seq_in_session, } for message in created_messages ] @@ -135,7 +135,7 @@ async def create_messages_with_file( "peer_name": message.peer_name, "created_at": message.created_at, "message_public_id": message.public_id, - "message_seq_in_session": message.seq_in_session, + "seq_in_session": message.seq_in_session, } for message in created_messages ] diff --git a/tests/integration/test_enqueue.py b/tests/integration/test_enqueue.py index 5b1870e8..6e5fda30 100644 --- a/tests/integration/test_enqueue.py +++ b/tests/integration/test_enqueue.py @@ -9,6 +9,7 @@ from sqlalchemy.ext.asyncio import AsyncSession from src import crud, models, schemas from src.deriver import enqueue +from src.deriver.enqueue import generate_queue_records from src.models import Peer, QueueItem, Workspace @@ -1379,3 +1380,146 @@ class TestAdvancedEnqueueEdgeCases: assert len(actual_payloads) == len(expected_payloads) for expected in expected_payloads: assert expected in actual_payloads + + +@pytest.mark.asyncio +class TestGenerateQueueRecordsSeqInSession: + """Unit tests for generate_queue_records function focusing on seq_in_session handling""" + + async def test_generate_queue_records_uses_seq_from_payload_not_crud( + self, + db_session: AsyncSession, + sample_data: tuple[Workspace, Peer], + ): + """ + Test that generate_queue_records uses seq_in_session from payload + instead of making a CRUD call to get_message_seq_in_session. + """ + + 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 payload with seq_in_session included + message_payload = { + "message_id": 12345, + "peer_name": test_peer.name, + "workspace_name": test_workspace.name, + "session_name": test_session.name, + "content": "Test message", + "seq_in_session": 20, # Multiple of MESSAGES_PER_SHORT_SUMMARY to trigger summary creation + "created_at": datetime.now(timezone.utc), # Required by create_payload + } + + # Mock the CRUD function to track if it's called + # Also enable summary generation in settings + with ( + patch("src.deriver.enqueue.crud.get_message_seq_in_session") as mock_crud, + patch("src.deriver.enqueue.settings.SUMMARY.ENABLED", new=True), + ): + mock_crud.return_value = 200 + mock_db_session = AsyncMock() + + peers_config: dict[str, list[Any]] = { + test_peer.name: [ + {"observe_me": True}, + {"observe_others": True}, + ] + } + records = await generate_queue_records( + db_session=mock_db_session, + message=message_payload, + peers_with_configuration=peers_config, + session_id=test_session.id, + deriver_disabled=False, + ) + + mock_crud.assert_not_called() + + assert len(records) > 0 + + summary_records = [r for r in records if r["task_type"] == "summary"] + assert len(summary_records) > 0, "Expected summary records to be created" + for record in summary_records: + assert ( + record["payload"]["message_seq_in_session"] + != mock_crud.return_value + ) + assert record["payload"]["message_seq_in_session"] == 20 + + async def test_generate_queue_records_falls_back_to_crud_when_seq_missing( + self, + db_session: AsyncSession, + sample_data: tuple[Workspace, Peer], + ): + """ + Test that generate_queue_records falls back to CRUD call + when seq_in_session is missing from payload. + + This is the fallback behavior for backward compatibility. + """ + + 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 payload WITHOUT seq_in_session + message_payload = { + "message_id": 12345, + "peer_name": test_peer.name, + "workspace_name": test_workspace.name, + "session_name": test_session.name, + "content": "Test message", + "created_at": datetime.now(timezone.utc), + # seq_in_session is MISSING + } + + # Mock the CRUD function and enable summary generation in settings + with ( + patch("src.deriver.enqueue.crud.get_message_seq_in_session") as mock_crud, + patch("src.deriver.enqueue.settings.SUMMARY.ENABLED", True), + ): + mock_crud.return_value = ( + 60 # Multiple of MESSAGES_PER_LONG_SUMMARY to trigger summary creation + ) + + mock_db_session = AsyncMock() + + peers_config: dict[str, list[Any]] = { + test_peer.name: [ + {"observe_me": True}, + {"observe_others": True}, + ] + } + records = await generate_queue_records( + db_session=mock_db_session, + message=message_payload, + peers_with_configuration=peers_config, + session_id=test_session.id, + deriver_disabled=False, + ) + + # The CRUD function SHOULD have been called as fallback + mock_crud.assert_called_once_with( + mock_db_session, + workspace_name=test_workspace.name, + session_name=test_session.name, + message_id=12345, + ) + + # Verify that records were created with the fallback value + summary_records = [r for r in records if r["task_type"] == "summary"] + assert len(summary_records) > 0, "Expected summary records to be created" + for record in summary_records: + # Should use the value from CRUD fallback (60) + assert record["payload"]["message_seq_in_session"] == 60 From a25be3c13645a76d92140a27e6254c71326bc50c Mon Sep 17 00:00:00 2001 From: Vineeth Voruganti <13438633+VVoruganti@users.noreply.github.com> Date: Tue, 4 Nov 2025 10:41:58 -0500 Subject: [PATCH 17/17] chore: Readme & Changelog Updates (#263) * chore: Readme & Changelog Updates * chore: (typos) Clarify Changelog --- CHANGELOG.md | 13 +++++++++++++ README.md | 2 +- docs/changelog/compatibility-guide.mdx | 5 +++-- docs/changelog/introduction.mdx | 16 +++++++++++++++- docs/docs.json | 2 +- pyproject.toml | 11 ++++++++++- src/main.py | 2 +- uv.lock | 2 +- 8 files changed, 45 insertions(+), 8 deletions(-) diff --git a/CHANGELOG.md b/CHANGELOG.md index 1fdd2e90..68b7e316 100644 --- a/CHANGELOG.md +++ b/CHANGELOG.md @@ -5,6 +5,19 @@ All notable changes to this project will be documented in this file. The format is based on [Keep a Changelog](http://keepachangelog.com/) and this project adheres to [Semantic Versioning](http://semver.org/). +## [2.4.2] - 2025-11-03 + +### Fixed + +- Langfuse tracing to have readable waterfalls +- Alembic Migrations to match models.py +- message_in_seq correctly included in webhook payload + +### Changed + +- Alembic to always use a session pooler +- Statement timeout during alembic operations to 5 min + ## [2.4.1] - 2025-10-24 ### Added diff --git a/README.md b/README.md index 452f487b..baac9e39 100644 --- a/README.md +++ b/README.md @@ -8,7 +8,7 @@ --- -![Static Badge](https://img.shields.io/badge/Version-2.4.1-blue) +![Static Badge](https://img.shields.io/badge/Version-2.4.2-blue) [![PyPI version](https://img.shields.io/pypi/v/honcho-ai.svg)](https://pypi.org/project/honcho-ai/) [![NPM version](https://img.shields.io/npm/v/@honcho-ai/sdk.svg)](https://npmjs.org/package/@honcho-ai/sdk) [![Discord](https://img.shields.io/discord/1016845111637839922?style=flat&logo=discord&logoColor=23ffffff&label=Plastic%20Labs&labelColor=235865F2)](https://discord.gg/plasticlabs) diff --git a/docs/changelog/compatibility-guide.mdx b/docs/changelog/compatibility-guide.mdx index edcc8fdd..2b9f67a7 100644 --- a/docs/changelog/compatibility-guide.mdx +++ b/docs/changelog/compatibility-guide.mdx @@ -8,7 +8,7 @@ This guide helps you understand which versions of Honcho's API are compatible wi ## Version Compatibility -### Honcho API v2.4.1 (Current) +### Honcho API v2.4.2 (Current) @@ -35,7 +35,8 @@ This guide helps you understand which versions of Honcho's API are compatible wi | Honcho API Version | TypeScript SDK | Python SDK | |-------------------|---------------|------------| -| v2.4.1 (Current) | v1.5.0 | v1.5.0 | +| v2.4.2 (Current) | v1.5.0 | v1.5.0 | +| v2.4.1 | v1.5.0 | v1.5.0 | | v2.4.0 | v1.5.0 | v1.5.0 | | v2.3.3 | v1.4.1 | v1.4.1 | | v2.3.2 | v1.4.0 | v1.4.0 | diff --git a/docs/changelog/introduction.mdx b/docs/changelog/introduction.mdx index 4edab9dc..ff9946a5 100644 --- a/docs/changelog/introduction.mdx +++ b/docs/changelog/introduction.mdx @@ -27,7 +27,21 @@ Welcome to the Honcho changelog! This section documents all notable changes to t ### Honcho API and SDK Changelogs - + + ### Fixed + + - Langfuse tracing to have readable waterfalls + - Alembic Migrations to match models.py + - message_in_seq correctly included in webhook payload + + + ### Changed + + - Alembic to always use a session pooler + - Statement timeout during alembic operations to 5 min + + + ### Added - Alembic migration validation test suite diff --git a/docs/docs.json b/docs/docs.json index 4a9a6412..74f0f139 100644 --- a/docs/docs.json +++ b/docs/docs.json @@ -19,7 +19,7 @@ "navigation": { "versions": [ { - "version": "v2.4.1", + "version": "v2.4.2", "api": { "openapi": [ "openapi.documented.yml" diff --git a/pyproject.toml b/pyproject.toml index 6c9b037c..f39dcec5 100644 --- a/pyproject.toml +++ b/pyproject.toml @@ -1,6 +1,6 @@ [project] name = "honcho" -version = "2.4.1" +version = "2.4.2" description = "Honcho Server" authors = [ {name = "Plastic Labs", email = "hello@plasticlabs.ai"}, @@ -85,6 +85,15 @@ addopts = "--strict-markers --cov=src/ --cov=sdks/python/src/honcho --cov-report testpaths = ["tests"] pythonpath = ["src"] +[tool.coverage.report] +exclude_lines = [ + "pragma: no cover", + "def __repr__", + "raise AssertionError", + "raise NotImplementedError", + "if __name__ == .__main__.:", + "if TYPE_CHECKING:", +] [tool.basedpyright] # BasedPyright currently seems like the best type checker option, much faster diff --git a/src/main.py b/src/main.py index 412822d2..0ea88272 100644 --- a/src/main.py +++ b/src/main.py @@ -123,7 +123,7 @@ app = FastAPI( title="Honcho API", summary="The Identity Layer for the Agentic World", description="""Honcho is a platform for giving agents user-centric memory and social cognition""", - version="2.4.1", + version="2.4.2", contact={ "name": "Plastic Labs", "url": "https://honcho.dev", diff --git a/uv.lock b/uv.lock index dfe736ee..43347634 100644 --- a/uv.lock +++ b/uv.lock @@ -673,7 +673,7 @@ wheels = [ [[package]] name = "honcho" -version = "2.4.1" +version = "2.4.2" source = { virtual = "." } dependencies = [ { name = "alembic" },