Merge branch 'main' into vineeth/dev-1193
This commit is contained in:
commit
7d2962434a
|
|
@ -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)
|
||||
|
|
|
|||
29
CHANGELOG.md
29
CHANGELOG.md
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -8,7 +8,7 @@
|
|||
|
||||
---
|
||||
|
||||

|
||||

|
||||
[](https://pypi.org/project/honcho-ai/)
|
||||
[](https://npmjs.org/package/@honcho-ai/sdk)
|
||||
[](https://discord.gg/plasticlabs)
|
||||
|
|
|
|||
|
|
@ -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 |
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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": [
|
||||
|
|
|
|||
|
|
@ -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():
|
||||
|
|
|
|||
|
|
@ -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"}
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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 ###
|
||||
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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):
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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")),
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
@ -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.
|
||||
|
|
|
|||
|
|
@ -1,3 +1,4 @@
|
|||
#!/usr/bin/env uv run python
|
||||
"""
|
||||
Script to generate embeddings for existing messages that don't already have embeddings.
|
||||
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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"],
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
60
src/main.py
60
src/main.py
|
|
@ -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",
|
||||
|
|
|
|||
175
src/models.py
175
src/models.py
|
|
@ -51,13 +51,25 @@ session_peers_table = Table(
|
|||
nullable=False,
|
||||
),
|
||||
Column("peer_name", TEXT, primary_key=True, nullable=False),
|
||||
Column("configuration", JSONB, default=dict),
|
||||
Column("internal_metadata", JSONB, default=dict),
|
||||
Column(
|
||||
"configuration",
|
||||
JSONB,
|
||||
default=dict,
|
||||
nullable=False,
|
||||
server_default=text("'{}'::jsonb"),
|
||||
),
|
||||
Column(
|
||||
"internal_metadata",
|
||||
JSONB,
|
||||
default=dict,
|
||||
nullable=False,
|
||||
server_default=text("'{}'::jsonb"),
|
||||
),
|
||||
Column(
|
||||
"joined_at",
|
||||
DateTime(timezone=True),
|
||||
nullable=False,
|
||||
default=func.now(),
|
||||
server_default=func.now(),
|
||||
),
|
||||
Column(
|
||||
"left_at",
|
||||
|
|
@ -81,22 +93,28 @@ session_peers_table = Table(
|
|||
class Workspace(Base):
|
||||
__tablename__: str = "workspaces"
|
||||
id: Mapped[str] = mapped_column(TEXT, default=generate_nanoid, primary_key=True)
|
||||
name: Mapped[str] = mapped_column(TEXT, index=True, unique=True)
|
||||
name: Mapped[str] = mapped_column(TEXT, unique=True)
|
||||
peers = relationship("Peer", back_populates="workspace")
|
||||
webhook_endpoints = relationship("WebhookEndpoint", back_populates="workspace")
|
||||
created_at: Mapped[datetime.datetime] = mapped_column(
|
||||
DateTime(timezone=True), index=True, default=func.now()
|
||||
DateTime(timezone=True), server_default=func.now()
|
||||
)
|
||||
h_metadata: Mapped[dict[str, Any]] = mapped_column(
|
||||
"metadata", JSONB, default=dict, server_default=text("'{}'::jsonb")
|
||||
)
|
||||
h_metadata: Mapped[dict[str, Any]] = mapped_column("metadata", JSONB, default=dict)
|
||||
internal_metadata: Mapped[dict[str, Any]] = mapped_column(
|
||||
"internal_metadata", JSONB, default=dict
|
||||
)
|
||||
configuration: Mapped[dict[str, Any]] = mapped_column(JSONB, default=dict)
|
||||
configuration: Mapped[dict[str, Any]] = mapped_column(
|
||||
JSONB, default=dict, server_default=text("'{}'::jsonb")
|
||||
)
|
||||
|
||||
__table_args__ = (
|
||||
CheckConstraint("length(id) = 21", name="id_length"),
|
||||
CheckConstraint("length(name) <= 512", name="name_length"),
|
||||
CheckConstraint("id ~ '^[A-Za-z0-9_-]+$'", name="id_format"),
|
||||
Index("ix_workspaces_created_at", "created_at"),
|
||||
Index("ix_workspaces_name", "name"),
|
||||
)
|
||||
|
||||
|
||||
|
|
@ -104,18 +122,22 @@ class Workspace(Base):
|
|||
class Peer(Base):
|
||||
__tablename__: str = "peers"
|
||||
id: Mapped[str] = mapped_column(TEXT, default=generate_nanoid, primary_key=True)
|
||||
name: Mapped[str] = mapped_column(TEXT, index=True)
|
||||
h_metadata: Mapped[dict[str, Any]] = mapped_column("metadata", JSONB, default=dict)
|
||||
name: Mapped[str] = mapped_column(TEXT, nullable=False)
|
||||
h_metadata: Mapped[dict[str, Any]] = mapped_column(
|
||||
"metadata", JSONB, default=dict, server_default=text("'{}'::jsonb")
|
||||
)
|
||||
internal_metadata: Mapped[dict[str, Any]] = mapped_column(
|
||||
"internal_metadata", JSONB, default=dict
|
||||
"internal_metadata", JSONB, default=dict, server_default=text("'{}'::jsonb")
|
||||
)
|
||||
created_at: Mapped[datetime.datetime] = mapped_column(
|
||||
DateTime(timezone=True), index=True, default=func.now()
|
||||
DateTime(timezone=True), server_default=func.now()
|
||||
)
|
||||
workspace_name: Mapped[str] = mapped_column(
|
||||
ForeignKey("workspaces.name"), index=True
|
||||
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")
|
||||
|
|
|
|||
|
|
@ -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
|
||||
]
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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}'")
|
||||
|
|
@ -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]]]
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
|
|
|
|||
|
|
@ -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:
|
||||
|
|
|
|||
|
|
@ -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"`
|
||||
|
|
@ -0,0 +1,5 @@
|
|||
"""Test package for Alembic migration scenarios."""
|
||||
|
||||
from . import revisions
|
||||
|
||||
__all__ = ["revisions"]
|
||||
|
|
@ -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()
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
|
@ -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",
|
||||
]
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
)
|
||||
|
|
@ -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
|
||||
|
|
@ -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"
|
||||
|
|
@ -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
|
||||
|
|
@ -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
|
||||
|
|
@ -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")])
|
||||
|
|
@ -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")
|
||||
|
|
@ -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
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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"
|
||||
|
|
@ -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."""
|
||||
|
|
@ -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)
|
||||
|
|
@ -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)
|
||||
|
|
@ -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"
|
||||
|
|
@ -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()
|
||||
|
|
@ -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)
|
||||
|
|
@ -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"]
|
||||
|
|
@ -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
|
||||
|
||||
|
|
|
|||
|
|
@ -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])
|
||||
|
||||
|
|
|
|||
|
|
@ -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,
|
||||
},
|
||||
]
|
||||
|
||||
|
|
|
|||
|
|
@ -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),
|
||||
)
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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()
|
||||
|
|
|
|||
|
|
@ -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())
|
||||
|
|
|
|||
|
|
@ -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]
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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"""
|
||||
|
|
|
|||
Loading…
Reference in New Issue