Compare commits
1
Commits
master
...
35ffd590ba
| Author | SHA1 | Date | |
|---|---|---|---|
|
|
35ffd590ba |
@@ -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",
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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"
|
||||
),
|
||||
),
|
||||
),
|
||||
]
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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]
|
||||
|
||||
@@ -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):
|
||||
|
||||
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user