Compare commits

...
1 Commits
Author SHA1 Message Date
westfarn 35ffd590ba Add opt-in use_conversation_context flag for chat history (#33).
CI / test (pull_request) Successful in 11s
Unit Tests / test (pull_request) Successful in 11s
Default-off user preference gates prior-turn history in LLM/RAG prompts;
expose via conversation_preferences GET/POST for the frontend toggle.
2026-08-04 06:26:41 -07:00
13 changed files with 213 additions and 40 deletions
+1
View File
@@ -68,6 +68,7 @@ class CustomUserAdmin(admin.ModelAdmin):
"has_usable_password", "has_usable_password",
"deleted", "deleted",
"has_signed_tos", "has_signed_tos",
"use_conversation_context",
"last_login", "last_login",
"slug", "slug",
"get_set_password_url", "get_set_password_url",
+9 -1
View File
@@ -508,7 +508,12 @@ class ChatConsumerAgain(AsyncWebsocketConsumer):
) )
await emit_status("refining") await emit_status("refining")
return service.generate_response( 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: elif prompt_type == PromptType.DATA_ANALYSIS:
@@ -565,6 +570,9 @@ class ChatConsumerAgain(AsyncWebsocketConsumer):
messages=messages, messages=messages,
model_name=input_dict.get("model_name"), model_name=input_dict.get("model_name"),
conversation_id=conversation_id, conversation_id=conversation_id,
use_conversation_context=bool(
getattr(chat_user, "use_conversation_context", False)
),
) )
if grounded.error: if grounded.error:
return grounded.error return grounded.error
+12 -1
View File
@@ -305,7 +305,14 @@ async def generation_node(state: ChatState) -> ChatState:
service = AsyncRAGService() service = AsyncRAGService()
workspace = await get_workspace(conversation_id, user=chat_user) workspace = await get_workspace(conversation_id, user=chat_user)
await emit_status("retrieving_docs") 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") await emit_status("refining")
return {"response_generator": generator} return {"response_generator": generator}
@@ -320,11 +327,15 @@ async def generation_node(state: ChatState) -> ChatState:
else: else:
# GENERAL_CHAT / SEARCH / UNKNOWN — always-on grounding (#62). # GENERAL_CHAT / SEARCH / UNKNOWN — always-on grounding (#62).
# FAST selects a smaller model; it no longer skips search. # FAST selects a smaller model; it no longer skips search.
chat_user = state.get("chat_user")
grounded = await prepare_grounded_chat( grounded = await prepare_grounded_chat(
message=state["message"], message=state["message"],
messages=messages, messages=messages,
model_name=state.get("model_name"), model_name=state.get("model_name"),
conversation_id=conversation_id, conversation_id=conversation_id,
use_conversation_context=bool(
getattr(chat_user, "use_conversation_context", False)
),
) )
if grounded.error: if grounded.error:
return { 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( conversation_order = models.BooleanField(
default=True, help_text="How the conversations should display" 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): def get_set_password_url(self):
from django.conf import settings from django.conf import settings
+10 -3
View File
@@ -78,6 +78,7 @@ async def prepare_grounded_chat(
model_name: str | None, model_name: str | None,
conversation_id: int, conversation_id: int,
on_status: StatusCallback | None = None, on_status: StatusCallback | None = None,
use_conversation_context: bool = False,
) -> GroundedTurnResult: ) -> GroundedTurnResult:
"""Decide grounding, retrieve, and return an AsyncLLMService generator. """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; ``on_status(stage, detail)`` emits live activity frames (#96) when provided;
otherwise uses the context-bound emitter from :mod:`status_context`. 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) internet = getattr(settings, "ALLOW_INTERNET_ACCESS", False)
if not internet: if not internet:
await _emit(on_status, "refining") await _emit(on_status, "refining")
service = build_chat_service(model_name=model_name, grounded=False) service = build_chat_service(model_name=model_name, grounded=False)
return GroundedTurnResult( return GroundedTurnResult(
generator=service.generate_response( generator=service.generate_response(
messages, message, conversation_id messages, message, conversation_id, **gen_kwargs
), ),
model_name=service.model_name, model_name=service.model_name,
) )
@@ -106,7 +111,7 @@ async def prepare_grounded_chat(
service = build_chat_service(model_name=model_name, grounded=False) service = build_chat_service(model_name=model_name, grounded=False)
return GroundedTurnResult( return GroundedTurnResult(
generator=service.generate_response( generator=service.generate_response(
messages, message, conversation_id messages, message, conversation_id, **gen_kwargs
), ),
decision=decision, decision=decision,
model_name=service.model_name, model_name=service.model_name,
@@ -149,7 +154,9 @@ async def prepare_grounded_chat(
sources_block=sources_block, sources_block=sources_block,
) )
return GroundedTurnResult( 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), citations=_citations_from_results(results),
grounded=True, grounded=True,
decision=decision, decision=decision,
+8 -1
View File
@@ -173,9 +173,16 @@ Response:"""
block (when grounded) is never trimmed. block (when grounded) is never trimmed.
""" """
sources = self.sources_block or kwargs.get("sources_block", "") or "" 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 = ( reserved = (
estimate_tokens(ASSISTANT_SYSTEM_PROMPT) 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(sources)
+ estimate_tokens(query) + estimate_tokens(query)
+ 256 # response headroom / instructions + 256 # response headroom / instructions
+9 -1
View File
@@ -490,11 +490,19 @@ class AsyncRAGService(RAGService):
"""Generate response with streaming support.""" """Generate response with streaming support."""
if workspace is None: if workspace is None:
raise ValueError("workspace is required for RAG generation") 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 = { chain_input = {
"query": query, "query": query,
"conversation": conversation, "conversation": conversation,
"workspace": workspace, "workspace": workspace,
"recent_conversation": await self._format_history(conversation), "recent_conversation": recent,
} }
async for chunk in self.rag_chain.astream(chain_input): 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.deleted)
self.assertFalse(user.has_signed_tos) self.assertFalse(user.has_signed_tos)
self.assertTrue(user.conversation_order) self.assertTrue(user.conversation_order)
self.assertFalse(user.use_conversation_context)
class ConversationAndPromptTestCase(TestCase): class ConversationAndPromptTestCase(TestCase):
+30 -3
View File
@@ -25,7 +25,10 @@ class AsyncLLMServiceTestCase(SimpleTestCase):
chunks = [ chunks = [
chunk chunk
async for chunk in self.service.generate_response( 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"]) self.service.conversation_chain = FakeChain(chunks=["ok"])
messages = conversation(4) # 8 messages 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 pass
payload = self.service.conversation_chain.calls[0] 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. # 8 messages → drop last (query) → 7 prior lines max in window.
self.assertLessEqual(len(payload["history"].splitlines()), 7) 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): async def test_grounded_service_includes_sources(self):
service = AsyncLLMService(grounded=True, sources_block='[1] "T" — x.com — undated') service = AsyncLLMService(grounded=True, sources_block='[1] "T" — x.com — undated')
service.conversation_chain = FakeChain(chunks=["ok"]) 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 pass
payload = service.conversation_chain.calls[0] payload = service.conversation_chain.calls[0]
+23 -1
View File
@@ -335,7 +335,10 @@ class RAGServiceTestCase(TransactionTestCase):
chunks = [ chunks = [
chunk chunk
async for chunk in self.service.generate_response( 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["workspace"], self.workspace)
self.assertEqual(payload["recent_conversation"], "User: what is our policy?") 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): async def test_get_documents_helper_scopes_by_workspace(self):
document = await sync_to_async(self._text_document)() document = await sync_to_async(self._text_document)()
other_workspace = await sync_to_async(make_workspace)( 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.assertEqual(response.status_code, status.HTTP_200_OK)
self.assertTrue(response.data["order"]) self.assertTrue(response.data["order"])
self.assertFalse(response.data["use_conversation_context"])
def test_post_toggles_and_persists_order(self): def test_post_toggles_and_persists_order(self):
response = self.client.post(self.url) response = self.client.post(self.url)
@@ -66,6 +67,33 @@ class ConversationPreferencesTestCase(APITestCase):
self.user.refresh_from_db() self.user.refresh_from_db()
self.assertFalse(self.user.conversation_order) 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): class ConversationDetailViewTestCase(APITestCase):
def setUp(self): def setUp(self):
+27 -3
View File
@@ -584,13 +584,37 @@ class ConversationsView(APIView):
class ConversationPreferences(APIView): class ConversationPreferences(APIView):
def get(self, request, format="json"): def get(self, request, format="json"):
user = request.user 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"): def post(self, request, format="json"):
user = request.user 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.conversation_order = not user.conversation_order
user.save() update_fields.append("conversation_order")
return Response({"order": user.conversation_order}, status=status.HTTP_200_OK)
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): class ConversationDetailView(APIView):