Merge branch 'main' into vineeth/dev-1193

This commit is contained in:
Vineeth Voruganti 2025-11-05 17:00:00 -05:00
commit 7d2962434a
84 changed files with 5207 additions and 897 deletions

View File

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

View File

@ -5,6 +5,35 @@ 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
- 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

View File

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

View File

@ -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.2 (Current)
<CardGroup cols={2}>
<Card title="TypeScript SDK" icon="js">
@ -34,7 +34,9 @@ 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.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 |
| v2.3.1 | v1.4.0 | v1.4.0 |

View File

@ -27,7 +27,37 @@ Welcome to the Honcho changelog! This section documents all notable changes to t
### Honcho API and SDK Changelogs
<Tabs>
<Tab title="Honcho API">
<Update label="v2.4.0 (Current)">
<Update label="v2.4.2 (Current)">
### 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
</Update>
<Update label="v2.4.1">
### 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
</Update>
<Update label="v2.4.0">
### Added
- Unified `Representation` class
@ -335,7 +365,7 @@ Welcome to the Honcho changelog! This section documents all notable changes to t
<Tab title="Python SDK">
[Python SDK](https://pypi.org/project/honcho-ai/)
<Update label="v1.5.0 (Current)">
<Update label="v1.5.0">
### Added
- Delete workspace method
@ -409,7 +439,7 @@ Welcome to the Honcho changelog! This section documents all notable changes to t
<Tab title="TypeScript SDK">
[TypeScript SDK](https://www.npmjs.com/package/@honcho-ai/sdk)
<Update label="v1.5.0 (Current)">
<Update label="v1.5.0">
### Added
- Delete workspace method

View File

@ -14,9 +14,9 @@
"navigation": {
"versions": [
{
"version": "v2.4.0",
"version": "v2.4.2",
"api": {
"openapi": ["openapi.documented.json"]
"openapi": ["openapi.documented.yml"]
},
"tabs": [
{
@ -132,7 +132,6 @@
"v2/api-reference/endpoint/messages/create-messages-with-file"
]
},
{
"group": "webhooks",
"pages": [

View File

@ -1,16 +1,21 @@
import logging
import logging # noqa: I001
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
# 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)
@ -59,7 +64,7 @@ def run_migrations_offline() -> None:
script output.
"""
url = get_url()
url = ensure_session_pooler(get_url())
context.configure(
url=url,
@ -74,6 +79,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,14 +137,17 @@ 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,
connect_args={"prepare_threshold": None},
echo=True,
connect_args={
"prepare_threshold": None,
"options": "-c statement_timeout=300000", # 5 minutes in milliseconds
},
)
with connectable.connect() as connection:
@ -117,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():

View File

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

View File

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

View File

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

View File

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

View File

@ -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."""
@ -50,21 +53,32 @@ 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},
)
# 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
while True:
result = connection.execute(
text(
f"""
WITH batch AS (
SELECT id
FROM {schema}.documents
WHERE session_name IS NULL
LIMIT :batch_size
)
UPDATE {schema}.documents d
SET session_name = '__global_observations__'
FROM batch
WHERE d.id = batch.id
"""
),
{"batch_size": BATCH_SIZE},
)
if result.rowcount == 0:
break
op.alter_column("documents", "session_name", nullable=False, schema=schema)
@ -84,29 +98,119 @@ 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:
collection_ids = [row.id for row in collections_to_delete]
# Delete in smaller chunks of actual documents to avoid locking issues
while True:
result = connection.execute(
text(
f"""
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, "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:
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.
# 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(
f"""
WITH batch AS (
SELECT id
FROM {schema}.collections
WHERE observer IS NULL OR observed IS NULL
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
# 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)
@ -130,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
@ -147,17 +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
"""
),
{"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)
@ -224,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"],
@ -238,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"],
@ -252,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"],
@ -265,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"],
@ -277,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"],
@ -322,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)
@ -392,12 +519,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"],
@ -415,24 +540,36 @@ 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
while True:
result = connection.execute(
text(
f"""
WITH batch AS (
SELECT id
FROM {schema}.documents
WHERE collection_name IS NULL
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)
# 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",
@ -440,24 +577,38 @@ 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
while True:
result = connection.execute(
text(
f"""
WITH batch AS (
SELECT id
FROM {schema}.documents
WHERE peer_name IS NULL
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)
# 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"],
@ -478,26 +629,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,
@ -525,14 +678,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"],
@ -542,17 +695,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,

View File

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

View File

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

View File

@ -34,16 +34,50 @@ 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 = 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(
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' != ''
LIMIT :batch_size
)
UPDATE {schema}.documents d
SET session_name = d.internal_metadata->>'session_name'
FROM batch b
WHERE d.id = b.id
"""
),
{"batch_size": batch_size},
)
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):

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@ -1,6 +1,6 @@
[project]
name = "honcho"
version = "2.4.0"
version = "2.4.2"
description = "Honcho Server"
authors = [
{name = "Plastic Labs", email = "hello@plasticlabs.ai"},
@ -81,10 +81,19 @@ 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"]
[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

View File

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

View File

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

View File

@ -1,3 +1,4 @@
#!/usr/bin/env uv run python
"""
Script to generate embeddings for existing messages that don't already have embeddings.

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@ -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:
@ -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"],

View File

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

View File

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

View File

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

View File

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

View File

@ -7,18 +7,15 @@ 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 src import prometheus
from src.utils.logging import get_route_template
if TYPE_CHECKING:
from sentry_sdk._types import Event, Hint
from pydantic import ValidationError
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 +28,11 @@ 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
if TYPE_CHECKING:
from sentry_sdk._types import Event, Hint
def get_log_level() -> int:
@ -68,30 +70,31 @@ 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:
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",
@ -100,6 +103,7 @@ if SENTRY_ENABLED:
transaction_style="endpoint",
),
],
before_send=before_send,
)
@ -119,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.2",
contact={
"name": "Plastic Labs",
"url": "https://honcho.dev",

View File

@ -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
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
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,25 +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)
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__ = (
@ -216,12 +251,23 @@ 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",
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
@ -239,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
@ -271,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"),
)
@ -279,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__ = (
@ -314,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"),
)
@ -322,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__ = (
@ -376,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"),
)
@ -388,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:
@ -411,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()
)
@ -423,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")

View File

@ -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,
"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,
"seq_in_session": message.seq_in_session,
}
for message in created_messages
]

View File

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

46
src/sentry.py Normal file
View File

@ -0,0 +1,46 @@
"""Sentry initialization and configuration."""
from __future__ import annotations
import logging
from collections.abc import Sequence
from typing import TYPE_CHECKING
import sentry_sdk
from src.config import settings
if TYPE_CHECKING:
from sentry_sdk._types import EventProcessor
from sentry_sdk.integrations import Integration
logger = logging.getLogger(__name__)
# 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[Integration],
before_send: EventProcessor | None = None,
) -> None:
"""Initialize Sentry SDK with project settings.
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,
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=before_send,
integrations=integrations,
)

View File

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

View File

@ -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}'")

View File

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

View File

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

View File

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

29
tests/alembic/README.md Normal file
View File

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

View File

@ -0,0 +1,5 @@
"""Test package for Alembic migration scenarios."""
from . import revisions
__all__ = ["revisions"]

107
tests/alembic/conftest.py Normal file
View File

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

73
tests/alembic/registry.py Normal file
View File

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

View File

@ -0,0 +1,37 @@
"""Register revision-specific hooks for migration verification."""
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,
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_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",
"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",
]

View File

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

View File

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

View File

@ -0,0 +1,363 @@
"""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_<i>')
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_<i>_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",
"fk_documents_observer_observed_workspace_name_collections",
"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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

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

148
tests/alembic/scaffold.py Normal file
View File

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

View File

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

144
tests/alembic/verifier.py Normal file
View File

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

View File

@ -492,6 +492,11 @@ 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),
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

View File

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

View File

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

View File

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

View File

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

View File

@ -4,11 +4,12 @@ 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
from src.deriver import enqueue
from src.deriver.enqueue import generate_queue_records
from src.models import Peer, QueueItem, Workspace
@ -26,6 +27,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 +44,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 +1096,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 +1113,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}"},
)
@ -1359,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

View File

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

View File

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

View File

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

View File

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

View File

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

View File

@ -673,7 +673,7 @@ wheels = [
[[package]]
name = "honcho"
version = "2.4.0"
version = "2.4.2"
source = { virtual = "." }
dependencies = [
{ name = "alembic" },