Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
16 changes: 16 additions & 0 deletions py/core/main/api/v3/conversations_router.py
Original file line number Diff line number Diff line change
Expand Up @@ -499,9 +499,14 @@ async def update_conversation(
This endpoint updates the name of an existing conversation
identified by its UUID.
"""
requesting_user_id = (
None if auth_user.is_superuser else [auth_user.id]
)

return await self.services.management.update_conversation( # type: ignore
conversation_id=id,
name=name,
user_ids=requesting_user_id,
)

@self.router.delete(
Expand Down Expand Up @@ -651,11 +656,16 @@ async def add_message(
if role not in ["user", "assistant", "system"]:
raise R2RException("Invalid role", status_code=400)
message = Message(role=role, content=content)
requesting_user_id = (
None if auth_user.is_superuser else [auth_user.id]
)

return await self.services.management.add_message( # type: ignore
conversation_id=id,
content=message,
parent_id=parent_id,
metadata=metadata,
user_ids=requesting_user_id,
)

@self.router.post(
Expand Down Expand Up @@ -730,8 +740,14 @@ async def update_message(
This endpoint updates the content of an existing message in a
conversation.
"""
requesting_user_id = (
None if auth_user.is_superuser else [auth_user.id]
)

return await self.services.management.edit_message( # type: ignore
conversation_id=id,
message_id=message_id,
new_content=content,
additional_metadata=metadata,
user_ids=requesting_user_id,
)
15 changes: 13 additions & 2 deletions py/core/main/services/management_service.py
Original file line number Diff line number Diff line change
Expand Up @@ -800,33 +800,44 @@ async def add_message(
content: Message,
parent_id: Optional[UUID] = None,
metadata: Optional[dict] = None,
user_ids: Optional[list[UUID]] = None,
) -> MessageResponse:
return await self.providers.database.conversations_handler.add_message(
conversation_id=conversation_id,
content=content,
parent_id=parent_id,
metadata=metadata,
filter_user_ids=user_ids,
)

async def edit_message(
self,
message_id: UUID,
new_content: Optional[str] = None,
additional_metadata: Optional[dict] = None,
conversation_id: Optional[UUID] = None,
user_ids: Optional[list[UUID]] = None,
) -> dict[str, Any]:
return (
await self.providers.database.conversations_handler.edit_message(
message_id=message_id,
new_content=new_content,
additional_metadata=additional_metadata or {},
conversation_id=conversation_id,
filter_user_ids=user_ids,
)
)

async def update_conversation(
self, conversation_id: UUID, name: str
self,
conversation_id: UUID,
name: str,
user_ids: Optional[list[UUID]] = None,
) -> ConversationResponse:
return await self.providers.database.conversations_handler.update_conversation(
conversation_id=conversation_id, name=name
conversation_id=conversation_id,
name=name,
filter_user_ids=user_ids,
)

async def delete_conversation(
Expand Down
53 changes: 44 additions & 9 deletions py/core/providers/database/conversations.py
Original file line number Diff line number Diff line change
Expand Up @@ -211,6 +211,7 @@ async def add_message(
parent_id: Optional[UUID] = None,
metadata: Optional[dict] = None,
max_image_size_bytes: int = 5 * 1024 * 1024, # 5MB default
filter_user_ids: Optional[list[UUID]] = None,
) -> MessageResponse:
# Validate image size
try:
Expand All @@ -226,12 +227,18 @@ async def add_message(
) from e

# 1) Validate that conversation and parent exist (existing code)
conditions = ["id = $1"]
params: list[Any] = [conversation_id]
if filter_user_ids:
conditions.append("user_id = ANY($2)")
params.append(filter_user_ids)

conv_check_query = f"""
SELECT 1 FROM {self._get_table_name("conversations")}
WHERE id = $1
WHERE {" AND ".join(conditions)}
"""
conv_row = await self.connection_manager.fetchrow_query(
conv_check_query, [conversation_id]
conv_check_query, params
)
if not conv_row:
raise R2RException(
Expand Down Expand Up @@ -298,14 +305,28 @@ async def edit_message(
message_id: UUID,
new_content: str | None = None,
additional_metadata: dict | None = None,
conversation_id: UUID | None = None,
filter_user_ids: Optional[list[UUID]] = None,
) -> dict[str, Any]:
# Get the original message
conditions = ["m.id = $1"]
params: list[Any] = [message_id]
if conversation_id:
conditions.append(f"m.conversation_id = ${len(params) + 1}")
params.append(conversation_id)
if filter_user_ids:
conditions.append(f"c.user_id = ANY(${len(params) + 1})")
params.append(filter_user_ids)

query = f"""
SELECT conversation_id, parent_id, content, metadata, created_at
FROM {self._get_table_name("messages")}
WHERE id = $1
SELECT m.conversation_id, m.parent_id, m.content, m.metadata,
m.created_at
FROM {self._get_table_name("messages")} m
JOIN {self._get_table_name("conversations")} c
ON c.id = m.conversation_id
WHERE {" AND ".join(conditions)}
"""
row = await self.connection_manager.fetchrow_query(query, [message_id])
row = await self.connection_manager.fetchrow_query(query, params)
if not row:
raise R2RException(
status_code=404,
Expand Down Expand Up @@ -487,13 +508,25 @@ async def get_conversation(
return response_messages

async def update_conversation(
self, conversation_id: UUID, name: str
self,
conversation_id: UUID,
name: str,
filter_user_ids: Optional[list[UUID]] = None,
) -> ConversationResponse:
try:
# Check if conversation exists
conv_query = f"SELECT 1 FROM {self._get_table_name('conversations')} WHERE id = $1"
conditions = ["id = $1"]
params: list[Any] = [conversation_id]
if filter_user_ids:
conditions.append("user_id = ANY($2)")
params.append(filter_user_ids)

conv_query = f"""
SELECT 1 FROM {self._get_table_name("conversations")}
WHERE {" AND ".join(conditions)}
"""
conv_row = await self.connection_manager.fetchrow_query(
conv_query, [conversation_id]
conv_query, params
)
if not conv_row:
raise R2RException(
Expand All @@ -515,6 +548,8 @@ async def update_conversation(
user_id=updated_row["user_id"] or None,
name=name,
)
except R2RException:
raise
except Exception as e:
raise HTTPException(
status_code=500,
Expand Down
Loading