fix: session cloning working
This commit is contained in:
parent
4cdf7222d2
commit
fffbcebf05
|
|
@ -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)
|
||||
|
|
|
|||
37
src/crud.py
37
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()
|
||||
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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
|
||||
|
|
|
|||
|
|
@ -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(
|
||||
|
|
|
|||
Loading…
Reference in New Issue