## Summary - **#59** — Persist and expose Drive sync progress (`sync_total`, `sync_processed`, `sync_added`, `sync_updated`, `sync_failed`) during sync so the FE can show a progress bar. - **#60** — Paginate/search/sort `GET /documents/` with `workspace=personal|company` scope isolation. - **#94 (API)** — Add `GET /analytics/user_prompt_heatmap/?tz=` weekday × hour bins for user-entered prompts. ## Test plan - [x] `manage.py test` drive sync / documents / analytics / drive_tasks suites - [ ] Manual: trigger Drive sync and poll connections → progress fields update while pending - [ ] Manual: `GET /documents/?workspace=personal&page=1&search=…&ordering=name` - [ ] Deploy FE companion PR for #92/#93/#94 UI Closes #59 Closes #60Reviewed-on: #61
1143 lines
40 KiB
Python
1143 lines
40 KiB
Python
from channels.layers import get_channel_layer
|
||
from channels.db import database_sync_to_async
|
||
from rest_framework_simplejwt.views import TokenObtainPairView
|
||
from rest_framework_simplejwt.tokens import RefreshToken
|
||
from rest_framework import permissions, status
|
||
from .serializers import (
|
||
MyTokenObtainPairSerializer,
|
||
CustomUserSerializer,
|
||
SelfServeRegistrationSerializer,
|
||
BasicUserSerializer,
|
||
AnnouncmentSerializer,
|
||
CompanySerializer,
|
||
ConversationSerializer,
|
||
PromptSerializer,
|
||
FeedbackSerializer,
|
||
DocumentWorkspaceSerializer,
|
||
DocumentSerializer,
|
||
)
|
||
|
||
from rest_framework.views import APIView
|
||
from rest_framework.response import Response
|
||
from .models import (
|
||
CustomUser,
|
||
Announcement,
|
||
Conversation,
|
||
Prompt,
|
||
Feedback,
|
||
PromptMetric,
|
||
DocumentWorkspace,
|
||
Document,
|
||
UserAuthEvent,
|
||
)
|
||
from django.views.decorators.cache import never_cache
|
||
from django.http import JsonResponse
|
||
from .client import LlamaClient
|
||
from asgiref.sync import sync_to_async, async_to_sync
|
||
from channels.generic.websocket import AsyncWebsocketConsumer
|
||
from langchain_ollama.llms import OllamaLLM
|
||
from langchain_core.prompts import ChatPromptTemplate
|
||
from langchain_core.messages import HumanMessage, AIMessage
|
||
from langchain_classic.chains import RetrievalQA
|
||
import re
|
||
import os
|
||
from django.conf import settings
|
||
import json
|
||
import base64
|
||
import pandas as pd
|
||
import io
|
||
from chat_backend.services.assistant_identity import ASSISTANT_SYSTEM_PROMPT
|
||
|
||
from django.utils import timezone
|
||
from django.core.files import File
|
||
from django.core.files.base import ContentFile
|
||
from django.db.models import Q
|
||
import math
|
||
import datetime
|
||
import pytz
|
||
from langchain_ollama import OllamaEmbeddings
|
||
|
||
from dateutil.relativedelta import relativedelta
|
||
import requests
|
||
|
||
from .utils import last_day_of_month
|
||
from .email_tasks import (
|
||
send_feedback_email,
|
||
send_invite_email,
|
||
send_password_reset_email,
|
||
)
|
||
from finance.services.quotas import FeatureNotAllowed, assert_feature_allowed
|
||
from .services.llm_service import AsyncLLMService
|
||
from .services.rag_services import AsyncRAGService
|
||
from .services.chat_tenant_scope import (
|
||
ensure_company_workspace,
|
||
ensure_personal_workspace,
|
||
ensure_workspace_for_user,
|
||
)
|
||
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 .ollama_config import ollama_llm_kwargs, ollama_model
|
||
|
||
|
||
from langchain_classic.chains import create_retrieval_chain
|
||
from langchain_classic.chains.combine_documents import create_stuff_documents_chain
|
||
from langchain_ollama import ChatOllama
|
||
|
||
import logging
|
||
|
||
logger = logging.getLogger(__name__)
|
||
|
||
|
||
CHANNEL_NAME: str = "llm_messages"
|
||
MODEL_NAME: str = ollama_model()
|
||
|
||
|
||
def _client_ip(request):
|
||
forwarded = request.META.get("HTTP_X_FORWARDED_FOR")
|
||
if forwarded:
|
||
return forwarded.split(",")[0].strip()
|
||
return request.META.get("REMOTE_ADDR")
|
||
|
||
|
||
# Create your views here.
|
||
class CustomObtainTokenView(TokenObtainPairView):
|
||
permission_classes = (permissions.AllowAny,)
|
||
serializer_class = MyTokenObtainPairSerializer
|
||
|
||
|
||
class PublicSettingsView(APIView):
|
||
"""Non-secret feature flags / public config for the SPA."""
|
||
|
||
permission_classes = (permissions.AllowAny,)
|
||
authentication_classes = ()
|
||
|
||
def get(self, request):
|
||
from .views_oauth import oauth_public_flags
|
||
|
||
return Response(
|
||
{
|
||
"enable_account_registration": settings.ENABLE_ACCOUNT_REGISTRATION,
|
||
"oauth": oauth_public_flags(),
|
||
}
|
||
)
|
||
|
||
|
||
class CustomUserCreate(APIView):
|
||
"""Self-serve registration. Gated by ENABLE_ACCOUNT_REGISTRATION (default off)."""
|
||
|
||
permission_classes = (permissions.AllowAny,)
|
||
authentication_classes = ()
|
||
|
||
def post(self, request, format="json"):
|
||
if not settings.ENABLE_ACCOUNT_REGISTRATION:
|
||
return Response(
|
||
{"detail": "Account registration is disabled."},
|
||
status=status.HTTP_403_FORBIDDEN,
|
||
)
|
||
|
||
serializer = SelfServeRegistrationSerializer(data=request.data)
|
||
if not serializer.is_valid():
|
||
return Response(serializer.errors, status=status.HTTP_400_BAD_REQUEST)
|
||
|
||
user = serializer.save()
|
||
refresh = RefreshToken.for_user(user)
|
||
from finance.services.plans import needs_checkout
|
||
|
||
return Response(
|
||
{
|
||
"email": user.email,
|
||
"first_name": user.first_name,
|
||
"last_name": user.last_name,
|
||
"access": str(refresh.access_token),
|
||
"refresh": str(refresh),
|
||
"needs_checkout": needs_checkout(user),
|
||
},
|
||
status=status.HTTP_201_CREATED,
|
||
)
|
||
|
||
|
||
class CustomUserInvite(APIView):
|
||
http_method_names = ["post"]
|
||
|
||
def post(self, request, format="json"):
|
||
def valid_email(email_string):
|
||
regex = r"^[a-z0-9]+[\._]?[a-z0-9]+[@]\w+[.]\w+$"
|
||
if re.match(regex, email_string):
|
||
return True
|
||
else:
|
||
return False
|
||
|
||
email_to_invite = request.data["email"]
|
||
|
||
if (
|
||
len(email_to_invite) == 0
|
||
or not valid_email(email_to_invite)
|
||
or not request.user.is_company_manager
|
||
):
|
||
return Response(status=status.HTTP_400_BAD_REQUEST)
|
||
# make sure there isn't a user with this email already
|
||
existing_users = CustomUser.objects.filter(email=email_to_invite)
|
||
if len(existing_users) > 0:
|
||
return Response(status=status.HTTP_400_BAD_REQUEST)
|
||
# create the object and send the email
|
||
user = CustomUser.objects.create(
|
||
email=email_to_invite,
|
||
username=email_to_invite,
|
||
company=request.user.company,
|
||
)
|
||
|
||
send_invite_email(user.slug, email_to_invite, user=user)
|
||
UserAuthEvent.log(
|
||
user,
|
||
UserAuthEvent.EventType.INVITE_SENT,
|
||
detail=f"Invited by {request.user.email}",
|
||
ip_address=_client_ip(request),
|
||
)
|
||
|
||
return Response(status=status.HTTP_201_CREATED)
|
||
|
||
|
||
class ResetUserPassword(APIView):
|
||
"""Request a password-reset email. Invalidates the current password when sent."""
|
||
|
||
http_method_names = ["post"]
|
||
permission_classes = (permissions.AllowAny,)
|
||
authentication_classes = ()
|
||
|
||
def post(self, request, format="json"):
|
||
logger.info("Password reset requested")
|
||
email = request.data.get("email")
|
||
token = request.data.get("recaptchaToken")
|
||
if not email:
|
||
return Response(status=status.HTTP_400_BAD_REQUEST)
|
||
|
||
payload = {
|
||
"secret": settings.CAPTCHA_SECRET_KEY,
|
||
"response": token,
|
||
}
|
||
try:
|
||
captcha_response = requests.post(
|
||
"https://www.google.com/recaptcha/api/siteverify",
|
||
data=payload,
|
||
timeout=10,
|
||
)
|
||
result = captcha_response.json()
|
||
except requests.RequestException as exc:
|
||
logger.error("Captcha verification request failed: %s", exc)
|
||
return Response(status=status.HTTP_400_BAD_REQUEST)
|
||
|
||
# v2 invisible returns success only; v3 also returns a score.
|
||
if not result.get("success"):
|
||
logger.error("Captcha verification failed: %s", result)
|
||
return Response(status=status.HTTP_400_BAD_REQUEST)
|
||
score = result.get("score")
|
||
if score is not None and score < 0.5:
|
||
logger.error("Captcha score too low: %s", score)
|
||
return Response(status=status.HTTP_400_BAD_REQUEST)
|
||
|
||
user = CustomUser.objects.filter(email=email).first()
|
||
if user:
|
||
user.set_unusable_password()
|
||
user.save(update_fields=["password"])
|
||
send_password_reset_email(user.slug, email, user=user)
|
||
UserAuthEvent.log(
|
||
user,
|
||
UserAuthEvent.EventType.PASSWORD_RESET_REQUESTED,
|
||
detail="Password reset email queued",
|
||
ip_address=_client_ip(request),
|
||
)
|
||
|
||
# Always 200 after valid captcha to avoid email enumeration.
|
||
return Response(status=status.HTTP_200_OK)
|
||
|
||
|
||
class SetUserPassword(APIView):
|
||
http_method_names = ["post", "get"]
|
||
permission_classes = (permissions.AllowAny,)
|
||
authentication_classes = ()
|
||
|
||
def get(self, request, slug):
|
||
try:
|
||
user = CustomUser.objects.get(slug=slug)
|
||
except CustomUser.DoesNotExist:
|
||
return Response(status=status.HTTP_404_NOT_FOUND)
|
||
if user.has_usable_password():
|
||
return Response(status=status.HTTP_401_UNAUTHORIZED)
|
||
return Response(status=status.HTTP_200_OK)
|
||
|
||
def post(self, request, slug, format="json"):
|
||
try:
|
||
user = CustomUser.objects.get(slug=slug)
|
||
except CustomUser.DoesNotExist:
|
||
return Response(status=status.HTTP_404_NOT_FOUND)
|
||
if user.has_usable_password():
|
||
return Response(status=status.HTTP_401_UNAUTHORIZED)
|
||
|
||
password = request.data.get("password")
|
||
if not password or len(password) < 8:
|
||
return Response(
|
||
{"password": "Password must be at least 8 characters."},
|
||
status=status.HTTP_400_BAD_REQUEST,
|
||
)
|
||
|
||
user.set_password(password)
|
||
user.save()
|
||
UserAuthEvent.log(
|
||
user,
|
||
UserAuthEvent.EventType.PASSWORD_SET,
|
||
detail="Password set via email link",
|
||
ip_address=_client_ip(request),
|
||
)
|
||
return Response(status=status.HTTP_200_OK)
|
||
|
||
|
||
class CustomUserGet(APIView):
|
||
http_method_names = ["get", "head", "post"]
|
||
|
||
def get(self, request, format="json"):
|
||
|
||
email = request.user.email
|
||
username = request.user.username
|
||
user = CustomUser.objects.filter(email=email).last()
|
||
logger.info(f"Getting the user: {user}")
|
||
try:
|
||
serializer = CustomUserSerializer(user)
|
||
logger.debug(f"serializer: {serializer}")
|
||
logger.debug(serializer.data)
|
||
return Response(serializer.data, status=status.HTTP_200_OK)
|
||
except Exception as e:
|
||
logger.error(f"Exception: {e}")
|
||
return Response({}, status=status.HTTP_400_BAD_REQUEST)
|
||
|
||
|
||
class CustomUserSelfDeleteView(APIView):
|
||
"""
|
||
Soft-delete the authenticated user's own account (#34).
|
||
|
||
Frontend contract:
|
||
- Method/path: ``DELETE /api/user/``
|
||
- Optional body: ``{"refresh_token": "<current refresh>"}`` to blacklist
|
||
the active session immediately (outstanding tokens are also blacklisted).
|
||
- Success: ``200`` with ``{"detail": "Account deleted.", "deleted": true}``
|
||
- After success: clear local tokens, redirect to sign-in. Subsequent
|
||
``/token/obtain/`` and authenticated calls fail (``is_active=False``,
|
||
``deleted=True``).
|
||
- Privacy v1: soft-delete only (no anonymization / hard purge).
|
||
"""
|
||
|
||
http_method_names = ["delete", "head", "options"]
|
||
|
||
def delete(self, request, format="json"):
|
||
from chat_backend.services.account_deletion import (
|
||
AccountDeletionError,
|
||
soft_delete_account,
|
||
)
|
||
|
||
refresh_token = request.data.get("refresh_token")
|
||
try:
|
||
soft_delete_account(
|
||
request.user,
|
||
refresh_token=refresh_token,
|
||
ip_address=_client_ip(request),
|
||
)
|
||
except AccountDeletionError as exc:
|
||
return Response(
|
||
{"detail": exc.detail, "code": exc.code},
|
||
status=status.HTTP_400_BAD_REQUEST,
|
||
)
|
||
|
||
return Response(
|
||
{"detail": "Account deleted.", "deleted": True},
|
||
status=status.HTTP_200_OK,
|
||
)
|
||
|
||
|
||
class FeedbackView(APIView):
|
||
http_method_names = ["post", "get"]
|
||
|
||
def post(self, request, format="json"):
|
||
serializer = FeedbackSerializer(data=request.data)
|
||
logger.debug(request.data)
|
||
if serializer.is_valid():
|
||
|
||
feedback_obj = serializer.save()
|
||
feedback_obj.user = request.user
|
||
|
||
feedback_obj.save()
|
||
send_feedback_email(
|
||
feedback_obj.title, feedback_obj.text, user=request.user
|
||
)
|
||
return Response(serializer.data, status=status.HTTP_201_CREATED)
|
||
else:
|
||
logger.error(serializer.errors)
|
||
return Response(serializer.errors, status=status.HTTP_400_BAD_REQUEST)
|
||
|
||
def get(self, request, format="json"):
|
||
feedback_objs = Feedback.objects.filter(user=request.user)
|
||
serializer = FeedbackSerializer(feedback_objs, many=True)
|
||
return Response(serializer.data, status=status.HTTP_200_OK)
|
||
|
||
|
||
class AcknowledgeTermsOfService(APIView):
|
||
http_method_names = ["post"]
|
||
|
||
def post(self, request, format="json"):
|
||
request.user.has_signed_tos = True
|
||
request.user.save()
|
||
return Response(status=status.HTTP_200_OK)
|
||
|
||
|
||
class CompanyUsersView(APIView):
|
||
def get(self, request, format="json"):
|
||
# TODO: make sure you are a manager of that company
|
||
if request.user.is_company_manager:
|
||
users = CustomUser.objects.filter(company_id=request.user.company.id)
|
||
serializer = BasicUserSerializer(users, many=True)
|
||
return Response(serializer.data, status=status.HTTP_200_OK)
|
||
else:
|
||
return Response(status=status.HTTP_401_UNAUTHORIZED)
|
||
|
||
def post(self, request, format="json"):
|
||
if request.user.is_company_manager:
|
||
user = CustomUser.objects.get(email=request.data.get("email"))
|
||
if request.user.company_id == user.company_id:
|
||
field = request.data.get("field")
|
||
data = {}
|
||
if field == "is_active":
|
||
data.update({"is_active": not user.is_active})
|
||
elif field == "company_manager":
|
||
data.update({"is_company_manager": not user.is_company_manager})
|
||
elif field == "has_password":
|
||
if user.has_usable_password():
|
||
user.set_unusable_password()
|
||
serializer = CustomUserSerializer(user, data, partial=True)
|
||
if serializer.is_valid():
|
||
serializer.save()
|
||
return Response(status=status.HTTP_200_OK)
|
||
return Response(status=status.HTTP_400_BAD_REQUEST)
|
||
return Response(status=status.HTTP_401_UNAUTHORIZED)
|
||
|
||
def delete(self, request, format="json"):
|
||
if request.user.is_company_manager:
|
||
user = CustomUser.objects.get(email=request.data.get("email"))
|
||
if request.user.company_id == user.company_id:
|
||
user.delete()
|
||
return Response(status=status.HTTP_200_OK)
|
||
return Response(status=status.HTTP_401_UNAUTHORIZED)
|
||
|
||
|
||
class AnnouncmentView(APIView):
|
||
permission_classes = (permissions.AllowAny,)
|
||
serializer_class = AnnouncmentSerializer
|
||
|
||
def get(self, request, format="json"):
|
||
announcements = Announcement.objects.all()
|
||
serializer = AnnouncmentSerializer(announcements, many=True)
|
||
|
||
return Response(serializer.data, status=status.HTTP_200_OK)
|
||
|
||
|
||
class LogoutAndBlacklistRefreshTokenForUserView(APIView):
|
||
permission_classes = (permissions.AllowAny,)
|
||
authentication_classes = ()
|
||
|
||
def post(self, request):
|
||
try:
|
||
refresh_token = request.data["refresh_token"]
|
||
token = RefreshToken(refresh_token)
|
||
token.blacklist()
|
||
return Response(status=status.HTTP_205_RESET_CONTENT)
|
||
except Exception as e:
|
||
return Response(status=status.HTTP_400_BAD_REQUEST)
|
||
|
||
|
||
@never_cache
|
||
def is_authenticated(request):
|
||
if request.user.is_authenticated:
|
||
return JsonResponse({}, status=status.HTTP_200_OK)
|
||
return JsonResponse({}, status=status.HTTP_401_UNAUTHORIZED)
|
||
|
||
|
||
class ConversationsView(APIView):
|
||
def get(self, request, format="json"):
|
||
order = "created" if request.user.conversation_order else "-created"
|
||
conversations = Conversation.objects.filter(
|
||
user=request.user, deleted=False
|
||
).order_by(order)
|
||
serializer = ConversationSerializer(conversations, many=True)
|
||
return Response(serializer.data, status=status.HTTP_200_OK)
|
||
|
||
def post(self, request, format="json"):
|
||
"""
|
||
Create a blank conversation and return the title and id number
|
||
"""
|
||
title = request.data.get("name")
|
||
conversation = Conversation.objects.create(title=title)
|
||
conversation.save()
|
||
conversation.user_id = request.user.id
|
||
conversation.save()
|
||
|
||
# TODO: when we are smart enough to create a conversation when a prompt is sent
|
||
# client = LlamaClient()
|
||
# inital_message = request.data.get("prompt")
|
||
# title = client.generate_conversation_title(inital_message)
|
||
# title = title if title else "New Conversation"
|
||
# conversation = Conversation.objects.create(title=title)
|
||
# conversation.save()
|
||
# conversation.user_id = request.user.id
|
||
# conversation.save()
|
||
|
||
return Response(
|
||
{"title": title, "id": conversation.id}, status=status.HTTP_201_CREATED
|
||
)
|
||
|
||
|
||
class ConversationPreferences(APIView):
|
||
def get(self, request, format="json"):
|
||
user = request.user
|
||
return Response({"order": user.conversation_order}, 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)
|
||
|
||
|
||
class ConversationDetailView(APIView):
|
||
def get(self, request, format="json"):
|
||
conversation_id = request.query_params.get("conversation_id")
|
||
if not Conversation.objects.filter(
|
||
id=conversation_id, user=request.user, deleted=False
|
||
).exists():
|
||
return Response(
|
||
{"detail": "Conversation not found."},
|
||
status=status.HTTP_404_NOT_FOUND,
|
||
)
|
||
prompts = Prompt.objects.filter(
|
||
conversation__id=conversation_id, conversation__user=request.user
|
||
)
|
||
serailzer = PromptSerializer(prompts, many=True)
|
||
return Response(serailzer.data, status=status.HTTP_200_OK)
|
||
|
||
def post(self, request, format="json"):
|
||
logger.info("In the post")
|
||
# Add the prompt to the database
|
||
# make sure there is a conversation for it
|
||
# if there is not a conversation create a title for it
|
||
|
||
# make sure that our model exists and it is running
|
||
prompt = request.data.get("prompt")
|
||
if not isinstance(prompt, str) or not prompt.strip():
|
||
return Response(
|
||
{"detail": "Message text cannot be empty."},
|
||
status=status.HTTP_400_BAD_REQUEST,
|
||
)
|
||
prompt = prompt.strip()
|
||
|
||
conversation_id = request.data.get("conversation_id")
|
||
is_user = bool(request.data.get("is_user"))
|
||
|
||
try:
|
||
conversation = Conversation.objects.get(
|
||
id=conversation_id, user=request.user, deleted=False
|
||
)
|
||
|
||
# add the prompt to the conversation
|
||
serializer = PromptSerializer(
|
||
data={
|
||
"message": prompt,
|
||
"user_created": is_user,
|
||
"created": timezone.now(),
|
||
}
|
||
)
|
||
if serializer.is_valid():
|
||
prompt_instance = serializer.save()
|
||
prompt_instance.conversation_id = conversation.id
|
||
prompt_instance = serializer.save()
|
||
|
||
# set up the streaming response if it is from the user
|
||
logger.info(f"Do we have a valid user? {is_user}")
|
||
if is_user:
|
||
messages = []
|
||
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",
|
||
}
|
||
)
|
||
|
||
channel_layer = get_channel_layer()
|
||
logger.info(f"Sending to the channel: {CHANNEL_NAME}")
|
||
async_to_sync(channel_layer.group_send)(
|
||
CHANNEL_NAME, {"type": "receive", "content": messages}
|
||
)
|
||
except:
|
||
logger.error(
|
||
f"Error trying to submit to conversation_id: {conversation_id} with request.data: {request.data}"
|
||
)
|
||
pass
|
||
|
||
return Response(status=status.HTTP_200_OK)
|
||
|
||
def delete(self, request, format="json"):
|
||
conversation_id = request.data.get("conversation_id")
|
||
conversation = Conversation.objects.get(id=conversation_id, user=request.user)
|
||
conversation.deleted = True
|
||
conversation.save()
|
||
return Response(status=status.HTTP_202_ACCEPTED)
|
||
|
||
|
||
class UserPromptAnalytics(APIView):
|
||
def get(self, request, format="json"):
|
||
now = timezone.now()
|
||
result = []
|
||
|
||
number_of_months = 3
|
||
company_user_ids = CustomUser.objects.filter(
|
||
company=request.user.company
|
||
).values_list("id", flat=True)
|
||
for i in range(number_of_months):
|
||
next_year = now.year
|
||
next_month = now.month - i
|
||
while next_month < 1:
|
||
next_year -= 1
|
||
next_month += 12
|
||
|
||
start_date = datetime.datetime(next_year, next_month, 1)
|
||
end_date = last_day_of_month(start_date)
|
||
total_conversations = Conversation.objects.filter(
|
||
created__gte=start_date, created__lte=end_date
|
||
)
|
||
total_prompts = Prompt.objects.filter(
|
||
conversation__id__in=total_conversations,
|
||
created__gte=start_date,
|
||
created__lte=end_date,
|
||
)
|
||
total_users = len(CustomUser.objects.all())
|
||
my_conversations = Conversation.objects.filter(user=request.user)
|
||
my_prompts = Prompt.objects.filter(
|
||
conversation__in=my_conversations,
|
||
created__gte=start_date,
|
||
created__lte=end_date,
|
||
)
|
||
company_conversations = Conversation.objects.filter(
|
||
user__id__in=company_user_ids
|
||
)
|
||
company_prompts = Prompt.objects.filter(
|
||
conversation__in=company_conversations,
|
||
created__gte=start_date,
|
||
created__lte=end_date,
|
||
)
|
||
|
||
result.append(
|
||
{
|
||
"month": start_date.strftime("%B"),
|
||
"you": len(my_prompts),
|
||
"others": len(company_prompts) / len(company_user_ids),
|
||
"all": len(total_prompts) / total_users,
|
||
}
|
||
)
|
||
|
||
return Response(result[::-1], status=status.HTTP_200_OK)
|
||
|
||
|
||
class UserConversationAnalytics(APIView):
|
||
def get(self, request, format="json"):
|
||
now = timezone.now()
|
||
result = []
|
||
|
||
number_of_months = 3
|
||
company_user_ids = CustomUser.objects.filter(
|
||
company=request.user.company
|
||
).values_list("id", flat=True)
|
||
for i in range(number_of_months):
|
||
next_year = now.year
|
||
next_month = now.month - i
|
||
while next_month < 1:
|
||
next_year -= 1
|
||
next_month += 12
|
||
|
||
start_date = datetime.datetime(next_year, next_month, 1)
|
||
end_date = last_day_of_month(start_date)
|
||
total_conversations = len(
|
||
Conversation.objects.filter(
|
||
created__gte=start_date, created__lte=end_date
|
||
)
|
||
)
|
||
total_users = len(CustomUser.objects.all())
|
||
company_conversations = len(
|
||
Conversation.objects.filter(
|
||
user__id__in=company_user_ids,
|
||
created__gte=start_date,
|
||
created__lte=end_date,
|
||
)
|
||
)
|
||
|
||
result.append(
|
||
{
|
||
"month": start_date.strftime("%B"),
|
||
"you": len(
|
||
Conversation.objects.filter(
|
||
user=request.user,
|
||
created__gte=start_date,
|
||
created__lte=end_date,
|
||
)
|
||
),
|
||
"others": company_conversations / len(company_user_ids),
|
||
"all": total_conversations / total_users,
|
||
}
|
||
)
|
||
|
||
return Response(result[::-1], status=status.HTTP_200_OK)
|
||
|
||
|
||
class CompanyUsageAnalytics(APIView):
|
||
def get(self, request, format="json"):
|
||
now = timezone.now()
|
||
result = []
|
||
|
||
number_of_months = 3
|
||
company_user_ids = CustomUser.objects.filter(
|
||
company=request.user.company
|
||
).values_list("id", flat=True)
|
||
|
||
for i in range(number_of_months):
|
||
next_year = now.year
|
||
next_month = now.month - i
|
||
while next_month < 1:
|
||
next_year -= 1
|
||
next_month += 12
|
||
|
||
start_date = datetime.datetime(next_year, next_month, 1)
|
||
end_date = last_day_of_month(start_date)
|
||
conversations = Conversation.objects.filter(
|
||
user__id__in=company_user_ids,
|
||
created__gte=start_date,
|
||
created__lte=end_date,
|
||
)
|
||
|
||
conversation_user_ids = conversations.values_list(
|
||
"user__id", flat=True
|
||
).distinct()
|
||
result.append(
|
||
{
|
||
"month": start_date.strftime("%B"),
|
||
"used": len(conversation_user_ids),
|
||
"not_used": len(company_user_ids) - len(conversation_user_ids),
|
||
}
|
||
)
|
||
return Response(result[::-1], status=status.HTTP_200_OK)
|
||
|
||
|
||
return Response(result[::-1], status=status.HTTP_200_OK)
|
||
|
||
|
||
class UserPromptHeatmap(APIView):
|
||
"""Weekday × hour bins of user-entered prompts (local timezone) (#94)."""
|
||
|
||
DAY_LABELS = ("Mon", "Tue", "Wed", "Thu", "Fri", "Sat", "Sun")
|
||
|
||
def get(self, request, format="json"):
|
||
tz_name = request.query_params.get("tz", "UTC")
|
||
try:
|
||
user_tz = pytz.timezone(tz_name)
|
||
except pytz.UnknownTimeZoneError:
|
||
user_tz = pytz.UTC
|
||
tz_name = "UTC"
|
||
|
||
matrix = [[0 for _ in range(24)] for _ in range(7)]
|
||
total = 0
|
||
prompts = Prompt.objects.filter(
|
||
conversation__user=request.user,
|
||
conversation__deleted=False,
|
||
user_created=True,
|
||
).only("created")
|
||
|
||
for prompt in prompts.iterator():
|
||
created = prompt.created
|
||
if timezone.is_naive(created):
|
||
created = timezone.make_aware(created, datetime.timezone.utc)
|
||
local = created.astimezone(user_tz)
|
||
matrix[local.weekday()][local.hour] += 1
|
||
total += 1
|
||
|
||
max_count = max((count for row in matrix for count in row), default=0)
|
||
most_active_day = None
|
||
most_active_hour = None
|
||
peak_cell = None
|
||
if max_count > 0:
|
||
day_totals = [sum(row) for row in matrix]
|
||
most_active_day = self.DAY_LABELS[day_totals.index(max(day_totals))]
|
||
hour_totals = [sum(matrix[day][hour] for day in range(7)) for hour in range(24)]
|
||
most_active_hour = hour_totals.index(max(hour_totals))
|
||
peak_day = 0
|
||
peak_hour = 0
|
||
for day_idx, row in enumerate(matrix):
|
||
for hour_idx, count in enumerate(row):
|
||
if count > matrix[peak_day][peak_hour]:
|
||
peak_day = day_idx
|
||
peak_hour = hour_idx
|
||
peak_cell = {
|
||
"day": self.DAY_LABELS[peak_day],
|
||
"hour": peak_hour,
|
||
"count": matrix[peak_day][peak_hour],
|
||
}
|
||
|
||
return Response(
|
||
{
|
||
"tz": tz_name,
|
||
"total": total,
|
||
"max": max_count,
|
||
"days": list(self.DAY_LABELS),
|
||
"hours": list(range(24)),
|
||
"matrix": matrix,
|
||
"most_active_day": most_active_day,
|
||
"most_active_hour": most_active_hour,
|
||
"peak_cell": peak_cell,
|
||
},
|
||
status=status.HTTP_200_OK,
|
||
)
|
||
|
||
|
||
class AdminAnalytics(APIView):
|
||
def get(self, request, format="json"):
|
||
number_of_months = 3
|
||
result = []
|
||
now = timezone.now()
|
||
|
||
for i in range(number_of_months):
|
||
next_year = now.year
|
||
next_month = now.month - i
|
||
while next_month < 1:
|
||
next_year -= 1
|
||
next_month += 12
|
||
|
||
start_date = datetime.datetime(next_year, next_month, 1)
|
||
end_date = last_day_of_month(start_date)
|
||
durations = [
|
||
item.get_duration()
|
||
for item in PromptMetric.objects.filter(
|
||
created__gte=start_date, created__lte=end_date
|
||
)
|
||
]
|
||
if len(durations) == 0:
|
||
result.append(
|
||
{
|
||
"month": start_date.strftime("%B"),
|
||
"range": [0, 0],
|
||
"avg": 0,
|
||
"median": 0,
|
||
}
|
||
)
|
||
continue
|
||
|
||
average = sum(durations) / len(durations)
|
||
min_value = min(durations)
|
||
max_value = max(durations)
|
||
durations.sort()
|
||
median = durations[len(durations) // 2]
|
||
result.append(
|
||
{
|
||
"month": start_date.strftime("%B"),
|
||
"range": [min_value, max_value],
|
||
"avg": average,
|
||
"median": median,
|
||
}
|
||
)
|
||
|
||
return Response(result[::-1], status=status.HTTP_200_OK)
|
||
|
||
|
||
prompt = ChatPromptTemplate.from_messages(
|
||
[("system", ASSISTANT_SYSTEM_PROMPT), ("user", "{input}")]
|
||
)
|
||
|
||
llm = OllamaLLM(**ollama_llm_kwargs(model=MODEL_NAME))
|
||
|
||
# output_parser = StrOutputParser()
|
||
# # Chain
|
||
# chain = prompt | llm.with_config({"run_name": "model"}) | output_parser.with_config({"run_name": "Assistant"})
|
||
|
||
|
||
# Document Views
|
||
def _feature_gate_response(exc: FeatureNotAllowed) -> Response:
|
||
return Response(
|
||
{"code": exc.code, "error": exc.message, "details": exc.details},
|
||
status=status.HTTP_403_FORBIDDEN,
|
||
)
|
||
|
||
|
||
class DocumentWorkspaceView(APIView):
|
||
# permission_classes = [permissions.IsAuthenticated]
|
||
|
||
def get(self, request):
|
||
try:
|
||
assert_feature_allowed(request.user, "rag")
|
||
except FeatureNotAllowed as exc:
|
||
return _feature_gate_response(exc)
|
||
if request.user.company_id:
|
||
workspaces = DocumentWorkspace.objects.filter(
|
||
company=request.user.company, user__isnull=True
|
||
)
|
||
else:
|
||
workspaces = DocumentWorkspace.objects.filter(
|
||
user=request.user, company__isnull=True
|
||
)
|
||
serializer = DocumentWorkspaceSerializer(workspaces, many=True)
|
||
return Response(serializer.data)
|
||
|
||
def post(self, request):
|
||
try:
|
||
assert_feature_allowed(request.user, "rag")
|
||
except FeatureNotAllowed as exc:
|
||
return _feature_gate_response(exc)
|
||
serializer = DocumentWorkspaceSerializer(data=request.data)
|
||
if serializer.is_valid():
|
||
if request.user.company_id:
|
||
serializer.save(company=request.user.company, user=None)
|
||
else:
|
||
serializer.save(company=None, user=request.user)
|
||
return Response(serializer.data, status=status.HTTP_201_CREATED)
|
||
return Response(serializer.errors, status=status.HTTP_400_BAD_REQUEST)
|
||
|
||
|
||
class DocumentUploadView(APIView):
|
||
# permission_classes = [permissions.IsAuthenticated]Z
|
||
|
||
DOCUMENT_SORT_FIELDS = {
|
||
"name": "file",
|
||
"-name": "-file",
|
||
"created": "created",
|
||
"-created": "-created",
|
||
"date_uploaded": "created",
|
||
"-date_uploaded": "-created",
|
||
"processed": "processed",
|
||
"-processed": "-processed",
|
||
"active": "active",
|
||
"-active": "-active",
|
||
}
|
||
|
||
def _resolve_list_workspace(self, request):
|
||
"""Return (workspace, scope) for personal|company document lists (#60)."""
|
||
raw = (request.query_params.get("workspace") or "").strip().lower()
|
||
if not raw:
|
||
raw = "company" if request.user.company_id else "personal"
|
||
|
||
if raw == "personal":
|
||
return ensure_personal_workspace(request.user), "personal"
|
||
if raw == "company":
|
||
if not request.user.company_id:
|
||
return None, "company"
|
||
return ensure_company_workspace(request.user.company), "company"
|
||
return False, raw
|
||
|
||
def get(self, request):
|
||
logger.debug(f"request_3: {request}")
|
||
try:
|
||
assert_feature_allowed(request.user, "rag")
|
||
except FeatureNotAllowed as exc:
|
||
return _feature_gate_response(exc)
|
||
|
||
workspace, scope = self._resolve_list_workspace(request)
|
||
if workspace is False:
|
||
return Response(
|
||
{"error": "workspace must be 'personal' or 'company'."},
|
||
status=status.HTTP_400_BAD_REQUEST,
|
||
)
|
||
if workspace is None:
|
||
return Response(
|
||
{"error": "Company workspace requires a company membership."},
|
||
status=status.HTTP_400_BAD_REQUEST,
|
||
)
|
||
|
||
queryset = Document.objects.filter(workspace=workspace)
|
||
|
||
search = (request.query_params.get("search") or request.query_params.get("q") or "").strip()
|
||
if search:
|
||
queryset = queryset.filter(
|
||
Q(file__icontains=search) | Q(remote_name__icontains=search)
|
||
)
|
||
|
||
ordering = (request.query_params.get("ordering") or "-created").strip()
|
||
order_by = self.DOCUMENT_SORT_FIELDS.get(ordering, "-created")
|
||
queryset = queryset.order_by(order_by, "id")
|
||
|
||
try:
|
||
page = max(int(request.query_params.get("page", 1)), 1)
|
||
except (TypeError, ValueError):
|
||
page = 1
|
||
try:
|
||
page_size = int(request.query_params.get("page_size", 20))
|
||
except (TypeError, ValueError):
|
||
page_size = 20
|
||
page_size = min(max(page_size, 1), 100)
|
||
|
||
total = queryset.count()
|
||
start = (page - 1) * page_size
|
||
end = start + page_size
|
||
serializer = DocumentSerializer(queryset[start:end], many=True)
|
||
return Response(
|
||
{
|
||
"count": total,
|
||
"page": page,
|
||
"page_size": page_size,
|
||
"scope": scope,
|
||
"results": serializer.data,
|
||
},
|
||
status=status.HTTP_200_OK,
|
||
)
|
||
|
||
def post(self, request):
|
||
logger.debug(f"request: {request}")
|
||
|
||
try:
|
||
assert_feature_allowed(request.user, "rag")
|
||
except FeatureNotAllowed as exc:
|
||
return _feature_gate_response(exc)
|
||
|
||
scope = (request.query_params.get("workspace") or request.data.get("workspace") or "").strip().lower()
|
||
if scope == "personal":
|
||
workspace = ensure_personal_workspace(request.user)
|
||
elif scope == "company":
|
||
if not request.user.company_id:
|
||
return Response(
|
||
{"error": "Company workspace requires a company membership."},
|
||
status=status.HTTP_400_BAD_REQUEST,
|
||
)
|
||
workspace = ensure_company_workspace(request.user.company)
|
||
else:
|
||
workspace = ensure_workspace_for_user(request.user)
|
||
|
||
logger.info(request.FILES)
|
||
file = request.FILES.get("file")
|
||
if not file:
|
||
return Response(
|
||
{"error": "No file provided"}, status=status.HTTP_400_BAD_REQUEST
|
||
)
|
||
|
||
logger.info("have the workspace and the file")
|
||
|
||
document = Document.objects.create(workspace=workspace, file=file)
|
||
|
||
# process the document inthe background
|
||
self.process_document(document)
|
||
|
||
serializer = DocumentSerializer(document)
|
||
return Response(serializer.data, status=status.HTTP_201_CREATED)
|
||
|
||
def process_document(self, document):
|
||
# File bytes live in DB (DatabaseStorage); RAG materializes a temp path.
|
||
document.processed = True
|
||
document.active = True
|
||
document.save()
|
||
service = AsyncRAGService()
|
||
service.add_files_to_store(
|
||
[
|
||
(
|
||
document.file,
|
||
document.file.name,
|
||
document.workspace_id,
|
||
document.id,
|
||
document.active,
|
||
)
|
||
],
|
||
workspace_id=document.workspace_id,
|
||
)
|
||
|
||
|
||
class DocumentDetailView(APIView):
|
||
# permission_classes = [permissions.IsAuthenticated]
|
||
|
||
def _get_document(self, request, document_id):
|
||
user = request.user
|
||
if user.company_id:
|
||
return (
|
||
Document.objects.filter(id=document_id)
|
||
.filter(
|
||
Q(workspace__user=user, workspace__company__isnull=True)
|
||
| Q(
|
||
workspace__company_id=user.company_id,
|
||
workspace__user__isnull=True,
|
||
)
|
||
)
|
||
.first()
|
||
)
|
||
return Document.objects.filter(
|
||
id=document_id,
|
||
workspace__user=user,
|
||
workspace__company__isnull=True,
|
||
).first()
|
||
|
||
def get(self, request, document_id):
|
||
logger.info(f"request: {request}")
|
||
try:
|
||
assert_feature_allowed(request.user, "rag")
|
||
except FeatureNotAllowed as exc:
|
||
return _feature_gate_response(exc)
|
||
|
||
document = self._get_document(request, document_id)
|
||
if document is None:
|
||
return Response(
|
||
{"error": "Document not found"}, status=status.HTTP_404_NOT_FOUND
|
||
)
|
||
|
||
serializer = DocumentSerializer(document)
|
||
return Response(serializer.data, status=status.HTTP_200_OK)
|
||
|
||
def patch(self, request, document_id):
|
||
"""Toggle a document's ``active`` flag (#44) and its vector metadata (#45)."""
|
||
try:
|
||
assert_feature_allowed(request.user, "rag")
|
||
except FeatureNotAllowed as exc:
|
||
return _feature_gate_response(exc)
|
||
|
||
document = self._get_document(request, document_id)
|
||
if document is None:
|
||
return Response(
|
||
{"error": "Document not found"}, status=status.HTTP_404_NOT_FOUND
|
||
)
|
||
|
||
if "active" not in request.data:
|
||
return Response(
|
||
{"error": "active is required"}, status=status.HTTP_400_BAD_REQUEST
|
||
)
|
||
|
||
active = request.data.get("active")
|
||
if isinstance(active, str):
|
||
active = active.strip().lower() in {"1", "true", "yes"}
|
||
document.active = bool(active)
|
||
document.save(update_fields=["active"])
|
||
|
||
try:
|
||
AsyncRAGService().set_document_active(document.id, document.active)
|
||
except Exception as exc:
|
||
logger.warning(
|
||
"Failed to update vector active metadata for document %s: %s",
|
||
document.id,
|
||
exc,
|
||
)
|
||
|
||
serializer = DocumentSerializer(document)
|
||
return Response(serializer.data, status=status.HTTP_200_OK)
|
||
|
||
def delete(self, request, document_id):
|
||
try:
|
||
assert_feature_allowed(request.user, "rag")
|
||
except FeatureNotAllowed as exc:
|
||
return _feature_gate_response(exc)
|
||
|
||
document = self._get_document(request, document_id)
|
||
if document is None:
|
||
return Response(
|
||
{"error": "Document not found"}, status=status.HTTP_404_NOT_FOUND
|
||
)
|
||
|
||
document.delete()
|
||
return Response(status=status.HTTP_204_NO_CONTENT)
|