fix: session cloning working

This commit is contained in:
Vineeth Voruganti 2025-04-02 17:24:26 -04:00
parent 4cdf7222d2
commit fffbcebf05
5 changed files with 80 additions and 22 deletions

View File

@ -124,12 +124,10 @@ async def get_latest_user_representation(
) -> str:
stmt = (
select(models.Metamessage)
.join(models.Message, models.Message.public_id == models.Metamessage.message_id)
.join(models.Session, models.Message.session_id == models.Session.public_id)
.join(models.User, models.User.public_id == models.Session.user_id)
.join(models.User, models.User.public_id == models.Metamessage.user_id)
.join(models.App, models.App.public_id == models.User.app_id)
.where(models.App.public_id == app_id)
.where(models.User.public_id == user_id)
.where(models.Metamessage.user_id == user_id)
.where(models.Metamessage.metamessage_type == "user_representation")
.order_by(models.Metamessage.id.desc()) # get the most recent
.limit(1)

View File

@ -571,31 +571,46 @@ async def clone_session(
)
# Handle metamessages if deep copy is requested
if deep_copy and message_id_map:
# Fetch all metamessages in a single query
if deep_copy:
# Fetch all metamessages tied to the session in a single query
stmt = select(models.Metamessage).where(
models.Metamessage.message_id.in_(message_id_map.keys())
models.Metamessage.session_id == original_session_id
)
if cutoff_message_id is not None and cutoff_message is not None:
# Only get metamessages related to messages we're cloning
message_ids = [message.public_id for message in messages_to_clone]
stmt = stmt.where(
(models.Metamessage.message_id.is_(None)) |
(models.Metamessage.message_id.in_(message_ids))
)
metamessages_result = await db.scalars(stmt)
metamessages = metamessages_result.all()
if metamessages:
# Prepare bulk insert data for metamessages
new_metamessages = [
{
"user_id": user_id,
new_metamessages = []
for meta in metamessages:
# Base metamessage data
meta_data = {
"user_id": meta.user_id, # Preserve original user
"session_id": new_session.public_id,
"message_id": message_id_map[meta.message_id],
"metamessage_type": meta.metamessage_type,
"content": meta.content,
"h_metadata": meta.h_metadata,
}
for meta in metamessages
]
# If the metamessage was tied to a message, tie it to the corresponding new message
if meta.message_id is not None and meta.message_id in message_id_map:
meta_data["message_id"] = message_id_map[meta.message_id]
new_metamessages.append(meta_data)
# Bulk insert metamessages using modern insert syntax
stmt = insert(models.Metamessage)
await db.execute(stmt, new_metamessages)
if new_metamessages:
stmt = insert(models.Metamessage)
await db.execute(stmt, new_metamessages)
await db.commit()

View File

@ -66,7 +66,7 @@ async def setup_test_database(db_url):
Returns:
engine: SQLAlchemy engine
"""
engine = create_async_engine(str(db_url))
engine = create_async_engine(str(db_url), echo=True)
async with engine.connect() as conn:
try:
logger.info("Attempting to create pgvector extension...")
@ -94,7 +94,10 @@ async def db_engine():
create_test_database(TEST_DB_URL)
engine = await setup_test_database(TEST_DB_URL)
# Drop all tables first to ensure clean state
async with engine.begin() as conn:
await conn.run_sync(Base.metadata.drop_all)
# Then create all tables with current models
await conn.run_sync(Base.metadata.create_all)
yield engine

View File

@ -333,8 +333,8 @@ async def test_deep_clone_session(client, db_session, sample_data):
assert data["items"][1]["metadata"] == {"key": "value2"}
response = client.post(
f"/v1/apps/{test_app.public_id}/users/{test_user.public_id}/sessions/{cloned_session_id}/metamessages/list",
json={},
f"/v1/apps/{test_app.public_id}/users/{test_user.public_id}/metamessages/list",
json={"session_id": cloned_session_id},
)
assert response.status_code == 200
@ -445,8 +445,8 @@ async def test_partial_deep_clone_session(client, db_session, sample_data):
assert data["items"][0]["metadata"] == {"key": "value"}
response = client.post(
f"/v1/apps/{test_app.public_id}/users/{test_user.public_id}/sessions/{cloned_session_id}/metamessages/list",
json={},
f"/v1/apps/{test_app.public_id}/users/{test_user.public_id}/metamessages/list",
json={"session_id": cloned_session_id},
)
assert response.status_code == 200

View File

@ -281,11 +281,12 @@ def test_metamessage_validations_api(client, sample_data):
# Test content too long
response = client.post(
f"/v1/apps/{test_app.public_id}/users/{test_user.public_id}/sessions/{session_id}/metamessages",
f"/v1/apps/{test_app.public_id}/users/{test_user.public_id}/metamessages",
json={
"metamessage_type": "test_type",
"content": "a" * 50001,
"message_id": message_id,
"session_id": session_id,
"metadata": {},
},
)
@ -371,7 +372,48 @@ def test_session_validations_api(client, sample_data):
assert response.status_code == 422
def test_agent_query_validations_api(client, sample_data):
def test_agent_query_validations_api(client, sample_data, monkeypatch):
# Mock the functions in agent.py that are causing the database issues
async def mock_get_user_representation(*args, **kwargs):
return "Mock user representation"
async def mock_chat_history(*args, **kwargs):
return "Mock chat history"
# Mock the Dialectic.call method
def mock_dialectic_call(self):
# Create a mock response that will work with line 179 in agent.py:
# return schemas.AgentChat(content=response[0].text)
class MockText:
def __init__(self):
self.text = "Mock response"
# Return a list with MockText object at index 0
return [MockText()]
# Mock the Dialectic.stream method
def mock_dialectic_stream(self):
class MockStream:
def __enter__(self):
return self
def __exit__(self, *args):
pass
@property
def text_stream(self):
yield "Mock streamed response"
return MockStream()
# Apply the monkeypatches
monkeypatch.setattr(
"src.agent.get_latest_user_representation", mock_get_user_representation
)
monkeypatch.setattr("src.agent.chat_history", mock_chat_history)
monkeypatch.setattr("src.agent.Dialectic.call", mock_dialectic_call)
monkeypatch.setattr("src.agent.Dialectic.stream", mock_dialectic_stream)
test_app, test_user = sample_data
# Create a session first since agent queries are likely session-based
session_response = client.post(