Ship the Phase 4 accuracy eval suite with a manual Gitea workflow, emit versioned WS status frames during grounded chat, and introduce opt-in agent infrastructure (Redis/Celery, AgentRun/Step, tools, orchestrator) gated by ALLOW_AGENTIC_TASKS so default chat behaviour stays unchanged.
590 lines
22 KiB
Python
590 lines
22 KiB
Python
import json
|
|
import base64
|
|
import logging
|
|
import pandas as pd
|
|
from datetime import datetime
|
|
from typing import TypedDict, Annotated, List, Union, Dict, Any
|
|
from django.utils import timezone
|
|
from django.conf import settings
|
|
from django.core.files.base import ContentFile
|
|
from asgiref.sync import sync_to_async
|
|
from channels.generic.websocket import AsyncWebsocketConsumer
|
|
from channels.db import database_sync_to_async
|
|
from langchain_core.messages import HumanMessage, AIMessage, BaseMessage
|
|
from langgraph.graph import StateGraph, END
|
|
|
|
from .models import Conversation, Prompt, PromptMetric, DocumentWorkspace, CustomUser
|
|
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
|
|
from .services.data_analysis_service import AsyncDataAnalysisService
|
|
from .services.grounded_chat import prepare_grounded_chat
|
|
from .services.status_context import (
|
|
emit_status,
|
|
reset_status_emitter,
|
|
set_status_emitter,
|
|
)
|
|
from .services.ws_frames import citations_frame, status_frame
|
|
from chat_backend.ollama_config import ollama_model_for_role, resolve_chat_role
|
|
from .utils import (
|
|
TokenUsageCollector,
|
|
aiter_text_chunks,
|
|
extract_token_usage,
|
|
has_usable_user_prompt,
|
|
is_heartbeat_payload,
|
|
normalize_user_message,
|
|
)
|
|
from monetization.services.quotas import FeatureNotAllowed, QuotaExceeded, check_generation_allowed
|
|
|
|
logger = logging.getLogger(__name__)
|
|
|
|
CHANNEL_NAME: str = "llm_messages"
|
|
PROMPT_CLASSIFIER = PromptClassifier()
|
|
|
|
# --- Database Helpers (Reused) ---
|
|
|
|
@database_sync_to_async
|
|
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, 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
|
|
def enforce_generation_gates(user, feature="text_generation"):
|
|
return check_generation_allowed(user, feature=feature)
|
|
|
|
|
|
@database_sync_to_async
|
|
def enforce_feature_gate(user, feature):
|
|
from monetization.services.quotas import assert_feature_allowed
|
|
|
|
assert_feature_allowed(user, feature)
|
|
|
|
@database_sync_to_async
|
|
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 = ""):
|
|
messages = []
|
|
conversation = Conversation.objects.get(id=conversation_id)
|
|
|
|
serializer = PromptSerializer(
|
|
data={
|
|
"message": prompt,
|
|
"user_created": True,
|
|
"created": timezone.now(),
|
|
}
|
|
)
|
|
if serializer.is_valid(raise_exception=True):
|
|
prompt_instance = serializer.save()
|
|
prompt_instance.conversation_id = conversation.id
|
|
prompt_instance.save()
|
|
if file_string:
|
|
file_name = f"prompt_{prompt_instance.id}_data.{file_type}"
|
|
f = ContentFile(file_string, name=file_name)
|
|
prompt_instance.file.save(file_name, f)
|
|
prompt_instance.file_type = file_type
|
|
prompt_instance.save()
|
|
|
|
for prompt_obj in Prompt.objects.filter(conversation__id=conversation_id):
|
|
messages.append(
|
|
{
|
|
"content": prompt_obj.message,
|
|
"role": "user" if prompt_obj.user_created else "assistant",
|
|
"has_file": prompt_obj.file_exists(),
|
|
"file": prompt_obj.file if prompt_obj.file_exists() else None,
|
|
"file_type": prompt_obj.file_type if prompt_obj.file_exists() else None,
|
|
}
|
|
)
|
|
|
|
transformed_messages = []
|
|
for message in messages:
|
|
if message["has_file"] and message["file_type"] != None:
|
|
# Simplified handling compared to original, as we rely on services to handle files now
|
|
# But we keep the structure for context
|
|
altered_message = message["content"]
|
|
else:
|
|
altered_message = message["content"]
|
|
|
|
transformed_message = (
|
|
AIMessage(content=altered_message)
|
|
if message["role"] == "assistant"
|
|
else HumanMessage(content=altered_message)
|
|
)
|
|
transformed_messages.append(transformed_message)
|
|
|
|
return transformed_messages, prompt_instance
|
|
|
|
@database_sync_to_async
|
|
def save_generated_message(conversation_id, message, citations=None):
|
|
conversation = Conversation.objects.get(id=conversation_id)
|
|
serializer = PromptSerializer(
|
|
data={
|
|
"message": message,
|
|
"user_created": False,
|
|
"created": timezone.now(),
|
|
}
|
|
)
|
|
if serializer.is_valid():
|
|
prompt_instance = serializer.save()
|
|
prompt_instance.conversation_id = conversation.id
|
|
prompt_instance.save()
|
|
if citations is not None:
|
|
Prompt.objects.filter(pk=prompt_instance.pk).update(citations=citations)
|
|
else:
|
|
print(serializer.errors)
|
|
|
|
@database_sync_to_async
|
|
def create_prompt_metric(prompt_id, prompt, has_file, file_type, model_name, conversation_id, tokens_in=None):
|
|
prompt_metric = PromptMetric.objects.create(
|
|
prompt_id=prompt_id,
|
|
start_time=timezone.now(),
|
|
prompt_length=len(prompt),
|
|
tokens_in=tokens_in,
|
|
has_file=has_file,
|
|
file_type=file_type,
|
|
model_name=model_name,
|
|
conversation_id=conversation_id,
|
|
)
|
|
return prompt_metric
|
|
|
|
@database_sync_to_async
|
|
def finish_prompt_metric(prompt_metric, response_length, tokens_in=None, tokens_out=None):
|
|
prompt_metric.end_time = timezone.now()
|
|
prompt_metric.reponse_length = response_length
|
|
prompt_metric.event = "FINISHED"
|
|
update_fields = ["end_time", "reponse_length", "event"]
|
|
if tokens_in is not None:
|
|
prompt_metric.tokens_in = tokens_in
|
|
update_fields.append("tokens_in")
|
|
if tokens_out is not None:
|
|
prompt_metric.tokens_out = tokens_out
|
|
update_fields.append("tokens_out")
|
|
prompt_metric.save(update_fields=update_fields)
|
|
|
|
async def get_conversation_file_async(conversation_id):
|
|
try:
|
|
prompt_with_file = await Prompt.objects.filter(
|
|
conversation_id=conversation_id
|
|
).exclude(file='').order_by('created').afirst()
|
|
|
|
if prompt_with_file and prompt_with_file.file:
|
|
# Opening a DatabaseStorage file hits the DB, so read inside the thread.
|
|
file_data = await sync_to_async(lambda: prompt_with_file.file.read())()
|
|
file_type = prompt_with_file.file_type
|
|
return file_data, file_type
|
|
except Exception as e:
|
|
logger.error(f"Error retrieving file from conversation history: {e}")
|
|
return None, None
|
|
|
|
|
|
# --- LangGraph State ---
|
|
|
|
class ChatState(TypedDict):
|
|
message: str
|
|
conversation_id: int
|
|
decoded_file: Union[bytes, None]
|
|
file_type: Union[str, None]
|
|
messages: List[BaseMessage]
|
|
prompt_instance: Any # Django model instance
|
|
moderation_label: Union[ModerationLabel, None]
|
|
prompt_type: Union[PromptType, None]
|
|
response_generator: Any # AsyncGenerator or dict
|
|
error: Union[str, None]
|
|
model_name: str
|
|
chat_user: Any
|
|
citations: List[Dict[str, Any]]
|
|
resolved_model: str
|
|
|
|
|
|
# --- LangGraph Nodes ---
|
|
|
|
async def moderation_node(state: ChatState) -> ChatState:
|
|
await emit_status("moderating")
|
|
msg = state["message"]
|
|
label = await moderation_classifier.classify_async(msg)
|
|
return {"moderation_label": label}
|
|
|
|
async def classification_node(state: ChatState) -> ChatState:
|
|
if state.get("moderation_label") == ModerationLabel.NSFW:
|
|
return {"prompt_type": None}
|
|
|
|
msg = state["message"]
|
|
decoded_file = state.get("decoded_file")
|
|
|
|
prompt_type = await PROMPT_CLASSIFIER.classify_async(msg)
|
|
|
|
# Override logic
|
|
if decoded_file and (prompt_type == PromptType.DATA_ANALYSIS or 'analyze' in msg.lower() or 'data' in msg.lower()):
|
|
prompt_type = PromptType.DATA_ANALYSIS
|
|
elif decoded_file:
|
|
prompt_type = PromptType.GENERAL_CHAT
|
|
|
|
return {"prompt_type": prompt_type}
|
|
|
|
async def generation_node(state: ChatState) -> ChatState:
|
|
if state.get("moderation_label") == ModerationLabel.NSFW:
|
|
response = "Prompt has been marked as NSFW. If this is in error, submit a feedback with the prompt text."
|
|
return {"response_generator": {"type": "error", "content": response}}
|
|
|
|
prompt_type = state["prompt_type"]
|
|
messages = state["messages"]
|
|
prompt_instance = state["prompt_instance"]
|
|
conversation_id = state["conversation_id"]
|
|
decoded_file = state.get("decoded_file")
|
|
file_type = state.get("file_type")
|
|
|
|
# Feature Flag + plan gate: Image Generation
|
|
if prompt_type == PromptType.IMAGE_GENERATION:
|
|
if not getattr(settings, "ALLOW_IMAGE_GENERATION", False):
|
|
return {"response_generator": {"type": "text", "content": "Image Generation is disabled."}}
|
|
chat_user = state.get("chat_user")
|
|
if chat_user is not None:
|
|
try:
|
|
await enforce_feature_gate(chat_user, "image_generation")
|
|
except FeatureNotAllowed as exc:
|
|
return {
|
|
"response_generator": {
|
|
"type": "error",
|
|
"code": exc.code,
|
|
"content": exc.message,
|
|
}
|
|
}
|
|
return {"response_generator": {"type": "text", "content": "Image Generation is not supported at this time, but it will be soon."}}
|
|
|
|
# Feature Flag: Internet Access / always-on grounding handled below for chat.
|
|
if prompt_type == PromptType.RAG:
|
|
chat_user = state.get("chat_user")
|
|
if chat_user is not None:
|
|
try:
|
|
await enforce_feature_gate(chat_user, "rag")
|
|
except FeatureNotAllowed as exc:
|
|
return {
|
|
"response_generator": {
|
|
"type": "error",
|
|
"code": exc.code,
|
|
"content": exc.message,
|
|
}
|
|
}
|
|
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)
|
|
await emit_status("refining")
|
|
return {"response_generator": generator}
|
|
|
|
elif prompt_type == PromptType.DATA_ANALYSIS:
|
|
service = AsyncDataAnalysisService()
|
|
if not decoded_file:
|
|
return {"response_generator": {"type": "text", "content": "Please upload a file to perform data analysis."}}
|
|
await emit_status("analysing")
|
|
generator = service.generate_response(prompt_instance.message, decoded_file, file_type)
|
|
return {"response_generator": generator}
|
|
|
|
else:
|
|
# GENERAL_CHAT / SEARCH / UNKNOWN — always-on grounding (#62).
|
|
# FAST selects a smaller model; it no longer skips search.
|
|
grounded = await prepare_grounded_chat(
|
|
message=state["message"],
|
|
messages=messages,
|
|
model_name=state.get("model_name"),
|
|
conversation_id=conversation_id,
|
|
)
|
|
if grounded.error:
|
|
return {
|
|
"response_generator": grounded.error,
|
|
"citations": [],
|
|
"resolved_model": grounded.model_name or "",
|
|
}
|
|
return {
|
|
"response_generator": grounded.generator,
|
|
"citations": grounded.citations,
|
|
"resolved_model": grounded.model_name,
|
|
}
|
|
|
|
|
|
# --- LangGraph Definition ---
|
|
|
|
workflow = StateGraph(ChatState)
|
|
|
|
workflow.add_node("moderation", moderation_node)
|
|
workflow.add_node("classification", classification_node)
|
|
workflow.add_node("generation", generation_node)
|
|
|
|
workflow.set_entry_point("moderation")
|
|
|
|
workflow.add_edge("moderation", "classification")
|
|
workflow.add_edge("classification", "generation")
|
|
workflow.add_edge("generation", END)
|
|
|
|
app = workflow.compile()
|
|
|
|
|
|
# --- Consumer ---
|
|
|
|
class ChatConsumerGraph(AsyncWebsocketConsumer):
|
|
async def connect(self):
|
|
await self.accept()
|
|
|
|
async def disconnect(self, close_code):
|
|
# Connection already closing — do not call self.close() again
|
|
# (triggers ASGI 'websocket.close' after close completed).
|
|
pass
|
|
|
|
async def send_json_message(self, data_str):
|
|
try:
|
|
json.loads(data_str)
|
|
await self.send(data_str)
|
|
except (json.JSONDecodeError, TypeError):
|
|
await self.send(data_str)
|
|
|
|
async def receive(self, text_data=None, bytes_data=None):
|
|
logger.debug(f"Text Data: {text_data}")
|
|
print("Text Data: ", text_data)
|
|
if text_data:
|
|
data = json.loads(text_data)
|
|
# Keepalive frames must not create conversations or hit the LLM.
|
|
if is_heartbeat_payload(data):
|
|
return
|
|
|
|
model = data.get("modelName", "Turbo")
|
|
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", "")
|
|
|
|
if not has_usable_user_prompt(message, file):
|
|
logger.info("Ignoring websocket payload with empty message")
|
|
await self.send_json_message(
|
|
json.dumps(
|
|
{
|
|
"type": "error",
|
|
"content": "Message text cannot be empty.",
|
|
}
|
|
)
|
|
)
|
|
return
|
|
|
|
chat_user = await resolve_chat_user(
|
|
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(
|
|
json.dumps(
|
|
{
|
|
"type": "error",
|
|
"code": "user_not_found",
|
|
"content": "Unable to resolve user for this chat session.",
|
|
}
|
|
)
|
|
)
|
|
return
|
|
|
|
try:
|
|
await enforce_generation_gates(chat_user, feature="text_generation")
|
|
except (QuotaExceeded, FeatureNotAllowed) as exc:
|
|
await self.send_json_message(
|
|
json.dumps(
|
|
{
|
|
"type": "error",
|
|
"code": exc.code,
|
|
"content": exc.message,
|
|
"details": getattr(exc, "details", {}),
|
|
}
|
|
)
|
|
)
|
|
return
|
|
|
|
if not conversation_id:
|
|
title = await title_generator.generate_async(message)
|
|
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)
|
|
decoded_file = None
|
|
if file:
|
|
decoded_file = base64.b64decode(file)
|
|
if "csv" in file_type: file_type = "csv"
|
|
elif "xmlformats-officedocument" in file_type: file_type = "xlsx"
|
|
elif "word" in file_type: file_type = "docx"
|
|
elif "pdf" in file_type: file_type = "pdf"
|
|
elif "text" in file_type: file_type = "txt"
|
|
else: file_type = "Not Sure"
|
|
|
|
# Pre-fetch messages and file
|
|
messages, prompt_instance = await get_messages(
|
|
conversation_id, message, decoded_file, file_type
|
|
)
|
|
print("Messages: ", messages)
|
|
if not decoded_file:
|
|
decoded_file, file_type = await get_conversation_file_async(conversation_id)
|
|
|
|
resolved_model = ollama_model_for_role(resolve_chat_role(model))
|
|
prompt_metric = await create_prompt_metric(
|
|
prompt_instance.id,
|
|
prompt_instance.message,
|
|
True if file else False,
|
|
file_type,
|
|
resolved_model,
|
|
conversation_id,
|
|
)
|
|
|
|
# Initialize State
|
|
initial_state = {
|
|
"message": message,
|
|
"conversation_id": conversation_id,
|
|
"decoded_file": decoded_file,
|
|
"file_type": file_type,
|
|
"messages": messages,
|
|
"prompt_instance": prompt_instance,
|
|
"moderation_label": None,
|
|
"prompt_type": None,
|
|
"response_generator": None,
|
|
"error": None,
|
|
"model_name": model,
|
|
"chat_user": chat_user,
|
|
"citations": [],
|
|
"resolved_model": resolved_model,
|
|
}
|
|
print("Initial State: ", initial_state)
|
|
|
|
# Stream markers early so status frames reach the client (#96).
|
|
await self.send("CONVERSATION_ID")
|
|
await self.send(str(conversation_id))
|
|
await self.send("START_OF_THE_STREAM_ENDER_GAME_42")
|
|
|
|
async def _send_status(stage, detail=None):
|
|
await self.send_json_message(
|
|
json.dumps(status_frame(stage, detail=detail))
|
|
)
|
|
|
|
status_token = set_status_emitter(_send_status)
|
|
try:
|
|
await emit_status("queued")
|
|
# Run Graph (moderation emits moderating; grounding emits evaluating/…)
|
|
final_state = await app.ainvoke(initial_state)
|
|
print("Final State: ", final_state)
|
|
|
|
response_generator_or_dict = final_state["response_generator"]
|
|
print("Response Generator: ", response_generator_or_dict)
|
|
|
|
full_response = ""
|
|
tokens_in = tokens_out = None
|
|
|
|
if isinstance(response_generator_or_dict, dict):
|
|
content = response_generator_or_dict.get("content", "")
|
|
await self.send_json_message(
|
|
json.dumps(response_generator_or_dict)
|
|
)
|
|
full_response = content
|
|
tokens_in, tokens_out = extract_token_usage(
|
|
response_generator_or_dict
|
|
)
|
|
else:
|
|
await emit_status("writing")
|
|
usage = TokenUsageCollector()
|
|
async for chunk in aiter_text_chunks(
|
|
response_generator_or_dict, usage
|
|
):
|
|
full_response += chunk
|
|
await self.send_json_message(chunk)
|
|
tokens_in, tokens_out = usage.pair
|
|
|
|
await self.send("END_OF_THE_STREAM_ENDER_GAME_42")
|
|
|
|
citations = final_state.get("citations") or []
|
|
if citations:
|
|
await self.send_json_message(
|
|
json.dumps(citations_frame(citations))
|
|
)
|
|
|
|
final_model = final_state.get("resolved_model") or resolved_model
|
|
if final_model and final_model != prompt_metric.model_name:
|
|
prompt_metric.model_name = final_model
|
|
await database_sync_to_async(prompt_metric.save)(
|
|
update_fields=["model_name"]
|
|
)
|
|
|
|
await save_generated_message(
|
|
conversation_id, full_response, citations=citations
|
|
)
|
|
await finish_prompt_metric(
|
|
prompt_metric,
|
|
len(full_response),
|
|
tokens_in=tokens_in,
|
|
tokens_out=tokens_out,
|
|
)
|
|
finally:
|
|
reset_status_emitter(status_token)
|