feat: add session_metadata JSON column to message table (#12255)

* feat: add session_metadata JSON column to message table with PostgreSQL indexes

* [autofix.ci] apply automated fixes

* fix: add EXPAND phase marker to session_metadata migration

* feat: add session_metadata field to Message schema and wire through to MessageTable

* add tests

* [autofix.ci] apply automated fixes

* [autofix.ci] apply automated fixes (attempt 2/3)

* update description

* address review comments

* [autofix.ci] apply automated fixes

* [autofix.ci] apply automated fixes (attempt 2/3)

---------

Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
Himavarsha
2026-03-25 11:00:59 -04:00
committed by GitHub
parent 098f3b4a7b
commit 677f16ba48
5 changed files with 407 additions and 5 deletions

View File

@ -0,0 +1,57 @@
"""Add session_metadata column to message table
Phase: EXPAND
Adds a flexible JSON column to store enterprise session context including
tenant_id, user_id, region, policies, retention profiles, and data flags.
This enables client-driven metadata injection for enterprise session management.
Revision ID: ef4b036b585d
Revises: 0e6138e7a0c2
Create Date: 2026-03-19 10:32:05.048791
"""
from typing import Sequence, Union
from alembic import op
import sqlalchemy as sa
revision: str = 'ef4b036b585d'
down_revision: Union[str, None] = '0e6138e7a0c2'
branch_labels: Union[str, Sequence[str], None] = None
depends_on: Union[str, Sequence[str], None] = None
def upgrade() -> None:
conn = op.get_bind()
with op.batch_alter_table('message', schema=None) as batch_op:
batch_op.add_column(sa.Column('session_metadata', sa.JSON(), nullable=True))
if conn.dialect.name == 'postgresql':
op.create_index(
'ix_message_session_metadata_tenant',
'message',
[sa.text("(session_metadata->>'tenant_id')")],
postgresql_using='btree',
if_not_exists=True
)
op.create_index(
'ix_message_session_metadata_user',
'message',
[sa.text("(session_metadata->>'user_id')")],
postgresql_using='btree',
if_not_exists=True
)
def downgrade() -> None:
conn = op.get_bind()
if conn.dialect.name == 'postgresql':
op.drop_index('ix_message_session_metadata_user', table_name='message', if_exists=True)
op.drop_index('ix_message_session_metadata_tenant', table_name='message', if_exists=True)
with op.batch_alter_table('message', schema=None) as batch_op:
batch_op.drop_column('session_metadata')

View File

@ -126,6 +126,7 @@ class MessageBase(SQLModel):
properties=properties,
category=message.category,
content_blocks=content_blocks,
session_metadata=getattr(message, "session_metadata", None),
)
@ -148,6 +149,13 @@ class MessageTable(MessageBase, table=True): # type: ignore[call-arg]
sa_column=Column(JSON),
)
# Enterprise session metadata - flexible JSON column for client-provided context
session_metadata: dict | None = Field(
default=None,
sa_column=Column(JSON),
description="Session context data (e.g., user roles, custom tags, or analytics data).",
)
@field_validator("flow_id", mode="before")
@classmethod
def validate_flow_id(cls, value):
@ -173,7 +181,7 @@ class MessageTable(MessageBase, table=True): # type: ignore[call-arg]
return value
@field_validator("properties", "content_blocks", mode="before")
@field_validator("properties", "content_blocks", "session_metadata", mode="before")
@classmethod
def validate_properties_or_content_blocks(cls, value):
if isinstance(value, list):
@ -185,11 +193,13 @@ class MessageTable(MessageBase, table=True): # type: ignore[call-arg]
return cls._sanitize_json(value)
@field_serializer("properties", "content_blocks")
@field_serializer("properties", "content_blocks", "session_metadata")
@classmethod
def serialize_properties_or_content_blocks(cls, value) -> dict | list[dict]:
def serialize_properties_or_content_blocks(cls, value) -> dict | list[dict] | None:
# Redundant sanitization here acts as a defensive measure for rows
# already in the database that might contain NaN/Infinity values.
if value is None:
return None
if isinstance(value, list):
value = [cls.serialize_properties_or_content_blocks(item) for item in value]
elif hasattr(value, "model_dump"):
@ -203,10 +213,11 @@ class MessageTable(MessageBase, table=True): # type: ignore[call-arg]
class MessageRead(MessageBase):
id: UUID
flow_id: UUID | None = Field()
session_metadata: dict | None = None
class MessageCreate(MessageBase):
pass
session_metadata: dict | None = None
class MessageUpdate(SQLModel):
@ -219,3 +230,4 @@ class MessageUpdate(SQLModel):
edit: bool | None = None
error: bool | None = None
properties: Properties | None = None
session_metadata: dict | None = None

