Add opt-in use_conversation_context flag for chat history (#33) (#73)
Unit Tests / test (push) Successful in 11s
Deploy Beta / unit-tests (push) Successful in 11s
Deploy Beta / docker (push) Successful in 22s
Deploy Beta / deploy-beta (push) Successful in 41s

## Summary
- Closes [#33](#33) — per-user `use_conversation_context` boolean on `CustomUser` (default `false`, opt-in).
- Migration `0034_customuser_use_conversation_context`.
- `GET/POST /api/conversation_preferences` reads/writes the flag; POST sets absolute boolean without toggling `conversation_order` when only this field is sent.
- Chat + document RAG paths skip prior-turn history when the flag is off (`AsyncLLMService`, `AsyncRAGService`, both WS consumers).

## Frontend contract
- Field: `use_conversation_context` (boolean, default false)
- Read: `GET /api/conversation_preferences` → `{ order, use_conversation_context }`
- Write: `POST /api/conversation_preferences` with `{ "use_conversation_context": true|false }`
- Tooltip copy: use previous conversations to better customize the experience
- Counterpart: [chat_web_app#66](ai_ml_operations/chat_web_app#66)

## Test plan
- [ ] `manage.py test chat_backend.tests.test_models chat_backend.tests.test_views_conversations.ConversationPreferencesTestCase chat_backend.tests.test_services_llm chat_backend.tests.test_services_rag`
- [ ] New user / unset flag → `use_conversation_context` is false; history empty in LLM/RAG prompts
- [ ] Enable via preferences API → prior turns appear in history
- [ ] Disable again → history skipped; order preference unchanged when posting only the context flagReviewed-on: #73
This commit was merged in pull request #73.
This commit is contained in:
2026-08-04 06:27:44 -07:00
parent bef151c92c
commit ab3dfffa6b
13 changed files with 213 additions and 40 deletions
+1
View File
@@ -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",
+9 -1
View File
@@ -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
+12 -1
View File
@@ -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 {
@@ -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"
),
),
),
]
+7
View File
@@ -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
+10 -3
View File
@@ -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,
+8 -1
View File
@@ -173,9 +173,16 @@ Response:"""
block (when grounded) is never trimmed.
"""
sources = self.sources_block or kwargs.get("sources_block", "") or ""
use_conversation_context = bool(
kwargs.get("use_conversation_context", False)
)
history_text = ""
if use_conversation_context:
reserved = (
estimate_tokens(ASSISTANT_SYSTEM_PROMPT)
+ estimate_tokens(GROUNDED_ANSWER_INSTRUCTIONS if self.grounded else "")
+ estimate_tokens(
GROUNDED_ANSWER_INSTRUCTIONS if self.grounded else ""
)
+ estimate_tokens(sources)
+ estimate_tokens(query)
+ 256 # response headroom / instructions
+9 -1
View File
@@ -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):
+1
View File
@@ -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):
+30 -3
View File
@@ -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]
+23 -1
View File
@@ -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)(
@@ -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):
+27 -3
View File
@@ -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
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
user.save()
return Response({"order": user.conversation_order}, status=status.HTTP_200_OK)
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):