Files
chat_backend/llm_be/chat_backend/views.py
T
westfarn 57a2350b8e
Deploy Beta / unit-tests (push) Successful in 11s
Unit Tests / test (push) Successful in 10s
Deploy Beta / docker (push) Successful in 21s
Deploy Beta / deploy-beta (push) Successful in 48s
Drive sync progress + documents list API + prompt heatmap (#59, #60, #94) (#61)
## 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
2026-08-02 04:49:23 -07:00

1143 lines
40 KiB
Python
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
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)