View File

@ -0,0 +1,324 @@
"""Unit tests for session_metadata functionality in Message and MessageTable."""
from uuid import uuid4
import pytest
from langflow.memory import aadd_messages, aget_messages, astore_message
from langflow.schema.message import Message
from langflow.services.database.models.message import MessageCreate, MessageRead
from langflow.services.database.models.message.model import MessageTable
from langflow.services.deps import session_scope
@pytest.fixture
def sample_session_metadata():
"""Sample session metadata for testing."""
return {
"tenant_id": "tenant-123",
"user_id": "user-456",
"region": "us-east-1",
"retention_profile": "standard",
"data_flags": {"pii": True, "sensitive": False},
"custom_fields": {"department": "engineering", "project": "langflow"},
}
@pytest.fixture
def minimal_session_metadata():
"""Minimal session metadata for testing."""
return {
"tenant_id": "tenant-789",
"user_id": "user-012",
}
@pytest.mark.usefixtures("client")
async def test_message_with_session_metadata(sample_session_metadata):
"""Test creating a Message with session_metadata."""
message = Message(
text="Test message with metadata",
sender="User",
sender_name="Test User",
session_id="test_session_1",
session_metadata=sample_session_metadata,
)
assert message.session_metadata == sample_session_metadata
assert message.session_metadata["tenant_id"] == "tenant-123"
assert message.session_metadata["user_id"] == "user-456"
assert message.session_metadata["region"] == "us-east-1"
@pytest.mark.usefixtures("client")
async def test_message_without_session_metadata():
"""Test creating a Message without session_metadata (backward compatibility)."""
message = Message(
text="Test message without metadata",
sender="User",
sender_name="Test User",
session_id="test_session_2",
)
assert message.session_metadata is None
@pytest.mark.usefixtures("client")
async def test_store_message_with_session_metadata(sample_session_metadata):
"""Test storing a message with session_metadata."""
session_id = f"stored_session_{uuid4()}"
message = Message(
text="Stored message with metadata",
sender="User",
sender_name="Test User",
session_id=session_id,
session_metadata=sample_session_metadata,
)
await astore_message(message)
# Retrieve and verify
stored_messages = await aget_messages(sender="User", session_id=session_id)
assert len(stored_messages) == 1
assert stored_messages[0].text == "Stored message with metadata"
assert stored_messages[0].session_metadata == sample_session_metadata
@pytest.mark.usefixtures("client")
async def test_store_message_without_session_metadata():
"""Test storing a message without session_metadata (backward compatibility)."""
session_id = f"stored_session_{uuid4()}"
message = Message(
text="Stored message without metadata",
sender="User",
sender_name="Test User",
session_id=session_id,
)
await astore_message(message)
# Retrieve and verify
stored_messages = await aget_messages(sender="User", session_id=session_id)
assert len(stored_messages) == 1
assert stored_messages[0].text == "Stored message without metadata"
assert stored_messages[0].session_metadata is None
@pytest.mark.usefixtures("client")
async def test_add_messages_with_session_metadata(sample_session_metadata, minimal_session_metadata):
"""Test adding multiple messages with different session_metadata."""
session_id = f"batch_session_{uuid4()}"
messages = [
Message(
text="Message 1 with full metadata",
sender="User",
sender_name="User 1",
session_id=session_id,
session_metadata=sample_session_metadata,
),
Message(
text="Message 2 with minimal metadata",
sender="User",
sender_name="User 2",
session_id=session_id,
session_metadata=minimal_session_metadata,
),
Message(
text="Message 3 without metadata",
sender="User",
sender_name="User 3",
session_id=session_id,
),
]
added_messages = await aadd_messages(messages)
assert len(added_messages) == 3
assert added_messages[0].session_metadata == sample_session_metadata
assert added_messages[1].session_metadata == minimal_session_metadata
assert added_messages[2].session_metadata is None
@pytest.mark.usefixtures("client")
async def test_messagetable_from_message_with_metadata(sample_session_metadata):
"""Test MessageTable.from_message() extracts session_metadata correctly."""
message = Message(
text="Test message",
sender="User",
sender_name="Test User",
session_id="test_session_3",
session_metadata=sample_session_metadata,
)
message_table = MessageTable.from_message(message, flow_id=uuid4())
assert message_table.session_metadata == sample_session_metadata
assert message_table.session_metadata["tenant_id"] == "tenant-123"
assert message_table.session_metadata["user_id"] == "user-456"
@pytest.mark.usefixtures("client")
async def test_messagetable_from_message_without_metadata():
"""Test MessageTable.from_message() handles missing session_metadata."""
message = Message(
text="Test message",
sender="User",
sender_name="Test User",
session_id="test_session_4",
)
message_table = MessageTable.from_message(message, flow_id=uuid4())
assert message_table.session_metadata is None
@pytest.mark.usefixtures("client")
async def test_messagecreate_with_session_metadata(sample_session_metadata):
"""Test MessageCreate schema with session_metadata."""
message_create = MessageCreate(
text="Test message",
sender="User",
sender_name="Test User",
session_id="test_session_5",
session_metadata=sample_session_metadata,
)
assert message_create.session_metadata == sample_session_metadata
@pytest.mark.usefixtures("client")
async def test_messagecreate_without_session_metadata():
"""Test MessageCreate schema without session_metadata."""
message_create = MessageCreate(
text="Test message",
sender="User",
sender_name="Test User",
session_id="test_session_6",
)
assert message_create.session_metadata is None
@pytest.mark.usefixtures("client")
async def test_messageread_with_session_metadata(sample_session_metadata):
"""Test MessageRead schema includes session_metadata."""
async with session_scope() as session:
message_create = MessageCreate(
text="Test message",
sender="User",
sender_name="Test User",
session_id=f"test_session_{uuid4()}",
session_metadata=sample_session_metadata,
)
message_table = MessageTable.model_validate(message_create, from_attributes=True)
session.add(message_table)
await session.commit()
await session.refresh(message_table)
message_read = MessageRead.model_validate(message_table, from_attributes=True)
assert message_read.session_metadata == sample_session_metadata
@pytest.mark.usefixtures("client")
async def test_session_metadata_persistence_and_retrieval(sample_session_metadata):
"""Test full cycle: create, store, retrieve, and verify session_metadata."""
session_id = f"full_cycle_session_{uuid4()}"
# Create and store
message = Message(
text="Full cycle test message",
sender="User",
sender_name="Test User",
session_id=session_id,
session_metadata=sample_session_metadata,
)
await astore_message(message)
# Retrieve
retrieved_messages = await aget_messages(sender="User", session_id=session_id)
# Verify
assert len(retrieved_messages) == 1
retrieved = retrieved_messages[0]
assert retrieved.text == "Full cycle test message"
assert retrieved.session_metadata is not None
assert retrieved.session_metadata["tenant_id"] == "tenant-123"
assert retrieved.session_metadata["user_id"] == "user-456"
assert retrieved.session_metadata["region"] == "us-east-1"
assert retrieved.session_metadata["retention_profile"] == "standard"
assert retrieved.session_metadata["data_flags"]["pii"] is True
assert retrieved.session_metadata["custom_fields"]["department"] == "engineering"
@pytest.mark.usefixtures("client")
async def test_session_metadata_json_serialization():
"""Test that session_metadata is properly serialized as JSON."""
session_id = f"json_test_session_{uuid4()}"
metadata = {
"tenant_id": "tenant-json",
"nested": {"key1": "value1", "key2": [1, 2, 3]},
"array": ["item1", "item2"],
"number": 42,
"boolean": True,
}
message = Message(
text="JSON serialization test",
sender="User",
sender_name="Test User",
session_id=session_id,
session_metadata=metadata,
)
await astore_message(message)
# Retrieve and verify complex JSON structure
retrieved_messages = await aget_messages(sender="User", session_id=session_id)
assert len(retrieved_messages) == 1
retrieved_metadata = retrieved_messages[0].session_metadata
assert retrieved_metadata["tenant_id"] == "tenant-json"
assert retrieved_metadata["nested"]["key1"] == "value1"
assert retrieved_metadata["nested"]["key2"] == [1, 2, 3]
assert retrieved_metadata["array"] == ["item1", "item2"]
assert retrieved_metadata["number"] == 42
assert retrieved_metadata["boolean"] is True
@pytest.mark.usefixtures("client")
async def test_empty_session_metadata():
"""Test storing message with empty dict as session_metadata."""
session_id = f"empty_metadata_session_{uuid4()}"
message = Message(
text="Empty metadata test",
sender="User",
sender_name="Test User",
session_id=session_id,
session_metadata={},
)
await astore_message(message)
retrieved_messages = await aget_messages(sender="User", session_id=session_id)
assert len(retrieved_messages) == 1
assert retrieved_messages[0].session_metadata == {}
@pytest.mark.usefixtures("client")
async def test_session_metadata_retrieval():
"""Test retrieving session_metadata from stored messages."""
session_id = f"retrieval_metadata_session_{uuid4()}"
# Create initial message
initial_metadata = {"tenant_id": "tenant-initial", "user_id": "user-initial"}
message = Message(
text="Initial message",
sender="User",
sender_name="Test User",
session_id=session_id,
session_metadata=initial_metadata,
)
await astore_message(message)
# Retrieve and verify
messages = await aget_messages(sender="User", session_id=session_id)
assert len(messages) == 1
assert messages[0].session_metadata == initial_metadata

