diff --git a/llm_be/chat_backend/admin.py b/llm_be/chat_backend/admin.py index dc86700..54676fb 100644 --- a/llm_be/chat_backend/admin.py +++ b/llm_be/chat_backend/admin.py @@ -68,6 +68,7 @@ class CustomUserAdmin(admin.ModelAdmin): "has_usable_password", "deleted", "has_signed_tos", + "use_conversation_context", "last_login", "slug", "get_set_password_url", diff --git a/llm_be/chat_backend/consumers.py b/llm_be/chat_backend/consumers.py index 74cbbaa..be63e77 100644 --- a/llm_be/chat_backend/consumers.py +++ b/llm_be/chat_backend/consumers.py @@ -508,7 +508,12 @@ class ChatConsumerAgain(AsyncWebsocketConsumer): ) await emit_status("refining") return service.generate_response( - messages, prompt_instance.message, workspace + messages, + prompt_instance.message, + workspace, + use_conversation_context=bool( + getattr(chat_user, "use_conversation_context", False) + ), ) elif prompt_type == PromptType.DATA_ANALYSIS: @@ -565,6 +570,9 @@ class ChatConsumerAgain(AsyncWebsocketConsumer): messages=messages, model_name=input_dict.get("model_name"), conversation_id=conversation_id, + use_conversation_context=bool( + getattr(chat_user, "use_conversation_context", False) + ), ) if grounded.error: return grounded.error diff --git a/llm_be/chat_backend/consumers_graph.py b/llm_be/chat_backend/consumers_graph.py index 7bcf195..c574c3f 100644 --- a/llm_be/chat_backend/consumers_graph.py +++ b/llm_be/chat_backend/consumers_graph.py @@ -305,7 +305,14 @@ async def generation_node(state: ChatState) -> ChatState: service = AsyncRAGService() workspace = await get_workspace(conversation_id, user=chat_user) await emit_status("retrieving_docs") - generator = service.generate_response(messages, prompt_instance.message, workspace) + generator = service.generate_response( + messages, + prompt_instance.message, + workspace, + use_conversation_context=bool( + getattr(chat_user, "use_conversation_context", False) + ), + ) await emit_status("refining") return {"response_generator": generator} @@ -320,11 +327,15 @@ async def generation_node(state: ChatState) -> ChatState: else: # GENERAL_CHAT / SEARCH / UNKNOWN — always-on grounding (#62). # FAST selects a smaller model; it no longer skips search. + chat_user = state.get("chat_user") grounded = await prepare_grounded_chat( message=state["message"], messages=messages, model_name=state.get("model_name"), conversation_id=conversation_id, + use_conversation_context=bool( + getattr(chat_user, "use_conversation_context", False) + ), ) if grounded.error: return { diff --git a/llm_be/chat_backend/migrations/0034_customuser_use_conversation_context.py b/llm_be/chat_backend/migrations/0034_customuser_use_conversation_context.py new file mode 100644 index 0000000..8822e0a --- /dev/null +++ b/llm_be/chat_backend/migrations/0034_customuser_use_conversation_context.py @@ -0,0 +1,22 @@ +from django.db import migrations, models + + +class Migration(migrations.Migration): + + dependencies = [ + ("chat_backend", "0033_agentrun_agentstep"), + ] + + operations = [ + migrations.AddField( + model_name="customuser", + name="use_conversation_context", + field=models.BooleanField( + default=False, + help_text=( + "When enabled, prior turns in the conversation are used as " + "LLM/RAG context for a more tailored experience" + ), + ), + ), + ] diff --git a/llm_be/chat_backend/models.py b/llm_be/chat_backend/models.py index 9cf359f..c39941e 100644 --- a/llm_be/chat_backend/models.py +++ b/llm_be/chat_backend/models.py @@ -73,6 +73,13 @@ class CustomUser(AbstractUser): conversation_order = models.BooleanField( default=True, help_text="How the conversations should display" ) + use_conversation_context = models.BooleanField( + default=False, + help_text=( + "When enabled, prior turns in the conversation are used as " + "LLM/RAG context for a more tailored experience" + ), + ) def get_set_password_url(self): from django.conf import settings diff --git a/llm_be/chat_backend/services/grounded_chat.py b/llm_be/chat_backend/services/grounded_chat.py index 2b99097..0c5718b 100644 --- a/llm_be/chat_backend/services/grounded_chat.py +++ b/llm_be/chat_backend/services/grounded_chat.py @@ -78,6 +78,7 @@ async def prepare_grounded_chat( model_name: str | None, conversation_id: int, on_status: StatusCallback | None = None, + use_conversation_context: bool = False, ) -> GroundedTurnResult: """Decide grounding, retrieve, and return an AsyncLLMService generator. @@ -86,14 +87,18 @@ async def prepare_grounded_chat( ``on_status(stage, detail)`` emits live activity frames (#96) when provided; otherwise uses the context-bound emitter from :mod:`status_context`. + + ``use_conversation_context`` gates prior-turn history in the LLM prompt + (#33); default off (opt-in). """ + gen_kwargs = {"use_conversation_context": use_conversation_context} internet = getattr(settings, "ALLOW_INTERNET_ACCESS", False) if not internet: await _emit(on_status, "refining") service = build_chat_service(model_name=model_name, grounded=False) return GroundedTurnResult( generator=service.generate_response( - messages, message, conversation_id + messages, message, conversation_id, **gen_kwargs ), model_name=service.model_name, ) @@ -106,7 +111,7 @@ async def prepare_grounded_chat( service = build_chat_service(model_name=model_name, grounded=False) return GroundedTurnResult( generator=service.generate_response( - messages, message, conversation_id + messages, message, conversation_id, **gen_kwargs ), decision=decision, model_name=service.model_name, @@ -149,7 +154,9 @@ async def prepare_grounded_chat( sources_block=sources_block, ) return GroundedTurnResult( - generator=service.generate_response(messages, message, conversation_id), + generator=service.generate_response( + messages, message, conversation_id, **gen_kwargs + ), citations=_citations_from_results(results), grounded=True, decision=decision, diff --git a/llm_be/chat_backend/services/llm_service.py b/llm_be/chat_backend/services/llm_service.py index 337e74d..c473a38 100644 --- a/llm_be/chat_backend/services/llm_service.py +++ b/llm_be/chat_backend/services/llm_service.py @@ -173,35 +173,42 @@ Response:""" block (when grounded) is never trimmed. """ sources = self.sources_block or kwargs.get("sources_block", "") or "" - reserved = ( - estimate_tokens(ASSISTANT_SYSTEM_PROMPT) - + estimate_tokens(GROUNDED_ANSWER_INSTRUCTIONS if self.grounded else "") - + estimate_tokens(sources) - + estimate_tokens(query) - + 256 # response headroom / instructions + use_conversation_context = bool( + kwargs.get("use_conversation_context", False) ) - # Leave ~40% of ctx for the completion. - history_budget = max(512, int(self.num_ctx * 0.55) - reserved) - # Drop any prior "Search Results:" blobs from history — sources are - # passed separately now so we don't double-inject. - clean = [ - m - for m in conversation - if not ( - getattr(m, "type", "") == "human" - and str(getattr(m, "content", "")).startswith("Search Results:") + history_text = "" + if use_conversation_context: + reserved = ( + estimate_tokens(ASSISTANT_SYSTEM_PROMPT) + + estimate_tokens( + GROUNDED_ANSWER_INSTRUCTIONS if self.grounded else "" + ) + + estimate_tokens(sources) + + estimate_tokens(query) + + 256 # response headroom / instructions ) - and not ( - getattr(m, "type", "") == "human" - and str(getattr(m, "content", "")).startswith("Live sources:") + # Leave ~40% of ctx for the completion. + history_budget = max(512, int(self.num_ctx * 0.55) - reserved) + # Drop any prior "Search Results:" blobs from history — sources are + # passed separately now so we don't double-inject. + clean = [ + m + for m in conversation + if not ( + getattr(m, "type", "") == "human" + and str(getattr(m, "content", "")).startswith("Search Results:") + ) + and not ( + getattr(m, "type", "") == "human" + and str(getattr(m, "content", "")).startswith("Live sources:") + ) + ] + # Exclude the latest user turn from history (it's in {query}). + prior = clean[:-1] if clean else [] + windowed = window_history( + prior, budget_tokens=history_budget, reserved_tokens=0 ) - ] - # Exclude the latest user turn from history (it's in {query}). - prior = clean[:-1] if clean else [] - windowed = window_history( - prior, budget_tokens=history_budget, reserved_tokens=0 - ) - history_text = format_history(windowed) + history_text = format_history(windowed) chain_input = { "query": query, diff --git a/llm_be/chat_backend/services/rag_services.py b/llm_be/chat_backend/services/rag_services.py index ce82d1c..960753e 100644 --- a/llm_be/chat_backend/services/rag_services.py +++ b/llm_be/chat_backend/services/rag_services.py @@ -490,11 +490,19 @@ class AsyncRAGService(RAGService): """Generate response with streaming support.""" if workspace is None: raise ValueError("workspace is required for RAG generation") + use_conversation_context = bool( + kwargs.get("use_conversation_context", False) + ) + recent = ( + await self._format_history(conversation) + if use_conversation_context + else "" + ) chain_input = { "query": query, "conversation": conversation, "workspace": workspace, - "recent_conversation": await self._format_history(conversation), + "recent_conversation": recent, } async for chunk in self.rag_chain.astream(chain_input): diff --git a/llm_be/chat_backend/tests/test_models.py b/llm_be/chat_backend/tests/test_models.py index 50fffee..eef159d 100644 --- a/llm_be/chat_backend/tests/test_models.py +++ b/llm_be/chat_backend/tests/test_models.py @@ -81,6 +81,7 @@ class CompanyAndUserTestCase(TestCase): self.assertFalse(user.deleted) self.assertFalse(user.has_signed_tos) self.assertTrue(user.conversation_order) + self.assertFalse(user.use_conversation_context) class ConversationAndPromptTestCase(TestCase): diff --git a/llm_be/chat_backend/tests/test_services_llm.py b/llm_be/chat_backend/tests/test_services_llm.py index 5e6f438..d085ff9 100644 --- a/llm_be/chat_backend/tests/test_services_llm.py +++ b/llm_be/chat_backend/tests/test_services_llm.py @@ -25,7 +25,10 @@ class AsyncLLMServiceTestCase(SimpleTestCase): chunks = [ chunk async for chunk in self.service.generate_response( - conversation(1), "hello", conversation_id=1 + conversation(1), + "hello", + conversation_id=1, + use_conversation_context=True, ) ] @@ -35,7 +38,9 @@ class AsyncLLMServiceTestCase(SimpleTestCase): self.service.conversation_chain = FakeChain(chunks=["ok"]) messages = conversation(4) # 8 messages - async for _ in self.service.generate_response(messages, "latest", 1): + async for _ in self.service.generate_response( + messages, "latest", 1, use_conversation_context=True + ): pass payload = self.service.conversation_chain.calls[0] @@ -47,11 +52,33 @@ class AsyncLLMServiceTestCase(SimpleTestCase): # 8 messages → drop last (query) → 7 prior lines max in window. self.assertLessEqual(len(payload["history"].splitlines()), 7) + async def test_generate_response_skips_history_when_context_disabled(self): + self.service.conversation_chain = FakeChain(chunks=["ok"]) + messages = conversation(4) + + async for _ in self.service.generate_response( + messages, "latest", 1, use_conversation_context=False + ): + pass + + payload = self.service.conversation_chain.calls[0] + self.assertEqual(payload["history"], "") + + async def test_generate_response_defaults_to_no_history(self): + self.service.conversation_chain = FakeChain(chunks=["ok"]) + + async for _ in self.service.generate_response(conversation(2), "q", 1): + pass + + self.assertEqual(self.service.conversation_chain.calls[0]["history"], "") + async def test_grounded_service_includes_sources(self): service = AsyncLLMService(grounded=True, sources_block='[1] "T" — x.com — undated') service.conversation_chain = FakeChain(chunks=["ok"]) - async for _ in service.generate_response([], "q", 1): + async for _ in service.generate_response( + [], "q", 1, use_conversation_context=True + ): pass payload = service.conversation_chain.calls[0] diff --git a/llm_be/chat_backend/tests/test_services_rag.py b/llm_be/chat_backend/tests/test_services_rag.py index 0c94331..3220f64 100644 --- a/llm_be/chat_backend/tests/test_services_rag.py +++ b/llm_be/chat_backend/tests/test_services_rag.py @@ -335,7 +335,10 @@ class RAGServiceTestCase(TransactionTestCase): chunks = [ chunk async for chunk in self.service.generate_response( - conversation, "what is our policy?", self.workspace + conversation, + "what is our policy?", + self.workspace, + use_conversation_context=True, ) ] @@ -345,6 +348,25 @@ class RAGServiceTestCase(TransactionTestCase): self.assertEqual(payload["workspace"], self.workspace) self.assertEqual(payload["recent_conversation"], "User: what is our policy?") + async def test_generate_response_skips_history_when_context_disabled(self): + self.service.rag_chain = FakeChain(chunks=["ok"]) + conversation = [ + HumanMessage(content="prior"), + AIMessage(content="answer"), + HumanMessage(content="what is our policy?"), + ] + + async for _ in self.service.generate_response( + conversation, + "what is our policy?", + self.workspace, + use_conversation_context=False, + ): + pass + + payload = self.service.rag_chain.calls[0] + self.assertEqual(payload["recent_conversation"], "") + async def test_get_documents_helper_scopes_by_workspace(self): document = await sync_to_async(self._text_document)() other_workspace = await sync_to_async(make_workspace)( diff --git a/llm_be/chat_backend/tests/test_views_conversations.py b/llm_be/chat_backend/tests/test_views_conversations.py index 182638b..90708ac 100644 --- a/llm_be/chat_backend/tests/test_views_conversations.py +++ b/llm_be/chat_backend/tests/test_views_conversations.py @@ -57,6 +57,7 @@ class ConversationPreferencesTestCase(APITestCase): self.assertEqual(response.status_code, status.HTTP_200_OK) self.assertTrue(response.data["order"]) + self.assertFalse(response.data["use_conversation_context"]) def test_post_toggles_and_persists_order(self): response = self.client.post(self.url) @@ -66,6 +67,33 @@ class ConversationPreferencesTestCase(APITestCase): self.user.refresh_from_db() self.assertFalse(self.user.conversation_order) + def test_post_sets_use_conversation_context_without_toggling_order(self): + before_order = self.user.conversation_order + + response = self.client.post( + self.url, {"use_conversation_context": True}, format="json" + ) + + self.assertEqual(response.status_code, status.HTTP_200_OK) + self.assertTrue(response.data["use_conversation_context"]) + self.assertEqual(response.data["order"], before_order) + self.user.refresh_from_db() + self.assertTrue(self.user.use_conversation_context) + self.assertEqual(self.user.conversation_order, before_order) + + def test_post_can_disable_use_conversation_context(self): + self.user.use_conversation_context = True + self.user.save(update_fields=["use_conversation_context"]) + + response = self.client.post( + self.url, {"use_conversation_context": False}, format="json" + ) + + self.assertEqual(response.status_code, status.HTTP_200_OK) + self.assertFalse(response.data["use_conversation_context"]) + self.user.refresh_from_db() + self.assertFalse(self.user.use_conversation_context) + class ConversationDetailViewTestCase(APITestCase): def setUp(self): diff --git a/llm_be/chat_backend/views.py b/llm_be/chat_backend/views.py index f9b42cd..73effa6 100644 --- a/llm_be/chat_backend/views.py +++ b/llm_be/chat_backend/views.py @@ -584,13 +584,37 @@ class ConversationsView(APIView): class ConversationPreferences(APIView): def get(self, request, format="json"): user = request.user - return Response({"order": user.conversation_order}, status=status.HTTP_200_OK) + return Response( + { + "order": user.conversation_order, + "use_conversation_context": user.use_conversation_context, + }, + status=status.HTTP_200_OK, + ) def post(self, request, format="json"): user = request.user - user.conversation_order = not user.conversation_order - user.save() - return Response({"order": user.conversation_order}, status=status.HTTP_200_OK) + data = request.data + update_fields = [] + + if "use_conversation_context" in data: + user.use_conversation_context = bool(data.get("use_conversation_context")) + update_fields.append("use_conversation_context") + + # Legacy: bare POST or POST with ``order`` toggles conversation_order. + # POST that only sets use_conversation_context leaves order unchanged. + if "order" in data or "use_conversation_context" not in data: + user.conversation_order = not user.conversation_order + update_fields.append("conversation_order") + + user.save(update_fields=update_fields) + return Response( + { + "order": user.conversation_order, + "use_conversation_context": user.use_conversation_context, + }, + status=status.HTTP_200_OK, + ) class ConversationDetailView(APIView):