From 2e9e95e16c847de6b9ea8df755d6736e1fbf672f Mon Sep 17 00:00:00 2001 From: Ryan Westfall Date: Sat, 1 Aug 2026 12:35:00 -0700 Subject: [PATCH] Stove-pipe RAG retrieval to prevent cross-tenant leakage (#40) (#41) ## Summary - Closes [#40](https://git.aimloperations.com/ai_ml_operations/chat_backend/issues/40) - Aligns chat/RAG with the abc_worker stove-pipe pattern ([b13cec8](https://git.aimloperations.com/GIS/abc_worker/commit/b13cec88f9730289f245b66b95aa11ee33cb4e8e)): immutable `ChatCompanyScope` per turn, conversation ownership validation, fail-closed Chroma filters - Prefer ASGI/JWT identity over client email; never bind identity from `conversation_id` alone - Close `ConversationDetailView` IDOR (prompts only for `request.user`) ## Changes - New `services/chat_tenant_scope.py` with frozen `ChatCompanyScope` + ownership checks - WebSocket consumers (`consumers.py` / `consumers_graph.py`) validate scope before `get_messages` / RAG - `search_documents` requires a workspace (no more `filter: None` over the shared collection) - Ingest writes `company_id` metadata (retrieval still keys on `workspace_id` for back-compat) - Legacy `get_retriever` always applies a workspace filter ## Test plan - [x] `manage.py test chat_backend.tests.test_chat_tenant_scope chat_backend.tests.test_consumers chat_backend.tests.test_services_rag chat_backend.tests.test_views_conversations` - [ ] Manual: user A cannot stream RAG context from user B `conversation_id` - [ ] Manual: RAG still returns own-company docs after deploy (existing vectors with `workspace_id` only) - [ ] Follow-up: FE can send JWT `token`/`access` on WS payloads for stronger identity bindingReviewed-on: https://git.aimloperations.com/ai_ml_operations/chat_backend/pulls/41 --- llm_be/chat_backend/consumers.py | 131 +++++++++---- llm_be/chat_backend/consumers_graph.py | 98 +++++++--- .../services/chat_tenant_scope.py | 182 ++++++++++++++++++ llm_be/chat_backend/services/rag_services.py | 53 +++-- .../tests/test_chat_tenant_scope.py | 64 ++++++ llm_be/chat_backend/tests/test_consumers.py | 28 ++- .../chat_backend/tests/test_services_rag.py | 16 +- .../tests/test_views_conversations.py | 10 + llm_be/chat_backend/views.py | 15 +- 9 files changed, 507 insertions(+), 90 deletions(-) create mode 100644 llm_be/chat_backend/services/chat_tenant_scope.py create mode 100644 llm_be/chat_backend/tests/test_chat_tenant_scope.py diff --git a/llm_be/chat_backend/consumers.py b/llm_be/chat_backend/consumers.py index 7f97115..2dc900a 100644 --- a/llm_be/chat_backend/consumers.py +++ b/llm_be/chat_backend/consumers.py @@ -23,6 +23,14 @@ from .models import Conversation, Prompt, PromptMetric, DocumentWorkspace, Docum from .serializers import PromptSerializer from .services.llm_service import AsyncLLMService from .services.rag_services import AsyncRAGService +from .services.chat_tenant_scope import ( + ChatTenantScopeError, + asgi_user_or_none, + create_conversation_for_user, + get_workspace_for_scope, + resolve_chat_company_scope, + resolve_chat_user as resolve_chat_user_sync, +) from .services.title_generator import title_generator from .services.moderation_classifier import moderation_classifier, ModerationLabel from .services.prompt_classifier.prompt_classifier import PromptClassifier, PromptType @@ -47,33 +55,27 @@ CHANNEL_NAME: str = "llm_messages" MODEL_NAME: str = "llama3.2" PROMPT_CLASSIFIER = PromptClassifier() -@database_sync_to_async -def create_conversation(prompt, email, title): - # return the conversation id - conversation = Conversation.objects.create(title=title) - conversation.save() - user = CustomUser.objects.get(email=email) - conversation.user_id = user.id - conversation.save() - return conversation.id +@database_sync_to_async +def create_conversation(prompt, email, title, user=None): + """Create a conversation for ``user`` (preferred) or legacy ``email``.""" + if user is None: + user = CustomUser.objects.get(email=email) + return create_conversation_for_user(user, title) @database_sync_to_async -def resolve_chat_user(email=None, conversation_id=None): - if email: - user = CustomUser.objects.filter(email__iexact=email).first() - if user: - return user - if conversation_id: - conversation = ( - Conversation.objects.select_related("user") - .filter(id=conversation_id) - .first() - ) - if conversation and conversation.user_id: - return conversation.user - return None +def resolve_chat_user( + email=None, conversation_id=None, token=None, authenticated_user=None +): + # conversation_id intentionally unused for identity — ownership is checked + # via resolve_chat_company_scope after the principal is known. + return resolve_chat_user_sync( + email=email, + token=token, + authenticated_user=authenticated_user, + conversation_id=conversation_id, + ) @database_sync_to_async @@ -89,9 +91,20 @@ def enforce_feature_gate(user, feature): @database_sync_to_async -def get_workspace(conversation_id): - conversation = Conversation.objects.get(id=conversation_id) - return DocumentWorkspace.objects.get(company=conversation.user.company) +def get_workspace(conversation_id, user=None): + """Resolve workspace only after conversation ownership is validated.""" + if user is None: + raise ChatTenantScopeError( + "Authenticated chat user is required.", + code="user_not_found", + ) + scope = resolve_chat_company_scope(user, conversation_id) + return get_workspace_for_scope(scope) + + +@database_sync_to_async +def resolve_tenant_scope(user, conversation_id=None): + return resolve_chat_company_scope(user, conversation_id) @database_sync_to_async @@ -223,12 +236,17 @@ def finish_prompt_metric(prompt_metric, response_length, tokens_in=None, tokens_ @database_sync_to_async -def get_retriever(conversation_id): +def get_retriever(conversation_id, user=None): + """Legacy helper — always applies a workspace metadata filter (fail closed).""" + if user is None: + raise ChatTenantScopeError( + "Authenticated chat user is required.", + code="user_not_found", + ) logger.info(f"getting workspace from conversation: {conversation_id}") - conversation = Conversation.objects.get(id=conversation_id) - logger.info(f"Got conversation: {conversation}") - workspace = DocumentWorkspace.objects.get(company=conversation.user.company) - logger.info(f"Got workspace: {conversation}") + scope = resolve_chat_company_scope(user, conversation_id) + workspace = get_workspace_for_scope(scope) + logger.info(f"Got workspace: {workspace.id} company={scope.company_id}") persist_directory = getattr( django_settings, "CHROMA_PERSIST_DIRECTORY", "./chroma_db/" ) @@ -236,7 +254,10 @@ def get_retriever(conversation_id): persist_directory=persist_directory, embedding=OllamaEmbeddings(**ollama_embeddings_kwargs()), ) - return vectorstore.as_retriever() + return vectorstore.as_retriever( + search_type="similarity", + search_kwargs={"k": 4, "filter": {"workspace_id": workspace.id}}, + ) async def get_conversation_file_async(conversation_id): try: @@ -289,6 +310,7 @@ class ChatConsumerAgain(AsyncWebsocketConsumer): message = normalize_user_message(data.get("message", None)) conversation_id = data.get("conversation_id", None) email = data.get("email", None) + token = data.get("token") or data.get("access") file = data.get("file", None) file_type = data.get("fileType", "") model = data.get("modelName", "Turbo") @@ -306,7 +328,10 @@ class ChatConsumerAgain(AsyncWebsocketConsumer): return chat_user = await resolve_chat_user( - email=email, conversation_id=conversation_id + email=email, + conversation_id=conversation_id, + token=token, + authenticated_user=asgi_user_or_none(self.scope.get("user")), ) if chat_user is None: await self.send_json_message( @@ -339,7 +364,37 @@ class ChatConsumerAgain(AsyncWebsocketConsumer): # we need to create a new conversation # we will generate a name for it too title = await title_generator.generate_async(message) - conversation_id = await create_conversation(message, email, title) + conversation_id = await create_conversation( + message, email, title, user=chat_user + ) + + try: + tenant_scope = await resolve_tenant_scope(chat_user, conversation_id) + except ChatTenantScopeError as exc: + logger.warning( + "websocket tenant validation failed conversation_id=%s user_id=%s code=%s", + conversation_id, + chat_user.id, + exc.code, + ) + await self.send_json_message( + json.dumps( + { + "type": "error", + "code": exc.code, + "content": exc.message, + } + ) + ) + return + + logger.info( + "chat_scope_validated conversation_id=%s user_id=%s company_id=%s workspace_id=%s", + tenant_scope.conversation_id, + tenant_scope.user_id, + tenant_scope.company_id, + tenant_scope.workspace_id, + ) if conversation_id: decoded_file = None @@ -443,8 +498,12 @@ class ChatConsumerAgain(AsyncWebsocketConsumer): if prompt_type == PromptType.RAG: service = AsyncRAGService() - workspace = await get_workspace(conversation_id) - return service.generate_response(messages, prompt_instance.message, workspace) + workspace = await get_workspace( + conversation_id, user=chat_user + ) + return service.generate_response( + messages, prompt_instance.message, workspace + ) elif prompt_type == PromptType.DATA_ANALYSIS: service = AsyncDataAnalysisService() diff --git a/llm_be/chat_backend/consumers_graph.py b/llm_be/chat_backend/consumers_graph.py index 2f6102c..10329e7 100644 --- a/llm_be/chat_backend/consumers_graph.py +++ b/llm_be/chat_backend/consumers_graph.py @@ -18,6 +18,14 @@ from .models import Conversation, Prompt, PromptMetric, DocumentWorkspace, Custo from .serializers import PromptSerializer from .services.llm_service import AsyncLLMService from .services.rag_services import AsyncRAGService +from .services.chat_tenant_scope import ( + ChatTenantScopeError, + asgi_user_or_none, + create_conversation_for_user, + get_workspace_for_scope, + resolve_chat_company_scope, + resolve_chat_user as resolve_chat_user_sync, +) from .services.title_generator import title_generator from .services.moderation_classifier import moderation_classifier, ModerationLabel from .services.prompt_classifier.prompt_classifier import PromptClassifier, PromptType @@ -41,29 +49,22 @@ PROMPT_CLASSIFIER = PromptClassifier() # --- Database Helpers (Reused) --- @database_sync_to_async -def create_conversation(prompt, email, title): - conversation = Conversation.objects.create(title=title) - user = CustomUser.objects.get(email=email) - conversation.user_id = user.id - conversation.save() - return conversation.id +def create_conversation(prompt, email, title, user=None): + if user is None: + user = CustomUser.objects.get(email=email) + return create_conversation_for_user(user, title) @database_sync_to_async -def resolve_chat_user(email=None, conversation_id=None): - if email: - user = CustomUser.objects.filter(email__iexact=email).first() - if user: - return user - if conversation_id: - conversation = ( - Conversation.objects.select_related("user") - .filter(id=conversation_id) - .first() - ) - if conversation and conversation.user_id: - return conversation.user - return None +def resolve_chat_user( + email=None, conversation_id=None, token=None, authenticated_user=None +): + return resolve_chat_user_sync( + email=email, + token=token, + authenticated_user=authenticated_user, + conversation_id=conversation_id, + ) @database_sync_to_async @@ -78,9 +79,19 @@ def enforce_feature_gate(user, feature): assert_feature_allowed(user, feature) @database_sync_to_async -def get_workspace(conversation_id): - conversation = Conversation.objects.get(id=conversation_id) - return DocumentWorkspace.objects.get(company=conversation.user.company) +def get_workspace(conversation_id, user=None): + if user is None: + raise ChatTenantScopeError( + "Authenticated chat user is required.", + code="user_not_found", + ) + scope = resolve_chat_company_scope(user, conversation_id) + return get_workspace_for_scope(scope) + + +@database_sync_to_async +def resolve_tenant_scope(user, conversation_id=None): + return resolve_chat_company_scope(user, conversation_id) @database_sync_to_async def get_messages(conversation_id, prompt, file_string: str = None, file_type: str = ""): @@ -284,7 +295,8 @@ async def generation_node(state: ChatState) -> ChatState: if prompt_type == PromptType.RAG: service = AsyncRAGService() - workspace = await get_workspace(conversation_id) + chat_user = state.get("chat_user") + workspace = await get_workspace(conversation_id, user=chat_user) generator = service.generate_response(messages, prompt_instance.message, workspace) return {"response_generator": generator} @@ -349,6 +361,7 @@ class ChatConsumerGraph(AsyncWebsocketConsumer): message = normalize_user_message(data.get("message", None)) conversation_id = data.get("conversation_id", None) email = data.get("email", None) + token = data.get("token") or data.get("access") file = data.get("file", None) file_type = data.get("fileType", "") @@ -365,7 +378,10 @@ class ChatConsumerGraph(AsyncWebsocketConsumer): return chat_user = await resolve_chat_user( - email=email, conversation_id=conversation_id + email=email, + conversation_id=conversation_id, + token=token, + authenticated_user=asgi_user_or_none(self.scope.get("user")), ) if chat_user is None: await self.send_json_message( @@ -396,7 +412,37 @@ class ChatConsumerGraph(AsyncWebsocketConsumer): if not conversation_id: title = await title_generator.generate_async(message) - conversation_id = await create_conversation(message, email, title) + conversation_id = await create_conversation( + message, email, title, user=chat_user + ) + + try: + tenant_scope = await resolve_tenant_scope(chat_user, conversation_id) + except ChatTenantScopeError as exc: + logger.warning( + "websocket tenant validation failed conversation_id=%s user_id=%s code=%s", + conversation_id, + chat_user.id, + exc.code, + ) + await self.send_json_message( + json.dumps( + { + "type": "error", + "code": exc.code, + "content": exc.message, + } + ) + ) + return + + logger.info( + "chat_scope_validated conversation_id=%s user_id=%s company_id=%s workspace_id=%s", + tenant_scope.conversation_id, + tenant_scope.user_id, + tenant_scope.company_id, + tenant_scope.workspace_id, + ) if conversation_id: print("Conversation ID: ", conversation_id) diff --git a/llm_be/chat_backend/services/chat_tenant_scope.py b/llm_be/chat_backend/services/chat_tenant_scope.py new file mode 100644 index 0000000..aa9d9e2 --- /dev/null +++ b/llm_be/chat_backend/services/chat_tenant_scope.py @@ -0,0 +1,182 @@ +"""Immutable per-turn company/workspace scope for chat + RAG. + +Mirrors the abc_worker ChatTenantScope stove-pipe: resolve identity once, +validate conversation ownership, never derive tenant from an untrusted +conversation_id alone. +""" + +from __future__ import annotations + +from dataclasses import dataclass +from typing import Optional + +from django.contrib.auth.models import AnonymousUser +from rest_framework_simplejwt.exceptions import TokenError +from rest_framework_simplejwt.tokens import AccessToken + +from chat_backend.models import Conversation, CustomUser, DocumentWorkspace + + +class ChatTenantScopeError(Exception): + """Raised when chat tenant resolution or ownership validation fails.""" + + def __init__(self, message: str, *, code: str = "tenant_scope_denied"): + super().__init__(message) + self.message = message + self.code = code + + +@dataclass(frozen=True) +class ChatCompanyScope: + """Frozen tenant identity for one websocket turn / RAG retrieval.""" + + user_id: int + company_id: int + workspace_id: int + conversation_id: Optional[int] = None + + +def user_from_access_token(token: str) -> Optional[CustomUser]: + """Resolve an active user from a SimpleJWT access token string.""" + if not token or not isinstance(token, str): + return None + try: + access = AccessToken(token) + user_id = access.get("user_id") + if not user_id: + return None + return CustomUser.objects.filter(id=user_id, is_active=True).first() + except TokenError: + return None + + +def resolve_chat_user( + *, + email: Optional[str] = None, + token: Optional[str] = None, + authenticated_user=None, + conversation_id: Optional[int] = None, +) -> Optional[CustomUser]: + """ + Resolve the chat principal for a websocket turn. + + Preference order: + 1. Authenticated ASGI/session user + 2. JWT access token (payload or query) + 3. Client email (legacy FE path) + + Does not fall back to conversation.user — that would bind identity to an + attacker-chosen conversation_id. ``conversation_id`` is accepted for API + compatibility but ignored for identity resolution. + """ + del conversation_id + if authenticated_user is not None and getattr( + authenticated_user, "is_authenticated", False + ): + if isinstance(authenticated_user, CustomUser): + return authenticated_user + user = CustomUser.objects.filter( + id=authenticated_user.pk, is_active=True + ).first() + if user: + return user + + token_user = user_from_access_token(token) if token else None + if token_user: + return token_user + + if email: + return CustomUser.objects.filter(email__iexact=email, is_active=True).first() + + return None + + +def resolve_chat_company_scope( + user: CustomUser, + conversation_id: Optional[int] = None, +) -> ChatCompanyScope: + """ + Build an immutable company/workspace scope for ``user``. + + When ``conversation_id`` is set, require ``conversation.user_id == user.id`` + and that the conversation owner's company matches the user's company. + """ + if user is None or not getattr(user, "id", None): + raise ChatTenantScopeError( + "Authenticated chat user is required.", + code="user_not_found", + ) + if not getattr(user, "company_id", None): + raise ChatTenantScopeError( + "User is not attached to a company workspace.", + code="company_missing", + ) + + if conversation_id is not None: + conversation = ( + Conversation.objects.select_related("user") + .filter(id=conversation_id, deleted=False) + .first() + ) + if conversation is None: + raise ChatTenantScopeError( + "Conversation was not found.", + code="conversation_not_found", + ) + if conversation.user_id != user.id: + raise ChatTenantScopeError( + "Conversation does not belong to the authenticated user.", + code="conversation_forbidden", + ) + owner_company_id = getattr(conversation.user, "company_id", None) + if owner_company_id != user.company_id: + raise ChatTenantScopeError( + "Conversation company does not match the authenticated user.", + code="conversation_forbidden", + ) + + workspace = ( + DocumentWorkspace.objects.filter(company_id=user.company_id) + .order_by("id") + .first() + ) + if workspace is None: + raise ChatTenantScopeError( + "No document workspace exists for this company.", + code="workspace_missing", + ) + + return ChatCompanyScope( + user_id=user.id, + company_id=user.company_id, + workspace_id=workspace.id, + conversation_id=conversation_id, + ) + + +def create_conversation_for_user(user: CustomUser, title: str) -> int: + """Create a conversation owned by ``user`` and return its id.""" + conversation = Conversation.objects.create(title=title, user=user) + return conversation.id + + +def get_workspace_for_scope(scope: ChatCompanyScope) -> DocumentWorkspace: + """Load workspace rows only when they match the frozen scope keys.""" + try: + return DocumentWorkspace.objects.get( + id=scope.workspace_id, company_id=scope.company_id + ) + except DocumentWorkspace.DoesNotExist as exc: + raise ChatTenantScopeError( + "Scoped document workspace was not found.", + code="workspace_missing", + ) from exc + + +def asgi_user_or_none(scope_user): + """Return an authenticated user from Channels scope, else None.""" + if scope_user is None or isinstance(scope_user, AnonymousUser): + return None + if getattr(scope_user, "is_authenticated", False): + return scope_user + return None diff --git a/llm_be/chat_backend/services/rag_services.py b/llm_be/chat_backend/services/rag_services.py index 94bd645..01d6dd4 100644 --- a/llm_be/chat_backend/services/rag_services.py +++ b/llm_be/chat_backend/services/rag_services.py @@ -118,6 +118,7 @@ class RAGService(BaseService): metadata={ "source": doc.file.name, "workspace_id": doc.workspace_id, + "company_id": doc.workspace.company_id, "document_id": doc.id, }, ) @@ -132,10 +133,11 @@ class RAGService(BaseService): def ingest_documents(self, workspace: DocumentWorkspace | None = None) -> None: """Ingest documents from a workspace into the vector store.""" print(f"Getting the Document via the workspace: {workspace}") + qs = Document.objects.select_related("workspace") if workspace: - documents = [doc for doc in Document.objects.filter(workspace=workspace)] + documents = list(qs.filter(workspace=workspace)) else: - documents = [doc for doc in Document.objects.all()] + documents = list(qs.all()) print(f"Processing the documents : {documents}") self._prepare_documents(documents) @@ -192,9 +194,17 @@ class RAGService(BaseService): tmp_created = self._materialize_file_field(file_ref) file_path = tmp_created + company_id = None + if ws_id is not None: + company_id = ( + DocumentWorkspace.objects.filter(id=ws_id) + .values_list("company_id", flat=True) + .first() + ) metadata = { "source": original_name, "workspace_id": ws_id, + "company_id": company_id, "original_filename": original_name, "file_path": original_name, } @@ -220,6 +230,17 @@ class RAGService(BaseService): self.vector_store.persist() return results + def _workspace_filter(self, workspace: DocumentWorkspace) -> Dict[str, Any]: + """Build a fail-closed Chroma metadata filter for one workspace. + + ``company_id`` is written on ingest for defense-in-depth / future dual + filters, but retrieval keys on ``workspace_id`` so older vectors without + ``company_id`` metadata still match after deploy. + """ + if workspace is None or getattr(workspace, "id", None) is None: + raise ValueError("workspace is required for RAG retrieval") + return {"workspace_id": workspace.id} + class SyncRAGService(RAGService): """Synchronous RAG service implementation.""" @@ -265,9 +286,12 @@ class SyncRAGService(RAGService): def _retriever_with_history(self, input_dict: Dict[str, Any]) -> str: """Retrieve documents considering conversation history.""" query = input_dict["query"] - conversation = input_dict["conversation"] + workspace = input_dict.get("workspace") + if workspace is None: + conversation = input_dict.get("conversation") + workspace = getattr(conversation, "workspace", None) - relevant_docs = self.search_documents(query, conversation.workspace) + relevant_docs = self.search_documents(query, workspace) if not relevant_docs: print("didn't find any relevant docs") return relevant_docs @@ -277,11 +301,9 @@ class SyncRAGService(RAGService): def search_documents( self, query: str, workspace: Optional[DocumentWorkspace] = None, k: int = 4 ) -> List[Document]: - """Search relevant documents from the vector store.""" - filter_dict = {} - if workspace: - filter_dict["workspace_id"] = workspace.id - search_kwargs = {"k": k, "filter": filter_dict if filter_dict else None} + """Search relevant documents from the vector store (workspace required).""" + filter_dict = self._workspace_filter(workspace) + search_kwargs = {"k": k, "filter": filter_dict} print(f"search_kwargs: {search_kwargs}") retriever = self.vector_store.as_retriever( search_type="similarity", @@ -358,17 +380,14 @@ class AsyncRAGService(RAGService): async def search_documents( self, query: str, workspace: Optional[DocumentWorkspace] = None, k: int = 4 ) -> List[Document]: - """Search relevant documents from the vector store.""" - filter_dict = {} + """Search relevant documents from the vector store (workspace required).""" + filter_dict = self._workspace_filter(workspace) print(f"Do we have a workspace: {workspace}") - if workspace: - filter_dict["workspace_id"] = workspace.id - search_kwargs = {"k": k, "filter": filter_dict if filter_dict else None} - print(f"search_kwargs: {search_kwargs}") + print(f"search_kwargs: {{'k': {k}, 'filter': {filter_dict}}}") retriever = self.vector_store.as_retriever( search_type="mmr", - search_kwargs={"k": k, "filter": filter_dict if filter_dict else None}, + search_kwargs={"k": k, "filter": filter_dict}, ) return await retriever.aget_relevant_documents(query) @@ -380,6 +399,8 @@ class AsyncRAGService(RAGService): **kwargs, ) -> AsyncGenerator[str, None]: """Generate response with streaming support.""" + if workspace is None: + raise ValueError("workspace is required for RAG generation") chain_input = { "query": query, "conversation": conversation, diff --git a/llm_be/chat_backend/tests/test_chat_tenant_scope.py b/llm_be/chat_backend/tests/test_chat_tenant_scope.py new file mode 100644 index 0000000..848b1d2 --- /dev/null +++ b/llm_be/chat_backend/tests/test_chat_tenant_scope.py @@ -0,0 +1,64 @@ +from django.test import TestCase +from rest_framework_simplejwt.tokens import RefreshToken + +from chat_backend.services.chat_tenant_scope import ( + ChatCompanyScope, + ChatTenantScopeError, + resolve_chat_company_scope, + resolve_chat_user, + user_from_access_token, +) + +from .factories import make_company, make_conversation, make_user, make_workspace + + +class ChatTenantScopeTestCase(TestCase): + def setUp(self): + self.company = make_company() + self.user = make_user(company=self.company) + self.workspace = make_workspace(self.company) + self.conversation = make_conversation(user=self.user) + + def test_resolve_scope_for_owned_conversation(self): + scope = resolve_chat_company_scope(self.user, self.conversation.id) + + self.assertEqual( + scope, + ChatCompanyScope( + user_id=self.user.id, + company_id=self.company.id, + workspace_id=self.workspace.id, + conversation_id=self.conversation.id, + ), + ) + + def test_resolve_scope_rejects_cross_user_conversation(self): + other = make_user( + email="other-tenant@example.com", company=make_company("OtherCo") + ) + foreign = make_conversation(user=other) + make_workspace(other.company) + + with self.assertRaises(ChatTenantScopeError) as ctx: + resolve_chat_company_scope(self.user, foreign.id) + + self.assertEqual(ctx.exception.code, "conversation_forbidden") + + def test_resolve_chat_user_does_not_bind_identity_to_conversation(self): + other = make_user(company=make_company("VictimCo"), email="victim@example.com") + foreign = make_conversation(user=other) + + resolved = resolve_chat_user(conversation_id=foreign.id) + + self.assertIsNone(resolved) + + def test_resolve_chat_user_prefers_jwt_over_email(self): + token = str(RefreshToken.for_user(self.user).access_token) + spoof = make_user(company=make_company("Spoof"), email="spoof@example.com") + + resolved = resolve_chat_user(email=spoof.email, token=token) + + self.assertEqual(resolved.id, self.user.id) + + def test_user_from_access_token_rejects_garbage(self): + self.assertIsNone(user_from_access_token("not-a-jwt")) diff --git a/llm_be/chat_backend/tests/test_consumers.py b/llm_be/chat_backend/tests/test_consumers.py index 52e23f4..fd03501 100644 --- a/llm_be/chat_backend/tests/test_consumers.py +++ b/llm_be/chat_backend/tests/test_consumers.py @@ -47,10 +47,35 @@ class DatabaseHelperTestCase(TransactionTestCase): @parameterized.expand([("websocket", consumers), ("langgraph", consumers_graph)]) async def test_get_workspace(self, _name, module): - workspace = await module.get_workspace(self.conversation.id) + workspace = await module.get_workspace( + self.conversation.id, user=self.user + ) self.assertEqual(workspace.id, self.workspace.id) + @parameterized.expand([("websocket", consumers), ("langgraph", consumers_graph)]) + async def test_get_workspace_rejects_other_users_conversation(self, _name, module): + other_company = await sync_to_async(make_company)("OtherCo") + other = await sync_to_async(make_user)( + email="other-tenant@example.com", company=other_company + ) + other_conversation = await sync_to_async(make_conversation)(user=other) + await sync_to_async(make_workspace)(other_company) + + with self.assertRaises(consumers.ChatTenantScopeError): + await module.get_workspace(other_conversation.id, user=self.user) + + @parameterized.expand([("websocket", consumers), ("langgraph", consumers_graph)]) + async def test_resolve_tenant_scope_binds_company_and_workspace( + self, _name, module + ): + scope = await module.resolve_tenant_scope(self.user, self.conversation.id) + + self.assertEqual(scope.user_id, self.user.id) + self.assertEqual(scope.company_id, self.company.id) + self.assertEqual(scope.workspace_id, self.workspace.id) + self.assertEqual(scope.conversation_id, self.conversation.id) + @parameterized.expand([("websocket", consumers), ("langgraph", consumers_graph)]) async def test_get_messages_stores_prompt_and_returns_history(self, _name, module): messages, prompt_instance = await module.get_messages( @@ -210,6 +235,7 @@ class GraphNodeTestCase(TransactionTestCase): "response_generator": None, "error": None, "model_name": "Turbo", + "chat_user": self.user, } state.update(overrides) return state diff --git a/llm_be/chat_backend/tests/test_services_rag.py b/llm_be/chat_backend/tests/test_services_rag.py index a14fbdb..03839aa 100644 --- a/llm_be/chat_backend/tests/test_services_rag.py +++ b/llm_be/chat_backend/tests/test_services_rag.py @@ -62,7 +62,8 @@ class RAGServiceTestCase(TransactionTestCase): self.addCleanup(reset_singletons) self.service = AsyncRAGService() - self.workspace = make_workspace(make_company()) + self.company = make_company() + self.workspace = make_workspace(self.company) def _patch(self, target): patcher = mock.patch(target) @@ -137,6 +138,7 @@ class RAGServiceTestCase(TransactionTestCase): added = self.service.vector_store.add_documents.call_args[0][0] self.assertIn("ingest me", added[0].page_content) self.assertEqual(added[0].metadata["workspace_id"], self.workspace.id) + self.assertEqual(added[0].metadata["company_id"], self.company.id) self.assertEqual(added[0].metadata["document_id"], document.id) def test_ingest_documents_deletes_the_materialized_temp_file(self): @@ -229,15 +231,11 @@ class RAGServiceTestCase(TransactionTestCase): search_kwargs={"k": 2, "filter": {"workspace_id": self.workspace.id}}, ) - async def test_search_documents_without_workspace_has_no_filter(self): - retriever = self.service.vector_store.as_retriever.return_value - retriever.aget_relevant_documents = mock.AsyncMock(return_value=[]) + async def test_search_documents_without_workspace_fails_closed(self): + with self.assertRaises(ValueError): + await self.service.search_documents("revenue") - await self.service.search_documents("revenue") - - self.service.vector_store.as_retriever.assert_called_with( - search_type="mmr", search_kwargs={"k": 4, "filter": None} - ) + self.service.vector_store.as_retriever.assert_not_called() async def test_format_history_labels_speakers(self): history = await self.service._format_history( diff --git a/llm_be/chat_backend/tests/test_views_conversations.py b/llm_be/chat_backend/tests/test_views_conversations.py index ef20fee..182638b 100644 --- a/llm_be/chat_backend/tests/test_views_conversations.py +++ b/llm_be/chat_backend/tests/test_views_conversations.py @@ -87,6 +87,16 @@ class ConversationDetailViewTestCase(APITestCase): [("hello", True), ("hi there", False)], ) + def test_get_rejects_other_users_conversation(self): + other = make_user(email="other-tenant@example.com", company=make_company("Other")) + foreign = make_conversation(user=other) + make_prompt(foreign, message="secret") + + response = self.client.get(self.url, {"conversation_id": foreign.id}) + + self.assertEqual(response.status_code, status.HTTP_404_NOT_FOUND) + self.assertEqual(Prompt.objects.filter(conversation=foreign).count(), 1) + def test_post_stores_assistant_prompt(self): response = self.client.post( self.url, diff --git a/llm_be/chat_backend/views.py b/llm_be/chat_backend/views.py index 1c23dfb..9edef51 100644 --- a/llm_be/chat_backend/views.py +++ b/llm_be/chat_backend/views.py @@ -502,7 +502,16 @@ class ConversationPreferences(APIView): class ConversationDetailView(APIView): def get(self, request, format="json"): conversation_id = request.query_params.get("conversation_id") - prompts = Prompt.objects.filter(conversation__id=conversation_id) + if not Conversation.objects.filter( + id=conversation_id, user=request.user, deleted=False + ).exists(): + return Response( + {"detail": "Conversation not found."}, + status=status.HTTP_404_NOT_FOUND, + ) + prompts = Prompt.objects.filter( + conversation__id=conversation_id, conversation__user=request.user + ) serailzer = PromptSerializer(prompts, many=True) return Response(serailzer.data, status=status.HTTP_200_OK) @@ -525,7 +534,9 @@ class ConversationDetailView(APIView): is_user = bool(request.data.get("is_user")) try: - conversation = Conversation.objects.get(id=conversation_id) + conversation = Conversation.objects.get( + id=conversation_id, user=request.user, deleted=False + ) # add the prompt to the conversation serializer = PromptSerializer(