Skip to content

Commit 0510a72

Browse files
committed
fix: enforce conversation ownership on writes
1 parent 9c5a94d commit 0510a72

4 files changed

Lines changed: 299 additions & 11 deletions

File tree

py/core/main/api/v3/conversations_router.py

Lines changed: 16 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -499,9 +499,14 @@ async def update_conversation(
499499
This endpoint updates the name of an existing conversation
500500
identified by its UUID.
501501
"""
502+
requesting_user_id = (
503+
None if auth_user.is_superuser else [auth_user.id]
504+
)
505+
502506
return await self.services.management.update_conversation( # type: ignore
503507
conversation_id=id,
504508
name=name,
509+
user_ids=requesting_user_id,
505510
)
506511

507512
@self.router.delete(
@@ -651,11 +656,16 @@ async def add_message(
651656
if role not in ["user", "assistant", "system"]:
652657
raise R2RException("Invalid role", status_code=400)
653658
message = Message(role=role, content=content)
659+
requesting_user_id = (
660+
None if auth_user.is_superuser else [auth_user.id]
661+
)
662+
654663
return await self.services.management.add_message( # type: ignore
655664
conversation_id=id,
656665
content=message,
657666
parent_id=parent_id,
658667
metadata=metadata,
668+
user_ids=requesting_user_id,
659669
)
660670

661671
@self.router.post(
@@ -730,8 +740,14 @@ async def update_message(
730740
This endpoint updates the content of an existing message in a
731741
conversation.
732742
"""
743+
requesting_user_id = (
744+
None if auth_user.is_superuser else [auth_user.id]
745+
)
746+
733747
return await self.services.management.edit_message( # type: ignore
748+
conversation_id=id,
734749
message_id=message_id,
735750
new_content=content,
736751
additional_metadata=metadata,
752+
user_ids=requesting_user_id,
737753
)

py/core/main/services/management_service.py

Lines changed: 13 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -800,33 +800,44 @@ async def add_message(
800800
content: Message,
801801
parent_id: Optional[UUID] = None,
802802
metadata: Optional[dict] = None,
803+
user_ids: Optional[list[UUID]] = None,
803804
) -> MessageResponse:
804805
return await self.providers.database.conversations_handler.add_message(
805806
conversation_id=conversation_id,
806807
content=content,
807808
parent_id=parent_id,
808809
metadata=metadata,
810+
filter_user_ids=user_ids,
809811
)
810812

811813
async def edit_message(
812814
self,
813815
message_id: UUID,
814816
new_content: Optional[str] = None,
815817
additional_metadata: Optional[dict] = None,
818+
conversation_id: Optional[UUID] = None,
819+
user_ids: Optional[list[UUID]] = None,
816820
) -> dict[str, Any]:
817821
return (
818822
await self.providers.database.conversations_handler.edit_message(
819823
message_id=message_id,
820824
new_content=new_content,
821825
additional_metadata=additional_metadata or {},
826+
conversation_id=conversation_id,
827+
filter_user_ids=user_ids,
822828
)
823829
)
824830

825831
async def update_conversation(
826-
self, conversation_id: UUID, name: str
832+
self,
833+
conversation_id: UUID,
834+
name: str,
835+
user_ids: Optional[list[UUID]] = None,
827836
) -> ConversationResponse:
828837
return await self.providers.database.conversations_handler.update_conversation(
829-
conversation_id=conversation_id, name=name
838+
conversation_id=conversation_id,
839+
name=name,
840+
filter_user_ids=user_ids,
830841
)
831842

832843
async def delete_conversation(

py/core/providers/database/conversations.py

Lines changed: 44 additions & 9 deletions
Original file line numberDiff line numberDiff line change
@@ -211,6 +211,7 @@ async def add_message(
211211
parent_id: Optional[UUID] = None,
212212
metadata: Optional[dict] = None,
213213
max_image_size_bytes: int = 5 * 1024 * 1024, # 5MB default
214+
filter_user_ids: Optional[list[UUID]] = None,
214215
) -> MessageResponse:
215216
# Validate image size
216217
try:
@@ -226,12 +227,18 @@ async def add_message(
226227
) from e
227228

228229
# 1) Validate that conversation and parent exist (existing code)
230+
conditions = ["id = $1"]
231+
params: list[Any] = [conversation_id]
232+
if filter_user_ids:
233+
conditions.append("user_id = ANY($2)")
234+
params.append(filter_user_ids)
235+
229236
conv_check_query = f"""
230237
SELECT 1 FROM {self._get_table_name("conversations")}
231-
WHERE id = $1
238+
WHERE {" AND ".join(conditions)}
232239
"""
233240
conv_row = await self.connection_manager.fetchrow_query(
234-
conv_check_query, [conversation_id]
241+
conv_check_query, params
235242
)
236243
if not conv_row:
237244
raise R2RException(
@@ -298,14 +305,28 @@ async def edit_message(
298305
message_id: UUID,
299306
new_content: str | None = None,
300307
additional_metadata: dict | None = None,
308+
conversation_id: UUID | None = None,
309+
filter_user_ids: Optional[list[UUID]] = None,
301310
) -> dict[str, Any]:
302311
# Get the original message
312+
conditions = ["m.id = $1"]
313+
params: list[Any] = [message_id]
314+
if conversation_id:
315+
conditions.append(f"m.conversation_id = ${len(params) + 1}")
316+
params.append(conversation_id)
317+
if filter_user_ids:
318+
conditions.append(f"c.user_id = ANY(${len(params) + 1})")
319+
params.append(filter_user_ids)
320+
303321
query = f"""
304-
SELECT conversation_id, parent_id, content, metadata, created_at
305-
FROM {self._get_table_name("messages")}
306-
WHERE id = $1
322+
SELECT m.conversation_id, m.parent_id, m.content, m.metadata,
323+
m.created_at
324+
FROM {self._get_table_name("messages")} m
325+
JOIN {self._get_table_name("conversations")} c
326+
ON c.id = m.conversation_id
327+
WHERE {" AND ".join(conditions)}
307328
"""
308-
row = await self.connection_manager.fetchrow_query(query, [message_id])
329+
row = await self.connection_manager.fetchrow_query(query, params)
309330
if not row:
310331
raise R2RException(
311332
status_code=404,
@@ -487,13 +508,25 @@ async def get_conversation(
487508
return response_messages
488509

489510
async def update_conversation(
490-
self, conversation_id: UUID, name: str
511+
self,
512+
conversation_id: UUID,
513+
name: str,
514+
filter_user_ids: Optional[list[UUID]] = None,
491515
) -> ConversationResponse:
492516
try:
493517
# Check if conversation exists
494-
conv_query = f"SELECT 1 FROM {self._get_table_name('conversations')} WHERE id = $1"
518+
conditions = ["id = $1"]
519+
params: list[Any] = [conversation_id]
520+
if filter_user_ids:
521+
conditions.append("user_id = ANY($2)")
522+
params.append(filter_user_ids)
523+
524+
conv_query = f"""
525+
SELECT 1 FROM {self._get_table_name("conversations")}
526+
WHERE {" AND ".join(conditions)}
527+
"""
495528
conv_row = await self.connection_manager.fetchrow_query(
496-
conv_query, [conversation_id]
529+
conv_query, params
497530
)
498531
if not conv_row:
499532
raise R2RException(
@@ -515,6 +548,8 @@ async def update_conversation(
515548
user_id=updated_row["user_id"] or None,
516549
name=name,
517550
)
551+
except R2RException:
552+
raise
518553
except Exception as e:
519554
raise HTTPException(
520555
status_code=500,

0 commit comments

Comments
 (0)