View File

@ -1812,6 +1812,7 @@
"sender": null,
"sender_name": null,
"session_id": "",
"session_metadata": null,
"text": "http://localhost:11434"
},
"default_value": "",
@ -2223,6 +2224,7 @@
"sender": null,
"sender_name": null,
"session_id": "",
"session_metadata": null,
"text": "http://localhost:11434"
},
"default_value": "",
@ -2629,6 +2631,7 @@
"sender": null,
"sender_name": null,
"session_id": "",
"session_metadata": null,
"text": "http://localhost:11434"
},
"default_value": "",
@ -118887,6 +118890,6 @@
"num_components": 360,
"num_modules": 97
},
"sha256": "a7dbe8e05fcd66e2b9c2cc1be11bc19f5a06e77103fea3337a7e405c99e7fc3d",
"sha256": "75bc084b1d3f8e8a508e5883ff630167793df395ed3a574d109b2d35e908a1b2",
"version": "0.4.0"
}

View File

@ -74,6 +74,7 @@ class Message(Data):
category: Literal["message", "error", "warning", "info"] | None = "message"
content_blocks: list[ContentBlock] = Field(default_factory=list)
duration: int | None = None
session_metadata: dict | None = None
@field_validator("flow_id", mode="before")
@classmethod
@ -231,6 +232,7 @@ class Message(Data):
flow_id=data.flow_id,
error=data.error,
edit=data.edit,
session_metadata=getattr(data, "session_metadata", None),
)
@field_serializer("text", mode="plain")
@ -430,6 +432,7 @@ class MessageResponse(DefaultModel):
properties: Properties | None = None
category: str | None = None
content_blocks: list[ContentBlock] | None = None
session_metadata: dict | None = None
@field_validator("content_blocks", mode="before")
@classmethod
@ -483,6 +486,7 @@ class MessageResponse(DefaultModel):
files=message.files or [],
timestamp=message.timestamp,
flow_id=flow_id,
session_metadata=getattr(message, "session_metadata", None),
)
@ -534,6 +538,7 @@ class ErrorMessage(Message):
source: Source | None = None,
trace_name: str | None = None,
flow_id: UUID | str | None = None,
session_metadata: dict | None = None,
) -> None:
# This is done to avoid circular imports
if exception.__class__.__name__ == "ExceptionWithMessageError" and exception.__cause__ is not None:
@ -580,6 +585,7 @@ class ErrorMessage(Message):
)
],
flow_id=flow_id,
session_metadata=session_metadata,
)