From fffbcebf059df4896e60b1b5fa581c744c930e83 Mon Sep 17 00:00:00 2001 From: Vineeth Voruganti <13438633+VVoruganti@users.noreply.github.com> Date: Wed, 2 Apr 2025 17:24:26 -0400 Subject: [PATCH] fix: session cloning working --- src/agent.py | 6 ++-- src/crud.py | 37 ++++++++++++++++------- tests/conftest.py | 5 +++- tests/routes/test_sessions.py | 8 ++--- tests/routes/test_validation_api.py | 46 +++++++++++++++++++++++++++-- 5 files changed, 80 insertions(+), 22 deletions(-) diff --git a/src/agent.py b/src/agent.py index 45d9bbb5..a3fac1c5 100644 --- a/src/agent.py +++ b/src/agent.py @@ -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) diff --git a/src/crud.py b/src/crud.py index 492dbf00..3c0248db 100644 --- a/src/crud.py +++ b/src/crud.py @@ -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() diff --git a/tests/conftest.py b/tests/conftest.py index e0eea897..a0d9950f 100644 --- a/tests/conftest.py +++ b/tests/conftest.py @@ -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 diff --git a/tests/routes/test_sessions.py b/tests/routes/test_sessions.py index f4a75624..6e51fa34 100644 --- a/tests/routes/test_sessions.py +++ b/tests/routes/test_sessions.py @@ -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 diff --git a/tests/routes/test_validation_api.py b/tests/routes/test_validation_api.py index 8b0b515e..7fef7bb4 100644 --- a/tests/routes/test_validation_api.py +++ b/tests/routes/test_validation_api.py @@ -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(